cute(7)simple_gemm

鱿鱼圈 Lv4

前置阅读

7 简单 GEMM(不用 Shared Memory)

对应代码:07_simple_gemm.cu 需要 GPU

核心概念

用 TiledMMA 实现真正的矩阵乘法:

  • 直接从 global memory 读到寄存器(不经过 shared memory)
  • partition_A / partition_B / partition_C:把 global tensor 按 MMA 分到每个线程
  • partition_fragment_A / B / C:创建寄存器级别的 fragment
  • cute::gemm:自动遍历所有子块执行 mma.sync

代码实现

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
template <int kTileM, int kTileN, int kTileK, typename TiledMMA>
__global__ void simple_gemm_kernel(half_t* Dptr, const half_t* Aptr,
const half_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);

// MMA
cute::gemm(tiled_mma, tCrD, tCrA, tCrB, tCrD);
}

// 写回 global
cute::copy(tCrD, tDgD);
}

kernel加入调试信息

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
// 打印调试信息(仅线程 0,仅第一个 block)
if (threadIdx.x == 0 && blockIdx.x == 0 && blockIdx.y == 0) {
printf("\n === Tensor Shapes ===\n");
printf(" gA : "); print(gA.layout()); printf("\n");
printf(" gB : "); print(gB.layout()); printf("\n");
printf(" gD : "); print(gD.layout()); printf("\n");

printf("\n === partition_A/B/C (global tensor -> 线程级别) ===\n");
printf(" tAgA shape: "); print(shape(tAgA)); printf(" layout: "); print(tAgA.layout()); printf("\n");
printf(" tBgB shape: "); print(shape(tBgB)); printf(" layout: "); print(tBgB.layout()); printf("\n");
printf(" tDgD shape: "); print(shape(tDgD)); printf(" layout: "); print(tDgD.layout()); printf("\n");

printf("\n === partition_fragment (寄存器 fragment) ===\n");
printf(" tCrA shape: "); print(shape(tCrA)); printf("\n");
printf(" tCrB shape: "); print(shape(tCrB)); printf("\n");
printf(" tCrD shape: "); print(shape(tCrD)); printf("\n");

printf("\n === TiledMMA info ===\n");
printf(" TiledMMA size (threads): %d\n", int(size(tiled_mma)));
printf(" num_tile_k: %d\n", int(size<3>(tAgA)));

// 打印 tAgA:线程 0 读 A 的哪些位置
printf("\n === tAgA 线程 0 的元素位置 (m, k) ===\n");
printf(" 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)));

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");
}
}
}

1. 整体数据流

1
2
3
4
5
D(M,N) = A(M,K) @ B(N,K)^T

Global(A,B) ──直接──> Register(tCrA,tCrB) ──mma──> Register(tCrD) ──直接──> Global(D)
↑ ↑
没有 smem 没有 smem

每个 block 处理输出矩阵 D 的一个 (128×128) tile,K 方向按 kTileK=16 循环 4 轮(K=64)。


2. TiledMMA 的二级扩展

1
2
3
4
5
6
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));
  • A, B 是 row-major,D = A @ B^T
  • _(下划线)保留 K 维度不切,所以 gA 多出第 3 维 num_tile_k
  • gD 不保留 K 维度(输出不需要)

打印信息如下

1
2
3
4
5
=== Tensor Shapes ===
gA : (_128,_16,4):(64,_1,_16)
gB : (_128,_16,4):(64,_1,_16)
gD : (_128,_128):(256,_1)

1
2
3
4
5
gA: (128, 16, 4) : (64, 1, 16)
↑M ↑K ↑num_tile_k

gD: (128, 128) : (256, 1)
↑M ↑N

3.2 partition_A / partition_B / partition_C

1
2
3
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)

把 global tensor 按 TiledMMA 的分配规则切到线程级别。每个线程只看到自己要读/写的部分。

打印信息如下

