From 107dcffcbdafec4dd69f6c004bc7972ae75bcb3b Mon Sep 17 00:00:00 2001 From: lzllx123 <1803100521@qq.com> Date: Thu, 30 Jul 2026 10:00:50 +0800 Subject: [PATCH 1/9] [TLE]Add a node parameter to tle.remote --- .../experimental/tle/language/distributed.py | 168 ++++++++++++++++-- .../TritonNVIDIAGPUToLLVM/TritonGPUToLLVM.cpp | 2 + .../TleToLLVM/LocalPointersOpToLLVM.h | 4 + third_party/tle/dialect/include/IR/TleOps.td | 18 ++ .../TleToLLVM/LocalPointersOpToLLVM.cpp | 109 ++++++++++-- third_party/tle/dialect/lib/IR/Ops.cpp | 45 +++++ third_party/tle/triton_tle.cc | 19 +- 7 files changed, 340 insertions(+), 25 deletions(-) diff --git a/python/triton/experimental/tle/language/distributed.py b/python/triton/experimental/tle/language/distributed.py index d9ae4db190..afa7dd2288 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): @@ -907,11 +917,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 +927,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 @@ -948,6 +952,120 @@ def _remote_pointer( return _create_remote_pointers_tensor(tensor, shard_id_tensor, _semantic, dtype=dtype, space=space, offset=offset) +# dstoffset / srcoffset / nelems -> scalar i64 tl.tensor +# dstoffset and srcoffset 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_put_coop_kind(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 _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) + + +def _node_put(dst, shard_id, src, scope, dtype, offset, dstoffset, srcoffset, + nelems, coopkind, _semantic) -> None: + if dstoffset is None: + raise TypeError('tle.remote(..., space="node") requires dstoffset') + 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') + if offset is not None: + raise ValueError('tle.remote(..., space="node") does not accept offset; use dstoffset and srcoffset') + + builder = _semantic.builder + if not hasattr(builder, "create_node_put"): + raise RuntimeError("node put requires TLE node_put support in the active Triton build") + + peer = _normalize_node_peer(shard_id, scope, _semantic) + elem_bytes = _normalize_node_elem_bytes(dtype) + + dstoffset = _normalize_node_i64(dstoffset, "dstoffset", must_be_positive=False, _semantic=_semantic) + if srcoffset is None: + srcoffset = dstoffset + else: + srcoffset = _normalize_node_i64(srcoffset, "srcoffset", must_be_positive=False, _semantic=_semantic) + nelems = _normalize_node_i64(nelems, "nelems", must_be_positive=True, _semantic=_semantic) + coop_kind = _normalize_put_coop_kind(coopkind) + + dst_mem_handle = _parse_node_context(builder, dst, "dst", 0) + dst_comm_handle = _parse_node_context(builder, dst, "dst", 1) + src_mem_handle = (dst_mem_handle if src is None else + _parse_node_context(builder, src, "src", 0)) + + builder.create_node_put(dst_mem_handle, src_mem_handle, dst_comm_handle, + peer.handle, dstoffset.handle, + srcoffset.handle, nelems.handle, elem_bytes, + coop_kind) + return None + @tl.builtin def remote( @@ -957,6 +1075,11 @@ def remote( space: str = "cluster", dtype: tl.dtype = None, offset: int | tl.tensor | None = None, + src=None, + dstoffset: int | tl.tensor | None = None, + srcoffset: int | tl.tensor | None = None, + nelems: int | tl.tensor | None = None, + coopkind: GroupKind | str | None = None, _semantic=None, ): """ @@ -969,8 +1092,24 @@ def remote( pointer directly. `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). + 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 destination `DistributedRtContext`. + The optional `src` is another `DistributedRtContext`; it defaults to + `tensor`, providing same registered-buffer transfer by default. + `dtype`, `dstoffset`, `nelems`, and `coopkind` must be explicit, while + `srcoffset` defaults to `dstoffset`. 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. `dstoffset`, + `srcoffset`, 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 put without flush, completion + notification, or a remote-visibility guarantee. `offset` is an optional scalar element offset relative to the target shard's memory base address. It is only supported for `space="device"` @@ -978,8 +1117,13 @@ def remote( `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) + shard_id = _unwrap_remote_shard_id(shard_id) scope = tl._unwrap_if_constexpr(scope) + if space == "node": + return _node_put(tensor, shard_id, src, scope, dtype, offset, + dstoffset, srcoffset, nelems, coopkind, + _semantic) 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 +1131,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/third_party/nvidia/lib/TritonNVIDIAGPUToLLVM/TritonGPUToLLVM.cpp b/third_party/nvidia/lib/TritonNVIDIAGPUToLLVM/TritonGPUToLLVM.cpp index 552dc79edf..cf2c43efc3 100644 --- a/third_party/nvidia/lib/TritonNVIDIAGPUToLLVM/TritonGPUToLLVM.cpp +++ b/third_party/nvidia/lib/TritonNVIDIAGPUToLLVM/TritonGPUToLLVM.cpp @@ -178,6 +178,8 @@ struct ConvertTritonGPUToLLVM typeConverter, patterns, benefit); mlir::triton::tle::populateLocalPointersOpToLLVMPatterns( typeConverter, targetInfo, patterns, benefit); + mlir::triton::tle::populateNodePutOpToLLVMPatterns(typeConverter, + patterns, benefit); mlir::triton::tle::populateExtractTileOpToLLVMPatterns( typeConverter, patterns, targetInfo, benefit); mlir::triton::tle::populateInsertTileOpToLLVMPatterns( diff --git a/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h b/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h index 3eb5d9b974..92b19dce6c 100644 --- a/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h +++ b/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h @@ -32,6 +32,10 @@ void populateLocalPointersOpToLLVMPatterns( mlir::LLVMTypeConverter &typeConverter, const TargetInfoBase &targetInfo, RewritePatternSet &patterns, PatternBenefit benefit); +void populateNodePutOpToLLVMPatterns( + mlir::LLVMTypeConverter &typeConverter, 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..c16c0a7bb4 100644 --- a/third_party/tle/dialect/include/IR/TleOps.td +++ b/third_party/tle/dialect/include/IR/TleOps.td @@ -367,6 +367,24 @@ def Tle_RemotePointersOp : Tle_Op<"remote_pointers", [Pure, AttrSizedOperandSegm let results = (outs Tle_LocalPointerResultType:$result); let hasVerifier = 1; } + +def Tle_NodePutOp + : Tle_Op<"node_put", [MemoryEffects<[MemRead, MemWrite]>]> { + + let arguments = (ins + I64:$dst_mem, + I64:$src_mem, + I64:$comm, + I32:$peer, + I64:$dst_offset, + I64:$src_offset, + I64:$nelems, + I64Attr:$elem_bytes, + I32Attr:$put_coop_kind + ); + 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/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp index 9f50d8e09e..abd93822b1 100644 --- a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp +++ b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp @@ -106,6 +106,98 @@ 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; +} + +struct NodePutOpConversion : public ConvertOpToLLVMPattern { + using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; + + LogicalResult + matchAndRewrite(tle::NodePutOp op, OpAdaptor adaptor, + ConversionPatternRewriter &rewriter) const override { + Location loc = op.getLoc(); + MLIRContext *ctx = rewriter.getContext(); + ModuleOp module = op->getParentOfType(); + if (!module) + return rewriter.notifyMatchFailure(op, "expected a parent module"); + + 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.getSrcMem()); + Value comm = + rewriter.create(loc, ptrTy, adaptor.getComm()); + + Value dstByteOffset = adaptor.getDstOffset(); + Value srcByteOffset = adaptor.getSrcOffset(); + Value byteCount = adaptor.getNelems(); + if (op.getElemBytes() != 1) { + Value elemBytes = rewriter.create( + loc, op.getElemBytes(), 64); + dstByteOffset = rewriter.create( + loc, adaptor.getDstOffset(), elemBytes); + srcByteOffset = rewriter.create( + loc, adaptor.getSrcOffset(), elemBytes); + byteCount = + rewriter.create(loc, adaptor.getNelems(), elemBytes); + } + + LLVM::LLVMFuncOp getNet = getOrInsertNetFromComm(module, ctx); + LLVM::LLVMFuncOp put = getOrInsertNetPut(module, ctx); + Value netIdx = rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(0)); + auto getNetCall = rewriter.create( + loc, TypeRange{ptrTy}, FlatSymbolRefAttr::get(getNet), + ValueRange{comm, netIdx}); + Value teamKind = rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(2)); + Value coopKind = rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(op.getPutCoopKind())); + rewriter.create( + loc, TypeRange{}, FlatSymbolRefAttr::get(put), + ValueRange{getNetCall.getResult(), comm, teamKind, adaptor.getPeer(), + dstMem, dstByteOffset, srcMem, srcByteOffset, byteCount, + coopKind}); + rewriter.eraseOp(op); + return success(); + } +}; + struct LocalPointersOpConversion : public ConvertOpToLLVMPattern { LocalPointersOpConversion(LLVMTypeConverter &typeConverter, @@ -468,13 +560,6 @@ 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 -} - Value getDistDevicePtr(tle::RemotePointersOp op, SmallVector &srcElems) { if (!srcElems.empty()) return srcElems[0]; @@ -549,10 +634,6 @@ 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()); } @@ -571,6 +652,12 @@ void tle::populateLocalPointersOpToLLVMPatterns( patterns.add(typeConverter, targetInfo, benefit); } +void tle::populateNodePutOpToLLVMPatterns(LLVMTypeConverter &typeConverter, + RewritePatternSet &patterns, + PatternBenefit benefit) { + patterns.add(typeConverter, benefit); +} + void tle::populateRemotePointersOpToLLVMPatterns( LLVMTypeConverter &typeConverter, const TargetInfoBase &targetInfo, RewritePatternSet &patterns, PatternBenefit benefit) { diff --git a/third_party/tle/dialect/lib/IR/Ops.cpp b/third_party/tle/dialect/lib/IR/Ops.cpp index b72b66c47e..5727034a8e 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" @@ -44,6 +46,16 @@ namespace { constexpr int kSharedMemoryAddressSpace = 3; // Cluster-shared pointers map to LLVM address space 7 (NVVM shared::cluster). constexpr int kClusterSharedMemoryAddressSpace = 7; + +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 // ============================================================================ @@ -921,6 +933,9 @@ LogicalResult DistributedBarrierOp::verify() { LogicalResult RemotePointersOp::verify() { auto spaceAttr = getSpace(); + if (spaceAttr != "cluster" && spaceAttr != "device") + return emitOpError() + << "expects space to be either 'cluster' or 'device'"; if (spaceAttr == "device") { if (failed(RemotePointers::verifyDeviceSpace(getSrc(), getResult()))) return failure(); @@ -1020,4 +1035,34 @@ LogicalResult RemotePointersOp::verify() { return success(); } +LogicalResult NodePutOp::verify() { + if (getElemBytes() <= 0) + return emitOpError() << "expects elem_bytes to be > 0"; + + int64_t coopKind = getPutCoopKind(); + if (coopKind < 0 || coopKind > 2) + return emitOpError() + << "expects put_coop_kind 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 emitOpError() + << "expects constant " << name << " to be >= 0"; + return success(); + }; + + if (failed(verifyNonNegativeConstant(getPeer(), "peer")) || + failed(verifyNonNegativeConstant(getDstOffset(), "dst_offset")) || + failed(verifyNonNegativeConstant(getSrcOffset(), "src_offset"))) + return failure(); + + if (std::optional nelems = getConstantIntValue(getNelems()); + nelems && *nelems <= 0) + return emitOpError() << "expects constant nelems to be > 0"; + + return success(); +} + } // namespace mlir::triton::tle diff --git a/third_party/tle/triton_tle.cc b/third_party/tle/triton_tle.cc index 7f467c280a..8aae5c4dc6 100644 --- a/third_party/tle/triton_tle.cc +++ b/third_party/tle/triton_tle.cc @@ -606,11 +606,11 @@ void init_triton_tle_ir(py::module &&m) { std::optional &offset) -> OpState { auto &builder = self.getBuilder(); static const std::unordered_set valid = { - "cluster", "device", "node"}; + "cluster", "device"}; if (valid.find(space) == valid.end()) { throw std::invalid_argument( "Invalid space: " + space + - ". Expected one of: cluster, device, node."); + ". Expected one of: cluster, device."); } auto space_attr = builder.getStringAttr(space); @@ -620,6 +620,21 @@ void init_triton_tle_ir(py::module &&m) { }, py::arg("resultTy"), py::arg("src") = py::none(), py::arg("shardId"), py::arg("space"), py::arg("offset") = py::none()) + .def( + "create_node_put", + [](TritonOpBuilder &self, Value dstMem, Value srcMem, Value comm, + Value peer, Value dstOffset, Value srcOffset, Value nelems, + int64_t elemBytes, int32_t putCoopKind) -> void { + auto &builder = self.getBuilder(); + self.create( + dstMem, srcMem, comm, peer, dstOffset, srcOffset, nelems, + builder.getI64IntegerAttr(elemBytes), + builder.getI32IntegerAttr(putCoopKind)); + }, + py::arg("dst_mem"), py::arg("src_mem"), py::arg("comm"), + py::arg("peer"), py::arg("dst_offset"), py::arg("src_offset"), + py::arg("nelems"), py::arg("elem_bytes"), + py::arg("put_coop_kind")) .def("get_device_id", [](TritonOpBuilder &self, Type resultTy, std::optional src) -> Value { From aabeb889ce736313cf2318e2f16695223edab446 Mon Sep 17 00:00:00 2001 From: lzllx123 <1803100521@qq.com> Date: Mon, 3 Aug 2026 17:49:25 +0800 Subject: [PATCH 2/9] [TLE]Change the way tle.remote is used for communication between nodes --- .../experimental/tle/language/distributed.py | 258 +++++++++++++++--- python/triton/language/core.py | 27 +- .../TritonNVIDIAGPUToLLVM/TritonGPUToLLVM.cpp | 2 - .../TleToLLVM/LocalPointersOpToLLVM.h | 5 - third_party/tle/dialect/include/IR/TleOps.td | 38 ++- .../tle/dialect/include/IR/VerfiyUtils.h | 1 + .../tle/dialect/lib/Analysis/AxisInfoExt.cpp | 6 +- .../TleToLLVM/LocalPointersOpToLLVM.cpp | 127 +++++---- third_party/tle/dialect/lib/IR/Ops.cpp | 90 +++--- .../tle/dialect/lib/IR/VerfiyUtils.cpp | 77 ++++++ .../lib/Transforms/TleSelectEncodings.cpp | 7 +- third_party/tle/triton_tle.cc | 56 ++-- 12 files changed, 481 insertions(+), 213 deletions(-) diff --git a/python/triton/experimental/tle/language/distributed.py b/python/triton/experimental/tle/language/distributed.py index afa7dd2288..7999a3aa8d 100644 --- a/python/triton/experimental/tle/language/distributed.py +++ b/python/triton/experimental/tle/language/distributed.py @@ -952,8 +952,8 @@ def _remote_pointer( return _create_remote_pointers_tensor(tensor, shard_id_tensor, _semantic, dtype=dtype, space=space, offset=offset) -# dstoffset / srcoffset / nelems -> scalar i64 tl.tensor -# dstoffset and srcoffset must be >= 0. +# offset / srcoffset / nelems -> scalar i64 tl.tensor +# offset and srcoffset 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) @@ -1019,6 +1019,27 @@ def _normalize_put_coop_kind(coopkind) -> int: 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) @@ -1027,45 +1048,195 @@ def _parse_node_context(builder, value, label: str, index: int): return _parse_src_arg(builder, value, index) -def _node_put(dst, shard_id, src, scope, dtype, offset, dstoffset, srcoffset, - nelems, coopkind, _semantic) -> None: - if dstoffset is None: - raise TypeError('tle.remote(..., space="node") requires dstoffset') +class _node_remote_destination_type(tl.base_type): + + def __init__(self, field_types, dtype: tl.dtype, elem_bytes: int, coop_kind: int): + # dst_mem, src_mem, comm, peer, dst_offset, src_offset, 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.store') + + +class _node_remote_destination(tl.base_value): + + def __init__(self, dst_mem: tl.tensor, src_mem: tl.tensor, comm: tl.tensor, + peer: tl.tensor, dst_offset: tl.tensor, src_offset: tl.tensor, + nelems: tl.tensor, net_idx: tl.tensor, *, dtype: tl.dtype, + elem_bytes: int, coop_kind: int): + super().__init__() + self.dst_mem = dst_mem + self.src_mem = src_mem + self.comm = comm + self.peer = peer + self.dst_offset = dst_offset + self.src_offset = src_offset + 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.dst_mem, self.src_mem, self.comm, self.peer, + self.dst_offset, self.src_offset, 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.dst_mem, self.src_mem, self.comm, self.peer, + self.dst_offset, self.src_offset, self.nelems, + self.net_idx): + field._flatten_ir(handles) + + def _unsupported_pointer_operation(self): + raise ValueError('tle.remote(..., space="node") destinations only support 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): + self._unsupported_pointer_operation() + + 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; " + "pass the source context to tle.remote(..., src=...) and call tl.store(remote_dst)") + 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, + None, + self.peer.handle, + "node", + self.dst_offset.handle, + self.dst_mem.handle, + self.src_mem.handle, + self.comm.handle, + self.src_offset.handle, + self.nelems.handle, + self.net_idx.handle, + self.elem_bytes, + self.coop_kind, + ) + return _semantic.tensor(None, tl.void) + + +def _create_node_remote_destination(dst, shard_id, scope, dtype, offset, src, + srcoffset, 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') - if offset is not None: - raise ValueError('tle.remote(..., space="node") does not accept offset; use dstoffset and srcoffset') builder = _semantic.builder - if not hasattr(builder, "create_node_put"): - raise RuntimeError("node put requires TLE node_put support in the active Triton build") + if not hasattr(builder, "create_remote_pointers"): + raise RuntimeError( + "node put 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) - dstoffset = _normalize_node_i64(dstoffset, "dstoffset", must_be_positive=False, _semantic=_semantic) + offset = _normalize_node_i64( + offset, "offset", must_be_positive=False, _semantic=_semantic) if srcoffset is None: - srcoffset = dstoffset + srcoffset = offset else: - srcoffset = _normalize_node_i64(srcoffset, "srcoffset", must_be_positive=False, _semantic=_semantic) - nelems = _normalize_node_i64(nelems, "nelems", must_be_positive=True, _semantic=_semantic) + srcoffset = _normalize_node_i64( + srcoffset, "srcoffset", 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_put_coop_kind(coopkind) - dst_mem_handle = _parse_node_context(builder, dst, "dst", 0) - dst_comm_handle = _parse_node_context(builder, dst, "dst", 1) - src_mem_handle = (dst_mem_handle if src is None else - _parse_node_context(builder, src, "src", 0)) - - builder.create_node_put(dst_mem_handle, src_mem_handle, dst_comm_handle, - peer.handle, dstoffset.handle, - srcoffset.handle, nelems.handle, elem_bytes, - coop_kind) - return None - + if src is None: + src = dst + dst_mem = tl.tensor(_parse_node_context(builder, dst, "dst", 0), tl.int64) + src_mem = tl.tensor(_parse_node_context(builder, src, "src", 0), tl.int64) + dst_comm = tl.tensor(_parse_node_context(builder, dst, "dst", 1), tl.int64) + return _node_remote_destination( + dst_mem, + src_mem, + dst_comm, + peer, + offset, + srcoffset, + nelems, + net_idx, + dtype=dtype, + elem_bytes=elem_bytes, + coop_kind=coop_kind, + ) @tl.builtin def remote( @@ -1076,10 +1247,10 @@ def remote( dtype: tl.dtype = None, offset: int | tl.tensor | None = None, src=None, - dstoffset: int | tl.tensor | None = None, srcoffset: int | tl.tensor | None = None, nelems: int | tl.tensor | None = None, coopkind: GroupKind | str | None = None, + netidx: int | tl.tensor = 0, _semantic=None, ): """ @@ -1090,40 +1261,45 @@ 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 store-only remote + destination consumed by `tl.store`. `shard_id` is the target block id inside the current thread block cluster. 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 destination `DistributedRtContext`. - The optional `src` is another `DistributedRtContext`; it defaults to - `tensor`, providing same registered-buffer transfer by default. - `dtype`, `dstoffset`, `nelems`, and `coopkind` must be explicit, while - `srcoffset` defaults to `dstoffset`. The cooperative kind accepts only - `GroupKind.THREAD`, `GroupKind.WARP`, `GroupKind.BLOCK`, or their strings. + For `space="node"`, `tensor` is the destination `DistributedRtContext` + and `src` is the source registered-memory context. `src` defaults to + `tensor`. This function returns a store-only destination triggered with + `tl.store(remote_dst)`; node destinations do not accept a store value. + `dtype`, `offset`, `nelems`, and `coopkind` must be explicit, while + `srcoffset` 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. `dstoffset`, + addressing and does not alter the CUDA cluster launch. `offset`, `srcoffset`, 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 put without flush, completion notification, or a remote-visibility guarantee. - `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 == ()). + `offset` is a scalar element offset relative to the target shard's memory + base address. It is required for `space="node"` and optional for + `space="device"`. The device path 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 == ()). """ space = tl._unwrap_if_constexpr(space) shard_id = _unwrap_remote_shard_id(shard_id) scope = tl._unwrap_if_constexpr(scope) if space == "node": - return _node_put(tensor, shard_id, src, scope, dtype, offset, - dstoffset, srcoffset, nelems, coopkind, - _semantic) + return _create_node_remote_destination( + tensor, shard_id, scope, dtype, offset, src, srcoffset, nelems, + coopkind, netidx, _semantic) 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: diff --git a/python/triton/language/core.py b/python/triton/language/core.py index 0a229d0d2a..77e491f7f7 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,7 @@ 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 +2239,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 +2257,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/nvidia/lib/TritonNVIDIAGPUToLLVM/TritonGPUToLLVM.cpp b/third_party/nvidia/lib/TritonNVIDIAGPUToLLVM/TritonGPUToLLVM.cpp index cf2c43efc3..552dc79edf 100644 --- a/third_party/nvidia/lib/TritonNVIDIAGPUToLLVM/TritonGPUToLLVM.cpp +++ b/third_party/nvidia/lib/TritonNVIDIAGPUToLLVM/TritonGPUToLLVM.cpp @@ -178,8 +178,6 @@ struct ConvertTritonGPUToLLVM typeConverter, patterns, benefit); mlir::triton::tle::populateLocalPointersOpToLLVMPatterns( typeConverter, targetInfo, patterns, benefit); - mlir::triton::tle::populateNodePutOpToLLVMPatterns(typeConverter, - patterns, benefit); mlir::triton::tle::populateExtractTileOpToLLVMPatterns( typeConverter, patterns, targetInfo, benefit); mlir::triton::tle::populateInsertTileOpToLLVMPatterns( diff --git a/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h b/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h index 92b19dce6c..b8a36ec24d 100644 --- a/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h +++ b/third_party/tle/dialect/include/Conversion/TleToLLVM/LocalPointersOpToLLVM.h @@ -31,11 +31,6 @@ namespace mlir::triton::tle { void populateLocalPointersOpToLLVMPatterns( mlir::LLVMTypeConverter &typeConverter, const TargetInfoBase &targetInfo, RewritePatternSet &patterns, PatternBenefit benefit); - -void populateNodePutOpToLLVMPatterns( - mlir::LLVMTypeConverter &typeConverter, 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 c16c0a7bb4..39b62dfdf9 100644 --- a/third_party/tle/dialect/include/IR/TleOps.td +++ b/third_party/tle/dialect/include/IR/TleOps.td @@ -356,32 +356,30 @@ 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:$src_mem, + Optional:$comm, TT_Int:$shard_id, StrAttr:$space, - Optional:$offset + Optional:$offset, + Optional:$src_offset, + Optional:$nelems, + Optional:$net_idx, + OptionalAttr:$elem_bytes, + OptionalAttr:$put_coop_kind ); - let results = (outs Tle_LocalPointerResultType:$result); - let hasVerifier = 1; -} - -def Tle_NodePutOp - : Tle_Op<"node_put", [MemoryEffects<[MemRead, MemWrite]>]> { - - let arguments = (ins - I64:$dst_mem, - I64:$src_mem, - I64:$comm, - I32:$peer, - I64:$dst_offset, - I64:$src_offset, - I64:$nelems, - I64Attr:$elem_bytes, - I32Attr:$put_coop_kind - ); + let results = (outs Optional:$result); + let extraClassDeclaration = [{ + Speculation::Speculatability getSpeculatability(); + }]; let hasVerifier = 1; } diff --git a/third_party/tle/dialect/include/IR/VerfiyUtils.h b/third_party/tle/dialect/include/IR/VerfiyUtils.h index b24a346842..d268926783 100644 --- a/third_party/tle/dialect/include/IR/VerfiyUtils.h +++ b/third_party/tle/dialect/include/IR/VerfiyUtils.h @@ -38,6 +38,7 @@ namespace mlir::triton::tle { namespace RemotePointers { llvm::LogicalResult verifyDeviceSpace(mlir::Value src, mlir::Value result); +llvm::LogicalResult verifyNodeSpace(RemotePointersOp op); } namespace DistributedBarrier { diff --git a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp index 4b141cd4a8..38e89b9a75 100644 --- a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp +++ b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp @@ -233,7 +233,8 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { getAxisInfo(Operation *op, ArrayRef *> operands) override { auto remote = dyn_cast(op); - if (!remote || operands.empty()) + if (!remote || !remote.getResult() || remote.getSpace() == "node" || + operands.empty()) return AxisInfo(); const AxisInfo &baseInfo = operands[0]->getValue(); @@ -285,7 +286,8 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { } bool match(Operation *op) override { - return isa(op); + auto remote = dyn_cast(op); + return remote && remote.getResult() && 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 abd93822b1..ea2602ab84 100644 --- a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp +++ b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp @@ -142,62 +142,6 @@ static LLVM::LLVMFuncOp getOrInsertNetPut(ModuleOp module, return func; } -struct NodePutOpConversion : public ConvertOpToLLVMPattern { - using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern; - - LogicalResult - matchAndRewrite(tle::NodePutOp op, OpAdaptor adaptor, - ConversionPatternRewriter &rewriter) const override { - Location loc = op.getLoc(); - MLIRContext *ctx = rewriter.getContext(); - ModuleOp module = op->getParentOfType(); - if (!module) - return rewriter.notifyMatchFailure(op, "expected a parent module"); - - 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.getSrcMem()); - Value comm = - rewriter.create(loc, ptrTy, adaptor.getComm()); - - Value dstByteOffset = adaptor.getDstOffset(); - Value srcByteOffset = adaptor.getSrcOffset(); - Value byteCount = adaptor.getNelems(); - if (op.getElemBytes() != 1) { - Value elemBytes = rewriter.create( - loc, op.getElemBytes(), 64); - dstByteOffset = rewriter.create( - loc, adaptor.getDstOffset(), elemBytes); - srcByteOffset = rewriter.create( - loc, adaptor.getSrcOffset(), elemBytes); - byteCount = - rewriter.create(loc, adaptor.getNelems(), elemBytes); - } - - LLVM::LLVMFuncOp getNet = getOrInsertNetFromComm(module, ctx); - LLVM::LLVMFuncOp put = getOrInsertNetPut(module, ctx); - Value netIdx = rewriter.create( - loc, i32Ty, rewriter.getI32IntegerAttr(0)); - auto getNetCall = rewriter.create( - loc, TypeRange{ptrTy}, FlatSymbolRefAttr::get(getNet), - ValueRange{comm, netIdx}); - Value teamKind = rewriter.create( - loc, i32Ty, rewriter.getI32IntegerAttr(2)); - Value coopKind = rewriter.create( - loc, i32Ty, rewriter.getI32IntegerAttr(op.getPutCoopKind())); - rewriter.create( - loc, TypeRange{}, FlatSymbolRefAttr::get(put), - ValueRange{getNetCall.getResult(), comm, teamKind, adaptor.getPeer(), - dstMem, dstByteOffset, srcMem, srcByteOffset, byteCount, - coopKind}); - rewriter.eraseOp(op); - return success(); - } -}; - struct LocalPointersOpConversion : public ConvertOpToLLVMPattern { LocalPointersOpConversion(LLVMTypeConverter &typeConverter, @@ -560,6 +504,58 @@ LogicalResult lowerDeviceSpace(Location loc, Value mem_ptr, return success(); } +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.getSrcMem()); + Value comm = + rewriter.create(loc, ptrTy, adaptor.getComm()); + + Value dstByteOffset = adaptor.getOffset(); + Value srcByteOffset = adaptor.getSrcOffset(); + Value byteCount = adaptor.getNelems(); + int64_t elemBytes = + op->getAttrOfType("elem_bytes").getInt(); + if (elemBytes != 1) { + Value elemBytesValue = rewriter.create( + loc, elemBytes, 64); + dstByteOffset = rewriter.create( + loc, adaptor.getOffset(), elemBytesValue); + srcByteOffset = rewriter.create( + loc, adaptor.getSrcOffset(), elemBytesValue); + byteCount = rewriter.create( + loc, adaptor.getNelems(), elemBytesValue); + } + + LLVM::LLVMFuncOp getNet = getOrInsertNetFromComm(module, ctx); + LLVM::LLVMFuncOp put = getOrInsertNetPut(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 putCoopKind = + op->getAttrOfType("put_coop_kind").getInt(); + Value coopKind = rewriter.create( + loc, i32Ty, rewriter.getI32IntegerAttr(putCoopKind)); + rewriter.create( + loc, TypeRange{}, FlatSymbolRefAttr::get(put), + ValueRange{getNetCall.getResult(), comm, teamKind, adaptor.getShardId(), + dstMem, dstByteOffset, srcMem, srcByteOffset, byteCount, + coopKind}); + return success(); +} + Value getDistDevicePtr(tle::RemotePointersOp op, SmallVector &srcElems) { if (!srcElems.empty()) return srcElems[0]; @@ -585,9 +581,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); @@ -612,6 +614,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, @@ -621,7 +624,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); @@ -638,7 +641,7 @@ struct RemotePointersOpConversion 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(); } @@ -652,12 +655,6 @@ void tle::populateLocalPointersOpToLLVMPatterns( patterns.add(typeConverter, targetInfo, benefit); } -void tle::populateNodePutOpToLLVMPatterns(LLVMTypeConverter &typeConverter, - RewritePatternSet &patterns, - PatternBenefit benefit) { - patterns.add(typeConverter, benefit); -} - void tle::populateRemotePointersOpToLLVMPatterns( LLVMTypeConverter &typeConverter, const TargetInfoBase &targetInfo, RewritePatternSet &patterns, PatternBenefit benefit) { diff --git a/third_party/tle/dialect/lib/IR/Ops.cpp b/third_party/tle/dialect/lib/IR/Ops.cpp index 5727034a8e..0c2f96eb52 100644 --- a/third_party/tle/dialect/lib/IR/Ops.cpp +++ b/third_party/tle/dialect/lib/IR/Ops.cpp @@ -46,16 +46,6 @@ namespace { constexpr int kSharedMemoryAddressSpace = 3; // Cluster-shared pointers map to LLVM address space 7 (NVVM shared::cluster). constexpr int kClusterSharedMemoryAddressSpace = 7; - -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 // ============================================================================ @@ -931,15 +921,53 @@ 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(); - if (spaceAttr != "cluster" && spaceAttr != "device") + 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 putCoopKindAttr = + (*this)->getAttrOfType("put_coop_kind"); + + if (getDstMem() || getSrcMem() || getComm() || getSrcOffset() || + getNelems() || getNetIdx() || elemBytesAttr || putCoopKindAttr) + return emitOpError() + << "cluster/device space does not accept node put operands or " + "attributes"; + if (!getResult()) return emitOpError() - << "expects space to be either 'cluster' or 'device'"; + << "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, @@ -990,8 +1018,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 " @@ -1008,9 +1035,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) @@ -1035,34 +1059,4 @@ LogicalResult RemotePointersOp::verify() { return success(); } -LogicalResult NodePutOp::verify() { - if (getElemBytes() <= 0) - return emitOpError() << "expects elem_bytes to be > 0"; - - int64_t coopKind = getPutCoopKind(); - if (coopKind < 0 || coopKind > 2) - return emitOpError() - << "expects put_coop_kind 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 emitOpError() - << "expects constant " << name << " to be >= 0"; - return success(); - }; - - if (failed(verifyNonNegativeConstant(getPeer(), "peer")) || - failed(verifyNonNegativeConstant(getDstOffset(), "dst_offset")) || - failed(verifyNonNegativeConstant(getSrcOffset(), "src_offset"))) - return failure(); - - if (std::optional nelems = getConstantIntValue(getNelems()); - nelems && *nelems <= 0) - return emitOpError() << "expects constant nelems to be > 0"; - - return success(); -} - } // namespace mlir::triton::tle diff --git a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp index 098a404d9b..e2a2340252 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,69 @@ llvm::LogicalResult verifyDeviceSpace(mlir::Value src, mlir::Value result) { } return success(); } + +llvm::LogicalResult verifyNodeSpace(RemotePointersOp op) { + if (op.getSrc()) + return op.emitOpError() + << "node space does not accept a pointer source operand"; + if (op.getResult()) + return op.emitOpError() << "node space must not produce a result"; + + auto requireOperand = [&](Value value, StringRef name) -> LogicalResult { + if (!value) + return op.emitOpError() + << "node space requires " << name << " operand"; + return success(); + }; + if (failed(requireOperand(op.getDstMem(), "dst_mem")) || + failed(requireOperand(op.getSrcMem(), "src_mem")) || + failed(requireOperand(op.getComm(), "comm")) || + failed(requireOperand(op.getOffset(), "offset")) || + failed(requireOperand(op.getSrcOffset(), "src_offset")) || + failed(requireOperand(op.getNelems(), "nelems")) || + failed(requireOperand(op.getNetIdx(), "net_idx"))) + return failure(); + + auto elemBytesAttr = op->getAttrOfType("elem_bytes"); + if (!elemBytesAttr || elemBytesAttr.getInt() <= 0) + return op.emitOpError() << "expects elem_bytes to be > 0"; + + auto putCoopKindAttr = op->getAttrOfType("put_coop_kind"); + if (!putCoopKindAttr || putCoopKindAttr.getInt() < 0 || + putCoopKindAttr.getInt() > 2) + return op.emitOpError() + << "expects put_coop_kind 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(), "dst_offset")) || + failed(verifyNonNegativeConstant(op.getSrcOffset(), "src_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..09792d1fb9 100644 --- a/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp +++ b/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp @@ -462,7 +462,8 @@ collectConsumerEncodingVotes(Value root, continue; } if (auto remote = dyn_cast(owner)) { - enqueue(remote.getResult()); + if (Value result = remote.getResult()) + enqueue(result); continue; } } @@ -847,6 +848,8 @@ class SelectEncodingsPass continue; } if (auto remote = dyn_cast(owner)) { + if (!remote.getResult()) + continue; auto remoteResultTy = dyn_cast(remote.getResult().getType()); if (!remoteResultTy) @@ -903,6 +906,8 @@ class SelectEncodingsPass // passes can reason about remote operands without dialect-specific // visitors. module.walk([&](triton::tle::RemotePointersOp op) { + if (op.getSpace() == "node" || !op.getResult() || !op.getSrc()) + 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 8aae5c4dc6..1406ecd83f 100644 --- a/third_party/tle/triton_tle.cc +++ b/third_party/tle/triton_tle.cc @@ -601,40 +601,46 @@ 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 srcMem, + std::optional comm, std::optional srcOffset, + std::optional nelems, std::optional netIdx, + std::optional elemBytes, + std::optional putCoopKind) -> OpState { auto &builder = self.getBuilder(); static const std::unordered_set valid = { - "cluster", "device"}; + "cluster", "device", "node"}; if (valid.find(space) == valid.end()) { throw std::invalid_argument( "Invalid space: " + space + - ". Expected one of: cluster, device."); + ". 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 putCoopKindAttr = + putCoopKind ? builder.getI32IntegerAttr(*putCoopKind) + : IntegerAttr(); 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()), srcMem.value_or(Value()), + comm.value_or(Value()), shardId, spaceAttr, + offset.value_or(Value()), srcOffset.value_or(Value()), + nelems.value_or(Value()), netIdx.value_or(Value()), + elemBytesAttr, putCoopKindAttr); }, - py::arg("resultTy"), py::arg("src") = py::none(), py::arg("shardId"), - py::arg("space"), py::arg("offset") = py::none()) - .def( - "create_node_put", - [](TritonOpBuilder &self, Value dstMem, Value srcMem, Value comm, - Value peer, Value dstOffset, Value srcOffset, Value nelems, - int64_t elemBytes, int32_t putCoopKind) -> void { - auto &builder = self.getBuilder(); - self.create( - dstMem, srcMem, comm, peer, dstOffset, srcOffset, nelems, - builder.getI64IntegerAttr(elemBytes), - builder.getI32IntegerAttr(putCoopKind)); - }, - py::arg("dst_mem"), py::arg("src_mem"), py::arg("comm"), - py::arg("peer"), py::arg("dst_offset"), py::arg("src_offset"), - py::arg("nelems"), py::arg("elem_bytes"), - py::arg("put_coop_kind")) + py::arg("resultTy"), py::arg("src") = py::none(), + py::arg("shardId"), py::arg("space"), + py::arg("offset") = py::none(), py::arg("dst_mem") = py::none(), + py::arg("src_mem") = py::none(), py::arg("comm") = py::none(), + py::arg("src_offset") = py::none(), + py::arg("nelems") = py::none(), py::arg("net_idx") = py::none(), + py::arg("elem_bytes") = py::none(), + py::arg("put_coop_kind") = py::none()) .def("get_device_id", [](TritonOpBuilder &self, Type resultTy, std::optional src) -> Value { From 6dfca1927c4192b789df98bcbc476b27eee7a53b Mon Sep 17 00:00:00 2001 From: flagtree-bot Date: Wed, 5 Aug 2026 03:52:59 +0000 Subject: [PATCH 3/9] Apply code-format changes --- .../experimental/tle/language/distributed.py | 60 ++++++++----------- python/triton/language/core.py | 3 +- .../tle/dialect/include/IR/VerfiyUtils.h | 2 +- .../TleToLLVM/LocalPointersOpToLLVM.cpp | 51 ++++++++-------- third_party/tle/dialect/lib/IR/Ops.cpp | 6 +- .../tle/dialect/lib/IR/VerfiyUtils.cpp | 10 ++-- third_party/tle/triton_tle.cc | 20 +++---- 7 files changed, 66 insertions(+), 86 deletions(-) diff --git a/python/triton/experimental/tle/language/distributed.py b/python/triton/experimental/tle/language/distributed.py index 7999a3aa8d..cb90a41c7b 100644 --- a/python/triton/experimental/tle/language/distributed.py +++ b/python/triton/experimental/tle/language/distributed.py @@ -952,6 +952,7 @@ def _remote_pointer( return _create_remote_pointers_tensor(tensor, shard_id_tensor, _semantic, dtype=dtype, space=space, offset=offset) + # offset / srcoffset / nelems -> scalar i64 tl.tensor # offset and srcoffset must be >= 0. # nelems must be > 0. @@ -1010,8 +1011,7 @@ def _normalize_put_coop_kind(coopkind) -> int: 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") + 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: @@ -1023,18 +1023,15 @@ 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}") + 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}") + 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}") + 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 @@ -1091,10 +1088,9 @@ def scalar(self): class _node_remote_destination(tl.base_value): - def __init__(self, dst_mem: tl.tensor, src_mem: tl.tensor, comm: tl.tensor, - peer: tl.tensor, dst_offset: tl.tensor, src_offset: tl.tensor, - nelems: tl.tensor, net_idx: tl.tensor, *, dtype: tl.dtype, - elem_bytes: int, coop_kind: int): + def __init__(self, dst_mem: tl.tensor, src_mem: tl.tensor, comm: tl.tensor, peer: tl.tensor, dst_offset: tl.tensor, + src_offset: tl.tensor, nelems: tl.tensor, net_idx: tl.tensor, *, dtype: tl.dtype, elem_bytes: int, + coop_kind: int): super().__init__() self.dst_mem = dst_mem self.src_mem = src_mem @@ -1110,8 +1106,8 @@ def __init__(self, dst_mem: tl.tensor, src_mem: tl.tensor, comm: tl.tensor, @property def type(self): - fields = (self.dst_mem, self.src_mem, self.comm, self.peer, - self.dst_offset, self.src_offset, self.nelems, self.net_idx) + fields = (self.dst_mem, self.src_mem, self.comm, self.peer, self.dst_offset, self.src_offset, self.nelems, + self.net_idx) return _node_remote_destination_type( tuple(field.type for field in fields), self.dtype, @@ -1120,8 +1116,7 @@ def type(self): ) def _flatten_ir(self, handles) -> None: - for field in (self.dst_mem, self.src_mem, self.comm, self.peer, - self.dst_offset, self.src_offset, self.nelems, + for field in (self.dst_mem, self.src_mem, self.comm, self.peer, self.dst_offset, self.src_offset, self.nelems, self.net_idx): field._flatten_ir(handles) @@ -1143,15 +1138,14 @@ def __rsub__(self, other): 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): + def __triton_load__(self, mask, other, boundary_check, padding_option, cache_modifier, eviction_policy, volatile, + flagtree_hints, _semantic=None): self._unsupported_pointer_operation() 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; " - "pass the source context to tle.remote(..., src=...) and call tl.store(remote_dst)") + raise TypeError("tl.store to a node remote destination does not accept a value; " + "pass the source context to tle.remote(..., src=...) and call tl.store(remote_dst)") 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) @@ -1166,8 +1160,7 @@ def __triton_store__(self, value, mask, boundary_check, cache_modifier, eviction builder = _semantic.builder if not hasattr(builder, "create_remote_pointers"): - raise RuntimeError( - "node put requires TLE remote_pointers support in the active Triton build") + raise RuntimeError("node put requires TLE remote_pointers support in the active Triton build") builder.create_remote_pointers( None, None, @@ -1186,8 +1179,7 @@ def __triton_store__(self, value, mask, boundary_check, cache_modifier, eviction return _semantic.tensor(None, tl.void) -def _create_node_remote_destination(dst, shard_id, scope, dtype, offset, src, - srcoffset, nelems, coopkind, netidx, +def _create_node_remote_destination(dst, shard_id, scope, dtype, offset, src, srcoffset, nelems, coopkind, netidx, _semantic) -> _node_remote_destination: if offset is None: raise TypeError('tle.remote(..., space="node") requires offset') @@ -1200,22 +1192,18 @@ def _create_node_remote_destination(dst, shard_id, scope, dtype, offset, src, builder = _semantic.builder if not hasattr(builder, "create_remote_pointers"): - raise RuntimeError( - "node put requires TLE remote_pointers support in the active Triton build") + raise RuntimeError("node put 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) + offset = _normalize_node_i64(offset, "offset", must_be_positive=False, _semantic=_semantic) if srcoffset is None: srcoffset = offset else: - srcoffset = _normalize_node_i64( - srcoffset, "srcoffset", must_be_positive=False, _semantic=_semantic) - nelems = _normalize_node_i64( - nelems, "nelems", must_be_positive=True, _semantic=_semantic) + srcoffset = _normalize_node_i64(srcoffset, "srcoffset", 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_put_coop_kind(coopkind) @@ -1238,6 +1226,7 @@ def _create_node_remote_destination(dst, shard_id, scope, dtype, offset, src, coop_kind=coop_kind, ) + @tl.builtin def remote( tensor=tl.tensor | None, @@ -1297,9 +1286,8 @@ def remote( shard_id = _unwrap_remote_shard_id(shard_id) scope = tl._unwrap_if_constexpr(scope) if space == "node": - return _create_node_remote_destination( - tensor, shard_id, scope, dtype, offset, src, srcoffset, nelems, - coopkind, netidx, _semantic) + return _create_node_remote_destination(tensor, shard_id, scope, dtype, offset, src, srcoffset, nelems, coopkind, + netidx, _semantic) 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: diff --git a/python/triton/language/core.py b/python/triton/language/core.py index 77e491f7f7..d7a54e0f44 100644 --- a/python/triton/language/core.py +++ b/python/triton/language/core.py @@ -2215,7 +2215,8 @@ def store_tensor_descriptor(desc: tensor_descriptor_base, offsets: Sequence[cons @_tensor_member_fn @builtin -def store(pointer, value=_STORE_VALUE_UNSET, 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`. diff --git a/third_party/tle/dialect/include/IR/VerfiyUtils.h b/third_party/tle/dialect/include/IR/VerfiyUtils.h index d268926783..9d694b1ee0 100644 --- a/third_party/tle/dialect/include/IR/VerfiyUtils.h +++ b/third_party/tle/dialect/include/IR/VerfiyUtils.h @@ -39,7 +39,7 @@ 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/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp index ea2602ab84..6dd806f87a 100644 --- a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp +++ b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp @@ -116,14 +116,13 @@ static LLVM::LLVMFuncOp getOrInsertNetFromComm(ModuleOp module, 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); + auto func = + builder.create(module.getLoc(), funcName, funcTy); func.setLinkage(LLVM::Linkage::External); return func; } -static LLVM::LLVMFuncOp getOrInsertNetPut(ModuleOp module, - MLIRContext *ctx) { +static LLVM::LLVMFuncOp getOrInsertNetPut(ModuleOp module, MLIRContext *ctx) { const char *funcName = "flagcxDevNetPutS"; if (auto func = module.lookupSymbol(funcName)) return func; @@ -136,8 +135,8 @@ static LLVM::LLVMFuncOp getOrInsertNetPut(ModuleOp module, 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); + auto func = + builder.create(module.getLoc(), funcName, funcTy); func.setLinkage(LLVM::Linkage::External); return func; } @@ -514,27 +513,25 @@ LogicalResult lowerNodeSpace(Location loc, tle::RemotePointersOp op, 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.getSrcMem()); - Value comm = - rewriter.create(loc, ptrTy, adaptor.getComm()); + Value dstMem = + rewriter.create(loc, ptrTy, adaptor.getDstMem()); + Value srcMem = + rewriter.create(loc, ptrTy, adaptor.getSrcMem()); + Value comm = rewriter.create(loc, ptrTy, adaptor.getComm()); Value dstByteOffset = adaptor.getOffset(); Value srcByteOffset = adaptor.getSrcOffset(); Value byteCount = adaptor.getNelems(); - int64_t elemBytes = - op->getAttrOfType("elem_bytes").getInt(); + int64_t elemBytes = op->getAttrOfType("elem_bytes").getInt(); if (elemBytes != 1) { - Value elemBytesValue = rewriter.create( - loc, elemBytes, 64); - dstByteOffset = rewriter.create( - loc, adaptor.getOffset(), elemBytesValue); - srcByteOffset = rewriter.create( - loc, adaptor.getSrcOffset(), elemBytesValue); - byteCount = rewriter.create( - loc, adaptor.getNelems(), elemBytesValue); + Value elemBytesValue = + rewriter.create(loc, elemBytes, 64); + dstByteOffset = rewriter.create(loc, adaptor.getOffset(), + elemBytesValue); + srcByteOffset = rewriter.create(loc, adaptor.getSrcOffset(), + elemBytesValue); + byteCount = rewriter.create(loc, adaptor.getNelems(), + elemBytesValue); } LLVM::LLVMFuncOp getNet = getOrInsertNetFromComm(module, ctx); @@ -548,11 +545,11 @@ LogicalResult lowerNodeSpace(Location loc, tle::RemotePointersOp op, op->getAttrOfType("put_coop_kind").getInt(); Value coopKind = rewriter.create( loc, i32Ty, rewriter.getI32IntegerAttr(putCoopKind)); - rewriter.create( - loc, TypeRange{}, FlatSymbolRefAttr::get(put), - ValueRange{getNetCall.getResult(), comm, teamKind, adaptor.getShardId(), - dstMem, dstByteOffset, srcMem, srcByteOffset, byteCount, - coopKind}); + rewriter.create(loc, TypeRange{}, FlatSymbolRefAttr::get(put), + ValueRange{getNetCall.getResult(), comm, + teamKind, adaptor.getShardId(), + dstMem, dstByteOffset, srcMem, + srcByteOffset, byteCount, coopKind}); return success(); } diff --git a/third_party/tle/dialect/lib/IR/Ops.cpp b/third_party/tle/dialect/lib/IR/Ops.cpp index 0c2f96eb52..87e492822b 100644 --- a/third_party/tle/dialect/lib/IR/Ops.cpp +++ b/third_party/tle/dialect/lib/IR/Ops.cpp @@ -937,8 +937,7 @@ Speculation::Speculatability RemotePointersOp::getSpeculatability() { LogicalResult RemotePointersOp::verify() { StringRef spaceAttr = getSpace(); - if (spaceAttr != "cluster" && spaceAttr != "device" && - spaceAttr != "node") + if (spaceAttr != "cluster" && spaceAttr != "device" && spaceAttr != "node") return emitOpError() << "expects space to be 'cluster', 'device', or 'node'"; @@ -949,8 +948,7 @@ LogicalResult RemotePointersOp::verify() { return RemotePointers::verifyNodeSpace(*this); auto elemBytesAttr = (*this)->getAttrOfType("elem_bytes"); - auto putCoopKindAttr = - (*this)->getAttrOfType("put_coop_kind"); + auto putCoopKindAttr = (*this)->getAttrOfType("put_coop_kind"); if (getDstMem() || getSrcMem() || getComm() || getSrcOffset() || getNelems() || getNetIdx() || elemBytesAttr || putCoopKindAttr) diff --git a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp index e2a2340252..a5445a38e3 100644 --- a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp +++ b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp @@ -75,8 +75,7 @@ llvm::LogicalResult verifyNodeSpace(RemotePointersOp op) { auto requireOperand = [&](Value value, StringRef name) -> LogicalResult { if (!value) - return op.emitOpError() - << "node space requires " << name << " operand"; + return op.emitOpError() << "node space requires " << name << " operand"; return success(); }; if (failed(requireOperand(op.getDstMem(), "dst_mem")) || @@ -98,12 +97,11 @@ llvm::LogicalResult verifyNodeSpace(RemotePointersOp op) { return op.emitOpError() << "expects put_coop_kind to be THREAD(0), WARP(1), or BLOCK(2)"; - auto verifyNonNegativeConstant = - [&](Value value, StringRef name) -> LogicalResult { + 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 op.emitOpError() << "expects constant " << name << " to be >= 0"; return success(); }; diff --git a/third_party/tle/triton_tle.cc b/third_party/tle/triton_tle.cc index 1406ecd83f..595e8ae27c 100644 --- a/third_party/tle/triton_tle.cc +++ b/third_party/tle/triton_tle.cc @@ -602,12 +602,11 @@ void init_triton_tle_ir(py::module &&m) { .def( "create_remote_pointers", [](TritonOpBuilder &self, std::optional resultTy, - std::optional src, Value shardId, - const std::string &space, std::optional offset, - std::optional dstMem, std::optional srcMem, - std::optional comm, std::optional srcOffset, - std::optional nelems, std::optional netIdx, - std::optional elemBytes, + std::optional src, Value shardId, const std::string &space, + std::optional offset, std::optional dstMem, + std::optional srcMem, std::optional comm, + std::optional srcOffset, std::optional nelems, + std::optional netIdx, std::optional elemBytes, std::optional putCoopKind) -> OpState { auto &builder = self.getBuilder(); static const std::unordered_set valid = { @@ -633,11 +632,10 @@ void init_triton_tle_ir(py::module &&m) { nelems.value_or(Value()), netIdx.value_or(Value()), elemBytesAttr, putCoopKindAttr); }, - py::arg("resultTy"), py::arg("src") = py::none(), - py::arg("shardId"), py::arg("space"), - py::arg("offset") = py::none(), py::arg("dst_mem") = py::none(), - py::arg("src_mem") = py::none(), py::arg("comm") = py::none(), - py::arg("src_offset") = py::none(), + py::arg("resultTy"), py::arg("src") = py::none(), py::arg("shardId"), + py::arg("space"), py::arg("offset") = py::none(), + py::arg("dst_mem") = py::none(), py::arg("src_mem") = py::none(), + py::arg("comm") = py::none(), py::arg("src_offset") = py::none(), py::arg("nelems") = py::none(), py::arg("net_idx") = py::none(), py::arg("elem_bytes") = py::none(), py::arg("put_coop_kind") = py::none()) From 3f30178ae7dd98ff7bac70031c43b83d6901ff7d Mon Sep 17 00:00:00 2001 From: lzllx123 <1803100521@qq.com> Date: Thu, 6 Aug 2026 15:28:26 +0800 Subject: [PATCH 4/9] [TLE]Update how tle.remote is used for communication between nodes --- .../experimental/tle/language/distributed.py | 122 +++++++++++------- third_party/tle/dialect/include/IR/TleOps.td | 3 +- .../tle/dialect/lib/Analysis/AxisInfoExt.cpp | 14 +- .../TleToLLVM/LocalPointersOpToLLVM.cpp | 33 ++--- third_party/tle/dialect/lib/IR/Ops.cpp | 2 +- .../tle/dialect/lib/IR/VerfiyUtils.cpp | 17 +-- .../lib/Transforms/TleSelectEncodings.cpp | 12 +- third_party/tle/triton_tle.cc | 27 ++-- 8 files changed, 137 insertions(+), 93 deletions(-) diff --git a/python/triton/experimental/tle/language/distributed.py b/python/triton/experimental/tle/language/distributed.py index cb90a41c7b..ecc289af01 100644 --- a/python/triton/experimental/tle/language/distributed.py +++ b/python/triton/experimental/tle/language/distributed.py @@ -881,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) @@ -952,9 +954,8 @@ def _remote_pointer( return _create_remote_pointers_tensor(tensor, shard_id_tensor, _semantic, dtype=dtype, space=space, offset=offset) - -# offset / srcoffset / nelems -> scalar i64 tl.tensor -# offset and srcoffset must be >= 0. +# 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) @@ -1048,7 +1049,7 @@ def _parse_node_context(builder, value, label: str, index: int): class _node_remote_destination_type(tl.base_type): def __init__(self, field_types, dtype: tl.dtype, elem_bytes: int, coop_kind: int): - # dst_mem, src_mem, comm, peer, dst_offset, src_offset, nelems, net_idx. + # 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 @@ -1088,16 +1089,17 @@ def scalar(self): class _node_remote_destination(tl.base_value): - def __init__(self, dst_mem: tl.tensor, src_mem: tl.tensor, comm: tl.tensor, peer: tl.tensor, dst_offset: tl.tensor, - src_offset: tl.tensor, nelems: tl.tensor, net_idx: tl.tensor, *, dtype: tl.dtype, elem_bytes: int, - coop_kind: int): + 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.dst_mem = dst_mem self.src_mem = src_mem + self.dst_mem = dst_mem self.comm = comm self.peer = peer - self.dst_offset = dst_offset - self.src_offset = src_offset + self.offset = offset + self.dstoffset = dstoffset self.nelems = nelems self.net_idx = net_idx self.dtype = dtype @@ -1106,8 +1108,8 @@ def __init__(self, dst_mem: tl.tensor, src_mem: tl.tensor, comm: tl.tensor, peer @property def type(self): - fields = (self.dst_mem, self.src_mem, self.comm, self.peer, self.dst_offset, self.src_offset, self.nelems, - self.net_idx) + 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, @@ -1116,7 +1118,8 @@ def type(self): ) def _flatten_ir(self, handles) -> None: - for field in (self.dst_mem, self.src_mem, self.comm, self.peer, self.dst_offset, self.src_offset, self.nelems, + 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) @@ -1144,8 +1147,10 @@ def __triton_load__(self, mask, other, boundary_check, padding_option, cache_mod 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; " - "pass the source context to tle.remote(..., src=...) and call tl.store(remote_dst)") + 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) @@ -1163,14 +1168,13 @@ def __triton_store__(self, value, mask, boundary_check, cache_modifier, eviction raise RuntimeError("node put requires TLE remote_pointers support in the active Triton build") builder.create_remote_pointers( None, - None, + self.src_mem.handle, self.peer.handle, "node", - self.dst_offset.handle, + self.offset.handle, self.dst_mem.handle, - self.src_mem.handle, self.comm.handle, - self.src_offset.handle, + self.dstoffset.handle, self.nelems.handle, self.net_idx.handle, self.elem_bytes, @@ -1179,7 +1183,8 @@ def __triton_store__(self, value, mask, boundary_check, cache_modifier, eviction return _semantic.tensor(None, tl.void) -def _create_node_remote_destination(dst, shard_id, scope, dtype, offset, src, srcoffset, nelems, coopkind, netidx, +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') @@ -1198,27 +1203,31 @@ def _create_node_remote_destination(dst, shard_id, scope, dtype, offset, src, sr 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 srcoffset is None: - srcoffset = offset + offset = _normalize_node_i64( + offset, "offset", must_be_positive=False, _semantic=_semantic) + if dstoffset is None: + dstoffset = offset else: - srcoffset = _normalize_node_i64(srcoffset, "srcoffset", must_be_positive=False, _semantic=_semantic) - nelems = _normalize_node_i64(nelems, "nelems", must_be_positive=True, _semantic=_semantic) + 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_put_coop_kind(coopkind) - if src is None: - src = dst - dst_mem = tl.tensor(_parse_node_context(builder, dst, "dst", 0), tl.int64) src_mem = tl.tensor(_parse_node_context(builder, src, "src", 0), tl.int64) - dst_comm = tl.tensor(_parse_node_context(builder, dst, "dst", 1), 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( - dst_mem, src_mem, - dst_comm, + dst_mem, + comm, peer, offset, - srcoffset, + dstoffset, nelems, net_idx, dtype=dtype, @@ -1235,8 +1244,8 @@ def remote( space: str = "cluster", dtype: tl.dtype = None, offset: int | tl.tensor | None = None, - src=None, - srcoffset: 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, @@ -1258,36 +1267,55 @@ def remote( 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 destination `DistributedRtContext` - and `src` is the source registered-memory context. `src` defaults to - `tensor`. This function returns a store-only destination triggered with + For `space="node"`, `tensor` is the source registered-memory + `DistributedRtContext` and `dst` is the destination context. `dst` + defaults to `tensor`. This function returns a store-only destination with `tl.store(remote_dst)`; node destinations do not accept a store value. `dtype`, `offset`, `nelems`, and `coopkind` must be explicit, while - `srcoffset` defaults to `offset` and `netidx` defaults to zero. The + `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`, - `srcoffset`, and `nelems` are scalar element counts normalized to i64; + `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 put without flush, completion notification, or a remote-visibility guarantee. - `offset` is a scalar element offset relative to the target shard's memory - base address. It is required for `space="node"` and optional for - `space="device"`. The device path 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 == ()). + In node mode, `offset` is relative to the source memory base and + `dstoffset` is relative to the target shard's memory base. 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 == ()). """ 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, shard_id, scope, dtype, offset, src, srcoffset, nelems, coopkind, - netidx, _semantic) + 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: diff --git a/third_party/tle/dialect/include/IR/TleOps.td b/third_party/tle/dialect/include/IR/TleOps.td index 39b62dfdf9..b44d89e12b 100644 --- a/third_party/tle/dialect/include/IR/TleOps.td +++ b/third_party/tle/dialect/include/IR/TleOps.td @@ -364,12 +364,11 @@ def Tle_RemotePointersOp : Tle_Op<"remote_pointers", [ let arguments = (ins Optional:$src, Optional:$dst_mem, - Optional:$src_mem, Optional:$comm, TT_Int:$shard_id, StrAttr:$space, Optional:$offset, - Optional:$src_offset, + Optional:$dst_offset, Optional:$nelems, Optional:$net_idx, OptionalAttr:$elem_bytes, diff --git a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp index 38e89b9a75..6963b1d28f 100644 --- a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp +++ b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp @@ -233,8 +233,10 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { getAxisInfo(Operation *op, ArrayRef *> operands) override { auto remote = dyn_cast(op); - if (!remote || !remote.getResult() || remote.getSpace() == "node" || - operands.empty()) + // RemotePointersOpAxisInfoVisitor 只处理有 src/result 的 + // RemotePointersOp,且 space 不是 "node"。 + if (!remote || !remote.getSrc() || !remote.getResult() || + remote.getSpace() == "node" || operands.empty()) return AxisInfo(); const AxisInfo &baseInfo = operands[0]->getValue(); @@ -286,8 +288,14 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { } bool match(Operation *op) override { + // 只有当 op 同时满足以下四个条件时,TleRemotePointersOpAxisInfoVisitor 才愿意处理它: + // 它是 RemotePointersOp + // 它有 src + // 它有 result + // 它的 space 不是 "node" auto remote = dyn_cast(op); - return remote && remote.getResult() && remote.getSpace() != "node"; + return remote && remote.getSrc() && remote.getResult() && + 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 6dd806f87a..25493a0818 100644 --- a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp +++ b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp @@ -513,25 +513,26 @@ LogicalResult lowerNodeSpace(Location loc, tle::RemotePointersOp op, 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.getSrcMem()); - Value comm = rewriter.create(loc, ptrTy, adaptor.getComm()); - - Value dstByteOffset = adaptor.getOffset(); - Value srcByteOffset = adaptor.getSrcOffset(); + 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); - dstByteOffset = rewriter.create(loc, adaptor.getOffset(), - elemBytesValue); - srcByteOffset = rewriter.create(loc, adaptor.getSrcOffset(), - elemBytesValue); - byteCount = rewriter.create(loc, adaptor.getNelems(), - elemBytesValue); + 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); diff --git a/third_party/tle/dialect/lib/IR/Ops.cpp b/third_party/tle/dialect/lib/IR/Ops.cpp index 87e492822b..b1af76087b 100644 --- a/third_party/tle/dialect/lib/IR/Ops.cpp +++ b/third_party/tle/dialect/lib/IR/Ops.cpp @@ -950,7 +950,7 @@ LogicalResult RemotePointersOp::verify() { auto elemBytesAttr = (*this)->getAttrOfType("elem_bytes"); auto putCoopKindAttr = (*this)->getAttrOfType("put_coop_kind"); - if (getDstMem() || getSrcMem() || getComm() || getSrcOffset() || + if (getDstMem() || getComm() || getDstOffset() || getNelems() || getNetIdx() || elemBytesAttr || putCoopKindAttr) return emitOpError() << "cluster/device space does not accept node put operands or " diff --git a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp index a5445a38e3..eba16ba052 100644 --- a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp +++ b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp @@ -67,9 +67,6 @@ llvm::LogicalResult verifyDeviceSpace(mlir::Value src, mlir::Value result) { } llvm::LogicalResult verifyNodeSpace(RemotePointersOp op) { - if (op.getSrc()) - return op.emitOpError() - << "node space does not accept a pointer source operand"; if (op.getResult()) return op.emitOpError() << "node space must not produce a result"; @@ -78,15 +75,19 @@ llvm::LogicalResult verifyNodeSpace(RemotePointersOp op) { return op.emitOpError() << "node space requires " << name << " operand"; return success(); }; - if (failed(requireOperand(op.getDstMem(), "dst_mem")) || - failed(requireOperand(op.getSrcMem(), "src_mem")) || + 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.getSrcOffset(), "src_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"; @@ -106,8 +107,8 @@ llvm::LogicalResult verifyNodeSpace(RemotePointersOp op) { }; if (failed(verifyNonNegativeConstant(op.getShardId(), "peer")) || - failed(verifyNonNegativeConstant(op.getOffset(), "dst_offset")) || - failed(verifyNonNegativeConstant(op.getSrcOffset(), "src_offset")) || + failed(verifyNonNegativeConstant(op.getOffset(), "offset")) || + failed(verifyNonNegativeConstant(op.getDstOffset(), "dst_offset")) || failed(verifyNonNegativeConstant(op.getNetIdx(), "net_idx"))) return failure(); diff --git a/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp b/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp index 09792d1fb9..e492be2c1e 100644 --- a/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp +++ b/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp @@ -462,8 +462,11 @@ collectConsumerEncodingVotes(Value root, continue; } if (auto remote = dyn_cast(owner)) { - if (Value result = remote.getResult()) - enqueue(result); + // Node mode reuses src but has no pointer result to propagate. + if (remote.getSpace() != "node") { + if (Value result = remote.getResult()) + enqueue(result); + } continue; } } @@ -848,7 +851,8 @@ class SelectEncodingsPass continue; } if (auto remote = dyn_cast(owner)) { - if (!remote.getResult()) + // Node mode reuses src but has no result encoding to update. + if (remote.getSpace() == "node" || !remote.getResult()) continue; auto remoteResultTy = dyn_cast(remote.getResult().getType()); @@ -906,6 +910,8 @@ 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" || !op.getResult() || !op.getSrc()) return; module->setAttr(kTleEnableEncodingRematerializationAttr, diff --git a/third_party/tle/triton_tle.cc b/third_party/tle/triton_tle.cc index 595e8ae27c..593f8ed845 100644 --- a/third_party/tle/triton_tle.cc +++ b/third_party/tle/triton_tle.cc @@ -602,11 +602,12 @@ void init_triton_tle_ir(py::module &&m) { .def( "create_remote_pointers", [](TritonOpBuilder &self, std::optional resultTy, - std::optional src, Value shardId, const std::string &space, - std::optional offset, std::optional dstMem, - std::optional srcMem, std::optional comm, - std::optional srcOffset, std::optional nelems, - std::optional netIdx, std::optional elemBytes, + 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 putCoopKind) -> OpState { auto &builder = self.getBuilder(); static const std::unordered_set valid = { @@ -626,16 +627,16 @@ void init_triton_tle_ir(py::module &&m) { : IntegerAttr(); return self.create( resultTy.value_or(Type()), src.value_or(Value()), - dstMem.value_or(Value()), srcMem.value_or(Value()), - comm.value_or(Value()), shardId, spaceAttr, - offset.value_or(Value()), srcOffset.value_or(Value()), - nelems.value_or(Value()), netIdx.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, putCoopKindAttr); }, - py::arg("resultTy"), py::arg("src") = py::none(), py::arg("shardId"), - py::arg("space"), py::arg("offset") = py::none(), - py::arg("dst_mem") = py::none(), py::arg("src_mem") = py::none(), - py::arg("comm") = py::none(), py::arg("src_offset") = py::none(), + py::arg("resultTy"), py::arg("src") = py::none(), + py::arg("shardId"), 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("put_coop_kind") = py::none()) From 86c5e25916be114451b1a3b9aea784603e109f7f Mon Sep 17 00:00:00 2001 From: flagtree-bot Date: Thu, 6 Aug 2026 07:33:26 +0000 Subject: [PATCH 5/9] Apply code-format changes --- .../experimental/tle/language/distributed.py | 50 ++++++++----------- .../tle/dialect/lib/Analysis/AxisInfoExt.cpp | 8 ++- .../TleToLLVM/LocalPointersOpToLLVM.cpp | 27 +++++----- third_party/tle/dialect/lib/IR/Ops.cpp | 4 +- third_party/tle/triton_tle.cc | 21 ++++---- 5 files changed, 47 insertions(+), 63 deletions(-) diff --git a/python/triton/experimental/tle/language/distributed.py b/python/triton/experimental/tle/language/distributed.py index ecc289af01..16e8337824 100644 --- a/python/triton/experimental/tle/language/distributed.py +++ b/python/triton/experimental/tle/language/distributed.py @@ -954,6 +954,7 @@ 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. @@ -1089,10 +1090,9 @@ def scalar(self): 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): + 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 @@ -1108,8 +1108,8 @@ def __init__(self, src_mem: tl.tensor, dst_mem: tl.tensor, comm: tl.tensor, @property def type(self): - fields = (self.src_mem, self.dst_mem, self.comm, self.peer, - self.offset, self.dstoffset, self.nelems, self.net_idx) + 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, @@ -1118,8 +1118,7 @@ def type(self): ) 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, + 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) @@ -1147,10 +1146,9 @@ def __triton_load__(self, mask, other, boundary_check, padding_option, cache_mod 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") + 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) @@ -1183,8 +1181,7 @@ def __triton_store__(self, value, mask, boundary_check, cache_modifier, eviction return _semantic.tensor(None, tl.void) -def _create_node_remote_destination(src, dst, shard_id, scope, dtype, offset, - dstoffset, nelems, coopkind, netidx, +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') @@ -1203,15 +1200,12 @@ def _create_node_remote_destination(src, dst, shard_id, scope, dtype, offset, 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) + 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) + 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_put_coop_kind(coopkind) @@ -1296,26 +1290,22 @@ def remote( 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}") + 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) + 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)) + 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)}') + 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: diff --git a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp index 6963b1d28f..9503753fd3 100644 --- a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp +++ b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp @@ -288,11 +288,9 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { } bool match(Operation *op) override { - // 只有当 op 同时满足以下四个条件时,TleRemotePointersOpAxisInfoVisitor 才愿意处理它: - // 它是 RemotePointersOp - // 它有 src - // 它有 result - // 它的 space 不是 "node" + // 只有当 op 同时满足以下四个条件时,TleRemotePointersOpAxisInfoVisitor + // 才愿意处理它: 它是 RemotePointersOp 它有 src 它有 result 它的 space 不是 + // "node" auto remote = dyn_cast(op); return remote && remote.getSrc() && remote.getResult() && 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 25493a0818..ee114700dd 100644 --- a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp +++ b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp @@ -513,26 +513,25 @@ LogicalResult lowerNodeSpace(Location loc, tle::RemotePointersOp op, 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 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); + 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); diff --git a/third_party/tle/dialect/lib/IR/Ops.cpp b/third_party/tle/dialect/lib/IR/Ops.cpp index b1af76087b..736132f4bc 100644 --- a/third_party/tle/dialect/lib/IR/Ops.cpp +++ b/third_party/tle/dialect/lib/IR/Ops.cpp @@ -950,8 +950,8 @@ LogicalResult RemotePointersOp::verify() { auto elemBytesAttr = (*this)->getAttrOfType("elem_bytes"); auto putCoopKindAttr = (*this)->getAttrOfType("put_coop_kind"); - if (getDstMem() || getComm() || getDstOffset() || - getNelems() || getNetIdx() || elemBytesAttr || putCoopKindAttr) + if (getDstMem() || getComm() || getDstOffset() || getNelems() || + getNetIdx() || elemBytesAttr || putCoopKindAttr) return emitOpError() << "cluster/device space does not accept node put operands or " "attributes"; diff --git a/third_party/tle/triton_tle.cc b/third_party/tle/triton_tle.cc index 593f8ed845..d48459ad90 100644 --- a/third_party/tle/triton_tle.cc +++ b/third_party/tle/triton_tle.cc @@ -602,10 +602,9 @@ void init_triton_tle_ir(py::module &&m) { .def( "create_remote_pointers", [](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 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 putCoopKind) -> OpState { @@ -630,15 +629,13 @@ void init_triton_tle_ir(py::module &&m) { 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, putCoopKindAttr); + netIdx.value_or(Value()), elemBytesAttr, putCoopKindAttr); }, - py::arg("resultTy"), py::arg("src") = py::none(), - py::arg("shardId"), 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("resultTy"), py::arg("src") = py::none(), py::arg("shardId"), + 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("put_coop_kind") = py::none()) .def("get_device_id", [](TritonOpBuilder &self, Type resultTy, From 13f88b282ba978096dd08500d0a26cf232562a48 Mon Sep 17 00:00:00 2001 From: lzllx123 <1803100521@qq.com> Date: Thu, 6 Aug 2026 15:37:23 +0800 Subject: [PATCH 6/9] [TLE]Update how tle.remote is used for communication between nodes2 --- third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp | 5 ----- 1 file changed, 5 deletions(-) diff --git a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp index 9503753fd3..a4db55bc10 100644 --- a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp +++ b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp @@ -233,8 +233,6 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { getAxisInfo(Operation *op, ArrayRef *> operands) override { auto remote = dyn_cast(op); - // RemotePointersOpAxisInfoVisitor 只处理有 src/result 的 - // RemotePointersOp,且 space 不是 "node"。 if (!remote || !remote.getSrc() || !remote.getResult() || remote.getSpace() == "node" || operands.empty()) return AxisInfo(); @@ -288,9 +286,6 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { } bool match(Operation *op) override { - // 只有当 op 同时满足以下四个条件时,TleRemotePointersOpAxisInfoVisitor - // 才愿意处理它: 它是 RemotePointersOp 它有 src 它有 result 它的 space 不是 - // "node" auto remote = dyn_cast(op); return remote && remote.getSrc() && remote.getResult() && remote.getSpace() != "node"; From 770139eff85dc291fbc6d50ec6c5a4bf3111e066 Mon Sep 17 00:00:00 2001 From: lzllx123 <1803100521@qq.com> Date: Thu, 6 Aug 2026 17:32:15 +0800 Subject: [PATCH 7/9] [TLE]Update how tle.remote is used for communication between nodes3 --- third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp | 6 ++---- .../tle/dialect/lib/Transforms/TleSelectEncodings.cpp | 10 ++++------ 2 files changed, 6 insertions(+), 10 deletions(-) diff --git a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp index a4db55bc10..1eb2dcfeb0 100644 --- a/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp +++ b/third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp @@ -233,8 +233,7 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { getAxisInfo(Operation *op, ArrayRef *> operands) override { auto remote = dyn_cast(op); - if (!remote || !remote.getSrc() || !remote.getResult() || - remote.getSpace() == "node" || operands.empty()) + if (!remote || remote.getSpace() == "node" || operands.empty()) return AxisInfo(); const AxisInfo &baseInfo = operands[0]->getValue(); @@ -287,8 +286,7 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor { bool match(Operation *op) override { auto remote = dyn_cast(op); - return remote && remote.getSrc() && remote.getResult() && - remote.getSpace() != "node"; + return remote && remote.getSpace() != "node"; } }; diff --git a/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp b/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp index e492be2c1e..a083235ab6 100644 --- a/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp +++ b/third_party/tle/dialect/lib/Transforms/TleSelectEncodings.cpp @@ -463,10 +463,8 @@ collectConsumerEncodingVotes(Value root, } if (auto remote = dyn_cast(owner)) { // Node mode reuses src but has no pointer result to propagate. - if (remote.getSpace() != "node") { - if (Value result = remote.getResult()) - enqueue(result); - } + if (remote.getSpace() != "node") + enqueue(remote.getResult()); continue; } } @@ -852,7 +850,7 @@ class SelectEncodingsPass } if (auto remote = dyn_cast(owner)) { // Node mode reuses src but has no result encoding to update. - if (remote.getSpace() == "node" || !remote.getResult()) + if (remote.getSpace() == "node") continue; auto remoteResultTy = dyn_cast(remote.getResult().getType()); @@ -912,7 +910,7 @@ class SelectEncodingsPass 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" || !op.getResult() || !op.getSrc()) + if (op.getSpace() == "node") return; module->setAttr(kTleEnableEncodingRematerializationAttr, UnitAttr::get(module.getContext())); From 3ae811f973ec17b51f4d152f09e9a8cc3b560e22 Mon Sep 17 00:00:00 2001 From: lzllx123 <1803100521@qq.com> Date: Thu, 13 Aug 2026 10:11:52 +0800 Subject: [PATCH 8/9] [TLE]Add a GET route for node-to-node communication in tle.remote --- .../experimental/tle/language/distributed.py | 76 +++++++++++++++---- third_party/tle/dialect/include/IR/TleOps.td | 3 +- .../TleToLLVM/LocalPointersOpToLLVM.cpp | 47 +++++++++--- third_party/tle/dialect/lib/IR/Ops.cpp | 7 +- .../tle/dialect/lib/IR/VerfiyUtils.cpp | 14 +++- third_party/tle/triton_tle.cc | 18 +++-- 6 files changed, 128 insertions(+), 37 deletions(-) diff --git a/python/triton/experimental/tle/language/distributed.py b/python/triton/experimental/tle/language/distributed.py index 16e8337824..eb5c0763f5 100644 --- a/python/triton/experimental/tle/language/distributed.py +++ b/python/triton/experimental/tle/language/distributed.py @@ -1008,7 +1008,7 @@ def _normalize_node_peer(shard_id, scope, _semantic) -> tl.tensor: return _normalize_runtime_remote_shard_id_tensor(shard_id) -def _normalize_put_coop_kind(coopkind) -> int: +def _normalize_coopkind(coopkind) -> int: coopkind = tl._unwrap_if_constexpr(coopkind) if isinstance(coopkind, GroupKind): coopkind = coopkind.value @@ -1085,7 +1085,7 @@ def __str__(self) -> str: @property def scalar(self): - raise ValueError('tle.remote(..., space="node") destinations only support tl.store') + raise ValueError('tle.remote(..., space="node") destinations only support tl.load/tl.store') class _node_remote_destination(tl.base_value): @@ -1123,7 +1123,7 @@ def _flatten_ir(self, handles) -> None: field._flatten_ir(handles) def _unsupported_pointer_operation(self): - raise ValueError('tle.remote(..., space="node") destinations only support tl.store') + raise ValueError('tle.remote(..., space="node") destinations only support tl.load/tl.store') def __add__(self, other): self._unsupported_pointer_operation() @@ -1142,7 +1142,50 @@ def __getitem__(self, index): def __triton_load__(self, mask, other, boundary_check, padding_option, cache_modifier, eviction_policy, volatile, flagtree_hints, _semantic=None): - self._unsupported_pointer_operation() + 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: @@ -1177,6 +1220,7 @@ def __triton_store__(self, value, mask, boundary_check, cache_modifier, eviction self.net_idx.handle, self.elem_bytes, self.coop_kind, + "put", ) return _semantic.tensor(None, tl.void) @@ -1194,7 +1238,7 @@ def _create_node_remote_destination(src, dst, shard_id, scope, dtype, offset, ds builder = _semantic.builder if not hasattr(builder, "create_remote_pointers"): - raise RuntimeError("node put requires TLE remote_pointers support in the active Triton build") + 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) @@ -1207,7 +1251,7 @@ def _create_node_remote_destination(src, dst, shard_id, scope, dtype, offset, ds 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_put_coop_kind(coopkind) + coop_kind = _normalize_coopkind(coopkind) src_mem = tl.tensor(_parse_node_context(builder, src, "src", 0), tl.int64) if dst is None: @@ -1253,8 +1297,8 @@ 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 store-only remote - destination consumed by `tl.store`. + - 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. For cluster/device pointer paths, when `scope` is provided, launch cluster @@ -1263,8 +1307,10 @@ def remote( For `space="node"`, `tensor` is the source registered-memory `DistributedRtContext` and `dst` is the destination context. `dst` - defaults to `tensor`. This function returns a store-only destination with - `tl.store(remote_dst)`; node destinations do not accept a store value. + 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`, @@ -1275,12 +1321,14 @@ def remote( 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 put without flush, completion - notification, or a remote-visibility guarantee. + 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 target shard's memory base. For device - mode, `offset` is the remote-memory offset. It is required for + `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 diff --git a/third_party/tle/dialect/include/IR/TleOps.td b/third_party/tle/dialect/include/IR/TleOps.td index b44d89e12b..01b6fa7fec 100644 --- a/third_party/tle/dialect/include/IR/TleOps.td +++ b/third_party/tle/dialect/include/IR/TleOps.td @@ -372,7 +372,8 @@ def Tle_RemotePointersOp : Tle_Op<"remote_pointers", [ Optional:$nelems, Optional:$net_idx, OptionalAttr:$elem_bytes, - OptionalAttr:$put_coop_kind + OptionalAttr:$coopkind, + OptionalAttr:$transfer_kind ); let results = (outs Optional:$result); diff --git a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp index ee114700dd..153a910af7 100644 --- a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp +++ b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp @@ -141,6 +141,25 @@ static LLVM::LLVMFuncOp getOrInsertNetPut(ModuleOp module, MLIRContext *ctx) { 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, @@ -535,21 +554,31 @@ LogicalResult lowerNodeSpace(Location loc, tle::RemotePointersOp op, } LLVM::LLVMFuncOp getNet = getOrInsertNetFromComm(module, ctx); - LLVM::LLVMFuncOp put = getOrInsertNetPut(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 putCoopKind = - op->getAttrOfType("put_coop_kind").getInt(); + int64_t coopKindValue = op.getCoopkindAttr().getInt(); Value coopKind = rewriter.create( - loc, i32Ty, rewriter.getI32IntegerAttr(putCoopKind)); - rewriter.create(loc, TypeRange{}, FlatSymbolRefAttr::get(put), - ValueRange{getNetCall.getResult(), comm, - teamKind, adaptor.getShardId(), - dstMem, dstByteOffset, srcMem, - srcByteOffset, byteCount, coopKind}); + 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(); } diff --git a/third_party/tle/dialect/lib/IR/Ops.cpp b/third_party/tle/dialect/lib/IR/Ops.cpp index 736132f4bc..f64aaea0fe 100644 --- a/third_party/tle/dialect/lib/IR/Ops.cpp +++ b/third_party/tle/dialect/lib/IR/Ops.cpp @@ -948,12 +948,13 @@ LogicalResult RemotePointersOp::verify() { return RemotePointers::verifyNodeSpace(*this); auto elemBytesAttr = (*this)->getAttrOfType("elem_bytes"); - auto putCoopKindAttr = (*this)->getAttrOfType("put_coop_kind"); + auto coopKindAttr = getCoopkindAttr(); + auto transferKindAttr = (*this)->getAttrOfType("transfer_kind"); if (getDstMem() || getComm() || getDstOffset() || getNelems() || - getNetIdx() || elemBytesAttr || putCoopKindAttr) + getNetIdx() || elemBytesAttr || coopKindAttr || transferKindAttr) return emitOpError() - << "cluster/device space does not accept node put operands or " + << "cluster/device space does not accept node transfer operands or " "attributes"; if (!getResult()) return emitOpError() diff --git a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp index eba16ba052..911c151a0f 100644 --- a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp +++ b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp @@ -70,6 +70,12 @@ 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"; @@ -92,11 +98,11 @@ llvm::LogicalResult verifyNodeSpace(RemotePointersOp op) { if (!elemBytesAttr || elemBytesAttr.getInt() <= 0) return op.emitOpError() << "expects elem_bytes to be > 0"; - auto putCoopKindAttr = op->getAttrOfType("put_coop_kind"); - if (!putCoopKindAttr || putCoopKindAttr.getInt() < 0 || - putCoopKindAttr.getInt() > 2) + auto coopKindAttr = op.getCoopkindAttr(); + if (!coopKindAttr || coopKindAttr.getInt() < 0 || + coopKindAttr.getInt() > 2) return op.emitOpError() - << "expects put_coop_kind to be THREAD(0), WARP(1), or BLOCK(2)"; + << "expects coopkind to be THREAD(0), WARP(1), or BLOCK(2)"; auto verifyNonNegativeConstant = [&](Value value, StringRef name) -> LogicalResult { diff --git a/third_party/tle/triton_tle.cc b/third_party/tle/triton_tle.cc index d48459ad90..1558151d12 100644 --- a/third_party/tle/triton_tle.cc +++ b/third_party/tle/triton_tle.cc @@ -607,7 +607,8 @@ void init_triton_tle_ir(py::module &&m) { std::optional comm, std::optional dstOffset, std::optional nelems, std::optional netIdx, std::optional elemBytes, - std::optional putCoopKind) -> OpState { + std::optional coopKind, + std::optional transferKind) -> OpState { auto &builder = self.getBuilder(); static const std::unordered_set valid = { "cluster", "device", "node"}; @@ -621,22 +622,27 @@ void init_triton_tle_ir(py::module &&m) { IntegerAttr elemBytesAttr = elemBytes ? builder.getI64IntegerAttr(*elemBytes) : IntegerAttr(); - IntegerAttr putCoopKindAttr = - putCoopKind ? builder.getI32IntegerAttr(*putCoopKind) - : IntegerAttr(); + IntegerAttr coopKindAttr = + coopKind ? builder.getI32IntegerAttr(*coopKind) + : IntegerAttr(); + StringAttr transferKindAttr = + transferKind ? builder.getStringAttr(*transferKind) + : StringAttr(); return self.create( 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, putCoopKindAttr); + 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("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("put_coop_kind") = 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 { From 2690155a9484fde55f6804e7de5345221967589f Mon Sep 17 00:00:00 2001 From: flagtree-bot Date: Thu, 13 Aug 2026 02:15:53 +0000 Subject: [PATCH 9/9] Apply code-format changes --- .../TleToLLVM/LocalPointersOpToLLVM.cpp | 15 +++++++-------- third_party/tle/dialect/lib/IR/VerfiyUtils.cpp | 8 +++----- third_party/tle/triton_tle.cc | 6 ++---- 3 files changed, 12 insertions(+), 17 deletions(-) diff --git a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp index 153a910af7..7b0450780b 100644 --- a/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp +++ b/third_party/tle/dialect/lib/Conversion/TleToLLVM/LocalPointersOpToLLVM.cpp @@ -562,22 +562,21 @@ LogicalResult lowerNodeSpace(Location loc, tle::RemotePointersOp op, int64_t coopKindValue = op.getCoopkindAttr().getInt(); Value coopKind = rewriter.create( loc, i32Ty, rewriter.getI32IntegerAttr(coopKindValue)); - auto transferKind = - op->getAttrOfType("transfer_kind").getValue(); + 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}); + 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}); + ValueRange{getNetCall.getResult(), comm, teamKind, adaptor.getShardId(), + srcMem, srcByteOffset, dstMem, dstByteOffset, byteCount, + coopKind}); } return success(); } diff --git a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp index 911c151a0f..5e32085487 100644 --- a/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp +++ b/third_party/tle/dialect/lib/IR/VerfiyUtils.cpp @@ -71,9 +71,8 @@ llvm::LogicalResult verifyNodeSpace(RemotePointersOp op) { return op.emitOpError() << "node space must not produce a result"; auto transferKindAttr = op->getAttrOfType("transfer_kind"); - if (!transferKindAttr || - (transferKindAttr.getValue() != "put" && - transferKindAttr.getValue() != "get")) + 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 { @@ -99,8 +98,7 @@ llvm::LogicalResult verifyNodeSpace(RemotePointersOp op) { return op.emitOpError() << "expects elem_bytes to be > 0"; auto coopKindAttr = op.getCoopkindAttr(); - if (!coopKindAttr || coopKindAttr.getInt() < 0 || - coopKindAttr.getInt() > 2) + if (!coopKindAttr || coopKindAttr.getInt() < 0 || coopKindAttr.getInt() > 2) return op.emitOpError() << "expects coopkind to be THREAD(0), WARP(1), or BLOCK(2)"; diff --git a/third_party/tle/triton_tle.cc b/third_party/tle/triton_tle.cc index 1558151d12..4fc7de9cfc 100644 --- a/third_party/tle/triton_tle.cc +++ b/third_party/tle/triton_tle.cc @@ -606,8 +606,7 @@ void init_triton_tle_ir(py::module &&m) { 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 elemBytes, std::optional coopKind, std::optional transferKind) -> OpState { auto &builder = self.getBuilder(); static const std::unordered_set valid = { @@ -623,8 +622,7 @@ void init_triton_tle_ir(py::module &&m) { elemBytes ? builder.getI64IntegerAttr(*elemBytes) : IntegerAttr(); IntegerAttr coopKindAttr = - coopKind ? builder.getI32IntegerAttr(*coopKind) - : IntegerAttr(); + coopKind ? builder.getI32IntegerAttr(*coopKind) : IntegerAttr(); StringAttr transferKindAttr = transferKind ? builder.getStringAttr(*transferKind) : StringAttr();