[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