Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 5 additions & 2 deletions python/tvm/backend/cuda/cpp/asm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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,
)


# =============================================================================
Expand Down
40 changes: 28 additions & 12 deletions python/tvm/backend/cuda/op.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
----------
Expand All @@ -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):
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
35 changes: 34 additions & 1 deletion src/backend/cuda/codegen/codegen_cuda.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<bool> lazy_args(num_args, false);
if (const auto* attrs = op->attrs.as<DictAttrsNode>()) {
if (auto indices = attrs->dict.Get("lazy_args")) {
for (int64_t index : indices.value().cast<ffi::Array<int64_t>>()) {
TVM_FFI_ICHECK_GE(index, 0) << "cuda_func_call lazy_args index is out of range";
TVM_FFI_ICHECK_LT(static_cast<size_t>(index), num_args)
<< "cuda_func_call lazy_args index is out of range";
lazy_args[index] = true;
}
}
}
std::vector<std::string> 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<StringImmNode>()->value;
std::string func_name = op->args[0].as<StringImmNode>()->value;
Expand Down
16 changes: 15 additions & 1 deletion src/tirx/script/printer/expr.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<OpNode>();
const auto* dict_attrs = call->attrs.as<DictAttrsNode>();
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<ExprDoc> call_args;
int n_args = call->args.size();
call_args.reserve(n_args);
Expand Down Expand Up @@ -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<ExprDoc> 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<ffi::Array<int64_t>>()) {
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");
Expand Down
63 changes: 63 additions & 0 deletions tests/python/tirx/codegen/test_codegen_cuda.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down
50 changes: 50 additions & 0 deletions tests/python/tirx/codegen/test_cuda_wait_until.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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
# =============================================================================
Expand Down
21 changes: 21 additions & 0 deletions tests/python/tirx/test_parser_printer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down