CUDA 算子工程:手写 FlashAttention v2 之路
第 12 章 Tensor Core GEMM:mma.sync 与 ldmatrix
FP32 SIMT 这条赛道的天花板是 67 TFLOPs,FP16 Tensor Core 是 989——差 15 倍。 换赛道的入场券是三样东西:mma.sync 的矩阵语义、fragment 布局、ldmatrix 与 swizzle。
12.1 为什么 Tensor Core 是必经之路
第 11 章我们用 SMEM 与寄存器分块把 SIMT GEMM 推到了 FP32 通路的合理水平。但回顾 Hopper 算力(H100 SXM5 官方规格,Tensor Core 为稠密口径,2:4 稀疏再翻倍):
FP32 SIMT 峰值: 67 TFLOPs/s
FP16 Tensor Core 峰值: 989 TFLOPs/s
FP8 Tensor Core 峰值: 1979 TFLOPs/s
Tensor Core 比 SIMT 快 15× 到 30×。任何严肃的 LLM 训练 / 推理都必须用 Tensor Core——这不是优化选项,是入场券。
但 Tensor Core 不是一个"快版本的 FMA 指令"——它是一个全新的编程模型:
- 指令是矩阵级的:一条
mma.sync算 16×8×16 矩阵乘,不是单个浮点。 - 数据需要特殊布局:mma 输入要按 NVIDIA 定义的 fragment 格式排列。
- 加载需要专用指令:
ldmatrix一次性把 16×16 数据从 SMEM 拉成 fragment。 - 输出是分布式的:累加结果分布在 32 个线程的寄存器里,不是连续存储。
这一章我们把这套新的编程模型彻底讲透。
12.2 mma.sync:一条指令算一个矩阵乘
Ampere+ 上的核心 Tensor Core 指令是:
mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32
D, A, B, C
含义:
- m16n8k16:算 D = A @ B + C,其中 A 是 16×16,B 是 16×8,D = C = 16×8。
- row.col:A 是行优先,B 是列优先(即 K 维连续);m16n8k16 只支持这一种组合。
- f32.f16.f16.f32:D 和 C 是 fp32,A 和 B 是 fp16。
- D, A, B, C:四组寄存器(不是矩阵指针!)
每条指令的浮点操作数:
反推一下 H100 的峰值:989 TFLOPs ÷ 132 SM ÷ 1.83 GHz ≈ 4096 FLOPs / SM / 周期——正好是一条 m16n8k16 的量。也就是说按峰值折算,一个 SM 的 4 个 Tensor Core 合起来每周期交付一条 mma.sync 的算力(每个 sub-core 每 4 个周期完成一条),而不是每个 Tensor Core 每周期一条。(这只是峰值折算:Hopper 上 mma.sync 实际跑不满,见 12.5 节。)
12.2.1 Fragment 布局
最反直觉的部分:mma.sync 的 A、B、C、D 不是单个寄存器,而是一组寄存器,分布在 32 个线程上:
A (16×16, FP16) 共 256 个 fp16 = 512 字节 = 128 个 32-bit 寄存器。 分布在 32 lane 上,每 lane 4 个寄存器(128 / 32 = 4)。
具体的分布模式很复杂,由 NVIDIA 硬件规定:
A 的 fragment layout (m16n8k16, row-major; 每格 = 同一 lane 持有的 2 个相邻 fp16):
k=0..7 k=8..15
┌──────────────────┐ ┌──────────────────┐
m=0: │ T0 T1 T2 T3 │ │ T0 T1 T2 T3 │
m=1: │ T4 T5 T6 T7 │ │ T4 T5 T6 T7 │
... │ ... │ │ ... │
m=7: │ T28 T29 T30 T31 │ │ T28 T29 T30 T31 │
m=8: │ T0 T1 T2 T3 │ │ T0 T1 T2 T3 │
... │ ... │ │ ... │
m=15: │ T28 T29 T30 T31 │ │ T28 T29 T30 T31 │
└──────────────────┘ └──────────────────┘
每 lane 持有 4 个 fp16 每 lane 持有 4 个 fp16
按 PTX 手册对 m16n8k16 的规定(A 是 16 行 × 16 列,行=m、列=k),lane 持有的 8 个 fp16 是:
组 0 (a0,a1): 行 l/4, 列 (l%4)*2 + {0,1}
组 1 (a2,a3): 行 l/4 + 8, 列 (l%4)*2 + {0,1}
组 2 (a4,a5): 行 l/4, 列 (l%4)*2 + 8 + {0,1}
组 3 (a6,a7): 行 l/4 + 8, 列 (l%4)*2 + 8 + {0,1}
也就是 lane 0 持有 A[0,0..1]、A[8,0..1]、A[0,8..9]、A[8,8..9]——同一个 lane 拿到的是两行的四个片段,既不连续也不同行。
读者完全不需要记这个表——下一节的 ldmatrix 会自动按这个布局排好。但重要的是理解:fragment 不是连续存储,而是分布式存储。
12.2.2 Inline PTX
CUDA C++ 写 mma.sync 用 inline PTX:
unsigned A[4]; // 4 个 32-bit, 每个 = 2 个 fp16, 共 8 个 fp16 (本 lane 分到的 A fragment)
unsigned B[2]; // 2 个 32-bit = 4 个 fp16 (本 lane 分到的 B fragment)
float C[4]; // 4 个 fp32 (本 lane 分到的累加器 fragment)
asm("mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 "
"{%0, %1, %2, %3}, "
"{%4, %5, %6, %7}, "
"{%8, %9}, "
"{%0, %1, %2, %3};\n"
: "+f"(C[0]), "+f"(C[1]), "+f"(C[2]), "+f"(C[3])
: "r"(A[0]), "r"(A[1]), "r"(A[2]), "r"(A[3]),
"r"(B[0]), "r"(B[1]));
或者用 CUDA 9 起提供的 nvcuda::wmma C++ API(对应 PTX 的 wmma.* 指令,sm_70 起可用;形状是 m16n16k16 这类整块,fragment 内部布局不公开,更高级但灵活性差)。CUTLASS 用 inline PTX(如 cutlass-4.7.0/include/cute/arch/mma_sm80.hpp:173)。
12.3 ldmatrix:把 SMEM 数据加载成 fragment
mma.sync 要求 fragment 已经在寄存器里,且按特定布局排列。怎么把 SMEM 数据装进 fragment?
最朴素的方式是每个线程自己 load:
unsigned A[4];
A[0] = reinterpret_cast<unsigned*>(&sA[m + lane_id / 4][k + (lane_id % 4) * 2])[0];
// ... 算地址再 load 4 次
地址计算超复杂,且每线程独立 load 会触发 bank conflict。
NVIDIA 提供了 ldmatrix 指令——一条指令把 SMEM 中一个 16×16 子块加载到 32 个 lane 的 fragment:
unsigned A[4];
// ldmatrix 取的是 shared 地址空间的 32-bit 地址, 约束是 "r" 不是 "l"
unsigned smem_addr = static_cast<unsigned>(__cvta_generic_to_shared(smem_ptr));
asm("ldmatrix.sync.aligned.m8n8.x4.shared.b16 "
"{%0, %1, %2, %3}, [%4];\n"
: "=r"(A[0]), "=r"(A[1]), "=r"(A[2]), "=r"(A[3])
: "r"(smem_addr));
ldmatrix.x4 一次加载 4 个 8×8 fp16 子块(合计 16×16),输出 4 个寄存器/线程。32 lane × 4 寄存器 = 128 个寄存器 = 256 fp16 = 16×16 矩阵。注意 smem_ptr 每个 lane 各不相同:每个 lane 提供一行(8 个 fp16 = 16 字节)的起始地址,lane 0–7、8–15、16–23、24–31 依次给第 0–3 个子块的 8 行。让这 4 个子块依次是 A 的左上、左下、右上、右下(即 lane 指向第 l % 16 行、第 (l / 16) * 8 列),4 个输出寄存器就恰好是 a0a1、a2a3、a4a5、a6a7——完美匹配 mma.sync 的输入 fragment 布局。
ldmatrix 还有一个 .trans 变种(如 ldmatrix.sync.aligned.m8n8.x4.trans.shared.b16)——加载时就地转置。这对加载 B 矩阵特别有用:mma 的 .col 要求每个 lane 拿到的是沿 K 方向相邻的两个元素,若 B 在 SMEM 里按 K×N 行优先(N 维连续)存放,用 .trans 加载就能直接得到这种排列;若 B 本来就按 N×K(K 维连续)存放,则不需要 .trans。
12.4 SMEM Layout 与 Swizzle
ldmatrix 对 SMEM 的排布很敏感。如果 SMEM 是简单的 row-major,ldmatrix 会触发严重的 bank conflict——它按 8×8 子块分阶段访问,每个子块的 8 行各 16 字节;行跨度是 128 字节的整数倍时,这 8 行全落在同一组 4 个 bank 上,是 8 路冲突(行跨度 64 字节时是 4 路),该阶段吞吐随之降到 1/8(1/4)。
通行的解法是 swizzle layout(CUTLASS 的标准做法;Hopper 的 TMA 还把它做进了硬件,见第 4 章),让 ldmatrix 访问的地址自动错开 bank:
flowchart TB
subgraph LinearLayout [Row-major Layout]
L1["行 0:16B 块 0, 1, 2, 3, 4, 5, 6, 7"]
L2["行 1:16B 块 0, 1, 2, 3, 4, 5, 6, 7(同列与行 0 同 bank)"]
end
subgraph SwizzleLayout [Swizzled Layout:块号 XOR 行号低 3 位]
S1["行 0:16B 块 0, 1, 2, 3, 4, 5, 6, 7(XOR 0,不变)"]
S2["行 1:16B 块 1, 0, 3, 2, 5, 4, 7, 6(XOR 1)"]
end
简单说,swizzle 把每行的 16 字节块按一个 XOR 函数重排:
__device__ int swizzle_idx(int row, int col) {
// col 以 fp16 元素计, 假设每行 64 个 fp16 (128 字节)
// 8 个 fp16 = 16 字节是 ldmatrix 一行的单位, 块内不打散;
// 用 row 的低 3 位去 XOR 块号 (col / 8) 的低 3 位
return (((col >> 3) ^ (row & 0x7)) << 3) | (col & 0x7);
}
具体实现还有几种变体,但核心思想都是用 row 的低位去 XOR 列方向的 16 字节块号,让同一个 8×8 子块的 8 行落到不同 bank。注意 XOR 的单位必须是 16 字节块:只在 8 个元素内部重排的话,ldmatrix 每行读的还是同一组 bank,冲突一点不少。
CUTLASS 把这套东西抽象成了 CuTe 的 Swizzle<BBits, MBase, SShift>(cutlass-4.7.0/include/cute/swizzle.hpp:55),三个参数的含义是:MBase 指定"不打散的基本单元"占多少位(单元大小 ),BBits 指定拿多少位去做 XOR,SShift 指定从哪一段位取这个 XOR 源。核心那一行就是
// cutlass-4.7.0/include/cute/swizzle.hpp 的 Swizzle::apply
return offset ^ shiftr(offset & yyy_msk{}, msk_sft{}); // ZZZ ^= YYY
跟上面手写的 swizzle_idx 是同一个东西,只是把"哪几位 XOR 哪几位"参数化了:作用在元素偏移 row * 64 + col 上,它就是 Swizzle<3,3,3>。GMMA 实际用的四档写在 cutlass-4.7.0/include/cute/atom/mma_traits_sm90_gmma.hpp:75-:84:Swizzle<0,4,3>(不 swizzle)、Swizzle<1,4,3>(32B)、Swizzle<2,4,3>(64B)、Swizzle<3,4,3>(128B),正好对应第 4 章讲的 TMA descriptor 那四个 swizzle 档位。注意这几个 atom 的 Layout 部分以比特为单位(名字里的 _Bits),但经 smem_ptr_flag 组合后 swizzle 作用在 SMEM 字节地址上(cute/pointer_swizzle.hpp:87):MBase = 4 即 16 字节一块,Swizzle<3,4,3> 用字节地址第 7–9 位去 XOR 第 4–6 位,形成 128B 的 swizzle 模式。第 13 章会展开。
12.5 完整的 Tensor Core GEMM 骨架
把 mma + ldmatrix + swizzle 拼起来,给一个 HGEMM kernel 骨架(假设 sA 按 BM×BK、sB 按 BK×BN 行优先存放):
template <int BM = 128, int BN = 128, int BK = 32>
__global__ void hgemm_tensorcore(
const half* A, const half* B, half* C,
int M, int N, int K
) {
__shared__ half sA[BM * BK]; // 8 KB (128*32*2 byte)
__shared__ half sB[BN * BK]; // 8 KB
const int tid = threadIdx.x;
const int warp_id = tid / 32;
const int lane_id = tid % 32;
const int warp_m = warp_id / 4; // 2 warps in M
const int warp_n = warp_id % 4; // 4 warps in N
// 一个 block 8 warp (256 线程), 处理 BM × BN = 128×128
// 每 warp 64 × 32 (2 × 4 = 8 个 warp 刚好铺满 128×128)
constexpr int WM = 64, WN = 32;
constexpr int MMAS_M = WM / 16; // 4 (mma 的 M 是 16)
constexpr int MMAS_N = WN / 8; // 4 (mma 的 N 是 8)
// 累加器 fragment
float c_frag[MMAS_M][MMAS_N][4] = {0};
const int block_row = blockIdx.y * BM;
const int block_col = blockIdx.x * BN;
for (int k_step = 0; k_step < K; k_step += BK) {
// 1. cp.async 加载 A_tile, B_tile 到 sA / sB (使用 swizzle layout)
cp_async_load_a_tile(sA, A, block_row, k_step);
cp_async_load_b_tile(sB, B, block_col, k_step);
cp_async_commit_and_wait();
__syncthreads();
// 2. 内层 K (BK / 16 个 mma 步)
for (int kk = 0; kk < BK; kk += 16) {
// 用 ldmatrix 加载 A fragments
unsigned a_frag[MMAS_M][4];
#pragma unroll
for (int i = 0; i < MMAS_M; ++i) {
int row_offset = warp_m * WM + i * 16;
ldmatrix_x4(sA, row_offset, kk, &a_frag[i]);
}
// ldmatrix 加载 B fragments (sB 是 K×N 行优先, 所以用 .trans)
unsigned b_frag[MMAS_N][2];
#pragma unroll
for (int j = 0; j < MMAS_N; ++j) {
int col_offset = warp_n * WN + j * 8;
ldmatrix_x2_trans(sB, col_offset, kk, &b_frag[j]);
}
// 3. mma.sync 累加
#pragma unroll
for (int i = 0; i < MMAS_M; ++i)
#pragma unroll
for (int j = 0; j < MMAS_N; ++j) {
asm("mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 "
"{%0, %1, %2, %3},"
"{%4, %5, %6, %7},"
"{%8, %9},"
"{%0, %1, %2, %3};\n"
: "+f"(c_frag[i][j][0]), "+f"(c_frag[i][j][1]),
"+f"(c_frag[i][j][2]), "+f"(c_frag[i][j][3])
: "r"(a_frag[i][0]), "r"(a_frag[i][1]),
"r"(a_frag[i][2]), "r"(a_frag[i][3]),
"r"(b_frag[j][0]), "r"(b_frag[j][1]));
}
}
__syncthreads();
}
// 4. 写 C (epilogue: fp32 -> fp16, 写回 HBM)
#pragma unroll
for (int i = 0; i < MMAS_M; ++i)
#pragma unroll
for (int j = 0; j < MMAS_N; ++j) {
int row = block_row + warp_m * WM + i * 16;
int col = block_col + warp_n * WN + j * 8;
// 每 lane 持有 4 个 fp32: c0,c1 在第 lane/4 行, c2,c3 在其下 8 行
int my_row = row + lane_id / 4;
int my_col = col + (lane_id % 4) * 2;
half2 v;
v.x = __float2half(c_frag[i][j][0]);
v.y = __float2half(c_frag[i][j][1]);
*reinterpret_cast<half2*>(&C[my_row * N + my_col]) = v;
v.x = __float2half(c_frag[i][j][2]);
v.y = __float2half(c_frag[i][j][3]);
*reinterpret_cast<half2*>(&C[(my_row + 8) * N + my_col]) = v;
}
}
这段代码省略了细节(cp.async 与 ldmatrix 辅助函数、地址计算、swizzle 实现、边界处理),但骨架就是这样。补上这几个辅助函数后,用 nvcc 13.4 以 sm_80 / sm_90a 编译,SASS 里每个 k_step 每 warp 是 32 条 HMMA.16816.F32、8 条 LDSM.16.M88.4(A)和 8 条 LDSM.16.MT88.2(B,带转置),与 4×4 个 mma × 2 个 kk 步对得上(不同 CUDA 版本可能不同)。完整可工作的代码在 CUTLASS 中:cutlass-4.7.0/include/cutlass/gemm/threadblock/mma_pipelined.h(GitHub 上的同一份)。注意这是 CUTLASS 2.x 风格的 threadblock 层 API,在 4.7.0 里仍然保留着,读起来比 3.x 的 CollectiveMma 直白得多,适合对照本节骨架;但 Hopper 上真正在跑的是 3.x 那一套,第 13 章讲。
天花板对比(H100 SXM5 官方规格;实测可达比例引自已发表的微基准):
FP32 SIMT 通路峰值(第 10、11 章): 67 TFLOPS
FP16 Tensor Core 稠密峰值(本章): 989 TFLOPS (约 15 倍)
mma.sync m16n8k16 实测可达(H800,见下文): 约 64.9% 峰值 ≈ 640 TFLOPS
这里换的是赛道,不是工艺:同样是"写得不算特别精细"的手写实现,天花板从 67 TFLOPS 跳到 989 TFLOPS,靠的全是 Tensor Core。要说明的是,Hopper 上 mma.sync 这条指令本身就跑不满峰值——Luo et al. 2024 在 H800 上实测 m16n8k16(FP16 输入、FP32 累加)只到理论峰值的 64.9%(论文 Table VII),全部 mma 形状平均 62.9%;本节这个单缓冲骨架连这个上限都未必摸得到。离峰值剩下的差距,首先要靠换成 Hopper 的 WGMMA(12.6 节),其次才是 double buffer 流水深度、CUTLASS 级别的细致 fragment 调度、PTX 微优化——那是第 13 章的话题。
12.6 Hopper 升级:WGMMA
Hopper 引入 WGMMA(Warp-Group MMA)后,mma 指令的粒度从 warp-level 提升到 warp-group-level:
mma.sync.m16n8k16: 16×8×16 = 2048 次乘加 = 4096 FLOPs, warp 级 (32 线程)
wgmma.mma_async.m64n128k16: 64×128×16 = 131072 次乘加 = 262144 FLOPs, warp-group 级 (128 线程)
(注意"乘加数 × 2 = FLOPs"这个换算,12.2 节算 m16n8k16 的 4096 FLOPs 用的是同一个口径。)
操作数来源也变了:wgmma 的 B 必须放在 SMEM 里,通过一个 64 位的矩阵描述符(descriptor)交给指令;A 既可以同样走 SMEM 描述符,也可以放在寄存器里(CUTLASS 里分别对应 _SS 和 _RS 两族 atom,如 cutlass-4.7.0/include/cute/arch/mma_sm90_gmma.hpp:1632 的 MMA_64x128x16_F32F16F16_SS)。也就是说 B 不再需要 ldmatrix 搬进寄存器。用 nvcc 13.4 以 sm_90a 编译这个 atom,SASS 里是一条 HGMMA.64x128x16.F32。必须用 -arch=sm_90a:只写 sm_90 时,直接写的 wgmma 内联汇编会被 ptxas 拒绝,而 CuTe 这个 atom 会走进 CUTE_INVALID_CONTROL_PATH 的报错分支(cutlass-4.7.0/include/cute/config.hpp:158),SASS 里一条 HGMMA 都没有(nvcc 13.4 编译所得,不同 CUDA 版本可能不同)。
WGMMA 单条指令的计算量是 mma.sync 的 64 倍。不过一条 wgmma 由 warp-group 里的 4 个 warp 各自发射,折到每个 warp 调度器,达到同样吞吐所需的指令发射次数少 16 倍——取指与调度压力大幅下降,更易跑满 Tensor Core。
WGMMA 还是异步指令:
wgmma.fence.sync.aligned; // 声明累加器寄存器即将交给异步 wgmma(第一条 wgmma 之前必须有)
wgmma.mma_async.sync.aligned...; // 发起异步矩阵乘
wgmma.commit_group.sync.aligned; // 把已发出的 wgmma 打包成一个 group
... 做别的事 ...
wgmma.wait_group.sync.aligned 0; // 等到未完成的 group 不超过 N 个(这里 N=0,即全部完成)
这四条助记符可以在 CUTLASS 里逐条对上:
cutlass-4.7.0/include/cute/arch/mma_sm90_gmma.hpp:53(fence)、:67(wait_group,模板参数就是那个N)、:80(commit_group)。wgmma.fence最容易被漏掉——它不是可选的性能提示,而是告诉硬件"累加器寄存器从现在起由异步的 wgmma 读写",漏掉就是数据竞争。
发起 wgmma 之后 warp 可以继续做别的事(比如 TMA 加载下一个 tile),等需要结果时再同步。这是 Hopper GEMM 性能跃迁的核心机制——算和拷贝真正流水起来。
完整的 Hopper WGMMA GEMM 框架第 13 章 CUTLASS 部分会展开,第 17 章 FA2 SOTA 会用到。
12.7 这一章的小结与下一章
Tensor Core 是 GEMM 性能跃迁的关键:
- mma.sync 是矩阵级指令:单条指令算 16×8×16 = 2048 次乘加 = 4096 FLOPs。
- ldmatrix 是配套的矩阵 load 指令:把 SMEM 中的 16×16 子块加载到 fragment。
- Fragment 是分布式寄存器布局:32 lane 协作持有矩阵。
- SMEM Swizzle 防 bank conflict:CUTLASS 的标准 swizzle layout 解决了相邻行同列的 conflict 问题。
- WGMMA 是 Hopper 的升级:单条指令 64 倍计算量(折到每个 warp 调度器是 16 倍的发射压力下降)+ 操作数直接取自 SMEM + 异步执行。
到这里,读者已经能写出一个 mma.sync 版的 HGEMM;在 Hopper 上它的天花板约是峰值的六成多,再往上要靠 WGMMA。下一步是把这套手艺工业化——CUTLASS 把所有这些技巧抽象成可组合的 C++ 模板,让 NVIDIA 和工业界能用统一的工具构建各种 GEMM 变体(包括 FA2 内的 QK^T 和 PV)。
第 13 章我们剖析 CUTLASS 3.x 的设计哲学——CollectiveOp、CuTe Layout、Hopper Kernel Schedule。读完第 13 章读者会理解为什么 CUTLASS 的代码"看起来很复杂但实际上很优雅",并学会怎么读 CUTLASS 源码。
本章动手练习:
- 实现一个最简版 mma.sync HGEMM(小尺寸 M=N=K=64),亲手写 inline PTX,体验 fragment 布局。
- 阅读 CUTLASS 的
cutlass-4.7.0/include/cutlass/gemm/threadblock/mma_pipelined.h,看双缓冲 + ldmatrix + mma 是怎么组装的。- 在 H100 上跑 cuBLAS HGEMM 和你的版本,用 Nsight Compute 看
sm__pipe_tensor_cycles_active系列指标(如.avg.pct_of_peak_sustained_active,衡量 Tensor Core 管线忙碌的周期占比;sm__inst_executed_pipe_tensor数的是指令条数,不等于利用率,指标名以所用 ncu 版本的--query-metrics为准)——你的 kernel Tensor Core 利用率是多少?