[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