Skip to content

Commit dd3e3f2

Browse files
author
zxy.monado
committed
test(source_base): follow parallel cell API changes
1 parent 208f494 commit dd3e3f2

3 files changed

Lines changed: 21 additions & 16 deletions

File tree

source/source_base/test_parallel/parallel_device_test.cpp

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,8 @@
11
#ifdef __MPI
22
#include "source_base/parallel_device.h"
33

4-
#include "source_base/communication_domain.h"
4+
#include "source_base/parallel_cell.h"
5+
#include "source_base/parallel_comm.h"
56

67
#include "gtest/gtest.h"
78
#include <complex>
@@ -51,7 +52,8 @@ void exercise_mpi_wrappers(const ModuleBase::CommunicationDomain& domain)
5152
{
5253
MPI_Comm communicator = domain.communicator();
5354
const int rank = domain.rank();
54-
const int size = domain.size();
55+
MPICommGroup world_group(communicator);
56+
const int size = world_group.gsize;
5557
const int count = 2;
5658

5759
T sent[count] = {static_cast<T>(rank + 1), static_cast<T>(rank + 2)};

source/source_base/test_parallel/parallel_domain_grid_test.cpp

Lines changed: 13 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
1-
#include "source_base/communication_domain.h"
21
#include "source_base/global_variable.h"
2+
#include "source_base/parallel_cell.h"
33
#include "source_base/parallel_comm.h"
44
#include "source_base/parallel_grid.h"
55

@@ -12,32 +12,32 @@ TEST(CommunicationDomainTest, ReportsDefaultAndWorldDomains)
1212
{
1313
const ModuleBase::CommunicationDomain local_domain;
1414
EXPECT_EQ(local_domain.rank(), 0);
15-
EXPECT_EQ(local_domain.size(), 1);
15+
EXPECT_EQ(local_domain.communicator(), MPI_COMM_NULL);
1616

1717
const ModuleBase::CommunicationDomain world_domain = ModuleBase::world_communication_domain();
18+
MPICommGroup world_group(world_domain.communicator());
1819
EXPECT_EQ(world_domain.communicator(), MPI_COMM_WORLD);
1920
EXPECT_GE(world_domain.rank(), 0);
20-
EXPECT_LT(world_domain.rank(), world_domain.size());
21+
EXPECT_LT(world_domain.rank(), world_group.gsize);
2122

22-
const ModuleBase::CommunicationDomain null_domain(MPI_COMM_NULL);
23+
ModuleBase::CommunicationDomain null_domain;
24+
null_domain.initialize(MPI_COMM_NULL);
2325
EXPECT_EQ(null_domain.communicator(), MPI_COMM_NULL);
2426
EXPECT_EQ(null_domain.rank(), 0);
25-
EXPECT_EQ(null_domain.size(), 1);
2627
}
2728

2829
TEST(MPICommGroupTest, DividesWorldIntoEvenGroups)
2930
{
3031
const ModuleBase::CommunicationDomain world_domain = ModuleBase::world_communication_domain();
3132
MPICommGroup group(MPI_COMM_WORLD);
32-
EXPECT_EQ(group.gsize, world_domain.size());
3333
EXPECT_EQ(group.grank, world_domain.rank());
3434

35-
const int group_count = world_domain.size() > 1 ? 2 : 1;
35+
const int group_count = group.gsize > 1 ? 2 : 1;
3636
group.divide_group_comm(group_count);
3737

3838
EXPECT_TRUE(group.is_even);
3939
EXPECT_EQ(group.ngroups, group_count);
40-
EXPECT_EQ(group.nprocs_in_group, world_domain.size() / group_count);
40+
EXPECT_EQ(group.nprocs_in_group, group.gsize / group_count);
4141
EXPECT_EQ(group.my_group, world_domain.rank() / group.nprocs_in_group);
4242
EXPECT_EQ(group.rank_in_group, world_domain.rank() % group.nprocs_in_group);
4343
EXPECT_NE(group.group_comm, MPI_COMM_NULL);
@@ -47,25 +47,26 @@ TEST(MPICommGroupTest, DividesWorldIntoEvenGroups)
4747
TEST(ParallelGridTest, BroadcastsAndReducesDistributedGrid)
4848
{
4949
const ModuleBase::CommunicationDomain world_domain = ModuleBase::world_communication_domain();
50+
MPICommGroup world_group(world_domain.communicator());
5051
const int nx = 2;
5152
const int ny = 1;
52-
const int nz = world_domain.size();
53+
const int nz = world_group.gsize;
5354
const int local_nz = 1;
5455
const int local_size = nx * ny * local_nz;
5556

5657
legacy_global::KPAR = 1;
5758
legacy_global::MY_POOL = 0;
58-
legacy_global::NPROC = world_domain.size();
59+
legacy_global::NPROC = world_group.gsize;
5960
legacy_global::MY_RANK = world_domain.rank();
60-
legacy_global::NPROC_IN_POOL = world_domain.size();
61+
legacy_global::NPROC_IN_POOL = world_group.gsize;
6162
legacy_global::RANK_IN_POOL = world_domain.rank();
6263
legacy_global::RANK_IN_BPGROUP = world_domain.rank();
6364
POOL_WORLD = MPI_COMM_WORLD;
6465
INT_BGROUP = MPI_COMM_WORLD;
6566
KP_WORLD = MPI_COMM_NULL;
6667

6768
Parallel_Grid grid;
68-
grid.init(nx, ny, nz, local_nz, local_size, nz, 1, world_domain.size());
69+
grid.init(nx, ny, nz, local_nz, local_size, nz, 1, world_group.gsize);
6970
EXPECT_EQ(grid.get_nx(), nx);
7071
EXPECT_EQ(grid.get_ny(), ny);
7172
EXPECT_EQ(grid.get_nz(), nz);

source/source_base/test_parallel/test_para_gemm.cpp

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
1-
#include "../communication_domain.h"
21
#include "../kernels/math_kernel_op.h"
32
#include "../para_gemm.h"
3+
#include "../parallel_cell.h"
4+
#include "../parallel_comm.h"
45

56
#include <gtest/gtest.h>
67
#include <iostream>
@@ -76,7 +77,8 @@ void test_additional_type_paths()
7677
const ModuleBase::CommunicationDomain domain = ModuleBase::world_communication_domain();
7778
MPI_Comm world = domain.communicator();
7879
const int rank = domain.rank();
79-
const int size = domain.size();
80+
MPICommGroup world_group(world);
81+
const int size = world_group.gsize;
8082
const T alpha = static_cast<T>(1);
8183
const T beta = static_cast<T>(0);
8284
const T a[1] = {static_cast<T>(rank + 1)};

0 commit comments

Comments
 (0)