Initial K3 snapshot: 0.5B KDA/MLA/MoE train path

Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+173
View File
@@ -0,0 +1,173 @@
% teach:
% gap: 读者知道标准 MHA 但不知道 MLA 怎么压缩 KV、矩阵吸收怎么避免解压
% takeaway: MLA 把 KV 压成低秩 latent c, 通过吸收 W_UK 进 q 直接在 latent 空间算 attention
% jump: 为什么可以先在 latent 加权再乘 W_UV?因为矩阵乘和加权求和可交换
% omit: RoPE (K3 用 NoPE)
\section{Gated MLA(矩阵吸收版)}
\splabel{C4}
\subsection{标准 MHA 的 KV cache 问题}
标准 MHA 推理时需要缓存所有历史 token 的 $K, V$,cache 大小 $\propto T \cdot H \cdot d$。
MLA 的想法:把 $K, V$ 压缩成一个低秩 latent $c$,cache 大小 $\propto T \cdot r$,
其中 $r \ll H \cdot d$。
\subsection{低秩压缩}
\[
c = \mathrm{RMSNorm}(W_{\downarrow} \cdot x) \qquad \shape{B, T, r}
\]
推理时只缓存 $c$,不缓存解压后的 $K, V$。
解压矩阵 $W_{\mathrm{KV}\uparrow}$ 包含两部分:
\[
W_{\mathrm{KV}\uparrow} = \begin{bmatrix} W_{UK} \\ W_{UV} \end{bmatrix}
\qquad \shape{H \cdot (d_q + d_v), r}
\]
拆开:$W_{UK} \in \mathbb{R}^{H \times d_q \times r}$(key 解压),
$W_{UV} \in \mathbb{R}^{H \times d_v \times r}$(value 解压)。
\subsection{矩阵吸收的核心思路}
\textbf{不解压} $K$ 和 $V$。标准做法会先解压再算 attention:
\begin{center}
\textit{标准}:$k_h = c \cdot W_{UK,h}^T$ \shape{B,T,d_q},
$\mathrm{score} = q_h \cdot k_h^T$
\end{center}
矩阵吸收反过来:把 $W_{UK}$ 吸收进 $q$:
\begin{center}
\textit{吸收}:$q_{\mathrm{abs},h} = q_h \cdot W_{UK,h}$ \shape{B,T,r},
$\mathrm{score} = q_{\mathrm{abs},h} \cdot c^T$
\end{center}
\begin{importantbox}{如果你只记一件事}
$(q \cdot W_{UK}^T) \cdot c^T = q \cdot (W_{UK}^T \cdot c^T) = q_{\mathrm{abs}} \cdot c^T$
吸收后,attention 直接在 latent 空间 $r$ 维上算,永不解压到 $H \cdot d_q$ 维。
\end{importantbox}
\subsection{完整计算流(四步)}
\begin{enumerate}[leftmargin=2em]
\item \textbf{Q 低秩路径}(NoPE,只有 nope 段):
\[
q = W_{q\uparrow} \cdot \mathrm{RMSNorm}(W_{q\downarrow} \cdot x)
\qquad \shape{B, T, H, d_q}
\]
\item \textbf{吸收 $W_{UK}$ + 打分}:
\[
q_{\mathrm{abs}} = q \cdot W_{UK} \quad
\xrightarrow{\texttt{einsum('bthd,hdj->bthj')}} \quad \shape{B, T, H, r}
\]
\[
\mathrm{score} = q_{\mathrm{abs}} \cdot c^T \quad
\xrightarrow{\texttt{einsum('bthj,bsj->bhts')}} \quad \shape{B, H, T, T}
\]
\[
\mathrm{attn} = \mathrm{softmax}(\mathrm{causal\_mask}(\mathrm{score}))
\qquad \shape{B, H, T, T}
\]
\item \textbf{先在 latent 加权,再乘 $W_{UV}^T$}:
\[
\tilde{o}_{\mathrm{lat}} = \mathrm{attn} \cdot c \quad
\xrightarrow{\texttt{einsum('bhts,bsj->bhtj')}} \quad \shape{B, H, T, r}
\]
\[
\tilde{o} = \tilde{o}_{\mathrm{lat}} \cdot W_{UV}^T \quad
\xrightarrow{\texttt{einsum('bhtj,hvj->bhtv')}} \quad \shape{B, H, T, d_v}
\]
\item \textbf{输出门 + 投影}:
\[
y = W_o \big[ \sigma(W_g \cdot x) \odot \tilde{o}_{\mathrm{flat}} \big]
\qquad \shape{B, T, D}
\]
\end{enumerate}
\subsection{代码对照}
\begin{codemathtop}{layers/mla.py — GatedMLA.forward}
\begin{lstlisting}
def forward(self, x): # x: [B, T, D]
B, T, _ = x.shape
H, r = self.num_heads, self.kv_up.in_features
# Step 1: latent + query
c = self.kv_norm(self.kv_down(x)) # [B, T, r]
q = self.q_up(self.q_norm(self.q_down(x))) # [B, T, H*d_q]
q = q.view(B, T, H, self.qk_nope_head_dim) # [B, T, H, d_q]
# Split W_UK, W_UV from kv_up.weight
w = self.kv_up.weight # [H*(d_q+d_v), r]
w_uk = w[:H*d_q].view(H, d_q, r) # [H, d_q, r]
w_uv = w[H*d_q:].view(H, d_v, r) # [H, d_v, r]
# Step 2: absorb W_UK, score
q_absorb = einsum('bthd,hdj->bthj', q, w_uk) # [B,T,H,r]
scores = einsum('bthj,bsj->bhts', q_absorb, c) # [B,H,T,T]
scores = scores.masked_fill(causal_mask, -inf)
attn = softmax(scores, dim=-1) # [B,H,T,T]
# Step 3: latent-space weighted sum, then W_UV
latent_out = einsum('bhts,bsj->bhtj', attn, c) # [B,H,T,r]
o_heads = einsum('bhtj,hvj->bhtv', latent_out, w_uv) # [B,H,T,d_v]
# Step 4: output gate
o_heads = o_heads.transpose(1,2).reshape(B,T, H*d_v)
gate = sigmoid(self.gate(x)) # [B,T,H*d_v]
return self.o_proj(gate * o_heads) # [B,T,D]
\end{lstlisting}
\end{codemathtop}
\begin{warningbox}{为什么可以先加权再乘 $W_{UV}$?}
标准做法:$o = \mathrm{attn} \cdot V = \mathrm{attn} \cdot (c \cdot W_{UV}^T)$
交换顺序:$o = (\mathrm{attn} \cdot c) \cdot W_{UV}^T$
这能成立是因为矩阵乘法的结合律:$A(BC) = (AB)C$。
$\mathrm{attn} \cdot c$ 先在 latent 空间 $r$ 维上加权求和,
得到的 \shape{B,H,T,r} 再乘 $W_{UV}^T$ 还原到 $d_v$ 维。
全程不需要显式构造 $H \cdot T$ 大小的 $V$ 矩阵。
\end{warningbox}
\subsection{形状与参数对比}
\begin{center}
\begin{tabular}{lll}
\toprule
参数 & 形状 & 说明 \\
\midrule
\texttt{kv\_down.weight} & \shape{r, D} & KV latent 压缩 \\
\texttt{kv\_up.weight} & \shape{H \cdot (d_q+d_v), r} & 包含 $W_{UK}$ 和 $W_{UV}$ \\
\texttt{q\_down.weight} & \shape{r_q, D} & Q 低秩 \\
\texttt{q\_up.weight} & \shape{H \cdot d_q, r_q} & Q 解压 \\
\texttt{gate.weight} & \shape{H \cdot d_v, D} & 输出门 \\
\texttt{o\_proj.weight} & \shape{D, H \cdot d_v} & 输出投影 \\
\bottomrule
\end{tabular}
\end{center}
\noindent KV cache 大小对比(推理时):
\begin{center}
\begin{tabular}{ll}
\toprule
方法 & Cache 大小 per token \\
\midrule
标准 MHA & $2 \times H \times d = 2 H d$ \\
MLA (latent) & $r$(只存 $c$) \\
\bottomrule
\end{tabular}
\end{center}
\subsection{本章小结}
Gated MLA 把 KV 压缩到低秩 latent $c$ \shape{B,T,r},通过矩阵吸收
($q_{\mathrm{abs}} = q \cdot W_{UK}$)直接在 latent 空间打分和加权,
永不解压 K/V。输出通过 sigmoid 门控。NoPE:不使用 RoPE,位置感交给夹层 KDA 的 decay/gate。