CUDA 算子工程:手写 FlashAttention v2 之路

第 12 章 Tensor Core GEMM:mma.sync 与 ldmatrix

作者 杨艺韬 · 3,959 字 · 发布于 · 更新于

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:四组寄存器(不是矩阵指针!)

每条指令的浮点操作数:

16×8×16×2=4096 FLOPs16 \times 8 \times 16 \times 2 = 4096 \text{ FLOPs}

反推一下 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 节。)

4096×132×1.83 GHz≈989 TFLOPs4096 \times 132 \times 1.83\,\text{GHz} \approx 989\ \text{TFLOPs}

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 ll 持有的 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 ll 指向第 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 指定"不打散的基本单元"占多少位(单元大小 2MBase2^{\text{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 性能跃迁的关键:

  1. mma.sync 是矩阵级指令:单条指令算 16×8×16 = 2048 次乘加 = 4096 FLOPs。
  2. ldmatrix 是配套的矩阵 load 指令:把 SMEM 中的 16×16 子块加载到 fragment。
  3. Fragment 是分布式寄存器布局:32 lane 协作持有矩阵。
  4. SMEM Swizzle 防 bank conflict:CUTLASS 的标准 swizzle layout 解决了相邻行同列的 conflict 问题。
  5. 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 源码。

本章动手练习:

  1. 实现一个最简版 mma.sync HGEMM(小尺寸 M=N=K=64),亲手写 inline PTX,体验 fragment 布局。
  2. 阅读 CUTLASS 的 cutlass-4.7.0/include/cutlass/gemm/threadblock/mma_pipelined.h,看双缓冲 + ldmatrix + mma 是怎么组装的。
  3. 在 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 利用率是多少?