MLIR 24.0.0git
MemRefToLLVM.cpp
Go to the documentation of this file.
1//===- MemRefToLLVM.cpp - MemRef to LLVM dialect conversion ---------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
10
23#include "mlir/IR/AffineMap.h"
25#include "mlir/IR/IRMapping.h"
26#include "mlir/Pass/Pass.h"
27#include "llvm/Support/DebugLog.h"
28#include "llvm/Support/MathExtras.h"
29
30#include <optional>
31
32#define DEBUG_TYPE "memref-to-llvm"
33
34namespace mlir {
35#define GEN_PASS_DEF_FINALIZEMEMREFTOLLVMCONVERSIONPASS
36#include "mlir/Conversion/Passes.h.inc"
37} // namespace mlir
38
39using namespace mlir;
40
41/// Returns GEP no-wrap flags for a memref load/store.
42/// inbounds is always valid when indices are in-bounds per the memref spec.
43/// nuw requires every index*stride term to not unsigned-wrap, which holds iff
44/// all strides are statically non-negative. Negative strides would make the
45/// intermediate mul nuw overflow (e.g., idx * (-1 as u64) wraps for idx > 0).
46static LLVM::GEPNoWrapFlags getLoadStoreNoWrapFlags(MemRefType type) {
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;
51 }))
52 flags = flags | LLVM::GEPNoWrapFlags::nuw;
53 return flags;
54}
55
56namespace {
57
58static bool isStaticStrideOrOffset(int64_t strideOrOffset) {
59 return ShapedType::isStatic(strideOrOffset);
60}
61
62static FailureOr<LLVM::LLVMFuncOp>
63getFreeFn(OpBuilder &b, const LLVMTypeConverter *typeConverter,
64 Operation *module, SymbolTableCollection *symbolTables) {
65 bool useGenericFn = typeConverter->getOptions().useGenericFunctions;
66
67 if (useGenericFn)
68 return LLVM::lookupOrCreateGenericFreeFn(b, module, symbolTables);
69
70 return LLVM::lookupOrCreateFreeFn(b, module, symbolTables);
71}
72
73static FailureOr<LLVM::LLVMFuncOp>
74getNotalignedAllocFn(OpBuilder &b, const LLVMTypeConverter *typeConverter,
75 Operation *module, Type indexType,
76 SymbolTableCollection *symbolTables) {
77 bool useGenericFn = typeConverter->getOptions().useGenericFunctions;
78 if (useGenericFn)
79 return LLVM::lookupOrCreateGenericAllocFn(b, module, indexType,
80 symbolTables);
81
82 return LLVM::lookupOrCreateMallocFn(b, module, indexType, symbolTables);
83}
84
85static FailureOr<LLVM::LLVMFuncOp>
86getAlignedAllocFn(OpBuilder &b, const LLVMTypeConverter *typeConverter,
87 Operation *module, Type indexType,
88 SymbolTableCollection *symbolTables) {
89 bool useGenericFn = typeConverter->getOptions().useGenericFunctions;
90
91 if (useGenericFn)
92 return LLVM::lookupOrCreateGenericAlignedAllocFn(b, module, indexType,
93 symbolTables);
94
95 return LLVM::lookupOrCreateAlignedAllocFn(b, module, indexType, symbolTables);
96}
97
98/// Computes the aligned value for 'input' as follows:
99/// bumped = input + alignement - 1
100/// aligned = bumped - bumped % alignment
101static Value createAligned(ConversionPatternRewriter &rewriter, Location loc,
102 Value input, Value alignment) {
103 Value one =
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);
110}
111
112/// Computes the byte size for the MemRef element type.
113static unsigned getMemRefEltSizeInBytes(const LLVMTypeConverter *typeConverter,
114 MemRefType memRefType, Operation *op,
115 const DataLayout *defaultLayout) {
116 const DataLayout *layout = defaultLayout;
117 if (const DataLayoutAnalysis *analysis =
118 typeConverter->getDataLayoutAnalysis()) {
119 layout = &analysis->getAbove(op);
120 }
121 Type elementType = memRefType.getElementType();
122 if (auto memRefElementType = dyn_cast<MemRefType>(elementType))
123 return typeConverter->getMemRefDescriptorSize(memRefElementType, *layout);
124 if (auto memRefElementType = dyn_cast<UnrankedMemRefType>(elementType))
125 return typeConverter->getUnrankedMemRefDescriptorSize(memRefElementType,
126 *layout);
127 return layout->getTypeSize(elementType);
128}
129
130static Value castAllocFuncResult(ConversionPatternRewriter &rewriter,
131 Location loc, Value allocatedPtr,
132 MemRefType memRefType, Type elementPtrType,
133 const LLVMTypeConverter &typeConverter) {
134 auto allocatedPtrTy = cast<LLVM::LLVMPointerType>(allocatedPtr.getType());
135 FailureOr<unsigned> maybeMemrefAddrSpace =
136 typeConverter.getMemRefAddressSpace(memRefType);
137 assert(succeeded(maybeMemrefAddrSpace) && "unsupported address space");
138 unsigned memrefAddrSpace = *maybeMemrefAddrSpace;
139 if (allocatedPtrTy.getAddressSpace() != memrefAddrSpace)
140 allocatedPtr = LLVM::AddrSpaceCastOp::create(
141 rewriter, loc,
142 LLVM::LLVMPointerType::get(rewriter.getContext(), memrefAddrSpace),
143 allocatedPtr);
144 return allocatedPtr;
145}
146
147class AllocOpLowering : public ConvertOpToLLVMPattern<memref::AllocOp> {
148 SymbolTableCollection *symbolTables = nullptr;
149
150public:
151 explicit AllocOpLowering(const LLVMTypeConverter &typeConverter,
152 SymbolTableCollection *symbolTables = nullptr,
153 PatternBenefit benefit = 1)
154 : ConvertOpToLLVMPattern<memref::AllocOp>(typeConverter, benefit),
155 symbolTables(symbolTables) {}
156
157 LogicalResult
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");
164
165 // Get or insert alloc function into the module.
166 FailureOr<LLVM::LLVMFuncOp> allocFuncOp =
167 getNotalignedAllocFn(rewriter, getTypeConverter(),
168 op->getParentWithTrait<OpTrait::SymbolTable>(),
169 getIndexType(), symbolTables);
170 if (failed(allocFuncOp))
171 return failure();
172
173 // Get actual sizes of the memref as values: static sizes are constant
174 // values and dynamic sizes are passed to 'alloc' as operands. In case of
175 // zero-dimensional memref, assume a scalar (size 1).
176 SmallVector<Value, 4> sizes;
177 SmallVector<Value, 4> strides;
178 Value sizeBytes;
179
180 this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
181 rewriter, sizes, strides, sizeBytes, true);
182
183 Value alignment = getAlignment(rewriter, loc, op);
184 if (alignment) {
185 // Adjust the allocation size to consider alignment.
186 sizeBytes = LLVM::AddOp::create(rewriter, loc, sizeBytes, alignment);
187 }
188
189 // Allocate the underlying buffer.
190 Type elementPtrType = this->getElementPtrType(memRefType);
191 assert(elementPtrType && "could not compute element ptr type");
192 auto results =
193 LLVM::CallOp::create(rewriter, loc, allocFuncOp.value(), sizeBytes);
194
195 Value allocatedPtr =
196 castAllocFuncResult(rewriter, loc, results.getResult(), memRefType,
197 elementPtrType, *getTypeConverter());
198 Value alignedPtr = allocatedPtr;
199 if (alignment) {
200 // Compute the aligned pointer.
201 Value allocatedInt =
202 LLVM::PtrToIntOp::create(rewriter, loc, getIndexType(), allocatedPtr);
203 Value alignmentInt =
204 createAligned(rewriter, loc, allocatedInt, alignment);
205 alignedPtr =
206 LLVM::IntToPtrOp::create(rewriter, loc, elementPtrType, alignmentInt);
207 }
208
209 // Create the MemRef descriptor.
210 auto memRefDescriptor = this->createMemRefDescriptor(
211 loc, memRefType, allocatedPtr, alignedPtr, sizes, strides, rewriter);
212
213 // Return the final value of the descriptor.
214 rewriter.replaceOp(op, {memRefDescriptor});
215 return success();
216 }
217
218 /// Computes the alignment for the given memory allocation op.
219 template <typename OpType>
220 Value getAlignment(ConversionPatternRewriter &rewriter, Location loc,
221 OpType op) const {
222 MemRefType memRefType = op.getType();
223 Value alignment;
224 if (auto alignmentAttr = op.getAlignment()) {
225 Type indexType = getIndexType();
226 alignment =
227 createIndexAttrConstant(rewriter, loc, indexType, *alignmentAttr);
228 } else if (!memRefType.getElementType().isSignlessIntOrIndexOrFloat()) {
229 // In the case where no alignment is specified, we may want to override
230 // `malloc's` behavior. `malloc` typically aligns at the size of the
231 // biggest scalar on a target HW. For non-scalars, use the natural
232 // alignment of the LLVM type given by the LLVM DataLayout.
233 alignment = getSizeInBytes(loc, memRefType.getElementType(), rewriter);
234 }
235 return alignment;
236 }
237};
238
239class AlignedAllocOpLowering : public ConvertOpToLLVMPattern<memref::AllocOp> {
240 SymbolTableCollection *symbolTables = nullptr;
241
242public:
243 explicit AlignedAllocOpLowering(const LLVMTypeConverter &typeConverter,
244 SymbolTableCollection *symbolTables = nullptr,
245 PatternBenefit benefit = 1)
246 : ConvertOpToLLVMPattern<memref::AllocOp>(typeConverter, benefit),
247 symbolTables(symbolTables) {}
248
249 LogicalResult
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");
256
257 // Get or insert alloc function into module.
258 FailureOr<LLVM::LLVMFuncOp> allocFuncOp =
259 getAlignedAllocFn(rewriter, getTypeConverter(),
260 op->getParentWithTrait<OpTrait::SymbolTable>(),
261 getIndexType(), symbolTables);
262 if (failed(allocFuncOp))
263 return failure();
264
265 // Get actual sizes of the memref as values: static sizes are constant
266 // values and dynamic sizes are passed to 'alloc' as operands. In case of
267 // zero-dimensional memref, assume a scalar (size 1).
268 SmallVector<Value, 4> sizes;
269 SmallVector<Value, 4> strides;
270 Value sizeBytes;
271
272 this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
273 rewriter, sizes, strides, sizeBytes, !false);
274
275 int64_t alignment = alignedAllocationGetAlignment(op, &defaultLayout);
276
277 Value allocAlignment =
278 createIndexAttrConstant(rewriter, loc, getIndexType(), alignment);
279
280 // Function aligned_alloc requires size to be a multiple of alignment; we
281 // pad the size to the next multiple if necessary.
282 if (!isMemRefSizeMultipleOf(memRefType, alignment, op, &defaultLayout))
283 sizeBytes = createAligned(rewriter, loc, sizeBytes, allocAlignment);
284
285 Type elementPtrType = this->getElementPtrType(memRefType);
286 auto results =
287 LLVM::CallOp::create(rewriter, loc, allocFuncOp.value(),
288 ValueRange({allocAlignment, sizeBytes}));
289
290 Value ptr =
291 castAllocFuncResult(rewriter, loc, results.getResult(), memRefType,
292 elementPtrType, *getTypeConverter());
293
294 // Create the MemRef descriptor.
295 auto memRefDescriptor = this->createMemRefDescriptor(
296 loc, memRefType, ptr, ptr, sizes, strides, rewriter);
297
298 // Return the final value of the descriptor.
299 rewriter.replaceOp(op, {memRefDescriptor});
300 return success();
301 }
302
303 /// The minimum alignment to use with aligned_alloc (has to be a power of 2).
304 static constexpr uint64_t kMinAlignedAllocAlignment = 16UL;
305
306 /// Computes the alignment for aligned_alloc used to allocate the buffer for
307 /// the memory allocation op.
308 ///
309 /// Aligned_alloc requires the allocation size to be a power of two, and the
310 /// allocation size to be a multiple of the alignment.
311 int64_t alignedAllocationGetAlignment(memref::AllocOp op,
312 const DataLayout *defaultLayout) const {
313 if (std::optional<uint64_t> alignment = op.getAlignment())
314 return *alignment;
315
316 // Whenever we don't have alignment set, we will use an alignment
317 // consistent with the element type; since the allocation size has to be a
318 // power of two, we will bump to the next power of two if it isn't.
319 unsigned eltSizeBytes = getMemRefEltSizeInBytes(
320 getTypeConverter(), op.getType(), op, defaultLayout);
321 return std::max(kMinAlignedAllocAlignment,
322 llvm::PowerOf2Ceil(eltSizeBytes));
323 }
324
325 /// Returns true if the memref size in bytes is known to be a multiple of
326 /// factor.
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))
333 continue;
334 sizeDivisor = sizeDivisor * type.getDimSize(i);
335 }
336 return sizeDivisor % factor == 0;
337 }
338
339private:
340 /// Default layout to use in absence of the corresponding analysis.
341 DataLayout defaultLayout;
342};
343
344struct AllocaOpLowering : public ConvertOpToLLVMPattern<memref::AllocaOp> {
345 using ConvertOpToLLVMPattern<memref::AllocaOp>::ConvertOpToLLVMPattern;
346
347 /// Allocates the underlying buffer using the right call. `allocatedBytePtr`
348 /// is set to null for stack allocations. `accessAlignment` is set if
349 /// alignment is needed post allocation (for eg. in conjunction with malloc).
350 LogicalResult
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");
357
358 // Get actual sizes of the memref as values: static sizes are constant
359 // values and dynamic sizes are passed to 'alloc' as operands. In case of
360 // zero-dimensional memref, assume a scalar (size 1).
361 SmallVector<Value, 4> sizes;
362 SmallVector<Value, 4> strides;
363 Value size;
364
365 this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
366 rewriter, sizes, strides, size, !true);
367
368 // With alloca, one gets a pointer to the element type right away.
369 // For stack allocations.
370 auto elementType =
371 typeConverter->convertType(op.getType().getElementType());
372 FailureOr<unsigned> maybeAddressSpace =
373 getTypeConverter()->getMemRefAddressSpace(op.getType());
374 assert(succeeded(maybeAddressSpace) && "unsupported address space");
375 unsigned addrSpace = *maybeAddressSpace;
376 auto elementPtrType =
377 LLVM::LLVMPointerType::get(rewriter.getContext(), addrSpace);
378
379 auto allocatedElementPtr =
380 LLVM::AllocaOp::create(rewriter, loc, elementPtrType, elementType, size,
381 op.getAlignment().value_or(0));
382
383 // Create the MemRef descriptor.
384 auto memRefDescriptor = this->createMemRefDescriptor(
385 loc, memRefType, allocatedElementPtr, allocatedElementPtr, sizes,
386 strides, rewriter);
387
388 // Return the final value of the descriptor.
389 rewriter.replaceOp(op, {memRefDescriptor});
390 return success();
391 }
392};
393
394struct AllocaScopeOpLowering
395 : public ConvertOpToLLVMPattern<memref::AllocaScopeOp> {
396 using ConvertOpToLLVMPattern<memref::AllocaScopeOp>::ConvertOpToLLVMPattern;
397
398 LogicalResult
399 matchAndRewrite(memref::AllocaScopeOp allocaScopeOp, OpAdaptor adaptor,
400 ConversionPatternRewriter &rewriter) const override {
401 OpBuilder::InsertionGuard guard(rewriter);
402 Location loc = allocaScopeOp.getLoc();
403
404 // Split the current block before the AllocaScopeOp to create the inlining
405 // point.
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;
412 } else {
413 continueBlock = rewriter.createBlock(
414 remainingOpsBlock, allocaScopeOp.getResultTypes(),
415 SmallVector<Location>(allocaScopeOp->getNumResults(),
416 allocaScopeOp.getLoc()));
417 LLVM::BrOp::create(rewriter, loc, ValueRange(), remainingOpsBlock);
418 }
419
420 // Inline body region.
421 Block *beforeBody = &allocaScopeOp.getBodyRegion().front();
422 Block *afterBody = &allocaScopeOp.getBodyRegion().back();
423 rewriter.inlineRegionBefore(allocaScopeOp.getBodyRegion(), continueBlock);
424
425 // Save stack and then branch into the body of the region.
426 rewriter.setInsertionPointToEnd(currentBlock);
427 auto stackSaveOp = LLVM::StackSaveOp::create(rewriter, loc, getPtrType());
428 LLVM::BrOp::create(rewriter, loc, ValueRange(), beforeBody);
429
430 // Replace the alloca_scope return with a branch that jumps out of the body.
431 // Stack restore before leaving the body region.
432 rewriter.setInsertionPointToEnd(afterBody);
433 auto returnOp =
434 cast<memref::AllocaScopeReturnOp>(afterBody->getTerminator());
435 auto branchOp = rewriter.replaceOpWithNewOp<LLVM::BrOp>(
436 returnOp, returnOp.getResults(), continueBlock);
437
438 // Insert stack restore before jumping out the body of the region.
439 rewriter.setInsertionPoint(branchOp);
440 LLVM::StackRestoreOp::create(rewriter, loc, stackSaveOp);
441
442 // Replace the op with values return from the body region.
443 rewriter.replaceOp(allocaScopeOp, continueBlock->getArguments());
444
445 return success();
446 }
447};
448
449struct AssumeAlignmentOpLowering
450 : public ConvertOpToLLVMPattern<memref::AssumeAlignmentOp> {
451 using ConvertOpToLLVMPattern<
452 memref::AssumeAlignmentOp>::ConvertOpToLLVMPattern;
453 explicit AssumeAlignmentOpLowering(const LLVMTypeConverter &converter)
454 : ConvertOpToLLVMPattern<memref::AssumeAlignmentOp>(converter) {}
455
456 LogicalResult
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();
462
463 auto srcMemRefType = cast<MemRefType>(op.getMemref().getType());
464 Value ptr = getStridedElementPtr(rewriter, loc, srcMemRefType, memref,
465 /*indices=*/{});
466
467 // Emit llvm.assume(true) ["align"(memref, alignment)].
468 // This is more direct than ptrtoint-based checks, is explicitly supported,
469 // and works with non-integral address spaces.
470 Value trueCond =
471 LLVM::ConstantOp::create(rewriter, loc, rewriter.getBoolAttr(true));
472 Value alignmentConst =
473 createIndexAttrConstant(rewriter, loc, getIndexType(), alignment);
474 LLVM::AssumeOp::create(rewriter, loc, trueCond, LLVM::AssumeAlignTag(), ptr,
475 alignmentConst);
476 rewriter.replaceOp(op, memref);
477 return success();
478 }
479};
480
481struct DistinctObjectsOpLowering
482 : public ConvertOpToLLVMPattern<memref::DistinctObjectsOp> {
483 using ConvertOpToLLVMPattern<
484 memref::DistinctObjectsOp>::ConvertOpToLLVMPattern;
485 explicit DistinctObjectsOpLowering(const LLVMTypeConverter &converter)
486 : ConvertOpToLLVMPattern<memref::DistinctObjectsOp>(converter) {}
487
488 LogicalResult
489 matchAndRewrite(memref::DistinctObjectsOp op, OpAdaptor adaptor,
490 ConversionPatternRewriter &rewriter) const override {
491 ValueRange operands = adaptor.getOperands();
492 if (operands.size() <= 1) {
493 // Fast path.
494 rewriter.replaceOp(op, operands);
495 return success();
496 }
497
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(),
505 memrefType);
506 ptrs.push_back(ptr);
507 }
508
509 auto cond =
510 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI1Type(), 1);
511 // Generate separate_storage assumptions for each pair of pointers.
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);
518 }
519 }
520
521 rewriter.replaceOp(op, operands);
522 return success();
523 }
524};
525
526// A `dealloc` is converted into a call to `free` on the underlying data buffer.
527// The memref descriptor being an SSA value, there is no need to clean it up
528// in any way.
529class DeallocOpLowering : public ConvertOpToLLVMPattern<memref::DeallocOp> {
530 SymbolTableCollection *symbolTables = nullptr;
531
532public:
533 explicit DeallocOpLowering(const LLVMTypeConverter &typeConverter,
534 SymbolTableCollection *symbolTables = nullptr,
535 PatternBenefit benefit = 1)
536 : ConvertOpToLLVMPattern<memref::DeallocOp>(typeConverter, benefit),
537 symbolTables(symbolTables) {}
538
539 LogicalResult
540 matchAndRewrite(memref::DeallocOp op, OpAdaptor adaptor,
541 ConversionPatternRewriter &rewriter) const override {
542 // Insert the `free` declaration if it is not already present.
543 FailureOr<LLVM::LLVMFuncOp> freeFunc =
544 getFreeFn(rewriter, getTypeConverter(),
545 op->getParentWithTrait<OpTrait::SymbolTable>(), symbolTables);
546 if (failed(freeFunc))
547 return failure();
548 Value allocatedPtr;
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()),
557 elementPtrTy);
558 } else {
559 allocatedPtr = MemRefDescriptor(adaptor.getMemref())
560 .allocatedPtr(rewriter, op.getLoc());
561 }
562 rewriter.replaceOpWithNewOp<LLVM::CallOp>(op, freeFunc.value(),
563 allocatedPtr);
564 return success();
565 }
566};
567
568// A `dim` is converted to a constant for static sizes and to an access to the
569// size stored in the memref descriptor for dynamic sizes.
570struct DimOpLowering : public ConvertOpToLLVMPattern<memref::DimOp> {
571 using ConvertOpToLLVMPattern<memref::DimOp>::ConvertOpToLLVMPattern;
572
573 LogicalResult
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))
581 return failure();
582 rewriter.replaceOp(dimOp, {*extractedSize});
583 return success();
584 }
585 if (isa<MemRefType>(operandType)) {
586 rewriter.replaceOp(
587 dimOp, {extractSizeOfRankedMemRef(operandType, dimOp,
588 adaptor.getOperands(), rewriter)});
589 return success();
590 }
591 llvm_unreachable("expected MemRefType or UnrankedMemRefType");
592 }
593
594private:
595 FailureOr<Value>
596 extractSizeOfUnrankedMemRef(Type operandType, memref::DimOp dimOp,
597 OpAdaptor adaptor,
598 ConversionPatternRewriter &rewriter) const {
599 Location loc = dimOp.getLoc();
600
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 "
608 "address space");
609 return failure();
610 }
611 unsigned addressSpace = *maybeAddressSpace;
612
613 // Extract pointer to the underlying ranked descriptor and bitcast it to a
614 // memref<element_type> descriptor pointer to minimize the number of GEP
615 // operations.
616 UnrankedMemRefDescriptor unrankedDesc(adaptor.getSource());
617 Value underlyingRankedDesc = unrankedDesc.memRefDescPtr(rewriter, loc);
618
619 Type elementType = typeConverter->convertType(scalarMemRefType);
620
621 // Get pointer to offset field of memref<element_type> descriptor.
622 auto indexPtrTy =
623 LLVM::LLVMPointerType::get(rewriter.getContext(), addressSpace);
624 Value offsetPtr =
625 LLVM::GEPOp::create(rewriter, loc, indexPtrTy, elementType,
626 underlyingRankedDesc, ArrayRef<LLVM::GEPArg>{0, 2});
627
628 // The size value that we have to extract can be obtained using GEPop with
629 // `dimOp.index() + 1` index argument.
630 Value idxPlusOne = LLVM::AddOp::create(
631 rewriter, loc,
632 createIndexAttrConstant(rewriter, loc, getIndexType(), 1),
633 adaptor.getIndex());
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)
639 .getResult();
640 }
641
642 std::optional<int64_t> getConstantDimIndex(memref::DimOp dimOp) const {
643 if (auto idx = dimOp.getConstantIndex())
644 return idx;
645
646 if (auto constantOp = dimOp.getIndex().getDefiningOp<LLVM::ConstantOp>())
647 return cast<IntegerAttr>(constantOp.getValue()).getValue().getSExtValue();
648
649 return std::nullopt;
650 }
651
652 Value extractSizeOfRankedMemRef(Type operandType, memref::DimOp dimOp,
653 OpAdaptor adaptor,
654 ConversionPatternRewriter &rewriter) const {
655 Location loc = dimOp.getLoc();
656
657 // Take advantage if index is constant.
658 MemRefType memRefType = cast<MemRefType>(operandType);
659 Type indexType = getIndexType();
660 if (std::optional<int64_t> index = getConstantDimIndex(dimOp)) {
661 int64_t i = *index;
662 if (i >= 0 && i < memRefType.getRank()) {
663 if (memRefType.isDynamicDim(i)) {
664 // extract dynamic size from the memref descriptor.
665 MemRefDescriptor descriptor(adaptor.getSource());
666 return descriptor.size(rewriter, loc, i);
667 }
668 // Use constant for static size.
669 int64_t dimSize = memRefType.getDimSize(i);
670 return createIndexAttrConstant(rewriter, loc, indexType, dimSize);
671 }
672 }
673 Value index = adaptor.getIndex();
674 int64_t rank = memRefType.getRank();
675 MemRefDescriptor memrefDescriptor(adaptor.getSource());
676 return memrefDescriptor.size(rewriter, loc, index, rank);
677 }
678};
679
680/// Common base for load and store operations on MemRefs. Restricts the match
681/// to supported MemRef types. Provides functionality to emit code accessing a
682/// specific element of the underlying data buffer.
683template <typename Derived>
684struct LoadStoreOpLowering : public ConvertOpToLLVMPattern<Derived> {
685 using ConvertOpToLLVMPattern<Derived>::ConvertOpToLLVMPattern;
686 using ConvertOpToLLVMPattern<Derived>::isConvertibleAndHasIdentityMaps;
687 using Base = LoadStoreOpLowering<Derived>;
688};
689
690/// Wrap a llvm.cmpxchg operation in a while loop so that the operation can be
691/// retried until it succeeds in atomically storing a new value into memory.
692///
693/// +---------------------------------+
694/// | <code before the AtomicRMWOp> |
695/// | <compute initial %loaded> |
696/// | cf.br loop(%loaded) |
697/// +---------------------------------+
698/// |
699/// -------| |
700/// | v v
701/// | +--------------------------------+
702/// | | loop(%loaded): |
703/// | | <body contents> |
704/// | | %pair = cmpxchg |
705/// | | %ok = %pair[0] |
706/// | | %new = %pair[1] |
707/// | | cf.cond_br %ok, end, loop(%new) |
708/// | +--------------------------------+
709/// | | |
710/// |----------- |
711/// v
712/// +--------------------------------+
713/// | end: |
714/// | <code after the AtomicRMWOp> |
715/// +--------------------------------+
716///
717struct GenericAtomicRMWOpLowering
718 : public LoadStoreOpLowering<memref::GenericAtomicRMWOp> {
719 using Base::Base;
720
721 LogicalResult
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());
726
727 // `llvm.cmpxchg` only supports integer or pointer operands. For
728 // floating-point element types, perform the CAS on a same-width integer
729 // and bitcast at the boundaries.
730 bool needsBitcast = isa<FloatType>(valueType);
731 Type cmpxchgType = valueType;
732 if (needsBitcast) {
733 unsigned bitWidth = cast<FloatType>(valueType).getWidth();
734 cmpxchgType = rewriter.getIntegerType(bitWidth);
735 }
736
737 // Split the block into initial, loop, and ending parts.
738 auto *initBlock = rewriter.getInsertionBlock();
739 auto *loopBlock = rewriter.splitBlock(initBlock, Block::iterator(atomicOp));
740 loopBlock->addArgument(cmpxchgType, loc);
741
742 auto *endBlock =
743 rewriter.splitBlock(loopBlock, Block::iterator(atomicOp)++);
744
745 // Compute the loaded value and branch to the loop block.
746 rewriter.setInsertionPointToEnd(initBlock);
747 auto memRefType = cast<MemRefType>(atomicOp.getMemref().getType());
748 auto dataPtr = getStridedElementPtr(
749 rewriter, loc, memRefType, adaptor.getMemref(), adaptor.getIndices());
750 Value init = LLVM::LoadOp::create(
751 rewriter, loc, typeConverter->convertType(memRefType.getElementType()),
752 dataPtr);
753 if (needsBitcast)
754 init = LLVM::BitcastOp::create(rewriter, loc, cmpxchgType, init);
755 LLVM::BrOp::create(rewriter, loc, init, loopBlock);
756
757 // Prepare the body of the loop block.
758 rewriter.setInsertionPointToStart(loopBlock);
759
760 // Clone the GenericAtomicRMWOp region and extract the result.
761 Value loopArgument = loopBlock->getArgument(0);
762 Value loopArgForBody = loopArgument;
763 if (needsBitcast)
764 loopArgForBody =
765 LLVM::BitcastOp::create(rewriter, loc, valueType, loopArgument);
766 IRMapping mapping;
767 mapping.map(atomicOp.getCurrentValue(), loopArgForBody);
768 Block &entryBlock = atomicOp.body().front();
769 for (auto &nestedOp : entryBlock.without_terminator()) {
770 Operation *clone = rewriter.clone(nestedOp, mapping);
771 mapping.map(nestedOp.getResults(), clone->getResults());
772 }
773
774 Value result =
775 mapping.lookupOrNull(entryBlock.getTerminator()->getOperand(0));
776 if (!result) {
777 return atomicOp.emitError("result not defined in region");
778 }
779 if (needsBitcast)
780 result = LLVM::BitcastOp::create(rewriter, loc, cmpxchgType, result);
781
782 // Prepare the epilog of the loop block.
783 // Append the cmpxchg op to the end of the loop block.
784 auto successOrdering = LLVM::AtomicOrdering::acq_rel;
785 auto failureOrdering = LLVM::AtomicOrdering::monotonic;
786 auto cmpxchg =
787 LLVM::AtomicCmpXchgOp::create(rewriter, loc, dataPtr, loopArgument,
788 result, successOrdering, failureOrdering);
789 // Extract the %new_loaded and %ok values from the pair.
790 Value newLoaded = LLVM::ExtractValueOp::create(rewriter, loc, cmpxchg, 0);
791 Value ok = LLVM::ExtractValueOp::create(rewriter, loc, cmpxchg, 1);
792
793 // Conditionally branch to the end or back to the loop depending on %ok.
794 LLVM::CondBrOp::create(rewriter, loc, ok, endBlock, ArrayRef<Value>(),
795 loopBlock, newLoaded);
796
797 // The 'result' of the atomic_rmw op is the newly loaded value. Bitcast
798 // back to the float type if needed. Insert at the start of `endBlock` so
799 // the bitcast precedes the existing terminator (split into endBlock).
800 if (needsBitcast) {
801 rewriter.setInsertionPointToStart(endBlock);
802 newLoaded = LLVM::BitcastOp::create(rewriter, loc, valueType, newLoaded);
803 }
804 rewriter.setInsertionPointToEnd(endBlock);
805 rewriter.replaceOp(atomicOp, {newLoaded});
806
807 return success();
808 }
809};
810
811/// Returns the LLVM type of the global variable given the memref type `type`.
812static Type
813convertGlobalMemrefTypeToLLVM(MemRefType type,
814 const LLVMTypeConverter &typeConverter) {
815 // LLVM type for a global memref will be a multi-dimension array. For
816 // declarations or uninitialized global memrefs, we can potentially flatten
817 // this to a 1D array. However, for memref.global's with an initial value,
818 // we do not intend to flatten the ElementsAttribute when going from std ->
819 // LLVM dialect, so the LLVM type needs to me a multi-dimension array.
820 Type elementType = typeConverter.convertType(type.getElementType());
821 Type arrayTy = elementType;
822 // Shape has the outermost dim at index 0, so need to walk it backwards
823 for (int64_t dim : llvm::reverse(type.getShape()))
824 arrayTy = LLVM::LLVMArrayType::get(arrayTy, dim);
825 return arrayTy;
826}
827
828/// GlobalMemrefOp is lowered to a LLVM Global Variable.
829class GlobalMemrefOpLowering : public ConvertOpToLLVMPattern<memref::GlobalOp> {
830 SymbolTableCollection *symbolTables = nullptr;
831
832public:
833 explicit GlobalMemrefOpLowering(const LLVMTypeConverter &typeConverter,
834 SymbolTableCollection *symbolTables = nullptr,
835 PatternBenefit benefit = 1)
836 : ConvertOpToLLVMPattern<memref::GlobalOp>(typeConverter, benefit),
837 symbolTables(symbolTables) {}
838
839 LogicalResult
840 matchAndRewrite(memref::GlobalOp global, OpAdaptor adaptor,
841 ConversionPatternRewriter &rewriter) const override {
842 MemRefType type = global.getType();
843 if (!isConvertibleAndHasIdentityMaps(type))
844 return failure();
845
846 Type arrayTy = convertGlobalMemrefTypeToLLVM(type, *getTypeConverter());
847
848 LLVM::Linkage linkage =
849 global.isPublic() ? LLVM::Linkage::External : LLVM::Linkage::Private;
850 bool isExternal = global.isExternal();
851 bool isUninitialized = global.isUninitialized();
852
853 Attribute initialValue = nullptr;
854 if (!isExternal && !isUninitialized) {
855 auto elementsAttr = llvm::cast<ElementsAttr>(*global.getInitialValue());
856 initialValue = elementsAttr;
857
858 // For scalar memrefs, the global variable created is of the element type,
859 // so unpack the elements attribute to extract the value.
860 if (type.getRank() == 0)
861 initialValue = elementsAttr.getSplatValue<Attribute>();
862 }
863
864 uint64_t alignment = global.getAlignment().value_or(0);
865 FailureOr<unsigned> addressSpace =
866 getTypeConverter()->getMemRefAddressSpace(type);
867 if (failed(addressSpace))
868 return global.emitOpError(
869 "memory space cannot be converted to an integer address space");
870
871 // Remove old operation from symbol table.
872 SymbolTable *symbolTable = nullptr;
873 if (symbolTables) {
874 Operation *symbolTableOp =
875 global->getParentWithTrait<OpTrait::SymbolTable>();
876 symbolTable = &symbolTables->getSymbolTable(symbolTableOp);
877 symbolTable->remove(global);
878 }
879
880 // Create new operation.
881 auto newGlobal = rewriter.replaceOpWithNewOp<LLVM::GlobalOp>(
882 global, arrayTy, global.getConstant(), linkage, global.getSymName(),
883 initialValue, alignment, *addressSpace);
884
885 // Insert new operation into symbol table.
886 if (symbolTable)
887 symbolTable->insert(newGlobal, rewriter.getInsertionPoint());
888
889 if (!isExternal && isUninitialized) {
890 rewriter.createBlock(&newGlobal.getInitializerRegion());
891 Value undef[] = {
892 LLVM::UndefOp::create(rewriter, newGlobal.getLoc(), arrayTy)};
893 LLVM::ReturnOp::create(rewriter, newGlobal.getLoc(), undef);
894 }
895 return success();
896 }
897};
898
899/// GetGlobalMemrefOp is lowered into a Memref descriptor with the pointer to
900/// the first element stashed into the descriptor. This reuses
901/// `AllocLikeOpLowering` to reuse the Memref descriptor construction.
902struct GetGlobalMemrefOpLowering
903 : public ConvertOpToLLVMPattern<memref::GetGlobalOp> {
904 using ConvertOpToLLVMPattern<memref::GetGlobalOp>::ConvertOpToLLVMPattern;
905
906 /// Buffer "allocation" for memref.get_global op is getting the address of
907 /// the global variable referenced.
908 LogicalResult
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");
915
916 // Get actual sizes of the memref as values: static sizes are constant
917 // values and dynamic sizes are passed to 'alloc' as operands. In case of
918 // zero-dimensional memref, assume a scalar (size 1).
919 SmallVector<Value, 4> sizes;
920 SmallVector<Value, 4> strides;
921 Value sizeBytes;
922
923 this->getMemRefDescriptorSizes(loc, memRefType, adaptor.getOperands(),
924 rewriter, sizes, strides, sizeBytes, !false);
925
926 MemRefType type = cast<MemRefType>(op.getResult().getType());
927
928 // This is called after a type conversion, which would have failed if this
929 // call fails.
930 FailureOr<unsigned> maybeAddressSpace =
931 getTypeConverter()->getMemRefAddressSpace(type);
932 assert(succeeded(maybeAddressSpace) && "unsupported address space");
933 unsigned memSpace = *maybeAddressSpace;
934
935 Type arrayTy = convertGlobalMemrefTypeToLLVM(type, *getTypeConverter());
936 auto ptrTy = LLVM::LLVMPointerType::get(rewriter.getContext(), memSpace);
937 auto addressOf =
938 LLVM::AddressOfOp::create(rewriter, loc, ptrTy, op.getName());
939
940 // Get the address of the first element in the array by creating a GEP with
941 // the address of the GV as the base, and (rank + 1) number of 0 indices.
942 auto gep =
943 LLVM::GEPOp::create(rewriter, loc, ptrTy, arrayTy, addressOf,
944 SmallVector<LLVM::GEPArg>(type.getRank() + 1, 0));
945
946 // We do not expect the memref obtained using `memref.get_global` to be
947 // ever deallocated. Set the allocated pointer to be known bad value to
948 // help debug if that ever happens.
949 auto intPtrType = getIntPtrType(memSpace);
950 Value deadBeefConst =
951 createIndexAttrConstant(rewriter, op->getLoc(), intPtrType, 0xdeadbeef);
952 auto deadBeefPtr =
953 LLVM::IntToPtrOp::create(rewriter, loc, ptrTy, deadBeefConst);
954
955 // Both allocated and aligned pointers are same. We could potentially stash
956 // a nullptr for the allocated pointer since we do not expect any dealloc.
957 // Create the MemRef descriptor.
958 auto memRefDescriptor = this->createMemRefDescriptor(
959 loc, memRefType, deadBeefPtr, gep, sizes, strides, rewriter);
960
961 // Return the final value of the descriptor.
962 rewriter.replaceOp(op, {memRefDescriptor});
963 return success();
964 }
965};
966
967// Load operation is lowered to obtaining a pointer to the indexed element
968// and loading it.
969struct LoadOpLowering : public LoadStoreOpLowering<memref::LoadOp> {
970 using Base::Base;
971
972 LogicalResult
973 matchAndRewrite(memref::LoadOp loadOp, OpAdaptor adaptor,
974 ConversionPatternRewriter &rewriter) const override {
975 auto type = loadOp.getMemRefType();
976
977 // Per memref.load spec, the indices must be in-bounds:
978 // 0 <= idx < dim_size, and additionally all offsets are non-negative,
979 // hence inbounds and nuw are used when lowering to llvm.getelementptr.
980 Value dataPtr = getStridedElementPtr(
981 rewriter, loadOp.getLoc(), type, adaptor.getMemref(),
982 adaptor.getIndices(), getLoadStoreNoWrapFlags(type));
983 rewriter.replaceOpWithNewOp<LLVM::LoadOp>(
984 loadOp, typeConverter->convertType(type.getElementType()), dataPtr,
985 loadOp.getAlignment().value_or(0), false, loadOp.getNontemporal(),
986 /*isInvariant=*/loadOp.getInvariant());
987 return success();
988 }
989};
990
991// Store operation is lowered to obtaining a pointer to the indexed element,
992// and storing the given value to it.
993struct StoreOpLowering : public LoadStoreOpLowering<memref::StoreOp> {
994 using Base::Base;
995
996 LogicalResult
997 matchAndRewrite(memref::StoreOp op, OpAdaptor adaptor,
998 ConversionPatternRewriter &rewriter) const override {
999 auto type = op.getMemRefType();
1000
1001 // Per memref.store spec, the indices must be in-bounds:
1002 // 0 <= idx < dim_size, and additionally all offsets are non-negative,
1003 // hence inbounds and nuw are used when lowering to llvm.getelementptr.
1004 Value dataPtr = getStridedElementPtr(
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());
1010 return success();
1011 }
1012};
1013
1014// The prefetch operation is lowered in a way similar to the load operation
1015// except that the llvm.prefetch operation is used for replacement.
1016struct PrefetchOpLowering : public LoadStoreOpLowering<memref::PrefetchOp> {
1017 using Base::Base;
1018
1019 LogicalResult
1020 matchAndRewrite(memref::PrefetchOp prefetchOp, OpAdaptor adaptor,
1021 ConversionPatternRewriter &rewriter) const override {
1022 auto type = prefetchOp.getMemRefType();
1023 auto loc = prefetchOp.getLoc();
1024
1025 Value dataPtr = getStridedElementPtr(
1026 rewriter, loc, type, adaptor.getMemref(), adaptor.getIndices());
1027
1028 // Replace with llvm.prefetch.
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);
1035 return success();
1036 }
1037};
1038
1039struct RankOpLowering : public ConvertOpToLLVMPattern<memref::RankOp> {
1040 using ConvertOpToLLVMPattern<memref::RankOp>::ConvertOpToLLVMPattern;
1041
1042 LogicalResult
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)});
1050 return success();
1051 }
1052 if (auto rankedMemRefType = dyn_cast<MemRefType>(operandType)) {
1053 Type indexType = getIndexType();
1054 rewriter.replaceOp(op,
1055 {createIndexAttrConstant(rewriter, loc, indexType,
1056 rankedMemRefType.getRank())});
1057 return success();
1058 }
1059 return failure();
1060 }
1061};
1062
1063struct MemRefCastOpLowering : public ConvertOpToLLVMPattern<memref::CastOp> {
1065
1066 LogicalResult
1067 matchAndRewrite(memref::CastOp memRefCastOp, OpAdaptor adaptor,
1068 ConversionPatternRewriter &rewriter) const override {
1069 Type srcType = memRefCastOp.getOperand().getType();
1070 Type dstType = memRefCastOp.getType();
1071
1072 // memref::CastOp reduce to bitcast in the ranked MemRef case and can be
1073 // used for type erasure. For now they must preserve underlying element type
1074 // and require source and result type to have the same rank. Therefore,
1075 // perform a sanity check that the underlying structs are the same. Once op
1076 // semantics are relaxed we can revisit.
1077 if (isa<MemRefType>(srcType) && isa<MemRefType>(dstType))
1078 if (typeConverter->convertType(srcType) !=
1079 typeConverter->convertType(dstType))
1080 return failure();
1081
1082 // Unranked to unranked cast is disallowed
1083 if (isa<UnrankedMemRefType>(srcType) && isa<UnrankedMemRefType>(dstType))
1084 return failure();
1085
1086 auto targetStructType = typeConverter->convertType(memRefCastOp.getType());
1087 auto loc = memRefCastOp.getLoc();
1088
1089 // For ranked/ranked case, just keep the original descriptor.
1090 if (isa<MemRefType>(srcType) && isa<MemRefType>(dstType)) {
1091 rewriter.replaceOp(memRefCastOp, {adaptor.getSource()});
1092 return success();
1093 }
1094
1095 if (isa<MemRefType>(srcType) && isa<UnrankedMemRefType>(dstType)) {
1096 // Casting ranked to unranked memref type
1097 // Set the rank in the destination from the memref type
1098 // Allocate space on the stack and copy the src memref descriptor
1099 // Set the ptr in the destination to the stack space
1100 auto srcMemRefType = cast<MemRefType>(srcType);
1101 int64_t rank = srcMemRefType.getRank();
1102 // ptr = AllocaOp sizeof(MemRefDescriptor)
1103 auto ptr = getTypeConverter()->promoteOneMemRefDescriptor(
1104 loc, adaptor.getSource(), rewriter);
1105
1106 // rank = ConstantOp srcRank
1107 auto rankVal =
1108 createIndexAttrConstant(rewriter, loc, getIndexType(), rank);
1109 // poison = PoisonOp
1110 UnrankedMemRefDescriptor memRefDesc =
1111 UnrankedMemRefDescriptor::poison(rewriter, loc, targetStructType);
1112 // d1 = InsertValueOp poison, rank, 0
1113 memRefDesc.setRank(rewriter, loc, rankVal);
1114 // d2 = InsertValueOp d1, ptr, 1
1115 memRefDesc.setMemRefDescPtr(rewriter, loc, ptr);
1116 rewriter.replaceOp(memRefCastOp, (Value)memRefDesc);
1117
1118 } else if (isa<UnrankedMemRefType>(srcType) && isa<MemRefType>(dstType)) {
1119 // Casting from unranked type to ranked.
1120 // The operation is assumed to be doing a correct cast. If the destination
1121 // type mismatches the unranked the type, it is undefined behavior.
1122 UnrankedMemRefDescriptor memRefDesc(adaptor.getSource());
1123 // ptr = ExtractValueOp src, 1
1124 auto ptr = memRefDesc.memRefDescPtr(rewriter, loc);
1125
1126 // struct = LoadOp ptr
1127 auto loadOp = LLVM::LoadOp::create(rewriter, loc, targetStructType, ptr);
1128 rewriter.replaceOp(memRefCastOp, loadOp.getResult());
1129 } else {
1130 llvm_unreachable("Unsupported unranked memref to unranked memref cast");
1131 }
1132
1133 return success();
1134 }
1135};
1136
1137/// Pattern to lower a `memref.copy` to llvm.
1138///
1139/// For memrefs with identity layouts, the copy is lowered to the llvm
1140/// `memcpy` intrinsic. For non-identity layouts, the copy is lowered to a call
1141/// to the generic `MemrefCopyFn`.
1142class MemRefCopyOpLowering : public ConvertOpToLLVMPattern<memref::CopyOp> {
1143 SymbolTableCollection *symbolTables = nullptr;
1144
1145public:
1146 explicit MemRefCopyOpLowering(const LLVMTypeConverter &typeConverter,
1147 SymbolTableCollection *symbolTables = nullptr,
1148 PatternBenefit benefit = 1)
1149 : ConvertOpToLLVMPattern<memref::CopyOp>(typeConverter, benefit),
1150 symbolTables(symbolTables) {}
1151
1152 LogicalResult
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());
1157
1158 MemRefDescriptor srcDesc(adaptor.getSource());
1159
1160 // Compute number of elements.
1161 Value numElements =
1162 createIndexAttrConstant(rewriter, loc, getIndexType(), 1);
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);
1166 }
1167
1168 // Get element size.
1169 auto sizeInBytes = getSizeInBytes(loc, srcType.getElementType(), rewriter);
1170 // Compute total.
1171 Value totalSize =
1172 LLVM::MulOp::create(rewriter, loc, numElements, sizeInBytes);
1173
1174 Type elementType = typeConverter->convertType(srcType.getElementType());
1175
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);
1183 Value targetPtr =
1184 LLVM::GEPOp::create(rewriter, loc, targetBasePtr.getType(), elementType,
1185 targetBasePtr, targetOffset);
1186 LLVM::MemcpyOp::create(rewriter, loc, targetPtr, srcPtr, totalSize,
1187 /*isVolatile=*/false);
1188 rewriter.eraseOp(op);
1189
1190 return success();
1191 }
1192
1193 LogicalResult
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());
1199
1200 // First make sure we have an unranked memref descriptor representation.
1201 auto makeUnranked = [&, this](Value ranked, MemRefType type) {
1202 auto rank = LLVM::ConstantOp::create(rewriter, loc, getIndexType(),
1203 type.getRank());
1204 auto *typeConverter = getTypeConverter();
1205 auto ptr =
1206 typeConverter->promoteOneMemRefDescriptor(loc, ranked, rewriter);
1207
1208 auto unrankedType =
1209 UnrankedMemRefType::get(type.getElementType(), type.getMemorySpace());
1211 rewriter, loc, *typeConverter, unrankedType, ValueRange{rank, ptr});
1212 };
1213
1214 // Save stack position before promoting descriptors
1215 auto stackSaveOp = LLVM::StackSaveOp::create(rewriter, loc, getPtrType());
1216
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();
1225
1226 // Now promote the unranked descriptors to the stack.
1227 auto one = createIndexAttrConstant(rewriter, loc, getIndexType(), 1);
1228 auto promote = [&](Value desc) {
1229 auto ptrType = LLVM::LLVMPointerType::get(rewriter.getContext());
1230 auto allocated =
1231 LLVM::AllocaOp::create(rewriter, loc, ptrType, desc.getType(), one);
1232 LLVM::StoreOp::create(rewriter, loc, desc, allocated);
1233 return allocated;
1234 };
1235
1236 auto sourcePtr = promote(unrankedSource);
1237 auto targetPtr = promote(unrankedTarget);
1238
1239 // Derive size from llvm.getelementptr which will account for any
1240 // potential alignment
1241 auto elemSize = getSizeInBytes(loc, srcType.getElementType(), rewriter);
1242 auto copyFn = LLVM::lookupOrCreateMemRefCopyFn(
1243 rewriter, op->getParentOfType<ModuleOp>(), getIndexType(),
1244 sourcePtr.getType(), symbolTables);
1245 if (failed(copyFn))
1246 return failure();
1247 LLVM::CallOp::create(rewriter, loc, copyFn.value(),
1248 ValueRange{elemSize, sourcePtr, targetPtr});
1249
1250 // Restore stack used for descriptors
1251 LLVM::StackRestoreOp::create(rewriter, loc, stackSaveOp);
1252
1253 rewriter.eraseOp(op);
1254
1255 return success();
1256 }
1257
1258 LogicalResult
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());
1263
1264 auto isContiguousMemrefType = [&](BaseMemRefType type) {
1265 auto memrefType = dyn_cast<mlir::MemRefType>(type);
1266 // We can use memcpy for memrefs if they have an identity layout or are
1267 // contiguous with an arbitrary offset. Ignore empty memrefs, which is a
1268 // special case handled by memrefCopy.
1269 return memrefType &&
1270 (memrefType.getLayout().isIdentity() ||
1271 (memrefType.hasStaticShape() && memrefType.getNumElements() > 0 &&
1273 };
1274
1275 if (isContiguousMemrefType(srcType) && isContiguousMemrefType(targetType))
1276 return lowerToMemCopyIntrinsic(op, adaptor, rewriter);
1277
1278 return lowerToMemCopyFunctionCall(op, adaptor, rewriter);
1279 }
1280};
1281
1282struct MemorySpaceCastOpLowering
1283 : public ConvertOpToLLVMPattern<memref::MemorySpaceCastOp> {
1284 using ConvertOpToLLVMPattern<
1285 memref::MemorySpaceCastOp>::ConvertOpToLLVMPattern;
1286
1287 LogicalResult
1288 matchAndRewrite(memref::MemorySpaceCastOp op, OpAdaptor adaptor,
1289 ConversionPatternRewriter &rewriter) const override {
1290 Location loc = op.getLoc();
1291
1292 Type resultType = op.getDest().getType();
1293 if (auto resultTypeR = dyn_cast<MemRefType>(resultType)) {
1294 auto convertedType =
1295 typeConverter->convertType<LLVM::LLVMStructType>(resultTypeR);
1296 if (!convertedType)
1297 return rewriter.notifyMatchFailure(op, "memref type conversion failed");
1298 Type newPtrType = convertedType.getBody()[0];
1299
1300 SmallVector<Value> descVals;
1301 MemRefDescriptor::unpack(rewriter, loc, adaptor.getSource(), resultTypeR,
1302 descVals);
1303 descVals[0] =
1304 LLVM::AddrSpaceCastOp::create(rewriter, loc, newPtrType, descVals[0]);
1305 descVals[1] =
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);
1310 return success();
1311 }
1312 if (auto resultTypeU = dyn_cast<UnrankedMemRefType>(resultType)) {
1313 // Since the type converter won't be doing this for us, get the address
1314 // space.
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;
1328
1329 UnrankedMemRefDescriptor sourceDesc(adaptor.getSource());
1330 Value rank = sourceDesc.rank(rewriter, loc);
1331 Value sourceUnderlyingDesc = sourceDesc.memRefDescPtr(rewriter, loc);
1332
1333 // Create and allocate storage for new memref descriptor.
1335 rewriter, loc, typeConverter->convertType(resultTypeU));
1336 result.setRank(rewriter, loc, rank);
1337 Value resultUnderlyingSize = UnrankedMemRefDescriptor::computeSize(
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);
1343
1344 // Copy pointers, performing address space casts.
1345 auto sourceElemPtrType =
1346 LLVM::LLVMPointerType::get(rewriter.getContext(), sourceAddrSpace);
1347 auto resultElemPtrType =
1348 LLVM::LLVMPointerType::get(rewriter.getContext(), resultAddrSpace);
1349
1350 Value allocatedPtr = sourceDesc.allocatedPtr(
1351 rewriter, loc, sourceUnderlyingDesc, sourceElemPtrType);
1352 Value alignedPtr =
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);
1359
1360 result.setAllocatedPtr(rewriter, loc, resultUnderlyingDesc,
1361 resultElemPtrType, allocatedPtr);
1362 result.setAlignedPtr(rewriter, loc, *getTypeConverter(),
1363 resultUnderlyingDesc, resultElemPtrType, alignedPtr);
1364
1365 // Copy all the index-valued operands.
1366 Value sourceIndexVals =
1367 sourceDesc.offsetBasePtr(rewriter, loc, *getTypeConverter(),
1368 sourceUnderlyingDesc, sourceElemPtrType);
1369 Value resultIndexVals =
1370 result.offsetBasePtr(rewriter, loc, *getTypeConverter(),
1371 resultUnderlyingDesc, resultElemPtrType);
1372
1373 int64_t bytesToSkip =
1374 2 * llvm::divideCeil(
1375 getTypeConverter()->getPointerBitwidth(resultAddrSpace), 8);
1376 Value bytesToSkipConst =
1377 createIndexAttrConstant(rewriter, loc, getIndexType(), bytesToSkip);
1378 Value copySize =
1379 LLVM::SubOp::create(rewriter, loc, getIndexType(),
1380 resultUnderlyingSize, bytesToSkipConst);
1381 LLVM::MemcpyOp::create(rewriter, loc, resultIndexVals, sourceIndexVals,
1382 copySize, /*isVolatile=*/false);
1383
1384 rewriter.replaceOp(op, ValueRange{result});
1385 return success();
1386 }
1387 return rewriter.notifyMatchFailure(loc, "unexpected memref type");
1388 }
1389};
1390
1391/// Extracts allocated, aligned pointers and offset from a ranked or unranked
1392/// memref type. In unranked case, the fields are extracted from the underlying
1393/// ranked descriptor.
1394static void extractPointersAndOffset(Location loc,
1395 ConversionPatternRewriter &rewriter,
1396 const LLVMTypeConverter &typeConverter,
1397 Value originalOperand,
1398 Value convertedOperand,
1399 Value *allocatedPtr, Value *alignedPtr,
1400 Value *offset = nullptr) {
1401 Type operandType = originalOperand.getType();
1402 if (isa<MemRefType>(operandType)) {
1403 MemRefDescriptor desc(convertedOperand);
1404 *allocatedPtr = desc.allocatedPtr(rewriter, loc);
1405 *alignedPtr = desc.alignedPtr(rewriter, loc);
1406 if (offset != nullptr)
1407 *offset = desc.offset(rewriter, loc);
1408 return;
1409 }
1410
1411 // These will all cause assert()s on unconvertible types.
1412 unsigned memorySpace = *typeConverter.getMemRefAddressSpace(
1413 cast<UnrankedMemRefType>(operandType));
1414 auto elementPtrType =
1415 LLVM::LLVMPointerType::get(rewriter.getContext(), memorySpace);
1416
1417 // Extract pointer to the underlying ranked memref descriptor and cast it to
1418 // ElemType**.
1419 UnrankedMemRefDescriptor unrankedDesc(convertedOperand);
1420 Value underlyingDescPtr = unrankedDesc.memRefDescPtr(rewriter, loc);
1421
1423 rewriter, loc, underlyingDescPtr, elementPtrType);
1425 rewriter, loc, typeConverter, underlyingDescPtr, elementPtrType);
1426 if (offset != nullptr) {
1428 rewriter, loc, typeConverter, underlyingDescPtr, elementPtrType);
1429 }
1430}
1431
1432struct MemRefReinterpretCastOpLowering
1433 : public ConvertOpToLLVMPattern<memref::ReinterpretCastOp> {
1434 using ConvertOpToLLVMPattern<
1435 memref::ReinterpretCastOp>::ConvertOpToLLVMPattern;
1436
1437 LogicalResult
1438 matchAndRewrite(memref::ReinterpretCastOp castOp, OpAdaptor adaptor,
1439 ConversionPatternRewriter &rewriter) const override {
1440 Type srcType = castOp.getSource().getType();
1441
1442 Value descriptor;
1443 if (failed(convertSourceMemRefToDescriptor(rewriter, srcType, castOp,
1444 adaptor, &descriptor)))
1445 return failure();
1446 rewriter.replaceOp(castOp, {descriptor});
1447 return success();
1448 }
1449
1450private:
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)
1460 return failure();
1461
1462 // Create descriptor.
1463 Location loc = castOp.getLoc();
1464 auto desc = MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy);
1465
1466 // Set allocated and aligned pointers.
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);
1473
1474 // Set offset.
1475 if (castOp.isDynamicOffset(0))
1476 desc.setOffset(rewriter, loc, adaptor.getOffsets()[0]);
1477 else
1478 desc.setConstantOffset(rewriter, loc, castOp.getStaticOffset(0));
1479
1480 // Set sizes and strides.
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++]);
1486 else
1487 desc.setConstantSize(rewriter, loc, i, castOp.getStaticSize(i));
1488
1489 if (castOp.isDynamicStride(i))
1490 desc.setStride(rewriter, loc, i, adaptor.getStrides()[dynStrideId++]);
1491 else
1492 desc.setConstantStride(rewriter, loc, i, castOp.getStaticStride(i));
1493 }
1494 *descriptor = desc;
1495 return success();
1496 }
1497};
1498
1499struct MemRefReshapeOpLowering
1500 : public ConvertOpToLLVMPattern<memref::ReshapeOp> {
1501 using ConvertOpToLLVMPattern<memref::ReshapeOp>::ConvertOpToLLVMPattern;
1502
1503 LogicalResult
1504 matchAndRewrite(memref::ReshapeOp reshapeOp, OpAdaptor adaptor,
1505 ConversionPatternRewriter &rewriter) const override {
1506 Type srcType = reshapeOp.getSource().getType();
1507
1508 Value descriptor;
1509 if (failed(convertSourceMemRefToDescriptor(rewriter, srcType, reshapeOp,
1510 adaptor, &descriptor)))
1511 return failure();
1512 rewriter.replaceOp(reshapeOp, {descriptor});
1513 return success();
1514 }
1515
1516private:
1517 LogicalResult
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)
1529 return failure();
1530
1531 // Create descriptor.
1532 Location loc = reshapeOp.getLoc();
1533 auto desc =
1534 MemRefDescriptor::poison(rewriter, loc, llvmTargetDescriptorTy);
1535
1536 // Set allocated and aligned pointers.
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);
1543
1544 // Extract the offset and strides from the type.
1545 int64_t offset;
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");
1550
1551 if (!isStaticStrideOrOffset(offset))
1552 return rewriter.notifyMatchFailure(reshapeOp,
1553 "dynamic offset is unsupported");
1554
1555 desc.setConstantOffset(rewriter, loc, offset);
1556
1557 assert(targetMemRefType.getLayout().isIdentity() &&
1558 "Identity layout map is a precondition of a valid reshape op");
1559
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])) {
1565 // If the stride for this dimension is dynamic, then use the product
1566 // of the sizes of the inner dimensions.
1567 stride =
1568 createIndexAttrConstant(rewriter, loc, indexType, strides[i]);
1569 } else if (!stride) {
1570 // `stride` is null only in the first iteration of the loop. However,
1571 // since the target memref has an identity layout, we can safely set
1572 // the innermost stride to 1.
1573 stride = createIndexAttrConstant(rewriter, loc, indexType, 1);
1574 }
1575
1576 Value dimSize;
1577 // If the size of this dimension is dynamic, then load it at runtime
1578 // from the shape operand.
1579 if (!targetMemRefType.isDynamicDim(i)) {
1580 dimSize = createIndexAttrConstant(rewriter, loc, indexType,
1581 targetMemRefType.getDimSize(i));
1582 } else {
1583 Value shapeOp = reshapeOp.getShape();
1584 Value index = createIndexAttrConstant(rewriter, loc, indexType, i);
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");
1591 }
1592
1593 desc.setSize(rewriter, loc, i, dimSize);
1594 desc.setStride(rewriter, loc, i, stride);
1595
1596 // Prepare the stride value for the next dimension.
1597 stride = LLVM::MulOp::create(rewriter, loc, stride, dimSize);
1598 }
1599
1600 *descriptor = desc;
1601 return success();
1602 }
1603
1604 // The shape is a rank-1 tensor with unknown length.
1605 Location loc = reshapeOp.getLoc();
1606 MemRefDescriptor shapeDesc(adaptor.getShape());
1607 Value resultRank = shapeDesc.size(rewriter, loc, 0);
1608
1609 // Extract address space and element type.
1610 auto targetType = cast<UnrankedMemRefType>(reshapeOp.getResult().getType());
1611 unsigned addressSpace =
1612 *getTypeConverter()->getMemRefAddressSpace(targetType);
1613
1614 // Create the unranked memref descriptor that holds the ranked one. The
1615 // inner descriptor is allocated on stack.
1616 auto targetDesc = UnrankedMemRefDescriptor::poison(
1617 rewriter, loc, typeConverter->convertType(targetType));
1618 targetDesc.setRank(rewriter, loc, resultRank);
1619 Value allocationSize = UnrankedMemRefDescriptor::computeSize(
1620 rewriter, loc, *getTypeConverter(), targetDesc, addressSpace);
1621 Value underlyingDescPtr = LLVM::AllocaOp::create(
1622 rewriter, loc, getPtrType(), IntegerType::get(getContext(), 8),
1623 allocationSize);
1624 targetDesc.setMemRefDescPtr(rewriter, loc, underlyingDescPtr);
1625
1626 // Extract pointers and offset from the source memref.
1627 Value allocatedPtr, alignedPtr, offset;
1628 extractPointersAndOffset(loc, rewriter, *getTypeConverter(),
1629 reshapeOp.getSource(), adaptor.getSource(),
1630 &allocatedPtr, &alignedPtr, &offset);
1631
1632 // Set pointers and offset.
1633 auto elementPtrType =
1634 LLVM::LLVMPointerType::get(rewriter.getContext(), addressSpace);
1635
1636 UnrankedMemRefDescriptor::setAllocatedPtr(rewriter, loc, underlyingDescPtr,
1637 elementPtrType, allocatedPtr);
1638 UnrankedMemRefDescriptor::setAlignedPtr(rewriter, loc, *getTypeConverter(),
1639 underlyingDescPtr, elementPtrType,
1640 alignedPtr);
1641 UnrankedMemRefDescriptor::setOffset(rewriter, loc, *getTypeConverter(),
1642 underlyingDescPtr, elementPtrType,
1643 offset);
1644
1645 // Use the offset pointer as base for further addressing. Copy over the new
1646 // shape and compute strides. For this, we create a loop from rank-1 to 0.
1647 Value targetSizesBase = UnrankedMemRefDescriptor::sizeBasePtr(
1648 rewriter, loc, *getTypeConverter(), underlyingDescPtr, elementPtrType);
1649 Value targetStridesBase = UnrankedMemRefDescriptor::strideBasePtr(
1650 rewriter, loc, *getTypeConverter(), targetSizesBase, resultRank);
1651 Value shapeOperandPtr = shapeDesc.alignedPtr(rewriter, loc);
1652 Value oneIndex = createIndexAttrConstant(rewriter, loc, getIndexType(), 1);
1653 Value resultRankMinusOne =
1654 LLVM::SubOp::create(rewriter, loc, resultRank, oneIndex);
1655
1656 Block *initBlock = rewriter.getInsertionBlock();
1657 Type indexType = getTypeConverter()->getIndexType();
1658 Block::iterator remainingOpsIt = std::next(rewriter.getInsertionPoint());
1659
1660 Block *condBlock = rewriter.createBlock(initBlock->getParent(), {},
1661 {indexType, indexType}, {loc, loc});
1662
1663 // Move the remaining initBlock ops to condBlock.
1664 Block *remainingBlock = rewriter.splitBlock(initBlock, remainingOpsIt);
1665 rewriter.mergeBlocks(remainingBlock, condBlock, ValueRange());
1666
1667 rewriter.setInsertionPointToEnd(initBlock);
1668 LLVM::BrOp::create(rewriter, loc,
1669 ValueRange({resultRankMinusOne, oneIndex}), condBlock);
1670 rewriter.setInsertionPointToStart(condBlock);
1671 Value indexArg = condBlock->getArgument(0);
1672 Value strideArg = condBlock->getArgument(1);
1673
1674 Value zeroIndex = createIndexAttrConstant(rewriter, loc, indexType, 0);
1675 Value pred = LLVM::ICmpOp::create(
1676 rewriter, loc, IntegerType::get(rewriter.getContext(), 1),
1677 LLVM::ICmpPredicate::sge, indexArg, zeroIndex);
1678
1679 Block *bodyBlock =
1680 rewriter.splitBlock(condBlock, rewriter.getInsertionPoint());
1681 rewriter.setInsertionPointToStart(bodyBlock);
1682
1683 // Copy size from shape to descriptor.
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);
1690 UnrankedMemRefDescriptor::setSize(rewriter, loc, *getTypeConverter(),
1691 targetSizesBase, indexArg, size);
1692
1693 // Write stride value and compute next one.
1694 UnrankedMemRefDescriptor::setStride(rewriter, loc, *getTypeConverter(),
1695 targetStridesBase, indexArg, strideArg);
1696 Value nextStride = LLVM::MulOp::create(rewriter, loc, strideArg, size);
1697
1698 // Decrement loop counter and branch back.
1699 Value decrement = LLVM::SubOp::create(rewriter, loc, indexArg, oneIndex);
1700 LLVM::BrOp::create(rewriter, loc, ValueRange({decrement, nextStride}),
1701 condBlock);
1702
1703 Block *remainder =
1704 rewriter.splitBlock(bodyBlock, rewriter.getInsertionPoint());
1705
1706 // Hook up the cond exit to the remainder.
1707 rewriter.setInsertionPointToEnd(condBlock);
1708 LLVM::CondBrOp::create(rewriter, loc, pred, bodyBlock, ValueRange(),
1709 remainder, ValueRange());
1710
1711 // Reset position to beginning of new remainder block.
1712 rewriter.setInsertionPointToStart(remainder);
1713
1714 *descriptor = targetDesc;
1715 return success();
1716 }
1717};
1718
1719/// RessociatingReshapeOp must be expanded before we reach this stage.
1720/// Report that information.
1721template <typename ReshapeOp>
1722class ReassociatingReshapeOpConversion
1723 : public ConvertOpToLLVMPattern<ReshapeOp> {
1724public:
1725 using ConvertOpToLLVMPattern<ReshapeOp>::ConvertOpToLLVMPattern;
1726 using ReshapeOpAdaptor = typename ReshapeOp::Adaptor;
1727
1728 LogicalResult
1729 matchAndRewrite(ReshapeOp reshapeOp, typename ReshapeOp::Adaptor adaptor,
1730 ConversionPatternRewriter &rewriter) const override {
1731 return rewriter.notifyMatchFailure(
1732 reshapeOp,
1733 "reassociation operations should have been expanded beforehand");
1734 }
1735};
1736
1737/// Subviews must be expanded before we reach this stage.
1738/// Report that information.
1739struct SubViewOpLowering : public ConvertOpToLLVMPattern<memref::SubViewOp> {
1740 using ConvertOpToLLVMPattern<memref::SubViewOp>::ConvertOpToLLVMPattern;
1741
1742 LogicalResult
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");
1747 }
1748};
1749
1750/// Conversion pattern that transforms a transpose op into:
1751/// 1. A function entry `alloca` operation to allocate a ViewDescriptor.
1752/// 2. A load of the ViewDescriptor from the pointer allocated in 1.
1753/// 3. Updates to the ViewDescriptor to introduce the data ptr, offset, size
1754/// and stride. Size and stride are permutations of the original values.
1755/// 4. A store of the resulting ViewDescriptor to the alloca'ed pointer.
1756/// The transpose op is replaced by the alloca'ed pointer.
1757class TransposeOpLowering : public ConvertOpToLLVMPattern<memref::TransposeOp> {
1758public:
1759 using ConvertOpToLLVMPattern<memref::TransposeOp>::ConvertOpToLLVMPattern;
1760
1761 LogicalResult
1762 matchAndRewrite(memref::TransposeOp transposeOp, OpAdaptor adaptor,
1763 ConversionPatternRewriter &rewriter) const override {
1764 auto loc = transposeOp.getLoc();
1765 MemRefDescriptor viewMemRef(adaptor.getIn());
1766
1767 // No permutation, early exit.
1768 if (transposeOp.getPermutation().isIdentity())
1769 return rewriter.replaceOp(transposeOp, {viewMemRef}), success();
1770
1771 auto targetMemRef = MemRefDescriptor::poison(
1772 rewriter, loc,
1773 typeConverter->convertType(transposeOp.getIn().getType()));
1774
1775 // Copy the base and aligned pointers from the old descriptor to the new
1776 // one.
1777 targetMemRef.setAllocatedPtr(rewriter, loc,
1778 viewMemRef.allocatedPtr(rewriter, loc));
1779 targetMemRef.setAlignedPtr(rewriter, loc,
1780 viewMemRef.alignedPtr(rewriter, loc));
1781
1782 // Copy the offset pointer from the old descriptor to the new one.
1783 targetMemRef.setOffset(rewriter, loc, viewMemRef.offset(rewriter, loc));
1784
1785 // Iterate over the dimensions and apply size/stride permutation:
1786 // When enumerating the results of the permutation map, the enumeration
1787 // index is the index into the target dimensions and the DimExpr points to
1788 // the dimension of the source memref.
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));
1797 }
1798
1799 rewriter.replaceOp(transposeOp, {targetMemRef});
1800 return success();
1801 }
1802};
1803
1804/// Conversion pattern that transforms an op into:
1805/// 1. An `llvm.mlir.undef` operation to create a memref descriptor
1806/// 2. Updates to the descriptor to introduce the data ptr, offset, size
1807/// and stride.
1808/// The view op is replaced by the descriptor.
1809struct ViewOpLowering : public ConvertOpToLLVMPattern<memref::ViewOp> {
1810 using ConvertOpToLLVMPattern<memref::ViewOp>::ConvertOpToLLVMPattern;
1811
1812 // Build and return the value for the idx^th shape dimension, either by
1813 // returning the constant shape dimension or counting the proper dynamic size.
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]))
1819 return createIndexAttrConstant(rewriter, loc, indexType, shape[idx]);
1820 // Count the number of dynamic dims in range [0, idx]
1821 unsigned nDynamic =
1822 llvm::count_if(shape.take_front(idx), ShapedType::isDynamic);
1823 return dynamicSizes[nDynamic];
1824 }
1825
1826 // Build and return the idx^th stride, either by returning the constant stride
1827 // or by computing the dynamic stride from the current `runningStride` and
1828 // `nextSize`. The caller should keep a running stride and update it with the
1829 // result returned by this function.
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]))
1835 return createIndexAttrConstant(rewriter, loc, indexType, strides[idx]);
1836 if (nextSize)
1837 return runningStride
1838 ? LLVM::MulOp::create(rewriter, loc, runningStride, nextSize)
1839 : nextSize;
1840 assert(!runningStride);
1841 return createIndexAttrConstant(rewriter, loc, indexType, 1);
1842 }
1843
1844 LogicalResult
1845 matchAndRewrite(memref::ViewOp viewOp, OpAdaptor adaptor,
1846 ConversionPatternRewriter &rewriter) const override {
1847 auto loc = viewOp.getLoc();
1848
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"),
1857 failure();
1858
1859 int64_t offset;
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");
1865
1866 // Target memref must be contiguous in memory (innermost stride is 1), or
1867 // empty (special case when at least one of the memref dimensions is 0).
1868 if (!strides.empty() && (strides.back() != 1 && strides.back() != 0))
1869 return viewOp.emitWarning("cannot cast to non-contiguous shape"),
1870 failure();
1871
1872 // Create the descriptor.
1873 MemRefDescriptor sourceMemRef(adaptor.getSource());
1874 auto targetMemRef = MemRefDescriptor::poison(rewriter, loc, targetDescTy);
1875
1876 // Field 1: Copy the allocated pointer, used for malloc/free.
1877 Value allocatedPtr = sourceMemRef.allocatedPtr(rewriter, loc);
1878 auto srcMemRefType = cast<MemRefType>(viewOp.getSource().getType());
1879 targetMemRef.setAllocatedPtr(rewriter, loc, allocatedPtr);
1880
1881 // Field 2: Copy the actual aligned pointer to payload.
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());
1887
1888 targetMemRef.setAlignedPtr(rewriter, loc, alignedPtr);
1889
1890 Type indexType = getIndexType();
1891 // Field 3: The offset in the resulting type must be 0. This is
1892 // because of the type change: an offset on srcType* may not be
1893 // expressible as an offset on dstType*.
1894 targetMemRef.setOffset(
1895 rewriter, loc,
1896 createIndexAttrConstant(rewriter, loc, indexType, offset));
1897
1898 // Early exit for 0-D corner case.
1899 if (viewMemRefType.getRank() == 0)
1900 return rewriter.replaceOp(viewOp, {targetMemRef}), success();
1901
1902 // Fields 4 and 5: Update sizes and strides.
1903 Value stride = nullptr, nextSize = nullptr;
1904 for (int i = viewMemRefType.getRank() - 1; i >= 0; --i) {
1905 // Update size.
1906 Value size = getSize(rewriter, loc, viewMemRefType.getShape(),
1907 adaptor.getSizes(), i, indexType);
1908 targetMemRef.setSize(rewriter, loc, i, size);
1909 // Update stride.
1910 stride =
1911 getStride(rewriter, loc, strides, nextSize, stride, i, indexType);
1912 targetMemRef.setStride(rewriter, loc, i, stride);
1913 nextSize = size;
1914 }
1915
1916 rewriter.replaceOp(viewOp, {targetMemRef});
1917 return success();
1918 }
1919};
1920
1921//===----------------------------------------------------------------------===//
1922// AtomicRMWOpLowering
1923//===----------------------------------------------------------------------===//
1924
1925/// Try to match the kind of a memref.atomic_rmw to determine whether to use a
1926/// lowering to llvm.atomicrmw or fallback to llvm.cmpxchg.
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:
1937 // TODO: remove this by end of 2025.
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:
1948 // TODO: remove this by end of 2025.
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;
1964 default:
1965 return std::nullopt;
1966 }
1967 llvm_unreachable("Invalid AtomicRMWKind");
1968}
1969
1970struct AtomicRMWOpLowering : public LoadStoreOpLowering<memref::AtomicRMWOp> {
1971 using Base::Base;
1972
1973 LogicalResult
1974 matchAndRewrite(memref::AtomicRMWOp atomicOp, OpAdaptor adaptor,
1975 ConversionPatternRewriter &rewriter) const override {
1976 auto maybeKind = matchSimpleAtomicOp(atomicOp);
1977 if (!maybeKind)
1978 return failure();
1979 auto memRefType = atomicOp.getMemRefType();
1980 SmallVector<int64_t> strides;
1981 int64_t offset;
1982 if (failed(memRefType.getStridesAndOffset(strides, offset)))
1983 return failure();
1984 auto dataPtr =
1985 getStridedElementPtr(rewriter, atomicOp.getLoc(), memRefType,
1986 adaptor.getMemref(), adaptor.getIndices());
1987 rewriter.replaceOpWithNewOp<LLVM::AtomicRMWOp>(
1988 atomicOp, *maybeKind, dataPtr, adaptor.getValue(),
1989 LLVM::AtomicOrdering::acq_rel);
1990 return success();
1991 }
1992};
1993
1994/// Unpack the pointer returned by a memref.extract_aligned_pointer_as_index.
1995class ConvertExtractAlignedPointerAsIndex
1996 : public ConvertOpToLLVMPattern<memref::ExtractAlignedPointerAsIndexOp> {
1997public:
1998 using ConvertOpToLLVMPattern<
1999 memref::ExtractAlignedPointerAsIndexOp>::ConvertOpToLLVMPattern;
2000
2001 LogicalResult
2002 matchAndRewrite(memref::ExtractAlignedPointerAsIndexOp extractOp,
2003 OpAdaptor adaptor,
2004 ConversionPatternRewriter &rewriter) const override {
2005 BaseMemRefType sourceTy = extractOp.getSource().getType();
2006
2007 Value alignedPtr;
2008 if (sourceTy.hasRank()) {
2009 MemRefDescriptor desc(adaptor.getSource());
2010 alignedPtr = desc.alignedPtr(rewriter, extractOp->getLoc());
2011 } else {
2012 auto elementPtrTy = LLVM::LLVMPointerType::get(
2013 rewriter.getContext(), sourceTy.getMemorySpaceAsInt());
2014
2015 UnrankedMemRefDescriptor desc(adaptor.getSource());
2016 Value descPtr = desc.memRefDescPtr(rewriter, extractOp->getLoc());
2017
2019 rewriter, extractOp->getLoc(), *getTypeConverter(), descPtr,
2020 elementPtrTy);
2021 }
2022
2023 rewriter.replaceOpWithNewOp<LLVM::PtrToIntOp>(
2024 extractOp, getTypeConverter()->getIndexType(), alignedPtr);
2025 return success();
2026 }
2027};
2028
2029/// Materialize the MemRef descriptor represented by the results of
2030/// ExtractStridedMetadataOp.
2031class ExtractStridedMetadataOpLowering
2032 : public ConvertOpToLLVMPattern<memref::ExtractStridedMetadataOp> {
2033public:
2034 using ConvertOpToLLVMPattern<
2035 memref::ExtractStridedMetadataOp>::ConvertOpToLLVMPattern;
2036
2037 LogicalResult
2038 matchAndRewrite(memref::ExtractStridedMetadataOp extractStridedMetadataOp,
2039 OpAdaptor adaptor,
2040 ConversionPatternRewriter &rewriter) const override {
2041
2042 if (!LLVM::isCompatibleType(adaptor.getOperands().front().getType()))
2043 return failure();
2044
2045 // Create the descriptor.
2046 MemRefDescriptor sourceMemRef(adaptor.getSource());
2047 Location loc = extractStridedMetadataOp.getLoc();
2048 Value source = extractStridedMetadataOp.getSource();
2049
2050 auto sourceMemRefType = cast<MemRefType>(source.getType());
2051 int64_t rank = sourceMemRefType.getRank();
2052 SmallVector<Value> results;
2053 results.reserve(2 + rank * 2);
2054
2055 // Base buffer.
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);
2063
2064 // Offset.
2065 results.push_back(sourceMemRef.offset(rewriter, loc));
2066
2067 // Sizes.
2068 for (unsigned i = 0; i < rank; ++i)
2069 results.push_back(sourceMemRef.size(rewriter, loc, i));
2070 // Strides.
2071 for (unsigned i = 0; i < rank; ++i)
2072 results.push_back(sourceMemRef.stride(rewriter, loc, i));
2073
2074 rewriter.replaceOp(extractStridedMetadataOp, results);
2075 return success();
2076 }
2077};
2078
2079} // namespace
2080
2082 const LLVMTypeConverter &converter, RewritePatternSet &patterns,
2083 SymbolTableCollection *symbolTables) {
2084 // clang-format off
2085 patterns.add<
2086 AllocaOpLowering,
2087 AllocaScopeOpLowering,
2088 AssumeAlignmentOpLowering,
2089 AtomicRMWOpLowering,
2090 ConvertExtractAlignedPointerAsIndex,
2091 DimOpLowering,
2092 DistinctObjectsOpLowering,
2093 ExtractStridedMetadataOpLowering,
2094 GenericAtomicRMWOpLowering,
2095 GetGlobalMemrefOpLowering,
2096 LoadOpLowering,
2097 MemRefCastOpLowering,
2098 MemRefReinterpretCastOpLowering,
2099 MemRefReshapeOpLowering,
2100 MemorySpaceCastOpLowering,
2101 PrefetchOpLowering,
2102 RankOpLowering,
2103 ReassociatingReshapeOpConversion<memref::CollapseShapeOp>,
2104 ReassociatingReshapeOpConversion<memref::ExpandShapeOp>,
2105 StoreOpLowering,
2106 SubViewOpLowering,
2108 ViewOpLowering>(converter);
2109 // clang-format on
2110 patterns.add<GlobalMemrefOpLowering, MemRefCopyOpLowering>(converter,
2111 symbolTables);
2112 auto allocLowering = converter.getOptions().allocLowering;
2114 patterns.add<AlignedAllocOpLowering, DeallocOpLowering>(converter,
2115 symbolTables);
2116 else if (allocLowering == LowerToLLVMOptions::AllocLowering::Malloc)
2117 patterns.add<AllocOpLowering, DeallocOpLowering>(converter, symbolTables);
2118}
2119
2120namespace {
2121struct FinalizeMemRefToLLVMConversionPass
2122 : public impl::FinalizeMemRefToLLVMConversionPassBase<
2123 FinalizeMemRefToLLVMConversionPass> {
2124 using FinalizeMemRefToLLVMConversionPassBase::
2125 FinalizeMemRefToLLVMConversionPassBase;
2126
2127 void runOnOperation() override {
2128 Operation *op = getOperation();
2129 const auto &dataLayoutAnalysis = getAnalysis<DataLayoutAnalysis>();
2131 dataLayoutAnalysis.getAtOrAbove(op));
2132 options.allocLowering =
2135
2136 options.useGenericFunctions = useGenericFunctions;
2137
2138 if (indexBitwidth != kDeriveIndexBitwidthFromDataLayout)
2139 options.overrideIndexBitwidth(indexBitwidth);
2140
2141 LLVMTypeConverter typeConverter(&getContext(), options,
2142 &dataLayoutAnalysis);
2143 RewritePatternSet patterns(&getContext());
2144 SymbolTableCollection symbolTables;
2145 populateFinalizeMemRefToLLVMConversionPatterns(typeConverter, patterns,
2146 &symbolTables);
2148 target.addLegalOp<func::FuncOp>();
2149 if (failed(applyPartialConversion(op, target, std::move(patterns))))
2150 signalPassFailure();
2151 }
2152};
2153
2154/// Implement the interface to convert MemRef to LLVM.
2155struct MemRefToLLVMDialectInterface : public ConvertToLLVMPatternInterface {
2156 MemRefToLLVMDialectInterface(Dialect *dialect)
2157 : ConvertToLLVMPatternInterface(dialect) {}
2158
2159 void loadDependentDialects(MLIRContext *context) const final {
2160 context->loadDialect<LLVM::LLVMDialect>();
2161 }
2162
2163 /// Hook for derived dialect interface to provide conversion patterns
2164 /// and mark dialect legal for the conversion target.
2165 void populateConvertToLLVMConversionPatterns(
2166 ConversionTarget &target, LLVMTypeConverter &typeConverter,
2167 RewritePatternSet &patterns) const final {
2168 populateFinalizeMemRefToLLVMConversionPatterns(typeConverter, patterns);
2169 }
2170};
2171
2172} // namespace
2173
2175 registry.addExtension(+[](MLIRContext *ctx, memref::MemRefDialect *dialect) {
2176 dialect->addInterfaces<MemRefToLLVMDialectInterface>();
2177 });
2178}
return success()
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
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
Definition Block.h:164
BlockArgument getArgument(unsigned i)
Definition Block.h:153
Region * getParent() const
Provide a 'getParent' method for ilist_node_with_parent methods.
Definition Block.cpp:27
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
BlockArgListType getArguments()
Definition Block.h:111
iterator_range< iterator > without_terminator()
Return an iterator range over the operation within this block excluding the terminator operation at t...
Definition Block.h:236
Utility class for operation conversions targeting the LLVM dialect that match exactly one source oper...
Definition Pattern.h:233
ConvertOpToLLVMPattern(const LLVMTypeConverter &typeConverter, PatternBenefit benefit=1)
Definition Pattern.h:239
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'.
Definition IRMapping.h:30
auto lookupOrNull(T from) const
Lookup a mapped value within the map.
Definition IRMapping.h:58
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...
Definition Location.h:76
Options to control the LLVM lowering.
@ 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.
Definition MLIRContext.h:63
Helper class to produce LLVM dialect operations extracting or inserting elements of a MemRef descript...
This class helps build Operations.
Definition Builders.h:210
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
Value getOperand(unsigned idx)
Definition Operation.h:375
result_range getResults()
Definition Operation.h:440
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...
Definition Types.h:74
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...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
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...
Definition Pattern.cpp:620
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...
Definition Pattern.cpp:58
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.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:732
detail::InFlightRemark analysis(Location loc, RemarkOpts opts)
Report an optimization analysis remark.
Definition Remarks.h:738
void promote(RewriterBase &rewriter, scf::ForallOp forallOp)
Promotes the loop body of a scf::ForallOp to its containing block.
Definition SCF.cpp:753
Include the generated interface declarations.
void registerConvertMemRefToLLVMInterface(DialectRegistry &registry)
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)