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 44060f8d7c9b03caf1f7909a052657d48d520f0b
Author: Tianqi Chen <[email protected]>
AuthorDate: Wed Sep 23 09:05:16 2026 +0000

    Preserve cluster launch query destination indices in sentinel stores
---
 python/tvm/backend/cuda/lang/clc.py            |  2 +-
 tests/python/tirx/codegen/test_codegen_cuda.py | 34 ++++++++++++++++++++++++++
 2 files changed, 35 insertions(+), 1 deletion(-)

diff --git a/python/tvm/backend/cuda/lang/clc.py 
b/python/tvm/backend/cuda/lang/clc.py
index e56d17d4fc..65577963a7 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)
-    T.buffer_store(first_ctaid_x.source, T.uint32(0xFFFFFFFF), 0)
+    T.buffer_store(first_ctaid_x.source, T.uint32(0xFFFFFFFF), 
first_ctaid_x.indices)
     T.ptx.clusterlaunchcontrol.query_cancel.get_first_ctaid__x.b32.b128(
         first_ctaid_x, response, pred=canceled
     )
diff --git a/tests/python/tirx/codegen/test_codegen_cuda.py 
b/tests/python/tirx/codegen/test_codegen_cuda.py
index 04e872c833..0500bee7c7 100644
--- a/tests/python/tirx/codegen/test_codegen_cuda.py
+++ b/tests/python/tirx/codegen/test_codegen_cuda.py
@@ -20,6 +20,7 @@ import re
 
 import numpy as np
 import pytest
+import tvm_ffi
 
 import tvm
 import tvm.testing
@@ -714,6 +715,39 @@ def test_ptx_cp_async_bulk_non_tma_form_codegen():
     assert 'asm volatile("cp.async.bulk.wait_group 1;" :  :  : "memory");' in 
src
 
 
[email protected]("shape,indices", [((3,), (1,)), ((2, 3), (1, 2))])
+def test_clc_query_preserves_output_indices(shape, indices):
+    @T.prim_func
+    def main(A: T.Buffer(shape, "uint32")):
+        response = T.alloc_buffer((4,), "uint32", scope="shared", align=16)
+        query_cancel_first_ctaid_x(A[indices], response.ptr_to([0]))
+
+    stores = []
+    queries = []
+
+    def collect(node):
+        if isinstance(node, tvm.tirx.BufferStore):
+            stores.append(node)
+        if isinstance(node, tvm.ir.Call) and any(
+            isinstance(arg, tvm.ir.StringImm) and arg.value == 
"get_first_ctaid::x"
+            for arg in node.args
+        ):
+            queries.append(node)
+
+    tvm_ffi.structural_walk(main.body, collect)
+    assert len(stores) == len(queries) == 1
+    sentinel = stores[0]
+    destination = queries[0].args[0]
+    assert sentinel.buffer.same_as(main.params[0])
+    assert destination.source.same_as(sentinel.buffer)
+    assert int(sentinel.value) == 0xFFFFFFFF
+    assert tuple(int(index) for index in sentinel.indices) == indices
+    assert len(sentinel.indices) == len(destination.indices)
+    assert all(
+        stored.same_as(queried) for stored, queried in zip(sentinel.indices, 
destination.indices)
+    )
+
+
 def test_ptx_sync_and_clc_codegen():
     @T.prim_func
     def main(A: T.Buffer((1,), "uint32")):

Reply via email to