Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
158 lines
5.1 KiB
TeX
158 lines
5.1 KiB
TeX
% 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 路径)。
|