From 6fed8670a57238e2336f1139b4395d03de956fc4 Mon Sep 17 00:00:00 2001 From: Hongyi Jin <3231950289@qq.com> Date: Thu, 24 Sep 2026 18:05:20 -0400 Subject: [PATCH] [FIX][TIRx][CUDA] Reevaluate conditional wait predicates after each load Keep statements emitted while printing wait_until predicates inside the repeated macro expression. Preserve lazy conditional evaluation and leave ordinary predicates as direct expressions. Add bounded-process GPU regressions and correct the load-first API documentation. [FIX][TIRx][CUDA] Make lazy macro arguments explicit Record zero-based lazy_args in cuda_func_call attributes and preserve their evaluation sites in CUDA codegen. Let wait_until mark its predicate instead of relying on a wait-specific printer branch. Preserve TVMScript round-trip and cover per-argument, repeated, nested conditional evaluation on GPU. [TEST][TIRx][CUDA] Remove redundant wait predicate dtype variants Keep one int32 regression for each backoff macro form. The previous dtype matrix used identical small positive values and did not exercise distinct signedness or width boundaries. --- python/tvm/backend/cuda/cpp/asm.py | 7 ++- python/tvm/backend/cuda/op.py | 40 ++++++++---- src/backend/cuda/codegen/codegen_cuda.cc | 35 ++++++++++- src/tirx/script/printer/expr.cc | 16 ++++- .../python/tirx/codegen/test_codegen_cuda.py | 63 +++++++++++++++++++ .../tirx/codegen/test_cuda_wait_until.py | 50 +++++++++++++++ tests/python/tirx/test_parser_printer.py | 21 +++++++ 7 files changed, 216 insertions(+), 16 deletions(-) diff --git a/python/tvm/backend/cuda/cpp/asm.py b/python/tvm/backend/cuda/cpp/asm.py index ccbe995be146..8054ef572cf7 100644 --- a/python/tvm/backend/cuda/cpp/asm.py +++ b/python/tvm/backend/cuda/cpp/asm.py @@ -160,7 +160,7 @@ def _wait_until_forward(spelling, *args): @register_codegen("cuda_wait_until") def cuda_wait_until(dst, ptr, condition, scope, space, ptx_type, backoff_ns): - """Lower a declared wait to a pre-tested loop around one scoped load.""" + """Lower a declared wait to a load-first loop around one scoped load.""" scope, space, ptx_type = (parse_str(x) for x in (scope, space, ptx_type)) dtype = _wait_until_thread_local_scalar(dst, "wait_until") suffix = _wait_until_word_suffix(ptr, ptx_type, "wait_until") @@ -260,7 +260,10 @@ def cuda_wait_until(dst, ptr, condition, scope, space, ptx_type, backoff_ns): f"{closing} }} while (0)\n" ) operands = (condition, backoff_ns) - return cuda_func_call(name, *load_call.args[1:-1], *operands, source_code=source), tags + return ( + cuda_func_call(name, *load_call.args[1:-1], *operands, source_code=source, lazy_args=(2,)), + tags, + ) # ============================================================================= diff --git a/python/tvm/backend/cuda/op.py b/python/tvm/backend/cuda/op.py index 534bc625f9a3..56f9f32672e3 100644 --- a/python/tvm/backend/cuda/op.py +++ b/python/tvm/backend/cuda/op.py @@ -88,8 +88,8 @@ def cuda_iket_official_event(event_id, source_code="", payload=None): return call_intrin("uint32", "tirx.cuda.iket_official_event", event_id, source_code) -def cuda_func_call(func_name, *args, source_code, return_type="void"): - """TVM intrinsic to call a CUDA function. Source code is provided as a string. +def cuda_func_call(func_name, *args, source_code, return_type="void", lazy_args=()): + """TVM intrinsic to call a CUDA function or macro supplied as source code. Parameters ---------- @@ -104,8 +104,24 @@ def cuda_func_call(func_name, *args, source_code, return_type="void"): return_type: str The return type of the CUDA function. + + lazy_args: Sequence[int] + Zero-based indices into ``args`` whose generated statements must remain + inside the argument expression. Use this for value arguments evaluated + conditionally or repeatedly by a macro. The macro determines when and + how often they run; an ordinary function still evaluates its arguments + before entering its body. Output arguments that require an lvalue should + not be marked lazy. The default preserves ordinary argument codegen. """ - return call_intrin(return_type, "tirx.cuda.func_call", func_name, *args, source_code) + lazy_args = tuple(lazy_args) + if any(isinstance(i, bool) or not isinstance(i, int) for i in lazy_args): + raise TypeError("cuda_func_call lazy_args must contain integer argument indices") + if any(i < 0 or i >= len(args) for i in lazy_args): + raise ValueError("cuda_func_call lazy_args index is out of range") + attrs = {"lazy_args": tuple(sorted(set(lazy_args)))} if lazy_args else None + return call_intrin( + return_type, "tirx.cuda.func_call", func_name, *args, source_code, attrs=attrs + ) def cuda_warp_reduce(value, op, width=32): @@ -459,12 +475,12 @@ def cuda_wait_until( word: the checker judges every access to that address against the protocol the wait names, rather than as an ordinary pair of memory accesses. - ``dst`` is an initialized thread-local scalar; its current value is tested - first, so an already satisfied predicate performs no load. ``predicate`` is - a trace-time callable taking the current value, or the boolean expression - itself. It is re-evaluated on every iteration, so it may test ``dst`` - against a loop-carried scalar such as a barrier's phase: the loop body only - loads, and nothing it does can move that scalar. + ``dst`` is a writable thread-local scalar. Each poll loads into ``dst`` + before testing the predicate, so its initial value is not used. + ``predicate`` is a trace-time callable taking the current value, or the + boolean expression itself. It is re-evaluated on every iteration, so it may + test ``dst`` against a loop-carried scalar such as a barrier's phase: the + loop body only loads, and nothing it does can move that scalar. The wait always synchronizes with the contributions that made the predicate hold, so data those threads published elsewhere is visible when @@ -492,9 +508,9 @@ def cuda_wait_until( only ``global``. ``backoff_ns`` puts a ``__nanosleep`` before each retry, as a contended - wait is ordinarily written. It goes before the load, so a predicate that - holds on entry still performs no load and no sleep, and a wait whose first - poll succeeds pays nothing. A kernel that spells the backoff itself writes + wait is ordinarily written. It goes before each retry's load, so a wait + whose first poll succeeds performs one load and no sleep before the + closing acquire. A kernel that spells the backoff itself writes ``ld`` once and then waits, which is the same instruction sequence. The backoff is the only thing a wait carries besides its own load, and it diff --git a/src/backend/cuda/codegen/codegen_cuda.cc b/src/backend/cuda/codegen/codegen_cuda.cc index e7a06977c973..ffe022218985 100644 --- a/src/backend/cuda/codegen/codegen_cuda.cc +++ b/src/backend/cuda/codegen/codegen_cuda.cc @@ -1033,9 +1033,42 @@ void CodeGenCUDA::Dispatch_(const CallNode* op, std::ostream& os) { auto print_cuda_func_call = [&](const CallNode* op, std::ostream& os) { TVM_FFI_ICHECK_GE(op->args.size(), 2U); size_t num_args = op->args.size() - 2; + std::vector lazy_args(num_args, false); + if (const auto* attrs = op->attrs.as()) { + if (auto indices = attrs->dict.Get("lazy_args")) { + for (int64_t index : indices.value().cast>()) { + TVM_FFI_ICHECK_GE(index, 0) << "cuda_func_call lazy_args index is out of range"; + TVM_FFI_ICHECK_LT(static_cast(index), num_args) + << "cuda_func_call lazy_args index is out of range"; + lazy_args[index] = true; + } + } + } std::vector args; for (size_t i = 1; i < num_args + 1; i++) { - args.push_back(this->PrintExpr(op->args[i])); + if (lazy_args[i - 1]) { + // PrintExpr can emit statements (e.g. for if_then_else). Keep them + // inside this argument so the macro controls their evaluation. + std::ostringstream outer_stream; + stream.swap(outer_stream); + int scope = BeginScope(); + std::string value = PrintExpr(op->args[i]); + bool has_statements = !stream.str().empty(); + if (has_statements) { + PrintIndent(); + stream << "return " << value << ";\n"; + } + EndScope(scope); + if (has_statements) { + PrintIndent(); + stream << "}())"; + value = "([&]() {\n" + stream.str(); + } + stream.swap(outer_stream); + args.push_back(value); + } else { + args.push_back(this->PrintExpr(op->args[i])); + } } std::string source_code = op->args[num_args + 1].as()->value; std::string func_name = op->args[0].as()->value; diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc index f694b844edce..4f5ce7232314 100644 --- a/src/tirx/script/printer/expr.cc +++ b/src/tirx/script/printer/expr.cc @@ -371,7 +371,11 @@ Doc PrintTIRCall(Call call, AccessPath call_p, IRDocsifier d) { // Annotation spellings such as None for an empty tuple are not type values. return d->AddMetadata(call->ty); }; - if (call->attrs.defined()) { + const auto* call_op = call->op.as(); + const auto* dict_attrs = call->attrs.as(); + bool has_cuda_lazy_args = call_op && call_op->name == "tirx.cuda.func_call" && dict_attrs && + dict_attrs->dict.size() == 1 && dict_attrs->dict.count("lazy_args"); + if (call->attrs.defined() && !has_cuda_lazy_args) { ffi::Array call_args; int n_args = call->args.size(); call_args.reserve(n_args); @@ -472,6 +476,16 @@ Doc PrintTIRCall(Call call, AccessPath call_p, IRDocsifier d) { ExprDoc src = LiteralDoc::Str(src_str->value, call_p->Attr("args")->ArrayItem(n_args - 1)); kw_keys.push_back("source_code"); kw_vals.push_back(src); + if (has_cuda_lazy_args) { + ffi::Array lazy_args; + auto indices_p = call_p->Attr("attrs")->Attr("__dict__")->MapItem("lazy_args"); + int i = 0; + for (int64_t index : dict_attrs->dict.at("lazy_args").cast>()) { + lazy_args.push_back(LiteralDoc::Int(index, indices_p->ArrayItem(i++))); + } + kw_keys.push_back("lazy_args"); + kw_vals.push_back(TupleDoc(lazy_args)); + } // If non-void return type, print return_type keyword. if (!call_prim_type || !call_prim_type.value().IsVoid()) { kw_keys.push_back("return_type"); diff --git a/tests/python/tirx/codegen/test_codegen_cuda.py b/tests/python/tirx/codegen/test_codegen_cuda.py index 04e872c8338b..49af851a8b5f 100644 --- a/tests/python/tirx/codegen/test_codegen_cuda.py +++ b/tests/python/tirx/codegen/test_codegen_cuda.py @@ -1033,6 +1033,69 @@ def run_and_check(): test_print() +@pytest.mark.parametrize("lazy_args", [(), (1,), (2,), (1, 2)]) +@pytest.mark.gpu +@pytest.mark.skipif(not env.has_cuda(), reason="need cuda") +def test_cuda_func_call_lazy_args(lazy_args): + """Each marked argument observes updates at its own macro evaluation site.""" + source = r""" +#define evaluate_values(dst, first, second) do { \ + (dst) = 3; \ + int first_value = (first); \ + (dst) = 7; \ + (dst) = first_value + (first) + (second); \ +} while (0) +""" + + @T.prim_func + def main(out: T.Buffer((32,), "int32")): + T.device_entry() + T.cta_id([1]) + lane = T.thread_id([32]) + value = T.alloc_local((1,), "int32") + value[0] = 1 + T.cuda.func_call( + "evaluate_values", + value[0], + T.if_then_else( + lane % 2 == 0, + T.if_then_else(lane % 4 == 0, value[0] * 2, value[0] * 4), + value[0] * 3, + ), + T.if_then_else( + lane % 2 == 0, + T.if_then_else(lane % 4 == 0, value[0] * 2, value[0] * 4), + value[0] * 3, + ), + source_code=source, + lazy_args=lazy_args, + ) + out[lane] = value[0] + + # Attributes must survive serialization and the normal lowering pipeline. + _, mod = _get_source(tvm.ir.load_json(tvm.ir.save_json(main))) + lanes = np.arange(32) + scale = np.where(lanes % 4 == 0, 2, np.where(lanes % 2 == 0, 4, 3)) + expected = scale * ((3 + 7 if 1 in lazy_args else 1 + 1) + (7 if 2 in lazy_args else 1)) + + def run_and_check(): + output = tvm.runtime.tensor(np.zeros(32, dtype="int32"), device=tvm.cuda()) + mod(output) + np.testing.assert_array_equal(output.numpy(), expected) + + tvm.testing.run_with_gpu_lock(run_and_check) + + +@pytest.mark.parametrize( + "lazy_args, error", [((-1,), ValueError), ((1,), ValueError), ((0.5,), TypeError)] +) +def test_cuda_func_call_invalid_lazy_args(lazy_args, error): + from tvm.backend.cuda.op import cuda_func_call + + with pytest.raises(error, match="lazy_args"): + cuda_func_call("macro", 1, source_code="", lazy_args=lazy_args) + + @pytest.mark.gpu @pytest.mark.skipif(not env.has_cuda(), reason="need cuda") def test_warp_shuffle_xor_sync(): diff --git a/tests/python/tirx/codegen/test_cuda_wait_until.py b/tests/python/tirx/codegen/test_cuda_wait_until.py index 20be1e8da927..c53d1d76350b 100644 --- a/tests/python/tirx/codegen/test_cuda_wait_until.py +++ b/tests/python/tirx/codegen/test_cuda_wait_until.py @@ -28,10 +28,12 @@ caller's destination. The protocols exercised are the shapes real kernels use. """ +import numpy as np import pytest import tvm from tvm.script import tirx as T +from tvm.support.popen_pool import PopenWorker def build(func): @@ -99,6 +101,54 @@ def _wait_macro(source): ) +@pytest.mark.gpu +@pytest.mark.parametrize("backoff_ns", [None, 40]) +def test_conditional_predicate_observes_loaded_value(backoff_ns): + """Statements needed by a lazy predicate must execute after the polling load.""" + + def run(): + @T.prim_func + def kernel(state: T.Buffer((32,), "int32"), out: T.Buffer((32,), "int32")): + T.device_entry() + T.cta_id([1]) + lane = T.thread_id([32]) + observed = T.alloc_local((1,), "int32") + observed[0] = 999 + if lane < 17: + T.cuda.wait_until( + observed[0], + state.ptr_to([lane]), + lambda current: T.if_then_else( + lane % 2 == 0, + current // 2 == (lane + 17) // 2, + T.bitwise_and(current, 255) == lane + 17, + ), + backoff_ns=backoff_ns, + ) + out[lane] = observed[0] + + target = tvm.target.Target({"kind": "cuda", "arch": "sm_100a"}) + with target: + executable = tvm.compile( + tvm.IRModule({"main": kernel}), target=target, tir_pipeline="tirx" + ) + state = tvm.runtime.tensor(np.arange(17, 49, dtype="int32"), device=tvm.cuda(0)) + output = tvm.runtime.tensor(np.zeros(32, dtype="int32"), device=tvm.cuda(0)) + executable(state, output) + expected = np.zeros(32, dtype="int32") + expected[:17] = np.arange(17, 34, dtype="int32") + np.testing.assert_array_equal(output.numpy(), expected) + + # A stale false predicate spins forever. Isolate the CUDA context so a + # regression fails with a timeout and cannot leave a kernel running. + worker = PopenWorker() + try: + worker.send(run, timeout=60) + worker.recv() + finally: + worker.kill() + + # ============================================================================= # The emitted loop # ============================================================================= diff --git a/tests/python/tirx/test_parser_printer.py b/tests/python/tirx/test_parser_printer.py index 01d5a1b12d7a..fdd81c9ba0c3 100644 --- a/tests/python/tirx/test_parser_printer.py +++ b/tests/python/tirx/test_parser_printer.py @@ -2712,6 +2712,27 @@ def func(): assert_structural_equal(func, from_source(code)) +def test_roundtrip_cuda_func_call_lazy_args(): + """Preserve lazy argument positions, embedded source, and the return type.""" + source = "\n#define choose_values(first, second) ((first) + (second))\n" + + @T.prim_func + def func(A: T.Buffer((2,), "int32")): + T.device_entry() + A[0] = T.cuda.func_call( + "choose_values", + T.if_then_else(A[0] > 0, A[0], A[1]), + T.if_then_else(A[1] > 0, A[1], A[0]), + source_code=source, + return_type="int32", + lazy_args=(0, 1), + ) + + code = func.script() + assert from_source(code).script() == code + assert_structural_equal(func, from_source(code)) + + def test_roundtrip_cp_async_bulk_tensor_g2s_cluster(): """The TMA load composite [tensorMap, coords] operand must round-trip."""