MCPcopy Create free account
hub / github.com/apache/tvm / mma_schedule

Function mma_schedule

python/tvm/testing/tir.py:21–130  ·  view source on GitHub ↗

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",
)

Source from the content-addressed store, hash-verified

19
20
21def 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")

Callers 2

get_mma_scheduleFunction · 0.90
run_testFunction · 0.90

Calls 15

get_sblockMethod · 0.95
get_loopsMethod · 0.95
splitMethod · 0.95
reorderMethod · 0.95
blockizeMethod · 0.95
fuseMethod · 0.95
bindMethod · 0.95
cache_readMethod · 0.95
compute_atMethod · 0.95
cache_writeMethod · 0.95
reverse_compute_atMethod · 0.95
decompose_reductionMethod · 0.95

Tested by 2

get_mma_scheduleFunction · 0.72
run_testFunction · 0.72

Used in the wild real call sites across dependent graphs

searching dependent graphs…