🏠 总目录📚 本教程 FlashAttention ← →
📑 本页目录(点开跳转)

08 · FlashAttention

⏱ 30 分钟 | ⭐ 这个领域最漂亮的一个优化


🎯 一句话

它没有减少任何一次浮点运算,却快了好几倍。 秘密只有一句:别把那个 n×n 的中间矩阵写回显存。 —— 这是第 3 章"少搬点"思想最成功的一次实践。

朴素:整个 N×N 都写进显存显存 ∝ N²N=8k 时就是 64M 个数 × 层数FlashAttention:一次只算一小块显存 ∝ N块在片上 SRAM 里算完就丢同样的数学结果,但中间矩阵【从来没有落到显存过】⭐ 它不是「近似」——结果和朴素实现完全一致,省的是【搬运】不是计算这就是为什么它同时更快又更省显存:算力本来就没吃满,瓶颈一直是带宽
同样的数学结果,但那个 N×N 的中间矩阵从来没有落到显存过。⭐ 它不是「近似」—— 结果和朴素实现完全一致,省的是搬运不是计算。这就是它同时更快又更省显存的原因:瓶颈本来就是带宽不是算力。

😖 一、标准 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 还短 ⭐

🔨 四、怎么用(几乎不用改代码)

# 🧩 骨架:`q` 来自你自己的代码,这一段只看写法
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 优化

✅ 检查点

  1. 标准 attention 的搬运量花在哪了?占比多少?
  2. FlashAttention 减少了浮点运算次数吗?它到底优化了什么?
  3. 核心思想是什么?靠哪一级存储实现?
  4. softmax 需要全局最大值,怎么分块?代价是什么?
  5. 显存复杂度从多少降到多少?这对长上下文意味着什么?
  6. 反向传播需要 P 矩阵但没保存,怎么解决?为什么这样划算?
  7. 哪五个条件不满足会静默回退?最常见的坑是哪个?
  8. FA-2 的关键洞察是什么?它提醒了什么?
  9. FlashAttention 给出的三条一般性启发是什么?
👀 答案
  1. 花在那个 n×n 的中间矩阵上。n=4096 时中间矩阵 33.5MB,HBM 往返约 134MB,而 Q/K/V/O 本身只有 4MB——97% 的搬运量花在一个根本不需要保存的东西上。
  2. 没有减少任何一次浮点运算(实际还多做了 10-20%)。它优化的是 HBM 访问量。
  3. 分块 + 在片上完成:把 Q/K/V 切成能装进 Shared Memory 的小块,在片上算完 softmax 和乘 V,只把最终的 O 写回 HBM——那个 n×n 矩阵从来没进过 HBM。靠 Shared Memory(~228KB,比 HBM 快 5 倍)。
  4. online softmax(增量修正):维护 running max 和 running sum,每遇到更大的最大值就把之前累积的结果乘 e^(m−m_new) 打折再加新块贡献。代价:多做指数和乘法,计算量增加约 10-20%。
  5. 从 O(n²) 降到 O(n)。意味着长上下文能做到 128K——标准 attention 在 n=128K 时光中间矩阵就要 32GB(单头单样本),根本不可能。
  6. 重算:反向时用保存的 O、ℓ、m 重新算出 P。划算是因为处于带宽瓶颈,重算的时间比从 HBM 读还短。
  7. ①必须 FP16/BF16(FP32 不支持)②头维度 ≤256 ③张量连续 ④不能用任意 attention mask ⑤Ampere 及以上。最常见的坑是自定义 attention mask——手写 [B,1,n,n] 的 mask 直接禁用 FlashAttention 且静默回退。解法:用 is_causal=True,变长用 packing + cu_seqlens。
  8. Tensor Core 做矩阵乘极快,但非矩阵乘操作(指数、除法)慢得多。FA-2 重排顺序把这类操作减到最少。提醒:优化到后期,非矩阵乘操作会变成新瓶颈。
  9. ①FLOP 数不变也能大幅提速(优化不等于少算,更多是少搬)②用计算换带宽是划算的 ③要意识到内存层次的存在——同样的算法在哪一层算决定几倍的差距。

🛑 可以停在这里

⚡ 走神救援

⭐⭐ 它没有减少任何一次浮点运算,却快了几倍——秘密是别把那个 n×n 的中间矩阵写回显存。

把账摊开就很清楚:序列一长,中间矩阵的搬运量能占到总量的绝大部分,⭐ 而它是一个根本不需要保存的东西;算术强度因此远低于平衡点,是彻底的带宽瓶颈。

⭐ 核心思想:分块,在片上高速存储里把 softmax 和乘 V 一次算完,只把结果写回——那个 n×n 矩阵从来没进过显存。难点是 softmax 需要全局最大值,解法是 ⭐ 在线 softmax:维护当前的最大值和累加和,遇到更大的最大值就把已累积的结果按比例打折,再加上新块。

⚠️ 代价是多算一两成——⭐ 用计算换带宽,在带宽瓶颈下极其划算。

收益里最该记的一条不是速度:⭐⭐ 显存从随序列平方增长变成线性增长——这才是长上下文做得到的原因(标准实现在超长序列下,单头单样本的中间矩阵就要几十 GB)。⭐ 而且它数值上完全一致,不是近似。反向需要那个中间结果却没存,于是重算——比从显存读回来还快。

⚠️ 五个会让它静默回退的条件里最常踩的是 ⭐ 自定义 attention mask——手写一个稠密 mask 就直接禁用了它。处方:用框架提供的因果标志,变长序列用打包而不是 padding mask;⭐ 还可以强制指定后端,让它报错而不是静默退回慢路径。

⭐ 三条可迁移的启发:FLOP 不变也能提速(优化是少搬,不是少算)、用计算换带宽常常划算、同样的算法「在哪一层算」能差好几倍。

下一节 👉 09-显存优化全家桶.md

打卡记录保存在你的浏览器里,首页能看到总进度