diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index 5669e012ae..31584650aa 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -318,6 +318,11 @@ topk_routing_group: -1 # number of top groups to route inputs. For EP, # all-to-all communication with compute. Currently only implemented with DeepSeek sparse layers. use_batch_split_schedule: false # a flag if splitting batch into micro-batches to hide communications that yields performance benefits. batch_split_factor: 1 # the factor by which to split the batch. Only used if use_batch_split_schedule is true. +use_lineage: false # a flag to use Lineage DeepSeek-V3 execution. +lineage_attention_sharding: "head" # attention sharding strategy for Lineage ('head' or 'sequence'). +lineage_activation_checkpointing: true # a flag to use activation checkpointing in Lineage layers. +lineage_capacity_factor: 0.5 # positive capacity factor determining the destination buffer size for Lineage sparse dispatch. If <= 0, falls back to capacity_factor or ragged_buffer_factor. +lineage_mesh_axes_mapping: {} # custom mapping from Lineage logical axes to mesh physical axes. # For complex architectures like llama4 there are repeated sets of # inhomogeneous layers. E.g. maverick uses [dense+rope, moe+rope, dense+rope, moe+nope] diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index 353656c665..90d2af4ef8 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -236,6 +236,7 @@ class ProfilerType(str, Enum): "deepseek3-671b", "deepseek3-671b-2dfsdp", "deepseek3-671b-batchsplit", + "deepseek3-671b-lineage", "deepseek3-test", "deepseek3-tiny", "deepseek3.2-671b", @@ -1127,6 +1128,42 @@ class DeepSeekMoE(BaseModel): 1, description="Factor by which to split the batch into micro-batches. Only used if use_batch_split_schedule is True.", ) + use_lineage: bool = Field( + False, + description="Whether to use Lineage DeepSeek-V3 execution.", + ) + lineage_attention_sharding: Literal["head", "sequence"] = Field( + "head", + description=("Attention sharding strategy for Lineage ('head' or 'sequence')."), + ) + lineage_activation_checkpointing: bool = Field( + True, + description="Whether to use activation checkpointing in Lineage layers.", + ) + lineage_capacity_factor: float = Field( + 0.5, + description=( + "Positive capacity factor determining the destination buffer size for" + " Lineage sparse dispatch. If <= 0, falls back to capacity_factor or" + " ragged_buffer_factor." + ), + ) + lineage_mesh_axes_mapping: dict[str, Any] = Field( + default_factory=dict, + description=("Custom mapping from Lineage logical axes to mesh physical axes."), + ) + + @model_validator(mode="after") + def validate_lineage(self) -> "DeepSeekMoE": + """Validates that Lineage DeepSeek-V3 execution requirements are met.""" + if self.use_lineage: + scan_layers = getattr(self, "scan_layers", None) + if scan_layers is not None and not scan_layers: + raise ValueError("use_lineage=True requires scan_layers=True.") + decoder_block = getattr(self, "decoder_block", None) + if decoder_block is not None and decoder_block != DecoderBlockType.DEEPSEEK: + raise ValueError(f"use_lineage=True requires decoder_block='deepseek', got decoder_block={decoder_block!r}.") + return self class Qwen3Next(BaseModel): diff --git a/src/maxtext/layers/decoders.py b/src/maxtext/layers/decoders.py index a9bacf692f..0198299d00 100644 --- a/src/maxtext/layers/decoders.py +++ b/src/maxtext/layers/decoders.py @@ -63,6 +63,11 @@ qwen3_5, simple_layer, ) + +try: + from maxtext.models import lineage_adapter +except ImportError: + lineage_adapter = None from maxtext.multimodal import utils as mm_utils from maxtext.utils.sharding import create_sharding from maxtext.utils import max_logging @@ -958,105 +963,19 @@ def __call__( else: if cfg.scan_layers: if cfg.decoder_block == DecoderBlockType.DEEPSEEK: - assert len(RemattedBlockLayers) == 2, "Scanned layers must have a length of 2 using deepseek." - layer_call_kwargs = { - "previous_chunk": previous_chunk, - "slot": slot, - } - dense_layer = RemattedBlockLayers[0] - moe_layer = RemattedBlockLayers[1] - if cfg.engram_layers: - original_dense_call = dense_layer.__call__ - original_moe_call = moe_layer.__call__ - dense_layer.__call__ = functools.partial(dense_layer.__call__, **layer_call_kwargs) - moe_layer.__call__ = functools.partial(moe_layer.__call__, **layer_call_kwargs) - - common_kwargs = { - "dense_layer": dense_layer, - "moe_layer": moe_layer, - "original_dense_call": original_dense_call, - "original_moe_call": original_moe_call, - "layer_call_kwargs": layer_call_kwargs, - "decoder_segment_ids": decoder_segment_ids, - "decoder_positions": decoder_positions, - "deterministic": deterministic, - "model_mode": model_mode, - "decoder_input_tokens": decoder_input_tokens, - "broadcast_args": broadcast_args, - } - - # Apply Dense Layers - y = self._apply_interleaved_scanned_layers( - y, - layer_type="dense", - start_idx=0, - end_idx=cfg.first_num_dense_layers, - engram_indices=cfg.engram_layers, - **common_kwargs, - ) - - # Apply MoE Layers - y = self._apply_interleaved_scanned_layers( - y, - layer_type="moe", - start_idx=cfg.first_num_dense_layers, - end_idx=cfg.num_decoder_layers, - engram_indices=cfg.engram_layers, - **common_kwargs, - ) - else: - dense_layer.__call__ = functools.partial(dense_layer.__call__, **layer_call_kwargs) - y, _ = self.scan_decoder_layers( - cfg, - dense_layer, - cfg.first_num_dense_layers, - "dense_layers", - mesh, - in_axes_tuple=(nn.broadcast,) * len(broadcast_args), - model_mode=model_mode, - )(y, *broadcast_args) - moe_layer.__call__ = functools.partial(moe_layer.__call__, **layer_call_kwargs) - num_moe_layers = cfg.num_decoder_layers - cfg.first_num_dense_layers - - # If batch-split schedule is used and initialization is complete, - # as detected by immutable params, use deepseek_batchsplit custom - # scan with initialized parameters. - if cfg.use_batch_split_schedule and not self.is_mutable_collection("params"): - # old version of batch-split that fully uses qwix quantization. - if cfg.quantization and cfg.use_qwix_quantization and not cfg.use_manual_quantization: - y = deepseek_batchsplit_fp8.scan_batch_split_layers( - y, - self.variables["params"]["moe_layers"], - decoder_positions, - decoder_segment_ids, - model_mode=model_mode, - mesh=mesh, - quant=self.quant, - cfg=cfg, - policy=policy, - ) - else: - # bf16 and fp8 code path for pure-JAX batch-split. - # fp8 code path supports both manual quantization and qwix - # quantization. - y = deepseek_batchsplit.scan_batch_split_layers( - y, - self.variables["params"]["moe_layers"], - decoder_positions, - mesh=mesh, - cfg=cfg, - num_layers=num_moe_layers, - ) - else: - y, _ = self.scan_decoder_layers( - cfg, - moe_layer, - num_moe_layers, - "moe_layers", - mesh, - in_axes_tuple=(nn.broadcast,) * len(broadcast_args), - model_mode=model_mode, - )(y, *broadcast_args) + y = self._apply_deepseek_scanned_blocks( + y, + RemattedBlockLayers, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + decoder_input_tokens, + previous_chunk, + slot, + broadcast_args, + policy, + ) elif cfg.decoder_block == DecoderBlockType.GEMMA3: bidirectional_mask_value = multimodal_input.bidirectional_mask if multimodal_input is not None else None y = self._apply_gemma3_scanned_blocks( @@ -1344,6 +1263,138 @@ def __call__( # and the raw hidden state needed for auxiliary tasks. return logits, hidden_state, kv_caches + def _apply_deepseek_scanned_blocks( + self, + y, + RemattedBlockLayers, + decoder_segment_ids, + decoder_positions, + deterministic, + model_mode, + decoder_input_tokens, + previous_chunk, + slot, + broadcast_args, + policy, + ): + """Applies DeepSeek scanned decoder blocks, handling dense and MoE layers.""" + cfg = self.config + mesh = self.mesh + + assert len(RemattedBlockLayers) == 2, "Scanned layers must have a length of 2 using deepseek." + layer_call_kwargs = { + "previous_chunk": previous_chunk, + "slot": slot, + } + dense_layer = RemattedBlockLayers[0] + moe_layer = RemattedBlockLayers[1] + if cfg.engram_layers: + original_dense_call = dense_layer.__call__ + original_moe_call = moe_layer.__call__ + dense_layer.__call__ = functools.partial(dense_layer.__call__, **layer_call_kwargs) + moe_layer.__call__ = functools.partial(moe_layer.__call__, **layer_call_kwargs) + + common_kwargs = { + "dense_layer": dense_layer, + "moe_layer": moe_layer, + "original_dense_call": original_dense_call, + "original_moe_call": original_moe_call, + "layer_call_kwargs": layer_call_kwargs, + "decoder_segment_ids": decoder_segment_ids, + "decoder_positions": decoder_positions, + "deterministic": deterministic, + "model_mode": model_mode, + "decoder_input_tokens": decoder_input_tokens, + "broadcast_args": broadcast_args, + } + + # Apply Dense Layers + y = self._apply_interleaved_scanned_layers( + y, + layer_type="dense", + start_idx=0, + end_idx=cfg.first_num_dense_layers, + engram_indices=cfg.engram_layers, + **common_kwargs, + ) + + # Apply MoE Layers + y = self._apply_interleaved_scanned_layers( + y, + layer_type="moe", + start_idx=cfg.first_num_dense_layers, + end_idx=cfg.num_decoder_layers, + engram_indices=cfg.engram_layers, + **common_kwargs, + ) + else: + num_moe_layers = cfg.num_decoder_layers - cfg.first_num_dense_layers + if getattr(cfg, "use_lineage", False) and not self.is_mutable_collection("params"): + y = lineage_adapter.run_lineage_dsv3( + inputs=y, + dense_params=self.variables["params"]["dense_layers"], + sparse_params=self.variables["params"]["moe_layers"], + decoder_positions=decoder_positions, + mesh=mesh, + cfg=cfg, + decoder_segment_ids=decoder_segment_ids, + num_dense_layers=cfg.first_num_dense_layers, + num_sparse_layers=num_moe_layers, + ) + else: + dense_layer.__call__ = functools.partial(dense_layer.__call__, **layer_call_kwargs) + y, _ = self.scan_decoder_layers( + cfg, + dense_layer, + cfg.first_num_dense_layers, + "dense_layers", + mesh, + in_axes_tuple=(nn.broadcast,) * len(broadcast_args), + model_mode=model_mode, + )(y, *broadcast_args) + moe_layer.__call__ = functools.partial(moe_layer.__call__, **layer_call_kwargs) + + # If batch-split schedule is used and initialization is complete, + # as detected by immutable params, use deepseek_batchsplit custom + # scan with initialized parameters. + if cfg.use_batch_split_schedule and not self.is_mutable_collection("params"): + # old version of batch-split that fully uses qwix quantization. + if cfg.quantization and cfg.use_qwix_quantization and not cfg.use_manual_quantization: + y = deepseek_batchsplit_fp8.scan_batch_split_layers( + y, + self.variables["params"]["moe_layers"], + decoder_positions, + decoder_segment_ids, + model_mode=model_mode, + mesh=mesh, + quant=self.quant, + cfg=cfg, + policy=policy, + ) + else: + # bf16 and fp8 code path for pure-JAX batch-split. + # fp8 code path supports both manual quantization and qwix + # quantization. + y = deepseek_batchsplit.scan_batch_split_layers( + y, + self.variables["params"]["moe_layers"], + decoder_positions, + mesh=mesh, + cfg=cfg, + num_layers=num_moe_layers, + ) + else: + y, _ = self.scan_decoder_layers( + cfg, + moe_layer, + num_moe_layers, + "moe_layers", + mesh, + in_axes_tuple=(nn.broadcast,) * len(broadcast_args), + model_mode=model_mode, + )(y, *broadcast_args) + return y + def _apply_gemma3_scanned_blocks( self, y, diff --git a/src/maxtext/layers/nnx_decoders.py b/src/maxtext/layers/nnx_decoders.py index 1e06020e72..6b9dfbd2c4 100644 --- a/src/maxtext/layers/nnx_decoders.py +++ b/src/maxtext/layers/nnx_decoders.py @@ -69,6 +69,11 @@ qwen3_custom, simple_layer, ) + +try: + from maxtext.models import lineage_adapter +except ImportError: + lineage_adapter = None from maxtext.multimodal import utils as mm_utils from maxtext.utils import max_logging, max_utils, maxtext_utils, maxtext_utils_nnx, sharding from maxtext.utils.sharding import create_sharding @@ -1925,77 +1930,15 @@ def __call__( ) elif cfg.scan_layers: if self.is_deepseek: - if cfg.engram_layers: - common_kwargs = { - "layer_kwargs": layer_kwargs, - "decoder_input_tokens": decoder_input_tokens, - } - - y = self._apply_interleaved_scanned_layers( - y, - "dense_layers", - 0, - cfg.first_num_dense_layers, - cfg.engram_layers, - *layer_args, - **common_kwargs, - ) - - y = self._apply_interleaved_scanned_layers( - y, - "moe_layers", - cfg.first_num_dense_layers, - cfg.num_decoder_layers, - cfg.engram_layers, - *layer_args, - **common_kwargs, - ) - else: - y, self.dense_layers, _ = self._apply_layers_sequentially( - self.dense_layers, - y, - *layer_args, - length=cfg.first_num_dense_layers, - **layer_kwargs, - ) - - num_moe = cfg.num_decoder_layers - cfg.first_num_dense_layers - - if cfg.use_batch_split_schedule: - policy = self.get_remat_policy() - mock_params = self._build_linen_params(self.moe_layers) - - if cfg.quantization and cfg.use_qwix_quantization and not cfg.use_manual_quantization: - y = deepseek_batchsplit_fp8.scan_batch_split_layers( - y, - mock_params, - decoder_positions, - decoder_segment_ids, - model_mode=model_mode, - mesh=self.mesh, - quant=self.quant, - cfg=cfg, - policy=policy, - ) - else: - # bf16 code path - y = deepseek_batchsplit.scan_batch_split_layers( - y, - mock_params, - decoder_positions, - mesh=self.mesh, - cfg=cfg, - num_layers=num_moe, - ) - else: - y, self.moe_layers, _ = self._apply_layers_sequentially( - self.moe_layers, - y, - *layer_args, - length=num_moe, - **layer_kwargs, - ) - + y = self._apply_deepseek_scanned_blocks( + y, + layer_args, + layer_kwargs, + decoder_positions, + decoder_segment_ids, + decoder_input_tokens, + model_mode, + ) elif self.is_deepseek4: y = self._apply_deepseek4_scanned_blocks( y, @@ -2277,6 +2220,104 @@ def _path_sort_key(path): return logits, hidden_state, kv_caches, expert_indices return logits, hidden_state, kv_caches + def _apply_deepseek_scanned_blocks( + self, + y, + layer_args, + layer_kwargs, + decoder_positions, + decoder_segment_ids, + decoder_input_tokens, + model_mode, + ): + """Applies DeepSeek V3 scanned decoder blocks, handling dense and MoE layers.""" + cfg = self.config + if cfg.engram_layers: + common_kwargs = { + "layer_kwargs": layer_kwargs, + "decoder_input_tokens": decoder_input_tokens, + } + + y = self._apply_interleaved_scanned_layers( + y, + "dense_layers", + 0, + cfg.first_num_dense_layers, + cfg.engram_layers, + *layer_args, + **common_kwargs, + ) + + y = self._apply_interleaved_scanned_layers( + y, + "moe_layers", + cfg.first_num_dense_layers, + cfg.num_decoder_layers, + cfg.engram_layers, + *layer_args, + **common_kwargs, + ) + else: + num_moe = cfg.num_decoder_layers - cfg.first_num_dense_layers + if getattr(cfg, "use_lineage", False): + dense_params = self.dense_layers + moe_params = self._build_linen_params(self.moe_layers) + y = lineage_adapter.run_lineage_dsv3( + inputs=y, + dense_params=dense_params, + sparse_params=moe_params, + decoder_positions=decoder_positions, + mesh=self.mesh, + cfg=cfg, + decoder_segment_ids=decoder_segment_ids, + num_dense_layers=cfg.first_num_dense_layers, + num_sparse_layers=num_moe, + ) + else: + y, self.dense_layers, _ = self._apply_layers_sequentially( + self.dense_layers, + y, + *layer_args, + length=cfg.first_num_dense_layers, + **layer_kwargs, + ) + + if cfg.use_batch_split_schedule: + policy = self.get_remat_policy() + mock_params = self._build_linen_params(self.moe_layers) + + if cfg.quantization and cfg.use_qwix_quantization and not cfg.use_manual_quantization: + y = deepseek_batchsplit_fp8.scan_batch_split_layers( + y, + mock_params, + decoder_positions, + decoder_segment_ids, + model_mode=model_mode, + mesh=self.mesh, + quant=self.quant, + cfg=cfg, + policy=policy, + ) + else: + # bf16 code path + y = deepseek_batchsplit.scan_batch_split_layers( + y, + mock_params, + decoder_positions, + mesh=self.mesh, + cfg=cfg, + num_layers=num_moe, + ) + else: + y, self.moe_layers, _ = self._apply_layers_sequentially( + self.moe_layers, + y, + *layer_args, + length=num_moe, + **layer_kwargs, + ) + return y + def _apply_deepseek4_scanned_blocks( self, y, diff --git a/src/maxtext/trainers/pre_train/train.py b/src/maxtext/trainers/pre_train/train.py index 31f13ba34c..712ee61fe3 100644 --- a/src/maxtext/trainers/pre_train/train.py +++ b/src/maxtext/trainers/pre_train/train.py @@ -712,7 +712,7 @@ def move(path, value): # The update from the scan is (num_moe_layers, num_experts) and must be transposed. decoder_layer = getattr(new_state.model.decoder, "moe_layers", new_state.model.decoder) decoder_bias = _find_gate_bias(decoder_layer) - if decoder_bias is not None: + if decoder_bias is not None and moe_bias_updates is not None: decoder_bias.value = decoder_bias.value + jnp.array(moe_bias_updates[0]) # 2. Update auxiliary MTP MoE layers (if enabled). diff --git a/tests/unit/configs_test.py b/tests/unit/configs_test.py index 2a7bd0f660..fdee2e43af 100644 --- a/tests/unit/configs_test.py +++ b/tests/unit/configs_test.py @@ -208,6 +208,9 @@ def test_gpt_configs(config_file): os.path.join(CONFIGS_DIR, "models", "deepseek3-671b-2dfsdp.yml"), os.path.join(CONFIGS_DIR, "models", "deepseek3-671b-batchsplit.yml"), ] +_LINEAGE_CONFIG = os.path.join(CONFIGS_DIR, "models", "deepseek3-671b-lineage.yml") +if os.path.exists(_LINEAGE_CONFIG): + DEEPSEEK_CONFIGS.append(_LINEAGE_CONFIG) @pytest.mark.parametrize("config_file", DEEPSEEK_CONFIGS) diff --git a/tests/unit/train_compile_test.py b/tests/unit/train_compile_test.py index 4719e90119..b6f467c025 100644 --- a/tests/unit/train_compile_test.py +++ b/tests/unit/train_compile_test.py @@ -543,6 +543,41 @@ def test_moe_deepseek_scanned_bf16(self): ) ) + @parameterized.named_parameters( + {"testcase_name": "tpu7x_64", "compile_topology": "tpu7x-64"}, + {"testcase_name": "tpu7x_128", "compile_topology": "tpu7x-128"}, + ) + def test_moe_deepseek_batchsplit_lineage(self, compile_topology): + try: + from maxtext.models import lineage_adapter # pylint: disable=g-import-not-at-top,import-outside-toplevel + + if lineage_adapter is None: + self.skipTest("Lineage adapter not available in open-source export.") + except ImportError: + self.skipTest("Lineage adapter not available in open-source export.") + temp_dir = gettempdir() + compiled_trainstep_file = os.path.join( + temp_dir, + f"test_moe_deepseek_batchsplit_lineage_{compile_topology}.pickle", + ) + train_compile_main( + ( + "", + get_test_config_path(), + f"compiled_trainstep_file={compiled_trainstep_file}", + f"compile_topology={compile_topology}", + "compile_topology_num_slices=1", + "model_name=deepseek3-test", + "use_batch_split_schedule=True", + "use_lineage=True", + "per_device_batch_size=2", + "max_target_length=128", + "dtype=bfloat16", + "weight_dtype=bfloat16", + "scan_layers=True", + ) + ) + def test_moe_emb_chunking(self): temp_dir = gettempdir() compiled_trainstep_file = os.path.join(temp_dir, "test_moe_emb_chunking.pickle") diff --git a/tests/unit/train_nnx_test.py b/tests/unit/train_nnx_test.py index 16155c6c7b..f06bbd6377 100644 --- a/tests/unit/train_nnx_test.py +++ b/tests/unit/train_nnx_test.py @@ -181,6 +181,14 @@ def __call__(self, decoder_input_tokens, decoder_positions, **kwargs): return out +class _TinyDecoderGateBiasNoSow(_TinyDecoder): + """`_TinyDecoder` with GateLogit bias but no sown moe_bias_updates.""" + + def __init__(self, vocab_size: int, hidden: int, rngs: nnx.Rngs): + super().__init__(vocab_size, hidden, rngs=rngs) + self.decoder = nnx.Dict({"gate": GateLogit(bias_shape=(3, 2))}) + + from maxtext.layers.attention_mla import indexer_losses @@ -544,6 +552,27 @@ def test_routed_bias_disabled_returns_none(self): _, aux = pre_train.loss_fn(model, cfg, data, None, None, is_train=True) self.assertIsNone(aux["moe_bias_updates"]) + def test_train_step_with_routed_bias_enabled_but_no_updates(self): + """Ensures train_step does not crash when routed_bias is True but no bias updates were sown.""" + cfg = _Cfg() + cfg.routed_bias = True + cfg.routed_bias_update_rate = 0.001 + model = _TinyDecoderGateBiasNoSow(cfg.vocab_size, hidden=4, rngs=nnx.Rngs(0)) + optimizer = nnx.Optimizer(model, optax.sgd(0.01), wrt=nnx.Param) + ts = train_state_nnx.TrainStateNNX(model, optimizer) + state_graphdef, state_pure = nnx.split(ts) + data = _make_data(batch=cfg.micro_batch_size_to_train_on, vocab=cfg.vocab_size) + + new_state, _ = pre_train.train_step( + state_graphdef, + cfg, + state_mesh_shardings=None, + params_shardings=None, + state=state_pure, + data=data, + ) + self.assertIsNotNone(new_state) + class TestRecordActivationMetricsParity(unittest.TestCase): """record_activation_metrics must yield identical metrics for Linen- and NNX-shaped intermediates.