FlashAttention 入门:从 FA1 到 FA4 的 IO 感知注意力

FlashAttention 入门:从 FA1 到 FA4 的 IO 感知注意力

注意力层是把 Transformer 扩展到长序列时的主要瓶颈:它的时间与内存开销都随序列长度平方增长。FlashAttention 是一族让注意力保持精确(exact)的同时显著提速、省内存的算法,其核心理念是把 GPU 的内存流量——而不只是 FLOPs——当作需要优化的资源。本文用原创的图示与推导讲清核心思想:GPU 内存层级、分块(tiling)、在线 softmax(online softmax)、反向重计算(recomputation),然后是 FA1 → FA2 → FA3 → FA4 的时间线、独立 flash-attn 包与 PyTorch scaled_dot_product_attention 的区别,以及实践中真正要紧的硬件注意事项。

本文有两条规则。第一,所有性能数字均为引用,来自原始论文、Dao-AILab 官方仓库与 PyTorch 官方文档,每个数字都绑定其测量硬件;本文没有任何本地基准结果。第二,本文写作所用的工作站 GPU 是 TITAN RTX(Turing 架构,计算能力 7.5)——这不是巧合,因为官方的 FlashAttention-2/3/4 CUDA 内核无法在 Turing 上运行,最后一节把这个限制转化为一份诚实的验证方案。

文中所有”截至”表述均指 2026-08-09

注意力为什么是平方复杂度#

以单头为例,设 query、key、value 矩阵为 $Q, K, V \in \mathbb{R}^{N \times d}$($N$ 个 token,头维度 $d$),缩放点积注意力为:

$$S = \frac{Q K^{\top}}{\sqrt{d}}, \qquad P = \operatorname{softmax}(S) \text{(按行)}, \qquad O = P V,$$

其中 $O \in \mathbb{R}^{N \times d}$ 是输出。算术开销是 $O(N^2 d)$ 次 FLOPs:两次 $N \times N \times d$ 的矩阵乘法。这部分是固有的——真正让人头疼的是内存开销。得分矩阵 $S$ 和概率矩阵 $P$ 各有 $N^2$ 个元素,而标准实现会把两者都物化(materialize)在 HBM(GPU 主存)里。

一个纯算术的算例:$N = 8192$ 个 token、头维度 $d = 128$、fp16(2 字节)时,单头仅存放 $S$ 和 $P$ 就需要 $2 \times 8192^2 \times 2$ 字节 ≈ 268 MB。32 个头就是 ≈ 8.6 GB——这还没算模型权重、KV 缓存和激活值。到 $N = 64,K$(长上下文场景),单头需要 ≈ 17 GB,32 个头 ≈ 550 GB,任何现有 GPU 都装不下。这正是稀疏、低秩、哈希等近似注意力方法流行的原因:它们想砍掉 $N^2$ 成本。FlashAttention 走了另一条路——它计算相同的矩阵,但绝不让它们出现在 HBM 里。

GPU 内存层级与 IO 感知#

GPU 有两个相关的内存层级。HBM 容量大(几十 GB)但相对慢:A100 约 1.5 TB/s、H100 约 3.35 TB/s(NVIDIA 数据手册),TITAN RTX 约 672 GB/s。SRAM 是紧挨着计算单元的片上内存:每个流式多处理器(SM)只有几十到几百 KB(FlashAttention v1 的 README 给出 T4 为 64 KB;NVIDIA 数据手册列出 A100 每 block 可用共享内存最高 164 KB、H100 每 SM 为 228 KB)。SRAM 比 HBM 快一个数量级,但装不下 $N \times N$ 矩阵。

图 1. GPU 内存层级与 FlashAttention 的数据流。

