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
+144 -40
View File
@@ -1,8 +1,9 @@
% teach:
% gap: 读者知道 MoE 的 top-k 路由但不知道 LatentMoE 的 latent 接口和 SiTU-GLU
% takeaway: LatentMoE 通过 latent 接口把 routed 专家限制在 ℓ=d/2 上算, SiTU-GLU 用软上限防溢出
% jump: 为什么 routed 专家在 latent 空间而 shared 在全宽?省参数
% omit: load balancing loss
% gap: 读者知道 MoE 的 top-k 路由但不知道 LatentMoE 的 latent 接口、稀疏 dispatch 和负载均衡损失
% takeaway: LatentMoE = shared 全宽 + routed 半宽 latent;路由用 K3 sigmoid-TopK + L1 归一化;
% 执行用 permute-dispatch + padded bmm(每个 token 只算 k 个专家);训练加 Switch/GShard aux + z-loss
% jump: 为什么不用 dense stack 算全部专家?稀疏 dispatch 让算力只随 k 不随 R 涨
% omit: none
\section{SiTU-GLU 与 Stable LatentMoE}
\splabel{C5}
@@ -54,7 +55,7 @@ class SiTU(nn.Module):
\textbf{Shared 专家} & $n_{\mathrm{shared}}$ 个 SiTU,全宽 $d \to d$,所有 token 都经过 \\
\textbf{Routed 专家} & $n_{\mathrm{routed}}$ 个 SiTU,半宽 $\ell \to \ell$($\ell = d/2$) \\
\textbf{Latent 接口} & $W_\downarrow: d \to \ell$, $W_\uparrow: \ell \to d$(压缩/还原) \\
\textbf{Router} & $W_r: d \to n_{\mathrm{routed}}$,Top-k 选择 + softmax 归一化 \\
\textbf{Router} & $W_r: d \to n_{\mathrm{routed}}$,K3:$s=\sigma(W_r x)$,Top-$k(s+b)$,$p_i=s_i/\sum_{j\in T}s_j$ \\
\bottomrule
\end{tabular}
\end{center}
@@ -67,29 +68,31 @@ class SiTU(nn.Module):
z = W_\downarrow \cdot x \qquad \shape{B, T, \ell}
\]
\item \textbf{Routing}:
\item \textbf{Routing}(K3 eq.13):
\[
\mathrm{logits} = W_r \cdot x \qquad \shape{B, T, n_{\mathrm{routed}}}
\]
\[
\mathrm{ids}, \mathrm{probs} = \mathrm{TopK}(\mathrm{logits}, k)
\qquad \mathrm{ids}: \shape{B, T, k}, \;\; \mathrm{probs}: \shape{B, T, k}
s = \sigma(W_r \cdot x),\quad
T = \mathrm{TopK}(s+b, k),\quad
p_i = \frac{s_i}{\sum_{j\in T} s_j}
\]
\item \textbf{Routed 专家}(在 latent 空间 $\ell$ 上):
\item \textbf{Routed 专家}(在 latent 空间 $\ell$ 上,稀疏执行,见 \S7.4):
\[
u = \sum_{i \in \mathrm{Top\text{-}k}} p_i \cdot E_i^{\mathrm{rt}}(z)
\qquad \shape{B, T, \ell}
\]
数学上是"对选中的 $k$ 个专家加权求和",但\textbf{实现上不是}用
\texttt{stack([e(z) for e in experts])} 把所有专家都算一遍——
而是每个 token 只被送进它选中的 $k$ 个专家(permute-dispatch + padded bmm)。
\item \textbf{Shared 专家}(全宽 $d$):
\[
s = \sum_j E_j^{\mathrm{sh}}(x) \qquad \shape{B, T, d}
s_{\mathrm{sh}} = \sum_j E_j^{\mathrm{sh}}(x) \qquad \shape{B, T, d}
\]
\item \textbf{合并}:
\[
y = s + W_\uparrow \cdot \mathrm{RMSNorm}(u) \qquad \shape{B, T, d}
y = s_{\mathrm{sh}} + W_\uparrow \cdot \mathrm{RMSNorm}(u) \qquad \shape{B, T, d}
\]
\end{enumerate}
@@ -99,38 +102,132 @@ Routed 专家只在 $\ell = d/2$ 的 latent 空间操作,
Shared 专家保持全宽 $d$,提供基础表达能力。
\end{importantbox}
\subsection{代码对照}
\subsection{代码对照:路由}
\begin{codemathtop}{layers/latent\_moe.py — LatentMoE.forward}
\begin{codemathtop}{layers/latent\_moe.py — \_route(K3 eq.13)}
\begin{lstlisting}
def forward(self, x): # [B, T, d]
z = self.down(x) # [B, T, ell]
logits = self.router(x) # [B, T, n_routed]
topk = torch.topk(logits, self.top_k, dim=-1)
ids = topk.indices # [B, T, k]
probs = F.softmax(topk.values, dim=-1) # [B, T, k]
# All expert outputs (vectorized)
all_out = stack([e(z) for e in self.experts]) # [R, B, T, ell]
# Gather top-k and weighted sum
u = zeros(B, T, ell)
for i in range(self.top_k):
idx = ids[:,:,i].reshape(B*T)
sel = all_out[arange, idx]
u += probs[:,:,i:i+1] * sel.reshape(B, T, ell)
shared_out = stack([e(x) for e in self.shared]).sum(0) # [B, T, d]
return shared_out + self.up(self.norm(u)) # [B, T, d]
def _route(self, logits): # logits: [B, T, n_routed]
scores = sigmoid(logits) # s = σ(W_r x)
ids = topk(scores + self.expert_bias, k).indices # T = TopK(s+b)
selected = scores.gather(-1, ids)
probs = selected / selected.sum(-1).clamp_min(1e-9) # p_i = s_i / Σ_{j∈T} s_j
return ids, probs
\end{lstlisting}
\end{codemathtop}
注意三点:
\begin{enumerate}[leftmargin=2em]
\item \textbf{sigmoid 代替 softmax}:K3 的 router 对每个专家独立打分
$s_i = \sigma(w_i \cdot x)$,不再是 softmax 归一化。这样"专家之间"不互相竞争
归一化预算,便于用 bias 做负载调节。
\item \textbf{bias 只进选择、不进权重}:$\texttt{expert\_bias}$ 是
\texttt{register\_buffer(..., persistent=False)} 的非持久 buffer(初始化为 $0$,
不进 \texttt{state\_dict}),只参与 $\mathrm{TopK}(s+b)$ 的选择,
归一化 $p_i$ 仍用原始 $s_i$。
\item \textbf{L1 归一化}:$p_i = s_i / \sum_{j \in T} s_j$(在选中的 $k$ 个上做),
权重之和为 1,等价于"选中的 sigmoid 分数重新归一化"。
\end{enumerate}
\subsection{稀疏执行:permute-dispatch + padded bmm}\label{sec:sparse-dispatch}
\begin{codemathtop}{layers/latent\_moe.py — \_routed\_u}
\begin{lstlisting}
def _routed_u(self, z, ids, probs): # z: [B,T,ell] ids/probs: [B,T,k]
N = B * T; R, k = self.n_routed, self.top_k
tok = arange(N).unsqueeze(1).expand(N, k).reshape(-1) # 每个 token 重复 k 次
eid, pw = ids.reshape(-1), probs.reshape(-1)
order = eid.argsort(stable=True) # 按专家 id 排序 → 同专家连续
tok, eid, pw = tok[order], eid[order], pw[order]
counts = bincount(eid, minlength=R) # 每个专家的 token 数
offsets = counts.cumsum(0) - counts
local_pos = arange(N*k) - offsets[eid] # 专家内局部位置
C = int(counts.max().item()) # capacity = 最大负载
gathered = z.reshape(N, ell)[tok]
padded = index_put(zeros(R, C, ell), (eid, local_pos), gathered) # [R, C, ell]
w_g = stack([e.w_g.weight for e in self.experts]) # [R, ff, ell]
w_u = stack([e.w_u.weight for e in self.experts])
w_o = stack([e.w_o.weight for e in self.experts]) # [R, ell, ff]
wg = bmm(padded, w_g.transpose(-1,-2)) # [R, C, ff] grouped GEMM
g = beta1 * tanh(wg / beta1) * sigmoid(wg)
wu = bmm(padded, w_u.transpose(-1,-2))
h = beta2 * tanh(wu / beta2)
out = bmm(g * h, w_o.transpose(-1,-2)) # [R, C, ell]
weighted = pw.unsqueeze(-1) * out[eid, local_pos]
return index_add(zeros(N, ell), 0, tok, weighted).view(B, T, ell) # scatter-add
\end{lstlisting}
\end{codemathtop}
三步走:\textbf{① permute-dispatch}(按专家排序 + pad 到 $[R, C, \ell]$)→
\textbf{② padded bmm}(专家参数堆成 batch 维,三次 batched GEMM 一次算完 $R$ 个专家)→
\textbf{③ scatter-add}(\texttt{index\_add} 把加权输出按 \texttt{tok} 累加回 $u$)。
\begin{importantbox}{为什么不用 dense stack?}
朴素写法 \texttt{stack([e(z) for e in experts])} 会让每个专家都算全部 $B\cdot T$ 个 token,
FLOPs 是 $R \cdot N$,退化成 dense,失去 MoE 的加速。稀疏 dispatch 把每个 token 只送进
它选中的 $k$ 个专家,FLOPs 是 $R \cdot C$($C = \max_e \text{count}_e \approx k \cdot N / R$),
当 $k \ll R$ 时远小于 $R \cdot N$。padding 槽位不被 \texttt{index\_add} 收集,
贡献零梯度;空专家仍留在堆叠权重里(padded grouped GEMM)。
\end{importantbox}
\begin{warningbox}{为什么 router 用 $x$(全宽)而不是 $z$(latent)?}
路由需要看到 token 的完整表示才能做好选择。
如果用 $z$ 路由,压缩过程可能丢失路由需要的信息。
K3 论文里也是用全宽 $x$ 做 Top-k,然后在 $\ell$ 空间计算。
\end{warningbox}
\subsection{负载均衡损失:Switch/GShard aux + z-loss}\label{sec:load-balancing}
Top-k 路由容易"塌缩"到少数专家(router 学出永远选某几个专家),导致负载不均衡、
专家利用率低。训练时加两个损失,由 train loop 加到 CE 上:
\[
f_e = \frac{\#\{\text{路由到 } e\}}{N \cdot k},\qquad
P_e = \frac{1}{N}\sum_{t} \sigma(W_r x_t)_e
\]
\[
\mathcal{L}_{\mathrm{aux}} = n_{\mathrm{routed}} \sum_e f_e \cdot P_e,\qquad
\mathcal{L}_{z} = \frac{1}{N}\sum_t \left(\log\!\sum_j e^{l_{tj}}\right)^2
\]
\begin{center}
\begin{tabular}{ll}
\toprule
项 & 作用 \\
\midrule
$\mathcal{L}_{\mathrm{aux}}$ & Switch/GShard 风格:$f_e$ 是专家 $e$ 被路由到的 token 占比,$P_e$ 是它的平均 sigmoid 分数。塌缩时 $f$ 集中到单个专家、损失变大,逼着路由均匀化 \\
$\mathcal{L}_{z}$ & z-loss:对 raw logits 的 logsumexp 求平方,压住 logits 幅度、防 router 分数爆炸 \\
\bottomrule
\end{tabular}
\end{center}
\begin{codemathtop}{layers/latent\_moe.py — \_balancing\_losses}
\begin{lstlisting}
def _balancing_losses(self, logits, ids):
flat = logits.reshape(-1, self.n_routed).float()
scores = sigmoid(flat)
counts = bincount(ids.reshape(-1), minlength=self.n_routed).float()
frac = counts / counts.sum().clamp_min(1.0) # f_e
prob_mean = scores.mean(dim=0) # P_e
aux = self.n_routed * (frac * prob_mean).sum() # N Σ f_e P_e
z_loss = logsumexp(flat, dim=-1).square().mean() # mean (logsumexp)^2
return self.aux_loss_coef * aux, self.z_loss_coef * z_loss
\end{lstlisting}
\end{codemathtop}
两个损失只在 \texttt{self.training} 且系数非零时计算;系数默认
$\alpha_{\mathrm{aux}} = 10^{-2}$、$\alpha_z = 10^{-3}$(\texttt{K3Config.moe\_aux\_loss\_coef} /
\texttt{moe\_z\_loss\_coef},可用 \texttt{--moe-aux-coef} / \texttt{--moe-z-coef} 覆盖)。
train loop 里 \texttt{moe\_router\_losses(model)} 把所有 LatentMoE 层的损失求和,
\texttt{loss = task + aux + z\_loss} 一起反传。aux/z 只更新 router 参数,
不碰专家权重(\texttt{ids} 已 \texttt{.detach()})。
\subsection{形状总览}
\begin{center}
@@ -140,12 +237,17 @@ K3 论文里也是用全宽 $x$ 做 Top-k,然后在 $\ell$ 空间计算。
\midrule
$x$ & \shape{B, T, d} & 输入 \\
$z$ & \shape{B, T, \ell} & latent($\ell = d/2$)\\
logits & \shape{B, T, n_r} & router 输出 \\
logits & \shape{B, T, n_r} & router logits \\
$s$ & \shape{B, T, n_r} & $\sigma(\mathrm{logits})$ \\
$b$ & \shape{n_r} & expert bias(非持久,只进 TopK 不进 $p$)\\
ids & \shape{B, T, k} & Top-k 专家索引 \\
probs & \shape{B, T, k} & Top-k softmax 权重 \\
\texttt{all\_out} & \shape{n_r, B, T, \ell} & 所有 routed 专家输出 \\
probs & \shape{B, T, k} & K3 sigmoid-L1 权重 \\
padded & \shape{n_r, C, \ell} & dispatch 后 pad 到容量 $C$ 的张量 \\
$C$ & 标量 & 最大专家负载(pad 宽度)\\
$u$ & \shape{B, T, \ell} & 加权求和后的 routed 输出 \\
\texttt{shared\_out} & \shape{B, T, d} & shared 专家求和 \\
$f_e$, $P_e$ & 标量 & aux loss 的负载占比 / 平均分数 \\
$\mathcal{L}_{\mathrm{aux}}$, $\mathcal{L}_z$ & 标量 & 负载均衡 / z-loss \\
$y$ & \shape{B, T, d} & 最终输出 \\
\bottomrule
\end{tabular}
@@ -154,5 +256,7 @@ $y$ & \shape{B, T, d} & 最终输出 \\
\subsection{本章小结}
LatentMoE 把 routed 专家限制在 $\ell = d/2$ 的 latent 空间,省参数。
SiTU-GLU 给 gate 和 up 加 $\tanh$ 软上限($\beta_1=4, \beta_2=25$),
防止低精度溢出。Shared 专家全宽,提供基础能力;routed 专家通过 Top-k 路由提供专业化能力。
SiTU-GLU 给 gate 和 up 加 $\tanh$ 软上限($\beta_1=4, \beta_2=25$),防止低精度溢出。
路由走 K3 eq.13(sigmoid-TopK + L1 归一化),执行用稀疏 permute-dispatch + padded bmm
(每个 token 只算 $k$ 个专家,FLOPs $R \cdot C$ 而非 $R \cdot N$),训练加
Switch/GShard aux loss + router z-loss 防路由塌缩。
+14 -4
View File
@@ -30,6 +30,9 @@ $n_r$ & routed 专家数 & 16 \\
$k$ & Top-$k$ & 2 \\
$n_s$ & shared 专家数 & 2 \\
$d_{\mathrm{ff}}$ & 专家中间维度 & 96 \\
$C$ & MoE 专家容量(pad 宽度) & 动态 \\
$\alpha_{\mathrm{aux}}$ & Switch/GShard aux 系数 & $10^{-2}$ \\
$\alpha_z$ & router z-loss 系数 & $10^{-3}$ \\
$N$ & AttnRes 原子层数 ($= 2L$) & 8 \\
$S$ & AttnRes 块大小(原子层) & 2--24 \\
\bottomrule
@@ -105,12 +108,19 @@ gate & \shape{B, T, H \cdot d_v} & $\sigma(W_g x)$ \\
\midrule
$x$ & \shape{B, T, D} & 输入 \\
$z$ & \shape{B, T, \ell} & latent ($\ell = D/2$) \\
logits & \shape{B, T, n_r} & router logits \\
logits & \shape{B, T, n_r} & router logits $W_r x$ \\
$s$ & \shape{B, T, n_r} & sigmoid 分数 $\sigma(\mathrm{logits})$ \\
$b$ & \shape{n_r} & expert bias(非持久,只进 TopK) \\
ids & \shape{B, T, k} & Top-$k$ 专家索引 \\
probs & \shape{B, T, k} & softmax 权重 \\
$p_i$ & \shape{B, T, k} & sigmoid-L1 权重 $s_i/\sum_{j\in T}s_j$ \\
padded & \shape{n_r, C, \ell} & dispatch 后 pad 到容量 $C$ \\
$C$ & 标量 & 最大专家负载(pad 宽度) \\
$u$ & \shape{B, T, \ell} & routed 加权输出 \\
$s$ & \shape{B, T, D} & shared 专家求和 \\
$y$ & \shape{B, T, D} & $s + W_\uparrow \mathrm{RMSNorm}(u)$ \\
$s_{\mathrm{sh}}$ & \shape{B, T, D} & shared 专家求和 \\
$y$ & \shape{B, T, D} & $s_{\mathrm{sh}} + W_\uparrow \mathrm{RMSNorm}(u)$ \\
$f_e, P_e$ & 标量 & aux loss 负载占比 / 平均分数 \\
$\mathcal{L}_{\mathrm{aux}}, \mathcal{L}_z$ & 标量 & 负载均衡 / z-loss \\
$\alpha_{\mathrm{aux}}, \alpha_z$ & 标量 & 对应系数($10^{-2}$ / $10^{-3}$) \\
\bottomrule
\end{tabular}
\end{center}