1
2
3
4
=== partition_A/B/C (global tensor -> 线程级别) ===
tAgA shape: ((_2,_2,_2),_4,_1,4) layout: ((_2,_2,_2),_4,_1,4):((_1,512,_8),2048,_0,_16)
tBgB shape: ((_2,_2),(_2,_4),_1,4) layout: ((_2,_2),(_2,_4),_1,4):((_1,_8),(1024,2048),_0,_16)
tDgD shape: ((_2,_2),_4,_8) layout: ((_2,_2),_4,_8):((_1,2048),8192,_16)

3.3 partition_fragment_A / B / C

1
2
3
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)

在寄存器中创建 fragment。tCrA/B 是数据缓冲区,tCrD 是累加器。

打印信息如下

1
2
3
4
=== partition_fragment (寄存器 fragment) ===
tCrA shape: ((_2,_2,_2),_4,_1)
tCrB shape: ((_2,_2),(_2,_4),_1)
tCrD shape: ((_2,_2),_4,_8)

3.4 K 方向循环

1
2
3
4
5
6
7
int num_tile_k = size<3>(tAgA);   // = 4
for (int ik = 0; ik < num_tile_k; ++ik) {
cute::copy(tAgA(_, _, _, ik), tCrA); // Global → Register
cute::copy(tBgB(_, _, _, ik), tCrB);
cute::gemm(tiled_mma, tCrD, tCrA, tCrB, tCrD); // MMA 累加
}
cute::copy(tCrD, tDgD); // Register → Global

4. tAgA 详解

1
2
3
4
// 切分 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)

打印信息如下

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
=== TiledMMA info ===
TiledMMA size (threads): 128
num_tile_k: 4

=== tAgA 线程 0 的元素位置 (m, k) ===
shape: (MMA=8, MMA_M=4, MMA_K=1, num_tile_k=4)

--- K tile 0 (ik=0) ---
MMA_M=0, MMA_K=0: A[ 0][ 0] A[ 0][ 1] A[ 8][ 0] A[ 8][ 1] A[ 0][ 8] A[ 0][ 9] A[ 8][ 8] A[ 8][ 9]
MMA_M=1, MMA_K=0: A[ 32][ 0] A[ 32][ 1] A[ 40][ 0] A[ 40][ 1] A[ 32][ 8] A[ 32][ 9] A[ 40][ 8] A[ 40][ 9]
MMA_M=2, MMA_K=0: A[ 64][ 0] A[ 64][ 1] A[ 72][ 0] A[ 72][ 1] A[ 64][ 8] A[ 64][ 9] A[ 72][ 8] A[ 72][ 9]
MMA_M=3, MMA_K=0: A[ 96][ 0] A[ 96][ 1] A[104][ 0] A[104][ 1] A[ 96][ 8] A[ 96][ 9] A[104][ 8] A[104][ 9]

--- K tile 1 (ik=1) ---
MMA_M=0, MMA_K=0: A[ 0][16] A[ 0][17] A[ 8][16] A[ 8][17] A[ 0][24] A[ 0][25] A[ 8][24] A[ 8][25]
MMA_M=1, MMA_K=0: A[ 32][16] A[ 32][17] A[ 40][16] A[ 40][17] A[ 32][24] A[ 32][25] A[ 40][24] A[ 40][25]
MMA_M=2, MMA_K=0: A[ 64][16] A[ 64][17] A[ 72][16] A[ 72][17] A[ 64][24] A[ 64][25] A[ 72][24] A[ 72][25]
MMA_M=3, MMA_K=0: A[ 96][16] A[ 96][17] A[104][16] A[104][17] A[ 96][24] A[ 96][25] A[104][24] A[104][25]

--- K tile 2 (ik=2) ---
MMA_M=0, MMA_K=0: A[ 0][32] A[ 0][33] A[ 8][32] A[ 8][33] A[ 0][40] A[ 0][41] A[ 8][40] A[ 8][41]
MMA_M=1, MMA_K=0: A[ 32][32] A[ 32][33] A[ 40][32] A[ 40][33] A[ 32][40] A[ 32][41] A[ 40][40] A[ 40][41]
MMA_M=2, MMA_K=0: A[ 64][32] A[ 64][33] A[ 72][32] A[ 72][33] A[ 64][40] A[ 64][41] A[ 72][40] A[ 72][41]
MMA_M=3, MMA_K=0: A[ 96][32] A[ 96][33] A[104][32] A[104][33] A[ 96][40] A[ 96][41] A[104][40] A[104][41]

