Featured image of post FlashAttention:不是少算,而是少搬

FlashAttention:不是少算,而是少搬

FlashAttention 没有近似注意力,也没有消除二次计算;它通过分块、在线 softmax 与重计算避免物化巨大的注意力矩阵。本文从 GPU 内存层级解释它为何更快,以及 FA2、FA3 如何继续榨出硬件效率。

FlashAttention 有一个很反直觉的事实:它计算的仍是精确 attention,浮点运算量仍随序列长度平方增长,反向传播甚至会主动重算一部分结果;但它通常更快,也把额外显存从二次量级降到线性量级。

如果只用“大 O 复杂度”理解算法,这件事很难解释。答案不在公式少了哪一项,而在 GPU 真正为哪一种成本付时间:现代加速器能很快地做矩阵乘法,却不喜欢把巨大的中间张量反复写入高带宽显存(HBM),再读回来交给下一个 kernel。FlashAttention 的贡献,是把 attention 从一串按数学步骤排列的算子,重排成一条适配内存层级的数据流。

它没有改变模型学什么,而是改变同一公式怎样落到机器上。也正因为如此,这项工作比某个 CUDA 技巧更有解释价值:它提醒我们,模型效率既取决于计算图,也取决于数据在硬件中的旅程。

慢的未必是乘法,而是中间矩阵的往返

单个注意力头的计算可以写成:

$$S=QK^\top,\qquad P=\operatorname{softmax}(S),\qquad O=PV$$

若序列长度为 $N$、head dimension 为 $d$,$Q、K、V、O$ 的尺寸都是 $N\times d$,而分数矩阵 $S$ 与概率矩阵 $P$ 都是 $N\times N$。标准实现通常让不同 kernel 依次完成矩阵乘、mask、softmax、dropout 和第二次矩阵乘:先把 $S$ 写入 HBM,softmax 再把它读出并写回 $P$,最后第三次读出 $P$ 与 $V$ 相乘。

问题不只是 $N^2$ 个元素占空间。每跨过一个 kernel 边界,中间结果都可能要经过一次“片上计算单元 ↔ HBM”的长途搬运。FlashAttention 原论文用 A100 举例:每个流多处理器只有约 192 KB 片上 SRAM,但带宽估算约 19 TB/s;整卡 HBM 有 40–80 GB,带宽却只有 1.5–2.0 TB/s。SRAM 快一个数量级,却小了许多个数量级。于是关键问题变成:能否让小块数据留在片上完成整段计算,而不是让完整 $N\times N$ 矩阵落地到 HBM?

图 1:标准实现物化 S、P 并多次访问 HBM;FlashAttention 只让小分块进入 SRAM

这里需要区分两个经常混在一起的“内存”。FlashAttention 并不会让输入 $Q、K、V$ 或最终输出凭空消失;它减少的是注意力内部中间量的 HBM 读写与保存。原论文证明,在常见 SRAM 大小范围内,标准 attention 的 HBM 访问量为 $\Theta(Nd+N^2)$,FlashAttention 则为 $\Theta(N^2d^2/M)$,其中 $M$ 是片上 SRAM 容量。计算量仍是 $O(N^2d)$,但数据搬运明显减少。

难点不是分块,而是 softmax 看起来需要整行

矩阵乘法天然适合 tiling:把 $Q、K、V$ 切成能放进 SRAM 的小块,逐块相乘即可。softmax 却把一整行耦合在一起。对分数向量 $x$,稳定 softmax 需要先知道全局最大值,再计算所有指数的总和:

$$m=\max_i x_i,\qquad \ell=\sum_i e^{x_i-m},\qquad \operatorname{softmax}(x)_i=\frac{e^{x_i-m}}{\ell}$$

