hgemm_mma_m16n8k16_naive_kernel

鱿鱼圈 Lv4

hgemm_mma_m16n8k16_naive_kernel详解

仓库链接:xlite-dev/HGEMM at eee72be829545bd6bd115a4b252b5068c9f61597

代码链接: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
#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); }


// only 1 warp per block(32 threads), m16n8k16. A, B, C: all row_major.
template<const int MMA_M=16, const int MMA_N=8, const int MMA_K=16>
__global__ void hgemm_mma_m16n8k16_naive_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; // 16
constexpr int BN = MMA_N; // 8
constexpr int BK = MMA_K; // 16

__shared__ half s_a[MMA_M][MMA_K]; // 16x16
__shared__ half s_b[MMA_K][MMA_N]; // 16x8
__shared__ half s_c[MMA_M][MMA_N]; // 16x8

const int tid = threadIdx.y * blockDim.x + threadIdx.x; // within block
const int lane_id = tid % WARP_SIZE; // 0~31

// s_a[16][16], 每行16,每线程load 8,需要2线程,共16行,需2x16=32线程
const int load_smem_a_m = tid / 2; // row 0~15
const int load_smem_a_k = (tid % 2) * 8; // col 0,8
// s_b[16][8], 每行8,每线程load 8,需要1线程,共16行,需16线程,只需一半线程加载
const int load_smem_b_k = tid; // row 0~31, but only use 0~15
const int load_smem_b_n = 0; // col 0
const int load_gmem_a_m = by * BM + load_smem_a_m; // global m
const int load_gmem_b_n = bx * BN + load_smem_b_n; // global n
if (load_gmem_a_m >= M && load_gmem_b_n >= N) return;

uint32_t RC[2] = {0, 0};

#pragma unroll
for (int k = 0; k < NUM_K_TILES; ++k) {
// gmem_a -> smem_a
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;
LDST128BITS(s_a[load_smem_a_m][load_smem_a_k]) = (
LDST128BITS(A[load_gmem_a_addr]));

// gmem_b -> smem_b
if (lane_id < MMA_K) {
int load_gmem_b_k = k * MMA_K + 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]));
}
__syncthreads();

uint32_t RA[4];
uint32_t RB[2];

// ldmatrix for s_a, ldmatrix.trans for s_b.
// s_a: (0,1)*8 -> 0,8 -> [(0~15),(0,8)]
uint32_t load_smem_a_ptr = __cvta_generic_to_shared(
&s_a[lane_id % 16][(lane_id / 16) * 8]);
LDMATRIX_X4(RA[0], RA[1], RA[2], RA[3], load_smem_a_ptr);
uint32_t load_smem_b_ptr = __cvta_generic_to_shared(
&s_b[lane_id % 16][0]);
LDMATRIX_X2_T(RB[0], RB[1], load_smem_b_ptr);

HMMA16816(RC[0], RC[1], RA[0], RA[1], RA[2], RA[3], RB[0], RB[1], RC[0], RC[1]);

__syncthreads();
}

// s_c[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]
LDST32BITS(s_c[lane_id / 4 ][(lane_id % 4) * 2]) = LDST32BITS(RC[0]);
LDST32BITS(s_c[lane_id / 4 + 8][(lane_id % 4) * 2]) = LDST32BITS(RC[1]);

__syncthreads();

// store s_c[16][8]
if (lane_id < MMA_M) {
// store 128 bits per memory issue.
int store_gmem_c_m = by * BM + lane_id;
int store_gmem_c_n = bx * BN;
int store_gmem_c_addr = store_gmem_c_m * N + store_gmem_c_n;
LDST128BITS(C[store_gmem_c_addr]) = (LDST128BITS(s_c[lane_id][0]));
}
}

1. 什么是 MMA 指令?

MMA (Matrix Multiply-Accumulate) 是 NVIDIA Tensor Core 提供的矩阵乘累加指令,可以在一条指令中完成小规模矩阵乘法。

1.1 m16n8k16 指令规格

1
2
3
4
5
6
7
8
9
mma.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16
│ │ │ │ │ │ │ └─ D 类型 (f16)
│ │ │ │ │ │ └─ C 类型 (f16)
│ │ │ │ │ └─ B 类型 (f16)
│ │ │ │ └─ A 类型 (f16)
│ │ │ └─ B 布局 (col-major)
│ │ └─ A 布局 (row-major)
│ └─ 矩阵维度: M=16, N=8, K=16
└─ 同步执行

