Skip to content

Commit 3a792f1

Browse files
Refactor NNX decoder layers structure and fix checkpoint loading (#4504)
- Fix NNX sequential layer naming and PyTree structure to match Flax Linen decoder structure when scan_layers=False. This fixes checkpoint loading for unscanned models. - Refactor layer retrieval to use a unified get_layers() method in both ToNNX wrapper and NNX_Decoder, avoiding manual listing of layer names. - Add unit tests for ToNNX.get_layers(). - Fix a TypeError in Engram integration (_apply_single_engram_layer) by removing duplicate decoder_input_tokens from key list. - Copy tests/ directory in maxtext_runner.Dockerfile to support running verification tests inside the container. TAG=agy CONV=b1cd074b-f03b-4919-9868-597278555171
1 parent 80646f7 commit 3a792f1

11 files changed

Lines changed: 346 additions & 65 deletions

File tree

src/maxtext/inference/maxengine/maxengine.py

Lines changed: 31 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,7 @@
2929
if jax.__version_info__ >= (0, 6, 3):
3030
from jax.experimental.layout import Layout as DLL # type: ignore
3131
else:
32-
from jax.experimental.layout import DeviceLocalLayout as DLL # type: ignore
32+
from jax.experimental.layout import DeviceLocalLayout as DLL # type: ignore # pylint: disable=no-name-in-module
3333

3434
from flax import linen as nn
3535
from flax import nnx
@@ -486,7 +486,10 @@ def _overlay(dst, src):
486486
lambda x: jax.sharding.NamedSharding(self._mesh, jax.sharding.PartitionSpec(None, *x.spec)),
487487
self.prefill_kv_cache_shardings,
488488
)
489-
self.prefill_kv_cache_shardings = {"decoder": {"layers": self.prefill_kv_cache_shardings["decoder"]["layers"][0]}}
489+
decoder_dict = self.prefill_kv_cache_shardings["decoder"]
490+
first_key = next(k for k in decoder_dict.keys() if k.endswith("layers_0"))
491+
first_layer_sharding = decoder_dict[first_key]
492+
self.prefill_kv_cache_shardings = {"decoder": {"layers": first_layer_sharding}}
490493
# scan_layers=True is already stacked on axis 0; shardings stay as-is and stack/unstack are no-ops.
491494
# AR-mode abstract model so axis names use CACHE_BATCH (not CACHE_BATCH_PREFILL);
492495
# bulk_insert / _insert_jit search for "cache_batch" in the per-leaf logical axes.
@@ -610,10 +613,17 @@ def _maybe_stack_prefill_result_cache(self, cache):
610613
if self.config.scan_layers:
611614
# scan_layers already stacks the per-layer KV cache on axis 0; nothing to restack.
612615
return cache
613-
# scan_layers=False: stack the per-layer subtrees under decoder/layers into one
616+
# scan_layers=False: stack the per-layer subtrees under decoder into one
614617
# subtree with a leading layer axis (matching the scan_layers=True shape).
615-
layers = cache["decoder"]["layers"]
616-
stacked = jax.tree.map(lambda *c: jnp.stack(c), *[layers[i] for i in range(self.config.num_decoder_layers)])
618+
if "dense_layers_0" in cache["decoder"] or "moe_layers_0" in cache["decoder"]:
619+
first_dense = self.config.first_num_dense_layers
620+
num_moe = self.config.num_decoder_layers - first_dense
621+
layer_keys = [f"dense_layers_{i}" for i in range(first_dense)] + [f"moe_layers_{i}" for i in range(num_moe)]
622+
else:
623+
layer_keys = [f"layers_{i}" for i in range(self.config.num_decoder_layers)]
624+
625+
layer_cache = [cache["decoder"][key] for key in layer_keys]
626+
stacked = jax.tree.map(lambda *c: jnp.stack(c), *layer_cache)
617627
return {"decoder": {"layers": stacked}}
618628

