|
9 | 9 | #include <KokkosComm/concepts.hpp> |
10 | 10 | #include <KokkosComm/traits.hpp> |
11 | 11 | #include <KokkosComm/datatype.hpp> |
| 12 | +#include "mpi_space.hpp" |
| 13 | +#include "req.hpp" |
12 | 14 |
|
13 | | -#include "impl/pack_traits.hpp" |
14 | 15 | #include "impl/error_handling.hpp" |
| 16 | +#include "impl/pack_traits.hpp" |
15 | 17 |
|
16 | 18 | namespace KokkosComm::mpi { |
17 | 19 |
|
| 20 | +template <KokkosExecutionSpace ExecSpace, KokkosView SView, KokkosView RView> |
| 21 | +auto ireduce(const ExecSpace &space, const SView sv, RView rv, MPI_Op op, int root, MPI_Comm comm) -> Req<MpiSpace> { |
| 22 | + using ST = typename SView::non_const_value_type; |
| 23 | + using RT = typename RView::non_const_value_type; |
| 24 | + using SPkr = typename PackTraits<SView>::packer_type; |
| 25 | + using RPkr = typename PackTraits<RView>::packer_type; |
| 26 | + static_assert(std::is_same_v<ST, RT>, "KokkosComm::mpi::ireduce: View value types must be identical"); |
| 27 | + static_assert(rank<SView>() <= 1 and rank<RView>() <= 1, |
| 28 | + "KokkosComm::mpi::ireduce: Views with rank higher than 1 are not supported"); |
| 29 | + Kokkos::Tools::pushRegion("KokkosComm::mpi::ireduce"); |
| 30 | + |
| 31 | + const int rank = [=]() { |
| 32 | + int _r; |
| 33 | + MPI_Comm_rank(comm, &_r); |
| 34 | + return _r; |
| 35 | + }(); |
| 36 | + |
| 37 | + Req<MpiSpace> req; |
| 38 | + if (is_contiguous(sv)) { |
| 39 | + if (rank == root and not is_contiguous(rv)) { |
| 40 | + auto pkd_rv = RPkr::allocate_packed_for(space, "KC::mpi::ireduce c_sv pkd_rv", rv); |
| 41 | + space.fence("fence allocation before MPI call"); |
| 42 | + MPI_Ireduce(data_handle(sv), data_handle(pkd_rv.view), span(sv), datatype<MpiSpace, ST>(), op, root, comm, |
| 43 | + &req.mpi_request()); |
| 44 | + RPkr::unpack_into(space, rv, pkd_rv.view); |
| 45 | + req.extend_view_lifetime(pkd_rv.view); |
| 46 | + } else { |
| 47 | + space.fence("fence before MPI call"); |
| 48 | + MPI_Ireduce(data_handle(sv), data_handle(rv), span(sv), datatype<MpiSpace, ST>(), op, root, comm, |
| 49 | + &req.mpi_request()); |
| 50 | + } |
| 51 | + } else { |
| 52 | + auto send_args = SPkr::pack(space, sv); |
| 53 | + if (rank == root and not is_contiguous(rv)) { |
| 54 | + auto pkd_rv = RPkr::allocate_packed_for(space, "KC::mpi::ireduce nc_sv pkd_rv", rv); |
| 55 | + space.fence("fence allocation before MPI call"); |
| 56 | + MPI_Ireduce(data_handle(send_args.view), data_handle(pkd_rv.view), send_args.count, send_args.datatype, op, root, |
| 57 | + comm, &req.mpi_request()); |
| 58 | + RPkr::unpack_into(space, rv, pkd_rv.view); // no-op |
| 59 | + req.extend_view_lifetime(pkd_rv.view); |
| 60 | + } else { |
| 61 | + space.fence("fence before MPI call"); |
| 62 | + MPI_Ireduce(data_handle(send_args.view), data_handle(rv), send_args.count, datatype<MpiSpace, ST>(), op, root, |
| 63 | + comm, &req.mpi_request()); |
| 64 | + } |
| 65 | + req.extend_view_lifetime(send_args.view); |
| 66 | + } |
| 67 | + req.extend_view_lifetime(sv); |
| 68 | + req.extend_view_lifetime(rv); |
| 69 | + |
| 70 | + Kokkos::Tools::popRegion(); |
| 71 | + return req; |
| 72 | +} |
| 73 | + |
18 | 74 | template <KokkosView SendView, KokkosView RecvView> |
19 | 75 | void reduce(const SendView &sv, const RecvView &rv, MPI_Op op, int root, MPI_Comm comm) { |
20 | 76 | Kokkos::Tools::pushRegion("KokkosComm::mpi::reduce"); |
|
0 commit comments