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
+10 -7
View File
@@ -124,16 +124,17 @@ class LatentMoE(nn.Module):
if C == 0 or eid.numel() == 0:
return u_flat.view(B, T, ell)
dtype = z.dtype
gathered = z.reshape(N, ell)[tok]
padded = torch.index_put(
gathered.new_zeros(R, C, ell), (eid, local_pos), gathered
)
w_g = torch.stack([e.w_g.weight for e in self.experts]) # [R, ff, ℓ]
w_u = torch.stack([e.w_u.weight for e in self.experts])
w_o = torch.stack([e.w_o.weight for e in self.experts]) # [R, ℓ, ff]
beta1 = self.experts[0].beta1
beta2 = self.experts[0].beta2
w_g = torch.stack([e.w_g.weight for e in self.experts]).to(dtype)
w_u = torch.stack([e.w_u.weight for e in self.experts]).to(dtype)
w_o = torch.stack([e.w_o.weight for e in self.experts]).to(dtype)
beta1 = padded.new_tensor(self.experts[0].beta1)
beta2 = padded.new_tensor(self.experts[0].beta2)
wg = torch.bmm(padded, w_g.transpose(-1, -2)) # [R, C, ff]
g = beta1 * torch.tanh(wg / beta1) * torch.sigmoid(wg)
@@ -141,14 +142,16 @@ class LatentMoE(nn.Module):
hidden = beta2 * torch.tanh(wu / beta2)
out = torch.bmm(g * hidden, w_o.transpose(-1, -2)) # [R, C, ℓ]
weighted = pw.unsqueeze(-1) * out[eid, local_pos]
weighted = pw.to(dtype).unsqueeze(-1) * out[eid, local_pos]
u_flat = u_flat.index_add(0, tok, weighted)
return u_flat.view(B, T, ell)
def _route(self, logits: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""K3: s=σ(l), T=TopK(s+b), p_i = s_i / Σ_{j∈T} s_j. Bias does not enter p."""
scores = torch.sigmoid(logits)
ids = torch.topk(scores + self.expert_bias, self.top_k, dim=-1).indices
ids = torch.topk(
scores + self.expert_bias.to(dtype=scores.dtype), self.top_k, dim=-1
).indices
selected = scores.gather(-1, ids)
probs = selected / selected.sum(dim=-1, keepdim=True).clamp_min(1e-9)
return ids, probs
+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(