TileLang 实战:FlashAttention 前向 Kernel
TileLang 实战:FlashAttention 前向 Kernel
FlashAttention 没有发明新的数学。它是 online softmax 递推与两个 GEMM 在同一个 tile 循环里的交织–分数块 算出来就地做 softmax 得到 , 就地乘上 ,中间矩阵全程驻留在 SRAM,HBM 流量从 降到 。本文用 TileLang(v0.1.13)把 FA-2 论文的 Algorithm 1 写成可运行的 kernel:先补齐数学地基,再讲清楚"怎么切",然后逐行拆解主循环,最后做数值验证、测速与调参。只覆盖前向;TileLang 五原语与 T.Pipelined 流水线机制见上一篇《TileLang 编程基本知识点》,本文直接使用其结论。
一、为什么标准 Attention 慢:IO 才是瓶颈
标准 scaled dot-product attention 的流程是 。计算本身没有问题,问题在中间矩阵: 和 都是 ,必须写进显存再读回来。 时单个 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 是 ,标准实现与 FlashAttention 完全相同–快慢差别全部来自中间矩阵走没走 HBM:
| 指标 | 标准 Attention | FlashAttention |
|---|---|---|
| 中间矩阵内存 | ,S/P 落显存 | ,S/P̃ 留在 SRAM |
| HBM 往返 | 3 次(写 S → 读写 P → 读 P 做 PV) | 1 次(读 Q/K/V → 写 O) |
| FLOPs | 相同 | |
| 精确性 | 精确 | 精确(非近似) |
这就是论文标题里 “IO-aware” 的含义:不省计算,省搬运。
一句话总结:attention 是访存受限的算子,FlashAttention 的全部收益来自让 的中间矩阵不落 HBM。
二、数学地基:softmax 的三级台阶
2.1 平移不变性 → 安全 softmax
softmax 对输入加任意常数 不变(分子分母的 约掉)。这个自由度是数值安全的救命稻草:fp16 上限只有 65504, 时 直接 inf,inf/inf 变 NaN。减去行最大值 后指数上限恰为 0:
代价是计算顺序被强制:先扫一遍求 ,再扫一遍求分母与输出–两遍扫描。
2.2 分块场景下,"两遍"成为灾难
SRAM 装不下整行 K/V,只能按块流入。设处理完第 1 块时按基准 攒好部分和;第 2 块冒出更大分数,基准换成 –旧 exp 全部基准错误。三条路:
| 方案 | 做法 | 代价 |
|---|---|---|
| 存下全部 S | 等全局 max 再统一归一化 | 落显存,爆 |
| 两遍重扫 | pass1 求 、;pass2 重算加权 | K/V 从 HBM 读两次,带宽 ×2 |
| online softmax | 边收块边维护"以当前 max 为基准"的部分和 | 一遍过,数学无损 |
2.3 online softmax 递推
每个 query 行维护三个状态,初值 、、。新块 到来时:
循环结束后做全程唯一一次归一化:。
为什么无损:缩放因子连乘时指数项逐项相消,,循环结束时每项都被精确换算到最终 max 的基准下。中间任何时刻 都是合法但未归一化的加权和,数值永远有界(当前最大项的 exp 恰为 1)。
数值走一遍(一行 query,三块各来一个分数 1、2、3,正确答案 = softmax(1, 2, 3) 的权重):
1 | init : m=−∞ ℓ=0 O=0 |
盯住 的系数: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 | for j in 1..Tc: # K/V 方向遍历 |
与 的生成和消费全部在 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 | seq ──-> |
为什么这样分?softmax 对分数矩阵做逐行归一化,O 的第 行只由 Q 的第 行决定,与 Q 的其他行无关。因此按行把 Q/O 切成 block_M 大小的块、分给不同的 CUDA block 并行计算,块间不需要任何通信。K/V 则不同:任何一行的 softmax 分母都要对整条 seq 求和,每一块 Q 都需要全部的 K/V。K/V 无法划归某个 block 独占,只能作为内层循环的遍历方向,按 block_N 逐块加载、逐块累积。
对比普通 GEMM 就能看出差别: 的输出块 只依赖 的第 个行块与 的第 个列块,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 需求近似为:
通常取 。A100 上 KB、 时 ;再考虑 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 |
|
三层结构是 TileLang 的标准混用写法:外层 @autotune 管配置搜索,中间 @tilelang.jit(out_idx=[3]) 声明输出形参,内部 @T.prim_func 管精确签名。
causal 截断值得推一遍:Q 块 bx 的行最远到 ,key 位置超过它的全被掩码,对应的 K/V 块整块不进循环。、块 128 时对角线右侧一半块直接消失,省一半计算。注意非 causal 时这个 T.min 退化为总块数,掩码换成右边界越界判断(见下节)。
4.2 掩码写进累加器:顺序是精髓
1 | for k in T.Pipelined(loop_range, num_stages=num_stages): |
掩码写在 GEMM 之前能成立,靠的是 T.gemm 的累加语义(,所以纯 GEMM 例子里要先 T.clear):$-\infty + $ 有限值 ,被掩码的位置在 GEMM 里"存活"下来,之后 ,对分数和、对输出零贡献。省掉一遍 GEMM 后的掩码 pass。
两个分支的条件不同,处理的边界不同:
- causal:query 全局位置 ≥ key 全局位置才保留(只许看过去);
- 非 causal:
k*block_N + j >= seq_len置 −inf,处理的是 seq 不整除 block_N 时最后一块里的"假 key"。
4.3 online softmax 七步
1 | # ③ online softmax 更新 |
对应 2.3 节的递推式,逐步可查。两个容易踩的坑:
- max 与 scale 的次序:先由历史 max 和本块 max 定出 ,再由这一对 max 算出 。max 是 scale 的来源,不是被更新对象;
scores_sum与ell别混:前者只是本轮 rowsum(P̃) 的临时量,每次迭代被覆盖;ell是带折扣链跨轮累积的总分母。归一化除的是ell,除成某一块的局部和就错了。
4.4 第二个 GEMM 与收尾
1 | # ④ 先打折旧成果,再加新项 |
acc_o *= scores_scale 必须在第二个 GEMM 之前:先加新项再打折,新项就被错误地折旧了。归一化只在循环外做一次,中途除会错–分母没定型。
写回走 fragment → O_shared → Output 两跳:fragment 里的数据按 MMA 布局散在各线程寄存器,先落 SMEM 重排布局,再合并写 global。
4.5 工程细节清单
- exp2 技巧:把 预乘 折进
scale,全部指数运算用exp2硬件指令(比exp快),FA 实现的标配; - P̃ 的 fp16 cast 是全 kernel 唯一降精度点:MMA 只吃低精度操作数,
acc_s/acc_o/ell等累加器与状态全程 fp32; - FullRow warp policy:
T.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 | Square(默认):每 warp 近似方块 FullRow:按 M 横切,每 warp 拿全宽整段行 |
- 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 的逐行实现。伪代码符号 ↔ 代码变量:
| 伪代码 | 内容 | 代码对应 |
|---|---|---|
| 分块大小 | block_M × block_N |
|
for i ≤ Tᵣ(外层 Q 块) |
外层循环进 grid | with T.Kernel(...) as (bx, by, bz) |
| 进 SRAM | 装载后驻留 | T.copy(Q[...], Q_shared) |
| 初始化三状态 | T.fill(acc_o/ell, 0)、T.fill(scores_max, -inf) |
|
for j ≤ T_c(内层 K/V 块) |
内层循环 | for k in T.Pipelined(loop_range, ...) |
| 分数块 | T.gemm(Q_shared, K_shared, acc_s, transpose_B=True) |
|
| 更新、、 更新 | online softmax 一族 | reduce_max → scores_scale → exp2 → ell 递推 |
| 先打折再加新项 | acc_o *= scores_scale → gemm(acc_s_cast, V_shared, acc_o) |
|
| (第 12-13 行) | 循环外归一化 | acc_o /= ell |
四处伪代码没写、真实 kernel 必须处理的差异:
- 底数:论文 e 底,实现全用 exp2 硬件指令,scale 预乘 ;
- scale 折叠: 不显式乘,折进指数表达式省一遍逐元素乘;
- acc_s_cast:P̃ 从 fp32 fragment cast 成 fp16 才能进 MMA;
- 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 | kernel = flashattn(batch, heads, seq_len, dim, is_causal, |
容差 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 按 个 matmul 计算,causal 乘 0.5:
1 | flops = 2.0 * batch * heads * seq_len * seq_len * dim # 单个 matmul |
调参交给 @autotune:把 block_M/block_N/num_stages/threads 写成带默认值的参数,搜索空间在 get_configs() 里列全(笛卡尔积),机器自己选。注意首跑很慢(每组都要编译 + 测速),调通阶段先裁成单组。
七、总结
- 数学在前:softmax 平移不变性 → 安全 softmax → 分块下的换基准问题 → online softmax 递推。FlashAttention 只是这条递推式与两个 GEMM 的交织,精确、非近似;
- 切法是结构:Q/O 与 K/V 都沿 seq 轴切块,block_M 控制的 Q 块进 grid(决定并行度),block_N 控制的 K/V 块进内层循环(决定遍历长度);softmax 的逐行归约使 seq 方向对 K/V 产生全局依赖,只能作为循环方向,这是 FA 与 GEMM 的本质结构差异;
- 实现靠 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 论文整理,代码注释为本文所加。