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