Skip to content

Commit f218929

Browse files
committed
fix: fix missing implementation and add enhancement
Includes - Correctly positioned multimodal_input parameter in NNXDecoder.__call__ to match Linen signature. - Pass seq_len (y.shape[1]) to positional_embedding in NNXDecoder instead of passing y directly. - Dynamically extract metadata_axis_name via nnx.PARTITION_NAME for scan axes in _apply_layers_sequentially instead of hardcoding 'layers'.
1 parent 4a932fb commit f218929

6 files changed

Lines changed: 1231 additions & 106 deletions

File tree

.vscode/settings.json

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
{
22
"python.testing.pytestArgs": [],
3-
"python.testing.cwd": "${workspaceFolder}/MaxText",
3+
"python.testing.cwd": "${workspaceFolder}",
44
"python.testing.unittestEnabled": false,
55
"python.testing.pytestEnabled": true
66
}

src/maxtext/layers/decoders.py

Lines changed: 37 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)