% 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] = """ 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 逐位相同。