Eigen  5.0.1
 
Loading...
Searching...
No Matches
EventCount.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_EVENTCOUNT_H
12#define EIGEN_THREADPOOL_EVENTCOUNT_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19// EventCount allows to wait for arbitrary predicates in non-blocking
20// algorithms. Think of condition variable, but wait predicate does not need to
21// be protected by a mutex. Usage:
22// Waiting thread does:
23//
24// if (predicate)
25// return act();
26// EventCount::Waiter& w = waiters[my_index];
27// ec.Prewait(&w);
28// if (predicate) {
29// ec.CancelWait(&w);
30// return act();
31// }
32// ec.CommitWait(&w);
33//
34// Notifying thread does:
35//
36// predicate = true;
37// ec.Notify(true);
38//
39// Notify is cheap if there are no waiting threads. Prewait/CommitWait are not
40// cheap, but they are executed only if the preceding predicate check has
41// failed.
42//
43// Algorithm outline:
44// There are two main variables: predicate (managed by user) and state_.
45// Operation closely resembles Dekker mutual algorithm:
46// https://en.wikipedia.org/wiki/Dekker%27s_algorithm
47// Waiting thread sets state_ then checks predicate, Notifying thread sets
48// predicate then checks state_. Due to seq_cst fences in between these
49// operations it is guaranteed that either waiter will see predicate change
50// and won't block, or notifying thread will see state_ change and will unblock
51// the waiter, or both. But it can't happen that both threads don't see each
52// other changes, which would lead to deadlock.
53class EventCount {
54 public:
55 class Waiter;
56
57 EventCount(MaxSizeVector<Waiter>& waiters) : state_(kStackMask), waiters_(waiters) {
58 eigen_plain_assert(waiters.size() < (1 << kWaiterBits) - 1);
59 }
60
61 EventCount(const EventCount&) = delete;
62 void operator=(const EventCount&) = delete;
63
64 ~EventCount() {
65 // Ensure there are no waiters.
66 eigen_plain_assert(state_.load() == kStackMask);
67 }
68
69 // Prewait prepares for waiting.
70 // After calling Prewait, the thread must re-check the wait predicate
71 // and then call either CancelWait or CommitWait.
72 void Prewait() {
73 uint64_t state = state_.load(std::memory_order_relaxed);
74 for (;;) {
75 CheckState(state);
76 uint64_t newstate = state + kWaiterInc;
77 CheckState(newstate);
78 if (state_.compare_exchange_weak(state, newstate, std::memory_order_seq_cst)) return;
79 }
80 }
81
82 // CommitWait commits waiting after Prewait.
83 void CommitWait(Waiter* w) {
84 eigen_plain_assert((w->epoch & ~kEpochMask) == 0);
85 w->state = Waiter::kNotSignaled;
86 const uint64_t me = (w - &waiters_[0]) | w->epoch;
87 uint64_t state = state_.load(std::memory_order_seq_cst);
88 for (;;) {
89 CheckState(state, true);
90 uint64_t newstate;
91 if ((state & kSignalMask) != 0) {
92 // Consume the signal and return immediately.
93 newstate = state - kWaiterInc - kSignalInc;
94 } else {
95 // Remove this thread from pre-wait counter and add to the waiter stack.
96 newstate = ((state & kWaiterMask) - kWaiterInc) | me;
97 w->next.store(state & (kStackMask | kEpochMask), std::memory_order_relaxed);
98 }
99 CheckState(newstate);
100 if (state_.compare_exchange_weak(state, newstate, std::memory_order_acq_rel)) {
101 if ((state & kSignalMask) == 0) {
102 w->epoch += kEpochInc;
103 Park(w);
104 }
105 return;
106 }
107 }
108 }
109
110 // CancelWait cancels effects of the previous Prewait call.
111 void CancelWait() {
112 uint64_t state = state_.load(std::memory_order_relaxed);
113 for (;;) {
114 CheckState(state, true);
115 uint64_t newstate = state - kWaiterInc;
116 // We don't know if the thread was also notified or not,
117 // so we should not consume a signal unconditionally.
118 // Only if number of waiters is equal to number of signals,
119 // we know that the thread was notified and we must take away the signal
120 // and forward it to any other waiting thread.
121 const bool notify = (((state & kWaiterMask) >> kWaiterShift) == ((state & kSignalMask) >> kSignalShift));
122 if (notify) {
123 newstate -= kSignalInc;
124 }
125 CheckState(newstate);
126 if (state_.compare_exchange_weak(state, newstate, std::memory_order_acq_rel)) {
127 if (notify) {
128 Notify(false);
129 }
130 return;
131 }
132 }
133 }
134
135 // Notify wakes one or all waiting threads.
136 // Must be called after changing the associated wait predicate.
137 void Notify(bool notifyAll) {
138 std::atomic_thread_fence(std::memory_order_seq_cst);
139 uint64_t state = state_.load(std::memory_order_acquire);
140 for (;;) {
141 CheckState(state);
142 const uint64_t waiters = (state & kWaiterMask) >> kWaiterShift;
143 const uint64_t signals = (state & kSignalMask) >> kSignalShift;
144 // Easy case: no waiters.
145 if ((state & kStackMask) == kStackMask && waiters == signals) return;
146 uint64_t newstate;
147 if (notifyAll) {
148 // Empty wait stack and set signal to number of pre-wait threads.
149 newstate = (state & kWaiterMask) | (waiters << kSignalShift) | kStackMask;
150 } else if (signals < waiters) {
151 // There is a thread in pre-wait state, unblock it.
152 newstate = state + kSignalInc;
153 } else {
154 // Pop a waiter from list and unpark it.
155 Waiter* w = &waiters_[state & kStackMask];
156 uint64_t next = w->next.load(std::memory_order_relaxed);
157 newstate = (state & (kWaiterMask | kSignalMask)) | next;
158 }
159 CheckState(newstate);
160 if (state_.compare_exchange_weak(state, newstate, std::memory_order_acq_rel)) {
161 if (!notifyAll && (signals < waiters)) return; // unblocked pre-wait thread
162 if ((state & kStackMask) == kStackMask) return;
163 Waiter* w = &waiters_[state & kStackMask];
164 if (!notifyAll) w->next.store(kStackMask, std::memory_order_relaxed);
165 Unpark(w);
166 return;
167 }
168 }
169 }
170
171 private:
172 // State_ layout:
173 // - low kWaiterBits is a stack of waiters committed wait
174 // (indexes in waiters_ array are used as stack elements,
175 // kStackMask means empty stack).
176 // - next kWaiterBits is count of waiters in prewait state.
177 // - next kWaiterBits is count of pending signals.
178 // - remaining bits are ABA counter for the stack.
179 // (stored in Waiter node and incremented on push).
180 static const uint64_t kWaiterBits = 14;
181 static const uint64_t kStackMask = (1ull << kWaiterBits) - 1;
182 static const uint64_t kWaiterShift = kWaiterBits;
183 static const uint64_t kWaiterMask = ((1ull << kWaiterBits) - 1) << kWaiterShift;
184 static const uint64_t kWaiterInc = 1ull << kWaiterShift;
185 static const uint64_t kSignalShift = 2 * kWaiterBits;
186 static const uint64_t kSignalMask = ((1ull << kWaiterBits) - 1) << kSignalShift;
187 static const uint64_t kSignalInc = 1ull << kSignalShift;
188 static const uint64_t kEpochShift = 3 * kWaiterBits;
189 static const uint64_t kEpochBits = 64 - kEpochShift;
190 static const uint64_t kEpochMask = ((1ull << kEpochBits) - 1) << kEpochShift;
191 static const uint64_t kEpochInc = 1ull << kEpochShift;
192
193 public:
194 class Waiter {
195 friend class EventCount;
196
197 enum State {
198 kNotSignaled,
199 kWaiting,
200 kSignaled,
201 };
202
203 EIGEN_ALIGN_TO_AVOID_FALSE_SHARING std::atomic<uint64_t> next{kStackMask};
204 EIGEN_MUTEX mu;
205 EIGEN_CONDVAR cv;
206 uint64_t epoch{0};
207 unsigned state{kNotSignaled};
208 };
209
210 private:
211 static void CheckState(uint64_t state, bool waiter = false) {
212 static_assert(kEpochBits >= 20, "not enough bits to prevent ABA problem");
213 const uint64_t waiters = (state & kWaiterMask) >> kWaiterShift;
214 const uint64_t signals = (state & kSignalMask) >> kSignalShift;
215 eigen_plain_assert(waiters >= signals);
216 eigen_plain_assert(waiters < (1 << kWaiterBits) - 1);
217 eigen_plain_assert(!waiter || waiters > 0);
218 (void)waiter;
219 (void)waiters;
220 (void)signals;
221 }
222
223 void Park(Waiter* w) {
224 EIGEN_MUTEX_LOCK lock(w->mu);
225 while (w->state != Waiter::kSignaled) {
226 w->state = Waiter::kWaiting;
227 w->cv.wait(lock);
228 }
229 }
230
231 void Unpark(Waiter* w) {
232 for (Waiter* next; w; w = next) {
233 uint64_t wnext = w->next.load(std::memory_order_relaxed) & kStackMask;
234 next = wnext == kStackMask ? nullptr : &waiters_[internal::convert_index<size_t>(wnext)];
235 unsigned state;
236 {
237 EIGEN_MUTEX_LOCK lock(w->mu);
238 state = w->state;
239 w->state = Waiter::kSignaled;
240 }
241 // Avoid notifying if it wasn't waiting.
242 if (state == Waiter::kWaiting) w->cv.notify_one();
243 }
244 }
245
246 std::atomic<uint64_t> state_;
247 MaxSizeVector<Waiter>& waiters_;
248};
249
250} // namespace Eigen
251
252#endif // EIGEN_THREADPOOL_EVENTCOUNT_H