Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
83 changes: 43 additions & 40 deletions src/maxtext/layers/decoders.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,7 +33,12 @@
from maxtext.layers import mhc
from maxtext.layers import normalizations
from maxtext.layers import pipeline
from maxtext.layers.nnx_decoders import NNXDecoderLayer, NNXSequentialPipelineStage, NNXScannedPipelineStage
from maxtext.layers.nnx_decoders import (
NNXDecoderLayer,
NNXSequentialPipelineStage,
NNXScannedPipelineStage,
_make_single_layer_remat_stage_cls,
)
from maxtext.layers import quantizations
from maxtext.layers.attentions import attention_as_linen
from maxtext.layers.embeddings import attend_on_embedding, embed_as_linen, positional_embedding_as_linen
Expand Down Expand Up @@ -526,6 +531,7 @@ def get_scannable(normal_cls, scannable_cls):
DecoderBlockType.SIMPLE: [simple_layer.SimpleDecoderLayer],
DecoderBlockType.SIMPLE_MLP: [simple_layer.SimpleMlpDecoderLayer],
DecoderBlockType.DEEPSEEK: [deepseek.DeepSeekDenseLayer, deepseek.DeepSeekMoELayer],
DecoderBlockType.DEEPSEEK4: get_scannable(deepseek4.DeepSeek4DecoderLayer, deepseek4.DeepSeek4ScannableBlock),
DecoderBlockType.LLAMA4: get_scannable(llama4.Llama4DecoderLayer, llama4.Llama4ScannableBlock),
DecoderBlockType.OLMO3: get_scannable(olmo3.Olmo3DecoderLayer, olmo3.Olmo3ScannableBlock),
}
Expand Down Expand Up @@ -567,52 +573,49 @@ def _build_nnx_pipeline_stage(self, decoder_blocks, rngs):
cfg = self.config
base_stage_cls = decoder_blocks[1] if cfg.decoder_block == DecoderBlockType.DEEPSEEK else decoder_blocks[0]

# Per-stage-layer remat (+ params-only host-offload inside the stage) when the flag is set.
# apply_per_stage_remat is the boolean decision; per_stage_remat is the policy value
# (None == full remat for remat_policy='full', matching Linen nn.remat(policy=None)).
apply_per_stage_remat = cfg.set_remat_policy_on_layers_per_stage
per_stage_remat = self.get_remat_policy() if apply_per_stage_remat else None

if cfg.num_layers_per_pipeline_stage == 1:
if apply_per_stage_remat:
# Linen nn.remat parity: keep params TOP-LEVEL (no 'layers_0' nesting).
stage_cls = _make_single_layer_remat_stage_cls(base_stage_cls)
return stage_cls(
config=cfg,
mesh=self.mesh,
quant=self.quant,
model_mode=self.model_mode,
rngs=rngs,
remat_policy=per_stage_remat,
apply_remat=True,
)
return base_stage_cls(config=cfg, mesh=self.mesh, quant=self.quant, model_mode=self.model_mode, rngs=rngs)
elif cfg.scan_layers_per_stage:
return NNXScannedPipelineStage(
base_stage_cls, cfg.num_layers_per_pipeline_stage, cfg, self.mesh, self.quant, self.model_mode, rngs=rngs
)
return NNXSequentialPipelineStage(
base_stage_cls, cfg.num_layers_per_pipeline_stage, cfg, self.mesh, self.quant, self.model_mode, rngs=rngs
)

def get_pipeline_stage_module(self, decoder_blocks):
"""get pipeline stage module"""

def get_layer_to_pipeline(blocks, cfg):
if cfg.decoder_block == DecoderBlockType.DEEPSEEK:
return blocks[1] # return the sparse block
else:
return blocks[0]

cfg = self.config
base_stage = get_layer_to_pipeline(decoder_blocks, cfg)
if cfg.set_remat_policy_on_layers_per_stage:
policy = self.get_remat_policy()
base_stage = self.set_remat_policy([base_stage], policy)[0]
if cfg.num_layers_per_pipeline_stage == 1:
stage_module = base_stage(config=cfg, mesh=self.mesh, quant=self.quant, model_mode=self.model_mode)
elif cfg.scan_layers_per_stage:
stage_module = self.scan_decoder_layers(
cfg,
base_stage,
base_stage_cls,
cfg.num_layers_per_pipeline_stage,
"layers_per_stage",
cfg,
self.mesh,
in_axes_tuple=(nn.broadcast,) * 4,
model_mode=self.model_mode,
self.quant,
self.model_mode,
rngs=rngs,
remat_policy=per_stage_remat,
apply_remat=apply_per_stage_remat,
)
else:
stage_module = SequentialBlockDecoderLayers(
decoder_layer=base_stage,
num_decoder_layers=cfg.num_layers_per_pipeline_stage,
config=cfg,
mesh=self.mesh,
quant=self.quant,
model_mode=self.model_mode,
)
return stage_module
return NNXSequentialPipelineStage(
base_stage_cls,
cfg.num_layers_per_pipeline_stage,
cfg,
self.mesh,
self.quant,
self.model_mode,
rngs=rngs,
remat_policy=per_stage_remat,
apply_remat=apply_per_stage_remat,
)

def get_norm_layer(self, num_features: int):
"""get normalization layer (return type inherits from nn.Module)"""
Expand Down
Loading
Loading