619629
layer_keys = []
@@ -636,8 +646,22 @@ def _maybe_unstack_prefill_result_cache(self, cache):
636646
return cache
637647
# scan_layers=False: split the leading layer axis back into per-layer subtrees.
638648
stacked = cache["decoder"]["layers"]
639-
layers = {i: jax.tree.map(lambda x, i=i: x[i], stacked) for i in range(self.config.num_decoder_layers)}
640-
return {"decoder": {"layers": layers}}
649+
res_cache = {"decoder": {}}
650+
is_deepseek = (
651+
getattr(self.model, "is_deepseek", False)
652+
or (hasattr(self.model, "decoder") and getattr(self.model.decoder, "is_deepseek", False))
653+
or (hasattr(self.config, "decoder_block") and str(self.config.decoder_block).lower() == "deepseek")
654+
)
655+
if is_deepseek:
656+
first_dense = self.config.first_num_dense_layers
657+
num_moe = self.config.num_decoder_layers - first_dense
658+
layer_keys = [f"dense_layers_{i}" for i in range(first_dense)] + [f"moe_layers_{i}" for i in range(num_moe)]
659+
else:
660+
layer_keys = [f"layers_{i}" for i in range(self.config.num_decoder_layers)]
661+
662+
for idx, key in enumerate(layer_keys):
663+
res_cache["decoder"][key] = jax.tree.map(lambda x, i=idx: x[i], stacked)
664+
return res_cache
641665

642666
flat_cache, treedef = jax.tree.flatten(cache)
643667
layer_cache = [jax.tree.unflatten(treedef, flat_cache_vars) for flat_cache_vars in zip(*flat_cache, strict=True)]

src/maxtext/layers/learn_to_init_layer.py

Lines changed: 1 addition & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -388,14 +388,7 @@ def apply_lti_model_update(student_model, student_config):
388388
if student_config.attn_module_name is None:
389389
return
390390

391-
if getattr(student_config, "scan_layers", True):
392-
layer_modules = [student_model.decoder.layers]
393-
else:
394-
layer_modules = []
395-
# Collect all possible layer names (e.g. layers_0, dense_layers_0, moe_layers_0)
396-
for name, module in vars(student_model.decoder).items():
397-
if name.startswith(LTI_LAYER_PATH_PREFIXES):
398-
layer_modules.append(module)
391+
layer_modules = student_model.decoder.get_layers()
399392

400393
for layer_module in layer_modules:
401394
attn_state_dict = layer_module.get(student_config.attn_module_name)

src/maxtext/layers/nnx_decoders.py

Lines changed: 113 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -502,13 +502,13 @@ def _init_pipeline_deepseek(self, decoder_block_classes, rngs):
502502
else:
503503
self.num_dense_layers = config.first_num_dense_layers
504504
for i in range(self.num_dense_layers):
505-
self._create_and_register_named_layer(dense_cls, rngs, "dense_layers", i)
505+
self._create_and_register_layer(dense_cls, rngs, "dense_layers", i)
506506
self.num_moe_outside_pipeline = (
507507
config.num_decoder_layers - config.first_num_dense_layers
508508
) - config.pipeline_parallel_layers
509509
if self.num_moe_outside_pipeline > 0:
510510
for i in range(self.num_moe_outside_pipeline):
511-
self._create_and_register_named_layer(moe_cls, rngs, "moe_layers_outside_pipeline", i)
511+
self._create_and_register_layer(moe_cls, rngs, "moe_layers_outside_pipeline", i)
512512

513513
def _init_pipeline_generic(self, decoder_block_classes, rngs):
514514
"""Initializes generic decoder layers outside pipeline."""
@@ -526,7 +526,7 @@ def _init_pipeline_generic(self, decoder_block_classes, rngs):
526526
else:
527527
self.num_layers_outside_pipeline = remaining_layers
528528
for i in range(self.num_layers_outside_pipeline):
529-
self._create_and_register_named_layer(base_cls, rngs, "layers_outside_pipeline", i)
529+
self._create_and_register_layer(base_cls, rngs, "layers_outside_pipeline", i)
530530

531531
def _init_scanned_layers(self, decoder_block_classes, rngs, mesh):
532532
"""Initializes decoder layers with scanning (non-pipeline)."""
@@ -693,12 +693,9 @@ def _init_scanned_generic(self, decoder_block_classes, rngs):
693693
rngs=rngs,
694694
**layer_kwargs,
695695
)
696-
else:
697-
self.layers = nnx.List([])
698696

699697
def _init_sequential_layers(self, decoder_block_classes, rngs):
700698
"""Initializes decoder layers sequentially (no scanning)."""
701-
self.layers = nnx.List([])
702699

