VM2D 1.14
Vortex methods for 2D flows simulation
Loading...
Searching...
No Matches
mpi_utils.h
Go to the documentation of this file.
1#pragma once
2#include "../common/utils.h"
3#include <utility>
4#include <iostream>
5
6#ifndef __NVCC__
7#include <oneapi/tbb/parallel_for.h>
8#endif
9
10#ifdef FMM_MPI
11#include "mpi.h"
12
13namespace fmm {
14
15inline int RootID()
16{
17 return 0;
18}
19
20inline int MyID()
21{
22 int rank;
23 MPI_Comm_rank(MPI_COMM_WORLD, &rank);
24 return rank;
25}
26
27inline bool IAmRoot()
28{
29 int rank;
30 MPI_Comm_rank(MPI_COMM_WORLD, &rank);
31 return rank == RootID();
32}
33
34inline int NProc()
35{
36 int size;
37 MPI_Comm_size(MPI_COMM_WORLD, &size);
38 return size;
39}
40
41inline std::pair<size_t, size_t> LocalPart(size_t begin, size_t end)
42{
43 int rank;
44 MPI_Comm_rank(MPI_COMM_WORLD, &rank);
45 int N = NProc();
46 size_t size = end - begin;
47 return { rank * size / N, (rank + 1) * size / N };
48}
49
50inline int LocalPart(size_t begin, size_t end, int rank)
51{
52 int N = NProc();
53 size_t size = end - begin;
54 return (rank + 1) * size / N - rank * size / N;
55}
56
57#ifndef __NVCC__
58template <typename T>
59inline void AllReduce(T* data, int size)
60{
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>)
65 {
66 MPI_Type_contiguous(3, MPI_DOUBLE, &mpi_type);
67 MPI_Type_commit(&mpi_type);
68 }
69 if constexpr (std::is_same_v<T, std::complex<double>>)
70 mpi_type = MPI_COMPLEX16;
71
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];
79
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);
85
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];
91 });
92 MPI_Allgatherv(buf.get(), local_size, mpi_type, data, send_counts.data(), send_displs.data(), mpi_type, MPI_COMM_WORLD);
93}
94#endif
95
96}
97#else
98
99namespace fmm {
100
101inline bool IAmRoot()
102{
103 return true;
104}
105
106}
107
108#endif
Definition avx.h:5
bool IAmRoot()
Definition mpi_utils.h:101