[llvm] [SCEV] Fix infinite recursion in A + zext(-A + B) -> zext(B) fold (PR #227690)
via llvm-commits
llvm-commits at lists.llvm.org
Wed Sep 30 05:53:45 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-llvm-transforms
Author: Marco Bartoli (wsxarcher)
<details>
<summary>Changes</summary>
The A + zext(-A + B) -> zext(B) fold in getAddExpr, added in d74d841b65dc, undoes the zext(C + X) -> zext(D) + zext((C - D) + X) split in getZeroExtendExprImpl. It called getTruncateExpr, getZeroExtendExpr and getAddExpr without passing Depth, so each round trip between the two folds reset the recursion depth to zero.
Pass Depth + 1 to all recursive calls in the fold, as the analogous C * zext(A + B) fold in getMulExpr already does, so the existing depth limits bound the recursion.
Fixes llvm/llvm-project#<!-- -->227664
Fixes llvm/llvm-project#<!-- -->184947
---
Full diff: https://github.com/llvm/llvm-project/pull/227690.diff
3 Files Affected:
- (modified) llvm/lib/Analysis/ScalarEvolution.cpp (+8-3)
- (added) llvm/test/Transforms/LoopStrengthReduce/zext-add-fold-depth-limit.ll (+106)
- (modified) llvm/unittests/Analysis/ScalarEvolutionTest.cpp (+45)
``````````diff
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index 5b8e30985ee7f..66f968233e7a9 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -2694,15 +2694,20 @@ SCEVUse ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
// Try to push the constant operand into a ZExt: A + zext (-A + B) -> zext
// (B), if trunc (A) + -A + B does not unsigned-wrap.
+ // This undoes the zext(C + X) -> zext(D) + zext((C - D) + X) split in
+ // getZeroExtendExprImpl, so propagate Depth to bound the mutual recursion
+ // when (C - D) + X is not folded (e.g. created past the depth limit).
const SCEVAddExpr *InnerAdd;
if (match(B, m_scev_ZExt(m_scev_Add(InnerAdd)))) {
- const SCEV *NarrowA = getTruncateExpr(A, InnerAdd->getType());
+ const SCEV *NarrowA = getTruncateExpr(A, InnerAdd->getType(), Depth + 1);
if (NarrowA == getNegativeSCEV(InnerAdd->getOperand(0)) &&
- getZeroExtendExpr(NarrowA, B->getType()) == A &&
+ getZeroExtendExpr(NarrowA, B->getType(), Depth + 1) == A &&
hasFlags(StrengthenNoWrapFlags(this, scAddExpr, {NarrowA, InnerAdd},
SCEV::FlagNone),
SCEV::FlagNUW)) {
- return getZeroExtendExpr(getAddExpr(NarrowA, InnerAdd), B->getType());
+ return getZeroExtendExpr(
+ getAddExpr(NarrowA, InnerAdd, SCEV::FlagNone, Depth + 1),
+ B->getType(), Depth + 1);
}
}
}
diff --git a/llvm/test/Transforms/LoopStrengthReduce/zext-add-fold-depth-limit.ll b/llvm/test/Transforms/LoopStrengthReduce/zext-add-fold-depth-limit.ll
new file mode 100644
index 0000000000000..dbd12433998f4
--- /dev/null
+++ b/llvm/test/Transforms/LoopStrengthReduce/zext-add-fold-depth-limit.ll
@@ -0,0 +1,106 @@
+; NOTE: Assertions have been autogenerated by utils/update_test_checks.py UTC_ARGS: --version 6
+; RUN: opt -loop-reduce -scalar-evolution-max-arith-depth=4 -S < %s | FileCheck %s
+
+; Make sure SCEV does not recurse infinitely between the
+; zext(C + X) -> zext(D) + zext((C - D) + X) fold in getZeroExtendExpr and the
+; A + zext(-A + B) -> zext(B) fold in getAddExpr, when (C - D) + X was created
+; unsimplified past the arithmetic depth limit.
+; https://github.com/llvm/llvm-project/issues/227664
+;
+; This uses the legacy pass manager on purpose: it visits the loops in reverse
+; program order, which is needed to reach the problematic SCEV cache state.
+; The new pass manager (-passes=loop-reduce) does not reproduce the crash.
+
+target datalayout = "e-m:o-p270:32:32-p271:32:32-p272:64:64-i64:64-i128:128-n32:64-S128-Fn32"
+
+define i1 @f(i1 %c, ptr %q) {
+; CHECK-LABEL: define i1 @f(
+; CHECK-SAME: i1 [[C:%.*]], ptr [[Q:%.*]]) {
+; CHECK-NEXT: [[ENTRY:.*:]]
+; CHECK-NEXT: [[START:%.*]] = select i1 [[C]], i64 0, i64 2
+; CHECK-NEXT: [[TMP0:%.*]] = sub i64 0, [[START]]
+; CHECK-NEXT: [[TMP1:%.*]] = mul nsw i64 [[TMP0]], -1
+; CHECK-NEXT: [[TMP2:%.*]] = sub i64 4, [[TMP0]]
+; CHECK-NEXT: [[TMP3:%.*]] = add nsw i64 [[TMP2]], -1
+; CHECK-NEXT: br label %[[L0:.*]]
+; CHECK: [[L0]]:
+; CHECK-NEXT: [[CMP0:%.*]] = icmp eq ptr null, null
+; CHECK-NEXT: br i1 [[CMP0]], label %[[L1_PREHEADER:.*]], label %[[L0]]
+; CHECK: [[L1_PREHEADER]]:
+; CHECK-NEXT: br label %[[L1:.*]]
+; CHECK: [[L1]]:
+; CHECK-NEXT: [[LSR_IV25:%.*]] = phi i64 [ -1, %[[L1_PREHEADER]] ], [ [[LSR_IV_NEXT26:%.*]], %[[L1]] ]
+; CHECK-NEXT: [[LSR_IV_NEXT26]] = add i64 [[LSR_IV25]], 1
+; CHECK-NEXT: br i1 [[C]], label %[[L2_PREHEADER:.*]], label %[[L1]]
+; CHECK: [[L2_PREHEADER]]:
+; CHECK-NEXT: [[INDVAR_LCSSA:%.*]] = phi i64 [ [[LSR_IV_NEXT26]], %[[L1]] ]
+; CHECK-NEXT: br label %[[L2:.*]]
+; CHECK: [[L2]]:
+; CHECK-NEXT: [[LSR_IV23:%.*]] = phi i64 [ -1, %[[L2_PREHEADER]] ], [ [[LSR_IV_NEXT24:%.*]], %[[L2]] ]
+; CHECK-NEXT: [[LSR_IV_NEXT24]] = add i64 [[LSR_IV23]], 1
+; CHECK-NEXT: br i1 [[C]], label %[[L3_PREHEADER:.*]], label %[[L2]]
+; CHECK: [[L3_PREHEADER]]:
+; CHECK-NEXT: [[INDVAR21_LCSSA:%.*]] = phi i64 [ [[LSR_IV_NEXT24]], %[[L2]] ]
+; CHECK-NEXT: [[TMP4:%.*]] = add i64 [[INDVAR_LCSSA]], [[TMP3]]
+; CHECK-NEXT: [[TMP5:%.*]] = add i64 [[INDVAR21_LCSSA]], [[TMP4]]
+; CHECK-NEXT: br label %[[L3:.*]]
+; CHECK: [[L3]]:
+; CHECK-NEXT: [[LSR_IV19:%.*]] = phi i64 [ 0, %[[L3_PREHEADER]] ], [ [[LSR_IV_NEXT20:%.*]], %[[L3]] ]
+; CHECK-NEXT: [[LSR_IV_NEXT20]] = add i64 [[LSR_IV19]], -1
+; CHECK-NEXT: [[EC3:%.*]] = icmp eq i64 [[LSR_IV_NEXT20]], 0
+; CHECK-NEXT: br i1 [[EC3]], label %[[L4_PREHEADER:.*]], label %[[L3]]
+; CHECK: [[L4_PREHEADER]]:
+; CHECK-NEXT: br label %[[L4:.*]]
+; CHECK: [[L4]]:
+; CHECK-NEXT: [[LSR_IV9:%.*]] = phi i64 [ 0, %[[L4_PREHEADER]] ], [ [[LSR_IV_NEXT10:%.*]], %[[L4]] ]
+; CHECK-NEXT: [[LSR_IV7:%.*]] = phi i64 [ [[TMP5]], %[[L4_PREHEADER]] ], [ [[LSR_IV_NEXT8:%.*]], %[[L4]] ]
+; CHECK-NEXT: [[LSR_IV_NEXT8]] = add i64 [[LSR_IV7]], 1
+; CHECK-NEXT: [[LSR_IV_NEXT10]] = add i64 [[LSR_IV9]], -1
+; CHECK-NEXT: [[EC4:%.*]] = icmp eq i64 [[LSR_IV_NEXT10]], 0
+; CHECK-NEXT: br i1 [[EC4]], label %[[EXIT:.*]], label %[[L4]]
+; CHECK: [[EXIT]]:
+; CHECK-NEXT: [[V:%.*]] = load i64, ptr [[Q]], align 8
+; CHECK-NEXT: [[R:%.*]] = icmp eq i64 [[V]], [[LSR_IV_NEXT8]]
+; CHECK-NEXT: ret i1 [[R]]
+;
+entry:
+ %start = select i1 %c, i64 0, i64 2
+ br label %l0
+
+l0:
+ %iv0 = phi i64 [ %iv0.next, %l0 ], [ %start, %entry ]
+ %iv0.next = add i64 %iv0, 1
+ %cmp0 = icmp eq ptr null, null
+ br i1 %cmp0, label %l1, label %l0
+
+l1:
+ %iv1 = phi i64 [ %iv1.next, %l1 ], [ %iv0.next, %l0 ]
+ %iv1.next = add i64 %iv1, 1
+ br i1 %c, label %l2, label %l1
+
+l2:
+ %iv2 = phi i64 [ %iv2.next, %l2 ], [ %iv1.next, %l1 ]
+ %iv2.next = add i64 %iv2, 1
+ br i1 %c, label %l3, label %l2
+
+l3:
+ %iv3 = phi i64 [ %iv3.next, %l3 ], [ %iv2.next, %l2 ]
+ %p3 = phi ptr [ %p3.next, %l3 ], [ null, %l2 ]
+ %iv3.next = add i64 %iv3, 1
+ %p3.next = getelementptr i8, ptr %p3, i64 1
+ %ec3 = icmp eq ptr %p3.next, null
+ br i1 %ec3, label %l4, label %l3
+
+l4:
+ %iv4 = phi i64 [ %iv4.next, %l4 ], [ %iv3.next, %l3 ]
+ %p4 = phi ptr [ %p4.next, %l4 ], [ null, %l3 ]
+ %iv4.next = add i64 %iv4, 1
+ %p4.next = getelementptr i8, ptr %p4, i64 1
+ %ec4 = icmp eq ptr %p4.next, null
+ br i1 %ec4, label %exit, label %l4
+
+exit:
+ %v = load i64, ptr %q, align 8
+ %r = icmp eq i64 %iv4.next, %v
+ ret i1 %r
+}
diff --git a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
index 1fd2eaa5eb72f..33213ed1f6a3a 100644
--- a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
+++ b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
@@ -2657,4 +2657,49 @@ TEST_F(ScalarEvolutionsTest, AddRecExprUseFlags) {
#endif
});
}
+
+// A + zext(-A + B) -> zext(B) in getAddExpr undoes the
+// zext(C + X) -> zext(D) + zext((C - D) + X) split in getZeroExtendExpr. If
+// (C - D) + X is an unsimplified add created past the depth limit, the two
+// folds used to recurse infinitely.
+// https://github.com/llvm/llvm-project/issues/227664
+TEST_F(ScalarEvolutionsTest, ZExtOfAddWithUnsimplifiedResidual) {
+ LLVMContext C;
+ SMDiagnostic Err;
+ std::unique_ptr<Module> M = parseAssemblyString(
+ R"(define void @f(i1 %c) {
+ entry:
+ %sel = select i1 %c, i64 0, i64 2
+ 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 *Sel = SE.getSCEV(getInstructionByName(F, "sel"));
+ Type *I64 = Sel->getType();
+ Type *I128 = Type::getInt128Ty(F.getContext());
+ const SCEV *MinusOne = SE.getMinusOne(I64);
+
+ // Create Y = (-1 + (32 + %sel)) and (-1 + Y) past the arithmetic depth
+ // limit, so they are neither flattened nor constant folded.
+ SmallVector<SCEVUse, 2> Ops = {MinusOne,
+ SE.getAddExpr(SE.getConstant(I64, 32), Sel)};
+ const SCEV *Y = SE.getAddExpr(Ops, SCEV::FlagNone, /*Depth=*/100);
+ Ops = {MinusOne, Y};
+ const SCEV *YMinusOne = SE.getAddExpr(Ops, SCEV::FlagNone, /*Depth=*/100);
+ ASSERT_TRUE(
+ match(YMinusOne, m_scev_Add(m_scev_AllOnes(), m_scev_Specific(Y))));
+
+ // zext(Y) splits off D = 1 and gets 1 + zext(-1 + Y) from the cached
+ // (-1 + Y), which folds back to zext(Y).
+ const SCEV *ZExt = SE.getZeroExtendExpr(Y, I128);
+ EXPECT_EQ(ZExt->getType(), I128);
+ });
+}
} // end namespace llvm
``````````
</details>
https://github.com/llvm/llvm-project/pull/227690
More information about the llvm-commits
mailing list