Skip to content

Commit a8c0eb4

Browse files
linpeizePeizeLin
andauthored
Refactor: simplify Mixing DM (#7844)
* 1. add OpenMP in XC_Functional_Libxc::v_xc_libxc() 2. add OpenMP in RI_2D_Comm::split_m2D_ktoR() * Refactor: simplify Mixing DM --------- Co-authored-by: linpz <linpz@mail.ustc.edu.cn>
1 parent ee763df commit a8c0eb4

11 files changed

Lines changed: 189 additions & 477 deletions

source/Makefile.Objects

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -601,7 +601,6 @@ OBJS_MODULE_RI=conv_coulomb_pot_k.o\
601601
Matrix_Orbs22.o\
602602
RI_2D_Comm.o\
603603
Mix_DMk_2D.o\
604-
Mix_Matrix.o\
605604
symmetry_rotation.o\
606605
symmetry_irreducible_sector.o\
607606

source/module_ri/CMakeLists.txt

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,6 @@ if (ENABLE_LIBRI)
88
Matrix_Orbs22.cpp
99
RI_2D_Comm.cpp
1010
Mix_DMk_2D.cpp
11-
Mix_Matrix.cpp
1211
)
1312

1413
if(ENABLE_LCAO)

source/module_ri/Exx_LRI_interface.h

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -120,7 +120,8 @@ class Exx_LRI_Interface
120120
std::shared_ptr<Exx_LRI<Tdata>> exx_ptr;
121121

122122
private:
123-
Mix_DMk_2D mix_DMk_2D;
123+
124+
Mix_DMk_2D<T> mix_DMk_2D;
124125

125126
bool exx_spacegroup_symmetry = false;
126127
ModuleSymmetry::Symmetry_rotation symrot_;

source/module_ri/Exx_LRI_interface.hpp

Lines changed: 42 additions & 42 deletions
Original file line numberDiff line numberDiff line change
@@ -62,9 +62,7 @@ void Exx_LRI_Interface<T, Tdata>::cal_exx_ions(const UnitCell& ucell, const bool
6262
ModuleBase::TITLE("Exx_LRI_Interface","cal_exx_ions");
6363
if(!this->flag_finish.init)
6464
{ throw std::runtime_error("Exx init unfinished when "+std::string(__FILE__)+" line "+std::to_string(__LINE__)); }
65-
6665
this->exx_ptr->cal_exx_ions(ucell, write_cv);
67-
6866
this->flag_finish.ions = true;
6967
}
7068

@@ -79,7 +77,6 @@ void Exx_LRI_Interface<T, Tdata>::cal_exx_elec(const std::vector<std::map<TA, st
7977
{ throw std::runtime_error("Exx init unfinished when "+std::string(__FILE__)+" line "+std::to_string(__LINE__)); }
8078

8179
this->exx_ptr->cal_exx_elec(Ds, ucell, pv, p_symrot);
82-
8380
this->flag_finish.elec = true;
8481
}
8582

@@ -93,7 +90,6 @@ void Exx_LRI_Interface<T, Tdata>::cal_exx_force(const int& nat)
9390
{ throw std::runtime_error("Exx Hamiltonian unfinished when "+std::string(__FILE__)+" line "+std::to_string(__LINE__)); }
9491

9592
this->exx_ptr->cal_exx_force(nat);
96-
9793
this->flag_finish.force = true;
9894
}
9995

@@ -107,12 +103,14 @@ void Exx_LRI_Interface<T, Tdata>::cal_exx_stress(const double& omega, const doub
107103
{ throw std::runtime_error("Exx Hamiltonian unfinished when "+std::string(__FILE__)+" line "+std::to_string(__LINE__)); }
108104

109105
this->exx_ptr->cal_exx_stress(omega, lat0);
110-
111106
this->flag_finish.stress = true;
112107
}
113108

