@@ -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
20732155def decoder_as_linen (
20742156 config : Config ,
0 commit comments