Eigen  5.0.1
 
Loading...
Searching...
No Matches
NonBlockingThreadPool.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2016 Dmitry Vyukov <dvyukov@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_NONBLOCKING_THREAD_POOL_H
12#define EIGEN_THREADPOOL_NONBLOCKING_THREAD_POOL_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19template <typename Environment>
20class ThreadPoolTempl : public Eigen::ThreadPoolInterface {
21 public:
22 using Thread = typename Environment::EnvThread;
23 using Task = typename Environment::Task;
24 using Queue = RunQueue<Task, 1024>;
25
26 struct PerThread {
27 ThreadPoolTempl* pool = nullptr; // Parent pool, or null for normal threads.
28 uint64_t rand = 0; // Random generator state.
29 int thread_id = -1; // Worker thread index in pool.
30 };
31
32 struct ThreadData {
33 constexpr ThreadData() = default;
34 std::unique_ptr<Thread> thread;
35 std::atomic<unsigned> steal_partition{0};
36 Queue queue;
37 };
38
39 ThreadPoolTempl(int num_threads, Environment env = Environment()) : ThreadPoolTempl(num_threads, true, env) {}
40
41 ThreadPoolTempl(int num_threads, bool allow_spinning, Environment env = Environment())
42 : env_(env),
43 num_threads_(num_threads),
44 allow_spinning_(allow_spinning),
45 spin_count_(
46 // TODO(dvyukov,rmlarsen): The time spent in NonEmptyQueueIndex() is proportional to num_threads_ and
47 // we assume that new work is scheduled at a constant rate, so we divide `kSpinCount` by number of
48 // threads and number of spinning threads. The constant was picked based on a fair dice roll, tune it.
49 allow_spinning && num_threads > 0 ? kSpinCount / kMaxSpinningThreads / num_threads : 0),
50 thread_data_(num_threads),
51 all_coprimes_(num_threads),
52 waiters_(num_threads),
53 global_steal_partition_(EncodePartition(0, num_threads_)),
54 spinning_state_(0),
55 blocked_(0),
56 done_(false),
57 cancelled_(false),
58 ec_(waiters_) {
59 waiters_.resize(num_threads_);
60 // Calculate coprimes of all numbers [1, num_threads].
61 // Coprimes are used for random walks over all threads in Steal
62 // and NonEmptyQueueIndex. Iteration is based on the fact that if we take
63 // a random starting thread index t and calculate num_threads - 1 subsequent
64 // indices as (t + coprime) % num_threads, we will cover all threads without
65 // repetitions (effectively getting a pseudo-random permutation of thread
66 // indices).
67 eigen_plain_assert(num_threads_ < kMaxThreads);
68 for (int i = 1; i <= num_threads_; ++i) {
69 all_coprimes_.emplace_back(i);
70 ComputeCoprimes(i, &all_coprimes_.back());
71 }
72#ifndef EIGEN_THREAD_LOCAL
73 init_barrier_ = std::make_unique<Barrier>(num_threads_);
74#endif
75 thread_data_.resize(num_threads_);
76 for (int i = 0; i < num_threads_; i++) {
77 SetStealPartition(i, EncodePartition(0, num_threads_));
78 thread_data_[i].thread.reset(env_.CreateThread([this, i]() { WorkerLoop(i); }));
79 }
80#ifndef EIGEN_THREAD_LOCAL
81 // Wait for workers to initialize per_thread_map_. Otherwise we might race
82 // with them in Schedule or CurrentThreadId.
83 init_barrier_->Wait();
84#endif
85 }
86
87 ~ThreadPoolTempl() {
88 done_ = true;
89
90 // Now if all threads block without work, they will start exiting.
91 // But note that threads can continue to work arbitrary long,
92 // block, submit new work, unblock and otherwise live full life.
93 if (!cancelled_) {
94 ec_.Notify(true);
95 } else {
96 // Since we were cancelled, there might be entries in the queues.
97 // Empty them to prevent their destructor from asserting.
98 for (size_t i = 0; i < thread_data_.size(); i++) {
99 thread_data_[i].queue.Flush();
100 }
101 }
102 // Join threads explicitly (by destroying) to avoid destruction order within
103 // this class.
104 for (size_t i = 0; i < thread_data_.size(); ++i) thread_data_[i].thread.reset();
105 }
106
107 void SetStealPartitions(const std::vector<std::pair<unsigned, unsigned>>& partitions) {
108 eigen_plain_assert(partitions.size() == static_cast<std::size_t>(num_threads_));
109
110 // Pass this information to each thread queue.
111 for (int i = 0; i < num_threads_; i++) {
112 const auto& pair = partitions[i];
113 unsigned start = pair.first, end = pair.second;
114 AssertBounds(start, end);
115 unsigned val = EncodePartition(start, end);
116 SetStealPartition(i, val);
117 }
118 }
119
120 void Schedule(std::function<void()> fn) EIGEN_OVERRIDE { ScheduleWithHint(std::move(fn), 0, num_threads_); }
121
122 void ScheduleWithHint(std::function<void()> fn, int start, int limit) override {
123 Task t = env_.CreateTask(std::move(fn));
124 PerThread* pt = GetPerThread();
125 if (pt->pool == this) {
126 // Worker thread of this pool, push onto the thread's queue.
127 Queue& q = thread_data_[pt->thread_id].queue;
128 t = q.PushFront(std::move(t));
129 } else {
130 // A free-standing thread (or worker of another pool), push onto a random
131 // queue.
132 eigen_plain_assert(start < limit);
133 eigen_plain_assert(limit <= num_threads_);
134 int num_queues = limit - start;
135 int rnd = Rand(&pt->rand) % num_queues;
136 eigen_plain_assert(start + rnd < limit);
137 Queue& q = thread_data_[start + rnd].queue;
138 t = q.PushBack(std::move(t));
139 }
140 // Note: below we touch this after making t available to worker threads.
141 // Strictly speaking, this can lead to a racy-use-after-free. Consider that
142 // Schedule is called from a thread that is neither main thread nor a worker
143 // thread of this pool. Then, execution of t directly or indirectly
144 // completes overall computations, which in turn leads to destruction of
145 // this. We expect that such scenario is prevented by program, that is,
146 // this is kept alive while any threads can potentially be in Schedule.
147 if (!t.f) {
148 if (IsNotifyParkedThreadRequired()) {
149 ec_.Notify(false);
150 }
151 } else {
152 env_.ExecuteTask(t); // Push failed, execute directly.
153 }
154 }
155
156 // Tries to assign work to the current task.
157 void MaybeGetTask(Task* t) {
158 PerThread* pt = GetPerThread();
159 const int thread_id = pt->thread_id;
160 // If we are not a worker thread of this pool, we can't get any work.
161 if (thread_id < 0) return;
162 Queue& q = thread_data_[thread_id].queue;
163 *t = q.PopFront();
164 if (t->f) return;
165 if (num_threads_ == 1) {
166 // For num_threads_ == 1 there is no point in going through the expensive
167 // steal loop. Moreover, since NonEmptyQueueIndex() calls PopBack() on the
168 // victim queues it might reverse the order in which ops are executed
169 // compared to the order in which they are scheduled, which tends to be
170 // counter-productive for the types of I/O workloads single thread pools
171 // tend to be used for.
172 for (int i = 0; i < spin_count_ && !t->f; ++i) *t = q.PopFront();
173 } else {
174 if (EIGEN_PREDICT_FALSE(!t->f)) *t = LocalSteal();
175 if (EIGEN_PREDICT_FALSE(!t->f)) *t = GlobalSteal();
176 if (EIGEN_PREDICT_FALSE(!t->f)) {
177 if (allow_spinning_ && StartSpinning()) {
178 for (int i = 0; i < spin_count_ && !t->f; ++i) *t = GlobalSteal();
179 // Notify `spinning_state_` that we are no longer spinning.
180 bool has_no_notify_task = StopSpinning();
181 // A suppressed notification can belong to a different task than the
182 // one we just stole. Pass it on before executing a potentially blocking task.
183 if (has_no_notify_task) {
184 if (t->f) {
185 if (NonEmptyQueueIndex() != -1) ec_.Notify(false);
186 } else {
187 *t = GlobalSteal();
188 }
189 }
190 }
191 }
192 }
193 }
194
195 void Cancel() EIGEN_OVERRIDE {
196 cancelled_ = true;
197 done_ = true;
198
199 // Let each thread know it's been cancelled.
200#ifdef EIGEN_THREAD_ENV_SUPPORTS_CANCELLATION
201 for (size_t i = 0; i < thread_data_.size(); i++) {
202 thread_data_[i].thread->OnCancel();
203 }
204#endif
205
206 // Wake up the threads without work to let them exit on their own.
207 ec_.Notify(true);
208 }
209
210 int NumThreads() const EIGEN_FINAL { return num_threads_; }
211
212 int CurrentThreadId() const EIGEN_FINAL {
213 const PerThread* pt = const_cast<ThreadPoolTempl*>(this)->GetPerThread();
214 if (pt->pool == this) {
215 return pt->thread_id;
216 } else {
217 return -1;
218 }
219 }
220
221 private:
222 // Create a single atomic<int> that encodes start and limit information for
223 // each thread.
224 // We expect num_threads_ < 65536, so we can store them in a single
225 // std::atomic<unsigned>.
226 // Exposed publicly as static functions so that external callers can reuse
227 // this encode/decode logic for maintaining their own thread-safe copies of
228 // scheduling and steal domain(s).
229 static constexpr int kMaxPartitionBits = 16;
230 static constexpr int kMaxThreads = 1 << kMaxPartitionBits;
231
232 inline unsigned EncodePartition(unsigned start, unsigned limit) { return (start << kMaxPartitionBits) | limit; }
233
234 inline void DecodePartition(unsigned val, unsigned* start, unsigned* limit) {
235 *limit = val & (kMaxThreads - 1);
236 val >>= kMaxPartitionBits;
237 *start = val;
238 }
239
240 void AssertBounds(int start, int end) {
241 eigen_plain_assert(start >= 0);
242 eigen_plain_assert(start < end); // non-zero sized partition
243 eigen_plain_assert(end <= num_threads_);
244 (void)start;
245 (void)end;
246 }
247
248 inline void SetStealPartition(size_t i, unsigned val) {
249 thread_data_[i].steal_partition.store(val, std::memory_order_relaxed);
250 }
251
252 inline unsigned GetStealPartition(int i) { return thread_data_[i].steal_partition.load(std::memory_order_relaxed); }
253
254 void ComputeCoprimes(int N, MaxSizeVector<unsigned>* coprimes) {
255 for (int i = 1; i <= N; i++) {
256 unsigned a = i;
257 unsigned b = N;
258 // If GCD(a, b) == 1, then a and b are coprimes.
259 while (b != 0) {
260 unsigned tmp = a;
261 a = b;
262 b = tmp % b;
263 }
264 if (a == 1) {
265 coprimes->push_back(i);
266 }
267 }
268 }
269
270 // Maximum number of threads that can spin in steal loop.
271 static constexpr int kMaxSpinningThreads = 1;
272
273 // The number of steal loop spin iterations before parking (this number is
274 // divided by the number of threads, to get spin count for each thread).
275 static constexpr int kSpinCount = 5000;
276
277 // If there are enough active threads with empty pending-task queues, a thread
278 // that runs out of work can just be parked without spinning, because these
279 // active threads will go into a steal loop after finishing their current
280 // tasks.
281 //
282 // In the worst case when all active threads are executing long/expensive
283 // tasks, the next Schedule() will have to wait until one of the parked
284 // threads will be unparked, however this should be very rare in practice.
285 static constexpr int kMinActiveThreadsToStartSpinning = 4;
286
287 struct SpinningState {
288 // Spinning state layout:
289 //
290 // - Low 32 bits encode the number of threads that are spinning in steal
291 // loop.
292 //
293 // - High 32 bits encode the number of tasks that were submitted to the pool
294 // without a call to `ec_.Notify()`. This number can't be larger than
295 // the number of spinning threads. Each spinning thread, when it exits the
296 // spin loop must check if this number is greater than zero, and maybe
297 // make another attempt to steal a task and decrement it by one.
298 static constexpr uint64_t kNumSpinningMask = 0x00000000FFFFFFFF;
299 static constexpr uint64_t kNumNoNotifyMask = 0xFFFFFFFF00000000;
300 static constexpr uint64_t kNumNoNotifyShift = 32;
301
302 uint64_t num_spinning; // number of spinning threads
303 uint64_t num_no_notification; // number of tasks submitted without
304 // notifying waiting threads
305
306 // Decodes `spinning_state_` value.
307 static SpinningState Decode(uint64_t state) {
308 uint64_t num_spinning = (state & kNumSpinningMask);
309 uint64_t num_no_notification = (state & kNumNoNotifyMask) >> kNumNoNotifyShift;
310
311 eigen_plain_assert(num_no_notification <= num_spinning);
312 return {num_spinning, num_no_notification};
313 }
314
315 // Encodes as `spinning_state_` value.
316 uint64_t Encode() const {
317 eigen_plain_assert(num_no_notification <= num_spinning);
318 return (num_no_notification << kNumNoNotifyShift) | num_spinning;
319 }
320 };
321
322 Environment env_;
323 const int num_threads_;
324 const bool allow_spinning_;
325 const int spin_count_;
326 MaxSizeVector<ThreadData> thread_data_;
327 MaxSizeVector<MaxSizeVector<unsigned>> all_coprimes_;
328 MaxSizeVector<EventCount::Waiter> waiters_;
329 unsigned global_steal_partition_;
330 std::atomic<uint64_t> spinning_state_;
331 std::atomic<unsigned> blocked_;
332 std::atomic<bool> done_;
333 std::atomic<bool> cancelled_;
334 EventCount ec_;
335#ifndef EIGEN_THREAD_LOCAL
336 std::unique_ptr<Barrier> init_barrier_;
337 EIGEN_MUTEX per_thread_map_mutex_; // Protects per_thread_map_.
338 std::unordered_map<uint64_t, std::unique_ptr<PerThread>> per_thread_map_;
339#endif
340
341 unsigned NumActiveThreads() const { return num_threads_ - blocked_.load(); }
342
343 // Main worker thread loop.
344 void WorkerLoop(int thread_id) {
345#ifndef EIGEN_THREAD_LOCAL
346 auto new_pt = std::make_unique<PerThread>();
347 per_thread_map_mutex_.lock();
348 bool insertOK = per_thread_map_.emplace(GlobalThreadIdHash(), std::move(new_pt)).second;
349 eigen_plain_assert(insertOK);
350 EIGEN_UNUSED_VARIABLE(insertOK);
351 per_thread_map_mutex_.unlock();
352 init_barrier_->Notify();
353 init_barrier_->Wait();
354#endif
355 PerThread* pt = GetPerThread();
356 pt->pool = this;
357 pt->rand = GlobalThreadIdHash();
358 pt->thread_id = thread_id;
359 Task t;
360 while (!cancelled_.load(std::memory_order_relaxed)) {
361 MaybeGetTask(&t);
362 // If we still don't have a task, wait for one. Return if thread pool is
363 // in cancelled state.
364 if (EIGEN_PREDICT_FALSE(!t.f)) {
365 EventCount::Waiter* waiter = &waiters_[pt->thread_id];
366 if (!WaitForWork(waiter, &t)) return;
367 }
368 if (EIGEN_PREDICT_TRUE(t.f)) env_.ExecuteTask(t);
369 }
370 }
371
372 // Steal tries to steal work from other worker threads in the range [start,
373 // limit) in best-effort manner.
374 Task Steal(unsigned start, unsigned limit) {
375 PerThread* pt = GetPerThread();
376 const size_t size = limit - start;
377 unsigned r = Rand(&pt->rand);
378 // Reduce r into [0, size) range, this utilizes trick from
379 // https://lemire.me/blog/2016/06/27/a-fast-alternative-to-the-modulo-reduction/
380 eigen_plain_assert(all_coprimes_[size - 1].size() < (1 << 30));
381 unsigned victim = ((uint64_t)r * (uint64_t)size) >> 32;
382 unsigned index = ((uint64_t)all_coprimes_[size - 1].size() * (uint64_t)r) >> 32;
383 unsigned inc = all_coprimes_[size - 1][index];
384
385 for (unsigned i = 0; i < size; i++) {
386 eigen_plain_assert(start + victim < limit);
387 Task t = thread_data_[start + victim].queue.PopBack();
388 if (t.f) {
389 return t;
390 }
391 victim += inc;
392 if (victim >= size) {
393 victim -= static_cast<unsigned int>(size);
394 }
395 }
396 return Task();
397 }
398
399 // Steals work within threads belonging to the partition.
400 Task LocalSteal() {
401 PerThread* pt = GetPerThread();
402 unsigned partition = GetStealPartition(pt->thread_id);
403 // If thread steal partition is the same as global partition, there is no
404 // need to go through the steal loop twice.
405 if (global_steal_partition_ == partition) return Task();
406 unsigned start, limit;
407 DecodePartition(partition, &start, &limit);
408 AssertBounds(start, limit);
409
410 return Steal(start, limit);
411 }
412
413 // Steals work from any other thread in the pool.
414 Task GlobalSteal() { return Steal(0, num_threads_); }
415
416 // WaitForWork blocks until new work is available (returns true), or if it is
417 // time to exit (returns false). Can optionally return a task to execute in t
418 // (in such case t.f != nullptr on return).
419 bool WaitForWork(EventCount::Waiter* waiter, Task* t) {
420 eigen_plain_assert(!t->f);
421 // We already did best-effort emptiness check in Steal, so prepare for
422 // blocking.
423 ec_.Prewait();
424 // Now do a reliable emptiness check.
425 int victim = NonEmptyQueueIndex();
426 if (victim != -1) {
427 ec_.CancelWait();
428 if (cancelled_) {
429 return false;
430 } else {
431 *t = thread_data_[victim].queue.PopBack();
432 return true;
433 }
434 }
435 // Number of blocked threads is used as termination condition.
436 // If we are shutting down and all worker threads blocked without work,
437 // that's we are done.
438 blocked_++;
439 // TODO: is blocked_ required to be unsigned?
440 if (done_ && blocked_ == static_cast<unsigned>(num_threads_)) {
441 ec_.CancelWait();
442 // Almost done, but need to re-check queues.
443 // Consider that all queues are empty and all worker threads are preempted
444 // right after incrementing blocked_ above. Now a free-standing thread
445 // submits work and calls destructor (which sets done_). If we don't
446 // re-check queues, we will exit leaving the work unexecuted.
447 if (NonEmptyQueueIndex() != -1) {
448 // Note: we must not pop from queues before we decrement blocked_,
449 // otherwise the following scenario is possible. Consider that instead
450 // of checking for emptiness we popped the only element from queues.
451 // Now other worker threads can start exiting, which is bad if the
452 // work item submits other work. So we just check emptiness here,
453 // which ensures that all worker threads exit at the same time.
454 blocked_--;
455 return true;
456 }
457 // Reached stable termination state.
458 ec_.Notify(true);
459 return false;
460 }
461 ec_.CommitWait(waiter);
462 blocked_--;
463 return true;
464 }
465
466 int NonEmptyQueueIndex() {
467 PerThread* pt = GetPerThread();
468 // We intentionally design NonEmptyQueueIndex to steal work from
469 // anywhere in the queue so threads don't block in WaitForWork() forever
470 // when all threads in their partition go to sleep. Steal is still local.
471 const size_t size = thread_data_.size();
472 unsigned r = Rand(&pt->rand);
473 unsigned inc = all_coprimes_[size - 1][r % all_coprimes_[size - 1].size()];
474 unsigned victim = r % size;
475 for (unsigned i = 0; i < size; i++) {
476 if (!thread_data_[victim].queue.Empty()) {
477 return victim;
478 }
479 victim += inc;
480 if (victim >= size) {
481 victim -= static_cast<unsigned int>(size);
482 }
483 }
484 return -1;
485 }
486
487 // StartSpinning() checks if the number of threads in the spin loop is less
488 // than the allowed maximum. If so, increments the number of spinning threads
489 // by one and returns true (caller must enter the spin loop). Otherwise
490 // returns false, and the caller must not enter the spin loop.
491 bool StartSpinning() {
492 if (NumActiveThreads() > kMinActiveThreadsToStartSpinning) return false;
493
494 uint64_t spinning = spinning_state_.load(std::memory_order_relaxed);
495 for (;;) {
496 SpinningState state = SpinningState::Decode(spinning);
497
498 if ((state.num_spinning - state.num_no_notification) >= kMaxSpinningThreads) {
499 return false;
500 }
501
502 // Increment the number of spinning threads.
503 ++state.num_spinning;
504
505 if (spinning_state_.compare_exchange_weak(spinning, state.Encode(), std::memory_order_relaxed)) {
506 return true;
507 }
508 }
509 }
510
511 // StopSpinning() decrements the number of spinning threads by one. It also
512 // checks if there were any tasks submitted into the pool without notifying
513 // parked threads, and decrements the count by one. Returns true if the number
514 // of tasks submitted without notification was decremented. In this case,
515 // caller must either steal again or notify another worker.
516 bool StopSpinning() {
517 uint64_t spinning = spinning_state_.load(std::memory_order_relaxed);
518 for (;;) {
519 SpinningState state = SpinningState::Decode(spinning);
520
521 // Decrement the number of spinning threads.
522 --state.num_spinning;
523
524 // Maybe decrement the number of tasks submitted without notification.
525 bool has_no_notify_task = state.num_no_notification > 0;
526 if (has_no_notify_task) --state.num_no_notification;
527
528 // Acquire the queue publication paired with the suppressed notification.
529 if (spinning_state_.compare_exchange_weak(spinning, state.Encode(), std::memory_order_acquire)) {
530 return has_no_notify_task;
531 }
532 }
533 }
534
535 // IsNotifyParkedThreadRequired() returns true if parked thread must be
536 // notified about new added task. If there are threads spinning in the steal
537 // loop, there is no need to unpark any of the waiting threads, the task will
538 // be picked up by one of the spinning threads.
539 bool IsNotifyParkedThreadRequired() {
540 uint64_t spinning = spinning_state_.load(std::memory_order_relaxed);
541 for (;;) {
542 SpinningState state = SpinningState::Decode(spinning);
543
544 // If the number of tasks submitted without notifying parked threads is
545 // equal to the number of spinning threads, we must wake up one of the
546 // parked threads.
547 if (state.num_no_notification == state.num_spinning) return true;
548
549 // Increment the number of tasks submitted without notification.
550 ++state.num_no_notification;
551
552 // Publish the queued task to the spinner that consumes this notification.
553 if (spinning_state_.compare_exchange_weak(spinning, state.Encode(), std::memory_order_release,
554 std::memory_order_relaxed)) {
555 return false;
556 }
557 }
558 }
559
560 static EIGEN_STRONG_INLINE uint64_t GlobalThreadIdHash() {
561 return std::hash<std::thread::id>()(std::this_thread::get_id());
562 }
563
564 EIGEN_STRONG_INLINE PerThread* GetPerThread() {
565#ifndef EIGEN_THREAD_LOCAL
566 static PerThread dummy;
567 auto it = per_thread_map_.find(GlobalThreadIdHash());
568 if (it == per_thread_map_.end()) {
569 return &dummy;
570 } else {
571 return it->second.get();
572 }
573#else
574 EIGEN_THREAD_LOCAL PerThread per_thread_;
575 PerThread* pt = &per_thread_;
576 return pt;
577#endif
578 }
579
580 static EIGEN_STRONG_INLINE unsigned Rand(uint64_t* state) {
581 uint64_t current = *state;
582 // Update the internal state
583 *state = current * 6364136223846793005ULL + 0xda3e39cb94b95bdbULL;
584 // Generate the random output (using the PCG-XSH-RS scheme)
585 return static_cast<unsigned>((current ^ (current >> 22)) >> (22 + (current >> 61)));
586 }
587};
588
589using ThreadPool = ThreadPoolTempl<StlThreadEnvironment>;
590
591} // namespace Eigen
592
593#endif // EIGEN_THREADPOOL_NONBLOCKING_THREAD_POOL_H