Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 4 additions & 1 deletion include/exec/thread_pool_base.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -38,11 +38,13 @@ import stdexec;
# include "../stdexec/__detail/__transform_completion_signatures.hpp"
# include "../stdexec/__detail/__type_traits.hpp"

# include <algorithm>
# include <atomic>
# include <concepts>
# include <cstdint>
# include <exception>
# include <tuple>
# include <type_traits>
# include <utility>
# include <variant>
#endif
Expand Down Expand Up @@ -230,14 +232,15 @@ namespace experimental::execution
[[nodiscard]]
auto num_agents_required() const -> std::uint32_t
{
using common_shape_t = std::common_type_t<Shape, std::uint32_t>;
auto const parallelism = parallelize_ ? pool_.available_parallelism()
: static_cast<std::uint32_t>(1);

// With work stealing, is std::min necessary, or can we feel free to
// ask for more agents (tasks) than we can actually deal with at one
// time?
return Shape{} < shape_
? static_cast<std::uint32_t>((std::min) (shape_, static_cast<Shape>(parallelism)))
? static_cast<std::uint32_t>((std::min<common_shape_t>) (shape_, parallelism))
: 0;
}

Expand Down
136 changes: 131 additions & 5 deletions test/exec/test_thread_pool_base.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -18,8 +18,11 @@
#include <stdexec/execution.hpp>
#include <test_common/catch2.hpp>

#include <array>
#include <cstdint>
#include <exception>
#include <limits>
#include <utility>

namespace ex = STDEXEC;

Expand All @@ -32,7 +35,7 @@ namespace
[[nodiscard]]
auto available_parallelism() const noexcept -> std::uint32_t
{
return 1;
return parallelism_;
}

[[nodiscard]]
Expand All @@ -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
Expand Down Expand Up @@ -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<int, 5> 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<std::uint32_t>(shape) + 1);
for (std::size_t i = 0; i < visits.size(); ++i)
{
CHECK(visits[i] == (i < static_cast<std::size_t>(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<int, 255> 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<std::uint64_t>::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);
Expand Down