GEMM 工程教程:从三重循环到 BF16 / FP8 Tensor Core

这是一份工程视角的 GEMM 教程。我们先从最普通的矩阵乘法写法出发,解释数据复用、tile、片上存储、BF16/FP8 matrix-core 路径和 profiling。硬件细节以前十二章的 CUDA/NVIDIA 主线展开,最后再对照 Triton 与 AMD ROCm 的实现和工具。

BF16 主线 GPU Kernel Tiling Tensor Core FP8 扩展

0. 学习地图

GEMM 的优化可以看成一个不断减少“无效数据移动”、提高“计算单元饱和度”的过程。

三重循环
Naive GPU
Block Tiling
Shared + Register
Tensor Core Pipeline

如果只记一条主线:先让每个输出 tile 复用 A/B 数据,再让每个线程保留多个 accumulator,最后让 Tensor Core 在数据搬运期间持续工作。

Global / HBM容量大、延迟高、带宽珍贵。GEMM 优化的第一目标是减少重复读,并让读写连续、对齐、可合并。
Shared MemoryCTA 内共享,延迟低很多,但容量小,并且有 bank conflict、padding、同步和 occupancy 成本。
Register每个线程私有,最快,但数量有限。accumulator 通常放这里,寄存器多了会降低活跃 warp 数。

所以一个成熟 GEMM kernel 的数据流一般是:HBM 读入 A/B tile,到 shared memory 做 CTA 级复用,再进入 warp fragment / register,最后用 Tensor Core 累加到 register accumulator。

0.1 阅读路径

这份文档同时服务新手和已经有 CUDA 基础的读者。建议先按自己的目标选路径,不需要第一次就从头读到尾。

读者路径阅读顺序前置条件可以先跳过
新手0 → 1 → 2 → 3 → 3.3 demo → 4 → 5 → 6 → 7知道矩阵乘法和基本 CUDA thread/block 概念第一次先跳过 10、11、12
有 CUDA 基础0 → 4 → 4.6 → 5 → 6 → 7 → 8 → 10 → 11 → 12 → 13能读 CUDA kernel,知道 shared memory 和 warp1-3 可快速扫过
只想调性能3 → 3.3 demo → 4.4 → 4.6 → 8 → 10 → 13 → 14已经有一个能跑的 GEMM 或 matmul 调用FP8 和 CuTe 先不用看;profile 时可用 --skip-check 避免 CPU reference 污染时间线
关注 FP87 → 9 → 10 → 11 → 12 → 13 → 14理解 BF16 Tensor Core 和 epilogueCPU/naive 部分只作背景

0.2 读前准备

这份文档的目标不是让你背 CUDA 名词,而是让你能看懂一个生产级 GEMM kernel 的性能结构,并能判断一个优化点为什么有效。

先看代码块标签再读代码:Concept sketch / Teaching kernel 用于解释机制,不承诺可直接编译;Production command / Production profiling command 可在匹配的环境中运行;Production reference 使用真实工程 API,但仍依赖对应的库版本和 GPU 架构。标签后的补充文字只说明语言或场景。
读完应该掌握
  • 能从三重循环推导出 block/shared/register/tensor core tiling。
  • 能解释 BM/BN/BK、stage、occupancy、bank conflict 的取舍。
  • 能区分 BF16 GEMM 和 FP8 GEMM 的数值路径。
  • 能读 CUTLASS/CuTe GEMM 示例,知道 mainloop 和 epilogue 在哪里。
默认前置知识
  • 知道 CUDA thread、warp、CTA/block、SM 的基本含义。
  • 知道 global memory、shared memory、register 的大致层级。
  • 能读 C++/Python 代码,不要求已经写过 CUTLASS kernel。
  • 知道 BF16/FP16/FP8 是低精度浮点格式。

0.2.1 术语速查

术语本文里的含义容易混淆点
CTACUDA thread block,一个 block 负责一个或多个输出 tile很多资料会把 CTA 和 block 混用
BM/BN/BKCTA 级 GEMM tile 的 M/N/K 尺寸不是固定经验值,要结合资源和 shape profile
Warp tileCTA tile 内分给单个 warp 的子 tile不等于 MMA 指令一次覆盖的 tile
MMA tileTensor Core 指令级小矩阵形状由架构和指令族决定
Epilogueaccumulator 完成后到输出写回之间的融合逻辑FP8 下 scale/amax/requantize 常在这里变成瓶颈
CuTeCUTLASS 里的 layout/tensor/MMA/copy DSL既有 C++ CuTe,也有 Python CuTe DSL

1. 数学定义和形状

标准 GEMM 写作:

Concept sketch
C = alpha * A * B + beta * C

教学里先忽略 alphabeta,关注核心乘法:

Concept sketch
A: [M, K]
B: [K, N]
C: [M, N]

C[m, n] = sum(A[m, k] * B[k, n] for k in 0..K-1)
M输出矩阵的行数,常对应 token 数、batch 维或 row tile。
N输出矩阵的列数,常对应 hidden/out feature 维。
Kreduce 维,每个输出元素都要沿 K 做点积。

计算量近似为 2 * M * N * K FLOPs,因为一次 multiply-add 通常算 2 FLOPs。

1.1 为什么 GEMM 值得专门优化

GEMM 的计算量随 M*N*K 增长,但输入输出数据量只随 M*K + K*N + M*N 增长。如果 tile 设计合理,同一份 A/B 数据可以被重复使用很多次,kernel 就更容易接近计算峰值。

Arithmetic intensity ≈ FLOPs / bytes moved

这也是 roofline 模型的核心:当 arithmetic intensity 低时,瓶颈是内存带宽;当它足够高时,瓶颈才可能变成 Tensor Core / CUDA core 吞吐。

1.2 Layout 先于 kernel

矩阵在内存里通常是 row-major 或 column-major。对 GPU 来说,线程束访问连续地址时才能合并成更少的 memory transaction。一个数学上等价的转置或 layout 选择,可能直接决定 global load 是否 coalesced。

好访问
warp 中相邻 lane 访问相邻地址,并尽量完整利用当前架构的 32B global-memory transaction;单个线程的 vectorized load 还要满足自身宽度对齐,例如 16B load 至少对齐到 16B。GEMM tile 起点常进一步按 128B 组织,但这不是 warp 的通用硬性对齐要求。
坏访问
warp 中 lane 跨 stride 访问,或者每个 lane 只拿很小的不连续片段,transaction 多且有效带宽低。

2. 一般写法:三重循环

最直接的写法如下:

Concept sketch
for m in range(M):
  for n in range(N):
    acc = 0.0
    for k in range(K):
      acc += A[m, k] * B[k, n]
    C[m, n] = acc
1它正确,但没有利用数据复用

同一个 A[m, k] 会被多个 n 使用,同一个 B[k, n] 会被多个 m 使用。朴素循环没有显式把这种复用组织起来。

2它也没有匹配 GPU 的执行模型

GPU 需要大量线程并行,且访问模式要连续、对齐、可合并。三重循环只是算法表达,不是高性能实现。

2.1 循环顺序会改变缓存行为

即使在 CPU 上,m-n-km-k-nk-m-n 的性能也可能差很多,因为 B 的访问是连续还是跨 stride,会影响 cache line 利用率。GPU 上这个问题更尖锐:warp 是否 coalesced 比单个线程的局部性更重要。

Concept sketch
// row-major B[K, N] 下,n 连续时 B[k, n] 连续
for m:
  for k:
    a = A[m, k]
    for n:
      C[m, n] += a * B[k, n]

这段 CPU 伪代码已经开始表达“复用 A、连续读 B”的思想。GPU kernel 的 tiling 只是把这个思想推进到 CTA、warp、register 三个层级。

3. Naive GPU Kernel:一个线程算一个输出元素

第一版 GPU 化通常是:每个 thread 负责一个 C[m,n]

Concept sketch
__global__ void matmul_naive_bf16(A, B, C, M, N, K) {
  int m = blockIdx.y * blockDim.y + threadIdx.y;
  int n = blockIdx.x * blockDim.x + threadIdx.x;
  if (m >= M || n >= N) return;

  float acc = 0.0f;
  for (int k = 0; k < K; ++k) {
    acc += float(A[m * K + k]) * float(B[k * N + n]);
  }
  C[m * N + n] = bf16(acc);
}

这比 CPU 三重循环更并行,但离高性能仍很远。主要问题通常不是乘加数量,而是 A/B 数据缺少 CTA 内复用,导致相对计算量产生了过多 HBM 访问。

Naive GPU kernel 的核心问题:每个 thread 独立从 global memory 扫 A 的一行和 B 的一列,block 内大量重复读取,数据复用几乎丢失。

3.1 Naive kernel 的访存问题

假设一个 block 有很多线程同时计算相邻的 C[m,n]。它们在同一个 k 上往往会重复读取同一个 A[m,k],并读取一段连续的 B[k,n]。B 的访问可能还能合并,A 的读取却会在多个线程之间重复发生。

访问对象Naive 行为后续优化目标
A[m,k]同一行的元素被多个 n 重复读取放入 shared,让 CTA 内多个列复用
B[k,n]相邻 n 可能连续,但跨 m 会重复读取放入 shared,让 CTA 内多行复用
C[m,n]每个线程只维护一个 acc每个线程或 warp 维护多个 acc,提高计算密度

3.2 第一个性能判断

如果 profile 里看到 global load throughput 很高,但 Tensor Core 或 FMA 利用率很低,通常说明 kernel 被内存喂不饱。这时不要急着调 block size,先问:A/B 是否被重复从 HBM 读了太多次?

3.3 从零能跑的 CUDA GEMM demo

只看伪代码很难建立性能直觉。这里提供一个单文件 CUDA demo:同一个 FP32 GEMM 同时实现 naive kernel 和 shared-memory tiled kernel,并用 CUDA event 计时。它不是生产级 Tensor Core kernel,而是用来亲眼看到“重复读 HBM”和“CTA 内复用 A/B tile”的差别。

Production command
nvcc -O3 -std=c++17 -arch=sm_80 gemm-demo.cu -o gemm-demo
./gemm-demo
./gemm-demo 1024 1024 1024 50

-arch=sm_80 只是示例,应该替换成你的 GPU 架构。这个 demo 只依赖 CUDA runtime,不依赖 CUTLASS、Triton 或 cuBLAS。

Concept sketch
__global__ void matmul_tiled_kernel(const float* A, const float* B, float* C,
                                    int M, int N, int K) {
  __shared__ float As[16][16];
  __shared__ float Bs[16][16];
  int n = blockIdx.x * 16 + threadIdx.x;
  int m = blockIdx.y * 16 + threadIdx.y;
  float acc = 0.0f;

  for (int k0 = 0; k0 < K; k0 += 16) {
    As[threadIdx.y][threadIdx.x] = in_bounds_A ? A[m * K + k0 + threadIdx.x] : 0.0f;
    Bs[threadIdx.y][threadIdx.x] = in_bounds_B ? B[(k0 + threadIdx.y) * N + n] : 0.0f;
    __syncthreads();
    for (int kk = 0; kk < 16; ++kk) acc += As[threadIdx.y][kk] * Bs[kk][threadIdx.x];
    __syncthreads();
  }
  if (m < M && n < N) C[m * N + n] = acc;
}
查看完整 gemm-demo.cu 源码
Production reference · full source
// gemm-demo.cu
// Minimal CUDA GEMM lab: naive one-thread-per-output vs shared-memory tiling.
//
// Build on a CUDA machine:
//   nvcc -O3 -std=c++17 -arch=sm_80 gemm-demo.cu -o gemm-demo
//
// Run:
//   ./gemm-demo
//   ./gemm-demo 1024 1024 1024 50
//   ./gemm-demo 1024 1024 1024 50 --skip-check

#include <cuda_runtime.h>

#include <algorithm>
#include <cmath>
#include <cstdio>
#include <cstdlib>
#include <cstring>
#include <iostream>
#include <vector>

#define CUDA_CHECK(call)                                                     \
  do {                                                                       \
    cudaError_t err__ = (call);                                               \
    if (err__ != cudaSuccess) {                                               \
      std::cerr << "CUDA error at " << __FILE__ << ":" << __LINE__ << ": " \
                << cudaGetErrorString(err__) << std::endl;                   \
      std::exit(EXIT_FAILURE);                                                \
    }                                                                        \
  } while (0)

constexpr int TILE = 16;

__global__ void matmul_naive_kernel(const float* __restrict__ A,
                                    const float* __restrict__ B,
                                    float* __restrict__ C,
                                    int M,
                                    int N,
                                    int K) {
  int n = blockIdx.x * blockDim.x + threadIdx.x;
  int m = blockIdx.y * blockDim.y + threadIdx.y;

  if (m >= M || n >= N) {
    return;
  }

  float acc = 0.0f;
  for (int k = 0; k < K; ++k) {
    acc += A[m * K + k] * B[k * N + n];
  }
  C[m * N + n] = acc;
}

__global__ void matmul_tiled_kernel(const float* __restrict__ A,
                                    const float* __restrict__ B,
                                    float* __restrict__ C,
                                    int M,
                                    int N,
                                    int K) {
  __shared__ float As[TILE][TILE];
  __shared__ float Bs[TILE][TILE];

  int tx = threadIdx.x;
  int ty = threadIdx.y;
  int n = blockIdx.x * TILE + tx;
  int m = blockIdx.y * TILE + ty;

  float acc = 0.0f;
  int k_tiles = (K + TILE - 1) / TILE;

  for (int tile = 0; tile < k_tiles; ++tile) {
    int a_col = tile * TILE + tx;
    int b_row = tile * TILE + ty;

    As[ty][tx] = (m < M && a_col < K) ? A[m * K + a_col] : 0.0f;
    Bs[ty][tx] = (b_row < K && n < N) ? B[b_row * N + n] : 0.0f;
    __syncthreads();

    #pragma unroll
    for (int kk = 0; kk < TILE; ++kk) {
      acc += As[ty][kk] * Bs[kk][tx];
    }
    __syncthreads();
  }

  if (m < M && n < N) {
    C[m * N + n] = acc;
  }
}

void init_matrix(std::vector<float>& x, int rows, int cols, int salt) {
  for (int r = 0; r < rows; ++r) {
    for (int c = 0; c < cols; ++c) {
      int v = (r * 17 + c * 13 + salt) % 97;
      x[r * cols + c] = static_cast<float>(v - 48) / 48.0f;
    }
  }
}

void cpu_reference(const std::vector<float>& A,
                   const std::vector<float>& B,
                   std::vector<float>& C,
                   int M,
                   int N,
                   int K) {
  for (int m = 0; m < M; ++m) {
    for (int n = 0; n < N; ++n) {
      float acc = 0.0f;
      for (int k = 0; k < K; ++k) {
        acc += A[m * K + k] * B[k * N + n];
      }
      C[m * N + n] = acc;
    }
  }
}

float max_abs_error(const std::vector<float>& ref, const std::vector<float>& got) {
  float err = 0.0f;
  for (size_t i = 0; i < ref.size(); ++i) {
    err = std::max(err, std::abs(ref[i] - got[i]));
  }
  return err;
}

template <typename Kernel>
float time_kernel(Kernel kernel,
                  dim3 grid,
                  dim3 block,
                  const float* d_A,
                  const float* d_B,
                  float* d_C,
                  int M,
                  int N,
                  int K,
                  int iters) {
  cudaEvent_t start;
  cudaEvent_t stop;
  CUDA_CHECK(cudaEventCreate(&start));
  CUDA_CHECK(cudaEventCreate(&stop));

  kernel<<<grid, block>>>(d_A, d_B, d_C, M, N, K);
  CUDA_CHECK(cudaGetLastError());
  CUDA_CHECK(cudaDeviceSynchronize());

  CUDA_CHECK(cudaEventRecord(start));
  for (int i = 0; i < iters; ++i) {
    kernel<<<grid, block>>>(d_A, d_B, d_C, M, N, K);
  }
  CUDA_CHECK(cudaEventRecord(stop));
  CUDA_CHECK(cudaEventSynchronize(stop));
  CUDA_CHECK(cudaGetLastError());

  float ms = 0.0f;
  CUDA_CHECK(cudaEventElapsedTime(&ms, start, stop));
  CUDA_CHECK(cudaEventDestroy(start));
  CUDA_CHECK(cudaEventDestroy(stop));
  return ms / static_cast<float>(iters);
}

double gflops(int M, int N, int K, float ms) {
  double flops = 2.0 * static_cast<double>(M) * static_cast<double>(N) * static_cast<double>(K);
  return flops / (static_cast<double>(ms) * 1.0e6);
}

int parse_arg(char** argv, int index, int fallback) {
  if (argv[index] == nullptr) {
    return fallback;
  }
  int value = std::atoi(argv[index]);
  return value > 0 ? value : fallback;
}

bool has_flag(int argc, char** argv, const char* flag) {
  for (int i = 1; i < argc; ++i) {
    if (std::strcmp(argv[i], flag) == 0) {
      return true;
    }
  }
  return false;
}

