[llvm] [TailRecElim] Introduce support for shift accumulator optimization (PR #181331)
Antonio Frighetto via llvm-commits
llvm-commits at lists.llvm.org
Thu Jun 4 07:02:32 PDT 2026
================
@@ -374,29 +374,136 @@ static bool canMoveAboveCall(Instruction *I, CallInst *CI, AliasAnalysis *AA) {
return !is_contained(I->operands(), CI);
}
-static bool canTransformAccumulatorRecursion(Instruction *I, CallInst *CI) {
- if (!I->isAssociative() || !I->isCommutative())
+// While shifts are neither associative nor commutative, a chain of shifts by a
+// constant amount C is equivalent to a single shift by the sum of the amounts:
+// ... (Base << C) << C) ... << C == Base << (C * Iterations)
+// This relation applies to left shifts as well as arithmetic/logical right
+// shifts when the shift amount is a constant.
+static bool isPseudoAssociative(Instruction *I) {
+ if (!I->isShift())
return false;
+ return isa<ConstantInt>(I->getOperand(1));
+}
+
+// Find the base-case return value for function F: examine all
+// return instructions and pick return values that do not depend on a
+// recursive call to F. If there is exactly one distinct such value,
+// return it. If there are none or more than one distinct value, return
+// nullptr to indicate failure.
+//
+// FIXME: There is a room for improvement here in the future, e.g., consider
+// non-constant values and multiple base cases -- e.g., we want to be able to
+// handle code like:
+// ```
+// int f(int x) {
+// if (x == 1) return 1;
+// if (x == 10) return 10;
+// return f(x-1) << 1;
+// }
+// ```
+static Constant *getReturnValue(Function &F) {
+ Constant *BaseCaseVal = nullptr;
+
+ // Local lambda with conservative bail-out: It isn't so trivial to perform TRE
+ // if the return value depends on the result of a recursive call, because the
+ // return value will be different for different iterations of the recursion.
+ // So we ignore return values that depend on recursive calls.
+ auto TryGetConstantBaseCase = [&](Value *RV) -> Constant * {
+ // Case 1: direct constant return.
+ if (auto *C = dyn_cast<Constant>(RV))
+ return C;
+
+ // Case 2: PHI with one constant incoming.
+ if (auto *PN = dyn_cast<PHINode>(RV)) {
+ if (PN->getNumIncomingValues() != 2)
+ return nullptr;
+ for (unsigned I = 0; I < 2; ++I) {
+ if (auto *C = dyn_cast<Constant>(PN->getIncomingValue(I)))
+ return C;
+ }
+ return nullptr;
+ }
+
+ // Case 3: select with a constant arm.
+ if (auto *SI = dyn_cast<SelectInst>(RV)) {
+ if (auto *CTrue = dyn_cast<Constant>(SI->getTrueValue()))
+ return CTrue;
+ if (auto *CFalse = dyn_cast<Constant>(SI->getFalseValue()))
+ return CFalse;
+ return nullptr;
+ }
+
+ return nullptr;
+ };
+
+ for (BasicBlock &BB : F) {
+ if (auto *RI = dyn_cast<ReturnInst>(BB.getTerminator())) {
+ Value *RV = RI->getReturnValue();
+ Constant *Candidate = TryGetConstantBaseCase(RV);
+ if (!Candidate)
+ return nullptr;
+
+ if (!BaseCaseVal)
+ BaseCaseVal = Candidate;
+ else if (BaseCaseVal != Candidate)
+ return nullptr;
+ }
+ }
+
+ return BaseCaseVal;
+}
+
+namespace {
+struct AccumulatorRecursionInfo {
+ bool CanTransform = false;
+ Constant *BaseCaseConst = nullptr;
+};
----------------
antoniofrighetto wrote:
I don't think we need any additional struct here. What I meant with the previous refactoring opportunity was to set `BaseConstValue` in both paths in `canTransformAccumulatorRecursion` (the new one `IsPseudoAssoc`, and the pre-existing if we have an identity).
Renaming `BaseConstValue` field to, say, `AccumulatorInitValue`, then `canTransformAccumulatorRecursion` would look as follows:
```cpp
Constant *AccInitVal = nullptr;
if (IsPseudoAssoc) {
if (I->getOperand(0) != CI)
return nullptr;
AccInitVal = getReturnValue(*CI->getCalledFunction());
if (!AccInitVal)
return nullptr;
} else {
AccInitVal = ConstantExpr::getIdentity(I, I->getType());
if (!AccInitVal)
return nullptr;
if ((I->getOperand(0) == CI && I->getOperand(1) == CI) ||
(I->getOperand(0) != CI && I->getOperand(1) != CI))
return nullptr;
}
// ...
return AccInitVal;
```
Then the call-site in `eliminateCall` would be simplified to:
```cpp
Constant *AccInitVal = canTransformAccumulatorRecursion(&*BBI, CI);
if (AccPN || !AccInitVal)
return false; // We cannot eliminate the tail recursion!
AccRecInstr = &*BBI;
AccumulatorInitialValue = AccInitVal;
```
Some simplification in `insertAccumulator` too:
```cpp
for (pred_iterator PI = PB; PI != PE; ++PI) {
BasicBlock *P = *PI;
if (P == &F.getEntryBlock())
AccPN->addIncoming(AccumulatorInitialValue, P);
else
AccPN->addIncoming(AccPN, P);
}
```
We would only need some extra care at the end in `cleanupAndFinalize`, but I think we can reuse `isPseudoAssociative` for it:
```cpp
if (isPseudoAssociative(AccRecInstr)) {
// Base-case initialization.
RI->setOperand(0, AccPN);
} else {
// Existing logic to materialize the accumulator...
}
```
https://github.com/llvm/llvm-project/pull/181331
More information about the llvm-commits
mailing list