+------------------------------------------------------------------+
| HBM —— 主存,几十 GB                                              |
|   Q, K, V, O    (A100 ~1.5 TB/s,H100 ~3.35 TB/s,              |
|                  TITAN RTX ~672 GB/s)                            |
|                                                                  |
|      |  按块加载(Q_i, K_j, V_j)        ^  写回 O 分块            |
|      v                                 |                         |
|   +----------------------------------------------+               |
|   | SRAM —— 片上,每 SM,约 64–228 KB                |            |
|   |   S_ij、P_ij、运行中的 m、l、O_i 都放在这里        |           |
|   |   注意力逐块在本地算完                          |            |
|   +----------------------------------------------+               |
+------------------------------------------------------------------+

FLOPs 数量与标准注意力完全相同;改变的是中间矩阵存放在哪里。FA1 论文的关键观察是:很多运算是内存受限(memory-bound)的——它们的运行时间由内存访问次数决定,而不是由算术量决定。标准注意力的一次前向正是这种情况:

图 2. HBM 流量对比:标准注意力 vs FlashAttention(前向)。

标准注意力:                               FlashAttention:
  Q, K  --HBM-->  S = QK^T / sqrt(d)         Q、K、V 按块流入;
  S    --HBM-->   逐行求最大/求和(softmax)   S_ij 与 P_ij 在 SRAM 内算完
  P    --HBM-->   O = P V                    即被消费,从不写回
  HBM 流量: Theta(Nd + N^2)                 HBM 流量:Theta(N^2 d^2 / M)
  额外内存: O(N^2)                          额外内存:O(N)

标准注意力要把 $S$ 写进 HBM、读回来做 softmax、再写 $P$、再读回来乘 $V$——四次平方级的往返。FlashAttention 的 IO 复杂度分析(2205.14135 的定理 2)证明标准注意力需要 $\Theta(Nd + N^2)$ 次 HBM 访问,而 FlashAttention 只需 $\Theta(N^2 d^2 / M)$,其中 $M$ 是 SRAM 大小。对常见的 $d = 64$–$128$、$M$ 在几十到几百 KB 的量级,$M \gg d^2$,于是 FlashAttention 的 HBM 访问次数少很多倍。论文还给出了匹配的下界:任何精确注意力算法都不可能在所有 SRAM 大小上渐近地优于这一结果。

FlashAttention-1:分块与在线 softmax#

把一切融合进单个内核的障碍是 softmax:$P$ 的第 $i$ 行依赖该行全部 $N$ 个得分的最大值总和,朴素实现必须看完整行才能输出任何结果。FlashAttention(2205.14135 的算法 1)把两个经典技巧组合起来解决这个问题——分块(tiling)与在线(online)softmax——反向传播再用重计算(下一节)。

分块。把 $Q$ 切成大小为 $B_r \times d$ 的行块 $Q_i$,把 $K, V$ 切成大小为 $B_c \times d$ 的列块 $K_j, V_j$。论文令 $B_c = \lceil M / 4d \rceil$、$B_r = \min(\lceil M / 4d \rceil, d)$,使 $Q_i, K_j, V_j, S_{ij}$(分别是 $B_r \times d$、$B_c \times d$、$B_c \times d$、$B_r \times B_c$)四块能同时放进 SRAM。内核遍历所有 $(i, j)$ 分块,计算 $S_{ij} = Q_i K_j^{\top} / \sqrt{d}$ 并汇入运行中的输出——$N \times N$ 矩阵从头到尾不存在。

图 3. 注意力计算的分块。每个分块 S_ij 在 SRAM 内算完、softmax、消费。

        K 分块(每个 B_c x d)->
   +----------------------------------+
   |  S_11  S_12  S_13  ...  S_1,nc  |     单个分块内部:
   |  S_21  S_22  S_23  ...  S_2,nc  |       K_j (B_c x d)   V_j (B_c x d)
   |  ...                            |           \             /
   |  S_nr1  S_nr2  ...   S_nr,nc    |      S_ij = Q_i K_j^T / sqrt(d)
   +----------------------------------+      P_ij = softmax(S_ij)
   Q 分块(每个 B_r x d)                     O_i += P_ij V_j

