TileLang 实战:单位下三角矩阵求逆

解法是分块递归:对角块的逆互不依赖、可以并行求,块间耦合退化成矩阵乘。C=32C = 32 时依赖链从 31 步压到 17 步,代价是算术量涨到 1.88 倍;C=64C = 64 多切一层,63 步压到 19 步。本文以 Qwen 团队 flash_qla 里的 kkt_solve 为样本,把这个 kernel 从数学恒等式一路拆到 bank conflict,包括三个不看代码想不到的工程细节,以及块长 CC 到底该取 32 还是 64。


1. 问题从哪来

DeltaNet 系列(DeltaNet / Gated DeltaNet / KDA)的 chunkwise 形式里,块内 CC 个 Householder 变换连乘会被 UT 变换整理成一个三角系统。系数矩阵长这样:

A=I+strictLower(diag(β)KK)RC×C\mathbf{A} = \mathbf{I} + \operatorname{strictLower}\big(\operatorname{diag}(\beta)\,\mathbf{K}\mathbf{K}^{\intercal}\big) \in \mathbb{R}^{C\times C}

K\mathbf{K} 是块内 CC 个 key 排成的矩阵(C×dkC \times d_k,论文里 k\bm{k} 做过 L2 归一化),βt(0,1)\beta_t \in (0,1) 是写入强度。要求的是 T=A1\mathbf{T} = \mathbf{A}^{-1},GDN 论文 (§3.3) 用它构造两个量:

W[t]=Tdiag(β)K,U~[t]=Tdiag(β)V\mathbf{W}_{[t]} = \mathbf{T}\operatorname{diag}(\beta)\mathbf{K}, \qquad \widetilde{\mathbf{U}}_{[t]} = \mathbf{T}\operatorname{diag}(\beta)\mathbf{V}

推导细节见《线性注意力的线性代数前置知识》§4 与《TileLang 实战:KDA 从零到一–Gated DeltaNet》§2.2,本文只关心一件事:这个 T\mathbf{T} 怎么在 GPU 上算得快

论文 §3.3 的标准形式里 T\mathbf{T} 被用两次(W\mathbf{W}U~\widetilde{\mathbf{U}}),不物化就得做两次前代、两条串行链。kkt_solve 于是选择物化:一个独立 kernel 把 T\mathbf{T} 算出来写回 HBM,下游两处都退化成 T.gemm。(对照组:GDN 那篇 §4.2 的右端项只有一个,于是把前代直接作用在它行上,不物化 T\mathbf{T}。)

既然要物化,就得正面回答「怎么快速求一个单位下三角矩阵的逆」。


2. 两种算法

2.1 逐格前代:O(n3/6)O(n^3/6)nn 步串行

X=L1\mathbf{X} = \mathbf{L}^{-1} 即解 LX=I\mathbf{L}\mathbf{X} = \mathbf{I}。因为 L\mathbf{L} 单位对角,X\mathbf{X} 也是单位下三角–对角上的 1 和上三角的 0 不用算,直接就是答案。真正要算的只有严格下三角部分。

按行写出递推(ek\bm{e}_k 是第 kk 个标准基向量):

X[k,:]=ekj<kL[k,j]X[j,:]\mathbf{X}[k,:] = \bm{e}_k^{\intercal} - \sum_{j<k}\mathbf{L}[k,j]\,\mathbf{X}[j,:]

算第 kk 行要等前 k1k-1 行全部算完,nn 行就是 n1n-1 步串行。

2.2 分块递归:把长链换成短链 + GEMM

先看 C=32C = 32:把 L\mathbf{L}2×22\times2 切块(每块 16×1616\times16):

L=[L000L10L11]L1=[L0010L111L10L001L111]\mathbf{L} = \begin{bmatrix}\mathbf{L}_{00} & \mathbf{0}\\ \mathbf{L}_{10} & \mathbf{L}_{11}\end{bmatrix} \qquad\Longrightarrow\qquad \mathbf{L}^{-1} = \begin{bmatrix} \mathbf{L}_{00}^{-1} & \mathbf{0}\\ -\mathbf{L}_{11}^{-1}\mathbf{L}_{10}\mathbf{L}_{00}^{-1} & \mathbf{L}_{11}^{-1} \end{bmatrix}

