跳到正文

目录

Flash Attention:把注意力从 HBM 带宽瓶颈里捞出来

Flash Attention:把注意力从 HBM 带宽瓶颈里捞出来

面向大模型训练/推理工程师、CUDA 内核爱好者。前置知识:标准 Attention 公式、GPU 内存层级(SRAM/HBM)、PyTorch 基础。

读完本文能说清:标准 Attention 在 HBM 带宽上的瓶颈位置;FA 用 tiling + online softmax 把 O(N²) 内存压到 O(N) 的机制;FA1/FA2/FA3/FA4 历代各自解决的瓶颈层次;在主流框架里正确启用 FA 并排查常见报错;读 benchmark 时区分"测的是什么"和"不能推出什么";判断自己的场景是否适合用 FA,以及什么时候该换其他方案。

目录

  1. 先给判断 — FA 是什么、不是什么、历代各自解决什么瓶颈
  2. 标准 Attention 的瓶颈在哪里 — HBM 带宽而不是 FLOPs
  3. Tiling + Online Softmax — O(N²) 内存怎么压到 O(N)
  4. 项目快照 — 仓库、版本、作者
  5. 历代演进 — FA1/FA2/FA3/FA4 各自解决什么
  6. 安装与环境验证 — pip/源码/Docker,CUDA 版本对齐
  7. API 与典型调用 — flash_attn_func / qkvpacked / varlen
  8. 与主流框架集成 — HuggingFace / xFormers / Megatron-LM
  9. Benchmark 怎么读 — 测的是什么、不能推出什么
  10. 训练场景的注意点 — 自定义层、反向传播、DDP/FSDP
  11. 推理场景的注意点 — prefill/decode 阶段的差异
  12. 常见报错与排查 — ImportError / kernel image / 数值误差 / OOM
  13. 与近似注意力算法的边界 — 什么时候该用 Reformer/Linformer
  14. 采用顺序与决策建议 — 新项目与迁移路径
  15. 自测题 — 检验理解程度

先给判断

Flash Attention 不是近似注意力算法。它重排标准 Attention 的计算,让 GPU 内存层级能高效处理:在片上 SRAM 里完成 softmax 与加权求和,避免把 N×N 的中间矩阵写回 HBM 再读回来。同一组数学运算,内存复杂度从 O(N²) 降到 O(N),A100 上 2-4 倍墙钟时间加速,与标准 Attention 在数学上等价(FP16 下误差通常 < 1e-3)。

瓶颈在 HBM 带宽,不在 FLOPs。FA1/FA2/FA3/FA4 四代的演进都围绕这一点展开:tiling 必须配 online softmax 才能在分块下保持全局归一化;FA2 把序列维也纳入并行网格,warp 分工改成按 Q 切片,A100 前向利用率最高推到理论峰值的 73%;FA3 在 H100 上靠 warp-specialization 把 matmul 和 softmax 重叠起来,补上 Hopper 架构的 Tensor Core 空窗;FA4 转向 Blackwell 的非对称扩张——Tensor Core 吞吐翻倍,而共享内存带宽和指数单元(SFU,Special Function Unit)没跟上,于是用软件模拟 exp、条件性 softmax rescale、全异步 MMA 加张量内存,把非 matmul 的开销削掉。

本文覆盖 FA1/FA2/FA3/FA4 四代的原理差异、安装与 API 调用、与主流框架的集成、benchmark 解读、训练与推理场景的注意点、常见报错排查。FA4 论文 2026 年 3 月发布(arXiv:2603.05451),单独装 flash-attn-4,面向 Hopper 与 Blackwell(H100/B200),一并纳入讨论。CUTLASS 内核细节、Triton 实现版本不在范围内。

建议先读"标准 Attention 的瓶颈在哪里"和"Tiling + Online Softmax"两节建立直觉,再按需跳到安装、API、集成等实操章节。

标准 Attention 的瓶颈在哪里

先看标准 Attention 的实现:

import torch
import torch.nn.functional as F

def standard_attention(Q, K, V, scale=None):
    """
    Q, K, V: (batch, seq_len, d_k)
    """
    d_k = Q.size(-1)
    if scale is None:
        scale = d_k ** -0.5

    # Step 1: 计算注意力分数
    scores = torch.matmul(Q, K.transpose(-2, -1)) * scale
    # scores: (batch, seq_len, seq_len) — 完整 N×N 矩阵驻留 HBM

    # Step 2: Softmax
    attn_weights = F.softmax(scores, dim=-1)
    # attn_weights: 同样是 N×N,再次读写 HBM

    # Step 3: 加权求和
    outputs = torch.matmul(attn_weights, V)

    return outputs

三行代码里,scoresattn_weights 都是 (batch, seq_len, seq_len) 的张量。对 Llama-7B 训练时常见的 seq_len=4096batch=8heads=32head_dim=128 配置,单个 attn_weights 就是 8 × 32 × 4096 × 4096 × 2 bytes ≈ 8 GB 的 FP16 矩阵,要写一次、读一次,再写一次。

A100 的 HBM 带宽是 2.0 TB/s(SXM4 版本;80GB PCIe 版为 1.9 TB/s),H100 SXM5 是 3.35 TB/s。除以一次 attention 里要搬运的 N² 数据量,墙钟时间就上去了。FLOPs 反而不是瓶颈——A100 SXM4 的 Tensor Core FP16 算力是 156 TFLOPS(稠密;2:4 结构化稀疏可翻倍到 312 TFLOPS,但 attention 的矩阵碰不上这种稀疏),算 QK^T 和 PV 的 FLOPs 用不了那么多时间。

