% teach: % gap: 读者会用 autograd 但不知道手写 KDA backward 的具体展开 % takeaway: backward = 逆序遍历时间步, 每步求 dq/dk/dv/dg/dbeta + 累积 dS; GVA 反传 = view+sum % jump: 为什么 dS_dec 要加 dS_acc 和 -k⊗dr 两项 % omit: Triton backward 优化 \section{反向传播推导} \subsection{Forward 回顾} 逐步写下 forward(省略 batch/head 下标): \begin{align} S_{\mathrm{dec}} &= \exp(g_t) \odot S_{t-1} \tag{F1} \\ p_t &= k_t^T S_{\mathrm{dec}} \tag{F2} \\ r_t &= v_t - p_t \tag{F3} \\ a_t &= \beta_t \cdot k_t \tag{F4} \\ S_t &= S_{\mathrm{dec}} + a_t \otimes r_t \tag{F5} \\ o_t &= q_t^T S_t \tag{F6} \end{align} \subsection{反传公式(BPTT,$T \to 0$)} 设 $dS_{\mathrm{acc}}$ 是从时间步 $t$ 开始累积到 $S_t$ 上的梯度。逆序遍历: \paragraph{Step 1: $o_t = q_t^T S_t$} \begin{align} dS_{\mathrm{acc}} &\mathrel{+}= q_t \otimes do_t & \xrightarrow{\texttt{einsum('bhk,bhv->bhkv')}} & \quad \shape{B, HV, K, V} \\ dq_t &= do_t^T S_t & \xrightarrow{\texttt{einsum('bhv,bhkv->bhk')}} & \quad \shape{B, HV, K} \end{align} \paragraph{Step 2: $S_t = S_{\mathrm{dec}} + a_t \otimes r_t$} 外积的反传:$d(a \otimes r) = (\cdot)$,分解为: \begin{align} da_t &= \sum_v r_{t,v} \cdot dS_{\mathrm{acc},\cdot,v} & \xrightarrow{\texttt{einsum('bhv,bhkv->bhk')}} & \quad \shape{B, HV, K} \\ dr_t &= \sum_k a_{t,k} \cdot dS_{\mathrm{acc},k,\cdot} & \xrightarrow{\texttt{einsum('bhk,bhkv->bhv')}} & \quad \shape{B, HV, V} \end{align} \paragraph{Step 3: $a_t = \beta_t \cdot k_t$} \begin{align} d\beta_t &= \sum_k k_{t,k} \cdot da_{t,k} & \xrightarrow{\texttt{einsum('bhk,bhk->bh')}} & \quad \shape{B, HV} \\ dk_t^{(a)} &= \beta_t \cdot da_t & & \quad \shape{B, HV, K} \end{align} \paragraph{Step 4: $r_t = v_t - p_t = v_t - k_t^T S_{\mathrm{dec}}$} \begin{align} dv_t &= dr_t & & \shape{B, HV, V} \\ dk_t^{(r)} &= -S_{\mathrm{dec}}^T \cdot dr_t & \xrightarrow{\texttt{einsum('bhv,bhkv->bhk')}} & \quad \shape{B, HV, K} \\ dS_{\mathrm{dec}}^{(r)} &= -k_t \otimes dr_t & \xrightarrow{\texttt{einsum('bhv,bhk->bhkv')}} & \quad \shape{B, HV, K, V} \end{align} \paragraph{Step 5: 合并 $dS_{\mathrm{dec}}$ 并传递 $dg_t$, $dS_{t-1}$} \[ dS_{\mathrm{dec}}^{\mathrm{total}} = dS_{\mathrm{acc}} + dS_{\mathrm{dec}}^{(r)} = dS_{\mathrm{acc}} - k_t \otimes dr_t \] 因为 $S_{\mathrm{dec}} = \exp(g_t) \odot S_{t-1}$: \begin{align} dg_t &= S_{\mathrm{dec}} \odot dS_{\mathrm{dec}}^{\mathrm{total}} & \xrightarrow{\texttt{einsum('bhkv,bhkv->bhk')}} & \quad \shape{B, HV, K} \\ dS_{t-1} &= \exp(g_t) \odot dS_{\mathrm{dec}}^{\mathrm{total}} & & \quad \shape{B, HV, K, V} \end{align} \paragraph{Step 6: 合并 $dk_t$ 和 GVA 归约} \[ dk_t = dk_t^{(a)} + dk_t^{(r)} = \beta_t \cdot da_t - S_{\mathrm{dec}}^T \cdot dr_t \] GVA 反传($H_V \to H$): \[ dq_H = dq_{H_V}.\texttt{view}(B, T, H, G, K).\texttt{sum}(\text{dim}=3) \cdot \mathrm{scale} \] \[ dk_H = dk_{H_V}.\texttt{view}(B, T, H, G, K).\texttt{sum}(\text{dim}=3) \] \subsection{代码对照} \begin{codemathtop}{ops/reference/recurrent.py — KDAFunction.backward} \begin{lstlisting} for t in range(T - 1, -1, -1): q_t, k_t, b_t = q_ts[:,t], k_ts[:,t], b_ts[:,t] S_dec, r_t, a_t = S_decs[:,t], r_ts[:,t], a_ts[:,t] exp_g_t, do_t = exp_g_ts[:,t], do[:,t] # Step 1: o_t = q_t . S_t S_t = S_dec + einsum('bhk,bhv->bhkv', a_t, r_t) dS_acc += einsum('bhk,bhv->bhkv', q_t, do_t) dq_e[:,t] = einsum('bhv,bhkv->bhk', do_t, S_t) # Step 2: outer product grads da_t = einsum('bhv,bhkv->bhk', r_t, dS_acc) dr_t = einsum('bhk,bhkv->bhv', a_t, dS_acc) # Step 3: a_t = beta_t * k_t dbeta[:,t] = einsum('bhk,bhk->bh', k_t, da_t) dk_t_a = b_t.unsqueeze(-1) * da_t # Step 4: r_t = v_t - k_t . S_dec dv[:,t] = dr_t dS_dec_from_r = -einsum('bhv,bhk->bhkv', dr_t, k_t) dk_t_r = -einsum('bhv,bhkv->bhk', dr_t, S_dec) # Step 5: S_dec = exp(g) * S_{t-1} dS_dec_total = dS_acc + dS_dec_from_r dk_e[:,t] = dk_t_a + dk_t_r dg[:,t] = einsum('bhkv,bhkv->bhk', S_dec, dS_dec_total) dS_acc = exp_g_t.unsqueeze(-1) * dS_dec_total # Step 6: GVA reduce dq_H = dq_e.view(B,T,H,G,K).sum(dim=3) * scale dk_H = dk_e.view(B,T,H,G,K).sum(dim=3) \end{lstlisting} \end{codemathtop} \begin{warningbox}{$dS_{\mathrm{dec}}^{\mathrm{total}}$ 为什么包含两项?} $S_t = S_{\mathrm{dec}} + a_t \otimes r_t$,$S_{\mathrm{dec}}$ 同时参与了: \begin{enumerate}[nosep] \item 直接传递到 $dS_{\mathrm{acc}}$(作为 $S_t$ 的一部分被读出) \item 通过 $r_t = v_t - k_t \cdot S_{\mathrm{dec}}$ 参与 delta rule \end{enumerate} 所以 $dS_{\mathrm{dec}}^{\mathrm{total}} = dS_{\mathrm{acc}} + dS_{\mathrm{dec}}^{(r)}$, 两条路径的梯度要\textbf{加}起来(chain rule 分叉处求和)。 \end{warningbox} \subsection{本章小结} KDA backward 是 BPTT 展开:逆序遍历时间步,每步 6 个 einsum + 一次 $dS$ 累积更新。 GVA 反传在最后做 \texttt{view+sum}。手写 backward 的关键是正确处理 $dS_{\mathrm{dec}}$ 的两条梯度路径(直接传递 + 通过 $r_t$ 的 delta rule 路径)。