Fenix @develop
 
Loading...
Searching...
No Matches
tasks.hpp
1#ifndef FENIX_MPIXX_TASKS_HPP
2#define FENIX_MPIXX_TASKS_HPP
3
4#include <type_traits>
5#include <utility>
6#include <vector>
7
8#include <mpi.h>
9
10#include "fenix/mpixx/status.hpp"
11#include "fenix/mpixx/request.hpp"
12#include "fenix/mpixx/util.hpp"
13#include "fenix/tasks/task.hpp"
14
15namespace fenix::mpixx {
16
17using MPITask = fenix::tasks::Task<Status>;
18
19template <typename T>
20MPITask recv(T* b, int n, MPI_Datatype d, int r, int t, MPI_Comm c) {
21 MPI_Request request;
22 Status ret = MPI_Irecv(b, n, d, r, t, c, &request);
23 if (ret) ret = co_await request;
24 co_return ret;
25}
26template <typename T>
27auto recv(T* b, int n, int r, int t, MPI_Comm c) {
28 return recv(b, datatype_count(b, n), datatype(b), r, t, c);
29}
30template <typename T>
31auto recv(T& b, int r, int t, MPI_Comm c) {
32 return recv(&b, 1, r, t, c);
33}
34template <typename T, typename A>
35auto recv(std::vector<T, A>& v, int r, int t, MPI_Comm c) {
36 return recv(v.data(), v.size(), r, t, c);
37}
38
39template <typename T>
40MPITask send(const T* b, int n, MPI_Datatype d, int r, int t, MPI_Comm c) {
41 MPI_Request request;
42 Status ret = MPI_Isend(b, n, d, r, t, c, &request);
43 if (ret) ret = co_await request;
44 co_return ret;
45}
46template <typename T>
47auto send(const T* b, int n, int r, int t, MPI_Comm c) {
48 return send(b, datatype_count(b, n), datatype(b), r, t, c);
49}
50template <typename T>
51auto send(const T& b, int r, int t, MPI_Comm c) {
52 return send(&b, 1, r, t, c);
53}
54template <typename T, typename A>
55auto send(const std::vector<T>& v, int r, int t, MPI_Comm c) {
56 return send(v.data(), v.size(), r, t, c);
57}
58
59template <typename ST, typename RT>
60MPITask sendrecv(
61 const ST* sb, int sn, MPI_Datatype sd, int sr, int st,
62 RT* rb, int rn, MPI_Datatype rd, int rr, int rt, MPI_Comm c
63) {
64 auto recv_task = recv(rb, rn, rd, rr, rt, c);
65 // ensure lazily-evaluated recv_task actually begins
66 recv_task.resume();
67 co_await send(sb, sn, sd, sr, st, c);
68 co_return co_await recv_task;
69}
70template <typename ST, typename RT>
71auto sendrecv(
72 const ST* sb, int sn, int sr, int st,
73 RT* rb, int rn, int rr, int rt, MPI_Comm c
74) {
75 return sendrecv(
76 sb, datatype_count(sb, sn), datatype(sb), sr, st,
77 rb, datatype_count(rb, rn), datatype(rb), rr, rt, c
78 );
79}
80template <typename ST, typename RT>
81auto sendrecv(
82 const ST& sb, int sr, int st,
83 RT& rb, int rr, int rt, MPI_Comm c
84) {
85 return sendrecv(&sb, 1, sr, st, &rb, 1, rr, rt, c);
86}
87template <typename ST, typename SA, typename RT, typename RA>
88auto sendrecv(
89 const std::vector<ST, SA>& sv, int sr, int st,
90 std::vector<RT, RA>& rv, int rr, int rt, MPI_Comm c
91) {
92 return sendrecv(&sv[0], sv.size(), sr, st, &rv[0], rv.size(), rr, rt, c);
93}
94
95template <typename T>
96MPITask allreduce(
97 const void* sb, T& rb, int n, MPI_Datatype d, MPI_Op o, MPI_Comm c
98) {
99 MPI_Request request;
100 Status ret = MPI_Iallreduce(sb, &rb, n, d, o, c, &request);
101 if (ret) ret = co_await request;
102 co_return ret;
103}
104template <typename T>
105auto allreduce(const T* sb, T& rb, int n, MPI_Op o, MPI_Comm c) {
106 return allreduce(sb, rb, datatype_count(sb, n), datatype(sb), o, c);
107}
108template <typename T>
109auto allreduce(const T& sb, T& rb, MPI_Op o, MPI_Comm c) {
110 return allreduce(&sb, rb, 1, o, c);
111}
112template <typename T, typename A>
113auto allreduce(const std::vector<T, A>& sv, T& rb, MPI_Op o, MPI_Comm c) {
114 return allreduce(&sv[0], rb, sv.size(), o, c);
115}
116
117// Template for pointer types - pass pointer directly to MPI
118template <typename T>
119MPITask reduce(
120 const void* sb, T* rb, int n, MPI_Datatype d, MPI_Op o, int r, MPI_Comm c
121) {
122 MPI_Request request;
123 Status ret = MPI_Ireduce(sb, rb, n, d, o, r, c, &request);
124 if (ret) ret = co_await request;
125 co_return ret;
126}
127
128// Template for reference types - take address of reference
129template <typename T>
130MPITask reduce(
131 const void* sb, T& rb, int n, MPI_Datatype d, MPI_Op o, int r, MPI_Comm c
132) {
133 MPI_Request request;
134 Status ret = MPI_Ireduce(sb, &rb, n, d, o, r, c, &request);
135 if (ret) ret = co_await request;
136 co_return ret;
137}
138template <typename T>
139auto reduce(const T* sb, T& rb, int n, MPI_Op o, int r, MPI_Comm c) {
140 return reduce(sb, rb, datatype_count(sb, n), datatype(sb), o, r, c);
141}
142template <typename T>
143auto reduce(const T& sb, T& rb, MPI_Op o, int r, MPI_Comm c) {
144 return reduce(&sb, rb, 1, o, r, c);
145}
146template <typename T, typename A>
147auto reduce(const std::vector<T, A>& sv, T& rb, MPI_Op o, int r, MPI_Comm c) {
148 return reduce(&sv[0], rb, sv.size(), o, r, c);
149}
150
151template <typename T>
152MPITask bcast(T* b, int n, MPI_Datatype d, int r, MPI_Comm c) {
153 MPI_Request request;
154 Status ret = MPI_Ibcast(b, n, d, r, c, &request);
155 if (ret) ret = co_await request;
156 co_return ret;
157}
158template <typename T>
159auto bcast(T* b, int n, int r, MPI_Comm c) {
160 return bcast(b, datatype_count(b, n), datatype(b), r, c);
161}
162template <typename T>
163auto bcast(T& b, int r, MPI_Comm c) {
164 return bcast(b, 1, r, c);
165}
166template <typename T, typename A>
167auto bcast(std::vector<T, A>& v, int r, MPI_Comm c) {
168 return bcast(&v[0], v.size(), r, c);
169}
170
171template <typename ST, typename RT>
172MPITask allgather(
173 const ST* sb, int sn, MPI_Datatype sd, RT* rb, int rn, MPI_Datatype rd,
174 MPI_Comm c
175) {
176 MPI_Request request;
177 Status ret = MPI_Iallgather(sb, sn, sd, rb, rn, rd, c, &request);
178 if (ret) ret = co_await request;
179 co_return ret;
180}
181template <typename ST, typename RT>
182auto allgather(const ST* sb, int sn, RT* rb, int rn, MPI_Comm c) {
183 return allgather(
184 sb, datatype_count(sb, sn), datatype(sb), rb, datatype_count(rb, rn),
185 datatype(rb), c
186 );
187}
188
189template <typename ST, typename RT>
190MPITask allgatherv(
191 const ST* sb, int sn, MPI_Datatype sd, RT* rb, const int* rn,
192 const int* displs, MPI_Datatype rd, MPI_Comm c
193) {
194 MPI_Request request;
195 Status ret = MPI_Iallgatherv(sb, sn, sd, rb, rn, displs, rd, c, &request);
196 if (ret) ret = co_await request;
197 co_return ret;
198}
199template <typename ST, typename RT>
200auto allgatherv(
201 const ST* sb, int sn, RT* rb, const int* rn, const int* displs, MPI_Comm c
202) {
203 return allgatherv(
204 sb, datatype_count(sb, sn), datatype(sb), rb, rn, displs, datatype(rb), c
205 );
206}
207
208inline MPITask probe(int src, int tag, MPI_Comm comm) {
209 int found;
210 Status ret;
211 do {
212 int found;
213 ret = MPI_Iprobe(src, tag, comm, &found, ret);
214 if (found || !ret) co_return ret;
215 co_await std::suspend_always{};
216 } while (true);
217}
218
219} // namespace fenix::mpixx
220
221#endif // FENIX_MPIXX_TASKS_HPP
Definition task.hpp:15