H100 · WGMMA Part 2(分组、stmatrix、FP8、稀疏)

目录 · ← l6 · l8 →

H100 · WGMMA Part 2(分组、stmatrix、FP8、稀疏)

本文档基于课程讲义《7. Wgmma part 2.pdf》(Lesson 7,43 页)整理, 承接 H100-WGMMA.md(Part 1),深入 WGMMA 的分组同步、结果写出(stmatrix)、FP8 与稀疏。 前置阅读:H100-WGMMA.mdH100-异步与屏障.mdH100-cuTensorMap.md


目录

  1. 分组:commit_group / wait_group
  2. Commit/Wait 流水线时间线
  3. wgmma.fence 的双重角色
  4. 结果写出:stmatrix
  5. stmatrix 的协作寻址
  6. stmatrix 与 FP32 的兼容性
  7. Swizzle 的 atom(16B)
  8. FP8:为何能翻倍算力
  9. e4m3 与 e5m2
  10. 饱和(Saturation)
  11. 缩放因子(Scaling Factors)
  12. FP8 打包(x2 / x4)
  13. FP8 转换与量化
  14. FP8 WGMMA 指令
  15. FP8 的 A 在寄存器、K-Major 规则与精度陷阱
  16. 稀疏 WGMMA(Sparse)
  17. sp-sel 与 sp-meta 的配置
  18. 总结与学习衔接

1. 分组:commit_group / wait_group

  • WGMMA 启动是异步的:wgmma.mma_async 把工作”放进在飞队列”并立即返回。
  • 逐条跟踪每个 MMA,要么付沉重的 scoreboard 开销,要么被迫用”等一切”的过度保守屏障。
  • 分组(group)给了我们流水线友好的粒度
    • wgmma.commit_group = “关闭当前批次”(使其可被跟踪)。
    • wgmma.wait_group N = “仅在在飞批次过多时才停顿”。
  • 这带来重叠:group g 在算时,你可以为 group g+1 做准备(地址计算、TMA、staging 等)。

1.1 wgmma.commit_group.sync.aligned

  • 只是把所有尚未提交的 wgmma.mma_async 打成包。
  • 硬件能同时跟踪多个组;分组可降低”逐条跟踪每个矩阵乘”的开销
  • 通常:发一组算出一个输出 tile(如 64×64)的 WGMMA 指令,然后把该 tile 提交为一个组。

1.2 wgmma.wait_group.sync.aligned N

  • 同步点:”暂停本线程,直到只剩 N 个已提交组仍在运行。”
    • wait_group 0:等一切完成——用于”即将读寄存器里的最终结果并写回 global”之前。
    • wait_group N:让 N 个组在后台继续跑,同时你去准备下一组数据。

2. Commit/Wait 流水线时间线

不是”先 commit 再 wait”,而是流水化

1. 发一批 wgmma.mma_async(这是 1 个输出 tile)
2. wgmma.commit_group.sync.aligned;
3. 开始准备下一个 tile(TMA、指针计算等)
4. wgmma.wait_group.sync.aligned N;   // 保留 N 个组在跑
5. 排空(drain)
6. wgmma.wait_group.sync.aligned 0;   // 等全部完成
7. 现在安全读累加器 D 并写出

3. wgmma.fence 的双重角色

wgmma.fence.sync.aligned 是 WGMMA 流水线的 warpgroup fence,承担两个紧密相关的角色:

  1. WGMMA 发射的”到达/组边界”标记:WGMMA 操作按组发射,fence 用于标记一个组发射窗口的起点。
  2. 操作数的寄存器排序 / 冒险控制wgmma.mma_async 是异步的、用 warpgroup 级流水线; fence 防止涉及累加器寄存器寄存器驻留的操作数片段(尤其是寄存器来源的 A)的重排/冒险。

注意:它不是跨 proxy fence。若你读的是”被 TMA 修改过”的寄存器数据,需要跨 proxy fence: fence.proxy.async(或等价的 acquire/release 发布协议)。


4. 结果写出:stmatrix

  • WGMMA 指令结束时,结果矩阵 C 在累加器寄存器里。结果不是连续矩阵(行 0、行 1……), 而是以高度不透明的方式分散在 128 个线程的寄存器里。
  • 普通 st.shared:谁写数据谁给地址——这需要手工管理地址,不行。
  • stmatrix:像 ldmatrix 一样,“数据提供者”与”地址提供者”解耦

