TileLang 实战:单位下三角矩阵求逆
解法是分块递归:对角块的逆互不依赖、可以并行求,块间耦合退化成矩阵乘。 时依赖链从 31 步压到 17 步,代价是算术量涨到 1.88 倍; 多切一层,63 步压到 19 步。本文以 Qwen 团队 flash_qla 里的 kkt_solve 为样本,把这个 kernel 从数学恒等式一路拆到 bank conflict,包括三个不看代码想不到的工程细节,以及块长 到底该取 32 还是 64。
解法是分块递归:对角块的逆互不依赖、可以并行求,块间耦合退化成矩阵乘。C=32 时依赖链从 31 步压到 17 步,代价是算术量涨到 1.88 倍;C=64 多切一层,63 步压到 19 步。本文以 Qwen 团队 flash_qla 里的 kkt_solve 为样本,把这个 kernel 从数学恒等式一路拆到 bank conflict,包括三个不看代码想不到的工程细节,以及块长 C 到底该取 32 还是 64。
DeltaNet / KDA 这一系的 chunkwise 算法,数学上就建在三件事上:Householder 变换、WY 表示、以及把串行递推改写成闭式的 UT 变换。本文把这三层从定义推到可用于 kernel 的形式,并附两个可交互演示(广义 Householder 的 β 系数、单位下三角求逆的前代法)。后续四篇 TileLang 实战可以直接回查本文的结论。
前三篇分别实现了无遗忘的 chunked 线性注意力、标量衰减、以及带删除项的 Gated DeltaNet。本篇是系列收尾:把 GDN 的标量门 αt 换成逐通道向量门 at∈Rdk。
St=St−1Diag(at)(I−βtktkt⊺)+βtvtkt⊺
改动只有一处:αt 变成了 Diag(at)。但这一处把前三篇积累的所有便利拆掉了–衰减不再能从矩阵里外提,Γ 从 C×C 变成 C×C×dk,而累积衰减积的通道离散度会直接把 fp16 打穿。
本篇聚焦实现。KDA 的数学推导(递推式、逐通道 WY 表示、UT 变换、下界衰减与满秩门控的动机)见《KDA 的来龙去脉》§3–§4,这里不重复;本文只做一件事:把那些公式落成能跑的 kernel,并量化每一处数值边界。
前三篇见 Chunked 线性注意力、标量衰减、Gated DeltaNet。