Files
dela a2c4217dae Document LatentMoE sigmoid routing, sparse dispatch, and K3 block figures
Ledger C10–C12 match the permute-pad-bmm path and Switch aux/z-loss.
Section 8 adds overview and component TikZ; MoE capacity is C_moe so it
does not collide with KDA chunk size.
2026-08-26 14:43:58 +08:00

264 lines
11 KiB
TeX
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.
% teach:
% 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}
\subsection{SiTU-GLU:带软上限的激活}
SwiGLU 在低精度(fp16/bf16)训练时可能溢出:$\mathrm{silu}(x) \cdot x$ 没有上限。
SiTU-GLU 用 $\tanh$ 给门控和上投影加软上限:
\[
\mathrm{SiTU}(x) = W_o \big[\underbrace{\beta_1 \tanh\!\left(\frac{W_g x}{\beta_1}\right) \cdot \sigma(W_g x)}_{\text{gate}} \;\cdot\; \underbrace{\beta_2 \tanh\!\left(\frac{W_u x}{\beta_2}\right)}_{\text{up}}\big]
\]
\begin{center}
\begin{tabular}{lp{8cm}}
\toprule
性质 & 说明 \\
\midrule
输出上限 & $\|\mathrm{SiTU}\|_\infty \leq \beta_1 \cdot \beta_2 = 4 \times 25 = 100$ \\
原点附近 & $\tanh(x/\beta) \approx x/\beta$,所以 $\beta \cdot \tanh(x/\beta) \approx x$,退化为 SwiGLU \\
远端 & 软饱和,防 fp16 溢出 \\
\bottomrule
\end{tabular}
\end{center}
\begin{codemathtop}{layers/latent\_moe.py — SiTU}
\begin{lstlisting}
class SiTU(nn.Module):
def __init__(self, dim_in, dim_ff, beta1=4.0, beta2=25.0):
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): # [*, dim_in]
wg = self.w_g(x)
g = self.beta1 * tanh(wg / self.beta1) * sigmoid(wg) # gate
u = self.beta2 * tanh(self.w_u(x) / self.beta2) # up
return self.w_o(g * u) # [*, dim_in]
\end{lstlisting}
\end{codemathtop}
\subsection{LatentMoE 架构}
\begin{center}
\begin{tabular}{rl}
\toprule
组件 & 说明 \\
\midrule
\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}}$,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}
\subsection{计算流(五步)}
\begin{enumerate}[leftmargin=2em]
\item \textbf{Latent 压缩}:
\[
z = W_\downarrow \cdot x \qquad \shape{B, T, \ell}
\]
\item \textbf{Routing}(K3 eq.13):
\[
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$ 上,稀疏执行,见 \ref{sec:sparse-dispatch} 节):
\[
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_{\mathrm{sh}} = \sum_j E_j^{\mathrm{sh}}(x) \qquad \shape{B, T, d}
\]
\item \textbf{合并}:
\[
y = s_{\mathrm{sh}} + W_\uparrow \cdot \mathrm{RMSNorm}(u) \qquad \shape{B, T, d}
\]
\end{enumerate}
\begin{importantbox}{如果你只记一件事}
Routed 专家只在 $\ell = d/2$ 的 latent 空间操作,
参数量是全宽专家的 $1/4$($\ell^2$ vs $d^2$)。
Shared 专家保持全宽 $d$,提供基础表达能力。
\end{importantbox}
\subsection{代码对照:路由}
\begin{codemathtop}{layers/latent\_moe.py — \_route(K3 eq.13)}
\begin{lstlisting}
def _route(self, logits): # logits: [B, T, n_routed]
scores = sigmoid(logits) # s = sigma(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 / sum_{j in 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{(1) permute-dispatch}(按专家排序 + pad 到 $[R, C, \ell]$)$\to$
\textbf{(2) padded bmm}(专家参数堆成 batch 维,三次 batched GEMM 一次算完 $R$ 个专家)$\to$
\textbf{(3) 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}{lp{11.5cm}}
\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 * sum_e 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} 的
\texttt{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} 一起反传。
两个损失只依赖 router 输出 \texttt{logits} 与不可微的索引 \texttt{ids},
所以梯度只流回 router 的 $W_r$,不碰专家权重。
\subsection{形状总览}
\begin{center}
\begin{tabular}{llll}
\toprule
变量 & 形状 & 说明 \\
\midrule
$x$ & \shape{B, T, d} & 输入 \\
$z$ & \shape{B, T, \ell} & latent($\ell = d/2$)\\
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} & 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}
\end{center}
\subsection{本章小结}
LatentMoE 把 routed 专家限制在 $\ell = d/2$ 的 latent 空间,省参数。
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 防路由塌缩。