From b074035bf8fda9ae9a8b2a143e6d0afd271fb5c7 Mon Sep 17 00:00:00 2001 From: Shuwen-Fang Date: Wed, 2 Sep 2026 02:08:57 +0000 Subject: [PATCH 1/2] Prevent MoE token dropping via per-layer dropless fallback RoutedMoE now falls back to a dropless (worst-case) ragged buffer and redoes just that layer's route+compute via jax.lax.cond when the tuned ragged_buffer_factor would otherwise drop tokens, gated behind retry_when_tokens_dropped. Also validates against num_moe_emb_chunks>0 (not supported), and surfaces overflow via a gated log line in training_loop_iteration. --- src/maxtext/configs/base.yml | 1 + src/maxtext/configs/types.py | 21 +++ src/maxtext/layers/moe.py | 169 ++++++++++++------- src/maxtext/trainers/pre_train/train.py | 16 ++ tests/unit/maxtext_utils_test.py | 9 + tests/unit/moe_test.py | 214 ++++++++++++++++++++++++ tests/unit/train_nnx_test.py | 45 +++++ 7 files changed, 413 insertions(+), 62 deletions(-) diff --git a/src/maxtext/configs/base.yml b/src/maxtext/configs/base.yml index 518ef8fc56..11899bad16 100644 --- a/src/maxtext/configs/base.yml +++ b/src/maxtext/configs/base.yml @@ -233,6 +233,7 @@ ragged_buffer_factor: -1.0 # a factor to determine the size of the ragged buffer # When set to 1.0 this buffer if set to the size assuming perfectly balanced. If the routing dictates # a size larger than this then tokens are dropped. # In general if ragged_buffer_factor > 0, the ragged_buffer_size is balanced_size * ragged_buffer_factor. +retry_when_tokens_dropped: false # retry a layer with worst-case ragged buffer size if tokens would otherwise be dropped. moe_expert_input_dim: -1 # feature dimension of the tokens entering the MoE expert blocks. base_moe_mlp_dim: -1 # intermediate dimension at MoE layer. For a fully MoE model, base_mlp_dim must be equal to base_moe_mlp_dim. load_balance_loss_weight: 0.0 # weight for the load balance loss diff --git a/src/maxtext/configs/types.py b/src/maxtext/configs/types.py index 362a97fa3e..90f432598f 100644 --- a/src/maxtext/configs/types.py +++ b/src/maxtext/configs/types.py @@ -873,6 +873,11 @@ class MoEGeneral(BaseModel): num_experts: PositiveInt = Field(1, description="The total number of experts in each MoE layer.") num_experts_per_tok: PositiveInt = Field(1, description="The number of experts to route each token to.") capacity_factor: float = Field(-1.0, description="Expert capacity factor. If < 0, no token dropping.") + retry_when_tokens_dropped: bool = Field( + False, + description="Whether a MoE layer falls back to a dropless (worst-case) buffer, redoing just that " + "layer's route+compute, if its ragged sort buffer would otherwise drop tokens.", + ) ragged_buffer_factor: float = Field( -1.0, description="Ragged buffer factor. If < 0, ragged buffer is worst case size.", @@ -3248,6 +3253,20 @@ def _validate_check_vma_is_supported(self): f"Found other ICI axes enabled: {active}." ) + def validate_retry_when_tokens_dropped(self): + """Validates prerequisites for the per-layer dropless fallback.""" + if self.retry_when_tokens_dropped: + if self.num_experts <= 1: + raise ValueError("retry_when_tokens_dropped=True requires num_experts > 1.") + if self.ragged_buffer_factor <= 0: + raise ValueError("retry_when_tokens_dropped=True requires ragged_buffer_factor > 0.0.") + if not self.use_ring_of_experts: + raise ValueError("retry_when_tokens_dropped=True is currently only supported with use_ring_of_experts=True.") + if not self.use_ragged_sort: + raise ValueError("retry_when_tokens_dropped=True requires use_ragged_sort=True.") + if self.num_moe_emb_chunks > 0: + raise ValueError("retry_when_tokens_dropped=True does not support num_moe_emb_chunks > 0.") + def validate_ragged_buffer_factor(self): if self.ragged_buffer_factor <= 0: return # Not using a ragged buffer factor @@ -3443,6 +3462,7 @@ def set_derived_and_validate_values(self) -> "MaxTextConfig": _ep_disabled_flags = { "use_random_routing": False, "use_ragged_sort": False, + "retry_when_tokens_dropped": False, "ragged_buffer_factor": -1.0, "use_ring_of_experts": False, "num_moe_emb_chunks": 0, @@ -4231,6 +4251,7 @@ def calculate_global_batch_sizes(per_device_batch_size, expansion_factor, num_de if self.model_name.startswith("deepseek4") and self.first_num_hash_layers > 0 and self.use_ring_of_experts: raise ValueError("DeepSeek V4 hash routing is currently not supported with ring of experts.") self.validate_ragged_buffer_factor() + self.validate_retry_when_tokens_dropped() self.validate_num_moe_emb_chunks() if self.enable_diloco and not self.pure_nnx: diff --git a/src/maxtext/layers/moe.py b/src/maxtext/layers/moe.py index 82d4b01e47..471f144295 100644 --- a/src/maxtext/layers/moe.py +++ b/src/maxtext/layers/moe.py @@ -102,6 +102,7 @@ class RouteOutput: bias_updates: Optional[jax.Array] # Shape [local experts], tracks number of local tokens routed to every local expert. local_group_sizes: Optional[jax.Array] = None + has_overflow: jax.Array = struct.field(default_factory=lambda: jnp.bool_(False)) def _truncate_matrix(all_shards_group_sizes: jax.Array, buffer_size: int) -> jax.Array: @@ -947,6 +948,7 @@ def permute( roll_to_expert_id=None, input_ids=None, forced_routed_experts=None, + force_dropless_buffer=False, ): """Permute tokens to group by expert to fit gmm call.""" # reshape inputs (batch, sequence, emb) to (batch * sequence, emb) @@ -999,7 +1001,7 @@ def permute( self.config.num_experts, ) # roll_to_expert_id is not directly used in the kernel, ep axis id is directly called - if self.config.ragged_buffer_factor > 0.0: + if not force_dropless_buffer and self.config.ragged_buffer_factor > 0.0: balanced_size = (bsz_times_seq_len // num_expert_parallelism) * self.num_experts_per_tok buffer_size = self.get_ragged_buffer_size( balanced_size, @@ -1058,7 +1060,7 @@ def permute( num_tokens = bsz_times_seq_len * self.num_experts_per_tok use_truncated_buffer = use_ragged_in_permute and buffer_size is not None and buffer_size < num_tokens - + has_overflow = jnp.bool_(False) if use_truncated_buffer: local_num_experts = self.config.num_experts // num_expert_parallelism shard_idx = jax.lax.axis_index(self._expert_parallelism_name) if num_expert_parallelism > 1 else 0 @@ -1069,9 +1071,14 @@ def permute( local_num_experts, axis=0, ) + local_overflow = (jnp.sum(local_group_size) > buffer_size).astype(jnp.int32) # Clamp local_group_size to buffer_size to ensure we don't exceed buffer # capacity by leveraging the helper _truncate_matrix. local_group_size = _truncate_matrix(local_group_size[:, None], buffer_size)[:, 0] + if num_expert_parallelism > 1: + has_overflow = jax.lax.psum(local_overflow, self._expert_parallelism_name) > 0 + else: + has_overflow = local_overflow > 0 expert_indices = jnp.arange(local_num_experts) sorted_experts = jnp.repeat( expert_indices, @@ -1096,6 +1103,7 @@ def permute( lb_loss, bias_updates, local_group_size, + has_overflow, ) def unpermute( @@ -1821,6 +1829,7 @@ def roe_ag_and_route( rngs, input_ids=None, forced_routed_experts=None, + force_dropless_buffer=False, ): # The ring-of-experts strategy first duplicates the inputs to all # expert shards, and then routes within each shard. @@ -1849,6 +1858,7 @@ def roe_ag_and_route( lb_loss, bias_updates, local_group_sizes, + has_overflow, ) = self.permute( x, logits, @@ -1858,6 +1868,7 @@ def roe_ag_and_route( rngs=rngs, input_ids=input_ids, forced_routed_experts=forced_routed_experts, + force_dropless_buffer=force_dropless_buffer, ) return ( x, @@ -1869,6 +1880,7 @@ def roe_ag_and_route( lb_loss=lb_loss, bias_updates=bias_updates, local_group_sizes=local_group_sizes, + has_overflow=has_overflow, ), RouteMetadata( expert_shard_id=expert_shard_id, @@ -1900,6 +1912,7 @@ def ra2a_and_route( lb_loss, bias_updates, local_group_sizes, + _, ) = self.permute( x, logits, @@ -1996,6 +2009,7 @@ def route( rngs, input_ids=None, forced_routed_experts=None, + force_dropless_buffer=False, ): """Performs both across device and within device token routing/sorting""" num_ep = self.get_expert_parallelism_size() @@ -2011,6 +2025,7 @@ def route( rngs, input_ids=input_ids, forced_routed_experts=forced_routed_experts, + force_dropless_buffer=force_dropless_buffer, ) else: return ra2a_and_route( @@ -2354,6 +2369,7 @@ def _moe_body( sharded_input_ids, rngs, forced_routed_experts=None, + force_dropless_buffer=False, ): batch_size, sequence_length, embed_dim = x.shape if self.config.num_moe_emb_chunks > 0: @@ -2379,6 +2395,7 @@ def _moe_body( rngs, input_ids=sharded_input_ids, forced_routed_experts=forced_routed_experts, + force_dropless_buffer=force_dropless_buffer, ) mask = jnp.arange(x.shape[0]) < valid_token_count(x, routing, route_metadata) @@ -2436,7 +2453,7 @@ def _moe_body( scatter_dimension=0, tiled=True, ) - return output, routing.lb_loss, routing.bias_updates + return output, routing.lb_loss, routing.bias_updates, routing.has_overflow if self.get_expert_parallelism_size() > 1: original_inputs_first_dim = batch_size * sequence_length * self.config.num_experts_per_tok @@ -2468,7 +2485,7 @@ def _moe_body( group_sizes=routing.group_sizes, ) - return output, routing.lb_loss, routing.bias_updates + return output, routing.lb_loss, routing.bias_updates, routing.has_overflow @functools.partial( jax.shard_map, @@ -2494,6 +2511,7 @@ def _moe_body( output_pspec, P(), # Handle None or replicate the output P(), # Handle None or replicate the output + P(), # has_overflow: replicated scalar, already all-reduced across expert shards ), check_vma=self.config.check_vma, ) @@ -2516,65 +2534,90 @@ def sparse_matmul_route_and_compute( # drops fsdp -> GSPMD inserts the boundary all-gather) and reused across all # chunks of the ring-of-experts pipeline below. n_chunks = self.config.num_moe_token_chunks - if n_chunks <= 1 or not self.config.use_ring_of_experts: - return _moe_body( - x, - logits, - pre_bias_logits, - w0, - w1, - wo, - w0_bias, - w1_bias, - wo_bias, - sharded_input_ids, - rngs, - forced_routed_experts, - ) - # Chunked ring-of-experts pipeline: split the per-shard tokens along the - # sequence dim into `n_chunks` data-independent chunks. Each chunk runs the - # full route -> GMM -> combine path; with no barrier between them XLA is - # free to overlap chunk (c+1)'s EP all-gather and chunk (c-1)'s - # reduce-scatter with chunk c's GMM compute. Token routing is per-token, so - # the main (lm) output is identical to n_chunks=1; only the aggregate - # load-balance loss / bias updates are averaged across chunks. - seq_len = x.shape[1] - chunk = seq_len // n_chunks - outs, lb_losses, bias_updates_list = [], [], [] - _prev = None - for c in range(n_chunks): - sl = slice(c * chunk, (c + 1) * chunk) - x_c = x[:, sl, :] - # Fence each chunk's input on the previous chunk's output to control XLA's - # scheduling and prevent it from interleaving/fusing the chunks -- forces - # sequential pipelining. Math is unchanged (the barrier is identity), so - # loss stays bit-exact. - if self.config.moe_chunk_barrier and _prev is not None: - x_c, _prev = jax.lax.optimization_barrier((x_c, _prev)) - out_c, lb_c, bu_c = _moe_body( - x_c, - logits[:, sl, :], - None if pre_bias_logits is None else pre_bias_logits[:, sl, :], - w0, - w1, - wo, - w0_bias, - w1_bias, - wo_bias, - None if sharded_input_ids is None else sharded_input_ids[:, sl], - rngs, - None if forced_routed_experts is None else forced_routed_experts[:, sl, :], + def _route_and_compute(force_dropless_buffer): + """Runs route+compute once; force_dropless_buffer=True redoes all n_chunks, not just the overflowing one(s).""" + if n_chunks <= 1 or not self.config.use_ring_of_experts: + return _moe_body( + x, + logits, + pre_bias_logits, + w0, + w1, + wo, + w0_bias, + w1_bias, + wo_bias, + sharded_input_ids, + rngs, + forced_routed_experts, + force_dropless_buffer=force_dropless_buffer, + ) + + # Chunked ring-of-experts pipeline: split the per-shard tokens along the + # sequence dim into `n_chunks` data-independent chunks. Each chunk runs the + # full route -> GMM -> combine path; with no barrier between them XLA is + # free to overlap chunk (c+1)'s EP all-gather and chunk (c-1)'s + # reduce-scatter with chunk c's GMM compute. Token routing is per-token, so + # the main (lm) output is identical to n_chunks=1; only the aggregate + # load-balance loss / bias updates are averaged across chunks. + seq_len = x.shape[1] + chunk = seq_len // n_chunks + outs, lb_losses, bias_updates_list, has_overflows = [], [], [], [] + _prev = None + for c in range(n_chunks): + sl = slice(c * chunk, (c + 1) * chunk) + x_c = x[:, sl, :] + # Fence each chunk's input on the previous chunk's output to control XLA's + # scheduling and prevent it from interleaving/fusing the chunks -- forces + # sequential pipelining. Math is unchanged (the barrier is identity), so + # loss stays bit-exact. + if self.config.moe_chunk_barrier and _prev is not None: + x_c, _prev = jax.lax.optimization_barrier((x_c, _prev)) + out_c, lb_c, bu_c, ov_c = _moe_body( + x_c, + logits[:, sl, :], + None if pre_bias_logits is None else pre_bias_logits[:, sl, :], + w0, + w1, + wo, + w0_bias, + w1_bias, + wo_bias, + None if sharded_input_ids is None else sharded_input_ids[:, sl], + rngs, + None if forced_routed_experts is None else forced_routed_experts[:, sl, :], + force_dropless_buffer=force_dropless_buffer, + ) + if self.config.moe_chunk_barrier: + _prev = out_c + outs.append(out_c) + lb_losses.append(lb_c) + bias_updates_list.append(bu_c) + has_overflows.append(ov_c) + output = jnp.concatenate(outs, axis=1) + lb_loss = None if lb_losses[0] is None else sum(lb_losses) / n_chunks + bias_updates = None if bias_updates_list[0] is None else sum(bias_updates_list) / n_chunks + has_overflow = jnp.any(jnp.stack(has_overflows)) + return output, lb_loss, bias_updates, has_overflow + + out, lb_loss, bias_updates, has_overflow = _route_and_compute(force_dropless_buffer=False) + if self.config.retry_when_tokens_dropped: + + def _retry_dropless(_): + retried_out, retried_lb_loss, retried_bias_updates, _ = _route_and_compute(force_dropless_buffer=True) + return retried_out, retried_lb_loss, retried_bias_updates + + def _use_tight_buffer_result(_): + return out, lb_loss, bias_updates + + out, lb_loss, bias_updates = jax.lax.cond( + has_overflow, + _retry_dropless, + _use_tight_buffer_result, + None, ) - if self.config.moe_chunk_barrier: - _prev = out_c - outs.append(out_c) - lb_losses.append(lb_c) - bias_updates_list.append(bu_c) - output = jnp.concatenate(outs, axis=1) - lb_loss = None if lb_losses[0] is None else sum(lb_losses) / n_chunks - bias_updates = None if bias_updates_list[0] is None else sum(bias_updates_list) / n_chunks - return output, lb_loss, bias_updates + return out, lb_loss, bias_updates, has_overflow if self.config.moe_fsdp_use_two_stage_all_gather: # Unshard on fsdp axis @@ -2630,7 +2673,7 @@ def sparse_matmul_route_and_compute( if wo_bias is not None: wo_bias = self._maybe_shard_with_pspec(wo_bias, wo_bias_pspec) - return sparse_matmul_route_and_compute( + output, lb_loss, bias_updates, has_overflow = sparse_matmul_route_and_compute( inputs, gate_logits, pre_bias_logits, @@ -2644,6 +2687,8 @@ def sparse_matmul_route_and_compute( self.rngs, forced_routed_experts, ) + self.sow(nnx.Intermediate, "moe_has_overflow", has_overflow) + return output, lb_loss, bias_updates def reshape_and_update_weights(self, weights, indices, safe_updates=False): """Reshape and update weights. diff --git a/src/maxtext/trainers/pre_train/train.py b/src/maxtext/trainers/pre_train/train.py index 692690db76..52e9cdc5c6 100644 --- a/src/maxtext/trainers/pre_train/train.py +++ b/src/maxtext/trainers/pre_train/train.py @@ -424,12 +424,19 @@ def loss_fn(model, config, data, dropout_rng, params, sparsity_state=None, is_tr if not is_train and config.mtp_eval_target_module > 0: intermediate_outputs["logits"] = logits + has_moe_overflow = False + if config.retry_when_tokens_dropped: + moe_overflows = maxtext_utils.collect_intermediates_by_suffix(intermediate_outputs, "moe_has_overflow") + if moe_overflows: + has_moe_overflow = jnp.any(jnp.array([jnp.any(x) for x in moe_overflows])) + aux = { "intermediate_outputs": intermediate_outputs, "xent_sum": xent_sum, "z_loss": total_z_loss, "total_weights": total_weights, "moe_lb_loss": moe_lb_loss, + "has_moe_overflow": has_moe_overflow, "indexer_loss": indexer_loss, "moe_bias_updates": moe_bias_updates, "mtp_moe_bias_updates": mtp_moe_bias_updates, @@ -576,6 +583,7 @@ def diff_wrapper(curr_params, custom_params, rest, config, data): xent_sum = aux["xent_sum"] total_weights = aux["total_weights"] moe_lb_loss = aux["moe_lb_loss"] + has_moe_overflow = aux.get("has_moe_overflow", False) indexer_loss = aux.get("indexer_loss", 0.0) z_loss = aux.get("z_loss", 0.0) moe_bias_updates = aux.get("moe_bias_updates") @@ -764,6 +772,7 @@ def move(path, value): metrics = { "scalar": scalar_metrics, "scalars": {}, + "has_moe_overflow": has_moe_overflow, } if getattr(config, "record_internal_nn_metrics", False): record_activation_metrics(metrics, intermediate_outputs, config) @@ -881,6 +890,13 @@ def training_loop_iteration( state = sharding.maybe_shard_with_name(state, state_mesh_shardings, shard_mode) state, metrics = p_train_step(state, example_batch, *step_rng_args) + if config.retry_when_tokens_dropped: + # DEBUG: remove the else branch once verified. + if bool(metrics.get("has_moe_overflow", False)): + max_logging.log(f"Step {step}: MoE ragged buffer overflow, layer(s) fell back to a dropless buffer.") + else: + max_logging.log(f"Step {step}: no MoE ragged buffer overflow.") + step_time_delta = datetime.datetime.now() - last_step_completion last_step_completion = datetime.datetime.now() diff --git a/tests/unit/maxtext_utils_test.py b/tests/unit/maxtext_utils_test.py index 3256230694..cc0611c059 100644 --- a/tests/unit/maxtext_utils_test.py +++ b/tests/unit/maxtext_utils_test.py @@ -1088,6 +1088,15 @@ def test_donate_argnums_is_zero(self): ) self.assertEqual(donate_argnums, 0) + def test_donate_argnums_is_zero_with_retry_when_tokens_dropped(self): + step = self._make_mock_step() + cfg = self._make_mock_config() + cfg.retry_when_tokens_dropped = True + _, _, _, _, donate_argnums = maxtext_utils.get_functional_train_with_signature( + step, "data_sharding", "state_shardings", "model", cfg + ) + self.assertEqual(donate_argnums, 0) + def test_functional_train_is_partial(self): """functional_train should partially apply model and config.""" received = {} diff --git a/tests/unit/moe_test.py b/tests/unit/moe_test.py index ad3618f8f9..228ee266db 100644 --- a/tests/unit/moe_test.py +++ b/tests/unit/moe_test.py @@ -1064,6 +1064,220 @@ def test_ragged_sort_single_sparsecore_ring_of_experts(self): def test_ragged_sort_single_sparsecore_no_ring_of_experts(self): self._run_ragged_sort_loss_and_grad(use_ring_of_experts=False, ragged_sort_use_single_sparsecore=True) + @pytest.mark.tpu_only + @parameterized.named_parameters( + ("overflow", 0.1, True), + ("dropless", -1.0, False), + ) + def test_ragged_sort_overflow_detection(self, ragged_buffer_factor, expect_overflow): + cfg = pyconfig.initialize( + [None, get_test_config_path()], + run_name=f"moe_overflow_detection_{ragged_buffer_factor}", + enable_checkpointing=False, + model_name="mixtral-8x7b", + override_model_config=True, + base_emb_dim=7168, + base_mlp_dim=256, + base_moe_mlp_dim=256, + dtype="bfloat16", + megablox=True, + sparse_matmul=True, + per_device_batch_size=4, + ici_expert_parallelism=2, + use_ring_of_experts=True, + max_target_length=128, + float32_gate_logits=True, + use_ragged_sort=True, + ragged_buffer_factor=ragged_buffer_factor, + ) + + rng = jax.random.PRNGKey(2345) + rng_model, rng_hidden_states = jax.random.split(rng) + device_count = jax.device_count() + hidden_states = jax.random.uniform( + rng_hidden_states, + (int(cfg.per_device_batch_size) * device_count, cfg.max_target_length, cfg.base_emb_dim), + dtype=cfg.dtype, + ) + + devices_array = maxtext_utils.create_device_mesh(cfg) + mesh = Mesh(devices_array, cfg.mesh_axes) + model = moe.get_routed_moe( + name="MoeBlock", + config=cfg, + num_experts=cfg.num_experts, + num_experts_per_tok=cfg.num_experts_per_tok, + mesh=mesh, + kernel_init=nd_dense_init(1.0, "fan_in", "truncated_normal"), + kernel_axes=("embed", "mlp"), + intermediate_dim=cfg.mlp_dim, + dtype=cfg.dtype, + ) + + with jax.set_mesh(mesh), nn_partitioning.axis_rules(cfg.logical_axis_rules): + variables = model.init({"params": rng_model, "dropout": rng_model}, hidden_states) + _, mutated = model.apply({"params": variables["params"]}, hidden_states, mutable=["intermediates"]) + + has_overflow = maxtext_utils.collect_intermediates_by_suffix(mutated, "moe_has_overflow") + self.assertTrue(has_overflow, "Expected a moe_has_overflow intermediate to be sown.") + any_overflow = bool(jnp.any(jnp.array([jnp.any(x) for x in has_overflow]))) + + if expect_overflow: + self.assertTrue(any_overflow, "Expected has_overflow=True with a tiny ragged_buffer_factor.") + else: + self.assertFalse(any_overflow, "Expected has_overflow=False with a dropless (worst-case) buffer.") + + def _build_retry_test_mesh(self): + cfg = pyconfig.initialize( + [None, get_test_config_path()], + run_name="moe_retry_test_probe", + enable_checkpointing=False, + model_name="mixtral-8x7b", + override_model_config=True, + ici_expert_parallelism=2, + ) + return Mesh(maxtext_utils.create_device_mesh(cfg), cfg.mesh_axes) + + def _build_retry_test_model(self, mesh, ragged_buffer_factor, retry_when_tokens_dropped): + """Builds a mixtral-8x7b RoutedMoE with the given ragged buffer/retry settings.""" + cfg = pyconfig.initialize( + [None, get_test_config_path()], + run_name=f"moe_retry_test_{ragged_buffer_factor}_{retry_when_tokens_dropped}", + enable_checkpointing=False, + model_name="mixtral-8x7b", + override_model_config=True, + base_emb_dim=7168, + base_mlp_dim=256, + base_moe_mlp_dim=256, + dtype="bfloat16", + megablox=True, + sparse_matmul=True, + per_device_batch_size=4, + ici_expert_parallelism=2, + use_ring_of_experts=True, + max_target_length=128, + float32_gate_logits=True, + use_ragged_sort=True, + ragged_buffer_factor=ragged_buffer_factor, + retry_when_tokens_dropped=retry_when_tokens_dropped, + ) + model = moe.get_routed_moe( + name="MoeBlock", + config=cfg, + num_experts=cfg.num_experts, + num_experts_per_tok=cfg.num_experts_per_tok, + mesh=mesh, + kernel_init=nd_dense_init(1.0, "fan_in", "truncated_normal"), + kernel_axes=("embed", "mlp"), + intermediate_dim=cfg.mlp_dim, + dtype=cfg.dtype, + ) + return cfg, model + + def _out_loss_and_grad(self, model, params, hidden_states): + def loss_fn(p): + out, lb_loss, _ = model.apply({"params": p}, hidden_states) + loss = jnp.mean(out.astype(jnp.float32) ** 2) + return loss + (lb_loss.astype(jnp.float32) if lb_loss is not None else 0.0), out + + (loss, out), grads = jax.jit(jax.value_and_grad(loss_fn, has_aux=True))(params) + return out, loss, grads + + @pytest.mark.tpu_only + def test_retry_when_tokens_dropped_layer_level_prevents_drops(self): + """A tiny buffer + the flag should match a dropless model's output and gradients; without it, they should diverge.""" + mesh = self._build_retry_test_mesh() + rng = jax.random.PRNGKey(2345) + rng_model, rng_hidden_states = jax.random.split(rng) + device_count = jax.device_count() + + cfg_dropless, model_dropless = self._build_retry_test_model( + mesh, ragged_buffer_factor=-1.0, retry_when_tokens_dropped=False + ) + hidden_states = jax.random.uniform( + rng_hidden_states, + ( + int(cfg_dropless.per_device_batch_size) * device_count, + cfg_dropless.max_target_length, + cfg_dropless.base_emb_dim, + ), + dtype=cfg_dropless.dtype, + ) + with jax.set_mesh(mesh), nn_partitioning.axis_rules(cfg_dropless.logical_axis_rules): + variables = model_dropless.init({"params": rng_model, "dropout": rng_model}, hidden_states) + out_dropless, _, grads_dropless = self._out_loss_and_grad(model_dropless, variables["params"], hidden_states) + + _, model_retry = self._build_retry_test_model(mesh, ragged_buffer_factor=0.1, retry_when_tokens_dropped=True) + with jax.set_mesh(mesh), nn_partitioning.axis_rules(cfg_dropless.logical_axis_rules): + out_retry, _, grads_retry = self._out_loss_and_grad(model_retry, variables["params"], hidden_states) + + _, model_no_retry = self._build_retry_test_model(mesh, ragged_buffer_factor=0.1, retry_when_tokens_dropped=False) + with jax.set_mesh(mesh), nn_partitioning.axis_rules(cfg_dropless.logical_axis_rules): + out_no_retry, _, _ = self._out_loss_and_grad(model_no_retry, variables["params"], hidden_states) + + # Sanity check: the buffer must actually force drops without the flag, else this test proves nothing. + self.assertFalse( + jnp.allclose(out_no_retry.astype(jnp.float32), out_dropless.astype(jnp.float32), rtol=1e-2, atol=1e-2), + msg="retry_when_tokens_dropped=False unexpectedly matches dropless -- buffer isn't forcing an overflow.", + ) + + assert_moe_close(out_retry, out_dropless, cfg_dropless.dtype) + for g_retry, g_dropless in zip(jax.tree_util.tree_leaves(grads_retry), jax.tree_util.tree_leaves(grads_dropless)): + assert_moe_close(g_retry, g_dropless, cfg_dropless.dtype) + + @pytest.mark.tpu_only + def test_retry_when_tokens_dropped_asymmetric_shard_overflow(self): + """Retry must fire correctly when only one EP shard overflows, not both.""" + mesh = self._build_retry_test_mesh() + rng = jax.random.PRNGKey(2345) + rng_model, rng_hidden_states = jax.random.split(rng) + device_count = jax.device_count() + + cfg_dropless, model_dropless = self._build_retry_test_model( + mesh, ragged_buffer_factor=-1.0, retry_when_tokens_dropped=False + ) + hidden_states = jax.random.uniform( + rng_hidden_states, + ( + int(cfg_dropless.per_device_batch_size) * device_count, + cfg_dropless.max_target_length, + cfg_dropless.base_emb_dim, + ), + dtype=cfg_dropless.dtype, + ) + # num_experts=8, ici_expert_parallelism=2 -> shard 0 owns experts [0, 4). + # Route every token to experts 0 and 1: shard 0 gets everything and + # overflows a tiny buffer; shard 1 gets nothing and never overflows. + forced_routed_experts = jnp.broadcast_to( + jnp.arange(cfg_dropless.num_experts_per_tok, dtype=jnp.int32), + hidden_states.shape[:2] + (cfg_dropless.num_experts_per_tok,), + ) + + with jax.set_mesh(mesh), nn_partitioning.axis_rules(cfg_dropless.logical_axis_rules): + variables = model_dropless.init( + {"params": rng_model, "dropout": rng_model}, hidden_states, forced_routed_experts=forced_routed_experts + ) + out_dropless, _, _ = model_dropless.apply( + {"params": variables["params"]}, hidden_states, forced_routed_experts=forced_routed_experts + ) + + _, model_no_retry = self._build_retry_test_model(mesh, ragged_buffer_factor=0.1, retry_when_tokens_dropped=False) + with jax.set_mesh(mesh), nn_partitioning.axis_rules(cfg_dropless.logical_axis_rules): + out_no_retry, _, _ = model_no_retry.apply( + {"params": variables["params"]}, hidden_states, forced_routed_experts=forced_routed_experts + ) + self.assertFalse( + jnp.allclose(out_no_retry.astype(jnp.float32), out_dropless.astype(jnp.float32), rtol=1e-2, atol=1e-2), + msg="retry_when_tokens_dropped=False unexpectedly matches dropless -- forced routing isn't overflowing shard 0.", + ) + + _, model_retry = self._build_retry_test_model(mesh, ragged_buffer_factor=0.1, retry_when_tokens_dropped=True) + with jax.set_mesh(mesh), nn_partitioning.axis_rules(cfg_dropless.logical_axis_rules): + out_retry, _, _ = model_retry.apply( + {"params": variables["params"]}, hidden_states, forced_routed_experts=forced_routed_experts + ) + assert_moe_close(out_retry, out_dropless, cfg_dropless.dtype) + @pytest.mark.tpu_only def test_moe_fsdp_two_stage_parallelism_tpu_only(self): # Use an imperative skip inside the test method instead of a static decorator. diff --git a/tests/unit/train_nnx_test.py b/tests/unit/train_nnx_test.py index 91d13e35fb..8e7c9f9d67 100644 --- a/tests/unit/train_nnx_test.py +++ b/tests/unit/train_nnx_test.py @@ -51,6 +51,7 @@ class _Cfg: indexer_loss_scaling_factor: float = 0.0 num_vocab_tiling: int = 1 num_experts: int = 1 + retry_when_tokens_dropped: bool = False routed_bias: bool = False routed_bias_update_rate: float = 0.0 mtp_num_layers: int = 0 @@ -169,6 +170,19 @@ def __call__(self, decoder_input_tokens, decoder_positions, **kwargs): return out +class _TinyDecoderMoEOverflow(_TinyDecoder): + """`_TinyDecoder` with a layer that sows moe_has_overflow, like RoutedMoE.sparse_matmul.""" + + def __init__(self, vocab_size: int, hidden: int, rngs: nnx.Rngs, has_overflow: bool): + super().__init__(vocab_size, hidden, rngs=rngs) + self.has_overflow = has_overflow + + def __call__(self, decoder_input_tokens, decoder_positions, **kwargs): + out = super().__call__(decoder_input_tokens, decoder_positions, **kwargs) + self.sow(nnx.Intermediate, "moe_has_overflow", jnp.bool_(self.has_overflow)) + return out + + from maxtext.layers.attention_mla import indexer_losses @@ -369,6 +383,37 @@ def test_train_step_with_gradient_clipping(self): self.assertTrue(jnp.isfinite(metrics["scalar"]["learning/loss"])) +class TestMoeOverflowLoggingNNX(unittest.TestCase): + """Covers train_step surfacing has_moe_overflow for training_loop_iteration's log line.""" + + def _build_state(self, has_overflow, retry_when_tokens_dropped): + cfg = _Cfg(retry_when_tokens_dropped=retry_when_tokens_dropped) + model = _TinyDecoderMoEOverflow(cfg.vocab_size, hidden=4, rngs=nnx.Rngs(0), has_overflow=has_overflow) + optimizer = nnx.Optimizer(model, optax.sgd(0.01), wrt=nnx.Param) + return cfg, train_state_nnx.TrainStateNNX(model, optimizer) + + def _metrics(self, has_overflow, retry_when_tokens_dropped): + cfg, ts = self._build_state(has_overflow, retry_when_tokens_dropped) + state_graphdef, state_pure = nnx.split(ts) + data = _make_data(batch=cfg.micro_batch_size_to_train_on, vocab=cfg.vocab_size) + _, metrics = pre_train.train_step( + state_graphdef, cfg, state_mesh_shardings=None, params_shardings=None, state=state_pure, data=data + ) + return metrics + + def test_surfaces_overflow_when_flag_on(self): + metrics = self._metrics(has_overflow=True, retry_when_tokens_dropped=True) + self.assertTrue(bool(metrics["has_moe_overflow"])) + + def test_no_overflow_when_flag_on(self): + metrics = self._metrics(has_overflow=False, retry_when_tokens_dropped=True) + self.assertFalse(bool(metrics["has_moe_overflow"])) + + def test_ignores_overflow_when_flag_off(self): + metrics = self._metrics(has_overflow=True, retry_when_tokens_dropped=False) + self.assertFalse(bool(metrics["has_moe_overflow"])) + + class TestEvalStepNNX(unittest.TestCase): """Cover the NNX branch of eval_step (lines 568-570).""" From 6abb4a202f6a764f0bc93b5d50bd5b650bf3ea13 Mon Sep 17 00:00:00 2001 From: Shuwen-Fang Date: Wed, 2 Sep 2026 21:08:52 +0000 Subject: [PATCH 2/2] Fix CI regressions from the MoE overflow logging plumbing - Only add has_moe_overflow to aux/metrics when retry_when_tokens_dropped is set, instead of unconditionally: an always-present key changed every model's compiled train_step output (breaking hlo_diff_test for dense models), and use getattr() for the config read so minimal test configs without the field don't hit an AttributeError. - Cast has_moe_overflow to float32: a raw bool broke DiLoCo's generic metric aggregation (jaxpr verifier type mismatch, i32 vs i1). - Update the shard_map mock in test_sparse_matmul_repairs_batch_specs_only_ without_expert_parallelism to the new 4-tuple sparse_matmul_route_and_ compute return (added has_overflow), and stub fake_moe.sow(). - Update NNX overflow-logging tests for the new absent-when-off semantics. Co-Authored-By: Claude Sonnet 5 --- src/maxtext/trainers/pre_train/train.py | 16 +++++++--------- tests/unit/moe_test.py | 3 ++- tests/unit/train_nnx_test.py | 8 ++++++-- 3 files changed, 15 insertions(+), 12 deletions(-) diff --git a/src/maxtext/trainers/pre_train/train.py b/src/maxtext/trainers/pre_train/train.py index 52e9cdc5c6..8dcf60a4a0 100644 --- a/src/maxtext/trainers/pre_train/train.py +++ b/src/maxtext/trainers/pre_train/train.py @@ -424,25 +424,22 @@ def loss_fn(model, config, data, dropout_rng, params, sparsity_state=None, is_tr if not is_train and config.mtp_eval_target_module > 0: intermediate_outputs["logits"] = logits - has_moe_overflow = False - if config.retry_when_tokens_dropped: - moe_overflows = maxtext_utils.collect_intermediates_by_suffix(intermediate_outputs, "moe_has_overflow") - if moe_overflows: - has_moe_overflow = jnp.any(jnp.array([jnp.any(x) for x in moe_overflows])) - aux = { "intermediate_outputs": intermediate_outputs, "xent_sum": xent_sum, "z_loss": total_z_loss, "total_weights": total_weights, "moe_lb_loss": moe_lb_loss, - "has_moe_overflow": has_moe_overflow, "indexer_loss": indexer_loss, "moe_bias_updates": moe_bias_updates, "mtp_moe_bias_updates": mtp_moe_bias_updates, "mtp_loss": mtp_loss, "batch_stats": (intermediate_outputs.get("batch_stats", None) if hasattr(intermediate_outputs, "get") else None), } + if getattr(config, "retry_when_tokens_dropped", False): + moe_overflows = maxtext_utils.collect_intermediates_by_suffix(intermediate_outputs, "moe_has_overflow") + has_moe_overflow = jnp.any(jnp.array([jnp.any(x) for x in moe_overflows])) if moe_overflows else False + aux["has_moe_overflow"] = jnp.asarray(has_moe_overflow, dtype=jnp.float32) return loss, aux @@ -583,7 +580,7 @@ def diff_wrapper(curr_params, custom_params, rest, config, data): xent_sum = aux["xent_sum"] total_weights = aux["total_weights"] moe_lb_loss = aux["moe_lb_loss"] - has_moe_overflow = aux.get("has_moe_overflow", False) + has_moe_overflow = aux.get("has_moe_overflow") indexer_loss = aux.get("indexer_loss", 0.0) z_loss = aux.get("z_loss", 0.0) moe_bias_updates = aux.get("moe_bias_updates") @@ -772,8 +769,9 @@ def move(path, value): metrics = { "scalar": scalar_metrics, "scalars": {}, - "has_moe_overflow": has_moe_overflow, } + if has_moe_overflow is not None: + metrics["has_moe_overflow"] = has_moe_overflow if getattr(config, "record_internal_nn_metrics", False): record_activation_metrics(metrics, intermediate_outputs, config) diff --git a/tests/unit/moe_test.py b/tests/unit/moe_test.py index 228ee266db..b49e7f3a25 100644 --- a/tests/unit/moe_test.py +++ b/tests/unit/moe_test.py @@ -466,6 +466,7 @@ def test_sparse_matmul_repairs_batch_specs_only_without_expert_parallelism(exper *(original_batch_partition if axis == "activation_batch" else None for axis in logical_axes) ) fake_moe._maybe_shard_with_pspec = lambda value, _pspec, **_kwargs: value # pylint: disable=protected-access + fake_moe.sow = lambda *_args, **_kwargs: None inputs = SimpleNamespace(shape=(4, 1024, 2048)) gate_logits = SimpleNamespace(shape=(4, 1024, 256)) @@ -478,7 +479,7 @@ def fake_shard_map(function, *, mesh, in_specs, out_specs, check_vma): del function, mesh, check_vma captured["in_specs"] = in_specs captured["out_specs"] = out_specs - return lambda x, *_args: (x, None, None) + return lambda x, *_args: (x, None, None, jnp.bool_(False)) with mock.patch.object(jax, "shard_map", side_effect=fake_shard_map): output, _, _ = moe.RoutedMoE.sparse_matmul( diff --git a/tests/unit/train_nnx_test.py b/tests/unit/train_nnx_test.py index 8e7c9f9d67..a064b98c4b 100644 --- a/tests/unit/train_nnx_test.py +++ b/tests/unit/train_nnx_test.py @@ -409,9 +409,13 @@ def test_no_overflow_when_flag_on(self): metrics = self._metrics(has_overflow=False, retry_when_tokens_dropped=True) self.assertFalse(bool(metrics["has_moe_overflow"])) - def test_ignores_overflow_when_flag_off(self): + def test_key_absent_when_flag_off(self): + """Key must be absent (not just False) when off: an always-present key would + + change every model's compiled train_step output, even non-MoE ones. + """ metrics = self._metrics(has_overflow=True, retry_when_tokens_dropped=False) - self.assertFalse(bool(metrics["has_moe_overflow"])) + self.assertNotIn("has_moe_overflow", metrics) class TestEvalStepNNX(unittest.TestCase):