cute(9)multistage_gemm

鱿鱼圈 Lv4

前置阅读

9 多阶段流水线 GEMM

对应代码:09_multistage_gemm.cu 需要 GPU

核心概念

在 08 的基础上引入多 stage 环形缓冲 + 流水线重叠

  • 多 stage shared memory:同时保存 kStage 份 K tile,形成环形缓冲
  • Prologue(预填充):kernel 开始前先填满 kStage-1 个 stage
  • Mainloop(主循环):G2S、S2R、MMA 三级流水线交错执行
  • Epilogue(收尾):累加器写回 Global Memory
  • cp_async_fence / cp_async_wait:控制异步搬运批次的同步点

代码实现

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
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
template <typename Config>
__global__ void multistage_gemm_kernel(half_t* Dptr, const half_t* Aptr,
const half_t* Bptr, int M, int N, int K) {
using T = half_t;
using SmemLayoutA = typename Config::SmemLayoutA;
using SmemLayoutB = typename Config::SmemLayoutB;
using TiledMMA = typename Config::MMA;
using S2RCopyAtomA = typename Config::S2RCopyAtomA;
using S2RCopyAtomB = typename Config::S2RCopyAtomB;
using G2SCopyA = typename Config::G2SCopyA;
using G2SCopyB = typename Config::G2SCopyB;

constexpr int kTileM = Config::kTileM;
constexpr int kTileN = Config::kTileN;
constexpr int kTileK = Config::kTileK;
constexpr int kStage = Config::kStage;

extern __shared__ T shm_data[];
T* Ashm = shm_data;
T* Bshm = shm_data + cosize(SmemLayoutA{});

int idx = threadIdx.x;
int bx = blockIdx.x;
int by = blockIdx.y;

// 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>{}));

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

// Shared Memory(多 stage)
auto sA = make_tensor(make_smem_ptr(Ashm), SmemLayoutA{}); // (kTileM, kTileK, kStage)
auto sB = make_tensor(make_smem_ptr(Bshm), SmemLayoutB{}); // (kTileN, kTileK, kStage)

// MMA
TiledMMA tiled_mma;
auto thr_mma = tiled_mma.get_slice(idx);
auto tCrA = thr_mma.partition_fragment_A(gA(_, _, 0));
auto tCrB = thr_mma.partition_fragment_B(gB(_, _, 0));
auto tCrD = thr_mma.partition_fragment_C(gD);
clear(tCrD);

// S2R Copy
auto s2r_tiled_copy_a = make_tiled_copy_A(S2RCopyAtomA{}, tiled_mma);
auto s2r_thr_copy_a = s2r_tiled_copy_a.get_slice(idx);
auto tAsA = s2r_thr_copy_a.partition_S(sA);
auto tCrA_view = s2r_thr_copy_a.retile_D(tCrA);

auto s2r_tiled_copy_b = make_tiled_copy_B(S2RCopyAtomB{}, tiled_mma);
auto s2r_thr_copy_b = s2r_tiled_copy_b.get_slice(idx);
auto tBsB = s2r_thr_copy_b.partition_S(sB);
auto tCrB_view = s2r_thr_copy_b.retile_D(tCrB);

// G2S Copy
G2SCopyA g2s_tiled_copy_a;
auto g2s_thr_copy_a = g2s_tiled_copy_a.get_slice(idx);
auto tAgA_copy = g2s_thr_copy_a.partition_S(gA);
auto tAsA_copy = g2s_thr_copy_a.partition_D(sA);

G2SCopyB g2s_tiled_copy_b;
auto g2s_thr_copy_b = g2s_tiled_copy_b.get_slice(idx);
auto tBgB_copy = g2s_thr_copy_b.partition_S(gB);
auto tBsB_copy = g2s_thr_copy_b.partition_D(sB);

// ============================================================
// 流水线状态
// ============================================================
int itile_to_read = 0;
int ismem_read = 0;
int ismem_write = 0;

// ============================================================
// Prologue:预填充 kStage-1 个 stage
// ============================================================
#pragma unroll
for (int istage = 0; istage < kStage - 1; ++istage) {
cute::copy(g2s_tiled_copy_a, tAgA_copy(_, _, _, istage),
tAsA_copy(_, _, _, istage));
cute::copy(g2s_tiled_copy_b, tBgB_copy(_, _, _, istage),
tBsB_copy(_, _, _, istage));
cp_async_fence();
++itile_to_read;
++ismem_write;
}

// 等待第一个 stage 就绪
cp_async_wait<kStage - 2>();
__syncthreads();

// 首次 S->R 预取
int ik = 0;
cute::copy(s2r_tiled_copy_a, tAsA(_, _, ik, ismem_read), tCrA_view(_, _, ik));
cute::copy(s2r_tiled_copy_b, tBsB(_, _, ik, ismem_read), tCrB_view(_, _, ik));

// ============================================================
// Mainloop:K 方向主循环
// ============================================================
int ntile = K / kTileK;
#pragma unroll 1
for (int itile = 0; itile < ntile; ++itile) {
int nk = size<2>(tCrA);

#pragma unroll
for (int ik = 0; ik < nk; ++ik) {
int ik_next = (ik + 1) % nk;

// 当前 tile 最后一个 ik:等待下一个 stage 就绪,切换读缓冲
if (ik == nk - 1) {
cp_async_wait<kStage - 2>();
__syncthreads();
ismem_read = (ismem_read + 1) % kStage;
}

// 预取下一步的 S->R
cute::copy(s2r_tiled_copy_a, tAsA(_, _, ik_next, ismem_read),
tCrA_view(_, _, ik_next));
cute::copy(s2r_tiled_copy_b, tBsB(_, _, ik_next, ismem_read),
tCrB_view(_, _, ik_next));

// 第一个 ik:发起下一个 tile 的 G->S
if (ik == 0) {
if (itile_to_read < ntile) {
cute::copy(g2s_tiled_copy_a, tAgA_copy(_, _, _, itile_to_read),
tAsA_copy(_, _, _, ismem_write));
cute::copy(g2s_tiled_copy_b, tBgB_copy(_, _, _, itile_to_read),
tBsB_copy(_, _, _, ismem_write));
++itile_to_read;
ismem_write = (ismem_write + 1) % kStage;
}
cp_async_fence();
}

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

// ============================================================
// Epilogue:写回 global(简化版)
// ============================================================
auto tDgD = thr_mma.partition_C(gD);
cute::copy(tCrD, tDgD);
}

struct MultiStageConfig {
using T = half_t;
static constexpr int kTileM = 128;
static constexpr int kTileN = 128;
static constexpr int kTileK = 32;
static constexpr int kStage = 3;

// Swizzle SmemLayout
using SmemLayoutAtom = decltype(composition(
Swizzle<3, 3, 3>{},
make_layout(make_shape(Int<8>{}, Int<kTileK>{}),
make_stride(Int<kTileK>{}, Int<1>{}))));
using SmemLayoutA = decltype(tile_to_shape(
SmemLayoutAtom{}, make_shape(Int<kTileM>{}, Int<kTileK>{}, Int<kStage>{})));
using SmemLayoutB = decltype(tile_to_shape(
SmemLayoutAtom{}, make_shape(Int<kTileN>{}, Int<kTileK>{}, Int<kStage>{})));

// MMA
using mma_op = SM80_16x8x16_F16F16F16F16_TN;
using mma_atom = MMA_Atom<MMA_Traits<mma_op>>;
using MMA = decltype(make_tiled_mma(
mma_atom{},
make_layout(make_shape(Int<2>{}, Int<2>{}, Int<1>{})),
Tile<Int<32>, Int<32>, Int<16>>{}));

// G2S Copy: cp.async 128bit
using g2s_copy_op = SM80_CP_ASYNC_CACHEGLOBAL<cute::uint128_t>;
using g2s_copy_atom = Copy_Atom<Copy_Traits<g2s_copy_op>, T>;
using G2SCopyA = decltype(make_tiled_copy(
g2s_copy_atom{},
make_layout(make_shape(Int<32>{}, Int<4>{}), make_stride(Int<4>{}, Int<1>{})),
make_layout(make_shape(Int<1>{}, Int<8>{}))));
using G2SCopyB = G2SCopyA;

// S2R Copy: ldmatrix
using s2r_copy_atom = Copy_Atom<Copy_Traits<SM75_U32x4_LDSM_N>, T>;
using S2RCopyAtomA = s2r_copy_atom;
using S2RCopyAtomB = s2r_copy_atom;

static constexpr int kThreadNum = size(MMA{});
static constexpr int kShmSize =
(cosize(SmemLayoutA{}) + cosize(SmemLayoutB{})) * sizeof(T);
};

1. 08 vs 09:从"等着用"到"提前搬"

1.1 08 的问题

08 的 K 循环是串行的:

1
2
3
4
5
6
08 时间线 (每个 K tile):
┌──────────┐ ┌──────┐ ┌──────┐ ┌──────┐ ┌──────┐ ┌──────┐
│ G→S 搬运 │ │ wait │ │ S→R │ │ MMA │ │ S→R │ │ MMA │
└──────────┘ └──────┘ └──────┘ └──────┘ └──────┘ └──────┘
↑ ↑
搬运中 GPU 闲着 等 smem 就绪才能算

问题:G2S 搬运的时候,MMA 完全空闲;MMA 计算的时候,G2S 也没在搬下一个。

1.2 09 的改进

1
2
3
4
5
6
09 时间线 (流水线重叠):
G→S: [tile0] [tile1] [tile2] [tile3] [tile4] ...
S→R: [tile0] [tile1] [tile2] [tile3] ...
MMA: [tile0] [tile1] [tile2] [tile3] ...

G→S 搬 tile1 的同时,MMA 在算 tile0 → 重叠!

关键:用多份 smem(kStage=3 份),让 G2S 写新 stage 的同时,MMA 读老 stage,互不冲突。


2. Config 对比 08

2.1 增加了 kStage

1
2
3
4
static constexpr int kTileM = 128;   // 同 08
static constexpr int kTileN = 128; // 同 08
static constexpr int kTileK = 32; // 同 08
static constexpr int kStage = 3; // ← 新增!3 个 stage

2.2 SmemLayout 多了一维

1
2
3
4
5
6
7
8
9
// 08:
using SmemLayoutA = decltype(tile_to_shape(
SmemLayoutAtom{}, make_shape(Int<kTileM>{}, Int<kTileK>{})));
// (128, 32) → 二维

// 09:
using SmemLayoutA = decltype(tile_to_shape(
SmemLayoutAtom{}, make_shape(Int<kTileM>{}, Int<kTileK>{}, Int<kStage>{})));
// (128, 32, 3) → 三维!

第三维就是 stage 维度。smem 大小变成 3 倍:

1
2
08: (128 × 32) × 2 sides × 2 bytes = 16 KB
09: (128 × 32 × 3) × 2 sides × 2 bytes = 48 KB = 49152 bytes

2.3 SmemLayoutAtom 中的参数

1
2
3
4
using SmemLayoutAtom = decltype(composition(
Swizzle<3, 3, 3>{},
make_layout(make_shape (Int<8>{}, Int<kTileK>{}),
make_stride(Int<kTileK>{}, Int<1>{}))));

这是一个 8 行 × 32 列 的行优先矩阵布局,配合 Swizzle:

参数 含义
8 行数 Swizzle<3,3,3> 的 B=3 → 2^3=8 行一个 XOR 周期
kTileK=32 列数 K 方向一个 tile 的宽度
Int<kTileK>{} 行 stride 行优先:相邻行间隔 32 元素
Int<1>{} 列 stride 行优先:同行相邻列间隔 1 元素

为什么恰好是 8 行——三个原因汇聚:

  1. Swizzle 周期 = 8:B=3 → 2^3=8 行后 XOR 模式重复
  2. ldmatrix 一次读 8 行ldmatrix.x4 由 8 个线程各提供一个地址,读 8 行
  3. 128 bit = 8 个 half_t:M=3 → 2^3=8 个 half_t 为一组,恰好是 ldmatrix 一次取的宽度

tile_to_shape 把这个 8×32 的 atom 沿 M 方向重复 128/8=16 次,沿 stage 方向重复 3 次,铺满 (128, 32, 3)。

2.4 其他部分

MMA、G2S Copy、S2R Copy 与 08 完全一样,详见 08 文档 的 2.3~2.5 节。


3. Kernel 结构总览

1
2
3
4
5
6
7
8
multistage_gemm_kernel:

1. 创建 Global/Smem/Reg tensor
2. 创建 G2S / S2R / MMA 的 partition
3. 流水线状态变量初始化
4. Prologue: 预填充 kStage-1 个 stage
5. Mainloop: K 方向主循环 (G2S + S2R + MMA 交错)
6. Epilogue: 累加器 → Global

4. 流水线状态变量

4.1 三个关键变量

1
2
3
int itile_to_read = 0;   // 下一个要从 global 读的 K tile 编号
int ismem_read = 0; // 当前从哪个 stage 读 (S→R 消费)
int ismem_write = 0; // 下一个要写入哪个 stage (G→S 生产)

它们各自追踪不同的东西:

1
2
3
4
5
6
7
8
9
Global K tiles:   [tile0] [tile1] [tile2] [tile3] [tile4] ...

itile_to_read = 4
"还没搬的从这里开始"

Smem stages: [stage0] [stage1] [stage2]
↑ ↑
ismem_read ismem_write
"MMA正在用" "G2S下次写这里"

4.2 环形缓冲

三个 stage 形成环形缓冲,用取模实现:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
时间线:
stage0 stage1 stage2
Prologue:
istage=0: 写入 ← ismem_write: 0→1
istage=1: 写入 ← ismem_write: 1→2

Mainloop itile=0:
读(S→R+MMA): 读取 ← ismem_read: 0
写(G→S): 写入 ← ismem_write: 2→0 (环绕)
ik==nk-1: ismem_read: 0→1

Mainloop itile=1:
读(S→R+MMA): 读取 ← ismem_read: 1
写(G→S): 写入 ← ismem_write: 0→1
ik==nk-1: ismem_read: 1→2

Mainloop itile=2:
读(S→R+MMA): 读取 ← ismem_read: 2
写(G→S): 写入 ← ismem_write: 1→2
ik==nk-1: ismem_read: 2→0
...

4.3 与 tAsA 坐标的关系

S2R 的 smem tensor tAsA 有四个维度:

1
2
3
tAsA(_, _, ik, ismem_read)
// ↑ ↑ ↑ ↑
// val M K内层 stage编号
坐标 含义 取值范围
第 0 维 _ 每线程搬的 8 个值 自动
第 1 维 _ M 方向分块 自动
ik / ik_next K 方向内层循环(32/16=2 轮) 0 或 1
ismem_read 从哪个 stage 读 0, 1, 2 (环形)

G2S 的 global tensor 和 smem tensor:

1
2
tAgA_copy(_, _, _, itile_to_read)   // global 侧:第 4 维是全局 K tile 编号
tAsA_copy(_, _, _, ismem_write) // smem 侧:第 4 维是 stage 编号
坐标 含义
itile_to_read 全局第几个 K tile(0, 1, 2, … ntile-1)
ismem_write 写入 smem 的哪个 stage(0, 1, 2 环形)

5. Prologue(预填充)

1
2
3
4
5
6
7
8
9
10
11
12
13
#pragma unroll
for (int istage = 0; istage < kStage - 1; ++istage) {
cute::copy(g2s_tiled_copy_a, tAgA_copy(_, _, _, istage),
tAsA_copy(_, _, _, istage));
cute::copy(g2s_tiled_copy_b, tBgB_copy(_, _, _, istage),
tBsB_copy(_, _, _, istage));
cp_async_fence();
++itile_to_read;
++ismem_write;
}

cp_async_wait<kStage - 2>(); // 等待第一个 stage 就绪
__syncthreads();

5.1 为什么填 kStage-1 个

1
2
3
4
5
6
7
kStage = 3 → 填 2 个 stage

stage0 stage1 stage2
Prologue后: [tile0] [tile1] [空]
↑ ↑
ismem_read=0 ismem_write=2
"准备消费" "下次 G2S 写这里"

只填 kStage-1=2 个,留一个空 stage 给 mainloop 里的 G2S 写。如果填满 3 个:

  • mainloop 第一次 G2S 要写 stage0,但 stage0 还没被 MMA 消费完
  • 新数据会覆盖正在用的老数据 → 计算错误

5.2 cp_async_fence 的作用

每次 cp_async_fence() 标记一个"批次边界":

1
2
3
4
copy tile0 的所有 cp.async 指令
cp_async_fence() ← 批次 0 到此结束
copy tile1 的所有 cp.async 指令
cp_async_fence() ← 批次 1 到此结束

后续 cp_async_wait<N>() 按批次等待,不按指令等。

5.3 cp_async_wait<kStage - 2>

1
2
3
cp_async_wait<kStage - 2>()  即  cp_async_wait<1>()

含义:等到最多还有 1 个批次在飞行中

Prologue 发了 2 个批次(tile0 和 tile1)。wait<1> 确保至少 tile0 完成了(tile1 可能还在搬)。

为什么不用 wait<0>:wait<0> 会等所有批次都完成,包括 tile1。但我们只需要 tile0 就绪就可以开始算了,让 tile1 继续异步搬,实现重叠。

5.4 首次 S→R 预取

1
2
3
int ik = 0;
cute::copy(s2r_tiled_copy_a, tAsA(_, _, ik, ismem_read), tCrA_view(_, _, ik));
cute::copy(s2r_tiled_copy_b, tBsB(_, _, ik, ismem_read), tCrB_view(_, _, ik));

在进入 mainloop 之前,先把 stage0 的第一个 K 片段(ik=0)搬进寄存器。这样 mainloop 一开始就能直接算 MMA,不用等 S→R。


6. Mainloop(主循环)详解

6.1 整体结构

1
2
3
4
5
6
7
8
9
10
int ntile = K / kTileK;           // K=256, kTileK=32 → 8 轮
#pragma unroll 1
for (int itile = 0; itile < ntile; ++itile) {
int nk = size<2>(tCrA); // kTileK/MMA_K = 32/16 = 2

#pragma unroll
for (int ik = 0; ik < nk; ++ik) {
// ... 见下方分解
}
}

外层循环遍历全局 K tile,内层循环处理一个 smem stage 内的 K 片段。

6.2 内层循环逐段分析

1
2
for (int ik = 0; ik < nk; ++ik) {
int ik_next = (ik + 1) % nk;

ik_next下一步 S→R 预取的 K 片段索引。当 ik=nk-1(最后一个片段)时,ik_next=0,指向下一个 stage 的第 0 个 K 片段。

第一段:等待下一 stage 就绪(仅最后一个 ik)

1
2
3
4
5
if (ik == nk - 1) {
cp_async_wait<kStage - 2>(); // wait<1>
__syncthreads();
ismem_read = (ismem_read + 1) % kStage;
}

为什么 ismem_read 只在最后一个 ik 更新?

一个 stage 包含 kTileK=32 列的数据,分 nk=2 个 K 片段。必须把这个 stage 的所有片段都消费完,才能切换到下一个 stage:

1
2
stage0: [K片段0 (col 0~15)] [K片段1 (col 16~31)]
ik=0 ik=1 ← 用完了,才切 stage

如果 ik=0 就切:K 片段 1 还没用,直接跳到下一个 stage 了 → 漏算。

wait<kStage-2> 即 wait<1>:确保至少一个旧批次完成。因为 pipeline 中最多有 kStage-1=2 个批次在飞行:

  • 当前正在消费的 stage(已完成)
  • 刚发出的 G2S(可能还在飞)
  • wait<1> 确保下一个要消费的 stage 已经到达

第二段:S→R 预取

1
2
3
4
cute::copy(s2r_tiled_copy_a, tAsA(_, _, ik_next, ismem_read),
tCrA_view(_, _, ik_next));
cute::copy(s2r_tiled_copy_b, tBsB(_, _, ik_next, ismem_read),
tCrB_view(_, _, ik_next));

提前把下一步要用的数据搬进寄存器。这样当前的 MMA 和下一步的 S→R 可以重叠执行。

注意:当 ik=nk-1 时,ismem_read 已经更新了,ik_next=0,所以这里预取的是新 stage 的第 0 个 K 片段

第三段:G2S 发起新搬运(仅第一个 ik)

1
2
3
4
5
6
7
8
9
10
11
if (ik == 0) {
if (itile_to_read < ntile) {
cute::copy(g2s_tiled_copy_a, tAgA_copy(_, _, _, itile_to_read),
tAsA_copy(_, _, _, ismem_write));
cute::copy(g2s_tiled_copy_b, tBgB_copy(_, _, _, itile_to_read),
tBsB_copy(_, _, _, ismem_write));
++itile_to_read;
ismem_write = (ismem_write + 1) % kStage;
}
cp_async_fence();
}

为什么 G2S 放在 ik==0

  • 一个 stage 有 nk=2 个 ik 循环。G2S 只需发起一次(搬整个 32 列 tile)
  • 放在 ik=0 最早发起,给 cp.async 最多的时间异步搬运
  • 等到 ik=nk-1 时(消费完当前 stage),新 stage 大概率已经搬完了

itile_to_read < ntile 的检查:K 方向的 tile 总共 ntile 个,Prologue 已搬了 kStage-1 个,mainloop 每轮再搬 1 个,到后期所有 tile 都搬完了,不再发起 G2S。

cp_async_fence() 无条件执行:即使没有实际搬运(itile_to_read >= ntile),也要 fence,确保 wait 的批次计数正确。

第四段:MMA 计算

1
cute::gemm(tiled_mma, tCrD, tCrA(_, _, ik), tCrB(_, _, ik), tCrD);

用当前 ik 的寄存器数据做 MMA。注意这里用 tCrA(MMA 视角),不是 tCrA_view(Copy 视角)。

6.3 完整时间线(K=256, kTileK=32, kStage=3)

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
ntile = 256/32 = 8 个 K tile
nk = 32/16 = 2 个 K 片段/stage

Prologue:
G2S tile0 → stage0, fence
G2S tile1 → stage1, fence
wait<1> (tile0 就绪)
S2R stage0[ik=0] → reg

Mainloop itile=0 (消费 stage0):
ik=0:
S2R stage0[ik=1] → reg ← 预取下一个 K 片段
G2S tile2 → stage2, fence ← 发起搬运
MMA(ik=0) ← 算当前 K 片段
ik=1 (=nk-1):
wait<1>, sync ← 等 stage1 就绪
ismem_read: 0→1 ← 切到 stage1
S2R stage1[ik=0] → reg ← 预取新 stage 首个 K 片段
MMA(ik=1) ← 算 stage0 最后一个 K 片段

Mainloop itile=1 (消费 stage1):
ik=0:
S2R stage1[ik=1] → reg
G2S tile3 → stage0, fence ← stage0 已被消费完,可以覆盖
MMA(ik=0)
ik=1:
wait<1>, sync
ismem_read: 1→2
S2R stage2[ik=0] → reg
MMA(ik=1)

... 以此类推 ...

Mainloop itile=7 (最后一轮,消费最后一个 stage):
ik=0:
S2R 预取
itile_to_read=10 > ntile=8, 不搬了, 但仍 fence
MMA(ik=0)
ik=1:
wait<1>, sync
ismem_read 更新
S2R 预取 (下一轮不会用到)
MMA(ik=1)

6.4 流水线重叠效果

1
2
3
4
5
6
7
8
9
10
不重叠(08 风格):
[G→S 搬运] [等待] [S→R] [MMA] [S→R] [MMA] [G→S 搬运] [等待] ...
────────── ──── ──── ──── ──── ──── ────────── ────

重叠(09 流水线):
G→S: [tile0][tile1] [tile2] [tile3] [tile4] ...
S→R: [ ][ ] [ ][ ] [ ][ ] [ ][ ] ...
MMA: [ ][ ] [ ][ ] [ ][ ] [ ][ ] ...

G2S 和 MMA 同时进行!

G2S 的 cp.async 是真正异步的:发完指令后 GPU 就能去做别的(S2R + MMA),数据搬运由硬件后台完成。


7. Epilogue(收尾)

1
2
auto tDgD = thr_mma.partition_C(gD);
cute::copy(tCrD, tDgD);

7.1 partition_C vs partition_fragment_C

API 返回 存储位置
partition_fragment_C(gD) tCrD — 累加器片段 寄存器
partition_C(gD) tDgD — 对应的 global 位置 Global Memory

两者的 shape 完全一样,元素一一对应:

1
2
tCrD[i]  →  tDgD[i]
寄存器 Global Memory 对应位置

cute::copy(tCrD, tDgD) 就是把每个线程的累加器结果逐元素写回 global memory。

7.2 为什么叫"简化版"

当前的 Epilogue 是 Reg → Global 直写。问题在于:

1
2
3
4
5
6
MMA 的寄存器布局是为计算优化的,不是为内存访问优化的。

线程 0 写: D[0][0], D[0][1], D[8][0], D[8][1], ... ← 不连续!
线程 1 写: D[1][0], D[1][1], D[9][0], D[9][1], ... ← 不连续!

这些写操作不是 coalesced(合并访问)的,性能不好。

高性能的 Epilogue 会走 Reg → Shared → Global

1
2
Reg → Shared:  按 MMA 布局写入 smem
Shared → Global: 按连续顺序读出 smem 写入 global (coalesced)

用 smem 做中转,重排数据,保证 global 写回是连续合并的。


8. cp_async_wait<N> 详解

8.1 语义

1
2
3
cp_async_wait<N>()

等待直到"仍在飞行中的 cp.async 批次数 ≤ N"

每次 cp_async_fence() 把之前的 cp.async 指令打包成一个"批次"。wait<N> 按批次等待。

8.2 为什么用 wait<kStage-2> 而不是 wait<kStage-1>

1
2
3
4
5
kStage = 3

wait<kStage-2> = wait<1>: 允许最多 1 批在飞 → 确保倒数第 2 批完成
wait<kStage-1> = wait<2>: 允许最多 2 批在飞 → 基本不等(可能没完成就继续了)
wait<0>: 等所有批完成 → 太保守,失去流水线优势

wait<1> 的逻辑

pipeline 中最多同时存在 kStage-1=2 个 in-flight 批次。当要消费一个 stage 时:

  • 允许 1 个批次还在飞
  • 至少 1 个批次已完成 → 就是我们要消费的那个

如果用 wait<2>:可能要消费的 stage 还没搬完 → 数据不对。 如果用 wait<0>:所有批都等完了 → 不需要 3 个 stage 了 → 退化成 08。

8.3 能不能改成 kStage-1?

1
2
3
4
5
6
7
cp_async_wait<kStage - 1>  即  wait<2>

允许 2 个批次在飞行中。
Prologue 发了 2 个批次,wait<2> 基本不等。
→ 可能 stage0 都没搬完就开始 S→R 了
→ 读到垃圾数据
→ 计算错误!

不能改。 wait 的 N 必须保证要读的 stage 已经完成。


9. ismem_read 的更新时机

9.1 为什么不在 ik=0 就更新?

一个 smem stage 包含 kTileK=32 列,分 nk=2 个 K 片段消费:

1
2
3
4
5
6
7
8
9
10
stage0 (32 列):
[K 片段 0: col 0~15] [K 片段 1: col 16~31]
ik=0 ik=1

如果 ik=0 就更新 ismem_read:
ik=0: MMA 用 stage0[ik=0] ✓
然后 ismem_read → 1,切到 stage1
ik=1: 预取用的是 stage1[ik=1] ← 错了!应该是 stage0[ik=1]

stage0 的后半段没用就被跳过了 → 漏算 → 结果错误

9.2 正确的更新逻辑

1
2
3
4
5
6
7
8
9
10
ik=0:
用 tCrA(ik=0) 做 MMA ← 用的是 Prologue 或上一轮末尾预取的
预取 tCrA(ik=1) ← 从当前 ismem_read 的 stage 读
不更新 ismem_read ← 当前 stage 还没消费完

ik=1 (=nk-1):
wait + sync ← 等下一个 stage 就绪
更新 ismem_read ← 当前 stage 消费完了,切换
预取 tCrA(ik_next=0) ← 从新的 ismem_read 读,准备下一轮
用 tCrA(ik=1) 做 MMA ← 用的是上一步预取的

一句话:消费完当前 stage 的所有 K 片段后,才切到下一个 stage。


10. #pragma unroll 策略

1
2
3
4
#pragma unroll 1
for (int itile = 0; itile < ntile; ++itile) { // 外层:不展开
#pragma unroll
for (int ik = 0; ik < nk; ++ik) { // 内层:完全展开
循环 策略 原因
外层 itile #pragma unroll 1(不展开) ntile 可能很大(如 8),展开会膨胀指令缓存
内层 ik #pragma unroll(完全展开) nk=2,很小,展开后 compiler 能更好调度指令

11. 性能测试

11.1 测试配置

1
2
3
4
5
6
7
8
M = 81920, N = 256, K = 256
A(M,K) 行优先, B(N,K) 行优先, D(M,N) 行优先
D = A @ B^T

Tile: 128×128×32, Stage: 3
Grid: (N/128, M/128) = (2, 640) = 1280 blocks
Threads: 128/block
ShmSize: 49152 bytes = 48 KB

11.2 结果

1
2
3
Custom kernel:  0.1923 ms  (55.83 TFLOPS)
cuBLAS: 0.0851 ms (126.20 TFLOPS)
Custom/cuBLAS: 44.24%

11.3 分析

  • 正确性:PASS,0 errors
  • Custom 达到 cuBLAS 的 44%
  • 这个 tall-skinny 形状(M=81920, N=256, K=256)N 方向只有 2 个 block,wave 利用率低
  • cuBLAS 有专门的优化路径(如 splitK、自适应 tile 大小)
  • 当前 Epilogue 是简化版(Reg→Global 直写),非 coalesced 写回影响性能
  • 没有做 thread block 调度优化(如 swizzle block 顺序提高 L2 命中率)

12. 完整数据流图

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
     ┌──────────────────────────────┐
│ Global Memory │
│ A(M,K) B(N,K) D(M,N) │
└──┬───────┬──────────┬────────┘
│cp.async │
┌───────▼───────▼──────┐ │
│ Shared Memory │ │
│ sA(128,32,3) │ │
│ sB(128,32,3) │ │
│ ┌──────┬──────┬────┐ │ │
│ │stg 0 │stg 1 │stg2│ │ │
│ └──┬───┴──────┴────┘ │ │
└─────┼─────────────────┘ │
│ldmatrix │
┌─────▼─────────────────┐ │
│ Registers │ │
│ tCrA: MMA 视角 │ │
│ tCrA_view: Copy 视角 │ │
│ tCrB / tCrB_view │ │
│ tCrD: 累加器 │ │
└─────┬─────────────┬───┘ │
│mma.sync │ │
┌─────▼─────┐ │ │
│ tCrD 累加 │ │ │
│ (FMA) │ │ │
└─────┬─────┘ │ │
│ │ │
└─────────────┘ │
cute::copy │
┌────────────────────▼─┐
│ D (Global) │
└──────────────────────┘

13. 与 07/08 的完整对比

07(Simple GEMM) 08(Smem GEMM) 09(Multi-Stage)
数据搬运 Global→Reg Global→Smem→Reg Global→Smem(×3)→Reg
Shared Memory 有,单份 有,3 份环形缓冲
Swizzle 不需要
K 循环 1 层 2 层(外=G2S,内=S2R+MMA) 2 层 + 流水线交错
cp.async 不用 用,同步等待 用,异步重叠
G2S 与 MMA 串行 并行
S2R 预取 没有 有(提前搬下一个 ik)
Epilogue Reg→Global Reg→Global Reg→Global(简化版)
Prologue 有(预填充 kStage-1 个 stage)
smem 大小 0 16 KB 48 KB
make_tiled_copy_A 不需要 需要 需要
retile_D 不需要 需要 需要
核心思想 partition + gemm 两级搬运 + Swizzle 流水线重叠隐藏延迟

14. API 总结

API 作用 09 中的使用
cp_async_fence() 标记 cp.async 批次边界 Prologue 每 stage 一次,Mainloop 每 itile 一次
cp_async_wait<N>() 等到至多 N 批在飞行 wait<kStage-2> = wait<1>
__syncthreads() block 内线程同步 wait 之后必须 sync
tile_to_shape(atom, shape) 扩展 layout atom 到目标 shape 加了第 3 维 kStage
make_tiled_copy(...) G2S copy,线程布局自选 同 08
make_tiled_copy_A(atom, mma) S2R copy,线程布局匹配 MMA 同 08
retile_D(fragment) MMA 视角 → Copy 视角 同 08
partition_C(gD) 按 MMA C 布局标出 global 位置 Epilogue 写回
partition_fragment_C(gD) 分配寄存器累加器 MMA 累加
  • 标题: cute(9)multistage_gemm
  • 作者: 鱿鱼圈
  • 创建于 : 2026-06-16 22:13:32
  • 更新于 : 2026-06-14 22:17:27
  • 链接: https://yuyanqi.com/2026/06/16/cute(9)multistage_gemm/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论