跳到主要内容

数据截至 (上游 commit 0251105a2fb1)

05 · kernel 工程与四代演化

这一章讲什么: 第 03、04 章讲的是「一个 kernel」长什么样;这一章讲的是「一个 kernel 怎么变成一个库」——模板实例化的组合爆炸怎么管、新硬件出来怎么跟、以及 Python 侧怎么把它接到 PyTorch 生态里。


1. 它要解决的小问题

两个工程难题:

  • 组合爆炸:最优 tile 尺寸随 head_dim、dtype、GPU 型号、是否 causal/dropout 而变;每个组合都要单独编译,因为它们决定 smem 用量和寄存器压力——都必须是编译期常量
  • 硬件换代:Ampere 的最佳形态(cp.async + 16x8x16 MMA)到 Hopper(TMA + WGMMA)到 Blackwell(tcgen05)全部失效,kernel 要按代重写。

2. 实例化矩阵:一个算法,几百个编译单元

2.1 组合从哪来

代码生成器 csrc/flash_attn/src/generate_kernels.py 给出了基础维度(csrc/flash_attn/src/generate_kernels.py:13-15):

维度取值
架构sm80
head_dim32, 64, 96, 128, 192, 256
dtypefp16, bf16
causalfalse, true

再乘上前向/反向/split-KV 三类 kernel,就是 csrc/flash_attn/src/ 下一百多个 flash_*_hdim*_sm80.cu 文件——每个文件只实例化一个组合,以便并行编译(setup.py:351 起逐个列出)。

2.2 运行时 Bool 也是编译期常量

Is_even_MN(序列长度是否整除 tile)、Is_dropoutHas_alibiIs_softcap 这些运行时才知道的标志,也全部模板化。机制是 BOOL_SWITCH(csrc/flash_attn/src/static_switch.h:17-24):运行时用 if 把 bool 分叉成两个 constexpr 分支,每个分支实例化一份 kernel。run_flash_fwd 里嵌套了五层这样的 switch(csrc/flash_attn/src/flash_fwd_launch_template.h:76-88)。

2.3 组合怎么收敛

全开会编出天文数字,于是 launch 模板里手工裁剪——例如 Is_even_MN 只在 IsEvenKConst && !Is_local && !Has_alibi && !ReturnSoftmaxConst && kHeadDim <= 128 时才保留 true 分支,ReturnSoftmaxConst 被收窄到 Is_dropout && !Is_softcap(csrc/flash_attn/src/flash_fwd_launch_template.h:84-86)。裁剪规则本身就是工程知识:稀有组合回退到通用(慢)路径,常见组合才有专属 kernel。

2.4 代价:编译时间

README 自己给了数字:无 ninja 串行编译可达 2 小时;64 核 + ninja 也要 3–5 分钟(README.md:111-116)。内存不够时用 MAX_JOBS 限流(README.md:129-131)。这就是为什么官方发预编译 wheel,也是「改一行 kernel 重编一下午」这个社区梗的来源。


3. traits:把「硬件知识」集中到一处

所有 tile 尺寸、smem 布局、MMA 指令的选择都收在 kernel_traits.h:

决策内容锚点
MMA 指令sm80+:fp16 用 SM80_16x8x16_F32F16F16F32_TN;老架构退到 SM75 atomcsrc/flash_attn/src/kernel_traits.h:30-36
smem→寄存器拷贝ldmatrix(SM75_U32x4_LDSM_N / U16x8_LDSM_T)csrc/flash_attn/src/kernel_traits.h:40-43
warp 排布kNWarps 个 warp 沿 M 维一字排开csrc/flash_attn/src/kernel_traits.h:74-80
smem 防 bank conflictSwizzle 异或打乱 + 拷贝线程布局调优(注释记录 d=128 快 6–10%)csrc/flash_attn/src/kernel_traits.h:72csrc/flash_attn/src/kernel_traits.h:114-119
smem 总账Q、K、V 布局之和(Q/K 可共享则取 max)csrc/flash_attn/src/kernel_traits.h:109

读 traits 文件是理解「这个 kernel 为什么这么摆数据」的最短路径。


4. FA3:为 Hopper 重写一遍

Hopper(H100)给了三件新玩具:TMA(硬件异步大块搬运,不占线程)、WGMMA(warpgroup 级异步矩阵乘,也不占线程)、更大的 smem。FA3(hopper/)围绕「让三种硬件单元同时忙」重组了 kernel。

4.1 warp specialization:生产者/消费者分工

一个 CTA 内的 warpgroup 分两种角色(hopper/flash_fwd_kernel_sm90.h:307 起):

生产者 warpgroup ×1 消费者 warpgroup ×1–2
TMA 搬 K/V → SRAM 多级流水 ──► WGMMA 算 GEMM
(寄存器减配到 24–56) pipeline CUDA core 同时做 softmax
barrier (寄存器加配到 160–256)

寄存器是不平均分的:生产者干的是「发 TMA 指令」的轻活,warpgroup_reg_dealloc 把寄存器让到 24–56;消费者要同时握住 Q、O、S 的累加器,warpgroup_reg_alloc 加到 160–256(hopper/flash_fwd_kernel_sm90.h:309hopper/flash_fwd_kernel_sm90.h:361,配额常量见 hopper/flash_fwd_kernel_sm90.h:82-85)。

4.2 GEMM 与 softmax 交叠

WGMMA 是异步的:发完矩阵乘指令,CUDA core 就空出来了。FA3 的 mainloop 把迭代软件流水化——第 n 块的 GEMM 在跑时,同时做第 n−1 块的 softmax。注释写得很凝练:「Each step does gemm0 for iter n_block, gemm1 for iter n_block + 1, and softmax for iter n_block」(hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:1169);两个消费者 warpgroup 之间再做 pingpong,让 tensor core 几乎没有空泡。

