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 0f78d0ce501f5ba302db2d7988076e8943c329a3 Author: Tianqi Chen <[email protected]> AuthorDate: Wed Sep 23 08:48:05 2026 +0000 Declare mutable kernel and transform storage explicitly --- python/tvm/backend/cuda/lang/clc.py | 2 +- python/tvm/relax/frontend/tflite/tflite_frontend.py | 4 ++-- python/tvm/s_tir/tensor_intrin/x86.py | 10 +++++----- .../python/tirx-transform/test_tir_transform_bf16_legalize.py | 1 + .../test_tir_transform_force_narrow_index_to_i32.py | 8 ++++---- tests/python/tirx/iket/test_iket_profiler.py | 10 +++++----- tests/python/tirx/transform/test_transform_lower_tirx.py | 2 ++ 7 files changed, 20 insertions(+), 17 deletions(-) diff --git a/python/tvm/backend/cuda/lang/clc.py b/python/tvm/backend/cuda/lang/clc.py index bc79b85ec1..e56d17d4fc 100644 --- a/python/tvm/backend/cuda/lang/clc.py +++ b/python/tvm/backend/cuda/lang/clc.py @@ -41,7 +41,7 @@ def query_cancel_first_ctaid_x(first_ctaid_x, handle, *, use_ld_acquire=True): T.ptx[f"ld{'.acquire.cta' if use_ld_acquire else ''}.shared.b128"](response, handle) T.ptx.clusterlaunchcontrol.query_cancel.is_canceled.pred.b128(canceled, response) - first_ctaid_x = T.uint32(0xFFFFFFFF) + T.buffer_store(first_ctaid_x.source, T.uint32(0xFFFFFFFF), 0) T.ptx.clusterlaunchcontrol.query_cancel.get_first_ctaid__x.b32.b128( first_ctaid_x, response, pred=canceled ) diff --git a/python/tvm/relax/frontend/tflite/tflite_frontend.py b/python/tvm/relax/frontend/tflite/tflite_frontend.py index b6704660de..7129198aa4 100644 --- a/python/tvm/relax/frontend/tflite/tflite_frontend.py +++ b/python/tvm/relax/frontend/tflite/tflite_frontend.py @@ -8343,8 +8343,8 @@ def _build_tflite_rfft2d_primfunc(input_shape, output_pair_shape): for b_idx, out_y, out_x in T.grid(batch, height, out_width): with T.sblock("rfft2d"): v_b, v_oy, v_ox = T.axis.remap("SSS", [b_idx, out_y, out_x]) - real_sum = T.float32(0) - imag_sum = T.float32(0) + real_sum: T.float32 = T.float32(0) + imag_sum: T.float32 = T.float32(0) input_base = v_b * height * width for in_y, in_x in T.grid(height, width): phase_y = T.Cast("float32", v_oy) * T.Cast("float32", in_y) / T.float32(height) diff --git a/python/tvm/s_tir/tensor_intrin/x86.py b/python/tvm/s_tir/tensor_intrin/x86.py index 2fad505104..1db0cc51e3 100644 --- a/python/tvm/s_tir/tensor_intrin/x86.py +++ b/python/tvm/s_tir/tensor_intrin/x86.py @@ -51,12 +51,12 @@ def dot_product_16x4_u8i8i32_vnni( T.reads(C[0:16], A[0:4], B[0:16, 0:4]) T.writes(C[0:16]) - A_u8x4 = A.vload([0], "uint8x4") - A_i32 = T.reinterpret(A_u8x4, dtype="int32") + A_u8x4: T.uint8x4 = A.vload([0], "uint8x4") + A_i32: T.int32 = T.reinterpret(A_u8x4, dtype="int32") - B_i8x64 = B.vload([0, 0], dtype="int8x64") - B_i32x16 = T.reinterpret(B_i8x64, dtype="int32x16") - C_i32x16 = C.vload([0], dtype="int32x16") + B_i8x64: T.int8x64 = B.vload([0, 0], dtype="int8x64") + B_i32x16: T.int32x16 = T.reinterpret(B_i8x64, dtype="int32x16") + C_i32x16: T.int32x16 = C.vload([0], dtype="int32x16") C[T.ramp(T.int32(0), 1, 16)] = T.call_llvm_pure_intrin( T.llvm_lookup_intrinsic_id("llvm.x86.avx512.vpdpbusd.512"), diff --git a/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py b/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py index 8591f70754..2127b2255c 100644 --- a/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py +++ b/tests/python/tirx-transform/test_tir_transform_bf16_legalize.py @@ -118,6 +118,7 @@ def test_bf16_masked_load_store_will_legalize(): A = T.decl_buffer((16,), "bfloat16", data=Aptr) B = T.decl_buffer((16,), "bfloat16") C = T.decl_buffer((16,), "bfloat16", data=Cptr) + mask = T.local_scalar("boolx4") mask = T.Broadcast(T.bool(True), 4) T.evaluate( T.call_intrin( diff --git a/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py b/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py index bb99a48363..2206f69283 100644 --- a/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py +++ b/tests/python/tirx-transform/test_tir_transform_force_narrow_index_to_i32.py @@ -301,7 +301,7 @@ def test_conditional_index_mixed_width_branches(): class Before: @T.prim_func(s_tir=True) def main(A: T.Buffer((T.int64(4),), "float32"), B: T.Buffer((4,), "float32"), n: T.int64): - opaque_index = T.call_extern("opaque_index", n, dtype="int64") + opaque_index: T.int64 = T.call_extern("opaque_index", n, dtype="int64") B[0] = A[T.if_then_else(n < T.int64(0), opaque_index, n)] B[1] = A[T.if_then_else(n < T.int64(0), n, opaque_index)] B[2] = A[T.Select(n < T.int64(0), opaque_index, n)] @@ -311,7 +311,7 @@ def test_conditional_index_mixed_width_branches(): class Expected: @T.prim_func(s_tir=True) def main(A: T.Buffer((4,), "float32"), B: T.Buffer((4,), "float32"), n: T.int32): - opaque_index = T.call_extern("opaque_index", n, dtype="int64") + opaque_index: T.int64 = T.call_extern("opaque_index", n, dtype="int64") B[0] = A[T.if_then_else(n < 0, opaque_index, T.Cast("int64", n))] B[1] = A[T.if_then_else(n < 0, T.Cast("int64", n), opaque_index)] B[2] = A[T.Select(n < 0, opaque_index, T.Cast("int64", n))] @@ -445,7 +445,7 @@ def test_let_binding(): def main(buf: T.handle): n = T.int64() Buf = T.match_buffer(buf, [n], "int32") - ceil_log2 = T.Cast("int64", T.ceil(T.log2(T.Cast("float32", n)))) + ceil_log2: T.int64 = T.Cast("int64", T.ceil(T.log2(T.Cast("float32", n)))) for i in T.serial(ceil_log2): T.evaluate(0) @@ -458,7 +458,7 @@ def test_let_binding(): # The pass narrows indexing variables (n, the For extent) but leaves # an explicitly-typed `T.Cast("int64", ...)` storage alone; a Cast to # int32 is inserted at the use site (the For iter) instead. - ceil_log2 = T.Cast("int64", T.ceil(T.log2(T.Cast("float32", n)))) + ceil_log2: T.int64 = T.Cast("int64", T.ceil(T.log2(T.Cast("float32", n)))) for i in range(T.Cast("int32", ceil_log2)): T.evaluate(0) diff --git a/tests/python/tirx/iket/test_iket_profiler.py b/tests/python/tirx/iket/test_iket_profiler.py index ba6f30f3d7..05eb45ebe4 100644 --- a/tests/python/tirx/iket/test_iket_profiler.py +++ b/tests/python/tirx/iket/test_iket_profiler.py @@ -83,7 +83,7 @@ def token_loop(n: T.int32, out: T.Buffer((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) - token = iket.sentinel_token("sentinel") + token: T.uint32 = iket.sentinel_token("sentinel") for i in T.serial(n, unroll=False): iket.range_end(token) if i % 2 == 0: @@ -119,7 +119,7 @@ def payload_types(n: T.int64, out: T.Buffer((32,), "int32")): iket.mark("u64", T.uint64(64)) iket.mark("f32", T.float32(-3.25)) iket.mark("f64", T.float64(6.5)) - token = iket.range_start("token_payload", T.int32(-7)) + token: T.uint32 = iket.range_start("token_payload", T.int32(-7)) iket.range_end(token, T.int32(9)) iket.range_push("stack_payload", T.float32(1.5)) iket.range_pop() @@ -131,7 +131,7 @@ def payload_presence_mismatch(out: T.Buffer((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) - token = iket.range_start("mismatch", tx) + token: T.uint32 = iket.range_start("mismatch", tx) iket.range_end(token) out[tx] = tx @@ -141,7 +141,7 @@ def payload_type_mismatch(out: T.Buffer((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) - token = iket.range_start("mismatch", tx) + token: T.uint32 = iket.range_start("mismatch", tx) iket.range_end(token, T.uint32(tx)) out[tx] = tx @@ -151,7 +151,7 @@ def sentinel_only_payload(out: T.Buffer((32,), "int32")): T.device_entry() iket = IketProfiler() tx = T.thread_id([32]) - token = iket.sentinel_token("not-a-declaration") + token: T.uint32 = iket.sentinel_token("not-a-declaration") iket.range_end(token, out[tx]) out[tx] = tx diff --git a/tests/python/tirx/transform/test_transform_lower_tirx.py b/tests/python/tirx/transform/test_transform_lower_tirx.py index f65050f587..6233438f35 100644 --- a/tests/python/tirx/transform/test_transform_lower_tirx.py +++ b/tests/python/tirx/transform/test_transform_lower_tirx.py @@ -21,6 +21,7 @@ import tvm_ffi import tvm import tvm.testing from tvm.script import tirx as T +from tvm.script.parser.protocol import register_mutable_var_decl from tvm.script.tirx import tile as Tx from tvm.tirx.function import PrimFunc from tvm.tirx.layout import laneid, warpid, wg_local_layout @@ -1491,6 +1492,7 @@ def test_lower_alloc_decl_buffer_outside_of_parser(): self.B = T.alloc_local([1], "float16") self.C = T.decl_buffer([1], "float16", smem, elem_offset=0, scope="shared.dyn") + @register_mutable_var_decl def int_var1(val): buf = T.local_scalar("int32") if val is not None:
