kokkos-execution 0.0.1
Loading...
Searching...
No Matches
when_all.hpp
Go to the documentation of this file.
1#ifndef KOKKOS_EXECUTION_GRAPH_WHEN_ALL_HPP
2#define KOKKOS_EXECUTION_GRAPH_WHEN_ALL_HPP
3
5
6#if defined(KOKKOS_EXECUTION_ENABLE_DEBUG_LOGGING)
7# include "plog/Log.h"
8#endif
9
10#include "Kokkos_Graph.hpp"
11
13
24
26
33template <Kokkos::ExecutionSpace Exec, stdexec::receiver Rcvr, stdexec::sender... Sndrs>
35 : public Impl::Immovable
36 , public OpStateBase<Exec, Rcvr> {
37 using operation_state_concept = stdexec::operation_state_tag;
38
40 using execution_space = Exec;
41
43 using root_t = typename state_t::graph_t::root_t;
44
46 struct WhenAllChildReceiver : public Impl::Receiver<WhenAllOpState, stdexec::env_of_t<Rcvr>> {
52 [[nodiscard]]
53 auto query(get_node_t) const & noexcept -> const root_t& {
54 return this->parent_op->root;
55 }
56 };
57
58 using children_op_states_t = stdexec::__tuple<stdexec::connect_result_t<Sndrs, WhenAllChildReceiver>...>;
59
60#if defined(KOKKOS_ENABLE_DEBUG)
62 static_assert(
63 stdexec::__mapply<
64 stdexec::__mall_of<stdexec::__q<Impl::queryable_for<get_node_t>::type>>,
66 >::value,
67 "Child senders of the 'when_all' must lead to 'get_node_t' queryable operation states.");
68#endif
69
72
73 using node_t = decltype(stdexec::__apply(
74 [](const auto&... ops) { return Kokkos::Experimental::when_all(ops.query(get_node)...); },
75 std::declval<const children_op_states_t&>()));
76
82 std::atomic<size_t> count = sizeof...(Sndrs);
83
85 WhenAllOpState(stdexec::__tuple<Sndrs...>&& sndrs_, Rcvr&& rcvr_)
86 : base_t(std::move(rcvr_))
93 , state{Kokkos::Experimental::get_device_handle(execution_space{})}
94 , root(state.graph.root_node())
96 stdexec::__apply(
97 [this](auto&&... children) -> children_op_states_t {
98 static_assert((std::same_as<decltype(children), Sndrs&&> && ...));
100 stdexec::connect(std::forward<Sndrs>(children), WhenAllChildReceiver{this})...};
101 },
102 std::move(sndrs_)))
103 , node(
104 stdexec::__apply(
105 [](const auto&... child_op) {
106 auto agg = Kokkos::Experimental::when_all(child_op.query(get_node)...);
107 graph_add_aggregate_node_event(agg, child_op.query(get_node)...);
108 return agg;
109 },
110 children_op_states)) {
111 }
112
113 [[nodiscard]]
114 auto query(get_node_t) const & noexcept -> const node_t& {
115 return node;
116 }
117
118 [[nodiscard]]
119 auto query(get_graph_t) const & noexcept -> const typename state_t::graph_t& {
120 return state.graph;
121 }
122
147 void submit() & noexcept requires(!as_one)
148 {
149 if (count.fetch_sub(1) == 1) {
150 this->submit_graph();
151 }
152 }
153
154 void submit() & noexcept requires(as_one)
155 {
156 this->submit_graph();
157 }
158
159 void complete(stdexec::set_value_t) & noexcept {
160 this->submit();
161 }
162
163 template <typename Tag, typename... Args>
164 requires(!std::same_as<Tag, stdexec::set_value_t>)
165 void complete(Tag, Args&&... args) & noexcept {
166 base_t::complete(Tag{}, std::forward<Args>(args)...);
167 }
168
169 void start() & noexcept requires(!as_one)
170 {
171#if defined(KOKKOS_EXECUTION_ENABLE_DEBUG_LOGGING)
172 PLOG_INFO << "Starting all branches before submission.";
173#endif
174 stdexec::__apply([](auto&... ops) -> void { (stdexec::start(ops), ...); }, children_op_states);
175 }
176
178 void start() & noexcept requires(as_one)
179 {
180#if defined(KOKKOS_EXECUTION_ENABLE_DEBUG_LOGGING)
181 PLOG_INFO << "Submit the graph directly without starting the branches.";
182#endif
183 this->submit();
184 }
185
186 void submit_graph() & noexcept {
187#if defined(KOKKOS_EXECUTION_ENABLE_DEBUG_LOGGING)
188 PLOG_INFO << "Submitting graph " << get_graph_impl_ptr(state.graph.root_node()) << " on "
189 << Kokkos::Tools::Experimental::device_id(state.get_device_handle().m_exec) << '.';
190#endif
191 try {
192 Kokkos::Execution::GraphImpl::submit_graph(state.graph, state.get_device_handle().m_exec);
193 } catch (...) {
194 stdexec::set_error(std::move(this->completion_signal.rcvr), std::current_exception());
195 return;
196 }
197 this->completion_signal.propagate(state.get_device_handle().m_exec);
198 }
199
201};
202
204template <stdexec::operation_state OpState, Kokkos::ExecutionSpace Exec>
205requires(
206 stdexec::__is_instance_of<OpState, Kokkos::Execution::GraphImpl::WhenAllOpState>
207 && std::same_as<typename OpState::execution_space, Exec>)
208struct GraphOperationStateFor<OpState, Exec> : public std::true_type { };
209
211template <stdexec::operation_state OpState, Kokkos::ExecutionSpace Exec>
212requires(
214 && stdexec::__is_instance_of<OpState, Kokkos::Execution::GraphImpl::WhenAllOpState>)
216 template <stdexec::operation_state ChildOpState>
218
219 static constexpr bool value = stdexec::__mapply<
220 stdexec::__mall_of<stdexec::__q<RemainsOnGraphForChild>>,
221 typename OpState::children_op_states_t
222 >::value;
223};
224
226template <Kokkos::ExecutionSpace Exec, stdexec::sender... Sndrs>
228 using sender_concept = stdexec::sender_tag;
229
230 using sndrs_t = stdexec::__tuple<Sndrs...>;
231
232 struct attrs {
233 template <typename... Env>
234 [[nodiscard]]
235 constexpr auto
236 query(stdexec::get_completion_domain_t<stdexec::set_value_t>, const Env&...) const noexcept -> Domain {
237 return {};
238 }
239 };
240
242 template <typename Self, typename... Env>
243 static consteval auto get_completion_signatures() {
244 return stdexec::completion_signatures<stdexec::set_value_t(), stdexec::set_error_t(std::exception_ptr)>{};
245 }
246
247 template <stdexec::receiver Rcvr>
248 stdexec::operation_state auto connect(Rcvr rcvr) && noexcept(
249 std::is_nothrow_constructible_v<WhenAllOpState<Exec, Rcvr, Sndrs...>, sndrs_t&&, Rcvr&&>) {
250 return WhenAllOpState<Exec, Rcvr, Sndrs...>(std::move(sndrs), std::move(rcvr));
251 }
252
253 constexpr auto get_env() const noexcept -> attrs {
254 return {};
255 }
256
258};
259
260struct BECAUSE_THE_EXECUTION_SPACE_TYPE_IS_NOT_HOMOGENEOUS;
261
262template <size_t Index, typename Sndr>
264
265template <size_t Index, typename Sndr>
267
268template <>
269struct TransformSenderFor<stdexec::when_all_t> {
270 template <typename Env, typename... Sndrs>
271 using trnsfrmd_sndr_t = WhenAllSender<Impl::exec_of_t<stdexec::__m_at_c<0, Sndrs...>, Env>, Sndrs...>;
272
273 template <typename Env, typename... Sndrs>
274 auto operator()(const Env&, stdexec::when_all_t, stdexec::__ignore, Sndrs&&... sndrs) const
275 noexcept(std::is_nothrow_constructible_v<typename trnsfrmd_sndr_t<Env, Sndrs...>::sndrs_t, Sndrs&&...>) {
276 if constexpr ((graph_completing_sender<Sndrs, Env> && ...)) {
277 using execution_space = Impl::exec_of_t<stdexec::__m_at_c<0, Sndrs...>, Env>;
278
280 if constexpr ((std::same_as<Impl::exec_of_t<Sndrs, Env>, execution_space> && ...)) {
281 return trnsfrmd_sndr_t<Env, Sndrs...>{.sndrs = {std::forward<Sndrs>(sndrs)...}};
282 } else {
284 STDEXEC_CONSTEXPR_LOCAL bool map[] = {
285 !std::same_as<Impl::exec_of_t<stdexec::__m_at_c<0, Sndrs>, Env>, execution_space>...};
286 STDEXEC_CONSTEXPR_LOCAL std::size_t index = stdexec::__pos_of(map, map + sizeof...(Sndrs));
287 using invalid_sndr_t = stdexec::__m_at_c<index, Sndrs...>;
288 return stdexec::__not_a_sender<
289 stdexec::_WHAT_(CANNOT_DISPATCH_THIS_ALGORITHM_TO_THE_GRAPH_SCHEDULER),
290 stdexec::_WHY_(BECAUSE_THE_EXECUTION_SPACE_TYPE_IS_NOT_HOMOGENEOUS),
291 stdexec::_WHERE_(stdexec::_IN_ALGORITHM_, stdexec::when_all_t),
293 stdexec::_WITH_PRETTY_SENDERS_<Sndrs...>,
294 stdexec::_WITH_ENVIRONMENT_(Env)
295 >{};
296 }
297 } else {
299 STDEXEC_CONSTEXPR_LOCAL bool map[] = {!graph_completing_sender<Sndrs, Env>...};
300 STDEXEC_CONSTEXPR_LOCAL std::size_t index = stdexec::__pos_of(map, map + sizeof...(Sndrs));
301 using invalid_sndr_t = stdexec::__m_at_c<index, Sndrs...>;
303 }
304 }
305};
306
307} // namespace Kokkos::Execution::GraphImpl
308
309// NOLINTBEGIN(bugprone-reserved-identifier)
310namespace stdexec::__detail {
311template <typename... Sndrs>
312extern __mtype<Kokkos::Execution::GraphImpl::WhenAllSender<__demangle_t<Sndrs>...>>
313 __demangle_v<Kokkos::Execution::GraphImpl::WhenAllSender<Sndrs...>>;
314} // namespace stdexec::__detail
315// NOLINTEND(bugprone-reserved-identifier)
316
317#endif // KOKKOS_EXECUTION_GRAPH_WHEN_ALL_HPP
Concept for a sender whose completion scheduler is Kokkos::Execution::GraphImpl::Scheduler.
#define KOKKOS_EXECUTION_GET_ENV(_type_, _obj_)
Retrieve the environment of _obj_. // NOLINTNEXTLINE(cppcoreguidelines-macro-usage).
Definition env.hpp:14
void graph_add_aggregate_node_event(const NodeType &aggregate, const Predecessors &... predecessors)
Record an event for an aggregate node added after predecessors.
Definition events.hpp:148
WITH_SENDER_AT_INDEX< Index, stdexec::__demangle_t< Sndr > > WITH_PRETTY_SENDER_AT_INDEX
Definition when_all.hpp:266
constexpr get_node_t get_node
Definition get_node.hpp:15
auto * get_graph_impl_ptr(const NodeType &node) noexcept
Retrieve the raw graph pointer from a node.
Definition events.hpp:94
void submit_graph(const Kokkos::Experimental::Graph< Exec > &graph, const Exec &exec)
Submit a graph and record the associated event with graph_submit_event.
Definition events.hpp:179
auto no_graph_scheduler_in_env() noexcept
Show a better compile diagnostic when there is no Kokkos::Execution::GraphImpl::Scheduler found.
typename ExecOf< Args... >::type exec_of_t
Definition get_exec.hpp:37
void complete(stdexec::set_error_t, Error &&error) noexcept
constexpr OpStateBase(Rcvr rcvr) noexcept(std::is_nothrow_constructible_v< completion_signal_t, Rcvr && >)
RemainsOnGraphFor< ChildOpState, Exec > RemainsOnGraphForChild
Definition when_all.hpp:217
WhenAllSender< Impl::exec_of_t< stdexec::__m_at_c< 0, Sndrs... >, Env >, Sndrs... > trnsfrmd_sndr_t
Definition when_all.hpp:271
auto operator()(const Env &, stdexec::when_all_t, stdexec::__ignore, Sndrs &&... sndrs) const noexcept(std::is_nothrow_constructible_v< typename trnsfrmd_sndr_t< Env, Sndrs... >::sndrs_t, Sndrs &&... >)
Definition when_all.hpp:274
auto query(get_node_t) const &noexcept -> const root_t &
Definition when_all.hpp:53
Operation state for stdexec::when_all.
Definition when_all.hpp:36
WhenAllOpState(stdexec::__tuple< Sndrs... > &&sndrs_, Rcvr &&rcvr_)
Definition when_all.hpp:85
State< GraphComposition::Create, execution_space > state_t
Definition when_all.hpp:42
stdexec::operation_state_tag operation_state_concept
Definition when_all.hpp:37
static constexpr bool as_one
Determine if all branches remain fully on the graph, if connected to WhenAllChildReceiver.
Definition when_all.hpp:71
decltype(stdexec::__apply([](const auto &... ops) { return Kokkos::Experimental::when_all(ops.query(get_node)...);}, std::declval< const children_op_states_t & >())) node_t
Definition when_all.hpp:73
void start() &noexcept
If as_one is true, there is no need to start the branches.
Definition when_all.hpp:178
typename state_t::graph_t::root_t root_t
Definition when_all.hpp:43
stdexec::__tuple< stdexec::connect_result_t< Sndrs, WhenAllChildReceiver >... > children_op_states_t
Definition when_all.hpp:58
auto query(get_graph_t) const &noexcept -> const typename state_t::graph_t &
Definition when_all.hpp:119
void complete(Tag, Args &&... args) &noexcept
Definition when_all.hpp:165
void complete(stdexec::set_value_t) &noexcept
Definition when_all.hpp:159
auto query(get_node_t) const &noexcept -> const node_t &
Definition when_all.hpp:114
constexpr auto query(stdexec::get_completion_domain_t< stdexec::set_value_t >, const Env &...) const noexcept -> Domain
Definition when_all.hpp:236
Sender for stdexec::when_all.
Definition when_all.hpp:227
constexpr auto get_env() const noexcept -> attrs
Definition when_all.hpp:253
static consteval auto get_completion_signatures()
Definition when_all.hpp:243
stdexec::__tuple< Sndrs... > sndrs_t
Definition when_all.hpp:230
stdexec::operation_state auto connect(Rcvr rcvr) &&noexcept(std::is_nothrow_constructible_v< WhenAllOpState< Exec, Rcvr, Sndrs... >, sndrs_t &&, Rcvr && >)
Definition when_all.hpp:248
Receiver for an object parent_op that implements complete.
Definition receiver.hpp:13
Kokkos::DefaultExecutionSpace execution_space