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.
This commit is contained in:
+15
-14
@@ -75,7 +75,7 @@ class SiTU(nn.Module):
|
||||
p_i = \frac{s_i}{\sum_{j\in T} s_j}
|
||||
\]
|
||||
|
||||
\item \textbf{Routed 专家}(在 latent 空间 $\ell$ 上,稀疏执行,见 \S7.4):
|
||||
\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}
|
||||
@@ -107,10 +107,10 @@ Shared 专家保持全宽 $d$,提供基础表达能力。
|
||||
\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)
|
||||
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 / Σ_{j∈T} s_j
|
||||
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}
|
||||
@@ -138,7 +138,7 @@ def _routed_u(self, z, ids, probs): # z: [B,T,ell] ids/probs: [B,T,
|
||||
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 排序 → 同专家连续
|
||||
order = eid.argsort(stable=True) # 按专家 id 排序 -> 同专家连续
|
||||
tok, eid, pw = tok[order], eid[order], pw[order]
|
||||
|
||||
counts = bincount(eid, minlength=R) # 每个专家的 token 数
|
||||
@@ -164,9 +164,9 @@ def _routed_u(self, z, ids, probs): # z: [B,T,ell] ids/probs: [B,T,
|
||||
\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$)。
|
||||
三步走:\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,
|
||||
@@ -197,7 +197,7 @@ Top-k 路由容易"塌缩"到少数专家(router 学出永远选某几个专
|
||||
\]
|
||||
|
||||
\begin{center}
|
||||
\begin{tabular}{ll}
|
||||
\begin{tabular}{lp{11.5cm}}
|
||||
\toprule
|
||||
项 & 作用 \\
|
||||
\midrule
|
||||
@@ -215,18 +215,19 @@ def _balancing_losses(self, logits, ids):
|
||||
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
|
||||
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.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()})。
|
||||
$\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{形状总览}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user