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
10 changes: 5 additions & 5 deletions axlearn/common/mixture_of_experts.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
88 changes: 88 additions & 0 deletions axlearn/common/mixture_of_experts_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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."""
Expand Down