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 即可集成。