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:
+144
-40
@@ -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 防路由塌缩。
|
||||
|
||||
Reference in New Issue
Block a user