红色标记的两个 N×N 矩阵是 HBM 带宽的主要消耗者。Flash Attention 把这两个矩阵的搬运压进 SRAM。

Tiling + Online Softmax:怎么把 O(N²) 内存压到 O(N)

既然把 N×N 矩阵留在 HBM 是问题所在,能不能根本不 materialize 这个矩阵?难点在 softmax。softmax 的分母是 sum(exp(scores)),要算这个分母必须看到整行 scores。如果按块(tile)算 QK^T,每块算完就丢,怎么保证 softmax 的全局归一化?

Online softmax 的数学

标准 softmax 对一行 scores $s_1, \dots, s_N$ 的定义是:

$$ \text{softmax}(s_i) = \frac{e^{s_i}}{\sum_{j=1}^{N} e^{s_j}} $$

直接算会数值溢出,工程实现里先减去行内最大值 $m = \max_j s_j$:

$$ \text{softmax}(s_i) = \frac{e^{s_i - m}}{\sum_{j=1}^{N} e^{s_j - m}} $$

现在把 scores 拆成两块 $s^{(1)} = [s_1, s_2]$ 和 $s^{(2)} = [s_3, s_4]$(以 N=4 为例),分两次处理。第一块算完后得到局部最大值 $m^{(1)} = \max(s_1, s_2)$ 和局部和 $l^{(1)} = e^{s_1 - m^{(1)}} + e^{s_2 - m^{(1)}}$。第二块到来时,全局最大值更新为 $m = \max(m^{(1)}, m^{(2)})$,其中 $m^{(2)} = \max(s_3, s_4)$。此时第一块的局部和需要 rescale:

$$ l = e^{m^{(1)} - m} \cdot l^{(1)} + l^{(2)} $$

输出累加器 $O$ 同步 rescale:$O \leftarrow e^{m^{(1)} - m} \cdot O^{(1)} + O^{(2)}$。所有块处理完后,$O / l$ 就是完整的 softmax 加权结果。整个过程里,完整的 N×N 矩阵从未出现,每个 block 的 $S_{ij}$ 和 $P_{ij}$ 算完即丢。

数据流伪代码

以下伪代码展示 attention 计算在 SRAM/HBM 间的流转。真实实现还要处理 warp 分工、共享内存 bank conflict、异步搬运(Ampere 上的 cp.async、Hopper 上的 TMA)等细节,这里只看数据流。

def flash_attention_tiled(Q, K, V, block_size=64):
    """
    Flash Attention 数据流伪代码(非真实 CUDA 实现)
    关键点:N×N 矩阵永远不离开 SRAM
    """
    batch_size, seq_len, d_k = Q.shape

    # 输出和归一化因子都驻留 HBM,但只有 O(N) 大小
    outputs = torch.zeros_like(Q)
    l = torch.zeros((batch_size, seq_len, 1))      # running sum
    m = torch.full((batch_size, seq_len, 1), -float('inf'))  # running max

    for i in range(0, seq_len, block_size):
        Q_block = Q[:, i:i+block_size, :]          # (B, Br, d)
        m_i = m[:, i:i+block_size, :]              # 当前 block 的 running max
        l_i = l[:, i:i+block_size, :]              # 当前 block 的 running sum
        O_i = outputs[:, i:i+block_size, :]        # 当前 block 的 running output

        for j in range(0, seq_len, block_size):
            K_block = K[:, j:j+block_size, :]
            V_block = V[:, j:j+block_size, :]

            # === 以下全部在 SRAM 内完成 ===
            S_ij = torch.matmul(Q_block, K_block.transpose(-2, -1)) / (d_k ** 0.5)

            # online softmax 更新
            m_ij = torch.maximum(m_i, S_ij.amax(dim=-1, keepdim=True))
            P_ij = torch.exp(S_ij - m_ij)
            l_i = l_i * torch.exp(m_i - m_ij) + P_ij.sum(dim=-1, keepdim=True)

            # rescale 之前的输出,加上当前 block 的贡献
            O_i = O_i * torch.exp(m_i - m_ij) + torch.matmul(P_ij, V_block)
            m_i = m_ij
            # === SRAM 部分结束 ===

        # 只把最终结果写回 HBM,N×N 矩阵从未离开 SRAM
        outputs[:, i:i+block_size, :] = O_i / l_i
        m[:, i:i+block_size, :] = m_i
        l[:, i:i+block_size, :] = l_i

    return outputs

伪代码末尾把 ml 写回 HBM 不是画蛇添足:真实内核会保存每行的 log-sum-exp(softmax 的归一化统计量,O(N) 大小)。反向传播重算出 S 之后,用它一步恢复 P,不再需要归一化扫描,也不用存 N×N 矩阵。

具体数据流追踪

seq_len=8block_size=4 走一遍 i=0 的外层循环:

  1. 加载 Q_0:从 HBM 读 Q[0:4](4 行)进 SRAM,初始化 m_0=-infl_0=0O_0=0
  2. j=0:从 HBM 读 K[0:4]V[0:4] 进 SRAM。在 SRAM 内算 S_00 = Q_0 · K_0^T(4×4),更新 m_0 = max(S_00),算 P_00 = exp(S_00 - m_0)l_0 = sum(P_00)O_0 = P_00 · V_0S_00P_00 用完即丢,不写回 HBM。
  3. j=4:从 HBM 读 K[4:8]V[4:8]。在 SRAM 内算 S_01 = Q_0 · K_1^T,更新全局 max m_new = max(m_0, max(S_01)),把 l_0O_0 乘以 exp(m_0 - m_new) 做 rescale,再累加 P_01 · V_1
  4. 写回:把 O_0 / l_0m_0l_0 写回 HBM。完整 8×8 矩阵从未在 HBM 出现过,SRAM 里同时存在的只有 4×4 的 block。

