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:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user