4.3 持久化调度与多级流水

  • K/V 的搬运走 PipelineTmaAsync 多级流水(hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:283-286),smem 分 stage 循环使用;
  • tile 分配不再是一个 CTA 一个 tile,而是 DynamicPersistentTileScheduler(hopper/tile_scheduler.hpp:220):CTA 数量 ≈ SM 数,循环领工单——长序列负载均衡、省 launch 开销。

4.4 FP8 前向

FA3 支持 FP8 前向(README.md:43)。一个细节:FP8 的指数值域太窄,softmax 时把 max 再减 8,让 exp2 的结果落在 [0, 256] 而不是 [0, 1],「use more of the FP8 range to reduce underflow」(hopper/softmax.h:65-70 注释)。


5. FA4:用 Python 写 kernel(CuTeDSL)

flash_attn/cute/ 是第四代:kernel 不再用 C++ 写,而是用 NVIDIA 的 CuTeDSL(Python 前端、编译到 PTX)写。

对比项FA2(CUDA C++)FA4(CuTeDSL)
kernel 语言C++ 模板Python + @cute.jit 装饰器(flash_attn/cute/flash_fwd.py:311)
编译时机安装期,全量预编译首次调用时 JIT,按 compile_key 缓存(flash_attn/cute/interface.py:1389)
架构覆盖sm80–90sm80/90/100/120,按 arch 选实现类(flash_attn/cute/interface.py:681)
实现组织一套模板打天下每代一个类:FlashAttentionForwardBase(flash_attn/cute/flash_fwd.py:40)、...Sm90(flash_attn/cute/flash_fwd_sm90.py:52)、...Sm100(flash_attn/cute/flash_fwd_sm100.py:136)
交付主包 flash-attn独立包 flash-attn-4(flash_attn/cute/README.md:8)

意义不止是「Python 更好写」:JIT 让「按 shape/特性组合生成 kernel」从安装期的排列组合变成运行期按需编译,第 2 节的组合爆炸在模型上换了打法。它还顺势吃进 Blackwell——同一个仓库里第四份实现,但维护成本远低于再写一遍 CUDA C++(inferred)。


6. Python 接口层:接进 PyTorch 生态

6.1 三层封装

flash_attn_func() 用户API,docstring 即合同
└─ FlashAttnFunc(autograd.Function) 定 forward/backward 存什么
└─ torch custom op 让 torch.compile 看得见
└─ pybind11 mha_fwd 进 C++
  • pybind 导出在 PYBIND11_MODULE(csrc/flash_attn/flash_api.cpp:1535-1541):fwdvarlen_fwdbwdvarlen_bwdfwd_kvcache
  • PyTorch ≥2.4 时注册为正式 custom op,并配 register_fake 版本——torch.compile 的 fake tensor 推导因此能穿过这个 op(flash_attn/flash_attn_interface.py:63-82 的兼容层、flash_attn/flash_attn_interface.py:118_flash_attn_forward_fake)。
  • head_dim 不是 8 的倍数时,Python 侧先 pad 再进 kernel,返回时切回(flash_attn/flash_attn_interface.py:849-853)——约束在 kernel,容错在封装。

6.2 层与配套件

flash_attn/modules/mha.py:373MHA 是即插即用的 nn.Module(自/交叉注意力、rotary、KV cache 推理路径 _update_kv_cacheflash_attn/modules/mha.py:344);flash_attn/ops/ 还有 layer norm、fused dense 等周边 kernel;flash_attn/flash_attn_triton.py 是同算法的 Triton 版,行数少一个量级,读算法先看它

6.3 生态位置

下游几乎人手一份:vLLM、SGLang 等推理引擎的注意力后端,HuggingFace Transformers 的 attn_implementation="flash_attention_2",以及 Megatron 系训练框架(采用清单见仓库 usage.md)。在本书架上可对照 vLLMSGLangPyTorchUnsloth 四篇 teardown——前两者是调用方,PyTorch 是把同一思想收编为标准 API 的一方,Unsloth 是用 Triton 复刻的一方。


7. 关键细节 / 坑

  1. 两份 block size 表必须手工同步。 Python 侧 _get_block_size_n 的注释明说「should match the block sizes in the CUDA kernel」(flash_attn/flash_attn_interface.py:31)——没有单一事实源,改一边忘另一边就是 bug。
  2. FA3 的安装是独立的。 cd hopper && python setup.py install,包名 flash_attn_3,接口与 FA2 分开(README.md:49-52);FA4 也是独立包。flash_attn/__init__.py:3-4extend_path 让 FA2 和 FA4 可以共存在一个命名空间下。
  3. 编译期裁剪是隐性文档。 想知道「哪个组合有快路径」,答案不在 README,在 launch 模板那几行收窄条件里(csrc/flash_attn/src/flash_fwd_launch_template.h:84-88)。
  4. 四代并存,选型看硬件。 A100 → FA2;H100 追求极限 → FA3;B200 或想要 JIT/新特性(score_mod、learnable sink 等,见 flash_attn/cute/interface.py:3371 的参数表)→ FA4。

8. 本章小结

  • 组合爆炸的解法:编译期模板 + 代码生成 + 手工裁剪 + 并行编译单元;代价是编译时间。
  • 硬件换代的解法:每代重写(FA3 的 warp specialization、FA4 的 CuTeDSL),算法不变,组织方式全变。
  • 生态对接的解法:pybind → custom op(+fake)→ autograd.Function → nn.Module,一层一个职责。
  • 回看全系列:IO 分析(01)给出方向,online softmax(02)给出数学,前向/反向(03/04)给出实现,本章给出它成为「基础设施」的工程形态。