Files
K3/notes/sections/sec-08.tex
T
dela a2c4217dae 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.
2026-08-26 14:43:58 +08:00

329 lines
13 KiB
TeX
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
% teach:
% gap: 读者已知各组件但不知道怎么组装成完整模型
% takeaway: K3 = Hybrid(3 KDA + 1 MLA) × DecoderBlock(attn + MoE), 末层强制 MLA
% jump: 为什么每 4 层才放一次 MLA?位置感知只需要周期性提供
% omit: 0.5b preset 的训练超参
\section{K3 混合架构}
\subsection{整体结构}
\begin{center}
\texttt{Embedding} $\to$ \texttt{DecoderBlock} $\times L$ $\to$ \texttt{RMSNorm} $\to$ \texttt{LM Head}
\end{center}
每个 \texttt{DecoderBlock} 是 Pre-Norm 残差:
\begin{lstlisting}
def forward(self, x):
x = x + self.attn(self.attn_norm(x)) # mixing
return x + self.ffn(self.ffn_norm(x)) # channel
\end{lstlisting}
\subsection{Hybrid Attention Pattern}
K3 用两种 attention 层交替:
\begin{center}
\begin{tabular}{ccccccccc}
\toprule
层 & 0 & 1 & 2 & 3 & 4 & 5 & 6 & 7 \\
\midrule
Attn & KDA & KDA & KDA & \textbf{MLA} & KDA & KDA & KDA & \textbf{MLA} \\
FFN & MoE & MoE & MoE & MoE & MoE & MoE & MoE & MoE \\
\bottomrule
\end{tabular}
\end{center}
\noindent 规则:每 4 层放 1 次 MLA(0-based 层 3, 7, 11, ...),\textbf{末层强制 MLA}。
\begin{codemathtop}{models/k3\_config.py — layer\_types}
\begin{lstlisting}
def layer_types(self) -> list[str]:
"""Hybrid: 3 KDA + 1 MLA per group, last always MLA."""
types = ["kda"] * self.num_hidden_layers
for i in range(self.num_hidden_layers):
if i % 4 == 3:
types[i] = "mla"
types[-1] = "mla" # last layer forced
return types
def layer_specs(self):
return [(kind, "moe") for kind in self.layer_types()]
\end{lstlisting}
\end{codemathtop}
\begin{knowledgebox}{为什么 KDA 不需要 RoPE?}
KDA 的 gate/decay 机制天然提供位置感知:
远的 token 衰减更多,近的保留更多。
但 softmax attention(MLA)没有这个机制,所以真实 K3 用 NoPE
(本复现的 MLA 也是 NoPE)。
位置感知从 KDA 层``渗透''到 MLA 层——3:1 的比例足够了。
\end{knowledgebox}
\subsection{CausalLM 完整数据流}
\begin{codemathtop}{models/causal\_lm.py — CausalLM}
\begin{lstlisting}
class CausalLM(nn.Module):
def __init__(self, config):
self.embedding = nn.Embedding(vocab_size, D) # [V, D]
self.blocks = ModuleList([
DecoderBlock.from_spec(config, attn, ffn)
for attn, ffn in config.layer_specs()
])
self.norm = RMSNorm(D)
self.lm_head = nn.Linear(D, vocab_size) # [V, D]
if config.tie_word_embeddings:
self.lm_head.weight = self.embedding.weight
def forward(self, input_ids, labels=None):
x = self.embedding(input_ids) # [B,T] -> [B,T,D]
if self.mixer is None: # attnres="off"
for block in self.blocks:
x = block(x) # [B,T,D] -> [B,T,D]
else:
x = self.mixer(x) # AttnRes 深度残差, 见 §9
logits = self.lm_head(self.norm(x)) # [B,T,D] -> [B,T,V]
if labels is None:
return logits
# Shifted CE: predict next token
return cross_entropy(logits[:,:-1], labels[:,1:])
\end{lstlisting}
\end{codemathtop}
\begin{knowledgebox}{残差流是可替换的}
上面的 \texttt{DecoderBlock} 逐层堆叠(\texttt{x = x + sublayer(norm(x))})
是 \texttt{config.attnres="off"} 时的默认路径。
置为 \texttt{"full"} / \texttt{"block"} 时,\texttt{CausalLM} 会把每个 block
拆成 attn / ffn 两个原子子层交给 \texttt{mixer},用\textbf{深度维注意力}
代替等权残差加法——见 \S9。真实 K3 用的是 \texttt{block} 模式。
\end{knowledgebox}
\subsection{两种配置}
\begin{center}
\begin{tabular}{lll}
\toprule
& \textbf{KDAConfig}(纯 KDA) & \textbf{K3Config}(混合) \\
\midrule
Attn & KDA only & 3 KDA + 1 MLA \\
FFN & SwiGLU & LatentMoE \\
典型规模 & \textasciitilde8M (toy) & \textasciitilde8M (toy) / \textasciitilde500M (0.5b) \\
\texttt{layer\_specs()} & \texttt{[("kda","swiglu")] * L} & \texttt{[(kind,"moe") for kind in ...]} \\
\bottomrule
\end{tabular}
\end{center}
\subsection{K3 toy 尺寸}
\begin{center}
\begin{tabular}{llll}
\toprule
参数 & 真实 K3 & toy 复现 & 缩比 \\
\midrule
$D$ & 7168 & 256 & 28$\times$ \\
$L$ & 93 & 4 & 23$\times$ \\
$H = H_V$ & 96 & 8 & 12$\times$ \\
$K = V$ & 128 & 16 & 8$\times$ \\
kv\_lora\_rank & 512 & 32 & 16$\times$ \\
q\_lora\_rank & 1536 & 64 & 24$\times$ \\
$\ell$ (MoE latent) & 3584 & 128 & 28$\times$ \\
$n_{\mathrm{routed}}$ / Top-$k$ & 896/16 & 16/2 & 56$\times$ / 8$\times$ \\
\bottomrule
\end{tabular}
\end{center}
\subsection{DecoderBlock 构建}
\begin{codemathtop}{layers/block.py — build\_attn / build\_ffn}
\begin{lstlisting}
def build_attn(config, kind: str) -> nn.Module:
if kind == "kda": return KDAAttention.from_config(config)
if kind == "mla": return GatedMLA.from_config(config)
def build_ffn(config, kind: str) -> nn.Module:
if kind == "swiglu": return SwiGLUMLP.from_config(config)
if kind == "moe": return LatentMoE.from_config(config)
class DecoderBlock(nn.Module):
def forward(self, x):
x = x + self.attn(self.attn_norm(x))
return x + self.ffn(self.ffn_norm(x))
\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。
KDA 层提供线性复杂度的序列混合和位置感知(通过 decay),
MLA 层提供全局 softmax attention(NoPE,利用 KDA 渗透的位置信息)。
每层默认是 Pre-Norm 残差 DecoderBlock;\texttt{config.attnres} 可以把这条
等权残差流换成 AttnRes 深度注意力(\S9)。