Skip to content

Commit 409ab83

Browse files
txs19991tangxinsheng.txs
andauthored
[AMD] support fp8 T.gemm (#804)
* [AMD] support fp8 T.gemm * format --------- Co-authored-by: tangxinsheng.txs <tangxinsheng.txs@alibaba-inc.com>
1 parent b62a0b4 commit 409ab83

6 files changed

Lines changed: 239 additions & 53 deletions

File tree

Lines changed: 137 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,137 @@
1+
import torch
2+
import tilelang
3+
import tilelang.language as T
4+
from tilelang.utils.tensor import torch_assert_close
5+
import itertools
6+
7+
8+
def ref_program(A, B):
9+
return (A.half() @ B.half().T).to(dtype=torch.float32)
10+
11+
12+
def manual_check_prog(C, C_ref):
13+
torch_assert_close(C[0], C_ref[0], rtol=0.01, atol=0.1)
14+
15+
16+
def supply_prog(args):
17+
a_param, b_param = args
18+
M, K = a_param.shape
19+
N, _ = b_param.shape
20+
a = (torch.randn(M, K, dtype=torch.float16, device='cuda') *
21+
0.01).to(dtype=torch.float8_e4m3fnuz)
22+
b = (torch.randn(N, K, dtype=torch.float16, device='cuda') *
23+
0.01).to(dtype=torch.float8_e4m3fnuz)
24+
return [a, b]
25+
26+
27+
def get_configs():
28+
block_Ms = [32, 64, 128]
29+
block_Ns = [32, 64, 128]
30+
block_Ks = [64, 128]
31+
num_stages = [0]
32+
num_threads = [256]
33+
k_packs = [1, 2]
34+
gemm_types = ["ss", "rs"]
35+
36+
valid_configs = []
37+
38+
for m, n, k, stages, t, kp, gemm_type in itertools.product(block_Ms, block_Ns, block_Ks,
39+
num_stages, num_threads, k_packs,
40+
gemm_types):
41+
valid_configs.append({
42+
"block_M": m,
43+
"block_N": n,
44+
"block_K": k,
45+
"num_stages": stages,
46+
"num_threads": t,
47+
"k_pack": kp,
48+
"gemm_type": gemm_type,
49+
})
50+
return valid_configs
51+
52+
53+
@tilelang.autotune(
54+
configs=get_configs(),
55+
cache_input_tensors=True,
56+
ref_prog=ref_program,
57+
manual_check_prog=manual_check_prog,
58+
supply_prog=supply_prog)
59+
@tilelang.jit(out_idx=[-1])
60+
def fp8_matmul(M, N, K, block_M, block_N, block_K, num_stages, num_threads, k_pack, gemm_type):
61+
dtype = "float8_e4m3fnuz"
62+
accum_dtype = "float"
63+
64+
@T.prim_func
65+
def gemm_fp8_rs(
66+
A: T.Tensor((M, K), dtype),
67+
B: T.Tensor((N, K), dtype),
68+
C: T.Tensor((M, N), accum_dtype),
69+
):
70+
with T.Kernel(
71+
T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=num_threads) as (bx, by):
72+
A_local = T.alloc_fragment((block_M, block_K), dtype)
73+
B_shared = T.alloc_shared((block_N, block_K), dtype)
74+
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
75+
76+
T.clear(C_local)
77+
for k in T.Pipelined(T.ceildiv(K, block_K), num_stages=num_stages):
78+
T.copy(A[by * block_M, k * block_K], A_local)
79+
T.copy(B[bx * block_N, k * block_K], B_shared)
80+
T.gemm(
81+
A_local,
82+
B_shared,
83+
C_local,
84+
transpose_B=True,
85+
k_pack=k_pack,
86+
policy=T.GemmWarpPolicy.FullRow)
87+
88+
T.copy(C_local, C[by * block_M, bx * block_N])
89+
90+
@T.prim_func
91+
def gemm_fp8_ss(
92+
A: T.Tensor((M, K), dtype),
93+
B: T.Tensor((N, K), dtype),
94+
C: T.Tensor((M, N), accum_dtype),
95+
):
96+
with T.Kernel(
97+
T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=num_threads) as (bx, by):
98+
A_shared = T.alloc_shared((block_M, block_K), dtype)
99+
B_shared = T.alloc_shared((block_N, block_K), dtype)
100+
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)
101+
102+
T.clear(C_local)
103+
for k in T.Pipelined(T.ceildiv(K, block_K), num_stages=num_stages):
104+
T.copy(A[by * block_M, k * block_K], A_shared)
105+
T.copy(B[bx * block_N, k * block_K], B_shared)
106+
T.gemm(
107+
A_shared,
108+
B_shared,
109+
C_local,
110+
transpose_B=True,
111+
k_pack=k_pack,
112+
policy=T.GemmWarpPolicy.FullRow)
113+
114+
T.copy(C_local, C[by * block_M, bx * block_N])
115+
116+
if gemm_type == "ss":
117+
return gemm_fp8_ss
118+
elif gemm_type == "rs":
119+
return gemm_fp8_rs
120+
else:
121+
raise ValueError(f"Invalid gemm_type: {gemm_type}")
122+
123+
124+
def test_gemm_fp8(M, N, K):
125+
kernel = fp8_matmul(M, N, K)
126+
a = (torch.randn(M, K, dtype=torch.float16, device='cuda') *
127+
0.01).to(dtype=torch.float8_e4m3fnuz)
128+
b = (torch.randn(N, K, dtype=torch.float16, device='cuda') *
129+
0.01).to(dtype=torch.float8_e4m3fnuz)
130+
c = kernel(a, b)
131+
ref_c = ref_program(a, b)
132+
torch_assert_close(c, ref_c, rtol=1e-2, atol=1e-2)
133+
print("passed~")
134+
135+
136+
if __name__ == "__main__":
137+
test_gemm_fp8(512, 512, 512)