在线 softmax。标准技巧是把 softmax 写成”减去最大值”的形式,只维护运行中的最大值与归一化量。设处理完 key 分块 $1 \dots j$ 后,$m^{(j)}$ 是运行中的行最大值,$\ell^{(j)}$ 是 $\exp(s_k - m^{(j)})$ 的运行和,$O^{(j)}$ 是运行中的未归一化输出 $\sum_{k \le j} \exp(s_k - m^{(j)}) v_k$。不变量是:任何一步的”真输出”都等于 $O^{(j)} / \ell^{(j)}$,最后一步除一次即可。当带得分 $s_{j+1}$ 的新分块到达时,先抬升最大值 $m’ = \max(m^{(j)}, \max s_{j+1})$,再用 $\alpha = e^{m^{(j)} - m’}$ 缩放旧状态、用 $\beta = e^{\max s_{j+1} - m’}$ 缩放新分块:

$$\ell^{(j+1)} = \alpha, \ell^{(j)} + \beta, \ell_{j+1}, \qquad O^{(j+1)} = \alpha, O^{(j)} + \beta, \tilde{P}{j+1} V{j+1},$$

其中 $\tilde{P}{j+1} = \exp(S{j+1} - \max s_{j+1})$,$\ell_{j+1} = \sum \tilde{P}{j+1}$(按行)。为什么这样是对的:旧行的指数项此前都相对旧最大值 $m^{(j)}$ 缩放了 $e^{m^{(j)}}$ 倍;把参考最大值从 $m^{(j)}$ 换成 $m’$,恰好把每个旧指数项乘以 $\alpha$、每个新分块指数项乘以 $\beta$。输出 $O$ 是 value 向量以这些指数为权重的加权和,因此必须做完全相同的缩放。对 $j$ 做简短归纳即可确认不变量:任意前缀之后 $O^{(j)} = \sum{k \le j} \exp(s_k - m^{(j)}) v_k$,故最终的 $O / \ell$ 精确等于 $\operatorname{softmax}(S) V$。

