diff --git a/include/exec/thread_pool_base.hpp b/include/exec/thread_pool_base.hpp index 1317b0d92..7ce4899ce 100644 --- a/include/exec/thread_pool_base.hpp +++ b/include/exec/thread_pool_base.hpp @@ -38,11 +38,13 @@ import stdexec; # include "../stdexec/__detail/__transform_completion_signatures.hpp" # include "../stdexec/__detail/__type_traits.hpp" +# include # include # include # include # include # include +# include # include # include #endif @@ -230,6 +232,7 @@ namespace experimental::execution [[nodiscard]] auto num_agents_required() const -> std::uint32_t { + using common_shape_t = std::common_type_t; auto const parallelism = parallelize_ ? pool_.available_parallelism() : static_cast(1); @@ -237,7 +240,7 @@ namespace experimental::execution // ask for more agents (tasks) than we can actually deal with at one // time? return Shape{} < shape_ - ? static_cast((std::min) (shape_, static_cast(parallelism))) + ? static_cast((std::min) (shape_, parallelism)) : 0; } diff --git a/test/exec/test_thread_pool_base.cpp b/test/exec/test_thread_pool_base.cpp index 68fe40a08..58aac6c9c 100644 --- a/test/exec/test_thread_pool_base.cpp +++ b/test/exec/test_thread_pool_base.cpp @@ -18,8 +18,11 @@ #include #include +#include #include #include +#include +#include namespace ex = STDEXEC; @@ -32,7 +35,7 @@ namespace [[nodiscard]] auto available_parallelism() const noexcept -> std::uint32_t { - return 1; + return parallelism_; } [[nodiscard]] @@ -47,7 +50,8 @@ namespace task->execute_(task, tid); } - std::uint32_t enqueued_ = 0; + std::uint32_t parallelism_ = 1; + std::uint32_t enqueued_ = 0; }; struct value_capture_error @@ -126,15 +130,137 @@ namespace CHECK(pool.enqueued_ == 1); } - TEST_CASE("thread_pool_base bulk does not invoke the function with a negative shape", + TEMPLATE_TEST_CASE("thread_pool_base bulk handles pool sizes wider than the shape", + "[thread_pool_base][bulk]", + std::uint8_t, + std::int8_t, + std::uint16_t, + std::int16_t, + std::int32_t) + { + inline_test_thread_pool pool; + pool.parallelism_ = GENERATE(256u, 257u, 65'536u, 65'537u, 0x80000000u, 0xffffffffu); + completion_state state; + std::array visits{}; + auto const shape = GENERATE(TestType{1}, TestType{3}, TestType{5}); + + SECTION("chunked") + { + auto sndr = ex::schedule(pool.get_scheduler()) + | ex::bulk_chunked(ex::par, + shape, + [&](TestType begin, TestType end) noexcept + { + for (auto i = begin; i < end; ++i) + { + ++visits[i]; + } + }); + auto op = ex::connect(std::move(sndr), counting_receiver{&state}); + ex::start(op); + } + + SECTION("unchunked") + { + auto sndr = ex::schedule(pool.get_scheduler()) + | ex::bulk_unchunked(ex::par, shape, [&](TestType i) noexcept { ++visits[i]; }); + auto op = ex::connect(std::move(sndr), counting_receiver{&state}); + ex::start(op); + } + + REQUIRE(state.completions_ == 1); + CHECK_FALSE(state.error_); + CHECK(pool.enqueued_ == static_cast(shape) + 1); + for (std::size_t i = 0; i < visits.size(); ++i) + { + CHECK(visits[i] == (i < static_cast(shape) ? 1 : 0)); + } + } + + TEST_CASE("thread_pool_base bulk covers a maximal uint8 shape", "[thread_pool_base][bulk]") + { + inline_test_thread_pool pool; + pool.parallelism_ = 256; + completion_state state; + std::array visits{}; + + auto sndr = ex::schedule(pool.get_scheduler()) + | ex::bulk_unchunked(ex::par, + std::uint8_t{255}, + [&](std::uint8_t i) noexcept { ++visits[i]; }); + auto op = ex::connect(std::move(sndr), counting_receiver{&state}); + ex::start(op); + + CHECK(state.completions_ == 1); + CHECK_FALSE(state.error_); + CHECK(pool.enqueued_ == 256); + for (int count: visits) + { + CHECK(count == 1); + } + } + + TEST_CASE("thread_pool_base bulk preserves shapes wider than the worker count", + "[thread_pool_base][bulk]") + { + inline_test_thread_pool pool; + pool.parallelism_ = 2; + completion_state state; + auto const shape = GENERATE(std::uint64_t{1} << 32, + (std::numeric_limits::max)()); + std::uint64_t visited = 0; + + auto sndr = ex::schedule(pool.get_scheduler()) + | ex::bulk_chunked(ex::par, + shape, + [&](std::uint64_t begin, std::uint64_t end) noexcept + { visited += end - begin; }); + auto op = ex::connect(std::move(sndr), counting_receiver{&state}); + ex::start(op); + + CHECK(state.completions_ == 1); + CHECK_FALSE(state.error_); + CHECK(visited == shape); + CHECK(pool.enqueued_ == 3); + } + + TEST_CASE("thread_pool_base sequential bulk uses one worker with a narrow shape", + "[thread_pool_base][bulk]") + { + inline_test_thread_pool pool; + pool.parallelism_ = 256; + completion_state state; + int visited = 0; + int bulk_calls = 0; + + auto sndr = ex::schedule(pool.get_scheduler()) + | ex::bulk_chunked(ex::seq, + std::uint8_t{3}, + [&](std::uint8_t begin, std::uint8_t end) noexcept + { + visited += end - begin; + ++bulk_calls; + }); + auto op = ex::connect(std::move(sndr), counting_receiver{&state}); + ex::start(op); + + CHECK(state.completions_ == 1); + CHECK_FALSE(state.error_); + CHECK(visited == 3); + CHECK(bulk_calls == 1); + CHECK(pool.enqueued_ == 2); + } + + TEST_CASE("thread_pool_base bulk does not invoke the function with a nonpositive shape", "[thread_pool_base][bulk]") { inline_test_thread_pool pool; completion_state state; int bulk_calls = 0; - auto sndr = ex::schedule(pool.get_scheduler()) - | ex::bulk_chunked(ex::par, -1, [&](int, int) noexcept { ++bulk_calls; }); + auto const shape = GENERATE(-1, 0); + auto sndr = ex::schedule(pool.get_scheduler()) + | ex::bulk_chunked(ex::par, shape, [&](int, int) noexcept { ++bulk_calls; }); auto op = ex::connect(std::move(sndr), counting_receiver{&state}); ex::start(op);