From 9cf910132fc1487cdd89358da05885a4603e438e Mon Sep 17 00:00:00 2001 From: haochong zhang Date: Wed, 22 Jul 2026 08:30:32 +0000 Subject: [PATCH] refactor: introduce simulation context for LCAO workflows --- source/CMakeLists.txt | 8 + source/source_context/CMakeLists.txt | 19 + source/source_context/context_types.h | 321 +++++++++++ .../exx_state_field_mapping.inc | 42 ++ .../source_context/globalv_field_mapping.inc | 19 + source/source_context/input_field_mapping.inc | 505 ++++++++++++++++++ source/source_context/orchestration_context.h | 17 + .../source_context/restart_field_mapping.inc | 9 + source/source_context/simulation_context.h | 66 +++ .../simulation_context_binding.cpp | 48 ++ .../simulation_context_binding.h | 28 + .../simulation_context_builder.cpp | 370 +++++++++++++ .../simulation_context_builder.h | 34 ++ .../source_context/system_field_mapping.inc | 46 ++ source/source_context/test/CMakeLists.txt | 30 ++ .../test/check_context_mapping.py | 360 +++++++++++++ .../test/simulation_context_binding_test.cpp | 85 +++ .../test/simulation_context_builder_test.cpp | 250 +++++++++ source/source_esolver/CMakeLists.txt | 3 +- source/source_esolver/esolver.cpp | 62 +-- source/source_esolver/esolver.h | 9 + source/source_esolver/esolver_gets.cpp | 37 +- .../source_esolver/esolver_ks_lcao_tddft.cpp | 7 +- source/source_esolver/esolver_ks_lcaopw.cpp | 22 +- source/source_io/CMakeLists.txt | 19 +- .../source_io/module_ctrl/ctrl_iter_lcao.cpp | 6 +- .../module_ctrl/ctrl_runner_lcao.cpp | 44 +- .../source_io/module_ctrl/ctrl_scf_lcao.cpp | 104 ++-- source/source_io/module_dhs/write_dH.cpp | 3 +- source/source_io/module_dhs/write_dH.h | 3 + .../module_energy/write_eband_terms.hpp | 29 +- source/source_io/module_hs/cal_pLpR.cpp | 9 +- source/source_io/module_hs/cal_pLpR.h | 7 +- .../source_io/module_hs/cal_r_overlap_R.cpp | 134 +++-- source/source_io/module_hs/cal_r_overlap_R.h | 41 +- .../source_io/module_hs/output_mat_sparse.cpp | 120 ++--- .../source_io/module_hs/output_mat_sparse.h | 27 +- source/source_io/module_hs/single_R_io.cpp | 18 +- source/source_io/module_hs/single_R_io.h | 4 +- source/source_io/module_hs/write_HS.h | 14 +- source/source_io/module_hs/write_HS.hpp | 66 +-- source/source_io/module_hs/write_HS_R.cpp | 177 ++++-- source/source_io/module_hs/write_HS_R.h | 61 ++- .../source_io/module_hs/write_HS_sparse.cpp | 111 ++-- source/source_io/module_hs/write_HS_sparse.h | 16 +- source/source_io/module_hs/write_H_terms.cpp | 177 ++++-- source/source_io/module_hs/write_H_terms.h | 58 +- source/source_io/module_hs/write_vxc.hpp | 30 +- source/source_io/module_hs/write_vxc_lip.hpp | 23 +- source/source_io/module_hs/write_vxc_r.hpp | 10 +- .../module_restart/restart_exx_csr.h | 12 +- .../module_restart/restart_exx_csr.hpp | 6 +- source/source_io/test/single_R_io_test.cpp | 55 +- .../source_io/test/write_hs_r_compat_test.cpp | 138 +++-- source/source_lcao/module_lr/CMakeLists.txt | 10 +- .../module_lr/esolver_lrtd_lcao.cpp | 3 + source/source_lcao/module_lr/lr_spectrum.h | 6 + .../module_lr/lr_spectrum_velocity.cpp | 20 +- .../source_lcao/module_ri/Exx_LRI_interface.h | 4 +- .../module_ri/Exx_LRI_interface.hpp | 5 +- source/source_lcao/module_rt/td_info.cpp | 17 +- source/source_lcao/module_rt/td_info.h | 6 + .../test/snap_psibeta_half_tddft_test.cpp | 7 +- source/source_lcao/module_rt/velocity_op.cpp | 13 +- source/source_lcao/module_rt/velocity_op.h | 9 +- source/source_main/driver.cpp | 2 + source/source_main/driver.h | 8 + source/source_main/driver_run.cpp | 16 + .../agent_governance_check.py | 92 ++++ .../test_agent_governance_check.py | 52 ++ 70 files changed, 3595 insertions(+), 594 deletions(-) create mode 100644 source/source_context/CMakeLists.txt create mode 100644 source/source_context/context_types.h create mode 100644 source/source_context/exx_state_field_mapping.inc create mode 100644 source/source_context/globalv_field_mapping.inc create mode 100644 source/source_context/input_field_mapping.inc create mode 100644 source/source_context/orchestration_context.h create mode 100644 source/source_context/restart_field_mapping.inc create mode 100644 source/source_context/simulation_context.h create mode 100644 source/source_context/simulation_context_binding.cpp create mode 100644 source/source_context/simulation_context_binding.h create mode 100644 source/source_context/simulation_context_builder.cpp create mode 100644 source/source_context/simulation_context_builder.h create mode 100644 source/source_context/system_field_mapping.inc create mode 100644 source/source_context/test/CMakeLists.txt create mode 100644 source/source_context/test/check_context_mapping.py create mode 100644 source/source_context/test/simulation_context_binding_test.cpp create mode 100644 source/source_context/test/simulation_context_builder_test.cpp diff --git a/source/CMakeLists.txt b/source/CMakeLists.txt index 7273c6b8d72..364f7298ffa 100644 --- a/source/CMakeLists.txt +++ b/source/CMakeLists.txt @@ -460,6 +460,7 @@ add_subdirectory(source_basis/module_ao) add_subdirectory(source_basis/module_nao) add_subdirectory(source_md) add_subdirectory(source_basis/module_pw) +add_subdirectory(source_context) add_subdirectory(source_esolver) add_subdirectory(source_lcao/module_gint) add_subdirectory(source_io) @@ -475,6 +476,11 @@ add_library( source_main/driver.cpp source_main/driver_run.cpp) +target_compile_definitions( + driver + PRIVATE + ABACUS_CAN_BIND_SIMULATION_CONTEXT) + list(APPEND device_srcs source_pw/module_pwdft/kernels/nonlocal_op.cpp source_pw/module_pwdft/kernels/veff_op.cpp @@ -592,6 +598,7 @@ target_link_libraries( PRIVATE # Internal ABACUS targets: consumers before providers. driver + context esolver hsolver hamilt_general @@ -603,6 +610,7 @@ target_link_libraries( xc_ vdw relax + io_orchestration io_advanced io_basic io_input diff --git a/source/source_context/CMakeLists.txt b/source/source_context/CMakeLists.txt new file mode 100644 index 00000000000..b035d8f9f36 --- /dev/null +++ b/source/source_context/CMakeLists.txt @@ -0,0 +1,19 @@ +add_library( + context + OBJECT + simulation_context_builder.cpp + simulation_context_binding.cpp) + +target_compile_definitions( + context + PRIVATE + ABACUS_CAN_BIND_SIMULATION_CONTEXT + ABACUS_CAN_READ_SIMULATION_CONTEXT) + +if(ENABLE_COVERAGE) + add_coverage(context) +endif() + +if(BUILD_TESTING) + add_subdirectory(test) +endif() diff --git a/source/source_context/context_types.h b/source/source_context/context_types.h new file mode 100644 index 00000000000..bc04636c2a1 --- /dev/null +++ b/source/source_context/context_types.h @@ -0,0 +1,321 @@ +#ifndef ABACUS_SOURCE_CONTEXT_CONTEXT_TYPES_H +#define ABACUS_SOURCE_CONTEXT_CONTEXT_TYPES_H + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace ModuleContext +{ + +class InputFieldValue +{ + private: + struct HolderBase + { + virtual ~HolderBase() {} + virtual const std::type_info& type() const = 0; + }; + + template + struct Holder : HolderBase + { + explicit Holder(const T& source) : value(source) {} + const std::type_info& type() const override { return typeid(T); } + T value; + }; + + public: + InputFieldValue() {} + + template + explicit InputFieldValue(const T& value) : holder_(new Holder(value)) + { + } + + bool empty() const { return !holder_; } + const std::type_info& type() const { return holder_ ? holder_->type() : typeid(void); } + + template + const T& get() const + { + if (!holder_ || holder_->type() != typeid(T)) + { + throw std::bad_cast(); + } + return static_cast&>(*holder_).value; + } + + private: + std::shared_ptr holder_; +}; + +struct FieldProjection +{ + template + void project(const std::string& legacy_name, const T& value) + { + mapped_values[legacy_name] = InputFieldValue(value); + } + + template + const T& value(const std::string& legacy_name) const + { + const std::map::const_iterator it = mapped_values.find(legacy_name); + if (it == mapped_values.end()) + { + throw std::out_of_range("unknown projected INPUT field: " + legacy_name); + } + return it->second.get(); + } + + std::size_t mapped_field_count() const { return mapped_values.size(); } + + std::map mapped_values; +}; + +struct InputConfigDomain : FieldProjection +{ +}; + +struct RunControl : InputConfigDomain +{ + std::time_t start_time = 0; + std::string calculation; + std::string esolver_type; + std::string output_level; + bool cal_force = false; + bool cal_stress = false; +}; + +struct FileSystemLayout +{ + std::string input_card; + std::string structure_file; + std::string output_directory; + std::string readin_directory; + std::string structure_directory; + std::string matrix_directory; + std::string wavefunction_directory; + std::string mlkedf_descriptor_directory; + std::string deepks_label_directory; + std::string log_file; +}; + +struct ParallelTopology +{ + int world_size = 1; + int world_rank = 0; + int kpoint_pool_count = 1; + int pool_index = 0; + int processes_in_pool = 1; + int rank_in_pool = 0; + int band_group_index = 0; + int processes_in_band_group = 1; + int rank_in_band_group = 0; + int diagonalization_rank = 0; + int diagonalization_size = 1; + int diagonalization_color = 0; + int grid_rank = 0; + int grid_size = 1; + int threads_per_process = 1; +}; + +struct LogStreams +{ + std::ostream* running = nullptr; + std::ostream* warning = nullptr; + std::ostream* information = nullptr; + std::ostream* device = nullptr; +}; + +struct BasisInfo +{ + int nlocal = 0; + int npol = 1; + bool gamma_only_pw = false; + bool gamma_only_local = false; +}; + +struct GridInfo +{ + int charge_nx = 0; + int charge_ny = 0; + int charge_nz = 0; + bool double_grid = false; + double radial_dq = 0.0; + int radial_nqx = 0; + int radial_nqxq = 0; +}; + +struct ElectronicRuntimeInfo +{ + bool two_fermi = false; + bool use_uspp = false; + bool dos_minimum_is_explicit = false; + bool dos_maximum_is_explicit = false; + int local_band_count = 0; + bool ks_run = false; + bool all_ks_run = true; + bool has_double_data = true; + bool has_float_data = false; +}; + +struct FeatureRuntimeInfo +{ + bool deepks_setorb = false; + bool search_periodic_boundaries = true; + int lcao_kpoint_pool_count = 1; +}; + +struct SpinConfig : InputConfigDomain +{ + int nspin = 1; + bool spin_orbit = false; + bool noncollinear = false; + bool domag = false; + bool domag_z = false; +}; + +struct SolverConfig : InputConfigDomain +{ + std::string ks_solver; + std::string basis_type; +}; + +struct MatrixOutputRequest +{ + bool enabled = false; + int precision = 8; + int mode = 0; +}; + +struct MatrixOutputConfig : InputConfigDomain +{ + MatrixOutputRequest hs_k; + MatrixOutputRequest hs_r; + MatrixOutputRequest kinetic_k; + MatrixOutputRequest kinetic_r; + MatrixOutputRequest position_r; + MatrixOutputRequest dh; + MatrixOutputRequest ds; + MatrixOutputRequest h_t; + MatrixOutputRequest h_vnl; + MatrixOutputRequest h_vl; + MatrixOutputRequest h_vh; + MatrixOutputRequest h_vxc; + MatrixOutputRequest h_exx; + MatrixOutputRequest dh_t; + MatrixOutputRequest dh_vnl; + MatrixOutputRequest dh_vl; + MatrixOutputRequest dh_vh; + MatrixOutputRequest dh_vxc; + MatrixOutputRequest dh_exx; + MatrixOutputRequest vxc_r; + bool vxc_k = false; + bool band_energy_terms = false; + bool append = true; + int digits = 8; +}; + +struct GeneralOutputConfig : InputConfigDomain +{ + int electronic_frequency = 0; + int ionic_frequency = 0; + int tddft_frequency = 0; + bool all_logs = false; + bool molecular_dynamics_control = false; +}; + +struct DftUConfig : InputConfigDomain +{ + bool enabled = false; + bool dmft_enabled = false; + double ramping_ry = -10.0 / 13.6; + std::vector hubbard_u_ry; +}; + +struct ExactExchangeConfig : InputConfigDomain +{ + bool separate_loop = true; + int hybrid_step = 0; + double mixing_beta = 1.0; + bool symmetry_realspace = true; +}; + +struct ExactExchangeState : FieldProjection +{ + bool enabled = false; + bool real_number = false; + double hybrid_alpha = 0.0; + double hse_omega = 0.0; +}; + +struct RestartState +{ + bool save_charge = false; + bool save_hamiltonian = false; + bool load_charge = false; + bool load_charge_finished = false; + bool load_hamiltonian = false; + bool load_hamiltonian_finished = false; + bool restart_exact_exchange = false; + std::string folder; +}; + +class ExactExchangeStateService +{ + public: + virtual ~ExactExchangeStateService() {} + virtual ExactExchangeState snapshot() const = 0; +}; + +class RestartService +{ + public: + virtual ~RestartService() {} + virtual RestartState snapshot() const = 0; + virtual void update(const RestartState& state) = 0; +}; + +// The legacy parser remains the only source of defaults and validation. The +// complete one-to-one field projection is retained in these domain containers; +// stable named members are added only when a migrated consumer needs them. +struct CellConfig : InputConfigDomain {}; +struct FileInputConfig : InputConfigDomain {}; +struct ParallelConfig : InputConfigDomain {}; +struct DeviceConfig : InputConfigDomain {}; +struct ElectronicStructureConfig : InputConfigDomain {}; +struct ScfMixingConfig : InputConfigDomain {}; +struct LcaoConfig : InputConfigDomain {}; +struct RelaxationConfig : InputConfigDomain {}; +struct MolecularDynamicsConfig : InputConfigDomain {}; +struct OfdftConfig : InputConfigDomain {}; +struct StochasticDftConfig : InputConfigDomain {}; +struct DeepksConfig : InputConfigDomain {}; +struct RealTimeTddftConfig : InputConfigDomain {}; +struct LinearResponseTddftConfig : InputConfigDomain {}; +struct ChargeWavefunctionOutputConfig : InputConfigDomain {}; +struct RestartOutputConfig : InputConfigDomain {}; +struct PostprocessConfig : InputConfigDomain {}; +struct ModelConfig : InputConfigDomain {}; +struct VdwConfig : InputConfigDomain {}; +struct SpinConstraintConfig : InputConfigDomain {}; +struct QuasiAtomicOrbitalConfig : InputConfigDomain {}; +struct PexsiConfig : InputConfigDomain {}; +struct TestConfig : InputConfigDomain {}; +struct RdmftConfig : InputConfigDomain {}; +struct ExternalXcConfig : InputConfigDomain {}; +struct TimeDependentOfdftConfig : InputConfigDomain {}; +struct UncommonHardwareConfig : InputConfigDomain {}; + +} // namespace ModuleContext + +#endif diff --git a/source/source_context/exx_state_field_mapping.inc b/source/source_context/exx_state_field_mapping.inc new file mode 100644 index 00000000000..275f770f3f3 --- /dev/null +++ b/source/source_context/exx_state_field_mapping.inc @@ -0,0 +1,42 @@ +// MODULE_CONTEXT_LEGACY_SCHEMA_SHA256(3a575c0a888da54b0ea45d936a2ea8e21608fb08dfa88d885f003ac67e48c9c5) +MODULE_CONTEXT_EXX_STATE_FIELD(info_global, cal_exx) +MODULE_CONTEXT_EXX_STATE_FIELD(info_global, coulomb_param) +MODULE_CONTEXT_EXX_STATE_FIELD(info_global, ccp_type) +MODULE_CONTEXT_EXX_STATE_FIELD(info_global, hybrid_alpha) +MODULE_CONTEXT_EXX_STATE_FIELD(info_global, hse_omega) +MODULE_CONTEXT_EXX_STATE_FIELD(info_global, mixing_beta_for_loop1) +MODULE_CONTEXT_EXX_STATE_FIELD(info_global, separate_loop) +MODULE_CONTEXT_EXX_STATE_FIELD(info_global, hybrid_step) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, coulomb_param) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, real_number) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, coul_moment) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, rotate_abfs) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, pca_threshold) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, files_abfs) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, files_shrink_abfs) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, C_threshold) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, V_threshold) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, dm_threshold) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, C_grad_threshold) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, V_grad_threshold) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, C_grad_R_threshold) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, V_grad_R_threshold) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, ccp_rmesh_times) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, exx_symmetry_realspace) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, kmesh_times) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, Cs_inv_thr) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, shrink_abfs_pca_thr) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, shrink_LU_inv_thr) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, multip_moments_threshold) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, exx_cs_inv_thr) +MODULE_CONTEXT_EXX_STATE_FIELD(info_ri, abfs_Lmax) +MODULE_CONTEXT_EXX_STATE_FIELD(info_lip, ccp_type) +MODULE_CONTEXT_EXX_STATE_FIELD(info_lip, hse_omega) +MODULE_CONTEXT_EXX_STATE_FIELD(info_lip, lambda) +MODULE_CONTEXT_EXX_STATE_FIELD(info_opt_abfs, abfs_Lmax) +MODULE_CONTEXT_EXX_STATE_FIELD(info_opt_abfs, ecut_exx) +MODULE_CONTEXT_EXX_STATE_FIELD(info_opt_abfs, tolerence) +MODULE_CONTEXT_EXX_STATE_FIELD(info_opt_abfs, files_jles) +MODULE_CONTEXT_EXX_STATE_FIELD(info_opt_abfs, pca_threshold) +MODULE_CONTEXT_EXX_STATE_FIELD(info_opt_abfs, files_abfs) +MODULE_CONTEXT_EXX_STATE_FIELD(info_opt_abfs, kmesh_times) diff --git a/source/source_context/globalv_field_mapping.inc b/source/source_context/globalv_field_mapping.inc new file mode 100644 index 00000000000..5c4cd4c2091 --- /dev/null +++ b/source/source_context/globalv_field_mapping.inc @@ -0,0 +1,19 @@ +// MODULE_CONTEXT_LEGACY_SCHEMA_SHA256(a515d20576517a4222dc29569298031b75cfba96d73682c3647df04b6278aa16) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, NPROC) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, KPAR) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, MY_RANK) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, MY_POOL) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, MY_BNDGROUP) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, NPROC_IN_POOL) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, NPROC_IN_BNDGROUP) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, RANK_IN_POOL) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, RANK_IN_BPGROUP) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, DRANK) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, DSIZE) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, DCOLOR) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, GRANK) +MODULE_CONTEXT_GLOBALV_FIELD(parallel, GSIZE) +MODULE_CONTEXT_GLOBALV_FIELD(logs, ofs_running) +MODULE_CONTEXT_GLOBALV_FIELD(logs, ofs_warning) +MODULE_CONTEXT_GLOBALV_FIELD(logs, ofs_info) +MODULE_CONTEXT_GLOBALV_FIELD(logs, ofs_device) diff --git a/source/source_context/input_field_mapping.inc b/source/source_context/input_field_mapping.inc new file mode 100644 index 00000000000..57f943e343a --- /dev/null +++ b/source/source_context/input_field_mapping.inc @@ -0,0 +1,505 @@ +// MODULE_CONTEXT_LEGACY_SCHEMA_SHA256(c7b32ee9ca7fe372db633204d5e03dc8ba84bb2efcace1877e31d86ec88d560d) +MODULE_CONTEXT_INPUT_FIELD(cell, suffix) +MODULE_CONTEXT_INPUT_FIELD(cell, ntype) +MODULE_CONTEXT_INPUT_FIELD(run, calculation) +MODULE_CONTEXT_INPUT_FIELD(run, esolver_type) +MODULE_CONTEXT_INPUT_FIELD(cell, symmetry) +MODULE_CONTEXT_INPUT_FIELD(cell, symmetry_prec) +MODULE_CONTEXT_INPUT_FIELD(cell, symmetry_autoclose) +MODULE_CONTEXT_INPUT_FIELD(run, cal_force) +MODULE_CONTEXT_INPUT_FIELD(run, cal_stress) +MODULE_CONTEXT_INPUT_FIELD(parallel_input, kpar) +MODULE_CONTEXT_INPUT_FIELD(parallel_input, bndpar) +MODULE_CONTEXT_INPUT_FIELD(cell, latname) +MODULE_CONTEXT_INPUT_FIELD(cell, assume_isolated) +MODULE_CONTEXT_INPUT_FIELD(cell, ecutwfc) +MODULE_CONTEXT_INPUT_FIELD(cell, ecutrho) +MODULE_CONTEXT_INPUT_FIELD(cell, nx) +MODULE_CONTEXT_INPUT_FIELD(cell, ny) +MODULE_CONTEXT_INPUT_FIELD(cell, nz) +MODULE_CONTEXT_INPUT_FIELD(cell, ndx) +MODULE_CONTEXT_INPUT_FIELD(cell, ndy) +MODULE_CONTEXT_INPUT_FIELD(cell, ndz) +MODULE_CONTEXT_INPUT_FIELD(cell, cell_factor) +MODULE_CONTEXT_INPUT_FIELD(cell, erf_ecut) +MODULE_CONTEXT_INPUT_FIELD(cell, erf_height) +MODULE_CONTEXT_INPUT_FIELD(cell, erf_sigma) +MODULE_CONTEXT_INPUT_FIELD(cell, fft_mode) +MODULE_CONTEXT_INPUT_FIELD(cell, init_wfc) +MODULE_CONTEXT_INPUT_FIELD(cell, pw_seed) +MODULE_CONTEXT_INPUT_FIELD(cell, init_chg) +MODULE_CONTEXT_INPUT_FIELD(cell, dm_to_rho) +MODULE_CONTEXT_INPUT_FIELD(cell, chg_extrap) +MODULE_CONTEXT_INPUT_FIELD(cell, init_vel) +MODULE_CONTEXT_INPUT_FIELD(file_input, input_file) +MODULE_CONTEXT_INPUT_FIELD(file_input, stru_file) +MODULE_CONTEXT_INPUT_FIELD(file_input, kpoint_file) +MODULE_CONTEXT_INPUT_FIELD(file_input, pseudo_dir) +MODULE_CONTEXT_INPUT_FIELD(file_input, orbital_dir) +MODULE_CONTEXT_INPUT_FIELD(file_input, read_file_dir) +MODULE_CONTEXT_INPUT_FIELD(file_input, restart_load) +MODULE_CONTEXT_INPUT_FIELD(file_input, wannier_card) +MODULE_CONTEXT_INPUT_FIELD(cell, mem_saver) +MODULE_CONTEXT_INPUT_FIELD(parallel_input, diago_proc) +MODULE_CONTEXT_INPUT_FIELD(cell, nbspline) +MODULE_CONTEXT_INPUT_FIELD(cell, kspacing) +MODULE_CONTEXT_INPUT_FIELD(cell, koffset) +MODULE_CONTEXT_INPUT_FIELD(cell, kmesh_type) +MODULE_CONTEXT_INPUT_FIELD(cell, min_dist_coef) +MODULE_CONTEXT_INPUT_FIELD(device, device) +MODULE_CONTEXT_INPUT_FIELD(device, precision) +MODULE_CONTEXT_INPUT_FIELD(device, gint_precision) +MODULE_CONTEXT_INPUT_FIELD(device, timer_enable_nvtx) +MODULE_CONTEXT_INPUT_FIELD(solver, ks_solver) +MODULE_CONTEXT_INPUT_FIELD(solver, basis_type) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, nbands) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, nelec) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, nelec_delta) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, nupdown) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, dft_functional) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, xc_temperature) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, pseudo_rcut) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, pseudo_mesh) +MODULE_CONTEXT_INPUT_FIELD(spin, nspin) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, pw_diag_nmax) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, pw_diag_thr) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, diago_smooth_ethr) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, pw_diag_ndim) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, diago_cg_prec) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, diag_subspace) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, use_k_continuity) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, smearing_method) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, smearing_sigma) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_mode) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_beta) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_ndim) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_restart) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_gg0) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_beta_mag) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_gg0_mag) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_gg0_min) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_angle) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_tau) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_dftu) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, mixing_dmr) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, gamma_only) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, scf_nmax) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, scf_thr) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, scf_ene_thr) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, scf_thr_type) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, scf_os_stop) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, scf_os_thr) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, scf_os_ndim) +MODULE_CONTEXT_INPUT_FIELD(scf_mixing, sc_os_ndim) +MODULE_CONTEXT_INPUT_FIELD(spin, lspinorb) +MODULE_CONTEXT_INPUT_FIELD(spin, noncolin) +MODULE_CONTEXT_INPUT_FIELD(spin, soc_lambda) +MODULE_CONTEXT_INPUT_FIELD(electronic_structure, dfthalf_type) +MODULE_CONTEXT_INPUT_FIELD(lcao, nb2d) +MODULE_CONTEXT_INPUT_FIELD(lcao, lmaxmax) +MODULE_CONTEXT_INPUT_FIELD(lcao, lcao_ecut) +MODULE_CONTEXT_INPUT_FIELD(lcao, lcao_dk) +MODULE_CONTEXT_INPUT_FIELD(lcao, lcao_dr) +MODULE_CONTEXT_INPUT_FIELD(lcao, lcao_rmax) +MODULE_CONTEXT_INPUT_FIELD(lcao, search_radius) +MODULE_CONTEXT_INPUT_FIELD(lcao, bx) +MODULE_CONTEXT_INPUT_FIELD(lcao, by) +MODULE_CONTEXT_INPUT_FIELD(lcao, bz) +MODULE_CONTEXT_INPUT_FIELD(lcao, elpa_num_thread) +MODULE_CONTEXT_INPUT_FIELD(lcao, nstream) +MODULE_CONTEXT_INPUT_FIELD(lcao, bessel_nao_ecut) +MODULE_CONTEXT_INPUT_FIELD(lcao, bessel_nao_tolerence) +MODULE_CONTEXT_INPUT_FIELD(lcao, bessel_nao_rcuts) +MODULE_CONTEXT_INPUT_FIELD(lcao, bessel_nao_smooth) +MODULE_CONTEXT_INPUT_FIELD(lcao, bessel_nao_sigma) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax_method) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax_new) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax_scale_force) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax_nmax) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax_cg_thr) +MODULE_CONTEXT_INPUT_FIELD(relaxation, force_thr) +MODULE_CONTEXT_INPUT_FIELD(relaxation, force_thr_ev) +MODULE_CONTEXT_INPUT_FIELD(relaxation, force_zero_out) +MODULE_CONTEXT_INPUT_FIELD(relaxation, stress_thr) +MODULE_CONTEXT_INPUT_FIELD(relaxation, press1) +MODULE_CONTEXT_INPUT_FIELD(relaxation, press2) +MODULE_CONTEXT_INPUT_FIELD(relaxation, press3) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax_bfgs_w1) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax_bfgs_w2) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax_bfgs_rmax) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax_bfgs_rmin) +MODULE_CONTEXT_INPUT_FIELD(relaxation, relax_bfgs_init) +MODULE_CONTEXT_INPUT_FIELD(relaxation, fixed_axes) +MODULE_CONTEXT_INPUT_FIELD(relaxation, fixed_ibrav) +MODULE_CONTEXT_INPUT_FIELD(relaxation, fixed_atoms) +MODULE_CONTEXT_INPUT_FIELD(molecular_dynamics, mdp) +MODULE_CONTEXT_INPUT_FIELD(molecular_dynamics, ref_cell_factor) +MODULE_CONTEXT_INPUT_FIELD(molecular_dynamics, cal_syns) +MODULE_CONTEXT_INPUT_FIELD(molecular_dynamics, dmax) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_kinetic) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_method) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_conv) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_tole) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_tolp) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_tf_weight) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_vw_weight) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_wt_alpha) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_wt_beta) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_extwt_kappa) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_wt_rho0) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_hold_rho0) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_lkt_a) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_full_pw) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_full_pw_dim) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_read_kernel) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_kernel_file) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_xwm_kappa) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_xwm_rho_ref) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_gene_data) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_device) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_feg) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_nkernel) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_kernel) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_kernel_scaling) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_yukawa_alpha) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_kernel_file) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_gamma) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_p) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_q) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_tanhp) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_tanhq) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_chi_p) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_chi_q) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_gammanl) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_pnl) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_qnl) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_xi) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_tanhxi) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_tanhxi_nl) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_tanh_pnl) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_tanh_qnl) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_tanhp_nl) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_tanhq_nl) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_chi_xi) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_chi_pnl) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_chi_qnl) +MODULE_CONTEXT_INPUT_FIELD(ofdft, of_ml_local_test) +MODULE_CONTEXT_INPUT_FIELD(stochastic_dft, method_sto) +MODULE_CONTEXT_INPUT_FIELD(stochastic_dft, npart_sto) +MODULE_CONTEXT_INPUT_FIELD(stochastic_dft, nbands_sto) +MODULE_CONTEXT_INPUT_FIELD(stochastic_dft, nche_sto) +MODULE_CONTEXT_INPUT_FIELD(stochastic_dft, emin_sto) +MODULE_CONTEXT_INPUT_FIELD(stochastic_dft, emax_sto) +MODULE_CONTEXT_INPUT_FIELD(stochastic_dft, seed_sto) +MODULE_CONTEXT_INPUT_FIELD(stochastic_dft, initsto_ecut) +MODULE_CONTEXT_INPUT_FIELD(stochastic_dft, initsto_freq) +MODULE_CONTEXT_INPUT_FIELD(stochastic_dft, ml_exx) +MODULE_CONTEXT_INPUT_FIELD(deepks, deepks_out_labels) +MODULE_CONTEXT_INPUT_FIELD(deepks, deepks_out_freq_elec) +MODULE_CONTEXT_INPUT_FIELD(deepks, deepks_out_base) +MODULE_CONTEXT_INPUT_FIELD(deepks, deepks_scf) +MODULE_CONTEXT_INPUT_FIELD(deepks, deepks_bandgap) +MODULE_CONTEXT_INPUT_FIELD(deepks, deepks_band_range) +MODULE_CONTEXT_INPUT_FIELD(deepks, deepks_v_delta) +MODULE_CONTEXT_INPUT_FIELD(deepks, deepks_equiv) +MODULE_CONTEXT_INPUT_FIELD(deepks, deepks_out_unittest) +MODULE_CONTEXT_INPUT_FIELD(deepks, deepks_model) +MODULE_CONTEXT_INPUT_FIELD(deepks, bessel_descriptor_lmax) +MODULE_CONTEXT_INPUT_FIELD(deepks, bessel_descriptor_ecut) +MODULE_CONTEXT_INPUT_FIELD(deepks, bessel_descriptor_tolerence) +MODULE_CONTEXT_INPUT_FIELD(deepks, bessel_descriptor_rcut) +MODULE_CONTEXT_INPUT_FIELD(deepks, bessel_descriptor_smooth) +MODULE_CONTEXT_INPUT_FIELD(deepks, bessel_descriptor_sigma) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_dt) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, estep_per_md) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_force_dt) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_vext) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_vext_dire) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, init_vecpot_file) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_print_eij) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_edm) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, propagator) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_stype) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_ttype) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_tstart) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_tend) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_lcut1) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_lcut2) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_gauss_freq) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_gauss_phase) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_gauss_sigma) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_gauss_t0) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_gauss_amp) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trape_freq) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trape_phase) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trape_t1) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trape_t2) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trape_t3) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trape_amp) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trigo_freq1) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trigo_freq2) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trigo_phase1) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trigo_phase2) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_trigo_amp) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_heavi_t0) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, td_heavi_amp) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, ocp) +MODULE_CONTEXT_INPUT_FIELD(real_time_tddft, ocp_kb) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, lr_nstates) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, lr_init_xc_kernel) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, nocc) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, nvirt) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, xc_kernel) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, lr_solver) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, lr_thr) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, out_wfc_lr) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, lr_unrestricted) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, abs_wavelen_range) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, abs_broadening) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, abs_gauge) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, ri_hartree_benchmark) +MODULE_CONTEXT_INPUT_FIELD(linear_response_tddft, aims_nbasis) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_stru) +MODULE_CONTEXT_INPUT_FIELD(output, out_freq_elec) +MODULE_CONTEXT_INPUT_FIELD(output, out_freq_ion) +MODULE_CONTEXT_INPUT_FIELD(output, out_freq_td) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_chg) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_xc_r) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_pot) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_wfc_pw) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_band) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_dos) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_ldos) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_mul) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_proj_band) +MODULE_CONTEXT_INPUT_FIELD(run, out_level) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_dmr) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_dmk) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_bandgap) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_hs) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_tk) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_l) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_hs2) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_h_t) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_h_vnl) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_h_vl) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_h_vh) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_h_vxc) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_h_exx) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_dh) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_dh_t) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_dh_vl) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_dh_vnl) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_dh_vh) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_dh_vxc) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_dh_exx) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_ds) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_xc) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_xc2) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_eband_terms) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_hr_npz) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_hsr_npz) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_dm_npz) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_interval) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_app_flag) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_ndigits) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_t) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_element_info) +MODULE_CONTEXT_INPUT_FIELD(matrix_output, out_mat_r) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_wfc_lcao) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_dipole) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_efield) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_current) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_current_k) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_vecpot) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, cal_symm_repr) +MODULE_CONTEXT_INPUT_FIELD(restart_output, restart_save) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, rpa) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_pchg) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_wfc_norm) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_wfc_re_im) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, if_separate_k) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_elf) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, out_spillage) +MODULE_CONTEXT_INPUT_FIELD(charge_wavefunction_output, spillage_outdir) +MODULE_CONTEXT_INPUT_FIELD(postprocess, dos_emin_ev) +MODULE_CONTEXT_INPUT_FIELD(postprocess, dos_emax_ev) +MODULE_CONTEXT_INPUT_FIELD(postprocess, dos_edelta_ev) +MODULE_CONTEXT_INPUT_FIELD(postprocess, dos_scale) +MODULE_CONTEXT_INPUT_FIELD(postprocess, dos_sigma) +MODULE_CONTEXT_INPUT_FIELD(postprocess, dos_nche) +MODULE_CONTEXT_INPUT_FIELD(postprocess, stm_bias) +MODULE_CONTEXT_INPUT_FIELD(postprocess, ldos_line) +MODULE_CONTEXT_INPUT_FIELD(postprocess, cal_cond) +MODULE_CONTEXT_INPUT_FIELD(postprocess, cond_che_thr) +MODULE_CONTEXT_INPUT_FIELD(postprocess, cond_dw) +MODULE_CONTEXT_INPUT_FIELD(postprocess, cond_wcut) +MODULE_CONTEXT_INPUT_FIELD(postprocess, cond_dt) +MODULE_CONTEXT_INPUT_FIELD(postprocess, cond_dtbatch) +MODULE_CONTEXT_INPUT_FIELD(postprocess, cond_smear) +MODULE_CONTEXT_INPUT_FIELD(postprocess, cond_fwhm) +MODULE_CONTEXT_INPUT_FIELD(postprocess, cond_nonlocal) +MODULE_CONTEXT_INPUT_FIELD(postprocess, cond_mgga_vel) +MODULE_CONTEXT_INPUT_FIELD(postprocess, berry_phase) +MODULE_CONTEXT_INPUT_FIELD(postprocess, gdir) +MODULE_CONTEXT_INPUT_FIELD(postprocess, towannier90) +MODULE_CONTEXT_INPUT_FIELD(postprocess, nnkpfile) +MODULE_CONTEXT_INPUT_FIELD(postprocess, wannier_spin) +MODULE_CONTEXT_INPUT_FIELD(postprocess, wannier_method) +MODULE_CONTEXT_INPUT_FIELD(postprocess, out_wannier_mmn) +MODULE_CONTEXT_INPUT_FIELD(postprocess, out_wannier_amn) +MODULE_CONTEXT_INPUT_FIELD(postprocess, out_wannier_unk) +MODULE_CONTEXT_INPUT_FIELD(postprocess, out_wannier_eig) +MODULE_CONTEXT_INPUT_FIELD(postprocess, out_wannier_wvfn_formatted) +MODULE_CONTEXT_INPUT_FIELD(model, efield_flag) +MODULE_CONTEXT_INPUT_FIELD(model, dip_cor_flag) +MODULE_CONTEXT_INPUT_FIELD(model, efield_dir) +MODULE_CONTEXT_INPUT_FIELD(model, efield_pos_max) +MODULE_CONTEXT_INPUT_FIELD(model, efield_pos_dec) +MODULE_CONTEXT_INPUT_FIELD(model, efield_amp) +MODULE_CONTEXT_INPUT_FIELD(model, gate_flag) +MODULE_CONTEXT_INPUT_FIELD(model, zgate) +MODULE_CONTEXT_INPUT_FIELD(model, block) +MODULE_CONTEXT_INPUT_FIELD(model, block_down) +MODULE_CONTEXT_INPUT_FIELD(model, block_up) +MODULE_CONTEXT_INPUT_FIELD(model, block_height) +MODULE_CONTEXT_INPUT_FIELD(model, imp_sol) +MODULE_CONTEXT_INPUT_FIELD(model, eb_k) +MODULE_CONTEXT_INPUT_FIELD(model, tau) +MODULE_CONTEXT_INPUT_FIELD(model, sigma_k) +MODULE_CONTEXT_INPUT_FIELD(model, nc_k) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_method) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_s6) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_s8) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_a1) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_a2) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_d) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_abc) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_C6_file) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_C6_unit) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_R0_file) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_R0_unit) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_cutoff_type) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_cutoff_radius) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_radius_unit) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_cn_thr) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_cn_thr_unit) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_d4_xc) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_d4_model) +MODULE_CONTEXT_INPUT_FIELD(vdw, vdw_cutoff_period) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_fock_alpha) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_fock_lambda) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_erfc_alpha) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_erfc_omega) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_separate_loop) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_singularity_correction) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_hybrid_step) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_mixing_beta) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_real_number) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_pca_threshold) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_c_threshold) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_v_threshold) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_dm_threshold) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_c_grad_threshold) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_v_grad_threshold) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_c_grad_r_threshold) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_v_grad_r_threshold) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_ccp_rmesh_times) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_opt_orb_lmax) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_opt_orb_ecut) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_opt_orb_tolerence) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_symmetry_realspace) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, rpa_ccp_rmesh_times) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_cs_inv_thr) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, shrink_abfs_pca_thr) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, shrink_LU_inv_thr) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, out_ri_cv) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, out_unshrinked_v) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_coul_moment) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_rotate_abfs) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_multip_moments_threshold) +MODULE_CONTEXT_INPUT_FIELD(dftu, dft_plus_u) +MODULE_CONTEXT_INPUT_FIELD(dftu, dft_plus_dmft) +MODULE_CONTEXT_INPUT_FIELD(dftu, yukawa_potential) +MODULE_CONTEXT_INPUT_FIELD(dftu, yukawa_lambda) +MODULE_CONTEXT_INPUT_FIELD(dftu, uramping_eV) +MODULE_CONTEXT_INPUT_FIELD(dftu, omc) +MODULE_CONTEXT_INPUT_FIELD(dftu, onsite_radius) +MODULE_CONTEXT_INPUT_FIELD(dftu, hubbard_u_eV) +MODULE_CONTEXT_INPUT_FIELD(dftu, orbital_corr) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, sc_mag_switch) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, decay_grad_switch) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, sc_thr) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, nsc) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, nsc_min) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, alpha_trial) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, sccut) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, sc_scf_thr) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, sc_drop_thr) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, sc_lambda_strategy) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, sc_direction_only) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, sc_scan_lambda_start) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, sc_scan_lambda_end) +MODULE_CONTEXT_INPUT_FIELD(spin_constraint, sc_scan_steps) +MODULE_CONTEXT_INPUT_FIELD(quasi_atomic_orbital, qo_switch) +MODULE_CONTEXT_INPUT_FIELD(quasi_atomic_orbital, qo_basis) +MODULE_CONTEXT_INPUT_FIELD(quasi_atomic_orbital, qo_thr) +MODULE_CONTEXT_INPUT_FIELD(quasi_atomic_orbital, qo_strategy) +MODULE_CONTEXT_INPUT_FIELD(quasi_atomic_orbital, qo_screening_coeff) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_npole) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_inertia) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_nmax) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_comm) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_storage) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_ordering) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_row_ordering) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_nproc) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_symm) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_trans) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_method) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_nproc_pole) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_temp) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_gap) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_delta_e) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_mu_lower) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_mu_upper) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_mu) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_mu_thr) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_mu_expand) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_mu_guard) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_elec_thr) +MODULE_CONTEXT_INPUT_FIELD(pexsi, pexsi_zero_thr) +MODULE_CONTEXT_INPUT_FIELD(output, out_alllog) +MODULE_CONTEXT_INPUT_FIELD(test, nurse) +MODULE_CONTEXT_INPUT_FIELD(test, t_in_h) +MODULE_CONTEXT_INPUT_FIELD(test, vl_in_h) +MODULE_CONTEXT_INPUT_FIELD(test, vnl_in_h) +MODULE_CONTEXT_INPUT_FIELD(test, vh_in_h) +MODULE_CONTEXT_INPUT_FIELD(test, vion_in_h) +MODULE_CONTEXT_INPUT_FIELD(test, test_force) +MODULE_CONTEXT_INPUT_FIELD(test, test_stress) +MODULE_CONTEXT_INPUT_FIELD(test, test_skip_ewald) +MODULE_CONTEXT_INPUT_FIELD(test, test_atom_input) +MODULE_CONTEXT_INPUT_FIELD(test, test_symmetry) +MODULE_CONTEXT_INPUT_FIELD(test, test_wf) +MODULE_CONTEXT_INPUT_FIELD(test, test_grid) +MODULE_CONTEXT_INPUT_FIELD(test, test_charge) +MODULE_CONTEXT_INPUT_FIELD(test, test_energy) +MODULE_CONTEXT_INPUT_FIELD(test, test_gridt) +MODULE_CONTEXT_INPUT_FIELD(test, test_pseudo_cell) +MODULE_CONTEXT_INPUT_FIELD(test, test_pp) +MODULE_CONTEXT_INPUT_FIELD(test, test_relax_method) +MODULE_CONTEXT_INPUT_FIELD(test, test_deconstructor) +MODULE_CONTEXT_INPUT_FIELD(rdmft, rdmft) +MODULE_CONTEXT_INPUT_FIELD(rdmft, rdmft_power_alpha) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exxace) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_gamma_extrapolation) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_thr_type) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, exx_ene_thr) +MODULE_CONTEXT_INPUT_FIELD(exact_exchange, ecutexx) +MODULE_CONTEXT_INPUT_FIELD(external_xc, xc_exch_ext) +MODULE_CONTEXT_INPUT_FIELD(external_xc, xc_corr_ext) +MODULE_CONTEXT_INPUT_FIELD(time_dependent_ofdft, of_cd) +MODULE_CONTEXT_INPUT_FIELD(time_dependent_ofdft, of_mCD_alpha) +MODULE_CONTEXT_INPUT_FIELD(uncommon_hardware, dsp_count) diff --git a/source/source_context/orchestration_context.h b/source/source_context/orchestration_context.h new file mode 100644 index 00000000000..fd692c4cb13 --- /dev/null +++ b/source/source_context/orchestration_context.h @@ -0,0 +1,17 @@ +#ifndef ABACUS_SOURCE_CONTEXT_ORCHESTRATION_CONTEXT_H +#define ABACUS_SOURCE_CONTEXT_ORCHESTRATION_CONTEXT_H + +#ifndef ABACUS_CAN_READ_SIMULATION_CONTEXT +#error "orchestration_context.h is restricted to L0/L1 orchestration targets" +#endif + +#include "source_context/simulation_context.h" + +namespace ModuleContext +{ + +const SimulationContext& current_simulation_context(); + +} // namespace ModuleContext + +#endif diff --git a/source/source_context/restart_field_mapping.inc b/source/source_context/restart_field_mapping.inc new file mode 100644 index 00000000000..2c7d4a76e4e --- /dev/null +++ b/source/source_context/restart_field_mapping.inc @@ -0,0 +1,9 @@ +// MODULE_CONTEXT_LEGACY_SCHEMA_SHA256(9ad81b8fc2e5be2e155a2e96c07a6f89a6a4ded6be9246d1c323d3175abd2e55) +MODULE_CONTEXT_RESTART_FIELD(info_save.save_charge) +MODULE_CONTEXT_RESTART_FIELD(info_save.save_H) +MODULE_CONTEXT_RESTART_FIELD(info_load.load_charge) +MODULE_CONTEXT_RESTART_FIELD(info_load.load_charge_finish) +MODULE_CONTEXT_RESTART_FIELD(info_load.load_H) +MODULE_CONTEXT_RESTART_FIELD(info_load.load_H_finish) +MODULE_CONTEXT_RESTART_FIELD(info_load.restart_exx) +MODULE_CONTEXT_RESTART_FIELD(folder) diff --git a/source/source_context/simulation_context.h b/source/source_context/simulation_context.h new file mode 100644 index 00000000000..a06fca94af4 --- /dev/null +++ b/source/source_context/simulation_context.h @@ -0,0 +1,66 @@ +#ifndef ABACUS_SOURCE_CONTEXT_SIMULATION_CONTEXT_H +#define ABACUS_SOURCE_CONTEXT_SIMULATION_CONTEXT_H + +#if !defined(ABACUS_CAN_READ_SIMULATION_CONTEXT) && !defined(ABACUS_CAN_BIND_SIMULATION_CONTEXT) +#error "simulation_context.h is restricted to L0/L1 Context-capable targets" +#endif + +#include "source_context/context_types.h" + +namespace ModuleContext +{ + +struct SimulationContext +{ + RunControl run; + FileSystemLayout files; + ParallelTopology parallel; + LogStreams logs; + BasisInfo basis; + GridInfo grid; + ElectronicRuntimeInfo electronic_runtime; + FeatureRuntimeInfo features; + SpinConfig spin; + SolverConfig solver; + MatrixOutputConfig matrix_output; + GeneralOutputConfig output; + DftUConfig dftu; + ExactExchangeConfig exact_exchange; + + std::shared_ptr exact_exchange_state; + std::shared_ptr restart; + + // Full input-domain type skeleton. The root aggregate is restricted to + // Driver/ESolver orchestration and must not be passed into leaf modules. + CellConfig cell; + FileInputConfig file_input; + ParallelConfig parallel_input; + DeviceConfig device; + ElectronicStructureConfig electronic_structure; + ScfMixingConfig scf_mixing; + LcaoConfig lcao; + RelaxationConfig relaxation; + MolecularDynamicsConfig molecular_dynamics; + OfdftConfig ofdft; + StochasticDftConfig stochastic_dft; + DeepksConfig deepks; + RealTimeTddftConfig real_time_tddft; + LinearResponseTddftConfig linear_response_tddft; + ChargeWavefunctionOutputConfig charge_wavefunction_output; + RestartOutputConfig restart_output; + PostprocessConfig postprocess; + ModelConfig model; + VdwConfig vdw; + SpinConstraintConfig spin_constraint; + QuasiAtomicOrbitalConfig quasi_atomic_orbital; + PexsiConfig pexsi; + TestConfig test; + RdmftConfig rdmft; + ExternalXcConfig external_xc; + TimeDependentOfdftConfig time_dependent_ofdft; + UncommonHardwareConfig uncommon_hardware; +}; + +} // namespace ModuleContext + +#endif diff --git a/source/source_context/simulation_context_binding.cpp b/source/source_context/simulation_context_binding.cpp new file mode 100644 index 00000000000..72f00ea10f8 --- /dev/null +++ b/source/source_context/simulation_context_binding.cpp @@ -0,0 +1,48 @@ +#include "source_context/orchestration_context.h" +#include "source_context/simulation_context_binding.h" + +#include +#include + +namespace +{ + +std::atomic active_context(nullptr); + +} // namespace + +namespace ModuleContext +{ + +ScopedSimulationContextBinding::ScopedSimulationContextBinding(const SimulationContext& context) : context_(&context) +{ + const SimulationContext* expected = nullptr; + if (!active_context.compare_exchange_strong(expected, + context_, + std::memory_order_release, + std::memory_order_relaxed)) + { + throw std::logic_error("a SimulationContext is already bound"); + } +} + +ScopedSimulationContextBinding::~ScopedSimulationContextBinding() +{ + const SimulationContext* expected = context_; + active_context.compare_exchange_strong(expected, + nullptr, + std::memory_order_release, + std::memory_order_relaxed); +} + +const SimulationContext& current_simulation_context() +{ + const SimulationContext* context = active_context.load(std::memory_order_acquire); + if (context == nullptr) + { + throw std::logic_error("SimulationContext has not been bound"); + } + return *context; +} + +} // namespace ModuleContext diff --git a/source/source_context/simulation_context_binding.h b/source/source_context/simulation_context_binding.h new file mode 100644 index 00000000000..24beb84120e --- /dev/null +++ b/source/source_context/simulation_context_binding.h @@ -0,0 +1,28 @@ +#ifndef ABACUS_SOURCE_CONTEXT_SIMULATION_CONTEXT_BINDING_H +#define ABACUS_SOURCE_CONTEXT_SIMULATION_CONTEXT_BINDING_H + +#ifndef ABACUS_CAN_BIND_SIMULATION_CONTEXT +#error "simulation_context_binding.h is restricted to the assembly layer" +#endif + +namespace ModuleContext +{ + +struct SimulationContext; + +class ScopedSimulationContextBinding +{ + public: + explicit ScopedSimulationContextBinding(const SimulationContext& context); + ~ScopedSimulationContextBinding(); + + ScopedSimulationContextBinding(const ScopedSimulationContextBinding&) = delete; + ScopedSimulationContextBinding& operator=(const ScopedSimulationContextBinding&) = delete; + + private: + const SimulationContext* context_; +}; + +} // namespace ModuleContext + +#endif diff --git a/source/source_context/simulation_context_builder.cpp b/source/source_context/simulation_context_builder.cpp new file mode 100644 index 00000000000..0d4b480d608 --- /dev/null +++ b/source/source_context/simulation_context_builder.cpp @@ -0,0 +1,370 @@ +#include "source_context/simulation_context_builder.h" + +#include "source_base/global_variable.h" +#include "source_hamilt/module_xc/exx_info.h" +#include "source_io/module_parameter/input_parameter.h" +#include "source_io/module_parameter/system_parameter.h" +#include "source_io/module_restart/restart.h" + +#include +#include + +namespace ModuleContext +{ +namespace +{ + +MatrixOutputRequest request_from(const std::vector& value) +{ + MatrixOutputRequest request; + if (!value.empty()) + { + request.enabled = value[0] != 0; + } + if (value.size() > 1) + { + request.precision = value[1]; + } + if (value.size() > 2) + { + request.mode = value[2]; + } + return request; +} + +bool field_is(const char* field, const char* expected) +{ + return std::strcmp(field, expected) == 0; +} + +template +void project_input_field(Domain& domain, const char* field, const T& value) +{ + domain.project(field, value); +} + +template +void project_input_field(RunControl& domain, const char* field, const T& value) +{ + if (field_is(field, "calculation") || field_is(field, "esolver_type") || field_is(field, "out_level") + || field_is(field, "cal_force") || field_is(field, "cal_stress")) + { + return; + } + domain.project(field, value); +} + +template +void project_input_field(SpinConfig& domain, const char* field, const T& value) +{ + if (field_is(field, "nspin") || field_is(field, "lspinorb") || field_is(field, "noncolin")) + { + return; + } + domain.project(field, value); +} + +template +void project_input_field(SolverConfig& domain, const char* field, const T& value) +{ + if (field_is(field, "ks_solver") || field_is(field, "basis_type")) + { + return; + } + domain.project(field, value); +} + +template +void project_input_field(MatrixOutputConfig& domain, const char* field, const T& value) +{ + if (field_is(field, "out_mat_hs") || field_is(field, "out_mat_hs2") + || field_is(field, "out_mat_tk") || field_is(field, "out_mat_t") || field_is(field, "out_mat_r") + || field_is(field, "out_mat_dh") || field_is(field, "out_mat_ds") + || field_is(field, "out_mat_h_t") || field_is(field, "out_mat_h_vnl") + || field_is(field, "out_mat_h_vl") || field_is(field, "out_mat_h_vh") + || field_is(field, "out_mat_h_vxc") || field_is(field, "out_mat_h_exx") + || field_is(field, "out_mat_dh_t") || field_is(field, "out_mat_dh_vnl") + || field_is(field, "out_mat_dh_vl") || field_is(field, "out_mat_dh_vh") + || field_is(field, "out_mat_dh_vxc") || field_is(field, "out_mat_dh_exx") + || field_is(field, "out_mat_xc2") || field_is(field, "out_mat_xc") + || field_is(field, "out_eband_terms") || field_is(field, "out_app_flag") + || field_is(field, "out_ndigits")) + { + return; + } + domain.project(field, value); +} + +template +void project_input_field(GeneralOutputConfig& domain, const char* field, const T& value) +{ + if (field_is(field, "out_freq_elec") || field_is(field, "out_freq_ion") + || field_is(field, "out_freq_td") || field_is(field, "out_alllog")) + { + return; + } + domain.project(field, value); +} + +template +void project_input_field(DftUConfig& domain, const char* field, const T& value) +{ + if (field_is(field, "dft_plus_u") || field_is(field, "dft_plus_dmft")) + { + return; + } + domain.project(field, value); +} + +template +void project_input_field(ExactExchangeConfig& domain, const char* field, const T& value) +{ + if (field_is(field, "exx_separate_loop") || field_is(field, "exx_hybrid_step") + || field_is(field, "exx_mixing_beta") || field_is(field, "exx_symmetry_realspace")) + { + return; + } + domain.project(field, value); +} + +template +void project_exx_state_field(ExactExchangeState& state, const char* field, const T& value) +{ + if (field_is(field, "info_global.cal_exx") || field_is(field, "info_global.hybrid_alpha") + || field_is(field, "info_global.hse_omega") || field_is(field, "info_ri.real_number")) + { + return; + } + state.project(field, value); +} + +class LegacyExactExchangeStateService : public ExactExchangeStateService +{ + public: + ExactExchangeState snapshot() const override + { + ExactExchangeState state; +#define MODULE_CONTEXT_EXX_STATE_FIELD(section, field) \ + project_exx_state_field(state, #section "." #field, GlobalC::exx_info.section.field); +#include "source_context/exx_state_field_mapping.inc" +#undef MODULE_CONTEXT_EXX_STATE_FIELD + state.enabled = GlobalC::exx_info.info_global.cal_exx; + state.real_number = GlobalC::exx_info.info_ri.real_number; + state.hybrid_alpha = GlobalC::exx_info.info_global.hybrid_alpha; + state.hse_omega = GlobalC::exx_info.info_global.hse_omega; + return state; + } +}; + +class LegacyRestartService : public RestartService +{ + public: + RestartState snapshot() const override + { + RestartState state; + state.save_charge = GlobalC::restart.info_save.save_charge; + state.save_hamiltonian = GlobalC::restart.info_save.save_H; + state.load_charge = GlobalC::restart.info_load.load_charge; + state.load_charge_finished = GlobalC::restart.info_load.load_charge_finish; + state.load_hamiltonian = GlobalC::restart.info_load.load_H; + state.load_hamiltonian_finished = GlobalC::restart.info_load.load_H_finish; + state.restart_exact_exchange = GlobalC::restart.info_load.restart_exx; + state.folder = GlobalC::restart.folder; + return state; + } + + void update(const RestartState& state) override + { + GlobalC::restart.info_save.save_charge = state.save_charge; + GlobalC::restart.info_save.save_H = state.save_hamiltonian; + GlobalC::restart.info_load.load_charge = state.load_charge; + GlobalC::restart.info_load.load_charge_finish = state.load_charge_finished; + GlobalC::restart.info_load.load_H = state.load_hamiltonian; + GlobalC::restart.info_load.load_H_finish = state.load_hamiltonian_finished; + GlobalC::restart.info_load.restart_exx = state.restart_exact_exchange; + GlobalC::restart.folder = state.folder; + } +}; + +void capture_input(SimulationContext& context, const Input_para& input) +{ +#define MODULE_CONTEXT_INPUT_FIELD(domain, field) project_input_field(context.domain, #field, input.field); +#include "source_context/input_field_mapping.inc" +#undef MODULE_CONTEXT_INPUT_FIELD + + context.run.calculation = input.calculation; + context.run.esolver_type = input.esolver_type; + context.run.output_level = input.out_level; + context.run.cal_force = input.cal_force; + context.run.cal_stress = input.cal_stress; + + context.spin.nspin = input.nspin; + context.spin.spin_orbit = input.lspinorb; + context.spin.noncollinear = input.noncolin; + + context.solver.ks_solver = input.ks_solver; + context.solver.basis_type = input.basis_type; + + context.matrix_output.hs_k = request_from(input.out_mat_hs); + context.matrix_output.hs_r = request_from(input.out_mat_hs2); + context.matrix_output.kinetic_k = request_from(input.out_mat_tk); + context.matrix_output.kinetic_r = request_from(input.out_mat_t); + context.matrix_output.position_r = request_from(input.out_mat_r); + context.matrix_output.dh = request_from(input.out_mat_dh); + context.matrix_output.ds = request_from(input.out_mat_ds); + context.matrix_output.h_t = request_from(input.out_mat_h_t); + context.matrix_output.h_vnl = request_from(input.out_mat_h_vnl); + context.matrix_output.h_vl = request_from(input.out_mat_h_vl); + context.matrix_output.h_vh = request_from(input.out_mat_h_vh); + context.matrix_output.h_vxc = request_from(input.out_mat_h_vxc); + context.matrix_output.h_exx = request_from(input.out_mat_h_exx); + context.matrix_output.dh_t = request_from(input.out_mat_dh_t); + context.matrix_output.dh_vnl = request_from(input.out_mat_dh_vnl); + context.matrix_output.dh_vl = request_from(input.out_mat_dh_vl); + context.matrix_output.dh_vh = request_from(input.out_mat_dh_vh); + context.matrix_output.dh_vxc = request_from(input.out_mat_dh_vxc); + context.matrix_output.dh_exx = request_from(input.out_mat_dh_exx); + context.matrix_output.vxc_r = request_from(input.out_mat_xc2); + context.matrix_output.vxc_k = input.out_mat_xc; + context.matrix_output.band_energy_terms = input.out_eband_terms; + context.matrix_output.append = input.out_app_flag; + context.matrix_output.digits = input.out_ndigits; + + context.output.electronic_frequency = input.out_freq_elec; + context.output.ionic_frequency = input.out_freq_ion; + context.output.tddft_frequency = input.out_freq_td; + context.output.all_logs = input.out_alllog; + + context.dftu.enabled = input.dft_plus_u != 0; + context.dftu.dmft_enabled = input.dft_plus_dmft; + + context.exact_exchange.separate_loop = input.exx_separate_loop; + context.exact_exchange.hybrid_step = input.exx_hybrid_step; + context.exact_exchange.mixing_beta = input.exx_mixing_beta; + context.exact_exchange.symmetry_realspace = input.exx_symmetry_realspace; +} + +void capture_system(SimulationContext& context, const System_para& system) +{ +#define MODULE_CONTEXT_SYSTEM_FIELD(domain, field) static_cast(system.field); +#include "source_context/system_field_mapping.inc" +#undef MODULE_CONTEXT_SYSTEM_FIELD + + context.run.start_time = system.start_time; + context.files.input_card = system.global_in_card; + context.files.structure_file = system.global_in_stru; + context.files.output_directory = system.global_out_dir; + context.files.readin_directory = system.global_readin_dir; + context.files.structure_directory = system.global_stru_dir; + context.files.matrix_directory = system.global_matrix_dir; + context.files.wavefunction_directory = system.global_wfc_dir; + context.files.mlkedf_descriptor_directory = system.global_mlkedf_descriptor_dir; + context.files.deepks_label_directory = system.global_deepks_label_elec_dir; + context.files.log_file = system.log_file; + + context.parallel.world_rank = system.myrank; + context.parallel.world_size = system.nproc; + context.parallel.pool_index = system.mypool; + context.parallel.kpoint_pool_count = system.npool; + context.parallel.processes_in_pool = system.nproc_in_pool; + context.parallel.threads_per_process = system.nthread_per_proc; + + context.basis.nlocal = system.nlocal; + context.basis.npol = system.npol; + context.basis.gamma_only_pw = system.gamma_only_pw; + context.basis.gamma_only_local = system.gamma_only_local; + + context.spin.domag = system.domag; + context.spin.domag_z = system.domag_z; + + context.grid.charge_nx = system.ncx; + context.grid.charge_ny = system.ncy; + context.grid.charge_nz = system.ncz; + context.grid.double_grid = system.double_grid; + context.grid.radial_dq = system.dq; + context.grid.radial_nqx = system.nqx; + context.grid.radial_nqxq = system.nqxq; + + context.electronic_runtime.two_fermi = system.two_fermi; + context.electronic_runtime.use_uspp = system.use_uspp; + context.electronic_runtime.dos_minimum_is_explicit = system.dos_setemin; + context.electronic_runtime.dos_maximum_is_explicit = system.dos_setemax; + context.electronic_runtime.local_band_count = system.nbands_l; + context.electronic_runtime.ks_run = system.ks_run; + context.electronic_runtime.all_ks_run = system.all_ks_run; + context.electronic_runtime.has_double_data = system.has_double_data; + context.electronic_runtime.has_float_data = system.has_float_data; + + context.features.deepks_setorb = system.deepks_setorb; + context.features.search_periodic_boundaries = system.search_pbc; + context.features.lcao_kpoint_pool_count = system.kpar_lcao; + + context.output.molecular_dynamics_control = system.out_md_control; + context.dftu.ramping_ry = system.uramping; + context.dftu.hubbard_u_ry = system.hubbard_u; +} + +} // namespace + +SimulationContextBuilder::SimulationContextBuilder(const Input_para& input, const System_para& system) +{ + capture_input(context_, input); + capture_system(context_, system); + + context_.logs.running = &GlobalV::ofs_running; + context_.logs.warning = &GlobalV::ofs_warning; + context_.logs.information = &GlobalV::ofs_info; + context_.logs.device = &GlobalV::ofs_device; + context_.exact_exchange_state.reset(new LegacyExactExchangeStateService()); + context_.restart.reset(new LegacyRestartService()); +} + +void SimulationContextBuilder::capture_runtime(const System_para& system) +{ + if (finalized_) + { + throw std::logic_error("SimulationContextBuilder::capture_runtime called after finalize"); + } + capture_system(context_, system); + runtime_captured_ = true; +} + +SimulationContext SimulationContextBuilder::finalize(const System_para& system) +{ + if (finalized_) + { + throw std::logic_error("SimulationContextBuilder::finalize called more than once"); + } + if (!runtime_captured_) + { + throw std::logic_error("SimulationContextBuilder::finalize called before runtime initialization"); + } + if (system.myrank != GlobalV::MY_RANK || system.nproc != GlobalV::NPROC) + { + throw std::logic_error("legacy PARAM and GlobalV process topology disagree"); + } + + capture_system(context_, system); +#define MODULE_CONTEXT_GLOBALV_FIELD(domain, field) static_cast(GlobalV::field); +#include "source_context/globalv_field_mapping.inc" +#undef MODULE_CONTEXT_GLOBALV_FIELD + context_.parallel.world_rank = GlobalV::MY_RANK; + context_.parallel.world_size = GlobalV::NPROC; + context_.parallel.kpoint_pool_count = GlobalV::KPAR; + context_.parallel.pool_index = GlobalV::MY_POOL; + context_.parallel.processes_in_pool = GlobalV::NPROC_IN_POOL; + context_.parallel.rank_in_pool = GlobalV::RANK_IN_POOL; + context_.parallel.band_group_index = GlobalV::MY_BNDGROUP; + context_.parallel.processes_in_band_group = GlobalV::NPROC_IN_BNDGROUP; + context_.parallel.rank_in_band_group = GlobalV::RANK_IN_BPGROUP; + context_.parallel.diagonalization_rank = GlobalV::DRANK; + context_.parallel.diagonalization_size = GlobalV::DSIZE; + context_.parallel.diagonalization_color = GlobalV::DCOLOR; + context_.parallel.grid_rank = GlobalV::GRANK; + context_.parallel.grid_size = GlobalV::GSIZE; + + finalized_ = true; + return context_; +} + +} // namespace ModuleContext diff --git a/source/source_context/simulation_context_builder.h b/source/source_context/simulation_context_builder.h new file mode 100644 index 00000000000..8a0dbd09769 --- /dev/null +++ b/source/source_context/simulation_context_builder.h @@ -0,0 +1,34 @@ +#ifndef ABACUS_SOURCE_CONTEXT_SIMULATION_CONTEXT_BUILDER_H +#define ABACUS_SOURCE_CONTEXT_SIMULATION_CONTEXT_BUILDER_H + +#ifndef ABACUS_CAN_BIND_SIMULATION_CONTEXT +#error "simulation_context_builder.h is restricted to the assembly layer" +#endif + +#include "source_context/simulation_context.h" + +struct Input_para; +struct System_para; + +namespace ModuleContext +{ + +class SimulationContextBuilder +{ + public: + SimulationContextBuilder(const Input_para& input, const System_para& system); + + void capture_runtime(const System_para& system); + SimulationContext finalize(const System_para& system); + bool runtime_captured() const { return runtime_captured_; } + bool finalized() const { return finalized_; } + + private: + SimulationContext context_; + bool runtime_captured_ = false; + bool finalized_ = false; +}; + +} // namespace ModuleContext + +#endif diff --git a/source/source_context/system_field_mapping.inc b/source/source_context/system_field_mapping.inc new file mode 100644 index 00000000000..7477ac48b1e --- /dev/null +++ b/source/source_context/system_field_mapping.inc @@ -0,0 +1,46 @@ +// MODULE_CONTEXT_LEGACY_SCHEMA_SHA256(398ec9ce4683dd10de5bd6f2796708fa0a0b6033ab393affec311bd931775a41) +MODULE_CONTEXT_SYSTEM_FIELD(parallel, myrank) +MODULE_CONTEXT_SYSTEM_FIELD(parallel, nproc) +MODULE_CONTEXT_SYSTEM_FIELD(parallel, nthread_per_proc) +MODULE_CONTEXT_SYSTEM_FIELD(parallel, mypool) +MODULE_CONTEXT_SYSTEM_FIELD(parallel, npool) +MODULE_CONTEXT_SYSTEM_FIELD(parallel, nproc_in_pool) +MODULE_CONTEXT_SYSTEM_FIELD(run, start_time) +MODULE_CONTEXT_SYSTEM_FIELD(basis, nlocal) +MODULE_CONTEXT_SYSTEM_FIELD(electronic_runtime, two_fermi) +MODULE_CONTEXT_SYSTEM_FIELD(electronic_runtime, use_uspp) +MODULE_CONTEXT_SYSTEM_FIELD(electronic_runtime, dos_setemin) +MODULE_CONTEXT_SYSTEM_FIELD(electronic_runtime, dos_setemax) +MODULE_CONTEXT_SYSTEM_FIELD(grid, dq) +MODULE_CONTEXT_SYSTEM_FIELD(grid, nqx) +MODULE_CONTEXT_SYSTEM_FIELD(grid, nqxq) +MODULE_CONTEXT_SYSTEM_FIELD(grid, ncx) +MODULE_CONTEXT_SYSTEM_FIELD(grid, ncy) +MODULE_CONTEXT_SYSTEM_FIELD(grid, ncz) +MODULE_CONTEXT_SYSTEM_FIELD(output, out_md_control) +MODULE_CONTEXT_SYSTEM_FIELD(basis, gamma_only_pw) +MODULE_CONTEXT_SYSTEM_FIELD(basis, gamma_only_local) +MODULE_CONTEXT_SYSTEM_FIELD(files, global_in_card) +MODULE_CONTEXT_SYSTEM_FIELD(files, global_in_stru) +MODULE_CONTEXT_SYSTEM_FIELD(files, global_out_dir) +MODULE_CONTEXT_SYSTEM_FIELD(files, global_readin_dir) +MODULE_CONTEXT_SYSTEM_FIELD(files, global_stru_dir) +MODULE_CONTEXT_SYSTEM_FIELD(files, global_matrix_dir) +MODULE_CONTEXT_SYSTEM_FIELD(files, global_wfc_dir) +MODULE_CONTEXT_SYSTEM_FIELD(files, global_mlkedf_descriptor_dir) +MODULE_CONTEXT_SYSTEM_FIELD(files, global_deepks_label_elec_dir) +MODULE_CONTEXT_SYSTEM_FIELD(files, log_file) +MODULE_CONTEXT_SYSTEM_FIELD(features, deepks_setorb) +MODULE_CONTEXT_SYSTEM_FIELD(basis, npol) +MODULE_CONTEXT_SYSTEM_FIELD(spin, domag) +MODULE_CONTEXT_SYSTEM_FIELD(spin, domag_z) +MODULE_CONTEXT_SYSTEM_FIELD(grid, double_grid) +MODULE_CONTEXT_SYSTEM_FIELD(dftu, uramping) +MODULE_CONTEXT_SYSTEM_FIELD(dftu, hubbard_u) +MODULE_CONTEXT_SYSTEM_FIELD(features, kpar_lcao) +MODULE_CONTEXT_SYSTEM_FIELD(electronic_runtime, nbands_l) +MODULE_CONTEXT_SYSTEM_FIELD(electronic_runtime, ks_run) +MODULE_CONTEXT_SYSTEM_FIELD(electronic_runtime, all_ks_run) +MODULE_CONTEXT_SYSTEM_FIELD(electronic_runtime, has_double_data) +MODULE_CONTEXT_SYSTEM_FIELD(electronic_runtime, has_float_data) +MODULE_CONTEXT_SYSTEM_FIELD(features, search_pbc) diff --git a/source/source_context/test/CMakeLists.txt b/source/source_context/test/CMakeLists.txt new file mode 100644 index 00000000000..916b1da8d7d --- /dev/null +++ b/source/source_context/test/CMakeLists.txt @@ -0,0 +1,30 @@ +AddTest( + TARGET MODULE_CONTEXT_builder_test + LIBS context parameter base device + SOURCES simulation_context_builder_test.cpp +) + +target_compile_definitions( + MODULE_CONTEXT_builder_test + PRIVATE + ABACUS_CAN_BIND_SIMULATION_CONTEXT + ABACUS_CAN_READ_SIMULATION_CONTEXT) + +AddTest( + TARGET MODULE_CONTEXT_binding_test + LIBS context parameter base device + SOURCES simulation_context_binding_test.cpp +) + +target_compile_definitions( + MODULE_CONTEXT_binding_test + PRIVATE + ABACUS_CAN_BIND_SIMULATION_CONTEXT + ABACUS_CAN_READ_SIMULATION_CONTEXT) + +find_package(Python3 REQUIRED COMPONENTS Interpreter) +add_test( + NAME MODULE_CONTEXT_mapping_manifest_test + COMMAND ${Python3_EXECUTABLE} + ${CMAKE_CURRENT_SOURCE_DIR}/check_context_mapping.py + --repo-root ${PROJECT_SOURCE_DIR}) diff --git a/source/source_context/test/check_context_mapping.py b/source/source_context/test/check_context_mapping.py new file mode 100644 index 00000000000..7985e78f6a4 --- /dev/null +++ b/source/source_context/test/check_context_mapping.py @@ -0,0 +1,360 @@ +#!/usr/bin/env python3 +"""Verify that the legacy-state manifests cover every declared data member.""" + +from __future__ import print_function + +import argparse +import collections +import hashlib +import re +import sys +from pathlib import Path + + +def strip_comments(text): + result = [] + index = 0 + state = "code" + quote = "" + while index < len(text): + char = text[index] + following = text[index + 1] if index + 1 < len(text) else "" + if state == "line_comment": + if char == "\n": + result.append(char) + state = "code" + index += 1 + continue + if state == "block_comment": + if char == "*" and following == "/": + state = "code" + index += 2 + else: + if char == "\n": + result.append(char) + index += 1 + continue + if state == "string": + result.append(char) + if char == "\\": + if following: + result.append(following) + index += 2 + continue + elif char == quote: + state = "code" + index += 1 + continue + if char == "/" and following == "/": + state = "line_comment" + index += 2 + elif char == "/" and following == "*": + state = "block_comment" + index += 2 + else: + result.append(char) + if char in ('"', "'"): + state = "string" + quote = char + index += 1 + return "".join(result) + + +def block_body(text, keyword, name): + clean = strip_comments(text) + match = re.search(r"\b" + re.escape(keyword) + r"\s+" + re.escape(name) + r"\b[^\{]*\{", clean) + if not match: + raise ValueError("cannot find {} {}".format(keyword, name)) + start = match.end() + depth = 1 + index = start + quote = "" + while index < len(clean): + char = clean[index] + if quote: + if char == "\\": + index += 2 + continue + if char == quote: + quote = "" + elif char in ('"', "'"): + quote = char + elif char == "{": + depth += 1 + elif char == "}": + depth -= 1 + if depth == 0: + return clean[start:index] + index += 1 + raise ValueError("unterminated {} {}".format(keyword, name)) + + +def top_level_parts(text, separator): + parts = [] + start = 0 + braces = brackets = parentheses = angles = 0 + quote = "" + index = 0 + while index < len(text): + char = text[index] + if quote: + if char == "\\": + index += 2 + continue + if char == quote: + quote = "" + elif char in ('"', "'"): + quote = char + elif char == "{": + braces += 1 + elif char == "}": + braces -= 1 + elif char == "[": + brackets += 1 + elif char == "]": + brackets -= 1 + elif char == "(": + parentheses += 1 + elif char == ")": + parentheses -= 1 + elif char == "<": + angles += 1 + elif char == ">" and angles: + angles -= 1 + elif char == separator and not any((braces, brackets, parentheses, angles)): + parts.append(text[start:index]) + start = index + 1 + index += 1 + parts.append(text[start:]) + return parts + + +def normalize_type(type_name): + normalized = re.sub(r"\s+", " ", type_name.strip()) + normalized = re.sub(r"\s*([<>,*&:\[\]])\s*", r"\1", normalized) + return re.sub(r"^(?:extern|mutable)\s+", "", normalized) + + +def data_member_schema(text, kind, name): + body = block_body(text, kind, name) + body = "\n".join(line for line in body.splitlines() if not line.lstrip().startswith("#")) + members = [] + for statement in top_level_parts(body, ";"): + statement = statement.strip() + if not statement or statement.endswith(":"): + continue + if statement.startswith(("using ", "typedef ", "static_assert", "struct ", "class ")): + continue + if "(" in top_level_parts(statement, "=")[0]: + continue + declarators = top_level_parts(statement, ",") + base_type = "" + for index, declarator in enumerate(declarators): + declaration = top_level_parts(declarator, "=")[0].strip() + match = re.search(r"([A-Za-z_]\w*)\s*(?:\[[^\]]*\]\s*)*$", declaration) + if match: + prefix = declaration[: match.start()].strip() + if index == 0: + base_type = prefix + member_type = base_type + else: + member_type = base_type + (" " + prefix if prefix else "") + members.append((match.group(1), normalize_type(member_type))) + return members + + +def data_members(text, kind, name): + return [field for field, _ in data_member_schema(text, kind, name)] + + +def manifest_fields(text, macro, path_field=False): + if path_field: + pattern = re.compile(r"\b" + re.escape(macro) + r"\(\s*([^\)]+?)\s*\)") + return [match.group(1).strip() for match in pattern.finditer(text)] + pattern = re.compile( + r"\b" + re.escape(macro) + r"\(\s*([A-Za-z_]\w*)\s*,\s*([A-Za-z_]\w*)\s*\)" + ) + return [match.group(2) for match in pattern.finditer(text)] + + +def compare(label, declared, mapped): + errors = [] + declared_counts = collections.Counter(declared) + mapped_counts = collections.Counter(mapped) + duplicate_declarations = sorted(name for name, count in declared_counts.items() if count != 1) + duplicate_mappings = sorted(name for name, count in mapped_counts.items() if count != 1) + missing = sorted(set(declared_counts) - set(mapped_counts)) + extra = sorted(set(mapped_counts) - set(declared_counts)) + if duplicate_declarations: + errors.append("{} duplicate declarations: {}".format(label, ", ".join(duplicate_declarations))) + if duplicate_mappings: + errors.append("{} duplicate mappings: {}".format(label, ", ".join(duplicate_mappings))) + if missing: + errors.append("{} fields missing from manifest: {}".format(label, ", ".join(missing))) + if extra: + errors.append("{} manifest fields not declared by legacy type: {}".format(label, ", ".join(extra))) + return errors + + +def schema_digest(schema): + canonical = "".join("{}:{}\n".format(field, type_name) for field, type_name in sorted(schema)) + return hashlib.sha256(canonical.encode("utf-8")).hexdigest() + + +def manifest_schema_digest(text): + match = re.search(r"MODULE_CONTEXT_LEGACY_SCHEMA_SHA256\(([0-9a-f]{64})\)", text) + return match.group(1) if match else None + + +def compare_schema_digest(label, schema, manifest): + expected = manifest_schema_digest(manifest) + actual = schema_digest(schema) + if expected == actual: + return [] + if expected is None: + return ["{} manifest has no legacy schema type fingerprint (actual {})".format(label, actual)] + return [ + "{} legacy field types changed: manifest fingerprint {}, actual {}".format(label, expected, actual) + ] + + +def read(root, relative): + return (root / relative).read_text(encoding="utf-8") + + +def collect_schemas(root): + context = root / "source/source_context" + input_header = read(root, "source/source_io/module_parameter/input_parameter.h") + input_manifest = read(context, "input_field_mapping.inc") + input_schema = data_member_schema(input_header, "struct", "Input_para") + + system_header = read(root, "source/source_io/module_parameter/system_parameter.h") + system_manifest = read(context, "system_field_mapping.inc") + system_schema = data_member_schema(system_header, "struct", "System_para") + + global_header = read(root, "source/source_base/global_variable.h") + global_manifest = read(context, "globalv_field_mapping.inc") + global_schema = data_member_schema(global_header, "namespace", "GlobalV") + + exx_sections = { + "info_global": ("source/source_hamilt/module_xc/exx_info_global.h", "Exx_Info_Global"), + "info_lip": ("source/source_hamilt/module_xc/exx_info_lip.h", "Exx_Info_Lip"), + "info_ri": ("source/source_hamilt/module_xc/exx_info_ri.h", "Exx_Info_RI"), + "info_opt_abfs": ("source/source_hamilt/module_xc/exx_info_opt_abfs.h", "Exx_Info_Opt_ABFs"), + } + exx_schema = [] + for section, (header_path, type_name) in exx_sections.items(): + exx_schema.extend( + (section + "." + field, type_name_value) + for field, type_name_value in data_member_schema(read(root, header_path), "struct", type_name) + ) + exx_manifest = read(context, "exx_state_field_mapping.inc") + + restart_header = read(root, "source/source_io/module_restart/restart.h") + restart_body = block_body(restart_header, "class", "Restart") + save_member = re.search(r"\bInfo_Save\s+([A-Za-z_]\w*)\s*;", strip_comments(restart_body)) + load_member = re.search(r"\bInfo_Load\s+([A-Za-z_]\w*)\s*;", strip_comments(restart_body)) + restart_schema = [] + if save_member and load_member: + restart_schema.extend( + (save_member.group(1) + "." + field, type_name_value) + for field, type_name_value in data_member_schema(restart_header, "struct", "Info_Save") + ) + restart_schema.extend( + (load_member.group(1) + "." + field, type_name_value) + for field, type_name_value in data_member_schema(restart_header, "struct", "Info_Load") + ) + restart_schema.append(("folder", "std::string")) + restart_manifest = read(context, "restart_field_mapping.inc") + + return { + "Input_para": (input_schema, input_manifest), + "System_para": (system_schema, system_manifest), + "GlobalV": (global_schema, global_manifest), + "GlobalC::exx_info": (exx_schema, exx_manifest), + "GlobalC::restart": (restart_schema, restart_manifest), + }, save_member, load_member + + +def check(root): + errors = [] + schemas, save_member, load_member = collect_schemas(root) + + input_schema, input_manifest = schemas["Input_para"] + errors.extend( + compare( + "Input_para", + [field for field, _ in input_schema], + manifest_fields(input_manifest, "MODULE_CONTEXT_INPUT_FIELD"), + ) + ) + errors.extend(compare_schema_digest("Input_para", input_schema, input_manifest)) + + system_schema, system_manifest = schemas["System_para"] + errors.extend( + compare( + "System_para", + [field for field, _ in system_schema], + manifest_fields(system_manifest, "MODULE_CONTEXT_SYSTEM_FIELD"), + ) + ) + errors.extend(compare_schema_digest("System_para", system_schema, system_manifest)) + + global_schema, global_manifest = schemas["GlobalV"] + errors.extend( + compare( + "GlobalV", + [field for field, _ in global_schema], + manifest_fields(global_manifest, "MODULE_CONTEXT_GLOBALV_FIELD"), + ) + ) + errors.extend(compare_schema_digest("GlobalV", global_schema, global_manifest)) + + exx_schema, exx_manifest = schemas["GlobalC::exx_info"] + declared_exx = [field for field, _ in exx_schema] + mapped_exx = [] + for match in re.finditer( + r"\bMODULE_CONTEXT_EXX_STATE_FIELD\(\s*([A-Za-z_]\w*)\s*,\s*([A-Za-z_]\w*)\s*\)", + exx_manifest, + ): + mapped_exx.append(match.group(1) + "." + match.group(2)) + errors.extend(compare("GlobalC::exx_info", declared_exx, mapped_exx)) + errors.extend(compare_schema_digest("GlobalC::exx_info", exx_schema, exx_manifest)) + + restart_schema, restart_manifest = schemas["GlobalC::restart"] + if not save_member or not load_member: + errors.append("Restart state container declarations could not be identified") + else: + declared_restart = [field for field, _ in restart_schema] + errors.extend( + compare( + "GlobalC::restart", + declared_restart, + manifest_fields(restart_manifest, "MODULE_CONTEXT_RESTART_FIELD", path_field=True), + ) + ) + errors.extend(compare_schema_digest("GlobalC::restart", restart_schema, restart_manifest)) + return errors + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--repo-root", type=Path, required=True) + parser.add_argument("--print-digests", action="store_true") + args = parser.parse_args() + if args.print_digests: + schemas, _, _ = collect_schemas(args.repo_root.resolve()) + for label, (schema, _) in sorted(schemas.items()): + print("{} {}".format(label, schema_digest(schema))) + return 0 + errors = check(args.repo_root.resolve()) + if errors: + for error in errors: + print("context mapping error: " + error, file=sys.stderr) + return 1 + print("Context mapping manifests cover all legacy fields exactly once.") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/source/source_context/test/simulation_context_binding_test.cpp b/source/source_context/test/simulation_context_binding_test.cpp new file mode 100644 index 00000000000..13b926cef53 --- /dev/null +++ b/source/source_context/test/simulation_context_binding_test.cpp @@ -0,0 +1,85 @@ +#include "source_context/orchestration_context.h" +#include "source_context/simulation_context_binding.h" +#include "source_hamilt/module_xc/exx_info.h" +#include "source_io/module_restart/restart.h" + +#include + +#include +#include + +#ifdef _OPENMP +#include +#endif + +namespace GlobalC +{ +Exx_Info exx_info; +Restart restart; +} // namespace GlobalC + +TEST(SimulationContextBinding, RejectsReadOutsideBindingLifetime) +{ + EXPECT_THROW(ModuleContext::current_simulation_context(), std::logic_error); + + ModuleContext::SimulationContext context; + { + ModuleContext::ScopedSimulationContextBinding binding(context); + EXPECT_EQ(&context, &ModuleContext::current_simulation_context()); + } + + EXPECT_THROW(ModuleContext::current_simulation_context(), std::logic_error); +} + +TEST(SimulationContextBinding, RejectsDuplicateBinding) +{ + ModuleContext::SimulationContext first; + ModuleContext::SimulationContext second; + ModuleContext::ScopedSimulationContextBinding binding(first); + + EXPECT_THROW(ModuleContext::ScopedSimulationContextBinding duplicate(second), std::logic_error); + EXPECT_EQ(&first, &ModuleContext::current_simulation_context()); +} + +TEST(SimulationContextBinding, SupportsSequentialRuns) +{ + ModuleContext::SimulationContext first; + ModuleContext::SimulationContext second; + { + ModuleContext::ScopedSimulationContextBinding binding(first); + EXPECT_EQ(&first, &ModuleContext::current_simulation_context()); + } + { + ModuleContext::ScopedSimulationContextBinding binding(second); + EXPECT_EQ(&second, &ModuleContext::current_simulation_context()); + } +} + +TEST(SimulationContextBinding, SharesTheDriverOwnedObjectWithWorkerThreads) +{ + ModuleContext::SimulationContext context; + ModuleContext::ScopedSimulationContextBinding binding(context); + + std::future worker + = std::async(std::launch::async, []() { return &ModuleContext::current_simulation_context(); }); + EXPECT_EQ(&context, worker.get()); +} + +#ifdef _OPENMP +TEST(SimulationContextBinding, IsVisibleToOpenMpWorkers) +{ + ModuleContext::SimulationContext context; + ModuleContext::ScopedSimulationContextBinding binding(context); + int matching_workers = 0; + int worker_count = 0; +#pragma omp parallel reduction(+ : matching_workers, worker_count) + { + ++worker_count; + if (&ModuleContext::current_simulation_context() == &context) + { + ++matching_workers; + } + } + EXPECT_EQ(worker_count, matching_workers); +} +#endif diff --git a/source/source_context/test/simulation_context_builder_test.cpp b/source/source_context/test/simulation_context_builder_test.cpp new file mode 100644 index 00000000000..b10eecd54d1 --- /dev/null +++ b/source/source_context/test/simulation_context_builder_test.cpp @@ -0,0 +1,250 @@ +#include "source_context/simulation_context_builder.h" + +#include "source_base/global_variable.h" +#include "source_hamilt/module_xc/exx_info.h" +#include "source_io/module_parameter/input_parameter.h" +#include "source_io/module_parameter/system_parameter.h" +#include "source_io/module_restart/restart.h" + +#include + +#include +#include +#include + +namespace GlobalC +{ +Exx_Info exx_info; +Restart restart; +} // namespace GlobalC + +namespace +{ + +class SimulationContextBuilderTest : public ::testing::Test +{ + protected: + void SetUp() override + { + GlobalV::NPROC = 1; + GlobalV::KPAR = 1; + GlobalV::MY_RANK = 0; + GlobalV::MY_POOL = 0; + GlobalV::MY_BNDGROUP = 0; + GlobalV::NPROC_IN_POOL = 1; + GlobalV::NPROC_IN_BNDGROUP = 1; + GlobalV::RANK_IN_POOL = 0; + GlobalV::RANK_IN_BPGROUP = 0; + GlobalV::DRANK = 0; + GlobalV::DSIZE = 1; + GlobalV::DCOLOR = 0; + GlobalV::GRANK = 0; + GlobalV::GSIZE = 1; + GlobalC::exx_info = Exx_Info(); + GlobalC::restart = Restart(); + } + + ModuleContext::SimulationContext freeze(Input_para& input, System_para& system) + { + ModuleContext::SimulationContextBuilder builder(input, system); + builder.capture_runtime(system); + return builder.finalize(system); + } +}; + +TEST_F(SimulationContextBuilderTest, ProjectsEveryInputFieldExactlyOnce) +{ + Input_para input; + input.calculation = "nscf"; + input.suffix = "projection-test"; + input.pseudo_dir = "pseudo"; + input.device = "cpu"; + input.nelec = 12.5; + input.mixing_beta = 0.37; + input.nspin = 4; + input.ks_solver = "cg"; + input.lcao_dr = 0.02; + input.relax_nmax = 9; + input.ref_cell_factor = 1.5; + input.of_method = "cg1"; + input.nbands_sto = 32; + input.deepks_model = "model.pt"; + input.td_dt = 0.1; + input.lr_nstates = 7; + input.out_level = "m"; + input.out_mat_hs2 = {1, 12, 2}; + input.out_chg = {1, 6}; + input.restart_save = true; + input.dos_sigma = 0.12; + input.efield_amp = 0.3; + input.vdw_method = "d4"; + input.exx_hybrid_step = 13; + input.dft_plus_u = 1; + input.sc_thr = 2.0e-6; + input.qo_switch = true; + input.pexsi_npole = 28; + input.nurse = 3; + input.rdmft_power_alpha = 0.7; + input.xc_exch_ext = {101.0, 0.9}; + input.of_cd = true; + input.dsp_count = 6; + + System_para system; + const ModuleContext::SimulationContext context = freeze(input, system); + + std::set mapped_names; +#define MODULE_CONTEXT_INPUT_FIELD(domain, field) \ + do \ + { \ + EXPECT_TRUE(mapped_names.insert(#field).second) << "duplicate INPUT field " << #field; \ + } while (false); +#include "source_context/input_field_mapping.inc" +#undef MODULE_CONTEXT_INPUT_FIELD + EXPECT_EQ(504u, mapped_names.size()); + + EXPECT_EQ("projection-test", context.cell.value("suffix")); + EXPECT_EQ("pseudo", context.file_input.value("pseudo_dir")); + EXPECT_DOUBLE_EQ(0.37, context.scf_mixing.value("mixing_beta")); + EXPECT_EQ("model.pt", context.deepks.value("deepks_model")); + EXPECT_EQ("d4", context.vdw.value("vdw_method")); + EXPECT_EQ(6, context.uncommon_hardware.value("dsp_count")); + EXPECT_EQ(0u, context.run.mapped_values.count("calculation")); + EXPECT_EQ(0u, context.spin.mapped_values.count("nspin")); + EXPECT_EQ(0u, context.matrix_output.mapped_values.count("out_mat_hs2")); + + EXPECT_EQ("nscf", context.run.calculation); + EXPECT_EQ(4, context.spin.nspin); + EXPECT_TRUE(context.matrix_output.hs_r.enabled); + EXPECT_EQ(12, context.matrix_output.hs_r.precision); + EXPECT_EQ(2, context.matrix_output.hs_r.mode); +} + +TEST_F(SimulationContextBuilderTest, RequiresRuntimeCaptureBeforeFreeze) +{ + Input_para input; + System_para system; + ModuleContext::SimulationContextBuilder builder(input, system); + + EXPECT_THROW(builder.finalize(system), std::logic_error); + EXPECT_FALSE(builder.finalized()); + + system.nlocal = 17; + system.npol = 2; + builder.capture_runtime(system); + EXPECT_TRUE(builder.runtime_captured()); + const ModuleContext::SimulationContext context = builder.finalize(system); + EXPECT_EQ(17, context.basis.nlocal); + EXPECT_EQ(2, context.basis.npol); + EXPECT_TRUE(builder.finalized()); + EXPECT_THROW(builder.finalize(system), std::logic_error); + EXPECT_THROW(builder.capture_runtime(system), std::logic_error); +} + +TEST_F(SimulationContextBuilderTest, SourceMappingManifestsAreCompleteAndUnique) +{ + System_para system; + std::set system_fields; +#define MODULE_CONTEXT_SYSTEM_FIELD(domain, field) \ + static_cast(system.field); \ + EXPECT_TRUE(system_fields.insert(#field).second) << "duplicate System_para field " << #field; +#include "source_context/system_field_mapping.inc" +#undef MODULE_CONTEXT_SYSTEM_FIELD + EXPECT_EQ(45u, system_fields.size()); + + std::set globalv_fields; +#define MODULE_CONTEXT_GLOBALV_FIELD(domain, field) \ + static_cast(GlobalV::field); \ + EXPECT_TRUE(globalv_fields.insert(#field).second) << "duplicate GlobalV field " << #field; +#include "source_context/globalv_field_mapping.inc" +#undef MODULE_CONTEXT_GLOBALV_FIELD + EXPECT_EQ(18u, globalv_fields.size()); + + std::set restart_fields; +#define MODULE_CONTEXT_RESTART_FIELD(field) \ + static_cast(GlobalC::restart.field); \ + EXPECT_TRUE(restart_fields.insert(#field).second) << "duplicate restart field " << #field; +#include "source_context/restart_field_mapping.inc" +#undef MODULE_CONTEXT_RESTART_FIELD + EXPECT_EQ(8u, restart_fields.size()); +} + +TEST_F(SimulationContextBuilderTest, RejectsInconsistentLegacyWorldTopology) +{ + Input_para input; + System_para system; + system.nproc = 1; + GlobalV::NPROC = 2; + + ModuleContext::SimulationContextBuilder builder(input, system); + builder.capture_runtime(system); + EXPECT_THROW(builder.finalize(system), std::logic_error); +} + +TEST_F(SimulationContextBuilderTest, UsesLegacyGlobalVAsTheRuntimePoolTopology) +{ + Input_para input; + System_para system; + system.npool = 2; + system.nproc_in_pool = 1; + GlobalV::KPAR = 3; + GlobalV::MY_POOL = 1; + GlobalV::NPROC_IN_POOL = 4; + + ModuleContext::SimulationContextBuilder builder(input, system); + builder.capture_runtime(system); + const ModuleContext::SimulationContext context = builder.finalize(system); + + EXPECT_EQ(3, context.parallel.kpoint_pool_count); + EXPECT_EQ(1, context.parallel.pool_index); + EXPECT_EQ(4, context.parallel.processes_in_pool); +} + +TEST_F(SimulationContextBuilderTest, AdaptersUseTheSingleLegacyStateInstances) +{ + GlobalC::exx_info.info_global.cal_exx = true; + GlobalC::exx_info.info_global.hybrid_alpha = 0.25; + GlobalC::exx_info.info_global.hse_omega = 0.11; + GlobalC::exx_info.info_ri.real_number = true; + GlobalC::restart.info_save.save_charge = true; + GlobalC::restart.info_load.load_H = true; + GlobalC::restart.folder = "restart-a/"; + + Input_para input; + System_para system; + const ModuleContext::SimulationContext context = freeze(input, system); + + const ModuleContext::ExactExchangeState exx = context.exact_exchange_state->snapshot(); + std::set exx_names; +#define MODULE_CONTEXT_EXX_STATE_FIELD(section, field) \ + do \ + { \ + const std::string name = #section "." #field; \ + EXPECT_TRUE(exx_names.insert(name).second) << "duplicate EXX state field " << name; \ + } while (false); +#include "source_context/exx_state_field_mapping.inc" +#undef MODULE_CONTEXT_EXX_STATE_FIELD + EXPECT_EQ(41u, exx_names.size()); + EXPECT_TRUE(exx.enabled); + EXPECT_TRUE(exx.real_number); + EXPECT_DOUBLE_EQ(0.25, exx.hybrid_alpha); + EXPECT_DOUBLE_EQ(0.11, exx.hse_omega); + EXPECT_EQ(0u, exx.mapped_values.count("info_global.hybrid_alpha")); + EXPECT_EQ(1u, exx.mapped_values.count("info_ri.C_threshold")); + + ModuleContext::RestartState restart = context.restart->snapshot(); + EXPECT_TRUE(restart.save_charge); + EXPECT_TRUE(restart.load_hamiltonian); + EXPECT_EQ("restart-a/", restart.folder); + + restart.save_charge = false; + restart.load_hamiltonian_finished = true; + restart.restart_exact_exchange = true; + restart.folder = "restart-b/"; + context.restart->update(restart); + EXPECT_FALSE(GlobalC::restart.info_save.save_charge); + EXPECT_TRUE(GlobalC::restart.info_load.load_H_finish); + EXPECT_TRUE(GlobalC::restart.info_load.restart_exx); + EXPECT_EQ("restart-b/", GlobalC::restart.folder); +} + +} // namespace diff --git a/source/source_esolver/CMakeLists.txt b/source/source_esolver/CMakeLists.txt index f0f91cc3953..1bb87b933da 100644 --- a/source/source_esolver/CMakeLists.txt +++ b/source/source_esolver/CMakeLists.txt @@ -34,6 +34,8 @@ add_library( ../source_pw/module_pwdft/exx_helper.h ) +target_compile_definitions(esolver PRIVATE ABACUS_CAN_READ_SIMULATION_CONTEXT) + if(ENABLE_COVERAGE) add_coverage(esolver) endif() @@ -43,4 +45,3 @@ if(BUILD_TESTING) add_subdirectory(test) endif() endif() - diff --git a/source/source_esolver/esolver.cpp b/source/source_esolver/esolver.cpp index fc3aab0aa46..be9f8b78fff 100644 --- a/source/source_esolver/esolver.cpp +++ b/source/source_esolver/esolver.cpp @@ -285,46 +285,20 @@ ESolver* init_esolver(const Input_para& inp, UnitCell& ucell) } else if (esolver_type == "ksdft_lr_lcao") { - // initialize the 1st ESolver_KS - ModuleESolver::ESolver* p_esolver = nullptr; + // Driver runs this ground-state solver after the new read-only Context + // has been finalized and bound, then transitions it to LR below. if (PARAM.globalv.gamma_only_local) { - p_esolver = new ESolver_KS_LCAO(); + return new ESolver_KS_LCAO(); } else if (PARAM.inp.nspin < 4) { - p_esolver = new ESolver_KS_LCAO, double>(); + return new ESolver_KS_LCAO, double>(); } else { - p_esolver = new ESolver_KS_LCAO, std::complex>(); + return new ESolver_KS_LCAO, std::complex>(); } - p_esolver->before_all_runners(ucell, inp); - p_esolver->runner(ucell, 0); // scf-only - - // force and stress is not needed currently, - // they will be supported after the analytical gradient - // of LR-TDDFT is implemented. - std::cout << " PREPARING FOR EXCITED STATES." << std::endl; - // initialize the 2nd ESolver_LR at the temporary pointer - ModuleESolver::ESolver* p_esolver_lr = nullptr; - if (PARAM.globalv.gamma_only_local) - { - p_esolver_lr = new LR::ESolver_LR( - std::move(*dynamic_cast*>(p_esolver)), - inp, - ucell); - } - else - { - p_esolver_lr = new LR::ESolver_LR, double>( - std::move(*dynamic_cast, double>*>(p_esolver)), - inp, - ucell); - } - // clean the 1st ESolver_KS and swap the pointer - delete p_esolver; - return p_esolver_lr; } #endif else if (esolver_type == "ofdft") @@ -351,5 +325,31 @@ ESolver* init_esolver(const Input_para& inp, UnitCell& ucell) + " line " + std::to_string(__LINE__)); } +ESolver* transition_ksdft_to_lr(ESolver* p_esolver, const Input_para& inp, UnitCell& ucell) +{ +#ifdef __LCAO + std::cout << " PREPARING FOR EXCITED STATES." << std::endl; + ESolver* p_esolver_lr = nullptr; + if (PARAM.globalv.gamma_only_local) + { + p_esolver_lr = new LR::ESolver_LR( + std::move(*dynamic_cast*>(p_esolver)), inp, ucell); + } + else + { + p_esolver_lr = new LR::ESolver_LR, double>( + std::move(*dynamic_cast, double>*>(p_esolver)), inp, ucell); + } + delete p_esolver; + return p_esolver_lr; +#else + static_cast(p_esolver); + static_cast(inp); + static_cast(ucell); + ModuleBase::WARNING_QUIT("ESolver", "LR-LCAO requires an LCAO build"); + return nullptr; +#endif +} + } // namespace ModuleESolver diff --git a/source/source_esolver/esolver.h b/source/source_esolver/esolver.h index d1f2b1ae782..1fee24ea86c 100644 --- a/source/source_esolver/esolver.h +++ b/source/source_esolver/esolver.h @@ -70,6 +70,15 @@ std::string determine_type(); */ ESolver* init_esolver(const Input_para& inp, UnitCell& ucell); +/** + * @brief Replace the completed ground-state KS-LCAO solver with the LR solver. + * + * The caller must have run the KS solver once. Keeping this transition outside + * init_esolver() lets Driver bind the finalized SimulationContext before that + * first runner executes. + */ +ESolver* transition_ksdft_to_lr(ESolver* p_esolver, const Input_para& inp, UnitCell& ucell); + } // namespace ModuleESolver diff --git a/source/source_esolver/esolver_gets.cpp b/source/source_esolver/esolver_gets.cpp index 44bbd830814..7e684b98c30 100644 --- a/source/source_esolver/esolver_gets.cpp +++ b/source/source_esolver/esolver_gets.cpp @@ -11,6 +11,7 @@ #include "source_io/module_hs/cal_r_overlap_R.h" #include "source_io/module_output/print_info.h" #include "source_io/module_hs/write_HS_R.h" +#include "source_context/orchestration_context.h" namespace ModuleESolver { @@ -156,16 +157,35 @@ void ESolver_GetS::runner(UnitCell& ucell, const int istep) } } - const std::string fn = PARAM.globalv.global_out_dir + "sr_nao.csr"; + const ModuleContext::SimulationContext& context = ModuleContext::current_simulation_context(); + const std::string fn = context.files.output_directory + "sr_nao.csr"; auto* hamilt_ptr = static_cast>*>(this->p_hamilt); - ModuleIO::output_SR(pv, gd, hamilt_ptr, fn); + ModuleIO::output_SR(pv, + gd, + hamilt_ptr, + fn, + false, + 1.0e-10, + 16, + context.files, + context.parallel, + context.logs, + context.spin); if (PARAM.inp.out_mat_r[0]) { cal_r_overlap_R r_matrix; - r_matrix.init(ucell, pv, orb_); - r_matrix.out_rR(ucell, gd, istep, PARAM.inp.out_mat_r[1]); + r_matrix.init(ucell, pv, orb_, context.run, context.basis); + r_matrix.out_rR(ucell, + gd, + istep, + PARAM.inp.out_mat_r[1], + context.run, + context.files, + context.parallel, + context.basis, + context.matrix_output); } if (PARAM.inp.out_mat_ds[0]) @@ -182,7 +202,14 @@ void ESolver_GetS::runner(UnitCell& ucell, const int istep) kv, false, 1e-10, - PARAM.inp.out_mat_ds[1]); + PARAM.inp.out_mat_ds[1], + context.run, + context.files, + context.parallel, + context.logs, + context.basis, + context.spin, + context.matrix_output); } ModuleBase::timer::end("ESolver_GetS", "runner"); diff --git a/source/source_esolver/esolver_ks_lcao_tddft.cpp b/source/source_esolver/esolver_ks_lcao_tddft.cpp index c3872481db9..6a07be3e8b3 100644 --- a/source/source_esolver/esolver_ks_lcao_tddft.cpp +++ b/source/source_esolver/esolver_ks_lcao_tddft.cpp @@ -15,6 +15,7 @@ #include "source_hsolver/hsolver_lcao.h" #include "source_lcao/module_rt/evolve_elec.h" #include "source_lcao/rho_tau_lcao.h" +#include "source_context/orchestration_context.h" #ifdef __EXX #include "source_lcao/module_ri/Exx_LRI_interface.h" #endif @@ -102,7 +103,9 @@ void ESolver_KS_LCAO_TDDFT::runner(UnitCell& ucell, const int istep) //---------------------------------------------------------------- // 1) before_scf (electronic iteration loops) //---------------------------------------------------------------- + const ModuleContext::SimulationContext& context = ModuleContext::current_simulation_context(); this->before_scf(ucell, istep); // From ESolver_KS_LCAO + td_p->initialize_r_calculator(ucell, this->pv, this->orb_, context.run, context.basis); td_p->initialize_phase_hybrid(ucell, dynamic_cast, TR>*>(this->p_hamilt)->getHR()); td_p->calculate_grad_overlap(this->pv, ucell, this->gd, this->orb_.cutoffs(), this->two_center_bundle_.overlap_orb.get()); // Initialize the moving spatial gauge @@ -132,7 +135,9 @@ void ESolver_KS_LCAO_TDDFT::runner(UnitCell& ucell, const int istep) &(this->gd), &this->pv, this->orb_, - this->two_center_bundle_.overlap_orb.get()); + this->two_center_bundle_.overlap_orb.get(), + context.run, + context.basis); // calculate velocity operator velocity_mat->calculate_grad_term(); velocity_mat->calculate_vcomm_r(); diff --git a/source/source_esolver/esolver_ks_lcaopw.cpp b/source/source_esolver/esolver_ks_lcaopw.cpp index e5cfb0dc24f..eb53aa122ab 100644 --- a/source/source_esolver/esolver_ks_lcaopw.cpp +++ b/source/source_esolver/esolver_ks_lcaopw.cpp @@ -27,6 +27,7 @@ #include #ifdef __LCAO #include "source_io/module_hs/write_vxc_lip.hpp" +#include "source_context/orchestration_context.h" #endif namespace ModuleESolver @@ -227,16 +228,12 @@ namespace ModuleESolver #ifdef __LCAO if (PARAM.inp.out_mat_xc) { -#ifdef __EXX - bool cal_exx = GlobalC::exx_info.info_global.cal_exx; - double hybrid_alpha = GlobalC::exx_info.info_global.hybrid_alpha; -#else - bool cal_exx = false; - double hybrid_alpha = 0.0; -#endif + const ModuleContext::SimulationContext& context = ModuleContext::current_simulation_context(); + const ModuleContext::ExactExchangeState exx_state + = context.exact_exchange_state ? context.exact_exchange_state->snapshot() + : ModuleContext::ExactExchangeState(); ModuleIO::write_Vxc(PARAM.inp.nspin, PARAM.globalv.nlocal, - GlobalV::DRANK, *this->stp.template get_psi_t(), ucell, this->sf, @@ -248,8 +245,13 @@ namespace ModuleESolver this->chr, this->kv, this->pelec->wg, - cal_exx, - hybrid_alpha + context.files, + context.parallel, + context.basis, + context.solver, + context.matrix_output, + exx_state.enabled, + exx_state.hybrid_alpha #ifdef __EXX , *this->exx_lip diff --git a/source/source_io/CMakeLists.txt b/source/source_io/CMakeLists.txt index 4de0b6e88f3..012bf8f05c4 100644 --- a/source/source_io/CMakeLists.txt +++ b/source/source_io/CMakeLists.txt @@ -3,9 +3,6 @@ list(APPEND objects module_parameter/input_conv.cpp - module_ctrl/ctrl_output_fp.cpp - module_ctrl/ctrl_output_pw.cpp - module_ctrl/ctrl_output_td.cpp module_bessel/bessel_basis.cpp module_output/cal_test.cpp module_dos/cal_dos.cpp @@ -47,6 +44,12 @@ list(APPEND objects module_output/ucell_io.cpp ) +list(APPEND orchestration_objects + module_ctrl/ctrl_output_fp.cpp + module_ctrl/ctrl_output_pw.cpp + module_ctrl/ctrl_output_td.cpp +) + list(APPEND objects_advanced module_unk/unk_overlap_pw.cpp module_unk/berryphase.cpp @@ -93,6 +96,8 @@ if(ENABLE_LCAO) module_hs/rr_sparse_writer.cpp module_hs/cal_r_overlap_R.cpp module_hs/output_mat_sparse.cpp + ) + list(APPEND orchestration_objects module_ctrl/ctrl_scf_lcao.cpp module_ctrl/ctrl_runner_lcao.cpp module_ctrl/ctrl_iter_lcao.cpp @@ -138,6 +143,14 @@ add_library( ${objects_advanced} ) +add_library( + io_orchestration + OBJECT + ${orchestration_objects} +) + +target_compile_definitions(io_orchestration PRIVATE ABACUS_CAN_READ_SIMULATION_CONTEXT) + if(ENABLE_COVERAGE) add_coverage(io_basic) endif() diff --git a/source/source_io/module_ctrl/ctrl_iter_lcao.cpp b/source/source_io/module_ctrl/ctrl_iter_lcao.cpp index 7fd133d414b..347d7821d25 100644 --- a/source/source_io/module_ctrl/ctrl_iter_lcao.cpp +++ b/source/source_io/module_ctrl/ctrl_iter_lcao.cpp @@ -6,6 +6,7 @@ #include "source_lcao/module_deepks/LCAO_deepks_interface.h" #endif #include "source_io/module_restart/restart.h" +#include "source_context/orchestration_context.h" #ifdef __EXX #include "source_lcao/module_ri/Exx_LRI_interface.h" #endif @@ -35,6 +36,7 @@ void ctrl_iter_lcao(UnitCell& ucell, // unit cell * { ModuleBase::TITLE("ModuleIO", "ctrl_iter_lcao"); ModuleBase::timer::start("ModuleIO", "ctrl_iter_lcao"); + const ModuleContext::SimulationContext& context = ModuleContext::current_simulation_context(); // save charge density // Peize Lin add 2020.04.04 @@ -56,9 +58,9 @@ void ctrl_iter_lcao(UnitCell& ucell, // unit cell * { real_number ? exx_nao.exd->exx_iter_finish(kv, ucell, *p_hamilt, *pelec, &dm, - *p_chgmix, scf_ene_thr, iter, istep, conv_esolver) : + *p_chgmix, scf_ene_thr, iter, istep, conv_esolver, context.parallel) : exx_nao.exc->exx_iter_finish(kv, ucell, *p_hamilt, *pelec, &dm, - *p_chgmix, scf_ene_thr, iter, istep, conv_esolver); + *p_chgmix, scf_ene_thr, iter, istep, conv_esolver, context.parallel); } } #endif diff --git a/source/source_io/module_ctrl/ctrl_runner_lcao.cpp b/source/source_io/module_ctrl/ctrl_runner_lcao.cpp index 0fa74574c64..34107405612 100644 --- a/source/source_io/module_ctrl/ctrl_runner_lcao.cpp +++ b/source/source_io/module_ctrl/ctrl_runner_lcao.cpp @@ -2,7 +2,7 @@ #include "source_estate/elecstate_lcao.h" // use elecstate::ElecState #include "source_lcao/hamilt_lcao.h" // use hamilt::HamiltLCAO -#include "source_hamilt/module_xc/exx_info.h" +#include "source_context/orchestration_context.h" #include "../module_energy/write_proj_band_lcao.h" // projcted band structure #include "../module_dos/cal_ldos.h" // cal LDOS @@ -39,6 +39,7 @@ void ctrl_runner_lcao(UnitCell& ucell, // unitcell { ModuleBase::TITLE("ModuleIO", "ctrl_runner_lcao"); ModuleBase::timer::start("ModuleIO", "ctrl_runner_lcao"); + const ModuleContext::SimulationContext& context = ModuleContext::current_simulation_context(); // 1) write projected band structure if (inp.out_proj_band) @@ -56,10 +57,11 @@ void ctrl_runner_lcao(UnitCell& ucell, // unitcell // 3) print out exchange-correlation potential if (inp.out_mat_xc) { - bool cal_exx = GlobalC::exx_info.info_global.cal_exx; + const ModuleContext::ExactExchangeState exx_state + = context.exact_exchange_state ? context.exact_exchange_state->snapshot() + : ModuleContext::ExactExchangeState(); ModuleIO::write_Vxc(inp.nspin, - PARAM.globalv.nlocal, - GlobalV::DRANK, + context.basis.nlocal, &pv, *psi, ucell, @@ -73,7 +75,13 @@ void ctrl_runner_lcao(UnitCell& ucell, // unitcell orb.cutoffs(), pelec->wg, gd, - cal_exx + context.files, + context.parallel, + context.basis, + context.solver, + context.matrix_output, + context.dftu, + exx_state.enabled #ifdef __EXX , exx_nao.exd ? &exx_nao.exd->get_Hexxs() : nullptr, @@ -84,9 +92,9 @@ void ctrl_runner_lcao(UnitCell& ucell, // unitcell if (inp.out_mat_xc2[0]) { - bool cal_exx = GlobalC::exx_info.info_global.cal_exx; - double hybrid_alpha = GlobalC::exx_info.info_global.hybrid_alpha; - bool real_number = GlobalC::exx_info.info_ri.real_number; + const ModuleContext::ExactExchangeState exx_state + = context.exact_exchange_state ? context.exact_exchange_state->snapshot() + : ModuleContext::ExactExchangeState(); ModuleIO::write_Vxc_R(inp.nspin, &pv, ucell, @@ -99,9 +107,11 @@ void ctrl_runner_lcao(UnitCell& ucell, // unitcell kv, orb.cutoffs(), gd, - cal_exx, - hybrid_alpha, - real_number + context.files, + context.parallel, + exx_state.enabled, + exx_state.hybrid_alpha, + exx_state.real_number #ifdef __EXX , exx_nao.exd ? &exx_nao.exd->get_Hexxs() : nullptr, @@ -116,7 +126,6 @@ void ctrl_runner_lcao(UnitCell& ucell, // unitcell { ModuleIO::write_eband_terms(inp.nspin, PARAM.globalv.nlocal, - GlobalV::DRANK, &pv, *psi, ucell, @@ -130,7 +139,16 @@ void ctrl_runner_lcao(UnitCell& ucell, // unitcell pelec->wg, gd, orb.cutoffs(), - two_center_bundle + two_center_bundle, + context.files, + context.parallel, + context.basis, + context.solver, + context.matrix_output, + context.dftu, + context.exact_exchange_state + ? context.exact_exchange_state->snapshot() + : ModuleContext::ExactExchangeState() #ifdef __EXX , exx_nao.exd ? &exx_nao.exd->get_Hexxs() : nullptr, diff --git a/source/source_io/module_ctrl/ctrl_scf_lcao.cpp b/source/source_io/module_ctrl/ctrl_scf_lcao.cpp index aafeb18e6aa..da4442a74cc 100644 --- a/source/source_io/module_ctrl/ctrl_scf_lcao.cpp +++ b/source/source_io/module_ctrl/ctrl_scf_lcao.cpp @@ -38,16 +38,21 @@ #include "source_lcao/module_rdmft/rdmft.h" // use RDMFT codes #include "source_lcao/rho_tau_lcao.h" // mohan add 2025-10-24 #include "source_lcao/module_operator_lcao/overlap.h" // use hamilt::Overlap for NAMD +#include "source_context/orchestration_context.h" #ifdef __EXX template -void setup_exx_dh_params(ModuleIO::WriteDHParams& dh_params, Exx_NAO& exx_nao) +void setup_exx_dh_params(ModuleIO::WriteDHParams& dh_params, + Exx_NAO& exx_nao, + const ModuleContext::ExactExchangeState& exact_exchange_state) {} template <> -void setup_exx_dh_params(ModuleIO::WriteDHParams& dh_params, Exx_NAO& exx_nao) +void setup_exx_dh_params(ModuleIO::WriteDHParams& dh_params, + Exx_NAO& exx_nao, + const ModuleContext::ExactExchangeState& exact_exchange_state) { - if (GlobalC::exx_info.info_global.cal_exx) + if (exact_exchange_state.enabled) { if (exx_nao.exd) { dh_params.exd = exx_nao.exd.get(); } if (exx_nao.exc) { dh_params.exc = exx_nao.exc.get(); } @@ -55,7 +60,9 @@ void setup_exx_dh_params(ModuleIO::WriteDHParams& dh_params, Exx_NAO -void setup_exx_h_params(ModuleIO::WriteHParams& h_params, Exx_NAO& exx_nao) +void setup_exx_h_params(ModuleIO::WriteHParams& h_params, + Exx_NAO& exx_nao, + const ModuleContext::ExactExchangeState& exact_exchange_state) { // Only the gamma-only (TK==double) specialization below actually writes V^EXX(R). // This generic body is instantiated for the multi-k (TK==std::complex) path, where the @@ -67,13 +74,14 @@ void setup_exx_h_params(ModuleIO::WriteHParams& h_params, Exx_NAO& exx_nao) } template <> -void setup_exx_h_params(ModuleIO::WriteHParams& h_params, Exx_NAO& exx_nao) +void setup_exx_h_params(ModuleIO::WriteHParams& h_params, + Exx_NAO& exx_nao, + const ModuleContext::ExactExchangeState& exact_exchange_state) { - if (GlobalC::exx_info.info_global.cal_exx) + if (exact_exchange_state.enabled) { if (exx_nao.exd) { h_params.exd = exx_nao.exd.get(); } if (exx_nao.exc) { h_params.exc = exx_nao.exc.get(); } - ModuleIO::write_h_exx(h_params); } } #endif @@ -107,6 +115,10 @@ void ModuleIO::ctrl_scf_lcao(UnitCell& ucell, { ModuleBase::TITLE("ModuleIO", "ctrl_scf_lcao"); ModuleBase::timer::start("ModuleIO", "ctrl_scf_lcao"); + const ModuleContext::SimulationContext& context = ModuleContext::current_simulation_context(); + const ModuleContext::ExactExchangeState exact_exchange_state + = context.exact_exchange_state ? context.exact_exchange_state->snapshot() + : ModuleContext::ExactExchangeState(); //***** // if istep_in = -1, istep will not appear in file name @@ -201,18 +213,20 @@ void ModuleIO::ctrl_scf_lcao(UnitCell& ucell, //------------------------------------------------------------------ if (inp.out_mat_hs[0]) { - ModuleIO::write_hsk(global_out_dir, - nspin, + ModuleIO::write_hsk(nspin, kv.get_nks(), kv.get_nkstot(), kv.ik2iktot, kv.isk, p_hamilt, pv, - gamma_only, - out_app_flag, istep, - GlobalV::ofs_running); + context.files, + context.parallel, + context.logs, + context.basis, + context.solver, + context.matrix_output); } //------------------------------------------------------------------ @@ -269,8 +283,17 @@ void ModuleIO::ctrl_scf_lcao(UnitCell& ucell, std::vector*> hr_vec = p_hamilt->getHR_vector(); const hamilt::HContainer* sr = p_hamilt->getSR(); - ModuleIO::write_hsr(hr_vec, sr, &ucell, precision, pv, - out_app_flag, ucell.get_iat2iwt(), ucell.nat, istep); + ModuleIO::write_hsr(hr_vec, + sr, + &ucell, + precision, + pv, + out_app_flag, + ucell.get_iat2iwt(), + ucell.nat, + istep, + context.files, + context.parallel); } //------------------------------------------------------------------ @@ -330,7 +353,14 @@ void ModuleIO::ctrl_scf_lcao(UnitCell& ucell, gd, kv, p_ham_tk, - &dftu); + &dftu, + context.run, + context.files, + context.parallel, + context.logs, + context.basis, + context.spin, + context.matrix_output); //------------------------------------------------------------------ //! 7c) Output atomic dH components (dT/dτ, dV^NL/dτ, dV^L/dτ, dV^H/dτ, dV^XC/dτ), only for nspin =1, 2 now @@ -347,6 +377,8 @@ void ModuleIO::ctrl_scf_lcao(UnitCell& ucell, dh_params.v_eff = &pelec->pot->get_eff_v(); dh_params.pot = pelec->pot; dh_params.chg = pelec->charge; + dh_params.parallel = &context.parallel; + dh_params.solver = &context.solver; // pelec->pot->get_eff_v() is the SUM V^L + V^H + V^XC; feeding it to cal_dH would // give the wrong potential for the separated V^L / V^H / V^XC outputs. Build one // dedicated Potential per term with exactly one component registered (see write_vxc.hpp). @@ -399,7 +431,7 @@ void ModuleIO::ctrl_scf_lcao(UnitCell& ucell, #ifdef __EXX // dV^EXX/dR output is wired for the gamma (TK==double) exx interfaces. exd/exc are // mutually exclusive (real vs complex Hexx); write_dH_exx picks by info_ri.real_number. - setup_exx_dh_params(dh_params, exx_nao); + setup_exx_dh_params(dh_params, exx_nao, exact_exchange_state); #endif ModuleIO::write_dH_components(dh_params); delete pot_vl; @@ -430,29 +462,37 @@ void ModuleIO::ctrl_scf_lcao(UnitCell& ucell, h_params.nat = ucell.nat; if (inp.out_mat_h_t[0]) { - ModuleIO::write_h_t(h_params); + ModuleIO::write_h_t(h_params, context.run, context.files, context.parallel, context.basis, context.solver, context.matrix_output); } if (inp.out_mat_h_vnl[0]) { - ModuleIO::write_h_vnl(h_params); + ModuleIO::write_h_vnl(h_params, context.run, context.files, context.parallel, context.basis, context.solver, context.matrix_output); } if (inp.out_mat_h_vl[0]) { - ModuleIO::write_h_vl(h_params); + ModuleIO::write_h_vl(h_params, context.run, context.files, context.parallel, context.basis, context.solver, context.matrix_output); } if (inp.out_mat_h_vh[0]) { - ModuleIO::write_h_vh(h_params); + ModuleIO::write_h_vh(h_params, context.run, context.files, context.parallel, context.basis, context.solver, context.matrix_output); } if (inp.out_mat_h_vxc[0]) { - ModuleIO::write_h_vxc(h_params); + ModuleIO::write_h_vxc(h_params, context.run, context.files, context.parallel, context.basis, context.spin, context.solver, context.matrix_output); } #ifdef __EXX - if (inp.out_mat_h_exx[0] && GlobalC::exx_info.info_global.cal_exx) + if (inp.out_mat_h_exx[0] && exact_exchange_state.enabled) { // V^EXX(R) output is wired for the gamma (TK==double) exx interfaces. - setup_exx_h_params(h_params, exx_nao); + setup_exx_h_params(h_params, exx_nao, exact_exchange_state); + ModuleIO::write_h_exx(h_params, + context.run, + context.files, + context.parallel, + context.basis, + context.solver, + context.matrix_output, + exact_exchange_state); } #endif } @@ -500,7 +540,8 @@ void ModuleIO::ctrl_scf_lcao(UnitCell& ucell, inp.out_app_flag, t_fn, pv, - GlobalV::DRANK); + context.parallel.diagonalization_rank, + context.solver); } delete ekinetic; @@ -519,7 +560,9 @@ void ModuleIO::ctrl_scf_lcao(UnitCell& ucell, inp.test_atom_input, PARAM.globalv.search_pbc, &GlobalV::ofs_running, - GlobalV::MY_RANK); + GlobalV::MY_RANK, + context.run, + context.basis); mylcalculator.calculate(inp.suffix, global_out_dir, ucell, inp.out_mat_l[1], GlobalV::MY_RANK); } @@ -611,21 +654,18 @@ void ModuleIO::ctrl_scf_lcao(UnitCell& ucell, //! 15) Output Hexx matrix in LCAO basis // (see `out_chg` in docs/advanced/input_files/input-main.md) //------------------------------------------------------------------ - bool cal_exx = GlobalC::exx_info.info_global.cal_exx; - bool real_number = GlobalC::exx_info.info_ri.real_number; - if (inp.out_chg[0]) { - if (cal_exx && inp.calculation != "nscf") // Peize Lin add if 2022.11.14 + if (exact_exchange_state.enabled && inp.calculation != "nscf") // Peize Lin add if 2022.11.14 { const std::string file_name_exx = global_out_dir + "HexxR" + std::to_string(GlobalV::MY_RANK); - if (real_number) + if (exact_exchange_state.real_number) { - ModuleIO::write_Hexxs_csr(file_name_exx, ucell, exx_nao.exd->get_Hexxs()); + ModuleIO::write_Hexxs_csr(file_name_exx, ucell, exx_nao.exd->get_Hexxs(), context.parallel); } else { - ModuleIO::write_Hexxs_csr(file_name_exx, ucell, exx_nao.exc->get_Hexxs()); + ModuleIO::write_Hexxs_csr(file_name_exx, ucell, exx_nao.exc->get_Hexxs(), context.parallel); } } } diff --git a/source/source_io/module_dhs/write_dH.cpp b/source/source_io/module_dhs/write_dH.cpp index ea062bf45a4..156cfec51d5 100644 --- a/source/source_io/module_dhs/write_dH.cpp +++ b/source/source_io/module_dhs/write_dH.cpp @@ -121,7 +121,8 @@ void write_dh_perI(WriteDHParams& params, out_app_flag, fk, pv, - GlobalV::DRANK); + params.parallel->diagonalization_rank, + *params.solver); } } } diff --git a/source/source_io/module_dhs/write_dH.h b/source/source_io/module_dhs/write_dH.h index 0fc944f426e..596e424a2f2 100644 --- a/source/source_io/module_dhs/write_dH.h +++ b/source/source_io/module_dhs/write_dH.h @@ -7,6 +7,7 @@ #include "source_estate/module_pot/potential_new.h" #include "source_lcao/LCAO_domain.h" #include "source_lcao/module_hcontainer/hcontainer.h" +#include "source_context/context_types.h" #include #include @@ -53,6 +54,8 @@ struct WriteDHParams // terms: V^H needs the total density (sum over spins), V^XC the spin-resolved densities. std::vector*> dmR; const Charge* chg = nullptr; // ground-state charge for XC Hellmann-Feynman (FDM) + const ModuleContext::ParallelTopology* parallel = nullptr; + const ModuleContext::SolverConfig* solver = nullptr; #ifdef __EXX // The gamma-only (TK==double) exx interfaces used by write_dH_exx. // Deliberately NOT templated on TK, for two reasons: diff --git a/source/source_io/module_energy/write_eband_terms.hpp b/source/source_io/module_energy/write_eband_terms.hpp index 45a51bf125c..daf2686cea4 100644 --- a/source/source_io/module_energy/write_eband_terms.hpp +++ b/source/source_io/module_energy/write_eband_terms.hpp @@ -2,7 +2,6 @@ #define WRITE_EBAND_TERMS_HPP #include "source_io/module_hs/write_vxc.hpp" -#include "source_hamilt/module_xc/exx_info.h" #include "source_lcao/module_operator_lcao/ekinetic.h" #include "source_lcao/module_operator_lcao/nonlocal.h" #include "source_basis/module_nao/two_center_bundle.h" @@ -12,7 +11,6 @@ namespace ModuleIO template void write_eband_terms(const int nspin, const int nbasis, - const int drank, const Parallel_Orbitals* pv, const psi::Psi& psi, const UnitCell& ucell, @@ -26,7 +24,14 @@ void write_eband_terms(const int nspin, const ModuleBase::matrix& wg, Grid_Driver& gd, const std::vector& orb_cutoff, - const TwoCenterBundle& two_center_bundle + const TwoCenterBundle& two_center_bundle, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output, + const ModuleContext::DftUConfig& dftu, + const ModuleContext::ExactExchangeState& exact_exchange_state #ifdef __EXX , std::vector>>>* Hexxd = nullptr, @@ -91,7 +96,7 @@ void write_eband_terms(const int nspin, cVc(kinetic_k_ao.get_hk(), &psi(ik, 0, 0), nbasis, nbands, *pv, p2d), p2d)); } - write_orb_energy(kv, nspin0, nbands, e_orb_kinetic, "kinetic", ""); + write_orb_energy(kv, nspin0, nbands, e_orb_kinetic, files, "kinetic", ""); } // 2. pp: local @@ -123,7 +128,7 @@ void write_eband_terms(const int nspin, e_orb_pp_local.emplace_back(orbital_energy(ik, nbands, cVc(v_pp_local_k_ao.get_hk(), &psi(ik, 0, 0), nbasis, nbands, *pv, p2d), p2d)); } - write_orb_energy(kv, nspin0, nbands, e_orb_pp_local, "vpp_local", ""); + write_orb_energy(kv, nspin0, nbands, e_orb_pp_local, files, "vpp_local", ""); } // 3. pp: nonlocal @@ -143,7 +148,7 @@ void write_eband_terms(const int nspin, e_orb_pp_nonlocal.emplace_back(orbital_energy(ik, nbands, cVc(v_pp_nonlocal_k_ao.get_hk(), &psi(ik, 0, 0), nbasis, nbands, *pv, p2d), p2d)); } - write_orb_energy(kv, nspin0, nbands, e_orb_pp_nonlocal, "vpp_nonlocal", ""); + write_orb_energy(kv, nspin0, nbands, e_orb_pp_nonlocal, files, "vpp_nonlocal", ""); } // 4. hartree @@ -176,16 +181,14 @@ void write_eband_terms(const int nspin, cVc(v_hartree_k_ao.get_hk(), &psi(ik, 0, 0), nbasis, nbands, *pv, p2d), p2d)); } for (auto& op : v_hartree_op) { delete op; } - write_orb_energy(kv, nspin0, nbands, e_orb_hartree, "vhartree", ""); + write_orb_energy(kv, nspin0, nbands, e_orb_hartree, files, "vhartree", ""); } // 5. xc (including exx) if (!PARAM.inp.out_mat_xc) // avoid duplicate output { - bool cal_exx = GlobalC::exx_info.info_global.cal_exx; write_Vxc(nspin, nbasis, - drank, pv, psi, ucell, @@ -199,7 +202,13 @@ void write_eband_terms(const int nspin, orb_cutoff, wg, gd, - cal_exx + files, + parallel, + basis, + solver, + output, + dftu, + exact_exchange_state.enabled #ifdef __EXX , Hexxd, diff --git a/source/source_io/module_hs/cal_pLpR.cpp b/source/source_io/module_hs/cal_pLpR.cpp index 1f62d58cce1..a169f67001e 100644 --- a/source/source_io/module_hs/cal_pLpR.cpp +++ b/source/source_io/module_hs/cal_pLpR.cpp @@ -10,7 +10,6 @@ #include "source_basis/module_nao/two_center_integrator.h" #include "source_cell/module_neighbor/sltk_grid_driver.h" #include "source_cell/module_neighbor/sltk_atom_arrange.h" -#include "source_io/module_parameter/parameter.h" #include "source_io/module_hs/cal_pLpR.h" #include "source_base/formatter.h" #include "source_base/parallel_common.h" @@ -177,7 +176,9 @@ ModuleIO::AngularMomentumCalculator::AngularMomentumCalculator( const int tatom, const bool searchpbc, std::ofstream* ptr_log, - const int rank) + const int rank, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis) { this->ofs_ = ptr_log; @@ -244,10 +245,10 @@ ModuleIO::AngularMomentumCalculator::AngularMomentumCalculator( // we don't really set, but use std::max to mask :) } temp = atom_arrange::set_sr_NL(*ofs_, - PARAM.inp.out_level, + run.output_level, std::max(search_radius, rcut_max), ucell.infoNL->get_rcutmax_Beta(), - PARAM.globalv.gamma_only_local); + basis.gamma_only_local); temp = std::max(temp, std::max(search_radius, rcut_max)); this->neighbor_searcher_ = std::unique_ptr(new Grid_Driver(tdestructor, tgrid)); atom_arrange::search(searchpbc, diff --git a/source/source_io/module_hs/cal_pLpR.h b/source/source_io/module_hs/cal_pLpR.h index c4694596cec..881933dead5 100644 --- a/source/source_io/module_hs/cal_pLpR.h +++ b/source/source_io/module_hs/cal_pLpR.h @@ -94,6 +94,7 @@ #include "source_basis/module_nao/two_center_integrator.h" #include "source_cell/module_neighbor/sltk_grid_driver.h" #include "source_cell/module_neighbor/sltk_atom_arrange.h" +#include "source_context/context_types.h" namespace ModuleIO { @@ -208,8 +209,10 @@ namespace ModuleIO const int tgrid, const int tatom, const bool searchpbc, - std::ofstream* ptr_log = nullptr, - const int rank = 0); + std::ofstream* ptr_log, + const int rank, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis); ~AngularMomentumCalculator() = default; void calculate(const std::string& prefix, diff --git a/source/source_io/module_hs/cal_r_overlap_R.cpp b/source/source_io/module_hs/cal_r_overlap_R.cpp index db0d98de249..03f3faacbb9 100644 --- a/source/source_io/module_hs/cal_r_overlap_R.cpp +++ b/source/source_io/module_hs/cal_r_overlap_R.cpp @@ -7,7 +7,6 @@ #include "source_base/timer.h" #include "source_base/tool_quit.h" #include "source_cell/module_neighbor/sltk_grid_driver.h" -#include "source_io/module_parameter/parameter.h" #include "source_cell/nonlocal_info_base.h" cal_r_overlap_R::cal_r_overlap_R() @@ -40,7 +39,10 @@ void cal_r_overlap_R::initialize_orb_table(const UnitCell& ucell, const LCAO_Orb MGT.init_Gaunt(Lmax); } -void cal_r_overlap_R::construct_orbs_and_orb_r(const UnitCell& ucell, const LCAO_Orbitals& orb) +void cal_r_overlap_R::construct_orbs_and_orb_r(const UnitCell& ucell, + const LCAO_Orbitals& orb, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis) { int orb_r_ntype = 0; int mat_Nr = orb.Phi[0].PhiLN(0, 0).getNr(); @@ -77,7 +79,7 @@ void cal_r_overlap_R::construct_orbs_and_orb_r(const UnitCell& ucell, const LCAO orb_origin.getDruniform(), false, true, - PARAM.inp.cal_force); + run.cal_force); } } } @@ -96,7 +98,7 @@ void cal_r_overlap_R::construct_orbs_and_orb_r(const UnitCell& ucell, const LCAO orbs[orb_r_ntype][0][0].getDruniform(), false, true, - PARAM.inp.cal_force); + run.cal_force); for (int TA = 0; TA < orb.get_ntype(); ++TA) { @@ -180,7 +182,7 @@ void cal_r_overlap_R::construct_orbs_and_orb_r(const UnitCell& ucell, const LCAO } } - int map_size = PARAM.globalv.nlocal; + int map_size = basis.nlocal; int required_orbitals = 0; for (int it = 0; it < ucell.ntype; ++it) { @@ -218,7 +220,10 @@ void cal_r_overlap_R::construct_orbs_and_orb_r(const UnitCell& ucell, const LCAO } } -void cal_r_overlap_R::construct_orbs_and_nonlocal_and_orb_r(const UnitCell& ucell, const LCAO_Orbitals& orb) +void cal_r_overlap_R::construct_orbs_and_nonlocal_and_orb_r(const UnitCell& ucell, + const LCAO_Orbitals& orb, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis) { const NonlocalInfoBase& infoNL_ = *ucell.infoNL; @@ -257,7 +262,7 @@ void cal_r_overlap_R::construct_orbs_and_nonlocal_and_orb_r(const UnitCell& ucel orb_origin.getDruniform(), false, true, - PARAM.inp.cal_force); + run.cal_force); } } } @@ -276,7 +281,7 @@ void cal_r_overlap_R::construct_orbs_and_nonlocal_and_orb_r(const UnitCell& ucel orbs[orb_r_ntype][0][0].getDruniform(), false, true, - PARAM.inp.cal_force); + run.cal_force); orbs_nonlocal.resize(orb.get_ntype()); for (int T = 0; T < orb.get_ntype(); ++T) @@ -344,7 +349,7 @@ void cal_r_overlap_R::construct_orbs_and_nonlocal_and_orb_r(const UnitCell& ucel infoNL_.get_proj_dr_uniform(T, ip), false, true, - PARAM.inp.cal_force); + run.cal_force); delete[] rad; delete[] rab; @@ -424,7 +429,7 @@ void cal_r_overlap_R::construct_orbs_and_nonlocal_and_orb_r(const UnitCell& ucel } } - int map_size = PARAM.globalv.nlocal; + int map_size = basis.nlocal; int required_orbitals = 0; for (int it = 0; it < ucell.ntype; ++it) { @@ -462,27 +467,35 @@ void cal_r_overlap_R::construct_orbs_and_nonlocal_and_orb_r(const UnitCell& ucel } } -void cal_r_overlap_R::init(const UnitCell& ucell, const Parallel_Orbitals& pv, const LCAO_Orbitals& orb) +void cal_r_overlap_R::init(const UnitCell& ucell, + const Parallel_Orbitals& pv, + const LCAO_Orbitals& orb, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis) { ModuleBase::TITLE("cal_r_overlap_R", "init"); ModuleBase::timer::start("cal_r_overlap_R", "init"); this->ParaV = &pv; initialize_orb_table(ucell, orb); - construct_orbs_and_orb_r(ucell, orb); + construct_orbs_and_orb_r(ucell, orb, run, basis); ModuleBase::timer::end("cal_r_overlap_R", "init"); return; } -void cal_r_overlap_R::init_nonlocal(const UnitCell& ucell, const Parallel_Orbitals& pv, const LCAO_Orbitals& orb) +void cal_r_overlap_R::init_nonlocal(const UnitCell& ucell, + const Parallel_Orbitals& pv, + const LCAO_Orbitals& orb, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis) { ModuleBase::TITLE("cal_r_overlap_R", "init_nonlocal"); ModuleBase::timer::start("cal_r_overlap_R", "init_nonlocal"); this->ParaV = &pv; initialize_orb_table(ucell, orb); - construct_orbs_and_nonlocal_and_orb_r(ucell, orb); + construct_orbs_and_nonlocal_and_orb_r(ucell, orb, run, basis); ModuleBase::timer::end("cal_r_overlap_R", "init_nonlocal"); return; @@ -638,7 +651,15 @@ void cal_r_overlap_R::get_psi_r_beta(const UnitCell& ucell, } } -void cal_r_overlap_R::out_rR(const UnitCell& ucell, const Grid_Driver& gd, const int& istep, const int precision) +void cal_r_overlap_R::out_rR(const UnitCell& ucell, + const Grid_Driver& gd, + const int& istep, + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("cal_r_overlap_R", "out_rR"); ModuleBase::timer::start("cal_r_overlap_R", "out_rR"); @@ -676,13 +697,13 @@ void cal_r_overlap_R::out_rR(const UnitCell& ucell, const Grid_Driver& gd, const single_R_options.binary = binary; single_R_options.precision = precision; single_R_options.reduce = true; - single_R_options.temp_dir = PARAM.globalv.global_out_dir; + single_R_options.temp_dir = files.output_directory; std::stringstream tem1; - tem1 << PARAM.globalv.global_out_dir << "tmp-rr.csr"; + tem1 << files.output_directory << "tmp-rr.csr"; std::ofstream ofs_tem1; - if (GlobalV::DRANK == 0) + if (parallel.diagonalization_rank == 0) { if (binary) { @@ -709,22 +730,22 @@ void cal_r_overlap_R::out_rR(const UnitCell& ucell, const Grid_Driver& gd, const ModuleBase::Vector3 R_car = ModuleBase::Vector3(dRx, dRy, dRz) * ucell.latvec; int ir, ic; - for (int iw1 = 0; iw1 < PARAM.globalv.nlocal; iw1++) + for (int iw1 = 0; iw1 < basis.nlocal; iw1++) { ir = this->ParaV->global2local_row(iw1); if (ir >= 0) { - for (int iw2 = 0; iw2 < PARAM.globalv.nlocal; iw2++) + for (int iw2 = 0; iw2 < basis.nlocal; iw2++) { ic = this->ParaV->global2local_col(iw2); if (ic >= 0) { - int orb_index_row = iw1 / PARAM.globalv.npol; - int orb_index_col = iw2 / PARAM.globalv.npol; + int orb_index_row = iw1 / basis.npol; + int orb_index_col = iw2 / basis.npol; // The off-diagonal term in SOC calculaiton is zero, and the two diagonal terms are the same int new_index - = iw1 - PARAM.globalv.npol * orb_index_row + (iw2 - PARAM.globalv.npol * orb_index_col) * PARAM.globalv.npol; + = iw1 - basis.npol * orb_index_row + (iw2 - basis.npol * orb_index_col) * basis.npol; if (new_index == 0 || new_index == 3) { @@ -805,7 +826,7 @@ void cal_r_overlap_R::out_rR(const UnitCell& ucell, const Grid_Driver& gd, const { output_R_number++; - if (GlobalV::DRANK == 0) + if (parallel.diagonalization_rank == 0) { if (binary) { @@ -821,7 +842,7 @@ void cal_r_overlap_R::out_rR(const UnitCell& ucell, const Grid_Driver& gd, const for (int direction = 0; direction < 3; ++direction) { - if (GlobalV::DRANK == 0) + if (parallel.diagonalization_rank == 0) { if (binary) { @@ -835,7 +856,11 @@ void cal_r_overlap_R::out_rR(const UnitCell& ucell, const Grid_Driver& gd, const if (rR_nonzero_num[direction]) { - ModuleIO::output_single_R(ofs_tem1, psi_r_psi_sparse[direction], *(this->ParaV), single_R_options); + ModuleIO::output_single_R(ofs_tem1, + psi_r_psi_sparse[direction], + *(this->ParaV), + single_R_options, + parallel); } else { @@ -845,26 +870,26 @@ void cal_r_overlap_R::out_rR(const UnitCell& ucell, const Grid_Driver& gd, const } } - if (GlobalV::DRANK == 0) + if (parallel.diagonalization_rank == 0) { std::stringstream ssr; - if (PARAM.inp.calculation == "md" && !PARAM.inp.out_app_flag) + if (run.calculation == "md" && !output.append) { - ssr << PARAM.globalv.global_matrix_dir << "rrg" << step << ".csr"; + ssr << files.matrix_directory << "rrg" << step << ".csr"; } else { - ssr << PARAM.globalv.global_out_dir << "rr.csr"; + ssr << files.output_directory << "rr.csr"; } ofs_tem1.close(); ModuleIO::detail::finalize_rr_sparse_file(ssr.str(), tem1.str(), step, - PARAM.globalv.nlocal, + basis.nlocal, output_R_number, binary, - PARAM.inp.calculation == "md" && PARAM.inp.out_app_flag && step, + run.calculation == "md" && output.append && step, "cal_r_overlap_R::out_rR"); std::remove(tem1.str().c_str()); @@ -877,7 +902,12 @@ void cal_r_overlap_R::out_rR(const UnitCell& ucell, const Grid_Driver& gd, const void cal_r_overlap_R::out_rR_other(const UnitCell& ucell, const int& istep, const std::set>& output_R_coor, - const int precision) + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("cal_r_overlap_R", "out_rR_other"); ModuleBase::timer::start("cal_r_overlap_R", "out_rR_other"); @@ -893,12 +923,12 @@ void cal_r_overlap_R::out_rR_other(const UnitCell& ucell, single_R_options.binary = binary; single_R_options.precision = precision; single_R_options.reduce = true; - single_R_options.temp_dir = PARAM.globalv.global_out_dir; + single_R_options.temp_dir = files.output_directory; std::stringstream tem1; - tem1 << PARAM.globalv.global_out_dir << "tmp-rr-other.csr"; + tem1 << files.output_directory << "tmp-rr-other.csr"; std::ofstream ofs_tem1; - if (GlobalV::DRANK == 0) + if (parallel.diagonalization_rank == 0) { if (binary) { @@ -915,13 +945,13 @@ void cal_r_overlap_R::out_rR_other(const UnitCell& ucell, } std::stringstream ssr; - if (PARAM.inp.calculation == "md" && !PARAM.inp.out_app_flag) + if (run.calculation == "md" && !output.append) { - ssr << PARAM.globalv.global_matrix_dir << "rrg" << step << ".csr"; + ssr << files.matrix_directory << "rrg" << step << ".csr"; } else { - ssr << PARAM.globalv.global_out_dir << "rr.csr"; + ssr << files.output_directory << "rr.csr"; } for (auto& R_coor: output_R_coor) @@ -936,22 +966,22 @@ void cal_r_overlap_R::out_rR_other(const UnitCell& ucell, int ir = 0; int ic = 0; - for (int iw1 = 0; iw1 < PARAM.globalv.nlocal; iw1++) + for (int iw1 = 0; iw1 < basis.nlocal; iw1++) { ir = this->ParaV->global2local_row(iw1); if (ir >= 0) { - for (int iw2 = 0; iw2 < PARAM.globalv.nlocal; iw2++) + for (int iw2 = 0; iw2 < basis.nlocal; iw2++) { ic = this->ParaV->global2local_col(iw2); if (ic >= 0) { - int orb_index_row = iw1 / PARAM.globalv.npol; - int orb_index_col = iw2 / PARAM.globalv.npol; + int orb_index_row = iw1 / basis.npol; + int orb_index_col = iw2 / basis.npol; // The off-diagonal term in SOC calculaiton is zero, and the two diagonal terms are the same int new_index - = iw1 - PARAM.globalv.npol * orb_index_row + (iw2 - PARAM.globalv.npol * orb_index_col) * PARAM.globalv.npol; + = iw1 - basis.npol * orb_index_row + (iw2 - basis.npol * orb_index_col) * basis.npol; if (new_index == 0 || new_index == 3) { @@ -1034,7 +1064,7 @@ void cal_r_overlap_R::out_rR_other(const UnitCell& ucell, } output_R_number++; - if (GlobalV::DRANK == 0) + if (parallel.diagonalization_rank == 0) { if (binary) // .dat { @@ -1050,7 +1080,7 @@ void cal_r_overlap_R::out_rR_other(const UnitCell& ucell, for (int direction = 0; direction < 3; ++direction) { - if (GlobalV::DRANK == 0) + if (parallel.diagonalization_rank == 0) { if (binary) { @@ -1064,7 +1094,11 @@ void cal_r_overlap_R::out_rR_other(const UnitCell& ucell, if (rR_nonzero_num[direction]) { - ModuleIO::output_single_R(ofs_tem1, psi_r_psi_sparse[direction], *(this->ParaV), single_R_options); + ModuleIO::output_single_R(ofs_tem1, + psi_r_psi_sparse[direction], + *(this->ParaV), + single_R_options, + parallel); } else { @@ -1073,16 +1107,16 @@ void cal_r_overlap_R::out_rR_other(const UnitCell& ucell, } } - if (GlobalV::DRANK == 0) + if (parallel.diagonalization_rank == 0) { ofs_tem1.close(); ModuleIO::detail::finalize_rr_sparse_file(ssr.str(), tem1.str(), step, - PARAM.globalv.nlocal, + basis.nlocal, output_R_number, binary, - PARAM.inp.calculation == "md" && PARAM.inp.out_app_flag && step, + run.calculation == "md" && output.append && step, "cal_r_overlap_R::out_rR_other"); std::remove(tem1.str().c_str()); } diff --git a/source/source_io/module_hs/cal_r_overlap_R.h b/source/source_io/module_hs/cal_r_overlap_R.h index 1ac5c0f47f1..2cb4262416d 100644 --- a/source/source_io/module_hs/cal_r_overlap_R.h +++ b/source/source_io/module_hs/cal_r_overlap_R.h @@ -15,6 +15,8 @@ #include "source_lcao/center2_orb.h" #include "source_lcao/module_ri/abfs-vector3_order.h" +#include "source_context/context_types.h" + #include #include #include @@ -31,8 +33,16 @@ class cal_r_overlap_R double sparse_threshold = 1e-10; bool binary = false; - void init(const UnitCell& ucell, const Parallel_Orbitals& pv, const LCAO_Orbitals& orb); - void init_nonlocal(const UnitCell& ucell, const Parallel_Orbitals& pv, const LCAO_Orbitals& orb); + void init(const UnitCell& ucell, + const Parallel_Orbitals& pv, + const LCAO_Orbitals& orb, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis); + void init_nonlocal(const UnitCell& ucell, + const Parallel_Orbitals& pv, + const LCAO_Orbitals& orb, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis); ModuleBase::Vector3 get_psi_r_psi(const ModuleBase::Vector3& R1, const int& T1, const int& L1, @@ -64,16 +74,35 @@ class cal_r_overlap_R const int& N1, const ModuleBase::Vector3& R2, const int& T2); - void out_rR(const UnitCell& ucell, const Grid_Driver& gd, const int& istep, const int precision = 16); + void out_rR(const UnitCell& ucell, + const Grid_Driver& gd, + const int& istep, + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::MatrixOutputConfig& output); void out_rR_other(const UnitCell& ucell, const int& istep, const std::set>& output_R_coor, - const int precision = 16); + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::MatrixOutputConfig& output); private: void initialize_orb_table(const UnitCell& ucell, const LCAO_Orbitals& orb); - void construct_orbs_and_orb_r(const UnitCell& ucell, const LCAO_Orbitals& orb); - void construct_orbs_and_nonlocal_and_orb_r(const UnitCell& ucell, const LCAO_Orbitals& orb); + void construct_orbs_and_orb_r(const UnitCell& ucell, + const LCAO_Orbitals& orb, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis); + void construct_orbs_and_nonlocal_and_orb_r(const UnitCell& ucell, + const LCAO_Orbitals& orb, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis); std::vector iw2ia; std::vector iw2iL; diff --git a/source/source_io/module_hs/output_mat_sparse.cpp b/source/source_io/module_hs/output_mat_sparse.cpp index d55409ba481..c7ed62a2aa2 100644 --- a/source/source_io/module_hs/output_mat_sparse.cpp +++ b/source/source_io/module_hs/output_mat_sparse.cpp @@ -16,7 +16,14 @@ void output_mat_sparse(const MatSparseOutputOptions& options, const Grid_Driver& grid, const K_Vectors& kv, hamilt::Hamilt* p_ham, - Plus_U* p_dftu) + Plus_U* p_dftu, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SpinConfig& spin, + const ModuleContext::MatrixOutputConfig& output) { LCAO_HS_Arrays HS_Arrays; // store sparse arrays @@ -33,7 +40,12 @@ void output_mat_sparse(const MatSparseOutputOptions& options, "trs1_nao.csr", options.binary, options.sparse_threshold, - options.t_precision); + options.t_precision, + run, + files, + parallel, + logs, + output); } //! generate a file containing the derivatives of the Hamiltonian matrix (in Ry/Bohr) @@ -50,7 +62,14 @@ void output_mat_sparse(const MatSparseOutputOptions& options, kv, options.binary, options.sparse_threshold, - options.dh_precision); + options.dh_precision, + run, + files, + parallel, + logs, + basis, + spin, + output); } //! generate a file containing the derivatives of the overlap matrix (in Ry/Bohr) if (options.out_mat_ds) @@ -65,7 +84,14 @@ void output_mat_sparse(const MatSparseOutputOptions& options, kv, options.binary, options.sparse_threshold, - options.ds_precision); + options.ds_precision, + run, + files, + parallel, + logs, + basis, + spin, + output); } // add by jingan for out r_R matrix 2019.8.14 @@ -74,77 +100,13 @@ void output_mat_sparse(const MatSparseOutputOptions& options, cal_r_overlap_R r_matrix; r_matrix.binary = options.binary; r_matrix.sparse_threshold = options.sparse_threshold; - r_matrix.init(ucell, pv, orb); - r_matrix.out_rR(ucell, grid, istep, options.r_precision); + r_matrix.init(ucell, pv, orb, run, basis); + r_matrix.out_rR(ucell, grid, istep, options.r_precision, run, files, parallel, basis, output); } return; } -template -void output_mat_sparse(const bool& out_mat_dh, - const bool& out_mat_ds, - const bool& out_mat_t, - const bool& out_mat_r, - const int& istep, - const ModuleBase::matrix& v_eff, - const Parallel_Orbitals& pv, - const TwoCenterBundle& two_center_bundle, - const LCAO_Orbitals& orb, - UnitCell& ucell, - const Grid_Driver& grid, - const K_Vectors& kv, - hamilt::Hamilt* p_ham, - Plus_U* p_dftu) -{ - MatSparseOutputOptions options; - options.out_mat_dh = out_mat_dh; - options.out_mat_ds = out_mat_ds; - options.out_mat_t = out_mat_t; - options.out_mat_r = out_mat_r; - output_mat_sparse(options, - istep, - v_eff, - pv, - two_center_bundle, - orb, - ucell, - grid, - kv, - p_ham, - p_dftu); -} - -template void output_mat_sparse(const bool& out_mat_dh, - const bool& out_mat_ds, - const bool& out_mat_t, - const bool& out_mat_r, - const int& istep, - const ModuleBase::matrix& v_eff, - const Parallel_Orbitals& pv, - const TwoCenterBundle& two_center_bundle, - const LCAO_Orbitals& orb, - UnitCell& ucell, - const Grid_Driver& grid, - const K_Vectors& kv, - hamilt::Hamilt* p_ham, - Plus_U* p_dftu); - -template void output_mat_sparse>(const bool& out_mat_dh, - const bool& out_mat_ds, - const bool& out_mat_t, - const bool& out_mat_r, - const int& istep, - const ModuleBase::matrix& v_eff, - const Parallel_Orbitals& pv, - const TwoCenterBundle& two_center_bundle, - const LCAO_Orbitals& orb, - UnitCell& ucell, - const Grid_Driver& grid, - const K_Vectors& kv, - hamilt::Hamilt>* p_ham, - Plus_U* p_dftu); - template void output_mat_sparse(const MatSparseOutputOptions& options, const int& istep, const ModuleBase::matrix& v_eff, @@ -155,7 +117,14 @@ template void output_mat_sparse(const MatSparseOutputOptions& options, const Grid_Driver& grid, const K_Vectors& kv, hamilt::Hamilt* p_ham, - Plus_U* p_dftu); + Plus_U* p_dftu, + const ModuleContext::RunControl&, + const ModuleContext::FileSystemLayout&, + const ModuleContext::ParallelTopology&, + const ModuleContext::LogStreams&, + const ModuleContext::BasisInfo&, + const ModuleContext::SpinConfig&, + const ModuleContext::MatrixOutputConfig&); template void output_mat_sparse>(const MatSparseOutputOptions& options, const int& istep, @@ -167,6 +136,13 @@ template void output_mat_sparse>(const MatSparseOutputOptio const Grid_Driver& grid, const K_Vectors& kv, hamilt::Hamilt>* p_ham, - Plus_U* p_dftu); + Plus_U* p_dftu, + const ModuleContext::RunControl&, + const ModuleContext::FileSystemLayout&, + const ModuleContext::ParallelTopology&, + const ModuleContext::LogStreams&, + const ModuleContext::BasisInfo&, + const ModuleContext::SpinConfig&, + const ModuleContext::MatrixOutputConfig&); } // namespace ModuleIO diff --git a/source/source_io/module_hs/output_mat_sparse.h b/source/source_io/module_hs/output_mat_sparse.h index cfea51ca73e..fabc158f89a 100644 --- a/source/source_io/module_hs/output_mat_sparse.h +++ b/source/source_io/module_hs/output_mat_sparse.h @@ -7,6 +7,7 @@ #include "source_hamilt/hamilt.h" #include "source_cell/module_neighbor/sltk_grid_driver.h" #include "source_lcao/module_dftu/dftu.h" // mohan add 20251107 +#include "source_context/context_types.h" namespace ModuleIO { @@ -36,24 +37,14 @@ void output_mat_sparse(const MatSparseOutputOptions& options, const Grid_Driver& grid, const K_Vectors& kv, hamilt::Hamilt* p_ham, - Plus_U* p_dftu); - -/// @brief legacy bool-only interface kept for source compatibility -template -void output_mat_sparse(const bool& out_mat_dh, - const bool& out_mat_ds, - const bool& out_mat_t, - const bool& out_mat_r, - const int& istep, - const ModuleBase::matrix& v_eff, - const Parallel_Orbitals& pv, - const TwoCenterBundle& two_center_bundle, - const LCAO_Orbitals& orb, - UnitCell& ucell, - const Grid_Driver& grid, - const K_Vectors& kv, - hamilt::Hamilt* p_ham, - Plus_U* p_dftu); + Plus_U* p_dftu, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SpinConfig& spin, + const ModuleContext::MatrixOutputConfig& output); } // namespace ModuleIO #endif // OUTPUT_MAT_SPARSE_H diff --git a/source/source_io/module_hs/single_R_io.cpp b/source/source_io/module_hs/single_R_io.cpp index 174742dc994..4efb84be1f2 100644 --- a/source/source_io/module_hs/single_R_io.cpp +++ b/source/source_io/module_hs/single_R_io.cpp @@ -1,7 +1,6 @@ #include "single_R_io.h" #include "source_base/parallel_reduce.h" #include "source_base/global_function.h" -#include "source_base/global_variable.h" #include #include @@ -24,7 +23,8 @@ template void ModuleIO::output_single_R(std::ofstream& ofs, const SparseRBlock& XR, const Parallel_Orbitals& pv, - const SparseWriteOptions& options) + const SparseWriteOptions& options, + const ModuleContext::ParallelTopology& parallel) { const int nlocal = pv.get_global_row_size(); if (nlocal <= 0) @@ -38,12 +38,12 @@ void ModuleIO::output_single_R(std::ofstream& ofs, indptr.push_back(0); std::stringstream tem1; - tem1 << options.temp_dir << std::to_string(GlobalV::DRANK) + tem1 << options.temp_dir << std::to_string(parallel.diagonalization_rank) << "temp_sparse_indices.dat"; std::ofstream ofs_tem1; std::ifstream ifs_tem1; - if (!options.reduce || GlobalV::DRANK == 0) + if (!options.reduce || parallel.diagonalization_rank == 0) { if (options.binary) { @@ -88,7 +88,7 @@ void ModuleIO::output_single_R(std::ofstream& ofs, Parallel_Reduce::reduce_all(line.data(), nlocal); } - if (!options.reduce || GlobalV::DRANK == 0) + if (!options.reduce || parallel.diagonalization_rank == 0) { long long nonzeros_count = 0; for (int col = 0; col < nlocal; ++col) @@ -116,7 +116,7 @@ void ModuleIO::output_single_R(std::ofstream& ofs, } } - if (!options.reduce || GlobalV::DRANK == 0) + if (!options.reduce || parallel.diagonalization_rank == 0) { if (options.binary) { @@ -161,9 +161,11 @@ void ModuleIO::output_single_R(std::ofstream& ofs, template void ModuleIO::output_single_R(std::ofstream& ofs, const SparseRBlock& XR, const Parallel_Orbitals& pv, - const SparseWriteOptions& options); + const SparseWriteOptions& options, + const ModuleContext::ParallelTopology& parallel); template void ModuleIO::output_single_R>(std::ofstream& ofs, const SparseRBlock>& XR, const Parallel_Orbitals& pv, - const SparseWriteOptions& options); + const SparseWriteOptions& options, + const ModuleContext::ParallelTopology& parallel); diff --git a/source/source_io/module_hs/single_R_io.h b/source/source_io/module_hs/single_R_io.h index 405a5bee4e2..269efb74e8a 100644 --- a/source/source_io/module_hs/single_R_io.h +++ b/source/source_io/module_hs/single_R_io.h @@ -2,6 +2,7 @@ #define SINGLE_R_IO_H #include "write_HS_sparse.h" +#include "source_context/context_types.h" #include @@ -11,7 +12,8 @@ namespace ModuleIO void output_single_R(std::ofstream& ofs, const SparseRBlock& XR, const Parallel_Orbitals& pv, - const SparseWriteOptions& options); + const SparseWriteOptions& options, + const ModuleContext::ParallelTopology& parallel); } #endif diff --git a/source/source_io/module_hs/write_HS.h b/source/source_io/module_hs/write_HS.h index 43cd65b6d85..0915a416ace 100644 --- a/source/source_io/module_hs/write_HS.h +++ b/source/source_io/module_hs/write_HS.h @@ -4,10 +4,9 @@ #include #include -//#include "source_base/global_function.h" -//#include "source_base/global_variable.h" #include "source_basis/module_ao/parallel_orbitals.h" // use Parallel_Orbitals #include "source_hamilt/hamilt.h" +#include "source_context/context_types.h" // mohan add this file 2010-09-10 @@ -15,7 +14,6 @@ namespace ModuleIO { template void write_hsk( - const std::string &global_out_dir, const int nspin, const int nks, const int nkstot, @@ -23,10 +21,13 @@ namespace ModuleIO const std::vector &isk, hamilt::Hamilt* p_hamilt, const Parallel_Orbitals &pv, - const bool gamma_only, - const bool out_app_flag, const int istep, - std::ofstream &ofs_running); + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output); /// @brief save a square matrix, such as H(k) and S(k) /// @param[in] istep : the step of the calculation @@ -48,6 +49,7 @@ namespace ModuleIO const std::string& file_name, const Parallel_2D& pv, const int drank, + const ModuleContext::SolverConfig& solver, const bool reduce = true); } diff --git a/source/source_io/module_hs/write_HS.hpp b/source/source_io/module_hs/write_HS.hpp index 04c677f98b1..469329ca36c 100644 --- a/source/source_io/module_hs/write_HS.hpp +++ b/source/source_io/module_hs/write_HS.hpp @@ -1,6 +1,5 @@ #include "write_HS.h" -#include "source_io/module_parameter/parameter.h" #include "source_base/parallel_reduce.h" #include "source_base/timer.h" #include "source_base/tool_quit.h" @@ -10,31 +9,33 @@ template void ModuleIO::write_hsk( - const std::string &global_out_dir, - const int nspin, + const int nspin, const int nks, const int nkstot, const std::vector &ik2iktot, const std::vector &isk, - hamilt::Hamilt* p_hamilt, - const Parallel_Orbitals &pv, - const bool gamma_only, - const bool out_app_flag, - const int istep, - std::ofstream &ofs_running) + hamilt::Hamilt* p_hamilt, + const Parallel_Orbitals &pv, + const int istep, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output) { - ofs_running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" + *logs.running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" ">>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; - ofs_running << " | " + *logs.running << " | " " |" << std::endl; - ofs_running << " | Write Hamiltonian matrix H(k) or overlap matrix S(k) in numerical |" << std::endl; - ofs_running << " | atomic orbitals at each k-point. |" << std::endl; - ofs_running << " | " + *logs.running << " | Write Hamiltonian matrix H(k) or overlap matrix S(k) in numerical |" << std::endl; + *logs.running << " | atomic orbitals at each k-point. |" << std::endl; + *logs.running << " | " " |" << std::endl; - ofs_running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" + *logs.running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" ">>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; - ofs_running << "\n WRITE H(k) OR S(k)" << std::endl; + *logs.running << "\n WRITE H(k) OR S(k)" << std::endl; for (int ik = 0; ik < nks; ++ik) { @@ -50,20 +51,21 @@ void ModuleIO::write_hsk( const int out_label=1; // 1: .txt, 2: .dat - std::string h_fn = ModuleIO::filename_output(global_out_dir, + std::string h_fn = ModuleIO::filename_output(files.output_directory, "hk","nao",ik,ik2iktot,nspin,nkstot, - out_label,out_app_flag,gamma_only,istep); + out_label,output.append,basis.gamma_only_local,istep); ModuleIO::save_mat(istep, h_mat.p, - PARAM.globalv.nlocal, + basis.nlocal, bit, - PARAM.inp.out_mat_hs[1], + output.hs_k.precision, 1, - out_app_flag, + output.append, h_fn, pv, - GlobalV::DRANK); + parallel.diagonalization_rank, + solver); // mohan note 2025-06-02 // for overlap matrix, the two spin channels yield the same matrix @@ -74,22 +76,23 @@ void ModuleIO::write_hsk( continue; } - std::string s_fn = ModuleIO::filename_output(global_out_dir, + std::string s_fn = ModuleIO::filename_output(files.output_directory, "sk","nao",ik,ik2iktot,nspin,nkstot, - out_label,out_app_flag,gamma_only,istep); + out_label,output.append,basis.gamma_only_local,istep); - ofs_running << " The output filename is " << s_fn << std::endl; + *logs.running << " The output filename is " << s_fn << std::endl; ModuleIO::save_mat(istep, s_mat.p, - PARAM.globalv.nlocal, + basis.nlocal, bit, - PARAM.inp.out_mat_hs[1], + output.hs_k.precision, 1, - out_app_flag, + output.append, s_fn, pv, - GlobalV::DRANK); + parallel.diagonalization_rank, + solver); } // end ik } @@ -106,6 +109,7 @@ void ModuleIO::save_mat(const int istep, const std::string& filename, const Parallel_2D& pv, const int drank, + const ModuleContext::SolverConfig& solver, const bool reduce) { ModuleBase::TITLE("ModuleIO", "save_mat"); @@ -147,7 +151,7 @@ void ModuleIO::save_mat(const int istep, if (ic >= 0) { int iic; - if (ModuleBase::GlobalFunc::IS_COLUMN_MAJOR_KS_SOLVER(PARAM.inp.ks_solver)) + if (ModuleBase::GlobalFunc::IS_COLUMN_MAJOR_KS_SOLVER(solver.ks_solver)) { iic = ir + ic * pv.nrow; } @@ -247,7 +251,7 @@ void ModuleIO::save_mat(const int istep, if (ic >= 0) { int iic=0; - if (ModuleBase::GlobalFunc::IS_COLUMN_MAJOR_KS_SOLVER(PARAM.inp.ks_solver)) + if (ModuleBase::GlobalFunc::IS_COLUMN_MAJOR_KS_SOLVER(solver.ks_solver)) { iic = ir + ic * pv.nrow; } diff --git a/source/source_io/module_hs/write_HS_R.cpp b/source/source_io/module_hs/write_HS_R.cpp index f9ebe893ec8..c31cbcac5ba 100644 --- a/source/source_io/module_hs/write_HS_R.cpp +++ b/source/source_io/module_hs/write_HS_R.cpp @@ -2,7 +2,6 @@ #include "source_base/timer.h" #include "source_base/tool_quit.h" -#include "source_io/module_parameter/parameter.h" #include "source_lcao/LCAO_HS_arrays.hpp" #include "source_lcao/spar_dh.h" #include "source_lcao/spar_hsr.h" @@ -24,7 +23,14 @@ void ModuleIO::output_dSR(const int& istep, const K_Vectors& kv, const bool& binary, const double& sparse_thr, - const int precision) + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SpinConfig& spin, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("ModuleIO", "output_dSR"); ModuleBase::timer::start("ModuleIO", "output_dSR"); @@ -32,7 +38,20 @@ void ModuleIO::output_dSR(const int& istep, sparse_format::cal_dS(ucell, pv, HS_Arrays, grid, two_center_bundle, orb, sparse_thr); // mohan update 2024-04-01 - ModuleIO::save_dH_sparse(istep, pv, HS_Arrays, sparse_thr, binary, "s", precision); + ModuleIO::save_dH_sparse(istep, + pv, + HS_Arrays, + sparse_thr, + binary, + "s", + precision, + run, + files, + parallel, + logs, + basis, + spin, + output); sparse_format::destroy_dH_R_sparse(HS_Arrays); @@ -51,18 +70,25 @@ void ModuleIO::output_dHR(const int& istep, const K_Vectors& kv, const bool& binary, const double& sparse_thr, - const int precision) + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SpinConfig& spin, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("ModuleIO", "output_dHR"); ModuleBase::timer::start("ModuleIO", "output_dHR"); - GlobalV::ofs_running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; - GlobalV::ofs_running << " | |" << std::endl; - GlobalV::ofs_running << " | #Print out dH/dR# |" << std::endl; - GlobalV::ofs_running << " | |" << std::endl; - GlobalV::ofs_running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; + *logs.running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; + *logs.running << " | |" << std::endl; + *logs.running << " | #Print out dH/dR# |" << std::endl; + *logs.running << " | |" << std::endl; + *logs.running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; - const int nspin = PARAM.inp.nspin; + const int nspin = spin.nspin; if (nspin == 1 || nspin == 4) { @@ -79,7 +105,20 @@ void ModuleIO::output_dHR(const int& istep, } } // mohan update 2024-04-01 - ModuleIO::save_dH_sparse(istep, pv, HS_Arrays, sparse_thr, binary, "h", precision); + ModuleIO::save_dH_sparse(istep, + pv, + HS_Arrays, + sparse_thr, + binary, + "h", + precision, + run, + files, + parallel, + logs, + basis, + spin, + output); sparse_format::destroy_dH_R_sparse(HS_Arrays); @@ -94,19 +133,23 @@ void ModuleIO::output_SR(Parallel_Orbitals& pv, const std::string& SR_filename, const bool& binary, const double& sparse_thr, - const int precision) + const int precision, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::SpinConfig& spin) { ModuleBase::TITLE("ModuleIO", "output_SR"); ModuleBase::timer::start("ModuleIO", "output_SR"); - GlobalV::ofs_running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; - GlobalV::ofs_running << " | |" << std::endl; - GlobalV::ofs_running << " | #Print out overlap matrix S(R)# |" << std::endl; - GlobalV::ofs_running << " | |" << std::endl; - GlobalV::ofs_running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; + *logs.running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; + *logs.running << " | |" << std::endl; + *logs.running << " | #Print out overlap matrix S(R)# |" << std::endl; + *logs.running << " | |" << std::endl; + *logs.running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; std::cout << " Overlap matrix file is in " << SR_filename << std::endl; - GlobalV::ofs_running << " Overlap matrix file is in " << SR_filename << std::endl; + *logs.running << " Overlap matrix file is in " << SR_filename << std::endl; LCAO_HS_Arrays HS_Arrays; @@ -127,21 +170,23 @@ void ModuleIO::output_SR(Parallel_Orbitals& pv, options.precision = precision; options.istep = istep; options.reduce = true; - options.temp_dir = PARAM.globalv.global_out_dir; + options.temp_dir = files.output_directory; - if (PARAM.inp.nspin == 4) + if (spin.nspin == 4) { ModuleIO::save_sparse(HS_Arrays.SR_soc_sparse, HS_Arrays.all_R_coor, pv, - options); + options, + parallel); } else { ModuleIO::save_sparse(HS_Arrays.SR_sparse, HS_Arrays.all_R_coor, pv, - options); + options, + parallel); } sparse_format::destroy_HS_R_sparse(HS_Arrays); @@ -160,27 +205,32 @@ void ModuleIO::output_TR(const int istep, const std::string& TR_filename, const bool& binary, const double& sparse_thr, - const int precision) + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("ModuleIO", "output_TR"); ModuleBase::timer::start("ModuleIO", "output_TR"); - GlobalV::ofs_running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; - GlobalV::ofs_running << " | |" << std::endl; - GlobalV::ofs_running << " | #Print out kinetic energy term matrix T(R)# |" << std::endl; - GlobalV::ofs_running << " | |" << std::endl; - GlobalV::ofs_running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; + *logs.running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; + *logs.running << " | |" << std::endl; + *logs.running << " | #Print out kinetic energy term matrix T(R)# |" << std::endl; + *logs.running << " | |" << std::endl; + *logs.running << " >>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>>" << std::endl; std::stringstream sst; - if (PARAM.inp.calculation == "md" && !PARAM.inp.out_app_flag) + if (run.calculation == "md" && !output.append) { - sst << PARAM.globalv.global_matrix_dir << TR_filename << "g" << istep; - GlobalV::ofs_running << " T(R) data are in file: " << sst.str() << std::endl; + sst << files.matrix_directory << TR_filename << "g" << istep; + *logs.running << " T(R) data are in file: " << sst.str() << std::endl; } else { - sst << PARAM.globalv.global_out_dir << TR_filename; - GlobalV::ofs_running << " T(R) data are in file: " << sst.str() << std::endl; + sst << files.output_directory << TR_filename; + *logs.running << " T(R) data are in file: " << sst.str() << std::endl; } sparse_format::cal_TR(ucell, pv, HS_Arrays, grid, two_center_bundle, orb, sparse_thr); @@ -192,12 +242,14 @@ void ModuleIO::output_TR(const int istep, options.precision = precision; options.istep = istep; options.reduce = true; - options.temp_dir = PARAM.globalv.global_out_dir; + options.append = run.calculation == "md" && output.append; + options.temp_dir = files.output_directory; ModuleIO::save_sparse(HS_Arrays.TR_sparse, HS_Arrays.all_R_coor, pv, - options); + options, + parallel); sparse_format::destroy_T_R_sparse(HS_Arrays); @@ -211,14 +263,22 @@ template void ModuleIO::output_SR(Parallel_Orbitals& pv, const std::string& SR_filename, const bool& binary, const double& sparse_thr, - const int precision); + const int precision, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::SpinConfig& spin); template void ModuleIO::output_SR>(Parallel_Orbitals& pv, const Grid_Driver& grid, hamilt::Hamilt>* p_ham, const std::string& SR_filename, const bool& binary, const double& sparse_thr, - const int precision); + const int precision, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::SpinConfig& spin); #include "source_lcao/module_hcontainer/hcontainer_funcs.h" #include "source_lcao/module_hcontainer/output_hcontainer.h" @@ -304,7 +364,9 @@ void ModuleIO::write_hsr(const std::vector*>& hr_vec, const bool append, const int* iat2iwt, const int nat, - const int istep) + const int istep, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel) { const int nspin = hr_vec.size(); assert(nspin > 0); @@ -325,10 +387,9 @@ void ModuleIO::write_hsr(const std::vector*>& hr_vec, hamilt::HContainer hr_serial(*hr_vec[ispin]); #endif - if (GlobalV::MY_RANK == 0) + if (parallel.world_rank == 0) { - std::string fname = PARAM.globalv.global_out_dir - + hsr_gen_fname("hrs", ispin, append, istep); + std::string fname = files.output_directory + hsr_gen_fname("hrs", ispin, append, istep); write_hcontainer_csr(fname, ucell, precision, &hr_serial, istep, ispin, nspin, "H"); } } @@ -348,10 +409,9 @@ void ModuleIO::write_hsr(const std::vector*>& hr_vec, hamilt::HContainer sr_serial(*sr); #endif - if (GlobalV::MY_RANK == 0) + if (parallel.world_rank == 0) { - std::string fname = PARAM.globalv.global_out_dir - + hsr_gen_fname("srs", 0, append, istep); + std::string fname = files.output_directory + hsr_gen_fname("srs", 0, append, istep); write_hcontainer_csr(fname, ucell, precision, &sr_serial, istep, 0, 1, "S"); } } @@ -369,12 +429,14 @@ template void ModuleIO::write_hsr( const std::vector*>&, const hamilt::HContainer*, const UnitCell*, const int, const Parallel_2D&, - const bool, const int*, const int, const int); + const bool, const int*, const int, const int, + const ModuleContext::FileSystemLayout&, const ModuleContext::ParallelTopology&); template void ModuleIO::write_hsr>( const std::vector>*>&, const hamilt::HContainer>*, const UnitCell*, const int, const Parallel_2D&, - const bool, const int*, const int, const int); + const bool, const int*, const int, const int, + const ModuleContext::FileSystemLayout&, const ModuleContext::ParallelTopology&); template @@ -387,7 +449,10 @@ void ModuleIO::write_matrix_r(const std::string& matrix_label, const bool append, const int* iat2iwt, const int nat, - const int istep) + const int istep, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel) { const int nspin = matrices.size(); assert(nspin > 0); @@ -398,13 +463,13 @@ void ModuleIO::write_matrix_r(const std::string& matrix_label, // Generate filename std::string fname = dhr_gen_fname(matrix_label, ispin, append, istep); - if (PARAM.inp.calculation == "md" && !PARAM.inp.out_app_flag) + if (run.calculation == "md" && !append) { - fname = PARAM.globalv.global_matrix_dir + fname; + fname = files.matrix_directory + fname; } else { - fname = PARAM.globalv.global_out_dir + fname; + fname = files.output_directory + fname; } // Gather parallel matrix to serial @@ -417,7 +482,7 @@ void ModuleIO::write_matrix_r(const std::string& matrix_label, hamilt::HContainer matrix_serial(&serialV); hamilt::gatherParallels(*matrices[ispin], &matrix_serial, 0); - if (GlobalV::MY_RANK == 0) + if (parallel.world_rank == 0) { write_hcontainer_csr(fname, ucell, precision, &matrix_serial, istep, ispin, nspin, description); } @@ -438,7 +503,10 @@ template void ModuleIO::write_matrix_r( const bool, const int*, const int, - const int); + const int, + const ModuleContext::RunControl&, + const ModuleContext::FileSystemLayout&, + const ModuleContext::ParallelTopology&); template void ModuleIO::write_matrix_r>( const std::string&, @@ -450,4 +518,7 @@ template void ModuleIO::write_matrix_r>( const bool, const int*, const int, - const int); + const int, + const ModuleContext::RunControl&, + const ModuleContext::FileSystemLayout&, + const ModuleContext::ParallelTopology&); diff --git a/source/source_io/module_hs/write_HS_R.h b/source/source_io/module_hs/write_HS_R.h index 603fbd9ac05..dbff343584f 100644 --- a/source/source_io/module_hs/write_HS_R.h +++ b/source/source_io/module_hs/write_HS_R.h @@ -7,6 +7,7 @@ #include "source_hamilt/hamilt.h" #include "source_lcao/LCAO_HS_arrays.hpp" #include "source_lcao/module_dftu/dftu.h" // mohan add 20251107 +#include "source_context/context_types.h" #ifdef __EXX #include "RI/global/Tensor.h" // for RI::Tensor @@ -23,9 +24,16 @@ void output_dHR(const int& istep, const TwoCenterBundle& two_center_bundle, const LCAO_Orbitals& orb, const K_Vectors& kv, - const bool& binary = false, - const double& sparse_threshold = 1e-10, - const int precision = 16); + const bool& binary, + const double& sparse_threshold, + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SpinConfig& spin, + const ModuleContext::MatrixOutputConfig& output); void output_dSR(const int& istep, const UnitCell& ucell, @@ -35,9 +43,16 @@ void output_dSR(const int& istep, const TwoCenterBundle& two_center_bundle, const LCAO_Orbitals& orb, const K_Vectors& kv, - const bool& binary = false, - const double& sparse_thr = 1e-10, - const int precision = 16); + const bool& binary, + const double& sparse_thr, + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SpinConfig& spin, + const ModuleContext::MatrixOutputConfig& output); void output_TR(const int istep, const UnitCell& ucell, @@ -46,19 +61,28 @@ void output_TR(const int istep, const Grid_Driver& grid, const TwoCenterBundle& two_center_bundle, const LCAO_Orbitals& orb, - const std::string& TR_filename = "trs1_nao.csr", - const bool& binary = false, - const double& sparse_threshold = 1e-10, - const int precision = 16); + const std::string& TR_filename, + const bool& binary, + const double& sparse_threshold, + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::MatrixOutputConfig& output); template void output_SR(Parallel_Orbitals& pv, const Grid_Driver& grid, hamilt::Hamilt* p_ham, - const std::string& SR_filename = "srs1_nao.csr", - const bool& binary = false, - const double& sparse_threshold = 1e-10, - const int precision = 16); + const std::string& SR_filename, + const bool& binary, + const double& sparse_threshold, + const int precision, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::SpinConfig& spin); /// Generate filename for HR/SR CSR output. std::string hsr_gen_fname(const std::string& prefix, @@ -93,7 +117,9 @@ void write_hsr(const std::vector*>& hr_vec, const bool append, const int* iat2iwt, const int nat, - const int istep); + const int istep, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel); /// Write real-space matrix in CSR format (generic interface). template @@ -106,7 +132,10 @@ void write_matrix_r(const std::string& matrix_label, const bool append, const int* iat2iwt, const int nat, - const int istep); + const int istep, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel); } // namespace ModuleIO diff --git a/source/source_io/module_hs/write_HS_sparse.cpp b/source/source_io/module_hs/write_HS_sparse.cpp index 8fffba237e3..c6750f70a8f 100644 --- a/source/source_io/module_hs/write_HS_sparse.cpp +++ b/source/source_io/module_hs/write_HS_sparse.cpp @@ -3,7 +3,6 @@ #include "source_base/global_function.h" #include "source_base/parallel_reduce.h" #include "source_base/timer.h" -#include "source_io/module_parameter/parameter.h" #include "source_lcao/module_rt/td_info.h" #include "single_R_io.h" @@ -69,7 +68,7 @@ void open_sparse_file(std::ofstream& ofs, const ModuleIO::SparseWriteOptions& op { mode |= std::ios::binary; } - if (PARAM.inp.calculation == "md" && PARAM.inp.out_app_flag && options.istep) + if (options.append && options.istep) { mode |= std::ios::app; } @@ -143,7 +142,14 @@ void ModuleIO::save_dH_sparse(const int& istep, const double& sparse_thr, const bool& binary, const std::string& fileflag, - const int precision) { + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SpinConfig& spin, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("ModuleIO", "save_dH_sparse"); ModuleBase::timer::start("ModuleIO", "save_dH_sparse"); SparseWriteOptions single_R_options; @@ -151,7 +157,7 @@ void ModuleIO::save_dH_sparse(const int& istep, single_R_options.binary = binary; single_R_options.precision = precision; single_R_options.reduce = true; - single_R_options.temp_dir = PARAM.globalv.global_out_dir; + single_R_options.temp_dir = files.output_directory; auto& all_R_coor_ptr = HS_Arrays.all_R_coor; auto& output_R_coor_ptr = HS_Arrays.output_R_coor; @@ -170,11 +176,11 @@ void ModuleIO::save_dH_sparse(const int& istep, int step = istep; int spin_loop = 1; - if (PARAM.inp.nspin == 2) { + if (spin.nspin == 2) { spin_loop = 2; } - if (PARAM.inp.nspin != 4) + if (spin.nspin != 4) { for (int ispin = 0; ispin < spin_loop; ++ispin) { @@ -215,42 +221,42 @@ void ModuleIO::save_dH_sparse(const int& istep, std::stringstream sshy[2]; std::stringstream sshz[2]; - if (PARAM.inp.calculation == "md" && !PARAM.inp.out_app_flag) + if (run.calculation == "md" && !output.append) { - sshx[0] << PARAM.globalv.global_matrix_dir + sshx[0] << files.matrix_directory << "d"<(dHx_nonzero_num[ispin][count]); @@ -396,42 +402,48 @@ void ModuleIO::save_dH_sparse(const int& istep, for (int ispin = 0; ispin < spin_loop; ++ispin) { if (dHx_nonzero_num[ispin][count] > 0) { - if (PARAM.inp.nspin != 4) { + if (spin.nspin != 4) { output_single_R(g1x[ispin], dHRx_sparse_ptr[ispin][R_coor], pv, - single_R_options); + single_R_options, + parallel); } else { output_single_R(g1x[ispin], dHRx_soc_sparse_ptr[R_coor], pv, - single_R_options); + single_R_options, + parallel); } } if (dHy_nonzero_num[ispin][count] > 0) { - if (PARAM.inp.nspin != 4) { + if (spin.nspin != 4) { output_single_R(g1y[ispin], dHRy_sparse_ptr[ispin][R_coor], pv, - single_R_options); + single_R_options, + parallel); } else { output_single_R(g1y[ispin], dHRy_soc_sparse_ptr[R_coor], pv, - single_R_options); + single_R_options, + parallel); } } if (dHz_nonzero_num[ispin][count] > 0) { - if (PARAM.inp.nspin != 4) { + if (spin.nspin != 4) { output_single_R(g1z[ispin], dHRz_sparse_ptr[ispin][R_coor], pv, - single_R_options); + single_R_options, + parallel); } else { output_single_R(g1z[ispin], dHRz_soc_sparse_ptr[R_coor], pv, - single_R_options); + single_R_options, + parallel); } } } @@ -439,7 +451,7 @@ void ModuleIO::save_dH_sparse(const int& istep, count++; } - if (GlobalV::DRANK == 0) { + if (parallel.diagonalization_rank == 0) { for (int ispin = 0; ispin < spin_loop; ++ispin) { g1x[ispin].close(); } @@ -460,7 +472,8 @@ void ModuleIO::save_sparse( const SparseRMatrix& smat, const std::set& all_R_coor, const Parallel_Orbitals& pv, - const SparseWriteOptions& options) { + const SparseWriteOptions& options, + const ModuleContext::ParallelTopology& parallel) { ModuleBase::TITLE("ModuleIO", "save_sparse"); ModuleBase::timer::start("ModuleIO", "save_sparse"); const int nlocal = pv.get_global_row_size(); @@ -474,7 +487,7 @@ void ModuleIO::save_sparse( = count_nonzeros_by_R(smat, all_R_coor, options.threshold, options.reduce); const int output_R_number = count_output_R(nonzero_num); std::ofstream ofs; - if (!options.reduce || GlobalV::DRANK == 0) + if (!options.reduce || parallel.diagonalization_rank == 0) { open_sparse_file(ofs, options); write_sparse_header(ofs, options, nlocal, output_R_number); @@ -489,23 +502,23 @@ void ModuleIO::save_sparse( continue; } - if (!options.reduce || GlobalV::DRANK == 0) + if (!options.reduce || parallel.diagonalization_rank == 0) { write_R_record(ofs, R_coor, nonzero_num[count], options.binary); } if (smat.count(R_coor)) { - output_single_R(ofs, smat.at(R_coor), pv, options); + output_single_R(ofs, smat.at(R_coor), pv, options, parallel); } else { SparseRBlock empty_map; - output_single_R(ofs, empty_map, pv, options); + output_single_R(ofs, empty_map, pv, options, parallel); } ++count; } - if (!options.reduce || GlobalV::DRANK == 0) + if (!options.reduce || parallel.diagonalization_rank == 0) { ofs.close(); } @@ -517,10 +530,12 @@ template void ModuleIO::save_sparse( const SparseRMatrix&, const std::set&, const Parallel_Orbitals&, - const SparseWriteOptions&); + const SparseWriteOptions&, + const ModuleContext::ParallelTopology&); template void ModuleIO::save_sparse>( const SparseRMatrix>&, const std::set&, const Parallel_Orbitals&, - const SparseWriteOptions&); + const SparseWriteOptions&, + const ModuleContext::ParallelTopology&); diff --git a/source/source_io/module_hs/write_HS_sparse.h b/source/source_io/module_hs/write_HS_sparse.h index f6271ebefb7..4362fd5706c 100644 --- a/source/source_io/module_hs/write_HS_sparse.h +++ b/source/source_io/module_hs/write_HS_sparse.h @@ -3,6 +3,7 @@ #include "source_basis/module_ao/parallel_orbitals.h" #include "source_lcao/LCAO_HS_arrays.hpp" +#include "source_context/context_types.h" #include #include @@ -28,6 +29,7 @@ struct SparseWriteOptions int precision = 16; int istep = -1; bool reduce = true; + bool append = false; std::string temp_dir; }; @@ -36,14 +38,22 @@ void save_dH_sparse(const int& istep, LCAO_HS_Arrays& HS_Arrays, const double& sparse_thr, const bool& binary, - const std::string& fileflag = "h", - const int precision = 16); + const std::string& fileflag, + const int precision, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::LogStreams& logs, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SpinConfig& spin, + const ModuleContext::MatrixOutputConfig& output); template void save_sparse(const SparseRMatrix& smat, const std::set& all_R_coor, const Parallel_Orbitals& pv, - const SparseWriteOptions& options); + const SparseWriteOptions& options, + const ModuleContext::ParallelTopology& parallel); } // namespace ModuleIO #endif diff --git a/source/source_io/module_hs/write_H_terms.cpp b/source/source_io/module_hs/write_H_terms.cpp index a94e234a9a2..67bad8e9aa3 100644 --- a/source/source_io/module_hs/write_H_terms.cpp +++ b/source/source_io/module_hs/write_H_terms.cpp @@ -8,7 +8,6 @@ #include "source_io/module_hs/write_HS_R.h" #include "source_io/module_output/filename.h" #include "source_io/module_output/ucell_io.h" -#include "source_io/module_parameter/parameter.h" #include "source_lcao/module_gint/gint_interface.h" #include "source_lcao/module_hcontainer/hcontainer_funcs.h" #include "source_lcao/module_hcontainer/output_hcontainer.h" @@ -73,7 +72,10 @@ static void gather_and_write(const std::string& prefix, const int istep, const bool append, const int* iat2iwt, - const int nat) + const int nat, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel) { const int nbasis = hR.get_nbasis(); #ifdef __MPI @@ -83,17 +85,17 @@ static void gather_and_write(const std::string& prefix, serialV.set_atomic_trace(iat2iwt, nat, nbasis); hamilt::HContainer hr_serial(&serialV); hamilt::gatherParallels(hR, &hr_serial, 0); - if (GlobalV::MY_RANK == 0) + if (parallel.world_rank == 0) #endif { std::string fname; - if (PARAM.inp.calculation == "md" && !PARAM.inp.out_app_flag) + if (run.calculation == "md" && !append) { - fname = PARAM.globalv.global_matrix_dir + hsr_gen_fname(prefix, ispin, append, istep); + fname = files.matrix_directory + hsr_gen_fname(prefix, ispin, append, istep); } else { - fname = PARAM.globalv.global_out_dir + hsr_gen_fname(prefix, ispin, append, istep); + fname = files.output_directory + hsr_gen_fname(prefix, ispin, append, istep); } #ifdef __MPI write_hcontainer_csr(fname, &ucell, 8, &hr_serial, istep, ispin, nspin, label); @@ -112,14 +114,15 @@ static void write_hk_common(hamilt::HContainer& hR, const int istep, const bool append, const int* iat2iwt, - const int nat) + const int nat, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver) { const int nspin_k = (nspin == 2 ? 2 : 1); const int nks = kv.get_nks() / nspin_k; - const int nlocal = PARAM.globalv.nlocal; - const bool gamma_only = PARAM.globalv.gamma_only_local; - const std::string global_out_dir = PARAM.globalv.global_out_dir; - const bool out_app_flag = PARAM.inp.out_app_flag; + const int nlocal = basis.nlocal; for (int ik = 0; ik < nks; ++ik) { @@ -129,7 +132,7 @@ static void write_hk_common(hamilt::HContainer& hR, hamilt::folding_HR(hR, hk_global.data(), kvec_d, nlocal, 0); const int out_label = 1; - std::string fname = ModuleIO::filename_output(global_out_dir, + std::string fname = ModuleIO::filename_output(files.output_directory, prefix, "nao", ik, @@ -137,8 +140,8 @@ static void write_hk_common(hamilt::HContainer& hR, nspin, kv.get_nkstot(), out_label, - out_app_flag, - gamma_only, + append, + basis.gamma_only_local, istep); ModuleIO::save_mat(istep, hk_global.data(), @@ -146,14 +149,21 @@ static void write_hk_common(hamilt::HContainer& hR, false, 8, false, - out_app_flag, + append, fname, pv, - GlobalV::DRANK); + parallel.diagonalization_rank, + solver); } } -void write_h_t(WriteHParams& params) +void write_h_t(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("ModuleIO", "write_h_t"); ModuleBase::timer::start("ModuleIO", "write_h_t"); @@ -166,7 +176,7 @@ void write_h_t(WriteHParams& params) const K_Vectors& kv = *params.kv; const int nspin = params.nspin; const int istep = params.istep; - const bool append = params.append; + const bool append = output.append; const int* iat2iwt = params.iat2iwt; const int nat = params.nat; const bool also_hR = params.also_hR; @@ -182,18 +192,24 @@ void write_h_t(WriteHParams& params) tmp_ekinetic(nullptr, kv.kvec_d, &hR_tmp, &ucell, orb_cutoff, &gd, two_center_bundle.kinetic_orb.get()); tmp_ekinetic.contributeHR(); - write_hk_common(hR_tmp, "tk", ucell, pv, kv, nspin, istep, append, iat2iwt, nat); + write_hk_common(hR_tmp, "tk", ucell, pv, kv, nspin, istep, append, iat2iwt, nat, files, parallel, basis, solver); if (also_hR) { - gather_and_write("t", "T", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat); + gather_and_write("t", "T", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat, run, files, parallel); } } ModuleBase::timer::end("ModuleIO", "write_h_t"); } -void write_h_vnl(WriteHParams& params) +void write_h_vnl(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("ModuleIO", "write_h_vnl"); ModuleBase::timer::start("ModuleIO", "write_h_vnl"); @@ -206,7 +222,7 @@ void write_h_vnl(WriteHParams& params) const K_Vectors& kv = *params.kv; const int nspin = params.nspin; const int istep = params.istep; - const bool append = params.append; + const bool append = output.append; const int* iat2iwt = params.iat2iwt; const int nat = params.nat; const bool also_hR = params.also_hR; @@ -227,18 +243,24 @@ void write_h_vnl(WriteHParams& params) two_center_bundle.overlap_orb_beta.get()); tmp_nonlocal.contributeHR(); - write_hk_common(hR_tmp, "vnlk", ucell, pv, kv, nspin, istep, append, iat2iwt, nat); + write_hk_common(hR_tmp, "vnlk", ucell, pv, kv, nspin, istep, append, iat2iwt, nat, files, parallel, basis, solver); if (also_hR) { - gather_and_write("vnl", "V^NL", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat); + gather_and_write("vnl", "V^NL", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat, run, files, parallel); } } ModuleBase::timer::end("ModuleIO", "write_h_vnl"); } -void write_h_vl(WriteHParams& params) +void write_h_vl(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("ModuleIO", "write_h_vl"); ModuleBase::timer::start("ModuleIO", "write_h_vl"); @@ -251,7 +273,7 @@ void write_h_vl(WriteHParams& params) const K_Vectors& kv = *params.kv; const int nspin = params.nspin; const int istep = params.istep; - const bool append = params.append; + const bool append = output.append; const int* iat2iwt = params.iat2iwt; const int nat = params.nat; const bool also_hR = params.also_hR; @@ -267,18 +289,24 @@ void write_h_vl(WriteHParams& params) const double* v_local = pot->get_fixed_v(); // local pp, no Hxc ModuleGint::cal_gint_vl(v_local, &hR_tmp); - write_hk_common(hR_tmp, "vlk", ucell, pv, kv, nspin, istep, append, iat2iwt, nat); + write_hk_common(hR_tmp, "vlk", ucell, pv, kv, nspin, istep, append, iat2iwt, nat, files, parallel, basis, solver); if (also_hR) { - gather_and_write("vl", "V^L", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat); + gather_and_write("vl", "V^L", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat, run, files, parallel); } } ModuleBase::timer::end("ModuleIO", "write_h_vl"); } -void write_h_vh(WriteHParams& params) +void write_h_vh(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("ModuleIO", "write_h_vh"); ModuleBase::timer::start("ModuleIO", "write_h_vh"); @@ -292,7 +320,7 @@ void write_h_vh(WriteHParams& params) const K_Vectors& kv = *params.kv; const int nspin = params.nspin; const int istep = params.istep; - const bool append = params.append; + const bool append = output.append; const int* iat2iwt = params.iat2iwt; const int nat = params.nat; const bool also_hR = params.also_hR; @@ -310,18 +338,25 @@ void write_h_vh(WriteHParams& params) ModuleGint::cal_gint_vl(&v_h(ispin, 0), &hR_tmp); - write_hk_common(hR_tmp, "vhk", ucell, pv, kv, nspin, istep, append, iat2iwt, nat); + write_hk_common(hR_tmp, "vhk", ucell, pv, kv, nspin, istep, append, iat2iwt, nat, files, parallel, basis, solver); if (also_hR) { - gather_and_write("vh", "V^H", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat); + gather_and_write("vh", "V^H", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat, run, files, parallel); } } ModuleBase::timer::end("ModuleIO", "write_h_vh"); } -void write_h_vxc(WriteHParams& params) +void write_h_vxc(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SpinConfig& spin, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output) { ModuleBase::TITLE("ModuleIO", "write_h_vxc"); ModuleBase::timer::start("ModuleIO", "write_h_vxc"); @@ -335,7 +370,7 @@ void write_h_vxc(WriteHParams& params) const K_Vectors& kv = *params.kv; const int nspin = params.nspin; const int istep = params.istep; - const bool append = params.append; + const bool append = output.append; const int* iat2iwt = params.iat2iwt; const int nat = params.nat; const bool also_hR = params.also_hR; @@ -351,7 +386,8 @@ void write_h_vxc(WriteHParams& params) #else const double hse_omega = 0.0; #endif - std::tie(etxc, vtxc, v_xc) = XC_Functional::v_xc(nrxx, chg, &ucell, PARAM.inp.nspin, PARAM.globalv.domag, PARAM.globalv.domag_z, hybrid_alpha, hse_omega); + std::tie(etxc, vtxc, v_xc) + = XC_Functional::v_xc(nrxx, chg, &ucell, spin.nspin, spin.domag, spin.domag_z, hybrid_alpha, hse_omega); for (int ispin = 0; ispin < nspin_out; ispin++) { @@ -360,11 +396,11 @@ void write_h_vxc(WriteHParams& params) ModuleGint::cal_gint_vl(&v_xc(ispin, 0), &hR_tmp); - write_hk_common(hR_tmp, "vxck", ucell, pv, kv, nspin, istep, append, iat2iwt, nat); + write_hk_common(hR_tmp, "vxck", ucell, pv, kv, nspin, istep, append, iat2iwt, nat, files, parallel, basis, solver); if (also_hR) { - gather_and_write("vxc", "V^XC", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat); + gather_and_write("vxc", "V^XC", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat, run, files, parallel); } } @@ -383,11 +419,17 @@ static void write_h_exx_impl(const UnitCell& ucell, const bool append, const int* iat2iwt, const int nat, - const bool also_hR) + const bool also_hR, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::ExactExchangeState& exact_exchange_state) { const auto& Hexxs = ex->get_Hexxs(); // vector over spin of map> const int nspin_out = (nspin == 2 ? 2 : 1); - const double alpha = GlobalC::exx_info.info_global.hybrid_alpha; + const double alpha = exact_exchange_state.hybrid_alpha; for (int ispin = 0; ispin < nspin_out; ispin++) { @@ -395,47 +437,84 @@ static void write_h_exx_impl(const UnitCell& ucell, // add_HexxR only fills existing matrices, so first allocate the atom-pair structure // from the exx-form data (native cells, consistent with the nullptr cell_nearest below). hamilt::reallocate_hcontainer(Hexxs, &hR_tmp); - RI_2D_Comm::add_HexxR(ispin, alpha, Hexxs, pv, PARAM.globalv.npol, hR_tmp, nullptr); + RI_2D_Comm::add_HexxR(ispin, alpha, Hexxs, pv, basis.npol, hR_tmp, nullptr); - write_hk_common(hR_tmp, "vexxk", ucell, pv, kv, nspin, istep, append, iat2iwt, nat); + write_hk_common(hR_tmp, "vexxk", ucell, pv, kv, nspin, istep, append, iat2iwt, nat, files, parallel, basis, solver); if (also_hR) { - gather_and_write("vexx", "V^EXX", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat); + gather_and_write("vexx", "V^EXX", hR_tmp, ucell, pv, nspin, ispin, istep, append, iat2iwt, nat, run, files, parallel); } } } -void write_h_exx(WriteHParams& params) +void write_h_exx(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output, + const ModuleContext::ExactExchangeState& exact_exchange_state) { ModuleBase::TITLE("ModuleIO", "write_h_exx"); ModuleBase::timer::start("ModuleIO", "write_h_exx"); - // Multi-k out_mat_h_exx is rejected upstream at the call site (setup_exx_h_params in - // ctrl_scf_lcao.cpp); this function is only reached on the gamma-only path. + // Multi-k out_mat_h_exx is rejected upstream at the call site; this + // function is only reached on the gamma-only path. const UnitCell& ucell = *params.ucell; const Parallel_Orbitals& pv = *params.pv; const K_Vectors& kv = *params.kv; const int nspin = params.nspin; const int istep = params.istep; - const bool append = params.append; + const bool append = output.append; const int* iat2iwt = params.iat2iwt; const int nat = params.nat; const bool also_hR = params.also_hR; // exd (real Hexx) and exc (complex Hexx) are mutually exclusive; pick by real_number. - if (GlobalC::exx_info.info_ri.real_number) + if (exact_exchange_state.real_number) { if (params.exd != nullptr) { - write_h_exx_impl(ucell, pv, params.exd, kv, nspin, istep, append, iat2iwt, nat, also_hR); + write_h_exx_impl(ucell, + pv, + params.exd, + kv, + nspin, + istep, + append, + iat2iwt, + nat, + also_hR, + run, + files, + parallel, + basis, + solver, + exact_exchange_state); } } else { if (params.exc != nullptr) { - write_h_exx_impl(ucell, pv, params.exc, kv, nspin, istep, append, iat2iwt, nat, also_hR); + write_h_exx_impl(ucell, + pv, + params.exc, + kv, + nspin, + istep, + append, + iat2iwt, + nat, + also_hR, + run, + files, + parallel, + basis, + solver, + exact_exchange_state); } } diff --git a/source/source_io/module_hs/write_H_terms.h b/source/source_io/module_hs/write_H_terms.h index 7dbede427a8..2dd698aaad6 100644 --- a/source/source_io/module_hs/write_H_terms.h +++ b/source/source_io/module_hs/write_H_terms.h @@ -9,6 +9,7 @@ #include "source_estate/module_pot/potential_new.h" #include "source_lcao/LCAO_domain.h" #include "source_lcao/module_hcontainer/hcontainer.h" +#include "source_context/context_types.h" #include #include @@ -38,31 +39,64 @@ struct WriteHParams int nat = 0; bool also_hR = false; // H(k) is always written; H(R) (CSR) only when this is true #ifdef __EXX - // The gamma-only (TK==double) exx interfaces used by write_h_exx. - // Deliberately NOT templated on TK, because it would force WriteHParams, WriteDHParams and - // every free function taking them to become templates as well -- a large, purely - // mechanical change for a case nobody needs. - // Multi-k + EXX is therefore rejected up front (see write_h_exx) - // instead of silently producing output with the EXX term missing. + // Gamma EXX interfaces; the explicit ExactExchangeState selects the active representation. Exx_LRI_Interface* exd = nullptr; Exx_LRI_Interface>* exc = nullptr; #endif }; -void write_h_t(WriteHParams& params); +void write_h_t(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output); -void write_h_vnl(WriteHParams& params); +void write_h_vnl(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output); -void write_h_vl(WriteHParams& params); +void write_h_vl(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output); -void write_h_vh(WriteHParams& params); +void write_h_vh(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output); -void write_h_vxc(WriteHParams& params); +void write_h_vxc(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SpinConfig& spin, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output); #ifdef __EXX // Build V^EXX(R) into a real HContainer via add_HexxR (from exd/exc->get_Hexxs()) and write it. // exd (real Hexx) and exc (complex Hexx) are mutually exclusive; picked by info_ri.real_number. -void write_h_exx(WriteHParams& params); +void write_h_exx(WriteHParams& params, + const ModuleContext::RunControl& run, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output, + const ModuleContext::ExactExchangeState& exact_exchange_state); #endif } // namespace ModuleIO diff --git a/source/source_io/module_hs/write_vxc.hpp b/source/source_io/module_hs/write_vxc.hpp index d4719a17874..c289c517a33 100644 --- a/source/source_io/module_hs/write_vxc.hpp +++ b/source/source_io/module_hs/write_vxc.hpp @@ -1,6 +1,5 @@ #ifndef __WRITE_VXC_H_ #define __WRITE_VXC_H_ -#include "source_io/module_parameter/parameter.h" #include "source_base/parallel_reduce.h" #include "source_base/module_container/base/third_party/blas.h" #include "source_base/module_external/scalapack_connector.h" @@ -113,12 +112,13 @@ std::vector orbital_energy(const int ik, const int nbands, const std::ve inline void write_orb_energy(const K_Vectors& kv, const int nspin0, const int nbands, const std::vector>& e_orb, + const ModuleContext::FileSystemLayout& files, const std::string& term, const std::string& label, const bool app = false) { assert(e_orb.size() == kv.get_nks()); const int nk = kv.get_nks() / nspin0; std::ofstream ofs; - ofs.open(PARAM.globalv.global_out_dir + term + "_" + (label == "" ? "out.dat" : label + "_out.dat"), + ofs.open(files.output_directory + term + "_" + (label == "" ? "out.dat" : label + "_out.dat"), app ? std::ios::app : std::ios::out); ofs << nk << "\n" << nspin0 << "\n" << nbands << "\n"; ofs << std::scientific << std::setprecision(16); @@ -139,7 +139,6 @@ inline void write_orb_energy(const K_Vectors& kv, template void write_Vxc(const int nspin, const int nbasis, - const int drank, const Parallel_Orbitals* pv, const psi::Psi& psi, const UnitCell& ucell, @@ -153,6 +152,12 @@ void write_Vxc(const int nspin, const std::vector& orb_cutoff, const ModuleBase::matrix& wg, Grid_Driver& gd, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output, + const ModuleContext::DftUConfig& dftu, bool cal_exx #ifdef __EXX , @@ -235,7 +240,7 @@ void write_Vxc(const int nspin, e_orb_exx.emplace_back(orbital_energy(ik, nbands, vexx_k_mo, p2d)); } #endif - if (PARAM.inp.dft_plus_u) + if (dftu.enabled) { vdftu_op_ao.contributeHk(ik); } @@ -248,10 +253,10 @@ void write_Vxc(const int nspin, const int istep = -1; const int out_label = 1; // 1 means .txt while 2 means .dat const bool out_app_flag = 0; - const bool gamma_only = PARAM.globalv.gamma_only_local; + const bool gamma_only = basis.gamma_only_local; std::string vxc_file = ModuleIO::filename_output( - PARAM.globalv.global_out_dir, + files.output_directory, "vxc","nao",ik,kv.ik2iktot,nspin,kv.get_nkstot(), out_label,out_app_flag,gamma_only,istep); @@ -259,12 +264,13 @@ void write_Vxc(const int nspin, vxc_tot_k_mo.data(), nbands, false /*binary*/, - PARAM.inp.out_ndigits, + output.digits, true /*triangle*/, out_app_flag /*append*/, vxc_file, p2d, - drank); + parallel.diagonalization_rank, + solver); // ======test======= // total_energy += all_band_energy(ik, vxc_tot_k_mo, p2d, wg); // ======test======= @@ -283,14 +289,14 @@ void write_Vxc(const int nspin, delete vxcs_op_ao[is]; } - if (GlobalV::MY_RANK == 0) + if (parallel.world_rank == 0) { - write_orb_energy(kv, nspin0, nbands, e_orb_tot, "vxc", ""); + write_orb_energy(kv, nspin0, nbands, e_orb_tot, files, "vxc", ""); #ifdef __EXX if (cal_exx) { - write_orb_energy(kv, nspin0, nbands, e_orb_locxc, "vxc", "local"); - write_orb_energy(kv, nspin0, nbands, e_orb_exx, "vxc", "exx"); + write_orb_energy(kv, nspin0, nbands, e_orb_locxc, files, "vxc", "local"); + write_orb_energy(kv, nspin0, nbands, e_orb_exx, files, "vxc", "exx"); } #endif } diff --git a/source/source_io/module_hs/write_vxc_lip.hpp b/source/source_io/module_hs/write_vxc_lip.hpp index 5909094a2f7..9faa43b5913 100644 --- a/source/source_io/module_hs/write_vxc_lip.hpp +++ b/source/source_io/module_hs/write_vxc_lip.hpp @@ -1,6 +1,5 @@ #ifndef __WRITE_VXC_LIP_H_ #define __WRITE_VXC_LIP_H_ -#include "source_io/module_parameter/parameter.h" #include "source_base/parallel_reduce.h" #include "source_base/module_container/base/third_party/blas.h" #include "source_pw/module_pwdft/op_pw_veff.h" @@ -103,7 +102,6 @@ namespace ModuleIO template void write_Vxc(int nspin, int naos, - int drank, const psi::Psi>& psi_pw, const UnitCell& ucell, Structure_Factor& sf, @@ -115,6 +113,11 @@ namespace ModuleIO const Charge& chg, const K_Vectors& kv, const ModuleBase::matrix& wg, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, + const ModuleContext::BasisInfo& basis, + const ModuleContext::SolverConfig& solver, + const ModuleContext::MatrixOutputConfig& output, bool cal_exx, double hybrid_alpha #ifdef __EXX @@ -209,17 +212,17 @@ namespace ModuleIO const int istep = -1; const int out_label = 1; // 1 means .txt while 2 means .dat const bool out_app_flag = 0; - const bool gamma_only = PARAM.globalv.gamma_only_local; + const bool gamma_only = basis.gamma_only_local; std::string vxc_file = ModuleIO::filename_output( - PARAM.globalv.global_out_dir, + files.output_directory, "vxc","nao",ik,kv.ik2iktot,nspin,kv.get_nkstot(), out_label,out_app_flag,gamma_only,istep); ModuleIO::save_mat(istep, vxc_tot_k_mo.data(), nbands, - false, PARAM.inp.out_ndigits, true, + false, output.digits, true, out_app_flag, vxc_file, - p2d_serial, drank, false); + p2d_serial, parallel.diagonalization_rank, solver, false); e_orb_tot.emplace_back(orbital_energy(ik, nbands, vxc_tot_k_mo)); } @@ -235,20 +238,20 @@ namespace ModuleIO // for (int ir = 0;ir < potxc->get_veff_smooth().nc;++ir) // exc_by_rho += potxc->get_veff_smooth()(0, ir) * chg.rho[0][ir]; // Parallel_Reduce::reduce_all(exc_by_rho); - // exc_by_rho *= ((FPTYPE)ucell.omega * (FPTYPE)GlobalV::NPROC / (FPTYPE)potxc->get_veff_smooth().nc); + // The grid integral is normalized by the process count in its caller. // std::cout << "xc all-bands energy by rho =" << exc_by_rho << std::endl; //===== test total xc energy ======= //===== test total exx energy ======= //===== test total exx energy ======= // write the orbital energy for xc and exx in LibRPA format const int nspin0 = (nspin == 2) ? 2 : 1; - auto write_orb_energy = [&kv, &nspin0, &nbands](const std::vector>& e_orb, + auto write_orb_energy = [&kv, &nspin0, &nbands, &files](const std::vector>& e_orb, const std::string& label, const bool app = false) { assert(e_orb.size() == kv.get_nks()); const int nk = kv.get_nks() / nspin0; std::ofstream ofs; - ofs.open(PARAM.globalv.global_out_dir + "vxc_" + (label == "" ? "out.dat" : label + "_out.dat"), + ofs.open(files.output_directory + "vxc_" + (label == "" ? "out.dat" : label + "_out.dat"), app ? std::ios::app : std::ios::out); ofs << nk << "\n" << nspin0 << "\n" << nbands << "\n"; ofs << std::scientific << std::setprecision(16); @@ -264,7 +267,7 @@ namespace ModuleIO } }; - if (GlobalV::MY_RANK == 0) + if (parallel.world_rank == 0) { write_orb_energy(e_orb_tot, ""); #if((defined __LCAO)&&(defined __EXX) && !(defined __CUDA)&& !(defined __ROCM)) diff --git a/source/source_io/module_hs/write_vxc_r.hpp b/source/source_io/module_hs/write_vxc_r.hpp index 81d4b38a102..624d554e3a5 100644 --- a/source/source_io/module_hs/write_vxc_r.hpp +++ b/source/source_io/module_hs/write_vxc_r.hpp @@ -1,6 +1,5 @@ #ifndef __WRITE_VXC_R_H_ #define __WRITE_VXC_R_H_ -#include "source_io/module_parameter/parameter.h" #include "source_io/module_hs/write_HS_sparse.h" #include "source_lcao/module_operator_lcao/op_dftu_lcao.h" #include "source_lcao/module_operator_lcao/veff_lcao.h" @@ -36,6 +35,8 @@ void write_Vxc_R(const int nspin, const K_Vectors& kv, const std::vector& orb_cutoff, Grid_Driver& gd, + const ModuleContext::FileSystemLayout& files, + const ModuleContext::ParallelTopology& parallel, bool cal_exx, double hybrid_alpha, bool real_number @@ -143,17 +144,18 @@ void write_Vxc_R(const int nspin, std::set> all_R_coor = sparse_format::get_R_range(vxcs_R_ao[is]); const std::string filename = "Vxc_R_spin" + std::to_string(is); ModuleIO::SparseWriteOptions options; - options.filename = PARAM.globalv.global_out_dir + filename + ".csr"; + options.filename = files.output_directory + filename + ".csr"; options.label = filename; options.threshold = sparse_thr; options.binary = false; options.istep = -1; options.reduce = true; - options.temp_dir = PARAM.globalv.global_out_dir; + options.temp_dir = files.output_directory; ModuleIO::save_sparse(cal_HR_sparse(vxcs_R_ao[is], sparse_thr), all_R_coor, *pv, - options); + options, + parallel); } } } // namespace ModuleIO diff --git a/source/source_io/module_restart/restart_exx_csr.h b/source/source_io/module_restart/restart_exx_csr.h index 6980e9ab6a6..e9cdc94321d 100644 --- a/source/source_io/module_restart/restart_exx_csr.h +++ b/source/source_io/module_restart/restart_exx_csr.h @@ -6,6 +6,11 @@ #include #include +namespace ModuleContext +{ +struct ParallelTopology; +} + namespace ModuleIO { using TC = std::array; @@ -28,15 +33,16 @@ void read_Hexxs_cereal(const std::string& file_name, template void write_Hexxs_csr(const std::string& file_name, const UnitCell& ucell, - const std::map>>& Hexxs); + const std::vector>>>& Hexxs, + const ModuleContext::ParallelTopology& parallel); /// calculate CSR sparse matrix from the global matrix stored with RI::Tensor /// the return type is same as SR_sparse, HR_sparse, etc. template std::map, std::map>> calculate_RI_Tensor_sparse( const double& sparse_threshold, - const std::vector>>>& Hexxs, + const std::map>>& Hexxs, const UnitCell& ucell); } // namespace ModuleIO -#include "restart_exx_csr.hpp" \ No newline at end of file +#include "restart_exx_csr.hpp" diff --git a/source/source_io/module_restart/restart_exx_csr.hpp b/source/source_io/module_restart/restart_exx_csr.hpp index 568ed952fcb..d5853359fc1 100644 --- a/source/source_io/module_restart/restart_exx_csr.hpp +++ b/source/source_io/module_restart/restart_exx_csr.hpp @@ -112,7 +112,8 @@ namespace ModuleIO template void write_Hexxs_csr(const std::string& file_name, const UnitCell& ucell, - const std::vector>>>& Hexxs) + const std::vector>>>& Hexxs, + const ModuleContext::ParallelTopology& parallel) { ModuleBase::TITLE("ModuleIO", "write_Hexxs_csr"); std::set> all_R_coor; @@ -157,7 +158,8 @@ namespace ModuleIO calculate_RI_Tensor_sparse(sparse_threshold, Hexxs[is], ucell), all_R_coor, pv, - options); + options, + parallel); } } } diff --git a/source/source_io/test/single_R_io_test.cpp b/source/source_io/test/single_R_io_test.cpp index fcadfaff91d..8315b788252 100644 --- a/source/source_io/test/single_R_io_test.cpp +++ b/source/source_io/test/single_R_io_test.cpp @@ -1,10 +1,6 @@ #include "gtest/gtest.h" #include "gmock/gmock.h" -#define private public -#include "source_io/module_parameter/parameter.h" -#undef private #include "source_io/module_hs/single_R_io.h" -#include "source_base/global_variable.h" #include "source_basis/module_ao/parallel_orbitals.h" #include #include @@ -49,15 +45,15 @@ TEST(ModuleIOTest, OutputSingleR) { // Create temporary output file std::stringstream ofs_filename; - GlobalV::DRANK=0; - ofs_filename << "test_output_single_R_" << GlobalV::DRANK << ".dat"; + ModuleContext::ParallelTopology parallel; + parallel.diagonalization_rank = 0; + ofs_filename << "test_output_single_R_0.dat"; std::ofstream ofs(ofs_filename.str()); // Define input parameters const double sparse_threshold = 1e-8; const bool binary = false; Parallel_Orbitals pv; - PARAM.sys.nlocal = 99; pv.set_serial(5, 5); std::map> XR = { {0, {{1, 0.5}, {3, 0.3}}}, @@ -71,7 +67,7 @@ TEST(ModuleIOTest, OutputSingleR) options.temp_dir = "./"; // Call function under test - ModuleIO::output_single_R(ofs, XR, pv, options); + ModuleIO::output_single_R(ofs, XR, pv, options, parallel); // Close output file and open it for reading ofs.close(); @@ -125,7 +121,8 @@ TEST(ModuleIOTest, OutputSingleRComplexKeepsHighPrecision) { const std::string filename = "test_output_single_R_complex.dat"; std::remove(filename.c_str()); - GlobalV::DRANK = 0; + ModuleContext::ParallelTopology parallel; + parallel.diagonalization_rank = 0; std::ofstream ofs(filename); Parallel_Orbitals pv; @@ -139,7 +136,7 @@ TEST(ModuleIOTest, OutputSingleRComplexKeepsHighPrecision) options.reduce = false; options.temp_dir = "./"; - ModuleIO::output_single_R(ofs, XR, pv, options); + ModuleIO::output_single_R(ofs, XR, pv, options, parallel); ofs.close(); std::ifstream ifs(filename); @@ -154,7 +151,8 @@ TEST(ModuleIOTest, OutputSingleRUsesConfiguredPrecision) { const std::string filename = "test_output_single_R_precision.dat"; std::remove(filename.c_str()); - GlobalV::DRANK = 0; + ModuleContext::ParallelTopology parallel; + parallel.diagonalization_rank = 0; std::ofstream ofs(filename); Parallel_Orbitals pv; @@ -169,7 +167,7 @@ TEST(ModuleIOTest, OutputSingleRUsesConfiguredPrecision) options.reduce = false; options.temp_dir = "./"; - ModuleIO::output_single_R(ofs, XR, pv, options); + ModuleIO::output_single_R(ofs, XR, pv, options, parallel); ofs.close(); std::ifstream ifs(filename); @@ -180,9 +178,36 @@ TEST(ModuleIOTest, OutputSingleRUsesConfiguredPrecision) std::remove(filename.c_str()); } +TEST(ModuleIOTest, OutputSingleRDoesNotWriteOnNonRootDuringReduction) +{ + const std::string filename = "test_output_single_R_nonroot.dat"; + std::remove(filename.c_str()); + ModuleContext::ParallelTopology parallel; + parallel.diagonalization_rank = 1; + std::ofstream ofs(filename); + + Parallel_Orbitals pv; + pv.set_serial(5, 5); + ModuleIO::SparseRBlock XR = {{0, {{1, 2.0}}}}; + ModuleIO::SparseWriteOptions options; + options.threshold = 1e-12; + options.binary = false; + options.reduce = true; + options.temp_dir = "./"; + + ModuleIO::output_single_R(ofs, XR, pv, options, parallel); + ofs.close(); + + std::ifstream ifs(filename, std::ios::binary | std::ios::ate); + ASSERT_TRUE(ifs.is_open()); + EXPECT_EQ(0, ifs.tellg()); + std::remove(filename.c_str()); +} + void write_out_of_range_sparse_column(const char* filename) { - GlobalV::DRANK = 0; + ModuleContext::ParallelTopology parallel; + parallel.diagonalization_rank = 0; std::ofstream ofs(filename); Parallel_Orbitals pv; pv.set_serial(5, 5); @@ -193,7 +218,7 @@ void write_out_of_range_sparse_column(const char* filename) options.binary = false; options.reduce = false; options.temp_dir = "/tmp/"; - ModuleIO::output_single_R(ofs, XR, pv, options); + ModuleIO::output_single_R(ofs, XR, pv, options, parallel); } TEST(ModuleIOTest, OutputSingleRRejectsOutOfRangeColumn) @@ -212,8 +237,6 @@ int main(int argc, char **argv) #ifdef __MPI MPI_Init(&argc, &argv); - MPI_Comm_size(MPI_COMM_WORLD,&GlobalV::NPROC); - MPI_Comm_rank(MPI_COMM_WORLD,&GlobalV::MY_RANK); #endif testing::InitGoogleTest(&argc, argv); diff --git a/source/source_io/test/write_hs_r_compat_test.cpp b/source/source_io/test/write_hs_r_compat_test.cpp index 8c7c5821243..f31dbb9e161 100644 --- a/source/source_io/test/write_hs_r_compat_test.cpp +++ b/source/source_io/test/write_hs_r_compat_test.cpp @@ -1,11 +1,6 @@ #include "gmock/gmock.h" #include "gtest/gtest.h" -#define private public -#include "source_io/module_parameter/parameter.h" -#undef private - -#include "source_base/global_variable.h" #include "source_cell/module_neighbor/sltk_grid_driver.h" #include "source_io/module_dm/write_dmr.h" #include "source_io/module_hs/output_mat_sparse.h" @@ -190,15 +185,32 @@ void fill_matrix(hamilt::HContainer& matrix, Parallel_Orbitals& pv, doub matrix.insert_pair(pair); } -void init_sparse_output_globals(const int nspin = 1) +struct HsContextSlices { - GlobalV::DRANK = 0; - PARAM.input.nspin = nspin; - PARAM.input.calculation = "scf"; - PARAM.input.out_app_flag = false; - PARAM.sys.global_out_dir = "./"; - PARAM.sys.global_matrix_dir = "./"; - PARAM.sys.nlocal = 2; + ModuleContext::RunControl run; + ModuleContext::FileSystemLayout files; + ModuleContext::ParallelTopology parallel; + ModuleContext::LogStreams logs; + ModuleContext::BasisInfo basis; + ModuleContext::SpinConfig spin; + ModuleContext::MatrixOutputConfig output; +}; + +std::ostringstream sparse_log; + +HsContextSlices make_hs_context(const int nspin = 1) +{ + HsContextSlices context; + context.run.calculation = "scf"; + context.files.output_directory = "./"; + context.files.matrix_directory = "./"; + context.parallel.diagonalization_rank = 0; + context.logs.running = &sparse_log; + context.basis.nlocal = 2; + context.basis.npol = 1; + context.spin.nspin = nspin; + context.output.append = false; + return context; } void remove_derivative_files(const std::string& fileflag) @@ -330,10 +342,6 @@ TEST(WriteHsRCompatibility, LegacySparseHeaderKeepsStepStyle) const std::string filename = "write_hs_r_legacy_s.csr"; std::remove(filename.c_str()); - GlobalV::DRANK = 0; - PARAM.sys.global_out_dir = "./"; - PARAM.sys.nlocal = 99; - Parallel_Orbitals pv; init_serial_orbitals(pv); const Abfs::Vector3_Order r_vector(0, 0, 0); @@ -351,7 +359,8 @@ TEST(WriteHsRCompatibility, LegacySparseHeaderKeepsStepStyle) options.istep = 0; options.reduce = false; options.temp_dir = "./"; - ModuleIO::save_sparse(sparse_matrix, all_R_coor, pv, options); + const ModuleContext::ParallelTopology parallel; + ModuleIO::save_sparse(sparse_matrix, all_R_coor, pv, options, parallel); const std::string output = read_file(filename); EXPECT_TRUE(starts_with(output, "STEP: 0\n")); @@ -367,9 +376,6 @@ TEST(WriteHsRCompatibility, LegacySparseTextCountsOnlyValuesAboveThreshold) const std::string filename = "write_hs_r_threshold_s.csr"; std::remove(filename.c_str()); - GlobalV::DRANK = 0; - PARAM.sys.global_out_dir = "./"; - Parallel_Orbitals pv; init_serial_orbitals(pv); const Abfs::Vector3_Order r_vector(0, 0, 0); @@ -389,7 +395,8 @@ TEST(WriteHsRCompatibility, LegacySparseTextCountsOnlyValuesAboveThreshold) options.istep = 2; options.reduce = false; options.temp_dir = "./"; - ModuleIO::save_sparse(sparse_matrix, all_R_coor, pv, options); + const ModuleContext::ParallelTopology parallel; + ModuleIO::save_sparse(sparse_matrix, all_R_coor, pv, options, parallel); const std::vector lines = read_lines(filename); ASSERT_GE(lines.size(), 6); @@ -434,10 +441,6 @@ TEST(WriteHsRCompatibility, LegacySparseBinaryHeaderWritesConcreteStep) const std::string filename = "write_hs_r_legacy_binary_s.csr"; std::remove(filename.c_str()); - GlobalV::DRANK = 0; - PARAM.sys.global_out_dir = "./"; - PARAM.sys.nlocal = 99; - Parallel_Orbitals pv; init_serial_orbitals(pv); const Abfs::Vector3_Order r_vector(0, 0, 0); @@ -455,7 +458,8 @@ TEST(WriteHsRCompatibility, LegacySparseBinaryHeaderWritesConcreteStep) options.istep = 3; options.reduce = false; options.temp_dir = "./"; - ModuleIO::save_sparse(sparse_matrix, all_R_coor, pv, options); + const ModuleContext::ParallelTopology parallel; + ModuleIO::save_sparse(sparse_matrix, all_R_coor, pv, options, parallel); const std::vector header_and_r = read_binary_ints(filename, 7); EXPECT_THAT(header_and_r, testing::ElementsAre(3, 2, 1, 0, 0, 0, 2)); @@ -468,9 +472,6 @@ TEST(WriteHsRCompatibility, LegacySparseBinaryCountsOnlyValuesAboveThreshold) const std::string filename = "write_hs_r_threshold_binary_s.csr"; std::remove(filename.c_str()); - GlobalV::DRANK = 0; - PARAM.sys.global_out_dir = "./"; - Parallel_Orbitals pv; init_serial_orbitals(pv); const Abfs::Vector3_Order r_vector(0, 0, 0); @@ -490,7 +491,8 @@ TEST(WriteHsRCompatibility, LegacySparseBinaryCountsOnlyValuesAboveThreshold) options.istep = 4; options.reduce = false; options.temp_dir = "./"; - ModuleIO::save_sparse(sparse_matrix, all_R_coor, pv, options); + const ModuleContext::ParallelTopology parallel; + ModuleIO::save_sparse(sparse_matrix, all_R_coor, pv, options, parallel); std::ifstream ifs(filename.c_str(), std::ios::binary); ASSERT_TRUE(ifs.is_open()); @@ -515,7 +517,7 @@ TEST(WriteHsRCompatibility, LegacySparseBinaryCountsOnlyValuesAboveThreshold) TEST(WriteHsRCompatibility, SaveDHSparseTextCountsOnlyValuesAboveThreshold) { remove_derivative_files("h"); - init_sparse_output_globals(); + const HsContextSlices context = make_hs_context(); Parallel_Orbitals pv; init_serial_orbitals(pv); @@ -527,7 +529,9 @@ TEST(WriteHsRCompatibility, SaveDHSparseTextCountsOnlyValuesAboveThreshold) arrays.dHRx_sparse[0][r_vector][1][0] = 0.0; arrays.dHRx_sparse[0][r_vector][1][1] = -2.0; - ModuleIO::save_dH_sparse(5, pv, arrays, 1e-10, false, "h", 8); + ModuleIO::save_dH_sparse(5, pv, arrays, 1e-10, false, "h", 8, + context.run, context.files, context.parallel, context.logs, + context.basis, context.spin, context.output); const std::vector lines = read_lines("dhrxs1_nao.csr"); ASSERT_GE(lines.size(), 7); @@ -566,7 +570,7 @@ TEST(WriteHsRCompatibility, SaveDHSparseTextCountsOnlyValuesAboveThreshold) TEST(WriteHsRCompatibility, SaveDHSparseBinaryCountsOnlyValuesAboveThreshold) { remove_derivative_files("h"); - init_sparse_output_globals(); + const HsContextSlices context = make_hs_context(); Parallel_Orbitals pv; init_serial_orbitals(pv); @@ -578,7 +582,9 @@ TEST(WriteHsRCompatibility, SaveDHSparseBinaryCountsOnlyValuesAboveThreshold) arrays.dHRx_sparse[0][r_vector][1][0] = 0.0; arrays.dHRx_sparse[0][r_vector][1][1] = -2.0; - ModuleIO::save_dH_sparse(6, pv, arrays, 1e-10, true, "h", 8); + ModuleIO::save_dH_sparse(6, pv, arrays, 1e-10, true, "h", 8, + context.run, context.files, context.parallel, context.logs, + context.basis, context.spin, context.output); std::ifstream ifs("dhrxs1_nao.csr", std::ios::binary); ASSERT_TRUE(ifs.is_open()); @@ -600,10 +606,64 @@ TEST(WriteHsRCompatibility, SaveDHSparseBinaryCountsOnlyValuesAboveThreshold) remove_derivative_files("h"); } +TEST(WriteHsRCompatibility, SaveDHSparseSpinTwoWritesBothSpinChannels) +{ + remove_derivative_files("h"); + const HsContextSlices context = make_hs_context(2); + + Parallel_Orbitals pv; + init_serial_orbitals(pv); + LCAO_HS_Arrays arrays; + const Abfs::Vector3_Order r_vector(0, 0, 0); + arrays.all_R_coor.insert(r_vector); + arrays.dHRx_sparse[0][r_vector][0][0] = 1.25; + arrays.dHRx_sparse[1][r_vector][1][1] = -2.5; + + ModuleIO::save_dH_sparse(8, pv, arrays, 1e-10, false, "h", 8, + context.run, context.files, context.parallel, context.logs, + context.basis, context.spin, context.output); + + const std::string spin1 = read_file("dhrxs1_nao.csr"); + const std::string spin2 = read_file("dhrxs2_nao.csr"); + EXPECT_THAT(spin1, testing::HasSubstr("Matrix number of dHx(R): 1\n0 0 0 1\n")); + EXPECT_THAT(spin2, testing::HasSubstr("Matrix number of dHx(R): 1\n0 0 0 1\n")); + EXPECT_THAT(spin1, testing::HasSubstr("1.25000000e+00")); + EXPECT_THAT(spin2, testing::HasSubstr("-2.50000000e+00")); + + remove_derivative_files("h"); +} + +TEST(WriteHsRCompatibility, SaveDHSparseMdNonAppendUsesStepFileName) +{ + const std::string filename = "dhrxs1g9_nao.csr"; + std::remove(filename.c_str()); + HsContextSlices context = make_hs_context(1); + context.run.calculation = "md"; + context.output.append = false; + + Parallel_Orbitals pv; + init_serial_orbitals(pv); + LCAO_HS_Arrays arrays; + const Abfs::Vector3_Order r_vector(0, 0, 0); + arrays.all_R_coor.insert(r_vector); + arrays.dHRx_sparse[0][r_vector][0][0] = 3.0; + + ModuleIO::save_dH_sparse(9, pv, arrays, 1e-10, false, "h", 8, + context.run, context.files, context.parallel, context.logs, + context.basis, context.spin, context.output); + + const std::string output = read_file(filename); + EXPECT_THAT(output, testing::HasSubstr("STEP: 9")); + EXPECT_THAT(output, testing::HasSubstr("3.00000000e+00")); + std::remove(filename.c_str()); + std::remove("dhrys1g9_nao.csr"); + std::remove("dhrzs1g9_nao.csr"); +} + TEST(WriteHsRCompatibility, SaveDSSparseSocWritesAllDirections) { remove_derivative_files("s"); - init_sparse_output_globals(4); + const HsContextSlices context = make_hs_context(4); Parallel_Orbitals pv; init_serial_orbitals(pv); @@ -614,7 +674,9 @@ TEST(WriteHsRCompatibility, SaveDSSparseSocWritesAllDirections) arrays.dHRy_soc_sparse[r_vector][0][1] = std::complex(2.0, -1.0); arrays.dHRz_soc_sparse[r_vector][1][1] = std::complex(-3.0, 0.5); - ModuleIO::save_dH_sparse(7, pv, arrays, 1e-10, false, "s", 8); + ModuleIO::save_dH_sparse(7, pv, arrays, 1e-10, false, "s", 8, + context.run, context.files, context.parallel, context.logs, + context.basis, context.spin, context.output); const std::string x_output = read_file("dsrxs1_nao.csr"); const std::string y_output = read_file("dsrys1_nao.csr"); @@ -796,8 +858,6 @@ int main(int argc, char** argv) { #ifdef __MPI MPI_Init(&argc, &argv); - MPI_Comm_size(MPI_COMM_WORLD, &GlobalV::NPROC); - MPI_Comm_rank(MPI_COMM_WORLD, &GlobalV::MY_RANK); #endif ::testing::InitGoogleTest(&argc, argv); diff --git a/source/source_lcao/module_lr/CMakeLists.txt b/source/source_lcao/module_lr/CMakeLists.txt index 331e7dc22e8..3a4b7c70c9c 100644 --- a/source/source_lcao/module_lr/CMakeLists.txt +++ b/source/source_lcao/module_lr/CMakeLists.txt @@ -26,4 +26,12 @@ add_library( ${objects} ) -endif() \ No newline at end of file +# The LR ESolver is an L1 orchestration exception inside this mixed directory. +# Limit root Context access to that translation unit; LR algorithms only see +# explicitly forwarded small context structures. +set_source_files_properties( + esolver_lrtd_lcao.cpp + PROPERTIES COMPILE_DEFINITIONS ABACUS_CAN_READ_SIMULATION_CONTEXT +) + +endif() diff --git a/source/source_lcao/module_lr/esolver_lrtd_lcao.cpp b/source/source_lcao/module_lr/esolver_lrtd_lcao.cpp index 1b157e00dac..d97cf421060 100644 --- a/source/source_lcao/module_lr/esolver_lrtd_lcao.cpp +++ b/source/source_lcao/module_lr/esolver_lrtd_lcao.cpp @@ -18,6 +18,7 @@ #include "source_io/module_parameter/parameter.h" #include "source_lcao/module_lr/ri_benchmark/ri_benchmark.h" #include "source_lcao/module_lr/operator_casida/operator_lr_diag.h" // for precondition +#include "source_context/orchestration_context.h" #ifdef __EXX #include "source_lcao/module_ri/Exx_LRI_interface.h" #endif @@ -545,6 +546,7 @@ void LR::ESolver_LR::after_all_runners(UnitCell& ucell) { ModuleBase::TITLE("ESolver_LR", "after_all_runners"); if (input.ri_hartree_benchmark != "none") { return; } //no need to calculate the spectrum in the benchmark routine + const ModuleContext::SimulationContext& context = ModuleContext::current_simulation_context(); //cal spectrum std::vector freq(100); std::vector abs_wavelen_range({ 20, 200 });//default range @@ -562,6 +564,7 @@ void LR::ESolver_LR::after_all_runners(UnitCell& ucell) this->ucell, this->kv, this->gd, this->orb_cutoff_, this->two_center_bundle_, this->paraX_, this->paraC_, this->paraMat_, &this->pelec->ekb.c[is * nstates], this->X[is].template data(), nstates, openshell, + context.run, context.basis, LR_Util::tolower(input.abs_gauge)); spectrum.transition_analysis(spin_types[is]); if (spin_types[is] != "triplet") // triplets has no transition dipole and no contribution to the spectrum diff --git a/source/source_lcao/module_lr/lr_spectrum.h b/source/source_lcao/module_lr/lr_spectrum.h index 5ed4678dc9c..67d29ec1c9c 100644 --- a/source/source_lcao/module_lr/lr_spectrum.h +++ b/source/source_lcao/module_lr/lr_spectrum.h @@ -5,6 +5,7 @@ #include "source_lcao/module_lr/utils/lr_util.h" #include "source_basis/module_nao/two_center_bundle.h" #include "source_lcao/module_rt/velocity_op.h" +#include "source_context/context_types.h" namespace LR { template @@ -17,11 +18,14 @@ namespace LR const TwoCenterBundle& two_center_bundle_, const std::vector& pX_in, const Parallel_2D& pc_in, const Parallel_Orbitals& pmat_in, const double* eig, const T* X, const int& nstate, const bool& openshell, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis, const std::string& gauge = "length") : nspin_x(openshell ? 2 : 1), naos(naos), nocc(nocc), nvirt(nvirt), nk(kv_in.get_nks() / nspin_global), rho_basis(rho_basis), ucell(ucell), kv(kv_in), gd_(gd), orb_cutoff_(orb_cutoff), two_center_bundle_(two_center_bundle_), pX(pX_in), pc(pc_in), pmat(pmat_in), + run_(run), basis_(basis), eig(eig), X(X), nstate(nstate), ldim(nk* (nspin_x == 2 ? pX_in[0].get_local_size() + pX_in[1].get_local_size() : pX_in[0].get_local_size())), gdim(nk* std::inner_product(nocc.begin(), nocc.end(), nvirt.begin(), 0)) @@ -79,6 +83,8 @@ namespace LR const UnitCell& ucell; const std::vector& orb_cutoff_; const TwoCenterBundle& two_center_bundle_; + const ModuleContext::RunControl& run_; + const ModuleContext::BasisInfo& basis_; void cal_gint_rho(double** rho, const int& nrxx); std::map get_pair_info(const int i); ///< given the index in X, return its ispin, ik, iocc, ivirt diff --git a/source/source_lcao/module_lr/lr_spectrum_velocity.cpp b/source/source_lcao/module_lr/lr_spectrum_velocity.cpp index 8297a14cbea..4e0aa0bd609 100644 --- a/source/source_lcao/module_lr/lr_spectrum_velocity.cpp +++ b/source/source_lcao/module_lr/lr_spectrum_velocity.cpp @@ -10,7 +10,9 @@ namespace LR inline Velocity_op> get_velocity_matrix_R(const UnitCell& ucell, const Grid_Driver& gd, const Parallel_Orbitals& pmat, - const TwoCenterBundle& two_center_bundle) + const TwoCenterBundle& two_center_bundle, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis) { // convert the orbital object to the old class for Velocity_op LCAO_Orbitals orb; @@ -18,7 +20,13 @@ namespace LR two_center_bundle.to_LCAO_Orbitals(orb, inp.lcao_ecut, inp.lcao_dk, inp.lcao_dr, inp.lcao_rmax, inp.out_element_info, inp.cal_force); // actually this class calculates the velocity matrix v(R) at A=0 - Velocity_op> vR(&ucell, &gd, &pmat, orb, two_center_bundle.overlap_orb.get()); + Velocity_op> vR(&ucell, + &gd, + &pmat, + orb, + two_center_bundle.overlap_orb.get(), + run, + basis); vR.calculate_vcomm_r(); // $<\mu, 0|[Vnl, r]|\nu, R>$ vR.calculate_grad_term(); // $<\mu, 0|\nabla|\nu, R>$ return vR; @@ -98,7 +106,8 @@ namespace LR template void LR::LR_Spectrum::cal_transition_dipoles_velocity() { - const Velocity_op>& vR = get_velocity_matrix_R(ucell, gd_, pmat, two_center_bundle_); // velocity matrix v(R) + const Velocity_op>& vR + = get_velocity_matrix_R(ucell, gd_, pmat, two_center_bundle_, run_, basis_); // velocity matrix v(R) transition_dipole_.resize(nstate); this->mean_squared_transition_dipole_.resize(nstate); for (int istate = 0;istate < nstate;++istate) @@ -149,7 +158,8 @@ namespace LR void LR::LR_Spectrum::test_transition_dipoles_velocity_ks(const double* const ks_eig) { // velocity matrix v(R) - const Velocity_op>& vR = get_velocity_matrix_R(ucell, gd_, pmat, two_center_bundle_); + const Velocity_op>& vR + = get_velocity_matrix_R(ucell, gd_, pmat, two_center_bundle_, run_, basis_); // (e_c-e_v) of KS eigenvalues std::vector eig_ks_diff(this->ldim); for (int is = 0;is < this->nspin_x;++is) @@ -184,4 +194,4 @@ namespace LR } } template class LR::LR_Spectrum; -template class LR::LR_Spectrum>; \ No newline at end of file +template class LR::LR_Spectrum>; diff --git a/source/source_lcao/module_ri/Exx_LRI_interface.h b/source/source_lcao/module_ri/Exx_LRI_interface.h index 810a3090771..bcd00d47385 100644 --- a/source/source_lcao/module_ri/Exx_LRI_interface.h +++ b/source/source_lcao/module_ri/Exx_LRI_interface.h @@ -8,6 +8,7 @@ #include "source_estate/module_dm/density_matrix.h" // mohan add 2025-11-04 #include "source_hamilt/hamilt.h" #include "source_hamilt/module_xc/exx_info_global.h" +#include "source_context/context_types.h" #include class LCAO_Matrix; @@ -121,7 +122,8 @@ class Exx_LRI_Interface const double& scf_ene_thr, int& iter, const int istep, - bool& conv_esolver); + bool& conv_esolver, + const ModuleContext::ParallelTopology& parallel); /// @brief: in do_after_converge: add exx operators; do DM mixing if seperate loop bool exx_after_converge(const UnitCell& ucell, hamilt::Hamilt& hamilt, diff --git a/source/source_lcao/module_ri/Exx_LRI_interface.hpp b/source/source_lcao/module_ri/Exx_LRI_interface.hpp index 479d4c1691f..9c42707db46 100644 --- a/source/source_lcao/module_ri/Exx_LRI_interface.hpp +++ b/source/source_lcao/module_ri/Exx_LRI_interface.hpp @@ -262,7 +262,8 @@ void Exx_LRI_Interface::exx_iter_finish(const K_Vectors& kv, const double& scf_ene_thr, int& iter, const int istep, - bool& conv_esolver) + bool& conv_esolver, + const ModuleContext::ParallelTopology& parallel) { ModuleBase::TITLE("Exx_LRI_Interface","exx_iter_finish"); if (GlobalC::restart.info_save.save_H && (this->two_level_step > 0 || istep > 0) @@ -287,7 +288,7 @@ void Exx_LRI_Interface::exx_iter_finish(const K_Vectors& kv, }*/ ////////// for Add_Hexx_Type:R const std::string& restart_HR_path = GlobalC::restart.folder + "HexxR" + std::to_string(GlobalV::MY_RANK); - ModuleIO::write_Hexxs_csr(restart_HR_path, ucell, this->get_Hexxs()); + ModuleIO::write_Hexxs_csr(restart_HR_path, ucell, this->get_Hexxs(), parallel); if (GlobalV::MY_RANK == 0) { diff --git a/source/source_lcao/module_rt/td_info.cpp b/source/source_lcao/module_rt/td_info.cpp index 916dc3b0810..99e6704c072 100644 --- a/source/source_lcao/module_rt/td_info.cpp +++ b/source/source_lcao/module_rt/td_info.cpp @@ -39,11 +39,20 @@ TD_info::TD_info(const UnitCell* ucell_in,const Parallel_Orbitals& pv, const LCA //std::cout<<"estep_shift"<istep += estep_shift; - if(out_current==2||elecstate::H_TDDFT_pw::stype == 2) + return; +} + +void TD_info::initialize_r_calculator(const UnitCell& ucell, + const Parallel_Orbitals& pv, + const LCAO_Orbitals& orb, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis) +{ + if (!r_calculator_initialized_ && (out_current == 2 || elecstate::H_TDDFT_pw::stype == 2)) { - r_calculator.init(*ucell_in, pv, orb); + r_calculator.init(ucell, pv, orb, run, basis); + r_calculator_initialized_ = true; } - return; } TD_info::~TD_info() { @@ -422,4 +431,4 @@ void TD_info::calculate_grad_overlap(const Parallel_Orbitals& paraV, template void TD_info::initialize_phase_hybrid>(const UnitCell& ucell, const hamilt::HContainer>* hR); template -void TD_info::initialize_phase_hybrid(const UnitCell& ucell, const hamilt::HContainer* hR); \ No newline at end of file +void TD_info::initialize_phase_hybrid(const UnitCell& ucell, const hamilt::HContainer* hR); diff --git a/source/source_lcao/module_rt/td_info.h b/source/source_lcao/module_rt/td_info.h index d14c67daf70..70ecdd2aa67 100644 --- a/source/source_lcao/module_rt/td_info.h +++ b/source/source_lcao/module_rt/td_info.h @@ -70,6 +70,11 @@ class TD_info const Grid_Driver& GridD, const std::vector& orb_cutoff, const TwoCenterIntegrator* intor); + void initialize_r_calculator(const UnitCell& ucell, + const Parallel_Orbitals& pv, + const LCAO_Orbitals& orb, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis); std::vector*> get_grad_overlap() const { return this->grad_overlap; @@ -128,6 +133,7 @@ class TD_info /// @brief store kinetic hamilton hamilt::HContainer>* velocity_HR = nullptr; + bool r_calculator_initialized_ = false; }; #endif diff --git a/source/source_lcao/module_rt/test/snap_psibeta_half_tddft_test.cpp b/source/source_lcao/module_rt/test/snap_psibeta_half_tddft_test.cpp index 47ed4f09878..29e2b14978b 100644 --- a/source/source_lcao/module_rt/test/snap_psibeta_half_tddft_test.cpp +++ b/source/source_lcao/module_rt/test/snap_psibeta_half_tddft_test.cpp @@ -134,7 +134,12 @@ class SnapPsibetaHalfTddftTest : public ::testing::Test void initialize_r_overlap_reference() { - r_calculator.init_nonlocal(ucell, pv, orb); + ModuleContext::RunControl run; + run.cal_force = true; + ModuleContext::BasisInfo basis; + basis.nlocal = pv.get_global_row_size(); + basis.npol = ucell.get_npol(); + r_calculator.init_nonlocal(ucell, pv, orb, run, basis); } ComparisonStats compare_zero_vector_potential(const int radial_grid_num, const int lebedev_grid_points) diff --git a/source/source_lcao/module_rt/velocity_op.cpp b/source/source_lcao/module_rt/velocity_op.cpp index 2986aedae15..acfb4fc37f2 100644 --- a/source/source_lcao/module_rt/velocity_op.cpp +++ b/source/source_lcao/module_rt/velocity_op.cpp @@ -16,13 +16,15 @@ Velocity_op::Velocity_op(const UnitCell* ucell_in, const Grid_Driver* GridD_in, const Parallel_Orbitals* paraV, const LCAO_Orbitals& orb, - const TwoCenterIntegrator* intor) + const TwoCenterIntegrator* intor, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis) : ucell(ucell_in), paraV(paraV) , orb_(orb), intor_(intor) { // for length gauge, the A(t) = 0 for all the time. this->cart_At = ModuleBase::Vector3(0,0,0); this->initialize_grad_term(GridD_in, paraV); - this->initialize_vcomm_r(GridD_in, paraV); + this->initialize_vcomm_r(GridD_in, paraV, run, basis); } template Velocity_op::~Velocity_op() @@ -34,13 +36,16 @@ Velocity_op::~Velocity_op() } //allocate space for current_term template -void Velocity_op::initialize_vcomm_r(const Grid_Driver* GridD, const Parallel_Orbitals* paraV) +void Velocity_op::initialize_vcomm_r(const Grid_Driver* GridD, + const Parallel_Orbitals* paraV, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis) { ModuleBase::TITLE("Velocity_op", "initialize_vcomm_r"); ModuleBase::timer::start("Velocity_op", "initialize_vcomm_r"); if(!init_done) { - r_calculator.init_nonlocal(*ucell, *paraV, orb_); + r_calculator.init_nonlocal(*ucell, *paraV, orb_, run, basis); init_done = true; } diff --git a/source/source_lcao/module_rt/velocity_op.h b/source/source_lcao/module_rt/velocity_op.h index 69f184872e7..c8b45c19c59 100644 --- a/source/source_lcao/module_rt/velocity_op.h +++ b/source/source_lcao/module_rt/velocity_op.h @@ -18,7 +18,9 @@ class Velocity_op const Grid_Driver* GridD_in, const Parallel_Orbitals* paraV, const LCAO_Orbitals& orb, - const TwoCenterIntegrator* intor); + const TwoCenterIntegrator* intor, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis); ~Velocity_op(); hamilt::HContainer>* get_current_term_pointer(const int& i)const @@ -46,7 +48,10 @@ class Velocity_op * HContainer is used to store the non-local pseudopotential matrix with specific atom-pairs * the size of HR will be fixed after initialization */ - void initialize_vcomm_r(const Grid_Driver* GridD_in, const Parallel_Orbitals* paraV); + void initialize_vcomm_r(const Grid_Driver* GridD_in, + const Parallel_Orbitals* paraV, + const ModuleContext::RunControl& run, + const ModuleContext::BasisInfo& basis); void initialize_grad_term(const Grid_Driver* GridD_in, const Parallel_Orbitals* paraV); /** diff --git a/source/source_main/driver.cpp b/source/source_main/driver.cpp index 8d88fa69f70..ec40de4f1f1 100644 --- a/source/source_main/driver.cpp +++ b/source/source_main/driver.cpp @@ -12,6 +12,7 @@ #include "source_io/module_parameter/parameter.h" #include "source_base/version.h" #include "source_base/parallel_global.h" +#include "source_context/simulation_context_builder.h" #ifdef __DSP #include "source_base/module_device/memory_op.h" #include "source_base/module_external/blas_connector.h" @@ -183,6 +184,7 @@ void Driver::reading() GlobalV::RANK_IN_POOL, GlobalV::MY_POOL); #endif + this->context_builder_.reset(new ModuleContext::SimulationContextBuilder(PARAM.inp, PARAM.globalv)); ModuleBase::timer::end("Driver", "reading"); } diff --git a/source/source_main/driver.h b/source/source_main/driver.h index e0204d079ca..a0853c3e4b6 100644 --- a/source/source_main/driver.h +++ b/source/source_main/driver.h @@ -1,6 +1,12 @@ #ifndef DRIVER_H #define DRIVER_H +#include + +namespace ModuleContext +{ +class SimulationContextBuilder; +} class Driver { @@ -40,6 +46,8 @@ class Driver // Init harewares according to Input parameters void init_hardware(); void finalize_hardware(); + + std::unique_ptr context_builder_; }; #endif diff --git a/source/source_main/driver_run.cpp b/source/source_main/driver_run.cpp index 9dc2935c64d..d45d854697d 100644 --- a/source/source_main/driver_run.cpp +++ b/source/source_main/driver_run.cpp @@ -11,6 +11,9 @@ #include "source_base/module_device/memory_op.h" #include "source_base/kernels/math_kernel_op.h" #include "source_hsolver/kernels/hegvd_op.h" +#include "source_context/simulation_context.h" +#include "source_context/simulation_context_binding.h" +#include "source_context/simulation_context_builder.h" #include #include @@ -69,6 +72,19 @@ void Driver::driver_run() //! 3: initialize Esolver and fill json-structure p_esolver->before_all_runners(ucell, PARAM.inp); + // Runtime-derived fields such as nlocal are final after pseudopotential + // and basis initialization in before_all_runners(). + this->context_builder_->capture_runtime(PARAM.globalv); + ModuleContext::SimulationContext context = this->context_builder_->finalize(PARAM.globalv); + ModuleContext::ScopedSimulationContextBinding context_binding(context); + + if (PARAM.inp.basis_type == "lcao" && PARAM.inp.esolver_type == "ks-lr") + { + p_esolver->runner(ucell, 0); // ground-state SCF, now inside the Context lifetime + p_esolver = ModuleESolver::transition_ksdft_to_lr(p_esolver, PARAM.inp, ucell); + p_esolver->before_all_runners(ucell, PARAM.inp); + } + // this Json part should be moved to before_all_runners, mohan 2024-05-12 #ifdef __RAPIDJSON Json::gen_stru_wrapper(&ucell); diff --git a/tools/03_code_analysis/agent_governance_check.py b/tools/03_code_analysis/agent_governance_check.py index 48975a728a6..d81e781cd53 100644 --- a/tools/03_code_analysis/agent_governance_check.py +++ b/tools/03_code_analysis/agent_governance_check.py @@ -244,6 +244,97 @@ def check_line_endings( GLOBAL_DEPENDENCY_RE = re.compile(r"\b(GlobalV::|GlobalC::|PARAM(?:\.|->|::|\b))") +MODULE_HS_PREFIX = "source/source_io/module_hs/" +MODULE_HS_FORBIDDEN = ( + (re.compile(r"\bPARAM\s*\."), "legacy PARAM access"), + (re.compile(r"\bGlobalV::"), "legacy GlobalV access"), + (re.compile(r"\bGlobalC::"), "legacy GlobalC access"), + (re.compile(r"\bSimulationContext\b"), "root SimulationContext dependency"), + (re.compile(r"\bcurrent_simulation_context\s*\("), "root Context accessor"), + (re.compile(r"(?:^|/)module_parameter/parameter\.h[\"\>]"), "parameter.h dependency"), + (re.compile(r"(?:^|/)global_variable\.h[\"\>]"), "global_variable.h dependency"), + (re.compile(r"(?:^|/)module_xc/exx_info\.h[\"\>]"), "global EXX definition dependency"), +) +CONTEXT_ACCESSOR_RE = re.compile(r"\bcurrent_simulation_context\s*\(") +CAPABILITY_DEFINE_RE = re.compile( + r"^\s*#\s*define\s+ABACUS_CAN_(?:READ|BIND)_SIMULATION_CONTEXT\b", re.MULTILINE +) +CONTEXT_ACCESS_ALLOWED = ( + "source/source_main/", + "source/source_esolver/", + "source/source_io/module_ctrl/", + "source/source_context/", +) +CONTEXT_ACCESS_ALLOWED_FILES = { + "source/source_lcao/module_lr/esolver_lrtd_lcao.cpp", +} + + +def repository_paths(root: Path, args: argparse.Namespace) -> List[str]: + if args.staged: + return git(["ls-files"], root).stdout.splitlines() + if args.head: + return git(["ls-tree", "-r", "--name-only", args.head], root).stdout.splitlines() + return [str(path.relative_to(root)) for path in root.rglob("*") if path.is_file()] + + +def read_repository_text(root: Path, path: str, args: argparse.Namespace) -> str: + return read_changed_file_bytes(root, path, args).decode("utf-8", errors="replace") + + +def check_context_access_boundaries( + findings: List[Finding], root: Path, args: argparse.Namespace +) -> None: + for path in repository_paths(root, args): + if Path(path).suffix.lower() not in SOURCE_REVIEW_EXTENSIONS: + continue + try: + content = read_repository_text(root, path, args) + except OSError: + continue + + if path.startswith(MODULE_HS_PREFIX): + for pattern, dependency in MODULE_HS_FORBIDDEN: + match = pattern.search(content) + if match: + add_finding( + findings, + "module_hs Context boundary", + BLOCK, + path, + content.count("\n", 0, match.start()) + 1, + f"module_hs contains forbidden {dependency}.", + "Pass the required ModuleContext small structure by const reference.", + allow_exception=False, + ) + + accessor = CONTEXT_ACCESSOR_RE.search(content) + allowed_accessor = path in CONTEXT_ACCESS_ALLOWED_FILES or path.startswith(CONTEXT_ACCESS_ALLOWED) + if accessor and not allowed_accessor: + add_finding( + findings, + "SimulationContext layer capability", + BLOCK, + path, + content.count("\n", 0, accessor.start()) + 1, + "Root Context accessor is used outside an L0/L1 orchestration file.", + "Move the read to L0/L1 and forward only the required small structures.", + allow_exception=False, + ) + + capability = CAPABILITY_DEFINE_RE.search(content) + if capability: + add_finding( + findings, + "SimulationContext capability ownership", + BLOCK, + path, + content.count("\n", 0, capability.start()) + 1, + "Source code defines a Context capability macro directly.", + "Grant capabilities from the owning CMake target or source property.", + allow_exception=False, + ) + def is_global_dependency_check_path(path: str) -> bool: if path.startswith("tools/03_code_analysis/"): @@ -731,6 +822,7 @@ def collect_findings(root: Path, args: argparse.Namespace) -> List[Finding]: check_line_endings(findings, root, changed, statuses, args) check_global_dependencies(findings, lines, removed_lines) + check_context_access_boundaries(findings, root, args) check_default_parameters(findings, lines) check_hpp_warnings(findings, statuses, lines) check_header_include_warnings(findings, lines) diff --git a/tools/03_code_analysis/test_agent_governance_check.py b/tools/03_code_analysis/test_agent_governance_check.py index 092da9faf9c..ff43ba4ed72 100644 --- a/tools/03_code_analysis/test_agent_governance_check.py +++ b/tools/03_code_analysis/test_agent_governance_check.py @@ -148,6 +148,58 @@ def test_allows_global_names_in_governance_checker_tests(self): self.assertEqual(result.returncode, 0, result.stdout + result.stderr) + def test_blocks_legacy_global_access_anywhere_in_module_hs(self): + self.write("source/source_io/module_hs/forbidden.h", "inline int rank() { return GlobalV::MY_RANK; }\n") + head = self.commit_change() + + result = self.run_checker("--base", self.base, "--head", head) + + self.assert_blocked_by(result, "module_hs Context boundary") + self.assertIn("legacy GlobalV access", result.stdout) + + def test_blocks_root_context_dependency_in_module_hs(self): + self.write("source/source_io/module_hs/forbidden.h", "const SimulationContext& context();\n") + head = self.commit_change() + + result = self.run_checker("--base", self.base, "--head", head) + + self.assert_blocked_by(result, "module_hs Context boundary") + self.assertIn("root SimulationContext dependency", result.stdout) + + def test_blocks_context_accessor_below_orchestration_layers(self): + self.write( + "source/source_lcao/module_operator_lcao/forbidden.h", + "inline void f() { current_simulation_context(); }\n", + ) + head = self.commit_change() + + result = self.run_checker("--base", self.base, "--head", head) + + self.assert_blocked_by(result, "SimulationContext layer capability") + + def test_allows_context_accessor_in_esolver(self): + self.write( + "source/source_esolver/allowed.h", + "inline void f() { current_simulation_context(); }\n", + ) + head = self.commit_change() + + result = self.run_checker("--base", self.base, "--head", head) + + self.assertEqual(result.returncode, 0, result.stdout + result.stderr) + self.assertNotIn("SimulationContext layer capability", result.stdout) + + def test_blocks_source_defined_context_capability(self): + self.write( + "source/source_esolver/forbidden.h", + "#define ABACUS_CAN_READ_SIMULATION_CONTEXT\n", + ) + head = self.commit_change() + + result = self.run_checker("--base", self.base, "--head", head) + + self.assert_blocked_by(result, "SimulationContext capability ownership") + def test_warns_for_default_parameters_added_to_headers(self): self.write("source/source_base/defaults.h", "void update_solver(int step = 0);\n") head = self.commit_change()