int main(int argc, char** argv) {
  int M = argc > 1 ? parse_arg(argv, 1, 512) : 512;
  int N = argc > 2 ? parse_arg(argv, 2, 512) : 512;
  int K = argc > 3 ? parse_arg(argv, 3, 512) : 512;
  int iters = argc > 4 ? parse_arg(argv, 4, 20) : 20;
  bool skip_check = has_flag(argc, argv, "--skip-check");

  std::cout << "GEMM demo shape: M=" << M << " N=" << N << " K=" << K
            << " iters=" << iters << std::endl;
  if (skip_check) {
    std::cout << "correctness: skipped (--skip-check)" << std::endl;
  }

  std::vector<float> h_A(static_cast<size_t>(M) * K);
  std::vector<float> h_B(static_cast<size_t>(K) * N);
  std::vector<float> h_ref(skip_check ? 0 : static_cast<size_t>(M) * N);
  std::vector<float> h_naive(static_cast<size_t>(M) * N);
  std::vector<float> h_tiled(static_cast<size_t>(M) * N);

  init_matrix(h_A, M, K, 7);
  init_matrix(h_B, K, N, 19);
  if (!skip_check) {
    cpu_reference(h_A, h_B, h_ref, M, N, K);
  }

  float* d_A = nullptr;
  float* d_B = nullptr;
  float* d_C = nullptr;
  size_t bytes_A = h_A.size() * sizeof(float);
  size_t bytes_B = h_B.size() * sizeof(float);
  size_t bytes_C = static_cast<size_t>(M) * N * sizeof(float);

  CUDA_CHECK(cudaMalloc(&d_A, bytes_A));
  CUDA_CHECK(cudaMalloc(&d_B, bytes_B));
  CUDA_CHECK(cudaMalloc(&d_C, bytes_C));
  CUDA_CHECK(cudaMemcpy(d_A, h_A.data(), bytes_A, cudaMemcpyHostToDevice));
  CUDA_CHECK(cudaMemcpy(d_B, h_B.data(), bytes_B, cudaMemcpyHostToDevice));

  dim3 block(TILE, TILE);
  dim3 grid((N + TILE - 1) / TILE, (M + TILE - 1) / TILE);

  float naive_ms = time_kernel(matmul_naive_kernel, grid, block, d_A, d_B, d_C, M, N, K, iters);
  CUDA_CHECK(cudaMemcpy(h_naive.data(), d_C, bytes_C, cudaMemcpyDeviceToHost));
  float naive_err = skip_check ? NAN : max_abs_error(h_ref, h_naive);

  float tiled_ms = time_kernel(matmul_tiled_kernel, grid, block, d_A, d_B, d_C, M, N, K, iters);
  CUDA_CHECK(cudaMemcpy(h_tiled.data(), d_C, bytes_C, cudaMemcpyDeviceToHost));
  float tiled_err = skip_check ? NAN : max_abs_error(h_ref, h_tiled);

  std::cout << "kernel, avg_ms, GFLOP/s, max_abs_error" << std::endl;
  std::cout << "naive, " << naive_ms << ", " << gflops(M, N, K, naive_ms)
            << ", " << naive_err << std::endl;
  std::cout << "tiled, " << tiled_ms << ", " << gflops(M, N, K, tiled_ms)
            << ", " << tiled_err << std::endl;

  CUDA_CHECK(cudaFree(d_A));
  CUDA_CHECK(cudaFree(d_B));
  CUDA_CHECK(cudaFree(d_C));
  return 0;
}

运行输出会包含 correctness、max absolute error、平均耗时和 GFLOP/s。做 profiling 时可以加 --skip-check 跳过 CPU reference,避免时间线里混入 CPU 端三重循环。不要把 demo 数字当成硬件上限:naive/tiled 只是教学 baseline,真正的 BF16/FP8 高性能路径还要进入 Tensor Core、pipeline、epilogue 和库级调参。

Production profiling command
nsys profile -o gemm_demo_timeline --trace=cuda,nvtx,osrt ./gemm-demo 1024 1024 1024 20 --skip-check
ncu --set full --kernel-name regex:matmul_ ./gemm-demo 1024 1024 1024 20 --skip-check
如果 tiled kernel 明显快于 naive kernel,主要原因通常不是少做了计算,而是同一份 A/B tile 被 CTA 内多个线程复用,HBM traffic 相对下降。后面的 §4-§8 会把这个思路推进到 Tensor Core 级别。

4. Block Tiling:让一个 CTA 负责一个 C tile

优化的第一步是把 C 切成 tile。比如一个 CTA 负责 128 x 128 的 C tile。

Concept sketch
for block_m in range(0, M, BM):
  for block_n in range(0, N, BN):
    C_tile[BM, BN] = 0
    for block_k in range(0, K, BK):
      A_tile = A[block_m:block_m+BM, block_k:block_k+BK]
      B_tile = B[block_k:block_k+BK, block_n:block_n+BN]
      C_tile += A_tile @ B_tile

这里的关键变化是:A_tile 会被同一个 C tile 的多个列复用,B_tile 会被多个行复用。

参数典型含义选择影响
BMC tile 的行数影响 CTA 数量和 A 的复用
BNC tile 的列数影响 B 的复用和输出写回
BKK 维每次处理的块大小影响 shared memory 占用和 pipeline 粒度

4.1 CTA tile、warp tile、MMA tile 是三层结构

不要把 BM/BN/BK 理解成一个孤立参数。实际 kernel 通常会把一个 CTA tile 再切给多个 warp,每个 warp 再执行多个 MMA tile。

Concept sketch
CTA tile:   128 x 128 x 64
Warp tile:   64 x  64 x 64   // 一个 CTA 里多个 warp 分工
MMA tile:    16 x   8 x 16   // 具体 Tensor Core 指令形状示意
层级谁负责和下一层的关系
CTA tile一个 CUDA block / CTA切成多个 warp tile,决定 shared memory 中 A/B tile 的总体大小
Warp tile一个 warp 或 warp group由多次 MMA tile 拼出来,决定每个 warp 维护多少 accumulator fragment
MMA tile一次或一组 Tensor Core 指令最小计算原语,形状由架构和 dtype 指令族决定

CTA tile 太小,数据复用少,调度开销相对变大;CTA tile 太大,shared memory 和 register 使用量会上升,occupancy 下降。高性能 GEMM 本质上是在这些约束之间找平衡点。

4.2 为什么 BK 很关键

BK 决定每个 mainloop stage 搬多少 K 维数据。BK 大会增加每次搬运的数据和 shared memory 占用,也可能让单个 stage 的计算更饱满;BK 小则 pipeline 粒度细,但 loop overhead 和同步占比可能更高。

Shared bytes per CTA ≈ (BM * BK + BK * BN) * bytes_per_element * stages

例如 BF16 每个元素 2 bytes,BM=128, BN=128, BK=64, stages=2 时,只存 A/B tile 就大约需要 (128*64 + 64*128)*2*2 = 65536 bytes shared memory,还没算 padding 或特殊 layout。

4.3 BM / BN / BK 怎么量化选择

这确实是 GEMM 最难的参数之一,因为它不是单目标优化。一个 tile shape 同时决定数据复用、CTA 数量、shared memory 占用、寄存器数量、tail 浪费、warp 分工和 Tensor Core 指令排布。

比较实用的方法不是“背一个固定组合”,而是先用资源约束过滤候选,再用 profile 选择。

约束近似判断太大时的问题
shared memory(BM*BK + BK*BN) * bytes * stagesCTA 驻留数下降,occupancy 降低
accumulator 寄存器warp_tile_m * warp_tile_n 拆到每个 lane寄存器溢出,local memory 出现
CTA 数量ceil(M/BM) * ceil(N/BN)并行度不足,尤其小 M 场景
tail 浪费M % BMN % BNK % BK大量线程处理无效元素
MMA 对齐BM/BN/BK 尽量是 MMA tile 的整数倍需要额外 mask 或 fallback path

4.4 一个可执行的选型流程

  1. 先根据 dtype 和架构选 MMA 指令族,例如 BF16 Tensor Core 的基本 MMA 形状。
  2. BM/BN/BK 都对齐到 MMA 形状和 vectorized load 粒度,避免主路径里到处都是 tail。
  3. 估算 shared memory:把 stages 算进去,再留出 padding / swizzle 的空间。
  4. 估算寄存器:accumulator fragment 是大头,再加 A/B fragment、地址、predicate、epilogue 临时变量。
  5. 检查 CTA 并行度:如果 ceil(M/BM) * ceil(N/BN) 太小,即使单 CTA 很快,整卡也可能吃不满。
  6. 用 profile 比较候选,观察 Tensor Core 利用率、HBM 带宽、shared conflict、occupancy 和 local memory。
Concept sketch
// 候选生成的思路,不是固定答案
for BM, BN in [(64, 64), (64, 128), (128, 64), (128, 128), (128, 256)]:
  for BK in [32, 64, 128]:
    if not mma_aligned(BM, BN, BK): continue
    if smem_bytes(BM, BN, BK, stages) > smem_budget: continue
    if estimated_registers(BM, BN, BK) > register_budget: continue
    benchmark(BM, BN, BK)

4.5 不同 shape 的倾向

场景tile 倾向原因
大 M、大 N、大 K较大的 BM/BN,例如 128x128 或更大候选数据复用充分,CTA 数量也足够
小 M、大 N,例如 decode较小 BM,或让一个 CTA 覆盖更多 N避免 M 维浪费,同时提高 CTA 数量
N 很大、M 中等BN 可以偏大,但要看 C 写回和寄存器B tile 复用好,但 accumulator 会变多
K 很大BK 和 stage 数要一起调需要覆盖加载延迟,但 shared 占用不能爆
MoE / grouped GEMM更小或更灵活的 tile每组 shape 不规则,tail 和调度开销更明显
经验值只能当候选起点。BM/BN/BK 的最终选择必须和具体 shape 分布、GPU 架构、dtype、epilogue 融合内容一起 profile。

4.6 Split-K:当 M/N 并行度不够时,拆 K 维

前面的 block tiling 默认一个 C_tile[BM, BN] 由一个 CTA 沿完整 K 维算完。这个设计在大 M、大 N 的 GEMM 上通常很好,因为 ceil(M/BM) * ceil(N/BN) 已经能提供足够多 CTA。但在小 M、大 K,或者 M/N tile 数量很少的场景,整张卡可能只有很少 CTA 可调度,Tensor Core 明明很快,SM 却吃不满。

Split-K 的想法是:不要让一个 CTA 独占完整 K 维,而是把 K 维切成多个 slice,让多个 CTA 同时计算同一个 C tile 的不同 partial sum,最后再把 partial sum reduce 成最终 C。

Concept sketch
// 普通 tiling:一个 CTA 算完整 K
for cta_m, cta_n:
  acc = 0
  for k0 in range(0, K, BK):
    acc += A[cta_m, k0] @ B[k0, cta_n]
  C[cta_m, cta_n] = acc

// Split-K:多个 CTA 分别算 K slice,然后再 reduce
for cta_m, cta_n, split_id:
  k_begin, k_end = split_range(split_id, K)
  partial = 0
  for k0 in range(k_begin, k_end, BK):
    partial += A[cta_m, k0] @ B[k0, cta_n]
  Workspace[split_id, cta_m, cta_n] = partial

reduce Workspace[:, cta_m, cta_n] -> C[cta_m, cta_n]

4.6.1 Split-K 解决的不是复用,而是并行度

Split-K 不会神奇地减少总计算量。它通常还会增加额外的 workspace 写读和 reduction 成本。它真正解决的是 CTA 数量不足:当 M/N tile 太少时,把 K 拆开可以制造更多独立 work,让更多 SM 同时工作。

场景普通 GEMM 的问题Split-K 的收益
小 M、大 K、大 NM 维 tile 很少,CTA 数不够沿 K 增加 CTA 数,提高设备并行度和 SM 覆盖率
batched / grouped GEMM 中某些 group 很小单个 group 吃不满 GPU给大 K 的 group 拆更多 split,减少尾部空转
单个输出 tile 计算很重少数 CTA 占用时间长,负载不均把长 K 分摊给多个 CTA,改善调度粒度

4.6.2 两种 reduction 路线

Split-K 的关键不是前半段 partial GEMM,而是后半段怎么合并 partial sums。常见有两条路线。

Workspace + second kernel

每个 split 把 partial C 写到 workspace,第二个 reduction kernel 再沿 split 维求和。

优点:逻辑清楚,数值顺序相对可控,适合 split 数较多或 C tile 较大。

代价:多一次 global memory 写读,多一次 kernel launch。

Atomic accumulate

每个 split 直接 atomic add 到最终 C。

优点:少一个 reduction kernel,代码直观。

代价:atomic contention、写回顺序不稳定、浮点结果可能不 deterministic。

Split-K 数据流
A/B tiles
K slice 0
K slice 1
K slice 2
Reduce to C

4.6.3 Split-K 的选择条件

Split-K 应该是一个 profile-driven 的开关,而不是默认打开。一个简单判断是先算普通 tiling 的 CTA 数:

num_cta = ceil(M / BM) * ceil(N / BN)

如果 num_cta 已经远大于 SM 数,Split-K 往往只会引入 reduction overhead;如果 num_cta 明显小于 SM 数,且 K 很大,Split-K 才可能值得尝试。

问题判断方式调参方向
SM 空闲但少数 CTA 执行很久nsys 先排除 launch gap;ncu 再看 SM throughput、grid waves 和 active warps增加 split 数,或减小 BM/BN 产生更多 CTA
HBM 写读暴涨workspace traffic 增加,L2/HBM store/load 上升降低 split 数,或换 atomic / fused reduction 策略
结果不稳定多次运行 low-bit 差异,尤其 atomic accumulate使用固定 reduction 顺序,或接受非 deterministic 误差边界
tail 变多K 不能均匀切分,最后一个 split 工作量小让 split 数匹配 K/BK,避免太细的 K slice

4.6.4 Split-K 和 Stream-K 不一样

有些库文档还会提到 Stream-K。它和 Split-K 都是在 K 维和 tile 调度上做文章,但目标更偏向负载均衡和持续填满 SM。Split-K 更容易理解为“把一个 C tile 的 K 维拆给多个 worker”;Stream-K 通常会把全局 GEMM 的 tile/K 工作以更流式的方式分配给 worker,减少尾部不均衡。学习时可以先掌握 Split-K,再把 Stream-K 看成更工程化的调度策略。

Split-K 的常见误用:看到 K 很大就打开。K 大只是必要条件之一;如果 M/N tile 已经足够多,Split-K 的 reduction 成本可能直接吃掉收益。

5. Shared Memory:把 A/B tile 搬到片上

Block tiling 只是逻辑复用。要真正减少 HBM 访问,通常需要把当前 A_tileB_tile 搬到 shared memory。

Concept sketch
// 伪代码
for k0 in range(0, K, BK):
  shared_A = load_global_to_shared(A_tile)
  shared_B = load_global_to_shared(B_tile)
  __syncthreads()

  for kk in range(BK):
    acc += shared_A[row, kk] * shared_B[kk, col]
  __syncthreads()

这样 block 内线程会从 shared memory 复用 A/B 数据,而不是每次都打到 HBM。

shared memory 不是免费午餐。要关注 bank conflict、对齐、padding、load vectorization,以及 shared memory 占用对 occupancy 的影响。

5.1 Coalesced global load

从 HBM 搬到 shared 时,首先要让 global memory load 合并。常见做法是让每个线程搬连续的 16 bytes,例如 float4int4 或等价 vectorized load。这里的 16B 是单线程指令宽度和自然对齐要求;32 个 lane 合起来会请求 512B,但这 512B 不会成为一个单独的 512B transaction。当前 CUDA 文档对 Compute Capability 6.0 及之后的通用解释,是硬件按覆盖到的 32B segment 数量合并请求。

Concept sketch
// 概念示意:每个线程搬 16B
int lane = threadIdx.x % 32;       // 当前线程在 warp 内的编号
int vec_id = block_thread_id;      // CTA 内连续的线性线程编号

const char* src = reinterpret_cast<const char*>(global_ptr) + vec_id * 16;
char* dst = reinterpret_cast<char*>(shared_ptr) + vec_id * 16;

// debug 时检查实际地址,而不只是检查 cudaMalloc 返回的基地址
assert((reinterpret_cast<uintptr_t>(src) & 0xF) == 0);
assert((reinterpret_cast<uintptr_t>(dst) & 0xF) == 0);

uint4 x = *reinterpret_cast<const uint4*>(src);  // 16B global load
*reinterpret_cast<uint4*>(dst) = x;              // 16B shared store

uint4 由 4 个 32-bit 无符号整数组成,因此 sizeof(uint4) = 4 * 4B = 16Bvec_id 每增加 1,地址增加 16B;只要起始地址 16B 对齐,连续线程的每次 vector load 就都满足 16B 对齐。上面的 lane 在简化代码里没有参与地址计算,它只是提示真实 kernel 通常还会用 lane id 设计 warp 内的数据分工。

addr(vec_id) = base_addr + vec_id * 16;16B aligned ⇔ addr(vec_id) % 16 == 0
reinterpret_cast<uint4*> 不会把一个未对齐地址“变对齐”。它只是要求编译器把该地址当作 uint4* 使用。真正要检查的是加入 row、leading dimension、tile offset 后的有效地址;不满足自然对齐时,应走 scalar/masked fallback,或者调整 layout 和 padding。

5.1.1 从矩阵下标判断 16B 对齐

对 row-major 矩阵,线程从 A[row, col] 开始搬运时,有效字节地址是:

effective_addr = base(A) + (row * lda + col) * sizeof(T)

如果 base(A) 已经 16B 对齐,那么只需继续判断 (row * lda + col) * sizeof(T) 是否为 16 的倍数。注意:设备分配返回的基地址通常已经满足较强对齐,但 subview、切片、tile 起点和非整齐的 leading dimension 仍然可能破坏有效地址的对齐。

数据类型元素大小一次 16B load元素 offset 条件
FP32 / INT324B4 个元素(row * lda + col) % 4 == 0
FP16 / BF162B8 个元素(row * lda + col) % 8 == 0
FP8 / INT81B16 个元素(row * lda + col) % 16 == 0
Concept sketch · BF16 alignment
// A 是 row-major BF16,base(A) 已满足 16B 对齐
const __nv_bfloat16* src = A + row * lda + k0;
bool aligned_16b =
    (reinterpret_cast<uintptr_t>(src) & 0xF) == 0;

// lda=128、k0 是 8 的倍数:任意 row 都满足 16B 对齐
// (row * 128 + k0) * 2B 是 16B 的倍数

// lda=130、row=1、k0=0:第二行行首没有 16B 对齐
// (1 * 130 + 0) * 2B = 260B,260 % 16 = 4

因此,高性能 kernel 往往会把 lda padding 到 vector width 的整数倍,让 mainloop 走对齐的 16B load;对矩阵开头、结尾或 K-tail 中不满足条件的元素,再使用 masked load、标量 load,或者单独的 fallback kernel。

5.1.2 16B、32B、128B 和 512B 分别是什么

