MLIR 24.0.0git
LoopRangeFolding.cpp
Go to the documentation of this file.
1//===- LoopRangeFolding.cpp - Code to perform loop range folding-----------===//
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//
9// This file implements loop range folding.
10//
11//===----------------------------------------------------------------------===//
12
14
20#include "mlir/IR/IRMapping.h"
21
22namespace mlir {
23#define GEN_PASS_DEF_SCFFORLOOPRANGEFOLDING
24#include "mlir/Dialect/SCF/Transforms/Passes.h.inc"
25} // namespace mlir
26
27using namespace mlir;
28using namespace mlir::scf;
29
30namespace {
31struct ForLoopRangeFolding
32 : public impl::SCFForLoopRangeFoldingBase<ForLoopRangeFolding> {
33 void runOnOperation() override;
34};
35} // namespace
36
37void ForLoopRangeFolding::runOnOperation() {
38 getOperation()->walk([&](ForOp op) {
39 Value indVar = op.getInductionVar();
40
41 auto canBeFolded = [&](Value value) {
42 return op.isDefinedOutsideOfLoop(value) || value == indVar;
43 };
44
45 // Fold until a fixed point is reached
46 while (true) {
47
48 // If the induction variable is used by more than one operation, we can't
49 // fold its arith ops into the loop range. Note that a single operation
50 // may use the induction variable several times (e.g. `arith.addi %i,
51 // %i`), so check for a single user rather than a single use.
52 if (indVar.use_empty())
53 break;
54
55 Operation *user = *indVar.getUsers().begin();
56 if (!llvm::all_of(indVar.getUsers(),
57 [&](Operation *u) { return u == user; }))
58 break;
59
60 if (!isa<arith::AddIOp, arith::MulIOp>(user))
61 break;
62
63 if (!llvm::all_of(user->getOperands(), canBeFolded))
64 break;
65
66 OpBuilder b(op);
67 IRMapping lbMap;
68 lbMap.map(indVar, op.getLowerBound());
69 IRMapping ubMap;
70 ubMap.map(indVar, op.getUpperBound());
71 IRMapping stepMap;
72 stepMap.map(indVar, op.getStep());
73
74 if (auto addOp = dyn_cast<arith::AddIOp>(user)) {
75 Operation *lbFold = b.clone(*user, lbMap);
76 Operation *ubFold = b.clone(*user, ubMap);
77
78 op.setLowerBound(lbFold->getResult(0));
79 op.setUpperBound(ubFold->getResult(0));
80
81 // `arith.addi %i, %i` is `2 * %i`, so the step has to be doubled too.
82 if (addOp.getLhs() == indVar && addOp.getRhs() == indVar) {
83 Operation *stepFold = b.clone(*user, stepMap);
84 op.setStep(stepFold->getResult(0));
85 }
86
87 } else if (auto mulOp = dyn_cast<arith::MulIOp>(user)) {
88 // Only fold if the multiplier is a known strictly positive constant.
89 // Multiplying by zero or a negative value would produce an invalid
90 // step (scf.for requires a strictly positive step).
91 Value multiplier =
92 (mulOp.getLhs() == indVar) ? mulOp.getRhs() : mulOp.getLhs();
93 std::optional<int64_t> multiplierVal = getConstantIntValue(multiplier);
94 if (!multiplierVal || *multiplierVal <= 0)
95 break;
96
97 Operation *lbFold = b.clone(*user, lbMap);
98 Operation *ubFold = b.clone(*user, ubMap);
99 Operation *stepFold = b.clone(*user, stepMap);
100
101 op.setLowerBound(lbFold->getResult(0));
102 op.setUpperBound(ubFold->getResult(0));
103 op.setStep(stepFold->getResult(0));
104 }
105
106 ValueRange wrapIndvar(indVar);
107 user->replaceAllUsesWith(wrapIndvar);
108 user->erase();
109 }
110 });
111}
112
113std::unique_ptr<Pass> mlir::createForLoopRangeFoldingPass() {
114 return std::make_unique<ForLoopRangeFolding>();
115}
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
Definition IRMapping.h:30
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
operand_range getOperands()
Returns an iterator on the underlying Value's.
Definition Operation.h:403
void replaceAllUsesWith(ValuesT &&values)
Replace all uses of results of this operation with the provided 'values'.
Definition Operation.h:297
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.
Definition Value.h:208
user_range getUsers() const
Definition Value.h:218
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.