图 4. 新 key 分块到达时的在线 softmax 状态更新。

  更新前状态:                  新分块 j+1:              更新后:
  m  = 运行中最大值             m' = max(m, max s_j+1)    m  <- m'
  l  = sum exp(s_k - m)         a = exp(m - m')           l  <- a*l + b*l_j+1
  O  = sum exp(s_k - m) v_k     b = exp(max s_j+1 - m')   O  <- a*O + b*P~ V_j+1

因果掩码几乎免费:因果注意力只需跳过 $j > i$ 的分块(或用 $-\infty$ 掩掉),工作量减半。FA1 论文还把方案扩展到块稀疏(block-sparse)注意力,论文报告它在保留分块内精确性的同时,比已知的近似注意力方法都快。

因此前向只使用线性额外内存——论文定理 1:算法 1 在输入输出之外只用 $O(N)$ 额外内存返回 $\operatorname{softmax}(QK^{\top})V$——且 HBM 访问为 $\Theta(N^2 d^2 / M)$ 而非 $\Theta(Nd + N^2)$。

反向传播:用重计算代替存储#

训练还需要梯度。朴素反向需要再次用到 $P$(或 $S$),存储它要 $O(N^2)$ 内存。FlashAttention 只保存小的逐行统计量——$O$ 与每行的 log-sum-exp $m_i$——然后在反向内核里重计算每个 $S_{ij}$ 分块。设输出梯度为 $dO$,且 $\tilde{P}{ij} = \exp(S{ij} - m_i)$:

$$dV_j = \tilde{P}{ij}^{\top} dO_i, \qquad dS{ij} = \tilde{P}{ij} \odot \left( dO_i V_j^{\top} \right), \qquad dQ_i \mathrel{+}= dS{ij} K_j, \qquad dK_j \mathrel{+}= dS_{ij}^{\top} Q_i.$$

图 5. 反向传播:重计算,而不是存储。

  前向只存:  O、m(逐行 log-sum-exp)        —— 不存 N x N 矩阵
  反向:      重新加载 Q、K、V 分块
              S_ij = Q_i K_j^T / sqrt(d)          (重计算)
              P_ij = exp(S_ij - m_i)              (重计算)
              dV_j  = P_ij^T dO_i
              dS_ij = P_ij o (dO_i V_j^T)
              dQ_i += dS_ij K_j,  dK_j += dS_ij^T Q_i

这是用少量额外算术换取大幅减少的 HBM 访问;论文明确指出重计算版反向更快而非更慢,正是因为 HBM 流量主导了运行时间。官方仓库的 flash_attn_func 文档给出了确定性结论:前向永远是确定性的;默认反向则跨线程块用原子加(atomic add)累加 $dQ$/$dK$,逐次运行在比特级不确定。仓库提供 deterministic=True 开关,文档注明”略慢且占用更多内存”。FlashAttention-4 的论文后来也把”减少反向传播中的原子加”列为明确的设计目标之一。

FlashAttention-2:更好的工作划分#

FA2 论文(2307.08691 )从一个令人清醒的测量出发:FA1 只达到 GPU 理论 FLOPs/s 的 25–40%——远低于优化 GEMM 的水平。诊断结果是:线程块与线程束(warp)之间的工作划分欠佳,造成占用率低和不必要的共享内存流量。三个修复直接对应:

  1. 减少非 matmul 的 FLOPs:用更廉价的数学缩放运行中的统计量(例如能除就不重新求幂),让更多周期花在张量核的 matmul 上。
  2. 即使单个头也按序列维度并行:FA1 中一个头的输出只能由有限的线程块计算;FA2 让网格覆盖 $(batch, heads, \text{seqlen}_q \text{ 分块})$,长序列因此获得更多线程块(更高占用率)——而且前向不需要原子操作,因为每个输出分块只属于一个线程块。
  3. 更合理的块内线程束划分:不再按输出行切分线程束(那会强迫共享内存往返),而是把 $K, V$ 的列分给不同线程束,使各线程束的部分结果可以用更少通信合并。
图 6. FA2 前向:覆盖 (batch, heads, seqlen_q 分块) 的线程块网格。

   网格:batch x heads x (N / B_r) 个线程块
        seqlen_q 分块 ->
      +----+----+----+----+----+
      | C  | C  | C  | C  | C  |     每个 C 独占一个 B_r x d 输出分块,
      | C  | C  | C  | C  | C  |     流式地把 K/V 分块送进 SRAM,
      +----+----+----+----+----+     前向无需原子操作

论文摘要报告”相比 FlashAttention 约 2 倍加速,在 A100 上达到理论峰值 FLOPs/s 的 50–73%“,端到端训练 GPT 风格模型达到每张 A100 225 TFLOPs/s(72% 模型 FLOPs 利用率)。官方 README 的更新日志显示这条发布线随后长成了完整的推理工具箱:变长(varlen,非填充)入口、滑窗注意力、ALiBi、确定性反向、分页 KV 缓存、softcapping 与 torch.compile 兼容。硬件方面 README 说得很直白:FA2 CUDA 内核支持 Ampere、Ada 与 Hopper(sm80+);对 Turing 则指向一个独立的社区仓库 ssiu/flash-attention-turing ,”它在 Turing 上支持 FlashAttention 核心功能子集”。

FlashAttention-3:Hopper 异步与 FP8#

FA3(2407.08608 )面向 H100,在 H100 上 FA2 只有约 35% 的利用率。论文把原因归结为 FA2 没有利用 Hopper 的新硬件能力,并提出三项技术:

  • warp 专用化 + TMA:让一部分 warp 专职通过张量内存加速器(TMA,异步批量拷贝,HBM → 共享内存)搬运数据,另一部分 warp 专职张量核计算,使内存传输与计算重叠而非交替。
  • matmul 与 softmax 交错:把一次迭代中的两次 GEMM 与前一次的 softmax 软件流水化,张量核不必等待指数/softmax 单元。
  • 带块量化与非相干处理的 FP8:按块量化得分,量化前先施加随机(非相干)旋转,论文报告这显著降低了 FP8 数值误差;FA3 的 FP8 经验证比基线 FP8 注意力的误差低 2.6 倍。
图 7. FA3 在 Hopper 上的 warp 专用化流水线。

  生产者 warp:  TMA:HBM -> SRAM (在……的同时取 K_{j+1}, V_{j+1})
  消费者 warp:  MMA:S_ij = Q_i K_j^T   |   softmax(S_ij)
                MMA:O_i += P_ij V_j    |   (交错、流水化)
  时间 ->        重叠:算第 j 块的同时取第 j+1 块

报告数字(H100,出自摘要):相比 FA2 加速 1.5–2.0 倍,FP16 最高 740 TFLOPs/s(75% 利用率),FP8 接近 1.2 PFLOPs/s。截至 2026-08-09,官方仓库仍把 FA3 标注为测试版(beta):要求 H100/H800、CUDA ≥ 12.3(推荐 12.8),目前只发布 FP16/BF16 前向+反向、FP8 仅前向。

FlashAttention-4:Blackwell 软硬协同设计#

FA4(2603.05451 ,2026 年 3 月)是对硬件变化的回应:在 Blackwell(B200/GB200)上,张量核吞吐大约翻倍,而共享内存带宽与指数单元几乎没有增长——FA3 为 Hopper 调优的流水线因此把新瓶颈晾在一边。FA4 的技术直接映射这个不对称性:

  • 重新设计流水线:完全异步的 MMA 操作与更大的分块,让翻倍的张量核保持饱和,不必等待共享内存。
  • 软件模拟指数与条件式 softmax 重缩放,把工作从几乎没增长的专用函数单元上挪走。
  • 张量内存(tensor memory)与 2-CTA MMA 模式,削减共享内存流量,并——明确地——削减反向传播中的原子加。

论文报告在 B200 上、BF16 下最高比 cuDNN 9.13 快 1.3 倍、比 Triton 基线快 2.7 倍,达到 1613 TFLOPs/s(71% 利用率)。一个值得注意的工程转变:FA4 全部用 CuTeDSL(一种内嵌于 Python 的 DSL)实现,论文报告编译时间比 C++ 模板内核快 20–30 倍;官方 README 的安装方式是 pip install flash-attn-4,面向 Hopper 与 Blackwell(可选 cu13 extra 适配 CUDA 13)。

表 1 汇总时间线。所有数字均出自对应论文的引用,本文未复现。

表 1. FlashAttention 时间线(所有数字均为引用,未复现)。

世代 论文 / 发布 目标硬件 核心思想 报告性能
FA1 2205.14135 ,NeurIPS 2022 Ampere;v1 发布版也支持 Turing 分块、在线 softmax、反向重计算、块稀疏扩展 A100 上比标准注意力快 2–4 倍(序列 128–4K);BERT-large 端到端比 MLPerf 1.1 纪录快 15%;每张 A100 189 TFLOPs/s(60.6% MFU)
FA2 2307.08691 ,ICLR 2024 Ampere、Ada、Hopper(sm80+) 工作划分、更少非 matmul FLOPs、序列维度并行 约 2 倍于 FA1;A100 峰值 FLOPs/s 的 50–73%;每张 A100 225 TFLOPs/s(72% MFU)
FA3 2407.08608 ,2024(beta) Hopper H100/H800 TMA、warp 专用化、matmul/softmax 交错、FP8 H100 上 1.5–2.0 倍于 FA2;FP16 740 TFLOPs/s(75%);FP8 约 1.2 PFLOPs/s
FA4 2603.05451 ,2026 Hopper + Blackwell(B200/GB200) CuTeDSL、异步 MMA 流水线、模拟 exp/重缩放、张量内存、2-CTA MMA B200 BF16 下 1.3 倍于 cuDNN 9.13、2.7 倍于 Triton;1613 TFLOPs/s(71%)

独立 flash-attn 与 PyTorch SDPA 的区别#

两种接触 FlashAttention 风格内核的途径完全不同,却经常被混为一谈:

  • 独立包Dao-AILab/flash-attention )是论文背后的内核库。它提供 flash_attn_funcflash_attn_qkvpacked_funcflash_attn_varlen_funcflash_attn_with_kvcache 等接口,README 记载了很长的功能清单:dropout、因果与滑窗掩码、ALiBi、MQA/GQA、分页 KV 缓存、融合旋转位置编码,以及 deterministic 反向开关。安装方式是编译 CUDA 内核(pip install flash-attn --no-build-isolation;需要 CUDA/ROCm 工具链、PyTorch ≥ 2.2、Linux);FA3 与 FA4 以独立包发布(hopper/ 子目录下的 flash-attn-3flash-attn-4)。
  • PyTorch SDPAtorch.nn.functional.scaled_dot_product_attention)是框架级 API,带自动分派器。PyTorch 2.13 文档 列出三个后端——FlashAttention-2 风格内核、内存高效(memory-efficient,xFormers 风格)内核、C++ math 实现——较新版本还有 cuDNN 后端。后端自动选择,也可以用 torch.nn.attention.sdpa_kernel()torch.backends.cuda.enable_flash_sdp() / enable_mem_efficient_sdp() / enable_math_sdp() 强制或禁用。flash 后端是 PyTorch 自己的实现(源码甚至注释着”FlashAttentionV2 requires that head dimension be a multiple of 8”),不是 Dao-AILab 包,且有自己的一套限制:不接受显式 attn_mask(只有 is_causal 标志)、因果 + 非方形序列长度会被拒绝、头维度 ≤ 256。

