[llvm] [SCEV] Return a SCEVUse from getAddRecExpr and propagate use flags. (PR #221606)

Florian Hahn via llvm-commits llvm-commits at lists.llvm.org
Mon Sep 7 01:12:02 PDT 2026


https://github.com/fhahn updated https://github.com/llvm/llvm-project/pull/221606

>From 458ced930911144e993b44c5c9fcd17872eb6b3b Mon Sep 17 00:00:00 2001
From: Florian Hahn <flo at fhahn.com>
Date: Mon, 31 Aug 2026 10:30:19 +0100
Subject: [PATCH] [SCEV] Return a SCEVUse from getAddRecExpr and propagate use
 flags.

---
 llvm/include/llvm/Analysis/ScalarEvolution.h  |  17 +--
 llvm/lib/Analysis/ScalarEvolution.cpp         |  37 +++---
 llvm/lib/Transforms/Vectorize/VPlanUtils.cpp  |   3 +-
 .../Analysis/ScalarEvolutionTest.cpp          | 106 ++++++++++++++++++
 4 files changed, 142 insertions(+), 21 deletions(-)

diff --git a/llvm/include/llvm/Analysis/ScalarEvolution.h b/llvm/include/llvm/Analysis/ScalarEvolution.h
index 86bdc85147a20..203e74da71147 100644
--- a/llvm/include/llvm/Analysis/ScalarEvolution.h
+++ b/llvm/include/llvm/Analysis/ScalarEvolution.h
@@ -783,14 +783,17 @@ class ScalarEvolution {
   LLVM_ABI const SCEV *getUDivExpr(SCEVUse LHS, SCEVUse RHS);
   LLVM_ABI const SCEV *getUDivExactExpr(SCEVUse LHS, SCEVUse RHS);
   LLVM_ABI const SCEV *getURemExpr(SCEVUse LHS, SCEVUse RHS);
-  LLVM_ABI const SCEV *getAddRecExpr(SCEVUse Start, SCEVUse Step, const Loop *L,
-                                     SCEV::NoWrapFlags Flags);
-  LLVM_ABI const SCEV *getAddRecExpr(SmallVectorImpl<SCEVUse> &Operands,
-                                     const Loop *L, SCEV::NoWrapFlags Flags);
-  const SCEV *getAddRecExpr(const SmallVectorImpl<SCEVUse> &Operands,
-                            const Loop *L, SCEV::NoWrapFlags Flags) {
+  LLVM_ABI SCEVUse getAddRecExpr(
+      SCEVUse Start, SCEVUse Step, const Loop *L, SCEV::NoWrapFlags Flags,
+      SCEV::NoWrapFlags UseFlags = SCEV::FlagAnyWrap);
+  LLVM_ABI SCEVUse getAddRecExpr(
+      SmallVectorImpl<SCEVUse> &Operands, const Loop *L,
+      SCEV::NoWrapFlags Flags, SCEV::NoWrapFlags UseFlags = SCEV::FlagAnyWrap);
+  SCEVUse getAddRecExpr(const SmallVectorImpl<SCEVUse> &Operands, const Loop *L,
+                        SCEV::NoWrapFlags Flags,
+                        SCEV::NoWrapFlags UseFlags = SCEV::FlagAnyWrap) {
     SmallVector<SCEVUse, 4> NewOp(Operands.begin(), Operands.end());
-    return getAddRecExpr(NewOp, L, Flags);
+    return getAddRecExpr(NewOp, L, Flags, UseFlags);
   }
 
   /// Checks if \p SymbolicPHI can be rewritten as an AddRecExpr under some
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index d3b02b34caf27..d6770e3999c00 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -3684,26 +3684,30 @@ const SCEV *ScalarEvolution::getUDivExactExpr(SCEVUse LHS, SCEVUse RHS) {
 
 /// Get an add recurrence expression for the specified loop.  Simplify the
 /// expression as much as possible.
-const SCEV *ScalarEvolution::getAddRecExpr(SCEVUse Start, SCEVUse Step,
-                                           const Loop *L,
-                                           SCEV::NoWrapFlags Flags) {
+SCEVUse ScalarEvolution::getAddRecExpr(SCEVUse Start, SCEVUse Step,
+                                       const Loop *L, SCEV::NoWrapFlags Flags,
+                                       SCEV::NoWrapFlags UseFlags) {
   SmallVector<SCEVUse, 4> Operands;
   Operands.push_back(Start);
   if (const SCEVAddRecExpr *StepChrec = dyn_cast<SCEVAddRecExpr>(Step))
     if (StepChrec->getLoop() == L) {
       append_range(Operands, StepChrec->operands());
+      // The flags describe the two-operand recurrence, not the flattened one
+      // built here, so drop them just like the shared ones.
       return getAddRecExpr(Operands, L, maskFlags(Flags, SCEV::FlagNW));
     }
 
   Operands.push_back(Step);
-  return getAddRecExpr(Operands, L, Flags);
+  return getAddRecExpr(Operands, L, Flags, UseFlags);
 }
 
 /// Get an add recurrence expression for the specified loop.  Simplify the
 /// expression as much as possible.
-const SCEV *ScalarEvolution::getAddRecExpr(SmallVectorImpl<SCEVUse> &Operands,
-                                           const Loop *L,
-                                           SCEV::NoWrapFlags Flags) {
+SCEVUse ScalarEvolution::getAddRecExpr(SmallVectorImpl<SCEVUse> &Operands,
+                                       const Loop *L, SCEV::NoWrapFlags Flags,
+                                       SCEV::NoWrapFlags UseFlags) {
+  assert(UseFlags == maskFlags(UseFlags, SCEV::FlagNUW | SCEV::FlagNSW) &&
+         "only nuw or nsw allowed");
   if (Operands.size() == 1) return Operands[0];
 #ifndef NDEBUG
   Type *ETy = getEffectiveSCEVType(Operands[0]->getType());
@@ -3716,6 +3720,10 @@ const SCEV *ScalarEvolution::getAddRecExpr(SmallVectorImpl<SCEVUse> &Operands,
     assert(isAvailableAtLoopEntry(Op, L) &&
            "SCEVAddRecExpr operand is not available at loop entry!");
 #endif
+  // Keep track of original operands, if use-specific flags have been provided.
+  SmallVector<SCEVUse, 4> OrigOperands;
+  if (UseFlags != SCEV::FlagAnyWrap)
+    OrigOperands.assign(Operands.begin(), Operands.end());
 
   if (Operands.back()->isZero()) {
     Operands.pop_back();
@@ -3738,6 +3746,7 @@ const SCEV *ScalarEvolution::getAddRecExpr(SmallVectorImpl<SCEVUse> &Operands,
             : (!NestedLoop->contains(L) &&
                DT.dominates(L->getHeader(), NestedLoop->getHeader()))) {
       SmallVector<SCEVUse, 4> NestedOperands(NestedAR->operands());
+      SCEVUse OrigStart = Operands[0];
       Operands[0] = NestedAR->getStart();
       // AddRecs require their operands be loop-invariant with respect to their
       // loops. Don't perform this transformation if it would break this
@@ -3769,13 +3778,15 @@ const SCEV *ScalarEvolution::getAddRecExpr(SmallVectorImpl<SCEVUse> &Operands,
         }
       }
       // Reset Operands to its original state.
-      Operands[0] = NestedAR;
+      Operands[0] = OrigStart;
     }
   }
 
   // Okay, it looks like we really DO need an addrec expr.  Check to see if we
   // already have one, otherwise create a new one.
-  return getOrCreateAddRecExpr(Operands, L, Flags);
+  assert((UseFlags == SCEV::FlagAnyWrap || equal(OrigOperands, Operands)) &&
+         "Tried to add SCEVUse flags after operands changed");
+  return {getOrCreateAddRecExpr(Operands, L, Flags), UseFlags};
 }
 
 const SCEV *ScalarEvolution::getGEPExpr(GEPOperator *GEP,
@@ -5622,7 +5633,7 @@ ScalarEvolution::createAddRecFromPHIWithCastsImpl(const SCEVUnknown *SymbolicPHI
   // which the casts had been folded away. The caller can rewrite SymbolicPHI
   // into NewAR if it will also add the runtime overflow checks specified in
   // Predicates.
-  auto *NewAR = getAddRecExpr(StartVal, Accum, L, SCEV::FlagAnyWrap);
+  const SCEV *NewAR = getAddRecExpr(StartVal, Accum, L, SCEV::FlagAnyWrap);
 
   std::pair<const SCEV *, SmallVector<const SCEVPredicate *, 3>> PredRewrite =
       std::make_pair(NewAR, Predicates);
@@ -13372,9 +13383,9 @@ ScalarEvolution::howManyLessThans(const SCEV *LHS, const SCEV *RHS,
           // if we'd been able to infer the fact just above at that time.
           const SCEV *Step = AR->getStepRecurrence(*this);
           Type *Ty = ZExt->getType();
-          auto *S = getAddRecExpr(
-            getExtendAddRecStart<SCEVZeroExtendExpr>(AR, Ty, this, 0),
-            getZeroExtendExpr(Step, Ty, 0), L, AR->getNoWrapFlags());
+          const SCEV *S = getAddRecExpr(
+              getExtendAddRecStart<SCEVZeroExtendExpr>(AR, Ty, this, 0),
+              getZeroExtendExpr(Step, Ty, 0), L, AR->getNoWrapFlags());
           IV = dyn_cast<SCEVAddRecExpr>(S);
         }
       }
diff --git a/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp b/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp
index 573f539b2a6ef..8e3637e00d9e0 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp
@@ -322,7 +322,8 @@ const SCEV *vputils::getSCEVExprForVPValue(const VPValue *V,
               return SE.getTruncateExpr(AddRec, R->getScalarType());
             return AddRec;
           })
-          .Case([&SE, &PSE, L](const VPWidenPointerInductionRecipe *R) {
+          .Case([&SE, &PSE,
+                 L](const VPWidenPointerInductionRecipe *R) -> const SCEV * {
             const SCEV *Start =
                 getSCEVExprForVPValue(R->getStartValue(), PSE, L);
             if (!L || isa<SCEVCouldNotCompute>(Start))
diff --git a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
index a50566cb64093..450b3bfcd32b9 100644
--- a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
+++ b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
@@ -2374,4 +2374,110 @@ TEST_F(ScalarEvolutionsTest, ExtendFoldCacheKeysUseFlags) {
     EXPECT_EQ(cast<SCEVZeroExtendExpr>(SExtPlain)->getOperand(), Mul);
   });
 }
+
+TEST_F(ScalarEvolutionsTest, AddRecExprUseFlags) {
+  LLVMContext C;
+  SMDiagnostic Err;
+  std::unique_ptr<Module> M = parseAssemblyString(
+      R"(define void @f(i32 %a, i32 %b, i32 %c) {
+      entry:
+        br label %loop.1
+
+      loop.1:
+        %iv.1 = phi i32 [ 0, %entry ], [ %iv.1.next, %loop.1 ]
+        %iv.1.next = add i32 %iv.1, 1
+        %cond.1 = icmp ult i32 %iv.1.next, 10
+        br i1 %cond.1, label %loop.1, label %loop.2
+
+      loop.2:
+        %iv.2 = phi i32 [ 0, %loop.1 ], [ %iv.2.next, %loop.2 ]
+        %iv.2.next = add i32 %iv.2, 1
+        %cond.2 = icmp ult i32 %iv.2.next, 10
+        br i1 %cond.2, label %loop.2, label %exit
+
+      exit:
+        ret void
+      })",
+      Err, C);
+
+  if (!M) {
+    Err.print("ScalarEvolutionTest", errs());
+    ASSERT_TRUE(M && "Could not parse module?");
+  }
+  ASSERT_TRUE(!verifyModule(*M, &errs()) && "Must have been well formed!");
+
+  runWithSE(*M, "f", [](Function &F, LoopInfo &LI, ScalarEvolution &SE) {
+    const SCEV *A = SE.getSCEV(getArgByName(F, "a"));
+    const SCEV *B = SE.getSCEV(getArgByName(F, "b"));
+    const SCEV *Cc = SE.getSCEV(getArgByName(F, "c"));
+    Type *I32 = A->getType();
+    const Loop *L1 =
+        LI.getLoopFor(getInstructionByName(F, "iv.1")->getParent());
+    const Loop *L2 =
+        LI.getLoopFor(getInstructionByName(F, "iv.2")->getParent());
+    ASSERT_NE(L1, nullptr);
+    ASSERT_NE(L2, nullptr);
+
+    // The recurrence is built without simplifications and carries the
+    // use-specific flags.
+    SCEVUse AR = SE.getAddRecExpr(A, B, L1, SCEV::FlagAnyWrap, SCEV::FlagNUW);
+    EXPECT_TRUE(AR.hasUseFlags());
+    EXPECT_EQ(AR.getUseNoWrapFlags(), SCEV::FlagNUW | SCEV::FlagNW);
+
+    const SCEV *BareAR = SE.getAddRecExpr(A, B, L1, SCEV::FlagAnyWrap);
+    EXPECT_EQ(AR.getPointer(), BareAR);
+    EXPECT_EQ(cast<SCEVAddRecExpr>(BareAR)->getNoWrapFlags(SCEV::FlagNUW),
+              SCEV::FlagAnyWrap);
+    EXPECT_EQ(AR.getCanonical(), BareAR);
+
+    // Operands are ordered, swapping them describes a different AddRec, with
+    // different flags.
+    SCEVUse Swapped =
+        SE.getAddRecExpr(B, A, L1, SCEV::FlagAnyWrap, SCEV::FlagNSW);
+    EXPECT_NE(Swapped.getPointer(), AR.getPointer());
+    EXPECT_TRUE(Swapped.hasUseFlags());
+    EXPECT_EQ(Swapped.getUseNoWrapFlags(), SCEV::FlagNSW | SCEV::FlagNW);
+
+    // The same holds for recurrences with more than two operands.
+    SmallVector<SCEVUse, 3> Ops = {A, B, Cc};
+    SCEVUse AR3 = SE.getAddRecExpr(Ops, L1, SCEV::FlagAnyWrap, SCEV::FlagNSW);
+    EXPECT_EQ(AR3.getUseNoWrapFlags(), SCEV::FlagNSW | SCEV::FlagNW);
+    EXPECT_EQ(AR3.getCanonical(), SE.getAddRecExpr(Ops, L1, SCEV::FlagAnyWrap));
+
+    // Flags the expression already carries add nothing to the use.
+    SCEVUse NUWAR = SE.getAddRecExpr(A, Cc, L1, SCEV::FlagNUW, SCEV::FlagNUW);
+    ASSERT_TRUE(cast<SCEVAddRecExpr>(NUWAR.getPointer())->hasNoUnsignedWrap());
+    EXPECT_FALSE(NUWAR.hasUseFlags());
+
+    auto CheckNoUseFlags = [](SCEVUse U) {
+      EXPECT_FALSE(U.hasUseFlags());
+      EXPECT_EQ(U.getUseNoWrapFlags(), SCEV::FlagAnyWrap);
+    };
+
+    // A zero step folds the recurrence away to its start value.
+    CheckNoUseFlags(SE.getAddRecExpr(A, SE.getZero(I32), L1, SCEV::FlagAnyWrap,
+                                     SCEV::FlagNUW));
+
+    // A zero trailing operand shortens the recurrence.
+    SmallVector<SCEVUse, 3> OpsWithZeroStep = {A, B, SE.getZero(I32)};
+    CheckNoUseFlags(SE.getAddRecExpr(OpsWithZeroStep, L1, SCEV::FlagAnyWrap,
+                                     SCEV::FlagNUW));
+
+    // A step that is itself a recurrence in the same loop gets inlined.
+    const SCEV *StepAR = SE.getAddRecExpr(B, Cc, L1, SCEV::FlagAnyWrap);
+    CheckNoUseFlags(
+        SE.getAddRecExpr(A, StepAR, L1, SCEV::FlagAnyWrap, SCEV::FlagNUW));
+
+    // A recurrence with a zero step can be folded to a different AddRec.
+    ASSERT_EQ(cast<SCEVAddRecExpr>(BareAR)->getLoop(), L1);
+    CheckNoUseFlags(SE.getAddRecExpr(BareAR, SE.getZero(I32), L2,
+                                     SCEV::FlagAnyWrap, SCEV::FlagNUW));
+
+#ifndef NDEBUG
+    EXPECT_DEATH(
+        (void)SE.getAddRecExpr(A, B, L1, SCEV::FlagAnyWrap, SCEV::FlagNW),
+        "only nuw or nsw allowed");
+#endif
+  });
+}
 }  // end namespace llvm



More information about the llvm-commits mailing list