TileLang 实战:KDA 从零到一--Chunked 线性注意力
TileLang 实战:KDA 从零到一–Chunked 线性注意力
KDA(Kimi Delta Attention)的递推式一行就写完,但直接照着它写 kernel 需要同时处理四个机制:逐通道门控、delta rule 的三角求解、log 域 cumsum、跨 chunk 状态传递。四个机制耦合在一个 kernel 里,数值出错时无法判断误差来自哪一层。 本文采用递进式实现路径,每一级只引入一个新机制、每一级都可独立运行并做数值验证。本文覆盖第一级–移除全部衰减因子与删除因子,只保留 。
这一级的价值不在性能,而在于将 chunkwise 分解恒等式单独隔离验证。本文同时给出一个反直觉的结论:当前实现采用的逐块独立重算策略会使 FLOPs 退回 ,与因果 FlashAttention 同量级–线性注意力的线性复杂度在这一级尚未兑现。
数学背景见《KDA 的来龙去脉》,TileLang 语言基础见《TileLang 编程基本知识点》,本文复用的 tile 切分与 fragment 累加模式见《TileLang 实战:FlashAttention 前向 Kernel》。
1. 递进路径:把 KDA 拆成可验证的增量
KDA 的完整递推式(,):
按机制拆解,每一级只放开一个自由度:
| 级别 | 递推式 | 新增机制 | 新出现的实现结构 |
|---|---|---|---|
| 第一级(本文) | – | 分块恒等式本身、块内因果掩码 | |
| 第二级 | 标量衰减 | 块内权重从 0/1 变 指数下三角 | |
| 第三级 | 逐 token 门控 | log 域 cumsum、exp2 硬件指令 | |
| 第四级 | delta rule | UT 变换、三角求解(wy_fast 雏形) |
|
| 第五级 | 门控 + delta rule | 二者耦合 | 五阶段流水线、跨 kernel 状态传递 |
第五级即完整 GDN / KDA,结构对标 flash-linear-attention 的 cumsum -> chunk_scaled_dot_kkt -> wy_fast -> chunk_delta_h -> chunk_o。
这样拆分的收益是误差定位能力:第三级数值不符,可以确定问题在 log 域 cumsum 或 exp2 精度,与 delta rule 无关,因为第四级尚未引入。反之若直接实现第五级,UT 变换的三角求解误差与逐通道衰减的溢出问题会互相掩盖。
一句话总结:递进式实现不是教学冗余,而是把多机制耦合 kernel 的调试搜索空间从乘法降为加法。
2. 数学推导:分块恒等式
本级递推不含遗忘项,状态单调累加:
展开为显式求和,因果且包含当前 token:
这是标准线性注意力。与 softmax attention 的唯一区别是分数未经 softmax,因此求和顺序可自由交换–这是下述分块重写成立的前提。
2.1 将求和拆分为跨块与块内两部分
序列按 切分为 个 chunk。对第 块内第 个 token,将 的求和范围拆为两部分–落在此前完整块内的,与落在当前块内的:
关键在第一项:括号内的量对整个 chunk 是同一个矩阵,记作 ,形状 ,含义是处理该 chunk 之前已累积的状态。于是跨块贡献退化为一次矩阵乘 ,块内贡献是一个 的下三角掩码矩阵乘:
这条恒等式是后续各级的公共基础,第二级至第五级都是在它的两项上分别插入衰减权重,张量形状与 GEMM 调用次序保持不变。
2.2 掩码可以直接置零的原因
FlashAttention 中掩码必须填 ,因为后续要经过 exp, 才能使被掩位置不产生贡献。本级无 softmax,掩码位置直接写 0 即可:
1 | for i, j in T.Parallel(BC, BC): |
这一行是本级实现与 FA 的第一处分岔,也是「无 softmax」在代码上的全部体现–不需要 online 重标定、不需要 与 两个行状态、不需要输出阶段的最终除法。
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 | def ref_recurrent(Q, K, V): |
3.2 参考 B:分块向量化
1 | def ref_chunked(Q, K, V, BC): |
states_prev 的右移拼接对应 中的严格小于号:cumsum 给出的是含自身的前缀和,右移一格补零才是处理该 chunk 之前的累积状态。这是分块线性注意力最易出错的一行–不右移等价于把当前块的 重复计入,块内贡献会被计算两次。该错误的量级实测为相对 L2 误差 ,属于结构性错误而非精度问题,容易识别。
3.3 参考 C:kernel 控制流镜像
1 | def ref_kernel_mimic(Q, K, V, BC): |
对照参考 B 与参考 C 可以看出两种前缀状态求法的区别:B 用 cumsum 一次性算出全部 个前缀状态、总代价 ;C 的每个 独立重算、总代价 。kernel 采用的是 C 的策略,原因与代价见 §6。
3.4 三层参考的一致性验证
,fp64(numpy 复现):
| 比较 | max abs 误差 | 相对 L2 |
|---|---|---|
| B 分块向量化 vs A 逐 token 递归 | ||
| C kernel 结构镜像 vs A 逐 token 递归 | ||
| C vs B | ||
| 参考 B 去掉右移(错误实现) |
前三行误差均在 fp64 机器精度量级(),恒等式与三份实现均无误。第四行是刻意引入的错误,用于确认该验证流程对结构性错误敏感。
3.5 einsum 输出下标的约束
参考 B 中若把 torch.einsum("bhcnd,bhcnv->bhcdv", ...) 误写为 ->bhcdd(意图表达「输出是 方阵」),会直接抛出异常:
1 | ValueError: einstein sum subscripts string includes output subscript 'd' multiple times |
einsum 的输出下标不允许重复–重复下标在输出侧的语义是取对角线。而 的两个维度虽然长度均为 ,语义上分别是 key 维与 value 维,必须使用不同字母 d 与 v。同理 bhcnd,bhcdd->bhcnv 也不合法,因为输出的 v 从未在输入中出现。 时长度相同掩盖了语义差异,一旦 KDA 中 ,该疏忽会立即表现为形状错误。
4. TileLang kernel:三步实现
grid 划分与 FA 一致–按 Q 块切分所有权,每个 block 负责输出一个 的 tile:
1 |
|
4.1 步骤①:流式累加前缀状态
对应参考 C 中的 for c in range(bx):
1 | T.clear(S_f) |
T.Pipelined(bx) 的上界是 block 索引而非编译期常量–不同 block 的循环次数不同,第 0 块一次都不执行(),最后一块需执行 次。TileLang 支持 runtime 上界的流水线,代价是各 block 负载严重不均衡,尾部 block 构成整个 kernel 的关键路径。
4.2 步骤②③:两项贡献求和
1 | # ② 跨块:O = Q_c @ S_prev |
步骤③开头必须重新 T.copy 当前块的 K/V。步骤①的流水线循环结束时,K_s / V_s 中残留的是第 块的数据,且在 num_stages > 1 时具体残留哪一块取决于流水线的展开方式,不可假设。复用 shared buffer 节省了显存,但必须显式重载–这是 tile 编程中典型的隐式状态陷阱。参考 C 里没有这个问题,因为 Python 每轮都重新切片,不存在缓冲区复用。
T.copy(A, A_cast) 这次 f32→f16 转换与 FA 中 喂入第二个 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,维护 、 两个行状态,每块重标定 | 无,掩码直接置 0 | 无 exp,求和顺序可交换 |
| 循环携带的量 | 无跨块状态,每个 Q 块独立 | 是跨迭代累加的状态 | 递推式本身带状态 |
T.gemm policy |
必须 FullRow–行归约要求整行在同一 warp |
默认 Square 即可 |
无按行归约 |
第三点值得展开:FA 需对分数块做 rowmax / rowsum,若一行被切分到多个 warp,归约就需跨 warp 通信,因此必须用 FullRow policy 强制整行不拆。本级的块内矩阵 计算完成后仅做逐元素掩码,无任何跨列归约,warp 划分方式不影响正确性,编译器可自由选择寄存器分布最优方案。
第二点是从本级走向第五级的主要矛盾。当前实现用逐块独立重算把状态依赖隐藏了,但第五级的 chunk_delta_h 必须真正在 chunk 之间传递状态–届时状态需在 fragment(累加)与 shared(作为下一个 GEMM 输入)之间每轮拷贝一次,并需拆成两个 kernel 协作。
6. 代价账本: 仍在,线性复杂度尚未兑现
逐块独立重算的冗余可以精确计算。第 块执行 次 ,全部 block 合计 次,而理想情况每块只需计算一次、共 次(即参考 B 的 cumsum 策略)。冗余系数为 。
单次 的 FLOPs 为 ,取 ,步骤①的总量:
对比因果注意力的 (两个 GEMM 各占一半三角),比值为:
该比值与 无关–即本级实现的计算量与因果 FlashAttention 属同一量级,仅差一个由 决定的常数。实测账本(,):
| 冗余系数 | 步骤① FLOPs | 理想(参考 B 策略) | 因果 FA | ①/因果 FA | ||
|---|---|---|---|---|---|---|
| 512 | 8 | 3.5x | 0.015 GFLOP | 0.004 GFLOP | 0.034 GFLOP | 0.438 |
| 2048 | 32 | 15.5x | 0.260 GFLOP | 0.017 GFLOP | 0.537 GFLOP | 0.484 |
| 8192 | 128 | 63.5x | 4.261 GFLOP | 0.067 GFLOP | 8.590 GFLOP | 0.496 |
| 16384 | 256 | 127.5x | 17.113 GFLOP | 0.134 GFLOP | 34.360 GFLOP | 0.498 |
时冗余系数达 127.5 倍,而步骤②③合计始终是 ( 时分别为 0.067 与 0.134 GFLOP,比步骤①小两个数量级)。结论是:本级实现 99% 以上的开销耗在冗余状态重算上,线性注意力的 复杂度在这一级完全没有兑现。
的取舍随之明确:
| 含义 | ||
|---|---|---|
| 32 | 1.000 | 与因果 FA 计算量持平,块过小无收益 |
| 64 | 0.500 | 默认值,寄存器压力可控 |
| 128 | 0.250 | 冗余减半,但 占 f32 fragment,每线程约 128 个寄存器,易溢出 |
| 256 | 0.125 | 理论最省,实际 shared memory 与寄存器均无法容纳 |
增大 能线性减少冗余,但 是 的 f32 fragment, 时寄存器压力已接近溢出边界。真正的解法不是调 ,而是替换逐块独立重算策略–改用两个 kernel 协作,第一个 kernel 按参考 B 的方式算出所有 chunk 的前缀状态写回 HBM,第二个 kernel 读取。这正是第五级 chunk_delta_h 的结构,也解释了官方实现为何拆成五个 kernel 而非一个。
一句话总结:本级用 倍冗余计算换取单 kernel 闭环与零跨 block 同步,该取舍在教学上成立、在生产上不成立–它把线性注意力的核心优势整体交换掉了。
7. 数值验证
四层验证的组织方式:
- 参考 A vs B vs C(均 fp64):验证 §2.1 分块恒等式与 kernel 控制流,与 GPU 无关,误差应在 量级;
- 刻意引入的错误实现:确认验证流程对结构性错误敏感(去掉右移,相对 L2 应达 量级);
- kernel vs 参考 C:验证 TileLang 实现,fp16 输入 + f32 累加,阈值取相对 L2 ;
- 延迟对比走 CUDA event 中位数,取 50 次采样。
第 3 层阈值定在 而非更严,原因是 T.copy(S_f, S_s) 把 f32 状态降至 f16 落 shared。该降精度是必要的–Tensor Core MMA 的输入必须是低精度–但状态 是多个块累加的结果,越靠后的 block 累加项越多,误差随 单调增长。这是本级精度的主导误差源,影响远大于块内 的那次降精度。
第 1、2 层验证已在 numpy fp64 上完成(见 §3.4 表格)。第 3、4 层需要 CUDA 设备,本文未给出实测数字–待真卡跑通后单独补充实测数据,此处不做性能推测。
运行方式:
1 | python gdn_ladder_L1_linear_attn.py --seq_len 512 --dim 64 --blk 64 |
使用 parse_known_args() 而非 parse_args():Jupyter / IPython 会向 sys.argv 注入 kernel 连接参数,用 parse_args() 会直接报错退出。
8. 总结
- 本级移除了 KDA 的全部衰减因子与删除因子,只保留 ,目的是将 chunkwise 分解恒等式 单独隔离验证。该恒等式是后续四级的公共基础,每级只在它的两项上插入衰减权重,GEMM 的形状与调用次序不变。
- 三层 PyTorch 参考构成从数学到 kernel 的完整链条:A 逐 token 递归验证递推式定义,B 分块向量化验证恒等式,C 逐块独立重算镜像 kernel 控制流。kernel 与 C 的差异仅剩两项–f32→f16 显式降精度与 shared buffer 中转,这也是全部数值差异的来源。
- 与 FA 的差异集中在三点:无 online softmax(掩码置 0 而非 )、循环携带 状态、无按行归约(
T.gemmpolicy 用默认Square即可)。 - 实测数字:三层 fp64 参考互验的相对 L2 误差均在 ,恒等式无误;刻意去掉右移的错误实现相对 L2 为 ,验证流程对结构性错误敏感。代价账本显示 时逐块独立重算的冗余系数达 127.5 倍,步骤① FLOPs 是因果 FlashAttention 的 0.498 倍,且该比值 与 无关–线性复杂度在本级尚未兑现。
- 两个实现陷阱:einsum 输出下标不可重复(
->bhcdd会抛异常,key 维与 value 维必须用不同字母);T.Pipelined之后 shared buffer 的残留内容取决于流水线展开方式,复用前必须显式重载。
下一篇加入标量衰减 。改动集中在块内那一项–下三角掩码从 0/1 矩阵变为 的指数权重,跨块那一项则需给 乘上整块的衰减 。 一旦从标量升级为逐通道向量,指数外提就会引入 的数值溢出,那是这条实现路径上第一个真正棘手的问题。
可迁移的启示:多机制耦合的 kernel 不应一次写完。每级只放开一个自由度,并且每级都保留一份可独立运行的 fp64 参考–其中至少一份要与 kernel 的控制流同构。参考实现的成本远低于在五个机制中定位一个误差源的成本。
参考:
- Kimi Linear / KDA 论文:arXiv:2510.26692
flash-linear-attention:fla/ops/gated_delta_rule/、fla/ops/kda/naive.py- TileLang
examples/gdn(第五级的模板底本) - 本站《KDA 的来龙去脉》《TileLang 编程基本知识点》《TileLang 实战:FlashAttention 前向 Kernel》
本文的分块恒等式验证与 FLOPs 账本在 numpy fp64 上复现;GPU 实测数据待补。