MLIR 24.0.0git
Bufferize.cpp
Go to the documentation of this file.
1//===- Bufferize.cpp - Bufferization utilities ----------------------------===//
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
10
18#include "mlir/IR/Diagnostics.h"
19#include "mlir/IR/Operation.h"
23#include "llvm/Support/DebugLog.h"
24#include <optional>
25
26namespace mlir {
27namespace bufferization {
28#define GEN_PASS_DEF_ONESHOTBUFFERIZEPASS
29#include "mlir/Dialect/Bufferization/Transforms/Passes.h.inc"
30} // namespace bufferization
31} // namespace mlir
32
33#define DEBUG_TYPE "bufferize"
34
35using namespace mlir;
36using namespace mlir::bufferization;
37
38namespace {
39
41parseHeuristicOption(const std::string &s) {
42 if (s == "bottom-up")
44 if (s == "top-down")
46 if (s == "bottom-up-from-terminators")
47 return OneShotBufferizationOptions::AnalysisHeuristic::
48 BottomUpFromTerminators;
49 if (s == "fuzzer")
51 llvm_unreachable("invalid analysisheuristic option");
52}
53
54struct OneShotBufferizePass
55 : public bufferization::impl::OneShotBufferizePassBase<
56 OneShotBufferizePass> {
57 using Base::Base;
58
59 void runOnOperation() override {
61 if (!options) {
62 // Make new bufferization options if none were provided when creating the
63 // pass.
64 opt.allowReturnAllocsFromLoops = allowReturnAllocsFromLoops;
65 opt.allowUnknownOps = allowUnknownOps;
66 opt.analysisFuzzerSeed = analysisFuzzerSeed;
67 opt.analysisHeuristic = parseHeuristicOption(analysisHeuristic);
68 opt.copyBeforeWrite = copyBeforeWrite;
69 opt.dumpAliasSets = dumpAliasSets;
70 opt.setFunctionBoundaryTypeConversion(functionBoundaryTypeConversion);
71
72 if (mustInferMemorySpace && useEncodingForMemorySpace) {
73 emitError(getOperation()->getLoc())
74 << "only one of 'must-infer-memory-space' and "
75 "'use-encoding-for-memory-space' are allowed in "
76 << getArgument();
77 return signalPassFailure();
78 }
79
80 if (mustInferMemorySpace) {
81 opt.defaultMemorySpaceFn =
82 [](TensorLikeType t) -> std::optional<Attribute> {
83 return std::nullopt;
84 };
85 }
86
87 if (useEncodingForMemorySpace) {
88 opt.defaultMemorySpaceFn =
89 [](TensorLikeType t) -> std::optional<Attribute> {
90 if (auto rtt = dyn_cast<RankedTensorType>(t))
91 return rtt.getEncoding();
92 return std::nullopt;
93 };
94 }
95
96 opt.printConflicts = printConflicts;
97 opt.bufferAlignment = bufferAlignment;
98 opt.testAnalysisOnly = testAnalysisOnly;
99 opt.bufferizeFunctionBoundaries = bufferizeFunctionBoundaries;
100 if (!mayHaveParallelRegions)
101 opt.mayHaveParallelRegions = false;
102 opt.noAnalysisFuncFilter = noAnalysisFuncFilter;
103
104 // Configure type converter.
105 LayoutMapOption unknownTypeConversionOption = unknownTypeConversion;
106 if (unknownTypeConversionOption == LayoutMapOption::InferLayoutMap) {
107 emitError(UnknownLoc::get(&getContext()),
108 "Invalid option: 'infer-layout-map' is not a valid value for "
109 "'unknown-type-conversion'");
110 return signalPassFailure();
111 }
112 opt.unknownTypeConverterFn = [=](TensorLikeType type,
113 Attribute memorySpace,
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,
125 memorySpace));
126 };
127
128 // Configure op filter.
129 OpFilter::Entry::FilterFn filterFn = [&](Operation *op) {
130 // Filter may be specified via options.
131 if (this->dialectFilter.hasValue() && !(*this->dialectFilter).empty())
132 return llvm::is_contained(this->dialectFilter,
133 op->getDialect()->getNamespace());
134 // No filter specified: All other ops are allowed.
135 return true;
136 };
137 opt.opFilter.allowOperation(filterFn);
138 } else {
139 opt = *options;
140 }
141
142 if (opt.copyBeforeWrite && opt.testAnalysisOnly) {
143 // These two flags do not make sense together: "copy-before-write"
144 // indicates that copies should be inserted before every memory write,
145 // but "test-analysis-only" indicates that only the analysis should be
146 // tested. (I.e., no IR is bufferized.)
147 emitError(UnknownLoc::get(&getContext()),
148 "Invalid option: 'copy-before-write' cannot be used with "
149 "'test-analysis-only'");
150 return signalPassFailure();
151 }
152
153 if (opt.printConflicts && !opt.testAnalysisOnly) {
154 emitError(
155 UnknownLoc::get(&getContext()),
156 "Invalid option: 'print-conflicts' requires 'test-analysis-only'");
157 return signalPassFailure();
158 }
159
160 if (opt.dumpAliasSets && !opt.testAnalysisOnly) {
161 emitError(
162 UnknownLoc::get(&getContext()),
163 "Invalid option: 'dump-alias-sets' requires 'test-analysis-only'");
164 return signalPassFailure();
165 }
166
167 BufferizationState state;
168 BufferizationStatistics statistics;
169 ModuleOp moduleOp = getOperation();
170 if (opt.bufferizeFunctionBoundaries) {
171 if (failed(
172 runOneShotModuleBufferize(moduleOp, opt, state, &statistics))) {
173 signalPassFailure();
174 return;
175 }
176 } else {
177 if (!opt.noAnalysisFuncFilter.empty()) {
178 emitError(UnknownLoc::get(&getContext()),
179 "Invalid option: 'no-analysis-func-filter' requires "
180 "'bufferize-function-boundaries'");
181 return signalPassFailure();
182 }
183 if (failed(runOneShotBufferize(moduleOp, opt, state, &statistics))) {
184 signalPassFailure();
185 return;
186 }
187 }
188
189 // Set pass statistics.
190 this->numBufferAlloc = statistics.numBufferAlloc;
191 this->numTensorInPlace = statistics.numTensorInPlace;
192 this->numTensorOutOfPlace = statistics.numTensorOutOfPlace;
193 }
194
195private:
196 std::optional<OneShotBufferizationOptions> options;
197};
198} // namespace
199
200//===----------------------------------------------------------------------===//
201// BufferizableOpInterface-based Bufferization
202//===----------------------------------------------------------------------===//
203
204namespace {
205/// A rewriter that keeps track of extra information during bufferization.
206class BufferizationRewriter : public IRRewriter, public RewriterBase::Listener {
207public:
208 BufferizationRewriter(MLIRContext *ctx, DenseSet<Operation *> &erasedOps,
209 DenseSet<Operation *> &toBufferOps,
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) {
215 setListener(this);
216 }
217
218protected:
219 void notifyOperationErased(Operation *op) override {
220 erasedOps.insert(op);
221 // Erase if present.
222 toBufferOps.erase(op);
223 }
224
225 void notifyOperationInserted(Operation *op, InsertPoint previous) override {
226 // We only care about newly created ops.
227 if (previous.isSet())
228 return;
229
230 erasedOps.erase(op);
231
232 // Gather statistics about allocs.
233 if (statistics) {
234 if (auto sideEffectingOp = dyn_cast<MemoryEffectOpInterface>(op))
235 statistics->numBufferAlloc += static_cast<int64_t>(
236 sideEffectingOp.hasEffect<MemoryEffects::Allocate>());
237 }
238
239 // Keep track of to_buffer ops.
240 if (isa<ToBufferOp>(op)) {
241 toBufferOps.insert(op);
242 return;
243 }
244
245 // Skip to_tensor ops.
246 if (isa<ToTensorOp>(op))
247 return;
248
249 // Skip non-tensor ops.
250 if (!hasTensorSemantics(op))
251 return;
252
253 // Skip ops that are not allowed to be bufferized.
254 auto const &options = analysisState.getOptions();
255 if (!options.isOpAllowed(op))
256 return;
257
258 // Add op to worklist.
259 worklist.push_back(op);
260 }
261
262private:
263 /// A set of all erased ops.
264 DenseSet<Operation *> &erasedOps;
265
266 /// A set of all to_buffer ops.
267 DenseSet<Operation *> &toBufferOps;
268
269 /// The worklist of ops to be bufferized.
270 SmallVector<Operation *> &worklist;
271
272 /// The analysis state. Used for debug assertions and access to the
273 /// bufferization options.
274 const AnalysisState analysisState;
275
276 /// Bufferization statistics for debugging.
277 BufferizationStatistics *statistics;
278};
279} // namespace
280
283 BufferizationState &bufferizationState,
284 BufferizationStatistics *statistics) {
285 if (options.copyBeforeWrite) {
286 AnalysisState analysisState(options);
287 if (failed(insertTensorCopies(op, analysisState, bufferizationState)))
288 return failure();
289 }
290
291 // Keep track of to_buffer ops.
292 DenseSet<Operation *> toBufferOps;
293 op->walk([&](ToBufferOp toBufferOp) { toBufferOps.insert(toBufferOp); });
294
295 // Gather all bufferizable ops in top-to-bottom order.
296 //
297 // We should ideally know the exact memref type of all operands when
298 // bufferizing an op. (This is the case when bufferizing top-to-bottom.)
299 // Otherwise, we have to use a memref type with a fully dynamic layout map to
300 // avoid copies. We are currently missing patterns for layout maps to
301 // canonicalize away (or canonicalize to more precise layouts).
303 op->walk<WalkOrder::PostOrder>([&](Operation *op) {
304 if (options.isOpAllowed(op) && hasTensorSemantics(op))
305 worklist.push_back(op);
306 });
307
308 // Keep track of all erased ops.
309 DenseSet<Operation *> erasedOps;
310
311 // Bufferize all ops.
312 BufferizationRewriter rewriter(op->getContext(), erasedOps, toBufferOps,
313 worklist, options, statistics);
314 for (unsigned i = 0; i < worklist.size(); ++i) {
315 Operation *nextOp = worklist[i];
316 // Skip ops that were erased.
317 if (erasedOps.contains(nextOp))
318 continue;
319 // Skip ops that are not bufferizable or not allowed.
320 auto bufferizableOp = options.dynCastBufferizableOp(nextOp);
321 if (!bufferizableOp)
322 continue;
323 // Skip ops that no longer have tensor semantics.
324 if (!hasTensorSemantics(nextOp))
325 continue;
326 // Check for unsupported unstructured control flow.
327 if (!bufferizableOp.supportsUnstructuredControlFlow())
328 for (Region &r : nextOp->getRegions())
329 if (r.getBlocks().size() > 1)
330 return nextOp->emitOpError(
331 "op or BufferizableOpInterface implementation does not support "
332 "unstructured control flow, but at least one region has multiple "
333 "blocks");
334
335 // Bufferize the op.
336 LDBG(3) << "//===-------------------------------------------===//\n"
337 << "IR after bufferizing: " << nextOp->getName();
338 rewriter.setInsertionPoint(nextOp);
339 if (failed(
340 bufferizableOp.bufferize(rewriter, options, bufferizationState))) {
341 LDBG(2) << "failed to bufferize\n"
342 << "//===-------------------------------------------===//";
343 return nextOp->emitError("failed to bufferize op");
344 }
345 LDBG(3) << *op << "\n//===-------------------------------------------===//";
346 }
347
348 // Return early if the top-level op is entirely gone.
349 if (erasedOps.contains(op))
350 return success();
351
352 // Fold all to_buffer(to_tensor(x)) pairs. Snapshot the set first:
353 // `foldToBufferToTensorPair` can erase ops, and the rewriter listener
354 // mutates `toBufferOps` from inside that call, which would invalidate
355 // any DenseSet iterator held across it.
356 SmallVector<Operation *> toBufferOpsSnapshot = llvm::to_vector(toBufferOps);
357 for (Operation *op : toBufferOpsSnapshot) {
358 if (erasedOps.contains(op))
359 continue;
360 rewriter.setInsertionPoint(op);
362 rewriter, cast<ToBufferOp>(op), options);
363 }
364
365 // Remove all dead to_tensor ops.
366 op->walk<WalkOrder::PostOrder>([&](ToTensorOp toTensorOp) {
367 if (toTensorOp->getUses().empty()) {
368 rewriter.eraseOp(toTensorOp);
369 return WalkResult::skip();
370 }
371 return WalkResult::advance();
372 });
373
374 /// Check the result of bufferization. Return an error if an op was not
375 /// bufferized, unless partial bufferization is allowed.
376 if (options.allowUnknownOps)
377 return success();
378
379 for (Operation *op : worklist) {
380 // Skip ops that are entirely gone.
381 if (erasedOps.contains(op))
382 continue;
383 // Ops that no longer have tensor semantics (because they were updated
384 // in-place) are allowed.
385 if (!hasTensorSemantics(op))
386 continue;
387 // Continue ops that are not allowed.
388 if (!options.isOpAllowed(op))
389 continue;
390 // Ops without any uses and no side effects will fold away.
391 if (op->getUses().empty() && isMemoryEffectFree(op))
392 continue;
393 // ToTensorOps/ToBufferOps are allowed in the output.
394 if (isa<ToTensorOp, ToBufferOp>(op))
395 continue;
396 return op->emitError("op was not bufferized");
397 }
398
399 return success();
400}
401
402LogicalResult
405 BufferizationState &state) {
406 OpBuilder::InsertionGuard g(rewriter);
407 auto bufferizableOp = options.dynCastBufferizableOp(block->getParentOp());
408 if (!bufferizableOp)
409 return failure();
410
411 // Compute the new signature.
412 SmallVector<Type> newTypes;
413 for (BlockArgument &bbArg : block->getArguments()) {
414 auto tensorType = dyn_cast<TensorLikeType>(bbArg.getType());
415 if (!tensorType) {
416 newTypes.push_back(bbArg.getType());
417 continue;
418 }
419
420 FailureOr<BufferLikeType> bufferType =
421 bufferization::getBufferType(bbArg, options, state);
422 if (failed(bufferType))
423 return failure();
424 newTypes.push_back(*bufferType);
425 }
426
427 // Change the type of all block arguments.
428 for (auto [bbArg, type] : llvm::zip(block->getArguments(), newTypes)) {
429 if (bbArg.getType() == type)
430 continue;
431
432 // Collect all uses of the bbArg.
433 SmallVector<OpOperand *> bbArgUses;
434 for (OpOperand &use : bbArg.getUses())
435 bbArgUses.push_back(&use);
436
437 Type tensorType = bbArg.getType();
438 // Change the bbArg type to memref.
439 bbArg.setType(type);
440
441 // Replace all uses of the original tensor bbArg.
442 rewriter.setInsertionPointToStart(block);
443 if (!bbArgUses.empty()) {
444 Value toTensorOp = bufferization::ToTensorOp::create(
445 rewriter, bbArg.getLoc(), tensorType, bbArg);
446 for (OpOperand *use : bbArgUses)
447 use->set(toTensorOp);
448 }
449 }
450
451 // Bufferize callers of the block. Iterate BlockOperands so each successor
452 // edge is handled once, including when one branch op targets this block more
453 // than once.
454 for (BlockOperand &blockOperand : block->getUses()) {
455 Operation *op = blockOperand.getOwner();
456 auto branchOp = dyn_cast<BranchOpInterface>(op);
457 if (!branchOp)
458 return op->emitOpError("cannot bufferize ops with block references that "
459 "do not implement BranchOpInterface");
460
461 SuccessorOperands operands =
462 branchOp.getSuccessorOperands(blockOperand.getOperandNumber());
463 SmallVector<Value> newOperands;
464 for (auto [operand, type] :
465 llvm::zip(operands.getForwardedOperands(), newTypes)) {
466 if (operand.getType() == type) {
467 // Not a tensor type. Nothing to do for this operand.
468 newOperands.push_back(operand);
469 continue;
470 }
471 FailureOr<BufferLikeType> operandBufferType =
472 bufferization::getBufferType(operand, options, state);
473 if (failed(operandBufferType))
474 return failure();
475 rewriter.setInsertionPointAfterValue(operand);
476 Value bufferizedOperand = bufferization::ToBufferOp::create(
477 rewriter, operand.getLoc(), *operandBufferType, operand);
478 // A cast is needed if the operand and the block argument have different
479 // bufferized types.
480 if (type != *operandBufferType) {
481 bufferizedOperand = *options.castFn(rewriter, operand.getLoc(), type,
482 bufferizedOperand);
483 }
484 newOperands.push_back(bufferizedOperand);
485 }
486 operands.getMutableForwardedOperands().assign(newOperands);
487 }
488
489 return success();
490}
return success()
b getContext())
static llvm::ManagedStatic< PassManagerOptions > options
Base class for generic analysis states.
Attributes are known-constant values of operations.
Definition Attributes.h:25
This class represents an argument of a Block.
Definition Value.h:306
A block operand represents an operand that holds a reference to a Block, e.g.
Block represents an ordered list of Operations.
Definition Block.h:34
BlockArgListType getArguments()
Definition Block.h:112
Operation * getParentOp()
Returns the closest surrounding operation that contains this block.
Definition Block.cpp:31
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.
Definition Builders.h:351
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
Definition Builders.h:434
void setInsertionPointAfterValue(Value val)
Sets the insertion point to the node after the specified value.
Definition Builders.h:424
This class represents an operand of an operation.
Definition Value.h:254
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
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.
Definition Operation.h:115
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
Definition Operation.h:729
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),...
Definition Operation.h:849
use_range getUses()
Returns a range of all uses, which is useful for iterating over all uses.
Definition Operation.h:898
MLIRContext * getContext()
Return the context this operation is associated with.
Definition Operation.h:233
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.
Definition Region.h:26
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...
Definition Types.h:74
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
static WalkResult skip()
Definition WalkResult.h:48
static WalkResult advance()
Definition WalkResult.h:47
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
Definition LLVM.h:122
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.
Definition Bufferize.h:35
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.