diff --git a/python/triton/experimental/tle/language/distributed.py b/python/triton/experimental/tle/language/distributed.py index d9ae4db190..eb5c0763f5 100644 --- a/python/triton/experimental/tle/language/distributed.py +++ b/python/triton/experimental/tle/language/distributed.py @@ -766,11 +766,21 @@ def distributed_barrier(mesh: device_mesh | None = None, device_dptr=None, space return None +def _unwrap_remote_shard_id(shard_id: Any): + shard_id = tl._unwrap_if_constexpr(shard_id) + # Tuple literals in JIT functions are represented as tl.tuple even when + # every coordinate is compile-time constant. Convert them back to a Python + # tuple so the shared compile-time coordinate path can process them. + if isinstance(shard_id, tl.tuple): + shard_id = tuple(shard_id) + return shard_id + + def _normalize_remote_shard_id( shard_id: Any, scope: device_mesh | None, ) -> int: - shard_id = tl._unwrap_if_constexpr(shard_id) + shard_id = _unwrap_remote_shard_id(shard_id) scope = tl._unwrap_if_constexpr(scope) if isinstance(shard_id, int): @@ -871,6 +881,8 @@ def _create_remote_pointers_tensor( # automatic injection (e.g. inside this helper). if offset_tensor.dtype != tl.int64: offset_tensor = tl.cast(offset_tensor, tl.int64, _semantic=_semantic) + # create_dist_tensor exposes registered memory at index 0 and the + # communicator at index 1. ptr = _parse_src_arg(builder, tensor, 0) remote_op = builder.create_remote_pointers(remote_type, ptr, shard_id_tensor.handle, space, offset_tensor.handle) @@ -907,11 +919,6 @@ def _check_device_remote_pointer(tensor: tl.tensor, shard_id: int | tuple[int, . ... -def _check_node_remote_pointer(tensor: tl.tensor, shard_id: int | tuple[int, ...] | list[int], - scope: device_mesh | None) -> None: - ... - - def _remote_pointer( tensor: tl.tensor, shard_id, @@ -922,14 +929,13 @@ def _remote_pointer( _semantic=None, ) -> tl.tensor: - if not isinstance(tensor, tl.tensor) and space not in ("device", "node"): + if not isinstance(tensor, tl.tensor) and space != "device": raise TypeError(f"tensor must be tl.tensor, got {type(tensor).__name__}") space = tl._unwrap_if_constexpr(space) res = { "cluster": _check_cluster_remote_pointer, "device": _check_device_remote_pointer, - "node": _check_node_remote_pointer, }[space](tensor, shard_id, scope) if isinstance(res, tl.tensor): return res @@ -949,6 +955,325 @@ def _remote_pointer( return _create_remote_pointers_tensor(tensor, shard_id_tensor, _semantic, dtype=dtype, space=space, offset=offset) +# offset / dstoffset / nelems -> scalar i64 tl.tensor +# offset and dstoffset must be >= 0. +# nelems must be > 0. +def _normalize_node_i64(value, label: str, *, must_be_positive: bool, _semantic) -> tl.tensor: + value = tl._unwrap_if_constexpr(value) + if isinstance(value, int): + if must_be_positive and value <= 0: + raise ValueError(f"node space {label} must be > 0, got {value}") + if not must_be_positive and value < 0: + raise ValueError(f"node space {label} must be >= 0, got {value}") + + value_tensor = value if isinstance(value, tl.tensor) else _semantic.to_tensor(value) + if not value_tensor.dtype.is_int(): + raise TypeError(f"node space {label} must be an integer scalar, got {value_tensor.dtype}") + if value_tensor.shape != (): + raise ValueError(f"node space {label} must be scalar, got shape {value_tensor.shape}") + if value_tensor.dtype != tl.int64: + value_tensor = tl.cast(value_tensor, tl.int64, _semantic=_semantic) + return value_tensor + + +def _normalize_node_elem_bytes(dtype) -> int: + dtype = tl._unwrap_if_constexpr(dtype) + if not isinstance(dtype, tl.dtype): + raise TypeError(f"node space dtype must be a scalar Triton dtype, got {type(dtype).__name__}") + elem_bytes = dtype.itemsize + if elem_bytes <= 0: + raise ValueError(f"node space dtype must be byte-addressable, got {dtype}") + return elem_bytes + + +def _normalize_node_peer(shard_id, scope, _semantic) -> tl.tensor: + shard_id = _unwrap_remote_shard_id(shard_id) + scope = tl._unwrap_if_constexpr(scope) + if scope is not None and not isinstance(scope, device_mesh): + raise TypeError(f"node space scope must be device_mesh or None, got {type(scope).__name__}") + + if isinstance(shard_id, (int, tuple, list)): + is_coordinate = isinstance(shard_id, (tuple, list)) + peer = _normalize_compile_time_remote_shard_id(shard_id, scope) + if is_coordinate: + # Coordinates are relative to the selected mesh. Resolve through + # physical_ids so coordinates on a sliced submesh still produce + # the corresponding world rank rather than a submesh-local rank. + peer = scope.physical_ids[peer] + if peer > 0x7FFFFFFF: + raise ValueError(f"node space world rank {peer} exceeds int32 range") + shard_id = _semantic.to_tensor(peer) + elif not isinstance(shard_id, tl.tensor): + shard_id = _semantic.to_tensor(shard_id) + return _normalize_runtime_remote_shard_id_tensor(shard_id) + + +def _normalize_coopkind(coopkind) -> int: + coopkind = tl._unwrap_if_constexpr(coopkind) + if isinstance(coopkind, GroupKind): + coopkind = coopkind.value + if not isinstance(coopkind, str): + raise TypeError("node space coopkind must be GroupKind.THREAD/WARP/BLOCK or the corresponding string") + mapping = {"thread": 0, "warp": 1, "block": 2} + normalized = coopkind.lower() + if normalized not in mapping: + raise ValueError("node space coopkind must be THREAD, WARP, or BLOCK") + return mapping[normalized] + + +def _normalize_node_netidx(netidx, _semantic) -> tl.tensor: + netidx = tl._unwrap_if_constexpr(netidx) + if isinstance(netidx, int): + if netidx < 0 or netidx > 0x7FFFFFFF: + raise ValueError(f"node space netidx must be in int32 range [0, 2147483647], got {netidx}") + netidx = _semantic.to_tensor(netidx) + elif not isinstance(netidx, tl.tensor): + netidx = _semantic.to_tensor(netidx) + + if not netidx.dtype.is_int(): + raise TypeError(f"node space netidx must be an integer scalar, got {netidx.dtype}") + if netidx.shape != (): + raise ValueError(f"node space netidx must be scalar, got shape {netidx.shape}") + if netidx.dtype != tl.int32: + netidx = tl.cast(netidx, tl.int32, _semantic=_semantic) + return netidx + + +def _parse_node_context(builder, value, label: str, index: int): + from triton.runtime import DistributedRtContext + value = tl._unwrap_if_constexpr(value) + if not isinstance(value, DistributedRtContext): + raise TypeError(f"node space {label} must be DistributedRtContext, got {type(value).__name__}") + return _parse_src_arg(builder, value, index) + + +class _node_remote_destination_type(tl.base_type): + + def __init__(self, field_types, dtype: tl.dtype, elem_bytes: int, coop_kind: int): + # src_mem, dst_mem, comm, peer, offset, dstoffset, nelems, net_idx. + self.field_types = tuple(field_types) + self.dtype = dtype + self.elem_bytes = elem_bytes + self.coop_kind = coop_kind + + def _unflatten_ir(self, handles, cursor): + fields = [] + for field_type in self.field_types: + field, cursor = field_type._unflatten_ir(handles, cursor) + fields.append(field) + return _node_remote_destination( + *fields, + dtype=self.dtype, + elem_bytes=self.elem_bytes, + coop_kind=self.coop_kind, + ), cursor + + def _flatten_ir_types(self, builder, out) -> None: + for field_type in self.field_types: + field_type._flatten_ir_types(builder, out) + + def mangle(self) -> str: + fields = "_".join(field_type.mangle() for field_type in self.field_types) + return f"node_remote_dst_{self.dtype.mangle()}_e{self.elem_bytes}_c{self.coop_kind}_{fields}" + + def __eq__(self, other) -> bool: + return (type(self) is type(other) and self.field_types == other.field_types and self.dtype == other.dtype + and self.elem_bytes == other.elem_bytes and self.coop_kind == other.coop_kind) + + def __str__(self) -> str: + return f"node_remote_destination<{self.dtype}, coop_kind={self.coop_kind}>" + + @property + def scalar(self): + raise ValueError('tle.remote(..., space="node") destinations only support tl.load/tl.store') + + +class _node_remote_destination(tl.base_value): + + def __init__(self, src_mem: tl.tensor, dst_mem: tl.tensor, comm: tl.tensor, peer: tl.tensor, offset: tl.tensor, + dstoffset: tl.tensor, nelems: tl.tensor, net_idx: tl.tensor, *, dtype: tl.dtype, elem_bytes: int, + coop_kind: int): + super().__init__() + self.src_mem = src_mem + self.dst_mem = dst_mem + self.comm = comm + self.peer = peer + self.offset = offset + self.dstoffset = dstoffset + self.nelems = nelems + self.net_idx = net_idx + self.dtype = dtype + self.elem_bytes = elem_bytes + self.coop_kind = coop_kind + + @property + def type(self): + fields = (self.src_mem, self.dst_mem, self.comm, self.peer, self.offset, self.dstoffset, self.nelems, + self.net_idx) + return _node_remote_destination_type( + tuple(field.type for field in fields), + self.dtype, + self.elem_bytes, + self.coop_kind, + ) + + def _flatten_ir(self, handles) -> None: + for field in (self.src_mem, self.dst_mem, self.comm, self.peer, self.offset, self.dstoffset, self.nelems, + self.net_idx): + field._flatten_ir(handles) + + def _unsupported_pointer_operation(self): + raise ValueError('tle.remote(..., space="node") destinations only support tl.load/tl.store') + + def __add__(self, other): + self._unsupported_pointer_operation() + + def __radd__(self, other): + self._unsupported_pointer_operation() + + def __sub__(self, other): + self._unsupported_pointer_operation() + + def __rsub__(self, other): + self._unsupported_pointer_operation() + + def __getitem__(self, index): + self._unsupported_pointer_operation() + + def __triton_load__(self, mask, other, boundary_check, padding_option, cache_modifier, eviction_policy, volatile, + flagtree_hints, _semantic=None): + if mask is not None: + raise ValueError("tl.load from a node remote destination does not support mask") + if other is not None: + raise ValueError("tl.load from a node remote destination does not support other") + boundary_check = tl._unwrap_if_constexpr(boundary_check) + if isinstance(boundary_check, tl.tuple): + boundary_check = tuple(boundary_check) + if boundary_check: + raise ValueError("tl.load from a node remote destination does not support boundary_check") + padding_option = tl._unwrap_if_constexpr(padding_option) + if padding_option: + raise ValueError("tl.load from a node remote destination does not support padding_option") + cache_modifier = tl._unwrap_if_constexpr(cache_modifier) + if cache_modifier: + raise ValueError("tl.load from a node remote destination does not support cache_modifier") + eviction_policy = tl._unwrap_if_constexpr(eviction_policy) + if eviction_policy: + raise ValueError("tl.load from a node remote destination does not support eviction_policy") + volatile = tl._unwrap_if_constexpr(volatile) + if volatile: + raise ValueError("tl.load from a node remote destination does not support volatile") + flagtree_hints = tl._unwrap_if_constexpr(flagtree_hints) + if flagtree_hints: + raise ValueError("tl.load from a node remote destination does not support flagtree_hints") + + builder = _semantic.builder + if not hasattr(builder, "create_remote_pointers"): + raise RuntimeError("node get requires TLE remote_pointers support in the active Triton build") + builder.create_remote_pointers( + None, + self.src_mem.handle, + self.peer.handle, + "node", + self.offset.handle, + self.dst_mem.handle, + self.comm.handle, + self.dstoffset.handle, + self.nelems.handle, + self.net_idx.handle, + self.elem_bytes, + self.coop_kind, + "get", + ) + return _semantic.tensor(None, tl.void) + + def __triton_store__(self, value, mask, boundary_check, cache_modifier, eviction_policy, _semantic=None): + if value is not tl._STORE_VALUE_UNSET: + raise TypeError("tl.store to a node remote destination does not accept a value; " + "source and destination are configured by tle.remote(...) " + "(dst defaults to tensor), so call tl.store(remote_dst) without a value") + if mask is not None: + raise ValueError("tl.store to a node remote destination does not support mask") + boundary_check = tl._unwrap_if_constexpr(boundary_check) + if isinstance(boundary_check, tl.tuple): + boundary_check = tuple(boundary_check) + if boundary_check: + raise ValueError("tl.store to a node remote destination does not support boundary_check") + if cache_modifier: + raise ValueError("tl.store to a node remote destination does not support cache_modifier") + if eviction_policy: + raise ValueError("tl.store to a node remote destination does not support eviction_policy") + + builder = _semantic.builder + if not hasattr(builder, "create_remote_pointers"): + raise RuntimeError("node put requires TLE remote_pointers support in the active Triton build") + builder.create_remote_pointers( + None, + self.src_mem.handle, + self.peer.handle, + "node", + self.offset.handle, + self.dst_mem.handle, + self.comm.handle, + self.dstoffset.handle, + self.nelems.handle, + self.net_idx.handle, + self.elem_bytes, + self.coop_kind, + "put", + ) + return _semantic.tensor(None, tl.void) + + +def _create_node_remote_destination(src, dst, shard_id, scope, dtype, offset, dstoffset, nelems, coopkind, netidx, + _semantic) -> _node_remote_destination: + if offset is None: + raise TypeError('tle.remote(..., space="node") requires offset') + if nelems is None: + raise TypeError('tle.remote(..., space="node") requires nelems') + if coopkind is None: + raise TypeError('tle.remote(..., space="node") requires coopkind') + if dtype is None: + raise TypeError('tle.remote(..., space="node") requires dtype') + + builder = _semantic.builder + if not hasattr(builder, "create_remote_pointers"): + raise RuntimeError("node transfer requires TLE remote_pointers support in the active Triton build") + + peer = _normalize_node_peer(shard_id, scope, _semantic) + dtype = tl._unwrap_if_constexpr(dtype) + elem_bytes = _normalize_node_elem_bytes(dtype) + + offset = _normalize_node_i64(offset, "offset", must_be_positive=False, _semantic=_semantic) + if dstoffset is None: + dstoffset = offset + else: + dstoffset = _normalize_node_i64(dstoffset, "dstoffset", must_be_positive=False, _semantic=_semantic) + nelems = _normalize_node_i64(nelems, "nelems", must_be_positive=True, _semantic=_semantic) + net_idx = _normalize_node_netidx(netidx, _semantic) + coop_kind = _normalize_coopkind(coopkind) + + src_mem = tl.tensor(_parse_node_context(builder, src, "src", 0), tl.int64) + if dst is None: + dst_mem = src_mem + else: + dst_mem = tl.tensor(_parse_node_context(builder, dst, "dst", 0), tl.int64) + comm = tl.tensor(_parse_node_context(builder, src, "src", 1), tl.int64) + return _node_remote_destination( + src_mem, + dst_mem, + comm, + peer, + offset, + dstoffset, + nelems, + net_idx, + dtype=dtype, + elem_bytes=elem_bytes, + coop_kind=coop_kind, + ) + + @tl.builtin def remote( tensor=tl.tensor | None, @@ -957,6 +1282,11 @@ def remote( space: str = "cluster", dtype: tl.dtype = None, offset: int | tl.tensor | None = None, + dst=None, + dstoffset: int | tl.tensor | None = None, + nelems: int | tl.tensor | None = None, + coopkind: GroupKind | str | None = None, + netidx: int | tl.tensor = 0, _semantic=None, ): """ @@ -967,19 +1297,63 @@ def remote( should then use `tle.gpu.local_ptr(...)` to materialize remote pointers. - tl.tensor shared-memory pointer (scalar or tensor): returns remote pointer directly. + - DistributedRtContext with `space="node"`: returns a transfer-only remote + destination consumed by `tl.load` or `tl.store`. `shard_id` is the target block id inside the current thread block cluster. - When `scope` is provided, launch cluster dimensions are inferred from that - mesh and this mode requires `num_ctas=1` (one program maps to one block). - - `offset` is an optional scalar element offset relative to the target - shard's memory base address. It is only supported for `space="device"` - and is internally converted to a byte offset before being passed to - `flagcxGetIntraPointerC`. It may be a Python `int` (compile-time constant) - or a scalar `tl.tensor` (runtime value, shape == ()). + For cluster/device pointer paths, when `scope` is provided, launch cluster + dimensions are inferred from that mesh and this mode requires `num_ctas=1` + (one program maps to one block). + + For `space="node"`, `tensor` is the source registered-memory + `DistributedRtContext` and `dst` is the destination context. `dst` + defaults to `tensor`. `tl.store(remote_dst)` issues a put from the local + source to the remote destination, while `tl.load(remote_dst)` issues a get + from the remote source into the local destination. Both operations are + transfer-only and return no data; node stores do not accept a store value. + `dtype`, `offset`, `nelems`, and `coopkind` must be explicit, while + `dstoffset` defaults to `offset` and `netidx` defaults to zero. The + cooperative kind accepts only `GroupKind.THREAD`, `GroupKind.WARP`, + `GroupKind.BLOCK`, or their strings. + `shard_id` may be a scalar i32 world rank. With `scope=device_mesh`, a + compile-time tuple/list coordinate is also accepted and resolved through + the mesh's physical ids to a world rank. Node scope is used only for peer + addressing and does not alter the CUDA cluster launch. `offset`, + `dstoffset`, and `nelems` are scalar element counts normalized to i64; + lowering multiplies all three by `dtype.itemsize` before calling FlagCX. + This first version emits a network transfer without flush, completion + notification, or a visibility guarantee. + + In node mode, `offset` is relative to the source memory base and + `dstoffset` is relative to the destination memory base. For put, the + source is local and the destination is remote; for get, the source is + remote and the destination is local. For device mode, `offset` is the + remote-memory offset. It is required for + `space="node"` and optional for `space="device"`. Lowering converts it to a + byte offset before passing it to `flagcxGetIntraPointerC`. It may be a + Python `int` (compile-time constant) or a scalar `tl.tensor` (runtime + value, shape == ()). """ - shard_id = tl._unwrap_if_constexpr(shard_id) + space = tl._unwrap_if_constexpr(space) + if not isinstance(space, str): + raise TypeError(f"space must be str, got {type(space).__name__}") + if space not in ("cluster", "device", "node"): + raise ValueError(f"space must be 'cluster', 'device', or 'node', got {space!r}") + shard_id = _unwrap_remote_shard_id(shard_id) scope = tl._unwrap_if_constexpr(scope) + if space == "node": + return _create_node_remote_destination(tensor, dst, shard_id, scope, dtype, offset, dstoffset, nelems, coopkind, + netidx, _semantic) + node_only_args = [ + name for name, value in (("dst", dst), ("dstoffset", dstoffset), ("nelems", nelems), ("coopkind", coopkind)) + if value is not None + ] + unwrapped_netidx = tl._unwrap_if_constexpr(netidx) + if not isinstance(unwrapped_netidx, int) or unwrapped_netidx != 0: + node_only_args.append("netidx") + if node_only_args: + raise TypeError(f'{space} space does not accept node-only argument(s): ' + f'{", ".join(node_only_args)}') if scope is not None and not isinstance(scope, device_mesh): raise TypeError(f"scope must be device_mesh or None, got {type(scope).__name__}") if scope is not None: @@ -987,7 +1361,7 @@ def remote( # Direct pointer path: support local_ptr scalar/tensor values and return # remote pointer with preserved shape. - if isinstance(tensor, tl.tensor) or (space in ("device", "node")): + if isinstance(tensor, tl.tensor) or space == "device": return _remote_pointer(tensor, shard_id, scope=scope, space=space, _semantic=_semantic, dtype=dtype, offset=offset) diff --git a/python/triton/language/core.py b/python/triton/language/core.py index 0a229d0d2a..d7a54e0f44 100644 --- a/python/triton/language/core.py +++ b/python/triton/language/core.py @@ -41,6 +41,7 @@ T = TypeVar('T') TRITON_BUILTIN = "__triton_builtin__" +_STORE_VALUE_UNSET = object() PropagateNan = ir.PROPAGATE_NAN @@ -2177,6 +2178,11 @@ def load(pointer, mask=None, other=None, boundary_check=(), padding_option="", c :param flagtree_hints: flagtree hints :type flagtree_hints: str, optional """ + custom_load = getattr(pointer, "__triton_load__", None) + if custom_load is not None: + return custom_load(mask, other, boundary_check, padding_option, cache_modifier, eviction_policy, volatile, + flagtree_hints, _semantic=_semantic) + # `mask` and `other` can be constexpr mask = _unwrap_if_constexpr(mask) other = _unwrap_if_constexpr(other) @@ -2209,7 +2215,8 @@ def store_tensor_descriptor(desc: tensor_descriptor_base, offsets: Sequence[cons @_tensor_member_fn @builtin -def store(pointer, value, mask=None, boundary_check=(), cache_modifier="", eviction_policy="", _semantic=None): +def store(pointer, value=_STORE_VALUE_UNSET, mask=None, boundary_check=(), cache_modifier="", eviction_policy="", + _semantic=None): """ Store a tensor of data into memory locations defined by `pointer`. @@ -2233,6 +2240,9 @@ def store(pointer, value, mask=None, boundary_check=(), cache_modifier="", evict `value` is implicitly broadcast to `pointer.shape` and typecast to `pointer.dtype.element_ty`. + Experimental store-only destinations may omit `value`; ordinary pointers + still require it. + :param pointer: The memory location where the elements of `value` are stored :type pointer: `triton.PointerType`, or block of `dtype=triton.PointerType` :param value: The tensor of elements to be stored @@ -2248,13 +2258,23 @@ def store(pointer, value, mask=None, boundary_check=(), cache_modifier="", evict :param eviction_policy: changes eviction policy in NVIDIA PTX :type eviction_policy: str, optional, should be one of {"", "evict_first", "evict_last"} """ + mask = _unwrap_if_constexpr(mask) + cache_modifier = _unwrap_if_constexpr(cache_modifier) + eviction_policy = _unwrap_if_constexpr(eviction_policy) + + # Experimental pointer-like destinations can intercept stores before + # `value` is validated or converted to a tensor. + custom_store = getattr(pointer, "__triton_store__", None) + if custom_store is not None: + return custom_store(value, mask, boundary_check, cache_modifier, eviction_policy, _semantic=_semantic) + + if value is _STORE_VALUE_UNSET: + raise TypeError("tl.store() missing required argument 'value' for an ordinary pointer") + # `value` can be constexpr value = _semantic.to_tensor(value) - mask = _unwrap_if_constexpr(mask) if mask is not None: mask = _semantic.to_tensor(mask) - cache_modifier = _unwrap_if_constexpr(cache_modifier) - eviction_policy = _unwrap_if_constexpr(eviction_policy) return _semantic.store(pointer, value, mask, boundary_check, cache_modifier, eviction_policy) diff --git a/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h b/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h index 3eb5d9b974..b8a36ec24d 100644 --- a/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h +++ b/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h @@ -31,7 +31,6 @@ namespace mlir::triton::tle { void populateLocalPointersOpToLLVMPatterns( mlir::LLVMTypeConverter &typeConverter, const TargetInfoBase &targetInfo, RewritePatternSet &patterns, PatternBenefit benefit); - void populateRemotePointersOpToLLVMPatterns( mlir::LLVMTypeConverter &typeConverter, const TargetInfoBase &targetInfo, RewritePatternSet &patterns, PatternBenefit benefit); diff --git a/third_party/tle/dialect/include/IR/TleOps.td b/third_party/tle/dialect/include/IR/TleOps.td index 23a3a5201b..01b6fa7fec 100644 --- a/third_party/tle/dialect/include/IR/TleOps.td +++ b/third_party/tle/dialect/include/IR/TleOps.td @@ -356,17 +356,33 @@ def Tle_DistributedBarrierOp : Tle_Op<"distributed_barrier", let hasVerifier = 1; } -def Tle_RemotePointersOp : Tle_Op<"remote_pointers", [Pure, AttrSizedOperandSegments]> { +def Tle_RemotePointersOp : Tle_Op<"remote_pointers", [ + AttrSizedOperandSegments, + DeclareOpInterfaceMethods, + ConditionallySpeculatable +]> { let arguments = (ins Optional:$src, + Optional:$dst_mem, + Optional:$comm, TT_Int:$shard_id, StrAttr:$space, - Optional:$offset + Optional:$offset, + Optional:$dst_offset, + Optional:$nelems, + Optional:$net_idx, + OptionalAttr:$elem_bytes, + OptionalAttr:$coopkind, + OptionalAttr:$transfer_kind ); - let results = (outs Tle_LocalPointerResultType:$result); + let results = (outs Optional:$result); + let extraClassDeclaration = [{ + Speculation::Speculatability getSpeculatability(); + }]; let hasVerifier = 1; } + def Tle_GetNumPesOp : Tle_Op<"get_num_pes"> { let summary = "Get nume pes"; let arguments = (ins diff --git a/third_party/tle/dialect/include/IR/VerfiyUtils.h b/third_party/tle/dialect/include/IR/VerfiyUtils.h index b24a346842..9d694b1ee0 100644 --- a/third_party/tle/dialect/include/IR/VerfiyUtils.h +++ b/third_party/tle/dialect/include/IR/VerfiyUtils.h @@ -38,7 +38,8 @@ namespace mlir::triton::tle { namespace RemotePointers { llvm::LogicalResult verifyDeviceSpace(mlir::Value src, mlir::Value result); -} +llvm::LogicalResult verifyNodeSpace(RemotePointersOp op); +} // namespace RemotePointers namespace DistributedBarrier { llvm::LogicalResult verifyDeviceSpace(mlir::Operation *op, mlir::Value src); diff --git a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp index 4b141cd4a8..1eb2dcfeb0 100644 --- a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp +++ b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp @@ -233,7 +233,7 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { getAxisInfo(Operation *op, ArrayRef *> operands) override { auto remote = dyn_cast(op); - if (!remote || operands.empty()) + if (!remote || remote.getSpace() == "node" || operands.empty()) return AxisInfo(); const AxisInfo &baseInfo = operands[0]->getValue(); @@ -285,7 +285,8 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { } bool match(Operation *op) override { - return isa(op); + auto remote = dyn_cast(op); + return remote && remote.getSpace() != "node"; } }; diff --git a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp index 9f50d8e09e..7b0450780b 100644 --- a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp +++ b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp @@ -106,6 +106,60 @@ static LLVM::LLVMFuncOp getOrInsertGetPeerPointer(ModuleOp module, return func; } +static LLVM::LLVMFuncOp getOrInsertNetFromComm(ModuleOp module, + MLIRContext *ctx) { + const char *funcName = "flagcxDevNetGetFromCommS"; + if (auto func = module.lookupSymbol(funcName)) + return func; + + auto ptrTy = LLVM::LLVMPointerType::get(ctx); + auto i32Ty = IntegerType::get(ctx, 32); + auto funcTy = LLVM::LLVMFunctionType::get(ptrTy, {ptrTy, i32Ty}, false); + OpBuilder builder(module.getBodyRegion()); + auto func = + builder.create(module.getLoc(), funcName, funcTy); + func.setLinkage(LLVM::Linkage::External); + return func; +} + +static LLVM::LLVMFuncOp getOrInsertNetPut(ModuleOp module, MLIRContext *ctx) { + const char *funcName = "flagcxDevNetPutS"; + if (auto func = module.lookupSymbol(funcName)) + return func; + + auto voidTy = LLVM::LLVMVoidType::get(ctx); + auto ptrTy = LLVM::LLVMPointerType::get(ctx); + auto i32Ty = IntegerType::get(ctx, 32); + auto i64Ty = IntegerType::get(ctx, 64); + SmallVector argTypes{ptrTy, ptrTy, i32Ty, i32Ty, ptrTy, + i64Ty, ptrTy, i64Ty, i64Ty, i32Ty}; + auto funcTy = LLVM::LLVMFunctionType::get(voidTy, argTypes, false); + OpBuilder builder(module.getBodyRegion()); + auto func = + builder.create(module.getLoc(), funcName, funcTy); + func.setLinkage(LLVM::Linkage::External); + return func; +} + +static LLVM::LLVMFuncOp getOrInsertNetGet(ModuleOp module, MLIRContext *ctx) { + const char *funcName = "flagcxDevNetGetS"; + if (auto func = module.lookupSymbol(funcName)) + return func; + + auto voidTy = LLVM::LLVMVoidType::get(ctx); + auto ptrTy = LLVM::LLVMPointerType::get(ctx); + auto i32Ty = IntegerType::get(ctx, 32); + auto i64Ty = IntegerType::get(ctx, 64); + SmallVector argTypes{ptrTy, ptrTy, i32Ty, i32Ty, ptrTy, + i64Ty, ptrTy, i64Ty, i64Ty, i32Ty}; + auto funcTy = LLVM::LLVMFunctionType::get(voidTy, argTypes, false); + OpBuilder builder(module.getBodyRegion()); + auto func = + builder.create(module.getLoc(), funcName, funcTy); + func.setLinkage(LLVM::Linkage::External); + return func; +} + struct LocalPointersOpConversion : public ConvertOpToLLVMPattern { LocalPointersOpConversion(LLVMTypeConverter &typeConverter, @@ -468,11 +522,63 @@ LogicalResult lowerDeviceSpace(Location loc, Value mem_ptr, return success(); } -LogicalResult lowerNodeSpace(Location loc, ValueRange srcElems, - ValueRange shardElems, - ConversionPatternRewriter &rewriter, - SmallVectorImpl &resultPtrs) { - return failure(); // Not implemented yet +LogicalResult lowerNodeSpace(Location loc, tle::RemotePointersOp op, + tle::RemotePointersOp::Adaptor adaptor, + ConversionPatternRewriter &rewriter) { + ModuleOp module = op->getParentOfType(); + if (!module) + return rewriter.notifyMatchFailure(op, "expected a parent module"); + + MLIRContext *ctx = rewriter.getContext(); + auto ptrTy = LLVM::LLVMPointerType::get(ctx); + auto i32Ty = rewriter.getI32Type(); + Value dstMem = + rewriter.create(loc, ptrTy, adaptor.getDstMem()); + Value srcMem = + rewriter.create(loc, ptrTy, adaptor.getSrc()); + Value comm = rewriter.create(loc, ptrTy, adaptor.getComm()); + + Value srcByteOffset = adaptor.getOffset(); + Value dstByteOffset = adaptor.getDstOffset(); + Value byteCount = adaptor.getNelems(); + int64_t elemBytes = op->getAttrOfType("elem_bytes").getInt(); + if (elemBytes != 1) { + Value elemBytesValue = + rewriter.create(loc, elemBytes, 64); + srcByteOffset = rewriter.create(loc, adaptor.getOffset(), + elemBytesValue); + dstByteOffset = rewriter.create(loc, adaptor.getDstOffset(), + elemBytesValue); + byteCount = rewriter.create(loc, adaptor.getNelems(), + elemBytesValue); + } + + LLVM::LLVMFuncOp getNet = getOrInsertNetFromComm(module, ctx); + auto getNetCall = rewriter.create( + loc, TypeRange{ptrTy}, FlatSymbolRefAttr::get(getNet), + ValueRange{comm, adaptor.getNetIdx()}); + Value teamKind = rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(2)); + int64_t coopKindValue = op.getCoopkindAttr().getInt(); + Value coopKind = rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(coopKindValue)); + auto transferKind = op->getAttrOfType("transfer_kind").getValue(); + if (transferKind == "put") { + LLVM::LLVMFuncOp put = getOrInsertNetPut(module, ctx); + rewriter.create( + loc, TypeRange{}, FlatSymbolRefAttr::get(put), + ValueRange{getNetCall.getResult(), comm, teamKind, adaptor.getShardId(), + dstMem, dstByteOffset, srcMem, srcByteOffset, byteCount, + coopKind}); + } else { + LLVM::LLVMFuncOp get = getOrInsertNetGet(module, ctx); + rewriter.create( + loc, TypeRange{}, FlatSymbolRefAttr::get(get), + ValueRange{getNetCall.getResult(), comm, teamKind, adaptor.getShardId(), + srcMem, srcByteOffset, dstMem, dstByteOffset, byteCount, + coopKind}); + } + return success(); } Value getDistDevicePtr(tle::RemotePointersOp op, SmallVector &srcElems) { @@ -500,9 +606,15 @@ struct RemotePointersOpConversion return rewriter.notifyMatchFailure(op, msg); }; - SmallVector srcElems; auto space = adaptor.getSpace(); + if (space == "node") { + if (failed(lowerNodeSpace(loc, op, adaptor, rewriter))) + return reportFailure("node lowering failed"); + rewriter.eraseOp(op); + return success(); + } + SmallVector srcElems; if (auto src = adaptor.getSrc()) srcElems = unpackLLElements(loc, adaptor.getSrc(), rewriter); @@ -527,6 +639,7 @@ struct RemotePointersOpConversion SmallVector mappedPtrs; auto mem = getDistDevicePtr(op, srcElems); + auto resultType = op.getResult().getType(); if (space == "cluster") { if (failed(lowerClusterSpace(loc, srcElems, shardElems, rewriter, @@ -536,7 +649,7 @@ struct RemotePointersOpConversion } else if (space == "device") { int elemBytes = 1; if (offsetVal) { - Type resultPointeeTy = getRemotePointeeType(op.getType()); + Type resultPointeeTy = getRemotePointeeType(resultType); if (!resultPointeeTy) return reportFailure("result must be tt.ptr or tensor"); auto elemBits = getScalarBitWidth(resultPointeeTy); @@ -549,15 +662,11 @@ struct RemotePointersOpConversion rewriter, mappedPtrs))) { return rewriter.notifyMatchFailure(op, "device lowering failed"); } - } else if (space == "node") { - if (failed(lowerNodeSpace(loc, mem, shardElems, rewriter, mappedPtrs))) { - return rewriter.notifyMatchFailure(op, "node lowering failed"); - } } else { return reportFailure("unsupported remote space: " + space.str()); } Value packed = - packLLElements(loc, typeConverter, mappedPtrs, rewriter, op.getType()); + packLLElements(loc, typeConverter, mappedPtrs, rewriter, resultType); rewriter.replaceOp(op, packed); return success(); } diff --git a/third_party/tle/dialect/lib/IR/Ops.cpp b/third_party/tle/dialect/lib/IR/Ops.cpp index b72b66c47e..f64aaea0fe 100644 --- a/third_party/tle/dialect/lib/IR/Ops.cpp +++ b/third_party/tle/dialect/lib/IR/Ops.cpp @@ -21,6 +21,7 @@ * SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. */ +#include "mlir/Dialect/Arith/IR/Arith.h" #include "mlir/Dialect/LLVMIR/LLVMTypes.h" #include "mlir/IR/Builders.h" #include "mlir/Interfaces/SideEffectInterfaces.h" @@ -32,6 +33,7 @@ #include "llvm/ADT/SmallSet.h" #include #include +#include #include "tle/dialect/include/IR/VerfiyUtils.h" #include "triton/Dialect/TritonGPU/IR/Dialect.h" @@ -919,12 +921,52 @@ LogicalResult DistributedBarrierOp::verify() { return success(); } +void RemotePointersOp::getEffects( + SmallVectorImpl> + &effects) { + if (getSpace() != "node") + return; + effects.emplace_back(MemoryEffects::Read::get()); + effects.emplace_back(MemoryEffects::Write::get()); +} + +Speculation::Speculatability RemotePointersOp::getSpeculatability() { + return getSpace() == "node" ? Speculation::NotSpeculatable + : Speculation::Speculatable; +} + LogicalResult RemotePointersOp::verify() { - auto spaceAttr = getSpace(); + StringRef spaceAttr = getSpace(); + if (spaceAttr != "cluster" && spaceAttr != "device" && spaceAttr != "node") + return emitOpError() + << "expects space to be 'cluster', 'device', or 'node'"; + + if (!getShardId().getType().isInteger(32)) + return emitOpError() << "expects shard_id to be i32"; + + if (spaceAttr == "node") + return RemotePointers::verifyNodeSpace(*this); + + auto elemBytesAttr = (*this)->getAttrOfType("elem_bytes"); + auto coopKindAttr = getCoopkindAttr(); + auto transferKindAttr = (*this)->getAttrOfType("transfer_kind"); + + if (getDstMem() || getComm() || getDstOffset() || getNelems() || + getNetIdx() || elemBytesAttr || coopKindAttr || transferKindAttr) + return emitOpError() + << "cluster/device space does not accept node transfer operands or " + "attributes"; + if (!getResult()) + return emitOpError() + << "cluster/device space must produce a remote pointer result"; + if (spaceAttr == "device") { if (failed(RemotePointers::verifyDeviceSpace(getSrc(), getResult()))) return failure(); } else { + if (!getSrc()) + return emitOpError() << "cluster space requires a source pointer"; + Type srcTy = getSrc().getType(); Type resultTy = getResult().getType(); auto getPtrInfo = [&](Type ty, triton::PointerType &ptr, bool &isTensor, @@ -975,8 +1017,7 @@ LogicalResult RemotePointersOp::verify() { "match"; if (srcEncoding && resultEncoding && srcEncoding != resultEncoding) return emitOpError() - << "expects src/result pointer tensor encodings to " - "match"; + << "expects src/result pointer tensor encodings to match"; } if (srcPtrTy.getPointeeType() != resultPtrTy.getPointeeType()) return emitOpError() << "expects src/result pointer pointee types to " @@ -993,9 +1034,6 @@ LogicalResult RemotePointersOp::verify() { "(addrspace=7)"; } - if (!getShardId().getType().isInteger(32)) - return emitOpError() << "expects shard_id to be i32"; - bool hasOffset = getOffset() != nullptr; if (spaceAttr == "device") { if (!hasOffset) diff --git a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp index 098a404d9b..5e32085487 100644 --- a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp +++ b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp @@ -21,6 +21,7 @@ * SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. */ +#include "mlir/Dialect/Arith/IR/Arith.h" #include "mlir/Dialect/LLVMIR/LLVMTypes.h" #include "mlir/IR/Builders.h" #include "tle/dialect/include/IR/Dialect.h" @@ -36,8 +37,21 @@ #include "triton/Dialect/TritonGPU/IR/Dialect.h" #include "triton/Dialect/TritonGPU/IR/LinearLayoutConversions.h" #include +#include namespace mlir::triton::tle { +namespace { +std::optional getConstantIntValue(Value value) { + auto constant = value.getDefiningOp(); + if (!constant) + return std::nullopt; + auto integer = dyn_cast(constant.getValue()); + if (!integer) + return std::nullopt; + return integer.getInt(); +} +} // namespace + namespace RemotePointers { llvm::LogicalResult verifyDeviceSpace(mlir::Value src, mlir::Value result) { if (!src) @@ -51,6 +65,72 @@ llvm::LogicalResult verifyDeviceSpace(mlir::Value src, mlir::Value result) { } return success(); } + +llvm::LogicalResult verifyNodeSpace(RemotePointersOp op) { + if (op.getResult()) + return op.emitOpError() << "node space must not produce a result"; + + auto transferKindAttr = op->getAttrOfType("transfer_kind"); + if (!transferKindAttr || (transferKindAttr.getValue() != "put" && + transferKindAttr.getValue() != "get")) + return op.emitOpError() << "expects transfer_kind to be 'put' or 'get'"; + + auto requireOperand = [&](Value value, StringRef name) -> LogicalResult { + if (!value) + return op.emitOpError() << "node space requires " << name << " operand"; + return success(); + }; + if (failed(requireOperand(op.getSrc(), "src")) || + failed(requireOperand(op.getDstMem(), "dst_mem")) || + failed(requireOperand(op.getComm(), "comm")) || + failed(requireOperand(op.getOffset(), "offset")) || + failed(requireOperand(op.getDstOffset(), "dst_offset")) || + failed(requireOperand(op.getNelems(), "nelems")) || + failed(requireOperand(op.getNetIdx(), "net_idx"))) + return failure(); + + if (!op.getSrc().getType().isSignlessInteger(64)) + return op.emitOpError() + << "expects node source to be an i64 registered-memory handle"; + + auto elemBytesAttr = op->getAttrOfType("elem_bytes"); + if (!elemBytesAttr || elemBytesAttr.getInt() <= 0) + return op.emitOpError() << "expects elem_bytes to be > 0"; + + auto coopKindAttr = op.getCoopkindAttr(); + if (!coopKindAttr || coopKindAttr.getInt() < 0 || coopKindAttr.getInt() > 2) + return op.emitOpError() + << "expects coopkind to be THREAD(0), WARP(1), or BLOCK(2)"; + + auto verifyNonNegativeConstant = [&](Value value, + StringRef name) -> LogicalResult { + if (std::optional constant = getConstantIntValue(value); + constant && *constant < 0) + return op.emitOpError() << "expects constant " << name << " to be >= 0"; + return success(); + }; + + if (failed(verifyNonNegativeConstant(op.getShardId(), "peer")) || + failed(verifyNonNegativeConstant(op.getOffset(), "offset")) || + failed(verifyNonNegativeConstant(op.getDstOffset(), "dst_offset")) || + failed(verifyNonNegativeConstant(op.getNetIdx(), "net_idx"))) + return failure(); + + if (std::optional nelems = getConstantIntValue(op.getNelems()); + nelems && *nelems <= 0) + return op.emitOpError() << "expects constant nelems to be > 0"; + + Type offsetTy = op.getOffset().getType(); + if (auto tensorTy = dyn_cast(offsetTy)) { + if (!tensorTy.getShape().empty()) + return op.emitOpError() << "expects offset to be a scalar"; + offsetTy = tensorTy.getElementType(); + } + if (!offsetTy.isSignlessInteger(64)) + return op.emitOpError() << "expects offset to be i64"; + + return success(); +} } // namespace RemotePointers namespace DistributedBarrier { diff --git a/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp b/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp index 2a4e4acd9b..a083235ab6 100644 --- a/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp +++ b/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp @@ -462,7 +462,9 @@ collectConsumerEncodingVotes(Value root, continue; } if (auto remote = dyn_cast(owner)) { - enqueue(remote.getResult()); + // Node mode reuses src but has no pointer result to propagate. + if (remote.getSpace() != "node") + enqueue(remote.getResult()); continue; } } @@ -847,6 +849,9 @@ class SelectEncodingsPass continue; } if (auto remote = dyn_cast(owner)) { + // Node mode reuses src but has no result encoding to update. + if (remote.getSpace() == "node") + continue; auto remoteResultTy = dyn_cast(remote.getResult().getType()); if (!remoteResultTy) @@ -903,6 +908,10 @@ class SelectEncodingsPass // passes can reason about remote operands without dialect-specific // visitors. module.walk([&](triton::tle::RemotePointersOp op) { + // Node mode also has src now, but it has no pointer result whose axis + // properties could be propagated. + if (op.getSpace() == "node") + return; module->setAttr(kTleEnableEncodingRematerializationAttr, UnitAttr::get(module.getContext())); diff --git a/third_party/tle/triton_tle.cc b/third_party/tle/triton_tle.cc index 7f467c280a..4fc7de9cfc 100644 --- a/third_party/tle/triton_tle.cc +++ b/third_party/tle/triton_tle.cc @@ -601,9 +601,13 @@ void init_triton_tle_ir(py::module &&m) { py::arg("group_mask")) .def( "create_remote_pointers", - [](TritonOpBuilder &self, Type resultTy, std::optional &src, - Value shardId, const std::string &space, - std::optional &offset) -> OpState { + [](TritonOpBuilder &self, std::optional resultTy, + std::optional src, Value shardId, const std::string &space, + std::optional offset, std::optional dstMem, + std::optional comm, std::optional dstOffset, + std::optional nelems, std::optional netIdx, + std::optional elemBytes, std::optional coopKind, + std::optional transferKind) -> OpState { auto &builder = self.getBuilder(); static const std::unordered_set valid = { "cluster", "device", "node"}; @@ -612,14 +616,31 @@ void init_triton_tle_ir(py::module &&m) { "Invalid space: " + space + ". Expected one of: cluster, device, node."); } - auto space_attr = builder.getStringAttr(space); + auto spaceAttr = builder.getStringAttr(space); + IntegerAttr elemBytesAttr = + elemBytes ? builder.getI64IntegerAttr(*elemBytes) + : IntegerAttr(); + IntegerAttr coopKindAttr = + coopKind ? builder.getI32IntegerAttr(*coopKind) : IntegerAttr(); + StringAttr transferKindAttr = + transferKind ? builder.getStringAttr(*transferKind) + : StringAttr(); return self.create( - resultTy, src.value_or(Value()), shardId, space_attr, - offset.value_or(Value())); + resultTy.value_or(Type()), src.value_or(Value()), + dstMem.value_or(Value()), comm.value_or(Value()), shardId, + spaceAttr, offset.value_or(Value()), + dstOffset.value_or(Value()), nelems.value_or(Value()), + netIdx.value_or(Value()), elemBytesAttr, coopKindAttr, + transferKindAttr); }, py::arg("resultTy"), py::arg("src") = py::none(), py::arg("shardId"), - py::arg("space"), py::arg("offset") = py::none()) + py::arg("space"), py::arg("offset") = py::none(), + py::arg("dst_mem") = py::none(), py::arg("comm") = py::none(), + py::arg("dst_offset") = py::none(), py::arg("nelems") = py::none(), + py::arg("net_idx") = py::none(), py::arg("elem_bytes") = py::none(), + py::arg("coopkind") = py::none(), + py::arg("transfer_kind") = py::none()) .def("get_device_id", [](TritonOpBuilder &self, Type resultTy, std::optional src) -> Value {