Skip to content

Commit d88b719

Browse files
lijianing-sudoclaudemohanchen
authored
Refactor/merge openmp (#7446)
* optimize: add OpenMP parallelization to MD per-atom loops Add #pragma omp parallel for to major per-atom loops in MD module, enabling multi-threaded execution for NEP/DPMD potentials and thermostat/integrator operations. Scope (23 files): - source/source_md/: md_base, md_func, fire, msst, nhchain, verlet, run_md, md_statistics.h - source/source_esolver/: esolver_nep, esolver_dp - source/source_md/test/: 7 unit tests + md_test_fixture.h Strategy: schedule(static) with if(nat>=256), reduction clauses, atomic/critical for shared accumulators. LJ esolver excluded (upstream refactored to UnitCellLite API). Rebased onto deepmodeling/develop. Co-Authored-By: Claude <noreply@anthropic.com> * Resolve merge conflicts and fix Verlet CSVR test# Please enter the commit message for your changes. Lines starting Resolve merge conflicts and fix Verlet CSVR test# * fix(md): restore CSVR thermostat after conflict resolution --------- Co-authored-by: Claude <noreply@anthropic.com> Co-authored-by: Mohan Chen <mohanchen@pku.edu.cn>
1 parent 995bca7 commit d88b719

22 files changed

Lines changed: 467 additions & 403 deletions

source/source_esolver/esolver_dp.cpp

Lines changed: 47 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -37,9 +37,27 @@ void ESolver_DP::before_all_runners(BaseCell& basecell, const Input_para& inp)
3737
dp_potential = 0;
3838
dp_force.create(ucell.nat, 3);
3939
dp_virial.create(3, 3);
40+
dp_cell.resize(9);
41+
dp_coord.resize(3 * ucell.nat);
42+
dp_model_force.clear();
43+
dp_model_virial.clear();
4044

4145
atype.resize(ucell.nat);
4246

47+
// Build flat atom index for OpenMP coordinate fill in runner()
48+
atom_type_index.resize(ucell.nat);
49+
atom_local_index.resize(ucell.nat);
50+
int iat = 0;
51+
for (int it = 0; it < ucell.ntype; ++it)
52+
{
53+
for (int ia = 0; ia < ucell.atoms[it].na; ++ia)
54+
{
55+
atom_type_index[iat] = it;
56+
atom_local_index[iat] = ia;
57+
iat++;
58+
}
59+
}
60+
4361
rescaling = inp.mdp.dp_rescaling;
4462
fparam = inp.mdp.dp_fparam;
4563
aparam = inp.mdp.dp_aparam;
@@ -58,38 +76,36 @@ void ESolver_DP::runner(BaseCell& basecell, const int istep)
5876
ModuleBase::TITLE("ESolver_DP", "runner");
5977
ModuleBase::timer::start("ESolver_DP", "runner");
6078

61-
std::vector<double> cell(9, 0.0);
62-
cell[0] = ucell.latvec.e11 * ucell.lat0_angstrom;
63-
cell[1] = ucell.latvec.e12 * ucell.lat0_angstrom;
64-
cell[2] = ucell.latvec.e13 * ucell.lat0_angstrom;
65-
cell[3] = ucell.latvec.e21 * ucell.lat0_angstrom;
66-
cell[4] = ucell.latvec.e22 * ucell.lat0_angstrom;
67-
cell[5] = ucell.latvec.e23 * ucell.lat0_angstrom;
68-
cell[6] = ucell.latvec.e31 * ucell.lat0_angstrom;
69-
cell[7] = ucell.latvec.e32 * ucell.lat0_angstrom;
70-
cell[8] = ucell.latvec.e33 * ucell.lat0_angstrom;
71-
72-
std::vector<double> coord(3 * ucell.nat, 0.0);
73-
int iat = 0;
74-
for (int it = 0; it < ucell.ntype; ++it)
79+
dp_cell[0] = ucell.latvec.e11 * ucell.lat0_angstrom;
80+
dp_cell[1] = ucell.latvec.e12 * ucell.lat0_angstrom;
81+
dp_cell[2] = ucell.latvec.e13 * ucell.lat0_angstrom;
82+
dp_cell[3] = ucell.latvec.e21 * ucell.lat0_angstrom;
83+
dp_cell[4] = ucell.latvec.e22 * ucell.lat0_angstrom;
84+
dp_cell[5] = ucell.latvec.e23 * ucell.lat0_angstrom;
85+
dp_cell[6] = ucell.latvec.e31 * ucell.lat0_angstrom;
86+
dp_cell[7] = ucell.latvec.e32 * ucell.lat0_angstrom;
87+
dp_cell[8] = ucell.latvec.e33 * ucell.lat0_angstrom;
88+
89+
dp_coord.resize(3 * ucell.nat);
90+
const int nat = ucell.nat;
91+
#pragma omp parallel for schedule(static) if (nat >= 256)
92+
for (int iat = 0; iat < nat; ++iat)
7593
{
76-
for (int ia = 0; ia < ucell.atoms[it].na; ++ia)
77-
{
78-
coord[3 * iat] = ucell.atoms[it].tau[ia].x * ucell.lat0_angstrom;
79-
coord[3 * iat + 1] = ucell.atoms[it].tau[ia].y * ucell.lat0_angstrom;
80-
coord[3 * iat + 2] = ucell.atoms[it].tau[ia].z * ucell.lat0_angstrom;
81-
iat++;
82-
}
94+
const int it = atom_type_index[iat];
95+
const int ia = atom_local_index[iat];
96+
dp_coord[3 * iat] = ucell.atoms[it].tau[ia].x * ucell.lat0_angstrom;
97+
dp_coord[3 * iat + 1] = ucell.atoms[it].tau[ia].y * ucell.lat0_angstrom;
98+
dp_coord[3 * iat + 2] = ucell.atoms[it].tau[ia].z * ucell.lat0_angstrom;
8399
}
84-
assert(ucell.nat == iat);
85100

86101
#ifdef __DPMD
87-
std::vector<double> f, v;
88102
dp_potential = 0;
89103
dp_force.zero_out();
90104
dp_virial.zero_out();
105+
dp_model_force.clear();
106+
dp_model_virial.clear();
91107

92-
dp.compute(dp_potential, f, v, coord, atype, cell, fparam, aparam);
108+
dp.compute(dp_potential, dp_model_force, dp_model_virial, dp_coord, atype, dp_cell, fparam, aparam);
93109

94110
// rescale the energy, force, and stress
95111
const double fact_e = rescaling / ModuleBase::Ry_to_eV;
@@ -100,18 +116,20 @@ void ESolver_DP::runner(BaseCell& basecell, const int istep)
100116
GlobalV::ofs_running << " #TOTAL ENERGY# " << std::setprecision(11) << dp_potential * ModuleBase::Ry_to_eV << " eV"
101117
<< std::endl;
102118

103-
for (int i = 0; i < ucell.nat; ++i)
119+
const int nat_f = ucell.nat;
120+
#pragma omp parallel for schedule(static) if (nat_f >= 256)
121+
for (int i = 0; i < nat_f; ++i)
104122
{
105-
dp_force(i, 0) = f[3 * i] * fact_f;
106-
dp_force(i, 1) = f[3 * i + 1] * fact_f;
107-
dp_force(i, 2) = f[3 * i + 2] * fact_f;
123+
dp_force(i, 0) = dp_model_force[3 * i] * fact_f;
124+
dp_force(i, 1) = dp_model_force[3 * i + 1] * fact_f;
125+
dp_force(i, 2) = dp_model_force[3 * i + 2] * fact_f;
108126
}
109127

110128
for (int i = 0; i < 3; ++i)
111129
{
112130
for (int j = 0; j < 3; ++j)
113131
{
114-
dp_virial(i, j) = v[3 * i + j] * fact_v;
132+
dp_virial(i, j) = dp_model_virial[3 * i + j] * fact_v;
115133
}
116134
}
117135
#else

source/source_esolver/esolver_dp.h

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -109,12 +109,18 @@ class ESolver_DP : public ESolver
109109

110110
std::string dp_file; ///< directory of DP model file
111111
std::vector<int> atype = {}; ///< atom type corresponding to DP model
112+
std::vector<int> atom_type_index; ///< type index (it) for each global atom iat
113+
std::vector<int> atom_local_index; ///< local index (ia) within type for each global atom iat
112114
std::vector<double> fparam = {}; ///< frame parameter for dp potential: dim_fparam
113115
std::vector<double> aparam = {}; ///< atomic parameter for dp potential: natoms x dim_aparam
114116
double rescaling = 1.0; ///< rescaling factor for DP model
115117
double dp_potential = 0.0; ///< computed potential energy
116118
ModuleBase::matrix dp_force; ///< computed atomic forces
117119
ModuleBase::matrix dp_virial; ///< computed lattice virials
120+
std::vector<double> dp_cell; ///< DP cell buffer in Angstrom
121+
std::vector<double> dp_coord; ///< DP coordinate buffer in Angstrom
122+
std::vector<double> dp_model_force; ///< raw force buffer returned by DP
123+
std::vector<double> dp_model_virial; ///< raw virial buffer returned by DP
118124
};
119125