红色标记的 S_ij 和 P_ij 是关键:这两个 N×N 的中间量只在 SRAM 内存在一个 block 的时间,算完就丢,从不写回 HBM。O(N²) → O(N) 内存复杂度就是从这里来的。

tiling 把 N×N 矩阵的生命周期压缩到一个 block 内,降内存,不再需要 HBM 来暂存——计算量(FLOPs)并没有减少。标准 Attention 的墙钟时间主要花在 HBM 读写 N×N 矩阵上,FA 把这部分读写消掉,加速就出来了。

项目快照

属性
仓库github.com/Dao-AILab/flash-attention
Stars24.8k(2026-09 快照)
Forks3.0k
贡献者199
最新版本主包 2.8.3.post1(pip install flash-attn);FA3 为 beta(hopper/ 目录);FA4 论文 2026-03 发布,单独 pip install flash-attn-4
许可证BSD-3-Clause
语言占比Python 71.2% / C++ 21.5% / CUDA 7.2%(FA4 的 CuTeDSL 实现按 .py 文件计入 Python)
作者Tri Dao(Stanford 博士,Together AI 首席科学家,2024 年 9 月起任 Princeton 计算机科学助理教授)

快照数字随时间变化,抓图的日期不同会有波动;这里标记 2026-09 是为了让读者知道采集口径。

Stars 反映生态接受度,和性能没有直接关系——性能要看后面 benchmark 段的测量条件。

历代演进:FA1 → FA4

四代 FA 各自瞄准不同层面的瓶颈。

版本主要瓶颈解决方式相对前代加速
FA1HBM 带宽(N×N 矩阵来回搬运)Tiling + online softmax2-4x vs 标准 Attention
FA2GPU 占用率低(只在 batch × heads 上并行,序列维由单个 block 顺序扫描)把序列维(行块)也纳入并行网格,重划 warp 分工约 2x vs FA1
FA3H100 的 Tensor Core 利用率低(~35%)warp-specialization 重叠 matmul 与 softmax,TMA 异步搬运,支持 FP81.5-2x vs FA2(FP16)
FA4Blackwell 非对称扩张:Tensor Core 吞吐翻倍,SFU 与共享内存带宽没跟上warp-specialization + 全异步 MMA 流水线、软件模拟 exp、条件性 softmax rescale(CuTeDSL)B200 BF16 前向 71% 利用率,≤1.3x vs cuDNN 9.13、2.7x vs Triton

FA1 解决 HBM 带宽后,FA2 面对的是 GPU 占用率——FA1 的并行只覆盖 batch × heads,序列维由单个 block 顺序扫描,长序列时 GPU 尾部大量闲置。FA2 把序列维(行块)也纳入并行网格,warp 分工同时换掉:FA1 把 K、V 切到 4 个 warp(split-K),warp 之间要靠共享内存对齐 rescale 结果;FA2 改成把 Q 切到 4 个 warp、K/V 全员可见,每个 warp 算完自己的 QK^T 小块直接乘同一份 V,warp 间不再需要通信。这套改法把 A100 上的前向利用率推到理论峰值的最高 73%,反向最高 63%。FA3 要处理 Hopper 下的 Tensor Core 利用率:FA2 换到 H100 后只能跑到 ~35% 的理论 FP16 峰值,FA3 通过 warp-specialization(一部分 warp 做 matmul,另一部分做 softmax,两者重叠)和异步数据搬运(TMA 指令),把 FP16 利用率推到 ~75%(740 TFLOPS/s),FP8 接近 1.2 PFLOPS/s。

FA4 面对的是 Blackwell(B200/GB200)的非对称扩容:Tensor Core 吞吐翻了一倍,但共享内存带宽、指数单元这类"配套部件"几乎原地踏步,纯粹的访存优化不再够用。它的对策分三层:计算上,把 softmax 的指数换成 FMA 单元的软件模拟 exp,不再把指数运算压给稀缺的 SFU,再配条件性 rescale——running max 只有变化足够大时才重缩放输出,砍掉低效的逐块缩放;流水线上,全异步 MMA 让一个 tile 的矩阵乘与相邻 tile 的 softmax 重叠;存储上,把累加器搬进 Blackwell 新增的张量内存,并用 2-CTA MMA 让一对 CTA 协同算一个 tile,减少共享内存流量和反向传播里的原子加。结果是 B200 上 BF16 前向冲到 1613 TFLOPS/s(71% 利用率),最多领先 cuDNN 9.13 约 1.3x、领先 Triton 约 2.7x。整个 FA4 用 CuTeDSL(Python 内嵌的 DSL)写成,编译时间比传统 C++ 模板实现快 20-30 倍。

FA3 仍然是精确算法。它的 FP8 模式因为低精度量化会引入数值误差,但与 Linformer、Performer 那类通过数学近似降低复杂度的算法属于不同类别。FA3 的 FP16/BF16 路径与标准 Attention 数学等价。FA4 同理:它的 BF16/FP16 路径仍是精确注意力,软件模拟 exp 只是换了指数实现方式,属于有限精度下的舍入差异,不改变算法的时间复杂度类别。

安装与环境验证

环境要求

要求说明
GPUNVIDIA(Ampere 及以上:A100、RTX 3090/4090、H100 等;FA1 额外支持 V100);AMD(MI200/MI250/MI300/MI355、RDNA 3/4)
CUDA / ROCmCUDA 12.0+(FA3 beta 建议 12.3+,最佳性能用 12.8+);AMD 走 ROCm 6.0+
PyTorch2.2+
Python3.9+

