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.
264 lines
11 KiB
TeX
264 lines
11 KiB
TeX
% 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 防路由塌缩。
|