实践上的区别:SDPA 一次调用就会为你的输入、在任何 GPU 上选一个安全的内核;独立包则给你完整的内核功能集。两者不能当作可互换的 drop-in——文档警告”选择不同后端内核时,本函数的输出可能不同”,而且 math 后端对 fp16/bf16 输入会以 fp32 保存中间量,这就是 SDPA 结果在不同后端、不同 GPU 之间会略有差异的原因。

硬件、数据类型、确定性与填充注意事项#

表 2 汇总了”什么支持什么”,依据是官方 README 与 PyTorch v2.13 的分派代码(sdp_utils.cpp ,其中把 flash 后端限制在 sm80–sm121,把内存高效后端限制在 sm50–sm121)。

表 2. FlashAttention 家族内核的支持矩阵(截至 2026-08-09)。

实现 最低 GPU(CC) fp16 bf16 fp8 头维度 备注
FA1 v1(官方) Turing sm75+ 支持 支持,sm80+ 不支持 ≤ 128,8 的倍数;反向 > 64 需 A100/H100 varlen API;128/256 token 分块
FA2(官方 CUDA) Ampere sm80+ 支持 支持 不支持 最高 256 Turing 用户被指向社区 Turing 仓库
FA3(官方,beta) Hopper H100/H800 支持 支持 仅前向 视内核而定 CUDA ≥ 12.3,推荐 12.8
FA4(官方) Hopper + Blackwell 支持(已基准) 视内核而定 pip install flash-attn-4;CuTeDSL
PyTorch SDPA flash sm80–sm121 支持 支持 sm90+,若启用 FA3 ≤ 256 attn_mask;拒绝非方形因果
PyTorch SDPA 内存高效 sm50–sm121 支持 sm80+ 不支持 对齐 8(fp16)/ 4(fp32) sm80 以下仅限 fp16/fp32
PyTorch SDPA math 任意(含 CPU) 支持 支持 不支持 任意 fp32 累积;支持 fp64
ssiu/flash-attention-turing sm75 支持 见仓库 不支持 核心子集 官方 README 点名的社区移植