4.1 stmatrix 指令

  • 每个线程输入自己从 wgmma 拿到的寄存器数据,并给出写目标指针。
  • .m8n8:原子单元仍是 8×8 tile。
  • .x4:每线程写 4 个寄存器(共 16×16 tile)。
  • .x1 / .x2:也支持(更小 tile)。
  • .shared:目标永远是 shared memory
  • .sync.aligned:所有 warp 必须同步、每个线程都发起该指令。
  • .m16n8:仅在 Blackwell 的 .b8 类型下有效。
  • .b16:表示所存元素的数据类型与大小。

5. stmatrix 的协作寻址

  • ldmatrix 一样,stmatrix协作的:warp 内每个线程(0–31)提供自己的指针(数据所属的 shared 地址)。
  • 因为线程持有的数据是特定的非线性模式,它们生成的地址也必须同样非线性,才能让数据在 shared memory 里落地成连续、整洁的矩阵
  • 执行 stmatrix 时,指令把这些元素按行主序(连续行)从 [addr] 开始写入 SMEM

5.1 m8n8 tile 结构

  • 指令操作一个 8×8 的 16 位元素 tile
  • .x1:每线程持有 1 个 8×8 的一部分;.x4:每线程持有 4 个 8×8。
  • ldmatrix 一样,特定线程持有特定行(线程 0–7 持有第 0 行,等等)。

5.2 Phase 1:地址责任(谁指哪)

ldmatrix 一致(16×16 的 x4 布局,四个 8×8 象限):

线程指向
0–7左上块(行 0–7,列 0–7)
8–15左下块(行 8–15,列 0–7)
16–23右上块(行 0–7,列 8–15)
24–31右下块(行 8–15,列 8–15)

5.3 寄存器来源(谁持有什么)—— “条纹”布局

  • 线程 0 是整个 tile 的”第 0、1 列拥有者”:它看到的不是连续行,而是竖直切片
  • 线程 0 的寄存器映射(x4):
    • r0:矩阵 A(左上)→ 行 0,列 [0,1]
    • r1:矩阵 B(左下)→ 行 8,列 [0,1]
    • r2:矩阵 C(右上)→ 行 0,列 [8,9]
    • r3:矩阵 D(右下)→ 行 8,列 [8,9]

5.4 序列化循环

  • stmatrix.x4 每次只消费每线程 4 个寄存器,但一个输出 tile 的累加器更大。
  • 因此要对累加器的”切片”做循环
    切片 i → 打包进 {r0,r1,r2,r3} → stmatrix(...)
    切片 i+1 → 打包 → stmatrix(...)
    
  • 这就是为什么 epilogue(尾声)通常是一个循环,而不是单次 store

6. stmatrix 与 FP32 的兼容性