不支持 CPU。AMD GPU 走官方 ROCm 后端:默认 composable_kernel,可选 Triton,fp16/bf16 都覆盖,CK 后端 head_dim 最大支持 256。V100(sm_70)只能跑 FA1(对应 v1.x 老包);FA2 起要求 Ampere(sm_80)及以上;FA3 需要 Hopper(sm_90)。FA3 目前以 beta 形式发布在仓库的 hopper/ 目录,需要单独编译(cd hopper && python setup.py install),从 flash_attn_3 包导入(from flash_attn_3 import flash_attn_interface),与主包 flash_attn 是不同入口。FA4 面向 Hopper 和 Blackwell,已单独发布为 pip install flash-attn-4,用法是 from flash_attn.cute import flash_attn_func

安装方式

# 方式一:pip 安装(推荐,预编译 wheel)
pip install flash-attn

# 方式二:从源码安装(需要 CUDA toolkit,编译耗时 10-30 分钟)
git clone https://github.com/Dao-AILab/flash-attention.git
cd flash-attention
pip install .

不同 GPU 架构的安装差异主要在 wheel 来源:

# RTX 3090 / A100 (sm_80 / sm_86) — 标准 pip 即可,自动匹配预编译 wheel
pip install flash-attn --no-build-isolation

# 找不到匹配 wheel 时,到 GitHub Releases 下载预编译包再本地安装。
# 以 2.8.3 + CUDA 12 + PyTorch 2.8 + Python 3.12 为例,
# cxx11abiTRUE/FALSE 要与 torch.compiled_with_cxx11_abi() 的返回值一致
pip install ./flash_attn-2.8.3+cu12torch2.8cxx11abiTRUE-cp312-cp312-linux_x86_64.whl

# H100 / B200 — FA4 单独一个包;CUDA 13 环境建议加 cu13 extra 拿最佳性能
pip install flash-attn-4
# pip install "flash-attn-4[cu13]"

--no-build-isolation 让 pip 用当前环境里已装的 PyTorch 来编译扩展,而不是新建隔离环境去拉 PyTorch——后者经常因版本不匹配导致编译失败。flash-attn-4 与主包 flash-attn 是两套安装,可共存,使用 CuTeDSL 时从 flash_attn.cute 导入。官方没有发布预构建 Docker 镜像;想隔离本地 CUDA 环境,用 NVIDIA NGC 的 PyTorch 容器(nvcr.io/nvidia/pytorch)或 ROCm 的 rocm/pytorch 容器,进去之后按上面方式安装。

验证安装

import torch
from flash_attn import flash_attn_func

# 检查版本
import flash_attn
print(flash_attn.__version__)  # 期望: 2.8.x(主包);FA3 beta 用 flash_attn_3;FA4 用 flash_attn.cute

# 检查 CUDA 可用性
print(torch.cuda.is_available())           # True
print(torch.cuda.get_device_name(0))       # NVIDIA A100-SXM4-80GB / H100 / ...

# 跑一次最小用例,确认 kernel 能加载
Q = torch.randn(2, 64, 32, dtype=torch.float16, device='cuda')
K = torch.randn(2, 64, 32, dtype=torch.float16, device='cuda')
V = torch.randn(2, 64, 32, dtype=torch.float16, device='cuda')
out = flash_attn_func(Q, K, V)
print(out.shape)  # torch.Size([2, 64, 32])

如果 flash_attn_func 导入失败但 pip list 显示已安装,多半是 CUDA 版本和编译时的 CUDA 版本不匹配。nvidia-smi 看驱动支持的 CUDA 版本,nvcc --version 看编译器版本,两者要兼容。

API 与典型调用

基础调用:flash_attn_func

import torch
from flash_attn import flash_attn_func

# 张量形状: (batch_size, seq_len, num_heads, head_dim)
Q = torch.randn(2, 64, 8, 64, dtype=torch.float16, device='cuda')
K = torch.randn(2, 64, 8, 64, dtype=torch.float16, device='cuda')
V = torch.randn(2, 64, 8, 64, dtype=torch.float16, device='cuda')

# 前向计算
output = flash_attn_func(Q, K, V, dropout_p=0.0, causal=False)
print(output.shape)  # torch.Size([2, 64, 8, 64])

张量形状是 (batch, seq, heads, head_dim),不是 PyTorch 常见的 (batch, heads, seq, head_dim)。这是 FA 的约定,调用前别 transpose 错。

与标准 Attention 的误差对比

def standard_attention(Q, K, V):
    # Q, K, V: (batch, seq, heads, head_dim) → 转成 (batch, heads, seq, head_dim)
    Q = Q.transpose(1, 2)
    K = K.transpose(1, 2)
    V = V.transpose(1, 2)
    d_k = Q.size(-1)
    scores = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5)
    return torch.matmul(torch.softmax(scores, dim=-1), V).transpose(1, 2)

standard_out = standard_attention(Q, K, V)
flash_out = flash_attn_func(Q, K, V)

diff = (flash_out.float() - standard_out.float()).abs().max()
print(f"Max difference: {diff.item():.6f}")  # 通常 < 1e-3

误差来源是 FP16 累加顺序与舍入不同,算法本身是精确的。BF16 动态范围更大,长序列下不容易因中间值过大而溢出,但尾数位更少(7 位 vs FP16 的 10 位),逐次乘加的相对精度其实低于 FP16;两种精度在什么形状、什么数据上误差更大,以实测为准。FP32 下误差会趋近于 0,但 FA 不直接支持 FP32。

