28#include "llvm/Support/Debug.h"
32#define DEBUG_TYPE "linalg-hoisting"
34#define DBGS() (dbgs() << '[' << DEBUG_TYPE << "] ")
49 Value newYieldValue) {
52 auto inits = llvm::to_vector(loop.getInits());
55 assert(
index < inits.size());
56 inits[
index] = newInitOperand;
58 scf::ForOp newLoop = scf::ForOp::create(
59 rewriter, loop.getLoc(), loop.getLowerBound(), loop.getUpperBound(),
61 loop.getUnsignedCmp());
64 auto yieldOp = cast<scf::YieldOp>(loop.getBody()->getTerminator());
65 yieldOp.setOperand(
index, newYieldValue);
68 rewriter.
mergeBlocks(loop.getBody(), newLoop.getBody(),
69 newLoop.getBody()->getArguments());
72 rewriter.
replaceOp(loop.getOperation(), newLoop->getResults());
101 root->
walk([&](vector::ExtractOp extractOp) {
102 LLVM_DEBUG(
DBGS() <<
"Candidate for hoisting: "
103 << *extractOp.getOperation() <<
"\n");
105 auto loop = dyn_cast<scf::ForOp>(extractOp->getParentOp());
110 auto blockArg = dyn_cast<BlockArgument>(extractOp.getSource());
115 OpOperand *initArg = loop.getTiedLoopInit(blockArg);
121 if (!blockArg.hasOneUse())
124 unsigned index = blockArg.getArgNumber() - loop.getNumInductionVars();
128 loop.getTiedLoopYieldedValue(blockArg)->get().getDefiningOp();
129 auto broadcast = dyn_cast<vector::BroadcastOp>(yieldedVal);
133 LLVM_DEBUG(
DBGS() <<
"Candidate broadcast: " <<
broadcast <<
"\n");
136 if (broadcastInputType != extractOp.getType())
141 for (
auto operand : extractOp.getDynamicPosition())
142 if (!loop.isDefinedOutsideOfLoop(operand))
146 extractOp.getSourceMutable().assign(initArg->
get());
148 loop.moveOutOfLoop(extractOp);
152 rewriter, loop, extractOp.getResult(),
index,
broadcast.getSource());
154 LLVM_DEBUG(
DBGS() <<
"New loop: " << newLoop <<
"\n");
167 bool verifyNonZeroTrip) {
181 if (verifyNonZeroTrip) {
182 root->
walk([&](LoopLikeOpInterface loopLike) {
183 std::optional<SmallVector<OpFoldResult>> lbs =
184 loopLike.getLoopLowerBounds();
185 std::optional<SmallVector<OpFoldResult>> ubs =
186 loopLike.getLoopUpperBounds();
194 for (
auto [lb,
ub] : llvm::zip_equal(lbs.value(), ubs.value())) {
195 FailureOr<int64_t> maxLb =
202 FailureOr<int64_t> minUb =
207 if (minUb.value() <= maxLb.value())
209 definiteNonZeroTripCountLoops.insert(loopLike);
214 root->
walk([&](vector::TransferReadOp transferRead) {
215 if (!isa<MemRefType>(transferRead.getShapedType()))
218 LLVM_DEBUG(
DBGS() <<
"Candidate for hoisting: "
219 << *transferRead.getOperation() <<
"\n");
220 auto loop = dyn_cast<LoopLikeOpInterface>(transferRead->getParentOp());
221 LLVM_DEBUG(
DBGS() <<
"Parent op: " << *transferRead->getParentOp()
223 if (!isa_and_nonnull<scf::ForOp, affine::AffineForOp>(loop))
226 if (verifyNonZeroTrip && !definiteNonZeroTripCountLoops.contains(loop)) {
227 LLVM_DEBUG(
DBGS() <<
"Loop may have zero trip count: " << *loop
232 LLVM_DEBUG(
DBGS() <<
"Candidate read: " << *transferRead.getOperation()
240 vector::TransferWriteOp transferWrite;
241 for (
auto *sliceOp : llvm::reverse(forwardSlice)) {
242 auto candidateWrite = dyn_cast<vector::TransferWriteOp>(sliceOp);
243 if (!candidateWrite ||
244 candidateWrite.getBase() != transferRead.getBase())
246 transferWrite = candidateWrite;
250 for (
auto operand : transferRead.getOperands())
251 if (!loop.isDefinedOutsideOfLoop(operand))
256 if (!transferWrite) {
262 loop.moveOutOfLoop(transferRead);
266 LLVM_DEBUG(
DBGS() <<
"Candidate: " << *transferWrite.getOperation()
279 if (transferRead.getIndices() != transferWrite.getIndices() ||
280 transferRead.getVectorType() != transferWrite.getVectorType() ||
281 transferRead.getPermutationMap() != transferWrite.getPermutationMap())
286 auto base = transferRead.getBase();
289 std::optional<bool> viewAliasingIsSafeCache;
290 auto viewAliasingIsSafe = [&]() {
291 if (!viewAliasingIsSafeCache) {
292 Operation *hoistedPair[] = {transferRead, transferWrite};
293 viewAliasingIsSafeCache =
296 return *viewAliasingIsSafeCache;
298 auto *source = base.getDefiningOp();
311 if (
auto assume = dyn_cast<memref::AssumeAlignmentOp>(source)) {
312 Value memPreAlignment = assume.getMemref();
314 llvm::count_if(base.getUses(), [&loop](
OpOperand &use) {
315 return loop->isAncestor(use.getOwner());
318 if (numInLoopUses && memPreAlignment.
hasOneUse())
321 if (isa_and_nonnull<ViewLikeOpInterface>(source) &&
322 !viewAliasingIsSafe())
326 if (llvm::any_of(base.getUsers(), llvm::IsaPred<ViewLikeOpInterface>) &&
327 !viewAliasingIsSafe())
336 for (
auto &use : transferRead.getBase().getUses()) {
337 if (!loop->isAncestor(use.getOwner()))
339 if (use.getOwner() == transferRead.getOperation() ||
340 use.getOwner() == transferWrite.getOperation())
342 if (
auto transferWriteUse =
343 dyn_cast<vector::TransferWriteOp>(use.getOwner())) {
345 cast<VectorTransferOpInterface>(*transferWrite),
346 cast<VectorTransferOpInterface>(*transferWriteUse),
349 }
else if (
auto transferReadUse =
350 dyn_cast<vector::TransferReadOp>(use.getOwner())) {
352 cast<VectorTransferOpInterface>(*transferWrite),
353 cast<VectorTransferOpInterface>(*transferReadUse),
364 loop.moveOutOfLoop(transferRead);
367 transferWrite->moveAfter(loop);
371 IRRewriter rewriter(transferRead.getContext());
377 auto maybeNewLoop = loop.replaceWithAdditionalYields(
378 rewriter, transferRead.getVector(),
380 if (failed(maybeNewLoop))
383 transferWrite.getValueToStoreMutable().assign(
384 maybeNewLoop->getOperation()->getResults().back());
static scf::ForOp replaceWithDifferentYield(RewriterBase &rewriter, scf::ForOp loop, Value newInitOperand, unsigned index, Value newYieldValue)
Replace loop with a new loop that has a different init operand at position index.
static Value broadcast(Location loc, Value toBroadcast, unsigned numElements, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value to vector with numElements number of elements.
A class for computing basic dominance information.
bool properlyDominates(Operation *a, Operation *b, bool enclosingOpOk=true) const
Return true if operation A properly dominates operation B, i.e.
IRValueT get() const
Return the current value being used by this operand.
This class coordinates rewriting a piece of IR outside of a pattern rewrite, providing a way to keep ...
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
This class represents an operand of an operation.
Operation is the basic unit of execution within MLIR.
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),...
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...
void mergeBlocks(Block *source, Block *dest, ValueRange argValues={})
Inline the operations of block 'source' into the end of block 'dest'.
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
void moveOpAfter(Operation *op, Operation *existingOp)
Unlink this operation from its current block and insert it right after existingOp which may be in the...
virtual void replaceAllUsesWith(Value from, Value to)
Find uses of from and replace them with to.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
static FailureOr< int64_t > computeConstantBound(presburger::BoundType type, const Variable &var, const StopConditionFn &stopCondition=nullptr, ValueBoundsOptions options={})
Compute a constant bound for the given variable.
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...
bool hasOneUse() const
Returns true if this value has exactly one use.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static WalkResult advance()
static WalkResult interrupt()
void hoistRedundantVectorBroadcasts(RewriterBase &rewriter, Operation *root)
Hoist vector.extract/vector.broadcast pairs out of immediately enclosing scf::ForOp iteratively,...
void hoistRedundantVectorTransfers(Operation *root, bool verifyNonZeroTrip=false)
Hoist vector.transfer_read/vector.transfer_write on buffers pairs out of immediately enclosing scf::F...
bool hasNoAliasingAccessInScope(Value base, Operation *scope, ArrayRef< Operation * > excludedOps={}, bool readsAreSafe=false)
Return "true" when no other access nested in scope can alias base.
bool isDisjointTransferSet(VectorTransferOpInterface transferA, VectorTransferOpInterface transferB, bool testDynamicValueUsingBounds=false)
Return true if we can prove that the transfer operations access disjoint memory, requiring the operat...
Include the generated interface declarations.
std::function< SmallVector< Value >( OpBuilder &b, Location loc, ArrayRef< BlockArgument > newBbArgs)> NewYieldValuesFn
A function that returns the additional yielded values during replaceWithAdditionalYields.
llvm::SetVector< T, Vector, Set, N > SetVector
size_t moveLoopInvariantCode(ArrayRef< Region * > regions, function_ref< bool(Value, Region *)> isDefinedOutsideRegion, function_ref< bool(Operation *, Region *)> shouldMoveOutOfRegion, function_ref< void(Operation *, Region *)> moveOutOfRegion)
Given a list of regions, perform loop-invariant code motion.
void getForwardSlice(Operation *op, SetVector< Operation * > *forwardSlice, const ForwardSliceOptions &options={})
Fills forwardSlice with the computed forward slice (i.e.
Options that control value bound computation.