🏠 总目录📚 本教程 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 还短 ⭐

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

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 中间矩阵写回显存。标准 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

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