-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathcommon.h
More file actions
177 lines (152 loc) · 5.4 KB
/
Copy pathcommon.h
File metadata and controls
177 lines (152 loc) · 5.4 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
#pragma once
#include <cstdint>
#include <cudaTypedefs.h>
constexpr int WARP_SIZE = 32;
__host__ __device__ inline
constexpr int cdiv(int a, int b) { return (a + b - 1) / b; }
__device__ inline
uint32_t elect_sync() {
uint32_t pred = 0;
asm volatile(
"{\n\t"
".reg .pred %%px;\n\t"
"elect.sync _|%%px, %1;\n\t"
"@%%px mov.s32 %0, 1;\n\t"
"}"
: "+r"(pred)
: "r"(0xFFFFFFFF)
);
return pred;
}
__device__ inline
void mbarrier_init(int mbar_addr, int count) {
asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
}
__device__ inline
void mbarrier_wait(int mbar_addr, int phase) {
uint32_t ticks = 0x989680; // this is optional
asm volatile(
"{\n\t"
".reg .pred P1;\n\t"
"LAB_WAIT:\n\t"
"mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\n\t"
"@P1 bra.uni DONE;\n\t"
"bra.uni LAB_WAIT;\n\t"
"DONE:\n\t"
"}"
:: "r"(mbar_addr), "r"(phase), "r"(ticks)
);
}
__device__ inline
void mbarrier_arrive_expect_tx(int mbar_addr, int size) {
asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cluster.b64 _, [%0], %1;"
:: "r"(mbar_addr), "r"(size) : "memory");
}
__device__ inline
void mbarrier_arrive(int mbar_addr) {
asm volatile("mbarrier.arrive.release.cta.shared::cluster.b64 _, [%0];" :: "r"(mbar_addr) : "memory");
}
__device__ __forceinline__ void mbarrier_arrive_cluster(int mbar_addr, int target_cta)
{
int remote_addr;
asm volatile("mapa.shared::cluster.u32 %0, %1, %2;"
: "=r"(remote_addr) : "r"(mbar_addr), "r"(target_cta));
asm volatile("mbarrier.arrive.shared::cluster.b64 _, [%0];"
:: "r"(remote_addr) : "memory");
}
__device__ inline
void tma_2d_g2s(int dst, const void *tmap_ptr, int x, int y, int mbar_addr) {
asm volatile("cp.async.bulk.tensor.2d.shared::cta.global.mbarrier::complete_tx::bytes [%0], [%1, {%2, %3}], [%4];"
:: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(mbar_addr) : "memory");
}
template <int CTA_GROUP = 1>
__device__ inline
void tma_2d_g2s_cluster(int dst, const void *tmap_ptr, int x, int y, int mbar_addr) {
asm volatile("cp.async.bulk.tensor.2d.shared::cluster.global.mbarrier::complete_tx::bytes.cta_group::%5 "
"[%0], [%1, {%2, %3}], [%4];"
:: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(mbar_addr), "n"(CTA_GROUP)
: "memory");
}
template <int CTA_GROUP = 1>
__device__ inline
void tcgen05_alloc(int smem_addr, int size) {
asm volatile("tcgen05.alloc.cta_group::%2.sync.aligned.shared::cta.b32 [%0], %1;"
:: "r"(smem_addr), "r"(size), "n"(CTA_GROUP));
}
template <int CTA_GROUP = 1>
__device__ inline
void tcgen05_dealloc(int taddr, int size) {
asm volatile("tcgen05.dealloc.cta_group::%2.sync.aligned.b32 %0, %1;"
:: "r"(taddr), "r"(size), "n"(CTA_GROUP));
}
template <int CTA_GROUP = 1>
__device__ inline
void tcgen05_mma_f16(int taddr, uint64_t a_desc, uint64_t b_desc, uint32_t i_desc, int enable_input_d) {
asm volatile(
"{\n\t"
".reg .pred p;\n\t" // predicate register enable-input-d
"setp.ne.b32 p, %4, 0;\n\t"
"tcgen05.mma.cta_group::%5.kind::f16 [%0], %1, %2, %3, p;\n\t"
"}"
:: "r"(taddr), "l"(a_desc), "l"(b_desc), "r"(i_desc), "r"(enable_input_d), "n"(CTA_GROUP)
);
}
template <int CTA_GROUP = 1>
__device__ inline
void tcgen05_commit(int mbar_addr) {
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.b64 [%0], %1;"
:: "r"(mbar_addr), "n"(CTA_GROUP) : "memory");
}
template <int CTA_GROUP = 1>
__device__ inline
void tcgen05_commit_mcast(int mbar_addr, int16_t cta_mask) {
asm volatile("tcgen05.commit.cta_group::%2.mbarrier::arrive::one.shared::cluster.multicast::cluster.b64 [%0], %1;"
:: "r"(mbar_addr), "h"(cta_mask), "n"(CTA_GROUP) : "memory");
}
__device__ inline
constexpr uint64_t desc_encode(uint64_t x) { return (x & 0x3'FFFFULL) >> 4ULL; };
inline
void init_tmap_2d_simple(
CUtensorMap *tmap,
const nv_bfloat16 *ptr,
uint64_t global_height, uint64_t global_width,
uint32_t shared_height, uint32_t shared_width,
CUtensorMapSwizzle swizzle
)
{
constexpr uint32_t rank = 2;
uint64_t globalDim[rank] = {global_width, global_height};
uint64_t globalStrides[rank-1] = {global_width * sizeof(nv_bfloat16)}; // in bytes
uint32_t boxDim[rank] = {shared_width, shared_height};
uint32_t elementStrides[rank] = {1, 1};
auto err = cuTensorMapEncodeTiled(
tmap, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, rank, (void *)ptr,
globalDim, globalStrides, boxDim, elementStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE,
swizzle,
CU_TENSOR_MAP_L2_PROMOTION_NONE,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
);
}
inline
void init_tmap_2d_128B_l2_256B(
CUtensorMap *tmap,
const nv_bfloat16 *ptr,
uint64_t global_height, uint64_t global_width,
uint32_t shared_height, uint32_t shared_width
)
{
constexpr uint32_t rank = 2;
uint64_t globalDim[rank] = {global_width, global_height};
uint64_t globalStrides[rank-1] = {global_width * sizeof(nv_bfloat16)};
uint32_t boxDim[rank] = {shared_width, shared_height};
uint32_t elementStrides[rank] = {1, 1};
auto err = cuTensorMapEncodeTiled(
tmap, CU_TENSOR_MAP_DATA_TYPE_BFLOAT16, rank, (void *)ptr,
globalDim, globalStrides, boxDim, elementStrides,
CU_TENSOR_MAP_INTERLEAVE_NONE,
CU_TENSOR_MAP_SWIZZLE_128B,
CU_TENSOR_MAP_L2_PROMOTION_L2_256B,
CU_TENSOR_MAP_FLOAT_OOB_FILL_NAN_REQUEST_ZERO_FMA
);
}