9#ifndef MLIR_DIALECT_BUFFERIZATION_IR_BUFFERIZABLEOPINTERFACE_H_
10#define MLIR_DIALECT_BUFFERIZATION_IR_BUFFERIZABLEOPINTERFACE_H_
15#include "llvm/ADT/DenseMapInfoVariant.h"
16#include "llvm/ADT/SetVector.h"
19#include "mlir/Dialect/Bufferization/IR/BufferizationEnums.h.inc"
31class BufferizableOpInterface;
34enum class BufferRelation {
43struct AliasingOpOperand {
44 AliasingOpOperand(OpOperand *opOperand, BufferRelation relation,
45 bool isDefinite =
true)
46 : opOperand(opOperand), relation(relation), isDefinite(isDefinite) {}
49 BufferRelation relation;
56 AliasingValue(Value value, BufferRelation relation,
bool isDefinite =
true)
57 : value(value), relation(relation), isDefinite(isDefinite) {}
60 BufferRelation relation;
68 AliasList() =
default;
71 AliasList(std::initializer_list<T> elems) {
77 AliasList(SmallVector<T> &&aliases) : aliases(std::move(aliases)) {}
79 ArrayRef<T> getAliases()
const {
return aliases; }
81 size_t getNumAliases()
const {
return aliases.size(); }
83 void addAlias(T alias) { aliases.push_back(alias); }
85 auto begin()
const {
return aliases.begin(); }
86 auto end()
const {
return aliases.end(); }
90 SmallVector<T> aliases;
95using AliasingOpOperandList = AliasList<AliasingOpOperand>;
99using AliasingValueList = AliasList<AliasingValue>;
107 using FilterFn = std::function<bool(Operation *)>;
110 enum FilterType : int8_t { DENY = 0, ALLOW = 1 };
122 bool isOpAllowed(Operation *op)
const;
127 template <
typename... DialectTs>
128 void allowDialect() {
131 (allowDialectImpl<DialectTs>(), ...);
137 template <
typename... DialectTs>
139 (denyDialectImpl<DialectTs>(), ...);
145 void allowDialect(StringRef dialectNamespace) {
146 Entry::FilterFn filterFn = [=](Operation *op) {
147 return op->getName().getDialectNamespace() == dialectNamespace;
149 entries.push_back(Entry{filterFn, Entry::FilterType::ALLOW});
155 void denyDialect(StringRef dialectNamespace) {
156 Entry::FilterFn filterFn = [=](Operation *op) {
157 return op->getName().getDialectNamespace() == dialectNamespace;
159 entries.push_back(Entry{filterFn, Entry::FilterType::DENY});
165 template <
typename... OpTys>
166 void allowOperation() {
167 (allowOperationImpl<OpTys>(), ...);
173 template <
typename... OpTys>
174 void denyOperation() {
175 (denyOperationImpl<OpTys>(), ...);
181 void allowOperation(StringRef opName) {
182 Entry::FilterFn filterFn = [=](Operation *op) {
183 return op->getName().getStringRef() == opName;
185 allowOperation(filterFn);
191 void denyOperation(StringRef opName) {
192 Entry::FilterFn filterFn = [=](Operation *op) {
193 return op->getName().getStringRef() == opName;
195 denyOperation(filterFn);
201 void allowOperation(Entry::FilterFn fn) {
202 entries.push_back(Entry{fn, Entry::FilterType::ALLOW});
208 void denyOperation(Entry::FilterFn fn) {
209 entries.push_back(Entry{fn, Entry::FilterType::DENY});
214 bool hasAllowRule()
const {
215 for (
const Entry &e : entries)
216 if (e.type == Entry::FilterType::ALLOW)
222 template <
typename DialectT>
223 void allowDialectImpl() {
224 allowDialect(DialectT::getDialectNamespace());
228 template <
typename DialectT>
229 void denyDialectImpl() {
230 denyDialect(DialectT::getDialectNamespace());
234 template <
typename OpTy>
235 void allowOperationImpl() {
236 allowOperation(OpTy::getOperationName());
240 template <
typename OpTy>
241 void denyOperationImpl() {
242 denyOperation(OpTy::getOperationName());
249 SmallVector<Entry> entries;
253struct BufferizationOptions {
257 OpBuilder &, Location, MemRefType,
ValueRange,
unsigned int)>;
260 std::function<LogicalResult(OpBuilder &, Location, Value, Value)>;
266 std::function<FailureOr<Value>(OpBuilder &, Location, Type, Value)>;
268 using AnalysisStateInitFn = std::function<void(AnalysisState &)>;
271 using FunctionArgTypeConverterFn =
272 std::function<BufferLikeType(TensorLikeType, Attribute memorySpace,
273 func::FuncOp,
const BufferizationOptions &)>;
276 using UnknownTypeConverterFn = std::function<BufferLikeType(
277 TensorLikeType, Attribute memorySpace,
const BufferizationOptions &)>;
279 using DefaultMemorySpaceFn =
280 std::function<std::optional<Attribute>(TensorLikeType t)>;
286 using ReconcileBufferTypeMismatchFn = std::function<FailureOr<BufferLikeType>(
287 BufferLikeType, BufferLikeType,
const BufferizationOptions &)>;
289 BufferizationOptions();
293 BufferizableOpInterface dynCastBufferizableOp(Operation *op)
const;
297 BufferizableOpInterface dynCastBufferizableOp(Value value)
const;
304 bool isOpAllowed(Operation *op)
const;
309 bool allowUnknownOps =
false;
313 bool bufferizeFunctionBoundaries =
false;
319 std::optional<bool> mayHaveParallelRegions = std::nullopt;
338 void setFunctionBoundaryTypeConversion(LayoutMapOption layoutMapOption);
347 CastFn castFn =
nullptr;
357 FunctionArgTypeConverterFn functionArgTypeConverterFn =
nullptr;
364 bool inferFunctionResultLayout =
true;
369 UnknownTypeConverterFn unknownTypeConverterFn =
nullptr;
375 DefaultMemorySpaceFn defaultMemorySpaceFn =
376 [](TensorLikeType t) -> std::optional<Attribute> {
return Attribute(); };
386 ReconcileBufferTypeMismatchFn reconcileBufferTypeMismatchFn =
nullptr;
390 bool copyBeforeWrite =
false;
394 bool testAnalysisOnly =
false;
398 bool printConflicts =
false;
401 unsigned int bufferAlignment = 64;
405 SmallVector<AnalysisStateInitFn> stateInitializers;
409struct TraversalConfig {
412 bool alwaysIncludeLeaves =
true;
415 bool followInPlaceOnly =
false;
418 bool followEquivalentOnly =
false;
422 bool followUnknownOps =
false;
426 bool followSameTypeOrCastsOnly =
false;
430 bool revisitAlreadyVisitedValues =
false;
440 AliasingOpOperandList getAliasingOpOperands(Value value)
const;
444 AliasingValueList getAliasingValues(OpOperand &opOperand)
const;
448 bool bufferizesToMemoryRead(OpOperand &opOperand)
const;
452 bool bufferizesToMemoryWrite(OpOperand &opOperand)
const;
457 bool bufferizesToMemoryWrite(Value value)
const;
461 bool bufferizesToAliasOnly(OpOperand &opOperand)
const;
465 bool canOmitTensorCopy(OpOperand &opOperand)
const;
470 bool isValueRead(Value value)
const;
499 SetVector<Value> findValueInReverseUseDefChain(
500 OpOperand *opOperand, llvm::function_ref<
bool(Value)> condition,
501 TraversalConfig config = TraversalConfig(),
502 llvm::DenseSet<OpOperand *> *visitedOpOperands =
nullptr)
const;
536 SetVector<Value> findDefinitions(OpOperand *opOperand)
const;
539 virtual bool isInPlace(OpOperand &opOperand)
const;
542 virtual bool areEquivalentBufferizedValues(Value v1, Value v2)
const;
545 virtual bool areAliasingBufferizedValues(Value v1, Value v2)
const;
548 virtual bool hasUndefinedContents(OpOperand *opOperand)
const;
551 const BufferizationOptions &getOptions()
const {
return options; }
553 AnalysisState(
const BufferizationOptions &
options);
556 AnalysisState(
const AnalysisState &) =
delete;
558 virtual ~AnalysisState() =
default;
560 static bool classof(
const AnalysisState *base) {
return true; }
562 TypeID
getType()
const {
return type; }
566 const BufferizationOptions &
options);
571 const BufferizationOptions &
options);
575 const BufferizationOptions &
options);
577 virtual void resetCache();
585 AnalysisState(
const BufferizationOptions &
options, TypeID type);
589 const BufferizationOptions &
options;
595 DenseMap<std::variant<Operation *, Block *, Region *, Value>, Region *>
596 enclosingRepetitiveRegionCache;
600 DenseMap<std::pair<Operation *, Operation *>,
bool>
601 insideMutuallyExclusiveRegionsCache;
606class BufferizationState {
609 SymbolTableCollection &getSymbolTables();
611 SymbolTableCollection &getSymbolTables()
const;
617 mutable SymbolTableCollection symbolTables;
624allocateTensorForShapedValue(OpBuilder &
b, Location loc, Value shapedValue,
625 const BufferizationOptions &
options,
626 const BufferizationState &state,
bool copy =
true);
631FailureOr<Value> getBuffer(RewriterBase &rewriter, Value value,
632 const BufferizationOptions &
options,
633 const BufferizationState &state);
644 const BufferizationOptions &
options,
645 const BufferizationState &state);
658 const BufferizationOptions &
options,
659 const BufferizationState &state,
660 SmallVector<Value> &invocationStack);
666bool hasTensorSemantics(Operation *op);
670void replaceOpWithBufferizedValues(RewriterBase &rewriter, Operation *op,
675template <
typename OpTy,
typename... Args>
676OpTy replaceOpWithNewBufferizedOp(RewriterBase &rewriter, Operation *op,
679 OpTy::create(rewriter, op->getLoc(), std::forward<Args>(args)...);
680 replaceOpWithBufferizedValues(rewriter, op, newOp->getResults());
687getMemRefTypeWithFullyDynamicLayout(TensorType tensorType,
688 Attribute memorySpace =
nullptr);
693getMemRefTypeWithStaticIdentityLayout(TensorType tensorType,
694 Attribute memorySpace =
nullptr);
702Region *getNextEnclosingRepetitiveRegion(Region *region,
703 const BufferizationOptions &
options);
711Region *getParallelRegion(Region *region,
const BufferizationOptions &
options);
717AliasingOpOperandList defaultGetAliasingOpOperands(Value value,
718 const AnalysisState &state);
723FailureOr<BufferLikeType>
724defaultGetBufferType(Value value,
const BufferizationOptions &
options,
725 const BufferizationState &state,
726 SmallVector<Value> &invocationStack);
731bool defaultResultBufferizesToMemoryWrite(OpResult opResult,
732 const AnalysisState &state);
737bool defaultIsRepetitiveRegion(BufferizableOpInterface bufferizableOp,
742AliasingOpOperandList unknownGetAliasingOpOperands(Value value);
746AliasingValueList unknownGetAliasingValues(OpOperand &opOperand);
750bool defaultHasTensorSemantics(Operation *op);
758FailureOr<BaseMemRefType> asMemRefType(FailureOr<BufferLikeType> bufferType);
763bool typesMatchAfterBufferization(Operation &op, Value tensor, Value buffer);
775#include "mlir/Dialect/Bufferization/IR/BufferizableOpInterface.h.inc"
bufferization::BufferResultsToOutParamsOpts::AllocationFn AllocationFn
bufferization::BufferResultsToOutParamsOpts::MemCpyFn MemCpyFn
static void copy(Location loc, Value dst, Value src, Value size, OpBuilder &builder)
Copies the given number of bytes from src to dst pointers.
static llvm::ManagedStatic< PassManagerOptions > options
static RankedTensorType getBufferType(const SparseTensorType &stt, bool needTmpCOO)
#define MLIR_DECLARE_EXPLICIT_TYPE_ID(CLASS_NAME)
static Operation * getOwnerOfValue(Value value)
This class helps build Operations.
Include the generated interface declarations.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
bool insideMutuallyExclusiveRegions(Operation *a, Operation *b)
Return true if a and b are in mutually exclusive regions as per RegionBranchOpInterface.
Region * getEnclosingRepetitiveRegion(Operation *op)
Return the first enclosing region of the given op that may be executed repetitively as per RegionBran...