H200// wgmma

H200 · HOPPER GH100 · 特性专题

wgmma · 异步 Tensor Core (mma_async)

★ Hopper 计算主线 · warp group 级异步 MMA,MMA/搬运/softmax 真重叠

01

是什么

DEFINITION

wgmma.mma_async 是 Hopper 的异步 warp-group 级矩阵乘累加指令——一个 warp group(4 warp / 128 线程)协同发起一次 D = A·B + C,异步发射、完成需显式等待。这是 Hopper Tensor Core 的第四代核心指令。

02

为什么需要(vs A100 mma.sync m16n8)

RATIONALE · VS A100 MMA.SYNC
维度A100 mma.syncHopper wgmma
发起粒度1 warp(32 线程)1 warp group(128 线程)
tile Mm16m64(4× 宽)
同步模型同步——发射即阻塞异步——async proxy 执行,commit/wait 控制
可重叠性差(warp 被占住)发射后 warp 可做 epilogue/softmax/加载
操作数来源寄存器B 恒 smem descriptor;A 可寄存器或 smem

异步性是关键:它让 MMA 与 softmax / 数据搬运真正重叠(FlashAttention-3 的核心收益),CUTLASS 的 producer/consumer 双缓冲流水线也建立在 wgmma 异步之上。配合 TMA 喂数据,搬运与计算才第一次能完全并行。

03

关键数字

KEY NUMBERS · SHAPES
CAUTION · 最易记错的两个数

M 恒为 64;N ∈ {8,16,32,64,96,128,192,256}(注意不是"8 的任意倍数",CUTLASS 只为这 8 个 N 各定义一个 atom)。K 由 dtype 决定——下表的 K 值是踩坑高发区:

dtype(A·B)K累加器 D/C
FP16 / BF1616f16 或 f32最常用 GEMM/Attention
TF328f32非 16,常被错记
FP8 (e4m3/e5m2)32f16 或 f32每寄存器打包 4 元素
INT832s32同 FP8 的 K
FP6416f64非 8,HPC 场景常被错记

异步控制三件套

指令作用
wgmma.fence.sync.aligned前序对 RMEM/smem 的写对后续 wgmma 可见
wgmma.commit_group.sync.aligned提交一组已发射的 wgmma
wgmma.wait_group.sync.aligned N等待至多 N 组未完成(N∈[0,7],0=全等完)

smem 操作数 = 64-bit descriptor(base+stride+swizzle);swizzle 四选一:no-swizzle(16B) / 32B / 64B / 128B。若 smem 由普通 st.shared 写入(非 TMA),还需 fence.proxy.async

04

可编译示例

COMPILABLE EXAMPLE · VERBATIM

逐字摘自 CUTLASS include/cute/arch/mma_sm90_gmma.hpp

① 异步控制三件套

asm volatile("wgmma.fence.sync.aligned;\n"        ::: "memory");  // warpgroup_arrive
asm volatile("wgmma.commit_group.sync.aligned;\n"  ::: "memory");  // warpgroup_commit_batch
asm volatile("wgmma.wait_group.sync.aligned %0;\n" :: "n"(0) : "memory"); // 全等完

② wgmma 本体:64×64×16 FP16 SS(A、B 均来自 smem descriptor)

// d00..d15 = 该线程持有的 64×64 输出分片(RMEM);desc_a/desc_b 为 64-bit smem descriptor
asm volatile(
"{\n"
  ".reg .pred p;\n"
  "setp.ne.b32 p, %18, 0;\n"                    // scale_D==0 → 不累加(不读 C)
  "wgmma.mma_async.sync.aligned.m64n64k16.f16.f16.f16 "
  "{%0,%1,%2,%3,%4,%5,%6,%7,%8,%9,%10,%11,%12,%13,%14,%15},"  // 16 个累加器 reg
  " %16,"   // desc_a  ("l" 64-bit)
  " %17,"   // desc_b  ("l" 64-bit)
  " p, %19, %20, %21, %22;\n"                    // p=scaleD; scaleA/B,tnspA/B 立即数
"}\n"
  : "+r"(d00),"+r"(d01),"+r"(d02),"+r"(d03),"+r"(d04),"+r"(d05),"+r"(d06),"+r"(d07),
    "+r"(d08),"+r"(d09),"+r"(d10),"+r"(d11),"+r"(d12),"+r"(d13),"+r"(d14),"+r"(d15)
  : "l"(desc_a), "l"(desc_b),
    "r"(int32_t(scale_D)), "n"(int32_t(scaleA)), "n"(int32_t(scaleB)),
    "n"(int32_t(tnspA)), "n"(int32_t(tnspB)));

其它 dtype 助记符:.m64nNk16.f32.f16.f16(FP32累加)/ .m64nNk8.tf32.tf32 / .m64nNk32.f32.e4m3.e4m3 / .m64nNk32.s32.s8.s8。N 取 {8,16,32,64,96,128,192,256}。

05

注意事项 / 陷阱

PITFALLS · CHECKLIST
CAUTION · 陷阱清单
  • 128 线程必须全部到达sync.aligned 是集体指令,warpgroup 内所有 128 线程必须执行同一条 wgmma(统一谓词),否则 UB。
  • A 在寄存器时必须 K-major(CUTLASS static_assert);非 16-bit dtype 的 smem 操作数也只能 K-major。
  • smem 布局必须匹配 descriptor+swizzle——用 CUTLASS GMMA::Layout 构造,别手撸 stride
  • 同步顺序不能省:读累加器前 commit_group+wait_group<0>;wgmma 前对 RMEM/smem 有写要 wgmma.fence;普通 st.shared 写的 smem 还要 fence.proxy.async
  • 寄存器压力:每线程持有一片 64×N 输出,N=256 时累加器吃掉大量寄存器——大 tile 常拆 2 个 warpgroup 沿 M 分摊,并主动降占用度换 ILP。
06

数据源

SOURCES