QKV 打包格式:flash_attn_qkvpacked_func

当 Q、K、V 来自同一个输入投影且 head_dim 相同时,打包成单个张量更高效——少一次 kernel launch,少一次 HBM 读写。

from flash_attn import flash_attn_qkvpacked_func

# qkv: (batch, seq, 3, heads, head_dim)
qkv = torch.randn(2, 64, 3, 8, 64, dtype=torch.float16, device='cuda')

output = flash_attn_qkvpacked_func(qkv, dropout_p=0.0, causal=False)
print(output.shape)  # torch.Size([2, 64, 8, 64])

变长序列:flash_attn_varlen_func

训练时一个 batch 里序列长度不一,常规做法是 pad 到最长再算,padding 部分浪费算力。flash_attn_varlen_funccu_seqlens(累积长度)把多个变长序列拼成一个长张量,跳过 padding。

from flash_attn import flash_attn_varlen_func

# 假设 batch 内有 2 个序列,长度分别为 3 和 5,拼成 8 长的张量
# cu_seqlens 是累积长度的首尾哨兵,类似 CSR 格式的行指针
cu_seqlens_q = torch.tensor([0, 3, 8], dtype=torch.int32, device='cuda')
cu_seqlens_k = torch.tensor([0, 3, 8], dtype=torch.int32, device='cuda')

Q = torch.randn(8, 8, 64, dtype=torch.float16, device='cuda')  # (total_seq, heads, head_dim)
K = torch.randn(8, 8, 64, dtype=torch.float16, device='cuda')
V = torch.randn(8, 8, 64, dtype=torch.float16, device='cuda')

output = flash_attn_varlen_func(
    Q, K, V,
    cu_seqlens_q=cu_seqlens_q,
    cu_seqlens_k=cu_seqlens_k,
    max_seqlen_q=5,
    max_seqlen_k=5,
    dropout_p=0.0,
    causal=False,
)
print(output.shape)  # torch.Size([8, 8, 64])

cu_seqlens 的语义:[0, 3, 8] 表示第 0 个序列占索引 0-2(长度 3),第 1 个序列占索引 3-7(长度 5)。max_seqlen_q 是 batch 内最长序列长度,kernel 启动时用它确定 grid 划分。

与主流框架集成

Hugging Face Transformers

Transformers 内置了 Flash Attention 后端,通过 attn_implementation 参数切换:

from transformers import AutoModelForCausalLM, AutoTokenizer
import torch

model_name = "meta-llama/Llama-2-7b-hf"

tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(
    model_name,
    torch_dtype=torch.float16,
    device_map="auto",
    attn_implementation="flash_attention_2",  # 显式指定 FA2 后端
)

inputs = tokenizer("Hello, world!", return_tensors="pt").to("cuda")
outputs = model.generate(**inputs, max_new_tokens=100)

attn_implementation 的可选值:eager(标准 PyTorch 实现)、sdpa(PyTorch 2.0 的 scaled_dot_product_attention,内部可能调 FA)、flash_attention_2(显式 FA2)。如果 flash-attn 包没装,指定 flash_attention_2 会直接报错;如果只指定 sdpa,PyTorch 会根据硬件自动选择后端。

xFormers

xFormers 的 memory_efficient_attention 是另一条路径,内部会根据硬件和输入形状选择 FA 或自家的 cutlass 内核:

from xformers.ops import memory_efficient_attention

# xFormers 期望 (batch, heads, seq, head_dim) 形状
Q = Q.transpose(1, 2)
K = K.transpose(1, 2)
V = V.transpose(1, 2)

output = memory_efficient_attention(Q, K, V, attn_bias=None, p=0.0)

xFormers 和 flash-attn 是两套独立内核,互不依赖,也不会自动互相切换。Transformers 侧走不走 FA 始终看 attn_implementation:不指定时默认 sdpa(模型支持且 PyTorch ≥ 2.1.1),其次 eager,永远不会因为装了某个包就自动换成 FA2。

Megatron-LM

新版 Megatron Core 通过 TransformerConfig.attention_backend 选择 attention 后端,默认留空(None),交给 Transformer Engine 自动决策;要钉死 FA 时才显式指定,另有 flash_attention_version 控制内核代数。老版本走的是 use_flash_attn 布尔开关。

# Megatron Core 的 TransformerConfig
# attention_backend = None   # 默认:留给 Transformer Engine 自动选后端
# flash_attention_version = 2

字段名和取值随版本变化,以所用版本的 megatron/core/transformer/transformer_config.py 为准。

Benchmark 怎么读

FA 的加速比是相对值,随 GPU、序列长度、batch、head 数和软件版本变化。论文给的是量级和趋势,精确数字必须在自己的环境里实测——本节末尾附可运行的测量脚本。

速度对比(量级与趋势)

对比加速范围说明
FA1 vs 标准 Attention(A100, FP16)2-4xFA1 论文结论;序列越长越接近上限
FA2 vs FA1(A100)约 2x(实测 1.7-3.0x)FA2 论文摘要结论
FA3 vs FA2(H100, FP16)1.5-2xFA3 论文结论;依赖 Hopper 专属指令

趋势上,序列越长、batch 越大,加速比越高——N×N 矩阵的 HBM 读写占 attention 总耗时的比例随 N 增大而增大。别把某个形状下的单点数字当成横跨所有配置的常数。

内存对比(batch=8, heads=32, head_dim=128, FP16,只算 attention 阶段峰值)

