编程 Tensor Core 三种调用层次:m16n8k16 的 ldmatrix 对不上,以及 Hopper wgmma 的同步与布局约束

2026-09-28 00:04:03

Tensor Core 三种调用层次:m16n8k16 的 ldmatrix 对不上,以及 Hopper wgmma 的同步与布局约束

自 Volta 起,NVIDIA 在 GPU 中加入 Tensor Core,专门做小块矩阵乘加 D = A × B + C,支持 fp16/bf16/tf32/fp8 等低精度以获得更高吞吐。调用方式可以分成三层:封装库、CUDA WMMA、PTX MMA。三者的控制粒度和踩坑点差别很大。

1. 三层调用:库 / WMMA / PTX MMA

  • 封装库:cuBLAS 做 GEMM,cuDNN 做卷积/RNN,TensorRT 做推理。库层不直接暴露 fragment 布局,调用简单但难以插入自定义流水。
  • WMMA:nvcc 提供 warp 级 fragment 抽象,API 如 load_matrix_sync、store_matrix_sync、mma_sync。编程简单,但控制粗糙。
  • PTX MMA:直接写 PTX 汇编,面向寄存器,用 mma.sync。控制精细,但寄存器布局和 shared memory 地址必须自己对齐,容易出错。

WMMA 场景里,A/B 通常是 16x16 fp16,C/D 可以是 fp16 或 fp32。Tensor Core 不能直接使用 SMEM 或本地内存中的数据,必须先 load 到 fragment(wmma.matrix_a / wmma.matrix_b)与累加器(wmma.accumulator)。Tensor Core 只在 Volta 及以上 GPU 支持,例如 V100、T4、RTX-20X0、A100、RTX-30X0;Colab 里多数 GPU 太旧,不支持 Tensor Core。入门教程见 MLC Tensor Core 教程,另有一个入门仓库 Tensor_Core_Learning,其中每个 Tensor Core 提供 4x4x4 矩阵阵列做 D=A*B+C。

2. ldmatrix:行地址与寄存器映射

ldmatrix 指令形状如下:

ldmatrix.sync.aligned.shape.num{.trans}{.ss}.type r, [p];
.shape={.m8n8,.m16n16};
.num={.x1,.x2,.x4};
.ss={.shared{::cta}};
.type={.b16,.b8};

例子:

ldmatrix.sync.aligned.m8n8.x1.shared.b16 {%0}, [%1];

语义是一个 warp(32 线程)从 shared memory 的 [p] 处加载矩阵到寄存器。m8n8 时 64 个 bf16 元素等于 128 字节,只要求行内连续,行间可以不连续。地址寄存器需要由 thread0-7 填 8 个行首地址:.x1 用 0-7,.x2 用 0-15,.x4 用 0-31。每 4 个连续线程读取连续一行,每线程 2 个元素 4 字节,线程 0 取头两个元素。

.m16n16 且 b8 时,每线程 8 个元素 8 字节,需要 2 个寄存器,%0 存第一个 m8n8,%1 存第二个,并且同一线程的两个寄存器分别落在第 0 行与第 8 行。CUDA 里可以用 __asm__ volatile 内嵌:

#define LDMATRIX_X2_T(R0, R1, addr) \
  asm volatile("ldmatrix.sync.aligned.x2.trans.m8n8.shared.b16 {%0, %1}, [%2];\n" \
               : "=r"(R0), "=r"(R1) : "r"(addr))

注意 warp 内每个线程都要按需准备寄存器,否则会出错;shared 地址要用 __cvta_generic_to_shared 转成 shared 空间偏移。m8n8 与 m16n16、是否 .trans、.x1/.x2/.x4 的组合,会直接改变每个 lane 拿到的元素位置。

3. m16n8k16:A/B/C 的寄存器切片

PTX 中 m16n8k16 的典型形式:

mma.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16 {%0,%1},{%2,%3,%4,%5},{%6,%7},{%8,%9};

常见形状包括:m8n8k4(alayout/blayout,f16)、m16n8k8、m16n8k16(f16/bf16/tf32)、m16n8k32(fp8/fp6/fp4 混合)。.ctype / .dtype={.f16,.f32}。tf32 用 m16n8k4 / m16n8k8,bf16 用 m16n8k16。

以 m16n8k16 bf16 为例:A 是 16x16(M×K)共 256 元素,每线程 8 元素,需要 4 个寄存器,每寄存器 2 个元素。A[0][0]、A[0][1] 在 %2,A[8][0]、A[8][1] 在 %3,依此类推。B 必须列主序 16x8(K×N):B[0][0]、B[1][0] 在 %6,B[8][0]、B[9][0] 在 %7。C/D 同为 16x8:C[0][0]、C[0][1] 在 %8,C[8][0]、C[8][1] 在 %9。PTX 文档里有完整矩阵 fragment 说明:Matrix fragments for mma.m16n8k16 with floating point type。

