H100 · WGMMA 入门(Warp Group Matrix Multiply Accumulate)

目录 · ← l5 · l7 →

H100 · WGMMA 入门(Warp Group Matrix Multiply Accumulate)

本文档基于课程讲义《6. WGMMA-1.pdf》(Lesson 6,51 页)整理, 系统介绍 Hopper 的 WGMMA(Warp Group MMA):第四代 Tensor Core 的异步矩阵乘加指令。 前置阅读:H100-架构介绍.md(Tensor Core)、H100-异步与屏障.md(异步)、 H100-cuTensorMap.md(swizzle)、H100-内联PTX 相关知识。


目录

  1. Hopper MMA 的四大创新
  2. 范式转变:warp → warp group
  3. WGMMA 流水线
  4. wgmma.fence.sync.aligned
  5. wgmma.mma_async 指令
  6. 形状 m64nXkY:M 为何固定为 64
  7. 操作数位置与 scale 参数
  8. A 的位置:寄存器 vs 共享内存
  9. ldmatrix:把 A 装进寄存器
  10. 16×16 tile 的装载:Split-Warp 策略
  11. Swizzle 地址计算
  12. BF16 打包(Packing)
  13. WGMMA 描述符(Descriptor)
  14. wgmma.mma_async 的结果写回
  15. 总结与学习衔接

1. Hopper MMA 的四大创新

除了更快的 Tensor Core,Hopper 的 MMA 还有四个重要创新:

  1. 完全异步:Tensor Core 运算现在是非阻塞的——支持多个在飞(in-flight)操作, 让 Tensor Core 活跃更久。
  2. warp → warp group:从单 warp 转到 warp group,可使用大得多的 tile
  3. 直接从 shared memory / 寄存器异步取数:Ampere 里线程要等数据进寄存器才能发 mma.sync; 而 WGMMA 可以跳过加载寄存器,直接发起数学运算
  4. FP8 与稀疏的专用硬件支持

2. 范式转变:warp → warp group

  • Hopper 把延续十年的”warp 是执行单元“改成”warp group 是执行单元“。
  • wgmma 时,你把 4 个 warp 融合成一个在 Tensor Core 上运算的计算实体—— 最大化 Tensor Core 利用率,省去流水化多条 mma.sync 的麻烦。
  • 4 个 warp 一起发起 wgmma 指令;warp 调度器检查所有 warp 是否在该指令(同一 PC)汇聚, 然后调度器把它们融合、把数学命令派发给 Tensor Core。
  • 命令一旦交给 Tensor Core,线程就继续执行下一条指令(异步)。

3. WGMMA 流水线

1. 异步加载 tile 数据(必要时预加载下一个 buffer)
     - A 装寄存器 → 用 ldmatrix + 一堆复杂寻址;B 用 wgmma 描述符加载
     - 或(A/B 都放 shared)→ 只为 A、B 各建一个描述符
2. 发起 wgmma 运算
3. 从寄存器收集结果 → 用 wgmma.fence.sync.align 让寄存器写可见 → 复杂地址计算 → 搬到 shared memory