先给结论:warp 在这段代码里总共请求 512B,但 512B 只是 payload,不是 transaction size,所以没有“warp 首地址必须 512B 对齐”的规则。对齐要求应该跟它所属的层级一起读。

单个 lane16Buint4 load 的数据宽度和自然对齐。每个线程的实际地址必须是 16 的倍数。
Warp payload32 × 16B = 512B32 个 lane 在一条指令里合计需要的数据量。它描述“要多少数据”,不描述“一次搬多少”。
Global transaction32B当前 CUDA 文档对 Compute Capability 6.0+ 的通用 coalescing 模型。warp 请求会拆成覆盖这些地址所需的若干 32B transaction。
Cache line / tile常见 128B可用于 L1/cache-line、旧架构 segment 或 GEMM tile/stage 的强化对齐,但不是 warp 的统一硬性要求。
5.1.2.1 为什么 512B payload 不要求 512B 对齐

一条 warp load 发出后,coalescer 会统计 32 个 lane 的地址覆盖了多少个 transaction segment,而不是把所有 lane 拼成一个 512B 原子 transaction。对于连续的 16B-per-lane 访问:

lane L 访问 [base + 16L, base + 16L + 15];warp 覆盖 [base, base + 511]
32B transaction 数 = floor((base + 511) / 32) - floor(base / 32) + 1

只要 base % 32 == 0,512B 连续范围就正好覆盖 512 / 32 = 16 个 transaction。把 base 从 128B 对齐进一步提高到 512B 对齐,并不会把 16 个 transaction 继续减少。

Warp 首地址满足的对齐访问范围32B transaction覆盖的 128B line
0x08016B / 32B / 128B;不是 512B0x080...0x27F164
0x20016B / 32B / 128B / 512B0x200...0x3FF164
0x02016B / 32B;不是 128B0x020...0x21F165
0x010只有 16B;不是 32B0x010...0x20F175
Case A:base = 0x080,128B 对齐但不是 512B 对齐
512B payload 正好占 16 个 32B segment;每 4 格组成一个 128B line,共 4 条 line。
T0
080-09F
T1
0A0-0BF
T2
0C0-0DF
T3
0E0-0FF
T4
100-11F
T5
120-13F
T6
140-15F
T7
160-17F
T8
180-19F
T9
1A0-1BF
T10
1C0-1DF
T11
1E0-1FF
T12
200-21F
T13
220-23F
T14
240-25F
T15
260-27F
16 × 32B transaction4 × 128B line0 个额外 segment
Case B:base = 0x010,只满足每线程 16B 对齐
lane 0 从一个 32B segment 的中间开始,warp 末尾也落在另一个 segment 的前半段,所以首尾各有半个 transaction 没被充分利用。
T0
000-01F
只用后16B
T1
020-03F
T2
040-05F
T3
060-07F
T4
080-09F
T5
0A0-0BF
T6
0C0-0DF
T7
0E0-0FF
T8
100-11F
T9
120-13F
T10
140-15F
T11
160-17F
T12
180-19F
T13
1A0-1BF
T14
1C0-1DF
T15
1E0-1FF
T16
200-21F
只用前16B
17 × 32B transaction5 × 128B line多覆盖 32B segment
因此答案是:0x080 虽然不是 512B 对齐,但已经落在 32B transaction 边界上,所以和 0x200 一样只需要 16 个 transaction。512B 对齐没有进一步减少请求数量。
5.1.2.2 那为什么 GEMM 里仍经常看到 128B 对齐

128B 属于更高一级的布局选择,而不是由 32 lanes × 16B 推导出来的 warp 对齐要求。它常见有三个原因:

  1. Cache-line / segment 边界:部分旧架构和文档用 128B L1 segment 解释 coalescing;即使当前数据访问单元按 32B 计算,128B 对齐也能让四个连续 32B sector 自然落入同一条 128B line。
  2. Tile 和 pipeline stage 易于管理:GEMM 会反复搬运 A/B tile。让每行或每个 stage 从 128B 边界开始,地址计算、cache-line 利用、shared layout 和异步流水线通常更规整。
  3. 历史 API 性能指导:部分旧版 cuda::memcpy_async 文档曾建议 global/shared 都按 128B 对齐。当前 Cooperative Groups 文档则推荐 16B,因此这不是跨版本、跨 API 的统一硬规则。

base=0x020 是一个有用的反例:它已经 32B 对齐,因此仍是最少的 16 个 32B transaction;但 512B 范围会跨 5 条 128B line。也就是说,128B 对齐可能改善 cache-line/tile 组织,却不等于它是 coalescing 的最低门槛。

5.1.2.3 Async copy 应该检查哪种对齐

低层 cp.async 家族按线程支持 4B、8B、16B copy,指针至少要满足所选 copy size 的自然对齐。对于本节每线程 16B 的写法,首先必须证明 global source 和 shared destination 都满足 16B 对齐。

必须满足:src_lane % 16 == 0 且 dst_lane % 16 == 0

至于 tile 是否进一步 32B、128B 或更高对齐,要看具体 CUDA 版本、API、架构以及使用的是 cp.async、Cooperative Groups、TMA 还是其他 copy atom。不要把某个 API 某一版文档的“最佳性能建议”改写成所有 warp load 的硬件要求。

5.1.2.4 实际写 kernel 时按这个顺序判断
  1. 检查每个 lane 的实际地址:16B vector load 必须满足 addr % 16 == 0
  2. 检查 lane 间地址是否连续,确保请求能被 coalescer 合并。
  3. 在当前 32B transaction 模型下,检查 warp 首地址是否满足 base % 32 == 0,避免首尾多覆盖一个 segment。
  4. 把 128B 当作常见的 tile/cache-line 强化对齐,再通过 profile 判断它是否有额外收益。
  5. 只有具体 copy API、TMA descriptor 或自定义 allocator 明确要求时,才继续要求 256B、512B 等更强对齐;不要从 warp payload 大小反推对齐要求。

5.2 Bank conflict 是什么

Bank conflict 不是“多个线程同时访问 shared memory”这么简单。真正的判断条件是:同一个 warp 发出一次 shared memory 指令时,是否有多个 lane 访问了同一个 bank 中的不同地址。如果是,硬件就要把这次请求拆成多轮;如果 32 个 lane 的地址落到 32 个不同 bank,则这些 bank 可以并行服务。

下面使用 NVIDIA Compute Capability 5.x 及之后常用的教学模型:shared memory 有 32 个 bank,每个 bank 每周期提供一个 32-bit word 的带宽,连续 32-bit word 依次映射到 bank 0...31。更宽的 load 会被拆成多个 transaction 分析,具体端口和调度细节随架构变化。AMD LDS 的思想相近,但 bank 数、wave 大小和一次 LDS 指令的分相规则可能不同,不能把本节的 32-lane 公式逐字照搬。

5.2.1 Bank 到底是什么:先看 shared memory 的硬件模型

从 CUDA 程序看,shared memory 是一个普通的线性地址空间;你声明的是数组或 byte buffer,并不会手工选择 bank。硬件为了让一个 warp 的多个地址能够并行访问,把这片片上 SRAM 在逻辑上组织成 32 个可同时工作的带宽分区,这些分区就是 bank

可以把一次 shared load 想成下面的流水:warp 提供最多 32 个 lane address;地址解码器从每个地址里算出 bank id 和 bank 内位置;路由网络把请求送到 32 个 bank。不同 bank 可以并行返回,同一 bank 的不同位置则需要拆成多轮。

概念硬件图:Warp → 地址解码与路由 → 32 个 Shared Memory Bank
Warp 发出一条 shared-memory 指令
L0 addr
L1 addr
L2 addr
L3 addr
L4 addr
L5 addr
L6 addr
L7 addr
...
L29 addr
L30 addr
L31 addr
地址解码器 + 32-way 路由
word = addr >> 2
bank = word & 0x1F
depth = word >> 5

低位选择 bank,高位选择该 bank 内的 32-bit word

32 个可并行服务的 bank
B0
B1
B2
B3
B4
B5
B6
B7
B8
B9
B10
B11
B12
B13
B14
B15
B16
B17
B18
B19
B20
B21
B22
B23
B24
B25
B26
B27
B28
B29
B30
B31

教学模型:每个 bank 每周期提供一个 32-bit word 的带宽

Bank 不是只保存一个 word。每个 bank 内部还有很多位置。连续地址会横向铺到 B0...B31,铺满后再回到 B0 的下一个 depth。下面只画前 8 个 bank;真实硬件同一行还有 B8...B31。

Shared memory 的“bank × depth”视图:橙色列都是 Bank 0,但地址不同
内部位置
Bank 0
Bank 1
Bank 2
Bank 3
Bank 4
Bank 5
Bank 6
Bank 7
depth 0
word 0
addr 0B
word 1
addr 4B
word 2
addr 8B
word 3
addr 12B
word 4
addr 16B
word 5
addr 20B
word 6
addr 24B
word 7
addr 28B
depth 1
word 32
addr 128B
word 33
addr 132B
word 34
addr 136B
word 35
addr 140B
word 36
addr 144B
word 37
addr 148B
word 38
addr 152B
word 39
addr 156B
depth 2
word 64
addr 256B
word 65
addr 260B
word 66
addr 264B
word 67
addr 268B
word 68
addr 272B
word 69
addr 276B
word 70
addr 280B
word 71
addr 284B
depth 3
word 96
addr 384B
word 97
addr 388B
word 98
addr 392B
word 99
addr 396B
word 100
addr 400B
word 101
addr 404B
word 102
addr 408B
word 103
addr 412B
32-bit 对齐地址中:address bits [1:0] 是 byte-in-word;bits [6:2] 选择 32 个 bank;bits [7+] 选择 bank 内 depth

BF16/FP16 元素只有 2B,所以两个相邻元素位于同一个 32-bit bank word;FP8/INT8 则一个 word 可容纳四个元素。真实 kernel 常使用 packed/vector load、ldmatrix 或更宽 transaction,分析时要把实际指令拆成它触及的 32-bit bank word,不能机械地把“一个元素”当成“一个 bank”。

图里的 depth 是为了理解硬件地址映射而使用的概念名称,不是 GEMM 矩阵的 row,也不是 CUDA API 暴露出来的可编程维度。程序仍然只看到一段线性 shared address。

把路由结果分成三种情况,Bank conflict 的定义就很直观:

无冲突:不同 bank
L0: word0B0 / depth0
L1: word1B1 / depth0
L2: word2B2 / depth0

每个 bank 只收到一个地址,可以并行服务。

Conflict:同 bank,不同地址
L0: word0B0 / depth0
L1: word32B0 / depth1
L2: word64B0 / depth2

B0 一次收到多个不同位置,需要拆分或 replay。

Broadcast:同 bank,同地址
L0: word0B0 / depth0
L1: word0B0 / depth0
L2: word0B0 / depth0

地址完全相同,读取结果可以广播给多个 lane。

一句话心智模型:Bank 像 32 条可以并行工作的 shared-memory 服务通道。冲突不是“走了同一条通道”,而是“同一周期要这条通道返回多个不同位置”;请求同一个位置则可以广播。

5.2.2 地址怎样映射到 bank

word_index = byte_address / 4;bank_id = word_index % 32

因此前 32 个 32-bit word 分别进入 bank 0 到 bank 31;第 32 个 word 又回到 bank 0。Bank 是交错编址,不是“bank 0 存一整段地址、bank 1 再存下一整段”。

连续 32-bit word 在 32 个 bank 中轮转
word 0
B0
word 1
B1
word 2
B2
word 3
B3
word 4
B4
word 5
B5
word 6
B6
word 7
B7
word 8
B8
word 9
B9
word 10
B10
word 11
B11
word 12
B12
word 13
B13
word 14
B14
word 15
B15
word 16
B16
word 17
B17
word 18
B18
word 19
B19
word 20
B20
word 21
B21
word 22
B22
word 23
B23
word 24
B24
word 25
B25
word 26
B26
word 27
B27
word 28
B28
word 29
B29
word 30
B30
word 31
B31
word 32
又回到 B0
word 33
又回到 B1
Conflict
多个 lane 落到同一 bank,且读取的是不同 word。例如 lane 0 读 word 0、lane 1 读 word 32:两者都是 B0,但地址不同,需要拆开。
Broadcast 例外
多个 lane 读取完全相同的 shared address 时,硬件可以广播该 word,不算 bank conflict。“同 bank”不够,还要看是不是“不同地址”。

5.2.3 把问题放回 GEMM:一个 warp 读取 A 的同一列

考虑一个教学用的 warp 映射。把两个 BF16 打包成一个 32-bit word,shared A tile 记作 A_pack[32][32];它在 K 方向覆盖 64 个 BF16。warp 的 lane r 负责输出列 n0 上的不同行,因此在某个固定 k_pair 上执行:

lane r:C[r, n0] += A_pack[r, k_pair] * B_pack[k_pair, n0]

k_pair=0 时,lane 0...31 要读取 A tile 的同一逻辑列 A_pack[0...31][0]。B 侧所有 lane 读取的是完全相同的 B_pack[0][n0],可以 broadcast;真正出问题的是 A 侧“同列、不同地址”的读取。

一个 K step:A 的纵向读取 × B 的广播 → 更新 C 的一列
A_pack[32][32]:每个 lane 读一行的 k0
lanek0k1...k31
L0 / r0A[0,0]A[0,1]...A[0,31]
L1 / r1A[1,0]A[1,1]...A[1,31]
L2 / r2A[2,0]A[2,1]...A[2,31]
L3 / r3A[3,0]A[3,1]...A[3,31]
...............
L31 / r31A[31,0]A[31,1]...A[31,31]
×
B_pack:同一个 word 广播
kn0n1...n31
k0B[0,n0]B[0,n1]...B[0,n31]
k1B[1,n0]B[1,n1]...B[1,n31]
...............
k31B[31,n0]B[31,n1]...B[31,n31]
C:每个 lane 更新不同输出行
lanen0n1...n31
L0 / r0C[0,n0]C[0,n1]...C[0,n31]
L1 / r1C[1,n0]C[1,n1]...C[1,n31]
L2 / r2C[2,n0]C[2,n1]...C[2,n31]
L3 / r3C[3,n0]C[3,n1]...C[3,n31]
...............
L31 / r31C[31,n0]C[31,n1]...C[31,n31]

5.2.4 普通 row-major layout 为什么产生 32-way conflict

A_pack[32][32] 的一行正好是 32 个 32-bit word。普通 row-major 下:

word_index(L) = L * 32 + k_pair;bank(L) = (L * 32 + k_pair) % 32 = k_pair

固定 k_pair=0 后,无论 lane 是多少,bank 都是 0。每个 lane 读的是不同 row,因此地址不同,不能 broadcast;硬件必须把一次 warp request 拆成最多 32 轮 conflict-free request,这就是 32-way bank conflict。

Lane逻辑元素32-bit word index相对 byte addressBank
L0A_pack[0][0]00BB0
L1A_pack[1][0]32128BB0
L2A_pack[2][0]64256BB0
L3A_pack[3][0]96384BB0
............B0
L31A_pack[31][0]9923968BB0
普通 layout:32 个 lane 全部落到 B0,但访问 32 个不同地址
L0 → B0
L1 → B0
L2 → B0
L3 → B0
L4 → B0
L5 → B0
L6 → B0
L7 → B0
L8 → B0
L9 → B0
L10 → B0
L11 → B0
L12 → B0
L13 → B0
L14 → B0
L15 → B0
L16 → B0
L17 → B0
L18 → B0
L19 → B0
L20 → B0
L21 → B0
L22 → B0
L23 → B0
L24 → B0
L25 → B0
L26 → B0
L27 → B0
L28 → B0
L29 → B0
L30 → B0
L31 → B0
Bank 0
32 个不同 word
L0@0BL1@128BL2@256BL3@384B...L31@3968B
不要把这个 A 访问和 B 的读取混为一谈:32 个 lane 都读同一个 B_pack[0][n0] 时是同地址 broadcast;32 个 lane 读 A_pack[0...31][0] 时虽然 bank 相同,但地址不同,所以是 conflict。

5.2.5 Swizzle:不改逻辑矩阵,只改物理地址

Swizzle 的核心是对 shared memory 的地址做一个可逆置换。写入和读取都使用同一映射,所以数学上仍然是原来的 A[row][col];只是逻辑列不再总是落在相同的物理列和 bank。为了把原理画清楚,先使用最简单的 32-bit word 级 XOR:

Concept sketch · word-level XOR swizzle
// 教学映射:tile width 是 32 个 32-bit word
physical_col = logical_col ^ row;

// global -> shared:把逻辑元素写到置换后的物理位置
smem[row][physical_col] = A_pack[row][logical_col];

// shared -> register:用同一映射找回原来的逻辑元素
x = smem[row][logical_col ^ row];
8×8 裁剪图:橙色都表示逻辑 k0;Swizzle 后它从竖直一列变成物理对角线
普通 row-major:logical k0 永远位于 physical p0
p0
p1
p2
p3
p4
p5
p6
p7
r0
k0
k1
k2
k3
k4
k5
k6
k7
r1
k0
k1
k2
k3
k4
k5
k6
k7
r2
k0
k1
k2
k3
k4
k5
k6
k7
r3
k0
k1
k2
k3
k4
k5
k6
k7
r4
k0
k1
k2
k3
k4
k5
k6
k7
r5
k0
k1
k2
k3
k4
k5
k6
k7
r6
k0
k1
k2
k3
k4
k5
k6
k7
r7
k0
k1
k2
k3
k4
k5
k6
k7
XOR swizzle:physical_col = logical_col ^ row
p0
p1
p2
p3
p4
p5
p6
p7
r0
k0
k1
k2
k3
k4
k5
k6
k7
r1
k1
k0
k3
k2
k5
k4
k7
k6
r2
k2
k3
k0
k1
k6
k7
k4
k5
r3
k3
k2
k1
k0
k7
k6
k5
k4
r4
k4
k5
k6
k7
k0
k1
k2
k3
r5
k5
k4
k7
k6
k1
k0
k3
k2
r6
k6
k7
k4
k5
k2
k3
k0
k1
r7
k7
k6
k5
k4
k3
k2
k1
k0

