hgemm_mma_m16n8k16_mma2x4_warp4x4_kernel

鱿鱼圈 Lv4

hgemm_mma_m16n8k16_mma2x4_warp4x4_kernel 详解

本文档详细解析 HGEMM MMA Kernel 的实现原理,包括数据布局、线程组织、数据加载、MMA 计算和结果存储的完整流程。

代码链接:HGEMM/kernels/hgemm/mma/hgemm_mma.cu at eee72be829545bd6bd115a4b252b5068c9f61597 · xlite-dev/HGEMM

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
#define WARP_SIZE 32
#define DEVICE_INLINE __device__ inline
#define HOST_DEVICE_INLINE __device__ __host__ inline
#define INT4(value) (reinterpret_cast<int4*>(&(value))[0])
#define FLOAT4(value) (reinterpret_cast<float4*>(&(value))[0])
#define HALF2(value) (reinterpret_cast<half2*>(&(value))[0])
#define BFLOAT2(value) (reinterpret_cast<__nv_bfloat162*>(&(value))[0])
#define LDST32BITS(value) (reinterpret_cast<half2*>(&(value))[0])
#define LDST64BITS(value) (reinterpret_cast<float2*>(&(value))[0])
#define LDST128BITS(value) (reinterpret_cast<float4*>(&(value))[0])
#define CP_ASYNC_COMMIT_GROUP() asm volatile("cp.async.commit_group;\n" ::)
#define CP_ASYNC_WAIT_ALL() asm volatile("cp.async.wait_all;\n" ::)
#define CP_ASYNC_WAIT_GROUP(n) asm volatile("cp.async.wait_group %0;\n" ::"n"(n))
// ca(cache all, L1 + L2): support 4, 8, 16 bytes, cg(cache global, L2): only support 16 bytes.
#define CP_ASYNC_CA(dst, src, bytes) asm volatile("cp.async.ca.shared.global.L2::128B [%0], [%1], %2;\n" ::"r"(dst), "l"(src), "n"(bytes))
#define CP_ASYNC_CG(dst, src, bytes) asm volatile("cp.async.cg.shared.global.L2::128B [%0], [%1], %2;\n" ::"r"(dst), "l"(src), "n"(bytes))
#define LDMATRIX_X1(R, addr) asm volatile("ldmatrix.sync.aligned.x1.m8n8.shared.b16 {%0}, [%1];\n" : "=r"(R) : "r"(addr))
#define LDMATRIX_X2(R0, R1, addr) asm volatile("ldmatrix.sync.aligned.x2.m8n8.shared.b16 {%0, %1}, [%2];\n" : "=r"(R0), "=r"(R1) : "r"(addr))
#define LDMATRIX_X4(R0, R1, R2, R3, addr) asm volatile("ldmatrix.sync.aligned.x4.m8n8.shared.b16 {%0, %1, %2, %3}, [%4];\n" : "=r"(R0), "=r"(R1), "=r"(R2), "=r"(R3) : "r"(addr))
#define LDMATRIX_X1_T(R, addr) asm volatile("ldmatrix.sync.aligned.x1.trans.m8n8.shared.b16 {%0}, [%1];\n" : "=r"(R) : "r"(addr))
#define LDMATRIX_X2_T(R0, R1, addr) asm volatile("ldmatrix.sync.aligned.x2.trans.m8n8.shared.b16 {%0, %1}, [%2];\n" : "=r"(R0), "=r"(R1) : "r"(addr))
#define LDMATRIX_X4_T(R0, R1, R2, R3, addr) asm volatile("ldmatrix.sync.aligned.x4.trans.m8n8.shared.b16 {%0, %1, %2, %3}, [%4];\n" : "=r"(R0), "=r"(R1), "=r"(R2), "=r"(R3) : "r"(addr))
#define HMMA16816(RD0, RD1, RA0, RA1, RA2, RA3, RB0, RB1, RC0, RC1) asm volatile("mma.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16 {%0, %1}, {%2, %3, %4, %5}, {%6, %7}, {%8, %9};\n" : "=r"(RD0), "=r"(RD1) : "r"(RA0), "r"(RA1), "r"(RA2), "r"(RA3), "r"(RB0), "r"(RB1), "r"(RC0), "r"(RC1))

HOST_DEVICE_INLINE
int div_ceil(int a, int b) { return (a % b != 0) ? (a / b + 1) : (a / b); }

