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

128 lines
4.4 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: 读者知道 softmax attention 但不知道线性注意力怎么维护状态矩阵
% takeaway: KDA 用 delta rule 逐步更新 [K,V] 状态矩阵, 写入=擦旧写新, 每步 O(KV)
% jump: 为什么 r_t = v - k·S 而不是直接用 v?delta rule 的"先擦再写"
% omit: KDA 论文的 related work、实验细节
\section{KDA 递归核心}
\splabel{C1}
\subsection{动机:从 softmax 到状态矩阵}
标准 attention 每个 token 都要回看所有历史,复杂度 $O(T^2)$。
线性注意力换掉 softmax,把 $\sum_j v_j k_j^T$ 压成一个 $K \times V$ 的状态矩阵 $S$,
每步只做 $o_t = q_t \cdot S$,复杂度降到 $O(T \cdot K \cdot V)$。
但裸线性注意力的问题是:$S$ 只能加,不能改。写进去的信息永远在那里。
KDA 的核心想法是给 $S$ 加两个操作:\textbf{衰减}(逐渐忘记旧信息)和
\textbf{delta rule}(先擦旧的,再写新的)。
\begin{importantbox}{如果你只记一件事}
KDA 的状态更新 = 衰减旧状态 + $\beta_t k_t$ 写入 $(v_t - k_t \cdot S_{\mathrm{dec}})$。
减去 $k_t \cdot S_{\mathrm{dec}}$ 就是"先把 $k_t$ 方向的旧预测擦掉"。
\end{importantbox}
\subsection{逐步公式}
\noindent\textbf{输入张量:}
\begin{center}
\begin{tabular}{lll}
\toprule
符号 & 形状 & 含义 \\
\midrule
$q_t$ & \shape{B, HV, K} & query(已经 repeat\_interleave 到 HV) \\
$k_t$ & \shape{B, HV, K} & key(同上) \\
$v_t$ & \shape{B, HV, V} & value \\
$g_t$ & \shape{B, HV, K} & gate(log-space 衰减率,逐维) \\
$\beta_t$ & \shape{B, HV} & 写入强度标量 \\
$S_{t-1}$ & \shape{B, HV, K, V} & 上一步的 KV 状态 \\
\bottomrule
\end{tabular}
\end{center}
\noindent\textbf{四步更新:}
\begin{enumerate}[leftmargin=2em]
\item \textbf{衰减旧状态}(逐元素,$g_t$ 是 log-space 所以取 exp):
\[
S_{\mathrm{dec}} = \exp(g_t) \odot S_{t-1}
\qquad \shape{B, HV, K, V}
\]
\item \textbf{计算残差}(先用 $k_t$ 查旧状态,得到"旧预测",再减掉):
\[
p_t = \sum_k k_{t,k} \cdot S_{\mathrm{dec},k,\cdot}
= \texttt{einsum('bhk, bhkv -> bhv')}
\qquad \shape{B, HV, V}
\]
\[
r_t = v_t - p_t \qquad \shape{B, HV, V}
\]
\item \textbf{写入状态}(外积 rank-1 更新):
\[
a_t = \beta_t \cdot k_t \qquad \shape{B, HV, K}
\]
\[
S_t = S_{\mathrm{dec}} + a_t \otimes r_t
= S_{\mathrm{dec}} + \texttt{einsum('bhk, bhv -> bhkv')}
\qquad \shape{B, HV, K, V}
\]
\item \textbf{读出}:
\[
o_t = \frac{1}{\sqrt{K}} \cdot q_t \cdot S_t
= \texttt{einsum('bhk, bhkv -> bhv')}
\qquad \shape{B, HV, V}
\]
\end{enumerate}
\subsection{代码对照}
\begin{codemathtop}{ops/reference/recurrent.py — naive\_kda\_fwd (核心循环)}
\begin{lstlisting}
for t in range(T):
q_t = qe[:, t] # [B, HV, K]
k_t = ke[:, t] # [B, HV, K]
v_t = v[:, t] # [B, HV, V]
g_t = g[:, t] # [B, HV, K]
b_t = beta[:, t] # [B, HV]
# Step 1: decay
S_dec = S * g_t.exp().unsqueeze(-1) # [B,HV,K,V]
# Step 2: residual (delta rule)
p_t = einsum('bhk, bhkv -> bhv', k_t, S_dec)
r_t = v_t - p_t # [B,HV,V]
# Step 3: write (rank-1 update)
a_t = b_t.unsqueeze(-1) * k_t # [B,HV,K]
S = S_dec + einsum('bhk, bhv -> bhkv', a_t, r_t)
# Step 4: read
o[:, t] = einsum('bhk, bhkv -> bhv', q_t, S)
\end{lstlisting}
\end{codemathtop}
\begin{warningbox}{为什么 exp(g\_t) 要 unsqueeze(-1)?}
$g_t$ 的形状是 \shape{B, HV, K},而 $S$ 是 \shape{B, HV, K, V}。
衰减是在 $K$ 维上逐元素(同一 $k$ 索引的所有 $v$ 维度共享同一个衰减率),
所以 \texttt{exp(g\_t).unsqueeze(-1)} 把 K 维 broadcast 到 $K \times V$。
\end{warningbox}
\subsection{Delta rule 的直觉}
\begin{knowledgebox}{为什么减去 $k_t \cdot S_{\mathrm{dec}}$?}
把 $S$ 想象成一个 $K \to V$ 的线性映射。用 $k_t$ 去查它,得到的 $p_t = k_t^T S$
就是``旧状态对 $k_t$ 方向的预测''。如果 $p_t$ 已经很接近 $v_t$,说明这个方向的信息
已经写好了,不需要再写。$r_t = v_t - p_t$ 就是``需要修正的量''。
这就是 Widrow-Hoff delta rule:不是盲目地加,而是只修正误差。
\end{knowledgebox}
\subsection{本章小结}
KDA 的状态更新 = 衰减旧状态 + $\beta_t k_t$ 写入残差 $(v_t - k_t \cdot S_{\mathrm{dec}})$。
每步复杂度 $O(K \cdot V)$(两次矩阵-向量乘 + 一次外积),不需要 softmax。