计算: D[16×8] = A[16×16] × B[16×8] + C[16×8]

1.2 性能对比

方式 一条指令计算量 说明
CUDA Core (HFMA) 1×1×1 = 2 FLOPs 单个 FMA
Tensor Core (m16n8k16) 16×8×16×2 = 4096 FLOPs 一个 warp 协作

加速比: 4096 / (32 × 2) = 64倍 (理论上)


2. Kernel 概述

1
2
3
template<const int MMA_M=16, const int MMA_N=8, const int MMA_K=16>
__global__ void hgemm_mma_m16n8k16_naive_kernel(half* A, half* B, half* C,
int M, int N, int K)

2.1 Launch 配置

1
2
dim3 block(WARP_SIZE);  // 32 threads = 1 warp
dim3 grid(N/8, M/16); // 每个 block 计算 16×8 的输出

关键: MMA 指令需要整个 warp (32 threads) 协作执行!

2.2 Shared Memory 布局

1
2
3
__shared__ half s_a[MMA_M][MMA_K]; // 16×16 = 512 bytes
__shared__ half s_b[MMA_K][MMA_N]; // 16×8 = 256 bytes
__shared__ half s_c[MMA_M][MMA_N]; // 16×8 = 256 bytes

3. 数据加载 (Global → Shared)

3.1 加载 A 矩阵

1
2
3
4
5
6
7
// s_a[16][16]: 每行 16 个 half,每线程加载 8 个,需要 2 线程/行
// 共 16 行,需要 16×2 = 32 线程 (刚好 1 warp)
const int load_smem_a_m = tid / 2; // 行号: 0~15
const int load_smem_a_k = (tid % 2) * 8; // 列号: 0 或 8

// 加载 128 bits (8 × half)
LDST128BITS(s_a[load_smem_a_m][load_smem_a_k]) = LDST128BITS(A[addr]);

图示:

1
2
3
4
Thread 0,1   → s_a[0][0:7], s_a[0][8:15]   (第 0 行)
Thread 2,3 → s_a[1][0:7], s_a[1][8:15] (第 1 行)
...
Thread 30,31 → s_a[15][0:7], s_a[15][8:15] (第 15 行)

3.2 加载 B 矩阵

1
2
3
4
5
// s_b[16][8]: 每行 8 个 half,每线程加载 8 个,只需 1 线程/行
// 共 16 行,需要 16 线程 (只用一半 warp)
if (lane_id < MMA_K) { // lane_id < 16
LDST128BITS(s_b[load_smem_b_k][0]) = LDST128BITS(B[addr]);
}

4. ldmatrix 指令详解 ⭐

ldmatrix 是专为 MMA 设计的数据加载指令,从 Shared Memory 加载数据到寄存器,并自动重排数据布局以匹配 MMA 指令的要求。

4.1 ldmatrix.x4 加载 A

1
2
3
4
// 每个线程提供一个地址,协作加载 4 个 8×8 矩阵
uint32_t load_smem_a_ptr = __cvta_generic_to_shared(
&s_a[lane_id % 16][(lane_id / 16) * 8]);
LDMATRIX_X4(RA[0], RA[1], RA[2], RA[3], load_smem_a_ptr);

地址映射 (32 个线程的地址):

1
2
3
4
5
6
7
8
9
10
lane_id  |  row (lane_id % 16)  |  col ((lane_id/16)*8)  |  地址
---------|---------------------|----------------------|--------
0 | 0 | 0 | s_a[0][0]
1 | 1 | 0 | s_a[1][0]
...
15 | 15 | 0 | s_a[15][0]
16 | 0 | 8 | s_a[0][8]
17 | 1 | 8 | s_a[1][8]
...
31 | 15 | 8 | s_a[15][8]

结果: 每个线程获得 4 个 32-bit 寄存器 (RA[0]~RA[3]),包含 8 个 half 值。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
s_a[16][16] 在 Shared Memory 中:

col 0 col 8
↓ ↓
row 0 → [████████|████████] ← T0 指向 [0][0], T16 指向 [0][8]
row 1 → [████████|████████] ← T1 指向 [1][0], T17 指向 [1][8]
row 2 → [████████|████████] ← T2 指向 [2][0], T18 指向 [2][8]
...
row 15 → [████████|████████] ← T15 指向 [15][0], T31 指向 [15][8]
└──8 half──┘
s_a[16][16]:
col 0~7 col 8~15
┌─────────┬─────────┐
row 0 │ T0 ──► │ T16 ──► │ 每个线程提供一行的起始地址
row 1 │ T1 ──► │ T17 ──► │ ldmatrix 会从该地址读取 8 个 half (128 bits)
... │ ... │ ... │
row 15│ T15 ──► │ T31 ──► │
└─────────┴─────────┘

LDMATRIX_X4 做了什么?

  • 加载 4 个 8×8 子矩阵(共 16×16)
  • 硬件自动重新分发数据到各线程的寄存器
  • 输出:每线程得到 RA[0…3],共 4 个 u32(8 个 half)

4.2 ldmatrix.x2.trans 加载 B (带转置)

1
2
uint32_t load_smem_b_ptr = __cvta_generic_to_shared(&s_b[lane_id % 16][0]);
LDMATRIX_X2_T(RB[0], RB[1], load_smem_b_ptr);

地址计算: s_b[16][8] 是 B 矩阵的 shared memory

lane_id lane_id % 16 (行) 地址
0~15 0~15 s_b[0~15][0]
16~31 0~15 s_b[0~15][0] (重复)

图示:

1
2
3
4
5
6
7
8
9
10
11
s_b[16][8]:  (存储是行主序)
col 0~7
┌─────────┐
row 0 │ T0 ──► │
row 1 │ T1 ──► │
... │ ... │
row 15│ T15 ──► │
└─────────┘

LDMATRIX_X2_T 加载 2 个 8×8 矩阵,并转置!
转置后 B 变成列主序,满足 MMA 指令要求

为什么需要转置?

  • B 矩阵在 shared memory 中是行主序 s_b[K][N]
  • MMA 指令要求 B 是列主序
  • ldmatrix.trans 在加载时自动完成转置

4.3 ldmatrix 数据布局图解(硬件重新分发)

输入地址 和 输出寄存器 的对应关系不是直接的:

1
2
❌ 错误理解:T0 提供的地址 → T0 的寄存器
✅ 正确理解:32个线程提供32个地址 → 硬件收集所有数据 → 按MMA要求重新分配到各线程

具体过程

Step 1: 每个线程提供一个地址

1
2
3
4
5
6
7
8
&s_a[lane_id % 16][(lane_id / 16) * 8]
T0 → s_a[0][0] (读取 s_a[0][0:7], 8个half)
T1 → s_a[1][0] (读取 s_a[1][0:7], 8个half)
...
T15 → s_a[15][0] (读取 s_a[15][0:7], 8个half)
T16 → s_a[0][8] (读取 s_a[0][8:15], 8个half)
...
T31 → s_a[15][8] (读取 s_a[15][8:15], 8个half)

合计: 32 × 8 = 256 个 half = 整个 16×16 矩阵

Step 2: 硬件收集所有数据

ldmatrix 指令让硬件一次性从 shared memory 读取全部 256 个 half:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
┌─────────────────────────────────────────┐
│ Shared Memory s_a[16][16] │
│ ┌───┬───┬───┬───┬───┬───┬───┬───┬...┐ │
│ │a00│a01│a02│a03│a04│a05│a06│a07│...│ │ ← T0 提供的地址读这行
│ │a10│a11│a12│a13│a14│a15│a16│a17│...│ │ ← T1 提供的地址读这行
│ │... │ │
│ └───────────────────────────────────┘ │
└─────────────────────────────────────────┘

▼ 硬件收集全部 256 个 half
┌──────────┐
│ ldmatrix │
│ 硬件逻辑 │
└──────────┘

Step 3: 按 MMA 要求重新分发

关键:MMA 指令对输入数据的排布有特定要求!

关键公式(来自官方文档)

1
2
groupID           = %laneid >> 2      // lane_id / 4
threadID_in_group = %laneid % 4 // lane_id % 4