--- K tile 3 (ik=3) ---
MMA_M=0, MMA_K=0: A[ 0][48] A[ 0][49] A[ 8][48] A[ 8][49] A[ 0][56] A[ 0][57] A[ 8][56] A[ 8][57]
MMA_M=1, MMA_K=0: A[ 32][48] A[ 32][49] A[ 40][48] A[ 40][49] A[ 32][56] A[ 32][57] A[ 40][56] A[ 40][57]
MMA_M=2, MMA_K=0: A[ 64][48] A[ 64][49] A[ 72][48] A[ 72][49] A[ 64][56] A[ 64][57] A[ 72][56] A[ 72][57]
MMA_M=3, MMA_K=0: A[ 96][48] A[ 96][49] A[104][48] A[104][49] A[ 96][56] A[ 96][57] A[104][56] A[104][57]

4.1 Shape 和 Stride

1
2
3
tAgA shape:  ((_2,_2,_2),  _4,    _1,     4   )
tAgA stride: ((_1,512,_8), 2048, _0, _16 )
第①层 第②层 第③层 第④层

4.2 从外往里逐层解读

第④层:num_tile_k = 4,stride = 16

K 方向的外层循环,每轮所有地址往右平移 16 列:

1
2
3
4
ik=0: 读 K 的第 0~15 列   (基地址 + 0)
ik=1: 读 K 的第 16~31 列 (基地址 + 16)
ik=2: 读 K 的第 32~47 列 (基地址 + 32)
ik=3: 读 K 的第 48~63 列 (基地址 + 48)

第③层:MMA_K = 1,stride = 0

kTileK=16,MMA 的 K 也是 16,一次就覆盖。只有 1 个元素,不需要循环。

第②层:MMA_M = 4,stride = 2048

stride = 2048 = 32 行 × 64(行 stride),每次往下跳 32 行:

1
2
3
4
MMA_M=0: m = 0~31    (基地址 + 0)
MMA_M=1: m = 32~63 (基地址 + 2048)
MMA_M=2: m = 64~95 (基地址 + 4096)
MMA_M=3: m = 96~127 (基地址 + 6144)

第①层:MMA = (2,2,2),stride = (1, 512, 8)

一次 MMA 指令中,线程 0 要读的 8 个元素,位置由硬件规定:

子维度 shape stride 含义
最内层 2 1 同一行相邻两列 (k, k+1)
中间层 2 512 = 8×64 隔 8 行 (m 和 m+8)
最外层 2 8 跳 8 列 (前半 K 和后半 K)

4.3 嵌套括号的阅读方法

CuTe 的括号嵌套规则:括号 = 分组,不影响计算。

地址计算永远是:把所有嵌套展平,一一对应相乘再相加。

1
offset = a×1 + b×512 + c×8 + mma_m×2048 + mma_k×0 + ik×16

括号只表示语义分组——告诉你"这几个维度逻辑上是一组":

1
2
3
4
shape:  ((2,2,2),   4,     1,     4    )
↑ ↑ ↑ ↑
Atom硬件 EU扩展 K内部 K外层
8个值 128/32 16/16 64/16

不同的括号方式,计算完全一样,但语义不同:

写法 维度数 语义
((2,2,2), 4, 1, 4) 4 维 MMA内部 / M重复 / K内 / K外
((2,2,2), (4,1,4)) 2 维 MMA内部 / 其他全部
(2,2,2,4,1,4) 6 维 全部展平,失去分组信息

4.4 线程 0 在 ik=0 时读的 A 位置

1
2
3
4
MMA_M=0: A[  0][ 0] A[  0][ 1] A[  8][ 0] A[  8][ 1] A[  0][ 8] A[  0][ 9] A[  8][ 8] A[  8][ 9]
MMA_M=1: A[ 32][ 0] A[ 32][ 1] A[ 40][ 0] A[ 40][ 1] A[ 32][ 8] A[ 32][ 9] A[ 40][ 8] A[ 40][ 9]
MMA_M=2: A[ 64][ 0] A[ 64][ 1] A[ 72][ 0] A[ 72][ 1] A[ 64][ 8] A[ 64][ 9] A[ 72][ 8] A[ 72][ 9]
MMA_M=3: A[ 96][ 0] A[ 96][ 1] A[104][ 0] A[104][ 1] A[ 96][ 8] A[ 96][ 9] A[104][ 8] A[104][ 9]

