Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion source/source_esolver/esolver_ks.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -289,7 +289,7 @@ void ESolver_KS::iter_finish(UnitCell& ucell, const int istep, int& iter, bool &

// print energies
elecstate::print_etot(ucell.magnet, *pelec, conv_esolver, iter, drho,
dkin, duration, diag_ethr, 0, true, this->ds_rms_);
dkin, duration, *this->inp_, PARAM.globalv.two_fermi, diag_ethr, 0, true, this->ds_rms_);


#ifdef __JSON
Expand Down
30 changes: 16 additions & 14 deletions source/source_estate/elecstate_print.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -175,6 +175,8 @@ void print_etot(const Magnetism& magnet,
const double& scf_thr,
const double& scf_thr_kin,
const double& duration,
const Input_para& inp,
const bool two_fermi,
const double& pw_diag_thr,
const double& avg_iter,
const bool print,
Expand All @@ -190,7 +192,7 @@ void print_etot(const Magnetism& magnet,

GlobalV::ofs_running << " Electron density deviation " << scf_thr << std::endl;

if (PARAM.inp.basis_type == "pw")
if (inp.basis_type == "pw")
{
ModuleBase::GlobalFunc::OUT(GlobalV::ofs_running, "Diago Threshold", pw_diag_thr);
}
Expand All @@ -199,7 +201,7 @@ void print_etot(const Magnetism& magnet,
std::vector<double> energies_Ry;
std::vector<double> energies_eV;

if ((iter % PARAM.inp.out_freq_elec == 0) || converged || iter == PARAM.inp.scf_nmax)
if ((iter % inp.out_freq_elec == 0) || converged || iter == inp.scf_nmax)
{
int n_order = std::max(0, Occupy::gaussian_type);

Expand Down Expand Up @@ -248,7 +250,7 @@ void print_etot(const Magnetism& magnet,
energies_Ry.push_back(elec.f_en.e_local_pp);

//! vdw energy
std::string vdw_method = PARAM.inp.vdw_method;
std::string vdw_method = inp.vdw_method;
if (vdw_method == "d2") // Peize Lin add 2014-04, update 2021-03-09
{
titles.push_back("E_vdwD2");
Expand All @@ -266,7 +268,7 @@ void print_etot(const Magnetism& magnet,
}

// mohan add 20251108
if (PARAM.inp.dft_plus_u)
if (inp.dft_plus_u)
{
titles.push_back("E_plusU");
energies_Ry.push_back(elec.f_en.edftu);
Expand All @@ -277,7 +279,7 @@ void print_etot(const Magnetism& magnet,
energies_Ry.push_back(elec.f_en.exx);

//! solvation energy
if (PARAM.inp.imp_sol)
if (inp.imp_sol)
{
titles.push_back("E_sol_el");
energies_Ry.push_back(elec.f_en.esol_el);
Expand All @@ -286,27 +288,27 @@ void print_etot(const Magnetism& magnet,
}

//! electric field energy
if (PARAM.inp.efield_flag)
if (inp.efield_flag)
{
titles.push_back("E_efield");
energies_Ry.push_back(elecstate::Efield::etotefield);
}

//! gate energy
if (PARAM.inp.gate_flag)
if (inp.gate_flag)
{
titles.push_back("E_gatefield");
energies_Ry.push_back(elecstate::Gatefield::etotgatefield);
}

//! deepks energy
#ifdef __MLALGO
if (PARAM.inp.deepks_scf)
if (inp.deepks_scf)
{
titles.push_back("E_DeePKS");
energies_Ry.push_back(elec.f_en.edeepks_delta);
}
if (PARAM.inp.ml_exx)
if (inp.ml_exx)
{
titles.push_back("E_ML-EXX");
energies_Ry.push_back(elec.f_en.ml_exx);
Expand All @@ -322,7 +324,7 @@ void print_etot(const Magnetism& magnet,
}

// print out the Fermi energy if needed
if (PARAM.globalv.two_fermi)
if (two_fermi)
{
titles.push_back("E_Fermi_up");
energies_Ry.push_back(elec.eferm.ef_up);
Expand All @@ -336,7 +338,7 @@ void print_etot(const Magnetism& magnet,
}

// print out the band gap if needed
if (!PARAM.globalv.two_fermi)
if (!two_fermi)
{
titles.push_back("E_gap(k)"); // gap of given k-points
energies_Ry.push_back(elec.bandgap);
Expand All @@ -362,10 +364,10 @@ void print_etot(const Magnetism& magnet,

GlobalV::ofs_running << table.str() << std::endl;

if (PARAM.inp.out_level == "ie" || PARAM.inp.out_level == "m")
if (inp.out_level == "ie" || inp.out_level == "m")
{
std::vector<double> mag;
switch (PARAM.inp.nspin)
switch (inp.nspin)
{
case 2:
mag = {magnet.tot_mag, magnet.abs_mag};
Expand All @@ -384,7 +386,7 @@ void print_etot(const Magnetism& magnet,
}
// Pure SDFT (nbands=0) uses Chebyshev trace (CT) since no H diagonalization is performed.
// Mixed SDFT (nbands>0) still diagonalizes KS orbitals, so use the actual ks_solver label.
const std::string iter_label = (PARAM.inp.esolver_type == "sdft" && PARAM.inp.nbands == 0) ? "sdft" : PARAM.inp.ks_solver;
const std::string iter_label = (inp.esolver_type == "sdft" && inp.nbands == 0) ? "sdft" : inp.ks_solver;
elecstate::print_scf_iterinfo(iter_label,
iter,
4,
Expand Down
6 changes: 6 additions & 0 deletions source/source_estate/elecstate_print.h
Original file line number Diff line number Diff line change
Expand Up @@ -8,13 +8,19 @@ namespace elecstate
void print_format(const std::string& name,
const double& value);

/// @param inp the INPUT parameters whose flags decide which energy terms and
/// headers are printed
/// @param two_fermi whether the run keeps two Fermi levels; derived, so it
/// does not live in Input_para
void print_etot(const Magnetism& magnet,
const ElecState& elec,
const bool converged,
const int& iter_in,
const double& scf_thr,
const double& scf_thr_kin,
const double& duration,
const Input_para& inp,
const bool two_fermi,
const double& pw_diag_thr = 0,
const double& avg_iter = 0,
bool print = true,
Expand Down
98 changes: 50 additions & 48 deletions source/source_estate/test/elecstate_print_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,6 @@

#include "gmock/gmock.h"
#include "gtest/gtest.h"
#define private public
#include "source_cell/klist.h"
#include "source_estate/elecstate.h"
#include "source_estate/module_charge/charge.h"
Expand All @@ -11,7 +10,6 @@
#include "source_hamilt/module_xc/xc_functional.h"
#include "source_io/module_parameter/parameter.h"
#include "source_estate/elecstate_print.h"
#undef private
/***************************************************************
* mock functions
****************************************************************/
Expand Down Expand Up @@ -58,6 +56,11 @@ class ElecStatePrintTest : public ::testing::Test
protected:
elecstate::ElecState elecstate;
UnitCell ucell;
/// print_etot() takes the INPUT parameters and the two-Fermi flag as
/// arguments, so the fixture owns them instead of writing the global
/// parameter singleton.
Input_para inp;
bool two_fermi = false;
std::string output;
std::ifstream ifs;
std::ofstream ofs;
Expand Down Expand Up @@ -94,8 +97,7 @@ class ElecStatePrintTest : public ::testing::Test
ucell.magnet.tot_mag_nc[0] = 3.3;
ucell.magnet.tot_mag_nc[1] = 4.4;
ucell.magnet.tot_mag_nc[2] = 5.5;
PARAM.input.ks_solver = "dav";
PARAM.sys.log_file = "test.dat";
inp.ks_solver = "dav";
}
void TearDown()
{
Expand Down Expand Up @@ -129,56 +131,56 @@ TEST_F(ElecStatePrintTest, PrintEtot)
elecstate.charge = new Charge;
elecstate.charge->nrxx = 100;
elecstate.charge->nxyz = 1000;
PARAM.input.out_freq_elec = 1;
PARAM.input.imp_sol = true;
PARAM.input.efield_flag = true;
PARAM.input.gate_flag = true;
PARAM.sys.two_fermi = true;
inp.out_freq_elec = 1;
inp.imp_sol = true;
inp.efield_flag = true;
inp.gate_flag = true;
two_fermi = true;
GlobalV::MY_RANK = 0;
PARAM.input.basis_type = "pw";
PARAM.input.nspin = 2;
inp.basis_type = "pw";
inp.nspin = 2;

// iteration of different vdw_method
std::vector<std::string> vdw_methods = {"d2", "d3_0", "d3_bj"};
for (int i = 0; i < vdw_methods.size(); i++)
{
PARAM.input.vdw_method = vdw_methods[i];
inp.vdw_method = vdw_methods[i];
elecstate::print_etot(ucell.magnet,elecstate, converged, iter, scf_thr,
scf_thr_kin, duration, pw_diag_thr, avg_iter, false);
scf_thr_kin, duration, inp, two_fermi, pw_diag_thr, avg_iter, false);
}

// iteration of different ks_solver
std::vector<std::string> ks_solvers = {"cg", "lapack", "genelpa", "dav", "scalapack_gvx", "cusolver"};
for (int i = 0; i < ks_solvers.size(); i++)
{
PARAM.input.ks_solver = ks_solvers[i];
inp.ks_solver = ks_solvers[i];
testing::internal::CaptureStdout();

elecstate::print_etot(ucell.magnet,elecstate,converged, iter, scf_thr,
scf_thr_kin, duration, pw_diag_thr, avg_iter, print);
scf_thr_kin, duration, inp, two_fermi, pw_diag_thr, avg_iter, print);

output = testing::internal::GetCapturedStdout();
if (PARAM.input.ks_solver == "cg")
if (inp.ks_solver == "cg")
{
EXPECT_THAT(output, testing::HasSubstr("CG"));
}
else if (PARAM.input.ks_solver == "lapack")
else if (inp.ks_solver == "lapack")
{
EXPECT_THAT(output, testing::HasSubstr("LA"));
}
else if (PARAM.input.ks_solver == "genelpa")
else if (inp.ks_solver == "genelpa")
{
EXPECT_THAT(output, testing::HasSubstr("GE"));
}
else if (PARAM.input.ks_solver == "dav")
else if (inp.ks_solver == "dav")
{
EXPECT_THAT(output, testing::HasSubstr("DA"));
}
else if (PARAM.input.ks_solver == "scalapack_gvx")
else if (inp.ks_solver == "scalapack_gvx")
{
EXPECT_THAT(output, testing::HasSubstr("GV"));
}
else if (PARAM.input.ks_solver == "cusolver")
else if (inp.ks_solver == "cusolver")
{
EXPECT_THAT(output, testing::HasSubstr("CU"));
}
Expand Down Expand Up @@ -214,16 +216,16 @@ TEST_F(ElecStatePrintTest, PrintEtotColorS2)
elecstate.charge->nrxx = 100;
elecstate.charge->nxyz = 1000;

PARAM.input.out_freq_elec = 1;
PARAM.input.imp_sol = true;
PARAM.input.efield_flag = true;
PARAM.input.gate_flag = true;
PARAM.sys.two_fermi = true;
PARAM.input.nspin = 2;
inp.out_freq_elec = 1;
inp.imp_sol = true;
inp.efield_flag = true;
inp.gate_flag = true;
two_fermi = true;
inp.nspin = 2;
GlobalV::MY_RANK = 0;

elecstate::print_etot(ucell.magnet,elecstate,converged, iter, scf_thr,
scf_thr_kin, duration, pw_diag_thr, avg_iter, print);
scf_thr_kin, duration, inp, two_fermi, pw_diag_thr, avg_iter, print);

delete elecstate.charge;
}
Expand All @@ -243,17 +245,17 @@ TEST_F(ElecStatePrintTest, PrintEtotColorS4)
elecstate.charge->nrxx = 100;
elecstate.charge->nxyz = 1000;

PARAM.input.out_freq_elec = 1;
PARAM.input.imp_sol = true;
PARAM.input.efield_flag = true;
PARAM.input.gate_flag = true;
PARAM.sys.two_fermi = true;
PARAM.input.nspin = 4;
PARAM.input.noncolin = true;
inp.out_freq_elec = 1;
inp.imp_sol = true;
inp.efield_flag = true;
inp.gate_flag = true;
two_fermi = true;
inp.nspin = 4;
inp.noncolin = true;
GlobalV::MY_RANK = 0;

elecstate::print_etot(ucell.magnet,elecstate, converged, iter, scf_thr, scf_thr_kin,
duration, pw_diag_thr, avg_iter, print);
duration, inp, two_fermi, pw_diag_thr, avg_iter, print);

delete elecstate.charge;
}
Expand All @@ -272,17 +274,17 @@ TEST_F(ElecStatePrintTest, PrintEtotSDFTPure)
elecstate.charge->nrxx = 100;
elecstate.charge->nxyz = 1000;

PARAM.input.out_freq_elec = 1;
PARAM.input.nspin = 1;
inp.out_freq_elec = 1;
inp.nspin = 1;
GlobalV::MY_RANK = 0;
// Pure SDFT: nbands=0, no KS diagonalization -> ITER column should show CT
PARAM.input.esolver_type = "sdft";
PARAM.input.nbands = 0;
PARAM.input.ks_solver = "cg";
inp.esolver_type = "sdft";
inp.nbands = 0;
inp.ks_solver = "cg";

testing::internal::CaptureStdout();
elecstate::print_etot(ucell.magnet, elecstate, converged, iter, scf_thr,
scf_thr_kin, duration, pw_diag_thr, avg_iter, print);
scf_thr_kin, duration, inp, two_fermi, pw_diag_thr, avg_iter, print);
output = testing::internal::GetCapturedStdout();
EXPECT_THAT(output, testing::HasSubstr("CT"));

Expand All @@ -303,17 +305,17 @@ TEST_F(ElecStatePrintTest, PrintEtotSDFTMixed)
elecstate.charge->nrxx = 100;
elecstate.charge->nxyz = 1000;

PARAM.input.out_freq_elec = 1;
PARAM.input.nspin = 1;
inp.out_freq_elec = 1;
inp.nspin = 1;
GlobalV::MY_RANK = 0;
// Mixed SDFT: nbands>0, still diagonalizes KS orbitals -> ITER column shows ks_solver label
PARAM.input.esolver_type = "sdft";
PARAM.input.nbands = 5;
PARAM.input.ks_solver = "dav";
inp.esolver_type = "sdft";
inp.nbands = 5;
inp.ks_solver = "dav";

testing::internal::CaptureStdout();
elecstate::print_etot(ucell.magnet, elecstate, converged, iter, scf_thr,
scf_thr_kin, duration, pw_diag_thr, avg_iter, print);
scf_thr_kin, duration, inp, two_fermi, pw_diag_thr, avg_iter, print);
output = testing::internal::GetCapturedStdout();
EXPECT_THAT(output, testing::HasSubstr("DA"));

Expand Down
Loading