From b4d2a3eb22f673ea432790eec7084cb1fc84b5a3 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Sat, 3 Oct 2026 01:13:43 +0000 Subject: [PATCH 1/4] [REFACTOR][IR] Initialize Python Op APIs from registered namespaces --- python/tvm/backend/cuda/__init__.py | 6 +- python/tvm/backend/cuda/op.py | 551 +----------------- python/tvm/backend/cuda/script.py | 200 +++---- python/tvm/ir/op.py | 88 +++ python/tvm/tirx/script/ir_builder/op.py | 10 +- src/backend/cuda/op/iket.cc | 13 +- src/backend/cuda/op/target_builtin.cc | 43 +- src/tirx/script/printer/expr.cc | 68 ++- tests/python/ir/test_op_api.py | 345 +++++++++++ .../tirx/script/test_tirx_script_printer.py | 6 +- .../python/tirx/test_op_namespace_cleanup.py | 8 +- 11 files changed, 631 insertions(+), 707 deletions(-) create mode 100644 tests/python/ir/test_op_api.py diff --git a/python/tvm/backend/cuda/__init__.py b/python/tvm/backend/cuda/__init__.py index 627ff1a0ffbd..31621f2f2def 100644 --- a/python/tvm/backend/cuda/__init__.py +++ b/python/tvm/backend/cuda/__init__.py @@ -73,7 +73,7 @@ def register_backend(): pass register_device_target_detector("cuda", _detect_target_from_device) for name, namespace in script_namespaces().items(): - builder_op.register_script_namespace(name, namespace) + builder_op.register_script_namespace(name, namespace, register_printer_names=name != "cuda") # script_namespaces() above pulls in ptx, which only imports the shared # codegen layer -- not the device-helper modules. This import is the sole @@ -85,16 +85,16 @@ def register_backend(): def script_namespaces(**_): """Return CUDA-owned TVMScript namespaces.""" + from . import script # pylint: disable=import-outside-toplevel from .ptx import PTXNamespace # pylint: disable=import-outside-toplevel from .script import ( # pylint: disable=import-outside-toplevel - CUDANamespace, NVSHMEMNamespace, PTXLegacyNamespace, STIRNamespace, ) return { - "cuda": CUDANamespace(), + "cuda": script, "nvshmem": NVSHMEMNamespace(), "ptx_legacy": PTXLegacyNamespace(), "ptx": PTXNamespace(), diff --git a/python/tvm/backend/cuda/op.py b/python/tvm/backend/cuda/op.py index 534bc625f9a3..65d8df68c784 100644 --- a/python/tvm/backend/cuda/op.py +++ b/python/tvm/backend/cuda/op.py @@ -21,6 +21,7 @@ from tvm import tirx from tvm.ir import Call, Op, StringImm +from tvm.ir.op import _init_op_api from tvm.ir.type import PointerType, PrimType from tvm.runtime import const from tvm.tirx.op import bitwise_and, call_intrin, tvm_access_ptr @@ -69,16 +70,6 @@ def cuda_iket_range_push(name, payload=None): return call_intrin("", "tirx.cuda.iket_range_push", name) -def cuda_iket_range_pop(): - """Create an NVIDIA IKET stack-range pop annotation.""" - return call_intrin("", "tirx.cuda.iket_range_pop") - - -def cuda_iket_sentinel_token(name): - """Create a no-op NVIDIA IKET range token for warp-uniform control flow.""" - return call_intrin("uint32", "tirx.cuda.iket_sentinel_token", name) - - def cuda_iket_official_event(event_id, source_code="", payload=None): """Create an NVIDIA IKET official range-end event.""" if payload is not None: @@ -194,194 +185,6 @@ def cuda_cta_min(value, num_warps, scratch): return cuda_cta_reduce(value, "min", num_warps, scratch) -def cuda_warp_sync(): - """TVM intrinsic to synchronize threads within the current warp. - - This lowers to a CUDA `__syncwarp()` call. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.warp_sync") - - -def cuda_cta_sync(): - """TVM intrinsic to call CUDA syncthreads (block-wide barrier) - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.cta_sync") - - -def cuda_grid_sync(): - """TVM intrinsic to call CUDA grid-wide sync (cooperative groups) - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.grid_sync") - - -def cuda_cluster_sync(): - """TVM intrinsic to call CUDA cluster-wide barrier sync - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.cluster_sync") - - -def cuda_thread_rank(): - """TVM intrinsic that returns ``cooperative_groups::thread_rank()`` - for the enclosing CTA -- the linear thread index within the block. - - Useful for building "single thread of CTA" predicates without - referencing user-declared scope_id vars. For example, the idiomatic - mbarrier.init leader predicate is:: - - T.cuda.thread_rank() == 0 - - Returns - ------- - call : Expr - The call expression (``int32``). - """ - return call_intrin("int32", "tirx.cuda.thread_rank") - - -def cuda_half2float(src): - """TVM intrinsic to convert half to float - - Parameters - ---------- - src : Expr - Source pointer. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("float32", "tirx.cuda.half2float", src) - - -def cuda_bfloat162float(src): - """TVM intrinsic to convert bfloat16 to float - - Parameters - ---------- - src : Expr - Source pointer. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("float32", "tirx.cuda.bfloat162float", src) - - -def cuda_float22half2(dst, src): - """TVM intrinsic to convert float2 to half2 with rounding - - Parameters - ---------- - dst : Expr - Destination pointer. - - src : Expr - Source pointer. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.float22half2", dst, src) - - -def cuda_trap_when_assert_failed(cond): - """TVM intrinsic to trap when assertion failed (cond == false) - - Parameters - ---------- - cond : Expr - Condition to check. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.trap_when_assert_failed", cond) - - -def cuda_runtime_instr_desc(desc, sf_id): - """TVM intrinsic to update runtime instruction descriptor - - Parameters - ---------- - desc : Expr - Pointer to the descriptor (uint32*). - - sf_id : Expr - The subfragment id. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.runtime_instr_desc", desc, sf_id) - - -def cuda_half8tofloat8(src_addr, dst_addr): - """TVM intrinsic to convert 8 half2s to 8 float2s - - Parameters - ---------- - src_addr : Expr - Source pointer. - - dst_addr : Expr - Destination pointer. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.half8tofloat8", src_addr, dst_addr) - - -def cuda_float8tohalf8(src_addr, dst_addr): - """TVM intrinsic to convert 8 float2s to 8 half2s - - Parameters - ---------- - src_addr : Expr - Source pointer. - - dst_addr : Expr - Destination pointer. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.float8tohalf8", src_addr, dst_addr) - - _WAIT_UNTIL_SCOPE = ("cta", "cluster", "gpu", "sys") # Global only, and that is the whole surface a declared word needs. A protocol # that waits inside a CTA or a cluster has `mbarrier`, which is the hardware's @@ -534,42 +337,6 @@ def _validate_mbarrier_arrive_attrs(sem, scope, space, remote): raise ValueError("remote mbarrier.arrive requires space='shared::cluster'") -def cuda_mbarrier_wait(bar, phase): - """Retry ``mbarrier.try_wait.parity.acquire.cta`` until it returns true. - - Parameters - ---------- - bar : Var - The pointer to barrier variable. - - phase : int - The phase of the barrier. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.mbarrier_wait", bar, phase) - - -def cuda_mbarrier_wait_acquire_cluster(bar, phase): - """``mbarrier.try_wait.parity.acquire.cluster`` retry loop. - - Cluster-scope acquire wait — used to wait on a barrier that a remote CTA in - the cluster arrives on (a group cluster wait). - - Parameters - ---------- - bar : Var - The pointer to barrier variable. - - phase : int - The phase of the barrier. - """ - return call_intrin("", "tirx.cuda.mbarrier_wait_acquire_cluster", bar, phase) - - def ptx_cp_async_legacy(*all_args): """Legacy ``ptx_cp_async`` API taking explicit src/dst offsets. @@ -610,11 +377,6 @@ def _is_static_unicast_cta_mask(cta_mask): return False -def cuda_elect_sync(): - """TVM intrinsic to call elect.sync""" - return call_intrin("uint32", "tirx.cuda.elect_sync") - - def cuda_mov_sreg(bits, reg_name): """TVM intrinsic to tvm instrinsics to fetch PTX pre-defined registers @@ -841,72 +603,6 @@ def ptx_legacy_ldmatrix(*all_args): ) -def cuda_wgmma_encode_matrix_descriptor(desc, addr, ldo, sdo, swizzle): - """TVM intrinsic to create memory descriptor for wgmma instructions - - Parameters - ---------- - desc : Expr - The pointer to the shared memory descriptor. - - addr : Expr - The address of the matrix. - - ldo : Expr - The leading dimension offset. - - sdo : Expr - The stride dimension offset. - - swizzle : int - The swizzle value (CUtensorMapSwizzle_enum). - """ - return call_intrin( - "", "tirx.cuda.wgmma_encode_matrix_descriptor", desc, addr, ldo, sdo, swizzle - ) - - -def cuda_wgmma_noop_barrier(reg): - """TVM intrinsic to call "" : "+{format}"(reg)::"memory" - - Parameters - ---------- - reg : Expr - The register to fence. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.wgmma_noop_barrier", reg) - - -def cuda_tcgen05_encode_matrix_descriptor(desc, addr, ldo, sdo, swizzle): - """TVM intrinsic to create memory descriptor for tcgen05 instructions - - Parameters - ---------- - desc : Expr - The pointer to the shared memory descriptor. - - addr : Expr - The address of the matrix. - - ldo : Expr - The leading dimension offset. - - sdo : Expr - The stride dimension offset. - - swizzle : int - The swizzle value (CUtensorMapSwizzle_enum). - """ - return call_intrin( - "", "tirx.cuda.tcgen05_encode_matrix_descriptor", desc, addr, ldo, sdo, swizzle - ) - - def cuda_tcgen05_encode_instr_descriptor( desc, *, @@ -1324,104 +1020,6 @@ def cuda_atomic_add(res_addr, value): return call_intrin(value.ty, "tirx.cuda.atomic_add", res_addr, value) -def cuda_thread_fence(): - """TVM intrinsic to call cuda thread fence instruction - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.thread_fence") - - -def cuda_warpgroup_sync(bar_no): - """TVM intrinsic to synchronize a CUDA warpgroup via a named barrier. - - Parameters - ---------- - bar_no : Expr - The named barrier id to use for the warpgroup. - - Notes - ----- - Synchronizes 128 threads in a warpgroup using `bar.sync bar_no, 128`. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.warpgroup_sync", bar_no) - - -def cuda_syncthreads_and(cond): - """TVM intrinsic to call cuda syncthreads_and instruction - - Parameters - ---------- - cond: Expr - The condition. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("int64", "tirx.cuda.syncthreads_and", cond) - - -def cuda_syncthreads_or(cond): - """TVM intrinsic to call cuda syncthreads_or instruction - - Parameters - ---------- - cond: Expr - The condition. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("int64", "tirx.cuda.syncthreads_or", cond) - - -def cuda_nano_sleep(time): - """TVM intrinsic to call cuda nano sleep instruction - - Parameters - ---------- - time: Expr - The time to sleep. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.nano_sleep", time) - - -def cuda_printf(fmt, *args): - """TVM intrinsic to call cuda printf instruction - - Parameters - ---------- - fmt: str - The format string. - - *args: list - The arguments to the format string. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("", "tirx.cuda.printf", fmt, *args) - - def cuda_ldg(addr, dtype, *, dst=None, vec=""): """TVM intrinsic to call CUDA C++ ``__ldg()``. @@ -1456,76 +1054,12 @@ def cuda_ldg(addr, dtype, *, dst=None, vec=""): return call_intrin("", "tirx.cuda.ldg", *dst, addr, dtype, vec, vec_len) -def cuda_fdividef(x, y): - """TVM intrinsic to call CUDA C++ ``__fdividef`` fast float division.""" - return call_intrin("float32", "tirx.cuda.fdividef", x, y) - - -def cuda_get_tmem_addr(addr, row_offset, col_offset): - """TVM intrinsic to call cuda tmem address calculation - - Parameters - ---------- - addr: Expr - The memory address to calculate. - - row_offset: Expr - The row offset to calculate. - - col_offset: Expr - The column offset to calculate. - - Returns - ------- - call : Expr - The call expression. - """ - return call_intrin("uint32", "tirx.cuda.get_tmem_addr", addr, row_offset, col_offset) - - -def cuda_cvta_generic_to_shared(ptr): - """Convert a generic pointer to a shared-memory address (uint32). - - Wraps ``__cvta_generic_to_shared(ptr)``. Used by op-wrappers that - precompute the shared-memory address at the wrapper layer instead of - inside the asm helper body. - """ - return call_intrin("uint32", "tirx.cuda.cvta_generic_to_shared", ptr) - - -def cuda_smem_addr_from_uint64(cluster_addr): - """Narrow a 64-bit cluster-mapped SMEM address to a 32-bit SMEM address. - - Wraps ``static_cast(cluster_addr)``. Used by - cp.async.bulk.shared::cluster.* op-wrappers. - """ - return call_intrin("uint32", "tirx.cuda.smem_addr_from_uint64", cluster_addr) - - def cuda_sm100_2sm_leader_smem_addr(ptr): """Return the SM100 2SM leader CTA shared-address operand. The input is a generic pointer to shared memory. """ - return bitwise_and(cuda_cvta_generic_to_shared(ptr), const(0xFEFFFFFF, dtype="uint32")) - - -def cuda_any_sync(mask, pred): - """TVM intrinsic for PTX warp-wide any predicate (__any_sync) - - Parameters - ---------- - mask : Expr - The thread mask (uint32). - pred : Expr - The predicate value (int32). - - Returns - ------- - call : Expr - The call expression returning 1 if any thread in mask has pred != 0. - """ - return call_intrin("int32", "tirx.cuda.any_sync", mask, pred) + return bitwise_and(globals()["cvta_generic_to_shared"](ptr), const(0xFEFFFFFF, dtype="uint32")) _PTX_CVT_TYPES = { @@ -1646,78 +1180,6 @@ def _validate_ptx_address(addr, space, op_name): ) -def cuda_uint_as_float(bits): - return call_intrin("float32", "tirx.cuda.uint_as_float", bits) - - -def cuda_float_as_uint(x): - return call_intrin("uint32", "tirx.cuda.float_as_uint", x) - - -def cuda_ballot_sync(mask, pred): - return call_intrin("uint32", "tirx.cuda.ballot_sync", mask, pred) - - -def cuda_ffs_u32(value): - return call_intrin("int32", "tirx.cuda.ffs_u32", value) - - -def cuda_reduce_add_sync_u32(mask, value): - return call_intrin("uint32", "tirx.cuda.reduce_add_sync_u32", mask, value) - - -def cuda_reduce_min_sync_u32(mask, value): - return call_intrin("uint32", "tirx.cuda.reduce_min_sync_u32", mask, value) - - -def cuda_clock64(): - return call_intrin("uint64", "tirx.cuda.clock64") - - -def cuda_make_float2(x, y): - return call_intrin("uint64", "tirx.cuda.make_float2", x, y) - - -def cuda_float2_x(packed): - return call_intrin("float32", "tirx.cuda.float2_x", packed) - - -def cuda_float2_y(packed): - return call_intrin("float32", "tirx.cuda.float2_y", packed) - - -def cuda_fmul2_rn(a, b): - return call_intrin("uint64", "tirx.cuda.fmul2_rn", a, b) - - -def cuda_fadd2_rn(a, b): - return call_intrin("uint64", "tirx.cuda.fadd2_rn", a, b) - - -def cuda_float22bfloat162_rn(v0, v1): - return call_intrin("uint32", "tirx.cuda.float22bfloat162_rn", v0, v1) - - -def cuda_float22bfloat162_rn_from_float2(packed): - return call_intrin("uint32", "tirx.cuda.float22bfloat162_rn_from_float2", packed) - - -def cuda_bfloat1622float2(packed): - return call_intrin("uint64", "tirx.cuda.bfloat1622float2", packed) - - -def cuda_hmin2(a, b): - return call_intrin("uint32", "tirx.cuda.hmin2", a, b) - - -def cuda_hmax2(a, b): - return call_intrin("uint32", "tirx.cuda.hmax2", a, b) - - -def cuda_fp8x4_e4m3_from_float4(x, y, z, w): - return call_intrin("uint32", "tirx.cuda.fp8x4_e4m3_from_float4", x, y, z, w) - - def cuda_atomic_cas(ptr, old_val, new_val): """TVM intrinsic to call cuda atomic cas instruction @@ -2125,3 +1587,12 @@ def nvshmem_barrier_all(): """ return call_intrin("", "tirx.nvshmem.barrier_all") + + +# Canonical Op builders also supply the historical direct-import aliases. +_init_op_api("tirx.cuda", __name__) +for _name in Op.list_op_names(): + if _name.startswith("tirx.cuda."): + _suffix = _name.removeprefix("tirx.cuda.") + if "." not in _suffix: + globals().setdefault("cuda_" + _suffix, globals()[_suffix]) diff --git a/python/tvm/backend/cuda/script.py b/python/tvm/backend/cuda/script.py index 28f6109b9929..34952123b3f6 100644 --- a/python/tvm/backend/cuda/script.py +++ b/python/tvm/backend/cuda/script.py @@ -21,6 +21,7 @@ from collections.abc import Callable from typing import Any +from tvm import ir as _ir from tvm.backend.cuda import op as _cuda_op from tvm.tirx import is_buffer_var from tvm.tirx import op as _tir_op @@ -120,128 +121,80 @@ def __init__(self): self.official_event = _op_wrapper(_cuda_op.cuda_iket_official_event) -class CUDANamespace: - """The CUDA intrinsics submodule.""" - - def __init__(self): - self.iket = IketNamespace() - self.wgmma = CudaWgmmaNamespace() - self.tcgen05 = CudaTcgen05Namespace() - self.any_sync = _op_wrapper(_cuda_op.cuda_any_sync) - # elect.sync plus the predicated mov that materializes its predicate: - # a multi-statement asm block, so it belongs here rather than T.ptx. - # The warp-specialization passes match this op to build predicates. - self.elect_sync: Callable[..., Any] = _op_wrapper(_cuda_op.cuda_elect_sync) - # `mov.u32 d, %sreg` -- one PTX instruction, but the special-register - # name is baked into the asm text, so it is a helper per register - # rather than a ptx entry with a register operand. - self.mov_sreg: Callable[..., Any] = _op_wrapper(_cuda_op.cuda_mov_sreg) - # Spin-until-ready mbarrier waits: label-loop asm blocks, not single - # PTX instructions -- which is why they live here and not in T.ptx. - # One declared synchronization word: every access a protocol makes to - # it goes through these, so a checker can separate them from a stray - # access and read the word's write history off the declaration. - self.wait_until = _op_wrapper(_cuda_op.cuda_wait_until) - self.mbarrier_wait = _op_wrapper(_cuda_op.cuda_mbarrier_wait) - self.mbarrier_wait_acquire_cluster = _op_wrapper( - _cuda_op.cuda_mbarrier_wait_acquire_cluster - ) - self.atomic_add = _op_wrapper(_cuda_op.cuda_atomic_add) - self.thread_fence = _op_wrapper(_cuda_op.cuda_thread_fence) - self.warpgroup_sync = _op_wrapper(_cuda_op.cuda_warpgroup_sync) - self.warp_sync = _op_wrapper(_cuda_op.cuda_warp_sync) - self.warp_reduce = _op_wrapper(_cuda_op.cuda_warp_reduce) - self.warp_sum = _op_wrapper(_cuda_op.cuda_warp_sum) - self.warp_max = _op_wrapper(_cuda_op.cuda_warp_max) - self.warp_min = _op_wrapper(_cuda_op.cuda_warp_min) - self.cta_reduce = _op_wrapper(_cuda_op.cuda_cta_reduce) - self.cta_sum = _op_wrapper(_cuda_op.cuda_cta_sum) - self.cta_max = _op_wrapper(_cuda_op.cuda_cta_max) - self.cta_min = _op_wrapper(_cuda_op.cuda_cta_min) - self.cta_sync = _op_wrapper(_cuda_op.cuda_cta_sync) - self.grid_sync = _op_wrapper(_cuda_op.cuda_grid_sync) - self.cluster_sync = _op_wrapper(_cuda_op.cuda_cluster_sync) - self.thread_rank = _op_wrapper(_cuda_op.cuda_thread_rank) - self.trap_when_assert_failed = _op_wrapper(_cuda_op.cuda_trap_when_assert_failed) - self.runtime_instr_desc = _op_wrapper(_cuda_op.cuda_runtime_instr_desc) - self.half2float = _op_wrapper(_cuda_op.cuda_half2float) - self.bfloat162float = _op_wrapper(_cuda_op.cuda_bfloat162float) - self.float22half2 = _op_wrapper(_cuda_op.cuda_float22half2) - self.half8tofloat8 = _op_wrapper(_cuda_op.cuda_half8tofloat8) - self.float8tohalf8 = _op_wrapper(_cuda_op.cuda_float8tohalf8) - self.syncthreads_and = _op_wrapper(_cuda_op.cuda_syncthreads_and) - self.syncthreads_or = _op_wrapper(_cuda_op.cuda_syncthreads_or) - self.nano_sleep = _op_wrapper(_cuda_op.cuda_nano_sleep) - self.atomic_cas = _op_wrapper(_cuda_op.cuda_atomic_cas) - self.func_call = _op_wrapper(_cuda_op.cuda_func_call) - self.printf = _op_wrapper(_cuda_op.cuda_printf) - self.ldg = _op_wrapper(_cuda_op.cuda_ldg) - self.fdividef = _op_wrapper(_cuda_op.cuda_fdividef) - self.get_tmem_addr = _op_wrapper(_cuda_op.cuda_get_tmem_addr) - self.cvta_generic_to_shared = _op_wrapper(_cuda_op.cuda_cvta_generic_to_shared) - self.smem_addr_from_uint64 = _op_wrapper(_cuda_op.cuda_smem_addr_from_uint64) - self.sm100_2sm_leader_smem_addr = _op_wrapper(_cuda_op.cuda_sm100_2sm_leader_smem_addr) - self.uint_as_float = _op_wrapper(_cuda_op.cuda_uint_as_float) - self.float_as_uint = _op_wrapper(_cuda_op.cuda_float_as_uint) - self.ballot_sync = _op_wrapper(_cuda_op.cuda_ballot_sync) - self.ffs_u32 = _op_wrapper(_cuda_op.cuda_ffs_u32) - self.reduce_add_sync_u32 = _op_wrapper(_cuda_op.cuda_reduce_add_sync_u32) - self.reduce_min_sync_u32 = _op_wrapper(_cuda_op.cuda_reduce_min_sync_u32) - self.clock64 = _op_wrapper(_cuda_op.cuda_clock64) - self.make_float2 = _op_wrapper(_cuda_op.cuda_make_float2) - self.float2_x = _op_wrapper(_cuda_op.cuda_float2_x) - self.float2_y = _op_wrapper(_cuda_op.cuda_float2_y) - self.fmul2_rn = _op_wrapper(_cuda_op.cuda_fmul2_rn) - self.fadd2_rn = _op_wrapper(_cuda_op.cuda_fadd2_rn) - self.float22bfloat162_rn = _op_wrapper(_cuda_op.cuda_float22bfloat162_rn) - self.float22bfloat162_rn_from_float2 = _op_wrapper( - _cuda_op.cuda_float22bfloat162_rn_from_float2 - ) - self.bfloat1622float2 = _op_wrapper(_cuda_op.cuda_bfloat1622float2) - self.hmin2 = _op_wrapper(_cuda_op.cuda_hmin2) - self.hmax2 = _op_wrapper(_cuda_op.cuda_hmax2) - self.fp8x4_e4m3_from_float4 = _op_wrapper(_cuda_op.cuda_fp8x4_e4m3_from_float4) - self.timer_init = _op_wrapper(_cuda_op.timer_init_cuda) - self.timer_start = _op_wrapper(_cuda_op.timer_start_cuda) - self.timer_end = _op_wrapper(_cuda_op.timer_end_cuda) - self.timer_finalize = _op_wrapper(_cuda_op.timer_finalize_cuda) - self.mma_store = _dtype_forward(_cuda_op.mma_store) - self.mma_fill = _dtype_forward(_cuda_op.mma_fill) - self.mma_store_legacy = _dtype_forward(_cuda_op.mma_store_legacy) - self.mma_fill_legacy = _dtype_forward(_cuda_op.mma_fill_legacy) - setattr(self, "__shfl_sync", self._shfl_sync) - setattr(self, "__shfl_up_sync", self._shfl_up_sync) - setattr(self, "__shfl_down_sync", self._shfl_down_sync) - setattr(self, "__shfl_xor_sync", self._shfl_xor_sync) - setattr(self, "__activemask", self._activemask) - - @staticmethod - def _shfl_sync(mask, var, lane, width): - if is_buffer_var(var): - var = var[0] - return _tir_op.call_intrin(var.ty, "tirx.cuda.__shfl_sync", mask, var, lane, width) - - @staticmethod - def _shfl_up_sync(mask, var, delta, width): - if is_buffer_var(var): - var = var[0] - return _tir_op.call_intrin(var.ty, "tirx.cuda.__shfl_up_sync", mask, var, delta, width) - - @staticmethod - def _shfl_down_sync(mask, var, delta, width): - if is_buffer_var(var): - var = var[0] - return _tir_op.call_intrin(var.ty, "tirx.cuda.__shfl_down_sync", mask, var, delta, width) - - @staticmethod - def _shfl_xor_sync(mask, var, lane_mask, width): - if is_buffer_var(var): - var = var[0] - return _tir_op.call_intrin(var.ty, "tirx.cuda.__shfl_xor_sync", mask, var, lane_mask, width) - - @staticmethod - def _activemask(): - return _tir_op.call_intrin("uint32", "tirx.cuda.__activemask") +def _shfl_sync(mask, var, lane, width): + if is_buffer_var(var): + var = var[0] + return _tir_op.call_intrin(var.ty, "tirx.cuda.__shfl_sync", mask, var, lane, width) + + +def _shfl_up_sync(mask, var, delta, width): + if is_buffer_var(var): + var = var[0] + return _tir_op.call_intrin(var.ty, "tirx.cuda.__shfl_up_sync", mask, var, delta, width) + + +def _shfl_down_sync(mask, var, delta, width): + if is_buffer_var(var): + var = var[0] + return _tir_op.call_intrin(var.ty, "tirx.cuda.__shfl_down_sync", mask, var, delta, width) + + +def _shfl_xor_sync(mask, var, lane_mask, width): + if is_buffer_var(var): + var = var[0] + return _tir_op.call_intrin(var.ty, "tirx.cuda.__shfl_xor_sync", mask, var, lane_mask, width) + + +def _activemask(): + return _tir_op.call_intrin("uint32", "tirx.cuda.__activemask") + + +iket = IketNamespace() +wgmma = CudaWgmmaNamespace() +tcgen05 = CudaTcgen05Namespace() +mov_sreg: Callable[..., Any] = _cuda_op.cuda_mov_sreg +wait_until = _cuda_op.cuda_wait_until +atomic_add = _cuda_op.cuda_atomic_add +warp_reduce = _cuda_op.cuda_warp_reduce +warp_sum = _cuda_op.cuda_warp_sum +warp_max = _cuda_op.cuda_warp_max +warp_min = _cuda_op.cuda_warp_min +cta_reduce = _cuda_op.cuda_cta_reduce +cta_sum = _cuda_op.cuda_cta_sum +cta_max = _cuda_op.cuda_cta_max +cta_min = _cuda_op.cuda_cta_min +atomic_cas = _cuda_op.cuda_atomic_cas +func_call = _cuda_op.cuda_func_call +ldg = _cuda_op.cuda_ldg +sm100_2sm_leader_smem_addr_composed = _cuda_op.cuda_sm100_2sm_leader_smem_addr +timer_init = _cuda_op.timer_init_cuda +timer_start = _cuda_op.timer_start_cuda +timer_end = _cuda_op.timer_end_cuda +timer_finalize = _cuda_op.timer_finalize_cuda +mma_store = _dtype_forward(_cuda_op.mma_store) +mma_fill = _dtype_forward(_cuda_op.mma_fill) +mma_store_legacy = _dtype_forward(_cuda_op.mma_store_legacy) +mma_fill_legacy = _dtype_forward(_cuda_op.mma_fill_legacy) +mov_sreg.__tvm_op__ = _ir.Op.get("tirx.cuda.mov_sreg") +wait_until.__tvm_op__ = _ir.Op.get("tirx.cuda.wait_until") +atomic_add.__tvm_op__ = _ir.Op.get("tirx.cuda.atomic_add") +warp_reduce.__tvm_op__ = _ir.Op.get("tirx.cuda.warp_reduce") +cta_reduce.__tvm_op__ = _ir.Op.get("tirx.cuda.cta_reduce") +atomic_cas.__tvm_op__ = _ir.Op.get("tirx.cuda.atomic_cas") +func_call.__tvm_op__ = _ir.Op.get("tirx.cuda.func_call") +ldg.__tvm_op__ = _ir.Op.get("tirx.cuda.ldg") +__shfl_sync = _shfl_sync +__shfl_sync.__tvm_op__ = _ir.Op.get("tirx.cuda.__shfl_sync") +__shfl_up_sync = _shfl_up_sync +__shfl_up_sync.__tvm_op__ = _ir.Op.get("tirx.cuda.__shfl_up_sync") +__shfl_down_sync = _shfl_down_sync +__shfl_down_sync.__tvm_op__ = _ir.Op.get("tirx.cuda.__shfl_down_sync") +__shfl_xor_sync = _shfl_xor_sync +__shfl_xor_sync.__tvm_op__ = _ir.Op.get("tirx.cuda.__shfl_xor_sync") +__activemask = _activemask +__activemask.__tvm_op__ = _ir.Op.get("tirx.cuda.__activemask") + +_ir.op._init_op_api("tirx.cuda", __name__) class NVSHMEMNamespace: @@ -300,6 +253,3 @@ def __call__(self, *args, **kwds): # __call__ corresponds to nvshmem_putmem_signal_nbi __tir_call_op_name__ = "nvshmem_putmem_signal_nbi" - - -__all__ = ["CUDANamespace", "NVSHMEMNamespace", "PTXLegacyNamespace", "STIRNamespace"] diff --git a/python/tvm/ir/op.py b/python/tvm/ir/op.py index 6dc519d9b9fb..26574bbdc3d2 100644 --- a/python/tvm/ir/op.py +++ b/python/tvm/ir/op.py @@ -17,7 +17,10 @@ # pylint: disable=invalid-name """Primitive operators in the TVM IR.""" +import keyword +import sys from collections.abc import Sequence +from types import ModuleType, SimpleNamespace import tvm_ffi @@ -25,6 +28,91 @@ from .expr import Expr +def _make_op_api(op, module_name): + """Build a callable whose operands and result are governed by its Op.""" + from .expr import Call, reinfer_type # pylint: disable=import-outside-toplevel + + def call(*args, attrs=None, ty_args=None, span=None, ret_ty=None, **kwargs): + # Bind named operands using the registered signature, without inventing + # defaults or interpreting arbitrary Python wrapper signatures. + operands = list(args) + for index, info in enumerate(op.args_info): + if index < len(args): + if info.name in kwargs: + raise TypeError(f"{op.name}: multiple values for {info.name!r}") + elif info.name in kwargs: + operands.append(kwargs.pop(info.name)) + else: + raise TypeError(f"{op.name}: missing operand {info.name!r}") + if kwargs: + raise TypeError(f"{op.name}: unexpected keyword operands {tuple(kwargs)}") + if ret_ty is None and ( + op.get_attr("TFixedReturnType") is not None or op.get_attr("FInferType") is not None + ): + provisional = Call.unchecked(op, operands, attrs=attrs, ty_args=ty_args, span=span) + ret_ty = reinfer_type(provisional) + return Call(op, operands, attrs=attrs, ty_args=ty_args, span=span, ret_ty=ret_ty) + + call.__name__ = op.name.rsplit(".", 1)[-1] + call.__module__ = module_name + call.__doc__ = op.doc or f"Construct a call to {op.name}." + call.__tvm_op__ = op + return call + + +def _init_op_api(namespace, target_module_name=None): + """Initialize registered Op callables in an already loaded Python module. + + Like :func:`tvm_ffi.init_ffi_api`, the registry prefix comes first and the + target defaults to that module name. For example, a backend module uses + ``_init_op_api("tirx.cuda", __name__)``. Dotted suffixes require existing + module or SimpleNamespace containers; underscores are ordinary name parts. + + Generated functions accept registered positional/named operands plus + ``attrs``, ``ty_args``, ``span`` and ``ret_ty``. Omitting the result invokes + an available Op inference hook; without one, Call retains a missing type. + Explicit results and inference/validation errors are preserved. + + Existing generated functions are reused. A deliberate wrapper may declare + ``__tvm_op__ = Op.get(name)`` to retain ownership of that name. This declares + identity, not semantic equivalence: the wrapper must accept printed calls + or have an appropriate exceptional printer hook. Other collisions fail + before any functions are installed. Returns None. + """ + if not namespace or any( + not part.isidentifier() or keyword.iskeyword(part) for part in namespace.split(".") + ): + raise ValueError(f"Invalid Op namespace {namespace!r}") + target = sys.modules[target_module_name or namespace] + pending = [] + destinations = set() + for name in sorted(Op.list_op_names()): + if not name.startswith(namespace + "."): + continue + parts = name[len(namespace) + 1 :].split(".") + if any(not part.isidentifier() or keyword.iskeyword(part) for part in parts): + raise ValueError(f"Op {name!r} has no Python attribute spelling") + container = target + for part in parts[:-1]: + container = vars(container).get(part) + if not isinstance(container, ModuleType | SimpleNamespace): + raise ValueError(f"Op {name!r} requires an existing namespace at {part!r}") + destination = (id(container), parts[-1]) + if destination in destinations: + raise ValueError(f"Op {name!r} aliases another exposure destination") + destinations.add(destination) + op = Op.get(name) + if parts[-1] in vars(container): + current = vars(container)[parts[-1]] + identity = getattr(current, "__tvm_op__", None) + if not callable(current) or not isinstance(identity, Op) or not identity.same_as(op): + raise ValueError(f"Op {name!r} conflicts with an existing Python attribute") + else: + pending.append((container, parts[-1], _make_op_api(op, target.__name__))) + for container, name, call in pending: + setattr(container, name, call) + + @tvm_ffi.register_object("ir.Op") class Op(Expr): """Primitive operator in the IR.""" diff --git a/python/tvm/tirx/script/ir_builder/op.py b/python/tvm/tirx/script/ir_builder/op.py index f2173f361f40..64665b893e77 100644 --- a/python/tvm/tirx/script/ir_builder/op.py +++ b/python/tvm/tirx/script/ir_builder/op.py @@ -297,7 +297,9 @@ def _get_script_namespace(name: str) -> object: raise AttributeError(f"No script namespace {name!r}") -def register_script_namespace(name: str, namespace: object, override: bool = False) -> object: +def register_script_namespace( + name: str, namespace: object, override: bool = False, *, register_printer_names: bool = True +) -> object: """Register a construction namespace and return it. Parameters @@ -309,6 +311,9 @@ def register_script_namespace(name: str, namespace: object, override: bool = Fal override : bool, optional Replace differing operator printer names if True. Existing equal names are reused; other duplicates raise ValueError. + register_printer_names : bool, optional + Discover printer names from legacy wrappers. Backends exposing canonical + registered Op names disable this to retain registration-owned spelling. """ _SCRIPT_NAMESPACES[name] = namespace globals()[name] = namespace @@ -330,7 +335,8 @@ def register_script_namespace(name: str, namespace: object, override: bool = Fal if isinstance(module_all, list) and name not in module_all: module_all.append(name) - _register_script_namespace_printer_names(namespace, f"tirx.{name}", override) + if register_printer_names: + _register_script_namespace_printer_names(namespace, f"tirx.{name}", override) return namespace diff --git a/src/backend/cuda/op/iket.cc b/src/backend/cuda/op/iket.cc index 2e8c4fdcb454..fbd427a84215 100644 --- a/src/backend/cuda/op/iket.cc +++ b/src/backend/cuda/op/iket.cc @@ -33,7 +33,6 @@ TVM_FFI_STATIC_INIT_BLOCK() { .set_attr("TFixedReturnType", PrimType::Void()) .set_attr("TIRxOpCategory", ffi::String("device_intrin")) .set_attr("TDeviceIntrinsicNamespace", ffi::String("cuda")) - .set_attr("TScriptPrinterName", ffi::String("tirx.cuda.iket.mark")) .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); OpDef("tirx.cuda.iket_range_start") @@ -41,7 +40,6 @@ TVM_FFI_STATIC_INIT_BLOCK() { .set_attr("TFixedReturnType", PrimType::UInt(32)) .set_attr("TIRxOpCategory", ffi::String("device_intrin")) .set_attr("TDeviceIntrinsicNamespace", ffi::String("cuda")) - .set_attr("TScriptPrinterName", ffi::String("tirx.cuda.iket.range_start")) .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); OpDef("tirx.cuda.iket_range_end") @@ -49,37 +47,34 @@ TVM_FFI_STATIC_INIT_BLOCK() { .set_attr("TFixedReturnType", PrimType::Void()) .set_attr("TIRxOpCategory", ffi::String("device_intrin")) .set_attr("TDeviceIntrinsicNamespace", ffi::String("cuda")) - .set_attr("TScriptPrinterName", ffi::String("tirx.cuda.iket.range_end")) .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); OpDef("tirx.cuda.iket_range_push") + .set_attr("TFixedReturnType", PrimType::Void()) .signature(sig::arg("name", "The name."), sig::var_args("args")) .set_attr("TIRxOpCategory", ffi::String("device_intrin")) .set_attr("TDeviceIntrinsicNamespace", ffi::String("cuda")) - .set_attr("TScriptPrinterName", ffi::String("tirx.cuda.iket.range_push")) .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); OpDef("tirx.cuda.iket_range_pop") + .set_attr("TFixedReturnType", PrimType::Void()) .set_attr("TIRxOpCategory", ffi::String("device_intrin")) .set_attr("TDeviceIntrinsicNamespace", ffi::String("cuda")) - .set_attr("TScriptPrinterName", ffi::String("tirx.cuda.iket.range_pop")) .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); OpDef("tirx.cuda.iket_sentinel_token") + .set_attr("TFixedReturnType", PrimType::UInt(32)) .signature(sig::arg("name", "The name.")) .set_attr("TIRxOpCategory", ffi::String("device_intrin")) .set_attr("TDeviceIntrinsicNamespace", ffi::String("cuda")) - .set_attr("TScriptPrinterName", - ffi::String("tirx.cuda.iket.sentinel_token")) .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); OpDef("tirx.cuda.iket_official_event") + .set_attr("TFixedReturnType", PrimType::UInt(32)) .signature(sig::arg("event_id", "The event identifier."), sig::arg("source_code", "The source code."), sig::var_args("args")) .set_attr("TIRxOpCategory", ffi::String("device_intrin")) .set_attr("TDeviceIntrinsicNamespace", ffi::String("cuda")) - .set_attr("TScriptPrinterName", - ffi::String("tirx.cuda.iket.official_event")) .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); } diff --git a/src/backend/cuda/op/target_builtin.cc b/src/backend/cuda/op/target_builtin.cc index 4ac2b7b8ee7e..9be8f0a0f380 100644 --- a/src/backend/cuda/op/target_builtin.cc +++ b/src/backend/cuda/op/target_builtin.cc @@ -240,16 +240,10 @@ TVM_FFI_NO_INLINE DeviceIntrinsicNames MakeDeviceIntrinsicNames(const char* op_n if (suffix.rfind(prefix, 0) == 0) { suffix = suffix.substr(prefix.size()); } - std::string canonical = "tirx." + namespace_name + "." + suffix; - // Match the nested construction namespaces at the canonical registration site. - if (namespace_name == "cuda" && - (suffix.rfind("tcgen05_", 0) == 0 || suffix.rfind("wgmma_", 0) == 0 || - suffix.rfind("iket_", 0) == 0)) { - suffix[suffix.find('_')] = '.'; - } else if (namespace_name == "nvshmem" && - ((suffix.size() >= 6 && suffix.compare(suffix.size() - 6, 6, "_block") == 0) || - (suffix.size() >= 5 && suffix.compare(suffix.size() - 5, 5, "_warp") == 0))) { + if (namespace_name == "nvshmem" && + ((suffix.size() >= 6 && suffix.compare(suffix.size() - 6, 6, "_block") == 0) || + (suffix.size() >= 5 && suffix.compare(suffix.size() - 5, 5, "_warp") == 0))) { suffix[suffix.rfind('_')] = '.'; } return {std::move(canonical), "tirx." + namespace_name + "." + suffix}; @@ -260,8 +254,10 @@ TVM_FFI_NO_INLINE void RegisterDeviceIntrinsicAttrs(OpDef& def, const char* op_n const std::string& printer_name) { def.set_attr("TIRxOpCategory", ffi::String("device_intrin")) .set_attr("TDeviceIntrinsicNamespace", ffi::String(op_namespace)) - .set_attr("TCallEffectKind", static_cast(effect_kind)) - .set_attr("TScriptPrinterName", ffi::String(printer_name)); + .set_attr("TCallEffectKind", static_cast(effect_kind)); + if (std::string(op_namespace) != "cuda") { + def.set_attr("TScriptPrinterName", ffi::String(printer_name)); + } } template @@ -438,6 +434,31 @@ void RegisterDeviceIntrinsicAliases() { sig::arg("b_offset"), sig::arg("acc_ptr"), sig::arg("c_offset"), sig::arg("saturate"), sig::var_args("args")); + for (const char* name : {"tirx.cuda.thread_rank", "tirx.cuda.any_sync", "tirx.cuda.ffs_u32"}) { + OpDef(name).set_attr("TFixedReturnType", PrimType::Int(32)); + } + for (const char* name : {"tirx.cuda.half2float", "tirx.cuda.bfloat162float", "tirx.cuda.fdividef", + "tirx.cuda.uint_as_float", "tirx.cuda.float2_x", "tirx.cuda.float2_y"}) { + OpDef(name).set_attr("TFixedReturnType", PrimType::Float(32)); + } + for (const char* name : + {"tirx.cuda.float22half2", "tirx.cuda.trap_when_assert_failed", + "tirx.cuda.runtime_instr_desc", "tirx.cuda.half8tofloat8", "tirx.cuda.float8tohalf8", + "tirx.cuda.mbarrier_wait_acquire_cluster", "tirx.cuda.warpgroup_sync"}) { + OpDef(name).set_attr("TFixedReturnType", PrimType::Void()); + } + for (const char* name : + {"tirx.cuda.get_tmem_addr", "tirx.cuda.cvta_generic_to_shared", + "tirx.cuda.smem_addr_from_uint64", "tirx.cuda.float_as_uint", "tirx.cuda.ballot_sync", + "tirx.cuda.reduce_add_sync_u32", "tirx.cuda.reduce_min_sync_u32", + "tirx.cuda.float22bfloat162_rn", "tirx.cuda.float22bfloat162_rn_from_float2", + "tirx.cuda.hmin2", "tirx.cuda.hmax2", "tirx.cuda.fp8x4_e4m3_from_float4"}) { + OpDef(name).set_attr("TFixedReturnType", PrimType::UInt(32)); + } + for (const char* name : {"tirx.cuda.clock64", "tirx.cuda.make_float2", "tirx.cuda.fmul2_rn", + "tirx.cuda.fadd2_rn", "tirx.cuda.bfloat1622float2"}) { + OpDef(name).set_attr("TFixedReturnType", PrimType::UInt(64)); + } for (const char* name : { "tirx.cuda.tcgen05_encode_matrix_descriptor", "tirx.cuda.wgmma_encode_matrix_descriptor", diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc index 4b106b672b2c..fbf08f0323ed 100644 --- a/src/tirx/script/printer/expr.cc +++ b/src/tirx/script/printer/expr.cc @@ -510,16 +510,58 @@ ffi::Optional TIRCallDocTranslate(DocTranslatorObj* d, const CallNode* // their stored Call argument list. Keep the lossless I.Call form until a // dedicated translation covers each signature. bool incompatible_signature = (op.value()->name == "tirx.cuda.ldg" && call->args.size() != 2) || - op.value()->name == "tirx.cuda.wait_until"; - // IKET lowering accepts a variadic IR signature, but its named constructors - // expose just the event/token and an optional payload. - if (op.value()->name == "tirx.cuda.iket_mark" || - op.value()->name == "tirx.cuda.iket_range_start" || - op.value()->name == "tirx.cuda.iket_range_end") { - incompatible_signature = call->args.empty() || call->args.size() > 2; + op.value()->name == "tirx.cuda.wait_until" || + op.value()->name == "tirx.cuda.mov_sreg"; + // Meaningful wrappers choose their result from operands, independently of + // customizable inference hooks. Only use them when that choice is lossless. + const std::string op_name = op.value()->name; + int result_operand = -1; + if (op_name == "tirx.cuda.atomic_add" || op_name == "tirx.cuda.atomic_cas" || + op_name == "tirx.cuda.__shfl_sync" || op_name == "tirx.cuda.__shfl_up_sync" || + op_name == "tirx.cuda.__shfl_down_sync" || op_name == "tirx.cuda.__shfl_xor_sync") { + result_operand = 1; + } else if (op_name == "tirx.cuda.warp_reduce" || op_name == "tirx.cuda.cta_reduce") { + result_operand = 0; } - if (names.count(op.value()) && !incompatible_signature) { - std::string name = names[op.value()]; + if (result_operand >= 0 && + (call->args.size() <= static_cast(result_operand) || + !ffi::StructuralEqual()(call->ty, call->args[result_operand]->ty))) { + incompatible_signature = true; + } + if (op_name == "tirx.cuda.__activemask" && + !ffi::StructuralEqual()(call->ty, PrimType::UInt(32))) { + incompatible_signature = true; + } + if (op_name == "tirx.cuda.ldg" && call->args.size() == 2) { + auto dtype = call->args[1].as(); + auto result = call->ty.as(); + if (!dtype || !result || ffi::DLDataTypeToString(result.value()->dtype) != dtype->value) { + incompatible_signature = true; + } + } + // CUDA exposes every registered canonical name. Other backends still own + // their explicit printer-name registrations until their surfaces migrate. + bool canonical_cuda = op_name.find("tirx.cuda.") == 0; + if (canonical_cuda) { + // Generated APIs validate construction; preserve provisional or invalid + // Calls through the explicit unchecked reconstruction surface instead. + try { + op.value().Validate(call); + } catch (const ffi::Error&) { + return RawCall(d, call, false); + } + if (op_name.find("tirx.cuda.__shfl") == 0 && call->args.size() > 1 && + call->args[1].as() && call->args[1]->ty.as()) { + return RawCall(d, call, false); + } + if ((result_operand >= 0 || op_name == "tirx.cuda.ldg") && + std::any_of(call->args.begin(), call->args.end(), + [](const Expr& arg) { return arg.as() != nullptr; })) { + return RawCall(d, call, false); + } + } + if ((names.count(op.value()) || canonical_cuda) && !incompatible_signature) { + std::string name = names.get(op.value(), op_name); if (!name.empty()) { ExprDoc callee = NamedCallCallee(name); ffi::Array named_args; @@ -532,6 +574,10 @@ ffi::Optional TIRCallDocTranslate(DocTranslatorObj* d, const CallNode* named_args.push_back(LiteralDoc::Str(string->value, std::nullopt)); } else { ExprDoc argument = args[i]; + // Canonical CUDA APIs preserve structured IR operands. + if (op.value()->name.find("tirx.cuda.") == 0) { + argument = MaterializeCallArgument(d, call->args[i], argument); + } if (const auto* address = call->args[i].as(); is_ptx && address && IsPTXAddressCall(address)) { // Use the address constructor only when it preserves its operands. @@ -541,7 +587,9 @@ ffi::Optional TIRCallDocTranslate(DocTranslatorObj* d, const CallNode* } bool reads_operand_type = is_ptx || - (i == 0 && (op.value()->name == "tirx.selector" || + (i == 0 && (op.value()->name == "tirx.cuda.warp_reduce" || + op.value()->name == "tirx.cuda.cta_reduce" || + op.value()->name == "tirx.selector" || op.value()->name == "tirx.webgpu.subgroup_shuffle" || op.value()->name == "tirx.webgpu.subgroup_shuffle_up" || op.value()->name == "tirx.webgpu.subgroup_shuffle_down" || diff --git a/tests/python/ir/test_op_api.py b/tests/python/ir/test_op_api.py new file mode 100644 index 000000000000..4fd0217612d3 --- /dev/null +++ b/tests/python/ir/test_op_api.py @@ -0,0 +1,345 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +"""Registered operator exposure and canonical CUDA script construction.""" + +import importlib +import os +import re +import subprocess +import sys +from types import ModuleType, SimpleNamespace + +import pytest + +import tvm +from tvm import ir +from tvm.ir.op import _init_op_api +from tvm.script import tirx as T + + +def register(name, args=(), **kwargs): + ir._ffi_api.RegisterOp(name, "Initializer regression") + op = ir.Op.get(name) + op.set_signature(args, **kwargs) + op.set_attr("TCallEffectKind", 3) + if name.startswith("tirx.cuda."): + op.set_attr("TIRxOpCategory", "device_intrin") + op.set_attr("TDeviceIntrinsicNamespace", "cuda") + elif name.startswith("tirx."): + op.set_attr("TIRxOpCategory", "builtin") + return op + + +@pytest.fixture +def target(request, monkeypatch): + name = "test_op_api_" + re.sub(r"\W", "_", request.node.name) + module = ModuleType(name) + monkeypatch.setitem(sys.modules, name, module) + return module + + +def test_prefix_nested_reinitialization(target): + first = register(target.__name__ + ".__first") + register(target.__name__ + "_other.ignored") + nested = register(target.__name__ + ".nested.leaf") + target.nested = SimpleNamespace() + assert _init_op_api(target.__name__) is None + assert target.__first.__tvm_op__.same_as(first) + assert target.nested.leaf.__tvm_op__.same_as(nested) + assert not hasattr(target, "ignored") + original = target.__first + late = register(target.__name__ + ".late") + _init_op_api(target.__name__) + assert target.__first is original + assert target.late.__tvm_op__.same_as(late) + + +def test_explicit_module_and_wrapper(target): + op = register(target.__name__ + ".wrapped") + + def wrapper(): + return ir.Call(op, []) + + wrapper.__tvm_op__ = op + target.wrapped = wrapper + _init_op_api(target.__name__, target.__name__) + assert target.wrapped is wrapper + other = ModuleType(target.__name__ + "_target") + sys.modules[other.__name__] = other + try: + _init_op_api(target.__name__, other.__name__) + assert other.wrapped.__tvm_op__.same_as(op) + finally: + del sys.modules[other.__name__] + + +@pytest.mark.parametrize("occupant", [42, lambda: None]) +def test_collision_preflight(target, occupant): + register(target.__name__ + ".a_new") + register(target.__name__ + ".z_conflict") + target.z_conflict = occupant + with pytest.raises(ValueError, match="conflicts"): + _init_op_api(target.__name__) + assert not hasattr(target, "a_new") + + +def test_wrong_identity(target): + register(target.__name__ + ".leaf") + target.leaf = lambda: None + target.leaf.__tvm_op__ = ir.Op.get("tirx.cuda.clock64") + with pytest.raises(ValueError, match="conflicts"): + _init_op_api(target.__name__) + + +@pytest.mark.parametrize("suffix", ["nested.leaf", "bad..leaf", "class"]) +def test_missing_or_malformed_container(target, suffix): + register(target.__name__ + ".a_new") + register(target.__name__ + "." + suffix) + with pytest.raises(ValueError): + _init_op_api(target.__name__) + assert not hasattr(target, "a_new") + + +def test_aliased_container_collision(target): + register(target.__name__ + ".a.leaf") + register(target.__name__ + ".b.leaf") + target.a = target.b = SimpleNamespace() + with pytest.raises(ValueError, match="aliases"): + _init_op_api(target.__name__) + assert not vars(target.a) + + +def test_leaf_container_conflict(target): + register(target.__name__ + ".leaf") + register(target.__name__ + ".leaf.child") + with pytest.raises(ValueError, match="existing namespace"): + _init_op_api(target.__name__) + assert not hasattr(target, "leaf") + + +def test_missing_target(): + with pytest.raises(KeyError): + _init_op_api("test_op_api_unloaded") + with pytest.raises(ValueError, match="Invalid"): + _init_op_api("invalid..prefix") + + +def test_call_fields_and_customization(target): + op = register(target.__name__ + ".call", ["x"], ty_args=["T"]) + _init_op_api(target.__name__) + x = tvm.tirx.Var("x", "int16") + ty = ir.PrimType("int16") + span = ir.Span(ir.SourceName("api"), 1, 1, 1, 2) + call = target.call(x, attrs={"flag": 1}, ty_args=[ty], span=span) + assert call.ty.is_missing() + assert call.args[0].same_as(x) + assert call.span.same_as(span) + assert call.attrs["flag"] == 1 + ir.assert_structural_equal(call.ty_args[0], ty) + op.set_attr("FInferType", lambda c: c.args[0].ty) + ir.assert_structural_equal(target.call(x, ty_args=[ty]).ty, ty) + op.set_attr("FInferType", lambda c: ir.PrimType("int64"), override=True) + assert target.call(x=x, ty_args=[ty]).ty.dtype == "int64" + assert target.call(x, ty_args=[ty], ret_ty="float32").ty.dtype == "float32" + op.set_attr("TFixedReturnType", ir.PrimType("uint32")) + assert target.call(x, ty_args=[ty]).ty.dtype == "uint32" + with pytest.raises(TypeError): + target.call(x) # Required type argument is still validated. + with pytest.raises(TypeError, match="multiple values"): + target.call(x, x=x) + with pytest.raises(TypeError, match="unexpected"): + target.call(x, unknown=x) + + +def test_inference_failure_propagates(target): + op = register(target.__name__ + ".call") + + def fail(call): + raise ValueError("inference sentinel") + + op.set_attr("FInferType", fail) + _init_op_api(target.__name__) + with pytest.raises(ValueError, match="inference sentinel"): + target.call() + + +def test_cuda_inventory_and_import_identity(): + script = importlib.import_module("tvm.backend.cuda.script") + + assert T.cuda is script + for name in ir.Op.list_op_names(): + if name.startswith("tirx.cuda."): + current = script + for part in name.removeprefix("tirx.cuda.").split("."): + current = getattr(current, part) + assert callable(current) + assert current.__tvm_op__.same_as(ir.Op.get(name)), name + first = script.clock64 + wrapper = script.func_call + _init_op_api("tirx.cuda", script.__name__) + assert script.clock64 is first + assert script.func_call is wrapper + assert script.clock64().ty.dtype == "uint64" + assert callable(script.iket.mark) + assert callable(script.wgmma.encode_matrix_descriptor) + assert callable(script.tcgen05.encode_instr_descriptor) + assert script.sm100_2sm_leader_smem_addr(0).op.name == "tirx.cuda.sm100_2sm_leader_smem_addr" + + +@pytest.mark.parametrize("first", ["tvm.backend.cuda.script", "tvm.script.tirx"]) +def test_import_order(first): + code = f""" +import importlib +importlib.import_module({first!r}) +script = importlib.import_module("tvm.backend.cuda.script") +from tvm.script import tirx as T +assert T.cuda is script +assert T.cuda.clock64().ty.dtype == 'uint64' +""" + subprocess.run( + [sys.executable, "-S", "-c", code], + env=dict(os.environ, PYTHONPATH=os.pathsep.join(sys.path)), + check=True, + ) + + +def roundtrip(call, params=()): + func = tvm.tirx.PrimFunc(list(params), tvm.tirx.Evaluate(call)) + code = func.script() + reparsed = tvm.script.from_source(code, extra_vars={"T": T, "I": tvm.script.ir}) + ir.assert_structural_equal(func, reparsed) + return code + + +@pytest.mark.parametrize( + "name,args", + [ + ("clock64", ()), + ("iket_mark", ("mark", 1, 2, 3)), + ("iket_range_start", ("range", 1, 2)), + ("iket_range_end", (1, 2, 3)), + ("iket_range_push", ("range", 1, 2)), + ("iket_range_pop", ()), + ("iket_sentinel_token", ("token",)), + ("iket_official_event", (1, "source", 2, 3)), + ("wgmma_noop_barrier", (1,)), + ("__shfl_sync", (1, tvm.tirx.const(2, "int32"), 0, 32)), + ], +) +def test_canonical_roundtrip(name, args): + call = getattr(T.cuda, name)(*args) + code = roundtrip(call) + assert "T.cuda." + name + "(" in code + + +def test_retained_wrappers_and_raw_fields(): + p = tvm.tirx.Var("p", "handle") + for call in [ + T.cuda.atomic_add(p, 1), + T.cuda.atomic_cas(p, 1, 2), + T.cuda.ldg(p, "float32"), + T.cuda.func_call( + "f", 1, source_code="__device__ int f(int x) { return x; }", return_type="int32" + ), + T.cuda.tcgen05_encode_matrix_descriptor(p, p, 1, 2, 0), + T.cuda.wgmma_encode_matrix_descriptor(p, p, 1, 2, 0), + ]: + roundtrip(call, [p]) + for kwargs in [{"attrs": {"flag": 1}}, {"ret_ty": "float32"}, {"ret_ty": ir.Type.missing()}]: + assert "I.Call(" in roundtrip(T.cuda.clock64(**kwargs)) + + +def test_wrapper_custom_inference_is_lossless(): + op = ir.Op.get("tirx.cuda.atomic_add") + original = op.get_attr("FInferType") + p = tvm.tirx.Var("p", "handle") + try: + op.set_attr("FInferType", lambda c: ir.PrimType("int64"), override=True) + call = ir.Call(op, [p, tvm.tirx.const(1, "int32")], ret_ty="int64") + assert "I.Call(" in roundtrip(call, [p]) + finally: + op.set_attr("FInferType", original, override=True) + + +def test_generated_full_call_roundtrip(target): + op = register(target.__name__ + ".call", ["x"], ty_args=["T"]) + _init_op_api(target.__name__) + call = target.call(1, ty_args=[ir.PrimType("int16")], attrs={"flag": 1}, ret_ty="int32") + assert "I.Call(" in roundtrip(call) + assert op.get_attr("TScriptPrinterName") is None + + +def test_explicit_printer_alias_and_late_cuda_registration(): + module = importlib.import_module("tvm.backend.cuda.script") + op = register("tirx.cuda.test_op_api_alias") + op.set_attr("TFixedReturnType", ir.PrimType("int32")) + _init_op_api("tirx.cuda", module.__name__) + generated = module.test_op_api_alias + assert "T.cuda.test_op_api_alias(" in roundtrip(generated()) + module.test_op_api_custom_name = generated + op.set_attr("TScriptPrinterName", "tirx.cuda.test_op_api_custom_name") + try: + assert "T.cuda.test_op_api_custom_name(" in roundtrip(generated()) + finally: + op.set_attr("TScriptPrinterName", op.name, override=True) + del module.test_op_api_custom_name + + +def test_legacy_module_namespace_printer_discovery(): + from tvm.tirx.script.ir_builder.op import register_script_namespace + + op = register("tirx.test_op_api_legacy") + module = ModuleType("test_op_api_legacy") + module.wrapper = lambda: None + module.wrapper.__tir_op_name__ = "test_op_api_legacy" + register_script_namespace("test_op_api_legacy_namespace", module) + assert op.get_attr("TScriptPrinterName") == "tirx.test_op_api_legacy_namespace.wrapper" + + +def test_wait_and_vector_load_fallbacks(): + p = tvm.tirx.Var("p", "handle") + calls = [ + T.cuda.wait_until(p, "eq", 1), + T.cuda.ldg(p, "float32", dst=[p, p], vec="v2"), + ] + for call in calls: + assert "I.Call(" in roundtrip(call, [p]) + + +@pytest.mark.parametrize( + "name,args,ret_ty", + [ + ("clock64", [1], "uint64"), + ("atomic_add", [1, 2, 3], "int32"), + ("__shfl_sync", [], "int32"), + ], +) +def test_unchecked_calls_keep_lossless_fallback(name, args, ret_ty): + call = ir.Call.unchecked("tirx.cuda." + name, args, ret_ty=ret_ty) + assert "I.Call(" in roundtrip(call) + + +def test_retained_wrapper_tensor_region_fallback(): + buffer = tvm.tirx.decl_buffer((4,), "int32", name="A") + call = ir.Call("tirx.cuda.atomic_add", [buffer[0:1], 1], ret_ty="int32") + assert "I.Call(" in roundtrip(call, [buffer]) + + +def test_shuffle_buffer_operand_fallback(): + buffer = tvm.tirx.decl_buffer((1,), "int32", name="A") + call = ir.Call("tirx.cuda.__shfl_sync", [1, buffer, 0, 32], ret_ty=buffer.ty) + assert "I.Call(" in roundtrip(call, [buffer]) diff --git a/tests/python/tirx/script/test_tirx_script_printer.py b/tests/python/tirx/script/test_tirx_script_printer.py index b55e9623230f..c72eae858d5f 100644 --- a/tests/python/tirx/script/test_tirx_script_printer.py +++ b/tests/python/tirx/script/test_tirx_script_printer.py @@ -748,7 +748,7 @@ def test_printer_ptx_more(): _assert_namespace_print( cuda_op.cuda_tcgen05_encode_matrix_descriptor(d, a, 1, 2, 0), "d: T.handle = T.handle()\na: T.handle = T.handle()\n" - "T.cuda.tcgen05.encode_matrix_descriptor(d, a, 1, 2, 0)", + "T.cuda.tcgen05_encode_matrix_descriptor(d, a, 1, 2, 0)", ) _assert_namespace_print( cuda_op.cuda_tcgen05_encode_instr_descriptor( @@ -1025,6 +1025,6 @@ def test_printer_ptx_mma_and_wgmma(): _assert_namespace_print( cuda_op.cuda_wgmma_encode_matrix_descriptor(d, a, 1, 1, 0), "d: T.handle = T.handle()\na: T.handle = T.handle()\n" - "T.cuda.wgmma.encode_matrix_descriptor(d, a, 1, 1, 0)", + "T.cuda.wgmma_encode_matrix_descriptor(d, a, 1, 1, 0)", ) - _assert_namespace_print(cuda_op.cuda_wgmma_noop_barrier(0), "T.cuda.wgmma.noop_barrier(0)") + _assert_namespace_print(cuda_op.cuda_wgmma_noop_barrier(0), "T.cuda.wgmma_noop_barrier(0)") diff --git a/tests/python/tirx/test_op_namespace_cleanup.py b/tests/python/tirx/test_op_namespace_cleanup.py index c2337421b5fb..cec94791fa8c 100644 --- a/tests/python/tirx/test_op_namespace_cleanup.py +++ b/tests/python/tirx/test_op_namespace_cleanup.py @@ -147,9 +147,7 @@ def tile_aliases(A: T.Buffer((16,), "float32"), B: T.Buffer((16,), "float32")): def test_device_intrinsic_namespaces_are_canonical_and_classified(): - from tvm.backend.cuda.script import ( - CUDANamespace as BackendCUDANamespace, - ) + cuda_script = importlib.import_module("tvm.backend.cuda.script") from tvm.backend.cuda.script import ( NVSHMEMNamespace as BackendNVSHMEMNamespace, ) @@ -160,7 +158,7 @@ def test_device_intrinsic_namespaces_are_canonical_and_classified(): from tvm.backend.trn.script import NKINamespace as BackendNKINamespace from tvm.tirx.script.ir_builder import op as builder_op - assert isinstance(builder_op.cuda, BackendCUDANamespace) + assert builder_op.cuda is cuda_script assert isinstance(builder_op.s_tir, BackendSTIRNamespace) assert isinstance(builder_op.nvshmem, BackendNVSHMEMNamespace) assert isinstance(builder_op.metal, BackendMetalNamespace) @@ -386,6 +384,8 @@ def test_registered_tirx_ops_have_exactly_one_category(): elif category == "device_intrin": assert device_namespace in device_namespaces, op_name printer_name = _op_attr(op_name, "TScriptPrinterName") + if device_namespace == "cuda" and printer_name is None: + printer_name = op_name assert printer_name is not None, op_name assert printer_name.startswith("tirx." + device_namespace + "."), op_name assert _has_path(T, printer_name.removeprefix("tirx.")), op_name From 048d80d903159e55a068db2c9fd75943c876f807 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Sat, 3 Oct 2026 08:41:41 +0000 Subject: [PATCH 2/4] [REFACTOR][IR] Keep Op API coverage in existing suites --- tests/python/ir/test_op_api.py | 354 --------------------------------- 1 file changed, 354 deletions(-) delete mode 100644 tests/python/ir/test_op_api.py diff --git a/tests/python/ir/test_op_api.py b/tests/python/ir/test_op_api.py deleted file mode 100644 index 0ff0176c2c29..000000000000 --- a/tests/python/ir/test_op_api.py +++ /dev/null @@ -1,354 +0,0 @@ -# Licensed to the Apache Software Foundation (ASF) under one -# or more contributor license agreements. See the NOTICE file -# distributed with this work for additional information -# regarding copyright ownership. The ASF licenses this file -# to you under the Apache License, Version 2.0 (the -# "License"); you may not use this file except in compliance -# with the License. You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, -# software distributed under the License is distributed on an -# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY -# KIND, either express or implied. See the License for the -# specific language governing permissions and limitations -# under the License. -"""Registered operator exposure and canonical CUDA script construction.""" - -import importlib -import os -import re -import subprocess -import sys -from types import ModuleType, SimpleNamespace - -import pytest - -import tvm -from tvm import ir -from tvm.ir.op import _init_op_api -from tvm.script import tirx as T - - -def register(name, args=(), **kwargs): - ir._ffi_api.RegisterOp(name, "Initializer regression") - op = ir.Op.get(name) - op.set_signature(args, **kwargs) - op.set_attr("TCallEffectKind", 3) - if name.startswith("tirx.cuda."): - op.set_attr("TIRxOpCategory", "device_intrin") - op.set_attr("TDeviceIntrinsicNamespace", "cuda") - elif name.startswith("tirx."): - op.set_attr("TIRxOpCategory", "builtin") - return op - - -@pytest.fixture -def target(request, monkeypatch): - name = "test_op_api_" + re.sub(r"\W", "_", request.node.name) - module = ModuleType(name) - monkeypatch.setitem(sys.modules, name, module) - return module - - -def test_prefix_nested_reinitialization(target): - first = register(target.__name__ + ".__first") - register(target.__name__ + "_other.ignored") - nested = register(target.__name__ + ".nested.leaf") - target.nested = SimpleNamespace() - assert _init_op_api(target.__name__) is None - assert target.__first.__tvm_op__.same_as(first) - assert target.nested.leaf.__tvm_op__.same_as(nested) - assert not hasattr(target, "ignored") - original = target.__first - late = register(target.__name__ + ".late") - _init_op_api(target.__name__) - assert target.__first is original - assert target.late.__tvm_op__.same_as(late) - - -def test_explicit_module_and_wrapper(target): - op = register(target.__name__ + ".wrapped") - - def wrapper(): - return ir.Call(op, []) - - wrapper.__tvm_op__ = op - target.wrapped = wrapper - _init_op_api(target.__name__, target.__name__) - assert target.wrapped is wrapper - other = ModuleType(target.__name__ + "_target") - sys.modules[other.__name__] = other - try: - _init_op_api(target.__name__, other.__name__) - assert other.wrapped.__tvm_op__.same_as(op) - finally: - del sys.modules[other.__name__] - - -@pytest.mark.parametrize("occupant", [42, lambda: None]) -def test_collision_preflight(target, occupant): - register(target.__name__ + ".a_new") - register(target.__name__ + ".z_conflict") - target.z_conflict = occupant - with pytest.raises(ValueError, match="conflicts"): - _init_op_api(target.__name__) - assert not hasattr(target, "a_new") - - -def test_wrong_identity(target): - register(target.__name__ + ".leaf") - target.leaf = lambda: None - target.leaf.__tvm_op__ = ir.Op.get("tirx.cuda.clock64") - with pytest.raises(ValueError, match="conflicts"): - _init_op_api(target.__name__) - - -@pytest.mark.parametrize("suffix", ["nested.leaf", "bad..leaf", "class"]) -def test_missing_or_malformed_container(target, suffix): - register(target.__name__ + ".a_new") - register(target.__name__ + "." + suffix) - with pytest.raises(ValueError): - _init_op_api(target.__name__) - assert not hasattr(target, "a_new") - - -def test_aliased_container_collision(target): - register(target.__name__ + ".a.leaf") - register(target.__name__ + ".b.leaf") - target.a = target.b = SimpleNamespace() - with pytest.raises(ValueError, match="aliases"): - _init_op_api(target.__name__) - assert not vars(target.a) - - -def test_leaf_container_conflict(target): - register(target.__name__ + ".leaf") - register(target.__name__ + ".leaf.child") - with pytest.raises(ValueError, match="existing namespace"): - _init_op_api(target.__name__) - assert not hasattr(target, "leaf") - - -def test_missing_target(): - with pytest.raises(KeyError): - _init_op_api("test_op_api_unloaded") - with pytest.raises(ValueError, match="Invalid"): - _init_op_api("invalid..prefix") - - -def test_call_fields_and_customization(target): - op = register(target.__name__ + ".call", ["x"], ty_args=["T"]) - _init_op_api(target.__name__) - x = tvm.tirx.Var("x", "int16") - ty = ir.PrimType("int16") - span = ir.Span(ir.SourceName("api"), 1, 1, 1, 2) - call = target.call(x, attrs={"flag": 1}, ty_args=[ty], span=span) - assert call.ty.is_missing() - assert call.args[0].same_as(x) - assert call.span.same_as(span) - assert call.attrs["flag"] == 1 - ir.assert_structural_equal(call.ty_args[0], ty) - op.set_attr("FInferType", lambda c: c.args[0].ty) - ir.assert_structural_equal(target.call(x, ty_args=[ty]).ty, ty) - op.set_attr("FInferType", lambda c: ir.PrimType("int64"), override=True) - assert target.call(x=x, ty_args=[ty]).ty.dtype == "int64" - assert target.call(x, ty_args=[ty], ret_ty="float32").ty.dtype == "float32" - op.set_attr("TFixedReturnType", ir.PrimType("uint32")) - assert target.call(x, ty_args=[ty]).ty.dtype == "uint32" - with pytest.raises(TypeError): - target.call(x) # Required type argument is still validated. - with pytest.raises(TypeError, match="multiple values"): - target.call(x, x=x) - with pytest.raises(TypeError, match="unexpected"): - target.call(x, unknown=x) - - -def test_inference_failure_propagates(target): - op = register(target.__name__ + ".call") - - def fail(call): - raise ValueError("inference sentinel") - - op.set_attr("FInferType", fail) - _init_op_api(target.__name__) - with pytest.raises(ValueError, match="inference sentinel"): - target.call() - - -def test_cuda_inventory_and_import_identity(): - script = importlib.import_module("tvm.backend.cuda.script") - - assert T.cuda is script - for name in ir.Op.list_op_names(): - if name.startswith("tirx.cuda."): - current = script - for part in name.removeprefix("tirx.cuda.").split("."): - current = getattr(current, part) - assert callable(current) - assert current.__tvm_op__.same_as(ir.Op.get(name)), name - first = script.clock64 - wrapper = script.func_call - _init_op_api("tirx.cuda", script.__name__) - assert script.clock64 is first - assert script.func_call is wrapper - assert script.clock64().ty.dtype == "uint64" - assert callable(script.iket.mark) - assert callable(script.wgmma.encode_matrix_descriptor) - assert callable(script.tcgen05.encode_instr_descriptor) - assert script.sm100_2sm_leader_smem_addr(0).op.name == "tirx.cuda.sm100_2sm_leader_smem_addr" - - -@pytest.mark.parametrize("first", ["tvm.backend.cuda.script", "tvm.script.tirx"]) -def test_import_order(first): - code = f""" -import importlib -importlib.import_module({first!r}) -script = importlib.import_module("tvm.backend.cuda.script") -from tvm.script import tirx as T -assert T.cuda is script -assert T.cuda.clock64().ty.dtype == 'uint64' -""" - subprocess.run( - [sys.executable, "-S", "-c", code], - env=dict(os.environ, PYTHONPATH=os.pathsep.join(sys.path)), - check=True, - ) - - -def roundtrip(call, params=()): - func = tvm.tirx.PrimFunc(list(params), tvm.tirx.Evaluate(call)) - code = func.script() - reparsed = tvm.script.from_source(code, extra_vars={"T": T, "I": tvm.script.ir}) - ir.assert_structural_equal(func, reparsed) - return code - - -@pytest.mark.parametrize( - "name,args", - [ - ("clock64", ()), - ("iket_mark", ("mark", 1, 2, 3)), - ("iket_range_start", ("range", 1, 2)), - ("iket_range_end", (1, 2, 3)), - ("iket_range_push", ("range", 1, 2)), - ("iket_range_pop", ()), - ("iket_sentinel_token", ("token",)), - ("iket_official_event", (1, "source", 2, 3)), - ("wgmma_noop_barrier", (1,)), - ("__shfl_sync", (1, tvm.tirx.const(2, "int32"), 0, 32)), - ], -) -def test_canonical_roundtrip(name, args): - call = getattr(T.cuda, name)(*args) - code = roundtrip(call) - assert "T.cuda." + name + "(" in code - - -def test_retained_wrappers_and_raw_fields(): - p = tvm.tirx.Var("p", "handle") - for call in [ - T.cuda.atomic_add(p, 1), - T.cuda.atomic_cas(p, 1, 2), - T.cuda.ldg(p, "float32"), - T.cuda.func_call( - "f", 1, source_code="__device__ int f(int x) { return x; }", return_type="int32" - ), - T.cuda.tcgen05_encode_matrix_descriptor(p, p, 1, 2, 0), - T.cuda.wgmma_encode_matrix_descriptor(p, p, 1, 2, 0), - ]: - roundtrip(call, [p]) - for kwargs in [{"attrs": {"flag": 1}}, {"ret_ty": "float32"}, {"ret_ty": ir.Type.missing()}]: - assert "I.Call.unchecked(" in roundtrip(T.cuda.clock64(**kwargs)) - - -def test_wrapper_custom_inference_is_lossless(): - op = ir.Op.get("tirx.cuda.atomic_add") - original = op.get_attr("FInferType") - p = tvm.tirx.Var("p", "handle") - try: - op.set_attr("FInferType", lambda c: ir.PrimType("int64"), override=True) - call = ir.Call(op, [p, tvm.tirx.const(1, "int32")], ty="int64") - assert "I.Call.unchecked(" in roundtrip(call, [p]) - finally: - op.set_attr("FInferType", original, override=True) - - -def test_generated_full_call_roundtrip(target): - op = register(target.__name__ + ".call", ["x"], ty_args=["T"]) - _init_op_api(target.__name__) - call = target.call(1, ty_args=[ir.PrimType("int16")], attrs={"flag": 1}, ret_ty="int32") - assert "I.Call.unchecked(" in roundtrip(call) - assert op.get_attr("TScriptPrinterName") is None - - -def test_explicit_printer_alias_and_late_cuda_registration(): - module = importlib.import_module("tvm.backend.cuda.script") - op = register("tirx.cuda.test_op_api_alias") - op.set_attr("TFixedReturnType", ir.PrimType("int32")) - # Registration or generic initialization alone does not publish a name - # through the script namespace registration contract. - call = ir.Call(op, [], ty="int32") - assert "I.Call.unchecked(" in roundtrip(call) - _init_op_api("tirx.cuda", module.__name__) - generated = module.test_op_api_alias - assert "I.Call.unchecked(" in roundtrip(generated()) - from tvm.tirx.script.ir_builder.op import register_script_namespace - - register_script_namespace("cuda", module, canonical_op_names=True) - assert "T.cuda.test_op_api_alias(" in roundtrip(generated()) - module.test_op_api_custom_name = generated - op.set_attr("TScriptPrinterName", "tirx.cuda.test_op_api_custom_name", override=True) - try: - register_script_namespace("cuda", module, canonical_op_names=True) - assert "T.cuda.test_op_api_custom_name(" in roundtrip(generated()) - finally: - op.set_attr("TScriptPrinterName", op.name, override=True) - del module.test_op_api_custom_name - - -def test_legacy_module_namespace_printer_discovery(): - from tvm.tirx.script.ir_builder.op import register_script_namespace - - op = register("tirx.test_op_api_legacy") - module = ModuleType("test_op_api_legacy") - module.wrapper = lambda: None - module.wrapper.__tir_op_name__ = "test_op_api_legacy" - register_script_namespace("test_op_api_legacy_namespace", module) - assert op.get_attr("TScriptPrinterName") == "tirx.test_op_api_legacy_namespace.wrapper" - - -def test_wait_and_vector_load_fallbacks(): - p = tvm.tirx.Var("p", "handle") - calls = [ - T.cuda.wait_until(p, "eq", 1), - T.cuda.ldg(p, "float32", dst=[p, p], vec="v2"), - ] - for call in calls: - assert "I.Call.unchecked(" in roundtrip(call, [p]) - - -@pytest.mark.parametrize( - "name,args,ret_ty", - [ - ("clock64", [1], "uint64"), - ("atomic_add", [1, 2, 3], "int32"), - ("__shfl_sync", [], "int32"), - ], -) -def test_unchecked_calls_keep_lossless_fallback(name, args, ret_ty): - call = ir.Call.unchecked("tirx.cuda." + name, args, ty=ret_ty) - assert "I.Call.unchecked(" in roundtrip(call) - - -def test_retained_wrapper_tensor_region_fallback(): - buffer = tvm.tirx.decl_tensor((4,), "int32", name="A") - call = ir.Call("tirx.cuda.atomic_add", [buffer[0:1], 1], ty="int32") - assert "I.Call.unchecked(" in roundtrip(call, [buffer]) - - -def test_shuffle_buffer_operand_fallback(): - buffer = tvm.tirx.decl_tensor((1,), "int32", name="A") - call = ir.Call("tirx.cuda.__shfl_sync", [1, buffer, 0, 32], ty=buffer.ty) - assert "I.Call.unchecked(" in roundtrip(call, [buffer]) From 1d925a03f135003a703d66aeabaea259ddfd7ec7 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Sat, 3 Oct 2026 08:41:41 +0000 Subject: [PATCH 3/4] [REFACTOR][CUDA] Chain metadata on intrinsic registrations Accept a temporary OpDef and return its reference for chained result metadata. Keep canonical names explicit and colocate fixed and inferred result types with each signature. --- src/backend/cuda/op/target_builtin.cc | 488 +++++++++++++------------- 1 file changed, 241 insertions(+), 247 deletions(-) diff --git a/src/backend/cuda/op/target_builtin.cc b/src/backend/cuda/op/target_builtin.cc index 9be8f0a0f380..144f04bd003a 100644 --- a/src/backend/cuda/op/target_builtin.cc +++ b/src/backend/cuda/op/target_builtin.cc @@ -29,16 +29,14 @@ #include #include -#include #include -#include namespace tvm { namespace tirx { namespace builtin { namespace { -void RegisterDeviceIntrinsicAliases(); +void RegisterDeviceIntrinsics(); Type InferTypeCudaLdg(const CallNode* call) { TVM_FFI_CHECK_GE(call->args.size(), 2U, ValueError) @@ -221,37 +219,21 @@ void RegisterCudaTargetBuiltins() { .set_attr("TIRxOpCategory", ffi::String("builtin")) .set_attr("TCallEffectKind", static_cast(CallEffectKind::kOpaque)); - RegisterDeviceIntrinsicAliases(); + RegisterDeviceIntrinsics(); } namespace { -struct DeviceIntrinsicNames { - std::string canonical; - std::string printer; -}; - -TVM_FFI_NO_INLINE DeviceIntrinsicNames MakeDeviceIntrinsicNames(const char* op_name, - const char* op_namespace) { - std::string name(op_name); - std::string namespace_name(op_namespace); - std::string prefix = namespace_name + "_"; - std::string suffix = name; - if (suffix.rfind(prefix, 0) == 0) { - suffix = suffix.substr(prefix.size()); - } - std::string canonical = "tirx." + namespace_name + "." + suffix; - if (namespace_name == "nvshmem" && - ((suffix.size() >= 6 && suffix.compare(suffix.size() - 6, 6, "_block") == 0) || - (suffix.size() >= 5 && suffix.compare(suffix.size() - 5, 5, "_warp") == 0))) { - suffix[suffix.rfind('_')] = '.'; - } - return {std::move(canonical), "tirx." + namespace_name + "." + suffix}; -} - TVM_FFI_NO_INLINE void RegisterDeviceIntrinsicAttrs(OpDef& def, const char* op_namespace, - CallEffectKind effect_kind, - const std::string& printer_name) { + CallEffectKind effect_kind) { + std::string printer_name = def.op()->name; + if (std::string(op_namespace) == "nvshmem" && + ((printer_name.size() >= 6 && + printer_name.compare(printer_name.size() - 6, 6, "_block") == 0) || + (printer_name.size() >= 5 && + printer_name.compare(printer_name.size() - 5, 5, "_warp") == 0))) { + printer_name[printer_name.rfind('_')] = '.'; + } def.set_attr("TIRxOpCategory", ffi::String("device_intrin")) .set_attr("TDeviceIntrinsicNamespace", ffi::String(op_namespace)) .set_attr("TCallEffectKind", static_cast(effect_kind)); @@ -261,247 +243,259 @@ TVM_FFI_NO_INLINE void RegisterDeviceIntrinsicAttrs(OpDef& def, const char* op_n } template -void RegisterDeviceIntrinsic(const char* op_name, const char* op_namespace, - CallEffectKind effect_kind, const Specs&... specs) { - DeviceIntrinsicNames names = MakeDeviceIntrinsicNames(op_name, op_namespace); - OpDef def(names.canonical); +OpDef& RegisterDeviceIntrinsic(OpDef&& def, const char* op_namespace, CallEffectKind effect_kind, + const Specs&... specs) { def.signature(specs...); - RegisterDeviceIntrinsicAttrs(def, op_namespace, effect_kind, names.printer); + RegisterDeviceIntrinsicAttrs(def, op_namespace, effect_kind); + return def; } -void RegisterDeviceIntrinsicAliases() { - RegisterDeviceIntrinsic("cuda_any_sync", "cuda", CallEffectKind::kPure, sig::arg("mask"), - sig::arg("pred")); - RegisterDeviceIntrinsic("cuda_atomic_add", "cuda", CallEffectKind::kOpaque, sig::arg("res_addr"), - sig::arg("value")); - RegisterDeviceIntrinsic("cuda_atomic_cas", "cuda", CallEffectKind::kOpaque, sig::arg("ptr"), - sig::arg("old_val"), sig::arg("new_val")); - RegisterDeviceIntrinsic("cuda_wait_until", "cuda", CallEffectKind::kOpaque, sig::arg("dst"), - sig::arg("ptr"), sig::arg("condition"), sig::arg("scope"), - sig::arg("space"), sig::arg("ptx_type"), sig::arg("backoff_ns")); - RegisterDeviceIntrinsic("cuda_ballot_sync", "cuda", CallEffectKind::kOpaque, - sig::arg("mask"), sig::arg("pred")); - RegisterDeviceIntrinsic("cuda_bfloat1622float2", "cuda", CallEffectKind::kOpaque, - sig::arg("packed")); - RegisterDeviceIntrinsic("cuda_bfloat162float", "cuda", CallEffectKind::kOpaque, sig::arg("src")); - RegisterDeviceIntrinsic("cuda_clock64", "cuda", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("cuda_cluster_sync", "cuda", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("cuda_cta_reduce", "cuda", CallEffectKind::kOpaque, sig::arg("value"), - sig::arg("op"), sig::arg("num_warps"), sig::arg("scratch")); - RegisterDeviceIntrinsic("cuda_cta_sync", "cuda", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("cuda_cvta_generic_to_shared", "cuda", CallEffectKind::kOpaque, - sig::arg("ptr")); - RegisterDeviceIntrinsic("cuda_elect_sync", "cuda", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("cuda_fadd2_rn", "cuda", CallEffectKind::kOpaque, sig::arg("a"), - sig::arg("b")); - RegisterDeviceIntrinsic("cuda_fdividef", "cuda", CallEffectKind::kPure, sig::arg("x"), - sig::arg("y")); - RegisterDeviceIntrinsic("cuda_ffs_u32", "cuda", CallEffectKind::kOpaque, - sig::arg("value")); - RegisterDeviceIntrinsic("cuda_float22bfloat162_rn", "cuda", CallEffectKind::kOpaque, - sig::arg("v0"), sig::arg("v1")); - RegisterDeviceIntrinsic("cuda_float22bfloat162_rn_from_float2", "cuda", CallEffectKind::kOpaque, - sig::arg("packed")); - RegisterDeviceIntrinsic("cuda_float22half2", "cuda", CallEffectKind::kOpaque, sig::arg("dst"), - sig::arg("src")); - RegisterDeviceIntrinsic("cuda_float2_x", "cuda", CallEffectKind::kOpaque, sig::arg("packed")); - RegisterDeviceIntrinsic("cuda_float2_y", "cuda", CallEffectKind::kOpaque, sig::arg("packed")); - RegisterDeviceIntrinsic("cuda_float8tohalf8", "cuda", CallEffectKind::kOpaque, - sig::arg("src_addr"), sig::arg("dst_addr")); - RegisterDeviceIntrinsic("cuda_float_as_uint", "cuda", CallEffectKind::kOpaque, sig::arg("x")); - RegisterDeviceIntrinsic("cuda_fmul2_rn", "cuda", CallEffectKind::kOpaque, sig::arg("a"), - sig::arg("b")); - RegisterDeviceIntrinsic("cuda_fp8x4_e4m3_from_float4", "cuda", CallEffectKind::kOpaque, - sig::arg("x"), sig::arg("y"), sig::arg("z"), sig::arg("w")); - RegisterDeviceIntrinsic("cuda_func_call", "cuda", CallEffectKind::kOpaque, sig::arg("func_name"), - sig::var_args("args")); - RegisterDeviceIntrinsic("cuda_get_tmem_addr", "cuda", CallEffectKind::kOpaque, sig::arg("addr"), - sig::arg("row_offset"), sig::arg("col_offset")); - RegisterDeviceIntrinsic("cuda_grid_sync", "cuda", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("cuda_half2float", "cuda", CallEffectKind::kOpaque, sig::arg("src")); - RegisterDeviceIntrinsic("cuda_half8tofloat8", "cuda", CallEffectKind::kOpaque, - sig::arg("src_addr"), sig::arg("dst_addr")); - RegisterDeviceIntrinsic("cuda_hmax2", "cuda", CallEffectKind::kOpaque, sig::arg("a"), - sig::arg("b")); - RegisterDeviceIntrinsic("cuda_hmin2", "cuda", CallEffectKind::kOpaque, sig::arg("a"), - sig::arg("b")); - RegisterDeviceIntrinsic("cuda_ldg", "cuda", CallEffectKind::kOpaque, sig::var_args("args")); - RegisterDeviceIntrinsic("cuda_make_float2", "cuda", CallEffectKind::kOpaque, sig::arg("x"), - sig::arg("y")); - RegisterDeviceIntrinsic("cuda_mbarrier_wait", "cuda", CallEffectKind::kOpaque, sig::arg("bar"), - sig::arg("phase")); - RegisterDeviceIntrinsic("cuda_mbarrier_wait_acquire_cluster", "cuda", CallEffectKind::kOpaque, - sig::arg("bar"), sig::arg("phase")); - RegisterDeviceIntrinsic("cuda_mov_sreg", "cuda", CallEffectKind::kPure, sig::arg("bits"), - sig::arg("reg_name")); - RegisterDeviceIntrinsic("cuda_nano_sleep", "cuda", CallEffectKind::kOpaque, - sig::arg("time")); - RegisterDeviceIntrinsic("cuda_printf", "cuda", CallEffectKind::kOpaque, sig::arg("fmt"), - sig::var_args("args")); - RegisterDeviceIntrinsic("cuda_reduce_add_sync_u32", "cuda", CallEffectKind::kOpaque, - sig::arg("mask"), sig::arg("value")); - RegisterDeviceIntrinsic("cuda_reduce_min_sync_u32", "cuda", CallEffectKind::kOpaque, - sig::arg("mask"), sig::arg("value")); - RegisterDeviceIntrinsic("cuda_runtime_instr_desc", "cuda", CallEffectKind::kOpaque, - sig::arg("desc"), sig::arg("sf_id")); - RegisterDeviceIntrinsic("cuda_sm100_2sm_leader_smem_addr", "cuda", CallEffectKind::kOpaque, - sig::arg("ptr")); - RegisterDeviceIntrinsic("cuda_smem_addr_from_uint64", "cuda", CallEffectKind::kOpaque, - sig::arg("cluster_addr")); - RegisterDeviceIntrinsic("cuda_syncthreads_and", "cuda", CallEffectKind::kOpaque, - sig::arg("cond")); - RegisterDeviceIntrinsic("cuda_syncthreads_or", "cuda", CallEffectKind::kOpaque, sig::arg("cond")); - RegisterDeviceIntrinsic("cuda_tcgen05_encode_instr_descriptor", "cuda", CallEffectKind::kOpaque, - sig::arg("desc"), sig::arg("d_dtype"), sig::arg("a_dtype"), - sig::arg("b_dtype"), sig::arg("M"), sig::arg("N"), - sig::arg("K"), sig::arg("trans_a"), sig::arg("trans_b"), - sig::arg("n_cta_groups"), sig::arg("neg_a"), sig::arg("neg_b"), - sig::arg("sat_d"), sig::arg("is_sparse")); - RegisterDeviceIntrinsic("cuda_tcgen05_encode_instr_descriptor_block_scaled", "cuda", +void RegisterDeviceIntrinsics() { + RegisterDeviceIntrinsic(OpDef("tirx.cuda.any_sync"), "cuda", CallEffectKind::kPure, + sig::arg("mask"), sig::arg("pred")) + .set_attr("TFixedReturnType", PrimType::Int(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.atomic_add"), "cuda", CallEffectKind::kOpaque, + sig::arg("res_addr"), sig::arg("value")) + .set_attr("FInferType", FInferType::FromNative<&InferTypeReturnArgType<1>>()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.atomic_cas"), "cuda", CallEffectKind::kOpaque, + sig::arg("ptr"), sig::arg("old_val"), sig::arg("new_val")) + .set_attr("FInferType", FInferType::FromNative<&InferTypeReturnArgType<1>>()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.wait_until"), "cuda", CallEffectKind::kOpaque, + sig::arg("dst"), sig::arg("ptr"), sig::arg("condition"), + sig::arg("scope"), sig::arg("space"), sig::arg("ptx_type"), + sig::arg("backoff_ns")); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.ballot_sync"), "cuda", CallEffectKind::kOpaque, + sig::arg("mask"), sig::arg("pred")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.bfloat1622float2"), "cuda", CallEffectKind::kOpaque, + sig::arg("packed")) + .set_attr("TFixedReturnType", PrimType::UInt(64)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.bfloat162float"), "cuda", CallEffectKind::kOpaque, + sig::arg("src")) + .set_attr("TFixedReturnType", PrimType::Float(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.clock64"), "cuda", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::UInt(64)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.cluster_sync"), "cuda", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.cta_reduce"), "cuda", CallEffectKind::kOpaque, + sig::arg("value"), sig::arg("op"), sig::arg("num_warps"), + sig::arg("scratch")); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.cta_sync"), "cuda", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.cvta_generic_to_shared"), "cuda", + CallEffectKind::kOpaque, sig::arg("ptr")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.elect_sync"), "cuda", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.fadd2_rn"), "cuda", CallEffectKind::kOpaque, + sig::arg("a"), sig::arg("b")) + .set_attr("TFixedReturnType", PrimType::UInt(64)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.fdividef"), "cuda", CallEffectKind::kPure, sig::arg("x"), + sig::arg("y")) + .set_attr("TFixedReturnType", PrimType::Float(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.ffs_u32"), "cuda", CallEffectKind::kOpaque, + sig::arg("value")) + .set_attr("TFixedReturnType", PrimType::Int(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.float22bfloat162_rn"), "cuda", CallEffectKind::kOpaque, + sig::arg("v0"), sig::arg("v1")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.float22bfloat162_rn_from_float2"), "cuda", + CallEffectKind::kOpaque, sig::arg("packed")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.float22half2"), "cuda", CallEffectKind::kOpaque, + sig::arg("dst"), sig::arg("src")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.float2_x"), "cuda", CallEffectKind::kOpaque, + sig::arg("packed")) + .set_attr("TFixedReturnType", PrimType::Float(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.float2_y"), "cuda", CallEffectKind::kOpaque, + sig::arg("packed")) + .set_attr("TFixedReturnType", PrimType::Float(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.float8tohalf8"), "cuda", CallEffectKind::kOpaque, + sig::arg("src_addr"), sig::arg("dst_addr")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.float_as_uint"), "cuda", CallEffectKind::kOpaque, + sig::arg("x")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.fmul2_rn"), "cuda", CallEffectKind::kOpaque, + sig::arg("a"), sig::arg("b")) + .set_attr("TFixedReturnType", PrimType::UInt(64)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.fp8x4_e4m3_from_float4"), "cuda", + CallEffectKind::kOpaque, sig::arg("x"), sig::arg("y"), sig::arg("z"), + sig::arg("w")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.func_call"), "cuda", CallEffectKind::kOpaque, + sig::arg("func_name"), sig::var_args("args")); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.get_tmem_addr"), "cuda", CallEffectKind::kOpaque, + sig::arg("addr"), sig::arg("row_offset"), + sig::arg("col_offset")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.grid_sync"), "cuda", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.half2float"), "cuda", CallEffectKind::kOpaque, + sig::arg("src")) + .set_attr("TFixedReturnType", PrimType::Float(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.half8tofloat8"), "cuda", CallEffectKind::kOpaque, + sig::arg("src_addr"), sig::arg("dst_addr")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.hmax2"), "cuda", CallEffectKind::kOpaque, sig::arg("a"), + sig::arg("b")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.hmin2"), "cuda", CallEffectKind::kOpaque, sig::arg("a"), + sig::arg("b")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.ldg"), "cuda", CallEffectKind::kOpaque, + sig::var_args("args")) + .set_attr("FInferType", FInferType::FromNative<&InferTypeCudaLdg>()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.make_float2"), "cuda", CallEffectKind::kOpaque, + sig::arg("x"), sig::arg("y")) + .set_attr("TFixedReturnType", PrimType::UInt(64)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.mbarrier_wait"), "cuda", CallEffectKind::kOpaque, + sig::arg("bar"), sig::arg("phase")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.mbarrier_wait_acquire_cluster"), "cuda", + CallEffectKind::kOpaque, sig::arg("bar"), sig::arg("phase")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.mov_sreg"), "cuda", CallEffectKind::kPure, + sig::arg("bits"), sig::arg("reg_name")); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.nano_sleep"), "cuda", CallEffectKind::kOpaque, + sig::arg("time")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.printf"), "cuda", CallEffectKind::kOpaque, + sig::arg("fmt"), sig::var_args("args")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.reduce_add_sync_u32"), "cuda", CallEffectKind::kOpaque, + sig::arg("mask"), sig::arg("value")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.reduce_min_sync_u32"), "cuda", CallEffectKind::kOpaque, + sig::arg("mask"), sig::arg("value")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.runtime_instr_desc"), "cuda", CallEffectKind::kOpaque, + sig::arg("desc"), sig::arg("sf_id")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.sm100_2sm_leader_smem_addr"), "cuda", + CallEffectKind::kOpaque, sig::arg("ptr")); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.smem_addr_from_uint64"), "cuda", CallEffectKind::kOpaque, + sig::arg("cluster_addr")) + .set_attr("TFixedReturnType", PrimType::UInt(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.syncthreads_and"), "cuda", CallEffectKind::kOpaque, + sig::arg("cond")) + .set_attr("TFixedReturnType", PrimType::Int(64)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.syncthreads_or"), "cuda", CallEffectKind::kOpaque, + sig::arg("cond")) + .set_attr("TFixedReturnType", PrimType::Int(64)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.tcgen05_encode_instr_descriptor"), "cuda", + CallEffectKind::kOpaque, sig::arg("desc"), sig::arg("d_dtype"), + sig::arg("a_dtype"), sig::arg("b_dtype"), sig::arg("M"), + sig::arg("N"), sig::arg("K"), sig::arg("trans_a"), + sig::arg("trans_b"), sig::arg("n_cta_groups"), sig::arg("neg_a"), + sig::arg("neg_b"), sig::arg("sat_d"), sig::arg("is_sparse")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.tcgen05_encode_instr_descriptor_block_scaled"), "cuda", CallEffectKind::kOpaque, sig::arg("desc"), sig::arg("d_dtype"), sig::arg("a_dtype"), sig::arg("b_dtype"), sig::arg("sfa_dtype"), sig::arg("sfb_dtype"), sig::arg("sfa_tmem_addr"), sig::arg("sfb_tmem_addr"), sig::arg("M"), sig::arg("N"), sig::arg("K"), sig::arg("trans_a"), sig::arg("trans_b"), sig::arg("n_cta_groups"), sig::arg("neg_a"), sig::arg("neg_b"), - sig::arg("is_sparse")); - RegisterDeviceIntrinsic("cuda_tcgen05_encode_matrix_descriptor", "cuda", CallEffectKind::kOpaque, - sig::arg("desc"), sig::arg("addr"), sig::arg("ldo"), - sig::arg("sdo"), sig::arg("swizzle")); - RegisterDeviceIntrinsic("cuda_thread_fence", "cuda", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("cuda_thread_rank", "cuda", CallEffectKind::kPure); - RegisterDeviceIntrinsic("cuda_trap_when_assert_failed", "cuda", CallEffectKind::kOpaque, - sig::arg("cond")); - RegisterDeviceIntrinsic("cuda_uint_as_float", "cuda", CallEffectKind::kOpaque, - sig::arg("bits")); - RegisterDeviceIntrinsic("cuda_warp_reduce", "cuda", CallEffectKind::kOpaque, sig::arg("value"), - sig::arg("op"), sig::arg("width")); - RegisterDeviceIntrinsic("cuda_warp_sync", "cuda", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("cuda_warpgroup_sync", "cuda", CallEffectKind::kOpaque, - sig::arg("bar_no")); - RegisterDeviceIntrinsic("cuda_wgmma_encode_matrix_descriptor", "cuda", CallEffectKind::kOpaque, - sig::arg("desc"), sig::arg("addr"), sig::arg("ldo"), - sig::arg("sdo"), sig::arg("swizzle")); - RegisterDeviceIntrinsic("cuda_wgmma_noop_barrier", "cuda", CallEffectKind::kOpaque, - sig::arg("reg")); - RegisterDeviceIntrinsic("nvshmem_barrier_all", "nvshmem", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("nvshmem_fence", "nvshmem", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("nvshmem_getmem_nbi", "nvshmem", CallEffectKind::kOpaque, sig::arg("dst"), - sig::arg("src"), sig::arg("nelems"), sig::arg("pe")); - RegisterDeviceIntrinsic("nvshmem_getmem_nbi_block", "nvshmem", CallEffectKind::kOpaque, - sig::arg("dst"), sig::arg("src"), sig::arg("nelems"), - sig::arg("pe")); - RegisterDeviceIntrinsic("nvshmem_getmem_nbi_warp", "nvshmem", CallEffectKind::kOpaque, - sig::arg("dst"), sig::arg("src"), sig::arg("nelems"), - sig::arg("pe")); - RegisterDeviceIntrinsic("nvshmem_my_pe", "nvshmem", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("nvshmem_n_pes", "nvshmem", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("nvshmem_putmem_nbi", "nvshmem", CallEffectKind::kOpaque, sig::arg("dst"), - sig::arg("src"), sig::arg("nelems"), sig::arg("pe")); - RegisterDeviceIntrinsic("nvshmem_putmem_nbi_block", "nvshmem", CallEffectKind::kOpaque, - sig::arg("dst"), sig::arg("src"), sig::arg("nelems"), - sig::arg("pe")); - RegisterDeviceIntrinsic("nvshmem_putmem_nbi_warp", "nvshmem", CallEffectKind::kOpaque, + sig::arg("is_sparse")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.tcgen05_encode_matrix_descriptor"), "cuda", + CallEffectKind::kOpaque, sig::arg("desc"), sig::arg("addr"), + sig::arg("ldo"), sig::arg("sdo"), + sig::arg("swizzle")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.thread_fence"), "cuda", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.thread_rank"), "cuda", CallEffectKind::kPure) + .set_attr("TFixedReturnType", PrimType::Int(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.trap_when_assert_failed"), "cuda", + CallEffectKind::kOpaque, sig::arg("cond")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.uint_as_float"), "cuda", CallEffectKind::kOpaque, + sig::arg("bits")) + .set_attr("TFixedReturnType", PrimType::Float(32)); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.warp_reduce"), "cuda", CallEffectKind::kOpaque, + sig::arg("value"), sig::arg("op"), sig::arg("width")); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.warp_sync"), "cuda", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.warpgroup_sync"), "cuda", CallEffectKind::kOpaque, + sig::arg("bar_no")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.wgmma_encode_matrix_descriptor"), "cuda", + CallEffectKind::kOpaque, sig::arg("desc"), sig::arg("addr"), + sig::arg("ldo"), sig::arg("sdo"), + sig::arg("swizzle")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.cuda.wgmma_noop_barrier"), "cuda", CallEffectKind::kOpaque, + sig::arg("reg")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.barrier_all"), "nvshmem", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.fence"), "nvshmem", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.getmem_nbi"), "nvshmem", CallEffectKind::kOpaque, sig::arg("dst"), sig::arg("src"), sig::arg("nelems"), - sig::arg("pe")); - RegisterDeviceIntrinsic("nvshmem_putmem_signal_nbi", "nvshmem", CallEffectKind::kOpaque, + sig::arg("pe")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.getmem_nbi_block"), "nvshmem", + CallEffectKind::kOpaque, sig::arg("dst"), sig::arg("src"), + sig::arg("nelems"), sig::arg("pe")); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.getmem_nbi_warp"), "nvshmem", CallEffectKind::kOpaque, sig::arg("dst"), sig::arg("src"), sig::arg("nelems"), - sig::arg("sig_addr"), sig::arg("signal"), sig::arg("sig_op"), - sig::arg("pe")); - RegisterDeviceIntrinsic("nvshmem_putmem_signal_nbi_block", "nvshmem", CallEffectKind::kOpaque, + sig::arg("pe")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.my_pe"), "nvshmem", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::Int(32)); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.n_pes"), "nvshmem", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::Int(32)); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.putmem_nbi"), "nvshmem", CallEffectKind::kOpaque, sig::arg("dst"), sig::arg("src"), sig::arg("nelems"), - sig::arg("sig_addr"), sig::arg("signal"), sig::arg("sig_op"), - sig::arg("pe")); - RegisterDeviceIntrinsic("nvshmem_putmem_signal_nbi_warp", "nvshmem", CallEffectKind::kOpaque, + sig::arg("pe")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.putmem_nbi_block"), "nvshmem", + CallEffectKind::kOpaque, sig::arg("dst"), sig::arg("src"), + sig::arg("nelems"), sig::arg("pe")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.putmem_nbi_warp"), "nvshmem", CallEffectKind::kOpaque, sig::arg("dst"), sig::arg("src"), sig::arg("nelems"), + sig::arg("pe")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.putmem_signal_nbi"), "nvshmem", + CallEffectKind::kOpaque, sig::arg("dst"), sig::arg("src"), + sig::arg("nelems"), sig::arg("sig_addr"), + sig::arg("signal"), sig::arg("sig_op"), sig::arg("pe")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.putmem_signal_nbi_block"), "nvshmem", + CallEffectKind::kOpaque, sig::arg("dst"), sig::arg("src"), + sig::arg("nelems"), sig::arg("sig_addr"), + sig::arg("signal"), sig::arg("sig_op"), sig::arg("pe")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.putmem_signal_nbi_warp"), "nvshmem", + CallEffectKind::kOpaque, sig::arg("dst"), sig::arg("src"), + sig::arg("nelems"), sig::arg("sig_addr"), + sig::arg("signal"), sig::arg("sig_op"), sig::arg("pe")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.quiet"), "nvshmem", CallEffectKind::kOpaque) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.signal_op"), "nvshmem", CallEffectKind::kOpaque, sig::arg("sig_addr"), sig::arg("signal"), sig::arg("sig_op"), - sig::arg("pe")); - RegisterDeviceIntrinsic("nvshmem_quiet", "nvshmem", CallEffectKind::kOpaque); - RegisterDeviceIntrinsic("nvshmem_signal_op", "nvshmem", CallEffectKind::kOpaque, - sig::arg("sig_addr"), sig::arg("signal"), sig::arg("sig_op"), - sig::arg("pe")); - RegisterDeviceIntrinsic("nvshmem_wait_until", "nvshmem", CallEffectKind::kOpaque, + sig::arg("pe")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.nvshmem.wait_until"), "nvshmem", CallEffectKind::kOpaque, sig::arg("ivar"), sig::arg("cmp"), sig::arg("cmp_value"), - sig::arg("type")); - RegisterDeviceIntrinsic("ptx_legacy_ldmatrix", "ptx_legacy", CallEffectKind::kOpaque, + sig::arg("type")) + .set_attr("TFixedReturnType", PrimType::Void()); + RegisterDeviceIntrinsic(OpDef("tirx.ptx_legacy.ldmatrix"), "ptx_legacy", CallEffectKind::kOpaque, sig::arg("trans"), sig::arg("num"), sig::arg("dtype"), sig::arg("local_ptr"), sig::arg("local_offset"), sig::arg("smem_ptr"), sig::arg("smem_offset")); RegisterDeviceIntrinsic( - "ptx_legacy_mma", "ptx_legacy", CallEffectKind::kOpaque, sig::arg("shape"), + OpDef("tirx.ptx_legacy.mma"), "ptx_legacy", CallEffectKind::kOpaque, sig::arg("shape"), sig::arg("a_layout"), sig::arg("b_layout"), sig::arg("a_dtype"), sig::arg("b_dtype"), sig::arg("c_dtype"), sig::arg("a_ptr"), sig::arg("a_offset"), sig::arg("b_ptr"), sig::arg("b_offset"), sig::arg("acc_ptr"), sig::arg("c_offset"), sig::arg("saturate"), sig::var_args("args")); - - for (const char* name : {"tirx.cuda.thread_rank", "tirx.cuda.any_sync", "tirx.cuda.ffs_u32"}) { - OpDef(name).set_attr("TFixedReturnType", PrimType::Int(32)); - } - for (const char* name : {"tirx.cuda.half2float", "tirx.cuda.bfloat162float", "tirx.cuda.fdividef", - "tirx.cuda.uint_as_float", "tirx.cuda.float2_x", "tirx.cuda.float2_y"}) { - OpDef(name).set_attr("TFixedReturnType", PrimType::Float(32)); - } - for (const char* name : - {"tirx.cuda.float22half2", "tirx.cuda.trap_when_assert_failed", - "tirx.cuda.runtime_instr_desc", "tirx.cuda.half8tofloat8", "tirx.cuda.float8tohalf8", - "tirx.cuda.mbarrier_wait_acquire_cluster", "tirx.cuda.warpgroup_sync"}) { - OpDef(name).set_attr("TFixedReturnType", PrimType::Void()); - } - for (const char* name : - {"tirx.cuda.get_tmem_addr", "tirx.cuda.cvta_generic_to_shared", - "tirx.cuda.smem_addr_from_uint64", "tirx.cuda.float_as_uint", "tirx.cuda.ballot_sync", - "tirx.cuda.reduce_add_sync_u32", "tirx.cuda.reduce_min_sync_u32", - "tirx.cuda.float22bfloat162_rn", "tirx.cuda.float22bfloat162_rn_from_float2", - "tirx.cuda.hmin2", "tirx.cuda.hmax2", "tirx.cuda.fp8x4_e4m3_from_float4"}) { - OpDef(name).set_attr("TFixedReturnType", PrimType::UInt(32)); - } - for (const char* name : {"tirx.cuda.clock64", "tirx.cuda.make_float2", "tirx.cuda.fmul2_rn", - "tirx.cuda.fadd2_rn", "tirx.cuda.bfloat1622float2"}) { - OpDef(name).set_attr("TFixedReturnType", PrimType::UInt(64)); - } - for (const char* name : { - "tirx.cuda.tcgen05_encode_matrix_descriptor", - "tirx.cuda.wgmma_encode_matrix_descriptor", - "tirx.cuda.tcgen05_encode_instr_descriptor", - "tirx.cuda.tcgen05_encode_instr_descriptor_block_scaled", - "tirx.cuda.wgmma_noop_barrier", - "tirx.cuda.cluster_sync", - "tirx.cuda.cta_sync", - "tirx.cuda.grid_sync", - "tirx.cuda.warp_sync", - "tirx.cuda.thread_fence", - "tirx.cuda.mbarrier_wait", - "tirx.cuda.printf", - "tirx.cuda.nano_sleep", - "tirx.nvshmem.fence", - "tirx.nvshmem.quiet", - "tirx.nvshmem.barrier_all", - "tirx.nvshmem.getmem_nbi", - "tirx.nvshmem.getmem_nbi_warp", - "tirx.nvshmem.putmem_nbi", - "tirx.nvshmem.putmem_nbi_warp", - "tirx.nvshmem.putmem_nbi_block", - "tirx.nvshmem.signal_op", - "tirx.nvshmem.wait_until", - "tirx.nvshmem.putmem_signal_nbi", - "tirx.nvshmem.putmem_signal_nbi_warp", - "tirx.nvshmem.putmem_signal_nbi_block", - }) { - OpDef(name).set_attr("TFixedReturnType", PrimType::Void()); - } - for (const char* name : {"tirx.cuda.syncthreads_and", "tirx.cuda.syncthreads_or"}) { - OpDef(name).set_attr("TFixedReturnType", PrimType::Int(64)); - } - for (const char* name : {"tirx.nvshmem.my_pe", "tirx.nvshmem.n_pes"}) { - OpDef(name).set_attr("TFixedReturnType", PrimType::Int(32)); - } - OpDef("tirx.cuda.elect_sync").set_attr("TFixedReturnType", PrimType::UInt(32)); - OpDef("tirx.cuda.ldg") - .set_attr("FInferType", FInferType::FromNative<&InferTypeCudaLdg>()); - OpDef("tirx.cuda.atomic_add") - .set_attr("FInferType", FInferType::FromNative<&InferTypeReturnArgType<1>>()); - OpDef("tirx.cuda.atomic_cas") - .set_attr("FInferType", FInferType::FromNative<&InferTypeReturnArgType<1>>()); } } // namespace From b793a05c92f6f966e4418e682953df5fcb0e0e90 Mon Sep 17 00:00:00 2001 From: Tianqi Chen Date: Sun, 4 Oct 2026 01:13:45 +0000 Subject: [PATCH 4/4] [FIX][CUDA] Preserve composed leader-address construction Keep the documented leader shared-address helper as the existing conversion and mask expression, removing its unused native Op placeholder. Align IKET printer expectations with canonical registered names. --- python/tvm/backend/cuda/script.py | 2 +- src/backend/cuda/op/target_builtin.cc | 2 -- tests/python/tirx/iket/test_iket_profiler.py | 8 ++++---- 3 files changed, 5 insertions(+), 7 deletions(-) diff --git a/python/tvm/backend/cuda/script.py b/python/tvm/backend/cuda/script.py index 39e3f1da9f84..bfe448c663c5 100644 --- a/python/tvm/backend/cuda/script.py +++ b/python/tvm/backend/cuda/script.py @@ -166,7 +166,7 @@ def _activemask(): atomic_cas = _cuda_op.cuda_atomic_cas func_call = _cuda_op.cuda_func_call ldg = _cuda_op.cuda_ldg -sm100_2sm_leader_smem_addr_composed = _cuda_op.cuda_sm100_2sm_leader_smem_addr +sm100_2sm_leader_smem_addr = _cuda_op.cuda_sm100_2sm_leader_smem_addr timer_init = _cuda_op.timer_init_cuda timer_start = _cuda_op.timer_start_cuda timer_end = _cuda_op.timer_end_cuda diff --git a/src/backend/cuda/op/target_builtin.cc b/src/backend/cuda/op/target_builtin.cc index 144f04bd003a..7eaa19106e64 100644 --- a/src/backend/cuda/op/target_builtin.cc +++ b/src/backend/cuda/op/target_builtin.cc @@ -373,8 +373,6 @@ void RegisterDeviceIntrinsics() { RegisterDeviceIntrinsic(OpDef("tirx.cuda.runtime_instr_desc"), "cuda", CallEffectKind::kOpaque, sig::arg("desc"), sig::arg("sf_id")) .set_attr("TFixedReturnType", PrimType::Void()); - RegisterDeviceIntrinsic(OpDef("tirx.cuda.sm100_2sm_leader_smem_addr"), "cuda", - CallEffectKind::kOpaque, sig::arg("ptr")); RegisterDeviceIntrinsic(OpDef("tirx.cuda.smem_addr_from_uint64"), "cuda", CallEffectKind::kOpaque, sig::arg("cluster_addr")) .set_attr("TFixedReturnType", PrimType::UInt(32)); diff --git a/tests/python/tirx/iket/test_iket_profiler.py b/tests/python/tirx/iket/test_iket_profiler.py index ce7b8c9b4465..87067d982483 100644 --- a/tests/python/tirx/iket/test_iket_profiler.py +++ b/tests/python/tirx/iket/test_iket_profiler.py @@ -371,7 +371,7 @@ def test_public_interface_is_official_only(): assert callable(cuda_transforms.LowerIket) script = serial_a.script() - assert 'T.cuda.iket.mark("a")' in script + assert 'T.cuda.iket_mark("a")' in script assert "T.tirx.iket" not in script assert ( tvm.script.from_source( @@ -381,9 +381,9 @@ def test_public_interface_is_official_only(): ) payload_script = payload_types.script() - assert 'T.cuda.iket.mark("i8", T.int8(-8))' in payload_script - assert 'T.cuda.iket.range_start("token_payload", -7)' in payload_script - assert "T.cuda.iket.range_end(token, 9)" in payload_script + assert 'T.cuda.iket_mark("i8", T.int8(-8))' in payload_script + assert 'T.cuda.iket_range_start("token_payload", -7)' in payload_script + assert "T.cuda.iket_range_end(token, 9)" in payload_script assert ( tvm.script.from_source( payload_script, extra_vars={"I": tvm.script.ir, "T": tvm.script.tirx}