Files
K3/notes/sections/sec-04.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

115 lines
4.1 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: 读者不知道 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$ 维投影。