[llvm] [TailRecElim] Handle discarded-call return conflicts with a shift accumulator (PR #221694)

Federico Bruzzone via llvm-commits llvm-commits at lists.llvm.org
Tue Sep 8 03:06:10 PDT 2026


https://github.com/FedericoBruzzone updated https://github.com/llvm/llvm-project/pull/221694

>From 1c4793e0e74726f0e382de0e225a3bc55168101e Mon Sep 17 00:00:00 2001
From: Federico Bruzzone <federico.bruzzone.i at gmail.com>
Date: Mon, 7 Sep 2026 12:02:23 +0200
Subject: [PATCH 1/2] [TailRecElim] Handle discarded-call return conflicts with
 a shift accumulator

Signed-off-by: Federico Bruzzone <federico.bruzzone.i at gmail.com>
---
 .../Scalar/TailRecursionElimination.cpp       | 38 ++++++++++++++-----
 .../TailCallElim/shl-accumulator-opt.ll       | 36 ++++++++++++++++++
 2 files changed, 64 insertions(+), 10 deletions(-)

diff --git a/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp b/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
index b3df27433a9f5..99aafa3390eeb 100644
--- a/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
+++ b/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
@@ -50,6 +50,7 @@
 //===----------------------------------------------------------------------===//
 
 #include "llvm/Transforms/Scalar/TailRecursionElimination.h"
+#include "llvm/ADT/ArrayRef.h"
 #include "llvm/ADT/STLExtras.h"
 #include "llvm/ADT/SmallPtrSet.h"
 #include "llvm/ADT/Statistic.h"
@@ -452,6 +453,11 @@ static bool isUnaryAccumulatorRecurrence(Instruction *I) {
 // will be rewritten to return the accumulator, so all of them have to yield the
 // same base-case constant. Return that constant, or nullptr on failure.
 //
+// RetSelects are the selects already inserted for call sites eliminated via
+// the "found return value" mechanism instead of the accumulator one. Their
+// original `ret` is gone, so they'd otherwise be invisible to the scan above,
+// but they still have to agree on the same base-case constant.
+//
 // 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:
@@ -462,10 +468,18 @@ static bool isUnaryAccumulatorRecurrence(Instruction *I) {
 //  return f(x-1) << 1;
 // }
 // ```
-static Constant *findBaseCaseRetConstant(Function &F,
-                                         Instruction *AccRecInstr) {
+static Constant *findBaseCaseRetConstant(Function &F, Instruction *AccRecInstr,
+                                         ArrayRef<SelectInst *> RetSelects) {
   Constant *BaseCaseVal = nullptr;
 
+  // Records C as the base-case constant the first time it's seen, and
+  // otherwise checks that it agrees with the one already on record.
+  auto AgreesWithBaseCase = [&](Constant *C) {
+    if (!BaseCaseVal)
+      BaseCaseVal = C;
+    return BaseCaseVal == C;
+  };
+
   for (BasicBlock &BB : F) {
     auto *RI = dyn_cast<ReturnInst>(BB.getTerminator());
     if (!RI || !RI->getReturnValue())
@@ -483,12 +497,13 @@ static Constant *findBaseCaseRetConstant(Function &F,
     // not eliminated) must be rejected: returning the accumulator in its place
     // would drop that computation.
     auto *C = dyn_cast<Constant>(RV);
-    if (!C)
+    if (!C || !AgreesWithBaseCase(C))
       return nullptr;
+  }
 
-    if (!BaseCaseVal)
-      BaseCaseVal = C;
-    else if (BaseCaseVal != C)
+  for (SelectInst *SI : RetSelects) {
+    auto *C = dyn_cast<Constant>(SI->getFalseValue());
+    if (!C || !AgreesWithBaseCase(C))
       return nullptr;
   }
 
@@ -498,8 +513,9 @@ static Constant *findBaseCaseRetConstant(Function &F,
 // 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) {
+static Constant *
+canTransformAccumulatorRecursion(Instruction *I, CallInst *CI,
+                                 ArrayRef<SelectInst *> RetSelects) {
   bool IsUnaryAccumulatorRecurrence = isUnaryAccumulatorRecurrence(I);
   if ((!I->isAssociative() || !I->isCommutative()) &&
       !IsUnaryAccumulatorRecurrence)
@@ -517,7 +533,8 @@ static Constant *canTransformAccumulatorRecursion(Instruction *I,
 
     // 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(), /*AccRecInstr=*/I);
+    AccInitVal = findBaseCaseRetConstant(*CI->getFunction(), /*AccRecInstr=*/I,
+                                         RetSelects);
     if (!AccInitVal)
       return nullptr;
   } else {
@@ -815,7 +832,8 @@ bool TailRecursionEliminator::eliminateCall(CallInst *CI) {
     // 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);
+    Constant *AccInitVal =
+        canTransformAccumulatorRecursion(&*BBI, CI, RetSelects);
 
     if (AccPN || !AccInitVal)
       return false; // We cannot eliminate the tail recursion!
diff --git a/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll b/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
index d9d0308dc163f..9104d3d93e02c 100644
--- a/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
+++ b/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
@@ -239,3 +239,39 @@ other:
   %res = sub i32 %call2, 3
   ret i32 %res
 }
+
+; Negative test: a sibling call site that discards its own recursive call and
+; unconditionally returns a fixed constant must still agree with the shift
+; accumulator's base case.
+; int f(int x) {
+;   if (x == 0) return 1;
+;   if (x == 3) { f(x - 1); return 2; }
+;   return f(x - 1) << 1;
+; }
+define i32 @test_neg_discarded_call_conflicting_base(i32 %x) {
+; CHECK-LABEL: define i32 @test_neg_discarded_call_conflicting_base(
+; CHECK-NOT: accumulator.tr
+; CHECK: select {{.*}}, i32 2
+;
+entry:
+  %isbase = icmp eq i32 %x, 0
+  br i1 %isbase, label %base, label %rec
+
+rec:
+  %issp = icmp eq i32 %x, 3
+  br i1 %issp, label %special, label %normal
+
+special:
+  %d = sub i32 %x, 1
+  %c1 = tail call i32 @test_neg_discarded_call_conflicting_base(i32 %d)
+  ret i32 2
+
+normal:
+  %dec = sub i32 %x, 1
+  %c = tail call i32 @test_neg_discarded_call_conflicting_base(i32 %dec)
+  %shl = shl i32 %c, 1
+  ret i32 %shl
+
+base:
+  ret i32 1
+}

>From 80b81e28999270adfe0c25330621ceab376d9e0f Mon Sep 17 00:00:00 2001
From: Federico Bruzzone <federico.bruzzone.i at gmail.com>
Date: Tue, 8 Sep 2026 12:05:40 +0200
Subject: [PATCH 2/2] Address comments from Antonio

Signed-off-by: Federico Bruzzone <federico.bruzzone.i at gmail.com>
---
 .../Scalar/TailRecursionElimination.cpp       | 195 +++++++++---------
 .../TailCallElim/shl-accumulator-opt.ll       |  31 ++-
 2 files changed, 125 insertions(+), 101 deletions(-)

diff --git a/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp b/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
index 99aafa3390eeb..6ba1b5d897369 100644
--- a/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
+++ b/llvm/lib/Transforms/Scalar/TailRecursionElimination.cpp
@@ -50,7 +50,6 @@
 //===----------------------------------------------------------------------===//
 
 #include "llvm/Transforms/Scalar/TailRecursionElimination.h"
-#include "llvm/ADT/ArrayRef.h"
 #include "llvm/ADT/STLExtras.h"
 #include "llvm/ADT/SmallPtrSet.h"
 #include "llvm/ADT/Statistic.h"
@@ -447,7 +446,94 @@ static bool isUnaryAccumulatorRecurrence(Instruction *I) {
   return isa<ConstantInt>(I->getOperand(1));
 }
 
-// Find the base-case return value for function F, given the accumulator
+namespace {
+class TailRecursionEliminator {
+  Function &F;
+  const TargetTransformInfo *TTI;
+  AliasAnalysis *AA;
+  OptimizationRemarkEmitter *ORE;
+  DomTreeUpdater &DTU;
+  BlockFrequencyInfo *const BFI;
+  ProfileSummaryInfo *const PSI;
+  const bool UpdateFunctionEntryCount;
+  const uint64_t OrigEntryBBFreq;
+  const uint64_t OrigEntryCount;
+
+  // The below are shared state we want to have available when eliminating any
+  // calls in the function. There values should be populated by
+  // createTailRecurseLoopHeader the first time we find a call we can eliminate.
+  BasicBlock *HeaderBB = nullptr;
+  SmallVector<PHINode *, 8> ArgumentPHIs;
+
+  // PHI node to store our return value.
+  PHINode *RetPN = nullptr;
+
+  // i1 PHI node to track if we have a valid return value stored in RetPN.
+  PHINode *RetKnownPN = nullptr;
+
+  // Vector of select instructions we insereted. These selects use RetKnownPN
+  // to either propagate RetPN or select a new return value.
+  SmallVector<SelectInst *, 8> RetSelects;
+
+  // The below are shared state needed when performing accumulator recursion.
+  // There values should be populated by insertAccumulator the first time we
+  // find an elimination that requires an accumulator.
+
+  // PHI node to store our current accumulated value.
+  PHINode *AccPN = nullptr;
+
+  // 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,
+                          ProfileSummaryInfo *PSI,
+                          bool UpdateFunctionEntryCount)
+      : F(F), TTI(TTI), AA(AA), ORE(ORE), DTU(DTU), BFI(BFI), PSI(PSI),
+        UpdateFunctionEntryCount(UpdateFunctionEntryCount),
+        OrigEntryBBFreq(
+            BFI ? BFI->getBlockFreq(&F.getEntryBlock()).getFrequency() : 0U),
+        OrigEntryCount(F.getEntryCount() ? *F.getEntryCount() : 0) {
+    if (BFI) {
+      // The assert is meant as API documentation for the caller.
+      assert(OrigEntryBBFreq != 0 &&
+             "If a BFI was provided, the function should have an entry "
+             "basic block with a non-zero frequency.");
+    }
+  }
+
+  Constant *findBaseCaseRetConstant(Instruction *AccRecInstr);
+
+  Constant *canTransformAccumulatorRecursion(Instruction *I, CallInst *CI);
+
+  CallInst *findTRECandidate(BasicBlock *BB);
+
+  void createTailRecurseLoopHeader(CallInst *CI);
+
+  void insertAccumulator(Instruction *AccRecInstr);
+
+  bool eliminateCall(CallInst *CI);
+
+  void cleanupAndFinalize();
+
+  bool processBlock(BasicBlock &BB);
+
+  void copyByValueOperandIntoLocalTemp(CallInst *CI, int OpndIdx);
+
+  void copyLocalTempOfByValueOperandIntoArguments(CallInst *CI, int OpndIdx);
+
+public:
+  static bool eliminate(Function &F, const TargetTransformInfo *TTI,
+                        AliasAnalysis *AA, OptimizationRemarkEmitter *ORE,
+                        DomTreeUpdater &DTU, BlockFrequencyInfo *BFI,
+                        ProfileSummaryInfo *PSI, bool UpdateFunctionEntryCount);
+};
+} // namespace
+
+// Find the base-case return value for the function, given the accumulator
 // recursion instruction AccRecInstr that is about to be eliminated. Every
 // return other than the one fed by AccRecInstr survives the transformation and
 // will be rewritten to return the accumulator, so all of them have to yield the
@@ -455,7 +541,7 @@ static bool isUnaryAccumulatorRecurrence(Instruction *I) {
 //
 // RetSelects are the selects already inserted for call sites eliminated via
 // the "found return value" mechanism instead of the accumulator one. Their
-// original `ret` is gone, so they'd otherwise be invisible to the scan above,
+// original `ret` is gone, so they'd otherwise be invisible to the scan below,
 // but they still have to agree on the same base-case constant.
 //
 // FIXME: There is a room for improvement here in the future, e.g., consider
@@ -468,13 +554,13 @@ static bool isUnaryAccumulatorRecurrence(Instruction *I) {
 //  return f(x-1) << 1;
 // }
 // ```
-static Constant *findBaseCaseRetConstant(Function &F, Instruction *AccRecInstr,
-                                         ArrayRef<SelectInst *> RetSelects) {
+Constant *
+TailRecursionEliminator::findBaseCaseRetConstant(Instruction *AccRecInstr) {
   Constant *BaseCaseVal = nullptr;
 
   // Records C as the base-case constant the first time it's seen, and
   // otherwise checks that it agrees with the one already on record.
-  auto AgreesWithBaseCase = [&](Constant *C) {
+  auto SetOrMatchBaseCase = [&](Constant *C) {
     if (!BaseCaseVal)
       BaseCaseVal = C;
     return BaseCaseVal == C;
@@ -497,13 +583,13 @@ static Constant *findBaseCaseRetConstant(Function &F, Instruction *AccRecInstr,
     // not eliminated) must be rejected: returning the accumulator in its place
     // would drop that computation.
     auto *C = dyn_cast<Constant>(RV);
-    if (!C || !AgreesWithBaseCase(C))
+    if (!C || !SetOrMatchBaseCase(C))
       return nullptr;
   }
 
   for (SelectInst *SI : RetSelects) {
     auto *C = dyn_cast<Constant>(SI->getFalseValue());
-    if (!C || !AgreesWithBaseCase(C))
+    if (!C || !SetOrMatchBaseCase(C))
       return nullptr;
   }
 
@@ -513,9 +599,8 @@ static Constant *findBaseCaseRetConstant(Function &F, Instruction *AccRecInstr,
 // 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,
-                                 ArrayRef<SelectInst *> RetSelects) {
+Constant *TailRecursionEliminator::canTransformAccumulatorRecursion(
+    Instruction *I, CallInst *CI) {
   bool IsUnaryAccumulatorRecurrence = isUnaryAccumulatorRecurrence(I);
   if ((!I->isAssociative() || !I->isCommutative()) &&
       !IsUnaryAccumulatorRecurrence)
@@ -533,8 +618,7 @@ canTransformAccumulatorRecursion(Instruction *I, CallInst *CI,
 
     // 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(), /*AccRecInstr=*/I,
-                                         RetSelects);
+    AccInitVal = findBaseCaseRetConstant(/*AccRecInstr=*/I);
     if (!AccInitVal)
       return nullptr;
   } else {
@@ -555,89 +639,6 @@ canTransformAccumulatorRecursion(Instruction *I, CallInst *CI,
   return AccInitVal;
 }
 
-namespace {
-class TailRecursionEliminator {
-  Function &F;
-  const TargetTransformInfo *TTI;
-  AliasAnalysis *AA;
-  OptimizationRemarkEmitter *ORE;
-  DomTreeUpdater &DTU;
-  BlockFrequencyInfo *const BFI;
-  ProfileSummaryInfo *const PSI;
-  const bool UpdateFunctionEntryCount;
-  const uint64_t OrigEntryBBFreq;
-  const uint64_t OrigEntryCount;
-
-  // The below are shared state we want to have available when eliminating any
-  // calls in the function. There values should be populated by
-  // createTailRecurseLoopHeader the first time we find a call we can eliminate.
-  BasicBlock *HeaderBB = nullptr;
-  SmallVector<PHINode *, 8> ArgumentPHIs;
-
-  // PHI node to store our return value.
-  PHINode *RetPN = nullptr;
-
-  // i1 PHI node to track if we have a valid return value stored in RetPN.
-  PHINode *RetKnownPN = nullptr;
-
-  // Vector of select instructions we insereted. These selects use RetKnownPN
-  // to either propagate RetPN or select a new return value.
-  SmallVector<SelectInst *, 8> RetSelects;
-
-  // The below are shared state needed when performing accumulator recursion.
-  // There values should be populated by insertAccumulator the first time we
-  // find an elimination that requires an accumulator.
-
-  // PHI node to store our current accumulated value.
-  PHINode *AccPN = nullptr;
-
-  // 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,
-                          ProfileSummaryInfo *PSI,
-                          bool UpdateFunctionEntryCount)
-      : F(F), TTI(TTI), AA(AA), ORE(ORE), DTU(DTU), BFI(BFI), PSI(PSI),
-        UpdateFunctionEntryCount(UpdateFunctionEntryCount),
-        OrigEntryBBFreq(
-            BFI ? BFI->getBlockFreq(&F.getEntryBlock()).getFrequency() : 0U),
-        OrigEntryCount(F.getEntryCount() ? *F.getEntryCount() : 0) {
-    if (BFI) {
-      // The assert is meant as API documentation for the caller.
-      assert(OrigEntryBBFreq != 0 &&
-             "If a BFI was provided, the function should have an entry "
-             "basic block with a non-zero frequency.");
-    }
-  }
-
-  CallInst *findTRECandidate(BasicBlock *BB);
-
-  void createTailRecurseLoopHeader(CallInst *CI);
-
-  void insertAccumulator(Instruction *AccRecInstr);
-
-  bool eliminateCall(CallInst *CI);
-
-  void cleanupAndFinalize();
-
-  bool processBlock(BasicBlock &BB);
-
-  void copyByValueOperandIntoLocalTemp(CallInst *CI, int OpndIdx);
-
-  void copyLocalTempOfByValueOperandIntoArguments(CallInst *CI, int OpndIdx);
-
-public:
-  static bool eliminate(Function &F, const TargetTransformInfo *TTI,
-                        AliasAnalysis *AA, OptimizationRemarkEmitter *ORE,
-                        DomTreeUpdater &DTU, BlockFrequencyInfo *BFI,
-                        ProfileSummaryInfo *PSI, bool UpdateFunctionEntryCount);
-};
-} // namespace
-
 CallInst *TailRecursionEliminator::findTRECandidate(BasicBlock *BB) {
   Instruction *TI = BB->getTerminator();
 
@@ -833,7 +834,7 @@ bool TailRecursionEliminator::eliminateCall(CallInst *CI) {
     // recursion elimination. Check to see if this is the case, and if so,
     // remember which instruction accumulates for later.
     Constant *AccInitVal =
-        canTransformAccumulatorRecursion(&*BBI, CI, RetSelects);
+        canTransformAccumulatorRecursion(&*BBI, CI);
 
     if (AccPN || !AccInitVal)
       return false; // We cannot eliminate the tail recursion!
diff --git a/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll b/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
index 9104d3d93e02c..ca71554a2c04f 100644
--- a/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
+++ b/llvm/test/Transforms/TailCallElim/shl-accumulator-opt.ll
@@ -250,8 +250,31 @@ other:
 ; }
 define i32 @test_neg_discarded_call_conflicting_base(i32 %x) {
 ; CHECK-LABEL: define i32 @test_neg_discarded_call_conflicting_base(
-; CHECK-NOT: accumulator.tr
-; CHECK: select {{.*}}, i32 2
+; CHECK-SAME: i32 [[X:%.*]]) {
+; CHECK-NEXT:  [[ENTRY:.*]]:
+; CHECK-NEXT:    br label %[[TAILRECURSE:.*]]
+; CHECK:       [[TAILRECURSE]]:
+; CHECK-NEXT:    [[X_TR:%.*]] = phi i32 [ [[X]], %[[ENTRY]] ], [ [[D:%.*]], %[[DISCARD:.*]] ]
+; CHECK-NEXT:    [[RET_TR:%.*]] = phi i32 [ poison, %[[ENTRY]] ], [ [[CURRENT_RET_TR:%.*]], %[[DISCARD]] ]
+; CHECK-NEXT:    [[RET_KNOWN_TR:%.*]] = phi i1 [ false, %[[ENTRY]] ], [ true, %[[DISCARD]] ]
+; CHECK-NEXT:    [[ISBASE:%.*]] = icmp eq i32 [[X_TR]], 0
+; CHECK-NEXT:    br i1 [[ISBASE]], label %[[BASE:.*]], label %[[REC:.*]]
+; CHECK:       [[REC]]:
+; CHECK-NEXT:    [[ISSP:%.*]] = icmp eq i32 [[X_TR]], 3
+; CHECK-NEXT:    br i1 [[ISSP]], label %[[DISCARD]], label %[[NORMAL:.*]]
+; CHECK:       [[DISCARD]]:
+; CHECK-NEXT:    [[D]] = sub i32 [[X_TR]], 1
+; CHECK-NEXT:    [[CURRENT_RET_TR]] = select i1 [[RET_KNOWN_TR]], i32 [[RET_TR]], i32 2
+; CHECK-NEXT:    br label %[[TAILRECURSE]]
+; CHECK:       [[NORMAL]]:
+; CHECK-NEXT:    [[DEC:%.*]] = sub i32 [[X_TR]], 1
+; CHECK-NEXT:    [[C:%.*]] = tail call i32 @test_neg_discarded_call_conflicting_base(i32 [[DEC]])
+; CHECK-NEXT:    [[SHL:%.*]] = shl i32 [[C]], 1
+; CHECK-NEXT:    [[CURRENT_RET_TR1:%.*]] = select i1 [[RET_KNOWN_TR]], i32 [[RET_TR]], i32 [[SHL]]
+; CHECK-NEXT:    ret i32 [[CURRENT_RET_TR1]]
+; CHECK:       [[BASE]]:
+; CHECK-NEXT:    [[CURRENT_RET_TR2:%.*]] = select i1 [[RET_KNOWN_TR]], i32 [[RET_TR]], i32 1
+; CHECK-NEXT:    ret i32 [[CURRENT_RET_TR2]]
 ;
 entry:
   %isbase = icmp eq i32 %x, 0
@@ -259,9 +282,9 @@ entry:
 
 rec:
   %issp = icmp eq i32 %x, 3
-  br i1 %issp, label %special, label %normal
+  br i1 %issp, label %discard, label %normal
 
-special:
+discard:
   %d = sub i32 %x, 1
   %c1 = tail call i32 @test_neg_discarded_call_conflicting_base(i32 %d)
   ret i32 2



More information about the llvm-commits mailing list