H200 · HOPPER GH100 · 特性专题
wgmma · 异步 Tensor Core (mma_async)
★ Hopper 计算主线 · warp group 级异步 MMA,MMA/搬运/softmax 真重叠
是什么
DEFINITIONwgmma.mma_async 是 Hopper 的异步 warp-group 级矩阵乘累加指令——一个 warp group(4 warp / 128 线程)协同发起一次 D = A·B + C,异步发射、完成需显式等待。这是 Hopper Tensor Core 的第四代核心指令。
为什么需要(vs A100 mma.sync m16n8)
RATIONALE · VS A100 MMA.SYNC| 维度 | A100 mma.sync | Hopper wgmma |
|---|---|---|
| 发起粒度 | 1 warp(32 线程) | 1 warp group(128 线程) |
| tile M | m16 | m64(4× 宽) |
| 同步模型 | 同步——发射即阻塞 | 异步——async proxy 执行,commit/wait 控制 |
| 可重叠性 | 差(warp 被占住) | 发射后 warp 可做 epilogue/softmax/加载 |
| 操作数来源 | 寄存器 | B 恒 smem descriptor;A 可寄存器或 smem |
异步性是关键:它让 MMA 与 softmax / 数据搬运真正重叠(FlashAttention-3 的核心收益),CUTLASS 的 producer/consumer 双缓冲流水线也建立在 wgmma 异步之上。配合 TMA 喂数据,搬运与计算才第一次能完全并行。
关键数字
KEY NUMBERS · SHAPESM 恒为 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 / BF16 | 16 | f16 或 f32 | 最常用 GEMM/Attention |
| TF32 | 8 | f32 | 非 16,常被错记 |
| FP8 (e4m3/e5m2) | 32 | f16 或 f32 | 每寄存器打包 4 元素 |
| INT8 | 32 | s32 | 同 FP8 的 K |
| FP64 | 16 | f64 | 非 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。
可编译示例
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}。
注意事项 / 陷阱
PITFALLS · CHECKLIST- 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。
数据源
SOURCES- CUTLASS mma_sm90_gmma.hpp(wgmma inline-asm + 全部 shape,逐字)
- PTX ISA — Asynchronous Warpgroup MMA(语义、六步同步、descriptor 位域、shape 表)
- Colfax — CUTLASS Tutorial: WGMMA on Hopper(warpgroup=128线程、swizzle、fence/commit/wait 动机)