画在 tile 上(以 MMA_M=0 为例):

1
2
3
4
5
6
7
8
9
A 矩阵 (128×16) 的前 16 行:

k= 0 1 2 3 4 5 6 7 8 9 10 ... 15
m= 0 [ ★ ★ . . . . . . ★ ★ . ... . ]
m= 1 [ . . . . . . . . . . . ... . ]
...
m= 8 [ ★ ★ . . . . . . ★ ★ . ... . ]
...
m= 15 [ . . . . . . . . . . . ... . ]

MMA_M=1 时同样的模式出现在 m=32,40 行,MMA_M=2 在 m=64,72 行,MMA_M=3 在 m=96,104 行。

4.5 每一维的来源总结

维度 shape 来源 意义
MMA (2,2,2)=8 mma.sync 硬件规定 线程 0 固定读 A 的第 0,8 行的第 0,1,8,9 列
MMA_M 4 tile_M / EU_M = 128/32 M 方向 4 段:0, 32, 64, 96
MMA_K 1 tile_K / atom_K = 16/16 K 方向一次 MMA 就覆盖
num_tile_k 4 K / tile_K = 64/16 外层 K 循环 4 轮

线程 0 总共读 8 × 4 × 1 × 4 = 128 个 A 元素


5. tCrA 详解(寄存器 Fragment)

1
2
3
4
// 创建寄存器 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)

打印信息如下

1
2
3
4
5
6
7
8
9
10
11
12
13
14
=== tCrA 寄存器 fragment (线程 0) ===
shape: ((_2,_2,_2),_4,_1)
layout: ((_2,_2,_2),_4,_1):((_1,_2,_4),_8,_0)
size: 32 个寄存器

tCrA vs tAgA 的对应关系:
tCrA shape: (MMA=8, MMA_M=4, MMA_K=1)
tAgA shape: (MMA=8, MMA_M=4, MMA_K=1, num_tile_k=4)

tCrA 中每个元素的寄存器编号和逻辑位置:
MMA_M=0, MMA_K=0: reg[ 0] reg[ 1] reg[ 2] reg[ 3] reg[ 4] reg[ 5] reg[ 6] reg[ 7]
MMA_M=1, MMA_K=0: reg[ 8] reg[ 9] reg[10] reg[11] reg[12] reg[13] reg[14] reg[15]
MMA_M=2, MMA_K=0: reg[16] reg[17] reg[18] reg[19] reg[20] reg[21] reg[22] reg[23]
MMA_M=3, MMA_K=0: reg[24] reg[25] reg[26] reg[27] reg[28] reg[29] reg[30] reg[31]

5.1 基本信息

1
2
3
4
5
shape:  ((_2,_2,_2), _4,   _1  )
stride: ((_1,_2,_4), _8, _0 )
MMA MMA_M MMA_K

总共 8 × 4 × 1 = 32 个寄存器

5.2 tCrA 与 tAgA 的关系

1
2
tCrA: (MMA=8, MMA_M=4, MMA_K=1)        ← 32 个寄存器,前三维
tAgA: (MMA=8, MMA_M=4, MMA_K=1, 4) ← 多了 num_tile_k=4

前三维完全相同,tCrA 就是 tAgA 去掉最后一维。tCrA 是寄存器缓冲区,每轮 K 循环被覆盖写入新数据:

1
2
3
4
for (int ik = 0; ik < 4; ++ik) {
cute::copy(tAgA(_, _, _, ik), tCrA); // 第 ik 轮的数据覆盖同一个 tCrA
cute::gemm(tiled_mma, tCrD, tCrA, tCrB, tCrD); // 用完即废,下轮覆盖
}

5.3 Stride 的含义

tCrA 的 stride (1, 2, 4, 8, 0)寄存器内部的连续编号,跟 global memory 的 stride 完全无关:

1
reg_idx = v0×1 + v1×2 + v2×4 + mma_m×8 + mma_k×0
stride 含义
v0 × 1 相邻元素 MMA 内第 0 层
v1 × 2 跳 2 MMA 内第 1 层
v2 × 4 跳 4 MMA 内第 2 层
mma_m × 8 跳 8 下一个 32 行段

本质是 32 个寄存器按 col-major 连续排列,没有任何 gap。

5.4 寄存器编号与 A 矩阵位置的完整对应

1
2
3
4
5
6
7
8
9
10
11
12
13
14
                   tAgA 在 global 中的位置 (ik=0)           tCrA 寄存器编号
────────────────────────────── ──────────────

MMA_M=0 的 8 值: A[0][0] A[0][1] A[8][0] A[8][1] A[0][8] A[0][9] A[8][8] A[8][9]
reg[0] reg[1] reg[2] reg[3] reg[4] reg[5] reg[6] reg[7]

MMA_M=1 的 8 值: A[32][0] A[32][1] A[40][0] A[40][1] A[32][8] A[32][9] A[40][8] A[40][9]
reg[8] reg[9] reg[10] reg[11] reg[12] reg[13] reg[14] reg[15]

MMA_M=2 的 8 值: A[64][0] A[64][1] A[72][0] A[72][1] A[64][8] A[64][9] A[72][8] A[72][9]
reg[16] reg[17] reg[18] reg[19] reg[20] reg[21] reg[22] reg[23]

MMA_M=3 的 8 值: A[96][0] A[96][1] A[104][0] A[104][1] A[96][8] A[96][9] A[104][8] A[104][9]
reg[24] reg[25] reg[26] reg[27] reg[28] reg[29] reg[30] reg[31]

5.5 总结

1
2
3
4
5
6
7
8
tCrA 的角色:

tAgA(_, _, _, ik) ──copy──> tCrA ──gemm──> tCrD(累加器)
global memory 寄存器 寄存器

- 32 个 half_t 的寄存器数组
- 每轮 K 循环被覆盖写一次,用完即废
- shape 与 tAgA 前三维一模一样,一一对应 copy

6. cute::gemm 做了什么

6.1 一次调用的计算范围

cute::gemm(tiled_mma, tCrD, tCrA, tCrB, tCrD) 一次调用计算的是整个 tile 的一个 K 切片,即 128(M) × 128(N) × 16(K),不是 32×32×16。

6.2 内部实现

cute::gemm 会遍历 tCrD 的 MMA_M × MMA_N 所有子块:

1
2
3
4
5
6
7
8
9
10
11
// cute::gemm 内部等价伪代码:
for (int mma_m = 0; mma_m < 4; ++mma_m) { // MMA_M=4
for (int mma_n = 0; mma_n < 8; ++mma_n) { // MMA_N=8
for (int mma_k = 0; mma_k < 1; ++mma_k) { // MMA_K=1
mma.sync(tCrD(_, mma_m, mma_n),
tCrA(_, mma_m, mma_k),
tCrB(_, mma_n, mma_k));
}
}
}
// 共 4 × 8 × 1 = 32 次 mma.sync 调用

6.3 32 次 atom 怎么拼成 128×128

每次 mma.sync 覆盖 32(M) × 16(N)(EU Repeat 后的大小),不是 32×32:

1
2
3
4
5
6
7
8
9
128×128 的输出 tile(32 个子块,每个 32×16):

N 方向: 128 列, 分 8 段, 每段 16 列
┌────┬────┬────┬────┬────┬────┬────┬────┐
m=0 │ ① │ ② │ ③ │ ④ │ ⑤ │ ⑥ │ ⑦ │ ⑧ │ 32行
m=1 │ ⑨ │ ⑩ │ ⑪ │ ⑫ │ ⑬ │ ⑭ │ ⑮ │ ⑯ │ 32行
m=2 │ ⑰ │ ⑱ │ ⑲ │ ⑳ │ ㉑ │ ㉒ │ ㉓ │ ㉔ │ 32行
m=3 │ ㉕ │ ㉖ │ ㉗ │ ㉘ │ ㉙ │ ㉚ │ ㉛ │ ㉜ │ 32行
└────┴────┴────┴────┴────┴────┴────┴────┘

