Files
dela 8442f92c58 Keep LatentMoE routed bmm in activation dtype under bf16 autocast
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.
2026-08-25 20:22:51 +08:00

226 lines
8.7 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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()