[llvm] [SCEV] Share logic of udiv-AR-overflow with willNotOverflow (PR #217133)

Ramkumar Ramachandra via llvm-commits llvm-commits at lists.llvm.org
Wed Aug 26 06:23:00 PDT 2026


https://github.com/artagnon updated https://github.com/llvm/llvm-project/pull/217133

>From ff851c61ea640fcb784be07ddf9d9e7a27053919 Mon Sep 17 00:00:00 2001
From: Ramkumar Ramachandra <artagnon at tenstorrent.com>
Date: Wed, 19 Aug 2026 08:53:15 +0100
Subject: [PATCH 1/4] [SCEV] Use willNotOverflow for AR-check in udiv (NFC)

Extend willNotOverflow to cover Instruction::PHI for AddRecs, and use it
to check the LHS AddRec's NUW in getUDivExpr.
---
 llvm/include/llvm/Analysis/ScalarEvolution.h | 14 ++--
 llvm/lib/Analysis/ScalarEvolution.cpp        | 69 +++++++++-----------
 2 files changed, 38 insertions(+), 45 deletions(-)

diff --git a/llvm/include/llvm/Analysis/ScalarEvolution.h b/llvm/include/llvm/Analysis/ScalarEvolution.h
index 4f59bc99cb697..89404ecee80c0 100644
--- a/llvm/include/llvm/Analysis/ScalarEvolution.h
+++ b/llvm/include/llvm/Analysis/ScalarEvolution.h
@@ -704,12 +704,14 @@ class ScalarEvolution {
   /// scAddRecExpr. The result will be cached in HasRecMap.
   LLVM_ABI bool containsAddRecurrence(const SCEV *S);
 
-  /// Is operation \p BinOp between \p LHS and \p RHS provably does not have
-  /// a signed/unsigned overflow (\p Signed)? If \p CtxI is specified, the
-  /// no-overflow fact should be true in the context of this instruction.
-  LLVM_ABI bool willNotOverflow(Instruction::BinaryOps BinOp, bool Signed,
-                                const SCEV *LHS, const SCEV *RHS,
-                                const Instruction *CtxI = nullptr);
+  /// Is operation with \p Opcode between \p LHS and \p RHS provably does not
+  /// have a signed/unsigned overflow (\p Signed)? If \p CtxI is specified, the
+  /// no-overflow fact should be true in the context of this instruction. If \p
+  /// Opcode is Instruction::PHI, the parameter \p L must be specified.
+  LLVM_ABI bool willNotOverflow(unsigned Opcode, bool Signed, const SCEV *LHS,
+                                const SCEV *RHS,
+                                const Instruction *CtxI = nullptr,
+                                const Loop *L = nullptr);
 
   /// Parse NSW/NUW flags from add/sub/mul IR binary operation \p Op into
   /// SCEV no-wrap flags, and deduce flag[s] that aren't known yet.
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index f6ad701ce819e..deff7f82d984c 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -2287,46 +2287,54 @@ static bool CollectAddOperandsWithScales(SmallDenseMap<SCEVUse, APInt, 16> &M,
   return Interesting;
 }
 
-bool ScalarEvolution::willNotOverflow(Instruction::BinaryOps BinOp, bool Signed,
+bool ScalarEvolution::willNotOverflow(unsigned Opcode, bool Signed,
                                       const SCEV *LHS, const SCEV *RHS,
-                                      const Instruction *CtxI) {
-  const SCEV *(ScalarEvolution::*Operation)(SCEVUse, SCEVUse, SCEV::NoWrapFlags,
-                                            unsigned);
-  switch (BinOp) {
+                                      const Instruction *CtxI, const Loop *L) {
+  using OpFnTy = const SCEV *(ScalarEvolution::*)(SCEVUse, SCEVUse,
+                                                  SCEV::NoWrapFlags, unsigned);
+  std::function<const SCEV *(SCEVUse, SCEVUse, SCEV::NoWrapFlags, unsigned)>
+      OperationFn;
+  switch (Opcode) {
   default:
     llvm_unreachable("Unsupported binary op");
   case Instruction::Add:
-    Operation = &ScalarEvolution::getAddExpr;
+    OperationFn = bind_front<OpFnTy>(&ScalarEvolution::getAddExpr, this);
     break;
   case Instruction::Sub:
-    Operation = &ScalarEvolution::getMinusSCEV;
+    OperationFn = bind_front<OpFnTy>(&ScalarEvolution::getMinusSCEV, this);
     break;
   case Instruction::Mul:
-    Operation = &ScalarEvolution::getMulExpr;
+    OperationFn = bind_front<OpFnTy>(&ScalarEvolution::getMulExpr, this);
+    break;
+  case Instruction::PHI:
+    assert(L && "Loop argument must be given for PHI");
+    OperationFn = [&](SCEVUse LHS, SCEVUse RHS, SCEV::NoWrapFlags NW,
+                      unsigned) { return getAddRecExpr(LHS, RHS, L, NW); };
     break;
   }
 
-  const SCEV *(ScalarEvolution::*Extension)(SCEVUse, Type *, unsigned) =
-      Signed ? &ScalarEvolution::getSignExtendExpr
-             : &ScalarEvolution::getZeroExtendExpr;
+  using ExtFnTy = const SCEV *(ScalarEvolution::*)(SCEVUse, Type *, unsigned);
+  std::function<const SCEV *(SCEVUse, Type *, unsigned)> ExtensionFn =
+      Signed ? bind_front<ExtFnTy>(&ScalarEvolution::getSignExtendExpr, this)
+             : bind_front<ExtFnTy>(&ScalarEvolution::getZeroExtendExpr, this);
 
   // Check ext(LHS op RHS) == ext(LHS) op ext(RHS)
   auto *NarrowTy = cast<IntegerType>(LHS->getType());
   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 *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 *A =
+      ExtensionFn(OperationFn(LHS, RHS, SCEV::FlagAnyWrap, 0), WideTy, 0);
+  const SCEV *LHSB = ExtensionFn(LHS, WideTy, 0);
+  const SCEV *RHSB = ExtensionFn(RHS, WideTy, 0);
+  const SCEV *B = OperationFn(LHSB, RHSB, SCEV::FlagAnyWrap, 0);
   if (A == B)
     return true;
   // Can we use context to prove the fact we need?
   if (!CtxI)
     return false;
   // TODO: Support mul.
-  if (BinOp == Instruction::Mul)
+  if (Opcode == Instruction::Mul)
     return false;
   auto *RHSC = dyn_cast<SCEVConstant>(RHS);
   // TODO: Lift this limitation.
@@ -2334,7 +2342,7 @@ bool ScalarEvolution::willNotOverflow(Instruction::BinaryOps BinOp, bool Signed,
     return false;
   APInt C = RHSC->getAPInt();
   unsigned NumBits = C.getBitWidth();
-  bool IsSub = (BinOp == Instruction::Sub);
+  bool IsSub = (Opcode == Instruction::Sub);
   bool IsNegativeConst = (Signed && C.isNegative());
   // Compute the direction and magnitude by which we need to check overflow.
   bool OverflowDown = IsSub ^ IsNegativeConst;
@@ -3470,27 +3478,16 @@ const SCEV *ScalarEvolution::getUDivExpr(SCEVUse LHS, SCEVUse RHS) {
     if (!RHSC->getValue()->isZero()) {
       // Determine if the division can be folded into the operands of
       // its operands.
-      // TODO: Generalize this to non-constants by using known-bits information.
-      Type *Ty = LHS->getType();
-      unsigned LZ = RHSC->getAPInt().countl_zero();
-      unsigned MaxShiftAmt = getTypeSizeInBits(Ty) - LZ - 1;
-      // For non-power-of-two values, effectively round the value up to the
-      // nearest power of two.
-      if (!RHSC->getAPInt().isPowerOf2())
-        ++MaxShiftAmt;
-      IntegerType *ExtTy =
-        IntegerType::get(getContext(), getTypeSizeInBits(Ty) + MaxShiftAmt);
       if (const SCEVAddRecExpr *AR = dyn_cast<SCEVAddRecExpr>(LHS))
         if (const SCEVConstant *Step =
             dyn_cast<SCEVConstant>(AR->getStepRecurrence(*this))) {
           // {X,+,N}/C --> {X/C,+,N/C} if safe and N/C can be folded.
           const APInt &StepInt = Step->getAPInt();
           const APInt &DivInt = RHSC->getAPInt();
-          if (!StepInt.urem(DivInt) &&
-              getZeroExtendExpr(AR, ExtTy) ==
-              getAddRecExpr(getZeroExtendExpr(AR->getStart(), ExtTy),
-                            getZeroExtendExpr(Step, ExtTy),
-                            AR->getLoop(), SCEV::FlagAnyWrap)) {
+          bool NoWrap = willNotOverflow(Instruction::PHI, /*Signed=*/false,
+                                        AR->getStart(), Step, /*CtxI=*/nullptr,
+                                        AR->getLoop());
+          if (!StepInt.urem(DivInt) && NoWrap) {
             SmallVector<SCEVUse, 4> Operands;
             for (const SCEV *Op : AR->operands())
               Operands.push_back(getUDivExpr(Op, RHS));
@@ -3501,12 +3498,6 @@ const SCEV *ScalarEvolution::getUDivExpr(SCEVUse LHS, SCEVUse RHS) {
           const APInt *StartRem;
           if (!DivInt.urem(StepInt) && match(getURemExpr(AR->getStart(), Step),
                                              m_scev_APInt(StartRem))) {
-            bool NoWrap =
-                getZeroExtendExpr(AR, ExtTy) ==
-                getAddRecExpr(getZeroExtendExpr(AR->getStart(), ExtTy),
-                              getZeroExtendExpr(Step, ExtTy), AR->getLoop(),
-                              SCEV::FlagAnyWrap);
-
             // With N <= C and both N, C as powers-of-2, the transformation
             // {X,+,N}/C => {(X - X%N),+,N}/C preserves division results even
             // if wrapping occurs, as the division results remain equivalent for

>From 85e0b8d8bdac2c46570ccd331d0638210ae17625 Mon Sep 17 00:00:00 2001
From: Ramkumar Ramachandra <artagnon at tenstorrent.com>
Date: Sun, 23 Aug 2026 17:27:18 +0100
Subject: [PATCH 2/4] [SCEV] Fix a couple of minor issues

---
 llvm/lib/Analysis/ScalarEvolution.cpp         |  5 +-
 .../uniform_across_vf_induction1_div_urem.ll  | 52 +++++++++++++++----
 2 files changed, 44 insertions(+), 13 deletions(-)

diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index deff7f82d984c..952b82b492ac1 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -2334,7 +2334,7 @@ bool ScalarEvolution::willNotOverflow(unsigned Opcode, bool Signed,
   if (!CtxI)
     return false;
   // TODO: Support mul.
-  if (Opcode == Instruction::Mul)
+  if (Opcode == Instruction::Mul || Opcode == Instruction::PHI)
     return false;
   auto *RHSC = dyn_cast<SCEVConstant>(RHS);
   // TODO: Lift this limitation.
@@ -3484,7 +3484,8 @@ const SCEV *ScalarEvolution::getUDivExpr(SCEVUse LHS, SCEVUse RHS) {
           // {X,+,N}/C --> {X/C,+,N/C} if safe and N/C can be folded.
           const APInt &StepInt = Step->getAPInt();
           const APInt &DivInt = RHSC->getAPInt();
-          bool NoWrap = willNotOverflow(Instruction::PHI, /*Signed=*/false,
+          bool NoWrap = !StepInt.urem(DivInt) &&
+                        willNotOverflow(Instruction::PHI, /*Signed=*/false,
                                         AR->getStart(), Step, /*CtxI=*/nullptr,
                                         AR->getLoop());
           if (!StepInt.urem(DivInt) && NoWrap) {
diff --git a/llvm/test/Transforms/LoopVectorize/uniform_across_vf_induction1_div_urem.ll b/llvm/test/Transforms/LoopVectorize/uniform_across_vf_induction1_div_urem.ll
index 3f40db5179407..7c14a9562394e 100644
--- a/llvm/test/Transforms/LoopVectorize/uniform_across_vf_induction1_div_urem.ll
+++ b/llvm/test/Transforms/LoopVectorize/uniform_across_vf_induction1_div_urem.ll
@@ -239,18 +239,48 @@ define void @ld_div8_urem3(ptr noalias %A, ptr noalias %B) {
 ; CHECK-NEXT:    br label [[VECTOR_BODY:%.*]]
 ; CHECK:       vector.body:
 ; CHECK-NEXT:    [[INDEX:%.*]] = phi i64 [ 0, [[VECTOR_PH]] ], [ [[INDEX_NEXT:%.*]], [[VECTOR_BODY]] ]
-; CHECK-NEXT:    [[TMP0:%.*]] = lshr i64 [[INDEX]], 3
-; CHECK-NEXT:    [[TMP1:%.*]] = urem i64 [[TMP0]], 3
-; CHECK-NEXT:    [[TMP2:%.*]] = getelementptr inbounds i64, ptr [[A]], i64 [[TMP1]]
-; CHECK-NEXT:    [[TMP3:%.*]] = load i64, ptr [[TMP2]], align 8
-; CHECK-NEXT:    [[TMP4:%.*]] = add nsw i64 [[TMP3]], 42
-; CHECK-NEXT:    [[BROADCAST_SPLATINSERT:%.*]] = insertelement <8 x i64> poison, i64 [[TMP4]], i64 0
-; CHECK-NEXT:    [[BROADCAST_SPLAT:%.*]] = shufflevector <8 x i64> [[BROADCAST_SPLATINSERT]], <8 x i64> poison, <8 x i32> zeroinitializer
-; CHECK-NEXT:    [[TMP5:%.*]] = getelementptr inbounds i64, ptr [[B]], i64 [[INDEX]]
-; CHECK-NEXT:    store <8 x i64> [[BROADCAST_SPLAT]], ptr [[TMP5]], align 8
+; CHECK-NEXT:    [[VEC_IND:%.*]] = phi <8 x i64> [ <i64 0, i64 1, i64 2, i64 3, i64 4, i64 5, i64 6, i64 7>, [[VECTOR_PH]] ], [ [[VEC_IND_NEXT:%.*]], [[VECTOR_BODY]] ]
+; CHECK-NEXT:    [[TMP0:%.*]] = lshr <8 x i64> [[VEC_IND]], splat (i64 3)
+; CHECK-NEXT:    [[TMP1:%.*]] = urem <8 x i64> [[TMP0]], splat (i64 3)
+; CHECK-NEXT:    [[TMP2:%.*]] = extractelement <8 x i64> [[TMP1]], i64 0
+; CHECK-NEXT:    [[TMP3:%.*]] = getelementptr inbounds i64, ptr [[A]], i64 [[TMP2]]
+; CHECK-NEXT:    [[TMP4:%.*]] = extractelement <8 x i64> [[TMP1]], i64 1
+; CHECK-NEXT:    [[TMP5:%.*]] = getelementptr inbounds i64, ptr [[A]], i64 [[TMP4]]
+; CHECK-NEXT:    [[TMP6:%.*]] = extractelement <8 x i64> [[TMP1]], i64 2
+; CHECK-NEXT:    [[TMP7:%.*]] = getelementptr inbounds i64, ptr [[A]], i64 [[TMP6]]
+; CHECK-NEXT:    [[TMP8:%.*]] = extractelement <8 x i64> [[TMP1]], i64 3
+; CHECK-NEXT:    [[TMP9:%.*]] = getelementptr inbounds i64, ptr [[A]], i64 [[TMP8]]
+; CHECK-NEXT:    [[TMP10:%.*]] = extractelement <8 x i64> [[TMP1]], i64 4
+; CHECK-NEXT:    [[TMP11:%.*]] = getelementptr inbounds i64, ptr [[A]], i64 [[TMP10]]
+; CHECK-NEXT:    [[TMP12:%.*]] = extractelement <8 x i64> [[TMP1]], i64 5
+; CHECK-NEXT:    [[TMP13:%.*]] = getelementptr inbounds i64, ptr [[A]], i64 [[TMP12]]
+; CHECK-NEXT:    [[TMP14:%.*]] = extractelement <8 x i64> [[TMP1]], i64 6
+; CHECK-NEXT:    [[TMP15:%.*]] = getelementptr inbounds i64, ptr [[A]], i64 [[TMP14]]
+; CHECK-NEXT:    [[TMP16:%.*]] = extractelement <8 x i64> [[TMP1]], i64 7
+; CHECK-NEXT:    [[TMP17:%.*]] = getelementptr inbounds i64, ptr [[A]], i64 [[TMP16]]
+; CHECK-NEXT:    [[TMP18:%.*]] = load i64, ptr [[TMP3]], align 8
+; CHECK-NEXT:    [[TMP19:%.*]] = load i64, ptr [[TMP5]], align 8
+; CHECK-NEXT:    [[TMP20:%.*]] = load i64, ptr [[TMP7]], align 8
+; CHECK-NEXT:    [[TMP21:%.*]] = load i64, ptr [[TMP9]], align 8
+; CHECK-NEXT:    [[TMP22:%.*]] = load i64, ptr [[TMP11]], align 8
+; CHECK-NEXT:    [[TMP23:%.*]] = load i64, ptr [[TMP13]], align 8
+; CHECK-NEXT:    [[TMP24:%.*]] = load i64, ptr [[TMP15]], align 8
+; CHECK-NEXT:    [[TMP25:%.*]] = load i64, ptr [[TMP17]], align 8
+; CHECK-NEXT:    [[TMP26:%.*]] = insertelement <8 x i64> poison, i64 [[TMP18]], i64 0
+; CHECK-NEXT:    [[TMP27:%.*]] = insertelement <8 x i64> [[TMP26]], i64 [[TMP19]], i64 1
+; CHECK-NEXT:    [[TMP28:%.*]] = insertelement <8 x i64> [[TMP27]], i64 [[TMP20]], i64 2
+; CHECK-NEXT:    [[TMP29:%.*]] = insertelement <8 x i64> [[TMP28]], i64 [[TMP21]], i64 3
+; CHECK-NEXT:    [[TMP30:%.*]] = insertelement <8 x i64> [[TMP29]], i64 [[TMP22]], i64 4
+; CHECK-NEXT:    [[TMP31:%.*]] = insertelement <8 x i64> [[TMP30]], i64 [[TMP23]], i64 5
+; CHECK-NEXT:    [[TMP32:%.*]] = insertelement <8 x i64> [[TMP31]], i64 [[TMP24]], i64 6
+; CHECK-NEXT:    [[TMP33:%.*]] = insertelement <8 x i64> [[TMP32]], i64 [[TMP25]], i64 7
+; CHECK-NEXT:    [[TMP34:%.*]] = add nsw <8 x i64> [[TMP33]], splat (i64 42)
+; CHECK-NEXT:    [[TMP35:%.*]] = getelementptr inbounds i64, ptr [[B]], i64 [[INDEX]]
+; CHECK-NEXT:    store <8 x i64> [[TMP34]], ptr [[TMP35]], align 8
 ; CHECK-NEXT:    [[INDEX_NEXT]] = add nuw i64 [[INDEX]], 8
-; CHECK-NEXT:    [[TMP6:%.*]] = icmp eq i64 [[INDEX_NEXT]], 1000
-; CHECK-NEXT:    br i1 [[TMP6]], label [[MIDDLE_BLOCK:%.*]], label [[VECTOR_BODY]], !llvm.loop [[LOOP5:![0-9]+]]
+; CHECK-NEXT:    [[VEC_IND_NEXT]] = add nsw <8 x i64> [[VEC_IND]], splat (i64 8)
+; CHECK-NEXT:    [[TMP36:%.*]] = icmp eq i64 [[INDEX_NEXT]], 1000
+; CHECK-NEXT:    br i1 [[TMP36]], label [[MIDDLE_BLOCK:%.*]], label [[VECTOR_BODY]], !llvm.loop [[LOOP5:![0-9]+]]
 ; CHECK:       middle.block:
 ; CHECK-NEXT:    br label [[EXIT:%.*]]
 ; CHECK:       exit:

>From b7ce01996f24bebbd79a4afdd92e3d3bd35c3f39 Mon Sep 17 00:00:00 2001
From: Ramkumar Ramachandra <artagnon at tenstorrent.com>
Date: Sun, 23 Aug 2026 22:09:15 +0100
Subject: [PATCH 3/4] [SCEV] Introduce willNotWrapByExtend, strip hack

---
 llvm/include/llvm/Analysis/ScalarEvolution.h | 14 ++--
 llvm/lib/Analysis/ScalarEvolution.cpp        | 83 ++++++++++++--------
 2 files changed, 55 insertions(+), 42 deletions(-)

diff --git a/llvm/include/llvm/Analysis/ScalarEvolution.h b/llvm/include/llvm/Analysis/ScalarEvolution.h
index 89404ecee80c0..4f59bc99cb697 100644
--- a/llvm/include/llvm/Analysis/ScalarEvolution.h
+++ b/llvm/include/llvm/Analysis/ScalarEvolution.h
@@ -704,14 +704,12 @@ class ScalarEvolution {
   /// scAddRecExpr. The result will be cached in HasRecMap.
   LLVM_ABI bool containsAddRecurrence(const SCEV *S);
 
-  /// Is operation with \p Opcode between \p LHS and \p RHS provably does not
-  /// have a signed/unsigned overflow (\p Signed)? If \p CtxI is specified, the
-  /// no-overflow fact should be true in the context of this instruction. If \p
-  /// Opcode is Instruction::PHI, the parameter \p L must be specified.
-  LLVM_ABI bool willNotOverflow(unsigned Opcode, bool Signed, const SCEV *LHS,
-                                const SCEV *RHS,
-                                const Instruction *CtxI = nullptr,
-                                const Loop *L = nullptr);
+  /// Is operation \p BinOp between \p LHS and \p RHS provably does not have
+  /// a signed/unsigned overflow (\p Signed)? If \p CtxI is specified, the
+  /// no-overflow fact should be true in the context of this instruction.
+  LLVM_ABI bool willNotOverflow(Instruction::BinaryOps BinOp, bool Signed,
+                                const SCEV *LHS, const SCEV *RHS,
+                                const Instruction *CtxI = nullptr);
 
   /// Parse NSW/NUW flags from add/sub/mul IR binary operation \p Op into
   /// SCEV no-wrap flags, and deduce flag[s] that aren't known yet.
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index 952b82b492ac1..b21753c542a4a 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -2287,54 +2287,66 @@ static bool CollectAddOperandsWithScales(SmallDenseMap<SCEVUse, APInt, 16> &M,
   return Interesting;
 }
 
-bool ScalarEvolution::willNotOverflow(unsigned Opcode, bool Signed,
+/// Checks that the application of binary function \p OperationFn to \p LHS and
+/// \p RHS does not wrap in the unsigned or signed (if \p Signed) manner.
+static bool
+willNotWrapByExtend(function_ref<const SCEV *(SCEVUse, SCEVUse)> OperationFn,
+                    const SCEV *LHS, const SCEV *RHS, bool Signed,
+                    ScalarEvolution *SE) {
+  auto *NarrowTy = LHS->getType();
+  auto *WideTy = IntegerType::get(NarrowTy->getContext(),
+                                  SE->getTypeSizeInBits(NarrowTy) * 2);
+
+  using ExtFnTy = const SCEV *(ScalarEvolution::*)(SCEVUse, Type *, unsigned);
+  std::function<const SCEV *(SCEVUse, Type *)> ExtensionFn =
+      Signed
+          ? bind_front(
+                bind_back<ExtFnTy>(&ScalarEvolution::getSignExtendExpr, 0), SE)
+          : bind_front(
+                bind_back<ExtFnTy>(&ScalarEvolution::getZeroExtendExpr, 0), SE);
+
+  // Check ExtensionFn(OperationFn(LHS, RHS)) == OperationFn(ExtensionFn(LHS),
+  // ExtensionFn(RHS))
+  const SCEV *A = ExtensionFn(OperationFn(LHS, RHS), WideTy);
+  const SCEV *LHSB = ExtensionFn(LHS, WideTy);
+  const SCEV *RHSB = ExtensionFn(RHS, WideTy);
+  const SCEV *B = OperationFn(LHSB, RHSB);
+  return A == B;
+}
+
+bool ScalarEvolution::willNotOverflow(Instruction::BinaryOps BinOp, bool Signed,
                                       const SCEV *LHS, const SCEV *RHS,
-                                      const Instruction *CtxI, const Loop *L) {
+                                      const Instruction *CtxI) {
   using OpFnTy = const SCEV *(ScalarEvolution::*)(SCEVUse, SCEVUse,
                                                   SCEV::NoWrapFlags, unsigned);
-  std::function<const SCEV *(SCEVUse, SCEVUse, SCEV::NoWrapFlags, unsigned)>
-      OperationFn;
-  switch (Opcode) {
+  std::function<const SCEV *(SCEVUse, SCEVUse)> OperationFn;
+  switch (BinOp) {
   default:
     llvm_unreachable("Unsupported binary op");
   case Instruction::Add:
-    OperationFn = bind_front<OpFnTy>(&ScalarEvolution::getAddExpr, this);
+    OperationFn = bind_front(
+        bind_back<OpFnTy>(&ScalarEvolution::getAddExpr, SCEV::FlagAnyWrap, 0),
+        this);
     break;
   case Instruction::Sub:
-    OperationFn = bind_front<OpFnTy>(&ScalarEvolution::getMinusSCEV, this);
+    OperationFn = bind_front(
+        bind_back<OpFnTy>(&ScalarEvolution::getMinusSCEV, SCEV::FlagAnyWrap, 0),
+        this);
     break;
   case Instruction::Mul:
-    OperationFn = bind_front<OpFnTy>(&ScalarEvolution::getMulExpr, this);
-    break;
-  case Instruction::PHI:
-    assert(L && "Loop argument must be given for PHI");
-    OperationFn = [&](SCEVUse LHS, SCEVUse RHS, SCEV::NoWrapFlags NW,
-                      unsigned) { return getAddRecExpr(LHS, RHS, L, NW); };
+    OperationFn = bind_front(
+        bind_back<OpFnTy>(&ScalarEvolution::getMulExpr, SCEV::FlagAnyWrap, 0),
+        this);
     break;
   }
 
-  using ExtFnTy = const SCEV *(ScalarEvolution::*)(SCEVUse, Type *, unsigned);
-  std::function<const SCEV *(SCEVUse, Type *, unsigned)> ExtensionFn =
-      Signed ? bind_front<ExtFnTy>(&ScalarEvolution::getSignExtendExpr, this)
-             : bind_front<ExtFnTy>(&ScalarEvolution::getZeroExtendExpr, this);
-
-  // Check ext(LHS op RHS) == ext(LHS) op ext(RHS)
-  auto *NarrowTy = cast<IntegerType>(LHS->getType());
-  auto *WideTy =
-      IntegerType::get(NarrowTy->getContext(), NarrowTy->getBitWidth() * 2);
-
-  const SCEV *A =
-      ExtensionFn(OperationFn(LHS, RHS, SCEV::FlagAnyWrap, 0), WideTy, 0);
-  const SCEV *LHSB = ExtensionFn(LHS, WideTy, 0);
-  const SCEV *RHSB = ExtensionFn(RHS, WideTy, 0);
-  const SCEV *B = OperationFn(LHSB, RHSB, SCEV::FlagAnyWrap, 0);
-  if (A == B)
+  if (willNotWrapByExtend(OperationFn, LHS, RHS, Signed, this))
     return true;
   // Can we use context to prove the fact we need?
   if (!CtxI)
     return false;
   // TODO: Support mul.
-  if (Opcode == Instruction::Mul || Opcode == Instruction::PHI)
+  if (BinOp == Instruction::Mul)
     return false;
   auto *RHSC = dyn_cast<SCEVConstant>(RHS);
   // TODO: Lift this limitation.
@@ -2342,7 +2354,7 @@ bool ScalarEvolution::willNotOverflow(unsigned Opcode, bool Signed,
     return false;
   APInt C = RHSC->getAPInt();
   unsigned NumBits = C.getBitWidth();
-  bool IsSub = (Opcode == Instruction::Sub);
+  bool IsSub = (BinOp == Instruction::Sub);
   bool IsNegativeConst = (Signed && C.isNegative());
   // Compute the direction and magnitude by which we need to check overflow.
   bool OverflowDown = IsSub ^ IsNegativeConst;
@@ -3485,9 +3497,12 @@ const SCEV *ScalarEvolution::getUDivExpr(SCEVUse LHS, SCEVUse RHS) {
           const APInt &StepInt = Step->getAPInt();
           const APInt &DivInt = RHSC->getAPInt();
           bool NoWrap = !StepInt.urem(DivInt) &&
-                        willNotOverflow(Instruction::PHI, /*Signed=*/false,
-                                        AR->getStart(), Step, /*CtxI=*/nullptr,
-                                        AR->getLoop());
+                        willNotWrapByExtend(
+                            [&](SCEVUse Start, SCEVUse Step) {
+                              return getAddRecExpr(Start, Step, AR->getLoop(),
+                                                   SCEV::FlagAnyWrap);
+                            },
+                            AR->getStart(), Step, /*Signed=*/false, this);
           if (!StepInt.urem(DivInt) && NoWrap) {
             SmallVector<SCEVUse, 4> Operands;
             for (const SCEV *Op : AR->operands())

>From 99f018f43dc138583180467e382e8aec58bf65cb Mon Sep 17 00:00:00 2001
From: Ramkumar Ramachandra <artagnon at tenstorrent.com>
Date: Wed, 26 Aug 2026 14:18:06 +0100
Subject: [PATCH 4/4] [SCEV] Use ExtendOpTraits in willNotWrapByExtend

---
 llvm/lib/Analysis/ScalarEvolution.cpp | 31 ++++++++++++---------------
 1 file changed, 14 insertions(+), 17 deletions(-)

diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index b21753c542a4a..009e909dd188e 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -2288,28 +2288,22 @@ static bool CollectAddOperandsWithScales(SmallDenseMap<SCEVUse, APInt, 16> &M,
 }
 
 /// Checks that the application of binary function \p OperationFn to \p LHS and
-/// \p RHS does not wrap in the unsigned or signed (if \p Signed) manner.
+/// \p RHS does not wrap in the unsigned or signed (depending on ExtendOpTy)
+/// manner.
+template <typename ExtendOpTy>
 static bool
 willNotWrapByExtend(function_ref<const SCEV *(SCEVUse, SCEVUse)> OperationFn,
-                    const SCEV *LHS, const SCEV *RHS, bool Signed,
-                    ScalarEvolution *SE) {
+                    const SCEV *LHS, const SCEV *RHS, ScalarEvolution *SE) {
   auto *NarrowTy = LHS->getType();
   auto *WideTy = IntegerType::get(NarrowTy->getContext(),
                                   SE->getTypeSizeInBits(NarrowTy) * 2);
 
-  using ExtFnTy = const SCEV *(ScalarEvolution::*)(SCEVUse, Type *, unsigned);
-  std::function<const SCEV *(SCEVUse, Type *)> ExtensionFn =
-      Signed
-          ? bind_front(
-                bind_back<ExtFnTy>(&ScalarEvolution::getSignExtendExpr, 0), SE)
-          : bind_front(
-                bind_back<ExtFnTy>(&ScalarEvolution::getZeroExtendExpr, 0), SE);
-
   // Check ExtensionFn(OperationFn(LHS, RHS)) == OperationFn(ExtensionFn(LHS),
   // ExtensionFn(RHS))
-  const SCEV *A = ExtensionFn(OperationFn(LHS, RHS), WideTy);
-  const SCEV *LHSB = ExtensionFn(LHS, WideTy);
-  const SCEV *RHSB = ExtensionFn(RHS, WideTy);
+  auto ExtensionFn = ExtendOpTraits<ExtendOpTy>::GetExtendExpr;
+  const SCEV *A = (SE->*ExtensionFn)(OperationFn(LHS, RHS), WideTy, 0);
+  const SCEV *LHSB = (SE->*ExtensionFn)(LHS, WideTy, 0);
+  const SCEV *RHSB = (SE->*ExtensionFn)(RHS, WideTy, 0);
   const SCEV *B = OperationFn(LHSB, RHSB);
   return A == B;
 }
@@ -2340,7 +2334,10 @@ bool ScalarEvolution::willNotOverflow(Instruction::BinaryOps BinOp, bool Signed,
     break;
   }
 
-  if (willNotWrapByExtend(OperationFn, LHS, RHS, Signed, this))
+  if (Signed
+          ? willNotWrapByExtend<SCEVSignExtendExpr>(OperationFn, LHS, RHS, this)
+          : willNotWrapByExtend<SCEVZeroExtendExpr>(OperationFn, LHS, RHS,
+                                                    this))
     return true;
   // Can we use context to prove the fact we need?
   if (!CtxI)
@@ -3497,12 +3494,12 @@ const SCEV *ScalarEvolution::getUDivExpr(SCEVUse LHS, SCEVUse RHS) {
           const APInt &StepInt = Step->getAPInt();
           const APInt &DivInt = RHSC->getAPInt();
           bool NoWrap = !StepInt.urem(DivInt) &&
-                        willNotWrapByExtend(
+                        willNotWrapByExtend<SCEVZeroExtendExpr>(
                             [&](SCEVUse Start, SCEVUse Step) {
                               return getAddRecExpr(Start, Step, AR->getLoop(),
                                                    SCEV::FlagAnyWrap);
                             },
-                            AR->getStart(), Step, /*Signed=*/false, this);
+                            AR->getStart(), Step, this);
           if (!StepInt.urem(DivInt) && NoWrap) {
             SmallVector<SCEVUse, 4> Operands;
             for (const SCEV *Op : AR->operands())



More information about the llvm-commits mailing list