跳到主要内容

数据截至 (上游 commit 0251105a2fb1)

FlashAttention — IO 感知的精确注意力 kernel

30 秒导读: FlashAttention 是 Tri Dao 等人的注意力 CUDA kernel 库(FA1/FA2 论文见 README.md:6)。它算的和普通注意力一模一样(精确,不是近似),但把 N×N 注意力矩阵永远关在片上 SRAM 里、不往显存写,于是注意力从「被显存带宽卡死」变成「被算力喂饱」——长上下文因此变便宜。它是今天几乎所有训练框架和推理引擎的默认注意力后端。


1. 这是什么(零基础也能懂)

一句话定义

FlashAttention 是一个精确注意力的高性能 GPU 实现:输入 Q、K、V,输出 softmax(QKᵀ·scale)·V,数学上与教科书定义逐位等价,但显存占用从 O(N²) 降到 O(N),速度提升数倍。

解决什么问题 / 给谁用

设想你在 A100 上训练一个序列长度 8K 的 Transformer。标准注意力要先在显存里物化一张 8K×8K 的分数矩阵(每个头 128MB),写完再读回来做 softmax,再写回再读回来做第二次矩阵乘。算力没用多少,时间全花在搬运上。

FlashAttention 把这三次往返合并成一次 kernel:分数矩阵分块在 SRAM 里边算边扔。于是:

  • 训练侧:序列变长、batch 变大、不用 activation checkpointing 重算注意力。
  • 推理侧:KV cache 越长,收益越大;flash_attn_with_kvcache 是解码期标配。

它能做什么

能力具体支持出处
精确注意力前向+反向fp16/bf16,head_dim ≤ 256csrc/flash_attn/flash_api.cpp:397csrc/flash_attn/flash_api.cpp:416
causal / 滑窗 / ALiBi / softcap编译期模板开关csrc/flash_attn/src/flash_fwd_launch_template.h:76
变长 batch(varlen)cu_seqlens 打包,无 padding 浪费flash_attn/flash_attn_interface.py:1391
推理 KV cache原地追加、paged KV、split-KV 解码flash_attn/flash_attn_interface.py:1485
MQA / GQAK/V 头数整除 Q 头数即可csrc/flash_attn/flash_api.cpp:418
FA3(Hopper)FP16/BF16 前后向 + FP8 前向,H100/H800README.md:30hopper/flash_attn_interface.py:809
FA4(CuTeDSL)Hopper + Blackwell,Python 写 kernelflash_attn/cute/README.md:3

用起来什么样

# 摘自 README.md「How to use」一节的调用形态
from flash_attn import flash_attn_func

# q, k, v: (batch, seqlen, nheads, headdim),fp16/bf16,CUDA
out = flash_attn_func(q, k, v, dropout_p=0.0, causal=True)

# 变长(batch 内各序列不同长,打包存储)
from flash_attn import flash_attn_varlen_func
out = flash_attn_varlen_func(q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k)

# 推理解码:KV cache 原地更新 + rotary 一步到位
from flash_attn import flash_attn_with_kvcache
out = flash_attn_with_kvcache(q, k_cache, v_cache, k, v, cache_seqlens=cache_seqlens, causal=True)

返回的就是注意力输出,可以直接替换 torch.nn.functional.scaled_dot_product_attention 的位置——但显存曲线会立刻变平。

一句话直觉

把显存(HBM)当成磁盘、把片上 SRAM 当成内存。 标准注意力每算一步都把中间结果「落盘」;FlashAttention 全程在「内存」里算完,只把最终结果写回去。数学一行没改,路少走了很多。


2. 顶层全景(它大概怎么转)

2.1 从 Python 到 CUDA kernel 的调用链

flash_attn_func(q, k, v, causal=True)


① torch custom op 包装 flash_attn/flash_attn_interface.py
(autograd.Function + torch.compile fake)


② pybind11 入口 mha_fwd csrc/flash_attn/flash_api.cpp
检查 dtype/shape,组装 Flash_fwd_params


③ 编译期派发 csrc/flash_attn/src/
dtype × head_dim × causal × 架构,上百个实例化 .cu


④ CUDA kernel csrc/flash_attn/src/flash_fwd_kernel.h
每个 CTA 吃一个 Q tile,K/V 分块流过 SRAM,online softmax