实践中会咬人的注意事项:

  • 数据类型。融合内核是 fp16/bf16 的天下。bf16 需要 sm80+(FA1 README 与 PyTorch 的数据类型门槛一致);sm80 以下,PyTorch 的融合内核只接受 fp16。fp32 只能走 math 或内存高效后端。FP8 是 FA3 时代的领域(e4m3),且官方 beta 只有前向;FA4 的头条数字是 BF16。
  • 确定性。独立包 FA 前向永远确定;默认反向在比特级不确定(原子累加),deterministic=True 是显式开关。PyTorch 文档提到输出随后端而异,cuDNN 路径可能选择非确定性算法。需要可复现运行时,就强制 math 后端或使用确定性开关——代价是在你的负载上实测出来的性能损失。
  • 填充与形状。融合内核要求头维度是 8 的倍数(FA1 直接断言;PyTorch 在 composite 层把头维度补齐到 8 的倍数)。FA1 把头维度上限设为 128(非 A100/H100 上反向为 64);FA2 与 PyTorch flash 后端允许到 256。PyTorch flash 后端拒绝因果 + seqlen_q != seqlen_k 的组合。独立包通过 flash_attn_varlen_func(以及 FA1 时代的 cu_seqlens API)支持变长序列、无需填充,FA4 的头条特性就是明确的非填充(un-padded)注意力;cuDNN SDPA 后端历史上要求 seq_kv 是 64 的倍数(cuDNN 8.9.6 之前)、头维度是 8 的倍数——于是某个后端能跑的填充形状,换一个后端可能悄悄回退。

