@@ -541,6 +541,7 @@ def get_scannable(normal_cls, scannable_cls):
541541 DecoderBlockType .SIMPLE : [simple_layer .SimpleDecoderLayer ],
542542 DecoderBlockType .SIMPLE_MLP : [simple_layer .SimpleMlpDecoderLayer ],
543543 DecoderBlockType .DEEPSEEK : [deepseek .DeepSeekDenseLayer , deepseek .DeepSeekMoELayer ],
544+ DecoderBlockType .DEEPSEEK4 : get_scannable (deepseek4 .DeepSeek4DecoderLayer , deepseek4 .DeepSeek4ScannableBlock ),
544545 DecoderBlockType .LLAMA4 : get_scannable (llama4 .Llama4DecoderLayer , llama4 .Llama4ScannableBlock ),
545546 DecoderBlockType .OLMO3 : get_scannable (olmo3 .Olmo3DecoderLayer , olmo3 .Olmo3ScannableBlock ),
546547 }
@@ -582,52 +583,49 @@ def _build_nnx_pipeline_stage(self, decoder_blocks, rngs):
582583 cfg = self .config
583584 base_stage_cls = decoder_blocks [1 ] if cfg .decoder_block == DecoderBlockType .DEEPSEEK else decoder_blocks [0 ]
584585
586+ # Per-stage-layer remat (+ params-only host-offload inside the stage) when the flag is set.
587+ # apply_per_stage_remat is the boolean decision; per_stage_remat is the policy value
588+ # (None == full remat for remat_policy='full', matching Linen nn.remat(policy=None)).
589+ apply_per_stage_remat = cfg .set_remat_policy_on_layers_per_stage
590+ per_stage_remat = self .get_remat_policy () if apply_per_stage_remat else None
591+
585592 if cfg .num_layers_per_pipeline_stage == 1 :
593+ if apply_per_stage_remat :
594+ return NNXSequentialPipelineStage (
595+ base_stage_cls ,
596+ 1 ,
597+ cfg ,
598+ self .mesh ,
599+ self .quant ,
600+ self .model_mode ,
601+ rngs = rngs ,
602+ remat_policy = per_stage_remat ,
603+ apply_remat = True ,
604+ )
586605 return base_stage_cls (config = cfg , mesh = self .mesh , quant = self .quant , model_mode = self .model_mode , rngs = rngs )
587606 elif cfg .scan_layers_per_stage :
588607 return NNXScannedPipelineStage (
589- base_stage_cls , cfg .num_layers_per_pipeline_stage , cfg , self .mesh , self .quant , self .model_mode , rngs = rngs
590- )
591- return NNXSequentialPipelineStage (
592- base_stage_cls , cfg .num_layers_per_pipeline_stage , cfg , self .mesh , self .quant , self .model_mode , rngs = rngs
593- )
594-
595- def get_pipeline_stage_module (self , decoder_blocks ):
596- """get pipeline stage module"""
597-
598- def get_layer_to_pipeline (blocks , cfg ):
599- if cfg .decoder_block == DecoderBlockType .DEEPSEEK :
600- return blocks [1 ] # return the sparse block
601- else :
602- return blocks [0 ]
603-
604- cfg = self .config
605- base_stage = get_layer_to_pipeline (decoder_blocks , cfg )
606- if cfg .set_remat_policy_on_layers_per_stage :
607- policy = self .get_remat_policy ()
608- base_stage = self .set_remat_policy ([base_stage ], policy )[0 ]
609- if cfg .num_layers_per_pipeline_stage == 1 :
610- stage_module = base_stage (config = cfg , mesh = self .mesh , quant = self .quant , model_mode = self .model_mode )
611- elif cfg .scan_layers_per_stage :
612- stage_module = self .scan_decoder_layers (
613- cfg ,
614- base_stage ,
608+ base_stage_cls ,
615609 cfg .num_layers_per_pipeline_stage ,
616- "layers_per_stage" ,
610+ cfg ,
617611 self .mesh ,
618- in_axes_tuple = (nn .broadcast ,) * 4 ,
619- model_mode = self .model_mode ,
620- )
621- else :
622- stage_module = SequentialBlockDecoderLayers (
623- decoder_layer = base_stage ,
624- num_decoder_layers = cfg .num_layers_per_pipeline_stage ,
625- config = cfg ,
626- mesh = self .mesh ,
627- quant = self .quant ,
628- model_mode = self .model_mode ,
612+ self .quant ,
613+ self .model_mode ,
614+ rngs = rngs ,
615+ remat_policy = per_stage_remat ,
616+ apply_remat = apply_per_stage_remat ,
629617 )
630- return stage_module
618+ return NNXSequentialPipelineStage (
619+ base_stage_cls ,
620+ cfg .num_layers_per_pipeline_stage ,
621+ cfg ,
622+ self .mesh ,
623+ self .quant ,
624+ self .model_mode ,
625+ rngs = rngs ,
626+ remat_policy = per_stage_remat ,
627+ apply_remat = apply_per_stage_remat ,
628+ )
631629
632630 def get_norm_layer (self , num_features : int ):
633631 """get normalization layer (return type inherits from nn.Module)"""
0 commit comments