Keep LatentMoE routed bmm in activation dtype under bf16 autocast

Python float scales and fp32 expert weights promoted SiTU outputs to
fp32, so index_add mixed BFloat16 dest with Float source and crashed
the 0.5b run. Cast packed weights and gate scalars to z.dtype.
This commit is contained in:
dela
2026-08-25 20:22:51 +08:00
parent 49aede9cb2
commit 8442f92c58
2 changed files with 23 additions and 7 deletions
+13
View File
@@ -187,6 +187,19 @@ def test_moe_unselected_experts_have_zero_grad():
assert torch.equal(param.grad, torch.zeros_like(param.grad))
def test_routed_u_bf16_index_add_matches_z_dtype():
"""Python float * bf16 promotes to fp32; index_add must still land in z.dtype."""
torch.manual_seed(0)
moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24)
z = torch.randn(2, 6, 16, dtype=torch.bfloat16)
logits = torch.randn(2, 6, 8, dtype=torch.bfloat16)
ids, probs = moe._route(logits)
u = moe._routed_u(z, ids, probs)
assert u.dtype == torch.bfloat16
u.float().square().mean().backward()
assert any(p.grad is not None and p.grad.abs().sum() > 0 for p in moe.experts[0].parameters())
def test_moe_aux_loss_penalizes_collapse():
torch.manual_seed(0)
moe = LatentMoE(