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
2 changes: 1 addition & 1 deletion include/triton/Dialect/Triton/IR/TritonOps.td
Original file line number Diff line number Diff line change
Expand Up @@ -202,7 +202,7 @@ def TT_AddPtrOp : TT_Op<"addptr",
let results = (outs TT_PtrLike:$result);

let assemblyFormat = "$ptr `,` $offset attr-dict `:` type($result) `,` type($offset)";
// let hasFolder = 1;
let hasFolder = 1;

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

need to fix conflicts first

}

def TT_AdvanceOp : TT_Op<"advance",
Expand Down
14 changes: 7 additions & 7 deletions lib/Dialect/Triton/IR/Ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1033,13 +1033,13 @@ void MakeTensorPtrOp::build(OpBuilder &builder, OperationState &state,
}

//-- AddPtrOp --
// OpFoldResult AddPtrOp::fold(FoldAdaptor adaptor) {
// // addptr(ptr, 0) -> ptr
// if (matchPattern(adaptor.getOffset(), m_Zero())) {
// return getPtr();
// }
// return {};
// }
OpFoldResult AddPtrOp::fold(FoldAdaptor adaptor) {
// addptr(ptr, 0) -> ptr
if (matchPattern(adaptor.getOffset(), m_Zero())) {
return getPtr();
}
return {};
}

//-- AdvanceOp --
OpFoldResult AdvanceOp::fold(FoldAdaptor adaptor) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,10 @@ class BlockDataParser {
ConversionPatternRewriter &rewriter,
llvm::SmallDenseMap<Value, BlockData> &known);

static FailureOr<Value>
materializePointer(Value ptr, ConversionPatternRewriter &rewriter,
llvm::SmallDenseMap<Value, BlockData> &known);
Comment thread
CHNJZ marked this conversation as resolved.
Comment thread
CHNJZ marked this conversation as resolved.

