diff --git a/kda/layers/latent_moe.py b/kda/layers/latent_moe.py index 52107cd..8b73e49 100644 --- a/kda/layers/latent_moe.py +++ b/kda/layers/latent_moe.py @@ -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 diff --git a/tests/integration/test_k3_arch.py b/tests/integration/test_k3_arch.py index a972a91..91649ab 100644 --- a/tests/integration/test_k3_arch.py +++ b/tests/integration/test_k3_arch.py @@ -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(