标准实现会把 scoresattn_weights 两个 N×N 矩阵完整驻留内存,峰值约 2 × batch × heads × N² × 2 字节;FA 的峰值内存只包含 O(N) 大小的输出、running max/sum 以及 Q/K/V 本身。

序列长度标准 Attention(两个 N×N 矩阵)Flash Attention(O(N) 激活)差距
2048~4.3 GB~0.5 GB约 8x
4096~17 GB~1.1 GB约 16x
8192~69 GB~2.2 GB约 32x

内存差距随序列长度近似线性拉大,这是 O(N²) 与 O(N) 的直接后果。注意这不算 KV cache——KV cache 同样随序列长度线性增长,超过 32k 后它才是主瓶颈。

这些数字测的是什么

  • 测的是:attention 内核本身的耗时或吞吐(FA1/FA2 论文的加速比含前向与反向;FA3 报告的是 H100 上前向峰值利用率),不包含 FFN、QKV 投影、优化器通信。
  • 反映的是:HBM 带宽利用率和 Tensor Core 占用率的综合表现。FA 的加速主要来自减少 HBM 读写,所以序列越长(N² 增长越快),加速比越明显。
  • 不能推出
    • 端到端训练速度提升。训练里 attention 只占总时间的一部分(通常 20-40%),FFN 和优化器通信也占大头。FA 把 attention 加速 4x,端到端可能只快 1.2-1.5x。
    • 推理场景的加速比。推理时 seq_len 短、batch 小,FA 的优势不明显,甚至可能因 kernel launch 开销变慢。
    • 跨 GPU 架构外推。H100 上 FA3 的高加速依赖 Hopper 专属指令(TMA、warp-specialization),A100 上跑 FA3 拿不到这个数。

自己测一次

import torch
import time
from flash_attn import flash_attn_func

def benchmark_attention(seq_len, batch_size=4, heads=16, head_dim=64, repeats=100):
    Q = torch.randn(batch_size, seq_len, heads, head_dim,
                    dtype=torch.float16, device='cuda')
    K = torch.randn(batch_size, seq_len, heads, head_dim,
                    dtype=torch.float16, device='cuda')
    V = torch.randn(batch_size, seq_len, heads, head_dim,
                    dtype=torch.float16, device='cuda')

    # Warmup — 必须做,第一次调用包含 JIT 编译和缓存加载
    for _ in range(10):
        _ = flash_attn_func(Q, K, V)
    torch.cuda.synchronize()

    # 测量
    start = time.time()
    for _ in range(repeats):
        _ = flash_attn_func(Q, K, V)
    torch.cuda.synchronize()
    elapsed = (time.time() - start) / repeats * 1000  # ms
    return elapsed

for seq_len in [512, 1024, 2048, 4096, 8192]:
    ms = benchmark_attention(seq_len)
    print(f"Seq len {seq_len:>5}: {ms:.2f} ms")

Warmup 这一步不能省。FA 第一次调用时会根据输入形状和硬件选择 kernel 配置,这部分时间不算在实际性能里。torch.cuda.synchronize() 也不能省,否则测的是 launch 时间而不是执行时间。

训练场景的注意点

自定义 Attention 层

在自己的模型里用 FA,需把 Q/K/V 投影后的张量形状整理成 FA 期望的 (batch, seq, heads, head_dim)

import torch
import torch.nn as nn
from flash_attn import flash_attn_func

class FlashAttentionLayer(nn.Module):
    def __init__(self, d_model, num_heads):
        super().__init__()
        self.d_model = d_model
        self.num_heads = num_heads
        self.head_dim = d_model // num_heads

        self.W_q = nn.Linear(d_model, d_model)
        self.W_k = nn.Linear(d_model, d_model)
        self.W_v = nn.Linear(d_model, d_model)
        self.W_o = nn.Linear(d_model, d_model)

    def forward(self, x, causal=False):
        batch_size, seq_len, _ = x.shape

        # 投影后直接 reshape 成 (batch, seq, heads, head_dim)
        # 不需要 transpose 到 (batch, heads, seq, head_dim)
        Q = self.W_q(x).view(batch_size, seq_len, self.num_heads, self.head_dim)
        K = self.W_k(x).view(batch_size, seq_len, self.num_heads, self.head_dim)
        V = self.W_v(x).view(batch_size, seq_len, self.num_heads, self.head_dim)

        attn_output = flash_attn_func(Q, K, V, dropout_p=0.0, causal=causal)

        # 恢复成 (batch, seq, d_model)
        attn_output = attn_output.view(batch_size, seq_len, self.d_model)
        return self.W_o(attn_output)

反向传播

FA 的反向传播也是 IO-aware 的,会重新计算前向的中间量(recomputation),而不是存 checkpoint。反向传播的 FLOPs 大约是前向的 2 倍,但不会为保存 N×N 中间量增加 HBM 读写——FA 在训练里也能加速,靠的就是这一点。代价是反向时多算一次 QK^T 和 softmax,但 Tensor Core 算这些很快,省下的 HBM 带宽远比多算的 FLOPs 值钱。

DDP 与 FSDP

FA 与 DDP/FSDP 完全兼容,它只是一个 attention kernel,不涉及梯度通信。分布式训练的注意点和标准 Attention 一样:梯度同步在 backward 之后自动触发,不需要为 FA 做特殊处理。

import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP

model = FlashAttentionLayer(d_model=4096, num_heads=32).cuda()
model = DDP(model)

for batch in dataloader:
    optimizer.zero_grad()
    output = model(batch)
    loss = loss_fn(output, target)
    loss.backward()  # FA 反向自动触发,DDP 梯度同步也自动触发
    optimizer.step()

