TileLang 实战:KDA 从零到一--Chunked 线性注意力

TileLang 实战:KDA 从零到一–Chunked 线性注意力

KDA(Kimi Delta Attention)的递推式一行就写完,但直接照着它写 kernel 需要同时处理四个机制:逐通道门控、delta rule 的三角求解、log 域 cumsum、跨 chunk 状态传递。四个机制耦合在一个 kernel 里,数值出错时无法判断误差来自哪一层。 本文采用递进式实现路径,每一级只引入一个新机制、每一级都可独立运行并做数值验证。本文覆盖第一级–移除全部衰减因子与删除因子,只保留 SS+KVS \leftarrow S + K^\top V

这一级的价值不在性能,而在于将 chunkwise 分解恒等式单独隔离验证。本文同时给出一个反直觉的结论:当前实现采用的逐块独立重算策略会使 FLOPs 退回 O(N2)O(N^2),与因果 FlashAttention 同量级–线性注意力的线性复杂度在这一级尚未兑现。

数学背景见《KDA 的来龙去脉》,TileLang 语言基础见《TileLang 编程基本知识点》,本文复用的 tile 切分与 fragment 累加模式见《TileLang 实战:FlashAttention 前向 Kernel》。


1. 递进路径:把 KDA 拆成可验证的增量

KDA 的完整递推式(Dt=diag(egt)D_t = \operatorname{diag}(e^{g_t})gtR<0dkg_t \in \mathbb{R}^{d_k}_{<0}):

St=St1Dt(Iβtktkt)+βtvtkt,ot=StqtS_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

按机制拆解,每一级只放开一个自由度:

级别 递推式 新增机制 新出现的实现结构
第一级(本文) SS+KVS \leftarrow S + K^\top V 分块恒等式本身、块内因果掩码
第二级 SγS+KVS \leftarrow \gamma S + K^\top V 标量衰减 块内权重从 0/1 变 γij\gamma^{i-j} 指数下三角
第三级 Sdiag(egt)S+S \leftarrow \operatorname{diag}(e^{g_t}) \cdot S + \cdots 逐 token 门控 log 域 cumsum、exp2 硬件指令
第四级 SS(Iβkk)+βvkS \leftarrow S(I - \beta k k^\top) + \beta v k^\top delta rule UT 变换、三角求解(wy_fast 雏形)
第五级 门控 + delta rule 二者耦合 五阶段流水线、跨 kernel 状态传递

第五级即完整 GDN / KDA,结构对标 flash-linear-attentioncumsum -> chunk_scaled_dot_kkt -> wy_fast -> chunk_delta_h -> chunk_o

这样拆分的收益是误差定位能力:第三级数值不符,可以确定问题在 log 域 cumsum 或 exp2 精度,与 delta rule 无关,因为第四级尚未引入。反之若直接实现第五级,UT 变换的三角求解误差与逐通道衰减的溢出问题会互相掩盖。

一句话总结:递进式实现不是教学冗余,而是把多机制耦合 kernel 的调试搜索空间从乘法降为加法。


2. 数学推导:分块恒等式

本级递推不含遗忘项,状态单调累加:

St=St1+ktvt,ot=qtStS_t = S_{t-1} + k_t v_t^\top, \qquad o_t = q_t^\top S_t

展开为显式求和,因果且包含当前 token:

oi=jiqi(kjvj)=ji(qikj)vjo_i = \sum_{j \le i} q_i^\top (k_j v_j^\top) = \sum_{j \le i} (q_i \cdot k_j)\, v_j

这是标准线性注意力。与 softmax attention 的唯一区别是分数未经 softmax,因此求和顺序可自由交换–这是下述分块重写成立的前提。

2.1 将求和拆分为跨块与块内两部分

序列按 BCBC 切分为 NC=N/BCNC = N / BC 个 chunk。对第 cc 块内第 ii 个 token,将 jij \le i 的求和范围拆为两部分–落在此前完整块内的,与落在当前块内的:

