从 FlashAttention 到 FlashAttention-4:把注意力算成一道内存题(附 8 个交互式动画)
注意力慢的根因不是算力不够,而是数据在 HBM 和 SRAM 之间反复搬运。本文梳理 FlashAttention 原版到 FlashAttention-4、以及解码侧的 Flash-Decoding / FlashDecoding++,讲清每一代把哪块短板按了下去,并配 8 个可以动手玩的 HTML 交互动画。
注意力是 Transformer 里最贵的一层:序列长度翻倍,计算和显存都翻四倍。2022 年之前,大家对付它的主流思路是「近似」——稀疏化、低秩、核方法,用模型质量换计算量。但一个尴尬的事实是:这些方法 FLOPs 降下来了,墙上时间却没怎么降。
FlashAttention 换了个问法:注意力到底是被算力卡住,还是被内存卡住? 答案是后者。它没有改变注意力的数学定义,一分近似都没做,只是把「数据怎么在显存和片上缓存之间搬」这件事重新安排了一遍,就在 GPT-2 上拿到 3× 的端到端加速,并把显存从平方级降到线性级。
这之后是一条延续四年的技术线:FA1 解决「要不要物化 N×N 矩阵」,FA2 解决「GPU 喂不饱」,FA3 解决「Hopper 的异步能力没用上」,FA4 解决「Blackwell 上瓶颈挪了位置」。另外还有一条推理侧的分支,专门对付解码时 query 只有 1 个的场景。
这篇把这条线捋一遍。核心机制我都做成了可以点的交互动画,建议边看边玩。
说明
一个前提:FlashAttention 全系列都是精确注意力,输出和标准
softmax(QKᵀ)V逐位一致(FA3 的 FP8 版本除外,那是低精度数值误差,不是算法近似)。它省的是显存和带宽,不是数学。
一、先看清起点:显存墙和带宽墙
GPU 的内存是分层的。以 FA1 论文里反复举的 A100 为例:
| 层级 | 容量 | 带宽 | 特点 |
|---|---|---|---|
| HBM(显存) | 40–80 GB | 1.5–2.0 TB/s | 大而慢,放得下整个 N×N 矩阵 |
| SRAM(片上,每个 SM) | 192 KB × 108 | ~19 TB/s | 小而快,一次只放得下一个块 |
两者带宽差了大约一个数量级,容量差了几个数量级。而标准注意力的写法是:
S = Q Kᵀ # N×N
P = softmax(S) # N×N
O = P V # N×d
问题就在中间那个 N×N。序列长度 N=8192、头维度 d=128 时,单是 S 和 P 两个中间矩阵,每个头就要占掉约 17 GB 的 fp16 显存——注意这是每个注意力头。于是整个计算被强行拆成一串 kernel:算 S、存回去、读出来做 softmax、再存回去、读出来乘 V。每一次往返都在 HBM 上。论文给出的量级是:标准注意力需要 Θ(Nd + N²) 次 HBM 访问,主导项是平方的。
拖动下面的滑块,看看这个数量级意味着什么。
这也解释了为什么近似方法常常不划算:它们减少的是 FLOPs,但瓶颈在带宽。除非同时把访存也降下来,否则省下的计算时间会被搬运时间吃掉。
二、FlashAttention(2022):分块、在线 softmax、重计算
FA1 的核心只有三句话:分块(tiling)、在线 softmax(online softmax)、重计算(recomputation)。
分块与在线 softmax
softmax 天生「不局部」:要算某一行的归一化,得先知道这一整行的最大值和指数和。而分块计算时,我们一次只看到一行的一部分。在线 softmax 解决的正是这个矛盾——它维护三个逐行状态,每读入一个新块就更新一次:
m:到目前为止见过的最大分数;ℓ:到目前为止的指数和;Õ:到目前为止未归一化的输出累加。
每来一个新块,就用新的最大值把旧状态重新缩放一遍,再把新块的贡献加进去。关键在于「重新缩放」这一步:它让各个块可以按任意顺序处理,最终结果与一次性算完全一致。
点开下面的动画,它会用一个 4×6 的小矩阵,把这三个状态逐块更新给你看,最后还会和一次性算出的 softmax 结果对拍。
重计算:反向传播不存中间矩阵
前向省下了 N×N,反向还有一块:标准的反向传播需要用到前向算出的注意力概率矩阵 P。FA1 的选择是不存,用的时候重算——反向时重新加载 Q、K、V 的块,在 SRAM 里把 P 再算一遍。
这看起来是「用计算换内存」,但因为在片内算比从 HBM 读回来快得多,实际是净赚。论文的结论是:即便因为重计算增加了 FLOPs,整体仍然更快(GPT-2 上最高 7.6× 的注意力加速),显存则从平方级降到线性级,额外显存只有 O(N)。
它到底省了多少
论文给了两个可直接引用的结论:
| 结论 | 内容 |
|---|---|
| Theorem 1 | 算法输出精确的 softmax(QKᵀ)V,FLOPs 仍是 O(N²d),但输入输出之外只需 O(N) 额外显存 |
| Theorem 2 | 标准注意力需 Θ(Nd + N²) 次 HBM 访问;FlashAttention 需 Θ(N²d²/M) 次(M 为 SRAM 大小),且在 d ≤ M ≤ Nd 范围内是最优的 |
由于 d² 通常远小于 M(d 取 64–128,M 约 100 KB 量级),FlashAttention 的 HBM 访问量比标准实现少一个数量级左右。论文报告的战绩是:BERT-large 端到端快 15%、GPT-2(序列 1K)快 3×、长程竞技场(1K–4K)快 2.4×,显存节省 10–20×。它还顺手做了一个 block-sparse 版本,比当时所有近似注意力方法都快。
提示
记住一个反直觉的点:FlashAttention 做的事不是减少计算,而是增加了计算(重计算)却大幅减少了显存访问。它是典型的内存受限优化,不是计算受限优化。
三、FlashAttention-2(2023):把 GPU 喂饱
FA1 虽然比基线快,但离 GEMM 的效率还差得远:注意力前向只跑到 A100 理论峰值的 30–50%,反向只有 25–35%,而优化过的 GEMM 能到 80–90%。
FA2 的诊断是「工作划分(work partitioning)不对」,然后开了三个药方。
第一,减少非 matmul 运算。 GPU 上矩阵乘有专用单元,非 matmul 运算(逐元素乘、除、指数)的吞吐可以低到 1/16。FA1 每处理一个块都要把输出除以 ℓ 做归一化;FA2 改成保留未归一化的 Õ,只在循环全部结束后除以一次 ℓ,同时只存 log-sum-exp 而不是同时存 m 和 ℓ。
第二,把序列维度也切开做并行。 FA1 只沿 batch×heads 并行——序列越长、batch 越小,能用的线程块就越少,GPU 大量 SM 闲着。FA2 把每个 query 块也当成独立任务,于是「batch×heads×query 块数」一起填满 GPU。
第三,线程块内按 Q 切分给 warpp,而不是按 K/V 切分。 FA1 是「split-K」:4 个 warp 各拿一片 K/V,算出各自的部分结果后,得写共享内存、同步、再相加。FA2 改成每个 warp 负责一片 Q 行,K/V 对所有 warp 可见,各算各的输出——不需要 warp 间通信。
结果:相比 FA1 约 2× 加速,达到 A100 理论峰值的 50–73%——终于接近 GEMM 的效率。端到端训练 GPT 类模型时最高 225 TFLOPs/s(72% MFU)。
四、FlashAttention-3(2024):Hopper 的异步与低精度
到了 H100,硬件多了三样东西:TMA(张量内存加速器,负责异步搬数据)、WGMMA(异步的 warpgroup 矩阵乘指令)、以及 FP8 Tensor Core。但 FA2 在 H100 上只跑到 35%——因为它还是「先搬数据、再算矩阵、再算 softmax」的顺序结构。
FA3 的三招都在解决「让不同硬件单元同时干活」。
生产者 / 消费者 warp 专门化。 把搬运(TMA)和计算(WGMMA)拆给不同的 warp,形成软件流水线,让「搬下一块」和「算这一块」重叠起来。
pingpong 调度。 这一招针对的是 softmax 里的指数运算。H100 的 FP16 矩阵乘有 989 TFLOPs/s,但指数这类特殊函数只有 3.9 TFLOPs/s——差 250 倍。在 head_dim=128 的前向里,指数运算能吃掉接近一半的周期。既然指数走的是独立的多功能单元,那就让它和 Tensor Core 同时忙:用同步屏障安排两个 warpgroup 的 GEMM 交替进行,一个在算矩阵乘时,另一个正好在算 softmax。
FP8 与 incoherent processing。 FP8 的显存和算力都更划算,但注意力矩阵里常有离群值,直接把整块量化会让刻度被离群值撑大、普通值被抹平。FA3 的做法是把 Q、K 都乘上一个随机 ±1 对角矩阵和 Hadamard 矩阵——正交变换不改变 QKᵀ,却能把离群值摊平,降低量化误差。这个变换可以用 O(d log d) 的快速算法实现,还能和 RoPE 融合,几乎不花额外时间。
结果:FP16 最高 740 TFLOPs/s(75% 利用率),FP8 接近 1.2 PFLOPs/s,比 FA2 快 1.5–2.0×;FP8 版本的数值误差比基线 FP8 注意力低 2.6×。这套实现建立在 CUTLASS 的 WGMMA / TMA 原语上。
五、FlashAttention-4(2026):Blackwell 的「非对称扩展」
Blackwell(B200 / GB200)带来的变化不是「什么都更快」,而是非对称的:Tensor Core 吞吐翻倍,但共享内存带宽和指数单元增长很慢,甚至原地不动。
后果很直接:在 Hopper 上已经不小的 softmax / 指数开销,在 Blackwell 上反过来超过了 MMA 计算 25–60%。把 FA3 直接搬到 Blackwell 也行不通——Hopper 的 MMA 指令在 Blackwell 上没有前向兼容。
FA4 于是重做了流水线,同样三招:
- 更大的 tile + 全异步 MMA。 Blackwell 的 MMA 块是 128×128(面积是 Hopper 64×128 的两倍),且异步 Tensor Core 直接写 TMEM。FA4 用更大的 tile 和完全异步的 MMA,把搬运、矩阵乘、softmax 尽量叠在同一段时间里。
- 把指数运算「软」下来。 用 FMA 单元上的多项式近似来模拟 exp,绕开吞吐最紧的多功能单元;再引入条件式 softmax 重缩放,跳过不必要的缩放操作。
- 降低共享内存流量。 反向把更多中间结果放进每个 SM 256 KB 的 TMEM;并使用 2-CTA MMA 模式,让每个 CTA 只加载一半的 B 操作数,减少共享内存流量和反向的原子加。
结果:B200 上 BF16 最高 1613 TFLOPs/s(71% 利用率),比 cuDNN 9.13 快约 1.3×、比 Triton 快约 2.7×。另一个值得注意的工程变化是,FA4 整个用嵌入 Python 的 CuTe-DSL 写成,编译时间比传统 C++ 模板快 20–30×。
说明
FA4 这条线还在往前走。2026 年 9 月的 Hardware-Aware FP4 FlashAttention-4 把概率本身也压到 FP4(Direct-P),GB200 上前向吞吐达到 BF16 的 2.13×。不过论文也诚实地报告:所有被测的 MXFP4 概率/值训练轨迹都发散了——低精度概率目前更适合推理。
六、推理侧的分支:Flash-Decoding 与 FlashDecoding++
上面讲的都是训练视角:并行维度是 batch、heads、query 块。但自回归解码时,每步只生成一个 token,query 长度是 1,并行维度一下子塌了。如果 batch 也很小(长上下文时为了放得下 KV cache,batch 往往只能开 1),FA 能用到的 SM 少得可怜——论文原话说,batch=1 时用不到 1% 的 GPU。
Flash-Decoding 增加了一个新的并行维度:K/V 序列长度。做法分三步:
- 把 K/V 切成若干块(只是视图,不复制数据);
- 每块并行跑一遍 FlashAttention,额外写出每行的 log-sum-exp;
- 用 log-sum-exp 重新缩放各块的部分输出,归约相加。
因为 softmax 本身就可以迭代地算,这个归约不会引入近似。代价只是一次很小的合并 kernel。
另一篇 FlashDecoding++ 更偏工程:它指出解码推理里还有三处浪费——部分 softmax 之间需要同步更新(约 20% 开销)、扁平形状的 GEMM 利用率损失超过 50%、固定 dataflow 在不同输入下最多损失 50%;对应给出统一 max 值的异步 softmax、带双缓冲的 flat GEMM 优化、以及随硬件自适应的启发式 dataflow。它在 NVIDIA / AMD 上相对 Hugging Face 最高快 4.86× / 2.18×,相对当时最优推理引擎平均 1.37×。
七、一张表看懂每一代
| 版本 | 时间 | 论文 | 要按下去的短板 | 关键手段 | 代表数字 |
|---|---|---|---|---|---|
| FlashAttention | 2022.05 | arXiv.14135 | N×N 中间矩阵的显存与访存 | 分块 + 在线 softmax + 反向重计算 | GPT-2 3×、BERT-large +15%、显存 10–20× |
| FlashAttention-2 | 2023.07 | arXiv.08691 | GPU 占用率低、非 matmul 与 warp 通信 | 序列维并行、按 Q 划分 warp、少做非 matmul | 比 FA1 快 2×、225 TFLOPs/s(72% MFU) |
| Flash-Decoding | 2023.10 | PyTorch Blog | 解码时 query=1、并行度塌缩 | 沿 KV 长度 split-KV + log-sum-exp 归约 | 长序列解码最高约 8× |
| FlashDecoding++ | 2023.11 | arXiv.01282 | 同步 softmax、扁平 GEMM、静态 dataflow | 统一 max 异步 softmax、flat GEMM 双缓冲、启发式 dataflow | 相对 HF 最高 4.86× / 2.18× |
| FlashAttention-3 | 2024.07 | arXiv.08608 | Hopper 的异步能力没用上、指数单元瓶颈 | warp 专门化 + pingpong 调度 + FP8 incoherent processing | FP16 740 TFLOPs/s、FP8 ~1.2 PFLOPs/s、比 FA2 快 1.5–2× |
| FlashAttention-4 | 2026.03 | arXiv.05451 | Blackwell 非对称扩展、非 matmul 反超 | 全异步 MMA + 大 tile、模拟 exp + 条件重缩放、TMEM + 2-CTA MMA | 1613 TFLOPs/s(71%)、cuDNN 1.3×、Triton 2.7× |
| FP4 FlashAttention-4 | 2026.09 | arXiv.04105 | FP4 下 softmax 转换与片上依赖 | Direct-P 直接把分数映射为 FP4 概率 | GB200 前向 2.13× BF16(推理) |
八、怎么选,以及几个常见误解
先分清训练还是推理。 训练用 FA2 / FA3 / FA4 这条主线;解码长上下文别忘了 split-KV(Flash-Decoding 那一支)。两者解决的瓶颈不是一个。
务必按硬件选版本。 FA3 的收益来自 Hopper 的 TMA / WGMMA / FP8,FA4 的收益来自 Blackwell 的 TMEM 与 2-CTA MMA。在 A100 上跑 FA3 的内核不会变快,FA4 的指令在 Hopper 上甚至不兼容。实际使用时,先确认你的卡和对应的 kernel 分支。
它不是近似方法。 FA1–FA4 的输出与标准注意力一致;FP8 / FP4 版本引入的是数值精度误差,而不是算法层面的近似——这也是它和稀疏、低秩、线性注意力那一类方法的根本区别。
它不减少 FLOPs,反而增加。 反向重计算是用算力换访存。只有当访存是瓶颈时这笔交易才划算——在 GPU 上几乎总是划算,因为算力增长速度长期快于内存带宽。
显存从平方降到线性,但计算仍是平方。 FlashAttention 解决的是「放不下」和「搬太多」,不是把 O(N²) 变成 O(N)。它让 Window Attention、稀疏注意力之类的方法有了可用的实现载体,但注意力本身的二次复杂度还在——要真正线性化,得改数学,那是另一条路线。
参考资料
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness — arXiv.14135(Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré, NeurIPS 2022)
- FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning — arXiv.08691(Tri Dao)
- Flash-Decoding for long-context inference — PyTorch Blog(Tri Dao, Daniel Haziza, Francisco Massa, Grigory Sizov)
- FlashDecoding++: Faster Large Language Model Inference on GPUs — arXiv.01282(Ke Hong 等)
- FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision — arXiv.08608(Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao)
- FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling — arXiv.05451(Ted Zadouri, Markus Hoehnerbach, Jay Shah, Timmy Liu, Vijay Thakkar, Tri Dao)
- Hardware-Aware FP4 FlashAttention-4 — arXiv.04105(Robert Hu)
说明
文中所有交互图都自托管在本站
public/blog/flash-attention/下,是纯 HTML/CSS/JS,不依赖任何第三方脚本;它们只在文章加载时通过 iframe 引入,不会拖慢其它页面。数字与结论均来自上列论文,图示中的时间占比与倍数用于说明机制,非逐项实测。
讨论
用 GitHub 账号留言;评论保存在公开仓库chengshu-blog-discussions的 Discussions 里。也可通过 RSS 订阅后续文章。