Files
K3/notes/sections/sec-03.tex
T
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

149 lines
5.5 KiB
TeX
Raw 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: 读者知道递归形式但不知道怎么在 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 逐位相同。