布局容易对不上,就是因为 ldmatrix 输出的 lane→元素映射,必须和 mma.sync 要求的 fragment 顺序完全一致。一个 naive 例子 hgemm_mma_m16n8k16 可以说明这件事:每个 block 一个 warp(32 线程),__shared__ s_a[16][16]、s_b[16][8]、s_c[16][8];A 从 gmem 按 tid/2 行、(tid%2)*8 列做 128bit 载入;B 只由 lane_id < 16 载入;用 __cvta_generic_to_shared 取地址,A 用 LDMATRIX_X4,B 用 LDMATRIX_X2_T,然后 HMMA16816 计算。C 写回 s_c[lane_id/4][(lane_id%4)*2] 与 s_c[lane_id/4+8][...],其中 rc[0] 对应前半行,rc[1] 对应后半行,lane_id/4 定行,(lane_id%4)*2 定列。

载入 A 时,索引 [lane_id%16][(lane_id/16)*8] 是为了让线程 0-15 读前 16 列,线程 16-31 读后 16 列,从而让 ldmatrix 之后的寄存器顺序和 MMA 期望顺序一致。如果这里行列索引、trans 标志或 .x2/.x4 选错,ldmatrix 本身不会报错,但后续 mma.sync 会按另一套 lane 映射解释寄存器,结果自然对不上。

4. Hopper wgmma:warp group 与同步原语

Hopper 引入 warp group MMA(WGMMA)。一个 warp group 等于 4 个连续 warp,也就是 128 个连续线程,第一个 warp 编号必须是 4 的倍数。指令 wgmma.mma_async 由 128 个线程共同执行,异步,以整个 SM 为粒度,可做远大于传统 MMA 的矩阵分块,是 H100 上 MatMul/Attention 的关键原语,CUTLASS 也使用它。

bf16 操作数形状为 m64nNk16,N ∈ {8,16,24,...,256}。工程里常用 m64n64k16,一般较大的 N 性能更好,前提是寄存器和 SMEM 足够。

WGMMA 的关键同步原语:

  • wgmma.fence:确保 wgmma.mma_async 访问寄存器内存前,之前的访问已完成,否则行为未定义。例外是累加器形状相同的多条 wgmma 可以共享同一累加器,无需 fence。
  • wgmma.commit_group:把尚未提交的 wgmma.mma_async 批量归入一个新的 wgmma 组。
  • wgmma.wait_group:等待 wgmma 组。
  • warpgroup_arrive / commit_batch / wait 是 CUTLASS/cute 的封装。

代码里常见:

asm volatile("wgmma.fence.sync.aligned;" ::: "memory");

5. GmmaDescriptor 与布局约束

当 A 从 SMEM 取时,WGMMA 使用矩阵描述符(GmmaDescriptor)而不是寄存器张量。描述符由 make_gmma_desc 根据 SMEM 张量布局计算:算出 LBO(首维字节偏移)与 SBO(跨步字节偏移),并确定 swizzle 模式。这里需要符合 8 种规范 GMMA 布局原子,再配合 tile_to_shape。K-major 布局要求 shape00==8、shape11%2==0。SBO/LBO 由 base_ptr>>4 计算,忽略低 4 位;swizzle 与无 swizzle 时二者的取值互换。TransB 对应封装参数 transpose_b。这些约束让 WGMMA 的 SMEM 布局不能随意排,描述符字段错一个,计算结果就会偏离。

6. TMA 流水中的 producer / consumer

生产内核不是单条指令,而是与 TMA 深度流水化。producer warp group(wg_idx==0)用 cp_async_bulk_tensor_2d_global_to_shared 把 A/B tile 经 TMA 搬入 SMEM,配合 mbarrier(empty/full 队列、barrier_arrive_tx 绑定字节数)做多级缓冲。consumer warp group(wg_idx>0)等待 full[qidx] 后,对缓冲区跑 Tensor Core MMA,完成后标记 empty 供 producer 复用;MatmulTileWriter 把寄存器 C 写回全局。相关实现可参考 CUTLASS。WGMMA 编程的更多拆解见 深入 NVIDIA GPU 高性能矩阵乘法算子解构(四),以及 AtomGit/GitCode 博客的 WGMMA 编程指南。Tensor Core 与 MMA 的基础整理见 Tensor Core 和 MMA。

复制全文 生成海报 tensor core 张量核心 WGMMA CUDA PTX

推荐文章

程序员茄子在线接单