4. wgmma.fence.sync.aligned

  • wgmma.fencewgmma.mma_async 的”寄存器排序屏障”,不是”完成屏障”
  • 作用:在后续 wgmma.mma_async 复用同样寄存器之前,让先前的 warpgroup 寄存器访问以正确顺序可见
  • 两个主要场景需要它:
    1. warpgroup 中第一次 wgmma.mma_async 之前;
    2. 任何时候某线程访问过”稍后的 wgmma.mma_async 将复用为累加器或 A 片段输入”的寄存器。
  • 它不负责排序 wgmma.mma_async 消费的矩阵描述符/数据的 shared-memory 写; 那需要 async proxy fence(fence.proxy.async 来排序”先前的 shared 写”与”后续 wgmma 读”。

5. wgmma.mma_async 指令

核心异步指令,让 Tensor Core 做矩阵乘加:

  • A 直接从 shared memory / 寄存器读B 直接从 shared memory 读C、D 在寄存器里
wgmma.mma_async.sync.aligned.<shape>.<dtypeD>.<dtypeA>.<dtypeB>
    d, a-desc, b-desc, scale-d, imm-scale-a, imm-scale-b, imm-trans-a, imm-trans-b;

5.1 .sync.aligned

  • .sync:SM 级屏障,要求 warpgroup 内所有参与线程(通常 128 线程 = 4 warp) 都到达该指令后,任何一个才能继续。
  • .aligned:断言 warpgroup 内所有线程在该指令地址(PC)处已汇聚(lockstep), warp group 之间没有线程分歧。

5.2 指令限定符

  • Shape:tile 形状,用 m64, n{}, k{} 表示。
  • dtypeD / dtypeA / dtypeB:D、A、B 张量的数据类型。

6. 形状 m64nXkY:M 为何固定为 64

对所有 wgmma 指令:

  • M 维固定为 64
  • K 维(内积维)严格由 A、B 的精度决定
  • N 维最灵活:决定处理 Matrix B(和 C)的多少列。
    • N 必须是底层内存分配块大小的倍数(通常为 8 或 16 的倍数,视类型而定)。
    • 合法 N 值范围:8 到 256

6.1 固定 M 的原因

  • 硬件设计为每线程 4 个寄存器 → 每个 warpgroup 可持 128 × 4 = 512 个寄存器
  • 若用最小 tile m=64, n=8,则 64 × 8 = 512 个元素,恰好等于一个 warpgroup 的容量
  • 硬件把 A 的数据物理上静态映射到 Tensor Core 输入 B;M=64 时有完美的静态映射,运行时零决策逻辑
  • 若 M 可配置,就需要一个巨大的 Crossbar Switch(复杂多路复用器)在寄存器与数学单元之间 按请求动态重路由——昂贵

7. 操作数位置与 scale 参数

7.1 操作数位置

矩阵位置
A寄存器 shared memory
B只能 shared memory
C/D只能 寄存器

7.2 scale_d

  • 指定是累加D = A × B + C)还是覆盖D = A × B):
    • scale_d = 1 → 累加;
    • scale_d = 0 → 覆盖。

7.3 scale_a / scale_b

  • A / B 元素的缩放因子:+1 或 −1

7.4 寄存器操作数 d

  • 输出存放的寄存器;在 PTX 指令串里是第一个操作数,必须声明为花括号 {} 包裹的向量(元组)
  • 输出操作数个数 = 输出总元素数 ÷ 总线程数

8. A 的位置:寄存器 vs 共享内存

A 的 tile 可放 shared memory 或寄存器,因此告诉 wgmma A 在哪有两种方式:

  • 寄存器:把寄存器直接传给 wgmma 指令。
  • 共享内存:传 wgmma 描述符

8.1 A 放寄存器(何时用)

  • 复用数据时是好选择:从 shared 同时读 A 和 B 会加倍 bank 压力、多次写 shared 代价高,不如从寄存器读。
  • 从 shared 搬数贵,从寄存器搬数便宜——把可复用数据放在能更快访问的地方。
  • 寄存器数据来自 shared memory:ldmatrix 取未 swizzle 的地址,把数据装进寄存器。
  • 不复用数据时,把 A 放寄存器没意义(有指令开销 + 寄存器压力)。

9. ldmatrix:把 A 装进寄存器

ldmatrix 是 WGMMA 流水线里把 Matrix A 数据装进寄存器的首要机制。

  • 与标准 ld.shared(加载线性数据)不同,ldmatrix不透明的寄存器模式加载数据, 物理上与 Tensor Core 的输入 lane 对齐。
  • 提供指针的线程,有时不是最终拿到数据的线程
  • .m8n8:从 shared memory 装进 warp 寄存器的矩阵 tile 的几何形状。
  • .x1 / .x2 / .x4向量化宽度,以及每条指令加载的矩阵片段数。

9.1 .x{1,2,4} 的含义

  • 指一次能把多少个 8×8 核心矩阵搬进寄存器:
    • .x4:每线程 4 个寄存器 = 4 个 8×8 核心矩阵。最常用——搬 16×16×4 = 1024 字节打满寄存器带宽
    • .x2:每线程 2 个寄存器 = 512 字节,用于较小 tile。
    • .x1:每线程 1 个寄存器,主要用于边界处理等。