Swizzle 后,在固定 k_pair 时:

physical_col(L) = k_pair ^ L;bank(L) = (L * 32 + (k_pair ^ L)) % 32 = k_pair ^ L

对固定的 k_pair,XOR 是一个一一映射:lane 0...31 会得到 0...31 的某种排列,不会有两个 lane 落到同一个 bank。以 k_pair=0 为例,physical column 就等于 lane id:

Swizzle 后:32 个 lane 分散到 32 个 bank
L0 → B0
L1 → B1
L2 → B2
L3 → B3
L4 → B4
L5 → B5
L6 → B6
L7 → B7
L8 → B8
L9 → B9
L10 → B10
L11 → B11
L12 → B12
L13 → B13
L14 → B14
L15 → B15
L16 → B16
L17 → B17
L18 → B18
L19 → B19
L20 → B20
L21 → B21
L22 → B22
L23 → B23
L24 → B24
L25 → B25
L26 → B26
L27 → B27
L28 → B28
L29 → B29
L30 → B30
L31 → B31
Lane逻辑元素Swizzle 后物理元素word indexBank
L0A_pack[0][0]smem[0][0]0B0
L1A_pack[1][0]smem[1][1]33B1
L2A_pack[2][0]smem[2][2]66B2
L3A_pack[3][0]smem[3][3]99B3
...............
L31A_pack[31][0]smem[31][31]1023B31
关键点:Swizzle 没有转置 A,也没有改变 C=A×B。它只改变逻辑坐标到 shared physical address 的映射。生产者写入和消费者读取都经过同一个可逆映射,因此取出的仍是同一个逻辑元素。

5.2.6 为什么生产 kernel 的 Swizzle 看起来更复杂

physical_col = logical_col ^ row 是为了讲清 bank 置换的 32-bit word 级模型。真实 Tensor Core GEMM 还要同时满足三类约束:

  1. global → shared 通常按 16B vector copy,不能随意打散一个 vector 内的元素。
  2. shared → register 可能使用 ldmatrix、TMA 或架构特定 fragment layout,lane 的读取模式不是简单的一 lane 一 word。
  3. A/B 的 dtype、major mode、MMA tile 和 stage layout 不同,参与 XOR 的地址 bit 也不同。

因此 CUTLASS/CuTe 的 Swizzle 通常表现为“保留最低若干对齐 bit,再把一组高位 bit XOR 到地址低位”。它本质上仍是同一件事:保持 vector alignment,同时把消费者会同时访问的地址分散到不同 bank

CuTe 思路:保留 MBase 个低位;取 BBits 个地址 bit;平移 SShift 后做 XOR
先用本节公式检查逻辑访问是否冲突,再根据具体 ldmatrix/MMA copy atom 选择 production swizzle。不要从网上抄一个 XOR 常数:同一个 swizzle 换了 dtype、tile shape 或 lane mapping,可能重新产生 conflict,甚至破坏 16B 对齐。

进一步阅读:CUDA C++ Best Practices Guide:Shared Memory and Memory Banks,以及 CUTLASS CuTe Swizzle API

5.3 Padding 什么时候有用

padding 是最简单的破冲突方法之一:如果每行长度刚好导致 stride 命中同一组 bank,就给 shared tile 的 leading dimension 多加几个元素。

Concept sketch
// 不是固定要 +1,而是用 padding 改变行间 stride
__shared__ bf16 smem_a[BM][BK + PAD];

padding 的代价是 shared memory 占用增加,并且可能影响 vectorized store/load 对齐。因此它适合教学和简单 kernel;在生产级 Tensor Core kernel 里,更常见的是专门设计 shared layout。

6. Register Tiling:每个线程算多个输出元素

如果一个 thread 只算一个 C[m,n],每次加载 A/B 后只产生一个 accumulator 的贡献,计算密度不够。Register tiling 会让每个 thread 持有多个 accumulator,例如 4 x 4 的小块。

Concept sketch
float acc[TM][TN] = {0};

for k0 in range(0, K, BK):
  load A/B tile
  for kk in range(BK):
    a_frag[TM] = load A rows for this thread
    b_frag[TN] = load B cols for this thread
    for i in range(TM):
      for j in range(TN):
        acc[i][j] += a_frag[i] * b_frag[j]

这一步提升 arithmetic intensity,但也会增加寄存器压力。寄存器太多会降低 occupancy,反而可能变慢。

6.1 accumulator 为什么放寄存器

每个 C[m,n] 都要沿 K 累加。如果每次 partial sum 都写回 shared 或 global,再读出来继续累加,带宽和延迟都会爆炸。所以高性能 GEMM 会让 accumulator 长时间留在寄存器里,直到 K 维全部处理完后再进入 epilogue。

每线程寄存器数 ≈ accumulator registers + A/B fragment registers + 地址/循环/临时变量

例如一个线程维护 8x4 个 FP32 accumulator,仅 accumulator 就需要 32 个寄存器。再加上 fragment 和临时变量,很容易达到 80、120 甚至更多寄存器。

6.2 寄存器越多不一定越快

更多 register tiling 会提高单线程计算密度,但会降低一个 SM 上能同时驻留的 warp 或 CTA 数。如果 occupancy 太低,kernel 对内存延迟、指令调度空洞、同步等待的隐藏能力会下降。

现象可能原因调参方向
Tensor Core 利用率低,occupancy 也低寄存器或 shared memory 太多,活跃 warp 不够减小 warp tile、stage 数或 accumulator 数
occupancy 高但 TFLOP/s 低每个线程工作太少,数据复用不足增大 tile 或 register tiling
local memory load/store 出现寄存器溢出到 local memory减少临时变量,拆小 tile,检查编译器 unroll

6.3 Register tiling 和 warp-level 分工

在 Tensor Core kernel 里,寄存器不只是普通标量数组。A/B fragment 会按照 MMA 指令要求分布在 warp 的多个 lane 上,accumulator fragment 也由整个 warp 共同表示。单个线程看到的是自己那部分 fragment,整个 warp 合起来才是一个完整的 MMA tile。

视角每个 lane 看到什么整个 warp 合起来是什么
普通 CUDA core tiling几个标量 A/B 和多个 acc[i][j]一个 warp 覆盖一片 C 子区域,但每个线程仍在做标量 FMA
Tensor Core MMAA/B fragment 的若干寄存器和 accumulator fragment 的一部分一次或多次 MMA tile,例如把多个 m16n8k16 拼成 warp tile
Epilogue 写回lane 持有的 accumulator fragment 要映射回全局 C 坐标warp/CTA 协作把输出重排成 coalesced store
理解 Tensor Core GEMM 时,要从“一个线程算一个元素”的模型切换到“一个 warp 协作算一个或多个小矩阵 tile”的模型。

7. BF16 Tensor Core:真正的高吞吐路径

在 NVIDIA GPU 上,BF16 GEMM 的高性能路径通常不是普通 CUDA core FMA,而是 Tensor Core MMA 指令。

Concept sketch
// 概念模型,不是完整可编译代码
for k0 in range(0, K, BK):
  load A/B tiles into shared memory
  for warp_k in range(0, BK, mma_k):
    a_frag = ldmatrix_or_load_matrix_A(shared_A)
    b_frag = ldmatrix_or_load_matrix_B(shared_B)
    acc_frag = mma_sync(acc_frag, a_frag, b_frag)

BF16 常见策略是:输入 A/B 为 BF16,Tensor Core 做 BF16 multiply,accumulator 使用 FP32。最后写回时再转成 BF16 或 FP32,取决于输出需求。

对象推荐 dtype原因
A/B inputBF16吞吐高,动态范围接近 FP32
AccumulatorFP32降低 K 维累加误差
C outputBF16 或 FP32取决于后续算子和精度要求
高性能 GEMM 的核心不是“写一个 for loop”,而是组织出适合 Tensor Core 消费的 tile、fragment 和 pipeline。

7.1 MMA 指令在做什么

Tensor Core 指令可以理解为 warp 级别的小矩阵乘法:从 A fragment 和 B fragment 读入一小块矩阵,乘加到 accumulator fragment。它不是每个线程独立做完整矩阵乘法,而是 warp 内 32 个 lane 共同完成。

Concept sketch
// 概念公式
D_fragment = A_fragment @ B_fragment + C_fragment

具体指令形状会随架构变化,例如 m16n8k16 这类形状表示一次 MMA 覆盖的 M/N/K 小块大小。kernel 的 warp tile 通常由很多个 MMA 指令拼起来。

7.2 ldmatrix 和 shared layout 的关系

Tensor Core 前的数据搬运通常不是普通标量 load,而是按照矩阵片段形式从 shared memory 取数。ldmatrix 一类指令要求 warp 中各 lane 以特定模式读取 shared 地址,所以第 5 章的 swizzle 和 padding 会直接影响 Tensor Core 喂数效率。

Global 到 Shared
目标是 coalesced、对齐、吞吐高。
Shared 到 Fragment
目标是 bank conflict 少,并匹配 MMA fragment 布局。

7.3 BF16 的数值路径

BF16 只有 7 个显式 fraction bit,但 exponent 范围与 FP32 相同。实践里的常见安全路径是 BF16 输入、FP32 accumulator,最后按需求 cast 到 BF16/FP16/FP32。除非目标硬件和误差预算明确允许更低精度累加,否则 K 维 reduction 通常应保留 FP32 accumulator。

Concept sketch
bf16 A, bf16 B
fp32 acc = mma_bf16(A, B, fp32 acc)
output = cast_or_fuse(acc)

8. Pipeline 和 Async Copy:让加载与计算重叠

如果每次都同步加载 A/B tile,再计算,再加载下一块,Tensor Core 会等待内存。高性能 kernel 会做 pipeline。

Prologue
Load S0
Steady state
Compute S0 + Load S1
Steady state
Compute S1 + Load S2
Drain
Compute last stage
Epilogue
Store C
Pipeline 时间线:计算当前 stage,同时加载下一 stage
t0
t1
t2
t3
Load pipe
load S0
load S1
load S2
load S3
MMA pipe
compute S0
compute S1
compute S2

实际实现里会使用 double buffering 或 multi-stage buffering:当 Tensor Core 计算当前 stage 时,异步加载下一 stage 的 A/B。

Concept sketch
// 概念模型:允许下一 stage 保持 in flight
prefetch stage 0
commit copy group

for k_stage in range(num_k_stages):
  if has_next_stage:
    prefetch next stage
    commit copy group

  wait until current stage is ready
  CTA barrier for current buffer
  mma current stage
  CTA barrier before reusing current buffer
  swap current / next buffer

wait for remaining copies
store C tile

把 mainloop 组织成“搬运下一块、计算当前块”的稳态流水,是 CUTLASS、Triton 生成代码和 BLASLt 类库中普遍存在的高性能方向;具体使用 cp.async、TMA 还是其他 copy primitive,取决于架构与实现。

8.1 Double buffering 的直觉

只有一个 shared buffer 时,kernel 必须“加载完再算,算完再加载”。double buffering 准备两份 shared buffer:计算 stage 0 时加载 stage 1,计算 stage 1 时加载 stage 2。这样内存延迟被 mainloop 的 MMA 指令部分隐藏。

Concept sketch
load smem[0]
for k_stage:
  load smem[next]       // async copy
  compute smem[current] // mma
  swap current/next

8.2 stage 数不是越多越好

更多 stage 可以覆盖更长内存延迟,但每多一个 stage 都会增加 shared memory 占用,也可能减少可驻留 CTA 数。生产 kernel 常在 2、3、4 stage 之间调参,取决于矩阵形状、GPU 架构和每个 CTA 的 tile 大小。

stage 选择优点风险
2 stagesshared 占用较低,实现直观可能盖不住 HBM 延迟
3-4 stages更容易让 Tensor Core 连续工作shared 占用上升,occupancy 可能下降
过多 stages理论上预取更深资源占用过高,收益递减

8.3 同步点要尽量少但必须正确

shared memory 被多个线程协作写入和读取,因此同步不能随意删。优化的方向通常是用异步 copy 的 wait group、barrier 或 pipeline abstraction 精确表达“哪一 stage 已经可读”,而不是在每个小步骤后粗暴同步。

Concept sketch
// 稳态示意:此时 current group 和 next group 都在队列中
cp_async_copy(smem[next], gmem[next]);
cp_async_commit_group();

cp_async_wait_group(1);  // 最多保留 1 组未完成:current 已就绪,next 可继续在途
__syncthreads();         // CTA 内线程现在可以安全读取 current buffer
mma_on(smem[current]);
__syncthreads();         // 覆盖这个 buffer 前,确保所有消费者都已读完

// pipeline drain 时再用 wait_group(0) 等完最后一组

wait_group(0) 会等待所有已提交 group;如果在发出 next stage 后立刻这样做,就会把预取也等完,失去加载与计算的重叠。稳态常用 wait_group(1) 只保证当前最老的 group 已完成,同时允许下一组继续在途;最后 drain 才等到 0。生产代码通常用 CUTLASS pipeline、CUDA barrier 或架构专用 primitive 管理这些状态。wait_group 管 copy group 的完成边界,barrier 管 CTA 线程之间的可见性和 buffer 复用边界,两者不能互相替代。

9. Epilogue:累加之后还没结束

mainloop 负责把 A @ B 累加到 FP32 accumulator,但真实模型里的 GEMM 往往还要做 bias、activation、scale、residual、cast、量化和写回。这部分通常叫 epilogue。

Concept sketch
// 概念 epilogue
float x = acc[m, n];
x = alpha * x + beta * old_c[m, n];
x = x + bias[n];
x = activation(x);
C[m, n] = cast_to_output_dtype(x);

9.1 为什么 epilogue 要融合

如果 mainloop 写出 FP32 C,再启动另一个 kernel 做 bias/activation/cast,就会多一次 global write 和 read。对大 GEMM 来说这可能还可以接受;对 LLM 推理里的小 M、decode、MoE grouped GEMM,epilogue 的内存流量和 kernel launch overhead 会非常明显。

高性能实现通常把 epilogue 融合在 GEMM kernel 末尾,让 accumulator 在寄存器里直接完成后处理,然后只写一次最终输出。

9.2 写回也要 coalesced

不要只关心 A/B 的 load。C 的 store 如果是分散写、非对齐写,也会浪费带宽。很多 kernel 会让 warp 或 CTA 的输出 layout 匹配全局内存中的连续 N 维,这样写回更容易合并。

10. Profiling 工具和方法

评估 GEMM 不要只看 wall time。wall time 告诉你“慢了”,profiling 才告诉你“为什么慢”。CUDA 性能分析通常分两层:先用 Nsight Systems 看时间线和 kernel 调度,再用 Nsight Compute 深挖单个 kernel 的硬件指标。

10.1 工具怎么选

工具回答的问题适合场景
Nsight Systems / nsys程序时间花在哪里,kernel 是否被 CPU launch、同步、通信、数据拷贝卡住端到端推理、多个 kernel、CPU/GPU overlap、NCCL/拷贝/调度分析
Nsight Compute / ncu某个 CUDA kernel 为什么没有打满硬件GEMM kernel 本身调优,Tensor Core、访存、occupancy、stall reason
PyTorch ProfilerPyTorch op、CUDA kernel、shape、调用栈的对应关系先定位哪个 op / kernel 慢,再下钻到 nsys/ncu
CUPTI程序化采集 trace/counter框架、服务、自动 benchmark 系统里集成 profiling

10.2 推荐工作流

  1. 先用业务 benchmark 固定 shape、batch、dtype、warmup、iteration,避免 profile 的对象不稳定。
  2. nsys 看时间线:确认慢的是 GEMM kernel 本身,还是 kernel launch、同步、memcpy、NCCL、CPU 调度。
  3. 挑出最耗时或最可疑的 GEMM kernel,用 ncu 采集详细指标。
  4. 用指标判断瓶颈:Tensor Core 没打满、HBM 压力高、shared bank conflict、寄存器溢出、occupancy 太低、tail 浪费。
  5. 每次只改一个变量:BM/BN/BK、stage 数、warp 数、layout、epilogue 融合、split-K 或 grouped 调度。

10.3 Nsight Systems 怎么用

nsys 适合先看全局时间线。它不会告诉你 shared memory bank conflict 细节,但能快速判断“是不是 GEMM kernel 自己慢”。

Production profiling command
# 采集端到端 timeline
nsys profile \
  -o gemm_timeline \
  --trace=cuda,nvtx,osrt,cublas,cudnn \
  --force-overwrite=true \
  ./your_benchmark

# 只有 benchmark 调用了 cudaProfilerStart/Stop 时,才按 API 范围截取
nsys profile -o gemm_range \
  --trace=cuda,nvtx,osrt \
  --capture-range=cudaProfilerApi \
  --capture-range-end=stop \
  ./your_benchmark

# 也可以直接 profile Python
nsys profile -o torch_gemm --trace=cuda,nvtx,osrt python bench.py
在 nsys 里看什么可能说明什么
GPU timeline 中 GEMM kernel 是否连续中间有空洞可能是 CPU launch、同步或依赖问题
Memcpy / Memset 是否夹在 GEMM 中间可能有隐式拷贝、临时 buffer、layout 转换
CUDA API 时间是否很长可能 CPU 端同步、allocator、driver 调用开销明显
NCCL 与 GEMM 是否 overlap分布式推理里通信可能盖不住或阻塞计算
kernel 数量是否碎片化小 kernel 太多,launch overhead 或 fusion 不足

10.4 Nsight Compute 怎么用

ncu 用来深挖单个 kernel。它采集 counter 会显著拖慢程序,所以一般只 profile 少量 iteration,并用 kernel 名称或 launch 序号过滤。

Production profiling command
# 基础采集:先用常用 section
ncu --set full \
  --target-processes all \
  --kernel-name regex:gemm \
  -o gemm_kernel \
  ./your_benchmark

# 更轻量:只采集常用分析 section
ncu \
  --section SpeedOfLight \
  --section MemoryWorkloadAnalysis \
  --section Occupancy \
  --section SchedulerStats \
  --section WarpStateStats \
  --target-processes all \
  ./your_benchmark
