23#define GEN_PASS_DEF_SCFFORLOOPRANGEFOLDING
24#include "mlir/Dialect/SCF/Transforms/Passes.h.inc"
31struct ForLoopRangeFolding
33 void runOnOperation()
override;
37void ForLoopRangeFolding::runOnOperation() {
38 getOperation()->walk([&](ForOp op) {
39 Value indVar = op.getInductionVar();
41 auto canBeFolded = [&](Value value) {
42 return op.isDefinedOutsideOfLoop(value) || value == indVar;
55 Operation *user = *indVar.
getUsers().begin();
57 [&](Operation *u) { return u == user; }))
60 if (!isa<arith::AddIOp, arith::MulIOp>(user))
63 if (!llvm::all_of(user->
getOperands(), canBeFolded))
68 lbMap.
map(indVar, op.getLowerBound());
70 ubMap.
map(indVar, op.getUpperBound());
72 stepMap.
map(indVar, op.getStep());
74 if (
auto addOp = dyn_cast<arith::AddIOp>(user)) {
75 Operation *lbFold =
b.clone(*user, lbMap);
76 Operation *ubFold =
b.clone(*user, ubMap);
82 if (addOp.getLhs() == indVar && addOp.getRhs() == indVar) {
83 Operation *stepFold =
b.clone(*user, stepMap);
87 }
else if (
auto mulOp = dyn_cast<arith::MulIOp>(user)) {
92 (mulOp.getLhs() == indVar) ? mulOp.getRhs() : mulOp.getLhs();
94 if (!multiplierVal || *multiplierVal <= 0)
97 Operation *lbFold =
b.
clone(*user, lbMap);
98 Operation *ubFold =
b.
clone(*user, ubMap);
99 Operation *stepFold =
b.
clone(*user, stepMap);
114 return std::make_unique<ForLoopRangeFolding>();
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
operand_range getOperands()
Returns an iterator on the underlying Value's.
void replaceAllUsesWith(ValuesT &&values)
Replace all uses of results of this operation with the provided 'values'.
Operation * clone(IRMapping &mapper, const CloneOptions &options=CloneOptions::all())
Create a deep copy of this operation, remapping any operands that use values outside of the operation...
void erase()
Remove this operation from its parent block and delete it.
bool use_empty() const
Returns true if this value has no uses.
user_range getUsers() const
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
std::unique_ptr< Pass > createForLoopRangeFoldingPass()
Creates a pass which folds arith ops on induction variable into loop range.