Skip to content

Commit 3e9dbbe

Browse files
committed
fix(deltaspin): reuse pre-initialized phsol in lambda loop for solver convergence
Previously HSolverLCAO was created on stack each inner step in cal_mw_from_lambda, losing solver state between calls. Now phsol is initialized once in ESolver_KS_LCAO::before_scf and passed through init_sc -> set_solver_parameters -> SpinConstrain::phsol, matching zdy-tmp behavior. Changes: - ESolver_KS_LCAO::phsol member: persistent solver pointer - before_scf: create phsol once, pass to init_sc - hamilt2rho_single: reuse phsol instead of stack solver - init_sc/set_solver_parameters: accept phsol_in parameter - cal_mw_from_lambda: use this->phsol when available
1 parent e05952b commit 3e9dbbe

9 files changed

Lines changed: 1074 additions & 11 deletions

File tree

source/source_esolver/esolver_ks_lcao.cpp

Lines changed: 11 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -153,13 +153,19 @@ void ESolver_KS_LCAO<TK, TR>::before_scf(UnitCell& ucell, const int istep)
153153
// since it depends on ionic positions
154154
this->deepks.build_overlap(ucell, orb_, pv, gd, *(two_center_bundle_.overlap_orb_alpha), PARAM.inp);
155155

156-
// 10) prepare sc calculation
156+
// 10) initialize HSolver once (reuse across SCF iterations and lambda loop)
157+
if (this->phsol == nullptr)
158+
{
159+
this->phsol = new hsolver::HSolverLCAO<TK>(&(this->pv), PARAM.inp.ks_solver);
160+
}
161+
162+
// 11) prepare sc calculation
157163
if (PARAM.inp.sc_mag_switch)
158164
{
159165
spinconstrain::SpinConstrain<TK>& sc = spinconstrain::SpinConstrain<TK>::getScInstance();
160166
sc.init_sc(PARAM.inp.sc_thr, PARAM.inp.nsc, PARAM.inp.nsc_min, PARAM.inp.alpha_trial,
161167
PARAM.inp.sccut, PARAM.inp.sc_drop_thr, ucell, &(this->pv),
162-
PARAM.inp.nspin, this->kv, this->p_hamilt, this->psi, this->dmat.dm, this->pelec);
168+
PARAM.inp.nspin, this->kv, this->p_hamilt, this->psi, this->dmat.dm, this->pelec, nullptr, this->phsol);
163169
// Set lambda update strategy
164170
if (PARAM.inp.sc_lambda_strategy == "linear_response")
165171
{
@@ -460,8 +466,9 @@ void ESolver_KS_LCAO<TK, TR>::hamilt2rho_single(UnitCell& ucell, int istep, int
460466
// 3) run Hsolver
461467
if (!skip_solve)
462468
{
463-
hsolver::HSolverLCAO<TK> hsolver_lcao_obj(&(this->pv), PARAM.inp.ks_solver);
464-
hsolver_lcao_obj.solve(static_cast<hamilt::Hamilt<TK>*>(this->p_hamilt), this->psi[0], this->pelec, *this->dmat.dm,
469+
hsolver::HSolverLCAO<TK>* hsolver_lcao_obj
470+
= static_cast<hsolver::HSolverLCAO<TK>*>(this->phsol);
471+
hsolver_lcao_obj->solve(static_cast<hamilt::Hamilt<TK>*>(this->p_hamilt), this->psi[0], this->pelec, *this->dmat.dm,
465472
this->chr, PARAM.inp.nspin, skip_charge);
466473
}
467474

source/source_esolver/esolver_ks_lcao.h

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -96,6 +96,9 @@ class ESolver_KS_LCAO : public ESolver_KS
9696
friend class LR::ESolver_LR<double, double>;
9797
friend class LR::ESolver_LR<std::complex<double>, double>;
9898

99+
//! Persistent solver pointer (for SpinConstrain reuse)
100+
void* phsol = nullptr;
101+
99102
// Temporarily store the stress to unify the interface with PW,
100103
// because it's hard to seperate force and stress calculation in LCAO.
101104
ModuleBase::matrix scs;

source/source_lcao/module_deltaspin/CMakeLists.txt

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@ list(APPEND objects
1212
template_helpers.cpp
1313
lambda_update_strategies.cpp
1414
lambda_strategy_integration.cpp
15+
lambda_solvers.cpp
1516
deltaspin_lcao.cpp
1617
)
1718

source/source_lcao/module_deltaspin/cal_mw_from_lambda.cpp

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -122,7 +122,6 @@ void spinconstrain::SpinConstrain<std::complex<double>>::cal_mw_from_lambda(int
122122
{
123123
psi::Psi<std::complex<double>>* psi_t = static_cast<psi::Psi<std::complex<double>>*>(this->psi);
124124
hamilt::Hamilt<std::complex<double>>* hamilt_t = static_cast<hamilt::Hamilt<std::complex<double>>*>(this->p_hamilt);
125-
hsolver::HSolverLCAO<std::complex<double>> hsolver_t(this->ParaV, PARAM.inp.ks_solver);
126125
if (PARAM.inp.nspin == 2)
127126
{
128127
dynamic_cast<hamilt::DeltaSpin<hamilt::OperatorLCAO<std::complex<double>, double>>*>(this->p_operator)->update_lambda();
@@ -131,7 +130,19 @@ void spinconstrain::SpinConstrain<std::complex<double>>::cal_mw_from_lambda(int
131130
{
132131
dynamic_cast<hamilt::DeltaSpin<hamilt::OperatorLCAO<std::complex<double>, std::complex<double>>>*>(this->p_operator)->update_lambda();
133132
}
134-
hsolver_t.solve(hamilt_t, psi_t[0], this->pelec, *this->dm_, *this->pelec->charge, PARAM.inp.nspin, true);
133+
// Use the pre-initialized solver from ESolver (same as zdy-tmp)
134+
if (this->phsol != nullptr)
135+
{
136+
hsolver::HSolverLCAO<std::complex<double>>* hsolver_t
137+
= static_cast<hsolver::HSolverLCAO<std::complex<double>>*>(this->phsol);
138+
hsolver_t->solve(hamilt_t, psi_t[0], this->pelec, *this->dm_, *this->pelec->charge, PARAM.inp.nspin, true);
139+
}
140+
else
141+
{
142+
// Fallback: create solver on stack (less efficient but functional)
143+
hsolver::HSolverLCAO<std::complex<double>> hsolver_t(this->ParaV, PARAM.inp.ks_solver);
144+
hsolver_t.solve(hamilt_t, psi_t[0], this->pelec, *this->dm_, *this->pelec->charge, PARAM.inp.nspin, true);
145+
}
135146
elecstate::calculate_weights(this->pelec->ekb, this->pelec->wg, this->pelec->klist,
136147
this->pelec->eferm, this->pelec->f_en, this->pelec->nelec_spin, this->pelec->skip_weights);
137148
elecstate::calEBand(this->pelec->ekb, this->pelec->wg, this->pelec->f_en);

source/source_lcao/module_deltaspin/init_sc.cpp

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,8 @@ void spinconstrain::SpinConstrain<TK>::init_sc(double sc_thr_in,
1818
elecstate::DensityMatrix<TK, double>* dm_in, // mohan add 2025-11-03
1919
#endif
2020
elecstate::ElecState* pelec_in,
21-
ModulePW::PW_Basis_K* pw_wfc_in)
21+
ModulePW::PW_Basis_K* pw_wfc_in,
22+
void* phsol_in)
2223
{
2324
this->set_input_parameters(sc_thr_in, nsc_in, nsc_min_in, alpha_trial_in, sccut_in, sc_drop_thr_in);
2425
this->set_atomCounts(ucell.get_atom_Counts());
@@ -33,7 +34,7 @@ void spinconstrain::SpinConstrain<TK>::init_sc(double sc_thr_in,
3334
this->pw_wfc_ = pw_wfc_in;
3435
this->set_decay_grad();
3536
if(ParaV_in != nullptr) this->set_ParaV(ParaV_in);
36-
this->set_solver_parameters(kv_in, p_hamilt_in, psi_in, pelec_in);
37+
this->set_solver_parameters(kv_in, p_hamilt_in, psi_in, pelec_in, phsol_in);
3738
#ifdef __LCAO
3839
this->dm_ = dm_in; // mohan add 2025-11-03
3940
#endif

0 commit comments

Comments
 (0)