// 128x128, mma2x4, warp4x4(64,32,16)
template<const int MMA_M=16,
const int MMA_N=8,
const int MMA_K=16,
const int MMA_TILE_M=2,
const int MMA_TILE_N=4,
const int WARP_TILE_M=4,
const int WARP_TILE_N=4,
const int A_PAD=0,
const int B_PAD=0>
__global__ void __launch_bounds__(256)
hgemm_mma_m16n8k16_mma2x4_warp4x4_kernel(
half* A, half* B, half* C, int M, int N, int K) {
const int bx = blockIdx.x;
const int by = blockIdx.y;
const int NUM_K_TILES = div_ceil(K, MMA_K);
constexpr int BM = MMA_M * MMA_TILE_M * WARP_TILE_M; // 16*2*4=128
constexpr int BN = MMA_N * MMA_TILE_N * WARP_TILE_N; // 8*4*4=128
constexpr int BK = MMA_K; // 16

__shared__ half s_a[BM][BK+A_PAD]; // 128*16*2=4KB
__shared__ half s_b[BK][BN+B_PAD]; // 16*128*2=4KB, 16*(128+16)*2=4.5KB

const int tid = threadIdx.y * blockDim.x + threadIdx.x; // within block
const int warp_id = tid / WARP_SIZE; // 0~7 warp_id within block
const int lane_id = tid % WARP_SIZE; // 0~31
const int warp_m = warp_id % 2; // 0,1
const int warp_n = warp_id / 2; // 0,1,2,3

// 先计算shared memory中的索引
// tid和需要加载的smem s_a[BM][BK] 之间的索引关系 BM=128 BK=16 按行读取 A行主序
// 对于s_a每行16个数据,每个线程读取8个,需要2个线程;总共128行,需要128x2刚好256线程
int load_smem_a_m = tid / 2; // row 0~127
int load_smem_a_k = (tid % 2 == 0) ? 0 : 8; // col 0,8
// tid和需要加载的smem s_b[BK][BN] 之间的索引关系 BK=16 BN=128 按行读取 B行主序
// 对于s_b每行128个数据,每个线程读8个数据,需要16个线程;总共16行,需要16x16=256个线程
int load_smem_b_k = tid / 16; // row 0~15
int load_smem_b_n = (tid % 16) * 8; // col 0,8,...,120
// 再计算全局内存中的索引
// 要加载到s_a中的元素对应到A全局内存中的行数 每个block负责出C中大小为BM*BN的块
int load_gmem_a_m = by * BM + load_smem_a_m; // global row of a and c
int load_gmem_b_n = bx * BN + load_smem_b_n; // global col of b and c
if (load_gmem_a_m >= M || load_gmem_b_n >= N) return;

uint32_t RC[WARP_TILE_M][WARP_TILE_N][2];
#pragma unroll
for (int i = 0; i < WARP_TILE_M; ++i) {
#pragma unroll
for (int j = 0; j < WARP_TILE_N; ++j) {
RC[i][j][0] = 0;
RC[i][j][1] = 0;
}
}

#pragma unroll
for (int k = 0; k < NUM_K_TILES; ++k) {
// gmem -> smem
int load_gmem_a_k = k * BK + load_smem_a_k; // global col of a
int load_gmem_a_addr = load_gmem_a_m * K + load_gmem_a_k;
int load_gmem_b_k = k * BK + load_smem_b_k; // global row of b
int load_gmem_b_addr = load_gmem_b_k * N + load_gmem_b_n;
LDST128BITS(s_b[load_smem_b_k][load_smem_b_n]) = (
LDST128BITS(B[load_gmem_b_addr]));
LDST128BITS(s_a[load_smem_a_m][load_smem_a_k]) = (
LDST128BITS(A[load_gmem_a_addr]));
__syncthreads();

// ldmatrix for s_a, ldmatrix.trans for s_b.
uint32_t RA[WARP_TILE_M][4];
uint32_t RB[WARP_TILE_N][2];

// smem -> reg
#pragma unroll
for (int i = 0; i < WARP_TILE_M; ++i) {
int warp_smem_a_m = warp_m * (MMA_M * WARP_TILE_M) + i * MMA_M;
int lane_smem_a_m = warp_smem_a_m + lane_id % 16; // 0~15
int lane_smem_a_k = (lane_id / 16) * 8; // 0,8
uint32_t lane_smem_a_ptr = __cvta_generic_to_shared(
&s_a[lane_smem_a_m][lane_smem_a_k]);
LDMATRIX_X4(RA[i][0], RA[i][1], RA[i][2], RA[i][3], lane_smem_a_ptr);
}

#pragma unroll
for (int j = 0; j < WARP_TILE_N; ++j) {
int warp_smem_b_n = warp_n * (MMA_N * WARP_TILE_N) + j * MMA_N;
int lane_smem_b_k = lane_id % 16; // 0~15
int lane_smem_b_n = warp_smem_b_n; // 0, MMA_N=8
uint32_t lane_smem_b_ptr = __cvta_generic_to_shared(
&s_b[lane_smem_b_k][lane_smem_b_n]);
LDMATRIX_X2_T(RB[j][0], RB[j][1], lane_smem_b_ptr);
}

// MMA compute
#pragma unroll
for (int i = 0; i < WARP_TILE_M; ++i) {
#pragma unroll
for (int j = 0; j < WARP_TILE_N; ++j) {
HMMA16816(RC[i][j][0], RC[i][j][1],
RA[i][0], RA[i][1], RA[i][2], RA[i][3],
RB[j][0], RB[j][1],
RC[i][j][0], RC[i][j][1]);
}
}
__syncthreads();
}

// reg -> gmem, MMA_MxMMA_N=16x8
#pragma unroll
for (int i = 0; i < WARP_TILE_M; ++i) {
#pragma unroll
for (int j = 0; j < WARP_TILE_N; ++j) {
int store_warp_smem_c_m = warp_m * (MMA_M * WARP_TILE_M) + i * MMA_M;
int store_warp_smem_c_n = warp_n * (MMA_N * WARP_TILE_N) + j * MMA_N;
// mapping lane smem index -> global index.
// [16][8], https://docs.nvidia.com/cuda/parallel-thread-execution/index.html
// #matrix-fragments-for-mma-m16n8k16-with-floating-point-type
// [0~7][0~3 u32 -> 0~7 f16], [8~15][0~3 u32 -> 0~7 f16]
int store_lane_gmem_c_m = by * BM + store_warp_smem_c_m + lane_id / 4;
int store_lane_gmem_c_n = bx * BN + store_warp_smem_c_n + (lane_id % 4) * 2;
int store_gmem_c_addr_0 = store_lane_gmem_c_m * N + store_lane_gmem_c_n;
int store_gmem_c_addr_1 = (store_lane_gmem_c_m + 8) * N + store_lane_gmem_c_n;
// TODO: how to use LDST128BITS here ? reverse the loop order ?
LDST32BITS(C[store_gmem_c_addr_0]) = LDST32BITS(RC[i][j][0]);
LDST32BITS(C[store_gmem_c_addr_1]) = LDST32BITS(RC[i][j][1]);
}
}
}

