Chengshu@skadai · 2026.10.06
RSS
3,507 字 · 1,071 词 · 约 14 分钟

从 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 个的场景。

这篇把这条线捋一遍。核心机制我都做成了可以点的交互动画,建议边看边玩。

动画 00:家族谱系。点时间轴上的节点,看每一代论文的问题、手段和代表数字。

说明

一个前提: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 访问,主导项是平方的。

拖动下面的滑块,看看这个数量级意味着什么。

动画 01:显存墙与带宽墙。拖动序列长度 N,对比中间矩阵显存与 HBM 访问量;曲线按论文 Theorem 2 的量级估算。

这也解释了为什么近似方法常常不划算:它们减少的是 FLOPs,但瓶颈在带宽。除非同时把访存也降下来,否则省下的计算时间会被搬运时间吃掉。

二、FlashAttention(2022):分块、在线 softmax、重计算

FA1 的核心只有三句话:分块(tiling)、在线 softmax(online softmax)、重计算(recomputation)。

分块与在线 softmax

softmax 天生「不局部」:要算某一行的归一化,得先知道这一整行的最大值和指数和。而分块计算时,我们一次只看到一行的一部分。在线 softmax 解决的正是这个矛盾——它维护三个逐行状态,每读入一个新块就更新一次:

  • m:到目前为止见过的最大分数;
  • ℓ:到目前为止的指数和;
  • Õ:到目前为止未归一化的输出累加。

每来一个新块,就用新的最大值把旧状态重新缩放一遍,再把新块的贡献加进去。关键在于「重新缩放」这一步:它让各个块可以按任意顺序处理,最终结果与一次性算完全一致。

点开下面的动画,它会用一个 4×6 的小矩阵,把这三个状态逐块更新给你看,最后还会和一次性算出的 softmax 结果对拍。

动画 02:分块 + 在线 softmax。点「下一步」逐块推进,观察 m、ℓ、Õ 如何更新,以及最终 Õ/ℓ 与真实 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 间通信。

动画 03:把 GPU 喂饱。切换 FA1 / FA2,调整 batch×heads 和序列长度,看 108 个 SM 被点亮多少;右侧对比两种 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 融合,几乎不花额外时间。

动画 04:Hopper 上的流水线。在 FA2 同步、warp 专门化、pingpong 调度之间切换,看 TMA、WGMMA、softmax 如何从串行变成重叠。
动画 05:FP8 与 incoherent processing。可以换一组随机数据,对比「直接量化」和「正交变换后再量化」的误差。

结果: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 于是重做了流水线,同样三招:

  1. 更大的 tile + 全异步 MMA。 Blackwell 的 MMA 块是 128×128(面积是 Hopper 64×128 的两倍),且异步 Tensor Core 直接写 TMEM。FA4 用更大的 tile 和完全异步的 MMA,把搬运、矩阵乘、softmax 尽量叠在同一段时间里。
  2. 把指数运算「软」下来。 用 FMA 单元上的多项式近似来模拟 exp,绕开吞吐最紧的多功能单元;再引入条件式 softmax 重缩放,跳过不必要的缩放操作。
  3. 降低共享内存流量。 反向把更多中间结果放进每个 SM 256 KB 的 TMEM;并使用 2-CTA MMA 模式,让每个 CTA 只加载一半的 B 操作数,减少共享内存流量和反向的原子加。
动画 06:Blackwell 的非对称扩展。切换 Hopper / Blackwell,看时间开销如何从「矩阵乘为主」变成「非 matmul 为主」,以及 FA4 的三组对策。

结果: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 序列长度。做法分三步:

  1. 把 K/V 切成若干块(只是视图,不复制数据);
  2. 每块并行跑一遍 FlashAttention,额外写出每行的 log-sum-exp;
  3. 用 log-sum-exp 重新缩放各块的部分输出,归约相加。

因为 softmax 本身就可以迭代地算,这个归约不会引入近似。代价只是一次很小的合并 kernel。

动画 07:解码期的 split-KV。切换两种模式、调整 batch×heads 与上下文长度,看 SM 占用如何被 KV 分块拉开。

另一篇 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 订阅后续文章。