kokkos-execution 0.0.1
Loading...
Searching...
No Matches
test_when_all.cpp
Go to the documentation of this file.
2PRAGMA_DIAGNOSTIC_PUSH
4#include "exec/single_thread_context.hpp"
5PRAGMA_DIAGNOSTIC_POP
6
9
11
22
34
35using host_execution_space = Kokkos::DefaultHostExecutionSpace;
36
38
39using namespace Kokkos::utils::callbacks;
40
53
55TEST(WhenAll, no_branch) {
56 auto sndr = stdexec::when_all();
57
58 static_assert(std::same_as<stdexec::tag_of_t<decltype(sndr)>, stdexec::just_t>);
59
60 static_assert(!stdexec::dependent_sender<decltype(sndr)>);
61
62 static_assert(
63 stdexec::get_completion_signatures<decltype(sndr)>()
64 == stdexec::completion_signatures<stdexec::set_value_t()>{});
65
66 ASSERT_TRUE(stdexec::sync_wait(std::move(sndr)).has_value()); // NOLINT(performance-move-const-arg)
67}
68
70TEST(WhenAll, single_branch) {
71 auto sndr = stdexec::just(42);
72
73 static_assert(std::same_as<decltype(stdexec::when_all(sndr)), decltype(sndr)>);
74 static_assert(std::same_as<
75 decltype(stdexec::when_all(std::move(sndr))), // NOLINT(performance-move-const-arg)
76 decltype(sndr)
77 >);
78
79 auto&& w_a = stdexec::when_all(sndr);
80 ASSERT_NE(std::addressof(w_a), std::addressof(sndr)); // w_a is a new copy, not a reference to sndr.
81 sndr = stdexec::just(56);
82 const auto result = stdexec::sync_wait(std::move(w_a)); // NOLINT(performance-move-const-arg)
83 ASSERT_TRUE(result.has_value());
84 ASSERT_EQ(std::get<0>(*result), 42);
85 ASSERT_EQ(
86 std::get<0>(*stdexec::sync_wait(stdexec::when_all(std::move(sndr)))), 56); // NOLINT(performance-move-const-arg)
87
88 // lvalues must be copyable; rvalues move.
89 using just_not_copy_constructible_t = decltype(stdexec::just(std::make_unique<int>(0)));
90 static_assert(!std::copy_constructible<just_not_copy_constructible_t>);
91 static_assert(!std::invocable<const stdexec::when_all_t&, just_not_copy_constructible_t&>);
92 static_assert(std::invocable<const stdexec::when_all_t&, just_not_copy_constructible_t&&>);
93}
94
107TEST_F(WhenAllTest, single_branch) {
108 const view_s_t data(Kokkos::view_alloc(exec, "data - shared space"));
109
110 const context_t esc{exec};
111
112 auto sndr = stdexec::when_all(stdexec::schedule(esc.get_scheduler()) | THEN_INCREMENT(data));
113
114 static_assert(std::same_as<
115 decltype(stdexec::get_completion_domain<stdexec::set_value_t>(stdexec::get_env(sndr))),
117 >);
118
119 static_assert(std::same_as<stdexec::tag_of_t<decltype(sndr)>, stdexec::then_t>);
121
122 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
123
124 ASSERT_THAT(
126 testing::ElementsAre(
129
130 ASSERT_EQ(data(), 1);
131}
132
143TEST_F(WhenAllTest, schedule_sender_and_single_branch) {
144 const view_s_t data(Kokkos::view_alloc(exec, "data - shared space"));
145
146 const context_t esc{exec};
147
148 auto sndr = stdexec::when_all(
149 stdexec::schedule(esc.get_scheduler()), stdexec::schedule(esc.get_scheduler()) | THEN_INCREMENT(data));
150
151 static_assert(std::same_as<
152 decltype(stdexec::get_completion_domain<stdexec::set_value_t>(stdexec::get_env(sndr))),
154 >);
155
159
160 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
161
162 const auto recorded_events = Tests::Utils::record_sync_wait<recorder_listener_t>(std::move(sndr));
163
166 ASSERT_THAT(recorded_events, [&]() {
168 return testing::ElementsAre(
171 MATCHER_FOR_WAIT_EVENT(recorded_events.at(1)));
172 } else {
173 return testing::ElementsAre(
175 MATCHER_FOR_BEGIN_FENCE(exec, dispatch_label(exec, "after dispatch")));
176 }
177 }());
178
179 ASSERT_EQ(data(), 1);
180}
181
193TEST_F(WhenAllTest, schedule_sender_and_single_branch_followed_by_self) {
194 const view_s_t data(Kokkos::view_alloc(exec, "data - shared space"));
195
196 const context_t esc{exec};
197
198 auto sndr = stdexec::when_all(
199 stdexec::schedule(esc.get_scheduler()),
200 stdexec::schedule(esc.get_scheduler()) | THEN_INCREMENT(data))
201 | stdexec::continues_on(esc.get_scheduler()) | THEN_INCREMENT(data);
202
203 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
204
205 ASSERT_THAT(
207 testing::ElementsAre(
211
212 ASSERT_EQ(data(), 2);
213}
214
227TEST_F(WhenAllTest, schedule_sender_and_single_mixed_branch_followed_by_self) {
228 const view_s_t data(Kokkos::view_alloc(exec, "data - shared space"));
229
230 const context_t esc{exec};
231 experimental::execution::single_thread_context stc{};
232
233 auto sndr = stdexec::when_all(
234 stdexec::schedule(stc.get_scheduler()),
235 stdexec::schedule(esc.get_scheduler()) | THEN_INCREMENT(data)
236 | stdexec::continues_on(stc.get_scheduler()) | THEN_INCREMENT(data))
237 | stdexec::continues_on(esc.get_scheduler()) | THEN_INCREMENT(data);
238
239 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
240
242
243 ASSERT_THAT(
245 testing::ElementsAre(
250
251 ASSERT_EQ(data(), 3);
252}
253
265TEST_F(WhenAllTest, schedule_sender_and_single_branch_followed_by_other_and_finish_on_self) {
266 const view_s_t data(Kokkos::view_alloc(exec, "data - shared space"));
267
268 const context_t esc{exec};
269 experimental::execution::single_thread_context stc{};
270
271 auto sndr = stdexec::when_all(
272 stdexec::schedule(esc.get_scheduler()),
273 stdexec::schedule(esc.get_scheduler()) | THEN_INCREMENT(data))
274 | stdexec::continues_on(stc.get_scheduler()) | THEN_INCREMENT(data)
275 | stdexec::continues_on(esc.get_scheduler()) | THEN_INCREMENT(data);
276
277 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
278
280
281 const auto recorded_events = Tests::Utils::record_sync_wait<recorder_listener_t>(std::move(sndr));
282
283 ASSERT_THAT(recorded_events, [&]() {
285 return testing::ElementsAre(
288 MATCHER_FOR_WAIT_EVENT(recorded_events.at(1)),
291 } else {
292 return testing::ElementsAre(
294 MATCHER_FOR_BEGIN_FENCE(exec, dispatch_label(exec, "after dispatch")),
297 }
298 }());
299
300 ASSERT_EQ(data(), 3);
301}
302
313TEST_F(WhenAllTest, two_mixed_branches_followed_by_self) {
314 const view_s_t data(Kokkos::view_alloc("data - shared space"));
315
316 const context_t esc{exec};
317 experimental::execution::single_thread_context stc{};
318
319 auto w_a = stdexec::when_all(
320 stdexec::schedule(esc.get_scheduler()) | THEN_INCREMENT_ATOMIC(System, data),
321 stdexec::schedule(stc.get_scheduler()) | THEN_INCREMENT_ATOMIC(System, data)
322 | stdexec::continues_on(esc.get_scheduler()));
323
324 static_assert(
325 std::same_as<
326 decltype(stdexec::get_completion_domain<stdexec::set_value_t>(stdexec::get_env(w_a), stdexec::env<>{})),
328 >);
329
330 auto sndr = std::move(w_a) | stdexec::continues_on(esc.get_scheduler()) | THEN_INCREMENT(data);
331
332 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
333
335
336 ASSERT_THAT(
338 testing::ElementsAre(
342
343 ASSERT_EQ(data(), 3);
344}
345
356TEST_F(WhenAllTest, two_branches_followed_by_self) {
357 const view_s_t data(Kokkos::view_alloc(exec, "data - shared space"));
358
359 const auto [exec_A, exec_B] = Kokkos::Experimental::partition_space(exec, 1, 1);
360
361 const context_t esc_A{exec_A}, esc_B{exec_B};
362
363 auto sndr = stdexec::when_all(
364 stdexec::schedule(esc_A.get_scheduler()) | THEN_INCREMENT_ATOMIC(Device, data),
365 stdexec::schedule(esc_B.get_scheduler()) | THEN_INCREMENT_ATOMIC(Device, data))
366 | stdexec::continues_on(esc_A.get_scheduler()) | THEN_INCREMENT(data);
367
368 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
369
370 const auto recorded_events = Tests::Utils::record_sync_wait<recorder_listener_t>(std::move(sndr));
371
372 if (Tests::Utils::are_same_instances(exec_A, exec_B)) {
373 ASSERT_THAT(
374 recorded_events,
375 testing::ElementsAre(
376 MATCHER_FOR_BEGIN_PFOR(exec_A, dispatch_label(exec_A, "then")),
377 MATCHER_FOR_BEGIN_PFOR(exec_B, dispatch_label(exec_B, "then")),
378 MATCHER_FOR_BEGIN_PFOR(exec_A, dispatch_label(exec_A, "then")),
379 MATCHER_FOR_BEGIN_FENCE(exec_A, dispatch_label(exec_A, "sync_wait"))));
380 } else {
381 ASSERT_THAT(recorded_events, [&]() {
383 return testing::ElementsAre(
387 MATCHER_FOR_WAIT_EVENT(recorded_events.at(2)),
389 MATCHER_FOR_BEGIN_FENCE(exec_A, dispatch_label(exec, "sync_wait")));
390
391 } else {
392 return testing::ElementsAre(
393 MATCHER_FOR_BEGIN_PFOR(exec_A, dispatch_label(exec_A, "then")),
394 MATCHER_FOR_BEGIN_PFOR(exec_B, dispatch_label(exec_B, "then")),
395 MATCHER_FOR_BEGIN_FENCE(exec_B, dispatch_label(exec_B, "after dispatch")),
396 MATCHER_FOR_BEGIN_PFOR(exec_A, dispatch_label(exec_A, "then")),
397 MATCHER_FOR_BEGIN_FENCE(exec_A, dispatch_label(exec_A, "sync_wait")));
398 }
399 }());
400 }
401
402 ASSERT_EQ(data(), 3);
403}
404
416TEST_F(WhenAllTest, two_branches_host_device_followed_by_device) {
418
419 const view_s_t data(Kokkos::view_alloc("data - shared space"));
420
421 const context_t esc{exec};
422
423 const host_execution_space exec_h{};
424 const context_h_t esc_h{exec_h};
425
426 auto sndr = stdexec::when_all(
427 stdexec::schedule(esc.get_scheduler()) | THEN_INCREMENT_ATOMIC(System, data),
428 stdexec::schedule(esc_h.get_scheduler()) | THEN_INCREMENT_ATOMIC(System, data))
429 | stdexec::continues_on(esc.get_scheduler()) | THEN_INCREMENT(data);
430
431 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
432
433 const auto recorded_events = Tests::Utils::record_sync_wait<recorder_listener_t>(std::move(sndr));
434
436 ASSERT_THAT(
437 recorded_events,
438 testing::ElementsAre(
440 MATCHER_FOR_BEGIN_PFOR(exec_h, dispatch_label(exec_h, "then")),
443 } else {
444 ASSERT_THAT(recorded_events, [&]() {
446 return testing::ElementsAre(
448 MATCHER_FOR_BEGIN_PFOR(exec_h, dispatch_label(exec_h, "then")),
450 MATCHER_FOR_WAIT_EVENT(recorded_events.at(2)),
453 } else {
454 return testing::ElementsAre(
456 MATCHER_FOR_BEGIN_PFOR(exec_h, dispatch_label(exec_h, "then")),
457 MATCHER_FOR_BEGIN_FENCE(exec_h, dispatch_label(exec_h, "after dispatch")),
460 }
461 }());
462 }
463
464 ASSERT_EQ(data(), 3);
465}
466
478TEST_F(WhenAllTest, two_mixed_branches_followed_by_other_and_finish_on_self) {
479 const view_s_t data(Kokkos::view_alloc("data - shared space"));
480
481 const context_t esc{exec};
482 experimental::execution::single_thread_context stc{};
483
484 auto sndr = stdexec::when_all(
485 stdexec::schedule(esc.get_scheduler()) | THEN_INCREMENT_ATOMIC(System, data)
486 | stdexec::continues_on(stc.get_scheduler()),
487 stdexec::schedule(stc.get_scheduler()) | THEN_INCREMENT_ATOMIC(System, data))
488 | stdexec::continues_on(stc.get_scheduler()) | THEN_INCREMENT(data)
489 | stdexec::continues_on(esc.get_scheduler()) | THEN_INCREMENT(data);
490
491 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
492
494
495 ASSERT_THAT(
497 testing::ElementsAre(
502
503 ASSERT_EQ(data(), 4);
504}
505
523TEST_F(WhenAllTest, nested_with_inner_followed_by_other) {
524 const view_s_t data(Kokkos::view_alloc("data - shared space"));
525
526 const context_t esc{exec};
527 experimental::execution::single_thread_context stc{};
528
529 auto sndr = stdexec::when_all(
530 stdexec::schedule(esc.get_scheduler()) | THEN_INCREMENT_ATOMIC(System, data),
531 stdexec::when_all(
532 stdexec::schedule(esc.get_scheduler()),
533 stdexec::schedule(esc.get_scheduler()) | THEN_INCREMENT_ATOMIC(System, data))
535 | stdexec::continues_on(stc.get_scheduler()) | THEN_INCREMENT_ATOMIC(System, data)
536 | stdexec::continues_on(esc.get_scheduler()))
537 | stdexec::continues_on(esc.get_scheduler()) | THEN_INCREMENT(data);
538
539 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
540
542
543 const auto recorded_events = Tests::Utils::record_sync_wait<recorder_listener_t>(std::move(sndr));
544
545 ASSERT_THAT(recorded_events, [&]() {
547 return testing::ElementsAre(
551 MATCHER_FOR_WAIT_EVENT(recorded_events.at(2)),
554 } else {
555 return testing::ElementsAre(
558 MATCHER_FOR_BEGIN_FENCE(exec, dispatch_label(exec, "after dispatch")),
561 }
562 }());
563
564 ASSERT_EQ(data(), 4);
565}
566
578TEST_F(WhenAllTest, nested_when_all_with_independent_branch) {
579 const auto [exec_A, exec_B, exec_C] = Kokkos::Experimental::partition_space(exec, 1, 1, 1);
580
581 const context_t esc_A{exec_A}, esc_B{exec_B}, esc_C{exec_C};
582
583 auto br_A = ::stdexec::schedule(esc_A.get_scheduler()) | THEN_LABELED_PFOR(TEST_EXECUTION_SPACE, 'A');
584 auto br_B = ::stdexec::schedule(esc_B.get_scheduler()) | THEN_LABELED_PFOR(TEST_EXECUTION_SPACE, 'B');
585 auto br_C = ::stdexec::schedule(esc_C.get_scheduler()) | THEN_LABELED_PFOR(TEST_EXECUTION_SPACE, 'C');
586
587 auto when_AB_then_D = ::stdexec::when_all(std::move(br_A), std::move(br_B))
588 | ::stdexec::continues_on(esc_A.get_scheduler()) | THEN_LABELED_PFOR(TEST_EXECUTION_SPACE, 'D');
589
590 auto sndr = ::stdexec::when_all(std::move(when_AB_then_D), std::move(br_C));
591
592 const auto recorded_events = Tests::Utils::record_sync_wait<recorder_listener_t>(std::move(sndr));
593
594 if (Tests::Utils::are_same_instances(exec_A, exec_B)) {
595 ASSERT_THAT(
596 recorded_events,
597 testing::ElementsAre(
598 MATCHER_FOR_BEGIN_PFOR(exec_A, "'A'"),
599 MATCHER_FOR_BEGIN_PFOR(exec_B, "'B'"),
600 MATCHER_FOR_BEGIN_PFOR(exec_A, "'D'"),
601 MATCHER_FOR_BEGIN_FENCE(exec_A, dispatch_label(exec_A, "after dispatch")),
602 MATCHER_FOR_BEGIN_PFOR(exec_C, "'C'"),
603 MATCHER_FOR_BEGIN_FENCE(exec_C, dispatch_label(exec_C, "after dispatch"))));
604 } else {
605 ASSERT_THAT(recorded_events, [&]() {
607 return testing::ElementsAre(
608 MATCHER_FOR_BEGIN_PFOR(exec_A, "'A'"),
609 MATCHER_FOR_BEGIN_PFOR(exec_B, "'B'"),
611 MATCHER_FOR_BEGIN_PFOR(exec_C, "'C'"),
613 MATCHER_FOR_WAIT_EVENT(recorded_events.at(2)),
614 MATCHER_FOR_BEGIN_PFOR(exec_A, "'D'"),
616 MATCHER_FOR_WAIT_EVENT(recorded_events.at(4)),
617 MATCHER_FOR_WAIT_EVENT(recorded_events.at(7)));
618 } else {
619 return testing::ElementsAre(
620 MATCHER_FOR_BEGIN_PFOR(exec_A, "'A'"),
621 MATCHER_FOR_BEGIN_PFOR(exec_B, "'B'"),
622 MATCHER_FOR_BEGIN_FENCE(exec_B, dispatch_label(exec_B, "after dispatch")),
623 MATCHER_FOR_BEGIN_PFOR(exec_A, "'D'"),
624 MATCHER_FOR_BEGIN_FENCE(exec_A, dispatch_label(exec_A, "after dispatch")),
625 MATCHER_FOR_BEGIN_PFOR(exec_C, "'C'"),
626 MATCHER_FOR_BEGIN_FENCE(exec_C, dispatch_label(exec_C, "after dispatch")));
627 }
628 }());
629 }
630}
631
644TEST_F(WhenAllTest, many_concurrent_branches) {
645 const view_s_t data(Kokkos::view_alloc(exec, "data - shared space"));
646
647 constexpr size_t num_branches = 6;
648
649 unsigned int counter_start = 0, counter_after_stc = 0;
650 std::array<unsigned int, num_branches> order_start{}, order_after_stc{};
651
652 const auto [exec_A, exec_B, exec_C, exec_D, exec_E, exec_F] =
653 Kokkos::Experimental::partition_space(exec, 1, 1, 1, 1, 1, 1);
654
655 using view_um_h_t = Kokkos::View<unsigned int, Kokkos::HostSpace, Kokkos::MemoryTraits<Kokkos::Unmanaged>>;
656 using increment_and_memorize_t = Tests::Utils::Functors::FetchIncrement<view_um_h_t, false>;
657
658#define DEFINE_ONE_BRANCH(_letter_, _id_) \
659 const context_t esc_##_letter_{exec_##_letter_}; \
660 experimental::execution::single_thread_context stc_##_letter_{}; \
661 auto br_##_letter_ = stdexec::just() \
662 | stdexec::then( \
663 increment_and_memorize_t{ \
664 .counter = view_um_h_t{std::addressof(counter_start)}, \
665 .value = view_um_h_t{std::addressof(order_start.at(_id_))}}) \
666 | stdexec::continues_on(stc_##_letter_.get_scheduler()) \
667 | stdexec::then( \
668 increment_and_memorize_t{ \
669 .counter = view_um_h_t{std::addressof(counter_after_stc)}, \
670 .value = view_um_h_t{std::addressof(order_after_stc.at(_id_))}}) \
671 | stdexec::continues_on(esc_##_letter_.get_scheduler()) | THEN_INCREMENT_ATOMIC(Device, data);
672
679
680 auto sndr = stdexec::when_all(
681 std::move(br_A), std::move(br_B), std::move(br_C), std::move(br_D), std::move(br_E), std::move(br_F));
682
683 ASSERT_EQ(data(), 0) << "Eager execution is not allowed.";
684
686
687 stdexec::sync_wait(std::move(sndr));
688
689 ASSERT_EQ(counter_start, num_branches);
690 ASSERT_EQ(counter_after_stc, num_branches);
691
692 ASSERT_EQ(data(), num_branches);
693
695 ASSERT_THAT(order_start, testing::ElementsAre(0, 1, 2, 3, 4, 5));
696
699 // NOLINTBEGIN(modernize-use-std-print)
700#define SHOW_ONE_BRANCH_ORDER_AFTER_STC(_letter_, _id_) \
701 SCOPED_TRACE(testing::Message() << "Branch " #_letter_ " order after 'stc' is " << order_after_stc.at(_id_));
702
709 // NOLINTEND(modernize-use-std-print)
710
711 ASSERT_THAT(order_after_stc, testing::UnorderedElementsAre(0, 1, 2, 3, 4, 5));
712}
713
714} // namespace Tests::ExecutionSpaceImpl
constexpr std::string dispatch_label(const Exec &, Label &&label)
Get the dispatch label from Exec and label.
#define MATCHER_FOR_WAIT_EVENT(_record_event_variant_)
#define MATCHER_FOR_BEGIN_PFOR(_exec_, _label_)
#define MATCHER_FOR_RECORD_EVENT(_exec_)
#define MATCHER_FOR_BEGIN_FENCE(_exec_, _label_)
RecorderListener< EventDiscardMatcher< TEST_EXECUTION_SPACE >, BeginFenceEvent, BeginParallelForEvent, Kokkos::Execution::Impl::RecordEvent, Kokkos::Execution::Impl::WaitEvent > recorder_listener_t
Concept for a sender whose completion scheduler is Kokkos::Execution::ExecutionSpaceImpl::Scheduler.
#define KOKKOS_EXECUTION_THREADS_THROWS_ON_SYNC_WAIT_ASSERT_AND_SKIP(_sndr_)
Definition context.hpp:69
Kokkos::DefaultHostExecutionSpace host_execution_space
#define DEFINE_ONE_BRANCH(_letter_, _id_)
#define SHOW_ONE_BRANCH_ORDER_AFTER_STC(_letter_, _id_)
#define KOKKOS_EXECUTION_STDEXEC_PRAGMA_DIAGNOSTIC_IGNORED
Basic list of ignored diagnostics when including anything from stdexec.
#define THEN_INCREMENT(_data_)
Add a then using Tests::Utils::Functors::Increment that may throw. // NOLINTNEXTLINE(cppcoreguideline...
Definition increment.hpp:59
#define THEN_INCREMENT_ATOMIC(_scope_, _data_)
Same as THEN_INCREMENT, using Tests::Utils::atomic_fetch_add. // NOLINTNEXTLINE(cppcoreguidelines-mac...
Definition increment.hpp:63
#define THEN_LABELED_PFOR(_exec_, _id_)
Add a Kokkos::Execution::parallel_for using Tests::Utils::Functors::Labeled. // NOLINTNEXTLINE(cppcor...
Definition labeled.hpp:21
constexpr check_rcvr_env_queryable_with_t< false, Queries... > check_rcvr_env_not_queryable_with
auto record_sync_wait(Sndr &&sndr)
Definition sync_wait.hpp:14
bool are_same_instances(const Exec &exec, const OtherExec &other_exec)
Definition kokkos.hpp:13
Matcher to filter out events that are just noise for tests.
Execution context using a Kokkos execution space under the hood.
auto get_scheduler() const noexcept -> ExecutionSpaceImpl::Scheduler< Exec >
Event to be sent to Kokkos::utils::callbacks::dispatch when calling record.
Definition event.hpp:54
Event to be sent to Kokkos::utils::callbacks::dispatch when calling wait.
Definition event.hpp:75
Kokkos::Execution::ExecutionSpaceContext< Exec > context_t
Definition context.hpp:27