CUDA 算子工程:手写 FlashAttention v2 之路
第 17 章 TMA + Warp Specialization 把 FA2 写到 SOTA
Hopper 不是"Tensor Core 更快的 Ampere",它是另一套编程模型:TMA、WGMMA、warp specialization。 FA2 在 H100 上只跑出约 35% 利用率,正是没用上这套模型的代价。
17.1 为什么 Hopper 上的写法不一样
第 15 章我们写的 FA2 forward 用的是 Ampere 时代的工具:cp.async 异步拷贝 + mma.sync.m16n8k16 矩阵乘。这套工具在 H100 上还能用,但把 Tensor Core 峰值只用出三成半左右——FA3 论文测的 FA2 在 H100 FP16 上约 35% 利用率,折算约 350 TFLOPs。
如果换成 Hopper 原生工具:
- TMA 替代
cp.async - WGMMA 替代
mma.sync - Warp Specialization 替代对称 warp
性能能提升到 ~740 TFLOPs(FP16,75% 利用率),FP8 则接近 1.2 PFLOPs。
数据来源:Shah et al., FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision, 2024(arXiv:2407.08608)摘要与 §4.1。同一篇论文给出的 FA2 基线是 ~35% 利用率——本专栏
00-preface.md用的也是这个口径。
为什么差距这么大?三个原因:
- TMA 比 cp.async 更高效:单线程发起、专用硬件、原生 swizzle、不占 SIMT 算术单元。
- WGMMA 是异步指令:发完不阻塞,warp 可以继续算/拷下一份。
- Warp Specialization 把"算"和"拷"真正分离:Producer 永远在拷,Consumer 永远在算,硬件流水拉满。在此之上 FA3 还有两层调度上的重叠(ping-pong 与 warp-group 内 GEMM-softmax 流水),见 17.5 节末。
17.2 Producer / Consumer 的角色分配
把 FA2 搬到 Hopper(也就是 FA3)的核心 idea 是按 warp-group(128 线程 = 4 warp) 划分角色。FA3 官方实现里这一划分写在 flash-attn/hopper/flash_fwd_kernel_sm90.h:74:
static constexpr uint32_t NumLoadWarpGroups = 1;
static constexpr uint32_t NumMmaWarpGroups =
CUTE_STATIC_V(size(TiledMmaPV{})) / cutlass::NumThreadsPerWarpGroup;
static_assert(NumMmaWarpGroups == 1 || NumMmaWarpGroups == 2 || NumMmaWarpGroups == 3);
也就是:固定 1 个 producer warp-group + 1~3 个 consumer warp-group。consumer 的个数等于 kBlockM / 64(flash-attn/hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:89 的 AtomLayoutQK),headdim 128 的 kBlockM 是 128,所以典型是 2 个,整个 block 384 线程。producer warp-group 虽然占了 128 线程的名额,但真正发 TMA 的只有它的 warp 0(flash-attn/hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:798 里再 elect_one_sync() 选出一个线程);Q/K/V 都走 TMA 且不需要转置 V 时,其余 3 个 warp 直接退出(源码里的 SingleProducerWarp 分支,flash-attn/hopper/flash_fwd_kernel_sm90.h:319),FP8 且 V 为行主序时它们要留下来在 SMEM 里转置 V(见 17.7):
flowchart TB
subgraph PRO [Producer Warp Group · 128 线程, 只有 warp 0 发 TMA]
P1[发起 TMA: K, V tile]
P2[mbarrier.arrive.expect_tx 登记字节数, TMA 落地后自动 complete_tx]
end
subgraph CONS [Consumer Warp Groups · 典型 2 个, 256 线程]
C1[mbarrier.wait 等数据]
C2[WGMMA 发起 S = Q @ K^T]
C3[Online softmax]
C4[WGMMA 发起 O += P @ V]
end
PRO -->|信号: tile k ready| CONS
CONS -->|信号: tile k consumed| PRO
Producer 和 Consumer 在物理上是同一个 thread block 的不同 warp-group,通过 mbarrier 同步。它们各自专注自己的事,硬件层面真正异步并行。
17.3 TMA Descriptor 的构建
TMA 的关键是预先构建 TMA Descriptor(CUDA 里的类型叫 CUtensorMap)——一份描述张量布局、stride、swizzle 模式的元数据,每条 TMA 指令都要引用它。CUDA Programming Guide 推荐的传法是把它作为 const __grid_constant__ 的 kernel 参数按值传进去;另两种做法是用 cudaMemcpyToSymbol 拷进 __constant__ 变量,或放在 global memory 里(此时每个 block 首次使用前要加 tensormap proxy fence,而且可能更慢)。
构建 TMA descriptor 在 host 端完成:
// Host 端构建 TMA descriptor(driver API,链接时加 -lcuda)
// K 在 HBM 里按 [B, H, N, d] 行主序存放,维度从最内层(连续)往外列
CUtensorMap tma_desc_K;
cuuint64_t gdim[4] = {d, N, H, B};
cuuint64_t gstride[3] = {d * 2, N * d * 2, H * N * d * 2}; // 字节; 最内层 stride 隐含为元素大小, 所以只给 rank-1 个
cuuint32_t box[4] = {d, Bc, 1, 1}; // 每次拷贝的 tile
cuuint32_t estride[4] = {1, 1, 1, 1};
cuTensorMapEncodeTiled(
&tma_desc_K,
CU_TENSOR_MAP_DATA_TYPE_FLOAT16,
/*tensorRank=*/4,
K_global_ptr,
gdim, gstride, box, estride,
CU_TENSOR_MAP_INTERLEAVE_NONE,
CU_TENSOR_MAP_SWIZZLE_128B, // 128B swizzle
CU_TENSOR_MAP_L2_PROMOTION_L2_128B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
// tma_desc_K 直接作为 const __grid_constant__ 参数传给 kernel, 不需要另拷到 device memory
cuTensorMapEncodeTiled 是 CUDA 12.0+ 的 API,专门用来构建 TMA descriptor。注意 swizzle 对 box 有约束:box 最内层的字节数不能超过 swizzle 的跨度。这里 FP16、d = 64 正好 128 字节;d = 128 时一行 256 字节,不能再用一个 {d, Bc} 的 box,要把 d 方向拆成两个 64 列的 box 分两次发。
Kernel 端用 descriptor 发 TMA:
__global__ void fa3_fwd(
const __grid_constant__ CUtensorMap tma_desc_Q,
const __grid_constant__ CUtensorMap tma_desc_K,
const __grid_constant__ CUtensorMap tma_desc_V,
half* O,
...
) {
// 128B swizzle 的 pattern 每 1024 字节重复, 目标地址按 1024 对齐;
// 注意写成 alignas(1024) extern __shared__ ... 或 __align__, 把 alignas 放在 half 前面 nvcc 会报错
extern __shared__ __align__(1024) half smem[];
half* sQ = smem;
half* sK[STAGES]; // 多 stage pipeline
half* sV[STAGES];
// 设置 sK[i], sV[i] 指针 ...
__shared__ alignas(8) uint64_t mbar[STAGES * 2]; // 每 stage 两个 barrier (full/empty)
if (threadIdx.x == 0) {
for (int s = 0; s < STAGES * 2; ++s) {
mbarrier_init(&mbar[s], /*count=*/...);
}
fence_mbarrier_init(); // fence.mbarrier_init: 让初始化对 TMA 所在的异步代理可见
}
__syncthreads();
// 角色分工: 按 warp-group (128 线程) 切
int warp_group_idx = threadIdx.x / 128;
if (warp_group_idx == 0) {
// Producer warp group (只有它的 warp 0 真正发 TMA)
producer_main(tma_desc_K, tma_desc_V, sK, sV, mbar);
} else {
// Consumer warp groups
consumer_main(sQ, sK, sV, mbar, /* O accumulator */);
}
}
__grid_constant__ 不是 Hopper 专属(nvcc 13.4 以 -arch=sm_80 编译带它的 kernel 照样通过),它的作用是让整个 grid 通过同一个地址直接读这个 const 参数,编译器不再给每个线程拷一份副本——TMA 指令要的正是 descriptor 的地址,所以 Programming Guide 推荐用它传 CUtensorMap。
17.4 Producer Warp 的工作
__device__ void producer_main(
const CUtensorMap& tma_desc_K,
const CUtensorMap& tma_desc_V,
half** sK, half** sV,
uint64_t* mbar
) {
// (setmaxnreg.dec 要在这里之前由整个 warp-group 执行, 见下文)
if (threadIdx.x % 128 != 0) return; // 只留 warp 0 的 lane 0 发 TMA
int n_k_tiles = N / Bc;
for (int k = 0; k < n_k_tiles; ++k) {
int stage = k % STAGES;
// 等当前 stage 被 consumer 消费完 (empty barrier); 前 STAGES 轮缓冲本来就空, 不用等
if (k >= STAGES) mbarrier_wait(&mbar[stage * 2 + 1], /* phase = ... */);
// 先在 full barrier 上 arrive 并登记本 stage 要到账的字节数 (K + V 各一个 tile)
mbarrier_arrive_expect_tx(&mbar[stage * 2 + 0], 2 * Bc * d * sizeof(half));
// 发起 K[k], V[k] 的 TMA; 坐标顺序与 descriptor 的维度一致 {d, N, H, B}
cp_async_bulk_tensor_4d(
sK[stage], &tma_desc_K, /*coords=*/0, k * Bc, head_idx, batch_idx,
&mbar[stage * 2 + 0] // 数据落地时硬件在 full barrier 上 complete_tx
);
cp_async_bulk_tensor_4d(
sV[stage], &tma_desc_V, 0, k * Bc, head_idx, batch_idx, &mbar[stage * 2 + 0]
);
}
}
Producer 的循环极其简单——它只有一件事:发起 TMA、等 consumer 消费完、发下一个。full barrier 的翻相条件有两个:到达数凑够(这里 producer 一次 arrive),且 expect_tx 登记的字节数被 TMA 的 complete_tx 扣到零——所以 consumer 等到的一定是整块数据已写进 SMEM,漏了 expect_tx 这一步,barrier 会在数据到齐之前就翻相。整个 warp-group 128 线程,但只有 warp 0 的 lane 0 真正发指令,其余线程闲置。这看起来浪费,但因为 producer 不做计算(不占 ALU),实际硬件资源浪费很小——Hopper 引入的 setmaxnreg 让 producer 把寄存器配额还回去,给 consumer 用。
// Producer warp group 在开始时把寄存器配额降到最小
// .sync.aligned: warp-group 的 4 个 warp 必须都执行这同一条, 所以要放在只留单线程的 return 之前
asm("setmaxnreg.dec.sync.aligned.u32 24;\n");
setmaxnreg 只能在 sm_90a(以及 sm_100a 等带 a 后缀的目标)上用,只写 -arch=sm_90 会被 ptxas 拒绝;配额必须是 24~256 之间 8 的倍数。(nvcc 13.4 实测还有一层:-arch=sm_90a 直接编目标文件或可执行文件时,会顺带生成一份通用的 compute_90 PTX,里面的setmaxnreg照样被 ptxas 拒绝。解法是改写成 -gencode arch=compute_90a,code=sm_90a,或者像 CUTLASS 那样把这段 asm 包进 #if defined(__CUDA_ARCH_FEAT_SM90_ALL),见 cutlass-4.7.0/include/cutlass/arch/config.h:48;只出 -cubin 时没有这个问题。不同 CUDA 版本可能不同。)
FA3 源码把这两个配额写成了随 warp-group 数切换的编译期常量(flash-attn/hopper/flash_fwd_kernel_sm90.h:82):
static constexpr uint32_t LoadRegisterRequirement =
NumMmaWarpGroups == 1 ? 56 : (NumMmaWarpGroups == 2 ? (Use_TMA_KV ? 24 : 40) : 32);
static constexpr uint32_t MmaRegisterRequirement =
NumMmaWarpGroups == 1 ? 256 : (NumMmaWarpGroups == 2 ? (Use_TMA_KV ? 240 : 232) : 160);
注意它不是一对固定数字,而是三档:1 个 MMA warp-group 时是 56 / 256,2 个时是 24 / 240(KV 走 TMA)或 40 / 232(不走 TMA),3 个时是 32 / 160。取最常见的那一档——2 个 MMA warp-group + TMA 搬 KV——producer 降到 24,consumer 升到 240——每 SM 65536 个 32-bit 寄存器,384 线程平均只能分到 170 个左右(按 8 的分配粒度实际是 168,按本章结构写一个带 __launch_bounds__(384, 1) 的骨架 kernel,nvcc 13.4 以 -Xptxas -v 编译报的就是 168,不同 CUDA 版本可能不同),而 24×128 + 240×256 = 64512,正好装得下。正是靠这次再分配,consumer 才拿得到 240 个来放 O 累加器和各路 fragment。源码里调用它们的是 CUTLASS 的包装 cutlass::arch::warpgroup_reg_dealloc<>() / warpgroup_reg_alloc<>()(flash-attn/hopper/flash_fwd_kernel_sm90.h:309)。
17.5 Consumer Warp Group 的工作
__device__ void consumer_main(
half* sQ, half** sK, half** sV,
uint64_t* mbar,
/* O accumulator */ float* O_acc, float* row_max, float* row_sum
) {
// Consumer warp group 把寄存器配额提到 240
asm("setmaxnreg.inc.sync.aligned.u32 240;\n");
int n_k_tiles = N / Bc;
for (int k = 0; k < n_k_tiles; ++k) {
int stage = k % STAGES;
// 等 producer 拷完当前 stage (full barrier)
mbarrier_wait(&mbar[stage * 2 + 0], /* phase = ... */);
// ============ S = Q @ K^T (WGMMA) ============
float S_acc[MMAS_M * MMAS_N * 4] = {0};
wgmma_fence();
for (int kk = 0; kk < d; kk += 16) {
wgmma_mma_async_m64n64k16( // A、B 都从 SMEM 描述符取 (SS)
S_acc, sQ + kk_offset, sK[stage] + kk_offset,
/*scale_d=*/kk > 0 // 第一条覆盖 S_acc, 之后累加
);
}
wgmma_commit_group();
wgmma_wait_group(/*N=*/0); // 等 WGMMA 完成
// ============ Online softmax ============
// 与第 15 章一样, 但用 fragment level reduce
update_softmax_state(S_acc, row_max, row_sum, alpha);
scale_O_by_alpha(O_acc, alpha);
// ============ O += P @ V (WGMMA) ============
// 把 S_acc cast 为 fp16 的 P, 留在寄存器里直接当 A 操作数 (RS)
cast_S_to_P_fp16(S_acc, P_frag);
wgmma_fence(); // O_acc 刚被 rescale、P_frag 刚写过, 都是寄存器访问, 要 fence
wgmma_mma_async_m64n_d_k16_RS(
O_acc, P_frag, sV[stage], /*scale_d=*/1 // B (V) 仍必须来自 SMEM 描述符
);
wgmma_commit_group();
wgmma_wait_group(0);
// 通知 producer 这个 stage 已消费完
mbarrier_arrive(&mbar[stage * 2 + 1]);
}
// 最后归一化 O_acc /= row_sum, 写到 O HBM
finalize_and_write_O(O_acc, row_sum, ...);
}
几个关键点:
- WGMMA 是异步的:
wgmma_mma_async发完不阻塞,需要wgmma_commit_group+wgmma_wait_group显式同步。 - wgmma_fence 在每组 wgmma 之前调用,确保前面对累加器和 A 操作数寄存器的读写排在 wgmma 之前。它只管寄存器:如果操作数是线程刚用普通 store 写进 SMEM 的,还要另加
fence.proxy.async才能被 wgmma(异步代理)看见;TMA 写进来的 K/V 由 mbarrier 保证可见,不需要这一步。 - 操作数来源:WGMMA 的 B 必须在 SMEM(通过描述符给),A 可以在 SMEM 也可以在寄存器。FA3 在多数配置下(比如 headdim 128)把 P 留在寄存器里做 A,省掉一次写回 SMEM;少数配置(比如 headdim≤64 且非 causal)走 SS,P 先写进 SMEM。这一开关是
tile_size_fwd_sm90返回的MmaPV_is_RS(flash-attn/hopper/tile_size.h:10),FP8 强制走 RS(flash-attn/hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:83的static_assert)。 - mbarrier 同步:consumer 用
mbarrier_wait等 producer,用mbarrier_arrive通知 producer。
上面这版骨架每次都 wait_group(0) 等 WGMMA 做完才去算 softmax,Tensor Core 和做 exp 的 MUFU 仍是串行轮流忙。FA3 在 warp specialization 之上又做了两层重叠(论文 §3.1、§3.2):
- Ping-pong 调度(warp-group 之间):用
bar.sync强制 warp-group 1 的两次 GEMM 排在 warp-group 2 之前,于是一个 warp-group 做 softmax 时另一个在做 GEMM,然后角色互换。论文给的例子是 headdim 128、序列 8192 的 FP16 前向从 570 提到 620~640 TFLOPs。源码里是 mainloop 的warp_scheduler_barrier_sync()/warp_scheduler_barrier_arrive()(flash-attn/hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:915、:922),由UseSchedulerBarrier(同文件 :352)控制。 - Warp-group 内的 GEMM-softmax 流水:本轮先发下一块的
S_next = Q·K_{j}^T、再发O += P_cur·V_{j-1},都只 commit 不 wait,趁它们在 Tensor Core 上跑时算上一块的 softmax(论文 Algorithm 2)。代价是多一份 S 的寄存器。源码开关是IntraWGOverlap(tile_size_fwd_sm90返回的第 4 项)。
论文 §4.2 的消融(非 causal FP16,batch 4、序列 8448、16 头、headdim 128):完整 FA3 为 661 TFLOPs,去掉 GEMM-softmax 流水只有 582,去掉 warp specialization 只有 570——这两项对 FA3 的提速是实打实的。
17.6 Pipeline Depth 的选择
STAGES(pipeline 深度)是关键参数:
- STAGES=2:经典 double buffer。SMEM 占用最小,要求拷贝时延不超过计算时延才能完美重叠。
- STAGES=3/4:多一两份缓冲能容忍 producer/consumer 速度的不对齐,代价是 SMEM 线性增长。
更多 stage 意味着更多 SMEM 占用:
每 stage SMEM = 2 * (Bc * d * 2 byte) // K 和 V 各一份
Bc=64, d=64: 每 stage 16 KB
STAGES=2: 32 KB SMEM 仅 K/V 缓冲
+ Q tile、V 转置暂存、epilogue 复用区等,总 SMEM 数十 KB
(H100 每 SM 256 KB 统一 L1/SMEM,可配给 SMEM 的上限 228 KB)
值得注意的是:FA3 前向在 sm90 上取的其实是 2 stage,不是想当然的 4。flash-attn/hopper/flash_fwd_launch_template.h:47 写得很直白:
static constexpr int kStages = Arch >= 90 ? 2 : std::get<3>(kBlockMN_kNWarps_Stages_RS);
这与 Hopper 上 block tile 本来就大有关(flash-attn/hopper/tile_size.h:25:headdim≤64、FP16/BF16、非 causal / 非 local / 非 paged 的那一档返回的就是 {192, 192};一旦是 causal 或 local,use_blockN_128 置真、N 方向收到 128),同文件的注释里反复出现「hits the limit of smem」(如 headdim 128、192 两档),可见 SMEM 已被大 tile 用满;两级流水配合大 tile 来掩盖 TMA 延迟,再加 stage 就得缩小 tile。"stage 越深越好"是错的——stage 数要和 tile 大小一起放进同一个 SMEM 预算里权衡。
17.7 FP8 的特殊处理
FA3 的另一个关键创新是支持 FP8 GEMM:
wgmma_mma_async_e4m3_e4m3_f32_m64n64k32(
accumulator,
fp8_a_smem, fp8_b_smem,
/*scale_d=*/0
);
FP8 WGMMA 的 K 维一次性算 32(FP16 是 16)——单条指令算力加倍。代价是 e4m3 只有 3 位尾数,动态范围靠 scale 撑。
这里要澄清一个常见误解:FA3 的 FP8 路径上 Q/K/V 是同一种元素类型,Q 并不保持 FP16。源码里的判据就是元素类型本身(flash-attn/hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:44):
static constexpr bool Is_FP8 =
cute::is_same_v<Element, cutlass::float_e4m3_t> ||
cute::is_same_v<Element, cutlass::float_e5m2_t>;
WGMMA 也不支持 fp16×fp8 的混合输入。另外 FP8 WGMMA 只接受 K-major 的 SMEM 操作数(PTX 的转置位只对 f16/bf16 有效),而第二个 GEMM 里 V 按行主序存放恰好不是 K-major——FA3 的做法是 producer warp-group 把 V 的 tile 载入 SMEM 后在 kernel 内转置(论文 §3.3;源码 Transpose_V = Is_FP8 && !V_colmajor,flash-attn/hopper/mainloop_fwd_sm90_tma_gmma_ws.hpp:55),这就是 17.2 节里 producer 其余 3 个 warp 要留下来的原因。FA3 保精度靠的是论文里的两招:
- Block quantization:scale 按块给(Q 按 B_r×d、K/V 按 B_c×d 一块一个),而不是全张量一个 scale,避免异常值把整个 scale 拉坏。
- Incoherent processing:量化前给 Q、K 同乘一个随机正交矩阵 M(论文取 ±1 随机对角阵与 Hadamard 矩阵的乘积,可 O(d log d) 计算),因为 (QM)(KM)^T = QK^T,结果不变,异常值却被"摊平"到各维度上。
这两招是论文的做法,量化都在 attention 之前完成(论文建议与 rotary embedding 融合)。pin 的源码里 kernel 接收的 q_descale / k_descale / v_descale 形状是 (batch, num_heads_k)(flash-attn/hopper/flash_api.cpp:694-:696),flash-attn 整个仓库里也没有 Hadamard 变换的实现——这两步都要调用方在量化时自己做。
此外 FP8 的 fragment 排布和 FP16 不同,源码里能看到一串 permute_Aregs_fp8(flash-attn/hopper/utils.h:455)/ permute_Cregs_fp8(:489)/ permute_output_fp8(:509)专门做寄存器重排。
这些细节让 FA3 的 FP8 路径达到接近 1.2 PFLOPs——H100 FP8 稠密峰值 1979 TFLOPs,约合 60% 利用率。
17.8 性能跃迁:论文口径的数字
本专栏没有 H100 可跑,这一节的数字全部来自 FA3 论文,H100 SXM、长序列,除 FP8 一行外都是 FP16:
| 实现 | TFLOPs | % Tensor Core peak(989 稠密 / FP8 1979) |
|---|---|---|
| FA2 (cp.async + mma.sync) | ~350 | ~35% |
| FA3 (TMA + WGMMA + Warp Spec) | ~740 | 75% |
| FA3 FP8 | 接近 1.2 PFLOPs | ~60% of FP8 peak |
来源:Shah et al., FlashAttention-3, 2024(arXiv:2407.08608)摘要与 §4.1 的 H100 FP16/FP8 基准。论文对 FA2 的定性是"只用到 H100 理论算力的 35%"。
FA2 → FA3 约 2×——与 FA3 论文摘要给的 1.5–2.0× 加速一致,而且 attention 的数学没有变(仍是精确 attention):提速来自对 Hopper 异步硬件的利用,以及围绕它重排的调度(warp specialization、ping-pong、warp-group 内 GEMM-softmax 流水,见 17.5 节末的消融)。
17.9 Hopper → Blackwell 迁移
Blackwell(B200)相比 Hopper 的关键变化:
- 第 5 代 Tensor Core:增加 FP4 支持,FP4 算力是 FP8 的 2×。
- 第二代 TMA:PTX 在 sm_100 上给
cp.async.bulk.tensor新增了.tile::gather4(一次拼 4 行不连续的数据)、im2col::w等模式和服务 CTA pair 的.cta_group修饰。 - CTA Pair:同一 cluster 里
%cluster_ctarank只差最低位的两个 CTA 组成一对,多数 tcgen05 操作可以按 CTA pair 粒度执行,一次访问两个 CTA 各自的 Tensor Memory(PTX ISA「CTA Pair」一节)。
迁移策略:
- TMA 描述符接口几乎一样:换枚举值即可。
- WGMMA 让位给 Blackwell 的 tcgen05 系列 MMA:仍是异步 + 描述符驱动,但累加器搬到了专门的 Tensor Memory,形状与同步语义都要重学。
- Warp Specialization 框架不变:producer/consumer 模式继续用。
- CTA Pair 是新东西:可以让 attention kernel 进一步增大有效 tile。
手写 kernel 的迁移量取决于对 mainloop 的耦合程度,通常不小。好在 CUTLASS 4.x 已经把 Blackwell(sm100)适配做进去了——cutlass-4.7.0/include/cutlass/arch/mma_sm100.h 与 sm100_* 的一整套 collective/tile scheduler 都在——用 CUTLASS 写 GEMM/Attention 能省掉大部分重写(Blackwell attention 的参考实现在 cutlass-4.7.0/examples/77_blackwell_fmha/)。
17.10 这一章的小结与下一章
第 17 章是本专栏技术深度的高峰:
- TMA 替代 cp.async:单线程发起、专用硬件、不占 ALU。
- WGMMA 替代 mma.sync:一条指令由整个 warp-group 发出,形状 m64nNk16(FP16,N 为 8~256 中 8 的倍数,最大 64×256×16),异步执行;B 必在 SMEM,A 可在 SMEM 或寄存器。
- Warp Specialization:1 个 producer warp-group + 1~3 个 consumer warp-group(典型 2 个,整块 384 线程),物理硬件并行。
- mbarrier 同步:producer/consumer 之间用 phase 切换的同步机制。
- setmaxnreg:动态调整 warp 寄存器配额,让 consumer 拿到更多寄存器。
- 流水深度要和 tile 大小一起权衡:FA3 前向在 sm90 上取的是 2 stage,靠大 tile 而不是深流水掩盖延迟。
- FA3 在 H100 上做到 ~740 TFLOPs(FP16,75% 利用率)/ 接近 1.2 PFLOPs(FP8)——相对 FA2 的 ~35% 利用率约 2×。
第 18 章我们回到一个更广的话题——Persistent Kernel。Persistent kernel 是另一种"永远活着"的 kernel 模式:grid_size 固定为 SM 数,每个 block 通过 grid-stride loop 处理多个 tile。这种模式对小 tile 工作负载(比如 LLM 推理 decoding 阶段的小 batch)特别有效。读完第 18 章,第四篇结束,读者就完成了从基础 kernel 到 SOTA FA2 的完整训练。
本章动手练习:
- 构建一个 TMA descriptor,发起一次 TMA 拷贝,观察 SMEM 中的 swizzle 布局。
- 实现一个最简化的 Producer/Consumer kernel(单 K tile,纯 GEMM),熟悉 mbarrier 同步。
- 阅读 FA3 官方实现
flash-attn/hopper/flash_fwd_kernel_sm90.h,对照本章描述的概念找代码位置(提示:角色划分 :74、producer 单 warp :319、寄存器再分配 :309)。