oi=qi(c<cKcVc)跨块:与 i 无关+jchunk cji(qikj)vj块内:需因果掩码o_i = \underbrace{q_i^\top \Big( \sum_{c' < c} K_{c'}^\top V_{c'} \Big)}_{\text{跨块:与 } i \text{ 无关}} + \underbrace{\sum_{\substack{j \in \text{chunk } c \\ j \le i}} (q_i \cdot k_j) v_j}_{\text{块内:需因果掩码}}

关键在第一项:括号内的量对整个 chunk cc同一个矩阵,记作 Scprev=c<cKcVcS_c^{\text{prev}} = \sum_{c' < c} K_{c'}^\top V_{c'},形状 dk×dvd_k \times d_v,含义是处理该 chunk 之前已累积的状态。于是跨块贡献退化为一次矩阵乘 QcScprevQ_c S_c^{\text{prev}},块内贡献是一个 BC×BCBC \times BC 的下三角掩码矩阵乘:

Oc=QcScprev+tril(QcKc)VcO_c = Q_c S_c^{\text{prev}} + \operatorname{tril}(Q_c K_c^\top) V_c

这条恒等式是后续各级的公共基础,第二级至第五级都是在它的两项上分别插入衰减权重,张量形状与 GEMM 调用次序保持不变。

2.2 掩码可以直接置零的原因

FlashAttention 中掩码必须填 -\infty,因为后续要经过 expe=0e^{-\infty} = 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) # 无 exp,掩码直接清零

这一行是本级实现与 FA 的第一处分岔,也是「无 softmax」在代码上的全部体现–不需要 online 重标定、不需要 mm\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) # 每块自身的 K^T V
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) # 块内分数 [BC, BC]
mask = torch.tril(torch.ones(BC, BC, dtype=torch.bool, device=Q.device))
A = A.masked_fill(~mask, 0.0) # 线性注意力:置 0,非 -inf
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} 中的严格小于号:cumsum 给出的是含自身的前缀和,右移一格补零才是处理该 chunk 之前的累积状态。这是分块线性注意力最易出错的一行–不右移等价于把当前块的 KVK^\top V 重复计入,块内贡献会被计算两次。该错误的量级实测为相对 L2 误差 7.78×1017.78 \times 10^{-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): # grid.z ← batch
for by in range(H): # grid.y ← head
for bx in range(NC): # grid.x ← Q 块
sl = slice(bx * BC, (bx + 1) * BC)

# ① T.clear(S_f); for c in T.Pipelined(bx): T.gemm(..., transpose_A=True)
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()

# ② T.gemm(Q_s, S_s, acc_o)
Qb = Q[bz, sl, by, :].double()
acc = Qb @ S

# ③ T.gemm(Q_s,K_s,A,transpose_B=True) → 掩码 → T.gemm(A_cast,V_s,acc_o)
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 一次性算出全部 NCNC 个前缀状态、总代价 O(NC)O(NC);C 的每个 bxbx 独立重算、总代价 O(NC2)O(NC^2)kernel 采用的是 C 的策略,原因与代价见 §6。

3.4 三层参考的一致性验证

B=2,H=2,N=12,D=4,BC=4B=2, H=2, N=12, D=4, BC=4,fp64(numpy 复现):

比较 max abs 误差 相对 L2
B 分块向量化 vs A 逐 token 递归 3.55×10153.55 \times 10^{-15} 1.62×10161.62 \times 10^{-16}
C kernel 结构镜像 vs A 逐 token 递归 7.11×10157.11 \times 10^{-15} 1.86×10161.86 \times 10^{-16}
C vs B 4.44×10154.44 \times 10^{-15} 1.78×10161.78 \times 10^{-16}
参考 B 去掉右移(错误实现) 1.97×1011.97 \times 10^{1} 7.78×1017.78 \times 10^{-1}

前三行误差均在 fp64 机器精度量级(ε2.2×1016\varepsilon \approx 2.2 \times 10^{-16}),恒等式与三份实现均无误。第四行是刻意引入的错误,用于确认该验证流程对结构性错误敏感。

3.5 einsum 输出下标的约束

参考 B 中若把 torch.einsum("bhcnd,bhcnv->bhcdv", ...) 误写为 ->bhcdd(意图表达「输出是 d×dd \times d 方阵」),会直接抛出异常:

1
ValueError: einstein sum subscripts string includes output subscript 'd' multiple times

einsum 的输出下标不允许重复–重复下标在输出侧的语义是取对角线。而 KVK^\top V 的两个维度虽然长度均为 DD语义上分别是 key 维与 value 维,必须使用不同字母 dv。同理 bhcnd,bhcdd->bhcnv 也不合法,因为输出的 v 从未在输入中出现。dk=dvd_k = d_v 时长度相同掩盖了语义差异,一旦 KDA 中 dkdvd_k \ne d_v,该疏忽会立即表现为形状错误。


4. TileLang kernel:三步实现

grid 划分与 FA 一致–按 Q 块切分所有权,每个 block 负责输出一个 BC×dBC \times 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) # 状态降精度落地,喂第二个 gemm
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): # 循环上界是 runtime 的 bx
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) # S += K_c^T V_c

T.copy(S_f, S_s) # f32 → f16,本 kernel 主要精度损失点

T.Pipelined(bx) 的上界是 block 索引而非编译期常量–不同 block 的循环次数不同,第 0 块一次都不执行(S=0S = 0),最后一块需执行 NC1NC-1 次。TileLang 支持 runtime 上界的流水线,代价是各 block 负载严重不均衡,尾部 block 构成整个 kernel 的关键路径。

4.2 步骤②③:两项贡献求和

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
# ② 跨块:O = Q_c @ S_prev
T.copy(Q[bz, bx * BC:(bx + 1) * BC, by, :], Q_s)
T.clear(acc_o)
T.gemm(Q_s, S_s, acc_o)

# ③ 块内:O += tril(Q_c K_c^T) @ V_c
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) # f32 → f16 喂 MMA
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 中残留的是第 bx1bx-1 块的数据,且在 num_stages > 1 时具体残留哪一块取决于流水线的展开方式,不可假设。复用 shared buffer 节省了显存,但必须显式重载–这是 tile 编程中典型的隐式状态陷阱。参考 C 里没有这个问题,因为 Python 每轮都重新切片,不存在缓冲区复用。

T.copy(A, A_cast) 这次 f32→f16 转换与 FA 中 P~\tilde{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,维护 mm\ell 两个行状态,每块重标定 无,掩码直接置 0 exp,求和顺序可交换
循环携带的量 无跨块状态,每个 Q 块独立 S[dk,dv]S[d_k, d_v] 是跨迭代累加的状态 递推式本身带状态
T.gemm policy 必须 FullRow–行归约要求整行在同一 warp 默认 Square 即可 无按行归约

第三点值得展开:FA 需对分数块做 rowmax / rowsum,若一行被切分到多个 warp,归约就需跨 warp 通信,因此必须用 FullRow policy 强制整行不拆。本级的块内矩阵 AA 计算完成后仅做逐元素掩码,无任何跨列归约,warp 划分方式不影响正确性,编译器可自由选择寄存器分布最优方案。

第二点是从本级走向第五级的主要矛盾。当前实现用逐块独立重算把状态依赖隐藏了,但第五级的 chunk_delta_h 必须真正在 chunk 之间传递状态–届时状态需在 fragment(累加)与 shared(作为下一个 GEMM 输入)之间每轮拷贝一次,并需拆成两个 kernel 协作。


6. 代价账本:O(N2)O(N^2) 仍在,线性复杂度尚未兑现

逐块独立重算的冗余可以精确计算。第 bxbx 块执行 bxbxKVK^\top V,全部 block 合计 c=0NC1c=NC(NC1)/2\sum_{c=0}^{NC-1} c = NC(NC-1)/2 次,而理想情况每块只需计算一次、共 NCNC 次(即参考 B 的 cumsum 策略)。冗余系数为 (NC1)/2(NC-1)/2

单次 KVK^\top V 的 FLOPs 为 2BCdkdv2 \cdot BC \cdot d_k \cdot d_v,取 dk=dv=Dd_k = d_v = D,步骤①的总量:

FLOPs=NC(NC1)22BCD2N2D2BC\text{FLOPs}_① = \frac{NC(NC-1)}{2} \cdot 2\, BC\, D^2 \approx \frac{N^2 D^2}{BC}

对比因果注意力的 2N2D2N^2 D(两个 GEMM 各占一半三角),比值为:

FLOPsFLOPscausal FA=D2BC\frac{\text{FLOPs}_①}{\text{FLOPs}_{\text{causal FA}}} = \frac{D}{2\,BC}

该比值与 NN 无关–即本级实现的计算量与因果 FlashAttention 属同一量级,仅差一个由 D/BCD / BC 决定的常数。实测账本(D=64D = 64BC=64BC = 64):

NN NCNC 冗余系数 步骤① 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

N=16384N = 16384 时冗余系数达 127.5 倍,而步骤②③合计始终是 O(N)O(N)N=8192N = 8192 时分别为 0.067 与 0.134 GFLOP,比步骤①小两个数量级)。结论是:本级实现 99% 以上的开销耗在冗余状态重算上,线性注意力的 O(N)O(N) 复杂度在这一级完全没有兑现。

BCBC 的取舍随之明确:

BCBC D/(2BC)D/(2BC) 含义
32 1.000 与因果 FA 计算量持平,块过小无收益
64 0.500 默认值,寄存器压力可控
128 0.250 冗余减半,但 AA128×128128 \times 128 f32 fragment,每线程约 128 个寄存器,易溢出
256 0.125 理论最省,实际 shared memory 与寄存器均无法容纳

增大 BCBC 能线性减少冗余,但 AABC×BCBC \times BC 的 f32 fragment,BC=128BC = 128 时寄存器压力已接近溢出边界。真正的解法不是调 BCBC,而是替换逐块独立重算策略–改用两个 kernel 协作,第一个 kernel 按参考 B 的方式算出所有 chunk 的前缀状态写回 HBM,第二个 kernel 读取。这正是第五级 chunk_delta_h 的结构,也解释了官方实现为何拆成五个 kernel 而非一个。

一句话总结:本级用 (NC1)/2(NC-1)/2 倍冗余计算换取单 kernel 闭环与零跨 block 同步,该取舍在教学上成立、在生产上不成立–它把线性注意力的核心优势整体交换掉了。


7. 数值验证

四层验证的组织方式:

  1. 参考 A vs B vs C(均 fp64):验证 §2.1 分块恒等式与 kernel 控制流,与 GPU 无关,误差应在 101510^{-15} 量级;
  2. 刻意引入的错误实现:确认验证流程对结构性错误敏感(去掉右移,相对 L2 应达 10110^{-1} 量级);
  3. kernel vs 参考 C:验证 TileLang 实现,fp16 输入 + f32 累加,阈值取相对 L2 <2×102< 2 \times 10^{-2}
  4. 延迟对比走 CUDA event 中位数,取 50 次采样。

第 3 层阈值定在 2×1022 \times 10^{-2} 而非更严,原因是 T.copy(S_f, S_s) 把 f32 状态降至 f16 落 shared。该降精度是必要的–Tensor Core MMA 的输入必须是低精度–但状态 SS 是多个块累加的结果,越靠后的 block 累加项越多,误差随 bxbx 单调增长。这是本级精度的主导误差源,影响远大于块内 AA 的那次降精度。

第 1、2 层验证已在 numpy fp64 上完成(见 §3.4 表格)。第 3、4 层需要 CUDA 设备,本文未给出实测数字–待真卡跑通后单独补充实测数据,此处不做性能推测。

运行方式:

1
2
python gdn_ladder_L1_linear_attn.py --seq_len 512 --dim 64 --blk 64
python gdn_ladder_L1_linear_attn.py --seq_len 2048 --blk 128 --bench

使用 parse_known_args() 而非 parse_args():Jupyter / IPython 会向 sys.argv 注入 kernel 连接参数,用 parse_args() 会直接报错退出。


8. 总结

  1. 本级移除了 KDA 的全部衰减因子与删除因子,只保留 SS+KVS \leftarrow S + K^\top V,目的是将 chunkwise 分解恒等式 Oc=QcScprev+tril(QcKc)VcO_c = Q_c S_c^{\text{prev}} + \operatorname{tril}(Q_c K_c^\top) V_c 单独隔离验证。该恒等式是后续四级的公共基础,每级只在它的两项上插入衰减权重,GEMM 的形状与调用次序不变。
  2. 三层 PyTorch 参考构成从数学到 kernel 的完整链条:A 逐 token 递归验证递推式定义,B 分块向量化验证恒等式,C 逐块独立重算镜像 kernel 控制流。kernel 与 C 的差异仅剩两项–f32→f16 显式降精度与 shared buffer 中转,这也是全部数值差异的来源。
  3. 与 FA 的差异集中在三点:无 online softmax(掩码置 0 而非 -\infty)、循环携带 dk×dvd_k \times d_v 状态、无按行归约(T.gemm policy 用默认 Square 即可)。
  4. 实测数字:三层 fp64 参考互验的相对 L2 误差均在 1.61.9×10161.6\text{--}1.9 \times 10^{-16},恒等式无误;刻意去掉右移的错误实现相对 L2 为 7.78×1017.78 \times 10^{-1},验证流程对结构性错误敏感。代价账本显示 N=16384N = 16384 时逐块独立重算的冗余系数达 127.5 倍,步骤① FLOPs 是因果 FlashAttention 的 0.498 倍,且该比值 D/(2BC)D/(2BC)NN 无关–线性复杂度在本级尚未兑现。
  5. 两个实现陷阱:einsum 输出下标不可重复(->bhcdd 会抛异常,key 维与 value 维必须用不同字母);T.Pipelined 之后 shared buffer 的残留内容取决于流水线展开方式,复用前必须显式重载。

下一篇加入标量衰减 SγS+KVS \leftarrow \gamma S + K^\top V。改动集中在块内那一项–下三角掩码从 0/1 矩阵变为 γij\gamma^{i-j} 的指数权重,跨块那一项则需给 SS 乘上整块的衰减 γBC\gamma^{BC}γ\gamma 一旦从标量升级为逐通道向量,指数外提就会引入 1/Γ1/\Gamma 的数值溢出,那是这条实现路径上第一个真正棘手的问题。

可迁移的启示:多机制耦合的 kernel 不应一次写完。每级只放开一个自由度,并且每级都保留一份可独立运行的 fp64 参考–其中至少一份要与 kernel 的控制流同构。参考实现的成本远低于在五个机制中定位一个误差源的成本。


参考

本文的分块恒等式验证与 FLOPs 账本在 numpy fp64 上复现;GPU 实测数据待补。