KDA(Kimi Delta Attention)的递推式一行就写完,但直接照着它写 kernel 需要同时处理四个机制:逐通道门控、delta rule 的三角求解、log 域 cumsum、跨 chunk 状态传递。四个机制耦合在一个 kernel 里,数值出错时无法判断误差来自哪一层。 本文采用递进式实现路径,每一级只引入一个新机制、每一级都可独立运行并做数值验证。本文覆盖第一级–移除全部衰减因子与删除因子,只保留 S ← S + K ⊤ V S \leftarrow S + K^\top V S ← S + K ⊤ V 。
这一级的价值不在性能,而在于将 chunkwise 分解恒等式单独隔离验证。本文同时给出一个反直觉的结论:线性注意力的全部优势是把 O ( N 2 D ) O(N^2 D) O ( N 2 D ) 换成 O ( N D 2 ) O(N D^2) O ( N D 2 ) ,但本文采用的逐块独立重算策略会让 FLOPs 退回 O ( N 2 ) O(N^2) O ( N 2 ) –与因果 FlashAttention 同量级。§6 会说明这个退化的根源在于把序列轴放进了 grid,并对照 TileLang 官方 chunk_delta_h 的做法:序列轴不进 grid,递推就退回单 block 内的顺序循环,N N N 的次数才能保住 。
数学背景见《KDA 的来龙去脉 》,TileLang 语言基础见《TileLang 编程基本知识点 》,本文复用的 tile 切分与 fragment 累加模式见《TileLang 实战:FlashAttention 前向 Kernel 》。
1. 递进路径:把 KDA 拆成可验证的增量
KDA 的完整递推式(D t = diag ( e g t ) D_t = \operatorname{diag}(e^{g_t}) D t = diag ( e g t ) ,g t ∈ R < 0 d k g_t \in \mathbb{R}^{d_k}_{<0} g t ∈ R < 0 d k ):
S t = S t − 1 D t ( I − β t k t k t ⊤ ) + β t v t k t ⊤ , o t = S t q t S_t = S_{t-1} D_t (I - \beta_t k_t k_t^\top) + \beta_t v_t k_t^\top, \qquad o_t = S_t q_t
S t = S t − 1 D t ( I − β t k t k t ⊤ ) + β t v t k t ⊤ , o t = S t q t
记号说明 :Kimi Linear 论文(arXiv:2510.26692)原文写作 S t = ( I − β t k t k t ⊤ ) Diag ( α t ) S t − 1 + β t k t v t ⊤ S_t = (I - \beta_t k_t k_t^\top)\operatorname{Diag}(\alpha_t) S_{t-1} + \beta_t k_t v_t^\top S t = ( I − β t k t k t ⊤ ) Diag ( α t ) S t − 1 + β t k t v t ⊤ ,α t ∈ [ 0 , 1 ] d k \alpha_t \in [0,1]^{d_k} α t ∈ [ 0 , 1 ] d k 。本文取其转置形式以匹配 kernel 中 S S S 的 ( d v , d k ) (d_v, d_k) ( d v , d k ) 内存布局,并沿用 GDN / Mamba 的 log 域记号 α t = e g t \alpha_t = e^{g_t} α t = e g t ,因为实现里门控本来就存在 log 域。
按机制拆解,每一级只放开一个自由度:
级别
递推式
新增机制
新出现的实现结构
第一级(本文)
S ← S + K ⊤ V S \leftarrow S + K^\top V S ← S + K ⊤ V
–
分块恒等式本身、块内因果掩码
第二级
S ← γ S + K ⊤ V S \leftarrow \gamma S + K^\top V S ← γ S + K ⊤ V
标量衰减
块内权重从 0/1 变 γ i − j \gamma^{i-j} γ i − j 指数下三角
第三级
S ← diag ( e g t ) ⋅ S + ⋯ S \leftarrow \operatorname{diag}(e^{g_t}) \cdot S + \cdots S ← diag ( e g t ) ⋅ S + ⋯
逐 token 门控
log 域 cumsum、exp2 硬件指令
第四级
S ← S ( I − β k k ⊤ ) + β v k ⊤ S \leftarrow S(I - \beta k k^\top) + \beta v k^\top S ← S ( I − β k k ⊤ ) + β v k ⊤
delta rule
UT 变换、三角求解(wy_fast 雏形)
第五级
门控 + delta rule
二者耦合
五阶段流水线、跨 kernel 状态传递
最后一级即完整 GDN / KDA,对标 flash-linear-attention 中的 chunkwise 实现。
这样拆分的收益是误差定位能力 :任何一级数值对不上,怀疑对象只有这一级新引入的那一个机制,前面几级已经验证过了。本文对应第一级,一个机制都还没引入,因此这里能验证的只有分块恒等式本身。
2. 数学推导:分块恒等式
本级递推不含遗忘项,状态单调累加:
S t = S t − 1 + k t v t ⊤ , o t = q t ⊤ S t S_t = S_{t-1} + k_t v_t^\top, \qquad o_t = q_t^\top S_t
S t = S t − 1 + k t v t ⊤ , o t = q t ⊤ S t
展开为显式求和,因果且包含当前 token:
o i = ∑ j ≤ i q i ⊤ ( k j v j ⊤ ) = ∑ j ≤ i ( q i ⋅ k j ) v j o_i = \sum_{j \le i} q_i^\top (k_j v_j^\top) = \sum_{j \le i} (q_i \cdot k_j)\, v_j
o i = j ≤ i ∑ q i ⊤ ( k j v j ⊤ ) = j ≤ i ∑ ( q i ⋅ k j ) v j
这是标准线性注意力。与 softmax attention 的唯一区别是分数未经 softmax,因此求和顺序可自由交换–这是下述分块重写成立的前提。
2.1 将求和拆分为跨块与块内两部分
序列按 B C BC B C 切分为 N C = N / B C NC = N / BC N C = N / B C 个 chunk。拆分的依据不是下标落在哪里,而是因果约束是否需要逐 token 判断 :
当前块之前的 chunk :其内每一个 j j j 对当前块里的每一个 i i i 都满足 j ≤ i j \le i j ≤ i ,全部可见。因果约束在这些块上恒真、与 i i i 无关,所以整块直接相乘就行。
当前块 :内部的 j j j 才真正受 j ≤ i j \le i j ≤ i 约束,只能取 k j , v j k_j, v_j k j , v j 中 j ≤ i j \le i j ≤ i 的那一半。
o i = q i ⊤ ( ∑ c ′ < c K c ′ ⊤ V c ′ ) ⏟ 跨块:全部可见,无需掩码 + ∑ j ∈ chunk c j ≤ i ( q i ⋅ k j ) v j ⏟ 块内:需因果掩码 o_i = \underbrace{q_i^\top \Big( \sum_{c' < c} K_{c'}^\top V_{c'} \Big)}_{\text{跨块:全部可见,无需掩码}} + \underbrace{\sum_{\substack{j \in \text{chunk } c \\ j \le i}} (q_i \cdot k_j) v_j}_{\text{块内:需因果掩码}}
o i = 跨块:全部可见,无需掩码 q i ⊤ ( c ′ < c ∑ K c ′ ⊤ V c ′ ) + 块内:需因果掩码 j ∈ chunk c j ≤ i ∑ ( q i ⋅ k j ) v j
关键在第一项:正因为它的因果判断与 i i i 无关,括号内的量对整个 chunk c c c 才能是同一个矩阵 ,记作 S c prev = ∑ c ′ < c K c ′ ⊤ V c ′ S_c^{\text{prev}} = \sum_{c' < c} K_{c'}^\top V_{c'} S c prev = ∑ c ′ < c K c ′ ⊤ V c ′ ,形状 d k × d v d_k \times d_v d k × d v ,含义是处理该 chunk 之前已累积的状态。于是跨块贡献退化为一次矩阵乘 Q c S c prev Q_c S_c^{\text{prev}} Q c S c prev ,块内贡献是一个 B C × B C BC \times BC B C × B C 的下三角掩码矩阵乘:
O c = Q c S c prev + tril ( Q c K c ⊤ ) V c O_c = Q_c S_c^{\text{prev}} + \operatorname{tril}(Q_c K_c^\top) V_c
O c = Q c S c prev + tril ( Q c K c ⊤ ) V c
图上半部分是前序 chunk 的注意力方阵——它自身也带因果三角,但这些计算在处理 chunk c c c 时已经完成,结果被汇总进一个 d k × d v d_k \times d_v d k × d v 的状态矩阵 S c prev S_c^{\text{prev}} S c prev ,不必再逐 token 展开。下半部分是本文要算的两项:左边整块打满阴影,表示 Q c Q_c Q c 的每一行都能无条件地乘上完整的历史状态,没有掩码;右边只有下三角有阴影,表示块内必须逐 token 判断 j ≤ i j \le i j ≤ i 。跨块项的历史长度随序列增长,但它被压进固定形状的 S c prev S_c^{\text{prev}} S c prev ——这正是线性复杂度的来源;块内三角形的边长恒为 B C BC B C ,与序列长度无关。
这条恒等式是后续各级的公共基础 ,之后每一级都只是在它的两项上分别插入衰减权重,张量形状与 GEMM 调用次序保持不变。
2.2 掩码可以直接置零的原因
FlashAttention 中掩码必须填 − ∞ -\infty − ∞ ,因为后续要经过 exp,e − ∞ = 0 e^{-\infty} = 0 e − ∞ = 0 才能使被掩位置不产生贡献。线性注意力没有 softmax,掩码位置直接写 0 即可:
1 2 for i, j in T.Parallel(BC, BC): A[i, j] = T.if_then_else(j <= i, A[i, j], 0.0 )
这一行是线性注意力与 FA 的第一处分岔,也是「无 softmax」在代码上的全部体现–不需要 online 重标定、不需要 m m m 与 ℓ \ell ℓ 两个行状态、不需要输出阶段的最终除法 。
3. 四层 PyTorch 参考实现
验证 kernel 之前需要建立可信参考。本文构造三份 PyTorch fp64 实现,从「最直白」逐步过渡到「与 kernel 控制流同构」,相邻两层互相验证:
参考
实现方式
验证目标
A
逐 token 递归,双重 for 循环
递推式定义本身,几乎不可能写错
B
分块向量化,cumsum 求前缀状态
§2.1 分块恒等式的正确性
C
逐块独立重算,模拟 grid 与 T.Pipelined
kernel 的控制流与访存次序
参考 C 是连接 PyTorch 与 TileLang 的关键一层–它的循环结构与 kernel 逐句对应 ,若 kernel 与 C 不符,问题必在 TileLang 语法或内存层级使用上,与数学无关。
3.1 参考 A:逐 token 递归
1 2 3 4 5 6 7 8 9 10 11 def ref_recurrent (Q, K, V ): """最直白的实现:严格照抄递推式 S += k v^T; o = q S""" B, N, H, D = Q.shape O = torch.empty(B, N, H, D, dtype=torch.float64, device=Q.device) for b in range (B): for h in range (H): S = torch.zeros(D, D, dtype=torch.float64, device=Q.device) for t in range (N): S += torch.outer(K[b, t, h].double(), V[b, t, h].double()) O[b, t, h] = Q[b, t, h].double() @ S return O
3.2 参考 B:分块向量化
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 def ref_chunked (Q, K, V, BC ): """分块向量化:用 cumsum 一次算出所有 chunk 的前缀状态""" B, N, H, D = Q.shape NC = N // BC Qc = Q.permute(0 , 2 , 1 , 3 ).double().reshape(B, H, NC, BC, D) Kc = K.permute(0 , 2 , 1 , 3 ).double().reshape(B, H, NC, BC, D) Vc = V.permute(0 , 2 , 1 , 3 ).double().reshape(B, H, NC, BC, D) outer = torch.einsum("bhcnd,bhcnv->bhcdv" , Kc, Vc) states = torch.cumsum(outer, dim=2 ) states_prev = torch.cat( [torch.zeros_like(states[:, :, :1 ]), states[:, :, :-1 ]], dim=2 ) O_inter = torch.einsum("bhcnd,bhcdv->bhcnv" , Qc, states_prev) A = torch.einsum("bhcnd,bhcmd->bhcnm" , Qc, Kc) mask = torch.tril(torch.ones(BC, BC, dtype=torch.bool , device=Q.device)) A = A.masked_fill(~mask, 0.0 ) O_intra = torch.einsum("bhcnm,bhcmv->bhcnv" , A, Vc) return (O_inter + O_intra).reshape(B, H, N, D).permute(0 , 2 , 1 , 3 ).contiguous()
states_prev 的右移拼接对应 ∑ c ′ < c \sum_{c' < c} ∑ c ′ < c 中的严格小于号:cumsum 给出的是含自身的前缀和,右移一格补零才是处理该 chunk 之前的累积状态。这是分块线性注意力最易出错的一行 –不右移等价于把当前块的 K ⊤ V K^\top V K ⊤ V 重复计入,块内贡献会被计算两次。该错误的量级实测为相对 L2 误差 7.78 × 10 − 1 7.78 \times 10^{-1} 7.78 × 1 0 − 1 ,属于结构性错误而非精度问题,容易识别。
3.3 参考 C:kernel 控制流镜像
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 def ref_kernel_mimic (Q, K, V, BC ): """逐块独立重算:外层三重循环模拟 grid,内层 range(bx) 模拟 T.Pipelined(bx)""" B, N, H, D = Q.shape NC = N // BC O = torch.zeros(B, N, H, D, dtype=torch.float64, device=Q.device) mask = torch.tril(torch.ones(BC, BC, dtype=torch.float64, device=Q.device)) for bz in range (B): for by in range (H): for bx in range (NC): sl = slice (bx * BC, (bx + 1 ) * BC) S = torch.zeros(D, D, dtype=torch.float64, device=Q.device) for c in range (bx): cs = slice (c * BC, (c + 1 ) * BC) S += K[bz, cs, by, :].double().T @ V[bz, cs, by, :].double() Qb = Q[bz, sl, by, :].double() acc = Qb @ S Kb = K[bz, sl, by, :].double() Vb = V[bz, sl, by, :].double() acc += ((Qb @ Kb.T) * mask) @ Vb O[bz, sl, by, :] = acc return O
对照参考 B 与参考 C 可以看出两种前缀状态求法的区别:B 用 cumsum 一次性算出全部 N C NC N C 个前缀状态、总代价 O ( N C ) O(NC) O ( N C ) ;C 的每个 b x bx b x 独立重算、总代价 O ( N C 2 ) O(NC^2) O ( N C 2 ) 。kernel 采用的是 C 的策略 ,原因与代价见 §6。
3.4 四层参考的一致性验证
B = 2 , H = 2 , N = 12 , D = 4 , B C = 4 B=2, H=2, N=12, D=4, BC=4 B = 2 , H = 2 , N = 12 , D = 4 , B C = 4 ,fp64(numpy 复现):
比较
max abs 误差
相对 L2
B 分块向量化 vs A 逐 token 递归
3.55 × 10 − 15 3.55 \times 10^{-15} 3.55 × 1 0 − 15
1.62 × 10 − 16 1.62 \times 10^{-16} 1.62 × 1 0 − 16
C kernel 结构镜像 vs A 逐 token 递归
7.11 × 10 − 15 7.11 \times 10^{-15} 7.11 × 1 0 − 15
1.86 × 10 − 16 1.86 \times 10^{-16} 1.86 × 1 0 − 16
C vs B
4.44 × 10 − 15 4.44 \times 10^{-15} 4.44 × 1 0 − 15
1.78 × 10 − 16 1.78 \times 10^{-16} 1.78 × 1 0 − 16
参考 B 去掉右移(错误实现)
1.97 × 10 1 1.97 \times 10^{1} 1.97 × 1 0 1
7.78 × 10 − 1 7.78 \times 10^{-1} 7.78 × 1 0 − 1
前三行误差均在 fp64 机器精度量级(ε ≈ 2.2 × 10 − 16 \varepsilon \approx 2.2 \times 10^{-16} ε ≈ 2.2 × 1 0 − 16 ),恒等式与三份实现均无误。第四行是刻意引入的错误,用于确认该验证流程对结构性错误敏感。
参考 D(§6.4.2 的序列内循环形式)另外验证一件事–换 grid 划分、切 DV 后算的还是同一个恒等式 。B = 2 , S = 256 , H = 3 , D K = 64 , D V = 128 , C = 64 B{=}2, S{=}256, H{=}3, DK{=}64, DV{=}128, C{=}64 B = 2 , S = 256 , H = 3 , D K = 64 , D V = 128 , C = 64 ,fp64:
对象
相对 L2(vs 参考 A)
参考 D,block D V = 32 \text{block}_{DV} = 32 block D V = 32
5.60 × 10 − 16 5.60 \times 10^{-16} 5.60 × 1 0 − 16
参考 D,block D V = 64 \text{block}_{DV} = 64 block D V = 64
5.60 × 10 − 16 5.60 \times 10^{-16} 5.60 × 1 0 − 16
参考 D,block D V = 128 \text{block}_{DV} = 128 block D V = 128 (不切)
5.60 × 10 − 16 5.60 \times 10^{-16} 5.60 × 1 0 − 16
参考 D,状态更新提到写回之前
6.98 × 10 − 1 6.98 \times 10^{-1} 6.98 × 1 0 − 1
三个 block D V \text{block}_{DV} block D V 误差完全相同,这就是「DV 切块零依赖」的直接证据 –切不切、切多细,算出来的是比特级相同的东西,因为每个 dv 竖条的浮点累加顺序本就不受其他竖条影响。对照最后一行:仅仅把状态更新从循环尾部提到开头,误差就跳到 70%–右移语义是硬要求。
以上均为 numpy fp64 实测。TileLang kernel 本身需要 CUDA 设备,本文未给出实测数字–待真卡跑通后单独补充 ,此处不做性能推测。预期的主导误差源是 T.copy(S_f, S_s) 把 f32 状态降至 f16 落 shared(Tensor Core MMA 的输入必须是低精度);状态 S S S 是多个块累加的结果,越靠后的 block 累加项越多,误差随 b x bx b x 单调增长,影响远大于块内 A A A 的那次降精度。
3.5 einsum 输出下标的约束
参考 B 中若把 torch.einsum("bhcnd,bhcnv->bhcdv", ...) 误写为 ->bhcdd(意图表达「输出是 d × d d \times d d × d 方阵」),会直接抛出异常:
1 ValueError: einstein sum subscripts string includes output subscript 'd' multiple times
einsum 的输出下标不允许重复–重复下标在输出侧的语义是取对角线。而 K ⊤ V K^\top V K ⊤ V 的两个维度虽然长度均为 D D D ,语义上分别是 key 维与 value 维 ,必须使用不同字母 d 与 v。同理 bhcnd,bhcdd->bhcnv 也不合法,因为输出的 v 从未在输入中出现。d k = d v d_k = d_v d k = d v 时长度相同掩盖了语义差异,一旦 KDA 中 d k ≠ d v d_k \ne d_v d k = d v ,该疏忽会立即表现为形状错误。
4. TileLang kernel:三步实现
grid 划分与 FA 一致–按 Q 块切分所有权,每个 block 负责输出一个 B C × d BC \times d B C × d 的 tile:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 @tilelang.jit(out_idx=[3 ] ) def linattn_chunk (batch, heads, seq_len, dim, blk, num_stages=2 , dtype=T.float16, accum_dtype=T.float32 ): BC = blk @T.prim_func def main ( Q: T.Tensor([batch, seq_len, heads, dim], dtype ), K: T.Tensor([batch, seq_len, heads, dim], dtype ), V: T.Tensor([batch, seq_len, heads, dim], dtype ), O: T.Tensor([batch, seq_len, heads, dim], dtype ), ): with T.Kernel(T.ceildiv(seq_len, BC), heads, batch, threads=128 ) as (bx, by, bz): Q_s = T.alloc_shared([BC, dim], dtype) K_s = T.alloc_shared([BC, dim], dtype) V_s = T.alloc_shared([BC, dim], dtype) S_s = T.alloc_shared([dim, dim], dtype) O_s = T.alloc_shared([BC, dim], dtype) S_f = T.alloc_fragment([dim, dim], accum_dtype) acc_o = T.alloc_fragment([BC, dim], accum_dtype) A = T.alloc_fragment([BC, BC], accum_dtype) A_cast = T.alloc_fragment([BC, BC], dtype)
4.1 步骤①:流式累加前缀状态
对应参考 C 中的 for c in range(bx):
1 2 3 4 5 6 7 T.clear(S_f) for c in T.Pipelined(bx, num_stages=num_stages): T.copy(K[bz, c * BC:(c + 1 ) * BC, by, :], K_s) T.copy(V[bz, c * BC:(c + 1 ) * BC, by, :], V_s) T.gemm(K_s, V_s, S_f, transpose_A=True ) T.copy(S_f, S_s)
代码里看不到 +=,但这个循环确实是在累加:T.clear(S_f) 先把 fragment 清零,之后每次 T.gemm 都在 S_f 原地累加,因此循环等价于 S f = ∑ c < b x K c ⊤ V c S_f = \sum_{c < bx} K_c^\top V_c S f = ∑ c < b x K c ⊤ V c 。累加语义来自 Tensor Core MMA 的基本形式–MMA 指令算的是 d = a @ b + c,累加器既是输入也是输出,T.gemm 默认沿用这一行为;若要覆盖而非累加,需显式传 clear_accum=True。
T.Pipelined(bx) 的上界是 block 索引而非编译期常量–不同 block 的循环次数不同 ,第 0 块一次都不执行(S = 0 S = 0 S = 0 ),最后一块需执行 N C − 1 NC-1 N C − 1 次。TileLang 支持 runtime 上界的流水线,代价是各 block 负载严重不均衡,尾部 block 构成整个 kernel 的关键路径。
这里每个 block 都在重算自己需要的前缀状态,而递推式本身只需一次加法。为什么本文要这么写、以及官方实现如何避开它,见 §6.2 与 §6.4。
4.2 步骤②③:两项贡献求和
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 T.copy(Q[bz, bx * BC:(bx + 1 ) * BC, by, :], Q_s) T.clear(acc_o) T.gemm(Q_s, S_s, acc_o) T.copy(K[bz, bx * BC:(bx + 1 ) * BC, by, :], K_s) T.copy(V[bz, bx * BC:(bx + 1 ) * BC, by, :], V_s) T.clear(A) T.gemm(Q_s, K_s, A, transpose_B=True ) for i, j in T.Parallel(BC, BC): A[i, j] = T.if_then_else(j <= i, A[i, j], 0.0 ) T.copy(A, A_cast) T.gemm(A_cast, V_s, acc_o) T.copy(acc_o, O_s) T.copy(O_s, O[bz, bx * BC:(bx + 1 ) * BC, by, :])
步骤③开头必须重新 T.copy 当前块的 K/V。步骤①的流水线循环结束时,K_s / V_s 中残留的是第 b x − 1 bx-1 b x − 1 块的数据,且在 num_stages > 1 时具体残留哪一块取决于流水线的展开方式,不可假设 。复用 shared buffer 节省了显存,但必须显式重载–这是 tile 编程中典型的隐式状态陷阱。参考 C 里没有这个问题,因为 Python 每轮都重新切片,不存在缓冲区复用。
T.copy(A, A_cast) 这次 f32→f16 转换与 FA 中 P ~ \tilde{P} P ~ 喂入第二个 GEMM 前的降精度是同一操作–Tensor Core 的 MMA 输入必须是低精度,仅累加器为 f32。
4.3 kernel 与参考 C 的逐句对应
参考 C(PyTorch)
kernel(TileLang)
说明
for bz / by / bx 三重循环
T.Kernel(ceildiv(N,BC), heads, batch)
循环变并行 grid
S = torch.zeros(D, D)
T.clear(S_f)
状态初始化,fragment 常驻寄存器
for c in range(bx)
T.Pipelined(bx, num_stages)
顺序循环变软件流水线
S += K[cs].T @ V[cs]
T.gemm(K_s, V_s, S_f, transpose_A=True)
显式 shared 暂存 + Tensor Core
acc = Qb @ S
T.gemm(Q_s, S_s, acc_o)
需先 T.copy(S_f, S_s) 降精度
(Qb @ Kb.T) * mask
T.gemm(..., transpose_B=True) + T.Parallel 掩码
掩码从广播乘变逐元素条件
acc += (...) @ Vb
T.copy(A, A_cast) + T.gemm(A_cast, V_s, acc_o)
多一次 f32→f16 转换
O[bz, sl, by, :] = acc
T.copy(acc_o, O_s) + T.copy(O_s, O[...])
经 shared 中转写回 HBM
两处 PyTorch 中不存在的操作:f32→f16 显式降精度 (Tensor Core 输入约束)与shared buffer 中转 (内存层级手动管理)。这两项也正是 kernel 与参考 C 数值差异的全部来源。
5. 与 FlashAttention 的三点结构差异
本文实现复用了 FA 的全部 tile 模式,但有三处必须修改:
维度
FlashAttention
线性注意力
原因
归一化
online softmax,维护 m m m 、ℓ \ell ℓ 两个行状态,每块重标定
无,掩码直接置 0
无 exp,求和顺序可交换
循环携带的量
无跨块状态,每个 Q 块独立
S [ d k , d v ] S[d_k, d_v] S [ d k , d v ] 是跨迭代累加的状态
递推式本身带状态
T.gemm policy
必须 FullRow–行归约要求整行在同一 warp
默认 Square 即可
无按行归约
第三点值得展开:FA 需对分数块做 rowmax / rowsum,若一行被切分到多个 warp,归约就需跨 warp 通信,因此必须用 FullRow policy 强制整行不拆。线性注意力的块内矩阵 A A A 计算完成后仅做逐元素掩码,无任何跨列归约,warp 划分方式不影响正确性 ,编译器可自由选择寄存器分布最优方案。
第二点需要说明的是,本文实现并没有真正兑现这一行:逐块独立重算把状态依赖藏起来了,每个 block 都从零重算自己需要的前缀状态,块之间不传递任何东西。这不是偷懒,而是本文这种「序列轴占据 grid.x」的划分下,block 之间无法顺序传递状态。换一种 grid 划分就没有这个限制,见 §6.4。
6. 代价账本:线性复杂度从何而来,如何被丢掉,以及如何拿回
6.1 两条路线的规模差异
先把线性注意力的本质说清楚。同样的输出,有两条算法:
路线
计算方式
规模
softmax / FA
先算分数矩阵 Q K ⊤ QK^\top Q K ⊤ (N × N N \times N N × N ),再乘 V V V
O ( N 2 D ) O(N^2 D) O ( N 2 D )
线性注意力
先算状态 S = K ⊤ V S = K^\top V S = K ⊤ V (D × D D \times D D × D ),再乘 Q Q Q
O ( N D 2 ) O(N D^2) O ( N D 2 )
区别在结合律往哪边括:( Q K ⊤ ) V (QK^\top)V ( Q K ⊤ ) V 要物化一个 N × N N \times N N × N 的中间矩阵,Q ( K ⊤ V ) Q(K^\top V) Q ( K ⊤ V ) 物化的是 D × D D \times D D × D 。N N N 的次数从 2 降到 1,D D D 的次数从 1 升到 2 –这才是线性注意力唯一的、也是全部的优势来源,代价是 O ( D 2 ) O(D^2) O ( D 2 ) 的固定状态取代了 O ( N 2 ) O(N^2) O ( N 2 ) 的自由分数矩阵,表达能力随之受限。
chunkwise 形式是这两条路线的混合:跨块走 S S S 路线(每块一次 Q c S c prev Q_c S_c^{\text{prev}} Q c S c prev ,共 N ⋅ 2 D 2 N \cdot 2D^2 N ⋅ 2 D 2 ),块内走 Q K ⊤ QK^\top Q K ⊤ 路线(每块一个 B C × B C BC \times BC B C × B C 分数矩阵,共 N ⋅ 4 B C D N \cdot 4 BC D N ⋅ 4 B C D ),状态更新本身再花 N ⋅ 2 D 2 N \cdot 2D^2 N ⋅ 2 D 2 。理想总量:
FLOPs ideal = 4 N D 2 + 4 N B C D = 4 N D ( D + B C ) \text{FLOPs}_{\text{ideal}} = 4N D^2 + 4N\,BC\,D = 4ND(D + BC)
FLOPs ideal = 4 N D 2 + 4 N B C D = 4 N D ( D + B C )
N N N 是一次方 。与因果 FA 的 2 N 2 D 2N^2 D 2 N 2 D 相比:
FLOPs ideal FLOPs causal FA = 2 ( D + B C ) N \frac{\text{FLOPs}_{\text{ideal}}}{\text{FLOPs}_{\text{causal FA}}} = \frac{2(D + BC)}{N}
FLOPs causal FA FLOPs ideal = N 2 ( D + B C )
比值按 1 / N 1/N 1/ N 衰减–序列越长优势越大,这是线性注意力值得做的全部理由。
6.2 本文为何没有顺着递推式只加一次
理想账本里的 N C NC N C 次 K ⊤ V K^\top V K ⊤ V ,对应的是顺序递推:
S c = S c − 1 + K c − 1 ⊤ V c − 1 S_{c} = S_{c-1} + K_{c-1}^\top V_{c-1}
S c = S c − 1 + K c − 1 ⊤ V c − 1
每个 chunk 只做一次 GEMM,读上一步的结果、加上自己这一块。参考 B 就是这么算的(cumsum 一次扫完),CPU 上顺序执行毫无问题。
但 kernel 做不到这一点,原因在 grid 的划分方式。 §4 把所有权按 Q 块切分,b x bx b x 是 grid.x 的索引:所有 N C NC N C 个 block 由硬件并行调度,执行顺序不确定、彼此之间没有同步点、也没有共享的可写缓冲区 。block b x bx b x 若想读 S b x S_{bx} S b x ,就得等 block b x − 1 bx-1 b x − 1 算完并把结果落到某处——这两件事在单个 kernel 内都不成立:CUDA 不保证 block 间的执行次序(b x − 1 bx-1 b x − 1 可能还没启动),也没有跨 block 的 barrier 可用。
顺序递推要求的是"前一步已完成",而 grid 提供的是"所有步同时开始"。但这个矛盾是本文自己造出来的 –它成立的前提是"序列轴必须占据 grid 的一维"。放弃这个前提,矛盾就不存在了,见 §6.4。
本文仍保留逐块独立重算:它让 kernel 保持单文件闭环、控制流与参考 C 严格对应,适合把分块恒等式本身隔离出来验证。下面先算清这个选择的代价。
6.3 冗余重算把 N N N 的次数还了回去
第 b x bx b x 块执行 b x bx b x 次 K ⊤ V K^\top V K ⊤ V ,全部 block 合计 ∑ c = 0 N C − 1 c = N C ( N C − 1 ) / 2 \sum_{c=0}^{NC-1} c = NC(NC-1)/2 ∑ c = 0 N C − 1 c = N C ( N C − 1 ) /2 次,而顺序递推共 N C NC N C 次。冗余系数 ( N C − 1 ) / 2 (NC-1)/2 ( N C − 1 ) /2 ,且随 N N N 线性增长 –正是这个增长把 N N N 的次数从 1 顶回 2:
FLOPs ① = N C ( N C − 1 ) 2 ⋅ 2 B C D 2 ≈ N 2 D 2 B C \text{FLOPs}_① = \frac{NC(NC-1)}{2} \cdot 2\, BC\, D^2 \approx \frac{N^2 D^2}{BC}
FLOPs ① = 2 N C ( N C − 1 ) ⋅ 2 B C D 2 ≈ B C N 2 D 2
于是与因果 FA 的比值退化成常数:
FLOPs ① FLOPs causal FA = D 2 B C \frac{\text{FLOPs}_①}{\text{FLOPs}_{\text{causal FA}}} = \frac{D}{2\,BC}
FLOPs causal FA FLOPs ① = 2 B C D
这个比值与 N N N 无关,恰恰是失败的判据 :它说明本文实现与 FA 同属 O ( N 2 ) O(N^2) O ( N 2 ) ,1 / N 1/N 1/ N 的衰减优势被完全抹掉了。实测账本(D = 64 D = 64 D = 64 ,B C = 64 BC = 64 B C = 64 ):
N N N
N C NC N C
冗余系数
本文实现(步骤①)
理想线性注意力
因果 FA
本文/FA
理想/FA
512
8
3.5x
0.015 GFLOP
0.017 GFLOP
0.034 GFLOP
0.438
0.500
2048
32
15.5x
0.260 GFLOP
0.067 GFLOP
0.537 GFLOP
0.484
0.125
8192
128
63.5x
4.261 GFLOP
0.268 GFLOP
8.590 GFLOP
0.496
0.031
16384
256
127.5x
17.113 GFLOP
0.537 GFLOP
34.360 GFLOP
0.498
0.016
65536
1024
511.5x
274.609 GFLOP
2.147 GFLOP
549.756 GFLOP
0.500
0.004
看最后两列的走向:本文/FA 收敛到常数 D / ( 2 B C ) = 0.5 D/(2BC) = 0.5 D / ( 2 B C ) = 0.5 ,理想/FA 按 1 / N 1/N 1/ N 一路衰减到 0.004。 N = 65536 N = 65536 N = 65536 时理想实现只需 2.1 GFLOP,本文实现要 274.6 GFLOP–差 128 倍,而这个倍数还会随 N N N 继续涨。
B C BC B C 的取舍随之明确:
B C BC B C
D / ( 2 B C ) D/(2BC) D / ( 2 B C )
含义
32
1.000
与因果 FA 计算量持平,块过小无收益
64
0.500
默认值,寄存器压力可控
128
0.250
冗余减半,但 A A A 占 128 × 128 128 \times 128 128 × 128 f32 fragment,每线程约 128 个寄存器,易溢出
256
0.125
理论最省,实际 shared memory 与寄存器均无法容纳
但要注意这张表只是在常数上 打折,O ( N 2 ) O(N^2) O ( N 2 ) 的量级不变。增大 B C BC B C 能线性减少冗余,但 A A A 是 B C × B C BC \times BC B C × B C 的 f32 fragment,B C = 128 BC = 128 B C = 128 时寄存器压力已接近溢出边界。真正的解法不是调 B C BC B C ,见下一节。
6.4 更好的方案:把序列轴从 grid 里拿掉
TileLang 官方 examples/gdn 中的 chunk_delta_h 给出了另一种划分。它的 grid 只有两维:
1 2 3 4 5 6 7 8 with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh): ... for i_s in T.Pipelined(T.ceildiv(S, block_S), num_stages=num_stages): T.copy(b_h_shared, h[bb, i_s, bh, 0 :DK, bv * block_DV:(bv + 1 ) * block_DV]) ... T.gemm(K_shared, V_new_shared, b_h_fragment, transpose_A=True ) T.copy(b_h_fragment, b_h_shared)
序列不在 grid 里,而是 kernel 内部的一个顺序循环。 b_h_fragment 成为 loop-carried 的寄存器变量,跨 chunk 一路累加–这正是 S c = γ c S c − 1 + K c ⊤ V c S_c = \gamma_c S_{c-1} + K_c^\top V_c S c = γ c S c − 1 + K c ⊤ V c 的直接翻译,每个 chunk 只做一次 K ⊤ V K^\top V K ⊤ V ,一次都不重算。§6.2 里那个"block 间无法同步"的矛盾根本不会出现,因为递推的顺序性被限制在单个 block 内部,用寄存器解决了,从来不需要跨 block 通信 。
并行度从另外两个轴补回来:
轴
来源
数量(B = 1 , H = 32 , D V = 128 , b l o c k D V = 32 B=1,H=32,DV=128,block_{DV}=32 B = 1 , H = 32 , D V = 128 , b l oc k D V = 32 )
bv
DV 切块
128 / 32 = 4 128/32 = 4 128/32 = 4
bbh
batch × \times × head 融合
1 × 32 = 32 1 \times 32 = 32 1 × 32 = 32
合计 128 个 block,足够填满 SM。
6.4.1 为什么只切 DV、不切 DK
这一点容易给出错误的理由。若只看状态更新那一个 GEMM:
S c + = K c ⊤ V c , K c ⊤ : [ D K , C ] , V c : [ C , D V ] S_c \mathrel{+}= K_c^\top V_c,\qquad K_c^\top: [DK, C],\ V_c: [C, DV]
S c + = K c ⊤ V c , K c ⊤ : [ D K , C ] , V c : [ C , D V ]
它的收缩维是 chunk 长度 C C C ,D K DK D K 是 M 维、D V DV D V 是 N 维。而切 M 维只是输出 tiling,各 block 写各自的行,同样零依赖 –所以「D K DK D K 是 M 维、切开要跨 block 归约」是站不住的,单看这一步切哪边都行。
不对称性来自下游谁消费这个状态 。把一个 chunk 里三个 GEMM 的收缩维列出来:
GEMM
形状
收缩维
D K DK D K 的角色
D V DV D V 的角色
状态更新 K c ⊤ V c K_c^\top V_c K c ⊤ V c
[ D K , C ] × [ C , D V ] [DK,C] \times [C,DV] [ D K , C ] × [ C , D V ]
C C C
M(自由)
N(自由)
跨块读出 Q c S c prev Q_c S_c^{\text{prev}} Q c S c prev
[ C , D K ] × [ D K , D V ] [C,DK] \times [DK,DV] [ C , D K ] × [ D K , D V ]
D K DK D K
收缩
N(自由)
块内项 Q c K c ⊤ Q_c K_c^\top Q c K c ⊤
[ C , D K ] × [ D K , C ] [C,DK] \times [DK,C] [ C , D K ] × [ D K , C ]
D K DK D K
收缩
不出现
结论一句话:D V DV D V 在三个 GEMM 里始终是自由维,D K DK D K 在其中两个里是收缩维。
把两种切法的张量尺寸画出来,区别就是一个词:拼接还是相加 。
切 DV :block ( b v , b b h ) (bv, bbh) ( b v , bbh ) 独占 S [ : , dv ] S[:, \text{dv}] S [ : , dv ] 这一竖条,自己跑完整条递推,输出 O [ : , dv ] = Q c S prev [ : , dv ] + tril ( Q c K c ⊤ ) V c [ : , dv ] O[:, \text{dv}] = Q_c S^{\text{prev}}[:, \text{dv}] + \operatorname{tril}(Q_cK_c^\top)V_c[:, \text{dv}] O [ : , dv ] = Q c S prev [ : , dv ] + tril ( Q c K c ⊤ ) V c [ : , dv ] 是 C × D V / n C \times DV/n C × D V / n 的条带,已是最终值 ,直接写回。状态和输出两步都只有拼接,零累加、零跨块通信。
切 DK :状态更新那步确实也没事(S i = K i ⊤ V S^i = K_i^\top V S i = K i ⊤ V 是 D K / n × D V DK/n \times DV D K / n × D V 的行条带,也不重叠),但读出就崩了 :Q Q Q 被迫跟着切成 C × D K / n C \times DK/n C × D K / n ,算出的 o ~ i = Q i S i \tilde{o}^i = Q^i S^i o ~ i = Q i S i 是 C × D V C \times DV C × D V 全宽 –n n n 份 o ~ i \tilde{o}^i o ~ i 盖在同一块 O O O 上,各含 1/n 根收缩轴的贡献,O c = ∑ i o ~ i O_c = \sum_i \tilde{o}^i O c = ∑ i o ~ i ,少加任何一份结果就错。这就是 split-K:每个 chunk 每条序列都要归约一次 [ C , D V ] [C, DV] [ C , D V ] 的中间结果,摊销不掉,只能选 fp32 atomic add(非确定性 + 带宽)或另开 workspace 起第二个 kernel。
更麻烦的是块内项:Q c K c ⊤ Q_cK_c^\top Q c K c ⊤ 也沿 D K DK D K 收缩,只持有 dk \text{dk} dk 子块的 block 根本算不出完整的 [ C , C ] [C,C] [ C , C ] 转移矩阵(到了 DeltaNet 那级还要求 ( I + tril ( diag ( β ) K K ⊤ ) ) − 1 (I + \operatorname{tril}(\operatorname{diag}(\beta)K K^\top))^{-1} ( I + tril ( diag ( β ) K K ⊤ ) ) − 1 ,逻辑上无法切)。
图里“Q Q Q 也要切”那一行值得单独指出:DK 是吸收维,一旦切它,所有沿 DK 寻址的张量(Q Q Q 、K K K 、S S S 的行)都被动跟随,而产出反而变成全宽部分和 –输入维度变小、输出维度不变,这个尺寸上的不匹配就是归约的必然信号。切 DV 恰好相反:输入切窄一根,输出也跟着窄一根。
还有一个纯工程的理由:状态 fragment 是 [ D K , block D V ] [DK, \text{block}_{DV}] [ D K , block D V ] 的 f32,D K = D V = 128 DK{=}DV{=}128 D K = D V = 128 、128 线程时,不切 DV 要 128 × 128 / 128 = 128 128\times128/128 = 128 128 × 128/128 = 128 个寄存器/线程,直接溢出;切成 block D V = 32 \text{block}_{DV}{=}32 block D V = 32 降到 32 个。DV 切块同时解决了 SM 占用率和寄存器压力两件事,且不引入任何归约。
6.4.2 完整例子:按 ( b , h , d v ) (b, h, dv) ( b , h , d v ) 切 tile,沿 S S S 走 Pipelined
把上面的结论落成第一级(无衰减、无删除)的可运行形态。与 §4 相比只改一件事:grid 里的序列轴换成 DV 轴,序列退回 kernel 内部的顺序循环 。
所有权划分:block ( b v , b b h ) (bv, bbh) ( b v , bbh ) 负责 O [ b , : , h , dv 竖条 ] O[b, :, h, \ \text{dv 竖条}] O [ b , : , h , dv 竖条 ] –整条序列 的一个 value 通道子集。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 @tilelang.jit(out_idx=[3 ] ) def linattn_seq_in_loop (B, H, S, DK, DV, block_S=64 , block_DV=32 , num_stages=2 , threads=128 , dtype=T.float16, accum_dtype=T.float32 ): C = block_S NS = T.ceildiv(S, C) @T.prim_func def main ( Q: T.Tensor([B, S, H, DK], dtype ), K: T.Tensor([B, S, H, DK], dtype ), V: T.Tensor([B, S, H, DV], dtype ), O: T.Tensor([B, S, H, DV], dtype ), ): with T.Kernel(T.ceildiv(DV, block_DV), B * H, threads=threads) as (bv, bbh): bb = bbh // H bh = bbh % H dv0 = bv * block_DV Q_s = T.alloc_shared([C, DK], dtype) K_s = T.alloc_shared([C, DK], dtype) V_s = T.alloc_shared([C, block_DV], dtype) S_s = T.alloc_shared([DK, block_DV], dtype) O_s = T.alloc_shared([C, block_DV], dtype) S_f = T.alloc_fragment([DK, block_DV], accum_dtype) acc_o = T.alloc_fragment([C, block_DV], accum_dtype) A = T.alloc_fragment([C, C], accum_dtype) A_cast = T.alloc_fragment([C, C], dtype) T.clear(S_f) for i_s in T.Pipelined(NS, num_stages=num_stages): s0 = i_s * C T.copy(Q[bb, s0:s0 + C, bh, :], Q_s) T.copy(K[bb, s0:s0 + C, bh, :], K_s) T.copy(V[bb, s0:s0 + C, bh, dv0:dv0 + block_DV], V_s) T.copy(S_f, S_s) T.gemm(Q_s, S_s, acc_o, clear_accum=True ) T.gemm(Q_s, K_s, A, transpose_B=True , clear_accum=True ) for i, j in T.Parallel(C, C): A[i, j] = T.if_then_else(j <= i, A[i, j], 0.0 ) T.copy(A, A_cast) T.gemm(A_cast, V_s, acc_o) T.copy(acc_o, O_s) T.copy(O_s, O[bb, s0:s0 + C, bh, dv0:dv0 + block_DV]) T.gemm(K_s, V_s, S_f, transpose_A=True ) return main
四处必须讲清的细节:
循环体内的顺序就是右移语义。 T.gemm(K_s, V_s, S_f, ...) 必须排在输出写回之后:S f S_f S f 在读出时代表 S c prev = ∑ c ′ < c K c ′ ⊤ V c ′ S^{\text{prev}}_c = \sum_{c' < c} K_{c'}^\top V_{c'} S c prev = ∑ c ′ < c K c ′ ⊤ V c ′ ,本 chunk 自己的贡献由 tril \operatorname{tril} tril 那一项负责。把状态更新提到前面,就变成了 inclusive 前缀,块内项会被重复计入–这是本级唯一的结构性错误点,且数值上表现为「整体偏大」而非 NaN,很容易漏掉。§3.4 里刻意去掉右移的那个错误实现,对应的就是这里的顺序写反。
T.clear(S_f) 在循环外,clear_accum=True 在循环内。 状态要跨迭代累加,所以只能循环外清零一次;而 acc_o 和 A 每个 chunk 都是全新的,进了流水线循环就不能再用 T.clear(清零会被排到流水线的错误阶段),必须靠 clear_accum=True 在 MMA 那一刻覆盖累加器。
num_stages 只能盖住访存,盖不住计算。 S_f -> S_s -> T.gemm -> S_f 构成一条真正的循环依赖,编译器无法把相邻两个 chunk 的计算重叠;流水线的收益全部来自把下一个 chunk 的 Q/K/V 的 HBM→shared 搬运提前发出。所以 num_stages=2 基本够用,继续加只是多占 shared。
Q/K 被 dv 方向重复读。 每个 b v bv b v 都要读整份 Q Q Q 、K K K (各 D K DK D K 全宽),读放大系数 D V / block D V DV/\text{block}_{DV} D V / block D V ;V V V 和 O O O 则严格切分不重复。这是切 DV 唯一的代价–多读,但不归约 。相比之下切 DK 是要归约 ,这就是取舍的本质区别。
对应的 PyTorch 参考(延续 §3 的写法,可直接跟参考 A/B 对齐):
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 def ref_D_seq_in_loop (Q, K, V, C, block_DV ): B, S, H, DK = Q.shape DV = V.shape[-1 ] O = torch.zeros(B, S, H, DV, dtype=torch.float64) mask = torch.tril(torch.ones(C, C, dtype=torch.float64)) for b in range (B): for h in range (H): for dv0 in range (0 , DV, block_DV): dv = slice (dv0, dv0 + block_DV) Sm = torch.zeros(DK, block_DV, dtype=torch.float64) for i_s in range (S // C): sl = slice (i_s * C, (i_s + 1 ) * C) Qb, Kb, Vb = Q[b, sl, h, :], K[b, sl, h, :], V[b, sl, h, dv] O[b, sl, h, dv] = Qb @ Sm + (Qb @ Kb.T * mask) @ Vb Sm = Sm + Kb.T @ Vb return O
这份参考与参考 B 的相对 L2 误差在 fp64 下应落在 10 − 16 10^{-16} 1 0 − 16 量级:它算的是同一个恒等式,只是把「每块重算前缀」换成了「顺序携带前缀」,数学上完全等价,冗余系数从 ( N C − 1 ) / 2 (NC-1)/2 ( N C − 1 ) /2 降到 1。
FLOPs 回到 S S S 的一次方–每 chunk 每 dv-block 两次 GEMM,求和得 4 S D K D V 4\,S\,DK\,DV 4 S D K D V 。与本文实现对比(D = 64 D = 64 D = 64 ,B C = 64 BC = 64 B C = 64 ):
N N N
本文步骤①
序列出 grid
倍数
8192
4.261 GFLOP
0.067 GFLOP
63.5x
16384
17.113 GFLOP
0.134 GFLOP
127.5x
32768
68.585 GFLOP
0.268 GFLOP
255.5x
65536
274.609 GFLOP
0.537 GFLOP
511.5x
倍数恰好是冗余系数 ( N C − 1 ) / 2 (NC-1)/2 ( N C − 1 ) /2 ,随 N N N 线性增长。
6.4.3 融合还是拆成两个 kernel
值得注意的是,§6.4.2 那个 kernel 没有把状态快照写回 HBM –因为它在同一个循环里就把 O O O 算完了:block 持有完整的 D K DK D K ,Q c S prev Q_c S^{\text{prev}} Q c S prev 和 tril ( Q c K c ⊤ ) V c \operatorname{tril}(Q_cK_c^\top)V_c tril ( Q c K c ⊤ ) V c 都能就地完成,状态从头到尾只活在寄存器里。
官方 chunk_delta_h 却在循环开头写了 T.copy(b_h_shared, h[...]),把每个 chunk 的状态快照全部落盘。h 的形状是 ( B , S / b l o c k S , H , D K , D V ) (B, S/block_S, H, DK, DV) ( B , S / b l oc k S , H , D K , D V ) ,按官方 main() 的配置(B = 1 , S = 32768 , H = 32 , D K = D V = 128 B{=}1, S{=}32768, H{=}32, DK{=}DV{=}128 B = 1 , S = 32768 , H = 32 , D K = D V = 128 ,chunk 64,bf16)达 512 MiB –与 K K K 、V V V 输入之和等量。
方案
K ⊤ V K^\top V K ⊤ V 次数
HBM 额外开销
适用
§4 序列进 grid
N C ( N C − 1 ) / 2 NC(NC-1)/2 N C ( N C − 1 ) /2
无
隔离验证恒等式
§6.4.2 融合单 kernel
N C NC N C
无
无门控/无删除的前向、推理
官方两 kernel(chunk_delta_h + chunk_o)
N C NC N C
h h h 快照(示例 512 MiB)
训练(反向要 h h h )、DeltaNet 的 UT 变换
拆开的三个真实理由:反向传播 需要每个 chunk 的 S prev S^{\text{prev}} S prev ,重算不如存;DeltaNet 的 W W W 与 ( I + tril ( diag ( β ) K K ⊤ ) ) − 1 (I + \operatorname{tril}(\operatorname{diag}(\beta)KK^\top))^{-1} ( I + tril ( diag ( β ) K K ⊤ ) ) − 1 需要沿 D K DK D K 收缩的独立阶段,塞不进这个循环;chunk_o 有自己的 tile 划分自由度 –读一份现成的快照就能对所有 chunk 完全并行,不再受递推顺序约束。顺序性和并行性被分到两个 kernel 里,各自取所需;融合省 HBM,拆开换灵活性和反向所需的中间量。 第一级用不到后两者,所以融合是更好的起点。
还有两处细节值得对照本文的实现:
clear_accum=True 的用途 。T.gemm(W_shared, b_h_shared, V_new_fragment, clear_accum=True) 显式覆盖而非累加,因为 V_new_fragment 每轮都要重算,不能沿用上一 chunk 的残留。本文 §4.1 靠循环外的 T.clear(S_f) 达到同样目的,两种写法都行;进入流水线循环后就必须用 clear_accum。
门控在 log 域相减,从不物化比值 。官方实现写作 T.exp2((G_last_local - G_fragment[i_s2, i_v]) * 1.442695),其中 G G G 已是 logsigmoid 后的累积和,1.442695 = log 2 e 1.442695 = \log_2 e 1.442695 = log 2 e 把 e x e^x e x 转成硬件 exp2。G G G 单调递减保证 G l a s t − G i ≤ 0 G_{last} - G_i \le 0 G l a s t − G i ≤ 0 ,指数结果恒不大于 1,不存在溢出 。本文的第一级没有门控,这里仅作为对照记录。
一句话总结 :线性注意力的全部优势是把 O ( N 2 D ) O(N^2 D) O ( N 2 D ) 换成 O ( N D 2 ) O(N D^2) O ( N D 2 ) ,而这个优势能否落地,取决于序列轴放不放进 grid 。放进去,block 间无法传递状态,只能各自重算,( N C − 1 ) / 2 (NC-1)/2 ( N C − 1 ) /2 的冗余把 N N N 的次数顶回 2;不放进去,递推退回单 block 内的顺序循环、用寄存器承载状态,并行度改由 batch/head/DV 提供,代价是把状态快照写回 HBM。本文选前者换取实现的可隔离性,生产实现选后者。
7. 总结
本级移除了 KDA 的全部衰减因子与删除因子 ,只保留 S ← S + K ⊤ V S \leftarrow S + K^\top V S ← S + K ⊤ V ,目的是将 chunkwise 分解恒等式 O c = Q c S c prev + tril ( Q c K c ⊤ ) V c O_c = Q_c S_c^{\text{prev}} + \operatorname{tril}(Q_c K_c^\top) V_c O c = Q c S c prev + tril ( Q c K c ⊤ ) V c 单独隔离验证。该恒等式是后续各级的公共基础,每级只在它的两项上插入衰减权重,GEMM 的形状与调用次序不变。
四层 PyTorch 参考构成从数学到 kernel 的完整链条 :A 逐 token 递归验证递推式定义,B 分块向量化验证恒等式,C 逐块独立重算镜像 §4 kernel 控制流,D(§6.4.2)按 ( b , h , d v ) (b,h,dv) ( b , h , d v ) 切 tile + 序列内循环镜像生产型 kernel。kernel 与参考的差异仅剩两项–f32→f16 显式降精度与 shared buffer 中转,这也是全部数值差异的来源。
与 FA 的差异集中在三点 :无 online softmax(掩码置 0 而非 − ∞ -\infty − ∞ )、循环携带 d k × d v d_k \times d_v d k × d v 状态、无按行归约(T.gemm policy 用默认 Square 即可)。
线性注意力的优势与本文实现的退化 :优势来自结合律换边,( Q K ⊤ ) V (QK^\top)V ( Q K ⊤ ) V 的 O ( N 2 D ) O(N^2 D) O ( N 2 D ) 变成 Q ( K ⊤ V ) Q(K^\top V) Q ( K ⊤ V ) 的 O ( N D 2 ) O(N D^2) O ( N D 2 ) ,理想实现与因果 FA 的比值按 2 ( D + B C ) / N 2(D+BC)/N 2 ( D + B C ) / N 衰减。但本文的逐块独立重算使冗余系数 ( N C − 1 ) / 2 (NC-1)/2 ( N C − 1 ) /2 随 N N N 线性增长,把 N N N 的次数顶回 2–实测比值收敛到常数 D / ( 2 B C ) = 0.5 D/(2BC) = 0.5 D / ( 2 B C ) = 0.5 。比值与 N N N 无关正是失败的判据,不是中性描述。
退化的根源是 grid 划分,不是算法 :把序列轴放进 grid.x 后,block 间既无执行次序保证也无同步原语,前缀状态只能各自重算。换成 grid = ( D V / b l o c k D V , B ⋅ H ) = (DV/block_{DV},\ B \cdot H) = ( D V / b l oc k D V , B ⋅ H ) –序列轴不进 grid ,递推退回单 block 内的 T.Pipelined 顺序循环,状态作为 loop-carried fragment 常驻寄存器,每 chunk 只做一次 K ⊤ V K^\top V K ⊤ V 。并行度改由 batch/head/DV 提供(示例配置 128 个 block)。N = 65536 N = 65536 N = 65536 时两者相差 511.5 倍。§6.4.2 给出了完整可运行的融合版本。
只切 DV、不切 DK 的真正理由不是「K ⊤ V K^\top V K ⊤ V 的 M 维」 。K ⊤ V K^\top V K ⊤ V 的收缩维是 chunk 长度 C C C ,DK 和 DV 在这一步都是自由维,切哪边都不需归约。不对称性来自下游:D K DK D K 是 Q c S prev Q_c S^{\text{prev}} Q c S prev 和 Q c K c ⊤ Q_c K_c^\top Q c K c ⊤ 两个 GEMM 的收缩维,D V DV D V 在三个 GEMM 里始终是自由维 。切 DK 会把读出变成 split-K(每 chunk 归约一次 [ C , D V ] [C, DV] [ C , D V ] ),并且块内项根本算不出完整的 [ C , C ] [C,C] [ C , C ] 转移矩阵;切 DV 只付出Q/K 的读放大,多读而不归约 。实测三个 block D V \text{block}_{DV} block D V (32/64/128)相对 L2 完全相同(5.60 × 10 − 16 5.60 \times 10^{-16} 5.60 × 1 0 − 16 ,见 §3.4)。
循环体内的顺序就是右移语义 :状态更新 T.gemm(K_s, V_s, S_f, transpose_A=True) 必须排在输出写回之后,提前就变成 inclusive 前缀、块内项被重复计入,实测相对 L2 从 10 − 16 10^{-16} 1 0 − 16 跳到 6.98 × 10 − 1 6.98 \times 10^{-1} 6.98 × 1 0 − 1 (见 §3.4)。还有一条配套规则:T.clear 只能在流水线循环外给跨迭代的状态用,循环内那些每轮重算的累加器必须靠 clear_accum=True。
融合还是拆两个 kernel :本级无门控无删除,单 block 持有完整 D K DK D K ,状态可以全程待在寄存器里,不需要写 h h h 快照 。官方 chunk_delta_h + chunk_o 拆开的理由是反向传播需要 S prev S^{\text{prev}} S prev 、DeltaNet 的 UT 变换需要沿 D K DK D K 收缩的独立阶段、以及 chunk_o 想要自己的 tile 划分自由度,代价是示例配置下 512 MiB 的 HBM 写回。
实测数字 (均 numpy fp64,见 §3.4):四层参考互验的相对 L2 误差均在 1.6 – 5.6 × 10 − 16 1.6\text{--}5.6 \times 10^{-16} 1.6 – 5.6 × 1 0 − 16 ,恒等式无误;刻意去掉右移的错误实现相对 L2 为 7.78 × 10 − 1 7.78 \times 10^{-1} 7.78 × 1 0 − 1 (参考 B)与 6.98 × 10 − 1 6.98 \times 10^{-1} 6.98 × 1 0 − 1 (参考 D),验证流程对结构性错误敏感。TileLang kernel 本身需 CUDA 设备,实测待补。
两个实现陷阱 :einsum 输出下标不可重复(->bhcdd 会抛异常,key 维与 value 维必须用不同字母);T.Pipelined 之后 shared buffer 的残留内容取决于流水线展开方式,复用前必须显式重载。
参考 :
本文的分块恒等式验证与 FLOPs 账本在 numpy fp64 上复现;GPU 实测数据待补。