11#ifndef EIGEN_THREADPOOL_NONBLOCKING_THREAD_POOL_H
12#define EIGEN_THREADPOOL_NONBLOCKING_THREAD_POOL_H
15#include "./InternalHeaderCheck.h"
19template <
typename Environment>
20class ThreadPoolTempl :
public Eigen::ThreadPoolInterface {
22 using Thread =
typename Environment::EnvThread;
23 using Task =
typename Environment::Task;
24 using Queue = RunQueue<Task, 1024>;
27 ThreadPoolTempl* pool =
nullptr;
33 constexpr ThreadData() =
default;
34 std::unique_ptr<Thread> thread;
35 std::atomic<unsigned> steal_partition{0};
39 ThreadPoolTempl(
int num_threads, Environment env = Environment()) : ThreadPoolTempl(num_threads, true, env) {}
41 ThreadPoolTempl(
int num_threads,
bool allow_spinning, Environment env = Environment())
43 num_threads_(num_threads),
44 allow_spinning_(allow_spinning),
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_)),
59 waiters_.resize(num_threads_);
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());
72#ifndef EIGEN_THREAD_LOCAL
73 init_barrier_ = std::make_unique<Barrier>(num_threads_);
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); }));
80#ifndef EIGEN_THREAD_LOCAL
83 init_barrier_->Wait();
98 for (
size_t i = 0; i < thread_data_.size(); i++) {
99 thread_data_[i].queue.Flush();
104 for (
size_t i = 0; i < thread_data_.size(); ++i) thread_data_[i].thread.reset();
107 void SetStealPartitions(
const std::vector<std::pair<unsigned, unsigned>>& partitions) {
108 eigen_plain_assert(partitions.size() ==
static_cast<std::size_t
>(num_threads_));
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);
120 void Schedule(std::function<
void()> fn) EIGEN_OVERRIDE { ScheduleWithHint(std::move(fn), 0, num_threads_); }
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) {
127 Queue& q = thread_data_[pt->thread_id].queue;
128 t = q.PushFront(std::move(t));
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));
148 if (IsNotifyParkedThreadRequired()) {
157 void MaybeGetTask(Task* t) {
158 PerThread* pt = GetPerThread();
159 const int thread_id = pt->thread_id;
161 if (thread_id < 0)
return;
162 Queue& q = thread_data_[thread_id].queue;
165 if (num_threads_ == 1) {
172 for (
int i = 0; i < spin_count_ && !t->f; ++i) *t = q.PopFront();
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();
180 bool has_no_notify_task = StopSpinning();
183 if (has_no_notify_task) {
185 if (NonEmptyQueueIndex() != -1) ec_.Notify(
false);
195 void Cancel() EIGEN_OVERRIDE {
200#ifdef EIGEN_THREAD_ENV_SUPPORTS_CANCELLATION
201 for (
size_t i = 0; i < thread_data_.size(); i++) {
202 thread_data_[i].thread->OnCancel();
210 int NumThreads() const EIGEN_FINAL {
return num_threads_; }
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;
229 static constexpr int kMaxPartitionBits = 16;
230 static constexpr int kMaxThreads = 1 << kMaxPartitionBits;
232 inline unsigned EncodePartition(
unsigned start,
unsigned limit) {
return (start << kMaxPartitionBits) | limit; }
234 inline void DecodePartition(
unsigned val,
unsigned* start,
unsigned* limit) {
235 *limit = val & (kMaxThreads - 1);
236 val >>= kMaxPartitionBits;
240 void AssertBounds(
int start,
int end) {
241 eigen_plain_assert(start >= 0);
242 eigen_plain_assert(start < end);
243 eigen_plain_assert(end <= num_threads_);
248 inline void SetStealPartition(
size_t i,
unsigned val) {
249 thread_data_[i].steal_partition.store(val, std::memory_order_relaxed);
252 inline unsigned GetStealPartition(
int i) {
return thread_data_[i].steal_partition.load(std::memory_order_relaxed); }
254 void ComputeCoprimes(
int N, MaxSizeVector<unsigned>* coprimes) {
255 for (
int i = 1; i <= N; i++) {
265 coprimes->push_back(i);
271 static constexpr int kMaxSpinningThreads = 1;
275 static constexpr int kSpinCount = 5000;
285 static constexpr int kMinActiveThreadsToStartSpinning = 4;
287 struct SpinningState {
298 static constexpr uint64_t kNumSpinningMask = 0x00000000FFFFFFFF;
299 static constexpr uint64_t kNumNoNotifyMask = 0xFFFFFFFF00000000;
300 static constexpr uint64_t kNumNoNotifyShift = 32;
302 uint64_t num_spinning;
303 uint64_t num_no_notification;
307 static SpinningState Decode(uint64_t state) {
308 uint64_t num_spinning = (state & kNumSpinningMask);
309 uint64_t num_no_notification = (state & kNumNoNotifyMask) >> kNumNoNotifyShift;
311 eigen_plain_assert(num_no_notification <= num_spinning);
312 return {num_spinning, num_no_notification};
316 uint64_t Encode()
const {
317 eigen_plain_assert(num_no_notification <= num_spinning);
318 return (num_no_notification << kNumNoNotifyShift) | num_spinning;
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_;
335#ifndef EIGEN_THREAD_LOCAL
336 std::unique_ptr<Barrier> init_barrier_;
337 EIGEN_MUTEX per_thread_map_mutex_;
338 std::unordered_map<uint64_t, std::unique_ptr<PerThread>> per_thread_map_;
341 unsigned NumActiveThreads()
const {
return num_threads_ - blocked_.load(); }
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();
355 PerThread* pt = GetPerThread();
357 pt->rand = GlobalThreadIdHash();
358 pt->thread_id = thread_id;
360 while (!cancelled_.load(std::memory_order_relaxed)) {
364 if (EIGEN_PREDICT_FALSE(!t.f)) {
365 EventCount::Waiter* waiter = &waiters_[pt->thread_id];
366 if (!WaitForWork(waiter, &t))
return;
368 if (EIGEN_PREDICT_TRUE(t.f)) env_.ExecuteTask(t);
374 Task Steal(
unsigned start,
unsigned limit) {
375 PerThread* pt = GetPerThread();
376 const size_t size = limit - start;
377 unsigned r = Rand(&pt->rand);
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];
385 for (
unsigned i = 0; i < size; i++) {
386 eigen_plain_assert(start + victim < limit);
387 Task t = thread_data_[start + victim].queue.PopBack();
392 if (victim >= size) {
393 victim -=
static_cast<unsigned int>(size);
401 PerThread* pt = GetPerThread();
402 unsigned partition = GetStealPartition(pt->thread_id);
405 if (global_steal_partition_ == partition)
return Task();
406 unsigned start, limit;
407 DecodePartition(partition, &start, &limit);
408 AssertBounds(start, limit);
410 return Steal(start, limit);
414 Task GlobalSteal() {
return Steal(0, num_threads_); }
419 bool WaitForWork(EventCount::Waiter* waiter, Task* t) {
420 eigen_plain_assert(!t->f);
425 int victim = NonEmptyQueueIndex();
431 *t = thread_data_[victim].queue.PopBack();
440 if (done_ && blocked_ ==
static_cast<unsigned>(num_threads_)) {
447 if (NonEmptyQueueIndex() != -1) {
461 ec_.CommitWait(waiter);
466 int NonEmptyQueueIndex() {
467 PerThread* pt = GetPerThread();
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()) {
480 if (victim >= size) {
481 victim -=
static_cast<unsigned int>(size);
491 bool StartSpinning() {
492 if (NumActiveThreads() > kMinActiveThreadsToStartSpinning)
return false;
494 uint64_t spinning = spinning_state_.load(std::memory_order_relaxed);
496 SpinningState state = SpinningState::Decode(spinning);
498 if ((state.num_spinning - state.num_no_notification) >= kMaxSpinningThreads) {
503 ++state.num_spinning;
505 if (spinning_state_.compare_exchange_weak(spinning, state.Encode(), std::memory_order_relaxed)) {
516 bool StopSpinning() {
517 uint64_t spinning = spinning_state_.load(std::memory_order_relaxed);
519 SpinningState state = SpinningState::Decode(spinning);
522 --state.num_spinning;
525 bool has_no_notify_task = state.num_no_notification > 0;
526 if (has_no_notify_task) --state.num_no_notification;
529 if (spinning_state_.compare_exchange_weak(spinning, state.Encode(), std::memory_order_acquire)) {
530 return has_no_notify_task;
539 bool IsNotifyParkedThreadRequired() {
540 uint64_t spinning = spinning_state_.load(std::memory_order_relaxed);
542 SpinningState state = SpinningState::Decode(spinning);
547 if (state.num_no_notification == state.num_spinning)
return true;
550 ++state.num_no_notification;
553 if (spinning_state_.compare_exchange_weak(spinning, state.Encode(), std::memory_order_release,
554 std::memory_order_relaxed)) {
560 static EIGEN_STRONG_INLINE uint64_t GlobalThreadIdHash() {
561 return std::hash<std::thread::id>()(std::this_thread::get_id());
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()) {
571 return it->second.get();
574 EIGEN_THREAD_LOCAL PerThread per_thread_;
575 PerThread* pt = &per_thread_;
580 static EIGEN_STRONG_INLINE
unsigned Rand(uint64_t* state) {
581 uint64_t current = *state;
583 *state = current * 6364136223846793005ULL + 0xda3e39cb94b95bdbULL;
585 return static_cast<unsigned>((current ^ (current >> 22)) >> (22 + (current >> 61)));
589using ThreadPool = ThreadPoolTempl<StlThreadEnvironment>;