#include "mlir/Conversion/MemRefToLLVM/MemRefToLLVM.h"
#include "mlir/Analysis/DataLayoutAnalysis.h"
#include "mlir/Conversion/ConvertToLLVM/ToLLVMInterface.h"
#include "mlir/Conversion/LLVMCommon/ConversionTarget.h"
#include "mlir/Conversion/LLVMCommon/Pattern.h"
#include "mlir/Conversion/LLVMCommon/TypeConverter.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/LLVMIR/FunctionCallUtils.h"
#include "mlir/Dialect/LLVMIR/LLVMDialect.h"
#include "mlir/Dialect/LLVMIR/LLVMTypes.h"
#include "mlir/Dialect/MemRef/IR/MemRef.h"
#include "mlir/Dialect/MemRef/Utils/MemRefUtils.h"
#include "mlir/IR/AffineMap.h"
#include "mlir/IR/BuiltinTypes.h"
#include "mlir/IR/IRMapping.h"
#include "mlir/Pass/Pass.h"
#include "llvm/Support/DebugLog.h"
#include "llvm/Support/MathExtras.h"
#include <optional>
#define DEBUG_TYPE "memref-to-llvm"
namespace mlir {
#define GEN_PASS_DEF_FINALIZEMEMREFTOLLVMCONVERSIONPASS
#include "mlir/Conversion/Passes.h.inc"
}
using namespace mlir;
static constexpr LLVM::GEPNoWrapFlags kNoWrapFlags =
LLVM::GEPNoWrapFlags::inbounds | LLVM::GEPNoWrapFlags::nuw;
namespace {
static bool isStaticStrideOrOffset(int64_t strideOrOffset) {
return ShapedType::isStatic(strideOrOffset);
}
static FailureOr<LLVM::LLVMFuncOp>
getFreeFn(OpBuilder &b, const LLVMTypeConverter *typeConverter,
Operation *module, SymbolTableCollection *symbolTables) {
bool useGenericFn = typeConverter->getOptions().useGenericFunctions;
if (useGenericFn)
return LLVM::lookupOrCreateGenericFreeFn(b, module, symbolTables);
return LLVM::lookupOrCreateFreeFn(b, module, symbolTables);
}
static FailureOr<LLVM::LLVMFuncOp>
getNotalignedAllocFn(OpBuilder &b, const LLVMTypeConverter *typeConverter,
Operation *module, Type indexType,
SymbolTableCollection *symbolTables) {
bool useGenericFn = typeConverter->getOptions().useGenericFunctions;
if (useGenericFn)
return LLVM::lookupOrCreateGenericAllocFn(b, module, indexType,
symbolTables);
return LLVM::lookupOrCreateMallocFn(b, module, indexType, symbolTables);
}
static FailureOr<LLVM::LLVMFuncOp>
getAlignedAllocFn(OpBuilder &b, const LLVMTypeConverter *typeConverter,
Operation *module, Type indexType,
SymbolTableCollection *symbolTables) {
bool useGenericFn = typeConverter->getOptions().useGenericFunctions;
if (useGenericFn)
return LLVM::lookupOrCreateGenericAlignedAllocFn(b, module, indexType,
symbolTables);
return LLVM::lookupOrCreateAlignedAllocFn(b, module, indexType, symbolTables);
}
static Value createAligned(ConversionPatternRewriter &rewriter, Location loc,
Value input, Value alignment) {
Value one = LLVM::ConstantOp::create(rewriter, loc, alignment.getType(),
rewriter.getIndexAttr(1));
Value bump = LLVM::SubOp::create(rewriter, loc, alignment, one);
Value bumped = LLVM::AddOp::create(rewriter, loc, input, bump);
Value mod = LLVM::URemOp::create(rewriter, loc, bumped, alignment);
return LLVM::SubOp::create(rewriter, loc, bumped, mod);
}
static unsigned getMemRefEltSizeInBytes(const LLVMTypeConverter *typeConverter,
MemRefType memRefType, Operation *op,
const DataLayout *defaultLayout) {
const DataLayout *layout = defaultLayout;
if (const DataLayoutAnalysis *analysis =
typeConverter->getDataLayoutAnalysis()) {
layout = &analysis->getAbove(op);
}
Type elementType = memRefType.getElementType();
if (auto memRefElementType = dyn_cast<MemRefType>(elementType))
return typeConverter->getMemRefDescriptorSize(memRefElementType, *layout);
if (auto memRefElementType = dyn_cast<UnrankedMemRefType>(elementType))
return typeConverter->getUnrankedMemRefDescriptorSize(memRefElementType,
*layout);
return layout->getTypeSize(elementType);
}
static Value castAllocFuncResult(ConversionPatternRewriter &rewriter,
Location loc, Value allocatedPtr,
MemRefType memRefType, Type elementPtrType,
const LLVMTypeConverter &typeConverter) {
auto allocatedPtrTy = cast<LLVM::LLVMPointerType>(allocatedPtr.getType());
FailureOr<unsigned> maybeMemrefAddrSpace =
typeConverter.getMemRefAddressSpace(memRefType);
assert(succeeded(maybeMemrefAddrSpace) && "unsupported address space");
unsigned memrefAddrSpace = *maybeMemrefAddrSpace;
if (allocatedPtrTy.getAddressSpace() != memrefAddrSpace)
allocatedPtr = LLVM::AddrSpaceCastOp::create(
rewriter, loc,
LLVM::LLVMPointerType::get(rewriter.getContext(), memrefAddrSpace),
allocatedPtr);
return allocatedPtr;
}
class AllocOpLowering : public ConvertOpToLLVMPattern<memref::AllocOp> {
SymbolTableCollection *symbolTables = nullptr;
public:
explicit AllocOpLowering(const LLVMTypeConverter &typeConverter,
SymbolTableCollection *symbolTables = nullptr,
PatternBenefit benefit = 1)
: ConvertOpToLLVMPattern<memref::AllocOp>(typeConverter, benefit),
symbolTables(symbolTables) {}
LogicalResult
matchAndRewrite(memref::AllocOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto loc = op.getLoc();
MemRefType memRefType = op.getType();
if (!isConvertibleAndHasIdentityMaps(memRefType))
return rewriter.notifyMatchFailure(op, "incompatible memref type");
FailureOr<LLVM::LLVMFuncOp> allocFuncOp =
getNotalignedAllocFn(rewriter, getTypeConverter(),
op->getParentWithTrait<OpTrait::SymbolTable>(),
getIndexType(), symbolTables);
if (failed(allocFuncOp))
return failure();
SmallVector<Value, 4> sizes;
SmallVector<Value, 4> strides;
Value sizeBytes;
this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
rewriter, sizes, strides, sizeBytes, true);
Value alignment = getAlignment(rewriter, loc, op);
if (alignment) {
sizeBytes = LLVM::AddOp::create(rewriter, loc, sizeBytes, alignment);
}
Type elementPtrType = this->getElementPtrType(memRefType);
assert(elementPtrType && "could not compute element ptr type");
auto results =
LLVM::CallOp::create(rewriter, loc, allocFuncOp.value(), sizeBytes);
Value allocatedPtr =
castAllocFuncResult(rewriter, loc, results.getResult(), memRefType,
elementPtrType, *getTypeConverter());
Value alignedPtr = allocatedPtr;
if (alignment) {
Value allocatedInt =
LLVM::PtrToIntOp::create(rewriter, loc, getIndexType(), allocatedPtr);
Value alignmentInt =
createAligned(rewriter, loc, allocatedInt, alignment);
alignedPtr =
LLVM::IntToPtrOp::create(rewriter, loc, elementPtrType, alignmentInt);
}
auto memRefDescriptor = this->createMemRefDescriptor(
loc, memRefType, allocatedPtr, alignedPtr, sizes, strides, rewriter);
rewriter.replaceOp(op, {memRefDescriptor});
return success();
}
template <typename OpType>
Value getAlignment(ConversionPatternRewriter &rewriter, Location loc,
OpType op) const {
MemRefType memRefType = op.getType();
Value alignment;
if (auto alignmentAttr = op.getAlignment()) {
Type indexType = getIndexType();
alignment =
createIndexAttrConstant(rewriter, loc, indexType, *alignmentAttr);
} else if (!memRefType.getElementType().isSignlessIntOrIndexOrFloat()) {
alignment = getSizeInBytes(loc, memRefType.getElementType(), rewriter);
}
return alignment;
}
};
class AlignedAllocOpLowering : public ConvertOpToLLVMPattern<memref::AllocOp> {
SymbolTableCollection *symbolTables = nullptr;
public:
explicit AlignedAllocOpLowering(const LLVMTypeConverter &typeConverter,
SymbolTableCollection *symbolTables = nullptr,
PatternBenefit benefit = 1)
: ConvertOpToLLVMPattern<memref::AllocOp>(typeConverter, benefit),
symbolTables(symbolTables) {}
LogicalResult
matchAndRewrite(memref::AllocOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto loc = op.getLoc();
MemRefType memRefType = op.getType();
if (!isConvertibleAndHasIdentityMaps(memRefType))
return rewriter.notifyMatchFailure(op, "incompatible memref type");
FailureOr<LLVM::LLVMFuncOp> allocFuncOp =
getAlignedAllocFn(rewriter, getTypeConverter(),
op->getParentWithTrait<OpTrait::SymbolTable>(),
getIndexType(), symbolTables);
if (failed(allocFuncOp))
return failure();
SmallVector<Value, 4> sizes;
SmallVector<Value, 4> strides;
Value sizeBytes;
this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
rewriter, sizes, strides, sizeBytes, !false);
int64_t alignment = alignedAllocationGetAlignment(op, &defaultLayout);
Value allocAlignment =
createIndexAttrConstant(rewriter, loc, getIndexType(), alignment);
if (!isMemRefSizeMultipleOf(memRefType, alignment, op, &defaultLayout))
sizeBytes = createAligned(rewriter, loc, sizeBytes, allocAlignment);
Type elementPtrType = this->getElementPtrType(memRefType);
auto results =
LLVM::CallOp::create(rewriter, loc, allocFuncOp.value(),
ValueRange({allocAlignment, sizeBytes}));
Value ptr =
castAllocFuncResult(rewriter, loc, results.getResult(), memRefType,
elementPtrType, *getTypeConverter());
auto memRefDescriptor = this->createMemRefDescriptor(
loc, memRefType, ptr, ptr, sizes, strides, rewriter);
rewriter.replaceOp(op, {memRefDescriptor});
return success();
}
static constexpr uint64_t kMinAlignedAllocAlignment = 16UL;
int64_t alignedAllocationGetAlignment(memref::AllocOp op,
const DataLayout *defaultLayout) const {
if (std::optional<uint64_t> alignment = op.getAlignment())
return *alignment;
unsigned eltSizeBytes = getMemRefEltSizeInBytes(
getTypeConverter(), op.getType(), op, defaultLayout);
return std::max(kMinAlignedAllocAlignment,
llvm::PowerOf2Ceil(eltSizeBytes));
}
bool isMemRefSizeMultipleOf(MemRefType type, uint64_t factor, Operation *op,
const DataLayout *defaultLayout) const {
uint64_t sizeDivisor =
getMemRefEltSizeInBytes(getTypeConverter(), type, op, defaultLayout);
for (unsigned i = 0, e = type.getRank(); i < e; i++) {
if (type.isDynamicDim(i))
continue;
sizeDivisor = sizeDivisor * type.getDimSize(i);
}
return sizeDivisor % factor == 0;
}
private:
DataLayout defaultLayout;
};
struct AllocaOpLowering : public ConvertOpToLLVMPattern<memref::AllocaOp> {
using ConvertOpToLLVMPattern<memref::AllocaOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::AllocaOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto loc = op.getLoc();
MemRefType memRefType = op.getType();
if (!isConvertibleAndHasIdentityMaps(memRefType))
return rewriter.notifyMatchFailure(op, "incompatible memref type");
SmallVector<Value, 4> sizes;
SmallVector<Value, 4> strides;
Value size;
this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
rewriter, sizes, strides, size, !true);
auto elementType =
typeConverter->convertType(op.getType().getElementType());
FailureOr<unsigned> maybeAddressSpace =
getTypeConverter()->getMemRefAddressSpace(op.getType());
assert(succeeded(maybeAddressSpace) && "unsupported address space");
unsigned addrSpace = *maybeAddressSpace;
auto elementPtrType =
LLVM::LLVMPointerType::get(rewriter.getContext(), addrSpace);
auto allocatedElementPtr =
LLVM::AllocaOp::create(rewriter, loc, elementPtrType, elementType, size,
op.getAlignment().value_or(0));
auto memRefDescriptor = this->createMemRefDescriptor(
loc, memRefType, allocatedElementPtr, allocatedElementPtr, sizes,
strides, rewriter);
rewriter.replaceOp(op, {memRefDescriptor});
return success();
}
};
struct AllocaScopeOpLowering
: public ConvertOpToLLVMPattern<memref::AllocaScopeOp> {
using ConvertOpToLLVMPattern<memref::AllocaScopeOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::AllocaScopeOp allocaScopeOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
OpBuilder::InsertionGuard guard(rewriter);
Location loc = allocaScopeOp.getLoc();
auto *currentBlock = rewriter.getInsertionBlock();
auto *remainingOpsBlock =
rewriter.splitBlock(currentBlock, rewriter.getInsertionPoint());
Block *continueBlock;
if (allocaScopeOp.getNumResults() == 0) {
continueBlock = remainingOpsBlock;
} else {
continueBlock = rewriter.createBlock(
remainingOpsBlock, allocaScopeOp.getResultTypes(),
SmallVector<Location>(allocaScopeOp->getNumResults(),
allocaScopeOp.getLoc()));
LLVM::BrOp::create(rewriter, loc, ValueRange(), remainingOpsBlock);
}
Block *beforeBody = &allocaScopeOp.getBodyRegion().front();
Block *afterBody = &allocaScopeOp.getBodyRegion().back();
rewriter.inlineRegionBefore(allocaScopeOp.getBodyRegion(), continueBlock);
rewriter.setInsertionPointToEnd(currentBlock);
auto stackSaveOp = LLVM::StackSaveOp::create(rewriter, loc, getPtrType());
LLVM::BrOp::create(rewriter, loc, ValueRange(), beforeBody);
rewriter.setInsertionPointToEnd(afterBody);
auto returnOp =
cast<memref::AllocaScopeReturnOp>(afterBody->getTerminator());
auto branchOp = rewriter.replaceOpWithNewOp<LLVM::BrOp>(
returnOp, returnOp.getResults(), continueBlock);
rewriter.setInsertionPoint(branchOp);
LLVM::StackRestoreOp::create(rewriter, loc, stackSaveOp);
rewriter.replaceOp(allocaScopeOp, continueBlock->getArguments());
return success();
}
};
struct AssumeAlignmentOpLowering
: public ConvertOpToLLVMPattern<memref::AssumeAlignmentOp> {
using ConvertOpToLLVMPattern<
memref::AssumeAlignmentOp>::ConvertOpToLLVMPattern;
explicit AssumeAlignmentOpLowering(const LLVMTypeConverter &converter)
: ConvertOpToLLVMPattern<memref::AssumeAlignmentOp>(converter) {}
LogicalResult
matchAndRewrite(memref::AssumeAlignmentOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
Value memref = adaptor.getMemref();
unsigned alignment = op.getAlignment();
auto loc = op.getLoc();
auto srcMemRefType = cast<MemRefType>(op.getMemref().getType());
Value ptr = getStridedElementPtr(rewriter, loc, srcMemRefType, memref,
{});
Value trueCond =
LLVM::ConstantOp::create(rewriter, loc, rewriter.getBoolAttr(true));
Value alignmentConst =
createIndexAttrConstant(rewriter, loc, getIndexType(), alignment);
LLVM::AssumeOp::create(rewriter, loc, trueCond, LLVM::AssumeAlignTag(), ptr,
alignmentConst);
rewriter.replaceOp(op, memref);
return success();
}
};
struct DistinctObjectsOpLowering
: public ConvertOpToLLVMPattern<memref::DistinctObjectsOp> {
using ConvertOpToLLVMPattern<
memref::DistinctObjectsOp>::ConvertOpToLLVMPattern;
explicit DistinctObjectsOpLowering(const LLVMTypeConverter &converter)
: ConvertOpToLLVMPattern<memref::DistinctObjectsOp>(converter) {}
LogicalResult
matchAndRewrite(memref::DistinctObjectsOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
ValueRange operands = adaptor.getOperands();
if (operands.size() <= 1) {
rewriter.replaceOp(op, operands);
return success();
}
Location loc = op.getLoc();
SmallVector<Value> ptrs;
for (auto [origOperand, newOperand] :
llvm::zip_equal(op.getOperands(), operands)) {
auto memrefType = cast<MemRefType>(origOperand.getType());
MemRefDescriptor memRefDescriptor(newOperand);
Value ptr = memRefDescriptor.bufferPtr(rewriter, loc, *getTypeConverter(),
memrefType);
ptrs.push_back(ptr);
}
auto cond =
LLVM::ConstantOp::create(rewriter, loc, rewriter.getI1Type(), 1);
for (auto i : llvm::seq<size_t>(ptrs.size() - 1)) {
for (auto j : llvm::seq<size_t>(i + 1, ptrs.size())) {
Value ptr1 = ptrs[i];
Value ptr2 = ptrs[j];
LLVM::AssumeOp::create(rewriter, loc, cond,
LLVM::AssumeSeparateStorageTag{}, ptr1, ptr2);
}
}
rewriter.replaceOp(op, operands);
return success();
}
};
class DeallocOpLowering : public ConvertOpToLLVMPattern<memref::DeallocOp> {
SymbolTableCollection *symbolTables = nullptr;
public:
explicit DeallocOpLowering(const LLVMTypeConverter &typeConverter,
SymbolTableCollection *symbolTables = nullptr,
PatternBenefit benefit = 1)
: ConvertOpToLLVMPattern<memref::DeallocOp>(typeConverter, benefit),
symbolTables(symbolTables) {}
LogicalResult
matchAndRewrite(memref::DeallocOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
FailureOr<LLVM::LLVMFuncOp> freeFunc =
getFreeFn(rewriter, getTypeConverter(),
op->getParentWithTrait<OpTrait::SymbolTable>(), symbolTables);
if (failed(freeFunc))
return failure();
Value allocatedPtr;
if (auto unrankedTy =
llvm::dyn_cast<UnrankedMemRefType>(op.getMemref().getType())) {
auto elementPtrTy = LLVM::LLVMPointerType::get(
rewriter.getContext(), unrankedTy.getMemorySpaceAsInt());
allocatedPtr = UnrankedMemRefDescriptor::allocatedPtr(
rewriter, op.getLoc(),
UnrankedMemRefDescriptor(adaptor.getMemref())
.memRefDescPtr(rewriter, op.getLoc()),
elementPtrTy);
} else {
allocatedPtr = MemRefDescriptor(adaptor.getMemref())
.allocatedPtr(rewriter, op.getLoc());
}
rewriter.replaceOpWithNewOp<LLVM::CallOp>(op, freeFunc.value(),
allocatedPtr);
return success();
}
};
struct DimOpLowering : public ConvertOpToLLVMPattern<memref::DimOp> {
using ConvertOpToLLVMPattern<memref::DimOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::DimOp dimOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
Type operandType = dimOp.getSource().getType();
if (isa<UnrankedMemRefType>(operandType)) {
FailureOr<Value> extractedSize = extractSizeOfUnrankedMemRef(
operandType, dimOp, adaptor.getOperands(), rewriter);
if (failed(extractedSize))
return failure();
rewriter.replaceOp(dimOp, {*extractedSize});
return success();
}
if (isa<MemRefType>(operandType)) {
rewriter.replaceOp(
dimOp, {extractSizeOfRankedMemRef(operandType, dimOp,
adaptor.getOperands(), rewriter)});
return success();
}
llvm_unreachable("expected MemRefType or UnrankedMemRefType");
}
private:
FailureOr<Value>
extractSizeOfUnrankedMemRef(Type operandType, memref::DimOp dimOp,
OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const {
Location loc = dimOp.getLoc();
auto unrankedMemRefType = cast<UnrankedMemRefType>(operandType);
auto scalarMemRefType =
MemRefType::get({}, unrankedMemRefType.getElementType());
FailureOr<unsigned> maybeAddressSpace =
getTypeConverter()->getMemRefAddressSpace(unrankedMemRefType);
if (failed(maybeAddressSpace)) {
dimOp.emitOpError("memref memory space must be convertible to an integer "
"address space");
return failure();
}
unsigned addressSpace = *maybeAddressSpace;
UnrankedMemRefDescriptor unrankedDesc(adaptor.getSource());
Value underlyingRankedDesc = unrankedDesc.memRefDescPtr(rewriter, loc);
Type elementType = typeConverter->convertType(scalarMemRefType);
auto indexPtrTy =
LLVM::LLVMPointerType::get(rewriter.getContext(), addressSpace);
Value offsetPtr =
LLVM::GEPOp::create(rewriter, loc, indexPtrTy, elementType,
underlyingRankedDesc, ArrayRef<LLVM::GEPArg>{0, 2});
Value idxPlusOne = LLVM::AddOp::create(
rewriter, loc,
createIndexAttrConstant(rewriter, loc, getIndexType(), 1),
adaptor.getIndex());
Value sizePtr = LLVM::GEPOp::create(rewriter, loc, indexPtrTy,
getTypeConverter()->getIndexType(),
offsetPtr, idxPlusOne);
return LLVM::LoadOp::create(rewriter, loc,
getTypeConverter()->getIndexType(), sizePtr)
.getResult();
}
std::optional<int64_t> getConstantDimIndex(memref::DimOp dimOp) const {
if (auto idx = dimOp.getConstantIndex())
return idx;
if (auto constantOp = dimOp.getIndex().getDefiningOp<LLVM::ConstantOp>())
return cast<IntegerAttr>(constantOp.getValue()).getValue().getSExtValue();
return std::nullopt;
}
Value extractSizeOfRankedMemRef(Type operandType, memref::DimOp dimOp,
OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const {
Location loc = dimOp.getLoc();
MemRefType memRefType = cast<MemRefType>(operandType);
Type indexType = getIndexType();
if (std::optional<int64_t> index = getConstantDimIndex(dimOp)) {
int64_t i = *index;
if (i >= 0 && i < memRefType.getRank()) {
if (memRefType.isDynamicDim(i)) {
MemRefDescriptor descriptor(adaptor.getSource());
return descriptor.size(rewriter, loc, i);
}
int64_t dimSize = memRefType.getDimSize(i);
return createIndexAttrConstant(rewriter, loc, indexType, dimSize);
}
}
Value index = adaptor.getIndex();
int64_t rank = memRefType.getRank();
MemRefDescriptor memrefDescriptor(adaptor.getSource());
return memrefDescriptor.size(rewriter, loc, index, rank);
}
};
template <typename Derived>
struct LoadStoreOpLowering : public ConvertOpToLLVMPattern<Derived> {
using ConvertOpToLLVMPattern<Derived>::ConvertOpToLLVMPattern;
using ConvertOpToLLVMPattern<Derived>::isConvertibleAndHasIdentityMaps;
using Base = LoadStoreOpLowering<Derived>;
};
struct GenericAtomicRMWOpLowering
: public LoadStoreOpLowering<memref::GenericAtomicRMWOp> {
using Base::Base;
LogicalResult
matchAndRewrite(memref::GenericAtomicRMWOp atomicOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto loc = atomicOp.getLoc();
Type valueType = typeConverter->convertType(atomicOp.getResult().getType());
auto *initBlock = rewriter.getInsertionBlock();
auto *loopBlock = rewriter.splitBlock(initBlock, Block::iterator(atomicOp));
loopBlock->addArgument(valueType, loc);
auto *endBlock =
rewriter.splitBlock(loopBlock, Block::iterator(atomicOp)++);
rewriter.setInsertionPointToEnd(initBlock);
auto memRefType = cast<MemRefType>(atomicOp.getMemref().getType());
auto dataPtr = getStridedElementPtr(
rewriter, loc, memRefType, adaptor.getMemref(), adaptor.getIndices());
Value init = LLVM::LoadOp::create(
rewriter, loc, typeConverter->convertType(memRefType.getElementType()),
dataPtr);
LLVM::BrOp::create(rewriter, loc, init, loopBlock);
rewriter.setInsertionPointToStart(loopBlock);
auto loopArgument = loopBlock->getArgument(0);
IRMapping mapping;
mapping.map(atomicOp.getCurrentValue(), loopArgument);
Block &entryBlock = atomicOp.body().front();
for (auto &nestedOp : entryBlock.without_terminator()) {
Operation *clone = rewriter.clone(nestedOp, mapping);
mapping.map(nestedOp.getResults(), clone->getResults());
}
Value result = mapping.lookup(entryBlock.getTerminator()->getOperand(0));
auto successOrdering = LLVM::AtomicOrdering::acq_rel;
auto failureOrdering = LLVM::AtomicOrdering::monotonic;
auto cmpxchg =
LLVM::AtomicCmpXchgOp::create(rewriter, loc, dataPtr, loopArgument,
result, successOrdering, failureOrdering);
Value newLoaded = LLVM::ExtractValueOp::create(rewriter, loc, cmpxchg, 0);
Value ok = LLVM::ExtractValueOp::create(rewriter, loc, cmpxchg, 1);
LLVM::CondBrOp::create(rewriter, loc, ok, endBlock, ArrayRef<Value>(),
loopBlock, newLoaded);
rewriter.setInsertionPointToEnd(endBlock);
rewriter.replaceOp(atomicOp, {newLoaded});
return success();
}
};
static Type
convertGlobalMemrefTypeToLLVM(MemRefType type,
const LLVMTypeConverter &typeConverter) {
Type elementType = typeConverter.convertType(type.getElementType());
Type arrayTy = elementType;
for (int64_t dim : llvm::reverse(type.getShape()))
arrayTy = LLVM::LLVMArrayType::get(arrayTy, dim);
return arrayTy;
}
class GlobalMemrefOpLowering : public ConvertOpToLLVMPattern<memref::GlobalOp> {
SymbolTableCollection *symbolTables = nullptr;
public:
explicit GlobalMemrefOpLowering(const LLVMTypeConverter &typeConverter,
SymbolTableCollection *symbolTables = nullptr,
PatternBenefit benefit = 1)
: ConvertOpToLLVMPattern<memref::GlobalOp>(typeConverter, benefit),
symbolTables(symbolTables) {}
LogicalResult
matchAndRewrite(memref::GlobalOp global, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
MemRefType type = global.getType();
if (!isConvertibleAndHasIdentityMaps(type))
return failure();
Type arrayTy = convertGlobalMemrefTypeToLLVM(type, *getTypeConverter());
LLVM::Linkage linkage =
global.isPublic() ? LLVM::Linkage::External : LLVM::Linkage::Private;
bool isExternal = global.isExternal();
bool isUninitialized = global.isUninitialized();
Attribute initialValue = nullptr;
if (!isExternal && !isUninitialized) {
auto elementsAttr = llvm::cast<ElementsAttr>(*global.getInitialValue());
initialValue = elementsAttr;
if (type.getRank() == 0)
initialValue = elementsAttr.getSplatValue<Attribute>();
}
uint64_t alignment = global.getAlignment().value_or(0);
FailureOr<unsigned> addressSpace =
getTypeConverter()->getMemRefAddressSpace(type);
if (failed(addressSpace))
return global.emitOpError(
"memory space cannot be converted to an integer address space");
SymbolTable *symbolTable = nullptr;
if (symbolTables) {
Operation *symbolTableOp =
global->getParentWithTrait<OpTrait::SymbolTable>();
symbolTable = &symbolTables->getSymbolTable(symbolTableOp);
symbolTable->remove(global);
}
auto newGlobal = rewriter.replaceOpWithNewOp<LLVM::GlobalOp>(
global, arrayTy, global.getConstant(), linkage, global.getSymName(),
initialValue, alignment, *addressSpace);
if (symbolTable)
symbolTable->insert(newGlobal, rewriter.getInsertionPoint());
if (!isExternal && isUninitialized) {
rewriter.createBlock(&newGlobal.getInitializerRegion());
Value undef[] = {
LLVM::UndefOp::create(rewriter, newGlobal.getLoc(), arrayTy)};
LLVM::ReturnOp::create(rewriter, newGlobal.getLoc(), undef);
}
return success();
}
};
struct GetGlobalMemrefOpLowering
: public ConvertOpToLLVMPattern<memref::GetGlobalOp> {
using ConvertOpToLLVMPattern<memref::GetGlobalOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::GetGlobalOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto loc = op.getLoc();
MemRefType memRefType = op.getType();
if (!isConvertibleAndHasIdentityMaps(memRefType))
return rewriter.notifyMatchFailure(op, "incompatible memref type");
SmallVector<Value, 4> sizes;
SmallVector<Value, 4> strides;
Value sizeBytes;
this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
rewriter, sizes, strides, sizeBytes, !false);
MemRefType type = cast<MemRefType>(op.getResult().getType());
FailureOr<unsigned> maybeAddressSpace =
getTypeConverter()->getMemRefAddressSpace(type);
assert(succeeded(maybeAddressSpace) && "unsupported address space");
unsigned memSpace = *maybeAddressSpace;
Type arrayTy = convertGlobalMemrefTypeToLLVM(type, *getTypeConverter());
auto ptrTy = LLVM::LLVMPointerType::get(rewriter.getContext(), memSpace);
auto addressOf =
LLVM::AddressOfOp::create(rewriter, loc, ptrTy, op.getName());
auto gep =
LLVM::GEPOp::create(rewriter, loc, ptrTy, arrayTy, addressOf,
SmallVector<LLVM::GEPArg>(type.getRank() + 1, 0));
auto intPtrType = getIntPtrType(memSpace);
Value deadBeefConst =
createIndexAttrConstant(rewriter, op->getLoc(), intPtrType, 0xdeadbeef);
auto deadBeefPtr =
LLVM::IntToPtrOp::create(rewriter, loc, ptrTy, deadBeefConst);
auto memRefDescriptor = this->createMemRefDescriptor(
loc, memRefType, deadBeefPtr, gep, sizes, strides, rewriter);
rewriter.replaceOp(op, {memRefDescriptor});
return success();
}
};
struct LoadOpLowering : public LoadStoreOpLowering<memref::LoadOp> {
using Base::Base;
LogicalResult
matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto type = loadOp.getMemRefType();
Value dataPtr = getStridedElementPtr(rewriter, loadOp.getLoc(), type,
adaptor.getMemref(),
adaptor.getIndices(), kNoWrapFlags);
rewriter.replaceOpWithNewOp<LLVM::LoadOp>(
loadOp, typeConverter->convertType(type.getElementType()), dataPtr,
loadOp.getAlignment().value_or(0), false, loadOp.getNontemporal());
return success();
}
};
struct StoreOpLowering : public LoadStoreOpLowering<memref::StoreOp> {
using Base::Base;
LogicalResult
matchAndRewrite(memref::StoreOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto type = op.getMemRefType();
Value dataPtr =
getStridedElementPtr(rewriter, op.getLoc(), type, adaptor.getMemref(),
adaptor.getIndices(), kNoWrapFlags);
rewriter.replaceOpWithNewOp<LLVM::StoreOp>(op, adaptor.getValue(), dataPtr,
op.getAlignment().value_or(0),
false, op.getNontemporal());
return success();
}
};
struct PrefetchOpLowering : public LoadStoreOpLowering<memref::PrefetchOp> {
using Base::Base;
LogicalResult
matchAndRewrite(memref::PrefetchOp prefetchOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto type = prefetchOp.getMemRefType();
auto loc = prefetchOp.getLoc();
Value dataPtr = getStridedElementPtr(
rewriter, loc, type, adaptor.getMemref(), adaptor.getIndices());
IntegerAttr isWrite = rewriter.getI32IntegerAttr(prefetchOp.getIsWrite());
IntegerAttr localityHint = prefetchOp.getLocalityHintAttr();
IntegerAttr isData =
rewriter.getI32IntegerAttr(prefetchOp.getIsDataCache());
rewriter.replaceOpWithNewOp<LLVM::Prefetch>(prefetchOp, dataPtr, isWrite,
localityHint, isData);
return success();
}
};
struct RankOpLowering : public ConvertOpToLLVMPattern<memref::RankOp> {
using ConvertOpToLLVMPattern<memref::RankOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::RankOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
Location loc = op.getLoc();
Type operandType = op.getMemref().getType();
if (isa<UnrankedMemRefType>(operandType)) {
UnrankedMemRefDescriptor desc(adaptor.getMemref());
rewriter.replaceOp(op, {desc.rank(rewriter, loc)});
return success();
}
if (auto rankedMemRefType = dyn_cast<MemRefType>(operandType)) {
Type indexType = getIndexType();
rewriter.replaceOp(op,
{createIndexAttrConstant(rewriter, loc, indexType,
rankedMemRefType.getRank())});
return success();
}
return failure();
}
};
struct MemRefCastOpLowering : public ConvertOpToLLVMPattern<memref::CastOp> {
using ConvertOpToLLVMPattern::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::CastOp memRefCastOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
Type srcType = memRefCastOp.getOperand().getType();
Type dstType = memRefCastOp.getType();
if (isa<MemRefType>(srcType) && isa<MemRefType>(dstType))
if (typeConverter->convertType(srcType) !=
typeConverter->convertType(dstType))
return failure();
if (isa<UnrankedMemRefType>(srcType) && isa<UnrankedMemRefType>(dstType))
return failure();
auto targetStructType = typeConverter->convertType(memRefCastOp.getType());
auto loc = memRefCastOp.getLoc();
if (isa<MemRefType>(srcType) && isa<MemRefType>(dstType)) {
rewriter.replaceOp(memRefCastOp, {adaptor.getSource()});
return success();
}
if (isa<MemRefType>(srcType) && isa<UnrankedMemRefType>(dstType)) {
auto srcMemRefType = cast<MemRefType>(srcType);
int64_t rank = srcMemRefType.getRank();
auto ptr = getTypeConverter()->promoteOneMemRefDescriptor(
loc, adaptor.getSource(), rewriter);
auto rankVal = LLVM::ConstantOp::create(rewriter, loc, getIndexType(),
rewriter.getIndexAttr(rank));
UnrankedMemRefDescriptor memRefDesc =
UnrankedMemRefDescriptor::poison(rewriter, loc, targetStructType);
memRefDesc.setRank(rewriter, loc, rankVal);
memRefDesc.setMemRefDescPtr(rewriter, loc, ptr);
rewriter.replaceOp(memRefCastOp, (Value)memRefDesc);
} else if (isa<UnrankedMemRefType>(srcType) && isa<MemRefType>(dstType)) {
UnrankedMemRefDescriptor memRefDesc(adaptor.getSource());
auto ptr = memRefDesc.memRefDescPtr(rewriter, loc);
auto loadOp = LLVM::LoadOp::create(rewriter, loc, targetStructType, ptr);
rewriter.replaceOp(memRefCastOp, loadOp.getResult());
} else {
llvm_unreachable("Unsupported unranked memref to unranked memref cast");
}
return success();
}
};
class MemRefCopyOpLowering : public ConvertOpToLLVMPattern<memref::CopyOp> {
SymbolTableCollection *symbolTables = nullptr;
public:
explicit MemRefCopyOpLowering(const LLVMTypeConverter &typeConverter,
SymbolTableCollection *symbolTables = nullptr,
PatternBenefit benefit = 1)
: ConvertOpToLLVMPattern<memref::CopyOp>(typeConverter, benefit),
symbolTables(symbolTables) {}
LogicalResult
lowerToMemCopyIntrinsic(memref::CopyOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const {
auto loc = op.getLoc();
auto srcType = dyn_cast<MemRefType>(op.getSource().getType());
MemRefDescriptor srcDesc(adaptor.getSource());
Value numElements = LLVM::ConstantOp::create(rewriter, loc, getIndexType(),
rewriter.getIndexAttr(1));
for (int pos = 0; pos < srcType.getRank(); ++pos) {
auto size = srcDesc.size(rewriter, loc, pos);
numElements = LLVM::MulOp::create(rewriter, loc, numElements, size);
}
auto sizeInBytes = getSizeInBytes(loc, srcType.getElementType(), rewriter);
Value totalSize =
LLVM::MulOp::create(rewriter, loc, numElements, sizeInBytes);
Type elementType = typeConverter->convertType(srcType.getElementType());
Value srcBasePtr = srcDesc.alignedPtr(rewriter, loc);
Value srcOffset = srcDesc.offset(rewriter, loc);
Value srcPtr = LLVM::GEPOp::create(rewriter, loc, srcBasePtr.getType(),
elementType, srcBasePtr, srcOffset);
MemRefDescriptor targetDesc(adaptor.getTarget());
Value targetBasePtr = targetDesc.alignedPtr(rewriter, loc);
Value targetOffset = targetDesc.offset(rewriter, loc);
Value targetPtr =
LLVM::GEPOp::create(rewriter, loc, targetBasePtr.getType(), elementType,
targetBasePtr, targetOffset);
LLVM::MemcpyOp::create(rewriter, loc, targetPtr, srcPtr, totalSize,
false);
rewriter.eraseOp(op);
return success();
}
LogicalResult
lowerToMemCopyFunctionCall(memref::CopyOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const {
auto loc = op.getLoc();
auto srcType = cast<BaseMemRefType>(op.getSource().getType());
auto targetType = cast<BaseMemRefType>(op.getTarget().getType());
auto makeUnranked = [&, this](Value ranked, MemRefType type) {
auto rank = LLVM::ConstantOp::create(rewriter, loc, getIndexType(),
type.getRank());
auto *typeConverter = getTypeConverter();
auto ptr =
typeConverter->promoteOneMemRefDescriptor(loc, ranked, rewriter);
auto unrankedType =
UnrankedMemRefType::get(type.getElementType(), type.getMemorySpace());
return UnrankedMemRefDescriptor::pack(
rewriter, loc, *typeConverter, unrankedType, ValueRange{rank, ptr});
};
auto stackSaveOp = LLVM::StackSaveOp::create(rewriter, loc, getPtrType());
auto srcMemRefType = dyn_cast<MemRefType>(srcType);
Value unrankedSource =
srcMemRefType ? makeUnranked(adaptor.getSource(), srcMemRefType)
: adaptor.getSource();
auto targetMemRefType = dyn_cast<MemRefType>(targetType);
Value unrankedTarget =
targetMemRefType ? makeUnranked(adaptor.getTarget(), targetMemRefType)
: adaptor.getTarget();
auto one = LLVM::ConstantOp::create(rewriter, loc, getIndexType(),
rewriter.getIndexAttr(1));
auto promote = [&](Value desc) {
auto ptrType = LLVM::LLVMPointerType::get(rewriter.getContext());
auto allocated =
LLVM::AllocaOp::create(rewriter, loc, ptrType, desc.getType(), one);
LLVM::StoreOp::create(rewriter, loc, desc, allocated);
return allocated;
};
auto sourcePtr = promote(unrankedSource);
auto targetPtr = promote(unrankedTarget);
auto elemSize = getSizeInBytes(loc, srcType.getElementType(), rewriter);
auto copyFn = LLVM::lookupOrCreateMemRefCopyFn(
rewriter, op->getParentOfType<ModuleOp>(), getIndexType(),
sourcePtr.getType(), symbolTables);
if (failed(copyFn))
return failure();
LLVM::CallOp::create(rewriter, loc, copyFn.value(),
ValueRange{elemSize, sourcePtr, targetPtr});
LLVM::StackRestoreOp::create(rewriter, loc, stackSaveOp);
rewriter.eraseOp(op);
return success();
}
LogicalResult
matchAndRewrite(memref::CopyOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto srcType = cast<BaseMemRefType>(op.getSource().getType());
auto targetType = cast<BaseMemRefType>(op.getTarget().getType());
auto isContiguousMemrefType = [&](BaseMemRefType type) {
auto memrefType = dyn_cast<mlir::MemRefType>(type);
return memrefType &&
(memrefType.getLayout().isIdentity() ||
(memrefType.hasStaticShape() && memrefType.getNumElements() > 0 &&
memref::isStaticShapeAndContiguousRowMajor(memrefType)));
};
if (isContiguousMemrefType(srcType) && isContiguousMemrefType(targetType))
return lowerToMemCopyIntrinsic(op, adaptor, rewriter);
return lowerToMemCopyFunctionCall(op, adaptor, rewriter);
}
};
struct MemorySpaceCastOpLowering
: public ConvertOpToLLVMPattern<memref::MemorySpaceCastOp> {
using ConvertOpToLLVMPattern<
memref::MemorySpaceCastOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::MemorySpaceCastOp op, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
Location loc = op.getLoc();
Type resultType = op.getDest().getType();
if (auto resultTypeR = dyn_cast<MemRefType>(resultType)) {
auto resultDescType =
cast<LLVM::LLVMStructType>(typeConverter->convertType(resultTypeR));
Type newPtrType = resultDescType.getBody()[0];
SmallVector<Value> descVals;
MemRefDescriptor::unpack(rewriter, loc, adaptor.getSource(), resultTypeR,
descVals);
descVals[0] =
LLVM::AddrSpaceCastOp::create(rewriter, loc, newPtrType, descVals[0]);
descVals[1] =
LLVM::AddrSpaceCastOp::create(rewriter, loc, newPtrType, descVals[1]);
Value result = MemRefDescriptor::pack(rewriter, loc, *getTypeConverter(),
resultTypeR, descVals);
rewriter.replaceOp(op, result);
return success();
}
if (auto resultTypeU = dyn_cast<UnrankedMemRefType>(resultType)) {
auto sourceType = cast<UnrankedMemRefType>(op.getSource().getType());
FailureOr<unsigned> maybeSourceAddrSpace =
getTypeConverter()->getMemRefAddressSpace(sourceType);
if (failed(maybeSourceAddrSpace))
return rewriter.notifyMatchFailure(loc,
"non-integer source address space");
unsigned sourceAddrSpace = *maybeSourceAddrSpace;
FailureOr<unsigned> maybeResultAddrSpace =
getTypeConverter()->getMemRefAddressSpace(resultTypeU);
if (failed(maybeResultAddrSpace))
return rewriter.notifyMatchFailure(loc,
"non-integer result address space");
unsigned resultAddrSpace = *maybeResultAddrSpace;
UnrankedMemRefDescriptor sourceDesc(adaptor.getSource());
Value rank = sourceDesc.rank(rewriter, loc);
Value sourceUnderlyingDesc = sourceDesc.memRefDescPtr(rewriter, loc);
auto result = UnrankedMemRefDescriptor::poison(
rewriter, loc, typeConverter->convertType(resultTypeU));
result.setRank(rewriter, loc, rank);
Value resultUnderlyingSize = UnrankedMemRefDescriptor::computeSize(
rewriter, loc, *getTypeConverter(), result, resultAddrSpace);
Value resultUnderlyingDesc =
LLVM::AllocaOp::create(rewriter, loc, getPtrType(),
rewriter.getI8Type(), resultUnderlyingSize);
result.setMemRefDescPtr(rewriter, loc, resultUnderlyingDesc);
auto sourceElemPtrType =
LLVM::LLVMPointerType::get(rewriter.getContext(), sourceAddrSpace);
auto resultElemPtrType =
LLVM::LLVMPointerType::get(rewriter.getContext(), resultAddrSpace);
Value allocatedPtr = sourceDesc.allocatedPtr(
rewriter, loc, sourceUnderlyingDesc, sourceElemPtrType);
Value alignedPtr =
sourceDesc.alignedPtr(rewriter, loc, *getTypeConverter(),
sourceUnderlyingDesc, sourceElemPtrType);
allocatedPtr = LLVM::AddrSpaceCastOp::create(
rewriter, loc, resultElemPtrType, allocatedPtr);
alignedPtr = LLVM::AddrSpaceCastOp::create(rewriter, loc,
resultElemPtrType, alignedPtr);
result.setAllocatedPtr(rewriter, loc, resultUnderlyingDesc,
resultElemPtrType, allocatedPtr);
result.setAlignedPtr(rewriter, loc, *getTypeConverter(),
resultUnderlyingDesc, resultElemPtrType, alignedPtr);
Value sourceIndexVals =
sourceDesc.offsetBasePtr(rewriter, loc, *getTypeConverter(),
sourceUnderlyingDesc, sourceElemPtrType);
Value resultIndexVals =
result.offsetBasePtr(rewriter, loc, *getTypeConverter(),
resultUnderlyingDesc, resultElemPtrType);
int64_t bytesToSkip =
2 * llvm::divideCeil(
getTypeConverter()->getPointerBitwidth(resultAddrSpace), 8);
Value bytesToSkipConst = LLVM::ConstantOp::create(
rewriter, loc, getIndexType(), rewriter.getIndexAttr(bytesToSkip));
Value copySize =
LLVM::SubOp::create(rewriter, loc, getIndexType(),
resultUnderlyingSize, bytesToSkipConst);
LLVM::MemcpyOp::create(rewriter, loc, resultIndexVals, sourceIndexVals,
copySize, false);
rewriter.replaceOp(op, ValueRange{result});
return success();
}
return rewriter.notifyMatchFailure(loc, "unexpected memref type");
}
};
static void extractPointersAndOffset(Location loc,
ConversionPatternRewriter &rewriter,
const LLVMTypeConverter &typeConverter,
Value originalOperand,
Value convertedOperand,
Value *allocatedPtr, Value *alignedPtr,
Value *offset = nullptr) {
Type operandType = originalOperand.getType();
if (isa<MemRefType>(operandType)) {
MemRefDescriptor desc(convertedOperand);
*allocatedPtr = desc.allocatedPtr(rewriter, loc);
*alignedPtr = desc.alignedPtr(rewriter, loc);
if (offset != nullptr)
*offset = desc.offset(rewriter, loc);
return;
}
unsigned memorySpace = *typeConverter.getMemRefAddressSpace(
cast<UnrankedMemRefType>(operandType));
auto elementPtrType =
LLVM::LLVMPointerType::get(rewriter.getContext(), memorySpace);
UnrankedMemRefDescriptor unrankedDesc(convertedOperand);
Value underlyingDescPtr = unrankedDesc.memRefDescPtr(rewriter, loc);
*allocatedPtr = UnrankedMemRefDescriptor::allocatedPtr(
rewriter, loc, underlyingDescPtr, elementPtrType);
*alignedPtr = UnrankedMemRefDescriptor::alignedPtr(
rewriter, loc, typeConverter, underlyingDescPtr, elementPtrType);
if (offset != nullptr) {
*offset = UnrankedMemRefDescriptor::offset(
rewriter, loc, typeConverter, underlyingDescPtr, elementPtrType);
}
}
struct MemRefReinterpretCastOpLowering
: public ConvertOpToLLVMPattern<memref::ReinterpretCastOp> {
using ConvertOpToLLVMPattern<
memref::ReinterpretCastOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::ReinterpretCastOp castOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
Type srcType = castOp.getSource().getType();
Value descriptor;
if (failed(convertSourceMemRefToDescriptor(rewriter, srcType, castOp,
adaptor, &descriptor)))
return failure();
rewriter.replaceOp(castOp, {descriptor});
return success();
}
private:
LogicalResult convertSourceMemRefToDescriptor(
ConversionPatternRewriter &rewriter, Type srcType,
memref::ReinterpretCastOp castOp,
memref::ReinterpretCastOp::Adaptor adaptor, Value *descriptor) const {
MemRefType targetMemRefType =
cast<MemRefType>(castOp.getResult().getType());
auto llvmTargetDescriptorTy = dyn_cast_or_null<LLVM::LLVMStructType>(
typeConverter->convertType(targetMemRefType));
if (!llvmTargetDescriptorTy)
return failure();
Location loc = castOp.getLoc();
auto desc = MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy);
Value allocatedPtr, alignedPtr;
extractPointersAndOffset(loc, rewriter, *getTypeConverter(),
castOp.getSource(), adaptor.getSource(),
&allocatedPtr, &alignedPtr);
desc.setAllocatedPtr(rewriter, loc, allocatedPtr);
desc.setAlignedPtr(rewriter, loc, alignedPtr);
if (castOp.isDynamicOffset(0))
desc.setOffset(rewriter, loc, adaptor.getOffsets()[0]);
else
desc.setConstantOffset(rewriter, loc, castOp.getStaticOffset(0));
unsigned dynSizeId = 0;
unsigned dynStrideId = 0;
for (unsigned i = 0, e = targetMemRefType.getRank(); i < e; ++i) {
if (castOp.isDynamicSize(i))
desc.setSize(rewriter, loc, i, adaptor.getSizes()[dynSizeId++]);
else
desc.setConstantSize(rewriter, loc, i, castOp.getStaticSize(i));
if (castOp.isDynamicStride(i))
desc.setStride(rewriter, loc, i, adaptor.getStrides()[dynStrideId++]);
else
desc.setConstantStride(rewriter, loc, i, castOp.getStaticStride(i));
}
*descriptor = desc;
return success();
}
};
struct MemRefReshapeOpLowering
: public ConvertOpToLLVMPattern<memref::ReshapeOp> {
using ConvertOpToLLVMPattern<memref::ReshapeOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::ReshapeOp reshapeOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
Type srcType = reshapeOp.getSource().getType();
Value descriptor;
if (failed(convertSourceMemRefToDescriptor(rewriter, srcType, reshapeOp,
adaptor, &descriptor)))
return failure();
rewriter.replaceOp(reshapeOp, {descriptor});
return success();
}
private:
LogicalResult
convertSourceMemRefToDescriptor(ConversionPatternRewriter &rewriter,
Type srcType, memref::ReshapeOp reshapeOp,
memref::ReshapeOp::Adaptor adaptor,
Value *descriptor) const {
auto shapeMemRefType = cast<MemRefType>(reshapeOp.getShape().getType());
if (shapeMemRefType.hasStaticShape()) {
MemRefType targetMemRefType =
cast<MemRefType>(reshapeOp.getResult().getType());
auto llvmTargetDescriptorTy = dyn_cast_or_null<LLVM::LLVMStructType>(
typeConverter->convertType(targetMemRefType));
if (!llvmTargetDescriptorTy)
return failure();
Location loc = reshapeOp.getLoc();
auto desc =
MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy);
Value allocatedPtr, alignedPtr;
extractPointersAndOffset(loc, rewriter, *getTypeConverter(),
reshapeOp.getSource(), adaptor.getSource(),
&allocatedPtr, &alignedPtr);
desc.setAllocatedPtr(rewriter, loc, allocatedPtr);
desc.setAlignedPtr(rewriter, loc, alignedPtr);
int64_t offset;
SmallVector<int64_t> strides;
if (failed(targetMemRefType.getStridesAndOffset(strides, offset)))
return rewriter.notifyMatchFailure(
reshapeOp, "failed to get stride and offset exprs");
if (!isStaticStrideOrOffset(offset))
return rewriter.notifyMatchFailure(reshapeOp,
"dynamic offset is unsupported");
desc.setConstantOffset(rewriter, loc, offset);
assert(targetMemRefType.getLayout().isIdentity() &&
"Identity layout map is a precondition of a valid reshape op");
Type indexType = getIndexType();
Value stride = nullptr;
int64_t targetRank = targetMemRefType.getRank();
for (auto i : llvm::reverse(llvm::seq<int64_t>(0, targetRank))) {
if (ShapedType::isStatic(strides[i])) {
stride =
createIndexAttrConstant(rewriter, loc, indexType, strides[i]);
} else if (!stride) {
stride = createIndexAttrConstant(rewriter, loc, indexType, 1);
}
Value dimSize;
if (!targetMemRefType.isDynamicDim(i)) {
dimSize = createIndexAttrConstant(rewriter, loc, indexType,
targetMemRefType.getDimSize(i));
} else {
Value shapeOp = reshapeOp.getShape();
Value index = createIndexAttrConstant(rewriter, loc, indexType, i);
dimSize = memref::LoadOp::create(rewriter, loc, shapeOp, index);
Type indexType = getIndexType();
if (dimSize.getType() != indexType)
dimSize = typeConverter->materializeTargetConversion(
rewriter, loc, indexType, dimSize);
assert(dimSize && "Invalid memref element type");
}
desc.setSize(rewriter, loc, i, dimSize);
desc.setStride(rewriter, loc, i, stride);
stride = LLVM::MulOp::create(rewriter, loc, stride, dimSize);
}
*descriptor = desc;
return success();
}
Location loc = reshapeOp.getLoc();
MemRefDescriptor shapeDesc(adaptor.getShape());
Value resultRank = shapeDesc.size(rewriter, loc, 0);
auto targetType = cast<UnrankedMemRefType>(reshapeOp.getResult().getType());
unsigned addressSpace =
*getTypeConverter()->getMemRefAddressSpace(targetType);
auto targetDesc = UnrankedMemRefDescriptor::poison(
rewriter, loc, typeConverter->convertType(targetType));
targetDesc.setRank(rewriter, loc, resultRank);
Value allocationSize = UnrankedMemRefDescriptor::computeSize(
rewriter, loc, *getTypeConverter(), targetDesc, addressSpace);
Value underlyingDescPtr = LLVM::AllocaOp::create(
rewriter, loc, getPtrType(), IntegerType::get(getContext(), 8),
allocationSize);
targetDesc.setMemRefDescPtr(rewriter, loc, underlyingDescPtr);
Value allocatedPtr, alignedPtr, offset;
extractPointersAndOffset(loc, rewriter, *getTypeConverter(),
reshapeOp.getSource(), adaptor.getSource(),
&allocatedPtr, &alignedPtr, &offset);
auto elementPtrType =
LLVM::LLVMPointerType::get(rewriter.getContext(), addressSpace);
UnrankedMemRefDescriptor::setAllocatedPtr(rewriter, loc, underlyingDescPtr,
elementPtrType, allocatedPtr);
UnrankedMemRefDescriptor::setAlignedPtr(rewriter, loc, *getTypeConverter(),
underlyingDescPtr, elementPtrType,
alignedPtr);
UnrankedMemRefDescriptor::setOffset(rewriter, loc, *getTypeConverter(),
underlyingDescPtr, elementPtrType,
offset);
Value targetSizesBase = UnrankedMemRefDescriptor::sizeBasePtr(
rewriter, loc, *getTypeConverter(), underlyingDescPtr, elementPtrType);
Value targetStridesBase = UnrankedMemRefDescriptor::strideBasePtr(
rewriter, loc, *getTypeConverter(), targetSizesBase, resultRank);
Value shapeOperandPtr = shapeDesc.alignedPtr(rewriter, loc);
Value oneIndex = createIndexAttrConstant(rewriter, loc, getIndexType(), 1);
Value resultRankMinusOne =
LLVM::SubOp::create(rewriter, loc, resultRank, oneIndex);
Block *initBlock = rewriter.getInsertionBlock();
Type indexType = getTypeConverter()->getIndexType();
Block::iterator remainingOpsIt = std::next(rewriter.getInsertionPoint());
Block *condBlock = rewriter.createBlock(initBlock->getParent(), {},
{indexType, indexType}, {loc, loc});
Block *remainingBlock = rewriter.splitBlock(initBlock, remainingOpsIt);
rewriter.mergeBlocks(remainingBlock, condBlock, ValueRange());
rewriter.setInsertionPointToEnd(initBlock);
LLVM::BrOp::create(rewriter, loc,
ValueRange({resultRankMinusOne, oneIndex}), condBlock);
rewriter.setInsertionPointToStart(condBlock);
Value indexArg = condBlock->getArgument(0);
Value strideArg = condBlock->getArgument(1);
Value zeroIndex = createIndexAttrConstant(rewriter, loc, indexType, 0);
Value pred = LLVM::ICmpOp::create(
rewriter, loc, IntegerType::get(rewriter.getContext(), 1),
LLVM::ICmpPredicate::sge, indexArg, zeroIndex);
Block *bodyBlock =
rewriter.splitBlock(condBlock, rewriter.getInsertionPoint());
rewriter.setInsertionPointToStart(bodyBlock);
auto llvmIndexPtrType = LLVM::LLVMPointerType::get(rewriter.getContext());
Value sizeLoadGep = LLVM::GEPOp::create(
rewriter, loc, llvmIndexPtrType,
typeConverter->convertType(shapeMemRefType.getElementType()),
shapeOperandPtr, indexArg);
Value size = LLVM::LoadOp::create(rewriter, loc, indexType, sizeLoadGep);
UnrankedMemRefDescriptor::setSize(rewriter, loc, *getTypeConverter(),
targetSizesBase, indexArg, size);
UnrankedMemRefDescriptor::setStride(rewriter, loc, *getTypeConverter(),
targetStridesBase, indexArg, strideArg);
Value nextStride = LLVM::MulOp::create(rewriter, loc, strideArg, size);
Value decrement = LLVM::SubOp::create(rewriter, loc, indexArg, oneIndex);
LLVM::BrOp::create(rewriter, loc, ValueRange({decrement, nextStride}),
condBlock);
Block *remainder =
rewriter.splitBlock(bodyBlock, rewriter.getInsertionPoint());
rewriter.setInsertionPointToEnd(condBlock);
LLVM::CondBrOp::create(rewriter, loc, pred, bodyBlock, ValueRange(),
remainder, ValueRange());
rewriter.setInsertionPointToStart(remainder);
*descriptor = targetDesc;
return success();
}
};
template <typename ReshapeOp>
class ReassociatingReshapeOpConversion
: public ConvertOpToLLVMPattern<ReshapeOp> {
public:
using ConvertOpToLLVMPattern<ReshapeOp>::ConvertOpToLLVMPattern;
using ReshapeOpAdaptor = typename ReshapeOp::Adaptor;
LogicalResult
matchAndRewrite(ReshapeOp reshapeOp, typename ReshapeOp::Adaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
return rewriter.notifyMatchFailure(
reshapeOp,
"reassociation operations should have been expanded beforehand");
}
};
struct SubViewOpLowering : public ConvertOpToLLVMPattern<memref::SubViewOp> {
using ConvertOpToLLVMPattern<memref::SubViewOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::SubViewOp subViewOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
return rewriter.notifyMatchFailure(
subViewOp, "subview operations should have been expanded beforehand");
}
};
class TransposeOpLowering : public ConvertOpToLLVMPattern<memref::TransposeOp> {
public:
using ConvertOpToLLVMPattern<memref::TransposeOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::TransposeOp transposeOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto loc = transposeOp.getLoc();
MemRefDescriptor viewMemRef(adaptor.getIn());
if (transposeOp.getPermutation().isIdentity())
return rewriter.replaceOp(transposeOp, {viewMemRef}), success();
auto targetMemRef = MemRefDescriptor::poison(
rewriter, loc,
typeConverter->convertType(transposeOp.getIn().getType()));
targetMemRef.setAllocatedPtr(rewriter, loc,
viewMemRef.allocatedPtr(rewriter, loc));
targetMemRef.setAlignedPtr(rewriter, loc,
viewMemRef.alignedPtr(rewriter, loc));
targetMemRef.setOffset(rewriter, loc, viewMemRef.offset(rewriter, loc));
for (const auto &en :
llvm::enumerate(transposeOp.getPermutation().getResults())) {
int targetPos = en.index();
int sourcePos = cast<AffineDimExpr>(en.value()).getPosition();
targetMemRef.setSize(rewriter, loc, targetPos,
viewMemRef.size(rewriter, loc, sourcePos));
targetMemRef.setStride(rewriter, loc, targetPos,
viewMemRef.stride(rewriter, loc, sourcePos));
}
rewriter.replaceOp(transposeOp, {targetMemRef});
return success();
}
};
struct ViewOpLowering : public ConvertOpToLLVMPattern<memref::ViewOp> {
using ConvertOpToLLVMPattern<memref::ViewOp>::ConvertOpToLLVMPattern;
Value getSize(ConversionPatternRewriter &rewriter, Location loc,
ArrayRef<int64_t> shape, ValueRange dynamicSizes, unsigned idx,
Type indexType) const {
assert(idx < shape.size());
if (ShapedType::isStatic(shape[idx]))
return createIndexAttrConstant(rewriter, loc, indexType, shape[idx]);
unsigned nDynamic =
llvm::count_if(shape.take_front(idx), ShapedType::isDynamic);
return dynamicSizes[nDynamic];
}
Value getStride(ConversionPatternRewriter &rewriter, Location loc,
ArrayRef<int64_t> strides, Value nextSize,
Value runningStride, unsigned idx, Type indexType) const {
assert(idx < strides.size());
if (ShapedType::isStatic(strides[idx]))
return createIndexAttrConstant(rewriter, loc, indexType, strides[idx]);
if (nextSize)
return runningStride
? LLVM::MulOp::create(rewriter, loc, runningStride, nextSize)
: nextSize;
assert(!runningStride);
return createIndexAttrConstant(rewriter, loc, indexType, 1);
}
LogicalResult
matchAndRewrite(memref::ViewOp viewOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto loc = viewOp.getLoc();
auto viewMemRefType = viewOp.getType();
auto targetElementTy =
typeConverter->convertType(viewMemRefType.getElementType());
auto targetDescTy = typeConverter->convertType(viewMemRefType);
if (!targetDescTy || !targetElementTy ||
!LLVM::isCompatibleType(targetElementTy) ||
!LLVM::isCompatibleType(targetDescTy))
return viewOp.emitWarning("Target descriptor type not converted to LLVM"),
failure();
int64_t offset;
SmallVector<int64_t, 4> strides;
auto successStrides = viewMemRefType.getStridesAndOffset(strides, offset);
if (failed(successStrides))
return viewOp.emitWarning("cannot cast to non-strided shape"), failure();
assert(offset == 0 && "expected offset to be 0");
if (!strides.empty() && (strides.back() != 1 && strides.back() != 0))
return viewOp.emitWarning("cannot cast to non-contiguous shape"),
failure();
MemRefDescriptor sourceMemRef(adaptor.getSource());
auto targetMemRef = MemRefDescriptor::poison(rewriter, loc, targetDescTy);
Value allocatedPtr = sourceMemRef.allocatedPtr(rewriter, loc);
auto srcMemRefType = cast<MemRefType>(viewOp.getSource().getType());
targetMemRef.setAllocatedPtr(rewriter, loc, allocatedPtr);
Value alignedPtr = sourceMemRef.alignedPtr(rewriter, loc);
alignedPtr = LLVM::GEPOp::create(
rewriter, loc, alignedPtr.getType(),
typeConverter->convertType(srcMemRefType.getElementType()), alignedPtr,
adaptor.getByteShift());
targetMemRef.setAlignedPtr(rewriter, loc, alignedPtr);
Type indexType = getIndexType();
targetMemRef.setOffset(
rewriter, loc,
createIndexAttrConstant(rewriter, loc, indexType, offset));
if (viewMemRefType.getRank() == 0)
return rewriter.replaceOp(viewOp, {targetMemRef}), success();
Value stride = nullptr, nextSize = nullptr;
for (int i = viewMemRefType.getRank() - 1; i >= 0; --i) {
Value size = getSize(rewriter, loc, viewMemRefType.getShape(),
adaptor.getSizes(), i, indexType);
targetMemRef.setSize(rewriter, loc, i, size);
stride =
getStride(rewriter, loc, strides, nextSize, stride, i, indexType);
targetMemRef.setStride(rewriter, loc, i, stride);
nextSize = size;
}
rewriter.replaceOp(viewOp, {targetMemRef});
return success();
}
};
static std::optional<LLVM::AtomicBinOp>
matchSimpleAtomicOp(memref::AtomicRMWOp atomicOp) {
switch (atomicOp.getKind()) {
case arith::AtomicRMWKind::addf:
return LLVM::AtomicBinOp::fadd;
case arith::AtomicRMWKind::addi:
return LLVM::AtomicBinOp::add;
case arith::AtomicRMWKind::assign:
return LLVM::AtomicBinOp::xchg;
case arith::AtomicRMWKind::maximumf:
LDBG() << "the lowering of memref.atomicrmw maximumf changed "
"from fmax to fmaximum, expect more NaNs";
return LLVM::AtomicBinOp::fmaximum;
case arith::AtomicRMWKind::maxnumf:
return LLVM::AtomicBinOp::fmax;
case arith::AtomicRMWKind::maxs:
return LLVM::AtomicBinOp::max;
case arith::AtomicRMWKind::maxu:
return LLVM::AtomicBinOp::umax;
case arith::AtomicRMWKind::minimumf:
LDBG() << "the lowering of memref.atomicrmw minimum changed "
"from fmin to fminimum, expect more NaNs";
return LLVM::AtomicBinOp::fminimum;
case arith::AtomicRMWKind::minnumf:
return LLVM::AtomicBinOp::fmin;
case arith::AtomicRMWKind::mins:
return LLVM::AtomicBinOp::min;
case arith::AtomicRMWKind::minu:
return LLVM::AtomicBinOp::umin;
case arith::AtomicRMWKind::ori:
return LLVM::AtomicBinOp::_or;
case arith::AtomicRMWKind::xori:
return LLVM::AtomicBinOp::_xor;
case arith::AtomicRMWKind::andi:
return LLVM::AtomicBinOp::_and;
default:
return std::nullopt;
}
llvm_unreachable("Invalid AtomicRMWKind");
}
struct AtomicRMWOpLowering : public LoadStoreOpLowering<memref::AtomicRMWOp> {
using Base::Base;
LogicalResult
matchAndRewrite(memref::AtomicRMWOp atomicOp, OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
auto maybeKind = matchSimpleAtomicOp(atomicOp);
if (!maybeKind)
return failure();
auto memRefType = atomicOp.getMemRefType();
SmallVector<int64_t> strides;
int64_t offset;
if (failed(memRefType.getStridesAndOffset(strides, offset)))
return failure();
auto dataPtr =
getStridedElementPtr(rewriter, atomicOp.getLoc(), memRefType,
adaptor.getMemref(), adaptor.getIndices());
rewriter.replaceOpWithNewOp<LLVM::AtomicRMWOp>(
atomicOp, *maybeKind, dataPtr, adaptor.getValue(),
LLVM::AtomicOrdering::acq_rel);
return success();
}
};
class ConvertExtractAlignedPointerAsIndex
: public ConvertOpToLLVMPattern<memref::ExtractAlignedPointerAsIndexOp> {
public:
using ConvertOpToLLVMPattern<
memref::ExtractAlignedPointerAsIndexOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::ExtractAlignedPointerAsIndexOp extractOp,
OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
BaseMemRefType sourceTy = extractOp.getSource().getType();
Value alignedPtr;
if (sourceTy.hasRank()) {
MemRefDescriptor desc(adaptor.getSource());
alignedPtr = desc.alignedPtr(rewriter, extractOp->getLoc());
} else {
auto elementPtrTy = LLVM::LLVMPointerType::get(
rewriter.getContext(), sourceTy.getMemorySpaceAsInt());
UnrankedMemRefDescriptor desc(adaptor.getSource());
Value descPtr = desc.memRefDescPtr(rewriter, extractOp->getLoc());
alignedPtr = UnrankedMemRefDescriptor::alignedPtr(
rewriter, extractOp->getLoc(), *getTypeConverter(), descPtr,
elementPtrTy);
}
rewriter.replaceOpWithNewOp<LLVM::PtrToIntOp>(
extractOp, getTypeConverter()->getIndexType(), alignedPtr);
return success();
}
};
class ExtractStridedMetadataOpLowering
: public ConvertOpToLLVMPattern<memref::ExtractStridedMetadataOp> {
public:
using ConvertOpToLLVMPattern<
memref::ExtractStridedMetadataOp>::ConvertOpToLLVMPattern;
LogicalResult
matchAndRewrite(memref::ExtractStridedMetadataOp extractStridedMetadataOp,
OpAdaptor adaptor,
ConversionPatternRewriter &rewriter) const override {
if (!LLVM::isCompatibleType(adaptor.getOperands().front().getType()))
return failure();
MemRefDescriptor sourceMemRef(adaptor.getSource());
Location loc = extractStridedMetadataOp.getLoc();
Value source = extractStridedMetadataOp.getSource();
auto sourceMemRefType = cast<MemRefType>(source.getType());
int64_t rank = sourceMemRefType.getRank();
SmallVector<Value> results;
results.reserve(2 + rank * 2);
Value baseBuffer = sourceMemRef.allocatedPtr(rewriter, loc);
Value alignedBuffer = sourceMemRef.alignedPtr(rewriter, loc);
MemRefDescriptor dstMemRef = MemRefDescriptor::fromStaticShape(
rewriter, loc, *getTypeConverter(),
cast<MemRefType>(extractStridedMetadataOp.getBaseBuffer().getType()),
baseBuffer, alignedBuffer);
results.push_back((Value)dstMemRef);
results.push_back(sourceMemRef.offset(rewriter, loc));
for (unsigned i = 0; i < rank; ++i)
results.push_back(sourceMemRef.size(rewriter, loc, i));
for (unsigned i = 0; i < rank; ++i)
results.push_back(sourceMemRef.stride(rewriter, loc, i));
rewriter.replaceOp(extractStridedMetadataOp, results);
return success();
}
};
}
void mlir::populateFinalizeMemRefToLLVMConversionPatterns(
const LLVMTypeConverter &converter, RewritePatternSet &patterns,
SymbolTableCollection *symbolTables) {
patterns.add<
AllocaOpLowering,
AllocaScopeOpLowering,
AssumeAlignmentOpLowering,
AtomicRMWOpLowering,
ConvertExtractAlignedPointerAsIndex,
DimOpLowering,
DistinctObjectsOpLowering,
ExtractStridedMetadataOpLowering,
GenericAtomicRMWOpLowering,
GetGlobalMemrefOpLowering,
LoadOpLowering,
MemRefCastOpLowering,
MemRefReinterpretCastOpLowering,
MemRefReshapeOpLowering,
MemorySpaceCastOpLowering,
PrefetchOpLowering,
RankOpLowering,
ReassociatingReshapeOpConversion<memref::CollapseShapeOp>,
ReassociatingReshapeOpConversion<memref::ExpandShapeOp>,
StoreOpLowering,
SubViewOpLowering,
TransposeOpLowering,
ViewOpLowering>(converter);
patterns.add<GlobalMemrefOpLowering, MemRefCopyOpLowering>(converter,
symbolTables);
auto allocLowering = converter.getOptions().allocLowering;
if (allocLowering == LowerToLLVMOptions::AllocLowering::AlignedAlloc)
patterns.add<AlignedAllocOpLowering, DeallocOpLowering>(converter,
symbolTables);
else if (allocLowering == LowerToLLVMOptions::AllocLowering::Malloc)
patterns.add<AllocOpLowering, DeallocOpLowering>(converter, symbolTables);
}
namespace {
struct FinalizeMemRefToLLVMConversionPass
: public impl::FinalizeMemRefToLLVMConversionPassBase<
FinalizeMemRefToLLVMConversionPass> {
using FinalizeMemRefToLLVMConversionPassBase::
FinalizeMemRefToLLVMConversionPassBase;
void runOnOperation() override {
Operation *op = getOperation();
const auto &dataLayoutAnalysis = getAnalysis<DataLayoutAnalysis>();
LowerToLLVMOptions options(&getContext(),
dataLayoutAnalysis.getAtOrAbove(op));
options.allocLowering =
(useAlignedAlloc ? LowerToLLVMOptions::AllocLowering::AlignedAlloc
: LowerToLLVMOptions::AllocLowering::Malloc);
options.useGenericFunctions = useGenericFunctions;
if (indexBitwidth != kDeriveIndexBitwidthFromDataLayout)
options.overrideIndexBitwidth(indexBitwidth);
LLVMTypeConverter typeConverter(&getContext(), options,
&dataLayoutAnalysis);
RewritePatternSet patterns(&getContext());
SymbolTableCollection symbolTables;
populateFinalizeMemRefToLLVMConversionPatterns(typeConverter, patterns,
&symbolTables);
LLVMConversionTarget target(getContext());
target.addLegalOp<func::FuncOp>();
if (failed(applyPartialConversion(op, target, std::move(patterns))))
signalPassFailure();
}
};
struct MemRefToLLVMDialectInterface : public ConvertToLLVMPatternInterface {
using ConvertToLLVMPatternInterface::ConvertToLLVMPatternInterface;
void loadDependentDialects(MLIRContext *context) const final {
context->loadDialect<LLVM::LLVMDialect>();
}
void populateConvertToLLVMConversionPatterns(
ConversionTarget &target, LLVMTypeConverter &typeConverter,
RewritePatternSet &patterns) const final {
populateFinalizeMemRefToLLVMConversionPatterns(typeConverter, patterns);
}
};
}
void mlir::registerConvertMemRefToLLVMInterface(DialectRegistry ®istry) {
registry.addExtension(+[](MLIRContext *ctx, memref::MemRefDialect *dialect) {
dialect->addInterfaces<MemRefToLLVMDialectInterface>();
});
}