CUDA 算子工程:手写 FlashAttention v2 之路

第 17 章 TMA + Warp Specialization 把 FA2 写到 SOTA

作者 杨艺韬 · 4,700 字 · 发布于 · 更新于

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 用的也是这个口径。

为什么差距这么大?三个原因:

  1. TMA 比 cp.async 更高效:单线程发起、专用硬件、原生 swizzle、不占 SIMT 算术单元。
  2. WGMMA 是异步指令:发完不阻塞,warp 可以继续算/拷下一份。
  3. 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, ...);
}

几个关键点:

  1. WGMMA 是异步的:wgmma_mma_async 发完不阻塞,需要 wgmma_commit_group + wgmma_wait_group 显式同步。
  2. wgmma_fence 在每组 wgmma 之前调用,确保前面对累加器和 A 操作数寄存器的读写排在 wgmma 之前。它只管寄存器:如果操作数是线程刚用普通 store 写进 SMEM 的,还要另加 fence.proxy.async 才能被 wgmma(异步代理)看见;TMA 写进来的 K/V 由 mbarrier 保证可见,不需要这一步。
  3. 操作数来源: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)。
  4. 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 保精度靠的是论文里的两招:

  1. Block quantization:scale 按块给(Q 按 B_r×d、K/V 按 B_c×d 一块一个),而不是全张量一个 scale,避免异常值把整个 scale 拉坏。
  2. 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 的关键变化:

  1. 第 5 代 Tensor Core:增加 FP4 支持,FP4 算力是 FP8 的 2×。
  2. 第二代 TMA:PTX 在 sm_100 上给 cp.async.bulk.tensor 新增了 .tile::gather4(一次拼 4 行不连续的数据)、im2col::w 等模式和服务 CTA pair 的 .cta_group 修饰。
  3. 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 章是本专栏技术深度的高峰:

  1. TMA 替代 cp.async:单线程发起、专用硬件、不占 ALU。
  2. WGMMA 替代 mma.sync:一条指令由整个 warp-group 发出,形状 m64nNk16(FP16,N 为 8~256 中 8 的倍数,最大 64×256×16),异步执行;B 必在 SMEM,A 可在 SMEM 或寄存器。
  3. Warp Specialization:1 个 producer warp-group + 1~3 个 consumer warp-group(典型 2 个,整块 384 线程),物理硬件并行。
  4. mbarrier 同步:producer/consumer 之间用 phase 切换的同步机制。
  5. setmaxnreg:动态调整 warp 寄存器配额,让 consumer 拿到更多寄存器。
  6. 流水深度要和 tile 大小一起权衡:FA3 前向在 sm90 上取的是 2 stage,靠大 tile 而不是深流水掩盖延迟。
  7. 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 的完整训练。

本章动手练习:

  1. 构建一个 TMA descriptor,发起一次 TMA 拷贝,观察 SMEM 中的 swizzle 布局。
  2. 实现一个最简化的 Producer/Consumer kernel(单 K tile,纯 GEMM),熟悉 mbarrier 同步。
  3. 阅读 FA3 官方实现 flash-attn/hopper/flash_fwd_kernel_sm90.h,对照本章描述的概念找代码位置(提示:角色划分 :74、producer 单 warp :319、寄存器再分配 :309)。