6.4 为什么子块是 32×16 而不是 32×32

P Tile 的 N=32 不是让一次 mma.sync 覆盖 32 列。硬件 mma.sync 一次只能算 16×8×16,EU Repeat 分 4 个 warp 后一轮覆盖 32×16,这是硬件极限。

P Tile 说"覆盖 N=32"的实现方式是:在 cute::gemm 内部连续执行两次 mma.sync,共享同一份 A 寄存器

MMA_N = 8 的分解 (_2, _4) 体现了这一点:

1
2
3
4
5
6
7
8
MMA_N = (_2, _4)
↑ ↑
EU Repeat N方向 ×2
tile 128 / EU后覆盖32 = 4

8 次 N 方向迭代分成 4 组,每组 2 次:
n=0,1 共享 A │ n=2,3 共享 A │ n=4,5 共享 A │ n=6,7 共享 A
├── 32列 ──┤ ├── 32列 ──┤ ├── 32列 ──┤ ├── 32列 ──┤

7. Fragment 大小总结

1
2
3
4
5
tCrA: ((_2,_2,_2), _4, _1) = 32 个寄存器(A 的缓冲区)
tCrB: ((_2,_2), (_2,_4), _1) = 32 个寄存器(B 的缓冲区)
tCrD: ((_2,_2), _4, _8) = 128 个寄存器(累加器)

验证:每个线程 128 个 D 值 = 128×128 / 128 线程 = 128 ✓

线程 0 每轮读 32 个 A + 32 个 B,经过 4 轮 K 循环后,累加器 tCrD 持有 128 个输出值。


8. cuBLAS 对比

代码中加入了 cuBLAS 对比,包括正确性验证和性能计时。

8.1 cuBLAS 参数设置

D = A @ B^T,A(M,K) row-major,B(N,K) row-major,D(M,N) row-major。

cuBLAS 要求 col-major。row-major (M,K) stride=K 等价于 col-major (K,M) stride=K。所以:

1
2
3
4
5
6
7
D^T(N,M) = B(N,K) @ A^T(K,M)

cublasHgemm(handle,
CUBLAS_OP_T, // B 转置
CUBLAS_OP_N, // A 不转置(因为 row-major 已经等价于 A^T 的 col-major)
N, M, K,
&alpha, dB, K, dA, K, &beta, dD, N)

8.2 性能结果

07 是教学级 kernel(直接从 global memory 读寄存器,没有 smem 缓存,没有流水线),性能远低于 cuBLAS。这是正常的——目的是理解 partition 和 gemm 的机制,不是追求性能。


9. API 总结

API 作用
thr_mma.partition_A(gA) 把 A 的 global tensor 切分到线程级别
thr_mma.partition_B(gB) 把 B 的 global tensor 切分到线程级别
thr_mma.partition_C(gD) 把 D 的 global tensor 切分到线程级别
thr_mma.partition_fragment_A(tile) 创建 A 的寄存器 fragment
thr_mma.partition_fragment_B(tile) 创建 B 的寄存器 fragment
thr_mma.partition_fragment_C(tile) 创建 D 的寄存器 fragment(累加器)
cute::gemm(tiled_mma, D, A, B, D) 执行 MMA:遍历所有子块,调用 mma.sync
cute::copy(src, dst) global→register 或 register→global

10. 与 06/08 的对比

06(TiledMMA 理解) 07(Simple GEMM) 08(Smem GEMM)
计算 单 tile 演示 完整 GEMM 完整 GEMM
数据搬运 Global→Reg Global→Smem→Reg
K 循环
性能 N/A 很慢(global 延迟大) 快得多(smem 缓存)
学习目标 理解 MMA 分配 理解 partition + gemm 理解 G2S + S2R + Swizzle
  • 标题: cute(7)simple_gemm
  • 作者: 鱿鱼圈
  • 创建于 : 2026-06-14 22:13:32
  • 更新于 : 2026-06-14 17:55:16
  • 链接: https://yuyanqi.com/2026/06/14/cute(7)simple_gemm/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论