120126
} // namespace ModuleESolver

source/source_esolver/esolver_nep.cpp

Lines changed: 60 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@
2121
#include "source_io/module_output/output_log.h"
2222
#include "source_io/module_parameter/parameter.h"
2323

24-
#include <numeric>
24+
#include <algorithm>
2525
#include <unordered_map>
2626

2727
using namespace ModuleESolver;
@@ -35,6 +35,9 @@ void ESolver_NEP::before_all_runners(BaseCell& basecell, const Input_para& inp)
3535
nep_force.create(ucell.nat, 3);
3636
nep_virial.create(3, 3);
3737
atype.resize(ucell.nat);
38+
nep_cell.resize(9);
39+
nep_coord.resize(3 * ucell.nat);
40+
nep_virial_sum.resize(9);
3841
_e.resize(ucell.nat);
3942
_f.resize(3 * ucell.nat);
4043
_v.resize(9 * ucell.nat);
@@ -55,51 +58,54 @@ void ESolver_NEP::runner(BaseCell& basecell, const int istep)
5558

5659
// note that NEP are column major, thus a transpose is needed
5760
// cell
58-
std::vector<double> cell(9, 0.0);
59-
cell[0] = ucell.latvec.e11 * ucell.lat0_angstrom;
60-
cell[1] = ucell.latvec.e21 * ucell.lat0_angstrom;
61-
cell[2] = ucell.latvec.e31 * ucell.lat0_angstrom;
62-
cell[3] = ucell.latvec.e12 * ucell.lat0_angstrom;
63-
cell[4] = ucell.latvec.e22 * ucell.lat0_angstrom;
64-
cell[5] = ucell.latvec.e32 * ucell.lat0_angstrom;
65-
cell[6] = ucell.latvec.e13 * ucell.lat0_angstrom;
66-
cell[7] = ucell.latvec.e23 * ucell.lat0_angstrom;
67-
cell[8] = ucell.latvec.e33 * ucell.lat0_angstrom;
61+
nep_cell[0] = ucell.latvec.e11 * ucell.lat0_angstrom;
62+
nep_cell[1] = ucell.latvec.e21 * ucell.lat0_angstrom;
63+
nep_cell[2] = ucell.latvec.e31 * ucell.lat0_angstrom;
64+
nep_cell[3] = ucell.latvec.e12 * ucell.lat0_angstrom;
65+
nep_cell[4] = ucell.latvec.e22 * ucell.lat0_angstrom;
66+
nep_cell[5] = ucell.latvec.e32 * ucell.lat0_angstrom;
67+
nep_cell[6] = ucell.latvec.e13 * ucell.lat0_angstrom;
68+
nep_cell[7] = ucell.latvec.e23 * ucell.lat0_angstrom;
69+
nep_cell[8] = ucell.latvec.e33 * ucell.lat0_angstrom;
6870