profile 前一定要 warmup。第一次运行可能包含 JIT、cudaMalloc、cache cold start、autotune,这些会污染 GEMM kernel 本身的判断。

10.5 GEMM 里最该看的指标

至少看这些指标:

指标说明常见问题对应调参方向
TFLOP/s2MNK / time 算实际吞吐,再与该 dtype 的理论峰值比较Tensor Core 没打满,或计时/工作量口径有误回 §4/§7 检查 tile、MMA 对齐和是否走 Tensor Core 路径
HBM bandwidthglobal load/store 的实际带宽及其峰值占比tile 复用差、额外 workspace 或 epilogue 访存过多回 §4/§5 调 BM/BN/BK、coalesced load 和 shared 复用
Occupancy资源约束下可驻留 warp/CTA 的比例;不是性能分数寄存器或 shared memory 过多,也可能只是 grid 太小回 §6/§8 减小 register tile、stage 数或 shared footprint,并结合 grid waves 判断
Tail / tile utilization由 shape、tile 和有效元素比例推算;通常不是一个统一的 NCU counter非整除维度或小 M 让大量 lane/tile 空转回 §4.5/§10.9 考虑小 tile、split-K、persistent/grouped 调度
Shared bank conflictshared load/store 是否被序列化shared layout 不匹配 warp 读取模式回 §5/§7 调 swizzle、padding、ldmatrix 读取布局
Local memory寄存器溢出后的隐式内存访问register pressure 过高或 unroll 太激进回 §6/§9 减少 accumulator、临时数组和 epilogue 复杂度
Nsight Compute 的具体 metric 名称会随 GPU 架构和工具版本变化。先用 SpeedOfLight、MemoryWorkloadAnalysis、Occupancy、SchedulerStats 等 section 建立结论,再查该版本的 metric 定义;不要只凭一个百分比判断瓶颈。

10.6 指标怎么解读

Profiling signal优先怀疑下一步回到章节
Tensor Core utilization 低,HBM 带宽也低并行度不足、tail 多、kernel launch 间隙、依赖等待回 nsys 看 timeline,检查小 M、CTA 数量、split-K§4.5、§10.9
Tensor Core utilization 低,HBM 带宽高内存喂不饱,A/B tile 复用差,global load 不合并检查 BM/BN/BK、global load pattern、layout 转换§4、§5
shared load/store bank conflict 高shared layout 不适合 warp 读取或 ldmatrix 模式检查 swizzle、padding、shared leading dimension§5.2、§7.2
occupancy 很低寄存器、shared memory、cluster/stage 占用过大降低 tile、stage、epilogue 临时变量,查 local memory§6、§8、§9
local memory transaction 出现寄存器 spill减少 accumulator tile、unroll、临时数组或 epilogue 融合复杂度§6.2、§9
pipeline stall 或 barrier stall 高async copy 没盖住延迟,wait/barrier 放置过粗检查 stage 数、wait group 距离、barrier 粒度§8
epilogue 时间占比高FP8 scale、amax、bias、activation、requantize 过重分离 mainloop/epilogue profile,减少额外读写§9、§11.8

10.7 一个实用排查顺序

  1. 先确认结果正确,尤其是 tail block、非整除 K、转置 layout、dtype cast。
  2. 看是否走到 Tensor Core 指令路径,而不是退回普通 FMA。
  3. 看 HBM load/store 是否异常高,判断 tile 复用和 epilogue 融合是否有效。
  4. 看 shared bank conflict,判断 shared layout 是否喂得动 Tensor Core。
  5. 看寄存器数量、local memory、occupancy,判断 tile 是否过大。
  6. 最后再微调 tile size、stage 数、warp 数和 split-K 等策略。

10.8 PyTorch 场景怎么定位到 GEMM

如果 GEMM 来自 PyTorch 或推理框架,先用 PyTorch Profiler 或 NVTX 标记把高层 op 和底层 kernel 对上,再用 nsys/ncu 深挖。

Production reference
import torch
from torch.profiler import profile, ProfilerActivity

with profile(
    activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA],
    record_shapes=True,
    with_stack=True,
) as prof:
    run_model_once()

print(prof.key_averages().table(
    sort_by="cuda_time_total",
    row_limit=20,
))

看到慢 op 后,再在关键区域加 NVTX range,nsys 里就能按阶段定位。

Production reference
torch.cuda.nvtx.range_push("prefill_gemm")
run_prefill()
torch.cuda.nvtx.range_pop()

10.9 小 M 场景的特殊性

LLM decode 常见 M 很小,甚至 batch token 数只有几十或更少。这时大 CTA tile 可能浪费,输出 tile 数也可能少到无法铺满所有 SM。工程上可能需要 split-K、persistent kernel、grouped GEMM,或者把多个请求/专家合并调度。

10.9.1 Persistent kernel:CTA 算完一个 tile 后不退出

普通 GEMM 通常按输出 tile 启动大量 CTA,一个 CTA 算完自己负责的 tile 就结束。Persistent kernel 则只启动接近硬件可驻留数量的 CTA,让这些 CTA 在 kernel 生命周期内循环领取后续 tile;因此“persistent”表示 CTA 在一次 kernel launch 内持续工作,不是 kernel 永远不退出。

普通 tile-per-CTA
Concept sketch
CTA 0: tile 0 -> exit
CTA 1: tile 1 -> exit
CTA 2: tile 2 -> exit
...
下一批工作再次 launch
Persistent CTA
Concept sketch
while scheduler.has_work():
  tile = scheduler.next_tile()
  compute_gemm_tile(tile)

// 所有 tile 完成后才退出

当多个 tile、problem 或请求被合进同一次 launch 时,它可以摊薄 kernel launch 和 CTA 建立/回收开销;scheduler 也能把剩余 tile 交给先空闲的 CTA,改善尾部负载均衡。某些实现还可以跨多个 tile 复用调度状态或权重相关 metadata。对单次、单个 GEMM 而言,persistent 本身并不会减少 launch 次数。

Persistent 不等于自动解决小 M。如果整个 workload 真的只有两三个 tile,没有其他请求、group 或 K-split 可做,常驻 CTA 也不能凭空制造计算并行度。它更适合“单个问题很小,但同时有很多问题/tile 可以连续领取”的场景。

10.9.2 Grouped GEMM:把多个不同 GEMM 作为一组调度

Grouped GEMM 的输入不是一对 A/B,而是一组 GEMM problem。每个 problem 可以有自己的指针、M/N/K、stride 和输出地址;kernel 把所有 problem 的输出 tile 放进同一个逻辑工作集合,再分给 CTA。

Concept sketch · MoE grouped GEMM
// 四个 expert 实际收到的 token 数不同,因此 M 不同
GEMM 0: [ 1, K] @ [K, N] -> [ 1, N]
GEMM 1: [ 3, K] @ [K, N] -> [ 3, N]
GEMM 2: [17, K] @ [K, N] -> [17, N]
GEMM 3: [64, K] @ [K, N] -> [64, N]

grouped_gemm([problem0, problem1, problem2, problem3])

如果分别 launch 四个 kernel,前三个 GEMM 可能都只有很少 CTA,launch 和尾部空转占比很高。Grouped GEMM 用一次 launch 暴露整组 tile,scheduler 可以让空闲 CTA 从下一个 expert/problem 继续取工作。它尤其适合 MoE,因为每个 expert 的 token 数动态变化,problem shape 天然不规则。

Grouped + Persistent 的常见组合
G0 tiles
M=1
G1 tiles
M=3
G2 tiles
M=17
G3 tiles
M=64
统一 tile scheduler
Resident CTA 0
G0/T0 → G2/T1 → G3/T3 → ...
Resident CTA 1
G1/T0 → G2/T2 → G3/T4 → ...

10.9.3 两个概念是什么关系

概念回答的问题关键机制
Persistent kernel一个 CTA 算完当前 tile 后做什么?CTA 保持驻留,循环从 scheduler 领取多个 tile
Grouped GEMM哪些 GEMM problem 放在同一次调度里?把多个不同 shape/pointer 的 GEMM tile 合成一个工作集合
Grouped persistent GEMM如何持续处理一组不规则小 GEMM?少量驻留 CTA 跨 problem 不断领取 tile;CUTLASS grouped kernel 常采用这种方式
Batched GEMM如何重复执行一批通常较规则的 GEMM?常见接口假设 shape/stride 更一致;不等同于一般的 heterogeneous grouped GEMM

代价也很直接:device scheduler、problem metadata 读取和 tile 查找会产生额外开销;不同 problem 的 K 相差很大时,CTA 工作量仍可能失衡。因此 grouped/persistent 适合小而多、shape 不规则的 workload,不代表单个标准大 GEMM 也应该使用。

进一步阅读:CUTLASS Grouped Kernel Schedulers。CUTLASS 的 grouped kernel 会启动少于总 tile 数的 persistent threadblocks,每个 threadblock 循环查询 scheduler 并计算一个或多个 problem tile。
LLM 推理里的 GEMM 经常不是理想大方阵。小 M、MoE grouped GEMM、prefill/decode 混合形状都会让“库默认最优”不再绝对成立;应按 shape bucket 分别 profile。

11. FP8 GEMM 怎么做

FP8 不是简单地把 BF16 输入换成 FP8。FP8 的动态范围和精度都更紧,需要 scale 管理。

11.1 FP8 的基本形态

格式特点常见用途
E4M3mantissa 多,精度相对好,范围较小activation / weight 常用
E5M2exponent 多,范围更大,精度较低gradient 或范围更大的张量

11.2 scale 是 FP8 GEMM 的核心

通常会把真实值映射到 FP8 可表达范围:

Concept sketch
x_fp8 = quantize(x / scale)
x_real ≈ dequantize(x_fp8) * scale

scale 粒度可以不同:

  • per-tensor scale:实现简单,但容易被 outlier 影响。
  • per-row / per-column scale:必须注明广播轴;weight 常按输出通道,activation 常按 token/row。
  • per-block scale:例如每 32/64/128 个元素一个 scale,兼顾精度和吞吐。

11.3 FP8 GEMM 的计算路径

Concept sketch
// 简化路径:scale_a / scale_b 在该输出 tile 的整个 K reduction 上恒定
A_fp8, scale_A = quantize(A_bf16_or_fp32)
B_fp8, scale_B = quantize(B_bf16_or_fp32)

acc_fp32 = fp8_tensor_core_mma(A_fp8, B_fp8)
C = acc_fp32 * scale_A * scale_B
C_out = cast_or_quantize(C)

只有当 scale_A * scale_B 对整个 K reduction 都不变时,scale 才能从求和中提出,在 mainloop 结束后统一应用。高性能实现会把 scale load、MMA、scale apply、输出 cast 或 requantize 融合起来,减少额外内存读写。

关键边界:如果 scale 沿 K-block 变化,就不能先把所有 FP8 product 累加完再乘一次 scale。每个 K-block 的 partial accumulator 都要乘自己的 scale_A[kb] * scale_B[kb] 后再进入总和;§11.7 展开这条路径。

11.4 从 BF16 kernel 迁移到 FP8 时要改什么

  1. global memory layout 要容纳 FP8 数据和 scale 数据。
  2. load path 要加载 FP8 tile,同时加载对应 scale。
  3. MMA 指令路径从 BF16 MMA 换成 FP8 MMA。
  4. accumulator 通常仍用 FP32。
  5. 按 scale 粒度决定缩放位置:K 维恒定的 scale 可放到 epilogue;沿 K 变化的 block scale 必须进入 mainloop / block-scaled MMA。
  6. epilogue 还要处理 bias、activation、输出 dtype,以及可能的 output scale。
  7. 如果输出继续 FP8,还要做 requantize,并按约定读取或写出新的 output scale / amax。

11.5 scale 粒度会改变 kernel 设计

per-tensor scale 在整个 reduction 上恒定时,可以在 epilogue 里乘一个标量,最简单但最容易受 outlier 牵制。所谓 per-channel 必须说明轴:权重按输出通道量化时,scale 常沿 N 维;activation 按 token/row 量化时,scale 常沿 M 维。这些 scale 若不随 K 变化,仍可在输出 tile 上应用。per-block scale 一旦包含 K 维分块,scale load 和缩放就进入 mainloop,不能再被当作普通 epilogue 标量。

scale 粒度kernel 代价精度特点
per-tensor最小;scale 对 K 恒定时可在 epilogue 应用容易被全局 outlier 牵制
per-row / per-column按 M 或 N 广播 scale;输出 tile 需要对应向量常用于 token/activation 或输出通道/weight
per-block,且沿 K 分块每个 K stage 都要索引 scale,并缩放 partial sum 或由 block-scaled MMA 消费局部动态范围更好,工程复杂度更高

11.6 FP8 的 accumulator 和输出

FP8 GEMM 的常见安全路径仍使用 FP32 accumulator,但具体 accumulator 模式由硬件指令和库配置决定。原因和 BF16 一样:K 维累加误差需要受控。输出可以是 BF16/FP16/FP32,也可以重新量化成 FP8;如果输出 FP8,epilogue 还要决定 output scale,并处理 clipping 和 rounding。

Concept sketch
// 仍假设 input scale 在整个 K reduction 上恒定
acc_fp32 = fp8_mma(A_fp8, B_fp8)
y_fp32 = acc_fp32 * scale_a * scale_b + bias
out_fp8 = quantize(y_fp32 / scale_out)

即使 input scale 可以留到 K reduction 之后处理,scale apply、bias/activation、output scale 和再量化仍会增加 epilogue 的指令与访存;因此只比较 FP8/BF16 的 MMA 峰值,不能推导完整 kernel 的加速比。

11.7 per-block scale 的 layout 和索引

per-block scale 的难点不是多乘一个标量,而是 scale tensor 必须和 tile 形状、K-loop 顺序一起设计。下面用一个简化例子:A 的 scale 覆盖 64 x 64 子块,B 的 scale 也覆盖 64 x 64 子块。对某个固定的输出 scale-block (mb, nb),每个 K-block 都有不同的 scale pair。

Concept sketch
// mb / kb / nb 是 64x64 scale block 的索引
acc_fp32[mb, nb] = 0

for kb in range(num_k_scale_blocks):
  partial_fp32 = fp8_mma(
      A_fp8[mb, kb],
      B_fp8[kb, nb])

  scale_a = ScaleA[mb, kb]
  scale_b = ScaleB[kb, nb]
  acc_fp32[mb, nb] += partial_fp32 * scale_a * scale_b

y[mb, nb] = epilogue(acc_fp32[mb, nb])
C[mb, nb] = Σkb partial_fp32[mb, kb, nb] × ScaleA[mb, kb] × ScaleB[kb, nb]

公式的重点是 scale multiplication 位于 K-block 求和内部。实现上不一定真的生成一块独立的 partial_fp32:支持 block-scaled MMA 的架构可以把 data fragment 和 scale-factor fragment 一起交给 MMA;其他实现可能在 mainloop 中提升、缩放并累加 partial fragment。数学语义必须相同。

如果 scale block 和 CTA/MMA tile 不整齐对齐,kernel 会多出 predicate、额外 scale load,甚至跨 block 的 scale 拼接与 fragment 重排。实际工程常让 scale layout 和 mainloop 的消费粒度对齐,牺牲一点布局自由度,换取更简单的地址计算和更连续的 scale load。

11.8 requantize:clipping、rounding、amax

如果 GEMM 输出还要写成 FP8,epilogue 需要把 FP32 accumulator 重新映射回 FP8 范围。这个过程通常包含三件事:先按约定除以 output scale,再对目标 FP8 格式的可表达范围做 clipping,最后按硬件或库的策略 rounding。output scale 可以预先给定,也可以根据当前或历史 amax 动态更新;两种路径的同步和可复现性成本不同。

步骤为什么需要常见风险
选择 output scale决定 FP8 动态范围如何覆盖输出分布scale 太小会大量 clipping,scale 太大有效精度下降
clipping避免超出 E4M3/E5M2 可表达范围outlier 多时会损失模型精度
rounding把 FP32 映射到离散 FP8 值不同 rounding 策略会影响误差分布和可复现性
amax 统计给下一次 scale 更新提供动态范围信息额外 reduction 可能让 epilogue 变重

11.9 为什么 FP8 epilogue 可能成为瓶颈

BF16 GEMM 的 epilogue 可能只是 bias、activation 和 cast;FP8 epilogue 往往还要处理可从 K reduction 提出的 input scale、output scale、clipping/rounding,并可能写 amax 或 auxiliary tensor。沿 K 变化的 A/B block scale 属于 mainloop 成本,不应误记到 epilogue。mainloop 因为 FP8 Tensor Core 变快以后,剩余 epilogue 工作在总时间里的占比反而可能上升。

下面是 profile 的解释框架,不是某张 GPU 上的实测比例。对融合 kernel,mainloop/epilogue 占比通常需要结合 source-level counter、指令区间,或用“保留 mainloop、逐项关闭 epilogue 功能”的变体对照来估算,不能指望一个通用 NCU 指标直接给出百分比。

Profile 对比BF16 常见解释FP8 常见解释
Mainloop 占比高Tensor Core 或 shared pipeline 是主要瓶颈scale load 已被很好隐藏,主要仍是 MMA
Epilogue 占比高bias/activation/store 或小 M launch overhead 明显scale apply、requantize、amax、额外 store 压过 mainloop 收益
HBM store 高输出写回或未融合后处理过重除了输出,可能还写 scale/amax/aux,需检查 epilogue fusion
FP8 的性能收益来自更低的带宽需求、更高的 Tensor Core 吞吐和更小的存储占用;代价是 scale 管理与数值验证明显更复杂。

12. CuTe / CUTLASS DSL:NVIDIA 生产实现的一条主力路径

前面讲的是 GEMM kernel 的原理:tile、shared memory、register、Tensor Core、pipeline、epilogue。在 NVIDIA 生产实现中,常见路径不是从裸 CUDA 一行行维护所有地址计算,而是使用 CuTe / CUTLASS 3.x API 这类 DSL 和模板库。

12.1 CuTe 在解决什么问题