9.2 .sync.aligned.trans

  • .sync:迷你屏障——硬件要保证 T16 准备好让 T0 请求的数据覆盖其寄存器; 若 T16 还在忙别的,硬件写者会破坏 T16 的状态。
  • .aligned:所有线程必须一起执行它,硬件才知道 warp 已就绪。
  • .trans:加载时是否转置矩阵。因为线程的寄存器对他人不可见,这整条指令由硬件内部完成。

10. 16×16 tile 的装载:Split-Warp 策略

整个操作分两个阶段:

  • Address Phase(谁提供指针?= Gather)
  • Register Phase(谁持有结果?= Destination)

.x4 下我们处理的是 4 个 8×8 子块

10.1 Split-Warp(拆 warp)

  • 地址不是线性算的。因为 ldmatrix8 列块加载,把 warp 拆成两条”竖直条带”:
    • 左半(列 0–7):由前 16 个线程(lane 0–15)控制。
    • 右半(列 8–15):由后 16 个线程(lane 16–31)控制。

10.2 Phase 1:地址责任(谁指)

  • 硬件看每个线程的地址寄存器,决定从 shared memory 哪里读。warp 被拆成 4 组 × 8 线程

10.3 Phase 2:寄存器责任(谁持有)

  • 数据取到后,ldmatrix 把它按行条带化分配到 warp 里(每组 4 线程),准备好喂 Tensor Core。
  • 对任意 8×8 矩阵(M0/M1/M2/M3),行分布:
    • 第 0 行 → 线程 0–3;第 1 行 → 线程 4–7;第 2 行 → 线程 8–11;……;第 7 行 → 线程 28–31。

11. Swizzle 地址计算

tile 在 shared memory 里是 swizzled 布局(为避 bank conflict),所以”逻辑行”的字节在物理上不连续。

  • 每个线程从逻辑坐标出发(”我要第 r 行、第 c 列块”),用 swizzle 映射(常为 XOR 重映射)算物理 shared 地址:
    smem_addr = base + (row_offset ^ swizzle_mask) + col_offset
    
  • 因为指针经 swizzle 映射,访问落到不同 bank,warp 命中极少/零 bank conflict
  • ldmatrix 每线程加载 16 字节,并重排成分片布局。

11.1 地址位拆解

  • 位 [0–1]:4 字节字内的字节偏移(与 bank 选择无关)。
  • 位 [2–6]Bank Index——这 5 位决定数据落到 32 个 bank(0–31)中的哪一个。
    • 其中 位 [4–6] 控制 bank index 的高 3 位
  • 位 [7–9]Row Index——因 pitch 是 128 字节(2⁷),位 7 是换行时第一个变化的位
  • 目标:让 Bank(位 4–6)随 Row(位 7–9)变化

11.2 Swizzle 地址公式

Physical_Address = (Base_Address + Linear_Offset) ^ Swizzle_Mask;

// 128 字节 = 2^7,取位 [7,8,9] 移到位置 [4,5,6]
uint32_t mask = ((linear_offset >> 7) & 0x7) << 4;
uint32_t smem_ptr = (base_ptr + linear_offset) ^ mask;

12. BF16 打包(Packing)

  • NVIDIA GPU 没有 16 位寄存器,只有 32 位寄存器。所以 bf16 数据在寄存器里必须成对打包
  • 容器:一个 .b32(32 位)寄存器装两个 bf16。
  • 布局
    • 位 0–15(LO):元素 N(偶数下标);
    • 位 16–31(HI):元素 N+1(奇数下标)。
  • 搬数据时用 .b32 类型指令;计算时用 .bf16x2

12.1 ldmatrix 的打包

  • ldmatrix.sync.aligned.m8n8.x1.b16:每线程收 1 个 32 位寄存器(含一对打包 bf16)。
  • .x2 / .x4:得 2 或 4 个寄存器,每个含一对打包 bf16。
  • shared memory 地址须 16 字节(128 位)对齐
  • .trans:加载时把行主序数据直接转置成列主序(或反之),无需手动 shuffle。

13. WGMMA 描述符(Descriptor)

64 位描述符,包含数据存储信息、步长与 swizzle 布局。它打包 5 个字段: 地址、LBO、SBO、Matrix base offset、swizzle 布局。

