From 708a188795f26e0d1bffac2a54e4a8e557a46ce9 Mon Sep 17 00:00:00 2001 From: Daniel Mandragona Date: Fri, 28 Aug 2026 08:58:16 -0700 Subject: [PATCH] Integrate Lineage DeepSeek-V3 model (dsv3.py) into MaxText. Refactors the MaxText Lineage integration to use Lineage's high-level `dsv3.py` entry point as the unified integration point for DeepSeek-V3 execution (dense layers followed by sparse layers). Lineage Changes: * Visibility: Updated `learning/performance/lineage/BUILD` to grant package visibility to `//third_party/py/maxtext:__subpackages__`. * Launcher: Added `learning/performance/lineage/scripts/launch_dsv3_batchsplit_maxtext.sh` (XManager launch script for 512-core Ghostfish TPU v7x). * `ops.py`: Added recursive tuple translation and deduplication in `_map_axis` and `_map_set`. * `dsv3_mla.py`: Resolved redundant `physical_pspec` sharding transformations on physical specs. MaxText Changes: * `lineage_adapter.py`: Translates MaxText parameter structures into Lineage's typed dataclasses (`DSv3WeightsPytree` containing both dense MLA/MLP weights and sparse MoE weights). Manages physical mesh axis mappings, YaRN RoPE frequencies, Splash Attention kernel initialization, dynamic minimum capacity factor alignment for Pallas kernel tiles, and delegates execution directly to `dsv3.dsv3`. * Decoders: Integrated `LineageAdapter` (`run_lineage_dsv3`) into Linen `Decoder` and `NNXDecoder` to execute the full DeepSeek-V3 layer stack when `use_lineage=true`. * Cleanup: Removed vestigial remnants of the previous sparse-only integration (`lineage_sparse_adapter.py`, `use_lineage_sparse_layers` flag, and early-return branching in `deepseek.py`). Removed all defensive fallback branches to maintain a lean, canonical integration. * Configs: Updated `deepseek3-671b-lineage.yml`, `base.yml`, and `types.py` with `use_lineage` and related Lineage flags. * Export: Updated `copy.bara.sky` to exclude `lineage_adapter.py`, `deepseek3-671b-lineage.yml`, and Lineage tests from open-source export. * Tests: Added unit and TPU tests in `lineage_adapter_test.py` and `lineage_adapter_tpu_test.py`; updated existing MaxText test suites (`batchsplit_google_test`, `train_compile_google_test`). Tested: Verified via unit test suites (`lineage_adapter_test`, `batchsplit_google_test`, `train_compile_google_test`, `lineage_adapter_tpu_test`, `ops_test`), and XManager training runs xid/285946702 PiperOrigin-RevId: 972617925 --- src/maxtext/configs/base.yml | 5 + src/maxtext/configs/types.py | 37 ++++ src/maxtext/layers/decoders.py | 249 ++++++++++++++---------- src/maxtext/layers/nnx_decoders.py | 183 ++++++++++------- src/maxtext/trainers/pre_train/train.py | 2 +- tests/unit/configs_test.py | 3 + tests/unit/train_compile_test.py | 35 ++++ tests/unit/train_nnx_test.py | 29 +++ 8 files changed, 372 insertions(+), 171 deletions(-) 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.