[llvm] ed5f9f9 - [TailRecElim] Introduce support for shift accumulator optimization (#181331)

via llvm-commits llvm-commits at lists.llvm.org
Wed Aug 5 23:45:40 PDT 2026


Author: Federico Bruzzone
Date: 2026-08-06T08:45:35+02:00
New Revision: ed5f9f9dc15fd5ea18f0da1be5a1b8be8b21b762

URL: https://github.com/llvm/llvm-project/commit/ed5f9f9dc15fd5ea18f0da1be5a1b8be8b21b762
DIFF: https://github.com/llvm/llvm-project/commit/ed5f9f9dc15fd5ea18f0da1be5a1b8be8b21b762.diff

LOG: [TailRecElim] Introduce support for shift accumulator optimization (#181331)

This PR enables Tail Recursion Elimination (TRE) for functions where the
accumulator operation is a shift (`shl`, `lshr`, `ashr`) by a constant
amount -- i.e., pseudo-associative relation.

As pointed out in #178805, `InstCombine` often strength-reduces
multiplications (or `f(x-1) + f(x-1)`) into `shl`.
Currently, TRE strictly requires operations to be associative and
commutative:

https://github.com/llvm/llvm-project/blob/05e908609227e1e8d993659e604a63668dfd2825/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp#L377-L379
This prevents TRE from transforming recursive shifts into loops,
creating a phase-ordering problem where canonicalization blocks a
structural optimization.

This PR does **not** perform shift accumulator optimization when there
are multiple base cases: it is reserved for future work.

Fixes #178805.

Added: 
    llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll

Modified: 
    llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp

Removed: 
    


################################################################################
diff  --git a/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp b/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
index 4ce2f08524baa..14e99ffc600bc 100644
--- a/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
+++ b/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
@@ -415,29 +415,126 @@ 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())
+// Return true if I is a unary accumulator recurrence: a chain of
+// applications of a unary function `g` composed with itself,
+// `g(g(...g(Base)...))`, which is equivalent to a single application of the
+// N-times-composed function when `g` is pure. Neither associative nor
+// commutative, this 
diff ers from the ordinary accumulator recurrence handled
+// below, which requires I to be associative and commutative.
+//
+// TODO: Generalize this beyond shifts by a constant amount to arbitrary pure
+// unary functions (e.g., `f(x) = x == 0 ? Base : g(f(x - 1))` for any pure
+// unary `g`).
+static bool isUnaryAccumulatorRecurrence(Instruction *I) {
+  if (!I->isShift())
     return false;
 
+  // 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.
+  return isa<ConstantInt>(I->getOperand(1));
+}
+
+// Return true if V is a recursive call to F or an instruction directly using
+// the result of one. A depth-1 check is enough here: the value feeding a
+// return either uses the recursive call as an immediate operand (the
+// accumulator instruction, or the PHI merging it with the base case), or it
+// is rejected by findBaseCaseRetConstant below as a non-constant anyway.
+static bool usesRecursiveCall(Value *V, Function &F) {
+  auto IsRecursiveCall = [&F](Value *V) {
+    auto *CI = dyn_cast<CallInst>(V);
+    return CI && CI->getCalledFunction() == &F;
+  };
+  if (IsRecursiveCall(V))
+    return true;
+  auto *I = dyn_cast<Instruction>(V);
+  return I && llvm::any_of(I->operands(), IsRecursiveCall);
+}
+
+// Find the base-case return value for function F: examine all return
+// instructions, skipping those whose return value depends on a recursive call
+// to F (that value 
diff ers for each iteration of the recursion). If the
+// remaining returns yield exactly one distinct constant, return it; otherwise
+// 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 *findBaseCaseRetConstant(Function &F) {
+  Constant *BaseCaseVal = nullptr;
+
+  for (BasicBlock &BB : F) {
+    auto *RI = dyn_cast<ReturnInst>(BB.getTerminator());
+    if (!RI || !RI->getReturnValue())
+      continue;
+
+    Value *RV = RI->getReturnValue();
+    if (usesRecursiveCall(RV, F))
+      continue;
+
+    auto *C = dyn_cast<Constant>(RV);
+    if (!C)
+      return nullptr;
+
+    if (!BaseCaseVal)
+      BaseCaseVal = C;
+    else if (BaseCaseVal != C)
+      return nullptr;
+  }
+
+  return BaseCaseVal;
+}
+
+// This function checks whether the instruction I can be used
+// to perform accumulator recursion elimination for the
+// call instruction CI.
+static Constant *canTransformAccumulatorRecursion(Instruction *I,
+                                                  CallInst *CI) {
+  bool IsUnaryAccumulatorRecurrence = isUnaryAccumulatorRecurrence(I);
+  if ((!I->isAssociative() || !I->isCommutative()) &&
+      !IsUnaryAccumulatorRecurrence)
+    return nullptr;
+
   assert(I->getNumOperands() >= 2 &&
          "Associative/commutative operations should have at least 2 args!");
 
-  if (IntrinsicInst *II = dyn_cast<IntrinsicInst>(I)) {
-    // Accumulators must have an identity.
-    if (!ConstantExpr::getIntrinsicIdentity(II->getIntrinsicID(), I->getType()))
-      return false;
-  }
+  Constant *AccInitVal = nullptr;
+  if (IsUnaryAccumulatorRecurrence) {
+    // For unary accumulator recurrences, we require that the recursive call
+    // is always on the first operand.
+    if (I->getOperand(0) != CI)
+      return nullptr;
 
-  // Exactly one operand should be the result of the call instruction.
-  if ((I->getOperand(0) == CI && I->getOperand(1) == CI) ||
-      (I->getOperand(0) != CI && I->getOperand(1) != CI))
-    return false;
+    // findTRECandidate guarantees CI is a recursive call to its own
+    // function, so scan the enclosing function for the base-case return.
+    AccInitVal = findBaseCaseRetConstant(*CI->getFunction());
+    if (!AccInitVal)
+      return nullptr;
+  } else {
+    AccInitVal = ConstantExpr::getIdentity(I, I->getType());
+    if (!AccInitVal)
+      return nullptr;
+
+    // Exactly one operand should be the result of the call instruction.
+    if ((I->getOperand(0) == CI && I->getOperand(1) == CI) ||
+        (I->getOperand(0) != CI && I->getOperand(1) != CI))
+      return nullptr;
+  }
 
   // The only user of this instruction we allow is a single return instruction.
   if (!I->hasOneUse() || !isa<ReturnInst>(I->user_back()))
-    return false;
+    return nullptr;
 
-  return true;
+  return AccInitVal;
 }
 
 namespace {
@@ -479,6 +576,8 @@ class TailRecursionEliminator {
   // The instruction doing the accumulating.
   Instruction *AccumulatorRecursionInstr = nullptr;
 
+  Constant *AccumulatorInitialValue = nullptr;
+
   TailRecursionEliminator(Function &F, const TargetTransformInfo *TTI,
                           AliasAnalysis *AA, OptimizationRemarkEmitter *ORE,
                           DomTreeUpdater &DTU, BlockFrequencyInfo *BFI,
@@ -640,9 +739,7 @@ void TailRecursionEliminator::insertAccumulator(Instruction *AccRecInstr) {
   for (pred_iterator PI = PB; PI != PE; ++PI) {
     BasicBlock *P = *PI;
     if (P == &F.getEntryBlock()) {
-      Constant *Identity =
-          ConstantExpr::getIdentity(AccRecInstr, AccRecInstr->getType());
-      AccPN->addIncoming(Identity, P);
+      AccPN->addIncoming(AccumulatorInitialValue, P);
     } else {
       AccPN->addIncoming(AccPN, P);
     }
@@ -713,15 +810,22 @@ bool TailRecursionEliminator::eliminateCall(CallInst *CI) {
       continue;
 
     // If we can't move the instruction above the call, it might be because it
-    // is an associative and commutative operation that could be transformed
-    // using accumulator recursion elimination.  Check to see if this is the
-    // case, and if so, remember which instruction accumulates for later.
-    if (AccPN || !canTransformAccumulatorRecursion(&*BBI, CI))
+    // is an (associative and commutative) or unary accumulator recurrence
+    // arithmetic operation that could be transformed using accumulator
+    // recursion elimination. Check to see if this is the case, and if so,
+    // remember which instruction accumulates for later.
+    Constant *AccInitVal = canTransformAccumulatorRecursion(&*BBI, CI);
+
+    if (AccPN || !AccInitVal)
       return false; // We cannot eliminate the tail recursion!
 
     // Yes, this is accumulator recursion.  Remember which instruction
     // accumulates.
     AccRecInstr = &*BBI;
+
+    // Keep track of the base case (i.e., initial value) of the accumulator
+    // return value if any.
+    AccumulatorInitialValue = AccInitVal;
   }
 
   BasicBlock *BB = Ret->getParent();
@@ -847,6 +951,17 @@ void TailRecursionEliminator::cleanupAndFinalize() {
   }
 
   if (RetPN) {
+    Instruction *AccRecInstr = AccumulatorRecursionInstr;
+    auto MaterializeAccumulator = [&](Value *OtherVal,
+                                      BasicBlock::iterator InsertPt) {
+      Instruction *New = AccRecInstr->clone();
+      New->setName("accumulator.ret.tr");
+      New->setOperand(AccRecInstr->getOperand(0) == AccPN, OtherVal);
+      New->insertBefore(InsertPt);
+      New->dropLocation();
+      return New;
+    };
+
     if (RetSelects.empty()) {
       // If we didn't insert any select instructions, then we know we didn't
       // store a return value and we can remove the PHI nodes we inserted.
@@ -859,19 +974,23 @@ void TailRecursionEliminator::cleanupAndFinalize() {
       if (AccPN) {
         // We need to insert a copy of our accumulator instruction before any
         // return in the function, and return its result instead.
-        Instruction *AccRecInstr = AccumulatorRecursionInstr;
         for (BasicBlock &BB : F) {
           ReturnInst *RI = dyn_cast<ReturnInst>(BB.getTerminator());
           if (!RI)
             continue;
 
-          Instruction *AccRecInstrNew = AccRecInstr->clone();
-          AccRecInstrNew->setName("accumulator.ret.tr");
-          AccRecInstrNew->setOperand(AccRecInstr->getOperand(0) == AccPN,
-                                     RI->getOperand(0));
-          AccRecInstrNew->insertBefore(RI->getIterator());
-          AccRecInstrNew->dropLocation();
-          RI->setOperand(0, AccRecInstrNew);
+          if (isUnaryAccumulatorRecurrence(AccRecInstr)) {
+            // Base-case initialization: the accumulator PHI already holds the
+            // final result, so return it directly.
+            RI->setOperand(0, AccPN);
+          } else {
+            // Since the accumulator starts with the identity value, before the
+            // return we need to apply the accumulation instruction one more
+            // time to combine the last value with the result of the recursive
+            // call.
+            RI->setOperand(0, MaterializeAccumulator(RI->getOperand(0),
+                                                     RI->getIterator()));
+          }
         }
       }
     } else {
@@ -893,15 +1012,13 @@ void TailRecursionEliminator::cleanupAndFinalize() {
       if (AccPN) {
         // We need to insert a copy of our accumulator instruction before any
         // of the selects we inserted, and select its result instead.
-        Instruction *AccRecInstr = AccumulatorRecursionInstr;
         for (SelectInst *SI : RetSelects) {
-          Instruction *AccRecInstrNew = AccRecInstr->clone();
-          AccRecInstrNew->setName("accumulator.ret.tr");
-          AccRecInstrNew->setOperand(AccRecInstr->getOperand(0) == AccPN,
-                                     SI->getFalseValue());
-          AccRecInstrNew->insertBefore(SI->getIterator());
-          AccRecInstrNew->dropLocation();
-          SI->setFalseValue(AccRecInstrNew);
+          if (isUnaryAccumulatorRecurrence(AccRecInstr)) {
+            SI->setFalseValue(AccPN);
+          } else {
+            SI->setFalseValue(
+                MaterializeAccumulator(SI->getFalseValue(), SI->getIterator()));
+          }
         }
       }
     }
@@ -938,7 +1055,9 @@ bool TailRecursionEliminator::processBlock(BasicBlock &BB) {
 
     eliminateCall(CI);
     return true;
-  } else if (isa<ReturnInst>(TI)) {
+  }
+
+  if (isa<ReturnInst>(TI)) {
     CallInst *CI = findTRECandidate(&BB);
 
     if (CI)

diff  --git a/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll b/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
new file mode 100644
index 0000000000000..41db471f2dbf0
--- /dev/null
+++ b/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
@@ -0,0 +1,188 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 5
+; RUN: opt < %s -passes="tailcallelim" -verify-dom-info -S | FileCheck %s
+
+; NOTE: All the following test cases are generate from the underlying C code (-O1)
+;   before that the shift accumulator optimization was implemented
+
+
+
+; InstCombine strength-reduce `f(x-1) + f(x-1)` to shl:
+; int f(int x) {
+;     if (x == 1) return 7;
+;     return f(x-1) + f(x-1); // f(x-1) * 2
+; }
+define i32 @test_shl_const_accumulator(i32 %x) {
+; CHECK-LABEL: define i32 @test_shl_const_accumulator(
+; CHECK-SAME: i32 [[X:%.*]]) {
+; CHECK-NEXT:  [[TAILRECURSE:.*]]:
+; CHECK-NEXT:    br label %[[COMMON_RET:.*]]
+; CHECK:       [[COMMON_RET]]:
+; CHECK-NEXT:    [[ADD:%.*]] = phi i32 [ 7, %[[TAILRECURSE]] ], [ [[ADD1:%.*]], %[[IF_END1:.*]] ]
+; CHECK-NEXT:    [[X_TR:%.*]] = phi i32 [ [[X]], %[[TAILRECURSE]] ], [ [[SUB:%.*]], %[[IF_END1]] ]
+; CHECK-NEXT:    [[CMP:%.*]] = icmp eq i32 [[X_TR]], 1
+; CHECK-NEXT:    br i1 [[CMP]], label %[[IF_END:.*]], label %[[IF_END1]]
+; CHECK:       [[IF_END]]:
+; CHECK-NEXT:    ret i32 [[ADD]]
+; CHECK:       [[IF_END1]]:
+; CHECK-NEXT:    [[SUB]] = add nsw i32 [[X_TR]], -1
+; CHECK-NEXT:    [[ADD1]] = shl i32 [[ADD]], 1
+; CHECK-NEXT:    br label %[[COMMON_RET]]
+;
+entry:
+  %cmp = icmp eq i32 %x, 1
+  br i1 %cmp, label %common.ret, label %if.end
+
+common.ret:
+  %common.ret.op = phi i32 [ %add, %if.end ], [ 7, %entry ]
+  ret i32 %common.ret.op
+
+if.end:
+  %sub = add nsw i32 %x, -1
+  %call = tail call i32 @test_shl_const_accumulator(i32 %sub)
+  %add = shl nsw i32 %call, 1
+  br label %common.ret
+}
+
+
+; int f2(int x) {
+;     if (x == 1) return 14;
+;     return f2(x-1) >> 1;
+; }
+define i32 @test_ashr_const_accumulator(i32 %x) {
+; CHECK-LABEL: define i32 @test_ashr_const_accumulator(
+; CHECK-SAME: i32 [[X:%.*]]) {
+; CHECK-NEXT:  [[TAILRECURSE:.*]]:
+; CHECK-NEXT:    br label %[[TAILRECURSE1:.*]]
+; CHECK:       [[TAILRECURSE1]]:
+; CHECK-NEXT:    [[ACCUMULATOR_TR:%.*]] = phi i32 [ 14, %[[TAILRECURSE]] ], [ [[SHR:%.*]], %[[IF_END:.*]] ]
+; CHECK-NEXT:    [[X_TR:%.*]] = phi i32 [ [[X]], %[[TAILRECURSE]] ], [ [[SUB:%.*]], %[[IF_END]] ]
+; CHECK-NEXT:    [[CMP:%.*]] = icmp eq i32 [[X_TR]], 1
+; CHECK-NEXT:    br i1 [[CMP]], label %[[COMMON_RET:.*]], label %[[IF_END]]
+; CHECK:       [[COMMON_RET]]:
+; CHECK-NEXT:    ret i32 [[ACCUMULATOR_TR]]
+; CHECK:       [[IF_END]]:
+; CHECK-NEXT:    [[SUB]] = add nsw i32 [[X_TR]], -1
+; CHECK-NEXT:    [[SHR]] = ashr i32 [[ACCUMULATOR_TR]], 1
+; CHECK-NEXT:    br label %[[TAILRECURSE1]]
+;
+entry:
+  %cmp = icmp eq i32 %x, 1
+  br i1 %cmp, label %common.ret, label %if.end
+
+common.ret:
+  %common.ret.op = phi i32 [ %shr, %if.end ], [ 14, %entry ]
+  ret i32 %common.ret.op
+
+if.end:
+  %sub = add nsw i32 %x, -1
+  %call = tail call i32 @test_ashr_const_accumulator(i32 %sub)
+  %shr = ashr i32 %call, 1
+  br label %common.ret
+}
+
+
+; unsigned int f3(unsigned int x) {
+;     if (x <= 1) return 21;
+;     return f3(x - 1) >> 1;
+; }
+define i32 @test_lshr_const_unsigned(i32 %x) {
+; CHECK-LABEL: define i32 @test_lshr_const_unsigned(
+; CHECK-SAME: i32 [[X:%.*]]) {
+; CHECK-NEXT:  [[TAILRECURSE:.*]]:
+; CHECK-NEXT:    br label %[[TAILRECURSE1:.*]]
+; CHECK:       [[TAILRECURSE1]]:
+; CHECK-NEXT:    [[ACCUMULATOR_TR:%.*]] = phi i32 [ 21, %[[TAILRECURSE]] ], [ [[SHR:%.*]], %[[IF_END:.*]] ]
+; CHECK-NEXT:    [[X_TR:%.*]] = phi i32 [ [[X]], %[[TAILRECURSE]] ], [ [[SUB:%.*]], %[[IF_END]] ]
+; CHECK-NEXT:    [[CMP:%.*]] = icmp ult i32 [[X_TR]], 2
+; CHECK-NEXT:    br i1 [[CMP]], label %[[COMMON_RET:.*]], label %[[IF_END]]
+; CHECK:       [[COMMON_RET]]:
+; CHECK-NEXT:    ret i32 [[ACCUMULATOR_TR]]
+; CHECK:       [[IF_END]]:
+; CHECK-NEXT:    [[SUB]] = add i32 [[X_TR]], -1
+; CHECK-NEXT:    [[SHR]] = lshr i32 [[ACCUMULATOR_TR]], 1
+; CHECK-NEXT:    br label %[[TAILRECURSE1]]
+;
+entry:
+  %cmp = icmp ult i32 %x, 2
+  br i1 %cmp, label %common.ret, label %if.end
+
+common.ret:
+  %common.ret.op = phi i32 [ %shr, %if.end ], [ 21, %entry ]
+  ret i32 %common.ret.op
+
+if.end:
+  %sub = add i32 %x, -1
+  %call = tail call i32 @test_lshr_const_unsigned(i32 %sub)
+  %shr = lshr i32 %call, 1
+  br label %common.ret
+}
+
+; Negative test: non-constant shift amount prevents shift-accumulator optimization
+; int f4(int x, int k) {
+;   if (x == 1) return 7;
+;   return f4(x-1, k) << k; // variable shift amount
+; }
+define i32 @test_neg_variable_shift_amount(i32 %x, i32 %k) {
+; CHECK-LABEL: define i32 @test_neg_variable_shift_amount(
+; CHECK-SAME: i32 [[X:%.*]], i32 [[K:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    [[CMP:%.*]] = icmp eq i32 [[X]], 1
+; CHECK-NEXT:    br i1 [[CMP]], label %[[COMMON_RET:.*]], label %[[IF_END:.*]]
+; CHECK:       [[COMMON_RET]]:
+; CHECK-NEXT:    ret i32 7
+; CHECK:       [[IF_END]]:
+; CHECK-NEXT:    [[SUB:%.*]] = add nsw i32 [[X]], -1
+; CHECK-NEXT:    [[CALL:%.*]] = tail call i32 @test_neg_variable_shift_amount(i32 [[SUB]], i32 [[K]])
+; CHECK-NEXT:    [[SHL:%.*]] = shl i32 [[CALL]], [[K]]
+; CHECK-NEXT:    ret i32 [[SHL]]
+;
+entry:
+  %cmp = icmp eq i32 %x, 1
+  br i1 %cmp, label %common.ret, label %if.end
+
+common.ret:
+  %common.ret.op = phi i32 [ %shl, %if.end ], [ 7, %entry ]
+  ret i32 %common.ret.op
+
+if.end:
+  %sub = add nsw i32 %x, -1
+  %call = tail call i32 @test_neg_variable_shift_amount(i32 %sub, i32 %k)
+  %shl = shl i32 %call, %k
+  br label %common.ret
+}
+
+; Negative test: intervening operation on result prevents forming clean accumulator
+; int f5(int x) {
+;   if (x == 1) return 7;
+;   return (f5(x-1) + 1) << 1; // extra add breaks accumulator invariant
+; }
+define i32 @test_neg_intervening_op(i32 %x) {
+; CHECK-LABEL: define i32 @test_neg_intervening_op(
+; CHECK-SAME: i32 [[X:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*:]]
+; CHECK-NEXT:    [[CMP:%.*]] = icmp eq i32 [[X]], 1
+; CHECK-NEXT:    br i1 [[CMP]], label %[[COMMON_RET:.*]], label %[[IF_END:.*]]
+; CHECK:       [[COMMON_RET]]:
+; CHECK-NEXT:    ret i32 7
+; CHECK:       [[IF_END]]:
+; CHECK-NEXT:    [[SUB:%.*]] = add nsw i32 [[X]], -1
+; CHECK-NEXT:    [[CALL:%.*]] = tail call i32 @test_neg_intervening_op(i32 [[SUB]])
+; CHECK-NEXT:    [[TMP:%.*]] = add i32 [[CALL]], 1
+; CHECK-NEXT:    [[SHL:%.*]] = shl i32 [[TMP]], 1
+; CHECK-NEXT:    ret i32 [[SHL]]
+;
+entry:
+  %cmp = icmp eq i32 %x, 1
+  br i1 %cmp, label %common.ret, label %if.end
+
+common.ret:
+  %common.ret.op = phi i32 [ %shl, %if.end ], [ 7, %entry ]
+  ret i32 %common.ret.op
+
+if.end:
+  %sub = add nsw i32 %x, -1
+  %call = tail call i32 @test_neg_intervening_op(i32 %sub)
+  %tmp = add i32 %call, 1
+  %shl = shl i32 %tmp, 1
+  br label %common.ret
+}


        


More information about the llvm-commits mailing list