flash-attention
通过 IO感知 tiling 算法,将 Transformer 注意力计算提速 2-4 倍、内存占用降低 N 倍,成为 PyTorch/HuggingFace/DeepSpeed 事实标准的底层加速库
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
通过 IO感知 tiling 算法,将 Transformer 注意力计算提速 2-4 倍、内存占用降低 N 倍,成为 PyTorch/HuggingFace/DeepSpeed 事实标准的底层加速库
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
2022年,一个让整个 AI 圈头疼的问题摆在面前:训练一个 1750 亿参数的 GPT-3 模型,一次前向传播需要计算约 490 PetaFLOPs,但传统 Attention 机制的内存占用随着序列长度的增长呈平方级爆炸——处理 2048 token 时,注意力矩阵就要占用 16GB GPU 显存,根本喂不下更大的模型。
斯坦福博士生 Tri Dao 和他在 UC Berkeley 的导师们,决定从底层重新设计这个"所有人都知道太慢、但没人想过能彻底解决"的问题。他们在 2022 年发表论文《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》,同年衍生出 FA2(2023年),2024年又发布为 Hopper GPU 优化的 FA3 测试版。
今天,这个项目已收获 24,638 颗 GitHub Stars,被 PyTorch 核心、Hugging Face Transformers、DeepSpeed、Megatron-LM 等几乎所有主流 AI 框架直接内置采用,成为当代大模型训练基础设施的"隐形冠军"。

FlashAttention-2 在 A100 80GB SXM 上的 Forward + Backward 基准测试,相比标准 Attention 提速 3-5 倍
理解 FlashAttention 的核心,要从 GPU 的"内存墙"说起。
GPU 有两种内存:高速的 HBM(High Bandwidth Memory,即显存) 和容量极小的 SRAM(Shared Memory)。传统 Attention 的问题在于:它需要反复在 HBM 和计算单元之间来回搬运数据——对于长序列,这个数据流就像一条堵满卡车的公路,计算单元却经常在"等数据到达"。
Tri Dao 的核心洞察是:与其优化 FLOPs(计算量),不如优化 IO(数据搬运)。FlashAttention 借鉴了 OS 领域的 tiling 和 SRAM 缓存思路——把注意力矩阵切成小块(tiles),每次只把一小块从 HBM 加载到 SRAM,在 SRAM 里完成计算后再写回。由于 SRAM 访问速度远快于 HBM,整体 IO 量大幅降低,内存占用也随之从 O(N²) 降到 O(N)。
这还不是近似计算——FlashAttention 输出的结果与标准 Attention 完全一致(exact attention),没有精度损失,只是更快、更省显存。

FlashAttention 通过 tiling 策略将 Attention 的内存占用从 O(N²) 降低至 O(N),这是突破性的内存效率提升
光有理论不够,实际跑起来才是真本事。
在 A100 80GB GPU 上,标准 Attention 的 FLOPs 利用率只有约 35%,意味着大量计算资源在空转。FlashAttention-2 通过三大优化——更好的并行化、革命性的工作分区(work partitioning)策略、以及更精细的 warp 级别控制——将这一数字提升到 72%,提速 2-3 倍。
这个提升有多夸张?在 MLPerf 2.0(2022年6月)基准测试中,FlashAttention 让 BERT 训练速度创下了云端最快纪录。而在 MLPerf 2.1(2022年11月),与微软 Azure 和 NVIDIA 合作,将 BERT 训练时间压缩到 16台 A100 节点、2分钟以内 完成。

