ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

flash-linear-attention 中 DeltaNet 的分块并行化:WY 表示推导与 Triton 内核实现

2026/9/17 22:24:34 拓冰建站 浏览量
flash-linear-attention 中 DeltaNet 的分块并行化:WY 表示推导与 Triton 内核实现 flash-linear-attention 中 DeltaNet 的分块并行化WY 表示推导与 Triton 内核实现【免费下载链接】flash-linear-attention Efficient implementations for emerging model architectures项目地址: https://gitcode.com/GitHub_Trending/fl/flash-linear-attention本文以fla/ops/delta_rule/README.md中Chunkwise-form Parallelism of DeltaNet一节为核心完整继承其对 DeltaNet delta rule 递推式做分块展开、WY 表示与T矩阵求解的全部数学推导并逐公式对照仓库中的 Triton 内核wy_fast.py、chunk.py及公共内核讲清从prepare_wy_repr_fwd到状态更新与输出计算的完整调用链最后给出chunk_delta_rule的可用调用示例与测试验证方式。1. 背景DeltaNet 的 delta rule 递推与并行化难题DeltaNet 属于delta rule线性注意力族维护一个跨时间的状态矩阵 $\mathbf{S} \in \mathbb{R}^{d_k \times d_v}$每一步先按误差修正方式更新记忆再读取输出。仓库中的参考实现 delta_rule_recurrence 直接给出了这一递推for i in range(l): _v _v - (S.clone() * _k[..., None]).sum(-2) # 读出旧记忆 k 对应的值 _v _v * beta_i # 按 beta 缩放修正量 S S.clone() _k.unsqueeze(-1) * _v.unsqueeze(-2) # 外积更新 o[:, :, i] torch.einsum(bhd,bhdm-bhm, _q, S) # 读取由于 $\mathbf{S}$ 的更新严格依赖上一步朴素递推在序列维上无法并行只能逐步串行执行。README 一节扩展了 DeltaNet 论文arXiv:2406.06484附录 B 中的公式给出一种chunkwise分块并行化推导使得块内计算全部化归为矩阵乘法、可用 Tensor Core 加速。下文按原文档的推导顺序逐节展开。2. 分块展开$\mathbf{S}^r \mathbf{P}^r \mathbf{S}^0 \mathbf{H}^r$为减少符号负担推导聚焦第一个 chunk记 $\mathbf{S}^r \mathbf{S}_{[1]}^r$。对递推式做部分展开可得\begin{equation} \begin{aligned} \mathbf{S}^r \underbrace{\left(\prod_{i1}^r \mathbf{I} - \beta^i \mathbf{k}^i \mathbf{k}^{i\top} \right)}_{: \mathbf{P}^r} \cdot\mathbf{S}^{0} \overbrace{\sum_{i1}^{r} \underbrace{\left(\prod_{ji1}^r \mathbf{I} - \beta^j \mathbf{k}^j \mathbf{k}^{j\top} \right)}_{: \mathbf{P}_{i1}^r}\beta^i \mathbf{k}^i\mathbf{v}^{i\top}}^{:\mathbf{H}^r} \\ \mathbf{P}^r \cdot \mathbf{S}^{0} \mathbf{H}^r \end{aligned} \end{equation}其中 $\mathbf{P}_i^r$ 是一连串广义 Householder 型矩阵 $\mathbf{I} - \beta^i \mathbf{k}^i \mathbf{k}^{i\top}$ 的累积乘积$\mathbf{P}_1^r$ 简记为 $\mathbf{P}^r$。难点在于直接计算 $\mathbf{P}^r$ 需要 $r$ 次 $d_k \times d_k$ 矩阵乘且依赖性强、无法并行。2.1 $\mathbf{P}^r$ 的 WY 表示$\mathbf{P}^r$ 可以用经典的 WY 表示改写为低秩求和形式\begin{equation} \mathbf{P}^{r} \mathbf{I} - \sum_{i1}^{r}\mathbf{k}^i\mathbf{w}^{i\top} \in \mathbb{R}^{d_k \times d_k};\qquad \mathbf{w}^r \beta^r \left(\mathbf{k}^r - \sum_{i1}^{r-1} \left(\mathbf{k}^{r\top}\mathbf{k}^i \right)\mathbf{w}^i \right) \in \mathbb{R}^{d_k} \end{equation}原文档用数学归纳法证明了这一点其展开过程为\begin{align*} \mathbf{P}^{r} \prod_{i1}^r \mathbf{I} - \beta^i \mathbf{k}^i \mathbf{k}^{i\top} \\ \left(\mathbf{I} - \beta^r \mathbf{k}^r \mathbf{k}^{r\top}\right)\mathbf{P}^{r-1} \\ \left(\mathbf{I} - \beta^r \mathbf{k}^r \mathbf{k}^{r\top}\right)\left(\mathbf{I} - \sum_{i1}^{r-1}\mathbf{k}^i\mathbf{w}^{i\top}\right) \\ \mathbf{I} - \sum_{i1}^{r-1}\mathbf{k}^i\mathbf{w}^{i\top} - \beta^r \mathbf{k}^r \mathbf{k}^{r\top} \beta^r\mathbf{k}^r \mathbf{k}^{r\top} \left(\sum_{i1}^{r-1}\mathbf{k}^i\mathbf{w}^{i\top}\right) \\ \mathbf{I} - \sum_{i1}^{r-1}\mathbf{k}^i\mathbf{w}^{i\top} - \beta^r \mathbf{k}^r \left(\mathbf{k}^{r} - \left(\sum_{i1}^{r-1}\left(\mathbf{k}^{r\top} \mathbf{k}^i\right)\mathbf{w}^{i}\right) \right)^\top \\ \mathbf{I} - \sum_{i1}^{r}\mathbf{k}^i\mathbf{w}^{i\top} \end{align*}归纳的关键步骤是第三行到第四行把 $\mathbf{P}^{r-1}$ 代入后利用 $\left(\sum_i \mathbf{k}^i \mathbf{w}^{i\top}\right)\mathbf{k}^r \sum_i (\mathbf{k}^{i\top}\mathbf{k}^r)\mathbf{w}^i$ 这一标量内积结合律将 $-\beta^r \mathbf{k}^r \mathbf{k}^{r\top}$ 与交叉项合并成 $\mathbf{k}^r$ 与一个向量的外积从而保持 WY 形式封闭。2.2 $\mathbf{H}^r$ 的 WY 表示$\mathbf{H}^r$ 可以同构地表示为\begin{equation} \mathbf{H}^{r} \sum_{i1}^{r} \mathbf{k}^i \mathbf{u}^{i\top} \in \mathbb{R}^{d_k \times d_v};\qquad \mathbf{u}^r \beta^r \left(\mathbf{v}^r - \sum_{i1}^{r-1} \left(\mathbf{k}^{r\top}\mathbf{k}^i\right) \mathbf{u}^i \right)\in \mathbb{R}^{d_v} \end{equation}归纳证明与上节完全平行\begin{align*} \mathbf{H}^{r} \sum_{i1}^{r} \mathbf{P}_{i1}^r \beta^i \mathbf{k}^i \mathbf{v}^{i\top}\\ \left(\mathbf{I} - \beta^r \mathbf{k}^r \mathbf{k}^{r\top}\right) \mathbf{H}^{r-1} \beta^r \mathbf{k}^r \mathbf{v}^{r\top}\\ \sum_{i1}^{r-1}\mathbf{k}^i \mathbf{u}^{i\top} - \beta^r \mathbf{k}^r \mathbf{k}^{r\top} \sum_{i1}^{r-1}\mathbf{k}^i \mathbf{u}^{i\top} \beta^r \mathbf{k}^r \mathbf{v}^{r\top}\\ \sum_{i1}^{r-1}\mathbf{k}^i \mathbf{u}^{i\top} \mathbf{k}^r \left(\beta^r \mathbf{v}^{r\top}-\beta^r \mathbf{k}^{r\top} \sum_{i1}^{r-1}\mathbf{k}^i \mathbf{u}^{i\top}\right) \\ \sum_{i1}^{r-1} \mathbf{k}^i \mathbf{u}^{i\top} \mathbf{k}^r \beta^r\left(\mathbf{v}^r-\sum_{i1}^{r-1}\left(\mathbf{k}^{r\top}\mathbf{k}^{i}\right)\mathbf{u}^{i} \right)^\top \\ \sum_{i1}^{r} \mathbf{k}^i \mathbf{u}^{i\top} \end{align*}2.3 矩阵形式与 $\mathbf{T}$ 的求解把逐行向量堆叠成矩阵$\mathbf{P}$ 与 $\mathbf{H}$ 写为块内共 $C$ 个 token\begin{equation} \mathbf{P}\mathbf{I}-\mathbf{K}^\top\mathbf{W} \in \mathbb{R}^{d_k \times d_k}, \qquad \mathbf{H}\mathbf{K}^\top\mathbf{U} \in \mathbb{R}^{d_k\times d_v} \end{equation}对 $\mathbf{W}$ 的递推式取矩阵形式$\mathbf{W} \mathrm{diag}(\beta)\mathbf{K} - \mathrm{tril}(\mathrm{diag}(\beta) \mathbf{K}\mathbf{K}^\top, -1)\mathbf{W}$移项后得到下三角线性方程组\begin{align*} \mathbf{W} \mathrm{diag}(\beta) \mathbf{K} - \mathrm{tril}(\mathrm{diag}(\beta) \mathbf{K}\mathbf{K}^\top, -1)\mathbf{W}\\ \left(\mathbf{I} \mathrm{tril}(\mathrm{diag}(\beta) \mathbf{K}\mathbf{K}^\top, -1)\right) \mathbf{W} \mathrm{diag}(\beta) \mathbf{K} \end{align*}$\mathbf{U}$ 满足同样的方程右端换成 $\mathrm{diag}(\beta)\mathbf{V}$。因此定义块内 $C\times C$ 的下三角求逆矩阵 $\mathbf{T}$\begin{align*} \mathbf{T} \left(\mathbf{I} \mathrm{tril}\left(\mathrm{diag}(\beta)\mathbf{K} \mathbf{K}^\top,-1\right)\right)^{-1}\mathrm{diag}\left(\beta\right)\in \mathbb{R}^{C \times C}\\ \mathbf{W} \mathbf{T} \mathbf{K}\in \mathbb{R}^{C \times d_k}\\ \mathbf{U} \mathbf{T}\mathbf{V}\in \mathbb{R}^{C \times d_v} \end{align*}这正是分块并行的核心收益$\mathbf{T}$ 的规模只是块大小 $C$实现中默认 64其求逆是一次小规模稠密下三角操作而 $\mathbf{W}\mathbf{TK}$、$\mathbf{U}\mathbf{TV}$ 以及后续所有状态更新全部是 GEMM可充分利用 Tensor Core。2.4 最终的分块并行算法将 $\mathbf{W}$、$\mathbf{U}$ 代回 $\mathbf{S}\mathbf{P}\mathbf{S}^0\mathbf{H}$并利用 $\mathbf{P}\mathbf{I}-\mathbf{K}^\top\mathbf{W}$、$\mathbf{H}\mathbf{K}^\top\mathbf{U}$得到对硬件友好的最终形式\begin{equation} \begin{aligned} \mathbf{S} \mathbf{P}\cdot\mathbf{S}^0 \mathbf{H} \\ \mathbf{S}^0 \mathbf{K}^\top (\mathbf{U} -\mathbf{W} \mathbf{S}^0) \in \mathbb{R}^{d_k \times d_v}\\ \mathbf{O} \mathbf{Q} \mathbf{S}^0 (\mathbf{Q} \mathbf{K}^{\top} \odot \mathbf{M}) \left(\mathbf{U} - \mathbf{W} \mathbf{S}^0\right) \in \mathbb{R}^{C \times d_v} \end{aligned} \end{equation}其中 $\mathbf{M}$ 是块内严格上三角掩码保证因果性token 只能看到同块内其之前的 token。整个 chunk 的计算量可概括为计算块内 $\mathbf{K}\mathbf{K}^\top$ 并乘 $\mathrm{diag}(\beta)$取严格下三角对 $C\times C$ 下三角矩阵做求逆再乘 $\mathrm{diag}(\beta)$ 得到 $\mathbf{T}$$\mathbf{W}\mathbf{TK}$、$\mathbf{U}\mathbf{TV}$两次 GEMM跨块$\mathbf{V}{new} \mathbf{U} - \mathbf{W}\mathbf{S}^0$$\mathbf{S} \leftarrow \mathbf{S}^0 \mathbf{K}^\top\mathbf{V}{new}$输出$\mathbf{O} \mathbf{Q}\mathbf{S}^0 (\mathbf{Q}\mathbf{K}^\top \odot \mathbf{M})\mathbf{V}_{new}$。3. 源码实现推导公式与内核的一一对应fla/ops/delta_rule包导出三个入口chunk_delta_rule分块并行训练内核、fused_chunk_delta_rule与fused_recurrent_delta_rule逐 token 递推参考见 fla/ops/delta_rule/init.py。上文第 2 节推导中的每一步都对应到一个具体内核。3.1 第一步$\mathrm{tril}(\mathrm{diag}(\beta)\mathbf{K}\mathbf{K}^\top, -1)$ ——chunk_scaled_dot_kkt_fwdprepare_wy_repr_fwd 是整个 WY 表示的入口第一步调用 chunk_scaled_dot_kkt_fwd。该公共内核的文档字符串写明其职责即Compute beta * K * K^T内核核心逻辑为b_A tl.dot(b_k, tl.trans(b_k), b_A) # 块内 K·KᵀGEMM b_A * b_b[:, None] # 行方向乘 diag(β) m_A (o_t[:, None] o_t[None, :]) ... b_A tl.where(m_A, b_A, 0) # 只保留严格下三角tril, -1这与推导中 $\mathrm{tril}(\mathrm{diag}(\beta)\mathbf{K}\mathbf{K}^\top, -1)$ 完全一致且输出按output_dtypetorch.float32存储为后续求逆提供数值精度。3.2 第二步$(\mathbf{I} \cdot)^{-1}\mathrm{diag}(\beta)$ ——solve_tril紧接着 solve_tril 对 $C\times C$ 下三角矩阵做块求逆得到推导中的 $\mathbf{T}$ 矩阵。实现上采用 16×16 子块求逆再逐级合并如merge_16x16_to_32x32_inverse_kernel的分治策略且通过环境变量FLA_TRIL_PRECISION取值ieee/tf32/tf32x3默认ieee控制点积精度——因为该矩阵接近奇异精度选择直接影响数值稳定性。3.3 第三步$\mathbf{W}\mathbf{TK}$、$\mathbf{U}\mathbf{TV}$ ——recompute_w_u_fwdrecompute_w_u_fwd_kernel 对 $\mathbf{V}$ 和 $\mathbf{K}$ 两个维度分别做分块 GEMMb_u tl.dot(b_A.to(b_vb.dtype), b_vb, allow_tf32False) # U T·(β⊙V) ... b_w tl.dot(b_A.to(b_kb.dtype), b_kb, allow_tf32False) # W T·(β⊙K)注意内核中先做b_vb b_v * b_beta[:, None]再与 $\mathbf{T}$ 相乘即把 $\mathrm{diag}(\beta)$ 右乘吸收进了 $\mathbf{K}$、$\mathbf{V}$与推导 $\mathbf{T} (\mathbf{I} \mathrm{tril}(\cdot))^{-1}\mathrm{diag}(\beta)$ 的写法在数值上等价。3.4 跨块状态更新与输出chunk.py的三段流水线chunk_delta_rule_fwd 把上述结果接成完整前向w, u, A prepare_wy_repr_fwd(kk, vv, betabeta, ...) # §2.1–§2.3得到 W, U, T h, v_new, final_state chunk_gated_delta_rule_fwd_h( # §2.4V_new U - W·S⁰, S ← S⁰ KᵀV_new kk, ww, uu, initial_stateinitial_state, ...) o chunk_fwd_o(qq, kk, vv_new, hh, scalescale, ...) # O Q·S⁰ (QKᵀ⊙M)·V_new其中 chunk_gated_delta_rule_fwd_h 沿 chunk 串行但块内并行地推进状态每个 chunk 只依赖上一次的 $\mathbf{S}^0$串行深度仅为 $T/C$ 而非 $T$chunk_fwd_o 计算跨块项 $\mathbf{Q}\mathbf{S}^0$ 与块内项 $(\mathbf{Q}\mathbf{K}^\top\odot\mathbf{M})\mathbf{V}_{new}$。反向传播在 chunk_delta_rule_bwd 中通过recompute_w_u_fwd重算 $\mathbf{W}$、$\mathbf{U}$而非保存再依次求 $\mathbf{dV}$、$\mathbf{dS}$、$\mathbf{dQ}/\mathbf{dK}/\mathbf{dW}$最后由 prepare_wy_repr_bwd 把梯度传回 $\mathbf{dK}$、$\mathbf{dV}$、$\mathbf{d\beta}$——注意它对下三角线性系统 $\mathbf{T}$ 的转置做了对应的链式求导b_dA tl.dot(b_A, b_dA)形式的三角系统解。4. 参考实现对照naive.py中的分块版式推导fla/ops/delta_rule/naive.py 提供了不依赖 Triton 的纯 PyTorch 分块实现delta_rule_chunkwise是验证推导正确性最直观的可运行公式# 计算 (I - tri(diag(beta) KK^T))^{-1}逐行递推得到下三角求逆 attn -(k_beta k.transpose(-1, -2)).masked_fill(mask, 0) for i in range(1, chunk_size): attn[..., i, :i] attn[..., i, :i] (attn[..., i, :, None] * attn[..., :, :i]).sum(-2) attn attn torch.eye(chunk_size, dtypetorch.float, deviceq.device) u attn v # U T·V w attn k_beta # W T·(β⊙K) for i in range(0, l // chunk_size): attn (q_i k_i.transpose(-1, -2)).masked_fill_(mask) # QKᵀ ⊙ M上三角掩码 u_i u[:, :, i] - w[:, :, i] S # V_new U - W·S⁰ o_inter q_i S # Q·S⁰ o[:, :, i] o_inter attn u_i S S k_i.transpose(-1, -2) u_i # S ← S⁰ Kᵀ·V_new这段代码与 README 推导的最后两个公式逐行对应attn k_beta对应 $\mathbf{W}\mathbf{TK}$u_i u - w S对应 $\mathbf{U} - \mathbf{W}\mathbf{S}^0$o qS attnu_i对应 $\mathbf{O} \mathbf{Q}\mathbf{S}^0 (\mathbf{Q}\mathbf{K}^\top\odot\mathbf{M})(\mathbf{U}-\mathbf{W}\mathbf{S}^0)$。除naive.py外fla/ops/delta_rule/parallel.py 还实现了全并行$O(L^2)$版本delta_rule_parallel同样以 $\mathbf{T}$ 矩阵代码注释为 compute (I - tri(diag(beta) KK^T))^{-1}为中间量可用于不同复杂度权衡下的交叉验证。5. API 使用与参数说明训练时直接调用chunk_delta_rule签名与约束见 fla/ops/delta_rule/chunk.pyimport torch import torch.nn.functional as F from einops import rearrange from fla.ops.delta_rule import chunk_delta_rule # 等长输入 B, T, H, K, V 4, 2048, 4, 512, 512 q torch.randn(B, T, H, K, dtypetorch.bfloat16, devicecuda) k F.normalize(torch.randn(B, T, H, K, dtypetorch.bfloat16, devicecuda), p2, dim-1) v torch.randn(B, T, H, V, dtypetorch.bfloat16, devicecuda) beta torch.rand(B, T, H, dtypetorch.bfloat16, devicecuda).sigmoid() h0 torch.randn(B, H, K, V, dtypetorch.bfloat16, devicecuda) o, ht chunk_delta_rule( q, k, v, beta, initial_stateh0, # [N, H, K, V] 的初始状态可省略 output_final_stateTrue, # 输出最终状态 [N, H, K, V] ) # 变长输入B 必须为 1并传入 cu_seqlens q, k, v, beta map(lambda x: rearrange(x, b t ... - 1 (b t) ...), (q, k, v, beta)) cu_seqlens q.new_tensor([0, 2048, 4096, 6144, 8192], dtypetorch.long) o, ht chunk_delta_rule(q, k, v, beta, initial_stateh0, output_final_stateTrue, cu_seqlenscu_seqlens)关键参数与实现约束均出自chunk_delta_rule的文档字符串与断言逻辑参数说明取值 / 约束q, k, v形状[B, T, H, K]/[B, T, H, V]三者 dtype 必须一致且不支持 float32需用 bfloat16 等低精度见源码断言beta形状[B, T, H]即推导中的 $\beta^i$通常取 sigmoid 输出scale注意力缩放因子默认1 / sqrt(K)initial_state形状[N, H, K, V]等长序列时N B变长时N len(cu_seqlens) - 1output_final_state是否返回最终状态[N, H, K, V]默认Falseuse_qk_l2norm_in_kernel在内核内对 q/k 做 L2 归一化以省显存默认False见 l2norm 前向/反向的接入chunk_size分块大小 $C$对应推导中 $\mathbf{T}$ 的规模仅允许 16 / 32 / 64默认 64源码中有显式校验cu_seqlens/cu_seqlens_cpu变长训练的 FlashAttention 风格累计长度使用时要求q.shape[0] 1cu_seqlens_cpu为 CPU 副本以避免设备同步从源码结构看chunk_size直接决定 $\mathbf{T}\in\mathbb{R}^{C\times C}$ 的求逆成本与串行深度$C$ 越大单块内 GEMM 越高效但三角求逆越贵、跨块串行步数越少默认值 64 与prepare_wy_repr_fwd的默认参数一致。6. 正确性验证chunk 与递推参考实现的对比测试仓库测试 tests/ops/test_delta.py 用fused_recurrent_delta_rule逐 token 递推作为参考实现覆盖多种(B, T, H, D)、scale与use_qk_l2norm_in_kernel组合对前向输出、最终状态以及dq / dk / dv / dbeta / dh0五个梯度做逐张量对比assert_close容差 0.006~0.008并单独用chunk_size ∈ {16, 32, 64}的参数化用例test_chunk_with_chunk_size验证分块大小的敏感性。运行方式pytest tests/ops/test_delta.py需要 CUDA 环境与 Triton测试中对 k 默认做了F.normalize(..., p2, dim-1)的 L2 归一化或在开启use_qk_l2norm_in_kernel时交给内核完成两种路径都会进入同一组 Triton 内核。7. 小结fla/ops/delta_rule/README.md给出的推导回答了 DeltaNet 分块并行的核心问题如何把块内强依赖的 $r$ 次 Householder 乘积压成一次 $C\times C$ 下三角求逆 若干 GEMM。关键链条是递推展开得到 $\mathbf{S}^r \mathbf{P}^r\mathbf{S}^0 \mathbf{H}^r$WY 表示把 $\mathbf{P}^r$、$\mathbf{H}^r$ 化为低秩求和矩阵化后即 $\mathbf{P}\mathbf{I}-\mathbf{K}^\top\mathbf{W}$、$\mathbf{H}\mathbf{K}^\top\mathbf{U}$$\mathbf{W}\mathbf{TK}$、$\mathbf{U}\mathbf{TV}$其中 $\mathbf{T}$ 是块大小 $C$ 量级的下三角求逆最终算法 $\mathbf{S}\mathbf{S}^0\mathbf{K}^\top(\mathbf{U}-\mathbf{W}\mathbf{S}^0)$、$\mathbf{O}\mathbf{Q}\mathbf{S}^0(\mathbf{Q}\mathbf{K}^\top\odot\mathbf{M})(\mathbf{U}-\mathbf{W}\mathbf{S}^0)$ 全部由矩阵乘法构成。仓库中 prepare_wy_repr_fwdchunk_scaled_dot_kkt_fwd→solve_tril→recompute_w_u_fwd、chunk_delta_rule_fwd 与 chunk_delta_rule_bwd 正是这套公式的内核化落地naive.py 中的纯 PyTorch 分块实现则提供了最易读的同构验证。读者可据此在同一套符号下对照论文附录、README 推导与生产内核三处内容深入理解 DeltaNet 分块算法的每一行代码从何而来。【免费下载链接】flash-linear-attention Efficient implementations for emerging model architectures项目地址: https://gitcode.com/GitHub_Trending/fl/flash-linear-attention创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考