From 4cc8ebe244c7e14a6bc8ff803e6e25f7e979bb5f Mon Sep 17 00:00:00 2001 From: sunliang98 <1700011430@pku.edu.cn> Date: Sat, 28 Mar 2026 16:24:38 +0800 Subject: [PATCH 1/3] Fix: Fix the get_energy function of several KEDFs --- source/source_estate/module_pot/pot_ml_exx.cpp | 4 +++- source/source_pw/module_ofdft/kedf_lkt.cpp | 2 +- source/source_pw/module_ofdft/kedf_ml.cpp | 4 +++- source/source_pw/module_ofdft/kedf_tf.cpp | 2 +- source/source_pw/module_ofdft/kedf_vw.cpp | 2 +- source/source_pw/module_ofdft/kedf_wt.cpp | 2 +- source/source_pw/module_ofdft/kedf_xwm.cpp | 2 +- 7 files changed, 11 insertions(+), 7 deletions(-) diff --git a/source/source_estate/module_pot/pot_ml_exx.cpp b/source/source_estate/module_pot/pot_ml_exx.cpp index 37283be39c..53393b1e33 100644 --- a/source/source_estate/module_pot/pot_ml_exx.cpp +++ b/source/source_estate/module_pot/pot_ml_exx.cpp @@ -65,8 +65,10 @@ void ML_EXX::set_para(const Input_para& inp, const UnitCell* ucell_in, const Mod if (this->descriptor_type[i] == "gamma") feg_inpt[i] = 1.; } - if (PARAM.inp.of_ml_feg == 1) + if (PARAM.inp.of_ml_feg == 1) + { this->feg_net_F = torch::softplus(this->nn->forward(feg_inpt)).to(this->device_CPU).contiguous().data_ptr()[0]; + } else { this->feg_net_F = this->nn->forward(feg_inpt).to(this->device_CPU).contiguous().data_ptr()[0]; diff --git a/source/source_pw/module_ofdft/kedf_lkt.cpp b/source/source_pw/module_ofdft/kedf_lkt.cpp index 9980a306f2..e35dca6d90 100644 --- a/source/source_pw/module_ofdft/kedf_lkt.cpp +++ b/source/source_pw/module_ofdft/kedf_lkt.cpp @@ -52,7 +52,7 @@ double KEDF_LKT::get_energy(const double* const* prho, ModulePW::PW_Basis* pw_rh } delete[] nabla_rho; - return energy; + return this->lkt_energy; } /** diff --git a/source/source_pw/module_ofdft/kedf_ml.cpp b/source/source_pw/module_ofdft/kedf_ml.cpp index 8bddaa3ad3..6b2643a7b5 100644 --- a/source/source_pw/module_ofdft/kedf_ml.cpp +++ b/source/source_pw/module_ofdft/kedf_ml.cpp @@ -87,8 +87,10 @@ void KEDF_ML::set_para( if (this->descriptor_type[i] == "gamma") feg_inpt[i] = 1.; } - if (PARAM.inp.of_ml_feg == 1) + if (PARAM.inp.of_ml_feg == 1) + { this->feg_net_F = torch::softplus(this->nn->forward(feg_inpt)).to(this->device_CPU).contiguous().data_ptr()[0]; + } else { this->feg_net_F = this->nn->forward(feg_inpt).to(this->device_CPU).contiguous().data_ptr()[0]; diff --git a/source/source_pw/module_ofdft/kedf_tf.cpp b/source/source_pw/module_ofdft/kedf_tf.cpp index fc84df77f1..a84be35887 100644 --- a/source/source_pw/module_ofdft/kedf_tf.cpp +++ b/source/source_pw/module_ofdft/kedf_tf.cpp @@ -43,7 +43,7 @@ double KEDF_TF::get_energy(const double* const* prho) } this->tf_energy = energy; Parallel_Reduce::reduce_all(this->tf_energy); - return energy; + return this->tf_energy; } /** diff --git a/source/source_pw/module_ofdft/kedf_vw.cpp b/source/source_pw/module_ofdft/kedf_vw.cpp index 605beb5999..33bd3bb5a1 100644 --- a/source/source_pw/module_ofdft/kedf_vw.cpp +++ b/source/source_pw/module_ofdft/kedf_vw.cpp @@ -69,7 +69,7 @@ double KEDF_vW::get_energy(double** pphi, ModulePW::PW_Basis* pw_rho) delete[] tempPhi; delete[] LapPhi; - return energy; + return this->vw_energy; } /** diff --git a/source/source_pw/module_ofdft/kedf_wt.cpp b/source/source_pw/module_ofdft/kedf_wt.cpp index 72dcc741f5..f9d99aef09 100644 --- a/source/source_pw/module_ofdft/kedf_wt.cpp +++ b/source/source_pw/module_ofdft/kedf_wt.cpp @@ -110,7 +110,7 @@ double KEDF_WT::get_energy(const double* const* prho, ModulePW::PW_Basis* pw_rho } delete[] kernelRhoBeta; - return energy; + return this->wt_energy; } /** diff --git a/source/source_pw/module_ofdft/kedf_xwm.cpp b/source/source_pw/module_ofdft/kedf_xwm.cpp index 811822e3fc..f80e5044a2 100644 --- a/source/source_pw/module_ofdft/kedf_xwm.cpp +++ b/source/source_pw/module_ofdft/kedf_xwm.cpp @@ -103,7 +103,7 @@ double KEDF_XWM::get_energy(const double* const* prho, ModulePW::PW_Basis* pw_rh delete[] w1Rho5_6; delete[] w2Rho5_6; - return energy; + return this->xwm_energy; } /** From 94c35da8ed42679dc54d6399080fed3396630d49 Mon Sep 17 00:00:00 2001 From: sunliang98 <1700011430@pku.edu.cn> Date: Sat, 28 Mar 2026 16:25:11 +0800 Subject: [PATCH 2/3] Fix: Fix generate_descriptor --- source/source_io/module_ml/write_mlkedf_descriptors.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/source/source_io/module_ml/write_mlkedf_descriptors.cpp b/source/source_io/module_ml/write_mlkedf_descriptors.cpp index 840adfac7f..6678c3f8f1 100644 --- a/source/source_io/module_ml/write_mlkedf_descriptors.cpp +++ b/source/source_io/module_ml/write_mlkedf_descriptors.cpp @@ -171,7 +171,7 @@ void Write_MLKEDF_Descriptors::generate_descriptor( // p this->cal_tool->getP(prho, pw_rho, nablaRho, container); - npy::SaveArrayAsNumpy("p.npy", false, 1, cshape, container); + npy::SaveArrayAsNumpy(out_dir + "p.npy", false, 1, cshape, container); for (int ik = 0; ik < this->cal_tool->nkernel; ++ik) { From d1966e5c663b2ccec9fd35af869913e4f0ed05b0 Mon Sep 17 00:00:00 2001 From: Mohan Chen Date: Sat, 28 Mar 2026 17:16:33 +0800 Subject: [PATCH 3/3] Apply suggestions from code review Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- source/source_io/module_ml/write_mlkedf_descriptors.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/source/source_io/module_ml/write_mlkedf_descriptors.cpp b/source/source_io/module_ml/write_mlkedf_descriptors.cpp index 6678c3f8f1..59ee758bd7 100644 --- a/source/source_io/module_ml/write_mlkedf_descriptors.cpp +++ b/source/source_io/module_ml/write_mlkedf_descriptors.cpp @@ -171,7 +171,7 @@ void Write_MLKEDF_Descriptors::generate_descriptor( // p this->cal_tool->getP(prho, pw_rho, nablaRho, container); - npy::SaveArrayAsNumpy(out_dir + "p.npy", false, 1, cshape, container); + npy::SaveArrayAsNumpy(out_dir + "/p.npy", false, 1, cshape, container); for (int ik = 0; ik < this->cal_tool->nkernel; ++ik) {