A 矩阵 Fragment(RA[0…3],4 个 u32 = 8 个 half)

官方公式

1
2
3
4
5
row = groupID            for ai where  0 <= i < 2 || 4 <= i < 6
groupID + 8 Otherwise

col = (threadID_in_group * 2) + (i & 0x1) for ai where i < 4
(threadID_in_group * 2) + (i & 0x1) + 8 for ai where i >= 4

硬件按照 NVIDIA 定义的 fragment 布局 将数据分发到各线程:

每线程持有的 A 矩阵元素:

lane_id groupID threadID a0,a1 位置 a2,a3 位置 a4,a5 位置 a6,a7 位置
0 0 0 A[0][0:1] A[8][0:1] A[0][8:9] A[8][8:9]
1 0 1 A[0][2:3] A[8][2:3] A[0][10:11] A[8][10:11]
2 0 2 A[0][4:5] A[8][4:5] A[0][12:13] A[8][12:13]
3 0 3 A[0][6:7] A[8][6:7] A[0][14:15] A[8][14:15]
4 1 0 A[1][0:1] A[9][0:1] A[1][8:9] A[9][8:9]
31 7 3 A[7][6:7] A[15][6:7] A[7][14:15] A[15][14:15]

图示:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
A[16][16] 矩阵的分布:

col 0-1 2-3 4-5 6-7 8-9 10-11 12-13 14-15
┌─────┬─────┬─────┬─────┬─────┬─────┬─────┬─────┐
row 0 │ T0 │ T1 │ T2 │ T3 │ T0 │ T1 │ T2 │ T3 │ ← a0,a1 / a4,a5
row 1 │ T4 │ T5 │ T6 │ T7 │ T4 │ T5 │ T6 │ T7 │
... │ │ │ │ │ │ │ │ │
row 7 │ T28 │ T29 │ T30 │ T31 │ T28 │ T29 │ T30 │ T31 │
row 8 │ T0 │ T1 │ T2 │ T3 │ T0 │ T1 │ T2 │ T3 │ ← a2,a3 / a6,a7
row 9 │ T4 │ T5 │ T6 │ T7 │ T4 │ T5 │ T6 │ T7 │
... │ │ │ │ │ │ │ │ │
row 15 │ T28 │ T29 │ T30 │ T31 │ T28 │ T29 │ T30 │ T31 │
└─────┴─────┴─────┴─────┴─────┴─────┴─────┴─────┘
k < 8 k >= 8

官方图示:

img


B 矩阵 Fragment(RB[0…1],2 个 u32 = 4 个 half)

官方公式:

1
2
3
4
row = (threadID_in_group * 2) + (i & 0x1)           for bi where i <  2
(threadID_in_group * 2) + (i & 0x1) + 8 for bi where i >= 2

col = groupID

每线程持有的 B 矩阵元素:

lane_id groupID threadID b0,b1 位置 b2,b3 位置
0 0 0 B[0:1][0] B[8:9][0]
1 0 1 B[2:3][0] B[10:11][0]
2 0 2 B[4:5][0] B[12:13][0]
3 0 3 B[6:7][0] B[14:15][0]
4 1 0 B[0:1][1] B[8:9][1]

官方图示:

img


C/D 矩阵 Fragment(RC[0…1],2 个 u32 = 4 个 half)

官方公式:

1
2
3
4
row = groupID                 for ci where i <  2
groupID + 8 for ci where i >= 2

col = (threadID_in_group * 2) + (i & 0x1) for ci where i = {0,..,3}

每线程持有的 C/D 矩阵元素:

lane_id groupID threadID c0,c1 (RC[0]) c2,c3 (RC[1])
0 0 0 C[0][0:1] C[8][0:1]
1 0 1 C[0][2:3] C[8][2:3]
2 0 2 C[0][4:5] C[8][4:5]
3 0 3 C[0][6:7] C[8][6:7]
4 1 0 C[1][0:1] C[9][0:1]
31 7 3 C[7][6:7] C[15][6:7]

图示:

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

