📑 本页目录(点开跳转)
08 · FlashAttention
⏱ 30 分钟 | ⭐⭐ 这个领域最漂亮的一个优化
🎯 一句话
它没有减少任何一次浮点运算,却快了好几倍。 秘密只有一句:别把那个 n×n 的中间矩阵写回显存。 —— 这是第 3 章"少搬点"思想最成功的一次实践。
😖 一、标准 attention 的问题
Attention(Q,K,V) = softmax(QKᵀ / √d) · V
朴素实现:
① S = QKᵀ → [n, n] 矩阵,写回 HBM ⚠️
② P = softmax(S) → 读 [n,n],写 [n,n] ⚠️
③ O = PV → 读 [n,n] ⚠️
算一笔账(n = 4096,头维度 d = 128,BF16):
中间矩阵 S:4096 × 4096 × 2 字节 = 33.5 MB ← 【每个头、每个样本】
HBM 往返:写 S、读 S、写 P、读 P ≈ 4 × 33.5 = 134 MB
而 Q、K、V、O 本身只有 4 × 4096 × 128 × 2 = 4 MB
⭐ 搬运量的 97% 花在了那个【根本不需要保存】的中间矩阵上 💀
算术强度:
计算 ≈ 4 × n² × d = 8.6 GFLOP
搬运 ≈ 138 MB
→ 8.6e9 / 138e6 ≈ 62 FLOP/字节
A100 平衡点是 156 → 【带宽瓶颈】⚠️
🔑 注意问题的性质: 不是算得慢,是数据搬来搬去。 显存里那个 33.5MB 的矩阵, 生成出来只是为了立刻被读回去用一次,然后丢掉。
💡 二、核心思想:分块 + 在片上完成
⭐ 关键观察:Shared Memory 有 ~228 KB,比 HBM 快 5 倍(第 3 章)
→ 如果把 Q、K、V 切成小块,让每一块的 S 都能装进 Shared Memory,
就可以【在片上算完 softmax 和乘 V】,
只把最终的 O 写回 HBM ⭐
→ 那个 n×n 的矩阵【从来没有出现在 HBM 里】
for 每个 Q 块 Qᵢ: ← 外层
初始化 Oᵢ = 0, ℓᵢ = 0, mᵢ = -∞
for 每个 K,V 块 Kⱼ, Vⱼ: ← 内层
在 Shared Memory 里:
Sᵢⱼ = Qᵢ Kⱼᵀ ← 小块,装得下
更新 running max 和 running sum
Oᵢ ← 修正后的 Oᵢ + P̃ᵢⱼ Vⱼ ⭐ 增量累加
写出 Oᵢ ← 只写一次
⭐ 难点:softmax 需要全局的最大值和求和,怎么分块?
标准 softmax:softmax(x)ᵢ = e^(xᵢ - max) / Σ e^(xⱼ - max)
↑ 需要看完【所有】元素才知道
解法:online softmax(增量修正)
📐 增量修正怎么做(想看再点)
已经处理完前面的块,维护三个量: - $m$:目前见过的最大值 - $\ell$:目前的指数和(以 $m$ 为基准) - $O$:目前的输出累加
来了新块,它的最大值是 $m'$,指数和是 $\ell'$:
$$m^{new} = \max(m, m')$$
$$\ell^{new} = e^{m - m^{new}}\ell + e^{m' - m^{new}}\ell'$$
$$O^{new} = e^{m-m^{new}}O + e^{m'-m^{new}}\,\tilde P' V'$$
💡 人话:每次遇到更大的最大值,就把之前累积的结果 乘一个修正系数 $e^{m-m^{new}}$「打个折」,再加上新块的贡献。
代价:多做几次指数和乘法(计算量增加约 10–20%)。 收益:那个 n×n 矩阵永远不进 HBM。
⭐ 这就是"用计算换带宽"的典型 —— 而在带宽瓶颈下,这笔交易极其划算。
💡 online softmax 不是 FlashAttention 发明的(它更早), FlashAttention 的贡献是把它和分块、片上计算、反向重算组合成一个完整的 IO 感知算法。
📊 三、收益
HBM 访问量:
标准: O(n² + nd)
FlashAttention:O(n²d² / M) (M = Shared Memory 大小)
n=4096, d=128, M=228KB 时 → 减少约 【10-20 倍】的 HBM 访问 ⭐
| 指标 | 提升 |
|---|---|
| 速度 | 2–4 倍(序列越长越明显)⭐ |
| 显存 | 从 O(n²) 降到 O(n) ⭐⭐ |
| 数值精度 | 完全一致(不是近似算法!)⭐ |
🔑 最重要的是那个 O(n²) → O(n): 这才是长上下文能做到 128K 的原因。 标准 attention 在 n=128K 时,光中间矩阵就要 32 GB(单头单样本)—— 根本不可能。
💡 反向传播怎么办
反向需要 P 矩阵,但它没被保存
⭐ 解法:重算(recomputation)
反向时用保存下来的 O、ℓ、m 重新算出 P
→ 又是"用计算换显存"
→ 因为是带宽瓶颈,重算的时间比读 HBM 还短 ⭐
🔨 四、怎么用(几乎不用改代码)
import torch.nn.functional as F
# ⭐ 首选:PyTorch 内置,自动选最优后端(含 FlashAttention)
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
# 确认它真的走了 Flash 后端
from torch.nn.attention import sdpa_kernel, SDPBackend
with sdpa_kernel(SDPBackend.FLASH_ATTENTION):
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
# ⭐ 如果条件不满足会直接报错,而不是静默回退
⚠️ 什么情况下它会静默回退到慢路径
| 条件 | 要求 |
|---|---|
| 数据类型 | 必须 FP16 / BF16(FP32 不支持)⭐ |
| 头维度 | ≤ 256,且通常要是 8 的倍数 |
| 张量布局 | 必须连续、形状规范 |
| attention mask | 任意 mask 不支持;is_causal=True 支持 ⭐ |
| GPU 架构 | Ampere(A100)及以上 |
💥 最常见的一个坑:用了自定义 attention mask。 很多人手写一个
[B, 1, n, n]的 mask 传进去, 直接把 FlashAttention 禁用了,静默回退到最慢的实现。✅ 解法:能用
is_causal=True就用它; 变长序列用 packing + cu_seqlens(FlashAttention 原生支持),不要用 padding mask。
🧬 五、版本演进(知道区别就够)
| 版本 | 关键改进 |
|---|---|
| FA-1(2022) | 提出分块 + online softmax + 重算 |
| FA-2(2023) | 减少非矩阵乘操作、更好的并行切分 → 再快约 2 倍 ⭐ |
| FA-3(2024) | 针对 H100:异步、FP8、warp 特化 → 再快约 1.5–2 倍 |
💡 FA-2 的一个关键洞察: Tensor Core 做矩阵乘极快,但非矩阵乘操作(指数、除法)慢得多。 FA-1 里的 online softmax 有大量这类操作,FA-2 重新安排了顺序把它们减到最少。 这提醒我们:优化到后期,"非矩阵乘操作"会变成新的瓶颈。
🧠 六、它给你的一般性启发
⭐ FlashAttention 教给你的三件事:
① 【FLOP 数不变也能大幅提速】
→ 优化不等于"少算",更多时候是"少搬"
② 【用计算换带宽是划算的】
→ 在带宽瓶颈下,重算比读显存快
③ 【要意识到内存层次的存在】
→ 同样的算法,"在哪一层算"决定了几倍的差距
🔗 同一个思路的其他应用: - 算子融合(第 7 章)—— 中间结果留在片上 - 梯度检查点(第 9 章)—— 用重算换激活显存 - PagedAttention(第 17 章)—— 推理侧的同类思想
🔗 和站内其他章的关系
| 相关的地方 | 这里的位置 |
|---|---|
| 第 3 章 Shared Memory 快 5 倍 | FlashAttention 的全部依据 ⭐ |
| 第 3 章 attention 是带宽瓶颈 | 本章是它的解法 |
| 第 7 章 算子融合 | 同一思想,FlashAttention 是最成功的案例 |
| 全景导论第 2 章 attention 公式 | 本章优化的就是它 |
| 第 16 章 KV Cache | 推理侧的 attention 优化 |
✅ 检查点
- 标准 attention 的搬运量花在哪了?占比多少?
- FlashAttention 减少了浮点运算次数吗?它到底优化了什么?
- 核心思想是什么?靠哪一级存储实现?
- softmax 需要全局最大值,怎么分块?代价是什么?
- 显存复杂度从多少降到多少?这对长上下文意味着什么?
- 反向传播需要 P 矩阵但没保存,怎么解决?为什么这样划算?
- 哪五个条件不满足会静默回退?最常见的坑是哪个?
- FA-2 的关键洞察是什么?它提醒了什么?
- FlashAttention 给出的三条一般性启发是什么?
👀 答案
- 花在那个 n×n 的中间矩阵上。n=4096 时中间矩阵 33.5MB,HBM 往返约 134MB,而 Q/K/V/O 本身只有 4MB——97% 的搬运量花在一个根本不需要保存的东西上。
- 没有减少任何一次浮点运算(实际还多做了 10-20%)。它优化的是 HBM 访问量。
- 分块 + 在片上完成:把 Q/K/V 切成能装进 Shared Memory 的小块,在片上算完 softmax 和乘 V,只把最终的 O 写回 HBM——那个 n×n 矩阵从来没进过 HBM。靠 Shared Memory(~228KB,比 HBM 快 5 倍)。
- online softmax(增量修正):维护 running max 和 running sum,每遇到更大的最大值就把之前累积的结果乘 e^(m−m_new) 打折再加新块贡献。代价:多做指数和乘法,计算量增加约 10-20%。
- 从 O(n²) 降到 O(n)。意味着长上下文能做到 128K——标准 attention 在 n=128K 时光中间矩阵就要 32GB(单头单样本),根本不可能。
- 重算:反向时用保存的 O、ℓ、m 重新算出 P。划算是因为处于带宽瓶颈,重算的时间比从 HBM 读还短。
- ①必须 FP16/BF16(FP32 不支持)②头维度 ≤256 ③张量连续 ④不能用任意 attention mask ⑤Ampere 及以上。最常见的坑是自定义 attention mask——手写
[B,1,n,n]的 mask 直接禁用 FlashAttention 且静默回退。解法:用is_causal=True,变长用 packing + cu_seqlens。 - Tensor Core 做矩阵乘极快,但非矩阵乘操作(指数、除法)慢得多。FA-2 重排顺序把这类操作减到最少。提醒:优化到后期,非矩阵乘操作会变成新瓶颈。
- ①FLOP 数不变也能大幅提速(优化不等于少算,更多是少搬)②用计算换带宽是划算的 ③要意识到内存层次的存在——同样的算法在哪一层算决定几倍的差距。
🛑 可以停在这里
⚡ 走神救援
⭐⭐它没减少任何一次浮点运算却快了几倍——秘密是别把 n×n 中间矩阵写回显存。标准 attention 在 n=4096 时中间矩阵 33.5MB,HBM 往返 134MB,而 Q/K/V/O 本身只有 4MB → ⭐97% 的搬运量花在一个根本不需要保存的东西上,算术强度 62 < 平衡点 156 = 带宽瓶颈。⭐核心思想:分块 + 在 Shared Memory(228KB,比 HBM 快 5 倍)里算完 softmax 和乘 V,只把 O 写回——n×n 矩阵从来没进过 HBM。难点是 softmax 需要全局最大值,解法是 ⭐online softmax(维护 running max/sum,遇到更大的最大值就把已累积的结果乘 e^(m−m_new) 打折再加新块);代价是多算 10-20%——用计算换带宽,在带宽瓶颈下极其划算。收益:HBM 访问少 10-20 倍、速度 2-4 倍、⭐⭐显存从 O(n²) 降到 O(n)(这才是长上下文能做 128K 的原因——标准实现在 128K 时单头单样本的中间矩阵就要 32GB)、数值完全一致不是近似。反向需要 P 但没存 → 重算(比读 HBM 还快)。⚠️五个静默回退的条件:必须 FP16/BF16、头维度≤256、张量连续、不能用任意 attention mask、Ampere 以上;⭐最常见的坑是自定义 mask(手写 [B,1,n,n] 直接禁用 FA)→ 用
is_causal=True,变长序列用 packing + cu_seqlens 不要 padding mask。用法:F.scaled_dot_product_attention,可用sdpa_kernel(SDPBackend.FLASH_ATTENTION)强制报错而非静默回退。FA-2 快 2 倍的洞察:Tensor Core 矩阵乘极快但指数/除法慢 → 优化到后期非矩阵乘操作会变成新瓶颈。⭐三条启发:FLOP 不变也能提速(优化=少搬不是少算)、用计算换带宽划算、同样的算法"在哪一层算"决定几倍差距。
下一节 👉 09-显存优化全家桶.md