推理场景的注意点

推理和训练的瓶颈不同。训练时 seq_len 长、batch 大,attention 占比高,FA 收益明显。推理时(特别是单条请求的生成阶段)seq_len 短、batch=1,attention 占比低,kernel launch 开销可能比省下的 HBM 带宽还大。

推理中 FA 的实际使用场景:

  • Prefill 阶段:处理长 prompt 时,attention 是 N×N 的密集计算,FA 收益和训练一样明显。
  • Batch 推理:多个请求拼 batch,seq_lenbatch 都不小,FA 有收益。
  • 单条请求的 decode 阶段:每步只算一个 token 对所有历史 token 的 attention,N 很小,FA 可能比标准 attention 还慢。这种场景用 PagedAttention 或其他 KV-cache 优化更合适。

vLLM、SGLang 等推理框架已内置针对 prefill/decode 分别调优的 attention 路径,选对框架即可,无需手动指定。

常见报错与排查

以下报错按出现频率从高到低排列。

ImportError: cannot import name 'flash_attn_func'

pip list | grep flash 看包是否真的装了。如果装了但导入失败,多半是 CUDA 版本不匹配。python -c "import torch; print(torch.version.cuda)" 看 PyTorch 编译时的 CUDA 版本,nvcc --version 看系统 CUDA 版本,两者要兼容(同一主版本号)。

RuntimeError: CUDA error: no kernel image is available for execution on the device

GPU 架构和编译目标不匹配。比如在 H100(sm_90)上跑了为 sm_80 编译的 wheel。解决:到 GitHub Releases 下载与 GPU 架构、CUDA 版本、PyTorch 版本、Python 版本都匹配的 wheel 重装。

RuntimeError: qkv must be half precision or bfloat16

FA 只支持 FP16 和 BF16。如果输入是 FP32,先转换:

Q = Q.to(torch.bfloat16)  # 或 torch.float16

BF16 通常比 FP16 更稳,因为动态范围大,不容易溢出。训练时优先选 BF16。

数值误差过大

如果 FA 输出和标准 Attention 的误差远超 1e-3,检查:

  1. 输入是否含 NaN 或 Inf——FA 的 online softmax 对异常值更敏感。
  2. 输入是否混用了 FP16 和 BF16——FA 要求 Q/K/V 同精度,混用要么直接报错,要么因隐式转换引入额外误差。
  3. causal 参数是否一致——标准 Attention 实现里手动加 mask 容易出错。

OOM 在长序列时

FA 把 attention 的内存从 O(N²) 降到 O(N),但整个模型还有 FFN、KV cache、激活值。如果还是 OOM,检查:

  • KV cache 是否用了 FA 的 varlen 接口拼 batch。
  • 梯度检查点(gradient_checkpointing)是否开启——这能把激活值内存也压下来。
  • 优化器状态是否分片(FSDP 或 ZeRO-2/3)。

常见误区

  • “FA 能加速所有 attention 计算”:短序列(seq_len < 512)、batch=1 的推理场景下,kernel launch 开销可能比省下的 HBM 带宽还大。用 PyTorch 原生 scaled_dot_product_attention 更合适。
  • “FA3 是近似算法”:FA3 的 FP16/BF16 路径与标准 Attention 数学等价。FP8 模式有量化误差,但这是低精度计算的代价,不是算法近似。
  • “装了 flash-attn 就一定走 FA”:Transformers 不指定 attn_implementation 时默认是 sdpa(模型支持且 PyTorch ≥ 2.1.1),否则 eager,都不会自动切到 FA;要用 FA2 必须显式指定。sdpa 后端会根据硬件在 flash、mem-efficient 等内核间自动选择,但不等于 flash-attn 包。
  • “FA 输出和标准 Attention 完全一致”:FP16 下误差通常 < 1e-3,来自累加顺序不同。对数值精度敏感的业务(如金融、科学计算),需评估是否可接受。
  • “FA 能解决所有长上下文问题”:FA 把 attention 内存从 O(N²) 降到 O(N),但 KV cache 仍是 O(N)。序列长度超过 32k 后 KV cache 内存会成为新瓶颈,需配合 Ring Attention、PagedAttention 等方案。

与近似注意力算法的边界

FA 是精确算法,但有些场景下近似算法更合适。下表列出各自的适用场景:

算法精确度时间/内存复杂度适用场景
Flash Attention精确O(N²) 时间、O(N) 内存序列长度 < 32k,GPU 内存够装 KV
Reformer近似(LSH)O(N log N) / O(N)极长序列(>32k),可接受精度损失;可逆层还能省激活内存
Linformer近似(低秩投影)O(N) / O(N)序列长度固定,离线训练
Performer近似(随机特征)O(N) / O(N)想要线性复杂度的无偏 softmax 近似,对精度损失不敏感
Longformer / BigBird近似(稀疏模式)O(N) / O(N)文档级任务,有明确的局部+全局模式

FA 出现后,近似算法在生产环境的使用明显减少。在大多数实际序列长度(< 32k)下,FA 又精确又快,近似算法省下的计算量往往被精度调优成本抵消。超过 32k 的超长序列,FA 在 H100 上仍能跑到 128k+,但 KV cache 内存会成为新瓶颈,这时候 Ring Attention、PagedAttention 这类方案更合适。

采用顺序与决策建议