col 0-1 2-3 4-5 6-7
┌─────┬─────┬─────┬─────┐
row 0 │ T0 │ T1 │ T2 │ T3 │ ← RC[0]: c0,c1
row 1 │ T4 │ T5 │ T6 │ T7 │
row 2 │ T8 │ T9 │ T10 │ T11 │
... │ │ │ │ │
row 7 │ T28 │ T29 │ T30 │ T31 │
row 8 │ T0 │ T1 │ T2 │ T3 │ ← RC[1]: c2,c3
row 9 │ T4 │ T5 │ T6 │ T7 │
... │ │ │ │ │
row 15 │ T28 │ T29 │ T30 │ T31 │
└─────┴─────┴─────┴─────┘

官方图示:

img


5. MMA 指令执行

1
2
3
4
5
6
7
8
9
10
11
12
#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))

// 调用
HMMA16816(RC[0], RC[1],
RA[0], RA[1], RA[2], RA[3], // A: 4 个 32-bit 寄存器
RB[0], RB[1], // B: 2 个 32-bit 寄存器
RC[0], RC[1]); // C/D: 2 个 32-bit 寄存器

5.1 寄存器用量

矩阵 维度 每线程元素数 寄存器数
A 16×16 8 half 4 × 32-bit
B 16×8 4 half 2 × 32-bit
C/D 16×8 4 half 2 × 32-bit

整个 Warp:

  • A: 32 threads × 8 = 256 half = 16×16 ✓
  • B: 32 threads × 4 = 128 half = 16×8 ✓
  • C: 32 threads × 4 = 128 half = 16×8 ✓

寄存器分布(每线程持有的数据)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
A矩阵 (RA[0~3], 每个 u32 = 2 个 half):
┌─────────────────────────────────┐
│ 每线程持有 A 的某些行的部分元素 │
│ 具体位置由硬件 MMA 布局决定 │
└─────────────────────────────────┘

C/D矩阵 (RC[0], RC[1]):
┌───────────────────┐
│ RC[0]: 行 0~7 │ lane_id / 4 决定行
│ RC[1]: 行 8~15 │ lane_id % 4 决定列组
└───────────────────┘

输出映射 (RC → C[16][8]):
lane_id │ RC[0] 位置 │ RC[1] 位置
─────────┼───────────────────┼───────────────────
0 │ C[0][0:1] │ C[8][0:1]
1 │ C[0][2:3] │ C[8][2:3]
2 │ C[0][4:5] │ C[8][4:5]
3 │ C[0][6:7] │ C[8][6:7]
4 │ C[1][0:1] │ C[9][0:1]
... │ ... │ ...
31 │ C[7][6:7] │ C[15][6:7]

数据流

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
┌──────────────────────────────────────────────────────────┐
│ Global Memory │
│ A[M×K] B[K×N] │
└─────┬────────────────┬───────────────────────────────────┘
│ LDST128BITS │ LDST128BITS
▼ ▼
┌──────────────────────────────────────────────────────────┐
│ Shared Memory │
│ s_a[16][16] s_b[16][8] │
└─────┬────────────────┬───────────────────────────────────┘
│ LDMATRIX_X4 │ LDMATRIX_X2_T (带转置)
▼ ▼
┌──────────────────────────────────────────────────────────┐
│ Registers (每线程) │
│ RA[0..3] RB[0..1] RC[0..1] │
│ (4×u32) (2×u32) (2×u32) │
└─────┬────────────────┬────────────────┬─────────────────┘
└────────┬───────┘ │
▼ │
┌──────────┐ │
│ HMMA16816│ ◄─────────────────┘
│ Tensor │
│ Core │
└────┬─────┘


RC[0], RC[1] (累加结果)

6. 输出矩阵 C 的数据布局 ⭐

MMA 指令的输出分布在 32 个线程的寄存器中,布局如下:

6.1 官方文档布局

根据 PTX ISA:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
C/D Matrix [16][8] 在各线程寄存器中的分布:

col 0,1 col 2,3 col 4,5 col 6,7
─────────────────────────────────────────
row 0 │ T0:RC[0] │ T1:RC[0] │ T2:RC[0] │ T3:RC[0] │
row 1 │ T4:RC[0] │ T5:RC[0] │ T6:RC[0] │ T7:RC[0] │
...
row 7 │T28:RC[0] │T29:RC[0] │T30:RC[0] │T31:RC[0] │
─────────────────────────────────────────
row 8 │ T0:RC[1] │ T1:RC[1] │ T2:RC[1] │ T3:RC[1] │
row 9 │ T4:RC[1] │ T5:RC[1] │ T6:RC[1] │ T7:RC[1] │
...
row 15 │T28:RC[1] │T29:RC[1] │T30:RC[1] │T31:RC[1] │
─────────────────────────────────────────