out + softmax_lse ← N×N 矩阵全程不曾落地 HBM

怎么读这张图: ①②是薄封装,真正的复杂度全在③的「编译期组合爆炸」和④的「单 kernel 融合主循环」。训练时反向走同样一条路,只是 kernel 换成重算版(见第 04 章)。

2.2 部件职责

部件干什么在哪个文件
Python 接口层参数规整、autograd、torch.compile 注册flash_attn/flash_attn_interface.py:1156(flash_attn_func)
C++/pybind 入口输入校验、参数结构、grid/split 启发式csrc/flash_attn/flash_api.cpp:369(mha_fwd)
前向 kerneltiling 主循环 + online softmaxcsrc/flash_attn/src/flash_fwd_kernel.h:55(compute_attn_1rowblock)
softmax 原语running max/sum、rescale、LSEcsrc/flash_attn/src/softmax.h:129(struct Softmax)
反向 kernel重算 P、累 dK/dV、atomic dQcsrc/flash_attn/src/flash_bwd_kernel.h:81(compute_dq_dk_dv_1colblock)
kernel 配置tile 尺寸、smem 布局、MMA atomcsrc/flash_attn/src/kernel_traits.h:51(Flash_fwd_kernel_traits)
FA3(Hopper)warp specialization + TMA/WGMMAhopper/flash_fwd_kernel_sm90.h:29(FlashAttnFwdSm90)
FA4(Blackwell)CuTeDSL 纯 Python kernelflash_attn/cute/flash_fwd.py:40(FlashAttentionForwardBase)
层封装MHA / 自注意力 / 交叉注意力 nn.Moduleflash_attn/modules/mha.py:373(MHA)

2.3 主线走一遍:一次前向调用

  1. Python 侧:FlashAttnFunc.forward(flash_attn/flash_attn_interface.py:830)把 head_dim pad 到 8 的倍数,调用 custom op。
  2. C++ 侧:mha_fwd 检查「Ampere 或更新、fp16/bf16、head_dim ≤ 256 且 % 8 == 0」(csrc/flash_attn/flash_api.cpp:393-416),填 Flash_fwd_params(csrc/flash_attn/src/flash.h:47)。
  3. 派发:run_mha_fwd(csrc/flash_attn/flash_api.cpp:261)按 dtype/head_dim/causal 三层 switch 选实例化 kernel;序列短且 KV 长时走 split-KV 分支。
  4. kernel:grid = (Q 行块数, batch, head)(csrc/flash_attn/src/flash_fwd_launch_template.h:64)。每个 CTA 把 Q tile 固定在 SRAM/寄存器,循环流过 K/V tile,边算边 rescale 输出累加器。
  5. 收尾:除以行和、把 O 与 LSE(logsumexp,反向要用)写回 HBM。N×N 的 S、P 从未出现在显存里。
  6. 反向(若需梯度):autograd 只保存 q, k, v, out, softmax_lse, rng_state(flash_attn/flash_attn_interface.py:869),反向 kernel 现场重算 S/P 再求梯度。

2.4 四代同堂:一个仓库里的四份实现

目录目标硬件写法状态(本 commit)
FA1/FA2csrc/flash_attn/sm80-90(A100/3090/H100)CUDA C++ + CUTLASS/CuTe 模板主线,版本 2.8.4(flash_attn/__init__.py:5)
FA3hopper/sm90(H100/H800)CUTLASS 3 风格 collective + warp specializationbeta,独立安装(README.md:30)
FA4flash_attn/cute/sm80/90/100/120CuTeDSL,Python 即 kernel现行主推,pip install flash-attn-4
Triton 参考版flash_attn/flash_attn_triton.py实验性Triton教学/实验用

3. 阅读地图

按「由浅入深」排序,前五章是一条线;每章可独立阅读,但概念层层依赖。

讲什么适合谁
01-io-bottleneck.md为什么标准注意力慢:内存层级、算术强度、N² 读写所有人,先读
02-online-softmax.md分块 softmax 的数学与 softmax.h 的逐行对应想懂原理的人
03-forward-kernel.md前向 kernel 全貌:tiling、双缓冲、mask、split-KV想读 CUDA 源码的人
04-backward-kernel.md反向的重计算设计与非显然的工程取舍想读 CUDA 源码的人
05-kernel-engineering.md模板实例化、FA3/FA4 怎么压榨新硬件、Python 生态想做 kernel 工程的人

