template <int kTileM, int kTileN, int kTileK, typename TiledMMA> __global__ voidsimple_gemm_kernel(half_t* Dptr, consthalf_t* Aptr, consthalf_t* Bptr, int M, int N, int K){ // Global Tensor auto A = make_tensor(make_gmem_ptr(Aptr), make_shape(M, K), make_stride(K, Int<1>{})); auto B = make_tensor(make_gmem_ptr(Bptr), make_shape(N, K), make_stride(K, Int<1>{})); auto D = make_tensor(make_gmem_ptr(Dptr), make_shape(M, N), make_stride(N, Int<1>{}));
int bx = blockIdx.x; int by = blockIdx.y;
// 取当前 block 的 tile auto gA = local_tile(A, make_tile(Int<kTileM>{}, Int<kTileK>{}), make_coord(by, _)); auto gB = local_tile(B, make_tile(Int<kTileN>{}, Int<kTileK>{}), make_coord(bx, _)); auto gD = local_tile(D, make_tile(Int<kTileM>{}, Int<kTileN>{}), make_coord(by, bx));
TiledMMA tiled_mma; auto thr_mma = tiled_mma.get_slice(threadIdx.x);
// 切分 global tensor 到线程级别 auto tAgA = thr_mma.partition_A(gA); // (MMA, MMA_M, MMA_K, num_tile_k) auto tBgB = thr_mma.partition_B(gB); // (MMA, MMA_N, MMA_K, num_tile_k) auto tDgD = thr_mma.partition_C(gD); // (MMA, MMA_M, MMA_N)
// 创建寄存器 fragment auto tCrA = thr_mma.partition_fragment_A(gA(_, _, 0)); // (MMA, MMA_M, MMA_K) auto tCrB = thr_mma.partition_fragment_B(gB(_, _, 0)); // (MMA, MMA_N, MMA_K) auto tCrD = thr_mma.partition_fragment_C(gD); // (MMA, MMA_M, MMA_N)
// 清零累加器 clear(tCrD);
// K 方向循环 int num_tile_k = size<3>(tAgA); for (int ik = 0; ik < num_tile_k; ++ik) { // Global -> Register cute::copy(tAgA(_, _, _, ik), tCrA); cute::copy(tBgB(_, _, _, ik), tCrB);
for (int ik = 0; ik < size<3>(tAgA); ++ik) { printf("\n --- K tile %d (ik=%d) ---\n", ik, ik); for (int mma_m = 0; mma_m < size<1>(tAgA); ++mma_m) { for (int mma_k = 0; mma_k < size<2>(tAgA); ++mma_k) { printf(" MMA_M=%d, MMA_K=%d: ", mma_m, mma_k); for (int v = 0; v < size<0>(tAgA); ++v) { int off = int(&tAgA(v, mma_m, mma_k, ik) - Aptr); int m = off / K; int k = off % K; printf("A[%3d][%2d] ", m, k); } printf("\n"); } } }
// 打印 tCrA:寄存器 fragment 的 layout 和内容 printf("\n === tCrA 寄存器 fragment (线程 0) ===\n"); printf(" shape: "); print(shape(tCrA)); printf("\n"); printf(" layout: "); print(tCrA.layout()); printf("\n"); printf(" size: %d 个寄存器\n", int(size(tCrA))); printf("\n tCrA vs tAgA 的对应关系:\n"); printf(" tCrA shape: (MMA=%d, MMA_M=%d, MMA_K=%d)\n", int(size<0>(tCrA)), int(size<1>(tCrA)), int(size<2>(tCrA))); printf(" tAgA shape: (MMA=%d, MMA_M=%d, MMA_K=%d, num_tile_k=%d)\n", int(size<0>(tAgA)), int(size<1>(tAgA)), int(size<2>(tAgA)), int(size<3>(tAgA))); printf("\n tCrA 中每个元素的寄存器编号和逻辑位置:\n"); for (int mma_m = 0; mma_m < size<1>(tCrA); ++mma_m) { for (int mma_k = 0; mma_k < size<2>(tCrA); ++mma_k) { printf(" MMA_M=%d, MMA_K=%d: ", mma_m, mma_k); for (int v = 0; v < size<0>(tCrA); ++v) { // tCrA 的 layout 把 (v, mma_m, mma_k) 映射到寄存器中的 flat index int reg_idx = int(tCrA.layout()(v, mma_m, mma_k)); printf("reg[%2d] ", reg_idx); } printf("\n"); } } }
using mma_op = SM80_16x8x16_F16F16F16F16_TN; // Atom using mma_atom = MMA_Atom<MMA_Traits<mma_op>>; using TiledMMA = decltype(make_tiled_mma( mma_atom{}, make_layout(make_shape(Int<2>{}, Int<2>{}, Int<1>{})), // EU Repeat Tile<Int<32>, Int<32>, Int<16>>{})); // P Tile
2.1 扩展过程
1 2 3 4 5
M 方向 N 方向 K 方向 线程数 ───────── ──────── ──────── ──────── ① Atom (16×8×16) 16 行 8 列 16 列 32 ② EU Repeat (2,2,1) ×2 = 32 ×2 = 16 不变 ×4 = 128 ③ P Tile <32,32,16> 32/32=不变 32/16=×2 16/16=不变 不变
2.2 EU Repeat vs P Tile 的本质区别
EU Repeat
P Tile
线程数
增加(32→128)
不变(128)
mma.sync 覆盖面积
真的变大
不变
方式
多 warp 并行执行
同一线程多轮复用数据
EU Repeat 加人干活,P Tile 让同一批人多干活。
2.3 P Tile 对 A、B、C 的影响
矩阵
P Tile 影响
说明
A (M×K)
M: 32/32=1, K: 16/16=1, 不扩展
P Tile 对 A 没有影响
B (N×K)
N: 32/16=2, 扩展
每个线程在 N 方向多读 B
C/D (M×N)
N: 32/16=2, 扩展
每个线程多算 N 方向的输出
2.4 P Tile 的本质:寄存器级 A 复用
1 2 3 4 5 6
没有 P Tile (EU 后覆盖 32×16): D[32×16] = A[32×16] × B[16×16]^T ← A 用完就扔
有 P Tile N=32 (覆盖 32×32): D[32×16_左] = A[32×16] × B_左[16×16]^T ← A 加载到寄存器 D[32×16_右] = A[32×16] × B_右[16×16]^T ← A 还在寄存器里,不用重新加载
不是重复计算:左右两块的 B 不同,输出 D 也不同,是新计算
复用 A 寄存器:一份 A 配多份 B,省掉 A 的重复加载
代价:多占 B 和 D 的寄存器,可能降低 occupancy
3. Kernel 逐段解析
3.1 全局 Tensor + local_tile
1 2 3 4 5 6 7
auto A = make_tensor(make_gmem_ptr(Aptr), make_shape(M, K), make_stride(K, Int<1>{})); auto B = make_tensor(make_gmem_ptr(Bptr), make_shape(N, K), make_stride(K, Int<1>{})); auto D = make_tensor(make_gmem_ptr(Dptr), make_shape(M, N), make_stride(N, Int<1>{}));
auto gA = local_tile(A, make_tile(Int<kTileM>{}, Int<kTileK>{}), make_coord(by, _)); auto gB = local_tile(B, make_tile(Int<kTileN>{}, Int<kTileK>{}), make_coord(bx, _)); auto gD = local_tile(D, make_tile(Int<kTileM>{}, Int<kTileN>{}), make_coord(by, bx));