diff --git a/python/tvm/backend/cuda/__init__.py b/python/tvm/backend/cuda/__init__.py index 627ff1a0ffbd..df0b40dc799c 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, canonical_op_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 048805fb2324..bfe448c663c5 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 = _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..927b8edac716 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,92 @@ 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, 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. Script namespaces separately publish + their callable names to the printer after initialization. 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 7b145e0590d7..3db9cec13796 100644 --- a/python/tvm/tirx/script/ir_builder/op.py +++ b/python/tvm/tirx/script/ir_builder/op.py @@ -287,7 +287,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, *, canonical_op_names: bool = False +) -> object: """Register a construction namespace and return it. Parameters @@ -299,6 +301,10 @@ 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. + canonical_op_names : bool, optional + Publish registered Op names whose canonical attributes expose matching + callables. Preserve explicit printer aliases. Otherwise discover names + from legacy wrappers. Repeat registration to publish newly exposed Ops. """ _SCRIPT_NAMESPACES[name] = namespace globals()[name] = namespace @@ -320,7 +326,25 @@ 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 canonical_op_names: + prefix = f"tirx.{name}." + for op_name in _ir.Op.list_op_names(): + if not op_name.startswith(prefix): + continue + current = namespace + for part in op_name[len(prefix) :].split("."): + current = getattr(current, part, None) + op = _ir.Op.get(op_name) + identity = getattr(current, "__tvm_op__", None) + if ( + callable(current) + and isinstance(identity, _ir.Op) + and identity.same_as(op) + and op.get_attr("TScriptPrinterName") is None + ): + op.set_attr("TScriptPrinterName", op_name) + else: + _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..7eaa19106e64 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,266 +219,281 @@ 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; - // 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))) { - 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)) - .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 -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.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("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_getmem_nbi_warp", "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("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("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("pe")); - RegisterDeviceIntrinsic("nvshmem_putmem_nbi_warp", "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("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.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_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")); - 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.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 diff --git a/src/tirx/script/printer/expr.cc b/src/tirx/script/printer/expr.cc index 9baf881187ea..405729ea0211 100644 --- a/src/tirx/script/printer/expr.cc +++ b/src/tirx/script/printer/expr.cc @@ -510,13 +510,55 @@ 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 (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; + } + } + // Published CUDA callables use canonical names. Late registrations without + // a published callable retain lossless reconstruction. + 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); + } + 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); + } + 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); + } } if (names.count(op.value()) && !incompatible_signature) { std::string name = names[op.value()]; @@ -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/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} diff --git a/tests/python/tirx/script/test_tirx_script_printer.py b/tests/python/tirx/script/test_tirx_script_printer.py index e352d4368ee1..fe22590a9ac3 100644 --- a/tests/python/tirx/script/test_tirx_script_printer.py +++ b/tests/python/tirx/script/test_tirx_script_printer.py @@ -749,7 +749,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( @@ -1026,6 +1026,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 ba6d8d6f60a4..0590f0b191e7 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.Tensor((16,), "float32"), B: T.Tensor((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)