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

Florian Hahn via llvm-commits llvm-commits at lists.llvm.org
Fri Sep 11 03:16:39 PDT 2026


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

>From 928a7edcd2c66a1b9aca22cc9d320cc008b1af30 Mon Sep 17 00:00:00 2001
From: Florian Hahn <flo at fhahn.com>
Date: Mon, 31 Aug 2026 09:52:13 +0100
Subject: [PATCH 1/3] [SCEV] Return a SCEVUse from getAddExpr and propagate use
 flags.

Add option to pass SCEVUse-specific flags to getAddExpr and propagate
them through, if valid conservatively. That is, the final expression
adds the same operands (potentially in different order). For example, it
is not valid to propagate the use flags if other (sub-)expressions have
been inlined.

It also includes a few mechanical changes, to update users that still
expected const SCEV *.
---
 llvm/include/llvm/Analysis/ScalarEvolution.h  | 25 ++---
 llvm/lib/Analysis/LoopAccessAnalysis.cpp      |  3 +-
 llvm/lib/Analysis/ScalarEvolution.cpp         | 92 +++++++++----------
 .../Target/PowerPC/PPCLoopInstrFormPrep.cpp   |  1 +
 .../Scalar/InductiveRangeCheckElimination.cpp | 31 +++----
 .../Transforms/Scalar/LoopStrengthReduce.cpp  | 10 +-
 llvm/lib/Transforms/Vectorize/VPlanUtils.cpp  |  2 +-
 .../Analysis/ScalarEvolutionTest.cpp          | 78 ++++++++++++++++
 8 files changed, 158 insertions(+), 84 deletions(-)

diff --git a/llvm/include/llvm/Analysis/ScalarEvolution.h b/llvm/include/llvm/Analysis/ScalarEvolution.h
index 77208bffce21b9..dbac53884874b0 100644
--- a/llvm/include/llvm/Analysis/ScalarEvolution.h
+++ b/llvm/include/llvm/Analysis/ScalarEvolution.h
@@ -749,20 +749,23 @@ class ScalarEvolution {
   LLVM_ABI const SCEV *getCastExpr(SCEVTypes Kind, SCEVUse Op, Type *Ty);
   LLVM_ABI const SCEV *getAnyExtendExpr(SCEVUse Op, Type *Ty);
 
-  LLVM_ABI const SCEV *getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
-                                  SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
-                                  unsigned Depth = 0);
-  const SCEV *getAddExpr(SCEVUse LHS, SCEVUse RHS,
-                         SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
-                         unsigned Depth = 0) {
+  LLVM_ABI SCEVUse getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
+                              SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
+                              unsigned Depth = 0,
+                              SCEV::NoWrapFlags UseFlags = SCEV::FlagAnyWrap);
+  SCEVUse getAddExpr(SCEVUse LHS, SCEVUse RHS,
+                     SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
+                     unsigned Depth = 0,
+                     SCEV::NoWrapFlags UseFlags = SCEV::FlagAnyWrap) {
     SmallVector<SCEVUse, 2> Ops = {LHS, RHS};
-    return getAddExpr(Ops, Flags, Depth);
+    return getAddExpr(Ops, Flags, Depth, UseFlags);
   }