验证只需一次分块乘法:左下块 L10L001+L11(L111L10L001)=0\mathbf{L}_{10}\mathbf{L}_{00}^{-1} + \mathbf{L}_{11}\cdot(-\mathbf{L}_{11}^{-1}\mathbf{L}_{10}\mathbf{L}_{00}^{-1}) = \mathbf{0}

价值全在依赖结构上:

  • L001\mathbf{L}_{00}^{-1}L111\mathbf{L}_{11}^{-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=64C = 64:递归一层不够,切四份。 上面的 2×22\times2 只是 C=32C = 32 的形态。C=64C = 64 时对角块还是 32×3232\times32,本身仍需 31 步前代–所以要再往下递归一层,最终落到 4×44\times416×1616\times16 网格:

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×1616\times16 前代 15 4 4×16=644\times16 = 64
2 组 M1\mathbf{M}^{-1} 耦合,各 2 次 16316^3 GEMM +2 2 2×256=5122\times256 = 512
1 组块间耦合,2 次 32332^3 GEMM +2 1 32×32=102432\times32 = 1024
合计 19

M001\mathbf{M}_{00}^{-1}M111\mathbf{M}_{11}^{-1} 在阶段 ①② 全程互不相干,两条链完全并行;阶段 ③ 才把它们缝起来。依赖链 63 → 19,而每一步的可并行宽度从 32 涨到 1024CC 翻倍时 GEMM 段的规模按 C3C^3 膨胀,正好把大 chunk 多出来的算术量塞进 Tensor Core,不落在关键路径上。

值得对比的是另一条路:扁平的块前代–把 4×44\times4 网格当成一个 4×44\times4 的“标量”下三角矩阵,直接跑 §2.1 的行式递推,每个“元素”换成 16×1616\times16 GEMM:

C=64C = 64 方案 FMA 依赖链 16316^3 GEMM 次数
整块逐格前代 41664 63
扁平块前代(4×44\times4 网格) 67776 21 16
两层分块递归 84160 19 4 + 2 组 32332^3

扁平块前代省 20% 算术量,但块行之间仍是串行的(第 3 块行要等第 2 块行),依赖链 21 步且后期块行的并行度递减。分块递归多花那 20%,换来的是两个 32×3232\times32 子问题的完全独立–在 TileLang 里这直接对应两组互不同步的 warp。C=128C = 128 时差距进一步拉开:扁平块前代 29 步、分块递归 21 步。

2.3 两种算法的账

C=32C = 32、递归到 16×1616\times16 为止:

逐格前代(整块 32) 分块递归(2×16)
串行依赖链 31 步 15 步 + 2 次 GEMM
数学必需 FMA 4960 2×560+2×163=93122\times560 + 2\times16^3 = 9312
算术量比 1.00 1.88
前代段并行 lane 32 32(2 块 × 16 列)
GEMM 段并行 lane 256
主体算子 AXPY(level-2) GEMM(level-3)
额外存储 一个 16×1616\times16 中间量

多花 88% 的算术,换依赖链减半、且一半的工作量跑在 level-3 上。 小规模三角求解本来就填不满 SM,算术单元闲着,这笔账划算。

递归可以继续往下切:

CC 递归层数 对角 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×1616\times16 在做前代),每多一层递归只增加两次串联 GEMM 的深度。所以依赖链是 O(logC)O(\log C) 而非 O(C)O(C),而新增的 GEMM 段内部完全可并行。

2.4 chunk 级别的账:CC 该取 32 还是 64

上面是 kernel 内部的账,但选块长要放到整条 chunkwise 流水线里看。求逆的算术量按 O(C3)O(C^3) 涨、而它服务的 token 只有 CC 个,摊到每 token 是 O(C2)O(C^2)

CC 分块递归 FMA 摊到每 token 主 GEMM 每 token 求逆占比
32 9312 291 40960 0.7%
64 84160 1315 49152 2.6%
128 692608 5411 65536 7.6%

