Eigen  5.0.1
 
Loading...
Searching...
No Matches
ForkJoin.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2025 Weiwei Kong <weiweikong@google.com>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_THREADPOOL_FORKJOIN_H
12#define EIGEN_THREADPOOL_FORKJOIN_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19// ForkJoinScheduler provides implementations of various non-blocking ParallelFor algorithms for unary
20// and binary parallel tasks. More specifically, the implementations follow the binary tree-based
21// algorithm from the following paper:
22//
23// Lea, D. (2000, June). A java fork/join framework. *In Proceedings of the
24// ACM 2000 conference on Java Grande* (pp. 36-43).
25//
26// For a given binary task function `f(i,j)` and integers `num_threads`, `granularity`, `start`, and `end`,
27// the implemented parallel for algorithm schedules and executes at most `num_threads` of the functions
28// from the following set in parallel (either synchronously or asynchronously):
29//
30// f(start,start+s_1), f(start+s_1,start+s_2), ..., f(start+s_n,end)
31//
32// where `s_{j+1} - s_{j}` and `end - s_n` are roughly within a factor of two of `granularity`. For a unary
33// task function `g(k)`, the same operation is applied with
34//
35// f(i,j) = [&](){ for(Index k = i; k < j; ++k) g(k); };
36//
37// Note that the parameter `granularity` should be tuned by the user based on the trade-off of running the
38// given task function sequentially vs. scheduling individual tasks in parallel. An example of a partially
39// tuned `granularity` is in `Eigen::CoreThreadPoolDevice::parallelFor(...)` where the template
40// parameter `PacketSize` and float input `cost` are used to indirectly compute a granularity level for a
41// given task function.
42//
43// Example usage #1 (synchronous):
44// ```
45// ThreadPool thread_pool(num_threads);
46// ForkJoinScheduler::ParallelFor(0, num_tasks, granularity, std::move(parallel_task), &thread_pool);
47// ```
48//
49// Example usage #2 (executing multiple tasks asynchronously, each one parallelized with ParallelFor):
50// ```
51// ThreadPool thread_pool(num_threads);
52// Barrier barrier(num_async_calls);
53// auto done = [&](){ barrier.Notify(); };
54// for (Index k=0; k<num_async_calls; ++k) {
55// ForkJoinScheduler::ParallelForAsync(task_start[k], task_end[k], granularity[k], parallel_task[k], done,
56// &thread_pool);
57// }
58// barrier.Wait();
59// ```
60class ForkJoinScheduler {
61 public:
62 // Runs `do_func` asynchronously for the range [start, end) with a specified
63 // granularity. `do_func` should be of type `std::function<void(Index,
64 // Index)>`. `done()` is called exactly once after all tasks have been executed.
65 //
66 // WARNING: like `ParallelFor` below, scheduling nested `ParallelForAsync`
67 // calls (one task body invokes ParallelForAsync on the same pool) can deadlock
68 // because `ForkJoin`'s help-while-waiting loop is not reentrancy-aware.
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) {
72 if (start >= end) {
73 done();
74 return;
75 }
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);
79 done();
80 });
81 }
82
83 // Synchronous variant of ParallelForAsync.
84 // WARNING: Making nested calls to `ParallelFor`, e.g., calling `ParallelFor` inside a task passed into another
85 // `ParallelFor` call, may lead to deadlocks due to how task stealing is implemented.
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;
90 Barrier barrier(1);
91 auto done = [&barrier]() { barrier.Notify(); };
92 ParallelForAsync(start, end, granularity, std::forward<DoFnType>(do_func), done, thread_pool);
93 barrier.Wait();
94 }
95
96 private:
97 // Schedules `right_thunk`, runs `left_thunk`, and runs other tasks until `right_thunk` has finished.
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);
105 };
106 thread_pool->Schedule(std::move(execute_right));
107 std::forward<LeftType>(left_thunk)();
108 Task task;
109 while (!right_done.load(std::memory_order_acquire)) {
110 thread_pool->MaybeGetTask(&task);
111 if (task.f) task.f();
112 }
113 }
114
115 static Index ComputeMidpoint(Index start, Index end, Index granularity) {
116 // Typical workloads choose initial values of `{start, end, granularity}` such that `end - start` and
117 // `granularity` are powers of two. Since modern processors usually implement (2^x)-way
118 // set-associative caches, we minimize the number of cache misses by choosing midpoints that are not
119 // powers of two (to avoid having two addresses in the main memory pointing to the same point in the
120 // cache). More specifically, we choose the midpoint at (roughly) the 9/16 mark.
121 const Index size = end - start;
122 const Index offset = numext::round_down(9 * (size + 1) / 16, granularity);
123 return start + offset;
124 }
125
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) {
131 do_func(start, end);
132 return;
133 }
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);
138 },
139 thread_pool);
140 }
141};
142
143} // namespace Eigen
144
145#endif // EIGEN_THREADPOOL_FORKJOIN_H