[llvm] [SCEV] Look up uniqued nodes by using interned profile. (NFC) (PR #217571)

Florian Hahn via llvm-commits llvm-commits at lists.llvm.org
Fri Aug 21 06:36:04 PDT 2026


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

>From a81b0725e80515be69519ab91418b56c7f1dad1f Mon Sep 17 00:00:00 2001
From: Florian Hahn <flo at fhahn.com>
Date: Mon, 17 Aug 2026 14:30:06 +0100
Subject: [PATCH 1/2] [SCEV] Look up uniqued nodes by using interned profile.
 (NFC)

Each SCEV node already stores an interned profile/hash. Add an overload
for FindNodeOrInsertPos which takes a match callback; SCEV uses this to
match against the interned ID, which saves unnecessary pointer walks to
hash/compare.
---
 llvm/include/llvm/ADT/FoldingSet.h           | 46 ++++++++++++++++
 llvm/include/llvm/Analysis/ScalarEvolution.h |  6 +++
 llvm/lib/Analysis/ScalarEvolution.cpp        | 57 ++++++++++++--------
 llvm/lib/Support/FoldingSet.cpp              | 36 +++----------
 4 files changed, 93 insertions(+), 52 deletions(-)

diff --git a/llvm/include/llvm/ADT/FoldingSet.h b/llvm/include/llvm/ADT/FoldingSet.h
index ab4fa2712d4a5..5187fa6022db7 100644
--- a/llvm/include/llvm/ADT/FoldingSet.h
+++ b/llvm/include/llvm/ADT/FoldingSet.h
@@ -296,6 +296,8 @@ class FoldingSetNodeID {
 /// facilitate node removal.
 ///
 class FoldingSetBase {
+  friend class FoldingSetIteratorImpl;
+
 protected:
   /// Array of bucket chains.
   void **Buckets;
@@ -374,6 +376,25 @@ class FoldingSetBase {
   void GrowBucketCount(unsigned NewBucketCount, const FoldingSetInfo &Info);
 
 protected:
+  /// Return the hash bucket for \p Hash.
+  void **getBucketFor(unsigned Hash) const {
+    // NumBuckets is always a power of 2.
+    return Buckets + (Hash & (NumBuckets - 1));
+  }
+
+  /// In order to save space, each bucket is a singly-linked-list. In order to
+  /// make deletion more efficient, we make the list circular, so we can delete
+  /// a node without computing its hash. The problem with this is that the start
+  /// of the hash buckets are not Nodes. If \p NextInBucketPtr is a bucket
+  /// pointer, this method returns null.
+  static Node *GetNextPtr(void *NextInBucketPtr) {
+    // The low bit is set if this is the pointer back to the bucket.
+    if (reinterpret_cast<intptr_t>(NextInBucketPtr) & 1)
+      return nullptr;
+
+    return static_cast<Node *>(NextInBucketPtr);
+  }
+
   // The below methods are protected to encourage subclasses to provide a more
   // type-safe API.
 
@@ -504,6 +525,31 @@ class FoldingSetImpl : public FoldingSetBase, public Trait::ContextStorage {
   const_iterator begin() const { return const_iterator(Buckets); }
   const_iterator end() const { return const_iterator(Buckets + NumBuckets); }
 
+  /// Look up the node matching \p ID, using \p IsMatch to compare a candidate
+  /// node against it. If there is no such node, return null and set
+  /// \p InsertPos to the insertion token to pass to InsertNode(); the token is
+  /// only valid until the set is modified.
+  ///
+  /// This is an alternative to the FoldingSetInfo-based overload for clients
+  /// whose nodes can be compared against an ID directly.
+  template <typename FnT>
+  T *FindNodeOrInsertPos(const FoldingSetNodeID &ID, void *&InsertPos,
+                         FnT IsMatch) {
+    void **Bucket = getBucketFor(ID.ComputeHash());
+    for (void *Probe = *Bucket; Node *N = GetNextPtr(Probe);
+         Probe = N->getNextInBucket()) {
+      T *TN = static_cast<T *>(N);
+      if (IsMatch(*TN)) {
+        InsertPos = nullptr;
+        return TN;
+      }
+    }
+
+    // Didn't find the node, return null with the bucket as the InsertPos.
+    InsertPos = Bucket;
+    return nullptr;
+  }
+
   /// Grow the number of buckets so that we can hold at least \p EltCount
   /// nodes before rebucketing. May allocate more space than requested.
   void reserve(unsigned EltCount) {
diff --git a/llvm/include/llvm/Analysis/ScalarEvolution.h b/llvm/include/llvm/Analysis/ScalarEvolution.h
index d17c21ea3401e..b4ee92746af46 100644
--- a/llvm/include/llvm/Analysis/ScalarEvolution.h
+++ b/llvm/include/llvm/Analysis/ScalarEvolution.h
@@ -326,6 +326,9 @@ class SCEV : public FoldingSetNode {
   /// stream.  This should really only be used for debugging purposes.
   LLVM_ABI void print(raw_ostream &OS) const;
 
+  /// Return true if \p ID is this node's interned uniquing profile.
+  bool hasProfile(const FoldingSetNodeID &ID) const { return ID == FastID; }
+
   /// This method is used for debugging.
   LLVM_ABI void dump() const;
 
@@ -399,6 +402,9 @@ class SCEVPredicate : public FoldingSetNode {
 
   SCEVPredicateKind getKind() const { return Kind; }
 
+  /// Return true if \p ID is this node's interned uniquing profile.
+  bool hasProfile(const FoldingSetNodeID &ID) const { return ID == FastID; }
+
   /// Returns the estimated complexity of this predicate.  This is roughly
   /// measured in the number of run-time checks required.
   virtual unsigned getComplexity() const { return 1; }
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index ebd58c825faa5..755a661a2267f 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -506,6 +506,16 @@ bool SCEVCouldNotCompute::classof(const SCEV *S) {
   return S->getSCEVType() == scCouldNotCompute;
 }
 
+/// Look up the node with profile \p ID in \p Set. On a miss, set \p InsertPos
+/// for a subsequent Set.InsertNode() and return nullptr.
+template <typename NodeTy>
+static NodeTy *findUniqued(FoldingSet<NodeTy> &Set, const FoldingSetNodeID &ID,
+                           void *&InsertPos) {
+  // The nodes intern their profile, so compare against it directly.
+  return Set.FindNodeOrInsertPos(
+      ID, InsertPos, [&ID](const NodeTy &N) { return N.hasProfile(ID); });
+}
+
 const SCEV *ScalarEvolution::getConstant(ConstantInt *V) {
   auto &Entry = ConstantSCEVs[V];
   if (Entry)
@@ -516,7 +526,7 @@ const SCEV *ScalarEvolution::getConstant(ConstantInt *V) {
   ID.AddPointer(V);
   void *IP = nullptr;
   if (SCEVConstant *S =
-          static_cast<SCEVConstant *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP)))
+          static_cast<SCEVConstant *>(findUniqued(UniqueSCEVs, ID, IP)))
     return Entry = S;
   SCEVConstant *S =
       new (SCEVAllocator) SCEVConstant(ID.Intern(SCEVAllocator), V);
@@ -543,7 +553,7 @@ const SCEV *ScalarEvolution::getVScale(Type *Ty) {
   ID.AddInteger(scVScale);
   ID.AddPointer(Ty);
   void *IP = nullptr;
-  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP))
+  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
     return S;
   SCEV *S = new (SCEVAllocator) SCEVVScale(ID.Intern(SCEVAllocator), Ty);
   UniqueSCEVs.InsertNode(S, IP);
@@ -1142,7 +1152,7 @@ const SCEV *ScalarEvolution::getPtrToAddrExpr(const SCEV *Op) {
         ID.AddPointer(U);
         ID.AddPointer(Ty);
         void *IP = nullptr;
-        if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP))
+        if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
           return S;
         SCEV *S = new (SCEVAllocator)
             SCEVPtrToAddrExpr(ID.Intern(SCEVAllocator), U, Ty);
@@ -1171,7 +1181,8 @@ const SCEV *ScalarEvolution::getTruncateExpr(SCEVUse Op, Type *Ty,
   ID.AddPointer(Op.getOpaqueValue());
   ID.AddPointer(Ty);
   void *IP = nullptr;
-  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) return S;
+  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
+    return S;
 
   // Fold if the operand is constant.
   if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(Op))
@@ -1225,7 +1236,7 @@ const SCEV *ScalarEvolution::getTruncateExpr(SCEVUse Op, Type *Ty,
     // Although we checked in the beginning that ID is not in the cache, it is
     // possible that during recursion and different modification ID was inserted
     // into the cache. So if we find it, just return it.
-    if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP))
+    if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
       return S;
   }
 
@@ -1497,7 +1508,7 @@ bool ScalarEvolution::proveNoWrapByVaryingStart(const SCEV *Start,
     ID.AddPointer(L);
     void *IP = nullptr;
     const auto *PreAR =
-      static_cast<SCEVAddRecExpr *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP));
+        static_cast<SCEVAddRecExpr *>(findUniqued(UniqueSCEVs, ID, IP));
 
     // Give up if we don't already have the add recurrence we need because
     // actually constructing an add recurrence is relatively expensive.
@@ -1628,7 +1639,8 @@ const SCEV *ScalarEvolution::getZeroExtendExprImpl(SCEVUse Op, Type *Ty,
   ID.AddPointer(Op.getOpaqueValue());
   ID.AddPointer(Ty);
   void *IP = nullptr;
-  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) return S;
+  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
+    return S;
   if (Depth > MaxCastDepth) {
     SCEV *S = new (SCEVAllocator) SCEVZeroExtendExpr(ID.Intern(SCEVAllocator),
                                                      Op, Ty);
@@ -1914,7 +1926,8 @@ const SCEV *ScalarEvolution::getZeroExtendExprImpl(SCEVUse Op, Type *Ty,
 
   // The cast wasn't folded; create an explicit cast node.
   // Recompute the insert position, as it may have been invalidated.
-  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) return S;
+  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
+    return S;
   SCEV *S = new (SCEVAllocator) SCEVZeroExtendExpr(ID.Intern(SCEVAllocator),
                                                    Op, Ty);
   UniqueSCEVs.InsertNode(S, IP);
@@ -1983,7 +1996,8 @@ const SCEV *ScalarEvolution::getSignExtendExprImpl(SCEVUse Op, Type *Ty,
   ID.AddPointer(Op.getOpaqueValue());
   ID.AddPointer(Ty);
   void *IP = nullptr;
-  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) return S;
+  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
+    return S;
   // Limit recursion depth.
   if (Depth > MaxCastDepth) {
     SCEV *S = new (SCEVAllocator) SCEVSignExtendExpr(ID.Intern(SCEVAllocator),
@@ -2177,7 +2191,8 @@ const SCEV *ScalarEvolution::getSignExtendExprImpl(SCEVUse Op, Type *Ty,
 
   // The cast wasn't folded; create an explicit cast node.
   // Recompute the insert position, as it may have been invalidated.
-  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) return S;
+  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
+    return S;
   SCEV *S = new (SCEVAllocator) SCEVSignExtendExpr(ID.Intern(SCEVAllocator),
                                                    Op, Ty);
   UniqueSCEVs.InsertNode(S, IP);
@@ -3032,8 +3047,7 @@ const SCEV *ScalarEvolution::getOrCreateAddExpr(ArrayRef<SCEVUse> Ops,
   for (SCEVUse Op : Ops)
     ID.AddPointer(Op.getOpaqueValue());
   void *IP = nullptr;
-  SCEVAddExpr *S =
-      static_cast<SCEVAddExpr *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP));
+  SCEVAddExpr *S = static_cast<SCEVAddExpr *>(findUniqued(UniqueSCEVs, ID, IP));
   if (!S) {
     SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
     llvm::uninitialized_copy(Ops, O);
@@ -3057,7 +3071,7 @@ const SCEV *ScalarEvolution::getOrCreateAddRecExpr(ArrayRef<SCEVUse> Ops,
   ID.AddPointer(L);
   void *IP = nullptr;
   SCEVAddRecExpr *S =
-      static_cast<SCEVAddRecExpr *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP));
+      static_cast<SCEVAddRecExpr *>(findUniqued(UniqueSCEVs, ID, IP));
   if (!S) {
     SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
     llvm::uninitialized_copy(Ops, O);
@@ -3079,8 +3093,7 @@ const SCEV *ScalarEvolution::getOrCreateMulExpr(ArrayRef<SCEVUse> Ops,
   for (SCEVUse Op : Ops)
     ID.AddPointer(Op.getOpaqueValue());
   void *IP = nullptr;
-  SCEVMulExpr *S =
-    static_cast<SCEVMulExpr *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP));
+  SCEVMulExpr *S = static_cast<SCEVMulExpr *>(findUniqued(UniqueSCEVs, ID, IP));
   if (!S) {
     SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
     llvm::uninitialized_copy(Ops, O);
@@ -3100,7 +3113,7 @@ const SCEV *ScalarEvolution::getOrCreateUDivExpr(SCEVUse LHS, SCEVUse RHS) {
   ID.AddPointer(LHS.getOpaqueValue());
   ID.AddPointer(RHS.getOpaqueValue());
   void *IP = nullptr;
-  SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
+  SCEV *S = findUniqued(UniqueSCEVs, ID, IP);
   if (!S) {
     S = new (SCEVAllocator) SCEVUDivExpr(ID.Intern(SCEVAllocator), LHS, RHS);
     UniqueSCEVs.InsertNode(S, IP);
@@ -3890,7 +3903,7 @@ SCEV *ScalarEvolution::findExistingSCEVInCache(SCEVTypes SCEVType,
   for (SCEVUse Op : Ops)
     ID.AddPointer(Op.getOpaqueValue());
   void *IP = nullptr;
-  return UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
+  return findUniqued(UniqueSCEVs, ID, IP);
 }
 
 const SCEV *ScalarEvolution::getAbsExpr(const SCEV *Op, bool IsNSW) {
@@ -4012,7 +4025,7 @@ const SCEV *ScalarEvolution::getMinMaxExpr(SCEVTypes Kind,
   for (SCEVUse Op : Ops)
     ID.AddPointer(Op.getOpaqueValue());
   void *IP = nullptr;
-  const SCEV *ExistingSCEV = UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
+  const SCEV *ExistingSCEV = findUniqued(UniqueSCEVs, ID, IP);
   if (ExistingSCEV)
     return ExistingSCEV;
   SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
@@ -4399,7 +4412,7 @@ ScalarEvolution::getSequentialMinMaxExpr(SCEVTypes Kind,
   for (SCEVUse Op : Ops)
     ID.AddPointer(Op.getOpaqueValue());
   void *IP = nullptr;
-  const SCEV *ExistingSCEV = UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
+  const SCEV *ExistingSCEV = findUniqued(UniqueSCEVs, ID, IP);
   if (ExistingSCEV)
     return ExistingSCEV;
 
@@ -4491,7 +4504,7 @@ const SCEV *ScalarEvolution::getUnknown(Value *V) {
   ID.AddInteger(scUnknown);
   ID.AddPointer(V);
   void *IP = nullptr;
-  if (SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) {
+  if (SCEV *S = findUniqued(UniqueSCEVs, ID, IP)) {
     assert(cast<SCEVUnknown>(S)->getValue() == V &&
            "Stale SCEVUnknown in uniquing map!");
     return S;
@@ -15168,7 +15181,7 @@ ScalarEvolution::getComparePredicate(const ICmpInst::Predicate Pred,
   ID.AddPointer(LHS);
   ID.AddPointer(RHS);
   void *IP = nullptr;
-  if (const auto *S = UniquePreds.FindNodeOrInsertPos(ID, IP))
+  if (const auto *S = findUniqued(UniquePreds, ID, IP))
     return S;
   SCEVComparePredicate *Eq = new (SCEVAllocator)
     SCEVComparePredicate(ID.Intern(SCEVAllocator), Pred, LHS, RHS);
@@ -15185,7 +15198,7 @@ const SCEVPredicate *ScalarEvolution::getWrapPredicate(
   ID.AddPointer(AR);
   ID.AddInteger(AddedFlags);
   void *IP = nullptr;
-  if (const auto *S = UniquePreds.FindNodeOrInsertPos(ID, IP))
+  if (const auto *S = findUniqued(UniquePreds, ID, IP))
     return S;
   auto *OF = new (SCEVAllocator)
       SCEVWrapPredicate(ID.Intern(SCEVAllocator), AR, AddedFlags);
diff --git a/llvm/lib/Support/FoldingSet.cpp b/llvm/lib/Support/FoldingSet.cpp
index d9ae1aca5fc4a..e082984bab27f 100644
--- a/llvm/lib/Support/FoldingSet.cpp
+++ b/llvm/lib/Support/FoldingSet.cpp
@@ -133,20 +133,6 @@ FoldingSetNodeID::Intern(BumpPtrAllocator &Allocator) const {
 //===----------------------------------------------------------------------===//
 /// Helper functions for FoldingSetBase.
 
-/// GetNextPtr - In order to save space, each bucket is a
-/// singly-linked-list. In order to make deletion more efficient, we make
-/// the list circular, so we can delete a node without computing its hash.
-/// The problem with this is that the start of the hash buckets are not
-/// Nodes. If NextInBucketPtr is a bucket pointer, this method returns null:
-/// use GetBucketPtr when this happens.
-static FoldingSetBase::Node *GetNextPtr(void *NextInBucketPtr) {
-  // The low bit is set if this is the pointer back to the bucket.
-  if (reinterpret_cast<intptr_t>(NextInBucketPtr) & 1)
-    return nullptr;
-
-  return static_cast<FoldingSetBase::Node *>(NextInBucketPtr);
-}
-
 /// GetBucketPtr - Provides a casting of a bucket pointer for isNode
 /// testing.
 static void **GetBucketPtr(void *NextInBucketPtr) {
@@ -155,14 +141,6 @@ static void **GetBucketPtr(void *NextInBucketPtr) {
   return reinterpret_cast<void **>(Ptr & ~intptr_t(1));
 }
 
-/// GetBucketFor - Hash the specified node ID and return the hash bucket for
-/// the specified ID.
-static void **GetBucketFor(unsigned Hash, void **Buckets, unsigned NumBuckets) {
-  // NumBuckets is always a power of 2.
-  unsigned BucketNum = Hash & (NumBuckets - 1);
-  return Buckets + BucketNum;
-}
-
 /// AllocateBuckets - Allocate initialized bucket memory.
 static void **AllocateBuckets(unsigned NumBuckets) {
   void **Buckets =
@@ -234,8 +212,7 @@ void FoldingSetBase::GrowBucketCount(unsigned NewBucketCount,
       // Insert the node into the new bucket, after recomputing the hash.
       Tmp.InsertNode(
           NodeInBucket,
-          GetBucketFor(Info.ComputeNodeHash(this, NodeInBucket, TempID),
-                       Tmp.Buckets, Tmp.NumBuckets),
+          Tmp.getBucketFor(Info.ComputeNodeHash(this, NodeInBucket, TempID)),
           Info);
       TempID.clear();
     }
@@ -256,7 +233,7 @@ void FoldingSetBase::reserve(unsigned EltCount, const FoldingSetInfo &Info) {
 FoldingSetBase::Node *FoldingSetBase::FindNodeOrInsertPos(
     const FoldingSetNodeID &ID, void *&InsertPos, const FoldingSetInfo &Info) {
   unsigned IDHash = ID.ComputeHash();
-  void **Bucket = GetBucketFor(IDHash, Buckets, NumBuckets);
+  void **Bucket = getBucketFor(IDHash);
   void *Probe = *Bucket;
 
   InsertPos = nullptr;
@@ -282,8 +259,7 @@ void FoldingSetBase::InsertNode(Node *N, void *InsertPos,
   if (NumNodes + 1 > capacity()) {
     GrowBucketCount(NumBuckets * 2, Info);
     FoldingSetNodeID TempID;
-    InsertPos = GetBucketFor(Info.ComputeNodeHash(this, N, TempID), Buckets,
-                             NumBuckets);
+    InsertPos = getBucketFor(Info.ComputeNodeHash(this, N, TempID));
   }
 
   ++NumNodes;
@@ -360,7 +336,7 @@ FoldingSetBase::GetOrInsertNode(Node *N, const FoldingSetInfo &Info) {
 FoldingSetIteratorImpl::FoldingSetIteratorImpl(void **Bucket) {
   // Skip to the first non-null non-self-cycle bucket.
   while (*Bucket != reinterpret_cast<void *>(-1) &&
-         (!*Bucket || !GetNextPtr(*Bucket)))
+         (!*Bucket || !FoldingSetBase::GetNextPtr(*Bucket)))
     ++Bucket;
 
   NodePtr = static_cast<FoldingSetNode *>(*Bucket);
@@ -370,7 +346,7 @@ void FoldingSetIteratorImpl::advance() {
   // If there is another link within this bucket, go to it.
   void *Probe = NodePtr->getNextInBucket();
 
-  if (FoldingSetNode *NextNodeInBucket = GetNextPtr(Probe))
+  if (FoldingSetNode *NextNodeInBucket = FoldingSetBase::GetNextPtr(Probe))
     NodePtr = NextNodeInBucket;
   else {
     // Otherwise, this is the last link in this bucket.
@@ -380,7 +356,7 @@ void FoldingSetIteratorImpl::advance() {
     do {
       ++Bucket;
     } while (*Bucket != reinterpret_cast<void *>(-1) &&
-             (!*Bucket || !GetNextPtr(*Bucket)));
+             (!*Bucket || !FoldingSetBase::GetNextPtr(*Bucket)));
 
     NodePtr = static_cast<FoldingSetNode *>(*Bucket);
   }

>From a42894bef3ecc09b380d6bd41d349b376abd0d1a Mon Sep 17 00:00:00 2001
From: Florian Hahn <flo at fhahn.com>
Date: Fri, 21 Aug 2026 14:31:54 +0100
Subject: [PATCH 2/2] !fixup inline-only

---
 llvm/include/llvm/ADT/FoldingSet.h           | 50 ++++++++---------
 llvm/include/llvm/Analysis/ScalarEvolution.h |  6 ---
 llvm/lib/Analysis/ScalarEvolution.cpp        | 57 ++++++++------------
 llvm/lib/Support/FoldingSet.cpp              | 22 --------
 4 files changed, 44 insertions(+), 91 deletions(-)

diff --git a/llvm/include/llvm/ADT/FoldingSet.h b/llvm/include/llvm/ADT/FoldingSet.h
index 5187fa6022db7..16061e9e949dc 100644
--- a/llvm/include/llvm/ADT/FoldingSet.h
+++ b/llvm/include/llvm/ADT/FoldingSet.h
@@ -412,9 +412,28 @@ class FoldingSetBase {
 
   /// Look up the node specified by ID.  If it exists, return it.  If not,
   /// return the insertion token that will make insertion faster.
-  LLVM_ABI Node *FindNodeOrInsertPos(const FoldingSetNodeID &ID,
-                                     void *&InsertPos,
-                                     const FoldingSetInfo &Info);
+  ///
+  /// This is defined here rather than out-of-line, so the node comparison,
+  /// which is trivial for some clients, can be inlined into the caller.
+  Node *FindNodeOrInsertPos(const FoldingSetNodeID &ID, void *&InsertPos,
+                            const FoldingSetInfo &Info) {
+    unsigned IDHash = ID.ComputeHash();
+    void **Bucket = getBucketFor(IDHash);
+
+    InsertPos = nullptr;
+
+    FoldingSetNodeID TempID;
+    for (void *Probe = *Bucket; Node *N = GetNextPtr(Probe);
+         Probe = N->getNextInBucket()) {
+      if (Info.NodeEquals(this, N, ID, IDHash, TempID))
+        return N;
+      TempID.clear();
+    }
+
+    // Didn't find the node, return null with the bucket as the InsertPos.
+    InsertPos = Bucket;
+    return nullptr;
+  }
 
   /// Insert the specified node into the folding set, knowing that
   /// it is not already in the folding set.  InsertPos must be obtained from
@@ -525,31 +544,6 @@ class FoldingSetImpl : public FoldingSetBase, public Trait::ContextStorage {
   const_iterator begin() const { return const_iterator(Buckets); }
   const_iterator end() const { return const_iterator(Buckets + NumBuckets); }
 
-  /// Look up the node matching \p ID, using \p IsMatch to compare a candidate
-  /// node against it. If there is no such node, return null and set
-  /// \p InsertPos to the insertion token to pass to InsertNode(); the token is
-  /// only valid until the set is modified.
-  ///
-  /// This is an alternative to the FoldingSetInfo-based overload for clients
-  /// whose nodes can be compared against an ID directly.
-  template <typename FnT>
-  T *FindNodeOrInsertPos(const FoldingSetNodeID &ID, void *&InsertPos,
-                         FnT IsMatch) {
-    void **Bucket = getBucketFor(ID.ComputeHash());
-    for (void *Probe = *Bucket; Node *N = GetNextPtr(Probe);
-         Probe = N->getNextInBucket()) {
-      T *TN = static_cast<T *>(N);
-      if (IsMatch(*TN)) {
-        InsertPos = nullptr;
-        return TN;
-      }
-    }
-
-    // Didn't find the node, return null with the bucket as the InsertPos.
-    InsertPos = Bucket;
-    return nullptr;
-  }
-
   /// Grow the number of buckets so that we can hold at least \p EltCount
   /// nodes before rebucketing. May allocate more space than requested.
   void reserve(unsigned EltCount) {
diff --git a/llvm/include/llvm/Analysis/ScalarEvolution.h b/llvm/include/llvm/Analysis/ScalarEvolution.h
index b4ee92746af46..d17c21ea3401e 100644
--- a/llvm/include/llvm/Analysis/ScalarEvolution.h
+++ b/llvm/include/llvm/Analysis/ScalarEvolution.h
@@ -326,9 +326,6 @@ class SCEV : public FoldingSetNode {
   /// stream.  This should really only be used for debugging purposes.
   LLVM_ABI void print(raw_ostream &OS) const;
 
-  /// Return true if \p ID is this node's interned uniquing profile.
-  bool hasProfile(const FoldingSetNodeID &ID) const { return ID == FastID; }
-
   /// This method is used for debugging.
   LLVM_ABI void dump() const;
 
@@ -402,9 +399,6 @@ class SCEVPredicate : public FoldingSetNode {
 
   SCEVPredicateKind getKind() const { return Kind; }
 
-  /// Return true if \p ID is this node's interned uniquing profile.
-  bool hasProfile(const FoldingSetNodeID &ID) const { return ID == FastID; }
-
   /// Returns the estimated complexity of this predicate.  This is roughly
   /// measured in the number of run-time checks required.
   virtual unsigned getComplexity() const { return 1; }
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index 755a661a2267f..ebd58c825faa5 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -506,16 +506,6 @@ bool SCEVCouldNotCompute::classof(const SCEV *S) {
   return S->getSCEVType() == scCouldNotCompute;
 }
 
-/// Look up the node with profile \p ID in \p Set. On a miss, set \p InsertPos
-/// for a subsequent Set.InsertNode() and return nullptr.
-template <typename NodeTy>
-static NodeTy *findUniqued(FoldingSet<NodeTy> &Set, const FoldingSetNodeID &ID,
-                           void *&InsertPos) {
-  // The nodes intern their profile, so compare against it directly.
-  return Set.FindNodeOrInsertPos(
-      ID, InsertPos, [&ID](const NodeTy &N) { return N.hasProfile(ID); });
-}
-
 const SCEV *ScalarEvolution::getConstant(ConstantInt *V) {
   auto &Entry = ConstantSCEVs[V];
   if (Entry)
@@ -526,7 +516,7 @@ const SCEV *ScalarEvolution::getConstant(ConstantInt *V) {
   ID.AddPointer(V);
   void *IP = nullptr;
   if (SCEVConstant *S =
-          static_cast<SCEVConstant *>(findUniqued(UniqueSCEVs, ID, IP)))
+          static_cast<SCEVConstant *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP)))
     return Entry = S;
   SCEVConstant *S =
       new (SCEVAllocator) SCEVConstant(ID.Intern(SCEVAllocator), V);
@@ -553,7 +543,7 @@ const SCEV *ScalarEvolution::getVScale(Type *Ty) {
   ID.AddInteger(scVScale);
   ID.AddPointer(Ty);
   void *IP = nullptr;
-  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
+  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP))
     return S;
   SCEV *S = new (SCEVAllocator) SCEVVScale(ID.Intern(SCEVAllocator), Ty);
   UniqueSCEVs.InsertNode(S, IP);
@@ -1152,7 +1142,7 @@ const SCEV *ScalarEvolution::getPtrToAddrExpr(const SCEV *Op) {
         ID.AddPointer(U);
         ID.AddPointer(Ty);
         void *IP = nullptr;
-        if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
+        if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP))
           return S;
         SCEV *S = new (SCEVAllocator)
             SCEVPtrToAddrExpr(ID.Intern(SCEVAllocator), U, Ty);
@@ -1181,8 +1171,7 @@ const SCEV *ScalarEvolution::getTruncateExpr(SCEVUse Op, Type *Ty,
   ID.AddPointer(Op.getOpaqueValue());
   ID.AddPointer(Ty);
   void *IP = nullptr;
-  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
-    return S;
+  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) return S;
 
   // Fold if the operand is constant.
   if (const SCEVConstant *SC = dyn_cast<SCEVConstant>(Op))
@@ -1236,7 +1225,7 @@ const SCEV *ScalarEvolution::getTruncateExpr(SCEVUse Op, Type *Ty,
     // Although we checked in the beginning that ID is not in the cache, it is
     // possible that during recursion and different modification ID was inserted
     // into the cache. So if we find it, just return it.
-    if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
+    if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP))
       return S;
   }
 
@@ -1508,7 +1497,7 @@ bool ScalarEvolution::proveNoWrapByVaryingStart(const SCEV *Start,
     ID.AddPointer(L);
     void *IP = nullptr;
     const auto *PreAR =
-        static_cast<SCEVAddRecExpr *>(findUniqued(UniqueSCEVs, ID, IP));
+      static_cast<SCEVAddRecExpr *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP));
 
     // Give up if we don't already have the add recurrence we need because
     // actually constructing an add recurrence is relatively expensive.
@@ -1639,8 +1628,7 @@ const SCEV *ScalarEvolution::getZeroExtendExprImpl(SCEVUse Op, Type *Ty,
   ID.AddPointer(Op.getOpaqueValue());
   ID.AddPointer(Ty);
   void *IP = nullptr;
-  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
-    return S;
+  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) return S;
   if (Depth > MaxCastDepth) {
     SCEV *S = new (SCEVAllocator) SCEVZeroExtendExpr(ID.Intern(SCEVAllocator),
                                                      Op, Ty);
@@ -1926,8 +1914,7 @@ const SCEV *ScalarEvolution::getZeroExtendExprImpl(SCEVUse Op, Type *Ty,
 
   // The cast wasn't folded; create an explicit cast node.
   // Recompute the insert position, as it may have been invalidated.
-  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
-    return S;
+  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) return S;
   SCEV *S = new (SCEVAllocator) SCEVZeroExtendExpr(ID.Intern(SCEVAllocator),
                                                    Op, Ty);
   UniqueSCEVs.InsertNode(S, IP);
@@ -1996,8 +1983,7 @@ const SCEV *ScalarEvolution::getSignExtendExprImpl(SCEVUse Op, Type *Ty,
   ID.AddPointer(Op.getOpaqueValue());
   ID.AddPointer(Ty);
   void *IP = nullptr;
-  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
-    return S;
+  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) return S;
   // Limit recursion depth.
   if (Depth > MaxCastDepth) {
     SCEV *S = new (SCEVAllocator) SCEVSignExtendExpr(ID.Intern(SCEVAllocator),
@@ -2191,8 +2177,7 @@ const SCEV *ScalarEvolution::getSignExtendExprImpl(SCEVUse Op, Type *Ty,
 
   // The cast wasn't folded; create an explicit cast node.
   // Recompute the insert position, as it may have been invalidated.
-  if (const SCEV *S = findUniqued(UniqueSCEVs, ID, IP))
-    return S;
+  if (const SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) return S;
   SCEV *S = new (SCEVAllocator) SCEVSignExtendExpr(ID.Intern(SCEVAllocator),
                                                    Op, Ty);
   UniqueSCEVs.InsertNode(S, IP);
@@ -3047,7 +3032,8 @@ const SCEV *ScalarEvolution::getOrCreateAddExpr(ArrayRef<SCEVUse> Ops,
   for (SCEVUse Op : Ops)
     ID.AddPointer(Op.getOpaqueValue());
   void *IP = nullptr;
-  SCEVAddExpr *S = static_cast<SCEVAddExpr *>(findUniqued(UniqueSCEVs, ID, IP));
+  SCEVAddExpr *S =
+      static_cast<SCEVAddExpr *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP));
   if (!S) {
     SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
     llvm::uninitialized_copy(Ops, O);
@@ -3071,7 +3057,7 @@ const SCEV *ScalarEvolution::getOrCreateAddRecExpr(ArrayRef<SCEVUse> Ops,
   ID.AddPointer(L);
   void *IP = nullptr;
   SCEVAddRecExpr *S =
-      static_cast<SCEVAddRecExpr *>(findUniqued(UniqueSCEVs, ID, IP));
+      static_cast<SCEVAddRecExpr *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP));
   if (!S) {
     SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
     llvm::uninitialized_copy(Ops, O);
@@ -3093,7 +3079,8 @@ const SCEV *ScalarEvolution::getOrCreateMulExpr(ArrayRef<SCEVUse> Ops,
   for (SCEVUse Op : Ops)
     ID.AddPointer(Op.getOpaqueValue());
   void *IP = nullptr;
-  SCEVMulExpr *S = static_cast<SCEVMulExpr *>(findUniqued(UniqueSCEVs, ID, IP));
+  SCEVMulExpr *S =
+    static_cast<SCEVMulExpr *>(UniqueSCEVs.FindNodeOrInsertPos(ID, IP));
   if (!S) {
     SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
     llvm::uninitialized_copy(Ops, O);
@@ -3113,7 +3100,7 @@ const SCEV *ScalarEvolution::getOrCreateUDivExpr(SCEVUse LHS, SCEVUse RHS) {
   ID.AddPointer(LHS.getOpaqueValue());
   ID.AddPointer(RHS.getOpaqueValue());
   void *IP = nullptr;
-  SCEV *S = findUniqued(UniqueSCEVs, ID, IP);
+  SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
   if (!S) {
     S = new (SCEVAllocator) SCEVUDivExpr(ID.Intern(SCEVAllocator), LHS, RHS);
     UniqueSCEVs.InsertNode(S, IP);
@@ -3903,7 +3890,7 @@ SCEV *ScalarEvolution::findExistingSCEVInCache(SCEVTypes SCEVType,
   for (SCEVUse Op : Ops)
     ID.AddPointer(Op.getOpaqueValue());
   void *IP = nullptr;
-  return findUniqued(UniqueSCEVs, ID, IP);
+  return UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
 }
 
 const SCEV *ScalarEvolution::getAbsExpr(const SCEV *Op, bool IsNSW) {
@@ -4025,7 +4012,7 @@ const SCEV *ScalarEvolution::getMinMaxExpr(SCEVTypes Kind,
   for (SCEVUse Op : Ops)
     ID.AddPointer(Op.getOpaqueValue());
   void *IP = nullptr;
-  const SCEV *ExistingSCEV = findUniqued(UniqueSCEVs, ID, IP);
+  const SCEV *ExistingSCEV = UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
   if (ExistingSCEV)
     return ExistingSCEV;
   SCEVUse *O = SCEVAllocator.Allocate<SCEVUse>(Ops.size());
@@ -4412,7 +4399,7 @@ ScalarEvolution::getSequentialMinMaxExpr(SCEVTypes Kind,
   for (SCEVUse Op : Ops)
     ID.AddPointer(Op.getOpaqueValue());
   void *IP = nullptr;
-  const SCEV *ExistingSCEV = findUniqued(UniqueSCEVs, ID, IP);
+  const SCEV *ExistingSCEV = UniqueSCEVs.FindNodeOrInsertPos(ID, IP);
   if (ExistingSCEV)
     return ExistingSCEV;
 
@@ -4504,7 +4491,7 @@ const SCEV *ScalarEvolution::getUnknown(Value *V) {
   ID.AddInteger(scUnknown);
   ID.AddPointer(V);
   void *IP = nullptr;
-  if (SCEV *S = findUniqued(UniqueSCEVs, ID, IP)) {
+  if (SCEV *S = UniqueSCEVs.FindNodeOrInsertPos(ID, IP)) {
     assert(cast<SCEVUnknown>(S)->getValue() == V &&
            "Stale SCEVUnknown in uniquing map!");
     return S;
@@ -15181,7 +15168,7 @@ ScalarEvolution::getComparePredicate(const ICmpInst::Predicate Pred,
   ID.AddPointer(LHS);
   ID.AddPointer(RHS);
   void *IP = nullptr;
-  if (const auto *S = findUniqued(UniquePreds, ID, IP))
+  if (const auto *S = UniquePreds.FindNodeOrInsertPos(ID, IP))
     return S;
   SCEVComparePredicate *Eq = new (SCEVAllocator)
     SCEVComparePredicate(ID.Intern(SCEVAllocator), Pred, LHS, RHS);
@@ -15198,7 +15185,7 @@ const SCEVPredicate *ScalarEvolution::getWrapPredicate(
   ID.AddPointer(AR);
   ID.AddInteger(AddedFlags);
   void *IP = nullptr;
-  if (const auto *S = findUniqued(UniquePreds, ID, IP))
+  if (const auto *S = UniquePreds.FindNodeOrInsertPos(ID, IP))
     return S;
   auto *OF = new (SCEVAllocator)
       SCEVWrapPredicate(ID.Intern(SCEVAllocator), AR, AddedFlags);
diff --git a/llvm/lib/Support/FoldingSet.cpp b/llvm/lib/Support/FoldingSet.cpp
index e082984bab27f..34380bbf62e96 100644
--- a/llvm/lib/Support/FoldingSet.cpp
+++ b/llvm/lib/Support/FoldingSet.cpp
@@ -230,28 +230,6 @@ void FoldingSetBase::reserve(unsigned EltCount, const FoldingSetInfo &Info) {
   GrowBucketCount(llvm::bit_floor(EltCount), Info);
 }
 
-FoldingSetBase::Node *FoldingSetBase::FindNodeOrInsertPos(
-    const FoldingSetNodeID &ID, void *&InsertPos, const FoldingSetInfo &Info) {
-  unsigned IDHash = ID.ComputeHash();
-  void **Bucket = getBucketFor(IDHash);
-  void *Probe = *Bucket;
-
-  InsertPos = nullptr;
-
-  FoldingSetNodeID TempID;
-  while (Node *NodeInBucket = GetNextPtr(Probe)) {
-    if (Info.NodeEquals(this, NodeInBucket, ID, IDHash, TempID))
-      return NodeInBucket;
-    TempID.clear();
-
-    Probe = NodeInBucket->getNextInBucket();
-  }
-
-  // Didn't find the node, return null with the bucket as the InsertPos.
-  InsertPos = Bucket;
-  return nullptr;
-}
-
 void FoldingSetBase::InsertNode(Node *N, void *InsertPos,
                                 const FoldingSetInfo &Info) {
   assert(!N->getNextInBucket());



More information about the llvm-commits mailing list