Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
174 lines
5.9 KiB
TeX
174 lines
5.9 KiB
TeX
% 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。
|