FlashAttention
FlashAttention 是一种用于计算 Transformer 模型中注意力机制的优化算法。它比标准注意力更快、内存效率更高、可扩展性更强。它首次发表于论文 FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness,此后已成为现代 LLM 在训练和推理中广泛采用的注意力后端。
FlashAttention 是一种内核级优化。其收益来自注意力计算本身在 GPU 上的执行方式,包括数据如何在 HBM 与 SRAM 之间移动、工作如何分块(tiling),以及运算如何融合到单个内核中。这就是为什么本页归属于内核优化章节。关于服务层面的技术,请参见推理优化。
为什么注意力一开始就很慢
当 LLM 阅读文本时,它必须查看每个 token,并把它与所有其他 token 进行比较,以理解它们之间的关系。这就是所谓的注意力。
标准注意力机制有一个根本问题:它是内存受限的,而不是计算受限的。要理解这一点,我们需要看看注意力计算过程中发生了什么。
如果你想了解 HBM、SRAM、warp 和分块等概念背后的底层 GPU 背景知识,请参见GPU 架构基础。
标准注意力计算:
朴素实现遵循以下步骤:
- 计算注意力分数:将 Q 与 K^T 相乘,得到一个 N×N 矩阵(其中 N 是序列长度)
- 应用 softmax:对分数进行归一化
- 乘以值:与 V 计算加权和
问题在于内存访问模式。现代 GPU 具有:
- 高带宽内存(HBM):容量大但速度慢(1-2 TB/s 带宽)
- SRAM(片上内存):容量小但速度快(10-20 TB/s 带宽)
标准实现需要将完整的 N×N 注意力矩阵写入 HBM,然后再读回来供下一步运算使用。对于 4096 个 token 的序列长度,这个注意力矩阵包含约 1600 万个元素。由于多次读写,算法大部分时间都花在等待内存传输上,而不是做实际计算。
随着序列长度增加:
- 内存流量主导运行时间
- GPU 利用率下降
- 长上下文窗口变得不切实际
例如,16K token 需要的内存是 1K token 的 256 倍。
FlashAttention 是如何工作的?
FlashAttention 通过减少内存流量来加速注意力计算。其核心思想是绝不在 HBM 中物化完整的注意力矩阵,而是采用两种关键技术:
- 分块与重计算:FlashAttention 将计算分解为能放进快速 SRAM 的块(tile):
- 将 Q、K、V 的分块从 HBM 加载到快速 SRAM
- 完全在 SRAM 中计算该分块的注意力
- 增量更新输出,并丢弃中间结果
- 内核融合:FlashAttention 不采用分离的运算(matmul → softmax → matmul),而是将所有运算融合到单个 GPU 内核中。这意味着:
- 不会将中间结果写入 HBM
- 没有单独的内核启动(这些启动有开销)
- 所有运算都在快速 SRAM 中完成

FlashAttention 使用分块技术来避免在 HBM 上物化大型 N×N 注意力矩阵。图片来源
简单来说,FlashAttention 让注意力计算更高效。它重新组织工作,使 GPU 花更少的时间等待内存,花更多的时间做实际计算。
想更全面地了解 Triton、CUDA、编译器栈和性能分析工具在此类工作中的定位,请参见内核优化工具。
FlashAttention 的优势
FlashAttention 在速度和可扩展性方面都带来了显著改进:
- 注意力计算快 2–4 倍
- 由于不存储 N×N 注意力矩阵,内存占用大幅降低
- 允许 LLM 处理更长的上下文窗口(例如 128K token)
- 更高的吞吐量和更好的 GPU 利用率
- 对话、编程、推理等场景的推理速度更快
目前,FlashAttention 被广泛应用于:
- 训练框架(PyTorch、DeepSpeed)
- 推理引擎(vLLM、SGLang、Hugging Face TGI、TensorRT-LLM)
- 支持长上下文的模型架构
FlashAttention 版本对比
FlashAttention 主线目前已有 4 个大版本。下面是并排对比,说明该算法在各版本中的演进。
| 版本 | 年份 | 关键改进 | 性能 | 备注 |
|---|---|---|---|---|
| FlashAttention-1 | 2022 | 引入了 IO 感知的分块注意力算法。融合 softmax + matmul 内核。避免物化完整注意力矩阵 | 注意力快 2–4 倍,内存最多降低 10 倍 | 第一个版本;支持实用的长上下文;精确注意力(无近似) |
| FlashAttention-2 | 2023 | 更好的 warp 间并行与工作划分;减少了非 matmul 的 FLOPs | 比 FA-1 快 2 倍,尤其在长序列上 | 支撑了许多长上下文 LLM;已广泛集成到推理/训练框架中 |
| FlashAttention-3 | 2024 | Tensor Core 加速(FP8/BF16);针对 Hopper GPU(如 H100)优化 | 比 FA-2 最多快 2 倍,H100 上达 740 TFLOPS(75% 利用率);FP8 数值误差降低 2.6 倍 | 利用 Hopper 异步执行与 warp 专业化。许多框架仍在优先从 FA-2 升级 |
| FlashAttention-4 | 2026 | 完全异步的 MMA、更大的分块、软件模拟的指数运算、条件 softmax 重缩放、张量内存(tensor memory)与 2-CTA MMA | 在 B200 BF16 基准上比 cuDNN 9.13 最多快 1.3 倍,比 Triton 最多快 2.7 倍,最高达 1613 TFLOPS/s(71% 利用率) | 使用 CuTeDSL 编写。官方实现通过 flash-attn-4 暴露,面向 Hopper 和 Blackwell GPU,如 H100 和 B200 |
FlashAttention-4 专门针对 NVIDIA Blackwell 架构进行了调优。其核心洞见是硬件的不对称扩展。Tensor Core(负责 QKᵀ、PV 等大型矩阵乘法)在 Blackwell 上大幅提速,但其他关键资源却没有同等程度地扩展:
- 共享内存带宽。
- 用于 softmax 中指数运算的特殊函数单元(SFU)。
- 寄存器压力和调度开销。
因此,在 B200 上瓶颈发生了转移。更多信息请参阅 FlashAttention-4 论文。
如何使用 FlashAttention
最简单的入门方式是通过官方包:
pip install flash-attn --no-build-isolation
较新的 PyTorch 版本在支持的情况下,会通过 scaled_dot_product_attention 自动分派到 FlashAttention。更多信息请阅读API 参考。
许多推理框架已经集成了 FlashAttention,包括 vLLM 和 SGLang,但其版本可能因发布周期而不同。
对于 FlashAttention-4,官方仓库单独提供了一个 CuTeDSL 包:
pip install flash-attn-4
在 CUDA 13 上,仓库推荐:
pip install "flash-attn-4[cu13]"
FA4 的 API 通过 flash_attn.cute 暴露:
from flash_attn.cute import flash_attn_func
out = flash_attn_func(q, k, v, causal=True)
在生产环境中依赖它之前,请查看当前的 FlashAttention 仓库,因为 FA4 包和框架集成迭代很快。
扩展阅读
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
- FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning
- FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision
- FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling
- Official FlashAttention repository