H100 · WGMMA 入门(Warp Group Matrix Multiply Accumulate)
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相关知识。
目录
- Hopper MMA 的四大创新
- 范式转变:warp → warp group
- WGMMA 流水线
- wgmma.fence.sync.aligned
- wgmma.mma_async 指令
- 形状 m64nXkY:M 为何固定为 64
- 操作数位置与 scale 参数
- A 的位置:寄存器 vs 共享内存
- ldmatrix:把 A 装进寄存器
- 16×16 tile 的装载:Split-Warp 策略
- Swizzle 地址计算
- BF16 打包(Packing)
- WGMMA 描述符(Descriptor)
- wgmma.mma_async 的结果写回
- 总结与学习衔接
1. Hopper MMA 的四大创新
除了更快的 Tensor Core,Hopper 的 MMA 还有四个重要创新:
- 完全异步:Tensor Core 运算现在是非阻塞的——支持多个在飞(in-flight)操作, 让 Tensor Core 活跃更久。
- warp → warp group:从单 warp 转到 warp group,可使用大得多的 tile。
- 直接从 shared memory / 寄存器异步取数:Ampere 里线程要等数据进寄存器才能发
mma.sync; 而 WGMMA 可以跳过加载寄存器,直接发起数学运算。 - 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.fence是wgmma.mma_async的”寄存器排序屏障”,不是”完成屏障”。- 作用:在后续
wgmma.mma_async复用同样寄存器之前,让先前的 warpgroup 寄存器访问以正确顺序可见。 - 两个主要场景需要它:
- warpgroup 中第一次
wgmma.mma_async之前; - 任何时候某线程访问过”稍后的
wgmma.mma_async将复用为累加器或 A 片段输入”的寄存器。
- warpgroup 中第一次
- 它不负责排序
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)
- 地址不是线性算的。因为
ldmatrix按 8 列块加载,把 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 offset | — | swizzle 模式每 128B 重复;起始不在边界时告诉硬件从哪个 128B 块开始 |
| Swizzle mode | 61–63 | 00=无 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 每个线程拿到多少寄存器
- 例:
wgmma用m64n256k16→ 每线程累加器有 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 group | 4 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 位 |
| swizzle | addr = (base + linear) ^ mask,bank 随 row 变 |
| packing | bf16 成对打包进 .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 页)。
