Skip to content

LoRA under scan_layers=True restacks the whole base model, once per adapter #5129

Description

@hodaaaaaaaaaa

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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions