TileLang 实战:FlashAttention 前向 Kernel

TileLang 实战:FlashAttention 前向 Kernel

FlashAttention 没有发明新的数学。它是 online softmax 递推与两个 GEMM 在同一个 tile 循环里的交织–分数块 SS 算出来就地做 softmax 得到 P~\tilde{P}P~\tilde{P} 就地乘上 VV,中间矩阵全程驻留在 SRAM,HBM 流量从 O(N2)O(N^2) 降到 O(Nd)O(N \cdot d)。本文用 TileLang(v0.1.13)把 FA-2 论文的 Algorithm 1 写成可运行的 kernel:先补齐数学地基,再讲清楚"怎么切",然后逐行拆解主循环,最后做数值验证、测速与调参。只覆盖前向;TileLang 五原语与 T.Pipelined 流水线机制见上一篇《TileLang 编程基本知识点》,本文直接使用其结论。


一、为什么标准 Attention 慢:IO 才是瓶颈

标准 scaled dot-product attention 的流程是 S=QK/dP=softmax(S)O=PVS = QK^\top / \sqrt{d} \to P = \mathrm{softmax}(S) \to O = PV。计算本身没有问题,问题在中间矩阵:SSPP 都是 N×NN \times N,必须写进显存再读回来。N=8192N = 8192 时单个 fp32 矩阵就是 256 MB,长序列下这个开销直接失控。

GPU 内存层级(A100 量级):

层级 容量 带宽 延迟量级
HBM(显存) 40~80 GB ~1.5-2 TB/s ~400-800 cycle
SRAM(每 SM) 164 KB ~19 TB/s ~20-30 cycle
寄存器(每 SM) 256 KB ~19 TB/s 接近 0

SRAM 比 HBM 快 10 倍以上,但只有一百多 KB。attention 的 FLOPs 是 O(N2d)O(N^2 d),标准实现与 FlashAttention 完全相同–快慢差别全部来自中间矩阵走没走 HBM

指标 标准 Attention FlashAttention
中间矩阵内存 O(N2)O(N^2),S/P 落显存 O(N)O(N),S/P̃ 留在 SRAM
HBM 往返 3 次(写 S → 读写 P → 读 P 做 PV) 1 次(读 Q/K/V → 写 O)
FLOPs O(N2d)O(N^2 d) 相同
精确性 精确 精确(非近似)

这就是论文标题里 “IO-aware” 的含义:不省计算,省搬运。

一句话总结:attention 是访存受限的算子,FlashAttention 的全部收益来自让 N×NN \times N 的中间矩阵不落 HBM。


二、数学地基:softmax 的三级台阶

2.1 平移不变性 → 安全 softmax

softmax 对输入加任意常数 cc 不变(分子分母的 ece^c 约掉)。这个自由度是数值安全的救命稻草:fp16 上限只有 65504,s=100s = 100e1001043e^{100} \approx 10^{43} 直接 inf,inf/inf 变 NaN。减去行最大值 mm 后指数上限恰为 0:

softmax(xi)=eximjexjm,m=maxjxj\mathrm{softmax}(x_i) = \frac{e^{x_i - m}}{\sum_j e^{x_j - m}}, \quad m = \max_j x_j

代价是计算顺序被强制:先扫一遍求 mm,再扫一遍求分母与输出–两遍扫描。

2.2 分块场景下,"两遍"成为灾难

SRAM 装不下整行 K/V,只能按块流入。设处理完第 1 块时按基准 m=3m=3 攒好部分和;第 2 块冒出更大分数,基准换成 m=7m=7–旧 exp 全部基准错误。三条路:

方案 做法 代价
存下全部 S 等全局 max 再统一归一化 N×NN \times N 落显存,爆
两遍重扫 pass1 求 mm\ell;pass2 重算加权 K/V 从 HBM 读两次,带宽 ×2
online softmax 边收块边维护"以当前 max 为基准"的部分和 一遍过,数学无损

2.3 online softmax 递推

每个 query 行维护三个状态,初值 m=m = -\infty=0\ell = 0O=0O = 0。新块 SnewS_{new} 到来时:

mnew=max(mold, rowmax(Snew))α=emoldmnew旧成果的打折系数new=αold+rowsum(eSnewmnew)Onew=αOold+eSnewmnewVnew\begin{aligned} m_{new} &= \max(m_{old},\ \mathrm{rowmax}(S_{new})) \\ \alpha &= e^{m_{old} - m_{new}} && \leftarrow \text{旧成果的打折系数} \\ \ell_{new} &= \alpha \cdot \ell_{old} + \mathrm{rowsum}(e^{S_{new} - m_{new}}) \\ O_{new} &= \alpha \cdot O_{old} + e^{S_{new} - m_{new}} \cdot V_{new} \end{aligned}

循环结束后做全程唯一一次归一化:Output=O/\mathrm{Output} = O / \ell

为什么无损:缩放因子连乘时指数项逐项相消,em1m2em2m3=em1m3e^{m_1 - m_2} \cdot e^{m_2 - m_3} = e^{m_1 - m_3},循环结束时每项都被精确换算到最终 max 的基准下。中间任何时刻 OO 都是合法但未归一化的加权和,数值永远有界(当前最大项的 exp 恰为 1)。

数值走一遍(一行 query,三块各来一个分数 1、2、3,正确答案 = softmax(1, 2, 3) 的权重):

1
2
3
4
5
6
7
8
init : m=−∞   ℓ=0      O=0
j=1 : m=1, α=e^(−∞)=0
O = 0×0 + e^0·v₁ = 1.000·v₁ ℓ = 1.000
j=2 : m=2, α=e^(1−2)=0.368
O = 0.368·v₁ + 1·v₂ ℓ = 1×0.368 + 1 = 1.368
j=3 : m=3, α=e^(2−3)=0.368
O = 0.135·v₁ + 0.368·v₂ + 1·v₃ ℓ = 1.503
收尾 : O/ℓ = (0.090, 0.245, 0.665)·v ✓

盯住 v1v_1 的系数:1.000 → 0.368 → 0.135,每来一个更大的 max,历史成果整体打折一次。这就是代码里 acc_o *= scores_scale 那一行的全部含义。

2.4 FlashAttention = online softmax × 两个 GEMM

online softmax(2018,Milakov & Gimelshein)只解决流式归一化,没碰矩阵乘。FlashAttention(2022)的洞察是:attention 恰好是 matmul → softmax → matmul 的三明治,三个部件可以共用同一个 tile 循环:

1
2
3
4
5
for j in 1..Tc:                                  # K/V 方向遍历
S_j = Q_tile @ K_j.T # gemm#1:喂料
m, alpha, P̃_j, ℓ = online_softmax_update(S_j) # 三件套
O_tile = alpha * O_tile + P̃_j @ V_j # gemm#2:消费
Output = O_tile / ℓ # 唯一一次归一化

SSP~\tilde{P} 的生成和消费全部在 SRAM/寄存器内完成,不写入 HBM。

一句话总结:先有 online softmax 这条递推式,才有 FlashAttention 这个 kernel;数学在前,工程在后。


三、切法:block_M 进 grid,block_N 进循环

3.1 都沿 seq 轴切:Q/O 进 grid,K/V 进循环

Q/O 与 K/V 都沿着 seq 轴切块,块大小分别由 block_M 与 block_N 控制,但两者的去向不同:

1
2
3
4
5
6
seq ──->
├── Q₁ ├── Q₂ ├── Q₃ ┤ 按 block_M 切,进 grid:每个 CUDA block 认领一块 Q,
├── O₁ ├── O₂ ├── O₃ ┤ 独立算出对应的 O,块与块之间零依赖
├── K₁ ├── K₂ ├── K₃ ┤ 按 block_N 切,进内层循环:每个 block 沿 K/V 块逐块扫描;
├── V₁ ├── V₂ ├── V₃ ┤ K 与 V 必须按同样的边界切块(P̃ 的列与 V_j 的行须是同一批 key)
└────────────────────┘

为什么这样分?softmax 对分数矩阵做逐行归一化,O 的第 ii 行只由 Q 的第 ii 行决定,与 Q 的其他行无关。因此按行把 Q/O 切成 block_M 大小的块、分给不同的 CUDA block 并行计算,块间不需要任何通信。K/V 则不同:任何一行的 softmax 分母都要对整条 seq 求和,每一块 Q 都需要全部的 K/V。K/V 无法划归某个 block 独占,只能作为内层循环的遍历方向,按 block_N 逐块加载、逐块累积。

对比普通 GEMM 就能看出差别:C=ABC = AB 的输出块 CijC_{ij} 只依赖 AA 的第 ii 个行块与 BB 的第 jj 个列块,M、N 两个方向都可以切块并行;attention 的输出在 seq 方向上对 K/V 有全局依赖(softmax 分母是全行求和),这个方向只能串行遍历。block_M 决定并行度(grid 大小),block_N 决定每个 block 的循环长度

3.2 grid 与布局

张量布局用 BSHD([batch, seq_len, heads, dim]),batch/head 在外层、seq 连续,tile 拷贝才能访存合并。grid 三维:

1
with T.Kernel(T.ceildiv(seq_len, block_M), heads, batch, threads=threads) as (bx, by, bz):

每个 CUDA block 认领一个 Q tile(bx)、一个 head(by)、一个 batch(bz),沿 K/V 方向循环。这正是 FA-2 的循环序:Q 装载一次驻留 SRAM 全程不动。FA-1 是反过来的(外层 K/V、内层 Q),每轮 K/V 块的结果要经 HBM 中转更新各 Q 块,多出大量中间读写–FA-2 把循环反过来之后这条路才彻底堵死。

3.3 块大小的约束

SMEM 需求近似为:

QiBr×d+KjBc×d+VjBc×d+SijBr×BcMSRAM\underbrace{Q_i}_{B_r \times d} + \underbrace{K_j}_{B_c \times d} + \underbrace{V_j}_{B_c \times d} + \underbrace{S_{ij}}_{B_r \times B_c} \le M_{\text{SRAM}}

通常取 Br=BcMSRAM/dB_r = B_c \approx \sqrt{M_{\text{SRAM}} / d}。A100 上 M100M \approx 100 KB、d=128d = 128B128B \approx 128;再考虑 fragment 状态(acc_s、acc_o 等)占的寄存器,128×128 是常见起点,seq 很长或 head_dim 很大时倾向减小 block_N。

3.4 缓冲区清单

缓冲区 位置 形状 角色
Q_shared / K_shared / V_shared SMEM [block_M/N, dim] GEMM 的 A/B 操作数走 SMEM 路径
O_shared SMEM [block_M, dim] 写回前的中转(fragment 布局重排)
acc_s fragment(fp32) [block_M, block_N] S = QK^T 分块累加器
acc_s_cast fragment(fp16) [block_M, block_N] softmax 后的 P̃,cast 给第二个 MMA
acc_o fragment(fp32) [block_M, dim] 输出累加器,全程不归一化
scores_max / _prev / _scale / _sum / ell fragment(fp32) [block_M] online softmax 的逐行状态

一句话总结:所有逐行状态都放 fragment(寄存器),只有 GEMM 操作数和最终输出走 SMEM–这是三级内存抽象在 FA 里的标准分工。


四、主循环逐行拆解

完整代码基于官方 examples/flash_attention/example_mha_fwd_bshd.py(注释为本文所加)。这是仅推理版:不产出 backward 需要的 logsumexp,分母变量直接叫 ell

4.1 三层结构与初始化

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
@autotune(configs=get_configs(), warmup=10, rep=10)
@tilelang.jit(
out_idx=[3], # 第 4 个形参 Output 是输出
pass_configs={tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True},
)
def flashattn(batch, heads, seq_len, dim, is_causal,
block_M=64, block_N=64, num_stages=1, threads=128):
# softmax 的 1/sqrt(dim),预先乘上 log2(e),后面全部用硬件更快的 exp2
scale = (1.0 / dim) ** 0.5 * 1.44269504

@T.prim_func
def main(Q: T.Tensor(shape, dtype), K: T.Tensor(shape, dtype),
V: T.Tensor(shape, dtype), Output: T.Tensor(shape, dtype)):
with T.Kernel(T.ceildiv(seq_len, block_M), heads, batch,
threads=threads) as (bx, by, bz):
... # 缓冲区分配见 3.4

# Q tile 装载一次,全程驻留
T.copy(Q[bz, bx*block_M:(bx+1)*block_M, by, :], Q_shared)
T.fill(acc_o, 0)
T.fill(ell, 0)
T.fill(scores_max, -T.infinity(accum_dtype))

# causal 截断:Q 块最远看到 (bx+1)*block_M - 1,右侧整块跳过
loop_range = (
T.min(T.ceildiv(seq_len, block_N),
T.ceildiv((bx+1)*block_M, block_N))
if is_causal else T.ceildiv(seq_len, block_N)
)

三层结构是 TileLang 的标准混用写法:外层 @autotune 管配置搜索,中间 @tilelang.jit(out_idx=[3]) 声明输出形参,内部 @T.prim_func 管精确签名。

causal 截断值得推一遍:Q 块 bx 的行最远到 (bx+1)blockM1(bx+1) \cdot block_M - 1,key 位置超过它的全被掩码,对应的 K/V 块整块不进循环。N=4096N = 4096、块 128 时对角线右侧一半块直接消失,省一半计算。注意非 causal 时这个 T.min 退化为总块数,掩码换成右边界越界判断(见下节)。

4.2 掩码写进累加器:顺序是精髓

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
for k in T.Pipelined(loop_range, num_stages=num_stages):
T.copy(K[bz, k*block_N:(k+1)*block_N, by, :], K_shared)

# ① 先写掩码,再做 GEMM
if is_causal:
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.if_then_else(
bx*block_M + i >= k*block_N + j, 0, -T.infinity(acc_s.dtype))
else:
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.if_then_else(
k*block_N + j >= seq_len, -T.infinity(acc_s.dtype), 0)

# ② S = Q @ K^T(tensor core)
T.gemm(Q_shared, K_shared, acc_s, transpose_B=True,
policy=T.GemmWarpPolicy.FullRow)

掩码写在 GEMM 之前能成立,靠的是 T.gemm 的累加语义(C+=ABC \mathrel{+}= A B,所以纯 GEMM 例子里要先 T.clear):$-\infty + $ 有限值 == -\infty,被掩码的位置在 GEMM 里"存活"下来,之后 exp2(m)=0\mathrm{exp2}(-\infty - m) = 0,对分数和、对输出零贡献。省掉一遍 GEMM 后的掩码 pass。

两个分支的条件不同,处理的边界不同:

  • causal:query 全局位置 ≥ key 全局位置才保留(只许看过去);
  • 非 causal:k*block_N + j >= seq_len 置 −inf,处理的是 seq 不整除 block_N 时最后一块里的"假 key"。

4.3 online softmax 七步

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
# ③ online softmax 更新
T.copy(scores_max, scores_max_prev) # 1. 存历史 max
T.fill(scores_max, -T.infinity(accum_dtype))
T.reduce_max(acc_s, scores_max, dim=1, clear=False) # 2. 本块行 max
for i in T.Parallel(block_M):
scores_max[i] = T.max(scores_max[i], scores_max_prev[i]) # 3. 合并成 m_new
for i in T.Parallel(block_M):
scores_scale[i] = T.exp2( # 4. α = exp2(m_old − m_new)
scores_max_prev[i]*scale - scores_max[i]*scale)
for i, j in T.Parallel(block_M, block_N):
acc_s[i, j] = T.exp2( # 5. P̃ = exp2(S·scale − m·scale)
acc_s[i, j]*scale - scores_max[i]*scale)
T.reduce_sum(acc_s, scores_sum, dim=1) # 6. 本块 rowsum(P̃)
for i in T.Parallel(block_M):
ell[i] = ell[i]*scores_scale[i] + scores_sum[i] # 7. ℓ ← α·ℓ + 本块部分和
T.copy(acc_s, acc_s_cast) # fp32 → fp16,喂第二个 MMA

对应 2.3 节的递推式,逐步可查。两个容易踩的坑:

  • max 与 scale 的次序:先由历史 max 和本块 max 定出 mnewm_{new}由这一对 max 算出 α\alpha。max 是 scale 的来源,不是被更新对象;
  • scores_sumell 别混:前者只是本轮 rowsum(P̃) 的临时量,每次迭代被覆盖;ell 是带折扣链跨轮累积的总分母。归一化除的是 ell,除成某一块的局部和就错了。

4.4 第二个 GEMM 与收尾

1
2
3
4
5
6
7
8
9
10
11
    # ④ 先打折旧成果,再加新项
for i, j in T.Parallel(block_M, dim):
acc_o[i, j] *= scores_scale[i]
T.copy(V[bz, k*block_N:(k+1)*block_N, by, :], V_shared)
T.gemm(acc_s_cast, V_shared, acc_o, policy=T.GemmWarpPolicy.FullRow)

# ⑤ 循环外:唯一一次归一化,然后写回
for i, j in T.Parallel(block_M, dim):
acc_o[i, j] /= ell[i]
T.copy(acc_o, O_shared)
T.copy(O_shared, Output[bz, bx*block_M:(bx+1)*block_M, by, :])

acc_o *= scores_scale 必须在第二个 GEMM 之前:先加新项再打折,新项就被错误地折旧了。归一化只在循环外做一次,中途除会错–分母没定型。

写回走 fragment → O_shared → Output 两跳:fragment 里的数据按 MMA 布局散在各线程寄存器,先落 SMEM 重排布局,再合并写 global。

4.5 工程细节清单

  • exp2 技巧:把 1/d1/\sqrt{d} 预乘 log2e1.4427\log_2 e \approx 1.4427 折进 scale,全部指数运算用 exp2 硬件指令(比 exp 快),FA 实现的标配;
  • P̃ 的 fp16 cast 是全 kernel 唯一降精度点:MMA 只吃低精度操作数,acc_s/acc_o/ell 等累加器与状态全程 fp32;
  • FullRow warp policyT.gemm 把输出 tile 切给 block 内各 warp 的切法。FA 的 gemm#1 后紧跟按行归约(reduce_max/reduce_sum)和按行缩放(acc_o *= scores_scale),FullRow 保证一行完整住在一个 warp 内,这些操作全部退化为 warp 内 shuffle;若用默认 Square,第 i 行的左右两半住在两个 warp 里,归约就得经 SMEM 中转外加 bar.sync–结果仍正确,但每轮迭代多两次往返:
1
2
3
4
5
6
7
8
Square(默认):每 warp 近似方块    FullRow:按 M 横切,每 warp 拿全宽整段行
┌─────────┬─────────┐ ┌──────────────────────┐
│ warp0 │ warp1 │ │ warp0 行 0..31 │
│ 64×64 │ 64×64 │ ├──────────────────────┤
├─────────┼─────────┤ │ warp1 行 32..63 │
│ warp2 │ warp3 │ ├──────────────────────┤
│ 64×64 │ 64×64 │ │ warp2/3 行 64..127│
└─────────┴─────────┘ └──────────────────────┘
  • forward 用 num_stages=1:循环体是两个相互依赖的 GEMM 夹一串 elementwise(S 等 K、P̃ 等 S 和 m),预取重叠窗口小而寄存器压力大,官方权衡后的选择。流水线什么时候真的有用,见上一篇的三因素模型与对照实验;
  • O 不中途归一化是 FA-2 的关键改动之一(见下节)。

一句话总结:主循环五行骨架–搬 K、写掩码、GEMM 喂料、online softmax 三件套、rescale 后 GEMM 消费;每一步都有"顺序不能错"的理由。


五、与 FA-2 论文 Algorithm 1 的对照

这份代码就是 FlashAttention-2 论文 Algorithm 1 的逐行实现。伪代码符号 ↔ 代码变量:

伪代码 内容 代码对应
Br×BcB_r \times B_c 分块大小 block_M × block_N
for i ≤ Tᵣ(外层 Q 块) 外层循环进 grid with T.Kernel(...) as (bx, by, bz)
QiQ_i 进 SRAM 装载后驻留 T.copy(Q[...], Q_shared)
O0=0, 0=0, m0=O^0 = 0,\ \ell^0 = 0,\ m^0 = -\infty 初始化三状态 T.fill(acc_o/ell, 0)T.fill(scores_max, -inf)
for j ≤ T_c(内层 K/V 块) 内层循环 for k in T.Pipelined(loop_range, ...)
Sij=QiKjS_{ij} = Q_i K_j^\top 分数块 T.gemm(Q_shared, K_shared, acc_s, transpose_B=True)
mm 更新、P~\tilde{P}\ell 更新 online softmax 一族 reduce_maxscores_scaleexp2ell 递推
Odiag(eΔm)1O+P~VjO \leftarrow \mathrm{diag}(e^{\Delta m})^{-1} O + \tilde{P} V_j 先打折再加新项 acc_o *= scores_scalegemm(acc_s_cast, V_shared, acc_o)
O/O/\ell(第 12-13 行) 循环外归一化 acc_o /= ell

四处伪代码没写、真实 kernel 必须处理的差异:

  1. 底数:论文 e 底,实现全用 exp2 硬件指令,scale 预乘 log2e\log_2 e
  2. scale 折叠1/d1/\sqrt{d} 不显式乘,折进指数表达式省一遍逐元素乘;
  3. acc_s_cast:P̃ 从 fp32 fragment cast 成 fp16 才能进 MMA;
  4. causal 截断:算法按稠密写,代码用 loop_range 让对角线右侧整块跳过。

