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