17#include "llvm/ADT/SmallVectorExtras.h"
30 auto srcType = llvm::cast<MemRefType>(value.
getType());
33 if (srcType.getElementType() != destType.getElementType())
35 if (srcType.getRank() != destType.getRank())
41 auto isGuaranteedCastCompatible = [](MemRefType source, MemRefType
target) {
42 int64_t sourceOffset, targetOffset;
44 if (failed(source.getStridesAndOffset(sourceStrides, sourceOffset)) ||
45 failed(
target.getStridesAndOffset(targetStrides, targetOffset)))
48 return ShapedType::isDynamic(a) && ShapedType::isStatic(
b);
50 if (dynamicToStatic(sourceOffset, targetOffset))
52 for (
auto it : zip(sourceStrides, targetStrides))
53 if (dynamicToStatic(std::get<0>(it), std::get<1>(it)))
60 if (!source.getLayout().isIdentity() &&
target.getLayout().isIdentity())
68 if (memref::CastOp::areCastCompatible(srcType, destType) &&
69 isGuaranteedCastCompatible(srcType, destType)) {
76 for (
int i = 0; i < destType.getRank(); ++i) {
77 if (destType.getShape()[i] != ShapedType::kDynamic)
79 Value size = memref::DimOp::create(
b, loc, value, i);
80 dynamicOperands.push_back(size);
84 b, loc, destType, dynamicOperands,
options.bufferAlignment);
97 auto bufferToTensor = toBuffer.getTensor().getDefiningOp<ToTensorOp>();
101 Type srcType = bufferToTensor.getBuffer().getType();
102 Type destType = toBuffer.getType();
105 if (srcType == destType) {
106 rewriter.
replaceOp(toBuffer, bufferToTensor.getBuffer());
110 if (!llvm::isa<BaseMemRefType>(srcType) ||
111 !llvm::isa<BaseMemRefType>(destType)) {
114 options.castFn(rewriter, bufferToTensor.getBuffer().getLoc(), destType,
115 bufferToTensor.getBuffer());
122 auto rankedSrcType = llvm::dyn_cast<MemRefType>(srcType);
123 auto rankedDestType = llvm::dyn_cast<MemRefType>(destType);
124 auto unrankedSrcType = llvm::dyn_cast<UnrankedMemRefType>(srcType);
127 if (rankedSrcType && rankedDestType) {
129 rewriter, bufferToTensor.getBuffer(), rankedDestType,
options);
139 if (unrankedSrcType && rankedDestType)
144 if (!memref::CastOp::areCastCompatible(srcType, destType))
148 bufferToTensor.getBuffer());
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));
161 assert(llvm::isa<RankedTensorType>(shapedType) &&
"expected tensor");
162 dynamicDims.push_back(tensor::DimOp::create(
b, loc, shapedValue, i));
172LogicalResult AllocTensorOp::verify() {
174 return emitError(
"dynamic sizes not needed when copying a tensor");
179 return emitError(
"expected that `copy` and return type match");
184 RankedTensorType type,
ValueRange dynamicSizes) {
185 build(builder,
result, type, dynamicSizes,
Value(),
191 RankedTensorType type,
ValueRange dynamicSizes,
199 IntegerAttr memorySpace) {
217 using OpRewritePattern<AllocTensorOp>::OpRewritePattern;
219 LogicalResult matchAndRewrite(AllocTensorOp op,
220 PatternRewriter &rewriter)
const override {
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))
229 Value value = op.getDynamicSizes()[dynValCounter++];
232 int64_t dim = intVal.getSExtValue();
234 newShape[i] = intVal.getSExtValue();
236 newDynamicSizes.push_back(value);
238 newDynamicSizes.push_back(value);
241 RankedTensorType newType = RankedTensorType::get(
242 newShape, op.getType().getElementType(), op.getType().getEncoding());
243 if (newType == op.getType())
245 auto newOp = AllocTensorOp::create(rewriter, op.getLoc(), newType,
246 newDynamicSizes, Value());
253 using OpRewritePattern<tensor::DimOp>::OpRewritePattern;
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)
261 if (*maybeConstantIndex < 0 ||
262 *maybeConstantIndex >= allocTensorOp.getType().getRank())
264 if (!allocTensorOp.getType().isDynamicDim(*maybeConstantIndex))
267 dimOp, allocTensorOp.getDynamicSize(rewriter, *maybeConstantIndex));
275 results.
add<FoldDimOfAllocTensorOp, ReplaceStaticShapeDims>(ctx);
278LogicalResult AllocTensorOp::reifyResultShapes(
281 llvm::map_to_vector<4>(llvm::seq<int64_t>(0,
getType().getRank()),
283 if (isDynamicDim(dim))
287 reifiedReturnShapes.emplace_back(std::move(shapes));
298 if (copyKeyword.succeeded())
304 if (sizeHintKeyword.succeeded())
309 if (AllocTensorOp::genericParseProperties(parser, parsedProperties))
311 auto propertyDictionary = dyn_cast_or_null<DictionaryAttr>(parsedProperties);
312 if (parsedProperties && !propertyDictionary)
314 "expected properties dictionary");
319 for (StringRef attrName : AllocTensorOp::getAttributeNames()) {
320 if (
result.attributes.get(attrName))
322 <<
"inherent attribute '" << attrName
323 <<
"' cannot be parsed from attr-dict when strict properties in "
324 "assembly format is enabled";
337 if (copyKeyword.succeeded())
340 if (sizeHintKeyword.succeeded())
344 NamedAttrList properties(propertyDictionary ? propertyDictionary
346 properties.set(AllocTensorOp::getOperandSegmentSizeAttr(),
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());
354 << propertyDictionary <<
" for op " <<
result.name.getStringRef()
357 if (
failed(AllocTensorOp::setPropertiesFromParsedAttr(
358 result.getOrAddProperties<Properties>(), propertyDictionary,
367 p <<
" copy(" << getCopy() <<
")";
369 p <<
" size_hint=" << getSizeHint();
370 AllocTensorOp::printProperties(
getContext(), p, getProperties(),
371 getOperandSegmentSizeAttr());
374 auto type = getResult().getType();
375 if (
auto validType = llvm::dyn_cast<::mlir::TensorType>(type))
382 assert(isDynamicDim(idx) &&
"expected dynamic dim");
384 return tensor::DimOp::create(
b, getLoc(), getCopy(), idx);
385 return getOperand(getIndexOfDynamicSize(idx));
401 using OpRewritePattern<CloneOp>::OpRewritePattern;
403 LogicalResult matchAndRewrite(CloneOp cloneOp,
404 PatternRewriter &rewriter)
const override {
405 if (cloneOp.use_empty()) {
410 Value source = cloneOp.getInput();
411 if (source.
getType() != cloneOp.getType() &&
412 !memref::CastOp::areCastCompatible({source.getType()},
413 {cloneOp.getType()}))
418 Value canonicalSource = source;
419 while (
auto iface = dyn_cast_or_null<ViewLikeOpInterface>(
421 if (canonicalSource != iface.getViewDest()) {
424 canonicalSource = iface.getViewSource();
427 std::optional<Operation *> maybeCloneDeallocOp =
430 if (!maybeCloneDeallocOp.has_value())
432 std::optional<Operation *> maybeSourceDeallocOp =
434 if (!maybeSourceDeallocOp.has_value())
436 Operation *cloneDeallocOp = *maybeCloneDeallocOp;
437 Operation *sourceDeallocOp = *maybeSourceDeallocOp;
441 if (cloneDeallocOp && sourceDeallocOp &&
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;
453 if (!redundantDealloc)
461 for (Operation *pos = cloneOp->getNextNode(); pos != redundantDealloc;
462 pos = pos->getNextNode()) {
466 auto effectInterface = dyn_cast<MemoryEffectOpInterface>(pos);
467 if (!effectInterface)
469 if (effectInterface.hasEffect<MemoryEffects::Free>())
473 if (source.
getType() != cloneOp.getType())
474 source = memref::CastOp::create(rewriter, cloneOp.getLoc(),
475 cloneOp.getType(), source);
477 rewriter.
eraseOp(redundantDealloc);
486 results.
add<SimplifyClones>(context);
493LogicalResult MaterializeInDestinationOp::reifyResultShapes(
495 if (getOperation()->getNumResults() == 1) {
496 assert(isa<TensorType>(getDest().
getType()) &&
"expected tensor type");
497 reifiedReturnShapes.resize(1,
499 reifiedReturnShapes[0] =
505Value MaterializeInDestinationOp::buildSubsetExtraction(
OpBuilder &builder,
507 if (isa<TensorType>(getDest().
getType())) {
520 assert(isa<BaseMemRefType>(getDest().
getType()) &&
"expected memref type");
521 assert(getRestrict() &&
522 "expected that ops with memrefs dest have 'restrict'");
524 return ToTensorOp::create(
527 true, getWritable());
530bool MaterializeInDestinationOp::isEquivalentSubset(
532 return equivalenceFn(getDest(), candidate);
536MaterializeInDestinationOp::getValuesNeededToBuildSubsetExtraction() {
540OpOperand &MaterializeInDestinationOp::getSourceOperand() {
541 return getOperation()->getOpOperand(0) ;
544bool MaterializeInDestinationOp::operatesOnEquivalentSubset(
545 SubsetOpInterface subsetOp,
550bool MaterializeInDestinationOp::operatesOnDisjointSubset(
551 SubsetOpInterface subsetOp,
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");
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");
574 ShapedType destType = cast<ShapedType>(getDest().
getType());
575 if (srcType.
hasRank() != destType.hasRank())
576 return emitOpError(
"source/destination shapes are incompatible");
581 for (
auto [src, dest] :
582 llvm::zip(srcType.
getShape(), destType.getShape())) {
583 if (src == ShapedType::kDynamic || dest == ShapedType::kDynamic) {
589 return emitOpError(
"source/destination shapes are incompatible");
595void MaterializeInDestinationOp::build(
OpBuilder &builder,
598 auto destTensorType = dyn_cast<TensorType>(dest.
getType());
599 build(builder, state, destTensorType ? destTensorType :
Type(),
604 return getDestMutable();
607void MaterializeInDestinationOp::getEffects(
610 if (isa<BaseMemRefType>(getDest().
getType()))
620 if (
auto toBuffer = getBuffer().getDefiningOp<ToBufferOp>())
623 if (toBuffer->getBlock() == this->getOperation()->getBlock() &&
624 toBuffer->getNextNode() == this->getOperation())
625 return toBuffer.getTensor();
631 using OpRewritePattern<tensor::DimOp>::OpRewritePattern;
633 LogicalResult matchAndRewrite(tensor::DimOp dimOp,
634 PatternRewriter &rewriter)
const override {
635 auto memrefToTensorOp = dimOp.getSource().getDefiningOp<ToTensorOp>();
636 if (!memrefToTensorOp)
640 dimOp, memrefToTensorOp.getBuffer(), dimOp.getIndex());
648 results.
add<DimOfToTensorFolder>(context);
656 if (
auto memrefToTensor = getTensor().getDefiningOp<ToTensorOp>())
657 if (memrefToTensor.getBuffer().getType() ==
getType())
658 return memrefToTensor.getBuffer();
666 using OpRewritePattern<ToBufferOp>::OpRewritePattern;
668 LogicalResult matchAndRewrite(ToBufferOp toBuffer,
669 PatternRewriter &rewriter)
const final {
670 auto tensorCastOperand =
671 toBuffer.getOperand().getDefiningOp<tensor::CastOp>();
672 if (!tensorCastOperand)
674 auto srcTensorType = llvm::dyn_cast<RankedTensorType>(
675 tensorCastOperand.getOperand().getType());
678 auto currentOutputMemRefType =
679 dyn_cast<BaseMemRefType>(toBuffer.getResult().getType());
680 if (!currentOutputMemRefType)
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());
697 using OpRewritePattern<ToBufferOp>::OpRewritePattern;
699 LogicalResult matchAndRewrite(ToBufferOp toBuffer,
700 PatternRewriter &rewriter)
const final {
710 using OpRewritePattern<memref::LoadOp>::OpRewritePattern;
712 LogicalResult matchAndRewrite(memref::LoadOp
load,
713 PatternRewriter &rewriter)
const override {
714 auto toBuffer =
load.getMemref().getDefiningOp<ToBufferOp>();
715 if (!toBuffer || !toBuffer.getReadOnly())
726 using OpRewritePattern<memref::DimOp>::OpRewritePattern;
728 LogicalResult matchAndRewrite(memref::DimOp dimOp,
729 PatternRewriter &rewriter)
const override {
730 auto castOp = dimOp.getSource().getDefiningOp<ToBufferOp>();
733 Value newSource = castOp.getOperand();
744 results.
add<DimOfCastOp, LoadOfToBuffer, ToBufferOfCast,
745 ToBufferToTensorFolding>(context);
748std::optional<Operation *> CloneOp::buildDealloc(
OpBuilder &builder,
750 return memref::DeallocOp::create(builder, alloc.
getLoc(), alloc)
754std::optional<Value> CloneOp::buildClone(
OpBuilder &builder,
Value alloc) {
755 return CloneOp::create(builder, alloc.
getLoc(), alloc).getResult();
762LogicalResult DeallocOp::inferReturnTypes(
763 MLIRContext *context, std::optional<::mlir::Location> location,
766 DeallocOpAdaptor adaptor(operands, attributes, properties, regions);
768 IntegerType::get(context, 1));
772LogicalResult DeallocOp::verify() {
773 if (getMemrefs().size() != getConditions().size())
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");
786 if (deallocOp.getMemrefs() == memrefs &&
787 deallocOp.getConditions() == conditions)
791 deallocOp.getMemrefsMutable().assign(memrefs);
792 deallocOp.getConditionsMutable().assign(conditions);
812struct DeallocRemoveDuplicateDeallocMemrefs
814 using OpRewritePattern<DeallocOp>::OpRewritePattern;
816 LogicalResult matchAndRewrite(DeallocOp deallocOp,
817 PatternRewriter &rewriter)
const override {
820 SmallVector<Value> newMemrefs, newConditions;
821 for (
auto [i, memref, cond] :
822 llvm::enumerate(deallocOp.getMemrefs(), deallocOp.getConditions())) {
823 if (memrefToCondition.count(memref)) {
826 Value &newCond = newConditions[memrefToCondition[memref]];
829 arith::OrIOp::create(rewriter, deallocOp.getLoc(), newCond, cond);
831 memrefToCondition.insert({memref, newConditions.size()});
832 newMemrefs.push_back(memref);
833 newConditions.push_back(cond);
854struct DeallocRemoveDuplicateRetainedMemrefs
856 using OpRewritePattern<DeallocOp>::OpRewritePattern;
858 LogicalResult matchAndRewrite(DeallocOp deallocOp,
859 PatternRewriter &rewriter)
const override {
862 SmallVector<Value> newRetained;
863 SmallVector<unsigned> resultReplacementIdx;
865 for (
auto retained : deallocOp.getRetained()) {
866 if (seen.count(retained)) {
867 resultReplacementIdx.push_back(seen[retained]);
872 newRetained.push_back(retained);
873 resultReplacementIdx.push_back(i++);
878 if (newRetained.size() == deallocOp.getRetained().size())
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];
890 rewriter.
replaceOp(deallocOp, replacements);
901 using OpRewritePattern<DeallocOp>::OpRewritePattern;
903 LogicalResult matchAndRewrite(DeallocOp deallocOp,
904 PatternRewriter &rewriter)
const override {
905 if (deallocOp.getMemrefs().empty()) {
906 Value constFalse = arith::ConstantOp::create(rewriter, deallocOp.getLoc(),
909 deallocOp, SmallVector<Value>(deallocOp.getUpdatedConditions().size(),
930 using OpRewritePattern<DeallocOp>::OpRewritePattern;
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())) {
938 newMemrefs.push_back(memref);
939 newConditions.push_back(cond);
967 using OpRewritePattern<DeallocOp>::OpRewritePattern;
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)
977 Value allocMemref = extractStridedOp.getOperand();
978 auto allocOp = allocMemref.getDefiningOp<MemoryEffectOpInterface>();
981 if (allocOp.getEffectOnValue<MemoryEffects::Allocate>(allocMemref))
987 deallocOp.getConditions(), rewriter);
1005struct RemoveAllocDeallocPairWhenNoOtherUsers
1007 using OpRewritePattern<DeallocOp>::OpRewritePattern;
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>()) {
1019 if (allocOp.getEffectOnValue<MemoryEffects::Allocate>(memref) &&
1021 memref.hasOneUse()) {
1022 toDelete.push_back(allocOp);
1027 newMemrefs.push_back(memref);
1028 newConditions.push_back(cond);
1035 for (Operation *op : toDelete)
1051 patterns.
add<DeallocRemoveDuplicateDeallocMemrefs,
1052 DeallocRemoveDuplicateRetainedMemrefs, EraseEmptyDealloc,
1053 EraseAlwaysFalseDealloc, SkipExtractMetadataOfAlloc,
1054 RemoveAllocDeallocPairWhenNoOtherUsers>(context);
1061#define GET_OP_CLASSES
1062#include "mlir/Dialect/Bufferization/IR/BufferizationOps.cpp.inc"
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.
*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.
This class is a general helper class for creating context-global objects like types,...
IntegerAttr getIndexAttr(int64_t value)
DenseI32ArrayAttr getDenseI32ArrayAttr(ArrayRef< int32_t > values)
BoolAttr getBoolAttr(bool value)
MLIRContext * getContext() const
DictionaryAttr getDictionaryAttr(ArrayRef< NamedAttribute > value)
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
This class provides a mutable adaptor for a range of operands.
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.
This class represents a single result from folding an operation.
This class represents an operand of an operation.
Block * getBlock()
Returns the operation block that contains this operation.
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.
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.
static DerivedEffect * get()
static DefaultResource * get()
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...
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
Location getLoc() const
Return the location of this value.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
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.
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".
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given tensor value.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
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...
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.
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.
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
llvm::function_ref< Fn > function_ref
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.