This is an automated email from the ASF dual-hosted git repository. tqchen pushed a commit to branch script/canonical-parser-df in repository https://gitbox.apache.org/repos/asf/tvm.git
commit 07124f2485bf4fd8e836fb44fe2d6b2f683430ca Author: Tianqi Chen <[email protected]> AuthorDate: Tue Sep 22 22:02:04 2026 +0000 [F][TVMScript] Adapt source helpers and tests to explicit construction Make host predicates and symbolic declarations explicit, retain public script imports, and migrate internal parser probes and dtype expectations to the canonical contracts. --- python/tvm/backend/cuda/lang/warp_role.py | 2 +- .../backend/trn/tile_primitive/binary/default.py | 2 +- .../trn/tile_primitive/compose_op/binary_reduce.py | 4 +- .../trn/tile_primitive/compose_op/unary_reduce.py | 4 +- .../tvm/backend/trn/tile_primitive/gemm/default.py | 2 +- .../tvm/backend/trn/tile_primitive/unary/utils.py | 6 +- python/tvm/relax/backend/gpu_generic/sampling.py | 2 +- python/tvm/relax/frontend/nn/llm/_page_kernels.py | 12 +- .../tvm/s_tir/tensor_intrin/dot_product_common.py | 2 +- python/tvm/s_tir/tensor_intrin/metal.py | 4 +- python/tvm/s_tir/tensor_intrin/rocm.py | 8 +- tests/python/codegen/test_target_codegen_vulkan.py | 20 +-- tests/python/relax/test_analysis_type_analysis.py | 2 - tests/python/relax/test_frontend_onnx.py | 24 +-- .../s_tir/dlight/test_gpu_matmul_tensorize.py | 180 ++++++++++++++------- .../test_meta_schedule_trace_apply.py | 14 +- tests/python/te/test_te_create_primfunc.py | 6 +- .../tirx-transform/test_tir_transform_vectorize.py | 2 +- .../operator/tile_primitive/trn/test_binary_trn.py | 38 ++--- .../operator/tile_primitive/trn/test_unary_trn.py | 12 +- tests/python/tirx/test_inline.py | 2 +- tests/python/tirx/test_jit.py | 22 +-- tests/python/tirx/test_op_namespace_cleanup.py | 3 +- tests/python/tirx/test_parser_printer.py | 2 +- .../tvmscript/test_tvmscript_error_report.py | 4 +- .../tvmscript/test_tvmscript_parser_evaluator.py | 14 +- .../tvmscript/test_tvmscript_parser_source.py | 6 +- .../python/tvmscript/test_tvmscript_parser_tir.py | 15 +- tests/python/tvmscript/test_tvmscript_roundtrip.py | 15 +- 29 files changed, 242 insertions(+), 187 deletions(-) diff --git a/python/tvm/backend/cuda/lang/warp_role.py b/python/tvm/backend/cuda/lang/warp_role.py index b61cd8822d..51974e2ee9 100644 --- a/python/tvm/backend/cuda/lang/warp_role.py +++ b/python/tvm/backend/cuda/lang/warp_role.py @@ -132,7 +132,7 @@ class WarpgroupRole: def __enter__(self): if isinstance(self.wg_id_val, tuple): start, stop = self.wg_id_val - self._if_frame = T.If(start <= self.wg_id_var and self.wg_id_var < stop) + self._if_frame = T.If(T.And(T.LE(start, self.wg_id_var), self.wg_id_var < stop)) else: self._if_frame = T.If(self.wg_id_var == self.wg_id_val) self._if_frame.__enter__() diff --git a/python/tvm/backend/trn/tile_primitive/binary/default.py b/python/tvm/backend/trn/tile_primitive/binary/default.py index 85f19d11c5..9cf669d892 100644 --- a/python/tvm/backend/trn/tile_primitive/binary/default.py +++ b/python/tvm/backend/trn/tile_primitive/binary/default.py @@ -84,7 +84,7 @@ def binary_trn( if inst_gen.make_guard(_dst): dst_indices = T.meta_var(inst_gen.generate_indices(_dst)) src1_indices = T.meta_var(inst_gen.generate_indices(_src1)) - if CONST is None: + if T.constexpr(CONST is None): src2_indices = T.meta_var(inst_gen.generate_indices(_src2)) T.evaluate( func( diff --git a/python/tvm/backend/trn/tile_primitive/compose_op/binary_reduce.py b/python/tvm/backend/trn/tile_primitive/compose_op/binary_reduce.py index a3a83bb3a4..fd00bfc0a1 100644 --- a/python/tvm/backend/trn/tile_primitive/compose_op/binary_reduce.py +++ b/python/tvm/backend/trn/tile_primitive/compose_op/binary_reduce.py @@ -113,7 +113,7 @@ def binary_reduce_trn(op: TilePrimitiveCall, sctx: DispatchContext) -> PrimFunc vec_dst_idx = T.meta_var(inst_gen.generate_indices(binary_output)) reduce_dst_idx = T.meta_var(inst_gen.generate_indices(reduce_output)) if inst_gen.make_guard(binary_output): - if CONST is None: + if T.constexpr(CONST is None): src_2_indices = T.meta_var(inst_gen.generate_indices(binary_input2)) # noqa: E501 T.nki.tensorscalar_reduce(dst2[tuple(reduce_dst_idx)], dst1[tuple(vec_dst_idx)], src1[tuple(src_1_indices)], src2[tuple(src_2_indices)], binary_opcode, reduce_opcode, reverse[0]) # noqa: E501 else: @@ -134,7 +134,7 @@ def binary_reduce_trn(op: TilePrimitiveCall, sctx: DispatchContext) -> PrimFunc if inst_gen.make_guard(binary_output): src_1_indices = T.meta_var(inst_gen.generate_indices(binary_input1)) # noqa: E501 vec_dst_idx = T.meta_var(inst_gen.generate_indices(binary_output)) # noqa: E501 - if CONST is None: + if T.constexpr(CONST is None): src_2_indices = T.meta_var(inst_gen.generate_indices(binary_input2)) # noqa: E501 T.nki.tensorscalar_reduce(intermediate_buffer[p_loop, reduction_b_loop], dst1[tuple(vec_dst_idx)], src1[tuple(src_1_indices)], src2[tuple(src_2_indices)], binary_opcode, reduce_opcode, reverse[0]) # noqa: E501 else: diff --git a/python/tvm/backend/trn/tile_primitive/compose_op/unary_reduce.py b/python/tvm/backend/trn/tile_primitive/compose_op/unary_reduce.py index 2cc80e57c2..8b56df56b1 100644 --- a/python/tvm/backend/trn/tile_primitive/compose_op/unary_reduce.py +++ b/python/tvm/backend/trn/tile_primitive/compose_op/unary_reduce.py @@ -112,7 +112,7 @@ def unary_reduce_trn(op: TilePrimitiveCall, sctx: DispatchContext) -> PrimFunc | dst_1_indices = T.meta_var(inst_gen.generate_indices(unary_output)) dst_2_indices = T.meta_var(inst_gen.generate_indices(reduce_output)) if inst_gen.make_guard(unary_output): - if isinstance(bias, TensorRegion): + if T.constexpr(isinstance(bias, TensorRegion)): src_bias_indices = T.meta_var(inst_gen.generate_indices(bias)) T.evaluate(T.nki.activation_reduce(dst2[tuple(dst_2_indices)], dst1[tuple(dst_1_indices)], src[tuple(src_1_indices)], unary_opcode, reduce_opcode, bias_buffer[tuple(src_bias_indices)], scale)) # noqa: E501 else: @@ -138,7 +138,7 @@ def unary_reduce_trn(op: TilePrimitiveCall, sctx: DispatchContext) -> PrimFunc | src_1_indices = T.meta_var(inst_gen.generate_indices(unary_input)) dst_1_indices = T.meta_var(inst_gen.generate_indices(unary_output)) if inst_gen.make_guard(unary_output): - if isinstance(bias, TensorRegion): + if T.constexpr(isinstance(bias, TensorRegion)): src_bias_indices = T.meta_var(inst_gen.generate_indices(bias)) # noqa: E501 T.evaluate(T.nki.activation_reduce(intermediate_buffer[p_loop, reduction_b_loop], dst1[tuple(dst_1_indices)], src[tuple(src_1_indices)], unary_opcode, reduce_opcode, bias_buffer[tuple(src_bias_indices)], scale)) # noqa: E501 else: diff --git a/python/tvm/backend/trn/tile_primitive/gemm/default.py b/python/tvm/backend/trn/tile_primitive/gemm/default.py index 9935a38970..133a314327 100644 --- a/python/tvm/backend/trn/tile_primitive/gemm/default.py +++ b/python/tvm/backend/trn/tile_primitive/gemm/default.py @@ -235,7 +235,7 @@ def matmul_trn(op: TilePrimitiveCall, sctx: DispatchContext) -> PrimFunc | None: rhs_indices = T.meta_var(inst_gen.generate_indices(B_buffer_region)) C_indices = T.meta_var(inst_gen.generate_indices(C_buffer_region)) if inst_gen.make_guard(A_buffer_region) and inst_gen.make_guard(B_buffer_region): # noqa: E501 - if C_as_output: + if T.constexpr(C_as_output): T.evaluate(T.nki.matmul(acc[C_indices], A[lhs_indices], B[rhs_indices])) # noqa: E501 else: T.evaluate(T.nki.matmul(acc[b_idx % max_psum_slots, lhs_f_loop, rhs_f_loop], A[lhs_indices], B[rhs_indices])) # noqa: E501 diff --git a/python/tvm/backend/trn/tile_primitive/unary/utils.py b/python/tvm/backend/trn/tile_primitive/unary/utils.py index 106648b17d..ff635350a0 100644 --- a/python/tvm/backend/trn/tile_primitive/unary/utils.py +++ b/python/tvm/backend/trn/tile_primitive/unary/utils.py @@ -177,13 +177,13 @@ def generate_unary_func( inst_gen.set_bind_map_all({p_var: p_loop, f_var: f_loop, b_var: b_loop}) dst_indices = T.meta_var(inst_gen.generate_indices(dst_buffer_region)) if inst_gen.make_guard(dst_buffer_region): - if unary_op == MapOpType.FILL: + if T.constexpr(unary_op == MapOpType.FILL): T.evaluate(T.nki.memset(dst[tuple(dst_indices)], _src)) else: src_indices = T.meta_var(inst_gen.generate_indices(_src)) - if unary_op == MapOpType.RECIPROCAL: + if T.constexpr(unary_op == MapOpType.RECIPROCAL): T.evaluate(T.nki.reciprocal(dst[tuple(dst_indices)], src[tuple(src_indices)])) # noqa: E501 - elif isinstance(bias, TensorRegion): + elif T.constexpr(isinstance(bias, TensorRegion)): bias_indices = T.meta_var(inst_gen.generate_indices(bias)) T.evaluate(T.nki.activation(dst[tuple(dst_indices)], src[tuple(src_indices)], opcode, scale=scale, bias=bias_buffer[tuple(bias_indices)])) # noqa: E501 else: diff --git a/python/tvm/relax/backend/gpu_generic/sampling.py b/python/tvm/relax/backend/gpu_generic/sampling.py index 487027ce7b..302b8b7e8c 100644 --- a/python/tvm/relax/backend/gpu_generic/sampling.py +++ b/python/tvm/relax/backend/gpu_generic/sampling.py @@ -174,7 +174,7 @@ def gpu_multinomial_from_uniform( local_sum[()] = T.Cast(dtype, init_value) for i in T.unroll(thread_elem): - if mask_local is not None: + if T.constexpr(mask_local is not None): if mask_local[i]: local_sum[()] = reduce_op(local_sum[()], data_local[i]) else: diff --git a/python/tvm/relax/frontend/nn/llm/_page_kernels.py b/python/tvm/relax/frontend/nn/llm/_page_kernels.py index 6682e5f18c..882de88e45 100644 --- a/python/tvm/relax/frontend/nn/llm/_page_kernels.py +++ b/python/tvm/relax/frontend/nn/llm/_page_kernels.py @@ -48,7 +48,7 @@ def _kv_cache_transpose_append(num_key_value_heads, head_dim, dtype, page_size: var_position_map: T.handle, ): T.func_attr({"tirx.noalias": True}) - ntoken = T.Var("num_tokens_excluding_cache", "int64") + ntoken = T.int64() num_pages = T.int64() pages_elem_offset = T.int64() position_map_elem_offset = T.int32() @@ -84,7 +84,7 @@ def _kv_cache_transpose_append_mla(d_qk: int, dtype, page_size: int = 16): var_position_map: T.handle, ): T.func_attr({"tirx.noalias": True}) - ntoken = T.Var("num_tokens_excluding_cache", "int64") + ntoken = T.int64() num_pages = T.int64() pages_elem_offset = T.int64() position_map_elem_offset = T.int32() @@ -115,8 +115,8 @@ def _kv_cache_debug_get_kv(num_hidden_layers, num_key_value_heads, head_dim, dty layer_id: T.int64, ): T.func_attr({"tirx.noalias": True}) - seqlen = T.Var("num_tokens_including_cache", "int64") - page_size = T.Var("page_size", "int64") + seqlen = T.int64() + page_size = T.int64() num_pages = T.int64() pages_elem_offset = T.int64() position_map_elem_offset = T.int64() @@ -147,8 +147,8 @@ def _kv_cache_debug_get_kv_mla(num_hidden_layers, d_qk, dtype): layer_id: T.int64, ): T.func_attr({"tirx.noalias": True}) - seqlen = T.Var("num_tokens_including_cache", "int64") - page_size = T.Var("page_size", "int64") + seqlen = T.int64() + page_size = T.int64() num_pages = T.int64() pages_elem_offset = T.int64() position_map_elem_offset = T.int64() diff --git a/python/tvm/s_tir/tensor_intrin/dot_product_common.py b/python/tvm/s_tir/tensor_intrin/dot_product_common.py index 7272477406..74b1acebf0 100644 --- a/python/tvm/s_tir/tensor_intrin/dot_product_common.py +++ b/python/tvm/s_tir/tensor_intrin/dot_product_common.py @@ -56,7 +56,7 @@ def get_dp4a_intrin(dtype_a, dtype_b, dtype_c): "__dp4a", A.vload([0], vec_type_a), B.vload([0], vec_type_b), - T.uint32(0) if dtype_c == "uint32" else T.int32(0), + T.uint32(0) if T.constexpr(dtype_c == "uint32") else T.int32(0), dtype=dtype_c, ) diff --git a/python/tvm/s_tir/tensor_intrin/metal.py b/python/tvm/s_tir/tensor_intrin/metal.py index 1750338044..8c63e9b46f 100644 --- a/python/tvm/s_tir/tensor_intrin/metal.py +++ b/python/tvm/s_tir/tensor_intrin/metal.py @@ -93,7 +93,7 @@ def get_simdgroup_load_intrin( for i, j in T.grid(col, row): with T.sblock("load"): vii, vjj = T.axis.remap("SS", [i, j]) - if transpose_matrix: + if T.constexpr(transpose_matrix): # C[vii, vjj] = A[vjj, vii] C[vjj, vii] = A[vii, vjj] else: @@ -157,7 +157,7 @@ def get_simdgroup_store_intrin( for i, j in T.grid(col, row): with T.sblock("store"): vii, vjj = T.axis.remap("SS", [i, j]) - if transpose_matrix: + if T.constexpr(transpose_matrix): C[vjj, vii] = A[vii, vjj] else: C[vii, vjj] = A[vii, vjj] diff --git a/python/tvm/s_tir/tensor_intrin/rocm.py b/python/tvm/s_tir/tensor_intrin/rocm.py index 29749dd443..60b94ce192 100644 --- a/python/tvm/s_tir/tensor_intrin/rocm.py +++ b/python/tvm/s_tir/tensor_intrin/rocm.py @@ -336,8 +336,8 @@ def get_mfma_intrin(k_dim, in_dtype="float32", out_dtype="float32", b_transposed T.launch_thread(tx, WARP_SIZE) C[tx, T.ramp(0, 1, local_size_out)] = T.call_llvm_pure_intrin( T.llvm_lookup_intrinsic_id(mfma_intrin), - A[tx, T.ramp(0, 1, local_size) if local_size > 1 else 0], - B[tx, T.ramp(0, 1, local_size) if local_size > 1 else 0], + A[tx, T.ramp(0, 1, local_size) if T.constexpr(local_size > 1) else 0], + B[tx, T.ramp(0, 1, local_size) if T.constexpr(local_size > 1) else 0], C[tx, T.ramp(0, 1, local_size_out)], T.int32(0), T.int32(0), @@ -366,12 +366,12 @@ def get_mfma_intrin(k_dim, in_dtype="float32", out_dtype="float32", b_transposed T.call_intrin( "int32", "tirx.reinterpret", - A[tx, T.ramp(0, 1, local_size) if local_size > 1 else 0], + A[tx, T.ramp(0, 1, local_size) if T.constexpr(local_size > 1) else 0], ), T.call_intrin( "int32", "tirx.reinterpret", - B[tx, T.ramp(0, 1, local_size) if local_size > 1 else 0], + B[tx, T.ramp(0, 1, local_size) if T.constexpr(local_size > 1) else 0], ), C[tx, T.ramp(0, 1, local_size_out)], T.int32(0), diff --git a/tests/python/codegen/test_target_codegen_vulkan.py b/tests/python/codegen/test_target_codegen_vulkan.py index d3213b4dcb..9faa2eb250 100644 --- a/tests/python/codegen/test_target_codegen_vulkan.py +++ b/tests/python/codegen/test_target_codegen_vulkan.py @@ -488,7 +488,7 @@ def test_cooperative_matrix(out_dtype): v_j_o = T.axis.spatial(1, 0) T.reads() T.writes(compute_wmma_accumulator[0:16, 0:16]) - C = T.match_buffer(compute_wmma_accumulator[0:16, 0:16], (16, 16), out_dtype, strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + C = T.match_buffer(compute_wmma_accumulator[0:16, 0:16], (16, 16), out_dtype, strides=("C_0_s0", "C_0_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 k_0 in range(2): for ax0_ax1_fused_0 in range(2): @@ -516,8 +516,8 @@ def test_cooperative_matrix(out_dtype): v1_o = T.axis.spatial(2, k_0 + ax1_0) T.reads(X_shared[0:16, v1_o * 16:v1_o * 16 + 16]) T.writes(X_shared_wmma_matrix_a[0:16, v1_o * 16:v1_o * 16 + 16]) - A = T.match_buffer(X_shared[0:16, v1_o * 16:v1_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="shared", offset_factor=16) - C = T.match_buffer(X_shared_wmma_matrix_a[0:16, v1_o * 16:v1_o * 16 + 16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_a", offset_factor=16) + A = T.match_buffer(X_shared[0:16, v1_o * 16:v1_o * 16 + 16], (16, 16), "float16", strides=("A_0_s0", "A_0_s1"), scope="shared", offset_factor=16) + C = T.match_buffer(X_shared_wmma_matrix_a[0:16, v1_o * 16:v1_o * 16 + 16], (16, 16), "float16", strides=("C_1_s0", "C_1_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(1): for ax1_0 in T.unroll(1): @@ -526,8 +526,8 @@ def test_cooperative_matrix(out_dtype): v1_o = T.axis.spatial(1, ax1_0) T.reads(W_shared[v0_o * 16:v0_o * 16 + 16, 0:16]) T.writes(W_shared_wmma_matrix_b[v0_o * 16:v0_o * 16 + 16, 0:16]) - A = T.match_buffer(W_shared[v0_o * 16:v0_o * 16 + 16, 0:16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="shared", offset_factor=16) - C = T.match_buffer(W_shared_wmma_matrix_b[v0_o * 16:v0_o * 16 + 16, 0:16], (16, 16), "float16", strides=("C_s0", "C_s1"), scope="wmma.matrix_b", offset_factor=16) + A = T.match_buffer(W_shared[v0_o * 16:v0_o * 16 + 16, 0:16], (16, 16), "float16", strides=("A_1_s0", "A_1_s1"), scope="shared", offset_factor=16) + C = T.match_buffer(W_shared_wmma_matrix_b[v0_o * 16:v0_o * 16 + 16, 0:16], (16, 16), "float16", strides=("C_2_s0", "C_2_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], "row_major") with T.sblock("compute_update_o"): v_i_o = T.axis.spatial(1, 0) @@ -535,17 +535,17 @@ def test_cooperative_matrix(out_dtype): v_k_o = T.axis.reduce(2, k_0) T.reads(compute_wmma_accumulator[0:16, 0:16], X_shared_wmma_matrix_a[0:16, v_k_o * 16:v_k_o * 16 + 16], W_shared_wmma_matrix_b[v_k_o * 16:v_k_o * 16 + 16, 0:16]) T.writes(compute_wmma_accumulator[0:16, 0:16]) - A = T.match_buffer(X_shared_wmma_matrix_a[0:16, v_k_o * 16:v_k_o * 16 + 16], (16, 16), "float16", strides=("A_s0", "A_s1"), scope="wmma.matrix_a", offset_factor=16) - B = T.match_buffer(W_shared_wmma_matrix_b[v_k_o * 16:v_k_o * 16 + 16, 0:16], (16, 16), "float16", strides=("B_s0", "B_s1"), scope="wmma.matrix_b", offset_factor=16) - C = T.match_buffer(compute_wmma_accumulator[0:16, 0:16], (16, 16), out_dtype, strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + A = T.match_buffer(X_shared_wmma_matrix_a[0:16, v_k_o * 16:v_k_o * 16 + 16], (16, 16), "float16", strides=("A_2_s0", "A_2_s1"), scope="wmma.matrix_a", offset_factor=16) + B = T.match_buffer(W_shared_wmma_matrix_b[v_k_o * 16:v_k_o * 16 + 16, 0:16], (16, 16), "float16", strides=("B_0_s0", "B_0_s1"), scope="wmma.matrix_b", offset_factor=16) + C = T.match_buffer(compute_wmma_accumulator[0:16, 0:16], (16, 16), out_dtype, strides=("C_3_s0", "C_3_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) with T.sblock("compute_wmma.accumulator_o"): v0_o = T.axis.spatial(1, 0) v1_o = T.axis.spatial(1, 0) T.reads(compute_wmma_accumulator[0:16, 0:16]) T.writes(compute[0:16, 0:16]) - A = T.match_buffer(compute_wmma_accumulator[0:16, 0:16], (16, 16), out_dtype, strides=("A_s0", "A_s1"), scope="wmma.accumulator", offset_factor=16) - C = T.match_buffer(compute[0:16, 0:16], (16, 16), out_dtype, strides=("C_s0", "C_s1"), offset_factor=16) + A = T.match_buffer(compute_wmma_accumulator[0:16, 0:16], (16, 16), out_dtype, strides=("A_3_s0", "A_3_s1"), scope="wmma.accumulator", offset_factor=16) + C = T.match_buffer(compute[0:16, 0:16], (16, 16), out_dtype, strides=("C_4_s0", "C_4_s1"), 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(out_dtype), C.data, C.elem_offset, C.strides[0] * 16, 2), C.strides[0], "row_major") # fmt: on diff --git a/tests/python/relax/test_analysis_type_analysis.py b/tests/python/relax/test_analysis_type_analysis.py index b0b6a54aa0..ca6a595320 100644 --- a/tests/python/relax/test_analysis_type_analysis.py +++ b/tests/python/relax/test_analysis_type_analysis.py @@ -649,8 +649,6 @@ def test_prim_type_lca(test_case): def _normalize_ty(ty): if isinstance(ty, tvm.relax.Type): return ty - elif isinstance(ty, tvm.script.parser.relax.entry.TypeProxy): - return ty.as_ty() elif callable(ty): return ty() else: diff --git a/tests/python/relax/test_frontend_onnx.py b/tests/python/relax/test_frontend_onnx.py index d7aa987c05..0efb8030a6 100644 --- a/tests/python/relax/test_frontend_onnx.py +++ b/tests/python/relax/test_frontend_onnx.py @@ -40,7 +40,7 @@ from onnx import ModelProto, TensorProto, helper, numpy_helper import tvm import tvm.testing -from tvm import relax +from tvm import relax, tirx from tvm.relax.frontend.onnx import from_onnx from tvm.script import ir as I from tvm.script import relax as R @@ -1053,15 +1053,14 @@ def _make_expected_broadcast_ir_min( Returns: Expected IR module for the Min operation. """ - output_shape = (x_shape[0], 4) @I.ir_module class ExpectedMin: @R.function def main( - x: R.Tensor(x_shape, dtype="float32"), - y: R.Tensor(y_shape, dtype="float32"), - ) -> R.Tensor(output_shape, dtype="float32"): + x: R.Tensor(("n", x_shape[1]), dtype="float32"), + y: R.Tensor(("n", y_shape[1]), dtype="float32"), + ) -> R.Tensor(("n", 4), dtype="float32"): n = T.int64() R.func_attr({"num_input": 2}) with R.dataflow(): @@ -1088,15 +1087,14 @@ def _make_expected_broadcast_ir_max( Returns: Expected IR module for the Max operation. """ - output_shape = (x_shape[0], 4) @I.ir_module class ExpectedMax: @R.function def main( - x: R.Tensor(x_shape, dtype="float32"), - y: R.Tensor(y_shape, dtype="float32"), - ) -> R.Tensor(output_shape, dtype="float32"): + x: R.Tensor(("n", x_shape[1]), dtype="float32"), + y: R.Tensor(("n", y_shape[1]), dtype="float32"), + ) -> R.Tensor(("n", 4), dtype="float32"): n = T.int64() R.func_attr({"num_input": 2}) with R.dataflow(): @@ -6846,7 +6844,7 @@ def _make_reduce_expected_ir( def expected_input_shape(shape): if not dynamic: return tuple(shape) - return tuple(f"reduce_dim_{i}" for i in range(len(shape))) + return tuple(tirx.Var(f"reduce_dim_{i}", "int64") for i in range(len(shape))) axis = None if not axes else tuple(axes) parser_vars = { @@ -8784,7 +8782,7 @@ def test_split(): shape = shape_tuple(shape) if not dynamic: return shape - return tuple(f"split_input_dim_{i}" for i in range(len(shape))) + return tuple(tirx.Var(f"split_input_dim_{i}", "int64") for i in range(len(shape))) dtype = np.dtype(fp_arith).name input_shape = expected_input_shape(indata_shape) @@ -9078,7 +9076,9 @@ def test_tile_dynamic_repeats(): def make_expected(dynamic_input, in_shape): rank = len(in_shape) input_shape = ( - tuple(f"tile_data_dim_{i}" for i in range(rank)) if dynamic_input else tuple(in_shape) + tuple(tirx.Var(f"tile_data_dim_{i}", "int64") for i in range(rank)) + if dynamic_input + else tuple(in_shape) ) if rank == 2: diff --git a/tests/python/s_tir/dlight/test_gpu_matmul_tensorize.py b/tests/python/s_tir/dlight/test_gpu_matmul_tensorize.py index d12d713d3b..f3b971d0e6 100644 --- a/tests/python/s_tir/dlight/test_gpu_matmul_tensorize.py +++ b/tests/python/s_tir/dlight/test_gpu_matmul_tensorize.py @@ -65,7 +65,8 @@ def test_matmul_tensorize(): 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) + C_s0, C_s1 = T.int32(), T.int32() + 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)) for ax3_0_0 in range(4, 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(4): @@ -101,8 +102,10 @@ def test_matmul_tensorize(): v2_o = T.axis.spatial(16, ax3_0_0 * 4 + 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) + A_s0, A_s1 = T.int32(), T.int32() + 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_1_s0, C_1_s1 = T.int32(), T.int32() + 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_1_s0, C_1_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): @@ -112,8 +115,10 @@ def test_matmul_tensorize(): v2_o = T.axis.spatial(16, ax3_0_0 * 4 + 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) + A_1_s0, A_1_s1 = T.int32(), T.int32() + 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_1_s0, A_1_s1), scope="shared.dyn", offset_factor=16) + C_2_s0, C_2_s1 = T.int32(), T.int32() + 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_2_s0, C_2_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.sblock("compute_o_update"): @@ -129,9 +134,12 @@ def test_matmul_tensorize(): 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) + A_2_s0, A_2_s1 = T.int32(), T.int32() + 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_2_s0, A_2_s1), scope="wmma.matrix_a", offset_factor=16) + B_s0, B_s1 = T.int32(), T.int32() + 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_3_s0, C_3_s1 = T.int32(), T.int32() + 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_3_s0, C_3_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.sblock("compute_reindex_shared.dyn_wmma.accumulator_o"): @@ -140,8 +148,10 @@ def test_matmul_tensorize(): 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), "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) + A_3_s0, A_3_s1 = T.int32(), T.int32() + 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_3_s0, A_3_s1), scope="wmma.accumulator", offset_factor=16) + C_4_s0, C_4_s1 = T.int32(), T.int32() + 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_4_s0, C_4_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(8): for ax0_ax1_fused_1 in T.thread_binding(32, thread="threadIdx.x"): @@ -329,7 +339,8 @@ def test_matmul_tensorize_epilogue(): 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) + C_s0, C_s1 = T.int32(), T.int32() + 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)) for ax3_0_0 in range(32, 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(4): @@ -365,8 +376,10 @@ def test_matmul_tensorize_epilogue(): v2_o = T.axis.spatial(128, ax3_0_0 * 4 + 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) + A_s0, A_s1 = T.int32(), T.int32() + 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_1_s0, C_1_s1 = T.int32(), T.int32() + 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_1_s0, C_1_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): @@ -376,8 +389,10 @@ def test_matmul_tensorize_epilogue(): v2_o = T.axis.spatial(128, ax3_0_0 * 4 + 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) + A_1_s0, A_1_s1 = T.int32(), T.int32() + 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_1_s0, A_1_s1), scope="shared.dyn", offset_factor=16) + C_2_s0, C_2_s1 = T.int32(), T.int32() + 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_2_s0, C_2_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.sblock("NT_matmul_o_update"): @@ -393,9 +408,12 @@ def test_matmul_tensorize_epilogue(): 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) + A_2_s0, A_2_s1 = T.int32(), T.int32() + 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_2_s0, A_2_s1), scope="wmma.matrix_a", offset_factor=16) + B_s0, B_s1 = T.int32(), T.int32() + 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_3_s0, C_3_s1 = T.int32(), T.int32() + 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_3_s0, C_3_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.sblock("var_NT_matmul_intermediate_reindex_pad_shared.dyn_wmma.accumulator_o"): @@ -404,8 +422,10 @@ def test_matmul_tensorize_epilogue(): 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(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) + A_3_s0, A_3_s1 = T.int32(), T.int32() + 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_3_s0, A_3_s1), scope="wmma.accumulator", offset_factor=16) + C_4_s0, C_4_s1 = T.int32(), T.int32() + 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_4_s0, C_4_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(8): for ax0_ax1_fused_1 in T.thread_binding(32, thread="threadIdx.x"): @@ -468,7 +488,8 @@ def test_matmul_int8_tensorize(): 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) + C_s0, C_s1 = T.int32(), T.int32() + 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)) 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): @@ -504,8 +525,10 @@ def test_matmul_int8_tensorize(): 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) + A_s0, A_s1 = T.int32(), T.int32() + 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_1_s0, C_1_s1 = T.int32(), T.int32() + 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_1_s0, C_1_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): @@ -515,8 +538,10 @@ def test_matmul_int8_tensorize(): 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) + A_1_s0, A_1_s1 = T.int32(), T.int32() + 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_1_s0, A_1_s1), scope="shared.dyn", offset_factor=16) + C_2_s0, C_2_s1 = T.int32(), T.int32() + 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_2_s0, C_2_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.sblock("compute_o_update"): @@ -532,9 +557,12 @@ def test_matmul_int8_tensorize(): 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) + A_2_s0, A_2_s1 = T.int32(), T.int32() + 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_2_s0, A_2_s1), scope="wmma.matrix_a", offset_factor=16) + B_s0, B_s1 = T.int32(), T.int32() + 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_3_s0, C_3_s1 = T.int32(), T.int32() + 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_3_s0, C_3_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.sblock("compute_reindex_shared.dyn_wmma.accumulator_o"): @@ -543,8 +571,10 @@ def test_matmul_int8_tensorize(): 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) + A_3_s0, A_3_s1 = T.int32(), T.int32() + 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_3_s0, A_3_s1), scope="wmma.accumulator", offset_factor=16) + C_4_s0, C_4_s1 = T.int32(), T.int32() + 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_4_s0, C_4_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(8): for ax0_ax1_fused_1 in T.thread_binding(32, thread="threadIdx.x"): @@ -612,7 +642,8 @@ def test_matmul_int8_tensorize_3d2d_dyn(): 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) + C_s0, C_s1 = T.int32(), T.int32() + 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)) 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): @@ -648,8 +679,10 @@ def test_matmul_int8_tensorize_3d2d_dyn(): 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) + A_s0, A_s1 = T.int32(), T.int32() + 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_1_s0, C_1_s1 = T.int32(), T.int32() + 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_1_s0, C_1_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): @@ -659,8 +692,10 @@ def test_matmul_int8_tensorize_3d2d_dyn(): 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) + A_1_s0, A_1_s1 = T.int32(), T.int32() + 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_1_s0, A_1_s1), scope="shared.dyn", offset_factor=16) + C_2_s0, C_2_s1 = T.int32(), T.int32() + 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_2_s0, C_2_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.sblock("matmul_o_update"): @@ -676,9 +711,12 @@ def test_matmul_int8_tensorize_3d2d_dyn(): 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) + A_2_s0, A_2_s1 = T.int32(), T.int32() + 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_2_s0, A_2_s1), scope="wmma.matrix_a", offset_factor=16) + B_s0, B_s1 = T.int32(), T.int32() + 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_3_s0, C_3_s1 = T.int32(), T.int32() + 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_3_s0, C_3_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 [...] for ax0_0, ax1_0 in T.grid(2, 2): with T.sblock("matmul_1_reindex_pad_shared.dyn_wmma.accumulator_o"): @@ -687,8 +725,10 @@ def test_matmul_int8_tensorize_3d2d_dyn(): 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) + A_3_s0, A_3_s1 = T.int32(), T.int32() + 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_3_s0, A_3_s1), scope="wmma.accumulator", offset_factor=16) + C_4_s0, C_4_s1 = T.int32(), T.int32() + 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_4_s0, C_4_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(8): for ax0_ax1_fused_1 in T.thread_binding(32, thread="threadIdx.x"): @@ -754,7 +794,8 @@ def test_matmul_metal(): v2_o = T.axis.spatial(3584, ax2_0 * 8 + ax2_1 * 2 + ax2_2_init + ax2_3_init_0) T.reads() T.writes(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("A_s0", "A_s1"), scope="metal.simdgroup", offset_factor=1) + A_s0, A_s1 = T.int32(), T.int32() + A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_s0, A_s1), scope="metal.simdgroup", offset_factor=1) T.metal.make_filled_simdgroup_matrix(A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, T.float32(0), 8, 8) for ax3_0 in range(128): for ax0_1, ax1_ax2_fused_0 in T.grid(1, 1): @@ -789,8 +830,10 @@ def test_matmul_metal(): v2_o = T.axis.spatial(512, ax3_0 * 4 + ax3_1 + ax1_0_1) T.reads(A_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) T.writes(A_reindex_pad_shared_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_1 = T.match_buffer(A_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("A_s0", "A_s1"), scope="shared", offset_factor=1) - C_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("C_s0", "C_s1"), scope="metal.simdgroup", offset_factor=1) + A_1_s0, A_1_s1 = T.int32(), T.int32() + A_1 = T.match_buffer(A_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_1_s0, A_1_s1), scope="shared", offset_factor=1) + C_s0, C_s1 = T.int32(), T.int32() + C_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_s0, C_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_load(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), A_1.data, A_1.elem_offset, A_1.strides[0] * 8, 1), A_1.strides[0], 8, 8, T.bool(False)) for ax0_0, ax1_0_1 in T.grid(2, 1): with T.sblock("B_reindex_shared_metal.simdgroup_o"): @@ -799,8 +842,10 @@ def test_matmul_metal(): v2_o = T.axis.spatial(512, ax3_0 * 4 + ax3_1 + ax1_0_1) T.reads(B_reindex_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) T.writes(B_reindex_shared_metal_simdgroup[v0_o, v2_o * 8:v2_o * 8 + 8, v1_o * 8:v1_o * 8 + 8]) - A_1 = T.match_buffer(B_reindex_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("A_s0", "A_s1"), scope="shared", offset_factor=1) - C_1 = T.match_buffer(B_reindex_shared_metal_simdgroup[v0_o, v2_o * 8:v2_o * 8 + 8, v1_o * 8:v1_o * 8 + 8], (8, 8), "float16", strides=("C_s0", "C_s1"), scope="metal.simdgroup", offset_factor=1) + A_2_s0, A_2_s1 = T.int32(), T.int32() + A_1 = T.match_buffer(B_reindex_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_2_s0, A_2_s1), scope="shared", offset_factor=1) + C_1_s0, C_1_s1 = T.int32(), T.int32() + C_1 = T.match_buffer(B_reindex_shared_metal_simdgroup[v0_o, v2_o * 8:v2_o * 8 + 8, v1_o * 8:v1_o * 8 + 8], (8, 8), "float16", strides=(C_1_s0, C_1_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_load(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), A_1.data, A_1.elem_offset, A_1.strides[0] * 8, 1), A_1.strides[0], 8, 8, T.bool(True)) for ax1_2, ax2_2 in T.grid(2, 2): with T.sblock("C_update_o"): @@ -810,9 +855,12 @@ def test_matmul_metal(): v3_o = T.axis.reduce(512, ax3_0 * 4 + ax3_1) T.reads(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], A_reindex_pad_shared_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v3_o * 8:v3_o * 8 + 8], B_reindex_shared_metal_simdgroup[0, v3_o * 8:v3_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) T.writes(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v3_o * 8:v3_o * 8 + 8], (8, 8), "float16", strides=("A_s0", "A_s1"), scope="metal.simdgroup", offset_factor=1) - B_1 = T.match_buffer(B_reindex_shared_metal_simdgroup[0, v3_o * 8:v3_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("B_s0", "B_s1"), scope="metal.simdgroup", offset_factor=1) - C_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("C_s0", "C_s1"), scope="metal.simdgroup", offset_factor=1) + A_3_s0, A_3_s1 = T.int32(), T.int32() + A_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v3_o * 8:v3_o * 8 + 8], (8, 8), "float16", strides=(A_3_s0, A_3_s1), scope="metal.simdgroup", offset_factor=1) + B_s0, B_s1 = T.int32(), T.int32() + B_1 = T.match_buffer(B_reindex_shared_metal_simdgroup[0, v3_o * 8:v3_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(B_s0, B_s1), scope="metal.simdgroup", offset_factor=1) + C_2_s0, C_2_s1 = T.int32(), T.int32() + C_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_2_s0, C_2_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_multiply_accumulate(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, B_1.data, B_1.elem_offset // B_1.strides[0] // 8 * (B_1.strides[0] // 8) + B_1.elem_offset % B_1.strides[0] // 8, C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] / [...] for ax0_1, ax1_0_1, ax2_0_1 in T.grid(1, 2, 2): with T.sblock("C_reindex_pad_metal.simdgroup_o"): @@ -821,8 +869,10 @@ def test_matmul_metal(): v2_o = T.axis.spatial(3584, ax2_0 * 8 + ax2_1 * 2 + ax2_0_1) T.reads(C_reindex_pad_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) T.writes(C_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("A_s0", "A_s1"), scope="metal.simdgroup", offset_factor=1) - C_1 = T.match_buffer(C_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("C_s0", "C_s1"), scope="shared", offset_factor=1) + A_4_s0, A_4_s1 = T.int32(), T.int32() + A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_4_s0, A_4_s1), scope="metal.simdgroup", offset_factor=1) + C_3_s0, C_3_s1 = T.int32(), T.int32() + C_1 = T.match_buffer(C_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_3_s0, C_3_s1), scope="shared", offset_factor=1) T.metal.simdgroup_store(A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), C_1.data, C_1.elem_offset, C_1.strides[0] * 8, 2), C_1.strides[0], 8, 8, T.bool(False)) for ax0_1, ax1_ax2_fused_0 in T.grid(1, 2): for ax1_ax2_fused_1 in T.thread_binding(4, thread="threadIdx.z"): @@ -899,7 +949,8 @@ def test_matmul_metal_int4_quant(): v2_o = T.axis.spatial(3584, ax2_0 * 8 + ax2_1 * 2 + ax2_2_init + ax2_3_init_0) T.reads() T.writes(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("A_s0", "A_s1"), scope="metal.simdgroup", offset_factor=1) + A_s0, A_s1 = T.int32(), T.int32() + A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_s0, A_s1), scope="metal.simdgroup", offset_factor=1) T.metal.make_filled_simdgroup_matrix(A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, T.float32(0), 8, 8) for ax3_0 in range(128): for ax0_1, ax1_ax2_fused_0 in T.grid(1, 1): @@ -934,8 +985,10 @@ def test_matmul_metal_int4_quant(): v2_o = T.axis.spatial(512, ax3_0 * 4 + ax3_1 + ax1_0_1) T.reads(A_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) T.writes(A_reindex_pad_shared_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_1 = T.match_buffer(A_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("A_s0", "A_s1"), scope="shared", offset_factor=1) - C_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("C_s0", "C_s1"), scope="metal.simdgroup", offset_factor=1) + A_1_s0, A_1_s1 = T.int32(), T.int32() + A_1 = T.match_buffer(A_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_1_s0, A_1_s1), scope="shared", offset_factor=1) + C_s0, C_s1 = T.int32(), T.int32() + C_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_s0, C_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_load(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), A_1.data, A_1.elem_offset, A_1.strides[0] * 8, 1), A_1.strides[0], 8, 8, T.bool(False)) for ax0_0, ax1_0_1 in T.grid(2, 1): with T.sblock("B_reindex_shared_metal.simdgroup_o"): @@ -944,8 +997,10 @@ def test_matmul_metal_int4_quant(): v2_o = T.axis.spatial(512, ax3_0 * 4 + ax3_1 + ax1_0_1) T.reads(B_reindex_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) T.writes(B_reindex_shared_metal_simdgroup[v0_o, v2_o * 8:v2_o * 8 + 8, v1_o * 8:v1_o * 8 + 8]) - A_1 = T.match_buffer(B_reindex_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("A_s0", "A_s1"), scope="shared", offset_factor=1) - C_1 = T.match_buffer(B_reindex_shared_metal_simdgroup[v0_o, v2_o * 8:v2_o * 8 + 8, v1_o * 8:v1_o * 8 + 8], (8, 8), "float16", strides=("C_s0", "C_s1"), scope="metal.simdgroup", offset_factor=1) + A_2_s0, A_2_s1 = T.int32(), T.int32() + A_1 = T.match_buffer(B_reindex_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_2_s0, A_2_s1), scope="shared", offset_factor=1) + C_1_s0, C_1_s1 = T.int32(), T.int32() + C_1 = T.match_buffer(B_reindex_shared_metal_simdgroup[v0_o, v2_o * 8:v2_o * 8 + 8, v1_o * 8:v1_o * 8 + 8], (8, 8), "float16", strides=(C_1_s0, C_1_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_load(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), A_1.data, A_1.elem_offset, A_1.strides[0] * 8, 1), A_1.strides[0], 8, 8, T.bool(True)) for ax1_2, ax2_2 in T.grid(2, 2): with T.sblock("NT_matmul_update_o"): @@ -955,9 +1010,12 @@ def test_matmul_metal_int4_quant(): v3_o = T.axis.reduce(512, ax3_0 * 4 + ax3_1) T.reads(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], A_reindex_pad_shared_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v3_o * 8:v3_o * 8 + 8], B_reindex_shared_metal_simdgroup[0, v3_o * 8:v3_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) T.writes(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v3_o * 8:v3_o * 8 + 8], (8, 8), "float16", strides=("A_s0", "A_s1"), scope="metal.simdgroup", offset_factor=1) - B = T.match_buffer(B_reindex_shared_metal_simdgroup[0, v3_o * 8:v3_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("B_s0", "B_s1"), scope="metal.simdgroup", offset_factor=1) - C_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("C_s0", "C_s1"), scope="metal.simdgroup", offset_factor=1) + A_3_s0, A_3_s1 = T.int32(), T.int32() + A_1 = T.match_buffer(A_reindex_pad_shared_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v3_o * 8:v3_o * 8 + 8], (8, 8), "float16", strides=(A_3_s0, A_3_s1), scope="metal.simdgroup", offset_factor=1) + B_s0, B_s1 = T.int32(), T.int32() + B = T.match_buffer(B_reindex_shared_metal_simdgroup[0, v3_o * 8:v3_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(B_s0, B_s1), scope="metal.simdgroup", offset_factor=1) + C_2_s0, C_2_s1 = T.int32(), T.int32() + C_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[0, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_2_s0, C_2_s1), scope="metal.simdgroup", offset_factor=1) T.metal.simdgroup_multiply_accumulate(C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.elem_offset % C_1.strides[0] // 8, A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, B.data, B.elem_offset // B.strides[0] // 8 * (B.strides[0] // 8) + B.elem_offset % B.strides[0] // 8, C_1.data, C_1.elem_offset // C_1.strides[0] // 8 * (C_1.strides[0] // 8) + C_1.e [...] for ax0_1, ax1_0_1, ax2_0_1 in T.grid(1, 2, 2): with T.sblock("C_reindex_pad_metal.simdgroup_o"): @@ -966,8 +1024,10 @@ def test_matmul_metal_int4_quant(): v2_o = T.axis.spatial(3584, ax2_0 * 8 + ax2_1 * 2 + ax2_0_1) T.reads(C_reindex_pad_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) T.writes(C_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8]) - A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("A_s0", "A_s1"), scope="metal.simdgroup", offset_factor=1) - C_1 = T.match_buffer(C_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=("C_s0", "C_s1"), scope="shared", offset_factor=1) + A_4_s0, A_4_s1 = T.int32(), T.int32() + A_1 = T.match_buffer(C_reindex_pad_metal_simdgroup[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(A_4_s0, A_4_s1), scope="metal.simdgroup", offset_factor=1) + C_3_s0, C_3_s1 = T.int32(), T.int32() + C_1 = T.match_buffer(C_reindex_pad_shared[v0_o, v1_o * 8:v1_o * 8 + 8, v2_o * 8:v2_o * 8 + 8], (8, 8), "float16", strides=(C_3_s0, C_3_s1), scope="shared", offset_factor=1) T.metal.simdgroup_store(A_1.data, A_1.elem_offset // A_1.strides[0] // 8 * (A_1.strides[0] // 8) + A_1.elem_offset % A_1.strides[0] // 8, T.tvm_access_ptr(T.type_annotation("float16"), C_1.data, C_1.elem_offset, C_1.strides[0] * 8, 2), C_1.strides[0], 8, 8, T.bool(False)) for ax0_1, ax1_ax2_fused_0 in T.grid(1, 2): for ax1_ax2_fused_1 in T.thread_binding(4, thread="threadIdx.z"): diff --git a/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py b/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py index 42ec6b8384..09695063e6 100644 --- a/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py +++ b/tests/python/s_tir/meta_schedule/test_meta_schedule_trace_apply.py @@ -689,7 +689,7 @@ class Conv2dInt8_tensorcore_scheduled: T.reads(pad_temp_reindex_shared[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16]) T.writes(pad_temp_reindex_shared_wmma_matrix_a[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16]) A = T.match_buffer(pad_temp_reindex_shared[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int8", strides=("A_s0", "A_s1"), scope="shared", offset_factor=16) - C = T.match_buffer(pad_temp_reindex_shared_wmma_matrix_a[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int8", strides=("C_s0", "C_s1"), scope="wmma.matrix_a", offset_factor=16) + C = T.match_buffer(pad_temp_reindex_shared_wmma_matrix_a[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int8", strides=("C_1_s0", "C_1_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, ax1, ax2_0, ax3_0 in T.grid(1, 1, 1, 2): with T.sblock("p1_reindex_shared_wmma.matrix_b_o"): @@ -698,8 +698,8 @@ class Conv2dInt8_tensorcore_scheduled: v3_o = T.axis.spatial(4, ax4_0_0 * 2 + ax3_0) T.reads(p1_reindex_shared[v0_o, v1_o, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) T.writes(p1_reindex_shared_wmma_matrix_b[v0_o, v1_o, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) - A = T.match_buffer(p1_reindex_shared[v0_o, v1_o, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=("A_s0", "A_s1"), scope="shared", offset_factor=16) - C = T.match_buffer(p1_reindex_shared_wmma_matrix_b[v0_o, v1_o, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=("C_s0", "C_s1"), scope="wmma.matrix_b", offset_factor=16) + A = T.match_buffer(p1_reindex_shared[v0_o, v1_o, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=("A_1_s0", "A_1_s1"), scope="shared", offset_factor=16) + C = T.match_buffer(p1_reindex_shared_wmma_matrix_b[v0_o, v1_o, v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int8", strides=("C_2_s0", "C_2_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 ax2_0_3, ax3_0_3, ax0_2, ax1_2, ax4_0_2, ax2_0_4, ax3_0_4 in T.grid(1, 1, 1, 1, 2, 1, 1): with T.sblock("conv2d_nhwc_o_update"): @@ -711,9 +711,9 @@ class Conv2dInt8_tensorcore_scheduled: T.reads(conv2d_nhwc_reindex_shared_wmma_accumulator[v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], pad_temp_reindex_shared_wmma_matrix_a[v2_o * 16:v2_o * 16 + 16, v4_o * 16:v4_o * 16 + 16], p1_reindex_shared_wmma_matrix_b[v0_o, v1_o, v3_o * 16:v3_o * 16 + 16, v4_o * 16:v4_o * 16 + 16]) T.writes(conv2d_nhwc_reindex_shared_wmma_accumulator[v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16]) T.sblock_attr({"meta_schedule.thread_extent_high_inclusive": 1024, "meta_schedule.thread_extent_low_inclusive": 32, "warp_execution": 1}) - A = T.match_buffer(pad_temp_reindex_shared_wmma_matrix_a[v2_o * 16:v2_o * 16 + 16, v4_o * 16:v4_o * 16 + 16], (16, 16), "int8", strides=("A_s0", "A_s1"), scope="wmma.matrix_a", offset_factor=16) + A = T.match_buffer(pad_temp_reindex_shared_wmma_matrix_a[v2_o * 16:v2_o * 16 + 16, v4_o * 16:v4_o * 16 + 16], (16, 16), "int8", strides=("A_2_s0", "A_2_s1"), scope="wmma.matrix_a", offset_factor=16) B = T.match_buffer(p1_reindex_shared_wmma_matrix_b[v0_o, v1_o, v3_o * 16:v3_o * 16 + 16, v4_o * 16:v4_o * 16 + 16], (16, 16), "int8", strides=("B_s0", "B_s1"), scope="wmma.matrix_b", offset_factor=16) - C = T.match_buffer(conv2d_nhwc_reindex_shared_wmma_accumulator[v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int32", strides=("C_s0", "C_s1"), scope="wmma.accumulator", offset_factor=16) + C = T.match_buffer(conv2d_nhwc_reindex_shared_wmma_accumulator[v2_o * 16:v2_o * 16 + 16, v3_o * 16:v3_o * 16 + 16], (16, 16), "int32", strides=("C_3_s0", "C_3_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(1, 1): with T.sblock("conv2d_nhwc_reindex_shared_wmma.accumulator_o"): @@ -721,8 +721,8 @@ class Conv2dInt8_tensorcore_scheduled: v1_o = T.axis.spatial(16, ax2_0_0_ax3_0_0_fused % 8 * 2 + ax2_0_2_ax3_0_2_fused % 2 + ax1_0) T.reads(conv2d_nhwc_reindex_shared_wmma_accumulator[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16]) T.writes(conv2d_nhwc_reindex_shared[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16]) - A = T.match_buffer(conv2d_nhwc_reindex_shared_wmma_accumulator[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int32", strides=("A_s0", "A_s1"), scope="wmma.accumulator", offset_factor=16) - C = T.match_buffer(conv2d_nhwc_reindex_shared[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int32", strides=("C_s0", "C_s1"), scope="shared", offset_factor=16) + A = T.match_buffer(conv2d_nhwc_reindex_shared_wmma_accumulator[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int32", strides=("A_3_s0", "A_3_s1"), scope="wmma.accumulator", offset_factor=16) + C = T.match_buffer(conv2d_nhwc_reindex_shared[v0_o * 16:v0_o * 16 + 16, v1_o * 16:v1_o * 16 + 16], (16, 16), "int32", strides=("C_4_s0", "C_4_s1"), scope="shared", 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_0 in T.grid(128, 2): for ax1_1 in T.thread_binding(16, thread="threadIdx.x"): diff --git a/tests/python/te/test_te_create_primfunc.py b/tests/python/te/test_te_create_primfunc.py index ee1bc60498..6b92d02272 100644 --- a/tests/python/te/test_te_create_primfunc.py +++ b/tests/python/te/test_te_create_primfunc.py @@ -242,9 +242,9 @@ def te_extern(): @T.prim_func(s_tir=True) def tir_extern(a: T.handle, b: T.handle, c: T.handle) -> None: T.func_attr({"global_symbol": "main", "tirx.noalias": True}) - off1 = te.var("elem_offset") - off2 = te.var("elem_offset_1") - off3 = te.var("elem_offset_2") + off1 = T.int32() + off2 = T.int32() + off3 = T.int32() A = T.match_buffer(a, (128, 128), elem_offset=off1) B = T.match_buffer(b, (128, 128), elem_offset=off2) C = T.match_buffer(c, (128, 128), elem_offset=off3) diff --git a/tests/python/tirx-transform/test_tir_transform_vectorize.py b/tests/python/tirx-transform/test_tir_transform_vectorize.py index 379bbfd4b2..73c62b8f8b 100644 --- a/tests/python/tirx-transform/test_tir_transform_vectorize.py +++ b/tests/python/tirx-transform/test_tir_transform_vectorize.py @@ -502,7 +502,7 @@ def test_illegal_extent(): class Mod: @T.prim_func(s_tir=True) def main(A: T.Buffer((25,), "int32")): - n = T.Var("n", ty="int32") + n = T.int32() for j in T.vectorized(n): A[j] = 3 diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py b/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py index b2b27bafbf..4ddcd5f6bc 100644 --- a/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py +++ b/tests/python/tirx/operator/tile_primitive/trn/test_binary_trn.py @@ -79,11 +79,13 @@ def test_simple_binary(op_type, operands_type): A_sbuf = T.alloc_buffer(src1_shape, "float32", scope="trn.sbuf", layout=src1_layout) B_sbuf = T.alloc_buffer(src2_shape, "float32", scope="trn.sbuf", layout=src2_layout) C_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) - if operands_type == "region_region" or operands_type.startswith("region_broadcast"): + if T.constexpr( + operands_type == "region_region" or operands_type.startswith("region_broadcast") + ): Tx_func(C_sbuf, A_sbuf, B_sbuf) - elif operands_type == "const_region": + elif T.constexpr(operands_type == "const_region"): Tx_func(C_sbuf, const, A_sbuf) - elif operands_type == "region_const": + elif T.constexpr(operands_type == "region_const"): Tx_func(C_sbuf, A_sbuf, const) @T.prim_func @@ -96,15 +98,15 @@ def test_simple_binary(op_type, operands_type): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): for f_loop in T.serial(0, 512, annotations={"nki_dim":"F"}): - if operands_type == "region_region": + if T.constexpr(operands_type == "region_region"): T.nki.tensortensor(C_sbuf[p_loop, f_loop], A_sbuf[p_loop, f_loop], B_sbuf[p_loop, f_loop], op_type) # noqa: E501 - elif operands_type == "region_const": + elif T.constexpr(operands_type == "region_const"): T.nki.tensorscalar(C_sbuf[p_loop, f_loop], A_sbuf[p_loop, f_loop], T.float32(3.0), op_type, T.bool(False)) # noqa: E501 - elif operands_type == "const_region": + elif T.constexpr(operands_type == "const_region"): T.nki.tensorscalar(C_sbuf[p_loop, f_loop], A_sbuf[p_loop, f_loop], T.float32(3.0), op_type, T.bool(True)) # noqa: E501 - elif operands_type == "region_broadcast_rhs": + elif T.constexpr(operands_type == "region_broadcast_rhs"): T.nki.tensorscalar(C_sbuf[p_loop, f_loop], A_sbuf[p_loop, f_loop], B_sbuf[p_loop, 0], op_type, T.bool(False)) # noqa: E501 - elif operands_type == "region_broadcast_lhs": + elif T.constexpr(operands_type == "region_broadcast_lhs"): T.nki.tensorscalar(C_sbuf[p_loop, f_loop], B_sbuf[p_loop, f_loop], A_sbuf[p_loop, 0], op_type, T.bool(True)) # noqa: E501 # fmt: on with target: @@ -156,15 +158,15 @@ def test_binary_complex(op_type, operands_type): B_sbuf_view = B_sbuf.view(*src2_view_shape) C_sbuf_view = C_sbuf.view(*dst_view_shape) for i in range(4): - if operands_type == "region_region": + if T.constexpr(operands_type == "region_region"): Tx_func(C_sbuf_view[:, i, :], A_sbuf_view[:, i * 2, :], B_sbuf_view[:, i, :]) - elif operands_type == "region_const": + elif T.constexpr(operands_type == "region_const"): Tx_func(C_sbuf_view[:, i, :], A_sbuf_view[:, i * 2, :], const) - elif operands_type == "const_region": + elif T.constexpr(operands_type == "const_region"): Tx_func(C_sbuf_view[:, i, :], const, A_sbuf_view[:, i * 2, :]) - elif operands_type == "region_broadcast_rhs": + elif T.constexpr(operands_type == "region_broadcast_rhs"): Tx_func(C_sbuf_view[:, i, :], A_sbuf_view[:, i * 2, :], B_sbuf_view[:, 0, :]) - elif operands_type == "region_broadcast_lhs": + elif T.constexpr(operands_type == "region_broadcast_lhs"): Tx_func(C_sbuf_view[:, i, :, :], A_sbuf_view[:, i*2,:, :], B_sbuf_view[:, i, :, :]) f_extent = 128 if operands_type == "region_broadcast_lhs" else 512 @@ -183,15 +185,15 @@ def test_binary_complex(op_type, operands_type): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): for f_loop in T.serial(0, f_extent, annotations={"nki_dim":"F"}): - if operands_type == "region_region": + if T.constexpr(operands_type == "region_region"): T.nki.tensortensor(C_sbuf_view[p_loop, i * 512 + f_loop], A_sbuf_view[p_loop, i * 1024 + f_loop], B_sbuf_view[p_loop, i * 512 + f_loop], op_type) # noqa: E501 - elif operands_type == "const_region": + elif T.constexpr(operands_type == "const_region"): T.nki.tensorscalar(C_sbuf_view[p_loop, i * 512 + f_loop], A_sbuf_view[p_loop, i * 1024 + f_loop], T.float32(3.0), op_type, T.bool(True)) # noqa: E501 - elif operands_type == "region_const": + elif T.constexpr(operands_type == "region_const"): T.nki.tensorscalar(C_sbuf_view[p_loop, i * 512 + f_loop], A_sbuf_view[p_loop, i * 1024 + f_loop], T.float32(3.0), op_type, T.bool(False)) # noqa: E501 - elif operands_type == "region_broadcast_lhs": + elif T.constexpr(operands_type == "region_broadcast_lhs"): T.nki.tensorscalar(C_sbuf_view[p_loop, i * 512 + b_loop * 128 + f_loop], B_sbuf_view[p_loop, i * 512 + b_loop * 128 + f_loop], A_sbuf_view[p_loop, i * 8 + b_loop], op_type, T.bool(True)) # noqa: E501 - elif operands_type == "region_broadcast_rhs": + elif T.constexpr(operands_type == "region_broadcast_rhs"): T.nki.tensortensor(C_sbuf_view[p_loop, i * 512 + f_loop], A_sbuf_view[p_loop, i * 1024 + f_loop], B_sbuf_view[p_loop, f_loop], op_type) # noqa: E501 # fmt: on diff --git a/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py b/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py index 077200fbe3..0773ea8132 100644 --- a/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py +++ b/tests/python/tirx/operator/tile_primitive/trn/test_unary_trn.py @@ -65,7 +65,7 @@ def test_simple_unary(op_type): T.device_entry() A_sbuf = T.alloc_buffer(src_shape, "float32", scope="trn.sbuf", layout=src_layout) B_sbuf = T.alloc_buffer(dst_shape, "float32", scope="trn.sbuf", layout=dst_layout) - if op_type == "memset": + if T.constexpr(op_type == "memset"): tx_func(B_sbuf, T.float32(0.0)) else: tx_func(B_sbuf, A_sbuf) @@ -79,11 +79,11 @@ def test_simple_unary(op_type): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): for f_loop in T.serial(0, 512, annotations={"nki_dim":"F"}): - if op_type == "reciprocal": + if T.constexpr(op_type == "reciprocal"): T.nki.reciprocal( B_sbuf[p_loop, f_loop], A_sbuf[p_loop, f_loop] ) - elif op_type == "memset": + elif T.constexpr(op_type == "memset"): T.nki.memset(B_sbuf[p_loop, f_loop], 0.0) # fmt: on with target: @@ -110,7 +110,7 @@ def test_unary_in_a_loop(op_type): A_sbuf_view = A_sbuf.view(128, 8, 512) B_sbuf_view = B_sbuf.view(128, 4, 512) for i in range(4): - if op_type == "memset": + if T.constexpr(op_type == "memset"): Tx_func(B_sbuf_view[:, i, :], T.float32(0.0)) else: Tx_func(B_sbuf_view[:, i, :], A_sbuf_view[:, i * 2, :]) @@ -126,9 +126,9 @@ def test_unary_in_a_loop(op_type): T.attr(0, "tensorized_nki_instruction", 1) for p_loop in T.serial(0, 128, annotations={"nki_dim":"P"}): for f_loop in T.serial(0, 512, annotations={"nki_dim":"F"}): - if op_type == "reciprocal": + if T.constexpr(op_type == "reciprocal"): T.nki.reciprocal(B_sbuf_view[p_loop, i * 512 + f_loop], A_sbuf_view[p_loop, i * 1024 + f_loop]) # noqa: E501 - elif op_type == "memset": + elif T.constexpr(op_type == "memset"): T.nki.memset(B_sbuf_view[p_loop, i * 512 + f_loop], 0.0) # fmt: on with target: diff --git a/tests/python/tirx/test_inline.py b/tests/python/tirx/test_inline.py index 438c187c6c..4e3c8f8aa2 100644 --- a/tests/python/tirx/test_inline.py +++ b/tests/python/tirx/test_inline.py @@ -207,7 +207,7 @@ def test_recursive_inline(): @T.inline def add(x, c): - if c > 0: + if T.constexpr(c > 0): add(x, c - 1) T.evaluate(x) diff --git a/tests/python/tirx/test_jit.py b/tests/python/tirx/test_jit.py index ca9a91846f..00af56df61 100644 --- a/tests/python/tirx/test_jit.py +++ b/tests/python/tirx/test_jit.py @@ -254,7 +254,7 @@ def test_optional_param_present_and_absent_ir(): @T.jit(private=True) def kernel(a: T.Optional(T.handle), out_h: T.handle): out = T.match_buffer(out_h, (1,), "int32") - if a is not None: + if T.constexpr(a is not None): A = T.match_buffer(a, (1,), "int32") out[0] = A[0] else: @@ -285,7 +285,7 @@ def test_optional_specialization_cache_includes_presence(): @T.jit(private=True) def kernel(a: T.Optional(T.handle), out_h: T.handle): out = T.match_buffer(out_h, (1,), "int32") - if a is not None: + if T.constexpr(a is not None): A = T.match_buffer(a, (1,), "int32") out[0] = A[0] else: @@ -310,10 +310,10 @@ def test_multiple_optional_params_preserve_runtime_order(): first = T.match_buffer(first_h, (1,), "int32") out = T.match_buffer(out_h, (1,), "int32") out[0] = first[0] * scale - if a is not None: + if T.constexpr(a is not None): A = T.match_buffer(a, (1,), "int32") out[0] = out[0] + A[0] - if b is not None: + if T.constexpr(b is not None): B = T.match_buffer(b, (1,), "int32") out[0] = out[0] + B[0] @@ -340,7 +340,7 @@ def test_multiple_optional_params_preserve_runtime_order(): def test_optional_only_accepts_none_at_specialization_time(): @T.jit(private=True) def kernel(a: T.Optional(T.handle), out_h: T.handle): - if a is not None: + if T.constexpr(a is not None): T.match_buffer(a, (1,), "int32") T.match_buffer(out_h, (1,), "int32") @@ -374,7 +374,7 @@ def test_t_optional_is_restricted_to_jit(): def test_compile_time_if_binding_uses_python_scope(): @T.jit(private=True) def kernel(a: T.Optional(T.handle), out_h: T.handle): - if a is None: + if T.constexpr(a is None): selected = T.match_buffer(out_h, (1,), "int32") else: selected = T.match_buffer(a, (1,), "int32") @@ -395,11 +395,11 @@ def test_compile_time_bool_ops_and_if_expression_short_circuit(): @T.jit(private=True) def kernel(a: T.Optional(T.handle), out_h: T.handle): out = T.match_buffer(out_h, (1,), "int32") - if a is None or fail_if_evaluated(): + if T.constexpr(a is None or fail_if_evaluated()): out[0] = 1 - if a is not None and fail_if_evaluated(): + if T.constexpr(a is not None and fail_if_evaluated()): out[0] = 2 - out[0] = 3 if a is None else fail_if_evaluated() + out[0] = 3 if T.constexpr(a is None) else fail_if_evaluated() absent = kernel.specialize(a=None) assert [param.name for param in absent.params] == ["out"] @@ -426,9 +426,9 @@ def test_runtime_tir_if_cannot_guard_absent_optional_param(): def test_unguarded_absent_optional_param_reports_source(operation, source_text): @T.jit(private=True) def kernel(a: T.Optional(T.handle)): - if operation == "subscript": + if T.constexpr(operation == "subscript"): a[10] - elif operation == "attribute": + elif T.constexpr(operation == "attribute"): a.ptr_to([0]) else: T.match_buffer(a, (1,), "int32") diff --git a/tests/python/tirx/test_op_namespace_cleanup.py b/tests/python/tirx/test_op_namespace_cleanup.py index ba60b7f74c..15fd01448c 100644 --- a/tests/python/tirx/test_op_namespace_cleanup.py +++ b/tests/python/tirx/test_op_namespace_cleanup.py @@ -229,7 +229,8 @@ def test_backend_specific_wrappers_are_not_root_exports(): def test_backend_load_updates_tirx_alias_and_script_facades(monkeypatch): - from tvm.tirx.script import builder, parser + from tvm.script.parser import tirx as parser + from tvm.tirx.script import builder from tvm.tirx.script.builder import ir as builder_ir backend_name = "unit_test_backend" diff --git a/tests/python/tirx/test_parser_printer.py b/tests/python/tirx/test_parser_printer.py index 01d5a1b12d..793e1d40c9 100644 --- a/tests/python/tirx/test_parser_printer.py +++ b/tests/python/tirx/test_parser_printer.py @@ -852,7 +852,7 @@ def test_macro_recursive(): @T.inline def add(x, c): - if c > 0: + if T.constexpr(c > 0): add(x, c - 1) T.evaluate(x) diff --git a/tests/python/tvmscript/test_tvmscript_error_report.py b/tests/python/tvmscript/test_tvmscript_error_report.py index a61350c1af..1d2ef09ba3 100644 --- a/tests/python/tvmscript/test_tvmscript_error_report.py +++ b/tests/python/tvmscript/test_tvmscript_error_report.py @@ -619,7 +619,7 @@ def test_format_source_snippet_multi_line(): """Unit-level check that _format_source_snippet renders every line in a multi-line span, with the underline covering start-col..EOL on the first line, full interior lines, and col-1..end-col on the last line.""" - from tvm.script.parser.core.diagnostics import _format_source_snippet + from tvm.script.parser.diagnostics import _format_source_snippet source_lines = [ "first ignored line\n", @@ -647,7 +647,7 @@ def test_format_source_snippet_multi_line(): def test_format_source_snippet_single_line_unchanged(): """A single-line span (end_lineno == lineno) underlines only the [col_offset, end_col_offset) columns on that one line.""" - from tvm.script.parser.core.diagnostics import _format_source_snippet + from tvm.script.parser.diagnostics import _format_source_snippet source_lines = ["ignored\n", " abc + def\n", "ignored\n"] # Underline just 'abc' (cols 5..8 exclusive) on line 2. diff --git a/tests/python/tvmscript/test_tvmscript_parser_evaluator.py b/tests/python/tvmscript/test_tvmscript_parser_evaluator.py index 463b0f6d29..1c02b1b06c 100644 --- a/tests/python/tvmscript/test_tvmscript_parser_evaluator.py +++ b/tests/python/tvmscript/test_tvmscript_parser_evaluator.py @@ -20,19 +20,17 @@ import pytest import tvm.testing -from tvm.script.parser.core.diagnostics import Source -from tvm.script.parser.core.evaluator import ExprEvaluator +from tvm.script.parser.frontend import Compiler +from tvm.tirx.script import builder as T def _calc(expr, extra_vars=None): if extra_vars is None: extra_vars = {} - source = Source(expr) - mod_ast = source.as_ast() - mod_body_ast = mod_ast.body - expr_stmt_ast = mod_body_ast[0] - expr_ast = expr_stmt_ast.value - return ExprEvaluator.eval(None, extra_vars, expr_ast) + compiler = Compiler("def evaluate():\n return " + expr + "\n", extra_vars) + return compiler.run_statements( + compiler.tree.body[0].body, T, compiler.env, set(), preserve_return=True + ) def test_evaluator_basic(): diff --git a/tests/python/tvmscript/test_tvmscript_parser_source.py b/tests/python/tvmscript/test_tvmscript_parser_source.py index 717e7bc5bf..6ab4dc552b 100644 --- a/tests/python/tvmscript/test_tvmscript_parser_source.py +++ b/tests/python/tvmscript/test_tvmscript_parser_source.py @@ -15,7 +15,7 @@ # specific language governing permissions and limitations # under the License. # ruff: noqa: F401 -"""Unittests for tvm.script.parser.core""" +"""Source and span tests for the canonical parser""" import inspect @@ -27,8 +27,8 @@ import tvm import tvm.testing from tvm.ir import Call, SequentialSpan, TensorLoad, assert_structural_equal from tvm.script import tirx as T -from tvm.script.parser.core import doc_core as doc -from tvm.script.parser.core.diagnostics import Source +import ast as doc +from tvm.script.parser.source import Source from tvm.script.tirx import tile as Tx from tvm.tirx.stmt import TilePrimitiveCall diff --git a/tests/python/tvmscript/test_tvmscript_parser_tir.py b/tests/python/tvmscript/test_tvmscript_parser_tir.py index a4baa9b456..e7ba984464 100644 --- a/tests/python/tvmscript/test_tvmscript_parser_tir.py +++ b/tests/python/tvmscript/test_tvmscript_parser_tir.py @@ -118,16 +118,19 @@ def main(A: T.Buffer(("n",), "float32"), n: T.int64): assert str(n.ty.dtype) == "int64" -def test_tir_string_defined_symbol_does_not_take_dtype_from_body(): - with pytest.raises(tvm.error.DiagnosticError): - tvm.script.from_source( - """ +def test_tir_string_defined_symbol_uses_prescanned_body_dtype(): + func = tvm.script.from_source( + """ @T.prim_func def main(A: T.Buffer(("n",), "float32")): n = T.int64() T.evaluate(n) """ - ) + ) + + n = func.params[0].ty.shape[0] + assert str(n.ty.dtype) == "int64" + assert func.body.value.same_as(n) def test_tir_direct_use_before_string_definition_is_undefined(): @@ -690,7 +693,7 @@ def test_deterministic_branch(): def create_func(predicate: bool): @T.prim_func(private=True, s_tir=True) def func() -> None: - if predicate: + if T.constexpr(predicate): T.evaluate(0) else: T.evaluate(1) diff --git a/tests/python/tvmscript/test_tvmscript_roundtrip.py b/tests/python/tvmscript/test_tvmscript_roundtrip.py index 3466246b69..9c2051b84f 100644 --- a/tests/python/tvmscript/test_tvmscript_roundtrip.py +++ b/tests/python/tvmscript/test_tvmscript_roundtrip.py @@ -3163,7 +3163,7 @@ def relax_extern_func(): def relax_match_cast_ty_proxy(): - """TypeProxy subclasses may be used as expressions + """Default type constructors may be used as expressions This is a regression test. The TVMScript parser allows Type to be specified using a default-constructible class @@ -3188,16 +3188,9 @@ def relax_match_cast_ty_proxy(): inner.__name__ = subclass.__name__ return inner - # Not all subclasses of TypeProxy are default-constructible. - # This list is a subset of `TypeProxy.__subclasses__()`, - # excluding `PrimProxy` and `DTensorProxy`. - subclasses = [ - tvm.script.parser.relax.entry.AnyProxy, - tvm.script.parser.relax.entry.TensorProxy, - tvm.script.parser.relax.entry.CallableProxy, - tvm.script.parser.relax.entry.TupleProxy, - tvm.script.parser.relax.entry.ShapeProxy, - ] + # Prim and DTensor require arguments; the remaining public type + # constructors also work as bare values in match_cast expressions. + subclasses = [R.Any, R.Tensor, R.Callable, R.Tuple, R.Shape] for subclass in subclasses: yield make_ir_generator(subclass)
