23#include "llvm/Support/DebugLog.h"
28#define GEN_PASS_DEF_ONESHOTBUFFERIZEPASS
29#include "mlir/Dialect/Bufferization/Transforms/Passes.h.inc"
33#define DEBUG_TYPE "bufferize"
41parseHeuristicOption(
const std::string &s) {
46 if (s ==
"bottom-up-from-terminators")
47 return OneShotBufferizationOptions::AnalysisHeuristic::
48 BottomUpFromTerminators;
51 llvm_unreachable(
"invalid analysisheuristic option");
54struct OneShotBufferizePass
55 :
public bufferization::impl::OneShotBufferizePassBase<
56 OneShotBufferizePass> {
59 void runOnOperation()
override {
65 opt.allowUnknownOps = allowUnknownOps;
68 opt.copyBeforeWrite = copyBeforeWrite;
70 opt.setFunctionBoundaryTypeConversion(functionBoundaryTypeConversion);
72 if (mustInferMemorySpace && useEncodingForMemorySpace) {
74 <<
"only one of 'must-infer-memory-space' and "
75 "'use-encoding-for-memory-space' are allowed in "
77 return signalPassFailure();
80 if (mustInferMemorySpace) {
81 opt.defaultMemorySpaceFn =
82 [](TensorLikeType t) -> std::optional<Attribute> {
87 if (useEncodingForMemorySpace) {
88 opt.defaultMemorySpaceFn =
89 [](TensorLikeType t) -> std::optional<Attribute> {
90 if (
auto rtt = dyn_cast<RankedTensorType>(t))
91 return rtt.getEncoding();
96 opt.printConflicts = printConflicts;
97 opt.bufferAlignment = bufferAlignment;
98 opt.testAnalysisOnly = testAnalysisOnly;
99 opt.bufferizeFunctionBoundaries = bufferizeFunctionBoundaries;
100 if (!mayHaveParallelRegions)
101 opt.mayHaveParallelRegions =
false;
105 LayoutMapOption unknownTypeConversionOption = unknownTypeConversion;
106 if (unknownTypeConversionOption == LayoutMapOption::InferLayoutMap) {
108 "Invalid option: 'infer-layout-map' is not a valid value for "
109 "'unknown-type-conversion'");
110 return signalPassFailure();
112 opt.unknownTypeConverterFn = [=](TensorLikeType type,
115 const auto tensorType = cast<TensorType>(type);
116 if (unknownTypeConversionOption == LayoutMapOption::IdentityLayoutMap)
117 return cast<bufferization::BufferLikeType>(
118 bufferization::getMemRefTypeWithStaticIdentityLayout(
119 tensorType, memorySpace));
120 assert(unknownTypeConversionOption ==
121 LayoutMapOption::FullyDynamicLayoutMap &&
122 "invalid layout map option");
123 return cast<bufferization::BufferLikeType>(
124 bufferization::getMemRefTypeWithFullyDynamicLayout(tensorType,
129 OpFilter::Entry::FilterFn filterFn = [&](
Operation *op) {
131 if (this->dialectFilter.hasValue() && !(*this->dialectFilter).empty())
132 return llvm::is_contained(this->dialectFilter,
133 op->getDialect()->getNamespace());
137 opt.opFilter.allowOperation(filterFn);
142 if (opt.copyBeforeWrite && opt.testAnalysisOnly) {
148 "Invalid option: 'copy-before-write' cannot be used with "
149 "'test-analysis-only'");
150 return signalPassFailure();
153 if (opt.printConflicts && !opt.testAnalysisOnly) {
156 "Invalid option: 'print-conflicts' requires 'test-analysis-only'");
157 return signalPassFailure();
163 "Invalid option: 'dump-alias-sets' requires 'test-analysis-only'");
164 return signalPassFailure();
167 BufferizationState state;
169 ModuleOp moduleOp = getOperation();
170 if (opt.bufferizeFunctionBoundaries) {
179 "Invalid option: 'no-analysis-func-filter' requires "
180 "'bufferize-function-boundaries'");
181 return signalPassFailure();
196 std::optional<OneShotBufferizationOptions>
options;
210 SmallVector<Operation *> &worklist,
211 const BufferizationOptions &
options,
212 BufferizationStatistics *statistics)
213 : IRRewriter(ctx), erasedOps(erasedOps), toBufferOps(toBufferOps),
214 worklist(worklist), analysisState(
options), statistics(statistics) {
219 void notifyOperationErased(Operation *op)
override {
220 erasedOps.insert(op);
222 toBufferOps.erase(op);
225 void notifyOperationInserted(Operation *op, InsertPoint previous)
override {
227 if (previous.isSet())
234 if (
auto sideEffectingOp = dyn_cast<MemoryEffectOpInterface>(op))
235 statistics->numBufferAlloc +=
static_cast<int64_t
>(
236 sideEffectingOp.hasEffect<MemoryEffects::Allocate>());
240 if (isa<ToBufferOp>(op)) {
241 toBufferOps.insert(op);
246 if (isa<ToTensorOp>(op))
250 if (!hasTensorSemantics(op))
254 auto const &
options = analysisState.getOptions();
259 worklist.push_back(op);
270 SmallVector<Operation *> &worklist;
274 const AnalysisState analysisState;
277 BufferizationStatistics *statistics;
283 BufferizationState &bufferizationState,
293 op->
walk([&](ToBufferOp toBufferOp) { toBufferOps.insert(toBufferOp); });
304 if (
options.isOpAllowed(op) && hasTensorSemantics(op))
305 worklist.push_back(op);
312 BufferizationRewriter rewriter(op->
getContext(), erasedOps, toBufferOps,
313 worklist,
options, statistics);
314 for (
unsigned i = 0; i < worklist.size(); ++i) {
317 if (erasedOps.contains(nextOp))
320 auto bufferizableOp =
options.dynCastBufferizableOp(nextOp);
324 if (!hasTensorSemantics(nextOp))
327 if (!bufferizableOp.supportsUnstructuredControlFlow())
329 if (r.getBlocks().size() > 1)
331 "op or BufferizableOpInterface implementation does not support "
332 "unstructured control flow, but at least one region has multiple "
336 LDBG(3) <<
"//===-------------------------------------------===//\n"
337 <<
"IR after bufferizing: " << nextOp->
getName();
338 rewriter.setInsertionPoint(nextOp);
340 bufferizableOp.bufferize(rewriter,
options, bufferizationState))) {
341 LDBG(2) <<
"failed to bufferize\n"
342 <<
"//===-------------------------------------------===//";
343 return nextOp->
emitError(
"failed to bufferize op");
345 LDBG(3) << *op <<
"\n//===-------------------------------------------===//";
349 if (erasedOps.contains(op))
357 for (
Operation *op : toBufferOpsSnapshot) {
358 if (erasedOps.contains(op))
360 rewriter.setInsertionPoint(op);
362 rewriter, cast<ToBufferOp>(op),
options);
367 if (toTensorOp->getUses().empty()) {
368 rewriter.eraseOp(toTensorOp);
381 if (erasedOps.contains(op))
385 if (!hasTensorSemantics(op))
394 if (isa<ToTensorOp, ToBufferOp>(op))
396 return op->
emitError(
"op was not bufferized");
405 BufferizationState &state) {
414 auto tensorType = dyn_cast<TensorLikeType>(bbArg.getType());
416 newTypes.push_back(bbArg.getType());
420 FailureOr<BufferLikeType> bufferType =
421 bufferization::getBufferType(bbArg,
options, state);
422 if (failed(bufferType))
424 newTypes.push_back(*bufferType);
428 for (
auto [bbArg, type] : llvm::zip(block->
getArguments(), newTypes)) {
429 if (bbArg.getType() == type)
435 bbArgUses.push_back(&use);
437 Type tensorType = bbArg.getType();
443 if (!bbArgUses.empty()) {
444 Value toTensorOp = bufferization::ToTensorOp::create(
445 rewriter, bbArg.getLoc(), tensorType, bbArg);
447 use->set(toTensorOp);
456 auto branchOp = dyn_cast<BranchOpInterface>(op);
458 return op->
emitOpError(
"cannot bufferize ops with block references that "
459 "do not implement BranchOpInterface");
462 branchOp.getSuccessorOperands(blockOperand.getOperandNumber());
464 for (
auto [operand, type] :
466 if (operand.getType() == type) {
468 newOperands.push_back(operand);
471 FailureOr<BufferLikeType> operandBufferType =
472 bufferization::getBufferType(operand,
options, state);
473 if (failed(operandBufferType))
476 Value bufferizedOperand = bufferization::ToBufferOp::create(
477 rewriter, operand.getLoc(), *operandBufferType, operand);
480 if (type != *operandBufferType) {
481 bufferizedOperand = *
options.castFn(rewriter, operand.getLoc(), type,
484 newOperands.push_back(bufferizedOperand);
static llvm::ManagedStatic< PassManagerOptions > options
Base class for generic analysis states.
Attributes are known-constant values of operations.
This class represents an argument of a Block.
A block operand represents an operand that holds a reference to a Block, e.g.
Block represents an ordered list of Operations.
BlockArgListType getArguments()
Operation * getParentOp()
Returns the closest surrounding operation that contains this block.
use_range getUses() const
Returns a range of all uses, which is useful for iterating over all uses.
This class coordinates rewriting a piece of IR outside of a pattern rewrite, providing a way to keep ...
void assign(ValueRange values)
Assign this range to the given values.
RAII guard to reset the insertion point of the builder when destroyed.
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
void setInsertionPointAfterValue(Value val)
Sets the insertion point to the node after the specified value.
This class represents an operand of an operation.
Operation is the basic unit of execution within MLIR.
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
OperationName getName()
The name of an operation is the key identifier for it.
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
std::enable_if_t< llvm::function_traits< std::decay_t< FnT > >::num_args==1, RetT > walk(FnT &&callback)
Walk the operation by calling the callback for each nested operation (including this one),...
use_range getUses()
Returns a range of all uses, which is useful for iterating over all uses.
MLIRContext * getContext()
Return the context this operation is associated with.
InFlightDiagnostic emitOpError(const Twine &message={})
Emit an error with the op name prefixed, like "'dim' op " which is convenient for verifiers.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
This class models how operands are forwarded to block arguments in control flow.
MutableOperandRange getMutableForwardedOperands() const
Get the range of operands that are simply forwarded to the successor.
OperandRange getForwardedOperands() const
Get the range of operands that are simply forwarded to the successor.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
static WalkResult advance()
LogicalResult bufferizeOp(Operation *op, const BufferizationOptions &options, BufferizationState &bufferizationState, BufferizationStatistics *statistics=nullptr)
Bufferize op and its nested ops that implement BufferizableOpInterface.
LogicalResult bufferizeBlockSignature(Block *block, RewriterBase &rewriter, const BufferizationOptions &options, BufferizationState &state)
Bufferize the signature of block and its callers (i.e., ops that have the given block as a successor)...
LogicalResult insertTensorCopies(Operation *op, const OneShotBufferizationOptions &options, const BufferizationState &bufferizationState, BufferizationStatistics *statistics=nullptr)
Resolve RaW and other conflicts by inserting bufferization.alloc_tensor ops.
LogicalResult runOneShotBufferize(Operation *op, const OneShotBufferizationOptions &options, BufferizationState &state, BufferizationStatistics *statistics=nullptr)
Run One-Shot Bufferize on the given op: Analysis + Bufferization.
LogicalResult foldToBufferToTensorPair(RewriterBase &rewriter, ToBufferOp toBuffer, const BufferizationOptions &options)
Try to fold to_buffer(to_tensor(x)).
llvm::LogicalResult runOneShotModuleBufferize(Operation *moduleOp, const bufferization::OneShotBufferizationOptions &options, BufferizationState &state, BufferizationStatistics *statistics=nullptr)
Run One-Shot Module Bufferization on the given SymbolTable.
Include the generated interface declarations.
llvm::DenseSet< ValueT, ValueInfoT > DenseSet
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
Bufferization statistics for debugging.
int64_t numTensorOutOfPlace
Options for analysis-enabled bufferization.
unsigned analysisFuzzerSeed
Seed for the analysis fuzzer.
bool dumpAliasSets
Specifies whether the tensor IR should be annotated with alias sets.
bool allowReturnAllocsFromLoops
Specifies whether returning newly allocated memrefs from loops should be allowed.
AnalysisHeuristic analysisHeuristic
The heuristic controls the order in which ops are traversed during the analysis.
llvm::ArrayRef< std::string > noAnalysisFuncFilter
Specify the functions that should not be analyzed.