% 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$ 上,稀疏执行,见 \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_{\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 = σ(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} \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 防路由塌缩。