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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
410 changes: 392 additions & 18 deletions python/triton/experimental/tle/language/distributed.py

Large diffs are not rendered by default.

28 changes: 24 additions & 4 deletions python/triton/language/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
T = TypeVar('T')

TRITON_BUILTIN = "__triton_builtin__"
_STORE_VALUE_UNSET = object()

PropagateNan = ir.PROPAGATE_NAN

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -2209,7 +2215,8 @@ def store_tensor_descriptor(desc: tensor_descriptor_base, offsets: Sequence[cons

@_tensor_member_fn
@builtin
def store(pointer, value, mask=None, boundary_check=(), cache_modifier="", eviction_policy="", _semantic=None):
def store(pointer, value=_STORE_VALUE_UNSET, mask=None, boundary_check=(), cache_modifier="", eviction_policy="",
_semantic=None):
"""
Store a tensor of data into memory locations defined by `pointer`.

Expand All @@ -2233,6 +2240,9 @@ def store(pointer, value, mask=None, boundary_check=(), cache_modifier="", evict

`value` is implicitly broadcast to `pointer.shape` and typecast to `pointer.dtype.element_ty`.

Experimental store-only destinations may omit `value`; ordinary pointers
still require it.

:param pointer: The memory location where the elements of `value` are stored
:type pointer: `triton.PointerType`, or block of `dtype=triton.PointerType`
:param value: The tensor of elements to be stored
Expand All @@ -2248,13 +2258,23 @@ def store(pointer, value, mask=None, boundary_check=(), cache_modifier="", evict
:param eviction_policy: changes eviction policy in NVIDIA PTX
:type eviction_policy: str, optional, should be one of {"", "evict_first", "evict_last"}
"""
mask = _unwrap_if_constexpr(mask)
cache_modifier = _unwrap_if_constexpr(cache_modifier)
eviction_policy = _unwrap_if_constexpr(eviction_policy)

# Experimental pointer-like destinations can intercept stores before
# `value` is validated or converted to a tensor.
custom_store = getattr(pointer, "__triton_store__", None)
if custom_store is not None:
return custom_store(value, mask, boundary_check, cache_modifier, eviction_policy, _semantic=_semantic)

if value is _STORE_VALUE_UNSET:
raise TypeError("tl.store() missing required argument 'value' for an ordinary pointer")

# `value` can be constexpr
value = _semantic.to_tensor(value)
mask = _unwrap_if_constexpr(mask)
if mask is not None:
mask = _semantic.to_tensor(mask)
cache_modifier = _unwrap_if_constexpr(cache_modifier)
eviction_policy = _unwrap_if_constexpr(eviction_policy)
return _semantic.store(pointer, value, mask, boundary_check, cache_modifier, eviction_policy)


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -31,7 +31,6 @@ namespace mlir::triton::tle {
void populateLocalPointersOpToLLVMPatterns(
mlir::LLVMTypeConverter &typeConverter, const TargetInfoBase &targetInfo,
RewritePatternSet &patterns, PatternBenefit benefit);

void populateRemotePointersOpToLLVMPatterns(
mlir::LLVMTypeConverter &typeConverter, const TargetInfoBase &targetInfo,
RewritePatternSet &patterns, PatternBenefit benefit);
Expand Down
22 changes: 19 additions & 3 deletions third_party/tle/dialect/include/IR/TleOps.td
Original file line number Diff line number Diff line change
Expand Up @@ -356,17 +356,33 @@ def Tle_DistributedBarrierOp : Tle_Op<"distributed_barrier",
let hasVerifier = 1;
}

def Tle_RemotePointersOp : Tle_Op<"remote_pointers", [Pure, AttrSizedOperandSegments]> {
def Tle_RemotePointersOp : Tle_Op<"remote_pointers", [
AttrSizedOperandSegments,
DeclareOpInterfaceMethods<MemoryEffectsOpInterface>,
ConditionallySpeculatable
]> {
let arguments = (ins
Optional<Tle_LocalPointerResultType>:$src,
Optional<I64>:$dst_mem,
Optional<I64>:$comm,
TT_Int:$shard_id,
StrAttr:$space,
Optional<TT_IntLike>:$offset
Optional<TT_IntLike>:$offset,
Optional<I64>:$dst_offset,
Optional<I64>:$nelems,
Optional<I32>:$net_idx,
OptionalAttr<I64Attr>:$elem_bytes,
OptionalAttr<I32Attr>:$coopkind,
OptionalAttr<StrAttr>:$transfer_kind
);

let results = (outs Tle_LocalPointerResultType:$result);
let results = (outs Optional<Tle_LocalPointerResultType>:$result);
let extraClassDeclaration = [{
Speculation::Speculatability getSpeculatability();
}];
let hasVerifier = 1;
}

def Tle_GetNumPesOp : Tle_Op<"get_num_pes"> {
let summary = "Get nume pes";
let arguments = (ins
Expand Down
3 changes: 2 additions & 1 deletion third_party/tle/dialect/include/IR/VerfiyUtils.h
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,8 @@ namespace mlir::triton::tle {

namespace RemotePointers {
llvm::LogicalResult verifyDeviceSpace(mlir::Value src, mlir::Value result);
}
llvm::LogicalResult verifyNodeSpace(RemotePointersOp op);
} // namespace RemotePointers

namespace DistributedBarrier {
llvm::LogicalResult verifyDeviceSpace(mlir::Operation *op, mlir::Value src);
Expand Down
5 changes: 3 additions & 2 deletions third_party/tle/dialect/lib/Analysis/AxisInfoExt.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -233,7 +233,7 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor {
getAxisInfo(Operation *op,
ArrayRef<const dataflow::Lattice<AxisInfo> *> operands) override {
auto remote = dyn_cast<triton::tle::RemotePointersOp>(op);
if (!remote || operands.empty())
if (!remote || remote.getSpace() == "node" || operands.empty())
return AxisInfo();

const AxisInfo &baseInfo = operands[0]->getValue();
Expand Down Expand Up @@ -285,7 +285,8 @@ class TleRemotePointersOpAxisInfoVisitor final : public AxisInfoVisitor {
}

bool match(Operation *op) override {
return isa<triton::tle::RemotePointersOp>(op);
auto remote = dyn_cast<triton::tle::RemotePointersOp>(op);
return remote && remote.getSpace() != "node";
}
};

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,60 @@ static LLVM::LLVMFuncOp getOrInsertGetPeerPointer(ModuleOp module,
return func;
}

static LLVM::LLVMFuncOp getOrInsertNetFromComm(ModuleOp module,
MLIRContext *ctx) {
const char *funcName = "flagcxDevNetGetFromCommS";
if (auto func = module.lookupSymbol<LLVM::LLVMFuncOp>(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<LLVM::LLVMFuncOp>(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<LLVM::LLVMFuncOp>(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<Type> 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<LLVM::LLVMFuncOp>(module.getLoc(), funcName, funcTy);
func.setLinkage(LLVM::Linkage::External);
return func;
}

static LLVM::LLVMFuncOp getOrInsertNetGet(ModuleOp module, MLIRContext *ctx) {
const char *funcName = "flagcxDevNetGetS";
if (auto func = module.lookupSymbol<LLVM::LLVMFuncOp>(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<Type> 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<LLVM::LLVMFuncOp>(module.getLoc(), funcName, funcTy);
func.setLinkage(LLVM::Linkage::External);
return func;
}

struct LocalPointersOpConversion
: public ConvertOpToLLVMPattern<tle::LocalPointersOp> {
LocalPointersOpConversion(LLVMTypeConverter &typeConverter,
Expand Down Expand Up @@ -468,11 +522,63 @@ LogicalResult lowerDeviceSpace(Location loc, Value mem_ptr,
return success();
}

LogicalResult lowerNodeSpace(Location loc, ValueRange srcElems,
ValueRange shardElems,
ConversionPatternRewriter &rewriter,
SmallVectorImpl<Value> &resultPtrs) {
return failure(); // Not implemented yet
LogicalResult lowerNodeSpace(Location loc, tle::RemotePointersOp op,
tle::RemotePointersOp::Adaptor adaptor,
ConversionPatternRewriter &rewriter) {
ModuleOp module = op->getParentOfType<ModuleOp>();
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<LLVM::IntToPtrOp>(loc, ptrTy, adaptor.getDstMem());
Value srcMem =
rewriter.create<LLVM::IntToPtrOp>(loc, ptrTy, adaptor.getSrc());
Value comm = rewriter.create<LLVM::IntToPtrOp>(loc, ptrTy, adaptor.getComm());

Value srcByteOffset = adaptor.getOffset();
Value dstByteOffset = adaptor.getDstOffset();
Value byteCount = adaptor.getNelems();
int64_t elemBytes = op->getAttrOfType<IntegerAttr>("elem_bytes").getInt();
if (elemBytes != 1) {
Value elemBytesValue =
rewriter.create<arith::ConstantIntOp>(loc, elemBytes, 64);
srcByteOffset = rewriter.create<arith::MulIOp>(loc, adaptor.getOffset(),
elemBytesValue);
dstByteOffset = rewriter.create<arith::MulIOp>(loc, adaptor.getDstOffset(),
elemBytesValue);
byteCount = rewriter.create<arith::MulIOp>(loc, adaptor.getNelems(),
elemBytesValue);
}

LLVM::LLVMFuncOp getNet = getOrInsertNetFromComm(module, ctx);
auto getNetCall = rewriter.create<LLVM::CallOp>(
loc, TypeRange{ptrTy}, FlatSymbolRefAttr::get(getNet),
ValueRange{comm, adaptor.getNetIdx()});
Value teamKind = rewriter.create<LLVM::ConstantOp>(
loc, i32Ty, rewriter.getI32IntegerAttr(2));
int64_t coopKindValue = op.getCoopkindAttr().getInt();
Value coopKind = rewriter.create<LLVM::ConstantOp>(
loc, i32Ty, rewriter.getI32IntegerAttr(coopKindValue));
auto transferKind = op->getAttrOfType<StringAttr>("transfer_kind").getValue();
if (transferKind == "put") {
LLVM::LLVMFuncOp put = getOrInsertNetPut(module, ctx);
rewriter.create<LLVM::CallOp>(
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<LLVM::CallOp>(
loc, TypeRange{}, FlatSymbolRefAttr::get(get),
ValueRange{getNetCall.getResult(), comm, teamKind, adaptor.getShardId(),
srcMem, srcByteOffset, dstMem, dstByteOffset, byteCount,
coopKind});
}
return success();
}

Value getDistDevicePtr(tle::RemotePointersOp op, SmallVector<Value> &srcElems) {
Expand Down Expand Up @@ -500,9 +606,15 @@ struct RemotePointersOpConversion
return rewriter.notifyMatchFailure(op, msg);
};

SmallVector<Value> 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<Value> srcElems;
if (auto src = adaptor.getSrc())
srcElems = unpackLLElements(loc, adaptor.getSrc(), rewriter);

Expand All @@ -527,6 +639,7 @@ struct RemotePointersOpConversion

SmallVector<Value> mappedPtrs;
auto mem = getDistDevicePtr(op, srcElems);
auto resultType = op.getResult().getType();

if (space == "cluster") {
if (failed(lowerClusterSpace(loc, srcElems, shardElems, rewriter,
Expand All @@ -536,7 +649,7 @@ struct RemotePointersOpConversion
} else if (space == "device") {
int elemBytes = 1;
if (offsetVal) {
Type resultPointeeTy = getRemotePointeeType(op.getType());
Type resultPointeeTy = getRemotePointeeType(resultType);
if (!resultPointeeTy)
return reportFailure("result must be tt.ptr or tensor<tt.ptr>");
auto elemBits = getScalarBitWidth(resultPointeeTy);
Expand All @@ -549,15 +662,11 @@ struct RemotePointersOpConversion
rewriter, mappedPtrs))) {
return rewriter.notifyMatchFailure(op, "device lowering failed");
}
} else if (space == "node") {
if (failed(lowerNodeSpace(loc, mem, shardElems, rewriter, mappedPtrs))) {
return rewriter.notifyMatchFailure(op, "node lowering failed");
}
} else {
return reportFailure("unsupported remote space: " + space.str());
}
Value packed =
packLLElements(loc, typeConverter, mappedPtrs, rewriter, op.getType());
packLLElements(loc, typeConverter, mappedPtrs, rewriter, resultType);
rewriter.replaceOp(op, packed);
return success();
}
Expand Down
Loading
Loading