diff --git a/kda/layers/latent_moe.py b/kda/layers/latent_moe.py index 1a719d4..31dfe1d 100644 --- a/kda/layers/latent_moe.py +++ b/kda/layers/latent_moe.py @@ -10,6 +10,7 @@ SiTU-GLU: gate = β1·tanh(W_g x/β1)⊙σ(W_g x); up = β2·tanh(W_u x/β2) E: R^in → R^in (内部中间维 d_ff). Router: Top-k logits 基于全宽 x (笔记 Topk(x)); 归一化权重取 softmax(topk). +Routed 执行: permute-dispatch, pad 到 [R, C, ℓ], 三次 bmm(SiTU 参数结构不变). """ from __future__ import annotations @@ -65,6 +66,7 @@ class LatentMoE(nn.Module): self.norm = RMSNorm(latent_size) self.up = nn.Linear(latent_size, hidden_size, bias=False) # W↑ self.last_route_ids: torch.Tensor | None = None + self.last_capacity: int = 0 @classmethod def from_config(cls, config) -> LatentMoE: @@ -79,6 +81,58 @@ class LatentMoE(nn.Module): config.situ_beta2, ) + def _routed_u( + self, z: torch.Tensor, ids: torch.Tensor, probs: torch.Tensor + ) -> torch.Tensor: + """Permute-dispatch + pad to [R, C, ℓ] + 3 bmm + scatter-add. + + ``C = max(counts)``: FLOPs are ``R·C``, not ``sum(counts)``. Padding + slots are not gathered, so they contribute zero gradient. Empty + experts stay in the stacked weights (padded grouped GEMM). + """ + B, T, ell = z.shape + N = B * T + R, k = self.n_routed, self.top_k + device = z.device + + tok = torch.arange(N, device=device).unsqueeze(1).expand(N, k).reshape(-1) + eid = ids.reshape(-1) + pw = probs.reshape(-1) + + order = eid.argsort(stable=True) + tok, eid, pw = tok[order], eid[order], pw[order] + + counts = torch.bincount(eid, minlength=R) + offsets = counts.cumsum(0) - counts + local_pos = torch.arange(N * k, device=device) - offsets[eid] + C = int(counts.max().item()) if eid.numel() else 0 + self.last_capacity = C + + u_flat = z.new_zeros(N, ell) + if C == 0 or eid.numel() == 0: + return u_flat.view(B, T, ell) + + 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 + + wg = torch.bmm(padded, w_g.transpose(-1, -2)) # [R, C, ff] + g = beta1 * torch.tanh(wg / beta1) * torch.sigmoid(wg) + wu = torch.bmm(padded, w_u.transpose(-1, -2)) + 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] + u_flat = u_flat.index_add(0, tok, weighted) + return u_flat.view(B, T, ell) + def forward(self, x: torch.Tensor): B, T, _ = x.shape z = self.down(x) # [B, T, ℓ] @@ -89,15 +143,7 @@ class LatentMoE(nn.Module): self.last_route_ids = ids.detach() probs = F.softmax(topk.values, dim=-1) # [B, T, k] - # 向量化 routed: 预计算全部专家输出, 按 token 的 Top-k id 取 - all_out = torch.stack([e(z) for e in self.experts]) # [R, B, T, ℓ] - all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, self.n_routed, self.latent_size) - u = torch.zeros(B, T, self.latent_size, device=x.device, dtype=x.dtype) - for i in range(self.top_k): - idx = ids[:, :, i].reshape(B * T) # [B*T] - sel = all_out[torch.arange(B * T, device=x.device), idx] # [B*T, ℓ] - u += probs[:, :, i : i + 1] * sel.reshape(B, T, self.latent_size) - + u = self._routed_u(z, ids, probs) shared_out = torch.stack([e(x) for e in self.shared]).sum(0) # [B, T, d] return shared_out + self.up(self.norm(u)) diff --git a/kda/models/k3_config.py b/kda/models/k3_config.py index ce09eea..e60038e 100644 --- a/kda/models/k3_config.py +++ b/kda/models/k3_config.py @@ -73,7 +73,7 @@ class K3Config: if name == "toy": return cls() if name in {"0.5b", "500m"}: - # H * head_dim == hidden. Routed 16: LatentMoE still runs every expert. + # H * head_dim == hidden. Routed 16 Top-2; LatentMoE padded bmm. # ~482M with tied Qwen3 embeddings. 6×(3 KDA + 1 MLA). return cls( hidden_size=768, diff --git a/tests/integration/test_k3_arch.py b/tests/integration/test_k3_arch.py index d80f895..cab5589 100644 --- a/tests/integration/test_k3_arch.py +++ b/tests/integration/test_k3_arch.py @@ -88,8 +88,26 @@ def test_preset_0_5b_schedule(): assert cfg.layer_specs()[3] == ("mla", "moe") +def _dense_moe_forward(moe: LatentMoE, x: torch.Tensor) -> torch.Tensor: + """Old dense path: run every routed expert, then gather top-k.""" + z = moe.down(x) + logits = moe.router(x) + topk = torch.topk(logits, moe.top_k, dim=-1) + ids = topk.indices + probs = F.softmax(topk.values, dim=-1) + all_out = torch.stack([expert(z) for expert in moe.experts]) + B, T, _ = x.shape + all_out = all_out.permute(1, 2, 0, 3).reshape(B * T, moe.n_routed, moe.latent_size) + u = z.new_zeros(B, T, moe.latent_size) + for i in range(moe.top_k): + idx = ids[:, :, i].reshape(B * T) + sel = all_out[torch.arange(B * T, device=x.device), idx] + u = u + probs[:, :, i : i + 1] * sel.reshape(B, T, moe.latent_size) + shared = torch.stack([expert(x) for expert in moe.shared]).sum(0) + return shared + moe.up(moe.norm(u)) + + def test_moe_router_activates_topk_only(): - from kda.layers.latent_moe import LatentMoE torch.manual_seed(3) moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24) x = torch.randn(2, 6, 32) @@ -112,6 +130,62 @@ def test_moe_router_activates_topk_only(): torch.testing.assert_close(y, expected_y, atol=1e-5, rtol=1e-5) assert moe.last_route_ids is not None assert moe.last_route_ids.shape[-1] == moe.top_k + counts = torch.bincount(moe.last_route_ids.reshape(-1), minlength=moe.n_routed) + assert moe.last_capacity == int(counts.max()) + + +def test_moe_sparse_matches_dense_fwd_bwd(): + torch.manual_seed(3) + moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24) + x = torch.randn(2, 6, 32) + y_sparse = moe(x) + y_dense = _dense_moe_forward(moe, x) + torch.testing.assert_close(y_sparse, y_dense, atol=1e-5, rtol=1e-5) + + moe.zero_grad(set_to_none=True) + xs = x.detach().requires_grad_(True) + moe(xs).square().mean().backward() + grads_s = { + name: param.grad.detach().clone() + for name, param in moe.named_parameters() + if param.grad is not None + } + dx_s = xs.grad.detach().clone() + + moe.zero_grad(set_to_none=True) + xd = x.detach().requires_grad_(True) + _dense_moe_forward(moe, xd).square().mean().backward() + grads_d = { + name: param.grad.detach().clone() + for name, param in moe.named_parameters() + if param.grad is not None + } + dx_d = xd.grad.detach().clone() + + torch.testing.assert_close(dx_s, dx_d, atol=1e-5, rtol=1e-5) + assert grads_s.keys() == grads_d.keys() + for name in grads_s: + torch.testing.assert_close(grads_s[name], grads_d[name], atol=1e-5, rtol=1e-5) + + +def test_moe_unselected_experts_have_zero_grad(): + torch.manual_seed(0) + moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24) + with torch.no_grad(): + moe.router.weight.zero_() + moe.router.weight[0] = 1.0 + moe.router.weight[1] = 0.5 + x = torch.ones(2, 4, 32) + moe(x).square().mean().backward() + selected = set(moe.last_route_ids.reshape(-1).tolist()) + assert selected == {0, 1} + for idx, expert in enumerate(moe.experts): + for param in (expert.w_g.weight, expert.w_u.weight, expert.w_o.weight): + assert param.grad is not None + if idx in selected: + assert param.grad.abs().sum() > 0 + else: + assert torch.equal(param.grad, torch.zeros_like(param.grad)) def test_k3_causal_future_does_not_change_past_logits():