114109
template<typename T, typename Tdata>
115-
void Exx_LRI_Interface<T, Tdata>::exx_before_all_runners(const K_Vectors& kv, const UnitCell& ucell, const Parallel_2D& pv)
110+
void Exx_LRI_Interface<T, Tdata>::exx_before_all_runners(
111+
const K_Vectors& kv,
112+
const UnitCell& ucell,
113+
const Parallel_2D& pv)
116114
{
117115
ModuleBase::TITLE("Exx_LRI_Interface","exx_before_all_runners");
118116
// initialize the rotation matrix in AO representation
@@ -166,14 +164,15 @@ void Exx_LRI_Interface<T, Tdata>::exx_beforescf(const int istep,
166164
if(GlobalC::exx_info.info_global.cal_exx)
167165
{
168166
if (this->exx_spacegroup_symmetry)
169-
{this->mix_DMk_2D.set_nks(kv.get_nkstot_full() * (PARAM.inp.nspin == 2 ? 2 : 1), PARAM.globalv.gamma_only_local);}
167+
{ this->mix_DMk_2D.set_nks(kv.get_nkstot_full() * (PARAM.inp.nspin == 2 ? 2 : 1)); }
170168
else
171-
{this->mix_DMk_2D.set_nks(kv.get_nks(), PARAM.globalv.gamma_only_local);}
169+
{ this->mix_DMk_2D.set_nks(kv.get_nks()); }
172170

173-
if(GlobalC::exx_info.info_global.separate_loop)
174-
{ this->mix_DMk_2D.set_mixing(nullptr); }
171+
if (GlobalC::exx_info.info_global.separate_loop)
172+
{ this->mix_DMk_2D.set_mixing_plain(GlobalC::exx_info.info_global.mixing_beta_for_loop1); }
175173
else
176174
{ this->mix_DMk_2D.set_mixing(chgmix.get_mixing()); }
175+
177176
// for exx two_level scf
178177
this->two_level_step = 0;
179178
}
@@ -190,40 +189,39 @@ void Exx_LRI_Interface<T, Tdata>::exx_eachiterinit(const int istep,
190189
ModuleBase::TITLE("Exx_LRI_Interface","exx_eachiterinit");
191190
if (GlobalC::exx_info.info_global.cal_exx)
192191
{
193-
if (!GlobalC::exx_info.info_global.separate_loop && (this->two_level_step || istep > 0 || PARAM.inp.init_wfc == "file") // non separate loop case
194-
|| (GlobalC::exx_info.info_global.separate_loop && PARAM.inp.init_wfc == "file" && this->two_level_step == 0 && iter == 1)) // the first iter in separate loop case
192+
if (!GlobalC::exx_info.info_global.separate_loop
193+
&& (this->two_level_step
194+
|| istep > 0
195+
|| PARAM.inp.init_wfc == "file") // non separate loop case
196+
|| (GlobalC::exx_info.info_global.separate_loop
197+
&& PARAM.inp.init_wfc == "file"
198+
&& this->two_level_step == 0
199+
&& iter == 1)
200+
) // the first iter in separate loop case
195201
{
196202
const bool flag_restart = (iter == 1) ? true : false;
197203
auto cal = [this, &ucell,&kv, &flag_restart](const elecstate::DensityMatrix<T, double>& dm_in)
198204
{
199205
if (this->exx_spacegroup_symmetry)
200-
{ this->mix_DMk_2D.mix(symrot_.restore_dm(kv,dm_in.get_DMK_vector(), *dm_in.get_paraV_pointer()), flag_restart); }
206+
{ this->mix_DMk_2D.mix(symrot_.restore_dm(kv, dm_in.get_DMK_vector(), *dm_in.get_paraV_pointer()), flag_restart); }
201207
else
202208
{ this->mix_DMk_2D.mix(dm_in.get_DMK_vector(), flag_restart); }
203-
const std::vector<std::map<TA, std::map<TAC, RI::Tensor<Tdata>>>>
204-
Ds = PARAM.globalv.gamma_only_local
205-
? RI_2D_Comm::split_m2D_ktoR<Tdata>(
206-
ucell,
207-
*this->exx_ptr->p_kv,
208-
this->mix_DMk_2D.get_DMk_gamma_out(),
209-
*dm_in.get_paraV_pointer(),
210-
PARAM.inp.nspin)
211-
: RI_2D_Comm::split_m2D_ktoR<Tdata>(
212-
ucell,
213-
*this->exx_ptr->p_kv,
214-
this->mix_DMk_2D.get_DMk_k_out(),
215-
*dm_in.get_paraV_pointer(),
216-
PARAM.inp.nspin,
217-
this->exx_spacegroup_symmetry);
218-
219-
if (this->exx_spacegroup_symmetry && GlobalC::exx_info.info_global.exx_symmetry_realspace)
209+
const std::vector<std::map<TA, std::map<TAC, RI::Tensor<Tdata>>>> Ds =
210+
RI_2D_Comm::split_m2D_ktoR<Tdata>(
211+
ucell,
212+
*this->exx_ptr->p_kv,
213+
this->mix_DMk_2D.get_DMk_out(),
214+
*dm_in.get_paraV_pointer(),
215+
PARAM.inp.nspin,
216+
this->exx_spacegroup_symmetry);
217+
if(this->exx_spacegroup_symmetry && GlobalC::exx_info.info_global.exx_symmetry_realspace)
220218
{ this->cal_exx_elec(Ds, ucell,*dm_in.get_paraV_pointer(), &this->symrot_); }
221219
else
222220
{ this->cal_exx_elec(Ds, ucell,*dm_in.get_paraV_pointer()); }
223221
};
224222

225223
if(istep > 0 && flag_restart)
226-
{ cal(*dm_last_step); }
224+
{ cal(*this->dm_last_step); }
227225
else
228226
{ cal(dm); }
229227
}
@@ -396,21 +394,23 @@ bool Exx_LRI_Interface<T, Tdata>::exx_after_converge(
396394
// if init_wfc == "file", DM is calculated in the 1st iter of the 1st two-level step, so we mix it here
397395
const bool flag_restart = (this->two_level_step == 0 && PARAM.inp.init_wfc != "file") ? true : false;
398396

399-
if (this->exx_spacegroup_symmetry)
400-
{this->mix_DMk_2D.mix(symrot_.restore_dm(kv, dm.get_DMK_vector(), *dm.get_paraV_pointer()), flag_restart);}
397+
if(this->exx_spacegroup_symmetry)
398+
{ this->mix_DMk_2D.mix(symrot_.restore_dm(kv, dm.get_DMK_vector(), *dm.get_paraV_pointer()), flag_restart); }
401399
else
402-
{this->mix_DMk_2D.mix(dm.get_DMK_vector(), flag_restart);}
403-
404-
// GlobalC::exx_lcao.cal_exx_elec(p_esolver->LOC, p_esolver->LOWF.wfc_k_grid);
405-
const std::vector<std::map<int, std::map<std::pair<int, std::array<int, 3>>, RI::Tensor<Tdata>>>>
406-
Ds = std::is_same<T, double>::value //gamma_only_local
407-
? RI_2D_Comm::split_m2D_ktoR<Tdata>(ucell,*this->exx_ptr->p_kv, this->mix_DMk_2D.get_DMk_gamma_out(), *dm.get_paraV_pointer(), nspin)
408-
: RI_2D_Comm::split_m2D_ktoR<Tdata>(ucell,*this->exx_ptr->p_kv, this->mix_DMk_2D.get_DMk_k_out(), *dm.get_paraV_pointer(), nspin, this->exx_spacegroup_symmetry);
409-
410-
if (this->exx_spacegroup_symmetry && GlobalC::exx_info.info_global.exx_symmetry_realspace)
400+
{ this->mix_DMk_2D.mix(dm.get_DMK_vector(), flag_restart); }
401+
const std::vector<std::map<TA, std::map<TAC, RI::Tensor<Tdata>>>> Ds =
402+
RI_2D_Comm::split_m2D_ktoR<Tdata>(
403+
ucell,
404+
*this->exx_ptr->p_kv,
405+
this->mix_DMk_2D.get_DMk_out(),
406+
*dm.get_paraV_pointer(),
407+
nspin,
408+
this->exx_spacegroup_symmetry);
409+
if(this->exx_spacegroup_symmetry && GlobalC::exx_info.info_global.exx_symmetry_realspace)
411410
{ this->cal_exx_elec(Ds, ucell, *dm.get_paraV_pointer(), &this->symrot_); }
412411
else
413412
{ this->cal_exx_elec(Ds, ucell, *dm.get_paraV_pointer()); } // restore DM but not Hexx
413+
414414
iter = 0;
415415
this->two_level_step++;
416416

source/module_ri/Mix_DMk_2D.cpp

Lines changed: 63 additions & 47 deletions
Original file line numberDiff line numberDiff line change
@@ -4,69 +4,85 @@
44
//=======================
55

66
#include "Mix_DMk_2D.h"
7+
#include "module_base/module_mixing/plain_mixing.h"
78
#include "module_base/tool_title.h"
89

9-
Mix_DMk_2D &Mix_DMk_2D::set_nks(const int nks, const bool gamma_only_in)
10+
#include <cassert>
11+
12+
template <typename Tdata>
13+
Mix_DMk_2D<Tdata>::~Mix_DMk_2D<Tdata>()
14+
{
15+
if(this->flag_del_mixing)
16+
delete this->mixing;
17+
}
18+
19+
template <typename Tdata>
20+
void Mix_DMk_2D<Tdata>::set_nks(const int nks)
1021
{
11-
ModuleBase::TITLE("Mix_DMk_2D", "set_nks");
12-
this->gamma_only = gamma_only_in;
13-
if (this->gamma_only)
14-
this->mix_DMk_gamma.resize(nks);
15-
else
16-
this->mix_DMk_k.resize(nks);
17-
return *this;
22+
this->mix_DMk.clear();
23+
this->mix_DMk.resize(nks);
1824
}
1925

20-
Mix_DMk_2D &Mix_DMk_2D::set_mixing(Base_Mixing::Mixing* mixing_in)
26+
template <typename Tdata>
27+
void Mix_DMk_2D<Tdata>::set_mixing(Base_Mixing::Mixing* mixing_in)
2128
{
22-
ModuleBase::TITLE("Mix_DMk_2D","set_mixing");
23-
if(this->gamma_only)
24-
for (Mix_Matrix<std::vector<double>>& mix_one : this->mix_DMk_gamma)
25-
mix_one.init(mixing_in);
26-
else
27-
for (Mix_Matrix<std::vector<std::complex<double>>>& mix_one : this->mix_DMk_k)
28-
mix_one.init(mixing_in);
29-
return *this;
29+
if(this->flag_del_mixing)
30+
delete this->mixing;
31+
this->mixing = mixing_in;
32+
this->flag_del_mixing = false;
3033
}
3134

32-
Mix_DMk_2D &Mix_DMk_2D::set_mixing_beta(const double mixing_beta)
35+
template <typename Tdata>
36+
void Mix_DMk_2D<Tdata>::set_mixing_plain(const double& mixing_beta)
3337
{
34-
ModuleBase::TITLE("Mix_DMk_2D","set_mixing_beta");
35-
if(this->gamma_only)
36-
for (Mix_Matrix<std::vector<double>>& mix_one : this->mix_DMk_gamma)
37-
mix_one.mixing_beta = mixing_beta;
38-
else
39-
for (Mix_Matrix<std::vector<std::complex<double>>>& mix_one : this->mix_DMk_k)
40-
mix_one.mixing_beta = mixing_beta;
41-
return *this;
38+
if(this->flag_del_mixing)
39+
delete this->mixing;
40+
this->mixing = new Base_Mixing::Plain_Mixing(mixing_beta);
41+
this->flag_del_mixing = true;
4242
}
4343

44-
void Mix_DMk_2D::mix(const std::vector<std::vector<double>>& dm, const bool flag_restart)
44+
template <typename Tdata>
45+
void Mix_DMk_2D<Tdata>::mix(const std::vector<std::vector<Tdata>>& dm, const bool flag_restart)
4546
{
46-
ModuleBase::TITLE("Mix_DMk_2D","mix");
47-
assert(this->mix_DMk_gamma.size() == dm.size());
48-
for(int ik=0; ik<dm.size(); ++ik)
49-
this->mix_DMk_gamma[ik].mix(dm[ik], flag_restart);
47+
ModuleBase::TITLE("Mix_DMk_2D", "mix");
48+
if (flag_restart)
49+
{ this->restart_all(dm); }
50+
else
51+
{ this->mix_all(dm); }
5052
}
51-
void Mix_DMk_2D::mix(const std::vector<std::vector<std::complex<double>>>& dm, const bool flag_restart)
53+
54+
template <typename Tdata>
55+
std::vector<const std::vector<Tdata>*> Mix_DMk_2D<Tdata>::get_DMk_out() const
5256
{
53-
ModuleBase::TITLE("Mix_DMk_2D","mix");
54-
assert(this->mix_DMk_k.size() == dm.size());
55-
for(int ik=0; ik<dm.size(); ++ik)
56-
this->mix_DMk_k[ik].mix(dm[ik], flag_restart);
57+
std::vector<const std::vector<Tdata>*> DMk_out(this->mix_DMk.size());
58+
for (int ik = 0; ik < this->mix_DMk.size(); ++ik)
59+
{ DMk_out[ik] = &this->mix_DMk[ik].data_out; }
60+
return DMk_out;
5761
}
5862

59-
std::vector<const std::vector<double>*> Mix_DMk_2D::get_DMk_gamma_out() const
63+
template <typename Tdata>
64+
void Mix_DMk_2D<Tdata>::restart_all(const std::vector<std::vector<Tdata>>& data_in)
6065
{
61-
std::vector<const std::vector<double>*> DMk_out(this->mix_DMk_gamma.size());
62-
for(int ik=0; ik<this->mix_DMk_gamma.size(); ++ik)
63-
DMk_out[ik] = &this->mix_DMk_gamma[ik].get_data_out();
64-
return DMk_out;
66+
assert(this->mix_DMk.size() == data_in.size());
67+
assert(this->mixing != nullptr);
68+
for (int ik = 0; ik < data_in.size(); ++ik)
69+
{
70+
this->mix_DMk[ik].data_out = data_in[ik];
71+
this->mixing->init_mixing_data(this->mix_DMk[ik].mixing_data, data_in[ik].size(), sizeof(Tdata));
72+
}
6573
}
66-
std::vector<const std::vector<std::complex<double>>*> Mix_DMk_2D::get_DMk_k_out() const
74+
75+
template <typename Tdata>
76+
void Mix_DMk_2D<Tdata>::mix_all(const std::vector<std::vector<Tdata>>& data_in)
6777
{
68-
std::vector<const std::vector<std::complex<double>>*> DMk_out(this->mix_DMk_k.size());
69-
for(int ik=0; ik<this->mix_DMk_k.size(); ++ik)
70-
DMk_out[ik] = &this->mix_DMk_k[ik].get_data_out();
71-
return DMk_out;
72-
}
78+
assert(this->mix_DMk.size() == data_in.size());
79+
assert(this->mixing != nullptr);
80+
for (int ik = 0; ik < data_in.size(); ++ik)
81+
{
82+
this->mixing->push_data(this->mix_DMk[ik].mixing_data, this->mix_DMk[ik].data_out.data(), data_in[ik].data(), nullptr, false);
83+
this->mixing->mix_data(this->mix_DMk[ik].mixing_data, this->mix_DMk[ik].data_out.data());
84+
}
85+
}
86+
87+
template class Mix_DMk_2D<double>;
88+
template class Mix_DMk_2D<std::complex<double>>;

0 commit comments

Comments
 (0)