1. Kernel 概述

1.1 功能描述

该 Kernel 实现半精度 (FP16) 矩阵乘法:C = A × B

  • 输入: A[M×K], B[K×N](均为行主序)
  • 输出: C[M×N](行主序)
  • 使用硬件: NVIDIA Tensor Core (MMA 指令)

1.2 核心特点

特性 说明
Block Tile 128×128 每个 Block 计算 C 的 128×128 子块
线程数/Block 256 8 个 Warp
Warp Tile 64×32 每个 Warp 计算 64×32 区域
MMA Tile 16×8 单次 MMA 指令计算 16×8
数据复用 2×4 Warp 级别的数据复用

2. 模板参数与 Block Tile 尺寸

2.1 模板参数定义

1
2
3
4
5
6
7
8
9
template<const int MMA_M=16,        // MMA 指令的 M 维度
const int MMA_N=8, // MMA 指令的 N 维度
const int MMA_K=16, // MMA 指令的 K 维度
const int MMA_TILE_M=2, // Warp 排列的 M 方向数量
const int MMA_TILE_N=4, // Warp 排列的 N 方向数量
const int WARP_TILE_M=4, // 每 Warp 在 M 方向执行的 MMA 次数
const int WARP_TILE_N=4, // 每 Warp 在 N 方向执行的 MMA 次数
const int A_PAD=0, // A 矩阵 Shared Memory Padding
const int B_PAD=0> // B 矩阵 Shared Memory Padding

2.2 Block Tile 尺寸计算

1
2
3
BM = MMA_M × MMA_TILE_M × WARP_TILE_M = 16 × 2 × 4 = 128
BN = MMA_N × MMA_TILE_N × WARP_TILE_N = 8 × 4 × 4 = 128
BK = MMA_K = 16

2.3 Shared Memory 分配

1
2
3
__shared__ half s_a[BM][BK+A_PAD];  // 128 × (16+8) × 2 = 6 KB (带 Padding)
__shared__ half s_b[BK][BN+B_PAD]; // 16 × (128+8) × 2 = 4.25 KB (带 Padding)
// 总计约 10.25 KB Shared Memory

2.4 参数层次关系图

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
┌─────────────────────────────────────────────────────────────────────────┐
│ Block Tile (128 × 128) │
│ ┌────────────────────────────────┬────────────────────────────────┐ │
│ │ MMA_TILE_M=2 warps │ │ │
│ │ ┌──────────────────────────┐ │ │ │
│ │ │ Warp Tile (64 × 32) │ │ MMA_TILE_N=4 warps │ │
│ │ │ ┌─────┬─────┬─────┬─────│ │ │ │
│ │ │ │16×8 │16×8 │16×8 │16×8 │ │ │ │
│ │ │ │ MMA │ MMA │ MMA │ MMA │ │ WARP_TILE_N=4 │ │
│ │ │ ├─────┼─────┼─────┼─────│ │ │ │
│ │ │ │ │ │ │ │ │ │ │
│ │ │ │ ... │ ... │ ... │ ... │ │ │ │
│ │ │ │ │ │ │ │ │ │ │
│ │ │ └─────┴─────┴─────┴─────│ │ WARP_TILE_M=4 │ │
│ │ └──────────────────────────┘ │ │ │
│ └────────────────────────────────┴────────────────────────────────┘ │
└─────────────────────────────────────────────────────────────────────────┘

