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.
329 lines
13 KiB
TeX
329 lines
13 KiB
TeX
% 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)。
|