diff --git a/examples/benchmark/common.hpp b/examples/benchmark/common.hpp index 30ebc10f4..edb95014a 100644 --- a/examples/benchmark/common.hpp +++ b/examples/benchmark/common.hpp @@ -192,4 +192,4 @@ void my_main(int argc, char** argv, exec::numa_policy policy = exec::get_numa_po auto [dur_ms, ops_per_sec, avg, max, min, stddev] = compute_perf(starts, ends, warmup, nRuns - 1, total_scheds); std::cout << avg << " | " << max << " | " << min << " | " << stddev << "\n"; -} \ No newline at end of file +} diff --git a/include/exec/static_thread_pool.hpp b/include/exec/static_thread_pool.hpp index f8f12a058..1a1444d8f 100644 --- a/include/exec/static_thread_pool.hpp +++ b/include/exec/static_thread_pool.hpp @@ -139,6 +139,12 @@ namespace experimental::execution std::size_t index_{(std::numeric_limits::max)()}; }; + enum class remote_poll_mode + { + speculative, + before_sleep + }; + struct remote_queue_list { private: @@ -165,13 +171,18 @@ namespace experimental::execution } } - auto pop_all_reversed(std::size_t tid) noexcept -> __intrusive_queue<&task_base::next_> + auto pop_all_reversed(std::size_t tid, remote_poll_mode mode) noexcept + -> __intrusive_queue<&task_base::next_> { remote_queue* head = head_.load(__std::memory_order_acquire); __intrusive_queue<&task_base::next_> tasks{}; while (head != nullptr) { - tasks.append(head->queues_[tid].pop_all_reversed()); + auto& queue = head->queues_[tid]; + if (mode == remote_poll_mode::before_sleep || !queue.empty()) + { + tasks.append(queue.pop_all_reversed()); + } head = head->next_; } return tasks; @@ -645,7 +656,7 @@ namespace experimental::execution }; auto try_pop() -> pop_result; - auto try_remote() -> pop_result; + auto try_remote(remote_poll_mode mode) -> pop_result; auto try_steal(std::span victims) -> pop_result; auto try_steal_near() -> pop_result; auto try_steal_any() -> pop_result; @@ -970,11 +981,11 @@ namespace experimental::execution tmp.clear(); } - inline auto - _static_thread_pool::thread_state::try_remote() -> _static_thread_pool::thread_state::pop_result + inline auto _static_thread_pool::thread_state::try_remote(remote_poll_mode mode) + -> _static_thread_pool::thread_state::pop_result { pop_result result{.task = nullptr, .queue_index = index_}; - __intrusive_queue<&task_base::next_> remotes = pool_->remotes_.pop_all_reversed(index_); + __intrusive_queue<&task_base::next_> remotes = pool_->remotes_.pop_all_reversed(index_, mode); pending_queue_.append(std::move(remotes)); if (!pending_queue_.empty()) { @@ -994,7 +1005,7 @@ namespace experimental::execution { return result; } - return try_remote(); + return try_remote(remote_poll_mode::speculative); } inline auto _static_thread_pool::thread_state::try_steal(std::span victims) @@ -1127,11 +1138,22 @@ namespace experimental::execution return result; } state expected = state::running; - if (state_.compare_exchange_weak(expected, state::sleeping, __std::memory_order_relaxed)) - { - result = try_remote(); + if (state_.compare_exchange_weak(expected, + state::sleeping, + __std::memory_order_relaxed, + __std::memory_order_relaxed)) + { + // The relaxed empty probe is safe during normal polling, but the + // running-to-sleeping boundary must perform the CAS dequeue so work + // published before the transition cannot be missed. + result = try_remote(remote_poll_mode::before_sleep); if (result.task) { + state expected_sleeping = state::sleeping; + state_.compare_exchange_strong(expected_sleeping, + state::running, + __std::memory_order_relaxed, + __std::memory_order_relaxed); return result; } set_sleeping(); @@ -1143,7 +1165,7 @@ namespace experimental::execution { lock.unlock(); } - state_.store(state::running, __std::memory_order_relaxed); + state_.exchange(state::running, __std::memory_order_acquire); result = try_pop(); } return result; @@ -1151,7 +1173,7 @@ namespace experimental::execution inline auto _static_thread_pool::thread_state::notify() -> bool { - if (state_.exchange(state::notified, __std::memory_order_relaxed) == state::sleeping) + if (state_.exchange(state::notified, __std::memory_order_release) == state::sleeping) { { std::lock_guard lock{mut_}; diff --git a/test/exec/test_static_thread_pool.cpp b/test/exec/test_static_thread_pool.cpp index 01109af5a..485745063 100644 --- a/test/exec/test_static_thread_pool.cpp +++ b/test/exec/test_static_thread_pool.cpp @@ -16,17 +16,22 @@ #include #include +#include #include #include #include // IWYU pragma: keep #include +#include #include +#include #include +#include #include #include #include #include +#include namespace ex = STDEXEC; namespace @@ -229,3 +234,107 @@ TEST_CASE("bulk on static_thread_pool executes on multiple threads, take 2", ex::sync_wait(std::move(sender)); REQUIRE(thread_ids.size() == num_of_threads); } + +namespace +{ + void run_remote_poll_stress(bool separate_schedulers) + { + constexpr std::size_t num_producers = 4; + constexpr std::size_t rounds = 10'000; + + std::latch ready{num_producers}; + std::atomic start{false}; + std::atomic stop{false}; + std::vector> completed(num_producers); + std::vector producers; + producers.reserve(num_producers); + for (auto& count: completed) + { + count.store(0, std::memory_order_relaxed); + } + + exec::static_thread_pool pool{1}; + using scheduler_t = decltype(pool.get_scheduler()); + std::optional shared_scheduler; + if (!separate_schedulers) + { + shared_scheduler.emplace(pool.get_scheduler()); + } + + for (std::size_t producer = 0; producer < num_producers; ++producer) + { + producers.emplace_back( + [&, producer] + { + auto scheduler = separate_schedulers ? pool.get_scheduler() : *shared_scheduler; + ready.count_down(); + while (!start.load(std::memory_order_acquire)) + { + std::this_thread::yield(); + } + + auto* const producer_completed = &completed[producer]; + std::size_t expected = 0; + for (std::size_t round = 0; round < rounds && !stop.load(std::memory_order_relaxed); + ++round) + { + std::size_t const batch_size = (round % 4 == 0) ? 2 : 1; + expected += batch_size; + for (std::size_t i = 0; i < batch_size; ++i) + { + exec::start_detached( + ex::schedule(scheduler) + | ex::then([producer_completed] + { producer_completed->fetch_add(1, std::memory_order_relaxed); })); + } + + while (!stop.load(std::memory_order_relaxed) + && producer_completed->load(std::memory_order_relaxed) < expected) + { + std::this_thread::yield(); + } + std::this_thread::yield(); + } + }); + } + + ready.wait(); + start.store(true, std::memory_order_release); + + auto const expected = num_producers * rounds + num_producers * ((rounds + 3) / 4); + auto const deadline = std::chrono::steady_clock::now() + std::chrono::seconds(10); + auto completed_total = [&] + { + std::size_t result = 0; + for (auto const & count: completed) + { + result += count.load(std::memory_order_relaxed); + } + return result; + }; + + while (completed_total() < expected && std::chrono::steady_clock::now() < deadline) + { + std::this_thread::yield(); + } + stop.store(true, std::memory_order_release); + for (auto& producer: producers) + { + producer.join(); + } + + CHECK(completed_total() == expected); + } +} // namespace + +TEST_CASE("static_thread_pool drains remote work from a shared scheduler", + "[types][static_thread_pool][stress]") +{ + run_remote_poll_stress(false); +} + +TEST_CASE("static_thread_pool drains remote work from producer schedulers", + "[types][static_thread_pool][stress]") +{ + run_remote_poll_stress(true); +} diff --git a/test/rrd/CMakeLists.txt b/test/rrd/CMakeLists.txt index 36746b5a6..d8feafc67 100644 --- a/test/rrd/CMakeLists.txt +++ b/test/rrd/CMakeLists.txt @@ -54,7 +54,7 @@ function(add_relacy_test target_name) endfunction() set(relacy_tests async_scope bwos_lifo_queue intrusive_mpsc_queue split - sync_wait) + static_thread_pool_remote_poll sync_wait) foreach(test ${relacy_tests}) add_relacy_test(${test}) diff --git a/test/rrd/static_thread_pool_remote_poll.cpp b/test/rrd/static_thread_pool_remote_poll.cpp new file mode 100644 index 000000000..1faca3213 --- /dev/null +++ b/test/rrd/static_thread_pool_remote_poll.cpp @@ -0,0 +1,248 @@ +/* + * Copyright (c) 2026 NVIDIA Corporation + * + * Licensed under the Apache License Version 2.0 with LLVM Exceptions + * (the "License"); you may not use this file except in compliance + * with the License. You may obtain a copy of the License at + * + * https://llvm.org/LICENSE.txt + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include + +#include + +struct task_node +{ + task_node* next_ = nullptr; +}; + +struct remote_queue +{ + remote_queue* next_ = nullptr; + std::atomic head_{nullptr}; +}; + +struct static_thread_pool_remote_poll : rl::test_suite +{ + static constexpr int running = 0; + static constexpr int sleeping = 1; + static constexpr int notified = 2; + + enum class remote_poll_mode + { + speculative, + before_sleep + }; + + struct poll_result + { + bool any = false; + bool second = false; + }; + + std::atomic state_{running}; + std::atomic remote_head_{nullptr}; + std::atomic first_notification_published_{false}; + remote_queue first_queue_{}; + remote_queue second_queue_{}; + task_node first_task_{}; + task_node first_extra_task_{}; + task_node second_task_{}; + bool worker_would_sleep_ = false; + + void before() + { + state_.store(running, std::memory_order_relaxed); + remote_head_.store(nullptr, std::memory_order_relaxed); + first_notification_published_.store(false, std::memory_order_relaxed); + first_queue_.next_ = nullptr; + second_queue_.next_ = nullptr; + first_queue_.head_.store(nullptr, std::memory_order_relaxed); + second_queue_.head_.store(nullptr, std::memory_order_relaxed); + first_task_.next_ = nullptr; + first_extra_task_.next_ = nullptr; + second_task_.next_ = nullptr; + worker_would_sleep_ = false; + } + + void publish_remote_queue(remote_queue& queue) + { + auto* old_head = remote_head_.load(std::memory_order_acquire); + do + { + queue.next_ = old_head; + } + while (!remote_head_.compare_exchange_weak(old_head, + &queue, + std::memory_order_acq_rel, + std::memory_order_acquire)); + } + + auto push(remote_queue& queue, task_node& task) -> bool + { + auto* old_head = queue.head_.load(std::memory_order_relaxed); + do + { + task.next_ = old_head; + } + while (!queue.head_.compare_exchange_weak(old_head, + &task, + std::memory_order_acq_rel, + std::memory_order_acquire)); + return old_head == nullptr; + } + + void notify() + { + state_.exchange(notified, std::memory_order_release); + } + + void enqueue(remote_queue& queue, task_node& task) + { + bool const was_empty = push(queue, task); + if (was_empty) + { + notify(); + } + } + + auto drain(remote_queue& queue) -> task_node* + { + auto* old_head = queue.head_.load(std::memory_order_relaxed); + while (!queue.head_.compare_exchange_weak(old_head, + nullptr, + std::memory_order_acq_rel, + std::memory_order_acquire)) + { + } + return old_head; + } + + auto poll_remote(remote_poll_mode mode) -> poll_result + { + poll_result result{}; + auto* queue = remote_head_.load(std::memory_order_acquire); + while (queue != nullptr) + { + if (mode == remote_poll_mode::before_sleep + || queue->head_.load(std::memory_order_relaxed) != nullptr) + { + for (auto* task = drain(*queue); task != nullptr; task = task->next_) + { + result.any = true; + result.second = result.second || task == &second_task_; + } + } + queue = queue->next_; + } + return result; + } + + void worker_poll() + { + auto result = poll_remote(remote_poll_mode::speculative); + if (result.second) + { + return; + } + if (result.any) + { + result = poll_remote(remote_poll_mode::speculative); + if (result.second) + { + return; + } + } + + int expected = running; + if (!state_.compare_exchange_weak(expected, + sleeping, + std::memory_order_relaxed, + std::memory_order_relaxed)) + { + state_.exchange(running, std::memory_order_acquire); + result = poll_remote(remote_poll_mode::speculative); + if (result.second) + { + return; + } + if (result.any) + { + result = poll_remote(remote_poll_mode::speculative); + if (result.second) + { + return; + } + } + + expected = running; + if (!state_.compare_exchange_weak(expected, + sleeping, + std::memory_order_relaxed, + std::memory_order_relaxed)) + { + return; + } + } + + result = poll_remote(remote_poll_mode::before_sleep); + if (result.any) + { + int expected_sleeping = sleeping; + state_.compare_exchange_strong(expected_sleeping, + running, + std::memory_order_relaxed, + std::memory_order_relaxed); + return; + } + + worker_would_sleep_ = true; + } + + void thread(unsigned thread_id) + { + if (thread_id == 0) + { + publish_remote_queue(first_queue_); + enqueue(first_queue_, first_task_); + enqueue(first_queue_, first_extra_task_); + first_notification_published_.store(true, std::memory_order_release); + } + else if (thread_id == 1) + { + while (!first_notification_published_.load(std::memory_order_acquire)) + { + } + publish_remote_queue(second_queue_); + enqueue(second_queue_, second_task_); + } + else + { + worker_poll(); + } + } + + void after() + { + if (worker_would_sleep_) + { + RL_ASSERT(state_.load(std::memory_order_acquire) != sleeping); + } + } +}; + +auto main() -> int +{ + rl::test_params p; + p.iteration_count = 50000; + p.execution_depth_limit = 10000; + p.search_type = rl::random_scheduler_type; + return rl::simulate(p) ? 0 : 1; +}