img

6.2 索引公式

1
2
3
4
5
6
7
// 每个线程的 RC[0] 对应的位置
row = lane_id / 4 // 0~7
col = (lane_id % 4) * 2 // 0,2,4,6

// 每个线程的 RC[1] 对应的位置
row = lane_id / 4 + 8 // 8~15
col = (lane_id % 4) * 2 // 0,2,4,6

6.3 存储到 Shared Memory

1
2
3
4
5
// RC[0] → s_c 的上半部分 (row 0~7)
LDST32BITS(s_c[lane_id / 4 ][(lane_id % 4) * 2]) = LDST32BITS(RC[0]);

// RC[1] → s_c 的下半部分 (row 8~15)
LDST32BITS(s_c[lane_id / 4 + 8][(lane_id % 4) * 2]) = LDST32BITS(RC[1]);

线程分布表

lane_id lane_id/4 lane_id%4 RC[0] 写入位置 RC[1] 写入位置
0 0 0 s_c[0][0:1] s_c[8][0:1]
1 0 1 s_c[0][2:3] s_c[8][2:3]
2 0 2 s_c[0][4:5] s_c[8][4:5]
3 0 3 s_c[0][6:7] s_c[8][6:7]
4 1 0 s_c[1][0:1] s_c[9][0:1]
5 1 1 s_c[1][2:3] s_c[9][2:3]
6 1 2 s_c[1][4:5] s_c[9][4:5]
7 1 3 s_c[1][6:7] s_c[9][6:7]
28 7 0 s_c[7][0:1] s_c[15][0:1]
29 7 1 s_c[7][2:3] s_c[15][2:3]
30 7 2 s_c[7][4:5] s_c[15][4:5]
31 7 3 s_c[7][6:7] s_c[15][6:7]

图示:s_c[16][8] 的填充

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
s_c[16][8] 矩阵:

col 0-1 col 2-3 col 4-5 col 6-7
┌─────────┬─────────┬─────────┬─────────┐
row 0 │ T0 RC[0]│ T1 RC[0]│ T2 RC[0]│ T3 RC[0]│ ← lane 0~3 写 RC[0]
row 1 │ T4 RC[0]│ T5 RC[0]│ T6 RC[0]│ T7 RC[0]│ ← lane 4~7 写 RC[0]
row 2 │ T8 RC[0]│ T9 RC[0]│T10 RC[0]│T11 RC[0]│
row 3 │T12 RC[0]│T13 RC[0]│T14 RC[0]│T15 RC[0]│
row 4 │T16 RC[0]│T17 RC[0]│T18 RC[0]│T19 RC[0]│
row 5 │T20 RC[0]│T21 RC[0]│T22 RC[0]│T23 RC[0]│
row 6 │T24 RC[0]│T25 RC[0]│T26 RC[0]│T27 RC[0]│
row 7 │T28 RC[0]│T29 RC[0]│T30 RC[0]│T31 RC[0]│ ← lane 28~31 写 RC[0]
├─────────┼─────────┼─────────┼─────────┤
row 8 │ T0 RC[1]│ T1 RC[1]│ T2 RC[1]│ T3 RC[1]│ ← lane 0~3 写 RC[1]
row 9 │ T4 RC[1]│ T5 RC[1]│ T6 RC[1]│ T7 RC[1]│ ← lane 4~7 写 RC[1]
row 10 │ T8 RC[1]│ T9 RC[1]│T10 RC[1]│T11 RC[1]│
row 11 │T12 RC[1]│T13 RC[1]│T14 RC[1]│T15 RC[1]│
row 12 │T16 RC[1]│T17 RC[1]│T18 RC[1]│T19 RC[1]│
row 13 │T20 RC[1]│T21 RC[1]│T22 RC[1]│T23 RC[1]│
row 14 │T24 RC[1]│T25 RC[1]│T26 RC[1]│T27 RC[1]│
row 15 │T28 RC[1]│T29 RC[1]│T30 RC[1]│T31 RC[1]│ ← lane 28~31 写 RC[1]
└─────────┴─────────┴─────────┴─────────┘

