[llvm] [SCEV] Allow replacing disjoint ors when reusing instructions (PR #224246)
Maksim Shelegov via llvm-commits
llvm-commits at lists.llvm.org
Mon Sep 28 03:36:32 PDT 2026
https://github.com/mshelego updated https://github.com/llvm/llvm-project/pull/224246
>From ffd1f0d60eb36f63e287420c9873b172b264db02 Mon Sep 17 00:00:00 2001
From: "Shelegov, Maksim" <maksim.shelegov at intel.com>
Date: Wed, 16 Sep 2026 13:46:57 +0200
Subject: [PATCH] [SCEV] Allow replacing disjoint ors when reusing instructions
SCEV models disjoint ors as adds, but instruction reuse may require
removing poison-generating annotations. Dropping the disjoint flag
is not sufficient: or and add differ when the operands overlap. This
can prevent IndVarSimplify from reusing an existing value and cause
it to expand an exit value into a redundant multiply chain.
Allow canReuseInstruction() to collect disjoint ors for replacement
with adds after the reuse check succeeds. This is valid even when the
operands overlap, because the original disjoint or produces poison
in that case. Keep the existing behavior unless the caller opts in.
The expander leaves replaced ors for the caller to delete after
clearing its caches, as they may still be used as insertion points
during nested expansion. Use WeakTrackingVH for values held across
nested expansions so they follow the replacement.
Enable this only for IndVarSimplify exit-value rewriting, which does
not need to roll back expansions: SCEVExpanderCleaner cannot undo
these replacements. Add unit tests for reuse and replacement during
expansion, and indvars and phase-ordering tests for the redundant
exit-value computation.
---
llvm/include/llvm/Analysis/ScalarEvolution.h | 13 +-
.../Utils/ScalarEvolutionExpander.h | 41 ++-
llvm/lib/Analysis/ScalarEvolution.cpp | 20 +-
llvm/lib/Transforms/Scalar/IndVarSimplify.cpp | 9 +-
.../Utils/ScalarEvolutionExpander.cpp | 96 +++--
.../reuse-disjoint-or-exit-value.ll | 56 +++
.../reuse-disjoint-or-exit-value.ll | 61 ++++
.../Analysis/ScalarEvolutionTest.cpp | 153 ++++++++
.../Utils/ScalarEvolutionExpanderTest.cpp | 344 ++++++++++++++++++
9 files changed, 751 insertions(+), 42 deletions(-)
create mode 100644 llvm/test/Transforms/IndVarSimplify/reuse-disjoint-or-exit-value.ll
create mode 100644 llvm/test/Transforms/PhaseOrdering/reuse-disjoint-or-exit-value.ll
diff --git a/llvm/include/llvm/Analysis/ScalarEvolution.h b/llvm/include/llvm/Analysis/ScalarEvolution.h
index 16739d0a3e5cd..eb44d76f902ab 100644
--- a/llvm/include/llvm/Analysis/ScalarEvolution.h
+++ b/llvm/include/llvm/Analysis/ScalarEvolution.h
@@ -1640,9 +1640,20 @@ class ScalarEvolution {
/// Check whether it is poison-safe to represent the expression S using the
/// instruction I. If such a replacement is performed, the poison flags of
/// instructions in DropPoisonGeneratingInsts must be dropped.
+ ///
+ /// SCEV models disjoint ors as adds, so dropping the disjoint flag alone is
+ /// not sufficient for reuse. If \p ReplaceDisjointOrs is non-null, collect
+ /// these ors for the caller to replace with adds before reusing I; otherwise
+ /// reject them. This replacement is valid because overlapping operands make
+ /// a disjoint or poison. A disjointness proof may depend on annotations that
+ /// reuse drops and cannot be used instead.
+ ///
+ /// Both vectors may contain partial results on failure. The caller must
+ /// discard these without modifying the IR.
LLVM_ABI bool canReuseInstruction(
const SCEV *S, Instruction *I,
- SmallVectorImpl<Instruction *> &DropPoisonGeneratingInsts);
+ SmallVectorImpl<Instruction *> &DropPoisonGeneratingInsts,
+ SmallVectorImpl<BinaryOperator *> *ReplaceDisjointOrs = nullptr);
class FoldID {
SCEVUse Op;
diff --git a/llvm/include/llvm/Transforms/Utils/ScalarEvolutionExpander.h b/llvm/include/llvm/Transforms/Utils/ScalarEvolutionExpander.h
index 48bdd0c5a1e00..38044e462ebd5 100644
--- a/llvm/include/llvm/Transforms/Utils/ScalarEvolutionExpander.h
+++ b/llvm/include/llvm/Transforms/Utils/ScalarEvolutionExpander.h
@@ -74,6 +74,9 @@ class SCEVExpander : public SCEVUseVisitor<SCEVExpander, Value *> {
/// Indicates whether LCSSA phis should be created for inserted values.
bool PreserveLCSSA;
+ /// Optional list for deferred deletion; see setDisjointOrReplacementSink().
+ SmallVectorImpl<WeakTrackingVH> *DisjointOrReplacementSink = nullptr;
+
// InsertedExpressions caches Values for reuse, so must track RAUW.
DenseMap<std::pair<SCEVUse, Instruction *>, TrackingVH<Value>>
InsertedExpressions;
@@ -420,6 +423,28 @@ class SCEVExpander : public SCEVUseVisitor<SCEVExpander, Value *> {
/// that had been serving as the insertion point may have been deleted.
void clearInsertPoint() { Builder.ClearInsertionPoint(); }
+ /// Allow instruction reuse to replace disjoint ors with adds. Disabled when
+ /// \p DeadInsts is null.
+ ///
+ /// Replaced ors are left in the IR and appended to \p DeadInsts because they
+ /// may still be used as insertion points or InsertedExpressions keys.
+ /// The caller must keep the list alive while registered and defer deletion
+ /// until expansion has finished and clear() or destruction has released the
+ /// expander's references. Only delete instructions that are still dead.
+ /// Passing nullptr disables further replacement but does not clear the list
+ /// or allow earlier deletion.
+ ///
+ /// Do not enable this for expansions that may be rolled back:
+ /// SCEVExpanderCleaner cannot undo the replacements.
+ ///
+ /// Values held across expansion must track RAUW, e.g. using WeakTrackingVH.
+ /// SCEVUnknown follows the replacement, so a raw pointer to the old or may no
+ /// longer match its SCEV even though the instruction has not been deleted.
+ void
+ setDisjointOrReplacementSink(SmallVectorImpl<WeakTrackingVH> *DeadInsts) {
+ DisjointOrReplacementSink = DeadInsts;
+ }
+
/// Set location information used by debugging information.
void SetCurrentDebugLocation(DebugLoc L) {
Builder.SetCurrentDebugLocation(std::move(L));
@@ -495,14 +520,18 @@ class SCEVExpander : public SCEVUseVisitor<SCEVExpander, Value *> {
/// Find a previous Value in ExprValueMap for expand.
/// DropPoisonGeneratingInsts is populated with instructions for which
/// poison-generating flags must be dropped if the value is reused.
+ /// If non-null, ReplaceDisjointOrs collects disjoint ors to replace with
+ /// adds. See canReuseInstruction().
Value *FindValueInExprValueMap(
SCEVUse S, const Instruction *InsertPt,
- SmallVectorImpl<Instruction *> &DropPoisonGeneratingInsts);
-
- /// Like FindValueInExprValueMap, but on a successful lookup also drops the
- /// poison-generating flags that reusing the value requires.
- Value *findExistingExpansionAndDropPoisonFlags(SCEVUse S,
- const Instruction *InsertPt);
+ SmallVectorImpl<Instruction *> &DropPoisonGeneratingInsts,
+ SmallVectorImpl<BinaryOperator *> *ReplaceDisjointOrs);
+
+ /// Find an existing value and apply the changes required for reuse: drop
+ /// poison-generating flags and, if enabled, replace disjoint ors with adds.
+ /// Return the replacement if the reused value is itself replaced.
+ Value *findExistingExpansionAndApplyReuseFixups(SCEVUse S,
+ const Instruction *InsertPt);
LLVM_ABI Value *expand(SCEVUse S);
Value *expand(SCEVUse S, BasicBlock::iterator I) {
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index ac705f52cd239..db4670e526cfa 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -4180,7 +4180,8 @@ void ScalarEvolution::getPoisonGeneratingValues(
bool ScalarEvolution::canReuseInstruction(
const SCEV *S, Instruction *I,
- SmallVectorImpl<Instruction *> &DropPoisonGeneratingInsts) {
+ SmallVectorImpl<Instruction *> &DropPoisonGeneratingInsts,
+ SmallVectorImpl<BinaryOperator *> *ReplaceDisjointOrs) {
// If the instruction cannot be poison, it's always safe to reuse.
if (programUndefinedIfPoison(I))
return true;
@@ -4213,12 +4214,16 @@ bool ScalarEvolution::canReuseInstruction(
if (!I)
return false;
- // Disjoint or instructions are interpreted as adds by SCEV. However, we
- // can't replace an arbitrary add with disjoint or, even if we drop the
- // flag. We would need to convert the or into an add.
- if (auto *PDI = dyn_cast<PossiblyDisjointInst>(I))
- if (PDI->isDisjoint())
+ // SCEV models disjoint ors as adds. Dropping the flag is not sufficient,
+ // so reject the or unless the caller can replace it with an add.
+ bool IsDisjointOr = false;
+ if (auto *PDI = dyn_cast<PossiblyDisjointInst>(I);
+ PDI && PDI->isDisjoint()) {
+ if (!ReplaceDisjointOrs)
return false;
+ ReplaceDisjointOrs->push_back(cast<BinaryOperator>(PDI));
+ IsDisjointOr = true;
+ }
// FIXME: Ignore vscale, even though it technically could be poison. Do this
// because SCEV currently assumes it can't be poison. Remove this special
@@ -4231,7 +4236,8 @@ bool ScalarEvolution::canReuseInstruction(
return false;
// If the instruction can't create poison, we can recurse to its operands.
- if (I->hasPoisonGeneratingAnnotations())
+ // Replaced ors do not need their annotations dropped separately.
+ if (!IsDisjointOr && I->hasPoisonGeneratingAnnotations())
DropPoisonGeneratingInsts.push_back(I);
llvm::append_range(Worklist, I->operands());
diff --git a/llvm/lib/Transforms/Scalar/IndVarSimplify.cpp b/llvm/lib/Transforms/Scalar/IndVarSimplify.cpp
index 4b1fb39ce9932..a46da9e4d0784 100644
--- a/llvm/lib/Transforms/Scalar/IndVarSimplify.cpp
+++ b/llvm/lib/Transforms/Scalar/IndVarSimplify.cpp
@@ -2096,8 +2096,13 @@ bool IndVarSimplify::run(Loop *L) {
// loop into any instructions outside of the loop that use the final values
// of the current expressions.
if (ReplaceExitValue != NeverRepl) {
- if (int Rewrites = rewriteLoopExitValues(L, LI, TLI, SE, TTI, Rewriter, DT,
- ReplaceExitValue, DeadInsts)) {
+ // Allow disjoint or replacement here: exit-value expansions are not rolled
+ // back, and DeadInsts is drained after Rewriter.clear().
+ Rewriter.setDisjointOrReplacementSink(&DeadInsts);
+ int Rewrites = rewriteLoopExitValues(L, LI, TLI, SE, TTI, Rewriter, DT,
+ ReplaceExitValue, DeadInsts);
+ Rewriter.setDisjointOrReplacementSink(nullptr);
+ if (Rewrites) {
NumReplaced += Rewrites;
Changed = true;
}
diff --git a/llvm/lib/Transforms/Utils/ScalarEvolutionExpander.cpp b/llvm/lib/Transforms/Utils/ScalarEvolutionExpander.cpp
index 7ce2542551fff..67a09ab45e0f6 100644
--- a/llvm/lib/Transforms/Utils/ScalarEvolutionExpander.cpp
+++ b/llvm/lib/Transforms/Utils/ScalarEvolutionExpander.cpp
@@ -378,11 +378,13 @@ Value *SCEVExpander::InsertBinop(Instruction::BinaryOps Opcode,
/// loop-invariant portions of expressions, after considering what
/// can be folded using target addressing modes.
///
-Value *SCEVExpander::expandAddToGEP(SCEVUse Offset, Value *V,
+Value *SCEVExpander::expandAddToGEP(SCEVUse Offset, Value *Ptr,
SCEV::NoWrapFlags Flags) {
- assert(!isa<Instruction>(V) ||
- SE.DT.dominates(cast<Instruction>(V), &*Builder.GetInsertPoint()));
+ assert(!isa<Instruction>(Ptr) ||
+ SE.DT.dominates(cast<Instruction>(Ptr), &*Builder.GetInsertPoint()));
+ // Nested expansion may replace V; see setDisjointOrReplacementSink().
+ WeakTrackingVH V = Ptr;
Value *Idx = expand(Offset);
GEPNoWrapFlags NW = any(Flags & SCEV::FlagNUW)
? GEPNoWrapFlags::noUnsignedWrap()
@@ -530,7 +532,7 @@ Value *SCEVExpander::visitAddExpr(SCEVUseT<const SCEVAddExpr *> S) {
const SCEV *URemLHS = nullptr;
const SCEV *URemRHS = nullptr;
if (match(S, m_scev_URem(m_SCEV(URemLHS), m_SCEV(URemRHS), SE))) {
- Value *LHS = expand(URemLHS);
+ WeakTrackingVH LHS = expand(URemLHS);
Value *RHS = expand(URemRHS);
return InsertBinop(Instruction::URem, LHS, RHS, SCEV::FlagNone,
/*IsSafeToHoist*/ false);
@@ -562,7 +564,7 @@ Value *SCEVExpander::visitAddExpr(SCEVUseT<const SCEVAddExpr *> S) {
// Emit instructions to add all the operands. Hoist as much as possible
// out of loops, and form meaningful getelementptrs where possible.
- Value *Sum = nullptr;
+ WeakTrackingVH Sum = nullptr;
for (auto I = OpsAndLoops.begin(), E = OpsAndLoops.end(); I != E;) {
const Loop *CurLoop = I->first;
SCEVUse Op = I->second;
@@ -597,10 +599,11 @@ Value *SCEVExpander::visitAddExpr(SCEVUseT<const SCEVAddExpr *> S) {
} else {
// A simple add.
Value *W = expand(Op);
+ Value *L = Sum;
// Canonicalize a constant to the RHS.
- if (isa<Constant>(Sum))
- std::swap(Sum, W);
- Sum = InsertBinop(Instruction::Add, Sum, W, S.getNoWrapFlags(),
+ if (isa<Constant>(L))
+ std::swap(L, W);
+ Sum = InsertBinop(Instruction::Add, L, W, S.getNoWrapFlags(),
/*IsSafeToHoist*/ true);
++I;
}
@@ -639,7 +642,7 @@ Value *SCEVExpander::visitMulExpr(SCEVUseT<const SCEVMulExpr *> S) {
// Emit instructions to mul all the operands. Hoist as much as possible
// out of loops.
- Value *Prod = nullptr;
+ WeakTrackingVH Prod = nullptr;
auto I = OpsAndLoops.begin();
// Expand the calculation of X pow N in the following manner:
@@ -694,8 +697,10 @@ Value *SCEVExpander::visitMulExpr(SCEVUseT<const SCEVMulExpr *> S) {
} else {
// A simple mul.
Value *W = ExpandOpBinPowN();
+ Value *L = Prod;
// Canonicalize a constant to the RHS.
- if (isa<Constant>(Prod)) std::swap(Prod, W);
+ if (isa<Constant>(L))
+ std::swap(L, W);
const APInt *RHS;
if (match(W, m_Power2(RHS))) {
// Canonicalize Prod*(1<<C) to Prod<<C.
@@ -704,11 +709,11 @@ Value *SCEVExpander::visitMulExpr(SCEVUseT<const SCEVMulExpr *> S) {
// clear nsw flag if shl will produce poison value.
if (RHS->logBase2() == RHS->getBitWidth() - 1)
NWFlags = ScalarEvolution::clearFlags(NWFlags, SCEV::FlagNSW);
- Prod = InsertBinop(Instruction::Shl, Prod,
+ Prod = InsertBinop(Instruction::Shl, L,
ConstantInt::get(Ty, RHS->logBase2()), NWFlags,
/*IsSafeToHoist*/ true);
} else {
- Prod = InsertBinop(Instruction::Mul, Prod, W, S.getNoWrapFlags(),
+ Prod = InsertBinop(Instruction::Mul, L, W, S.getNoWrapFlags(),
/*IsSafeToHoist*/ true);
}
}
@@ -718,7 +723,7 @@ Value *SCEVExpander::visitMulExpr(SCEVUseT<const SCEVMulExpr *> S) {
}
Value *SCEVExpander::visitUDivExpr(SCEVUseT<const SCEVUDivExpr *> S) {
- Value *LHS = expand(S->getLHS());
+ WeakTrackingVH LHS = expand(S->getLHS());
if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(S->getRHS())) {
const APInt &RHS = SC->getAPInt();
if (RHS.isPowerOf2())
@@ -1116,7 +1121,7 @@ SCEVExpander::getAddRecExprPHILiterally(const SCEVAddRecExpr *Normalized,
// Expand code for the start value into the loop preheader.
assert(L->getLoopPreheader() &&
"Can't expand add recurrences without a loop preheader!");
- Value *StartV =
+ WeakTrackingVH StartV =
expand(Normalized->getStart(), L->getLoopPreheader()->getTerminator());
// StartV must have been be inserted into L's preheader to dominate the new
@@ -1565,7 +1570,7 @@ Value *SCEVExpander::expandMinMaxExpr(SCEVUseT<const SCEVNAryExpr *> S,
bool IsSequential) {
bool PrevSafeMode = SafeUDivMode;
SafeUDivMode |= IsSequential;
- Value *LHS = expand(S->getOperand(S->getNumOperands() - 1));
+ WeakTrackingVH LHS = expand(S->getOperand(S->getNumOperands() - 1));
Type *Ty = LHS->getType();
if (IsSequential)
LHS = Builder.CreateFreeze(LHS);
@@ -1636,7 +1641,8 @@ Value *SCEVExpander::expandCodeFor(SCEVUse SH, Type *Ty) {
Value *SCEVExpander::FindValueInExprValueMap(
SCEVUse S, const Instruction *InsertPt,
- SmallVectorImpl<Instruction *> &DropPoisonGeneratingInsts) {
+ SmallVectorImpl<Instruction *> &DropPoisonGeneratingInsts,
+ SmallVectorImpl<BinaryOperator *> *ReplaceDisjointOrs) {
// If the expansion is not in CanonicalMode, and the SCEV contains any
// sub scAddRecExpr type SCEV, it is required to expand the SCEV literally.
if (!CanonicalMode && SE.containsAddRecurrence(S))
@@ -1661,19 +1667,52 @@ Value *SCEVExpander::FindValueInExprValueMap(
continue;
// Make sure reusing the instruction is poison-safe.
- if (SE.canReuseInstruction(S, EntInst, DropPoisonGeneratingInsts))
+ if (SE.canReuseInstruction(S, EntInst, DropPoisonGeneratingInsts,
+ ReplaceDisjointOrs))
return V;
+ // Discard partial results before trying the next candidate.
DropPoisonGeneratingInsts.clear();
+ if (ReplaceDisjointOrs)
+ ReplaceDisjointOrs->clear();
}
return nullptr;
}
-Value *SCEVExpander::findExistingExpansionAndDropPoisonFlags(
+Value *SCEVExpander::findExistingExpansionAndApplyReuseFixups(
SCEVUse S, const Instruction *InsertPt) {
SmallVector<Instruction *> DropPoisonGeneratingInsts;
- Value *V = FindValueInExprValueMap(S, InsertPt, DropPoisonGeneratingInsts);
+ SmallVector<BinaryOperator *, 2> ReplaceDisjointOrs;
+ Value *V = FindValueInExprValueMap(
+ S, InsertPt, DropPoisonGeneratingInsts,
+ DisjointOrReplacementSink ? &ReplaceDisjointOrs : nullptr);
if (!V)
return nullptr;
+
+ // Replacing a disjoint or with an add refines poison for overlapping
+ // operands.
+ for (BinaryOperator *Or : ReplaceDisjointOrs) {
+ // Insert before the or to dominate instructions later inserted at that
+ // point.
+ BinaryOperator *Add = BinaryOperator::CreateAdd(
+ Or->getOperand(0), Or->getOperand(1), "", Or->getIterator());
+ Add->setDebugLoc(Or->getDebugLoc());
+ Add->takeName(Or);
+ // Preserve metadata other than poison-generating annotations.
+ Add->copyMetadata(*Or);
+ Add->dropPoisonGeneratingAnnotations();
+ // Mark the add as inserted so hoisting skips it and preserves cache keys
+ // pointing to the or. Mark it reused to prevent cleanup from deleting it.
+ InsertedValues.insert(Add);
+ ReusedValues.insert(Add);
+ // Invalidate cached SCEVs for the or and its users before replacement.
+ SE.forgetValue(Or);
+ Or->replaceAllUsesWith(Add);
+ // Queue after RAUW so the handle tracks the dead or, not the replacement.
+ DisjointOrReplacementSink->emplace_back(Or);
+ if (V == Or)
+ V = Add;
+ }
+
for (Instruction *I : DropPoisonGeneratingInsts) {
rememberFlags(I);
dropPoisonGeneratingAnnotationsAndReinfer(SE, I);
@@ -1748,13 +1787,13 @@ Value *SCEVExpander::expand(SCEVUse S) {
Builder.SetInsertPoint(InsertPt->getParent(), InsertPt);
// Expand the expression into instructions.
- Value *V = findExistingExpansionAndDropPoisonFlags(S, &*InsertPt);
+ Value *V = findExistingExpansionAndApplyReuseFixups(S, &*InsertPt);
BasicBlock::iterator CacheAt = InsertPt;
if (!V && InsertPt != OrigInsertPt && PostIncLoops.empty()) {
// Hoisting the insertion point can move it above a value that already
// computes S. Such a value is still usable: it only has to dominate the
// point we were asked to expand at, which is where the result is used.
- V = findExistingExpansionAndDropPoisonFlags(S, &*OrigInsertPt);
+ V = findExistingExpansionAndApplyReuseFixups(S, &*OrigInsertPt);
if (V)
CacheAt = OrigInsertPt;
}
@@ -2032,8 +2071,13 @@ bool SCEVExpander::hasRelatedExistingExpansion(const SCEV *S,
// ExprValueMap. Note that we don't currently model the cost of
// needing to drop poison generating flags on the instruction if we
// want to reuse it. We effectively assume that has zero cost.
+ // Match expand()'s reuse policy without applying the collected changes.
SmallVector<Instruction *> DropPoisonGeneratingInsts;
- return FindValueInExprValueMap(S, At, DropPoisonGeneratingInsts) != nullptr;
+ SmallVector<BinaryOperator *, 2> ReplaceDisjointOrs;
+ return FindValueInExprValueMap(S, At, DropPoisonGeneratingInsts,
+ DisjointOrReplacementSink
+ ? &ReplaceDisjointOrs
+ : nullptr) != nullptr;
}
template<typename T> static InstructionCost costAndCollectOperands(
@@ -2293,7 +2337,7 @@ Value *SCEVExpander::expandCodeForPredicate(const SCEVPredicate *Pred,
Value *SCEVExpander::expandComparePredicate(const SCEVComparePredicate *Pred,
Instruction *IP) {
- Value *Expr0 = expand(Pred->getLHS(), IP);
+ WeakTrackingVH Expr0 = expand(Pred->getLHS(), IP);
Value *Expr1 = expand(Pred->getRHS(), IP);
Builder.SetInsertPoint(IP);
@@ -2327,13 +2371,13 @@ Value *SCEVExpander::generateOverflowCheck(const SCEVAddRecExpr *AR,
// and |Step| * Backedge doesn't unsigned overflow.
Builder.SetInsertPoint(Loc);
- Value *TripCountVal = expand(ExitCount, Loc);
+ WeakTrackingVH TripCountVal = expand(ExitCount, Loc);
IntegerType *Ty =
IntegerType::get(Loc->getContext(), SE.getTypeSizeInBits(ARTy));
- Value *StepValue = expand(Step, Loc);
- Value *NegStepValue = expand(SE.getNegativeSCEV(Step), Loc);
+ WeakTrackingVH StepValue = expand(Step, Loc);
+ WeakTrackingVH NegStepValue = expand(SE.getNegativeSCEV(Step), Loc);
Value *StartValue = expand(Start, Loc);
ConstantInt *Zero =
diff --git a/llvm/test/Transforms/IndVarSimplify/reuse-disjoint-or-exit-value.ll b/llvm/test/Transforms/IndVarSimplify/reuse-disjoint-or-exit-value.ll
new file mode 100644
index 0000000000000..3b5567d3b75d2
--- /dev/null
+++ b/llvm/test/Transforms/IndVarSimplify/reuse-disjoint-or-exit-value.ll
@@ -0,0 +1,56 @@
+; RUN: opt -passes='loop(indvars)' -S < %s | FileCheck %s
+
+; Reuse the exit value by replacing its disjoint or with an add. Without reuse,
+; SCEVExpander emits a multiply chain:
+;
+; %0 = mul i24 %hi, 257
+; %1 = add i24 %0, 513
+;
+; Reduced from Transforms/PhaseOrdering/reuse-disjoint-or-exit-value.ll at the
+; inner loop's indvars invocation. Keep the outer recurrence: without it,
+; re-expansion only needs a trunc and an add and is cheaper than reuse.
+;
+; Check that indvars reuses %hi.trunc without a multiply. The phase-ordering
+; test covers the subsequent folding and sinking of the reused chain.
+
+define i1 @reuse_disjoint_or_exit_value(i32 %itr) {
+; CHECK-LABEL: define i1 @reuse_disjoint_or_exit_value(
+; CHECK-NOT: mul
+;
+; Replace the or with an add and reuse %hi.trunc for the outer recurrence.
+; CHECK: [[OUTER_LOOPEXIT:.*]]:
+; CHECK: %or = add i24 %hi, 1
+; CHECK: %hi.trunc = trunc i32 %combined to i24
+; CHECK-NOT: mul
+;
+entry:
+ %outer.cmp = icmp eq i32 %itr, 0
+ br label %split
+
+split: ; preds = %entry, %outer.loopexit
+ %hi = phi i24 [ 0, %entry ], [ %hi.next.lcssa, %outer.loopexit ]
+ %or = or disjoint i24 %hi, 1
+ %or.ext = zext i24 %or to i32
+ %step = add nuw nsw i32 %or.ext, 1
+ %step.hi = shl i32 %step, 8
+ %lo.trunc = trunc i32 %step to i8
+ %combined = add i32 %step.hi, %or.ext
+ %hi.trunc = trunc i32 %combined to i24
+ br label %inner.latch
+
+inner.latch: ; preds = %inner.latch, %split
+ %hi.next = phi i24 [ %hi.trunc, %inner.latch ], [ 0, %split ]
+ %lo.next = phi i8 [ %lo.trunc, %inner.latch ], [ 0, %split ]
+ %inner.cmp = phi i1 [ false, %inner.latch ], [ true, %split ]
+ br i1 %inner.cmp, label %inner.latch, label %outer.loopexit
+
+outer.loopexit: ; preds = %inner.latch
+ %hi.next.lcssa = phi i24 [ %hi.next, %inner.latch ]
+ %lo.next.lcssa = phi i8 [ %lo.next, %inner.latch ]
+ br i1 %outer.cmp, label %split, label %exit
+
+exit: ; preds = %outer.loopexit
+ %lo.lcssa = phi i8 [ %lo.next.lcssa, %outer.loopexit ]
+ %res = icmp eq i8 %lo.lcssa, 0
+ ret i1 %res
+}
diff --git a/llvm/test/Transforms/PhaseOrdering/reuse-disjoint-or-exit-value.ll b/llvm/test/Transforms/PhaseOrdering/reuse-disjoint-or-exit-value.ll
new file mode 100644
index 0000000000000..18e419bece76e
--- /dev/null
+++ b/llvm/test/Transforms/PhaseOrdering/reuse-disjoint-or-exit-value.ll
@@ -0,0 +1,61 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 6
+; RUN: opt -passes='default<O1>' -S < %s | FileCheck %s
+
+; Reduced from an OpenCL instruction-latency kernel. SROA splits a loop-carried
+; value into i8 and i24 pieces, reassembled each iteration with or disjoint.
+;
+; Replacing the disjoint or with an add lets IndVarSimplify reuse the existing
+; exit value instead of expanding its closed form into a multiply chain.
+;
+; The full pipeline is needed: indvars alone does not rewrite this input.
+; Earlier passes collapse the inner loops and rotate the outer loop before
+; exit-value expansion becomes expensive.
+
+define i1 @reuse_disjoint_or(i32 %itr) {
+; CHECK-LABEL: define i1 @reuse_disjoint_or(
+; CHECK-SAME: i32 [[ITR:%.*]]) local_unnamed_addr #[[ATTR0:[0-9]+]] {
+; CHECK-NEXT: [[ENTRY:.*]]:
+; CHECK-NEXT: [[OUTER_CMP:%.*]] = icmp eq i32 [[ITR]], 0
+; CHECK-NEXT: br label %[[SPLIT:.*]]
+; CHECK: [[SPLIT]]:
+; CHECK-NEXT: [[HI1:%.*]] = phi i24 [ 0, %[[ENTRY]] ], [ [[OR:%.*]], %[[SPLIT]] ]
+; CHECK-NEXT: [[OR]] = add i24 [[HI1]], 1
+; CHECK-NEXT: br i1 [[OUTER_CMP]], label %[[SPLIT]], label %[[EXIT:.*]]
+; CHECK: [[EXIT]]:
+; CHECK-NEXT: [[TMP0:%.*]] = trunc i24 [[HI1]] to i8
+; CHECK-NEXT: [[RES:%.*]] = icmp eq i8 [[TMP0]], -2
+; CHECK-NEXT: ret i1 [[RES]]
+;
+entry:
+ br label %outer
+
+outer: ; preds = %inner.latch, %entry
+ %hi = phi i24 [ 0, %entry ], [ %hi.next, %inner.latch ]
+ %lo = phi i8 [ 0, %entry ], [ %lo.next, %inner.latch ]
+ %u = phi i32 [ 0, %entry ], [ %itr, %inner.latch ]
+ %outer.cmp = icmp eq i32 %u, 0
+ br i1 %outer.cmp, label %split, label %exit
+
+split: ; preds = %outer
+ %or = or disjoint i24 %hi, 1
+ %or.ext = zext i24 %or to i32
+ %step = add i32 %or.ext, 1
+ %step.hi = shl i32 %step, 8
+ br label %inner.latch
+
+inner.latch: ; preds = %inner.body, %split
+ %hi.next = phi i24 [ %hi.trunc, %inner.body ], [ 0, %split ]
+ %lo.next = phi i8 [ %lo.trunc, %inner.body ], [ 0, %split ]
+ %inner.cmp = phi i1 [ false, %inner.body ], [ true, %split ]
+ br i1 %inner.cmp, label %inner.body, label %outer
+
+inner.body: ; preds = %inner.latch
+ %lo.trunc = trunc i32 %step to i8
+ %combined = add i32 %step.hi, %or.ext
+ %hi.trunc = trunc i32 %combined to i24
+ br label %inner.latch
+
+exit: ; preds = %outer
+ %res = icmp eq i8 %lo, 0
+ ret i1 %res
+}
diff --git a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
index 1fd2eaa5eb72f..4aea64f57be40 100644
--- a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
+++ b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
@@ -24,6 +24,7 @@
#include "llvm/IR/Verifier.h"
#include "llvm/Support/SourceMgr.h"
#include "llvm/Support/raw_ostream.h"
+#include "gmock/gmock.h"
#include "gtest/gtest.h"
namespace llvm {
@@ -2657,4 +2658,156 @@ TEST_F(ScalarEvolutionsTest, AddRecExprUseFlags) {
#endif
});
}
+
+// Check disjoint-or collection without modifying the IR, including on failure.
+class ScalarEvolutionCanReuseInstructionTest : public ScalarEvolutionsTest {
+protected:
+ /// Run canReuseInstruction() on \p Name in \p Assembly, collecting disjoint
+ /// ors if \p AllowDisjointOrs. Verify that the IR is unchanged and pass the
+ /// results to \p Check.
+ void
+ runQuery(StringRef Assembly, StringRef Name, bool AllowDisjointOrs,
+ function_ref<void(Function &F, bool CanReuse,
+ ArrayRef<Instruction *> DropPoisonGeneratingInsts,
+ ArrayRef<BinaryOperator *> ReplaceDisjointOrs)>
+ Check) {
+ LLVMContext C;
+ SMDiagnostic Err;
+ std::unique_ptr<Module> M = parseAssemblyString(Assembly, Err, C);
+ ASSERT_TRUE(M) << "Bad assembly: " << Err.getMessage();
+ ASSERT_FALSE(verifyModule(*M, &errs()));
+
+ std::string Before;
+ raw_string_ostream(Before) << *M;
+
+ runWithSE(*M, "test", [&](Function &F, LoopInfo &LI, ScalarEvolution &SE) {
+ Instruction &I = *getInstructionByName(F, Name);
+ SmallVector<Instruction *> DropPoisonGeneratingInsts;
+ SmallVector<BinaryOperator *, 2> ReplaceDisjointOrs;
+ bool CanReuse = SE.canReuseInstruction(
+ SE.getSCEV(&I), &I, DropPoisonGeneratingInsts,
+ AllowDisjointOrs ? &ReplaceDisjointOrs : nullptr);
+
+ std::string After;
+ raw_string_ostream(After) << *F.getParent();
+ EXPECT_EQ(Before, After) << "the query must not mutate the IR";
+
+ Check(F, CanReuse, DropPoisonGeneratingInsts, ReplaceDisjointOrs);
+ });
+ }
+};
+
+// Reject disjoint ors by default: dropping the flag does not preserve the SCEV.
+TEST_F(ScalarEvolutionCanReuseInstructionTest, DisjointOrRefusedByDefault) {
+ const char *Assembly = R"(
+ define i64 @test(i64 %a, i64 %b) {
+ %or = or disjoint i64 %a, %b
+ ret i64 %or
+ }
+ )";
+ runQuery(Assembly, "or", /*AllowDisjointOrs=*/false,
+ [](Function &, bool CanReuse, ArrayRef<Instruction *> Drop,
+ ArrayRef<BinaryOperator *> Ors) {
+ EXPECT_FALSE(CanReuse);
+ EXPECT_TRUE(Ors.empty());
+ });
+}
+
+// Opting in makes the candidate reusable and reports the or for replacement.
+TEST_F(ScalarEvolutionCanReuseInstructionTest, DisjointOrRootCollected) {
+ const char *Assembly = R"(
+ define i64 @test(i64 %a, i64 %b) {
+ %or = or disjoint i64 %a, %b
+ ret i64 %or
+ }
+ )";
+ runQuery(Assembly, "or", /*AllowDisjointOrs=*/true,
+ [](Function &F, bool CanReuse, ArrayRef<Instruction *> Drop,
+ ArrayRef<BinaryOperator *> Ors) {
+ EXPECT_TRUE(CanReuse);
+ EXPECT_TRUE(Drop.empty());
+ ASSERT_EQ(Ors.size(), 1u);
+ EXPECT_EQ(Ors[0], getInstructionByName(F, "or"));
+ });
+}
+
+// Also collect disjoint ors in the candidate's operand graph.
+TEST_F(ScalarEvolutionCanReuseInstructionTest, DisjointOrOperandCollected) {
+ const char *Assembly = R"(
+ define i64 @test(i64 %a, i64 %b) {
+ %or = or disjoint i64 %a, %b
+ %shl = shl i64 %or, 2
+ ret i64 %shl
+ }
+ )";
+ runQuery(Assembly, "shl", /*AllowDisjointOrs=*/false,
+ [](Function &, bool CanReuse, ArrayRef<Instruction *> Drop,
+ ArrayRef<BinaryOperator *> Ors) { EXPECT_FALSE(CanReuse); });
+ runQuery(Assembly, "shl", /*AllowDisjointOrs=*/true,
+ [](Function &F, bool CanReuse, ArrayRef<Instruction *> Drop,
+ ArrayRef<BinaryOperator *> Ors) {
+ EXPECT_TRUE(CanReuse);
+ ASSERT_EQ(Ors.size(), 1u);
+ EXPECT_EQ(Ors[0], getInstructionByName(F, "or"));
+ });
+}
+
+// Collect all disjoint ors, including dependent ones.
+TEST_F(ScalarEvolutionCanReuseInstructionTest, DependentDisjointOrsCollected) {
+ const char *Assembly = R"(
+ define i64 @test(i64 %a, i64 %b, i64 %c) {
+ %or1 = or disjoint i64 %a, %b
+ %or2 = or disjoint i64 %or1, %c
+ ret i64 %or2
+ }
+ )";
+ runQuery(Assembly, "or2", /*AllowDisjointOrs=*/true,
+ [](Function &F, bool CanReuse, ArrayRef<Instruction *> Drop,
+ ArrayRef<BinaryOperator *> Ors) {
+ EXPECT_TRUE(CanReuse);
+ EXPECT_THAT(Ors, testing::UnorderedElementsAre(
+ getInstructionByName(F, "or1"),
+ getInstructionByName(F, "or2")));
+ });
+}
+
+// Collect annotation drops and replacements together, without also reporting
+// replaced ors for annotation dropping.
+TEST_F(ScalarEvolutionCanReuseInstructionTest, DropAnnotationsAndReplaceOr) {
+ const char *Assembly = R"(
+ define i64 @test(i32 %x, i64 %b) {
+ %z = zext nneg i32 %x to i64
+ %or = or disjoint i64 %z, %b
+ ret i64 %or
+ }
+ )";
+ runQuery(
+ Assembly, "or", /*AllowDisjointOrs=*/true,
+ [](Function &F, bool CanReuse, ArrayRef<Instruction *> Drop,
+ ArrayRef<BinaryOperator *> Ors) {
+ EXPECT_TRUE(CanReuse);
+ EXPECT_THAT(Drop, testing::ElementsAre(getInstructionByName(F, "z")));
+ EXPECT_THAT(Ors, testing::ElementsAre(getInstructionByName(F, "or")));
+ });
+}
+
+// Collect the root or before rejecting an operand. SCEV folds out %a, but it
+// can still make the or poison. Failure must leave the IR unchanged.
+TEST_F(ScalarEvolutionCanReuseInstructionTest, RefusedAfterCollectingOr) {
+ const char *Assembly = R"(
+ define i64 @test(i64 %a, i64 %b) {
+ %zero = sub i64 %a, %a
+ %or = or disjoint i64 %zero, %b
+ ret i64 %or
+ }
+ )";
+ runQuery(Assembly, "or", /*AllowDisjointOrs=*/true,
+ [](Function &F, bool CanReuse, ArrayRef<Instruction *> Drop,
+ ArrayRef<BinaryOperator *> Ors) {
+ EXPECT_FALSE(CanReuse);
+ // Ensure this failure leaves partial results to discard.
+ EXPECT_THAT(Ors,
+ testing::ElementsAre(getInstructionByName(F, "or")));
+ });
+}
} // end namespace llvm
diff --git a/llvm/unittests/Transforms/Utils/ScalarEvolutionExpanderTest.cpp b/llvm/unittests/Transforms/Utils/ScalarEvolutionExpanderTest.cpp
index 76207f7260b9e..851329a3d90b9 100644
--- a/llvm/unittests/Transforms/Utils/ScalarEvolutionExpanderTest.cpp
+++ b/llvm/unittests/Transforms/Utils/ScalarEvolutionExpanderTest.cpp
@@ -8,6 +8,7 @@
#include "llvm/Transforms/Utils/ScalarEvolutionExpander.h"
#include "llvm/ADT/SmallVector.h"
+#include "llvm/ADT/StringMap.h"
#include "llvm/Analysis/AssumptionCache.h"
#include "llvm/Analysis/LoopInfo.h"
#include "llvm/Analysis/ScalarEvolutionExpressions.h"
@@ -23,6 +24,8 @@
#include "llvm/IR/PatternMatch.h"
#include "llvm/IR/Verifier.h"
#include "llvm/Support/SourceMgr.h"
+#include "llvm/Transforms/Utils/Local.h"
+#include "gmock/gmock.h"
#include "gtest/gtest.h"
namespace llvm {
@@ -1035,4 +1038,345 @@ TEST_F(ScalarEvolutionExpanderTest, InsertBinopReuseShlWithMatchingFlags) {
EXPECT_TRUE(ShlInst->hasNoSignedWrap());
}
+// Reuse can replace a disjoint or during nested expansion, including through
+// a phi's incoming values. Check that insertion points remain valid and values
+// held across expansion follow RAUW.
+class SCEVExpanderDisjointOrReplacementTest
+ : public ScalarEvolutionExpanderTest {
+protected:
+ /// Run \p Body on \p Assembly with disjoint-or replacement enabled.
+ ///
+ /// Look up original instructions before expansion: replacements take their
+ /// names. Check that replaced ors remain in the IR until the expander is
+ /// released, then delete them and verify the function again.
+ void
+ runWithExpander(StringRef Assembly,
+ function_ref<void(Function &F, LoopInfo &LI,
+ ScalarEvolution &SE, SCEVExpander &Exp,
+ SmallVectorImpl<WeakTrackingVH> &DeadInsts)>
+ Body) {
+ LLVMContext C;
+ SMDiagnostic Err;
+ std::unique_ptr<Module> M = parseAssemblyString(Assembly, Err, C);
+ ASSERT_TRUE(M) << "Bad assembly: " << Err.getMessage();
+ ASSERT_FALSE(verifyModule(*M, &errs()));
+
+ runWithSE(*M, "test", [&](Function &F, LoopInfo &LI, ScalarEvolution &SE) {
+ SmallVector<WeakTrackingVH, 2> DeadInsts;
+ {
+ SCEVExpander Exp(SE, "test");
+ Exp.setDisjointOrReplacementSink(&DeadInsts);
+ Body(F, LI, SE, Exp, DeadInsts);
+ EXPECT_FALSE(verifyFunction(F, &errs()));
+
+ // Replaced ors must remain in the IR until expansion has finished.
+ for (const WeakTrackingVH &VH : DeadInsts) {
+ auto *Queued = dyn_cast_or_null<Instruction>(VH);
+ ASSERT_TRUE(Queued) << "queued entry was deleted";
+ EXPECT_EQ(Queued->getOpcode(), Instruction::Or)
+ << "queued after the RAUW, so the or and not its replacement";
+ EXPECT_NE(Queued->getParent(), nullptr) << "the or was erased";
+ EXPECT_TRUE(Queued->use_empty()) << "the RAUW left a use behind";
+ }
+
+ // Release the expander's references before deleting replaced ors.
+ Exp.clear();
+ }
+
+ for (WeakTrackingVH &VH : DeadInsts)
+ if (auto *I = dyn_cast_or_null<Instruction>(VH))
+ RecursivelyDeleteTriviallyDeadInstructions(I);
+ EXPECT_FALSE(verifyFunction(F, &errs()));
+ });
+ }
+
+ /// Reusing %index for {1,+,1} visits its disjoint-or increment, %next.
+ ///
+ /// Start at 1 so the first increment is poison (or disjoint 1, 1), while the
+ /// replacement add is 2. This makes stale references observable.
+ ///
+ /// Keep the exit condition independent of the recurrence. Otherwise poison
+ /// would imply UB and canReuseInstruction() would accept %index without
+ /// checking the or.
+ static constexpr const char *RecurrenceAssembly = R"(
+ define void @test(i1 %c) {
+ entry:
+ br label %loop
+ loop:
+ %index = phi i64 [ 1, %entry ], [ %next, %loop ]
+ %next = or disjoint i64 %index, 1
+ br i1 %c, label %loop, label %exit
+ exit:
+ ret void
+ }
+ )";
+};
+
+// Without a sink, reuse is rejected and the disjoint or is unchanged.
+TEST_F(SCEVExpanderDisjointOrReplacementTest, RefusedWithoutSink) {
+ LLVMContext C;
+ SMDiagnostic Err;
+ std::unique_ptr<Module> M = parseAssemblyString(RecurrenceAssembly, Err, C);
+ ASSERT_TRUE(M) << "Bad assembly: " << Err.getMessage();
+
+ runWithSE(*M, "test", [&](Function &F, LoopInfo &LI, ScalarEvolution &SE) {
+ auto &Index = GetInstByName(F, "index");
+ auto &Or = GetInstByName(F, "next");
+ SCEVExpander Exp(SE, "test");
+ Exp.expandCodeFor(
+ SE.getUDivExpr(SE.getSCEV(&Index), SE.getConstant(Index.getType(), 3)),
+ nullptr, &Or);
+ EXPECT_FALSE(verifyFunction(F, &errs()));
+ EXPECT_EQ(Or.getOpcode(), Instruction::Or);
+ EXPECT_TRUE(cast<PossiblyDisjointInst>(&Or)->isDisjoint());
+ });
+}
+
+// Expanding the denominator reuses %index and replaces the or used as the
+// enclosing expansion's insertion point. Its cache entry must remain usable.
+TEST_F(SCEVExpanderDisjointOrReplacementTest, ReplacedCacheLocation) {
+ runWithExpander(
+ RecurrenceAssembly,
+ [](Function &F, LoopInfo &LI, ScalarEvolution &SE, SCEVExpander &Exp,
+ SmallVectorImpl<WeakTrackingVH> &DeadInsts) {
+ auto *Index = &GetInstByName(F, "index");
+ auto *Or = &GetInstByName(F, "next");
+ const SCEV *S = SE.getUDivExpr(SE.getSCEV(Index),
+ SE.getConstant(Index->getType(), 3));
+
+ // Use the or as the insertion point and cache key.
+ Value *V = Exp.expandCodeFor(S, nullptr, Or);
+ ASSERT_EQ(DeadInsts.size(), 1u);
+ EXPECT_EQ(DeadInsts[0], Or);
+
+ // Reuse %index and keep the cache entry keyed on the or.
+ EXPECT_TRUE(match(V, m_UDiv(m_Specific(Index), m_SpecificInt(3))));
+ EXPECT_EQ(Exp.expandCodeFor(S, nullptr, Or), V);
+ EXPECT_EQ(DeadInsts.size(), 1u) << "replaced twice";
+ });
+}
+
+// visitUDivExpr() expands the numerator before the denominator. Expanding the
+// denominator replaces the or used by the numerator's SCEVUnknown. The saved
+// numerator must follow RAUW or it will be poison where its SCEV is 2.
+TEST_F(SCEVExpanderDisjointOrReplacementTest, RetainedOperandFollowsRAUW) {
+ runWithExpander(
+ RecurrenceAssembly,
+ [](Function &F, LoopInfo &LI, ScalarEvolution &SE, SCEVExpander &Exp,
+ SmallVectorImpl<WeakTrackingVH> &DeadInsts) {
+ auto *Index = cast<PHINode>(&GetInstByName(F, "index"));
+ auto *Or = &GetInstByName(F, "next");
+
+ // Check that the numerator's SCEVUnknown follows RAUW.
+ const auto *Numerator = cast<SCEVUnknown>(SE.getUnknown(Or));
+ ASSERT_EQ(Numerator->getValue(), Or);
+
+ Value *V =
+ Exp.expandCodeFor(SE.getUDivExpr(Numerator, SE.getSCEV(Index)),
+ nullptr, Or->getParent()->getTerminator());
+ ASSERT_EQ(DeadInsts.size(), 1u);
+ EXPECT_EQ(DeadInsts[0], Or);
+
+ // Both the increment and its SCEVUnknown must refer to the add.
+ auto *Add = cast<Instruction>(
+ Index->getIncomingValueForBlock(Index->getParent()));
+ ASSERT_EQ(Add->getOpcode(), Instruction::Add);
+ EXPECT_EQ(Numerator->getValue(), Add)
+ << "the expression did not follow the RAUW";
+
+ // The numerator must also refer to the add.
+ auto *UDiv = dyn_cast<Instruction>(V);
+ ASSERT_TRUE(UDiv);
+ ASSERT_EQ(UDiv->getOpcode(), Instruction::UDiv);
+ EXPECT_EQ(UDiv->getOperand(0), Add)
+ << "numerator still names the replaced or";
+ });
+}
+
+// Literal add-recurrence expansion saves the start before expanding the step
+// and building the phi. Expanding the step replaces the or used as the start;
+// the phi must use the replacement add, not the now-stale or.
+TEST_F(SCEVExpanderDisjointOrReplacementTest,
+ LiteralRecurrenceStartFollowsRAUW) {
+ // Omit the phi to force its expansion, and keep the exit condition
+ // independent.
+ const char *Assembly = R"(
+ define void @test(i64 %left, i64 %right, i1 %c) {
+ entry:
+ %start = or disjoint i64 %left, %right
+ %step = shl i64 %start, 2
+ br label %loop
+ loop:
+ br i1 %c, label %loop, label %exit
+ exit:
+ ret void
+ }
+ )";
+ runWithExpander(Assembly, [](Function &F, LoopInfo &LI, ScalarEvolution &SE,
+ SCEVExpander &Exp,
+ SmallVectorImpl<WeakTrackingVH> &DeadInsts) {
+ auto *Or = &GetInstByName(F, "start");
+ auto *Step = &GetInstByName(F, "step");
+ BasicBlock *Preheader = Or->getParent();
+ BasicBlock *Header = Preheader->getSingleSuccessor();
+
+ const auto *Start = cast<SCEVUnknown>(SE.getUnknown(Or));
+ const SCEV *AR = SE.getAddRecExpr(Start, SE.getSCEV(Step),
+ LI.getLoopFor(Header), SCEV::FlagAnyWrap);
+
+ // The literal path is only taken outside canonical mode.
+ Exp.disableCanonicalMode();
+ const SCEV *S = AR;
+ Value *V = Exp.expandCodeFor(S, nullptr, Header->getTerminator());
+
+ ASSERT_EQ(DeadInsts.size(), 1u) << "the step's reuse replaced no or";
+ EXPECT_EQ(DeadInsts[0], Or);
+
+ auto *Add = dyn_cast<Instruction>(Start->getValue());
+ ASSERT_TRUE(Add);
+ EXPECT_EQ(Add->getOpcode(), Instruction::Add)
+ << "the start expression did not follow the RAUW";
+
+ // The phi must start from the add, not the or the start expanded to.
+ auto *PN = dyn_cast<PHINode>(V);
+ ASSERT_TRUE(PN) << "expected a new recurrence phi";
+ EXPECT_EQ(PN->getIncomingValueForBlock(Preheader), Add)
+ << "recurrence initialised from the replaced or";
+
+ // Expanding again reuses the phi rather than building a second one.
+ EXPECT_EQ(Exp.expandCodeFor(S, nullptr, Header->getTerminator()), V);
+ EXPECT_EQ(DeadInsts.size(), 1u) << "replaced twice";
+ });
+}
+
+// Reusing the or itself must return the replacement add.
+TEST_F(SCEVExpanderDisjointOrReplacementTest, ReusedRootIsReplaced) {
+ const char *Assembly = R"(
+ define i64 @test(i64 %a, i64 %b) {
+ entry:
+ %or = or disjoint i64 %a, %b
+ %anchor = mul i64 %a, %b
+ ret i64 %anchor
+ }
+ )";
+ runWithExpander(Assembly, [](Function &F, LoopInfo &LI, ScalarEvolution &SE,
+ SCEVExpander &Exp,
+ SmallVectorImpl<WeakTrackingVH> &DeadInsts) {
+ auto *Or = &GetInstByName(F, "or");
+ Value *V =
+ Exp.expandCodeFor(SE.getSCEV(Or), nullptr, &GetInstByName(F, "anchor"));
+ ASSERT_EQ(DeadInsts.size(), 1u);
+ EXPECT_EQ(DeadInsts[0], Or);
+ EXPECT_NE(V, Or) << "reuse handed back the replaced or";
+ EXPECT_TRUE(
+ match(V, m_Add(m_Specific(F.getArg(0)), m_Specific(F.getArg(1)))));
+ });
+}
+
+// The outer add is created before the inner or is replaced. RAUW must update
+// its operand to the inner add.
+TEST_F(SCEVExpanderDisjointOrReplacementTest, DependentOrsReplaced) {
+ const char *Assembly = R"(
+ define i64 @test(i64 %a, i64 %b, i64 %c) {
+ entry:
+ %or1 = or disjoint i64 %a, %b
+ %or2 = or disjoint i64 %or1, %c
+ %anchor = mul i64 %a, %c
+ ret i64 %anchor
+ }
+ )";
+ runWithExpander(Assembly, [](Function &F, LoopInfo &LI, ScalarEvolution &SE,
+ SCEVExpander &Exp,
+ SmallVectorImpl<WeakTrackingVH> &DeadInsts) {
+ auto *Or1 = &GetInstByName(F, "or1");
+ auto *Or2 = &GetInstByName(F, "or2");
+ Value *V = Exp.expandCodeFor(SE.getSCEV(Or2), nullptr,
+ &GetInstByName(F, "anchor"));
+
+ EXPECT_THAT(DeadInsts, testing::UnorderedElementsAre(WeakTrackingVH(Or1),
+ WeakTrackingVH(Or2)));
+
+ auto *Outer = dyn_cast<Instruction>(V);
+ ASSERT_TRUE(Outer);
+ ASSERT_EQ(Outer->getOpcode(), Instruction::Add);
+ EXPECT_EQ(Outer->getOperand(1), F.getArg(2));
+ EXPECT_TRUE(match(Outer->getOperand(0),
+ m_Add(m_Specific(F.getArg(0)), m_Specific(F.getArg(1)))));
+ });
+}
+
+// Apply annotation drops and replacements together, including dropping nneg
+// from the zext.
+TEST_F(SCEVExpanderDisjointOrReplacementTest, DropsAnnotationsAndReplaces) {
+ const char *Assembly = R"(
+ define i64 @test(i32 %x, i64 %b) {
+ entry:
+ %z = zext nneg i32 %x to i64
+ %or = or disjoint i64 %z, %b
+ %anchor = mul i64 %b, %b
+ ret i64 %anchor
+ }
+ )";
+ runWithExpander(Assembly, [](Function &F, LoopInfo &LI, ScalarEvolution &SE,
+ SCEVExpander &Exp,
+ SmallVectorImpl<WeakTrackingVH> &DeadInsts) {
+ auto *Zext = cast<PossiblyNonNegInst>(&GetInstByName(F, "z"));
+ auto *Or = &GetInstByName(F, "or");
+ ASSERT_TRUE(Zext->hasNonNeg());
+
+ Value *V =
+ Exp.expandCodeFor(SE.getSCEV(Or), nullptr, &GetInstByName(F, "anchor"));
+ ASSERT_EQ(DeadInsts.size(), 1u);
+ EXPECT_EQ(DeadInsts[0], Or);
+ EXPECT_FALSE(Zext->hasNonNeg()) << "nneg was not dropped";
+ EXPECT_TRUE(match(V, m_Add(m_Specific(Zext), m_Specific(F.getArg(1)))));
+ });
+}
+
+// Discard both fixup lists for a rejected candidate before trying the next one.
+//
+// Candidates are insertion-ordered. Prime %bad before %good and ensure none of
+// %bad's operands has the same SCEV (b + c), so %bad is examined first.
+//
+// %bad collects an annotation drop for %z and an or replacement before %x
+// causes rejection: SCEV folds out %x, but it can still make the or poison.
+// %good only uses %b and %c and needs neither fixup.
+TEST_F(SCEVExpanderDisjointOrReplacementTest, RefusedCandidateAppliesNothing) {
+ const char *Assembly = R"(
+ define i64 @test(i32 %x, i64 %b, i64 %c) {
+ entry:
+ %z = zext nneg i32 %x to i64
+ %zero = mul i64 %z, 0
+ %b.plus = add i64 %b, %zero
+ %bad = or disjoint i64 %b.plus, %c
+ %good = add i64 %b, %c
+ %anchor = mul i64 %b, %c
+ ret i64 %anchor
+ }
+ )";
+ runWithExpander(Assembly, [](Function &F, LoopInfo &LI, ScalarEvolution &SE,
+ SCEVExpander &Exp,
+ SmallVectorImpl<WeakTrackingVH> &DeadInsts) {
+ auto *Bad = &GetInstByName(F, "bad");
+ auto *Good = &GetInstByName(F, "good");
+ auto *Zext = cast<PossiblyNonNegInst>(&GetInstByName(F, "z"));
+ ASSERT_TRUE(Zext->hasNonNeg());
+
+ // Prime candidates in lookup order.
+ // Both have to be candidates for the same expression.
+ const SCEV *S = SE.getSCEV(Bad);
+ ASSERT_EQ(S, SE.getSCEV(Good));
+
+ Value *V = Exp.expandCodeFor(S, nullptr, &GetInstByName(F, "anchor"));
+
+ EXPECT_EQ(V, Good) << "expected the second candidate to be reused";
+ EXPECT_TRUE(DeadInsts.empty()) << "a refused candidate was applied";
+ EXPECT_EQ(Bad->getOpcode(), Instruction::Or);
+ EXPECT_TRUE(cast<PossiblyDisjointInst>(Bad)->isDisjoint())
+ << "the refused candidate's or was replaced";
+ EXPECT_TRUE(Zext->hasNonNeg())
+ << "the refused candidate's annotation was dropped";
+ });
+}
+
} // end namespace llvm
More information about the llvm-commits
mailing list