gtsam
Loading...
Searching...
No Matches
ForestTraversal.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
26
27#pragma once
28
29#include <gtsam/config.h>
30#ifdef GTSAM_USE_TBB
32#include <gtsam/base/types.h>
33#include <tbb/global_control.h>
34#else
36#endif
37
38#include <atomic>
39#include <cassert>
40#include <exception>
41#include <functional>
42#include <future>
43#include <memory>
44#include <thread>
45#include <type_traits>
46#include <utility>
47
48namespace gtsam {
49
60template <typename Forest, typename Node>
62 private:
63 size_t threadCount_;
64 using SharedNode = std::shared_ptr<Node>;
65
66 public:
69 size_t numThreads = std::thread::hardware_concurrency())
70 : threadCount_(numThreads == 0 ? 1 : numThreads)
71#ifndef GTSAM_USE_TBB
72 ,
73 scheduler_(threadCount_)
74#endif
75 {
76 }
77
78#ifdef GTSAM_USE_TBB
79 template <typename Fn>
81 void runTopDown(Fn fn, int parallelThreshold = 10) {
82 withTbbTraversalControl([&] {
83 // The pre-order visitor runs before visiting children; parallel task
84 // scheduling is handled by treeTraversal helpers.
85 struct VisitorPre {
86 Fn* fn;
87 int operator()(const SharedNode& node, int&) const {
88 if (node) std::invoke(*fn, *node);
89 return 0;
90 }
91 };
92
93 int rootData = 0;
94 VisitorPre visitor{&fn};
95 auto visitorPost = [](const SharedNode&, int) {};
96 if (threadCount_ == 1) {
97 treeTraversal::DepthFirstForest(static_cast<Forest&>(*this), rootData,
98 visitor, visitorPost);
99 } else {
100 treeTraversal::DepthFirstForestParallel(static_cast<Forest&>(*this),
101 rootData, visitor, visitorPost,
102 parallelThreshold);
103 }
104 });
105 }
106
107 template <typename Fn>
109 void runBottomUp(Fn fn, int parallelThreshold = 10,
110 size_t leafAggregationProblemSize = 0) {
111 withTbbTraversalControl([&] {
112 // The bottom-up visitor runs after all children are processed;
113 // treeTraversal helpers orchestrate the parallelism.
114 struct VisitorPost {
115 Fn* fn;
116 void operator()(const SharedNode& node) const {
117 if (node) std::invoke(*fn, *node);
118 }
119 };
120
121 VisitorPost visitor{&fn};
122 if (threadCount_ == 1) {
123 treeTraversal::PostOrderForest(static_cast<Forest&>(*this), visitor);
124 } else {
125 treeTraversal::PostOrderForestParallel(static_cast<Forest&>(*this),
126 visitor, parallelThreshold,
127 leafAggregationProblemSize);
128 }
129 });
130 }
131
132 private:
134 template <typename Body>
135 void withTbbTraversalControl(Body&& body) {
136 // Set a cap on TBB threads and enter an OpenMP-compatible scope,
137 // then execute the provided traversal body.
138 tbb::global_control control(tbb::global_control::max_allowed_parallelism,
139 static_cast<int>(threadCount_));
140 TbbOpenMPMixedScope threadLimiter;
141 std::forward<Body>(body)();
142 }
143
144#else
145
147 template <typename Fn>
148 void runTopDown(Fn fn, int parallelThreshold = 10) {
149 const auto& roots = getRoots();
150 if (roots.empty()) return;
151 // Create shared traversal state and run from all roots.
152 State state;
153 state.runTraversal([&]() {
154 for (const auto& root : roots) {
155 assert(root);
156 Frame<Fn> frame{&scheduler_, *root, 0, fn, parallelThreshold, &state};
157 frame.topDownDispatch();
158 }
159 });
160 }
161
163 template <typename Fn>
164 void runBottomUp(Fn fn, int parallelThreshold = 10,
165 size_t leafAggregationProblemSize = 0) {
166 const auto& roots = getRoots();
167 if (roots.empty()) return;
168 // Create shared traversal state and run from all roots.
169 State state;
170 state.runTraversal([&]() {
171 for (const auto& root : roots) {
172 assert(root);
173 Frame<Fn> frame{&scheduler_,
174 *root,
175 0,
176 fn,
177 parallelThreshold,
178 &state,
179 leafAggregationProblemSize};
180 frame.bottomUpAsync([] {});
181 }
182 });
183 }
184
185 private:
186 TaskScheduler<void> scheduler_;
187
189 decltype(auto) getRoots() const {
190 const Forest& forest = static_cast<const Forest&>(*this);
191 if constexpr (std::is_member_function_pointer_v<decltype(&Forest::roots)>) {
192 return forest.roots();
193 } else {
194 return (forest.roots);
195 }
196 }
197
199 struct State {
200 std::atomic<int> pending{0};
201 std::atomic_flag exceptionClaim = ATOMIC_FLAG_INIT;
202 std::atomic<bool> hasException{false};
203 std::exception_ptr exception;
204 std::promise<void> done;
205
207 inline int incrementPending() {
208 return pending.fetch_add(1, std::memory_order_relaxed);
209 }
210
212 inline int decrementPending() {
213 return pending.fetch_sub(1, std::memory_order_relaxed);
214 }
215
217 template <typename Body>
218 void runTraversal(Body&& body) {
219 std::future<void> future = done.get_future();
220
221 // Seed pending count for the overall traversal scope, invoke body,
222 // then mark this seed work item as finished.
223 incrementPending();
224 std::forward<Body>(body)();
225 maybeFinish();
226
227 // Block until traversal resolves (success or exception propagation).
228 future.get();
229 }
230
233 void recordException(std::exception_ptr e) {
234 if (!exceptionClaim.test_and_set(std::memory_order_acq_rel)) {
235 exception = e;
236 hasException.store(true, std::memory_order_release);
237 }
238 }
239
242 void maybeFinish() {
243 if (decrementPending() == 1) {
244 if (hasException.load(std::memory_order_acquire)) {
245 try {
246 done.set_exception(exception);
247 } catch (...) { /* ignore */
248 }
249 } else {
250 try {
251 done.set_value();
252 } catch (...) { /* ignore */
253 }
254 }
255 }
256 }
257 };
258 using DoneFn = std::function<void()>;
259
261 struct MaybeFinish {
262 State* state;
263 ~MaybeFinish() { state->maybeFinish(); }
264 };
265
267 template <typename Fn>
268 struct Frame {
269 TaskScheduler<void>* scheduler;
270 Node& node;
271 int depth;
272 const Fn& fn;
273 int threshold;
274 State* state;
275 size_t leafAggregationProblemSize = 0;
276
278 auto& getChildren() const {
279 if constexpr (std::is_member_function_pointer<
280 decltype(&Node::children)>::value) {
281 return node.children();
282 } else {
283 return node.children;
284 }
285 }
286
289 bool shouldParallelize() const {
290 if (threshold <= 0) {
291 return true;
292 } else {
293 return static_cast<int>(node.problemSize()) >= threshold;
294 }
295 }
296
298 inline void topDownDispatch() const {
299 if (shouldParallelize()) {
300 topDownAsync();
301 } else {
302 topDownTraverse();
303 }
304 }
305
307 inline void topDownAsync() const {
308 auto task = [frame = *this]() {
309 MaybeFinish finish{frame.state};
310 frame.topDownTraverse();
311 };
313 state->incrementPending();
314 scheduler->enqueue(std::function<void()>(std::move(task)));
315 }
316
318 inline void topDownTraverse() const {
319 if (state->hasException.load(std::memory_order_acquire))
320 return; // Keep draining without doing new work.
321 try {
322 std::invoke(fn, node);
323 if (!state->hasException.load(std::memory_order_acquire)) {
324 auto&& children = getChildren();
325 for (const auto& child : children) {
326 assert(child);
327 Frame childFrame{scheduler, *child, depth + 1,
328 fn, threshold, state};
329 childFrame.topDownDispatch();
330 }
331 }
332 } catch (...) {
333 state->recordException(std::current_exception());
334 }
335 }
336
338 inline void bottomUpAsync(const DoneFn& onDone) const {
339 auto&& children = getChildren();
340 if (children.empty()) {
341 completeBottomUpNode(onDone);
342 return;
343 }
344 if (leafAggregationProblemSize == 0) {
345 bottomUpUnaggregated(children, onDone);
346 return;
347 }
348
349 // The sentinel prevents inline child completions from finishing this
350 // node before every sibling unit has been dispatched.
351 auto remaining = std::make_shared<std::atomic<int> >(1);
352 std::function<void()> childDone = [frame = *this, remaining,
353 onDone]() mutable {
354 if (remaining->fetch_sub(1, std::memory_order_relaxed) == 1) {
355 frame.completeBottomUpNode(onDone);
356 }
357 };
358
359 for (size_t begin = 0; begin < children.size();) {
360 const auto& child = children[begin];
361 assert(child);
362 Frame childFrame{scheduler,
363 *child,
364 depth + 1,
365 fn,
366 threshold,
367 state,
368 leafAggregationProblemSize};
369 const size_t end = childFrame.getChildren().empty()
370 ? leafBatchEnd(children, begin)
371 : begin + 1;
372
373 remaining->fetch_add(1, std::memory_order_relaxed);
374 if (end > begin + 1) {
375 scheduleLeafBatch(begin, end, childDone);
376 } else {
377 childFrame.bottomUpAsync(childDone);
378 }
379 begin = end;
380 }
381
382 childDone(); // Release the dispatch sentinel.
383 }
384
386 template <typename Children>
387 void bottomUpUnaggregated(Children& children, const DoneFn& onDone) const {
388 auto remaining = std::make_shared<std::atomic<int> >(
389 static_cast<int>(children.size()));
390 std::function<void()> childDone = [frame = *this, remaining,
391 onDone]() mutable {
392 if (remaining->fetch_sub(1, std::memory_order_relaxed) == 1) {
393 frame.completeBottomUpNode(onDone);
394 }
395 };
396
397 for (const auto& child : children) {
398 assert(child);
399 Frame childFrame{scheduler, *child, depth + 1, fn, threshold, state, 0};
400 childFrame.bottomUpAsync(childDone);
401 }
402 }
403
405 template <typename Children>
406 size_t leafBatchEnd(const Children& children, size_t begin) const {
407 size_t end = begin;
408 size_t batchProblemSize = 0;
409 while (end < children.size()) {
410 const auto& child = children[end];
411 assert(child);
412 Frame childFrame{scheduler,
413 *child,
414 depth + 1,
415 fn,
416 threshold,
417 state,
418 leafAggregationProblemSize};
419 if (!childFrame.getChildren().empty()) break;
420
421 const int childProblemSize = child->problemSize();
422 const size_t contribution =
423 childProblemSize > 0 ? static_cast<size_t>(childProblemSize) : 1;
424 if (end > begin &&
425 batchProblemSize + contribution > leafAggregationProblemSize) {
426 break;
427 }
428 batchProblemSize += contribution;
429 ++end;
430 }
431 return end;
432 }
433
435 void scheduleLeafBatch(size_t begin, size_t end,
436 const DoneFn& onDone) const {
437 auto task = [frame = *this, begin, end, onDone]() mutable {
438 MaybeFinish finish{frame.state};
439 if (!frame.state->hasException.load(std::memory_order_acquire)) {
440 try {
441 auto&& children = frame.getChildren();
442 for (size_t index = begin; index < end; ++index) {
443 if (frame.state->hasException.load(std::memory_order_acquire)) {
444 break;
445 }
446 const auto& child = children[index];
447 assert(child);
448 std::invoke(frame.fn, *child);
449 }
450 } catch (...) {
451 frame.state->recordException(std::current_exception());
452 }
453 }
454 frame.callOnDone(onDone);
455 };
456
457 state->incrementPending();
458 scheduler->enqueueOrRunInline(std::function<void()>(std::move(task)));
459 }
460
462 inline void completeBottomUpNode(const DoneFn& onDone) const {
463 if (state->hasException.load(std::memory_order_acquire)) {
464 callOnDone(onDone);
465 } else if (shouldParallelize()) {
466 scheduleBottomUpNode(onDone);
467 } else {
468 bottomUpWork(onDone);
469 }
470 }
471
473 inline void scheduleBottomUpNode(const DoneFn& onDone) const {
474 auto task = [frame = *this, onDone]() mutable {
475 MaybeFinish finish{frame.state};
476 if (frame.state->hasException.load(std::memory_order_acquire)) {
477 frame.callOnDone(onDone);
478 } else {
479 frame.bottomUpWork(onDone);
480 }
481 };
482
483 // Each scheduled task increments the pending counter.
484 state->incrementPending();
485
486 // Schedule a continuation or run it inline if already on a worker thread.
487 scheduler->enqueueOrRunInline(std::function<void()>(std::move(task)));
488 }
489
491 inline void bottomUpWork(const DoneFn& onDone) const {
492 try {
493 std::invoke(fn, node);
494 } catch (...) {
495 state->recordException(std::current_exception());
496 }
497 callOnDone(onDone);
498 }
499
501 void callOnDone(const DoneFn& onDone) const {
502 try {
503 onDone();
504 } catch (...) {
505 state->recordException(std::current_exception());
506 }
507 }
508
509 }; // end Frame
510
511#endif
512};
513
514} // namespace gtsam
Cooperative task scheduler.
Typedefs for easier changing of types.
Global functions in a separate testing namespace.
Definition chartTesting.h:28
Scheduler< Y, detail::TaskSchedulerPolicy > TaskScheduler
Thread pool scheduler that executes tasks without priority ordering.
Definition TaskScheduler.h:82
void PostOrderForest(FOREST &forest, VISITOR_POST &visitorPost)
Traverse a forest depth-first with post-order visits only.
Definition treeTraversal-inst.h:130
void DepthFirstForest(FOREST &forest, DATA &rootData, VISITOR_PRE &visitorPre, VISITOR_POST &visitorPost)
Traverse a forest depth-first with pre-order and post-order visits.
Definition treeTraversal-inst.h:78
void PostOrderForestParallel(FOREST &forest, VISITOR_POST &visitorPost, int problemSizeThreshold=10, size_t leafAggregationProblemSize=0)
Traverse a forest depth-first with post-order visits only (parallel if TBB).
Definition treeTraversal-inst.h:209
void DepthFirstForestParallel(FOREST &forest, DATA &rootData, VISITOR_PRE &visitorPre, VISITOR_POST &visitorPost, int problemSizeThreshold=10)
Traverse a forest depth-first with pre-order and post-order visits.
Definition treeTraversal-inst.h:181
void runTopDown(Fn fn, int parallelThreshold=10)
Scheduler-based top-down traversal.
Definition ForestTraversal.h:148
void runBottomUp(Fn fn, int parallelThreshold=10, size_t leafAggregationProblemSize=0)
Scheduler-based bottom-up traversal.
Definition ForestTraversal.h:164
ForestTraversal(size_t numThreads=std::thread::hardware_concurrency())
Construct a helper with a fixed thread budget (used by TBB when enabled).
Definition ForestTraversal.h:68