Route with σ(W_r x), Top-k(s+b), then L1-normalize over the selected set. Add Switch/GShard aux and router z-loss into train_k3 and train_sft. Wiki parquet URLs honor HF_ENDPOINT for mirrored downloads.
223 lines
8.6 KiB
Python
223 lines
8.6 KiB
Python
"""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 (K3 eq.13): s=σ(W_r x), Top-k(s+b), p_i = s_i / Σ_{j∈T} s_j.
|
||
Routed 执行: permute-dispatch, pad 到 [R, C, ℓ], 三次 bmm(SiTU 参数结构不变).
|
||
训练: Switch/GShard aux + router z-loss, 由 train loop 加到 CE 上.
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import torch
|
||
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,
|
||
aux_loss_coef: float = 1e-2,
|
||
z_loss_coef: float = 1e-3,
|
||
):
|
||
super().__init__()
|
||
self.latent_size = latent_size
|
||
self.n_routed = n_routed
|
||
self.top_k = top_k
|
||
self.aux_loss_coef = aux_loss_coef
|
||
self.z_loss_coef = z_loss_coef
|
||
|
||
self.down = nn.Linear(hidden_size, latent_size, bias=False) # W↓
|
||
self.router = nn.Linear(hidden_size, n_routed, bias=False) # logits → sigmoid
|
||
self.register_buffer("expert_bias", torch.zeros(n_routed), persistent=False)
|
||
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
|
||
self.last_aux_loss: torch.Tensor | None = None
|
||
self.last_z_loss: torch.Tensor | None = None
|
||
|
||
@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,
|
||
float(getattr(config, "moe_aux_loss_coef", 1e-2)),
|
||
float(getattr(config, "moe_z_loss_coef", 1e-3)),
|
||
)
|
||
|
||
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 _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
|
||
selected = scores.gather(-1, ids)
|
||
probs = selected / selected.sum(dim=-1, keepdim=True).clamp_min(1e-9)
|
||
return ids, probs
|
||
|
||
def _balancing_losses(
|
||
self, logits: torch.Tensor, ids: torch.Tensor
|
||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||
"""Switch/GShard aux on sigmoid scores; z-loss on raw logits."""
|
||
routed = self.n_routed
|
||
flat_logits = logits.reshape(-1, routed).float()
|
||
scores = torch.sigmoid(flat_logits)
|
||
counts = torch.bincount(ids.reshape(-1), minlength=routed).to(
|
||
dtype=scores.dtype
|
||
)
|
||
frac = counts / counts.sum().clamp_min(1.0)
|
||
prob_mean = scores.mean(dim=0)
|
||
aux = routed * (frac * prob_mean).sum()
|
||
z_loss = torch.logsumexp(flat_logits, dim=-1).square().mean()
|
||
return self.aux_loss_coef * aux, self.z_loss_coef * z_loss
|
||
|
||
def forward(self, x: torch.Tensor):
|
||
B, T, _ = x.shape
|
||
z = self.down(x) # [B, T, ℓ]
|
||
|
||
logits = self.router(x) # [B, T, n_routed]
|
||
ids, probs = self._route(logits)
|
||
self.last_route_ids = ids.detach()
|
||
if self.training and (self.aux_loss_coef != 0.0 or self.z_loss_coef != 0.0):
|
||
self.last_aux_loss, self.last_z_loss = self._balancing_losses(logits, ids)
|
||
else:
|
||
zero = logits.new_zeros(())
|
||
self.last_aux_loss = zero
|
||
self.last_z_loss = zero
|
||
|
||
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)
|
||
|
||
|
||
def moe_router_losses(model: nn.Module) -> tuple[torch.Tensor, torch.Tensor]:
|
||
"""Sum Switch aux and z-loss over every LatentMoE layer (0 if none)."""
|
||
auxes: list[torch.Tensor] = []
|
||
z_losses: list[torch.Tensor] = []
|
||
for module in model.modules():
|
||
if not isinstance(module, LatentMoE):
|
||
continue
|
||
if module.last_aux_loss is None or module.last_z_loss is None:
|
||
continue
|
||
auxes.append(module.last_aux_loss)
|
||
z_losses.append(module.last_z_loss)
|
||
if not auxes:
|
||
param = next(model.parameters(), None)
|
||
zero = param.new_zeros(()) if param is not None else torch.zeros(())
|
||
return zero, zero
|
||
return torch.stack(auxes).sum(), torch.stack(z_losses).sum()
|