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)

Reply via email to