最关键的一点:stmatrix 不支持 f32,只支持 .b16(打包 16 位)和 .b8(打包 8 位)。 这造成了处理 wgmma 输出的分歧:

  • FP32 累加器(.f32:累加器在未打包的 .f32 寄存器里,不能直接喂给 stmatrix; 必须先降精度并打包(会损失精度)。
  • FP16 累加器(.f16:累加器在打包的 .b32 寄存器(当作 .f16x2), 一个 32 位寄存器装两个元素,无需转换与打包

7. Swizzle 的 atom(16B)

  • 我们 swizzle 的 “atom” 是 16B
  • 16B 之所以重要:一个 m8n8 bf16 tile 的一行 = 8 元素 = 16 字节,所以每行天然是一个 16B 块。
  • 技巧在于:我们 swizzle 的是”行 atom”,不是单个元素
  • 即使 atom 被置换,每个 16B atom 内部的数据仍是行主序
  • 正因如此,TMA 才能”反 swizzle”回整洁的行主序矩阵

8. FP8:为何能翻倍算力

  • FP8 Tensor Core 理论上给出 BF16/FP16 两倍的 TFLOPS,H100 上约 3958 TFLOPS(含稀疏;稠密约一半)。
  • 位宽减半 → 降低 global memory 带宽压力,L2 可缓存 2 倍大的模型/batch。
  • 收益主要在矩阵乘密集的部分(线性层、注意力投影、MLP);许多流水线仍把 softmax/归一化/规约留在 BF16/FP16。

8.1 范式转变

  • 约 2022 年前,AI 标准是 FP32 或 FP16/BF16。
  • FP8 主要因 H100(Hopper) 而工业化,能用一半内存把训练提速两倍
  • 大模型实例:DeepSeek V3 用 FP8 训练Kimi K2 用 INT4 量化感知训练(QAT)
  • 与 INT8 不同,FP8 是”浮点”,能表示宽动态范围。8 位太挤,故拆成两个专用格式平衡”精度 vs 范围”:e4m3 和 e5m2

9. e4m3 与 e5m2

 e4m3e5m2
指数/尾数4 位指数 + 3 位尾数5 位指数 + 2 位尾数
范围≈ ±448(最大正常值 448,最小正常 0.015625)≈ ±57344(远大于 e4m3)
用途权重与激活(前向)——动态范围好,表示 NaN/Inf训练梯度与反向传播——能表示梯度更新需要的很大数值

10. 饱和(Saturation)

  • FP32 动态范围巨大(~10³⁸),几乎不会触顶。
  • FP8 e4m3 的”天花板”是 448:若矩阵乘结果是 450 而不处理,就饱和。把 float 降成 FP8 时,硬件要决定怎么处理 450,通常两种模式(由 intrinsic 控制):
    1. 钳位到最大可表示值(448):数据被裁剪(clipping),但数学”稳定”。
    2. 变 Inf(或 e4m3 里变 NaN,因 e4m3 无 Inf 表示):立刻毁掉你的训练。

11. 缩放因子(Scaling Factors)

既然改不了 8 位的物理,就只能”作弊”——用缩放因子。这是 H100 上 FP8 的核心课。

  • Tensor-wise:整个张量一个缩放因子(取 max 值)——最简单、开销最低。
  • Vector-wise:每行/每列一个缩放因子。
  • Block-wise(MXFP8 / 微缩放):把矩阵切成小 tile(如 16×16 或 32×32),每个 tile 一个 scale。
  • 其它:SmoothQuant、Delayed Scaling 等。

12. FP8 打包(x2 / x4)

  • FP8 在传输时不是独立对象:H100 内存控制器不会为小于 32 字节(一个 sector)的数据醒来。 若你要 1 字节(一个 FP8 值),GPU 会取 32 字节、给你那 1 个、丢掉其余 31 个。
  • 为消除浪费,用打包类型 e4m3x2/4e5m2x2/4(硬件理解的 PTX 类型)。

12.1 x4 Pack(寄存器填充型)

  • .e4m3x4 / .e5m2x4:装 4 个 FP8
  • 总大小 8 位 × 4 = 32 位 → 恰好填满一个标准 GPU 寄存器
  • 布局(小端):
    [Element 3 \| Element 2 \| Element 1 \| Element 0]
     <-- MSB (31)                    LSB (0) -->
    
  • 这是存储与传输的主格式;搬数据或做简单数学(如找 max)时用 x4。

12.2 x2 Pack(数学输入型)

  • .e4m3x2 / .e5m2x2:装 2 个 FP8,共 16 位,与 FP16 同大小
  • H100 被设计成帮人从 FP16 迁到 FP8:有专用逻辑把一个装两个 FP16 的 32 位寄存器直接压缩成装两个 FP8 的 16 位。
  • Tensor Core 常以 16 位块摄取数据

13. FP8 转换与量化

  • 源数据是 FP16/FP32 时,用 cvt(常用向量形式 .e4m3x2/.e5m2x2)转 FP8,通常带 .satfinite 让越界值钳位:
    cvt.rn.satfinite.e4m3x2.f16x2   // 取 2 个 fp16 → 饱和 → 舍入 → 打包
    
  • 在喂 WGMMA 前,先决定量化策略(舍入模式、是否 satfinite、per-tensor/per-channel 缩放)。
  • PTX 描述饱和行为:\|input\| 超过目标最大正常值时,结果变为保号的 max normal
  • PTX 为 sm_90+ 提供 cvt.satfinite.{e4m3x2,e5m2x2}.{f32,f16x2}

14. FP8 WGMMA 指令

wgmma.mma_async.sync.aligned.m64nNk32.f32.e4m3.e4m3
    d, a-desc, b-desc, scale-d, imm-scale-a, imm-scale-b;
  • m64nNk32:M 固定 64,K = 32(对 8 位),N 可变(8, 16, … 256)。
  • .f32:累加器 D 类型(单精度)。
  • .e4m3.e4m3:A、B 的输入类型。FP8 特殊在 A、B 可为不同 FP8 格式(如 .e4m3 × .e5m2 混合)。
  • scale_*:缩放/符号翻转因子(A、B 取 {−1, 1},scale_d 取 {0, 1})。
  • K = 32 是关键区别:一次吞掉 A 的 32 列。

14.1 Hopper 第四代 Tensor Core 的 FP8

  • 输入 8 位(FP8),但累加在 32 位(FP32)寄存器里以保证数值稳定。
  • 通过 WGMMA 原生支持 e4m3 与 e5m2
  • 一条 WGMMA 指令发出巨量数学,有效隐藏延迟。
  • 寄存器虽是 FP32,实际像 FP22(~8 位指数 + ~13 位尾数)——已知的硬件特性。
  • 本质:把 FP8 当压缩存储格式,在”伪 FP32”空间里做数学

15. FP8 的 A 在寄存器、K-Major 规则与精度陷阱

15.1 A 在寄存器(RS)

  • FP8 RS WGMMA 中,A 操作数是直接传给 WGMMA 的寄存器向量表达式
  • “显然”的办法(ldmatrix ... .b8)在 Hopper 上不可用(SM90a 的 PTX 说明 .b8 ldmatrix 仅 sm_100a+ 支持)。 Hopper 上不能用 ldmatrix 加载 FP8。
  • 因此大多数 Hopper FP8 WGMMA 实现二选一:
    • SS 路径:A、B 都用描述符;
    • RS 路径:用普通 shared load(ld.shared.b32 / 向量化)把 A 装进寄存器,事先按 WGMMA 想要的打包方式排好,而不是用 ldmatrix。

15.2 严格的 K-Major 规则

  • wgmma.mma_async 不暴露转置控制(imm-trans-a/b(不像 FP16/BF16 变体)。
  • 无转置时,WGMMA 按默认 K-major 规范布局解释 shared 操作数;喂 MN-major 则无法”重解释”,必须自己打包/转置。
  • 转置/打包步骤会增加指令、同步与 shared 流量,伤吞吐;MN-major 还会让 FP8 的 swizzle-atom 对齐/整除约束更复杂(128B atom)。
  • 结论:FP8 tile 用 K-major staging(或离线预打包);只有 staging 期间会转置时才用 MN-major。

15.3 精度陷阱

  • 即便文档说用 FP32 累加器,实测报告 FP8 Tensor Core 累加像”降精度 FP32”(~8 位指数 + ~13 位尾数,即”共 22 位”)——最低的 FP32 尾数位在累加时可能实际丢失。
  • 大点积(大 K)下,缩小的有效尾数会增大舍入/截断误差,且误差随规约深度累积。
  • 对策:用 K-slicing 和/或多阶段累加(先累加部分和,再以更高精度归约部分和),限制每个 WGMMA “chunk” 的累加深度,改善数值稳定。

16. 稀疏 WGMMA(Sparse)

  • 一种特殊 WGMMA:Tensor Core 做同样的 MMA,区别是 A 是结构化稀疏(50% 零)
  • 操作数:descA(或 RS 形式的 A 寄存器片段)、descBsp-meta(含打包索引的 .b32 寄存器)、 sp-sel(32 位常量,”稀疏选择器”,选择哪些线程贡献某组的 metadata),以及其它操作数。

16.1 重要点

  • 硬件严格按”稀疏 A × 稠密 B”运行:即使 B 数学上稀疏(含零),Tensor Core 也当它是稠密矩阵。
  • 需为每个线程建一个 sp_meta 寄存器;只在硬件会读的 lane 里给 sp_meta 赋有意义值, 其它 lane 通常置 sp_meta = 0(或任意值),因为那些 lane 的 metadata 被忽略。

16.2 打包与 Metadata(2:4 结构化稀疏)

  • NVIDIA 稀疏 Tensor Core 的物理规则是 2:4 结构化稀疏
    • 提供 4 个连续值(叫 Quartet),必须删掉其中恰好 2 个,只存两个幸存者(在 global 或经 shared 变换)。
    • 50% 带宽与存储
    • 但若只给 GPU [8.5, 3.2],它不知道它们原来在哪(8.5 来自索引 0、1 还是 2?)——这就是 metadata 的作用。
  • 因为从 4 个位置(0,1,2,3)里选 2 个,需要编码幸存者的位置

17. sp-sel 与 sp-meta 的配置

17.1 sp-sel(线程选择器)

  • 告诉硬件谁负责提供 metadata——即在每组 4 个连续线程(T0–T3)里,哪个线程对是 metadata 贡献者(只在”仅一对贡献”时)。
  • 各精度要求:
    • TF32 稀疏(.m64nNk16 .tf32:spSel 须为 0(T0,T1)或 1(T2,T3)。
    • FP16/BF16 稀疏(.m64nNk32 .f16/.bf16:spSel 须为 0 或 1。
    • FP8/INT8 稀疏(.m64nNk64 .e4m3/.e5m2/.s8/.u8所有线程都贡献 metadata,spSel 必须为 0(否则未定义)。

17.2 sp-sel 的两种选择规则

  • 选 0:从 T0、T1 读 metadata,忽略 T2、T3。
  • 选 1:从 T2、T3 读 metadata,忽略 T0、T1。
  • 规则 1(Replicated):若加载器把 metadata 广播给所有线程(常见)→ 始终 sp-sel=0(最简单)。
  • 规则 2(Sharded):若拆分 metadata 省寄存器 → 动态让 sp-sel 匹配持有数据的线程
  • 配置错误 → Tensor Core 从空寄存器读,把块当稠密或清零。

17.3 sp-meta(metadata 位域)

  • 告诉硬件 A 的非零在哪spMeta 是打包位域,规则:
    • 2:4 稀疏(FP16/BF16、FP8、INT8):A 每 4 个相邻元素有 2 个非零;只存 2 个非零,其位置(0..3)用两个 2 位索引编码进 metadata。
    • 1:2 稀疏(TF32):A 每 2 个相邻元素有 1 个非零;metadata 用 4 位索引指示 2 个位置中的哪个,只有两个特定位模式有意义,其它值未定义。

17.4 配置 sp-meta(FP16/BF16/INT8)

  • 一个寄存器装载的 metadata 恰好够 4 条连续 WGMMA
  • 规则:主循环展开 4 次,让 sp-meta 循环 0,1,2,3(第一个 K-tile 用 0,第二个用 1……)。
  • 该索引递增硬件内部指向”寄存器内下一组 2 位索引”的指针。
  • 不递增 → 硬件对所有计算重复用第一个 tile 的稀疏模式。

17.5 配置 sp-meta(TF32 的例外)

  • TF32 寄存器只够 2 条 WGMMA,策略特殊:不能用顺序索引(0,1),必须在硬编码位掩码 14 和 4 之间切换
    • 第 1 次 K=16 计算:sp-meta = 0b1110 (14)(解码低位)。
    • 第 2 次 K=16 计算:sp-meta = 0b0100 (4)(解码高位)。
  • 这些特定值是让硬件 swizzler 对齐非标准的 19 位/32 位 TF32 数据格式所必需的。
  • 用 0 或 1 这类普通索引 → 硬件错位、矩阵结果错误。

18. 总结与学习衔接

18.1 核心脉络速记

概念一句话
commit/wait_group把异步 WGMMA 分组;wait_group N 保留 N 组在跑,实现流水重叠
wgmma.fence组边界标记 + 寄存器排序(不是跨 proxy fence)
stmatrix协作写回(地址/数据解耦),只支持 .b16/.b8,FP32 要先降精度打包
swizzle atom16B;swizzle 的是”行 atom”,行内仍行主序,TMA 可反 swizzle
FP8e4m3(前向)/e5m2(梯度);靠缩放因子、打包(x2/x4)、cvt.satfinite
FP8 WGMMAK=32,A/B 可混合格式,K-major 强制,累加像 FP22
稀疏2:4 结构化;sp-sel 选 metadata 线程对;sp-meta 编码非零位置;TF32 用 14/4 特例

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

本文概念对应后续专题
流水线(commit/wait + TMA 供数)《8. Kernel Design》《8.1 Stream-K》
内核启动 / 集群协同《8.2 Kernel Launch》
多 GPU 扩展《9. Multi GPU》《10. Multi GPU Part 2》

18.3 一句话记忆

WGMMA Part 2 = 把异步 MMA”流水线化 + 正确写出 + 玩转低精度与稀疏”: commit/wait_group 控制重叠、stmatrix 把碎片化的寄存器结果写成整洁 SMEM、 FP8 靠缩放/打包/饱和翻倍算力、稀疏靠 sp-meta/sp-sel 再省一半带宽。


参考来源:7. Wgmma part 2.pdf(Lesson 7,43 页)。