vkmEngine 1.0.0
A C++ game engine · vkmengine.com
Loading...
Searching...
No Matches
thread_pool.h
1#pragma once
2
3#include <algorithm>
4#include <atomic>
5#include <condition_variable>
6#include <deque>
7#include <exception>
8#include <functional>
9#include <mutex>
10#include <thread>
11#include <type_traits>
12#include <vector>
13
14namespace Vkm::Engine {
15
24class ThreadPool {
25 public:
32 class Batch {
33 public:
41 template<typename Fn>
42 Batch(size_t count, const Fn& task)
43 : m_run([](const void* body, size_t index) { (*static_cast<const Fn*>(body))(index); })
44 , m_task(&task)
45 , m_count(count)
46 {}
47
49 template<typename Fn>
50 Batch(size_t count, const Fn && task) = delete;
51
52 ~Batch() = default;
53
54 Batch(const Batch& other) = delete;
55 Batch& operator=(const Batch& other) = delete;
56
57 Batch(Batch && other) = delete;
58 Batch& operator=(Batch && other) = delete;
59
60 private:
61 friend class ThreadPool;
62
63 private:
64 void (*m_run)(const void* task, size_t index);
65 const void* m_task;
66 size_t m_count;
67 size_t m_claimed = 0;
68 std::exception_ptr m_error;
69
70 std::atomic<size_t> m_pending{0};
71 };
72
73 public:
74 ThreadPool(const ThreadPool& other) = delete;
75 ThreadPool& operator=(const ThreadPool& other) = delete;
76
77 ThreadPool(ThreadPool && other) = delete;
78 ThreadPool& operator=(ThreadPool && other) = delete;
79
80 public:
86 static ThreadPool& get();
87
95 void addTask(std::function<void()> && task);
96
104 void addBatch(Batch& batch);
105
114 void waitForBatch(Batch& batch);
115
123 void shutdown();
124
133 static bool isWorkerThread();
134
135 size_t threadCount() const { return m_threads.size(); }
136
137 private:
143 struct QueuedTask {
144 std::function<void()> function;
145 Batch* batch = nullptr;
146 size_t index = 0;
147 };
148
149 private:
150 ThreadPool(size_t threadCount);
151 ~ThreadPool();
152
156 void process();
157
166 void runIndex(Batch& batch, size_t index);
167
176 void retire(std::atomic<size_t>& pending);
177
178 private:
179 std::atomic<bool> m_running;
180
181 std::vector<std::thread> m_threads;
182
183 std::deque<QueuedTask> m_frameTasks;
184 std::deque<QueuedTask> m_backgroundTasks;
185
186 std::mutex m_tasksMutex;
187 std::condition_variable m_tasksCV;
188 std::condition_variable m_doneCV;
189};
190
203template<class Function>
204void parallelFor(size_t count, size_t grain, Function&& function) {
205 if (count == 0) {
206 return;
207 }
208
209 grain = std::max(grain, size_t(1));
210
211 auto invokeAt = [&](size_t i) {
212 if constexpr (std::is_invocable_v<Function, size_t>) {
213 function(i);
214 } else {
215 function();
216 }
217 };
218
220 for (size_t i = 0; i < count; ++i) invokeAt(i);
221 return;
222 }
223
224 auto& pool = ThreadPool::get();
225
226 // No workers after shutdown (or with no cores reported): a submission would never retire.
227 if (pool.threadCount() == 0) {
228 for (size_t i = 0; i < count; ++i) invokeAt(i);
229 return;
230 }
231
232 // Chunk 0 is the caller's; batch index i is chunk i + 1.
233 const size_t chunks = (count - 1) / grain + 1;
234 const auto runChunk = [&](size_t chunk) {
235 const size_t end = std::min(count, (chunk + 1) * grain);
236 for (size_t index = chunk * grain; index < end; ++index) invokeAt(index);
237 };
238 const auto runQueued = [&](size_t task) { runChunk(task + 1); };
239 ThreadPool::Batch batch(chunks - 1, runQueued);
240
241 pool.addBatch(batch);
242
243 // Queued indices reach into this frame, so join before the caller's exception leaves it;
244 // that exception, not a queued chunk's, propagates.
245 try {
246 runChunk(0);
247 } catch (...) {
248 try {
249 pool.waitForBatch(batch);
250 } catch (...) {
251 }
252 throw;
253 }
254 pool.waitForBatch(batch);
255}
256
266template<class Function>
267void parallelFor(size_t count, Function&& function) {
268 auto& pool = ThreadPool::get();
269
270 // Below this the dispatch cost (mutex, notify_all, done-CV round trip) dwarfs the work;
271 // grain == count sweeps serially.
272 constexpr size_t MIN_PARALLEL = 2048;
273
274 // The +1 is the calling thread, which runs a chunk too.
275 const size_t grain = (count < MIN_PARALLEL)
276 ? count
277 : count / (pool.threadCount() + 1);
278
279 parallelFor(count, grain, function);
280}
281
282} // namespace Vkm::Engine
A run of tasks addressed by index, queued as one entry.
Definition thread_pool.h:32
Batch(size_t count, const Fn &task)
A batch of count tasks, each a call of task with its index.
Definition thread_pool.h:42
Batch(size_t count, const Fn &&task)=delete
Refused: a temporary task is gone before a worker reaches it.
Fixed-size pool of worker threads draining two task queues.
Definition thread_pool.h:24
static bool isWorkerThread()
True when called from a thread owned by the pool.
void shutdown()
Join the workers now, ahead of the pool's own destruction.
void addTask(std::function< void()> &&task)
Enqueue a single task and wake one worker.
void waitForBatch(Batch &batch)
Block the caller until every index of batch has retired.
static ThreadPool & get()
Access the process-wide thread pool, constructed on first use.
void addBatch(Batch &batch)
Queue every index of batch as one entry and wake all workers.