gtsam
Loading...
Searching...
No Matches
Scheduler.h
Go to the documentation of this file.
1/* ----------------------------------------------------------------------------
2 * GTSAM Copyright 2010, Georgia Tech Research Corporation,
3 * Atlanta, Georgia 30332-0415
4 * All Rights Reserved
5 * Authors: Frank Dellaert, et al. (see THANKS for the full author list)
6 * See LICENSE for the license information
7 * -------------------------------------------------------------------------- */
8
24
25#pragma once
26
27#include <atomic>
28#include <condition_variable>
29#include <exception>
30#include <functional>
31#include <future>
32#include <memory>
33#include <mutex>
34#include <stdexcept>
35#include <thread>
36#include <type_traits>
37#include <utility>
38#include <variant>
39#include <vector>
40
41namespace gtsam {
42
62template <typename Y, typename Policy>
63class Scheduler {
64 using Metadata = typename Policy::Metadata;
65
66 struct WorkItem {
67 Metadata metadata{};
68 std::function<void()> run;
69 };
70
71 using Container = typename Policy::template Container<WorkItem>;
72
73 struct WorkerQueue {
74 std::mutex mutex;
75 Container queue;
76 };
77
78 std::vector<std::unique_ptr<WorkerQueue>> queues_; // Per-worker queues.
79 std::atomic<size_t> queuedTasks_{0}; // Tracks queued work for wakeups.
80 std::vector<std::thread> workers_;
81 mutable std::mutex waitMutex_; // Dedicated mutex for condition_variable.
82 std::condition_variable condition_;
83 std::atomic<bool> stop_{false};
84 std::atomic<int> activeTasks_{0}; // In-flight tasks (queued + running).
85 inline static thread_local Scheduler* currentScheduler_ = nullptr;
86 inline static thread_local int workerIndex_ = -1;
87 std::atomic<size_t> nextWorker_{0}; // Round-robin distributor.
88
95 void worker_thread(size_t index) {
96 currentScheduler_ = this;
97 workerIndex_ = static_cast<int>(index);
98 while (true) {
99 WorkItem item;
100 if (!tryPopTask(index, item)) {
101 std::unique_lock<std::mutex> lock(waitMutex_);
102 condition_.wait(lock, [this] {
103 return stop_.load(std::memory_order_acquire) ||
104 queuedTasks_.load(std::memory_order_acquire) > 0;
105 });
106 if (stop_.load(std::memory_order_acquire) &&
107 queuedTasks_.load(std::memory_order_acquire) == 0) {
108 currentScheduler_ = nullptr;
109 workerIndex_ = -1;
110 return;
111 }
112 continue;
113 }
114
115 try {
116 item.run();
117 } catch (...) {
118 /* ignore */
119 }
120
121 // Notify waiters only when work transitions to "done".
122 if (activeTasks_.fetch_sub(1, std::memory_order_release) == 1) {
123 std::lock_guard<std::mutex> lock(waitMutex_);
124 condition_.notify_all();
125 }
126 }
127 }
128
130 bool tryPopTask(size_t index, WorkItem& item) {
131 // Prefer local queue, then steal to keep workers busy.
132 if (tryPopLocal(index, item)) return true;
133 if (trySteal(index, item)) return true;
134 return false;
135 }
136
138 bool tryPopLocal(size_t index, WorkItem& item) {
139 WorkerQueue& queue = *queues_[index];
140 std::lock_guard<std::mutex> lock(queue.mutex);
141 if (!Policy::template popLocal<WorkItem>(queue.queue, item)) return false;
142 queuedTasks_.fetch_sub(1, std::memory_order_release);
143 return true;
144 }
145
147 bool trySteal(size_t index, WorkItem& item) {
148 const size_t workerCount = queues_.size();
149 if (workerCount <= 1) return false;
150 // Steal from other workers if local queue is empty.
151 for (size_t offset = 1; offset < workerCount; ++offset) {
152 size_t target = (index + offset) % workerCount;
153 WorkerQueue& queue = *queues_[target];
154 std::unique_lock<std::mutex> lock(queue.mutex, std::try_to_lock);
155 if (!lock.owns_lock()) continue;
156 if (!Policy::template popSteal<WorkItem>(queue.queue, item)) continue;
157 queuedTasks_.fetch_sub(1, std::memory_order_release);
158 return true;
159 }
160 return false;
161 }
162
164 bool isWorkerThread() const { return currentScheduler_ == this; }
165
166 bool enqueueImpl(Metadata metadata, std::function<void()> work) {
167 if (stop_.load(std::memory_order_acquire)) return false;
168 WorkerQueue* targetQueue = nullptr;
169 if (isWorkerThread()) {
170 targetQueue = queues_[static_cast<size_t>(workerIndex_)].get();
171 } else {
172 const size_t target =
173 nextWorker_.fetch_add(1, std::memory_order_relaxed) % queues_.size();
174 targetQueue = queues_[target].get();
175 }
176
177 {
178 std::lock_guard<std::mutex> lock(targetQueue->mutex);
179 Policy::template push<WorkItem>(
180 targetQueue->queue, WorkItem{std::move(metadata), std::move(work)});
181 }
182
183 queuedTasks_.fetch_add(1, std::memory_order_release);
184 activeTasks_.fetch_add(1, std::memory_order_release);
185 {
186 std::lock_guard<std::mutex> lock(waitMutex_);
187 condition_.notify_one();
188 }
189 return true;
190 }
191
192 template <typename T = Y, typename = std::enable_if_t<std::is_void_v<T>>>
193 void scheduleOrRunInlineImpl(Metadata metadata, std::function<void()> job) {
194 if (stop_.load(std::memory_order_relaxed)) return;
195 if (!isWorkerThread()) {
196 enqueueImpl(std::move(metadata), std::move(job));
197 return;
198 }
199
200 activeTasks_.fetch_add(1, std::memory_order_release);
201 try {
202 job();
203 } catch (...) { /* ignore */
204 }
205 if (activeTasks_.fetch_sub(1, std::memory_order_release) == 1) {
206 std::lock_guard<std::mutex> lock(waitMutex_);
207 condition_.notify_all();
208 }
209 }
210
211 public:
218 explicit Scheduler(size_t numThreads = std::thread::hardware_concurrency()) {
219 if (numThreads == 0) numThreads = 1;
220 queues_.reserve(numThreads);
221 for (size_t i = 0; i < numThreads; ++i) {
222 queues_.push_back(std::make_unique<WorkerQueue>());
223 }
224 for (size_t i = 0; i < numThreads; ++i) {
225 workers_.emplace_back(&Scheduler::worker_thread, this, i);
226 }
227 }
228
236 stop_.store(true, std::memory_order_release);
237 condition_.notify_all();
238 for (std::thread& worker : workers_) {
239 if (worker.joinable()) worker.join();
240 }
241 }
242
251 template <typename P = Policy, typename = std::enable_if_t<std::is_same_v<
252 typename P::Metadata, std::monostate>>>
253 std::future<Y> schedule(std::function<Y()> job) {
254 if (stop_.load(std::memory_order_acquire)) {
255 std::promise<Y> err_promise;
256 err_promise.set_exception(std::make_exception_ptr(
257 std::runtime_error("Scheduler is stopping or stopped.")));
258 return err_promise.get_future();
259 }
260
261 auto promise = std::make_shared<std::promise<Y>>();
262 std::future<Y> future = promise->get_future();
263
264 auto work = [promise, job = std::move(job)]() mutable {
265 try {
266 if constexpr (std::is_void_v<Y>) {
267 job();
268 promise->set_value();
269 } else {
270 promise->set_value(job());
271 }
272 } catch (...) {
273 try {
274 promise->set_exception(std::current_exception());
275 } catch (...) { /* ignore */
276 }
277 }
278 };
279
280 if (!enqueueImpl(Metadata{}, std::move(work))) {
281 try {
282 promise->set_exception(std::make_exception_ptr(
283 std::runtime_error("Scheduler is stopping or stopped.")));
284 } catch (...) { /* ignore */
285 }
286 }
287 return future;
288 }
289
299 template <typename P = Policy, typename = std::enable_if_t<!std::is_same_v<
300 typename P::Metadata, std::monostate>>>
301 std::future<Y> schedule(Metadata metadata, std::function<Y()> job) {
302 if (stop_.load(std::memory_order_acquire)) {
303 std::promise<Y> err_promise;
304 err_promise.set_exception(std::make_exception_ptr(
305 std::runtime_error("Scheduler is stopping or stopped.")));
306 return err_promise.get_future();
307 }
308
309 auto promise = std::make_shared<std::promise<Y>>();
310 std::future<Y> future = promise->get_future();
311
312 auto work = [promise, job = std::move(job)]() mutable {
313 try {
314 if constexpr (std::is_void_v<Y>) {
315 job();
316 promise->set_value();
317 } else {
318 promise->set_value(job());
319 }
320 } catch (...) {
321 try {
322 promise->set_exception(std::current_exception());
323 } catch (...) { /* ignore */
324 }
325 }
326 };
327
328 if (!enqueueImpl(std::move(metadata), std::move(work))) {
329 try {
330 promise->set_exception(std::make_exception_ptr(
331 std::runtime_error("Scheduler is stopping or stopped.")));
332 } catch (...) { /* ignore */
333 }
334 }
335 return future;
336 }
337
343 template <typename T = Y,
344 typename = std::enable_if_t<
345 std::is_void_v<T> && std::is_same_v<Metadata, std::monostate>>>
346 void scheduleOrRunInline(std::function<void()> job) {
347 scheduleOrRunInlineImpl(Metadata{}, std::move(job));
348 }
349
355 template <typename T = Y,
356 typename = std::enable_if_t<
357 std::is_void_v<T> && !std::is_same_v<Metadata, std::monostate>>>
358 void scheduleOrRunInline(Metadata metadata, std::function<void()> job) {
359 scheduleOrRunInlineImpl(std::move(metadata), std::move(job));
360 }
361
367 template <typename T = Y,
368 typename = std::enable_if_t<
369 std::is_void_v<T> && std::is_same_v<Metadata, std::monostate>>>
370 void enqueue(std::function<void()> job) {
371 enqueueImpl(Metadata{}, std::move(job));
372 }
373
379 template <typename T = Y,
380 typename = std::enable_if_t<
381 std::is_void_v<T> && std::is_same_v<Metadata, std::monostate>>>
382 void enqueueOrRunInline(std::function<void()> job) {
383 scheduleOrRunInlineImpl(Metadata{}, std::move(job));
384 }
385
392 std::unique_lock<std::mutex> lock(waitMutex_);
393 condition_.wait(lock, [this] {
394 return stop_.load(std::memory_order_acquire) ||
395 (activeTasks_.load(std::memory_order_acquire) == 0 &&
396 queuedTasks_.load(std::memory_order_acquire) == 0);
397 });
398 }
399};
400
401} // namespace gtsam
Global functions in a separate testing namespace.
Definition chartTesting.h:28
void scheduleOrRunInline(std::function< void()> job)
Schedule or run inline when called from a worker thread.
Definition Scheduler.h:346
std::future< Y > schedule(std::function< Y()> job)
Enqueue a task for execution (no metadata).
Definition Scheduler.h:253
void scheduleOrRunInline(Metadata metadata, std::function< void()> job)
Schedule or run inline when called from a worker thread.
Definition Scheduler.h:358
Scheduler(size_t numThreads=std::thread::hardware_concurrency())
Construct a scheduler with a fixed number of worker threads.
Definition Scheduler.h:218
void enqueue(std::function< void()> job)
Enqueue a fire-and-forget task for execution.
Definition Scheduler.h:370
void enqueueOrRunInline(std::function< void()> job)
Enqueue a fire-and-forget task or run it inline on worker threads.
Definition Scheduler.h:382
~Scheduler()
Wait for all tasks to finish, then stop worker threads.
Definition Scheduler.h:234
std::future< Y > schedule(Metadata metadata, std::function< Y()> job)
Enqueue a task for execution with metadata.
Definition Scheduler.h:301
void waitForAllTasks()
Block until all queued and active tasks complete.
Definition Scheduler.h:391