Skip to content

Commit 7735228

Browse files
Avoid RMWs when polling empty remote queues (NVIDIA#2180)
* Avoid RMWs when polling empty remote queues * Fix remote poll before worker sleep * Harden remote queue sleep polling * Consume remote poll notifications with acquire RMW * Format static thread pool polling changes * Clarify remote polling modes * Strengthen remote polling coverage * clang-format --------- Co-authored-by: Eric Niebler <eniebler@nvidia.com>
1 parent 5625b5e commit 7735228

5 files changed

Lines changed: 393 additions & 14 deletions

File tree

examples/benchmark/common.hpp

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -192,4 +192,4 @@ void my_main(int argc, char** argv, exec::numa_policy policy = exec::get_numa_po
192192
auto [dur_ms, ops_per_sec, avg, max, min, stddev] =
193193
compute_perf(starts, ends, warmup, nRuns - 1, total_scheds);
194194
std::cout << avg << " | " << max << " | " << min << " | " << stddev << "\n";
195-
}
195+
}

include/exec/static_thread_pool.hpp

Lines changed: 34 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -142,6 +142,12 @@ namespace experimental::execution
142142
std::size_t index_{(std::numeric_limits<std::size_t>::max)()};
143143
};
144144

145+
enum class remote_poll_mode
146+
{
147+
speculative,
148+
before_sleep
149+
};
150+
145151
struct remote_queue_list
146152
{
147153
private:
@@ -168,13 +174,18 @@ namespace experimental::execution
168174
}
169175
}
170176

171-
auto pop_all_reversed(std::size_t tid) noexcept -> __intrusive_queue<&task_base::next_>
177+
auto pop_all_reversed(std::size_t tid, remote_poll_mode mode) noexcept
178+
-> __intrusive_queue<&task_base::next_>
172179
{
173180
remote_queue* head = head_.load(__std::memory_order_acquire);
174181
__intrusive_queue<&task_base::next_> tasks{};
175182
while (head != nullptr)
176183
{
177-
tasks.append(head->queues_[tid].pop_all_reversed());
184+
auto& queue = head->queues_[tid];
185+
if (mode == remote_poll_mode::before_sleep || !queue.empty())
186+
{
187+
tasks.append(queue.pop_all_reversed());
188+
}
178189
head = head->next_;
179190
}
180191
return tasks;
@@ -648,7 +659,7 @@ namespace experimental::execution
648659
};
649660

650661
auto try_pop() -> pop_result;
651-
auto try_remote() -> pop_result;
662+
auto try_remote(remote_poll_mode mode) -> pop_result;
652663
auto try_steal(std::span<workstealing_victim> victims) -> pop_result;
653664
auto try_steal_near() -> pop_result;
654665
auto try_steal_any() -> pop_result;
@@ -973,11 +984,11 @@ namespace experimental::execution
973984
tmp.clear();
974985
}
975986

976-
inline auto
977-
_static_thread_pool::thread_state::try_remote() -> _static_thread_pool::thread_state::pop_result
987+
inline auto _static_thread_pool::thread_state::try_remote(remote_poll_mode mode)
988+
-> _static_thread_pool::thread_state::pop_result
978989
{
979990
pop_result result{.task = nullptr, .queue_index = index_};
980-
__intrusive_queue<&task_base::next_> remotes = pool_->remotes_.pop_all_reversed(index_);
991+
__intrusive_queue<&task_base::next_> remotes = pool_->remotes_.pop_all_reversed(index_, mode);
981992
pending_queue_.append(std::move(remotes));
982993
if (!pending_queue_.empty())
983994
{
@@ -997,7 +1008,7 @@ namespace experimental::execution
9971008
{
9981009
return result;
9991010
}
1000-
return try_remote();
1011+
return try_remote(remote_poll_mode::speculative);
10011012
}
10021013

10031014
inline auto _static_thread_pool::thread_state::try_steal(std::span<workstealing_victim> victims)
@@ -1130,11 +1141,22 @@ namespace experimental::execution
11301141
return result;
11311142
}
11321143
state expected = state::running;
1133-
if (state_.compare_exchange_weak(expected, state::sleeping, __std::memory_order_relaxed))
1134-
{
1135-
result = try_remote();
1144+
if (state_.compare_exchange_weak(expected,
1145+
state::sleeping,
1146+
__std::memory_order_relaxed,
1147+
__std::memory_order_relaxed))
1148+
{
1149+
// The relaxed empty probe is safe during normal polling, but the
1150+
// running-to-sleeping boundary must perform the CAS dequeue so work
1151+
// published before the transition cannot be missed.
1152+
result = try_remote(remote_poll_mode::before_sleep);
11361153
if (result.task)
11371154
{
1155+
state expected_sleeping = state::sleeping;
1156+
state_.compare_exchange_strong(expected_sleeping,
1157+
state::running,
1158+
__std::memory_order_relaxed,
1159+
__std::memory_order_relaxed);
11381160
return result;
11391161
}
11401162
set_sleeping();
@@ -1146,15 +1168,15 @@ namespace experimental::execution
11461168
{
11471169
lock.unlock();
11481170
}
1149-
state_.store(state::running, __std::memory_order_relaxed);
1171+
state_.exchange(state::running, __std::memory_order_acquire);
11501172
result = try_pop();
11511173
}
11521174
return result;
11531175
}
11541176

