数据截至 (上游 commit 0251105a2fb1)
03 · 前向 CUDA kernel
这一章讲什么: 把第 01 章的 tiling 骨架 + 第 02 章的 online softmax,落到
csrc/flash_attn/src/flash_fwd_kernel.h的真实 CUDA kernel 上:grid 怎么分、主循环每一步干什么、mask 怎么省、以及推理期的 split-KV 是怎么回事。
1. 它要解决的小问题
数学正确只是入场券。kernel 还要回答三个工程问题:
- 并行度:上百万个 tile,怎么映射到上百个 SM?
- 吞吐:搬运(cp.async)和计算(tensor core)怎么互相掩护,谁都不等谁?
- 特判:causal mask、越界、变长序列,怎么不让边界检查拖慢主循环?
2. 网格划分:一个 CTA 一个 Q 行块
前向的 launch 配置在 run_flash_fwd(csrc/flash_attn/src/flash_fwd_launch_template.h:55),grid 就一行:
// csrc/flash_attn/src/flash_fwd_launch_template.h:63-64
const int num_m_block = (params.seqlen_q + Kernel_traits::kBlockM - 1) / Kernel_traits::kBlockM;
dim3 grid(num_m_block, params.b, params.h);
每个 CTA(thread block)负责一个 (Q 行块 × batch × head),在函数 compute_attn_1rowblock(csrc/flash_attn/src/flash_fwd_kernel.h:55)里完成:
- 把自己的 Q tile 一次载入 SRAM(
kBlockM常见 128/64 行); - 循环:K/V 按
kBlockN大小的列块依次流过 SRAM; - 输出 O tile 在寄存器里累加,循环结束才写回。
Q 不动、K/V 流动——这样每个 K/V 元素从 HBM 只读一次。