Create a tensorized schedule for GEMM with MMA intrinsics.
(
workload,
k_inner,
in_dtype,
b_transposed,
i_factors,
j_factors,
k_factors,
index_map_A,
index_map_B,
index_map_C,
ldmatrix_a_intrin,
ldmatrix_b_intrin,
mma_intrin,
mma_fill_intrin,
mma_store_intrin,
shared_scope="shared",
)
| 19 | |
| 20 | |
| 21 | def mma_schedule( |
| 22 | workload, |
| 23 | k_inner, |
| 24 | in_dtype, |
| 25 | b_transposed, |
| 26 | i_factors, |
| 27 | j_factors, |
| 28 | k_factors, |
| 29 | index_map_A, |
| 30 | index_map_B, |
| 31 | index_map_C, |
| 32 | ldmatrix_a_intrin, |
| 33 | ldmatrix_b_intrin, |
| 34 | mma_intrin, |
| 35 | mma_fill_intrin, |
| 36 | mma_store_intrin, |
| 37 | shared_scope="shared", |
| 38 | ): |
| 39 | """Create a tensorized schedule for GEMM with MMA intrinsics.""" |
| 40 | import tvm # pylint: disable=import-outside-toplevel |
| 41 | |
| 42 | ir_module = tvm.IRModule({"main": workload}) |
| 43 | sch = tvm.s_tir.Schedule(ir_module) |
| 44 | |
| 45 | block = sch.get_sblock("C") |
| 46 | i, j, k = sch.get_loops(block) |
| 47 | i, i_tc = sch.split(i, factors=[None, 16]) |
| 48 | j, j_tc = sch.split(j, factors=[None, 16]) |
| 49 | k, k_tc = sch.split(k, factors=[None, k_inner]) |
| 50 | |
| 51 | sch.reorder(i, j, k, i_tc, j_tc, k_tc) |
| 52 | |
| 53 | block_inner = sch.blockize(i_tc) |
| 54 | block_outer, block_inner = block_inner, block |
| 55 | |
| 56 | num_ty = i_factors[2] * j_factors[2] |
| 57 | |
| 58 | i0, i1, i2, i3, i4 = sch.split(i, factors=i_factors) |
| 59 | j0, j1, j2, j3, j4 = sch.split(j, factors=j_factors) |
| 60 | k0, k1, k2 = sch.split(k, k_factors) |
| 61 | |
| 62 | sch.reorder(i0, j0, i1, j1, j2, i2, k0, k1, i3, j3, k2, i4, j4) |
| 63 | |
| 64 | block_idx = sch.fuse(i0, j0) |
| 65 | block_idy = sch.fuse(i1, j1) |
| 66 | thread_idy = sch.fuse(j2, i2) |
| 67 | sch.bind(block_idx, "blockIdx.x") |
| 68 | sch.bind(block_idy, "blockIdx.y") |
| 69 | sch.bind(thread_idy, "threadIdx.y") |
| 70 | |
| 71 | def fetch_to_shared(block, idx, ndim): |
| 72 | block_read = sch.cache_read(block, idx, shared_scope) |
| 73 | sch.compute_at(block_read, k0) |
| 74 | vector_size = 16 if in_dtype == "int8" else 8 |
| 75 | warp_size = 32 |
| 76 | fused = sch.fuse(*sch.get_loops(block_read)[-ndim:]) |
| 77 | _, f_1, f_2, f_3 = sch.split(fused, factors=[None, num_ty, warp_size, vector_size]) |
| 78 | sch.bind(f_2, "threadIdx.x") |
searching dependent graphs…