解法是分块递归:对角块的逆互不依赖、可以并行求,块间耦合退化成矩阵乘。C=32 时依赖链从 31 步压到 17 步,代价是算术量涨到 1.88 倍;C=64 多切一层,63 步压到 19 步。本文以 Qwen 团队 flash_qla 里的 kkt_solve 为样本,把这个 kernel 从数学恒等式一路拆到 bank conflict,包括三个不看代码想不到的工程细节,以及块长 C 到底该取 32 还是 64。
1. 问题从哪来
DeltaNet 系列(DeltaNet / Gated DeltaNet / KDA)的 chunkwise 形式里,块内 C 个 Householder 变换连乘会被 UT 变换整理成一个三角系统。系数矩阵长这样:
A=I+strictLower(diag(β)KK⊺)∈RC×C
K 是块内 C 个 key 排成的矩阵(C×dk,论文里 k 做过 L2 归一化),βt∈(0,1) 是写入强度。要求的是 T=A−1,GDN 论文 (§3.3) 用它构造两个量:
W[t]=Tdiag(β)K,U[t]=Tdiag(β)V
推导细节见《线性注意力的线性代数前置知识》§4 与《TileLang 实战:KDA 从零到一–Gated DeltaNet》§2.2,本文只关心一件事:这个 T 怎么在 GPU 上算得快。
论文 §3.3 的标准形式里 T 被用两次(W 和 U),不物化就得做两次前代、两条串行链。kkt_solve 于是选择物化:一个独立 kernel 把 T 算出来写回 HBM,下游两处都退化成 T.gemm。(对照组:GDN 那篇 §4.2 的右端项只有一个,于是把前代直接作用在它行上,不物化 T。)
既然要物化,就得正面回答「怎么快速求一个单位下三角矩阵的逆」。
2. 两种算法
2.1 逐格前代:O(n3/6) 但 n 步串行
求 X=L−1 即解 LX=I。因为 L 单位对角,X 也是单位下三角–对角上的 1 和上三角的 0 不用算,直接就是答案。真正要算的只有严格下三角部分。
按行写出递推(ek 是第 k 个标准基向量):
X[k,:]=ek⊺−j<k∑L[k,j]X[j,:]
算第 k 行要等前 k−1 行全部算完,n 行就是 n−1 步串行。
2.2 分块递归:把长链换成短链 + GEMM
先看 C=32:把 L 按 2×2 切块(每块 16×16):
L=[L00L100L11]⟹L−1=[L00−1−L11−1L10L00−10L11−1]
验证只需一次分块乘法:左下块 L10L00−1+L11⋅(−L11−1L10L00−1)=0。
价值全在依赖结构上:
- L00−1 与 L11−1 互不依赖,可以并行求;
- 耦合项是两次矩阵乘,level-3、无串行依赖。
1 2 3 4 5 6 7 8 9 10 11 12
| L (32×32 单位下三角) │ 切成 2×2 块(16×16) ▼ ┌──────────┬──────────┐ │ L00 │ 0 │ ├──────────┼──────────┤ │ L10 │ L11 │ └──────────┴──────────┘ │ ├── ① L00⁻¹, L11⁻¹ ← 两块并行做 15 步前代 │ (原地覆盖,不占额外空间) └── ② -L11⁻¹·L10·L00⁻¹ ← 两次 16×16×16 GEMM,无依赖
|
C=64:递归一层不够,切四份。 上面的 2×2 只是 C=32 的形态。C=64 时对角块还是 32×32,本身仍需 31 步前代–所以要再往下递归一层,最终落到 4×4 的 16×16 网格:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16
| L (64×64) ← 第 1 层:切 2×2 的 32×32 ┌───────────────┬───────────────┐ │ M00 (32×32) │ 0 │ │ ┌─────┬─────┐│ │ ← 第 2 层:每个 M 再切 2×2 的 16×16 │ │ L00 │ 0 ││ │ │ ├─────┼─────┤│ │ │ │ L10 │ L11 ││ │ │ └─────┴─────┘│ │ ├───────────────┼───────────────┤ │ M10 (32×32) │ M11 (32×32) │ │ │ ┌─────┬─────┐│ │ 稠密 │ │ L44 │ 0 ││ │ │ ├─────┼─────┤│ │ │ │ L54 │ L55 ││ │ │ └─────┴─────┘│ └───────────────┴───────────────┘
|
三个阶段,并行度呈阶梯上升:
| 阶段 |
内容 |
深度 |
独立任务数 |
并行输出元素 |
| ① |
4 个 16×16 前代 |
15 |
4 |
4×16=64 |
| ② |
2 组 M−1 耦合,各 2 次 163 GEMM |
+2 |
2 |
2×256=512 |
| ③ |
1 组块间耦合,2 次 323 GEMM |
+2 |
1 |
32×32=1024 |
|
合计 |
19 |
|
|
M00−1 和 M11−1 在阶段 ①② 全程互不相干,两条链完全并行;阶段 ③ 才把它们缝起来。依赖链 63 → 19,而每一步的可并行宽度从 32 涨到 1024–C 翻倍时 GEMM 段的规模按 C3 膨胀,正好把大 chunk 多出来的算术量塞进 Tensor Core,不落在关键路径上。
值得对比的是另一条路:扁平的块前代–把 4×4 网格当成一个 4×4 的“标量”下三角矩阵,直接跑 §2.1 的行式递推,每个“元素”换成 16×16 GEMM:
| C=64 方案 |
FMA |
依赖链 |
163 GEMM 次数 |
| 整块逐格前代 |
41664 |
63 |
– |
| 扁平块前代(4×4 网格) |
67776 |
21 |
16 |
| 两层分块递归 |
84160 |
19 |
4 + 2 组 323 |
扁平块前代省 20% 算术量,但块行之间仍是串行的(第 3 块行要等第 2 块行),依赖链 21 步且后期块行的并行度递减。分块递归多花那 20%,换来的是两个 32×32 子问题的完全独立–在 TileLang 里这直接对应两组互不同步的 warp。C=128 时差距进一步拉开:扁平块前代 29 步、分块递归 21 步。
2.3 两种算法的账
C=32、递归到 16×16 为止:
|
逐格前代(整块 32) |
分块递归(2×16) |
| 串行依赖链 |
31 步 |
15 步 + 2 次 GEMM |
| 数学必需 FMA |
4960 |
2×560+2×163=9312 |
| 算术量比 |
1.00 |
1.88 |
| 前代段并行 lane |
32 |
32(2 块 × 16 列) |
| GEMM 段并行 lane |
– |
256 |
| 主体算子 |
AXPY(level-2) |
GEMM(level-3) |
| 额外存储 |
无 |
一个 16×16 中间量 |
多花 88% 的算术,换依赖链减半、且一半的工作量跑在 level-3 上。 小规模三角求解本来就填不满 SM,算术单元闲着,这笔账划算。
递归可以继续往下切:
| C |
递归层数 |
对角 16 块数 |
耦合 GEMM 对数 |
分块依赖链 |
整块依赖链 |
| 32 |
1 |
2 |
1 |
$15 + 2 = $ 17 |
31 |
| 64 |
2 |
4 |
3 |
$15 + 4 = $ 19 |
63 |
| 128 |
3 |
8 |
7 |
$15 + 6 = $ 21 |
127 |
前代段恒定 15 步(永远只有最底层的 16×16 在做前代),每多一层递归只增加两次串联 GEMM 的深度。所以依赖链是 O(logC) 而非 O(C),而新增的 GEMM 段内部完全可并行。
2.4 chunk 级别的账:C 该取 32 还是 64
上面是 kernel 内部的账,但选块长要放到整条 chunkwise 流水线里看。求逆的算术量按 O(C3) 涨、而它服务的 token 只有 C 个,摊到每 token 是 O(C2):
| C |
分块递归 FMA |
摊到每 token |
主 GEMM 每 token |
求逆占比 |
| 32 |
9312 |
291 |
40960 |
0.7% |
| 64 |
84160 |
1315 |
49152 |
2.6% |
| 128 |
692608 |
5411 |
65536 |
7.6% |
(主 GEMM 口径:块内 QK⊺ + AV 加块间 QS + 状态更新,即 2dkdv+C(dk+dv),取 dk=dv=128。)
另一头,块长变大能省的是状态传递的 HBM 流量–每个 chunk 只读写一次 dk×dv 状态,摊到每 token 是 2⋅2dkdv/C 字节:
| C |
state 读写 |
T 写回 |
一张 C2 fp32 中间表 |
cond(A) 实测上界 |
| 32 |
2048 B/token |
64 B/token |
4 KB |
2.0 |
| 64 |
1024 B/token |
128 B/token |
16 KB |
2.6 |
| 128 |
512 B/token |
256 B/token |
64 KB |
3.6 |
三条线交叉出来的结论:
- C 从 32 到 64:state 流量减半(2048 → 1024 B/token,这是 chunkwise 最大的一笔片外开销),代价是求逆占比从 0.7% 涨到 2.6%、C2 表从 4 KB 涨到 16 KB。这笔交换通常划算,所以 GDN/KDA 那几篇的主 kernel 都取 C=64。
- C 从 64 到 128:state 流量只再省 512 B,但 QK⊺、Γ、A 这些 C2 表同时膨胀到 64 KB 一张–单个 SM 的 shared memory 就装不下几张,寄存器也会超限(GDN 篇 §4.1 实测
Gam + A + Amat 三张表在 C=128 时直接爆 reg)。上限卡在 64。
kkt_solve 单独取 C=32,是因为它是独立 kernel:T 要写回 HBM 再被下游读回来,这笔 2C B/token 的流量在融合 kernel 里是不存在的。C 翻倍它也翻倍,而 T 只是个中间量–省不了 state 流量、只多付带宽。物化路线的最优块长因此比融合路线更小。
一句话:融合 kernel 里 C 由 shared/reg 容量顶到 64;物化 T 的独立 kernel 里 C 由中间量写回带宽压回 32。 同一个算法,落在流水线的不同位置,最优块长不一样。
2.5 到底是谁卡住了 chunk size
上一节说“容量顶住了”,但究竟是寄存器还是 shared memory,取决于 kernel 怎么写。TileLang 里常见的三种写法,瓶颈完全不在一个地方。
写法 A:单 warp-group(128 线程),中间量常驻 fragment。这是 GDN/KDA 前几篇的写法,C×C 的 Γ、A、Amat 都放在寄存器里。128 线程分一张 C2 fp32 表,每线程 C2/128 个 reg:C=64 是 32 reg/表,C=128 就是 128 reg/表–三张表 384 reg,直接超过 255 的硬上限。卡在寄存器,上限 64。
写法 B:warp-specialized(512 线程),双缓冲过 shared。生产者 warp 搜 HBM、消费者 warp 做 MMA,中间量全走 shared memory。以 dk=128、dv=blockDV=64、bf16 为例:
| 缓冲区 |
随什么长 |
C=32 |
C=64 |
q_shared / k_shared(双缓冲 2Cdk) |
∝C |
16 + 16 KB |
32 + 32 KB |
v_shared(双缓冲 2Cdv) |
∝C |
8 KB |
16 KB |
o / h / vd / vn(各 Cdv) |
∝C |
16 KB |
32 KB |
a_shared(双缓冲 2C2) |
∝C2 |
4 KB |
16 KB |
p_shared(C2) |
∝C2 |
2 KB |
8 KB |
| 门控表 + barrier |
O(C) |
~1 KB |
~1 KB |
| 合计 |
|
63 KB |
137 KB |
这张表里有个容易看反的地方:大头是 ∝C 的项而不是 C2 项–C=64 时线性项 112 KB、平方项才 24 KB。因为 dk=128>C,Q/K 的双缓冲才是大户。结果是容量几乎随 C 线性翻倍,很快撞上硬件墙:
| 架构 |
shared/SM |
C=32(63 KB) |
C=64(137 KB) |
| Ada / L40S / RTX 4090 (sm89) |
100 KB |
1 block/SM |
跑不起来 |
| A100 (sm80) |
164 KB |
2 block/SM |
1 block/SM |
| H100 / H800 (sm90) |
228 KB |
3 block/SM |
1 block/SM |
即使在 H100 上能跑,C=64 也只剩 1 block/SM–没有第二个 block 掩盖尾巴和同步空档,而 warp-specialized 本来就是靠内部流水掩盖延迟的。想保 C=64 又想多一个 block,只能去掉双缓冲(降到 89 KB)–而那恰恰把 warp-specialized 的意义抵消了。卡在 shared memory,且消费级卡比数据中心卡卡得更早。
写法 C:独立物化 kernel(就是 kkt_solve)。它只吃 K、只吐 T,shared 就一张 C×dk 加几张 16×16,容量根本不是问题。卡它的是 §2.4 说的 T 写回 HBM 的 2C B/token。
汇总成一张表:
|
A:单 warp-group 融合 |
B:warp-specialized 融合 |
C:独立物化 |
| 线程数 |
128 |
512 |
128 |
| 中间量住哪 |
fragment |
shared(双缓冲) |
fragment |
| 主瓶颈 |
寄存器 255/线程 |
shared/SM 容量 |
中间量 HBM 带宽 |
| 随 C 增长 |
C2/128 reg/表 |
近似线性(因 dk>C) |
线性 2C B/token |
| 实际上限 |
64 |
64(且 1 block/SM);sm89 上 32 |
32 |
| 次要约束 |
串行前代并行度 |
生产/消费 warp 配比 |
求逆 O(C2)/token |
三种写法的最优 C 差不多都落在 32~64,但卡住它们的是三种不同的资源。换硬件、换 dv、多加一张中间表,哪个先爆都会变–所以调 C 之前先弄清楚自己卡在哪一条线上,别拿别人的数字当结论。
3. TileLang 实现
3.1 内存布局
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20
| block_S = 32 k_shared = T.alloc_shared((block_S, DK), dtype=qkva_dtype) b_shared = T.alloc_shared((block_S), dtype=accum_dtype) a32_fragment = T.alloc_fragment((block_S, block_S), dtype=accum_dtype) a32_shared = T.alloc_shared((block_S, block_S), dtype=accum_dtype)
a16i_shared = T.alloc_shared((2, 17, 16), dtype=accum_dtype)
a16o_shared = T.alloc_shared((1, 17, 16), dtype=accum_dtype) a16o_fragment = T.alloc_fragment((1, 16, 16), dtype=accum_dtype)
a16i_row = T.alloc_fragment((2, 16), dtype=accum_dtype) a16i_sum = T.alloc_fragment((2, 16), dtype=accum_dtype)
T.annotate_layout({ a16i_shared: tilelang.layout.make_linear_layout(a16i_shared), a16o_shared: tilelang.layout.make_linear_layout(a16o_shared), })
|
shared 总量约 15.5 KiB,threads=128(4 warp),一个 SM 能塞下多个 block。accum_dtype="float32"、输出 a 用 k.dtype(fp16)–单位下三角阵行列式恒为 1、不需 pivoting,精度这里给的余量很大(§4 有实测)。
3.2 构造系数矩阵
1 2 3 4 5 6 7 8 9 10 11
| T.gemm(k_shared, k_shared, a32_fragment, transpose_B=True, clear_accum=True)
for j_s, j_t in T.Parallel(block_S, block_S): a32_fragment[j_s, j_t] *= b_shared[j_s]
for j_s, j_t in T.Parallel(block_S, block_S): if j_s < j_t: a32_fragment[j_s, j_t] = 0 elif j_s == j_t: a32_fragment[j_s, j_t] = 1
|
β 乘在行上(b_shared[j_s]),对应 diag(β)KK⊺ 的左乘。对角必须置 1 而不是保留 βi∥ki∥2=βi–UT 变换的系数矩阵是 I+strictLower(⋅),含对角就变成解 (I+Afull),βi 被重复计入。
3.3 拆块:提取时就取负
1 2 3 4 5
| for j_s, j_t in T.Parallel(block_S, block_S): if (j_s // 16) == (j_t // 16) + 1: a16o_shared[j_s // 32, j_s % 16, j_t % 16] = -a32_fragment[j_s, j_t] elif (j_s // 16) == (j_t // 16): a16i_shared[j_s // 16, j_s % 16, j_t % 16] = a32_fragment[j_s, j_t]
|
工程点一:次对角块存的是 −L10 而不是 L10。恒等式里的负号被提前吃掉,后面两次 GEMM 直接得到 −L11−1L10L00−1,省一遍 162 的取反。这种「负号往前提」的技巧在 kernel 里很常见–取反本身不贵,但它是一个独立的 pass,要多读写一次 shared。
3.4 对角块并行前代
1 2 3 4 5 6 7 8 9 10 11
| for k_s in T.unroll(1, 16): for j_s, k_t in T.Parallel(2, 16): if k_t < k_s: a16i_row[j_s, k_t] = a16i_shared[j_s, k_s, k_t] T.clear(a16i_sum) for k_r in T.unroll(k_s): for j_s, k_t in T.Parallel(2, 16): a16i_sum[j_s, k_t] -= a16i_shared[j_s, k_r, k_t] * a16i_row[j_s, k_r] for j_s, k_t in T.Parallel(2, 16): if k_t < k_s: a16i_shared[j_s, k_s, k_t] = a16i_sum[j_s, k_t]
|
三点值得说:
原地覆盖是免费的。 输入 L 和输出 X 都是单位下三角,对角的 1 和上三角的 0 两边完全一样,所以只需改写严格下三角部分(k_t < k_s),不用额外开一块 buffer。
a16i_row 不是优化,是必需。 第 ks 行马上要被结果覆盖,但它的原始值 L[ks,j] 还要当递推系数用。不备份就会读到刚写进去的新值,算出来的东西静默错误–这是最容易踩的坑,因为它不报错。
两个块并行是分块递归的收益兑现处。 T.Parallel(2, 16) 的 32 条 lane 里,j_s=0 处理 L00、j_s=1 处理 L11,两条 15 步的链同时走完。对比整块 32 要走 31 步、每步同样 32 条 lane–并行度没变,时间减半。
工程点二:(2, 17, 16) 里的 17。 fp32 下 shared memory 有 32 个 bank,每 bank 4 字节。若声明成 (2, 16, 16),两个块的起始偏移差是 16×16=256,而 256mod32=0–T.Parallel(2, 16) 展开成 32 条 lane 时,j_s=0 的第 i 条和 j_s=1 的第 i 条落在同一个 bank 上,32 条 lane 两两冲突,访存吞吐直接减半。
改成 17 之后偏移差变成 17×16=272,272mod32=16,两个块正好错开半个 bank 区,32 条 lane 无冲突:
| 声明 |
块间偏移 |
mod 32 |
结果 |
(2, 16, 16) |
256 |
0 |
❌ 32 lane 两两撞 bank |
(2, 17, 16) |
272 |
16 |
✅ 错开半区,无冲突 |
配套的 T.annotate_layout(make_linear_layout(...)) 同样关键:TileLang 默认可能给 shared 张量套 swizzle 布局来避免冲突,但那会重新排列元素位置,让手工 padding 失效。显式声明线性布局,才能保证 17 这个数按预期生效。
3.5 耦合项:两次 GEMM 与中间那次转置
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17
| T.clear(a16o_fragment) for k_r in T.unroll(16): for j_s, k_s, k_t in T.Parallel(1, 16, 16): a16o_fragment[j_s, k_s, k_t] += ( a16i_shared[j_s * 2 + 1, k_s, k_r] * a16o_shared[j_s, k_r, k_t])
for j_s, k_s, k_t in T.Parallel(1, 16, 16): a16o_shared[j_s, k_t, k_s] = a16o_fragment[j_s, k_s, k_t]
T.clear(a16o_fragment) for k_r in T.unroll(16): for j_s, k_s, k_t in T.Parallel(1, 16, 16): a16o_fragment[j_s, k_s, k_t] += ( a16o_shared[j_s, k_r, k_s] * a16i_shared[j_s * 2, k_r, k_t])
|
工程点三:中间那次转置不是多余的。 Step 2 要算 temp⋅L00−1,按 kr 累加时需要读 temp[:, k_r] 这一列。列方向在行主序里跨度是 16(或 padding 后 17),16×16 的 lane 会散在多个 bank 上。先转置写回 shared,读取就变成沿最内维连续–配合 §3.4 的 17-padding,两处优化是配套的,只做一个另一个就白费。
注意两次 GEMM 的并行度是 16×16=256 条 lane,比前代段的 32 条高 8 倍。这正是「拿算术换深度」换到的东西:多出来的 8192 次 FMA 跑在满并行度上,实际墙钟时间的增量远小于 1.88 倍。
3.6 拼装
1 2 3 4 5
| for k_s, k_t in T.Parallel(16, 16): a32_shared[k_s, k_t ] = a16i_shared[0, k_s, k_t] a32_shared[k_s, 16 + k_t] = 0 a32_shared[16 + k_s, k_t ] = a16o_fragment[0, k_s, k_t] a32_shared[16 + k_s, 16 + k_t] = a16i_shared[1, k_s, k_t]
|
(原代码写成四个独立的 T.Parallel(16,16) 循环,合并成一个可以少三次同步。)
右上角必须显式清零–a32_shared 是复用的 buffer,不清会带进上一个 chunk 的残留。
4. 正确性验证
用 numpy 逐段镜像 kernel 的写法(含原地覆盖、提取时取负、两步 GEMM),与 np.linalg.inv 对照:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22
| def row_forward_inplace(L): """镜像 §3.4:原地覆盖,只算严格下三角""" X = L.copy() for ks in range(1, len(L)): row = X[ks, :ks].copy() s = np.zeros(len(L)) for kr in range(ks): s -= X[kr, :] * row[kr] X[ks, :ks] = s[:ks] return X
def blocked_inv(L, base=16): """镜像 §3.5:分块递归""" n = len(L) if n <= base: return row_forward_inplace(L) h = n // 2 I00, I11 = blocked_inv(L[:h,:h], base), blocked_inv(L[h:,h:], base) out = np.zeros_like(L) out[:h,:h], out[h:,h:] = I00, I11 out[h:,:h] = I11 @ (-L[h:,:h]) @ I00 return out
|
结果:
| 用例 |
误差(ℓ∞) |
| 16×16 行式前代 |
2.22×10−15 |
| 32×32 分块递归 |
4.44×10−14 |
| 真实 A(200 组采样) |
1.39×10−16 |
| fp16 存储 T 的相对误差 |
5.72×10−5 |
真实场景误差比一般随机矩阵低两个数量级,原因是 k 的 L2 归一化 + β∈(0,1) 把严格下三角元素压到 0.358 以内,cond(A) 实测不到 2.2。
两个必做的退化检验:
- β→0:A→I,T 应该是单位阵。这检验「对角置 1」和拼装时的清零。
- 只有次对角块非零(L00=L11=I):T 的左下块应恰好是 −L10。这单独检验耦合路径和那个负号,其他情况下负号错了会被对角块的贡献掩盖。
5. 总结
- 单位下三角求逆的瓶颈是依赖链,不是算术量。C=32 只需 4960 次乘加,但逐格前代要走 31 步串行,每步只有 32 条 lane 可并行(恰好一个 warp),
threads=128 时剩下 96 条线程全程闲置。
- 分块递归把长链换成短链 + GEMM,递归深度随 C 变。C=32 切一层(2×2 的 16×16)就到底;C=64 要切两层、落到 4×4 网格,且两个 32×32 子问题全程独立。依赖链 31 → 17(63 → 19),算术量 1.88 倍,而并行宽度从 32 涨到 1024。
- 前代段恒定 15 步,依赖链是 O(logC)。每多一层递归只加两次串联 GEMM(C=32/64/128 → 17/19/21 步),而 GEMM 段可并行。对比“扁平的块前代”(把网格当标量下三角跑行式递推):C=64 时它省 20% 算术量但依赖链 21 步且后期并行度递减。所以块长不是被求逆算法卡住的。
- C 取 32 还是 64,取决于这个 kernel 在流水线的哪个位置(§2.4)。求逆摊到每 token 是 O(C2):C=64 时只占主 GEMM 的 2.6%,而 state 的 HBM 流量从 2048 降到 1024 B/token–融合 kernel 里这笔交换划算。但
kkt_solve 是独立 kernel,T 要写回 HBM 再读回来(2C B/token 随 C 线性涨),省不到 state 流量、只多付带宽,于是压回 32。
- 卡住 chunk size 的资源因写法而异(§2.5)。单 warp-group + fragment 卡寄存器(一张 C2 表 = C2/128 reg/线程,C=128 时三张表 384 reg 超限);warp-specialized + 双缓冲卡 shared memory(C=32 → 63 KB、C=64 → 137 KB,后者在 sm89 直接跑不起来、在 H100 上也只剩 1 block/SM);独立物化 kernel 卡中间量带宽。值得注意的是 shared 里大头是 ∝C 的 Q/K 双缓冲而非 C2 表(dk=128>C),所以容量近似线性翻倍。
- 原地覆盖免费,但要备份当前行。L 与 X 的对角 1 和上三角 0 完全一致,只需改写严格下三角;代价是第 ks 行的原始值必须先存进
a16i_row,否则读到已覆盖的新值–静默出错,不报错。
(2, 17, 16) 的 17 是 bank conflict 解药。(2,16,16) 时两块偏移差 256≡0(mod32),T.Parallel(2,16) 的 32 条 lane 两两撞 bank;17 使偏移差 272≡16,错开半区。必须配 annotate_layout(make_linear_layout),否则 TileLang 换成 swizzle 会让 padding 白做。
- 两次 GEMM 中间的转置不是多余。Step 2 要读中间量的列,转置后变成沿最内维连续,与 17-padding 是配套优化,缺一个另一个失效。
- 负号提前吃掉。提取次对角块时直接存 −L10,省一遍 162 取反 pass。
- 物化 T 与否取决于它被用几次。只用一次就把前代直接作用在右端项上;GDN §3.3 的 W、U 两处复用,才值得单独起 kernel。
- 这个问题良态得反常。k 的 L2 归一化 + β∈(0,1) 让 cond(A)≤2.18(200 组采样),fp32 累加、fp16 存储都够。但换个来源的三角矩阵要重算–n=128 的一般随机三角矩阵条件数是 107 量级。
- Γ 不在这个 kernel 里。带门控时系数是 diag(β)(Γ⊙KK⊺),折叠需要两份缩放不同的 K,而
T.gemm(k_shared, k_shared, ...) 左右同一份。照抄会得到「α≡1 全对、有衰减就错」。
可迁移的启示:小规模、强依赖的计算在 GPU 上,算术量不是成本,依赖深度才是。判断一个改写划不划算,先数依赖链长度和每步的并行 lane 数,再数 FLOP。本文这个 kernel 多做 88% 的乘加,但把一半工作量从 32 lane 的 level-2 挪到 256 lane 的 level-3 上–这才是提速的来源。
参考
- 样本代码:Qwen team, Alibaba Group,
flash_qla 中的 tilelang_kkt_solve(MIT License)
- Gated Delta Networks:Yang, Kautz & Hatamizadeh, Gated Delta Networks: Improving Mamba2 with Delta Rule, arXiv:2412.06464,ICLR 2025。§3.3 给出 chunkwise 算法与 UT 变换
- UT 变换:Joffrain, Low, Quintana-Ortí, van de Geijn & Van Zee, Accumulating Householder Transformations, Revisited, TOMS 2006
- 分块三角求逆:Golub & Van Loan, Matrix Computations 4th ed., §3.1;cuBLAS
trtri / CUTLASS 的实现路线
- 相关文章:《线性注意力的线性代数前置知识》§4.4–4.5、《TileLang 实战:KDA 从零到一–Gated DeltaNet》§2.2、《TileLang 编程基本知识点》