"""Stable LatentMoE (K3): shared 全宽 + routed 半宽专家 + SiTU-GLU + Top-k. 对照 learning/kimi-k3-notes §Stable LatentMoE: z = W_down(x) [B, T, ℓ] ℓ = d/2 latent 接口宽 u = Σ_{i∈Top-k(x)} p_i E_i^rt(z) [B, T, ℓ] routed 专家只在 ℓ 上算 y = Σ_j E_j^sh(x) + W_up RMSNorm(u) [B, T, d] shared 全宽 SiTU-GLU: gate = β1·tanh(W_g x/β1)⊙σ(W_g x); up = β2·tanh(W_u x/β2) ||SiTU-GLU||_∞ ≤ β1·β2 (=100), 原点附近≈SwiGLU, 远端软饱和防低精度溢出. 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 import torch import torch.nn.functional as F from torch import nn from .rmsnorm import RMSNorm class SiTU(nn.Module): """SiTU-GLU expert: gate 支软上限 β1, up 支软上限 β2, 输出回到输入维.""" def __init__(self, dim_in: int, dim_ff: int, beta1: float = 4.0, beta2: float = 25.0): super().__init__() self.beta1, self.beta2 = beta1, beta2 self.w_g = nn.Linear(dim_in, dim_ff, bias=False) self.w_u = nn.Linear(dim_in, dim_ff, bias=False) self.w_o = nn.Linear(dim_ff, dim_in, bias=False) def forward(self, x: torch.Tensor): wg = self.w_g(x) g = self.beta1 * torch.tanh(wg / self.beta1) * torch.sigmoid(wg) u = self.beta2 * torch.tanh(self.w_u(x) / self.beta2) return self.w_o(g * u) class LatentMoE(nn.Module): def __init__( self, hidden_size: int, latent_size: int, n_routed: int, top_k: int, n_shared: int, d_ff: int, beta1: float = 4.0, beta2: float = 25.0, ): super().__init__() self.latent_size = latent_size self.n_routed = n_routed self.top_k = top_k self.down = nn.Linear(hidden_size, latent_size, bias=False) # W↓ self.router = nn.Linear(hidden_size, n_routed, bias=False) # Top-k logits self.shared = nn.ModuleList( [SiTU(hidden_size, d_ff, beta1, beta2) for _ in range(n_shared)] ) self.experts = nn.ModuleList( [SiTU(latent_size, d_ff, beta1, beta2) for _ in range(n_routed)] ) 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: return cls( config.hidden_size, config.moe_latent_size, config.n_routed, config.top_k, config.n_shared, config.moe_d_ff, config.situ_beta1, 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, ℓ] logits = self.router(x) # [B, T, n_routed] topk = torch.topk(logits, self.top_k, dim=-1) ids = topk.indices # [B, T, k] self.last_route_ids = ids.detach() probs = F.softmax(topk.values, dim=-1) # [B, T, k] 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)) def moe_route_frac(model: nn.Module) -> torch.Tensor | None: """Mean expert occupancy over LatentMoE layers from the last forward.""" hists: list[torch.Tensor] = [] n_routed: int | None = None for module in model.modules(): if not isinstance(module, LatentMoE) or module.last_route_ids is None: continue n_routed = module.n_routed ids = module.last_route_ids.reshape(-1) hists.append(torch.bincount(ids, minlength=n_routed).float()) if not hists or n_routed is None: return None stacked = torch.stack(hists).sum(0) return stacked / stacked.sum().clamp_min(1.0)