3. 线程组织与 Warp 布局

3.1 线程索引计算

1
2
3
4
5
const int tid = threadIdx.y * blockDim.x + threadIdx.x;  // 0~255
const int warp_id = tid / WARP_SIZE; // 0~7,Block 内的 Warp ID
const int lane_id = tid % WARP_SIZE; // 0~31,Warp 内的线程 ID
const int warp_m = warp_id % 2; // 0 或 1,Warp 在 M 方向的位置
const int warp_n = warp_id / 2; // 0~3,Warp 在 N 方向的位置

3.2 Warp 布局 (2×4 排列)

1
2
3
4
5
6
7
8
9
              warp_n=0    warp_n=1    warp_n=2    warp_n=3
N: 0~31 N: 32~63 N: 64~95 N: 96~127
┌─────────┬─────────┬─────────┬─────────┐
warp_m=0 │ warp 0 │ warp 2 │ warp 4 │ warp 6 │ M: 0~63
M: 0~63 │ (0,0) │ (0,1) │ (0,2) │ (0,3) │
├─────────┼─────────┼─────────┼─────────┤
warp_m=1 │ warp 1 │ warp 3 │ warp 5 │ warp 7 │ M: 64~127
M: 64~127 │ (1,0) │ (1,1) │ (1,2) │ (1,3) │
└─────────┴─────────┴─────────┴─────────┘

3.3 Warp ID 与坐标映射表

warp_id warp_m warp_n 负责 C 的区域
0 0 0 C[0:64, 0:32]
1 1 0 C[64:128, 0:32]
2 0 1 C[0:64, 32:64]
3 1 1 C[64:128, 32:64]
4 0 2 C[0:64, 64:96]
5 1 2 C[64:128, 64:96]
6 0 3 C[0:64, 96:128]
7 1 3 C[64:128, 96:128]

4. Global Memory → Shared Memory 数据加载

4.1 加载 A 矩阵到 s_a[128][16]

任务分配: 128 行 × 16 列,每线程加载 8 个 half (128 bits)

1
2
int load_smem_a_m = tid / 2;                    // row 0~127
int load_smem_a_k = (tid % 2 == 0) ? 0 : 8; // col 0 或 8

分配示意:

1
2
3
4
5
6
7
8
9
10
11
s_a[128][16]:
K=0~7 K=8~15
┌────────────┬────────────┐
tid=0 │ row 0 │ row 0 │ ← tid 0 加载 row 0, col 0~7
tid=1 │ │ │ ← tid 1 加载 row 0, col 8~15
tid=2 │ row 1 │ row 1 │
tid=3 │ │ │
... │ ... │ ... │
tid=254 │ │ │
tid=255 │ row 127 │ row 127 │
└────────────┴────────────┘

4.2 加载 B 矩阵到 s_b[16][128]

任务分配: 16 行 × 128 列,每线程加载 8 个 half (128 bits)

1
2
int load_smem_b_k = tid / 16;         // row 0~15
int load_smem_b_n = (tid % 16) * 8; // col 0,8,16,...,120

分配示意:

1
2
3
4
5
6
7
8
s_b[16][128]:
N=0~7 8~15 16~23 ... 120~127
┌──────┬──────┬──────┬──────┬──────┐
tid 0~15│ row 0│ │ │ │ │
tid16~31│ row 1│ │ │ │ │
... │ ... │ │ │ │ │
tid240~ │row 15│ │ │ │ │
└──────┴──────┴──────┴──────┴──────┘

4.3 数据加载代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
#pragma unroll
for (int k = 0; k < NUM_K_TILES; ++k) {
// 计算全局地址
int load_gmem_a_k = k * BK + load_smem_a_k;
int load_gmem_a_addr = load_gmem_a_m * K + load_gmem_a_k;
int load_gmem_b_k = k * BK + load_smem_b_k;
int load_gmem_b_addr = load_gmem_b_k * N + load_gmem_b_n;

// 128-bit 向量化加载
LDST128BITS(s_b[load_smem_b_k][load_smem_b_n]) = LDST128BITS(B[load_gmem_b_addr]);
LDST128BITS(s_a[load_smem_a_m][load_smem_a_k]) = LDST128BITS(A[load_gmem_a_addr]);
__syncthreads();

// ... MMA 计算 ...

__syncthreads();
}

图示:

1
2
3
4
5
6
7
s_a[128][16] 加载:                s_b[16][128] 加载:
┌────┬────┐ ┌──────────────────────┐
│T0 │T1 │ row 0 │T0 T1 T2 ... T15 │ row 0
│T2 │T3 │ row 1 │T16 T17 ... T31 │ row 1
│... │... │ │... │
│T254│T255│ row 127 │T240 T241 ... T255 │ row 15
└────┴────┘ └──────────────────────┘

5. Shared Memory → Register 数据加载 (ldmatrix)

5.1 ldmatrix 指令简介

ldmatrix 是专门为 Tensor Core 设计的 warp 级指令,可以高效地从 Shared Memory 加载数据到寄存器,并自动完成 MMA 所需的数据重排。

指令 功能
ldmatrix.x4 加载 4 个 8×8 矩阵(16×16 共 256 个元素)
ldmatrix.x2.trans 加载 2 个 8×8 矩阵并转置(16×8 共 128 个元素)

5.2 加载 A 矩阵 (s_a → RA)

1
2
3
4
5
6
7
8
#pragma unroll
for (int i = 0; i < WARP_TILE_M; ++i) { // i = 0,1,2,3
int warp_smem_a_m = warp_m * (MMA_M * WARP_TILE_M) + i * MMA_M;
int lane_smem_a_m = warp_smem_a_m + lane_id % 16; // 0~15
int lane_smem_a_k = (lane_id / 16) * 8; // 0 或 8
uint32_t lane_smem_a_ptr = __cvta_generic_to_shared(&s_a[lane_smem_a_m][lane_smem_a_k]);
LDMATRIX_X4(RA[i][0], RA[i][1], RA[i][2], RA[i][3], lane_smem_a_ptr);
}

索引计算解析:

变量 公式 含义
warp_smem_a_m warp_m × 64 + i × 16 当前 MMA 块的起始行
lane_smem_a_m warp_smem_a_m + lane_id % 16 每个 lane 提供的行地址
lane_smem_a_k (lane_id / 16) × 8 lane 0~15→col 0, lane 16~31→col 8

warp_m 对 s_a 的划分:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
s_a[128][16]:
┌──────────────────────┐
│ warp_m=0 负责区域 │ 行 0~63
│ i=0: 行 0~15 │ ← RA[0]
│ i=1: 行 16~31 │ ← RA[1]
│ i=2: 行 32~47 │ ← RA[2]
│ i=3: 行 48~63 │ ← RA[3]
├──────────────────────┤
│ warp_m=1 负责区域 │ 行 64~127
│ i=0: 行 64~79 │ ← RA[0]
│ i=1: 行 80~95 │ ← RA[1]
│ i=2: 行 96~111 │ ← RA[2]
│ i=3: 行 112~127 │ ← RA[3]
└──────────────────────┘

ldmatrix.x4 地址分布 (以 warp_m=0, i=0 为例):

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
s_a[16][16] 中一个 MMA 块:

K=0~7 K=8~15
┌────────┬────────┐
lane 0 → │ row 0 │ │
lane 1 → │ row 1 │ │
... │ ... │ │
lane 15 → │ row 15 │ │
├────────┼────────┤
lane 16 → │ │ row 0 │
lane 17 → │ │ row 1 │
... │ │ ... │
lane 31 → │ │ row 15 │
└────────┴────────┘

输出: RA[i][0], RA[i][1] ← K=0~7 的数据
RA[i][2], RA[i][3] ← K=8~15 的数据

5.3 加载 B 矩阵 (s_b → RB)

1
2
3
4
5
6
7
8
#pragma unroll
for (int j = 0; j < WARP_TILE_N; ++j) { // j = 0,1,2,3
int warp_smem_b_n = warp_n * (MMA_N * WARP_TILE_N) + j * MMA_N;
int lane_smem_b_k = lane_id % 16; // 0~15
int lane_smem_b_n = warp_smem_b_n; // 固定列起始
uint32_t lane_smem_b_ptr = __cvta_generic_to_shared(&s_b[lane_smem_b_k][lane_smem_b_n]);
LDMATRIX_X2_T(RB[j][0], RB[j][1], lane_smem_b_ptr);
}

索引计算解析:

变量 公式 含义
warp_smem_b_n warp_n × 32 + j × 8 当前 MMA 块的起始列
lane_smem_b_k lane_id % 16 每个 lane 提供的行地址 (0~15)
lane_smem_b_n warp_smem_b_n 所有 lane 使用相同列起始

warp_n 对 s_b 的划分:

1
2
3
4
5
6
7
8
9
s_b[16][128]:
warp_n=0 warp_n=1 warp_n=2 warp_n=3
N=0~31 N=32~63 N=64~95 N=96~127
┌────┬────┬────┬────┬────┬────┬────┬────┐
│j=0 │j=1 │j=2 │j=3 │j=0 │j=1 │... │j=3 │
K=0~15│ 8 │ 8 │ 8 │ 8 │ 8 │ 8 │ │ 8 │
└────┴────┴────┴────┴────┴────┴────┴────┘
↑ ↑ ↑ ↑
RB[0]RB[1]RB[2]RB[3]