新项目:

  1. 训练:Ampere/Ada 直接用 FA2(attn_implementation="flash_attention_2")。H100 上可试 FA3(beta,从仓库 hopper/ 目录单独编译,导入入口 flash_attn_3,要求 CUDA 12.3+,建议 12.8+);面向 H100/B200 的新训练,直接评估 FA4(pip install flash-attn-4,B200 上 BF16 利用率可到 71%)。
  2. 推理:用 vLLM 或 SGLang,它们内部已根据 prefill/decode 阶段选了最优 attention 实现。
  3. 长上下文(>32k):先确认 KV cache 内存是否够,再考虑 Ring Attention 或序列并行。
  4. 非 NVIDIA GPU:AMD 用官方 ROCm 后端(composable_kernel 或 Triton,ROCm 6.0+);其余厂商没有官方支持,走 PyTorch SDPA 等通用路径。

已有项目迁移:

  1. 先在测试集上对比 FA 输出和原 attention 的误差,确认 < 1e-3。
  2. 小 batch 跑通训练循环,确认 loss 曲线一致。
  3. 再上大 batch 长序列,观察实际加速比——attention 部分预期 2-4x,端到端 1.2-1.5x。
  4. 如果加速比远低于预期,profile 看 FFN 或通信是否成了新瓶颈。

自测题

答案都在对应章节里,不另给标准答案。

原理层

  1. 标准 Attention 的墙钟时间主要花在 FLOPs 还是 HBM 读写上?为什么 A100 的 156 TFLOPS(稠密 FP16)算力用不满?
  2. Online softmax 为什么必须保留 running max $m$ 和 running sum $l$ 两个状态?只保留 $l$ 会出什么问题?
  3. Tiling 把 N×N 矩阵的生命周期压缩到一个 block 内,FLOPs 减少了吗?如果没有,加速从哪里来?
  4. FA 反向传播用 recomputation 而不是存 checkpoint,这两者的区别是什么?为什么 FA 选前者?

工程层

  1. flash_attn_func 期望的张量形状是 (batch, seq, heads, head_dim),和 PyTorch 常见的 (batch, heads, seq, head_dim) 不同。如果调用前忘了 transpose(或 transpose 错了),会报错还是静默给出错误结果?
  2. flash_attn_varlen_funccu_seqlens 语义是 CSR 风格的累积长度。给定 cu_seqlens=[0, 3, 8],batch 里有几个序列?各自长度多少?
  3. HuggingFace Transformers 里指定 attn_implementation="flash_attention_2" 但没装 flash-attn 包,会报错还是回退到 eager
  4. FA3 的 FP8 模式有量化误差,为什么仍然算"精确算法"而不是"近似注意力"?

场景判断层

  1. 单条请求的 decode 阶段(batch=1,每步只算 1 个 token 对 N 个历史 token 的 attention),FA 通常比标准 attention 慢。原因是什么?这种场景该用什么?
  2. 序列长度 64k,KV cache 内存成为新瓶颈,FA 还能用吗?需要配合什么方案?
  3. 训练时 attention 部分用 FA 加速了 4x,端到端训练速度为什么通常只快 1.2-1.5x?剩下的时间花在哪了?
  4. AMD MI300X 上能用官方 flash-attn 包吗?用什么后端,装之前要确认 ROCm 什么版本?
  5. FA4 的"条件性 rescale"相比标准 online softmax 减少了哪件事?为什么在 Blackwell 上这种做法是划算的?

想深入内核方向,可以按这个顺序:

  1. 读 FA2 论文第 3 节的 work partitioning 部分,对照 FA1 的 split-K 方案,理解为什么把 Q 切到 4 个 warp、K/V 全员共享之后,warp 间通信就消失了。
  2. 对照本文伪代码,在 csrc/flash_attn/flash_api.cppflash_fwd_kernel.h 里找到 online softmax rescale 的 CUDA 实现。
  3. 读 FA3 论文第 3 节,理解 cp.async 和 TMA 指令如何重叠数据搬运与计算。
  4. 想自己写 tiling kernel,从 Triton 的 flash_attention 教程入手,比直接读 CUTLASS 容易。
  5. 读 FA4 论文(arXiv:2603.05451),关注 2-CTA MMA 与条件性 rescale;想动手,直接从它的 CuTeDSL 实现看起,比 C++ CUTLASS 好读得多。FA4 仍带 beta 标记,生产前按自己的 B200/GB200 实测。

引用

@article{dao2022flashattention,
  title={FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness},
  author={Dao, Tri},
  journal={Advances in Neural Information Processing Systems},
  year={2022}
}

@article{dao2023flashattention2,
  title={FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning},
  author={Dao, Tri},
  journal={arXiv preprint arXiv:2307.08691},
  year={2023}
}

@article{shah2024flashattention3,
  title={FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision},
  author={Shah, Jay and Bikshandi, Ganesh and Zhang, Ying and Thakkar, Vijay and Ramani, Pradeep and Dao, Tri},
  journal={arXiv preprint arXiv:2407.08608},
  year={2024}
}

@article{zadouri2026flashattention4,
  title={FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling},
  author={Zadouri, Ted and Hoehnerbach, Markus and Shah, Jay and Liu, Timmy and Thakkar, Vijay and Dao, Tri},
  journal={arXiv preprint arXiv:2603.05451},
  year={2026}
}

相关资源

资源链接
GitHub 仓库https://github.com/Dao-AILab/flash-attention
FA1 论文https://arxiv.org/abs/2205.14135
FA2 论文https://arxiv.org/abs/2307.08691
FA3 论文https://arxiv.org/abs/2407.08608
FA4 论文https://arxiv.org/abs/2603.05451
Tri Dao 主页https://tridao.me

参与讨论

使用 GitHub 登录。欢迎补充事实、异议与实践。