[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