关键理解

1
2
3
4
5
6
7
8
9
10
11
12
32 个线程,每线程写 2 次(RC[0] 和 RC[1])
每次写 2 个 half(32 bits)

RC[0] 填充 row 0~7 (上半部分)
RC[1] 填充 row 8~15 (下半部分)

同一行的 4 个线程(同一 group)横向排列:
T0 写 col 0-1
T1 写 col 2-3
T2 写 col 4-5
T3 写 col 6-7
→ 刚好填满一行 8 列

与 MMA 输出 Fragment 的对应

这个存储模式完全匹配 MMA m16n8k16 的 C/D fragment 布局:

1
2
3
4
5
6
7
8
9
官方公式:
row = groupID for ci where i < 2 → RC[0]
groupID + 8 for ci where i >= 2 → RC[1]
col = (threadID_in_group * 2) + (i & 0x1)

代码:
row = lane_id / 4 → groupID
lane_id / 4 + 8 → groupID + 8
col = (lane_id % 4) * 2 → threadID_in_group * 2

两者完全一致! 代码正确地将 MMA 计算结果从寄存器写回 shared memory。

6.4 s_c → Global Memory C 的存储解释

1
2
3
4
5
6
if (lane_id < MMA_M) {  // MMA_M = 16
int store_gmem_c_m = by * BM + lane_id; // 全局行号
int store_gmem_c_n = bx * BN; // 全局列号起点
int store_gmem_c_addr = store_gmem_c_m * N + store_gmem_c_n;
LDST128BITS(C[store_gmem_c_addr]) = (LDST128BITS(s_c[lane_id][0]));
}

代码含义

部分 含义
lane_id < 16 只用前 16 个线程(32 个线程中有 16 个空闲)
LDST128BITS 一次写 128 bits = 8 个 half
s_c[lane_id][0] lane_id 行,从列 0 开始读 8 个元素

图示

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
s_c[16][8] → C[M][N] 的全局位置

s_c 在 shared memory:
col 0 1 2 3 4 5 6 7
┌──────────────────────────┐
row 0 │ ████████████████████████│ ← T0 写整行 (8 half = 128 bits)
row 1 │ ████████████████████████│ ← T1 写整行
row 2 │ ████████████████████████│ ← T2 写整行
... │ ... │
row 15 │ ████████████████████████│ ← T15 写整行
└──────────────────────────┘

写到 C 的全局位置:

C[M][N]:
┌─────────────────────────────────────────┐
│ │
│ bx*BN │
│ ↓ │
by*BM →├───────┬────────┬────────────────────────┤
│ │ 8 cols │ │
T0 → │ │████████│ ← C[by*BM+0][bx*BN:+8]│
T1 → │ │████████│ ← C[by*BM+1][bx*BN:+8]│
... │ │ ... │ │
T15 → │ │████████│ ← C[by*BM+15][bx*BN:+8]│
├───────┴────────┴────────────────────────┤
│ │
└─────────────────────────────────────────┘

线程分工

lane_id 负责 写入全局地址
0 s_c[0][0:7] → C[byBM+0][bxBN : bx*BN+8] 一整行
1 s_c[1][0:7] → C[byBM+1][bxBN : bx*BN+8] 一整行
15 s_c[15][0:7] → C[byBM+15][bxBN : bx*BN+8] 一整行
16~31 空闲 不参与

为什么只用 16 个线程?

  • s_c 是 16×8 矩阵,共 16 行
  • 每线程写一整行(8 个 half = 128 bits)
  • 16 线程刚好覆盖 16 行
  • 剩余 16 个线程(lane_id 16~31)被 if 过滤掉

效率说明

1
2
3
4
5
6
优点:每次访存 128 bits,合并访问效率高
缺点:一半线程空闲(利用率 50%)

这是 naive kernel 的简单实现,
后面的优化版本(mma2x4_warp4x4)会直接从寄存器写 gmem,
避免经过 shared memory 中转

