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