Files
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

158 lines
5.1 KiB
TeX
Raw Permalink 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: 读者会用 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 路径)。