MLIR 24.0.0git
MemRefOps.cpp
Go to the documentation of this file.
1//===----------------------------------------------------------------------===//
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
15#include "mlir/IR/AffineMap.h"
16#include "mlir/IR/Builders.h"
18#include "mlir/IR/Matchers.h"
26#include "llvm/ADT/STLExtras.h"
27#include "llvm/ADT/SmallBitVector.h"
28#include "llvm/ADT/SmallVectorExtras.h"
29
30using namespace mlir;
31using namespace mlir::memref;
32
33/// Materialize a single constant operation from a given attribute value with
34/// the desired resultant type.
35Operation *MemRefDialect::materializeConstant(OpBuilder &builder,
36 Attribute value, Type type,
37 Location loc) {
38 return arith::ConstantOp::materialize(builder, value, type, loc);
39}
40
41//===----------------------------------------------------------------------===//
42// Common canonicalization pattern support logic
43//===----------------------------------------------------------------------===//
44
45/// This is a common class used for patterns of the form
46/// "someop(memrefcast) -> someop". It folds the source of any memref.cast
47/// into the root operation directly.
48LogicalResult mlir::memref::foldMemRefCast(Operation *op, Value inner) {
49 bool folded = false;
50 for (OpOperand &operand : op->getOpOperands()) {
51 auto cast = operand.get().getDefiningOp<CastOp>();
52 if (cast && operand.get() != inner &&
53 !llvm::isa<UnrankedMemRefType>(cast.getOperand().getType())) {
54 operand.set(cast.getOperand());
55 folded = true;
56 }
57 }
58 return success(folded);
59}
60
61/// Return an unranked/ranked tensor type for the given unranked/ranked memref
62/// type.
64 if (auto memref = llvm::dyn_cast<MemRefType>(type))
65 return RankedTensorType::get(memref.getShape(), memref.getElementType());
66 if (auto memref = llvm::dyn_cast<UnrankedMemRefType>(type))
67 return UnrankedTensorType::get(memref.getElementType());
68 return NoneType::get(type.getContext());
69}
70
72 int64_t dim) {
73 auto memrefType = llvm::cast<MemRefType>(value.getType());
74 if (memrefType.isDynamicDim(dim))
75 return builder.createOrFold<memref::DimOp>(loc, value, dim);
76
77 return builder.getIndexAttr(memrefType.getDimSize(dim));
78}
79
81 Location loc, Value value) {
82 auto memrefType = llvm::cast<MemRefType>(value.getType());
84 for (int64_t i = 0; i < memrefType.getRank(); ++i)
85 result.push_back(getMixedSize(builder, loc, value, i));
86 return result;
87}
88
89//===----------------------------------------------------------------------===//
90// Utility functions for propagating static information
91//===----------------------------------------------------------------------===//
92
93/// Helper function that sets values[i] to constValues[i] if the latter is a
94/// static value, as indicated by ShapedType::kDynamic.
95///
96/// If constValues[i] is dynamic, tries to extract a constant value from
97/// value[i] to allow for additional folding opportunities. Also convertes all
98/// existing attributes to index attributes. (They may be i64 attributes.)
100 ArrayRef<int64_t> constValues) {
101 assert(constValues.size() == values.size() &&
102 "incorrect number of const values");
103 for (auto [i, cstVal] : llvm::enumerate(constValues)) {
104 Builder builder(values[i].getContext());
105 if (ShapedType::isStatic(cstVal)) {
106 // Constant value is known, use it directly.
107 values[i] = builder.getIndexAttr(cstVal);
108 continue;
109 }
110 if (std::optional<int64_t> cst = getConstantIntValue(values[i])) {
111 // Try to extract a constant or convert an existing to index.
112 values[i] = builder.getIndexAttr(*cst);
113 }
114 }
115}
116
117/// Helper function to retrieve a lossless memory-space cast, and the
118/// corresponding new result memref type.
119static std::tuple<MemorySpaceCastOpInterface, PtrLikeTypeInterface, Type>
121 MemorySpaceCastOpInterface castOp =
122 MemorySpaceCastOpInterface::getIfPromotableCast(src);
123
124 // Bail if the cast is not lossless.
125 if (!castOp)
126 return {};
127
128 // Transform the source and target type of `castOp` to have the same metadata
129 // as `resultTy`. Bail if not possible.
130 FailureOr<PtrLikeTypeInterface> srcTy = resultTy.clonePtrWith(
131 castOp.getSourcePtr().getType().getMemorySpace(), std::nullopt);
132 if (failed(srcTy))
133 return {};
134
135 FailureOr<PtrLikeTypeInterface> tgtTy = resultTy.clonePtrWith(
136 castOp.getTargetPtr().getType().getMemorySpace(), std::nullopt);
137 if (failed(tgtTy))
138 return {};
139
140 // Check if this is a valid memory-space cast.
141 if (!castOp.isValidMemorySpaceCast(*tgtTy, *srcTy))
142 return {};
143
144 return std::make_tuple(castOp, *tgtTy, *srcTy);
145}
146
147/// Implementation of `bubbleDownCasts` method for memref operations that
148/// return a single memref result.
149template <typename ConcreteOpTy>
150static FailureOr<std::optional<SmallVector<Value>>>
152 OpOperand &src) {
153 auto [castOp, tgtTy, resTy] = getMemorySpaceCastInfo(op.getType(), src.get());
154 // Bail if we cannot cast.
155 if (!castOp)
156 return failure();
157
158 // Create the new operands.
159 SmallVector<Value> operands;
160 llvm::append_range(operands, op->getOperands());
161 operands[src.getOperandNumber()] = castOp.getSourcePtr();
162
163 // Create the new op and results.
164 auto newOp = ConcreteOpTy::create(
165 builder, op.getLoc(), TypeRange(resTy), operands, op.getProperties(),
166 op->getDiscardableAttrDictionary().getValue());
167
168 // Insert a memory-space cast to the original memory space of the op.
169 MemorySpaceCastOpInterface result = castOp.cloneMemorySpaceCastOp(
170 builder, tgtTy,
171 cast<TypedValue<PtrLikeTypeInterface>>(newOp.getResult()));
172 return std::optional<SmallVector<Value>>(
173 SmallVector<Value>({result.getTargetPtr()}));
174}
175
176//===----------------------------------------------------------------------===//
177// AllocOp / AllocaOp
178//===----------------------------------------------------------------------===//
179
180void AllocOp::getAsmResultNames(
181 function_ref<void(Value, StringRef)> setNameFn) {
182 setNameFn(getResult(), "alloc");
183}
184
185void AllocaOp::getAsmResultNames(
186 function_ref<void(Value, StringRef)> setNameFn) {
187 setNameFn(getResult(), "alloca");
188}
189
190template <typename AllocLikeOp>
191static LogicalResult verifyAllocLikeOp(AllocLikeOp op) {
192 static_assert(llvm::is_one_of<AllocLikeOp, AllocOp, AllocaOp>::value,
193 "applies to only alloc or alloca");
194 auto memRefType = llvm::dyn_cast<MemRefType>(op.getResult().getType());
195 if (!memRefType)
196 return op.emitOpError("result must be a memref");
197
198 if (failed(verifyDynamicDimensionCount(op, memRefType, op.getDynamicSizes())))
199 return failure();
200
201 unsigned numSymbols = 0;
202 if (!memRefType.getLayout().isIdentity())
203 numSymbols = memRefType.getLayout().getAffineMap().getNumSymbols();
204 if (op.getSymbolOperands().size() != numSymbols)
205 return op.emitOpError("symbol operand count does not equal memref symbol "
206 "count: expected ")
207 << numSymbols << ", got " << op.getSymbolOperands().size();
208
209 return success();
210}
211
212LogicalResult AllocOp::verify() { return verifyAllocLikeOp(*this); }
213
214LogicalResult AllocaOp::verify() {
215 // An alloca op needs to have an ancestor with an allocation scope trait.
216 if (!(*this)->getParentWithTrait<OpTrait::AutomaticAllocationScope>())
217 return emitOpError(
218 "requires an ancestor op with AutomaticAllocationScope trait");
219
220 return verifyAllocLikeOp(*this);
221}
222
223namespace {
224/// Fold constant dimensions into an alloc like operation.
225template <typename AllocLikeOp>
226struct SimplifyAllocConst : public OpRewritePattern<AllocLikeOp> {
227 using OpRewritePattern<AllocLikeOp>::OpRewritePattern;
228
229 LogicalResult matchAndRewrite(AllocLikeOp alloc,
230 PatternRewriter &rewriter) const override {
231 // Check to see if any dimensions operands are constants. If so, we can
232 // substitute and drop them.
233 if (llvm::none_of(alloc.getDynamicSizes(), [](Value operand) {
234 APInt constSizeArg;
235 if (!matchPattern(operand, m_ConstantInt(&constSizeArg)))
236 return false;
237 return constSizeArg.isNonNegative();
238 }))
239 return failure();
240
241 auto memrefType = alloc.getType();
242
243 // Ok, we have one or more constant operands. Collect the non-constant ones
244 // and keep track of the resultant memref type to build.
245 SmallVector<int64_t, 4> newShapeConstants;
246 newShapeConstants.reserve(memrefType.getRank());
247 SmallVector<Value, 4> dynamicSizes;
248
249 unsigned dynamicDimPos = 0;
250 for (unsigned dim = 0, e = memrefType.getRank(); dim < e; ++dim) {
251 int64_t dimSize = memrefType.getDimSize(dim);
252 // If this is already static dimension, keep it.
253 if (ShapedType::isStatic(dimSize)) {
254 newShapeConstants.push_back(dimSize);
255 continue;
256 }
257 auto dynamicSize = alloc.getDynamicSizes()[dynamicDimPos];
258 APInt constSizeArg;
259 if (matchPattern(dynamicSize, m_ConstantInt(&constSizeArg)) &&
260 constSizeArg.isNonNegative()) {
261 // Dynamic shape dimension will be folded.
262 newShapeConstants.push_back(constSizeArg.getZExtValue());
263 } else {
264 // Dynamic shape dimension not folded; copy dynamicSize from old memref.
265 newShapeConstants.push_back(ShapedType::kDynamic);
266 dynamicSizes.push_back(dynamicSize);
267 }
268 dynamicDimPos++;
269 }
270
271 // Create new memref type (which will have fewer dynamic dimensions).
272 MemRefType newMemRefType =
273 MemRefType::Builder(memrefType).setShape(newShapeConstants);
274 assert(dynamicSizes.size() == newMemRefType.getNumDynamicDims());
275
276 // Create and insert the alloc op for the new memref.
277 auto newAlloc = AllocLikeOp::create(rewriter, alloc.getLoc(), newMemRefType,
278 dynamicSizes, alloc.getSymbolOperands(),
279 alloc.getAlignmentAttr());
280 // Insert a cast so we have the same type as the old alloc.
281 rewriter.replaceOpWithNewOp<CastOp>(alloc, alloc.getType(), newAlloc);
282 return success();
283 }
284};
285
286/// Fold alloc operations with no users or only store and dealloc uses.
287template <typename T>
288struct SimplifyDeadAlloc : public OpRewritePattern<T> {
289 using OpRewritePattern<T>::OpRewritePattern;
290
291 LogicalResult matchAndRewrite(T alloc,
292 PatternRewriter &rewriter) const override {
293 if (llvm::any_of(alloc->getUsers(), [&](Operation *op) {
294 if (auto storeOp = dyn_cast<StoreOp>(op))
295 return storeOp.getValue() == alloc;
296 return !isa<DeallocOp>(op);
297 }))
298 return failure();
299
300 for (Operation *user : llvm::make_early_inc_range(alloc->getUsers()))
301 rewriter.eraseOp(user);
302
303 rewriter.eraseOp(alloc);
304 return success();
305 }
306};
307} // namespace
308
309void AllocOp::getCanonicalizationPatterns(RewritePatternSet &results,
310 MLIRContext *context) {
311 results.add<SimplifyAllocConst<AllocOp>, SimplifyDeadAlloc<AllocOp>>(context);
312}
313
314void AllocaOp::getCanonicalizationPatterns(RewritePatternSet &results,
315 MLIRContext *context) {
316 results.add<SimplifyAllocConst<AllocaOp>, SimplifyDeadAlloc<AllocaOp>>(
317 context);
318}
319
320//===----------------------------------------------------------------------===//
321// ReallocOp
322//===----------------------------------------------------------------------===//
323
324LogicalResult ReallocOp::verify() {
325 auto sourceType = llvm::cast<MemRefType>(getOperand(0).getType());
326 MemRefType resultType = getType();
327
328 // The source memref should have identity layout (or none).
329 if (!sourceType.getLayout().isIdentity())
330 return emitError("unsupported layout for source memref type ")
331 << sourceType;
332
333 // The result memref should have identity layout (or none).
334 if (!resultType.getLayout().isIdentity())
335 return emitError("unsupported layout for result memref type ")
336 << resultType;
337
338 // The source memref and the result memref should be in the same memory space.
339 if (sourceType.getMemorySpace() != resultType.getMemorySpace())
340 return emitError("different memory spaces specified for source memref "
341 "type ")
342 << sourceType << " and result memref type " << resultType;
343
344 // The source memref and the result memref should have the same element type.
345 if (failed(verifyElementTypesMatch(*this, sourceType, resultType, "source",
346 "result")))
347 return failure();
348
349 // Verify that we have the dynamic dimension operand when it is needed.
350 if (resultType.getNumDynamicDims() && !getDynamicResultSize())
351 return emitError("missing dimension operand for result type ")
352 << resultType;
353 if (!resultType.getNumDynamicDims() && getDynamicResultSize())
354 return emitError("unnecessary dimension operand for result type ")
355 << resultType;
356
357 return success();
358}
359
360void ReallocOp::getCanonicalizationPatterns(RewritePatternSet &results,
361 MLIRContext *context) {
362 results.add<SimplifyDeadAlloc<ReallocOp>>(context);
363}
364
365//===----------------------------------------------------------------------===//
366// AllocaScopeOp
367//===----------------------------------------------------------------------===//
368
369void AllocaScopeOp::print(OpAsmPrinter &p) {
370 bool printBlockTerminators = false;
371
372 p << ' ';
373 if (!getResults().empty()) {
374 p << " -> (" << getResultTypes() << ")";
375 printBlockTerminators = true;
376 }
377 p << ' ';
378 p.printRegion(getBodyRegion(),
379 /*printEntryBlockArgs=*/false,
380 /*printBlockTerminators=*/printBlockTerminators);
381 p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary());
382}
383
384ParseResult AllocaScopeOp::parse(OpAsmParser &parser, OperationState &result) {
385 // Create a region for the body.
386 result.regions.reserve(1);
387 Region *bodyRegion = result.addRegion();
388
389 // Parse optional results type list.
390 if (parser.parseOptionalArrowTypeList(result.types))
391 return failure();
392
393 // Parse the body region.
394 if (parser.parseRegion(*bodyRegion, /*arguments=*/{}))
395 return failure();
396 AllocaScopeOp::ensureTerminator(*bodyRegion, parser.getBuilder(),
397 result.location);
398
399 // Parse the optional attribute list.
400 if (parser.parseOptionalAttrDict(result.attributes))
401 return failure();
402
403 return success();
404}
405
406void AllocaScopeOp::getSuccessorRegions(
408 if (!point.isParent()) {
409 regions.push_back(RegionSuccessor(getOperation()));
410 return;
411 }
412
413 regions.push_back(RegionSuccessor(&getBodyRegion()));
414}
415
416ValueRange AllocaScopeOp::getSuccessorInputs(RegionSuccessor successor) {
417 return successor.isOperation() ? ValueRange(getResults()) : ValueRange();
418}
419
420/// Given an operation, return whether this op is guaranteed to
421/// allocate an AutomaticAllocationScopeResource
423 MemoryEffectOpInterface interface = dyn_cast<MemoryEffectOpInterface>(op);
424 if (!interface)
425 return false;
426 for (auto res : op->getResults()) {
427 if (auto effect =
428 interface.getEffectOnValue<MemoryEffects::Allocate>(res)) {
429 if (isa<SideEffects::AutomaticAllocationScopeResource>(
430 effect->getResource()))
431 return true;
432 }
433 }
434 return false;
435}
436
437/// Given an operation, return whether this op itself could
438/// allocate an AutomaticAllocationScopeResource. Note that
439/// this will not check whether an operation contained within
440/// the op can allocate.
442 // This op itself doesn't create a stack allocation,
443 // the inner allocation should be handled separately.
445 return false;
446 MemoryEffectOpInterface interface = dyn_cast<MemoryEffectOpInterface>(op);
447 if (!interface)
448 return true;
449 for (auto res : op->getResults()) {
450 if (auto effect =
451 interface.getEffectOnValue<MemoryEffects::Allocate>(res)) {
452 if (isa<SideEffects::AutomaticAllocationScopeResource>(
453 effect->getResource()))
454 return true;
455 }
456 }
457 return false;
458}
459
460/// Return whether this op is the last non terminating op
461/// in a region. That is to say, it is in a one-block region
462/// and is only followed by a terminator. This prevents
463/// extending the lifetime of allocations.
465 return op->getBlock()->mightHaveTerminator() &&
466 op->getNextNode() == op->getBlock()->getTerminator() &&
468}
469
470/// Inline an AllocaScopeOp if either the direct parent is an allocation scope
471/// or it contains no allocation.
472struct AllocaScopeInliner : public OpRewritePattern<AllocaScopeOp> {
473 using OpRewritePattern<AllocaScopeOp>::OpRewritePattern;
474
475 LogicalResult matchAndRewrite(AllocaScopeOp op,
476 PatternRewriter &rewriter) const override {
477 bool hasPotentialAlloca =
478 op->walk<WalkOrder::PreOrder>([&](Operation *alloc) {
479 if (alloc == op)
480 return WalkResult::advance();
482 return WalkResult::interrupt();
483 if (alloc->hasTrait<OpTrait::AutomaticAllocationScope>())
484 return WalkResult::skip();
485 return WalkResult::advance();
486 }).wasInterrupted();
487
488 // If this contains no potential allocation, it is always legal to
489 // inline. Otherwise, consider two conditions:
490 if (hasPotentialAlloca) {
491 // If the parent isn't an allocation scope, or we are not the last
492 // non-terminator op in the parent, we will extend the lifetime.
493 if (!op->getParentOp()->hasTrait<OpTrait::AutomaticAllocationScope>())
494 return failure();
496 return failure();
497 }
498
499 Block *block = &op.getRegion().front();
500 Operation *terminator = block->getTerminator();
501 ValueRange results = terminator->getOperands();
502 rewriter.inlineBlockBefore(block, op);
503 rewriter.replaceOp(op, results);
504 rewriter.eraseOp(terminator);
505 return success();
506 }
507};
508
509/// Move allocations into an allocation scope, if it is legal to
510/// move them (e.g. their operands are available at the location
511/// the op would be moved to).
512struct AllocaScopeHoister : public OpRewritePattern<AllocaScopeOp> {
513 using OpRewritePattern<AllocaScopeOp>::OpRewritePattern;
514
515 LogicalResult matchAndRewrite(AllocaScopeOp op,
516 PatternRewriter &rewriter) const override {
517
518 if (!op->getParentWithTrait<OpTrait::AutomaticAllocationScope>())
519 return failure();
520
521 Operation *lastParentWithoutScope = op->getParentOp();
522
523 if (!lastParentWithoutScope ||
524 lastParentWithoutScope->hasTrait<OpTrait::AutomaticAllocationScope>())
525 return failure();
526
527 // Only apply to if this is this last non-terminator
528 // op in the block (lest lifetime be extended) of a one
529 // block region
530 if (!lastNonTerminatorInRegion(op) ||
531 !lastNonTerminatorInRegion(lastParentWithoutScope))
532 return failure();
533
534 while (!lastParentWithoutScope->getParentOp()
536 lastParentWithoutScope = lastParentWithoutScope->getParentOp();
537 if (!lastParentWithoutScope ||
538 !lastNonTerminatorInRegion(lastParentWithoutScope))
539 return failure();
540 }
541 assert(lastParentWithoutScope->getParentOp()
543
544 Region *containingRegion = nullptr;
545 for (auto &r : lastParentWithoutScope->getRegions()) {
546 if (r.isAncestor(op->getParentRegion())) {
547 assert(containingRegion == nullptr &&
548 "only one region can contain the op");
549 containingRegion = &r;
550 }
551 }
552 assert(containingRegion && "op must be contained in a region");
553
555 op->walk([&](Operation *alloc) {
557 return WalkResult::skip();
558
559 // If any operand is not defined before the location of
560 // lastParentWithoutScope (i.e. where we would hoist to), skip.
561 if (llvm::any_of(alloc->getOperands(), [&](Value v) {
562 return containingRegion->isAncestor(v.getParentRegion());
563 }))
564 return WalkResult::skip();
565 toHoist.push_back(alloc);
566 return WalkResult::advance();
567 });
568
569 if (toHoist.empty())
570 return failure();
571 rewriter.setInsertionPoint(lastParentWithoutScope);
572 for (auto *op : toHoist) {
573 auto *cloned = rewriter.clone(*op);
574 rewriter.replaceOp(op, cloned->getResults());
575 }
576 return success();
577 }
578};
579
580void AllocaScopeOp::getCanonicalizationPatterns(RewritePatternSet &results,
581 MLIRContext *context) {
582 results.add<AllocaScopeInliner, AllocaScopeHoister>(context);
583}
584
585//===----------------------------------------------------------------------===//
586// AssumeAlignmentOp
587//===----------------------------------------------------------------------===//
588
589LogicalResult AssumeAlignmentOp::verify() {
590 if (!llvm::isPowerOf2_32(getAlignment()))
591 return emitOpError("alignment must be power of 2");
592 return success();
593}
594
595void AssumeAlignmentOp::getAsmResultNames(
596 function_ref<void(Value, StringRef)> setNameFn) {
597 setNameFn(getResult(), "assume_align");
598}
599
600OpFoldResult AssumeAlignmentOp::fold(FoldAdaptor adaptor) {
601 auto source = getMemref().getDefiningOp<AssumeAlignmentOp>();
602 if (!source)
603 return {};
604 if (source.getAlignment() != getAlignment())
605 return {};
606 return getMemref();
607}
608
609FailureOr<std::optional<SmallVector<Value>>>
610AssumeAlignmentOp::bubbleDownCasts(OpBuilder &builder) {
611 return bubbleDownCastsPassthroughOpImpl(*this, builder, getMemrefMutable());
612}
613
614FailureOr<OpFoldResult> AssumeAlignmentOp::reifyDimOfResult(OpBuilder &builder,
615 int resultIndex,
616 int dim) {
617 assert(resultIndex == 0 && "AssumeAlignmentOp has a single result");
618 return getMixedSize(builder, getLoc(), getMemref(), dim);
619}
620
621//===----------------------------------------------------------------------===//
622// DistinctObjectsOp
623//===----------------------------------------------------------------------===//
624
625LogicalResult DistinctObjectsOp::verify() {
626 if (getOperandTypes() != getResultTypes())
627 return emitOpError("operand types and result types must match");
628
629 if (getOperandTypes().empty())
630 return emitOpError("expected at least one operand");
631
632 return success();
633}
634
635LogicalResult DistinctObjectsOp::inferReturnTypes(
636 MLIRContext * /*context*/, std::optional<Location> /*location*/,
637 ValueRange operands, DictionaryAttr /*attributes*/,
638 PropertyRef /*properties*/, RegionRange /*regions*/,
639 SmallVectorImpl<Type> &inferredReturnTypes) {
640 llvm::copy(operands.getTypes(), std::back_inserter(inferredReturnTypes));
641 return success();
642}
643
644//===----------------------------------------------------------------------===//
645// CastOp
646//===----------------------------------------------------------------------===//
647
648void CastOp::getAsmResultNames(function_ref<void(Value, StringRef)> setNameFn) {
649 setNameFn(getResult(), "cast");
650}
651
652/// Determines whether MemRef_CastOp casts to a more dynamic version of the
653/// source memref. This is useful to fold a memref.cast into a consuming op
654/// and implement canonicalization patterns for ops in different dialects that
655/// may consume the results of memref.cast operations. Such foldable memref.cast
656/// operations are typically inserted as `view` and `subview` ops are
657/// canonicalized, to preserve the type compatibility of their uses.
658///
659/// Returns true when all conditions are met:
660/// 1. source and result are ranked memrefs with strided semantics and same
661/// element type and rank.
662/// 2. each of the source's size, offset or stride has more static information
663/// than the corresponding result's size, offset or stride.
664///
665/// Example 1:
666/// ```mlir
667/// %1 = memref.cast %0 : memref<8x16xf32> to memref<?x?xf32>
668/// %2 = consumer %1 ... : memref<?x?xf32> ...
669/// ```
670///
671/// may fold into:
672///
673/// ```mlir
674/// %2 = consumer %0 ... : memref<8x16xf32> ...
675/// ```
676///
677/// Example 2:
678/// ```
679/// %1 = memref.cast %0 : memref<?x16xf32, affine_map<(i, j)->(16 * i + j)>>
680/// to memref<?x?xf32>
681/// consumer %1 : memref<?x?xf32> ...
682/// ```
683///
684/// may fold into:
685///
686/// ```
687/// consumer %0 ... : memref<?x16xf32, affine_map<(i, j)->(16 * i + j)>>
688/// ```
689bool CastOp::canFoldIntoConsumerOp(CastOp castOp) {
690 MemRefType sourceType =
691 llvm::dyn_cast<MemRefType>(castOp.getSource().getType());
692 MemRefType resultType = llvm::dyn_cast<MemRefType>(castOp.getType());
693
694 // Requires ranked MemRefType.
695 if (!sourceType || !resultType)
696 return false;
697
698 // Requires same elemental type.
699 if (sourceType.getElementType() != resultType.getElementType())
700 return false;
701
702 // Requires same rank.
703 if (sourceType.getRank() != resultType.getRank())
704 return false;
705
706 // Only fold casts between strided memref forms.
707 int64_t sourceOffset, resultOffset;
708 SmallVector<int64_t, 4> sourceStrides, resultStrides;
709 if (failed(sourceType.getStridesAndOffset(sourceStrides, sourceOffset)) ||
710 failed(resultType.getStridesAndOffset(resultStrides, resultOffset)))
711 return false;
712
713 // If cast is towards more static sizes along any dimension, don't fold.
714 for (auto it : llvm::zip(sourceType.getShape(), resultType.getShape())) {
715 auto ss = std::get<0>(it), st = std::get<1>(it);
716 if (ss != st)
717 if (ShapedType::isDynamic(ss) && ShapedType::isStatic(st))
718 return false;
719 }
720
721 // If cast is towards more static offset along any dimension, don't fold.
722 if (sourceOffset != resultOffset)
723 if (ShapedType::isDynamic(sourceOffset) &&
724 ShapedType::isStatic(resultOffset))
725 return false;
726
727 // If cast is towards more static strides along any dimension, don't fold.
728 for (auto it : llvm::zip(sourceStrides, resultStrides)) {
729 auto ss = std::get<0>(it), st = std::get<1>(it);
730 if (ss != st)
731 if (ShapedType::isDynamic(ss) && ShapedType::isStatic(st))
732 return false;
733 }
734
735 return true;
736}
737
738bool CastOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
739 if (inputs.size() != 1 || outputs.size() != 1)
740 return false;
741 if (inputs == outputs)
742 return true;
743 Type a = inputs.front(), b = outputs.front();
744 auto aT = llvm::dyn_cast<MemRefType>(a);
745 auto bT = llvm::dyn_cast<MemRefType>(b);
746
747 auto uaT = llvm::dyn_cast<UnrankedMemRefType>(a);
748 auto ubT = llvm::dyn_cast<UnrankedMemRefType>(b);
749
750 if (aT && bT) {
751 if (aT.getElementType() != bT.getElementType())
752 return false;
753 if (aT.getLayout() != bT.getLayout()) {
754 int64_t aOffset, bOffset;
755 SmallVector<int64_t, 4> aStrides, bStrides;
756 if (failed(aT.getStridesAndOffset(aStrides, aOffset)) ||
757 failed(bT.getStridesAndOffset(bStrides, bOffset)) ||
758 aStrides.size() != bStrides.size())
759 return false;
760
761 // Strides along a dimension/offset are compatible if the value in the
762 // source memref is static and the value in the target memref is the
763 // same. They are also compatible if either one is dynamic (see
764 // description of MemRefCastOp for details).
765 // Note that for dimensions of size 1, the stride can differ.
766 auto checkCompatible = [](int64_t a, int64_t b) {
767 return (ShapedType::isDynamic(a) || ShapedType::isDynamic(b) || a == b);
768 };
769 if (!checkCompatible(aOffset, bOffset))
770 return false;
771 for (const auto &[index, aStride] : enumerate(aStrides)) {
772 if (aT.getDimSize(index) == 1 || bT.getDimSize(index) == 1)
773 continue;
774 if (!checkCompatible(aStride, bStrides[index]))
775 return false;
776 }
777 }
778 if (aT.getMemorySpace() != bT.getMemorySpace())
779 return false;
780
781 // They must have the same rank, and any specified dimensions must match.
782 if (aT.getRank() != bT.getRank())
783 return false;
784
785 for (unsigned i = 0, e = aT.getRank(); i != e; ++i) {
786 int64_t aDim = aT.getDimSize(i), bDim = bT.getDimSize(i);
787 if (ShapedType::isStatic(aDim) && ShapedType::isStatic(bDim) &&
788 aDim != bDim)
789 return false;
790 }
791 return true;
792 } else {
793 if (!aT && !uaT)
794 return false;
795 if (!bT && !ubT)
796 return false;
797 // Unranked to unranked casting is unsupported
798 if (uaT && ubT)
799 return false;
800
801 auto aEltType = (aT) ? aT.getElementType() : uaT.getElementType();
802 auto bEltType = (bT) ? bT.getElementType() : ubT.getElementType();
803 if (aEltType != bEltType)
804 return false;
805
806 auto aMemSpace = (aT) ? aT.getMemorySpace() : uaT.getMemorySpace();
807 auto bMemSpace = (bT) ? bT.getMemorySpace() : ubT.getMemorySpace();
808 return aMemSpace == bMemSpace;
809 }
810
811 return false;
812}
813
814OpFoldResult CastOp::fold(FoldAdaptor adaptor) {
815 return succeeded(foldMemRefCast(*this)) ? getResult() : Value();
816}
817
818FailureOr<std::optional<SmallVector<Value>>>
819CastOp::bubbleDownCasts(OpBuilder &builder) {
820 return bubbleDownCastsPassthroughOpImpl(*this, builder, getSourceMutable());
821}
822
823//===----------------------------------------------------------------------===//
824// CopyOp
825//===----------------------------------------------------------------------===//
826
827namespace {
828
829/// Fold memref.copy(%x, %x).
830struct FoldSelfCopy : public OpRewritePattern<CopyOp> {
831 using OpRewritePattern<CopyOp>::OpRewritePattern;
832
833 LogicalResult matchAndRewrite(CopyOp copyOp,
834 PatternRewriter &rewriter) const override {
835 if (copyOp.getSource() != copyOp.getTarget())
836 return failure();
837
838 rewriter.eraseOp(copyOp);
839 return success();
840 }
841};
842
843struct FoldEmptyCopy final : public OpRewritePattern<CopyOp> {
844 using OpRewritePattern<CopyOp>::OpRewritePattern;
845
846 static bool isEmptyMemRef(BaseMemRefType type) {
847 return type.hasRank() && llvm::is_contained(type.getShape(), 0);
848 }
849
850 LogicalResult matchAndRewrite(CopyOp copyOp,
851 PatternRewriter &rewriter) const override {
852 if (isEmptyMemRef(copyOp.getSource().getType()) ||
853 isEmptyMemRef(copyOp.getTarget().getType())) {
854 rewriter.eraseOp(copyOp);
855 return success();
856 }
857
858 return failure();
859 }
860};
861} // namespace
862
863void CopyOp::getCanonicalizationPatterns(RewritePatternSet &results,
864 MLIRContext *context) {
865 results.add<FoldEmptyCopy, FoldSelfCopy>(context);
866}
867
868/// If the source/target of a CopyOp is a CastOp that does not modify the shape
869/// and element type, the cast can be skipped. Such CastOps only cast the layout
870/// of the type.
871static LogicalResult foldCopyOfCast(CopyOp op) {
872 for (OpOperand &operand : op->getOpOperands()) {
873 auto castOp = operand.get().getDefiningOp<memref::CastOp>();
874 if (castOp && memref::CastOp::canFoldIntoConsumerOp(castOp)) {
875 operand.set(castOp.getOperand());
876 return success();
877 }
878 }
879 return failure();
880}
881
882LogicalResult CopyOp::fold(FoldAdaptor adaptor,
883 SmallVectorImpl<OpFoldResult> &results) {
884
885 /// copy(memrefcast) -> copy
886 return foldCopyOfCast(*this);
887}
888
889//===----------------------------------------------------------------------===//
890// DeallocOp
891//===----------------------------------------------------------------------===//
892
893LogicalResult DeallocOp::fold(FoldAdaptor adaptor,
894 SmallVectorImpl<OpFoldResult> &results) {
895 /// dealloc(memrefcast) -> dealloc
896 return foldMemRefCast(*this);
897}
898
899//===----------------------------------------------------------------------===//
900// DimOp
901//===----------------------------------------------------------------------===//
902
903void DimOp::getAsmResultNames(function_ref<void(Value, StringRef)> setNameFn) {
904 setNameFn(getResult(), "dim");
905}
906
907void DimOp::build(OpBuilder &builder, OperationState &result, Value source,
908 int64_t index) {
909 auto loc = result.location;
910 Value indexValue = arith::ConstantIndexOp::create(builder, loc, index);
911 build(builder, result, source, indexValue);
912}
913
914std::optional<int64_t> DimOp::getConstantIndex() {
916}
917
918Speculation::Speculatability DimOp::getSpeculatability() {
919 auto constantIndex = getConstantIndex();
920 if (!constantIndex)
922
923 auto rankedSourceType = dyn_cast<MemRefType>(getSource().getType());
924 if (!rankedSourceType)
926
927 if (rankedSourceType.getRank() <= constantIndex)
929
931}
932
933void DimOp::inferResultRangesFromOptional(ArrayRef<IntegerValueRange> argRanges,
934 SetIntLatticeFn setResultRange) {
935 setResultRange(getResult(),
936 intrange::inferShapedDimOpInterface(*this, argRanges[1]));
937}
938
939/// Return a map with key being elements in `vals` and data being number of
940/// occurences of it. Use std::map, since the `vals` here are strides and the
941/// dynamic stride value is the same as the tombstone value for
942/// `DenseMap<int64_t>`.
943static std::map<int64_t, unsigned> getNumOccurences(ArrayRef<int64_t> vals) {
944 std::map<int64_t, unsigned> numOccurences;
945 for (auto val : vals)
946 numOccurences[val]++;
947 return numOccurences;
948}
949
950/// Returns the set of source dimensions that are dropped in a rank reduction.
951/// For each result dimension in order, matches the leftmost unmatched source
952/// dimension with the same size. Source dimensions not matched are dropped.
953///
954/// Example: memref<1x8x1x3> to memref<1x8x3>. Source sizes [1, 8, 1, 3], result
955/// [1, 8, 3]. Match result[0]=1 -> source dim 0, result[1]=8 -> source dim 1,
956/// result[2]=3 -> source dim 3. Source dim 2 is unmatched and dropped.
957static FailureOr<llvm::SmallBitVector>
959 MemRefType reducedType,
961 int64_t rankReduction = originalType.getRank() - reducedType.getRank();
962 if (rankReduction <= 0)
963 return llvm::SmallBitVector(originalType.getRank());
964
965 // Build source sizes from subview sizes (one per source dim).
966 SmallVector<int64_t> sourceSizes(originalType.getRank());
967 for (const auto &it : llvm::enumerate(sizes)) {
968 if (std::optional<int64_t> cst = getConstantIntValue(it.value()))
969 sourceSizes[it.index()] = *cst;
970 else
971 sourceSizes[it.index()] = ShapedType::kDynamic;
972 }
973
974 ArrayRef<int64_t> resultSizes = reducedType.getShape();
975 llvm::SmallBitVector usedSourceDims(originalType.getRank());
976 int64_t startJ = 0;
977 for (int64_t resultSize : resultSizes) {
978 bool matched = false;
979 for (int64_t j = startJ; j < originalType.getRank(); ++j) {
980 if (sourceSizes[j] == resultSize) {
981 usedSourceDims.set(j);
982 matched = true;
983 startJ = j + 1;
984 break;
985 }
986 }
987 if (!matched)
988 return failure();
989 }
990
991 llvm::SmallBitVector unusedDims(originalType.getRank());
992 for (int64_t i = 0; i < originalType.getRank(); ++i)
993 if (!usedSourceDims.test(i))
994 unusedDims.set(i);
995 return unusedDims;
996}
997
998/// Returns the set of source dimensions that are dropped in a rank reduction.
999/// A dimension is dropped if its stride is dropped; uses stride occurrence
1000/// counting to disambiguate when multiple unit dims exist.
1001///
1002/// Example: memref<1x1x?xf32, strided<[?, 4, 1]>> to memref<1x4xf32,
1003/// strided<[4, 1]>>. Source strides [?, 4, 1], candidate [4, 1]. Dim 0 (stride
1004/// ?) can be dropped; dim 1 (stride 4) must be kept. Source dim 0 is dropped.
1005static FailureOr<llvm::SmallBitVector> computeMemRefRankReductionMaskByStrides(
1006 MemRefType originalType, MemRefType reducedType,
1007 ArrayRef<int64_t> originalStrides, ArrayRef<int64_t> candidateStrides,
1008 llvm::SmallBitVector unusedDims) {
1009 // Track the number of occurences of the strides in the original type
1010 // and the candidate type. For each unused dim that stride should not be
1011 // present in the candidate type. Note that there could be multiple dimensions
1012 // that have the same size. We dont need to exactly figure out which dim
1013 // corresponds to which stride, we just need to verify that the number of
1014 // reptitions of a stride in the original + number of unused dims with that
1015 // stride == number of repititions of a stride in the candidate.
1016 std::map<int64_t, unsigned> currUnaccountedStrides =
1017 getNumOccurences(originalStrides);
1018 std::map<int64_t, unsigned> candidateStridesNumOccurences =
1019 getNumOccurences(candidateStrides);
1020 for (size_t dim = 0, e = unusedDims.size(); dim != e; ++dim) {
1021 if (!unusedDims.test(dim))
1022 continue;
1023 int64_t originalStride = originalStrides[dim];
1024 if (currUnaccountedStrides[originalStride] >
1025 candidateStridesNumOccurences[originalStride]) {
1026 // This dim can be treated as dropped.
1027 currUnaccountedStrides[originalStride]--;
1028 continue;
1029 }
1030 if (currUnaccountedStrides[originalStride] ==
1031 candidateStridesNumOccurences[originalStride]) {
1032 // The stride for this is not dropped. Keep as is.
1033 unusedDims.reset(dim);
1034 continue;
1035 }
1036 if (currUnaccountedStrides[originalStride] <
1037 candidateStridesNumOccurences[originalStride]) {
1038 // This should never happen. Cant have a stride in the reduced rank type
1039 // that wasnt in the original one.
1040 return failure();
1041 }
1042 }
1043 if (static_cast<int64_t>(unusedDims.count()) + reducedType.getRank() !=
1044 originalType.getRank())
1045 return failure();
1046 return unusedDims;
1047}
1048
1049/// Given the `originalType` and a `candidateReducedType` whose shape is assumed
1050/// to be a subset of `originalType` with some `1` entries erased, return the
1051/// set of indices that specifies which of the entries of `originalShape` are
1052/// dropped to obtain `reducedShape`.
1053/// This accounts for cases where there are multiple unit-dims, but only a
1054/// subset of those are dropped. For MemRefTypes these can be disambiguated
1055/// using the strides. If a dimension is dropped the stride must be dropped too.
1056static FailureOr<llvm::SmallBitVector>
1057computeMemRefRankReductionMask(MemRefType originalType, MemRefType reducedType,
1058 ArrayRef<OpFoldResult> sizes) {
1059 llvm::SmallBitVector unusedDims(originalType.getRank());
1060 if (originalType.getRank() == reducedType.getRank())
1061 return unusedDims;
1062
1063 for (const auto &dim : llvm::enumerate(sizes))
1064 if (auto attr = llvm::dyn_cast_if_present<Attribute>(dim.value()))
1065 if (llvm::cast<IntegerAttr>(attr).getInt() == 1)
1066 unusedDims.set(dim.index());
1067
1068 // Early exit for the case where the number of unused dims matches the number
1069 // of ranks reduced.
1070 if (static_cast<int64_t>(unusedDims.count()) + reducedType.getRank() ==
1071 originalType.getRank())
1072 return unusedDims;
1073
1074 SmallVector<int64_t> originalStrides, candidateStrides;
1075 int64_t originalOffset, candidateOffset;
1076 if (failed(
1077 originalType.getStridesAndOffset(originalStrides, originalOffset)) ||
1078 failed(
1079 reducedType.getStridesAndOffset(candidateStrides, candidateOffset)))
1080 return failure();
1081
1082 // Try stride-based first when we have meaningful static stride info
1083 // (preserves static strides). Fall back to position-based otherwise.
1084 auto hasNonTrivialStaticStride = [](ArrayRef<int64_t> strides) {
1085 // The innermost stride 1 is trivial for row-major and does not help
1086 // disambiguate.
1087 if (strides.size() <= 1)
1088 return false;
1089 return llvm::any_of(strides.drop_back(),
1090 [](int64_t s) { return !ShapedType::isDynamic(s); });
1091 };
1092 if (hasNonTrivialStaticStride(originalStrides) ||
1093 hasNonTrivialStaticStride(candidateStrides)) {
1094 FailureOr<llvm::SmallBitVector> strideBased =
1095 computeMemRefRankReductionMaskByStrides(originalType, reducedType,
1096 originalStrides,
1097 candidateStrides, unusedDims);
1098 if (succeeded(strideBased))
1099 return *strideBased;
1100 }
1101 return computeMemRefRankReductionMaskByPosition(originalType, reducedType,
1102 sizes);
1103}
1104
1105llvm::SmallBitVector SubViewOp::getDroppedDims() {
1106 MemRefType sourceType = getSourceType();
1107 MemRefType resultType = getType();
1108 FailureOr<llvm::SmallBitVector> unusedDims =
1109 computeMemRefRankReductionMask(sourceType, resultType, getMixedSizes());
1110 assert(succeeded(unusedDims) && "unable to find unused dims of subview");
1111 return *unusedDims;
1112}
1113
1114OpFoldResult DimOp::fold(FoldAdaptor adaptor) {
1115 // All forms of folding require a known index.
1116 std::optional<int64_t> index = getConstantIndex();
1117 if (!index)
1118 return {};
1119
1120 // Folding for unranked types (UnrankedMemRefType) is not supported.
1121 auto memrefType = llvm::dyn_cast<MemRefType>(getSource().getType());
1122 if (!memrefType)
1123 return {};
1124
1125 // Out of bound indices produce undefined behavior but are still valid IR.
1126 // Don't choke on them.
1127 int64_t indexVal = index.value();
1128 if (indexVal < 0 || indexVal >= memrefType.getRank())
1129 return {};
1130
1131 // Fold if the shape extent along the given index is known.
1132 if (!memrefType.isDynamicDim(indexVal)) {
1133 Builder builder(getContext());
1134 return builder.getIndexAttr(memrefType.getShape()[indexVal]);
1135 }
1136
1137 // The size at the given index is now known to be a dynamic size.
1138 // Fold dim to the size argument for an `AllocOp`, `ViewOp`, or `SubViewOp`.
1139 Operation *definingOp = getSource().getDefiningOp();
1140
1141 if (auto alloc = dyn_cast_or_null<AllocOp>(definingOp))
1142 return *(alloc.getDynamicSizes().begin() +
1143 memrefType.getDynamicDimIndex(indexVal));
1144
1145 if (auto alloca = dyn_cast_or_null<AllocaOp>(definingOp))
1146 return *(alloca.getDynamicSizes().begin() +
1147 memrefType.getDynamicDimIndex(indexVal));
1148
1149 if (auto view = dyn_cast_or_null<ViewOp>(definingOp))
1150 return *(view.getDynamicSizes().begin() +
1151 memrefType.getDynamicDimIndex(indexVal));
1152
1153 if (auto subview = dyn_cast_or_null<SubViewOp>(definingOp)) {
1154 // The result dim is dynamic (the static case was handled above). Dropped
1155 // dims always have static size 1, so dynamic source sizes are never
1156 // dropped and map in order to the dynamic result dims. Find the k-th
1157 // dynamic source size, where k is the dynamic dim index of the result dim.
1158 unsigned dynamicResultDimIdx = memrefType.getDynamicDimIndex(indexVal);
1159 unsigned dynamicIdx = 0;
1160 for (OpFoldResult size : subview.getMixedSizes()) {
1161 if (llvm::isa<Attribute>(size))
1162 continue;
1163 if (dynamicIdx == dynamicResultDimIdx)
1164 return size;
1165 dynamicIdx++;
1166 }
1167 return {};
1168 }
1169
1170 // dim(memrefcast) -> dim
1171 if (succeeded(foldMemRefCast(*this)))
1172 return getResult();
1173
1174 return {};
1175}
1176
1177namespace {
1178/// Fold dim of a memref reshape operation to a load into the reshape's shape
1179/// operand.
1180struct DimOfMemRefReshape : public OpRewritePattern<DimOp> {
1181 using OpRewritePattern<DimOp>::OpRewritePattern;
1182
1183 LogicalResult matchAndRewrite(DimOp dim,
1184 PatternRewriter &rewriter) const override {
1185 auto reshape = dim.getSource().getDefiningOp<ReshapeOp>();
1186
1187 if (!reshape)
1188 return rewriter.notifyMatchFailure(
1189 dim, "Dim op is not defined by a reshape op.");
1190
1191 // dim of a memref reshape can be folded if dim.getIndex() dominates the
1192 // reshape. Instead of using `DominanceInfo` (which is usually costly) we
1193 // cheaply check that either of the following conditions hold:
1194 // 1. dim.getIndex() is defined in the same block as reshape but before
1195 // reshape.
1196 // 2. dim.getIndex() is defined in a parent block of
1197 // reshape.
1198
1199 // Check condition 1
1200 if (dim.getIndex().getParentBlock() == reshape->getBlock()) {
1201 if (auto *definingOp = dim.getIndex().getDefiningOp()) {
1202 if (reshape->isBeforeInBlock(definingOp)) {
1203 return rewriter.notifyMatchFailure(
1204 dim,
1205 "dim.getIndex is not defined before reshape in the same block.");
1206 }
1207 } // else dim.getIndex is a block argument to reshape->getBlock and
1208 // dominates reshape
1209 } // Check condition 2
1210 else if (dim->getBlock() != reshape->getBlock() &&
1211 !dim.getIndex().getParentRegion()->isProperAncestor(
1212 reshape->getParentRegion())) {
1213 // If dim and reshape are in the same block but dim.getIndex() isn't, we
1214 // already know dim.getIndex() dominates reshape without calling
1215 // `isProperAncestor`
1216 return rewriter.notifyMatchFailure(
1217 dim, "dim.getIndex does not dominate reshape.");
1218 }
1219
1220 // Place the load directly after the reshape to ensure that the shape memref
1221 // was not mutated.
1222 rewriter.setInsertionPointAfter(reshape);
1223 Location loc = dim.getLoc();
1224 Value load =
1225 LoadOp::create(rewriter, loc, reshape.getShape(), dim.getIndex());
1226 if (load.getType() != dim.getType())
1227 load = arith::IndexCastOp::create(rewriter, loc, dim.getType(), load);
1228 rewriter.replaceOp(dim, load);
1229 return success();
1230 }
1231};
1232
1233} // namespace
1234
1235void DimOp::getCanonicalizationPatterns(RewritePatternSet &results,
1236 MLIRContext *context) {
1237 results.add<DimOfMemRefReshape>(context);
1238}
1239
1240// ---------------------------------------------------------------------------
1241// DmaStartOp
1242// ---------------------------------------------------------------------------
1243
1244void DmaStartOp::build(OpBuilder &builder, OperationState &result,
1245 Value srcMemRef, ValueRange srcIndices, Value destMemRef,
1246 ValueRange destIndices, Value numElements,
1247 Value tagMemRef, ValueRange tagIndices, Value stride,
1248 Value elementsPerStride) {
1249 result.addOperands(srcMemRef);
1250 result.addOperands(srcIndices);
1251 result.addOperands(destMemRef);
1252 result.addOperands(destIndices);
1253 result.addOperands({numElements, tagMemRef});
1254 result.addOperands(tagIndices);
1255 if (stride)
1256 result.addOperands({stride, elementsPerStride});
1257}
1258
1259void DmaStartOp::print(OpAsmPrinter &p) {
1260 p << " " << getSrcMemRef() << '[' << getSrcIndices() << "], "
1261 << getDstMemRef() << '[' << getDstIndices() << "], " << getNumElements()
1262 << ", " << getTagMemRef() << '[' << getTagIndices() << ']';
1263 if (isStrided())
1264 p << ", " << getStride() << ", " << getNumElementsPerStride();
1265
1266 p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary());
1267 p << " : " << getSrcMemRef().getType() << ", " << getDstMemRef().getType()
1268 << ", " << getTagMemRef().getType();
1269}
1270
1271// Parse DmaStartOp.
1272// Ex:
1273// %dma_id = dma_start %src[%i, %j], %dst[%k, %l], %size,
1274// %tag[%index], %stride, %num_elt_per_stride :
1275// : memref<3076 x f32, 0>,
1276// memref<1024 x f32, 2>,
1277// memref<1 x i32>
1278//
1279ParseResult DmaStartOp::parse(OpAsmParser &parser, OperationState &result) {
1280 OpAsmParser::UnresolvedOperand srcMemRefInfo;
1281 SmallVector<OpAsmParser::UnresolvedOperand, 4> srcIndexInfos;
1282 OpAsmParser::UnresolvedOperand dstMemRefInfo;
1283 SmallVector<OpAsmParser::UnresolvedOperand, 4> dstIndexInfos;
1284 OpAsmParser::UnresolvedOperand numElementsInfo;
1285 OpAsmParser::UnresolvedOperand tagMemrefInfo;
1286 SmallVector<OpAsmParser::UnresolvedOperand, 4> tagIndexInfos;
1287 SmallVector<OpAsmParser::UnresolvedOperand, 2> strideInfo;
1288
1289 SmallVector<Type, 3> types;
1290 auto indexType = parser.getBuilder().getIndexType();
1291
1292 // Parse and resolve the following list of operands:
1293 // *) source memref followed by its indices (in square brackets).
1294 // *) destination memref followed by its indices (in square brackets).
1295 // *) dma size in KiB.
1296 if (parser.parseOperand(srcMemRefInfo) ||
1297 parser.parseOperandList(srcIndexInfos, OpAsmParser::Delimiter::Square) ||
1298 parser.parseComma() || parser.parseOperand(dstMemRefInfo) ||
1299 parser.parseOperandList(dstIndexInfos, OpAsmParser::Delimiter::Square) ||
1300 parser.parseComma() || parser.parseOperand(numElementsInfo) ||
1301 parser.parseComma() || parser.parseOperand(tagMemrefInfo) ||
1302 parser.parseOperandList(tagIndexInfos, OpAsmParser::Delimiter::Square))
1303 return failure();
1304
1305 // Parse optional stride and elements per stride.
1306 if (parser.parseTrailingOperandList(strideInfo))
1307 return failure();
1308
1309 bool isStrided = strideInfo.size() == 2;
1310 if (!strideInfo.empty() && !isStrided) {
1311 return parser.emitError(parser.getNameLoc(),
1312 "expected two stride related operands");
1313 }
1314
1315 if (parser.parseColonTypeList(types))
1316 return failure();
1317 if (types.size() != 3)
1318 return parser.emitError(parser.getNameLoc(), "fewer/more types expected");
1319
1320 if (parser.resolveOperand(srcMemRefInfo, types[0], result.operands) ||
1321 parser.resolveOperands(srcIndexInfos, indexType, result.operands) ||
1322 parser.resolveOperand(dstMemRefInfo, types[1], result.operands) ||
1323 parser.resolveOperands(dstIndexInfos, indexType, result.operands) ||
1324 // size should be an index.
1325 parser.resolveOperand(numElementsInfo, indexType, result.operands) ||
1326 parser.resolveOperand(tagMemrefInfo, types[2], result.operands) ||
1327 // tag indices should be index.
1328 parser.resolveOperands(tagIndexInfos, indexType, result.operands))
1329 return failure();
1330
1331 if (isStrided) {
1332 if (parser.resolveOperands(strideInfo, indexType, result.operands))
1333 return failure();
1334 }
1335
1336 return success();
1337}
1338
1339LogicalResult DmaStartOp::verify() {
1340 unsigned numOperands = getNumOperands();
1341
1342 // Mandatory non-variadic operands are: src memref, dst memref, tag memref and
1343 // the number of elements.
1344 if (numOperands < 4)
1345 return emitOpError("expected at least 4 operands");
1346
1347 // Check types of operands. The order of these calls is important: the later
1348 // calls rely on some type properties to compute the operand position.
1349 // 1. Source memref.
1350 if (!llvm::isa<MemRefType>(getSrcMemRef().getType()))
1351 return emitOpError("expected source to be of memref type");
1352 if (numOperands < getSrcMemRefRank() + 4)
1353 return emitOpError() << "expected at least " << getSrcMemRefRank() + 4
1354 << " operands";
1355 if (!getSrcIndices().empty() &&
1356 !llvm::all_of(getSrcIndices().getTypes(),
1357 [](Type t) { return t.isIndex(); }))
1358 return emitOpError("expected source indices to be of index type");
1359
1360 // 2. Destination memref.
1361 if (!llvm::isa<MemRefType>(getDstMemRef().getType()))
1362 return emitOpError("expected destination to be of memref type");
1363 unsigned numExpectedOperands = getSrcMemRefRank() + getDstMemRefRank() + 4;
1364 if (numOperands < numExpectedOperands)
1365 return emitOpError() << "expected at least " << numExpectedOperands
1366 << " operands";
1367 if (!getDstIndices().empty() &&
1368 !llvm::all_of(getDstIndices().getTypes(),
1369 [](Type t) { return t.isIndex(); }))
1370 return emitOpError("expected destination indices to be of index type");
1371
1372 // 3. Number of elements.
1373 if (!getNumElements().getType().isIndex())
1374 return emitOpError("expected num elements to be of index type");
1375
1376 // 4. Tag memref.
1377 if (!llvm::isa<MemRefType>(getTagMemRef().getType()))
1378 return emitOpError("expected tag to be of memref type");
1379 numExpectedOperands += getTagMemRefRank();
1380 if (numOperands < numExpectedOperands)
1381 return emitOpError() << "expected at least " << numExpectedOperands
1382 << " operands";
1383 if (!getTagIndices().empty() &&
1384 !llvm::all_of(getTagIndices().getTypes(),
1385 [](Type t) { return t.isIndex(); }))
1386 return emitOpError("expected tag indices to be of index type");
1387
1388 // Optional stride-related operands must be either both present or both
1389 // absent.
1390 if (numOperands != numExpectedOperands &&
1391 numOperands != numExpectedOperands + 2)
1392 return emitOpError("incorrect number of operands");
1393
1394 // 5. Strides.
1395 if (isStrided()) {
1396 if (!getStride().getType().isIndex() ||
1397 !getNumElementsPerStride().getType().isIndex())
1398 return emitOpError(
1399 "expected stride and num elements per stride to be of type index");
1400 }
1401
1402 return success();
1403}
1404
1405LogicalResult DmaStartOp::fold(FoldAdaptor adaptor,
1406 SmallVectorImpl<OpFoldResult> &results) {
1407 /// dma_start(memrefcast) -> dma_start
1408 return foldMemRefCast(*this);
1409}
1410
1411void DmaStartOp::setMemrefsAndIndices(RewriterBase &rewriter, Value newSrc,
1412 ValueRange newSrcIndices, Value newDst,
1413 ValueRange newDstIndices) {
1414 /// dma_start has special handling for variadic rank
1415 SmallVector<Value> newOperands;
1416 newOperands.push_back(newSrc);
1417 llvm::append_range(newOperands, newSrcIndices);
1418 newOperands.push_back(newDst);
1419 llvm::append_range(newOperands, newDstIndices);
1420 newOperands.push_back(getNumElements());
1421 newOperands.push_back(getTagMemRef());
1422 llvm::append_range(newOperands, getTagIndices());
1423 if (isStrided()) {
1424 newOperands.push_back(getStride());
1425 newOperands.push_back(getNumElementsPerStride());
1426 }
1427
1428 rewriter.modifyOpInPlace(*this, [&]() { (*this)->setOperands(newOperands); });
1429}
1430
1431// ---------------------------------------------------------------------------
1432// DmaWaitOp
1433// ---------------------------------------------------------------------------
1434
1435LogicalResult DmaWaitOp::fold(FoldAdaptor adaptor,
1436 SmallVectorImpl<OpFoldResult> &results) {
1437 /// dma_wait(memrefcast) -> dma_wait
1438 return foldMemRefCast(*this);
1439}
1440
1441LogicalResult DmaWaitOp::verify() {
1442 // Check that the number of tag indices matches the tagMemRef rank.
1443 unsigned numTagIndices = getTagIndices().size();
1444 unsigned tagMemRefRank = getTagMemRefRank();
1445 if (numTagIndices != tagMemRefRank)
1446 return emitOpError() << "expected tagIndices to have the same number of "
1447 "elements as the tagMemRef rank, expected "
1448 << tagMemRefRank << ", but got " << numTagIndices;
1449 return success();
1450}
1451
1452//===----------------------------------------------------------------------===//
1453// ExtractAlignedPointerAsIndexOp
1454//===----------------------------------------------------------------------===//
1455
1456void ExtractAlignedPointerAsIndexOp::getAsmResultNames(
1457 function_ref<void(Value, StringRef)> setNameFn) {
1458 setNameFn(getResult(), "intptr");
1459}
1460
1461//===----------------------------------------------------------------------===//
1462// ExtractStridedMetadataOp
1463//===----------------------------------------------------------------------===//
1464
1465/// The number and type of the results are inferred from the
1466/// shape of the source.
1467LogicalResult ExtractStridedMetadataOp::inferReturnTypes(
1468 MLIRContext *context, std::optional<Location> location,
1469 ExtractStridedMetadataOp::Adaptor adaptor,
1470 SmallVectorImpl<Type> &inferredReturnTypes) {
1471 auto sourceType = llvm::dyn_cast<MemRefType>(adaptor.getSource().getType());
1472 if (!sourceType)
1473 return failure();
1474
1475 unsigned sourceRank = sourceType.getRank();
1476 IndexType indexType = IndexType::get(context);
1477 auto memrefType =
1478 MemRefType::get({}, sourceType.getElementType(),
1479 MemRefLayoutAttrInterface{}, sourceType.getMemorySpace());
1480 // Base.
1481 inferredReturnTypes.push_back(memrefType);
1482 // Offset.
1483 inferredReturnTypes.push_back(indexType);
1484 // Sizes and strides.
1485 for (unsigned i = 0; i < sourceRank * 2; ++i)
1486 inferredReturnTypes.push_back(indexType);
1487 return success();
1488}
1489
1490void ExtractStridedMetadataOp::getAsmResultNames(
1491 function_ref<void(Value, StringRef)> setNameFn) {
1492 setNameFn(getBaseBuffer(), "base_buffer");
1493 setNameFn(getOffset(), "offset");
1494 // For multi-result to work properly with pretty names and packed syntax `x:3`
1495 // we can only give a pretty name to the first value in the pack.
1496 if (!getSizes().empty()) {
1497 setNameFn(getSizes().front(), "sizes");
1498 setNameFn(getStrides().front(), "strides");
1499 }
1500}
1501
1502/// Helper function to perform the replacement of all constant uses of `values`
1503/// by a materialized constant extracted from `maybeConstants`.
1504/// `values` and `maybeConstants` are expected to have the same size.
1505template <typename Container>
1506static bool replaceConstantUsesOf(OpBuilder &rewriter, Location loc,
1507 Container values,
1508 ArrayRef<OpFoldResult> maybeConstants) {
1509 assert(values.size() == maybeConstants.size() &&
1510 " expected values and maybeConstants of the same size");
1511 bool atLeastOneReplacement = false;
1512 for (auto [maybeConstant, result] : llvm::zip(maybeConstants, values)) {
1513 // Don't materialize a constant if there are no uses: this would indice
1514 // infinite loops in the driver.
1515 if (result.use_empty() || maybeConstant == getAsOpFoldResult(result))
1516 continue;
1517 assert(isa<Attribute>(maybeConstant) &&
1518 "The constified value should be either unchanged (i.e., == result) "
1519 "or a constant");
1521 rewriter, loc,
1522 llvm::cast<IntegerAttr>(cast<Attribute>(maybeConstant)).getInt());
1523 for (Operation *op : llvm::make_early_inc_range(result.getUsers())) {
1524 // modifyOpInPlace: lambda cannot capture structured bindings in C++17
1525 // yet.
1526 op->replaceUsesOfWith(result, constantVal);
1527 atLeastOneReplacement = true;
1528 }
1529 }
1530 return atLeastOneReplacement;
1531}
1532
1533LogicalResult
1534ExtractStridedMetadataOp::fold(FoldAdaptor adaptor,
1535 SmallVectorImpl<OpFoldResult> &results) {
1536 OpBuilder builder(*this);
1537
1538 bool atLeastOneReplacement = replaceConstantUsesOf(
1539 builder, getLoc(), ArrayRef<TypedValue<IndexType>>(getOffset()),
1540 getConstifiedMixedOffset());
1541 atLeastOneReplacement |= replaceConstantUsesOf(builder, getLoc(), getSizes(),
1542 getConstifiedMixedSizes());
1543 atLeastOneReplacement |= replaceConstantUsesOf(
1544 builder, getLoc(), getStrides(), getConstifiedMixedStrides());
1545
1546 // extract_strided_metadata(cast(x)) -> extract_strided_metadata(x).
1547 if (auto prev = getSource().getDefiningOp<CastOp>())
1548 if (isa<MemRefType>(prev.getSource().getType())) {
1549 getSourceMutable().assign(prev.getSource());
1550 atLeastOneReplacement = true;
1551 }
1552
1553 return success(atLeastOneReplacement);
1554}
1555
1556SmallVector<OpFoldResult> ExtractStridedMetadataOp::getConstifiedMixedSizes() {
1557 SmallVector<OpFoldResult> values = getAsOpFoldResult(getSizes());
1558 constifyIndexValues(values, getSource().getType().getShape());
1559 return values;
1560}
1561
1562SmallVector<OpFoldResult>
1563ExtractStridedMetadataOp::getConstifiedMixedStrides() {
1564 SmallVector<OpFoldResult> values = getAsOpFoldResult(getStrides());
1565 SmallVector<int64_t> staticValues;
1566 int64_t unused;
1567 LogicalResult status =
1568 getSource().getType().getStridesAndOffset(staticValues, unused);
1569 (void)status;
1570 assert(succeeded(status) && "could not get strides from type");
1571 constifyIndexValues(values, staticValues);
1572 return values;
1573}
1574
1575OpFoldResult ExtractStridedMetadataOp::getConstifiedMixedOffset() {
1576 OpFoldResult offsetOfr = getAsOpFoldResult(getOffset());
1577 SmallVector<OpFoldResult> values(1, offsetOfr);
1578 SmallVector<int64_t> staticValues, unused;
1579 int64_t offset;
1580 LogicalResult status =
1581 getSource().getType().getStridesAndOffset(unused, offset);
1582 (void)status;
1583 assert(succeeded(status) && "could not get offset from type");
1584 staticValues.push_back(offset);
1585 constifyIndexValues(values, staticValues);
1586 return values[0];
1587}
1588
1589//===----------------------------------------------------------------------===//
1590// GenericAtomicRMWOp
1591//===----------------------------------------------------------------------===//
1592
1593void GenericAtomicRMWOp::build(OpBuilder &builder, OperationState &result,
1594 Value memref, ValueRange ivs) {
1595 OpBuilder::InsertionGuard g(builder);
1596 result.addOperands(memref);
1597 result.addOperands(ivs);
1598
1599 if (auto memrefType = llvm::dyn_cast<MemRefType>(memref.getType())) {
1600 Type elementType = memrefType.getElementType();
1601 result.addTypes(elementType);
1602
1603 Region *bodyRegion = result.addRegion();
1604 builder.createBlock(bodyRegion);
1605 bodyRegion->addArgument(elementType, memref.getLoc());
1606 }
1607}
1608
1609LogicalResult GenericAtomicRMWOp::verify() {
1610 auto &body = getRegion();
1611 if (body.getNumArguments() != 1)
1612 return emitOpError("expected single number of entry block arguments");
1613
1614 if (getResult().getType() != body.getArgument(0).getType())
1615 return emitOpError("expected block argument of the same type result type");
1616
1617 bool hasSideEffects =
1618 body.walk([&](Operation *nestedOp) {
1619 if (isMemoryEffectFree(nestedOp))
1620 return WalkResult::advance();
1621 nestedOp->emitError(
1622 "body of 'memref.generic_atomic_rmw' should contain "
1623 "only operations with no side effects");
1624 return WalkResult::interrupt();
1625 })
1626 .wasInterrupted();
1627 return hasSideEffects ? failure() : success();
1628}
1629
1630ParseResult GenericAtomicRMWOp::parse(OpAsmParser &parser,
1631 OperationState &result) {
1632 OpAsmParser::UnresolvedOperand memref;
1633 Type memrefType;
1634 SmallVector<OpAsmParser::UnresolvedOperand, 4> ivs;
1635
1636 Type indexType = parser.getBuilder().getIndexType();
1637 if (parser.parseOperand(memref) ||
1639 parser.parseColonType(memrefType) ||
1640 parser.resolveOperand(memref, memrefType, result.operands) ||
1641 parser.resolveOperands(ivs, indexType, result.operands))
1642 return failure();
1643
1644 Region *body = result.addRegion();
1645 if (parser.parseRegion(*body, {}) ||
1646 parser.parseOptionalAttrDict(result.attributes))
1647 return failure();
1648 result.types.push_back(llvm::cast<MemRefType>(memrefType).getElementType());
1649 return success();
1650}
1651
1652void GenericAtomicRMWOp::print(OpAsmPrinter &p) {
1653 p << ' ' << getMemref() << "[" << getIndices()
1654 << "] : " << getMemref().getType() << ' ';
1655 p.printRegion(getRegion());
1656 p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary());
1657}
1658
1659TypedValue<MemRefType> GenericAtomicRMWOp::getAccessedMemref() {
1660 return getMemref();
1661}
1662
1663std::optional<SmallVector<Value>> GenericAtomicRMWOp::updateMemrefAndIndices(
1664 RewriterBase &rewriter, Value newMemref, ValueRange newIndices) {
1665 rewriter.modifyOpInPlace(*this, [&]() {
1666 getMemrefMutable().assign(newMemref);
1667 getIndicesMutable().assign(newIndices);
1668 });
1669 return std::nullopt;
1670}
1671
1672//===----------------------------------------------------------------------===//
1673// AtomicYieldOp
1674//===----------------------------------------------------------------------===//
1675
1676LogicalResult AtomicYieldOp::verify() {
1677 Type parentType = (*this)->getParentOp()->getResultTypes().front();
1678 Type resultType = getResult().getType();
1679 if (parentType != resultType)
1680 return emitOpError() << "types mismatch between yield op: " << resultType
1681 << " and its parent: " << parentType;
1682 return success();
1683}
1684
1685//===----------------------------------------------------------------------===//
1686// GlobalOp
1687//===----------------------------------------------------------------------===//
1688
1690 TypeAttr type,
1691 Attribute initialValue) {
1692 p << type;
1693 if (!op.isExternal()) {
1694 p << " = ";
1695 if (op.isUninitialized())
1696 p << "uninitialized";
1697 else
1698 p.printAttributeWithoutType(initialValue);
1699 }
1700}
1701
1702static ParseResult
1704 Attribute &initialValue) {
1705 Type type;
1706 if (parser.parseType(type))
1707 return failure();
1708
1709 auto memrefType = llvm::dyn_cast<MemRefType>(type);
1710 if (!memrefType || !memrefType.hasStaticShape())
1711 return parser.emitError(parser.getNameLoc())
1712 << "type should be static shaped memref, but got " << type;
1713 typeAttr = TypeAttr::get(type);
1714
1715 if (parser.parseOptionalEqual())
1716 return success();
1717
1718 if (succeeded(parser.parseOptionalKeyword("uninitialized"))) {
1719 initialValue = UnitAttr::get(parser.getContext());
1720 return success();
1721 }
1722
1723 Type tensorType = getTensorTypeFromMemRefType(memrefType);
1724 if (parser.parseAttribute(initialValue, tensorType))
1725 return failure();
1726 if (!llvm::isa<ElementsAttr>(initialValue))
1727 return parser.emitError(parser.getNameLoc())
1728 << "initial value should be a unit or elements attribute";
1729 return success();
1730}
1731
1732LogicalResult GlobalOp::verify() {
1733 auto memrefType = llvm::dyn_cast<MemRefType>(getType());
1734 if (!memrefType || !memrefType.hasStaticShape())
1735 return emitOpError("type should be static shaped memref, but got ")
1736 << getType();
1737
1738 // Verify that the initial value, if present, is either a unit attribute or
1739 // an elements attribute.
1740 if (getInitialValue().has_value()) {
1741 Attribute initValue = getInitialValue().value();
1742 if (!llvm::isa<UnitAttr>(initValue) && !llvm::isa<ElementsAttr>(initValue))
1743 return emitOpError("initial value should be a unit or elements "
1744 "attribute, but got ")
1745 << initValue;
1746
1747 // Check that the type of the initial value is compatible with the type of
1748 // the global variable.
1749 if (auto elementsAttr = llvm::dyn_cast<ElementsAttr>(initValue)) {
1750 // Check the element types match.
1751 auto initElementType =
1752 cast<TensorType>(elementsAttr.getType()).getElementType();
1753 auto memrefElementType = memrefType.getElementType();
1754
1755 if (initElementType != memrefElementType)
1756 return emitOpError("initial value element expected to be of type ")
1757 << memrefElementType << ", but was of type " << initElementType;
1758
1759 // Check the shapes match, given that memref globals can only produce
1760 // statically shaped memrefs and elements literal type must have a static
1761 // shape we can assume both types are shaped.
1762 auto initShape = elementsAttr.getShapedType().getShape();
1763 auto memrefShape = memrefType.getShape();
1764 if (initShape != memrefShape)
1765 return emitOpError("initial value shape expected to be ")
1766 << memrefShape << " but was " << initShape;
1767 }
1768 }
1769
1770 // TODO: verify visibility for declarations.
1771 return success();
1772}
1773
1774ElementsAttr GlobalOp::getConstantInitValue() {
1775 auto initVal = getInitialValue();
1776 if (getConstant() && initVal.has_value())
1777 return llvm::cast<ElementsAttr>(initVal.value());
1778 return {};
1779}
1780
1781//===----------------------------------------------------------------------===//
1782// GetGlobalOp
1783//===----------------------------------------------------------------------===//
1784
1785LogicalResult
1786GetGlobalOp::verifySymbolUses(SymbolTableCollection &symbolTable) {
1787 // Verify that the result type is same as the type of the referenced
1788 // memref.global op.
1789 auto global =
1790 symbolTable.lookupNearestSymbolFrom<GlobalOp>(*this, getNameAttr());
1791 if (!global)
1792 return emitOpError("'")
1793 << getName() << "' does not reference a valid global memref";
1794
1795 Type resultType = getResult().getType();
1796 if (global.getType() != resultType)
1797 return emitOpError("result type ")
1798 << resultType << " does not match type " << global.getType()
1799 << " of the global memref @" << getName();
1800 return success();
1801}
1802
1803//===----------------------------------------------------------------------===//
1804// LoadOp
1805//===----------------------------------------------------------------------===//
1806
1807static ParseResult parseBoolAttr(OpAsmParser &parser, BoolAttr &result) {
1808 Attribute attr;
1809 if (parser.parseAttribute(attr))
1810 return failure();
1811 result = dyn_cast<BoolAttr>(attr);
1812 if (!result)
1813 return parser.emitError(parser.getCurrentLocation(),
1814 "expected boolean attribute");
1815 return success();
1816}
1817
1818static void printBoolAttr(OpAsmPrinter &printer, Operation *, BoolAttr attr) {
1819 printer.printAttribute(attr);
1820}
1821
1822OpFoldResult LoadOp::fold(FoldAdaptor adaptor) {
1823 /// load(memrefcast) -> load
1824 if (succeeded(foldMemRefCast(*this)))
1825 return getResult();
1826
1827 // Fold load from a global constant memref.
1828 auto getGlobalOp = getMemref().getDefiningOp<memref::GetGlobalOp>();
1829 if (!getGlobalOp)
1830 return {};
1831
1832 // Get to the memref.global defining the symbol.
1834 getGlobalOp, getGlobalOp.getNameAttr());
1835 if (!global)
1836 return {};
1837 // If it's a splat constant, we can fold irrespective of indices.
1838 auto splatAttr =
1839 dyn_cast_or_null<SplatElementsAttr>(global.getConstantInitValue());
1840 if (!splatAttr)
1841 return {};
1842
1843 return splatAttr.getSplatValue<Attribute>();
1844}
1845
1846TypedValue<MemRefType> LoadOp::getAccessedMemref() { return getMemref(); }
1847
1848std::optional<SmallVector<Value>>
1849LoadOp::updateMemrefAndIndices(RewriterBase &rewriter, Value newMemref,
1850 ValueRange newIndices) {
1851 rewriter.modifyOpInPlace(*this, [&]() {
1852 getMemrefMutable().assign(newMemref);
1853 getIndicesMutable().assign(newIndices);
1854 });
1855 return std::nullopt;
1856}
1857
1858FailureOr<std::optional<SmallVector<Value>>>
1859LoadOp::bubbleDownCasts(OpBuilder &builder) {
1861 getResult());
1862}
1863
1864//===----------------------------------------------------------------------===//
1865// MemorySpaceCastOp
1866//===----------------------------------------------------------------------===//
1867
1868void MemorySpaceCastOp::getAsmResultNames(
1869 function_ref<void(Value, StringRef)> setNameFn) {
1870 setNameFn(getResult(), "memspacecast");
1871}
1872
1873bool MemorySpaceCastOp::areCastCompatible(TypeRange inputs, TypeRange outputs) {
1874 if (inputs.size() != 1 || outputs.size() != 1)
1875 return false;
1876 Type a = inputs.front(), b = outputs.front();
1877 auto aT = llvm::dyn_cast<MemRefType>(a);
1878 auto bT = llvm::dyn_cast<MemRefType>(b);
1879
1880 auto uaT = llvm::dyn_cast<UnrankedMemRefType>(a);
1881 auto ubT = llvm::dyn_cast<UnrankedMemRefType>(b);
1882
1883 if (aT && bT) {
1884 if (aT.getElementType() != bT.getElementType())
1885 return false;
1886 if (aT.getLayout() != bT.getLayout())
1887 return false;
1888 if (aT.getShape() != bT.getShape())
1889 return false;
1890 return true;
1891 }
1892 if (uaT && ubT) {
1893 return uaT.getElementType() == ubT.getElementType();
1894 }
1895 return false;
1896}
1897
1898OpFoldResult MemorySpaceCastOp::fold(FoldAdaptor adaptor) {
1899 // memory_space_cast(memory_space_cast(v, t1), t2) -> memory_space_cast(v,
1900 // t2)
1901 if (auto parentCast = getSource().getDefiningOp<MemorySpaceCastOp>()) {
1902 getSourceMutable().assign(parentCast.getSource());
1903 return getResult();
1904 }
1905 return Value{};
1906}
1907
1908TypedValue<PtrLikeTypeInterface> MemorySpaceCastOp::getSourcePtr() {
1909 return getSource();
1910}
1911
1912TypedValue<PtrLikeTypeInterface> MemorySpaceCastOp::getTargetPtr() {
1913 return getDest();
1914}
1915
1916bool MemorySpaceCastOp::isValidMemorySpaceCast(PtrLikeTypeInterface tgt,
1917 PtrLikeTypeInterface src) {
1918 return isa<BaseMemRefType>(tgt) &&
1919 tgt.clonePtrWith(src.getMemorySpace(), std::nullopt) == src;
1920}
1921
1922MemorySpaceCastOpInterface MemorySpaceCastOp::cloneMemorySpaceCastOp(
1923 OpBuilder &b, PtrLikeTypeInterface tgt,
1925 assert(isValidMemorySpaceCast(tgt, src.getType()) && "invalid arguments");
1926 return MemorySpaceCastOp::create(b, getLoc(), tgt, src);
1927}
1928
1929/// The only cast we recognize as promotable is to the generic space.
1930bool MemorySpaceCastOp::isSourcePromotable() {
1931 return getDest().getType().getMemorySpace() == nullptr;
1932}
1933
1934//===----------------------------------------------------------------------===//
1935// PrefetchOp
1936//===----------------------------------------------------------------------===//
1937
1938void PrefetchOp::print(OpAsmPrinter &p) {
1939 p << " " << getMemref() << '[';
1941 p << ']' << ", " << (getIsWrite() ? "write" : "read");
1942 p << ", locality<" << getLocalityHint();
1943 p << ">, " << (getIsDataCache() ? "data" : "instr");
1945 (*this)->getDiscardableAttrDictionary(),
1946 /*elidedAttrs=*/{"localityHint", "isWrite", "isDataCache"});
1947 p << " : " << getMemRefType();
1948}
1949
1950ParseResult PrefetchOp::parse(OpAsmParser &parser, OperationState &result) {
1951 OpAsmParser::UnresolvedOperand memrefInfo;
1952 SmallVector<OpAsmParser::UnresolvedOperand, 4> indexInfo;
1953 IntegerAttr localityHint;
1954 MemRefType type;
1955 StringRef readOrWrite, cacheType;
1956
1957 auto indexTy = parser.getBuilder().getIndexType();
1958 auto i32Type = parser.getBuilder().getIntegerType(32);
1959 if (parser.parseOperand(memrefInfo) ||
1961 parser.parseComma() || parser.parseKeyword(&readOrWrite) ||
1962 parser.parseComma() || parser.parseKeyword("locality") ||
1963 parser.parseLess() ||
1964 parser.parseAttribute(localityHint, i32Type, "localityHint",
1965 result.attributes) ||
1966 parser.parseGreater() || parser.parseComma() ||
1967 parser.parseKeyword(&cacheType) || parser.parseColonType(type) ||
1968 parser.resolveOperand(memrefInfo, type, result.operands) ||
1969 parser.resolveOperands(indexInfo, indexTy, result.operands))
1970 return failure();
1971
1972 if (readOrWrite != "read" && readOrWrite != "write")
1973 return parser.emitError(parser.getNameLoc(),
1974 "rw specifier has to be 'read' or 'write'");
1975 result.addAttribute(PrefetchOp::getIsWriteAttrStrName(),
1976 parser.getBuilder().getBoolAttr(readOrWrite == "write"));
1977
1978 if (cacheType != "data" && cacheType != "instr")
1979 return parser.emitError(parser.getNameLoc(),
1980 "cache type has to be 'data' or 'instr'");
1981
1982 result.addAttribute(PrefetchOp::getIsDataCacheAttrStrName(),
1983 parser.getBuilder().getBoolAttr(cacheType == "data"));
1984
1985 return success();
1986}
1987
1988LogicalResult PrefetchOp::verify() {
1989 if (getNumOperands() != 1 + getMemRefType().getRank())
1990 return emitOpError("too few indices");
1991
1992 return success();
1993}
1994
1995LogicalResult PrefetchOp::fold(FoldAdaptor adaptor,
1996 SmallVectorImpl<OpFoldResult> &results) {
1997 // prefetch(memrefcast) -> prefetch
1998 return foldMemRefCast(*this);
1999}
2000
2001TypedValue<MemRefType> PrefetchOp::getAccessedMemref() { return getMemref(); }
2002
2003std::optional<SmallVector<Value>>
2004PrefetchOp::updateMemrefAndIndices(RewriterBase &rewriter, Value newMemref,
2005 ValueRange newIndices) {
2006 rewriter.modifyOpInPlace(*this, [&]() {
2007 getMemrefMutable().assign(newMemref);
2008 getIndicesMutable().assign(newIndices);
2009 });
2010 return std::nullopt;
2011}
2012
2013//===----------------------------------------------------------------------===//
2014// RankOp
2015//===----------------------------------------------------------------------===//
2016
2017OpFoldResult RankOp::fold(FoldAdaptor adaptor) {
2018 // Constant fold rank when the rank of the operand is known.
2019 auto type = getOperand().getType();
2020 auto shapedType = llvm::dyn_cast<ShapedType>(type);
2021 if (shapedType && shapedType.hasRank())
2022 return IntegerAttr::get(IndexType::get(getContext()), shapedType.getRank());
2023 return IntegerAttr();
2024}
2025
2026//===----------------------------------------------------------------------===//
2027// ReinterpretCastOp
2028//===----------------------------------------------------------------------===//
2029
2030namespace {
2031
2032struct PrintDynamicOrValue {
2033 int64_t value;
2034};
2035
2036Diagnostic &operator<<(Diagnostic &diag, PrintDynamicOrValue printed) {
2037 if (ShapedType::isDynamic(printed.value))
2038 return diag << "dynamic";
2039 return diag << printed.value;
2040}
2041
2042} // namespace
2043
2044void ReinterpretCastOp::getAsmResultNames(
2045 function_ref<void(Value, StringRef)> setNameFn) {
2046 setNameFn(getResult(), "reinterpret_cast");
2047}
2048
2049/// Build a ReinterpretCastOp with all dynamic entries: `staticOffsets`,
2050/// `staticSizes` and `staticStrides` are automatically filled with
2051/// source-memref-rank sentinel values that encode dynamic entries.
2052void ReinterpretCastOp::build(OpBuilder &b, OperationState &result,
2053 MemRefType resultType, Value source,
2054 OpFoldResult offset, ArrayRef<OpFoldResult> sizes,
2055 ArrayRef<OpFoldResult> strides,
2056 ArrayRef<NamedAttribute> attrs) {
2057 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides;
2058 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides;
2059 dispatchIndexOpFoldResults(offset, dynamicOffsets, staticOffsets);
2060 dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes);
2061 dispatchIndexOpFoldResults(strides, dynamicStrides, staticStrides);
2062 result.addAttributes(attrs);
2063 build(b, result, resultType, source, dynamicOffsets, dynamicSizes,
2064 dynamicStrides, b.getDenseI64ArrayAttr(staticOffsets),
2065 b.getDenseI64ArrayAttr(staticSizes),
2066 b.getDenseI64ArrayAttr(staticStrides));
2067}
2068
2069void ReinterpretCastOp::build(OpBuilder &b, OperationState &result,
2070 Value source, OpFoldResult offset,
2071 ArrayRef<OpFoldResult> sizes,
2072 ArrayRef<OpFoldResult> strides,
2073 ArrayRef<NamedAttribute> attrs) {
2074 auto sourceType = cast<BaseMemRefType>(source.getType());
2075 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides;
2076 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides;
2077 dispatchIndexOpFoldResults(offset, dynamicOffsets, staticOffsets);
2078 dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes);
2079 dispatchIndexOpFoldResults(strides, dynamicStrides, staticStrides);
2080 auto stridedLayout = StridedLayoutAttr::get(
2081 b.getContext(), staticOffsets.front(), staticStrides);
2082 auto resultType = MemRefType::get(staticSizes, sourceType.getElementType(),
2083 stridedLayout, sourceType.getMemorySpace());
2084 build(b, result, resultType, source, offset, sizes, strides, attrs);
2085}
2086
2087void ReinterpretCastOp::build(OpBuilder &b, OperationState &result,
2088 MemRefType resultType, Value source,
2089 int64_t offset, ArrayRef<int64_t> sizes,
2090 ArrayRef<int64_t> strides,
2091 ArrayRef<NamedAttribute> attrs) {
2092 SmallVector<OpFoldResult> sizeValues = llvm::map_to_vector<4>(
2093 sizes, [&](int64_t v) -> OpFoldResult { return b.getI64IntegerAttr(v); });
2094 SmallVector<OpFoldResult> strideValues =
2095 llvm::map_to_vector<4>(strides, [&](int64_t v) -> OpFoldResult {
2096 return b.getI64IntegerAttr(v);
2097 });
2098 build(b, result, resultType, source, b.getI64IntegerAttr(offset), sizeValues,
2099 strideValues, attrs);
2100}
2101
2102void ReinterpretCastOp::build(OpBuilder &b, OperationState &result,
2103 MemRefType resultType, Value source, Value offset,
2104 ValueRange sizes, ValueRange strides,
2105 ArrayRef<NamedAttribute> attrs) {
2106 SmallVector<OpFoldResult> sizeValues =
2107 llvm::map_to_vector<4>(sizes, [](Value v) -> OpFoldResult { return v; });
2108 SmallVector<OpFoldResult> strideValues = llvm::map_to_vector<4>(
2109 strides, [](Value v) -> OpFoldResult { return v; });
2110 build(b, result, resultType, source, offset, sizeValues, strideValues, attrs);
2111}
2112
2113// TODO: ponder whether we want to allow missing trailing sizes/strides that are
2114// completed automatically, like we have for subview and extract_slice.
2115LogicalResult ReinterpretCastOp::verify() {
2116 // The source and result memrefs should be in the same memory space.
2117 auto srcType = llvm::cast<BaseMemRefType>(getSource().getType());
2118 auto resultType = llvm::cast<MemRefType>(getType());
2119 if (srcType.getMemorySpace() != resultType.getMemorySpace())
2120 return emitError("different memory spaces specified for source type ")
2121 << srcType << " and result memref type " << resultType;
2122 if (failed(verifyElementTypesMatch(*this, srcType, resultType, "source",
2123 "result")))
2124 return failure();
2125
2126 // Match sizes in result memref type and in static_sizes attribute.
2127 for (auto [idx, resultSize, expectedSize] :
2128 llvm::enumerate(resultType.getShape(), getStaticSizes())) {
2129 if (resultSize != expectedSize)
2130 return emitError("expected result type with size = ")
2131 << PrintDynamicOrValue{expectedSize} << " instead of "
2132 << PrintDynamicOrValue{resultSize} << " in dim = " << idx;
2133 }
2134
2135 // Match offset and strides in static_offset and static_strides attributes. If
2136 // result memref type has no affine map specified, this will assume an
2137 // identity layout.
2138 int64_t resultOffset;
2139 SmallVector<int64_t, 4> resultStrides;
2140 if (failed(resultType.getStridesAndOffset(resultStrides, resultOffset)))
2141 return emitError("expected result type to have strided layout but found ")
2142 << resultType;
2143
2144 // Match offset in result memref type and in static_offsets attribute.
2145 int64_t expectedOffset = getStaticOffsets().front();
2146 if (resultOffset != expectedOffset)
2147 return emitError("expected result type with offset = ")
2148 << PrintDynamicOrValue{expectedOffset} << " instead of "
2149 << PrintDynamicOrValue{resultOffset};
2150
2151 // Match strides in result memref type and in static_strides attribute.
2152 for (auto [idx, resultStride, expectedStride] :
2153 llvm::enumerate(resultStrides, getStaticStrides())) {
2154 if (resultStride != expectedStride)
2155 return emitError("expected result type with stride = ")
2156 << PrintDynamicOrValue{expectedStride} << " instead of "
2157 << PrintDynamicOrValue{resultStride} << " in dim = " << idx;
2158 }
2159
2160 return success();
2161}
2162
2163OpFoldResult ReinterpretCastOp::fold(FoldAdaptor /*operands*/) {
2164 Value src = getSource();
2165 auto getPrevSrc = [&]() -> Value {
2166 // reinterpret_cast(reinterpret_cast(x)) -> reinterpret_cast(x).
2167 if (auto prev = src.getDefiningOp<ReinterpretCastOp>())
2168 return prev.getSource();
2169
2170 // reinterpret_cast(cast(x)) -> reinterpret_cast(x).
2171 if (auto prev = src.getDefiningOp<CastOp>())
2172 return prev.getSource();
2173
2174 // reinterpret_cast(subview(x)) -> reinterpret_cast(x) if subview offsets
2175 // are 0.
2176 if (auto prev = src.getDefiningOp<SubViewOp>())
2177 if (llvm::all_of(prev.getMixedOffsets(), isZeroInteger))
2178 return prev.getSource();
2179
2180 return nullptr;
2181 };
2182
2183 if (auto prevSrc = getPrevSrc()) {
2184 getSourceMutable().assign(prevSrc);
2185 return getResult();
2186 }
2187
2188 // reinterpret_cast(x) w/o offset/shape/stride changes -> x
2189 if (ShapedType::isStaticShape(getType().getShape()) &&
2190 src.getType() == getType() && getStaticOffsets().front() == 0) {
2191 return src;
2192 }
2193
2194 return nullptr;
2195}
2196
2197SmallVector<OpFoldResult> ReinterpretCastOp::getConstifiedMixedSizes() {
2198 SmallVector<OpFoldResult> values = getMixedSizes();
2200 return values;
2201}
2202
2203SmallVector<OpFoldResult> ReinterpretCastOp::getConstifiedMixedStrides() {
2204 SmallVector<OpFoldResult> values = getMixedStrides();
2205 SmallVector<int64_t> staticValues;
2206 int64_t unused;
2207 LogicalResult status = getType().getStridesAndOffset(staticValues, unused);
2208 (void)status;
2209 assert(succeeded(status) && "could not get strides from type");
2210 constifyIndexValues(values, staticValues);
2211 return values;
2212}
2213
2214OpFoldResult ReinterpretCastOp::getConstifiedMixedOffset() {
2215 SmallVector<OpFoldResult> values = getMixedOffsets();
2216 assert(values.size() == 1 &&
2217 "reinterpret_cast must have one and only one offset");
2218 SmallVector<int64_t> staticValues, unused;
2219 int64_t offset;
2220 LogicalResult status = getType().getStridesAndOffset(unused, offset);
2221 (void)status;
2222 assert(succeeded(status) && "could not get offset from type");
2223 staticValues.push_back(offset);
2224 constifyIndexValues(values, staticValues);
2225 return values[0];
2226}
2227
2228namespace {
2229/// Replace the sequence:
2230/// ```
2231/// base, offset, sizes, strides = extract_strided_metadata src
2232/// dst = reinterpret_cast base to offset, sizes, strides
2233/// ```
2234/// With
2235///
2236/// ```
2237/// dst = memref.cast src
2238/// ```
2239///
2240/// Note: The cast operation is only inserted when the type of dst and src
2241/// are not the same. E.g., when going from <4xf32> to <?xf32>.
2242///
2243/// This pattern also matches when the offset, sizes, and strides don't come
2244/// directly from the `extract_strided_metadata`'s results but it can be
2245/// statically proven that they would hold the same values.
2246///
2247/// For instance, the following sequence would be replaced:
2248/// ```
2249/// base, offset, sizes, strides =
2250/// extract_strided_metadata memref : memref<3x4xty>
2251/// dst = reinterpret_cast base to 0, [3, 4], strides
2252/// ```
2253/// Because we know (thanks to the type of the input memref) that variable
2254/// `offset` and `sizes` will respectively hold 0 and [3, 4].
2255///
2256/// Similarly, the following sequence would be replaced:
2257/// ```
2258/// c0 = arith.constant 0
2259/// c4 = arith.constant 4
2260/// base, offset, sizes, strides =
2261/// extract_strided_metadata memref : memref<3x4xty>
2262/// dst = reinterpret_cast base to c0, [3, c4], strides
2263/// ```
2264/// Because we know that `offset`and `c0` will hold 0
2265/// and `c4` will hold 4.
2266///
2267/// If the pattern above does not match, the input of the
2268/// extract_strided_metadata is always folded into the input of the
2269/// reinterpret_cast operator. This allows for dead code elimination to get rid
2270/// of the extract_strided_metadata in some cases.
2271struct ReinterpretCastOpExtractStridedMetadataFolder
2272 : public OpRewritePattern<ReinterpretCastOp> {
2273public:
2274 using OpRewritePattern<ReinterpretCastOp>::OpRewritePattern;
2275
2276 LogicalResult matchAndRewrite(ReinterpretCastOp op,
2277 PatternRewriter &rewriter) const override {
2278 auto extractStridedMetadata =
2279 op.getSource().getDefiningOp<ExtractStridedMetadataOp>();
2280 if (!extractStridedMetadata)
2281 return failure();
2282
2283 // Check if the reinterpret cast reconstructs a memref with the exact same
2284 // properties as the extract strided metadata.
2285 auto isReinterpretCastNoop = [&]() -> bool {
2286 // First, check that the strides are the same.
2287 if (!llvm::equal(extractStridedMetadata.getConstifiedMixedStrides(),
2288 op.getConstifiedMixedStrides()))
2289 return false;
2290
2291 // Second, check the sizes.
2292 if (!llvm::equal(extractStridedMetadata.getConstifiedMixedSizes(),
2293 op.getConstifiedMixedSizes()))
2294 return false;
2295
2296 // Finally, check the offset.
2297 assert(op.getMixedOffsets().size() == 1 &&
2298 "reinterpret_cast with more than one offset should have been "
2299 "rejected by the verifier");
2300 return extractStridedMetadata.getConstifiedMixedOffset() ==
2301 op.getConstifiedMixedOffset();
2302 };
2303
2304 if (!isReinterpretCastNoop()) {
2305 // If the extract_strided_metadata / reinterpret_cast pair can't be
2306 // completely folded, then we could fold the input of the
2307 // extract_strided_metadata into the input of the reinterpret_cast
2308 // input. For some cases (e.g., static dimensions) the
2309 // the extract_strided_metadata is eliminated by dead code elimination.
2310 //
2311 // reinterpret_cast(extract_strided_metadata(x)) -> reinterpret_cast(x).
2312 //
2313 // We can always fold the input of a extract_strided_metadata operator
2314 // to the input of a reinterpret_cast operator, because they point to
2315 // the same memory. Note that the reinterpret_cast does not use the
2316 // layout of its input memref, only its base memory pointer which is
2317 // the same as the base pointer returned by the extract_strided_metadata
2318 // operator and the base pointer of the extract_strided_metadata memref
2319 // input.
2320 rewriter.modifyOpInPlace(op, [&]() {
2321 op.getSourceMutable().assign(extractStridedMetadata.getSource());
2322 });
2323 return success();
2324 }
2325
2326 // At this point, we know that the back and forth between extract strided
2327 // metadata and reinterpret cast is a noop. However, the final type of the
2328 // reinterpret cast may not be exactly the same as the original memref.
2329 // E.g., it could be changing a dimension from static to dynamic. Check that
2330 // here and add a cast if necessary.
2331 Type srcTy = extractStridedMetadata.getSource().getType();
2332 if (srcTy == op.getResult().getType())
2333 rewriter.replaceOp(op, extractStridedMetadata.getSource());
2334 else
2335 rewriter.replaceOpWithNewOp<CastOp>(op, op.getType(),
2336 extractStridedMetadata.getSource());
2337
2338 return success();
2339 }
2340};
2341
2342struct ReinterpretCastOpConstantFolder
2343 : public OpRewritePattern<ReinterpretCastOp> {
2344public:
2345 using OpRewritePattern<ReinterpretCastOp>::OpRewritePattern;
2346
2347 LogicalResult matchAndRewrite(ReinterpretCastOp op,
2348 PatternRewriter &rewriter) const override {
2349 unsigned srcStaticCount = llvm::count_if(
2350 llvm::concat<OpFoldResult>(op.getMixedOffsets(), op.getMixedSizes(),
2351 op.getMixedStrides()),
2352 [](OpFoldResult ofr) { return isa<Attribute>(ofr); });
2353
2354 SmallVector<OpFoldResult> offsets = {op.getConstifiedMixedOffset()};
2355 SmallVector<OpFoldResult> sizes = op.getConstifiedMixedSizes();
2356 SmallVector<OpFoldResult> strides = op.getConstifiedMixedStrides();
2357
2358 // If the offset is a negative constant, we can't fold it because the
2359 // resulting memref type would be invalid. In that case, we keep the
2360 // original offset.
2361 if (auto cst = getConstantIntValue(offsets[0]))
2362 if (*cst < 0)
2363 offsets[0] = op.getMixedOffsets()[0];
2364
2365 // If the size is a negative constant, we can't fold it because the
2366 // resulting memref type would be invalid. In that case, we keep the
2367 // original size.
2368 for (auto it : llvm::zip(op.getMixedSizes(), sizes)) {
2369 auto &srcSizeOfr = std::get<0>(it);
2370 auto &sizeOfr = std::get<1>(it);
2371 if (auto cst = getConstantIntValue(sizeOfr))
2372 if (*cst < 0)
2373 sizeOfr = srcSizeOfr;
2374 }
2375
2376 // TODO: Using counting comparison instead of direct comparison because
2377 // getMixedValues (and therefore ReinterpretCastOp::getMixed...) returns
2378 // IntegerAttrs, while constifyIndexValues (and therefore
2379 // ReinterpretCastOp::getConstifiedMixed...) returns IndexAttrs.
2380 if (srcStaticCount ==
2381 llvm::count_if(llvm::concat<OpFoldResult>(offsets, sizes, strides),
2382 [](OpFoldResult ofr) { return isa<Attribute>(ofr); }))
2383 return failure();
2384
2385 auto newReinterpretCast = ReinterpretCastOp::create(
2386 rewriter, op->getLoc(), op.getSource(), offsets[0], sizes, strides);
2387
2388 rewriter.replaceOpWithNewOp<CastOp>(op, op.getType(), newReinterpretCast);
2389 return success();
2390 }
2391};
2392} // namespace
2393
2394void ReinterpretCastOp::getCanonicalizationPatterns(RewritePatternSet &results,
2395 MLIRContext *context) {
2396 results.add<ReinterpretCastOpExtractStridedMetadataFolder,
2397 ReinterpretCastOpConstantFolder>(context);
2398}
2399
2400FailureOr<std::optional<SmallVector<Value>>>
2401ReinterpretCastOp::bubbleDownCasts(OpBuilder &builder) {
2402 return bubbleDownCastsPassthroughOpImpl(*this, builder, getSourceMutable());
2403}
2404
2405//===----------------------------------------------------------------------===//
2406// Reassociative reshape ops
2407//===----------------------------------------------------------------------===//
2408
2409void CollapseShapeOp::getAsmResultNames(
2410 function_ref<void(Value, StringRef)> setNameFn) {
2411 setNameFn(getResult(), "collapse_shape");
2412}
2413
2414void ExpandShapeOp::getAsmResultNames(
2415 function_ref<void(Value, StringRef)> setNameFn) {
2416 setNameFn(getResult(), "expand_shape");
2417}
2418
2419LogicalResult ExpandShapeOp::reifyResultShapes(
2420 OpBuilder &builder, ReifiedRankedShapedTypeDims &reifiedResultShapes) {
2421 reifiedResultShapes = {
2422 getMixedValues(getStaticOutputShape(), getOutputShape(), builder)};
2423 return success();
2424}
2425
2426/// Helper function for verifying the shape of ExpandShapeOp and ResultShapeOp
2427/// result and operand. Layout maps are verified separately.
2428///
2429/// If `allowMultipleDynamicDimsPerGroup`, multiple dynamic dimensions are
2430/// allowed in a reassocation group.
2431static LogicalResult
2433 ArrayRef<int64_t> expandedShape,
2434 ArrayRef<ReassociationIndices> reassociation,
2435 bool allowMultipleDynamicDimsPerGroup) {
2436 // There must be one reassociation group per collapsed dimension.
2437 if (collapsedShape.size() != reassociation.size())
2438 return op->emitOpError("invalid number of reassociation groups: found ")
2439 << reassociation.size() << ", expected " << collapsedShape.size();
2440
2441 // The next expected expanded dimension index (while iterating over
2442 // reassociation indices).
2443 int64_t nextDim = 0;
2444 for (const auto &it : llvm::enumerate(reassociation)) {
2445 ReassociationIndices group = it.value();
2446 int64_t collapsedDim = it.index();
2447
2448 bool foundDynamic = false;
2449 for (int64_t expandedDim : group) {
2450 if (expandedDim != nextDim++)
2451 return op->emitOpError("reassociation indices must be contiguous");
2452
2453 if (expandedDim >= static_cast<int64_t>(expandedShape.size()))
2454 return op->emitOpError("reassociation index ")
2455 << expandedDim << " is out of bounds";
2456
2457 // Check if there are multiple dynamic dims in a reassociation group.
2458 if (ShapedType::isDynamic(expandedShape[expandedDim])) {
2459 if (foundDynamic && !allowMultipleDynamicDimsPerGroup)
2460 return op->emitOpError(
2461 "at most one dimension in a reassociation group may be dynamic");
2462 foundDynamic = true;
2463 }
2464 }
2465
2466 // ExpandShapeOp/CollapseShapeOp may not be used to cast dynamicity.
2467 if (ShapedType::isDynamic(collapsedShape[collapsedDim]) != foundDynamic)
2468 return op->emitOpError("collapsed dim (")
2469 << collapsedDim
2470 << ") must be dynamic if and only if reassociation group is "
2471 "dynamic";
2472
2473 // If all dims in the reassociation group are static, the size of the
2474 // collapsed dim can be verified.
2475 if (!foundDynamic) {
2476 int64_t groupSize = 1;
2477 for (int64_t expandedDim : group)
2478 groupSize *= expandedShape[expandedDim];
2479 if (groupSize != collapsedShape[collapsedDim])
2480 return op->emitOpError("collapsed dim size (")
2481 << collapsedShape[collapsedDim]
2482 << ") must equal reassociation group size (" << groupSize << ")";
2483 }
2484 }
2485
2486 if (collapsedShape.empty()) {
2487 // Rank 0: All expanded dimensions must be 1.
2488 for (int64_t d : expandedShape)
2489 if (d != 1)
2490 return op->emitOpError(
2491 "rank 0 memrefs can only be extended/collapsed with/from ones");
2492 } else if (nextDim != static_cast<int64_t>(expandedShape.size())) {
2493 // Rank >= 1: Number of dimensions among all reassociation groups must match
2494 // the result memref rank.
2495 return op->emitOpError("expanded rank (")
2496 << expandedShape.size()
2497 << ") inconsistent with number of reassociation indices (" << nextDim
2498 << ")";
2499 }
2500
2501 return success();
2502}
2503
2504SmallVector<AffineMap, 4> CollapseShapeOp::getReassociationMaps() {
2505 return getSymbolLessAffineMaps(getReassociationExprs());
2506}
2507
2508SmallVector<ReassociationExprs, 4> CollapseShapeOp::getReassociationExprs() {
2510 getReassociationIndices());
2511}
2512
2513SmallVector<AffineMap, 4> ExpandShapeOp::getReassociationMaps() {
2514 return getSymbolLessAffineMaps(getReassociationExprs());
2515}
2516
2517SmallVector<ReassociationExprs, 4> ExpandShapeOp::getReassociationExprs() {
2519 getReassociationIndices());
2520}
2521
2522/// Compute the layout map after expanding a given source MemRef type with the
2523/// specified reassociation indices.
2524static FailureOr<StridedLayoutAttr>
2525computeExpandedLayoutMap(MemRefType srcType, ArrayRef<int64_t> resultShape,
2526 ArrayRef<ReassociationIndices> reassociation) {
2527 int64_t srcOffset;
2528 SmallVector<int64_t> srcStrides;
2529 if (failed(srcType.getStridesAndOffset(srcStrides, srcOffset)))
2530 return failure();
2531 assert(srcStrides.size() == reassociation.size() && "invalid reassociation");
2532
2533 // 1-1 mapping between srcStrides and reassociation packs.
2534 // Each srcStride starts with the given value and gets expanded according to
2535 // the proper entries in resultShape.
2536 // Example:
2537 // srcStrides = [10000, 1 , 100 ],
2538 // reassociations = [ [0], [1], [2, 3, 4]],
2539 // resultSizes = [2, 5, 4, 3, 2] = [ [2], [5], [4, 3, 2]]
2540 // -> For the purpose of stride calculation, the useful sizes are:
2541 // [x, x, x, 3, 2] = [ [x], [x], [x, 3, 2]].
2542 // resultStrides = [10000, 1, 600, 200, 100]
2543 // Note that a stride does not get expanded along the first entry of each
2544 // shape pack.
2545 SmallVector<int64_t> reverseResultStrides;
2546 reverseResultStrides.reserve(resultShape.size());
2547 unsigned shapeIndex = resultShape.size() - 1;
2548 for (auto it : llvm::reverse(llvm::zip(reassociation, srcStrides))) {
2549 ReassociationIndices reassoc = std::get<0>(it);
2550 int64_t currentStrideToExpand = std::get<1>(it);
2551 for (unsigned idx = 0, e = reassoc.size(); idx < e; ++idx) {
2552 reverseResultStrides.push_back(currentStrideToExpand);
2553 currentStrideToExpand =
2554 (SaturatedInteger::wrap(currentStrideToExpand) *
2555 SaturatedInteger::wrap(resultShape[shapeIndex--]))
2556 .asInteger();
2557 }
2558 }
2559 auto resultStrides = llvm::to_vector<8>(llvm::reverse(reverseResultStrides));
2560 resultStrides.resize(resultShape.size(), 1);
2561 return StridedLayoutAttr::get(srcType.getContext(), srcOffset, resultStrides);
2562}
2563
2564FailureOr<MemRefType> ExpandShapeOp::computeExpandedType(
2565 MemRefType srcType, ArrayRef<int64_t> resultShape,
2566 ArrayRef<ReassociationIndices> reassociation) {
2567 if (srcType.getLayout().isIdentity()) {
2568 // If the source is contiguous (i.e., no layout map specified), so is the
2569 // result.
2570 MemRefLayoutAttrInterface layout;
2571 return MemRefType::get(resultShape, srcType.getElementType(), layout,
2572 srcType.getMemorySpace());
2573 }
2574
2575 // Source may not be contiguous. Compute the layout map.
2576 FailureOr<StridedLayoutAttr> computedLayout =
2577 computeExpandedLayoutMap(srcType, resultShape, reassociation);
2578 if (failed(computedLayout))
2579 return failure();
2580 return MemRefType::get(resultShape, srcType.getElementType(), *computedLayout,
2581 srcType.getMemorySpace());
2582}
2583
2584FailureOr<SmallVector<OpFoldResult>>
2585ExpandShapeOp::inferOutputShape(OpBuilder &b, Location loc,
2586 MemRefType expandedType,
2587 ArrayRef<ReassociationIndices> reassociation,
2588 ArrayRef<OpFoldResult> inputShape) {
2589 std::optional<SmallVector<OpFoldResult>> outputShape =
2590 inferExpandShapeOutputShape(b, loc, expandedType, reassociation,
2591 inputShape);
2592 if (!outputShape)
2593 return failure();
2594 return *outputShape;
2595}
2596
2597void ExpandShapeOp::build(OpBuilder &builder, OperationState &result,
2598 Type resultType, Value src,
2599 ArrayRef<ReassociationIndices> reassociation,
2600 ArrayRef<OpFoldResult> outputShape) {
2601 auto [staticOutputShape, dynamicOutputShape] =
2602 decomposeMixedValues(SmallVector<OpFoldResult>(outputShape));
2603 build(builder, result, llvm::cast<MemRefType>(resultType), src,
2604 getReassociationIndicesAttribute(builder, reassociation),
2605 dynamicOutputShape, staticOutputShape);
2606}
2607
2608void ExpandShapeOp::build(OpBuilder &builder, OperationState &result,
2609 Type resultType, Value src,
2610 ArrayRef<ReassociationIndices> reassociation) {
2611 SmallVector<OpFoldResult> inputShape =
2612 getMixedSizes(builder, result.location, src);
2613 MemRefType memrefResultTy = llvm::cast<MemRefType>(resultType);
2614 FailureOr<SmallVector<OpFoldResult>> outputShape = inferOutputShape(
2615 builder, result.location, memrefResultTy, reassociation, inputShape);
2616 // Failure of this assertion usually indicates presence of multiple
2617 // dynamic dimensions in the same reassociation group.
2618 assert(succeeded(outputShape) && "unable to infer output shape");
2619 build(builder, result, memrefResultTy, src, reassociation, *outputShape);
2620}
2621
2622void ExpandShapeOp::build(OpBuilder &builder, OperationState &result,
2623 ArrayRef<int64_t> resultShape, Value src,
2624 ArrayRef<ReassociationIndices> reassociation) {
2625 // Only ranked memref source values are supported.
2626 auto srcType = llvm::cast<MemRefType>(src.getType());
2627 FailureOr<MemRefType> resultType =
2628 ExpandShapeOp::computeExpandedType(srcType, resultShape, reassociation);
2629 // Failure of this assertion usually indicates a problem with the source
2630 // type, e.g., could not get strides/offset.
2631 assert(succeeded(resultType) && "could not compute layout");
2632 build(builder, result, *resultType, src, reassociation);
2633}
2634
2635void ExpandShapeOp::build(OpBuilder &builder, OperationState &result,
2636 ArrayRef<int64_t> resultShape, Value src,
2637 ArrayRef<ReassociationIndices> reassociation,
2638 ArrayRef<OpFoldResult> outputShape) {
2639 // Only ranked memref source values are supported.
2640 auto srcType = llvm::cast<MemRefType>(src.getType());
2641 FailureOr<MemRefType> resultType =
2642 ExpandShapeOp::computeExpandedType(srcType, resultShape, reassociation);
2643 // Failure of this assertion usually indicates a problem with the source
2644 // type, e.g., could not get strides/offset.
2645 assert(succeeded(resultType) && "could not compute layout");
2646 build(builder, result, *resultType, src, reassociation, outputShape);
2647}
2648
2649LogicalResult ExpandShapeOp::verify() {
2651 return failure();
2652
2653 MemRefType srcType = getSrcType();
2654 MemRefType resultType = getResultType();
2655
2656 if (srcType.getRank() > resultType.getRank()) {
2657 auto r0 = srcType.getRank();
2658 auto r1 = resultType.getRank();
2659 return emitOpError("has source rank ")
2660 << r0 << " and result rank " << r1 << ". This is not an expansion ("
2661 << r0 << " > " << r1 << ").";
2662 }
2663
2664 // Verify result shape.
2665 if (failed(verifyCollapsedShape(getOperation(), srcType.getShape(),
2666 resultType.getShape(),
2667 getReassociationIndices(),
2668 /*allowMultipleDynamicDimsPerGroup=*/true)))
2669 return failure();
2670
2671 // Compute expected result type (including layout map).
2672 FailureOr<MemRefType> expectedResultType = ExpandShapeOp::computeExpandedType(
2673 srcType, resultType.getShape(), getReassociationIndices());
2674 if (failed(expectedResultType))
2675 return emitOpError("invalid source layout map");
2676
2677 // Check actual result type.
2678 if (*expectedResultType != resultType)
2679 return emitOpError("expected expanded type to be ")
2680 << *expectedResultType << " but found " << resultType;
2681
2682 if ((int64_t)getStaticOutputShape().size() != resultType.getRank())
2683 return emitOpError("expected number of static shape bounds to be equal to "
2684 "the output rank (")
2685 << resultType.getRank() << ") but found "
2686 << getStaticOutputShape().size() << " inputs instead";
2687
2688 if ((int64_t)getOutputShape().size() !=
2689 llvm::count(getStaticOutputShape(), ShapedType::kDynamic))
2690 return emitOpError("mismatch in dynamic dims in output_shape and "
2691 "static_output_shape: static_output_shape has ")
2692 << llvm::count(getStaticOutputShape(), ShapedType::kDynamic)
2693 << " dynamic dims while output_shape has " << getOutputShape().size()
2694 << " values";
2695
2696 // Verify that the number of dynamic dims in output_shape matches the number
2697 // of dynamic dims in the result type.
2698 if (failed(verifyDynamicDimensionCount(getOperation(), resultType,
2699 getOutputShape())))
2700 return failure();
2701
2702 // Verify if provided output shapes are in agreement with output type.
2703 DenseI64ArrayAttr staticOutputShapes = getStaticOutputShapeAttr();
2704 ArrayRef<int64_t> resShape = getResult().getType().getShape();
2705 for (auto [pos, shape] : llvm::enumerate(resShape)) {
2706 if (ShapedType::isStatic(shape) && shape != staticOutputShapes[pos]) {
2707 return emitOpError("invalid output shape provided at pos ") << pos;
2708 }
2709 }
2710
2711 return success();
2712}
2713
2714struct ExpandShapeOpMemRefCastFolder : public OpRewritePattern<ExpandShapeOp> {
2715public:
2716 using OpRewritePattern<ExpandShapeOp>::OpRewritePattern;
2717
2718 LogicalResult matchAndRewrite(ExpandShapeOp op,
2719 PatternRewriter &rewriter) const override {
2720 auto cast = op.getSrc().getDefiningOp<CastOp>();
2721 if (!cast)
2722 return failure();
2723
2724 if (!CastOp::canFoldIntoConsumerOp(cast))
2725 return failure();
2726
2727 SmallVector<OpFoldResult> originalOutputShape = op.getMixedOutputShape();
2728 SmallVector<OpFoldResult> newOutputShape = originalOutputShape;
2729 SmallVector<int64_t> newOutputShapeSizes;
2730
2731 // Convert output shape dims from dynamic to static where possible.
2732 for (auto [dimIdx, dimSize] : enumerate(originalOutputShape)) {
2733 std::optional<int64_t> sizeOpt = getConstantIntValue(dimSize);
2734 if (!sizeOpt.has_value()) {
2735 newOutputShapeSizes.push_back(ShapedType::kDynamic);
2736 continue;
2737 }
2738
2739 newOutputShapeSizes.push_back(sizeOpt.value());
2740 newOutputShape[dimIdx] = rewriter.getIndexAttr(sizeOpt.value());
2741 }
2742
2743 Value castSource = cast.getSource();
2744 auto castSourceType = llvm::cast<MemRefType>(castSource.getType());
2745 SmallVector<ReassociationIndices> reassociationIndices =
2746 op.getReassociationIndices();
2747 for (auto [idx, group] : llvm::enumerate(reassociationIndices)) {
2748 auto newOutputShapeSizesSlice =
2749 ArrayRef(newOutputShapeSizes).slice(group.front(), group.size());
2750 bool newOutputDynamic =
2751 llvm::is_contained(newOutputShapeSizesSlice, ShapedType::kDynamic);
2752 if (castSourceType.isDynamicDim(idx) != newOutputDynamic)
2753 return rewriter.notifyMatchFailure(
2754 op, "folding cast will result in changing dynamicity in "
2755 "reassociation group");
2756 }
2757
2758 FailureOr<MemRefType> newResultTypeOrFailure =
2759 ExpandShapeOp::computeExpandedType(castSourceType, newOutputShapeSizes,
2760 reassociationIndices);
2761
2762 if (failed(newResultTypeOrFailure))
2763 return rewriter.notifyMatchFailure(
2764 op, "could not compute new expanded type after folding cast");
2765
2766 if (*newResultTypeOrFailure == op.getResultType()) {
2767 rewriter.modifyOpInPlace(
2768 op, [&]() { op.getSrcMutable().assign(castSource); });
2769 } else {
2770 Value newOp = ExpandShapeOp::create(rewriter, op->getLoc(),
2771 *newResultTypeOrFailure, castSource,
2772 reassociationIndices, newOutputShape);
2773 rewriter.replaceOpWithNewOp<CastOp>(op, op.getType(), newOp);
2774 }
2775 return success();
2776 }
2777};
2778
2779void ExpandShapeOp::getCanonicalizationPatterns(RewritePatternSet &results,
2780 MLIRContext *context) {
2781 results.add<
2782 ComposeReassociativeReshapeOps<ExpandShapeOp, ReshapeOpKind::kExpand>,
2783 ComposeExpandOfCollapseOp<ExpandShapeOp, CollapseShapeOp, CastOp>,
2784 ExpandShapeOpMemRefCastFolder>(context);
2785}
2786
2787FailureOr<std::optional<SmallVector<Value>>>
2788ExpandShapeOp::bubbleDownCasts(OpBuilder &builder) {
2789 return bubbleDownCastsPassthroughOpImpl(*this, builder, getSrcMutable());
2790}
2791
2792/// Compute the layout map after collapsing a given source MemRef type with the
2793/// specified reassociation indices.
2794///
2795/// Note: All collapsed dims in a reassociation group must be contiguous. It is
2796/// not possible to check this by inspecting a MemRefType in the general case.
2797/// If non-contiguity cannot be checked statically, the collapse is assumed to
2798/// be valid (and thus accepted by this function) unless `strict = true`.
2799static FailureOr<StridedLayoutAttr>
2800computeCollapsedLayoutMap(MemRefType srcType,
2801 ArrayRef<ReassociationIndices> reassociation,
2802 bool strict = false) {
2803 int64_t srcOffset;
2804 SmallVector<int64_t> srcStrides;
2805 auto srcShape = srcType.getShape();
2806 if (failed(srcType.getStridesAndOffset(srcStrides, srcOffset)))
2807 return failure();
2808
2809 // The result stride of a reassociation group is the stride of the last entry
2810 // of the reassociation. (TODO: Should be the minimum stride in the
2811 // reassociation because strides are not necessarily sorted. E.g., when using
2812 // memref.transpose.) Dimensions of size 1 should be skipped, because their
2813 // strides are meaningless and could have any arbitrary value.
2814 SmallVector<int64_t> resultStrides;
2815 resultStrides.reserve(reassociation.size());
2816 for (const ReassociationIndices &reassoc : reassociation) {
2817 ArrayRef<int64_t> ref = llvm::ArrayRef(reassoc);
2818 while (srcShape[ref.back()] == 1 && ref.size() > 1)
2819 ref = ref.drop_back();
2820 if (ShapedType::isStatic(srcShape[ref.back()]) || ref.size() == 1) {
2821 resultStrides.push_back(srcStrides[ref.back()]);
2822 } else {
2823 // Dynamically-sized dims may turn out to be dims of size 1 at runtime, so
2824 // the corresponding stride may have to be skipped. (See above comment.)
2825 // Therefore, the result stride cannot be statically determined and must
2826 // be dynamic.
2827 resultStrides.push_back(ShapedType::kDynamic);
2828 }
2829 }
2830
2831 // Validate that each reassociation group is contiguous.
2832 unsigned resultStrideIndex = resultStrides.size() - 1;
2833 for (const ReassociationIndices &reassoc : llvm::reverse(reassociation)) {
2834 auto trailingReassocs = ArrayRef<int64_t>(reassoc).drop_front();
2835 auto stride = SaturatedInteger::wrap(resultStrides[resultStrideIndex--]);
2836 for (int64_t idx : llvm::reverse(trailingReassocs)) {
2837 stride = stride * SaturatedInteger::wrap(srcShape[idx]);
2838
2839 // Dimensions of size 1 should be skipped, because their strides are
2840 // meaningless and could have any arbitrary value.
2841 if (srcShape[idx - 1] == 1)
2842 continue;
2843
2844 // Both source and result stride must have the same static value. In that
2845 // case, we can be sure, that the dimensions are collapsible (because they
2846 // are contiguous).
2847 // If `strict = false` (default during op verification), we accept cases
2848 // where one or both strides are dynamic. This is best effort: We reject
2849 // ops where obviously non-contiguous dims are collapsed, but accept ops
2850 // where we cannot be sure statically. Such ops may fail at runtime. See
2851 // the op documentation for details.
2852 auto srcStride = SaturatedInteger::wrap(srcStrides[idx - 1]);
2853 if (strict && (stride.saturated || srcStride.saturated))
2854 return failure();
2855
2856 if (!stride.saturated && !srcStride.saturated && stride != srcStride)
2857 return failure();
2858 }
2859 }
2860 return StridedLayoutAttr::get(srcType.getContext(), srcOffset, resultStrides);
2861}
2862
2863bool CollapseShapeOp::isGuaranteedCollapsible(
2864 MemRefType srcType, ArrayRef<ReassociationIndices> reassociation) {
2865 // MemRefs with identity layout are always collapsible.
2866 if (srcType.getLayout().isIdentity())
2867 return true;
2868
2869 return succeeded(computeCollapsedLayoutMap(srcType, reassociation,
2870 /*strict=*/true));
2871}
2872
2873MemRefType CollapseShapeOp::computeCollapsedType(
2874 MemRefType srcType, ArrayRef<ReassociationIndices> reassociation) {
2875 SmallVector<int64_t> resultShape;
2876 resultShape.reserve(reassociation.size());
2877 for (const ReassociationIndices &group : reassociation) {
2878 auto groupSize = SaturatedInteger::wrap(1);
2879 for (int64_t srcDim : group)
2880 groupSize =
2881 groupSize * SaturatedInteger::wrap(srcType.getDimSize(srcDim));
2882 resultShape.push_back(groupSize.asInteger());
2883 }
2884
2885 if (srcType.getLayout().isIdentity()) {
2886 // If the source is contiguous (i.e., no layout map specified), so is the
2887 // result.
2888 MemRefLayoutAttrInterface layout;
2889 return MemRefType::get(resultShape, srcType.getElementType(), layout,
2890 srcType.getMemorySpace());
2891 }
2892
2893 // Source may not be fully contiguous. Compute the layout map.
2894 // Note: Dimensions that are collapsed into a single dim are assumed to be
2895 // contiguous.
2896 FailureOr<StridedLayoutAttr> computedLayout =
2897 computeCollapsedLayoutMap(srcType, reassociation);
2898 assert(succeeded(computedLayout) &&
2899 "invalid source layout map or collapsing non-contiguous dims");
2900 return MemRefType::get(resultShape, srcType.getElementType(), *computedLayout,
2901 srcType.getMemorySpace());
2902}
2903
2904void CollapseShapeOp::build(OpBuilder &b, OperationState &result, Value src,
2905 ArrayRef<ReassociationIndices> reassociation,
2906 ArrayRef<NamedAttribute> attrs) {
2907 auto srcType = llvm::cast<MemRefType>(src.getType());
2908 MemRefType resultType =
2909 CollapseShapeOp::computeCollapsedType(srcType, reassociation);
2910 buildPropertiesAndDiscardableAttributes(result, attrs);
2911 result.getOrAddProperties<Properties>().reassociation =
2912 getReassociationIndicesAttribute(b, reassociation);
2913 result.addOperands(src);
2914 result.addTypes(resultType);
2915}
2916
2917LogicalResult CollapseShapeOp::verify() {
2919 return failure();
2920
2921 MemRefType srcType = getSrcType();
2922 MemRefType resultType = getResultType();
2923
2924 if (srcType.getRank() < resultType.getRank()) {
2925 auto r0 = srcType.getRank();
2926 auto r1 = resultType.getRank();
2927 return emitOpError("has source rank ")
2928 << r0 << " and result rank " << r1 << ". This is not a collapse ("
2929 << r0 << " < " << r1 << ").";
2930 }
2931
2932 // Verify result shape.
2933 if (failed(verifyCollapsedShape(getOperation(), resultType.getShape(),
2934 srcType.getShape(), getReassociationIndices(),
2935 /*allowMultipleDynamicDimsPerGroup=*/true)))
2936 return failure();
2937
2938 // Compute expected result type (including layout map).
2939 MemRefType expectedResultType;
2940 if (srcType.getLayout().isIdentity()) {
2941 // If the source is contiguous (i.e., no layout map specified), so is the
2942 // result.
2943 MemRefLayoutAttrInterface layout;
2944 expectedResultType =
2945 MemRefType::get(resultType.getShape(), srcType.getElementType(), layout,
2946 srcType.getMemorySpace());
2947 } else {
2948 // Source may not be fully contiguous. Compute the layout map.
2949 // Note: Dimensions that are collapsed into a single dim are assumed to be
2950 // contiguous.
2951 FailureOr<StridedLayoutAttr> computedLayout =
2952 computeCollapsedLayoutMap(srcType, getReassociationIndices());
2953 if (failed(computedLayout))
2954 return emitOpError(
2955 "invalid source layout map or collapsing non-contiguous dims");
2956 expectedResultType =
2957 MemRefType::get(resultType.getShape(), srcType.getElementType(),
2958 *computedLayout, srcType.getMemorySpace());
2959 }
2960
2961 if (expectedResultType != resultType)
2962 return emitOpError("expected collapsed type to be ")
2963 << expectedResultType << " but found " << resultType;
2964
2965 return success();
2966}
2967
2969 : public OpRewritePattern<CollapseShapeOp> {
2970public:
2971 using OpRewritePattern<CollapseShapeOp>::OpRewritePattern;
2972
2973 LogicalResult matchAndRewrite(CollapseShapeOp op,
2974 PatternRewriter &rewriter) const override {
2975 auto cast = op.getOperand().getDefiningOp<CastOp>();
2976 if (!cast)
2977 return failure();
2978
2979 if (!CastOp::canFoldIntoConsumerOp(cast))
2980 return failure();
2981
2982 Type newResultType = CollapseShapeOp::computeCollapsedType(
2983 llvm::cast<MemRefType>(cast.getOperand().getType()),
2984 op.getReassociationIndices());
2985
2986 if (newResultType == op.getResultType()) {
2987 rewriter.modifyOpInPlace(
2988 op, [&]() { op.getSrcMutable().assign(cast.getSource()); });
2989 } else {
2990 Value newOp =
2991 CollapseShapeOp::create(rewriter, op->getLoc(), cast.getSource(),
2992 op.getReassociationIndices());
2993 rewriter.replaceOpWithNewOp<CastOp>(op, op.getType(), newOp);
2994 }
2995 return success();
2996 }
2997};
2998
2999void CollapseShapeOp::getCanonicalizationPatterns(RewritePatternSet &results,
3000 MLIRContext *context) {
3001 results.add<
3002 ComposeReassociativeReshapeOps<CollapseShapeOp, ReshapeOpKind::kCollapse>,
3003 ComposeCollapseOfExpandOp<CollapseShapeOp, ExpandShapeOp, CastOp,
3004 memref::DimOp, MemRefType>,
3005 CollapseShapeOpMemRefCastFolder>(context);
3006}
3007
3008OpFoldResult ExpandShapeOp::fold(FoldAdaptor adaptor) {
3010 adaptor.getOperands());
3011}
3012
3013OpFoldResult CollapseShapeOp::fold(FoldAdaptor adaptor) {
3015 adaptor.getOperands());
3016}
3017
3018FailureOr<std::optional<SmallVector<Value>>>
3019CollapseShapeOp::bubbleDownCasts(OpBuilder &builder) {
3020 return bubbleDownCastsPassthroughOpImpl(*this, builder, getSrcMutable());
3021}
3022
3023//===----------------------------------------------------------------------===//
3024// ReshapeOp
3025//===----------------------------------------------------------------------===//
3026
3027void ReshapeOp::getAsmResultNames(
3028 function_ref<void(Value, StringRef)> setNameFn) {
3029 setNameFn(getResult(), "reshape");
3030}
3031
3032LogicalResult ReshapeOp::verify() {
3033 Type operandType = getSource().getType();
3034 Type resultType = getResult().getType();
3035
3036 Type operandElementType =
3037 llvm::cast<ShapedType>(operandType).getElementType();
3038 Type resultElementType = llvm::cast<ShapedType>(resultType).getElementType();
3039 if (operandElementType != resultElementType)
3040 return emitOpError("element types of source and destination memref "
3041 "types should be the same");
3042
3043 if (auto operandMemRefType = llvm::dyn_cast<MemRefType>(operandType))
3044 if (!operandMemRefType.getLayout().isIdentity())
3045 return emitOpError("source memref type should have identity affine map");
3046
3047 int64_t shapeSize =
3048 llvm::cast<MemRefType>(getShape().getType()).getDimSize(0);
3049 auto resultMemRefType = llvm::dyn_cast<MemRefType>(resultType);
3050 if (resultMemRefType) {
3051 if (!resultMemRefType.getLayout().isIdentity())
3052 return emitOpError("result memref type should have identity affine map");
3053 if (shapeSize == ShapedType::kDynamic)
3054 return emitOpError("cannot use shape operand with dynamic length to "
3055 "reshape to statically-ranked memref type");
3056 if (shapeSize != resultMemRefType.getRank())
3057 return emitOpError(
3058 "length of shape operand differs from the result's memref rank");
3059 }
3060 return success();
3061}
3062
3063FailureOr<std::optional<SmallVector<Value>>>
3064ReshapeOp::bubbleDownCasts(OpBuilder &builder) {
3065 return bubbleDownCastsPassthroughOpImpl(*this, builder, getSourceMutable());
3066}
3067
3068//===----------------------------------------------------------------------===//
3069// StoreOp
3070//===----------------------------------------------------------------------===//
3071
3072LogicalResult StoreOp::fold(FoldAdaptor adaptor,
3073 SmallVectorImpl<OpFoldResult> &results) {
3074 /// store(memrefcast) -> store
3075 return foldMemRefCast(*this, getValueToStore());
3076}
3077
3078TypedValue<MemRefType> StoreOp::getAccessedMemref() { return getMemref(); }
3079
3080std::optional<SmallVector<Value>>
3081StoreOp::updateMemrefAndIndices(RewriterBase &rewriter, Value newMemref,
3082 ValueRange newIndices) {
3083 rewriter.modifyOpInPlace(*this, [&]() {
3084 getMemrefMutable().assign(newMemref);
3085 getIndicesMutable().assign(newIndices);
3086 });
3087 return std::nullopt;
3088}
3089
3090FailureOr<std::optional<SmallVector<Value>>>
3091StoreOp::bubbleDownCasts(OpBuilder &builder) {
3093 ValueRange());
3094}
3095
3096//===----------------------------------------------------------------------===//
3097// SubViewOp
3098//===----------------------------------------------------------------------===//
3099
3100void SubViewOp::getAsmResultNames(
3101 function_ref<void(Value, StringRef)> setNameFn) {
3102 setNameFn(getResult(), "subview");
3103}
3104
3105/// A subview result type can be fully inferred from the source type and the
3106/// static representation of offsets, sizes and strides. Special sentinels
3107/// encode the dynamic case.
3108MemRefType SubViewOp::inferResultType(MemRefType sourceMemRefType,
3109 ArrayRef<int64_t> staticOffsets,
3110 ArrayRef<int64_t> staticSizes,
3111 ArrayRef<int64_t> staticStrides) {
3112 unsigned rank = sourceMemRefType.getRank();
3113 (void)rank;
3114 assert(staticOffsets.size() == rank && "staticOffsets length mismatch");
3115 assert(staticSizes.size() == rank && "staticSizes length mismatch");
3116 assert(staticStrides.size() == rank && "staticStrides length mismatch");
3117
3118 // Extract source offset and strides.
3119 auto [sourceStrides, sourceOffset] = sourceMemRefType.getStridesAndOffset();
3120
3121 // Compute target offset whose value is:
3122 // `sourceOffset + sum_i(staticOffset_i * sourceStrides_i)`.
3123 int64_t targetOffset = sourceOffset;
3124 for (auto it : llvm::zip(staticOffsets, sourceStrides)) {
3125 auto staticOffset = std::get<0>(it), sourceStride = std::get<1>(it);
3126 targetOffset = (SaturatedInteger::wrap(targetOffset) +
3127 SaturatedInteger::wrap(staticOffset) *
3128 SaturatedInteger::wrap(sourceStride))
3129 .asInteger();
3130 }
3131
3132 // Compute target stride whose value is:
3133 // `sourceStrides_i * staticStrides_i`.
3134 SmallVector<int64_t, 4> targetStrides;
3135 targetStrides.reserve(staticOffsets.size());
3136 for (auto it : llvm::zip(sourceStrides, staticStrides)) {
3137 auto sourceStride = std::get<0>(it), staticStride = std::get<1>(it);
3138 targetStrides.push_back((SaturatedInteger::wrap(sourceStride) *
3139 SaturatedInteger::wrap(staticStride))
3140 .asInteger());
3141 }
3142
3143 // The type is now known.
3144 return MemRefType::get(staticSizes, sourceMemRefType.getElementType(),
3145 StridedLayoutAttr::get(sourceMemRefType.getContext(),
3146 targetOffset, targetStrides),
3147 sourceMemRefType.getMemorySpace());
3148}
3149
3150MemRefType SubViewOp::inferResultType(MemRefType sourceMemRefType,
3151 ArrayRef<OpFoldResult> offsets,
3152 ArrayRef<OpFoldResult> sizes,
3153 ArrayRef<OpFoldResult> strides) {
3154 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides;
3155 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides;
3156 dispatchIndexOpFoldResults(offsets, dynamicOffsets, staticOffsets);
3157 dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes);
3158 dispatchIndexOpFoldResults(strides, dynamicStrides, staticStrides);
3159 if (!hasValidSizesOffsets(staticOffsets))
3160 return {};
3161 if (!hasValidSizesOffsets(staticSizes))
3162 return {};
3163 if (!hasValidStrides(staticStrides))
3164 return {};
3165 return SubViewOp::inferResultType(sourceMemRefType, staticOffsets,
3166 staticSizes, staticStrides);
3167}
3168
3169MemRefType SubViewOp::inferRankReducedResultType(
3170 ArrayRef<int64_t> resultShape, MemRefType sourceRankedTensorType,
3171 ArrayRef<int64_t> offsets, ArrayRef<int64_t> sizes,
3172 ArrayRef<int64_t> strides) {
3173 MemRefType inferredType =
3174 inferResultType(sourceRankedTensorType, offsets, sizes, strides);
3175 assert(inferredType.getRank() >= static_cast<int64_t>(resultShape.size()) &&
3176 "expected ");
3177 if (inferredType.getRank() == static_cast<int64_t>(resultShape.size()))
3178 return inferredType;
3179
3180 // Compute which dimensions are dropped.
3181 std::optional<llvm::SmallDenseSet<unsigned>> dimsToProject =
3182 computeRankReductionMask(inferredType.getShape(), resultShape);
3183 assert(dimsToProject.has_value() && "invalid rank reduction");
3184
3185 // Compute the layout and result type.
3186 auto inferredLayout = llvm::cast<StridedLayoutAttr>(inferredType.getLayout());
3187 SmallVector<int64_t> rankReducedStrides;
3188 rankReducedStrides.reserve(resultShape.size());
3189 for (auto [idx, value] : llvm::enumerate(inferredLayout.getStrides())) {
3190 if (!dimsToProject->contains(idx))
3191 rankReducedStrides.push_back(value);
3192 }
3193 return MemRefType::get(resultShape, inferredType.getElementType(),
3194 StridedLayoutAttr::get(inferredLayout.getContext(),
3195 inferredLayout.getOffset(),
3196 rankReducedStrides),
3197 inferredType.getMemorySpace());
3198}
3199
3200MemRefType SubViewOp::inferRankReducedResultType(
3201 ArrayRef<int64_t> resultShape, MemRefType sourceRankedTensorType,
3202 ArrayRef<OpFoldResult> offsets, ArrayRef<OpFoldResult> sizes,
3203 ArrayRef<OpFoldResult> strides) {
3204 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides;
3205 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides;
3206 dispatchIndexOpFoldResults(offsets, dynamicOffsets, staticOffsets);
3207 dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes);
3208 dispatchIndexOpFoldResults(strides, dynamicStrides, staticStrides);
3209 return SubViewOp::inferRankReducedResultType(
3210 resultShape, sourceRankedTensorType, staticOffsets, staticSizes,
3211 staticStrides);
3212}
3213
3214// Build a SubViewOp with mixed static and dynamic entries and custom result
3215// type. If the type passed is nullptr, it is inferred.
3216void SubViewOp::build(OpBuilder &b, OperationState &result,
3217 MemRefType resultType, Value source,
3218 ArrayRef<OpFoldResult> offsets,
3219 ArrayRef<OpFoldResult> sizes,
3220 ArrayRef<OpFoldResult> strides,
3221 ArrayRef<NamedAttribute> attrs) {
3222 SmallVector<int64_t> staticOffsets, staticSizes, staticStrides;
3223 SmallVector<Value> dynamicOffsets, dynamicSizes, dynamicStrides;
3224 dispatchIndexOpFoldResults(offsets, dynamicOffsets, staticOffsets);
3225 dispatchIndexOpFoldResults(sizes, dynamicSizes, staticSizes);
3226 dispatchIndexOpFoldResults(strides, dynamicStrides, staticStrides);
3227 auto sourceMemRefType = llvm::cast<MemRefType>(source.getType());
3228 // Structuring implementation this way avoids duplication between builders.
3229 if (!resultType) {
3230 resultType = SubViewOp::inferResultType(sourceMemRefType, staticOffsets,
3231 staticSizes, staticStrides);
3232 }
3233 result.addAttributes(attrs);
3234 build(b, result, resultType, source, dynamicOffsets, dynamicSizes,
3235 dynamicStrides, b.getDenseI64ArrayAttr(staticOffsets),
3236 b.getDenseI64ArrayAttr(staticSizes),
3237 b.getDenseI64ArrayAttr(staticStrides));
3238}
3239
3240// Build a SubViewOp with mixed static and dynamic entries and inferred result
3241// type.
3242void SubViewOp::build(OpBuilder &b, OperationState &result, Value source,
3243 ArrayRef<OpFoldResult> offsets,
3244 ArrayRef<OpFoldResult> sizes,
3245 ArrayRef<OpFoldResult> strides,
3246 ArrayRef<NamedAttribute> attrs) {
3247 build(b, result, MemRefType(), source, offsets, sizes, strides, attrs);
3248}
3249
3250// Build a SubViewOp with static entries and inferred result type.
3251void SubViewOp::build(OpBuilder &b, OperationState &result, Value source,
3252 ArrayRef<int64_t> offsets, ArrayRef<int64_t> sizes,
3253 ArrayRef<int64_t> strides,
3254 ArrayRef<NamedAttribute> attrs) {
3255 SmallVector<OpFoldResult> offsetValues =
3256 llvm::map_to_vector<4>(offsets, [&](int64_t v) -> OpFoldResult {
3257 return b.getI64IntegerAttr(v);
3258 });
3259 SmallVector<OpFoldResult> sizeValues = llvm::map_to_vector<4>(
3260 sizes, [&](int64_t v) -> OpFoldResult { return b.getI64IntegerAttr(v); });
3261 SmallVector<OpFoldResult> strideValues =
3262 llvm::map_to_vector<4>(strides, [&](int64_t v) -> OpFoldResult {
3263 return b.getI64IntegerAttr(v);
3264 });
3265 build(b, result, source, offsetValues, sizeValues, strideValues, attrs);
3266}
3267
3268// Build a SubViewOp with dynamic entries and custom result type. If the
3269// type passed is nullptr, it is inferred.
3270void SubViewOp::build(OpBuilder &b, OperationState &result,
3271 MemRefType resultType, Value source,
3272 ArrayRef<int64_t> offsets, ArrayRef<int64_t> sizes,
3273 ArrayRef<int64_t> strides,
3274 ArrayRef<NamedAttribute> attrs) {
3275 SmallVector<OpFoldResult> offsetValues =
3276 llvm::map_to_vector<4>(offsets, [&](int64_t v) -> OpFoldResult {
3277 return b.getI64IntegerAttr(v);
3278 });
3279 SmallVector<OpFoldResult> sizeValues = llvm::map_to_vector<4>(
3280 sizes, [&](int64_t v) -> OpFoldResult { return b.getI64IntegerAttr(v); });
3281 SmallVector<OpFoldResult> strideValues =
3282 llvm::map_to_vector<4>(strides, [&](int64_t v) -> OpFoldResult {
3283 return b.getI64IntegerAttr(v);
3284 });
3285 build(b, result, resultType, source, offsetValues, sizeValues, strideValues,
3286 attrs);
3287}
3288
3289// Build a SubViewOp with dynamic entries and custom result type. If the type
3290// passed is nullptr, it is inferred.
3291void SubViewOp::build(OpBuilder &b, OperationState &result,
3292 MemRefType resultType, Value source, ValueRange offsets,
3293 ValueRange sizes, ValueRange strides,
3294 ArrayRef<NamedAttribute> attrs) {
3295 SmallVector<OpFoldResult> offsetValues = llvm::map_to_vector<4>(
3296 offsets, [](Value v) -> OpFoldResult { return v; });
3297 SmallVector<OpFoldResult> sizeValues =
3298 llvm::map_to_vector<4>(sizes, [](Value v) -> OpFoldResult { return v; });
3299 SmallVector<OpFoldResult> strideValues = llvm::map_to_vector<4>(
3300 strides, [](Value v) -> OpFoldResult { return v; });
3301 build(b, result, resultType, source, offsetValues, sizeValues, strideValues);
3302}
3303
3304// Build a SubViewOp with dynamic entries and inferred result type.
3305void SubViewOp::build(OpBuilder &b, OperationState &result, Value source,
3306 ValueRange offsets, ValueRange sizes, ValueRange strides,
3307 ArrayRef<NamedAttribute> attrs) {
3308 build(b, result, MemRefType(), source, offsets, sizes, strides, attrs);
3309}
3310
3311/// For ViewLikeOpInterface.
3312Value SubViewOp::getViewSource() { return getSource(); }
3313
3314/// Return true if `t1` and `t2` have equal offsets (both dynamic or of same
3315/// static value).
3316static bool haveCompatibleOffsets(MemRefType t1, MemRefType t2) {
3317 int64_t t1Offset, t2Offset;
3318 SmallVector<int64_t> t1Strides, t2Strides;
3319 auto res1 = t1.getStridesAndOffset(t1Strides, t1Offset);
3320 auto res2 = t2.getStridesAndOffset(t2Strides, t2Offset);
3321 return succeeded(res1) && succeeded(res2) && t1Offset == t2Offset;
3322}
3323
3324/// Return true if `t1` and `t2` have equal strides (both dynamic or of same
3325/// static value). Dimensions of `t1` may be dropped in `t2`; these must be
3326/// marked as dropped in `droppedDims`.
3327static bool haveCompatibleStrides(MemRefType t1, MemRefType t2,
3328 const llvm::SmallBitVector &droppedDims) {
3329 assert(size_t(t1.getRank()) == droppedDims.size() &&
3330 "incorrect number of bits");
3331 assert(size_t(t1.getRank() - t2.getRank()) == droppedDims.count() &&
3332 "incorrect number of dropped dims");
3333 int64_t t1Offset, t2Offset;
3334 SmallVector<int64_t> t1Strides, t2Strides;
3335 auto res1 = t1.getStridesAndOffset(t1Strides, t1Offset);
3336 auto res2 = t2.getStridesAndOffset(t2Strides, t2Offset);
3337 if (failed(res1) || failed(res2))
3338 return false;
3339 for (int64_t i = 0, j = 0, e = t1.getRank(); i < e; ++i) {
3340 if (droppedDims[i])
3341 continue;
3342 if (t1Strides[i] != t2Strides[j])
3343 return false;
3344 ++j;
3345 }
3346 return true;
3347}
3348
3350 SubViewOp op, Type expectedType) {
3351 auto memrefType = llvm::cast<ShapedType>(expectedType);
3352 switch (result) {
3354 return success();
3356 return op->emitError("expected result rank to be smaller or equal to ")
3357 << "the source rank, but got " << op.getType();
3359 return op->emitError("expected result type to be ")
3360 << expectedType
3361 << " or a rank-reduced version. (mismatch of result sizes), but got "
3362 << op.getType();
3364 return op->emitError("expected result element type to be ")
3365 << memrefType.getElementType() << ", but got " << op.getType();
3367 return op->emitError(
3368 "expected result and source memory spaces to match, but got ")
3369 << op.getType();
3371 return op->emitError("expected result type to be ")
3372 << expectedType
3373 << " or a rank-reduced version. (mismatch of result layout), but "
3374 "got "
3375 << op.getType();
3376 }
3377 llvm_unreachable("unexpected subview verification result");
3378}
3379
3380/// Verifier for SubViewOp.
3381LogicalResult SubViewOp::verify() {
3382 MemRefType baseType = getSourceType();
3383 MemRefType subViewType = getType();
3384 ArrayRef<int64_t> staticOffsets = getStaticOffsets();
3385 ArrayRef<int64_t> staticSizes = getStaticSizes();
3386 ArrayRef<int64_t> staticStrides = getStaticStrides();
3387
3388 // The base memref and the view memref should be in the same memory space.
3389 if (baseType.getMemorySpace() != subViewType.getMemorySpace())
3390 return emitError("different memory spaces specified for base memref "
3391 "type ")
3392 << baseType << " and subview memref type " << subViewType;
3393
3394 // Verify that the base memref type has a strided layout map.
3395 if (!baseType.isStrided())
3396 return emitError("base type ") << baseType << " is not strided";
3397
3398 // Compute the expected result type, assuming that there are no rank
3399 // reductions.
3400 MemRefType expectedType = SubViewOp::inferResultType(
3401 baseType, staticOffsets, staticSizes, staticStrides);
3402
3403 // Verify all properties of a shaped type: rank, element type and dimension
3404 // sizes. This takes into account potential rank reductions.
3405 auto shapedTypeVerification = isRankReducedType(
3406 /*originalType=*/expectedType, /*candidateReducedType=*/subViewType);
3407 if (shapedTypeVerification != SliceVerificationResult::Success)
3408 return produceSubViewErrorMsg(shapedTypeVerification, *this, expectedType);
3409
3410 // Make sure that the memory space did not change.
3411 if (expectedType.getMemorySpace() != subViewType.getMemorySpace())
3413 *this, expectedType);
3414
3415 // Verify the offset of the layout map.
3416 if (!haveCompatibleOffsets(expectedType, subViewType))
3418 *this, expectedType);
3419
3420 // The only thing that's left to verify now are the strides. First, compute
3421 // the unused dimensions due to rank reductions. We have to look at sizes and
3422 // strides to decide which dimensions were dropped. This function also
3423 // partially verifies strides in case of rank reductions.
3424 auto unusedDims = computeMemRefRankReductionMask(expectedType, subViewType,
3425 getMixedSizes());
3426 if (failed(unusedDims))
3428 *this, expectedType);
3429
3430 // Strides must match.
3431 if (!haveCompatibleStrides(expectedType, subViewType, *unusedDims))
3433 *this, expectedType);
3434
3435 // Verify that offsets, sizes, strides do not run out-of-bounds with respect
3436 // to the base memref.
3437 SliceBoundsVerificationResult boundsResult =
3438 verifyInBoundsSlice(baseType.getShape(), staticOffsets, staticSizes,
3439 staticStrides, /*generateErrorMessage=*/true);
3440 if (!boundsResult.isValid)
3441 return getOperation()->emitError(boundsResult.errorMessage);
3442
3443 return success();
3444}
3445
3447 return os << "range " << range.offset << ":" << range.size << ":"
3448 << range.stride;
3449}
3450
3451/// Return the list of Range (i.e. offset, size, stride). Each Range
3452/// entry contains either the dynamic value or a ConstantIndexOp constructed
3453/// with `b` at location `loc`.
3454SmallVector<Range, 8> mlir::getOrCreateRanges(OffsetSizeAndStrideOpInterface op,
3455 OpBuilder &b, Location loc) {
3456 std::array<unsigned, 3> ranks = op.getArrayAttrMaxRanks();
3457 assert(ranks[0] == ranks[1] && "expected offset and sizes of equal ranks");
3458 assert(ranks[1] == ranks[2] && "expected sizes and strides of equal ranks");
3460 unsigned rank = ranks[0];
3461 res.reserve(rank);
3462 for (unsigned idx = 0; idx < rank; ++idx) {
3463 Value offset =
3464 op.isDynamicOffset(idx)
3465 ? op.getDynamicOffset(idx)
3466 : arith::ConstantIndexOp::create(b, loc, op.getStaticOffset(idx));
3467 Value size =
3468 op.isDynamicSize(idx)
3469 ? op.getDynamicSize(idx)
3470 : arith::ConstantIndexOp::create(b, loc, op.getStaticSize(idx));
3471 Value stride =
3472 op.isDynamicStride(idx)
3473 ? op.getDynamicStride(idx)
3474 : arith::ConstantIndexOp::create(b, loc, op.getStaticStride(idx));
3475 res.emplace_back(Range{offset, size, stride});
3476 }
3477 return res;
3478}
3479
3480/// Compute the canonical result type of a SubViewOp. Call `inferResultType`
3481/// to deduce the result type for the given `sourceType`. Additionally, reduce
3482/// the rank of the inferred result type if `currentResultType` is lower rank
3483/// than `currentSourceType`. Use this signature if `sourceType` is updated
3484/// together with the result type. In this case, it is important to compute
3485/// the dropped dimensions using `currentSourceType` whose strides align with
3486/// `currentResultType`.
3488 MemRefType currentResultType, MemRefType currentSourceType,
3489 MemRefType sourceType, ArrayRef<OpFoldResult> mixedOffsets,
3490 ArrayRef<OpFoldResult> mixedSizes, ArrayRef<OpFoldResult> mixedStrides) {
3491 MemRefType nonRankReducedType = SubViewOp::inferResultType(
3492 sourceType, mixedOffsets, mixedSizes, mixedStrides);
3493 FailureOr<llvm::SmallBitVector> unusedDims = computeMemRefRankReductionMask(
3494 currentSourceType, currentResultType, mixedSizes);
3495 if (failed(unusedDims))
3496 return nullptr;
3497
3498 auto layout = llvm::cast<StridedLayoutAttr>(nonRankReducedType.getLayout());
3499 SmallVector<int64_t> shape, strides;
3500 unsigned numDimsAfterReduction =
3501 nonRankReducedType.getRank() - unusedDims->count();
3502 shape.reserve(numDimsAfterReduction);
3503 strides.reserve(numDimsAfterReduction);
3504 for (const auto &[idx, size, stride] :
3505 llvm::zip(llvm::seq<unsigned>(0, nonRankReducedType.getRank()),
3506 nonRankReducedType.getShape(), layout.getStrides())) {
3507 if (unusedDims->test(idx))
3508 continue;
3509 shape.push_back(size);
3510 strides.push_back(stride);
3511 }
3512
3513 return MemRefType::get(shape, nonRankReducedType.getElementType(),
3514 StridedLayoutAttr::get(sourceType.getContext(),
3515 layout.getOffset(), strides),
3516 nonRankReducedType.getMemorySpace());
3517}
3518
3520 OpBuilder &b, Location loc, Value memref, ArrayRef<int64_t> targetShape) {
3521 auto memrefType = llvm::cast<MemRefType>(memref.getType());
3522 unsigned rank = memrefType.getRank();
3523 SmallVector<OpFoldResult> offsets(rank, b.getIndexAttr(0));
3525 SmallVector<OpFoldResult> strides(rank, b.getIndexAttr(1));
3526 MemRefType targetType = SubViewOp::inferRankReducedResultType(
3527 targetShape, memrefType, offsets, sizes, strides);
3528 return b.createOrFold<memref::SubViewOp>(loc, targetType, memref, offsets,
3529 sizes, strides);
3530}
3531
3532FailureOr<Value> SubViewOp::rankReduceIfNeeded(OpBuilder &b, Location loc,
3533 Value value,
3534 ArrayRef<int64_t> desiredShape) {
3535 auto sourceMemrefType = llvm::dyn_cast<MemRefType>(value.getType());
3536 assert(sourceMemrefType && "not a ranked memref type");
3537 auto sourceShape = sourceMemrefType.getShape();
3538 if (sourceShape.equals(desiredShape))
3539 return value;
3540 auto maybeRankReductionMask =
3541 mlir::computeRankReductionMask(sourceShape, desiredShape);
3542 if (!maybeRankReductionMask)
3543 return failure();
3544 return createCanonicalRankReducingSubViewOp(b, loc, value, desiredShape);
3545}
3546
3547/// Helper method to check if a `subview` operation is trivially a no-op. This
3548/// is the case if the all offsets are zero, all strides are 1, and the source
3549/// shape is same as the size of the subview. In such cases, the subview can
3550/// be folded into its source.
3551static bool isTrivialSubViewOp(SubViewOp subViewOp) {
3552 if (subViewOp.getSourceType().getRank() != subViewOp.getType().getRank())
3553 return false;
3554
3555 auto mixedOffsets = subViewOp.getMixedOffsets();
3556 auto mixedSizes = subViewOp.getMixedSizes();
3557 auto mixedStrides = subViewOp.getMixedStrides();
3558
3559 // Check offsets are zero.
3560 if (llvm::any_of(mixedOffsets, [](OpFoldResult ofr) {
3561 std::optional<int64_t> intValue = getConstantIntValue(ofr);
3562 return !intValue || intValue.value() != 0;
3563 }))
3564 return false;
3565
3566 // Check strides are one.
3567 if (llvm::any_of(mixedStrides, [](OpFoldResult ofr) {
3568 std::optional<int64_t> intValue = getConstantIntValue(ofr);
3569 return !intValue || intValue.value() != 1;
3570 }))
3571 return false;
3572
3573 // Check all size values are static and matches the (static) source shape.
3574 ArrayRef<int64_t> sourceShape = subViewOp.getSourceType().getShape();
3575 for (const auto &size : llvm::enumerate(mixedSizes)) {
3576 std::optional<int64_t> intValue = getConstantIntValue(size.value());
3577 if (!intValue || *intValue != sourceShape[size.index()])
3578 return false;
3579 }
3580 // All conditions met. The `SubViewOp` is foldable as a no-op.
3581 return true;
3582}
3583
3584namespace {
3585/// Pattern to rewrite a subview op with MemRefCast arguments.
3586/// This essentially pushes memref.cast past its consuming subview when
3587/// `canFoldIntoConsumerOp` is true.
3588///
3589/// Example:
3590/// ```
3591/// %0 = memref.cast %V : memref<16x16xf32> to memref<?x?xf32>
3592/// %1 = memref.subview %0[0, 0][3, 4][1, 1] :
3593/// memref<?x?xf32> to memref<3x4xf32, strided<[?, 1], offset: ?>>
3594/// ```
3595/// is rewritten into:
3596/// ```
3597/// %0 = memref.subview %V: memref<16x16xf32> to memref<3x4xf32, #[[map0]]>
3598/// %1 = memref.cast %0: memref<3x4xf32, strided<[16, 1], offset: 0>> to
3599/// memref<3x4xf32, strided<[?, 1], offset: ?>>
3600/// ```
3601class SubViewOpMemRefCastFolder final : public OpRewritePattern<SubViewOp> {
3602public:
3603 using OpRewritePattern<SubViewOp>::OpRewritePattern;
3604
3605 LogicalResult matchAndRewrite(SubViewOp subViewOp,
3606 PatternRewriter &rewriter) const override {
3607 // Any constant operand, just return to let SubViewOpConstantFolder kick
3608 // in.
3609 if (llvm::any_of(subViewOp.getOperands(), [](Value operand) {
3610 return matchPattern(operand, matchConstantIndex());
3611 }))
3612 return failure();
3613
3614 auto castOp = subViewOp.getSource().getDefiningOp<CastOp>();
3615 if (!castOp)
3616 return failure();
3617
3618 if (!CastOp::canFoldIntoConsumerOp(castOp))
3619 return failure();
3620
3621 // Compute the SubViewOp result type after folding the MemRefCastOp. Use
3622 // the MemRefCastOp source operand type to infer the result type and the
3623 // current SubViewOp source operand type to compute the dropped dimensions
3624 // if the operation is rank-reducing.
3625 auto resultType = getCanonicalSubViewResultType(
3626 subViewOp.getType(), subViewOp.getSourceType(),
3627 llvm::cast<MemRefType>(castOp.getSource().getType()),
3628 subViewOp.getMixedOffsets(), subViewOp.getMixedSizes(),
3629 subViewOp.getMixedStrides());
3630 if (!resultType)
3631 return failure();
3632
3633 Value newSubView = SubViewOp::create(
3634 rewriter, subViewOp.getLoc(), resultType, castOp.getSource(),
3635 subViewOp.getOffsets(), subViewOp.getSizes(), subViewOp.getStrides(),
3636 subViewOp.getStaticOffsets(), subViewOp.getStaticSizes(),
3637 subViewOp.getStaticStrides());
3638 rewriter.replaceOpWithNewOp<CastOp>(subViewOp, subViewOp.getType(),
3639 newSubView);
3640 return success();
3641 }
3642};
3643
3644/// Canonicalize subview ops that are no-ops. When the source shape is not
3645/// same as a result shape due to use of `affine_map`.
3646class TrivialSubViewOpFolder final : public OpRewritePattern<SubViewOp> {
3647public:
3648 using OpRewritePattern<SubViewOp>::OpRewritePattern;
3649
3650 LogicalResult matchAndRewrite(SubViewOp subViewOp,
3651 PatternRewriter &rewriter) const override {
3652 if (!isTrivialSubViewOp(subViewOp))
3653 return failure();
3654 if (subViewOp.getSourceType() == subViewOp.getType()) {
3655 rewriter.replaceOp(subViewOp, subViewOp.getSource());
3656 return success();
3657 }
3658 rewriter.replaceOpWithNewOp<CastOp>(subViewOp, subViewOp.getType(),
3659 subViewOp.getSource());
3660 return success();
3661 }
3662};
3663} // namespace
3664
3665/// Return the canonical type of the result of a subview.
3667 MemRefType operator()(SubViewOp op, ArrayRef<OpFoldResult> mixedOffsets,
3668 ArrayRef<OpFoldResult> mixedSizes,
3669 ArrayRef<OpFoldResult> mixedStrides) {
3670 // Infer a memref type without taking into account any rank reductions.
3671 MemRefType resTy = SubViewOp::inferResultType(
3672 op.getSourceType(), mixedOffsets, mixedSizes, mixedStrides);
3673 if (!resTy)
3674 return {};
3675 MemRefType nonReducedType = resTy;
3676
3677 // Directly return the non-rank reduced type if there are no dropped dims.
3678 llvm::SmallBitVector droppedDims = op.getDroppedDims();
3679 if (droppedDims.none())
3680 return nonReducedType;
3681
3682 // Take the strides and offset from the non-rank reduced type.
3683 auto [nonReducedStrides, offset] = nonReducedType.getStridesAndOffset();
3684
3685 // Drop dims from shape and strides.
3686 SmallVector<int64_t> targetShape;
3687 SmallVector<int64_t> targetStrides;
3688 for (int64_t i = 0; i < static_cast<int64_t>(mixedSizes.size()); ++i) {
3689 if (droppedDims.test(i))
3690 continue;
3691 targetStrides.push_back(nonReducedStrides[i]);
3692 targetShape.push_back(nonReducedType.getDimSize(i));
3693 }
3694
3695 return MemRefType::get(targetShape, nonReducedType.getElementType(),
3696 StridedLayoutAttr::get(nonReducedType.getContext(),
3697 offset, targetStrides),
3698 nonReducedType.getMemorySpace());
3699 }
3700};
3701
3702/// A canonicalizer wrapper to replace SubViewOps.
3704 void operator()(PatternRewriter &rewriter, SubViewOp op, SubViewOp newOp) {
3705 rewriter.replaceOpWithNewOp<CastOp>(op, op.getType(), newOp);
3706 }
3707};
3708
3709void SubViewOp::getCanonicalizationPatterns(RewritePatternSet &results,
3710 MLIRContext *context) {
3711 results
3712 .add<OpWithOffsetSizesAndStridesConstantArgumentFolder<
3713 SubViewOp, SubViewReturnTypeCanonicalizer, SubViewCanonicalizer>,
3714 SubViewOpMemRefCastFolder, TrivialSubViewOpFolder>(context);
3715}
3716
3717OpFoldResult SubViewOp::fold(FoldAdaptor adaptor) {
3718 MemRefType sourceMemrefType = getSource().getType();
3719 MemRefType resultMemrefType = getResult().getType();
3720 auto resultLayout =
3721 dyn_cast_if_present<StridedLayoutAttr>(resultMemrefType.getLayout());
3722
3723 if (resultMemrefType == sourceMemrefType &&
3724 resultMemrefType.hasStaticShape() &&
3725 (!resultLayout || resultLayout.hasStaticLayout())) {
3726 return getViewSource();
3727 }
3728
3729 // Fold subview(subview(x)), where both subviews have the same size and the
3730 // second subview's offsets are all zero. (I.e., the second subview is a
3731 // no-op.)
3732 if (auto srcSubview = getViewSource().getDefiningOp<SubViewOp>()) {
3733 auto srcSizes = srcSubview.getMixedSizes();
3734 auto sizes = getMixedSizes();
3735 auto offsets = getMixedOffsets();
3736 bool allOffsetsZero = llvm::all_of(offsets, isZeroInteger);
3737 auto strides = getMixedStrides();
3738 bool allStridesOne = llvm::all_of(strides, isOneInteger);
3739 bool allSizesSame = llvm::equal(sizes, srcSizes);
3740 if (allOffsetsZero && allStridesOne && allSizesSame &&
3741 resultMemrefType == sourceMemrefType)
3742 return getViewSource();
3743 }
3744
3745 return {};
3746}
3747
3748FailureOr<std::optional<SmallVector<Value>>>
3749SubViewOp::bubbleDownCasts(OpBuilder &builder) {
3750 return bubbleDownCastsPassthroughOpImpl(*this, builder, getSourceMutable());
3751}
3752
3753void SubViewOp::inferStridedMetadataRanges(
3754 ArrayRef<StridedMetadataRange> ranges, GetIntRangeFn getIntRange,
3755 SetStridedMetadataRangeFn setMetadata, int32_t indexBitwidth) {
3756 auto isUninitialized =
3757 +[](IntegerValueRange range) { return range.isUninitialized(); };
3758
3759 // Bail early if any of the operands metadata is not ready:
3760 SmallVector<IntegerValueRange> offsetOperands =
3761 getIntValueRanges(getMixedOffsets(), getIntRange, indexBitwidth);
3762 if (llvm::any_of(offsetOperands, isUninitialized))
3763 return;
3764
3765 SmallVector<IntegerValueRange> sizeOperands =
3766 getIntValueRanges(getMixedSizes(), getIntRange, indexBitwidth);
3767 if (llvm::any_of(sizeOperands, isUninitialized))
3768 return;
3769
3770 SmallVector<IntegerValueRange> stridesOperands =
3771 getIntValueRanges(getMixedStrides(), getIntRange, indexBitwidth);
3772 if (llvm::any_of(stridesOperands, isUninitialized))
3773 return;
3774
3775 StridedMetadataRange sourceRange =
3776 ranges[getSourceMutable().getOperandNumber()];
3777 if (sourceRange.isUninitialized())
3778 return;
3779
3780 ArrayRef<ConstantIntRanges> srcStrides = sourceRange.getStrides();
3781
3782 // Get the dropped dims.
3783 llvm::SmallBitVector droppedDims = getDroppedDims();
3784
3785 // Compute the new offset, strides and sizes.
3786 ConstantIntRanges offset = sourceRange.getOffsets()[0];
3787 SmallVector<ConstantIntRanges> strides, sizes;
3788
3789 for (size_t i = 0, e = droppedDims.size(); i < e; ++i) {
3790 bool dropped = droppedDims.test(i);
3791 // Compute the new offset.
3792 ConstantIntRanges off =
3793 intrange::inferMul({offsetOperands[i].getValue(), srcStrides[i]});
3794 offset = intrange::inferAdd({offset, off});
3795
3796 // Skip dropped dimensions.
3797 if (dropped)
3798 continue;
3799 // Multiply the strides.
3800 strides.push_back(
3801 intrange::inferMul({stridesOperands[i].getValue(), srcStrides[i]}));
3802 // Get the sizes.
3803 sizes.push_back(sizeOperands[i].getValue());
3804 }
3805
3806 setMetadata(getResult(),
3808 SmallVector<ConstantIntRanges>({std::move(offset)}),
3809 std::move(sizes), std::move(strides)));
3810}
3811
3812//===----------------------------------------------------------------------===//
3813// TransposeOp
3814//===----------------------------------------------------------------------===//
3815
3816void TransposeOp::getAsmResultNames(
3817 function_ref<void(Value, StringRef)> setNameFn) {
3818 setNameFn(getResult(), "transpose");
3819}
3820
3821/// Build a strided memref type by applying `permutationMap` to `memRefType`.
3822static MemRefType inferTransposeResultType(MemRefType memRefType,
3823 AffineMap permutationMap) {
3824 auto originalSizes = memRefType.getShape();
3825 auto [originalStrides, offset] = memRefType.getStridesAndOffset();
3826 assert(originalStrides.size() == static_cast<unsigned>(memRefType.getRank()));
3827
3828 // Compute permuted sizes and strides.
3829 auto sizes = applyPermutationMap<int64_t>(permutationMap, originalSizes);
3830 auto strides = applyPermutationMap<int64_t>(permutationMap, originalStrides);
3831
3832 return MemRefType::Builder(memRefType)
3833 .setShape(sizes)
3834 .setLayout(
3835 StridedLayoutAttr::get(memRefType.getContext(), offset, strides));
3836}
3837
3838Value TransposeOp::getViewSource() { return getIn(); }
3839
3840void TransposeOp::build(OpBuilder &b, OperationState &result, Value in,
3841 AffineMapAttr permutation,
3842 ArrayRef<NamedAttribute> attrs) {
3843 auto permutationMap = permutation.getValue();
3844 assert(permutationMap);
3845
3846 auto memRefType = llvm::cast<MemRefType>(in.getType());
3847 // Compute result type.
3848 MemRefType resultType = inferTransposeResultType(memRefType, permutationMap);
3849
3850 buildPropertiesAndDiscardableAttributes(result, attrs);
3851 result.getOrAddProperties<Properties>().permutation = permutation;
3852 result.addOperands(in);
3853 result.addTypes(resultType);
3854}
3855
3856// transpose $in $permutation attr-dict : type($in) `to` type(results)
3857void TransposeOp::print(OpAsmPrinter &p) {
3858 p << " " << getIn() << " " << getPermutation();
3859 p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary(),
3860 {getPermutationAttrStrName()});
3861 p << " : " << getIn().getType() << " to " << getType();
3862}
3863
3864ParseResult TransposeOp::parse(OpAsmParser &parser, OperationState &result) {
3865 OpAsmParser::UnresolvedOperand in;
3866 AffineMap permutation;
3867 MemRefType srcType, dstType;
3868 if (parser.parseOperand(in) || parser.parseAffineMap(permutation) ||
3869 parser.parseOptionalAttrDict(result.attributes) ||
3870 parser.parseColonType(srcType) ||
3871 parser.resolveOperand(in, srcType, result.operands) ||
3872 parser.parseKeywordType("to", dstType) ||
3873 parser.addTypeToList(dstType, result.types))
3874 return failure();
3875
3876 result.addAttribute(TransposeOp::getPermutationAttrStrName(),
3877 AffineMapAttr::get(permutation));
3878 return success();
3879}
3880
3881LogicalResult TransposeOp::verify() {
3882 if (!getPermutation().isPermutation())
3883 return emitOpError("expected a permutation map");
3884 if (getPermutation().getNumDims() != getIn().getType().getRank())
3885 return emitOpError("expected a permutation map of same rank as the input");
3886
3887 auto srcType = llvm::cast<MemRefType>(getIn().getType());
3888 auto resultType = llvm::cast<MemRefType>(getType());
3889 auto canonicalResultType = inferTransposeResultType(srcType, getPermutation())
3890 .canonicalizeStridedLayout();
3891
3892 if (resultType.canonicalizeStridedLayout() != canonicalResultType)
3893 return emitOpError("result type ")
3894 << resultType
3895 << " is not equivalent to the canonical transposed input type "
3896 << canonicalResultType;
3897 return success();
3898}
3899
3900OpFoldResult TransposeOp::fold(FoldAdaptor) {
3901 // First check for identity permutation, we can fold it away if input and
3902 // result types are identical already.
3903 if (getPermutation().isIdentity() && getType() == getIn().getType())
3904 return getIn();
3905 // Fold two consecutive memref.transpose Ops into one by composing their
3906 // permutation maps.
3907 if (auto otherTransposeOp = getIn().getDefiningOp<memref::TransposeOp>()) {
3908 AffineMap composedPermutation =
3909 getPermutation().compose(otherTransposeOp.getPermutation());
3910 getInMutable().assign(otherTransposeOp.getIn());
3911 setPermutation(composedPermutation);
3912 return getResult();
3913 }
3914 return {};
3915}
3916
3917FailureOr<std::optional<SmallVector<Value>>>
3918TransposeOp::bubbleDownCasts(OpBuilder &builder) {
3919 return bubbleDownCastsPassthroughOpImpl(*this, builder, getInMutable());
3920}
3921
3922//===----------------------------------------------------------------------===//
3923// ViewOp
3924//===----------------------------------------------------------------------===//
3925
3926void ViewOp::getAsmResultNames(function_ref<void(Value, StringRef)> setNameFn) {
3927 setNameFn(getResult(), "view");
3928}
3929
3930LogicalResult ViewOp::verify() {
3931 auto baseType = llvm::cast<MemRefType>(getOperand(0).getType());
3932 auto viewType = getType();
3933
3934 // The base memref should have identity layout map (or none).
3935 if (!baseType.getLayout().isIdentity())
3936 return emitError("unsupported map for base memref type ") << baseType;
3937
3938 // The result memref should have identity layout map (or none).
3939 if (!viewType.getLayout().isIdentity())
3940 return emitError("unsupported map for result memref type ") << viewType;
3941
3942 // The base memref and the view memref should be in the same memory space.
3943 if (baseType.getMemorySpace() != viewType.getMemorySpace())
3944 return emitError("different memory spaces specified for base memref "
3945 "type ")
3946 << baseType << " and view memref type " << viewType;
3947
3948 // Verify that we have the correct number of sizes for the result type.
3949 if (failed(verifyDynamicDimensionCount(getOperation(), viewType, getSizes())))
3950 return failure();
3951
3952 return success();
3953}
3954
3955Value ViewOp::getViewSource() { return getSource(); }
3956
3957OpFoldResult ViewOp::fold(FoldAdaptor adaptor) {
3958 MemRefType sourceMemrefType = getSource().getType();
3959 MemRefType resultMemrefType = getResult().getType();
3960
3961 if (resultMemrefType == sourceMemrefType &&
3962 resultMemrefType.hasStaticShape() && isZeroInteger(getByteShift()))
3963 return getViewSource();
3964
3965 return {};
3966}
3967
3968SmallVector<OpFoldResult> ViewOp::getMixedSizes() {
3969 SmallVector<OpFoldResult> result;
3970 unsigned ctr = 0;
3971 Builder b(getContext());
3972 for (int64_t dim : getType().getShape()) {
3973 if (ShapedType::isDynamic(dim)) {
3974 result.push_back(getSizes()[ctr++]);
3975 } else {
3976 result.push_back(b.getIndexAttr(dim));
3977 }
3978 }
3979 return result;
3980}
3981
3982namespace {
3983/// Given a memref type and a range of values that defines its dynamic
3984/// dimension sizes, turn all dynamic sizes that have a constant value into
3985/// static dimension sizes.
3986static MemRefType
3987foldDynamicToStaticDimSizes(MemRefType type, ValueRange dynamicSizes,
3988 SmallVectorImpl<Value> &foldedDynamicSizes) {
3989 SmallVector<int64_t> staticShape(type.getShape());
3990 assert(type.getNumDynamicDims() == dynamicSizes.size() &&
3991 "incorrect number of dynamic sizes");
3992
3993 // Compute new static and dynamic sizes.
3994 unsigned ctr = 0;
3995 for (auto [dim, dimSize] : llvm::enumerate(type.getShape())) {
3996 if (ShapedType::isStatic(dimSize))
3997 continue;
3998
3999 Value dynamicSize = dynamicSizes[ctr++];
4000 if (auto cst = getConstantIntValue(dynamicSize)) {
4001 // Dynamic size must be non-negative.
4002 if (cst.value() < 0) {
4003 foldedDynamicSizes.push_back(dynamicSize);
4004 continue;
4005 }
4006 staticShape[dim] = cst.value();
4007 } else {
4008 foldedDynamicSizes.push_back(dynamicSize);
4009 }
4010 }
4011
4012 return MemRefType::Builder(type).setShape(staticShape);
4013}
4014
4015/// Change the result type of a `memref.view` by making originally dynamic
4016/// dimensions static when their sizes come from `constant` ops.
4017/// Example:
4018/// ```
4019/// %c5 = arith.constant 5: index
4020/// %0 = memref.view %src[%offset][%c5] : memref<?xi8> to memref<?x4xf32>
4021/// ```
4022/// to
4023/// ```
4024/// %0 = memref.view %src[%offset][] : memref<?xi8> to memref<5x4xf32>
4025/// ```
4026struct ViewOpShapeFolder : public OpRewritePattern<ViewOp> {
4027 using Base::Base;
4028
4029 LogicalResult matchAndRewrite(ViewOp viewOp,
4030 PatternRewriter &rewriter) const override {
4031 SmallVector<Value> foldedDynamicSizes;
4032 MemRefType resultType = viewOp.getType();
4033 MemRefType foldedMemRefType = foldDynamicToStaticDimSizes(
4034 resultType, viewOp.getSizes(), foldedDynamicSizes);
4035
4036 // Stop here if no dynamic size was promoted to static.
4037 if (foldedMemRefType == resultType)
4038 return failure();
4039
4040 // Create new ViewOp.
4041 auto newViewOp = ViewOp::create(rewriter, viewOp.getLoc(), foldedMemRefType,
4042 viewOp.getSource(), viewOp.getByteShift(),
4043 foldedDynamicSizes);
4044 // Insert a cast so we have the same type as the old memref type.
4045 rewriter.replaceOpWithNewOp<CastOp>(viewOp, resultType, newViewOp);
4046 return success();
4047 }
4048};
4049
4050/// view(memref.cast(%source)) -> view(%source).
4051struct ViewOpMemrefCastFolder : public OpRewritePattern<ViewOp> {
4052 using Base::Base;
4053
4054 LogicalResult matchAndRewrite(ViewOp viewOp,
4055 PatternRewriter &rewriter) const override {
4056 auto memrefCastOp = viewOp.getSource().getDefiningOp<CastOp>();
4057 if (!memrefCastOp)
4058 return failure();
4059
4060 rewriter.replaceOpWithNewOp<ViewOp>(
4061 viewOp, viewOp.getType(), memrefCastOp.getSource(),
4062 viewOp.getByteShift(), viewOp.getSizes());
4063 return success();
4064 }
4065};
4066} // namespace
4067
4068void ViewOp::getCanonicalizationPatterns(RewritePatternSet &results,
4069 MLIRContext *context) {
4070 results.add<ViewOpShapeFolder, ViewOpMemrefCastFolder>(context);
4071}
4072
4073FailureOr<std::optional<SmallVector<Value>>>
4074ViewOp::bubbleDownCasts(OpBuilder &builder) {
4075 return bubbleDownCastsPassthroughOpImpl(*this, builder, getSourceMutable());
4076}
4077
4078//===----------------------------------------------------------------------===//
4079// AtomicRMWOp
4080//===----------------------------------------------------------------------===//
4081
4082LogicalResult AtomicRMWOp::verify() {
4083 switch (getKind()) {
4084 case arith::AtomicRMWKind::addf:
4085 case arith::AtomicRMWKind::maximumf:
4086 case arith::AtomicRMWKind::minimumf:
4087 case arith::AtomicRMWKind::mulf:
4088 if (!llvm::isa<FloatType>(getValue().getType()))
4089 return emitOpError() << "with kind '"
4090 << arith::stringifyAtomicRMWKind(getKind())
4091 << "' expects a floating-point type";
4092 break;
4093 case arith::AtomicRMWKind::addi:
4094 case arith::AtomicRMWKind::maxs:
4095 case arith::AtomicRMWKind::maxu:
4096 case arith::AtomicRMWKind::mins:
4097 case arith::AtomicRMWKind::minu:
4098 case arith::AtomicRMWKind::muli:
4099 case arith::AtomicRMWKind::ori:
4100 case arith::AtomicRMWKind::xori:
4101 case arith::AtomicRMWKind::andi:
4102 if (!llvm::isa<IntegerType>(getValue().getType()))
4103 return emitOpError() << "with kind '"
4104 << arith::stringifyAtomicRMWKind(getKind())
4105 << "' expects an integer type";
4106 break;
4107 default:
4108 break;
4109 }
4110 return success();
4111}
4112
4113OpFoldResult AtomicRMWOp::fold(FoldAdaptor adaptor) {
4114 /// atomicrmw(memrefcast) -> atomicrmw
4115 if (succeeded(foldMemRefCast(*this, getValue())))
4116 return getResult();
4117 return OpFoldResult();
4118}
4119
4120FailureOr<std::optional<SmallVector<Value>>>
4121AtomicRMWOp::bubbleDownCasts(OpBuilder &builder) {
4123 getResult());
4124}
4125
4126TypedValue<MemRefType> AtomicRMWOp::getAccessedMemref() { return getMemref(); }
4127
4128std::optional<SmallVector<Value>>
4129AtomicRMWOp::updateMemrefAndIndices(RewriterBase &rewriter, Value newMemref,
4130 ValueRange newIndices) {
4131 rewriter.modifyOpInPlace(*this, [&]() {
4132 getMemrefMutable().assign(newMemref);
4133 getIndicesMutable().assign(newIndices);
4134 });
4135 return std::nullopt;
4136}
4137
4138//===----------------------------------------------------------------------===//
4139// TableGen'd op method definitions
4140//===----------------------------------------------------------------------===//
4141
4142#define GET_OP_CLASSES
4143#include "mlir/Dialect/MemRef/IR/MemRefOps.cpp.inc"
return success()
getNumOperands() - 1))) return failure()
Given a list of lists of parsed operands, populates uniqueOperands with unique operands.
static bool hasSideEffects(Operation *op)
static bool isPermutation(const std::vector< PermutationTy > &permutation)
Definition IRAffine.cpp:60
static int64_t getNumElements(Type t)
Compute the total number of elements in the given type, also taking into account nested types.
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
auto load
static LogicalResult foldCopyOfCast(CopyOp op)
If the source/target of a CopyOp is a CastOp that does not modify the shape and element type,...
static void constifyIndexValues(SmallVectorImpl< OpFoldResult > &values, ArrayRef< int64_t > constValues)
Helper function that sets values[i] to constValues[i] if the latter is a static value,...
Definition MemRefOps.cpp:99
static void printGlobalMemrefOpTypeAndInitialValue(OpAsmPrinter &p, GlobalOp op, TypeAttr type, Attribute initialValue)
static LogicalResult verifyCollapsedShape(Operation *op, ArrayRef< int64_t > collapsedShape, ArrayRef< int64_t > expandedShape, ArrayRef< ReassociationIndices > reassociation, bool allowMultipleDynamicDimsPerGroup)
Helper function for verifying the shape of ExpandShapeOp and ResultShapeOp result and operand.
static bool isOpItselfPotentialAutomaticAllocation(Operation *op)
Given an operation, return whether this op itself could allocate an AutomaticAllocationScopeResource.
static MemRefType inferTransposeResultType(MemRefType memRefType, AffineMap permutationMap)
Build a strided memref type by applying permutationMap to memRefType.
static ParseResult parseBoolAttr(OpAsmParser &parser, BoolAttr &result)
static bool isGuaranteedAutomaticAllocation(Operation *op)
Given an operation, return whether this op is guaranteed to allocate an AutomaticAllocationScopeResou...
static FailureOr< llvm::SmallBitVector > computeMemRefRankReductionMaskByStrides(MemRefType originalType, MemRefType reducedType, ArrayRef< int64_t > originalStrides, ArrayRef< int64_t > candidateStrides, llvm::SmallBitVector unusedDims)
Returns the set of source dimensions that are dropped in a rank reduction.
static FailureOr< StridedLayoutAttr > computeExpandedLayoutMap(MemRefType srcType, ArrayRef< int64_t > resultShape, ArrayRef< ReassociationIndices > reassociation)
Compute the layout map after expanding a given source MemRef type with the specified reassociation in...
static bool haveCompatibleOffsets(MemRefType t1, MemRefType t2)
Return true if t1 and t2 have equal offsets (both dynamic or of same static value).
static void printBoolAttr(OpAsmPrinter &printer, Operation *, BoolAttr attr)
static bool replaceConstantUsesOf(OpBuilder &rewriter, Location loc, Container values, ArrayRef< OpFoldResult > maybeConstants)
Helper function to perform the replacement of all constant uses of values by a materialized constant ...
static LogicalResult produceSubViewErrorMsg(SliceVerificationResult result, SubViewOp op, Type expectedType)
static MemRefType getCanonicalSubViewResultType(MemRefType currentResultType, MemRefType currentSourceType, MemRefType sourceType, ArrayRef< OpFoldResult > mixedOffsets, ArrayRef< OpFoldResult > mixedSizes, ArrayRef< OpFoldResult > mixedStrides)
Compute the canonical result type of a SubViewOp.
static ParseResult parseGlobalMemrefOpTypeAndInitialValue(OpAsmParser &parser, TypeAttr &typeAttr, Attribute &initialValue)
static std::tuple< MemorySpaceCastOpInterface, PtrLikeTypeInterface, Type > getMemorySpaceCastInfo(BaseMemRefType resultTy, Value src)
Helper function to retrieve a lossless memory-space cast, and the corresponding new result memref typ...
static FailureOr< llvm::SmallBitVector > computeMemRefRankReductionMask(MemRefType originalType, MemRefType reducedType, ArrayRef< OpFoldResult > sizes)
Given the originalType and a candidateReducedType whose shape is assumed to be a subset of originalTy...
static bool isTrivialSubViewOp(SubViewOp subViewOp)
Helper method to check if a subview operation is trivially a no-op.
static bool lastNonTerminatorInRegion(Operation *op)
Return whether this op is the last non terminating op in a region.
static std::map< int64_t, unsigned > getNumOccurences(ArrayRef< int64_t > vals)
Return a map with key being elements in vals and data being number of occurences of it.
static bool haveCompatibleStrides(MemRefType t1, MemRefType t2, const llvm::SmallBitVector &droppedDims)
Return true if t1 and t2 have equal strides (both dynamic or of same static value).
static FailureOr< StridedLayoutAttr > computeCollapsedLayoutMap(MemRefType srcType, ArrayRef< ReassociationIndices > reassociation, bool strict=false)
Compute the layout map after collapsing a given source MemRef type with the specified reassociation i...
static FailureOr< std::optional< SmallVector< Value > > > bubbleDownCastsPassthroughOpImpl(ConcreteOpTy op, OpBuilder &builder, OpOperand &src)
Implementation of bubbleDownCasts method for memref operations that return a single memref result.
static FailureOr< llvm::SmallBitVector > computeMemRefRankReductionMaskByPosition(MemRefType originalType, MemRefType reducedType, ArrayRef< OpFoldResult > sizes)
Returns the set of source dimensions that are dropped in a rank reduction.
static LogicalResult verifyAllocLikeOp(AllocLikeOp op)
static Type getElementType(Type type, ArrayRef< int32_t > indices, function_ref< InFlightDiagnostic(StringRef)> emitErrorFn)
Walks the given type hierarchy with the given indices, potentially down to component granularity,...
Definition SPIRVOps.cpp:229
static RankedTensorType foldDynamicToStaticDimSizes(RankedTensorType type, ValueRange dynamicSizes, SmallVector< Value > &foldedDynamicSizes)
Given a ranked tensor type and a range of values that defines its dynamic dimension sizes,...
static llvm::SmallBitVector getDroppedDims(ArrayRef< int64_t > reducedShape, ArrayRef< OpFoldResult > mixedSizes)
Compute the dropped dimensions of a rank-reducing tensor.extract_slice op or rank-extending tensor....
static ArrayRef< int64_t > getShape(Type type)
Returns the shape of the given type.
Definition Traits.cpp:117
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
@ Square
Square brackets surrounding zero or more operands.
virtual ParseResult parseColonTypeList(SmallVectorImpl< Type > &result)=0
Parse a colon followed by a type list, which must have at least one type.
virtual Builder & getBuilder() const =0
Return a builder which provides useful access to MLIRContext, global objects like types and attribute...
virtual ParseResult parseOptionalAttrDict(NamedAttrList &result)=0
Parse a named dictionary into 'result' if it is present.
virtual ParseResult parseOptionalEqual()=0
Parse a = token if present.
virtual ParseResult parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
MLIRContext * getContext() const
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual ParseResult parseAffineMap(AffineMap &map)=0
Parse an affine map instance into 'map'.
ParseResult addTypeToList(Type type, SmallVectorImpl< Type > &result)
Add the specified type to the end of the specified type list and return success.
virtual ParseResult parseLess()=0
Parse a '<' token.
virtual ParseResult parseColonType(Type &result)=0
Parse a colon followed by a type.
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseGreater()=0
Parse a '>' token.
virtual ParseResult parseType(Type &result)=0
Parse a type.
virtual ParseResult parseComma()=0
Parse a , token.
virtual ParseResult parseOptionalArrowTypeList(SmallVectorImpl< Type > &result)=0
Parse an optional arrow followed by a type list.
ParseResult parseKeywordType(const char *keyword, Type &result)
Parse a keyword followed by a type.
ParseResult parseKeyword(StringRef keyword)
Parse a given keyword.
virtual ParseResult parseAttribute(Attribute &result, Type type={})=0
Parse an arbitrary attribute of a given type and return it in result.
virtual void printAttributeWithoutType(Attribute attr)
Print the given attribute without its type.
virtual void printAttribute(Attribute attr)
Attributes are known-constant values of operations.
Definition Attributes.h:25
This class provides a shared interface for ranked and unranked memref types.
ArrayRef< int64_t > getShape() const
Returns the shape of this memref type.
FailureOr< PtrLikeTypeInterface > clonePtrWith(Attribute memorySpace, std::optional< Type > elementType) const
Clone this type with the given memory space and element type.
bool hasRank() const
Returns if this type is ranked, i.e. it has a known number of dimensions.
Block represents an ordered list of Operations.
Definition Block.h:34
Operation & front()
Definition Block.h:178
Operation * getTerminator()
Get the terminator operation of this block.
Definition Block.cpp:249
bool mightHaveTerminator()
Return "true" if this block might have a terminator.
Definition Block.cpp:255
Special case of IntegerAttr to represent boolean integers, i.e., signless i1 integers.
This class is a general helper class for creating context-global objects like types,...
Definition Builders.h:51
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
IntegerType getIntegerType(unsigned width)
Definition Builders.cpp:75
BoolAttr getBoolAttr(bool value)
Definition Builders.cpp:108
IndexType getIndexType()
Definition Builders.cpp:59
IRValueT get() const
Return the current value being used by this operand.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
This is a builder type that keeps local references to arguments.
Builder & setShape(ArrayRef< int64_t > newShape)
Builder & setLayout(MemRefLayoutAttrInterface newLayout)
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
virtual ParseResult parseRegion(Region &region, ArrayRef< Argument > arguments={}, bool enableNameShadowing=false)=0
Parses a region.
ParseResult parseTrailingOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None)
Parse zero or more trailing SSA comma-separated trailing operand references with a specified surround...
virtual ParseResult resolveOperand(const UnresolvedOperand &operand, Type type, SmallVectorImpl< Value > &result)=0
Resolve an operand to an SSA value, emitting an error on failure.
ParseResult resolveOperands(Operands &&operands, Type type, SmallVectorImpl< Value > &result)
Resolve a list of operands to SSA values, emitting an error on failure, or appending the results to t...
virtual ParseResult parseOperand(UnresolvedOperand &result, bool allowResultNumber=true)=0
Parse a single SSA value operand name along with a result number if allowResultNumber is true.
virtual ParseResult parseOperandList(SmallVectorImpl< UnresolvedOperand > &result, Delimiter delimiter=Delimiter::None, bool allowResultNumber=true, int requiredOperandCount=-1)=0
Parse zero or more SSA comma-separated operand references with a specified surrounding delimiter,...
This is a pure-virtual base class that exposes the asmprinter hooks necessary to implement a custom p...
void printOperands(OperandRange operands)
Print a comma separated range of operation operands out of line to avoid instantiating the range iter...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
virtual void printRegion(Region &blocks, bool printEntryBlockArgs=true, bool printBlockTerminators=true, bool printEmptyBlock=false)=0
Prints a region.
This class helps build Operations.
Definition Builders.h:210
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
Definition Builders.cpp:439
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
Definition Builders.cpp:581
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
Definition Builders.h:528
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
Definition Builders.h:415
This class represents a single result from folding an operation.
This class represents an operand of an operation.
Definition Value.h:254
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
Definition Value.cpp:226
A trait of region holding operations that define a new scope for automatic allocations,...
This trait indicates that the memory effects of an operation includes the effects of operations neste...
type_range getType() const
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
void replaceUsesOfWith(Value from, Value to)
Replace any uses of 'from' with 'to' within this operation.
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:801
Block * getBlock()
Returns the operation block that contains this operation.
Definition Operation.h:230
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
MutableArrayRef< OpOperand > getOpOperands()
Definition Operation.h:408
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
Definition Operation.h:729
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
result_range getResults()
Definition Operation.h:440
Region * getParentRegion()
Returns the region to which the instruction belongs.
Definition Operation.h:247
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
Type-safe wrapper around a void* for passing properties, including the properties structs of operatio...
This class represents a point being branched from in the methods of the RegionBranchOpInterface.
bool isParent() const
Returns true if branching from the parent op.
This class provides an abstraction over the different types of ranges over Regions.
Definition Region.h:363
This class represents a successor of a region.
bool isOperation() const
Return true if the successor is an operation.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Definition Region.h:26
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
Definition Region.h:111
bool hasOneBlock()
Return true if this region has exactly one block.
Definition Region.h:68
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
virtual void inlineBlockBefore(Block *source, Block *dest, Block::iterator before, ValueRange argValues={})
Inline the operations of block 'source' into block 'dest' before the given position.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
static StridedMetadataRange getRanked(SmallVectorImpl< ConstantIntRanges > &&offsets, SmallVectorImpl< ConstantIntRanges > &&sizes, SmallVectorImpl< ConstantIntRanges > &&strides)
Returns a ranked strided metadata range.
ArrayRef< ConstantIntRanges > getStrides() const
Get the strides ranges.
bool isUninitialized() const
Returns whether the metadata is uninitialized.
ArrayRef< ConstantIntRanges > getOffsets() const
Get the offsets range.
virtual Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
static Operation * lookupNearestSymbolFrom(Operation *from, StringAttr symbol)
Returns the operation registered with the given symbol name within the closest parent operation of,...
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
MLIRContext * getContext() const
Return the MLIRContext in which this type was uniqued.
Definition Types.cpp:35
bool isIndex() const
Definition Types.cpp:56
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getTypes() const
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
Location getLoc() const
Return the location of this value.
Definition Value.cpp:24
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
static WalkResult skip()
Definition WalkResult.h:48
static WalkResult advance()
Definition WalkResult.h:47
static WalkResult interrupt()
Definition WalkResult.h:46
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
Speculatability
This enum is returned from the getSpeculatability method in the ConditionallySpeculatable op interfac...
constexpr auto Speculatable
constexpr auto NotSpeculatable
constexpr void enumerate(std::tuple< Tys... > &tuple, CallbackT &&callback)
Definition Matchers.h:344
FailureOr< std::optional< SmallVector< Value > > > bubbleDownInPlaceMemorySpaceCastImpl(OpOperand &operand, ValueRange results)
Tries to bubble-down inplace a MemorySpaceCastOpInterface operation referenced by operand.
ConstantIntRanges inferAdd(ArrayRef< ConstantIntRanges > argRanges, OverflowFlags ovfFlags=OverflowFlags::None)
ConstantIntRanges inferMul(ArrayRef< ConstantIntRanges > argRanges, OverflowFlags ovfFlags=OverflowFlags::None)
ConstantIntRanges inferShapedDimOpInterface(ShapedDimOpInterface op, const IntegerValueRange &maybeDim)
Returns the integer range for the result of a ShapedDimOpInterface given the optional inferred ranges...
Type getTensorTypeFromMemRefType(Type type)
Return an unranked/ranked tensor type for the given unranked/ranked memref type.
Definition MemRefOps.cpp:63
OpFoldResult getMixedSize(OpBuilder &builder, Location loc, Value value, int64_t dim)
Return the dimension of the given memref value.
Definition MemRefOps.cpp:71
LogicalResult foldMemRefCast(Operation *op, Value inner=nullptr)
This is a common utility used for patterns of the form "someop(memref.cast) -> someop".
Definition MemRefOps.cpp:48
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given memref value.
Definition MemRefOps.cpp:80
Value createCanonicalRankReducingSubViewOp(OpBuilder &b, Location loc, Value memref, ArrayRef< int64_t > targetShape)
Create a rank-reducing SubViewOp @[0 .
Operation::operand_range getIndices(Operation *op)
Get the indices that the given load/store operation is operating on.
Definition Utils.cpp:18
DynamicAPInt getIndex(const ConeV &cone)
Get the index of a cone, i.e., the volume of the parallelepiped spanned by its generators,...
Definition Barvinok.cpp:63
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Value constantIndex(OpBuilder &builder, Location loc, int64_t i)
Generates a constant of index type.
MemRefType getMemRefType(T &&t)
Convenience method to abbreviate casting getType().
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
SmallVector< OpFoldResult > getMixedValues(ArrayRef< int64_t > staticValues, ValueRange dynamicValues, MLIRContext *context)
Return a vector of OpFoldResults with the same size a staticValues, but all elements for which Shaped...
detail::constant_int_value_binder m_ConstantInt(IntegerAttr::ValueType *bind_value)
Matches a constant holding a scalar/vector/tensor integer (splat) and writes the integer value to bin...
Definition Matchers.h:527
SliceVerificationResult
Enum that captures information related to verifier error conditions on slice insert/extract type of o...
detail::DenseArrayAttrImpl< int64_t > DenseI64ArrayAttr
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
raw_ostream & operator<<(raw_ostream &os, const AliasResult &result)
llvm::function_ref< void(Value, const IntegerValueRange &)> SetIntLatticeFn
Similar to SetIntRangeFn, but operating on IntegerValueRange lattice values.
SliceBoundsVerificationResult verifyInBoundsSlice(ArrayRef< int64_t > shape, ArrayRef< int64_t > staticOffsets, ArrayRef< int64_t > staticSizes, ArrayRef< int64_t > staticStrides, bool generateErrorMessage=false)
Verify that the offsets/sizes/strides-style access into the given shape is in-bounds.
LogicalResult verifyDynamicDimensionCount(Operation *op, ShapedType type, ValueRange dynamicSizes)
Verify that the number of dynamic size operands matches the number of dynamic dimensions in the shape...
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
Definition Utils.cpp:311
SmallVector< Range, 8 > getOrCreateRanges(OffsetSizeAndStrideOpInterface op, OpBuilder &b, Location loc)
Return the list of Range (i.e.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
SmallVector< AffineMap, 4 > getSymbolLessAffineMaps(ArrayRef< ReassociationExprs > reassociation)
Constructs affine maps out of Array<Array<AffineExpr>>.
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
OpFoldResult foldReshapeOp(ReshapeOpTy reshapeOp, ArrayRef< Attribute > operands)
bool hasValidSizesOffsets(SmallVector< int64_t > sizesOrOffsets)
Helper function to check whether the passed in sizes or offsets are valid.
SmallVector< SmallVector< OpFoldResult > > ReifiedRankedShapedTypeDims
SmallVector< IntegerValueRange > getIntValueRanges(ArrayRef< OpFoldResult > values, GetIntRangeFn getIntRange, int32_t indexBitwidth)
Helper function to collect the integer range values of an array of op fold results.
std::conditional_t< std::is_same_v< Ty, mlir::Type >, mlir::Value, detail::TypedValue< Ty > > TypedValue
If Ty is mlir::Type this will select Value instead of having a wrapper around it.
Definition Value.h:494
bool isZeroInteger(OpFoldResult v)
Return "true" if v is an integer value/attribute with constant value 0.
bool hasValidStrides(SmallVector< int64_t > strides)
Helper function to check whether the passed in strides are valid.
void dispatchIndexOpFoldResults(ArrayRef< OpFoldResult > ofrs, SmallVectorImpl< Value > &dynamicVec, SmallVectorImpl< int64_t > &staticVec)
Helper function to dispatch multiple OpFoldResults according to the behavior of dispatchIndexOpFoldRe...
SmallVector< SmallVector< AffineExpr, 2 >, 2 > convertReassociationIndicesToExprs(MLIRContext *context, ArrayRef< ReassociationIndices > reassociationIndices)
Convert reassociation indices to affine expressions.
std::optional< SmallVector< OpFoldResult > > inferExpandShapeOutputShape(OpBuilder &b, Location loc, ShapedType expandedType, ArrayRef< ReassociationIndices > reassociation, ArrayRef< OpFoldResult > inputShape)
Infer the output shape for a {memref|tensor}.expand_shape when it is possible to do so.
Definition Utils.cpp:26
LogicalResult verifyElementTypesMatch(Operation *op, ShapedType lhs, ShapedType rhs, StringRef lhsName, StringRef rhsName)
Verify that two shaped types have matching element types.
SmallVector< T > applyPermutationMap(AffineMap map, llvm::ArrayRef< T > source)
Apply a permutation from map to source and return the result.
Definition AffineMap.h:675
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
function_ref< void(Value, const StridedMetadataRange &)> SetStridedMetadataRangeFn
Callback function type for setting the strided metadata of a value.
std::optional< llvm::SmallDenseSet< unsigned > > computeRankReductionMask(ArrayRef< int64_t > originalShape, ArrayRef< int64_t > reducedShape, bool matchDynamic=false)
Given an originalShape and a reducedShape assumed to be a subset of originalShape with some 1 entries...
SmallVector< int64_t, 2 > ReassociationIndices
Definition Utils.h:27
SliceVerificationResult isRankReducedType(ShapedType originalType, ShapedType candidateReducedType)
Check if originalType can be rank reduced to candidateReducedType type by dropping some dimensions wi...
ArrayAttr getReassociationIndicesAttribute(Builder &b, ArrayRef< ReassociationIndices > reassociation)
Wraps a list of reassociations in an ArrayAttr.
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
bool isOneInteger(OpFoldResult v)
Return true if v is an IntegerAttr with value 1.
std::pair< SmallVector< int64_t >, SmallVector< Value > > decomposeMixedValues(ArrayRef< OpFoldResult > mixedValues)
Decompose a vector of mixed static or dynamic values into the corresponding pair of arrays.
LogicalResult verifyReassociationIndicesNotEmpty(ReshapeOpTy op)
Verify that none of the reassociation groups is empty.
function_ref< IntegerValueRange(Value)> GetIntRangeFn
Helper callback type to get the integer range of a value.
Move allocations into an allocation scope, if it is legal to move them (e.g.
LogicalResult matchAndRewrite(AllocaScopeOp op, PatternRewriter &rewriter) const override
Inline an AllocaScopeOp if either the direct parent is an allocation scope or it contains no allocati...
LogicalResult matchAndRewrite(AllocaScopeOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(CollapseShapeOp op, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(ExpandShapeOp op, PatternRewriter &rewriter) const override
A canonicalizer wrapper to replace SubViewOps.
void operator()(PatternRewriter &rewriter, SubViewOp op, SubViewOp newOp)
Return the canonical type of the result of a subview.
MemRefType operator()(SubViewOp op, ArrayRef< OpFoldResult > mixedOffsets, ArrayRef< OpFoldResult > mixedSizes, ArrayRef< OpFoldResult > mixedStrides)
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
This represents an operation in an abstracted form, suitable for use with the builder APIs.
Represents a range (offset, size, and stride) where each element of the triple may be dynamic or stat...
OpFoldResult stride
OpFoldResult size
OpFoldResult offset
static SaturatedInteger wrap(int64_t v)
bool isValid
If set to "true", the slice bounds verification was successful.
std::string errorMessage
An error message that can be printed during op verification.
Eliminates variable at the specified position using Fourier-Motzkin variable elimination.