diff --git a/source/CMakeLists.txt b/source/CMakeLists.txt index 0c1653293c..1c0ba02346 100644 --- a/source/CMakeLists.txt +++ b/source/CMakeLists.txt @@ -678,7 +678,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 acbe4e2bb8..1a8ef82957 100644 --- a/source/Makefile.Objects +++ b/source/Makefile.Objects @@ -453,7 +453,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_esolver/esolver_ks_pw.cpp b/source/source_esolver/esolver_ks_pw.cpp index 474cad9168..3916485e29 100644 --- a/source/source_esolver/esolver_ks_pw.cpp +++ b/source/source_esolver/esolver_ks_pw.cpp @@ -92,7 +92,7 @@ void ESolver_KS_PW::before_all_runners(BaseCell& basecell, const Inpu 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 00b7509990..ffed81dfe3 100644 --- a/source/source_io/module_ctrl/ctrl_output_pw.h +++ b/source/source_io/module_ctrl/ctrl_output_pw.h @@ -4,7 +4,9 @@ #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 + +class pseudopot_cell_vnl; namespace ModuleIO { 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 aad6f47a85..6955d84041 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.cpp b/source/source_io/module_wannier/to_wannier90_lcao_in_pw.cpp index 48d3e7ae65..4eaa005ce1 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,8 +44,10 @@ 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>(); - this->psi_initer_->initialize(sf_ptr, wfcpw_ptr, &ucell, &kv, 1, nullptr, GlobalV::MY_RANK); + 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; const int nks_psi = (PARAM.inp.calculation == "nscf" && PARAM.inp.mem_saver == 1)? 1 : wfcpw->nks; 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 a6f204858f..bb0100c1b1 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 c936787ba6..9223c2d616 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 5a86681c21..120aec6516 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 8be3a2ba39..6c10bbf339 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 86% rename from source/source_psi/psi_initializer.cpp rename to source/source_psi/psi_base.cpp index bed67ccd1c..55154ce3c7 100644 --- a/source/source_psi/psi_initializer.cpp +++ b/source/source_psi/psi_base.cpp @@ -1,50 +1,62 @@ -#include "psi_initializer.h" +#include "psi_base.h" +#include +#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" // basic functions support #include "source_base/timer.h" #include "source_base/tool_quit.h" // three global variables definition #include "source_base/global_variable.h" -#include "source_io/module_parameter/parameter.h" #ifdef __MPI #include "source_base/parallel_reduce.h" #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 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& lmaxkb, + const int& rank, + const int& npol, + const int& nbands) { 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->lmaxkb_ = lmaxkb; + this->npol_ = npol; + this->nbands_ = nbands; } 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); 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. 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 @@ -184,7 +196,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; @@ -211,7 +223,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++) { @@ -230,7 +242,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"); } } @@ -240,8 +252,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 67% rename from source/source_psi/psi_initializer.h rename to source/source_psi/psi_base.h index 589355d271..fb7d54f648 100644 --- a/source/source_psi/psi_initializer.h +++ b/source/source_psi/psi_base.h @@ -1,22 +1,22 @@ -#ifndef PSI_INITIALIZER_H -#define PSI_INITIALIZER_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" +#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" // for psi data structure -// smart pointer for auto-memory management +#include "source_psi/psi.h" #include -// numerical algorithm support #ifdef __MPI #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 +Psi (planewave based wavefunction) base class Auther: Kirk0830 Institute: AI for Science Institute, BEIJING @@ -24,14 +24,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: @@ -39,30 +39,33 @@ 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 - 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 int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0); //< rank + 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 + 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() @@ -112,6 +115,7 @@ class psi_initializer } protected: + #ifdef __MPI // MPI additional implementation /// @brief mapping from (ix, iy) to is void stick_to_pool(Real* stick, //< stick @@ -123,17 +127,35 @@ class psi_initializer 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 - const K_Vectors* p_kv = nullptr; ///< Parallel_Kpoints - 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 + 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 fa110f970a..ebb83357d6 100644 --- a/source/source_psi/psi_init_atomic.cpp +++ b/source/source_psi/psi_init_atomic.cpp @@ -1,35 +1,18 @@ #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() { - // 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++) @@ -38,36 +21,65 @@ 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."); + } + 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 = 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::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 int& random_seed, //< random seed - const pseudopot_cell_vnl* p_pspot_nl, - const int& rank) +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; + this->params_prepared_ = true; +} + +template +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, + const int& lmaxkb, + const int& rank, + const int& npol, + const int& nbands) { ModuleBase::timer::start("psi_init_atomic", "initialize"); - if(p_pspot_nl == nullptr) + if (!this->params_prepared_) { - ModuleBase::WARNING_QUIT("psi_init_atomic::initialize", - "pseudopot_cell_vnl object cannot be nullptr for atomic, quit."); + ModuleBase::WARNING_QUIT("psi_init_atomic::initialize", + "prepare_params() must be called before initialize()"); } - // 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_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); @@ -90,7 +102,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; @@ -98,17 +110,17 @@ 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++) { - 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; + 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++) @@ -201,9 +213,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++) { @@ -232,12 +244,13 @@ 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; 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); @@ -273,17 +286,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++) @@ -297,7 +310,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++) { @@ -336,17 +349,17 @@ void psi_init_atomic::init_psig(T* psig, const int& ik) int ipswfc_noncolin_soc=0; /* J = L - 1/2 -> continue */ /* J = L + 1/2 */ - if(fabs(j - l + 0.5) < 1e-4) - { - continue; - } - chiaux.clear(); - chiaux.resize(npw); + if(fabs(j - l + 0.5) < 1e-4) + { + continue; + } + chiaux.clear(); + chiaux.resize(npw); /* L == 0 */ - if(l == 0) - { - std::memcpy(chiaux.data(), ovlp_pswfcjlg.data(), npw * sizeof(double)); - } + if(l == 0) + { + std::memcpy(chiaux.data(), ovlp_pswfcjlg.data(), npw * sizeof(double)); + } else { /* L != 0, scan pswfcs that have the same L and satisfy J(pswfc) = L - 0.5 */ @@ -366,7 +379,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); } @@ -472,14 +485,14 @@ void psi_init_atomic::init_psig(T* psig, const int& ik) } } } - delete [] sk; + delete [] sk; } } - /* complement the rest of bands if there are */ - if(this->nbands_complem() > 0) - { - this->random_t(psig, index, this->nbands_start_, ik); - } + /* complement the rest of bands if there are */ + if(this->nbands_complem() > 0) + { + this->random_t(psig, index, this->nbands_start_, ik); + } ModuleBase::timer::end("psi_init_atomic", "init_psig"); } @@ -487,4 +500,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 4cdfabdc6c..9c349b9677 100644 --- a/source/source_psi/psi_init_atomic.h +++ b/source/source_psi/psi_init_atomic.h @@ -1,16 +1,25 @@ #ifndef PSI_INIT_ATOMIC_H #define PSI_INIT_ATOMIC_H +#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; + int nqx_ = 0; + double dq_ = 0.0; + int nspin_ = 1; + bool domag_ = false; + bool domag_z_ = false; + bool pseudo_mesh_ = false; + bool params_prepared_ = false; public: psi_init_atomic() @@ -19,21 +28,48 @@ class psi_init_atomic : public psi_initializer } ~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 K_Vectors*, //< kpoints - const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + /** + * @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); + + 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; 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.cpp b/source/source_psi/psi_init_atomic_random.cpp index 91fd4c10bc..e8d848a244 100644 --- a/source/source_psi/psi_init_atomic_random.cpp +++ b/source/source_psi/psi_init_atomic_random.cpp @@ -1,17 +1,18 @@ #include "psi_init_atomic_random.h" -#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 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& lmaxkb, + 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, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); } template @@ -19,7 +20,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_atomic_random.h b/source/source_psi/psi_init_atomic_random.h index 2c8e49fc8d..47b20f199a 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" /* @@ -20,17 +19,19 @@ class psi_init_atomic_random : public psi_init_atomic } ~psi_init_atomic_random(){}; - /// @brief initialize the psi_initializer 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 K_Vectors*, //< kpoints - const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + 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; private: }; -#endif \ No newline at end of file +#endif diff --git a/source/source_psi/psi_init_file.cpp b/source/source_psi/psi_init_file.cpp index 2e633573bb..95d13caae4 100644 --- a/source/source_psi/psi_init_file.cpp +++ b/source/source_psi/psi_init_file.cpp @@ -1,48 +1,66 @@ #include "psi_init_file.h" +#include +#include +#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" 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& lmaxkb, + 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_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); + this->nbands_start_ = nbands; 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) { 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->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; - 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", - ik,this->p_kv->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, PARAM.globalv.npol, + ModuleIO::read_wfc_pw(fn, this->pw_wfc_, + 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 72fb18ed1e..f9f497d75e 100644 --- a/source/source_psi/psi_init_file.h +++ b/source/source_psi/psi_init_file.h @@ -1,17 +1,22 @@ #ifndef PSI_INIT_FILE_H #define PSI_INIT_FILE_H -#include "source_pw/module_pwdft/vnl_pw.h" -#include "psi_initializer.h" +#include +#include +#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; + int nspin_ = 1; + std::string global_readin_dir_; + int rank_in_pool_ = 0; + int nproc_in_pool_ = 1; public: psi_init_file() @@ -20,18 +25,25 @@ class psi_init_file : public psi_initializer }; ~psi_init_file(){}; - /// @brief initialize the psi_initializer 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 K_Vectors*, //< kpoints - const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + 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 /// @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_init_nao.cpp b/source/source_psi/psi_init_nao.cpp index e514ae063d..7651dd8168 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 @@ -42,6 +40,19 @@ 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; + this->params_prepared_ = true; +} + template void psi_init_nao::read_external_orbs(const std::string* orbital_files, const int& rank) { @@ -67,7 +78,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 @@ -150,15 +161,24 @@ 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) + const int& lmaxkb, + const int& rank, + const int& npol, + const int& nbands) { 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_initializer::initialize(sf, pw_wfc, p_ucell, p_kv_in, random_seed, p_pspot_nl, rank); + psi_base::initialize(sf, pw_wfc, p_ucell, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); // allocate this->allocate_ao_table(); @@ -177,18 +197,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"); @@ -199,10 +211,16 @@ 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(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()) @@ -224,7 +242,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++) { @@ -237,7 +255,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()); @@ -258,7 +276,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); @@ -299,7 +317,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 bb05962ec0..be074c4bba 100644 --- a/source/source_psi/psi_init_nao.h +++ b/source/source_psi/psi_init_nao.h @@ -3,17 +3,23 @@ #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 +#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; + int nqx_ = 0; + double dq_ = 0.0; + int nspin_ = 1; + std::string orbital_dir_; + bool params_prepared_ = false; public: psi_init_nao() @@ -22,63 +28,96 @@ class psi_init_nao : public psi_initializer }; ~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_initializer 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 K_Vectors*, //< kpoints - const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + 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); + virtual void tabulate() override; + + virtual void init_psig(T* psig, const int& ik) 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]; @@ -101,4 +140,4 @@ class psi_init_nao : public psi_initializer /// @brief useful for atomic-like methods ModuleBase::SphericalBesselTransformer sbt; }; -#endif \ No newline at end of file +#endif diff --git a/source/source_psi/psi_init_nao_random.cpp b/source/source_psi/psi_init_nao_random.cpp index e3e2b8e89c..38b8e78a47 100644 --- a/source/source_psi/psi_init_nao_random.cpp +++ b/source/source_psi/psi_init_nao_random.cpp @@ -1,17 +1,18 @@ #include "psi_init_nao_random.h" -#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, 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& lmaxkb, + 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, ik2iktot, nkstot, random_seed, lmaxkb, rank, npol, nbands); } template @@ -19,7 +20,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/psi_init_nao_random.h b/source/source_psi/psi_init_nao_random.h index d6b38bb556..55e18d46a7 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" /* @@ -20,14 +19,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 K_Vectors*, //< kpoints - const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + 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.cpp b/source/source_psi/psi_init_random.cpp index 697a4476ad..21668bc29b 100644 --- a/source/source_psi/psi_init_random.cpp +++ b/source/source_psi/psi_init_random.cpp @@ -1,20 +1,23 @@ #include "psi_init_random.h" -#include "source_io/module_parameter/parameter.h" +#include 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& lmaxkb, + 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_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()); - 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 8d66035ec6..53a3adfc1a 100644 --- a/source/source_psi/psi_init_random.h +++ b/source/source_psi/psi_init_random.h @@ -1,14 +1,14 @@ #ifndef PSI_INIT_RANDOM_H #define PSI_INIT_RANDOM_H -#include "source_pw/module_pwdft/vnl_pw.h" -#include "psi_initializer.h" +#include +#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; @@ -23,13 +23,15 @@ class psi_init_random : public psi_initializer /// @param ik kpoint index /// @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 K_Vectors*, //< kpoints - const int& = 1, //< random seed - const pseudopot_cell_vnl* = nullptr, //< nonlocal pseudopotential - const int& = 0) override; //< MPI rank + 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/psi_prepare.cpp b/source/source_psi/psi_prepare.cpp index eada5f13d1..7c0227bcda 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 { @@ -24,10 +25,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 pseudopot_cell_vnl& nlpp_in, + const std::vector& ik2iktot_in, + const int& nkstot_in, + const int& lmaxkb_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), 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; @@ -44,12 +46,19 @@ void PSIPrepare::prepare_init(const int& random_seed, const int istep 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()); + 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)) @@ -71,7 +80,7 @@ void PSIPrepare::prepare_init(const int& random_seed, const int istep << std::endl; } GlobalV::ofs_running << "\n Using RANDOM starting wave functions for all " << PARAM.inp.nbands << " bands\n"; - 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)) @@ -88,22 +97,54 @@ void PSIPrepare::prepare_init(const int& random_seed, const int istep 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"; } 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 @@ -111,7 +152,8 @@ void PSIPrepare::prepare_init(const int& random_seed, const int istep 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, lmaxkb, rank, + PARAM.globalv.npol, PARAM.inp.nbands); this->psi_initer->tabulate(); ModuleBase::timer::end("PSIPrepare", "prepare_init"); diff --git a/source/source_psi/psi_prepare.h b/source/source_psi/psi_prepare.h index 20ec2eeb1b..6e476b64ad 100644 --- a/source/source_psi/psi_prepare.h +++ b/source/source_psi/psi_prepare.h @@ -1,13 +1,12 @@ #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 { -// This class is used to prepare the wavefunction template class PSIPrepare : public PSIPrepareBase { @@ -18,9 +17,11 @@ class PSIPrepare : public PSIPrepareBase const int& rank, const UnitCell& ucell, const Structure_Factor& sf, - const K_Vectors& kv_in, - const pseudopot_cell_vnl& nlpp, + const std::vector& ik2iktot, + const int& nkstot, + const int& lmaxkb, const ModulePW::PW_Basis_K& pw_wfc); + ~PSIPrepare(){}; ///@brief prepare the wavefunction initialization @@ -29,14 +30,13 @@ class PSIPrepare : public PSIPrepareBase /// printed on the first step to avoid spamming relax output void prepare_init(const int& random_seed, const int istep); - //------------------------ only for psi_initializer -------------------- /** * @brief initialize the wavefunction * * @param psi store the wavefunction + * @param kspw_psi Kohn-Sham wavefunction in plane-wave basis * @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, @@ -49,13 +49,10 @@ class PSIPrepare : public PSIPrepareBase */ void initialize_lcao_in_pw(Psi* psi_local, std::ofstream& ofs_running); - // psi_initializer* 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 std::string init_wfc = "none"; @@ -68,8 +65,10 @@ class PSIPrepare : public PSIPrepareBase // pw basis const ModulePW::PW_Basis_K& pw_wfc; - // parallel kpoints - const K_Vectors& kv; + // local->global k-point mapping + const std::vector& ik2iktot_; + // total number of k-points + const int nkstot_; // unit cell const UnitCell& ucell; @@ -77,20 +76,18 @@ class PSIPrepare : public PSIPrepareBase // structure factor const Structure_Factor& sf; - // nonlocal pseudopotential - const pseudopot_cell_vnl& nlpp; + // max angular momentum for non-local projectors + const int lmaxkb; Device* ctx = {}; ///< device base_device::DEVICE_CPU* cpu_ctx = {}; ///< CPU device const int rank; ///< MPI 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 ba658a02de..cdb7a0c247 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 @@ -12,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) @@ -50,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 a4c0f11d3e..5ce3f1f2fe 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); diff --git a/source/source_psi/setup_psi_pw.cpp b/source/source_psi/setup_psi_pw.cpp index 816842d444..f75c8730fe 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,12 +15,12 @@ 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, 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); @@ -26,25 +30,38 @@ void Setup_Psi_pw::before_runner_impl( // printed from this initial setup call. p_psi_init->prepare_init(inp.pw_seed, 0); - 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)); } } @@ -54,30 +71,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); } } } @@ -91,10 +115,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)); } } @@ -206,13 +232,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); } @@ -266,11 +292,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>*); @@ -287,16 +313,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>; @@ -304,11 +330,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>*); @@ -329,14 +355,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 88e9d42bf1..de1cf8c8fc 100644 --- a/source/source_psi/setup_psi_pw.h +++ b/source/source_psi/setup_psi_pw.h @@ -6,7 +6,6 @@ #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" @@ -39,7 +38,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; //------------ @@ -51,7 +50,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 +70,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 +125,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 +141,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/CMakeLists.txt b/source/source_psi/test/CMakeLists.txt index 63af5799e1..aed5bea543 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 52% rename from source/source_psi/test/psi_initializer_unit_test.cpp rename to source/source_psi/test/psi_init_test.cpp index 342e04a5a6..ed7ceb5023 100644 --- a/source/source_psi/test/psi_initializer_unit_test.cpp +++ b/source/source_psi/test/psi_init_test.cpp @@ -1,55 +1,60 @@ #include -#define private public -#include "source_io/module_parameter/parameter.h" -#undef private -#include "../psi_initializer.h" +#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_cell/klist.h" #include "source_base/output.h" /* ========================= -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 */ @@ -64,10 +69,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() {} @@ -86,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; } @@ -96,13 +99,23 @@ 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; - K_Vectors* p_kv = nullptr; + int lmaxkb = 0; + std::vector ik2iktot_; + int nkstot_ = 0; int random_seed = 1; - psi_initializer>* psi_init; + psi_base>* psi_init; + + 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; - private: protected: void SetUp() override { @@ -110,19 +123,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(); - this->p_kv = new K_Vectors(); - // 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}; @@ -165,19 +165,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; -} + 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 < 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; @@ -190,6 +199,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(); @@ -201,64 +211,93 @@ 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; -} + 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; -} + 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->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->gcar != nullptr) { delete[] this->p_pw_wfc->gcar; -} + 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; -} + 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}; - this->p_pspot_vnl->lmaxkb = 1; + this->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 @@ -267,37 +306,41 @@ class PsiIntializerUnitTest : public ::testing::Test { delete this->p_sf; delete this->p_pw_wfc; delete this->p_ucell; - delete this->p_pspot_vnl; - delete this->p_kv; } }; -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}; @@ -309,224 +352,282 @@ TEST_F(PsiIntializerUnitTest, CastToT) { EXPECT_EQ(this->psi_init->template cast_to_T(cd), f); } -TEST_F(PsiIntializerUnitTest, CalPsigRandom) { - PARAM.input.init_wfc = "random"; +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->p_kv, + this->psi_init->initialize(this->p_sf, + this->p_pw_wfc, + this->p_ucell, + this->ik2iktot_, + this->nkstot_, this->random_seed, - this->p_pspot_vnl, - GlobalV::MY_RANK); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->lmaxkb, + GlobalV::MY_RANK, + 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); delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigAtomic) { - PARAM.input.init_wfc = "atomic"; - this->psi_init = new psi_init_atomic>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->p_kv, +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->p_pspot_vnl, - GlobalV::MY_RANK); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->lmaxkb, + GlobalV::MY_RANK, + 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); delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigAtomicSoc) { - PARAM.input.init_wfc = "atomic"; - PARAM.input.nspin = 4; - PARAM.sys.npol = 2; +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>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->p_kv, + 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->p_pspot_vnl, - GlobalV::MY_RANK); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->lmaxkb, + GlobalV::MY_RANK, + 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; +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>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->p_kv, + 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->p_pspot_vnl, - GlobalV::MY_RANK); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->lmaxkb, + GlobalV::MY_RANK, + 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>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->p_kv, +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->p_pspot_vnl, - GlobalV::MY_RANK); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->lmaxkb, + GlobalV::MY_RANK, + 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); delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigNao) { - PARAM.input.init_wfc = "nao"; - this->psi_init = new psi_init_nao>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->p_kv, +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->p_pspot_vnl, - GlobalV::MY_RANK); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->lmaxkb, + GlobalV::MY_RANK, + 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); delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigNaoRandom) { - PARAM.input.init_wfc = "nao+random"; - this->psi_init = new psi_init_nao_random>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->p_kv, +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->p_pspot_vnl, - GlobalV::MY_RANK); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->lmaxkb, + GlobalV::MY_RANK, + 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); delete psi; } -TEST_F(PsiIntializerUnitTest, CalPsigNaoSoc) { - PARAM.input.init_wfc = "nao"; - PARAM.input.nspin = 4; - PARAM.sys.npol = 2; +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; - PARAM.sys.domag = false; - PARAM.sys.domag_z = false; - this->psi_init = new psi_init_nao>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->p_kv, + 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->p_pspot_vnl, - GlobalV::MY_RANK); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->lmaxkb, + GlobalV::MY_RANK, + 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; +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; - PARAM.sys.domag = false; - PARAM.sys.domag_z = false; - this->psi_init = new psi_init_nao>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->p_kv, + 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->p_pspot_vnl, - GlobalV::MY_RANK); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->lmaxkb, + GlobalV::MY_RANK, + 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; +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; - PARAM.sys.domag = true; - PARAM.sys.domag_z = false; - this->psi_init = new psi_init_nao>(); - this->psi_init->initialize(this->p_sf, - this->p_pw_wfc, - this->p_ucell, - this->p_kv, + 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->p_pspot_vnl, - GlobalV::MY_RANK); - this->psi_init->tabulate(); // always: new, initialize, tabulate, allocate, proj_ao_onkG + this->lmaxkb, + GlobalV::MY_RANK, + 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; } @@ -547,4 +648,4 @@ int main(int argc, char** argv) #endif return result; -} \ No newline at end of file +} diff --git a/source/source_psi/test/support/atomic_new b/source/source_psi/test/support/atomic_new deleted file mode 100644 index af7a7d973a..0000000000 --- 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 b0b7789a2d..0000000000 --- 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 87a70829dd..0000000000 --- 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