FlashAttention-4 · 干货攻略
- 链接:https://x.com/tri_dao/status/2029569889858646344
- 分类:x-tips
- 来源:X @tri_dao
- 作者:Jay
- 更新:2026-07-18
- 仓库:Dao-AILab/flash-attention
这是什么
FlashAttention-4(FA4)是 Tri Dao 团队发布的第四代 IO-aware 精确注意力内核,针对 NVIDIA Blackwell 架构(B200/GB200)全面重新设计,同时保留对 Hopper(H100)的支持。FA4 以 CuTeDSL(NVIDIA CUTLASS DSL)编写,在运行时编译为 PTX/CUBIN,核心定位是解决 Blackwell 不对称硬件扩展带来的新型瓶颈。
从 FA1 到 FA3,每代核心优化都是减少 HBM 访问;FA4 的关键变化是算法与内核共同设计——不再只优化 GEMM,而是系统性地 overlapped 所有瓶颈资源(矩阵乘法、指数运算、共享内存带宽)。
为什么值得关注
背景:Blackwell 的不对称扩展问题
官方博客(tridao.me/blog/2026/flash4)给出了关键 feeds & speeds 数据。以 M=N=D=128 在 B200 上每 SM 为单位:
| 资源 | 每周期操作数 |
|---|---|
| BF16 张量核 | 8192 ops/cycle |
| 指数单元(MUFU.EX2) | 16 ops/cycle |
| 共享内存带宽 | 128 bytes/cycle |
H100 → B200:BF16 张量核算力从 1 PFLOP/s 提升至 2.25 PFLOP/s(2.25×),但指数单元数量和共享内存带宽完全未变。
这意味着: - FWD pass 瓶颈从 GEMM 转移到了指数运算(MUFU.EX2) - BWD pass 瓶颈从 GEMM 转移到了共享内存带宽
传统"两个 GEMM 加一个 softmax"的朴素优化思路在 Blackwell 上彻底失效。FA4 的设计正是围绕这一新现实展开。
X 上的分享点(@tri_dao)
@tri_dao 分享了 FA4 的 Blackwell 专项优化,并提到 FA4 包含完整 debug 文档——这对想深入内核级别调试的工程师是直接可用的资源。原帖还提及 Claude/Codex 参与了特定场景的调试,此说法未在官方文档中直接核验,仅供参考。
核验过程
官方来源
1. GitHub README(Dao-AILab/flash-attention)
- 安装命令:pip install flash-attn-4;CUDA 13 优化版:pip install "flash-attn-4[cu13]"
- 接口:from flash_attn.cute import flash_attn_func
- 依赖:CUDA toolkit / ROCm toolkit、PyTorch 2.2+、packaging、psutil、ninja
- FA4 使用 CuTeDSL 编写,编译为 PTX/CUBIN,运行时编译
2. 官方博客(tridao.me/blog/2026/flash4) - B200 BF16 实测:1,605 TFLOPs/s(71% 峰值算力利用率),论文版本(arXiv:2603.05451)给出 1,613 TFLOPs/s,数值一致 - FWD pass:3.6× vs FA2(seq_len=32,768);BWD pass:3.15× vs FA2 - cuDNN 对比:1.3× vs cuDNN 9.13;Triton 对比:2.7× vs Triton - FWD 瓶颈:指数运算 → 解决:MUFU.EX2 + FMA 软件模拟 exp 的流水线重叠;ping-pong Q tile 调度 - BWD 瓶颈:共享内存带宽 → 解决:中间结果存 TMEM + 2-CTA MMA 模式(原子操作减半)
3. arXiv:2603.05451(FA4 论文) - 确认 1.3× vs cuDNN 9.13、2.7× vs Triton 数字 - Blackwell B200 BF16 基准:TFLOPs/s 数字(1,613)与博客(1,605)差异在合理范围
4. GitHub CLAUDE.md(flash-attention)
- 确认 flash_bwd_sm100.py 包含 FlashAttentionBackwardSm100,支持 Blackwell 2-CTA MMA 模式和 block sparse
交叉验证
搜索结果中多个第三方来源(ascii.co.uk、blockchain.news、thecryptohodl.com 等 2026 年 1-2 月报道)一致引用 B200 3.6× FWD speedup、1,605 TFLOPs/s、71% 利用率、1.3× cuDNN、2.4× Triton 等数字,与官方博客和论文结论完全吻合,交叉验证通过。
关于 Claude/Codex 参与调试的说法:GitHub 仓库包含 CLAUDE.md(AI 辅助开发配置文件),但未在公开文档中找到具体 "2CTA forward deadlock" 调试过程的直接描述。此点原帖主张,未核验,攻略正文不引用此具体说法。
上手步骤
安装
# 基础安装
pip install flash-attn-4
# CUDA 13 优化版本(推荐)
pip install "flash-attn-4[cu13]"
依赖项(缺一不可):
pip install torch packaging psutil ninja
验证安装:
import torch
from flash_attn.cute import flash_attn_func
q = torch.randn(1, 8, 512, 64, dtype=torch.float16, device='cuda')
k = torch.randn(1, 8, 512, 64, dtype=torch.float16, device='cuda')
v = torch.randn(1, 8, 512, 64, dtype=torch.float16, device='cuda')
out = flash_attn_func(q, k, v, causal=True)
print(out.shape) # torch.Size([1, 8, 512, 64])
生产环境集成
FA4 已进入 cuDNN 9.14。使用 PyTorch JIT 模式或通过 SGLang / vLLM 调用时,张量核自动分派到 FA4,无需手动加载。
在 SGLang 中启用(server.py 启动参数):
python -m sglang.launch_server --enable-flash-attn-4 ...
从源码编译(如需 debug 或自定义)
git clone https://github.com/Dao-AILab/flash-attention.git
cd flash-attention
pip install -e .
# 限制并行编译任务(内存 < 96GB 机器)
MAX_JOBS=4 pip install --no-build-isolation .
编译耗时(有 ninja):3-5 分钟(64 核);无 ninja:可达 2 小时。
坑与适用边界
适用硬件
- FA4 主目标:Blackwell B200 / GB200(Hopper H100 亦可运行)
- FA1/FA2 仍为 Ampere/Ada/Turing 用户最优解
- Windows 兼容性:从 v2.3.2 起有社区正向反馈,但非官方支持
关键约束
- head_dim 上限 256(与其他版本一致)
- BWD pass head_dim > 192 需 A100/A800 或 H100/H800(消费级 GPU 不支持)
- head_dim 256 BWD 在无 dropout 时可在消费级 GPU 运行(flash-attn ≥ 2.5.5)
确定性执行
FA4 支持 deterministic mode,保证可复现的训练结果,适合需要 bit-exact 复现的评测场景,通过环境变量或内核参数开启。
2-CTA MMA 的 CTA pair 约束
2-CTA MMA 需 CTA pair 两端在 operation 进行期间保持 active,CTA group size(1 或 2)必须在整个 TMEM 和张量核操作内核中保持恒定。修改调度策略时需注意此约束,否则可能导致死锁。
精度注意
FWD pass 中软件模拟 exp(利用 FMA 而非 MUFU.EX2)使用 Cody-Waite 范围约简 + Horner 多项式近似(4 阶),论文声明精度损失在可接受范围内,适合训练场景。
一句话结论
FlashAttention-4 通过流水线重叠张量核/指数/内存操作和 Blackwell 专属 2-CTA MMA,在 B200 上实测 1,605 TFLOPs/s(71% 利用率),FWD 比 FA2 快 3.6 倍,在 cuDNN 和 Triton 基准上全面领先,是当前 Blackwell 训练/推理的注意力内核最优解——直接 pip install flash-attn-4 即可集成。