2#include "../common/utils.h"
7#include <oneapi/tbb/parallel_for.h>
23 MPI_Comm_rank(MPI_COMM_WORLD, &rank);
30 MPI_Comm_rank(MPI_COMM_WORLD, &rank);
31 return rank == RootID();
37 MPI_Comm_size(MPI_COMM_WORLD, &size);
41inline std::pair<size_t, size_t> LocalPart(
size_t begin,
size_t end)
44 MPI_Comm_rank(MPI_COMM_WORLD, &rank);
46 size_t size = end - begin;
47 return { rank * size / N, (rank + 1) * size / N };
50inline int LocalPart(
size_t begin,
size_t end,
int rank)
53 size_t size = end - begin;
54 return (rank + 1) * size / N - rank * size / N;
59inline void AllReduce(T* data,
int size)
61 MPI_Datatype mpi_type;
62 if constexpr (std::is_same_v<T, double>)
63 mpi_type = MPI_DOUBLE;
64 if constexpr (std::is_same_v<T, Vector3d>)
66 MPI_Type_contiguous(3, MPI_DOUBLE, &mpi_type);
67 MPI_Type_commit(&mpi_type);
69 if constexpr (std::is_same_v<T, std::complex<double>>)
70 mpi_type = MPI_COMPLEX16;
72 const int nproc = NProc();
73 std::vector<int> send_displs(nproc + 1);
74 for (
int i = 0; i <= nproc; ++i)
75 send_displs[i] = size * i / nproc;
76 std::vector<int> send_counts(nproc);
77 for (
int i = 0; i < nproc; ++i)
78 send_counts[i] = send_displs[i + 1] - send_displs[i];
80 const int local_size = send_counts[MyID()];
81 std::vector<int> recv_displs(nproc, 0);
82 for (
int i = 1; i < nproc; ++i)
83 recv_displs[i] = recv_displs[i - 1] + local_size;
84 std::vector<int> recv_counts(nproc, local_size);
86 std::unique_ptr<T[]> buf(
new T[local_size * nproc]);
87 MPI_Alltoallv(data, send_counts.data(), send_displs.data(), mpi_type, buf.get(), recv_counts.data(), recv_displs.data(), mpi_type, MPI_COMM_WORLD);
88 tbb::parallel_for(0, local_size, [&](
int i) {
89 for (
int j = 1; j < nproc; ++j)
90 buf[i] += buf[i + j * local_size];
92 MPI_Allgatherv(buf.get(), local_size, mpi_type, data, send_counts.data(), send_displs.data(), mpi_type, MPI_COMM_WORLD);