-  const SCEV *getAddExpr(SCEVUse Op0, SCEVUse Op1, SCEVUse Op2,
-                         SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
-                         unsigned Depth = 0) {
+  SCEVUse getAddExpr(SCEVUse Op0, SCEVUse Op1, SCEVUse Op2,
+                     SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
+                     unsigned Depth = 0,
+                     SCEV::NoWrapFlags UseFlags = SCEV::FlagAnyWrap) {
     SmallVector<SCEVUse, 3> Ops = {Op0, Op1, Op2};
-    return getAddExpr(Ops, Flags, Depth);
+    return getAddExpr(Ops, Flags, Depth, UseFlags);
   }
   LLVM_ABI const SCEV *getMulExpr(SmallVectorImpl<SCEVUse> &Ops,
                                   SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
diff --git a/llvm/lib/Analysis/LoopAccessAnalysis.cpp b/llvm/lib/Analysis/LoopAccessAnalysis.cpp
index 2d81fb270df632..7e6e478ce5f7e1 100644
--- a/llvm/lib/Analysis/LoopAccessAnalysis.cpp
+++ b/llvm/lib/Analysis/LoopAccessAnalysis.cpp
@@ -1220,7 +1220,8 @@ static void findForkedSCEVs(
     return get<1>(S);
   };
 
-  auto GetBinOpExpr = [&SE](unsigned Opcode, const SCEV *L, const SCEV *R) {
+  auto GetBinOpExpr = [&SE](unsigned Opcode, const SCEV *L,
+                            const SCEV *R) -> const SCEV * {
     switch (Opcode) {
     case Instruction::Add:
       return SE->getAddExpr(L, R);
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index 06f350fd2179be..cdd6e12fb86072 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -970,27 +970,6 @@ static const SCEV *BinomialCoefficient(const SCEV *It, unsigned K,
                        SE.getTruncateOrZeroExtend(DivResult, ResultTy));
 }
 
-/// Attach \p UseFlags to \p Res as use-specific flags, but only if \p Res
-/// really is the two-operand \p ExprT over \p LHS and \p RHS - in either order,
-/// as operands get sorted by complexity.
-///
-/// Flags established for that operation say nothing about any other expression:
-/// a folded-away operand, a flattened nested expression or a distributed
-/// constant all give a different computation. They must not be attached to it,
-/// because an n-ary expression's no-wrap flags have to hold for all subsets and
-/// orders of its operands, and SCEVExpander relies on that when it stamps them
-/// on every partial sum or product it builds.
-template <typename ExprT>
-static SCEVUse withUseFlagsIfNotFolded(const SCEV *Res, SCEVUse LHS,
-                                       SCEVUse RHS,
-                                       SCEV::NoWrapFlags UseFlags) {
-  auto *E = dyn_cast<ExprT>(Res);
-  if (E && (equal(E->operands(), ArrayRef<SCEVUse>({LHS, RHS})) ||
-            equal(E->operands(), ArrayRef<SCEVUse>({RHS, LHS}))))
-    return {Res, UseFlags};
-  return Res;
-}
-
 /// Return the value of this chain of recurrences at the specified iteration
 /// number.  We can evaluate this recurrence by multiplying each element in the
 /// chain by the binomial coefficient corresponding to it.  In other words, we
@@ -1020,8 +999,8 @@ SCEVUse SCEVAddRecExpr::evaluateAtIteration(ArrayRef<SCEVUse> Operands,
       return Coeff;
 
     const SCEV *Mul = SE.getMulExpr(Operands[i].getPointer(), Coeff);
-    Result = withUseFlagsIfNotFolded<SCEVAddExpr>(SE.getAddExpr(Result, Mul),
-                                                  Result, Mul, UseFlags);
+    Result = SE.getAddExpr(Result, Mul, SCEV::FlagAnyWrap, /*Depth=*/0,
+                           UseFlags);
   }
   return Result;
 }
@@ -2320,21 +2299,18 @@ static bool CollectAddOperandsWithScales(SmallDenseMap<SCEVUse, APInt, 16> &M,
 bool ScalarEvolution::willNotOverflow(Instruction::BinaryOps BinOp, bool Signed,
                                       const SCEV *LHS, const SCEV *RHS,
                                       const Instruction *CtxI) {
-  const SCEV *(ScalarEvolution::*Operation)(SCEVUse, SCEVUse, SCEV::NoWrapFlags,
-                                            unsigned);
-  switch (BinOp) {
-  default:
-    llvm_unreachable("Unsupported binary op");
-  case Instruction::Add:
-    Operation = &ScalarEvolution::getAddExpr;
-    break;
-  case Instruction::Sub:
-    Operation = &ScalarEvolution::getMinusSCEV;
-    break;
-  case Instruction::Mul:
-    Operation = &ScalarEvolution::getMulExpr;
-    break;
-  }
+  auto Operation = [this, BinOp](SCEVUse L, SCEVUse R) -> const SCEV * {
+    switch (BinOp) {
+    default:
+      llvm_unreachable("Unsupported binary op");
+    case Instruction::Add:
+      return getAddExpr(L, R);
+    case Instruction::Sub:
+      return getMinusSCEV(L, R);
+    case Instruction::Mul:
+      return getMulExpr(L, R);
+    }
+  };
 
   const SCEV *(ScalarEvolution::*Extension)(SCEVUse, Type *, unsigned) =
       Signed ? &ScalarEvolution::getSignExtendExpr
@@ -2345,11 +2321,10 @@ bool ScalarEvolution::willNotOverflow(Instruction::BinaryOps BinOp, bool Signed,
   auto *WideTy =
       IntegerType::get(NarrowTy->getContext(), NarrowTy->getBitWidth() * 2);
 
-  const SCEV *A = (this->*Extension)(
-      (this->*Operation)(LHS, RHS, SCEV::FlagAnyWrap, 0), WideTy, 0);
+  const SCEV *A = (this->*Extension)(Operation(LHS, RHS), WideTy, 0);
   const SCEV *LHSB = (this->*Extension)(LHS, WideTy, 0);
   const SCEV *RHSB = (this->*Extension)(RHS, WideTy, 0);
-  const SCEV *B = (this->*Operation)(LHSB, RHSB, SCEV::FlagAnyWrap, 0);
+  const SCEV *B = Operation(LHSB, RHSB);
   if (A == B)
     return true;
   // Can we use context to prove the fact we need?
@@ -2539,11 +2514,13 @@ bool ScalarEvolution::isAvailableAtLoopEntry(const SCEV *S, const Loop *L) {
 }
 
 /// Get a canonical add expression, or something simpler if possible.
-const SCEV *ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
-                                        SCEV::NoWrapFlags OrigFlags,
-                                        unsigned Depth) {
+SCEVUse ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
+                                    SCEV::NoWrapFlags OrigFlags, unsigned Depth,
+                                    SCEV::NoWrapFlags UseFlags) {
   assert(!(OrigFlags & ~(SCEV::FlagNUW | SCEV::FlagNSW)) &&
          "only nuw or nsw allowed");
+  assert(!(UseFlags & ~(SCEV::FlagNUW | SCEV::FlagNSW)) &&
+         "only nuw or nsw allowed");
   assert(!Ops.empty() && "Cannot get empty add!");
   if (Ops.size() == 1) return Ops[0];
 #ifndef NDEBUG
@@ -2555,6 +2532,10 @@ const SCEV *ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
       Ops, [](const SCEV *Op) { return Op->getType()->isPointerTy(); });
   assert(NumPtrs <= 1 && "add has at most one pointer operand");
 #endif
+  // Keep track of original ops, if use-specific flags have been provided.
+  SmallVector<SCEVUse, 8> OrigOps;
+  if (UseFlags != SCEV::FlagAnyWrap)
+    OrigOps.assign(Ops.begin(), Ops.end());
 
   const SCEV *Folded = constantFoldAndGroupOps(
       *this, LI, DT, Ops,
@@ -2564,6 +2545,15 @@ const SCEV *ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
   if (Folded)
     return Folded;
 
+  // Conservatively drop use-specific flags if operands changed after constant
+  // folding, i.e. we are building a different expression than the initial one,
+  // for which the use-specific flags hold.
+  // TODO: In some cases, this is overly conservative.
+  if (UseFlags != SCEV::FlagAnyWrap &&
+      !std::is_permutation(OrigOps.begin(), OrigOps.end(), Ops.begin(),
+                           Ops.end()))
+    UseFlags = SCEV::FlagAnyWrap;
+
   unsigned Idx = isa<SCEVConstant>(Ops[0]) ? 1 : 0;
 
   // Delay expensive flag strengthening until necessary.
@@ -2573,14 +2563,14 @@ const SCEV *ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
 
   // Limit recursion calls depth.
   if (Depth > MaxArithDepth || hasHugeExpression(Ops))
-    return getOrCreateAddExpr(Ops, ComputeFlags(Ops));
+    return {getOrCreateAddExpr(Ops, ComputeFlags(Ops)), UseFlags};
 
   if (SCEV *S = findExistingSCEVInCache(scAddExpr, Ops)) {
     // Don't strengthen flags if we have no new information.
     SCEVAddExpr *Add = static_cast<SCEVAddExpr *>(S);
     if (Add->getNoWrapFlags(OrigFlags) != OrigFlags)
       Add->setNoWrapFlags(ComputeFlags(Ops));
-    return S;
+    return {S, UseFlags};
   }
 
   // Okay, check to see if the same value occurs in the operand list more than
@@ -3002,7 +2992,11 @@ const SCEV *ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
 
   // Okay, it looks like we really DO need an add expr.  Check to see if we
   // already have one, otherwise create a new one.
-  return getOrCreateAddExpr(Ops, ComputeFlags(Ops));
+  assert((UseFlags == SCEV::FlagAnyWrap ||
+          std::is_permutation(OrigOps.begin(), OrigOps.end(), Ops.begin(),
+                              Ops.end())) &&
+         "Tried to add SCEVUse flags after operands changed");
+  return {getOrCreateAddExpr(Ops, ComputeFlags(Ops)), UseFlags};
 }
 
 const SCEV *ScalarEvolution::getOrCreateAddExpr(ArrayRef<SCEVUse> Ops,
@@ -3854,7 +3848,7 @@ const SCEV *ScalarEvolution::getGEPExpr(SCEVUse BaseExpr,
   bool NUW = NW.hasNoUnsignedWrap() ||
              (NW.hasNoUnsignedSignedWrap() && isKnownNonNegative(Offset));
   SCEV::NoWrapFlags BaseWrap = NUW ? SCEV::FlagNUW : SCEV::FlagAnyWrap;
-  auto *GEPExpr = getAddExpr(BaseExpr, Offset, BaseWrap);
+  const SCEV *GEPExpr = getAddExpr(BaseExpr, Offset, BaseWrap);
   assert(BaseExpr->getType() == GEPExpr->getType() &&
          "GEP should not change type mid-flight.");
   return GEPExpr;
@@ -13737,7 +13731,7 @@ ScalarEvolution::howManyLessThans(const SCEV *LHS, const SCEV *RHS,
         //
         // FIXME: Should isLoopEntryGuardedByCond do this for us?
         auto CondGT = IsSigned ? ICmpInst::ICMP_SGT : ICmpInst::ICMP_UGT;
-        auto *StartMinusOne =
+        const SCEV *StartMinusOne =
             getAddExpr(OrigStart, getMinusOne(OrigStart->getType()));
         return isLoopEntryGuardedByCond(L, CondGT, OrigRHS, StartMinusOne);
       };
diff --git a/llvm/lib/Target/PowerPC/PPCLoopInstrFormPrep.cpp b/llvm/lib/Target/PowerPC/PPCLoopInstrFormPrep.cpp
index c8f96bf6f30448..ea8de04d6c2227 100644
--- a/llvm/lib/Target/PowerPC/PPCLoopInstrFormPrep.cpp
+++ b/llvm/lib/Target/PowerPC/PPCLoopInstrFormPrep.cpp
@@ -553,6 +553,7 @@ bool PPCLoopInstrFormPrep::rewriteLoadStoresForCommoningChains(
     const SCEV *BaseSCEV =
         ChainIdx ? SE->getAddExpr(Bucket.BaseSCEV,
                                   Bucket.Elements[BaseElemIdx].Offset)
+                       .getPointer()
                  : Bucket.BaseSCEV;
     const SCEVAddRecExpr *BasePtrSCEV = cast<SCEVAddRecExpr>(BaseSCEV);
 
diff --git a/llvm/lib/Transforms/Scalar/InductiveRangeCheckElimination.cpp b/llvm/lib/Transforms/Scalar/InductiveRangeCheckElimination.cpp
index 6d2af61de000f4..6d4d16825e6fcf 100644
--- a/llvm/lib/Transforms/Scalar/InductiveRangeCheckElimination.cpp
+++ b/llvm/lib/Transforms/Scalar/InductiveRangeCheckElimination.cpp
@@ -425,22 +425,20 @@ bool InductiveRangeCheck::reassociateSubLHS(
   auto getExprScaledIfOverflow = [&](Instruction::BinaryOps BinOp,
                                      const SCEV *LHS,
                                      const SCEV *RHS) -> const SCEV * {
-    const SCEV *(ScalarEvolution::*Operation)(SCEVUse, SCEVUse,
-                                              SCEV::NoWrapFlags, unsigned);
-    switch (BinOp) {
-    default:
-      llvm_unreachable("Unsupported binary op");
-    case Instruction::Add:
-      Operation = &ScalarEvolution::getAddExpr;
-      break;
-    case Instruction::Sub:
-      Operation = &ScalarEvolution::getMinusSCEV;
-      break;
-    }
+    auto Operation = [&SE, BinOp](SCEVUse L, SCEVUse R) -> const SCEV * {
+      switch (BinOp) {
+      default:
+        llvm_unreachable("Unsupported binary op");
+      case Instruction::Add:
+        return SE.getAddExpr(L, R);
+      case Instruction::Sub:
+        return SE.getMinusSCEV(L, R);
+      }
+    };
 
     if (SE.willNotOverflow(BinOp, ICmpInst::isSigned(Pred), LHS, RHS,
                            cast<Instruction>(VariantLHS)))
-      return (SE.*Operation)(LHS, RHS, SCEV::FlagAnyWrap, 0);
+      return Operation(LHS, RHS);
 
     // We couldn't prove that the expression does not overflow.
     // Than scale it to a wider type to check overflow at runtime.
@@ -449,9 +447,8 @@ bool InductiveRangeCheck::reassociateSubLHS(
       return nullptr;
 
     auto WideTy = IntegerType::get(Ty->getContext(), Ty->getBitWidth() * 2);
-    return (SE.*Operation)(SE.getSignExtendExpr(LHS, WideTy),
-                           SE.getSignExtendExpr(RHS, WideTy), SCEV::FlagAnyWrap,
-                           0);
+    return Operation(SE.getSignExtendExpr(LHS, WideTy),
+                     SE.getSignExtendExpr(RHS, WideTy));
   };
 
   if (OffsetSubtracted)
@@ -767,7 +764,7 @@ InductiveRangeCheck::computeSafeIterationSpace(ScalarEvolution &SE,
   const SCEV *Zero = SE.getZero(M->getType());
 
   // This function returns SCEV equal to 1 if X is non-negative 0 otherwise.
-  auto SCEVCheckNonNegative = [&](const SCEV *X) {
+  auto SCEVCheckNonNegative = [&](const SCEV *X) -> const SCEV * {
     const Loop *L = IndVar->getLoop();
     const SCEV *Zero = SE.getZero(X->getType());
     const SCEV *One = SE.getOne(X->getType());
diff --git a/llvm/lib/Transforms/Scalar/LoopStrengthReduce.cpp b/llvm/lib/Transforms/Scalar/LoopStrengthReduce.cpp
index e2ca5f166dbb5c..0528f9877804ac 100644
--- a/llvm/lib/Transforms/Scalar/LoopStrengthReduce.cpp
+++ b/llvm/lib/Transforms/Scalar/LoopStrengthReduce.cpp
@@ -3506,8 +3506,9 @@ void LSRInstance::GenerateIVChain(const IVChain &Chain,
       // be signed.
       const SCEV *IncExpr = SE.getNoopOrSignExtend(Inc.IncExpr, IntTy);
       Accum = SE.getAddExpr(Accum, IncExpr);
-      LeftOverExpr = LeftOverExpr ?
-        SE.getAddExpr(LeftOverExpr, IncExpr) : IncExpr;
+      LeftOverExpr = LeftOverExpr
+                         ? SE.getAddExpr(LeftOverExpr, IncExpr).getPointer()
+                         : IncExpr;
     }
 
     // Look through each base to see if any can produce a nice addressing mode.
@@ -5971,9 +5972,8 @@ Value *LSRInstance::Expand(const LSRUse &LU, const LSRFixup &LF,
   }
 
   // Emit instructions summing all the operands.
-  const SCEV *FullS = Ops.empty() ?
-                      SE.getConstant(IntTy, 0) :
-                      SE.getAddExpr(Ops);
+  const SCEV *FullS =
+      Ops.empty() ? SE.getConstant(IntTy, 0) : SE.getAddExpr(Ops).getPointer();
   Value *FullV = Rewriter.expandCodeFor(FullS, Ty);
 
   // We're done expanding now, so reset the rewriter.
diff --git a/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp b/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp
index a39760f3a34d1c..9c2677d807e901 100644
--- a/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp
+++ b/llvm/lib/Transforms/Vectorize/VPlanUtils.cpp
@@ -335,7 +335,7 @@ const SCEV *vputils::getSCEVExprForVPValue(const VPValue *V,
               return SE.getCouldNotCompute();
             return SE.getAddRecExpr(Start, Step, L, SCEV::FlagAnyWrap);
           })
-          .Case([&SE, &PSE, L](const VPDerivedIVRecipe *R) {
+          .Case([&SE, &PSE, L](const VPDerivedIVRecipe *R) -> const SCEV * {
             const SCEV *Start = getSCEVExprForVPValue(R->getOperand(0), PSE, L);
             const SCEV *IV = getSCEVExprForVPValue(R->getOperand(1), PSE, L);
             const SCEV *Scale = getSCEVExprForVPValue(R->getOperand(2), PSE, L);
diff --git a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
index a50566cb640939..6e271ba5e40859 100644
--- a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
+++ b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
@@ -2374,4 +2374,82 @@ TEST_F(ScalarEvolutionsTest, ExtendFoldCacheKeysUseFlags) {
     EXPECT_EQ(cast<SCEVZeroExtendExpr>(SExtPlain)->getOperand(), Mul);
   });
 }
+
+TEST_F(ScalarEvolutionsTest, AddExprUseFlags) {
+  LLVMContext C;
+  SMDiagnostic Err;
+  std::unique_ptr<Module> M = parseAssemblyString(
+      R"(define void @f(i32 %a, i32 %b, i32 %c) {
+      entry:
+        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();
+
+    // The sum is built as-is, so the use carries the requested flags.
+    SCEVUse Sum = SE.getAddExpr(A, B, SCEV::FlagAnyWrap, 0, SCEV::FlagNUW);
+    EXPECT_TRUE(Sum.hasUseFlags());
+    EXPECT_EQ(Sum.getUseNoWrapFlags(), SCEV::FlagNUW | SCEV::FlagNW);
+
+    // Same as Sun, but without NUW use flags.
+    const SCEV *BareSum = SE.getAddExpr(A, B);
+    EXPECT_EQ(Sum.getPointer(), BareSum);
+    EXPECT_EQ(cast<SCEVAddExpr>(BareSum)->getNoWrapFlags(SCEV::FlagNUW),
+              SCEV::FlagAnyWrap);
+    EXPECT_EQ(Sum.getCanonical(), BareSum);
+    EXPECT_FALSE(SCEVUse(BareSum).hasUseFlags());
+
+    // Operands get sorted by complexity, so their order does not matter.
+    EXPECT_EQ(SE.getAddExpr(B, A, SCEV::FlagAnyWrap, 0, SCEV::FlagNUW), Sum);
+
+    // The same holds for sums of more than two operands.
+    SmallVector<SCEVUse, 3> Ops = {Cc, B, A};
+    SCEVUse Sum3 = SE.getAddExpr(Ops, SCEV::FlagAnyWrap, 0, SCEV::FlagNSW);
+    EXPECT_EQ(Sum3.getUseNoWrapFlags(), SCEV::FlagNSW | SCEV::FlagNW);
+    EXPECT_EQ(Sum3.getCanonical(), SE.getAddExpr(A, B, Cc));
+
+    // Flags the expression already carries add nothing to the use.
+    SCEVUse NUWSum = SE.getAddExpr(A, Cc, SCEV::FlagNUW, 0, SCEV::FlagNUW);
+    ASSERT_TRUE(cast<SCEVAddExpr>(NUWSum.getPointer())->hasNoUnsignedWrap());
+    EXPECT_FALSE(NUWSum.hasUseFlags());
+
+    // A folded-away operand, a flattened nested sum and a sum distributed into
+    // a product all describe a different computation than the requested sum,
+    // so none of them may carry its flags.
+    auto CheckNoUseFlags = [](SCEVUse U) {
+      EXPECT_FALSE(U.hasUseFlags());
+      EXPECT_EQ(U.getUseNoWrapFlags(), SCEV::FlagAnyWrap);
+    };
+    CheckNoUseFlags(SE.getAddExpr(SE.getConstant(APInt(32, 1)),
+                                  SE.getConstant(APInt(32, 2)),
+                                  SCEV::FlagAnyWrap, 0, SCEV::FlagNUW));
+    CheckNoUseFlags(
+        SE.getAddExpr(A, SE.getZero(I32), SCEV::FlagAnyWrap, 0, SCEV::FlagNUW));
+    CheckNoUseFlags(SE.getAddExpr(A, SE.getAddExpr(B, Cc), SCEV::FlagAnyWrap, 0,
+                                  SCEV::FlagNUW));
+    CheckNoUseFlags(SE.getAddExpr(A, A, SCEV::FlagAnyWrap, 0, SCEV::FlagNUW));
+
+    // Constants folding together still leave a sum, but not the requested one.
+    SmallVector<SCEVUse, 3> FoldedOps = {SE.getConstant(APInt(32, 1)),
+                                         SE.getConstant(APInt(32, 2)), A};
+    CheckNoUseFlags(
+        SE.getAddExpr(FoldedOps, SCEV::FlagAnyWrap, 0, SCEV::FlagNUW));
+
+#ifndef NDEBUG
+    EXPECT_DEATH((void)SE.getAddExpr(A, B, SCEV::FlagAnyWrap, 0, SCEV::FlagNW),
+                 "only nuw or nsw allowed");
+#endif
+  });
+}
 }  // end namespace llvm

>From 2049d4a8a9b19e6a9482bb00f7b63a92b75c0e56 Mon Sep 17 00:00:00 2001
From: Florian Hahn <flo at fhahn.com>
Date: Fri, 4 Sep 2026 11:09:23 +0100
Subject: [PATCH 2/3] !fixup fix formatting

---
 llvm/lib/Analysis/ScalarEvolution.cpp | 4 ++--
 1 file changed, 2 insertions(+), 2 deletions(-)

diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index cdd6e12fb86072..db9c214643cdf4 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -999,8 +999,8 @@ SCEVUse SCEVAddRecExpr::evaluateAtIteration(ArrayRef<SCEVUse> Operands,
       return Coeff;
 
     const SCEV *Mul = SE.getMulExpr(Operands[i].getPointer(), Coeff);
-    Result = SE.getAddExpr(Result, Mul, SCEV::FlagAnyWrap, /*Depth=*/0,
-                           UseFlags);
+    Result =
+        SE.getAddExpr(Result, Mul, SCEV::FlagAnyWrap, /*Depth=*/0, UseFlags);
   }
   return Result;
 }

>From 766a7b1fa5692eac3ac652725a1ddc06c9396a47 Mon Sep 17 00:00:00 2001
From: Florian Hahn <flo at fhahn.com>
Date: Tue, 8 Sep 2026 11:23:12 +0100
Subject: [PATCH 3/3] !fixup pass flags via struct, keep flags when constant
 folding

---
 llvm/include/llvm/Analysis/ScalarEvolution.h  | 33 +++++++++------
 llvm/lib/Analysis/LoopAccessAnalysis.cpp      | 10 ++---
 llvm/lib/Analysis/ScalarEvolution.cpp         | 25 +++++-------
 .../Analysis/ScalarEvolutionTest.cpp          | 40 +++++++++++--------
 4 files changed, 58 insertions(+), 50 deletions(-)

diff --git a/llvm/include/llvm/Analysis/ScalarEvolution.h b/llvm/include/llvm/Analysis/ScalarEvolution.h
index dbac53884874b0..33ba3a47e365ee 100644
--- a/llvm/include/llvm/Analysis/ScalarEvolution.h
+++ b/llvm/include/llvm/Analysis/ScalarEvolution.h
@@ -192,6 +192,21 @@ template <typename SCEVPtrT> SCEVUseT(SCEVPtrT) -> SCEVUseT<SCEVPtrT>;
 
 using SCEVUse = SCEVUseT<const SCEV *>;
 
+/// The no-wrap flags to apply when creating a SCEV expression, to the
+/// expression and use respectively.
+struct SCEVFlags {
+  /// Flags applied directly to a SCEV expression, must be valid wherever the
+  /// expression is valid.
+  SCEVNoWrapFlags ExprFlags;
+
+  /// Flags only applied to a SCEVUse.
+  SCEVNoWrapFlags UseFlags;
+
+  constexpr SCEVFlags(SCEVNoWrapFlags ExprFlags = SCEVNoWrapFlags::FlagAnyWrap,
+                      SCEVNoWrapFlags UseFlags = SCEVNoWrapFlags::FlagAnyWrap)
+      : ExprFlags(ExprFlags), UseFlags(UseFlags) {}
+};
+
 /// Provide PointerLikeTypeTraits for SCEVUse, so it can be used with
 /// SmallPtrSet, among others.
 template <> struct PointerLikeTypeTraits<SCEVUse> {
@@ -750,22 +765,16 @@ class ScalarEvolution {
   LLVM_ABI const SCEV *getAnyExtendExpr(SCEVUse Op, Type *Ty);
 
   LLVM_ABI SCEVUse getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
-                              SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
-                              unsigned Depth = 0,
-                              SCEV::NoWrapFlags UseFlags = SCEV::FlagAnyWrap);
-  SCEVUse getAddExpr(SCEVUse LHS, SCEVUse RHS,
-                     SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
-                     unsigned Depth = 0,
-                     SCEV::NoWrapFlags UseFlags = SCEV::FlagAnyWrap) {
+                              SCEVFlags Flags = {}, unsigned Depth = 0);
+  SCEVUse getAddExpr(SCEVUse LHS, SCEVUse RHS, SCEVFlags Flags = {},
+                     unsigned Depth = 0) {
     SmallVector<SCEVUse, 2> Ops = {LHS, RHS};
-    return getAddExpr(Ops, Flags, Depth, UseFlags);
+    return getAddExpr(Ops, Flags, Depth);
   }
   SCEVUse getAddExpr(SCEVUse Op0, SCEVUse Op1, SCEVUse Op2,
-                     SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
-                     unsigned Depth = 0,
-                     SCEV::NoWrapFlags UseFlags = SCEV::FlagAnyWrap) {
+                     SCEVFlags Flags = {}, unsigned Depth = 0) {
     SmallVector<SCEVUse, 3> Ops = {Op0, Op1, Op2};
-    return getAddExpr(Ops, Flags, Depth, UseFlags);
+    return getAddExpr(Ops, Flags, Depth);
   }
   LLVM_ABI const SCEV *getMulExpr(SmallVectorImpl<SCEVUse> &Ops,
                                   SCEV::NoWrapFlags Flags = SCEV::FlagAnyWrap,
diff --git a/llvm/lib/Analysis/LoopAccessAnalysis.cpp b/llvm/lib/Analysis/LoopAccessAnalysis.cpp
index 7e6e478ce5f7e1..6fde1f26938776 100644
--- a/llvm/lib/Analysis/LoopAccessAnalysis.cpp
+++ b/llvm/lib/Analysis/LoopAccessAnalysis.cpp
@@ -458,11 +458,11 @@ std::pair<const SCEV *, const SCEV *> llvm::getStartAndEndForAccess(
       ScStart = Start;
       // The highest address for the type saturates; adding EltSize to it would
       // wrap to the start of the address space.
-      ScEnd =
-          LastAddr
-              ? SE->getAddExpr(LastAddr, EltSizeSCEV)
-              : SE->getSCEV(ConstantExpr::getIntToPtr(
-                    Constant::getAllOnesValue(DL.getIndexType(PtrTy)), PtrTy));
+      if (LastAddr)
+        ScEnd = SE->getAddExpr(LastAddr, EltSizeSCEV);
+      else
+        ScEnd = SE->getSCEV(ConstantExpr::getIntToPtr(
+            Constant::getAllOnesValue(DL.getIndexType(PtrTy)), PtrTy));
     } else {
       if (!LastAddr)
         return {SE->getCouldNotCompute(), SE->getCouldNotCompute()};
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index db9c214643cdf4..5c626608699ad4 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -999,8 +999,7 @@ SCEVUse SCEVAddRecExpr::evaluateAtIteration(ArrayRef<SCEVUse> Operands,
       return Coeff;
 
     const SCEV *Mul = SE.getMulExpr(Operands[i].getPointer(), Coeff);
-    Result =
-        SE.getAddExpr(Result, Mul, SCEV::FlagAnyWrap, /*Depth=*/0, UseFlags);
+    Result = SE.getAddExpr(Result, Mul, {SCEV::FlagAnyWrap, UseFlags});
   }
   return Result;
 }
@@ -2515,8 +2514,9 @@ bool ScalarEvolution::isAvailableAtLoopEntry(const SCEV *S, const Loop *L) {
 
 /// Get a canonical add expression, or something simpler if possible.
 SCEVUse ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
-                                    SCEV::NoWrapFlags OrigFlags, unsigned Depth,
-                                    SCEV::NoWrapFlags UseFlags) {
+                                    SCEVFlags Flags, unsigned Depth) {
+  SCEV::NoWrapFlags OrigFlags = Flags.ExprFlags;
+  SCEV::NoWrapFlags UseFlags = Flags.UseFlags;
   assert(!(OrigFlags & ~(SCEV::FlagNUW | SCEV::FlagNSW)) &&
          "only nuw or nsw allowed");
   assert(!(UseFlags & ~(SCEV::FlagNUW | SCEV::FlagNSW)) &&
@@ -2532,10 +2532,6 @@ SCEVUse ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
       Ops, [](const SCEV *Op) { return Op->getType()->isPointerTy(); });
   assert(NumPtrs <= 1 && "add has at most one pointer operand");
 #endif
-  // Keep track of original ops, if use-specific flags have been provided.
-  SmallVector<SCEVUse, 8> OrigOps;
-  if (UseFlags != SCEV::FlagAnyWrap)
-    OrigOps.assign(Ops.begin(), Ops.end());
 
   const SCEV *Folded = constantFoldAndGroupOps(
       *this, LI, DT, Ops,
@@ -2545,14 +2541,11 @@ SCEVUse ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
   if (Folded)
     return Folded;
 
-  // Conservatively drop use-specific flags if operands changed after constant
-  // folding, i.e. we are building a different expression than the initial one,
-  // for which the use-specific flags hold.
-  // TODO: In some cases, this is overly conservative.
-  if (UseFlags != SCEV::FlagAnyWrap &&
-      !std::is_permutation(OrigOps.begin(), OrigOps.end(), Ops.begin(),
-                           Ops.end()))
-    UseFlags = SCEV::FlagAnyWrap;
+#ifndef NDEBUG
+  // Keep track of operands after constant folding, for verification when adding
+  // use-specific flags.
+  const SmallVector<SCEVUse, 8> OrigOps(Ops.begin(), Ops.end());
+#endif
 
   unsigned Idx = isa<SCEVConstant>(Ops[0]) ? 1 : 0;
 
diff --git a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
index 6e271ba5e40859..e3fa7dd9c92b5f 100644
--- a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
+++ b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
@@ -2398,7 +2398,7 @@ TEST_F(ScalarEvolutionsTest, AddExprUseFlags) {
     Type *I32 = A->getType();
 
     // The sum is built as-is, so the use carries the requested flags.
-    SCEVUse Sum = SE.getAddExpr(A, B, SCEV::FlagAnyWrap, 0, SCEV::FlagNUW);
+    SCEVUse Sum = SE.getAddExpr(A, B, {SCEV::FlagAnyWrap, SCEV::FlagNUW});
     EXPECT_TRUE(Sum.hasUseFlags());
     EXPECT_EQ(Sum.getUseNoWrapFlags(), SCEV::FlagNUW | SCEV::FlagNW);
 
@@ -2411,43 +2411,49 @@ TEST_F(ScalarEvolutionsTest, AddExprUseFlags) {
     EXPECT_FALSE(SCEVUse(BareSum).hasUseFlags());
 
     // Operands get sorted by complexity, so their order does not matter.
-    EXPECT_EQ(SE.getAddExpr(B, A, SCEV::FlagAnyWrap, 0, SCEV::FlagNUW), Sum);
+    EXPECT_EQ(SE.getAddExpr(B, A, {SCEV::FlagAnyWrap, SCEV::FlagNUW}), Sum);
 
     // The same holds for sums of more than two operands.
     SmallVector<SCEVUse, 3> Ops = {Cc, B, A};
-    SCEVUse Sum3 = SE.getAddExpr(Ops, SCEV::FlagAnyWrap, 0, SCEV::FlagNSW);
+    SCEVUse Sum3 = SE.getAddExpr(Ops, {SCEV::FlagAnyWrap, SCEV::FlagNSW});
     EXPECT_EQ(Sum3.getUseNoWrapFlags(), SCEV::FlagNSW | SCEV::FlagNW);
     EXPECT_EQ(Sum3.getCanonical(), SE.getAddExpr(A, B, Cc));
 
     // Flags the expression already carries add nothing to the use.
-    SCEVUse NUWSum = SE.getAddExpr(A, Cc, SCEV::FlagNUW, 0, SCEV::FlagNUW);
+    SCEVUse NUWSum = SE.getAddExpr(A, Cc, {SCEV::FlagNUW, SCEV::FlagNUW});
     ASSERT_TRUE(cast<SCEVAddExpr>(NUWSum.getPointer())->hasNoUnsignedWrap());
     EXPECT_FALSE(NUWSum.hasUseFlags());
 
-    // A folded-away operand, a flattened nested sum and a sum distributed into
-    // a product all describe a different computation than the requested sum,
-    // so none of them may carry its flags.
+    // A sum folded to a single operand, a flattened nested sum and a
+    // sum distributed into a product all describe a different computation than
+    // the requested sum, so none of them may carry its flags.
     auto CheckNoUseFlags = [](SCEVUse U) {
       EXPECT_FALSE(U.hasUseFlags());
       EXPECT_EQ(U.getUseNoWrapFlags(), SCEV::FlagAnyWrap);
     };
     CheckNoUseFlags(SE.getAddExpr(SE.getConstant(APInt(32, 1)),
                                   SE.getConstant(APInt(32, 2)),
-                                  SCEV::FlagAnyWrap, 0, SCEV::FlagNUW));
+                                  {SCEV::FlagAnyWrap, SCEV::FlagNUW}));
     CheckNoUseFlags(
-        SE.getAddExpr(A, SE.getZero(I32), SCEV::FlagAnyWrap, 0, SCEV::FlagNUW));
-    CheckNoUseFlags(SE.getAddExpr(A, SE.getAddExpr(B, Cc), SCEV::FlagAnyWrap, 0,
-                                  SCEV::FlagNUW));
-    CheckNoUseFlags(SE.getAddExpr(A, A, SCEV::FlagAnyWrap, 0, SCEV::FlagNUW));
-
-    // Constants folding together still leave a sum, but not the requested one.
+        SE.getAddExpr(A, SE.getZero(I32), {SCEV::FlagAnyWrap, SCEV::FlagNUW}));
+    CheckNoUseFlags(SE.getAddExpr(A, SE.getAddExpr(B, Cc),
+                                  {SCEV::FlagAnyWrap, SCEV::FlagNUW}));
+    CheckNoUseFlags(SE.getAddExpr(A, A, {SCEV::FlagAnyWrap, SCEV::FlagNUW}));
+
+    // SCEV flags are valid for all subsets and orders of the operands, so
+    // use-specific flags can be preserved when folding constants. the requested
+    // flags already cover, so it keeps carrying them.
     SmallVector<SCEVUse, 3> FoldedOps = {SE.getConstant(APInt(32, 1)),
                                          SE.getConstant(APInt(32, 2)), A};
-    CheckNoUseFlags(
-        SE.getAddExpr(FoldedOps, SCEV::FlagAnyWrap, 0, SCEV::FlagNUW));
+    SCEVUse FoldedSum = SE.getAddExpr(
+        FoldedOps, {SCEV::FlagAnyWrap, SCEV::FlagNUW | SCEV::FlagNSW});
+    EXPECT_EQ(FoldedSum.getCanonical(),
+              SE.getAddExpr(SE.getConstant(I32, 3), A));
+    EXPECT_EQ(FoldedSum.getUseNoWrapFlags(),
+              SCEV::FlagNUW | SCEV::FlagNSW | SCEV::FlagNW);
 
 #ifndef NDEBUG
-    EXPECT_DEATH((void)SE.getAddExpr(A, B, SCEV::FlagAnyWrap, 0, SCEV::FlagNW),
+    EXPECT_DEATH((void)SE.getAddExpr(A, B, {SCEV::FlagAnyWrap, SCEV::FlagNW}),
                  "only nuw or nsw allowed");
 #endif
   });



More information about the llvm-commits mailing list