6971
// coord
70-
std::vector<double> coord(3 * ucell.nat, 0.0);
71-
int iat = 0;
72+
nep_coord.resize(3 * ucell.nat);
7273
const int nat = ucell.nat;
73-
for (int it = 0; it < ucell.ntype; ++it)
74+
#pragma omp parallel for schedule(static) if (nat >= 256)
75+
for (int iat = 0; iat < nat; ++iat)
7476
{
75-
for (int ia = 0; ia < ucell.atoms[it].na; ++ia)
76-
{
77-
coord[iat] = ucell.atoms[it].tau[ia].x * ucell.lat0_angstrom;
78-
coord[iat + nat] = ucell.atoms[it].tau[ia].y * ucell.lat0_angstrom;
79-
coord[iat + 2 * nat] = ucell.atoms[it].tau[ia].z * ucell.lat0_angstrom;
80-
iat++;
81-
}
77+
const int it = atom_type_index[iat];
78+
const int ia = atom_local_index[iat];
79+
nep_coord[iat] = ucell.atoms[it].tau[ia].x * ucell.lat0_angstrom;
80+
nep_coord[iat + nat] = ucell.atoms[it].tau[ia].y * ucell.lat0_angstrom;
81+
nep_coord[iat + 2 * nat] = ucell.atoms[it].tau[ia].z * ucell.lat0_angstrom;
8282
}
83-
assert(ucell.nat == iat);
8483

8584
#ifdef __NEP
8685
nep_potential = 0.0;
8786
nep_force.zero_out();
8887
nep_virial.zero_out();
8988

90-
nep.compute(atype, cell, coord, _e, _f, _v);
89+
nep.compute(atype, nep_cell, nep_coord, _e, _f, _v);
9190

9291
// unit conversion
9392
const double fact_e = 1.0 / ModuleBase::Ry_to_eV;
9493
const double fact_f = 1.0 / (ModuleBase::Ry_to_eV * ModuleBase::ANGSTROM_AU);
9594
const double fact_v = 1.0 / (ucell.omega * ModuleBase::Ry_to_eV);
9695