(主 GEMM 口径:块内 QK\mathbf{Q}\mathbf{K}^\intercal + AV\mathbf{A}\mathbf{V} 加块间 QS\mathbf{Q}\mathbf{S} + 状态更新,即 2dkdv+C(dk+dv)2d_kd_v + C(d_k+d_v),取 dk=dv=128d_k = d_v = 128。)

另一头,块长变大能省的是状态传递的 HBM 流量–每个 chunk 只读写一次 dk×dvd_k \times d_v 状态,摊到每 token 是 22dkdv/C2\cdot 2d_kd_v/C 字节:

CC state 读写 T\mathbf{T} 写回 一张 C2C^2 fp32 中间表 cond(A)\operatorname{cond}(\mathbf{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

三条线交叉出来的结论:

  • CC 从 32 到 64:state 流量减半(2048 → 1024 B/token,这是 chunkwise 最大的一笔片外开销),代价是求逆占比从 0.7% 涨到 2.6%、C2C^2 表从 4 KB 涨到 16 KB。这笔交换通常划算,所以 GDN/KDA 那几篇的主 kernel 都取 C=64C = 64
  • CC 从 64 到 128:state 流量只再省 512 B,但 QK\mathbf{Q}\mathbf{K}^\intercalΓ\GammaA\mathbf{A} 这些 C2C^2 表同时膨胀到 64 KB 一张–单个 SM 的 shared memory 就装不下几张,寄存器也会超限(GDN 篇 §4.1 实测 Gam + A + Amat 三张表在 C=128C = 128 时直接爆 reg)。上限卡在 64。
  • kkt_solve 单独取 C=32C = 32,是因为它是独立 kernelT\mathbf{T} 要写回 HBM 再被下游读回来,这笔 2C2C B/token 的流量在融合 kernel 里是不存在的。CC 翻倍它也翻倍,而 T\mathbf{T} 只是个中间量–省不了 state 流量、只多付带宽。物化路线的最优块长因此比融合路线更小。

一句话:融合 kernel 里 CC 由 shared/reg 容量顶到 64;物化 T\mathbf{T} 的独立 kernel 里 CC 由中间量写回带宽压回 32。 同一个算法,落在流水线的不同位置,最优块长不一样。

2.5 到底是谁卡住了 chunk size

上一节说“容量顶住了”,但究竟是寄存器还是 shared memory,取决于 kernel 怎么写。TileLang 里常见的三种写法,瓶颈完全不在一个地方。

写法 A:单 warp-group(128 线程),中间量常驻 fragment。这是 GDN/KDA 前几篇的写法,C×CC\times CΓ\GammaA\mathbf{A}Amat\mathbf{A}^{\text{mat}} 都放在寄存器里。128 线程分一张 C2C^2 fp32 表,每线程 C2/128C^2/128 个 reg:C=64C = 64 是 32 reg/表,C=128C = 128 就是 128 reg/表–三张表 384 reg,直接超过 255 的硬上限。卡在寄存器,上限 64。

写法 B:warp-specialized(512 线程),双缓冲过 shared。生产者 warp 搜 HBM、消费者 warp 做 MMA,中间量全走 shared memory。以 dk=128d_k = 128dv=blockDV=64d_v = \text{block}_{DV} = 64、bf16 为例:

缓冲区 随什么长 C=32C = 32 C=64C = 64
q_shared / k_shared(双缓冲 2Cdk2Cd_k C\propto C 16 + 16 KB 32 + 32 KB
v_shared(双缓冲 2Cdv2Cd_v C\propto C 8 KB 16 KB
o / h / vd / vn(各 CdvCd_v C\propto C 16 KB 32 KB
a_shared(双缓冲 2C22C^2 C2\propto C^2 4 KB 16 KB
p_sharedC2C^2 C2\propto C^2 2 KB 8 KB
门控表 + barrier O(C)O(C) ~1 KB ~1 KB
合计 63 KB 137 KB

这张表里有个容易看反的地方:大头是 C\propto C 的项而不是 C2C^2C=64C = 64 时线性项 112 KB、平方项才 24 KB。因为 dk=128>Cd_k = 128 > C,Q/K 的双缓冲才是大户。结果是容量几乎随 CC 线性翻倍,很快撞上硬件墙:

架构 shared/SM C=32C = 32(63 KB) C=64C = 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=64C = 64 也只剩 1 block/SM–没有第二个 block 掩盖尾巴和同步空档,而 warp-specialized 本来就是靠内部流水掩盖延迟的。想保 C=64C = 64 又想多一个 block,只能去掉双缓冲(降到 89 KB)–而那恰恰把 warp-specialized 的意义抵消了。卡在 shared memory,且消费级卡比数据中心卡卡得更早。

写法 C:独立物化 kernel(就是 kkt_solve)。它只吃 K\mathbf{K}、只吐 T\mathbf{T},shared 就一张 C×dkC\times d_k 加几张 16×1616\times16,容量根本不是问题。卡它的是 §2.4 说的 T\mathbf{T} 写回 HBM 的 2C2C B/token。

汇总成一张表:

A:单 warp-group 融合 B:warp-specialized 融合 C:独立物化
线程数 128 512 128
中间量住哪 fragment shared(双缓冲) fragment
主瓶颈 寄存器 255/线程 shared/SM 容量 中间量 HBM 带宽
CC 增长 C2/128C^2/128 reg/表 近似线性(因 dk>Cd_k > C 线性 2C2C B/token
实际上限 64 64(且 1 block/SM);sm89 上 32 32
次要约束 串行前代并行度 生产/消费 warp 配比 求逆 O(C2)O(C^2)/token

三种写法的最优 CC 差不多都落在 32~64,但卡住它们的是三种不同的资源。换硬件、换 dvd_v、多加一张中间表,哪个先爆都会变–所以调 CC 之前先弄清楚自己卡在哪一条线上,别拿别人的数字当结论。


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                                    # chunk_size,硬编码
k_shared = T.alloc_shared((block_S, DK), dtype=qkva_dtype) # 8 KiB (fp16, DK=128)
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) # 4 KiB

# 2 个对角 16×16 块(padding 到 17,见 §3.4)
a16i_shared = T.alloc_shared((2, 17, 16), dtype=accum_dtype) # 2176 B
# 1 个次对角 16×16 耦合块
a16o_shared = T.alloc_shared((1, 17, 16), dtype=accum_dtype) # 1088 B
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"、输出 ak.dtype(fp16)–单位下三角阵行列式恒为 1、不需 pivoting,精度这里给的余量很大(§4 有实测)。

3.2 构造系数矩阵

1
2
3
4
5
6
7
8
9
10
11
# A = K @ K^T
T.gemm(k_shared, k_shared, a32_fragment, transpose_B=True, clear_accum=True)

# A = diag(β) A
for j_s, j_t in T.Parallel(block_S, block_S):
a32_fragment[j_s, j_t] *= b_shared[j_s]

# A = I + strictLower(A)
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

β\beta 乘在上(b_shared[j_s]),对应 diag(β)KK\operatorname{diag}(\beta)\mathbf{K}\mathbf{K}^\intercal 的左乘。对角必须置 1 而不是保留 βiki2=βi\beta_i\|\bm{k}_i\|^2 = \beta_i–UT 变换的系数矩阵是 I+strictLower()\mathbf{I} + \operatorname{strictLower}(\cdot),含对角就变成解 (I+Afull)(\mathbf{I}+\mathbf{A}_{\text{full}})βi\beta_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] # ← -L10
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-\mathbf{L}_{10} 而不是 L10\mathbf{L}_{10}。恒等式里的负号被提前吃掉,后面两次 GEMM 直接得到 L111L10L001-\mathbf{L}_{11}^{-1}\mathbf{L}_{10}\mathbf{L}_{00}^{-1},省一遍 16216^2 的取反。这种「负号往前提」的技巧在 kernel 里很常见–取反本身不贵,但它是一个独立的 pass,要多读写一次 shared。

3.4 对角块并行前代

1
2
3
4
5
6
7
8
9
10
11
for k_s in T.unroll(1, 16):                       # 15 步依赖链
for j_s, k_t in T.Parallel(2, 16): # 2 块 × 16 列 同时算
if k_t < k_s:
a16i_row[j_s, k_t] = a16i_shared[j_s, k_s, k_t] # 备份原始 L[k_s, :]
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\mathbf{L} 和输出 X\mathbf{X} 都是单位下三角,对角的 1 和上三角的 0 两边完全一样,所以只需改写严格下三角部分(k_t < k_s),不用额外开一块 buffer。

a16i_row 不是优化,是必需。ksk_s 行马上要被结果覆盖,但它的原始值 L[ks,j]\mathbf{L}[k_s, j] 还要当递推系数用。不备份就会读到刚写进去的新值,算出来的东西静默错误–这是最容易踩的坑,因为它不报错。

两个块并行是分块递归的收益兑现处。 T.Parallel(2, 16) 的 32 条 lane 里,j_s=0 处理 L00\mathbf{L}_{00}j_s=1 处理 L11\mathbf{L}_{11},两条 15 步的链同时走完。对比整块 32 要走 31 步、每步同样 32 条 lane–并行度没变,时间减半

工程点二:(2, 17, 16) 里的 17。 fp32 下 shared memory 有 32 个 bank,每 bank 4 字节。若声明成 (2, 16, 16),两个块的起始偏移差是 16×16=25616\times16 = 256,而 256mod32=0256 \bmod 32 = 0T.Parallel(2, 16) 展开成 32 条 lane 时,j_s=0 的第 ii 条和 j_s=1 的第 ii落在同一个 bank 上,32 条 lane 两两冲突,访存吞吐直接减半。

改成 17 之后偏移差变成 17×16=27217\times16 = 272272mod32=16272 \bmod 32 = 16,两个块正好错开半个 bank 区,32 条 lane 无冲突:

声明 块间偏移 mod 32\bmod\ 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
# Step 1: temp = L11⁻¹ @ (-L10)
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])

# 写回 shared 时转置
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]

# Step 2: result = temp @ L00⁻¹
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 要算 tempL001\text{temp}\cdot\mathbf{L}_{00}^{-1},按 krk_r 累加时需要读 temp[:, k_r] 这一。列方向在行主序里跨度是 16(或 padding 后 17),16×1616\times16 的 lane 会散在多个 bank 上。先转置写回 shared,读取就变成沿最内维连续–配合 §3.4 的 17-padding,两处优化是配套的,只做一个另一个就白费。

注意两次 GEMM 的并行度是 16×16=25616\times16 = 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] # 左上 L00⁻¹
a32_shared[k_s, 16 + k_t] = 0 # 右上 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] # 右下 L11⁻¹

(原代码写成四个独立的 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() # a16i_row 备份
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

结果:

用例 误差(\ell_\infty
16×1616\times16 行式前代 2.22×10152.22\times10^{-15}
32×3232\times32 分块递归 4.44×10144.44\times10^{-14}
真实 A\mathbf{A}(200 组采样) 1.39×10161.39\times10^{-16}
fp16 存储 T\mathbf{T} 的相对误差 5.72×1055.72\times10^{-5}

真实场景误差比一般随机矩阵低两个数量级,原因是 k\bm{k} 的 L2 归一化 + β(0,1)\beta \in (0,1) 把严格下三角元素压到 0.358 以内,cond(A)\operatorname{cond}(\mathbf{A}) 实测不到 2.2。

两个必做的退化检验

  1. β0\beta \to 0AI\mathbf{A} \to \mathbf{I}T\mathbf{T} 应该是单位阵。这检验「对角置 1」和拼装时的清零。
  2. 只有次对角块非零(L00=L11=I\mathbf{L}_{00} = \mathbf{L}_{11} = \mathbf{I}):T\mathbf{T} 的左下块应恰好是 L10-\mathbf{L}_{10}。这单独检验耦合路径和那个负号,其他情况下负号错了会被对角块的贡献掩盖。

5. 总结

  1. 单位下三角求逆的瓶颈是依赖链,不是算术量C=32C = 32 只需 4960 次乘加,但逐格前代要走 31 步串行,每步只有 32 条 lane 可并行(恰好一个 warp),threads=128 时剩下 96 条线程全程闲置。
  2. 分块递归把长链换成短链 + GEMM,递归深度随 CCC=32C = 32 切一层(2×22\times216×1616\times16)就到底;C=64C = 64 要切两层、落到 4×44\times4 网格,且两个 32×3232\times32 子问题全程独立。依赖链 31 → 17(63 → 19),算术量 1.88 倍,而并行宽度从 32 涨到 1024。
  3. 前代段恒定 15 步,依赖链是 O(logC)O(\log C)。每多一层递归只加两次串联 GEMM(C=32/64/128C=32/64/128 → 17/19/21 步),而 GEMM 段可并行。对比“扁平的块前代”(把网格当标量下三角跑行式递推):C=64C=64 时它省 20% 算术量但依赖链 21 步且后期并行度递减。所以块长不是被求逆算法卡住的。
  4. CC 取 32 还是 64,取决于这个 kernel 在流水线的哪个位置(§2.4)。求逆摊到每 token 是 O(C2)O(C^2)C=64C=64 时只占主 GEMM 的 2.6%,而 state 的 HBM 流量从 2048 降到 1024 B/token–融合 kernel 里这笔交换划算。但 kkt_solve 是独立 kernel,T\mathbf{T} 要写回 HBM 再读回来(2C2C B/token 随 CC 线性涨),省不到 state 流量、只多付带宽,于是压回 32。
  5. 卡住 chunk size 的资源因写法而异(§2.5)。单 warp-group + fragment 卡寄存器(一张 C2C^2 表 = C2/128C^2/128 reg/线程,C=128C=128 时三张表 384 reg 超限);warp-specialized + 双缓冲卡 shared memoryC=32C=32 → 63 KB、C=64C=64 → 137 KB,后者在 sm89 直接跑不起来、在 H100 上也只剩 1 block/SM);独立物化 kernel 卡中间量带宽。值得注意的是 shared 里大头是 C\propto C 的 Q/K 双缓冲而非 C2C^2 表(dk=128>Cd_k = 128 > C),所以容量近似线性翻倍。
  6. 原地覆盖免费,但要备份当前行L\mathbf{L}X\mathbf{X} 的对角 1 和上三角 0 完全一致,只需改写严格下三角;代价是第 ksk_s 行的原始值必须先存进 a16i_row,否则读到已覆盖的新值–静默出错,不报错
  7. (2, 17, 16) 的 17 是 bank conflict 解药(2,16,16) 时两块偏移差 2560(mod32)256 \equiv 0 \pmod{32}T.Parallel(2,16) 的 32 条 lane 两两撞 bank;17 使偏移差 27216272 \equiv 16,错开半区。必须配 annotate_layout(make_linear_layout),否则 TileLang 换成 swizzle 会让 padding 白做。
  8. 两次 GEMM 中间的转置不是多余。Step 2 要读中间量的列,转置后变成沿最内维连续,与 17-padding 是配套优化,缺一个另一个失效。
  9. 负号提前吃掉。提取次对角块时直接存 L10-\mathbf{L}_{10},省一遍 16216^2 取反 pass。
  10. 物化 T\mathbf{T} 与否取决于它被用几次。只用一次就把前代直接作用在右端项上;GDN §3.3 的 W\mathbf{W}U~\widetilde{\mathbf{U}} 两处复用,才值得单独起 kernel。
  11. 这个问题良态得反常k\bm{k} 的 L2 归一化 + β(0,1)\beta \in (0,1)cond(A)2.18\operatorname{cond}(\mathbf{A}) \le 2.18(200 组采样),fp32 累加、fp16 存储都够。但换个来源的三角矩阵要重算–n=128n=128 的一般随机三角矩阵条件数是 10710^7 量级。
  12. Γ\Gamma 不在这个 kernel 里。带门控时系数是 diag(β)(ΓKK)\operatorname{diag}(\beta)(\Gamma\odot\mathbf{K}\mathbf{K}^\intercal),折叠需要两份缩放不同的 K,而 T.gemm(k_shared, k_shared, ...) 左右同一份。照抄会得到「α1\alpha\equiv1 全对、有衰减就错」。

可迁移的启示:小规模、强依赖的计算在 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 编程基本知识点