% 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)。