From 30e0184d2180c2ef76617c6f2ae039bcbe599caa Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 13:17:47 +0800 Subject: [PATCH 01/22] remove some PARAM in source_psi --- source/source_psi/psi_init_atomic.cpp | 8 +++++--- source/source_psi/psi_init_atomic.h | 4 +++- source/source_psi/psi_init_atomic_random.cpp | 6 ++++-- source/source_psi/psi_init_atomic_random.h | 4 +++- source/source_psi/psi_init_file.cpp | 8 +++++--- source/source_psi/psi_init_file.h | 4 +++- source/source_psi/psi_init_nao.cpp | 18 ++++++------------ source/source_psi/psi_init_nao.h | 4 +++- source/source_psi/psi_init_nao_random.cpp | 6 ++++-- source/source_psi/psi_init_nao_random.h | 4 +++- source/source_psi/psi_init_random.cpp | 8 +++++--- source/source_psi/psi_init_random.h | 4 +++- source/source_psi/psi_initializer.cpp | 6 +++++- source/source_psi/psi_initializer.h | 6 +++++- source/source_psi/psi_prepare.cpp | 3 ++- 15 files changed, 59 insertions(+), 34 deletions(-) diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index fa110f970a6..8ed5e91ef99 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -53,7 +53,9 @@ void psi_init_atomic::initialize(const Structure_Factor* sf, //< stru const K_Vectors* p_kv_in, const int& random_seed, //< random seed const pseudopot_cell_vnl* p_pspot_nl, - const int& rank) + const int& rank, + const int& npol, + const int& nbands) { ModuleBase::timer::start("psi_init_atomic", "initialize"); @@ -63,8 +65,8 @@ void psi_init_atomic::initialize(const Structure_Factor* sf, //< stru "pseudopot_cell_vnl object cannot be nullptr for atomic, quit."); } // import - psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank); - this->nbands_start_ = std::max(this->p_ucell_->natomwfc, PARAM.inp.nbands); + psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); + this->nbands_start_ = std::max(this->p_ucell_->natomwfc, nbands); this->nbands_complem_ = this->nbands_start_ - this->p_ucell_->natomwfc; // allocate this->allocate_ps_table(); diff --git a/source/source_psi/psi_init_atomic.h b/source/source_psi/psi_init_atomic.h index 4cdfabdc6cd..791f2a0a510 100644 --- a/source/source_psi/psi_init_atomic.h +++ b/source/source_psi/psi_init_atomic.h @@ -26,7 +26,9 @@ class psi_init_atomic : public psi_initializer const K_Vectors*, //< kpoints const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + const int& = 0, //< MPI rank + const int& = 1, //< npol + const int& = 1) override; //< nbands virtual void tabulate() override; virtual void init_psig(T* psig, const int& ik) override; diff --git a/source/source_psi/psi_init_atomic_random.cpp b/source/source_psi/psi_init_atomic_random.cpp index 91fd4c10bc8..021467db9de 100644 --- a/source/source_psi/psi_init_atomic_random.cpp +++ b/source/source_psi/psi_init_atomic_random.cpp @@ -9,9 +9,11 @@ void psi_init_atomic_random::initialize(const Structure_Factor* sf, / const K_Vectors* p_kv_in, const int& random_seed, //< random seed const pseudopot_cell_vnl* p_pspot_nl, - const int& rank) + const int& rank, + const int& npol, + const int& nbands) { - psi_init_atomic::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank); + psi_init_atomic::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); } template diff --git a/source/source_psi/psi_init_atomic_random.h b/source/source_psi/psi_init_atomic_random.h index 2c8e49fc8d6..91c9789d298 100644 --- a/source/source_psi/psi_init_atomic_random.h +++ b/source/source_psi/psi_init_atomic_random.h @@ -27,7 +27,9 @@ class psi_init_atomic_random : public psi_init_atomic const K_Vectors*, //< kpoints const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + const int& = 0, //< MPI rank + const int& = 1, //< npol + const int& = 1) override; //< nbands virtual void init_psig(T* psig, const int& ik) override; diff --git a/source/source_psi/psi_init_file.cpp b/source/source_psi/psi_init_file.cpp index 2e633573bb4..2354088527c 100644 --- a/source/source_psi/psi_init_file.cpp +++ b/source/source_psi/psi_init_file.cpp @@ -13,10 +13,12 @@ void psi_init_file::initialize(const Structure_Factor* sf, const K_Vectors* p_kv_in, const int& random_seed, const pseudopot_cell_vnl* p_pspot_nl, - const int& rank) + const int& rank, + const int& npol, + const int& nbands) { - psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank); - this->nbands_start_ = PARAM.inp.nbands; + psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); + this->nbands_start_ = nbands; this->nbands_complem_ = 0; } diff --git a/source/source_psi/psi_init_file.h b/source/source_psi/psi_init_file.h index 72fb18ed1eb..a9207effff1 100644 --- a/source/source_psi/psi_init_file.h +++ b/source/source_psi/psi_init_file.h @@ -27,7 +27,9 @@ class psi_init_file : public psi_initializer const K_Vectors*, //< kpoints const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + const int& = 0, //< MPI rank + const int& = 1, //< npol + const int& = 1) override; //< nbands /// @brief calculate and output planewave wavefunction /// @param ik kpoint index diff --git a/source/source_psi/psi_init_nao.cpp b/source/source_psi/psi_init_nao.cpp index 916d87295a9..5cca485f29b 100644 --- a/source/source_psi/psi_init_nao.cpp +++ b/source/source_psi/psi_init_nao.cpp @@ -153,12 +153,14 @@ void psi_init_nao::initialize(const Structure_Factor* sf, const K_Vectors* p_kv_in, const int& random_seed, const pseudopot_cell_vnl* p_pspot_nl, - const int& rank) + const int& rank, + const int& npol, + const int& nbands) { ModuleBase::timer::start("psi_init_nao", "initialize"); // import - psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank); + psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); // allocate this->allocate_ao_table(); @@ -177,18 +179,10 @@ void psi_init_nao::initialize(const Structure_Factor* sf, /* EVERY ZETA FOR (2l+1) ORBS */ const int nchi = this->p_ucell_->atoms[it].l_nchi[l]; const int degen_l = (l == 0) ? 1 : 2 * l + 1; - nbands_local += nchi * degen_l * PARAM.globalv.npol * this->p_ucell_->atoms[it].na; - /* - non-rotate basis, nbands_local*=2 for PARAM.globalv.npol = 2 is enough - */ - // nbands_local += this->p_ucell_->atoms[it].l_nchi[l]*(2*l+1) * PARAM.globalv.npol; - /* - rotate basis, nbands_local*=4 for p, d, f,... orbitals, and nbands_local*=2 for s orbitals - risky when NSPIN = 4, problematic psi value, needed to be checked - */ + nbands_local += nchi * degen_l * npol * this->p_ucell_->atoms[it].na; } } - this->nbands_start_ = std::max(nbands_local, PARAM.inp.nbands); + this->nbands_start_ = std::max(nbands_local, nbands); this->nbands_complem_ = this->nbands_start_ - nbands_local; ModuleBase::timer::end("psi_init_nao", "initialize"); diff --git a/source/source_psi/psi_init_nao.h b/source/source_psi/psi_init_nao.h index bb05962ec07..a6a60243e69 100644 --- a/source/source_psi/psi_init_nao.h +++ b/source/source_psi/psi_init_nao.h @@ -31,7 +31,9 @@ class psi_init_nao : public psi_initializer const K_Vectors*, //< kpoints const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + const int& = 0, //< MPI rank + const int& = 1, //< npol + const int& = 1) override; //< nbands void read_external_orbs(const std::string* orbital_files, const int& rank); virtual void tabulate() override; diff --git a/source/source_psi/psi_init_nao_random.cpp b/source/source_psi/psi_init_nao_random.cpp index e3e2b8e89cd..50e28235042 100644 --- a/source/source_psi/psi_init_nao_random.cpp +++ b/source/source_psi/psi_init_nao_random.cpp @@ -9,9 +9,11 @@ void psi_init_nao_random::initialize(const Structure_Factor* sf, const K_Vectors* p_kv_in, const int& random_seed, const pseudopot_cell_vnl* p_pspot_nl, - const int& rank) + const int& rank, + const int& npol, + const int& nbands) { - psi_init_nao::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank); + psi_init_nao::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); } template diff --git a/source/source_psi/psi_init_nao_random.h b/source/source_psi/psi_init_nao_random.h index d6b38bb5568..1f9608e87a1 100644 --- a/source/source_psi/psi_init_nao_random.h +++ b/source/source_psi/psi_init_nao_random.h @@ -27,7 +27,9 @@ class psi_init_nao_random : public psi_init_nao const K_Vectors*, //< kpoints const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + const int& = 0, //< MPI rank + const int& = 1, //< npol + const int& = 1) override; //< nbands virtual void init_psig(T* psig, const int& ik) override; }; diff --git a/source/source_psi/psi_init_random.cpp b/source/source_psi/psi_init_random.cpp index 697a4476adb..f6ead43ac81 100644 --- a/source/source_psi/psi_init_random.cpp +++ b/source/source_psi/psi_init_random.cpp @@ -8,13 +8,15 @@ void psi_init_random::initialize(const Structure_Factor* sf, const K_Vectors* p_kv_in, const int& random_seed, const pseudopot_cell_vnl* p_pspot_nl, - const int& rank) + const int& rank, + const int& npol, + const int& nbands) { - psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank); + psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); this->ixy2is_.clear(); this->ixy2is_.resize(this->pw_wfc_->fftnxy); this->pw_wfc_->getfftixy2is(this->ixy2is_.data()); - this->nbands_start_ = PARAM.inp.nbands; + this->nbands_start_ = nbands; this->nbands_complem_ = 0; } diff --git a/source/source_psi/psi_init_random.h b/source/source_psi/psi_init_random.h index 8d66035ec60..1018585088f 100644 --- a/source/source_psi/psi_init_random.h +++ b/source/source_psi/psi_init_random.h @@ -30,6 +30,8 @@ class psi_init_random : public psi_initializer const K_Vectors*, //< kpoints const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + const int& = 0, //< MPI rank + const int& = 1, //< npol + const int& = 1) override; //< nbands }; #endif \ No newline at end of file diff --git a/source/source_psi/psi_initializer.cpp b/source/source_psi/psi_initializer.cpp index bed67ccd1c3..ae196cf05ad 100644 --- a/source/source_psi/psi_initializer.cpp +++ b/source/source_psi/psi_initializer.cpp @@ -19,7 +19,9 @@ void psi_initializer::initialize(const Structure_Factor* sf, const K_Vectors* p_kv_in, const int& random_seed, const pseudopot_cell_vnl* p_pspot_nl, - const int& rank) + const int& rank, + const int& npol, + const int& nbands) { this->sf_ = sf; this->pw_wfc_ = pw_wfc; @@ -27,6 +29,8 @@ void psi_initializer::initialize(const Structure_Factor* sf, this->p_kv = p_kv_in; this->random_seed_ = random_seed; this->p_pspot_nl_ = p_pspot_nl; + this->npol_ = npol; + this->nbands_ = nbands; } template diff --git a/source/source_psi/psi_initializer.h b/source/source_psi/psi_initializer.h index 589355d2710..3e53447c9d9 100644 --- a/source/source_psi/psi_initializer.h +++ b/source/source_psi/psi_initializer.h @@ -62,7 +62,9 @@ class psi_initializer const K_Vectors* = nullptr, //< parallel kpoints const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0); //< rank + const int& = 0, //< rank + const int& = 1, //< npol + const int& = 1); //< nbands /// @brief CENTRAL FUNCTION: calculate the interpolate table if needed virtual void tabulate() @@ -135,5 +137,7 @@ class psi_initializer int nbands_complem_ = 0; ///< complement number of bands, which is nbands_start_ - ucell.natomwfc double mixing_coef_ = 0; ///< mixing coefficient for atomic+random and nao+random int nbands_start_ = 0; ///< starting nbands, which is no less than PARAM.inp.nbands + int npol_ = 1; ///< number of polarizations + int nbands_ = 1; ///< number of bands }; #endif \ No newline at end of file diff --git a/source/source_psi/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index adc10eeb3fa..501e3793442 100644 --- a/source/source_psi/psi_prepare.cpp +++ b/source/source_psi/psi_prepare.cpp @@ -105,7 +105,8 @@ void PSIPrepare::prepare_init(const int& random_seed) ModuleBase::WARNING_QUIT("PSIInit::prepare_init", "for new psi initializer, init_wfc type not supported"); } - this->psi_initer->initialize(&sf, &pw_wfc, &ucell, &kv, random_seed, &nlpp, rank); + this->psi_initer->initialize(&sf, &pw_wfc, &ucell, &kv, random_seed, &nlpp, rank, + PARAM.globalv.npol, PARAM.inp.nbands); this->psi_initer->tabulate(); ModuleBase::timer::end("PSIPrepare", "prepare_init"); From d9a38f9d9ab1b0b804a91e44489341f4b5eac5ca Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 13:38:37 +0800 Subject: [PATCH 02/22] remove the dependency on klist.cpp --- .../to_wannier90_lcao_in_pw.cpp | 2 +- source/source_psi/psi_init_atomic.cpp | 5 ++- source/source_psi/psi_init_atomic.h | 5 ++- source/source_psi/psi_init_atomic_random.cpp | 5 ++- source/source_psi/psi_init_atomic_random.h | 3 +- source/source_psi/psi_init_file.cpp | 15 ++++--- source/source_psi/psi_init_file.h | 4 +- source/source_psi/psi_init_nao.cpp | 5 ++- source/source_psi/psi_init_nao.h | 3 +- source/source_psi/psi_init_nao_random.cpp | 5 ++- source/source_psi/psi_init_nao_random.h | 3 +- source/source_psi/psi_init_random.cpp | 6 ++- source/source_psi/psi_init_random.h | 4 +- source/source_psi/psi_initializer.cpp | 16 ++++++-- source/source_psi/psi_initializer.h | 10 +++-- source/source_psi/psi_prepare.cpp | 7 ++-- source/source_psi/psi_prepare.h | 8 ++-- source/source_psi/setup_psi_pw.cpp | 2 +- .../test/psi_initializer_unit_test.cpp | 41 +++++++++++-------- 19 files changed, 97 insertions(+), 52 deletions(-) diff --git a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp index 48d3e7ae65c..1e047119ff0 100644 --- a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp +++ b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp @@ -45,7 +45,7 @@ void toWannier90_LCAO_IN_PW::calculate( ModulePW::PW_Basis_K* wfcpw_ptr = const_cast(wfcpw); delete this->psi_initer_; this->psi_initer_ = new psi_init_nao>(); - this->psi_initer_->initialize(sf_ptr, wfcpw_ptr, &ucell, &kv, 1, nullptr, GlobalV::MY_RANK); + this->psi_initer_->initialize(sf_ptr, wfcpw_ptr, &ucell, kv.ik2iktot, kv.get_nkstot(), 1, nullptr, GlobalV::MY_RANK); this->psi_initer_->tabulate(); delete this->psi; const int nks_psi = (PARAM.inp.calculation == "nscf" && PARAM.inp.mem_saver == 1)? 1 : wfcpw->nks; diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index 8ed5e91ef99..319d9996b64 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -50,7 +50,8 @@ template void psi_init_atomic::initialize(const Structure_Factor* sf, //< structure factor const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell - const K_Vectors* p_kv_in, + const std::vector& ik2iktot, + const int& nkstot, const int& random_seed, //< random seed const pseudopot_cell_vnl* p_pspot_nl, const int& rank, @@ -65,7 +66,7 @@ void psi_init_atomic::initialize(const Structure_Factor* sf, //< stru "pseudopot_cell_vnl object cannot be nullptr for atomic, quit."); } // import - psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); + psi_initializer::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); this->nbands_start_ = std::max(this->p_ucell_->natomwfc, nbands); this->nbands_complem_ = this->nbands_start_ - this->p_ucell_->natomwfc; // allocate diff --git a/source/source_psi/psi_init_atomic.h b/source/source_psi/psi_init_atomic.h index 791f2a0a510..b9394cef5e1 100644 --- a/source/source_psi/psi_init_atomic.h +++ b/source/source_psi/psi_init_atomic.h @@ -1,5 +1,7 @@ #ifndef PSI_INIT_ATOMIC_H #define PSI_INIT_ATOMIC_H +#include +#include #include "source_base/realarray.h" #include "psi_initializer.h" @@ -23,7 +25,8 @@ class psi_init_atomic : public psi_initializer virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell - const K_Vectors*, //< kpoints + const std::vector& = {}, //< ik2iktot: local->global k-point mapping + const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential const int& = 0, //< MPI rank diff --git a/source/source_psi/psi_init_atomic_random.cpp b/source/source_psi/psi_init_atomic_random.cpp index 021467db9de..68455dee3c1 100644 --- a/source/source_psi/psi_init_atomic_random.cpp +++ b/source/source_psi/psi_init_atomic_random.cpp @@ -6,14 +6,15 @@ template void psi_init_atomic_random::initialize(const Structure_Factor* sf, //< structure factor const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell - const K_Vectors* p_kv_in, + const std::vector& ik2iktot, + const int& nkstot, const int& random_seed, //< random seed const pseudopot_cell_vnl* p_pspot_nl, const int& rank, const int& npol, const int& nbands) { - psi_init_atomic::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); + psi_init_atomic::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); } template diff --git a/source/source_psi/psi_init_atomic_random.h b/source/source_psi/psi_init_atomic_random.h index 91c9789d298..d9623169137 100644 --- a/source/source_psi/psi_init_atomic_random.h +++ b/source/source_psi/psi_init_atomic_random.h @@ -24,7 +24,8 @@ class psi_init_atomic_random : public psi_init_atomic virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell - const K_Vectors*, //< kpoints + const std::vector& = {}, //< ik2iktot: local->global k-point mapping + const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential const int& = 0, //< MPI rank diff --git a/source/source_psi/psi_init_file.cpp b/source/source_psi/psi_init_file.cpp index 2354088527c..3a7f5924a47 100644 --- a/source/source_psi/psi_init_file.cpp +++ b/source/source_psi/psi_init_file.cpp @@ -1,7 +1,9 @@ #include "psi_init_file.h" +#include +#include + #include "source_base/timer.h" -#include "source_cell/klist.h" #include "source_io/module_wf/read_wfc_pw.h" #include "source_io/module_output/filename.h" #include "source_io/module_parameter/parameter.h" @@ -10,14 +12,15 @@ template void psi_init_file::initialize(const Structure_Factor* sf, const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, - const K_Vectors* p_kv_in, + const std::vector& ik2iktot, + const int& nkstot, const int& random_seed, const pseudopot_cell_vnl* p_pspot_nl, const int& rank, const int& npol, const int& nbands) { - psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); + psi_initializer::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); this->nbands_start_ = nbands; this->nbands_complem_ = 0; } @@ -28,9 +31,9 @@ void psi_init_file::init_psig(T* psig, const int& ik) ModuleBase::timer::start("psi_init_file", "init_psig"); const int npol = PARAM.globalv.npol; const int nbasis = this->pw_wfc_->npwk_max * npol; - const int nkstot = this->p_kv->get_nkstot(); + const int nkstot = this->nkstot_; ModuleBase::ComplexMatrix wfcatom(this->nbands_start_, nbasis); - int ik_tot = this->p_kv->ik2iktot[ik]; + int ik_tot = this->ik2iktot_[ik]; // mohan update, this is for plane wave, 2025-05-17 const int out_type = 2; @@ -39,7 +42,7 @@ void psi_init_file::init_psig(T* psig, const int& ik) const int istep = -1; std::string fn = ModuleIO::filename_output(PARAM.globalv.global_readin_dir,"wf","pw", - ik,this->p_kv->ik2iktot,PARAM.inp.nspin,nkstot, + ik,this->ik2iktot_,PARAM.inp.nspin,nkstot, out_type,out_app_flag,gamma_only,istep); ModuleIO::read_wfc_pw(fn, this->pw_wfc_, diff --git a/source/source_psi/psi_init_file.h b/source/source_psi/psi_init_file.h index a9207effff1..217a0747c28 100644 --- a/source/source_psi/psi_init_file.h +++ b/source/source_psi/psi_init_file.h @@ -1,6 +1,7 @@ #ifndef PSI_INIT_FILE_H #define PSI_INIT_FILE_H +#include #include "source_pw/module_pwdft/vnl_pw.h" #include "psi_initializer.h" @@ -24,7 +25,8 @@ class psi_init_file : public psi_initializer virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell - const K_Vectors*, //< kpoints + const std::vector& = {}, //< ik2iktot: local->global k-point mapping + const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential const int& = 0, //< MPI rank diff --git a/source/source_psi/psi_init_nao.cpp b/source/source_psi/psi_init_nao.cpp index 5cca485f29b..8788557797e 100644 --- a/source/source_psi/psi_init_nao.cpp +++ b/source/source_psi/psi_init_nao.cpp @@ -150,7 +150,8 @@ template void psi_init_nao::initialize(const Structure_Factor* sf, const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, - const K_Vectors* p_kv_in, + const std::vector& ik2iktot, + const int& nkstot, const int& random_seed, const pseudopot_cell_vnl* p_pspot_nl, const int& rank, @@ -160,7 +161,7 @@ void psi_init_nao::initialize(const Structure_Factor* sf, ModuleBase::timer::start("psi_init_nao", "initialize"); // import - psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); + psi_initializer::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); // allocate this->allocate_ao_table(); diff --git a/source/source_psi/psi_init_nao.h b/source/source_psi/psi_init_nao.h index a6a60243e69..e6e98429414 100644 --- a/source/source_psi/psi_init_nao.h +++ b/source/source_psi/psi_init_nao.h @@ -28,7 +28,8 @@ class psi_init_nao : public psi_initializer virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell - const K_Vectors*, //< kpoints + const std::vector& = {}, //< ik2iktot: local->global k-point mapping + const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential const int& = 0, //< MPI rank diff --git a/source/source_psi/psi_init_nao_random.cpp b/source/source_psi/psi_init_nao_random.cpp index 50e28235042..f305ed24f58 100644 --- a/source/source_psi/psi_init_nao_random.cpp +++ b/source/source_psi/psi_init_nao_random.cpp @@ -6,14 +6,15 @@ template void psi_init_nao_random::initialize(const Structure_Factor* sf, const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, - const K_Vectors* p_kv_in, + const std::vector& ik2iktot, + const int& nkstot, const int& random_seed, const pseudopot_cell_vnl* p_pspot_nl, const int& rank, const int& npol, const int& nbands) { - psi_init_nao::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); + psi_init_nao::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); } template diff --git a/source/source_psi/psi_init_nao_random.h b/source/source_psi/psi_init_nao_random.h index 1f9608e87a1..498142a3692 100644 --- a/source/source_psi/psi_init_nao_random.h +++ b/source/source_psi/psi_init_nao_random.h @@ -24,7 +24,8 @@ class psi_init_nao_random : public psi_init_nao virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell - const K_Vectors*, //< kpoints + const std::vector& = {}, //< ik2iktot: local->global k-point mapping + const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential const int& = 0, //< MPI rank diff --git a/source/source_psi/psi_init_random.cpp b/source/source_psi/psi_init_random.cpp index f6ead43ac81..5acd36786ac 100644 --- a/source/source_psi/psi_init_random.cpp +++ b/source/source_psi/psi_init_random.cpp @@ -1,18 +1,20 @@ #include "psi_init_random.h" +#include #include "source_io/module_parameter/parameter.h" template void psi_init_random::initialize(const Structure_Factor* sf, const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, - const K_Vectors* p_kv_in, + const std::vector& ik2iktot, + const int& nkstot, const int& random_seed, const pseudopot_cell_vnl* p_pspot_nl, const int& rank, const int& npol, const int& nbands) { - psi_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank, npol, nbands); + psi_initializer::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); this->ixy2is_.clear(); this->ixy2is_.resize(this->pw_wfc_->fftnxy); this->pw_wfc_->getfftixy2is(this->ixy2is_.data()); diff --git a/source/source_psi/psi_init_random.h b/source/source_psi/psi_init_random.h index 1018585088f..4eac5f52996 100644 --- a/source/source_psi/psi_init_random.h +++ b/source/source_psi/psi_init_random.h @@ -1,6 +1,7 @@ #ifndef PSI_INIT_RANDOM_H #define PSI_INIT_RANDOM_H +#include #include "source_pw/module_pwdft/vnl_pw.h" #include "psi_initializer.h" @@ -27,7 +28,8 @@ class psi_init_random : public psi_initializer virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell - const K_Vectors*, //< kpoints + const std::vector& = {}, //< ik2iktot: local->global k-point mapping + const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential const int& = 0, //< MPI rank diff --git a/source/source_psi/psi_initializer.cpp b/source/source_psi/psi_initializer.cpp index ae196cf05ad..094b08d8397 100644 --- a/source/source_psi/psi_initializer.cpp +++ b/source/source_psi/psi_initializer.cpp @@ -1,5 +1,13 @@ #include "psi_initializer.h" +#include +#include +#include + +#include "source_pw/module_pwdft/structure_factor.h" +#include "source_pw/module_pwdft/vnl_pw.h" +#include "source_cell/unitcell.h" +#include "source_basis/module_pw/pw_basis_k.h" #include "source_base/parallel_global.h" // basic functions support #include "source_base/timer.h" @@ -16,7 +24,8 @@ template void psi_initializer::initialize(const Structure_Factor* sf, const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, - const K_Vectors* p_kv_in, + const std::vector& ik2iktot, + const int& nkstot, const int& random_seed, const pseudopot_cell_vnl* p_pspot_nl, const int& rank, @@ -26,7 +35,8 @@ void psi_initializer::initialize(const Structure_Factor* sf, this->sf_ = sf; this->pw_wfc_ = pw_wfc; this->p_ucell_ = p_ucell; - this->p_kv = p_kv_in; + this->ik2iktot_ = ik2iktot; + this->nkstot_ = nkstot; this->random_seed_ = random_seed; this->p_pspot_nl_ = p_pspot_nl; this->npol_ = npol; @@ -48,7 +58,7 @@ void psi_initializer::random_t(T* psi, const int iw_start, const int iw_end, if (this->random_seed_ > 0) // qianrui add 2021-8-13 { #ifdef __MPI - srand(unsigned(this->random_seed_ + this->p_kv->ik2iktot[ik])); + srand(unsigned(this->random_seed_ + this->ik2iktot_[ik])); #else srand(unsigned(this->random_seed_ + ik)); #endif diff --git a/source/source_psi/psi_initializer.h b/source/source_psi/psi_initializer.h index 3e53447c9d9..29f81087524 100644 --- a/source/source_psi/psi_initializer.h +++ b/source/source_psi/psi_initializer.h @@ -12,9 +12,11 @@ #include #endif #include "source_base/macros.h" -#include "source_cell/klist.h" +#include "source_cell/unitcell.h" #include +#include +using namespace std; /* Psi (planewave based wavefunction) initializer Auther: Kirk0830 @@ -59,7 +61,8 @@ class psi_initializer virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell - const K_Vectors* = nullptr, //< parallel kpoints + const std::vector& = {}, //< ik2iktot: local->global k-point mapping + const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential const int& = 0, //< rank @@ -128,8 +131,9 @@ class psi_initializer const Structure_Factor* sf_ = nullptr; ///< Structure_Factor const ModulePW::PW_Basis_K* pw_wfc_ = nullptr; ///< use |k+G>, |G>, getgpluskcar and so on in PW_Basis_K const UnitCell* p_ucell_ = nullptr; ///< UnitCell - const K_Vectors* p_kv = nullptr; ///< Parallel_Kpoints const pseudopot_cell_vnl* p_pspot_nl_ = nullptr; ///< pseudopot_cell_vnl + std::vector ik2iktot_; ///< local->global k-point mapping + int nkstot_ = 0; ///< total number of k-points int random_seed_ = 1; ///< random seed, shared by random, atomic+random, nao+random std::vector ixy2is_; ///< used by stick_to_pool function int mem_saver_ = 0; ///< if save memory, only for nscf diff --git a/source/source_psi/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index 501e3793442..e256306585d 100644 --- a/source/source_psi/psi_prepare.cpp +++ b/source/source_psi/psi_prepare.cpp @@ -24,10 +24,11 @@ PSIPrepare::PSIPrepare(const std::string& init_wfc_in, const int& rank_in, const UnitCell& ucell_in, const Structure_Factor& sf_in, - const K_Vectors& kv_in, + const std::vector& ik2iktot_in, + const int& nkstot_in, const pseudopot_cell_vnl& nlpp_in, const ModulePW::PW_Basis_K& pw_wfc_in) - : ucell(ucell_in), sf(sf_in), nlpp(nlpp_in), kv(kv_in), pw_wfc(pw_wfc_in), rank(rank_in) + : ucell(ucell_in), sf(sf_in), nlpp(nlpp_in), pw_wfc(pw_wfc_in), rank(rank_in), ik2iktot_(ik2iktot_in), nkstot_(nkstot_in) { this->init_wfc = init_wfc_in; this->ks_solver = ks_solver_in; @@ -105,7 +106,7 @@ void PSIPrepare::prepare_init(const int& random_seed) ModuleBase::WARNING_QUIT("PSIInit::prepare_init", "for new psi initializer, init_wfc type not supported"); } - this->psi_initer->initialize(&sf, &pw_wfc, &ucell, &kv, random_seed, &nlpp, rank, + this->psi_initer->initialize(&sf, &pw_wfc, &ucell, ik2iktot_, nkstot_, random_seed, &nlpp, rank, PARAM.globalv.npol, PARAM.inp.nbands); this->psi_initer->tabulate(); diff --git a/source/source_psi/psi_prepare.h b/source/source_psi/psi_prepare.h index 4b35f545214..bad674d9317 100644 --- a/source/source_psi/psi_prepare.h +++ b/source/source_psi/psi_prepare.h @@ -18,7 +18,8 @@ class PSIPrepare : public PSIPrepareBase const int& rank, const UnitCell& ucell, const Structure_Factor& sf, - const K_Vectors& kv_in, + const std::vector& ik2iktot, + const int& nkstot, const pseudopot_cell_vnl& nlpp, const ModulePW::PW_Basis_K& pw_wfc); ~PSIPrepare(){}; @@ -65,8 +66,9 @@ class PSIPrepare : public PSIPrepareBase // pw basis const ModulePW::PW_Basis_K& pw_wfc; - // parallel kpoints - const K_Vectors& kv; + // k-point mapping and total count + const std::vector& ik2iktot_; + const int nkstot_; // unit cell const UnitCell& ucell; diff --git a/source/source_psi/setup_psi_pw.cpp b/source/source_psi/setup_psi_pw.cpp index f5bc240292d..3e75e5c87b3 100644 --- a/source/source_psi/setup_psi_pw.cpp +++ b/source/source_psi/setup_psi_pw.cpp @@ -16,7 +16,7 @@ void Setup_Psi_pw::before_runner_impl( { this->p_psi_init = new psi::PSIPrepare(inp.init_wfc, inp.ks_solver, inp.basis_type, GlobalV::MY_RANK, ucell, - sf, kv, ppcell, pw_wfc); + sf, kv.ik2iktot, kv.get_nkstot(), ppcell, pw_wfc); allocate_psi(this->psi_cpu, kv.get_nks(), kv.ngk, PARAM.globalv.nbands_l, pw_wfc.npwk_max); diff --git a/source/source_psi/test/psi_initializer_unit_test.cpp b/source/source_psi/test/psi_initializer_unit_test.cpp index 342e04a5a60..7ae463a2cb8 100644 --- a/source/source_psi/test/psi_initializer_unit_test.cpp +++ b/source/source_psi/test/psi_initializer_unit_test.cpp @@ -9,7 +9,6 @@ #include "../psi_init_nao_random.h" #include "../psi_init_random.h" #include "source_pw/module_pwdft/vl_pw.h" -#include "source_cell/klist.h" #include "source_base/output.h" /* @@ -97,7 +96,8 @@ class PsiIntializerUnitTest : public ::testing::Test { ModulePW::PW_Basis_K* p_pw_wfc = nullptr; UnitCell* p_ucell = nullptr; pseudopot_cell_vnl* p_pspot_vnl = nullptr; - K_Vectors* p_kv = nullptr; + std::vector ik2iktot_; + int nkstot_ = 0; int random_seed = 1; psi_initializer>* psi_init; @@ -111,7 +111,6 @@ class PsiIntializerUnitTest : public ::testing::Test { this->p_pw_wfc = new ModulePW::PW_Basis_K(); this->p_ucell = new UnitCell(); this->p_pspot_vnl = new pseudopot_cell_vnl(); - this->p_kv = new K_Vectors(); // mock PARAM.input.nbands = 1; PARAM.input.nspin = 1; @@ -257,8 +256,9 @@ class PsiIntializerUnitTest : public ::testing::Test { this->p_pspot_vnl->lmaxkb = 1; - this->p_kv->ik2iktot.resize(1); - this->p_kv->ik2iktot[0] = 0; + this->ik2iktot_.resize(1); + this->ik2iktot_[0] = 0; + this->nkstot_ = 1; } void TearDown() override @@ -268,7 +268,6 @@ class PsiIntializerUnitTest : public ::testing::Test { delete this->p_pw_wfc; delete this->p_ucell; delete this->p_pspot_vnl; - delete this->p_kv; } }; @@ -315,7 +314,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigRandom) { this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, - this->p_kv, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->p_pspot_vnl, GlobalV::MY_RANK); @@ -334,7 +334,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomic) { this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, - this->p_kv, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->p_pspot_vnl, GlobalV::MY_RANK); @@ -357,7 +358,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) { this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, - this->p_kv, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->p_pspot_vnl, GlobalV::MY_RANK); @@ -384,7 +386,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) { this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, - this->p_kv, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->p_pspot_vnl, GlobalV::MY_RANK); @@ -407,7 +410,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) { this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, - this->p_kv, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->p_pspot_vnl, GlobalV::MY_RANK); @@ -426,7 +430,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNao) { this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, - this->p_kv, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->p_pspot_vnl, GlobalV::MY_RANK); @@ -445,7 +450,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) { this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, - this->p_kv, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->p_pspot_vnl, GlobalV::MY_RANK); @@ -469,7 +475,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) { this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, - this->p_kv, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->p_pspot_vnl, GlobalV::MY_RANK); @@ -493,7 +500,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSo) { this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, - this->p_kv, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->p_pspot_vnl, GlobalV::MY_RANK); @@ -517,7 +525,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSoDOMAG) { this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, - this->p_kv, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->p_pspot_vnl, GlobalV::MY_RANK); From d4c5e44ea689bbedd066f865fac0a0d6aa2e88cd Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 14:46:47 +0800 Subject: [PATCH 03/22] move psi_initializer to psi_base --- source/CMakeLists.txt | 2 +- source/Makefile.Objects | 2 +- .../read_input_item_postprocess.cpp | 2 +- .../module_wannier/to_wannier90_lcao_in_pw.h | 4 +- .../source_lcao/module_rdmft/CMakeLists.txt | 2 +- source/source_lcao/module_ri/exx_lip.hpp | 2 +- source/source_psi/CMakeLists.txt | 6 +-- .../{psi_initializer.cpp => psi_base.cpp} | 20 ++++---- .../{psi_initializer.h => psi_base.h} | 26 +++++----- source/source_psi/psi_init_atomic.cpp | 2 +- source/source_psi/psi_init_atomic.h | 4 +- source/source_psi/psi_init_atomic_random.h | 2 +- source/source_psi/psi_init_file.cpp | 2 +- source/source_psi/psi_init_file.h | 6 +-- source/source_psi/psi_init_nao.cpp | 2 +- source/source_psi/psi_init_nao.h | 6 +-- source/source_psi/psi_init_random.cpp | 2 +- source/source_psi/psi_init_random.h | 4 +- source/source_psi/psi_prepare.cpp | 14 +++--- source/source_psi/psi_prepare.h | 8 ++-- source/source_psi/setup_psi_pw.h | 2 +- source/source_psi/test/CMakeLists.txt | 7 ++- ...alizer_unit_test.cpp => psi_init_test.cpp} | 48 +++++++++---------- source/source_psi/test/support/atomic_new | 11 ----- source/source_psi/test/support/nao_new | 11 ----- source/source_psi/test/support/random_new | 11 ----- 26 files changed, 87 insertions(+), 121 deletions(-) rename source/source_psi/{psi_initializer.cpp => psi_base.cpp} (93%) rename source/source_psi/{psi_initializer.h => psi_base.h} (90%) rename source/source_psi/test/{psi_initializer_unit_test.cpp => psi_init_test.cpp} (95%) delete mode 100644 source/source_psi/test/support/atomic_new delete mode 100644 source/source_psi/test/support/nao_new delete mode 100644 source/source_psi/test/support/random_new diff --git a/source/CMakeLists.txt b/source/CMakeLists.txt index 7273c6b8d72..bbc6a9025f1 100644 --- a/source/CMakeLists.txt +++ b/source/CMakeLists.txt @@ -615,7 +615,7 @@ target_link_libraries( cell parameter psi_overall_init - psi_initializer + psi_init psi dftu deltaspin diff --git a/source/Makefile.Objects b/source/Makefile.Objects index 7d8b5ee56b7..59ebda0ea7f 100644 --- a/source/Makefile.Objects +++ b/source/Makefile.Objects @@ -443,7 +443,7 @@ OBJS_ORBITAL=ORB_atomic.o\ OBJS_PSI=psi.o\ -OBJS_PSI_INITIALIZER=psi_initializer.o\ +OBJS_PSI_INITIALIZER=psi_base.o\ psi_init_random.o\ psi_init_file.o\ psi_init_atomic.o\ diff --git a/source/source_io/module_parameter/read_input_item_postprocess.cpp b/source/source_io/module_parameter/read_input_item_postprocess.cpp index aad6f47a85a..6955d840410 100644 --- a/source/source_io/module_parameter/read_input_item_postprocess.cpp +++ b/source/source_io/module_parameter/read_input_item_postprocess.cpp @@ -403,7 +403,7 @@ void ReadInput::item_postprocess() In the future lcao_in_pw will have its own ESolver. - 2023/12/22 use new psi_initializer to expand numerical + 2023/12/22 use new psi_base to expand numerical atomic orbitals, ykhuang */ if (para.input.towannier90 && para.input.basis_type == "lcao_in_pw") diff --git a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.h b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.h index a6f204858f1..bb0100c1b1b 100644 --- a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.h +++ b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.h @@ -21,7 +21,7 @@ #ifdef __LCAO #include "source_basis/module_ao/parallel_orbitals.h" -#include "source_psi/psi_initializer.h" +#include "source_psi/psi_base.h" class toWannier90_LCAO_IN_PW : public toWannier90_PW { @@ -59,7 +59,7 @@ class toWannier90_LCAO_IN_PW : public toWannier90_PW protected: const Parallel_Orbitals* ParaV = nullptr; /// @brief psi initializer for expanding nao in planewave basis - psi_initializer>* psi_initer_ = nullptr; + psi_base>* psi_initer_ = nullptr; psi::Psi, base_device::DEVICE_CPU>* psi = nullptr; diff --git a/source/source_lcao/module_rdmft/CMakeLists.txt b/source/source_lcao/module_rdmft/CMakeLists.txt index c936787ba62..9223c2d6169 100644 --- a/source/source_lcao/module_rdmft/CMakeLists.txt +++ b/source/source_lcao/module_rdmft/CMakeLists.txt @@ -11,7 +11,7 @@ endif() # if(ENABLE_COVERAGE) # add_coverage(psi) -# add_coverage(psi_initializer) +# add_coverage(psi_init) # endif() # if (BUILD_TESTING) diff --git a/source/source_lcao/module_ri/exx_lip.hpp b/source/source_lcao/module_ri/exx_lip.hpp index 5a86681c21b..120aec65164 100644 --- a/source/source_lcao/module_ri/exx_lip.hpp +++ b/source/source_lcao/module_ri/exx_lip.hpp @@ -19,7 +19,7 @@ #include "source_estate/elecstate.h" #include "source_basis/module_pw/pw_basis_k.h" #include "source_cell/module_symmetry/symmetry.h" -#include "source_psi/psi_initializer.h" +#include "source_psi/psi_base.h" #include "source_pw/module_pwdft/structure_factor.h" #include "source_base/tool_title.h" #include "source_base/timer.h" diff --git a/source/source_psi/CMakeLists.txt b/source/source_psi/CMakeLists.txt index 8be3a2ba396..6c10bbf3390 100644 --- a/source/source_psi/CMakeLists.txt +++ b/source/source_psi/CMakeLists.txt @@ -13,9 +13,9 @@ add_library( ) add_library( - psi_initializer + psi_init OBJECT - psi_initializer.cpp + psi_base.cpp psi_init_random.cpp psi_init_file.cpp psi_init_atomic.cpp @@ -27,7 +27,7 @@ add_library( if(ENABLE_COVERAGE) add_coverage(psi) - add_coverage(psi_initializer) + add_coverage(psi_init) endif() if (BUILD_TESTING) diff --git a/source/source_psi/psi_initializer.cpp b/source/source_psi/psi_base.cpp similarity index 93% rename from source/source_psi/psi_initializer.cpp rename to source/source_psi/psi_base.cpp index 094b08d8397..773e813c824 100644 --- a/source/source_psi/psi_initializer.cpp +++ b/source/source_psi/psi_base.cpp @@ -1,4 +1,4 @@ -#include "psi_initializer.h" +#include "psi_base.h" #include #include @@ -21,7 +21,7 @@ #endif template -void psi_initializer::initialize(const Structure_Factor* sf, +void psi_base::initialize(const Structure_Factor* sf, const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, const std::vector& ik2iktot, @@ -44,7 +44,7 @@ void psi_initializer::initialize(const Structure_Factor* sf, } template -void psi_initializer::random_t(T* psi, const int iw_start, const int iw_end, const int ik, const int mode) +void psi_base::random_t(T* psi, const int iw_start, const int iw_end, const int ik, const int mode) { ModuleBase::timer::start("psi_init", "random_t"); assert(mode <= 1); @@ -198,7 +198,7 @@ void psi_initializer::random_t(T* psi, const int iw_start, const int iw_end, #ifdef __MPI template -void psi_initializer::stick_to_pool(Real* stick, const int& ir, Real* out) const +void psi_base::stick_to_pool(Real* stick, const int& ir, Real* out) const { ModuleBase::timer::start("psi_init", "stick_to_pool"); MPI_Status ierror; @@ -225,7 +225,7 @@ void psi_initializer::stick_to_pool(Real* stick, const int& ir, Real* out) co } else { - ModuleBase::WARNING_QUIT("psi_initializer", "stick_to_pool: Real type not supported"); + ModuleBase::WARNING_QUIT("psi_base", "stick_to_pool: Real type not supported"); } for (int iz = 0; iz < nz; iz++) { @@ -244,7 +244,7 @@ void psi_initializer::stick_to_pool(Real* stick, const int& ir, Real* out) co } else { - ModuleBase::WARNING_QUIT("psi_initializer", "stick_to_pool: Real type not supported"); + ModuleBase::WARNING_QUIT("psi_base", "stick_to_pool: Real type not supported"); } } @@ -254,8 +254,8 @@ void psi_initializer::stick_to_pool(Real* stick, const int& ir, Real* out) co #endif // explicit instantiation -template class psi_initializer>; -template class psi_initializer>; +template class psi_base>; +template class psi_base>; // gamma point calculation -template class psi_initializer; -template class psi_initializer; +template class psi_base; +template class psi_base; diff --git a/source/source_psi/psi_initializer.h b/source/source_psi/psi_base.h similarity index 90% rename from source/source_psi/psi_initializer.h rename to source/source_psi/psi_base.h index 29f81087524..6d8df9656f1 100644 --- a/source/source_psi/psi_initializer.h +++ b/source/source_psi/psi_base.h @@ -1,5 +1,5 @@ -#ifndef PSI_INITIALIZER_H -#define PSI_INITIALIZER_H +#ifndef PSI_BASE_H +#define PSI_BASE_H // data structure support #include "source_basis/module_pw/pw_basis_k.h" // for kpoint related data structure #include "source_pw/module_pwdft/vnl_pw.h" @@ -18,7 +18,7 @@ #include using namespace std; /* -Psi (planewave based wavefunction) initializer +Psi (planewave based wavefunction) base class Auther: Kirk0830 Institute: AI for Science Institute, BEIJING @@ -26,14 +26,14 @@ This class is used to allocate memory and give initial guess for psi therefore only double datatype is needed to be supported. Following methods are available: 1. file: use wavefunction file to initialize psi - implemented in psi_initializer_file.h + implemented in psi_init_file.h 2. random: use random number to initialize psi - implemented in psi_initializer_random.h + implemented in psi_init_random.h 3. atomic: use pseudo-wavefunction in pseudopotential file to initialize psi - implemented in psi_initializer_atomic.h + implemented in psi_init_atomic.h 4. atomic+random: mix 'atomic' with some random numbers to initialize psi 5. nao: use numerical orbitals to initialize psi - implemented in psi_initializer_nao.h + implemented in psi_init_nao.h 6. nao+random: mix 'nao' with some random numbers to initialize psi To use: @@ -41,23 +41,23 @@ To use: A practical example would be in ESolver_KS_PW, because polymorphism is achieved by pointer, while a raw pointer is risky, therefore std::unique_ptr is a better choice. -1. new a std::unique_ptr with specific derived class -2. initialize() to link psi_initializer with external data and methods +1. new a std::unique_ptr with specific derived class +2. initialize() to link psi_base with external data and methods 3. tabulate() to calculate the interpolate table 4. init_psig() to calculate projection of atomic radial function onto planewave basis In summary: new->initialize->tabulate->init_psig */ template -class psi_initializer +class psi_base { private: using Real = typename GetTypeReal::type; public: - psi_initializer(){}; - virtual ~psi_initializer(){}; - /// @brief initialize the psi_initializer with external data and methods + psi_base(){}; + virtual ~psi_base(){}; + /// @brief initialize the psi_base with external data and methods virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index 319d9996b64..bb7a3e5ac71 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -66,7 +66,7 @@ void psi_init_atomic::initialize(const Structure_Factor* sf, //< stru "pseudopot_cell_vnl object cannot be nullptr for atomic, quit."); } // import - psi_initializer::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); + psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); this->nbands_start_ = std::max(this->p_ucell_->natomwfc, nbands); this->nbands_complem_ = this->nbands_start_ - this->p_ucell_->natomwfc; // allocate diff --git a/source/source_psi/psi_init_atomic.h b/source/source_psi/psi_init_atomic.h index b9394cef5e1..27f5426138c 100644 --- a/source/source_psi/psi_init_atomic.h +++ b/source/source_psi/psi_init_atomic.h @@ -3,13 +3,13 @@ #include #include #include "source_base/realarray.h" -#include "psi_initializer.h" +#include "psi_base.h" /* Psi (planewave based wavefunction) initializer: atomic */ template -class psi_init_atomic : public psi_initializer +class psi_init_atomic : public psi_base { private: using Real = typename GetTypeReal::type; diff --git a/source/source_psi/psi_init_atomic_random.h b/source/source_psi/psi_init_atomic_random.h index d9623169137..ca09c02e74d 100644 --- a/source/source_psi/psi_init_atomic_random.h +++ b/source/source_psi/psi_init_atomic_random.h @@ -20,7 +20,7 @@ class psi_init_atomic_random : public psi_init_atomic } ~psi_init_atomic_random(){}; - /// @brief initialize the psi_initializer with external data and methods + /// @brief initialize the psi_base with external data and methods virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell diff --git a/source/source_psi/psi_init_file.cpp b/source/source_psi/psi_init_file.cpp index 3a7f5924a47..1386881d49e 100644 --- a/source/source_psi/psi_init_file.cpp +++ b/source/source_psi/psi_init_file.cpp @@ -20,7 +20,7 @@ void psi_init_file::initialize(const Structure_Factor* sf, const int& npol, const int& nbands) { - psi_initializer::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); + psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); this->nbands_start_ = nbands; this->nbands_complem_ = 0; } diff --git a/source/source_psi/psi_init_file.h b/source/source_psi/psi_init_file.h index 217a0747c28..1406c2a52b0 100644 --- a/source/source_psi/psi_init_file.h +++ b/source/source_psi/psi_init_file.h @@ -3,13 +3,13 @@ #include #include "source_pw/module_pwdft/vnl_pw.h" -#include "psi_initializer.h" +#include "psi_base.h" /* Psi (planewave based wavefunction) initializer: random method */ template -class psi_init_file : public psi_initializer +class psi_init_file : public psi_base { private: using Real = typename GetTypeReal::type; @@ -21,7 +21,7 @@ class psi_init_file : public psi_initializer }; ~psi_init_file(){}; - /// @brief initialize the psi_initializer with external data and methods + /// @brief initialize the psi_base with external data and methods virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell diff --git a/source/source_psi/psi_init_nao.cpp b/source/source_psi/psi_init_nao.cpp index 8788557797e..9c02be19119 100644 --- a/source/source_psi/psi_init_nao.cpp +++ b/source/source_psi/psi_init_nao.cpp @@ -161,7 +161,7 @@ void psi_init_nao::initialize(const Structure_Factor* sf, ModuleBase::timer::start("psi_init_nao", "initialize"); // import - psi_initializer::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); + psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); // allocate this->allocate_ao_table(); diff --git a/source/source_psi/psi_init_nao.h b/source/source_psi/psi_init_nao.h index e6e98429414..a3f2bca83e3 100644 --- a/source/source_psi/psi_init_nao.h +++ b/source/source_psi/psi_init_nao.h @@ -3,14 +3,14 @@ #include "source_base/cubic_spline.h" #include "source_base/realarray.h" #include "source_base/spherical_bessel_transformer.h" -#include "psi_initializer.h" +#include "psi_base.h" #include /* Psi (planewave based wavefunction) initializer: numerical atomic orbital method */ template -class psi_init_nao : public psi_initializer +class psi_init_nao : public psi_base { private: using Real = typename GetTypeReal::type; @@ -24,7 +24,7 @@ class psi_init_nao : public psi_initializer virtual void init_psig(T* psig, const int& ik) override; - /// @brief initialize the psi_initializer with external data and methods + /// @brief initialize the psi_base with external data and methods virtual void initialize(const Structure_Factor*, //< structure factor const ModulePW::PW_Basis_K*, //< planewave basis const UnitCell*, //< unit cell diff --git a/source/source_psi/psi_init_random.cpp b/source/source_psi/psi_init_random.cpp index 5acd36786ac..8a1df41c10d 100644 --- a/source/source_psi/psi_init_random.cpp +++ b/source/source_psi/psi_init_random.cpp @@ -14,7 +14,7 @@ void psi_init_random::initialize(const Structure_Factor* sf, const int& npol, const int& nbands) { - psi_initializer::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); + psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); this->ixy2is_.clear(); this->ixy2is_.resize(this->pw_wfc_->fftnxy); this->pw_wfc_->getfftixy2is(this->ixy2is_.data()); diff --git a/source/source_psi/psi_init_random.h b/source/source_psi/psi_init_random.h index 4eac5f52996..4589741802c 100644 --- a/source/source_psi/psi_init_random.h +++ b/source/source_psi/psi_init_random.h @@ -3,13 +3,13 @@ #include #include "source_pw/module_pwdft/vnl_pw.h" -#include "psi_initializer.h" +#include "psi_base.h" /* Psi (planewave based wavefunction) initializer: random method */ template -class psi_init_random : public psi_initializer +class psi_init_random : public psi_base { private: using Real = typename GetTypeReal::type; diff --git a/source/source_psi/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index e256306585d..9aaedceffc1 100644 --- a/source/source_psi/psi_prepare.cpp +++ b/source/source_psi/psi_prepare.cpp @@ -45,12 +45,12 @@ void PSIPrepare::prepare_init(const int& random_seed) this->psi_initer.reset(); if (this->init_wfc == "random") { - this->psi_initer = std::unique_ptr>(new psi_init_random()); + this->psi_initer = std::unique_ptr>(new psi_init_random()); GlobalV::ofs_running << "\n Using RANDOM starting wave functions for all " << PARAM.inp.nbands << " bands\n"; } else if (this->init_wfc == "file") { - this->psi_initer = std::unique_ptr>(new psi_init_file()); + this->psi_initer = std::unique_ptr>(new psi_init_file()); GlobalV::ofs_running << "\n Using FILE starting wave functions\n"; } else if ((this->init_wfc.substr(0, 6) == "atomic") && (this->ucell.natomwfc == 0)) @@ -66,7 +66,7 @@ void PSIPrepare::prepare_init(const int& random_seed) " 1) A pseudopotential file that includes atomic wavefunctions (with PP_PSWFC), or\n" " 2) Numerical atomic orbitals with 'init_wfc = nao' or 'nao+random' if available.\n" << std::endl; - this->psi_initer = std::unique_ptr>(new psi_init_random()); + this->psi_initer = std::unique_ptr>(new psi_init_random()); } else if (this->init_wfc == "atomic" || (this->init_wfc == "atomic+random" && this->ucell.natomwfc < PARAM.inp.nbands)) @@ -83,22 +83,22 @@ void PSIPrepare::prepare_init(const int& random_seed) GlobalV::ofs_running << "\n Using ATOMIC starting wave functions for all " << this->ucell.natomwfc << " atomic orbitals" << " (covers " << PARAM.inp.nbands << " bands)\n"; } - this->psi_initer = std::unique_ptr>(new psi_init_atomic()); + this->psi_initer = std::unique_ptr>(new psi_init_atomic()); } else if (this->init_wfc == "atomic+random") { - this->psi_initer = std::unique_ptr>(new psi_init_atomic_random()); + this->psi_initer = std::unique_ptr>(new psi_init_atomic_random()); GlobalV::ofs_running << "\n Using ATOMIC+RANDOM starting wave functions with " << this->ucell.natomwfc << " atomic orbitals\n"; } else if (this->init_wfc == "nao") { - this->psi_initer = std::unique_ptr>(new psi_init_nao()); + this->psi_initer = std::unique_ptr>(new psi_init_nao()); GlobalV::ofs_running << "\n Using NAO starting wave functions\n"; } else if (this->init_wfc == "nao+random") { - this->psi_initer = std::unique_ptr>(new psi_init_nao_random()); + this->psi_initer = std::unique_ptr>(new psi_init_nao_random()); GlobalV::ofs_running << "\n Using NAO+RANDOM starting wave functions\n"; } else diff --git a/source/source_psi/psi_prepare.h b/source/source_psi/psi_prepare.h index bad674d9317..87003e5fbe5 100644 --- a/source/source_psi/psi_prepare.h +++ b/source/source_psi/psi_prepare.h @@ -1,7 +1,7 @@ #ifndef PSI_PREPARE_H #define PSI_PREPARE_H #include "source_hamilt/hamilt.h" -#include "source_psi/psi_initializer.h" +#include "source_psi/psi_base.h" #include "source_psi/psi_prepare_base.h" namespace psi @@ -27,7 +27,7 @@ class PSIPrepare : public PSIPrepareBase ///@brief prepare the wavefunction initialization void prepare_init(const int& random_seed); - //------------------------ only for psi_initializer -------------------- + //------------------------ only for psi_base -------------------- /** * @brief initialize the wavefunction * @@ -47,11 +47,11 @@ class PSIPrepare : public PSIPrepareBase */ void initialize_lcao_in_pw(Psi* psi_local, std::ofstream& ofs_running); - // psi_initializer* psi_initer = nullptr; + // psi_base* psi_initer = nullptr; // change to use smart pointer to manage the memory, and avoid memory leak // while the std::make_unique() is not supported till C++14, // so use the new and std::unique_ptr to manage the memory, but this makes new-delete not symmetric - std::unique_ptr> psi_initer; + std::unique_ptr> psi_initer; private: // wavefunction initialization type diff --git a/source/source_psi/setup_psi_pw.h b/source/source_psi/setup_psi_pw.h index 88e9d42bf1c..c6087f9b2b4 100644 --- a/source/source_psi/setup_psi_pw.h +++ b/source/source_psi/setup_psi_pw.h @@ -39,7 +39,7 @@ class Setup_Psi_pw // for PW, we have psi_cpu psi::Psi, base_device::DEVICE_CPU>* psi_cpu = nullptr; - // psi_initializer controller + // psi_base controller psi::PSIPrepareBase* p_psi_init = nullptr; //------------ diff --git a/source/source_psi/test/CMakeLists.txt b/source/source_psi/test/CMakeLists.txt index 63af5799e1a..aed5bea543b 100644 --- a/source/source_psi/test/CMakeLists.txt +++ b/source/source_psi/test/CMakeLists.txt @@ -8,13 +8,12 @@ AddTest( if(ENABLE_LCAO) AddTest( - TARGET MODULE_PSI_initializer_unit_test - LIBS parameter base device psi psi_initializer planewave + TARGET MODULE_PSI_init_test + LIBS parameter base device psi psi_init planewave SOURCES - psi_initializer_unit_test.cpp + psi_init_test.cpp ../../source_pw/module_pwdft/soc.cpp ../../source_cell/atom_spec.cpp - ../../source_cell/parallel_kpoints.cpp ../../source_cell/test/support/mock_unitcell.cpp ../../source_io/module_output/orb_io.cpp ../../source_io/module_output/write_pao.cpp diff --git a/source/source_psi/test/psi_initializer_unit_test.cpp b/source/source_psi/test/psi_init_test.cpp similarity index 95% rename from source/source_psi/test/psi_initializer_unit_test.cpp rename to source/source_psi/test/psi_init_test.cpp index 7ae463a2cb8..56a9a9342c8 100644 --- a/source/source_psi/test/psi_initializer_unit_test.cpp +++ b/source/source_psi/test/psi_init_test.cpp @@ -2,7 +2,7 @@ #define private public #include "source_io/module_parameter/parameter.h" #undef private -#include "../psi_initializer.h" +#include "../psi_base.h" #include "../psi_init_atomic.h" #include "../psi_init_atomic_random.h" #include "../psi_init_nao.h" @@ -13,42 +13,42 @@ /* ========================= -psi initializer unit test +psi base unit test ========================= - Tested functions: - - psi_initializer_random::psi_initializer_random - - constructor of psi_initializer_random - - psi_initializer_atomic::psi_initializer_atomic - - constructor of psi_initializer_atomic - - psi_initializer_atomic_random::psi_initializer_atomic_random - - constructor of psi_initializer_atomic_random - - psi_initializer_nao::psi_initializer_nao - - constructor of psi_initializer_nao - - psi_initializer_nao_random::psi_initializer_nao_random - - constructor of psi_initializer_nao_random - - psi_initializer::cast_to_T (psi_initializer specialized as random) + - psi_init_random::psi_init_random + - constructor of psi_init_random + - psi_init_atomic::psi_init_atomic + - constructor of psi_init_atomic + - psi_init_atomic_random::psi_init_atomic_random + - constructor of psi_init_atomic_random + - psi_init_nao::psi_init_nao + - constructor of psi_init_nao + - psi_init_nao_random::psi_init_nao_random + - constructor of psi_init_nao_random + - psi_base::cast_to_T (psi_base specialized as random) - function cast std::complex to float, double, std::complex, std::complex - - psi_initializer_random::allocate + - psi_init_random::allocate - allocate wavefunctions with random-specific method - - psi_initializer_atomic::allocate + - psi_init_atomic::allocate - allocate wavefunctions with atomic-specific method - - psi_initializer_atomic_random::allocate + - psi_init_atomic_random::allocate - allocate wavefunctions with atomic-specific method - - psi_initializer_nao::allocate + - psi_init_nao::allocate - allocate wavefunctions with nao-specific method - - psi_initializer_nao_random::allocate + - psi_init_nao_random::allocate - allocate wavefunctions with nao-specific method - - psi_initializer_random::proj_ao_onkG + - psi_init_random::proj_ao_onkG - calculate wavefunction initial guess (before diagonalization) by randomly generating numbers - - psi_initializer_atomic::proj_ao_onkG + - psi_init_atomic::proj_ao_onkG - calculate wavefunction initial guess (before diagonalization) with atomic pseudo wavefunctions - nspin = 4 case - nspin = 4 with has_so case - - psi_initializer_atomic_random::proj_ao_onkG + - psi_init_atomic_random::proj_ao_onkG - calculate wavefunction initial guess (before diagonalization) with atomic pseudo wavefunctions and random numbers - - psi_initializer_nao::proj_ao_onkG + - psi_init_nao::proj_ao_onkG - calculate wavefunction initial guess (before diagonalization) with numerical atomic orbital wavefunctions - - psi_initializer_nao_random::proj_ao_onkG + - psi_init_nao_random::proj_ao_onkG - calculate wavefunction initial guess (before diagonalization) with numerical atomic orbital wavefunctions and random numbers */ @@ -100,7 +100,7 @@ class PsiIntializerUnitTest : public ::testing::Test { int nkstot_ = 0; int random_seed = 1; - psi_initializer>* psi_init; + psi_base>* psi_init; private: protected: diff --git a/source/source_psi/test/support/atomic_new b/source/source_psi/test/support/atomic_new deleted file mode 100644 index af7a7d973a4..00000000000 --- a/source/source_psi/test/support/atomic_new +++ /dev/null @@ -1,11 +0,0 @@ -INPUT_PARAMETERS -pseudo_dir . -orbital_dir . -basis_type pw -ecutwfc 60 -scf_thr 1e-8 -scf_nmax 1 - -init_wfc atomic -psi_initializer 1 -ks_solver dav \ No newline at end of file diff --git a/source/source_psi/test/support/nao_new b/source/source_psi/test/support/nao_new deleted file mode 100644 index b0b7789a2d5..00000000000 --- a/source/source_psi/test/support/nao_new +++ /dev/null @@ -1,11 +0,0 @@ -INPUT_PARAMETERS -pseudo_dir . -orbital_dir . -basis_type pw -ecutwfc 60 -scf_thr 1e-8 -scf_nmax 1 - -init_wfc nao -psi_initializer 1 -ks_solver dav \ No newline at end of file diff --git a/source/source_psi/test/support/random_new b/source/source_psi/test/support/random_new deleted file mode 100644 index 87a70829dd9..00000000000 --- a/source/source_psi/test/support/random_new +++ /dev/null @@ -1,11 +0,0 @@ -INPUT_PARAMETERS -pseudo_dir . -orbital_dir . -basis_type pw -ecutwfc 60 -scf_thr 1e-8 -scf_nmax 1 - -init_wfc random -psi_initializer 1 -ks_solver dav \ No newline at end of file From 97703050f6bd1a31415cdee5c0844b5bb12d1585 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 15:33:47 +0800 Subject: [PATCH 04/22] remove nonlocal pseudopotential dependency --- source/source_esolver/esolver_ks_pw.cpp | 2 +- source/source_io/module_ctrl/ctrl_output_pw.h | 3 +- .../to_wannier90_lcao_in_pw.cpp | 2 +- source/source_psi/psi_base.cpp | 5 +- source/source_psi/psi_base.h | 5 +- source/source_psi/psi_init_atomic.cpp | 12 +-- source/source_psi/psi_init_atomic.h | 2 +- source/source_psi/psi_init_atomic_random.cpp | 4 +- source/source_psi/psi_init_atomic_random.h | 3 +- source/source_psi/psi_init_file.cpp | 4 +- source/source_psi/psi_init_file.h | 3 +- source/source_psi/psi_init_nao.cpp | 4 +- source/source_psi/psi_init_nao.h | 2 +- source/source_psi/psi_init_nao_random.cpp | 4 +- source/source_psi/psi_init_nao_random.h | 3 +- source/source_psi/psi_init_random.cpp | 4 +- source/source_psi/psi_init_random.h | 2 +- source/source_psi/psi_prepare.cpp | 6 +- source/source_psi/psi_prepare.h | 47 +++----- source/source_psi/setup_psi.cpp | 1 + source/source_psi/setup_psi_pw.cpp | 102 +++++++++++------- source/source_psi/setup_psi_pw.h | 12 +-- source/source_psi/test/psi_init_test.cpp | 32 +++--- 23 files changed, 126 insertions(+), 138 deletions(-) diff --git a/source/source_esolver/esolver_ks_pw.cpp b/source/source_esolver/esolver_ks_pw.cpp index f08d2aa99f3..a62748386db 100644 --- a/source/source_esolver/esolver_ks_pw.cpp +++ b/source/source_esolver/esolver_ks_pw.cpp @@ -93,7 +93,7 @@ void ESolver_KS_PW::before_all_runners(UnitCell& ucell, const Input_p this->locpp, this->ppcell, this->vsep_cell, this->pw_wfc, this->pw_rho, this->pw_rhod, this->pw_big, this->solvent, inp); - this->stp.before_runner(ucell, this->kv, this->sf, *this->pw_wfc, this->ppcell, PARAM.inp); + this->stp.before_runner(ucell, this->kv, this->sf, *this->pw_wfc, this->ppcell.lmaxkb, PARAM.inp); ModuleBase::GlobalFunc::DONE(GlobalV::ofs_running, "INIT BASIS"); diff --git a/source/source_io/module_ctrl/ctrl_output_pw.h b/source/source_io/module_ctrl/ctrl_output_pw.h index 00b7509990c..f6739b6a921 100644 --- a/source/source_io/module_ctrl/ctrl_output_pw.h +++ b/source/source_io/module_ctrl/ctrl_output_pw.h @@ -4,7 +4,8 @@ #include "source_base/module_device/device.h" // use Device #include "source_psi/psi.h" // define psi #include "source_estate/elecstate_lcao.h" // use pelec -#include "source_psi/setup_psi_pw.h" // use Setup_Psi class +#include "source_psi/setup_psi_pw.h" // use Setup_Psi class +#include "source_pw/module_pwdft/vnl_pw.h" // use pseudopot_cell_vnl namespace ModuleIO { diff --git a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp index 1e047119ff0..3344dfb6fb1 100644 --- a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp +++ b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp @@ -45,7 +45,7 @@ void toWannier90_LCAO_IN_PW::calculate( ModulePW::PW_Basis_K* wfcpw_ptr = const_cast(wfcpw); delete this->psi_initer_; this->psi_initer_ = new psi_init_nao>(); - this->psi_initer_->initialize(sf_ptr, wfcpw_ptr, &ucell, kv.ik2iktot, kv.get_nkstot(), 1, nullptr, GlobalV::MY_RANK); + this->psi_initer_->initialize(sf_ptr, wfcpw_ptr, &ucell, kv.ik2iktot, kv.get_nkstot(), 1, 0, GlobalV::MY_RANK); this->psi_initer_->tabulate(); delete this->psi; const int nks_psi = (PARAM.inp.calculation == "nscf" && PARAM.inp.mem_saver == 1)? 1 : wfcpw->nks; diff --git a/source/source_psi/psi_base.cpp b/source/source_psi/psi_base.cpp index 773e813c824..152920f06ee 100644 --- a/source/source_psi/psi_base.cpp +++ b/source/source_psi/psi_base.cpp @@ -5,7 +5,6 @@ #include #include "source_pw/module_pwdft/structure_factor.h" -#include "source_pw/module_pwdft/vnl_pw.h" #include "source_cell/unitcell.h" #include "source_basis/module_pw/pw_basis_k.h" #include "source_base/parallel_global.h" @@ -27,7 +26,7 @@ void psi_base::initialize(const Structure_Factor* sf, const std::vector& ik2iktot, const int& nkstot, const int& random_seed, - const pseudopot_cell_vnl* p_pspot_nl, + const int& lmaxkb, const int& rank, const int& npol, const int& nbands) @@ -38,7 +37,7 @@ void psi_base::initialize(const Structure_Factor* sf, this->ik2iktot_ = ik2iktot; this->nkstot_ = nkstot; this->random_seed_ = random_seed; - this->p_pspot_nl_ = p_pspot_nl; + this->lmaxkb_ = lmaxkb; this->npol_ = npol; this->nbands_ = nbands; } diff --git a/source/source_psi/psi_base.h b/source/source_psi/psi_base.h index 6d8df9656f1..5457ba04646 100644 --- a/source/source_psi/psi_base.h +++ b/source/source_psi/psi_base.h @@ -2,7 +2,6 @@ #define PSI_BASE_H // data structure support #include "source_basis/module_pw/pw_basis_k.h" // for kpoint related data structure -#include "source_pw/module_pwdft/vnl_pw.h" #include "source_pw/module_pwdft/structure_factor.h" #include "source_psi/psi.h" // for psi data structure // smart pointer for auto-memory management @@ -64,7 +63,7 @@ class psi_base const std::vector& = {}, //< ik2iktot: local->global k-point mapping const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential + const int& = 0, //< lmaxkb: max angular momentum for non-local projectors const int& = 0, //< rank const int& = 1, //< npol const int& = 1); //< nbands @@ -131,7 +130,7 @@ class psi_base const Structure_Factor* sf_ = nullptr; ///< Structure_Factor const ModulePW::PW_Basis_K* pw_wfc_ = nullptr; ///< use |k+G>, |G>, getgpluskcar and so on in PW_Basis_K const UnitCell* p_ucell_ = nullptr; ///< UnitCell - const pseudopot_cell_vnl* p_pspot_nl_ = nullptr; ///< pseudopot_cell_vnl + int lmaxkb_ = 0; ///< max angular momentum for non-local projectors std::vector ik2iktot_; ///< local->global k-point mapping int nkstot_ = 0; ///< total number of k-points int random_seed_ = 1; ///< random seed, shared by random, atomic+random, nao+random diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index bb7a3e5ac71..06e29b3c11e 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -53,20 +53,14 @@ void psi_init_atomic::initialize(const Structure_Factor* sf, //< stru const std::vector& ik2iktot, const int& nkstot, const int& random_seed, //< random seed - const pseudopot_cell_vnl* p_pspot_nl, + const int& lmaxkb, const int& rank, const int& npol, const int& nbands) { ModuleBase::timer::start("psi_init_atomic", "initialize"); - if(p_pspot_nl == nullptr) - { - ModuleBase::WARNING_QUIT("psi_init_atomic::initialize", - "pseudopot_cell_vnl object cannot be nullptr for atomic, quit."); - } - // import - psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); + psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); this->nbands_start_ = std::max(this->p_ucell_->natomwfc, nbands); this->nbands_complem_ = this->nbands_start_ - this->p_ucell_->natomwfc; // allocate @@ -300,7 +294,7 @@ void psi_init_atomic::init_psig(T* psig, const int& ik) if(fabs(cg_coeffs[is]) > 1e-8) { /* GET COMPLEX SPHERICAL HARMONIC FUNCTION */ - const int ind = this->p_pspot_nl_->lmaxkb + soc.sph_ind(l,j,m,is); // ind can be l+m, l+m+1, l+m-1 + const int ind = this->lmaxkb_ + soc.sph_ind(l,j,m,is); // ind can be l+m, l+m+1, l+m-1 std::fill(aux.begin(), aux.end(), std::complex(0.0, 0.0)); for(int n1 = 0; n1 < 2*l+1; n1++) { diff --git a/source/source_psi/psi_init_atomic.h b/source/source_psi/psi_init_atomic.h index 27f5426138c..90a91fd8e19 100644 --- a/source/source_psi/psi_init_atomic.h +++ b/source/source_psi/psi_init_atomic.h @@ -28,7 +28,7 @@ class psi_init_atomic : public psi_base const std::vector& = {}, //< ik2iktot: local->global k-point mapping const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential + const int& = 0, //< lmaxkb: max angular momentum for non-local projectors const int& = 0, //< MPI rank const int& = 1, //< npol const int& = 1) override; //< nbands diff --git a/source/source_psi/psi_init_atomic_random.cpp b/source/source_psi/psi_init_atomic_random.cpp index 68455dee3c1..b1f358db474 100644 --- a/source/source_psi/psi_init_atomic_random.cpp +++ b/source/source_psi/psi_init_atomic_random.cpp @@ -9,12 +9,12 @@ void psi_init_atomic_random::initialize(const Structure_Factor* sf, / const std::vector& ik2iktot, const int& nkstot, const int& random_seed, //< random seed - const pseudopot_cell_vnl* p_pspot_nl, + const int& lmaxkb, const int& rank, const int& npol, const int& nbands) { - psi_init_atomic::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); + psi_init_atomic::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); } template diff --git a/source/source_psi/psi_init_atomic_random.h b/source/source_psi/psi_init_atomic_random.h index ca09c02e74d..17ff511ad8b 100644 --- a/source/source_psi/psi_init_atomic_random.h +++ b/source/source_psi/psi_init_atomic_random.h @@ -1,6 +1,5 @@ #ifndef PSI_INIT_ATOMIC_RANDOM_H #define PSI_INIT_ATOMIC_RANDOM_H -#include "source_pw/module_pwdft/vnl_pw.h" #include "psi_init_atomic.h" /* @@ -27,7 +26,7 @@ class psi_init_atomic_random : public psi_init_atomic const std::vector& = {}, //< ik2iktot: local->global k-point mapping const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential + const int& = 0, //< lmaxkb: max angular momentum for non-local projectors const int& = 0, //< MPI rank const int& = 1, //< npol const int& = 1) override; //< nbands diff --git a/source/source_psi/psi_init_file.cpp b/source/source_psi/psi_init_file.cpp index 1386881d49e..e97eed23260 100644 --- a/source/source_psi/psi_init_file.cpp +++ b/source/source_psi/psi_init_file.cpp @@ -15,12 +15,12 @@ void psi_init_file::initialize(const Structure_Factor* sf, const std::vector& ik2iktot, const int& nkstot, const int& random_seed, - const pseudopot_cell_vnl* p_pspot_nl, + const int& lmaxkb, const int& rank, const int& npol, const int& nbands) { - psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); + psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); this->nbands_start_ = nbands; this->nbands_complem_ = 0; } diff --git a/source/source_psi/psi_init_file.h b/source/source_psi/psi_init_file.h index 1406c2a52b0..9bca8eebf2c 100644 --- a/source/source_psi/psi_init_file.h +++ b/source/source_psi/psi_init_file.h @@ -2,7 +2,6 @@ #define PSI_INIT_FILE_H #include -#include "source_pw/module_pwdft/vnl_pw.h" #include "psi_base.h" /* @@ -28,7 +27,7 @@ class psi_init_file : public psi_base const std::vector& = {}, //< ik2iktot: local->global k-point mapping const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential + const int& = 0, //< lmaxkb: max angular momentum for non-local projectors const int& = 0, //< MPI rank const int& = 1, //< npol const int& = 1) override; //< nbands diff --git a/source/source_psi/psi_init_nao.cpp b/source/source_psi/psi_init_nao.cpp index 9c02be19119..6ca0e7e7184 100644 --- a/source/source_psi/psi_init_nao.cpp +++ b/source/source_psi/psi_init_nao.cpp @@ -153,7 +153,7 @@ void psi_init_nao::initialize(const Structure_Factor* sf, const std::vector& ik2iktot, const int& nkstot, const int& random_seed, - const pseudopot_cell_vnl* p_pspot_nl, + const int& lmaxkb, const int& rank, const int& npol, const int& nbands) @@ -161,7 +161,7 @@ void psi_init_nao::initialize(const Structure_Factor* sf, ModuleBase::timer::start("psi_init_nao", "initialize"); // import - psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); + psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); // allocate this->allocate_ao_table(); diff --git a/source/source_psi/psi_init_nao.h b/source/source_psi/psi_init_nao.h index a3f2bca83e3..cbd2c3c2ac1 100644 --- a/source/source_psi/psi_init_nao.h +++ b/source/source_psi/psi_init_nao.h @@ -31,7 +31,7 @@ class psi_init_nao : public psi_base const std::vector& = {}, //< ik2iktot: local->global k-point mapping const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential + const int& = 0, //< lmaxkb: max angular momentum for non-local projectors const int& = 0, //< MPI rank const int& = 1, //< npol const int& = 1) override; //< nbands diff --git a/source/source_psi/psi_init_nao_random.cpp b/source/source_psi/psi_init_nao_random.cpp index f305ed24f58..c7bd0e3bcac 100644 --- a/source/source_psi/psi_init_nao_random.cpp +++ b/source/source_psi/psi_init_nao_random.cpp @@ -9,12 +9,12 @@ void psi_init_nao_random::initialize(const Structure_Factor* sf, const std::vector& ik2iktot, const int& nkstot, const int& random_seed, - const pseudopot_cell_vnl* p_pspot_nl, + const int& lmaxkb, const int& rank, const int& npol, const int& nbands) { - psi_init_nao::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); + psi_init_nao::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); } template diff --git a/source/source_psi/psi_init_nao_random.h b/source/source_psi/psi_init_nao_random.h index 498142a3692..d91f27d598d 100644 --- a/source/source_psi/psi_init_nao_random.h +++ b/source/source_psi/psi_init_nao_random.h @@ -1,6 +1,5 @@ #ifndef PSI_INIT_NAO_RANDOM_H #define PSI_INIT_NAO_RANDOM_H -#include "source_pw/module_pwdft/vnl_pw.h" #include "psi_init_nao.h" /* @@ -27,7 +26,7 @@ class psi_init_nao_random : public psi_init_nao const std::vector& = {}, //< ik2iktot: local->global k-point mapping const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential + const int& = 0, //< lmaxkb: max angular momentum for non-local projectors const int& = 0, //< MPI rank const int& = 1, //< npol const int& = 1) override; //< nbands diff --git a/source/source_psi/psi_init_random.cpp b/source/source_psi/psi_init_random.cpp index 8a1df41c10d..2ca461ffc13 100644 --- a/source/source_psi/psi_init_random.cpp +++ b/source/source_psi/psi_init_random.cpp @@ -9,12 +9,12 @@ void psi_init_random::initialize(const Structure_Factor* sf, const std::vector& ik2iktot, const int& nkstot, const int& random_seed, - const pseudopot_cell_vnl* p_pspot_nl, + const int& lmaxkb, const int& rank, const int& npol, const int& nbands) { - psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, p_pspot_nl, rank, npol, nbands); + psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); this->ixy2is_.clear(); this->ixy2is_.resize(this->pw_wfc_->fftnxy); this->pw_wfc_->getfftixy2is(this->ixy2is_.data()); diff --git a/source/source_psi/psi_init_random.h b/source/source_psi/psi_init_random.h index 4589741802c..a830edac250 100644 --- a/source/source_psi/psi_init_random.h +++ b/source/source_psi/psi_init_random.h @@ -31,7 +31,7 @@ class psi_init_random : public psi_base const std::vector& = {}, //< ik2iktot: local->global k-point mapping const int& = 0, //< nkstot: total number of k-points const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential + const int& = 0, //< lmaxkb: max angular momentum for non-local projectors const int& = 0, //< MPI rank const int& = 1, //< npol const int& = 1) override; //< nbands diff --git a/source/source_psi/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index 9aaedceffc1..bf33edabf17 100644 --- a/source/source_psi/psi_prepare.cpp +++ b/source/source_psi/psi_prepare.cpp @@ -26,9 +26,9 @@ PSIPrepare::PSIPrepare(const std::string& init_wfc_in, const Structure_Factor& sf_in, const std::vector& ik2iktot_in, const int& nkstot_in, - const pseudopot_cell_vnl& nlpp_in, + const int& lmaxkb_in, const ModulePW::PW_Basis_K& pw_wfc_in) - : ucell(ucell_in), sf(sf_in), nlpp(nlpp_in), pw_wfc(pw_wfc_in), rank(rank_in), ik2iktot_(ik2iktot_in), nkstot_(nkstot_in) + : ucell(ucell_in), sf(sf_in), lmaxkb(lmaxkb_in), pw_wfc(pw_wfc_in), rank(rank_in), ik2iktot_(ik2iktot_in), nkstot_(nkstot_in) { this->init_wfc = init_wfc_in; this->ks_solver = ks_solver_in; @@ -106,7 +106,7 @@ void PSIPrepare::prepare_init(const int& random_seed) ModuleBase::WARNING_QUIT("PSIInit::prepare_init", "for new psi initializer, init_wfc type not supported"); } - this->psi_initer->initialize(&sf, &pw_wfc, &ucell, ik2iktot_, nkstot_, random_seed, &nlpp, rank, + this->psi_initer->initialize(&sf, &pw_wfc, &ucell, ik2iktot_, nkstot_, random_seed, lmaxkb, rank, PARAM.globalv.npol, PARAM.inp.nbands); this->psi_initer->tabulate(); diff --git a/source/source_psi/psi_prepare.h b/source/source_psi/psi_prepare.h index 87003e5fbe5..eba177b90f3 100644 --- a/source/source_psi/psi_prepare.h +++ b/source/source_psi/psi_prepare.h @@ -4,10 +4,13 @@ #include "source_psi/psi_base.h" #include "source_psi/psi_prepare_base.h" +class UnitCell; +class Structure_Factor; +namespace ModulePW { class PW_Basis_K; } + namespace psi { -// This class is used to prepare the wavefunction template class PSIPrepare : public PSIPrepareBase { @@ -20,76 +23,50 @@ class PSIPrepare : public PSIPrepareBase const Structure_Factor& sf, const std::vector& ik2iktot, const int& nkstot, - const pseudopot_cell_vnl& nlpp, + const int& lmaxkb, const ModulePW::PW_Basis_K& pw_wfc); + ~PSIPrepare(){}; - ///@brief prepare the wavefunction initialization void prepare_init(const int& random_seed); - //------------------------ only for psi_base -------------------- - /** - * @brief initialize the wavefunction - * - * @param psi store the wavefunction - * @param p_hamilt Hamiltonian operator - * @param ofs_running output stream for running information - * @param is_already_initpsi whether psi has been initialized - */ void initialize_psi(Psi>* psi, psi::Psi* kspw_psi, hamilt::Hamilt* p_hamilt, std::ofstream& ofs_running); - /** - * @brief initialize NAOs in plane wave basis, only for LCAO_IN_PW - * - */ void initialize_lcao_in_pw(Psi* psi_local, std::ofstream& ofs_running); - // psi_base* psi_initer = nullptr; - // change to use smart pointer to manage the memory, and avoid memory leak - // while the std::make_unique() is not supported till C++14, - // so use the new and std::unique_ptr to manage the memory, but this makes new-delete not symmetric std::unique_ptr> psi_initer; private: - // wavefunction initialization type + std::string init_wfc = "none"; - // Kohn-Sham solver type std::string ks_solver = "none"; - // basis type std::string basis_type = "none"; - // pw basis const ModulePW::PW_Basis_K& pw_wfc; - // k-point mapping and total count const std::vector& ik2iktot_; const int nkstot_; - // unit cell const UnitCell& ucell; - // structure factor const Structure_Factor& sf; - // nonlocal pseudopotential - const pseudopot_cell_vnl& nlpp; + const int lmaxkb; - Device* ctx = {}; ///< device - base_device::DEVICE_CPU* cpu_ctx = {}; ///< CPU device - const int rank; ///< MPI rank + Device* ctx = {}; + base_device::DEVICE_CPU* cpu_ctx = {}; + const int rank; - //-------------------------OP-------------------------------------------- using syncmem_complex_op = base_device::memory::synchronize_memory_op; using syncmem_h2d_op = base_device::memory::synchronize_memory_op; }; -///@brief allocate the wavefunction void allocate_psi(Psi>*& psi, const int& nks, const std::vector& ngk, const int& nbands, const int& npwx); } // namespace psi -#endif \ No newline at end of file +#endif diff --git a/source/source_psi/setup_psi.cpp b/source/source_psi/setup_psi.cpp index ba658a02de4..c5ca08caaea 100644 --- a/source/source_psi/setup_psi.cpp +++ b/source/source_psi/setup_psi.cpp @@ -1,4 +1,5 @@ #include "source_psi/setup_psi.h" +#include "source_cell/klist.h" #include "source_io/module_parameter/parameter.h" // use parameter template diff --git a/source/source_psi/setup_psi_pw.cpp b/source/source_psi/setup_psi_pw.cpp index 3e75e5c87b3..16be8ee790a 100644 --- a/source/source_psi/setup_psi_pw.cpp +++ b/source/source_psi/setup_psi_pw.cpp @@ -1,5 +1,9 @@ #include "source_psi/setup_psi_pw.h" -#include "source_io/module_parameter/parameter.h" // use parameter +#include "source_cell/klist.h" +#include "source_cell/unitcell.h" +#include "source_pw/module_pwdft/structure_factor.h" +#include "source_basis/module_pw/pw_basis_k.h" +#include "source_io/module_parameter/parameter.h" Setup_Psi_pw::Setup_Psi_pw(){} @@ -11,37 +15,50 @@ void Setup_Psi_pw::before_runner_impl( const K_Vectors &kv, const Structure_Factor &sf, const ModulePW::PW_Basis_K &pw_wfc, - const pseudopot_cell_vnl &ppcell, + const int &lmaxkb, const Input_para &inp) { this->p_psi_init = new psi::PSIPrepare(inp.init_wfc, inp.ks_solver, inp.basis_type, GlobalV::MY_RANK, ucell, - sf, kv.ik2iktot, kv.get_nkstot(), ppcell, pw_wfc); + sf, kv.ik2iktot, kv.get_nkstot(), lmaxkb, pw_wfc); allocate_psi(this->psi_cpu, kv.get_nks(), kv.ngk, PARAM.globalv.nbands_l, pw_wfc.npwk_max); auto* p_psi_init = static_cast*>(this->p_psi_init); p_psi_init->prepare_init(inp.pw_seed); - if (std::is_same::value) { + if (std::is_same::value) + { precision_type_ = PrecisionType::Float; - } else if (std::is_same::value) { + } + else if (std::is_same::value) + { precision_type_ = PrecisionType::Double; - } else if (std::is_same>::value) { + } + else if (std::is_same>::value) + { precision_type_ = PrecisionType::ComplexFloat; - } else { + } + else + { precision_type_ = PrecisionType::ComplexDouble; } - if (std::is_same::value) { + if (std::is_same::value) + { device_type_ = base_device::GpuDevice; - } else { + } + else + { device_type_ = base_device::CpuDevice; } - if (inp.device == "gpu" || inp.precision == "single") { + if (inp.device == "gpu" || inp.precision == "single") + { this->psi_t = static_cast(new psi::Psi(this->psi_cpu[0])); - } else { + } + else + { this->psi_t = static_cast(reinterpret_cast*>(this->psi_cpu)); } } @@ -51,30 +68,37 @@ void Setup_Psi_pw::before_runner( const K_Vectors &kv, const Structure_Factor &sf, const ModulePW::PW_Basis_K &pw_wfc, - const pseudopot_cell_vnl &ppcell, + const int &lmaxkb, const Input_para &inp) { const bool is_gpu = (inp.device == "gpu"); const bool is_single = (inp.precision == "single"); #if ((defined __CUDA) || (defined __ROCM)) - if (is_gpu) { - if (is_single) { + if (is_gpu) + { + if (is_single) + { before_runner_impl, base_device::DEVICE_GPU>( - ucell, kv, sf, pw_wfc, ppcell, inp); - } else { + ucell, kv, sf, pw_wfc, lmaxkb, inp); + } + else + { before_runner_impl, base_device::DEVICE_GPU>( - ucell, kv, sf, pw_wfc, ppcell, inp); + ucell, kv, sf, pw_wfc, lmaxkb, inp); } } else #endif { - if (is_single) { + if (is_single) + { before_runner_impl, base_device::DEVICE_CPU>( - ucell, kv, sf, pw_wfc, ppcell, inp); - } else { + ucell, kv, sf, pw_wfc, lmaxkb, inp); + } + else + { before_runner_impl, base_device::DEVICE_CPU>( - ucell, kv, sf, pw_wfc, ppcell, inp); + ucell, kv, sf, pw_wfc, lmaxkb, inp); } } } @@ -88,10 +112,12 @@ void Setup_Psi_pw::update_psi_d_impl() delete this->get_psi_d(); } - // Refresh this->psi_d - if (this->precision_type_ == PrecisionType::ComplexFloat) { + if (this->precision_type_ == PrecisionType::ComplexFloat) + { this->psi_d = static_cast(new psi::Psi, Device>(*this->get_psi_t())); - } else { + } + else + { this->psi_d = static_cast(reinterpret_cast, Device>*>(this->psi_t)); } } @@ -203,13 +229,13 @@ void Setup_Psi_pw::copy_d2h() } template -void Setup_Psi_pw::castmem_d2h_impl(std::complex* dst, const std::complex* src, const size_t size) +void Setup_Psi_pw::castmem_d2h_impl(std::complex* dst, const std::complex* src, const std::size_t size) { base_device::memory::cast_memory_op, std::complex, base_device::DEVICE_CPU, Device>()(dst, src, size); } template -void Setup_Psi_pw::castmem_d2h_impl(std::complex* dst, const std::complex* src, const size_t size) +void Setup_Psi_pw::castmem_d2h_impl(std::complex* dst, const std::complex* src, const std::size_t size) { base_device::memory::cast_memory_op, std::complex, base_device::DEVICE_CPU, Device>()(dst, src, size); } @@ -263,11 +289,11 @@ template class psi::PSIPrepare, base_device::DEVICE_CPU>; template void Setup_Psi_pw::before_runner_impl, base_device::DEVICE_CPU>( const UnitCell&, const K_Vectors&, const Structure_Factor&, - const ModulePW::PW_Basis_K&, const pseudopot_cell_vnl&, const Input_para&); + const ModulePW::PW_Basis_K&, const int&, const Input_para&); template void Setup_Psi_pw::before_runner_impl, base_device::DEVICE_CPU>( const UnitCell&, const K_Vectors&, const Structure_Factor&, - const ModulePW::PW_Basis_K&, const pseudopot_cell_vnl&, const Input_para&); + const ModulePW::PW_Basis_K&, const int&, const Input_para&); template void Setup_Psi_pw::init_impl, base_device::DEVICE_CPU>( hamilt::Hamilt, base_device::DEVICE_CPU>*); @@ -284,16 +310,16 @@ template void Setup_Psi_pw::clean_impl, base_device::DEVICE_ template void Setup_Psi_pw::clean_impl, base_device::DEVICE_CPU>(); template void Setup_Psi_pw::castmem_d2h_impl, base_device::DEVICE_CPU>( - std::complex*, const std::complex*, const size_t); + std::complex*, const std::complex*, const std::size_t); template void Setup_Psi_pw::castmem_d2h_impl, base_device::DEVICE_CPU>( - std::complex*, const std::complex*, const size_t); + std::complex*, const std::complex*, const std::size_t); template void Setup_Psi_pw::castmem_d2h_impl, base_device::DEVICE_CPU>( - std::complex*, const std::complex*, const size_t); + std::complex*, const std::complex*, const std::size_t); template void Setup_Psi_pw::castmem_d2h_impl, base_device::DEVICE_CPU>( - std::complex*, const std::complex*, const size_t); + std::complex*, const std::complex*, const std::size_t); #if ((defined __CUDA) || (defined __ROCM)) template class psi::PSIPrepare, base_device::DEVICE_GPU>; @@ -301,11 +327,11 @@ template class psi::PSIPrepare, base_device::DEVICE_GPU>; template void Setup_Psi_pw::before_runner_impl, base_device::DEVICE_GPU>( const UnitCell&, const K_Vectors&, const Structure_Factor&, - const ModulePW::PW_Basis_K&, const pseudopot_cell_vnl&, const Input_para&); + const ModulePW::PW_Basis_K&, const int&, const Input_para&); template void Setup_Psi_pw::before_runner_impl, base_device::DEVICE_GPU>( const UnitCell&, const K_Vectors&, const Structure_Factor&, - const ModulePW::PW_Basis_K&, const pseudopot_cell_vnl&, const Input_para&); + const ModulePW::PW_Basis_K&, const int&, const Input_para&); template void Setup_Psi_pw::init_impl, base_device::DEVICE_GPU>( hamilt::Hamilt, base_device::DEVICE_GPU>*); @@ -326,14 +352,14 @@ template void Setup_Psi_pw::clean_impl, base_device::DEVICE_ template void Setup_Psi_pw::clean_impl, base_device::DEVICE_GPU>(); template void Setup_Psi_pw::castmem_d2h_impl, base_device::DEVICE_GPU>( - std::complex*, const std::complex*, const size_t); + std::complex*, const std::complex*, const std::size_t); template void Setup_Psi_pw::castmem_d2h_impl, base_device::DEVICE_GPU>( - std::complex*, const std::complex*, const size_t); + std::complex*, const std::complex*, const std::size_t); template void Setup_Psi_pw::castmem_d2h_impl, base_device::DEVICE_GPU>( - std::complex*, const std::complex*, const size_t); + std::complex*, const std::complex*, const std::size_t); template void Setup_Psi_pw::castmem_d2h_impl, base_device::DEVICE_GPU>( - std::complex*, const std::complex*, const size_t); + std::complex*, const std::complex*, const std::size_t); #endif diff --git a/source/source_psi/setup_psi_pw.h b/source/source_psi/setup_psi_pw.h index c6087f9b2b4..1d05d804f65 100644 --- a/source/source_psi/setup_psi_pw.h +++ b/source/source_psi/setup_psi_pw.h @@ -6,10 +6,10 @@ #include "source_cell/klist.h" #include "source_pw/module_pwdft/structure_factor.h" #include "source_basis/module_pw/pw_basis_k.h" -#include "source_pw/module_pwdft/vnl_pw.h" #include "source_io/module_parameter/input_parameter.h" #include "source_base/module_device/device.h" #include "source_hamilt/hamilt.h" +#include "source_psi/psi_base.h" class Setup_Psi_pw { @@ -51,7 +51,7 @@ class Setup_Psi_pw const K_Vectors &kv, const Structure_Factor &sf, const ModulePW::PW_Basis_K &pw_wfc, - const pseudopot_cell_vnl &ppcell, + const int &lmaxkb, const Input_para &inp); void init(hamilt::HamiltBase* p_hamilt); @@ -71,7 +71,7 @@ class Setup_Psi_pw int get_nbands() const { return this->psi_cpu->get_nbands(); } int get_nk() const { return this->psi_cpu->get_nk(); } int get_nbasis() const { return this->psi_cpu->get_nbasis(); } - size_t size() const { return this->psi_cpu->size(); } + std::size_t size() const { return this->psi_cpu->size(); } // Get runtime type information base_device::AbacusDevice_t get_device_type() const { return device_type_; } @@ -126,7 +126,7 @@ class Setup_Psi_pw const K_Vectors &kv, const Structure_Factor &sf, const ModulePW::PW_Basis_K &pw_wfc, - const pseudopot_cell_vnl &ppcell, + const int &lmaxkb, const Input_para &inp); template @@ -142,10 +142,10 @@ class Setup_Psi_pw void copy_d2h_impl(); template - void castmem_d2h_impl(std::complex* dst, const std::complex* src, const size_t size); + void castmem_d2h_impl(std::complex* dst, const std::complex* src, const std::size_t size); template - void castmem_d2h_impl(std::complex* dst, const std::complex* src, const size_t size); + void castmem_d2h_impl(std::complex* dst, const std::complex* src, const std::size_t size); }; diff --git a/source/source_psi/test/psi_init_test.cpp b/source/source_psi/test/psi_init_test.cpp index 56a9a9342c8..bf2a18e8f5f 100644 --- a/source/source_psi/test/psi_init_test.cpp +++ b/source/source_psi/test/psi_init_test.cpp @@ -63,10 +63,6 @@ void Atom_pseudo::bcast_atom_pseudo() {} pseudo::pseudo() {} pseudo::~pseudo() {} -pseudopot_cell_vnl::pseudopot_cell_vnl() {} -pseudopot_cell_vnl::~pseudopot_cell_vnl() -{ -} pseudopot_cell_vl::pseudopot_cell_vl() {} pseudopot_cell_vl::~pseudopot_cell_vl() {} Magnetism::Magnetism() {} @@ -95,7 +91,7 @@ class PsiIntializerUnitTest : public ::testing::Test { Structure_Factor* p_sf = nullptr; ModulePW::PW_Basis_K* p_pw_wfc = nullptr; UnitCell* p_ucell = nullptr; - pseudopot_cell_vnl* p_pspot_vnl = nullptr; + int lmaxkb = 0; std::vector ik2iktot_; int nkstot_ = 0; int random_seed = 1; @@ -110,7 +106,6 @@ class PsiIntializerUnitTest : public ::testing::Test { this->p_sf = new Structure_Factor(); this->p_pw_wfc = new ModulePW::PW_Basis_K(); this->p_ucell = new UnitCell(); - this->p_pspot_vnl = new pseudopot_cell_vnl(); // mock PARAM.input.nbands = 1; PARAM.input.nspin = 1; @@ -254,7 +249,7 @@ class PsiIntializerUnitTest : public ::testing::Test { this->p_pw_wfc->kvec_d = new ModuleBase::Vector3[1]; this->p_pw_wfc->kvec_d[0] = {0.0, 0.0, 0.0}; - this->p_pspot_vnl->lmaxkb = 1; + this->lmaxkb = 1; this->ik2iktot_.resize(1); this->ik2iktot_[0] = 0; @@ -267,7 +262,6 @@ class PsiIntializerUnitTest : public ::testing::Test { delete this->p_sf; delete this->p_pw_wfc; delete this->p_ucell; - delete this->p_pspot_vnl; } }; @@ -317,7 +311,7 @@ TEST_F(PsiIntializerUnitTest, CalPsigRandom) { this->ik2iktot_, this->nkstot_, this->random_seed, - this->p_pspot_vnl, + this->lmaxkb, GlobalV::MY_RANK); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); @@ -337,7 +331,7 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomic) { this->ik2iktot_, this->nkstot_, this->random_seed, - this->p_pspot_vnl, + this->lmaxkb, GlobalV::MY_RANK); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); @@ -361,7 +355,7 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) { this->ik2iktot_, this->nkstot_, this->random_seed, - this->p_pspot_vnl, + this->lmaxkb, GlobalV::MY_RANK); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); @@ -389,7 +383,7 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) { this->ik2iktot_, this->nkstot_, this->random_seed, - this->p_pspot_vnl, + this->lmaxkb, GlobalV::MY_RANK); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); @@ -413,7 +407,7 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) { this->ik2iktot_, this->nkstot_, this->random_seed, - this->p_pspot_vnl, + this->lmaxkb, GlobalV::MY_RANK); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); @@ -433,7 +427,7 @@ TEST_F(PsiIntializerUnitTest, CalPsigNao) { this->ik2iktot_, this->nkstot_, this->random_seed, - this->p_pspot_vnl, + this->lmaxkb, GlobalV::MY_RANK); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); @@ -453,7 +447,7 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) { this->ik2iktot_, this->nkstot_, this->random_seed, - this->p_pspot_vnl, + this->lmaxkb, GlobalV::MY_RANK); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); @@ -478,7 +472,7 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) { this->ik2iktot_, this->nkstot_, this->random_seed, - this->p_pspot_vnl, + this->lmaxkb, GlobalV::MY_RANK); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); @@ -503,7 +497,7 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSo) { this->ik2iktot_, this->nkstot_, this->random_seed, - this->p_pspot_vnl, + this->lmaxkb, GlobalV::MY_RANK); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); @@ -528,7 +522,7 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSoDOMAG) { this->ik2iktot_, this->nkstot_, this->random_seed, - this->p_pspot_vnl, + this->lmaxkb, GlobalV::MY_RANK); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); @@ -556,4 +550,4 @@ int main(int argc, char** argv) #endif return result; -} \ No newline at end of file +} From 5ffbad032f5070a97e5ec32a3fffdb81179d264e Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 15:36:48 +0800 Subject: [PATCH 05/22] update --- source/source_io/module_ctrl/ctrl_output_pw.h | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/source/source_io/module_ctrl/ctrl_output_pw.h b/source/source_io/module_ctrl/ctrl_output_pw.h index f6739b6a921..ffed81dfe33 100644 --- a/source/source_io/module_ctrl/ctrl_output_pw.h +++ b/source/source_io/module_ctrl/ctrl_output_pw.h @@ -5,7 +5,8 @@ #include "source_psi/psi.h" // define psi #include "source_estate/elecstate_lcao.h" // use pelec #include "source_psi/setup_psi_pw.h" // use Setup_Psi class -#include "source_pw/module_pwdft/vnl_pw.h" // use pseudopot_cell_vnl + +class pseudopot_cell_vnl; namespace ModuleIO { From 759f288319455343f83437a824f06b19d4d80b57 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 15:43:31 +0800 Subject: [PATCH 06/22] update --- source/source_psi/psi_init_file.cpp | 12 ++++++------ source/source_psi/psi_init_random.cpp | 1 - 2 files changed, 6 insertions(+), 7 deletions(-) diff --git a/source/source_psi/psi_init_file.cpp b/source/source_psi/psi_init_file.cpp index e97eed23260..ad7c0917ddd 100644 --- a/source/source_psi/psi_init_file.cpp +++ b/source/source_psi/psi_init_file.cpp @@ -36,16 +36,16 @@ void psi_init_file::init_psig(T* psig, const int& ik) int ik_tot = this->ik2iktot_[ik]; // mohan update, this is for plane wave, 2025-05-17 - const int out_type = 2; - const bool out_app_flag = false; - const bool gamma_only = false; - const int istep = -1; + const int out_type = 2; + const bool out_app_flag = false; + const bool gamma_only = false; + const int istep = -1; - std::string fn = ModuleIO::filename_output(PARAM.globalv.global_readin_dir,"wf","pw", + std::string fn = ModuleIO::filename_output(PARAM.globalv.global_readin_dir,"wf","pw", ik,this->ik2iktot_,PARAM.inp.nspin,nkstot, out_type,out_app_flag,gamma_only,istep); - ModuleIO::read_wfc_pw(fn, this->pw_wfc_, + ModuleIO::read_wfc_pw(fn, this->pw_wfc_, GlobalV::RANK_IN_POOL, GlobalV::NPROC_IN_POOL, PARAM.inp.nbands, PARAM.globalv.npol, ik, ik_tot, nkstot, wfcatom); diff --git a/source/source_psi/psi_init_random.cpp b/source/source_psi/psi_init_random.cpp index 2ca461ffc13..21668bc29bb 100644 --- a/source/source_psi/psi_init_random.cpp +++ b/source/source_psi/psi_init_random.cpp @@ -1,6 +1,5 @@ #include "psi_init_random.h" #include -#include "source_io/module_parameter/parameter.h" template void psi_init_random::initialize(const Structure_Factor* sf, From 4bb4c6622c41a2f0104c8d699a7332d825aef737 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 15:54:58 +0800 Subject: [PATCH 07/22] some small updates --- source/source_psi/psi_base.h | 27 ++++++++--- source/source_psi/psi_init_atomic.cpp | 53 +++++++++------------- source/source_psi/psi_init_atomic.h | 5 +- source/source_psi/psi_init_atomic_random.h | 2 +- source/source_psi/psi_init_nao.h | 15 +++++- source/source_psi/psi_prepare.cpp | 1 + source/source_psi/setup_psi.cpp | 16 +++---- source/source_psi/setup_psi.h | 8 ++-- 8 files changed, 74 insertions(+), 53 deletions(-) diff --git a/source/source_psi/psi_base.h b/source/source_psi/psi_base.h index 5457ba04646..b0a50373a4f 100644 --- a/source/source_psi/psi_base.h +++ b/source/source_psi/psi_base.h @@ -1,12 +1,9 @@ #ifndef PSI_BASE_H #define PSI_BASE_H -// data structure support -#include "source_basis/module_pw/pw_basis_k.h" // for kpoint related data structure +#include "source_basis/module_pw/pw_basis_k.h" #include "source_pw/module_pwdft/structure_factor.h" -#include "source_psi/psi.h" // for psi data structure -// smart pointer for auto-memory management +#include "source_psi/psi.h" #include -// numerical algorithm support #ifdef __MPI #include #endif @@ -15,7 +12,9 @@ #include #include + using namespace std; + /* Psi (planewave based wavefunction) base class Auther: Kirk0830 @@ -116,6 +115,7 @@ class psi_base } protected: + #ifdef __MPI // MPI additional implementation /// @brief mapping from (ix, iy) to is void stick_to_pool(Real* stick, //< stick @@ -127,20 +127,35 @@ class psi_base const int iw_end, ///< iw_end, ending band index const int ik, ///< ik, kpoint index const int mode = 1); ///< mode, 0 for rr*exp(i*arg), 1 for rr/(1+gk2)*exp(i*arg) + const Structure_Factor* sf_ = nullptr; ///< Structure_Factor + const ModulePW::PW_Basis_K* pw_wfc_ = nullptr; ///< use |k+G>, |G>, getgpluskcar and so on in PW_Basis_K + const UnitCell* p_ucell_ = nullptr; ///< UnitCell + int lmaxkb_ = 0; ///< max angular momentum for non-local projectors + std::vector ik2iktot_; ///< local->global k-point mapping + int nkstot_ = 0; ///< total number of k-points + int random_seed_ = 1; ///< random seed, shared by random, atomic+random, nao+random + std::vector ixy2is_; ///< used by stick_to_pool function + int mem_saver_ = 0; ///< if save memory, only for nscf + std::string method_ = "none"; ///< method name + int nbands_complem_ = 0; ///< complement number of bands, which is nbands_start_ - ucell.natomwfc + double mixing_coef_ = 0; ///< mixing coefficient for atomic+random and nao+random + int nbands_start_ = 0; ///< starting nbands, which is no less than PARAM.inp.nbands + int npol_ = 1; ///< number of polarizations + int nbands_ = 1; ///< number of bands }; -#endif \ No newline at end of file +#endif diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index 06e29b3c11e..4870d612ee1 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -1,31 +1,15 @@ #include "psi_init_atomic.h" #include "source_pw/module_pwdft/soc.h" -// numerical algorithm support #include "source_base/math_integral.h" // for numerical integration #include "source_base/math_polyint.h" // for polynomial interpolation #include "source_base/math_ylmreal.h" // for real spherical harmonics #include "source_base/math_sphbes.h" // for spherical bessel functions -// basic functions support #include "source_base/tool_quit.h" #include "source_base/timer.h" -// global variables definition #include "source_base/global_variable.h" #include "source_io/module_parameter/parameter.h" -// io support #include "source_io/module_output/write_pao.h" -// free function, compared with common radial function normalization, it does not multiply r to function -// due to pswfc is already multiplied by r -// template -// void normalize(int n_rgrid, std::vector& pswfcr, double* rab) -// { -// std::vector pswfc2r2(pswfcr.size()); -// std::transform(pswfcr.begin(), pswfcr.end(), pswfc2r2.begin(), [](T pswfc) { return pswfc * pswfc; }); -// T norm = ModuleBase::Integral::simpson(n_rgrid, pswfc2r2.data(), rab); -// norm = sqrt(norm); -// std::transform(pswfcr.begin(), pswfcr.end(), pswfcr.begin(), [norm](T pswfc) { return pswfc / norm; }); -// } - template void psi_init_atomic::allocate_ps_table() { @@ -38,7 +22,8 @@ void psi_init_atomic::allocate_ps_table() } if (dim2 == 0) { - ModuleBase::WARNING_QUIT("psi_init_atomic::allocate_table", "there is not ANY pseudo atomic orbital read in present system, recommand other methods, quit."); + ModuleBase::WARNING_QUIT("psi_init_atomic::allocate_table", + "there is not ANY pseudo atomic orbital read in present system, recommand other methods, quit."); } int dim3 = PARAM.globalv.nqx; // allocate memory for ovlp_flzjlq @@ -48,23 +33,26 @@ void psi_init_atomic::allocate_ps_table() template void psi_init_atomic::initialize(const Structure_Factor* sf, //< structure factor - const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis - const UnitCell* p_ucell, //< unit cell - const std::vector& ik2iktot, - const int& nkstot, - const int& random_seed, //< random seed - const int& lmaxkb, - const int& rank, - const int& npol, - const int& nbands) + const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + const UnitCell* p_ucell, //< unit cell + const std::vector& ik2iktot, + const int& nkstot, + const int& random_seed, //< random seed + const int& lmaxkb, + const int& rank, + const int& npol, + const int& nbands) { ModuleBase::timer::start("psi_init_atomic", "initialize"); psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); + this->nbands_start_ = std::max(this->p_ucell_->natomwfc, nbands); this->nbands_complem_ = this->nbands_start_ - this->p_ucell_->natomwfc; + // allocate this->allocate_ps_table(); + // then for generate random number to fill in the wavefunction this->ixy2is_.clear(); this->ixy2is_.resize(this->pw_wfc_->fftnxy); @@ -87,7 +75,7 @@ void psi_init_atomic::tabulate() { max_msh = (this->p_ucell_->atoms[it].ncpp.msh > max_msh) ? this->p_ucell_->atoms[it].ncpp.msh : max_msh; } - ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running,"max mesh points in Pseudopotential",max_msh); + ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running,"max mesh points in Pseudopotential",max_msh); this->ovlp_pswfcjlq_.zero_out(); const int startq = 0; @@ -95,14 +83,14 @@ void psi_init_atomic::tabulate() std::vector aux(max_msh); std::vector vchi(max_msh); - ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running,"dq(describe PAO in reciprocal space)",PARAM.globalv.dq); - ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running,"max q",PARAM.globalv.nqx); + ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running,"dq(describe PAO in reciprocal space)",PARAM.globalv.dq); + ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running,"max q",PARAM.globalv.nqx); for (int it=0; itp_ucell_->ntype; it++) { - Atom* atom = &this->p_ucell_->atoms[it]; + Atom* atom = &this->p_ucell_->atoms[it]; - GlobalV::ofs_running<<"\n number of pseudo atomic orbitals for "<label<<" is "<< atom->ncpp.nchi << std::endl; + GlobalV::ofs_running<<"\n number of pseudo atomic orbitals for "<label<<" is "<< atom->ncpp.nchi << std::endl; // QE uses atom->ncpp.mesh const int n_rgrid = (PARAM.inp.pseudo_mesh) ? atom->ncpp.mesh : atom->ncpp.msh; @@ -229,6 +217,7 @@ template void psi_init_atomic::init_psig(T* psig, const int& ik) { ModuleBase::timer::start("psi_init_atomic", "init_psig"); + const int npw = this->pw_wfc_->npwk[ik]; const int npwk_max = this->pw_wfc_->npwk_max; int lmax = this->p_ucell_->lmax_ppwf; @@ -484,4 +473,4 @@ template class psi_init_atomic>; template class psi_init_atomic>; // gamma point calculation template class psi_init_atomic; -template class psi_init_atomic; \ No newline at end of file +template class psi_init_atomic; diff --git a/source/source_psi/psi_init_atomic.h b/source/source_psi/psi_init_atomic.h index 90a91fd8e19..a4e12d36cd0 100644 --- a/source/source_psi/psi_init_atomic.h +++ b/source/source_psi/psi_init_atomic.h @@ -36,9 +36,12 @@ class psi_init_atomic : public psi_base virtual void init_psig(T* psig, const int& ik) override; protected: + // allocate memory for overlap table void allocate_ps_table(); + std::vector pseudopot_files_; + ModuleBase::realArray ovlp_pswfcjlq_; }; -#endif \ No newline at end of file +#endif diff --git a/source/source_psi/psi_init_atomic_random.h b/source/source_psi/psi_init_atomic_random.h index 17ff511ad8b..7cea8580172 100644 --- a/source/source_psi/psi_init_atomic_random.h +++ b/source/source_psi/psi_init_atomic_random.h @@ -35,4 +35,4 @@ class psi_init_atomic_random : public psi_init_atomic private: }; -#endif \ No newline at end of file +#endif diff --git a/source/source_psi/psi_init_nao.h b/source/source_psi/psi_init_nao.h index cbd2c3c2ac1..1d1f663427f 100644 --- a/source/source_psi/psi_init_nao.h +++ b/source/source_psi/psi_init_nao.h @@ -37,51 +37,64 @@ class psi_init_nao : public psi_base const int& = 1) override; //< nbands void read_external_orbs(const std::string* orbital_files, const int& rank); + virtual void tabulate() override; + std::vector external_orbs() const { return orbital_files_; } + std::vector> nr() const { return nr_; } + std::vector nr(const int& itype) const { return nr_[itype]; } + int nr(const int& itype, const int& ichi) const { return nr_[itype][ichi]; } + std::vector>> chi() const { return chi_; } + std::vector> chi(const int& itype) const { return chi_[itype]; } + std::vector chi(const int& itype, const int& ichi) const { return chi_[itype][ichi]; } + double chi(const int& itype, const int& ichi, const int& ir) const { return chi_[itype][ichi][ir]; } + std::vector>> rgrid() const { return rgrid_; } + std::vector> rgrid(const int& itype) const { return rgrid_[itype]; } + std::vector rgrid(const int& itype, const int& ichi) const { return rgrid_[itype][ichi]; } + double rgrid(const int& itype, const int& ichi, const int& ir) const { return rgrid_[itype][ichi][ir]; @@ -104,4 +117,4 @@ class psi_init_nao : public psi_base /// @brief useful for atomic-like methods ModuleBase::SphericalBesselTransformer sbt; }; -#endif \ No newline at end of file +#endif diff --git a/source/source_psi/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index bf33edabf17..ca0461dd5b8 100644 --- a/source/source_psi/psi_prepare.cpp +++ b/source/source_psi/psi_prepare.cpp @@ -14,6 +14,7 @@ #include "source_psi/psi_init_nao.h" #include "source_psi/psi_init_nao_random.h" #include "source_psi/psi_init_random.h" + namespace psi { diff --git a/source/source_psi/setup_psi.cpp b/source/source_psi/setup_psi.cpp index c5ca08caaea..cdb7a0c247b 100644 --- a/source/source_psi/setup_psi.cpp +++ b/source/source_psi/setup_psi.cpp @@ -13,10 +13,10 @@ Setup_Psi::~Setup_Psi(){} // In that case, psi may change its size multiple times during SCF template void Setup_Psi::allocate_psi( - psi::Psi* &psi, - const K_Vectors &kv, - const Parallel_Orbitals ¶_orb, - const Input_para &inp) + psi::Psi* &psi, + const K_Vectors &kv, + const Parallel_Orbitals ¶_orb, + const Input_para &inp) { // init electronic wave function psi if (psi == nullptr) @@ -51,10 +51,10 @@ void Setup_Psi::allocate_psi( template void Setup_Psi::deallocate_psi(psi::Psi* &psi) { - if(psi!=nullptr) - { - delete psi; - } + if(psi!=nullptr) + { + delete psi; + } } template class Setup_Psi; diff --git a/source/source_psi/setup_psi.h b/source/source_psi/setup_psi.h index a4c0f11d3e3..5ce3f1f2fe8 100644 --- a/source/source_psi/setup_psi.h +++ b/source/source_psi/setup_psi.h @@ -14,11 +14,11 @@ class Setup_Psi Setup_Psi(); ~Setup_Psi(); - static void allocate_psi( - psi::Psi* &psi, - const K_Vectors &kv, + static void allocate_psi( + psi::Psi* &psi, + const K_Vectors &kv, const Parallel_Orbitals ¶_orb, - const Input_para &inp); + const Input_para &inp); static void deallocate_psi(psi::Psi* &psi); From 5bb12372a1e7f5dec8e352392d156e2db3a16571 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 16:33:42 +0800 Subject: [PATCH 08/22] fix a bug --- source/source_psi/psi_init_atomic.cpp | 2 +- source/source_psi/psi_init_atomic_random.cpp | 2 +- source/source_psi/psi_init_file.cpp | 4 ++-- source/source_psi/psi_init_nao.cpp | 2 +- source/source_psi/psi_init_nao_random.cpp | 2 +- source/source_psi/test/psi_init_test.cpp | 23 ++++++++++++-------- 6 files changed, 20 insertions(+), 15 deletions(-) diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index 4870d612ee1..f96f7c3d591 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -223,7 +223,7 @@ void psi_init_atomic::init_psig(T* psig, const int& ik) int lmax = this->p_ucell_->lmax_ppwf; const int total_lm = (lmax + 1) * (lmax + 1); ModuleBase::matrix ylm(total_lm, npw); - ModuleBase::GlobalFunc::ZEROS(psig, PARAM.globalv.npol * this->nbands_start_ * npwk_max); + ModuleBase::GlobalFunc::ZEROS(psig, this->npol_ * this->nbands_start_ * npwk_max); std::vector> aux(npw); std::vector chiaux(npw); diff --git a/source/source_psi/psi_init_atomic_random.cpp b/source/source_psi/psi_init_atomic_random.cpp index b1f358db474..7bff9deecb6 100644 --- a/source/source_psi/psi_init_atomic_random.cpp +++ b/source/source_psi/psi_init_atomic_random.cpp @@ -22,7 +22,7 @@ void psi_init_atomic_random::init_psig(T* psig, const int& ik) { double rm = this->mixing_coef_; psi_init_atomic::init_psig(psig, ik); - const int npol = PARAM.globalv.npol; + const int npol = this->npol_; const int nbasis = this->pw_wfc_->npwk_max * npol; psi::Psi psi_random(1, this->nbands_start_, nbasis, nbasis, true); psi_random.fix_k(0); diff --git a/source/source_psi/psi_init_file.cpp b/source/source_psi/psi_init_file.cpp index ad7c0917ddd..eb3d8dca878 100644 --- a/source/source_psi/psi_init_file.cpp +++ b/source/source_psi/psi_init_file.cpp @@ -29,7 +29,7 @@ template void psi_init_file::init_psig(T* psig, const int& ik) { ModuleBase::timer::start("psi_init_file", "init_psig"); - const int npol = PARAM.globalv.npol; + const int npol = this->npol_; const int nbasis = this->pw_wfc_->npwk_max * npol; const int nkstot = this->nkstot_; ModuleBase::ComplexMatrix wfcatom(this->nbands_start_, nbasis); @@ -47,7 +47,7 @@ void psi_init_file::init_psig(T* psig, const int& ik) ModuleIO::read_wfc_pw(fn, this->pw_wfc_, GlobalV::RANK_IN_POOL, GlobalV::NPROC_IN_POOL, - PARAM.inp.nbands, PARAM.globalv.npol, + PARAM.inp.nbands, this->npol_, ik, ik_tot, nkstot, wfcatom); assert(this->nbands_start_ <= wfcatom.nr); diff --git a/source/source_psi/psi_init_nao.cpp b/source/source_psi/psi_init_nao.cpp index 6ca0e7e7184..6798354dcb0 100644 --- a/source/source_psi/psi_init_nao.cpp +++ b/source/source_psi/psi_init_nao.cpp @@ -253,7 +253,7 @@ void psi_init_nao::init_psig(T* psig, const int& ik) const int npwk_max = this->pw_wfc_->npwk_max; const int total_lm = (this->p_ucell_->lmax + 1) * (this->p_ucell_->lmax + 1); ModuleBase::matrix ylm(total_lm, npw); - ModuleBase::GlobalFunc::ZEROS(psig, PARAM.globalv.npol * this->nbands_start_ * npwk_max); + ModuleBase::GlobalFunc::ZEROS(psig, this->npol_ * this->nbands_start_ * npwk_max); std::vector> aux(npw); std::vector qnorm(npw); diff --git a/source/source_psi/psi_init_nao_random.cpp b/source/source_psi/psi_init_nao_random.cpp index c7bd0e3bcac..23661f24d41 100644 --- a/source/source_psi/psi_init_nao_random.cpp +++ b/source/source_psi/psi_init_nao_random.cpp @@ -22,7 +22,7 @@ void psi_init_nao_random::init_psig(T* psig, const int& ik) { double rm = this->mixing_coef_; psi_init_nao::init_psig(psig, ik); - const int npol = PARAM.globalv.npol; + const int npol = this->npol_; const int nbasis = this->pw_wfc_->npwk_max * npol; psi::Psi psi_random(1, this->nbands_start_, nbasis, nbasis, true); psi_random.fix_k(0); diff --git a/source/source_psi/test/psi_init_test.cpp b/source/source_psi/test/psi_init_test.cpp index bf2a18e8f5f..aa22b08d923 100644 --- a/source/source_psi/test/psi_init_test.cpp +++ b/source/source_psi/test/psi_init_test.cpp @@ -219,14 +219,14 @@ class PsiIntializerUnitTest : public ::testing::Test { } this->p_pw_wfc->igl2isz_k = new int[1]; this->p_pw_wfc->igl2isz_k[0] = 0; + if(this->p_pw_wfc->igl2ig_k != nullptr) { delete[] this->p_pw_wfc->igl2ig_k; +} + this->p_pw_wfc->igl2ig_k = new int[1]; + this->p_pw_wfc->igl2ig_k[0] = 0; if(this->p_pw_wfc->gcar != nullptr) { delete[] this->p_pw_wfc->gcar; } this->p_pw_wfc->gcar = new ModuleBase::Vector3[1]; this->p_pw_wfc->gcar[0] = {0.0, 0.0, 0.0}; - if(this->p_pw_wfc->igl2isz_k != nullptr) { delete[] this->p_pw_wfc->igl2isz_k; -} - this->p_pw_wfc->igl2isz_k = new int[1]; - this->p_pw_wfc->igl2isz_k[0] = 0; if(this->p_pw_wfc->gk2 != nullptr) { delete[] this->p_pw_wfc->gk2; } this->p_pw_wfc->gk2 = new double[1]; @@ -356,7 +356,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) { this->nkstot_, this->random_seed, this->lmaxkb, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + PARAM.sys.npol); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -384,7 +385,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) { this->nkstot_, this->random_seed, this->lmaxkb, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + PARAM.sys.npol); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -473,7 +475,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) { this->nkstot_, this->random_seed, this->lmaxkb, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + PARAM.sys.npol); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -498,7 +501,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSo) { this->nkstot_, this->random_seed, this->lmaxkb, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + PARAM.sys.npol); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -523,7 +527,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSoDOMAG) { this->nkstot_, this->random_seed, this->lmaxkb, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + PARAM.sys.npol); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; From 7c9342d48a3eea14d2b0d1400c6323839eedd621 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 16:48:11 +0800 Subject: [PATCH 09/22] fix another bug --- .../to_wannier90_lcao_in_pw.cpp | 2 +- source/source_psi/psi_base.cpp | 2 +- source/source_psi/psi_base.h | 20 +++++------ source/source_psi/psi_init_atomic.h | 20 +++++------ source/source_psi/psi_init_atomic_random.h | 20 +++++------ source/source_psi/psi_init_file.h | 20 +++++------ source/source_psi/psi_init_nao.h | 20 +++++------ source/source_psi/psi_init_nao_random.h | 20 +++++------ source/source_psi/psi_init_random.h | 20 +++++------ source/source_psi/test/psi_init_test.cpp | 35 +++++++++++++------ 10 files changed, 97 insertions(+), 82 deletions(-) diff --git a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp index 3344dfb6fb1..a6441fab1d0 100644 --- a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp +++ b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp @@ -45,7 +45,7 @@ void toWannier90_LCAO_IN_PW::calculate( ModulePW::PW_Basis_K* wfcpw_ptr = const_cast(wfcpw); delete this->psi_initer_; this->psi_initer_ = new psi_init_nao>(); - this->psi_initer_->initialize(sf_ptr, wfcpw_ptr, &ucell, kv.ik2iktot, kv.get_nkstot(), 1, 0, GlobalV::MY_RANK); + this->psi_initer_->initialize(sf_ptr, wfcpw_ptr, &ucell, kv.ik2iktot, kv.get_nkstot(), 1, 0, GlobalV::MY_RANK, PARAM.globalv.npol, PARAM.inp.nbands); this->psi_initer_->tabulate(); delete this->psi; const int nks_psi = (PARAM.inp.calculation == "nscf" && PARAM.inp.mem_saver == 1)? 1 : wfcpw->nks; diff --git a/source/source_psi/psi_base.cpp b/source/source_psi/psi_base.cpp index 152920f06ee..8f1ed9de0a5 100644 --- a/source/source_psi/psi_base.cpp +++ b/source/source_psi/psi_base.cpp @@ -50,7 +50,7 @@ void psi_base::random_t(T* psi, const int iw_start, const int iw_end, const i assert(iw_start >= 0); const int ng = this->pw_wfc_->npwk[ik]; const int npwk_max = this->pw_wfc_->npwk_max; - const int npol = PARAM.globalv.npol; + const int npol = this->npol_; // If random seed is specified, then generate random wavefunction satisfying that // it can generate the same results using different number of processors. diff --git a/source/source_psi/psi_base.h b/source/source_psi/psi_base.h index b0a50373a4f..fb7d54f648a 100644 --- a/source/source_psi/psi_base.h +++ b/source/source_psi/psi_base.h @@ -56,16 +56,16 @@ class psi_base psi_base(){}; virtual ~psi_base(){}; /// @brief initialize the psi_base with external data and methods - virtual void initialize(const Structure_Factor*, //< structure factor - const ModulePW::PW_Basis_K*, //< planewave basis - const UnitCell*, //< unit cell - const std::vector& = {}, //< ik2iktot: local->global k-point mapping - const int& = 0, //< nkstot: total number of k-points - const int& = 1, //< random seed - const int& = 0, //< lmaxkb: max angular momentum for non-local projectors - const int& = 0, //< rank - const int& = 1, //< npol - const int& = 1); //< nbands + virtual void initialize(const Structure_Factor* sf, //< structure factor + const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + const UnitCell* p_ucell, //< unit cell + const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping + const int& nkstot, //< nkstot: total number of k-points + const int& random_seed, //< random seed + const int& lmaxkb, //< lmaxkb: max angular momentum for non-local projectors + const int& rank, //< rank + const int& npol, //< npol + const int& nbands); //< nbands /// @brief CENTRAL FUNCTION: calculate the interpolate table if needed virtual void tabulate() diff --git a/source/source_psi/psi_init_atomic.h b/source/source_psi/psi_init_atomic.h index a4e12d36cd0..5bff541393e 100644 --- a/source/source_psi/psi_init_atomic.h +++ b/source/source_psi/psi_init_atomic.h @@ -22,16 +22,16 @@ class psi_init_atomic : public psi_base ~psi_init_atomic(){}; /// @brief initialize the psi_init with external data and methods - virtual void initialize(const Structure_Factor*, //< structure factor - const ModulePW::PW_Basis_K*, //< planewave basis - const UnitCell*, //< unit cell - const std::vector& = {}, //< ik2iktot: local->global k-point mapping - const int& = 0, //< nkstot: total number of k-points - const int& = 1, //< random seed - const int& = 0, //< lmaxkb: max angular momentum for non-local projectors - const int& = 0, //< MPI rank - const int& = 1, //< npol - const int& = 1) override; //< nbands + virtual void initialize(const Structure_Factor* sf, //< structure factor + const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + const UnitCell* p_ucell, //< unit cell + const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping + const int& nkstot, //< nkstot: total number of k-points + const int& random_seed, //< random seed + const int& lmaxkb, //< lmaxkb: max angular momentum for non-local projectors + const int& rank, //< MPI rank + const int& npol, //< npol + const int& nbands) override; //< nbands virtual void tabulate() override; virtual void init_psig(T* psig, const int& ik) override; diff --git a/source/source_psi/psi_init_atomic_random.h b/source/source_psi/psi_init_atomic_random.h index 7cea8580172..085f67eb494 100644 --- a/source/source_psi/psi_init_atomic_random.h +++ b/source/source_psi/psi_init_atomic_random.h @@ -20,16 +20,16 @@ class psi_init_atomic_random : public psi_init_atomic ~psi_init_atomic_random(){}; /// @brief initialize the psi_base with external data and methods - virtual void initialize(const Structure_Factor*, //< structure factor - const ModulePW::PW_Basis_K*, //< planewave basis - const UnitCell*, //< unit cell - const std::vector& = {}, //< ik2iktot: local->global k-point mapping - const int& = 0, //< nkstot: total number of k-points - const int& = 1, //< random seed - const int& = 0, //< lmaxkb: max angular momentum for non-local projectors - const int& = 0, //< MPI rank - const int& = 1, //< npol - const int& = 1) override; //< nbands + virtual void initialize(const Structure_Factor* sf, //< structure factor + const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + const UnitCell* p_ucell, //< unit cell + const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping + const int& nkstot, //< nkstot: total number of k-points + const int& random_seed, //< random seed + const int& lmaxkb, //< lmaxkb: max angular momentum for non-local projectors + const int& rank, //< MPI rank + const int& npol, //< npol + const int& nbands) override; //< nbands virtual void init_psig(T* psig, const int& ik) override; diff --git a/source/source_psi/psi_init_file.h b/source/source_psi/psi_init_file.h index 9bca8eebf2c..12d44b8e44b 100644 --- a/source/source_psi/psi_init_file.h +++ b/source/source_psi/psi_init_file.h @@ -21,16 +21,16 @@ class psi_init_file : public psi_base ~psi_init_file(){}; /// @brief initialize the psi_base with external data and methods - virtual void initialize(const Structure_Factor*, //< structure factor - const ModulePW::PW_Basis_K*, //< planewave basis - const UnitCell*, //< unit cell - const std::vector& = {}, //< ik2iktot: local->global k-point mapping - const int& = 0, //< nkstot: total number of k-points - const int& = 1, //< random seed - const int& = 0, //< lmaxkb: max angular momentum for non-local projectors - const int& = 0, //< MPI rank - const int& = 1, //< npol - const int& = 1) override; //< nbands + virtual void initialize(const Structure_Factor* sf, //< structure factor + const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + const UnitCell* p_ucell, //< unit cell + const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping + const int& nkstot, //< nkstot: total number of k-points + const int& random_seed, //< random seed + const int& lmaxkb, //< lmaxkb: max angular momentum for non-local projectors + const int& rank, //< MPI rank + const int& npol, //< npol + const int& nbands) override; //< nbands /// @brief calculate and output planewave wavefunction /// @param ik kpoint index diff --git a/source/source_psi/psi_init_nao.h b/source/source_psi/psi_init_nao.h index 1d1f663427f..01dbf3dde0c 100644 --- a/source/source_psi/psi_init_nao.h +++ b/source/source_psi/psi_init_nao.h @@ -25,16 +25,16 @@ class psi_init_nao : public psi_base virtual void init_psig(T* psig, const int& ik) override; /// @brief initialize the psi_base with external data and methods - virtual void initialize(const Structure_Factor*, //< structure factor - const ModulePW::PW_Basis_K*, //< planewave basis - const UnitCell*, //< unit cell - const std::vector& = {}, //< ik2iktot: local->global k-point mapping - const int& = 0, //< nkstot: total number of k-points - const int& = 1, //< random seed - const int& = 0, //< lmaxkb: max angular momentum for non-local projectors - const int& = 0, //< MPI rank - const int& = 1, //< npol - const int& = 1) override; //< nbands + virtual void initialize(const Structure_Factor* sf, //< structure factor + const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + const UnitCell* p_ucell, //< unit cell + const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping + const int& nkstot, //< nkstot: total number of k-points + const int& random_seed, //< random seed + const int& lmaxkb, //< lmaxkb: max angular momentum for non-local projectors + const int& rank, //< MPI rank + const int& npol, //< npol + const int& nbands) override; //< nbands void read_external_orbs(const std::string* orbital_files, const int& rank); diff --git a/source/source_psi/psi_init_nao_random.h b/source/source_psi/psi_init_nao_random.h index d91f27d598d..e3e10e7fc03 100644 --- a/source/source_psi/psi_init_nao_random.h +++ b/source/source_psi/psi_init_nao_random.h @@ -20,16 +20,16 @@ class psi_init_nao_random : public psi_init_nao ~psi_init_nao_random(){}; /// @brief initialize the psi_init with external data and methods - virtual void initialize(const Structure_Factor*, //< structure factor - const ModulePW::PW_Basis_K*, //< planewave basis - const UnitCell*, //< unit cell - const std::vector& = {}, //< ik2iktot: local->global k-point mapping - const int& = 0, //< nkstot: total number of k-points - const int& = 1, //< random seed - const int& = 0, //< lmaxkb: max angular momentum for non-local projectors - const int& = 0, //< MPI rank - const int& = 1, //< npol - const int& = 1) override; //< nbands + virtual void initialize(const Structure_Factor* sf, //< structure factor + const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + const UnitCell* p_ucell, //< unit cell + const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping + const int& nkstot, //< nkstot: total number of k-points + const int& random_seed, //< random seed + const int& lmaxkb, //< lmaxkb: max angular momentum for non-local projectors + const int& rank, //< MPI rank + const int& npol, //< npol + const int& nbands) override; //< nbands virtual void init_psig(T* psig, const int& ik) override; }; diff --git a/source/source_psi/psi_init_random.h b/source/source_psi/psi_init_random.h index a830edac250..3a514588f02 100644 --- a/source/source_psi/psi_init_random.h +++ b/source/source_psi/psi_init_random.h @@ -25,15 +25,15 @@ class psi_init_random : public psi_base /// @return initialized planewave wavefunction (psi::Psi>*) virtual void init_psig(T* psig, const int& ik) override; /// @brief initialize the psi_init with external data and methods - virtual void initialize(const Structure_Factor*, //< structure factor - const ModulePW::PW_Basis_K*, //< planewave basis - const UnitCell*, //< unit cell - const std::vector& = {}, //< ik2iktot: local->global k-point mapping - const int& = 0, //< nkstot: total number of k-points - const int& = 1, //< random seed - const int& = 0, //< lmaxkb: max angular momentum for non-local projectors - const int& = 0, //< MPI rank - const int& = 1, //< npol - const int& = 1) override; //< nbands + virtual void initialize(const Structure_Factor* sf, //< structure factor + const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + const UnitCell* p_ucell, //< unit cell + const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping + const int& nkstot, //< nkstot: total number of k-points + const int& random_seed, //< random seed + const int& lmaxkb, //< lmaxkb: max angular momentum for non-local projectors + const int& rank, //< MPI rank + const int& npol, //< npol + const int& nbands) override; //< nbands }; #endif \ No newline at end of file diff --git a/source/source_psi/test/psi_init_test.cpp b/source/source_psi/test/psi_init_test.cpp index aa22b08d923..f776a88d73c 100644 --- a/source/source_psi/test/psi_init_test.cpp +++ b/source/source_psi/test/psi_init_test.cpp @@ -312,7 +312,9 @@ TEST_F(PsiIntializerUnitTest, CalPsigRandom) { this->nkstot_, this->random_seed, this->lmaxkb, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + PARAM.sys.npol, + PARAM.input.nbands); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -332,7 +334,9 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomic) { this->nkstot_, this->random_seed, this->lmaxkb, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + PARAM.sys.npol, + PARAM.input.nbands); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -357,7 +361,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol); + PARAM.sys.npol, + PARAM.input.nbands); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -386,7 +391,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol); + PARAM.sys.npol, + PARAM.input.nbands); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -410,7 +416,9 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) { this->nkstot_, this->random_seed, this->lmaxkb, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + PARAM.sys.npol, + PARAM.input.nbands); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -430,7 +438,9 @@ TEST_F(PsiIntializerUnitTest, CalPsigNao) { this->nkstot_, this->random_seed, this->lmaxkb, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + PARAM.sys.npol, + PARAM.input.nbands); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -450,7 +460,9 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) { this->nkstot_, this->random_seed, this->lmaxkb, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + PARAM.sys.npol, + PARAM.input.nbands); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -476,7 +488,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol); + PARAM.sys.npol, + PARAM.input.nbands); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -502,7 +515,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSo) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol); + PARAM.sys.npol, + PARAM.input.nbands); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; @@ -528,7 +542,8 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSoDOMAG) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol); + PARAM.sys.npol, + PARAM.input.nbands); this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG const int nbands_start = this->psi_init->nbands_start(); const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; From fe412e9e8bb6d837fa28d6bb1cfd7583c7c3b795 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 17:11:49 +0800 Subject: [PATCH 10/22] update --- source/source_psi/psi_init_file.cpp | 23 ++++++++++++++++++----- source/source_psi/psi_init_file.h | 10 ++++++++++ source/source_psi/psi_prepare.cpp | 9 ++++++++- 3 files changed, 36 insertions(+), 6 deletions(-) diff --git a/source/source_psi/psi_init_file.cpp b/source/source_psi/psi_init_file.cpp index eb3d8dca878..95d13caae4f 100644 --- a/source/source_psi/psi_init_file.cpp +++ b/source/source_psi/psi_init_file.cpp @@ -2,11 +2,12 @@ #include #include +#include +#include #include "source_base/timer.h" #include "source_io/module_wf/read_wfc_pw.h" #include "source_io/module_output/filename.h" -#include "source_io/module_parameter/parameter.h" template void psi_init_file::initialize(const Structure_Factor* sf, @@ -25,6 +26,18 @@ void psi_init_file::initialize(const Structure_Factor* sf, this->nbands_complem_ = 0; } +template +void psi_init_file::prepare_params(const int& nspin, + const std::string& global_readin_dir, + const int& rank_in_pool, + const int& nproc_in_pool) +{ + this->nspin_ = nspin; + this->global_readin_dir_ = global_readin_dir; + this->rank_in_pool_ = rank_in_pool; + this->nproc_in_pool_ = nproc_in_pool; +} + template void psi_init_file::init_psig(T* psig, const int& ik) { @@ -41,13 +54,13 @@ void psi_init_file::init_psig(T* psig, const int& ik) const bool gamma_only = false; const int istep = -1; - std::string fn = ModuleIO::filename_output(PARAM.globalv.global_readin_dir,"wf","pw", - ik,this->ik2iktot_,PARAM.inp.nspin,nkstot, + std::string fn = ModuleIO::filename_output(this->global_readin_dir_,"wf","pw", + ik,this->ik2iktot_,this->nspin_,nkstot, out_type,out_app_flag,gamma_only,istep); ModuleIO::read_wfc_pw(fn, this->pw_wfc_, - GlobalV::RANK_IN_POOL, GlobalV::NPROC_IN_POOL, - PARAM.inp.nbands, this->npol_, + this->rank_in_pool_, this->nproc_in_pool_, + this->nbands_start_, this->npol_, ik, ik_tot, nkstot, wfcatom); assert(this->nbands_start_ <= wfcatom.nr); diff --git a/source/source_psi/psi_init_file.h b/source/source_psi/psi_init_file.h index 12d44b8e44b..8aeeeb21435 100644 --- a/source/source_psi/psi_init_file.h +++ b/source/source_psi/psi_init_file.h @@ -2,6 +2,7 @@ #define PSI_INIT_FILE_H #include +#include #include "psi_base.h" /* @@ -12,6 +13,10 @@ class psi_init_file : public psi_base { private: using Real = typename GetTypeReal::type; + int nspin_ = 1; + std::string global_readin_dir_; + int rank_in_pool_ = 0; + int nproc_in_pool_ = 1; public: psi_init_file() @@ -36,5 +41,10 @@ class psi_init_file : public psi_base /// @param ik kpoint index /// @return initialized planewave wavefunction (psi::Psi>*) virtual void init_psig(T* psig, const int& ik) override; + + void prepare_params(const int& nspin, + const std::string& global_readin_dir, + const int& rank_in_pool, + const int& nproc_in_pool); }; #endif diff --git a/source/source_psi/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index ca0461dd5b8..cb07d1ecadf 100644 --- a/source/source_psi/psi_prepare.cpp +++ b/source/source_psi/psi_prepare.cpp @@ -51,7 +51,14 @@ void PSIPrepare::prepare_init(const int& random_seed) } else if (this->init_wfc == "file") { - this->psi_initer = std::unique_ptr>(new psi_init_file()); + psi_init_file* file_initer = new psi_init_file(); + file_initer->prepare_params( + PARAM.inp.nspin, + PARAM.globalv.global_readin_dir, + GlobalV::RANK_IN_POOL, + GlobalV::NPROC_IN_POOL + ); + this->psi_initer = std::unique_ptr>(file_initer); GlobalV::ofs_running << "\n Using FILE starting wave functions\n"; } else if ((this->init_wfc.substr(0, 6) == "atomic") && (this->ucell.natomwfc == 0)) From b54d8face42eb2a1adaf15cf8edd81ace831d88e Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 17:16:49 +0800 Subject: [PATCH 11/22] update atomic --- source/source_psi/psi_init_atomic.cpp | 36 +++++++++++++++++++-------- source/source_psi/psi_init_atomic.h | 13 ++++++++++ source/source_psi/psi_prepare.cpp | 22 ++++++++++++++-- 3 files changed, 59 insertions(+), 12 deletions(-) diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index f96f7c3d591..a7e0086d997 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -25,12 +25,28 @@ void psi_init_atomic::allocate_ps_table() ModuleBase::WARNING_QUIT("psi_init_atomic::allocate_table", "there is not ANY pseudo atomic orbital read in present system, recommand other methods, quit."); } - int dim3 = PARAM.globalv.nqx; + int dim3 = this->nqx_; // allocate memory for ovlp_flzjlq this->ovlp_pswfcjlq_.create(dim1, dim2, dim3); this->ovlp_pswfcjlq_.zero_out(); } +template +void psi_init_atomic::prepare_params(const int& nqx, + const double& dq, + const int& nspin, + const bool& domag, + const bool& domag_z, + const bool& pseudo_mesh) +{ + this->nqx_ = nqx; + this->dq_ = dq; + this->nspin_ = nspin; + this->domag_ = domag; + this->domag_z_ = domag_z; + this->pseudo_mesh_ = pseudo_mesh; +} + template void psi_init_atomic::initialize(const Structure_Factor* sf, //< structure factor const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis @@ -83,8 +99,8 @@ void psi_init_atomic::tabulate() std::vector aux(max_msh); std::vector vchi(max_msh); - ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running,"dq(describe PAO in reciprocal space)",PARAM.globalv.dq); - ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running,"max q",PARAM.globalv.nqx); + ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running,"dq(describe PAO in reciprocal space)",this->dq_); + ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running,"max q",this->nqx_); for (int it=0; itp_ucell_->ntype; it++) { @@ -93,7 +109,7 @@ void psi_init_atomic::tabulate() GlobalV::ofs_running<<"\n number of pseudo atomic orbitals for "<label<<" is "<< atom->ncpp.nchi << std::endl; // QE uses atom->ncpp.mesh - const int n_rgrid = (PARAM.inp.pseudo_mesh) ? atom->ncpp.mesh : atom->ncpp.msh; + const int n_rgrid = (this->pseudo_mesh_) ? atom->ncpp.mesh : atom->ncpp.msh; std::vector chi2(n_rgrid); for (int ic = 0; ic < atom->ncpp.nchi ;ic++) @@ -186,9 +202,9 @@ void psi_init_atomic::tabulate() } const int l = atom->ncpp.lchi[ic]; - for (int iq = startq; iq < PARAM.globalv.nqx; iq++) + for (int iq = startq; iq < this->nqx_; iq++) { - const double q = PARAM.globalv.dq * iq; + const double q = this->dq_ * iq; ModuleBase::Sphbes::Spherical_Bessel(atom->ncpp.msh, atom->ncpp.r.data(), q, l, aux.data()); for (int ir = 0; ir < atom->ncpp.msh; ir++) { @@ -259,17 +275,17 @@ void psi_init_atomic::init_psig(T* psig, const int& ik) { ovlp_pswfcjlg[ig] = ModuleBase::PolyInt::Polynomial_Interpolation( this->ovlp_pswfcjlq_, it, ipswfc, - PARAM.globalv.nqx, PARAM.globalv.dq, gk[ig].norm() * this->p_ucell_->tpiba ); + this->nqx_, this->dq_, gk[ig].norm() * this->p_ucell_->tpiba ); } /* NSPIN == 4 */ - if(PARAM.inp.nspin == 4) + if(this->nspin_ == 4) { if(this->p_ucell_->atoms[it].ncpp.has_so) { Soc soc; soc.rot_ylm(l + 1); const double j = this->p_ucell_->atoms[it].ncpp.jchi[ipswfc]; /* NOT NONCOLINEAR CASE, rotation matrix become identity */ - if (!(PARAM.globalv.domag||PARAM.globalv.domag_z)) + if (!(this->domag_||this->domag_z_)) { double cg_coeffs[2]; for(int m = -l-1; m < l+1; m++) @@ -352,7 +368,7 @@ void psi_init_atomic::init_psig(T* psig, const int& ik) chiaux[ig] = l * ModuleBase::PolyInt::Polynomial_Interpolation( this->ovlp_pswfcjlq_, it, ipswfc_noncolin_soc, - PARAM.globalv.nqx, PARAM.globalv.dq, gk[ig].norm() * this->p_ucell_->tpiba); + this->nqx_, this->dq_, gk[ig].norm() * this->p_ucell_->tpiba); chiaux[ig] += ovlp_pswfcjlg[ig] * (l + 1.0) ; chiaux[ig] *= 1/(2.0*l+1.0); } diff --git a/source/source_psi/psi_init_atomic.h b/source/source_psi/psi_init_atomic.h index 5bff541393e..ea3d6f294a8 100644 --- a/source/source_psi/psi_init_atomic.h +++ b/source/source_psi/psi_init_atomic.h @@ -13,6 +13,12 @@ class psi_init_atomic : public psi_base { private: using Real = typename GetTypeReal::type; + int nqx_ = 0; + double dq_ = 0.0; + int nspin_ = 1; + bool domag_ = false; + bool domag_z_ = false; + bool pseudo_mesh_ = false; public: psi_init_atomic() @@ -35,6 +41,13 @@ class psi_init_atomic : public psi_base virtual void tabulate() override; virtual void init_psig(T* psig, const int& ik) override; + void prepare_params(const int& nqx, + const double& dq, + const int& nspin, + const bool& domag, + const bool& domag_z, + const bool& pseudo_mesh); + protected: // allocate memory for overlap table diff --git a/source/source_psi/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index cb07d1ecadf..e5596252244 100644 --- a/source/source_psi/psi_prepare.cpp +++ b/source/source_psi/psi_prepare.cpp @@ -91,11 +91,29 @@ void PSIPrepare::prepare_init(const int& random_seed) GlobalV::ofs_running << "\n Using ATOMIC starting wave functions for all " << this->ucell.natomwfc << " atomic orbitals" << " (covers " << PARAM.inp.nbands << " bands)\n"; } - this->psi_initer = std::unique_ptr>(new psi_init_atomic()); + psi_init_atomic* atomic_initer = new psi_init_atomic(); + atomic_initer->prepare_params( + PARAM.globalv.nqx, + PARAM.globalv.dq, + PARAM.inp.nspin, + PARAM.globalv.domag, + PARAM.globalv.domag_z, + PARAM.inp.pseudo_mesh + ); + this->psi_initer = std::unique_ptr>(atomic_initer); } else if (this->init_wfc == "atomic+random") { - this->psi_initer = std::unique_ptr>(new psi_init_atomic_random()); + psi_init_atomic_random* atomic_rand_initer = new psi_init_atomic_random(); + atomic_rand_initer->prepare_params( + PARAM.globalv.nqx, + PARAM.globalv.dq, + PARAM.inp.nspin, + PARAM.globalv.domag, + PARAM.globalv.domag_z, + PARAM.inp.pseudo_mesh + ); + this->psi_initer = std::unique_ptr>(atomic_rand_initer); GlobalV::ofs_running << "\n Using ATOMIC+RANDOM starting wave functions with " << this->ucell.natomwfc << " atomic orbitals\n"; } From c640071d6c350464dc8165af6b11305d7fa8cfc9 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 17:32:02 +0800 Subject: [PATCH 12/22] update --- source/source_psi/psi_init_nao.cpp | 24 ++++++++++++++++++------ source/source_psi/psi_init_nao.h | 10 ++++++++++ source/source_psi/psi_prepare.cpp | 18 ++++++++++++++++-- 3 files changed, 44 insertions(+), 8 deletions(-) diff --git a/source/source_psi/psi_init_nao.cpp b/source/source_psi/psi_init_nao.cpp index 6798354dcb0..3fefd257287 100644 --- a/source/source_psi/psi_init_nao.cpp +++ b/source/source_psi/psi_init_nao.cpp @@ -42,6 +42,18 @@ void normalize(const std::vector& r, std::vector& flz) std::transform(flz.begin(), flz.end(), flz.begin(), [norm](double flz) { return flz / norm; }); } +template +void psi_init_nao::prepare_params(const int& nqx, + const double& dq, + const int& nspin, + const std::string& orbital_dir) +{ + this->nqx_ = nqx; + this->dq_ = dq; + this->nspin_ = nspin; + this->orbital_dir_ = orbital_dir; +} + template void psi_init_nao::read_external_orbs(const std::string* orbital_files, const int& rank) { @@ -67,7 +79,7 @@ void psi_init_nao::read_external_orbs(const std::string* orbital_files, const bool is_open = false; if (rank == 0) { - ifs_it.open(PARAM.inp.orbital_dir + this->orbital_files_[it]); + ifs_it.open(this->orbital_dir_ + this->orbital_files_[it]); is_open = ifs_it.is_open(); } #ifdef __MPI @@ -195,9 +207,9 @@ void psi_init_nao::tabulate() ModuleBase::timer::start("psi_init_nao", "tabulate"); // a uniformed qgrid - std::vector qgrid(PARAM.globalv.nqx); + std::vector qgrid(this->nqx_); std::iota(qgrid.begin(), qgrid.end(), 0); - std::for_each(qgrid.begin(), qgrid.end(), [this](double& q) { q = q * PARAM.globalv.dq; }); + std::for_each(qgrid.begin(), qgrid.end(), [this](double& q) { q = q * this->dq_; }); // only when needed, allocate memory for cubspl_ if (this->cubspl_.get()) @@ -219,7 +231,7 @@ void psi_init_nao::tabulate() ModuleBase::SphericalBesselTransformer sbt_(true); // bool: enable cache // tabulate the spherical bessel transform of numerical orbital function - std::vector Jlfq(PARAM.globalv.nqx, 0.0); + std::vector Jlfq(this->nqx_, 0.0); int i = 0; for (int it = 0; it < this->p_ucell_->ntype; it++) { @@ -232,7 +244,7 @@ void psi_init_nao::tabulate() this->nr_[it][ic], this->rgrid_[it][ic].data(), this->chi_[it][ic].data(), - PARAM.globalv.nqx, + this->nqx_, qgrid.data(), Jlfq.data()); this->cubspl_->add(Jlfq.data()); @@ -294,7 +306,7 @@ void psi_init_nao::init_psig(T* psig, const int& ik) this->cubspl_->eval(npw, qnorm.data(), Jlfq.data(), nullptr, nullptr, this->projmap_(it, L, N)); /* FOR EVERY NAO IN EACH ATOM */ - if (PARAM.inp.nspin == 4) + if (this->nspin_ == 4) { /* FOR EACH SPIN CHANNEL */ for (int is_N = 0; is_N < 2; is_N++) // rotate base diff --git a/source/source_psi/psi_init_nao.h b/source/source_psi/psi_init_nao.h index 01dbf3dde0c..2efb72b0319 100644 --- a/source/source_psi/psi_init_nao.h +++ b/source/source_psi/psi_init_nao.h @@ -6,6 +6,7 @@ #include "psi_base.h" #include +#include /* Psi (planewave based wavefunction) initializer: numerical atomic orbital method */ @@ -14,6 +15,10 @@ class psi_init_nao : public psi_base { private: using Real = typename GetTypeReal::type; + int nqx_ = 0; + double dq_ = 0.0; + int nspin_ = 1; + std::string orbital_dir_; public: psi_init_nao() @@ -40,6 +45,11 @@ class psi_init_nao : public psi_base virtual void tabulate() override; + void prepare_params(const int& nqx, + const double& dq, + const int& nspin, + const std::string& orbital_dir); + std::vector external_orbs() const { return orbital_files_; diff --git a/source/source_psi/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index e5596252244..7a61efb2d89 100644 --- a/source/source_psi/psi_prepare.cpp +++ b/source/source_psi/psi_prepare.cpp @@ -119,12 +119,26 @@ void PSIPrepare::prepare_init(const int& random_seed) } else if (this->init_wfc == "nao") { - this->psi_initer = std::unique_ptr>(new psi_init_nao()); + psi_init_nao* nao_initer = new psi_init_nao(); + nao_initer->prepare_params( + PARAM.globalv.nqx, + PARAM.globalv.dq, + PARAM.inp.nspin, + PARAM.inp.orbital_dir + ); + this->psi_initer = std::unique_ptr>(nao_initer); GlobalV::ofs_running << "\n Using NAO starting wave functions\n"; } else if (this->init_wfc == "nao+random") { - this->psi_initer = std::unique_ptr>(new psi_init_nao_random()); + psi_init_nao_random* nao_rand_initer = new psi_init_nao_random(); + nao_rand_initer->prepare_params( + PARAM.globalv.nqx, + PARAM.globalv.dq, + PARAM.inp.nspin, + PARAM.inp.orbital_dir + ); + this->psi_initer = std::unique_ptr>(nao_rand_initer); GlobalV::ofs_running << "\n Using NAO+RANDOM starting wave functions\n"; } else From 44a248d81367a35c704e6071c97269aef15c2d30 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Thu, 23 Jul 2026 17:35:25 +0800 Subject: [PATCH 13/22] remove parameter.h --- source/source_psi/psi_init_atomic.cpp | 1 - source/source_psi/psi_init_nao.cpp | 4 +--- 2 files changed, 1 insertion(+), 4 deletions(-) diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index a7e0086d997..7a0036d0a15 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -7,7 +7,6 @@ #include "source_base/tool_quit.h" #include "source_base/timer.h" #include "source_base/global_variable.h" -#include "source_io/module_parameter/parameter.h" #include "source_io/module_output/write_pao.h" template diff --git a/source/source_psi/psi_init_nao.cpp b/source/source_psi/psi_init_nao.cpp index 3fefd257287..45c0ee2cc05 100644 --- a/source/source_psi/psi_init_nao.cpp +++ b/source/source_psi/psi_init_nao.cpp @@ -17,9 +17,7 @@ #include "source_base/parallel_reduce.h" #endif #include "source_io/module_output/orb_io.h" -#include "source_io/module_parameter/parameter.h" -// GlobalV::NQX and GlobalV::DQ are here -#include "source_io/module_parameter/parameter.h" + #include #include From b2e9a0d9291f2cd392e1c3c95a2d45539a9e9ed2 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Fri, 24 Jul 2026 10:53:35 +0800 Subject: [PATCH 14/22] update psi init tests --- source/source_psi/test/psi_init_test.cpp | 282 ++++++++++++++--------- 1 file changed, 171 insertions(+), 111 deletions(-) diff --git a/source/source_psi/test/psi_init_test.cpp b/source/source_psi/test/psi_init_test.cpp index f776a88d73c..8652ae5ea30 100644 --- a/source/source_psi/test/psi_init_test.cpp +++ b/source/source_psi/test/psi_init_test.cpp @@ -1,14 +1,20 @@ #include -#define private public -#include "source_io/module_parameter/parameter.h" -#undef private +#include + +#include "source_pw/module_pwdft/vl_pw.h" +#include "source_pw/module_pwdft/structure_factor.h" +#include "source_pw/module_pwdft/parallel_grid.h" +#include "source_basis/module_pw/pw_basis.h" +#include "source_cell/pseudo.h" +#include "source_cell/atom_pseudo.h" +#include "source_cell/magnetism.h" +#include "source_cell/unitcell.h" #include "../psi_base.h" #include "../psi_init_atomic.h" #include "../psi_init_atomic_random.h" #include "../psi_init_nao.h" #include "../psi_init_nao_random.h" #include "../psi_init_random.h" -#include "source_pw/module_pwdft/vl_pw.h" #include "source_base/output.h" /* @@ -98,7 +104,16 @@ class PsiIntializerUnitTest : public ::testing::Test { psi_base>* psi_init; - private: + int nbands_ = 1; + int nspin_ = 1; + int npol_ = 1; + bool domag_ = false; + bool domag_z_ = false; + std::string orbital_dir_ = "./support/"; + int nqx_ = 100; + double dq_ = 0.01; + bool pseudo_mesh_ = false; + protected: void SetUp() override { @@ -106,17 +121,6 @@ class PsiIntializerUnitTest : public ::testing::Test { this->p_sf = new Structure_Factor(); this->p_pw_wfc = new ModulePW::PW_Basis_K(); this->p_ucell = new UnitCell(); - // mock - PARAM.input.nbands = 1; - PARAM.input.nspin = 1; - PARAM.input.orbital_dir = "./support/"; - PARAM.input.pseudo_dir = "./support/"; - PARAM.sys.npol = 1; - PARAM.input.calculation = "scf"; - PARAM.input.init_wfc = "random"; - PARAM.input.ks_solver = "cg"; - PARAM.sys.domag = false; - PARAM.sys.domag_z = false; // lattice this->p_ucell->a1 = {10.0, 0.0, 0.0}; this->p_ucell->a2 = {0.0, 10.0, 0.0}; @@ -159,19 +163,28 @@ class PsiIntializerUnitTest : public ::testing::Test { this->p_ucell->atoms[0].ncpp.mesh = 11; this->p_ucell->atoms[0].ncpp.msh = 11; this->p_ucell->atoms[0].ncpp.lmax = 2; - //if(this->p_ucell->atoms[0].ncpp.rab != nullptr) delete[] this->p_ucell->atoms[0].ncpp.rab; + this->p_ucell->atoms[0].ncpp.rab = std::vector(11, 0.0); - for(int i = 0; i < 11; ++i) { this->p_ucell->atoms[0].ncpp.rab[i] = 0.01; -} - //if(this->p_ucell->atoms[0].ncpp.r != nullptr) delete[] this->p_ucell->atoms[0].ncpp.r; + for(int i = 0; i < 11; ++i) + { + this->p_ucell->atoms[0].ncpp.rab[i] = 0.01; + } + this->p_ucell->atoms[0].ncpp.r = std::vector(11, 0.0); - for(int i = 0; i < 11; ++i) { this->p_ucell->atoms[0].ncpp.r[i] = 0.01*i; -} - this->p_ucell->atoms[0].ncpp.chi.create(2, 11); - for(int i = 0; i < 2; ++i) { for(int j = 0; j < 11; ++j) { this->p_ucell->atoms[0].ncpp.chi(i, j) = 0.01; -} -} - //if(this->p_ucell->atoms[0].ncpp.lchi != nullptr) delete[] this->p_ucell->atoms[0].ncpp.lchi; + for(int i = 0; i < 11; ++i) + { + this->p_ucell->atoms[0].ncpp.r[i] = 0.01*i; + } + + this->p_ucell->atoms[0].ncpp.chi.create(2, 11); + for(int i = 0; i < 2; ++i) + { + for(int j = 0; j < 11; ++j) + { + this->p_ucell->atoms[0].ncpp.chi(i, j) = 0.01; + } + } + this->p_ucell->atoms[0].ncpp.lchi = std::vector(2, 0); this->p_ucell->atoms[0].ncpp.lchi[0] = 0; this->p_ucell->atoms[0].ncpp.lchi[1] = 1; @@ -184,6 +197,7 @@ class PsiIntializerUnitTest : public ::testing::Test { this->p_ucell->atoms[0].ncpp.jchi = std::vector(2, 0.0); this->p_ucell->atoms[0].ncpp.jchi[0] = 0.5; this->p_ucell->atoms[0].ncpp.jchi[1] = 1.5; + // atom numerical orbital this->p_ucell->lmax = 2; p_ucell->orbital_fn.shrink_to_fit(); @@ -199,20 +213,30 @@ class PsiIntializerUnitTest : public ::testing::Test { // can support function PW_Basis::getfftixy2is this->p_pw_wfc->nks = 1; this->p_pw_wfc->npwk_max = 1; - if(this->p_pw_wfc->npwk != nullptr) { delete[] this->p_pw_wfc->npwk; -} - this->p_pw_wfc->npwk = new int[1]; - this->p_pw_wfc->npwk[0] = 1; - this->p_pw_wfc->fftnxy = 1; - this->p_pw_wfc->fftnz = 1; - this->p_pw_wfc->nst = 1; - this->p_pw_wfc->nz = 1; - if(this->p_pw_wfc->is2fftixy != nullptr) { delete[] this->p_pw_wfc->is2fftixy; -} - this->p_pw_wfc->is2fftixy = new int[1]; - this->p_pw_wfc->is2fftixy[0] = 0; - if(this->p_pw_wfc->fftixy2ip != nullptr) { delete[] this->p_pw_wfc->fftixy2ip; -} + if(this->p_pw_wfc->npwk != nullptr) + { + delete[] this->p_pw_wfc->npwk; + } + + this->p_pw_wfc->npwk = new int[1]; + this->p_pw_wfc->npwk[0] = 1; + this->p_pw_wfc->fftnxy = 1; + this->p_pw_wfc->fftnz = 1; + this->p_pw_wfc->nst = 1; + this->p_pw_wfc->nz = 1; + if(this->p_pw_wfc->is2fftixy != nullptr) + { + delete[] this->p_pw_wfc->is2fftixy; + } + + this->p_pw_wfc->is2fftixy = new int[1]; + this->p_pw_wfc->is2fftixy[0] = 0; + + if(this->p_pw_wfc->fftixy2ip != nullptr) + { + delete[] this->p_pw_wfc->fftixy2ip; + } + this->p_pw_wfc->fftixy2ip = new int[1]; this->p_pw_wfc->fftixy2ip[0] = 0; if(this->p_pw_wfc->igl2isz_k != nullptr) { delete[] this->p_pw_wfc->igl2isz_k; @@ -303,7 +327,6 @@ TEST_F(PsiIntializerUnitTest, CastToT) { } TEST_F(PsiIntializerUnitTest, CalPsigRandom) { - PARAM.input.init_wfc = "random"; this->psi_init = new psi_init_random>(); this->psi_init->initialize(this->p_sf, this->p_pw_wfc, @@ -313,11 +336,11 @@ TEST_F(PsiIntializerUnitTest, CalPsigRandom) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol, - PARAM.input.nbands); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->npol_, + this->nbands_); + this->psi_init->tabulate(); const int nbands_start = this->psi_init->nbands_start(); - const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; + const int nbasis = this->p_pw_wfc->npwk_max * this->npol_; psi::Psi>* psi = new psi::Psi>(1, nbands_start, nbasis, nbasis, true); this->psi_init->init_psig(psi->get_pointer(), 0); EXPECT_NEAR(-0.66187696761064307, psi->operator()(0,0,0).real(), 1e-4); @@ -325,8 +348,11 @@ TEST_F(PsiIntializerUnitTest, CalPsigRandom) { } TEST_F(PsiIntializerUnitTest, CalPsigAtomic) { - PARAM.input.init_wfc = "atomic"; this->psi_init = new psi_init_atomic>(); + psi_init_atomic>* atomic_initer = dynamic_cast>*>(this->psi_init); + if (atomic_initer) { + atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + } this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, @@ -335,11 +361,11 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomic) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol, - PARAM.input.nbands); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->npol_, + this->nbands_); + this->psi_init->tabulate(); const int nbands_start = this->psi_init->nbands_start(); - const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; + const int nbasis = this->p_pw_wfc->npwk_max * this->npol_; psi::Psi>* psi = new psi::Psi>(1, nbands_start, nbasis, nbasis, true); this->psi_init->init_psig(psi->get_pointer(), 0); EXPECT_NEAR(0, psi->operator()(0,0,0).real(), 1e-12); @@ -347,12 +373,17 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomic) { } TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) { - PARAM.input.init_wfc = "atomic"; - PARAM.input.nspin = 4; - PARAM.sys.npol = 2; + int nspin_save = this->nspin_; + int npol_save = this->npol_; + this->nspin_ = 4; + this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = false; this->p_ucell->natomwfc *= 2; this->psi_init = new psi_init_atomic>(); + psi_init_atomic>* atomic_initer = dynamic_cast>*>(this->psi_init); + if (atomic_initer) { + atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + } this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, @@ -361,28 +392,33 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol, - PARAM.input.nbands); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->npol_, + this->nbands_); + this->psi_init->tabulate(); const int nbands_start = this->psi_init->nbands_start(); - const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; + const int nbasis = this->p_pw_wfc->npwk_max * this->npol_; psi::Psi>* psi = new psi::Psi>(1, nbands_start, nbasis, nbasis, true); this->psi_init->init_psig(psi->get_pointer(), 0); EXPECT_NEAR(0, psi->operator()(0,0,0).real(), 1e-12); - PARAM.input.nspin = 1; - PARAM.sys.npol = 1; + this->nspin_ = nspin_save; + this->npol_ = npol_save; this->p_ucell->atoms[0].ncpp.has_so = false; this->p_ucell->natomwfc /= 2; delete psi; } TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) { - PARAM.input.init_wfc = "atomic"; - PARAM.input.nspin = 4; - PARAM.sys.npol = 2; + int nspin_save = this->nspin_; + int npol_save = this->npol_; + this->nspin_ = 4; + this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = true; this->p_ucell->natomwfc *= 2; this->psi_init = new psi_init_atomic>(); + psi_init_atomic>* atomic_initer = dynamic_cast>*>(this->psi_init); + if (atomic_initer) { + atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + } this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, @@ -391,24 +427,27 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol, - PARAM.input.nbands); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->npol_, + this->nbands_); + this->psi_init->tabulate(); const int nbands_start = this->psi_init->nbands_start(); - const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; + const int nbasis = this->p_pw_wfc->npwk_max * this->npol_; psi::Psi>* psi = new psi::Psi>(1, nbands_start, nbasis, nbasis, true); this->psi_init->init_psig(psi->get_pointer(), 0); EXPECT_NEAR(0, psi->operator()(0,0,0).real(), 1e-12); - PARAM.input.nspin = 1; - PARAM.sys.npol = 1; + this->nspin_ = nspin_save; + this->npol_ = npol_save; this->p_ucell->atoms[0].ncpp.has_so = false; this->p_ucell->natomwfc /= 2; delete psi; } TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) { - PARAM.input.init_wfc = "atomic+random"; this->psi_init = new psi_init_atomic_random>(); + psi_init_atomic_random>* atomic_rand_initer = dynamic_cast>*>(this->psi_init); + if (atomic_rand_initer) { + atomic_rand_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + } this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, @@ -417,11 +456,11 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol, - PARAM.input.nbands); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->npol_, + this->nbands_); + this->psi_init->tabulate(); const int nbands_start = this->psi_init->nbands_start(); - const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; + const int nbasis = this->p_pw_wfc->npwk_max * this->npol_; psi::Psi>* psi = new psi::Psi>(1, nbands_start, nbasis, nbasis, true); this->psi_init->init_psig(psi->get_pointer(), 0); EXPECT_NEAR(0, psi->operator()(0,0,0).real(), 1e-12); @@ -429,8 +468,11 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) { } TEST_F(PsiIntializerUnitTest, CalPsigNao) { - PARAM.input.init_wfc = "nao"; this->psi_init = new psi_init_nao>(); + psi_init_nao>* nao_initer = dynamic_cast>*>(this->psi_init); + if (nao_initer) { + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + } this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, @@ -439,11 +481,11 @@ TEST_F(PsiIntializerUnitTest, CalPsigNao) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol, - PARAM.input.nbands); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->npol_, + this->nbands_); + this->psi_init->tabulate(); const int nbands_start = this->psi_init->nbands_start(); - const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; + const int nbasis = this->p_pw_wfc->npwk_max * this->npol_; psi::Psi>* psi = new psi::Psi>(1, nbands_start, nbasis, nbasis, true); this->psi_init->init_psig(psi->get_pointer(), 0); EXPECT_NEAR(0, psi->operator()(0,0,0).real(), 1e-12); @@ -451,8 +493,11 @@ TEST_F(PsiIntializerUnitTest, CalPsigNao) { } TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) { - PARAM.input.init_wfc = "nao+random"; this->psi_init = new psi_init_nao_random>(); + psi_init_nao_random>* nao_rand_initer = dynamic_cast>*>(this->psi_init); + if (nao_rand_initer) { + nao_rand_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + } this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, @@ -461,11 +506,11 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol, - PARAM.input.nbands); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->npol_, + this->nbands_); + this->psi_init->tabulate(); const int nbands_start = this->psi_init->nbands_start(); - const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; + const int nbasis = this->p_pw_wfc->npwk_max * this->npol_; psi::Psi>* psi = new psi::Psi>(1, nbands_start, nbasis, nbasis, true); this->psi_init->init_psig(psi->get_pointer(), 0); EXPECT_NEAR(0, psi->operator()(0,0,0).real(), 1e-12); @@ -473,13 +518,16 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) { } TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) { - PARAM.input.init_wfc = "nao"; - PARAM.input.nspin = 4; - PARAM.sys.npol = 2; + int nspin_save = this->nspin_; + int npol_save = this->npol_; + this->nspin_ = 4; + this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = false; - PARAM.sys.domag = false; - PARAM.sys.domag_z = false; this->psi_init = new psi_init_nao>(); + psi_init_nao>* nao_initer = dynamic_cast>*>(this->psi_init); + if (nao_initer) { + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + } this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, @@ -488,25 +536,30 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol, - PARAM.input.nbands); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->npol_, + this->nbands_); + this->psi_init->tabulate(); const int nbands_start = this->psi_init->nbands_start(); - const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; + const int nbasis = this->p_pw_wfc->npwk_max * this->npol_; psi::Psi>* psi = new psi::Psi>(1, nbands_start, nbasis, nbasis, true); this->psi_init->init_psig(psi->get_pointer(), 0); EXPECT_NEAR(0, psi->operator()(0,0,0).real(), 1e-12); + this->nspin_ = nspin_save; + this->npol_ = npol_save; delete psi; } TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSo) { - PARAM.input.init_wfc = "nao"; - PARAM.input.nspin = 4; - PARAM.sys.npol = 2; + int nspin_save = this->nspin_; + int npol_save = this->npol_; + this->nspin_ = 4; + this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = true; - PARAM.sys.domag = false; - PARAM.sys.domag_z = false; this->psi_init = new psi_init_nao>(); + psi_init_nao>* nao_initer = dynamic_cast>*>(this->psi_init); + if (nao_initer) { + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + } this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, @@ -515,25 +568,30 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSo) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol, - PARAM.input.nbands); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->npol_, + this->nbands_); + this->psi_init->tabulate(); const int nbands_start = this->psi_init->nbands_start(); - const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; + const int nbasis = this->p_pw_wfc->npwk_max * this->npol_; psi::Psi>* psi = new psi::Psi>(1, nbands_start, nbasis, nbasis, true); this->psi_init->init_psig(psi->get_pointer(), 0); EXPECT_NEAR(0, psi->operator()(0,0,0).real(), 1e-12); + this->nspin_ = nspin_save; + this->npol_ = npol_save; delete psi; } TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSoDOMAG) { - PARAM.input.init_wfc = "nao"; - PARAM.input.nspin = 4; - PARAM.sys.npol = 2; + int nspin_save = this->nspin_; + int npol_save = this->npol_; + this->nspin_ = 4; + this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = true; - PARAM.sys.domag = true; - PARAM.sys.domag_z = false; this->psi_init = new psi_init_nao>(); + psi_init_nao>* nao_initer = dynamic_cast>*>(this->psi_init); + if (nao_initer) { + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + } this->psi_init->initialize(this->p_sf, this->p_pw_wfc, this->p_ucell, @@ -542,14 +600,16 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSoDOMAG) { this->random_seed, this->lmaxkb, GlobalV::MY_RANK, - PARAM.sys.npol, - PARAM.input.nbands); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->npol_, + this->nbands_); + this->psi_init->tabulate(); const int nbands_start = this->psi_init->nbands_start(); - const int nbasis = this->p_pw_wfc->npwk_max * PARAM.globalv.npol; + const int nbasis = this->p_pw_wfc->npwk_max * this->npol_; psi::Psi>* psi = new psi::Psi>(1, nbands_start, nbasis, nbasis, true); this->psi_init->init_psig(psi->get_pointer(), 0); EXPECT_NEAR(0, psi->operator()(0,0,0).real(), 1e-12); + this->nspin_ = nspin_save; + this->npol_ = npol_save; delete psi; } From 7ce9df225fde0162bae3224344414eb793e6b9da Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Fri, 24 Jul 2026 11:38:20 +0800 Subject: [PATCH 15/22] update --- .../to_wannier90_lcao_in_pw.cpp | 4 +- source/source_psi/psi_init_atomic.cpp | 24 +- source/source_psi/psi_init_atomic.h | 51 ++- source/source_psi/psi_init_atomic_random.h | 21 +- source/source_psi/psi_init_nao.cpp | 13 + source/source_psi/psi_init_nao.h | 45 ++- source/source_psi/psi_init_nao_random.h | 21 +- source/source_psi/test/psi_init_test.cpp | 356 +++++++++--------- 8 files changed, 343 insertions(+), 192 deletions(-) diff --git a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp index a6441fab1d0..4eaa005ce19 100644 --- a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp +++ b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp @@ -44,7 +44,9 @@ void toWannier90_LCAO_IN_PW::calculate( Structure_Factor* sf_ptr = const_cast(&sf); ModulePW::PW_Basis_K* wfcpw_ptr = const_cast(wfcpw); delete this->psi_initer_; - this->psi_initer_ = new psi_init_nao>(); + psi_init_nao>* nao_initer = new psi_init_nao>(); + nao_initer->prepare_params(PARAM.globalv.nqx, PARAM.globalv.dq, PARAM.inp.nspin, PARAM.inp.orbital_dir); + this->psi_initer_ = nao_initer; this->psi_initer_->initialize(sf_ptr, wfcpw_ptr, &ucell, kv.ik2iktot, kv.get_nkstot(), 1, 0, GlobalV::MY_RANK, PARAM.globalv.npol, PARAM.inp.nbands); this->psi_initer_->tabulate(); delete this->psi; diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index 7a0036d0a15..2f7689b5814 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -12,7 +12,7 @@ template void psi_init_atomic::allocate_ps_table() { - // find correct dimension for ovlp_flzjlq + // find correct dimension for ovlp_flzjlq int dim1 = this->p_ucell_->ntype; int dim2 = 0; // dim2 should be the maximum number of pseudo atomic orbitals for (int it = 0; it < this->p_ucell_->ntype; it++) @@ -22,7 +22,12 @@ void psi_init_atomic::allocate_ps_table() if (dim2 == 0) { ModuleBase::WARNING_QUIT("psi_init_atomic::allocate_table", - "there is not ANY pseudo atomic orbital read in present system, recommand other methods, quit."); + "there is not ANY pseudo atomic orbital read in present system, recommand other methods, quit."); + } + if (this->nqx_ <= 0) + { + ModuleBase::WARNING_QUIT("psi_init_atomic::allocate_ps_table", + "nqx_ must be greater than 0. Did you forget to call prepare_params() before initialize()?"); } int dim3 = this->nqx_; // allocate memory for ovlp_flzjlq @@ -44,15 +49,16 @@ void psi_init_atomic::prepare_params(const int& nqx, this->domag_ = domag; this->domag_z_ = domag_z; this->pseudo_mesh_ = pseudo_mesh; + this->params_prepared_ = true; } template -void psi_init_atomic::initialize(const Structure_Factor* sf, //< structure factor - const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis - const UnitCell* p_ucell, //< unit cell +void psi_init_atomic::initialize(const Structure_Factor* sf, + const ModulePW::PW_Basis_K* pw_wfc, + const UnitCell* p_ucell, const std::vector& ik2iktot, const int& nkstot, - const int& random_seed, //< random seed + const int& random_seed, const int& lmaxkb, const int& rank, const int& npol, @@ -60,6 +66,12 @@ void psi_init_atomic::initialize(const Structure_Factor* sf, //< stru { ModuleBase::timer::start("psi_init_atomic", "initialize"); + if (!this->params_prepared_) + { + ModuleBase::WARNING_QUIT("psi_init_atomic::initialize", + "prepare_params() must be called before initialize()"); + } + psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); this->nbands_start_ = std::max(this->p_ucell_->natomwfc, nbands); diff --git a/source/source_psi/psi_init_atomic.h b/source/source_psi/psi_init_atomic.h index ea3d6f294a8..3055458bcd0 100644 --- a/source/source_psi/psi_init_atomic.h +++ b/source/source_psi/psi_init_atomic.h @@ -19,6 +19,7 @@ class psi_init_atomic : public psi_base bool domag_ = false; bool domag_z_ = false; bool pseudo_mesh_ = false; + bool params_prepared_ = false; public: psi_init_atomic() @@ -27,7 +28,48 @@ class psi_init_atomic : public psi_base } ~psi_init_atomic(){}; - /// @brief initialize the psi_init with external data and methods + /** + * @brief Prepare parameters before initialization. + * + * This method must be called before initialize(). It sets up the necessary + * parameters for the psi initialization process. + * + * @param nqx Number of q-points for interpolation + * @param dq Spacing between q-points + * @param nspin Number of spin components + * @param domag Whether to use non-collinear magnetism + * @param domag_z Whether to use z-axis only non-collinear magnetism + * @param pseudo_mesh Whether to use pseudo mesh for radial grid + * + * @see initialize() + */ + void prepare_params(const int& nqx, + const double& dq, + const int& nspin, + const bool& domag, + const bool& domag_z, + const bool& pseudo_mesh); + + /** + * @brief Initialize the psi_init with external data and methods. + * + * This method must be called after prepare_params(). It initializes the + * psi initializer with the provided structure factor, planewave basis, + * and unit cell information. + * + * @param sf Structure factor + * @param pw_wfc Planewave basis + * @param p_ucell Unit cell + * @param ik2iktot Local->global k-point mapping + * @param nkstot Total number of k-points + * @param random_seed Random seed + * @param lmaxkb Max angular momentum for non-local projectors + * @param rank MPI rank + * @param npol Number of polarization components + * @param nbands Number of bands + * + * @see prepare_params() + */ virtual void initialize(const Structure_Factor* sf, //< structure factor const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell @@ -41,13 +83,6 @@ class psi_init_atomic : public psi_base virtual void tabulate() override; virtual void init_psig(T* psig, const int& ik) override; - void prepare_params(const int& nqx, - const double& dq, - const int& nspin, - const bool& domag, - const bool& domag_z, - const bool& pseudo_mesh); - protected: // allocate memory for overlap table diff --git a/source/source_psi/psi_init_atomic_random.h b/source/source_psi/psi_init_atomic_random.h index 085f67eb494..b0e45d6b768 100644 --- a/source/source_psi/psi_init_atomic_random.h +++ b/source/source_psi/psi_init_atomic_random.h @@ -19,7 +19,26 @@ class psi_init_atomic_random : public psi_init_atomic } ~psi_init_atomic_random(){}; - /// @brief initialize the psi_base with external data and methods + /** + * @brief Initialize the psi_init with external data and methods. + * + * This method must be called after prepare_params(). It initializes the + * psi initializer with the provided structure factor, planewave basis, + * and unit cell information. + * + * @param sf Structure factor + * @param pw_wfc Planewave basis + * @param p_ucell Unit cell + * @param ik2iktot Local->global k-point mapping + * @param nkstot Total number of k-points + * @param random_seed Random seed + * @param lmaxkb Max angular momentum for non-local projectors + * @param rank MPI rank + * @param npol Number of polarization components + * @param nbands Number of bands + * + * @see prepare_params() + */ virtual void initialize(const Structure_Factor* sf, //< structure factor const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell diff --git a/source/source_psi/psi_init_nao.cpp b/source/source_psi/psi_init_nao.cpp index 45c0ee2cc05..e325bbd8f83 100644 --- a/source/source_psi/psi_init_nao.cpp +++ b/source/source_psi/psi_init_nao.cpp @@ -50,6 +50,7 @@ void psi_init_nao::prepare_params(const int& nqx, this->dq_ = dq; this->nspin_ = nspin; this->orbital_dir_ = orbital_dir; + this->params_prepared_ = true; } template @@ -170,6 +171,12 @@ void psi_init_nao::initialize(const Structure_Factor* sf, { ModuleBase::timer::start("psi_init_nao", "initialize"); + if (!this->params_prepared_) + { + ModuleBase::WARNING_QUIT("psi_init_nao::initialize", + "prepare_params() must be called before initialize()"); + } + // import psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); @@ -204,6 +211,12 @@ void psi_init_nao::tabulate() { ModuleBase::timer::start("psi_init_nao", "tabulate"); + if (this->nqx_ <= 0) + { + ModuleBase::WARNING_QUIT("psi_init_nao::tabulate", + "nqx_ must be greater than 0. Did you forget to call prepare_params() with valid nqx?"); + } + // a uniformed qgrid std::vector qgrid(this->nqx_); std::iota(qgrid.begin(), qgrid.end(), 0); diff --git a/source/source_psi/psi_init_nao.h b/source/source_psi/psi_init_nao.h index 2efb72b0319..a389f933d99 100644 --- a/source/source_psi/psi_init_nao.h +++ b/source/source_psi/psi_init_nao.h @@ -19,6 +19,7 @@ class psi_init_nao : public psi_base double dq_ = 0.0; int nspin_ = 1; std::string orbital_dir_; + bool params_prepared_ = false; public: psi_init_nao() @@ -27,9 +28,44 @@ class psi_init_nao : public psi_base }; ~psi_init_nao(){}; - virtual void init_psig(T* psig, const int& ik) override; + /** + * @brief Prepare parameters before initialization. + * + * This method must be called before initialize(). It sets up the necessary + * parameters for the psi initialization process. + * + * @param nqx Number of q-points for interpolation + * @param dq Spacing between q-points + * @param nspin Number of spin components + * @param orbital_dir Directory containing orbital files + * + * @see initialize() + */ + void prepare_params(const int& nqx, + const double& dq, + const int& nspin, + const std::string& orbital_dir); - /// @brief initialize the psi_base with external data and methods + /** + * @brief Initialize the psi_init with external data and methods. + * + * This method must be called after prepare_params(). It initializes the + * psi initializer with the provided structure factor, planewave basis, + * and unit cell information. + * + * @param sf Structure factor + * @param pw_wfc Planewave basis + * @param p_ucell Unit cell + * @param ik2iktot Local->global k-point mapping + * @param nkstot Total number of k-points + * @param random_seed Random seed + * @param lmaxkb Max angular momentum for non-local projectors + * @param rank MPI rank + * @param npol Number of polarization components + * @param nbands Number of bands + * + * @see prepare_params() + */ virtual void initialize(const Structure_Factor* sf, //< structure factor const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell @@ -45,10 +81,7 @@ class psi_init_nao : public psi_base virtual void tabulate() override; - void prepare_params(const int& nqx, - const double& dq, - const int& nspin, - const std::string& orbital_dir); + virtual void init_psig(T* psig, const int& ik) override; std::vector external_orbs() const { diff --git a/source/source_psi/psi_init_nao_random.h b/source/source_psi/psi_init_nao_random.h index e3e10e7fc03..4e42964f37b 100644 --- a/source/source_psi/psi_init_nao_random.h +++ b/source/source_psi/psi_init_nao_random.h @@ -19,7 +19,26 @@ class psi_init_nao_random : public psi_init_nao }; ~psi_init_nao_random(){}; - /// @brief initialize the psi_init with external data and methods + /** + * @brief Initialize the psi_init with external data and methods. + * + * This method must be called after prepare_params(). It initializes the + * psi initializer with the provided structure factor, planewave basis, + * and unit cell information. + * + * @param sf Structure factor + * @param pw_wfc Planewave basis + * @param p_ucell Unit cell + * @param ik2iktot Local->global k-point mapping + * @param nkstot Total number of k-points + * @param random_seed Random seed + * @param lmaxkb Max angular momentum for non-local projectors + * @param rank MPI rank + * @param npol Number of polarization components + * @param nbands Number of bands + * + * @see prepare_params() + */ virtual void initialize(const Structure_Factor* sf, //< structure factor const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell diff --git a/source/source_psi/test/psi_init_test.cpp b/source/source_psi/test/psi_init_test.cpp index 8652ae5ea30..b8322cea27a 100644 --- a/source/source_psi/test/psi_init_test.cpp +++ b/source/source_psi/test/psi_init_test.cpp @@ -87,8 +87,10 @@ std::complex* Structure_Factor::get_sk(int ik, int it, int ia, ModulePW: { int npw = wfc_basis->npwk[ik]; std::complex *sk = new std::complex[npw]; - for(int ipw = 0; ipw < npw; ++ipw) { sk[ipw] = std::complex(0.0, 0.0); -} + for(int ipw = 0; ipw < npw; ++ipw) + { + sk[ipw] = std::complex(0.0, 0.0); + } return sk; } @@ -165,25 +167,25 @@ class PsiIntializerUnitTest : public ::testing::Test { this->p_ucell->atoms[0].ncpp.lmax = 2; this->p_ucell->atoms[0].ncpp.rab = std::vector(11, 0.0); - for(int i = 0; i < 11; ++i) - { - this->p_ucell->atoms[0].ncpp.rab[i] = 0.01; - } + for(int i = 0; i < 11; ++i) + { + this->p_ucell->atoms[0].ncpp.rab[i] = 0.01; + } this->p_ucell->atoms[0].ncpp.r = std::vector(11, 0.0); - for(int i = 0; i < 11; ++i) - { - this->p_ucell->atoms[0].ncpp.r[i] = 0.01*i; - } - - this->p_ucell->atoms[0].ncpp.chi.create(2, 11); - for(int i = 0; i < 2; ++i) - { - for(int j = 0; j < 11; ++j) - { - this->p_ucell->atoms[0].ncpp.chi(i, j) = 0.01; - } - } + for(int i = 0; i < 11; ++i) + { + this->p_ucell->atoms[0].ncpp.r[i] = 0.01*i; + } + + this->p_ucell->atoms[0].ncpp.chi.create(2, 11); + for(int i = 0; i < 2; ++i) + { + for(int j = 0; j < 11; ++j) + { + this->p_ucell->atoms[0].ncpp.chi(i, j) = 0.01; + } + } this->p_ucell->atoms[0].ncpp.lchi = std::vector(2, 0); this->p_ucell->atoms[0].ncpp.lchi[0] = 0; @@ -209,67 +211,85 @@ class PsiIntializerUnitTest : public ::testing::Test { this->p_ucell->atoms[0].l_nchi[1] = 2; this->p_ucell->atoms[0].l_nchi[2] = 1; - + // can support function PW_Basis::getfftixy2is this->p_pw_wfc->nks = 1; this->p_pw_wfc->npwk_max = 1; - if(this->p_pw_wfc->npwk != nullptr) - { - delete[] this->p_pw_wfc->npwk; - } - - this->p_pw_wfc->npwk = new int[1]; - this->p_pw_wfc->npwk[0] = 1; - this->p_pw_wfc->fftnxy = 1; - this->p_pw_wfc->fftnz = 1; - this->p_pw_wfc->nst = 1; - this->p_pw_wfc->nz = 1; - if(this->p_pw_wfc->is2fftixy != nullptr) - { - delete[] this->p_pw_wfc->is2fftixy; - } - - this->p_pw_wfc->is2fftixy = new int[1]; - this->p_pw_wfc->is2fftixy[0] = 0; - - if(this->p_pw_wfc->fftixy2ip != nullptr) - { - delete[] this->p_pw_wfc->fftixy2ip; - } + if(this->p_pw_wfc->npwk != nullptr) + { + delete[] this->p_pw_wfc->npwk; + } + + this->p_pw_wfc->npwk = new int[1]; + this->p_pw_wfc->npwk[0] = 1; + this->p_pw_wfc->fftnxy = 1; + this->p_pw_wfc->fftnz = 1; + this->p_pw_wfc->nst = 1; + this->p_pw_wfc->nz = 1; + if(this->p_pw_wfc->is2fftixy != nullptr) + { + delete[] this->p_pw_wfc->is2fftixy; + } + + this->p_pw_wfc->is2fftixy = new int[1]; + this->p_pw_wfc->is2fftixy[0] = 0; + + if(this->p_pw_wfc->fftixy2ip != nullptr) + { + delete[] this->p_pw_wfc->fftixy2ip; + } this->p_pw_wfc->fftixy2ip = new int[1]; this->p_pw_wfc->fftixy2ip[0] = 0; - if(this->p_pw_wfc->igl2isz_k != nullptr) { delete[] this->p_pw_wfc->igl2isz_k; -} + if(this->p_pw_wfc->igl2isz_k != nullptr) + { + delete[] this->p_pw_wfc->igl2isz_k; + } this->p_pw_wfc->igl2isz_k = new int[1]; this->p_pw_wfc->igl2isz_k[0] = 0; - if(this->p_pw_wfc->igl2ig_k != nullptr) { delete[] this->p_pw_wfc->igl2ig_k; -} + if(this->p_pw_wfc->igl2ig_k != nullptr) + { + delete[] this->p_pw_wfc->igl2ig_k; + } this->p_pw_wfc->igl2ig_k = new int[1]; this->p_pw_wfc->igl2ig_k[0] = 0; - if(this->p_pw_wfc->gcar != nullptr) { delete[] this->p_pw_wfc->gcar; -} + if(this->p_pw_wfc->gcar != nullptr) + { + delete[] this->p_pw_wfc->gcar; + } this->p_pw_wfc->gcar = new ModuleBase::Vector3[1]; this->p_pw_wfc->gcar[0] = {0.0, 0.0, 0.0}; - if(this->p_pw_wfc->gk2 != nullptr) { delete[] this->p_pw_wfc->gk2; -} + if(this->p_pw_wfc->gk2 != nullptr) + { + delete[] this->p_pw_wfc->gk2; + } this->p_pw_wfc->gk2 = new double[1]; this->p_pw_wfc->gk2[0] = 0.0; - this->p_pw_wfc->latvec.e11 = this->p_ucell->latvec.e11; this->p_pw_wfc->latvec.e12 = this->p_ucell->latvec.e12; this->p_pw_wfc->latvec.e13 = this->p_ucell->latvec.e13; - this->p_pw_wfc->latvec.e21 = this->p_ucell->latvec.e21; this->p_pw_wfc->latvec.e22 = this->p_ucell->latvec.e22; this->p_pw_wfc->latvec.e23 = this->p_ucell->latvec.e23; - this->p_pw_wfc->latvec.e31 = this->p_ucell->latvec.e31; this->p_pw_wfc->latvec.e32 = this->p_ucell->latvec.e32; this->p_pw_wfc->latvec.e33 = this->p_ucell->latvec.e33; + this->p_pw_wfc->latvec.e11 = this->p_ucell->latvec.e11; + this->p_pw_wfc->latvec.e12 = this->p_ucell->latvec.e12; + this->p_pw_wfc->latvec.e13 = this->p_ucell->latvec.e13; + this->p_pw_wfc->latvec.e21 = this->p_ucell->latvec.e21; + this->p_pw_wfc->latvec.e22 = this->p_ucell->latvec.e22; + this->p_pw_wfc->latvec.e23 = this->p_ucell->latvec.e23; + this->p_pw_wfc->latvec.e31 = this->p_ucell->latvec.e31; + this->p_pw_wfc->latvec.e32 = this->p_ucell->latvec.e32; + this->p_pw_wfc->latvec.e33 = this->p_ucell->latvec.e33; this->p_pw_wfc->G = this->p_ucell->G; this->p_pw_wfc->GT = this->p_ucell->GT; this->p_pw_wfc->GGT = this->p_ucell->GGT; this->p_pw_wfc->lat0 = this->p_ucell->lat0; this->p_pw_wfc->tpiba = 2.0 * M_PI / this->p_ucell->lat0; this->p_pw_wfc->tpiba2 = this->p_pw_wfc->tpiba * this->p_pw_wfc->tpiba; - if(this->p_pw_wfc->kvec_c != nullptr) { delete[] this->p_pw_wfc->kvec_c; -} + if(this->p_pw_wfc->kvec_c != nullptr) + { + delete[] this->p_pw_wfc->kvec_c; + } this->p_pw_wfc->kvec_c = new ModuleBase::Vector3[1]; this->p_pw_wfc->kvec_c[0] = {0.0, 0.0, 0.0}; - if(this->p_pw_wfc->kvec_d != nullptr) { delete[] this->p_pw_wfc->kvec_d; -} + if(this->p_pw_wfc->kvec_d != nullptr) + { + delete[] this->p_pw_wfc->kvec_d; + } this->p_pw_wfc->kvec_d = new ModuleBase::Vector3[1]; this->p_pw_wfc->kvec_d[0] = {0.0, 0.0, 0.0}; @@ -289,32 +309,38 @@ class PsiIntializerUnitTest : public ::testing::Test { } }; -TEST_F(PsiIntializerUnitTest, ConstructorRandom) { +TEST_F(PsiIntializerUnitTest, ConstructorRandom) +{ this->psi_init = new psi_init_random>(); EXPECT_EQ("random", this->psi_init->method()); } -TEST_F(PsiIntializerUnitTest, ConstructorAtomic) { +TEST_F(PsiIntializerUnitTest, ConstructorAtomic) +{ this->psi_init = new psi_init_atomic>(); EXPECT_EQ("atomic", this->psi_init->method()); } -TEST_F(PsiIntializerUnitTest, ConstructorAtomicRandom) { +TEST_F(PsiIntializerUnitTest, ConstructorAtomicRandom) +{ this->psi_init = new psi_init_atomic_random>(); EXPECT_EQ("atomic+random", this->psi_init->method()); } -TEST_F(PsiIntializerUnitTest, ConstructorNao) { +TEST_F(PsiIntializerUnitTest, ConstructorNao) +{ this->psi_init = new psi_init_nao>(); EXPECT_EQ("nao", this->psi_init->method()); } -TEST_F(PsiIntializerUnitTest, ConstructorNaoRandom) { +TEST_F(PsiIntializerUnitTest, ConstructorNaoRandom) +{ this->psi_init = new psi_init_nao_random>(); EXPECT_EQ("nao+random", this->psi_init->method()); } -TEST_F(PsiIntializerUnitTest, CastToT) { +TEST_F(PsiIntializerUnitTest, CastToT) +{ this->psi_init = new psi_init_random>(); std::complex cd = {1.0, 2.0}; std::complex cf = {1.0, 2.0}; @@ -326,13 +352,14 @@ TEST_F(PsiIntializerUnitTest, CastToT) { EXPECT_EQ(this->psi_init->template cast_to_T(cd), f); } -TEST_F(PsiIntializerUnitTest, CalPsigRandom) { +TEST_F(PsiIntializerUnitTest, CalPsigRandom) +{ this->psi_init = new psi_init_random>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->ik2iktot_, - this->nkstot_, + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->lmaxkb, GlobalV::MY_RANK, @@ -347,17 +374,16 @@ TEST_F(PsiIntializerUnitTest, CalPsigRandom) { delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigAtomic) { - this->psi_init = new psi_init_atomic>(); - psi_init_atomic>* atomic_initer = dynamic_cast>*>(this->psi_init); - if (atomic_initer) { - atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); - } - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->ik2iktot_, - this->nkstot_, +TEST_F(PsiIntializerUnitTest, CalPsigAtomic) +{ + psi_init_atomic>* atomic_initer = new psi_init_atomic>(); + atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + this->psi_init = atomic_initer; + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->lmaxkb, GlobalV::MY_RANK, @@ -372,23 +398,22 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomic) { delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) { +TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) +{ int nspin_save = this->nspin_; int npol_save = this->npol_; this->nspin_ = 4; this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = false; this->p_ucell->natomwfc *= 2; - this->psi_init = new psi_init_atomic>(); - psi_init_atomic>* atomic_initer = dynamic_cast>*>(this->psi_init); - if (atomic_initer) { - atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); - } - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->ik2iktot_, - this->nkstot_, + psi_init_atomic>* atomic_initer = new psi_init_atomic>(); + atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + this->psi_init = atomic_initer; + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->lmaxkb, GlobalV::MY_RANK, @@ -407,23 +432,22 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) { delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) { +TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) +{ int nspin_save = this->nspin_; int npol_save = this->npol_; this->nspin_ = 4; this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = true; this->p_ucell->natomwfc *= 2; - this->psi_init = new psi_init_atomic>(); - psi_init_atomic>* atomic_initer = dynamic_cast>*>(this->psi_init); - if (atomic_initer) { - atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); - } - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->ik2iktot_, - this->nkstot_, + psi_init_atomic>* atomic_initer = new psi_init_atomic>(); + atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + this->psi_init = atomic_initer; + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->lmaxkb, GlobalV::MY_RANK, @@ -442,17 +466,16 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) { delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) { - this->psi_init = new psi_init_atomic_random>(); - psi_init_atomic_random>* atomic_rand_initer = dynamic_cast>*>(this->psi_init); - if (atomic_rand_initer) { - atomic_rand_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); - } - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->ik2iktot_, - this->nkstot_, +TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) +{ + psi_init_atomic_random>* atomic_rand_initer = new psi_init_atomic_random>(); + atomic_rand_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + this->psi_init = atomic_rand_initer; + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->lmaxkb, GlobalV::MY_RANK, @@ -467,17 +490,16 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) { delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigNao) { - this->psi_init = new psi_init_nao>(); - psi_init_nao>* nao_initer = dynamic_cast>*>(this->psi_init); - if (nao_initer) { - nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); - } - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->ik2iktot_, - this->nkstot_, +TEST_F(PsiIntializerUnitTest, CalPsigNao) +{ + psi_init_nao>* nao_initer = new psi_init_nao>(); + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + this->psi_init = nao_initer; + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->lmaxkb, GlobalV::MY_RANK, @@ -492,17 +514,16 @@ TEST_F(PsiIntializerUnitTest, CalPsigNao) { delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) { - this->psi_init = new psi_init_nao_random>(); - psi_init_nao_random>* nao_rand_initer = dynamic_cast>*>(this->psi_init); - if (nao_rand_initer) { - nao_rand_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); - } - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->ik2iktot_, - this->nkstot_, +TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) +{ + psi_init_nao_random>* nao_rand_initer = new psi_init_nao_random>(); + nao_rand_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + this->psi_init = nao_rand_initer; + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->lmaxkb, GlobalV::MY_RANK, @@ -517,22 +538,21 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) { delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) { +TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) +{ int nspin_save = this->nspin_; int npol_save = this->npol_; this->nspin_ = 4; this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = false; - this->psi_init = new psi_init_nao>(); - psi_init_nao>* nao_initer = dynamic_cast>*>(this->psi_init); - if (nao_initer) { - nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); - } - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->ik2iktot_, - this->nkstot_, + psi_init_nao>* nao_initer = new psi_init_nao>(); + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + this->psi_init = nao_initer; + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->lmaxkb, GlobalV::MY_RANK, @@ -549,22 +569,21 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) { delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSo) { +TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSo) +{ int nspin_save = this->nspin_; int npol_save = this->npol_; this->nspin_ = 4; this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = true; - this->psi_init = new psi_init_nao>(); - psi_init_nao>* nao_initer = dynamic_cast>*>(this->psi_init); - if (nao_initer) { - nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); - } - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->ik2iktot_, - this->nkstot_, + psi_init_nao>* nao_initer = new psi_init_nao>(); + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + this->psi_init = nao_initer; + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->lmaxkb, GlobalV::MY_RANK, @@ -581,22 +600,21 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSo) { delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSoDOMAG) { +TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSoDOMAG) +{ int nspin_save = this->nspin_; int npol_save = this->npol_; this->nspin_ = 4; this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = true; - this->psi_init = new psi_init_nao>(); - psi_init_nao>* nao_initer = dynamic_cast>*>(this->psi_init); - if (nao_initer) { - nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); - } - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->ik2iktot_, - this->nkstot_, + psi_init_nao>* nao_initer = new psi_init_nao>(); + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + this->psi_init = nao_initer; + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, this->lmaxkb, GlobalV::MY_RANK, From 4e6e6abcc14ca14323f0eb7bd1c24bdcd5cb8834 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Fri, 24 Jul 2026 13:45:25 +0800 Subject: [PATCH 16/22] update --- .../to_wannier90_lcao_in_pw.cpp | 5 +- source/source_psi/psi_base.cpp | 5 +- source/source_psi/psi_base.h | 6 +-- source/source_psi/psi_init_atomic.cpp | 9 ++-- source/source_psi/psi_init_atomic.h | 15 +++--- source/source_psi/psi_init_atomic_random.cpp | 5 +- source/source_psi/psi_init_atomic_random.h | 7 +-- source/source_psi/psi_init_file.cpp | 5 +- source/source_psi/psi_init_file.h | 3 +- source/source_psi/psi_init_nao.cpp | 9 ++-- source/source_psi/psi_init_nao.h | 12 +++-- source/source_psi/psi_init_nao_random.cpp | 5 +- source/source_psi/psi_init_nao_random.h | 7 +-- source/source_psi/psi_init_random.cpp | 5 +- source/source_psi/psi_init_random.h | 3 +- source/source_psi/psi_prepare.cpp | 14 +++-- source/source_psi/test/psi_init_test.cpp | 54 +++++++++---------- 17 files changed, 80 insertions(+), 89 deletions(-) diff --git a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp index 4eaa005ce19..3826308f235 100644 --- a/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp +++ b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp @@ -41,13 +41,12 @@ void toWannier90_LCAO_IN_PW::calculate( { this->ParaV = pv; - Structure_Factor* sf_ptr = const_cast(&sf); ModulePW::PW_Basis_K* wfcpw_ptr = const_cast(wfcpw); delete this->psi_initer_; psi_init_nao>* nao_initer = new psi_init_nao>(); - nao_initer->prepare_params(PARAM.globalv.nqx, PARAM.globalv.dq, PARAM.inp.nspin, PARAM.inp.orbital_dir); + nao_initer->prepare_params(PARAM.globalv.nqx, PARAM.globalv.dq, PARAM.inp.nspin, PARAM.inp.orbital_dir, sf); this->psi_initer_ = nao_initer; - this->psi_initer_->initialize(sf_ptr, wfcpw_ptr, &ucell, kv.ik2iktot, kv.get_nkstot(), 1, 0, GlobalV::MY_RANK, PARAM.globalv.npol, PARAM.inp.nbands); + this->psi_initer_->initialize(wfcpw_ptr, &ucell, kv.ik2iktot, kv.get_nkstot(), 1, 0, GlobalV::MY_RANK, PARAM.globalv.npol, PARAM.inp.nbands); this->psi_initer_->tabulate(); delete this->psi; const int nks_psi = (PARAM.inp.calculation == "nscf" && PARAM.inp.mem_saver == 1)? 1 : wfcpw->nks; diff --git a/source/source_psi/psi_base.cpp b/source/source_psi/psi_base.cpp index 8f1ed9de0a5..f5bc4d44a36 100644 --- a/source/source_psi/psi_base.cpp +++ b/source/source_psi/psi_base.cpp @@ -4,7 +4,6 @@ #include #include -#include "source_pw/module_pwdft/structure_factor.h" #include "source_cell/unitcell.h" #include "source_basis/module_pw/pw_basis_k.h" #include "source_base/parallel_global.h" @@ -20,8 +19,7 @@ #endif template -void psi_base::initialize(const Structure_Factor* sf, - const ModulePW::PW_Basis_K* pw_wfc, +void psi_base::initialize(const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, const std::vector& ik2iktot, const int& nkstot, @@ -31,7 +29,6 @@ void psi_base::initialize(const Structure_Factor* sf, const int& npol, const int& nbands) { - this->sf_ = sf; this->pw_wfc_ = pw_wfc; this->p_ucell_ = p_ucell; this->ik2iktot_ = ik2iktot; diff --git a/source/source_psi/psi_base.h b/source/source_psi/psi_base.h index fb7d54f648a..e77d529c08f 100644 --- a/source/source_psi/psi_base.h +++ b/source/source_psi/psi_base.h @@ -1,7 +1,6 @@ #ifndef PSI_BASE_H #define PSI_BASE_H #include "source_basis/module_pw/pw_basis_k.h" -#include "source_pw/module_pwdft/structure_factor.h" #include "source_psi/psi.h" #include #ifdef __MPI @@ -56,8 +55,7 @@ class psi_base psi_base(){}; virtual ~psi_base(){}; /// @brief initialize the psi_base with external data and methods - virtual void initialize(const Structure_Factor* sf, //< structure factor - const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + virtual void initialize(const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping const int& nkstot, //< nkstot: total number of k-points @@ -128,8 +126,6 @@ class psi_base const int ik, ///< ik, kpoint index const int mode = 1); ///< mode, 0 for rr*exp(i*arg), 1 for rr/(1+gk2)*exp(i*arg) - const Structure_Factor* sf_ = nullptr; ///< Structure_Factor - const ModulePW::PW_Basis_K* pw_wfc_ = nullptr; ///< use |k+G>, |G>, getgpluskcar and so on in PW_Basis_K const UnitCell* p_ucell_ = nullptr; ///< UnitCell diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index 2f7689b5814..21ae88b3183 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -41,7 +41,8 @@ void psi_init_atomic::prepare_params(const int& nqx, const int& nspin, const bool& domag, const bool& domag_z, - const bool& pseudo_mesh) + const bool& pseudo_mesh, + const Structure_Factor& sf) { this->nqx_ = nqx; this->dq_ = dq; @@ -49,12 +50,12 @@ void psi_init_atomic::prepare_params(const int& nqx, this->domag_ = domag; this->domag_z_ = domag_z; this->pseudo_mesh_ = pseudo_mesh; + this->sf_ = &sf; this->params_prepared_ = true; } template -void psi_init_atomic::initialize(const Structure_Factor* sf, - const ModulePW::PW_Basis_K* pw_wfc, +void psi_init_atomic::initialize(const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, const std::vector& ik2iktot, const int& nkstot, @@ -72,7 +73,7 @@ void psi_init_atomic::initialize(const Structure_Factor* sf, "prepare_params() must be called before initialize()"); } - psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); + psi_base::initialize(pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); this->nbands_start_ = std::max(this->p_ucell_->natomwfc, nbands); this->nbands_complem_ = this->nbands_start_ - this->p_ucell_->natomwfc; diff --git a/source/source_psi/psi_init_atomic.h b/source/source_psi/psi_init_atomic.h index 3055458bcd0..53a4a0e4635 100644 --- a/source/source_psi/psi_init_atomic.h +++ b/source/source_psi/psi_init_atomic.h @@ -2,8 +2,12 @@ #define PSI_INIT_ATOMIC_H #include #include +#include #include "source_base/realarray.h" #include "psi_base.h" +#include "source_basis/module_pw/pw_basis_k.h" +#include "source_cell/unitcell.h" +#include "source_pw/module_pwdft/structure_factor.h" /* Psi (planewave based wavefunction) initializer: atomic @@ -20,6 +24,7 @@ class psi_init_atomic : public psi_base bool domag_z_ = false; bool pseudo_mesh_ = false; bool params_prepared_ = false; + const Structure_Factor* sf_ = nullptr; public: psi_init_atomic() @@ -48,16 +53,15 @@ class psi_init_atomic : public psi_base const int& nspin, const bool& domag, const bool& domag_z, - const bool& pseudo_mesh); + const bool& pseudo_mesh, + const Structure_Factor& sf); /** * @brief Initialize the psi_init with external data and methods. * * This method must be called after prepare_params(). It initializes the - * psi initializer with the provided structure factor, planewave basis, - * and unit cell information. + * psi initializer with the provided planewave basis and unit cell information. * - * @param sf Structure factor * @param pw_wfc Planewave basis * @param p_ucell Unit cell * @param ik2iktot Local->global k-point mapping @@ -70,8 +74,7 @@ class psi_init_atomic : public psi_base * * @see prepare_params() */ - virtual void initialize(const Structure_Factor* sf, //< structure factor - const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + virtual void initialize(const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping const int& nkstot, //< nkstot: total number of k-points diff --git a/source/source_psi/psi_init_atomic_random.cpp b/source/source_psi/psi_init_atomic_random.cpp index 7bff9deecb6..d7b51c8d3c2 100644 --- a/source/source_psi/psi_init_atomic_random.cpp +++ b/source/source_psi/psi_init_atomic_random.cpp @@ -3,8 +3,7 @@ #include "source_io/module_parameter/parameter.h" template -void psi_init_atomic_random::initialize(const Structure_Factor* sf, //< structure factor - const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis +void psi_init_atomic_random::initialize(const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell const std::vector& ik2iktot, const int& nkstot, @@ -14,7 +13,7 @@ void psi_init_atomic_random::initialize(const Structure_Factor* sf, / const int& npol, const int& nbands) { - psi_init_atomic::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); + psi_init_atomic::initialize(pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); } template diff --git a/source/source_psi/psi_init_atomic_random.h b/source/source_psi/psi_init_atomic_random.h index b0e45d6b768..b43c633d5e5 100644 --- a/source/source_psi/psi_init_atomic_random.h +++ b/source/source_psi/psi_init_atomic_random.h @@ -23,10 +23,8 @@ class psi_init_atomic_random : public psi_init_atomic * @brief Initialize the psi_init with external data and methods. * * This method must be called after prepare_params(). It initializes the - * psi initializer with the provided structure factor, planewave basis, - * and unit cell information. + * psi initializer with the provided planewave basis and unit cell information. * - * @param sf Structure factor * @param pw_wfc Planewave basis * @param p_ucell Unit cell * @param ik2iktot Local->global k-point mapping @@ -39,8 +37,7 @@ class psi_init_atomic_random : public psi_init_atomic * * @see prepare_params() */ - virtual void initialize(const Structure_Factor* sf, //< structure factor - const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + virtual void initialize(const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping const int& nkstot, //< nkstot: total number of k-points diff --git a/source/source_psi/psi_init_file.cpp b/source/source_psi/psi_init_file.cpp index 95d13caae4f..607bf800c03 100644 --- a/source/source_psi/psi_init_file.cpp +++ b/source/source_psi/psi_init_file.cpp @@ -10,8 +10,7 @@ #include "source_io/module_output/filename.h" template -void psi_init_file::initialize(const Structure_Factor* sf, - const ModulePW::PW_Basis_K* pw_wfc, +void psi_init_file::initialize(const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, const std::vector& ik2iktot, const int& nkstot, @@ -21,7 +20,7 @@ void psi_init_file::initialize(const Structure_Factor* sf, const int& npol, const int& nbands) { - psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); + psi_base::initialize(pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); this->nbands_start_ = nbands; this->nbands_complem_ = 0; } diff --git a/source/source_psi/psi_init_file.h b/source/source_psi/psi_init_file.h index 8aeeeb21435..f8b8ebd796b 100644 --- a/source/source_psi/psi_init_file.h +++ b/source/source_psi/psi_init_file.h @@ -26,8 +26,7 @@ class psi_init_file : public psi_base ~psi_init_file(){}; /// @brief initialize the psi_base with external data and methods - virtual void initialize(const Structure_Factor* sf, //< structure factor - const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + virtual void initialize(const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping const int& nkstot, //< nkstot: total number of k-points diff --git a/source/source_psi/psi_init_nao.cpp b/source/source_psi/psi_init_nao.cpp index e325bbd8f83..57008cc5265 100644 --- a/source/source_psi/psi_init_nao.cpp +++ b/source/source_psi/psi_init_nao.cpp @@ -44,12 +44,14 @@ template void psi_init_nao::prepare_params(const int& nqx, const double& dq, const int& nspin, - const std::string& orbital_dir) + const std::string& orbital_dir, + const Structure_Factor& sf) { this->nqx_ = nqx; this->dq_ = dq; this->nspin_ = nspin; this->orbital_dir_ = orbital_dir; + this->sf_ = &sf; this->params_prepared_ = true; } @@ -158,8 +160,7 @@ void psi_init_nao::allocate_ao_table() } template -void psi_init_nao::initialize(const Structure_Factor* sf, - const ModulePW::PW_Basis_K* pw_wfc, +void psi_init_nao::initialize(const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, const std::vector& ik2iktot, const int& nkstot, @@ -178,7 +179,7 @@ void psi_init_nao::initialize(const Structure_Factor* sf, } // import - psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); + psi_base::initialize(pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); // allocate this->allocate_ao_table(); diff --git a/source/source_psi/psi_init_nao.h b/source/source_psi/psi_init_nao.h index a389f933d99..5c98b9713c3 100644 --- a/source/source_psi/psi_init_nao.h +++ b/source/source_psi/psi_init_nao.h @@ -4,9 +4,14 @@ #include "source_base/realarray.h" #include "source_base/spherical_bessel_transformer.h" #include "psi_base.h" +#include "source_basis/module_pw/pw_basis_k.h" +#include "source_cell/unitcell.h" +#include "source_pw/module_pwdft/structure_factor.h" #include #include +#include +#include /* Psi (planewave based wavefunction) initializer: numerical atomic orbital method */ @@ -20,6 +25,7 @@ class psi_init_nao : public psi_base int nspin_ = 1; std::string orbital_dir_; bool params_prepared_ = false; + const Structure_Factor* sf_ = nullptr; public: psi_init_nao() @@ -44,7 +50,8 @@ class psi_init_nao : public psi_base void prepare_params(const int& nqx, const double& dq, const int& nspin, - const std::string& orbital_dir); + const std::string& orbital_dir, + const Structure_Factor& sf); /** * @brief Initialize the psi_init with external data and methods. @@ -66,8 +73,7 @@ class psi_init_nao : public psi_base * * @see prepare_params() */ - virtual void initialize(const Structure_Factor* sf, //< structure factor - const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + virtual void initialize(const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping const int& nkstot, //< nkstot: total number of k-points diff --git a/source/source_psi/psi_init_nao_random.cpp b/source/source_psi/psi_init_nao_random.cpp index 23661f24d41..aaabbfd8c3a 100644 --- a/source/source_psi/psi_init_nao_random.cpp +++ b/source/source_psi/psi_init_nao_random.cpp @@ -3,8 +3,7 @@ #include "source_io/module_parameter/parameter.h" template -void psi_init_nao_random::initialize(const Structure_Factor* sf, - const ModulePW::PW_Basis_K* pw_wfc, +void psi_init_nao_random::initialize(const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, const std::vector& ik2iktot, const int& nkstot, @@ -14,7 +13,7 @@ void psi_init_nao_random::initialize(const Structure_Factor* sf, const int& npol, const int& nbands) { - psi_init_nao::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); + psi_init_nao::initialize(pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); } template diff --git a/source/source_psi/psi_init_nao_random.h b/source/source_psi/psi_init_nao_random.h index 4e42964f37b..cbe042e68be 100644 --- a/source/source_psi/psi_init_nao_random.h +++ b/source/source_psi/psi_init_nao_random.h @@ -23,10 +23,8 @@ class psi_init_nao_random : public psi_init_nao * @brief Initialize the psi_init with external data and methods. * * This method must be called after prepare_params(). It initializes the - * psi initializer with the provided structure factor, planewave basis, - * and unit cell information. + * psi initializer with the provided planewave basis and unit cell information. * - * @param sf Structure factor * @param pw_wfc Planewave basis * @param p_ucell Unit cell * @param ik2iktot Local->global k-point mapping @@ -39,8 +37,7 @@ class psi_init_nao_random : public psi_init_nao * * @see prepare_params() */ - virtual void initialize(const Structure_Factor* sf, //< structure factor - const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + virtual void initialize(const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping const int& nkstot, //< nkstot: total number of k-points diff --git a/source/source_psi/psi_init_random.cpp b/source/source_psi/psi_init_random.cpp index 21668bc29bb..b509a57d869 100644 --- a/source/source_psi/psi_init_random.cpp +++ b/source/source_psi/psi_init_random.cpp @@ -2,8 +2,7 @@ #include template -void psi_init_random::initialize(const Structure_Factor* sf, - const ModulePW::PW_Basis_K* pw_wfc, +void psi_init_random::initialize(const ModulePW::PW_Basis_K* pw_wfc, const UnitCell* p_ucell, const std::vector& ik2iktot, const int& nkstot, @@ -13,7 +12,7 @@ void psi_init_random::initialize(const Structure_Factor* sf, const int& npol, const int& nbands) { - psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); + psi_base::initialize(pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); this->ixy2is_.clear(); this->ixy2is_.resize(this->pw_wfc_->fftnxy); this->pw_wfc_->getfftixy2is(this->ixy2is_.data()); diff --git a/source/source_psi/psi_init_random.h b/source/source_psi/psi_init_random.h index 3a514588f02..c4c9e979ac0 100644 --- a/source/source_psi/psi_init_random.h +++ b/source/source_psi/psi_init_random.h @@ -25,8 +25,7 @@ class psi_init_random : public psi_base /// @return initialized planewave wavefunction (psi::Psi>*) virtual void init_psig(T* psig, const int& ik) override; /// @brief initialize the psi_init with external data and methods - virtual void initialize(const Structure_Factor* sf, //< structure factor - const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis + virtual void initialize(const ModulePW::PW_Basis_K* pw_wfc, //< planewave basis const UnitCell* p_ucell, //< unit cell const std::vector& ik2iktot, //< ik2iktot: local->global k-point mapping const int& nkstot, //< nkstot: total number of k-points diff --git a/source/source_psi/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index 7a61efb2d89..5ab1c5186ef 100644 --- a/source/source_psi/psi_prepare.cpp +++ b/source/source_psi/psi_prepare.cpp @@ -98,7 +98,8 @@ void PSIPrepare::prepare_init(const int& random_seed) PARAM.inp.nspin, PARAM.globalv.domag, PARAM.globalv.domag_z, - PARAM.inp.pseudo_mesh + PARAM.inp.pseudo_mesh, + this->sf ); this->psi_initer = std::unique_ptr>(atomic_initer); } @@ -111,7 +112,8 @@ void PSIPrepare::prepare_init(const int& random_seed) PARAM.inp.nspin, PARAM.globalv.domag, PARAM.globalv.domag_z, - PARAM.inp.pseudo_mesh + PARAM.inp.pseudo_mesh, + this->sf ); this->psi_initer = std::unique_ptr>(atomic_rand_initer); GlobalV::ofs_running << "\n Using ATOMIC+RANDOM starting wave functions with " @@ -124,7 +126,8 @@ void PSIPrepare::prepare_init(const int& random_seed) PARAM.globalv.nqx, PARAM.globalv.dq, PARAM.inp.nspin, - PARAM.inp.orbital_dir + PARAM.inp.orbital_dir, + this->sf ); this->psi_initer = std::unique_ptr>(nao_initer); GlobalV::ofs_running << "\n Using NAO starting wave functions\n"; @@ -136,7 +139,8 @@ void PSIPrepare::prepare_init(const int& random_seed) PARAM.globalv.nqx, PARAM.globalv.dq, PARAM.inp.nspin, - PARAM.inp.orbital_dir + PARAM.inp.orbital_dir, + this->sf ); this->psi_initer = std::unique_ptr>(nao_rand_initer); GlobalV::ofs_running << "\n Using NAO+RANDOM starting wave functions\n"; @@ -146,7 +150,7 @@ void PSIPrepare::prepare_init(const int& random_seed) ModuleBase::WARNING_QUIT("PSIInit::prepare_init", "for new psi initializer, init_wfc type not supported"); } - this->psi_initer->initialize(&sf, &pw_wfc, &ucell, ik2iktot_, nkstot_, random_seed, lmaxkb, rank, + this->psi_initer->initialize(&pw_wfc, &ucell, ik2iktot_, nkstot_, random_seed, lmaxkb, rank, PARAM.globalv.npol, PARAM.inp.nbands); this->psi_initer->tabulate(); diff --git a/source/source_psi/test/psi_init_test.cpp b/source/source_psi/test/psi_init_test.cpp index b8322cea27a..b7c653b457e 100644 --- a/source/source_psi/test/psi_init_test.cpp +++ b/source/source_psi/test/psi_init_test.cpp @@ -1,5 +1,6 @@ #include #include +#include #include "source_pw/module_pwdft/vl_pw.h" #include "source_pw/module_pwdft/structure_factor.h" @@ -355,8 +356,7 @@ TEST_F(PsiIntializerUnitTest, CastToT) TEST_F(PsiIntializerUnitTest, CalPsigRandom) { this->psi_init = new psi_init_random>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, + this->psi_init->initialize(this->p_pw_wfc, this->p_ucell, this->ik2iktot_, this->nkstot_, @@ -377,10 +377,11 @@ TEST_F(PsiIntializerUnitTest, CalPsigRandom) TEST_F(PsiIntializerUnitTest, CalPsigAtomic) { psi_init_atomic>* atomic_initer = new psi_init_atomic>(); - atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, + this->domag_, this->domag_z_, this->pseudo_mesh_, + *this->p_sf); this->psi_init = atomic_initer; - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, + this->psi_init->initialize(this->p_pw_wfc, this->p_ucell, this->ik2iktot_, this->nkstot_, @@ -407,10 +408,10 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) this->p_ucell->atoms[0].ncpp.has_so = false; this->p_ucell->natomwfc *= 2; psi_init_atomic>* atomic_initer = new psi_init_atomic>(); - atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_, + *this->p_sf); this->psi_init = atomic_initer; - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, + this->psi_init->initialize(this->p_pw_wfc, this->p_ucell, this->ik2iktot_, this->nkstot_, @@ -441,10 +442,10 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) this->p_ucell->atoms[0].ncpp.has_so = true; this->p_ucell->natomwfc *= 2; psi_init_atomic>* atomic_initer = new psi_init_atomic>(); - atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + atomic_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_, + *this->p_sf); this->psi_init = atomic_initer; - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, + this->psi_init->initialize(this->p_pw_wfc, this->p_ucell, this->ik2iktot_, this->nkstot_, @@ -469,10 +470,10 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicSocHasSo) TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) { psi_init_atomic_random>* atomic_rand_initer = new psi_init_atomic_random>(); - atomic_rand_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_); + atomic_rand_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->domag_, this->domag_z_, this->pseudo_mesh_, + *this->p_sf); this->psi_init = atomic_rand_initer; - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, + this->psi_init->initialize(this->p_pw_wfc, this->p_ucell, this->ik2iktot_, this->nkstot_, @@ -493,10 +494,9 @@ TEST_F(PsiIntializerUnitTest, CalPsigAtomicRandom) TEST_F(PsiIntializerUnitTest, CalPsigNao) { psi_init_nao>* nao_initer = new psi_init_nao>(); - nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_, *this->p_sf); this->psi_init = nao_initer; - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, + this->psi_init->initialize(this->p_pw_wfc, this->p_ucell, this->ik2iktot_, this->nkstot_, @@ -517,10 +517,9 @@ TEST_F(PsiIntializerUnitTest, CalPsigNao) TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) { psi_init_nao_random>* nao_rand_initer = new psi_init_nao_random>(); - nao_rand_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + nao_rand_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_, *this->p_sf); this->psi_init = nao_rand_initer; - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, + this->psi_init->initialize(this->p_pw_wfc, this->p_ucell, this->ik2iktot_, this->nkstot_, @@ -546,10 +545,9 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = false; psi_init_nao>* nao_initer = new psi_init_nao>(); - nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_, *this->p_sf); this->psi_init = nao_initer; - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, + this->psi_init->initialize(this->p_pw_wfc, this->p_ucell, this->ik2iktot_, this->nkstot_, @@ -577,10 +575,9 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSo) this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = true; psi_init_nao>* nao_initer = new psi_init_nao>(); - nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_, *this->p_sf); this->psi_init = nao_initer; - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, + this->psi_init->initialize(this->p_pw_wfc, this->p_ucell, this->ik2iktot_, this->nkstot_, @@ -608,10 +605,9 @@ TEST_F(PsiIntializerUnitTest, CalPsigNaoSocHasSoDOMAG) this->npol_ = 2; this->p_ucell->atoms[0].ncpp.has_so = true; psi_init_nao>* nao_initer = new psi_init_nao>(); - nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_); + nao_initer->prepare_params(this->nqx_, this->dq_, this->nspin_, this->orbital_dir_, *this->p_sf); this->psi_init = nao_initer; - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, + this->psi_init->initialize(this->p_pw_wfc, this->p_ucell, this->ik2iktot_, this->nkstot_, From 53b0ab3b521b98e939d3312d5834bb4ad35650c1 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Fri, 24 Jul 2026 14:19:47 +0800 Subject: [PATCH 17/22] update fix bugs --- source/source_esolver/esolver_ks_pw.cpp | 24 ++++- source/source_psi/psi_prepare.cpp | 123 ++++++++++++++---------- source/source_psi/psi_prepare.h | 27 +++++- source/source_psi/psi_prepare_base.h | 22 +++-- source/source_psi/setup_psi_pw.cpp | 74 +++++++++++--- source/source_psi/setup_psi_pw.h | 22 ++++- 6 files changed, 214 insertions(+), 78 deletions(-) diff --git a/source/source_esolver/esolver_ks_pw.cpp b/source/source_esolver/esolver_ks_pw.cpp index a62748386db..743ac4c9c62 100644 --- a/source/source_esolver/esolver_ks_pw.cpp +++ b/source/source_esolver/esolver_ks_pw.cpp @@ -143,7 +143,18 @@ void ESolver_KS_PW::before_scf(UnitCell& ucell, const int istep) if (ucell.cell_parameter_updated) { - this->stp.p_psi_init->prepare_init(PARAM.inp.pw_seed); + this->stp.p_psi_init->prepare_init(PARAM.inp.pw_seed, + PARAM.globalv.nbands_l, + PARAM.inp.nspin, + PARAM.globalv.nqx, + PARAM.globalv.dq, + PARAM.globalv.domag, + PARAM.globalv.domag_z, + PARAM.inp.pseudo_mesh, + PARAM.inp.orbital_dir, + PARAM.globalv.global_readin_dir, + PARAM.globalv.ks_run, + PARAM.globalv.npol); } //! Init Hamiltonian (cell changed) @@ -168,7 +179,16 @@ void ESolver_KS_PW::before_scf(UnitCell& ucell, const int istep) this->pw_wfc, this->pw_rhod, PARAM.globalv.global_out_dir, PARAM.inp); // setup psi (electronic wave functions) - this->stp.init(this->p_hamilt); + this->stp.init(this->p_hamilt, + PARAM.inp.device, + PARAM.inp.precision, + PARAM.inp.ks_solver, + PARAM.inp.bndpar, + PARAM.inp.use_k_continuity, + PARAM.inp.calculation, + PARAM.inp.mem_saver, + PARAM.globalv.npol, + PARAM.globalv.ks_run); //! Setup EXX helper for Hamiltonian and psi exx_helper->before_scf(this->p_hamilt, this->stp.template get_psi_t(), PARAM.inp); diff --git a/source/source_psi/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index 5ab1c5186ef..3ab5c4d5756 100644 --- a/source/source_psi/psi_prepare.cpp +++ b/source/source_psi/psi_prepare.cpp @@ -7,7 +7,7 @@ #include "source_base/timer.h" #include "source_base/tool_quit.h" #include "source_hsolver/diago_iter_assist.h" -#include "source_io/module_parameter/parameter.h" + #include "source_psi/psi_init_atomic.h" #include "source_psi/psi_init_atomic_random.h" #include "source_psi/psi_init_file.h" @@ -37,7 +37,18 @@ PSIPrepare::PSIPrepare(const std::string& init_wfc_in, } template -void PSIPrepare::prepare_init(const int& random_seed) +void PSIPrepare::prepare_init(const int& random_seed, + const int& nbands, + const int& nspin, + const int& nqx, + const double& dq, + const bool& domag, + const bool& domag_z, + const bool& pseudo_mesh, + const std::string& orbital_dir, + const std::string& global_readin_dir, + const bool& ks_run, + const int& npol_in) { // under restriction of C++11, std::unique_ptr can not be allocate via std::make_unique @@ -47,14 +58,14 @@ void PSIPrepare::prepare_init(const int& random_seed) if (this->init_wfc == "random") { this->psi_initer = std::unique_ptr>(new psi_init_random()); - GlobalV::ofs_running << "\n Using RANDOM starting wave functions for all " << PARAM.inp.nbands << " bands\n"; + GlobalV::ofs_running << "\n Using RANDOM starting wave functions for all " << nbands << " bands\n"; } else if (this->init_wfc == "file") { psi_init_file* file_initer = new psi_init_file(); file_initer->prepare_params( - PARAM.inp.nspin, - PARAM.globalv.global_readin_dir, + nspin, + global_readin_dir, GlobalV::RANK_IN_POOL, GlobalV::NPROC_IN_POOL ); @@ -66,7 +77,7 @@ void PSIPrepare::prepare_init(const int& random_seed) std::cout << " WARNING: init_wfc = " + this->init_wfc + " requires atomic pseudo wavefunctions(PP_PSWFC),\n but none available." " Automatically switch to random initialization." << std::endl; - GlobalV::ofs_running << "\n Using RANDOM starting wave functions for all " << PARAM.inp.nbands << " bands\n"; + GlobalV::ofs_running << "\n Using RANDOM starting wave functions for all " << nbands << " bands\n"; GlobalV::ofs_running << "\n WARNING:\n init_wfc = " + this->init_wfc + " requires atomic pseudo wavefunctions(PP_PSWFC), but none available. \n" " Automatically switch to random initialization.\n" " Note: Random starting wavefunctions may slow down convergence.\n" @@ -77,28 +88,28 @@ void PSIPrepare::prepare_init(const int& random_seed) this->psi_initer = std::unique_ptr>(new psi_init_random()); } else if (this->init_wfc == "atomic" - || (this->init_wfc == "atomic+random" && this->ucell.natomwfc < PARAM.inp.nbands)) + || (this->init_wfc == "atomic+random" && this->ucell.natomwfc < nbands)) { - if (this->ucell.natomwfc < PARAM.inp.nbands) + if (this->ucell.natomwfc < nbands) { - int nrandom = PARAM.inp.nbands - this->ucell.natomwfc; + int nrandom = nbands - this->ucell.natomwfc; GlobalV::ofs_running << "\n Using ATOMIC starting wave functions with " << this->ucell.natomwfc << " atomic orbitals" << " + " << nrandom << " random orbitals" - << " (total " << PARAM.inp.nbands << " bands)\n"; + << " (total " << nbands << " bands)\n"; } else { GlobalV::ofs_running << "\n Using ATOMIC starting wave functions for all " << this->ucell.natomwfc << " atomic orbitals" - << " (covers " << PARAM.inp.nbands << " bands)\n"; + << " (covers " << nbands << " bands)\n"; } psi_init_atomic* atomic_initer = new psi_init_atomic(); atomic_initer->prepare_params( - PARAM.globalv.nqx, - PARAM.globalv.dq, - PARAM.inp.nspin, - PARAM.globalv.domag, - PARAM.globalv.domag_z, - PARAM.inp.pseudo_mesh, + nqx, + dq, + nspin, + domag, + domag_z, + pseudo_mesh, this->sf ); this->psi_initer = std::unique_ptr>(atomic_initer); @@ -107,12 +118,12 @@ void PSIPrepare::prepare_init(const int& random_seed) { psi_init_atomic_random* atomic_rand_initer = new psi_init_atomic_random(); atomic_rand_initer->prepare_params( - PARAM.globalv.nqx, - PARAM.globalv.dq, - PARAM.inp.nspin, - PARAM.globalv.domag, - PARAM.globalv.domag_z, - PARAM.inp.pseudo_mesh, + nqx, + dq, + nspin, + domag, + domag_z, + pseudo_mesh, this->sf ); this->psi_initer = std::unique_ptr>(atomic_rand_initer); @@ -123,10 +134,10 @@ void PSIPrepare::prepare_init(const int& random_seed) { psi_init_nao* nao_initer = new psi_init_nao(); nao_initer->prepare_params( - PARAM.globalv.nqx, - PARAM.globalv.dq, - PARAM.inp.nspin, - PARAM.inp.orbital_dir, + nqx, + dq, + nspin, + orbital_dir, this->sf ); this->psi_initer = std::unique_ptr>(nao_initer); @@ -136,10 +147,10 @@ void PSIPrepare::prepare_init(const int& random_seed) { psi_init_nao_random* nao_rand_initer = new psi_init_nao_random(); nao_rand_initer->prepare_params( - PARAM.globalv.nqx, - PARAM.globalv.dq, - PARAM.inp.nspin, - PARAM.inp.orbital_dir, + nqx, + dq, + nspin, + orbital_dir, this->sf ); this->psi_initer = std::unique_ptr>(nao_rand_initer); @@ -151,7 +162,7 @@ void PSIPrepare::prepare_init(const int& random_seed) } this->psi_initer->initialize(&pw_wfc, &ucell, ik2iktot_, nkstot_, random_seed, lmaxkb, rank, - PARAM.globalv.npol, PARAM.inp.nbands); + npol_in, nbands); this->psi_initer->tabulate(); ModuleBase::timer::end("PSIPrepare", "prepare_init"); @@ -161,9 +172,18 @@ template void PSIPrepare::initialize_psi(Psi>* psi, psi::Psi* kspw_psi, hamilt::Hamilt* p_hamilt, - std::ofstream& ofs_running) + std::ofstream& ofs_running, + const std::string& device, + const std::string& precision, + const std::string& ks_solver_in, + const int& bndpar, + const bool& use_k_continuity, + const std::string& calculation, + const int& mem_saver, + const int& npol_in, + const bool& ks_run_in) { - if (kspw_psi->get_nbands() == 0 || (!PARAM.globalv.ks_run)) + if (kspw_psi->get_nbands() == 0 || (!ks_run_in)) { return; } @@ -181,18 +201,18 @@ void PSIPrepare::initialize_psi(Psi>* psi, Psi* psi_cpu = reinterpret_cast*>(psi); Psi* psi_device = kspw_psi; - bool fill = PARAM.inp.ks_solver != "bpcg" || GlobalV::MY_BNDGROUP == 0; + bool fill = ks_solver_in != "bpcg" || GlobalV::MY_BNDGROUP == 0; if (fill) { if (not_equal) { psi_cpu = new Psi(1, nbands_start, nbasis, nbasis, true); - psi_device = PARAM.inp.device == "gpu" ? new psi::Psi(psi_cpu[0]) + psi_device = device == "gpu" ? new psi::Psi(psi_cpu[0]) : reinterpret_cast*>(psi_cpu); } - else if (PARAM.inp.precision == "single") + else if (precision == "single") { - if (PARAM.inp.device == "cpu") + if (device == "cpu") { psi_cpu = reinterpret_cast*>(kspw_psi); psi_device = kspw_psi; @@ -209,7 +229,7 @@ void PSIPrepare::initialize_psi(Psi>* psi, // like (1, nbands, npwx), in which npwx is the maximal npw of all kpoints for (int ik = 0; ik < this->pw_wfc.nks; ik++) { - if(PARAM.inp.use_k_continuity && ik > 0) continue; + if(use_k_continuity && ik > 0) continue; //! Fix the wavefunction to initialize at given kpoint psi->fix_k(ik); kspw_psi->fix_k(ik); @@ -259,21 +279,21 @@ void PSIPrepare::initialize_psi(Psi>* psi, } } #ifdef __MPI - if (PARAM.inp.ks_solver == "bpcg" && PARAM.inp.bndpar > 1) + if (ks_solver_in == "bpcg" && bndpar > 1) { - std::vector sendcounts(PARAM.inp.bndpar); - std::vector displs(PARAM.inp.bndpar); + std::vector sendcounts(bndpar); + std::vector displs(bndpar); MPI_Allgather(&nbands_l, 1, MPI_INT, sendcounts.data(), 1, MPI_INT, BP_WORLD); displs[0] = 0; sendcounts[0] *= nbasis; - for (int i = 1; i < PARAM.inp.bndpar; i++) + for (int i = 1; i < bndpar; i++) { sendcounts[i] *= nbasis; displs[i] = displs[i - 1] + sendcounts[i - 1]; } if (GlobalV::MY_BNDGROUP == 0) { - for (int ip = 1; ip < PARAM.inp.bndpar; ++ip) + for (int ip = 1; ip < bndpar; ++ip) { Parallel_Common::send_data(psi_cpu->get_pointer() + displs[ip], sendcounts[ip], ip, 0, BP_WORLD); } @@ -292,12 +312,12 @@ void PSIPrepare::initialize_psi(Psi>* psi, if (not_equal) { delete psi_cpu; - if (PARAM.inp.device == "gpu") + if (device == "gpu") { delete psi_device; } } - else if (PARAM.inp.precision == "single" && PARAM.inp.device == "gpu") + else if (precision == "single" && device == "gpu") { delete psi_cpu; } @@ -322,7 +342,10 @@ void allocate_psi(Psi>*& psi, const int& nks, const std::vector& ngk, const int& nbands, - const int& npwx) + const int& npwx, + const std::string& calculation, + const int& mem_saver, + const int& npol) { assert(npwx > 0); assert(nks > 0); @@ -330,12 +353,12 @@ void allocate_psi(Psi>*& psi, delete psi; int nks2 = nks; - if (PARAM.inp.calculation == "nscf" && PARAM.inp.mem_saver == 1) + if (calculation == "nscf" && mem_saver == 1) { nks2 = 1; } - psi = new psi::Psi>(nks2, nbands, npwx * PARAM.globalv.npol, ngk, true); - const size_t memory_cost = sizeof(std::complex) * nks2 * nbands * (PARAM.globalv.npol * npwx); + psi = new psi::Psi>(nks2, nbands, npwx * npol, ngk, true); + const size_t memory_cost = sizeof(std::complex) * nks2 * nbands * (npol * npwx); std::cout << " MEMORY FOR PSI (MB) : " << static_cast(memory_cost) / 1024.0 / 1024.0 << std::endl; ModuleBase::Memory::record("Psi_PW", memory_cost); } diff --git a/source/source_psi/psi_prepare.h b/source/source_psi/psi_prepare.h index eba177b90f3..7136e15171e 100644 --- a/source/source_psi/psi_prepare.h +++ b/source/source_psi/psi_prepare.h @@ -1,5 +1,6 @@ #ifndef PSI_PREPARE_H #define PSI_PREPARE_H +#include #include "source_hamilt/hamilt.h" #include "source_psi/psi_base.h" #include "source_psi/psi_prepare_base.h" @@ -28,12 +29,32 @@ class PSIPrepare : public PSIPrepareBase ~PSIPrepare(){}; - void prepare_init(const int& random_seed); + void prepare_init(const int& random_seed, + const int& nbands, + const int& nspin, + const int& nqx, + const double& dq, + const bool& domag, + const bool& domag_z, + const bool& pseudo_mesh, + const std::string& orbital_dir, + const std::string& global_readin_dir, + const bool& ks_run, + const int& npol); void initialize_psi(Psi>* psi, psi::Psi* kspw_psi, hamilt::Hamilt* p_hamilt, - std::ofstream& ofs_running); + std::ofstream& ofs_running, + const std::string& device, + const std::string& precision, + const std::string& ks_solver_in, + const int& bndpar, + const bool& use_k_continuity, + const std::string& calculation, + const int& mem_saver, + const int& npol_in, + const bool& ks_run_in); void initialize_lcao_in_pw(Psi* psi_local, std::ofstream& ofs_running); @@ -66,7 +87,7 @@ class PSIPrepare : public PSIPrepareBase using syncmem_h2d_op = base_device::memory::synchronize_memory_op; }; -void allocate_psi(Psi>*& psi, const int& nks, const std::vector& ngk, const int& nbands, const int& npwx); +void allocate_psi(Psi>*& psi, const int& nks, const std::vector& ngk, const int& nbands, const int& npwx, const std::string& calculation, const int& mem_saver, const int& npol); } // namespace psi #endif diff --git a/source/source_psi/psi_prepare_base.h b/source/source_psi/psi_prepare_base.h index c7a10718bf9..d7f5f953fbb 100644 --- a/source/source_psi/psi_prepare_base.h +++ b/source/source_psi/psi_prepare_base.h @@ -1,22 +1,28 @@ #ifndef PSI_PREPARE_BASE_H #define PSI_PREPARE_BASE_H +#include + namespace psi { -/** - * @brief Base class for PSIPrepare without template parameters. - * - * This class provides a non-template base class for PSIPrepare, - * allowing Setup_Psi_pw to store a base class pointer instead of a template pointer. - * This is part of the gradual refactoring to remove template parameters from Setup_Psi_pw. - */ class PSIPrepareBase { public: PSIPrepareBase() = default; virtual ~PSIPrepareBase() = default; - virtual void prepare_init(const int& random_seed) = 0; + virtual void prepare_init(const int& random_seed, + const int& nbands, + const int& nspin, + const int& nqx, + const double& dq, + const bool& domag, + const bool& domag_z, + const bool& pseudo_mesh, + const std::string& orbital_dir, + const std::string& global_readin_dir, + const bool& ks_run, + const int& npol) = 0; }; } // namespace psi diff --git a/source/source_psi/setup_psi_pw.cpp b/source/source_psi/setup_psi_pw.cpp index 16be8ee790a..eab1adb361a 100644 --- a/source/source_psi/setup_psi_pw.cpp +++ b/source/source_psi/setup_psi_pw.cpp @@ -22,10 +22,22 @@ void Setup_Psi_pw::before_runner_impl( inp.ks_solver, inp.basis_type, GlobalV::MY_RANK, ucell, sf, kv.ik2iktot, kv.get_nkstot(), lmaxkb, pw_wfc); - allocate_psi(this->psi_cpu, kv.get_nks(), kv.ngk, PARAM.globalv.nbands_l, pw_wfc.npwk_max); + allocate_psi(this->psi_cpu, kv.get_nks(), kv.ngk, PARAM.globalv.nbands_l, pw_wfc.npwk_max, + inp.calculation, inp.mem_saver, PARAM.globalv.npol); auto* p_psi_init = static_cast*>(this->p_psi_init); - p_psi_init->prepare_init(inp.pw_seed); + p_psi_init->prepare_init(inp.pw_seed, + PARAM.globalv.nbands_l, + inp.nspin, + PARAM.globalv.nqx, + PARAM.globalv.dq, + PARAM.globalv.domag, + PARAM.globalv.domag_z, + inp.pseudo_mesh, + inp.orbital_dir, + PARAM.globalv.global_readin_dir, + PARAM.globalv.ks_run, + PARAM.globalv.npol); if (std::is_same::value) { @@ -151,17 +163,37 @@ void Setup_Psi_pw::update_psi_d() } template -void Setup_Psi_pw::init_impl(hamilt::Hamilt* p_hamilt) +void Setup_Psi_pw::init_impl(hamilt::Hamilt* p_hamilt, + const std::string& device, + const std::string& precision, + const std::string& ks_solver, + const int& bndpar, + const bool& use_k_continuity, + const std::string& calculation, + const int& mem_saver, + const int& npol, + const bool& ks_run) { if (!this->already_initpsi) { auto* p_psi_init = static_cast*>(this->p_psi_init); - p_psi_init->initialize_psi(this->psi_cpu, this->get_psi_t(), p_hamilt, GlobalV::ofs_running); + p_psi_init->initialize_psi(this->psi_cpu, this->get_psi_t(), p_hamilt, GlobalV::ofs_running, + device, precision, ks_solver, bndpar, use_k_continuity, + calculation, mem_saver, npol, ks_run); this->already_initpsi = true; } } -void Setup_Psi_pw::init(hamilt::HamiltBase* p_hamilt) +void Setup_Psi_pw::init(hamilt::HamiltBase* p_hamilt, + const std::string& device, + const std::string& precision, + const std::string& ks_solver, + const int& bndpar, + const bool& use_k_continuity, + const std::string& calculation, + const int& mem_saver, + const int& npol, + const bool& ks_run) { if (this->already_initpsi) { @@ -174,12 +206,16 @@ void Setup_Psi_pw::init(hamilt::HamiltBase* p_hamilt) if (this->precision_type_ == PrecisionType::ComplexFloat) { init_impl, base_device::DEVICE_GPU>( - static_cast, base_device::DEVICE_GPU>*>(p_hamilt)); + static_cast, base_device::DEVICE_GPU>*>(p_hamilt), + device, precision, ks_solver, bndpar, use_k_continuity, + calculation, mem_saver, npol, ks_run); } else { init_impl, base_device::DEVICE_GPU>( - static_cast, base_device::DEVICE_GPU>*>(p_hamilt)); + static_cast, base_device::DEVICE_GPU>*>(p_hamilt), + device, precision, ks_solver, bndpar, use_k_continuity, + calculation, mem_saver, npol, ks_run); } } else @@ -188,12 +224,16 @@ void Setup_Psi_pw::init(hamilt::HamiltBase* p_hamilt) if (this->precision_type_ == PrecisionType::ComplexFloat) { init_impl, base_device::DEVICE_CPU>( - static_cast, base_device::DEVICE_CPU>*>(p_hamilt)); + static_cast, base_device::DEVICE_CPU>*>(p_hamilt), + device, precision, ks_solver, bndpar, use_k_continuity, + calculation, mem_saver, npol, ks_run); } else { init_impl, base_device::DEVICE_CPU>( - static_cast, base_device::DEVICE_CPU>*>(p_hamilt)); + static_cast, base_device::DEVICE_CPU>*>(p_hamilt), + device, precision, ks_solver, bndpar, use_k_continuity, + calculation, mem_saver, npol, ks_run); } } } @@ -296,10 +336,14 @@ template void Setup_Psi_pw::before_runner_impl, base_device const ModulePW::PW_Basis_K&, const int&, const Input_para&); template void Setup_Psi_pw::init_impl, base_device::DEVICE_CPU>( - hamilt::Hamilt, base_device::DEVICE_CPU>*); + hamilt::Hamilt, base_device::DEVICE_CPU>*, + const std::string&, const std::string&, const std::string&, + const int&, const bool&, const std::string&, const int&, const int&, const bool&); template void Setup_Psi_pw::init_impl, base_device::DEVICE_CPU>( - hamilt::Hamilt, base_device::DEVICE_CPU>*); + hamilt::Hamilt, base_device::DEVICE_CPU>*, + const std::string&, const std::string&, const std::string&, + const int&, const bool&, const std::string&, const int&, const int&, const bool&); template void Setup_Psi_pw::update_psi_d_impl, base_device::DEVICE_CPU>(); @@ -334,10 +378,14 @@ template void Setup_Psi_pw::before_runner_impl, base_device const ModulePW::PW_Basis_K&, const int&, const Input_para&); template void Setup_Psi_pw::init_impl, base_device::DEVICE_GPU>( - hamilt::Hamilt, base_device::DEVICE_GPU>*); + hamilt::Hamilt, base_device::DEVICE_GPU>*, + const std::string&, const std::string&, const std::string&, + const int&, const bool&, const std::string&, const int&, const int&, const bool&); template void Setup_Psi_pw::init_impl, base_device::DEVICE_GPU>( - hamilt::Hamilt, base_device::DEVICE_GPU>*); + hamilt::Hamilt, base_device::DEVICE_GPU>*, + const std::string&, const std::string&, const std::string&, + const int&, const bool&, const std::string&, const int&, const int&, const bool&); template void Setup_Psi_pw::update_psi_d_impl, base_device::DEVICE_GPU>(); diff --git a/source/source_psi/setup_psi_pw.h b/source/source_psi/setup_psi_pw.h index 1d05d804f65..f1c338a4338 100644 --- a/source/source_psi/setup_psi_pw.h +++ b/source/source_psi/setup_psi_pw.h @@ -54,7 +54,16 @@ class Setup_Psi_pw const int &lmaxkb, const Input_para &inp); - void init(hamilt::HamiltBase* p_hamilt); + void init(hamilt::HamiltBase* p_hamilt, + const std::string& device, + const std::string& precision, + const std::string& ks_solver, + const int& bndpar, + const bool& use_k_continuity, + const std::string& calculation, + const int& mem_saver, + const int& npol, + const bool& ks_run); void update_psi_d(); @@ -130,7 +139,16 @@ class Setup_Psi_pw const Input_para &inp); template - void init_impl(hamilt::Hamilt* p_hamilt); + void init_impl(hamilt::Hamilt* p_hamilt, + const std::string& device, + const std::string& precision, + const std::string& ks_solver, + const int& bndpar, + const bool& use_k_continuity, + const std::string& calculation, + const int& mem_saver, + const int& npol, + const bool& ks_run); template void update_psi_d_impl(); From 7990b8ba1a2c96097cf006c18440fcde528d76a0 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Fri, 24 Jul 2026 15:04:04 +0800 Subject: [PATCH 18/22] move soc to source_psi --- .../source_lcao/module_rt/test/CMakeLists.txt | 2 +- source/source_lcao/setup_nonlocal.cpp | 4 ++-- source/source_psi/CMakeLists.txt | 1 + source/source_psi/psi_init_atomic.cpp | 5 ++-- .../soc.cpp => source_psi/soc_base.cpp} | 12 +++++----- .../soc.h => source_psi/soc_base.h} | 15 +++++------- source/source_psi/test/CMakeLists.txt | 2 +- .../test/soc_base_test.cpp} | 24 +++++++++---------- source/source_pw/module_pwdft/CMakeLists.txt | 1 - .../module_pwdft/test/CMakeLists.txt | 2 +- source/source_pw/module_pwdft/vnl_pw.h | 4 ++-- 11 files changed, 35 insertions(+), 37 deletions(-) rename source/{source_pw/module_pwdft/soc.cpp => source_psi/soc_base.cpp} (94%) rename source/{source_pw/module_pwdft/soc.h => source_psi/soc_base.h} (91%) rename source/{source_pw/module_pwdft/test/soc_test.cpp => source_psi/test/soc_base_test.cpp} (92%) diff --git a/source/source_lcao/module_rt/test/CMakeLists.txt b/source/source_lcao/module_rt/test/CMakeLists.txt index 1cc5455b003..985793690a6 100644 --- a/source/source_lcao/module_rt/test/CMakeLists.txt +++ b/source/source_lcao/module_rt/test/CMakeLists.txt @@ -52,5 +52,5 @@ AddTest( ../../../source_io/module_hs/cal_r_overlap_R.cpp ../../../source_io/module_hs/single_R_io.cpp ../../../source_io/module_hs/rr_sparse_writer.cpp - ../../../source_pw/module_pwdft/soc.cpp + ../../../source_psi/soc_base.cpp ) diff --git a/source/source_lcao/setup_nonlocal.cpp b/source/source_lcao/setup_nonlocal.cpp index 7a93b65308e..0b1b9f91b3f 100644 --- a/source/source_lcao/setup_nonlocal.cpp +++ b/source/source_lcao/setup_nonlocal.cpp @@ -4,7 +4,7 @@ #include "source_io/module_parameter/parameter.h" #ifdef __LCAO -#include "source_pw/module_pwdft/soc.h" +#include "source_psi/soc_base.h" // mohan add 2013-08-02 // In order to get rid of the read in file .NONLOCAL. @@ -53,7 +53,7 @@ void InfoNonlocal::Set_NonLocal(const int& it, { lmaxkb = std::max(lmaxkb, atom->ncpp.lll[ibeta]); } - Soc soc; + SocBase soc; if (atom->ncpp.has_so) { soc.rot_ylm(lmaxkb); diff --git a/source/source_psi/CMakeLists.txt b/source/source_psi/CMakeLists.txt index 6c10bbf3390..8bdcf5354c9 100644 --- a/source/source_psi/CMakeLists.txt +++ b/source/source_psi/CMakeLists.txt @@ -2,6 +2,7 @@ add_library( psi OBJECT psi.cpp + soc_base.cpp ) add_library( diff --git a/source/source_psi/psi_init_atomic.cpp b/source/source_psi/psi_init_atomic.cpp index 21ae88b3183..446f8a9312a 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -1,5 +1,5 @@ #include "psi_init_atomic.h" -#include "source_pw/module_pwdft/soc.h" +#include "soc_base.h" #include "source_base/math_integral.h" // for numerical integration #include "source_base/math_polyint.h" // for polynomial interpolation #include "source_base/math_ylmreal.h" // for real spherical harmonics @@ -294,7 +294,8 @@ void psi_init_atomic::init_psig(T* psig, const int& ik) { if(this->p_ucell_->atoms[it].ncpp.has_so) { - Soc soc; soc.rot_ylm(l + 1); + SocBase soc; + soc.rot_ylm(l + 1); const double j = this->p_ucell_->atoms[it].ncpp.jchi[ipswfc]; /* NOT NONCOLINEAR CASE, rotation matrix become identity */ if (!(this->domag_||this->domag_z_)) diff --git a/source/source_pw/module_pwdft/soc.cpp b/source/source_psi/soc_base.cpp similarity index 94% rename from source/source_pw/module_pwdft/soc.cpp rename to source/source_psi/soc_base.cpp index a96c7ad2351..c4b65dea990 100644 --- a/source/source_pw/module_pwdft/soc.cpp +++ b/source/source_psi/soc_base.cpp @@ -1,4 +1,4 @@ -#include "soc.h" +#include "soc_base.h" Fcoef::~Fcoef() { @@ -46,7 +46,7 @@ void Fcoef::create(const int i1, const int i2, const int i3) return; } -Soc::~Soc() +SocBase::~SocBase() { if (this->p_rot != nullptr) { @@ -54,7 +54,7 @@ Soc::~Soc() } } -double Soc::spinor(const int l, const double j, const int m, const int spin) const +double SocBase::spinor(const int l, const double j, const int m, const int spin) const { if (spin != 0 && spin != 1) ModuleBase::WARNING_QUIT("spinor", "spin direction unknown"); @@ -99,7 +99,7 @@ double Soc::spinor(const int l, const double j, const int m, const int spin) con return spinor0; } -void Soc::rot_ylm(const int lmax) +void SocBase::rot_ylm(const int lmax) { // initialize the l_max_ and l2plus1_ this->l_max_ = lmax; @@ -129,7 +129,7 @@ void Soc::rot_ylm(const int lmax) return; } -int Soc::sph_ind(const int l, const double j, const int m, const int spin) const +int SocBase::sph_ind(const int l, const double j, const int m, const int spin) const { // This function calculates the m index of the spherical harmonic // in a spinor with orbital angular momentum l, total angular @@ -181,7 +181,7 @@ int Soc::sph_ind(const int l, const double j, const int m, const int spin) const return sph_ind0; } -void Soc::set_fcoef(const int &l1, +void SocBase::set_fcoef(const int &l1, const int &l2, const int &is1, const int &is2, diff --git a/source/source_pw/module_pwdft/soc.h b/source/source_psi/soc_base.h similarity index 91% rename from source/source_pw/module_pwdft/soc.h rename to source/source_psi/soc_base.h index 40a3b2e41ac..03b595a17f1 100644 --- a/source/source_pw/module_pwdft/soc.h +++ b/source/source_psi/soc_base.h @@ -1,7 +1,7 @@ -#ifndef SOC_H -#define SOC_H +#ifndef SOC_BASE_H +#define SOC_BASE_H -#include "source_base/global_function.h" +#include "../source_base/global_function.h" #include #include @@ -42,15 +42,12 @@ class Fcoef int ind3 = 2; }; -//----------------------- -// spin-orbital coupling -//----------------------- -class Soc +class SocBase { public: - Soc(){}; - ~Soc(); + SocBase(){}; + ~SocBase(); double spinor(const int l, const double j, const int m, const int spin) const; diff --git a/source/source_psi/test/CMakeLists.txt b/source/source_psi/test/CMakeLists.txt index aed5bea543b..87c89330382 100644 --- a/source/source_psi/test/CMakeLists.txt +++ b/source/source_psi/test/CMakeLists.txt @@ -12,7 +12,7 @@ AddTest( LIBS parameter base device psi psi_init planewave SOURCES psi_init_test.cpp - ../../source_pw/module_pwdft/soc.cpp + soc_base_test.cpp ../../source_cell/atom_spec.cpp ../../source_cell/test/support/mock_unitcell.cpp ../../source_io/module_output/orb_io.cpp diff --git a/source/source_pw/module_pwdft/test/soc_test.cpp b/source/source_psi/test/soc_base_test.cpp similarity index 92% rename from source/source_pw/module_pwdft/test/soc_test.cpp rename to source/source_psi/test/soc_base_test.cpp index 9ea9bb53a92..98b13819384 100644 --- a/source/source_pw/module_pwdft/test/soc_test.cpp +++ b/source/source_psi/test/soc_base_test.cpp @@ -5,16 +5,16 @@ #include /************************************************ - * unit test of class Soc and Fcoef + * unit test of class SocBase and Fcoef ***********************************************/ /** * - Tested Functions: * - Fcoef::create to create a 5 dimensional array of complex numbers - * - Soc::set_fcoef to set the fcoef array - * - Soc::spinor to calculate the spinor - * - Soc::rot_ylm to calculate the rotation matrix - * - Soc::sph_ind to calculate the m index of the spherical harmonics + * - SocBase::set_fcoef to set the fcoef array + * - SocBase::spinor to calculate the spinor + * - SocBase::rot_ylm to calculate the rotation matrix + * - SocBase::sph_ind to calculate the m index of the spherical harmonics */ //compare two complex by using EXPECT_DOUBLE_EQ() @@ -25,7 +25,7 @@ void EXPECT_COMPLEX_DOUBLE_EQ(const std::complex& a,const std::complex Date: Fri, 24 Jul 2026 15:12:25 +0800 Subject: [PATCH 19/22] fix bug --- source/source_estate/test/elecstate_pw_test.cpp | 2 +- source/source_lcao/module_deepks/test/CMakeLists.txt | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/source/source_estate/test/elecstate_pw_test.cpp b/source/source_estate/test/elecstate_pw_test.cpp index b1e9e354c30..6342ce67e4a 100644 --- a/source/source_estate/test/elecstate_pw_test.cpp +++ b/source/source_estate/test/elecstate_pw_test.cpp @@ -8,7 +8,7 @@ #include "source_hamilt/module_xc/xc_functional.h" #include "source_pw/module_pwdft/vl_pw.h" #include "source_pw/module_pwdft/vnl_pw.h" -#include "source_pw/module_pwdft/soc.h" +#include "source_psi/soc_base.h" #include "source_io/module_parameter/parameter.h" // mock functions for testing int XC_Functional::func_type = 1; diff --git a/source/source_lcao/module_deepks/test/CMakeLists.txt b/source/source_lcao/module_deepks/test/CMakeLists.txt index 9c474f42583..dc0c4546082 100644 --- a/source/source_lcao/module_deepks/test/CMakeLists.txt +++ b/source/source_lcao/module_deepks/test/CMakeLists.txt @@ -37,7 +37,7 @@ set(DEEPKS_UNIT_COMMON_SOURCES ../../../source_cell/read_pp_blps.cpp ../../../source_cell/sep.cpp ../../../source_cell/sep_cell.cpp - ../../../source_pw/module_pwdft/soc.cpp + ../../../source_psi/soc_base.cpp ../../../source_io/module_output/sparse_matrix.cpp ../../../source_estate/read_pseudo.cpp ../../../source_estate/param_update.cpp From 51ee5ad751f4a31051d03dbbf28166e203bbeb07 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Fri, 24 Jul 2026 15:29:56 +0800 Subject: [PATCH 20/22] fix bugs --- source/source_cell/test/klist_test.cpp | 2 +- source/source_cell/test/klist_test_para.cpp | 2 +- source/source_estate/test/elecstate_pw_test.cpp | 2 +- source/source_io/test/for_testing_input_conv.h | 2 +- source/source_io/test/for_testing_klist.h | 2 +- source/source_lcao/module_deepks/test/CMakeLists.txt | 1 - 6 files changed, 5 insertions(+), 6 deletions(-) diff --git a/source/source_cell/test/klist_test.cpp b/source/source_cell/test/klist_test.cpp index 79bea7c2433..c05744bb785 100644 --- a/source/source_cell/test/klist_test.cpp +++ b/source/source_cell/test/klist_test.cpp @@ -70,7 +70,7 @@ pseudopot_cell_vnl::pseudopot_cell_vnl() pseudopot_cell_vnl::~pseudopot_cell_vnl() { } -Soc::~Soc() +SocBase::~SocBase() { } Fcoef::~Fcoef() diff --git a/source/source_cell/test/klist_test_para.cpp b/source/source_cell/test/klist_test_para.cpp index d5417cc5492..3415cadc7e7 100644 --- a/source/source_cell/test/klist_test_para.cpp +++ b/source/source_cell/test/klist_test_para.cpp @@ -73,7 +73,7 @@ pseudopot_cell_vnl::pseudopot_cell_vnl() pseudopot_cell_vnl::~pseudopot_cell_vnl() { } -Soc::~Soc() +SocBase::~SocBase() { } Fcoef::~Fcoef() diff --git a/source/source_estate/test/elecstate_pw_test.cpp b/source/source_estate/test/elecstate_pw_test.cpp index 6342ce67e4a..837ffe37ce9 100644 --- a/source/source_estate/test/elecstate_pw_test.cpp +++ b/source/source_estate/test/elecstate_pw_test.cpp @@ -115,7 +115,7 @@ void pseudopot_cell_vnl::getvnl(base_device::DE std::complex*) const { } -Soc::~Soc() +SocBase::~SocBase() { } Fcoef::~Fcoef() diff --git a/source/source_io/test/for_testing_input_conv.h b/source/source_io/test/for_testing_input_conv.h index 893f94f76c1..917a36a2f93 100644 --- a/source/source_io/test/for_testing_input_conv.h +++ b/source/source_io/test/for_testing_input_conv.h @@ -107,7 +107,7 @@ pseudopot_cell_vnl::pseudopot_cell_vnl() pseudopot_cell_vnl::~pseudopot_cell_vnl() { } -Soc::~Soc() +SocBase::~SocBase() { } Fcoef::~Fcoef() diff --git a/source/source_io/test/for_testing_klist.h b/source/source_io/test/for_testing_klist.h index 0a0e2a14c39..48850444a72 100644 --- a/source/source_io/test/for_testing_klist.h +++ b/source/source_io/test/for_testing_klist.h @@ -30,7 +30,7 @@ pseudopot_cell_vl::pseudopot_cell_vl(){} pseudopot_cell_vl::~pseudopot_cell_vl(){} pseudopot_cell_vnl::pseudopot_cell_vnl(){} pseudopot_cell_vnl::~pseudopot_cell_vnl(){} -Soc::~Soc() +SocBase::~SocBase() { } Fcoef::~Fcoef() diff --git a/source/source_lcao/module_deepks/test/CMakeLists.txt b/source/source_lcao/module_deepks/test/CMakeLists.txt index dc0c4546082..6f0c8a16683 100644 --- a/source/source_lcao/module_deepks/test/CMakeLists.txt +++ b/source/source_lcao/module_deepks/test/CMakeLists.txt @@ -37,7 +37,6 @@ set(DEEPKS_UNIT_COMMON_SOURCES ../../../source_cell/read_pp_blps.cpp ../../../source_cell/sep.cpp ../../../source_cell/sep_cell.cpp - ../../../source_psi/soc_base.cpp ../../../source_io/module_output/sparse_matrix.cpp ../../../source_estate/read_pseudo.cpp ../../../source_estate/param_update.cpp From e6a51f2c7cdff9b0c18ccaef3f1e1d1b07689d0f Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Fri, 24 Jul 2026 15:36:17 +0800 Subject: [PATCH 21/22] update soc_base --- source/source_psi/soc_base.cpp | 20 +++++++++++++++++++- source/source_psi/soc_base.h | 6 ++++-- 2 files changed, 23 insertions(+), 3 deletions(-) diff --git a/source/source_psi/soc_base.cpp b/source/source_psi/soc_base.cpp index c4b65dea990..92a1b15e066 100644 --- a/source/source_psi/soc_base.cpp +++ b/source/source_psi/soc_base.cpp @@ -15,15 +15,21 @@ void Fcoef::create(const int i1, const int i2, const int i3) ind1 = i1; ind4 = i2; ind5 = i3; + if (this->p != nullptr) { delete[] this->p; this->p = nullptr; } + int tot = ind1 * ind2 * ind3 * ind4 * ind5; + this->p = new std::complex[tot]; + for (int i = 0; i < tot; i++) + { this->p[i] = std::complex(0.0, 0.0); + } } else { @@ -57,9 +63,13 @@ SocBase::~SocBase() double SocBase::spinor(const int l, const double j, const int m, const int spin) const { if (spin != 0 && spin != 1) + { ModuleBase::WARNING_QUIT("spinor", "spin direction unknown"); + } if (m < -l - 1 || m > l) + { ModuleBase::WARNING_QUIT("spinor", "m not allowed"); + } double den = 1.0 / (2.0 * l + 1.0); // denominator @@ -86,9 +96,13 @@ double SocBase::spinor(const int l, const double j, const int m, const int spin) else { if (spin == 0) + { spinor0 = sqrt((l - m + 1.0) * den); + } if (spin == 1) + { spinor0 = -sqrt((l + m) * den); + } } } else @@ -147,9 +161,13 @@ int SocBase::sph_ind(const int l, const double j, const int m, const int spin) c if (fabs(j - l - 0.5) < 1e-8) { if (spin == 0) + { sph_ind0 = m; + } if (spin == 1) + { sph_ind0 = m + 1; + } } else if (fabs(j - l + 0.5) < 1e-8) { @@ -201,4 +219,4 @@ void SocBase::set_fcoef(const int &l1, coeff += rotylm(m1, mi) * spinor(l1, j1, m, is1) * conj(rotylm(m2, mj)) * spinor(l2, j2, m, is2); } this->fcoef(it, is1, is2, ip1, ip2) = coeff; -} \ No newline at end of file +} diff --git a/source/source_psi/soc_base.h b/source/source_psi/soc_base.h index 03b595a17f1..a29930198dd 100644 --- a/source/source_psi/soc_base.h +++ b/source/source_psi/soc_base.h @@ -54,7 +54,6 @@ class SocBase int sph_ind(const int l, const double j, const int m, const int spin) const; void rot_ylm(const int lmax); - // std::complex **rotylm; const std::complex &rotylm(const int &i1, const int &i2) const { @@ -62,6 +61,7 @@ class SocBase } Fcoef fcoef; + void set_fcoef(const int &l1, const int &l2, const int &is1, @@ -74,10 +74,12 @@ class SocBase const int &ip1, const int &ip2); - // int npol; private: + int l_max_ = -1; // maximum orbital index, must initialized before set_fcoef + int l2plus1_ = -1; + std::complex *p_rot = nullptr; }; #endif From 2583b1d37b0a5ab92abd93a412d9f3c1a3dc7b87 Mon Sep 17 00:00:00 2001 From: abacus_fixer Date: Fri, 24 Jul 2026 16:18:59 +0800 Subject: [PATCH 22/22] add a threshold for 29 test in DeePKS --- tests/09_DeePKS/29_NO_GO_deepks_scf_nspin2/threshold | 1 + 1 file changed, 1 insertion(+) create mode 100644 tests/09_DeePKS/29_NO_GO_deepks_scf_nspin2/threshold diff --git a/tests/09_DeePKS/29_NO_GO_deepks_scf_nspin2/threshold b/tests/09_DeePKS/29_NO_GO_deepks_scf_nspin2/threshold new file mode 100644 index 00000000000..dc467b56d6f --- /dev/null +++ b/tests/09_DeePKS/29_NO_GO_deepks_scf_nspin2/threshold @@ -0,0 +1 @@ +threshold 1e-6