From 454fe9f6685819e9cc7943b8087b7f2af04327ff Mon Sep 17 00:00:00 2001 From: ventijing <304195737@qq.com> Date: Wed, 3 Sep 2025 14:40:49 +0800 Subject: [PATCH] [MetaXGPU]Add optimization strategies for MetaXGPU in the dlight --- NOTICE | 1 + python/tvm/dlight/gpu/matmul.py | 446 +++++++++++- tests/python/dlight/test_gpu_conv.py | 151 ++++ tests/python/dlight/test_gpu_matmul.py | 681 ++++++++++++++++++ .../dlight/test_gpu_matmul_tensorize.py | 681 ++++++++++++++++++ tests/scripts/task_python_unittest.sh | 1 + 6 files changed, 1960 insertions(+), 1 deletion(-) diff --git a/NOTICE b/NOTICE index 5c2f53f8dc6f..bc460950d8d4 100644 --- a/NOTICE +++ b/NOTICE @@ -2,6 +2,7 @@ The MetaX-MACA/mcTVM project is modified from apache/tvm (https://github.com/apa The following files may have been Modified by MetaX Integrated Circuits (Shanghai) Co., Ltd. in 2025. modified: .github/workflows/main.yml + modified: .gitmodules modified: .pre-commit-config.yaml modified: CMakeLists.txt modified: README.md diff --git a/python/tvm/dlight/gpu/matmul.py b/python/tvm/dlight/gpu/matmul.py index 368552c88d43..7377ede20c62 100644 --- a/python/tvm/dlight/gpu/matmul.py +++ b/python/tvm/dlight/gpu/matmul.py @@ -906,6 +906,432 @@ def tensorize_init_store_compute(): return sch +class MACAMatmulTensorization(GPUScheduleRule): + """ + The schedule rule for float16 mma computation. + func with attr 'dlight.do_not_tensorize' will not be tensorized. + """ + + def apply( # pylint: disable=too-many-locals,missing-docstring + self, + func: tir.PrimFunc, + target: Target, + _: bool, + ) -> Optional[tir.Schedule]: + from tvm.tir.tensor_intrin.maca import ( # pylint: disable=import-outside-toplevel + get_wmma_intrin_group, + ) + + if not isinstance(func, tir.PrimFunc) or not self.is_target_available(target): + return None + sch = tir.Schedule(func) + root_block = get_root_block(sch) + blocks = sch.get_child_blocks(root_block) + + if "dlight.do_not_tensorize" in func.attrs.keys(): + return None + + reduction_blocks = get_reduction_blocks(sch, blocks) + if reduction_blocks is None: + return None + + # Start Schedule + # Step 0. Get schedule config. + # NOTE: we can analyze the config by the hardware spec in the future + + # tensor core intrinsic size + micro_size_x = 16 + micro_size_y = 16 + micro_size_k = 16 + + warp_size = 64 + vector_size = 4 + + i_factors, j_factors, k_factors = ( + [None, 1, 2, 2], + [1, None, 2, 2], + [None, 2], + ) + + num_ty = i_factors[2] * j_factors[2] + x_pad_factor = i_factors[2] * i_factors[3] + y_pad_factor = j_factors[2] * j_factors[3] + k_pad_factor = k_factors[1] + + # Step 1. Normalize generic matmul to C[S, I, J] += A[S, I, K] * B[S, J, K] + # Reindex first and than analyze the index map + main_block = reduction_blocks[0] + reindex_a = sch.reindex(main_block, ("read", 0)) + reindex_b = sch.reindex(main_block, ("read", 1)) + reindex_c = sch.reindex(main_block, ("write", 0)) + + index_maps = get_index_map(sch.get(main_block)) + assert index_maps is not None + matmul_index_map, a_index_map, b_index_map, c_index_map = index_maps + + sch.transform_layout(reindex_a, ("write", 0), a_index_map) + sch.transform_layout(reindex_b, ("write", 0), b_index_map) + sch.transform_layout(reindex_c, ("read", 0), c_index_map) + sch.transform_block_layout(main_block, matmul_index_map) + + # Step 2. Padding for dynamic shape kernels + sch.pad_einsum( + main_block, + [ + 1, + micro_size_x * x_pad_factor, + micro_size_y * y_pad_factor, + micro_size_k * k_pad_factor, + ], + ) + + # Step 3. Schedule matmul to use tensor core + block = main_block + + batch, i, j, k = sch.get_loops(block) + + # inner loops for tensor core computation + i, i_inner = sch.split(i, factors=[None, micro_size_x]) + j, j_inner = sch.split(j, factors=[None, micro_size_y]) + k, k_inner = sch.split(k, factors=[None, micro_size_k]) + + sch.reorder(i, j, k, i_inner, j_inner, k_inner) + + block_inner = block + block_outer = sch.blockize(i_inner) + + i0, i1, i2, i3 = sch.split(i, factors=i_factors) + j0, j1, j2, j3 = sch.split(j, factors=j_factors) + k0, k1 = sch.split(k, k_factors) + sch.annotate(k0, "software_pipeline_order", [0, 3, 1, 4, 5, 2, 6]) + sch.annotate(k0, "software_pipeline_stage", [0, 0, 0, 0, 0, 1, 1]) + sch.annotate(k1, "software_pipeline_order", [0, 1, 2]) + sch.annotate(k1, "software_pipeline_stage", [0, 0, 1]) + + sch.reorder(i0, j0, i1, j1, j2, i2, k0, k1, i3, j3) + + block_idx = sch.fuse(i0, j0) + block_idy = sch.fuse(i1, j1) + thread_idy = sch.fuse(j2, i2) + sch.bind(batch, "blockIdx.z") + sch.bind(block_idx, "blockIdx.x") + sch.bind(block_idy, "blockIdx.y") + sch.bind(thread_idy, "threadIdx.y") + + def fetch_to_shared(block, idx, ndim): + block_read = sch.cache_read(block, idx, "shared.dyn") + sch.compute_at(block_read, k0) + fused = sch.fuse(*sch.get_loops(block_read)[-ndim:]) + + _, f_1, f_2, f_3 = sch.split(fused, factors=[None, num_ty, warp_size, vector_size]) + + sch.bind(f_2, "threadIdx.x") + sch.bind(f_1, "threadIdx.y") + sch.vectorize(f_3) + + sch.storage_align(block_read, 0, axis=-2, factor=16, offset=8) + sch.annotate(block_read, "tir.manifest_shared_memory_local_stage", 1) + sch.annotate(block_read, "double_buffer_scope", 0) + return block_read + + a_g2s = fetch_to_shared(block_outer, 0, 2) + b_g2s = fetch_to_shared(block_outer, 1, 2) + + auto_inline_producers(sch, a_g2s) + auto_inline_producers(sch, b_g2s) + + # create read cache to load matrix from shared memory to wmma fragments + A_mat = sch.cache_read(block_outer, 0, "wmma.matrix_a") + B_mat = sch.cache_read(block_outer, 1, "wmma.matrix_b") + sch.compute_at(A_mat, k1) + sch.compute_at(B_mat, k1) + + # create write cache to store matrix from wmma fragments to shared memory and global memory + accumulator_shared_to_global = sch.cache_write(block_outer, 0, "shared.dyn") + sch.storage_align(accumulator_shared_to_global, 0, -2, 16, 4) + + store = sch.cache_write(block_outer, 0, "wmma.accumulator") + sch.reverse_compute_at(store, thread_idy) + sch.reverse_compute_at(accumulator_shared_to_global, thread_idy) + + # split the store loop to match hardware intrinsic pattern + i, j = sch.get_loops(store)[-2:] + i0, i1 = sch.split(i, factors=[None, 16]) + j0, j1 = sch.split(j, factors=[None, 16]) + sch.reorder(i0, j0, i1, j1) + + block_init_c = sch.decompose_reduction(block_outer, k0) + block_init_c_inner = sch.get_child_blocks(block_init_c)[0] + + # Tensorization by hardware intrinsics + intrin_group = get_wmma_intrin_group( + load_scope="shared.dyn", + store_scope="shared.dyn", + in_dtype="float16", + out_dtype="float32", + trans_b=True, + ) + + try: + i, j = sch.get_loops(A_mat)[-2:] + i0, i1 = sch.split(i, factors=[None, 16]) + j0, j1 = sch.split(j, factors=[None, 16]) + sch.reorder(i0, j0, i1, j1) + sch.unroll(i0) + sch.unroll(j0) + sch.tensorize(i1, intrin_group["load_a"]) + + i, j = sch.get_loops(B_mat)[-2:] + i0, i1 = sch.split(i, factors=[None, 16]) + j0, j1 = sch.split(j, factors=[None, 16]) + sch.reorder(i0, j0, i1, j1) + sch.unroll(i0) + sch.unroll(j0) + sch.tensorize(i1, intrin_group["load_b"]) + except: # pylint: disable=bare-except + return None + + # Try to tensorize the init, store and compute block with f16 or f32 intrinsics + tensorize_success: bool = False + + def tensorize_init_store_compute(): + sch.tensorize(sch.get_loops(block_init_c_inner)[-2], intrin_group["init"]) + sch.tensorize(sch.get_loops(store)[-2], intrin_group["store"]) + sch.tensorize(sch.get_loops(block_inner)[-3], intrin_group["compute"]) + + try: + tensorize_init_store_compute() + tensorize_success = True + except: # pylint: disable=bare-except + intrin_group = get_wmma_intrin_group( + load_scope="shared.dyn", + store_scope="shared.dyn", + in_dtype="float16", + out_dtype="float16", + trans_b=True, + ) + + if not tensorize_success: + try: + tensorize_init_store_compute() + tensorize_success = True + except: # pylint: disable=bare-except + return None + auto_inline_consumer_chain(sch, accumulator_shared_to_global) + + fused = sch.fuse(*sch.get_loops(accumulator_shared_to_global)[-2:]) + _, f1, f2 = sch.split(fused, factors=[None, warp_size, vector_size]) + sch.bind(f1, "threadIdx.x") + sch.vectorize(f2) + + return sch if tensorize_success else None + + +class MACAMatmulInt8Tensorization(GPUScheduleRule): + """ + The schedule rule for int8 mma computation. + func with attr 'dlight.do_not_tensorize' will not be tensorized. + """ + + def apply( # pylint: disable=too-many-locals,missing-docstring + self, + func: tir.PrimFunc, + target: Target, + _: bool, + ) -> Optional[tir.Schedule]: + from tvm.tir.tensor_intrin.maca import ( # pylint: disable=import-outside-toplevel + get_wmma_intrin_group, + ) + + if not isinstance(func, tir.PrimFunc) or not self.is_target_available(target): + return None + sch = tir.Schedule(func) + root_block = get_root_block(sch) + blocks = sch.get_child_blocks(root_block) + + if "dlight.do_not_tensorize" in func.attrs.keys(): + return None + + reduction_blocks = get_reduction_blocks(sch, blocks) + if reduction_blocks is None: + return None + + # Start Schedule + # Step 0. Get schedule config. + # NOTE: we can analyze the config by the hardware spec in the future + + # tensor core intrinsic size + micro_size_x = 16 + micro_size_y = 16 + micro_size_k = 16 + + warp_size = 64 + vector_size = 4 + + i_factors, j_factors, k_factors = ( + [None, 1, 4, 2], + [1, None, 4, 2], + [None, 1], + ) + + num_ty = i_factors[2] * j_factors[2] + x_pad_factor = i_factors[2] * i_factors[3] + y_pad_factor = j_factors[2] * j_factors[3] + k_pad_factor = k_factors[1] + + # Step 1. Normalize generic matmul to C[S, I, J] += A[S, I, K] * B[S, J, K] + # Reindex first and than analyze the index map + main_block = reduction_blocks[0] + reindex_a = sch.reindex(main_block, ("read", 0)) + reindex_b = sch.reindex(main_block, ("read", 1)) + reindex_c = sch.reindex(main_block, ("write", 0)) + + index_maps = get_index_map(sch.get(main_block)) + assert index_maps is not None + matmul_index_map, a_index_map, b_index_map, c_index_map = index_maps + + sch.transform_layout(reindex_a, ("write", 0), a_index_map) + sch.transform_layout(reindex_b, ("write", 0), b_index_map) + sch.transform_layout(reindex_c, ("read", 0), c_index_map) + sch.transform_block_layout(main_block, matmul_index_map) + + # Step 2. Padding for dynamic shape kernels + sch.pad_einsum( + main_block, + [ + 1, + micro_size_x * x_pad_factor, + micro_size_y * y_pad_factor, + micro_size_k * k_pad_factor, + ], + ) + + # Step 3. Schedule matmul to use tensor core + block = main_block + + batch, i, j, k = sch.get_loops(block) + + # inner loops for tensor core computation + i, i_inner = sch.split(i, factors=[None, micro_size_x]) + j, j_inner = sch.split(j, factors=[None, micro_size_y]) + k, k_inner = sch.split(k, factors=[None, micro_size_k]) + + sch.reorder(i, j, k, i_inner, j_inner, k_inner) + + block_inner = block + block_outer = sch.blockize(i_inner) + + i0, i1, i2, i3 = sch.split(i, factors=i_factors) + j0, j1, j2, j3 = sch.split(j, factors=j_factors) + k0, k1 = sch.split(k, k_factors) + sch.annotate(k0, "software_pipeline_order", [0, 3, 1, 4, 5, 2, 6]) + sch.annotate(k0, "software_pipeline_stage", [0, 0, 0, 0, 0, 1, 1]) + sch.annotate(k1, "software_pipeline_order", [0, 1, 2]) + sch.annotate(k1, "software_pipeline_stage", [0, 0, 1]) + + sch.reorder(i0, j0, i1, j1, j2, i2, k0, k1, i3, j3) + + block_idx = sch.fuse(i0, j0) + block_idy = sch.fuse(i1, j1) + thread_idy = sch.fuse(j2, i2) + sch.bind(batch, "blockIdx.z") + sch.bind(block_idx, "blockIdx.x") + sch.bind(block_idy, "blockIdx.y") + sch.bind(thread_idy, "threadIdx.y") + + def fetch_to_shared(block, idx, ndim): + block_read = sch.cache_read(block, idx, "shared.dyn") + sch.compute_at(block_read, k0) + fused = sch.fuse(*sch.get_loops(block_read)[-ndim:]) + + _, f_1, f_2, f_3 = sch.split(fused, factors=[None, num_ty, warp_size, vector_size]) + + sch.bind(f_2, "threadIdx.x") + sch.bind(f_1, "threadIdx.y") + sch.vectorize(f_3) + + sch.storage_align(block_read, 0, axis=-2, factor=32, offset=16) + sch.annotate(block_read, "tir.manifest_shared_memory_local_stage", 1) + sch.annotate(block_read, "double_buffer_scope", 0) + return block_read + + a_g2s = fetch_to_shared(block_outer, 0, 2) + b_g2s = fetch_to_shared(block_outer, 1, 2) + + auto_inline_producers(sch, a_g2s) + auto_inline_producers(sch, b_g2s) + + # create read cache to load matrix from shared memory to wmma fragments + A_mat = sch.cache_read(block_outer, 0, "wmma.matrix_a") + B_mat = sch.cache_read(block_outer, 1, "wmma.matrix_b") + sch.compute_at(A_mat, k1) + sch.compute_at(B_mat, k1) + + # create write cache to store matrix from wmma fragments to shared memory and global memory + accumulator_shared_to_global = sch.cache_write(block_outer, 0, "shared.dyn") + sch.storage_align(accumulator_shared_to_global, 0, -2, 16, 4) + + store = sch.cache_write(block_outer, 0, "wmma.accumulator") + sch.reverse_compute_at(store, thread_idy) + sch.reverse_compute_at(accumulator_shared_to_global, thread_idy) + + # split the store loop to match hardware intrinsic pattern + i, j = sch.get_loops(store)[-2:] + i0, i1 = sch.split(i, factors=[None, 16]) + j0, j1 = sch.split(j, factors=[None, 16]) + sch.reorder(i0, j0, i1, j1) + + block_init_c = sch.decompose_reduction(block_outer, k0) + block_init_c_inner = sch.get_child_blocks(block_init_c)[0] + + # Tensorization by hardware intrinsics + intrin_group = get_wmma_intrin_group( + load_scope="shared.dyn", + store_scope="shared.dyn", + in_dtype="int8", + out_dtype="int32", + trans_b=True, + ) + + try: + i, j = sch.get_loops(A_mat)[-2:] + i0, i1 = sch.split(i, factors=[None, 16]) + j0, j1 = sch.split(j, factors=[None, 16]) + sch.reorder(i0, j0, i1, j1) + sch.unroll(i0) + sch.unroll(j0) + sch.tensorize(i1, intrin_group["load_a"]) + + i, j = sch.get_loops(B_mat)[-2:] + i0, i1 = sch.split(i, factors=[None, 16]) + j0, j1 = sch.split(j, factors=[None, 16]) + sch.reorder(i0, j0, i1, j1) + sch.unroll(i0) + sch.unroll(j0) + sch.tensorize(i1, intrin_group["load_b"]) + except: # pylint: disable=bare-except + return None + + def tensorize_init_store_compute(): + sch.tensorize(sch.get_loops(block_init_c_inner)[-2], intrin_group["init"]) + sch.tensorize(sch.get_loops(store)[-2], intrin_group["store"]) + sch.tensorize(sch.get_loops(block_inner)[-3], intrin_group["compute"]) + + try: + tensorize_init_store_compute() + except: # pylint: disable=bare-except + return None + + auto_inline_consumer_chain(sch, accumulator_shared_to_global) + + fused = sch.fuse(*sch.get_loops(accumulator_shared_to_global)[-2:]) + _, f1, f2 = sch.split(fused, factors=[None, warp_size, vector_size]) + sch.bind(f1, "threadIdx.x") + sch.vectorize(f2) + + return sch + + class Matmul(GPUScheduleRule): """The schedule rule for matmul-like computation""" @@ -926,7 +1352,7 @@ class Config: def get_configs(self, target: Target) -> Config: """Get the schedule config for the target""" - if target.kind.name == "cuda" or target.kind.name == "rocm": + if target.kind.name == "cuda" or target.kind.name == "rocm" or target.kind.name == "maca": return Matmul.Config( block_size_x=8, block_size_y=16, @@ -1038,6 +1464,24 @@ def is_inner_reduction(block_stmt, iter_infos): tensorize_sch = MatmulTensorization().apply(func, target, _) if tensorize_sch is not None: return tensorize_sch + elif target.kind.name == "maca": + apply_tensorization: bool = True + # the batch dimension is not taken into consideration. + for item_var in block_stmt.iter_vars[1:]: + extent = item_var.dom.extent + if isinstance(extent, tir.expr.IntImm): + if extent.value <= minimal_tensorize_threshold: + apply_tensorization = False + if apply_tensorization: + # Analyze read/write buffers and choose correct tensorizer: int8 or fp16. + in_dtype, out_dtype = get_in_out_dtypes(block_stmt) + tensorize_sch = None + if in_dtype == "int8" and out_dtype == "int32": + tensorize_sch = MACAMatmulInt8Tensorization().apply(func, target, _) + elif in_dtype == "float16" and out_dtype in ["float16", "float32"]: + tensorize_sch = MACAMatmulTensorization().apply(func, target, _) + if tensorize_sch is not None: + return tensorize_sch elif target.kind.name == "metal": try: return MetalMatmul().apply(func, target, _) diff --git a/tests/python/dlight/test_gpu_conv.py b/tests/python/dlight/test_gpu_conv.py index 90603a8bf293..3f8552f3267e 100644 --- a/tests/python/dlight/test_gpu_conv.py +++ b/tests/python/dlight/test_gpu_conv.py @@ -114,5 +114,156 @@ def expected(A: T.Buffer((14308, 3, 2, 14, 14), "float16"), W: T.Buffer((1280, 3 # fmt: on +class MACABeforeAfter(tvm.testing.CompareBeforeAfter): + @pytest.fixture + def transform(self): + def transform(mod): + with Target("maca"): + # Use Matmul rule for Conv for now + return dl.ApplyDefaultSchedule(dl.gpu.Matmul())(mod) + + return transform + + +@tvm.testing.requires_maca +class TestConv3dMACA(MACABeforeAfter): + # fmt: off + @T.prim_func + def before( + A: T.Buffer((14308, 3, 2, 14, 14), "float16"), + W: T.Buffer((1280, 3, 2, 14, 14), "float16"), + C: T.Buffer((14308, 1280, 1, 1, 1), "float16"), + ): + pad_A = T.alloc_buffer((14308, 3, 2, 14, 14), "float16") + for i0, i1, i2, i3, i4 in T.grid(14308, 3, 2, 14, 14): + with T.block("pad_A"): + v_i0, v_i1, v_i2, v_i3, v_i4 = T.axis.remap("SSSSS", [i0, i1, i2, i3, i4]) + pad_A[v_i0, v_i1, v_i2, v_i3, v_i4] = A[v_i0, v_i1, v_i2, v_i3, v_i4] + for nn, ff, yy, xx, zz, rc, ry, rx, rz in T.grid(14308, 1280, 1, 1, 1, 3, 2, 14, 14): + with T.block("C"): + v_nn, v_ff, v_yy, v_xx, v_zz, v_rc, v_ry, v_rx, v_rz = T.axis.remap("SSSSSRRRR", [nn, ff, yy, xx, zz, rc, ry, rx, rz]) + with T.init(): + C[v_nn, v_ff, v_yy, v_xx, v_zz] = T.float16(0.0) + C[v_nn, v_ff, v_yy, v_xx, v_zz] += pad_A[v_nn, v_rc, v_yy * 2 + v_ry, v_xx * 14 + v_rx, v_zz * 14 + v_rz]* W[v_ff, v_rc, v_ry, v_rx, v_rz] + + @T.prim_func + def expected(A: T.Buffer((14308, 3, 2, 14, 14), "float16"), W: T.Buffer((1280, 3, 2, 14, 14), "float16"), C: T.Buffer((14308, 1280, 1, 1, 1), "float16")): + T.func_attr({"global_symbol": "before", "tir.is_scheduled": True}) + # with T.block("root"): + pad_A_reindex_pad_shared_dyn = T.alloc_buffer((1, 14336, 1184), "float16", scope="shared.dyn") + W_reindex_pad_shared_dyn = T.alloc_buffer((1, 1280, 1184), "float16", scope="shared.dyn") + pad_A_reindex_pad_shared_dyn_wmma_matrix_a = T.alloc_buffer((1, 14336, 1184), "float16", scope="wmma.matrix_a") + W_reindex_pad_shared_dyn_wmma_matrix_b = T.alloc_buffer((1, 1280, 1184), "float16", scope="wmma.matrix_b") + C_reindex_pad_shared_dyn = T.alloc_buffer((1, 14336, 1280), "float16", scope="shared.dyn") + C_reindex_pad_shared_dyn_wmma_accumulator = T.alloc_buffer((1, 14336, 1280), "float16", scope="wmma.accumulator") + for ax0 in T.thread_binding(1, thread="blockIdx.z"): + for ax1_0_0_ax2_0_0_fused in T.thread_binding(224, thread="blockIdx.x"): + for ax1_0_1_ax2_0_1_fused in T.thread_binding(20, thread="blockIdx.y"): + for ax2_0_2_ax1_0_2_fused in T.thread_binding(4, thread="threadIdx.y"): + for ax1_0_3_init, ax2_0_3_init in T.grid(2, 2): + with T.block("C_o_init"): + v0_o = T.axis.spatial(1, ax0) + v1_o = T.axis.spatial(896, ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax1_0_3_init) + v2_o = T.axis.spatial(80, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax2_0_3_init) + T.reads() + T.writes(C_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + with T.block("C_init_o"): + v1_i_init_o = T.axis.spatial(1, 0) + v2_i_init_o = T.axis.spatial(1, 0) + T.reads() + T.writes(C_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + C_1 = T.match_buffer(C_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_fill_fragment(C_1.data, 16, 16, 16, C_1.elem_offset // C_1.strides[0] // 16 * (C_1.strides[0] // 16) + C_1.elem_offset % C_1.strides[0] // 16, T.float32(0.0)) + for ax3_0_0 in T.serial(37, annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): + for ax0_ax1_fused_0 in range(2): + for ax0_ax1_fused_1 in T.thread_binding(4, thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(4): + with T.block("pad_A_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(14336, ax1_0_0_ax2_0_0_fused * 64 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) // 32) + v2 = T.axis.spatial(1184, ax3_0_0 * 32 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) % 32) + T.reads(A[v1, v2 // 392, v2 // 196 % 2, v2 // 14 % 14, v2 % 14]) + T.writes(pad_A_reindex_pad_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 8]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + pad_A_reindex_pad_shared_dyn[v0, v1, v2] = T.if_then_else(v1 < 14308 and v2 < 1176, A[v1, v2 // 392, v2 // 196 % 2, v2 // 14 % 14, v2 % 14], T.float16(0.0)) + for ax0_ax1_fused_0 in range(2): + for ax0_ax1_fused_1 in T.thread_binding(4, thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(4): + with T.block("W_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(1280, ax1_0_1_ax2_0_1_fused * 64 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) // 32) + v2 = T.axis.spatial(1184, ax3_0_0 * 32 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) % 32) + T.reads(W[v1, v2 // 392, v2 // 196 % 2, v2 // 14 % 14, v2 % 14]) + T.writes(W_reindex_pad_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 8]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + W_reindex_pad_shared_dyn[v0, v1, v2] = T.if_then_else(v2 < 1176, W[v1, v2 // 392, v2 // 196 % 2, v2 // 14 % 14, v2 % 14], T.float16(0.0)) + for ax3_0_1 in T.serial(2, annotations={"software_pipeline_order": [0, 1, 2], "software_pipeline_stage": [0, 0, 1]}): + for ax0_0 in T.unroll(2): + for ax1_0 in T.unroll(1): + with T.block("pad_A_reindex_pad_shared.dyn_wmma.matrix_a_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(896, ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax0_0) + v2_o = T.axis.spatial(74, ax3_0_0 * 2 + ax3_0_1 + ax1_0) + T.reads(pad_A_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(pad_A_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A_1 = T.match_buffer(pad_A_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C_1 = T.match_buffer(pad_A_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_a", offset_factor=16) + T.tvm_load_matrix_sync(C_1.data, 16, 16, 16, C_1.elem_offset // C_1.strides[0] // 16 * (C_1.strides[0] // 16) + C_1.elem_offset % C_1.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), A_1.data, A_1.elem_offset, A_1.strides[0] * 16, 1), A_1.strides[0], "row_major") + for ax0_0 in T.unroll(2): + for ax1_0 in T.unroll(1): + with T.block("W_reindex_pad_shared.dyn_wmma.matrix_b_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(80, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax0_0) + v2_o = T.axis.spatial(74, ax3_0_0 * 2 + ax3_0_1 + ax1_0) + T.reads(W_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(W_reindex_pad_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A_1 = T.match_buffer(W_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C_1 = T.match_buffer(W_reindex_pad_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_b", offset_factor=16) + T.tvm_load_matrix_sync(C_1.data, 16, 16, 16, C_1.elem_offset // C_1.strides[0] // 16 * (C_1.strides[0] // 16) + C_1.elem_offset % C_1.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), A_1.data, A_1.elem_offset, A_1.strides[0] * 16, 1), A_1.strides[0], "col_major") + for ax1_0_3, ax2_0_3 in T.grid(2, 2): + with T.block("C_o_update"): + v0_o = T.axis.spatial(1, ax0) + v1_o = T.axis.spatial(896, ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax1_0_3) + v2_o = T.axis.spatial(80, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax2_0_3) + v3_o = T.axis.reduce(74, ax3_0_0 * 2 + ax3_0_1) + T.reads(C_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], pad_A_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], W_reindex_pad_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) + T.writes(C_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + with T.block("C_o"): + v1_i_o = T.axis.spatial(1, 0) + v2_i_o = T.axis.spatial(1, 0) + v3_i_o = T.axis.reduce(1, 0) + T.reads(C_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], pad_A_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], W_reindex_pad_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) + T.writes(C_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A_1 = T.match_buffer(pad_A_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="wmma.matrix_a", offset_factor=16) + B = T.match_buffer(W_reindex_pad_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "float16", strides=("B_s0", "B_s1"), scope="wmma.matrix_b", offset_factor=16) + C_1 = T.match_buffer(C_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_mma_sync(C_1.data, C_1.elem_offset // C_1.strides[0] // 16 * (C_1.strides[0] // 16) + C_1.elem_offset % C_1.strides[0] // 16, A_1.data, A_1.elem_offset // A_1.strides[0] // 16 * (A_1.strides[0] // 16) + A_1.elem_offset % A_1.strides[0] // 16, B.data, B.elem_offset // B.strides[0] // 16 * (B.strides[0] // 16) + B.elem_offset % B.strides[0] // 16, C_1.data, C_1.elem_offset // C_1.strides[0] // 16 * (C_1.strides[0] // 16) + C_1.elem_offset % C_1.strides[0] // 16) + for ax0_0, ax1_0 in T.grid(2, 2): + with T.block("C_reindex_pad_shared.dyn_wmma.accumulator_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(896, ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax0_0) + v2_o = T.axis.spatial(80, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax1_0) + T.reads(C_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(C_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A_1 = T.match_buffer(C_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="wmma.accumulator", offset_factor=16) + C_1 = T.match_buffer(C_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="shared.dyn", offset_factor=16) + T.tvm_store_matrix_sync(A_1.data, 16, 16, 16, A_1.elem_offset // A_1.strides[0] // 16 * (A_1.strides[0] // 16) + A_1.elem_offset % A_1.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), C_1.data, C_1.elem_offset, C_1.strides[0] * 16, 2), C_1.strides[0], "row_major") + for ax0_ax1_fused_0 in range(4): + for ax0_ax1_fused_1 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_2 in T.vectorized(4): + with T.block("C_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(14336, ax1_0_0_ax2_0_0_fused * 64 + ax2_0_2_ax1_0_2_fused % 2 * 32 + (ax0_ax1_fused_0 * 256 + ax0_ax1_fused_1 * 4 + ax0_ax1_fused_2) // 32) + v2 = T.axis.spatial(1280, ax1_0_1_ax2_0_1_fused * 64 + ax2_0_2_ax1_0_2_fused // 2 * 32 + (ax0_ax1_fused_0 * 256 + ax0_ax1_fused_1 * 4 + ax0_ax1_fused_2) % 32) + T.where(ax1_0_0_ax2_0_0_fused * 64 + ax2_0_2_ax1_0_2_fused % 2 * 32 + ((ax0_ax1_fused_0 * 64 + ax0_ax1_fused_1) * 4 + ax0_ax1_fused_2) // 32 < 14308) + T.reads(C_reindex_pad_shared_dyn[v0, v1, v2]) + T.writes(C[v1, v2, 0, 0, 0]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 4]]}) + C[v1, v2, 0, 0, 0] = C_reindex_pad_shared_dyn[v0, v1, v2] + # fmt: on + + if __name__ == "__main__": tvm.testing.main() diff --git a/tests/python/dlight/test_gpu_matmul.py b/tests/python/dlight/test_gpu_matmul.py index f27d9d370fce..214b74c67694 100644 --- a/tests/python/dlight/test_gpu_matmul.py +++ b/tests/python/dlight/test_gpu_matmul.py @@ -842,5 +842,686 @@ def expected(lv452: T.Buffer((T.int64(512), T.int64(12288)), "uint32"), lv453: T # fmt: on +class MACABeforeAfter(tvm.testing.CompareBeforeAfter): + @pytest.fixture + def transform(self): + def transform(mod): + with Target("maca"): + return dl.ApplyDefaultSchedule(dl.gpu.Matmul())(mod) + + return transform + + +@tvm.testing.requires_maca +class TestMatmulMACA(MACABeforeAfter): + # fmt: off + @T.prim_func + def before(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), var_matmul: T.handle): + m = T.int64() + inp0 = T.match_buffer(var_inp0, (T.int64(1), m, T.int64(4096))) + matmul = T.match_buffer(var_matmul, (T.int64(1), m, T.int64(4096))) + for i0, i1, i2, k in T.grid(T.int64(1), m, T.int64(4096), T.int64(4096)): + with T.block("matmul"): + v_i0, v_i1, v_i2, v_k = T.axis.remap("SSSR", [i0, i1, i2, k]) + with T.init(): + matmul[v_i0, v_i1, v_i2] = T.float32(0) + matmul[v_i0, v_i1, v_i2] = matmul[v_i0, v_i1, v_i2] + inp0[v_i0, v_i1, v_k] * inp1[v_k, v_i2] + + @T.prim_func + def expected(var_inp0: T.handle, inp1: T.Buffer((T.int64(4096), T.int64(4096)), "float32"), var_matmul: T.handle): + T.func_attr({"tir.is_scheduled": True}) + m = T.int64() + inp0 = T.match_buffer(var_inp0, (T.int64(1), m, T.int64(4096))) + matmul = T.match_buffer(var_matmul, (T.int64(1), m, T.int64(4096))) + # with T.block("root"): + matmul_reindex_pad_local = T.alloc_buffer((T.int64(1), (m + T.int64(31)) // T.int64(32) * T.int64(32), T.int64(4096)), scope="local") + inp0_reindex_pad_shared = T.alloc_buffer((T.int64(1), (m + T.int64(31)) // T.int64(32) * T.int64(32), T.int64(4096)), scope="shared") + inp1_reindex_shared = T.alloc_buffer((T.int64(1), T.int64(4096), T.int64(4096)), scope="shared") + for ax0_ax2_0_fused in T.thread_binding(T.int64(64), thread="blockIdx.y"): + for ax1_0 in T.thread_binding((m + T.int64(31)) // T.int64(32), thread="blockIdx.x"): + for ax2_1 in T.thread_binding(T.int64(1), thread="vthread.y"): + for ax1_1 in T.thread_binding(T.int64(1), thread="vthread.x"): + for ax2_2 in T.thread_binding(T.int64(16), thread="threadIdx.y"): + for ax1_2 in T.thread_binding(T.int64(8), thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}): + for ax1_3_init, ax2_3_0_init in T.grid(T.int64(4), T.int64(2)): + for ax2_3_1_init in T.vectorized(T.int64(2)): + with T.block("matmul_init"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial((m + T.int64(31)) // T.int64(32) * T.int64(32), ax1_0 * T.int64(32) + ax1_1 * T.int64(32) + ax1_2 * T.int64(4) + ax1_3_init) + v2 = T.axis.spatial(T.int64(4096), ax0_ax2_0_fused * T.int64(64) + ax2_1 * T.int64(64) + ax2_2 * T.int64(4) + ax2_3_0_init * T.int64(2) + ax2_3_1_init) + T.reads() + T.writes(matmul_reindex_pad_local[T.int64(0), v1, v2]) + matmul_reindex_pad_local[T.int64(0), v1, v2] = T.float32(0) + for ax3_0 in range(T.int64(256)): + for ax0_ax1_ax2_fused_0 in T.thread_binding(T.int64(16), thread="threadIdx.y"): + for ax0_ax1_ax2_fused_1 in T.thread_binding(T.int64(8), thread="threadIdx.x"): + for ax0_ax1_ax2_fused_2 in range(T.int64(2)): + for ax0_ax1_ax2_fused_3 in T.vectorized(T.int64(2)): + with T.block("inp0_reindex_pad_shared"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial((m + T.int64(31)) // T.int64(32) * T.int64(32), ax1_0 * T.int64(32) + (ax0_ax1_ax2_fused_0 * T.int64(32) + ax0_ax1_ax2_fused_1 * T.int64(4) + ax0_ax1_ax2_fused_2 * T.int64(2) + ax0_ax1_ax2_fused_3) // T.int64(16)) + v2 = T.axis.spatial(T.int64(4096), ax3_0 * T.int64(16) + (ax0_ax1_ax2_fused_0 * T.int64(32) + ax0_ax1_ax2_fused_1 * T.int64(4) + ax0_ax1_ax2_fused_2 * T.int64(2) + ax0_ax1_ax2_fused_3) % T.int64(16)) + T.reads(inp0[v0, v1, v2]) + T.writes(inp0_reindex_pad_shared[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 8, 2]]}) + inp0_reindex_pad_shared[v0, v1, v2] = T.if_then_else(v1 < m, inp0[v0, v1, v2], T.float32(0)) + for ax0_ax1_ax2_fused_0 in T.thread_binding(T.int64(16), thread="threadIdx.y"): + for ax0_ax1_ax2_fused_1 in T.thread_binding(T.int64(8), thread="threadIdx.x"): + for ax0_ax1_ax2_fused_2 in range(T.int64(4)): + for ax0_ax1_ax2_fused_3 in T.vectorized(T.int64(2)): + with T.block("inp1_reindex_shared"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial(T.int64(4096), ax0_ax2_0_fused * T.int64(64) + (ax0_ax1_ax2_fused_0 * T.int64(64) + ax0_ax1_ax2_fused_1 * T.int64(8) + ax0_ax1_ax2_fused_2 * T.int64(2) + ax0_ax1_ax2_fused_3) // T.int64(16)) + v2 = T.axis.spatial(T.int64(4096), ax3_0 * T.int64(16) + (ax0_ax1_ax2_fused_0 * T.int64(64) + ax0_ax1_ax2_fused_1 * T.int64(8) + ax0_ax1_ax2_fused_2 * T.int64(2) + ax0_ax1_ax2_fused_3) % T.int64(16)) + T.reads(inp1[v2, v1]) + T.writes(inp1_reindex_shared[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 8, 2]]}) + inp1_reindex_shared[v0, v1, v2] = inp1[v2, v1] + for ax3_1, ax1_3, ax2_3_0 in T.grid(T.int64(16), T.int64(4), T.int64(2)): + for ax2_3_1 in T.vectorized(T.int64(2)): + with T.block("matmul_update"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial((m + T.int64(31)) // T.int64(32) * T.int64(32), ax1_0 * T.int64(32) + ax1_1 * T.int64(32) + ax1_2 * T.int64(4) + ax1_3) + v2 = T.axis.spatial(T.int64(4096), ax0_ax2_0_fused * T.int64(64) + ax2_1 * T.int64(64) + ax2_2 * T.int64(4) + ax2_3_0 * T.int64(2) + ax2_3_1) + v3 = T.axis.reduce(T.int64(4096), ax3_0 * T.int64(16) + ax3_1) + T.reads(matmul_reindex_pad_local[T.int64(0), v1, v2], inp0_reindex_pad_shared[T.int64(0), v1, v3], inp1_reindex_shared[T.int64(0), v2, v3]) + T.writes(matmul_reindex_pad_local[T.int64(0), v1, v2]) + matmul_reindex_pad_local[T.int64(0), v1, v2] = matmul_reindex_pad_local[T.int64(0), v1, v2] + inp0_reindex_pad_shared[T.int64(0), v1, v3] * inp1_reindex_shared[T.int64(0), v2, v3] + for ax0, ax1, ax2_0 in T.grid(T.int64(1), T.int64(4), T.int64(2)): + for ax2_1_1 in T.vectorized(T.int64(2)): + with T.block("matmul_reindex_pad_local"): + v0 = T.axis.spatial(T.int64(1), ax0) + v1 = T.axis.spatial((m + T.int64(31)) // T.int64(32) * T.int64(32), ax1_0 * T.int64(32) + ax1_2 * T.int64(4) + ax1) + v2 = T.axis.spatial(T.int64(4096), ax0_ax2_0_fused * T.int64(64) + ax2_2 * T.int64(4) + ax2_0 * T.int64(2) + ax2_1_1) + T.where(ax1_0 * T.int64(32) + ax1_2 * T.int64(4) + ax1 < m) + T.reads(matmul_reindex_pad_local[v0, v1, v2]) + T.writes(matmul[T.int64(0), v1, v2]) + matmul[T.int64(0), v1, v2] = matmul_reindex_pad_local[v0, v1, v2] + # fmt: on + + +def test_matmul_int32(): + # fmt: off + @T.prim_func(private=True) + def func(var_inp0: T.handle, inp1: T.Buffer((4096, 4096), "float32"), var_matmul: T.handle): + m = T.int32() + inp0 = T.match_buffer(var_inp0, (1, m, 4096)) + matmul = T.match_buffer(var_matmul, (1, m, 4096)) + for i0, i1, i2, k in T.grid(1, m, 4096, 4096): + with T.block("matmul"): + v_i0, v_i1, v_i2, v_k = T.axis.remap("SSSR", [i0, i1, i2, k]) + with T.init(): + matmul[v_i0, v_i1, v_i2] = T.float32(0) + matmul[v_i0, v_i1, v_i2] = matmul[v_i0, v_i1, v_i2] + inp0[v_i0, v_i1, v_k] * inp1[v_k, v_i2] + + @T.prim_func(private=True) + def expected(var_inp0: T.handle, inp1: T.Buffer((4096, 4096), "float32"), var_matmul: T.handle): + T.func_attr({"tir.is_scheduled": True}) + m = T.int32() + inp0 = T.match_buffer(var_inp0, (1, m, 4096)) + matmul = T.match_buffer(var_matmul, (1, m, 4096)) + # with T.block("root"): + matmul_reindex_pad_local = T.alloc_buffer((1, (m + 31) // 32 * 32, 4096), scope="local") + inp0_reindex_pad_shared = T.alloc_buffer((1, (m + 31) // 32 * 32, 4096), scope="shared") + inp1_reindex_shared = T.alloc_buffer((1, 4096, 4096), scope="shared") + for ax0_ax2_0_fused in T.thread_binding(64, thread="blockIdx.y"): + for ax1_0 in T.thread_binding((m + 31) // 32, thread="blockIdx.x"): + for ax2_1 in T.thread_binding(1, thread="vthread.y"): + for ax1_1 in T.thread_binding(1, thread="vthread.x"): + for ax2_2 in T.thread_binding(16, thread="threadIdx.y"): + for ax1_2 in T.thread_binding(8, thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}): + for ax1_3_init, ax2_3_0_init in T.grid(4, 2): + for ax2_3_1_init in T.vectorized(2): + with T.block("matmul_init"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial((m + 31) // 32 * 32, ax1_0 * 32 + ax1_1 * 32 + ax1_2 * 4 + ax1_3_init) + v2 = T.axis.spatial(4096, ax0_ax2_0_fused * 64 + ax2_1 * 64 + ax2_2 * 4 + ax2_3_0_init * 2 + ax2_3_1_init) + T.reads() + T.writes(matmul_reindex_pad_local[0, v1, v2]) + matmul_reindex_pad_local[0, v1, v2] = T.float32(0) + for ax3_0 in range(256): + for ax0_ax1_ax2_fused_0 in T.thread_binding(16, thread="threadIdx.y"): + for ax0_ax1_ax2_fused_1 in T.thread_binding(8, thread="threadIdx.x"): + for ax0_ax1_ax2_fused_2 in range(2): + for ax0_ax1_ax2_fused_3 in T.vectorized(2): + with T.block("inp0_reindex_pad_shared"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial((m + 31) // 32 * 32, ax1_0 * 32 + (ax0_ax1_ax2_fused_0 * 32 + ax0_ax1_ax2_fused_1 * 4 + ax0_ax1_ax2_fused_2 * 2 + ax0_ax1_ax2_fused_3) // 16) + v2 = T.axis.spatial(4096, ax3_0 * 16 + (ax0_ax1_ax2_fused_0 * 32 + ax0_ax1_ax2_fused_1 * 4 + ax0_ax1_ax2_fused_2 * 2 + ax0_ax1_ax2_fused_3) % 16) + T.reads(inp0[v0, v1, v2]) + T.writes(inp0_reindex_pad_shared[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 8, 2]]}) + inp0_reindex_pad_shared[v0, v1, v2] = T.if_then_else(v1 < m, inp0[v0, v1, v2], T.float32(0)) + for ax0_ax1_ax2_fused_0 in T.thread_binding(16, thread="threadIdx.y"): + for ax0_ax1_ax2_fused_1 in T.thread_binding(8, thread="threadIdx.x"): + for ax0_ax1_ax2_fused_2 in range(4): + for ax0_ax1_ax2_fused_3 in T.vectorized(2): + with T.block("inp1_reindex_shared"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(4096, ax0_ax2_0_fused * 64 + (ax0_ax1_ax2_fused_0 * 64 + ax0_ax1_ax2_fused_1 * 8 + ax0_ax1_ax2_fused_2 * 2 + ax0_ax1_ax2_fused_3) // 16) + v2 = T.axis.spatial(4096, ax3_0 * 16 + (ax0_ax1_ax2_fused_0 * 64 + ax0_ax1_ax2_fused_1 * 8 + ax0_ax1_ax2_fused_2 * 2 + ax0_ax1_ax2_fused_3) % 16) + T.reads(inp1[v2, v1]) + T.writes(inp1_reindex_shared[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 8, 2]]}) + inp1_reindex_shared[v0, v1, v2] = inp1[v2, v1] + for ax3_1, ax1_3, ax2_3_0 in T.grid(16, 4, 2): + for ax2_3_1 in T.vectorized(2): + with T.block("matmul_update"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial((m + 31) // 32 * 32, ax1_0 * 32 + ax1_1 * 32 + ax1_2 * 4 + ax1_3) + v2 = T.axis.spatial(4096, ax0_ax2_0_fused * 64 + ax2_1 * 64 + ax2_2 * 4 + ax2_3_0 * 2 + ax2_3_1) + v3 = T.axis.reduce(4096, ax3_0 * 16 + ax3_1) + T.reads(matmul_reindex_pad_local[0, v1, v2], inp0_reindex_pad_shared[0, v1, v3], inp1_reindex_shared[0, v2, v3]) + T.writes(matmul_reindex_pad_local[0, v1, v2]) + matmul_reindex_pad_local[0, v1, v2] = matmul_reindex_pad_local[0, v1, v2] + inp0_reindex_pad_shared[0, v1, v3] * inp1_reindex_shared[0, v2, v3] + for ax0, ax1, ax2_0 in T.grid(1, 4, 2): + for ax2_1_1 in T.vectorized(2): + with T.block("matmul_reindex_pad_local"): + v0 = T.axis.spatial(1, ax0) + v1 = T.axis.spatial((m + 31) // 32 * 32, ax1_0 * 32 + ax1_2 * 4 + ax1) + v2 = T.axis.spatial(4096, ax0_ax2_0_fused * 64 + ax2_2 * 4 + ax2_0 * 2 + ax2_1_1) + T.where(ax1_0 * 32 + ax1_2 * 4 + ax1 < m) + T.reads(matmul_reindex_pad_local[v0, v1, v2]) + T.writes(matmul[0, v1, v2]) + matmul[0, v1, v2] = matmul_reindex_pad_local[v0, v1, v2] + # fmt: on + + mod = tvm.IRModule({"main": func}) + with Target("maca"): + mod = dl.ApplyDefaultSchedule(dl.gpu.Matmul())(mod) + tvm.ir.assert_structural_equal(mod["main"], expected) + + +@tvm.testing.requires_maca +class TestFusedMatmulMACA(MACABeforeAfter): + # fmt: off + + @T.prim_func + def before(W: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), S: T.Buffer((T.int64(128), T.int64(4096)), "uint32"), A: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32"), C: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32"), Out: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32")): + var_decode_intermediate = T.alloc_buffer((T.int64(4096), T.int64(4096))) + var_matmul_intermediate = T.alloc_buffer((T.int64(1), T.int64(32), T.int64(4096))) + for i, j in T.grid(T.int64(4096), T.int64(4096)): + with T.block("decode"): + v_i, v_j = T.axis.remap("SS", [i, j]) + T.reads(W[v_i // T.int64(8), v_j], S[v_i // T.int64(32), v_j]) + T.writes(var_decode_intermediate[v_i, v_j]) + var_decode_intermediate[v_i, v_j] = T.Cast("float32", T.bitwise_and(T.shift_right(W[v_i // T.int64(8), v_j], T.Cast("uint32", v_i % T.int64(8) * T.int64(4))), T.uint32(15))) * T.reinterpret("float32", T.shift_left(T.bitwise_and(S[v_i // T.int64(32), v_j], T.uint32(65535)), T.uint32(16))) + T.reinterpret("float32", T.shift_left(T.bitwise_and(T.shift_right(S[v_i // T.int64(32), v_j], T.uint32(16)), T.uint32(65535)), T.uint32(16))) + for i0, i1, i2, k in T.grid(T.int64(1), T.int64(32), T.int64(4096), T.int64(4096)): + with T.block("matmul"): + v_i0, v_i1, v_i2, v_k = T.axis.remap("SSSR", [i0, i1, i2, k]) + T.reads(A[v_i0, v_i1, v_k], var_decode_intermediate[v_k, v_i2]) + T.writes(var_matmul_intermediate[v_i0, v_i1, v_i2]) + with T.init(): + var_matmul_intermediate[v_i0, v_i1, v_i2] = T.float32(0) + var_matmul_intermediate[v_i0, v_i1, v_i2] = var_matmul_intermediate[v_i0, v_i1, v_i2] + A[v_i0, v_i1, v_k] * var_decode_intermediate[v_k, v_i2] + for ax0, ax1, ax2 in T.grid(T.int64(1), T.int64(32), T.int64(4096)): + with T.block("T_add"): + v_ax0, v_ax1, v_ax2 = T.axis.remap("SSS", [ax0, ax1, ax2]) + T.reads(C[v_ax0, v_ax1, v_ax2], var_matmul_intermediate[v_ax0, v_ax1, v_ax2]) + T.writes(Out[v_ax0, v_ax1, v_ax2]) + Out[v_ax0, v_ax1, v_ax2] = C[v_ax0, v_ax1, v_ax2] + var_matmul_intermediate[v_ax0, v_ax1, v_ax2] + + @T.prim_func + def expected(W: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), S: T.Buffer((T.int64(128), T.int64(4096)), "uint32"), A: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32"), C: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32"), Out: T.Buffer((T.int64(1), T.int64(32), T.int64(4096)), "float32")): + T.func_attr({"tir.is_scheduled": True}) + # with T.block("root"): + var_matmul_intermediate_reindex_local = T.alloc_buffer((T.int64(1), T.int64(32), T.int64(4096)), scope="local") + A_reindex_shared = T.alloc_buffer((T.int64(1), T.int64(32), T.int64(4096)), scope="shared") + var_decode_intermediate_reindex_shared = T.alloc_buffer((T.int64(1), T.int64(4096), T.int64(4096)), scope="shared") + for ax0_ax2_0_fused in T.thread_binding(T.int64(64), thread="blockIdx.y"): + for ax1_0 in T.thread_binding(T.int64(1), thread="blockIdx.x"): + for ax2_1 in T.thread_binding(T.int64(1), thread="vthread.y"): + for ax1_1 in T.thread_binding(T.int64(1), thread="vthread.x"): + for ax2_2 in T.thread_binding(T.int64(16), thread="threadIdx.y"): + for ax1_2 in T.thread_binding(T.int64(8), thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}): + for ax1_3_init, ax2_3_0_init in T.grid(T.int64(4), T.int64(2)): + for ax2_3_1_init in T.vectorized(T.int64(2)): + with T.block("matmul_init"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial(T.int64(32), ax1_0 * T.int64(32) + ax1_1 * T.int64(32) + ax1_2 * T.int64(4) + ax1_3_init) + v2 = T.axis.spatial(T.int64(4096), ax0_ax2_0_fused * T.int64(64) + ax2_1 * T.int64(64) + ax2_2 * T.int64(4) + ax2_3_0_init * T.int64(2) + ax2_3_1_init) + T.reads() + T.writes(var_matmul_intermediate_reindex_local[T.int64(0), v1, v2]) + var_matmul_intermediate_reindex_local[T.int64(0), v1, v2] = T.float32(0) + for ax3_0 in range(T.int64(256)): + for ax0_ax1_ax2_fused_0 in T.thread_binding(T.int64(16), thread="threadIdx.y"): + for ax0_ax1_ax2_fused_1 in T.thread_binding(T.int64(8), thread="threadIdx.x"): + for ax0_ax1_ax2_fused_2 in range(T.int64(2)): + for ax0_ax1_ax2_fused_3 in T.vectorized(T.int64(2)): + with T.block("A_reindex_shared"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial(T.int64(32), (ax0_ax1_ax2_fused_0 * T.int64(32) + ax0_ax1_ax2_fused_1 * T.int64(4) + ax0_ax1_ax2_fused_2 * T.int64(2) + ax0_ax1_ax2_fused_3) // T.int64(16)) + v2 = T.axis.spatial(T.int64(4096), ax3_0 * T.int64(16) + (ax0_ax1_ax2_fused_0 * T.int64(32) + ax0_ax1_ax2_fused_1 * T.int64(4) + ax0_ax1_ax2_fused_2 * T.int64(2) + ax0_ax1_ax2_fused_3) % T.int64(16)) + T.reads(A[v0, v1, v2]) + T.writes(A_reindex_shared[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 8, 2]]}) + A_reindex_shared[v0, v1, v2] = A[v0, v1, v2] + for ax0_ax1_ax2_fused_0 in T.thread_binding(T.int64(16), thread="threadIdx.y"): + for ax0_ax1_ax2_fused_1 in T.thread_binding(T.int64(8), thread="threadIdx.x"): + for ax0_ax1_ax2_fused_2 in range(T.int64(4)): + for ax0_ax1_ax2_fused_3 in T.vectorized(T.int64(2)): + with T.block("var_decode_intermediate_reindex_shared"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial(T.int64(4096), ax0_ax2_0_fused * T.int64(64) + (ax0_ax1_ax2_fused_0 * T.int64(64) + ax0_ax1_ax2_fused_1 * T.int64(8) + ax0_ax1_ax2_fused_2 * T.int64(2) + ax0_ax1_ax2_fused_3) // T.int64(16)) + v2 = T.axis.spatial(T.int64(4096), ax3_0 * T.int64(16) + (ax0_ax1_ax2_fused_0 * T.int64(64) + ax0_ax1_ax2_fused_1 * T.int64(8) + ax0_ax1_ax2_fused_2 * T.int64(2) + ax0_ax1_ax2_fused_3) % T.int64(16)) + T.reads(W[v2 // T.int64(8), v1], S[v2 // T.int64(32), v1]) + T.writes(var_decode_intermediate_reindex_shared[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 8, 2]]}) + var_decode_intermediate_reindex_shared[v0, v1, v2] = T.Cast("float32", T.bitwise_and(T.shift_right(W[v2 // T.int64(8), v1], T.Cast("uint32", v2 % T.int64(8) * T.int64(4))), T.uint32(15))) * T.reinterpret("float32", T.shift_left(T.bitwise_and(S[v2 // T.int64(32), v1], T.uint32(65535)), T.uint32(16))) + T.reinterpret("float32", T.shift_left(T.bitwise_and(T.shift_right(S[v2 // T.int64(32), v1], T.uint32(16)), T.uint32(65535)), T.uint32(16))) + for ax3_1, ax1_3, ax2_3_0 in T.grid(T.int64(16), T.int64(4), T.int64(2)): + for ax2_3_1 in T.vectorized(T.int64(2)): + with T.block("matmul_update"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial(T.int64(32), ax1_0 * T.int64(32) + ax1_1 * T.int64(32) + ax1_2 * T.int64(4) + ax1_3) + v2 = T.axis.spatial(T.int64(4096), ax0_ax2_0_fused * T.int64(64) + ax2_1 * T.int64(64) + ax2_2 * T.int64(4) + ax2_3_0 * T.int64(2) + ax2_3_1) + v3 = T.axis.reduce(T.int64(4096), ax3_0 * T.int64(16) + ax3_1) + T.reads(var_matmul_intermediate_reindex_local[T.int64(0), v1, v2], A_reindex_shared[T.int64(0), v1, v3], var_decode_intermediate_reindex_shared[T.int64(0), v2, v3]) + T.writes(var_matmul_intermediate_reindex_local[T.int64(0), v1, v2]) + var_matmul_intermediate_reindex_local[T.int64(0), v1, v2] = var_matmul_intermediate_reindex_local[T.int64(0), v1, v2] + A_reindex_shared[T.int64(0), v1, v3] * var_decode_intermediate_reindex_shared[T.int64(0), v2, v3] + for ax0, ax1, ax2_0 in T.grid(T.int64(1), T.int64(4), T.int64(2)): + for ax2_1_1 in T.vectorized(T.int64(2)): + with T.block("var_matmul_intermediate_reindex_local"): + v0 = T.axis.spatial(T.int64(1), ax0) + v1 = T.axis.spatial(T.int64(32), ax1_2 * T.int64(4) + ax1) + v2 = T.axis.spatial(T.int64(4096), ax0_ax2_0_fused * T.int64(64) + ax2_2 * T.int64(4) + ax2_0 * T.int64(2) + ax2_1_1) + T.reads(C[T.int64(0), v1, v2], var_matmul_intermediate_reindex_local[v0, v1, v2]) + T.writes(Out[T.int64(0), v1, v2]) + Out[T.int64(0), v1, v2] = C[T.int64(0), v1, v2] + var_matmul_intermediate_reindex_local[v0, v1, v2] + + # fmt: on + + +@tvm.testing.requires_maca +class TestSkipGEMVMACA(MACABeforeAfter): + # fmt: off + + @T.prim_func + def before(W: T.Buffer((T.int64(512), T.int64(4096)), "uint32"), S: T.Buffer((T.int64(128), T.int64(4096)), "uint32"), A: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float32"), C: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float32"), Out: T.Buffer((T.int64(1), T.int64(1), T.int64(4096)), "float32")): + T.func_attr({"tir.noalias": T.bool(True)}) + var_decode_intermediate = T.alloc_buffer((T.int64(4096), T.int64(4096))) + var_matmul_intermediate = T.alloc_buffer((T.int64(1), T.int64(1), T.int64(4096))) + for i, j in T.grid(T.int64(4096), T.int64(4096)): + with T.block("decode"): + v_i, v_j = T.axis.remap("SS", [i, j]) + T.reads(W[v_i // T.int64(8), v_j], S[v_i // T.int64(32), v_j]) + T.writes(var_decode_intermediate[v_i, v_j]) + var_decode_intermediate[v_i, v_j] = T.Cast("float32", T.bitwise_and(T.shift_right(W[v_i // T.int64(8), v_j], T.Cast("uint32", v_i % T.int64(8) * T.int64(4))), T.uint32(15))) * T.reinterpret("float32", T.shift_left(T.bitwise_and(S[v_i // T.int64(32), v_j], T.uint32(65535)), T.uint32(16))) + T.reinterpret("float32", T.shift_left(T.bitwise_and(T.shift_right(S[v_i // T.int64(32), v_j], T.uint32(16)), T.uint32(65535)), T.uint32(16))) + for i0, i1, i2, k in T.grid(T.int64(1), T.int64(1), T.int64(4096), T.int64(4096)): + with T.block("matmul"): + v_i0, v_i1, v_i2, v_k = T.axis.remap("SSSR", [i0, i1, i2, k]) + T.reads(A[v_i0, v_i1, v_k], var_decode_intermediate[v_k, v_i2]) + T.writes(var_matmul_intermediate[v_i0, v_i1, v_i2]) + with T.init(): + var_matmul_intermediate[v_i0, v_i1, v_i2] = T.float32(0) + var_matmul_intermediate[v_i0, v_i1, v_i2] = var_matmul_intermediate[v_i0, v_i1, v_i2] + A[v_i0, v_i1, v_k] * var_decode_intermediate[v_k, v_i2] + for ax0, ax1, ax2 in T.grid(T.int64(1), T.int64(1), T.int64(4096)): + with T.block("T_add"): + v_ax0, v_ax1, v_ax2 = T.axis.remap("SSS", [ax0, ax1, ax2]) + T.reads(C[v_ax0, v_ax1, v_ax2], var_matmul_intermediate[v_ax0, v_ax1, v_ax2]) + T.writes(Out[v_ax0, v_ax1, v_ax2]) + Out[v_ax0, v_ax1, v_ax2] = C[v_ax0, v_ax1, v_ax2] + var_matmul_intermediate[v_ax0, v_ax1, v_ax2] + + # fmt: on + + expected = before + + +@tvm.testing.requires_maca +class TestOutputFP32MACA(MACABeforeAfter): + # fmt: off + + @T.prim_func + def before(lv13: T.Buffer((T.int64(4096), T.int64(512)), "uint32"), lv14: T.Buffer((T.int64(4096), T.int64(128)), "float16"), p_lv48: T.handle, lv13_1: T.Buffer((T.int64(4096),), "float16"), p_lv3: T.handle, p_output0: T.handle): + T.func_attr({"tir.noalias": T.bool(True)}) + n = T.int64() + lv48 = T.match_buffer(p_lv48, (T.int64(1), n, T.int64(4096)), "float16") + lv3 = T.match_buffer(p_lv3, (T.int64(1), n, T.int64(4096)), "float16") + p_output0_intermediate = T.match_buffer(p_output0, (T.int64(1), n, T.int64(4096)), "float16") + # with T.block("root"): + p_output0_intermediate_1 = T.alloc_buffer((T.int64(4096), T.int64(4096)), "float16") + var_matmul_intermediate = T.alloc_buffer((T.int64(1), n, T.int64(4096))) + var_compute_intermediate = T.alloc_buffer((T.int64(4096),)) + var_T_add_intermediate = T.alloc_buffer((T.int64(1), n, T.int64(4096))) + var_compute_intermediate_1 = T.alloc_buffer((T.int64(1), n, T.int64(4096)), "float16") + for i, j in T.grid(T.int64(4096), T.int64(4096)): + with T.block("decode"): + v_i, v_j = T.axis.remap("SS", [i, j]) + T.reads(lv13[v_i, v_j // T.int64(8)], lv14[v_i, v_j // T.int64(32)]) + T.writes(p_output0_intermediate_1[v_i, v_j]) + p_output0_intermediate_1[v_i, v_j] = (T.Cast("float16", T.bitwise_and(T.shift_right(lv13[v_i, v_j // T.int64(8)], T.Cast("uint32", v_j % T.int64(8)) * T.uint32(4)), T.uint32(15))) - T.float16(7)) * lv14[v_i, v_j // T.int64(32)] + for i0, i1, i2, k in T.grid(T.int64(1), n, T.int64(4096), T.int64(4096)): + with T.block("matmul"): + v_i0, v_i1, v_i2, v_k = T.axis.remap("SSSR", [i0, i1, i2, k]) + T.reads(lv48[v_i0, v_i1, v_k], p_output0_intermediate_1[v_k, v_i2]) + T.writes(var_matmul_intermediate[v_i0, v_i1, v_i2]) + with T.init(): + var_matmul_intermediate[v_i0, v_i1, v_i2] = T.float32(0) + var_matmul_intermediate[v_i0, v_i1, v_i2] = var_matmul_intermediate[v_i0, v_i1, v_i2] + T.Cast("float32", lv48[v_i0, v_i1, v_k]) * T.Cast("float32", p_output0_intermediate_1[v_k, v_i2]) + for i0 in range(T.int64(4096)): + with T.block("compute"): + v_i0 = T.axis.spatial(T.int64(4096), i0) + T.reads(lv13_1[v_i0]) + T.writes(var_compute_intermediate[v_i0]) + var_compute_intermediate[v_i0] = T.Cast("float32", lv13_1[v_i0]) + for ax0, ax1, ax2 in T.grid(T.int64(1), n, T.int64(4096)): + with T.block("T_add"): + v_ax0, v_ax1, v_ax2 = T.axis.remap("SSS", [ax0, ax1, ax2]) + T.reads(var_matmul_intermediate[v_ax0, v_ax1, v_ax2], var_compute_intermediate[v_ax2]) + T.writes(var_T_add_intermediate[v_ax0, v_ax1, v_ax2]) + var_T_add_intermediate[v_ax0, v_ax1, v_ax2] = var_matmul_intermediate[v_ax0, v_ax1, v_ax2] + var_compute_intermediate[v_ax2] + for i0, i1, i2 in T.grid(T.int64(1), n, T.int64(4096)): + with T.block("compute_1"): + v_i0, v_i1, v_i2 = T.axis.remap("SSS", [i0, i1, i2]) + T.reads(var_T_add_intermediate[v_i0, v_i1, v_i2]) + T.writes(var_compute_intermediate_1[v_i0, v_i1, v_i2]) + var_compute_intermediate_1[v_i0, v_i1, v_i2] = T.Cast("float16", var_T_add_intermediate[v_i0, v_i1, v_i2]) + for ax0, ax1, ax2 in T.grid(T.int64(1), n, T.int64(4096)): + with T.block("T_add_1"): + v_ax0, v_ax1, v_ax2 = T.axis.remap("SSS", [ax0, ax1, ax2]) + T.reads(var_compute_intermediate_1[v_ax0, v_ax1, v_ax2], lv3[v_ax0, v_ax1, v_ax2]) + T.writes(p_output0_intermediate[v_ax0, v_ax1, v_ax2]) + p_output0_intermediate[v_ax0, v_ax1, v_ax2] = var_compute_intermediate_1[v_ax0, v_ax1, v_ax2] + lv3[v_ax0, v_ax1, v_ax2] + + @T.prim_func + def expected(lv13: T.Buffer((T.int64(4096), T.int64(512)), "uint32"), lv14: T.Buffer((T.int64(4096), T.int64(128)), "float16"), p_lv48: T.handle, lv13_1: T.Buffer((T.int64(4096),), "float16"), p_lv3: T.handle, p_output0: T.handle): + T.func_attr({"global_symbol": "before", "tir.is_scheduled": True, "tir.noalias": T.bool(True)}) + n = T.int64() + lv48 = T.match_buffer(p_lv48, (T.int64(1), n, T.int64(4096)), "float16") + lv3 = T.match_buffer(p_lv3, (T.int64(1), n, T.int64(4096)), "float16") + p_output0_intermediate = T.match_buffer(p_output0, (T.int64(1), n, T.int64(4096)), "float16") + # with T.block("root"): + lv48_reindex_pad_shared_dyn = T.alloc_buffer((T.int64(1), (n + T.int64(63)) // T.int64(64) * T.int64(64), T.int64(4096)), "float16", scope="shared.dyn") + p_output0_intermediate_1_reindex_shared_dyn = T.alloc_buffer((T.int64(1), T.int64(4096), T.int64(4096)), "float16", scope="shared.dyn") + lv48_reindex_pad_shared_dyn_wmma_matrix_a = T.alloc_buffer((T.int64(1), (n + T.int64(63)) // T.int64(64) * T.int64(64), T.int64(4096)), "float16", scope="wmma.matrix_a") + p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b = T.alloc_buffer((T.int64(1), T.int64(4096), T.int64(4096)), "float16", scope="wmma.matrix_b") + var_matmul_intermediate_reindex_pad_shared_dyn = T.alloc_buffer((T.int64(1), (n + T.int64(63)) // T.int64(64) * T.int64(64), T.int64(4096)), scope="shared.dyn") + var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator = T.alloc_buffer((T.int64(1), (n + T.int64(63)) // T.int64(64) * T.int64(64), T.int64(4096)), scope="wmma.accumulator") + for ax0 in T.thread_binding(T.int64(1), thread="blockIdx.z"): + for ax1_0_0_ax2_0_0_fused in T.thread_binding((n + T.int64(63)) // T.int64(64), thread="blockIdx.x"): + for ax1_0_1_ax2_0_1_fused in T.thread_binding(T.int64(64), thread="blockIdx.y"): + for ax2_0_2_ax1_0_2_fused in T.thread_binding(T.int64(4), thread="threadIdx.y"): + for ax1_0_3_init, ax2_0_3_init in T.grid(T.int64(2), T.int64(2)): + with T.block("matmul_o_init"): + v0_o = T.axis.spatial(T.int64(1), ax0) + v1_o = T.axis.spatial((n + T.int64(63)) // T.int64(64) * T.int64(4), ax1_0_0_ax2_0_0_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(2) + ax1_0_3_init) + v2_o = T.axis.spatial(T.int64(256), ax1_0_1_ax2_0_1_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused // T.int64(2) * T.int64(2) + ax2_0_3_init) + T.reads() + T.writes(var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + with T.block("matmul_init_o"): + v1_i_init_o = T.axis.spatial(T.int64(1), T.int64(0)) + v2_i_init_o = T.axis.spatial(T.int64(1), T.int64(0)) + T.reads() + T.writes(var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + C = T.match_buffer(var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // T.int64(16) * (C.strides[0] // T.int64(16)) + C.elem_offset % C.strides[0] // T.int64(16), T.float32(0.0)) + for ax3_0_0 in T.serial(T.int64(128), annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): + for ax0_ax1_fused_0 in range(T.int64(2)): + for ax0_ax1_fused_1 in T.thread_binding(T.int64(4), thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(T.int64(64), thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(T.int64(4)): + with T.block("lv48_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial((n + T.int64(63)) // T.int64(64) * T.int64(64), ax1_0_0_ax2_0_0_fused * T.int64(64) + (ax0_ax1_fused_0 * T.int64(1024) + ax0_ax1_fused_1 * T.int64(256) + ax0_ax1_fused_2 * T.int64(4) + ax0_ax1_fused_3) // T.int64(32)) + v2 = T.axis.spatial(T.int64(4096), ax3_0_0 * T.int64(32) + (ax0_ax1_fused_0 * T.int64(1024) + ax0_ax1_fused_1 * T.int64(256) + ax0_ax1_fused_2 * T.int64(4) + ax0_ax1_fused_3) % T.int64(32)) + T.reads(lv48[v0, v1, v2]) + T.writes(lv48_reindex_pad_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 8]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + lv48_reindex_pad_shared_dyn[v0, v1, v2] = T.if_then_else(v1 < n, lv48[v0, v1, v2], T.float16(0.0)) + for ax0_ax1_fused_0 in range(T.int64(2)): + for ax0_ax1_fused_1 in T.thread_binding(T.int64(4), thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(T.int64(64), thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(T.int64(4)): + with T.block("p_output0_intermediate_1_reindex_shared.dyn"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial(T.int64(4096), ax1_0_1_ax2_0_1_fused * T.int64(64) + (ax0_ax1_fused_0 * T.int64(1024) + ax0_ax1_fused_1 * T.int64(256) + ax0_ax1_fused_2 * T.int64(4) + ax0_ax1_fused_3) // T.int64(32)) + v2 = T.axis.spatial(T.int64(4096), ax3_0_0 * T.int64(32) + (ax0_ax1_fused_0 * T.int64(1024) + ax0_ax1_fused_1 * T.int64(256) + ax0_ax1_fused_2 * T.int64(4) + ax0_ax1_fused_3) % T.int64(32)) + T.reads(lv13[v2, v1 // T.int64(8)], lv14[v2, v1 // T.int64(32)]) + T.writes(p_output0_intermediate_1_reindex_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 8]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + p_output0_intermediate_1_reindex_shared_dyn[v0, v1, v2] = (T.Cast("float16", T.bitwise_and(T.shift_right(lv13[v2, v1 // T.int64(8)], T.Cast("uint32", v1 % T.int64(8)) * T.uint32(4)), T.uint32(15))) - T.float16(7.0)) * lv14[v2, v1 // T.int64(32)] + for ax3_0_1 in T.serial(T.int64(2), annotations={"software_pipeline_order": [0, 1, 2], "software_pipeline_stage": [0, 0, 1]}): + for ax0_0 in T.unroll(T.int64(2)): + for ax1_0 in T.unroll(T.int64(1)): + with T.block("lv48_reindex_pad_shared.dyn_wmma.matrix_a_o"): + v0_o = T.axis.spatial(T.int64(1), T.int64(0)) + v1_o = T.axis.spatial(T.int64(4) * ((n + T.int64(63)) // T.int64(64)), ax1_0_0_ax2_0_0_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(2) + ax0_0) + v2_o = T.axis.spatial(T.int64(256), ax3_0_0 * T.int64(2) + ax3_0_1 + ax1_0) + T.reads(lv48_reindex_pad_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + T.writes(lv48_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + A = T.match_buffer(lv48_reindex_pad_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(lv48_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_a", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // T.int64(16) * (C.strides[0] // T.int64(16)) + C.elem_offset % C.strides[0] // T.int64(16), T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * T.int64(16), 1), A.strides[0], "row_major") + for ax0_0 in T.unroll(T.int64(2)): + for ax1_0 in T.unroll(T.int64(1)): + with T.block("p_output0_intermediate_1_reindex_shared.dyn_wmma.matrix_b_o"): + v0_o = T.axis.spatial(T.int64(1), T.int64(0)) + v1_o = T.axis.spatial(T.int64(256), ax1_0_1_ax2_0_1_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused // T.int64(2) * T.int64(2) + ax0_0) + v2_o = T.axis.spatial(T.int64(256), ax3_0_0 * T.int64(2) + ax3_0_1 + ax1_0) + T.reads(p_output0_intermediate_1_reindex_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + T.writes(p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + A = T.match_buffer(p_output0_intermediate_1_reindex_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_b", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // T.int64(16) * (C.strides[0] // T.int64(16)) + C.elem_offset % C.strides[0] // T.int64(16), T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * T.int64(16), 1), A.strides[0], "col_major") + for ax1_0_3, ax2_0_3 in T.grid(T.int64(2), T.int64(2)): + with T.block("matmul_o_update"): + v0_o = T.axis.spatial(T.int64(1), ax0) + v1_o = T.axis.spatial((n + T.int64(63)) // T.int64(64) * T.int64(4), ax1_0_0_ax2_0_0_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(2) + ax1_0_3) + v2_o = T.axis.spatial(T.int64(256), ax1_0_1_ax2_0_1_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused // T.int64(2) * T.int64(2) + ax2_0_3) + v3_o = T.axis.reduce(T.int64(256), ax3_0_0 * T.int64(2) + ax3_0_1) + T.reads(var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], lv48_reindex_pad_shared_dyn_wmma_matrix_a[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)], p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[T.int64(0), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)]) + T.writes(var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + with T.block("matmul_o"): + v1_i_o = T.axis.spatial(T.int64(1), T.int64(0)) + v2_i_o = T.axis.spatial(T.int64(1), T.int64(0)) + v3_i_o = T.axis.reduce(T.int64(1), T.int64(0)) + T.reads(var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], lv48_reindex_pad_shared_dyn_wmma_matrix_a[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)], p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[T.int64(0), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)]) + T.writes(var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + A = T.match_buffer(lv48_reindex_pad_shared_dyn_wmma_matrix_a[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("A_s0", "A_s1"), scope="wmma.matrix_a", offset_factor=16) + B = T.match_buffer(p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[T.int64(0), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("B_s0", "B_s1"), scope="wmma.matrix_b", offset_factor=16) + C = T.match_buffer(var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // T.int64(16) * (C.strides[0] // T.int64(16)) + C.elem_offset % C.strides[0] // T.int64(16), A.data, A.elem_offset // A.strides[0] // T.int64(16) * (A.strides[0] // T.int64(16)) + A.elem_offset % A.strides[0] // T.int64(16), B.data, B.elem_offset // B.strides[0] // T.int64(16) * (B.strides[0] // T.int64(16)) + B.elem_offset % B.strides[0] // T.int64(16), C.data, C.elem_offset // C.strides[0] // T.int64(16) * (C.strides[0] // T.int64(16)) + C.elem_offset % C.strides[0] // T.int64(16)) + for ax0_0, ax1_0 in T.grid(T.int64(2), T.int64(2)): + with T.block("var_matmul_intermediate_reindex_pad_shared.dyn_wmma.accumulator_o"): + v0_o = T.axis.spatial(T.int64(1), T.int64(0)) + v1_o = T.axis.spatial(T.int64(4) * ((n + T.int64(63)) // T.int64(64)), ax1_0_0_ax2_0_0_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(2) + ax0_0) + v2_o = T.axis.spatial(T.int64(256), ax1_0_1_ax2_0_1_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused // T.int64(2) * T.int64(2) + ax1_0) + T.reads(var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + T.writes(var_matmul_intermediate_reindex_pad_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + A = T.match_buffer(var_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), strides=("A_s0", "A_s1"), scope="wmma.accumulator", offset_factor=16) + C = T.match_buffer(var_matmul_intermediate_reindex_pad_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), strides=("C_s0", "C_s1"), scope="shared.dyn", offset_factor=16) + T.tvm_store_matrix_sync(A.data, 16, 16, 16, A.elem_offset // A.strides[0] // T.int64(16) * (A.strides[0] // T.int64(16)) + A.elem_offset % A.strides[0] // T.int64(16), T.tvm_access_ptr(T.type_annotation("float32"), C.data, C.elem_offset, C.strides[0] * T.int64(16), 2), C.strides[0], "row_major") + for ax0_ax1_fused_0 in range(T.int64(4)): + for ax0_ax1_fused_1 in T.thread_binding(T.int64(64), thread="threadIdx.x"): + for ax0_ax1_fused_2 in T.vectorized(T.int64(4)): + with T.block("var_matmul_intermediate_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial((n + T.int64(63)) // T.int64(64) * T.int64(64), ax1_0_0_ax2_0_0_fused * T.int64(64) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(32) + (ax0_ax1_fused_0 * T.int64(256) + ax0_ax1_fused_1 * T.int64(4) + ax0_ax1_fused_2) // T.int64(32)) + v2 = T.axis.spatial(T.int64(4096), ax1_0_1_ax2_0_1_fused * T.int64(64) + ax2_0_2_ax1_0_2_fused // T.int64(2) * T.int64(32) + (ax0_ax1_fused_0 * T.int64(256) + ax0_ax1_fused_1 * T.int64(4) + ax0_ax1_fused_2) % T.int64(32)) + T.where(ax1_0_0_ax2_0_0_fused * T.int64(64) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(32) + ((ax0_ax1_fused_0 * T.int64(64) + ax0_ax1_fused_1) * T.int64(4) + ax0_ax1_fused_2) // T.int64(32) < n) + T.reads(var_matmul_intermediate_reindex_pad_shared_dyn[v0, v1, v2], lv13_1[v2], lv3[T.int64(0), v1, v2]) + T.writes(p_output0_intermediate[T.int64(0), v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 4]]}) + p_output0_intermediate[T.int64(0), v1, v2] = T.Cast("float16", var_matmul_intermediate_reindex_pad_shared_dyn[v0, v1, v2] + T.Cast("float32", lv13_1[v2])) + lv3[T.int64(0), v1, v2] + + # fmt: on + + +@tvm.testing.requires_maca +class TestInlineConsumerChainMACA(MACABeforeAfter): + # fmt: off + @T.prim_func(private=True) + def before(p_lv26: T.handle, lv9: T.Buffer((T.int64(2048), T.int64(2048)), "float16"), p_lv52: T.handle, p_output0: T.handle): + T.func_attr({"tir.noalias": T.bool(True)}) + n = T.int64() + lv26 = T.match_buffer(p_lv26, (n, T.int64(2048)), "float16") + lv52 = T.match_buffer(p_lv52, (T.int64(1), n, T.int64(2048))) + var_T_multiply_intermediate = T.match_buffer(p_output0, (n, T.int64(2048)), "float16") + # with T.block("root"): + var_NT_matmul_intermediate = T.alloc_buffer((n, T.int64(2048)), "float16") + compute = T.alloc_buffer((n, T.int64(2048)), "float16") + var_T_multiply_intermediate_1 = T.alloc_buffer((n, T.int64(2048)), "float16") + var_T_squeeze_intermediate = T.alloc_buffer((n, T.int64(2048))) + var_compute_intermediate = T.alloc_buffer((n, T.int64(2048)), "float16") + for i0, i1, k in T.grid(n, T.int64(2048), T.int64(2048)): + with T.block("NT_matmul"): + v_i0, v_i1, v_k = T.axis.remap("SSR", [i0, i1, k]) + T.reads(lv26[v_i0, v_k], lv9[v_i1, v_k]) + T.writes(var_NT_matmul_intermediate[v_i0, v_i1]) + with T.init(): + var_NT_matmul_intermediate[v_i0, v_i1] = T.float16(0) + var_NT_matmul_intermediate[v_i0, v_i1] = var_NT_matmul_intermediate[v_i0, v_i1] + lv26[v_i0, v_k] * lv9[v_i1, v_k] + for i0, i1 in T.grid(n, T.int64(2048)): + with T.block("compute"): + v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) + T.reads(var_NT_matmul_intermediate[v_i0, v_i1]) + T.writes(compute[v_i0, v_i1]) + compute[v_i0, v_i1] = T.sigmoid(var_NT_matmul_intermediate[v_i0, v_i1]) + for ax0, ax1 in T.grid(n, T.int64(2048)): + with T.block("T_multiply"): + v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1]) + T.reads(var_NT_matmul_intermediate[v_ax0, v_ax1], compute[v_ax0, v_ax1]) + T.writes(var_T_multiply_intermediate_1[v_ax0, v_ax1]) + var_T_multiply_intermediate_1[v_ax0, v_ax1] = var_NT_matmul_intermediate[v_ax0, v_ax1] * compute[v_ax0, v_ax1] + for ax0, ax1 in T.grid(n, T.int64(2048)): + with T.block("T_squeeze"): + v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1]) + T.reads(lv52[T.int64(0), v_ax0, v_ax1]) + T.writes(var_T_squeeze_intermediate[v_ax0, v_ax1]) + var_T_squeeze_intermediate[v_ax0, v_ax1] = lv52[T.int64(0), v_ax0, v_ax1] + for i0, i1 in T.grid(n, T.int64(2048)): + with T.block("compute_1"): + v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) + T.reads(var_T_squeeze_intermediate[v_i0, v_i1]) + T.writes(var_compute_intermediate[v_i0, v_i1]) + var_compute_intermediate[v_i0, v_i1] = T.Cast("float16", var_T_squeeze_intermediate[v_i0, v_i1]) + for ax0, ax1 in T.grid(n, T.int64(2048)): + with T.block("T_multiply_1"): + v_ax0, v_ax1 = T.axis.remap("SS", [ax0, ax1]) + T.reads(var_compute_intermediate[v_ax0, v_ax1], var_T_multiply_intermediate_1[v_ax0, v_ax1]) + T.writes(var_T_multiply_intermediate[v_ax0, v_ax1]) + var_T_multiply_intermediate[v_ax0, v_ax1] = var_compute_intermediate[v_ax0, v_ax1] * var_T_multiply_intermediate_1[v_ax0, v_ax1] + + @T.prim_func + def expected(p_lv26: T.handle, lv9: T.Buffer((T.int64(2048), T.int64(2048)), "float16"), p_lv52: T.handle, p_output0: T.handle): + T.func_attr({"tir.is_scheduled": True, "tir.noalias": T.bool(True)}) + n = T.int64() + lv26 = T.match_buffer(p_lv26, (n, T.int64(2048)), "float16") + lv52 = T.match_buffer(p_lv52, (T.int64(1), n, T.int64(2048))) + var_T_multiply_intermediate = T.match_buffer(p_output0, (n, T.int64(2048)), "float16") + # with T.block("root"): + lv26_reindex_pad_shared_dyn = T.alloc_buffer((T.int64(1), (n + T.int64(63)) // T.int64(64) * T.int64(64), T.int64(2048)), "float16", scope="shared.dyn") + lv9_reindex_shared_dyn = T.alloc_buffer((T.int64(1), T.int64(2048), T.int64(2048)), "float16", scope="shared.dyn") + lv26_reindex_pad_shared_dyn_wmma_matrix_a = T.alloc_buffer((T.int64(1), (n + T.int64(63)) // T.int64(64) * T.int64(64), T.int64(2048)), "float16", scope="wmma.matrix_a") + lv9_reindex_shared_dyn_wmma_matrix_b = T.alloc_buffer((T.int64(1), T.int64(2048), T.int64(2048)), "float16", scope="wmma.matrix_b") + var_NT_matmul_intermediate_reindex_pad_shared_dyn = T.alloc_buffer((T.int64(1), (n + T.int64(63)) // T.int64(64) * T.int64(64), T.int64(2048)), "float16", scope="shared.dyn") + var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator = T.alloc_buffer((T.int64(1), (n + T.int64(63)) // T.int64(64) * T.int64(64), T.int64(2048)), "float16", scope="wmma.accumulator") + for ax0 in T.thread_binding(T.int64(1), thread="blockIdx.z"): + for ax1_0_0_ax2_0_0_fused in T.thread_binding((n + T.int64(63)) // T.int64(64), thread="blockIdx.x"): + for ax1_0_1_ax2_0_1_fused in T.thread_binding(T.int64(32), thread="blockIdx.y"): + for ax2_0_2_ax1_0_2_fused in T.thread_binding(T.int64(4), thread="threadIdx.y"): + for ax1_0_3_init, ax2_0_3_init in T.grid(T.int64(2), T.int64(2)): + with T.block("NT_matmul_o_init"): + v0_o = T.axis.spatial(T.int64(1), ax0) + v1_o = T.axis.spatial((n + T.int64(63)) // T.int64(64) * T.int64(4), ax1_0_0_ax2_0_0_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(2) + ax1_0_3_init) + v2_o = T.axis.spatial(T.int64(128), ax1_0_1_ax2_0_1_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused // T.int64(2) * T.int64(2) + ax2_0_3_init) + T.reads() + T.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + with T.block("NT_matmul_init_o"): + v1_i_init_o = T.axis.spatial(T.int64(1), T.int64(0)) + v2_i_init_o = T.axis.spatial(T.int64(1), T.int64(0)) + T.reads() + T.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + C = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // T.int64(16) * (C.strides[0] // T.int64(16)) + C.elem_offset % C.strides[0] // T.int64(16), T.float32(0.0)) + for ax3_0_0 in T.serial(T.int64(64), annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): + for ax0_ax1_fused_0 in range(T.int64(2)): + for ax0_ax1_fused_1 in T.thread_binding(T.int64(4), thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(T.int64(64), thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(T.int64(4)): + with T.block("lv26_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial((n + T.int64(63)) // T.int64(64) * T.int64(64), ax1_0_0_ax2_0_0_fused * T.int64(64) + (ax0_ax1_fused_0 * T.int64(1024) + ax0_ax1_fused_1 * T.int64(256) + ax0_ax1_fused_2 * T.int64(4) + ax0_ax1_fused_3) // T.int64(32)) + v2 = T.axis.spatial(T.int64(2048), ax3_0_0 * T.int64(32) + (ax0_ax1_fused_0 * T.int64(1024) + ax0_ax1_fused_1 * T.int64(256) + ax0_ax1_fused_2 * T.int64(4) + ax0_ax1_fused_3) % T.int64(32)) + T.reads(lv26[v1, v2]) + T.writes(lv26_reindex_pad_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 8]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + lv26_reindex_pad_shared_dyn[v0, v1, v2] = T.if_then_else(v1 < n, lv26[v1, v2], T.float16(0.0)) + for ax0_ax1_fused_0 in range(T.int64(2)): + for ax0_ax1_fused_1 in T.thread_binding(T.int64(4), thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(T.int64(64), thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(T.int64(4)): + with T.block("lv9_reindex_shared.dyn"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial(T.int64(2048), ax1_0_1_ax2_0_1_fused * T.int64(64) + (ax0_ax1_fused_0 * T.int64(1024) + ax0_ax1_fused_1 * T.int64(256) + ax0_ax1_fused_2 * T.int64(4) + ax0_ax1_fused_3) // T.int64(32)) + v2 = T.axis.spatial(T.int64(2048), ax3_0_0 * T.int64(32) + (ax0_ax1_fused_0 * T.int64(1024) + ax0_ax1_fused_1 * T.int64(256) + ax0_ax1_fused_2 * T.int64(4) + ax0_ax1_fused_3) % T.int64(32)) + T.reads(lv9[v1, v2]) + T.writes(lv9_reindex_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 8]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + lv9_reindex_shared_dyn[v0, v1, v2] = lv9[v1, v2] + for ax3_0_1 in T.serial(T.int64(2), annotations={"software_pipeline_order": [0, 1, 2], "software_pipeline_stage": [0, 0, 1]}): + for ax0_0 in T.unroll(T.int64(2)): + for ax1_0 in T.unroll(T.int64(1)): + with T.block("lv26_reindex_pad_shared.dyn_wmma.matrix_a_o"): + v0_o = T.axis.spatial(T.int64(1), T.int64(0)) + v1_o = T.axis.spatial(T.int64(4) * ((n + T.int64(63)) // T.int64(64)), ax1_0_0_ax2_0_0_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(2) + ax0_0) + v2_o = T.axis.spatial(T.int64(128), ax3_0_0 * T.int64(2) + ax3_0_1 + ax1_0) + T.reads(lv26_reindex_pad_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + T.writes(lv26_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + A = T.match_buffer(lv26_reindex_pad_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(lv26_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_a", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // T.int64(16) * (C.strides[0] // T.int64(16)) + C.elem_offset % C.strides[0] // T.int64(16), T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * T.int64(16), 1), A.strides[0], "row_major") + for ax0_0 in T.unroll(T.int64(2)): + for ax1_0 in T.unroll(T.int64(1)): + with T.block("lv9_reindex_shared.dyn_wmma.matrix_b_o"): + v0_o = T.axis.spatial(T.int64(1), T.int64(0)) + v1_o = T.axis.spatial(T.int64(128), ax1_0_1_ax2_0_1_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused // T.int64(2) * T.int64(2) + ax0_0) + v2_o = T.axis.spatial(T.int64(128), ax3_0_0 * T.int64(2) + ax3_0_1 + ax1_0) + T.reads(lv9_reindex_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + T.writes(lv9_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + A = T.match_buffer(lv9_reindex_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(lv9_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_b", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // T.int64(16) * (C.strides[0] // T.int64(16)) + C.elem_offset % C.strides[0] // T.int64(16), T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * T.int64(16), 1), A.strides[0], "col_major") + for ax1_0_3, ax2_0_3 in T.grid(T.int64(2), T.int64(2)): + with T.block("NT_matmul_o_update"): + v0_o = T.axis.spatial(T.int64(1), ax0) + v1_o = T.axis.spatial((n + T.int64(63)) // T.int64(64) * T.int64(4), ax1_0_0_ax2_0_0_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(2) + ax1_0_3) + v2_o = T.axis.spatial(T.int64(128), ax1_0_1_ax2_0_1_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused // T.int64(2) * T.int64(2) + ax2_0_3) + v3_o = T.axis.reduce(T.int64(128), ax3_0_0 * T.int64(2) + ax3_0_1) + T.reads(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], lv26_reindex_pad_shared_dyn_wmma_matrix_a[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)], lv9_reindex_shared_dyn_wmma_matrix_b[T.int64(0), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)]) + T.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + with T.block("NT_matmul_o"): + v1_i_o = T.axis.spatial(T.int64(1), T.int64(0)) + v2_i_o = T.axis.spatial(T.int64(1), T.int64(0)) + v3_i_o = T.axis.reduce(T.int64(1), T.int64(0)) + T.reads(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], lv26_reindex_pad_shared_dyn_wmma_matrix_a[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)], lv9_reindex_shared_dyn_wmma_matrix_b[T.int64(0), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)]) + T.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + A = T.match_buffer(lv26_reindex_pad_shared_dyn_wmma_matrix_a[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("A_s0", "A_s1"), scope="wmma.matrix_a", offset_factor=16) + B = T.match_buffer(lv9_reindex_shared_dyn_wmma_matrix_b[T.int64(0), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16), v3_o * T.int64(16):v3_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("B_s0", "B_s1"), scope="wmma.matrix_b", offset_factor=16) + C = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[T.int64(0), v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // T.int64(16) * (C.strides[0] // T.int64(16)) + C.elem_offset % C.strides[0] // T.int64(16), A.data, A.elem_offset // A.strides[0] // T.int64(16) * (A.strides[0] // T.int64(16)) + A.elem_offset % A.strides[0] // T.int64(16), B.data, B.elem_offset // B.strides[0] // T.int64(16) * (B.strides[0] // T.int64(16)) + B.elem_offset % B.strides[0] // T.int64(16), C.data, C.elem_offset // C.strides[0] // T.int64(16) * (C.strides[0] // T.int64(16)) + C.elem_offset % C.strides[0] // T.int64(16)) + for ax0_0, ax1_0 in T.grid(T.int64(2), T.int64(2)): + with T.block("var_NT_matmul_intermediate_reindex_pad_shared.dyn_wmma.accumulator_o"): + v0_o = T.axis.spatial(T.int64(1), T.int64(0)) + v1_o = T.axis.spatial(T.int64(4) * ((n + T.int64(63)) // T.int64(64)), ax1_0_0_ax2_0_0_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(2) + ax0_0) + v2_o = T.axis.spatial(T.int64(128), ax1_0_1_ax2_0_1_fused * T.int64(4) + ax2_0_2_ax1_0_2_fused // T.int64(2) * T.int64(2) + ax1_0) + T.reads(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + T.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)]) + A = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("A_s0", "A_s1"), scope="wmma.accumulator", offset_factor=16) + C = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0_o, v1_o * T.int64(16):v1_o * T.int64(16) + T.int64(16), v2_o * T.int64(16):v2_o * T.int64(16) + T.int64(16)], (T.int64(16), T.int64(16)), "float16", strides=("C_s0", "C_s1"), scope="shared.dyn", offset_factor=16) + T.tvm_store_matrix_sync(A.data, 16, 16, 16, A.elem_offset // A.strides[0] // T.int64(16) * (A.strides[0] // T.int64(16)) + A.elem_offset % A.strides[0] // T.int64(16), T.tvm_access_ptr(T.type_annotation("float16"), C.data, C.elem_offset, C.strides[0] * T.int64(16), 2), C.strides[0], "row_major") + for ax0_ax1_fused_0 in range(T.int64(4)): + for ax0_ax1_fused_1 in T.thread_binding(T.int64(64), thread="threadIdx.x"): + for ax0_ax1_fused_2 in T.vectorized(T.int64(4)): + with T.block("var_NT_matmul_intermediate_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(T.int64(1), T.int64(0)) + v1 = T.axis.spatial((n + T.int64(63)) // T.int64(64) * T.int64(64), ax1_0_0_ax2_0_0_fused * T.int64(64) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(32) + (ax0_ax1_fused_0 * T.int64(256) + ax0_ax1_fused_1 * T.int64(4) + ax0_ax1_fused_2) // T.int64(32)) + v2 = T.axis.spatial(T.int64(2048), ax1_0_1_ax2_0_1_fused * T.int64(64) + ax2_0_2_ax1_0_2_fused // T.int64(2) * T.int64(32) + (ax0_ax1_fused_0 * T.int64(256) + ax0_ax1_fused_1 * T.int64(4) + ax0_ax1_fused_2) % T.int64(32)) + T.where(ax1_0_0_ax2_0_0_fused * T.int64(64) + ax2_0_2_ax1_0_2_fused % T.int64(2) * T.int64(32) + ((ax0_ax1_fused_0 * T.int64(64) + ax0_ax1_fused_1) * T.int64(4) + ax0_ax1_fused_2) // T.int64(32) < n) + T.reads(lv52[T.int64(0), v1, v2], var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0, v1, v2]) + T.writes(var_T_multiply_intermediate[v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 4]]}) + var_T_multiply_intermediate[v1, v2] = T.Cast("float16", lv52[T.int64(0), v1, v2]) * (var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0, v1, v2] * T.sigmoid(var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0, v1, v2])) + + if __name__ == "__main__": tvm.testing.main() diff --git a/tests/python/dlight/test_gpu_matmul_tensorize.py b/tests/python/dlight/test_gpu_matmul_tensorize.py index 261981c5e46c..947471ca7af9 100644 --- a/tests/python/dlight/test_gpu_matmul_tensorize.py +++ b/tests/python/dlight/test_gpu_matmul_tensorize.py @@ -977,5 +977,686 @@ def expected(B0: T.Buffer((28672, 512), "uint32"), B1: T.Buffer((28672, 128), "f C[v1, 0, v2] = C_reindex_pad_shared[v0, v1, v2] +class MACABeforeAfter(tvm.testing.CompareBeforeAfter): + @pytest.fixture + def transform(self): + def transform(mod): + with Target("maca"): + return dl.ApplyDefaultSchedule(dl.gpu.Matmul())(mod) + + return transform + + +@tvm.testing.requires_maca +class TestMatmulTensorizeMACA(MACABeforeAfter): + # fmt: off + + @T.prim_func + def before(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float16"), compute: T.Buffer((256, 256), "float16")): + T.func_attr({"global_symbol": "main", "tir.noalias": T.bool(True)}) + # with T.block("root"): + for i, j, k in T.grid(256, 256, 256): + with T.block("compute"): + v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) + T.reads(X[v_i, v_k], W[v_j, v_k]) + T.writes(compute[v_i, v_j]) + with T.init(): + compute[v_i, v_j] = T.float16(0) + compute[v_i, v_j] = compute[v_i, v_j] + X[v_i, v_k] * W[v_j, v_k] + + @T.prim_func + def expected(X: T.Buffer((256, 256), "float16"), W: T.Buffer((256, 256), "float16"), compute: T.Buffer((256, 256), "float16")): + T.func_attr({"tir.is_scheduled": True, "tir.noalias": T.bool(True)}) + # with T.block("root"): + X_reindex_shared_dyn = T.alloc_buffer((1, 256, 256), "float16", scope="shared.dyn") + W_reindex_shared_dyn = T.alloc_buffer((1, 256, 256), "float16", scope="shared.dyn") + X_reindex_shared_dyn_wmma_matrix_a = T.alloc_buffer((1, 256, 256), "float16", scope="wmma.matrix_a") + W_reindex_shared_dyn_wmma_matrix_b = T.alloc_buffer((1, 256, 256), "float16", scope="wmma.matrix_b") + compute_reindex_shared_dyn = T.alloc_buffer((1, 256, 256), "float16", scope="shared.dyn") + compute_reindex_shared_dyn_wmma_accumulator = T.alloc_buffer((1, 256, 256), "float16", scope="wmma.accumulator") + for ax0 in T.thread_binding(1, thread="blockIdx.z"): + for ax1_0_0_ax2_0_0_fused in T.thread_binding(4, thread="blockIdx.x"): + for ax1_0_1_ax2_0_1_fused in T.thread_binding(4, thread="blockIdx.y"): + for ax2_0_2_ax1_0_2_fused in T.thread_binding(4, thread="threadIdx.y"): + for ax1_0_3_init, ax2_0_3_init in T.grid(2, 2): + with T.block("compute_o_init"): + v0_o = T.axis.spatial(1, ax0) + v1_o = T.axis.spatial(16, ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax1_0_3_init) + v2_o = T.axis.spatial(16, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax2_0_3_init) + T.reads() + T.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + with T.block("compute_init_o"): + v1_i_init_o = T.axis.spatial(1, 0) + v2_i_init_o = T.axis.spatial(1, 0) + T.reads() + T.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + C = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.float32(0.0)) + for ax3_0_0 in T.serial(8, annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): + for ax0_ax1_fused_0 in range(2): + for ax0_ax1_fused_1 in T.thread_binding(4, thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(4): + with T.block("X_reindex_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(256, ax1_0_0_ax2_0_0_fused * 64 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) // 32) + v2 = T.axis.spatial(256, ax3_0_0 * 32 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) % 32) + T.reads(X[v1, v2]) + T.writes(X_reindex_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 8]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + X_reindex_shared_dyn[v0, v1, v2] = X[v1, v2] + for ax0_ax1_fused_0 in range(2): + for ax0_ax1_fused_1 in T.thread_binding(4, thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(4): + with T.block("W_reindex_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 64 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) // 32) + v2 = T.axis.spatial(256, ax3_0_0 * 32 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) % 32) + T.reads(W[v1, v2]) + T.writes(W_reindex_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 8]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + W_reindex_shared_dyn[v0, v1, v2] = W[v1, v2] + for ax3_0_1 in T.serial(2, annotations={"software_pipeline_order": [0, 1, 2], "software_pipeline_stage": [0, 0, 1]}): + for ax0_0 in T.unroll(2): + for ax1_0 in T.unroll(1): + with T.block("X_reindex_shared.dyn_wmma.matrix_a_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(16, ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax0_0) + v2_o = T.axis.spatial(16, ax3_0_0 * 2 + ax3_0_1 + ax1_0) + T.reads(X_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(X_reindex_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(X_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(X_reindex_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_a", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "row_major") + for ax0_0 in T.unroll(2): + for ax1_0 in T.unroll(1): + with T.block("W_reindex_shared.dyn_wmma.matrix_b_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(16, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax0_0) + v2_o = T.axis.spatial(16, ax3_0_0 * 2 + ax3_0_1 + ax1_0) + T.reads(W_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(W_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(W_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(W_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_b", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "col_major") + for ax1_0_3, ax2_0_3 in T.grid(2, 2): + with T.block("compute_o_update"): + v0_o = T.axis.spatial(1, ax0) + v1_o = T.axis.spatial(16, ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax1_0_3) + v2_o = T.axis.spatial(16, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax2_0_3) + v3_o = T.axis.reduce(16, ax3_0_0 * 2 + ax3_0_1) + T.reads(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], X_reindex_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], W_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) + T.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + with T.block("compute_o"): + v1_i_o = T.axis.spatial(1, 0) + v2_i_o = T.axis.spatial(1, 0) + v3_i_o = T.axis.reduce(1, 0) + T.reads(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], X_reindex_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], W_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) + T.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(X_reindex_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="wmma.matrix_a", offset_factor=16) + B = T.match_buffer(W_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "float16", strides=("B_s0", "B_s1"), scope="wmma.matrix_b", offset_factor=16) + C = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, A.data, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, B.data, B.elem_offset // B.strides[0] // 16 * (B.strides[0] // 16) + B.elem_offset % B.strides[0] // 16, C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16) + for ax0_0, ax1_0 in T.grid(2, 2): + with T.block("compute_reindex_shared.dyn_wmma.accumulator_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(16, ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax0_0) + v2_o = T.axis.spatial(16, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax1_0) + T.reads(compute_reindex_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(compute_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="wmma.accumulator", offset_factor=16) + C = T.match_buffer(compute_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="shared.dyn", offset_factor=16) + T.tvm_store_matrix_sync(A.data, 16, 16, 16, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), C.data, C.elem_offset, C.strides[0] * 16, 2), C.strides[0], "row_major") + for ax0_ax1_fused_0 in range(4): + for ax0_ax1_fused_1 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_2 in T.vectorized(4): + with T.block("compute_reindex_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(256, ax1_0_0_ax2_0_0_fused * 64 + ax2_0_2_ax1_0_2_fused % 2 * 32 + (ax0_ax1_fused_0 * 256 + ax0_ax1_fused_1 * 4 + ax0_ax1_fused_2) // 32) + v2 = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 64 + ax2_0_2_ax1_0_2_fused // 2 * 32 + (ax0_ax1_fused_0 * 256 + ax0_ax1_fused_1 * 4 + ax0_ax1_fused_2) % 32) + T.reads(compute_reindex_shared_dyn[v0, v1, v2]) + T.writes(compute[v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 4]]}) + compute[v1, v2] = compute_reindex_shared_dyn[v0, v1, v2] + # fmt: on + + +@tvm.testing.requires_maca +class TestMatmulTensorizeTooSmallMACA(MACABeforeAfter): + # fmt: off + + @T.prim_func + def before(var_X: T.handle, W: T.Buffer((15, 256), "float16"), var_compute: T.handle): + T.func_attr({"global_symbol": "main", "tir.noalias": T.bool(True)}) + m = T.int32() + X = T.match_buffer(var_X, (m, 256), "float16") + compute = T.match_buffer(var_compute, (m, 15)) + # with T.block("root"): + for i, j, k in T.grid(m, 15, 256): + with T.block("compute"): + v_i, v_j, v_k = T.axis.remap("SSR", [i, j, k]) + T.reads(X[v_i, v_k], W[v_j, v_k]) + T.writes(compute[v_i, v_j]) + with T.init(): + compute[v_i, v_j] = T.float32(0) + compute[v_i, v_j] = compute[v_i, v_j] + T.Cast("float32", X[v_i, v_k]) * T.Cast("float32", W[v_j, v_k]) + + @T.prim_func + def expected(var_X: T.handle, W: T.Buffer((15, 256), "float16"), var_compute: T.handle): + T.func_attr({"tir.is_scheduled": True, "tir.noalias": T.bool(True)}) + m = T.int32() + X = T.match_buffer(var_X, (m, 256), "float16") + compute = T.match_buffer(var_compute, (m, 15)) + # with T.block("root"): + compute_reindex_pad_local = T.alloc_buffer((1, (m + 31) // 32 * 32, 64), scope="local") + X_reindex_pad_shared = T.alloc_buffer((1, (m + 31) // 32 * 32, 256), "float16", scope="shared") + W_reindex_pad_shared = T.alloc_buffer((1, 64, 256), "float16", scope="shared") + for ax0_ax2_0_fused in T.thread_binding(1, thread="blockIdx.y"): + for ax1_0 in T.thread_binding((m + 31) // 32, thread="blockIdx.x"): + for ax2_1 in T.thread_binding(1, thread="vthread.y"): + for ax1_1 in T.thread_binding(1, thread="vthread.x"): + for ax2_2 in T.thread_binding(16, thread="threadIdx.y"): + for ax1_2 in T.thread_binding(8, thread="threadIdx.x", annotations={"pragma_auto_unroll_max_step": 256, "pragma_unroll_explicit": 1}): + for ax1_3_init, ax2_3_0_init in T.grid(4, 2): + for ax2_3_1_init in T.vectorized(2): + with T.block("compute_init"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial((m + 31) // 32 * 32, ax1_0 * 32 + ax1_1 * 32 + ax1_2 * 4 + ax1_3_init) + v2 = T.axis.spatial(64, ax2_1 * 64 + ax2_2 * 4 + ax2_3_0_init * 2 + ax2_3_1_init) + T.reads() + T.writes(compute_reindex_pad_local[0, v1, v2]) + compute_reindex_pad_local[0, v1, v2] = T.float32(0) + for ax3_0 in range(16): + for ax0_ax1_ax2_fused_0 in T.thread_binding(16, thread="threadIdx.y"): + for ax0_ax1_ax2_fused_1 in T.thread_binding(8, thread="threadIdx.x"): + for ax0_ax1_ax2_fused_2 in range(2): + for ax0_ax1_ax2_fused_3 in T.vectorized(2): + with T.block("X_reindex_pad_shared"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial((m + 31) // 32 * 32, ax1_0 * 32 + (ax0_ax1_ax2_fused_0 * 32 + ax0_ax1_ax2_fused_1 * 4 + ax0_ax1_ax2_fused_2 * 2 + ax0_ax1_ax2_fused_3) // 16) + v2 = T.axis.spatial(256, ax3_0 * 16 + (ax0_ax1_ax2_fused_0 * 32 + ax0_ax1_ax2_fused_1 * 4 + ax0_ax1_ax2_fused_2 * 2 + ax0_ax1_ax2_fused_3) % 16) + T.reads(X[v1, v2]) + T.writes(X_reindex_pad_shared[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 8, 2]]}) + X_reindex_pad_shared[v0, v1, v2] = T.if_then_else(v1 < m, X[v1, v2], T.float16(0)) + for ax0_ax1_ax2_fused_0 in T.thread_binding(16, thread="threadIdx.y"): + for ax0_ax1_ax2_fused_1 in T.thread_binding(8, thread="threadIdx.x"): + for ax0_ax1_ax2_fused_2 in range(4): + for ax0_ax1_ax2_fused_3 in T.vectorized(2): + with T.block("W_reindex_pad_shared"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(64, (ax0_ax1_ax2_fused_0 * 64 + ax0_ax1_ax2_fused_1 * 8 + ax0_ax1_ax2_fused_2 * 2 + ax0_ax1_ax2_fused_3) // 16) + v2 = T.axis.spatial(256, ax3_0 * 16 + (ax0_ax1_ax2_fused_0 * 64 + ax0_ax1_ax2_fused_1 * 8 + ax0_ax1_ax2_fused_2 * 2 + ax0_ax1_ax2_fused_3) % 16) + T.reads(W[v1, v2]) + T.writes(W_reindex_pad_shared[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 8, 2]]}) + W_reindex_pad_shared[v0, v1, v2] = T.if_then_else(v1 < 15, W[v1, v2], T.float16(0)) + for ax3_1, ax1_3, ax2_3_0 in T.grid(16, 4, 2): + for ax2_3_1 in T.vectorized(2): + with T.block("compute_update"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial((m + 31) // 32 * 32, ax1_0 * 32 + ax1_1 * 32 + ax1_2 * 4 + ax1_3) + v2 = T.axis.spatial(64, ax2_1 * 64 + ax2_2 * 4 + ax2_3_0 * 2 + ax2_3_1) + v3 = T.axis.reduce(256, ax3_0 * 16 + ax3_1) + T.reads(compute_reindex_pad_local[0, v1, v2], X_reindex_pad_shared[0, v1, v3], W_reindex_pad_shared[0, v2, v3]) + T.writes(compute_reindex_pad_local[0, v1, v2]) + compute_reindex_pad_local[0, v1, v2] = compute_reindex_pad_local[0, v1, v2] + T.Cast("float32", X_reindex_pad_shared[0, v1, v3]) * T.Cast("float32", W_reindex_pad_shared[0, v2, v3]) + for ax0, ax1, ax2_0 in T.grid(1, 4, 2): + for ax2_1_1 in T.vectorized(2): + with T.block("compute_reindex_pad_local"): + v0 = T.axis.spatial(1, ax0) + v1 = T.axis.spatial((m + 31) // 32 * 32, ax1_0 * 32 + ax1_2 * 4 + ax1) + v2 = T.axis.spatial(64, ax2_2 * 4 + ax2_0 * 2 + ax2_1_1) + T.where(ax1_0 * 32 + ax1_2 * 4 + ax1 < m and ax2_2 * 4 + ax2_0 * 2 + ax2_1_1 < 15) + T.reads(compute_reindex_pad_local[v0, v1, v2]) + T.writes(compute[v1, v2]) + compute[v1, v2] = compute_reindex_pad_local[v0, v1, v2] + # fmt: on + + +@tvm.testing.requires_maca +class TestMatmulTensorizeEpilogueMACA(MACABeforeAfter): + # fmt: off + + @T.prim_func + def before(lv686: T.Buffer((T.int32(4096), T.int32(256)), "uint32"), lv687: T.Buffer((T.int32(4096), T.int32(64)), "float16"), p_lv42: T.handle, p_lv3: T.handle, p_output0: T.handle): + T.func_attr({"tir.noalias": T.bool(True)}) + n = T.int32() + lv42 = T.match_buffer(p_lv42, (T.int32(1), n, T.int32(2048)), "float16") + lv3 = T.match_buffer(p_lv3, (T.int32(1), n, T.int32(4096)), "float16") + p_output0_intermediate = T.match_buffer(p_output0, (T.int32(1), n, T.int32(4096)), "float16") + # with T.block("root"): + p_output0_intermediate_1 = T.alloc_buffer((T.int32(4096), T.int32(2048)), "float16") + var_NT_matmul_intermediate = T.alloc_buffer((T.int32(1), n, T.int32(4096)), "float16") + var_T_divide_intermediate = T.alloc_buffer((T.int32(1), n, T.int32(4096)), "float16") + for i, j in T.grid(T.int32(4096), T.int32(2048)): + with T.block("decode"): + v_i, v_j = T.axis.remap("SS", [i, j]) + T.reads(lv686[v_i, v_j // T.int32(8)], lv687[v_i, v_j // T.int32(32)]) + T.writes(p_output0_intermediate_1[v_i, v_j]) + p_output0_intermediate_1[v_i, v_j] = (T.Cast("float16", T.bitwise_and(T.shift_right(lv686[v_i, v_j // T.int32(8)], T.Cast("uint32", v_j % T.int32(8)) * T.uint32(4)), T.uint32(15))) - T.float16(7)) * lv687[v_i, v_j // T.int32(32)] + for i0, i1, i2, k in T.grid(T.int32(1), n, T.int32(4096), T.int32(2048)): + with T.block("NT_matmul"): + v_i0, v_i1, v_i2, v_k = T.axis.remap("SSSR", [i0, i1, i2, k]) + T.reads(lv42[v_i0, v_i1, v_k], p_output0_intermediate_1[v_i2, v_k]) + T.writes(var_NT_matmul_intermediate[v_i0, v_i1, v_i2]) + with T.init(): + var_NT_matmul_intermediate[v_i0, v_i1, v_i2] = T.float16(0) + var_NT_matmul_intermediate[v_i0, v_i1, v_i2] = var_NT_matmul_intermediate[v_i0, v_i1, v_i2] + lv42[v_i0, v_i1, v_k] * p_output0_intermediate_1[v_i2, v_k] + for ax0, ax1, ax2 in T.grid(T.int32(1), n, T.int32(4096)): + with T.block("T_divide"): + v_ax0, v_ax1, v_ax2 = T.axis.remap("SSS", [ax0, ax1, ax2]) + T.reads(lv3[v_ax0, v_ax1, v_ax2]) + T.writes(var_T_divide_intermediate[v_ax0, v_ax1, v_ax2]) + var_T_divide_intermediate[v_ax0, v_ax1, v_ax2] = lv3[v_ax0, v_ax1, v_ax2] * T.float16(0.5) + for ax0, ax1, ax2 in T.grid(T.int32(1), n, T.int32(4096)): + with T.block("T_add"): + v_ax0, v_ax1, v_ax2 = T.axis.remap("SSS", [ax0, ax1, ax2]) + T.reads(var_T_divide_intermediate[v_ax0, v_ax1, v_ax2], var_NT_matmul_intermediate[v_ax0, v_ax1, v_ax2]) + T.writes(p_output0_intermediate[v_ax0, v_ax1, v_ax2]) + p_output0_intermediate[v_ax0, v_ax1, v_ax2] = var_T_divide_intermediate[v_ax0, v_ax1, v_ax2] + var_NT_matmul_intermediate[v_ax0, v_ax1, v_ax2] + + @T.prim_func + def expected(lv686: T.Buffer((4096, 256), "uint32"), lv687: T.Buffer((4096, 64), "float16"), p_lv42: T.handle, p_lv3: T.handle, p_output0: T.handle): + T.func_attr({"global_symbol": "before", "tir.is_scheduled": True, "tir.noalias": T.bool(True)}) + n = T.int32() + lv42 = T.match_buffer(p_lv42, (1, n, 2048), "float16") + lv3 = T.match_buffer(p_lv3, (1, n, 4096), "float16") + p_output0_intermediate = T.match_buffer(p_output0, (1, n, 4096), "float16") + # with T.block("root"): + lv42_reindex_pad_shared_dyn = T.alloc_buffer((1, (n + 63) // 64 * 64, 2048), "float16", scope="shared.dyn") + p_output0_intermediate_1_reindex_shared_dyn = T.alloc_buffer((1, 4096, 2048), "float16", scope="shared.dyn") + lv42_reindex_pad_shared_dyn_wmma_matrix_a = T.alloc_buffer((1, (n + 63) // 64 * 64, 2048), "float16", scope="wmma.matrix_a") + p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b = T.alloc_buffer((1, 4096, 2048), "float16", scope="wmma.matrix_b") + var_NT_matmul_intermediate_reindex_pad_shared_dyn = T.alloc_buffer((1, (n + 63) // 64 * 64, 4096), "float16", scope="shared.dyn") + var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator = T.alloc_buffer((1, (n + 63) // 64 * 64, 4096), "float16", scope="wmma.accumulator") + for ax0 in T.thread_binding(1, thread="blockIdx.z"): + for ax1_0_0_ax2_0_0_fused in T.thread_binding((n + 63) // 64, thread="blockIdx.x"): + for ax1_0_1_ax2_0_1_fused in T.thread_binding(64, thread="blockIdx.y"): + for ax2_0_2_ax1_0_2_fused in T.thread_binding(4, thread="threadIdx.y"): + for ax1_0_3_init, ax2_0_3_init in T.grid(2, 2): + with T.block("NT_matmul_o_init"): + v0_o = T.axis.spatial(1, ax0) + v1_o = T.axis.spatial((n + 63) // 64 * 4, ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax1_0_3_init) + v2_o = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax2_0_3_init) + T.reads() + T.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + with T.block("NT_matmul_init_o"): + v1_i_init_o = T.axis.spatial(1, 0) + v2_i_init_o = T.axis.spatial(1, 0) + T.reads() + T.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + C = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.float32(0.0)) + for ax3_0_0 in T.serial(64, annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): + for ax0_ax1_fused_0 in range(2): + for ax0_ax1_fused_1 in T.thread_binding(4, thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(4): + with T.block("lv42_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial((n + 63) // 64 * 64, ax1_0_0_ax2_0_0_fused * 64 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) // 32) + v2 = T.axis.spatial(2048, ax3_0_0 * 32 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) % 32) + T.reads(lv42[v0, v1, v2]) + T.writes(lv42_reindex_pad_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 8]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + lv42_reindex_pad_shared_dyn[v0, v1, v2] = T.if_then_else(v1 < n, lv42[v0, v1, v2], T.float16(0.0)) + for ax0_ax1_fused_0 in range(2): + for ax0_ax1_fused_1 in T.thread_binding(4, thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(4): + with T.block("p_output0_intermediate_1_reindex_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(4096, ax1_0_1_ax2_0_1_fused * 64 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) // 32) + v2 = T.axis.spatial(2048, ax3_0_0 * 32 + (ax0_ax1_fused_0 * 1024 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) % 32) + T.reads(lv686[v1, v2 // 8], lv687[v1, v2 // 32]) + T.writes(p_output0_intermediate_1_reindex_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 8]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + p_output0_intermediate_1_reindex_shared_dyn[v0, v1, v2] = (T.Cast("float16", T.bitwise_and(T.shift_right(lv686[v1, v2 // 8], T.Cast("uint32", v2 % 8) * T.uint32(4)), T.uint32(15))) - T.float16(7.0)) * lv687[v1, v2 // 32] + for ax3_0_1 in T.serial(2, annotations={"software_pipeline_order": [0, 1, 2], "software_pipeline_stage": [0, 0, 1]}): + for ax0_0 in T.unroll(2): + for ax1_0 in T.unroll(1): + with T.block("lv42_reindex_pad_shared.dyn_wmma.matrix_a_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(4 * ((n + 63) // 64), ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax0_0) + v2_o = T.axis.spatial(128, ax3_0_0 * 2 + ax3_0_1 + ax1_0) + T.reads(lv42_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(lv42_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(lv42_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(lv42_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_a", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "row_major") + for ax0_0 in T.unroll(2): + for ax1_0 in T.unroll(1): + with T.block("p_output0_intermediate_1_reindex_shared.dyn_wmma.matrix_b_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax0_0) + v2_o = T.axis.spatial(128, ax3_0_0 * 2 + ax3_0_1 + ax1_0) + T.reads(p_output0_intermediate_1_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(p_output0_intermediate_1_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_b", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "col_major") + for ax1_0_3, ax2_0_3 in T.grid(2, 2): + with T.block("NT_matmul_o_update"): + v0_o = T.axis.spatial(1, ax0) + v1_o = T.axis.spatial((n + 63) // 64 * 4, ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax1_0_3) + v2_o = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax2_0_3) + v3_o = T.axis.reduce(128, ax3_0_0 * 2 + ax3_0_1) + T.reads(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], lv42_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) + T.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + with T.block("NT_matmul_o"): + v1_i_o = T.axis.spatial(1, 0) + v2_i_o = T.axis.spatial(1, 0) + v3_i_o = T.axis.reduce(1, 0) + T.reads(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], lv42_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) + T.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(lv42_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="wmma.matrix_a", offset_factor=16) + B = T.match_buffer(p_output0_intermediate_1_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "float16", strides=("B_s0", "B_s1"), scope="wmma.matrix_b", offset_factor=16) + C = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, A.data, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, B.data, B.elem_offset // B.strides[0] // 16 * (B.strides[0] // 16) + B.elem_offset % B.strides[0] // 16, C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16) + for ax0_0, ax1_0 in T.grid(2, 2): + with T.block("var_NT_matmul_intermediate_reindex_pad_shared.dyn_wmma.accumulator_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(4 * ((n + 63) // 64), ax1_0_0_ax2_0_0_fused * 4 + ax2_0_2_ax1_0_2_fused % 2 * 2 + ax0_0) + v2_o = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 4 + ax2_0_2_ax1_0_2_fused // 2 * 2 + ax1_0) + T.reads(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="wmma.accumulator", offset_factor=16) + C = T.match_buffer(var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="shared.dyn", offset_factor=16) + T.tvm_store_matrix_sync(A.data, 16, 16, 16, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("float16"), C.data, C.elem_offset, C.strides[0] * 16, 2), C.strides[0], "row_major") + for ax0_ax1_fused_0 in range(4): + for ax0_ax1_fused_1 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_2 in T.vectorized(4): + with T.block("var_NT_matmul_intermediate_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial((n + 63) // 64 * 64, ax1_0_0_ax2_0_0_fused * 64 + ax2_0_2_ax1_0_2_fused % 2 * 32 + (ax0_ax1_fused_0 * 256 + ax0_ax1_fused_1 * 4 + ax0_ax1_fused_2) // 32) + v2 = T.axis.spatial(4096, ax1_0_1_ax2_0_1_fused * 64 + ax2_0_2_ax1_0_2_fused // 2 * 32 + (ax0_ax1_fused_0 * 256 + ax0_ax1_fused_1 * 4 + ax0_ax1_fused_2) % 32) + T.where(ax1_0_0_ax2_0_0_fused * 64 + ax2_0_2_ax1_0_2_fused % 2 * 32 + ((ax0_ax1_fused_0 * 64 + ax0_ax1_fused_1) * 4 + ax0_ax1_fused_2) // 32 < n) + T.reads(lv3[0, v1, v2], var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0, v1, v2]) + T.writes(p_output0_intermediate[0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 4]]}) + p_output0_intermediate[0, v1, v2] = lv3[0, v1, v2] * T.float16(0.5) + var_NT_matmul_intermediate_reindex_pad_shared_dyn[v0, v1, v2] + # fmt: on + + +@tvm.testing.requires_maca +class TestMatmulInt8TensorizeMACA(MACABeforeAfter): + # fmt: off + @T.prim_func + def before(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), compute: T.Buffer((256, 256), "int32")): + T.func_attr({"global_symbol": "main", "tir.noalias": T.bool(True)}) + # with T.block("root"): + for i, j, r in T.grid(256, 256, 256): + with T.block("compute"): + v_i, v_j, v_k = T.axis.remap("SSR", [i, j, r]) + T.reads(X[v_i, v_k], W[v_j, v_k]) + T.writes(compute[v_i, v_j]) + with T.init(): + compute[v_i, v_j] = 0 + compute[v_i, v_j] = compute[v_i, v_j] + T.Cast("int32", X[v_i, v_k]) * T.Cast("int32", W[v_j, v_k]) + + @T.prim_func + def expected(X: T.Buffer((256, 256), "int8"), W: T.Buffer((256, 256), "int8"), compute: T.Buffer((256, 256), "int32")): + T.func_attr({"tir.is_scheduled": True, "tir.noalias": T.bool(True)}) + # with T.block("root"): + X_reindex_shared_dyn = T.alloc_buffer((1, 256, 256), "int8", scope="shared.dyn") + W_reindex_shared_dyn = T.alloc_buffer((1, 256, 256), "int8", scope="shared.dyn") + X_reindex_shared_dyn_wmma_matrix_a = T.alloc_buffer((1, 256, 256), "int8", scope="wmma.matrix_a") + W_reindex_shared_dyn_wmma_matrix_b = T.alloc_buffer((1, 256, 256), "int8", scope="wmma.matrix_b") + compute_reindex_shared_dyn = T.alloc_buffer((1, 256, 256), "int32", scope="shared.dyn") + compute_reindex_shared_dyn_wmma_accumulator = T.alloc_buffer((1, 256, 256), "int32", scope="wmma.accumulator") + for ax0 in T.thread_binding(1, thread="blockIdx.z"): + for ax1_0_0_ax2_0_0_fused in T.thread_binding(2, thread="blockIdx.x"): + for ax1_0_1_ax2_0_1_fused in T.thread_binding(2, thread="blockIdx.y"): + for ax2_0_2_ax1_0_2_fused in T.thread_binding(16, thread="threadIdx.y"): + for ax1_0_3_init, ax2_0_3_init in T.grid(2, 2): + with T.block("compute_o_init"): + v0_o = T.axis.spatial(1, ax0) + v1_o = T.axis.spatial(16, ax1_0_0_ax2_0_0_fused * 8 + ax2_0_2_ax1_0_2_fused % 4 * 2 + ax1_0_3_init) + v2_o = T.axis.spatial(16, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax2_0_3_init) + T.reads() + T.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + with T.block("compute_init_o"): + v1_i_init_o = T.axis.spatial(1, 0) + v2_i_init_o = T.axis.spatial(1, 0) + T.reads() + T.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + C = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.float32(0.0)) + for ax3_0_0 in T.serial(16, annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): + for ax0_ax1_fused_0 in range(1): + for ax0_ax1_fused_1 in T.thread_binding(16, thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(4): + with T.block("X_reindex_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(256, ax1_0_0_ax2_0_0_fused * 128 + (ax0_ax1_fused_0 * 4096 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) // 16) + v2 = T.axis.spatial(256, ax3_0_0 * 16 + (ax0_ax1_fused_0 * 4096 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) % 16) + T.where(((ax0_ax1_fused_0 * 16 + ax0_ax1_fused_1) * 64 + ax0_ax1_fused_2) * 4 + ax0_ax1_fused_3 < 2048) + T.reads(X[v1, v2]) + T.writes(X_reindex_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 32, 16]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + X_reindex_shared_dyn[v0, v1, v2] = X[v1, v2] + for ax0_ax1_fused_0 in range(1): + for ax0_ax1_fused_1 in T.thread_binding(16, thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(4): + with T.block("W_reindex_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 128 + (ax0_ax1_fused_0 * 4096 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) // 16) + v2 = T.axis.spatial(256, ax3_0_0 * 16 + (ax0_ax1_fused_0 * 4096 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) % 16) + T.where(((ax0_ax1_fused_0 * 16 + ax0_ax1_fused_1) * 64 + ax0_ax1_fused_2) * 4 + ax0_ax1_fused_3 < 2048) + T.reads(W[v1, v2]) + T.writes(W_reindex_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 32, 16]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + W_reindex_shared_dyn[v0, v1, v2] = W[v1, v2] + for ax3_0_1 in T.serial(1, annotations={"software_pipeline_order": [0, 1, 2], "software_pipeline_stage": [0, 0, 1]}): + for ax0_0 in T.unroll(2): + for ax1_0 in T.unroll(1): + with T.block("X_reindex_shared.dyn_wmma.matrix_a_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(16, ax1_0_0_ax2_0_0_fused * 8 + ax2_0_2_ax1_0_2_fused % 4 * 2 + ax0_0) + v2_o = T.axis.spatial(16, ax3_0_0 + ax1_0) + T.reads(X_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(X_reindex_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(X_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(X_reindex_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=("C_s0", "C_s1"), scope="wmma.matrix_a", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int8"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "row_major") + for ax0_0 in T.unroll(2): + for ax1_0 in T.unroll(1): + with T.block("W_reindex_shared.dyn_wmma.matrix_b_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(16, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax0_0) + v2_o = T.axis.spatial(16, ax3_0_0 + ax1_0) + T.reads(W_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(W_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(W_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(W_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=("C_s0", "C_s1"), scope="wmma.matrix_b", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int8"), A.data, A.elem_offset, A.strides[0] * 16, 1), A.strides[0], "col_major") + for ax1_0_3, ax2_0_3 in T.grid(2, 2): + with T.block("compute_o_update"): + v0_o = T.axis.spatial(1, ax0) + v1_o = T.axis.spatial(16, ax1_0_0_ax2_0_0_fused * 8 + ax2_0_2_ax1_0_2_fused % 4 * 2 + ax1_0_3) + v2_o = T.axis.spatial(16, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax2_0_3) + v3_o = T.axis.reduce(16, ax3_0_0 + ax3_0_1) + T.reads(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], X_reindex_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], W_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) + T.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + with T.block("compute_o"): + v1_i_o = T.axis.spatial(1, 0) + v2_i_o = T.axis.spatial(1, 0) + v3_i_o = T.axis.reduce(1, 0) + T.reads(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], X_reindex_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], W_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) + T.writes(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(X_reindex_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=("A_s0", "A_s1"), scope="wmma.matrix_a", offset_factor=16) + B = T.match_buffer(W_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=("B_s0", "B_s1"), scope="wmma.matrix_b", offset_factor=16) + C = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, A.data, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, B.data, B.elem_offset // B.strides[0] // 16 * (B.strides[0] // 16) + B.elem_offset % B.strides[0] // 16, C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16) + for ax0_0, ax1_0 in T.grid(2, 2): + with T.block("compute_reindex_shared.dyn_wmma.accumulator_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(16, ax1_0_0_ax2_0_0_fused * 8 + ax2_0_2_ax1_0_2_fused % 4 * 2 + ax0_0) + v2_o = T.axis.spatial(16, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax1_0) + T.reads(compute_reindex_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(compute_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A = T.match_buffer(compute_reindex_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=("A_s0", "A_s1"), scope="wmma.accumulator", offset_factor=16) + C = T.match_buffer(compute_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=("C_s0", "C_s1"), scope="shared.dyn", offset_factor=16) + T.tvm_store_matrix_sync(A.data, 16, 16, 16, A.elem_offset // A.strides[0] // 16 * (A.strides[0] // 16) + A.elem_offset % A.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int32"), C.data, C.elem_offset, C.strides[0] * 16, 2), C.strides[0], "row_major") + for ax0_ax1_fused_0 in range(4): + for ax0_ax1_fused_1 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_2 in T.vectorized(4): + with T.block("compute_reindex_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(256, ax1_0_0_ax2_0_0_fused * 128 + ax2_0_2_ax1_0_2_fused % 4 * 32 + (ax0_ax1_fused_0 * 256 + ax0_ax1_fused_1 * 4 + ax0_ax1_fused_2) // 32) + v2 = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 128 + ax2_0_2_ax1_0_2_fused // 4 * 32 + (ax0_ax1_fused_0 * 256 + ax0_ax1_fused_1 * 4 + ax0_ax1_fused_2) % 32) + T.reads(compute_reindex_shared_dyn[v0, v1, v2]) + T.writes(compute[v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 4]]}) + compute[v1, v2] = compute_reindex_shared_dyn[v0, v1, v2] + # fmt: on + + +@tvm.testing.requires_maca +class TestMatmulInt8Tensorize3d2dDynMACA(MACABeforeAfter): + # fmt: off + @T.prim_func + def before(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T.handle): + T.func_attr({"op_pattern": 4, "tir.noalias": T.bool(True)}) + m = T.int32() + A = T.match_buffer(var_A, (1, m, 22016), "int8") + matmul_1 = T.match_buffer(var_matmul, (1, m, 4096), "int32") + # with T.block("root"): + for i0, i1, i2, k in T.grid(1, m, 4096, 22016): + with T.block("matmul"): + v_i0, v_i1, v_i2, v_k = T.axis.remap("SSSR", [i0, i1, i2, k]) + T.reads(A[v_i0, v_i1, v_k], B[v_i2, v_k]) + T.writes(matmul_1[v_i0, v_i1, v_i2]) + with T.init(): + matmul_1[v_i0, v_i1, v_i2] = 0 + matmul_1[v_i0, v_i1, v_i2] = matmul_1[v_i0, v_i1, v_i2] + T.Cast("int32", A[v_i0, v_i1, v_k]) * T.Cast("int32", B[v_i2, v_k]) + + @T.prim_func + def expected(var_A: T.handle, B: T.Buffer((4096, 22016), "int8"), var_matmul: T.handle): + T.func_attr({"global_symbol": "before", "op_pattern": 4, "tir.is_scheduled": True, "tir.noalias": T.bool(True)}) + m = T.int32() + A = T.match_buffer(var_A, (1, m, 22016), "int8") + matmul_1 = T.match_buffer(var_matmul, (1, m, 4096), "int32") + # with T.block("root"): + A_reindex_pad_shared_dyn = T.alloc_buffer((1, (m + 127) // 128 * 128, 22016), "int8", scope="shared.dyn") + B_reindex_shared_dyn = T.alloc_buffer((1, 4096, 22016), "int8", scope="shared.dyn") + A_reindex_pad_shared_dyn_wmma_matrix_a = T.alloc_buffer((1, (m + 127) // 128 * 128, 22016), "int8", scope="wmma.matrix_a") + B_reindex_shared_dyn_wmma_matrix_b = T.alloc_buffer((1, 4096, 22016), "int8", scope="wmma.matrix_b") + matmul_1_reindex_pad_shared_dyn = T.alloc_buffer((1, (m + 127) // 128 * 128, 4096), "int32", scope="shared.dyn") + matmul_1_reindex_pad_shared_dyn_wmma_accumulator = T.alloc_buffer((1, (m + 127) // 128 * 128, 4096), "int32", scope="wmma.accumulator") + for ax0 in T.thread_binding(1, thread="blockIdx.z"): + for ax1_0_0_ax2_0_0_fused in T.thread_binding((m + 127) // 128, thread="blockIdx.x"): + for ax1_0_1_ax2_0_1_fused in T.thread_binding(32, thread="blockIdx.y"): + for ax2_0_2_ax1_0_2_fused in T.thread_binding(16, thread="threadIdx.y"): + for ax1_0_3_init, ax2_0_3_init in T.grid(2, 2): + with T.block("matmul_o_init"): + v0_o = T.axis.spatial(1, ax0) + v1_o = T.axis.spatial((m + 127) // 128 * 8, ax1_0_0_ax2_0_0_fused * 8 + ax2_0_2_ax1_0_2_fused % 4 * 2 + ax1_0_3_init) + v2_o = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax2_0_3_init) + T.reads() + T.writes(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + with T.block("matmul_init_o"): + v1_i_init_o = T.axis.spatial(1, 0) + v2_i_init_o = T.axis.spatial(1, 0) + T.reads() + T.writes(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + C = T.match_buffer(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_fill_fragment(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.float32(0.0)) + for ax3_0_0 in T.serial(1376, annotations={"software_pipeline_order": [0, 3, 1, 4, 5, 2, 6], "software_pipeline_stage": [0, 0, 0, 0, 0, 1, 1]}): + for ax0_ax1_fused_0 in range(1): + for ax0_ax1_fused_1 in T.thread_binding(16, thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(4): + with T.block("A_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial((m + 127) // 128 * 128, ax1_0_0_ax2_0_0_fused * 128 + (ax0_ax1_fused_0 * 4096 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) // 16) + v2 = T.axis.spatial(22016, ax3_0_0 * 16 + (ax0_ax1_fused_0 * 4096 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) % 16) + T.where(((ax0_ax1_fused_0 * 16 + ax0_ax1_fused_1) * 64 + ax0_ax1_fused_2) * 4 + ax0_ax1_fused_3 < 2048) + T.reads(A[v0, v1, v2]) + T.writes(A_reindex_pad_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 32, 16]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + A_reindex_pad_shared_dyn[v0, v1, v2] = T.if_then_else(v1 < m, A[v0, v1, v2], T.int8(0)) + for ax0_ax1_fused_0 in range(1): + for ax0_ax1_fused_1 in T.thread_binding(16, thread="threadIdx.y"): + for ax0_ax1_fused_2 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_3 in T.vectorized(4): + with T.block("B_reindex_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial(4096, ax1_0_1_ax2_0_1_fused * 128 + (ax0_ax1_fused_0 * 4096 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) // 16) + v2 = T.axis.spatial(22016, ax3_0_0 * 16 + (ax0_ax1_fused_0 * 4096 + ax0_ax1_fused_1 * 256 + ax0_ax1_fused_2 * 4 + ax0_ax1_fused_3) % 16) + T.where(((ax0_ax1_fused_0 * 16 + ax0_ax1_fused_1) * 64 + ax0_ax1_fused_2) * 4 + ax0_ax1_fused_3 < 2048) + T.reads(B[v1, v2]) + T.writes(B_reindex_shared_dyn[v0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 32, 16]], "double_buffer_scope": 0, "tir.manifest_shared_memory_local_stage": 1}) + B_reindex_shared_dyn[v0, v1, v2] = B[v1, v2] + for ax3_0_1 in T.serial(1, annotations={"software_pipeline_order": [0, 1, 2], "software_pipeline_stage": [0, 0, 1]}): + for ax0_0 in T.unroll(2): + for ax1_0 in T.unroll(1): + with T.block("A_reindex_pad_shared.dyn_wmma.matrix_a_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(8 * ((m + 127) // 128), ax1_0_0_ax2_0_0_fused * 8 + ax2_0_2_ax1_0_2_fused % 4 * 2 + ax0_0) + v2_o = T.axis.spatial(1376, ax3_0_0 + ax1_0) + T.reads(A_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(A_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A_1 = T.match_buffer(A_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(A_reindex_pad_shared_dyn_wmma_matrix_a[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=("C_s0", "C_s1"), scope="wmma.matrix_a", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int8"), A_1.data, A_1.elem_offset, A_1.strides[0] * 16, 1), A_1.strides[0], "row_major") + for ax0_0 in T.unroll(2): + for ax1_0 in T.unroll(1): + with T.block("B_reindex_shared.dyn_wmma.matrix_b_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax0_0) + v2_o = T.axis.spatial(1376, ax3_0_0 + ax1_0) + T.reads(B_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(B_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A_1 = T.match_buffer(B_reindex_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=("A_s0", "A_s1"), scope="shared.dyn", offset_factor=16) + C = T.match_buffer(B_reindex_shared_dyn_wmma_matrix_b[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int8", strides=("C_s0", "C_s1"), scope="wmma.matrix_b", offset_factor=16) + T.tvm_load_matrix_sync(C.data, 16, 16, 16, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int8"), A_1.data, A_1.elem_offset, A_1.strides[0] * 16, 1), A_1.strides[0], "col_major") + for ax1_0_3, ax2_0_3 in T.grid(2, 2): + with T.block("matmul_o_update"): + v0_o = T.axis.spatial(1, ax0) + v1_o = T.axis.spatial((m + 127) // 128 * 8, ax1_0_0_ax2_0_0_fused * 8 + ax2_0_2_ax1_0_2_fused % 4 * 2 + ax1_0_3) + v2_o = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax2_0_3) + v3_o = T.axis.reduce(1376, ax3_0_0 + ax3_0_1) + T.reads(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], A_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], B_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) + T.writes(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + with T.block("matmul_o"): + v1_i_o = T.axis.spatial(1, 0) + v2_i_o = T.axis.spatial(1, 0) + v3_i_o = T.axis.reduce(1, 0) + T.reads(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], A_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], B_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) + T.writes(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A_1 = T.match_buffer(A_reindex_pad_shared_dyn_wmma_matrix_a[0, v1_o * 16:v1_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=("A_s0", "A_s1"), scope="wmma.matrix_a", offset_factor=16) + B_1 = T.match_buffer(B_reindex_shared_dyn_wmma_matrix_b[0, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=("B_s0", "B_s1"), scope="wmma.matrix_b", offset_factor=16) + C = T.match_buffer(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[0, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + T.tvm_mma_sync(C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16, A_1.data, A_1.elem_offset // A_1.strides[0] // 16 * (A_1.strides[0] // 16) + A_1.elem_offset % A_1.strides[0] // 16, B_1.data, B_1.elem_offset // B_1.strides[0] // 16 * (B_1.strides[0] // 16) + B_1.elem_offset % B_1.strides[0] // 16, C.data, C.elem_offset // C.strides[0] // 16 * (C.strides[0] // 16) + C.elem_offset % C.strides[0] // 16) + for ax0_0, ax1_0 in T.grid(2, 2): + with T.block("matmul_1_reindex_pad_shared.dyn_wmma.accumulator_o"): + v0_o = T.axis.spatial(1, 0) + v1_o = T.axis.spatial(8 * ((m + 127) // 128), ax1_0_0_ax2_0_0_fused * 8 + ax2_0_2_ax1_0_2_fused % 4 * 2 + ax0_0) + v2_o = T.axis.spatial(256, ax1_0_1_ax2_0_1_fused * 8 + ax2_0_2_ax1_0_2_fused // 4 * 2 + ax1_0) + T.reads(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + T.writes(matmul_1_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16]) + A_1 = T.match_buffer(matmul_1_reindex_pad_shared_dyn_wmma_accumulator[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=("A_s0", "A_s1"), scope="wmma.accumulator", offset_factor=16) + C = T.match_buffer(matmul_1_reindex_pad_shared_dyn[v0_o, v1_o * 16:v1_o * 16 + 16, v2_o * 16:v2_o * 16 + 16], (16, 16), "int32", strides=("C_s0", "C_s1"), scope="shared.dyn", offset_factor=16) + T.tvm_store_matrix_sync(A_1.data, 16, 16, 16, A_1.elem_offset // A_1.strides[0] // 16 * (A_1.strides[0] // 16) + A_1.elem_offset % A_1.strides[0] // 16, T.tvm_access_ptr(T.type_annotation("int32"), C.data, C.elem_offset, C.strides[0] * 16, 2), C.strides[0], "row_major") + for ax0_ax1_fused_0 in range(4): + for ax0_ax1_fused_1 in T.thread_binding(64, thread="threadIdx.x"): + for ax0_ax1_fused_2 in T.vectorized(4): + with T.block("matmul_1_reindex_pad_shared.dyn"): + v0 = T.axis.spatial(1, 0) + v1 = T.axis.spatial((m + 127) // 128 * 128, ax1_0_0_ax2_0_0_fused * 128 + ax2_0_2_ax1_0_2_fused % 4 * 32 + (ax0_ax1_fused_0 * 256 + ax0_ax1_fused_1 * 4 + ax0_ax1_fused_2) // 32) + v2 = T.axis.spatial(4096, ax1_0_1_ax2_0_1_fused * 128 + ax2_0_2_ax1_0_2_fused // 4 * 32 + (ax0_ax1_fused_0 * 256 + ax0_ax1_fused_1 * 4 + ax0_ax1_fused_2) % 32) + T.where(ax1_0_0_ax2_0_0_fused * 128 + ax2_0_2_ax1_0_2_fused % 4 * 32 + ((ax0_ax1_fused_0 * 64 + ax0_ax1_fused_1) * 4 + ax0_ax1_fused_2) // 32 < m) + T.reads(matmul_1_reindex_pad_shared_dyn[v0, v1, v2]) + T.writes(matmul_1[0, v1, v2]) + T.block_attr({"buffer_dim_align": [[0, 1, 16, 4]]}) + matmul_1[0, v1, v2] = matmul_1_reindex_pad_shared_dyn[v0, v1, v2] + + if __name__ == "__main__": tvm.testing.main() diff --git a/tests/scripts/task_python_unittest.sh b/tests/scripts/task_python_unittest.sh index 54170133530d..a54b02387120 100755 --- a/tests/scripts/task_python_unittest.sh +++ b/tests/scripts/task_python_unittest.sh @@ -52,6 +52,7 @@ TEST_FILES=( "tir-schedule" "tir-transform" "tvmscript" + "dlight" ) for TEST_FILE in ${TEST_FILES[@]}; do