FlashAttention-2 在 A100 80GB 上的 Forward + Backward 性能对比,越靠右越快
FlashAttention 不是一次性的项目,而是持续迭代的工程体系:
| 版本 | 主要优化 | 适用硬件 | 状态 |
|---|---|---|---|
| FA2 | 并行化 + warp 分区,3-4x 提速 | Ampere (A100/RTX 3090)、Ada (RTX 4090)、Hopper (H100) | 稳定版,pip install flash-attn |
| FA3 | TMA 张量内存加速器、FP8 支持 | Hopper (H100/H800),需 CUDA 12.3+ | Beta,已发布 |
| FA4 | CuTeDSL 编写,JIT 编译到 PTX/CUBIN | Hopper + Blackwell (H100/B200) | pip install flash-attn-4 |
FA4 是当前最活跃的开发分支,完全用 CuTeDSL(CUDA CUTLASS 的 DSL)编写,由 NVIDIA CUTLASS 团队合作开发。代码直接生成 PTX/CUBIN,通过 JIT 在运行时编译,针对 SM90(Hopper)和 SM100/SM110(Blackwell)架构做了专门优化。相比 FA2,FA4 的前向传播吞吐量在 H100 上再提升 1.5-2 倍。

FlashAttention-3 在 H100 80GB SXM5 上的 FP16 前向传播性能(Beta 版本)
除了核心 Attention 内核,项目还提供了完整的优化 Transformer 训练实现(training/ 目录),将 FlashAttention 集成进 GPT、ViT 等模型端到端训练流程。
benchmark 数据极具冲击力:

端到端训练效率对比:FlashAttention 集成后,GPT-2(左)和 GPT-3(右)在各种序列长度上均实现 3-5 倍训练加速
FlashAttention 最深远的影响,可能不是它自己的接口,而是它被各大框架直接内置的程度:
torch.nn.Transformer 自 PyTorch 2.0 起内置 FlashAttention这意味着,如果你用 PyTorch 2.0+ 跑 nn.MultiheadAttention,或者用 Hugging Face 的 BertModel,底层自动就在用 FlashAttention——无需任何额外代码,透明地享受 2-4 倍的注意力加速。
对普通用户来说,最简单的体验方式:
pip install flash-attn --no-build-isolation
不过这里有个门槛:pip 安装过程会在本地编译 CUDA 内核。没有 ninja 时,编译耗时可达 2 小时;有 ninja 时约 3-5 分钟(64核机器)。如果机器 RAM < 96GB,可以用 MAX_JOBS=4 限制并发编译任务数。
依赖要求:
Python 接口极其简洁:
from flash_attn import flash_attn_func
output = flash_attn_func(q, k, v, causal=True)
FA4 的接口略有不同(针对 Blackwell 优化):
from flash_attn.cute import flash_attn_func
out = flash_attn_func(q, k, v, causal=True)
冷静看待,FlashAttention 也有其局限:
FlashAttention 的意义远超一个"加速库"。它的出现证明了两件事:
第一,算法层面的 IO 优化可以和硬件无关地带来巨大收益。 通过重新设计计算流程而非近似计算,FlashAttention 在保证数学等价性的前提下实现了数量级的效率提升,为"精确注意力"这个赛道指明了方向。
第二,学术成果到工业落地可以非常快。 2022年论文发表,2022年 PyTorch 就完成了集成,2023年已经成为大模型训练的事实标准——这个速度在传统 HPC 领域是不可想象的。
从 BERT 到 GPT-3,从 LLaMA 到 ChatGLM,FlashAttention 隐藏在每一代大模型的训练流水线上,默默支撑着参数量的指数增长。今天 Meta 开源 LLaMA-3 能训练 405B 参数的模型,背后有 FlashAttention 系列不可磨灭的贡献。
| 项目 | 信息 |
|---|---|
| 仓库 | Dao-AILab/flash-attention |
| Stars | 24,638 ★ |
| 语言 | Python + CUDA C++ |
| 许可 | BSD-3-Clause |
| 主要贡献者 | Tri Dao(斯坦福),Daniel Y. Fu 等 |
| pip 包 | flash-attn(FA2)、flash-attn-4(FA4) |
| 集成框架 | PyTorch、Hugging Face、DeepSpeed、Megatron-LM 等 |
| 训练性能 | 189 TFLOPs/sec/A100,60.6% MFU |