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:

Reply via email to