Skip to content

Commit 655c636

Browse files
fix(nnx): optimize scan memory by excluding read-only parameters from scan outputs
When using `jax.lax.scan` in Flax NNX decoder layers (e.g. DeepSeek-V3 / Gemma 2), returning `nnx.state(layer)` inside `layer_fn` caused `jax.lax.scan` to stack read-only model parameters (`nnx.Param`) across all 58 MoE layers into Python scan outputs. This allocated ~12.7 GiB of extra HBM memory and triggered `RESOURCE_EXHAUSTED` OOM errors during large model execution. This change uses `_, _, updated_state = nnx.split(layer, nnx.Param, ...)` inside `layer_fn` to exclude read-only parameters from scan carry/output states, returning only mutable layer states to `jax.lax.scan`. Verification: - Local NNX decoder unit tests pass 100% (44/44 in `tests/unit/nnx_decoders_test.py`). - Pre-commit code quality checks pass 100% (`codespell`, `pylint`, `pyink`, `yamllint`). - End-to-end logit check on `mlperf-v5p` cluster (`v5p-128`) for DeepSeek-V3 671B ran forward pass to completion without OOM. TAG=agy CONV=0c977f36-644a-472a-981c-1856f8cac8a6
1 parent d557e88 commit 655c636

1 file changed

Lines changed: 10 additions & 3 deletions

File tree

src/maxtext/layers/nnx_decoders.py

Lines changed: 10 additions & 3 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)
@@ -957,7 +961,10 @@ def layer_fn(carry, scanned_vars):
957961
returned_params = updated_params
958962
new_current_state = nnx.State.merge(returned_params, updated_state)
959963
else:
960-
new_current_state = nnx.state(layer)
964+
# Avoid returning and stacking read-only parameters inside the scan body.
965+
# This prevents huge unnecessary memory allocation.
966+
_, _, updated_state = nnx.split(layer, nnx.Param, ...)
967+
new_current_state = updated_state
961968

962969
if use_kv:
963970
return new_carry, (new_current_state, updated_kv)

0 commit comments

Comments
 (0)