[llvm] [llvm] Don't assume non-erased DenseMap entries remain valid after erase. NFC (PR #198982)

Fangrui Song via llvm-commits llvm-commits at lists.llvm.org
Fri May 22 00:03:02 PDT 2026


https://github.com/MaskRay updated https://github.com/llvm/llvm-project/pull/198982

>From eae139cf8cf8dbae40c769b2b00b9e3de657b31e Mon Sep 17 00:00:00 2001
From: Fangrui Song <i at maskray.me>
Date: Mon, 11 May 2026 23:37:03 -0700
Subject: [PATCH 1/3] [llvm] Don't assume non-erased DenseMap entries remain
 valid after erase. NFC

In preparation for switching DenseMap from tombstone deletion to TAOCP
Algorithm 6.4R, improving performance, don't rely on the iterators of
non-erased elements being kept.

Switch to the collect-then-erase idiom.

Aided by Claude Opus 4.7
---
 llvm/include/llvm/ADT/SetOperations.h         | 27 ++++++++--------
 llvm/lib/Analysis/LoopAccessAnalysis.cpp      |  5 ++-
 llvm/lib/Analysis/ScalarEvolution.cpp         | 31 +++++++++----------
 llvm/lib/CodeGen/MachineCopyPropagation.cpp   |  2 +-
 llvm/lib/CodeGen/MachineLateInstrsCleanup.cpp | 15 +++++----
 llvm/lib/CodeGen/PeepholeOptimizer.cpp        | 13 +++++---
 .../CodeGen/RemoveRedundantDebugValues.cpp    |  5 ++-
 llvm/lib/IR/LegacyPassManager.cpp             | 29 +++++++++--------
 llvm/lib/MC/MCObjectStreamer.cpp              | 13 +++++---
 llvm/lib/Transforms/IPO/FunctionImport.cpp    | 12 +++----
 .../Utils/PromoteMemoryToRegister.cpp         | 14 ++++-----
 11 files changed, 90 insertions(+), 76 deletions(-)

diff --git a/llvm/include/llvm/ADT/SetOperations.h b/llvm/include/llvm/ADT/SetOperations.h
index 4d4ff4045f813..36943452cd422 100644
--- a/llvm/include/llvm/ADT/SetOperations.h
+++ b/llvm/include/llvm/ADT/SetOperations.h
@@ -16,6 +16,7 @@
 #define LLVM_ADT_SETOPERATIONS_H
 
 #include "llvm/ADT/STLExtras.h"
+#include "llvm/ADT/SmallVector.h"
 
 namespace llvm {
 
@@ -60,12 +61,12 @@ template <class S1Ty, class S2Ty> void set_intersect(S1Ty &S1, const S2Ty &S2) {
   if constexpr (detail::HasMemberRemoveIf<S1Ty, decltype(Pred)>) {
     S1.remove_if(Pred);
   } else {
-    typename S1Ty::iterator Next;
-    for (typename S1Ty::iterator I = S1.begin(); I != S1.end(); I = Next) {
-      Next = std::next(I);
-      if (!S2.count(*I))
-        S1.erase(I); // Erase element if not in S2
-    }
+    SmallVector<typename S1Ty::value_type> ToRemove;
+    for (const auto &E : S1)
+      if (!S2.count(E))
+        ToRemove.push_back(E);
+    for (const auto &E : ToRemove)
+      S1.erase(E);
   }
 }
 
@@ -117,13 +118,13 @@ template <class S1Ty, class S2Ty> void set_subtract(S1Ty &S1, const S2Ty &S2) {
       }
     } else if constexpr (detail::HasMemberEraseIter<S1Ty>) {
       if (S1.size() < S2.size()) {
-        typename S1Ty::iterator Next;
-        for (typename S1Ty::iterator SI = S1.begin(), SE = S1.end(); SI != SE;
-             SI = Next) {
-          Next = std::next(SI);
-          if (S2.contains(*SI))
-            S1.erase(SI);
-        }
+        // Collect-then-erase: see the matching comment in set_intersect.
+        SmallVector<typename S1Ty::value_type> ToRemove;
+        for (const auto &E : S1)
+          if (S2.contains(E))
+            ToRemove.push_back(E);
+        for (const auto &E : ToRemove)
+          S1.erase(E);
         return;
       }
     }
