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:
dela
2026-08-26 14:43:58 +08:00
parent ea7167b3f7
commit a2c4217dae
6 changed files with 239 additions and 18 deletions
+15 -14
View File
@@ -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{形状总览}
+166
View File
@@ -153,6 +153,172 @@ class DecoderBlock(nn.Module):
\end{lstlisting}
\end{codemathtop}
\subsection{架构图}
\begin{figure}[H]
\centering
\begin{subfigure}[t]{0.44\textwidth}
\centering
\begin{tikzpicture}[>=Stealth, node distance=4mm,
blk/.style={draw, rounded corners=2pt, minimum width=32mm, minimum height=6mm,
align=center, font=\small},
io/.style={font=\small\itshape}]
\node[io] (in) {Input tokens};
\node[blk, fill=gray!8, below=5mm of in] (emb) {Embedding};
\node[blk, fill=blue!10, draw=blue!40, below=5mm of emb] (l0) {KDA + MoE};
\node[blk, fill=blue!10, draw=blue!40, below=2mm of l0] (l1) {KDA + MoE};
\node[blk, fill=blue!10, draw=blue!40, below=2mm of l1] (l2) {KDA + MoE};
\node[blk, fill=orange!12, draw=orange!50, below=2mm of l2] (l3) {MLA + MoE};
\node[below=1mm of l3, font=\normalsize] (dots) {$\vdots$};
\node[blk, fill=orange!12, draw=orange!50, below=1mm of dots] (lL) {MLA + MoE};
\draw[decorate, decoration={brace, amplitude=5pt, mirror}]
([xshift=2mm]l0.north east) -- ([xshift=2mm]lL.south east)
node[midway, right=6pt, font=\small] {$\times L$};
\node[blk, fill=gray!8, below=5mm of lL] (fnorm) {RMSNorm};
\node[blk, fill=gray!8, below=of fnorm] (head) {LM Head};
\node[io, below=of head] (out) {Logits};
\foreach \a/\b in {in/emb, emb/l0, l0/l1, l1/l2, l2/l3, l3/dots, dots/lL,
lL/fnorm, fnorm/head, head/out}
\draw[->] (\a) -- (\b);
\node[left=1mm of l0, font=\scriptsize, text=gray] {0};
\node[left=1mm of l1, font=\scriptsize, text=gray] {1};
\node[left=1mm of l2, font=\scriptsize, text=gray] {2};
\node[left=1mm of l3, font=\scriptsize, text=gray] {3};
\node[left=1mm of lL, font=\scriptsize, text=gray] {$L{-}1$};
\node[right=3mm of lL, font=\tiny, text=orange!60!black] {(强制)};
\end{tikzpicture}
\caption{整体模型}
\end{subfigure}
\hfill
\begin{subfigure}[t]{0.44\textwidth}
\centering
\begin{tikzpicture}[>=Stealth, node distance=5mm,
blk/.style={draw, rounded corners=2pt, minimum width=26mm, minimum height=6mm,
align=center, font=\small},
add/.style={circle, draw, thick, inner sep=0pt, minimum size=5.5mm,
font=\small\bfseries},
io/.style={font=\small\itshape}]
\node[io] (x) {$x$};
\node[blk, fill=gray!8, below=8mm of x] (n1) {RMSNorm};
\node[blk, fill=blue!10, draw=blue!40, below=of n1] (attn) {Attention};
\node[add, below=8mm of attn] (a1) {$+$};
\node[blk, fill=gray!8, below=8mm of a1] (n2) {RMSNorm};
\node[blk, fill=green!10, draw=green!40, below=of n2] (ffn) {FFN};
\node[add, below=8mm of ffn] (a2) {$+$};
\node[io, below=8mm of a2] (y) {$y$};
\foreach \a/\b in {x/n1, n1/attn, attn/a1, a1/n2, n2/ffn, ffn/a2, a2/y}
\draw[->] (\a) -- (\b);
\draw[->, gray!50, rounded corners=3pt]
(x.east) -- ++(14mm,0) |- (a1.east);
\draw[->, gray!50, rounded corners=3pt]
(a1.west) -- ++(-14mm,0) |- (a2.west);
\node[right=9mm of attn, font=\tiny, text=blue!60!black, align=left]
{KDA\\[-1pt]or MLA};
\node[left=9mm of ffn, font=\tiny, text=green!50!black, align=right]
{LatentMoE\\[-1pt]or SwiGLU};
\end{tikzpicture}
\caption{DecoderBlock}
\end{subfigure}
\caption{K3 混合架构。(a)~整体模型:每 4 层 1 次 MLA(层 3, 7, 11, \ldots),末层强制 MLA,
所有 FFN 均为 LatentMoE。(b)~DecoderBlock:Pre-Norm 残差,两个子块各含
RMSNorm $\to$ 子层 $\to$ 残差加。}
\label{fig:k3-overview}
\end{figure}
\begin{figure}[H]
\centering
%% ---------- (a) KDA ----------
\begin{subfigure}[t]{0.28\textwidth}
\centering
\begin{tikzpicture}[>=Stealth, node distance=5mm,
blk/.style={draw, rounded corners=2pt, minimum width=24mm, minimum height=6mm,
align=center, font=\footnotesize},
io/.style={font=\footnotesize\itshape}]
\node[io] (x) {$x$};
\node[blk, fill=blue!8, below=5mm of x] (proj)
{5 投影\\[-1pt]{\tiny $q, k, v, g, \beta$}};
\node[blk, fill=blue!12, draw=blue!40, below=of proj] (gate)
{Gate 激活};
\node[blk, fill=blue!20, draw=blue!50, below=of gate, minimum height=9mm]
(kda) {\texttt{chunk\_kda}\\[-1pt]{\tiny decay $+$ delta rule}};
\node[blk, fill=blue!8, below=of kda] (op) {$W_o$};
\node[io, below=5mm of op] (y) {$y$};
\foreach \a/\b in {x/proj, proj/gate, gate/kda, kda/op, op/y}
\draw[->] (\a) -- (\b);
\end{tikzpicture}
\caption{KDA Attention}
\end{subfigure}
\hfill
%% ---------- (b) Gated MLA ----------
\begin{subfigure}[t]{0.35\textwidth}
\centering
\begin{tikzpicture}[>=Stealth, node distance=5mm,
blk/.style={draw, rounded corners=2pt, minimum width=24mm, minimum height=6mm,
align=center, font=\footnotesize},
mul/.style={circle, draw, inner sep=0pt, minimum size=5mm, font=\tiny},
io/.style={font=\footnotesize\itshape}]
\node[io] (x) {$x$};
\node[blk, fill=orange!8, below=5mm of x] (lr)
{Q / KV 低秩压缩\\[-1pt]{\tiny $q_\downarrow\!\!\to\!\mathrm{norm}\!\to\!q_\uparrow$\;;\;
$c\!=\!\mathrm{norm}(W_\downarrow x)$}};
\node[blk, fill=orange!15, draw=orange!50, below=of lr] (abs)
{矩阵吸收 + 打分\\[-1pt]{\tiny $q_{\mathrm{abs}}\!=\!q\!\cdot\!W_{UK}$\;;\;
$\mathrm{score}\!=\!q_{\mathrm{abs}}\!\cdot\!c^T$}};
\node[blk, fill=orange!10, below=of abs] (sm)
{Causal Softmax};
\node[blk, fill=orange!12, draw=orange!40, below=of sm] (wuv)
{$\mathrm{attn}\!\cdot\!c \;\to\; W_{UV}^T$};
\node[mul, below=6mm of wuv] (m) {$\odot$};
\node[blk, fill=orange!6, right=4mm of m, minimum width=13mm, minimum height=5mm]
(g) {\tiny $\sigma(W_g x)$};
\draw[->] (g) -- (m);
\node[blk, fill=orange!8, below=6mm of m, minimum width=16mm] (op) {$W_o$};
\node[io, below=5mm of op] (y) {$y$};
\foreach \a/\b in {x/lr, lr/abs, abs/sm, sm/wuv, wuv/m, m/op, op/y}
\draw[->] (\a) -- (\b);
\end{tikzpicture}
\caption{Gated MLA}
\end{subfigure}
\hfill
%% ---------- (c) LatentMoE ----------
\begin{subfigure}[t]{0.30\textwidth}
\centering
\begin{tikzpicture}[>=Stealth, node distance=5mm,
blk/.style={draw, rounded corners=2pt, minimum width=16mm, minimum height=6mm,
align=center, font=\footnotesize},
add/.style={circle, draw, inner sep=0pt, minimum size=5mm,
font=\scriptsize\bfseries},
io/.style={font=\footnotesize\itshape}]
\node[io] (x) at (0,0) {$x$};
\node[blk, fill=green!10] (sh) at (-1.1,-1.3)
{Shared\\[-1pt]{\tiny SiTU, $d\!\to\!d$}};
\node[blk, fill=green!8, minimum width=20mm] (dr) at (1.1,-1.3)
{$W_\downarrow$ + Router\\[-1pt]{\tiny $\sigma$-TopK}};
\draw[->] (x) -- (sh);
\draw[->] (x) -- (dr);
\node[blk, fill=green!15, draw=green!40, minimum width=20mm] (re) at (1.1,-2.7)
{Routed 专家\\[-1pt]{\tiny SiTU, $\ell\!\to\!\ell$}};
\draw[->] (dr) -- (re);
\node[blk, fill=green!8, minimum width=20mm] (up) at (1.1,-4.0)
{RMSNorm $\to$ $W_\uparrow$};
\draw[->] (re) -- (up);
\node[add] (a) at (0,-5.2) {$+$};
\draw[->, rounded corners=3pt] (sh.south) -- ++(0,-3mm) -| (a);
\draw[->, rounded corners=3pt] (up.south) -- ++(0,-3mm) -| (a);
\node[io] (y) at (0,-6.0) {$y$};
\draw[->] (a) -- (y);
\end{tikzpicture}
\caption{LatentMoE}
\end{subfigure}
\caption{K3 三大组件。
(a)~KDA:5 路投影 $\to$ gate 激活 $\to$ \texttt{chunk\_kda}(decay $+$ delta rule)
$\to$ 输出投影。
(b)~Gated MLA:$q$ 吸收 $W_{UK}$ 后在 latent $c$ 上打分(NoPE);输出经 sigmoid 门控。
(c)~LatentMoE:shared 全宽 $d$ + routed 半宽 $\ell\!=\!d/2$;sigmoid-TopK 路由,
padded bmm 稀疏执行。}
\label{fig:k3-components}
\end{figure}
\subsection{本章小结}
K3 架构 = Hybrid Attention(3 KDA + 1 MLA,末层强制 MLA)+ LatentMoE。
+11 -4
View File
@@ -30,7 +30,7 @@ $n_r$ & routed 专家数 & 16 \\
$k$ & Top-$k$ & 2 \\
$n_s$ & shared 专家数 & 2 \\
$d_{\mathrm{ff}}$ & 专家中间维度 & 96 \\
$C$ & MoE 专家容量(pad 宽度) & 动态 \\
$C_{\mathrm{moe}}$ & MoE 专家容量(pad 宽度) & 动态 \\
$\alpha_{\mathrm{aux}}$ & Switch/GShard aux 系数 & $10^{-2}$ \\
$\alpha_z$ & router z-loss 系数 & $10^{-3}$ \\
$N$ & AttnRes 原子层数 ($= 2L$) & 8 \\
@@ -113,8 +113,8 @@ $s$ & \shape{B, T, n_r} & sigmoid 分数 $\sigma(\mathrm{logits})$ \\
$b$ & \shape{n_r} & expert bias(非持久,只进 TopK) \\
ids & \shape{B, T, k} & Top-$k$ 专家索引 \\
$p_i$ & \shape{B, T, k} & sigmoid-L1 权重 $s_i/\sum_{j\in T}s_j$ \\
padded & \shape{n_r, C, \ell} & dispatch 后 pad 到容量 $C$ \\
$C$ & 标量 & 最大专家负载(pad 宽度) \\
padded & \shape{n_r, C_{\mathrm{moe}}, \ell} & dispatch 后 pad 到容量 $C_{\mathrm{moe}}$ \\
$C_{\mathrm{moe}}$ & 标量 & 最大专家负载(pad 宽度) \\
$u$ & \shape{B, T, \ell} & routed 加权输出 \\
$s_{\mathrm{sh}}$ & \shape{B, T, D} & shared 专家求和 \\
$y$ & \shape{B, T, D} & $s_{\mathrm{sh}} + W_\uparrow \mathrm{RMSNorm}(u)$ \\
@@ -164,6 +164,11 @@ MLA 解压 & \texttt{'bhtj,hvj->bhtv'} & $\tilde{o}$ \shape{B,H,T,d_v} \\
AttnRes 深度打分 & \texttt{'d,nbtd->nbt'} & $s_{l,i}$ \shape{n,B,T} \\
AttnRes 深度加权和 & \texttt{'nbt,nbtd->btd'} & $h_l$ \shape{B,T,D} \\
AttnRes 批量打分(inter) & \texttt{'qd,nbtd->qnbt'} & logits \shape{S,n,B,T} \\
MoE dispatch pad & \texttt{index\_put} & padded \shape{R,C_{\mathrm{moe}},\ell} \\
MoE gate 投影(grouped) & \texttt{bmm(padded, w\_g.T)} & $wg$ \shape{R,C_{\mathrm{moe}},ff} \\
MoE up 投影(grouped) & \texttt{bmm(padded, w\_u.T)} & $wu$ \shape{R,C_{\mathrm{moe}},ff} \\
MoE 输出投影(grouped) & \texttt{bmm(g$\odot$h, w\_o.T)} & out \shape{R,C_{\mathrm{moe}},\ell} \\
MoE scatter-add & \texttt{index\_add(0, tok, ...)} & $u$ \shape{N,\ell} \\
\bottomrule
\end{tabular}
\end{center}
@@ -177,7 +182,9 @@ AttnRes 批量打分(inter) & \texttt{'qd,nbtd->qnbt'} & logits \shape{S,n,B
\item \textbf{分块} = chunk 内下三角解 + chunk 间状态递推,等价于 naive recurrent
\item \textbf{GVA} = $H_V = G \cdot H$,forward repeat\_interleave / backward view+sum
\item \textbf{MLA} = 低秩 latent + 矩阵吸收,KV cache 从 $2Hd$ 降到 $r$
\item \textbf{LatentMoE} = shared 全宽 + routed 半宽 latent + SiTU-GLU 防溢出
\item \textbf{LatentMoE} = shared 全宽 + routed 半宽 latent + SiTU-GLU 防溢出;
K3 sigmoid-TopK 路由 + 稀疏 permute-dispatch(每 token 只算 $k$ 个专家)+
Switch/GShard aux \& z-loss 防塌缩
\item \textbf{K3 Hybrid} = 3 KDA + 1 MLA,KDA 提供位置感知
\item \textbf{AttnRes} = 深度维 softmax 残差,Block 版把源数压到 $O(N/S)$,
两阶段 = inter 批量 + intra online-softmax 合并