如果必须先保存完整一行才能得到 $m$ 与 $\ell$,$N\times N$ 中间矩阵仍然躲不掉。FlashAttention 使用 online softmax,把“看完整行再归一化”改成“每读一块,就更新足够的统计量”。处理新分块 $x_b$ 时,只维护到目前为止的最大值 $m$、归一化和 $\ell$,以及尚未最终归一化的输出累积量 $o$:

$$ \begin{aligned} m' &= \max(m,\max x_b)\\ \ell' &= e^{m-m'}\ell + \sum_j e^{x_{b,j}-m'}\\ o' &= e^{m-m'}o + \sum_j e^{x_{b,j}-m'}v_{b,j} \end{aligned} $$

最后用 $O=o/\ell$ 得到结果。新分块如果带来更大的最大值,旧累积量会按 $e^{m-m’}$ 重新缩放;因此无论分成多少块,最终值与整行 stable softmax 相同,只存在正常的浮点舍入差异。它不是近似 attention,也不是稀疏 attention。

图 2:online softmax 逐块维护 m、ℓ、o,最终仍得到完整精确 attention

在实际 kernel 中,一个 $K/V$ tile 被加载进 SRAM 后,会与若干 $Q$ tile 完成 $QK^\top$、mask、softmax 更新和与 $V$ 的乘法。局部的 score tile 用完即弃,只有输出块与少量归一化统计量需要保留。这样,FlashAttention 把原本跨多个 kernel 的流水线融合成一个 I/O 感知 kernel。

为什么“多算一点”反而更快

训练的反向传播通常需要前向时的 $S$ 和 $P$。标准做法把它们保存到 HBM;FlashAttention 只保存输出与每行的 softmax 归一化统计量,反向时在 $Q、K、V$ 分块已经进入 SRAM 后重新算出局部 $S$ 和 $P$。

这是一笔明确的交换:增加矩阵乘 FLOPs,换掉大规模 HBM 写入、读取与驻留。对一台矩阵乘吞吐远高于内存搬运效率的 GPU,这笔交换可能同时省时间和显存。原论文把额外内存需求从 $O(N^2)$ 降到 $O(N)$;FlashAttention-2 论文报告,反向阶段因不保存 $S、P$ 可节省约 10–20 倍内存,并获得 2–4 倍 wall-clock 加速。倍率依赖序列长度、head dimension、mask、数据类型与硬件,不能当作固定承诺。

这也解释了一个常见误区:重计算不必然更慢。只有当“算力”是主要瓶颈时,少算才一定占优;当算子受内存访问限制时,宁可在片上重算,也可能比从 HBM 取回旧结果更便宜。FlashAttention 的核心不是某一条 softmax 公式,而是用正确的成本模型选择保存什么、搬运什么、重算什么。

从 FA1 到 FA3:优化对象一层层靠近硬件

第一代解决了最昂贵的 HBM 往返后,attention 仍未接近 GPU 的矩阵乘峰值。后续版本没有改变“分块 + online softmax + 重计算”的主干,而是继续处理更细的闲置和同步成本。

FlashAttention-2 关注工作如何分给 GPU。A100 的矩阵乘吞吐远高于非矩阵运算,论文给出的理论峰值分别是 312 TFLOPs/s(FP16/BF16 matmul)与 19.5 TFLOPs/s(FP32 non-matmul)。因此 FA2 减少不必要的缩放等非 matmul 操作;除 batch 和 head 外,还沿序列长度把工作分给更多 thread blocks,改善长序列、小 batch 时的 occupancy;在一个 block 内重新分配 warps,减少 shared memory 通信。论文报告 FA2 相对 FA1 约快 2 倍,A100 前向最高达到理论峰值的 73%,GPT 风格模型训练最高达到每张 A100 225 TFLOPs/s。

FlashAttention-3 则针对 Hopper。H100 的 Tensor Core 矩阵乘、TMA 数据搬运与普通 CUDA core 可以异步工作;若软件仍按“搬完再算、算完再 softmax”的顺序执行,专用单元会互相等待。FA3 用 warp specialization 分开 producer 与 consumer,用 ping-pong pipeline 重叠下一块搬运、矩阵乘和当前块 softmax,并为 FP8 加入 block quantization 与 incoherent processing 来抑制离群值造成的量化误差。论文在 H100 上报告:FP16 前向较 FA2 快 1.5–2.0 倍、最高 740 TFLOPs/s;FP8 接近 1.2 PFLOPs/s,并在其测试中把基线 FP8 attention 的数值误差降低 2.6 倍。

图 3:FA1 减少 HBM 往返,FA2 改进并行分工,FA3 用异步流水线与低精度适配 Hopper

这条演进线很有代表性:先修正算法的数据路径,再修正 thread block/warp 的分工,最后把执行计划对齐到具体硬件的异步单元。性能不是“用了 GPU”自动得到的,而是算法与每一代 GPU 重新协商的结果。

它改变了什么,又没有改变什么

第一,FlashAttention 不是新的注意力模式。 模型权重、dense attention 的语义和输出都不需要改变;同一个模型可以在兼容条件下替换 kernel。它与 Linformer、稀疏 attention、线性 attention 的根本区别,是后者改变或近似了要计算的关系,FlashAttention 只改变执行顺序。

第二,它没有消除二次计算。 训练或长 prompt 的 prefill 仍要计算大量 token 两两关系。FlashAttention 让这些工作更接近硬件可承受的方式,并避免二次大小的中间激活,却不能让任意长上下文变成免费。上下文翻倍时,dense attention 的 FLOPs 仍大约变为四倍。

第三,它与 KV cache 优化解决的是不同账单。 在单 token 自回归 decode 中,不再形成大块 query-by-key 矩阵,读取不断增长的 KV cache 往往更突出。GQA 减少 KV head 数,MLA 压缩每个 token 的缓存表示,Mamba 改成固定状态;FlashAttention 主要优化 attention kernel 内部的数据流。它们可以互补,却不能互相替代。官方实现也不断增加 causal、GQA/MQA、variable length、paged attention 等路径,正说明“attention”并非一个固定形状的 kernel。

第四,峰值倍率不是应用端倍率。 一个训练 step 还包含 MLP、通信、优化器与数据加载;一次在线推理还包含调度、采样和服务框架。论文中的 attention microbenchmark、单卡 TFLOPs/s 与端到端吞吐回答的是不同问题。判断收益时必须同时说明硬件、精度、序列形状、因果 mask 和比较基线。

真正的突破,是把内存层级写进算法

FlashAttention 的公式结果并不新,分块和 online softmax 也各有前史。它的突破在于把这些技术组成一个完整、可证明、可实现的 I/O 感知 attention:不物化 $N\times N$ 中间矩阵,正向与反向都围绕片上 tile 组织,并用开源 kernel 把理论节省变成 wall-clock 收益。官方仓库后来覆盖 NVIDIA 与 AMD 后端,主流框架也把 fused scaled-dot-product attention 变成常用执行路径;这类采用比一次榜单领先更能说明其影响。

更深的一层判断是:长上下文能力从来不只是一条模型曲线。窗口能否真正扩展,还取决于激活是否放得下、数据能否喂得动、kernel 能否利用硬件、推理时 KV cache 如何管理。FlashAttention 没有解决所有这些问题,但它消掉了其中一个曾经极其昂贵的中间物。

所以“不是少算,而是少搬”并非一句性能口号。它代表一种设计方法:不要只问算法做了多少 FLOPs,还要问每个字节在哪一层内存、被读取几次、何时值得保存,以及能否用廉价重算换掉昂贵搬运。模型规模继续增长后,这种问题只会更重要。

参考资料

  1. Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness, NeurIPS 2022.
  2. Dao, FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning, ICLR 2024.
  3. Shah et al., FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision, 2024.
  4. Dao AI Lab, FlashAttention official implementation.
  5. Tri Dao, FlashAttention-3 technical blog, 2024.