顺带把 FA-1 → FA-2 的三个改动列清楚,这份代码全部站在 FA-2 一边:

FA-1(2022) FA-2(2023)
外层循环 K/V 块 Q 块
中间读写 各 Q 块的结果需经 HBM 中转更新 Q 驻留 SRAM,块间零 HBM 往返
O 的归一化 逐块保持已归一化(多两次乘除) 循环外一次性 /= ell
并行度 受 K/V 块数限制 Q 块 × head × batch,并行度更高
warp 分工 均分 K/V 切给不同 warp,减少同步(TileLang 里由 GemmWarpPolicy 表达)

六、验证、测速与调参

写完 kernel 的三个标准动作(数值验证 -> dump 生成码 -> 基准测试),在 FA 上一个不少:

① 数值验证profiler.assert_allclose 直接接受 PyTorch 参考实现做正确性比对,不需要自己写验证框架:

1
2
3
4
5
kernel = flashattn(batch, heads, seq_len, dim, is_causal,
block_M=128, block_N=128, num_stages=1, threads=128)
ref = partial(ref_program, is_causal=is_causal) # einsum 朴素注意力 + tril 掩码
profiler = kernel.get_profiler()
profiler.assert_allclose(ref, rtol=0.01, atol=0.01)

容差 0.01 是 fp16 输出的合理范围;更严的判定方法(与 fp32 参考比、大误差元素计数)见上一篇第九节。

② dump 生成码kernel.get_kernel_source() 打印完整 CUDA 源码。对 FA 值得确认三件事:T.copy 降成了 cp.async 还是 TMA、FullRow 下 warp 怎么切行、exp2 是否真的成了硬件指令。

③ 测速profiler.do_bench(warmup=500) 自动 warmup 多次取统计。TFlops 按 2BHN2d×22 \cdot B \cdot H \cdot N^2 \cdot d \times 2 个 matmul 计算,causal 乘 0.5:

1
2
3
4
flops = 2.0 * batch * heads * seq_len * seq_len * dim   # 单个 matmul
total_flops = 2 * flops * (0.5 if is_causal else 1.0)
latency = profiler.do_bench(warmup=500)
print(f"{total_flops / latency * 1e-9:.2f} TFlops")

调参交给 @autotune:把 block_M/block_N/num_stages/threads 写成带默认值的参数,搜索空间在 get_configs() 里列全(笛卡尔积),机器自己选。注意首跑很慢(每组都要编译 + 测速),调通阶段先裁成单组。


七、总结

  1. 数学在前:softmax 平移不变性 → 安全 softmax → 分块下的换基准问题 → online softmax 递推。FlashAttention 只是这条递推式与两个 GEMM 的交织,精确、非近似;
  2. 切法是结构:Q/O 与 K/V 都沿 seq 轴切块,block_M 控制的 Q 块进 grid(决定并行度),block_N 控制的 K/V 块进内层循环(决定遍历长度);softmax 的逐行归约使 seq 方向对 K/V 产生全局依赖,只能作为循环方向,这是 FA 与 GEMM 的本质结构差异;
  3. 实现靠 TileLang 五原语T.Kernel 定切法、T.copy 管搬运、T.gemm 喂两个 matmul、T.Pipelined 管 K/V 流水(FA 前向官方权衡后用 stages=1)、T.Parallel + reduce_* 承载 online softmax 三件套。掩码写进累加器、exp2 折叠、FullRow 行所有权、O 延迟归一化,四个细节决定了这份代码"像论文"还是"只是能跑"。

可迁移的启示:读 FA 代码的正确顺序是先读递推式再读循环–所有 tile 级 kernel 都是"一条数学递推 + 一个 tile 循环",TileLang 只是把后者写到了 30 行的量级。


参考

  • FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(Dao et al., NeurIPS 2022, arXiv:2205.14135)
  • FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning(Dao, 2023, arXiv:2307.08691)
  • Online normalizer calculation for softmax(Milakov & Gimelshein, 2018, arXiv:1805.02867)
  • TileLang GitHub(v0.1.13):examples/flash_attention/example_mha_fwd_bshd.py 为本文代码底本
  • 本站:《TileLang 编程基本知识点

本文基于 TileLang v0.1.13 官方示例与 FlashAttention-2 论文整理,代码注释为本文所加。