11#ifndef EIGEN_THREADPOOL_FORKJOIN_H
12#define EIGEN_THREADPOOL_FORKJOIN_H
15#include "./InternalHeaderCheck.h"
60class ForkJoinScheduler {
69 template <
typename DoFnType,
typename DoneFnType,
typename ThreadPoolEnv>
70 static void ParallelForAsync(Index start, Index end, Index granularity, DoFnType&& do_func, DoneFnType&& done,
71 ThreadPoolTempl<ThreadPoolEnv>* thread_pool) {
76 thread_pool->Schedule([start, end, granularity, thread_pool, do_func = std::forward<DoFnType>(do_func),
77 done = std::forward<DoneFnType>(done)]() {
78 RunParallelFor(start, end, granularity, do_func, thread_pool);
86 template <
typename DoFnType,
typename ThreadPoolEnv>
87 static void ParallelFor(Index start, Index end, Index granularity, DoFnType&& do_func,
88 ThreadPoolTempl<ThreadPoolEnv>* thread_pool) {
89 if (start >= end)
return;
91 auto done = [&barrier]() { barrier.Notify(); };
92 ParallelForAsync(start, end, granularity, std::forward<DoFnType>(do_func), done, thread_pool);
98 template <
typename LeftType,
typename RightType,
typename ThreadPoolEnv>
99 static void ForkJoin(LeftType&& left_thunk, RightType&& right_thunk, ThreadPoolTempl<ThreadPoolEnv>* thread_pool) {
100 using Task =
typename ThreadPoolTempl<ThreadPoolEnv>::Task;
101 std::atomic<bool> right_done(
false);
102 auto execute_right = [&right_thunk, &right_done]() {
103 std::forward<RightType>(right_thunk)();
104 right_done.store(
true, std::memory_order_release);
106 thread_pool->Schedule(std::move(execute_right));
107 std::forward<LeftType>(left_thunk)();
109 while (!right_done.load(std::memory_order_acquire)) {
110 thread_pool->MaybeGetTask(&task);
111 if (task.f) task.f();
115 static Index ComputeMidpoint(Index start, Index end, Index granularity) {
121 const Index size = end - start;
122 const Index offset = numext::round_down(9 * (size + 1) / 16, granularity);
123 return start + offset;
126 template <
typename DoFnType,
typename ThreadPoolEnv>
127 static void RunParallelFor(Index start, Index end, Index granularity, DoFnType&& do_func,
128 ThreadPoolTempl<ThreadPoolEnv>* thread_pool) {
129 Index mid = ComputeMidpoint(start, end, granularity);
130 if ((end - start) < granularity || mid == start || mid == end) {
134 ForkJoin([start, mid, granularity, &do_func,
135 thread_pool]() { RunParallelFor(start, mid, granularity, do_func, thread_pool); },
136 [mid, end, granularity, &do_func, thread_pool]() {
137 RunParallelFor(mid, end, granularity, do_func, thread_pool);