% teach: % gap: 读者不知道 q/k 和 v 为什么可以有不同的头数, 以及 repeat_interleave 的反传怎么做 % takeaway: GVA 让 G 组 value heads 共享一组 q/k, forward repeat_interleave, backward sum % jump: 论文没解释为什么反传是 sum 而不是 mean % omit: GQA 的历史 \section{GVA(分组值注意力)} \splabel{GVA} \subsection{为什么头数不一样?} 标准 MHA 里 $H_q = H_k = H_v$。GQA(Grouped Query Attention)让多组 q/k 共享同一组 v/k, 减少 KV cache。KDA 反过来做:$H$ 组 q/k 对应 $H_V = G \cdot H$ 组 value heads。 直觉:value 维度决定表达能力,多一点 value head 增加容量; q/k 主要负责路由(``看哪里''),可以共享。 \begin{center} \begin{tabular}{lll} \toprule & 标准头数 & GVA \\ \midrule $q, k$ & \shape{B, T, H, K} & \shape{B, T, H, K}(不变)\\ $v$ & \shape{B, T, H, V} & \shape{B, T, HV, V}($H_V = G \cdot H$) \\ $g, \beta$ & \shape{B, T, H, K/1} & \shape{B, T, HV, K/1} \\ $S$ & \shape{B, H, K, V} & \shape{B, HV, K, V} \\ \bottomrule \end{tabular} \end{center} \subsection{Forward: repeat\_interleave} 进入 KDA 核心前,$q$ 和 $k$ 从 $H$ 维复制到 $H_V$ 维: \begin{lstlisting} G = HV // H qe = q.repeat_interleave(G, dim=2) * scale # [B,T,H,K] -> [B,T,HV,K] ke = k.repeat_interleave(G, dim=2) # [B,T,H,K] -> [B,T,HV,K] \end{lstlisting} \noindent 例如 $H=4, G=2, H_V=8$:head 0 的 q/k 复制到 value head 0 和 1, head 1 复制到 value head 2 和 3,依此类推。 \subsection{Backward: view + sum} 反传时,$dq_e$ 和 $dk_e$ 的形状是 \shape{B, T, HV, K}(在 $H_V$ 维上计算的梯度)。 因为 forward 是复制,反传就是求和: \begin{lstlisting} # Backward: HV -> H dq_H = dq_e.view(B, T, H, G, K).sum(dim=3) # [B,T,HV,K] -> [B,T,H,K] dk_H = dk_e.view(B, T, H, G, K).sum(dim=3) \end{lstlisting} \begin{warningbox}{为什么是 sum 不是 mean?} \texttt{repeat\_interleave} 是\textbf{复制}:$y_0 = x_0, y_1 = x_0, y_2 = x_1, \ldots$ 对 $x_0$ 的梯度 = $\frac{\partial L}{\partial y_0} + \frac{\partial L}{\partial y_1}$ = \textbf{sum}(不是 mean)。 这和 \texttt{.expand()} 的反传一样:复制的反传是求和。 \end{warningbox} \subsection{scale 的处理} $q$ 在 repeat\_interleave 之后乘了 \texttt{scale = $1/\sqrt{K}$}。 反传时 chain rule 要求 $dq_{\mathrm{orig}} = dq_e \cdot \texttt{scale}$: \begin{lstlisting} # q 在 forward 内被乘过 scale, chain rule: dq_H = dq_H * scale \end{lstlisting} \subsection{KDAAttention 层中的投影} \begin{codemathtop}{layers/kda\_attn.py — forward} \begin{lstlisting} def forward(self, x): # x: [B, T, D] B, T, _ = x.shape H, HV, K, V = self.num_heads, self.num_value_heads, ... q = self.q_proj(x).view(B, T, H, K) # [B,T,D] -> [B,T,H*K] -> [B,T,H,K] k = self.k_proj(x).view(B, T, H, K) # 同上 v = self.v_proj(x).view(B, T, HV, V) # [B,T,D] -> [B,T,HV*V] -> [B,T,HV,V] g_raw = self.g_proj(x).view(B, T, HV, K) beta_raw = self.beta_proj(x).view(B, T, HV) o, _ = chunk_kda(q, k, v, g_raw, beta_raw, ...) return self.o_proj(o.reshape(B, T, HV * V)) # [B,T,HV,V] -> [B,T,D] \end{lstlisting} \end{codemathtop} \subsection{投影矩阵形状总览} \begin{center} \begin{tabular}{llll} \toprule 投影 & 权重形状 & 输入 & 输出 \\ \midrule \texttt{q\_proj} & \shape{H \cdot K, D} & \shape{B,T,D} & \shape{B,T,H,K} \\ \texttt{k\_proj} & \shape{H \cdot K, D} & \shape{B,T,D} & \shape{B,T,H,K} \\ \texttt{v\_proj} & \shape{HV \cdot V, D} & \shape{B,T,D} & \shape{B,T,HV,V} \\ \texttt{g\_proj} & \shape{HV \cdot K, D} & \shape{B,T,D} & \shape{B,T,HV,K} \\ \texttt{beta\_proj} & \shape{HV, D} & \shape{B,T,D} & \shape{B,T,HV} \\ \texttt{o\_proj} & \shape{D, HV \cdot V} & \shape{B,T,HV \cdot V} & \shape{B,T,D} \\ \bottomrule \end{tabular} \end{center} \subsection{本章小结} GVA 让 $H_V = G \cdot H$ 组 value heads 共享 $H$ 组 q/k。 Forward 用 \texttt{repeat\_interleave} 复制,backward 用 \texttt{view+sum} 归约。 $v, g, \beta$ 直接在 $H_V$ 维投影,q/k 在 $H$ 维投影。