MLIR 24.0.0git
BufferizationOps.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
16#include "mlir/IR/Matchers.h"
17#include "llvm/ADT/SmallVectorExtras.h"
18#include <optional>
19
20using namespace mlir;
21using namespace mlir::bufferization;
22
23//===----------------------------------------------------------------------===//
24// Helper functions
25//===----------------------------------------------------------------------===//
26
28 OpBuilder &b, Value value, MemRefType destType,
30 auto srcType = llvm::cast<MemRefType>(value.getType());
31
32 // Element type and rank must match.
33 if (srcType.getElementType() != destType.getElementType())
34 return failure();
35 if (srcType.getRank() != destType.getRank())
36 return failure();
37
38 // In case the affine maps are different, we may need to use a copy if we go
39 // from dynamic to static offset or stride (the canonicalization cannot know
40 // at this point that it is really cast compatible).
41 auto isGuaranteedCastCompatible = [](MemRefType source, MemRefType target) {
42 int64_t sourceOffset, targetOffset;
43 SmallVector<int64_t, 4> sourceStrides, targetStrides;
44 if (failed(source.getStridesAndOffset(sourceStrides, sourceOffset)) ||
45 failed(target.getStridesAndOffset(targetStrides, targetOffset)))
46 return false;
47 auto dynamicToStatic = [](int64_t a, int64_t b) {
48 return ShapedType::isDynamic(a) && ShapedType::isStatic(b);
49 };
50 if (dynamicToStatic(sourceOffset, targetOffset))
51 return false;
52 for (auto it : zip(sourceStrides, targetStrides))
53 if (dynamicToStatic(std::get<0>(it), std::get<1>(it)))
54 return false;
55 // A cast cannot safely zero out a non-zero offset. If the source has a
56 // non-identity layout (e.g., a subview with a dynamic offset) and the
57 // destination requires identity layout (zero offset, unit strides), a cast
58 // would silently strip the offset and cause incorrect loads in the callee.
59 // Fall through to the alloc+copy path instead.
60 if (!source.getLayout().isIdentity() && target.getLayout().isIdentity())
61 return false;
62 return true;
63 };
64
65 // Note: If `areCastCompatible`, a cast is valid, but may fail at runtime. To
66 // ensure that we only generate casts that always succeed at runtime, we check
67 // a few extra conditions in `isGuaranteedCastCompatible`.
68 if (memref::CastOp::areCastCompatible(srcType, destType) &&
69 isGuaranteedCastCompatible(srcType, destType)) {
70 Value casted = *options.castFn(b, value.getLoc(), destType, value);
71 return casted;
72 }
73
74 auto loc = value.getLoc();
75 SmallVector<Value, 4> dynamicOperands;
76 for (int i = 0; i < destType.getRank(); ++i) {
77 if (destType.getShape()[i] != ShapedType::kDynamic)
78 continue;
79 Value size = memref::DimOp::create(b, loc, value, i);
80 dynamicOperands.push_back(size);
81 }
82
83 FailureOr<Value> copy = options.allocationFn(
84 b, loc, destType, dynamicOperands, options.bufferAlignment);
85 if (failed(copy))
86 return failure();
87 if (failed(options.memCpyFn(b, loc, value, *copy)))
88 return failure();
89 return copy;
90}
91
92/// Try to fold to_buffer(to_tensor(x)). If x's type and the result type of the
93/// to_buffer op are different, a memref.cast is needed.
95 RewriterBase &rewriter, ToBufferOp toBuffer,
97 auto bufferToTensor = toBuffer.getTensor().getDefiningOp<ToTensorOp>();
98 if (!bufferToTensor)
99 return failure();
100
101 Type srcType = bufferToTensor.getBuffer().getType();
102 Type destType = toBuffer.getType();
103
104 // Directly rewrite if the type did not change.
105 if (srcType == destType) {
106 rewriter.replaceOp(toBuffer, bufferToTensor.getBuffer());
107 return success();
108 }
109
110 if (!llvm::isa<BaseMemRefType>(srcType) ||
111 !llvm::isa<BaseMemRefType>(destType)) {
112 // Non-builtin case: the best is to try the user-provided cast.
113 auto replacement =
114 options.castFn(rewriter, bufferToTensor.getBuffer().getLoc(), destType,
115 bufferToTensor.getBuffer());
116 if (failed(replacement))
117 return failure();
118 rewriter.replaceOp(toBuffer, *replacement);
119 return success();
120 }
121
122 auto rankedSrcType = llvm::dyn_cast<MemRefType>(srcType);
123 auto rankedDestType = llvm::dyn_cast<MemRefType>(destType);
124 auto unrankedSrcType = llvm::dyn_cast<UnrankedMemRefType>(srcType);
125
126 // Ranked memref -> Ranked memref cast.
127 if (rankedSrcType && rankedDestType) {
128 FailureOr<Value> replacement = castOrReallocMemRefValue(
129 rewriter, bufferToTensor.getBuffer(), rankedDestType, options);
130 if (failed(replacement))
131 return failure();
132
133 rewriter.replaceOp(toBuffer, *replacement);
134 return success();
135 }
136
137 // Unranked memref -> Ranked memref cast: May require a copy.
138 // TODO: Not implemented at the moment.
139 if (unrankedSrcType && rankedDestType)
140 return failure();
141
142 // Unranked/ranked memref -> unranked memref cast: No copy needed if the types
143 // are cast-compatible.
144 if (!memref::CastOp::areCastCompatible(srcType, destType))
145 return failure();
146
147 rewriter.replaceOpWithNewOp<memref::CastOp>(toBuffer, destType,
148 bufferToTensor.getBuffer());
149 return success();
150}
151
153 OpBuilder &b, Location loc, Value shapedValue,
154 SmallVector<Value> &dynamicDims) {
155 auto shapedType = llvm::cast<ShapedType>(shapedValue.getType());
156 for (int64_t i = 0; i < shapedType.getRank(); ++i) {
157 if (shapedType.isDynamicDim(i)) {
158 if (llvm::isa<MemRefType>(shapedType)) {
159 dynamicDims.push_back(memref::DimOp::create(b, loc, shapedValue, i));
160 } else {
161 assert(llvm::isa<RankedTensorType>(shapedType) && "expected tensor");
162 dynamicDims.push_back(tensor::DimOp::create(b, loc, shapedValue, i));
163 }
164 }
165 }
166}
167
168//===----------------------------------------------------------------------===//
169// AllocTensorOp
170//===----------------------------------------------------------------------===//
171
172LogicalResult AllocTensorOp::verify() {
173 if (getCopy() && !getDynamicSizes().empty())
174 return emitError("dynamic sizes not needed when copying a tensor");
175 if (!getCopy() && failed(verifyDynamicDimensionCount(
176 getOperation(), getType(), getDynamicSizes())))
177 return failure();
178 if (getCopy() && getCopy().getType() != getType())
179 return emitError("expected that `copy` and return type match");
180 return success();
181}
182
183void AllocTensorOp::build(OpBuilder &builder, OperationState &result,
184 RankedTensorType type, ValueRange dynamicSizes) {
185 build(builder, result, type, dynamicSizes, /*copy=*/Value(),
186 /*size_hint=*/Value(),
187 /*memory_space=*/IntegerAttr());
188}
189
190void AllocTensorOp::build(OpBuilder &builder, OperationState &result,
191 RankedTensorType type, ValueRange dynamicSizes,
192 Value copy) {
193 build(builder, result, type, dynamicSizes, copy, /*size_hint=*/Value(),
194 /*memory_space=*/IntegerAttr());
195}
196
197void AllocTensorOp::build(OpBuilder &builder, OperationState &result,
198 TensorType type, ValueRange dynamicSizes, Value copy,
199 IntegerAttr memorySpace) {
200 build(builder, result, type, dynamicSizes, copy, /*size_hint=*/Value(),
201 memorySpace);
202}
203
204namespace {
205/// Change the type of the result of a `bufferization.alloc_tensor` by making
206/// the result type statically sized along dimension that in the original
207/// operation where defined as dynamic, but the size was defined using a
208/// `constant` op. For example:
209///
210/// %c5 = arith.constant 5: index
211/// %0 = bufferization.alloc_tensor(%arg0, %c5) : tensor<?x?xf32>
212///
213/// to
214///
215/// %0 = bufferization.alloc_tensor(%arg0) : tensor<?x5xf32>
216struct ReplaceStaticShapeDims : OpRewritePattern<AllocTensorOp> {
217 using OpRewritePattern<AllocTensorOp>::OpRewritePattern;
218
219 LogicalResult matchAndRewrite(AllocTensorOp op,
220 PatternRewriter &rewriter) const override {
221 if (op.getCopy())
222 return failure();
223 SmallVector<int64_t> newShape = llvm::to_vector(op.getType().getShape());
224 SmallVector<Value> newDynamicSizes;
225 unsigned int dynValCounter = 0;
226 for (int64_t i = 0; i < op.getType().getRank(); ++i) {
227 if (!op.isDynamicDim(i))
228 continue;
229 Value value = op.getDynamicSizes()[dynValCounter++];
230 APInt intVal;
231 if (matchPattern(value, m_ConstantInt(&intVal))) {
232 int64_t dim = intVal.getSExtValue();
233 if (dim >= 0)
234 newShape[i] = intVal.getSExtValue();
235 else
236 newDynamicSizes.push_back(value);
237 } else {
238 newDynamicSizes.push_back(value);
239 }
240 }
241 RankedTensorType newType = RankedTensorType::get(
242 newShape, op.getType().getElementType(), op.getType().getEncoding());
243 if (newType == op.getType())
244 return failure();
245 auto newOp = AllocTensorOp::create(rewriter, op.getLoc(), newType,
246 newDynamicSizes, /*copy=*/Value());
247 rewriter.replaceOpWithNewOp<tensor::CastOp>(op, op.getType(), newOp);
248 return success();
249 }
250};
251
252struct FoldDimOfAllocTensorOp : public OpRewritePattern<tensor::DimOp> {
253 using OpRewritePattern<tensor::DimOp>::OpRewritePattern;
254
255 LogicalResult matchAndRewrite(tensor::DimOp dimOp,
256 PatternRewriter &rewriter) const override {
257 std::optional<int64_t> maybeConstantIndex = dimOp.getConstantIndex();
258 auto allocTensorOp = dimOp.getSource().getDefiningOp<AllocTensorOp>();
259 if (!allocTensorOp || !maybeConstantIndex)
260 return failure();
261 if (*maybeConstantIndex < 0 ||
262 *maybeConstantIndex >= allocTensorOp.getType().getRank())
263 return failure();
264 if (!allocTensorOp.getType().isDynamicDim(*maybeConstantIndex))
265 return failure();
266 rewriter.replaceOp(
267 dimOp, allocTensorOp.getDynamicSize(rewriter, *maybeConstantIndex));
268 return success();
269 }
270};
271} // namespace
272
273void AllocTensorOp::getCanonicalizationPatterns(RewritePatternSet &results,
274 MLIRContext *ctx) {
275 results.add<FoldDimOfAllocTensorOp, ReplaceStaticShapeDims>(ctx);
276}
277
278LogicalResult AllocTensorOp::reifyResultShapes(
279 OpBuilder &builder, ReifiedRankedShapedTypeDims &reifiedReturnShapes) {
280 auto shapes =
281 llvm::map_to_vector<4>(llvm::seq<int64_t>(0, getType().getRank()),
282 [&](int64_t dim) -> OpFoldResult {
283 if (isDynamicDim(dim))
284 return getDynamicSize(builder, dim);
285 return builder.getIndexAttr(getStaticSize(dim));
286 });
287 reifiedReturnShapes.emplace_back(std::move(shapes));
288 return success();
289}
290
291ParseResult AllocTensorOp::parse(OpAsmParser &parser, OperationState &result) {
293 if (parser.parseLParen() || parser.parseOperandList(dynamicSizesOperands) ||
294 parser.parseRParen())
295 return failure();
296 ParseResult copyKeyword = parser.parseOptionalKeyword("copy");
298 if (copyKeyword.succeeded())
299 if (parser.parseLParen() || parser.parseOperand(copyOperand) ||
300 parser.parseRParen())
301 return failure();
302 ParseResult sizeHintKeyword = parser.parseOptionalKeyword("size_hint");
303 OpAsmParser::UnresolvedOperand sizeHintOperand;
304 if (sizeHintKeyword.succeeded())
305 if (parser.parseEqual() || parser.parseOperand(sizeHintOperand))
306 return failure();
307
308 Attribute parsedProperties;
309 if (AllocTensorOp::genericParseProperties(parser, parsedProperties))
310 return failure();
311 auto propertyDictionary = dyn_cast_or_null<DictionaryAttr>(parsedProperties);
312 if (parsedProperties && !propertyDictionary)
313 return parser.emitError(parser.getNameLoc(),
314 "expected properties dictionary");
315
316 auto attrsLoc = parser.getCurrentLocation();
317 if (parser.parseOptionalAttrDict(result.attributes))
318 return failure();
319 for (StringRef attrName : AllocTensorOp::getAttributeNames()) {
320 if (result.attributes.get(attrName))
321 return parser.emitError(attrsLoc)
322 << "inherent attribute '" << attrName
323 << "' cannot be parsed from attr-dict when strict properties in "
324 "assembly format is enabled";
325 }
326 if (parser.parseColon())
327 return failure();
328
329 TensorType type;
330 if (parser.parseCustomTypeWithFallback(type))
331 return failure();
332 result.addTypes(type);
333
334 Type indexType = parser.getBuilder().getIndexType();
335 if (parser.resolveOperands(dynamicSizesOperands, indexType, result.operands))
336 return failure();
337 if (copyKeyword.succeeded())
338 if (parser.resolveOperand(copyOperand, type, result.operands))
339 return failure();
340 if (sizeHintKeyword.succeeded())
341 if (parser.resolveOperand(sizeHintOperand, indexType, result.operands))
342 return failure();
343 Builder &builder = parser.getBuilder();
344 NamedAttrList properties(propertyDictionary ? propertyDictionary
345 : builder.getDictionaryAttr({}));
346 properties.set(AllocTensorOp::getOperandSegmentSizeAttr(),
347 builder.getDenseI32ArrayAttr(
348 {static_cast<int32_t>(dynamicSizesOperands.size()),
349 static_cast<int32_t>(copyKeyword.succeeded()),
350 static_cast<int32_t>(sizeHintKeyword.succeeded())}));
351 propertyDictionary = properties.getDictionary(builder.getContext());
352 auto emitError = [&]() {
353 return mlir::emitError(result.location, "invalid properties ")
354 << propertyDictionary << " for op " << result.name.getStringRef()
355 << ": ";
356 };
357 if (failed(AllocTensorOp::setPropertiesFromParsedAttr(
358 result.getOrAddProperties<Properties>(), propertyDictionary,
359 emitError)))
360 return failure();
361 return success();
362}
363
364void AllocTensorOp::print(OpAsmPrinter &p) {
365 p << "(" << getDynamicSizes() << ")";
366 if (getCopy())
367 p << " copy(" << getCopy() << ")";
368 if (getSizeHint())
369 p << " size_hint=" << getSizeHint();
370 AllocTensorOp::printProperties(getContext(), p, getProperties(),
371 /*elidedProps=*/getOperandSegmentSizeAttr());
372 p.printOptionalAttrDict((*this)->getDiscardableAttrDictionary().getValue());
373 p << " : ";
374 auto type = getResult().getType();
375 if (auto validType = llvm::dyn_cast<::mlir::TensorType>(type))
376 p.printStrippedAttrOrType(validType);
377 else
378 p << type;
379}
380
381Value AllocTensorOp::getDynamicSize(OpBuilder &b, unsigned idx) {
382 assert(isDynamicDim(idx) && "expected dynamic dim");
383 if (getCopy())
384 return tensor::DimOp::create(b, getLoc(), getCopy(), idx);
385 return getOperand(getIndexOfDynamicSize(idx));
386}
387
388//===----------------------------------------------------------------------===//
389// CloneOp
390//===----------------------------------------------------------------------===//
391
392OpFoldResult CloneOp::fold(FoldAdaptor adaptor) {
393 return succeeded(memref::foldMemRefCast(*this)) ? getResult() : Value();
394}
395
396namespace {
397
398/// Merge the clone and its source (by converting the clone to a cast) when
399/// possible.
400struct SimplifyClones : public OpRewritePattern<CloneOp> {
401 using OpRewritePattern<CloneOp>::OpRewritePattern;
402
403 LogicalResult matchAndRewrite(CloneOp cloneOp,
404 PatternRewriter &rewriter) const override {
405 if (cloneOp.use_empty()) {
406 rewriter.eraseOp(cloneOp);
407 return success();
408 }
409
410 Value source = cloneOp.getInput();
411 if (source.getType() != cloneOp.getType() &&
412 !memref::CastOp::areCastCompatible({source.getType()},
413 {cloneOp.getType()}))
414 return failure();
415
416 // Aims to find the dealloc op for the canonical source
417 // which otherwise could prevent removal of unnecessary allocs.
418 Value canonicalSource = source;
419 while (auto iface = dyn_cast_or_null<ViewLikeOpInterface>(
420 canonicalSource.getDefiningOp())) {
421 if (canonicalSource != iface.getViewDest()) {
422 break;
423 }
424 canonicalSource = iface.getViewSource();
425 }
426
427 std::optional<Operation *> maybeCloneDeallocOp =
428 memref::findDealloc(cloneOp.getOutput());
429 // Skip if either of them has > 1 deallocate operations.
430 if (!maybeCloneDeallocOp.has_value())
431 return failure();
432 std::optional<Operation *> maybeSourceDeallocOp =
433 memref::findDealloc(canonicalSource);
434 if (!maybeSourceDeallocOp.has_value())
435 return failure();
436 Operation *cloneDeallocOp = *maybeCloneDeallocOp;
437 Operation *sourceDeallocOp = *maybeSourceDeallocOp;
438
439 // If both are deallocated in the same block, their in-block lifetimes
440 // might not fully overlap, so we cannot decide which one to drop.
441 if (cloneDeallocOp && sourceDeallocOp &&
442 cloneDeallocOp->getBlock() == sourceDeallocOp->getBlock())
443 return failure();
444
445 Block *currentBlock = cloneOp->getBlock();
446 Operation *redundantDealloc = nullptr;
447 if (cloneDeallocOp && cloneDeallocOp->getBlock() == currentBlock) {
448 redundantDealloc = cloneDeallocOp;
449 } else if (sourceDeallocOp && sourceDeallocOp->getBlock() == currentBlock) {
450 redundantDealloc = sourceDeallocOp;
451 }
452
453 if (!redundantDealloc)
454 return failure();
455
456 // Safety check that there are no other deallocations inbetween
457 // cloneOp and redundantDealloc, as otherwise we might deallocate an alias
458 // of source before the uses of the clone. With alias information, we could
459 // restrict this to only fail of the dealloc's operand is an alias
460 // of the source.
461 for (Operation *pos = cloneOp->getNextNode(); pos != redundantDealloc;
462 pos = pos->getNextNode()) {
463 // Bail if we run out of operations while looking for a deallocation op.
464 if (!pos)
465 return failure();
466 auto effectInterface = dyn_cast<MemoryEffectOpInterface>(pos);
467 if (!effectInterface)
468 continue;
469 if (effectInterface.hasEffect<MemoryEffects::Free>())
470 return failure();
471 }
472
473 if (source.getType() != cloneOp.getType())
474 source = memref::CastOp::create(rewriter, cloneOp.getLoc(),
475 cloneOp.getType(), source);
476 rewriter.replaceOp(cloneOp, source);
477 rewriter.eraseOp(redundantDealloc);
478 return success();
479 }
480};
481
482} // namespace
483
484void CloneOp::getCanonicalizationPatterns(RewritePatternSet &results,
485 MLIRContext *context) {
486 results.add<SimplifyClones>(context);
487}
488
489//===----------------------------------------------------------------------===//
490// MaterializeInDestinationOp
491//===----------------------------------------------------------------------===//
492
493LogicalResult MaterializeInDestinationOp::reifyResultShapes(
494 OpBuilder &builder, ReifiedRankedShapedTypeDims &reifiedReturnShapes) {
495 if (getOperation()->getNumResults() == 1) {
496 assert(isa<TensorType>(getDest().getType()) && "expected tensor type");
497 reifiedReturnShapes.resize(1,
499 reifiedReturnShapes[0] =
500 tensor::getMixedSizes(builder, getLoc(), getDest());
501 }
502 return success();
503}
504
505Value MaterializeInDestinationOp::buildSubsetExtraction(OpBuilder &builder,
506 Location loc) {
507 if (isa<TensorType>(getDest().getType())) {
508 // The subset is the entire destination tensor.
509 return getDest();
510 }
511
512 // The "restrict" attribute is transferred from this op to the newly created
513 // to_tensor op. If this op does not the "restrict" attribute, the subset
514 // extraction cannot be built because there is no guarantee that there is no
515 // pre-existing "restrict" to_tensor op with the same/an aliasing destination.
516 if (!getRestrict())
517 return {};
518
519 // Build a bufferization.to_tensor op.
520 assert(isa<BaseMemRefType>(getDest().getType()) && "expected memref type");
521 assert(getRestrict() &&
522 "expected that ops with memrefs dest have 'restrict'");
523 setRestrict(false);
524 return ToTensorOp::create(
525 builder, loc, memref::getTensorTypeFromMemRefType(getDest().getType()),
526 getDest(),
527 /*restrict=*/true, getWritable());
528}
529
530bool MaterializeInDestinationOp::isEquivalentSubset(
531 Value candidate, function_ref<bool(Value, Value)> equivalenceFn) {
532 return equivalenceFn(getDest(), candidate);
533}
534
536MaterializeInDestinationOp::getValuesNeededToBuildSubsetExtraction() {
537 return {getDest()};
538}
539
540OpOperand &MaterializeInDestinationOp::getSourceOperand() {
541 return getOperation()->getOpOperand(0) /*source*/;
542}
543
544bool MaterializeInDestinationOp::operatesOnEquivalentSubset(
545 SubsetOpInterface subsetOp,
546 function_ref<bool(Value, Value)> equivalenceFn) {
547 return false;
548}
549
550bool MaterializeInDestinationOp::operatesOnDisjointSubset(
551 SubsetOpInterface subsetOp,
552 function_ref<bool(Value, Value)> equivalenceFn) {
553 return false;
554}
555
556LogicalResult MaterializeInDestinationOp::verify() {
557 if (!isa<TensorType, BaseMemRefType>(getDest().getType()))
558 return emitOpError("'dest' must be a tensor or a memref");
559 if (auto destType = dyn_cast<TensorType>(getDest().getType())) {
560 if (getOperation()->getNumResults() != 1)
561 return emitOpError("tensor 'dest' implies exactly one tensor result");
562 if (destType != getResult().getType())
563 return emitOpError("result and 'dest' types must match");
564 }
565 if (isa<BaseMemRefType>(getDest().getType()) &&
566 getOperation()->getNumResults() != 0)
567 return emitOpError("memref 'dest' implies zero results");
568 if (getRestrict() && !isa<BaseMemRefType>(getDest().getType()))
569 return emitOpError("'restrict' is valid only for memref destinations");
570 if (getWritable() != isa<BaseMemRefType>(getDest().getType()))
571 return emitOpError("'writable' must be specified if and only if the "
572 "destination is of memref type");
573 TensorType srcType = getSource().getType();
574 ShapedType destType = cast<ShapedType>(getDest().getType());
575 if (srcType.hasRank() != destType.hasRank())
576 return emitOpError("source/destination shapes are incompatible");
577 if (srcType.hasRank()) {
578 if (failed(verifyRanksMatch(getOperation(), srcType, destType, "source",
579 "destination")))
580 return failure();
581 for (auto [src, dest] :
582 llvm::zip(srcType.getShape(), destType.getShape())) {
583 if (src == ShapedType::kDynamic || dest == ShapedType::kDynamic) {
584 // Cannot verify dynamic dimension size. Assume that that they match at
585 // runtime.
586 continue;
587 }
588 if (src != dest)
589 return emitOpError("source/destination shapes are incompatible");
590 }
591 }
592 return success();
593}
594
595void MaterializeInDestinationOp::build(OpBuilder &builder,
596 OperationState &state, Value source,
597 Value dest) {
598 auto destTensorType = dyn_cast<TensorType>(dest.getType());
599 build(builder, state, /*result=*/destTensorType ? destTensorType : Type(),
600 source, dest);
601}
602
603MutableOperandRange MaterializeInDestinationOp::getDpsInitsMutable() {
604 return getDestMutable();
605}
606
607void MaterializeInDestinationOp::getEffects(
609 &effects) {
610 if (isa<BaseMemRefType>(getDest().getType()))
611 effects.emplace_back(MemoryEffects::Write::get(), &getDestMutable(),
613}
614
615//===----------------------------------------------------------------------===//
616// ToTensorOp
617//===----------------------------------------------------------------------===//
618
619OpFoldResult ToTensorOp::fold(FoldAdaptor) {
620 if (auto toBuffer = getBuffer().getDefiningOp<ToBufferOp>())
621 // Approximate alias analysis by conservatively folding only when no there
622 // is no interleaved operation.
623 if (toBuffer->getBlock() == this->getOperation()->getBlock() &&
624 toBuffer->getNextNode() == this->getOperation())
625 return toBuffer.getTensor();
626 return {};
627}
628
629namespace {
630struct DimOfToTensorFolder : public OpRewritePattern<tensor::DimOp> {
631 using OpRewritePattern<tensor::DimOp>::OpRewritePattern;
632
633 LogicalResult matchAndRewrite(tensor::DimOp dimOp,
634 PatternRewriter &rewriter) const override {
635 auto memrefToTensorOp = dimOp.getSource().getDefiningOp<ToTensorOp>();
636 if (!memrefToTensorOp)
637 return failure();
638
639 rewriter.replaceOpWithNewOp<memref::DimOp>(
640 dimOp, memrefToTensorOp.getBuffer(), dimOp.getIndex());
641 return success();
642 }
643};
644} // namespace
645
646void ToTensorOp::getCanonicalizationPatterns(RewritePatternSet &results,
647 MLIRContext *context) {
648 results.add<DimOfToTensorFolder>(context);
649}
650
651//===----------------------------------------------------------------------===//
652// ToBufferOp
653//===----------------------------------------------------------------------===//
654
655OpFoldResult ToBufferOp::fold(FoldAdaptor) {
656 if (auto memrefToTensor = getTensor().getDefiningOp<ToTensorOp>())
657 if (memrefToTensor.getBuffer().getType() == getType())
658 return memrefToTensor.getBuffer();
659 return {};
660}
661
662namespace {
663
664/// Replace tensor.cast + to_buffer by to_buffer + memref.cast.
665struct ToBufferOfCast : public OpRewritePattern<ToBufferOp> {
666 using OpRewritePattern<ToBufferOp>::OpRewritePattern;
667
668 LogicalResult matchAndRewrite(ToBufferOp toBuffer,
669 PatternRewriter &rewriter) const final {
670 auto tensorCastOperand =
671 toBuffer.getOperand().getDefiningOp<tensor::CastOp>();
672 if (!tensorCastOperand)
673 return failure();
674 auto srcTensorType = llvm::dyn_cast<RankedTensorType>(
675 tensorCastOperand.getOperand().getType());
676 if (!srcTensorType)
677 return failure();
678 auto currentOutputMemRefType =
679 dyn_cast<BaseMemRefType>(toBuffer.getResult().getType());
680 if (!currentOutputMemRefType)
681 return failure();
682
683 auto memrefType = currentOutputMemRefType.cloneWith(
684 srcTensorType.getShape(), srcTensorType.getElementType());
685 Value memref = ToBufferOp::create(rewriter, toBuffer.getLoc(), memrefType,
686 tensorCastOperand.getOperand(),
687 toBuffer.getReadOnly());
688 rewriter.replaceOpWithNewOp<memref::CastOp>(toBuffer, toBuffer.getType(),
689 memref);
690 return success();
691 }
692};
693
694/// Canonicalize bufferization.to_tensor + bufferization.to_buffer. Insert a
695/// cast if necessary.
696struct ToBufferToTensorFolding : public OpRewritePattern<ToBufferOp> {
697 using OpRewritePattern<ToBufferOp>::OpRewritePattern;
698
699 LogicalResult matchAndRewrite(ToBufferOp toBuffer,
700 PatternRewriter &rewriter) const final {
701 BufferizationOptions options;
702 options.bufferAlignment = 0;
703 return foldToBufferToTensorPair(rewriter, toBuffer, options);
704 }
705};
706
707/// Fold a load on a to_buffer operation into an tensor.extract on the
708/// corresponding tensor.
709struct LoadOfToBuffer : public OpRewritePattern<memref::LoadOp> {
710 using OpRewritePattern<memref::LoadOp>::OpRewritePattern;
711
712 LogicalResult matchAndRewrite(memref::LoadOp load,
713 PatternRewriter &rewriter) const override {
714 auto toBuffer = load.getMemref().getDefiningOp<ToBufferOp>();
715 if (!toBuffer || !toBuffer.getReadOnly())
716 return failure();
717
718 rewriter.replaceOpWithNewOp<tensor::ExtractOp>(load, toBuffer.getTensor(),
719 load.getIndices());
720 return success();
721 }
722};
723
724/// Fold dim of a to_buffer into the dim of the tensor.
725struct DimOfCastOp : public OpRewritePattern<memref::DimOp> {
726 using OpRewritePattern<memref::DimOp>::OpRewritePattern;
727
728 LogicalResult matchAndRewrite(memref::DimOp dimOp,
729 PatternRewriter &rewriter) const override {
730 auto castOp = dimOp.getSource().getDefiningOp<ToBufferOp>();
731 if (!castOp)
732 return failure();
733 Value newSource = castOp.getOperand();
734 rewriter.replaceOpWithNewOp<tensor::DimOp>(dimOp, newSource,
735 dimOp.getIndex());
736 return success();
737 }
738};
739
740} // namespace
741
742void ToBufferOp::getCanonicalizationPatterns(RewritePatternSet &results,
743 MLIRContext *context) {
744 results.add<DimOfCastOp, LoadOfToBuffer, ToBufferOfCast,
745 ToBufferToTensorFolding>(context);
746}
747
748std::optional<Operation *> CloneOp::buildDealloc(OpBuilder &builder,
749 Value alloc) {
750 return memref::DeallocOp::create(builder, alloc.getLoc(), alloc)
751 .getOperation();
752}
753
754std::optional<Value> CloneOp::buildClone(OpBuilder &builder, Value alloc) {
755 return CloneOp::create(builder, alloc.getLoc(), alloc).getResult();
756}
757
758//===----------------------------------------------------------------------===//
759// DeallocOp
760//===----------------------------------------------------------------------===//
761
762LogicalResult DeallocOp::inferReturnTypes(
763 MLIRContext *context, std::optional<::mlir::Location> location,
764 ValueRange operands, DictionaryAttr attributes, PropertyRef properties,
765 RegionRange regions, SmallVectorImpl<Type> &inferredReturnTypes) {
766 DeallocOpAdaptor adaptor(operands, attributes, properties, regions);
767 inferredReturnTypes = SmallVector<Type>(adaptor.getRetained().size(),
768 IntegerType::get(context, 1));
769 return success();
770}
771
772LogicalResult DeallocOp::verify() {
773 if (getMemrefs().size() != getConditions().size())
774 return emitOpError(
775 "must have the same number of conditions as memrefs to deallocate");
776 if (getRetained().size() != getUpdatedConditions().size())
777 return emitOpError("must have the same number of updated conditions "
778 "(results) as retained operands");
779 return success();
780}
781
782static LogicalResult updateDeallocIfChanged(DeallocOp deallocOp,
783 ValueRange memrefs,
784 ValueRange conditions,
785 PatternRewriter &rewriter) {
786 if (deallocOp.getMemrefs() == memrefs &&
787 deallocOp.getConditions() == conditions)
788 return failure();
789
790 rewriter.modifyOpInPlace(deallocOp, [&]() {
791 deallocOp.getMemrefsMutable().assign(memrefs);
792 deallocOp.getConditionsMutable().assign(conditions);
793 });
794 return success();
795}
796
797namespace {
798
799/// Remove duplicate values in the list of memrefs to be deallocated. We need to
800/// make sure the corresponding condition value is updated accordingly since
801/// their two conditions might not cover the same set of cases. In that case, we
802/// have to combine them (by computing the disjunction of them).
803/// Example:
804/// ```mlir
805/// bufferization.dealloc (%arg0, %arg0 : ...) if (%arg1, %arg2)
806/// ```
807/// is canonicalized to
808/// ```mlir
809/// %0 = arith.ori %arg1, %arg2 : i1
810/// bufferization.dealloc (%arg0 : memref<2xi32>) if (%0)
811/// ```
812struct DeallocRemoveDuplicateDeallocMemrefs
813 : public OpRewritePattern<DeallocOp> {
814 using OpRewritePattern<DeallocOp>::OpRewritePattern;
815
816 LogicalResult matchAndRewrite(DeallocOp deallocOp,
817 PatternRewriter &rewriter) const override {
818 // Unique memrefs to be deallocated.
819 DenseMap<Value, unsigned> memrefToCondition;
820 SmallVector<Value> newMemrefs, newConditions;
821 for (auto [i, memref, cond] :
822 llvm::enumerate(deallocOp.getMemrefs(), deallocOp.getConditions())) {
823 if (memrefToCondition.count(memref)) {
824 // If the dealloc conditions don't match, we need to make sure that the
825 // dealloc happens on the union of cases.
826 Value &newCond = newConditions[memrefToCondition[memref]];
827 if (newCond != cond)
828 newCond =
829 arith::OrIOp::create(rewriter, deallocOp.getLoc(), newCond, cond);
830 } else {
831 memrefToCondition.insert({memref, newConditions.size()});
832 newMemrefs.push_back(memref);
833 newConditions.push_back(cond);
834 }
835 }
836
837 // Return failure if we don't change anything such that we don't run into an
838 // infinite loop of pattern applications.
839 return updateDeallocIfChanged(deallocOp, newMemrefs, newConditions,
840 rewriter);
841 }
842};
843
844/// Remove duplicate values in the list of retained memrefs. We need to make
845/// sure the corresponding result condition value is replaced properly.
846/// Example:
847/// ```mlir
848/// %0:2 = bufferization.dealloc retain (%arg3, %arg3 : ...)
849/// ```
850/// is canonicalized to
851/// ```mlir
852/// %0 = bufferization.dealloc retain (%arg3 : memref<2xi32>)
853/// ```
854struct DeallocRemoveDuplicateRetainedMemrefs
855 : public OpRewritePattern<DeallocOp> {
856 using OpRewritePattern<DeallocOp>::OpRewritePattern;
857
858 LogicalResult matchAndRewrite(DeallocOp deallocOp,
859 PatternRewriter &rewriter) const override {
860 // Unique retained values
862 SmallVector<Value> newRetained;
863 SmallVector<unsigned> resultReplacementIdx;
864 unsigned i = 0;
865 for (auto retained : deallocOp.getRetained()) {
866 if (seen.count(retained)) {
867 resultReplacementIdx.push_back(seen[retained]);
868 continue;
869 }
870
871 seen[retained] = i;
872 newRetained.push_back(retained);
873 resultReplacementIdx.push_back(i++);
874 }
875
876 // Return failure if we don't change anything such that we don't run into an
877 // infinite loop of pattern applications.
878 if (newRetained.size() == deallocOp.getRetained().size())
879 return failure();
880
881 // We need to create a new op because the number of results is always the
882 // same as the number of condition operands.
883 auto newDeallocOp =
884 DeallocOp::create(rewriter, deallocOp.getLoc(), deallocOp.getMemrefs(),
885 deallocOp.getConditions(), newRetained);
886 SmallVector<Value> replacements(
887 llvm::map_range(resultReplacementIdx, [&](unsigned idx) {
888 return newDeallocOp.getUpdatedConditions()[idx];
889 }));
890 rewriter.replaceOp(deallocOp, replacements);
891 return success();
892 }
893};
894
895/// Erase deallocation operations where the variadic list of memrefs to
896/// deallocate is empty. Example:
897/// ```mlir
898/// %0 = bufferization.dealloc retain (%arg0: memref<2xi32>)
899/// ```
900struct EraseEmptyDealloc : public OpRewritePattern<DeallocOp> {
901 using OpRewritePattern<DeallocOp>::OpRewritePattern;
902
903 LogicalResult matchAndRewrite(DeallocOp deallocOp,
904 PatternRewriter &rewriter) const override {
905 if (deallocOp.getMemrefs().empty()) {
906 Value constFalse = arith::ConstantOp::create(rewriter, deallocOp.getLoc(),
907 rewriter.getBoolAttr(false));
908 rewriter.replaceOp(
909 deallocOp, SmallVector<Value>(deallocOp.getUpdatedConditions().size(),
910 constFalse));
911 return success();
912 }
913 return failure();
914 }
915};
916
917/// Removes memrefs from the deallocation list if their associated condition is
918/// always 'false'.
919///
920/// Example:
921/// ```
922/// bufferization.dealloc (%arg0, %arg1 : memref<2xi32>, memref<2xi32>)
923/// if (%arg2, %false)
924/// ```
925/// becomes
926/// ```
927/// bufferization.dealloc (%arg0 : memref<2xi32>) if (%arg2)
928/// ```
929struct EraseAlwaysFalseDealloc : public OpRewritePattern<DeallocOp> {
930 using OpRewritePattern<DeallocOp>::OpRewritePattern;
931
932 LogicalResult matchAndRewrite(DeallocOp deallocOp,
933 PatternRewriter &rewriter) const override {
934 SmallVector<Value> newMemrefs, newConditions;
935 for (auto [memref, cond] :
936 llvm::zip(deallocOp.getMemrefs(), deallocOp.getConditions())) {
937 if (!matchPattern(cond, m_Zero())) {
938 newMemrefs.push_back(memref);
939 newConditions.push_back(cond);
940 }
941 }
942
943 return updateDeallocIfChanged(deallocOp, newMemrefs, newConditions,
944 rewriter);
945 }
946};
947
948/// The `memref.extract_strided_metadata` is often inserted to get the base
949/// memref if the operand is not already guaranteed to be the result of a memref
950/// allocation operation. This canonicalization pattern removes this extraction
951/// operation if the operand is now produced by an allocation operation (e.g.,
952/// due to other canonicalizations simplifying the IR).
953///
954/// Example:
955/// ```mlir
956/// %alloc = memref.alloc() : memref<2xi32>
957/// %base_memref, %offset, %size, %stride = memref.extract_strided_metadata
958/// %alloc : memref<2xi32> -> memref<i32>, index, index, index
959/// bufferization.dealloc (%base_memref : memref<i32>) if (%cond)
960/// ```
961/// is canonicalized to
962/// ```mlir
963/// %alloc = memref.alloc() : memref<2xi32>
964/// bufferization.dealloc (%alloc : memref<2xi32>) if (%cond)
965/// ```
966struct SkipExtractMetadataOfAlloc : public OpRewritePattern<DeallocOp> {
967 using OpRewritePattern<DeallocOp>::OpRewritePattern;
968
969 LogicalResult matchAndRewrite(DeallocOp deallocOp,
970 PatternRewriter &rewriter) const override {
971 SmallVector<Value> newMemrefs(
972 llvm::map_range(deallocOp.getMemrefs(), [&](Value memref) {
973 auto extractStridedOp =
974 memref.getDefiningOp<memref::ExtractStridedMetadataOp>();
975 if (!extractStridedOp)
976 return memref;
977 Value allocMemref = extractStridedOp.getOperand();
978 auto allocOp = allocMemref.getDefiningOp<MemoryEffectOpInterface>();
979 if (!allocOp)
980 return memref;
981 if (allocOp.getEffectOnValue<MemoryEffects::Allocate>(allocMemref))
982 return allocMemref;
983 return memref;
984 }));
985
986 return updateDeallocIfChanged(deallocOp, newMemrefs,
987 deallocOp.getConditions(), rewriter);
988 }
989};
990
991/// Removes pairs of `bufferization.dealloc` and alloc operations if there is no
992/// other user of the allocated value and the allocating operation can be safely
993/// removed. If the same value is present multiple times, this pattern relies on
994/// other canonicalization patterns to remove the duplicate first.
995///
996/// Example:
997/// ```mlir
998/// %alloc = memref.alloc() : memref<2xi32>
999/// bufferization.dealloc (%alloc, %arg0, : ...) if (%true, %true)
1000/// ```
1001/// is canonicalized to
1002/// ```mlir
1003/// bufferization.dealloc (%arg0 : ...) if (%true)
1004/// ```
1005struct RemoveAllocDeallocPairWhenNoOtherUsers
1006 : public OpRewritePattern<DeallocOp> {
1007 using OpRewritePattern<DeallocOp>::OpRewritePattern;
1008
1009 LogicalResult matchAndRewrite(DeallocOp deallocOp,
1010 PatternRewriter &rewriter) const override {
1011 SmallVector<Value> newMemrefs, newConditions;
1012 SmallVector<Operation *> toDelete;
1013 for (auto [memref, cond] :
1014 llvm::zip(deallocOp.getMemrefs(), deallocOp.getConditions())) {
1015 if (auto allocOp = memref.getDefiningOp<MemoryEffectOpInterface>()) {
1016 // Check that it is indeed an allocate effect, that the op has no other
1017 // side effects (which would not allow us to remove the op), and that
1018 // there are no other users.
1019 if (allocOp.getEffectOnValue<MemoryEffects::Allocate>(memref) &&
1021 memref.hasOneUse()) {
1022 toDelete.push_back(allocOp);
1023 continue;
1024 }
1025 }
1026
1027 newMemrefs.push_back(memref);
1028 newConditions.push_back(cond);
1029 }
1030
1031 if (failed(updateDeallocIfChanged(deallocOp, newMemrefs, newConditions,
1032 rewriter)))
1033 return failure();
1034
1035 for (Operation *op : toDelete)
1036 rewriter.eraseOp(op);
1037
1038 return success();
1039 }
1040};
1041
1042} // anonymous namespace
1043
1044void DeallocOp::getCanonicalizationPatterns(RewritePatternSet &results,
1045 MLIRContext *context) {
1047}
1048
1050 RewritePatternSet &patterns, MLIRContext *context) {
1051 patterns.add<DeallocRemoveDuplicateDeallocMemrefs,
1052 DeallocRemoveDuplicateRetainedMemrefs, EraseEmptyDealloc,
1053 EraseAlwaysFalseDealloc, SkipExtractMetadataOfAlloc,
1054 RemoveAllocDeallocPairWhenNoOtherUsers>(context);
1055}
1056
1057//===----------------------------------------------------------------------===//
1058// TableGen'd op method definitions
1059//===----------------------------------------------------------------------===//
1060
1061#define GET_OP_CLASSES
1062#include "mlir/Dialect/Bufferization/IR/BufferizationOps.cpp.inc"
return success()
static SmallVector< Value > getDynamicSize(Value memref, func::FuncOp funcOp)
Return the dynamic shapes of the memref based on the defining op.
static LogicalResult updateDeallocIfChanged(DeallocOp deallocOp, ValueRange memrefs, ValueRange conditions, PatternRewriter &rewriter)
static void copy(Location loc, Value dst, Value src, Value size, OpBuilder &builder)
Copies the given number of bytes from src to dst pointers.
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
b getContext())
auto load
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be the output argument nBegin is set to its * replacement(set to `begin` if no invalidation happens). Since outgoing *copies could have been inserted at `end`
static llvm::ManagedStatic< PassManagerOptions > options
template bool mlir::hasSingleEffect< MemoryEffects::Allocate >(Operation *)
static void getDynamicSizes(RankedTensorType tp, ValueRange sizes, SmallVectorImpl< Value > &dynSizes)
Collects the dynamic dimension sizes for tp with the assumption that sizes are the dimension sizes fo...
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 parseOptionalKeyword(StringRef keyword)=0
Parse the given keyword if present.
virtual ParseResult parseRParen()=0
Parse a ) token.
virtual InFlightDiagnostic emitError(SMLoc loc, const Twine &message={})=0
Emit a diagnostic at the specified location and return failure.
virtual ParseResult parseEqual()=0
Parse a = token.
virtual ParseResult parseCustomTypeWithFallback(Type &result, function_ref< ParseResult(Type &result)> parseType)=0
Parse a custom type with the provided callback, unless the next token is #, in which case the generic...
virtual SMLoc getCurrentLocation()=0
Get the location of the next token and store it into the argument.
virtual ParseResult parseColon()=0
Parse a : token.
virtual SMLoc getNameLoc() const =0
Return the location of the original name token.
virtual ParseResult parseLParen()=0
Parse a ( token.
void printStrippedAttrOrType(AttrOrType attrOrType)
Print the provided attribute in the context of an operation custom printer/parser: this will invoke d...
Attributes are known-constant values of operations.
Definition Attributes.h:25
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
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
Definition Builders.cpp:171
BoolAttr getBoolAttr(bool value)
Definition Builders.cpp:108
MLIRContext * getContext() const
Definition Builders.h:56
IndexType getIndexType()
Definition Builders.cpp:59
DictionaryAttr getDictionaryAttr(ArrayRef< NamedAttribute > value)
Definition Builders.cpp:112
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 class provides a mutable adaptor for a range of operands.
Definition ValueRange.h:119
NamedAttrList is array of NamedAttributes that tracks whether it is sorted and does some basic work t...
The OpAsmParser has methods for interacting with the asm parser: parsing things from it,...
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...
virtual void printOptionalAttrDict(ArrayRef< NamedAttribute > attrs, ArrayRef< StringRef > elidedAttrs={})=0
If the specified operation has attributes, print out an attribute dictionary with their values.
This class helps build Operations.
Definition Builders.h:210
This class represents a single result from folding an operation.
This class represents an operand of an operation.
Definition Value.h:254
Block * getBlock()
Returns the operation block that contains this operation.
Definition Operation.h:230
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 provides an abstraction over the different types of ranges over Regions.
Definition Region.h:363
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
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.
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...
This class represents a specific instance of an effect.
Tensor types represent multi-dimensional arrays, and have two variants: RankedTensorType and Unranked...
ArrayRef< int64_t > getShape() const
Returns the shape of this tensor type.
bool hasRank() const
Returns if this type is ranked, i.e. it has a known number of dimensions.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
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
void populateDeallocOpCanonicalizationPatterns(RewritePatternSet &patterns, MLIRContext *context)
Add the canonicalization patterns for bufferization.dealloc to the given pattern set to make them ava...
FailureOr< Value > castOrReallocMemRefValue(OpBuilder &b, Value value, MemRefType type, const BufferizationOptions &options)
Try to cast the given ranked MemRef-typed value to the given ranked MemRef type.
LogicalResult foldToBufferToTensorPair(RewriterBase &rewriter, ToBufferOp toBuffer, const BufferizationOptions &options)
Try to fold to_buffer(to_tensor(x)).
void populateDynamicDimSizes(OpBuilder &b, Location loc, Value shapedValue, SmallVector< Value > &dynamicDims)
Populate dynamicDims with tensor::DimOp / memref::DimOp results for all dynamic dimensions of the giv...
Type getTensorTypeFromMemRefType(Type type)
Return an unranked/ranked tensor type for the given unranked/ranked memref type.
Definition MemRefOps.cpp:63
std::optional< Operation * > findDealloc(Value allocValue)
Finds a single dealloc operation for the given allocated value.
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
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given tensor value.
Definition TensorOps.cpp:91
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
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
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
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
SmallVector< SmallVector< OpFoldResult > > ReifiedRankedShapedTypeDims
detail::constant_int_predicate_matcher m_Zero()
Matches a constant scalar / vector splat / tensor splat integer zero.
Definition Matchers.h:442
LogicalResult verifyRanksMatch(Operation *op, ShapedType lhs, ShapedType rhs, StringRef lhsName, StringRef rhsName)
Verify that two shaped types have matching ranks.
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
This is the representation of an operand reference.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
This represents an operation in an abstracted form, suitable for use with the builder APIs.