64 using SharedNode = std::shared_ptr<Node>;
69 size_t numThreads = std::thread::hardware_concurrency())
70 : threadCount_(numThreads == 0 ? 1 : numThreads)
73 scheduler_(threadCount_)
79 template <
typename Fn>
81 void runTopDown(Fn fn,
int parallelThreshold = 10) {
82 withTbbTraversalControl([&] {
87 int operator()(
const SharedNode& node,
int&)
const {
88 if (node) std::invoke(*fn, *node);
94 VisitorPre visitor{&fn};
95 auto visitorPost = [](
const SharedNode&, int) {};
96 if (threadCount_ == 1) {
98 visitor, visitorPost);
101 rootData, visitor, visitorPost,
107 template <
typename Fn>
109 void runBottomUp(Fn fn,
int parallelThreshold = 10,
110 size_t leafAggregationProblemSize = 0) {
111 withTbbTraversalControl([&] {
116 void operator()(
const SharedNode& node)
const {
117 if (node) std::invoke(*fn, *node);
121 VisitorPost visitor{&fn};
122 if (threadCount_ == 1) {
126 visitor, parallelThreshold,
127 leafAggregationProblemSize);
134 template <
typename Body>
135 void withTbbTraversalControl(Body&& 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)();
147 template <
typename Fn>
149 const auto& roots = getRoots();
150 if (roots.empty())
return;
153 state.runTraversal([&]() {
154 for (
const auto& root : roots) {
156 Frame<Fn> frame{&scheduler_, *root, 0, fn, parallelThreshold, &state};
157 frame.topDownDispatch();
163 template <
typename Fn>
165 size_t leafAggregationProblemSize = 0) {
166 const auto& roots = getRoots();
167 if (roots.empty())
return;
170 state.runTraversal([&]() {
171 for (
const auto& root : roots) {
173 Frame<Fn> frame{&scheduler_,
179 leafAggregationProblemSize};
180 frame.bottomUpAsync([] {});
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();
194 return (forest.roots);
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;
207 inline int incrementPending() {
208 return pending.fetch_add(1, std::memory_order_relaxed);
212 inline int decrementPending() {
213 return pending.fetch_sub(1, std::memory_order_relaxed);
217 template <
typename Body>
218 void runTraversal(Body&& body) {
219 std::future<void> future = done.get_future();
224 std::forward<Body>(body)();
233 void recordException(std::exception_ptr e) {
234 if (!exceptionClaim.test_and_set(std::memory_order_acq_rel)) {
236 hasException.store(
true, std::memory_order_release);
243 if (decrementPending() == 1) {
244 if (hasException.load(std::memory_order_acquire)) {
246 done.set_exception(exception);
258 using DoneFn = std::function<void()>;
263 ~MaybeFinish() { state->maybeFinish(); }
267 template <
typename Fn>
275 size_t leafAggregationProblemSize = 0;
278 auto& getChildren()
const {
279 if constexpr (std::is_member_function_pointer<
280 decltype(&Node::children)>::value) {
281 return node.children();
283 return node.children;
289 bool shouldParallelize()
const {
290 if (threshold <= 0) {
293 return static_cast<int>(node.problemSize()) >= threshold;
298 inline void topDownDispatch()
const {
299 if (shouldParallelize()) {
307 inline void topDownAsync()
const {
308 auto task = [frame = *
this]() {
309 MaybeFinish finish{frame.state};
310 frame.topDownTraverse();
313 state->incrementPending();
314 scheduler->enqueue(std::function<
void()>(std::move(task)));
318 inline void topDownTraverse()
const {
319 if (state->hasException.load(std::memory_order_acquire))
322 std::invoke(fn, node);
323 if (!state->hasException.load(std::memory_order_acquire)) {
324 auto&& children = getChildren();
325 for (
const auto& child : children) {
327 Frame childFrame{scheduler, *child, depth + 1,
328 fn, threshold, state};
329 childFrame.topDownDispatch();
333 state->recordException(std::current_exception());
338 inline void bottomUpAsync(
const DoneFn& onDone)
const {
339 auto&& children = getChildren();
340 if (children.empty()) {
341 completeBottomUpNode(onDone);
344 if (leafAggregationProblemSize == 0) {
345 bottomUpUnaggregated(children, onDone);
351 auto remaining = std::make_shared<std::atomic<int> >(1);
352 std::function<void()> childDone = [frame = *
this, remaining,
354 if (remaining->fetch_sub(1, std::memory_order_relaxed) == 1) {
355 frame.completeBottomUpNode(onDone);
359 for (
size_t begin = 0; begin < children.size();) {
360 const auto& child = children[begin];
362 Frame childFrame{scheduler,
368 leafAggregationProblemSize};
369 const size_t end = childFrame.getChildren().empty()
370 ? leafBatchEnd(children, begin)
373 remaining->fetch_add(1, std::memory_order_relaxed);
374 if (end > begin + 1) {
375 scheduleLeafBatch(begin, end, childDone);
377 childFrame.bottomUpAsync(childDone);
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,
392 if (remaining->fetch_sub(1, std::memory_order_relaxed) == 1) {
393 frame.completeBottomUpNode(onDone);
397 for (
const auto& child : children) {
399 Frame childFrame{scheduler, *child, depth + 1, fn, threshold, state, 0};
400 childFrame.bottomUpAsync(childDone);
405 template <
typename Children>
406 size_t leafBatchEnd(
const Children& children,
size_t begin)
const {
408 size_t batchProblemSize = 0;
409 while (end < children.size()) {
410 const auto& child = children[end];
412 Frame childFrame{scheduler,
418 leafAggregationProblemSize};
419 if (!childFrame.getChildren().empty())
break;
421 const int childProblemSize = child->problemSize();
422 const size_t contribution =
423 childProblemSize > 0 ?
static_cast<size_t>(childProblemSize) : 1;
425 batchProblemSize + contribution > leafAggregationProblemSize) {
428 batchProblemSize += contribution;
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)) {
441 auto&& children = frame.getChildren();
442 for (
size_t index = begin; index < end; ++index) {
443 if (frame.state->hasException.load(std::memory_order_acquire)) {
446 const auto& child = children[index];
448 std::invoke(frame.fn, *child);
451 frame.state->recordException(std::current_exception());
454 frame.callOnDone(onDone);
457 state->incrementPending();
458 scheduler->enqueueOrRunInline(std::function<
void()>(std::move(task)));
462 inline void completeBottomUpNode(
const DoneFn& onDone)
const {
463 if (state->hasException.load(std::memory_order_acquire)) {
465 }
else if (shouldParallelize()) {
466 scheduleBottomUpNode(onDone);
468 bottomUpWork(onDone);
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);
479 frame.bottomUpWork(onDone);
484 state->incrementPending();
487 scheduler->enqueueOrRunInline(std::function<
void()>(std::move(task)));
491 inline void bottomUpWork(
const DoneFn& onDone)
const {
493 std::invoke(fn, node);
495 state->recordException(std::current_exception());
501 void callOnDone(
const DoneFn& onDone)
const {
505 state->recordException(std::current_exception());