13.1 蓝图

  • 一个 64 位寄存器编码:基地址、K 步长、M/N 步长、矩阵基偏移、swizzle 模式。
  • 所有步长都预先除以 16,以塞进 14 位字段(低 4 位没用,因为 shared 指针 16 字节对齐); 硬件运行时再乘 16,保证 16 字节对齐、降低地址延迟。
  • 基地址必须 16 字节对齐(否则预处理步长无法被硬件正确解释)。

13.2 各字段

字段含义
Base address前 14 位要加载的 tile 的 shared memory 地址(从地址取 14 位)
LBO(Leading Byte Offset)16–29沿”leading”方向相邻两个 core-matrix 列的字节距离;编码 K 步长(÷16、掩 14 位、左移 16)
SBO(Stride Byte Offset)32–45另一个方向跳”一个 core-matrix 块”的 SMEM 字节距离
Matrix base offsetswizzle 模式每 128B 重复;起始不在边界时告诉硬件从哪个 128B 块开始
Swizzle mode61–6300=无 swizzle,01=128B,10=64B,11=32B

13.3 LBO / SBO 细节

  • LBO
    • K-major swizzled 布局不用 LBO(硬件假设为 1)。
    • MN-major 时,LBO = 从”前 (swizzle-byte-size/16) 行”到”下一 (swizzle-byte-size/16) 行”的偏移; 对 128B swizzle,128/16 = 8,即”跳 8 行”。
  • SBO
    • K-major:从”前 8 行”到”下一 8 行”的偏移;
    • MN-major:从”前 8 列”到”下一 8 列”的偏移。

13.4 设置 swizzle 的重要性

  • 若 shared memory 数据被 swizzle,必须正确设置 swizzle 布局,否则 wgmma 单元会把打散的数据线性地读错

14. wgmma.mma_async 的结果写回

  • wgmma.mma_async 异步把数据从 shared memory / 寄存器搬到 Tensor Core 做运算。
    • consumer warp group 的所有线程都调用它。
  • Tensor Core 开始算,结果按 warpId、laneId 写回 consumer warpgroup 的寄存器——开销低,硬件把数据”倒”到最近的存储。
  • 硬件想把结果写给物理上最靠近负责该块的数学单元的线程,所以线程被交错成小块, 每个线程在内存里持有的是非连续数据。

14.1 每个线程拿到多少寄存器

  • 例:wgmmam64n256k16 → 每线程累加器有 128 个元素
  • 每一条 WGMMA 指令由 warpgroup 里每个线程都发起!
  • 对 128 线程的 warpgroup:16 位累加器需 16 条 stmatrix(.x4),32 位累加器需 32 条 stmatrix
  • 提示:像 FlashAttention 3 这类 kernel,若用寄存器装 A tile,要小心 softmax-GEMM 流水线带来的寄存器压力

15. 总结与学习衔接

15.1 核心脉络速记

概念一句话
四大创新全异步、warp→warp group、直接 shared/reg 取数、FP8+稀疏
warp group4 warp = 128 线程 = 一个计算实体,M 固定 64
操作数位置A 寄存器/shared,B 只能 shared,C/D 只能寄存器
ldmatrix把 A 从 shared 装进寄存器,.x4 打满带宽,split-warp + swizzle 防冲突
描述符64 位,含 base/LBO/SBO/base-offset/swizzle,步长÷16 塞 14 位
swizzleaddr = (base + linear) ^ mask,bank 随 row 变
packingbf16 成对打包进 .b32(LO=偶,HI=奇)
写回异步、按 warpId/laneId 写寄存器、线程交错、stmatrix 搬出

15.2 与本课程其他内容的衔接

本文概念对应后续专题
WGMMA 高级用法 / 寄存器布局《7. Wgmma part 2》
cp.async.bulk 供数 + wgmma 计算《8. Kernel Design》(GEMM 软件流水)
Stream-K / FlashAttention《8.1 Stream-K》《8.2 Kernel Launch》

15.3 一句话记忆

WGMMA = 4 个 warp 拧成一个 128 线程的”计算实体”,用一条异步指令让 Tensor Core 直接 从 shared/寄存器取 A、从 shared 取 B、把 C 累加在寄存器里; ldmatrix 负责把 A 装进寄存器(swizzle 防 bank conflict),64 位描述符告诉硬件 A/B 在 shared 里怎么摆。


参考来源:6. WGMMA-1.pdf(Lesson 6,51 页)。