9796
// potential energy
98-
nep_potential = fact_e * std::accumulate(_e.begin(), _e.end(), 0.0);
97+
double energy_sum = 0.0;
98+
#pragma omp parallel for reduction(+:energy_sum) schedule(static) if (nat >= 256)
99+
for (int i = 0; i < nat; ++i)
100+
{
101+
energy_sum += _e[i];
102+
}
103+
nep_potential = fact_e * energy_sum;
99104
GlobalV::ofs_running << " #TOTAL ENERGY# " << std::setprecision(11) << nep_potential * ModuleBase::Ry_to_eV << " eV"
100105
<< std::endl;
101106

102107
// forces
108+
#pragma omp parallel for schedule(static) if (nat >= 256)
103109
for (int i = 0; i < nat; ++i)
104110
{
105111
nep_force(i, 0) = _f[i] * fact_f;
@@ -108,22 +114,44 @@ void ESolver_NEP::runner(BaseCell& basecell, const int istep)
108114
}
109115

110116
// virial
111-
std::vector<double> v_sum(9, 0.0);
112-
for (int j = 0; j < 9; ++j)
117+
double v0 = 0.0;
118+
double v1 = 0.0;
119+
double v2 = 0.0;
120+
double v3 = 0.0;
121+
double v4 = 0.0;
122+
double v5 = 0.0;
123+
double v6 = 0.0;
124+
double v7 = 0.0;
125+
double v8 = 0.0;
126+
#pragma omp parallel for reduction(+:v0, v1, v2, v3, v4, v5, v6, v7, v8) schedule(static) if (nat >= 256)
127+
for (int i = 0; i < nat; ++i)
113128
{
114-
for (int i = 0; i < nat; ++i)
115-
{
116-
int index = j * nat + i;
117-
v_sum[j] += _v[index];
118-
}
129+
v0 += _v[i];
130+
v1 += _v[nat + i];
131+
v2 += _v[2 * nat + i];
132+
v3 += _v[3 * nat + i];
133+
v4 += _v[4 * nat + i];
134+
v5 += _v[5 * nat + i];
135+
v6 += _v[6 * nat + i];
136+
v7 += _v[7 * nat + i];
137+
v8 += _v[8 * nat + i];
119138
}
139+
nep_virial_sum[0] = v0;
140+
nep_virial_sum[1] = v1;
141+
nep_virial_sum[2] = v2;
142+
nep_virial_sum[3] = v3;
143+
nep_virial_sum[4] = v4;
144+
nep_virial_sum[5] = v5;
145+
nep_virial_sum[6] = v6;
146+
nep_virial_sum[7] = v7;
147+
nep_virial_sum[8] = v8;
120148

121149
// virial -> stress
122150
for (int i = 0; i < 3; ++i)
123151
{
124152
for (int j = 0; j < 3; ++j)
125153
{
126-
nep_virial(i, j) = v_sum[3 * i + j] * fact_v;
154+
nep_virial(i, j) = nep_virial_sum[3 * i + j] * fact_v;
127155
}
128156
}
129157
#else

source/source_esolver/esolver_nep.h

Lines changed: 14 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -93,16 +93,21 @@ class ESolver_NEP : public ESolver
9393
NEP nep; ///< NEP object for NEP calculations
9494
#endif
9595

96-
std::string nep_file; ///< directory of NEP model file
97-
std::vector<int> atype = {}; ///< atom type mapping for NEP model
98-
double nep_potential; ///< computed potential energy
99-
ModuleBase::matrix nep_force; ///< computed atomic forces
100-
ModuleBase::matrix nep_virial; ///< computed lattice virials
101-
std::vector<double> _e; ///< temporary storage for energy computation
102-
std::vector<double> _f; ///< temporary storage for force computation
103-
std::vector<double> _v; ///< temporary storage for virial computation
96+
std::string nep_file; ///< directory of NEP model file
97+
std::vector<int> atype = {}; ///< atom type mapping for NEP model
98+
std::vector<int> atom_type_index; ///< global atom index to UnitCell atom type
99+
std::vector<int> atom_local_index; ///< global atom index to local index inside atom type
100+
double nep_potential; ///< computed potential energy
101+
ModuleBase::matrix nep_force; ///< computed atomic forces
102+
ModuleBase::matrix nep_virial; ///< computed lattice virials
103+
std::vector<double> nep_cell; ///< NEP cell buffer in Angstrom, column-major
104+
std::vector<double> nep_coord; ///< NEP coordinate buffer in Angstrom, column-major
105+
std::vector<double> nep_virial_sum; ///< summed per-atom virial components
106+
std::vector<double> _e; ///< temporary storage for energy computation
107+
std::vector<double> _f; ///< temporary storage for force computation
108+
std::vector<double> _v; ///< temporary storage for virial computation
104109
};
105110

106111
} // namespace ModuleESolver
107112

108-
#endif
113+
#endif

0 commit comments

Comments
 (0)