‎src/layout/gemm_layouts.cc‎

Lines changed: 39 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -59,21 +59,39 @@ From https://github.com/RadeonOpenCompute/amd_matrix_instruction_calculator
5959
./matrix_calculator.py --architecture cdna1 --instruction v_mfma_f32_16x16x16f16
6060
--detail-instruction
6161
*/
62-
Fragment makeGemmFragmentAB16x16CDNA() {
62+
Fragment makeGemmFragmentAB16x16CDNA(const int k_pack) {
6363
IterVar i = make_itervar("i", 16);
64+
IterVar j = make_itervar("j", 16 * k_pack);
65+
IterVar rep = make_itervar("rep", 1);
66+
PrimExpr forward_thread = 16 * FloorDiv(j->var, 4 * k_pack) + i;
67+
PrimExpr index = FloorMod(j->var, 4 * k_pack);
68+
return Fragment({i, j}, {index}, forward_thread, rep);
69+
}
70+
71+
Fragment makeGemmFragmentAB16x16CDNATransposed(const int k_pack) {
72+
IterVar i = make_itervar("i", 16 * k_pack);
6473
IterVar j = make_itervar("j", 16);
6574
IterVar rep = make_itervar("rep", 1);
66-
PrimExpr forward_thread = 16 * FloorDiv(j->var, 4) + i;
67-
PrimExpr index = FloorMod(j->var, 4);
75+
PrimExpr forward_thread = 16 * FloorDiv(i->var, 4 * k_pack) + j;
76+
PrimExpr index = FloorMod(i->var, 4 * k_pack);
6877
return Fragment({i, j}, {index}, forward_thread, rep);
6978
}
7079

71-
Fragment makeGemmFragmentAB16x16CDNATransposed() {
80+
Fragment makeGemmFragmentAB16x32CDNA(const int k_pack) {
7281
IterVar i = make_itervar("i", 16);
82+
IterVar j = make_itervar("j", 32 * k_pack);
83+
IterVar rep = make_itervar("rep", 1);
84+
PrimExpr forward_thread = 16 * FloorDiv(j->var, 8 * k_pack) + i;
85+
PrimExpr index = FloorMod(j->var, 8 * k_pack);
86+
return Fragment({i, j}, {index}, forward_thread, rep);
87+
}
88+
89+
Fragment makeGemmFragmentAB16x32CDNATransposed(const int k_pack) {
90+
IterVar i = make_itervar("i", 32 * k_pack);
7391
IterVar j = make_itervar("j", 16);
7492
IterVar rep = make_itervar("rep", 1);
75-
PrimExpr forward_thread = 16 * FloorDiv(i->var, 4) + j;
76-
PrimExpr index = FloorMod(i->var, 4);
93+
PrimExpr forward_thread = 16 * FloorDiv(i->var, 8 * k_pack) + j;
94+
PrimExpr index = FloorMod(i->var, 8 * k_pack);
7795
return Fragment({i, j}, {index}, forward_thread, rep);
7896
}
7997

@@ -224,27 +242,34 @@ Fragment makeGemmFragmentB(const int block_m, const int block_n,
224242
Fragment makeGemmFragmentACDNA(const int block_m, const int block_n,
225243
const int block_k, const int warp_m,
226244
const int warp_n, const int element_size,
227-
bool transposed) {
245+
const int k_pack, bool transposed) {
228246
// assume not transposed
229247
ICHECK(block_m % warp_m == 0);
230248
ICHECK(block_n % warp_n == 0);
231249
ICHECK(warp_m % 16 == 0);
232-
ICHECK(block_k % 16 == 0);
250+
const int mfma_k = k_pack * (element_size == 16 ? 16 : 32);
251+
ICHECK(block_k % mfma_k == 0);
233252
ICHECK(element_size == 8 || element_size == 16)
234253
<< "element bitwidth=" << element_size;
235254
if (transposed) {
236255
auto base_layout =
237-
makeGemmFragmentAB16x16CDNATransposed()->Repeat({1, 1}, false, false);
256+
element_size == 16
257+
? makeGemmFragmentAB16x16CDNATransposed(k_pack)->Repeat(
258+
{1, 1}, false, false)
259+
: makeGemmFragmentAB16x32CDNATransposed(k_pack)->Repeat(
260+
{1, 1}, false, false);
238261
auto warp_layout =
239-
base_layout->Repeat({block_k / 16, warp_m / 16}, false, true);
262+
base_layout->Repeat({block_k / mfma_k, warp_m / 16}, false, true);
240263
auto block_layout = warp_layout->Repeat({1, block_m / warp_m}, true, true)
241264
->Replicate(block_n / warp_n);
242265
return block_layout;
243266
} else {
244267
auto base_layout =
245-
makeGemmFragmentAB16x16CDNA()->Repeat({1, 1}, false, false);
268+
element_size == 16
269+
? makeGemmFragmentAB16x16CDNA(k_pack)->Repeat({1, 1}, false, false)
270+
: makeGemmFragmentAB16x32CDNA(k_pack)->Repeat({1, 1}, false, false);
246271
auto warp_layout =
247-
base_layout->Repeat({warp_m / 16, block_k / 16}, false, false);
272+
base_layout->Repeat({warp_m / 16, block_k / mfma_k}, false, false);
248273
auto block_layout = warp_layout->Repeat({block_m / warp_m, 1}, true, true)
249274
->Replicate(block_n / warp_n);
250275
return block_layout;
@@ -397,7 +422,7 @@ Layout makeMatrixCoreSwizzleLayout(int stride, int continuous, int element_size,
397422
const int numBanks = 32;
398423
const int bankBitWidth = 32;
399424
const int SIMDWidth = 16;
400-
const int vecSize = 4 * kPack;
425+
const int vecSize = (64 / element_size) * kPack;
401426
const int innerDimLength = continuous;
402427
const int typeWidthInBit = element_size;
403428

@@ -616,12 +641,7 @@ Layout makeGemmABLayoutHopper(int mat_stride, int mat_continuous,
616641

617642
Layout makeGemmABLayoutCDNA(int stride, int continuous, int element_size,
618643
int kPack) {
619-
int vector_size = 128 / element_size;
620-
if (continuous % (vector_size * 4) == 0)
621-
return makeMatrixCoreSwizzleLayout(stride, continuous, element_size, kPack);
622-
else {
623-
return makeGemmABLayoutPadded(stride, continuous, element_size);
624-
}
644+
return makeMatrixCoreSwizzleLayout(stride, continuous, element_size, kPack);
625645
}
626646
} // namespace tl
627647
} // namespace tvm

‎src/layout/layout.h‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -154,7 +154,7 @@ Fragment makeGemmFragmentB(const int block_m, const int block_n,
154154
Fragment makeGemmFragmentACDNA(const int block_m, const int block_n,
155155
const int block_k, const int warp_m,
156156
const int warp_n, const int element_size,
157-
bool transposed = false);
157+
const int k_pack, bool transposed = false);
158158

159159
// Default Memory Layout
160160
Layout makeGemmLayoutLinear(int stride, int continuous);

‎src/op/gemm.cc‎

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -582,7 +582,7 @@ LayoutMap GemmNode::InferLayout(const LayoutInferArgs &T,
582582
results.Set(A, shared_layout);
583583
} else if (A.scope() == "local.fragment") {
584584
auto fragment = makeGemmFragmentACDNA(M, N, K, M / warp_m, N / warp_n,
585-
A->dtype.bits(), trans_A);
585+
A->dtype.bits(), kPack, trans_A);
586586
results.Set(A, fragment->BindThreadRange(thread_range));
587587
} else {
588588
ICHECK(0);
@@ -594,10 +594,6 @@ LayoutMap GemmNode::InferLayout(const LayoutInferArgs &T,
594594
*as_const_int(B->shape[dim_B - 1]), B->dtype.bits(), kPack);
595595

596596
results.Set(B, shared_layout);
597-
} else if (B.scope() == "local.fragment") {
598-
auto fragment =
599-
makeGemmFragmentB(M, N, K, M / warp_m, N / warp_n, trans_B);
600-
results.Set(B, fragment->BindThreadRange(thread_range));
601597
} else {
602598
ICHECK(0);
603599
}

0 commit comments

Comments
 (0)