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.
This commit is contained in:
@@ -0,0 +1,127 @@
|
||||
% 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。
|
||||
Reference in New Issue
Block a user