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