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")):