CuTe 的核心抽象是 Tensor = pointer + layout,以及用 ShapeStrideLayoutTiledMMATiledCopy 描述数据如何被 CTA、warp 和 MMA 指令消费。它让你把“逻辑上的矩阵 tile”和“物理上的 shared/global/register 布局”分开表达。

名字容易混:CuTe 本体最早是 C++ template DSL,CUTLASS 3.x API 的 GEMM collective 建立在它上面;现在 NVIDIA 也提供 Python CuTe DSL,用 Python 语法写低层 GPU kernel,再 JIT/AOT 编译。因此 §12.3 的 C++ 代码更准确地说是“CUTLASS collective 代码,使用了 CuTe 类型”,不是手写 CuTe kernel。

12.1.1 CuTe 极简心智模型

读 CuTe 代码时,先不要从模板类型硬啃。把它看成三层:Tensor 保存指针和 layout,Layout 把逻辑坐标映射到物理 offset,TiledMMA 描述 warp 如何把 fragment 喂给 Tensor Core。

Concept sketch
// CuTe 风格示意,不是完整 kernel
Tensor gA = make_tensor(make_gmem_ptr(A), make_layout(make_shape(M, K), make_stride(K, 1)));
Tensor tileA = local_tile(gA, make_shape(BM, BK), make_coord(block_m, block_k));

// Layout 决定 (m, k) 如何变成地址 offset
// TiledMMA 决定 warp/lane 如何持有 A/B/C fragment
auto tiled_mma = make_tiled_mma(mma_atom, warp_layout);
auto thr_mma = tiled_mma.get_thread_slice(threadIdx.x);
CuTe 名词先按什么理解对应本文概念
Tensor数据指针 + layout 的视图global/shared/register tile
Layout坐标到 offset 的函数row-major、column-major、swizzle、padding
Shape逻辑 tile 尺寸BM/BN/BK、warp tile、MMA tile
TiledMMAwarp 到 MMA fragment 的分配规则§4.1 和 §6.3 的 warp-level 分工
裸 CUDA 难点
地址计算散落在代码里,swizzle、bank conflict、TMA、MMA fragment、epilogue 很容易互相缠住。
CuTe / CUTLASS 路径
用 layout 和 collective 描述 mainloop/epilogue,让编译期类型系统承载 tile shape、copy atom、MMA atom 和 schedule。
名字语言典型代码形态适合场景
CuTe C++C++ templatecute::Tensorcute::LayoutTiledMMATiledCopy理解 CUTLASS 内核、写极底层 kernel、精细控制布局
CUTLASS collectiveC++ templateCollectiveBuilderGemmUniversalAdapter工程里组合高性能 GEMM;§12.3 的代码属于这一层
Python CuTe DSLPythonimport cutlass.cute as cute@cute.kernel新 kernel 快速迭代、减少 C++ 模板负担、面向 CUTLASS 4.x 路线

12.2 CUTLASS 3.x 的分层

层级作用你通常关心什么
CuTelayout、tensor、copy、MMA 的底层 DSLtile 如何映射到线程、warp、shared、fragment
Collective Mainloop从 global/shared 到 MMA accumulatorA/B dtype、layout、TileShape、ClusterShape、stage、kernel schedule
Collective Epilogueaccumulator 到 D 输出scale、bias、activation、amax、aux、输出 dtype
Device Adapter把 kernel 包装成 host 可调用对象problem shape、stride、workspace、initialize、run

所以工程上常见的做法是:先用 CUTLASS 的 CollectiveBuilder 组合出一个 kernel,再对 tile shape、cluster shape、schedule 和 epilogue fusion 做调参。只有当 builder 覆盖不了需求时,才下沉到更底层的 CuTe kernel。

12.3 Hopper FP8 GEMM:完整 CUTLASS 3.x collective 代码

下面是一份教学版 Hopper FP8 GEMM。它保留了完整 host 侧流程:定义 CUTLASS/CuTe kernel 类型、分配 A/B/C/D、初始化输入、构造 Gemm::Arguments、申请 workspace、检查支持性、初始化并运行。真实项目应以 CUTLASS 仓库中的 examples/54_hopper_fp8_warp_specialized_gemmexamples/67_hopper_fp8_warp_specialized_gemm_with_blockwise_scaling 和 Python CuTe DSL 示例为准。

查看完整 Hopper FP8 CUTLASS C++ 代码
Production reference
// hopper_fp8_gemm.cu
// Teaching version: CUTLASS 3.x + CuTe shapes for Hopper FP8 GEMM.
// It is intentionally compact, but keeps the complete host-side run path.
//
// Typical compile shape:
//   nvcc -std=c++17 -arch=sm_90a -I${CUTLASS_ROOT}/include \
//        -I${CUTLASS_ROOT}/tools/util/include hopper_fp8_gemm.cu -o hopper_fp8_gemm

#include <iostream>

#include "cutlass/cutlass.h"
#include "cutlass/numeric_types.h"
#include "cutlass/util/host_tensor.h"
#include "cutlass/util/device_memory.h"
#include "cutlass/util/reference/host/tensor_fill.h"
#include "cutlass/util/reference/host/tensor_compare.h"
#include "cutlass/gemm/device/gemm_universal_adapter.h"
#include "cutlass/gemm/kernel/gemm_universal.hpp"
#include "cutlass/gemm/collective/collective_builder.hpp"
#include "cutlass/epilogue/collective/collective_builder.hpp"
#include "cutlass/epilogue/thread/activation.h"
#include "cutlass/epilogue/fusion/operations.hpp"
#include "cute/tensor.hpp"

using namespace cute;

using ElementA = cutlass::float_e4m3_t;
using ElementB = cutlass::float_e4m3_t;
using ElementC = cutlass::half_t;
using ElementD = cutlass::half_t;
using ElementAccumulator = float;
using ElementCompute = float;
using ElementAux = ElementD;
using ElementAmax = float;
using ElementBias = float;

using LayoutA = cutlass::layout::RowMajor;
using LayoutB = cutlass::layout::ColumnMajor;
using LayoutC = cutlass::layout::ColumnMajor;
using LayoutD = cutlass::layout::ColumnMajor;
using LayoutAux = LayoutD;

// Alignment is expressed in elements. For FP8, 128 bits means 16 elements.
constexpr int AlignmentA = 128 / cutlass::sizeof_bits<ElementA>::value;
constexpr int AlignmentB = 128 / cutlass::sizeof_bits<ElementB>::value;
constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;

using ArchTag = cutlass::arch::Sm90;
using OperatorClass = cutlass::arch::OpClassTensorOp;

// This is the CTA-level tile discussed earlier as BM/BN/BK.
using TileShape = Shape<_128, _128, _128>;
using ClusterShape = Shape<_1, _2, _1>;

using KernelSchedule =
    cutlass::gemm::KernelTmaWarpSpecializedCooperative;
using EpilogueSchedule =
    cutlass::epilogue::TmaWarpSpecializedCooperative;

using EpilogueTile = cutlass::epilogue::collective::EpilogueTileAuto;
using FusionOperation =
    cutlass::epilogue::fusion::ScaledLinCombPerRowBiasEltActAmaxAux<
        LayoutAux,
        cutlass::epilogue::thread::ReLU,
        ElementD,
        ElementCompute,
        ElementAux,
        ElementAmax,
        ElementBias,
        ElementC>;

using CollectiveEpilogue =
    typename cutlass::epilogue::collective::CollectiveBuilder<
        ArchTag, OperatorClass,
        TileShape, ClusterShape, EpilogueTile,
        ElementAccumulator, ElementCompute,
        ElementC, LayoutC, AlignmentC,
        ElementD, LayoutD, AlignmentD,
        EpilogueSchedule,
        FusionOperation
    >::CollectiveOp;

using CollectiveMainloop =
    typename cutlass::gemm::collective::CollectiveBuilder<
        ArchTag, OperatorClass,
        ElementA, LayoutA, AlignmentA,
        ElementB, LayoutB, AlignmentB,
        ElementAccumulator,
        TileShape, ClusterShape,
        cutlass::gemm::collective::StageCountAutoCarveout<
            static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))
        >,
        KernelSchedule
    >::CollectiveOp;

using GemmKernel =
    cutlass::gemm::kernel::GemmUniversal<
        Shape<int, int, int, int>,
        CollectiveMainloop,
        CollectiveEpilogue
    >;

using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
using StrideA = typename GemmKernel::StrideA;
using StrideB = typename GemmKernel::StrideB;
using StrideC = typename GemmKernel::StrideC;
using StrideD = typename GemmKernel::StrideD;

template <class Element, class Layout>
void fill_random(cutlass::HostTensor<Element, Layout>& tensor, int seed) {
  int bits = cutlass::sizeof_bits<Element>::value;
  double scope = bits <= 8 ? 2.0 : 8.0;
  cutlass::reference::host::TensorFillRandomUniform(
      tensor.host_view(), seed, scope, -scope, 0);
  tensor.sync_device();
}

int main() {
  int M = 4096;
  int N = 4096;
  int K = 4096;
  int L = 1;

  auto stride_A = cutlass::make_cute_packed_stride(
      StrideA{}, cute::make_shape(M, K, L));
  auto stride_B = cutlass::make_cute_packed_stride(
      StrideB{}, cute::make_shape(N, K, L));
  auto stride_C = cutlass::make_cute_packed_stride(
      StrideC{}, cute::make_shape(M, N, L));
  auto stride_D = cutlass::make_cute_packed_stride(
      StrideD{}, cute::make_shape(M, N, L));

  cutlass::HostTensor<ElementA, LayoutA> tensor_A({M, K});
  cutlass::HostTensor<ElementB, LayoutB> tensor_B({K, N});
  cutlass::HostTensor<ElementC, LayoutC> tensor_C({M, N});
  cutlass::HostTensor<ElementD, LayoutD> tensor_D({M, N});

  fill_random(tensor_A, 2026);
  fill_random(tensor_B, 2027);
  fill_random(tensor_C, 2028);
  tensor_D.sync_device();

  typename Gemm::Arguments args{
    cutlass::gemm::GemmUniversalMode::kGemm,
    {M, N, K, L},
    {tensor_A.device_data(), stride_A,
     tensor_B.device_data(), stride_B},
    {
      {}, // fusion args; filled below
      tensor_C.device_data(), stride_C,
      tensor_D.device_data(), stride_D
    }
  };

  // Epilogue computes roughly:
  // D = activation(alpha * scale_a * scale_b * accumulator
  //              + beta * scale_c * C + bias)
  auto& fusion = args.epilogue.thread;
  fusion.alpha = 1.0f;
  fusion.beta = 0.0f;
  fusion.scale_a = 1.0f;
  fusion.scale_b = 1.0f;
  fusion.scale_c = 1.0f;
  fusion.scale_d = 1.0f;
  fusion.scale_aux = 1.0f;
  fusion.bias_ptr = nullptr;
  fusion.aux_ptr = nullptr;
  fusion.amax_D_ptr = nullptr;
  fusion.amax_aux_ptr = nullptr;

  Gemm gemm;
  size_t workspace_size = Gemm::get_workspace_size(args);
  cutlass::device_memory::allocation<uint8_t> workspace(workspace_size);

  cutlass::Status status = gemm.can_implement(args);
  if (status != cutlass::Status::kSuccess) {
    std::cerr << "GEMM configuration is not supported\n";
    return -1;
  }

  status = gemm.initialize(args, workspace.get());
  if (status != cutlass::Status::kSuccess) {
    std::cerr << "GEMM initialization failed\n";
    return -1;
  }

  status = gemm.run();
  if (status != cutlass::Status::kSuccess) {
    std::cerr << "GEMM run failed\n";
    return -1;
  }

  // cudaDeviceSynchronize();
  tensor_D.sync_host();
  std::cout << "Hopper FP8 GEMM finished\n";
  return 0;
}

12.4 blockwise scaling FP8 要多什么

上面的代码是“普通 FP8 A/B 输入”的主线骨架。LLM 推理中更常见的是 blockwise scaling:A/B 除了 FP8 数据外,还要带 scale tensor。CUTLASS Hopper 示例会把 A 和 scale-A 组合成一个 tuple layout,把 B 和 scale-B 组合成一个 tuple layout,mainloop 在 MMA 前把 scale 带入。

Concept sketch
// 概念差异:普通 FP8
ElementA, LayoutA, AlignmentA
ElementB, LayoutB, AlignmentB

// blockwise scaling FP8
ElementA, cute::tuple<LayoutA, LayoutSFA>, AlignmentA
ElementB, cute::tuple<LayoutB, LayoutSFB>, AlignmentB

在 Hopper blockwise FP8 示例里,scale layout 通常由 tile shape 推导,例如 sm90_trivial_blockwise_scale_config(TileShape{}) 再得到 LayoutSFA/LayoutSFB。这样 scale 的分块方式和 BM/BN/BK 绑定,避免 epilogue 或 mainloop 里临时做复杂索引。

把 FP8 GEMM 写到生产级时,CuTe/CUTLASS 的价值不是“少写几行代码”,而是把 tile、layout、TMA、MMA、scale、epilogue 的约束放进类型和 collective 里,让 kernel 更容易系统性调参。

12.5 Python CuTe DSL:运行入口与源码版本

Python CuTe DSL 的完整 GEMM kernel 不是几十行教学代码,而是包含 TMA descriptor、shared layout、pipeline barrier、TMEM、MMA、epilogue、reference check、benchmark 和 cold-L2 workspace 的工程实现。本节保留可运行命令,并从固定版本中摘取关键路径;完整文件仍以官方仓库为准。

下文代码固定到 NVIDIA CUTLASS commit d01487d(2026-07-16),并遵循仓库的 BSD-3-Clause license。固定 commit 很重要:CuTe DSL 的 API 和示例目录仍在演进,直接阅读 main 可能和本文行号不一致。
目标官方完整文件说明
普通 FP8 dense GEMMexamples/python/CuTeDSL/cute/blackwell/kernel/dense_gemm/dense_gemm.pyA/B 同 dtype,可用 Float8E4M3FNFloat8E5M2,accumulator 通常设 Float32
blockscaled FP8 / MXF8 GEMMexamples/python/CuTeDSL/cute/blackwell/kernel/blockscaled_gemm/dense_blockscaled_gemm_persistent.pyA/B 数据和 SFA/SFB scale tensor 一起参与 mainloop,更接近低精度推理里的真实路线

普通 FP8 dense GEMM 的完整运行命令:

Production command
cd ${CUTLASS_ROOT}

python examples/python/CuTeDSL/cute/blackwell/kernel/dense_gemm/dense_gemm.py \
  --ab_dtype Float8E4M3FN \
  --c_dtype Float16 \
  --acc_dtype Float32 \
  --a_major k \
  --b_major n \
  --c_major n \
  --mma_tiler_mn 128,128 \
  --cluster_shape_mn 1,1 \
  --mnkl 4096,4096,4096,1 \
  --use_tma_store \
  --warmup_iterations 1 \
  --iterations 10

blockscaled FP8 / MXF8 GEMM 的完整运行命令:

Production command
cd ${CUTLASS_ROOT}

python examples/python/CuTeDSL/cute/blackwell/kernel/blockscaled_gemm/dense_blockscaled_gemm_persistent.py \
  --a_dtype Float8E4M3FN \
  --b_dtype Float8E4M3FN \
  --sf_dtype Float8E8M0FNU \
  --sf_vec_size 32 \
  --c_dtype Float16 \
  --a_major k \
  --b_major n \
  --c_major n \
  --mma_tiler_mn 128,128 \
  --cluster_shape_mn 1,1 \
  --mnkl 4096,4096,4096,1 \
  --warmup_iterations 1 \
  --iterations 10

这里的 --a_major k 表示 A 的 K 维连续,等价于常见 row-major A;--b_major n 表示 B 的 N 维连续,也等价于常见 row-major B。参数名写的是逻辑连续维,不是直接写 RowMajor/ColumnMajor,所以读命令时要先结合矩阵形状 A[M,K]B[K,N] 判断。CuTe DSL 仍在快速演进;实际运行前应先对当前 checkout 执行 python .../dense_gemm.py --help,确认参数名与支持的 dtype。

12.6 普通 FP8 dense kernel:从 host 到 device

dense_gemm.py 并不是“Python 调一个现成 GEMM”。Python 代码先把 dtype、layout、tile 和 cluster 组织成 CuTe 对象,cute.compile 再把这些静态信息编译成目标 GPU kernel。下面的命令把 A/B dtype 设成 FP8;同一个 DenseGemmKernel 类也能被其他受支持 dtype 专门化。

Python CuTe DSL:编译路径与运行路径
run()
准备 tensor 与参数
cute.compile(...)
驱动 JIT 专门化
@cute.jit __call__
生成 layout / TMA / launch stub
compiled_gemm(...)
提交 GPU launch
@cute.kernel
TMA → tcgen05 → epilogue

12.6.1 run():先实例化,再编译

Official excerpt · dense_gemm.py L1639-L1663
a_tensor, b_tensor, c_tensor, a_torch_cpu, b_torch_cpu, c_torch_cpu, c_torch_gpu = (
    create_tensors(l, m, n, k, a_major, b_major, c_major, ab_dtype, c_dtype)
)

# Build GEMM object
gemm = DenseGemmKernel(
    acc_dtype, use_2cta_instrs, mma_tiler_mn, cluster_shape_mn, use_tma_store
)

# Check if configuration can be implemented
can_implement = gemm.can_implement(a_tensor, b_tensor, c_tensor)
if not can_implement:
    raise ValueError(
        f"The current config which is invalid/unsupported: use_2cta_instrs = {use_2cta_instrs}, "
        f"mma_tiler_mn = {mma_tiler_mn}, cluster_shape_mn = {cluster_shape_mn}, "
        f"use_tma_store = {use_tma_store}"
    )
max_active_clusters = utils.HardwareInfo().get_max_active_clusters(
    cluster_shape_mn[0] * cluster_shape_mn[1]
)
compiled_gemm = cute.compile(gemm, a_tensor, b_tensor, c_tensor, current_stream)

if not skip_ref_check:
    compiled_gemm(a_tensor, b_tensor, c_tensor, current_stream)
    compare(a_torch_cpu, b_torch_cpu, c_torch_gpu, c_dtype, tolerance)

查看固定版本中的原始行

create_tensors 同时准备 PyTorch tensor 和带动态 layout 的 cute.Tensor view。随后构造的 DenseGemmKernel 还不是已编译 kernel;它保存 accumulator dtype、MMA tiler、cluster shape 和 store 策略。can_implement 先验证 dtype、连续维对齐、tile/cluster 组合与输出边界,cute.compile 才依据对象配置及 A/B/C tensor 的类型和 layout 生成可调用的 compiled_gemm

cute.compile(gemm, ...) 以带 @cute.jitgemm.__call__ 为编译入口:编译器据此专门化 layout、TMA atom、grid 和 kernel launch stub,并产出 compiled_gemm。第一次 compiled_gemm(...) 用来做 reference check,benchmark 则复用这一编译结果。也就是说,Python 在这里主要负责配置、编译和调度,矩阵主循环仍在 @cute.kernel 生成的 GPU 代码里。

这个固定版本在 dense 路径里计算了 max_active_clusters,但没有把它传给 DenseGemmKernel;persistent blockscaled 路径才用它限制 grid。读工程示例时应沿参数实际传递关系判断作用,不能只根据局部变量名推断。

12.6.2 @cute.jit __call__:把 tensor 和 tile 变成 TMA 描述

Official excerpt · dense_gemm.py L377-L413
# Setup TMA load for A
a_op = sm100_utils.cluster_shape_to_tma_atom_A(
    self.cluster_shape_mn, tiled_mma.thr_id
)
a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0))
tma_atom_a, tma_tensor_a = cute.nvgpu.make_tiled_tma_atom_A(
    a_op,
    a,
    a_smem_layout,
    self.mma_tiler,
    tiled_mma,
    self.cluster_layout_vmnk.shape,
    internal_type=(
        cutlass.TFloat32 if a.element_type is cutlass.Float32 else None
    ),
)

# Setup TMA load for B
b_op = sm100_utils.cluster_shape_to_tma_atom_B(
    self.cluster_shape_mn, tiled_mma.thr_id
)
b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0))
tma_atom_b, tma_tensor_b = cute.nvgpu.make_tiled_tma_atom_B(
    b_op,
    b,
    b_smem_layout,
    self.mma_tiler,
    tiled_mma,
    self.cluster_layout_vmnk.shape,
    internal_type=(
        cutlass.TFloat32 if b.element_type is cutlass.Float32 else None
    ),
)

a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout)
b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout)
self.num_tma_load_bytes = (a_copy_size + b_copy_size) * atom_thr_size

查看固定版本中的原始行

make_tiled_tma_atom_A/B 把五类信息绑定在一起:global tensor、单 stage 的 shared layout、MMA tile、TiledMMA 的线程组织,以及 cluster layout。返回的 tma_atom_a 是 copy 操作描述,tma_tensor_a 是按照该描述重解释后的 global tensor view;二者不是已经搬入 shared memory 的数据。

self.num_tma_load_bytes 也不是“预估带宽”。后面的 TMA pipeline 会把它作为 barrier 的 transaction byte count:只有该 stage 对应的 A/B 字节都到达,消费者才能把 stage 判为 ready。A/B multicast 是否启用,则由 cluster shape 和 TMA atom 的选择共同决定。

12.6.3 Mainloop:等待 stage,发出 tcgen05 MMA,再释放 stage

Official excerpt · dense_gemm.py L771-L799
    if is_leader_cta:
        # Conditionally wait for AB buffer full
        consumer_handle = ab_consumer.wait_and_advance(peek_ab_full_status)

        # tCtAcc += tCrA * tCrB
        num_kblks = cute.size(tCrA, mode=[2])
        for kblk_idx in cutlass.range(num_kblks, unroll_full=True):
            kblk_crd = (None, None, kblk_idx, consumer_handle.index)

            cute.gemm(
                tiled_mma, tCtAcc, tCrA[kblk_crd], tCrB[kblk_crd], tCtAcc
            )
            # Enable accumulate on tCtAcc after first kblock
            tiled_mma.set(tcgen05.Field.ACCUMULATE, True)

        # Async arrive AB buffer empty
        consumer_handle.release()

    if k_tile_idx + 1 < k_tile_cnt - prefetch_k_tile_cnt:
        # Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1
        peek_ab_empty_status = ab_producer.try_acquire()

    if k_tile_idx + 1 < k_tile_cnt and is_leader_cta:
        # Peek (try_wait) AB buffer full for k_tile = k_tile + 1
        peek_ab_full_status = ab_consumer.try_wait()

# Async arrive accumulator buffer full
if is_leader_cta:
    acc_pipeline.producer_commit(acc_producer_state)

查看固定版本中的原始行

这段代码位于 producer 已经发出 TMA load 之后。wait_and_advance 返回当前可消费的 stage;consumer_handle.index 选择 shared-memory ring buffer 中对应的 A/B fragment。内层 kblk_idx 再把一个 CTA K tile 拆成 tcgen05 MMA 能直接处理的 K fragment。

tCtAcc 的名字里是 t,因为 Blackwell 版本把 accumulator 放在 TMEM,而不是普通线程寄存器数组。第一次 cute.gemm 后设置 ACCUMULATE=True,后续 MMA 才读旧 accumulator 并继续累加。consumer_handle.release() 把当前 A/B stage 归还 producer;K-loop 完成后,acc_pipeline.producer_commit 再通知 epilogue accumulator 已就绪。

12.6.4 Epilogue:TMEM → register → shared → global

Official excerpt · dense_gemm.py L1011-L1049
subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3])
for subtile_idx in cutlass.range(subtile_cnt):
    #
    # Load accumulator from tensor memory buffer to register
    #
    tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)]
    cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc)

    #
    # Perform epilogue op on accumulator and convert to C type
    #
    acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load()
    acc_vec = epilogue_op(acc_vec.to(self.c_dtype))
    tRS_rC.store(acc_vec)

    #
    # Store C to shared memory
    #
    c_buffer = subtile_idx % self.num_c_stage
    cute.copy(tiled_copy_r2s, tRS_rC, tRS_sC[(None, None, None, c_buffer)])
    # Fence and barrier to make sure shared memory store is visible to TMA store
    cute.arch.fence_proxy(
        "async.shared",
        space="cta",
    )
    pipeline.sync(barrier_id=1)

    # TMA store C to global memory
    if warp_idx == 0:
        cute.copy(
            tma_atom_c, bSG_sC[(None, c_buffer)], bSG_gC[(None, subtile_idx)]
        )
        # Fence and barrier to make sure TMA store is completed to recollect C buffer
        c_pipeline.producer_commit()
        c_pipeline.producer_acquire()
    pipeline.sync(barrier_id=1)

# Wait for C store complete
c_pipeline.producer_tail()

查看固定版本中的原始行

TMA store 不能直接读取 TMEM accumulator,因此输出要经过两次显式重排:tiled_copy_t2r 先把一个 epilogue subtile 从 TMEM 搬到 register,tiled_copy_r2s 再把转换后的结果写入 shared C buffer。epilogue_op 和 output dtype 转换发生在 register vector 上。

fence_proxy("async.shared") 与 CTA barrier 保证 shared store 对 TMA 可见,warp 0 随后发出 shared-to-global TMA store。PipelineTmaStore 控制 C buffer 何时可以复用。这里也能看到:即使 mainloop 已经结束,输出重排、类型转换和异步 store 仍是一条独立流水线。

12.7 Blockscaled FP8:SFA/SFB 如何进入 mainloop

普通 dense 示例的 kernel 参数只有 A、B、C;如果模型使用沿 K 分块的 scale,SFA/SFB 就不能留到 epilogue。Blackwell blockscaled kernel 把 scale factor 作为独立 tensor,与 A/B 一起经过 TMA、shared pipeline 和 TMEM,再由 block-scaled tcgen05 MMA 在每个 K tile 内消费。

维度普通 dense FP8Blockscaled FP8
Kernel operandA、B、C;不含 block scale tensorA、B、SFA、SFB、C
每个 K stage 搬运A tile + B tileA + B + 对应的 SFA + SFB
MMA operandA_fragmentB_fragment[A_fragment, SFA][B_fragment, SFB]
Scale 应用位置kernel 内没有 K-block scale位于 K-loop 内,由 block-scaled MMA 消费
调度固定 grid 的 dense kernelpersistent tile scheduler + 专用 TMA/MMA/epilogue warps
Blockscaled persistent kernel 的数据流
TMA warp
A / B / SFA / SFB
Staged SMEM
同一 barrier 管理
MMA warp
SFA/SFB: SMEM → TMEM
tcgen05 MMA
[A,SFA] × [B,SFB]
Epilogue warps
TMEM → C

12.7.1 Warp specialization:192 个线程不是做同一件事

Official excerpt · blockscaled persistent L194-L218
self.acc_dtype = cutlass.Float32
self.sf_vec_size = sf_vec_size
self.use_2cta_instrs = mma_tiler_mn[0] == 256
self.cluster_shape_mn = cluster_shape_mn
# K dimension is deferred in _setup_attributes
self.mma_tiler = (*mma_tiler_mn, 1)

self.cta_group = (
    tcgen05.CtaGroup.TWO if self.use_2cta_instrs else tcgen05.CtaGroup.ONE
)

self.occupancy = 1
# Set specialized warp ids
self.epilog_warp_id = (
    0,
    1,
    2,
    3,
)
self.mma_warp_id = 4
self.tma_warp_id = 5
self.threads_per_warp = 32
self.threads_per_cta = self.threads_per_warp * len(
    (self.mma_warp_id, self.tma_warp_id, *self.epilog_warp_id)
)

查看固定版本中的原始行

一个 CTA 有 6 个 warp:warp 0-3 负责 epilogue,warp 4 专门发出 MMA,warp 5 专门发出 TMA。TMA warp 可以继续为后续 K stage 或后续 output tile 搬数据,而 MMA warp 消费当前 stage;四个 epilogue warp 则并行处理 TMEM accumulator 的输出 subtile。这是 warp specialization,不是把 192 个线程平均分给一个标量循环。

kernel 还是 persistent 的:TMA、MMA 和 epilogue 角色各自维护 StaticPersistentTileScheduler 状态,在完成一个 output tile 后继续领取下一 tile。三个角色通过 A/B pipeline、accumulator pipeline、TMEM allocation barrier 和 epilogue barrier 对齐生命周期。

12.7.2 SFA/SFB 有自己的逻辑 layout 和 TMA atom

Official excerpt · blockscaled persistent L478-L500
# Setup sfa/sfb tensor by filling A/B tensor to scale factor atom layout
# ((Atom_M, Rest_M),(Atom_K, Rest_K),RestL)
sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(
    a_tensor.shape, self.sf_vec_size
)
sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)

# ((Atom_N, Rest_N),(Atom_K, Rest_K),RestL)
sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(
    b_tensor.shape, self.sf_vec_size
)
sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)

tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma(
    self.a_dtype,
    self.b_dtype,
    self.a_major_mode,
    self.b_major_mode,
    self.sf_dtype,
    self.sf_vec_size,
    self.cta_group,
    self.mma_inst_shape_mn,
)

查看固定版本中的原始行

tile_atom_to_shape_SF 根据 A/B 的逻辑 shape 和 sf_vec_size 构造 scale tensor:SFA 与 A 的 M/K 分块对应,SFB 与 B 的 N/K 分块对应。它们不是把一个一维 scale 数组随便挂到 GEMM 参数上;layout 必须表达“哪个 scale 覆盖哪段 K 数据”,并匹配 block-scaled MMA 的 scale-factor atom。

Official excerpt · blockscaled persistent L548-L580
# Setup TMA load for SFA
sfa_op = sm100_utils.cluster_shape_to_tma_atom_A(
    self.cluster_shape_mn, tiled_mma.thr_id
)
sfa_smem_layout = cute.slice_(
    self.sfa_smem_layout_staged, (None, None, None, 0)
)
tma_atom_sfa, tma_tensor_sfa = cute.nvgpu.make_tiled_tma_atom_A(
    sfa_op,
    sfa_tensor,
    sfa_smem_layout,
    self.mma_tiler,
    tiled_mma,
    self.cluster_layout_vmnk.shape,
    internal_type=cutlass.Int16,
)

# Setup TMA load for SFB
sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB(
    self.cluster_shape_mn, tiled_mma.thr_id
)
sfb_smem_layout = cute.slice_(
    self.sfb_smem_layout_staged, (None, None, None, 0)
)
tma_atom_sfb, tma_tensor_sfb = cute.nvgpu.make_tiled_tma_atom_B(
    sfb_op,
    sfb_tensor,
    sfb_smem_layout,
    self.mma_tiler_sfb,
    tiled_mma_sfb,
    self.cluster_layout_sfb_vmnk.shape,
    internal_type=cutlass.Int16,
)

查看固定版本中的原始行

SFA 复用 A 方向的 cluster multicast 规则,SFB 使用面向 B/列方向的专用 cluster_shape_to_tma_atom_SFB。两者最终都生成独立 TMA atom 和 tensor view。这里的 internal_type=cutlass.Int16 描述 TMA/SMEM 搬运时的内部表示,不等于把 scale 的逻辑数值类型改成 INT16;逻辑类型仍由 self.sf_dtype 决定。

12.7.3 Producer:一个 stage 必须同时等到 A/B/SFA/SFB

Official excerpt · blockscaled persistent L1070-L1104
for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1):
    # Conditionally wait for AB buffer empty
    ab_pipeline.producer_acquire(
        ab_producer_state, peek_ab_empty_status
    )

    # TMA load A/B/SFA/SFB
    cute.copy(
        tma_atom_a,
        tAgA_slice[(None, ab_producer_state.count)],
        tAsA[(None, ab_producer_state.index)],
        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
        mcast_mask=a_full_mcast_mask,
    )
    cute.copy(
        tma_atom_b,
        tBgB_slice[(None, ab_producer_state.count)],
        tBsB[(None, ab_producer_state.index)],
        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
        mcast_mask=b_full_mcast_mask,
    )
    cute.copy(
        tma_atom_sfa,
        tAgSFA_slice[(None, ab_producer_state.count)],
        tAsSFA[(None, ab_producer_state.index)],
        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
        mcast_mask=sfa_full_mcast_mask,
    )
    cute.copy(
        tma_atom_sfb,
        tBgSFB_slice[(None, ab_producer_state.count)],
        tBsSFB[(None, ab_producer_state.index)],
        tma_bar_ptr=ab_pipeline.producer_get_barrier(ab_producer_state),
        mcast_mask=sfb_full_mcast_mask,
    )

查看固定版本中的原始行

四次 cute.copy 使用同一个 stage barrier。该 kernel 在创建 pipeline 时把 A + B + SFA + SFB 的字节数都计入 tx_count,因此消费者不会在“数据到了、scale 还没到”的中间状态开始 MMA。ab_producer_state.index 选择 ring buffer 的物理 stage,count 则选择逻辑上的下一个 K tile。

12.7.4 Consumer:scale 从 SMEM 进入 TMEM,并成为 MMA operand

Official excerpt · blockscaled persistent L1252-L1289
for k_tile in range(k_tile_cnt):
    if is_leader_cta:
        # Conditionally wait for AB buffer full
        ab_pipeline.consumer_wait(
            ab_consumer_state, peek_ab_full_status
        )

        #  Copy SFA/SFB from smem to tmem
        s2t_stage_coord = (
            None,
            None,
            None,
            None,
            ab_consumer_state.index,
        )
        tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[s2t_stage_coord]
        tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[s2t_stage_coord]
        cute.copy(
            tiled_copy_s2t_sfa,
            tCsSFA_compact_s2t_staged,
            tCtSFA_compact_s2t,
        )
        cute.copy(
            tiled_copy_s2t_sfb,
            tCsSFB_compact_s2t_staged,
            tCtSFB_compact_s2t,
        )

        # tCtAcc += tCrA * tCrSFA * tCrB * tCrSFB
        tiled_mma.set(tcgen05.Field.ACCUMULATE, k_tile != 0)
        tile_crd = (None, None, None, ab_consumer_state.index)
        cute.gemm(
            tiled_mma,
            tCtAcc,
            [tCrA[tile_crd], tCtSFA],
            [tCrB[tile_crd], tCtSFB_mma],
            tCtAcc,
        )

查看固定版本中的原始行

MMA warp 先等待当前 A/B stage ready,再用两次 S2T copy 把这个 stage 的 SFA/SFB 从 shared memory 放进 TMEM scale layout。随后 cute.gemm 接收的不是单独 A/B fragment,而是 [A, SFA][B, SFB] 两个 paired operand。

ACCUMULATEk_tile 更新,说明每一轮 K tile 都使用与该轮匹配的 scale factor,再累加到 FP32 TMEM accumulator。这正是 §11.7 公式 Σ partial[kb] × scale_A[kb] × scale_B[kb] 的硬件实现;把 SFA/SFB 留到完整 K reduction 之后再乘一次,在这里既不符合代码,也不符合数学语义。

12.8 读 Python CuTe DSL 时先分清这些对象

对象它实际表示什么常见误读
@cute.jit由 CuTe 编译器处理的 host/JIT 组装路径误以为只是普通 Python helper
@cute.kernelGPU device kernel 本体从第一行顺序阅读,忽略 warp specialization
cute.Tensoriterator/pointer 与 layout 的组合 view误以为调用时复制了一份 tensor
SMEM / RMEM / TMEMshared memory、线程寄存器、Blackwell Tensor Memory把所有 fragment 都叫“寄存器”
cute.copy由 copy atom、source/destination layout 决定的搬运误以为等价于同步的 Python copy
cute.gemm生成 tcgen05 MMA 的 DSL primitive误以为是高层框架的 eager matmul
Pipeline statestage index、phase、count 与 mbarrier 生命周期只看循环下标,漏掉 buffer 所有权转移
建议先沿“host compile → TMA producer → MMA consumer → epilogue”四条路径分别读,再回头看完整 kernel。不要逐行穿过 3000 多行文件:那会把 layout 构造、persistent scheduler、数值 reference 和 benchmark scaffolding 混成一条难以验证的叙事。

