-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathblackwell_gemm_v2.cu
More file actions
224 lines (158 loc) · 6.29 KB
/
Copy pathblackwell_gemm_v2.cu
File metadata and controls
224 lines (158 loc) · 6.29 KB
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
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
#include <cuda_bf16.h>
#include "common.h"
template <int BLOCK_M, int BLOCK_N, int BLOCK_K, int MMA_K, int NUM_STAGES, int TB_SIZE>
__global__ __launch_bounds__(TB_SIZE) void matmul_v2_kernel(
const __grid_constant__ CUtensorMap A_tmap,
const __grid_constant__ CUtensorMap B_tmap,
nv_bfloat16 *C_ptr, int M, int N, int K)
{
const int tid = threadIdx.x;
const int bid = blockIdx.x;
const int warp_id = tid / WARP_SIZE;
const int grid_n = N / BLOCK_N;
const int bid_m = bid / grid_n;
const int bid_n = bid % grid_n;
const int off_m = bid_m * BLOCK_M;
const int off_n = bid_n * BLOCK_N;
extern __shared__ __align__(1024) char smem_ptr[];
const int smem = static_cast<int>(__cvta_generic_to_shared(smem_ptr));
const int sa_size = BLOCK_M * BLOCK_K * sizeof(nv_bfloat16);
const int sb_size = BLOCK_N * BLOCK_K * sizeof(nv_bfloat16);
#pragma nv_diag_suppress static_var_with_dynamic_init
__shared__ uint64_t tma_mbars[NUM_STAGES];
__shared__ uint64_t mma_mbar[1];
__shared__ int tmem_addr[1];
const int tma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(tma_mbars));
const int mma_mbar_addr = static_cast<int>(__cvta_generic_to_shared(mma_mbar));
if (warp_id == 0 && elect_sync())
{
mbarrier_init(mma_mbar_addr, 1);
for (int i = 0; i < NUM_STAGES; i++)
{
mbarrier_init(tma_mbar_addr + i * 8, 1);
}
asm volatile("fence.mbarrier_init.release.cluster;");
}
else if (warp_id == 1)
{
const int addr = static_cast<int>(__cvta_generic_to_shared(tmem_addr));
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(addr), "r"(BLOCK_N));
}
__syncthreads();
const int taddr = tmem_addr[0];
int tma_phase = 0, mma_phase = 0;
constexpr uint32_t i_desc = (1U << 4U)
| (1U << 7U)
| (1U << 10U)
| ((uint32_t)BLOCK_N >> 3U << 17U)
| ((uint32_t)BLOCK_M >> 4U << 24U);
auto load = [&](int iter_k)
{
if (warp_id == 0 && elect_sync())
{
const int stage_id = iter_k % NUM_STAGES;
const int cur_tmbar_addr = tma_mbar_addr + stage_id * 8;
const int A_smem = smem + stage_id * (sa_size + sb_size);
const int B_smem = A_smem + sa_size;
for (int k = 0; k < BLOCK_K / 64; k++)
{
const int off_k = iter_k * BLOCK_K + k * 64;
tma_2d_g2s(A_smem + k * BLOCK_M * 128, &A_tmap, off_k, off_m, cur_tmbar_addr);
tma_2d_g2s(B_smem + k * BLOCK_N * 128, &B_tmap, off_k, off_n, cur_tmbar_addr);
}
constexpr int cp_size = (BLOCK_M + BLOCK_N) * BLOCK_K * sizeof(nv_bfloat16);
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
:: "r"(cur_tmbar_addr), "r"(cp_size) : "memory");
}
};
auto compute = [&](int iter_k)
{
const int stage_id = iter_k % NUM_STAGES;
const int cur_tmbar_addr = tma_mbar_addr + stage_id * 8;
mbarrier_wait(cur_tmbar_addr, tma_phase);
asm volatile("tcgen05.fence::after_thread_sync;");
const int A_smem = smem + stage_id * (sa_size + sb_size);
const int B_smem = A_smem + sa_size;
if (stage_id == NUM_STAGES - 1) tma_phase ^= 1;
if (warp_id == 0 && elect_sync())
{
auto make_desc = [](int addr) -> uint64_t
{
const int SBO = 8 * 128;
return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
};
{
tcgen05_mma_f16(taddr, make_desc(A_smem), make_desc(B_smem), i_desc, iter_k);
for (int k2 = 1; k2 < 64 / MMA_K; k2++)
{
uint64_t a_desc = make_desc(A_smem + k2 * 32);
uint64_t b_desc = make_desc(B_smem + k2 * 32);
tcgen05_mma_f16(taddr, a_desc, b_desc, i_desc, 1);
}
}
for (int k1 = 1; k1 < BLOCK_K / 64; k1++)
for (int k2 = 0; k2 < 64 / MMA_K; k2++)
{
uint64_t a_desc = make_desc(A_smem + k1 * BLOCK_M * 128 + k2 * 32);
uint64_t b_desc = make_desc(B_smem + k1 * BLOCK_N * 128 + k2 * 32);
tcgen05_mma_f16(taddr, a_desc, b_desc, i_desc, 1);
}
asm volatile("tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];"
:: "r"(mma_mbar_addr) : "memory");
}
};
const int num_iters = K / BLOCK_K;
for (int i = 0; i < NUM_STAGES - 1; i ++)
load(i);
for (int iter_k = 0; iter_k < num_iters - NUM_STAGES + 1; iter_k++)
{
load(iter_k + NUM_STAGES - 1);
compute(iter_k);
mbarrier_wait(mma_mbar_addr, mma_phase);
mma_phase ^= 1;
}
for (int iter_k = num_iters - NUM_STAGES + 1; iter_k < num_iters; iter_k++)
{
compute(iter_k);
mbarrier_wait(mma_mbar_addr, mma_phase);
mma_phase ^= 1;
}
asm volatile("tcgen05.fence::after_thread_sync;");
for (int n = 0; n < BLOCK_N / 8; n++)
{
float tmp[8];
const int addr = taddr + ((warp_id * 32) << 16) + (n * 8);
asm volatile("tcgen05.ld.sync.aligned.32x32b.x8.b32 {%0, %1, %2, %3, %4, %5, %6, %7}, [%8];"
: "=f"(tmp[0]), "=f"(tmp[1]), "=f"(tmp[2]), "=f"(tmp[3]),
"=f"(tmp[4]), "=f"(tmp[5]), "=f"(tmp[6]), "=f"(tmp[7])
: "r"(addr));
asm volatile("tcgen05.wait::ld.sync.aligned;");
nv_bfloat162 out[4];
for (int i = 0; i < 4; i++)
out[i] = __float22bfloat162_rn({tmp[i * 2], tmp[i * 2 + 1]});
nv_bfloat16 *out_ptr = C_ptr + (off_m + tid) * N + (off_n + n * 8);
reinterpret_cast<int4 *>(out_ptr)[0] = reinterpret_cast<int4 *>(out)[0];
}
__syncthreads();
if (warp_id == 0)
asm volatile("tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;" :: "r"(taddr), "r"(BLOCK_N));
}
void matmul_v2(const nv_bfloat16 *A_ptr, const nv_bfloat16 *B_ptr, nv_bfloat16 *C_ptr, int M, int N, int K)
{
CUtensorMap A_tmap, B_tmap;
constexpr int BLOCK_M = 128;
constexpr int BLOCK_N = 256, BLOCK_K = 128;
constexpr int NUM_WARPS = 4;
constexpr int TB_SIZE = NUM_WARPS * WARP_SIZE;
constexpr int MMA_K = 16;
constexpr int NUM_STAGES = 2;
init_tmap_2d_simple(&A_tmap, A_ptr, M, K, BLOCK_M, 64, CU_TENSOR_MAP_SWIZZLE_128B);
init_tmap_2d_simple(&B_tmap, B_ptr, N, K, BLOCK_N, 64, CU_TENSOR_MAP_SWIZZLE_128B);
int grid = (M / BLOCK_M) * (N / BLOCK_N);
int size_AB = (BLOCK_M + BLOCK_N) * BLOCK_K * NUM_STAGES;
int smem_size = size_AB * sizeof(nv_bfloat16);
auto this_kernel = matmul_v2_kernel<BLOCK_M, BLOCK_N, BLOCK_K, MMA_K, NUM_STAGES, TB_SIZE>;
if (smem_size > 48 * 1024)
cudaFuncSetAttribute(this_kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size);
this_kernel<<<grid, TB_SIZE, smem_size>>>(A_tmap, B_tmap, C_ptr, M, N, K);
}