From 7ca454ff10c3841feac065208881717a8b278bf5 Mon Sep 17 00:00:00 2001 From: Sh0rck_Wang Date: Fri, 4 Sep 2026 14:40:36 +0800 Subject: [PATCH] Guard zero-sum MoE gate normalization --- axlearn/common/mixture_of_experts.py | 10 +-- axlearn/common/mixture_of_experts_test.py | 88 +++++++++++++++++++++++ 2 files changed, 93 insertions(+), 5 deletions(-) diff --git a/axlearn/common/mixture_of_experts.py b/axlearn/common/mixture_of_experts.py index b110d6b5e..1543b9320 100644 --- a/axlearn/common/mixture_of_experts.py +++ b/axlearn/common/mixture_of_experts.py @@ -828,7 +828,7 @@ def _get_normalized_gates(self, raw_gates: Tensor) -> Tensor: # Default softmax function already makes raw_gates normalized. return raw_gates else: - return raw_gates / raw_gates.sum(axis=-1, keepdims=True) + return raw_gates / (raw_gates.sum(axis=-1, keepdims=True) + 1e-9) def _process_logits(self, logits: Tensor) -> Tensor: """Converts input logits into float32, caps values, and optionally adds noise.""" @@ -953,7 +953,7 @@ def forward(self, logits: Tensor) -> NestedTensor: # Calculate the normalization factor for the selected experts. # [O, G, S, 1] - denom = jnp.sum(gate_weights, axis=-1, keepdims=True) + denom = jnp.sum(gate_weights, axis=-1, keepdims=True) + 1e-9 # Reshape gate_assignment from [O, G, S, K] to [O, G, K, S] then flatten to [O, G, K*S] gate_assignment = jnp.swapaxes(gate_assignment, 2, 3) @@ -1225,7 +1225,7 @@ def forward( else: seq_load_balance_loss = 0 # Caculate the normalization factor. - denom = jnp.sum(gate_weights, axis=-1, keepdims=True) + denom = jnp.sum(gate_weights, axis=-1, keepdims=True) + 1e-9 # Renormalize the gates of the selected expert. # [B, S, K] expert_weights = gate_weights / denom @@ -1350,7 +1350,7 @@ def forward( self.add_summary("seq_load_balance_loss", seq_load_balance_loss) else: seq_load_balance_loss = 0 - denom = jnp.sum(gate_weights, axis=-1, keepdims=True) + denom = jnp.sum(gate_weights, axis=-1, keepdims=True) + 1e-9 expert_weights = gate_weights / denom if cfg.routed_scaling_factor != 1: expert_weights *= cfg.routed_scaling_factor @@ -1424,7 +1424,7 @@ def forward( else: seq_load_balance_loss = 0 # Caculate the normalization factor. - denom = jnp.sum(gate_weights, axis=-1, keepdims=True) + denom = jnp.sum(gate_weights, axis=-1, keepdims=True) + 1e-9 # Renormalize the gates of the selected expert. # [B x S, K] expert_weights = gate_weights / denom diff --git a/axlearn/common/mixture_of_experts_test.py b/axlearn/common/mixture_of_experts_test.py index 4cc06a686..d522f7415 100644 --- a/axlearn/common/mixture_of_experts_test.py +++ b/axlearn/common/mixture_of_experts_test.py @@ -72,6 +72,14 @@ ) +def zero_score_fn(): + def fn(logits: jnp.ndarray, axis: int = -1) -> jnp.ndarray: + del axis + return jnp.zeros_like(logits) + + return fn + + # pylint: disable=no-self-use,protected-access class TransformerFeedForwardMoETest(TestCase): @parameterized.parameters( @@ -1333,6 +1341,86 @@ def test_capacity_factor_gating( gate_1_combine_tensor = jnp.asarray(gate_1.combine_tensor).astype(jnp.float32) assert jnp.array_equal(gate_0_combine_tensor, gate_1_combine_tensor) + def test_zero_score_fn_produces_finite_combine_tensor(self): + cfg = TopKGating.default_config().set( + name="test", + score_fn=config_for_function(zero_score_fn), + num_experts=4, + num_experts_per_token=2, + ) + layer: TopKGating = cfg.instantiate(parent=None) + state = layer.initialize_parameters_recursively(prng_key=jax.random.PRNGKey(123)) + logits = jnp.zeros((1, 2, 8, 4), dtype=jnp.float32) + + output, _ = F( + layer, + is_training=False, + prng_key=jax.random.PRNGKey(123), + state=state, + inputs=dict(logits=logits), + ) + + self.assertTrue(bool(jnp.isfinite(output.combine_tensor).all())) + self.assertTrue(bool(jnp.isfinite(output.dispatch_tensor).all())) + self.assertTrue(bool(jnp.isfinite(output.load_balance_loss))) + self.assertNestedAllClose(output.combine_tensor, jnp.zeros_like(output.combine_tensor)) + + +class TopKDropFreeGatingTest(TestCase): + def test_zero_score_fn_produces_finite_expert_weights(self): + cfg = TopKDropFreeGating.default_config().set( + name="test", + score_fn=config_for_function(zero_score_fn), + num_experts=4, + num_experts_per_token=2, + ) + layer: TopKDropFreeGating = cfg.instantiate(parent=None) + state = layer.initialize_parameters_recursively(prng_key=jax.random.PRNGKey(123)) + logits = jnp.zeros((2, 8, 4), dtype=jnp.float32) + + output, _ = F( + layer, + is_training=False, + prng_key=jax.random.PRNGKey(123), + state=state, + inputs=dict(logits=logits, seq_load_balance_loss_weight=None), + ) + + self.assertTrue(bool(jnp.isfinite(output.expert_weights).all())) + self.assertTrue(bool(jnp.isfinite(output.load_balance_loss))) + self.assertNestedAllClose(output.expert_weights, jnp.zeros_like(output.expert_weights)) + + +class TopKBiasGatingTest(TestCase): + @parameterized.parameters(False, True) + def test_zero_score_fn_produces_finite_expert_weights(self, use_group_routing: bool): + cfg = TopKBiasGating.default_config().set( + name="test", + score_fn=config_for_function(zero_score_fn), + num_experts=4, + num_experts_per_token=2, + routed_scaling_factor=1, + ) + if use_group_routing: + cfg.num_group_of_experts = 2 + cfg.topk_group = 1 + + layer: TopKBiasGating = cfg.instantiate(parent=None) + state = layer.initialize_parameters_recursively(prng_key=jax.random.PRNGKey(123)) + logits = jnp.zeros((2, 8, 4), dtype=jnp.float32) + + output, _ = F( + layer, + is_training=False, + prng_key=jax.random.PRNGKey(123), + state=state, + inputs=dict(logits=logits, seq_load_balance_loss_weight=None), + ) + + self.assertTrue(bool(jnp.isfinite(output.expert_weights).all())) + self.assertTrue(bool(jnp.isfinite(output.load_balance_loss))) + self.assertNestedAllClose(output.expert_weights, jnp.zeros_like(output.expert_weights)) + class TransformerFeedForwardDropFreeMoETest(TestCase): """Tests for TransformerFeedForwardDropFreeMoE layer."""