本机(Turing)能跑什么,以及我们提议的诚实验证方案#

本文工作站的 GPU 是 TITAN RTX(TU102,计算能力 7.5,24 GB HBM2,按 NVIDIA 规格约 672 GB/s)。对照表 2 折叠到本机:

  • 能跑:官方 FA1 v1 内核(fp16;头维度 ≤ 128,反向 ≤ 64)——FA1 是唯一在官方 CUDA 内核清单里列出 Turing 支持的世代;PyTorch SDPA 的内存高效(fp16/fp32)与 math 后端;社区 ssiu/flash-attention-turing 移植版。
  • 不能跑:官方 FA2、FA3、FA4 CUDA 内核——FA2 需要 Ampere 或更新(sm80+),FA3 需要 H100/H800,FA4 需要 Hopper/Blackwell。PyTorch 的 SDPA flash 后端同样被限制在 sm80–sm121,在本 GPU 上不会被选中;bf16 融合内核也不可用(bf16 需要 sm80+)。

与其报告任何我们没有实际跑过的东西,不如给出四条面向 Turing 的提议验证方案,以及依据所引来源它们各自应当产生的结果。本文写作时没有执行其中任何一条;请把它们当作具体计划,而不是结果。

  1. 后端探测(PyTorch)。在本 GPU 上用 sdpa_kernel([SDPBackend.FLASH_ATTENTION]) 调用 F.scaled_dot_product_attention。预期:分派器拒绝,并给出 v2.13 源码中的警告原文——”Flash attention only supports gpu architectures in the range [sm80, sm121]“——然后改用 [SDPBackend.EFFICIENT_ATTENTION, SDPBackend.MATH] 重跑,验证输出与 math 后端在 fp16 容差内一致。
  2. 算法参照(eager PyTorch)。用纯 torch 算子实现算法 1(分块 + 在线 softmax + 最终重缩放,即上图 4 的流程),fp16、头维度 64/128、序列 512/2048、含与不含因果,与同一张量上的 SDPA math 后端对比;沿用官方仓库的测试准则:最大误差不超过基线实现误差的约 2 倍。这条可以在 Turing 上、不借助任何融合内核验证本文的推导。
  3. Turing 上的 Triton(固定旧版本)。改编 OpenAI 的 06-fused-attention.py 教程内核——官方 FA README 把它当作可读性最好的参考实现——搭配 Triton 2.1.x 发布版在 TITAN RTX 上运行;该版本 README 写明”NVIDIA GPUs (Compute Capability 7.0+)“。注意这个带日期的坑:当前的 Triton(main,2026 年)要求计算能力 8.0+,所以这条必须用旧版本;预期结果是 fp16 内核可运行、输出与 math 后端一致。
  4. 社区 Turing 内核。安装 ssiu/flash-attention-turing(官方 FA2 README 为 Turing 点名的仓库),先在 TITAN RTX 上验证其前向/反向与 SDPA 内存高效后端一致,然后才比较计时。

