28#define GEN_PASS_DEF_BUFFERHOISTINGPASS
29#define GEN_PASS_DEF_BUFFERLOOPHOISTINGPASS
30#define GEN_PASS_DEF_PROMOTEBUFFERSTOSTACKPASS
31#include "mlir/Dialect/Bufferization/Transforms/Passes.h.inc"
41 return isa<LoopLikeOpInterface, RegionBranchOpInterface>(op);
51 if (isa<LoopLikeOpInterface>(op))
56 auto regionInterface = dyn_cast<RegionBranchOpInterface>(op);
60 return regionInterface.hasLoop();
72 auto allocOp = dyn_cast<AllocationOpInterface>(op);
80 auto allocOp = dyn_cast<AllocationOpInterface>(op);
89 unsigned maxRankOfAllocatedMemRef) {
90 auto type = dyn_cast<ShapedType>(alloc.
getType());
93 if (!type.hasStaticShape()) {
99 if (type.getRank() <= maxRankOfAllocatedMemRef) {
102 return operand.getDefiningOp<memref::RankOp>();
108 Type elemType = type.getElementType();
110 !isa<ComplexType, IndexType, VectorType>(elemType) &&
111 !isa<DataLayoutTypeInterface>(elemType))
118 std::optional<int64_t> numElements = type.tryGetNumElements();
123 *numElements >
static_cast<int64_t>(maximumSizeInBytes * 8ULL / bitwidth))
125 return *numElements * bitwidth <=
126 static_cast<int64_t>(maximumSizeInBytes * 8ULL);
133 for (
Value alias : aliases) {
134 for (
auto *use : alias.getUsers()) {
138 if (isa<RegionBranchTerminatorOpInterface>(use) &&
139 use->getParentRegion() == parentRegion)
175struct BufferAllocationHoistingStateBase {
177 DominanceInfo *dominators;
183 Block *placementBlock;
186 BufferAllocationHoistingStateBase(DominanceInfo *dominators, Value allocValue,
187 Block *placementBlock)
188 : dominators(dominators), allocValue(allocValue),
189 placementBlock(placementBlock) {}
193template <
typename StateT>
196 BufferAllocationHoisting(Operation *op)
197 : BufferPlacementTransformationBase(op), dominators(op),
198 postDominators(op), scopeOp(op) {}
202 SmallVector<Value> allocsAndAllocas;
204 allocsAndAllocas.push_back(std::get<0>(entry));
205 scopeOp->walk([&](memref::AllocaOp op) {
206 allocsAndAllocas.push_back(op.getMemref());
209 for (
auto allocValue : allocsAndAllocas) {
210 if (!StateT::shouldHoistOpType(allocValue.getDefiningOp()))
212 Operation *definingOp = allocValue.getDefiningOp();
213 assert(definingOp &&
"No defining op");
217 if (!dominators.isReachableFromEntry(allocValue.getParentBlock()))
220 auto resultAliases = aliases.resolve(allocValue);
222 Block *dominatorBlock =
225 StateT state(&dominators, allocValue, allocValue.getParentBlock());
228 Block *dependencyBlock =
nullptr;
232 for (Value depValue : operands) {
233 Block *depBlock = depValue.getParentBlock();
234 if (!dependencyBlock || dominators.dominates(dependencyBlock, depBlock))
235 dependencyBlock = depBlock;
241 Block *placementBlock = findPlacementBlock(
242 state, state.computeUpperBound(dominatorBlock, dependencyBlock));
244 allocValue, placementBlock, liveness);
247 Operation *allocOperation = allocValue.getDefiningOp();
256 Block *findPlacementBlock(StateT &state,
Block *upperBound) {
257 Block *currentBlock = state.placementBlock;
268 (parentBlock = parentOp->
getBlock()) &&
270 dominators.properlyDominates(upperBound, currentBlock))) {
273 if (!dominators.isReachableFromEntry(currentBlock))
282 idom = dominators.getNode(currentBlock)->getIDom();
284 if (idom && dominators.properlyDominates(parentBlock, idom->getBlock())) {
287 currentBlock = idom->getBlock();
288 state.recordMoveToDominator(currentBlock);
292 if (parentOp == scopeOp ||
293 parentOp->
hasTrait<OpTrait::IsIsolatedFromAbove>() ||
295 !state.isLegalPlacement(parentOp))
299 currentBlock = parentBlock;
300 state.recordMoveToParent(currentBlock);
304 return state.placementBlock;
309 DominanceInfo dominators;
313 PostDominanceInfo postDominators;
316 llvm::DenseMap<Value, Block *> placementBlocks;
326struct BufferAllocationHoistingState : BufferAllocationHoistingStateBase {
327 using BufferAllocationHoistingStateBase::BufferAllocationHoistingStateBase;
330 Block *computeUpperBound(
Block *dominatorBlock,
Block *dependencyBlock) {
333 if (!dependencyBlock)
334 return dominatorBlock;
338 return dominators->properlyDominates(dominatorBlock, dependencyBlock)
344 bool isLegalPlacement(Operation *op) {
return !
isLoop(op); }
347 static bool shouldHoistOpType(Operation *op) {
352 void recordMoveToDominator(
Block *block) { placementBlock = block; }
355 void recordMoveToParent(
Block *block) { recordMoveToDominator(block); }
360struct BufferAllocationLoopHoistingState : BufferAllocationHoistingStateBase {
361 using BufferAllocationHoistingStateBase::BufferAllocationHoistingStateBase;
364 Block *aliasDominatorBlock =
nullptr;
367 Block *computeUpperBound(
Block *dominatorBlock,
Block *dependencyBlock) {
368 aliasDominatorBlock = dominatorBlock;
371 return dependencyBlock ? dependencyBlock :
nullptr;
379 bool isLegalPlacement(Operation *op) {
381 !dominators->dominates(aliasDominatorBlock, op->
getBlock());
385 static bool shouldHoistOpType(Operation *op) {
391 void recordMoveToDominator(
Block *block) {}
394 void recordMoveToParent(
Block *block) { placementBlock = block; }
404 BufferPlacementPromotion(Operation *op)
405 : BufferPlacementTransformationBase(op) {}
410 Value alloc = std::get<0>(entry);
411 Operation *dealloc = std::get<1>(entry);
416 if (!isSmallAlloc(alloc) || dealloc ||
424 OpBuilder builder(startOperation);
426 if (
auto allocInterface = dyn_cast<AllocationOpInterface>(allocOp)) {
427 std::optional<Operation *> alloca =
428 allocInterface.buildPromotedAlloc(builder, alloc);
445struct BufferHoistingPass
446 :
public bufferization::impl::BufferHoistingPassBase<BufferHoistingPass> {
448 void runOnOperation()
override {
450 BufferAllocationHoisting<BufferAllocationHoistingState> optimizer(
457struct BufferLoopHoistingPass
458 :
public bufferization::impl::BufferLoopHoistingPassBase<
459 BufferLoopHoistingPass> {
461 void runOnOperation()
override {
469class PromoteBuffersToStackPass
470 :
public bufferization::impl::PromoteBuffersToStackPassBase<
471 PromoteBuffersToStackPass> {
475 explicit PromoteBuffersToStackPass(std::function<
bool(Value)> isSmallAlloc)
476 : isSmallAlloc(std::move(isSmallAlloc)) {}
478 LogicalResult
initialize(MLIRContext *context)
override {
479 if (isSmallAlloc ==
nullptr) {
480 isSmallAlloc = [=](Value alloc) {
482 maxRankOfAllocatedMemRef);
488 void runOnOperation()
override {
490 BufferPlacementPromotion optimizer(getOperation());
491 optimizer.promote(isSmallAlloc);
495 std::function<bool(Value)> isSmallAlloc;
501 BufferAllocationHoisting<BufferAllocationLoopHoistingState> optimizer(op);
506 std::function<
bool(
Value)> isSmallAlloc) {
507 return std::make_unique<PromoteBuffersToStackPass>(std::move(isSmallAlloc));
static bool leavesAllocationScope(Region *parentRegion, const BufferViewFlowAnalysis::ValueSetT &aliases)
Checks whether the given aliases leave the allocation scope.
static bool isKnownControlFlowInterface(Operation *op)
Returns true if the given operation implements a known high-level region- based control-flow interfac...
static bool hasAllocationScope(Value alloc, const BufferViewFlowAnalysis &aliasAnalysis)
Checks, if an automated allocation scope for a given alloc value exists.
static bool isSequentialLoop(Operation *op)
Return whether the given operation is a loop with sequential execution semantics.
static bool isLoop(Operation *op)
Returns true if the given operation represents a loop by testing whether it implements the LoopLikeOp...
static bool allowAllocDominateBlockHoisting(Operation *op)
Returns true if the given operation implements the AllocationOpInterface and it supports the dominate...
static bool allowAllocLoopHoisting(Operation *op)
Returns true if the given operation implements the AllocationOpInterface and it supports the loop hoi...
static bool defaultIsSmallAlloc(Value alloc, unsigned maximumSizeInBytes, unsigned maxRankOfAllocatedMemRef)
Check if the size of the allocation is less than the given size.
LogicalResult initialize(unsigned origNumLoops, ArrayRef< ReassociationIndices > foldedIterationDims)
bool isEntryBlock()
Return if this block is the entry block in the parent region.
Operation * getParentOp()
Returns the closest surrounding operation that contains this block.
A straight-forward alias analysis which ensures that all dependencies of all values will be determine...
SmallPtrSet< Value, 16 > ValueSetT
ValueSetT resolve(Value value) const
Find all immediate and indirect views upon this value.
static DataLayout closest(Operation *op)
Returns the layout of the closest parent operation carrying layout info.
llvm::TypeSize getTypeSizeInBits(Type t) const
Returns the size in bits of the given type in the current scope.
A trait of region holding operations that define a new scope for automatic allocations,...
Operation is the basic unit of execution within MLIR.
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Block * getBlock()
Returns the operation block that contains this operation.
operand_range getOperands()
Returns an iterator on the underlying Value's.
void moveBefore(Operation *existingOp)
Unlink this operation from its current block and insert it right before existingOp which may be in th...
void replaceAllUsesWith(ValuesT &&values)
Replace all uses of results of this operation with the provided 'values'.
void erase()
Remove this operation from its parent block and delete it.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
Region * getParentRegion()
Return the region containing this region or nullptr if the region is attached to a top-level operatio...
Operation * getParentOp()
Return the parent operation this region is attached to.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
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.
Block * getParentBlock()
Return the Block in which this Value is defined.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Region * getParentRegion()
Return the Region in which this Value is defined.
static Operation * getStartOperation(Value allocValue, Block *placementBlock, const Liveness &liveness)
Get the start operation to place the given alloc value within the specified placement block.
std::tuple< Value, Operation * > AllocEntry
Represents a tuple of allocValue and deallocOperation.
void hoistBuffersFromLoops(Operation *op)
Within the given operation, hoist buffers from loops where possible.
std::unique_ptr< Pass > createPromoteBuffersToStackPass(std::function< bool(Value)> isSmallAlloc)
Creates a pass that promotes heap-based allocations to stack-based ones.
Block * findCommonDominator(Value value, const BufferViewFlowAnalysis::ValueSetT &values, const DominatorT &doms)
Finds a common dominator for the given value while taking the positions of the values in the value se...
void promote(RewriterBase &rewriter, scf::ForallOp forallOp)
Promotes the loop body of a scf::ForallOp to its containing block.
Include the generated interface declarations.
llvm::DomTreeNodeBase< Block > DominanceInfoNode
llvm::function_ref< Fn > function_ref