1#ifndef FENIX_MPIXX_TASKS_HPP
2#define FENIX_MPIXX_TASKS_HPP
10#include "fenix/mpixx/status.hpp"
11#include "fenix/mpixx/request.hpp"
12#include "fenix/mpixx/util.hpp"
13#include "fenix/tasks/task.hpp"
15namespace fenix::mpixx {
20MPITask recv(T* b,
int n, MPI_Datatype d,
int r,
int t, MPI_Comm c) {
22 Status ret = MPI_Irecv(b, n, d, r, t, c, &request);
23 if (ret) ret =
co_await request;
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);
31auto recv(T& b,
int r,
int t, MPI_Comm c) {
32 return recv(&b, 1, r, t, c);
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);
40MPITask send(
const T* b,
int n, MPI_Datatype d,
int r,
int t, MPI_Comm c) {
42 Status ret = MPI_Isend(b, n, d, r, t, c, &request);
43 if (ret) ret =
co_await request;
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);
51auto send(
const T& b,
int r,
int t, MPI_Comm c) {
52 return send(&b, 1, r, t, c);
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);
59template <
typename ST,
typename RT>
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
64 auto recv_task = recv(rb, rn, rd, rr, rt, c);
67 co_await send(sb, sn, sd, sr, st, c);
68 co_return co_await recv_task;
70template <
typename ST,
typename RT>
72 const ST* sb,
int sn,
int sr,
int st,
73 RT* rb,
int rn,
int rr,
int rt, MPI_Comm c
76 sb, datatype_count(sb, sn), datatype(sb), sr, st,
77 rb, datatype_count(rb, rn), datatype(rb), rr, rt, c
80template <
typename ST,
typename RT>
82 const ST& sb,
int sr,
int st,
83 RT& rb,
int rr,
int rt, MPI_Comm c
85 return sendrecv(&sb, 1, sr, st, &rb, 1, rr, rt, c);
87template <
typename ST,
typename SA,
typename RT,
typename RA>
89 const std::vector<ST, SA>& sv,
int sr,
int st,
90 std::vector<RT, RA>& rv,
int rr,
int rt, MPI_Comm c
92 return sendrecv(&sv[0], sv.size(), sr, st, &rv[0], rv.size(), rr, rt, c);
97 const void* sb, T& rb,
int n, MPI_Datatype d, MPI_Op o, MPI_Comm c
100 Status ret = MPI_Iallreduce(sb, &rb, n, d, o, c, &request);
101 if (ret) ret =
co_await request;
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);
109auto allreduce(
const T& sb, T& rb, MPI_Op o, MPI_Comm c) {
110 return allreduce(&sb, rb, 1, o, c);
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);
120 const void* sb, T* rb,
int n, MPI_Datatype d, MPI_Op o,
int r, MPI_Comm c
123 Status ret = MPI_Ireduce(sb, rb, n, d, o, r, c, &request);
124 if (ret) ret =
co_await request;
131 const void* sb, T& rb,
int n, MPI_Datatype d, MPI_Op o,
int r, MPI_Comm c
134 Status ret = MPI_Ireduce(sb, &rb, n, d, o, r, c, &request);
135 if (ret) ret =
co_await request;
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);
143auto reduce(
const T& sb, T& rb, MPI_Op o,
int r, MPI_Comm c) {
144 return reduce(&sb, rb, 1, o, r, c);
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);
152MPITask bcast(T* b,
int n, MPI_Datatype d,
int r, MPI_Comm c) {
154 Status ret = MPI_Ibcast(b, n, d, r, c, &request);
155 if (ret) ret =
co_await request;
159auto bcast(T* b,
int n,
int r, MPI_Comm c) {
160 return bcast(b, datatype_count(b, n), datatype(b), r, c);
163auto bcast(T& b,
int r, MPI_Comm c) {
164 return bcast(b, 1, r, c);
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);
171template <
typename ST,
typename RT>
173 const ST* sb,
int sn, MPI_Datatype sd, RT* rb,
int rn, MPI_Datatype rd,
177 Status ret = MPI_Iallgather(sb, sn, sd, rb, rn, rd, c, &request);
178 if (ret) ret =
co_await request;
181template <
typename ST,
typename RT>
182auto allgather(
const ST* sb,
int sn, RT* rb,
int rn, MPI_Comm c) {
184 sb, datatype_count(sb, sn), datatype(sb), rb, datatype_count(rb, rn),
189template <
typename ST,
typename RT>
191 const ST* sb,
int sn, MPI_Datatype sd, RT* rb,
const int* rn,
192 const int* displs, MPI_Datatype rd, MPI_Comm c
195 Status ret = MPI_Iallgatherv(sb, sn, sd, rb, rn, displs, rd, c, &request);
196 if (ret) ret =
co_await request;
199template <
typename ST,
typename RT>
201 const ST* sb,
int sn, RT* rb,
const int* rn,
const int* displs, MPI_Comm c
204 sb, datatype_count(sb, sn), datatype(sb), rb, rn, displs, datatype(rb), c
208inline MPITask probe(
int src,
int tag, MPI_Comm comm) {
213 ret = MPI_Iprobe(src, tag, comm, &found, ret);
214 if (found || !ret)
co_return ret;
215 co_await std::suspend_always{};