数据截至 (上游 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_dim | 32, 64, 96, 128, 192, 256 |
| dtype | fp16, bf16 |
| causal | false, true |
再乘上前向/反向/split-KV 三类 kernel,就是 csrc/flash_attn/src/ 下一百多个 flash_*_hdim*_sm80.cu 文件——每个文件只实例化一个组合,以便并行编译(setup.py:351 起逐个列出)。
2.2 运行时 Bool 也是编译期常量
Is_even_MN(序列长度是否整除 tile)、Is_dropout、Has_alibi、Is_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 atom | csrc/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 conflict | Swizzle 异或打乱 + 拷贝线程布局调优(注释记录 d=128 快 6–10%) | csrc/flash_attn/src/kernel_traits.h:72、csrc/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:309、hopper/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–90 | sm80/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):fwd、varlen_fwd、bwd、varlen_bwd、fwd_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:373 的 MHA 是即插即用的 nn.Module(自/交叉注意力、rotary、KV cache 推理路径 _update_kv_cache 见 flash_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)。在本书架上可对照 vLLM、SGLang、PyTorch、Unsloth 四篇 teardown——前两者是调用方,PyTorch 是把同一思想收编为标准 API 的一方,Unsloth 是用 Triton 复刻的一方。
7. 关键细节 / 坑
- 两份 block size 表必须手工同步。 Python 侧
_get_block_size_n的注释明说「should match the block sizes in the CUDA kernel」(flash_attn/flash_attn_interface.py:31)——没有单一事实源,改一边忘另一边就是 bug。 - FA3 的安装是独立的。
cd hopper && python setup.py install,包名flash_attn_3,接口与 FA2 分开(README.md:49-52);FA4 也是独立包。flash_attn/__init__.py:3-4用extend_path让 FA2 和 FA4 可以共存在一个命名空间下。 - 编译期裁剪是隐性文档。 想知道「哪个组合有快路径」,答案不在 README,在 launch 模板那几行收窄条件里(
csrc/flash_attn/src/flash_fwd_launch_template.h:84-88)。 - 四代并存,选型看硬件。 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)给出实现,本章给出它成为「基础设施」的工程形态。