自 2022 年 5 月 Tri Dao 等人的论文发表后,FlashAttention 几乎成了大模型训练与推理栈的默认注意力底座。它没有改动注意力的数学定义,也没做任何近似,却同时把速度和显存都改善了一个量级。这篇文章讲清它的来龙去脉:问题出在哪、它怎么改、边界在哪。概念层面的 Q、K、V 本文不再展开(站内有概念篇专文)。
它解决什么问题
标准自注意力是两步矩阵乘:先算 S = QKᵀ/√d 得到 N×N 的注意力分数矩阵,softmax 归一化后再乘 V 得到输出。计算量 O(N²d) 是明账,但慢的根源不在乘法本身,而在那个 N×N 的中间矩阵必须完整地写进显存(HBM),随后再读回来做第二次乘法。序列长度 4K、单头 128 维时,这一个矩阵就是 1600 万个元素;batch 一大,光是中间结果的「写出去再读回来」,访存量就远超 Q、K、V 本身的输入输出。
GPU 的算力这些年涨得比显存带宽快得多,而注意力恰好是算术强度(每读一字节做多少浮点运算)很低的活——瓶颈卡在 HBM 的读写带宽上,而不是 Tensor Core 的算力上。所以正确的优化目标不是减少 FLOPs(计算量一样是 O(N²d)),而是减少访存次数。FlashAttention 的出发点就一句话:算得快不如搬得少。
GPU 内存层级:快而小的 SRAM,大而慢的 HBM
FlashAttention 论文里给了一组 A100 的数字:108 个流式多处理器(SM),每个配 192KB 片上 SRAM,合计约 20MB,带宽约 19TB/s;HBM 显存 40-80GB,带宽 1.5-2.0TB/s。SRAM 比 HBM 快一个数量级,却小了三个数量级——这就是 GPU 上的「寄存器与磁盘」关系。
用 roofline 的直觉看:一个 kernel 的实际速度上限是 min(峰值算力, 带宽 × 算术强度)。矩阵乘的算术强度高,能把 Tensor Core 喂饱;标准注意力因为要反复物化 N×N 矩阵,算术强度被拉低,实际速度贴着「带宽 × 算术强度」这条线走。把计算搬进 SRAM、让数据少在 HBM 里打转,就成了唯一的出路。
核心机制:tiling 分块 + online softmax
困难看起来在于 softmax 的分母:softmax 要对整行分数做归一化,似乎必须先拿到全部 N 个分数,才算得出任何一个输出——这正是「注意力矩阵必须整体物化」的表面理由。
online softmax(最早见于 NVIDIA 的 Milakov 与 Gimelshein 2018 年的论文 Online normalizer calculation for softmax)打破了这一点:归一化所需的统计量可以滚动地算出来。对一行分数 x,数值稳定的 softmax 是:
m = max(x),ℓ = Σ e^(x_j − m),输出第 i 项 = e^(x_i − m) / ℓ
关键在于:如果分数分块到达,每处理一块就维护「目前为止的最大值 m 与指数和 ℓ」;当新块带来更大的最大值时,旧块积累的部分只需整体乘一个重标定因子 α = e^(m_旧 − m_新):
ℓ ← α·ℓ + Σ_j∈新块 e^(x_j − m_新)
直观理解:旧块的指数是按旧最大值 m_旧 缩放过的,现在基准换成更大的 m_新,统一补乘 e^(m_旧 − m_新) 就回到同一把尺子上。注意力还要对 V 加权求和,于是再维护一个输出累加器 O,重标定时 O ← α·O,然后累加新块的贡献,最后 O / ℓ 一步归一化。每一步都是恒等变形,没有任何元素被丢弃或近似,数学上与一次性算完整行 softmax 完全等价;浮点上唯一的差别是求和顺序,带来的是可忽略的微小数值差异,而不是精度损失——这是它与各种稀疏、近似注意力方法的本质区别。
FlashAttention 把它做成 kernel(tiling 分块计算):Q、K、V 不再整体落盘,Q 的每一块留在 SRAM,K、V 沿序列方向逐块流入,在片上完成「算分数 → online softmax → 加权累加」的整条流水,全程只把最终结果 O 写回 HBM。反向传播同样只需 O(N) 的额外空间:保存每行的 (m, ℓ) 统计量,需要时按统计量重算分数,不必存下 N×N 矩阵。
flowchart TD
A["HBM 中的 Q、K、V(只读一遍)"] --> B["取 Q 的第 i 块载入 SRAM"]
B --> C["取 K/V 的第 j 块载入 SRAM"]
C --> D["片上算分数块 S_ij = Q_i K_j^T"]
D --> E["online softmax:更新 m 与 ℓ,按 α = e^(m_旧 − m_新) 重标定 ℓ 与累加器 O,累加新块贡献"]
E --> F{"K/V 还有分块?"}
F -- 有 --> C
F -- 没有 --> G["O_i = O / ℓ,写回 HBM"]
G --> H{"Q 还有分块?"}
H -- 有 --> B
H -- 没有 --> I["输出 O,全程不物化 N×N 矩阵"]
效果:显存从 O(N²) 到 O(N)
显存上,注意力部分从 O(N²) 降到 O(N)。官方仓库 README 给的参考:序列 2K 时约省 10 倍显存,4K 时约 20 倍。速度上,按第一版论文(arXiv 2205.14135)的报告:BERT-large(序列 512)端到端提速 15%,GPT-2(序列 1K)提速 3 倍,Long-Range Arena(序列 1K-4K)提速 2.4 倍。
规律很明显:收益随序列长度增长——序列越长,N×N 矩阵的访存占比越大,省得越多;短序列下注意力在整个模型里占比有限,收益自然小。
演进:FA2 修并行度,FA3 修 Hopper
FlashAttention-2(arXiv 2307.08691,2023)解决的是「访存省下来了,算力却没吃满」:FA1 在 A100 上只用到理论峰值 FLOPs 的 25-40%。FA2 做了三件事:减少非矩阵乘 FLOPs(softmax 相关运算比矩阵乘慢得多);把并行铺到序列维度——即使单个 head 也能铺满所有 SM,长序列下尤其重要;重排 thread block 内 warp 之间的分工,减少走共享内存的通信。结果约 2 倍于 FA1,A100 上达到理论峰值的 50-73%,GPT 风格模型训练每卡 225 TFLOPs/s(72% 的模型 FLOPs 利用率)。
FlashAttention-3(arXiv 2407.08608,2024)面向 Hopper 架构(H100)。新硬件把 Tensor Core 和 TMA(Tensor Memory Accelerator)做成了异步的,而 FA2 的同步流水线在 H100 上只能用到 35% 的峰值利用率。FA3 用三招吃满异步能力:warp specialization(warp 分成生产者与消费者,数据搬运与计算重叠)、以块为单位交错执行 GEMM 与 softmax、启用 FP8 低精度(配合 incoherent processing 抑制量化误差)。H100 上达到 FA2 的 1.5-2 倍:BF16 约 740 TFLOPs/s(75% 利用率),FP8 接近 1.2 PFLOPs/s,且 FP8 误差比基线 FP8 attention 实现低 2.6 倍。
边界在哪
- 短序列收益小。序列 512 以内,注意力占整个模型的开销有限(BERT-large 也只有 15%),FlashAttention 不是万金油;它的主场是长序列 prefill 与训练。
- decode 阶段帮不上多少。自回归解码每步只有一个 query 位置,注意力退化为对整个 KV Cache 的访存密集小运算——本质上仍是带宽问题,但瓶颈在「读越来越长的 KV Cache」,这归显存管理管(那是 vLLM 与 PagedAttention 的故事,站内另有专文)。FlashAttention 优化的 Q 长的那一侧。
- 不省计算量。它是精确注意力,FLOPs 仍是 O(N²),省的是访存与显存;要把计算量真正降到亚二次,得走稀疏/线性注意力等近似路线,那是另一条路。
- kernel 实现复杂度高。三代实现全是手写 CUDA,且各自吃不同代 GPU 的架构特性;主流框架已经把它集成为默认选项屏蔽了大部分麻烦,但在新硬件上复刻这种 kernel 依然是重活。
- 非 NVIDIA 硬件支持有限。截至 2026-10-07,官方仓库的 CUDA 版要求 Ampere、Ada 或 Hopper(A100、RTX 3090/4090、H100),Turing 不在主仓库支持范围(另有仅支持功能子集的独立仓库);AMD 侧通过 ROCm 6.0+ 提供 FA2 级实现(Composable Kernel 与 Triton 两个后端,覆盖 MI200/MI300 系列与 RDNA 3/4 消费卡);FA3 目前只面向 H100/H800(beta 状态,FP8 仅前向)。
结语
FlashAttention 是一个教科书式的系统优化案例:算法没变,变的只是数据在内存层级里的流动方式。它把一个看似必须「先全算完再归一化」的操作,改写成可以分块流式完成的扫描——凭一个 2018 年就有的 online softmax 技巧,加上对 GPU 内存层级的清醒认识,成了之后几乎所有大模型基础设施的默认底座。读懂它,也就读懂了「性能瓶颈在访存不在算力」这条系统领域的老经验如何在 AI 时代重演。
参考资料
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(arXiv 2205.14135)
- FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning(arXiv 2307.08691)
- FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision(arXiv 2407.08608)
- Dao-AILab/flash-attention(GitHub 仓库,硬件支持与安装现状)
- Online normalizer calculation for softmax(arXiv 1805.02867,online softmax 出处)
读者留言
COMMENTS 暂无还没有留言,来说第一句?