Summary
Injecting LoRA into a scanned model allocates a full private copy of the decoder layer stack per adapter: on gemma2-2b, 1194 MB per adapter, for adapter weights of 6.1 MB. Nothing errors — the base just can't be shared, so multi-adapter serving on one worker is impossible with scan on.
MaxText's two scan appliers landed at opposite ends of the same tradeoff:
| returns from the scan body |
consequence for LoRA |
apply_scanned_layers before 5e3a353: no params |
adapters never escape — loud failure |
apply_scanned_layers today: only params created in the body |
correct |
_apply_layers_sequentially: all params, whenever the graph is changing |
adapters escape and drag the base with them — silent 196× blowup |
5e3a353 ("keep params created inside the layer scan") found the middle answer for apply_scanned_layers, which has two callers (Gemma 4, Qwen3-Next). Everything else scanned, gemma2 included, goes through _apply_layers_sequentially, which never got the fix because it doesn't crash.
This issue is that fix, transcribed to the other applier.
Reproduce
maxtext @ main 9abba7e · jax 0.11.1 · flax 0.12.9 · qwix 0.1.8 · CPU, no TPU needed
Build gemma2-2b with base_num_decoder_layers=2, scan_layers=True, an fsdp=4 × tensor=2 mesh; call lora_utils.apply_lora_to_model(model, mesh, config) twice; diff jax.live_arrays() around each call.
XLA_FLAGS=--xla_force_host_platform_device_count=8 JAX_PLATFORMS=cpu python repro.py
base = 1260.1 MB, lora = 6.1 MB
adapter A added 1194.2 MB
adapter B added 1194.2 MB
base array identity across adapters: 2/24 same object
The two that stay shared are the only Gemma 2 weights outside the layer stack.
Cause
NNXDecoder._apply_layers_sequentially, src/maxtext/layers/nnx_decoders.py. dynamic_graph_init (line 1081) reads disable_quant_stats_update, the flag qwix sets for its one init forward pass. When true, the body returns the full param state instead of taking the cheap branch just below:
if dynamic_graph_init:
new_graphdef, updated_params, updated_state = nnx.split(layer, nnx.Param, ...)
updated_graphdef[0] = new_graphdef
returned_params = updated_params # <-- the whole base, every injection
new_current_state = nnx.State.merge(returned_params, updated_state)
jax.lax.scan stacks every output into fresh arrays; line 1197 merges them into a new module that replaces self.layers. The branch is necessary — new LoRA params are created inside the body and this is how they get out. It just carries the base along.
Fix
Record the paths fed in as scan inputs, return only paths not in that set, read the base back off the pre-scan layers.
# before the body
carried_param_paths = {path for path, _ in nnx.to_flat_state(params)}
# in the body, replacing `returned_params = updated_params`
returned_params = nnx.from_flat_state(
[(path, value) for path, value in nnx.to_flat_state(updated_params)
if path not in carried_param_paths]
)
# after the scan
new_params, new_rest = scanned_state.split(nnx.Param, ...)
out_layers = nnx.merge(updated_graphdef[0], nnx.state(layers, nnx.Param), new_params, new_rest)
nnx.state(layers, nnx.Param) preserves array identity, so the base is genuinely shared rather than just smaller, and it returns at its original param_scan_axis layout.
One semantic change: under dynamic_graph_init, in-body mutations of carried params are no longer returned. The non-dynamic branch already discards them and apply_scanned_layers assumes the same — the applier already treats scanned params as read-only inside the body.
Result
|
before |
after |
| per-adapter memory |
1194.2 MB |
6.1 MB |
| base array identity |
2/24 same object |
24/24 |
| LoRA factors created |
28 |
28 |
nnx.jit forward |
passes |
passes |
Adapter outputs are bitwise identical patched vs unpatched. tests/unit/{nnx_decoders,nnx_decoder,nnx_scan,lora_utils_nnx,nnx_wrappers}_test.py: 75 passed, 2 skipped.
The PR includes a regression test (TestApplyLayersSequentiallyDynamicGraphInit) that drives the applier directly with a stacked dummy layer creating an nnx.LoRAParam while tracing, and asserts that the new param escapes and that the base array is the same object afterwards, at param_scan_axis 0 and 1. Fails on main, passes patched.
Why it's worth taking
lora_module_path.yml has a first-class gemma2 entry and sft.yml doesn't override scan_layers, so MaxText's own LoRA SFT path runs with scan on. Running with scan off instead costs ~4.2× compile time and 19–39% per step in our measurements.
5e3a353 already settled that scanned LoRA is supported; this applies the same decision to the applier serving every other model. The two appliers are slated to be unified (nnx_decoders.py:1024) — fixing this first means the merge won't have to reconcile two different behaviors.
Summary
Injecting LoRA into a scanned model allocates a full private copy of the decoder layer stack per adapter: on
gemma2-2b, 1194 MB per adapter, for adapter weights of 6.1 MB. Nothing errors — the base just can't be shared, so multi-adapter serving on one worker is impossible with scan on.MaxText's two scan appliers landed at opposite ends of the same tradeoff:
apply_scanned_layersbefore5e3a353: no paramsapply_scanned_layerstoday: only params created in the body_apply_layers_sequentially: all params, whenever the graph is changing5e3a353("keep params created inside the layer scan") found the middle answer forapply_scanned_layers, which has two callers (Gemma 4, Qwen3-Next). Everything else scanned,gemma2included, goes through_apply_layers_sequentially, which never got the fix because it doesn't crash.This issue is that fix, transcribed to the other applier.
Reproduce
maxtext@main9abba7e·jax0.11.1 ·flax0.12.9 ·qwix0.1.8 · CPU, no TPU neededBuild
gemma2-2bwithbase_num_decoder_layers=2,scan_layers=True, anfsdp=4 × tensor=2mesh; calllora_utils.apply_lora_to_model(model, mesh, config)twice; diffjax.live_arrays()around each call.The two that stay shared are the only Gemma 2 weights outside the layer stack.
Cause
NNXDecoder._apply_layers_sequentially,src/maxtext/layers/nnx_decoders.py.dynamic_graph_init(line 1081) readsdisable_quant_stats_update, the flag qwix sets for its one init forward pass. When true, the body returns the full param state instead of taking the cheap branch just below:jax.lax.scanstacks every output into fresh arrays; line 1197 merges them into a new module that replacesself.layers. The branch is necessary — new LoRA params are created inside the body and this is how they get out. It just carries the base along.Fix
Record the paths fed in as scan inputs, return only paths not in that set, read the base back off the pre-scan
layers.nnx.state(layers, nnx.Param)preserves array identity, so the base is genuinely shared rather than just smaller, and it returns at its originalparam_scan_axislayout.One semantic change: under
dynamic_graph_init, in-body mutations of carried params are no longer returned. The non-dynamic branch already discards them andapply_scanned_layersassumes the same — the applier already treats scanned params as read-only inside the body.Result
nnx.jitforwardAdapter outputs are bitwise identical patched vs unpatched.
tests/unit/{nnx_decoders,nnx_decoder,nnx_scan,lora_utils_nnx,nnx_wrappers}_test.py: 75 passed, 2 skipped.The PR includes a regression test (
TestApplyLayersSequentiallyDynamicGraphInit) that drives the applier directly with a stacked dummy layer creating annnx.LoRAParamwhile tracing, and asserts that the new param escapes and that the base array is the same object afterwards, atparam_scan_axis0 and 1. Fails onmain, passes patched.Why it's worth taking
lora_module_path.ymlhas a first-classgemma2entry andsft.ymldoesn't overridescan_layers, so MaxText's own LoRA SFT path runs with scan on. Running with scan off instead costs ~4.2× compile time and 19–39% per step in our measurements.5e3a353already settled that scanned LoRA is supported; this applies the same decision to the applier serving every other model. The two appliers are slated to be unified (nnx_decoders.py:1024) — fixing this first means the merge won't have to reconcile two different behaviors.