FlashAttention
11-Self-Attention 机制 算过一笔账:
它的思路和前面所有优化都不一样:不减 FLOPs,只减访存。
标准 Attention 的瓶颈在哪
| 指标 | 量级 | 说明 |
|---|---|---|
| 中间矩阵 | 各需 | |
| HBM 读写量 | 反复读写 | |
| 计算量 | 两次矩阵乘 |
算术强度:
标准 Attention 是彻头彻尾的访存受限操作,计算单元大部分时间在等数据。
两个关键技术,缺一不可
| 技术 | 解决什么 |
|---|---|
| Tiling(分块) | 如何不生成完整的 |
| Online Softmax | 分块计算时,Softmax 依赖全局信息怎么办 |
第二条是关键。 Safe Softmax 需要先扫一遍求 max、再扫一遍求 sum、最后才能归一化 —— 这个顺序依赖与分块策略直接矛盾。Online Softmax 的递推(见 07-Softmax 与 Online Softmax)让 max 与 sum 可以边加载新块边修正,这是 FlashAttention 能成立的前提。
FlashAttention 本质上就是「Online Softmax + 把
也融进同一遍循环」。 后者的具体形式是输出 的增量修正公式(见 07-Softmax 与 Online Softmax 里 Online Softmax 的一节)。
IO 复杂度:从 到
| 算法 | HBM 访问量 |
|---|---|
| 标准 Attention | |
| FlashAttention |
(
- 由块大小约束
,外循环共 次 - 每次外循环要遍历完整的
与 (共 个元素)→ 的总访问量 、 整体只读一次, ,相对前者可忽略
SRAM 越大(
越大),每次能装下的块越大,外循环次数越少,对 的重复读写就越少。 很直观:这是「片上空间换 HBM 流量」的直接表述。
而且这个量级是最优的 —— 该文证明了任何精确 Attention 算法的 HBM 访问量下界是
典型值:
反向传播靠重计算
前向不保存
重计算的代价是可以算清的:前向本来要算两次矩阵乘(
这与「梯度检查点」是同一个思想(见 01-GPU 硬件架构与存储层次),只是作用在 Attention 这个特定算子上。
两种算法把中间矩阵放在哪里:
标准 Attention:N×N 的中间矩阵必须落到 HBM
┌───────────── HBM ─────────────┐
│ Q、K、V │
│ S = QKᵀ (N×N,写回 HBM) │
│ P = softmax(S) (N×N) │
│ O = PV (N×d) │
└────────────────────────────────┘
HBM 读写量 O(N²);算术强度只有 O(d),而 d 通常只有 64–128
FlashAttention:N×N 从头到尾不落地
┌──────────── SRAM(片上)────────────┐
│ Q 块 ← 外循环沿 Q │
│ └─ 内循环遍历 K/V 块 │
│ └─ 算 S 块 → Online 更新(m,d)│
│ → 增量修正累积的 O 块 │
└─────────────────────────────────────┘
HBM 读写量 O(N²d²/M),M 是 SRAM 大小
该文证明了任何精确 Attention 的 HBM 访问下界就是 Ω(N²d²/M)
⇒ 标准 Attention 减的是 FLOPs 之外的东西减不掉,FlashAttention 只减访存。V1 的实际效果
| 序列长度 | 加速比 | 备注 |
|---|---|---|
| 1024 | 2.4× | |
| 2048 | 2.8× | |
| 4096 | 3.5× | |
| 8192 | — | 标准实现在此 OOM,FA 可运行 |
最后一行才是本质收益:价值在「原来跑不了的现在能跑」,速度只是附带。
但 V1 在 A100 上只达到了理论 FLOPS 的 25%–40% —— 对一个「减少 IO 后应该变成 Compute-Bound」的算子来说,这个利用率不理想。
V2 的三个改进
瓶颈按影响从大到小:
| 问题 | 原因 | 影响 |
|---|---|---|
| 循环顺序不佳 | 外循环遍历 | 额外 HBM 流量,且无法跨 |
| 非 GEMM 运算占比高 | Softmax rescaling 有大量逐元素操作 | Tensor Core 利用率低 |
| Warp 切分维度不当 | 沿 | 需要频繁的 Warp 间通信与归并 |
改进一:调换循环顺序
外循环沿
块加载一次就够 不用反复读 - 减少
的重复读写 —— 每个 Q 块对应自己的一行输出,在自己的外循环迭代内累积完 - 天然支持跨
并行 —— 不同的 Q 块可以独立分给不同 Thread Block,这才是并行度的来源
V2 把外循环从 K/V 换到 Q,一个改动三个收益:
V1:外循环沿 K/V V2:外循环沿 Q
┌───────────────────────────┐ ┌─────────────────────────────┐
│ for each K/V block: │ │ for each Q block: │
│ 载入 Q 的全部 │ │ 载入这一块 Q(一次就够) │
│ 更新 O 的全部(读 + 写) │ │ 内循环遍历 K/V │
│ ⇒ Q 被反复加载 │ │ 在这一轮里把 O 块累积完 │
│ ⇒ O 块被反复读写 │ │ ⇒ Q 只读一次 │
│ ⇒ 不同 Q 块之间有依赖, │ │ ⇒ O 块不重复读写 │
│ 无法跨 Q 并行 │ │ ⇒ 不同 Q 块完全独立,可并行 │
└───────────────────────────┘ └─────────────────────────────┘
⇒ 并行度的来源正是第三条:Q 块可以独立分给不同的 Thread Block改进二:延迟 Rescaling
V1 每个
V2 只维护未归一化的累积:
收益是精确可数的:V1 在内循环每次迭代要做「旧
除法在 GPU 上比乘法贵得多,而内循环迭代
改进三:Warp 沿 分配
V1 沿
V2 改成沿
这三条改进有一个共同点:都在减少「非必要的同步与数据移动」,而不是减少计算量。与前面的结论一致。
Causal Mask 的块级跳过
自回归模型里第
V2 按块判断三类情况:
| 块的位置 | 处理 |
|---|---|
| 完全在 mask 下方(Q 块最小行号 ≥ K 块最大列号) | 全可见,正常计算 |
| 跨越对角线 | 逐元素应用 mask |
| 完全在 mask 上方(Q 块最大行号 < K 块最小列号) | 直接跳过 |
对角线以下(全可见) 对角线块(部分 mask) 对角线以上(全跳过)
┌─────────────┐ ┌─────────────┐ ┌─────────────┐
│ ■ ■ ■ ■ ■ ■ │ │ ■ ■ ■ □ □ □ │ │ □ □ □ □ □ □ │
│ ■ ■ ■ ■ ■ ■ │ │ ■ ■ ■ ■ □ □ │ │ □ □ □ □ □ □ │
│ ■ ■ ■ ■ ■ ■ │ │ ■ ■ ■ ■ ■ □ │ │ □ □ □ □ □ □ │
└─────────────┘ └─────────────┘ └─────────────┘配合改进一:外循环沿
V2 的性能
TFLOPS(A100)
| 序列长度 | V1 | V2 | V2 利用率 |
|---|---|---|---|
| 1024 | 124 | 196 | 63% |
| 2048 | 136 | 218 | 70% |
| 4096 | 141 | 227 | 73% |
| 8192 | 138 | 222 | 71% |
V2 相对 V1 约 1.6×,利用率从 25–40% 提到 63–73%。
与其他实现对比(A100-80GB, , )
| 实现 | 前向 | 反向 |
|---|---|---|
| PyTorch 标准 | 1.0× | 1.0× |
| FlashAttention V1 | 2.8× | 2.5× |
| FlashAttention V2 | 4.3× | 3.9× |
| xFormers (cutlass) | 3.1× | 2.8× |
长序列的显存
这张表不是「单个矩阵的存储成本」
它估算的是训练时与 Attention 相关的激活显存总量(含中间张量与梯度缓存,按典型多头配置粗估)。列表的目的是看
| 序列长度 | 标准 Attention | FlashAttention V2 |
|---|---|---|
| 4K | 128 MB | 4 MB |
| 16K | 2 GB | 16 MB |
| 64K | 32 GB(OOM) | 64 MB |
反向的一处细节
V2 在反向里保存 logsumexp 而不是分离的
一个量代替两个量,省一份存储;且反向里处处用到的是
后续演进
V1 与 V2 是算法与并行策略的改进。往后的路线转向吃满新硬件特性:
| 方向 | 内容 |
|---|---|
| FlashAttention-3 | 利用 Hopper 的 wgmma(异步、操作数从 SMEM 直进 Tensor Core)、TMA、FP8 |
| Decode 阶段的加速 | Flash-Decoding / FlashDecoding++ / FlashInfer —— Decode 是 batch 小、序列长的形态,与训练的前向并行结构不同 |
| 与推理引擎结合 | PagedAttention 在 GPU 上的分页寻址,见 02-PagedAttention:KV Cache 的分页管理 |
训练侧(长序列前向 + 反向)与推理侧(Decode)的优化方向不同 —— 训练要的是吞吐与显存,Decode 要的是低延迟与小 batch 下的并行度。这一条区分是「一个 Attention 优化方案能不能用在 Serving 上」的判据。
相关
- 07-Softmax 与 Online Softmax —— 本篇的数学前提,含输出修正公式的完整推导
- 11-Self-Attention 机制 ——
矩阵的显存算例与算术强度 - 06-GEMM 性能优化 —— Attention 里两次矩阵乘的优化基础
- 01-GPU 硬件架构与存储层次 —— SRAM 容量、Roofline 与 Tensor Core
- 05-Attention 后端与图优化 —— 推理引擎里 Attention 后端的工程实现
参考
- https://caomaolufei.github.io/AIInfraGuide/guides/模块二-cuda编程与算子优化/61-flashattention-v1详解
- Dao et al., https://arxiv.org/abs/2205.14135
- https://arxiv.org/abs/2307.08691
- https://arxiv.org/abs/2407.08608
- https://arxiv.org/abs/1805.02867
YJ