Python float scales and fp32 expert weights promoted SiTU outputs to fp32, so index_add mixed BFloat16 dest with Float source and crashed the 0.5b run. Cast packed weights and gate scalars to z.dtype.
226 lines
8.7 KiB
Python
226 lines
8.7 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)
|
||
|
||
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]).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)
|
||
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.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.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
|
||
|
||
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()
|