AI-Hypercomputer / AI-Hypercomputer/maxtext

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

オープン
#5,129 コメント 0 件 リアクション 0 件 担当者 0 名 GitHub で見る
主要言語
Python
スター
2.4k
フォーク
607
平均マージ
2日 19時間
マージ済み PR(30日)
158

説明

## 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.

```bash
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:

```python
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`.

```python
# 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.

コントリビューションガイド

コントリビューションガイドを開く

評価

この issue はまだ評価されていません。

新しい issue をメールで受け取る

初心者向けの GitHub issue を短くまとめたダイジェスト。