diff --git a/source/Makefile.Objects b/source/Makefile.Objects index 6006ef304ef..a91d959c86b 100644 --- a/source/Makefile.Objects +++ b/source/Makefile.Objects @@ -346,7 +346,7 @@ OBJS_GINT=batch_biggrid.o\ OBJS_HAMILT=hamilt_pw.o\ hs_matrix_k.o\ - hamilt_sdft_pw.o\ + sto_hamilt_pw.o\ operator.o\ op_pw.o\ op_pw_exx.o\ @@ -443,7 +443,6 @@ OBJS_HSOLVER=diago_cg.o\ hsolver.o\ hsolver_pw.o\ hsolver_lcaopw.o\ - hsolver_pw_sdft.o\ diago_iter_assist.o\ hegvd_op.o\ bpcg_kernel_op.o\ @@ -881,6 +880,7 @@ OBJS_SRCPW=h_ewald_pw.o\ stru_fac.o\ stru_fac_k.o\ soc.o\ + sto_hsolver_pw.o\ sto_iter.o\ sto_che.o\ sto_wf.o\ diff --git a/source/source_esolver/esolver_sdft_pw.cpp b/source/source_esolver/esolver_sdft_pw.cpp index 34f4b48c72d..89b156754a1 100644 --- a/source/source_esolver/esolver_sdft_pw.cpp +++ b/source/source_esolver/esolver_sdft_pw.cpp @@ -7,6 +7,7 @@ #include "source_hsolver/diago_iter_assist.h" #include "source_hsolver/diago_params.h" #include "source_io/module_parameter/parameter.h" +#include "source_pw/module_stodft/sto_hsolver_pw.h" #include "source_pw/module_stodft/sto_dos.h" #include "source_pw/module_stodft/sto_elecond.h" #include "source_pw/module_stodft/sto_forces.h" @@ -100,15 +101,15 @@ void ESolver_SDFT_PW::before_scf(UnitCell& ucell, const int istep) ESolver_KS_PW::before_scf(ucell, istep); delete reinterpret_cast*>(this->p_hamilt); - this->p_hamilt = new hamilt::HamiltSdftPW(this->pelec->pot, - this->pw_wfc, - &this->kv, - &this->ppcell, - &ucell, - PARAM.globalv.npol, - &this->stoche.emin_sto, - &this->stoche.emax_sto); - this->p_hamilt_sto = static_cast*>(this->p_hamilt); + this->p_hamilt = new StoHamiltPW(this->pelec->pot, + this->pw_wfc, + &this->kv, + &this->ppcell, + &ucell, + PARAM.globalv.npol, + &this->stoche.emin_sto, + &this->stoche.emax_sto); + this->p_hamilt_sto = static_cast*>(this->p_hamilt); if (istep > 0 && this->inp_->nbands_sto != 0 && this->inp_->initsto_freq > 0 && istep % this->inp_->initsto_freq == 0) { @@ -153,43 +154,43 @@ void ESolver_SDFT_PW::hamilt2rho_single(UnitCell& ucell, int istep, i bool skip_charge = this->inp_->calculation == "nscf" ? true : false; // hsolver only exists in this function - hsolver::HSolverPW_SDFT hsolver_pw_sdft_obj(&this->kv, - this->pw_wfc, - this->stowf, - this->stoche, - this->p_hamilt_sto, - this->inp_->calculation, - this->inp_->basis_type, - this->inp_->ks_solver, - PARAM.globalv.use_uspp, - this->inp_->nspin, - hsolver::DiagoIterAssist::SCF_ITER, - hsolver::DiagoIterAssist::PW_DIAG_NMAX, - hsolver::DiagoIterAssist::PW_DIAG_THR, - hsolver::DiagoIterAssist::need_subspace, - this->inp_->nbands, - this->inp_->diago_smooth_ethr, - this->inp_->pw_diag_ndim, - this->inp_->diag_subspace, - this->inp_->nb2d, - PARAM.globalv.ks_run, - PARAM.globalv.all_ks_run, - this->inp_->bndpar); - - hsolver_pw_sdft_obj.solve(ucell, - static_cast*>(this->p_hamilt), - *this->stp.template get_psi_t(), - this->stp.psi_cpu[0], - this->pelec, - this->pw_wfc, - this->stowf, - istep, - iter, - GlobalV::ofs_running, - skip_charge); + StoHSolverPW sto_hsolver_pw_obj(&this->kv, + this->pw_wfc, + this->stowf, + this->stoche, + this->p_hamilt_sto, + this->inp_->calculation, + this->inp_->basis_type, + this->inp_->ks_solver, + PARAM.globalv.use_uspp, + this->inp_->nspin, + hsolver::DiagoIterAssist::SCF_ITER, + hsolver::DiagoIterAssist::PW_DIAG_NMAX, + hsolver::DiagoIterAssist::PW_DIAG_THR, + hsolver::DiagoIterAssist::need_subspace, + this->inp_->nbands, + this->inp_->diago_smooth_ethr, + this->inp_->pw_diag_ndim, + this->inp_->diag_subspace, + this->inp_->nb2d, + PARAM.globalv.ks_run, + PARAM.globalv.all_ks_run, + this->inp_->bndpar); + + sto_hsolver_pw_obj.solve(ucell, + static_cast*>(this->p_hamilt), + *this->stp.template get_psi_t(), + this->stp.psi_cpu[0], + this->pelec, + this->pw_wfc, + this->stowf, + istep, + iter, + GlobalV::ofs_running, + skip_charge); // set_diagethr need it - this->esolver_KS_ne = hsolver_pw_sdft_obj.stoiter.KS_ne; + this->esolver_KS_ne = sto_hsolver_pw_obj.stoiter.KS_ne; if (PARAM.globalv.ks_run) { diff --git a/source/source_esolver/esolver_sdft_pw.h b/source/source_esolver/esolver_sdft_pw.h index aa48f97db7e..2a5c4965c3a 100644 --- a/source/source_esolver/esolver_sdft_pw.h +++ b/source/source_esolver/esolver_sdft_pw.h @@ -2,7 +2,7 @@ #define ESOLVER_SDFT_PW_H #include "esolver_ks_pw.h" -#include "source_pw/module_stodft/hamilt_sdft_pw.h" +#include "source_pw/module_stodft/sto_hamilt_pw.h" #include "source_pw/module_stodft/sto_che.h" #include "source_pw/module_stodft/sto_iter.h" #include "source_pw/module_stodft/sto_wf.h" @@ -31,7 +31,7 @@ class ESolver_SDFT_PW : public ESolver_KS_PW public: Stochastic_WF stowf; StoChe stoche; - hamilt::HamiltSdftPW* p_hamilt_sto = nullptr; + StoHamiltPW* p_hamilt_sto = nullptr; protected: virtual void before_scf(UnitCell& ucell, const int istep) override; diff --git a/source/source_hsolver/CMakeLists.txt b/source/source_hsolver/CMakeLists.txt index a9f1f9142ec..1285b915e8e 100644 --- a/source/source_hsolver/CMakeLists.txt +++ b/source/source_hsolver/CMakeLists.txt @@ -7,7 +7,6 @@ list(APPEND objects para_lin_tf.cpp hsolver_pw.cpp hsolver_lcaopw.cpp - hsolver_pw_sdft.cpp diago_iter_assist.cpp hsolver.cpp diago_pxxxgvx.cpp diff --git a/source/source_hsolver/hsolver_pw_sdft.h b/source/source_hsolver/hsolver_pw_sdft.h deleted file mode 100644 index 7c88a5a2b3b..00000000000 --- a/source/source_hsolver/hsolver_pw_sdft.h +++ /dev/null @@ -1,84 +0,0 @@ -#ifndef HSOLVERPW_SDFT_H -#define HSOLVERPW_SDFT_H -#include "hsolver_pw.h" -#include "source_pw/module_stodft/hamilt_sdft_pw.h" -#include "source_pw/module_stodft/sto_iter.h" -namespace hsolver -{ -template -class HSolverPW_SDFT : public HSolverPW -{ - protected: - using Real = typename GetTypeReal::type; - - public: - HSolverPW_SDFT(K_Vectors* pkv, - ModulePW::PW_Basis_K* wfc_basis_in, - Stochastic_WF& stowf, - StoChe& stoche, - hamilt::HamiltSdftPW* p_hamilt_sto, - const std::string calculation_type_in, - const std::string basis_type_in, - const std::string method_in, - const bool use_uspp_in, - const int nspin_in, - const int scf_iter_in, - const int diag_iter_max_in, - const double diag_thr_in, - const bool need_subspace_in, - const int nbands_in, - const bool diago_smooth_ethr_in, - const int pw_diag_ndim_in, - const int diag_subspace_in, - const int nb2d_in, - const bool ks_run_in, - const bool all_ks_run_in, - const int bndpar_in) - : HSolverPW(wfc_basis_in, - calculation_type_in, - basis_type_in, - method_in, - use_uspp_in, - nspin_in, - scf_iter_in, - diag_iter_max_in, - diag_thr_in, - need_subspace_in, - nbands_in, - diago_smooth_ethr_in, - pw_diag_ndim_in, - diag_subspace_in, - nb2d_in), - ks_run(ks_run_in), all_ks_run(all_ks_run_in), bndpar(bndpar_in) - { - stoiter.init(pkv, wfc_basis_in, stowf, stoche, p_hamilt_sto); - } - - void solve(const UnitCell& ucell, - hamilt::Hamilt* pHamilt, - psi::Psi& psi, - psi::Psi& psi_cpu, - elecstate::ElecState* pes, - ModulePW::PW_Basis_K* wfc_basis, - Stochastic_WF& stowf, - const int istep, - const int iter, - std::ostream& log, - const bool skip_charge); - - Stochastic_Iter stoiter; - - protected: - const bool ks_run; // true if the current process runs the KS part of the SDFT calculation - const bool all_ks_run; // true if every process runs the KS part - const int bndpar; // number of band-parallel groups - - using setmem_complex_op = base_device::memory::set_memory_op; - using setmem_var_op = base_device::memory::set_memory_op; - using syncmem_h2d_op = base_device::memory::synchronize_memory_op; - using syncmem_d2h_op = base_device::memory::synchronize_memory_op; - using syncmem_var_h2d_op = base_device::memory::synchronize_memory_op; - using syncmem_var_d2h_op = base_device::memory::synchronize_memory_op; -}; -} // namespace hsolver -#endif diff --git a/source/source_hsolver/test/CMakeLists.txt b/source/source_hsolver/test/CMakeLists.txt index e387be12a7f..ae6001ecea0 100644 --- a/source/source_hsolver/test/CMakeLists.txt +++ b/source/source_hsolver/test/CMakeLists.txt @@ -67,13 +67,6 @@ if (ENABLE_MPI) ../../source_cell/klist.cpp ../../source_cell/klist_io.cpp ../../source_cell/parallel_kpoints.cpp ../../source_cell/reciprocal_grid.cpp ) - AddTest( - TARGET MODULE_HSOLVER_sdft - LIBS parameter psi device base container - SOURCES test_hsolver_sdft.cpp ../hsolver_pw_sdft.cpp ../hsolver_pw.cpp ../diago_bpcg.cpp ../diago_dav_subspace.cpp ../diag_const_nums.cpp ../diago_iter_assist.cpp ../para_lin_tf.cpp - ../../source_estate/elecstate_tools.cpp ../../source_estate/occupy.cpp ../../source_base/module_fft/fft_bundle.cpp ../../source_base/module_fft/fft_cpu.cpp - ) - if(ENABLE_LCAO) if(TARGET ELPA::ELPA) AddTest( diff --git a/source/source_pw/module_stodft/CMakeLists.txt b/source/source_pw/module_stodft/CMakeLists.txt index b0e160c5b21..e284ca8ad17 100644 --- a/source/source_pw/module_stodft/CMakeLists.txt +++ b/source/source_pw/module_stodft/CMakeLists.txt @@ -1,5 +1,6 @@ list(APPEND hamilt_stodft_srcs - hamilt_sdft_pw.cpp + sto_hamilt_pw.cpp + sto_hsolver_pw.cpp sto_iter.cpp sto_che.cpp sto_wf.cpp diff --git a/source/source_pw/module_stodft/hamilt_sdft_pw.cpp b/source/source_pw/module_stodft/hamilt_sdft_pw.cpp deleted file mode 100644 index 90151972285..00000000000 --- a/source/source_pw/module_stodft/hamilt_sdft_pw.cpp +++ /dev/null @@ -1,72 +0,0 @@ -#include "hamilt_sdft_pw.h" -#include "source_base/timer.h" -#include "kernels/hpsi_norm_op.h" - -namespace hamilt -{ - -template -HamiltSdftPW::HamiltSdftPW(elecstate::Potential* pot_in, - ModulePW::PW_Basis_K* wfc_basis, - K_Vectors* p_kv, - pseudopot_cell_vnl* nlpp, - const UnitCell* ucell, - const int& npol, - Real* emin_in, - Real* emax_in) - : HamiltPW(pot_in, wfc_basis, p_kv, nlpp, nullptr, ucell, nullptr), ngk(p_kv->ngk) -{ - this->classname = "HamiltSdftPW"; - this->npwk_max = wfc_basis->npwk_max; - this->npol = npol; - this->emin = emin_in; - this->emax = emax_in; -} - -template -void HamiltSdftPW::hPsi(const T* psi_in, T* hpsi, const int& nbands) -{ - auto call_act = [&, this](const Operator* op, const bool& is_first_node) -> void { - op->act(nbands, this->npwk_max, this->npol, psi_in, hpsi, this->ngk[op->get_ik()], is_first_node); - }; - - ModuleBase::timer::start("HamiltSdftPW", "hPsi"); - call_act(this->ops, true); // first node - Operator* node((Operator*)this->ops->next_op); - while (node != nullptr) - { - call_act(node, false); // other nodes - node = (Operator*)(node->next_op); - } - ModuleBase::timer::end("HamiltSdftPW", "hPsi"); - - return; -} - -template -void HamiltSdftPW::hPsi_norm(const T* psi_in, T* hpsi_norm, const int& nbands) -{ - ModuleBase::timer::start("HamiltSdftPW", "hPsi_norm"); - - this->hPsi(psi_in, hpsi_norm, nbands); - - const int ik = this->ops->get_ik(); - const int npwk_max = this->npwk_max; - const int npwk = this->ngk[ik]; - const Real emin = *this->emin; - const Real emax = *this->emax; - const Real Ebar = (emin + emax) / 2; - const Real DeltaE = (emax - emin) / 2; - - hpsi_norm_op()(this->ctx, nbands, npwk_max, npwk, Ebar, DeltaE, hpsi_norm, psi_in); - ModuleBase::timer::end("HamiltSdftPW", "hPsi_norm"); -} - -template class HamiltSdftPW, base_device::DEVICE_CPU>; -template class HamiltSdftPW, base_device::DEVICE_CPU>; -#if ((defined __CUDA) || (defined __ROCM)) -template class HamiltSdftPW, base_device::DEVICE_GPU>; -template class HamiltSdftPW, base_device::DEVICE_GPU>; -#endif - -} // namespace hamilt diff --git a/source/source_pw/module_stodft/sto_dos.cpp b/source/source_pw/module_stodft/sto_dos.cpp index 884679a8b83..7d9193d3f98 100644 --- a/source/source_pw/module_stodft/sto_dos.cpp +++ b/source/source_pw/module_stodft/sto_dos.cpp @@ -25,7 +25,7 @@ Sto_DOS::Sto_DOS(ModulePW::PW_Basis_K* p_wfcpw_in, this->p_elec = p_elec_in; this->p_psi = p_psi_in; this->p_hamilt = p_hamilt_in; - this->p_hamilt_sto = static_cast>*>(p_hamilt_in); + this->p_hamilt_sto = static_cast>*>(p_hamilt_in); this->p_stowf = p_stowf_in; this->nbands_ks = p_psi_in->get_nbands(); this->nbands_sto = p_stowf_in->nchi; @@ -51,7 +51,7 @@ void Sto_DOS::decide_param(const int& dos_nche, this->nbands_sto, this->p_kv, reinterpret_cast, Device>*>(this->p_stowf), - reinterpret_cast, Device>*>(this->p_hamilt_sto)); + reinterpret_cast, Device>*>(this->p_hamilt_sto)); if (dos_setemax) { this->emax = dos_emax_ev; @@ -124,7 +124,7 @@ void Sto_DOS::caldos(const double sigmain, const double de, cons p_stowf->chi0->fix_k(ik); pchi = p_stowf->chi0->get_pointer(); } - auto hchi_norm = std::bind(&hamilt::HamiltSdftPW>::hPsi_norm, + auto hchi_norm = std::bind(&StoHamiltPW>::hPsi_norm, p_hamilt_sto, std::placeholders::_1, std::placeholders::_2, diff --git a/source/source_pw/module_stodft/sto_dos.h b/source/source_pw/module_stodft/sto_dos.h index 68a433eb24c..ffd1cfc08c0 100644 --- a/source/source_pw/module_stodft/sto_dos.h +++ b/source/source_pw/module_stodft/sto_dos.h @@ -1,7 +1,7 @@ #ifndef STO_DOS #define STO_DOS #include "source_estate/elecstate.h" -#include "source_pw/module_stodft/hamilt_sdft_pw.h" +#include "source_pw/module_stodft/sto_hamilt_pw.h" #include "source_pw/module_stodft/sto_che.h" #include "source_pw/module_stodft/sto_func.h" #include "source_pw/module_stodft/sto_wf.h" @@ -65,7 +65,7 @@ class Sto_DOS = nullptr; ///< pointer to the stochastic wavefunctions Sto_Func stofunc; ///< functions - hamilt::HamiltSdftPW>* p_hamilt_sto = nullptr; ///< pointer to the Hamiltonian for sDFT + StoHamiltPW>* p_hamilt_sto = nullptr; ///< pointer to the Hamiltonian for sDFT }; #endif // STO_DOS \ No newline at end of file diff --git a/source/source_pw/module_stodft/sto_elecond.cpp b/source/source_pw/module_stodft/sto_elecond.cpp index 85516450f99..b705960abc7 100644 --- a/source/source_pw/module_stodft/sto_elecond.cpp +++ b/source/source_pw/module_stodft/sto_elecond.cpp @@ -33,7 +33,7 @@ Sto_EleCond::Sto_EleCond(UnitCell* p_ucell_in, : EleCond(p_ucell_in, p_kv_in, p_elec_in, p_wfcpw_in, p_psi_in, p_ppcell_in) { this->p_hamilt = p_hamilt_in; - this->p_hamilt_sto = static_cast, Device>*>(p_hamilt_in); + this->p_hamilt_sto = static_cast, Device>*>(p_hamilt_in); this->p_stowf = p_stowf_in; this->nbands_ks = p_psi_in->get_nbands(); this->nbands_sto = p_stowf_in->nchi; @@ -42,7 +42,7 @@ Sto_EleCond::Sto_EleCond(UnitCell* p_ucell_in, #ifdef __FLOAT_FFTW if(!std::is_same::value) { - this->hamilt_sto_ = new hamilt::HamiltSdftPW, Device>(p_elec_in->pot, p_wfcpw_in, p_kv_in, p_ppcell_in, p_ucell_in, 1, &this->low_emin_, &this->low_emax_); + this->hamilt_sto_ = new StoHamiltPW, Device>(p_elec_in->pot, p_wfcpw_in, p_kv_in, p_ppcell_in, p_ucell_in, 1, &this->low_emin_, &this->low_emax_); } #endif } @@ -149,33 +149,33 @@ void Sto_EleCond::decide_nche(const FPTYPE dt, } template -void Sto_EleCond::cal_jmatrix(hamilt::HamiltSdftPW, Device>* hamilt, - const psi::Psi, Device>& kspsi_all, - const psi::Psi, Device>& vkspsi, - const double* en, - const double* en_all, - std::complex* leftfact, - std::complex* rightfact, - psi::Psi, Device>& leftchi, - psi::Psi, Device>& rightchi, - psi::Psi, Device>& left_hchi, - psi::Psi, Device>& right_hchi, - psi::Psi, Device>& batch_vchi, - psi::Psi, Device>& batch_vhchi, +void Sto_EleCond::cal_jmatrix(StoHamiltPW, Device>* hamilt, + const psi::Psi, Device>& kspsi_all, + const psi::Psi, Device>& vkspsi, + const double* en, + const double* en_all, + std::complex* leftfact, + std::complex* rightfact, + psi::Psi, Device>& leftchi, + psi::Psi, Device>& rightchi, + psi::Psi, Device>& left_hchi, + psi::Psi, Device>& right_hchi, + psi::Psi, Device>& batch_vchi, + psi::Psi, Device>& batch_vhchi, #ifdef __MPI - psi::Psi, Device>& chi_all, - psi::Psi, Device>& hchi_all, - void* gatherinfo_ks, - void* gatherinfo_sto, + psi::Psi, Device>& chi_all, + psi::Psi, Device>& hchi_all, + void* gatherinfo_ks, + void* gatherinfo_sto, #endif - const int& bsize_psi, - std::complex* j1, - std::complex* j2, - std::complex* tmpj, - hamilt::Velocity& velop, - const int& ik, - const std::complex& factor, - const int bandinfo[6]) + const int& bsize_psi, + std::complex* j1, + std::complex* j2, + std::complex* tmpj, + hamilt::Velocity& velop, + const int& ik, + const std::complex& factor, + const int bandinfo[6]) { ModuleBase::timer::start("Sto_EleCond", "cal_jmatrix"); const std::complex float_factor = factor; @@ -556,14 +556,14 @@ void Sto_EleCond::sKG(const int& smear_type, this->low_emin_ = static_cast(*this->stofunc.Emin); this->low_emax_ = static_cast(*this->stofunc.Emax); lowfunc.set_E_range(&low_emin_, &low_emax_); - hamilt::HamiltSdftPW* p_low_hamilt = nullptr; + StoHamiltPW* p_low_hamilt = nullptr; if(hamilt_sto_ != nullptr) { p_low_hamilt = hamilt_sto_; } else { - p_low_hamilt = reinterpret_cast, Device>*>(this->p_hamilt_sto); + p_low_hamilt = reinterpret_cast, Device>*>(this->p_hamilt_sto); } // Init Chebyshev @@ -794,12 +794,12 @@ void Sto_EleCond::sKG(const int& smear_type, auto nroot_fd = std::bind(&Sto_Func::nroot_fd, &this->stofunc, std::placeholders::_1); che.calcoef_real(nroot_fd); - auto hchi_norm = std::bind(&hamilt::HamiltSdftPW, Device>::hPsi_norm, + auto hchi_norm = std::bind(&StoHamiltPW, Device>::hPsi_norm, p_hamilt_sto, std::placeholders::_1, std::placeholders::_2, std::placeholders::_3); - auto hchi_norm_low = std::bind(&hamilt::HamiltSdftPW::hPsi_norm, + auto hchi_norm_low = std::bind(&StoHamiltPW::hPsi_norm, p_low_hamilt, std::placeholders::_1, std::placeholders::_2, diff --git a/source/source_pw/module_stodft/sto_elecond.h b/source/source_pw/module_stodft/sto_elecond.h index c54a2682f69..4845525dcc6 100644 --- a/source/source_pw/module_stodft/sto_elecond.h +++ b/source/source_pw/module_stodft/sto_elecond.h @@ -2,8 +2,8 @@ #define STOELECOND_H #include "source_hamilt/hamilt.h" -#include "source_hsolver/hsolver_pw_sdft.h" #include "source_pw/module_pwdft/elecond.h" +#include "source_pw/module_stodft/sto_hsolver_pw.h" #include "source_pw/module_stodft/sto_wf.h" template @@ -82,8 +82,8 @@ class Sto_EleCond : protected EleCond Stochastic_WF, Device>* p_stowf = nullptr; ///< pointer to the stochastic wavefunctions Sto_Func stofunc; ///< functions - hamilt::HamiltSdftPW, Device>* p_hamilt_sto = nullptr; ///< pointer to the Hamiltonian for sDFT - hamilt::HamiltSdftPW, Device>* hamilt_sto_ = nullptr; ///< pointer to the Hamiltonian for sDFT + StoHamiltPW, Device>* p_hamilt_sto = nullptr; ///< pointer to the Hamiltonian for sDFT + StoHamiltPW, Device>* hamilt_sto_ = nullptr; ///< pointer to the Hamiltonian for sDFT lowTYPE low_emin_ = 0; ///< Emin of the Hamiltonian for sDFT lowTYPE low_emax_ = 0; ///< Emax of the Hamiltonian for sDFT protected: @@ -91,7 +91,7 @@ class Sto_EleCond : protected EleCond * @brief calculate Jmatrix * */ - void cal_jmatrix(hamilt::HamiltSdftPW, Device>* hamilt, + void cal_jmatrix(StoHamiltPW, Device>* hamilt, const psi::Psi, Device>& kspsi_all, const psi::Psi, Device>& vkspsi, const double* en, diff --git a/source/source_pw/module_stodft/sto_hamilt_pw.cpp b/source/source_pw/module_stodft/sto_hamilt_pw.cpp new file mode 100644 index 00000000000..2bb5040a361 --- /dev/null +++ b/source/source_pw/module_stodft/sto_hamilt_pw.cpp @@ -0,0 +1,67 @@ +#include "sto_hamilt_pw.h" +#include "source_base/timer.h" +#include "kernels/hpsi_norm_op.h" + +template +StoHamiltPW::StoHamiltPW(elecstate::Potential* pot_in, + ModulePW::PW_Basis_K* wfc_basis, + K_Vectors* p_kv, + pseudopot_cell_vnl* nlpp, + const UnitCell* ucell, + const int& npol, + Real* emin_in, + Real* emax_in) + : hamilt::HamiltPW(pot_in, wfc_basis, p_kv, nlpp, nullptr, ucell, nullptr), ngk(p_kv->ngk) +{ + this->classname = "StoHamiltPW"; + this->npwk_max = wfc_basis->npwk_max; + this->npol = npol; + this->emin = emin_in; + this->emax = emax_in; +} + +template +void StoHamiltPW::hPsi(const T* psi_in, T* hpsi, const int& nbands) +{ + auto call_act = [&, this](const hamilt::Operator* op, const bool& is_first_node) -> void { + op->act(nbands, this->npwk_max, this->npol, psi_in, hpsi, this->ngk[op->get_ik()], is_first_node); + }; + + ModuleBase::timer::start("StoHamiltPW", "hPsi"); + call_act(this->ops, true); // first node + hamilt::Operator* node((hamilt::Operator*)this->ops->next_op); + while (node != nullptr) + { + call_act(node, false); // other nodes + node = (hamilt::Operator*)(node->next_op); + } + ModuleBase::timer::end("StoHamiltPW", "hPsi"); + + return; +} + +template +void StoHamiltPW::hPsi_norm(const T* psi_in, T* hpsi_norm, const int& nbands) +{ + ModuleBase::timer::start("StoHamiltPW", "hPsi_norm"); + + this->hPsi(psi_in, hpsi_norm, nbands); + + const int ik = this->ops->get_ik(); + const int npwk_max = this->npwk_max; + const int npwk = this->ngk[ik]; + const Real emin = *this->emin; + const Real emax = *this->emax; + const Real Ebar = (emin + emax) / 2; + const Real DeltaE = (emax - emin) / 2; + + hamilt::hpsi_norm_op()(this->ctx, nbands, npwk_max, npwk, Ebar, DeltaE, hpsi_norm, psi_in); + ModuleBase::timer::end("StoHamiltPW", "hPsi_norm"); +} + +template class StoHamiltPW, base_device::DEVICE_CPU>; +template class StoHamiltPW, base_device::DEVICE_CPU>; +#if ((defined __CUDA) || (defined __ROCM)) +template class StoHamiltPW, base_device::DEVICE_GPU>; +template class StoHamiltPW, base_device::DEVICE_GPU>; +#endif diff --git a/source/source_pw/module_stodft/hamilt_sdft_pw.h b/source/source_pw/module_stodft/sto_hamilt_pw.h similarity index 70% rename from source/source_pw/module_stodft/hamilt_sdft_pw.h rename to source/source_pw/module_stodft/sto_hamilt_pw.h index 282ebec4247..31f5716af80 100644 --- a/source/source_pw/module_stodft/hamilt_sdft_pw.h +++ b/source/source_pw/module_stodft/sto_hamilt_pw.h @@ -1,18 +1,15 @@ -#ifndef HAMILTSDFTPW_H -#define HAMILTSDFTPW_H +#ifndef STO_HAMILT_PW_H +#define STO_HAMILT_PW_H #include "source_pw/module_pwdft/hamilt_pw.h" -namespace hamilt -{ - template -class HamiltSdftPW : public HamiltPW +class StoHamiltPW : public hamilt::HamiltPW { public: using Real = typename GetTypeReal::type; /** - * @brief Construct a new HamiltSdftPW object + * @brief Construct a new StoHamiltPW object * * @param pot_in potential * @param wfc_basis pw basis for wave functions @@ -21,19 +18,19 @@ class HamiltSdftPW : public HamiltPW * @param emin_in Emin of the Hamiltonian * @param emax_in Emax of the Hamiltonian */ - HamiltSdftPW(elecstate::Potential* pot_in, - ModulePW::PW_Basis_K* wfc_basis, - K_Vectors* p_kv, - pseudopot_cell_vnl* nlpp, - const UnitCell* ucell, - const int& npol, - Real* emin_in, - Real* emax_in); + StoHamiltPW(elecstate::Potential* pot_in, + ModulePW::PW_Basis_K* wfc_basis, + K_Vectors* p_kv, + pseudopot_cell_vnl* nlpp, + const UnitCell* ucell, + const int& npol, + Real* emin_in, + Real* emax_in); /** - * @brief Destroy the HamiltSdftPW object + * @brief Destroy the StoHamiltPW object * */ - ~HamiltSdftPW(){}; + ~StoHamiltPW(){}; /** * @brief Calculate \hat{H}|psi> @@ -62,6 +59,4 @@ class HamiltSdftPW : public HamiltPW std::vector& ngk; ///< number of G vectors }; -} // namespace hamilt - #endif diff --git a/source/source_hsolver/hsolver_pw_sdft.cpp b/source/source_pw/module_stodft/sto_hsolver_pw.cpp similarity index 74% rename from source/source_hsolver/hsolver_pw_sdft.cpp rename to source/source_pw/module_stodft/sto_hsolver_pw.cpp index f370e725d1e..dcb53ec9ddc 100644 --- a/source/source_hsolver/hsolver_pw_sdft.cpp +++ b/source/source_pw/module_stodft/sto_hsolver_pw.cpp @@ -1,4 +1,4 @@ -#include "hsolver_pw_sdft.h" +#include "sto_hsolver_pw.h" #include "source_base/global_function.h" #include "source_base/parallel_comm.h" @@ -11,23 +11,21 @@ #include -namespace hsolver -{ template -void HSolverPW_SDFT::solve(const UnitCell& ucell, - hamilt::Hamilt* pHamilt, - psi::Psi& psi, - psi::Psi& psi_cpu, - elecstate::ElecState* pes, - ModulePW::PW_Basis_K* wfc_basis, - Stochastic_WF& stowf, - const int istep, - const int iter, - std::ostream& log, - const bool skip_charge) +void StoHSolverPW::solve(const UnitCell& ucell, + hamilt::Hamilt* pHamilt, + psi::Psi& psi, + psi::Psi& psi_cpu, + elecstate::ElecState* pes, + ModulePW::PW_Basis_K* wfc_basis, + Stochastic_WF& stowf, + const int istep, + const int iter, + std::ostream& log, + const bool skip_charge) { - ModuleBase::TITLE("HSolverPW_SDFT", "solve"); - ModuleBase::timer::start("HSolverPW_SDFT", "solve"); + ModuleBase::TITLE("StoHSolverPW", "solve"); + ModuleBase::timer::start("StoHSolverPW", "solve"); // This override never calls HSolverPW::solve, which is where the base class // normally establishes the pool communication context. Set it up here so that @@ -59,7 +57,7 @@ void HSolverPW_SDFT::solve(const UnitCell& ucell, // part of KSDFT to get KS orbitals for (int ik = 0; ik < nks; ++ik) { - ModuleBase::timer::start("HSolverPW_SDFT", "solve_KS"); + ModuleBase::timer::start("StoHSolverPW", "solve_KS"); op.update_k(ik); if (nbands > 0 && this->ks_run) { @@ -79,7 +77,7 @@ void HSolverPW_SDFT::solve(const UnitCell& ucell, MPI_Bcast(&pes->ekb(ik, 0), nbands, MPI_DOUBLE, 0, BP_WORLD); } #endif - ModuleBase::timer::end("HSolverPW_SDFT", "solve_KS"); + ModuleBase::timer::end("StoHSolverPW", "solve_KS"); stoiter.orthog(ik, psi, stowf); stoiter.checkemm(ik, istep, iter, stowf); // check and reset emax & emin } @@ -116,7 +114,7 @@ void HSolverPW_SDFT::solve(const UnitCell& ucell, // for nscf, skip charge if (skip_charge) { - ModuleBase::timer::end("HSolverPW_SDFT", "solve"); + ModuleBase::timer::end("StoHSolverPW", "solve"); return; } @@ -131,14 +129,13 @@ void HSolverPW_SDFT::solve(const UnitCell& ucell, stoiter.cal_storho(ucell, stowf, pes_pw,wfc_basis); // will do rho symmetry and energy calculation in esolver - ModuleBase::timer::end("HSolverPW_SDFT", "solve"); + ModuleBase::timer::end("StoHSolverPW", "solve"); return; } -// template class HSolverPW_SDFT, base_device::DEVICE_CPU>; -template class HSolverPW_SDFT, base_device::DEVICE_CPU>; +// template class StoHSolverPW, base_device::DEVICE_CPU>; +template class StoHSolverPW, base_device::DEVICE_CPU>; #if ((defined __CUDA) || (defined __ROCM)) -// template class HSolverPW_SDFT, base_device::DEVICE_GPU>; -template class HSolverPW_SDFT, base_device::DEVICE_GPU>; +// template class StoHSolverPW, base_device::DEVICE_GPU>; +template class StoHSolverPW, base_device::DEVICE_GPU>; #endif -} // namespace hsolver diff --git a/source/source_pw/module_stodft/sto_hsolver_pw.h b/source/source_pw/module_stodft/sto_hsolver_pw.h new file mode 100644 index 00000000000..8c9b8cef9c2 --- /dev/null +++ b/source/source_pw/module_stodft/sto_hsolver_pw.h @@ -0,0 +1,81 @@ +#ifndef STO_HSOLVER_PW_H +#define STO_HSOLVER_PW_H +#include "source_hsolver/hsolver_pw.h" +#include "source_pw/module_stodft/sto_hamilt_pw.h" +#include "source_pw/module_stodft/sto_iter.h" +template +class StoHSolverPW : public hsolver::HSolverPW +{ + protected: + using Real = typename GetTypeReal::type; + + public: + StoHSolverPW(K_Vectors* pkv, + ModulePW::PW_Basis_K* wfc_basis_in, + Stochastic_WF& stowf, + StoChe& stoche, + StoHamiltPW* p_hamilt_sto, + const std::string calculation_type_in, + const std::string basis_type_in, + const std::string method_in, + const bool use_uspp_in, + const int nspin_in, + const int scf_iter_in, + const int diag_iter_max_in, + const double diag_thr_in, + const bool need_subspace_in, + const int nbands_in, + const bool diago_smooth_ethr_in, + const int pw_diag_ndim_in, + const int diag_subspace_in, + const int nb2d_in, + const bool ks_run_in, + const bool all_ks_run_in, + const int bndpar_in) + : hsolver::HSolverPW(wfc_basis_in, + calculation_type_in, + basis_type_in, + method_in, + use_uspp_in, + nspin_in, + scf_iter_in, + diag_iter_max_in, + diag_thr_in, + need_subspace_in, + nbands_in, + diago_smooth_ethr_in, + pw_diag_ndim_in, + diag_subspace_in, + nb2d_in), + ks_run(ks_run_in), all_ks_run(all_ks_run_in), bndpar(bndpar_in) + { + stoiter.init(pkv, wfc_basis_in, stowf, stoche, p_hamilt_sto); + } + + void solve(const UnitCell& ucell, + hamilt::Hamilt* pHamilt, + psi::Psi& psi, + psi::Psi& psi_cpu, + elecstate::ElecState* pes, + ModulePW::PW_Basis_K* wfc_basis, + Stochastic_WF& stowf, + const int istep, + const int iter, + std::ostream& log, + const bool skip_charge); + + Stochastic_Iter stoiter; + + protected: + const bool ks_run; // true if the current process runs the KS part of the SDFT calculation + const bool all_ks_run; // true if every process runs the KS part + const int bndpar; // number of band-parallel groups + + using setmem_complex_op = base_device::memory::set_memory_op; + using setmem_var_op = base_device::memory::set_memory_op; + using syncmem_h2d_op = base_device::memory::synchronize_memory_op; + using syncmem_d2h_op = base_device::memory::synchronize_memory_op; + using syncmem_var_h2d_op = base_device::memory::synchronize_memory_op; + using syncmem_var_d2h_op = base_device::memory::synchronize_memory_op; +}; +#endif diff --git a/source/source_pw/module_stodft/sto_iter.cpp b/source/source_pw/module_stodft/sto_iter.cpp index d4041a54ad2..6a121064866 100644 --- a/source/source_pw/module_stodft/sto_iter.cpp +++ b/source/source_pw/module_stodft/sto_iter.cpp @@ -41,7 +41,7 @@ void Stochastic_Iter::init(K_Vectors* pkv_in, ModulePW::PW_Basis_K* wfc_basis, Stochastic_WF& stowf, StoChe& stoche, - hamilt::HamiltSdftPW* p_hamilt_sto) + StoHamiltPW* p_hamilt_sto) { p_che = stoche.p_che.get(); spolyv = stoche.spolyv.get(); @@ -186,7 +186,7 @@ void Stochastic_Iter::checkemm(const int& ik, while (true) { bool converge; - auto hchi_norm = std::bind(&hamilt::HamiltSdftPW::hPsi_norm, + auto hchi_norm = std::bind(&StoHamiltPW::hPsi_norm, p_hamilt_sto, std::placeholders::_1, std::placeholders::_2, @@ -400,7 +400,7 @@ void Stochastic_Iter::calPn(const int& ik, Stochastic_WF& pchi = stowf.chi0->get_pointer(); } - auto hchi_norm = std::bind(&hamilt::HamiltSdftPW::hPsi_norm, + auto hchi_norm = std::bind(&StoHamiltPW::hPsi_norm, p_hamilt_sto, std::placeholders::_1, std::placeholders::_2, @@ -784,7 +784,7 @@ void Stochastic_Iter::calTnchi_ik(const int& ik, Stochastic_WFupdateHk(ik); // necessary, because itermu should be called before this function } - auto hchi_norm = std::bind(&hamilt::HamiltSdftPW::hPsi_norm, + auto hchi_norm = std::bind(&StoHamiltPW::hPsi_norm, p_hamilt_sto, std::placeholders::_1, std::placeholders::_2, diff --git a/source/source_pw/module_stodft/sto_iter.h b/source/source_pw/module_stodft/sto_iter.h index 97df2639310..60c3a7a56af 100644 --- a/source/source_pw/module_stodft/sto_iter.h +++ b/source/source_pw/module_stodft/sto_iter.h @@ -3,7 +3,7 @@ #include "source_base/math_chebyshev.h" #include "source_estate/elecstate_pw.h" #include "source_hamilt/hamilt.h" -#include "source_pw/module_stodft/hamilt_sdft_pw.h" +#include "source_pw/module_stodft/sto_hamilt_pw.h" #include "source_psi/psi.h" #include "sto_che.h" #include "sto_func.h" @@ -42,7 +42,7 @@ class Stochastic_Iter ModulePW::PW_Basis_K* wfc_basis, Stochastic_WF& stowf, StoChe& stoche, - hamilt::HamiltSdftPW* p_hamilt_sto); + StoHamiltPW* p_hamilt_sto); /** * @brief sum demet and eband energies for each k point and each band @@ -117,7 +117,7 @@ class Stochastic_Iter ModuleBase::Chebyshev* p_che = nullptr; Sto_Func stofunc; - hamilt::HamiltSdftPW* p_hamilt_sto = nullptr; + StoHamiltPW* p_hamilt_sto = nullptr; double mu0 = 0.0; // chemical potential; unit in Ry bool change = false; diff --git a/source/source_pw/module_stodft/sto_tool.cpp b/source/source_pw/module_stodft/sto_tool.cpp index 4c432a687c6..f63a8707236 100644 --- a/source/source_pw/module_stodft/sto_tool.cpp +++ b/source/source_pw/module_stodft/sto_tool.cpp @@ -18,7 +18,7 @@ void check_che_op::operator()(const int& nche_in, const int& nbands_sto, K_Vectors* p_kv, Stochastic_WF, Device>* p_stowf, - hamilt::HamiltSdftPW, Device>* p_hamilt_sto) + StoHamiltPW, Device>* p_hamilt_sto) { //------------------------------ // Convergence test @@ -78,7 +78,7 @@ void check_che_op::operator()(const int& nche_in, while (true) { bool converge; - auto hchi_norm = std::bind(&hamilt::HamiltSdftPW, Device>::hPsi_norm, + auto hchi_norm = std::bind(&StoHamiltPW, Device>::hPsi_norm, p_hamilt_sto, std::placeholders::_1, std::placeholders::_2, diff --git a/source/source_pw/module_stodft/sto_tool.h b/source/source_pw/module_stodft/sto_tool.h index a2eaf2a00bb..e2a04f07da9 100644 --- a/source/source_pw/module_stodft/sto_tool.h +++ b/source/source_pw/module_stodft/sto_tool.h @@ -1,7 +1,7 @@ #ifndef STO_TOOL_H #define STO_TOOL_H #include "source_cell/klist.h" -#include "source_pw/module_stodft/hamilt_sdft_pw.h" +#include "source_pw/module_stodft/sto_hamilt_pw.h" #include "source_pw/module_stodft/sto_wf.h" #include "source_base/module_device/memory_op.h" #include "source_psi/psi.h" @@ -24,7 +24,7 @@ struct check_che_op const int& nbands_sto, K_Vectors* p_kv, Stochastic_WF, Device>* p_stowf, - hamilt::HamiltSdftPW, Device>* p_hamilt_sto); + StoHamiltPW, Device>* p_hamilt_sto); }; /** diff --git a/source/source_pw/module_stodft/test/CMakeLists.txt b/source/source_pw/module_stodft/test/CMakeLists.txt index 836a1ac7dcd..3e933f7d1cd 100644 --- a/source/source_pw/module_stodft/test/CMakeLists.txt +++ b/source/source_pw/module_stodft/test/CMakeLists.txt @@ -1,4 +1,8 @@ abacus_disable_feature_definitions(__MPI) +# These tests exercise CPU code only and link no GPU kernels, so build them +# without the device instantiations (as source_hsolver/test does). +abacus_disable_feature_definitions(__CUDA) +abacus_disable_feature_definitions(__ROCM) AddTest( TARGET MODULE_PW_Sto_Tool_UTs @@ -9,6 +13,20 @@ AddTest( AddTest( TARGET MODULE_PW_Sto_Hamilt_UTs LIBS parameter psi base device planewave_serial symmetry - SOURCES ../hamilt_sdft_pw.cpp test_hamilt_sto.cpp ../../../source_hamilt/operator.cpp + SOURCES ../sto_hamilt_pw.cpp test_sto_hamilt_pw.cpp ../../../source_hamilt/operator.cpp ../../../source_cell/klist.cpp ../../../source_cell/klist_io.cpp ../../../source_cell/parallel_kpoints.cpp ../../../source_cell/reciprocal_grid.cpp +) + +AddTest( + TARGET MODULE_PW_Sto_HSolver_UTs + LIBS parameter psi device base container MPI::MPI_CXX + SOURCES test_sto_hsolver_pw.cpp ../sto_hsolver_pw.cpp + ../../../source_hsolver/hsolver_pw.cpp ../../../source_hsolver/diago_bpcg.cpp + ../../../source_hsolver/diago_dav_subspace.cpp ../../../source_hsolver/diag_const_nums.cpp + ../../../source_hsolver/diago_iter_assist.cpp ../../../source_hsolver/para_lin_tf.cpp + ../../../source_estate/elecstate_tools.cpp ../../../source_estate/occupy.cpp + ../../../source_base/module_fft/fft_bundle.cpp ../../../source_base/module_fft/fft_cpu.cpp + # This test calls MPI_Init in main() and its mocks take MPI_Comm arguments + # unconditionally, so it must keep __MPI even though this directory disables it. + KEEP_FEATURE_DEFINITIONS __MPI ) \ No newline at end of file diff --git a/source/source_pw/module_stodft/test/test_hamilt_sto.cpp b/source/source_pw/module_stodft/test/test_sto_hamilt_pw.cpp similarity index 93% rename from source/source_pw/module_stodft/test/test_hamilt_sto.cpp rename to source/source_pw/module_stodft/test/test_sto_hamilt_pw.cpp index b749cc67b47..7023d90719b 100644 --- a/source/source_pw/module_stodft/test/test_hamilt_sto.cpp +++ b/source/source_pw/module_stodft/test/test_sto_hamilt_pw.cpp @@ -1,4 +1,4 @@ -#include "../hamilt_sdft_pw.h" +#include "../sto_hamilt_pw.h" #include "source_pw/module_pwdft/dftu_base.h" #include "source_hamilt/operator.h" @@ -77,7 +77,7 @@ class TestHamiltSto : public ::testing::Test p_kv = new K_Vectors(); std::vector ngk = {2}; p_kv->ngk = ngk; - hamilt_sto = new hamilt::HamiltSdftPW, base_device::DEVICE_CPU>(pot, wfc_basis, p_kv, nullptr, nullptr, npol, &emin, &emax); + hamilt_sto = new StoHamiltPW, base_device::DEVICE_CPU>(pot, wfc_basis, p_kv, nullptr, nullptr, npol, &emin, &emax); hamilt_sto->ops = new TestOp, base_device::DEVICE_CPU>(); } @@ -92,7 +92,7 @@ class TestHamiltSto : public ::testing::Test elecstate::Potential* pot; ModulePW::PW_Basis_K* wfc_basis; K_Vectors* p_kv; - hamilt::HamiltSdftPW, base_device::DEVICE_CPU>* hamilt_sto; + StoHamiltPW, base_device::DEVICE_CPU>* hamilt_sto; double emin = -2.0; double emax = 2.0; }; diff --git a/source/source_hsolver/test/test_hsolver_sdft.cpp b/source/source_pw/module_stodft/test/test_sto_hsolver_pw.cpp similarity index 94% rename from source/source_hsolver/test/test_hsolver_sdft.cpp rename to source/source_pw/module_stodft/test/test_sto_hsolver_pw.cpp index ddfdce7cc55..b6dc43a8085 100644 --- a/source/source_hsolver/test/test_hsolver_sdft.cpp +++ b/source/source_pw/module_stodft/test/test_sto_hsolver_pw.cpp @@ -2,12 +2,12 @@ #include #include -#include "hsolver_pw_sup.h" -#include "hsolver_supplementary_mock.h" +#include "source_hsolver/test/hsolver_pw_sup.h" +#include "source_hsolver/test/hsolver_supplementary_mock.h" #include "source_base/parallel_comm.h" #include "source_estate/elecstate_pw.h" #include "source_hsolver/hsolver_pw.h" -#include "source_hsolver/hsolver_pw_sdft.h" +#include "source_pw/module_stodft/sto_hsolver_pw.h" // mock for module_sdft template @@ -92,7 +92,7 @@ void Stochastic_Iter::init(K_Vectors* pkv_in, ModulePW::PW_Basis_K* wfc_basis, Stochastic_WF& stowf, StoChe& stoche, - hamilt::HamiltSdftPW* p_hamilt_sto) + StoHamiltPW* p_hamilt_sto) { this->nchip = stowf.nchip; ; @@ -230,7 +230,7 @@ namespace ModulePW { const double factor) const; } /************************************************ - * unit test of HSolverPW_SDFT class + * unit test of StoHSolverPW class ***********************************************/ /** @@ -239,13 +239,13 @@ namespace ModulePW { * - with psi; * - without psi; * - skip charge; - * - 2. hsolver::HSolverPW_SDFT::diagethr (for cases below) + * - 2. StoHSolverPW::diagethr (for cases below) * - set_diagethr, for setting diagethr; */ -class TestHSolverPW_SDFT : public ::testing::Test +class TestStoHSolverPW : public ::testing::Test { public: - TestHSolverPW_SDFT() : elecstate_test(nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr) + TestStoHSolverPW() : elecstate_test(nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr) { stoche.nche = 8; stoche.method_sto = 1; @@ -254,9 +254,9 @@ class TestHSolverPW_SDFT : public ::testing::Test Stochastic_WF> stowf; K_Vectors kv; StoChe stoche; - hamilt::HamiltSdftPW>* p_hamilt_sto = nullptr; - hsolver::HSolverPW_SDFT, base_device::DEVICE_CPU> hs_d - = hsolver::HSolverPW_SDFT, base_device::DEVICE_CPU>( + StoHamiltPW>* p_hamilt_sto = nullptr; + StoHSolverPW, base_device::DEVICE_CPU> hs_d + = StoHSolverPW, base_device::DEVICE_CPU>( &kv, &pwbk, stowf, @@ -292,7 +292,7 @@ class TestHSolverPW_SDFT : public ::testing::Test std::ofstream temp_ofs; }; -// TEST_F(TestHSolverPW_SDFT, solve) +// TEST_F(TestStoHSolverPW, solve) // { // // initial memory and data // elecstate_test.ekb.create(1, 2); @@ -328,7 +328,7 @@ class TestHSolverPW_SDFT : public ::testing::Test // std::cout<<__FILE__<<__LINE__<<" "< void hamilt::HamiltPW::sPsi(T const*, T*, const int, const int, const int) const{} template -hamilt::HamiltSdftPW::HamiltSdftPW(elecstate::Potential* pot_in, - ModulePW::PW_Basis_K* wfc_basis, - K_Vectors* p_kv, - pseudopot_cell_vnl* nlpp, - const UnitCell* ucell, - const int& npol, - Real* emin_in, - Real* emax_in) - : HamiltPW(pot_in, wfc_basis, p_kv, nlpp, nullptr, ucell, nullptr), ngk(p_kv->ngk) +StoHamiltPW::StoHamiltPW(elecstate::Potential* pot_in, + ModulePW::PW_Basis_K* wfc_basis, + K_Vectors* p_kv, + pseudopot_cell_vnl* nlpp, + const UnitCell* ucell, + const int& npol, + Real* emin_in, + Real* emax_in) + : hamilt::HamiltPW(pot_in, wfc_basis, p_kv, nlpp, nullptr, ucell, nullptr), ngk(p_kv->ngk) { } template -void hamilt::HamiltSdftPW::hPsi_norm(const T* psi_in, T* hpsi, const int& nbands){} +void StoHamiltPW::hPsi_norm(const T* psi_in, T* hpsi, const int& nbands){} template class hamilt::HamiltPW, base_device::DEVICE_CPU>; -template class hamilt::HamiltSdftPW, base_device::DEVICE_CPU>; +template class StoHamiltPW, base_device::DEVICE_CPU>; template class hamilt::HamiltPW, base_device::DEVICE_CPU>; -template class hamilt::HamiltSdftPW, base_device::DEVICE_CPU>; +template class StoHamiltPW, base_device::DEVICE_CPU>; #if ((defined __CUDA) || (defined __ROCM)) template class hamilt::HamiltPW, base_device::DEVICE_GPU>; -template class hamilt::HamiltSdftPW, base_device::DEVICE_GPU>; +template class StoHamiltPW, base_device::DEVICE_GPU>; template class hamilt::HamiltPW, base_device::DEVICE_GPU>; -template class hamilt::HamiltSdftPW, base_device::DEVICE_GPU>; +template class StoHamiltPW, base_device::DEVICE_GPU>; #endif /**