LatentMoE: K3 sigmoid routing and Switch aux/z-loss

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.
This commit is contained in:
dela
2026-08-25 19:50:07 +08:00
parent d1da0816f2
commit 7a12f61de1
8 changed files with 381 additions and 82 deletions
+72 -14
View File
@@ -9,13 +9,14 @@ 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).
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
import torch.nn.functional as F
from torch import nn
from .rmsnorm import RMSNorm
@@ -24,7 +25,9 @@ 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):
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)
@@ -49,14 +52,19 @@ class LatentMoE(nn.Module):
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) # Top-k logits
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)]
)
@@ -64,9 +72,11 @@ class LatentMoE(nn.Module):
[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.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:
@@ -79,6 +89,8 @@ class LatentMoE(nn.Module):
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(
@@ -123,28 +135,56 @@ class LatentMoE(nn.Module):
beta1 = self.experts[0].beta1
beta2 = self.experts[0].beta2
wg = torch.bmm(padded, w_g.transpose(-1, -2)) # [R, C, ff]
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, ℓ]
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, ℓ]
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]
logits = self.router(x) # [B, T, n_routed]
ids, probs = self._route(logits)
self.last_route_ids = ids.detach()
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
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]
shared_out = torch.stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
return shared_out + self.up(self.norm(u))
@@ -162,3 +202,21 @@ def moe_route_frac(model: nn.Module) -> torch.Tensor | 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()