[llvm] [SCEV] Make operand use flags part of expression's identity. (PR #216604)
via llvm-commits
llvm-commits at lists.llvm.org
Sun Aug 16 14:08:18 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-llvm-analysis
Author: Florian Hahn (fhahn)
<details>
<summary>Changes</summary>
Update hashing for SCEVNodes to include the use-specific operand flags. This makes sure expression with different operand flags distinct. Going forward, this ensures that various maps that cache SCEV expressions handle use-specific operands correctly.
---
Full diff: https://github.com/llvm/llvm-project/pull/216604.diff
3 Files Affected:
- (modified) llvm/include/llvm/Analysis/ScalarEvolution.h (+3-1)
- (modified) llvm/lib/Analysis/ScalarEvolution.cpp (+17-27)
- (modified) llvm/unittests/Analysis/ScalarEvolutionTest.cpp (+57)
``````````diff
diff --git a/llvm/include/llvm/Analysis/ScalarEvolution.h b/llvm/include/llvm/Analysis/ScalarEvolution.h
index 0d7f9ae298e2a..2d2b711549881 100644
--- a/llvm/include/llvm/Analysis/ScalarEvolution.h
+++ b/llvm/include/llvm/Analysis/ScalarEvolution.h
@@ -141,6 +141,9 @@ struct SCEVUseT : private PointerIntPair<SCEVPtrT, 2> {
/// operands.
bool isCanonical() const { return getCanonical() == getOpaqueValue(); }
+ /// Returns true if this use itself carries use-specific no-wrap flags.
+ bool hasUseFlags() const { return getOpaqueValue() != getPointer(); }
+
/// Return the canonical SCEV for this SCEVUse.
const SCEV *getCanonical() const;
@@ -2525,7 +2528,6 @@ class ScalarEvolution {
/// Look for a SCEV expression with type `SCEVType` and operands `Ops` in
/// `UniqueSCEVs`. Return if found, else nullptr.
- SCEV *findExistingSCEVInCache(SCEVTypes SCEVType, ArrayRef<const SCEV *> Ops);
SCEV *findExistingSCEVInCache(SCEVTypes SCEVType, ArrayRef<SCEVUse> Ops);
/// Get reachable blocks in this function, making limited use of SCEV
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index 27a1a20bcdf79..98e838997dc59 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -277,7 +277,7 @@ void SCEV::computeAndSetCanonical(ScalarEvolution &SE) {
SmallVector<SCEVUse, 4> CanonOps;
for (SCEVUse Op : operands()) {
CanonOps.push_back(Op->getCanonical());
- Changed |= CanonOps.back() != Op.getPointer();
+ Changed |= CanonOps.back() != Op.getPointer() || Op.hasUseFlags();
}
if (!Changed) {
@@ -3033,8 +3033,8 @@ const SCEV *ScalarEvolution::getOrCreateAddExpr(ArrayRef<SCEVUse> Ops,
SCEV::NoWrapFlags Flags) {
FoldingSetNodeID ID;
ID.AddInteger(scAddExpr);
- for (const SCEV *Op : Ops)
- ID.AddPointer(Op);
+ for (SCEVUse Op : Ops)
+ ID.AddPointer(Op.getOpaqueValue());
void *IP = nullptr;
SCEVAddExpr *S =
static_cast<SCEVAddExpr *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP));
@@ -3056,8 +3056,8 @@ const SCEV *ScalarEvolution::getOrCreateAddRecExpr(ArrayRef<SCEVUse> Ops,
SCEV::NoWrapFlags Flags) {
FoldingSetNodeID ID;
ID.AddInteger(scAddRecExpr);
- for (const SCEV *Op : Ops)
- ID.AddPointer(Op);
+ for (SCEVUse Op : Ops)
+ ID.AddPointer(Op.getOpaqueValue());
ID.AddPointer(L);
void *IP = nullptr;
SCEVAddRecExpr *S =
@@ -3080,8 +3080,8 @@ const SCEV *ScalarEvolution::getOrCreateMulExpr(ArrayRef<SCEVUse> Ops,
SCEV::NoWrapFlags Flags) {
FoldingSetNodeID ID;
ID.AddInteger(scMulExpr);
- for (const SCEV *Op : Ops)
- ID.AddPointer(Op);
+ for (SCEVUse Op : Ops)
+ ID.AddPointer(Op.getOpaqueValue());
void *IP = nullptr;
SCEVMulExpr *S =
static_cast<SCEVMulExpr *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP));
@@ -3493,8 +3493,8 @@ const SCEV *ScalarEvolution::getUDivExpr(SCEVUse LHS, SCEVUse RHS) {
FoldingSetNodeID ID;
ID.AddInteger(scUDivExpr);
- ID.AddPointer(LHS);
- ID.AddPointer(RHS);
+ ID.AddPointer(LHS.getOpaqueValue());
+ ID.AddPointer(RHS.getOpaqueValue());
void *IP = nullptr;
if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP))
return S;
@@ -3571,8 +3571,8 @@ const SCEV *ScalarEvolution::getUDivExpr(SCEVUse LHS, SCEVUse RHS) {
// already cached.
ID.clear();
ID.AddInteger(scUDivExpr);
- ID.AddPointer(LHS);
- ID.AddPointer(RHS);
+ ID.AddPointer(LHS.getOpaqueValue());
+ ID.AddPointer(RHS.getOpaqueValue());
IP = nullptr;
if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP))
return S;
@@ -3914,22 +3914,12 @@ const SCEV *ScalarEvolution::getGEPExpr(SCEVUse BaseExpr,
return GEPExpr;
}
-SCEV *ScalarEvolution::findExistingSCEVInCache(SCEVTypes SCEVType,
- ArrayRef<const SCEV *> Ops) {
- FoldingSetNodeID ID;
- ID.AddInteger(SCEVType);
- for (const SCEV *Op : Ops)
- ID.AddPointer(Op);
- void *IP = nullptr;
- return UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
-}
-
SCEV *ScalarEvolution::findExistingSCEVInCache(SCEVTypes SCEVType,
ArrayRef<SCEVUse> Ops) {
FoldingSetNodeID ID;
ID.AddInteger(SCEVType);
- for (const SCEV *Op : Ops)
- ID.AddPointer(Op);
+ for (SCEVUse Op : Ops)
+ ID.AddPointer(Op.getOpaqueValue());
void *IP = nullptr;
return UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
}
@@ -4050,8 +4040,8 @@ const SCEV *ScalarEvolution::getMinMaxExpr(SCEVTypes Kind,
// already have one, otherwise create a new one.
FoldingSetNodeID ID;
ID.AddInteger(Kind);
- for (const SCEV *Op : Ops)
- ID.AddPointer(Op);
+ for (SCEVUse Op : Ops)
+ ID.AddPointer(Op.getOpaqueValue());
void *IP = nullptr;
const SCEV *ExistingSCEV = UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
if (ExistingSCEV)
@@ -4437,8 +4427,8 @@ ScalarEvolution::getSequentialMinMaxExpr(SCEVTypes Kind,
// already have one, otherwise create a new one.
FoldingSetNodeID ID;
ID.AddInteger(Kind);
- for (const SCEV *Op : Ops)
- ID.AddPointer(Op);
+ for (SCEVUse Op : Ops)
+ ID.AddPointer(Op.getOpaqueValue());
void *IP = nullptr;
const SCEV *ExistingSCEV = UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
if (ExistingSCEV)
diff --git a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
index 621c4897a39d7..75f2609c0d93f 100644
--- a/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
+++ b/llvm/unittests/Analysis/ScalarEvolutionTest.cpp
@@ -2064,4 +2064,61 @@ TEST_F(ScalarEvolutionsTest, SimplifyICmpOperands) {
});
}
+// An operand carrying use-specific no-wrap flags makes the expression built
+// from it distinct from the one built from the bare operand
+TEST_F(ScalarEvolutionsTest, OperandUseFlagsArePartOfIdentity) {
+ LLVMContext C;
+ SMDiagnostic Err;
+ std::unique_ptr<Module> M = parseAssemblyString(
+ R"(define void @f(i32 %x, i32 %y, i1 %c) {
+ entry:
+ br label %loop
+ loop:
+ br i1 %c, label %loop, label %exit
+ exit:
+ ret void
+ })",
+ Err, C);
+
+ if (!M) {
+ Err.print("ScalarEvolutionTest", errs());
+ ASSERT_TRUE(M && "Could not parse module?");
+ }
+
+ runWithSE(*M, "f", [](Function &F, LoopInfo &LI, ScalarEvolution &SE) {
+ const SCEV *X = SE.getSCEV(getArgByName(F, "x"));
+ const SCEV *Y = SE.getSCEV(getArgByName(F, "y"));
+ SCEVUse FlaggedX(X, SCEV::FlagNUW);
+ ASSERT_FALSE(LI.empty());
+ const Loop *L = *LI.begin();
+
+ // Each builder keys its uniquing on the operand uses, so the flagged
+ // operand yields a different expression for every kind of node.
+ SCEVUse FlaggedAdd = SE.getAddExpr(FlaggedX, Y);
+ EXPECT_NE(FlaggedAdd, SE.getAddExpr(X, Y));
+ EXPECT_NE(SE.getMulExpr(FlaggedX, Y), SE.getMulExpr(X, Y));
+ EXPECT_NE(SE.getUDivExpr(FlaggedX, Y), SE.getUDivExpr(X, Y));
+ EXPECT_NE(SE.getUMaxExpr(FlaggedX, Y), SE.getUMaxExpr(X, Y));
+ EXPECT_NE(SE.getAddRecExpr(FlaggedX, Y, L, SCEV::FlagAnyWrap),
+ SE.getAddRecExpr(X, Y, L, SCEV::FlagAnyWrap));
+ SmallVector<SCEVUse, 2> FlaggedSeqOps = {FlaggedX, Y};
+ SmallVector<SCEVUse, 2> BareSeqOps = {X, Y};
+ EXPECT_NE(SE.getUMinExpr(FlaggedSeqOps, /*Sequential=*/true),
+ SE.getUMinExpr(BareSeqOps, /*Sequential=*/true));
+
+ EXPECT_EQ(FlaggedAdd->getCanonical(), SE.getAddExpr(X, Y));
+ EXPECT_EQ(SE.getUDivExpr(FlaggedX, Y)->getCanonical(),
+ SE.getUDivExpr(X, Y));
+
+ // The flagged use is the operand of the expression built from it, while the
+ // canonical form's operands are all bare.
+ SmallVector<SCEVUse> FlaggedOps;
+ copy_if(FlaggedAdd->operands(), std::back_inserter(FlaggedOps),
+ [](SCEVUse Op) { return Op.hasUseFlags(); });
+ ASSERT_EQ(FlaggedOps.size(), 1u);
+ EXPECT_EQ(FlaggedOps[0], FlaggedX);
+ EXPECT_TRUE(none_of(FlaggedAdd->getCanonical()->operands(),
+ [](SCEVUse Op) { return Op.hasUseFlags(); }));
+ });
+}
} // end namespace llvm
``````````
</details>
https://github.com/llvm/llvm-project/pull/216604
More information about the llvm-commits
mailing list