为什么使用 .trans (转置)?

  • B 矩阵在 Shared Memory 中是行主序 s_b[K][N]
  • MMA 指令要求 B 以列主序方式组织
  • ldmatrix.trans 在加载时自动转置,避免显式转置开销

5.4 Warp 数据加载完整总表

warp_id warp_m warp_n 加载 s_a 行范围 加载 s_b 列范围
0 0 0 0~63 0~31
1 1 0 64~127 0~31
2 0 1 0~63 32~63
3 1 1 64~127 32~63
4 0 2 0~63 64~95
5 1 2 64~127 64~95
6 0 3 0~63 96~127
7 1 3 64~127 96~127

数据复用示意:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
                        s_b[16][128]
┌────┬────┬────┬────┐
│0~31│32~63│64~95│96~127│
└──┬─┴──┬─┴──┬─┴──┬─┘
│ │ │ │
s_a[128][16] ▼ ▼ ▼ ▼
┌────────┐ ┌────┬────┬────┬────┐
0~63 │ warp_m │──│ W0 │ W2 │ W4 │ W6 │ ← warp 0,2,4,6 共享 A 的行 0~63
│ =0 │ ├────┼────┼────┼────┤
├────────┤ │ │ │ │ │
64~127 │ warp_m │──│ W1 │ W3 │ W5 │ W7 │ ← warp 1,3,5,7 共享 A 的行 64~127
│ =1 │ └────┴────┴────┴────┘
└────────┘
↑ ↑ ↑ ↑
warp 0,1 共享 B 的列 0~31
warp 2,3 共享 B 的列 32~63
...

6. MMA 计算

6.1 MMA 指令说明

使用 mma.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16 指令:

1
D[16×8] = A[16×16] × B[16×8] + C[16×8]

每次 MMA 指令执行 16 × 8 × 16 × 2 = 4096 次 FP16 乘加运算。

6.2 MMA 计算代码

1
2
3
4
5
6
7
8
9
10
#pragma unroll
for (int i = 0; i < WARP_TILE_M; ++i) { // i = 0,1,2,3
#pragma unroll
for (int j = 0; j < WARP_TILE_N; ++j) { // j = 0,1,2,3
HMMA16816(RC[i][j][0], RC[i][j][1],
RA[i][0], RA[i][1], RA[i][2], RA[i][3],
RB[j][0], RB[j][1],
RC[i][j][0], RC[i][j][1]);
}
}

6.3 每 Warp 的 MMA 执行矩阵

每个 Warp 执行 WARP_TILE_M × WARP_TILE_N = 4 × 4 = 16 次 MMA:

1
2
3
4
5
6
7
8
9
10
11
12
13
Warp Tile (64×32) 中的 16 次 MMA:

j=0 j=1 j=2 j=3
┌────────┬────────┬────────┬────────┐
i=0 │RC[0][0]│RC[0][1]│RC[0][2]│RC[0][3]│ 16×8 × 4 = 16×32
│ 16×8 │ 16×8 │ 16×8 │ 16×8 │
├────────┼────────┼────────┼────────┤
i=1 │RC[1][0]│RC[1][1]│RC[1][2]│RC[1][3]│
├────────┼────────┼────────┼────────┤
i=2 │RC[2][0]│RC[2][1]│RC[2][2]│RC[2][3]│
├────────┼────────┼────────┼────────┤
i=3 │RC[3][0]│RC[3][1]│RC[3][2]│RC[3][3]│ 64×32 总计
└────────┴────────┴────────┴────────┘

6.4 计算量统计

层级 计算量
单次 MMA 4,096 FLOPs
单 Warp / K-tile 16 × 4,096 = 65,536 FLOPs
单 Block / K-tile 8 × 65,536 = 524,288 FLOPs
单 Block 总计 524,288 × NUM_K_TILES FLOPs

7. Register → Global Memory 结果存储

7.1 存储代码详解

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
#pragma unroll
for (int i = 0; i < WARP_TILE_M; ++i) {
#pragma unroll
for (int j = 0; j < WARP_TILE_N; ++j) {
// 计算当前 MMA 块在 Block Tile 中的位置
int store_warp_smem_c_m = warp_m * (MMA_M * WARP_TILE_M) + i * MMA_M;
int store_warp_smem_c_n = warp_n * (MMA_N * WARP_TILE_N) + j * MMA_N;

// 计算每个 lane 在全局内存中的地址
int store_lane_gmem_c_m = by * BM + store_warp_smem_c_m + lane_id / 4;
int store_lane_gmem_c_n = bx * BN + store_warp_smem_c_n + (lane_id % 4) * 2;

// 计算两个 8 行块的地址
int store_gmem_c_addr_0 = store_lane_gmem_c_m * N + store_lane_gmem_c_n;
int store_gmem_c_addr_1 = (store_lane_gmem_c_m + 8) * N + store_lane_gmem_c_n;

// 写入全局内存
LDST32BITS(C[store_gmem_c_addr_0]) = LDST32BITS(RC[i][j][0]);
LDST32BITS(C[store_gmem_c_addr_1]) = LDST32BITS(RC[i][j][1]);
}
}

