LatentMoE: sparse permute-dispatch + padded bmm

Replace dense all-expert forward (16 experts × all tokens) with
permute-dispatch: sort token-expert pairs by expert id, pad to
[R, C, ℓ] (C = max tokens per expert), run 3 bmm calls for the
batched SiTU-GLU activation, then scatter-add weighted results back.

Routed expert FLOPs drop from R·N to R·C (C ≈ N·k/R under uniform
routing). SiTU parameter structure unchanged; checkpoint compatible.

Tests: sparse-vs-dense fwd/bwd equivalence, unselected expert zero
grad, last_capacity tracking.
This commit is contained in:
dela
2026-08-25 17:47:44 +08:00
parent 584f7e9e73
commit d1da0816f2
3 changed files with 131 additions and 11 deletions
+55 -9
View File
@@ -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))