703700
if self.is_deepseek:
704701
self._init_sequential_deepseek(decoder_block_classes, rngs)
@@ -710,9 +707,9 @@ def _init_sequential_deepseek(self, decoder_block_classes, rngs):
710707
config = self.config
711708
dense_cls, moe_cls = decoder_block_classes
712709
for i in range(config.first_num_dense_layers):
713-
self._create_and_register_layer(dense_cls, rngs, "dense_layer", i)
710+
self._create_and_register_layer(dense_cls, rngs, "dense_layers", i)
714711
for i in range(config.num_decoder_layers - config.first_num_dense_layers):
715-
self._create_and_register_layer(moe_cls, rngs, "moe_layer", i)
712+
self._create_and_register_layer(moe_cls, rngs, "moe_layers", i)
716713

717714
def _init_sequential_generic(self, decoder_block_classes, rngs):
718715
"""Initializes sequential generic decoder layers with per-architecture layer_kwargs."""
@@ -753,7 +750,6 @@ def _init_gemma4_small_layers(self, rngs):
753750
``_create_and_register_layer``.
754751
"""
755752
cfg = self.config
756-
self.layers = nnx.List([])
757753
# Only register the PLE submodule when it exists (mirrors the optional position_embedder
758754
# pattern); assigning None first would make nnx treat the attribute as static.
759755
if cfg.hidden_size_per_layer_input > 0 and cfg.vocab_size_per_layer_input > 0:
@@ -771,7 +767,6 @@ def _init_gemma4_small_layers(self, rngs):
771767
rngs=rngs,
772768
)
773769
setattr(self, f"layers_{lyr}", layer)
774-
self.layers.append(layer)
775770

776771
def _get_pipeline_stage_module(self, decoder_blocks, rngs):
777772
"""Retrieves the wrapper module formatted for single pipeline stage execution."""
@@ -801,14 +796,7 @@ def _get_pipeline_stage_module(self, decoder_blocks, rngs):
801796
)
802797

803798
def _create_and_register_layer(self, layer_cls, rngs, base_name, i, **layer_kwargs):
804-
attr_name = f"{base_name}_{i}"
805-
layer = self._create_single_layer(layer_cls, rngs, **layer_kwargs)
806-
setattr(self, attr_name, layer)
807-
self.layers.append(layer)
808-
809-
def _create_and_register_named_layer(self, layer_cls, rngs, base_name, i, **layer_kwargs):
810-
"""Creates a layer registered ONLY via named attribute. Used by pipeline-outside paths
811-
to avoid double-registration when self.layers list is also tracked elsewhere."""
799+
"""Creates a layer registered ONLY via named attribute."""
812800
attr_name = f"{base_name}_{i}"
813801
layer = self._create_single_layer(layer_cls, rngs, **layer_kwargs)
814802
setattr(self, attr_name, layer)
@@ -1408,6 +1396,11 @@ def _apply_single_engram_layer(self, y, layer_name, *args, **kwargs):
14081396
decoder_input_tokens = kwargs.get("decoder_input_tokens")
14091397
layer_kwargs = kwargs.get("layer_kwargs", {})
14101398

1399+
# Create a copy of layer_kwargs and pop decoder_input_tokens if it exists
1400+
# to avoid passing it twice (once explicitly and once via **layer_kwargs).
1401+
layer_kwargs = dict(layer_kwargs)
1402+
layer_kwargs.pop("decoder_input_tokens", None)
1403+
14111404
out = layer(y, *args, decoder_input_tokens=decoder_input_tokens, **layer_kwargs)
14121405
if isinstance(out, tuple):
14131406
y = out[0]
@@ -1545,6 +1538,9 @@ def __call__(
15451538
if attention_metadata is not None:
15461539
layer_kwargs["attention_metadata"] = attention_metadata
15471540

1541+
if cfg.engram_layers and decoder_input_tokens is not None:
1542+
layer_kwargs["decoder_input_tokens"] = decoder_input_tokens
1543+
15481544
if getattr(cfg, "using_pipeline_parallelism", False):
15491545
logical_partition_spec = (
15501546
self.pipeline_module.get_weight_sharding()
@@ -1780,23 +1776,41 @@ def __call__(
17801776
)
17811777
else:
17821778
prevent_cse = maxtext_utils.should_prevent_cse_in_remat(cfg)
1779+
dynamic_graph_init = bool(getattr(self, "disable_quant_stats_update", False))
17831780

1784-
# Hoisted function to preserve XLA cache ID
1785-
def pure_layer_fn(graphdef, state_in, y_in, kv_in):
1786-
1781+
def pure_layer_fn(graphdef_in, state_in, y_in, kv_in):
17871782
if cfg.parameter_memory_host_offload:
17881783
state_in = jax.tree.map(
17891784
lambda x: jax.device_put(x, max_utils.device_space()),
17901785
state_in,
17911786
)
1792-
1793-
merged_layer = nnx.merge(graphdef, state_in)
1787+
merged_layer = nnx.merge(graphdef_in, state_in)
17941788
out_y, out_kv = merged_layer(y_in, *layer_args, kv_cache=kv_in, **layer_kwargs)
1795-
return out_y, out_kv, nnx.state(merged_layer)
1789+
state_out = nnx.state(merged_layer)
1790+
1791+
if dynamic_graph_init:
1792+
new_graphdef, _, _ = nnx.split(merged_layer, nnx.Param, ...)
1793+
return out_y, out_kv, state_out, new_graphdef
1794+
else:
1795+
return out_y, out_kv, state_out, graphdef_in
17961796

17971797
checkpointed_fn = jax.checkpoint(pure_layer_fn, policy=policy, prevent_cse=prevent_cse)
17981798

1799-
for lyr, layer in enumerate(self.layers):
1799+
for lyr in range(cfg.num_decoder_layers):
1800+
if self.is_deepseek:
1801+
if lyr < cfg.first_num_dense_layers:
1802+
layer = getattr(self, f"dense_layers_{lyr}", None)
1803+
else:
1804+
moe_idx = lyr - cfg.first_num_dense_layers
1805+
layer = getattr(self, f"moe_layers_{moe_idx}", None)
1806+
else:
1807+
layer = getattr(self, f"layers_{lyr}", None)
1808+
if layer is None and hasattr(self, "layers") and self.layers:
1809+
layer = self.layers[lyr]
1810+
1811+
if layer is None:
1812+
raise AttributeError(f"Could not locate decoder layer at index {lyr} in {self.__class__.__name__}")
1813+
18001814
graphdef, state = nnx.split(layer)
18011815
if kv_caches is not None:
18021816
if cfg.decoder_block in (DecoderBlockType.QWEN3_NEXT, DecoderBlockType.QWEN3_5):
@@ -1812,12 +1826,23 @@ def pure_layer_fn(graphdef, state_in, y_in, kv_in):
18121826
else:
18131827
kv_cache = None
18141828

1815-
input_tokens = decoder_input_tokens if cfg.engram_layers else None
1816-
if input_tokens is not None:
1817-
layer_kwargs["decoder_input_tokens"] = input_tokens
1829+
if cfg.remat_policy != "none":
1830+
y, kv_cache, new_state, new_graphdef = checkpointed_fn(graphdef, state, y, kv_cache)
1831+
else:
1832+
y, kv_cache, new_state, new_graphdef = pure_layer_fn(graphdef, state, y, kv_cache)
18181833

1819-
y, kv_cache, new_state = checkpointed_fn(graphdef, state, y, kv_cache)
1820-
nnx.update(layer, new_state)
1834+
if dynamic_graph_init:
1835+
new_layer = nnx.merge(new_graphdef, new_state)
1836+
if self.is_deepseek:
1837+
if lyr < cfg.first_num_dense_layers:
1838+
setattr(self, f"dense_layers_{lyr}", new_layer)
1839+
else:
1840+
moe_idx = lyr - cfg.first_num_dense_layers
1841+
setattr(self, f"moe_layers_{moe_idx}", new_layer)
1842+
else:
1843+
setattr(self, f"layers_{lyr}", new_layer)
1844+
else:
1845+
nnx.update(layer, new_state)
18211846

18221847
if kv_caches is not None and kv_cache is not None:
18231848
if cfg.decoder_block in (DecoderBlockType.QWEN3_NEXT, DecoderBlockType.QWEN3_5):
@@ -2024,7 +2049,7 @@ def _apply_gemma4_small_layers(
20242049
cache_index_of = gemma4_small.kv_cache_slot_map(layer_types, num_kv_shared)
20252050

20262051
for lyr in range(cfg.num_decoder_layers):
2027-
layer = self.layers[lyr]
2052+
layer = getattr(self, f"layers_{lyr}")
20282053
donor_idx = gemma4_small.kv_donor_layer_idx(lyr, layer_types, num_kv_shared)
20292054
is_donor = gemma4_small.is_kv_donor_layer(lyr, layer_types, num_kv_shared)
20302055

@@ -2069,6 +2094,63 @@ def _apply_gemma4_small_layers(
20692094

20702095
return y, kv_caches
20712096

2097+
def get_layers(self) -> list[nnx.Module]:
2098+
"""Returns all decoder layer modules/blocks in their forward-pass execution order."""
2099+
layers = []
2100+
seen = set()
2101+
2102+
def _add(module):
2103+
if module is not None and id(module) not in seen:
2104+
seen.add(id(module))
2105+
layers.append(module)
2106+
2107+
def _append_unscanned(prefix):
2108+
i = 0
2109+
while hasattr(self, f"{prefix}_{i}"):
2110+
_add(getattr(self, f"{prefix}_{i}"))
2111+
i += 1
2112+
2113+
def _append_scanned(name):
2114+
if hasattr(self, name):
2115+
val = getattr(self, name)
2116+
if name == "layers_remainder" and getattr(val, "num_of_layers", 0) == 0:
2117+
return
2118+
if isinstance(val, (nnx.Module, list)):
2119+
_add(val)
2120+
2121+
if self.is_deepseek:
2122+
_append_scanned("dense_layers")
2123+
_append_unscanned("dense_layers")
2124+
2125+
_append_scanned("moe_layers")
2126+
_append_unscanned("moe_layers")
2127+
2128+
_append_scanned("moe_layers_outside_pipeline")
2129+
_append_unscanned("moe_layers_outside_pipeline")
2130+
2131+
if hasattr(self, "pipeline_module"):
2132+
_add(getattr(self.pipeline_module, "layers", None))
2133+
else:
2134+
if hasattr(self, "pipeline_module"):
2135+
_add(getattr(self.pipeline_module, "layers", None))
2136+
2137+
_append_scanned("scanned_blocks") # Gemma 4
2138+
_append_scanned("layers")
2139+
_append_unscanned("layers")
2140+
2141+
_append_scanned("layers_remainder") # Gemma 3/4
2142+
2143+
_append_scanned("layers_outside_pipeline")
2144+
_append_unscanned("layers_outside_pipeline")
2145+
2146+
# Fallback for dynamic/chunked layer attributes (e.g. Engram dense_layers_0_3)
2147+
if not layers:
2148+
for k, m in vars(self).items():
2149+
if k.startswith(("dense_layers", "moe_layers", "layers", "scanned_blocks")) and isinstance(m, (nnx.Module, list)):
2150+
_add(m)
2151+
2152+
return layers
2153+
20722154

20732155
def decoder_as_linen(
20742156
config: Config,

src/maxtext/layers/nnx_wrappers.py

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -312,6 +312,57 @@ def __call__(
312312

313313
return out
314314

315+
def get_layers(self) -> list[Any]:
316+
"""Returns all decoder layer modules in execution order."""
317+
layers = []
318+
seen = set()
319+
320+
def _add(module):
321+
if module is not None and id(module) not in seen:
322+
seen.add(id(module))
323+
layers.append(module)
324+
325+
def _append_unscanned(prefix):
326+
i = 0
327+
while hasattr(self, f"{prefix}_{i}"):
328+
_add(getattr(self, f"{prefix}_{i}"))
329+
i += 1
330+
331+
def _append_scanned(name):
332+
if hasattr(self, name):
333+
val = getattr(self, name)
334+
if name == "layers_remainder" and getattr(val, "num_of_layers", 0) == 0:
335+
return
336+
_add(val)
337+
338+
_append_scanned("dense_layers")
339+
_append_unscanned("dense_layers")
340+
341+
_append_scanned("moe_layers")
342+
_append_unscanned("moe_layers")
343+
344+
_append_scanned("moe_layers_outside_pipeline")
345+
_append_unscanned("moe_layers_outside_pipeline")
346+
347+
if hasattr(self, "pipeline_module"):
348+
_add(getattr(self.pipeline_module, "layers", None))
349+
350+
_append_scanned("scanned_blocks") # Gemma 4
351+
_append_scanned("layers")
352+
_append_unscanned("layers")
353+
354+
_append_scanned("layers_remainder") # Gemma 3/4
355+
356+
_append_scanned("layers_outside_pipeline")
357+
_append_unscanned("layers_outside_pipeline")
358+
359+
if not layers:
360+
for k, m in vars(self).items():
361+
if k.startswith(("dense_layers", "moe_layers", "layers", "scanned_blocks")):
362+
_add(m)
363+
364+
return layers
365+
315366

316367
def linen_rngs_dict(linen_module: linen.Module, add_default: bool = False):
317368
"""Given a module, split out one of its every active RNG key collections."""

0 commit comments

Comments
 (0)