11551177
inline auto _static_thread_pool::thread_state::notify() -> bool
11561178
{
1157-
if (state_.exchange(state::notified, __std::memory_order_relaxed) == state::sleeping)
1179+
if (state_.exchange(state::notified, __std::memory_order_release) == state::sleeping)
11581180
{
11591181
{
11601182
std::lock_guard lock{mut_};

test/exec/test_static_thread_pool.cpp

Lines changed: 109 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,17 +16,22 @@
1616

1717
#include <exec/sequence/ignore_all_values.hpp>
1818
#include <exec/sequence/transform_each.hpp>
19+
#include <exec/start_detached.hpp>
1920
#include <exec/static_thread_pool.hpp>
2021
#include <stdexec/execution.hpp>
2122
#include <test_common/catch2.hpp> // IWYU pragma: keep
2223

2324
#include <atomic>
25+
#include <chrono>
2426
#include <exception>
27+
#include <latch>
2528
#include <mutex>
29+
#include <optional>
2630
#include <ranges>
2731
#include <stdexcept>
2832
#include <thread>
2933
#include <unordered_set>
34+
#include <vector>
3035
namespace ex = STDEXEC;
3136

3237
namespace
@@ -229,3 +234,107 @@ TEST_CASE("bulk on static_thread_pool executes on multiple threads, take 2",
229234
ex::sync_wait(std::move(sender));
230235
REQUIRE(thread_ids.size() == num_of_threads);
231236
}
237+
238+
namespace
239+
{
240+
void run_remote_poll_stress(bool separate_schedulers)
241+
{
242+
constexpr std::size_t num_producers = 4;
243+
constexpr std::size_t rounds = 10'000;
244+
245+
std::latch ready{num_producers};
246+
std::atomic<bool> start{false};
247+
std::atomic<bool> stop{false};
248+
std::vector<std::atomic<std::size_t>> completed(num_producers);
249+
std::vector<std::thread> producers;
250+
producers.reserve(num_producers);
251+
for (auto& count: completed)
252+
{
253+
count.store(0, std::memory_order_relaxed);
254+
}
255+
256+
exec::static_thread_pool pool{1};
257+
using scheduler_t = decltype(pool.get_scheduler());
258+
std::optional<scheduler_t> shared_scheduler;
259+
if (!separate_schedulers)
260+
{
261+
shared_scheduler.emplace(pool.get_scheduler());
262+
}
263+
264+
for (std::size_t producer = 0; producer < num_producers; ++producer)
265+
{
266+
producers.emplace_back(
267+
[&, producer]
268+
{
269+
auto scheduler = separate_schedulers ? pool.get_scheduler() : *shared_scheduler;
270+
ready.count_down();
271+
while (!start.load(std::memory_order_acquire))
272+
{
273+
std::this_thread::yield();
274+
}
275+
276+
auto* const producer_completed = &completed[producer];
277+
std::size_t expected = 0;
278+
for (std::size_t round = 0; round < rounds && !stop.load(std::memory_order_relaxed);
279+
++round)
280+
{
281+
std::size_t const batch_size = (round % 4 == 0) ? 2 : 1;
282+
expected += batch_size;
283+
for (std::size_t i = 0; i < batch_size; ++i)
284+
{
285+
exec::start_detached(
286+
ex::schedule(scheduler)
287+
| ex::then([producer_completed]
288+
{ producer_completed->fetch_add(1, std::memory_order_relaxed); }));
289+
}
290+
291+
while (!stop.load(std::memory_order_relaxed)
292+
&& producer_completed->load(std::memory_order_relaxed) < expected)
293+
{
294+
std::this_thread::yield();
295+
}
296+
std::this_thread::yield();
297+
}
298+
});
299+
}
300+
301+
ready.wait();
302+
start.store(true, std::memory_order_release);
303+
304+
auto const expected = num_producers * rounds + num_producers * ((rounds + 3) / 4);
305+
auto const deadline = std::chrono::steady_clock::now() + std::chrono::seconds(10);
306+
auto completed_total = [&]
307+
{
308+
std::size_t result = 0;
309+
for (auto const & count: completed)
310+
{
311+
result += count.load(std::memory_order_relaxed);
312+
}
313+
return result;
314+
};
315+
316+
while (completed_total() < expected && std::chrono::steady_clock::now() < deadline)
317+
{
318+
std::this_thread::yield();
319+
}
320+
stop.store(true, std::memory_order_release);
321+
for (auto& producer: producers)
322+
{
323+
producer.join();
324+
}
325+
326+
CHECK(completed_total() == expected);
327+
}
328+
} // namespace
329+
330+
TEST_CASE("static_thread_pool drains remote work from a shared scheduler",
331+
"[types][static_thread_pool][stress]")
332+
{
333+
run_remote_poll_stress(false);
334+
}
335+
336+
TEST_CASE("static_thread_pool drains remote work from producer schedulers",
337+
"[types][static_thread_pool][stress]")
338+
{
339+
run_remote_poll_stress(true);
340+
}

test/rrd/CMakeLists.txt

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -54,7 +54,7 @@ function(add_relacy_test target_name)
5454
endfunction()
5555

5656
set(relacy_tests async_scope bwos_lifo_queue intrusive_mpsc_queue split
57-
sync_wait)
57+
static_thread_pool_remote_poll sync_wait)
5858

5959
foreach(test ${relacy_tests})
6060
add_relacy_test(${test})

0 commit comments

Comments
 (0)