只想用:读第 1 章 + 本文 §1 即可。 想仿写一个 fused kernel:读 01 → 02 → 03,然后对照 flash_attn/flash_attn_triton.py(同一算法的 Triton 版,行数少一个量级)。 想知道 H100/B200 上还能怎么快:读 05。


4. 巧妙之处(可借鉴的技术)

每条先白话,再上锚点;细节在对应章节展开。

  1. 把 softmax 拆成可在线更新的两个标量。 不存整行分数,只维护 running max m 和 running sum l,每来一块 K 就把旧输出乘 e^(m_old − m_new) 修正——softmax 从「看全行才能算」变成「流过即算」。csrc/flash_attn/src/softmax.h:137(softmax_rescale_o)。见 02 章。

  2. 用 exp2 换掉 exp,顺手白赚一次 FMA。 分数预先乘 log2(e),exp(x−m)exp2(x·log2e − m·log2e),编译器能合成单条 FMA;硬件本身 exp2 也更快。csrc/flash_attn/src/softmax.h:67(scale_apply_exp2)。见 02 章。

  3. 反向不存 P、只存 LSE,现场重算。 反向需要 P=N×N,但存它比算它贵;于是只存每行的 logsumexp(O(N)),反向 kernel 用 P = exp2(S·scale − LSE) 重算。flash_attn/flash_attn_interface.py:869 + csrc/flash_attn/src/flash_bwd_kernel.h:536。见 04 章。

  4. 反向换一个并行轴。 前向按 Q 行块分 CTA;反向改成按 K/V 列块分,dK/dV 留在寄存器里累加,dQ 用 atomicAdd 甩给显存——把「最热的累加」留在片上,把「最冷的累加」交给硬件原子操作。csrc/flash_attn/src/flash_bwd_kernel.h:800csrc/flash_attn/src/flash_bwd_kernel.h:678。见 04 章。

  5. 编译期把「是否偶数长度」也模板化。 Is_even_MNIs_even_K 为真时 kernel 里所有边界检查整段消失;代价是实例化数量爆炸(于是有了代码生成器)。csrc/flash_attn/src/flash_fwd_launch_template.h:76csrc/flash_attn/src/generate_kernels.py:13。见 05 章。

  6. 短 Q 长 KV 时把「头」折进序列维。 解码期 seqlen_q=1 且 GQA 时,把 Q 的组维 reshape 成序列维,一行查询变多行,kernel 形状立刻变好。csrc/flash_attn/flash_api.cpp:431(seqlenq_ngroups_swapped)。见 03 章。

  7. FA3 让 GEMM 和 softmax 互相打掩护。 Hopper 上一个 warpgroup 的 WGMMA(异步矩阵乘)在跑,另一个 warpgroup 的 CUDA core 同时做 softmax——两种硬件单元互相填对方的空泡。hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:1169。见 05 章。


5. 边界与局限

诚实清单——它刻意不做什么、会在哪崩。

  • 是精确实现,不是近似/稀疏/线性注意力。 数学等于标准注意力,只改 IO。想要亚二次复杂度请找别的家族(这一点源卡片也点名强调)。
  • dtype 只有 fp16/bf16(FA3 前向另有 FP8)。fp32 直接拒绝:csrc/flash_attn/flash_api.cpp:397
  • head_dim ≤ 256 且必须是 8 的倍数(csrc/flash_attn/flash_api.cpp:416);> 256 没有实现。
  • 架构锁死:FA2 要 sm80+(Ampere 起,csrc/flash_attn/flash_api.cpp:393);Turing 被移交给社区 fork;FA3 只在 H100/H800;FA4 面向 Hopper/Blackwell。每一代 kernel 都是按硬件重写的,不存在「一份 kernel 处处最快」。
  • 反向默认不确定:atomicAdd 累 dQ,浮点加法顺序不定 → 两次反向结果可能逐位不同。传 deterministic=True 换确定性,更慢更费显存(README.md:245 接口说明;实现见 csrc/flash_attn/src/flash_bwd_launch_template.h:78)。
  • 数值容差靠测试兜底,不靠证明:测试标准是「误差 ≤ PyTorch 基线误差的两倍」(README.md:551 Tests 一节)。
  • 编译代价高:无 ninja 单线程编译可达 2 小时,64 核 + ninja 也要 3–5 分钟(README.md:114)。几百个模板实例化是根源。
  • dropout 组合受限:softcap 不与 dropout 共存(csrc/flash_attn/flash_api.cpp:420);return_attn_probs 只在 dropout>0 时可用(csrc/flash_attn/flash_api.cpp:469 附近)。
  • Windows 支持是「有人跑通过」级别(README.md:108)。