7.2 索引计算详解

7.2.1 Warp 级偏移

1
2
3
4
5
store_warp_smem_c_m = warp_m * (MMA_M * WARP_TILE_M) + i * MMA_M
= warp_m * 64 + i * 16

store_warp_smem_c_n = warp_n * (MMA_N * WARP_TILE_N) + j * MMA_N
= warp_n * 32 + j * 8
warp_m warp_n i j store_warp_smem_c_m store_warp_smem_c_n
0 0 0 0 0 0
0 0 0 1 0 8
0 0 1 0 16 0
1 2 3 3 112 88

7.2.2 Lane 级偏移 (MMA Fragment 布局)

根据 NVIDIA PTX 文档,mma.m16n8k16 的输出 Fragment 布局:

1
2
store_lane_gmem_c_m = ... + lane_id / 4    // 行偏移: 0~7
store_lane_gmem_c_n = ... + (lane_id % 4) * 2 // 列偏移: 0,2,4,6

Fragment 布局公式 (官方 PTX 文档):

  • groupID = lane_id >> 2 (即 lane_id / 4)
  • threadID_in_group = lane_id % 4
  • 行 = groupID (对于 RC[0]) 或 groupID + 8 (对于 RC[1])
  • 列 = threadID_in_group * 2 + (i & 0x1),其中 i 是 fragment 内的元素索引

7.3 MMA 输出 Fragment 的 Lane 分布

单个 MMA 输出 RC[i][j][2] 对应 16×8 的 C 子块:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
C[16][8] 的 lane 分布:

RC[i][j][0] 覆盖行 0~7:
col: 0 1 2 3 4 5 6 7
┌──┬──┬──┬──┬──┬──┬──┬──┐
row 0 │L0│L0│L1│L1│L2│L2│L3│L3│
row 1 │L4│L4│L5│L5│L6│L6│L7│L7│
row 2 │L8│L8│L9│L9│L10│L10│L11│L11│
row 3 │L12│L12│L13│L13│L14│L14│L15│L15│
row 4 │L16│L16│L17│L17│L18│L18│L19│L19│
row 5 │L20│L20│L21│L21│L22│L22│L23│L23│
row 6 │L24│L24│L25│L25│L26│L26│L27│L27│
row 7 │L28│L28│L29│L29│L30│L30│L31│L31│
└──┴──┴──┴──┴──┴──┴──┴──┘

RC[i][j][1] 覆盖行 8~15 (相同 lane 分布,行 +8)

每个 Lane 的存储:

  • RC[i][j][0]: 存储 2 个连续的 half 到 行 lane_id/4,列 (lane_id%4)*2(lane_id%4)*2+1
  • RC[i][j][1]: 存储 2 个连续的 half 到 行 lane_id/4 + 8,相同列

7.4 存储过程示意图

Warp 0 (warp_m=0, warp_n=0) 的 RC[0][0] 为例:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
Block 在全局内存 C 中的位置: C[by*128 : by*128+128, bx*128 : bx*128+128]

Warp 0 的 RC[0][0] (i=0, j=0) 对应:
store_warp_smem_c_m = 0 × 64 + 0 × 16 = 0
store_warp_smem_c_n = 0 × 32 + 0 × 8 = 0

各 Lane 写入的全局地址:
┌─────────────────────────────────────────────────┐
│ Lane 0: C[by*128+0, bx*128+0:2] ← RC[0][0][0]│
│ C[by*128+8, bx*128+0:2] ← RC[0][0][1]│
│ Lane 1: C[by*128+0, bx*128+2:4] ← RC[0][0][0]│
│ C[by*128+8, bx*128+2:4] ← RC[0][0][1]│
│ ... │
│ Lane 4: C[by*128+1, bx*128+0:2] ← RC[0][0][0]│
│ C[by*128+9, bx*128+0:2] ← RC[0][0][1]│
│ ... │
│ Lane 31: C[by*128+7, bx*128+6:8] ← RC[0][0][0]│
│ C[by*128+15, bx*128+6:8] ← RC[0][0][1]│
└─────────────────────────────────────────────────┘

7.5 为什么使用 LDST32BITS 而不是 LDST128BITS?

代码中有注释 TODO: how to use LDST128BITS here ? reverse the loop order ?

