Skip to content

Commit 18d124f

Browse files
fix(nnx): optimize scan memory by excluding read-only parameters from scan outputs
1 parent e190ed6 commit 18d124f

1 file changed

Lines changed: 12 additions & 4 deletions

File tree

src/maxtext/layers/nnx_decoders.py

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -356,13 +356,17 @@ def layer_fn(carry, scanned_vars):
356356
**kwargs,
357357
)
358358
new_carry = layer_out[0] if isinstance(layer_out, tuple) else layer_out
359-
return new_carry, nnx.state(layer)
359+
# Avoid returning and stacking read-only parameters inside the scan body.
360+
# This prevents huge unnecessary memory allocation.
361+
_, _, updated_state = nnx.split(layer, nnx.Param, ...)
362+
return new_carry, updated_state
360363

361364
final_carry, scanned_state = jax.lax.scan(layer_fn, inputs, (params, state))
362365

363366
if scan_axis != 0:
364367
scanned_params, scanned_other = scanned_state.split(nnx.Param, ...)
365-
scanned_params = jax.tree.map(lambda x: jnp.moveaxis(x, 0, scan_axis), scanned_params)
368+
if scanned_params:
369+
scanned_params = jax.tree.map(lambda x: jnp.moveaxis(x, 0, scan_axis), scanned_params)
366370
scanned_state = nnx.State.merge(scanned_params, scanned_other)
367371

368372
nnx.update(self.scanned_layers, scanned_state)
@@ -1031,7 +1035,10 @@ def layer_fn(carry, scanned_vars):
10311035
returned_params = updated_params
10321036
new_current_state = nnx.State.merge(returned_params, updated_state)
10331037
else:
1034-
new_current_state = nnx.state(layer)
1038+
# Avoid returning and stacking read-only parameters inside the scan body.
1039+
# This prevents huge unnecessary memory allocation.
1040+
_, _, updated_state = nnx.split(layer, nnx.Param, ...)
1041+
new_current_state = updated_state
10351042

10361043
if use_kv:
10371044
return new_carry, (new_current_state, updated_kv)
@@ -1077,7 +1084,8 @@ def layer_fn(carry, scanned_vars):
10771084

10781085
if scan_axis != 0:
10791086
new_params, new_rest = scanned_state.split(nnx.Param, ...)
1080-
new_params = maxtext_utils_nnx.nnx_sync_moveaxis(new_params, 0, scan_axis)
1087+
if new_params:
1088+
new_params = maxtext_utils_nnx.nnx_sync_moveaxis(new_params, 0, scan_axis)
10811089
scanned_state = nnx.merge_state(new_params, new_rest)
10821090

10831091
returned_kv_stacked = None

0 commit comments

Comments
 (0)