6. 横向对比

同一书架上,这几个项目和 FlashAttention 的关系各不相同:

项目与 FlashAttention 的关系取舍差异
PyTorchF.scaled_dot_product_attention 的上游;SDPA 的 flash 后端即本库的封装形态PyTorch 要通用后端选择(cudnn/mem-efficient/math),本库只要单点极致
vLLM下游用户:paged attention / prefill 里调用 FA 系 kernelvLLM 关心 KV cache 管理与调度,注意力本体交给 kernel 库
SGLang下游用户:RadixAttention 的底层注意力同样走 FA 系 kernel同上,竞争点在前缀复用与调度,不在注意力数学
Unsloth平级「训练加速」玩家:用 Triton 重写一批 kernel(含注意力路径)Unsloth 重「少显存 + 免编译安装」,FA 重「极限吞吐 + 手写 CUDA」

正交对照:xformers 的 memory-efficient attention(Rabe & Staats)与本库同一思想的不同实现,README 在 split-KV 一节公开致谢过 xformers 团队(README.md:459)。Triton 教程版 fused-attention 是读算法的最短路径,本仓库也自带一份 Triton 版(flash_attn/flash_attn_triton.py),适合先读它再读 CUDA。


7. 代码地图(导航索引)

按「我想知道 X 在哪」查。符号名可直接 grep。

主题文件路径符号名
Python 入口(前向)flash_attn/flash_attn_interface.pyflash_attn_funcFlashAttnFunc
推理 KV cache 入口flash_attn/flash_attn_interface.pyflash_attn_with_kvcache
C++ 入口/校验/派发csrc/flash_attn/flash_api.cppmha_fwdrun_mha_fwdnum_splits_heuristic
参数结构体csrc/flash_attn/src/flash.hFlash_fwd_paramsFlash_bwd_params
前向 kernel 主循环csrc/flash_attn/src/flash_fwd_kernel.hcompute_attn_1rowblockcompute_attn
online softmaxcsrc/flash_attn/src/softmax.hSoftmax::softmax_rescale_oscale_apply_exp2normalize_softmax_lse
mask(causal/local/越界)csrc/flash_attn/src/mask.hMask::apply_maskapply_mask_causal
split-KV 合并csrc/flash_attn/src/flash_fwd_kernel.hcombine_attn_seqk_parallel
反向 kernelcsrc/flash_attn/src/flash_bwd_kernel.hcompute_dq_dk_dv_1colblockcompute_dq_dk_dv
反向预处理(dO·O)csrc/flash_attn/src/flash_bwd_preprocess_kernel.hdot_do_o
反向 launch/确定性csrc/flash_attn/src/flash_bwd_launch_template.hrun_flash_bwd
tile 尺寸/smem 布局csrc/flash_attn/src/kernel_traits.hFlash_kernel_traitsFlash_fwd_kernel_traits
实例化代码生成csrc/flash_attn/src/generate_kernels.pyHEAD_DIMENSIONSget_fwd_template
FA3 kernel(SM90)hopper/flash_fwd_kernel_sm90.hFlashAttnFwdSm90
FA3 mainloop(TMA/WGMMA)hopper/mainloop_fwd_sm90_tma_gmma_ws.hppCollectiveMainloopFwdSm90
FA3 持久化调度hopper/tile_scheduler.hppDynamicPersistentTileScheduler
FA4(CuTeDSL)flash_attn/cute/flash_fwd.pyFlashAttentionForwardBaseFlashAttentionForwardSm80
MHA 层封装flash_attn/modules/mha.pyMHAFlashSelfAttention
数值正确性测试tests/test_flash_attn.pytest_flash_attn_output