static void
rewriteMakeTensorPtrOp(triton::MakeTensorPtrOp op, Value base,
ConversionPatternRewriter &rewriter,
Expand Down
95 changes: 61 additions & 34 deletions third_party/ascend/include/TritonToLinalg/LoadStoreConverter.h
Original file line number Diff line number Diff line change
Expand Up @@ -34,9 +34,10 @@
#include "triton/Dialect/Triton/IR/Dialect.h"

#include "mlir/Dialect/Arith/Utils/Utils.h"

#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h"

#include "ascend/include/TritonToLinalg/BlockPtrAnalysis.h"

namespace LoadStoreConverter {

using namespace mlir;
Expand All @@ -50,6 +51,35 @@ class AddPtrConverter : public OpConversionPattern<triton::AddPtrOp> {
ConversionPatternRewriter &rewriter) const override;
};

/// Materialize pointer tensor expressions directly instead of routing them
/// through a synthetic AddPtrOp with a zero offset. A higher benefit than the
/// generic value converters ensures pointer layout is handled by
/// BlockDataParser.
template <typename OpTy>
class MemoryPointerConverter : public OpConversionPattern<OpTy> {
public:
explicit MemoryPointerConverter(MLIRContext *context)
: OpConversionPattern<OpTy>(context, PatternBenefit(2)) {}

LogicalResult
matchAndRewrite(OpTy op, typename OpTy::Adaptor,
ConversionPatternRewriter &rewriter) const override {
auto resultTy = dyn_cast<RankedTensorType>(op->getResult(0).getType());
if (!resultTy || !isa<triton::PointerType>(resultTy.getElementType()))
return failure();

llvm::SmallDenseMap<Value, BlockData> known;
FailureOr<Value> memref =
BlockDataParser::materializePointer(op->getResult(0), rewriter, known);
if (failed(memref))
return rewriter.notifyMatchFailure(
op, "unsupported pointer expression for direct materialization");

rewriter.replaceOp(op, *memref);
return success();
}
};

class LoadConverter : public OpConversionPattern<triton::LoadOp> {
private:
void propagateWasBoolToInt8Attr(Operation *srcLoadOp, Operation *dstOp,
Expand Down Expand Up @@ -93,39 +123,36 @@ class LoadStoreCanonicalizer : public OpRewritePattern<OpTy> {
LogicalResult matchAndRewrite(OpTy op,
PatternRewriter &rewriter) const override {
Value ptrVal = op.getPtr();
Type ptrTy = ptrVal.getType();
auto ptrDefOp = ptrVal.getDefiningOp();

bool shouldAddZeros = false;
if (!isa<BlockArgument>(ptrVal))
shouldAddZeros = !isTensorPointerType(ptrTy) &&
!isa_and_nonnull<triton::AddPtrOp>(ptrDefOp);
else if (auto ptrType = dyn_cast<triton::PointerType>(ptrTy))
shouldAddZeros = ptrType.getPointeeType().isIntOrIndexOrFloat();

if (shouldAddZeros) {
if (isa_and_nonnull<triton::BitcastOp>(ptrDefOp)) {
auto castOp = cast<triton::BitcastOp>(ptrDefOp);
auto castSrc = castOp.getSrc();
if (!isa<BlockArgument>(castSrc)) {
auto castSrcDefOp = castSrc.getDefiningOp();
if (isa<triton::AddPtrOp>(castSrcDefOp)) {
return rewriter.notifyMatchFailure(
op, "BitcastCanonicalizer handles addptr->bitcast->load!");
}
}
}

Type zeroTy = getI32SameShape(ptrTy);
Value zeroVal =
createScalarOrSplatConstant(rewriter, op.getLoc(), zeroTy, 0);
Value addptrVal = rewriter.create<triton::AddPtrOp>(op.getLoc(), ptrTy,
ptrVal, zeroVal);
rewriter.modifyOpInPlace(
op, [&]() { op->replaceUsesOfWith(ptrVal, addptrVal); });
return success();
}
return failure();
auto ptrTy = dyn_cast<RankedTensorType>(ptrVal.getType());
if (!ptrTy || !isa<triton::PointerType>(ptrTy.getElementType()))
return failure();

// ReorderBroadcast may turn
// addptr(splat(base), splat(offset))
// into
// splat(addptr(base, offset)).
// Recover the former shape at the memory boundary so AddPtrConverter can
// keep using the real (non-synthetic) offset as its lowering anchor.
auto splatOp = ptrVal.getDefiningOp<triton::SplatOp>();
if (!splatOp)
return failure();

auto scalarAddPtr = splatOp.getSrc().getDefiningOp<triton::AddPtrOp>();
if (!scalarAddPtr)
return failure();

Value ptrSplat = rewriter.create<triton::SplatOp>(op.getLoc(), ptrTy,
scalarAddPtr.getPtr());
auto offsetTy =
ptrTy.cloneWith(std::nullopt, scalarAddPtr.getOffset().getType());
Value offsetSplat = rewriter.create<triton::SplatOp>(
op.getLoc(), offsetTy, scalarAddPtr.getOffset());
Value addptr = rewriter.create<triton::AddPtrOp>(op.getLoc(), ptrTy,
ptrSplat, offsetSplat);

rewriter.modifyOpInPlace(op,
[&]() { op->replaceUsesOfWith(ptrVal, addptr); });
return success();
}
};

Expand Down
107 changes: 98 additions & 9 deletions third_party/ascend/lib/TritonToLinalg/BlockPtrAnalysis.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -386,16 +386,11 @@ Value BlockDataParser::getScalarMemRef(Value ptr, Value memref,
const Location &loc,
ConversionPatternRewriter &rewriter) {
assert(isa<triton::PointerType>(ptr.getType()) && "expect a scalar pointer");
if (ptr.getDefiningOp<triton::AddPtrOp>()) {
if (auto castOp = memref.getDefiningOp<memref::ReinterpretCastOp>()) {
return castOp.getResult();
} else {
llvm_unreachable("pointer value is defined by an unexpected op");
}
}
if (auto castOp = memref.getDefiningOp<memref::ReinterpretCastOp>())
return castOp.getResult();

assert(isa<BlockArgument>(ptr) &&
"pointer should be produced by addptr or block argument");
assert(isa<BaseMemRefType>(memref.getType()) &&
"converted scalar pointer should be a memref");
BlockData data;
data.setSource(memref);
data.getOffsetsRef().push_back(rewriter.getIndexAttr(0));
Expand Down Expand Up @@ -1389,6 +1384,100 @@ void BlockDataParser::rewriteAddPtr(
rewriter.restoreInsertionPoint(insertPoint);
}

FailureOr<Value> BlockDataParser::materializePointer(
Value ptr, ConversionPatternRewriter &rewriter,
llvm::SmallDenseMap<Value, BlockData> &known) {
SmallVector<int64_t> resultShape;
if (auto resultTy = dyn_cast<RankedTensorType>(ptr.getType())) {
if (!isa<triton::PointerType>(resultTy.getElementType()))
return failure();
resultShape.append(resultTy.getShape().begin(), resultTy.getShape().end());
} else if (auto pointerTy = dyn_cast<triton::PointerType>(ptr.getType())) {
// A pointer to a shaped value is a block pointer, not a scalar element
// pointer. It is materialized by the make_tensor_ptr path instead.
if (isa<ShapedType>(pointerTy.getPointeeType()))
return failure();
resultShape.push_back(1);
} else {
return failure();
}

Operation *defOp = ptr.getDefiningOp();
if (!defOp)
return failure();

ConversionPatternRewriter::InsertionGuard guard(rewriter);
rewriter.setInsertionPoint(defOp);

BlockData data;
parse(ptr, data, ptr.getLoc(), rewriter, known);
if (!data.hasSource() || data.getMemAccType().isUnstructured())
return failure();

if (data.getSizesRef().empty()) {
data.getSizesRef().push_back(rewriter.getIndexAttr(1));
data.getStridesRef().push_back(rewriter.getIndexAttr(0));
data.getOffsetsRef().push_back(data.getScalarRef().isNull()
? OpFoldResult(rewriter.getIndexAttr(0))
: data.getScalarRef());
}

if (data.getRank() != static_cast<int64_t>(resultShape.size()))
return failure();

known[ptr] = data;

// A unit dimension with zero stride is represented as a normal contiguous
// unit dimension. Non-unit zero strides are intentional (for example a
// splatted base pointer) and must be preserved.
int64_t inferredSize = 1;
for (int64_t i = data.getRank() - 1; i >= 0; --i) {
auto strideConst = getConstantIntValue(data.getStridesRef()[i]);
auto sizeConst = getConstantIntValue(data.getSizesRef()[i]);
if (!sizeConst)
return failure();
if (sizeConst.value() == 1 && strideConst && strideConst.value() == 0)
data.getStridesRef()[i] = rewriter.getIndexAttr(inferredSize);
inferredSize *= sizeConst.value();
}

// Keep negative static offsets as SSA values. This mirrors rewriteAddPtr and
// avoids unsigned attribute handling in later reinterpret_cast lowering.
auto &offsets = data.getOffsetsRef();
for (OpFoldResult &offset : offsets) {
if (auto constVal = getConstantIntValue(offset);
constVal && constVal.value() < 0) {
offset =
rewriter
.create<arith::ConstantIndexOp>(ptr.getLoc(), constVal.value())
.getResult();
}
}

if (auto intToPtrOp = dyn_cast_or_null<triton::IntToPtrOp>(
data.getSourceRef().getDefiningOp())) {
auto pointerTy =
cast<triton::PointerType>(intToPtrOp.getResult().getType());
auto memrefTy =
MemRefType::get({ShapedType::kDynamic}, pointerTy.getPointeeType());
auto pointerCast = rewriter.create<hivm::PointerCastOp>(
intToPtrOp.getLoc(), memrefTy, ValueRange{intToPtrOp.getSrc()});
data.setSource(pointerCast.getResult());
}

if (data.hasResElemTy()) {
auto sourceTy = dyn_cast<BaseMemRefType>(data.getSourceRef().getType());
if (!sourceTy)
return failure();
auto castTy = sourceTy.cloneWith(std::nullopt, data.getResElemTyRef());
auto cast = rewriter.create<UnrealizedConversionCastOp>(
ptr.getLoc(), castTy, data.getSourceRef());
data.setSource(cast.getResult(0));
}

return data.createCastOp(resultShape, ptr.getLoc(), rewriter).getResult();
}

OpFoldResult
accumulatePotentialOffsetOnBase(triton::MakeTensorPtrOp op, Value base,
OpFoldResult offset,
Expand Down
59 changes: 58 additions & 1 deletion third_party/ascend/lib/TritonToLinalg/LoadStoreConverter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,39 @@ using namespace triton;
const std::string MayImplicitTransposeWithLastAxisTAG =
"MayImplicitTransposeWithLastAxis";

static bool shouldMaterializeCustomPointer(Value ptr) {
auto result = dyn_cast<OpResult>(ptr);
if (!result)
return false;

Operation *defOp = result.getOwner();
if (!defOp || !isDistributedTypeCustomOp(defOp))
return false;

auto srcIndices = defOp->getAttrOfType<DenseI32ArrayAttr>(
ConverterUtils::customSrcPtrIndexAttrName);
if (!srcIndices)
return false;

auto values = srcIndices.asArrayRef();
unsigned resultIdx = result.getResultNumber();
return resultIdx < values.size() && values[resultIdx] >= 0;
}

static FailureOr<Value>
resolveMemoryPointer(Value originalPtr, Value convertedPtr,
ConversionPatternRewriter &rewriter) {
llvm::SmallDenseMap<Value, BlockData> known;

if (shouldMaterializeCustomPointer(originalPtr))
return BlockDataParser::materializePointer(originalPtr, rewriter, known);

if (isa<MemRefType>(convertedPtr.getType()))
return convertedPtr;
Comment thread
CHNJZ marked this conversation as resolved.

return BlockDataParser::materializePointer(originalPtr, rewriter, known);
}

LogicalResult
AddPtrConverter::matchAndRewrite(triton::AddPtrOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const {
Expand Down Expand Up @@ -372,6 +405,12 @@ LoadConverter::matchAndRewrite(triton::LoadOp op, OpAdaptor adaptor,
auto other = op.getOther();
auto loc = op.getLoc();

FailureOr<Value> resolvedPtr =
resolveMemoryPointer(op.getPtr(), ptr, rewriter);
if (failed(resolvedPtr))
return rewriter.notifyMatchFailure(
op, "unable to materialize the load pointer as a memref");
ptr = *resolvedPtr;
// handling scalar
if (!isa<ShapedType>(op.getResult().getType())) {
auto scalarMemref =
Expand Down Expand Up @@ -728,8 +767,14 @@ AtomicRMWConverter::matchAndRewrite(triton::AtomicRMWOp op, OpAdaptor adaptor,
auto mask = op.getMask();
auto rmwOp = op.getAtomicRmwOp();
auto resType = dyn_cast<TensorType>(op.getResult().getType());
auto ptrType = dyn_cast<MemRefType>(ptr.getType());

FailureOr<Value> resolvedPtr =
resolveMemoryPointer(op.getPtr(), ptr, rewriter);
if (failed(resolvedPtr))
return rewriter.notifyMatchFailure(
op, "unable to materialize the atomic RMW pointer as a memref");
ptr = *resolvedPtr;
auto ptrType = dyn_cast<MemRefType>(ptr.getType());
if (!resType)
Comment thread
CHNJZ marked this conversation as resolved.
return rewriter.notifyMatchFailure(
op, "atomicRMWConverter: scalar will be handled by "
Expand Down Expand Up @@ -891,6 +936,12 @@ AtomicCASConverter::matchAndRewrite(triton::AtomicCASOp op, OpAdaptor adaptor,
auto val = op.getVal();
auto loc = op.getLoc();

FailureOr<Value> resolvedPtr =
resolveMemoryPointer(op.getPtr(), ptr, rewriter);
if (failed(resolvedPtr))
return rewriter.notifyMatchFailure(
op, "unable to materialize the atomic CAS pointer as a memref");
ptr = *resolvedPtr;
auto resType = dyn_cast<TensorType>(op.getResult().getType());
if (!resType) {
return rewriter.notifyMatchFailure(
Expand Down Expand Up @@ -1280,6 +1331,12 @@ StoreConverter::matchAndRewrite(triton::StoreOp op, OpAdaptor adaptor,
auto ptr = adaptor.getPtr();
auto val = adaptor.getValue();

FailureOr<Value> resolvedPtr =
resolveMemoryPointer(op.getPtr(), ptr, rewriter);
if (failed(resolvedPtr))
return rewriter.notifyMatchFailure(
op, "unable to materialize the store pointer as a memref");
ptr = *resolvedPtr;
// 1. boundary size check
auto boundaryCheck = op.getBoundaryCheck();
if (!boundaryCheck.empty()) {
Expand Down
9 changes: 9 additions & 0 deletions third_party/ascend/lib/TritonToLinalg/TritonOpConverter.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3293,6 +3293,15 @@ LogicalResult IndirectLoadConverter::matchAndRewrite(
auto res = op.getResult();
auto resTy = res.getType();

if (!isa<MemRefType>(src.getType())) {
llvm::SmallDenseMap<Value, BlockData> known;
FailureOr<Value> materialized =
BlockDataParser::materializePointer(op.getSrc(), rewriter, known);
if (failed(materialized))
return rewriter.notifyMatchFailure(
op, "unable to materialize indirect-load source as a memref");
src = *materialized;
}
// convert !tt.ptr<f32> to memref<?xf32>
auto srcTy = dyn_cast<MemRefType>(src.getType());
if (!srcTy) {
Expand Down
Loading
Loading