[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