原因:

  • 每个 lane 的 RC[i][j][0] 只包含 2 个 half (32 bits)
  • 要使用 128-bit 存储,需要让 4 个连续 lane 的数据合并
  • 当前循环顺序下,相邻 lane 写入的列不连续 (列间隔为 2)
  • 需要调整循环顺序或使用 warp shuffle 来合并数据

可能的优化方向:

  1. 调换 i、j 循环顺序
  2. 使用 __shfl_sync 在 lane 间交换数据
  3. 先写入 Shared Memory,重排后再写入 Global Memory

7.6 完整存储流程图

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
                    RC[4][4][2] (每 Warp)

┌────────────────┼────────────────┐
│ │ │
▼ ▼ ▼
RC[0][0~3] RC[1][0~3] ... RC[3][0~3]
i=0 的 4 个 i=1 的 4 个 i=3 的 4 个
│ │ │
▼ ▼ ▼
┌─────────────────────────────────────────────────┐
│ Global Memory C[M][N] │
│ ┌─────────────────────────────────────────┐ │
│ │ Block (by, bx) 的 128×128 区域 │ │
│ │ ┌───────────────────────────────────┐ │ │
│ │ │ Warp 0 的 64×32 区域 │ │ │
│ │ │ ┌───────┬───────┬───────┬───────│ │ │
│ │ │ │RC[0][0]RC[0][1]RC[0][2]RC[0][3]│ │ │
│ │ │ │ 16×8 │ 16×8 │ 16×8 │ 16×8 │ │ │
│ │ │ ├───────┼───────┼───────┼───────│ │ │
│ │ │ │RC[1][0] ... │ │ │
│ │ │ │ ... │ │ │
│ │ │ │RC[3][0] ... RC[3][3] │ │ │
│ │ │ └───────┴───────┴───────┴───────│ │ │
│ │ └───────────────────────────────────┘ │ │
│ └─────────────────────────────────────────┘ │
└─────────────────────────────────────────────────┘

8. 性能优化要点

8.1 本 Kernel 已实现的优化

优化技术 实现方式 效果
Warp Tiling 2×4 warp 布局 数据复用,减少访存
ldmatrix 专用指令加载 高效数据重排
向量化访存 LDST128BITS 最大化带宽利用
Shared Memory Padding A_PAD=8, B_PAD=8 消除 Bank Conflict
循环展开 #pragma unroll 减少循环开销

8.2 数据复用分析

1
2
3
4
数据复用比 = MMA_TILE_M × MMA_TILE_N = 2 × 4 = 8

每份 A 数据被 4 个 warp (warp_n=0~3) 复用
每份 B 数据被 2 个 warp (warp_m=0~1) 复用

8.3 访存量分析

单个 K-tile 迭代:

  • 加载 A: 128 × 16 × 2 = 4 KB
  • 加载 B: 16 × 128 × 2 = 4 KB
  • 计算量: 128 × 128 × 16 × 2 = 524,288 FLOPs

计算访存比: 524,288 / 8,192 = 64 FLOPs/Byte

8.4 可进一步优化的方向

  1. Double Buffering: 使用双缓冲隐藏访存延迟
  2. Async Copy (cp.async): 异步数据拷贝
  3. Multi-stage Pipeline: 多阶段流水线
  4. Swizzle: 进一步优化 Bank Conflict
  5. 128-bit 输出存储: 合并输出写入

附录: 关键宏定义

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
#define LDST32BITS(value) (reinterpret_cast<half2*>(&(value))[0])
#define LDST128BITS(value) (reinterpret_cast<float4*>(&(value))[0])

#define LDMATRIX_X4(R0, R1, R2, R3, addr) \
asm volatile("ldmatrix.sync.aligned.x4.m8n8.shared.b16 {%0, %1, %2, %3}, [%4];\n" \
: "=r"(R0), "=r"(R1), "=r"(R2), "=r"(R3) : "r"(addr))

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

#define HMMA16816(RD0, RD1, RA0, RA1, RA2, RA3, RB0, RB1, RC0, RC1) \
asm volatile("mma.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16 \
{%0, %1}, {%2, %3, %4, %5}, {%6, %7}, {%8, %9};\n" \
: "=r"(RD0), "=r"(RD1) \
: "r"(RA0), "r"(RA1), "r"(RA2), "r"(RA3), "r"(RB0), "r"(RB1), "r"(RC0), "r"(RC1))

参考资料

  1. NVIDIA PTX ISA - Matrix Fragments for mma.m16n8k16
  2. NVIDIA PTX ISA - ldmatrix
  3. CUTLASS Documentation
  • 标题: hgemm_mma_m16n8k16_mma2x4_warp4x4_kernel
  • 作者: 鱿鱼圈
  • 创建于 : 2026-03-03 23:50:00
  • 更新于 : 2026-06-05 23:01:32
  • 链接: https://yuyanqi.com/2026/03/03/hgemm_mma_m16n8k16_mma2x4_warp4x4_kernel/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论