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:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+148
View File
@@ -0,0 +1,148 @@
% teach:
% gap: 读者知道递归形式但不知道怎么在 GPU 上并行, 以为只能一步步算
% takeaway: chunk 内用下三角线性系统并行求解, chunk 间递推状态, 数值等价于 naive recurrent
% jump: 论文直接写了 triangular solve 但没解释为什么要 solve 而不是直接矩阵乘
% omit: Triton 实现细节
\section{分块并行计算(Chunkwise)}
\splabel{C3}
\subsection{为什么需要分块?}
Naive recurrent 一步一步算,$T$ 步串行,GPU 利用率低。
分块的想法是把序列切成 $T/C$ 个长度为 $C$ 的 chunk:
\begin{itemize}[nosep]
\item \textbf{chunk 内}:$C$ 个 token 之间的依赖可以用矩阵运算并行处理
\item \textbf{chunk 间}:状态 $S$ 从上一个 chunk 传到下一个,仍然是递推
\end{itemize}
\begin{importantbox}{如果你只记一件事}
Chunkwise = chunk 内并行 + chunk 间递推。数值结果与 naive recurrent 逐位一致。
\end{importantbox}
\subsection{chunk 内的 cumsum 与下三角解}
在每个 chunk 内,先对 $g$ 做 cumsum(前缀和),这样衰减就变成了相对距离的函数:
\[
g_{\mathrm{cum},i} = \sum_{j=0}^{i} g_j, \qquad
\text{token } i \text{ 对 token } j \text{ 的衰减} = \exp(g_{\mathrm{cum},i} - g_{\mathrm{cum},j})
\]
定义 \textbf{decayed dot} 矩阵(chunk 内 $C \times C$):
\[
A_{ij} = \langle x_i, \exp(g_{\mathrm{cum},i} - g_{\mathrm{cum},j}) \cdot k_j \rangle
\qquad \shape{..., C, C}
\]
这个矩阵的构造是 chunk 内计算的核心。用它可以构造一个下三角线性系统:
\[
M = I + \mathrm{tril}(A_{kk} \cdot \beta, \text{diagonal}=-1)
\qquad \shape{..., C, C}
\]
\[
M \cdot W = \exp(g_{\mathrm{cum}}) \cdot k \qquad \Rightarrow \qquad
W = M^{-1} (\exp(g_{\mathrm{cum}}) \cdot k)
\]
\[
M \cdot U = v \qquad \Rightarrow \qquad U = M^{-1} v
\]
\subsection{代码对照}
\begin{codemathtop}{ops/reference/chunkwise.py — naive\_chunk\_kda (核心)}
\begin{lstlisting}
# Rearrange: [B,T,H,K] -> [B,H,N,C,K] where N=T/C
q, k = [rearrange(x, 'b (n c) h d -> b h n c d', c=C)
.repeat_interleave(HV//H, dim=1) for x in (q, k)]
v, g = [rearrange(x, 'b (n c) h d -> b h n c d', c=C)
for x in (v, g)]
beta = rearrange(beta, 'b (n c) h -> b h n c', c=C)
q = q * scale
g = g.cumsum(dim=-2) # chunk 内 cumsum
# Construct triangular system
A_kk = _decayed_dot(k, k, g) # [B,HV,N,C,C]
M = eye + (A_kk * beta[...,None,:]).masked_fill(mask_upper, 0)
W = solve_triangular(M, g.exp() * k, upper=False)
U = solve_triangular(M, v, upper=False)
# A_qk: query 对 key 的 decayed dot (含对角线)
A_qk = (_decayed_dot(q, k, g) * beta[...,None,:])
.masked_fill(mask_strict_upper, 0)
\end{lstlisting}
\end{codemathtop}
\subsection{chunk 间递推}
每个 chunk 内算完后,用 $W$ 和 $U$ 来处理跨 chunk 的状态:
\begin{codemathtop}{ops/reference/chunkwise.py — chunk 间循环}
\begin{lstlisting}
S = zeros(B, HV, K, V) # inter-chunk state
for n in range(T // C):
# r = "local residual, adjusted by cross-chunk state"
r = U[:,:,n] - W[:,:,n] @ S # [B,HV,C,V]
# output: cross-chunk part + intra-chunk part
o[:,:,n] = (q_n * g_n.exp()) @ S + A_qk[:,:,n] @ r
# update cross-chunk state
decay = (g_n[:,:,-1:,:] - g_n).exp() # decay to chunk end
S = S * g_n[:,:,-1,:,None].exp() # decay old state
S = S + (decay * k_n).T @ (r * beta_n) # write new
\end{lstlisting}
\end{codemathtop}
\subsection{形状流水线}
\begin{center}
\begin{tabular}{lll}
\toprule
变量 & 形状 & 说明 \\
\midrule
\texttt{q, k} (chunked) & \shape{B, HV, N, C, K} & $N = T/C$ 个 chunk \\
\texttt{v, g} (chunked) & \shape{B, HV, N, C, V/K} & \\
\texttt{beta} (chunked) & \shape{B, HV, N, C} & \\
\texttt{A\_kk} & \shape{B, HV, N, C, C} & key-key decayed dot \\
\texttt{M} & \shape{B, HV, N, C, C} & 下三角系统 \\
\texttt{W} & \shape{B, HV, N, C, K} & $M^{-1}(\exp(g) \cdot k)$ \\
\texttt{U} & \shape{B, HV, N, C, V} & $M^{-1} v$ \\
\texttt{A\_qk} & \shape{B, HV, N, C, C} & query-key decayed dot \\
\texttt{S} & \shape{B, HV, K, V} & 跨 chunk 状态 \\
\texttt{r} & \shape{B, HV, C, V} & 调整后的残差 \\
\texttt{o (chunk n)} & \shape{B, HV, C, V} & 本 chunk 输出 \\
\bottomrule
\end{tabular}
\end{center}
\begin{warningbox}{为什么用 triangular solve 而不是直接矩阵乘?}
Delta rule 的"先擦再写"引入了 chunk 内 token 之间的递归依赖:
token $i$ 的写入依赖 token $j < i$ 的写入结果。
这个依赖关系恰好形成一个下三角线性系统 $M \cdot x = b$,
用 \texttt{solve\_triangular} 可以在 $O(C^2)$ 内并行求解,
而展开递归需要 $O(C)$ 步串行。
\end{warningbox}
\subsection{Decayed dot 函数}
\begin{codemathtop}{ops/reference/chunkwise.py — \_decayed\_dot}
\begin{lstlisting}
def _decayed_dot(x, k, g):
"""A[..., i, j] = <x_i, exp(g_i - g_j) * k_j>"""
C = x.shape[-2]
A = empty(*x.shape[:-2], C, C)
for i in range(C):
decay = (g[..., i:i+1, :] - g).exp() # [.., 1, K] - [.., C, K]
A[..., i, :] = einsum('...jk,...jk->...j',
x[..., i, None, :] * decay, k)
return A
\end{lstlisting}
\end{codemathtop}
\noindent 这是一个 $C \times C$ 的矩阵,每个元素 $(i,j)$ 是
$x_i$ 和 $\exp(g_i - g_j) \cdot k_j$ 的内积。Triton 实现会把这个双循环融合成一个 kernel。
\subsection{本章小结}
分块把 $T$ 步串行拆成 $T/C$ 个 chunk,chunk 内用下三角 solve 并行处理 delta rule 依赖,
chunk 间递推状态 $S$。最终输出与 naive recurrent 逐位相同。