13. Triton GEMM:用 Python 写 tile kernel

Triton 位于“直接调用库”和“手写 CUDA/CuTe kernel”之间:它保留了 tile、program id、mask、accumulator、autotune 这些关键控制点,但把线程级地址计算、LLVM/MLIR lowering 和具体后端代码生成隐藏掉。对教学来说,Triton 很适合把前面讲的 BM/BN/BK、tail mask、tl.dot 和 accumulator 连起来。

Triton 不是“更简单的 cuBLAS”。它适合 shape 特殊、需要融合、想快速试 kernel 结构的场景;如果只是标准大 GEMM,成熟库仍然经常更稳、更快。

13.1 一个 Triton matmul kernel 对应前文哪几层

本文概念Triton 写法CUDA / CUTLASS 类比
CTA 负责一个 C tiletl.program_id(0) 映射到 (pid_m, pid_n)blockIdx / thread block tile
BM/BN/BKBLOCK_MBLOCK_NBLOCK_KTileShape<BM, BN, BK>
global load + tailtl.load(..., mask=...)边界判断、predicated load
MMA / Tensor Coretl.dot(a, b)mma.sync / WGMMA / collective mainloop
accumulatortl.zeros((BM, BN), tl.float32)register accumulator fragment
调参BLOCK_*num_warpsnum_stages、autotune configstile shape、warp 数、pipeline stage、heuristic

13.2 教学版 Triton GEMM

下面这段代码保留了 Triton GEMM 的核心结构:一个 program 负责一个 BLOCK_M x BLOCK_N 的 C tile,沿 K 维分块循环,每轮加载 A/B 子块,用 tl.dot 累加,最后写回 C。真实项目还会加 grouped ordering、autotune、dtype 选择、epilogue fusion 和 benchmark。

查看完整 Triton GEMM 教学 kernel
Teaching kernel · Triton
import torch
import triton
import triton.language as tl


@triton.jit
def matmul_kernel(
    a_ptr, b_ptr, c_ptr,
    M, N, K,
    stride_am, stride_ak,
    stride_bk, stride_bn,
    stride_cm, stride_cn,
    BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr,
):
    pid = tl.program_id(axis=0)
    num_pid_n = tl.cdiv(N, BLOCK_N)
    pid_m = pid // num_pid_n
    pid_n = pid % num_pid_n

    offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M)
    offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N)
    offs_k = tl.arange(0, BLOCK_K)

    a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak
    b_ptrs = b_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn
    acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32)

    for k_block in range(0, tl.cdiv(K, BLOCK_K)):
        a = tl.load(
            a_ptrs,
            mask=(offs_m[:, None] < M) &
                 (offs_k[None, :] < K - k_block * BLOCK_K),
            other=0.0,
        )
        b = tl.load(
            b_ptrs,
            mask=(offs_k[:, None] < K - k_block * BLOCK_K) &
                 (offs_n[None, :] < N),
            other=0.0,
        )
        acc += tl.dot(a, b)
        a_ptrs += BLOCK_K * stride_ak
        b_ptrs += BLOCK_K * stride_bk

    c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn
    tl.store(c_ptrs, acc, mask=(offs_m[:, None] < M) & (offs_n[None, :] < N))


def matmul(a, b):
    assert a.is_cuda and b.is_cuda
    assert a.ndim == 2 and b.ndim == 2
    assert a.dtype == b.dtype
    assert a.dtype in (torch.float16, torch.bfloat16)
    M, K = a.shape
    K2, N = b.shape
    assert K == K2
    c = torch.empty((M, N), device=a.device, dtype=a.dtype)

    grid = (triton.cdiv(M, 128) * triton.cdiv(N, 128),)
    matmul_kernel[grid](
        a, b, c, M, N, K,
        a.stride(0), a.stride(1),
        b.stride(0), b.stride(1),
        c.stride(0), c.stride(1),
        BLOCK_M=128, BLOCK_N=128, BLOCK_K=64,
        num_warps=4, num_stages=4,
    )
    return c

M/N/K 和 stride 作为运行时参数传入,避免每个 shape 都生成一份专门化 kernel;BLOCK_M/N/Knum_warpsnum_stages 才是编译期 meta-parameter。tl.store 会按 c_ptr 的元素类型把 FP32 accumulator 转成 FP16 或 BF16。

为什么要保留 mask

真实 shape 往往不能整除 BM/BN/BK。Triton 的 tl.load / tl.store mask 把 tail 处理写进向量化表达式里,避免为边界 tile 单独写一个 kernel。

为什么 accumulator 用 FP32

BF16/FP16 输入通常仍以 FP32 累加,然后在 epilogue cast 到目标 dtype。低精度 GEMM 的数值质量很大程度取决于 accumulator dtype 和 epilogue 的转换策略。

13.3 从 program_id 到 C tile

program_id
一维 grid 编号
pid_m / pid_n
映射到 C 的 tile 坐标
offs_m / offs_n
生成 M/N 向量索引
K loop
每轮推进 BLOCK_K
tl.dot
累加到 C tile

上面的 grid 是一维的:pid 先除以 N 方向 tile 数得到 pid_m,再取模得到 pid_n。这和 CUDA 里用二维 blockIdx.y/blockIdx.x 没有本质区别,只是 Triton 教程常用一维 program id 来方便做 grouped ordering。

13.4 Autotune 调什么

参数影响profile 上常见信号
BLOCK_M / BLOCK_NC tile 面积、A/B 复用、accumulator 大小Tensor Core 利用率低、CTA 数不足、register pressure 高
BLOCK_K每轮 K 维深度、load 粒度、dot 输入块大小HBM 事务过碎、shared/register 压力变化
num_warps一个 program 内的 warp 资源occupancy、stall、每 tile 并行度
num_stages软件 pipeline 深度memory dependency stall、shared memory 用量、occupancy

Triton 官方 matmul 教程会给出多组 triton.Config,CUDA 和 HIP 后端的推荐候选也不同。经验上,不要把 NVIDIA 上的 config 原样搬到 AMD;同一个 BLOCK_M/N/K,寄存器、LDS、wavefront、MFMA 路径和编译器 lowering 都可能让最优点变化。

13.5 Triton 和 CUTLASS / 手写 CUDA 的取舍

路线优势代价适合
Triton迭代快,容易融合,Python 侧接入方便极限控制弱于 CUDA/CuTe,性能依赖编译器和后端成熟度shape 专用 kernel、融合算子、研究和线上快速试错
CUTLASS / CuTe对 NVIDIA 新架构特性覆盖深,mainloop/epilogue 可组合C++ 模板复杂,学习成本高生产级 NVIDIA GEMM、FP8、TMA/WGMMA、复杂 epilogue
手写 CUDA/HIP控制最细,可以做非常规调度维护成本最高,跨架构成本高库和 DSL 覆盖不了的极特殊 kernel

14. hipBLAS / hipBLASLt:AMD ROCm 里的 GEMM 库路线

在 ROCm 生态里,GEMM 不只有一种入口。hipBLAS 更像 BLAS API 的可移植门面:应用调用 hipBLAS,下面可以接 AMD 的 rocBLAS,也可以在 CUDA 平台接 cuBLAShipBLASLt 则是更灵活的 matmul API,重点放在 layout、algorithm、heuristic、workspace、tuning 和 epilogue 上。

术语边界:前文的 Tensor Core、warp、shared memory、ldmatrixcp.async 主要是 NVIDIA 语境;AMD 常见对应词是 Matrix Core / MFMA、wave、LDS 和 ds_read/ds_write。tile、复用、pipeline、bank conflict 这些设计原则可以迁移,指令形状、lane mapping 和 layout 不能一一照抄。
ROCm GEMM 的三条并列路径,不是一条串行调用链
Framework / Application
PyTorch · vLLM · 自研 runtime
按可移植性、融合需求和目标 shape 选择入口
hipBLAS标准 BLAS 兼容门面
rocBLAS / cuBLASAMD / NVIDIA 平台后端
hipBLASLtdescriptor · heuristic · epilogue
Optimized algorithm pool按 shape、dtype、workspace 选择 kernel
Triton / CK / AITERDSL、组件库或 AI kernel 集合
Custom GPU kernels为融合和热点 shape 特化

hipBLASLt 不是 rocBLAS 调用链的下一层,Triton/CK/AITER 也不是 hipBLASLt 的后端。它们是应用可以分别选择、并在同一框架中并存的实现路线。

14.1 先把几个名字分清

名字它是什么更接近的 NVIDIA 概念读者该怎么用
rocBLASROCm 的底层 BLAS 实现库cuBLAS标准 GEMM baseline,通常被上层框架间接调用
hipBLASBLAS marshaling / portability API,可转发到 rocBLAS 或 cuBLAS跨平台 BLAS 门面想写 HIP/CUDA 双后端基础 GEMM 时使用
hipBLASLt更灵活的 GEMM API,支持 descriptor、algorithm、heuristic、workspace、tuningcuBLASLt需要 layout/epilogue/tuning 时优先看
Composable KernelROCm 的 C++ kernel/template 组件库部分角色接近 CUTLASS需要构造或理解 AMD 定制 kernel 时看
AITER面向 AI workload 的 AMD 优化算子与 kernel 集合更接近领域 kernel library框架已有接入,或热点 shape/融合已被覆盖时对照测试

14.2 hipBLAS:标准 GEMM baseline

hipBLAS 的优势是接口熟悉、迁移成本低。注意 BLAS API 的传统约定偏 column-major;如果你的张量是 row-major,常见处理方式是交换 A/B、调整转置标志,或在上层 layout 上显式处理 stride。

Concept sketch · hipBLAS SGEMM
hipblasHandle_t handle;
hipblasCreate(&handle);

const float alpha = 1.0f;
const float beta = 0.0f;

// BLAS uses column-major convention. For row-major C[M, N] = A[M, K] * B[K, N],
// many examples call GEMM as column-major C^T[N, M] = B^T[N, K] * A^T[K, M].
hipblasSgemm(
    handle,
    HIPBLAS_OP_N, HIPBLAS_OP_N,
    N, M, K,
    &alpha,
    B, N,
    A, K,
    &beta,
    C, N);

hipblasDestroy(handle);

这段代码适合作 baseline,不适合解释高性能 GEMM 的全部细节。真实性能来自库内部选到了什么 kernel、对应什么 dtype、layout、stride、workspace 和架构,而不是来自这几行 API 本身。

14.3 hipBLASLt:更像生产 matmul 的入口

LLM 推理里的 GEMM 往往不是裸乘法:可能有 bias、activation、scale、输出 dtype 转换,也可能要为固定 shape 做离线调优。hipBLASLt 的心智模型就是把一次 matmul 拆成 descriptor、matrix layout、preference、heuristic result 和实际 launch。

Concept sketch · hipBLASLt API flow
create hipblasLt handle

create matmul descriptor
  - compute type
  - transpose flags
  - epilogue mode
  - bias / auxiliary pointer if needed

create matrix layouts for A, B, C, D
  - dtype
  - rows / cols
  - leading dimension
  - batch stride if batched

create preference
  - max workspace bytes
  - algorithm search constraints

query heuristic algorithms
allocate workspace
run hipblasLtMatmul(...)
destroy descriptors
hipBLASLt 比 hipBLAS 多了什么

更多 layout 描述、algorithm 选择、workspace 控制、logging/heuristic/tuning 工具,以及更接近生产推理的 epilogue 配置。

它仍然不是 CUTLASS

hipBLASLt 是库 API,重点是选择和调用已有算法;CUTLASS/CK 更像构造 kernel 的模板工具。两者解决的问题层级不同。

14.4 和 CUDA / Triton / CUTLASS 路线怎么比较

路线抽象层什么时候优先用主要限制
cuBLAS / rocBLAS / hipBLAS标准 BLAS 调用快速建立 baseline,标准 GEMM融合和 layout 控制有限
cuBLASLt / hipBLASLtdescriptor + heuristic + algorithm需要 epilogue、workspace、算法选择和 shape tuning仍受库已实现 kernel 覆盖范围约束
CUTLASS / CKC++ 模板 / kernel 组件要构造生产 kernel,控制 mainloop/epilogue/layout代码复杂,编译和调参成本高
TritonPython kernel DSL快速做定制融合和 shape 专用 kernel极限性能和新硬件特性依赖后端成熟度
手写 CUDA/HIP最低层 kernel需要完全自定义调度或验证硬件机制维护和跨架构迁移成本最高

14.5 在 LLM 推理里怎么选

需求建议路线理由
先确认 AMD/NVIDIA 上 GEMM 是否正常hipBLAS / rocBLAS / cuBLAS最小变量,方便排查 dtype、layout、stride 问题
标准大 GEMM 要稳定性能hipBLASLt / cuBLASLt库的 heuristic 和 tuned algorithms 往往先给出很强 baseline
小 M、MoE、decode shape、特殊 batchhipBLASLt tuning + Triton / CK 对照标准大 GEMM 的最优策略不一定适合 serving shape
需要融合 scale、bias、activation、requantizehipBLASLt epilogue / Triton / CK / CUTLASS瓶颈可能从 mainloop 转移到 epilogue,必须把融合路径纳入 benchmark
跨后端框架适配先用 hipBLAS 建公共路径,再为热点 shape 下沉baseline 保持可移植,热点 kernel 再按架构特化
ROCm 上做 GEMM 调优时,最实用的工作流不是“选一个库后永远不换”,而是先用 hipBLAS/rocBLAS 建正确性和 baseline,再用 hipBLASLt、Triton、CK 或 AITER 对热点 shape 做交叉 benchmark。

14.6 ROCm 上怎么 profile

与 §10 的 CUDA 工作流类似,ROCm 也应先分清“端到端时间线问题”和“单 kernel 硬件瓶颈”。rocprofv3 适合采集 HIP API、kernel dispatch、memory copy 和指定硬件 counter;rocprof-compute 在这些数据之上提供 Speed-of-Light、memory hierarchy、roofline 和结果对比。

Production profiling command
# 先看 HIP runtime、kernel dispatch 和 memory copy 时间线
rocprofv3 --runtime-trace -- ./your_benchmark

# 再采集单 kernel 分析;工具会多次运行 workload 收集 counter
rocprof-compute profile \
  --name gemm_profile \
  --no-roof \
  -- ./your_benchmark

# 分析采集结果;实际 SoC 子目录由工具生成
rocprof-compute analyze -p workloads/gemm_profile/<soc>/
CUDA 侧问题ROCm 对应观察GEMM 调参关联
Tensor Core 是否饱和MFMA / Matrix Core 指令与 compute throughput检查 dtype 路径、tile、wave 分工和 MFMA 指令形状
Shared bank conflictLDS access / bank conflict 与 ds_read 模式检查 LDS padding、XOR preshuffle、wave lane mapping
HBM/L2/L1 压力memory hierarchy throughput 与 roofline检查 A/B 复用、global load 合并、epilogue traffic
Occupancy / spillwave occupancy、VGPR/SGPR/LDS 用量减小 tile、stage 或临时 fragment,检查 scratch/spill
ROCm counter 名称和可同时采集的组合取决于 gfx 架构与 ROCm 版本。先运行 rocprofv3-avail list / rocprof-compute profile --list-available-metrics,再为目标 MI 系列选择指标,不要把另一代 GPU 的 counter 名称写死到脚本里。

15. 实战 Checklist

  • 先确认 shape:M/N/K 是否足够大,是否存在小 M 或 tail block。
  • 先用 cuBLAS / cuBLASLt / CUTLASS / Triton / hipBLAS / hipBLASLt 建 baseline。
  • BF16 先检查 Tensor Core / MFMA matrix-core path 是否生效,accumulator 是否符合精度预算。
  • 调 tile size 时同时看 shared memory、register pressure 和 occupancy。
  • FP8 先定 scale 粒度,再定 kernel layout;不要把 scale 当成事后补丁。
  • CUDA 先用 nsys 判断是不是 kernel 本身慢,再用 ncu 看单 kernel;ROCm 对应先用 rocprofv3,再用 rocprof-compute
  • profile 时把 mainloop 和 epilogue 分开看,很多 FP8 kernel 瓶颈会转移到 epilogue。
工程建议:自己写 kernel 前,先把成熟库的性能边界摸清楚。只有在 shape 特殊、融合需求强、或框架 overhead 明确时,手写 GEMM 才更划算。

16. 常见误区

误区为什么不对更好的判断
BM/BN/BK 有固定最佳值tile 受 shape、dtype、SM 资源、epilogue、架构共同约束把经验值当候选,用 profile 选
occupancy 越高越好高 occupancy 可能来自每个线程工作太少,Tensor Core 仍然吃不满同时看 Tensor Core 利用率、stall reason、HBM 带宽
shared memory 一定更快shared 也有 bank conflict、同步、容量和 occupancy 成本确认 shared layout 服务于后续 warp/MMA 读取
FP8 只是把 dtype 从 BF16 改成 FP8FP8 的 scale、amax、requantize、clipping 会改变 mainloop 和 epilogue先定 scale 粒度,再定 layout 和 kernel
看平均 TFLOP/s 就够LLM 推理有小 M、tail、MoE、decode/prefill 混合,平均数掩盖瓶颈按 shape bucket 分开 benchmark
CuTe/CUTLASS 能自动解决所有问题DSL 降低表达成本,但 tile、layout、schedule、epilogue 仍要选择把 DSL 当调参框架,不是黑盒
Triton 一定比库快Triton 迭代快,但标准大 GEMM 上库 kernel 和 heuristic 可能更成熟对同一 shape 同时 benchmark Triton、BLASLt 和框架内置 kernel
hipBLASLt 就是 AMD 版 CUTLASShipBLASLt 是可调库 API,CUTLASS/CK 是更接近 kernel 构造的模板/组件库按抽象层选择工具:调用库、调算法、还是写 kernel

17. 参考资料

继续深入时,优先看官方文档和官方示例,因为 CUDA/CUTLASS 的接口和硬件细节会随版本更新。