Files
K3/notes/sections/sec-06.tex
T
dela 584f7e9e73 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.
2026-08-25 14:43:17 +08:00

174 lines
5.9 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: 读者知道标准 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。