[llvm] [Support] Optimize parallel `TaskGroup` (PR #189196)

via llvm-commits llvm-commits at lists.llvm.org
Sat Mar 28 15:22:59 PDT 2026


llvmbot wrote:


<!--LLVM PR SUMMARY COMMENT-->

@llvm/pr-subscribers-llvm-support

Author: Fangrui Song (MaskRay)

<details>
<summary>Changes</summary>

Three improvements to reduce `TaskGroup::spawn()` overhead:

1. Replace mutex-based `Latch::inc()` with atomic `fetch_add`. `dec()`
   retains the mutex to prevent a race where `sync()` observes Count==0
   and destroys the Latch while `dec()` is still running.

2. Pass `Latch&` through `Executor::add()` so the worker calls `dec()`
   directly, eliminating the wrapper lambda that previously captured
   both the user's callable and the Latch reference. This avoids one
   `std::function` construction and potential heap allocation per spawn.

3. Remove the `Executor` abstract base class. `ThreadPoolExecutor` is
   the only implementation. `add()` and `getThreadCount()` are now
   direct calls, inlined by the compiler.


---
Full diff: https://github.com/llvm/llvm-project/pull/189196.diff


2 Files Affected:

- (modified) llvm/include/llvm/Support/Parallel.h (+7-11) 
- (modified) llvm/lib/Support/Parallel.cpp (+30-35) 


``````````diff
diff --git a/llvm/include/llvm/Support/Parallel.h b/llvm/include/llvm/Support/Parallel.h
index b0c9e8f29f970..e528f08433dca 100644
--- a/llvm/include/llvm/Support/Parallel.h
+++ b/llvm/include/llvm/Support/Parallel.h
@@ -58,31 +58,27 @@ inline size_t getThreadCount() { return 1; }
 
 namespace detail {
 class Latch {
-  uint32_t Count;
+  std::atomic<uint32_t> Count;
   mutable std::mutex Mutex;
   mutable std::condition_variable Cond;
 
 public:
   explicit Latch(uint32_t Count = 0) : Count(Count) {}
-  ~Latch() {
-    // Ensure at least that sync() was called.
-    assert(Count == 0);
-  }
+  ~Latch() { assert(Count.load(std::memory_order_relaxed) == 0); }
 
-  void inc() {
-    std::lock_guard<std::mutex> lock(Mutex);
-    ++Count;
-  }
+  void inc() { Count.fetch_add(1, std::memory_order_relaxed); }
 
+  // dec() must hold Mutex so that sync() cannot observe Count==0 and
+  // destroy the Latch while dec() is still running.
   void dec() {
     std::lock_guard<std::mutex> lock(Mutex);
-    if (--Count == 0)
+    if (Count.fetch_sub(1, std::memory_order_acq_rel) == 1)
       Cond.notify_all();
   }
 
   void sync() const {
     std::unique_lock<std::mutex> lock(Mutex);
-    Cond.wait(lock, [&] { return Count == 0; });
+    Cond.wait(lock, [&] { return Count.load(std::memory_order_acquire) == 0; });
   }
 };
 } // namespace detail
diff --git a/llvm/lib/Support/Parallel.cpp b/llvm/lib/Support/Parallel.cpp
index 8f1092e4630dd..dfb6b7d2ca2b4 100644
--- a/llvm/lib/Support/Parallel.cpp
+++ b/llvm/lib/Support/Parallel.cpp
@@ -39,19 +39,8 @@ namespace detail {
 
 namespace {
 
-/// An abstract class that takes closures and runs them asynchronously.
-class Executor {
-public:
-  virtual ~Executor() = default;
-  virtual void add(std::function<void()> func) = 0;
-  virtual size_t getThreadCount() const = 0;
-
-  static Executor *getDefaultExecutor();
-};
-
-/// An implementation of an Executor that runs closures on a thread pool
-///   in filo order.
-class ThreadPoolExecutor : public Executor {
+/// Runs closures on a thread pool in filo order.
+class ThreadPoolExecutor {
 public:
   explicit ThreadPoolExecutor(ThreadPoolStrategy S) {
     if (S.UseJobserver)
@@ -99,7 +88,7 @@ class ThreadPoolExecutor : public Executor {
         T.join();
   }
 
-  ~ThreadPoolExecutor() override { stop(); }
+  ~ThreadPoolExecutor() { stop(); }
 
   struct Creator {
     static void *call() { return new ThreadPoolExecutor(strategy); }
@@ -108,15 +97,15 @@ class ThreadPoolExecutor : public Executor {
     static void call(void *Ptr) { ((ThreadPoolExecutor *)Ptr)->stop(); }
   };
 
-  void add(std::function<void()> F) override {
+  void add(std::function<void()> F, detail::Latch &L) {
     {
       std::lock_guard<std::mutex> Lock(Mutex);
-      WorkStack.push_back(std::move(F));
+      WorkStack.push_back({std::move(F), &L});
     }
     Cond.notify_one();
   }
 
-  size_t getThreadCount() const override { return ThreadCount; }
+  size_t getThreadCount() const { return ThreadCount; }
 
 private:
   void work(ThreadPoolStrategy S, unsigned ThreadID) {
@@ -153,7 +142,7 @@ class ThreadPoolExecutor : public Executor {
             [&] { TheJobserver->release(std::move(Slot)); });
 
         while (true) {
-          std::function<void()> Task;
+          WorkItem Item;
           {
             std::unique_lock<std::mutex> Lock(Mutex);
             Cond.wait(Lock, [&] { return Stop || !WorkStack.empty(); });
@@ -161,26 +150,35 @@ class ThreadPoolExecutor : public Executor {
               return;
             if (WorkStack.empty())
               break;
-            Task = std::move(WorkStack.back());
+            Item = std::move(WorkStack.back());
             WorkStack.pop_back();
           }
-          Task();
+          Item.F();
+          Item.L->dec();
         }
       } else {
-        std::unique_lock<std::mutex> Lock(Mutex);
-        Cond.wait(Lock, [&] { return Stop || !WorkStack.empty(); });
-        if (Stop)
-          break;
-        auto Task = std::move(WorkStack.back());
-        WorkStack.pop_back();
-        Lock.unlock();
-        Task();
+        WorkItem Item;
+        {
+          std::unique_lock<std::mutex> Lock(Mutex);
+          Cond.wait(Lock, [&] { return Stop || !WorkStack.empty(); });
+          if (Stop)
+            break;
+          Item = std::move(WorkStack.back());
+          WorkStack.pop_back();
+        }
+        Item.F();
+        Item.L->dec();
       }
     }
   }
 
+  struct WorkItem {
+    std::function<void()> F;
+    detail::Latch *L;
+  };
+
   std::atomic<bool> Stop{false};
-  std::vector<std::function<void()>> WorkStack;
+  std::vector<WorkItem> WorkStack;
   std::mutex Mutex;
   std::condition_variable Cond;
   std::promise<void> ThreadsCreated;
@@ -190,7 +188,7 @@ class ThreadPoolExecutor : public Executor {
   JobserverClient *TheJobserver = nullptr;
 };
 
-Executor *Executor::getDefaultExecutor() {
+ThreadPoolExecutor *getDefaultExecutor() {
 #ifdef _WIN32
   // The ManagedStatic enables the ThreadPoolExecutor to be stopped via
   // llvm_shutdown() on Windows. This is important to avoid various race
@@ -214,7 +212,7 @@ Executor *Executor::getDefaultExecutor() {
 } // namespace detail
 
 size_t getThreadCount() {
-  return detail::Executor::getDefaultExecutor()->getThreadCount();
+  return detail::getDefaultExecutor()->getThreadCount();
 }
 #endif
 
@@ -239,10 +237,7 @@ void TaskGroup::spawn(std::function<void()> F) {
 #if LLVM_ENABLE_THREADS
   if (Parallel) {
     L.inc();
-    detail::Executor::getDefaultExecutor()->add([&, F = std::move(F)] {
-      F();
-      L.dec();
-    });
+    detail::getDefaultExecutor()->add(std::move(F), L);
     return;
   }
 #endif

``````````

</details>


https://github.com/llvm/llvm-project/pull/189196


More information about the llvm-commits mailing list