7. 完整执行流程

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
┌─────────────────────────────────────────────────────────────────┐
│ Step 1: Global Memory → Shared Memory │
│ ─────────────────────────────────────── │
│ A[M][K] → s_a[16][16] (32 threads, 128-bit load) │
│ B[K][N] → s_b[16][8] (16 threads, 128-bit load) │
│ __syncthreads() │
└─────────────────────────────────────────────────────────────────┘


┌─────────────────────────────────────────────────────────────────┐
│ Step 2: Shared Memory → Registers (ldmatrix) │
│ ───────────────────────────────────────────── │
│ s_a[16][16] → RA[4] (LDMATRIX_X4) │
│ s_b[16][8] → RB[2] (LDMATRIX_X2_T, 带转置) │
└─────────────────────────────────────────────────────────────────┘


┌─────────────────────────────────────────────────────────────────┐
│ Step 3: MMA Compute │
│ ─────────────────── │
│ HMMA16816: RC = RA × RB + RC │
│ (整个 warp 协作执行) │
└─────────────────────────────────────────────────────────────────┘


┌─────────────────────────────────────────────────────────────────┐
│ Step 4: Registers → Shared Memory → Global Memory │
│ ───────────────────────────────────────────────── │
│ RC[2] → s_c[16][8] → C[M][N] │
└─────────────────────────────────────────────────────────────────┘

8. K 维度循环

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
for (int k = 0; k < NUM_K_TILES; ++k) {
// 1. 加载当前 K tile 的数据
load A[by*16:(by+1)*16, k*16:(k+1)*16] → s_a
load B[k*16:(k+1)*16, bx*8:(bx+1)*8] → s_b
__syncthreads();

// 2. ldmatrix 加载到寄存器
LDMATRIX_X4(RA[0], RA[1], RA[2], RA[3], s_a_ptr);
LDMATRIX_X2_T(RB[0], RB[1], s_b_ptr);

// 3. MMA 累加
HMMA16816(RC[0], RC[1], RA[], RB[], RC[0], RC[1]);

__syncthreads();
}

计算量:

  • 每次迭代: 16 × 8 × 16 × 2 = 4096 FLOPs
  • K/16 次迭代
  • 总计: 16 × 8 × K × 2 FLOPs

9. 性能分析

9.1 理论计算量

对于 M×K × K×N 的矩阵乘法:

  • Grid: (N/8) × (M/16) blocks
  • 每 block: K/16 次 MMA
  • 总 MMA 数: (N/8) × (M/16) × (K/16) = M×N×K / 2048

9.2 限制因素

因素 影响
只有 1 warp/block 严重限制 SM 利用率
无 double buffering 无法隐藏内存延迟
小 block size Grid 巨大,调度开销大

这就是为什么需要 Level 2 (多 warp) 和 Level 3 (多阶段) 的优化!


10. 总结

10.1 关键概念

概念 说明
MMA 指令 一条指令完成 16×8×16 矩阵乘累加
Warp 协作 32 个线程共同执行一条 MMA
ldmatrix 专用加载指令,自动重排数据
ldmatrix.trans 加载时转置,适配列主序 B
C 布局 输出分散在 32 个线程的寄存器中

10.2 数据分布

1
2
3
4
5
6
7
8
9
10
11
12
13
14
             ┌──────────────────┐
│ 32 Threads │
│ (1 Warp) │
└────────┬─────────┘

┌────────────────┼────────────────┐
│ │ │
▼ ▼ ▼
┌─────────┐ ┌─────────┐ ┌─────────┐
│ A[16×16]│ │ B[16×8] │ │ C[16×8] │
│ 8 half │ │ 4 half │ │ 4 half │
│ per thd │ │ per thd │ │ per thd │
│ RA[4] │ │ RB[2] │ │ RC[2] │
└─────────┘ └─────────┘ └─────────┘

10.3 下一步

学完 naive kernel 后,继续学习 Level 2: hgemm_mma_m16n8k16_mma2x4_warp4x4_kernel,了解如何组织多个 MMA 操作和多个 Warp 协作。

  • 标题: hgemm_mma_m16n8k16_naive_kernel
  • 作者: 鱿鱼圈
  • 创建于 : 2026-03-02 23:50:00
  • 更新于 : 2026-06-05 23:02:25
  • 链接: https://yuyanqi.com/2026/03/02/hgemm_mma_m16n8k16_naive_kernel/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论