27#include "llvm/Support/DebugLog.h"
28#include "llvm/Support/MathExtras.h"
32#define DEBUG_TYPE "memref-to-llvm"
35#define GEN_PASS_DEF_FINALIZEMEMREFTOLLVMCONVERSIONPASS
36#include "mlir/Conversion/Passes.h.inc"
47 auto [strides, offset] = type.getStridesAndOffset();
48 LLVM::GEPNoWrapFlags flags = LLVM::GEPNoWrapFlags::inbounds;
49 if (llvm::all_of(strides, [](
int64_t s) {
50 return !ShapedType::isDynamic(s) && s >= 0;
52 flags = flags | LLVM::GEPNoWrapFlags::nuw;
58static bool isStaticStrideOrOffset(
int64_t strideOrOffset) {
59 return ShapedType::isStatic(strideOrOffset);
62static FailureOr<LLVM::LLVMFuncOp>
73static FailureOr<LLVM::LLVMFuncOp>
85static FailureOr<LLVM::LLVMFuncOp>
101static Value createAligned(ConversionPatternRewriter &rewriter,
Location loc,
104 LLVM::ConstantOp::create(rewriter, loc, alignment.
getType(),
105 rewriter.getIntegerAttr(alignment.
getType(), 1));
106 Value bump = LLVM::SubOp::create(rewriter, loc, alignment, one);
107 Value bumped = LLVM::AddOp::create(rewriter, loc, input, bump);
108 Value mod = LLVM::URemOp::create(rewriter, loc, bumped, alignment);
109 return LLVM::SubOp::create(rewriter, loc, bumped, mod);
121 Type elementType = memRefType.getElementType();
122 if (
auto memRefElementType = dyn_cast<MemRefType>(elementType))
124 if (
auto memRefElementType = dyn_cast<UnrankedMemRefType>(elementType))
130static Value castAllocFuncResult(ConversionPatternRewriter &rewriter,
132 MemRefType memRefType,
Type elementPtrType,
134 auto allocatedPtrTy = cast<LLVM::LLVMPointerType>(allocatedPtr.
getType());
135 FailureOr<unsigned> maybeMemrefAddrSpace =
137 assert(succeeded(maybeMemrefAddrSpace) &&
"unsupported address space");
138 unsigned memrefAddrSpace = *maybeMemrefAddrSpace;
139 if (allocatedPtrTy.getAddressSpace() != memrefAddrSpace)
140 allocatedPtr = LLVM::AddrSpaceCastOp::create(
142 LLVM::LLVMPointerType::get(rewriter.getContext(), memrefAddrSpace),
148 SymbolTableCollection *symbolTables =
nullptr;
151 explicit AllocOpLowering(
const LLVMTypeConverter &typeConverter,
152 SymbolTableCollection *symbolTables =
nullptr,
153 PatternBenefit benefit = 1)
154 : ConvertOpToLLVMPattern<memref::AllocOp>(typeConverter, benefit),
155 symbolTables(symbolTables) {}
158 matchAndRewrite(memref::AllocOp op, OpAdaptor adaptor,
159 ConversionPatternRewriter &rewriter)
const override {
160 auto loc = op.getLoc();
161 MemRefType memRefType = op.getType();
162 if (!isConvertibleAndHasIdentityMaps(memRefType))
163 return rewriter.notifyMatchFailure(op,
"incompatible memref type");
166 FailureOr<LLVM::LLVMFuncOp> allocFuncOp =
167 getNotalignedAllocFn(rewriter, getTypeConverter(),
168 op->getParentWithTrait<OpTrait::SymbolTable>(),
169 getIndexType(), symbolTables);
176 SmallVector<Value, 4> sizes;
177 SmallVector<Value, 4> strides;
180 this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
181 rewriter, sizes, strides, sizeBytes,
true);
183 Value alignment = getAlignment(rewriter, loc, op);
186 sizeBytes = LLVM::AddOp::create(rewriter, loc, sizeBytes, alignment);
190 Type elementPtrType = this->getElementPtrType(memRefType);
191 assert(elementPtrType &&
"could not compute element ptr type");
193 LLVM::CallOp::create(rewriter, loc, allocFuncOp.value(), sizeBytes);
196 castAllocFuncResult(rewriter, loc, results.getResult(), memRefType,
197 elementPtrType, *getTypeConverter());
198 Value alignedPtr = allocatedPtr;
202 LLVM::PtrToIntOp::create(rewriter, loc, getIndexType(), allocatedPtr);
204 createAligned(rewriter, loc, allocatedInt, alignment);
206 LLVM::IntToPtrOp::create(rewriter, loc, elementPtrType, alignmentInt);
210 auto memRefDescriptor = this->createMemRefDescriptor(
211 loc, memRefType, allocatedPtr, alignedPtr, sizes, strides, rewriter);
214 rewriter.replaceOp(op, {memRefDescriptor});
219 template <
typename OpType>
220 Value getAlignment(ConversionPatternRewriter &rewriter, Location loc,
222 MemRefType memRefType = op.getType();
224 if (
auto alignmentAttr = op.getAlignment()) {
225 Type indexType = getIndexType();
228 }
else if (!memRefType.getElementType().isSignlessIntOrIndexOrFloat()) {
233 alignment =
getSizeInBytes(loc, memRefType.getElementType(), rewriter);
240 SymbolTableCollection *symbolTables =
nullptr;
243 explicit AlignedAllocOpLowering(
const LLVMTypeConverter &typeConverter,
244 SymbolTableCollection *symbolTables =
nullptr,
245 PatternBenefit benefit = 1)
246 : ConvertOpToLLVMPattern<memref::AllocOp>(typeConverter, benefit),
247 symbolTables(symbolTables) {}
250 matchAndRewrite(memref::AllocOp op, OpAdaptor adaptor,
251 ConversionPatternRewriter &rewriter)
const override {
252 auto loc = op.getLoc();
253 MemRefType memRefType = op.getType();
254 if (!isConvertibleAndHasIdentityMaps(memRefType))
255 return rewriter.notifyMatchFailure(op,
"incompatible memref type");
258 FailureOr<LLVM::LLVMFuncOp> allocFuncOp =
259 getAlignedAllocFn(rewriter, getTypeConverter(),
260 op->getParentWithTrait<OpTrait::SymbolTable>(),
261 getIndexType(), symbolTables);
268 SmallVector<Value, 4> sizes;
269 SmallVector<Value, 4> strides;
272 this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
273 rewriter, sizes, strides, sizeBytes, !
false);
275 int64_t alignment = alignedAllocationGetAlignment(op, &defaultLayout);
277 Value allocAlignment =
282 if (!isMemRefSizeMultipleOf(memRefType, alignment, op, &defaultLayout))
283 sizeBytes = createAligned(rewriter, loc, sizeBytes, allocAlignment);
285 Type elementPtrType = this->getElementPtrType(memRefType);
287 LLVM::CallOp::create(rewriter, loc, allocFuncOp.value(),
291 castAllocFuncResult(rewriter, loc, results.getResult(), memRefType,
292 elementPtrType, *getTypeConverter());
295 auto memRefDescriptor = this->createMemRefDescriptor(
296 loc, memRefType, ptr, ptr, sizes, strides, rewriter);
299 rewriter.replaceOp(op, {memRefDescriptor});
304 static constexpr uint64_t kMinAlignedAllocAlignment = 16UL;
311 int64_t alignedAllocationGetAlignment(memref::AllocOp op,
312 const DataLayout *defaultLayout)
const {
313 if (std::optional<uint64_t> alignment = op.getAlignment())
319 unsigned eltSizeBytes = getMemRefEltSizeInBytes(
320 getTypeConverter(), op.getType(), op, defaultLayout);
321 return std::max(kMinAlignedAllocAlignment,
322 llvm::PowerOf2Ceil(eltSizeBytes));
327 bool isMemRefSizeMultipleOf(MemRefType type, uint64_t factor, Operation *op,
328 const DataLayout *defaultLayout)
const {
329 uint64_t sizeDivisor =
330 getMemRefEltSizeInBytes(getTypeConverter(), type, op, defaultLayout);
331 for (
unsigned i = 0, e = type.getRank(); i < e; i++) {
332 if (type.isDynamicDim(i))
334 sizeDivisor = sizeDivisor * type.getDimSize(i);
336 return sizeDivisor % factor == 0;
341 DataLayout defaultLayout;
345 using ConvertOpToLLVMPattern<memref::AllocaOp>::ConvertOpToLLVMPattern;
351 matchAndRewrite(memref::AllocaOp op, OpAdaptor adaptor,
352 ConversionPatternRewriter &rewriter)
const override {
353 auto loc = op.getLoc();
354 MemRefType memRefType = op.getType();
355 if (!isConvertibleAndHasIdentityMaps(memRefType))
356 return rewriter.notifyMatchFailure(op,
"incompatible memref type");
361 SmallVector<Value, 4> sizes;
362 SmallVector<Value, 4> strides;
365 this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
366 rewriter, sizes, strides, size, !
true);
371 typeConverter->convertType(op.getType().getElementType());
372 FailureOr<unsigned> maybeAddressSpace =
374 assert(succeeded(maybeAddressSpace) &&
"unsupported address space");
375 unsigned addrSpace = *maybeAddressSpace;
376 auto elementPtrType =
377 LLVM::LLVMPointerType::get(rewriter.getContext(), addrSpace);
379 auto allocatedElementPtr =
380 LLVM::AllocaOp::create(rewriter, loc, elementPtrType, elementType, size,
381 op.getAlignment().value_or(0));
384 auto memRefDescriptor = this->createMemRefDescriptor(
385 loc, memRefType, allocatedElementPtr, allocatedElementPtr, sizes,
389 rewriter.replaceOp(op, {memRefDescriptor});
394struct AllocaScopeOpLowering
396 using ConvertOpToLLVMPattern<memref::AllocaScopeOp>::ConvertOpToLLVMPattern;
399 matchAndRewrite(memref::AllocaScopeOp allocaScopeOp, OpAdaptor adaptor,
400 ConversionPatternRewriter &rewriter)
const override {
401 OpBuilder::InsertionGuard guard(rewriter);
402 Location loc = allocaScopeOp.getLoc();
406 auto *currentBlock = rewriter.getInsertionBlock();
407 auto *remainingOpsBlock =
408 rewriter.splitBlock(currentBlock, rewriter.getInsertionPoint());
409 Block *continueBlock;
410 if (allocaScopeOp.getNumResults() == 0) {
411 continueBlock = remainingOpsBlock;
413 continueBlock = rewriter.createBlock(
414 remainingOpsBlock, allocaScopeOp.getResultTypes(),
415 SmallVector<Location>(allocaScopeOp->getNumResults(),
416 allocaScopeOp.getLoc()));
417 LLVM::BrOp::create(rewriter, loc,
ValueRange(), remainingOpsBlock);
421 Block *beforeBody = &allocaScopeOp.getBodyRegion().front();
422 Block *afterBody = &allocaScopeOp.getBodyRegion().back();
423 rewriter.inlineRegionBefore(allocaScopeOp.getBodyRegion(), continueBlock);
426 rewriter.setInsertionPointToEnd(currentBlock);
427 auto stackSaveOp = LLVM::StackSaveOp::create(rewriter, loc, getPtrType());
428 LLVM::BrOp::create(rewriter, loc,
ValueRange(), beforeBody);
432 rewriter.setInsertionPointToEnd(afterBody);
434 cast<memref::AllocaScopeReturnOp>(afterBody->
getTerminator());
435 auto branchOp = rewriter.replaceOpWithNewOp<LLVM::BrOp>(
436 returnOp, returnOp.getResults(), continueBlock);
439 rewriter.setInsertionPoint(branchOp);
440 LLVM::StackRestoreOp::create(rewriter, loc, stackSaveOp);
443 rewriter.replaceOp(allocaScopeOp, continueBlock->
getArguments());
449struct AssumeAlignmentOpLowering
451 using ConvertOpToLLVMPattern<
452 memref::AssumeAlignmentOp>::ConvertOpToLLVMPattern;
453 explicit AssumeAlignmentOpLowering(
const LLVMTypeConverter &converter)
454 : ConvertOpToLLVMPattern<memref::AssumeAlignmentOp>(converter) {}
457 matchAndRewrite(memref::AssumeAlignmentOp op, OpAdaptor adaptor,
458 ConversionPatternRewriter &rewriter)
const override {
459 Value memref = adaptor.getMemref();
460 unsigned alignment = op.getAlignment();
461 auto loc = op.getLoc();
463 auto srcMemRefType = cast<MemRefType>(op.getMemref().getType());
471 LLVM::ConstantOp::create(rewriter, loc, rewriter.getBoolAttr(
true));
472 Value alignmentConst =
474 LLVM::AssumeOp::create(rewriter, loc, trueCond, LLVM::AssumeAlignTag(), ptr,
476 rewriter.replaceOp(op, memref);
481struct DistinctObjectsOpLowering
483 using ConvertOpToLLVMPattern<
484 memref::DistinctObjectsOp>::ConvertOpToLLVMPattern;
485 explicit DistinctObjectsOpLowering(
const LLVMTypeConverter &converter)
486 : ConvertOpToLLVMPattern<memref::DistinctObjectsOp>(converter) {}
489 matchAndRewrite(memref::DistinctObjectsOp op, OpAdaptor adaptor,
490 ConversionPatternRewriter &rewriter)
const override {
492 if (operands.size() <= 1) {
494 rewriter.replaceOp(op, operands);
498 Location loc = op.getLoc();
499 SmallVector<Value> ptrs;
500 for (
auto [origOperand, newOperand] :
501 llvm::zip_equal(op.getOperands(), operands)) {
502 auto memrefType = cast<MemRefType>(origOperand.getType());
503 MemRefDescriptor memRefDescriptor(newOperand);
504 Value ptr = memRefDescriptor.bufferPtr(rewriter, loc, *getTypeConverter(),
510 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI1Type(), 1);
512 for (
auto i : llvm::seq<size_t>(ptrs.size() - 1)) {
513 for (
auto j : llvm::seq<size_t>(i + 1, ptrs.size())) {
514 Value ptr1 = ptrs[i];
515 Value ptr2 = ptrs[j];
516 LLVM::AssumeOp::create(rewriter, loc, cond,
517 LLVM::AssumeSeparateStorageTag{}, ptr1, ptr2);
521 rewriter.replaceOp(op, operands);
530 SymbolTableCollection *symbolTables =
nullptr;
533 explicit DeallocOpLowering(
const LLVMTypeConverter &typeConverter,
534 SymbolTableCollection *symbolTables =
nullptr,
535 PatternBenefit benefit = 1)
536 : ConvertOpToLLVMPattern<memref::DeallocOp>(typeConverter, benefit),
537 symbolTables(symbolTables) {}
540 matchAndRewrite(memref::DeallocOp op, OpAdaptor adaptor,
541 ConversionPatternRewriter &rewriter)
const override {
543 FailureOr<LLVM::LLVMFuncOp> freeFunc =
544 getFreeFn(rewriter, getTypeConverter(),
545 op->getParentWithTrait<OpTrait::SymbolTable>(), symbolTables);
549 if (
auto unrankedTy =
550 llvm::dyn_cast<UnrankedMemRefType>(op.getMemref().getType())) {
551 auto elementPtrTy = LLVM::LLVMPointerType::get(
552 rewriter.getContext(), unrankedTy.getMemorySpaceAsInt());
554 rewriter, op.getLoc(),
555 UnrankedMemRefDescriptor(adaptor.getMemref())
556 .memRefDescPtr(rewriter, op.getLoc()),
559 allocatedPtr = MemRefDescriptor(adaptor.getMemref())
560 .allocatedPtr(rewriter, op.getLoc());
562 rewriter.replaceOpWithNewOp<LLVM::CallOp>(op, freeFunc.value(),
571 using ConvertOpToLLVMPattern<memref::DimOp>::ConvertOpToLLVMPattern;
574 matchAndRewrite(memref::DimOp dimOp, OpAdaptor adaptor,
575 ConversionPatternRewriter &rewriter)
const override {
576 Type operandType = dimOp.getSource().getType();
577 if (isa<UnrankedMemRefType>(operandType)) {
578 FailureOr<Value> extractedSize = extractSizeOfUnrankedMemRef(
579 operandType, dimOp, adaptor.getOperands(), rewriter);
580 if (
failed(extractedSize))
582 rewriter.replaceOp(dimOp, {*extractedSize});
585 if (isa<MemRefType>(operandType)) {
587 dimOp, {extractSizeOfRankedMemRef(operandType, dimOp,
588 adaptor.getOperands(), rewriter)});
591 llvm_unreachable(
"expected MemRefType or UnrankedMemRefType");
596 extractSizeOfUnrankedMemRef(Type operandType, memref::DimOp dimOp,
598 ConversionPatternRewriter &rewriter)
const {
599 Location loc = dimOp.getLoc();
601 auto unrankedMemRefType = cast<UnrankedMemRefType>(operandType);
602 auto scalarMemRefType =
603 MemRefType::get({}, unrankedMemRefType.getElementType());
604 FailureOr<unsigned> maybeAddressSpace =
605 getTypeConverter()->getMemRefAddressSpace(unrankedMemRefType);
606 if (
failed(maybeAddressSpace)) {
607 dimOp.emitOpError(
"memref memory space must be convertible to an integer "
611 unsigned addressSpace = *maybeAddressSpace;
616 UnrankedMemRefDescriptor unrankedDesc(adaptor.getSource());
617 Value underlyingRankedDesc = unrankedDesc.memRefDescPtr(rewriter, loc);
619 Type elementType = typeConverter->convertType(scalarMemRefType);
623 LLVM::LLVMPointerType::get(rewriter.getContext(), addressSpace);
625 LLVM::GEPOp::create(rewriter, loc, indexPtrTy, elementType,
626 underlyingRankedDesc, ArrayRef<LLVM::GEPArg>{0, 2});
630 Value idxPlusOne = LLVM::AddOp::create(
634 Value sizePtr = LLVM::GEPOp::create(rewriter, loc, indexPtrTy,
635 getTypeConverter()->getIndexType(),
636 offsetPtr, idxPlusOne);
637 return LLVM::LoadOp::create(rewriter, loc,
638 getTypeConverter()->getIndexType(), sizePtr)
642 std::optional<int64_t> getConstantDimIndex(memref::DimOp dimOp)
const {
643 if (
auto idx = dimOp.getConstantIndex())
646 if (
auto constantOp = dimOp.getIndex().getDefiningOp<LLVM::ConstantOp>())
647 return cast<IntegerAttr>(constantOp.getValue()).getValue().getSExtValue();
652 Value extractSizeOfRankedMemRef(Type operandType, memref::DimOp dimOp,
654 ConversionPatternRewriter &rewriter)
const {
655 Location loc = dimOp.getLoc();
658 MemRefType memRefType = cast<MemRefType>(operandType);
659 Type indexType = getIndexType();
660 if (std::optional<int64_t> index = getConstantDimIndex(dimOp)) {
662 if (i >= 0 && i < memRefType.getRank()) {
663 if (memRefType.isDynamicDim(i)) {
665 MemRefDescriptor descriptor(adaptor.getSource());
666 return descriptor.size(rewriter, loc, i);
669 int64_t dimSize = memRefType.getDimSize(i);
673 Value index = adaptor.getIndex();
674 int64_t rank = memRefType.getRank();
675 MemRefDescriptor memrefDescriptor(adaptor.getSource());
676 return memrefDescriptor.size(rewriter, loc, index, rank);
683template <
typename Derived>
685 using ConvertOpToLLVMPattern<Derived>::ConvertOpToLLVMPattern;
686 using ConvertOpToLLVMPattern<Derived>::isConvertibleAndHasIdentityMaps;
687 using Base = LoadStoreOpLowering<Derived>;
717struct GenericAtomicRMWOpLowering
718 :
public LoadStoreOpLowering<memref::GenericAtomicRMWOp> {
722 matchAndRewrite(memref::GenericAtomicRMWOp atomicOp, OpAdaptor adaptor,
723 ConversionPatternRewriter &rewriter)
const override {
724 auto loc = atomicOp.getLoc();
725 Type valueType = typeConverter->convertType(atomicOp.getResult().getType());
730 bool needsBitcast = isa<FloatType>(valueType);
731 Type cmpxchgType = valueType;
733 unsigned bitWidth = cast<FloatType>(valueType).getWidth();
734 cmpxchgType = rewriter.getIntegerType(bitWidth);
738 auto *initBlock = rewriter.getInsertionBlock();
739 auto *loopBlock = rewriter.splitBlock(initBlock,
Block::iterator(atomicOp));
740 loopBlock->addArgument(cmpxchgType, loc);
746 rewriter.setInsertionPointToEnd(initBlock);
747 auto memRefType = cast<MemRefType>(atomicOp.getMemref().getType());
749 rewriter, loc, memRefType, adaptor.getMemref(), adaptor.getIndices());
750 Value init = LLVM::LoadOp::create(
751 rewriter, loc, typeConverter->convertType(memRefType.getElementType()),
754 init = LLVM::BitcastOp::create(rewriter, loc, cmpxchgType, init);
755 LLVM::BrOp::create(rewriter, loc, init, loopBlock);
758 rewriter.setInsertionPointToStart(loopBlock);
761 Value loopArgument = loopBlock->getArgument(0);
762 Value loopArgForBody = loopArgument;
765 LLVM::BitcastOp::create(rewriter, loc, valueType, loopArgument);
767 mapping.
map(atomicOp.getCurrentValue(), loopArgForBody);
768 Block &entryBlock = atomicOp.body().front();
770 Operation *
clone = rewriter.clone(nestedOp, mapping);
777 return atomicOp.emitError(
"result not defined in region");
780 result = LLVM::BitcastOp::create(rewriter, loc, cmpxchgType,
result);
784 auto successOrdering = LLVM::AtomicOrdering::acq_rel;
785 auto failureOrdering = LLVM::AtomicOrdering::monotonic;
787 LLVM::AtomicCmpXchgOp::create(rewriter, loc, dataPtr, loopArgument,
788 result, successOrdering, failureOrdering);
790 Value newLoaded = LLVM::ExtractValueOp::create(rewriter, loc, cmpxchg, 0);
791 Value ok = LLVM::ExtractValueOp::create(rewriter, loc, cmpxchg, 1);
794 LLVM::CondBrOp::create(rewriter, loc, ok, endBlock, ArrayRef<Value>(),
795 loopBlock, newLoaded);
801 rewriter.setInsertionPointToStart(endBlock);
802 newLoaded = LLVM::BitcastOp::create(rewriter, loc, valueType, newLoaded);
804 rewriter.setInsertionPointToEnd(endBlock);
805 rewriter.replaceOp(atomicOp, {newLoaded});
813convertGlobalMemrefTypeToLLVM(MemRefType type,
820 Type elementType = typeConverter.convertType(type.getElementType());
821 Type arrayTy = elementType;
823 for (
int64_t dim : llvm::reverse(type.getShape()))
824 arrayTy = LLVM::LLVMArrayType::get(arrayTy, dim);
830 SymbolTableCollection *symbolTables =
nullptr;
833 explicit GlobalMemrefOpLowering(
const LLVMTypeConverter &typeConverter,
834 SymbolTableCollection *symbolTables =
nullptr,
835 PatternBenefit benefit = 1)
836 : ConvertOpToLLVMPattern<memref::GlobalOp>(typeConverter, benefit),
837 symbolTables(symbolTables) {}
840 matchAndRewrite(memref::GlobalOp global, OpAdaptor adaptor,
841 ConversionPatternRewriter &rewriter)
const override {
842 MemRefType type = global.
getType();
843 if (!isConvertibleAndHasIdentityMaps(type))
846 Type arrayTy = convertGlobalMemrefTypeToLLVM(type, *getTypeConverter());
848 LLVM::Linkage linkage =
849 global.isPublic() ? LLVM::Linkage::External : LLVM::Linkage::Private;
850 bool isExternal = global.isExternal();
851 bool isUninitialized = global.isUninitialized();
853 Attribute initialValue =
nullptr;
854 if (!isExternal && !isUninitialized) {
855 auto elementsAttr = llvm::cast<ElementsAttr>(*global.getInitialValue());
856 initialValue = elementsAttr;
860 if (type.getRank() == 0)
861 initialValue = elementsAttr.getSplatValue<Attribute>();
864 uint64_t alignment = global.getAlignment().value_or(0);
865 FailureOr<unsigned> addressSpace =
866 getTypeConverter()->getMemRefAddressSpace(type);
868 return global.emitOpError(
869 "memory space cannot be converted to an integer address space");
872 SymbolTable *symbolTable =
nullptr;
874 Operation *symbolTableOp =
875 global->getParentWithTrait<OpTrait::SymbolTable>();
876 symbolTable = &symbolTables->getSymbolTable(symbolTableOp);
877 symbolTable->remove(global);
881 auto newGlobal = rewriter.replaceOpWithNewOp<LLVM::GlobalOp>(
882 global, arrayTy, global.getConstant(), linkage, global.getSymName(),
883 initialValue, alignment, *addressSpace);
887 symbolTable->insert(newGlobal, rewriter.getInsertionPoint());
889 if (!isExternal && isUninitialized) {
890 rewriter.createBlock(&newGlobal.getInitializerRegion());
892 LLVM::UndefOp::create(rewriter, newGlobal.getLoc(), arrayTy)};
893 LLVM::ReturnOp::create(rewriter, newGlobal.getLoc(), undef);
902struct GetGlobalMemrefOpLowering
904 using ConvertOpToLLVMPattern<memref::GetGlobalOp>::ConvertOpToLLVMPattern;
909 matchAndRewrite(memref::GetGlobalOp op, OpAdaptor adaptor,
910 ConversionPatternRewriter &rewriter)
const override {
911 auto loc = op.getLoc();
912 MemRefType memRefType = op.getType();
913 if (!isConvertibleAndHasIdentityMaps(memRefType))
914 return rewriter.notifyMatchFailure(op,
"incompatible memref type");
919 SmallVector<Value, 4> sizes;
920 SmallVector<Value, 4> strides;
923 this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
924 rewriter, sizes, strides, sizeBytes, !
false);
926 MemRefType type = cast<MemRefType>(op.getResult().getType());
930 FailureOr<unsigned> maybeAddressSpace =
931 getTypeConverter()->getMemRefAddressSpace(type);
932 assert(succeeded(maybeAddressSpace) &&
"unsupported address space");
933 unsigned memSpace = *maybeAddressSpace;
935 Type arrayTy = convertGlobalMemrefTypeToLLVM(type, *getTypeConverter());
936 auto ptrTy = LLVM::LLVMPointerType::get(rewriter.getContext(), memSpace);
938 LLVM::AddressOfOp::create(rewriter, loc, ptrTy, op.getName());
943 LLVM::GEPOp::create(rewriter, loc, ptrTy, arrayTy, addressOf,
944 SmallVector<LLVM::GEPArg>(type.getRank() + 1, 0));
949 auto intPtrType = getIntPtrType(memSpace);
950 Value deadBeefConst =
953 LLVM::IntToPtrOp::create(rewriter, loc, ptrTy, deadBeefConst);
958 auto memRefDescriptor = this->createMemRefDescriptor(
959 loc, memRefType, deadBeefPtr, gep, sizes, strides, rewriter);
962 rewriter.replaceOp(op, {memRefDescriptor});
969struct LoadOpLowering :
public LoadStoreOpLowering<memref::LoadOp> {
973 matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
974 ConversionPatternRewriter &rewriter)
const override {
975 auto type = loadOp.getMemRefType();
981 rewriter, loadOp.getLoc(), type, adaptor.getMemref(),
983 rewriter.replaceOpWithNewOp<LLVM::LoadOp>(
984 loadOp, typeConverter->convertType(type.getElementType()), dataPtr,
985 loadOp.getAlignment().value_or(0),
false, loadOp.getNontemporal(),
986 loadOp.getInvariant());
993struct StoreOpLowering :
public LoadStoreOpLowering<memref::StoreOp> {
997 matchAndRewrite(memref::StoreOp op, OpAdaptor adaptor,
998 ConversionPatternRewriter &rewriter)
const override {
999 auto type = op.getMemRefType();
1005 rewriter, op.getLoc(), type, adaptor.getMemref(), adaptor.getIndices(),
1007 rewriter.replaceOpWithNewOp<LLVM::StoreOp>(op, adaptor.getValue(), dataPtr,
1008 op.getAlignment().value_or(0),
1009 false, op.getNontemporal());
1016struct PrefetchOpLowering :
public LoadStoreOpLowering<memref::PrefetchOp> {
1020 matchAndRewrite(memref::PrefetchOp prefetchOp, OpAdaptor adaptor,
1021 ConversionPatternRewriter &rewriter)
const override {
1022 auto type = prefetchOp.getMemRefType();
1023 auto loc = prefetchOp.getLoc();
1026 rewriter, loc, type, adaptor.getMemref(), adaptor.getIndices());
1029 IntegerAttr isWrite = rewriter.getI32IntegerAttr(prefetchOp.getIsWrite());
1030 IntegerAttr localityHint = prefetchOp.getLocalityHintAttr();
1031 IntegerAttr isData =
1032 rewriter.getI32IntegerAttr(prefetchOp.getIsDataCache());
1033 rewriter.replaceOpWithNewOp<LLVM::Prefetch>(prefetchOp, dataPtr, isWrite,
1034 localityHint, isData);
1040 using ConvertOpToLLVMPattern<memref::RankOp>::ConvertOpToLLVMPattern;
1043 matchAndRewrite(memref::RankOp op, OpAdaptor adaptor,
1044 ConversionPatternRewriter &rewriter)
const override {
1045 Location loc = op.getLoc();
1046 Type operandType = op.getMemref().getType();
1047 if (isa<UnrankedMemRefType>(operandType)) {
1048 UnrankedMemRefDescriptor desc(adaptor.getMemref());
1049 rewriter.replaceOp(op, {desc.rank(rewriter, loc)});
1052 if (
auto rankedMemRefType = dyn_cast<MemRefType>(operandType)) {
1053 Type indexType = getIndexType();
1054 rewriter.replaceOp(op,
1056 rankedMemRefType.getRank())});
1067 matchAndRewrite(memref::CastOp memRefCastOp, OpAdaptor adaptor,
1068 ConversionPatternRewriter &rewriter)
const override {
1069 Type srcType = memRefCastOp.getOperand().getType();
1070 Type dstType = memRefCastOp.getType();
1077 if (isa<MemRefType>(srcType) && isa<MemRefType>(dstType))
1078 if (typeConverter->convertType(srcType) !=
1079 typeConverter->convertType(dstType))
1083 if (isa<UnrankedMemRefType>(srcType) && isa<UnrankedMemRefType>(dstType))
1086 auto targetStructType = typeConverter->convertType(memRefCastOp.getType());
1087 auto loc = memRefCastOp.getLoc();
1090 if (isa<MemRefType>(srcType) && isa<MemRefType>(dstType)) {
1091 rewriter.replaceOp(memRefCastOp, {adaptor.getSource()});
1095 if (isa<MemRefType>(srcType) && isa<UnrankedMemRefType>(dstType)) {
1100 auto srcMemRefType = cast<MemRefType>(srcType);
1101 int64_t rank = srcMemRefType.getRank();
1103 auto ptr = getTypeConverter()->promoteOneMemRefDescriptor(
1104 loc, adaptor.getSource(), rewriter);
1110 UnrankedMemRefDescriptor memRefDesc =
1113 memRefDesc.
setRank(rewriter, loc, rankVal);
1116 rewriter.replaceOp(memRefCastOp, (Value)memRefDesc);
1118 }
else if (isa<UnrankedMemRefType>(srcType) && isa<MemRefType>(dstType)) {
1122 UnrankedMemRefDescriptor memRefDesc(adaptor.getSource());
1127 auto loadOp = LLVM::LoadOp::create(rewriter, loc, targetStructType, ptr);
1128 rewriter.replaceOp(memRefCastOp, loadOp.getResult());
1130 llvm_unreachable(
"Unsupported unranked memref to unranked memref cast");
1143 SymbolTableCollection *symbolTables =
nullptr;
1146 explicit MemRefCopyOpLowering(
const LLVMTypeConverter &typeConverter,
1147 SymbolTableCollection *symbolTables =
nullptr,
1148 PatternBenefit benefit = 1)
1149 : ConvertOpToLLVMPattern<memref::CopyOp>(typeConverter, benefit),
1150 symbolTables(symbolTables) {}
1153 lowerToMemCopyIntrinsic(memref::CopyOp op, OpAdaptor adaptor,
1154 ConversionPatternRewriter &rewriter)
const {
1155 auto loc = op.getLoc();
1156 auto srcType = dyn_cast<MemRefType>(op.getSource().getType());
1158 MemRefDescriptor srcDesc(adaptor.getSource());
1163 for (
int pos = 0; pos < srcType.getRank(); ++pos) {
1164 auto size = srcDesc.size(rewriter, loc, pos);
1165 numElements = LLVM::MulOp::create(rewriter, loc, numElements, size);
1169 auto sizeInBytes =
getSizeInBytes(loc, srcType.getElementType(), rewriter);
1172 LLVM::MulOp::create(rewriter, loc, numElements, sizeInBytes);
1174 Type elementType = typeConverter->convertType(srcType.getElementType());
1176 Value srcBasePtr = srcDesc.alignedPtr(rewriter, loc);
1177 Value srcOffset = srcDesc.offset(rewriter, loc);
1178 Value srcPtr = LLVM::GEPOp::create(rewriter, loc, srcBasePtr.
getType(),
1179 elementType, srcBasePtr, srcOffset);
1180 MemRefDescriptor targetDesc(adaptor.getTarget());
1181 Value targetBasePtr = targetDesc.alignedPtr(rewriter, loc);
1182 Value targetOffset = targetDesc.offset(rewriter, loc);
1184 LLVM::GEPOp::create(rewriter, loc, targetBasePtr.
getType(), elementType,
1185 targetBasePtr, targetOffset);
1186 LLVM::MemcpyOp::create(rewriter, loc, targetPtr, srcPtr, totalSize,
1188 rewriter.eraseOp(op);
1194 lowerToMemCopyFunctionCall(memref::CopyOp op, OpAdaptor adaptor,
1195 ConversionPatternRewriter &rewriter)
const {
1196 auto loc = op.getLoc();
1197 auto srcType = cast<BaseMemRefType>(op.getSource().getType());
1198 auto targetType = cast<BaseMemRefType>(op.getTarget().getType());
1201 auto makeUnranked = [&,
this](Value ranked, MemRefType type) {
1202 auto rank = LLVM::ConstantOp::create(rewriter, loc, getIndexType(),
1204 auto *typeConverter = getTypeConverter();
1209 UnrankedMemRefType::get(type.getElementType(), type.getMemorySpace());
1211 rewriter, loc, *typeConverter, unrankedType,
ValueRange{rank, ptr});
1215 auto stackSaveOp = LLVM::StackSaveOp::create(rewriter, loc, getPtrType());
1217 auto srcMemRefType = dyn_cast<MemRefType>(srcType);
1218 Value unrankedSource =
1219 srcMemRefType ? makeUnranked(adaptor.getSource(), srcMemRefType)
1220 : adaptor.getSource();
1221 auto targetMemRefType = dyn_cast<MemRefType>(targetType);
1222 Value unrankedTarget =
1223 targetMemRefType ? makeUnranked(adaptor.getTarget(), targetMemRefType)
1224 : adaptor.getTarget();
1228 auto promote = [&](Value desc) {
1229 auto ptrType = LLVM::LLVMPointerType::get(rewriter.getContext());
1231 LLVM::AllocaOp::create(rewriter, loc, ptrType, desc.getType(), one);
1232 LLVM::StoreOp::create(rewriter, loc, desc, allocated);
1236 auto sourcePtr =
promote(unrankedSource);
1237 auto targetPtr =
promote(unrankedTarget);
1241 auto elemSize =
getSizeInBytes(loc, srcType.getElementType(), rewriter);
1242 auto copyFn = LLVM::lookupOrCreateMemRefCopyFn(
1243 rewriter, op->getParentOfType<ModuleOp>(), getIndexType(),
1244 sourcePtr.getType(), symbolTables);
1247 LLVM::CallOp::create(rewriter, loc, copyFn.value(),
1251 LLVM::StackRestoreOp::create(rewriter, loc, stackSaveOp);
1253 rewriter.eraseOp(op);
1259 matchAndRewrite(memref::CopyOp op, OpAdaptor adaptor,
1260 ConversionPatternRewriter &rewriter)
const override {
1261 auto srcType = cast<BaseMemRefType>(op.getSource().getType());
1262 auto targetType = cast<BaseMemRefType>(op.getTarget().getType());
1264 auto isContiguousMemrefType = [&](BaseMemRefType type) {
1265 auto memrefType = dyn_cast<mlir::MemRefType>(type);
1269 return memrefType &&
1270 (memrefType.getLayout().isIdentity() ||
1271 (memrefType.hasStaticShape() && memrefType.getNumElements() > 0 &&
1275 if (isContiguousMemrefType(srcType) && isContiguousMemrefType(targetType))
1276 return lowerToMemCopyIntrinsic(op, adaptor, rewriter);
1278 return lowerToMemCopyFunctionCall(op, adaptor, rewriter);
1282struct MemorySpaceCastOpLowering
1284 using ConvertOpToLLVMPattern<
1285 memref::MemorySpaceCastOp>::ConvertOpToLLVMPattern;
1288 matchAndRewrite(memref::MemorySpaceCastOp op, OpAdaptor adaptor,
1289 ConversionPatternRewriter &rewriter)
const override {
1290 Location loc = op.getLoc();
1292 Type resultType = op.getDest().getType();
1293 if (
auto resultTypeR = dyn_cast<MemRefType>(resultType)) {
1294 auto convertedType =
1295 typeConverter->convertType<LLVM::LLVMStructType>(resultTypeR);
1297 return rewriter.notifyMatchFailure(op,
"memref type conversion failed");
1298 Type newPtrType = convertedType.getBody()[0];
1300 SmallVector<Value> descVals;
1301 MemRefDescriptor::unpack(rewriter, loc, adaptor.getSource(), resultTypeR,
1304 LLVM::AddrSpaceCastOp::create(rewriter, loc, newPtrType, descVals[0]);
1306 LLVM::AddrSpaceCastOp::create(rewriter, loc, newPtrType, descVals[1]);
1307 Value
result = MemRefDescriptor::pack(rewriter, loc, *getTypeConverter(),
1308 resultTypeR, descVals);
1309 rewriter.replaceOp(op,
result);
1312 if (
auto resultTypeU = dyn_cast<UnrankedMemRefType>(resultType)) {
1315 auto sourceType = cast<UnrankedMemRefType>(op.getSource().getType());
1316 FailureOr<unsigned> maybeSourceAddrSpace =
1317 getTypeConverter()->getMemRefAddressSpace(sourceType);
1318 if (
failed(maybeSourceAddrSpace))
1319 return rewriter.notifyMatchFailure(loc,
1320 "non-integer source address space");
1321 unsigned sourceAddrSpace = *maybeSourceAddrSpace;
1322 FailureOr<unsigned> maybeResultAddrSpace =
1323 getTypeConverter()->getMemRefAddressSpace(resultTypeU);
1324 if (
failed(maybeResultAddrSpace))
1325 return rewriter.notifyMatchFailure(loc,
1326 "non-integer result address space");
1327 unsigned resultAddrSpace = *maybeResultAddrSpace;
1329 UnrankedMemRefDescriptor sourceDesc(adaptor.getSource());
1330 Value rank = sourceDesc.rank(rewriter, loc);
1331 Value sourceUnderlyingDesc = sourceDesc.memRefDescPtr(rewriter, loc);
1335 rewriter, loc, typeConverter->convertType(resultTypeU));
1336 result.setRank(rewriter, loc, rank);
1338 rewriter, loc, *getTypeConverter(),
result, resultAddrSpace);
1339 Value resultUnderlyingDesc =
1340 LLVM::AllocaOp::create(rewriter, loc, getPtrType(),
1341 rewriter.getI8Type(), resultUnderlyingSize);
1342 result.setMemRefDescPtr(rewriter, loc, resultUnderlyingDesc);
1345 auto sourceElemPtrType =
1346 LLVM::LLVMPointerType::get(rewriter.getContext(), sourceAddrSpace);
1347 auto resultElemPtrType =
1348 LLVM::LLVMPointerType::get(rewriter.getContext(), resultAddrSpace);
1350 Value allocatedPtr = sourceDesc.allocatedPtr(
1351 rewriter, loc, sourceUnderlyingDesc, sourceElemPtrType);
1353 sourceDesc.alignedPtr(rewriter, loc, *getTypeConverter(),
1354 sourceUnderlyingDesc, sourceElemPtrType);
1355 allocatedPtr = LLVM::AddrSpaceCastOp::create(
1356 rewriter, loc, resultElemPtrType, allocatedPtr);
1357 alignedPtr = LLVM::AddrSpaceCastOp::create(rewriter, loc,
1358 resultElemPtrType, alignedPtr);
1360 result.setAllocatedPtr(rewriter, loc, resultUnderlyingDesc,
1361 resultElemPtrType, allocatedPtr);
1362 result.setAlignedPtr(rewriter, loc, *getTypeConverter(),
1363 resultUnderlyingDesc, resultElemPtrType, alignedPtr);
1366 Value sourceIndexVals =
1367 sourceDesc.offsetBasePtr(rewriter, loc, *getTypeConverter(),
1368 sourceUnderlyingDesc, sourceElemPtrType);
1369 Value resultIndexVals =
1370 result.offsetBasePtr(rewriter, loc, *getTypeConverter(),
1371 resultUnderlyingDesc, resultElemPtrType);
1373 int64_t bytesToSkip =
1374 2 * llvm::divideCeil(
1375 getTypeConverter()->getPointerBitwidth(resultAddrSpace), 8);
1376 Value bytesToSkipConst =
1379 LLVM::SubOp::create(rewriter, loc, getIndexType(),
1380 resultUnderlyingSize, bytesToSkipConst);
1381 LLVM::MemcpyOp::create(rewriter, loc, resultIndexVals, sourceIndexVals,
1387 return rewriter.notifyMatchFailure(loc,
"unexpected memref type");
1394static void extractPointersAndOffset(
Location loc,
1395 ConversionPatternRewriter &rewriter,
1397 Value originalOperand,
1398 Value convertedOperand,
1400 Value *offset =
nullptr) {
1402 if (isa<MemRefType>(operandType)) {
1404 *allocatedPtr = desc.allocatedPtr(rewriter, loc);
1405 *alignedPtr = desc.alignedPtr(rewriter, loc);
1406 if (offset !=
nullptr)
1407 *offset = desc.offset(rewriter, loc);
1413 cast<UnrankedMemRefType>(operandType));
1414 auto elementPtrType =
1415 LLVM::LLVMPointerType::get(rewriter.getContext(), memorySpace);
1420 Value underlyingDescPtr = unrankedDesc.memRefDescPtr(rewriter, loc);
1423 rewriter, loc, underlyingDescPtr, elementPtrType);
1425 rewriter, loc, typeConverter, underlyingDescPtr, elementPtrType);
1426 if (offset !=
nullptr) {
1428 rewriter, loc, typeConverter, underlyingDescPtr, elementPtrType);
1432struct MemRefReinterpretCastOpLowering
1434 using ConvertOpToLLVMPattern<
1435 memref::ReinterpretCastOp>::ConvertOpToLLVMPattern;
1438 matchAndRewrite(memref::ReinterpretCastOp castOp, OpAdaptor adaptor,
1439 ConversionPatternRewriter &rewriter)
const override {
1440 Type srcType = castOp.getSource().getType();
1443 if (
failed(convertSourceMemRefToDescriptor(rewriter, srcType, castOp,
1444 adaptor, &descriptor)))
1446 rewriter.replaceOp(castOp, {descriptor});
1451 LogicalResult convertSourceMemRefToDescriptor(
1452 ConversionPatternRewriter &rewriter, Type srcType,
1453 memref::ReinterpretCastOp castOp,
1454 memref::ReinterpretCastOp::Adaptor adaptor, Value *descriptor)
const {
1455 MemRefType targetMemRefType =
1456 cast<MemRefType>(castOp.getResult().getType());
1457 auto llvmTargetDescriptorTy =
1458 typeConverter->convertType<LLVM::LLVMStructType>(targetMemRefType);
1459 if (!llvmTargetDescriptorTy)
1463 Location loc = castOp.getLoc();
1464 auto desc = MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy);
1467 Value allocatedPtr, alignedPtr;
1468 extractPointersAndOffset(loc, rewriter, *getTypeConverter(),
1469 castOp.getSource(), adaptor.getSource(),
1470 &allocatedPtr, &alignedPtr);
1471 desc.setAllocatedPtr(rewriter, loc, allocatedPtr);
1472 desc.setAlignedPtr(rewriter, loc, alignedPtr);
1475 if (castOp.isDynamicOffset(0))
1476 desc.setOffset(rewriter, loc, adaptor.getOffsets()[0]);
1478 desc.setConstantOffset(rewriter, loc, castOp.getStaticOffset(0));
1481 unsigned dynSizeId = 0;
1482 unsigned dynStrideId = 0;
1483 for (
unsigned i = 0, e = targetMemRefType.getRank(); i < e; ++i) {
1484 if (castOp.isDynamicSize(i))
1485 desc.setSize(rewriter, loc, i, adaptor.getSizes()[dynSizeId++]);
1487 desc.setConstantSize(rewriter, loc, i, castOp.getStaticSize(i));
1489 if (castOp.isDynamicStride(i))
1490 desc.setStride(rewriter, loc, i, adaptor.getStrides()[dynStrideId++]);
1492 desc.setConstantStride(rewriter, loc, i, castOp.getStaticStride(i));
1499struct MemRefReshapeOpLowering
1501 using ConvertOpToLLVMPattern<memref::ReshapeOp>::ConvertOpToLLVMPattern;
1504 matchAndRewrite(memref::ReshapeOp reshapeOp, OpAdaptor adaptor,
1505 ConversionPatternRewriter &rewriter)
const override {
1506 Type srcType = reshapeOp.getSource().getType();
1509 if (
failed(convertSourceMemRefToDescriptor(rewriter, srcType, reshapeOp,
1510 adaptor, &descriptor)))
1512 rewriter.replaceOp(reshapeOp, {descriptor});
1518 convertSourceMemRefToDescriptor(ConversionPatternRewriter &rewriter,
1519 Type srcType, memref::ReshapeOp reshapeOp,
1520 memref::ReshapeOp::Adaptor adaptor,
1521 Value *descriptor)
const {
1522 auto shapeMemRefType = cast<MemRefType>(reshapeOp.getShape().getType());
1523 if (shapeMemRefType.hasStaticShape()) {
1524 MemRefType targetMemRefType =
1525 cast<MemRefType>(reshapeOp.getResult().getType());
1526 auto llvmTargetDescriptorTy =
1527 typeConverter->convertType<LLVM::LLVMStructType>(targetMemRefType);
1528 if (!llvmTargetDescriptorTy)
1532 Location loc = reshapeOp.getLoc();
1534 MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy);
1537 Value allocatedPtr, alignedPtr;
1538 extractPointersAndOffset(loc, rewriter, *getTypeConverter(),
1539 reshapeOp.getSource(), adaptor.getSource(),
1540 &allocatedPtr, &alignedPtr);
1541 desc.setAllocatedPtr(rewriter, loc, allocatedPtr);
1542 desc.setAlignedPtr(rewriter, loc, alignedPtr);
1546 SmallVector<int64_t> strides;
1547 if (
failed(targetMemRefType.getStridesAndOffset(strides, offset)))
1548 return rewriter.notifyMatchFailure(
1549 reshapeOp,
"failed to get stride and offset exprs");
1551 if (!isStaticStrideOrOffset(offset))
1552 return rewriter.notifyMatchFailure(reshapeOp,
1553 "dynamic offset is unsupported");
1555 desc.setConstantOffset(rewriter, loc, offset);
1557 assert(targetMemRefType.getLayout().isIdentity() &&
1558 "Identity layout map is a precondition of a valid reshape op");
1560 Type indexType = getIndexType();
1561 Value stride =
nullptr;
1562 int64_t targetRank = targetMemRefType.getRank();
1563 for (
auto i : llvm::reverse(llvm::seq<int64_t>(0, targetRank))) {
1564 if (ShapedType::isStatic(strides[i])) {
1569 }
else if (!stride) {
1579 if (!targetMemRefType.isDynamicDim(i)) {
1581 targetMemRefType.getDimSize(i));
1583 Value shapeOp = reshapeOp.getShape();
1585 dimSize = memref::LoadOp::create(rewriter, loc, shapeOp, index);
1586 Type indexType = getIndexType();
1587 if (dimSize.
getType() != indexType)
1588 dimSize = typeConverter->materializeTargetConversion(
1589 rewriter, loc, indexType, dimSize);
1590 assert(dimSize &&
"Invalid memref element type");
1593 desc.setSize(rewriter, loc, i, dimSize);
1594 desc.setStride(rewriter, loc, i, stride);
1597 stride = LLVM::MulOp::create(rewriter, loc, stride, dimSize);
1605 Location loc = reshapeOp.getLoc();
1606 MemRefDescriptor shapeDesc(adaptor.getShape());
1607 Value resultRank = shapeDesc.size(rewriter, loc, 0);
1610 auto targetType = cast<UnrankedMemRefType>(reshapeOp.getResult().getType());
1611 unsigned addressSpace =
1612 *getTypeConverter()->getMemRefAddressSpace(targetType);
1617 rewriter, loc, typeConverter->convertType(targetType));
1618 targetDesc.setRank(rewriter, loc, resultRank);
1620 rewriter, loc, *getTypeConverter(), targetDesc, addressSpace);
1621 Value underlyingDescPtr = LLVM::AllocaOp::create(
1622 rewriter, loc, getPtrType(), IntegerType::get(
getContext(), 8),
1624 targetDesc.setMemRefDescPtr(rewriter, loc, underlyingDescPtr);
1627 Value allocatedPtr, alignedPtr, offset;
1628 extractPointersAndOffset(loc, rewriter, *getTypeConverter(),
1629 reshapeOp.getSource(), adaptor.getSource(),
1630 &allocatedPtr, &alignedPtr, &offset);
1633 auto elementPtrType =
1634 LLVM::LLVMPointerType::get(rewriter.getContext(), addressSpace);
1637 elementPtrType, allocatedPtr);
1639 underlyingDescPtr, elementPtrType,
1642 underlyingDescPtr, elementPtrType,
1648 rewriter, loc, *getTypeConverter(), underlyingDescPtr, elementPtrType);
1650 rewriter, loc, *getTypeConverter(), targetSizesBase, resultRank);
1651 Value shapeOperandPtr = shapeDesc.alignedPtr(rewriter, loc);
1653 Value resultRankMinusOne =
1654 LLVM::SubOp::create(rewriter, loc, resultRank, oneIndex);
1656 Block *initBlock = rewriter.getInsertionBlock();
1657 Type indexType = getTypeConverter()->getIndexType();
1658 Block::iterator remainingOpsIt = std::next(rewriter.getInsertionPoint());
1660 Block *condBlock = rewriter.createBlock(initBlock->
getParent(), {},
1661 {indexType, indexType}, {loc, loc});
1664 Block *remainingBlock = rewriter.splitBlock(initBlock, remainingOpsIt);
1665 rewriter.mergeBlocks(remainingBlock, condBlock,
ValueRange());
1667 rewriter.setInsertionPointToEnd(initBlock);
1668 LLVM::BrOp::create(rewriter, loc,
1669 ValueRange({resultRankMinusOne, oneIndex}), condBlock);
1670 rewriter.setInsertionPointToStart(condBlock);
1675 Value pred = LLVM::ICmpOp::create(
1676 rewriter, loc, IntegerType::get(rewriter.getContext(), 1),
1677 LLVM::ICmpPredicate::sge, indexArg, zeroIndex);
1680 rewriter.splitBlock(condBlock, rewriter.getInsertionPoint());
1681 rewriter.setInsertionPointToStart(bodyBlock);
1684 auto llvmIndexPtrType = LLVM::LLVMPointerType::get(rewriter.getContext());
1685 Value sizeLoadGep = LLVM::GEPOp::create(
1686 rewriter, loc, llvmIndexPtrType,
1687 typeConverter->convertType(shapeMemRefType.getElementType()),
1688 shapeOperandPtr, indexArg);
1689 Value size = LLVM::LoadOp::create(rewriter, loc, indexType, sizeLoadGep);
1691 targetSizesBase, indexArg, size);
1695 targetStridesBase, indexArg, strideArg);
1696 Value nextStride = LLVM::MulOp::create(rewriter, loc, strideArg, size);
1699 Value decrement = LLVM::SubOp::create(rewriter, loc, indexArg, oneIndex);
1700 LLVM::BrOp::create(rewriter, loc,
ValueRange({decrement, nextStride}),
1704 rewriter.splitBlock(bodyBlock, rewriter.getInsertionPoint());
1707 rewriter.setInsertionPointToEnd(condBlock);
1708 LLVM::CondBrOp::create(rewriter, loc, pred, bodyBlock,
ValueRange(),
1712 rewriter.setInsertionPointToStart(remainder);
1714 *descriptor = targetDesc;
1721template <
typename ReshapeOp>
1722class ReassociatingReshapeOpConversion
1725 using ConvertOpToLLVMPattern<ReshapeOp>::ConvertOpToLLVMPattern;
1726 using ReshapeOpAdaptor =
typename ReshapeOp::Adaptor;
1729 matchAndRewrite(ReshapeOp reshapeOp,
typename ReshapeOp::Adaptor adaptor,
1730 ConversionPatternRewriter &rewriter)
const override {
1731 return rewriter.notifyMatchFailure(
1733 "reassociation operations should have been expanded beforehand");
1740 using ConvertOpToLLVMPattern<memref::SubViewOp>::ConvertOpToLLVMPattern;
1743 matchAndRewrite(memref::SubViewOp subViewOp, OpAdaptor adaptor,
1744 ConversionPatternRewriter &rewriter)
const override {
1745 return rewriter.notifyMatchFailure(
1746 subViewOp,
"subview operations should have been expanded beforehand");
1763 ConversionPatternRewriter &rewriter)
const override {
1764 auto loc = transposeOp.getLoc();
1765 MemRefDescriptor viewMemRef(adaptor.getIn());
1768 if (transposeOp.getPermutation().isIdentity())
1769 return rewriter.replaceOp(transposeOp, {viewMemRef}),
success();
1771 auto targetMemRef = MemRefDescriptor::poison(
1773 typeConverter->convertType(transposeOp.getIn().getType()));
1777 targetMemRef.setAllocatedPtr(rewriter, loc,
1778 viewMemRef.allocatedPtr(rewriter, loc));
1779 targetMemRef.setAlignedPtr(rewriter, loc,
1780 viewMemRef.alignedPtr(rewriter, loc));
1783 targetMemRef.setOffset(rewriter, loc, viewMemRef.offset(rewriter, loc));
1789 for (
const auto &en :
1790 llvm::enumerate(transposeOp.getPermutation().getResults())) {
1791 int targetPos = en.index();
1792 int sourcePos = cast<AffineDimExpr>(en.value()).getPosition();
1793 targetMemRef.setSize(rewriter, loc, targetPos,
1794 viewMemRef.size(rewriter, loc, sourcePos));
1795 targetMemRef.setStride(rewriter, loc, targetPos,
1796 viewMemRef.stride(rewriter, loc, sourcePos));
1799 rewriter.replaceOp(transposeOp, {targetMemRef});
1810 using ConvertOpToLLVMPattern<memref::ViewOp>::ConvertOpToLLVMPattern;
1814 Value getSize(ConversionPatternRewriter &rewriter, Location loc,
1815 ArrayRef<int64_t> shape,
ValueRange dynamicSizes,
unsigned idx,
1816 Type indexType)
const {
1817 assert(idx < shape.size());
1818 if (ShapedType::isStatic(shape[idx]))
1822 llvm::count_if(shape.take_front(idx), ShapedType::isDynamic);
1823 return dynamicSizes[nDynamic];
1830 Value getStride(ConversionPatternRewriter &rewriter, Location loc,
1831 ArrayRef<int64_t> strides, Value nextSize,
1832 Value runningStride,
unsigned idx, Type indexType)
const {
1833 assert(idx < strides.size());
1834 if (ShapedType::isStatic(strides[idx]))
1837 return runningStride
1838 ? LLVM::MulOp::create(rewriter, loc, runningStride, nextSize)
1840 assert(!runningStride);
1845 matchAndRewrite(memref::ViewOp viewOp, OpAdaptor adaptor,
1846 ConversionPatternRewriter &rewriter)
const override {
1847 auto loc = viewOp.getLoc();
1849 auto viewMemRefType = viewOp.getType();
1850 auto targetElementTy =
1851 typeConverter->convertType(viewMemRefType.getElementType());
1852 auto targetDescTy = typeConverter->convertType(viewMemRefType);
1853 if (!targetDescTy || !targetElementTy ||
1854 !LLVM::isCompatibleType(targetElementTy) ||
1855 !LLVM::isCompatibleType(targetDescTy))
1856 return viewOp.emitWarning(
"Target descriptor type not converted to LLVM"),
1860 SmallVector<int64_t, 4> strides;
1861 auto successStrides = viewMemRefType.getStridesAndOffset(strides, offset);
1862 if (
failed(successStrides))
1863 return viewOp.emitWarning(
"cannot cast to non-strided shape"), failure();
1864 assert(offset == 0 &&
"expected offset to be 0");
1868 if (!strides.empty() && (strides.back() != 1 && strides.back() != 0))
1869 return viewOp.emitWarning(
"cannot cast to non-contiguous shape"),
1873 MemRefDescriptor sourceMemRef(adaptor.getSource());
1874 auto targetMemRef = MemRefDescriptor::poison(rewriter, loc, targetDescTy);
1877 Value allocatedPtr = sourceMemRef.allocatedPtr(rewriter, loc);
1878 auto srcMemRefType = cast<MemRefType>(viewOp.getSource().getType());
1879 targetMemRef.setAllocatedPtr(rewriter, loc, allocatedPtr);
1882 Value alignedPtr = sourceMemRef.alignedPtr(rewriter, loc);
1883 alignedPtr = LLVM::GEPOp::create(
1884 rewriter, loc, alignedPtr.
getType(),
1885 typeConverter->convertType(srcMemRefType.getElementType()), alignedPtr,
1886 adaptor.getByteShift());
1888 targetMemRef.setAlignedPtr(rewriter, loc, alignedPtr);
1890 Type indexType = getIndexType();
1894 targetMemRef.setOffset(
1899 if (viewMemRefType.getRank() == 0)
1900 return rewriter.replaceOp(viewOp, {targetMemRef}),
success();
1903 Value stride =
nullptr, nextSize =
nullptr;
1904 for (
int i = viewMemRefType.getRank() - 1; i >= 0; --i) {
1906 Value size = getSize(rewriter, loc, viewMemRefType.getShape(),
1907 adaptor.getSizes(), i, indexType);
1908 targetMemRef.setSize(rewriter, loc, i, size);
1911 getStride(rewriter, loc, strides, nextSize, stride, i, indexType);
1912 targetMemRef.setStride(rewriter, loc, i, stride);
1916 rewriter.replaceOp(viewOp, {targetMemRef});
1927static std::optional<LLVM::AtomicBinOp>
1928matchSimpleAtomicOp(memref::AtomicRMWOp atomicOp) {
1929 switch (atomicOp.getKind()) {
1930 case arith::AtomicRMWKind::addf:
1931 return LLVM::AtomicBinOp::fadd;
1932 case arith::AtomicRMWKind::addi:
1933 return LLVM::AtomicBinOp::add;
1934 case arith::AtomicRMWKind::assign:
1935 return LLVM::AtomicBinOp::xchg;
1936 case arith::AtomicRMWKind::maximumf:
1938 LDBG() <<
"the lowering of memref.atomicrmw maximumf changed "
1939 "from fmax to fmaximum, expect more NaNs";
1940 return LLVM::AtomicBinOp::fmaximum;
1941 case arith::AtomicRMWKind::maxnumf:
1942 return LLVM::AtomicBinOp::fmax;
1943 case arith::AtomicRMWKind::maxs:
1944 return LLVM::AtomicBinOp::max;
1945 case arith::AtomicRMWKind::maxu:
1946 return LLVM::AtomicBinOp::umax;
1947 case arith::AtomicRMWKind::minimumf:
1949 LDBG() <<
"the lowering of memref.atomicrmw minimum changed "
1950 "from fmin to fminimum, expect more NaNs";
1951 return LLVM::AtomicBinOp::fminimum;
1952 case arith::AtomicRMWKind::minnumf:
1953 return LLVM::AtomicBinOp::fmin;
1954 case arith::AtomicRMWKind::mins:
1955 return LLVM::AtomicBinOp::min;
1956 case arith::AtomicRMWKind::minu:
1957 return LLVM::AtomicBinOp::umin;
1958 case arith::AtomicRMWKind::ori:
1959 return LLVM::AtomicBinOp::_or;
1960 case arith::AtomicRMWKind::xori:
1961 return LLVM::AtomicBinOp::_xor;
1962 case arith::AtomicRMWKind::andi:
1963 return LLVM::AtomicBinOp::_and;
1965 return std::nullopt;
1967 llvm_unreachable(
"Invalid AtomicRMWKind");
1970struct AtomicRMWOpLowering :
public LoadStoreOpLowering<memref::AtomicRMWOp> {
1974 matchAndRewrite(memref::AtomicRMWOp atomicOp, OpAdaptor adaptor,
1975 ConversionPatternRewriter &rewriter)
const override {
1976 auto maybeKind = matchSimpleAtomicOp(atomicOp);
1979 auto memRefType = atomicOp.getMemRefType();
1980 SmallVector<int64_t> strides;
1982 if (
failed(memRefType.getStridesAndOffset(strides, offset)))
1986 adaptor.getMemref(), adaptor.getIndices());
1987 rewriter.replaceOpWithNewOp<LLVM::AtomicRMWOp>(
1988 atomicOp, *maybeKind, dataPtr, adaptor.getValue(),
1989 LLVM::AtomicOrdering::acq_rel);
1995class ConvertExtractAlignedPointerAsIndex
1998 using ConvertOpToLLVMPattern<
1999 memref::ExtractAlignedPointerAsIndexOp>::ConvertOpToLLVMPattern;
2002 matchAndRewrite(memref::ExtractAlignedPointerAsIndexOp extractOp,
2004 ConversionPatternRewriter &rewriter)
const override {
2005 BaseMemRefType sourceTy = extractOp.getSource().getType();
2009 MemRefDescriptor desc(adaptor.getSource());
2010 alignedPtr = desc.alignedPtr(rewriter, extractOp->getLoc());
2012 auto elementPtrTy = LLVM::LLVMPointerType::get(
2015 UnrankedMemRefDescriptor desc(adaptor.getSource());
2016 Value descPtr = desc.memRefDescPtr(rewriter, extractOp->getLoc());
2019 rewriter, extractOp->getLoc(), *getTypeConverter(), descPtr,
2023 rewriter.replaceOpWithNewOp<LLVM::PtrToIntOp>(
2024 extractOp, getTypeConverter()->getIndexType(), alignedPtr);
2031class ExtractStridedMetadataOpLowering
2034 using ConvertOpToLLVMPattern<
2035 memref::ExtractStridedMetadataOp>::ConvertOpToLLVMPattern;
2038 matchAndRewrite(memref::ExtractStridedMetadataOp extractStridedMetadataOp,
2040 ConversionPatternRewriter &rewriter)
const override {
2042 if (!LLVM::isCompatibleType(adaptor.getOperands().front().getType()))
2046 MemRefDescriptor sourceMemRef(adaptor.getSource());
2047 Location loc = extractStridedMetadataOp.getLoc();
2048 Value source = extractStridedMetadataOp.getSource();
2050 auto sourceMemRefType = cast<MemRefType>(source.
getType());
2051 int64_t rank = sourceMemRefType.getRank();
2052 SmallVector<Value> results;
2053 results.reserve(2 + rank * 2);
2056 Value baseBuffer = sourceMemRef.allocatedPtr(rewriter, loc);
2057 Value alignedBuffer = sourceMemRef.alignedPtr(rewriter, loc);
2058 MemRefDescriptor dstMemRef = MemRefDescriptor::fromStaticShape(
2059 rewriter, loc, *getTypeConverter(),
2060 cast<MemRefType>(extractStridedMetadataOp.getBaseBuffer().getType()),
2061 baseBuffer, alignedBuffer);
2062 results.push_back((Value)dstMemRef);
2065 results.push_back(sourceMemRef.offset(rewriter, loc));
2068 for (
unsigned i = 0; i < rank; ++i)
2069 results.push_back(sourceMemRef.size(rewriter, loc, i));
2071 for (
unsigned i = 0; i < rank; ++i)
2072 results.push_back(sourceMemRef.stride(rewriter, loc, i));
2074 rewriter.replaceOp(extractStridedMetadataOp, results);
2087 AllocaScopeOpLowering,
2088 AssumeAlignmentOpLowering,
2089 AtomicRMWOpLowering,
2090 ConvertExtractAlignedPointerAsIndex,
2092 DistinctObjectsOpLowering,
2093 ExtractStridedMetadataOpLowering,
2094 GenericAtomicRMWOpLowering,
2095 GetGlobalMemrefOpLowering,
2097 MemRefCastOpLowering,
2098 MemRefReinterpretCastOpLowering,
2099 MemRefReshapeOpLowering,
2100 MemorySpaceCastOpLowering,
2103 ReassociatingReshapeOpConversion<memref::CollapseShapeOp>,
2104 ReassociatingReshapeOpConversion<memref::ExpandShapeOp>,
2108 ViewOpLowering>(converter);
2110 patterns.
add<GlobalMemrefOpLowering, MemRefCopyOpLowering>(converter,
2114 patterns.
add<AlignedAllocOpLowering, DeallocOpLowering>(converter,
2117 patterns.
add<AllocOpLowering, DeallocOpLowering>(converter, symbolTables);
2121struct FinalizeMemRefToLLVMConversionPass
2122 :
public impl::FinalizeMemRefToLLVMConversionPassBase<
2123 FinalizeMemRefToLLVMConversionPass> {
2124 using FinalizeMemRefToLLVMConversionPassBase::
2125 FinalizeMemRefToLLVMConversionPassBase;
2127 void runOnOperation()
override {
2129 const auto &dataLayoutAnalysis = getAnalysis<DataLayoutAnalysis>();
2131 dataLayoutAnalysis.getAtOrAbove(op));
2136 options.useGenericFunctions = useGenericFunctions;
2139 options.overrideIndexBitwidth(indexBitwidth);
2142 &dataLayoutAnalysis);
2148 target.addLegalOp<func::FuncOp>();
2149 if (failed(applyPartialConversion(op,
target, std::move(patterns))))
2150 signalPassFailure();
2155struct MemRefToLLVMDialectInterface :
public ConvertToLLVMPatternInterface {
2156 MemRefToLLVMDialectInterface(Dialect *dialect)
2157 : ConvertToLLVMPatternInterface(dialect) {}
2159 void loadDependentDialects(MLIRContext *context)
const final {
2160 context->loadDialect<LLVM::LLVMDialect>();
2165 void populateConvertToLLVMConversionPatterns(
2166 ConversionTarget &
target, LLVMTypeConverter &typeConverter,
2167 RewritePatternSet &patterns)
const final {
2176 dialect->addInterfaces<MemRefToLLVMDialectInterface>();
static LLVM::GEPNoWrapFlags getLoadStoreNoWrapFlags(MemRefType type)
Returns GEP no-wrap flags for a memref load/store.
static llvm::Value * getSizeInBytes(DataLayout &dl, const mlir::Type &type, Operation *clauseOp, llvm::Value *basePointer, llvm::Type *baseType, llvm::IRBuilderBase &builder, LLVM::ModuleTranslation &moduleTranslation)
static llvm::ManagedStatic< PassManagerOptions > options
Rewrite AVX2-specific vector.transpose, for the supported cases and depending on the TransposeLowerin...
LogicalResult matchAndRewrite(vector::TransposeOp op, PatternRewriter &rewriter) const override
unsigned getMemorySpaceAsInt() const
[deprecated] Returns the memory space in old raw integer representation.
bool hasRank() const
Returns if this type is ranked, i.e. it has a known number of dimensions.
OpListType::iterator iterator
BlockArgument getArgument(unsigned i)
Region * getParent() const
Provide a 'getParent' method for ilist_node_with_parent methods.
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgListType getArguments()
iterator_range< iterator > without_terminator()
Return an iterator range over the operation within this block excluding the terminator operation at t...
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
Stores data layout objects for each operation that specifies the data layout above and below the give...
The main mechanism for performing data layout queries.
llvm::TypeSize getTypeSize(Type t) const
Returns the size of the given type in the current scope.
The DialectRegistry maps a dialect namespace to a constructor for the matching dialect.
bool addExtension(TypeID extensionID, std::unique_ptr< DialectExtensionBase > extension)
Add the given extension to the registry.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
auto lookupOrNull(T from) const
Lookup a mapped value within the map.
Derived class that automatically populates legalization information for different LLVM ops.
Conversion from types to the LLVM IR dialect.
unsigned getUnrankedMemRefDescriptorSize(UnrankedMemRefType type, const DataLayout &layout) const
Returns the size of the unranked memref descriptor object in bytes.
Value promoteOneMemRefDescriptor(Location loc, Value operand, OpBuilder &builder) const
Promote the LLVM struct representation of one MemRef descriptor to stack and use pointer to struct to...
const LowerToLLVMOptions & getOptions() const
FailureOr< unsigned > getMemRefAddressSpace(BaseMemRefType type) const
Return the LLVM address space corresponding to the memory space of the memref type type or failure if...
const DataLayoutAnalysis * getDataLayoutAnalysis() const
Returns the data layout analysis to query during conversion.
unsigned getMemRefDescriptorSize(MemRefType type, const DataLayout &layout) const
Returns the size of the memref descriptor object in bytes.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Options to control the LLVM lowering.
AllocLowering allocLowering
@ Malloc
Use malloc for heap allocations.
@ AlignedAlloc
Use aligned_alloc for heap allocations.
MLIRContext is the top-level object for a collection of MLIR operations.
Helper class to produce LLVM dialect operations extracting or inserting elements of a MemRef descript...
This class helps build Operations.
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
result_range getResults()
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class represents a collection of SymbolTables.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
static void setOffset(OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter, Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType, Value offset)
Builds IR inserting the offset into the descriptor.
static Value allocatedPtr(OpBuilder &builder, Location loc, Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType)
TODO: The following accessors don't take alignment rules between elements of the descriptor struct in...
static Value computeSize(OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter, UnrankedMemRefDescriptor desc, unsigned addressSpace)
Builds and returns IR computing the size in bytes (suitable for opaque allocation).
void setRank(OpBuilder &builder, Location loc, Value value)
Builds IR setting the rank in the descriptor.
Value memRefDescPtr(OpBuilder &builder, Location loc) const
Builds IR extracting ranked memref descriptor ptr.
static void setAllocatedPtr(OpBuilder &builder, Location loc, Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType, Value allocatedPtr)
Builds IR inserting the allocated pointer into the descriptor.
static void setSize(OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter, Value sizeBasePtr, Value index, Value size)
Builds IR inserting the size[index] into the descriptor.
static Value pack(OpBuilder &builder, Location loc, const LLVMTypeConverter &converter, UnrankedMemRefType type, ValueRange values)
Builds IR populating an unranked MemRef descriptor structure from a list of individual constituent va...
static UnrankedMemRefDescriptor poison(OpBuilder &builder, Location loc, Type descriptorType)
Builds IR creating an undef value of the descriptor type.
static void setAlignedPtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter, Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType, Value alignedPtr)
Builds IR inserting the aligned pointer into the descriptor.
static Value offset(OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter, Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType)
Builds IR extracting the offset from the descriptor.
static Value strideBasePtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter, Value sizeBasePtr, Value rank)
Builds IR extracting the pointer to the first element of the stride array.
void setMemRefDescPtr(OpBuilder &builder, Location loc, Value value)
Builds IR setting ranked memref descriptor ptr.
static void setStride(OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter, Value strideBasePtr, Value index, Value stride)
Builds IR inserting the stride[index] into the descriptor.
static Value sizeBasePtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter, Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType)
Builds IR extracting the pointer to the first element of the size array.
static Value alignedPtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &typeConverter, Value memRefDescPtr, LLVM::LLVMPointerType elemPtrType)
Builds IR extracting the aligned pointer from the descriptor.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
FailureOr< LLVM::LLVMFuncOp > lookupOrCreateFreeFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
Value getStridedElementPtr(OpBuilder &builder, Location loc, const LLVMTypeConverter &converter, MemRefType type, Value memRefDesc, ValueRange indices, LLVM::GEPNoWrapFlags noWrapFlags=LLVM::GEPNoWrapFlags::none)
Performs the index computation to get to the element at indices of the memory pointed to by memRefDes...
Value createIndexAttrConstant(OpBuilder &builder, Location loc, Type resultType, int64_t value)
Creates an llvm.mlir.constant producing value as resultType, which is expected to be the converted in...
FailureOr< LLVM::LLVMFuncOp > lookupOrCreateGenericAlignedAllocFn(OpBuilder &b, Operation *moduleOp, Type indexType, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreateMallocFn(OpBuilder &b, Operation *moduleOp, Type indexType, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreateGenericAllocFn(OpBuilder &b, Operation *moduleOp, Type indexType, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreateAlignedAllocFn(OpBuilder &b, Operation *moduleOp, Type indexType, SymbolTableCollection *symbolTables=nullptr)
FailureOr< LLVM::LLVMFuncOp > lookupOrCreateGenericFreeFn(OpBuilder &b, Operation *moduleOp, SymbolTableCollection *symbolTables=nullptr)
bool isStaticShapeAndContiguousRowMajor(MemRefType type)
Returns true, if the memref type has static shapes and represents a contiguous chunk of memory.
void promote(RewriterBase &rewriter, scf::ForallOp forallOp)
Promotes the loop body of a scf::ForallOp to its containing block.
Include the generated interface declarations.
void registerConvertMemRefToLLVMInterface(DialectRegistry ®istry)
static constexpr unsigned kDeriveIndexBitwidthFromDataLayout
Value to pass as bitwidth for the index type when the converter is expected to derive the bitwidth fr...
void populateFinalizeMemRefToLLVMConversionPatterns(const LLVMTypeConverter &converter, RewritePatternSet &patterns, SymbolTableCollection *symbolTables=nullptr)
Collect a set of patterns to convert memory-related operations from the MemRef dialect to the LLVM di...
Operation * clone(OpBuilder &b, Operation *op, TypeRange newResultTypes, ValueRange newOperands)