如果你真的在 Turing GPU 上跑这些,最有意思的问题正是论文 IO 分析预测的那几个:在 SRAM 只有 64 KB 级、没有异步拷贝硬件的 GPU 上,加速比还剩多少(FA1 的 T4 实测提示前向 2.5–4.5 倍是合理区间);以及同等精度下,内存高效的 cutlass 内核与分块在线 softmax 公式相比表现如何。

总结#

  • 注意力的平方级内存开销(而非 FLOPs)才是最初的瓶颈;FlashAttention 通过 IO 感知保持注意力精确:分块 $Q, K, V$,在 SRAM 里算 $S$ 和 $P$,绝不物化 $N \times N$ 矩阵(HBM 流量 $\Theta(N^2 d^2 / M)$ 对比 $\Theta(Nd + N^2)$,见 FA1 论文定理 2)。
  • 在线 softmax(运行最大值、运行和、重缩放)让融合成为可能;重计算让反向省内存,并且尽管多了 FLOPs 反而更快。
  • FA2 修好了工作划分(约 2 倍于 FA1;A100 峰值的 50–73%);FA3 用足 Hopper 的 TMA/warp 专用化/FP8(H100 上 FP16 740 TFLOPs/s、FP8 约 1.2 PFLOPs/s);FA4 针对 Blackwell 的不对称扩展重新协同设计(B200 上 BF16 1613 TFLOPs/s、1.3 倍于 cuDNN 9.13,CuTeDSL 实现)。
  • 独立 flash-attn 是内核库;PyTorch SDPA 是自动分派的 API,自带 FA2 风格后端且限制更严。两者不是 drop-in 等价物。
  • 硬件现实:官方 FA2/FA3/FA4 内核需要 Ampere/Hopper/Blackwell。在本工作站的 TITAN RTX(Turing,sm75)上,FA1 v1 与 PyTorch 的内存高效/math 后端是诚实的选择;上面四条验证方案是提议、尚未执行——本文不声称任何本地基准数字。

参考资料#

  1. Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness . arXiv:2205.14135,NeurIPS 2022。
  2. Tri Dao. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning . arXiv:2307.08691,ICLR 2024。
  3. Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao. FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision . arXiv:2407.08608,2024。
  4. Ted Zadouri, Markus Hoehnerbach, Jay Shah, Timmy Liu, Vijay Thakkar, Tri Dao. FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling . arXiv:2603.05451,2026。
  5. Dao-AILab. flash-attention:FlashAttention 与 FlashAttention-2 官方实现(含 FA3/FA4 发布) . GitHub 仓库(本文引用其 README 与 v1.0.9 tag)。
  6. PyTorch. torch.nn.functional.scaled_dot_product_attention torch.nn.attention.sdpa_kernel . PyTorch 2.13 文档。
  7. PyTorch. aten/src/ATen/native/transformers/cuda/sdp_utils.cpp(v2.13.0) . 本文引用的后端能力门槛。
  8. Shengqi Chen(ssiu). flash-attention-turing:面向 Turing GPU 的 FlashAttention . 官方 FA2 README 点名的社区仓库。
  9. Triton. triton-lang/triton README 兼容性说明 (main:NVIDIA 计算能力 8.0+;v2.1.0 tag:7.0+)与 06-fused-attention.py 教程
字体
阴影
滤镜
圆角
主题色