diff --git a/llvm/lib/Analysis/LoopAccessAnalysis.cpp b/llvm/lib/Analysis/LoopAccessAnalysis.cpp
index 2b9efd22131c6..365cd630278d0 100644
--- a/llvm/lib/Analysis/LoopAccessAnalysis.cpp
+++ b/llvm/lib/Analysis/LoopAccessAnalysis.cpp
@@ -3229,12 +3229,15 @@ void LoopAccessInfoManager::clear() {
   // analyzed loop or SCEVs that may have been modified or invalidated. At the
   // moment, that is loops requiring memory or SCEV runtime checks, as those cache
   // SCEVs, e.g. for pointer expressions.
+  SmallVector<Loop *> ToRemove;
   for (const auto &[L, LAI] : LoopAccessInfoMap) {
     if (LAI->getRuntimePointerChecking()->getChecks().empty() &&
         LAI->getPSE().getPredicate().isAlwaysTrue())
       continue;
-    LoopAccessInfoMap.erase(L);
+    ToRemove.push_back(L);
   }
+  for (Loop *L : ToRemove)
+    LoopAccessInfoMap.erase(L);
 }
 
 bool LoopAccessInfoManager::invalidate(
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index 6dbcd82744f8b..c345f449acaf1 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -8749,8 +8749,9 @@ void ScalarEvolution::visitAndClearUsers(
     ValueExprMapType::iterator It =
         ValueExprMap.find_as(static_cast<Value *>(I));
     if (It != ValueExprMap.end()) {
+      SCEVUse Mapped = It->second;
       eraseValueFromMap(It->first);
-      ToForget.push_back(It->second);
+      ToForget.push_back(Mapped);
       if (PHINode *PN = dyn_cast<PHINode>(I))
         ConstantEvolutionLoopExitValue.erase(PN);
     }
@@ -8774,14 +8775,12 @@ void ScalarEvolution::forgetLoop(const Loop *L) {
     forgetBackedgeTakenCounts(CurrL, /* Predicated */ true);
 
     // Drop information about predicated SCEV rewrites for this loop.
-    for (auto I = PredicatedSCEVRewrites.begin();
-         I != PredicatedSCEVRewrites.end();) {
-      std::pair<const SCEV *, const Loop *> Entry = I->first;
-      if (Entry.second == CurrL)
-        PredicatedSCEVRewrites.erase(I++);
-      else
-        ++I;
-    }
+    SmallVector<std::pair<const SCEVUnknown *, const Loop *>> ToRemove;
+    for (const auto &Entry : PredicatedSCEVRewrites)
+      if (Entry.first.second == CurrL)
+        ToRemove.push_back(Entry.first);
+    for (const auto &K : ToRemove)
+      PredicatedSCEVRewrites.erase(K);
 
     auto LoopUsersItr = LoopUsers.find(CurrL);
     if (LoopUsersItr != LoopUsers.end())
@@ -14567,14 +14566,12 @@ void ScalarEvolution::forgetMemoizedResults(ArrayRef<SCEVUse> SCEVs) {
   for (const auto *S : ToForget)
     forgetMemoizedResultsImpl(S);
 
-  for (auto I = PredicatedSCEVRewrites.begin();
-       I != PredicatedSCEVRewrites.end();) {
-    std::pair<const SCEV *, const Loop *> Entry = I->first;
-    if (ToForget.count(Entry.first))
-      PredicatedSCEVRewrites.erase(I++);
-    else
-      ++I;
-  }
+  SmallVector<std::pair<const SCEVUnknown *, const Loop *>> ToRemove;
+  for (const auto &Entry : PredicatedSCEVRewrites)
+    if (ToForget.count(Entry.first.first))
+      ToRemove.push_back(Entry.first);
+  for (const auto &K : ToRemove)
+    PredicatedSCEVRewrites.erase(K);
 }
 
 void ScalarEvolution::forgetMemoizedResultsImpl(const SCEV *S) {
diff --git a/llvm/lib/CodeGen/MachineCopyPropagation.cpp b/llvm/lib/CodeGen/MachineCopyPropagation.cpp
index 1d8fdcab64909..6bfd21b215706 100644
--- a/llvm/lib/CodeGen/MachineCopyPropagation.cpp
+++ b/llvm/lib/CodeGen/MachineCopyPropagation.cpp
@@ -255,7 +255,7 @@ class CopyTracker {
         }
       }
       // Now we can erase the copy.
-      Copies.erase(I);
+      Copies.erase(Unit);
     }
   }
 
diff --git a/llvm/lib/CodeGen/MachineLateInstrsCleanup.cpp b/llvm/lib/CodeGen/MachineLateInstrsCleanup.cpp
index 4f281fa1361ca..946fa687b067c 100644
--- a/llvm/lib/CodeGen/MachineLateInstrsCleanup.cpp
+++ b/llvm/lib/CodeGen/MachineLateInstrsCleanup.cpp
@@ -243,15 +243,18 @@ bool MachineLateInstrsCleanup::processBlock(MachineBasicBlock *MBB) {
     }
 
     // Clear any entries in map that MI clobbers.
-    for (auto DefI : llvm::make_early_inc_range(MBBDefs)) {
-      Register Reg = DefI.first;
-      if (MI.modifiesRegister(Reg, TRI)) {
-        MBBDefs.erase(Reg);
-        MBBKills.erase(Reg);
-      } else if (MI.findRegisterUseOperandIdx(Reg, TRI, true /*isKill*/) != -1)
+    SmallVector<Register> Clobbered;
+    for (auto [Reg, DefMI] : MBBDefs) {
+      if (MI.modifiesRegister(Reg, TRI))
+        Clobbered.push_back(Reg);
+      else if (MI.findRegisterUseOperandIdx(Reg, TRI, true /*isKill*/) != -1)
         // Keep track of all instructions that fully or partially kills Reg.
         MBBKills[Reg].push_back(&MI);
     }
+    for (Register Reg : Clobbered) {
+      MBBDefs.erase(Reg);
+      MBBKills.erase(Reg);
+    }
 
     // Record this MI for potential later reuse.
     if (IsCandidate) {
diff --git a/llvm/lib/CodeGen/PeepholeOptimizer.cpp b/llvm/lib/CodeGen/PeepholeOptimizer.cpp
index cfbffb920ef36..8ba7a78c1d014 100644
--- a/llvm/lib/CodeGen/PeepholeOptimizer.cpp
+++ b/llvm/lib/CodeGen/PeepholeOptimizer.cpp
@@ -1813,13 +1813,16 @@ bool PeepholeOptimizer::run(MachineFunction &MF) {
             }
           } else if (MO.isRegMask()) {
             const uint32_t *RegMask = MO.getRegMask();
+            SmallVector<Register> ClobberedDefs;
             for (auto &RegMI : NAPhysToVirtMIs) {
               Register Def = RegMI.first;
-              if (MachineOperand::clobbersPhysReg(RegMask, Def)) {
-                LLVM_DEBUG(dbgs()
-                           << "NAPhysCopy: invalidating because of " << *MI);
-                NAPhysToVirtMIs.erase(Def);
-              }
+              if (MachineOperand::clobbersPhysReg(RegMask, Def))
+                ClobberedDefs.push_back(Def);
+            }
+            for (Register Def : ClobberedDefs) {
+              LLVM_DEBUG(dbgs()
+                         << "NAPhysCopy: invalidating because of " << *MI);
+              NAPhysToVirtMIs.erase(Def);
             }
           }
         }
diff --git a/llvm/lib/CodeGen/RemoveRedundantDebugValues.cpp b/llvm/lib/CodeGen/RemoveRedundantDebugValues.cpp
index 11468245f8400..c2fd20084efcd 100644
--- a/llvm/lib/CodeGen/RemoveRedundantDebugValues.cpp
+++ b/llvm/lib/CodeGen/RemoveRedundantDebugValues.cpp
@@ -128,11 +128,14 @@ static bool reduceDbgValsForwardScan(MachineBasicBlock &MBB) {
       continue;
 
     // Stop tracking any location that is clobbered by this instruction.
+    SmallVector<DebugVariable> Clobbered;
     for (auto &Var : VariableMap) {
       auto &LocOp = Var.second.first;
       if (MI.modifiesRegister(LocOp->getReg(), TRI))
-        VariableMap.erase(Var.first);
+        Clobbered.push_back(Var.first);
     }
+    for (const DebugVariable &Var : Clobbered)
+      VariableMap.erase(Var);
   }
 
   for (auto &Instr : DbgValsToBeRemoved) {
diff --git a/llvm/lib/IR/LegacyPassManager.cpp b/llvm/lib/IR/LegacyPassManager.cpp
index 7b9ad89038dc6..d1d5237e28146 100644
--- a/llvm/lib/IR/LegacyPassManager.cpp
+++ b/llvm/lib/IR/LegacyPassManager.cpp
@@ -903,20 +903,21 @@ void PMDataManager::removeNotPreservedAnalysis(Pass *P) {
     return;
 
   const AnalysisUsage::VectorType &PreservedSet = AnUsage->getPreservedSet();
-  for (auto I = AvailableAnalysis.begin(), E = AvailableAnalysis.end();
-       I != E;) {
-    auto Info = I++;
-    if (Info->second->getAsImmutablePass() == nullptr &&
-        !is_contained(PreservedSet, Info->first)) {
+  SmallVector<AnalysisID, 16> ToRemove;
+  for (auto &Entry : AvailableAnalysis) {
+    if (Entry.second->getAsImmutablePass() == nullptr &&
+        !is_contained(PreservedSet, Entry.first)) {
       // Remove this analysis
       if (PassDebugging >= Details) {
-        Pass *S = Info->second;
+        Pass *S = Entry.second;
         dbgs() << " -- '" <<  P->getPassName() << "' is not preserving '";
         dbgs() << S->getPassName() << "'\n";
       }
-      AvailableAnalysis.erase(Info);
+      ToRemove.push_back(Entry.first);
     }
   }
+  for (AnalysisID ID : ToRemove)
+    AvailableAnalysis.erase(ID);
 
   // Check inherited analysis also. If P is not preserving analysis
   // provided by parent manager then remove it here.
@@ -924,19 +925,21 @@ void PMDataManager::removeNotPreservedAnalysis(Pass *P) {
     if (!IA)
       continue;
 
-    for (auto I = IA->begin(), E = IA->end(); I != E;) {
-      auto Info = I++;
-      if (Info->second->getAsImmutablePass() == nullptr &&
-          !is_contained(PreservedSet, Info->first)) {
+    ToRemove.clear();
+    for (auto &Entry : *IA) {
+      if (Entry.second->getAsImmutablePass() == nullptr &&
+          !is_contained(PreservedSet, Entry.first)) {
         // Remove this analysis
         if (PassDebugging >= Details) {
-          Pass *S = Info->second;
+          Pass *S = Entry.second;
           dbgs() << " -- '" <<  P->getPassName() << "' is not preserving '";
           dbgs() << S->getPassName() << "'\n";
         }
-        IA->erase(Info);
+        ToRemove.push_back(Entry.first);
       }
     }
+    for (AnalysisID ID : ToRemove)
+      IA->erase(ID);
   }
 }
 
diff --git a/llvm/lib/MC/MCObjectStreamer.cpp b/llvm/lib/MC/MCObjectStreamer.cpp
index 88dafb94a4aaa..2bf5f05c1c315 100644
--- a/llvm/lib/MC/MCObjectStreamer.cpp
+++ b/llvm/lib/MC/MCObjectStreamer.cpp
@@ -272,12 +272,15 @@ void MCObjectStreamer::emitLabel(MCSymbol *Symbol, SMLoc Loc) {
 
 void MCObjectStreamer::emitPendingAssignments(MCSymbol *Symbol) {
   auto Assignments = pendingAssignments.find(Symbol);
-  if (Assignments != pendingAssignments.end()) {
-    for (const PendingAssignment &A : Assignments->second)
-      emitAssignment(A.Symbol, A.Value);
+  if (Assignments == pendingAssignments.end())
+    return;
 
-    pendingAssignments.erase(Assignments);
-  }
+  // emitAssignment can recursively re-enter emitPendingAssignments for
+  // other symbols, so move the list out and erase before iterating.
+  SmallVector<PendingAssignment, 1> Pending = std::move(Assignments->second);
+  pendingAssignments.erase(Assignments);
+  for (const PendingAssignment &A : Pending)
+    emitAssignment(A.Symbol, A.Value);
 }
 
 // Emit a label at a previously emitted fragment/offset position. This must be
diff --git a/llvm/lib/Transforms/IPO/FunctionImport.cpp b/llvm/lib/Transforms/IPO/FunctionImport.cpp
index 456a9b116cc30..6f96d464f0a60 100644
--- a/llvm/lib/Transforms/IPO/FunctionImport.cpp
+++ b/llvm/lib/Transforms/IPO/FunctionImport.cpp
@@ -1284,12 +1284,12 @@ void llvm::ComputeCrossModuleImport(
     // exporting module. We do this after the above insertion since we may hit
     // the same ref/call target multiple times in above loop, and it is more
     // efficient to avoid a set lookup each time.
-    for (auto EI = NewExports.begin(); EI != NewExports.end();) {
-      if (!DefinedGVSummaries.count(EI->getGUID()))
-        NewExports.erase(EI++);
-      else
-        ++EI;
-    }
+    SmallVector<ValueInfo, 8> ToRemove;
+    for (const auto &VI : NewExports)
+      if (!DefinedGVSummaries.count(VI.getGUID()))
+        ToRemove.push_back(VI);
+    for (const auto &VI : ToRemove)
+      NewExports.erase(VI);
     ELI.second.insert_range(NewExports);
   }
 
diff --git a/llvm/lib/Transforms/Utils/PromoteMemoryToRegister.cpp b/llvm/lib/Transforms/Utils/PromoteMemoryToRegister.cpp
index b635d805bf13d..631190a61cc6d 100644
--- a/llvm/lib/Transforms/Utils/PromoteMemoryToRegister.cpp
+++ b/llvm/lib/Transforms/Utils/PromoteMemoryToRegister.cpp
@@ -929,22 +929,20 @@ void PromoteMem2Reg::run() {
     // simplify and RAUW them as we go.  If it was not, we could add uses to
     // the values we replace with in a non-deterministic order, thus creating
     // non-deterministic def->use chains.
-    for (DenseMap<std::pair<unsigned, unsigned>, PHINode *>::iterator
-             I = NewPhiNodes.begin(),
-             E = NewPhiNodes.end();
-         I != E;) {
-      PHINode *PN = I->second;
+    SmallVector<std::pair<unsigned, unsigned>> ToRemove;
+    for (auto &Entry : NewPhiNodes) {
+      PHINode *PN = Entry.second;
 
       // If this PHI node merges one value and/or undefs, get the value.
       if (Value *V = simplifyInstruction(PN, SQ)) {
         PN->replaceAllUsesWith(V);
         PN->eraseFromParent();
-        NewPhiNodes.erase(I++);
+        ToRemove.push_back(Entry.first);
         EliminatedAPHI = true;
-        continue;
       }
-      ++I;
     }
+    for (auto &K : ToRemove)
+      NewPhiNodes.erase(K);
   }
 
   // At this point, the renamer has added entries to PHI nodes for all reachable

>From 51ba1a081c525679dd0ca449f78724ddd7b362f0 Mon Sep 17 00:00:00 2001
From: Fangrui Song <i at maskray.me>
Date: Thu, 21 May 2026 10:26:42 -0700
Subject: [PATCH 2/3] simplify

---
 llvm/lib/IR/LegacyPassManager.cpp          | 2 +-
 llvm/lib/MC/MCObjectStreamer.cpp           | 2 +-
 llvm/lib/Transforms/IPO/FunctionImport.cpp | 2 +-
 3 files changed, 3 insertions(+), 3 deletions(-)

diff --git a/llvm/lib/IR/LegacyPassManager.cpp b/llvm/lib/IR/LegacyPassManager.cpp
index d1d5237e28146..457f290feebbe 100644
--- a/llvm/lib/IR/LegacyPassManager.cpp
+++ b/llvm/lib/IR/LegacyPassManager.cpp
@@ -903,7 +903,7 @@ void PMDataManager::removeNotPreservedAnalysis(Pass *P) {
     return;
 
   const AnalysisUsage::VectorType &PreservedSet = AnUsage->getPreservedSet();
-  SmallVector<AnalysisID, 16> ToRemove;
+  SmallVector<AnalysisID> ToRemove;
   for (auto &Entry : AvailableAnalysis) {
     if (Entry.second->getAsImmutablePass() == nullptr &&
         !is_contained(PreservedSet, Entry.first)) {
diff --git a/llvm/lib/MC/MCObjectStreamer.cpp b/llvm/lib/MC/MCObjectStreamer.cpp
index 2bf5f05c1c315..e6b622331036a 100644
--- a/llvm/lib/MC/MCObjectStreamer.cpp
+++ b/llvm/lib/MC/MCObjectStreamer.cpp
@@ -277,7 +277,7 @@ void MCObjectStreamer::emitPendingAssignments(MCSymbol *Symbol) {
 
   // emitAssignment can recursively re-enter emitPendingAssignments for
   // other symbols, so move the list out and erase before iterating.
-  SmallVector<PendingAssignment, 1> Pending = std::move(Assignments->second);
+  SmallVector<PendingAssignment> Pending = std::move(Assignments->second);
   pendingAssignments.erase(Assignments);
   for (const PendingAssignment &A : Pending)
     emitAssignment(A.Symbol, A.Value);
diff --git a/llvm/lib/Transforms/IPO/FunctionImport.cpp b/llvm/lib/Transforms/IPO/FunctionImport.cpp
index 6f96d464f0a60..c55674495646d 100644
--- a/llvm/lib/Transforms/IPO/FunctionImport.cpp
+++ b/llvm/lib/Transforms/IPO/FunctionImport.cpp
@@ -1284,7 +1284,7 @@ void llvm::ComputeCrossModuleImport(
     // exporting module. We do this after the above insertion since we may hit
     // the same ref/call target multiple times in above loop, and it is more
     // efficient to avoid a set lookup each time.
-    SmallVector<ValueInfo, 8> ToRemove;
+    SmallVector<ValueInfo> ToRemove;
     for (const auto &VI : NewExports)
       if (!DefinedGVSummaries.count(VI.getGUID()))
         ToRemove.push_back(VI);

>From 90be26bb48afafe8bf7571b06e939bf9ad322056 Mon Sep 17 00:00:00 2001
From: Fangrui Song <i at maskray.me>
Date: Fri, 22 May 2026 00:02:50 -0700
Subject: [PATCH 3/3] add remove_if and simplify code with it

---
 llvm/include/llvm/ADT/DenseMap.h              | 27 +++++++++++
 llvm/include/llvm/ADT/DenseSet.h              | 10 +++++
 llvm/include/llvm/ADT/SetOperations.h         | 27 ++++++-----
 llvm/lib/Analysis/LoopAccessAnalysis.cpp      | 14 +++---
 llvm/lib/Analysis/ScalarEvolution.cpp         | 16 ++-----
 llvm/lib/CodeGen/MachineLateInstrsCleanup.cpp | 19 ++++----
 llvm/lib/CodeGen/PeepholeOptimizer.cpp        | 14 +++---
 .../CodeGen/RemoveRedundantDebugValues.cpp    | 11 ++---
 llvm/lib/IR/LegacyPassManager.cpp             | 43 ++++++------------
 llvm/lib/MC/MCObjectStreamer.cpp              |  2 +-
 llvm/lib/Transforms/IPO/FunctionImport.cpp    |  8 +---
 .../Utils/PromoteMemoryToRegister.cpp         | 12 ++---
 llvm/unittests/ADT/DenseMapTest.cpp           | 45 +++++++++++++++++++
 llvm/unittests/ADT/DenseSetTest.cpp           | 16 +++++++
 14 files changed, 157 insertions(+), 107 deletions(-)

diff --git a/llvm/include/llvm/ADT/DenseMap.h b/llvm/include/llvm/ADT/DenseMap.h
index b8b548a31acbc..e13b64b3e6bf4 100644
--- a/llvm/include/llvm/ADT/DenseMap.h
+++ b/llvm/include/llvm/ADT/DenseMap.h
@@ -344,6 +344,33 @@ class DenseMapBase : public DebugEpochBase {
     incrementNumTombstones();
   }
 
+  /// Remove entries that match the given predicate. \p Pred is invoked
+  /// with a reference to each live bucket and must not access the map being
+  /// modified. This is the safe replacement for erase-while-iterating.
+  ///
+  /// Returns whether anything was removed. If so, all iterators and references
+  /// into the map are invalidated.
+  template <typename Predicate> bool remove_if(Predicate Pred) {
+    const KeyT EmptyKey = KeyInfoT::getEmptyKey();
+    const KeyT TombstoneKey = KeyInfoT::getTombstoneKey();
+    bool Removed = false;
+    for (BucketT &B : buckets()) {
+      if (KeyInfoT::isEqual(B.getFirst(), EmptyKey) ||
+          KeyInfoT::isEqual(B.getFirst(), TombstoneKey))
+        continue;
+      if (Pred(B)) {
+        B.getSecond().~ValueT();
+        B.getFirst() = TombstoneKey;
+        decrementNumEntries();
+        incrementNumTombstones();
+        Removed = true;
+      }
+    }
+    if (Removed)
+      incrementEpoch();
+    return Removed;
+  }
+
   ValueT &operator[](const KeyT &Key) {
     return lookupOrInsertIntoBucket(Key).first->second;
   }
diff --git a/llvm/include/llvm/ADT/DenseSet.h b/llvm/include/llvm/ADT/DenseSet.h
index eec800d07b6df..645d6d1568f35 100644
--- a/llvm/include/llvm/ADT/DenseSet.h
+++ b/llvm/include/llvm/ADT/DenseSet.h
@@ -99,6 +99,16 @@ class DenseSetImpl {
 
   bool erase(const ValueT &V) { return TheMap.erase(V); }
 
+  /// Remove all elements for which \p Pred returns true.  This is the safe
+  /// replacement for erase-while-iterating; see DenseMap::remove_if.  The
+  /// predicate must not access the set being modified.  Returns whether
+  /// anything was removed; if so, all iterators are invalidated.
+  template <typename Predicate> bool remove_if(Predicate Pred) {
+    return TheMap.remove_if([&](const typename MapTy::value_type &KV) {
+      return Pred(KV.getFirst());
+    });
+  }
+
   void swap(DenseSetImpl &RHS) { TheMap.swap(RHS.TheMap); }
 
 private:
diff --git a/llvm/include/llvm/ADT/SetOperations.h b/llvm/include/llvm/ADT/SetOperations.h
index 36943452cd422..4d4ff4045f813 100644
--- a/llvm/include/llvm/ADT/SetOperations.h
+++ b/llvm/include/llvm/ADT/SetOperations.h
@@ -16,7 +16,6 @@
 #define LLVM_ADT_SETOPERATIONS_H
 
 #include "llvm/ADT/STLExtras.h"
-#include "llvm/ADT/SmallVector.h"
 
 namespace llvm {
 
@@ -61,12 +60,12 @@ template <class S1Ty, class S2Ty> void set_intersect(S1Ty &S1, const S2Ty &S2) {
   if constexpr (detail::HasMemberRemoveIf<S1Ty, decltype(Pred)>) {
     S1.remove_if(Pred);
   } else {
-    SmallVector<typename S1Ty::value_type> ToRemove;
-    for (const auto &E : S1)
-      if (!S2.count(E))
-        ToRemove.push_back(E);
-    for (const auto &E : ToRemove)
-      S1.erase(E);
+    typename S1Ty::iterator Next;
+    for (typename S1Ty::iterator I = S1.begin(); I != S1.end(); I = Next) {
+      Next = std::next(I);
+      if (!S2.count(*I))
+        S1.erase(I); // Erase element if not in S2
+    }
   }
 }
 
@@ -118,13 +117,13 @@ template <class S1Ty, class S2Ty> void set_subtract(S1Ty &S1, const S2Ty &S2) {
       }
     } else if constexpr (detail::HasMemberEraseIter<S1Ty>) {
       if (S1.size() < S2.size()) {
-        // Collect-then-erase: see the matching comment in set_intersect.
-        SmallVector<typename S1Ty::value_type> ToRemove;
-        for (const auto &E : S1)
-          if (S2.contains(E))
-            ToRemove.push_back(E);
-        for (const auto &E : ToRemove)
-          S1.erase(E);
+        typename S1Ty::iterator Next;
+        for (typename S1Ty::iterator SI = S1.begin(), SE = S1.end(); SI != SE;
+             SI = Next) {
+          Next = std::next(SI);
+          if (S2.contains(*SI))
+            S1.erase(SI);
+        }
         return;
       }
     }
diff --git a/llvm/lib/Analysis/LoopAccessAnalysis.cpp b/llvm/lib/Analysis/LoopAccessAnalysis.cpp
index 365cd630278d0..60e15b2a5bd82 100644
--- a/llvm/lib/Analysis/LoopAccessAnalysis.cpp
+++ b/llvm/lib/Analysis/LoopAccessAnalysis.cpp
@@ -3229,15 +3229,11 @@ void LoopAccessInfoManager::clear() {
   // analyzed loop or SCEVs that may have been modified or invalidated. At the
   // moment, that is loops requiring memory or SCEV runtime checks, as those cache
   // SCEVs, e.g. for pointer expressions.
-  SmallVector<Loop *> ToRemove;
-  for (const auto &[L, LAI] : LoopAccessInfoMap) {
-    if (LAI->getRuntimePointerChecking()->getChecks().empty() &&
-        LAI->getPSE().getPredicate().isAlwaysTrue())
-      continue;
-    ToRemove.push_back(L);
-  }
-  for (Loop *L : ToRemove)
-    LoopAccessInfoMap.erase(L);
+  LoopAccessInfoMap.remove_if([](const auto &Entry) {
+    const auto &LAI = Entry.second;
+    return !(LAI->getRuntimePointerChecking()->getChecks().empty() &&
+             LAI->getPSE().getPredicate().isAlwaysTrue());
+  });
 }
 
 bool LoopAccessInfoManager::invalidate(
diff --git a/llvm/lib/Analysis/ScalarEvolution.cpp b/llvm/lib/Analysis/ScalarEvolution.cpp
index c345f449acaf1..ec221e5b6ade5 100644
--- a/llvm/lib/Analysis/ScalarEvolution.cpp
+++ b/llvm/lib/Analysis/ScalarEvolution.cpp
@@ -8775,12 +8775,8 @@ void ScalarEvolution::forgetLoop(const Loop *L) {
     forgetBackedgeTakenCounts(CurrL, /* Predicated */ true);
 
     // Drop information about predicated SCEV rewrites for this loop.
-    SmallVector<std::pair<const SCEVUnknown *, const Loop *>> ToRemove;
-    for (const auto &Entry : PredicatedSCEVRewrites)
-      if (Entry.first.second == CurrL)
-        ToRemove.push_back(Entry.first);
-    for (const auto &K : ToRemove)
-      PredicatedSCEVRewrites.erase(K);
+    PredicatedSCEVRewrites.remove_if(
+        [&](const auto &Entry) { return Entry.first.second == CurrL; });
 
     auto LoopUsersItr = LoopUsers.find(CurrL);
     if (LoopUsersItr != LoopUsers.end())
@@ -14566,12 +14562,8 @@ void ScalarEvolution::forgetMemoizedResults(ArrayRef<SCEVUse> SCEVs) {
   for (const auto *S : ToForget)
     forgetMemoizedResultsImpl(S);
 
-  SmallVector<std::pair<const SCEVUnknown *, const Loop *>> ToRemove;
-  for (const auto &Entry : PredicatedSCEVRewrites)
-    if (ToForget.count(Entry.first.first))
-      ToRemove.push_back(Entry.first);
-  for (const auto &K : ToRemove)
-    PredicatedSCEVRewrites.erase(K);
+  PredicatedSCEVRewrites.remove_if(
+      [&](const auto &Entry) { return ToForget.count(Entry.first.first); });
 }
 
 void ScalarEvolution::forgetMemoizedResultsImpl(const SCEV *S) {
diff --git a/llvm/lib/CodeGen/MachineLateInstrsCleanup.cpp b/llvm/lib/CodeGen/MachineLateInstrsCleanup.cpp
index 946fa687b067c..811cc4fe65f3f 100644
--- a/llvm/lib/CodeGen/MachineLateInstrsCleanup.cpp
+++ b/llvm/lib/CodeGen/MachineLateInstrsCleanup.cpp
@@ -243,18 +243,17 @@ bool MachineLateInstrsCleanup::processBlock(MachineBasicBlock *MBB) {
     }
 
     // Clear any entries in map that MI clobbers.
-    SmallVector<Register> Clobbered;
-    for (auto [Reg, DefMI] : MBBDefs) {
-      if (MI.modifiesRegister(Reg, TRI))
-        Clobbered.push_back(Reg);
-      else if (MI.findRegisterUseOperandIdx(Reg, TRI, true /*isKill*/) != -1)
+    MBBDefs.remove_if([&](const auto &Entry) {
+      Register Reg = Entry.first;
+      if (MI.modifiesRegister(Reg, TRI)) {
+        MBBKills.erase(Reg);
+        return true;
+      }
+      if (MI.findRegisterUseOperandIdx(Reg, TRI, true /*isKill*/) != -1)
         // Keep track of all instructions that fully or partially kills Reg.
         MBBKills[Reg].push_back(&MI);
-    }
-    for (Register Reg : Clobbered) {
-      MBBDefs.erase(Reg);
-      MBBKills.erase(Reg);
-    }
+      return false;
+    });
 
     // Record this MI for potential later reuse.
     if (IsCandidate) {
diff --git a/llvm/lib/CodeGen/PeepholeOptimizer.cpp b/llvm/lib/CodeGen/PeepholeOptimizer.cpp
index 8ba7a78c1d014..19a81e3363086 100644
--- a/llvm/lib/CodeGen/PeepholeOptimizer.cpp
+++ b/llvm/lib/CodeGen/PeepholeOptimizer.cpp
@@ -1813,17 +1813,13 @@ bool PeepholeOptimizer::run(MachineFunction &MF) {
             }
           } else if (MO.isRegMask()) {
             const uint32_t *RegMask = MO.getRegMask();
-            SmallVector<Register> ClobberedDefs;
-            for (auto &RegMI : NAPhysToVirtMIs) {
-              Register Def = RegMI.first;
-              if (MachineOperand::clobbersPhysReg(RegMask, Def))
-                ClobberedDefs.push_back(Def);
-            }
-            for (Register Def : ClobberedDefs) {
+            NAPhysToVirtMIs.remove_if([&](const auto &RegMI) {
+              if (!MachineOperand::clobbersPhysReg(RegMask, RegMI.first))
+                return false;
               LLVM_DEBUG(dbgs()
                          << "NAPhysCopy: invalidating because of " << *MI);
-              NAPhysToVirtMIs.erase(Def);
-            }
+              return true;
+            });
           }
         }
       }
diff --git a/llvm/lib/CodeGen/RemoveRedundantDebugValues.cpp b/llvm/lib/CodeGen/RemoveRedundantDebugValues.cpp
index c2fd20084efcd..30057317a6383 100644
--- a/llvm/lib/CodeGen/RemoveRedundantDebugValues.cpp
+++ b/llvm/lib/CodeGen/RemoveRedundantDebugValues.cpp
@@ -128,14 +128,9 @@ static bool reduceDbgValsForwardScan(MachineBasicBlock &MBB) {
       continue;
 
     // Stop tracking any location that is clobbered by this instruction.
-    SmallVector<DebugVariable> Clobbered;
-    for (auto &Var : VariableMap) {
-      auto &LocOp = Var.second.first;
-      if (MI.modifiesRegister(LocOp->getReg(), TRI))
-        Clobbered.push_back(Var.first);
-    }
-    for (const DebugVariable &Var : Clobbered)
-      VariableMap.erase(Var);
+    VariableMap.remove_if([&](const auto &Var) {
+      return MI.modifiesRegister(Var.second.first->getReg(), TRI);
+    });
   }
 
   for (auto &Instr : DbgValsToBeRemoved) {
diff --git a/llvm/lib/IR/LegacyPassManager.cpp b/llvm/lib/IR/LegacyPassManager.cpp
index 457f290feebbe..b8efa7a399734 100644
--- a/llvm/lib/IR/LegacyPassManager.cpp
+++ b/llvm/lib/IR/LegacyPassManager.cpp
@@ -903,43 +903,26 @@ void PMDataManager::removeNotPreservedAnalysis(Pass *P) {
     return;
 
   const AnalysisUsage::VectorType &PreservedSet = AnUsage->getPreservedSet();
-  SmallVector<AnalysisID> ToRemove;
-  for (auto &Entry : AvailableAnalysis) {
-    if (Entry.second->getAsImmutablePass() == nullptr &&
-        !is_contained(PreservedSet, Entry.first)) {
-      // Remove this analysis
-      if (PassDebugging >= Details) {
-        Pass *S = Entry.second;
-        dbgs() << " -- '" <<  P->getPassName() << "' is not preserving '";
-        dbgs() << S->getPassName() << "'\n";
-      }
-      ToRemove.push_back(Entry.first);
+  auto IsNotPreserved = [&](const auto &Entry) {
+    if (Entry.second->getAsImmutablePass() != nullptr ||
+        is_contained(PreservedSet, Entry.first))
+      return false;
+    // Remove this analysis
+    if (PassDebugging >= Details) {
+      Pass *S = Entry.second;
+      dbgs() << " -- '" << P->getPassName() << "' is not preserving '";
+      dbgs() << S->getPassName() << "'\n";
     }
-  }
-  for (AnalysisID ID : ToRemove)
-    AvailableAnalysis.erase(ID);
+    return true;
+  };
+  AvailableAnalysis.remove_if(IsNotPreserved);
 
   // Check inherited analysis also. If P is not preserving analysis
   // provided by parent manager then remove it here.
   for (DenseMap<AnalysisID, Pass *> *IA : InheritedAnalysis) {
     if (!IA)
       continue;
-
-    ToRemove.clear();
-    for (auto &Entry : *IA) {
-      if (Entry.second->getAsImmutablePass() == nullptr &&
-          !is_contained(PreservedSet, Entry.first)) {
-        // Remove this analysis
-        if (PassDebugging >= Details) {
-          Pass *S = Entry.second;
-          dbgs() << " -- '" <<  P->getPassName() << "' is not preserving '";
-          dbgs() << S->getPassName() << "'\n";
-        }
-        ToRemove.push_back(Entry.first);
-      }
-    }
-    for (AnalysisID ID : ToRemove)
-      IA->erase(ID);
+    IA->remove_if(IsNotPreserved);
   }
 }
 
diff --git a/llvm/lib/MC/MCObjectStreamer.cpp b/llvm/lib/MC/MCObjectStreamer.cpp
index e6b622331036a..2bf5f05c1c315 100644
--- a/llvm/lib/MC/MCObjectStreamer.cpp
+++ b/llvm/lib/MC/MCObjectStreamer.cpp
@@ -277,7 +277,7 @@ void MCObjectStreamer::emitPendingAssignments(MCSymbol *Symbol) {
 
   // emitAssignment can recursively re-enter emitPendingAssignments for
   // other symbols, so move the list out and erase before iterating.
-  SmallVector<PendingAssignment> Pending = std::move(Assignments->second);
+  SmallVector<PendingAssignment, 1> Pending = std::move(Assignments->second);
   pendingAssignments.erase(Assignments);
   for (const PendingAssignment &A : Pending)
     emitAssignment(A.Symbol, A.Value);
diff --git a/llvm/lib/Transforms/IPO/FunctionImport.cpp b/llvm/lib/Transforms/IPO/FunctionImport.cpp
index c55674495646d..d305eadc12f35 100644
--- a/llvm/lib/Transforms/IPO/FunctionImport.cpp
+++ b/llvm/lib/Transforms/IPO/FunctionImport.cpp
@@ -1284,12 +1284,8 @@ void llvm::ComputeCrossModuleImport(
     // exporting module. We do this after the above insertion since we may hit
     // the same ref/call target multiple times in above loop, and it is more
     // efficient to avoid a set lookup each time.
-    SmallVector<ValueInfo> ToRemove;
-    for (const auto &VI : NewExports)
-      if (!DefinedGVSummaries.count(VI.getGUID()))
-        ToRemove.push_back(VI);
-    for (const auto &VI : ToRemove)
-      NewExports.erase(VI);
+    NewExports.remove_if(
+        [&](ValueInfo VI) { return !DefinedGVSummaries.count(VI.getGUID()); });
     ELI.second.insert_range(NewExports);
   }
 
diff --git a/llvm/lib/Transforms/Utils/PromoteMemoryToRegister.cpp b/llvm/lib/Transforms/Utils/PromoteMemoryToRegister.cpp
index 631190a61cc6d..ed0e864fd6905 100644
--- a/llvm/lib/Transforms/Utils/PromoteMemoryToRegister.cpp
+++ b/llvm/lib/Transforms/Utils/PromoteMemoryToRegister.cpp
@@ -929,20 +929,16 @@ void PromoteMem2Reg::run() {
     // simplify and RAUW them as we go.  If it was not, we could add uses to
     // the values we replace with in a non-deterministic order, thus creating
     // non-deterministic def->use chains.
-    SmallVector<std::pair<unsigned, unsigned>> ToRemove;
-    for (auto &Entry : NewPhiNodes) {
+    EliminatedAPHI = NewPhiNodes.remove_if([&](const auto &Entry) {
       PHINode *PN = Entry.second;
-
       // If this PHI node merges one value and/or undefs, get the value.
       if (Value *V = simplifyInstruction(PN, SQ)) {
         PN->replaceAllUsesWith(V);
         PN->eraseFromParent();
-        ToRemove.push_back(Entry.first);
-        EliminatedAPHI = true;
+        return true;
       }
-    }
-    for (auto &K : ToRemove)
-      NewPhiNodes.erase(K);
+      return false;
+    });
   }
 
   // At this point, the renamer has added entries to PHI nodes for all reachable
diff --git a/llvm/unittests/ADT/DenseMapTest.cpp b/llvm/unittests/ADT/DenseMapTest.cpp
index 553d159d33b1a..c38c0709f615a 100644
--- a/llvm/unittests/ADT/DenseMapTest.cpp
+++ b/llvm/unittests/ADT/DenseMapTest.cpp
@@ -1108,4 +1108,49 @@ TEST(DenseMapCustomTest, ValueDtor) {
   EXPECT_EQ(0u, CtorTester::getNumConstructed());
 }
 
+TEST(DenseMapCustomTest, RemoveIf) {
+  // Use enough entries to exercise the large representation and force the
+  // same-size rehash inside remove_if to restore the probe invariant.
+  DenseMap<int, int> Map;
+  for (int I = 0; I < 100; ++I)
+    Map[I] = I * 10;
+
+  // Remove all even keys.
+  EXPECT_TRUE(Map.remove_if([](const auto &E) { return E.first % 2 == 0; }));
+  EXPECT_EQ(Map.size(), 50u);
+  for (int I = 0; I < 100; ++I) {
+    auto It = Map.find(I);
+    if (I % 2 == 0) {
+      EXPECT_EQ(It, Map.end());
+    } else {
+      ASSERT_NE(It, Map.end());
+      EXPECT_EQ(It->second, I * 10);
+    }
+  }
+
+  // A predicate that matches nothing returns false and leaves the map alone.
+  EXPECT_FALSE(Map.remove_if([](const auto &) { return false; }));
+  EXPECT_EQ(Map.size(), 50u);
+
+  // Remove everything.
+  EXPECT_TRUE(Map.remove_if([](const auto &) { return true; }));
+  EXPECT_TRUE(Map.empty());
+}
+
+TEST(DenseMapCustomTest, RemoveIfValueDtor) {
+  // remove_if must destroy the values of removed entries exactly once, and the
+  // rehash must not leak or double-destroy surviving values.
+  EXPECT_EQ(0u, CtorTester::getNumConstructed());
+  {
+    DenseMap<int, CtorTester> Map;
+    for (int I = 0; I < 16; ++I)
+      Map.try_emplace(I, CtorTester(I));
+    EXPECT_EQ(16u, CtorTester::getNumConstructed());
+    EXPECT_TRUE(Map.remove_if([](const auto &E) { return E.first < 10; }));
+    EXPECT_EQ(6u, CtorTester::getNumConstructed());
+    EXPECT_EQ(Map.size(), 6u);
+  }
+  EXPECT_EQ(0u, CtorTester::getNumConstructed());
+}
+
 } // namespace
diff --git a/llvm/unittests/ADT/DenseSetTest.cpp b/llvm/unittests/ADT/DenseSetTest.cpp
index a2a062b151b67..9d214b8649f66 100644
--- a/llvm/unittests/ADT/DenseSetTest.cpp
+++ b/llvm/unittests/ADT/DenseSetTest.cpp
@@ -65,6 +65,22 @@ TEST(SmallDenseSetTest, InsertRange) {
   EXPECT_THAT(set, ::testing::UnorderedElementsAre(7, 8, 9));
 }
 
+TEST(DenseSetTest, RemoveIf) {
+  llvm::DenseSet<unsigned> set;
+  for (unsigned I = 0; I < 100; ++I)
+    set.insert(I);
+
+  EXPECT_TRUE(set.remove_if([](unsigned V) { return V % 2 == 0; }));
+  EXPECT_EQ(set.size(), 50u);
+  for (unsigned I = 0; I < 100; ++I)
+    EXPECT_EQ(set.contains(I), I % 2 == 1);
+
+  EXPECT_FALSE(set.remove_if([](unsigned) { return false; }));
+  EXPECT_EQ(set.size(), 50u);
+  EXPECT_TRUE(set.remove_if([](unsigned) { return true; }));
+  EXPECT_TRUE(set.empty());
+}
+
 struct TestDenseSetInfo {
   static inline unsigned getEmptyKey() { return ~0; }
   static inline unsigned getTombstoneKey() { return ~0U - 1; }



More information about the llvm-commits mailing list