From 5225369481289bef488061a65c6eed395304ee79 Mon Sep 17 00:00:00 2001 From: jianrui geng Date: Tue, 7 Jul 2026 23:26:33 +0800 Subject: [PATCH 1/9] Add fixed-cell i-PI socket driver --- docs/advanced/input_files/input-main.md | 8 +- docs/parameters.yaml | 8 +- source/CMakeLists.txt | 4 +- .../read_input_item_system.cpp | 15 +- .../test_serial/read_input_item_test.cpp | 5 + source/source_main/driver.h | 1 + source/source_main/driver_ipi.cpp | 507 ++++++++++++++++++ source/source_main/driver_run.cpp | 9 +- source/source_main/ipi_socket.cpp | 241 +++++++++ source/source_main/ipi_socket.h | 40 ++ 10 files changed, 829 insertions(+), 9 deletions(-) create mode 100644 source/source_main/driver_ipi.cpp create mode 100644 source/source_main/ipi_socket.cpp create mode 100644 source/source_main/ipi_socket.h diff --git a/docs/advanced/input_files/input-main.md b/docs/advanced/input_files/input-main.md index 334144f44c5..c6c1b732991 100644 --- a/docs/advanced/input_files/input-main.md +++ b/docs/advanced/input_files/input-main.md @@ -584,6 +584,7 @@ - relax: perform structure relaxation calculations, the relax_nmax parameter depicts the maximal number of ionic iterations - cell-relax: perform cell relaxation calculations - md: perform molecular dynamics simulations + - socket: run as a socket client for external drivers using the i-PI protocol - get_pchg: obtain partial (band-decomposed) charge densities (for LCAO basis only). See out_pchg for more information - get_wf: obtain real space wave functions (for LCAO basis only). See out_wfc_norm and out_wfc_re_im for more information - get_s: obtain the overlap matrix formed by localized orbitals (for LCAO basis with multiple k points). the file name is SR.csr with file format being the same as that generated by out_mat_hs2 @@ -840,7 +841,12 @@ ### chg_extrap - **Type**: String -- **Description**: Charge extrapolation method for MD and relaxation calculations. +- **Description**: Charge extrapolation method for MD, relaxation, and socket-driven calculations. + + When set to default, ABACUS chooses second-order for md, first-order for + relax/cell-relax/socket, and atomic for other calculations. Socket-driven + molecular dynamics can explicitly set second-order if the external driver + updates structures smoothly enough for second-order extrapolation. - **Default**: default ### nb2d diff --git a/docs/parameters.yaml b/docs/parameters.yaml index 00ef9b9dcf7..d183783ad81 100644 --- a/docs/parameters.yaml +++ b/docs/parameters.yaml @@ -28,6 +28,7 @@ parameters: * relax: perform structure relaxation calculations, the relax_nmax parameter depicts the maximal number of ionic iterations * cell-relax: perform cell relaxation calculations * md: perform molecular dynamics simulations + * socket: run as a socket client for external drivers using the i-PI protocol * get_pchg: obtain partial (band-decomposed) charge densities (for LCAO basis only). See out_pchg for more information * get_wf: obtain real space wave functions (for LCAO basis only). See out_wfc_norm and out_wfc_re_im for more information * get_s: obtain the overlap matrix formed by localized orbitals (for LCAO basis with multiple k points). the file name is SR.csr with file format being the same as that generated by out_mat_hs2 @@ -328,7 +329,12 @@ parameters: category: System variables type: String description: | - Charge extrapolation method for MD and relaxation calculations. + Charge extrapolation method for MD, relaxation, and socket-driven calculations. + + When set to default, ABACUS chooses second-order for md, first-order for + relax/cell-relax/socket, and atomic for other calculations. Socket-driven + molecular dynamics can explicitly set second-order if the external driver + updates structures smoothly enough for second-order extrapolation. default_value: default unit: "" availability: "" diff --git a/source/CMakeLists.txt b/source/CMakeLists.txt index 0337707b61a..a5b78bd3b50 100644 --- a/source/CMakeLists.txt +++ b/source/CMakeLists.txt @@ -477,7 +477,9 @@ add_library( driver OBJECT source_main/driver.cpp - source_main/driver_run.cpp) + source_main/driver_run.cpp + source_main/driver_ipi.cpp + source_main/ipi_socket.cpp) list(APPEND device_srcs source_pw/module_pwdft/kernels/nonlocal_op.cpp diff --git a/source/source_io/module_parameter/read_input_item_system.cpp b/source/source_io/module_parameter/read_input_item_system.cpp index 3b0ae011ec4..6ed572b21f9 100644 --- a/source/source_io/module_parameter/read_input_item_system.cpp +++ b/source/source_io/module_parameter/read_input_item_system.cpp @@ -75,7 +75,7 @@ void ReadInput::item_system() } { Input_Item item("calculation"); - item.annotation = "scf; relax; md; cell-relax; nscf; get_s; get_wf; get_pchg; gen_bessel; gen_opt_abfs; test_memory; test_neighbour"; + item.annotation = "scf; relax; md; socket; cell-relax; nscf; get_s; get_wf; get_pchg; gen_bessel; gen_opt_abfs; test_memory; test_neighbour"; item.category = "System variables"; item.type = "String"; item.description = R"(Specify the type of calculation. @@ -85,6 +85,7 @@ void ReadInput::item_system() * relax: perform structure relaxation calculations, the relax_nmax parameter depicts the maximal number of ionic iterations * cell-relax: perform cell relaxation calculations * md: perform molecular dynamics simulations +* socket: run as a socket client for external drivers using the i-PI protocol * get_pchg: obtain partial (band-decomposed) charge densities (for LCAO basis only). See out_pchg for more information * get_wf: obtain real space wave functions (for LCAO basis only). See out_wfc_norm and out_wfc_re_im for more information * get_s: obtain the overlap matrix formed by localized orbitals (for LCAO basis with multiple k points). the file name is SR.csr with file format being the same as that generated by out_mat_hs2 @@ -102,6 +103,7 @@ void ReadInput::item_system() std::vector callist = {"scf", "relax", "md", + "socket", "cell-relax", "nscf", "get_s", @@ -266,7 +268,7 @@ void ReadInput::item_system() item.description = "If set to True, calculate the force at the end of the electronic iteration."; item.default_value = "False"; item.reset_value = [](const Input_Item& item, Parameter& para) { - std::vector use_force = {"cell-relax", "relax", "md"}; + std::vector use_force = {"cell-relax", "relax", "md", "socket"}; std::vector not_use_force = {"get_wf", "get_pchg", "get_s"}; if (std::find(use_force.begin(), use_force.end(), para.input.calculation) != use_force.end()) { @@ -873,7 +875,12 @@ Available options are: item.annotation = "atomic; first-order; second-order; dm:coefficients of SIA"; item.category = "System variables"; item.type = "String"; - item.description = "Charge extrapolation method for MD and relaxation calculations."; + item.description = R"(Charge extrapolation method for MD, relaxation, and socket-driven calculations. + +When set to default, ABACUS chooses second-order for md, first-order for +relax/cell-relax/socket, and atomic for other calculations. Socket-driven +molecular dynamics can explicitly set second-order if the external driver +updates structures smoothly enough for second-order extrapolation.)"; item.default_value = "default"; read_sync_string(input.chg_extrap); item.reset_value = [](const Input_Item& item, Parameter& para) { @@ -882,7 +889,7 @@ Available options are: para.input.chg_extrap = "second-order"; } else if (para.input.chg_extrap == "default" - && (para.input.calculation == "relax" || para.input.calculation == "cell-relax")) + && (para.input.calculation == "relax" || para.input.calculation == "cell-relax" || para.input.calculation == "socket")) { para.input.chg_extrap = "first-order"; } diff --git a/source/source_io/test_serial/read_input_item_test.cpp b/source/source_io/test_serial/read_input_item_test.cpp index 61673a418c7..eb043fe1b03 100644 --- a/source/source_io/test_serial/read_input_item_test.cpp +++ b/source/source_io/test_serial/read_input_item_test.cpp @@ -362,6 +362,11 @@ TEST_F(InputTest, Item_test) it->second.reset_value(it->second, param); EXPECT_EQ(param.input.chg_extrap, "first-order"); + param.input.chg_extrap = "default"; + param.input.calculation = "socket"; + it->second.reset_value(it->second, param); + EXPECT_EQ(param.input.chg_extrap, "first-order"); + param.input.chg_extrap = "default"; param.input.calculation = "none"; it->second.reset_value(it->second, param); diff --git a/source/source_main/driver.h b/source/source_main/driver.h index e0204d079ca..5a51701e26c 100644 --- a/source/source_main/driver.h +++ b/source/source_main/driver.h @@ -36,6 +36,7 @@ class Driver // the actual calculations void driver_run(); + void driver_ipi_run(); // Init harewares according to Input parameters void init_hardware(); diff --git a/source/source_main/driver_ipi.cpp b/source/source_main/driver_ipi.cpp new file mode 100644 index 00000000000..7be2105a9b3 --- /dev/null +++ b/source/source_main/driver_ipi.cpp @@ -0,0 +1,507 @@ +#include "source_main/driver.h" + +#include "source_main/ipi_socket.h" +#include "source_base/global_function.h" +#include "source_base/mathzone.h" +#include "source_cell/check_atomic_stru.h" +#include "source_cell/update_cell.h" +#include "source_esolver/esolver.h" +#include "source_base/global_variable.h" +#include "source_io/module_json/para_json.h" +#include "source_io/module_output/print_info.h" +#include "source_io/module_parameter/parameter.h" + +#ifdef __MPI +#include +#endif + +#include +#include +#include +#include +#include +#include +#include + +namespace +{ +constexpr double RY_TO_HARTREE = 0.5; +constexpr int IPI_RANK_ROOT = 0; + +bool is_root() +{ + return GlobalV::MY_RANK == IPI_RANK_ROOT; +} + +void bcast_int(int& value) +{ +#ifdef __MPI + MPI_Bcast(&value, 1, MPI_INT, IPI_RANK_ROOT, MPI_COMM_WORLD); +#else + (void)value; +#endif +} + +void bcast_double_vector(std::vector& values) +{ +#ifdef __MPI + MPI_Bcast(values.data(), static_cast(values.size()), MPI_DOUBLE, IPI_RANK_ROOT, MPI_COMM_WORLD); +#else + (void)values; +#endif +} + +std::string bcast_string(std::string value) +{ + int nbytes = static_cast(value.size()); + bcast_int(nbytes); + if (nbytes < 0) + { + throw std::runtime_error("negative string length in i-PI broadcast"); + } + if (!is_root()) + { + value.assign(static_cast(nbytes), '\0'); + } +#ifdef __MPI + if (nbytes > 0) + { + MPI_Bcast(&value[0], nbytes, MPI_CHAR, IPI_RANK_ROOT, MPI_COMM_WORLD); + } +#endif + return value; +} + +void throw_if_root_io_failed(int root_failed, const std::string& root_message) +{ + bcast_int(root_failed); + const std::string message = bcast_string(root_message); + if (root_failed != 0) + { + throw std::runtime_error(message.empty() ? "i-PI socket I/O failed" : message); + } +} + +std::string bcast_header(const std::string& root_header) +{ + char buffer[13] = {' ', ' ', ' ', ' ', ' ', ' ', ' ', ' ', ' ', ' ', ' ', ' ', '\0'}; + if (is_root()) + { + const std::size_t n = root_header.size() > 12 ? 12 : root_header.size(); + for (std::size_t i = 0; i < n; ++i) + { + buffer[i] = root_header[i]; + } + } +#ifdef __MPI + MPI_Bcast(buffer, 12, MPI_CHAR, IPI_RANK_ROOT, MPI_COMM_WORLD); +#endif + std::string out(buffer, 12); + while (!out.empty() && out.back() == ' ') + { + out.pop_back(); + } + return out; +} + +std::string ipi_address() +{ + const char* env = std::getenv("ABACUS_IPI_ADDRESS"); + if (env == nullptr || std::string(env).empty()) + { + return "localhost:31415"; + } + return std::string(env); +} + +std::vector ipi_cell_bohr_from_unitcell(const UnitCell& ucell) +{ + const double lat0 = ucell.lat0; + // ASE/i-PI sends POSDATA cell as cell.T in C order. ABACUS stores + // lattice vectors as rows in latvec, so use the transposed order here. + return { + ucell.latvec.e11 * lat0, ucell.latvec.e21 * lat0, ucell.latvec.e31 * lat0, + ucell.latvec.e12 * lat0, ucell.latvec.e22 * lat0, ucell.latvec.e32 * lat0, + ucell.latvec.e13 * lat0, ucell.latvec.e23 * lat0, ucell.latvec.e33 * lat0, + }; +} + +double max_abs_delta(const std::vector& a, const std::vector& b) +{ + if (a.size() != b.size()) + { + return 1.0e99; + } + double out = 0.0; + for (std::size_t i = 0; i < a.size(); ++i) + { + out = std::max(out, std::abs(a[i] - b[i])); + } + return out; +} + +void set_positions_from_ipi_bohr(UnitCell& ucell, const std::vector& positions_bohr) +{ + if (positions_bohr.size() != static_cast(3 * ucell.nat)) + { + throw std::runtime_error("POSDATA atom count does not match STRU."); + } + int iat = 0; + for (int it = 0; it < ucell.ntype; ++it) + { + Atom* atom = &ucell.atoms[it]; + for (int ia = 0; ia < atom->na; ++ia) + { + const double tau_x = positions_bohr[3 * iat + 0] / ucell.lat0; + const double tau_y = positions_bohr[3 * iat + 1] / ucell.lat0; + const double tau_z = positions_bohr[3 * iat + 2] / ucell.lat0; + + double dx = 0.0; + double dy = 0.0; + double dz = 0.0; + ModuleBase::Mathzone::Cartesian_to_Direct(tau_x, + tau_y, + tau_z, + ucell.latvec.e11, + ucell.latvec.e12, + ucell.latvec.e13, + ucell.latvec.e21, + ucell.latvec.e22, + ucell.latvec.e23, + ucell.latvec.e31, + ucell.latvec.e32, + ucell.latvec.e33, + dx, + dy, + dz); + + atom->dis[ia].x = dx - atom->taud[ia].x; + atom->dis[ia].y = dy - atom->taud[ia].y; + atom->dis[ia].z = dz - atom->taud[ia].z; + atom->taud[ia].x = dx; + atom->taud[ia].y = dy; + atom->taud[ia].z = dz; + atom->tau[ia].x = tau_x; + atom->tau[ia].y = tau_y; + atom->tau[ia].z = tau_z; + ++iat; + } + } + unitcell::periodic_boundary_adjustment(ucell.atoms, ucell.latvec, ucell.ntype); + ucell.ionic_position_updated = true; + ucell.cell_parameter_updated = false; +} + +std::vector flatten_forces_hartree_per_bohr(const ModuleBase::matrix& force) +{ + std::vector out(static_cast(force.nr * force.nc)); + for (int iat = 0; iat < force.nr; ++iat) + { + for (int idir = 0; idir < force.nc; ++idir) + { + out[static_cast(3 * iat + idir)] = force(iat, idir) * RY_TO_HARTREE; + } + } + return out; +} + +class CalculationModeGuard +{ + public: + explicit CalculationModeGuard(const std::string& inner_calculation) + : outer_calculation_(PARAM.inp.calculation) + { + const_cast(PARAM.inp.calculation) = inner_calculation; + } + + ~CalculationModeGuard() + { + const_cast(PARAM.inp.calculation) = outer_calculation_; + } + + CalculationModeGuard(const CalculationModeGuard&) = delete; + CalculationModeGuard& operator=(const CalculationModeGuard&) = delete; + + private: + std::string outer_calculation_; +}; +} // namespace + +void Driver::driver_ipi_run() +{ + ModuleBase::TITLE("Driver", "driver_ipi_run"); + + // "socket" is an outer driver mode. The KS/LCAO ESolver internals use the + // standard SCF code path for each POSDATA request from the i-PI protocol. + CalculationModeGuard calculation_guard("scf"); + + UnitCell ucell; + ucell.setup(PARAM.inp.latname, PARAM.inp.ntype, PARAM.inp.lmaxmax, PARAM.inp.init_vel, PARAM.inp.fixed_axes); + ucell.setup_cell(PARAM.globalv.global_in_stru, GlobalV::ofs_running); + unitcell::check_atomic_stru(ucell, PARAM.inp.min_dist_coef); + + IpiSocket socket; + std::unique_ptr p_esolver; + bool hardware_initialized = false; + bool esolver_ready = false; + bool runner_completed = false; + std::string pending_error; + + try + { + this->init_hardware(); + hardware_initialized = true; + + p_esolver.reset(ModuleESolver::init_esolver(PARAM.inp, ucell)); + p_esolver->before_all_runners(ucell, PARAM.inp); + esolver_ready = true; + +#ifdef __RAPIDJSON + Json::gen_stru_wrapper(&ucell); +#endif + + int io_failed = 0; + std::string io_message; + if (is_root()) + { + try + { + const std::string address = ipi_address(); + GlobalV::ofs_running << " ABACUS socket driver connecting to i-PI endpoint " << address << std::endl; + socket.connect(address); + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + throw_if_root_io_failed(io_failed, io_message); + + bool isinit = false; + bool hasdata = false; + int istep = 0; + int nat_return = ucell.nat; + double energy_hartree = 0.0; + std::vector forces_hartree_bohr(static_cast(3 * ucell.nat), 0.0); + std::vector virial_hartree(9, 0.0); + + const std::vector reference_cell = ipi_cell_bohr_from_unitcell(ucell); + + while (true) + { + std::string header; + io_failed = 0; + io_message.clear(); + if (is_root()) + { + try + { + header = socket.read_header(); + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + throw_if_root_io_failed(io_failed, io_message); + header = bcast_header(header); + + if (header == "STATUS") + { + io_failed = 0; + io_message.clear(); + if (is_root()) + { + try + { + if (hasdata) + { + socket.write_header("HAVEDATA"); + } + else if (isinit) + { + socket.write_header("READY"); + } + else + { + socket.write_header("NEEDINIT"); + } + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + throw_if_root_io_failed(io_failed, io_message); + } + else if (header == "INIT") + { + int rid = 0; + int nbytes = 0; + std::string params; + io_failed = 0; + io_message.clear(); + if (is_root()) + { + try + { + rid = socket.read_int(); + nbytes = socket.read_int(); + if (nbytes < 0) + { + throw std::runtime_error("negative INIT payload length from i-PI socket"); + } + if (nbytes > 0) + { + params = socket.read_string(static_cast(nbytes)); + } + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + throw_if_root_io_failed(io_failed, io_message); + bcast_int(rid); + bcast_int(nbytes); + if (nbytes > 0 && is_root()) + { + GlobalV::ofs_running << " ABACUS socket INIT params " << params << std::endl; + } + isinit = true; + if (is_root()) + { + GlobalV::ofs_running << " ABACUS socket INIT replica " << rid << std::endl; + } + } + else if (header == "POSDATA") + { + std::vector cell(9, 0.0); + std::vector inv_cell(9, 0.0); + int nat_socket = 0; + std::vector positions; + io_failed = 0; + io_message.clear(); + if (is_root()) + { + try + { + cell = socket.read_doubles(9); + inv_cell = socket.read_doubles(9); + nat_socket = socket.read_int(); + if (nat_socket < 0) + { + throw std::runtime_error("negative POSDATA atom count from i-PI socket"); + } + positions = socket.read_doubles(static_cast(3 * nat_socket)); + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + throw_if_root_io_failed(io_failed, io_message); + bcast_double_vector(cell); + bcast_double_vector(inv_cell); + bcast_int(nat_socket); + if (!is_root()) + { + positions.assign(static_cast(3 * nat_socket), 0.0); + } + bcast_double_vector(positions); + + if (nat_socket != ucell.nat) + { + throw std::runtime_error("POSDATA atom count does not match STRU."); + } + const double max_cell_delta_bohr = max_abs_delta(cell, reference_cell); + if (max_cell_delta_bohr > 1.0e-6) + { + throw std::runtime_error("variable-cell socket updates are not supported yet."); + } + + set_positions_from_ipi_bohr(ucell, positions); + p_esolver->runner(ucell, istep); + runner_completed = true; + energy_hartree = p_esolver->cal_energy() * RY_TO_HARTREE; + ModuleBase::matrix force; + if (PARAM.inp.cal_force) + { + p_esolver->cal_force(ucell, force); + forces_hartree_bohr = flatten_forces_hartree_per_bohr(force); + } + ++istep; + hasdata = true; + } + else if (header == "GETFORCE") + { + io_failed = 0; + io_message.clear(); + if (is_root()) + { + try + { + socket.write_header("FORCEREADY"); + socket.write_double(energy_hartree); + socket.write_int(nat_return); + socket.write_doubles(forces_hartree_bohr); + socket.write_doubles(virial_hartree); + socket.write_int(0); + } + catch (const std::exception& exc) + { + io_failed = 1; + io_message = exc.what(); + } + } + throw_if_root_io_failed(io_failed, io_message); + isinit = false; + hasdata = false; + } + else + { + if (is_root()) + { + GlobalV::ofs_running << " ABACUS socket driver exiting on header " << header << std::endl; + } + break; + } + } + } + catch (const std::exception& exc) + { + pending_error = exc.what(); + if (is_root()) + { + GlobalV::ofs_running << " ABACUS socket driver ended with error: " << pending_error << std::endl; + } + } + + if (is_root()) + { + socket.close(); + } + if (esolver_ready && runner_completed && p_esolver) + { + p_esolver->after_all_runners(ucell); + } + p_esolver.reset(); + if (hardware_initialized) + { + this->finalize_hardware(); + } + +#ifdef __RAPIDJSON + Json::create_Json(&ucell, PARAM); +#endif + + if (!pending_error.empty()) + { + ModuleBase::WARNING_QUIT("ABACUS socket", pending_error); + } +} diff --git a/source/source_main/driver_run.cpp b/source/source_main/driver_run.cpp index 4911b133de0..9285f9b2330 100644 --- a/source/source_main/driver_run.cpp +++ b/source/source_main/driver_run.cpp @@ -38,6 +38,13 @@ void Driver::driver_run() { ModuleBase::TITLE("Driver", "driver_run"); + const std::string cal = PARAM.inp.calculation; + if (cal == "socket") + { + this->driver_ipi_run(); + return; + } + //! 1: setup cell and atom information // this warning should not be here, mohan 2024-05-22 #ifndef __LCAO @@ -71,8 +78,6 @@ void Driver::driver_run() Json::gen_stru_wrapper(&ucell); #endif - const std::string cal = PARAM.inp.calculation; - //! 4: different types of calculations if (cal == "md") { diff --git a/source/source_main/ipi_socket.cpp b/source/source_main/ipi_socket.cpp new file mode 100644 index 00000000000..eca053566ff --- /dev/null +++ b/source/source_main/ipi_socket.cpp @@ -0,0 +1,241 @@ +#include "ipi_socket.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace +{ +constexpr std::size_t IPI_HEADER_LEN = 12; + +std::string errno_message(const std::string& prefix) +{ + return prefix + ": " + std::strerror(errno); +} + +std::string trim_header(const char* data) +{ + std::string value(data, IPI_HEADER_LEN); + while (!value.empty() && value.back() == ' ') + { + value.pop_back(); + } + return value; +} + +std::string padded_header(const std::string& header) +{ + if (header.size() > IPI_HEADER_LEN) + { + throw std::runtime_error("i-PI header is longer than 12 bytes: " + header); + } + std::string out = header; + out.resize(IPI_HEADER_LEN, ' '); + return out; +} +} // namespace + +IpiSocket::~IpiSocket() +{ + this->close(); +} + +void IpiSocket::connect(const std::string& address) +{ + this->close(); + const std::size_t colon = address.rfind(':'); + if (colon == std::string::npos) + { + throw std::runtime_error("i-PI address must be host:port or path:UNIX, got " + address); + } + const std::string host = address.substr(0, colon); + const std::string service = address.substr(colon + 1); + + if (service == "UNIX") + { + fd_ = ::socket(AF_UNIX, SOCK_STREAM, 0); + if (fd_ < 0) + { + throw std::runtime_error(errno_message("failed to create UNIX socket")); + } + sockaddr_un addr; + std::memset(&addr, 0, sizeof(addr)); + addr.sun_family = AF_UNIX; + if (host.size() >= sizeof(addr.sun_path)) + { + this->close(); + throw std::runtime_error("UNIX socket path too long: " + host); + } + std::strncpy(addr.sun_path, host.c_str(), sizeof(addr.sun_path) - 1); + if (::connect(fd_, reinterpret_cast(&addr), sizeof(addr)) != 0) + { + const std::string msg = errno_message("failed to connect UNIX i-PI socket " + host); + this->close(); + throw std::runtime_error(msg); + } + return; + } + + addrinfo hints; + std::memset(&hints, 0, sizeof(hints)); + hints.ai_family = AF_UNSPEC; + hints.ai_socktype = SOCK_STREAM; + + addrinfo* result = nullptr; + const int gai = ::getaddrinfo(host.c_str(), service.c_str(), &hints, &result); + if (gai != 0) + { + throw std::runtime_error("failed to resolve i-PI socket " + address + ": " + ::gai_strerror(gai)); + } + + std::string last_error; + for (addrinfo* rp = result; rp != nullptr; rp = rp->ai_next) + { + fd_ = ::socket(rp->ai_family, rp->ai_socktype, rp->ai_protocol); + if (fd_ < 0) + { + last_error = errno_message("failed to create INET socket"); + continue; + } + if (::connect(fd_, rp->ai_addr, rp->ai_addrlen) == 0) + { + ::freeaddrinfo(result); + return; + } + last_error = errno_message("failed to connect INET i-PI socket " + address); + this->close(); + } + ::freeaddrinfo(result); + throw std::runtime_error(last_error.empty() ? "failed to connect i-PI socket " + address : last_error); +} + +void IpiSocket::close() +{ + if (fd_ >= 0) + { + ::close(fd_); + fd_ = -1; + } +} + +std::string IpiSocket::read_header() +{ + char header[IPI_HEADER_LEN]; + this->read_exact(header, sizeof(header)); + return trim_header(header); +} + +void IpiSocket::write_header(const std::string& header) +{ + const std::string padded = padded_header(header); + this->write_exact(padded.data(), padded.size()); +} + +int IpiSocket::read_int() +{ + int value = 0; + this->read_exact(&value, sizeof(value)); + return value; +} + +void IpiSocket::write_int(int value) +{ + this->write_exact(&value, sizeof(value)); +} + +double IpiSocket::read_double() +{ + double value = 0.0; + this->read_exact(&value, sizeof(value)); + return value; +} + +void IpiSocket::write_double(double value) +{ + this->write_exact(&value, sizeof(value)); +} + +std::vector IpiSocket::read_doubles(std::size_t n) +{ + std::vector values(n); + if (!values.empty()) + { + this->read_exact(values.data(), values.size() * sizeof(double)); + } + return values; +} + +void IpiSocket::write_doubles(const std::vector& values) +{ + if (!values.empty()) + { + this->write_exact(values.data(), values.size() * sizeof(double)); + } +} + +std::string IpiSocket::read_string(std::size_t nbytes) +{ + std::string value(nbytes, '\0'); + if (nbytes > 0) + { + this->read_exact(&value[0], nbytes); + } + return value; +} + +void IpiSocket::read_exact(void* data, std::size_t nbytes) +{ + char* cursor = static_cast(data); + std::size_t done = 0; + while (done < nbytes) + { + const ssize_t nread = ::recv(fd_, cursor + done, nbytes - done, 0); + if (nread == 0) + { + throw std::runtime_error("i-PI socket closed while reading"); + } + if (nread < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("i-PI socket read failed")); + } + done += static_cast(nread); + } +} + +void IpiSocket::write_exact(const void* data, std::size_t nbytes) +{ + const char* cursor = static_cast(data); + std::size_t done = 0; + while (done < nbytes) + { +#ifdef MSG_NOSIGNAL + const int flags = MSG_NOSIGNAL; +#else + const int flags = 0; +#endif + const ssize_t nwritten = ::send(fd_, cursor + done, nbytes - done, flags); + if (nwritten == 0) + { + throw std::runtime_error("i-PI socket closed while writing"); + } + if (nwritten < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("i-PI socket write failed")); + } + done += static_cast(nwritten); + } +} diff --git a/source/source_main/ipi_socket.h b/source/source_main/ipi_socket.h new file mode 100644 index 00000000000..3e1713e868d --- /dev/null +++ b/source/source_main/ipi_socket.h @@ -0,0 +1,40 @@ +#ifndef ABACUS_IPI_SOCKET_H +#define ABACUS_IPI_SOCKET_H + +#include +#include +#include + +class IpiSocket +{ + public: + IpiSocket() = default; + ~IpiSocket(); + + IpiSocket(const IpiSocket&) = delete; + IpiSocket& operator=(const IpiSocket&) = delete; + + void connect(const std::string& address); + void close(); + + std::string read_header(); + void write_header(const std::string& header); + + int read_int(); + void write_int(int value); + + double read_double(); + void write_double(double value); + + std::vector read_doubles(std::size_t n); + void write_doubles(const std::vector& values); + std::string read_string(std::size_t nbytes); + + private: + int fd_ = -1; + + void read_exact(void* data, std::size_t nbytes); + void write_exact(const void* data, std::size_t nbytes); +}; + +#endif From e0f70ca031422c4821d5304050c0035f1be85ec7 Mon Sep 17 00:00:00 2001 From: jianrui geng Date: Thu, 9 Jul 2026 09:13:53 +0800 Subject: [PATCH 2/9] Refactor i-PI socket driver into source_relax Route calculation=socket through Relax_Driver, keep ESolver lifecycle in driver_run, and add source_relax socket transport tests. --- source/CMakeLists.txt | 4 +- source/source_main/driver.h | 1 - source/source_main/driver_run.cpp | 22 +- source/source_relax/CMakeLists.txt | 2 + .../ipi_socket.cpp | 31 ++- .../ipi_socket.h | 7 + source/source_relax/relax_driver.cpp | 9 + .../socket_driver.cpp} | 243 ++++++----------- source/source_relax/socket_driver.h | 22 ++ source/source_relax/test/CMakeLists.txt | 6 + source/source_relax/test/ipi_socket_test.cpp | 256 ++++++++++++++++++ 11 files changed, 431 insertions(+), 172 deletions(-) rename source/{source_main => source_relax}/ipi_socket.cpp (87%) rename source/{source_main => source_relax}/ipi_socket.h (85%) rename source/{source_main/driver_ipi.cpp => source_relax/socket_driver.cpp} (64%) create mode 100644 source/source_relax/socket_driver.h create mode 100644 source/source_relax/test/ipi_socket_test.cpp diff --git a/source/CMakeLists.txt b/source/CMakeLists.txt index a5b78bd3b50..0337707b61a 100644 --- a/source/CMakeLists.txt +++ b/source/CMakeLists.txt @@ -477,9 +477,7 @@ add_library( driver OBJECT source_main/driver.cpp - source_main/driver_run.cpp - source_main/driver_ipi.cpp - source_main/ipi_socket.cpp) + source_main/driver_run.cpp) list(APPEND device_srcs source_pw/module_pwdft/kernels/nonlocal_op.cpp diff --git a/source/source_main/driver.h b/source/source_main/driver.h index 5a51701e26c..e0204d079ca 100644 --- a/source/source_main/driver.h +++ b/source/source_main/driver.h @@ -36,7 +36,6 @@ class Driver // the actual calculations void driver_run(); - void driver_ipi_run(); // Init harewares according to Input parameters void init_hardware(); diff --git a/source/source_main/driver_run.cpp b/source/source_main/driver_run.cpp index 9285f9b2330..3e7d669b6c2 100644 --- a/source/source_main/driver_run.cpp +++ b/source/source_main/driver_run.cpp @@ -39,11 +39,7 @@ void Driver::driver_run() ModuleBase::TITLE("Driver", "driver_run"); const std::string cal = PARAM.inp.calculation; - if (cal == "socket") - { - this->driver_ipi_run(); - return; - } + const bool socket_mode = (cal == "socket"); //! 1: setup cell and atom information // this warning should not be here, mohan 2024-05-22 @@ -66,12 +62,20 @@ void Driver::driver_run() unitcell::check_atomic_stru(ucell, PARAM.inp.min_dist_coef); //! 2: initialize the ESolver (depends on a set-up ucell after `setup_cell`) + Input_para esolver_inp = PARAM.inp; + Input_para driver_inp = PARAM.inp; + if (socket_mode) + { + esolver_inp.calculation = "scf"; + driver_inp.calculation = cal; + } + this->init_hardware(); - ModuleESolver::ESolver* p_esolver = ModuleESolver::init_esolver(PARAM.inp, ucell); + ModuleESolver::ESolver* p_esolver = ModuleESolver::init_esolver(esolver_inp, ucell); //! 3: initialize Esolver and fill json-structure - p_esolver->before_all_runners(ucell, PARAM.inp); + p_esolver->before_all_runners(ucell, esolver_inp); // this Json part should be moved to before_all_runners, mohan 2024-05-12 #ifdef __RAPIDJSON @@ -83,10 +87,10 @@ void Driver::driver_run() { Run_MD::md_line(ucell, p_esolver, PARAM); } - else if (cal == "scf" || cal == "relax" || cal == "cell-relax" || cal == "nscf") + else if (cal == "scf" || cal == "relax" || cal == "cell-relax" || cal == "nscf" || cal == "socket") { Relax_Driver rl_driver; - rl_driver.relax_driver(p_esolver, ucell, PARAM.inp, GlobalV::ofs_running); + rl_driver.relax_driver(p_esolver, ucell, driver_inp, GlobalV::ofs_running); } else if (cal == "get_s") { diff --git a/source/source_relax/CMakeLists.txt b/source/source_relax/CMakeLists.txt index 9e7ef96b0e2..84f69c06c80 100644 --- a/source/source_relax/CMakeLists.txt +++ b/source/source_relax/CMakeLists.txt @@ -2,6 +2,8 @@ add_library( relax OBJECT relax_data.cpp + ipi_socket.cpp + socket_driver.cpp cg_base.cpp relax_driver.cpp relax_sync.cpp diff --git a/source/source_main/ipi_socket.cpp b/source/source_relax/ipi_socket.cpp similarity index 87% rename from source/source_main/ipi_socket.cpp rename to source/source_relax/ipi_socket.cpp index eca053566ff..304a0181813 100644 --- a/source/source_main/ipi_socket.cpp +++ b/source/source_relax/ipi_socket.cpp @@ -1,4 +1,4 @@ -#include "ipi_socket.h" +#include "source_relax/ipi_socket.h" #include #include @@ -41,6 +41,10 @@ std::string padded_header(const std::string& header) } } // namespace +IpiSocketClosed::IpiSocketClosed(const std::string& message) : std::runtime_error(message) +{ +} + IpiSocket::~IpiSocket() { this->close(); @@ -127,7 +131,28 @@ void IpiSocket::close() std::string IpiSocket::read_header() { char header[IPI_HEADER_LEN]; - this->read_exact(header, sizeof(header)); + std::size_t done = 0; + while (done < sizeof(header)) + { + const ssize_t nread = ::recv(fd_, header + done, sizeof(header) - done, 0); + if (nread == 0) + { + if (done == 0) + { + throw IpiSocketClosed("i-PI socket closed before next header"); + } + throw std::runtime_error("i-PI socket closed while reading header"); + } + if (nread < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("i-PI socket header read failed")); + } + done += static_cast(nread); + } return trim_header(header); } @@ -198,7 +223,7 @@ void IpiSocket::read_exact(void* data, std::size_t nbytes) const ssize_t nread = ::recv(fd_, cursor + done, nbytes - done, 0); if (nread == 0) { - throw std::runtime_error("i-PI socket closed while reading"); + throw IpiSocketClosed("i-PI socket closed while reading"); } if (nread < 0) { diff --git a/source/source_main/ipi_socket.h b/source/source_relax/ipi_socket.h similarity index 85% rename from source/source_main/ipi_socket.h rename to source/source_relax/ipi_socket.h index 3e1713e868d..28ce2cb81ff 100644 --- a/source/source_main/ipi_socket.h +++ b/source/source_relax/ipi_socket.h @@ -2,9 +2,16 @@ #define ABACUS_IPI_SOCKET_H #include +#include #include #include +class IpiSocketClosed : public std::runtime_error +{ + public: + explicit IpiSocketClosed(const std::string& message); +}; + class IpiSocket { public: diff --git a/source/source_relax/relax_driver.cpp b/source/source_relax/relax_driver.cpp index 26128bf0408..0cada1bfc3e 100644 --- a/source/source_relax/relax_driver.cpp +++ b/source/source_relax/relax_driver.cpp @@ -1,4 +1,5 @@ #include "relax_driver.h" +#include "socket_driver.h" #include "source_base/global_file.h" #include "source_io/module_output/cif_io.h" #include "source_io/module_json/output_info.h" @@ -17,6 +18,14 @@ void Relax_Driver::relax_driver( ModuleBase::TITLE("Relax_Driver", "relax_driver"); ModuleBase::timer::start("Relax_Driver", "relax_driver"); + if (inp.calculation == "socket") + { + Socket_Driver socket_driver; + socket_driver.socket_driver(p_esolver, ucell, inp, ofs_running); + ModuleBase::timer::end("Relax_Driver", "relax_driver"); + return; + } + this->init_relax(ucell.nat, inp); // steps[0]: istep (main iteration step) diff --git a/source/source_main/driver_ipi.cpp b/source/source_relax/socket_driver.cpp similarity index 64% rename from source/source_main/driver_ipi.cpp rename to source/source_relax/socket_driver.cpp index 7be2105a9b3..02928fec9d3 100644 --- a/source/source_main/driver_ipi.cpp +++ b/source/source_relax/socket_driver.cpp @@ -1,25 +1,17 @@ -#include "source_main/driver.h" +#include "socket_driver.h" -#include "source_main/ipi_socket.h" +#include "source_relax/ipi_socket.h" #include "source_base/global_function.h" +#include "source_base/global_variable.h" #include "source_base/mathzone.h" -#include "source_cell/check_atomic_stru.h" +#include "source_base/parallel_common.h" +#include "source_base/timer.h" #include "source_cell/update_cell.h" -#include "source_esolver/esolver.h" -#include "source_base/global_variable.h" -#include "source_io/module_json/para_json.h" -#include "source_io/module_output/print_info.h" -#include "source_io/module_parameter/parameter.h" - -#ifdef __MPI -#include -#endif #include #include #include -#include -#include +#include #include #include @@ -33,75 +25,42 @@ bool is_root() return GlobalV::MY_RANK == IPI_RANK_ROOT; } -void bcast_int(int& value) -{ -#ifdef __MPI - MPI_Bcast(&value, 1, MPI_INT, IPI_RANK_ROOT, MPI_COMM_WORLD); -#else - (void)value; -#endif -} - void bcast_double_vector(std::vector& values) { -#ifdef __MPI - MPI_Bcast(values.data(), static_cast(values.size()), MPI_DOUBLE, IPI_RANK_ROOT, MPI_COMM_WORLD); -#else - (void)values; -#endif + if (!values.empty()) + { + Parallel_Common::bcast_double(values.data(), static_cast(values.size())); + } } -std::string bcast_string(std::string value) +void bcast_socket_string(std::string& value) { - int nbytes = static_cast(value.size()); - bcast_int(nbytes); - if (nbytes < 0) - { - throw std::runtime_error("negative string length in i-PI broadcast"); - } + int size = static_cast(value.size()); + Parallel_Common::bcast_int(size); if (!is_root()) { - value.assign(static_cast(nbytes), '\0'); + value.resize(static_cast(size)); } -#ifdef __MPI - if (nbytes > 0) + if (size > 0) { - MPI_Bcast(&value[0], nbytes, MPI_CHAR, IPI_RANK_ROOT, MPI_COMM_WORLD); + Parallel_Common::bcast_char(&value[0], size); } -#endif - return value; } -void throw_if_root_io_failed(int root_failed, const std::string& root_message) +void quit_if_root_io_failed(int root_failed, std::string root_message) { - bcast_int(root_failed); - const std::string message = bcast_string(root_message); + Parallel_Common::bcast_int(root_failed); + bcast_socket_string(root_message); if (root_failed != 0) { - throw std::runtime_error(message.empty() ? "i-PI socket I/O failed" : message); + ModuleBase::WARNING_QUIT("ABACUS socket", root_message.empty() ? "i-PI socket I/O failed" : root_message); } } -std::string bcast_header(const std::string& root_header) +std::string bcast_header(std::string header) { - char buffer[13] = {' ', ' ', ' ', ' ', ' ', ' ', ' ', ' ', ' ', ' ', ' ', ' ', '\0'}; - if (is_root()) - { - const std::size_t n = root_header.size() > 12 ? 12 : root_header.size(); - for (std::size_t i = 0; i < n; ++i) - { - buffer[i] = root_header[i]; - } - } -#ifdef __MPI - MPI_Bcast(buffer, 12, MPI_CHAR, IPI_RANK_ROOT, MPI_COMM_WORLD); -#endif - std::string out(buffer, 12); - while (!out.empty() && out.back() == ' ') - { - out.pop_back(); - } - return out; + bcast_socket_string(header); + return header; } std::string ipi_address() @@ -144,8 +103,9 @@ void set_positions_from_ipi_bohr(UnitCell& ucell, const std::vector& pos { if (positions_bohr.size() != static_cast(3 * ucell.nat)) { - throw std::runtime_error("POSDATA atom count does not match STRU."); + ModuleBase::WARNING_QUIT("ABACUS socket", "POSDATA atom count does not match STRU."); } + int iat = 0; for (int it = 0; it < ucell.ntype; ++it) { @@ -204,62 +164,29 @@ std::vector flatten_forces_hartree_per_bohr(const ModuleBase::matrix& fo } return out; } +} // namespace -class CalculationModeGuard +void Socket_Driver::socket_driver(ModuleESolver::ESolver* p_esolver, + UnitCell& ucell, + const Input_para& inp, + std::ofstream& ofs_running) { - public: - explicit CalculationModeGuard(const std::string& inner_calculation) - : outer_calculation_(PARAM.inp.calculation) + ModuleBase::TITLE("Socket_Driver", "socket_driver"); + ModuleBase::timer::start("Socket_Driver", "socket_driver"); + + if (p_esolver == nullptr) { - const_cast(PARAM.inp.calculation) = inner_calculation; + ModuleBase::WARNING_QUIT("ABACUS socket", "socket driver requires a valid ESolver."); } - - ~CalculationModeGuard() + if (!inp.cal_force) { - const_cast(PARAM.inp.calculation) = outer_calculation_; + ModuleBase::WARNING_QUIT("ABACUS socket", "socket calculation requires cal_force=1 for i-PI GETFORCE."); } - CalculationModeGuard(const CalculationModeGuard&) = delete; - CalculationModeGuard& operator=(const CalculationModeGuard&) = delete; - - private: - std::string outer_calculation_; -}; -} // namespace - -void Driver::driver_ipi_run() -{ - ModuleBase::TITLE("Driver", "driver_ipi_run"); - - // "socket" is an outer driver mode. The KS/LCAO ESolver internals use the - // standard SCF code path for each POSDATA request from the i-PI protocol. - CalculationModeGuard calculation_guard("scf"); - - UnitCell ucell; - ucell.setup(PARAM.inp.latname, PARAM.inp.ntype, PARAM.inp.lmaxmax, PARAM.inp.init_vel, PARAM.inp.fixed_axes); - ucell.setup_cell(PARAM.globalv.global_in_stru, GlobalV::ofs_running); - unitcell::check_atomic_stru(ucell, PARAM.inp.min_dist_coef); - IpiSocket socket; - std::unique_ptr p_esolver; - bool hardware_initialized = false; - bool esolver_ready = false; - bool runner_completed = false; - std::string pending_error; try { - this->init_hardware(); - hardware_initialized = true; - - p_esolver.reset(ModuleESolver::init_esolver(PARAM.inp, ucell)); - p_esolver->before_all_runners(ucell, PARAM.inp); - esolver_ready = true; - -#ifdef __RAPIDJSON - Json::gen_stru_wrapper(&ucell); -#endif - int io_failed = 0; std::string io_message; if (is_root()) @@ -267,7 +194,7 @@ void Driver::driver_ipi_run() try { const std::string address = ipi_address(); - GlobalV::ofs_running << " ABACUS socket driver connecting to i-PI endpoint " << address << std::endl; + ofs_running << " ABACUS socket driver connecting to i-PI endpoint " << address << std::endl; socket.connect(address); } catch (const std::exception& exc) @@ -276,12 +203,12 @@ void Driver::driver_ipi_run() io_message = exc.what(); } } - throw_if_root_io_failed(io_failed, io_message); + quit_if_root_io_failed(io_failed, io_message); bool isinit = false; bool hasdata = false; int istep = 0; - int nat_return = ucell.nat; + const int nat_return = ucell.nat; double energy_hartree = 0.0; std::vector forces_hartree_bohr(static_cast(3 * ucell.nat), 0.0); std::vector virial_hartree(9, 0.0); @@ -299,16 +226,28 @@ void Driver::driver_ipi_run() { header = socket.read_header(); } + catch (const IpiSocketClosed&) + { + header.clear(); + } catch (const std::exception& exc) { io_failed = 1; io_message = exc.what(); } } - throw_if_root_io_failed(io_failed, io_message); + quit_if_root_io_failed(io_failed, io_message); header = bcast_header(header); - if (header == "STATUS") + if (header.empty()) + { + if (is_root()) + { + ofs_running << " ABACUS socket driver exiting after peer closed connection" << std::endl; + } + break; + } + else if (header == "STATUS") { io_failed = 0; io_message.clear(); @@ -335,7 +274,7 @@ void Driver::driver_ipi_run() io_message = exc.what(); } } - throw_if_root_io_failed(io_failed, io_message); + quit_if_root_io_failed(io_failed, io_message); } else if (header == "INIT") { @@ -352,9 +291,10 @@ void Driver::driver_ipi_run() nbytes = socket.read_int(); if (nbytes < 0) { - throw std::runtime_error("negative INIT payload length from i-PI socket"); + io_failed = 1; + io_message = "negative INIT payload length from i-PI socket"; } - if (nbytes > 0) + else if (nbytes > 0) { params = socket.read_string(static_cast(nbytes)); } @@ -365,17 +305,17 @@ void Driver::driver_ipi_run() io_message = exc.what(); } } - throw_if_root_io_failed(io_failed, io_message); - bcast_int(rid); - bcast_int(nbytes); + quit_if_root_io_failed(io_failed, io_message); + Parallel_Common::bcast_int(rid); + Parallel_Common::bcast_int(nbytes); if (nbytes > 0 && is_root()) { - GlobalV::ofs_running << " ABACUS socket INIT params " << params << std::endl; + ofs_running << " ABACUS socket INIT params bytes " << nbytes << std::endl; } isinit = true; if (is_root()) { - GlobalV::ofs_running << " ABACUS socket INIT replica " << rid << std::endl; + ofs_running << " ABACUS socket INIT replica " << rid << std::endl; } } else if (header == "POSDATA") @@ -395,9 +335,13 @@ void Driver::driver_ipi_run() nat_socket = socket.read_int(); if (nat_socket < 0) { - throw std::runtime_error("negative POSDATA atom count from i-PI socket"); + io_failed = 1; + io_message = "negative POSDATA atom count from i-PI socket"; + } + else + { + positions = socket.read_doubles(static_cast(3 * nat_socket)); } - positions = socket.read_doubles(static_cast(3 * nat_socket)); } catch (const std::exception& exc) { @@ -405,10 +349,10 @@ void Driver::driver_ipi_run() io_message = exc.what(); } } - throw_if_root_io_failed(io_failed, io_message); + quit_if_root_io_failed(io_failed, io_message); bcast_double_vector(cell); bcast_double_vector(inv_cell); - bcast_int(nat_socket); + Parallel_Common::bcast_int(nat_socket); if (!is_root()) { positions.assign(static_cast(3 * nat_socket), 0.0); @@ -417,20 +361,27 @@ void Driver::driver_ipi_run() if (nat_socket != ucell.nat) { - throw std::runtime_error("POSDATA atom count does not match STRU."); + ModuleBase::WARNING_QUIT("ABACUS socket", "POSDATA atom count does not match STRU."); } const double max_cell_delta_bohr = max_abs_delta(cell, reference_cell); if (max_cell_delta_bohr > 1.0e-6) { - throw std::runtime_error("variable-cell socket updates are not supported yet."); + ModuleBase::WARNING_QUIT("ABACUS socket", "variable-cell socket updates are not supported yet."); } set_positions_from_ipi_bohr(ucell, positions); p_esolver->runner(ucell, istep); - runner_completed = true; - energy_hartree = p_esolver->cal_energy() * RY_TO_HARTREE; + const double energy_ry = p_esolver->cal_energy(); + energy_hartree = energy_ry * RY_TO_HARTREE; + if (is_root()) + { + ofs_running << " ABACUS socket return energy " + << energy_ry << " Ry, " + << energy_ry * ModuleBase::Ry_to_eV << " eV, " + << energy_hartree << " Ha" << std::endl; + } ModuleBase::matrix force; - if (PARAM.inp.cal_force) + if (inp.cal_force) { p_esolver->cal_force(ucell, force); forces_hartree_bohr = flatten_forces_hartree_per_bohr(force); @@ -459,7 +410,7 @@ void Driver::driver_ipi_run() io_message = exc.what(); } } - throw_if_root_io_failed(io_failed, io_message); + quit_if_root_io_failed(io_failed, io_message); isinit = false; hasdata = false; } @@ -467,7 +418,7 @@ void Driver::driver_ipi_run() { if (is_root()) { - GlobalV::ofs_running << " ABACUS socket driver exiting on header " << header << std::endl; + ofs_running << " ABACUS socket driver exiting on header " << header << std::endl; } break; } @@ -475,33 +426,13 @@ void Driver::driver_ipi_run() } catch (const std::exception& exc) { - pending_error = exc.what(); - if (is_root()) - { - GlobalV::ofs_running << " ABACUS socket driver ended with error: " << pending_error << std::endl; - } + ModuleBase::WARNING_QUIT("ABACUS socket", exc.what()); } if (is_root()) { socket.close(); } - if (esolver_ready && runner_completed && p_esolver) - { - p_esolver->after_all_runners(ucell); - } - p_esolver.reset(); - if (hardware_initialized) - { - this->finalize_hardware(); - } -#ifdef __RAPIDJSON - Json::create_Json(&ucell, PARAM); -#endif - - if (!pending_error.empty()) - { - ModuleBase::WARNING_QUIT("ABACUS socket", pending_error); - } + ModuleBase::timer::end("Socket_Driver", "socket_driver"); } diff --git a/source/source_relax/socket_driver.h b/source/source_relax/socket_driver.h new file mode 100644 index 00000000000..51b75289cb3 --- /dev/null +++ b/source/source_relax/socket_driver.h @@ -0,0 +1,22 @@ +#ifndef ABACUS_SOURCE_RELAX_SOCKET_DRIVER_H +#define ABACUS_SOURCE_RELAX_SOCKET_DRIVER_H + +#include "source_cell/unitcell.h" +#include "source_esolver/esolver.h" +#include "source_io/module_parameter/input_parameter.h" + +#include + +class Socket_Driver +{ + public: + Socket_Driver() = default; + ~Socket_Driver() = default; + + void socket_driver(ModuleESolver::ESolver* p_esolver, + UnitCell& ucell, + const Input_para& inp, + std::ofstream& ofs_running); +}; + +#endif diff --git a/source/source_relax/test/CMakeLists.txt b/source/source_relax/test/CMakeLists.txt index 2262fefb419..d1c867dc690 100644 --- a/source/source_relax/test/CMakeLists.txt +++ b/source/source_relax/test/CMakeLists.txt @@ -6,6 +6,12 @@ abacus_disable_feature_definitions(__ROCM) install(DIRECTORY support DESTINATION ${CMAKE_CURRENT_BINARY_DIR}) + +AddTest( + TARGET MODULE_RELAX_ipi_socket_test + SOURCES ipi_socket_test.cpp ../ipi_socket.cpp +) + AddTest( TARGET MODULE_RELAX_relax_new_line_search LIBS parameter diff --git a/source/source_relax/test/ipi_socket_test.cpp b/source/source_relax/test/ipi_socket_test.cpp new file mode 100644 index 00000000000..21db9be2b90 --- /dev/null +++ b/source/source_relax/test/ipi_socket_test.cpp @@ -0,0 +1,256 @@ +#include "../ipi_socket.h" + +#include "gtest/gtest.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace +{ +constexpr std::size_t IPI_HEADER_LEN = 12; + +std::string errno_message(const std::string& prefix) +{ + return prefix + ": " + std::strerror(errno); +} + +void send_all(int fd, const void* data, std::size_t nbytes) +{ + const char* cursor = static_cast(data); + std::size_t done = 0; + while (done < nbytes) + { + const ssize_t sent = ::send(fd, cursor + done, nbytes - done, 0); + if (sent < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("send failed")); + } + if (sent == 0) + { + throw std::runtime_error("send returned zero"); + } + done += static_cast(sent); + } +} + +void recv_all(int fd, void* data, std::size_t nbytes) +{ + char* cursor = static_cast(data); + std::size_t done = 0; + while (done < nbytes) + { + const ssize_t received = ::recv(fd, cursor + done, nbytes - done, 0); + if (received < 0) + { + if (errno == EINTR) + { + continue; + } + throw std::runtime_error(errno_message("recv failed")); + } + if (received == 0) + { + throw std::runtime_error("socket closed while receiving test data"); + } + done += static_cast(received); + } +} + +std::string padded_header(const std::string& header) +{ + std::string padded = header; + padded.resize(IPI_HEADER_LEN, ' '); + return padded; +} + +class UnixSocketServer +{ + public: + UnixSocketServer() + { + char dir_template[] = "/tmp/abacus_ipi_socket_test_XXXXXX"; + char* made_dir = ::mkdtemp(dir_template); + if (made_dir == nullptr) + { + throw std::runtime_error(errno_message("mkdtemp failed")); + } + dir_ = made_dir; + path_ = dir_ + "/ipi.sock"; + + listen_fd_ = ::socket(AF_UNIX, SOCK_STREAM, 0); + if (listen_fd_ < 0) + { + throw std::runtime_error(errno_message("socket failed")); + } + + sockaddr_un addr; + std::memset(&addr, 0, sizeof(addr)); + addr.sun_family = AF_UNIX; + std::strncpy(addr.sun_path, path_.c_str(), sizeof(addr.sun_path) - 1); + if (::bind(listen_fd_, reinterpret_cast(&addr), sizeof(addr)) != 0) + { + throw std::runtime_error(errno_message("bind failed")); + } + if (::listen(listen_fd_, 1) != 0) + { + throw std::runtime_error(errno_message("listen failed")); + } + } + + ~UnixSocketServer() + { + if (listen_fd_ >= 0) + { + ::close(listen_fd_); + } + if (!path_.empty()) + { + ::unlink(path_.c_str()); + } + if (!dir_.empty()) + { + ::rmdir(dir_.c_str()); + } + } + + UnixSocketServer(const UnixSocketServer&) = delete; + UnixSocketServer& operator=(const UnixSocketServer&) = delete; + + std::string address() const + { + return path_ + ":UNIX"; + } + + int accept_once() + { + const int fd = ::accept(listen_fd_, nullptr, nullptr); + if (fd < 0) + { + throw std::runtime_error(errno_message("accept failed")); + } + return fd; + } + + private: + int listen_fd_ = -1; + std::string dir_; + std::string path_; +}; + +void rethrow_thread_error(const std::exception_ptr& thread_error) +{ + if (thread_error) + { + std::rethrow_exception(thread_error); + } +} +} // namespace + +TEST(IpiSocketTest, WriteHeaderPadsToTwelveBytes) +{ + UnixSocketServer server; + std::string received; + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + char buffer[IPI_HEADER_LEN]; + recv_all(fd, buffer, sizeof(buffer)); + received.assign(buffer, sizeof(buffer)); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + socket.write_header("READY"); + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); + EXPECT_EQ(padded_header("READY"), received); +} + +TEST(IpiSocketTest, CleanPeerCloseBeforeNextHeaderThrowsDedicatedSignal) +{ + UnixSocketServer server; + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + const std::string header = padded_header("STATUS"); + send_all(fd, header.data(), header.size()); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + EXPECT_EQ("STATUS", socket.read_header()); + EXPECT_THROW(socket.read_header(), IpiSocketClosed); + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); +} + +TEST(IpiSocketTest, PartialHeaderCloseStaysRuntimeError) +{ + UnixSocketServer server; + std::exception_ptr thread_error; + std::thread peer([&]() { + try + { + const int fd = server.accept_once(); + const std::string partial = "STAT"; + send_all(fd, partial.data(), partial.size()); + ::close(fd); + } + catch (...) + { + thread_error = std::current_exception(); + } + }); + + IpiSocket socket; + socket.connect(server.address()); + try + { + static_cast(socket.read_header()); + FAIL() << "partial header EOF should throw"; + } + catch (const IpiSocketClosed&) + { + FAIL() << "partial header EOF must not be treated as clean peer close"; + } + catch (const std::runtime_error& exc) + { + EXPECT_NE(std::string::npos, std::string(exc.what()).find("closed while reading header")); + } + socket.close(); + + peer.join(); + rethrow_thread_error(thread_error); +} From fe2c2a3cc61995b3383aaf7f2a6abf5ab1882907 Mon Sep 17 00:00:00 2001 From: jianrui geng Date: Thu, 9 Jul 2026 11:12:06 +0800 Subject: [PATCH 3/9] Fix socket driver CI regressions --- source/source_main/driver_run.cpp | 15 ++-- source/source_relax/socket_driver.cpp | 99 +++++++++++++++++++++++++-- 2 files changed, 101 insertions(+), 13 deletions(-) diff --git a/source/source_main/driver_run.cpp b/source/source_main/driver_run.cpp index 3e7d669b6c2..408b507b902 100644 --- a/source/source_main/driver_run.cpp +++ b/source/source_main/driver_run.cpp @@ -38,9 +38,6 @@ void Driver::driver_run() { ModuleBase::TITLE("Driver", "driver_run"); - const std::string cal = PARAM.inp.calculation; - const bool socket_mode = (cal == "socket"); - //! 1: setup cell and atom information // this warning should not be here, mohan 2024-05-22 #ifndef __LCAO @@ -62,13 +59,17 @@ void Driver::driver_run() unitcell::check_atomic_stru(ucell, PARAM.inp.min_dist_coef); //! 2: initialize the ESolver (depends on a set-up ucell after `setup_cell`) - Input_para esolver_inp = PARAM.inp; - Input_para driver_inp = PARAM.inp; + const std::string cal = PARAM.inp.calculation; + const bool socket_mode = (cal == "socket"); + Input_para socket_esolver_inp = PARAM.inp; + Input_para socket_driver_inp = PARAM.inp; if (socket_mode) { - esolver_inp.calculation = "scf"; - driver_inp.calculation = cal; + socket_esolver_inp.calculation = "scf"; + socket_driver_inp.calculation = cal; } + const Input_para& esolver_inp = socket_mode ? socket_esolver_inp : PARAM.inp; + const Input_para& driver_inp = socket_mode ? socket_driver_inp : PARAM.inp; this->init_hardware(); diff --git a/source/source_relax/socket_driver.cpp b/source/source_relax/socket_driver.cpp index 02928fec9d3..bb0a6757f50 100644 --- a/source/source_relax/socket_driver.cpp +++ b/source/source_relax/socket_driver.cpp @@ -27,29 +27,52 @@ bool is_root() void bcast_double_vector(std::vector& values) { +#ifdef __MPI if (!values.empty()) { Parallel_Common::bcast_double(values.data(), static_cast(values.size())); } +#else + (void)values; +#endif +} + +void bcast_socket_int(int& value) +{ +#ifdef __MPI + Parallel_Common::bcast_int(value); +#else + (void)value; +#endif +} + +void bcast_socket_chars(char* value, const int size) +{ +#ifdef __MPI + Parallel_Common::bcast_char(value, size); +#else + (void)value; + (void)size; +#endif } void bcast_socket_string(std::string& value) { int size = static_cast(value.size()); - Parallel_Common::bcast_int(size); + bcast_socket_int(size); if (!is_root()) { value.resize(static_cast(size)); } if (size > 0) { - Parallel_Common::bcast_char(&value[0], size); + bcast_socket_chars(&value[0], size); } } void quit_if_root_io_failed(int root_failed, std::string root_message) { - Parallel_Common::bcast_int(root_failed); + bcast_socket_int(root_failed); bcast_socket_string(root_message); if (root_failed != 0) { @@ -85,6 +108,58 @@ std::vector ipi_cell_bohr_from_unitcell(const UnitCell& ucell) }; } +double max_wrapped_direct_delta_from_unitcell(const UnitCell& ucell, const std::vector& positions_bohr) +{ + if (positions_bohr.size() != static_cast(3 * ucell.nat)) + { + return 1.0e99; + } + + double out = 0.0; + int iat = 0; + for (int it = 0; it < ucell.ntype; ++it) + { + const Atom* atom = &ucell.atoms[it]; + for (int ia = 0; ia < atom->na; ++ia) + { + const double tau_x = positions_bohr[3 * iat + 0] / ucell.lat0; + const double tau_y = positions_bohr[3 * iat + 1] / ucell.lat0; + const double tau_z = positions_bohr[3 * iat + 2] / ucell.lat0; + + double dx = 0.0; + double dy = 0.0; + double dz = 0.0; + ModuleBase::Mathzone::Cartesian_to_Direct(tau_x, + tau_y, + tau_z, + ucell.latvec.e11, + ucell.latvec.e12, + ucell.latvec.e13, + ucell.latvec.e21, + ucell.latvec.e22, + ucell.latvec.e23, + ucell.latvec.e31, + ucell.latvec.e32, + ucell.latvec.e33, + dx, + dy, + dz); + + double ddx = dx - atom->taud[ia].x; + double ddy = dy - atom->taud[ia].y; + double ddz = dz - atom->taud[ia].z; + ddx -= std::round(ddx); + ddy -= std::round(ddy); + ddz -= std::round(ddz); + out = std::max(out, std::abs(ddx)); + out = std::max(out, std::abs(ddy)); + out = std::max(out, std::abs(ddz)); + ++iat; + } + } + return out; +} + double max_abs_delta(const std::vector& a, const std::vector& b) { if (a.size() != b.size()) @@ -214,6 +289,7 @@ void Socket_Driver::socket_driver(ModuleESolver::ESolver* p_esolver, std::vector virial_hartree(9, 0.0); const std::vector reference_cell = ipi_cell_bohr_from_unitcell(ucell); + bool checked_initial_positions = false; while (true) { @@ -306,8 +382,8 @@ void Socket_Driver::socket_driver(ModuleESolver::ESolver* p_esolver, } } quit_if_root_io_failed(io_failed, io_message); - Parallel_Common::bcast_int(rid); - Parallel_Common::bcast_int(nbytes); + bcast_socket_int(rid); + bcast_socket_int(nbytes); if (nbytes > 0 && is_root()) { ofs_running << " ABACUS socket INIT params bytes " << nbytes << std::endl; @@ -352,7 +428,7 @@ void Socket_Driver::socket_driver(ModuleESolver::ESolver* p_esolver, quit_if_root_io_failed(io_failed, io_message); bcast_double_vector(cell); bcast_double_vector(inv_cell); - Parallel_Common::bcast_int(nat_socket); + bcast_socket_int(nat_socket); if (!is_root()) { positions.assign(static_cast(3 * nat_socket), 0.0); @@ -368,6 +444,17 @@ void Socket_Driver::socket_driver(ModuleESolver::ESolver* p_esolver, { ModuleBase::WARNING_QUIT("ABACUS socket", "variable-cell socket updates are not supported yet."); } + if (!checked_initial_positions) + { + checked_initial_positions = true; + if (max_wrapped_direct_delta_from_unitcell(ucell, positions) > 1.0e-5 && is_root()) + { + ModuleBase::WARNING( + "ABACUS socket", + "first POSDATA positions are not PBC-equivalent to STRU atom order; " + "i-PI POSDATA carries no species, so the client atoms should use the same atom order as STRU."); + } + } set_positions_from_ipi_bohr(ucell, positions); p_esolver->runner(ucell, istep); From 36353412e892d8d3a955d58f057285e886d1659d Mon Sep 17 00:00:00 2001 From: jianrui geng Date: Thu, 9 Jul 2026 11:23:16 +0800 Subject: [PATCH 4/9] Reduce socket driver global dependencies --- source/source_main/driver_run.cpp | 9 +++------ source/source_relax/socket_driver.cpp | 12 ++++++++++-- source/source_relax/socket_driver.h | 12 ++++++++---- 3 files changed, 21 insertions(+), 12 deletions(-) diff --git a/source/source_main/driver_run.cpp b/source/source_main/driver_run.cpp index 408b507b902..e565871d604 100644 --- a/source/source_main/driver_run.cpp +++ b/source/source_main/driver_run.cpp @@ -59,17 +59,13 @@ void Driver::driver_run() unitcell::check_atomic_stru(ucell, PARAM.inp.min_dist_coef); //! 2: initialize the ESolver (depends on a set-up ucell after `setup_cell`) - const std::string cal = PARAM.inp.calculation; - const bool socket_mode = (cal == "socket"); Input_para socket_esolver_inp = PARAM.inp; - Input_para socket_driver_inp = PARAM.inp; + const bool socket_mode = (socket_esolver_inp.calculation == "socket"); if (socket_mode) { socket_esolver_inp.calculation = "scf"; - socket_driver_inp.calculation = cal; } const Input_para& esolver_inp = socket_mode ? socket_esolver_inp : PARAM.inp; - const Input_para& driver_inp = socket_mode ? socket_driver_inp : PARAM.inp; this->init_hardware(); @@ -84,6 +80,7 @@ void Driver::driver_run() #endif //! 4: different types of calculations + const std::string cal = PARAM.inp.calculation; if (cal == "md") { Run_MD::md_line(ucell, p_esolver, PARAM); @@ -91,7 +88,7 @@ void Driver::driver_run() else if (cal == "scf" || cal == "relax" || cal == "cell-relax" || cal == "nscf" || cal == "socket") { Relax_Driver rl_driver; - rl_driver.relax_driver(p_esolver, ucell, driver_inp, GlobalV::ofs_running); + rl_driver.relax_driver(p_esolver, ucell, PARAM.inp, GlobalV::ofs_running); } else if (cal == "get_s") { diff --git a/source/source_relax/socket_driver.cpp b/source/source_relax/socket_driver.cpp index bb0a6757f50..ecd4d1be27e 100644 --- a/source/source_relax/socket_driver.cpp +++ b/source/source_relax/socket_driver.cpp @@ -2,11 +2,13 @@ #include "source_relax/ipi_socket.h" #include "source_base/global_function.h" -#include "source_base/global_variable.h" #include "source_base/mathzone.h" #include "source_base/parallel_common.h" #include "source_base/timer.h" +#include "source_cell/unitcell.h" #include "source_cell/update_cell.h" +#include "source_esolver/esolver.h" +#include "source_io/module_parameter/input_parameter.h" #include #include @@ -22,7 +24,13 @@ constexpr int IPI_RANK_ROOT = 0; bool is_root() { - return GlobalV::MY_RANK == IPI_RANK_ROOT; +#ifdef __MPI + int rank = IPI_RANK_ROOT; + MPI_Comm_rank(MPI_COMM_WORLD, &rank); + return rank == IPI_RANK_ROOT; +#else + return true; +#endif } void bcast_double_vector(std::vector& values) diff --git a/source/source_relax/socket_driver.h b/source/source_relax/socket_driver.h index 51b75289cb3..86c180fc42e 100644 --- a/source/source_relax/socket_driver.h +++ b/source/source_relax/socket_driver.h @@ -1,11 +1,15 @@ #ifndef ABACUS_SOURCE_RELAX_SOCKET_DRIVER_H #define ABACUS_SOURCE_RELAX_SOCKET_DRIVER_H -#include "source_cell/unitcell.h" -#include "source_esolver/esolver.h" -#include "source_io/module_parameter/input_parameter.h" +#include -#include +class UnitCell; +struct Input_para; + +namespace ModuleESolver +{ +class ESolver; +} class Socket_Driver { From 5711ab3ceccf6b7df036ca5cf6a932de38eb22a3 Mon Sep 17 00:00:00 2001 From: jianrui geng Date: Sun, 19 Jul 2026 13:20:52 +0800 Subject: [PATCH 5/9] Address socket i-PI review feedback --- docs/advanced/input_files/input-main.md | 13 +- docs/advanced/interface/ase.md | 58 +++++ docs/parameters.yaml | 13 +- interfaces/ASE_interface/README.md | 1 + interfaces/ASE_interface/abacuslite/core.py | 216 ++++++++++++++++++ interfaces/ASE_interface/examples/socketio.py | 157 +++++++++++++ .../module_parameter/input_parameter.h | 1 + .../read_input_item_system.cpp | 30 ++- .../test_serial/read_input_item_test.cpp | 33 ++- source/source_main/driver_run.cpp | 14 +- source/source_relax/CMakeLists.txt | 2 +- source/source_relax/relax_driver.cpp | 2 +- source/source_relax/socket_driver.cpp | 10 +- .../{ipi_socket.cpp => socket_ipi.cpp} | 6 +- .../{ipi_socket.h => socket_ipi.h} | 4 +- source/source_relax/test/CMakeLists.txt | 4 +- ...pi_socket_test.cpp => socket_ipi_test.cpp} | 4 +- 17 files changed, 531 insertions(+), 37 deletions(-) create mode 100644 interfaces/ASE_interface/examples/socketio.py rename source/source_relax/{ipi_socket.cpp => socket_ipi.cpp} (97%) rename source/source_relax/{ipi_socket.h => socket_ipi.h} (94%) rename source/source_relax/test/{ipi_socket_test.cpp => socket_ipi_test.cpp} (98%) diff --git a/docs/advanced/input_files/input-main.md b/docs/advanced/input_files/input-main.md index 529e2657edb..ee1fb6c5b32 100644 --- a/docs/advanced/input_files/input-main.md +++ b/docs/advanced/input_files/input-main.md @@ -9,6 +9,7 @@ - [suffix](#suffix) - [ntype](#ntype) - [calculation](#calculation) + - [socket\_driver](#socket_driver) - [esolver\_type](#esolver_type) - [symmetry](#symmetry) - [symmetry\_prec](#symmetry_prec) @@ -583,7 +584,6 @@ - relax: perform structure relaxation calculations, the relax_nmax parameter depicts the maximal number of ionic iterations - cell-relax: perform cell relaxation calculations - md: perform molecular dynamics simulations - - socket: run as a socket client for external drivers using the i-PI protocol - get_pchg: obtain partial (band-decomposed) charge densities (for LCAO basis only). See out_pchg for more information - get_wf: obtain real space wave functions (for LCAO basis only). See out_wfc_norm and out_wfc_re_im for more information - get_s: obtain the overlap matrix formed by localized orbitals (for LCAO basis with multiple k points). the file name is SR.csr with file format being the same as that generated by out_mat_hs2 @@ -593,6 +593,15 @@ - test_neighbour: obtain information of neighboring atoms (for LCAO basis only), please specify a positive search_radius manually - **Default**: scf +### socket_driver + +- **Type**: Boolean +- **Availability**: *calculation==scf* +- **Description**: If set to True, ABACUS keeps the calculation type as scf and receives atomic positions from an external driver through the i-PI socket protocol. + + > Note: Use calculation = scf with socket_driver = True. Set ABACUS_SOCKET_ADDRESS to host:port or path:UNIX to choose the socket endpoint; if unset, localhost:31415 is used. +- **Default**: False + ### esolver_type - **Type**: String @@ -844,7 +853,7 @@ - **Description**: Charge extrapolation method for MD, relaxation, and socket-driven calculations. When set to default, ABACUS chooses second-order for md, first-order for - relax/cell-relax/socket, and atomic for other calculations. Socket-driven + relax/cell-relax and socket_driver calculations, and atomic for other calculations. Socket-driven molecular dynamics can explicitly set second-order if the external driver updates structures smoothly enough for second-order extrapolation. - **Default**: default diff --git a/docs/advanced/interface/ase.md b/docs/advanced/interface/ase.md index e9b7b062809..82289cbac33 100644 --- a/docs/advanced/interface/ase.md +++ b/docs/advanced/interface/ase.md @@ -103,6 +103,64 @@ In the new implementation, we limit the range of functionalties supported to mai Please read the examples in `interfaces/ASE_interface/examples/` for more details. +### Socket I/O with ASE + +For socket-driven ASE workflows, use the `AbacusSocketIO` calculator. ASE runs the i-PI socket server, while ABACUS keeps `calculation=scf` and is launched with `socket_driver=1` as the client. The protocol is simple: ASE sends atomic positions and cell data to ABACUS; ABACUS evaluates one SCF step for that structure and returns energy, forces, and virial. See the [ASE socket I/O documentation](https://ase-lib.org/ase/calculators/socketio/socketio.html) and the i-PI reference paper, [Ceriotti et al., Comput. Phys. Commun. 185, 1019-1026 (2014)](https://doi.org/10.1016/j.cpc.2013.10.027), for the protocol background. + +Build ABACUS as usual before using this interface. PW-only builds work with `basis_type=pw`; LCAO socket calculations require an LCAO-enabled executable. No extra socket library is required. + +With CMake, choose the executable according to the basis: + +```bash +cmake -S . -B build-pw -DENABLE_MPI=ON -DENABLE_LCAO=OFF +cmake --build build-pw --target abacus_pw_para -j + +cmake -S . -B build-lcao -DENABLE_MPI=ON -DENABLE_LCAO=ON +cmake --build build-lcao --target abacus_basic_para -j +``` + +With the ABACUS toolchain workflow, build the normal ABACUS executable with LCAO support when `basis_type=lcao` is needed, then pass that executable to `AbacusProfile(command=...)`. The ASE interface can be installed from this repository with: + +```bash +cd interfaces/ASE_interface +pip install . +``` + +A minimal socket calculator setup is: + +```python +from ase.optimize import BFGS +from abacuslite import AbacusProfile, AbacusSocketIO + +aprof = AbacusProfile( + command="mpirun -np 4 /path/to/abacus", + pseudo_dir="/path/to/pseudopotentials", + orbital_dir="/path/to/orbitals", + omp_num_threads=1, +) + +abacus = AbacusSocketIO( + profile=aprof, + directory="socketio", + unixsocket="abacus_si", + pseudopotentials={"Si": "Si_ONCV_PBE-1.0.upf"}, + basissets={"Si": "Si_gga_8au_100Ry_2s2p1d.orb"}, + inp={"calculation": "scf", "basis_type": "lcao", "kspacing": 0.1}, +) + +with abacus as calc: + atoms.calc = calc + BFGS(atoms).run(fmax=0.05) +``` + +`AbacusSocketIO` sets `socket_driver=1` and `cal_force=1` automatically. The socket endpoint is selected by the calculator arguments and passed to ABACUS through `ABACUS_SOCKET_ADDRESS`, for example `localhost:31415` or `/tmp/ipi_abacus_si:UNIX`. Calling `atoms.get_potential_energy()` is supported, but the ABACUS client still computes forces because the i-PI `GETFORCE` exchange returns energy, forces, and virial as one response. + +A socket calculator owns one ABACUS process initialized from one fixed `INPUT`/`STRU` setup. Reuse the same `AbacusSocketIO` instance only for position updates under the same electronic-structure settings and the same cell. Do not change `kpts`, `kspacing`, `nspin`, `basis_type`, `basissets`, pseudopotentials, species, atom count, cell, or other core `INPUT`/`STRU` parameters through an existing socket calculator; create a new `AbacusSocketIO` instance and a new ABACUS client process for those changes. `AbacusSocketIO` rejects cell changes before sending them to ABACUS, and the ABACUS socket driver also checks incoming POSDATA cells against the initial `STRU` cell and exits if they differ. + +The i-PI protocol does not transmit element symbols. `AbacusSocketIO` therefore sorts the internal socket atoms with the same first-occurrence species grouping used when writing `STRU`, and maps returned forces back to the original ASE `Atoms` order. This avoids silent force/atom mismatches when structures are read from CIF, extxyz, POSCAR, or other formats whose atom order is not already grouped for ABACUS. Users should not manually reorder atoms for socket I/O; pass the physical ASE `Atoms` object directly to the calculator. + +A complete fixed-cell validation and benchmark example is available in `interfaces/ASE_interface/examples/socketio.py`. + ## SPAP Analysis [SPAP](https://github.com/chuanxun/StructurePrototypeAnalysisPackage) (Structure Prototype Analysis Package) is written by Dr. Chuanxun Su to analyze symmetry and compare similarity of large amount of atomic structures. The coordination characterization function (CCF) is used to diff --git a/docs/parameters.yaml b/docs/parameters.yaml index 2298e10d111..9c9a250707b 100644 --- a/docs/parameters.yaml +++ b/docs/parameters.yaml @@ -28,7 +28,6 @@ parameters: * relax: perform structure relaxation calculations, the relax_nmax parameter depicts the maximal number of ionic iterations * cell-relax: perform cell relaxation calculations * md: perform molecular dynamics simulations - * socket: run as a socket client for external drivers using the i-PI protocol * get_pchg: obtain partial (band-decomposed) charge densities (for LCAO basis only). See out_pchg for more information * get_wf: obtain real space wave functions (for LCAO basis only). See out_wfc_norm and out_wfc_re_im for more information * get_s: obtain the overlap matrix formed by localized orbitals (for LCAO basis with multiple k points). the file name is SR.csr with file format being the same as that generated by out_mat_hs2 @@ -39,6 +38,16 @@ parameters: default_value: scf unit: "" availability: "" + - name: socket_driver + category: System variables + type: Boolean + description: | + If set to True, ABACUS keeps the calculation type as scf and receives atomic positions from an external driver through the i-PI socket protocol. + + [NOTE] Use calculation = scf with socket_driver = True. Set ABACUS_SOCKET_ADDRESS to host:port or path:UNIX to choose the socket endpoint; if unset, localhost:31415 is used. + default_value: "False" + unit: "" + availability: calculation==scf - name: esolver_type category: System variables type: String @@ -333,7 +342,7 @@ parameters: Charge extrapolation method for MD, relaxation, and socket-driven calculations. When set to default, ABACUS chooses second-order for md, first-order for - relax/cell-relax/socket, and atomic for other calculations. Socket-driven + relax/cell-relax and socket_driver calculations, and atomic for other calculations. Socket-driven molecular dynamics can explicitly set second-order if the external driver updates structures smoothly enough for second-order extrapolation. default_value: default diff --git a/interfaces/ASE_interface/README.md b/interfaces/ASE_interface/README.md index 824e69900de..a76e68e177e 100644 --- a/interfaces/ASE_interface/README.md +++ b/interfaces/ASE_interface/README.md @@ -32,6 +32,7 @@ Please refer to the example scripts in the `examples` folder. Recommended learni 7. **constraintmd.py** - Constrained molecular dynamics simulation 8. **metadynamics.py** - Metadynamics simulation 9. **neb.py** - Nudged Elastic Band (NEB) calculation +10. **socketio.py** - ASE optimization with the AbacusSocketIO calculator running ABACUS as an i-PI socket client More usage examples will be provided in future versions. diff --git a/interfaces/ASE_interface/abacuslite/core.py b/interfaces/ASE_interface/abacuslite/core.py index a1aad1e773b..e2c00d4c889 100644 --- a/interfaces/ASE_interface/abacuslite/core.py +++ b/interfaces/ASE_interface/abacuslite/core.py @@ -45,6 +45,7 @@ GenericFileIOCalculator, read_stdout ) +from ase.calculators.socketio import SocketIOCalculator from ase.atoms import Atoms from ase.dft.kpoints import BandPath from ase.io import read @@ -135,6 +136,21 @@ def get_calculator_command(self, inputfile) -> List[str]: # additional inputfile argument is not used. return [] + def socketio_argv_inet(self, port: Optional[int] = None) -> List[str]: + port = 31415 if port is None else port + return [ + 'env', + f'ABACUS_SOCKET_ADDRESS=localhost:{port}', + *self._split_command, + ] + + def socketio_argv_unix(self, socket: str) -> List[str]: + return [ + 'env', + f'ABACUS_SOCKET_ADDRESS=/tmp/ipi_{socket}:UNIX', + *self._split_command, + ] + def version(self) -> str: '''get the abacus version information''' cmd_ = [*self._split_command, '--version'] @@ -443,6 +459,17 @@ def __init__(self, directory=directory, ) + def write_input(self, atoms, properties=None, system_changes=None): + if properties is None: + properties = self.template.implemented_properties + self.template.write_input( + profile=self.profile, + directory=Path(self.directory), + atoms=atoms, + parameters=self.parameters, + properties=properties, + ) + @classmethod def restart(cls, profile=None, directory='.', **kwargs): '''instantiate one ABACUS calculator from an existing job directory, @@ -557,11 +584,200 @@ def band_structure(self, efermi=None): from ase.spectrum.band_structure import get_band_structure return get_band_structure(calc=self, reference=efermi) +class AbacusSocketIO(SocketIOCalculator): + """ASE socket I/O calculator that launches ABACUS as an i-PI client. + + A socket calculator owns one ABACUS process with one fixed INPUT/STRU + setup. The i-PI protocol can update positions, but electronic-structure + parameters such as k-points, spin, basis, pseudopotentials, and species + require a new calculator instance. Energy-only ASE calls are accepted, but + ABACUS still computes forces because i-PI GETFORCE returns energy, forces, + and virial together. + """ + + def __init__(self, + profile=None, + directory='.', + port=None, + unixsocket=None, + timeout=None, + log=None, + **kwargs): + inp = self._socket_inp(kwargs.pop('inp', {})) + self.abacus = Abacus( + profile=profile, + directory=directory, + inp=inp, + **kwargs, + ) + self._reference_cell = None + super().__init__( + port=port, + unixsocket=unixsocket, + timeout=timeout, + log=log, + launch_client=self._launch_client, + ) + + def calculate(self, atoms=None, properties=['energy'], system_changes=None): + from ase.calculators.calculator import ( + PropertyNotImplementedError, + all_changes, + ) + from ase.stress import full_3x3_to_voigt_6_stress + + if system_changes is None: + system_changes = all_changes + if atoms is None: + atoms = self.atoms + if atoms is None: + raise ValueError('AbacusSocketIO.calculate requires atoms') + + bad = [change for change in system_changes + if change not in self.supported_changes] + if self.atoms is not None and any(bad): + raise PropertyNotImplementedError( + 'Cannot change {} through IPI protocol. ' + 'Please create new socket calculator.' + .format(bad if len(bad) > 1 else bad[0])) + + self._check_fixed_cell(atoms) + order = self._socket_sort_indices(atoms) + socket_atoms = atoms[order] + self.atoms = atoms.copy() + + if self.server is None: + self.server = self.launch_server() + proc = self.launch_client(socket_atoms, properties, + port=self._port, + unixsocket=self._unixsocket) + self.server.proc = proc + + results = self.server.calculate(socket_atoms) + results['free_energy'] = results['energy'] + virial = results.pop('virial') + if self.atoms.cell.rank == 3 and any(self.atoms.pbc): + vol = atoms.get_volume() + results['stress'] = -full_3x3_to_voigt_6_stress(virial) / vol + if 'forces' in results: + results['forces'] = self._forces_to_input_order( + results['forces'], order) + self.results.update(results) + + def _check_fixed_cell(self, atoms): + from ase.calculators.calculator import PropertyNotImplementedError + + cell = atoms.cell.array.copy() + if self._reference_cell is None: + self._reference_cell = cell + return + max_delta = np.max(np.abs(cell - self._reference_cell)) + if max_delta > 1.0e-10: + raise PropertyNotImplementedError( + 'AbacusSocketIO is fixed-cell only; create a new socket ' + 'calculator for a changed cell, or use the normal Abacus ' + 'FileIO calculator for variable-cell workflows.' + ) + + def set(self, **kwargs): + if kwargs: + raise ValueError( + 'AbacusSocketIO input parameters are fixed after construction; ' + 'create a new AbacusSocketIO calculator to change k-points, ' + 'spin, basis, pseudopotentials, species, or other INPUT/STRU ' + 'settings.' + ) + return super().set(**kwargs) + + def _launch_client(self, atoms, properties=None, port=None, unixsocket=None): + from subprocess import Popen + + if properties is None: + properties = self.abacus.template.implemented_properties + + directory = Path(self.abacus.directory) + directory.mkdir(exist_ok=True, parents=True) + + if hasattr(self.abacus, 'write_inputfiles'): + self.abacus.write_inputfiles(atoms, properties) + else: + self.abacus.write_input(atoms, properties=properties) + + if unixsocket is not None: + argv = self.abacus.profile.socketio_argv_unix(socket=unixsocket) + else: + argv = self.abacus.profile.socketio_argv_inet(port=port) + + stdout = open(directory / self.abacus.template.outputname, 'w') + stderr = open(directory / self.abacus.template.errorname, 'w') + try: + return Popen(argv, cwd=directory, env=os.environ, + stdout=stdout, stderr=stderr) + finally: + stdout.close() + stderr.close() + + @staticmethod + def _socket_inp(inp): + inp = dict(inp) + calculation = inp.get('calculation', 'scf') + if calculation != 'scf': + raise ValueError('ABACUS socket I/O requires calculation="scf"') + inp.update({ + 'calculation': 'scf', + 'socket_driver': 1, + 'cal_force': 1, + }) + return inp + + @staticmethod + def _socket_sort_indices(atoms): + return species_group_indices(atoms.get_chemical_symbols()) + + @staticmethod + def _forces_to_input_order(forces, order): + reordered = np.empty_like(forces) + for sorted_index, original_index in enumerate(order): + reordered[original_index] = forces[sorted_index] + return reordered + + class TestAbacusCalculator(unittest.TestCase): here = Path(__file__).parent pporb = here.parent.parent.parent / 'tests' / 'PP_ORB' + def test_socketio_species_order_mapping(self): + atoms = Atoms(symbols=['Si', 'O', 'C', 'Si', 'O', 'C']) + order = AbacusSocketIO._socket_sort_indices(atoms) + self.assertEqual(order, [0, 3, 1, 4, 2, 5]) + + socket_forces = np.arange(18).reshape(6, 3) + input_forces = AbacusSocketIO._forces_to_input_order( + socket_forces, order) + + expected = np.empty_like(socket_forces) + for sorted_index, original_index in enumerate(order): + expected[original_index] = socket_forces[sorted_index] + np.testing.assert_array_equal(input_forces, expected) + + def test_socketio_rejects_parameter_changes(self): + calc = object.__new__(AbacusSocketIO) + with self.assertRaisesRegex(ValueError, 'fixed after construction'): + calc.set(kpts={'mode': 'mp-sampling', 'nk': [2, 2, 2]}) + + def test_socketio_rejects_cell_changes(self): + from ase.calculators.calculator import PropertyNotImplementedError + + calc = object.__new__(AbacusSocketIO) + calc.atoms = Atoms('Si', cell=[5.0, 5.0, 5.0], pbc=True) + calc._reference_cell = calc.atoms.cell.array.copy() + + changed = calc.atoms.copy() + changed.cell[0, 0] = 5.1 + with self.assertRaisesRegex(PropertyNotImplementedError, 'fixed-cell'): + calc._check_fixed_cell(changed) + def test_calculator_results(self): from ase.build.bulk import bulk silicon = bulk('Si', crystalstructure='diamond', a=5.43) diff --git a/interfaces/ASE_interface/examples/socketio.py b/interfaces/ASE_interface/examples/socketio.py new file mode 100644 index 00000000000..e3087453ffc --- /dev/null +++ b/interfaces/ASE_interface/examples/socketio.py @@ -0,0 +1,157 @@ +""" +This example validates and benchmarks ABACUS socket I/O from ASE. + +ASE runs as the i-PI socket server and ABACUS runs as the socket client. +ABACUS keeps calculation=scf and enables socket_driver internally. + +The script checks two PR-review relevant points: +1. socket SCF gives the same energy and forces as a normal non-socket SCF; +2. repeated socket calculations avoid relaunching ABACUS and are faster than + the normal FileIO calculator for a sequence of SCF force evaluations. + +The i-PI protocol does not carry element symbols. AbacusSocketIO handles +the required STRU/socket atom-order alignment internally and returns forces +in the original ASE Atoms order. +""" +import os +import shutil +import time +from pathlib import Path + +import numpy as np +from ase import Atoms +from abacuslite import Abacus, AbacusProfile, AbacusSocketIO + +here = Path(__file__).parent +pporb = here.parent.parent.parent / 'tests' / 'PP_ORB' + +aprof = AbacusProfile( + command=os.environ.get('ABACUS_COMMAND', 'mpirun -np 4 abacus'), + pseudo_dir=pporb, + orbital_dir=pporb, + omp_num_threads=1, +) + +common_kwargs = { + 'pseudopotentials': {'Si': 'Si_ONCV_PBE-1.0.upf'}, + 'basissets': {'Si': 'Si_gga_8au_100Ry_2s2p1d.orb'}, + 'inp': { + 'calculation': 'scf', + 'nspin': 1, + 'basis_type': 'lcao', + 'ks_solver': 'scalapack_gvx', + 'ecutwfc': 30, + 'symmetry': 0, + 'kspacing': 0.5, + 'scf_thr': 1e-8, + 'scf_nmax': 40, + 'chg_extrap': 'atomic', + 'cal_force': 1, + }, +} + +base_atoms = Atoms( + 'Si2', + positions=[[0.0, 0.0, 0.0], [1.25, 1.25, 1.25]], + cell=[5.43, 5.43, 5.43], + pbc=True, +) + + +def clean(directory): + shutil.rmtree(directory, ignore_errors=True) + + +def run_fileio(atoms, directory): + clean(directory) + calc = Abacus(profile=aprof, directory=str(directory), **common_kwargs) + atoms = atoms.copy() + atoms.calc = calc + forces = atoms.get_forces() + energy = atoms.get_potential_energy() + return energy, forces + + +def run_socketio(atoms, directory, socket_name): + clean(directory) + calc = AbacusSocketIO( + profile=aprof, + directory=str(directory), + unixsocket=socket_name, + timeout=120, + **common_kwargs, + ) + atoms = atoms.copy() + with calc: + atoms.calc = calc + energy = atoms.get_potential_energy() + forces = atoms.get_forces() + return energy, forces + + +def displaced_structures(): + structures = [] + for scale in (0.00, 0.03, -0.02, 0.05): + atoms = base_atoms.copy() + atoms.positions[1] += scale + structures.append(atoms) + return structures + + +fileio_dir = here / 'socketio_fileio_scf' +socket_dir = here / 'socketio_socket_scf' +bench_fileio_dir = here / 'socketio_bench_fileio' +bench_socket_dir = here / 'socketio_bench_socket' + +try: + reference_energy, reference_forces = run_fileio(base_atoms, fileio_dir) + socket_energy, socket_forces = run_socketio( + base_atoms, socket_dir, 'abacus_si_check') + + energy_diff = abs(socket_energy - reference_energy) + force_diff = np.max(np.abs(socket_forces - reference_forces)) + print(f'FileIO SCF energy: {reference_energy:.12f} eV') + print(f'Socket SCF energy: {socket_energy:.12f} eV') + print(f'|dE|: {energy_diff:.3e} eV') + print(f'max |dF|: {force_diff:.3e} eV/Angstrom') + assert energy_diff < 1e-4 + assert force_diff < 1e-5 + + structures = displaced_structures() + + clean(bench_fileio_dir) + fileio_calc = Abacus( + profile=aprof, + directory=str(bench_fileio_dir), + **common_kwargs, + ) + t0 = time.perf_counter() + for atoms in structures: + atoms = atoms.copy() + atoms.calc = fileio_calc + atoms.get_forces() + fileio_seconds = time.perf_counter() - t0 + + clean(bench_socket_dir) + socket_calc = AbacusSocketIO( + profile=aprof, + directory=str(bench_socket_dir), + unixsocket='abacus_si_bench', + timeout=120, + **common_kwargs, + ) + t0 = time.perf_counter() + with socket_calc as calc: + for atoms in structures: + atoms = atoms.copy() + atoms.calc = calc + atoms.get_forces() + socket_seconds = time.perf_counter() - t0 + + speedup = fileio_seconds / socket_seconds + print(f'FileIO repeated SCF force time: {fileio_seconds:.2f} s') + print(f'Socket repeated SCF force time: {socket_seconds:.2f} s') + print(f'Socket speedup vs FileIO: {speedup:.2f}x') +finally: + for directory in (fileio_dir, socket_dir, bench_fileio_dir, bench_socket_dir): + clean(directory) diff --git a/source/source_io/module_parameter/input_parameter.h b/source/source_io/module_parameter/input_parameter.h index 3045d663361..c7cb990da21 100644 --- a/source/source_io/module_parameter/input_parameter.h +++ b/source/source_io/module_parameter/input_parameter.h @@ -19,6 +19,7 @@ struct Input_para std::string calculation = "scf"; ///< "scf" : self consistent calculation. ///< "nscf" : non-self consistent calculation. ///< "relax" : cell relaxations + bool socket_driver = false; ///< run ABACUS as an i-PI socket client std::string esolver_type = "ksdft"; ///< the energy solver: ksdft, sdft, ofdft, tddft, lj, dp /* symmetry level: -1, no symmetry at all; diff --git a/source/source_io/module_parameter/read_input_item_system.cpp b/source/source_io/module_parameter/read_input_item_system.cpp index 6ed572b21f9..bc2d1c94d4f 100644 --- a/source/source_io/module_parameter/read_input_item_system.cpp +++ b/source/source_io/module_parameter/read_input_item_system.cpp @@ -75,7 +75,7 @@ void ReadInput::item_system() } { Input_Item item("calculation"); - item.annotation = "scf; relax; md; socket; cell-relax; nscf; get_s; get_wf; get_pchg; gen_bessel; gen_opt_abfs; test_memory; test_neighbour"; + item.annotation = "scf; relax; md; cell-relax; nscf; get_s; get_wf; get_pchg; gen_bessel; gen_opt_abfs; test_memory; test_neighbour"; item.category = "System variables"; item.type = "String"; item.description = R"(Specify the type of calculation. @@ -85,7 +85,6 @@ void ReadInput::item_system() * relax: perform structure relaxation calculations, the relax_nmax parameter depicts the maximal number of ionic iterations * cell-relax: perform cell relaxation calculations * md: perform molecular dynamics simulations -* socket: run as a socket client for external drivers using the i-PI protocol * get_pchg: obtain partial (band-decomposed) charge densities (for LCAO basis only). See out_pchg for more information * get_wf: obtain real space wave functions (for LCAO basis only). See out_wfc_norm and out_wfc_re_im for more information * get_s: obtain the overlap matrix formed by localized orbitals (for LCAO basis with multiple k points). the file name is SR.csr with file format being the same as that generated by out_mat_hs2 @@ -103,7 +102,6 @@ void ReadInput::item_system() std::vector callist = {"scf", "relax", "md", - "socket", "cell-relax", "nscf", "get_s", @@ -138,6 +136,24 @@ void ReadInput::item_system() sync_string(input.calculation); this->add_item(item); } + { + Input_Item item("socket_driver"); + item.annotation = "run as a socket client for external drivers using the i-PI protocol"; + item.category = "System variables"; + item.type = "Boolean"; + item.description = R"(If set to True, ABACUS keeps the calculation type as scf and receives atomic positions from an external driver through the i-PI socket protocol. + +[NOTE] Use calculation = scf with socket_driver = True. Set ABACUS_SOCKET_ADDRESS to host:port or path:UNIX to choose the socket endpoint; if unset, localhost:31415 is used.)"; + item.default_value = "False"; + read_sync_bool(input.socket_driver); + item.check_value = [](const Input_Item& item, const Parameter& para) { + if (para.input.socket_driver && para.input.calculation != "scf") + { + ModuleBase::WARNING_QUIT("ReadInput", "socket_driver is only supported with calculation = scf."); + } + }; + this->add_item(item); + } { Input_Item item("esolver_type"); item.annotation = "the energy solver: ksdft, sdft, ofdft, tdofdft, tddft, lj, dp, ks-lr, lr"; @@ -268,9 +284,9 @@ void ReadInput::item_system() item.description = "If set to True, calculate the force at the end of the electronic iteration."; item.default_value = "False"; item.reset_value = [](const Input_Item& item, Parameter& para) { - std::vector use_force = {"cell-relax", "relax", "md", "socket"}; + std::vector use_force = {"cell-relax", "relax", "md"}; std::vector not_use_force = {"get_wf", "get_pchg", "get_s"}; - if (std::find(use_force.begin(), use_force.end(), para.input.calculation) != use_force.end()) + if (para.input.socket_driver || std::find(use_force.begin(), use_force.end(), para.input.calculation) != use_force.end()) { if (!para.input.cal_force) { @@ -878,7 +894,7 @@ Available options are: item.description = R"(Charge extrapolation method for MD, relaxation, and socket-driven calculations. When set to default, ABACUS chooses second-order for md, first-order for -relax/cell-relax/socket, and atomic for other calculations. Socket-driven +relax/cell-relax and socket_driver calculations, and atomic for other calculations. Socket-driven molecular dynamics can explicitly set second-order if the external driver updates structures smoothly enough for second-order extrapolation.)"; item.default_value = "default"; @@ -889,7 +905,7 @@ updates structures smoothly enough for second-order extrapolation.)"; para.input.chg_extrap = "second-order"; } else if (para.input.chg_extrap == "default" - && (para.input.calculation == "relax" || para.input.calculation == "cell-relax" || para.input.calculation == "socket")) + && (para.input.calculation == "relax" || para.input.calculation == "cell-relax" || para.input.socket_driver)) { para.input.chg_extrap = "first-order"; } diff --git a/source/source_io/test_serial/read_input_item_test.cpp b/source/source_io/test_serial/read_input_item_test.cpp index eb043fe1b03..0a53fa13d85 100644 --- a/source/source_io/test_serial/read_input_item_test.cpp +++ b/source/source_io/test_serial/read_input_item_test.cpp @@ -64,6 +64,27 @@ TEST_F(InputTest, Item_test) EXPECT_EXIT(it->second.check_value(it->second, param), ::testing::ExitedWithCode(1), ""); output = testing::internal::GetCapturedStdout(); EXPECT_THAT(output, testing::HasSubstr("NOTICE")); + + param.input.calculation = "socket"; + testing::internal::CaptureStdout(); + EXPECT_EXIT(it->second.check_value(it->second, param), ::testing::ExitedWithCode(1), ""); + output = testing::internal::GetCapturedStdout(); + EXPECT_THAT(output, testing::HasSubstr("NOTICE")); + } + + { // socket_driver + auto it = find_label("socket_driver", readinput.input_lists); + param.input.socket_driver = true; + param.input.calculation = "nscf"; + testing::internal::CaptureStdout(); + EXPECT_EXIT(it->second.check_value(it->second, param), ::testing::ExitedWithCode(1), ""); + output = testing::internal::GetCapturedStdout(); + EXPECT_THAT(output, testing::HasSubstr("NOTICE")); + + param.input.socket_driver = true; + param.input.calculation = "scf"; + EXPECT_NO_THROW(it->second.check_value(it->second, param)); + param.input.socket_driver = false; } { // esolver_type @@ -246,6 +267,7 @@ TEST_F(InputTest, Item_test) auto it = find_label("cal_force", readinput.input_lists); param.input.calculation = "cell-relax"; param.input.cal_force = false; + param.input.socket_driver = false; it->second.reset_value(it->second, param); EXPECT_EQ(param.input.cal_force, true); @@ -253,6 +275,13 @@ TEST_F(InputTest, Item_test) param.input.cal_force = true; it->second.reset_value(it->second, param); EXPECT_EQ(param.input.cal_force, false); + + param.input.calculation = "scf"; + param.input.socket_driver = true; + param.input.cal_force = false; + it->second.reset_value(it->second, param); + EXPECT_EQ(param.input.cal_force, true); + param.input.socket_driver = false; } { // ecutrho auto it = find_label("ecutrho", readinput.input_lists); @@ -363,12 +392,14 @@ TEST_F(InputTest, Item_test) EXPECT_EQ(param.input.chg_extrap, "first-order"); param.input.chg_extrap = "default"; - param.input.calculation = "socket"; + param.input.calculation = "scf"; + param.input.socket_driver = true; it->second.reset_value(it->second, param); EXPECT_EQ(param.input.chg_extrap, "first-order"); param.input.chg_extrap = "default"; param.input.calculation = "none"; + param.input.socket_driver = false; it->second.reset_value(it->second, param); EXPECT_EQ(param.input.chg_extrap, "atomic"); diff --git a/source/source_main/driver_run.cpp b/source/source_main/driver_run.cpp index e565871d604..d78ea0058e8 100644 --- a/source/source_main/driver_run.cpp +++ b/source/source_main/driver_run.cpp @@ -59,20 +59,12 @@ void Driver::driver_run() unitcell::check_atomic_stru(ucell, PARAM.inp.min_dist_coef); //! 2: initialize the ESolver (depends on a set-up ucell after `setup_cell`) - Input_para socket_esolver_inp = PARAM.inp; - const bool socket_mode = (socket_esolver_inp.calculation == "socket"); - if (socket_mode) - { - socket_esolver_inp.calculation = "scf"; - } - const Input_para& esolver_inp = socket_mode ? socket_esolver_inp : PARAM.inp; - this->init_hardware(); - ModuleESolver::ESolver* p_esolver = ModuleESolver::init_esolver(esolver_inp, ucell); + ModuleESolver::ESolver* p_esolver = ModuleESolver::init_esolver(PARAM.inp, ucell); //! 3: initialize Esolver and fill json-structure - p_esolver->before_all_runners(ucell, esolver_inp); + p_esolver->before_all_runners(ucell, PARAM.inp); // this Json part should be moved to before_all_runners, mohan 2024-05-12 #ifdef __RAPIDJSON @@ -85,7 +77,7 @@ void Driver::driver_run() { Run_MD::md_line(ucell, p_esolver, PARAM); } - else if (cal == "scf" || cal == "relax" || cal == "cell-relax" || cal == "nscf" || cal == "socket") + else if (cal == "scf" || cal == "relax" || cal == "cell-relax" || cal == "nscf") { Relax_Driver rl_driver; rl_driver.relax_driver(p_esolver, ucell, PARAM.inp, GlobalV::ofs_running); diff --git a/source/source_relax/CMakeLists.txt b/source/source_relax/CMakeLists.txt index 84f69c06c80..77079eb6a1f 100644 --- a/source/source_relax/CMakeLists.txt +++ b/source/source_relax/CMakeLists.txt @@ -2,7 +2,7 @@ add_library( relax OBJECT relax_data.cpp - ipi_socket.cpp + socket_ipi.cpp socket_driver.cpp cg_base.cpp relax_driver.cpp diff --git a/source/source_relax/relax_driver.cpp b/source/source_relax/relax_driver.cpp index 0cada1bfc3e..922a230d882 100644 --- a/source/source_relax/relax_driver.cpp +++ b/source/source_relax/relax_driver.cpp @@ -18,7 +18,7 @@ void Relax_Driver::relax_driver( ModuleBase::TITLE("Relax_Driver", "relax_driver"); ModuleBase::timer::start("Relax_Driver", "relax_driver"); - if (inp.calculation == "socket") + if (inp.socket_driver) { Socket_Driver socket_driver; socket_driver.socket_driver(p_esolver, ucell, inp, ofs_running); diff --git a/source/source_relax/socket_driver.cpp b/source/source_relax/socket_driver.cpp index ecd4d1be27e..42fafda5e63 100644 --- a/source/source_relax/socket_driver.cpp +++ b/source/source_relax/socket_driver.cpp @@ -1,6 +1,6 @@ #include "socket_driver.h" -#include "source_relax/ipi_socket.h" +#include "source_relax/socket_ipi.h" #include "source_base/global_function.h" #include "source_base/mathzone.h" #include "source_base/parallel_common.h" @@ -94,9 +94,9 @@ std::string bcast_header(std::string header) return header; } -std::string ipi_address() +std::string socket_address() { - const char* env = std::getenv("ABACUS_IPI_ADDRESS"); + const char* env = std::getenv("ABACUS_SOCKET_ADDRESS"); if (env == nullptr || std::string(env).empty()) { return "localhost:31415"; @@ -263,7 +263,7 @@ void Socket_Driver::socket_driver(ModuleESolver::ESolver* p_esolver, } if (!inp.cal_force) { - ModuleBase::WARNING_QUIT("ABACUS socket", "socket calculation requires cal_force=1 for i-PI GETFORCE."); + ModuleBase::WARNING_QUIT("ABACUS socket", "socket_driver requires cal_force=1 for i-PI GETFORCE."); } IpiSocket socket; @@ -276,7 +276,7 @@ void Socket_Driver::socket_driver(ModuleESolver::ESolver* p_esolver, { try { - const std::string address = ipi_address(); + const std::string address = socket_address(); ofs_running << " ABACUS socket driver connecting to i-PI endpoint " << address << std::endl; socket.connect(address); } diff --git a/source/source_relax/ipi_socket.cpp b/source/source_relax/socket_ipi.cpp similarity index 97% rename from source/source_relax/ipi_socket.cpp rename to source/source_relax/socket_ipi.cpp index 304a0181813..2bb09385b0d 100644 --- a/source/source_relax/ipi_socket.cpp +++ b/source/source_relax/socket_ipi.cpp @@ -1,4 +1,4 @@ -#include "source_relax/ipi_socket.h" +#include "source_relax/socket_ipi.h" #include #include @@ -149,6 +149,10 @@ std::string IpiSocket::read_header() { continue; } + if (errno == ECONNRESET && done == 0) + { + throw IpiSocketClosed("i-PI socket peer reset before next header"); + } throw std::runtime_error(errno_message("i-PI socket header read failed")); } done += static_cast(nread); diff --git a/source/source_relax/ipi_socket.h b/source/source_relax/socket_ipi.h similarity index 94% rename from source/source_relax/ipi_socket.h rename to source/source_relax/socket_ipi.h index 28ce2cb81ff..eac15eb41ac 100644 --- a/source/source_relax/ipi_socket.h +++ b/source/source_relax/socket_ipi.h @@ -1,5 +1,5 @@ -#ifndef ABACUS_IPI_SOCKET_H -#define ABACUS_IPI_SOCKET_H +#ifndef ABACUS_SOCKET_IPI_H +#define ABACUS_SOCKET_IPI_H #include #include diff --git a/source/source_relax/test/CMakeLists.txt b/source/source_relax/test/CMakeLists.txt index d1c867dc690..51dc9a82d3b 100644 --- a/source/source_relax/test/CMakeLists.txt +++ b/source/source_relax/test/CMakeLists.txt @@ -8,8 +8,8 @@ install(DIRECTORY support DESTINATION ${CMAKE_CURRENT_BINARY_DIR}) AddTest( - TARGET MODULE_RELAX_ipi_socket_test - SOURCES ipi_socket_test.cpp ../ipi_socket.cpp + TARGET MODULE_RELAX_socket_ipi_test + SOURCES socket_ipi_test.cpp ../socket_ipi.cpp ) AddTest( diff --git a/source/source_relax/test/ipi_socket_test.cpp b/source/source_relax/test/socket_ipi_test.cpp similarity index 98% rename from source/source_relax/test/ipi_socket_test.cpp rename to source/source_relax/test/socket_ipi_test.cpp index 21db9be2b90..930ab1fe629 100644 --- a/source/source_relax/test/ipi_socket_test.cpp +++ b/source/source_relax/test/socket_ipi_test.cpp @@ -1,4 +1,4 @@ -#include "../ipi_socket.h" +#include "../socket_ipi.h" #include "gtest/gtest.h" @@ -80,7 +80,7 @@ class UnixSocketServer public: UnixSocketServer() { - char dir_template[] = "/tmp/abacus_ipi_socket_test_XXXXXX"; + char dir_template[] = "/tmp/abacus_socket_ipi_test_XXXXXX"; char* made_dir = ::mkdtemp(dir_template); if (made_dir == nullptr) { From faafda3855a1bd7af3ebba43c33fc79bb078d961 Mon Sep 17 00:00:00 2001 From: jianrui geng Date: Sun, 19 Jul 2026 14:00:19 +0800 Subject: [PATCH 6/9] Fix socket parameter docs sync --- docs/advanced/interface/ase.md | 2 +- docs/parameters.yaml | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/docs/advanced/interface/ase.md b/docs/advanced/interface/ase.md index 82289cbac33..9e4153e3572 100644 --- a/docs/advanced/interface/ase.md +++ b/docs/advanced/interface/ase.md @@ -153,7 +153,7 @@ with abacus as calc: BFGS(atoms).run(fmax=0.05) ``` -`AbacusSocketIO` sets `socket_driver=1` and `cal_force=1` automatically. The socket endpoint is selected by the calculator arguments and passed to ABACUS through `ABACUS_SOCKET_ADDRESS`, for example `localhost:31415` or `/tmp/ipi_abacus_si:UNIX`. Calling `atoms.get_potential_energy()` is supported, but the ABACUS client still computes forces because the i-PI `GETFORCE` exchange returns energy, forces, and virial as one response. +`AbacusSocketIO` sets `socket_driver=1` and `cal_force=1` automatically. The socket endpoint is selected by the calculator arguments and passed to ABACUS through `ABACUS_SOCKET_ADDRESS`. `unixsocket="abacus_si"` means that ASE listens on the local Unix-domain socket `/tmp/ipi_abacus_si`; abacuslite launches ABACUS with `ABACUS_SOCKET_ADDRESS=/tmp/ipi_abacus_si:UNIX`. For TCP sockets, use the `port` argument instead, which maps to an address such as `localhost:31415`. Calling `atoms.get_potential_energy()` is supported, but the ABACUS client still computes forces because the i-PI `GETFORCE` exchange returns energy, forces, and virial as one response. A socket calculator owns one ABACUS process initialized from one fixed `INPUT`/`STRU` setup. Reuse the same `AbacusSocketIO` instance only for position updates under the same electronic-structure settings and the same cell. Do not change `kpts`, `kspacing`, `nspin`, `basis_type`, `basissets`, pseudopotentials, species, atom count, cell, or other core `INPUT`/`STRU` parameters through an existing socket calculator; create a new `AbacusSocketIO` instance and a new ABACUS client process for those changes. `AbacusSocketIO` rejects cell changes before sending them to ABACUS, and the ABACUS socket driver also checks incoming POSDATA cells against the initial `STRU` cell and exits if they differ. diff --git a/docs/parameters.yaml b/docs/parameters.yaml index 9c9a250707b..74127d1b841 100644 --- a/docs/parameters.yaml +++ b/docs/parameters.yaml @@ -47,7 +47,7 @@ parameters: [NOTE] Use calculation = scf with socket_driver = True. Set ABACUS_SOCKET_ADDRESS to host:port or path:UNIX to choose the socket endpoint; if unset, localhost:31415 is used. default_value: "False" unit: "" - availability: calculation==scf + availability: "" - name: esolver_type category: System variables type: String From 9f26e9f40188afc082c7e0937f702e416dea7c37 Mon Sep 17 00:00:00 2001 From: jianrui geng Date: Sun, 19 Jul 2026 14:06:28 +0800 Subject: [PATCH 7/9] Clarify socket endpoint address formats --- docs/advanced/input_files/input-main.md | 7 ++++++- docs/advanced/interface/ase.md | 11 +++++++++-- docs/parameters.yaml | 5 ++++- .../module_parameter/read_input_item_system.cpp | 5 ++++- 4 files changed, 23 insertions(+), 5 deletions(-) diff --git a/docs/advanced/input_files/input-main.md b/docs/advanced/input_files/input-main.md index ee1fb6c5b32..dee59db41d1 100644 --- a/docs/advanced/input_files/input-main.md +++ b/docs/advanced/input_files/input-main.md @@ -599,7 +599,12 @@ - **Availability**: *calculation==scf* - **Description**: If set to True, ABACUS keeps the calculation type as scf and receives atomic positions from an external driver through the i-PI socket protocol. - > Note: Use calculation = scf with socket_driver = True. Set ABACUS_SOCKET_ADDRESS to host:port or path:UNIX to choose the socket endpoint; if unset, localhost:31415 is used. + > Note: Use calculation = scf with socket_driver = True. ABACUS connects to the external i-PI server selected by ABACUS_SOCKET_ADDRESS. If ABACUS_SOCKET_ADDRESS is unset, ABACUS uses localhost:31415. The value can use one of two forms: + + - host:port, for example localhost:31415 or 127.0.0.1:31415, opens a TCP connection to that host and port. Use this when the i-PI server listens on a TCP port. + - path:UNIX, for example /tmp/ipi_abacus_si:UNIX, opens a Unix-domain socket at the given filesystem path. The :UNIX suffix tells ABACUS that the preceding value is a local socket path rather than a TCP host name. This form only works on the same machine. + + When using the ASE AbacusSocketIO interface, this environment variable is set automatically from the port or unixsocket calculator argument. - **Default**: False ### esolver_type diff --git a/docs/advanced/interface/ase.md b/docs/advanced/interface/ase.md index 9e4153e3572..7e4ed1a6af5 100644 --- a/docs/advanced/interface/ase.md +++ b/docs/advanced/interface/ase.md @@ -119,7 +119,7 @@ cmake -S . -B build-lcao -DENABLE_MPI=ON -DENABLE_LCAO=ON cmake --build build-lcao --target abacus_basic_para -j ``` -With the ABACUS toolchain workflow, build the normal ABACUS executable with LCAO support when `basis_type=lcao` is needed, then pass that executable to `AbacusProfile(command=...)`. The ASE interface can be installed from this repository with: +With the ABACUS toolchain workflow, build the normal ABACUS executable with LCAO support when `basis_type=lcao` is needed, then pass that executable to `AbacusProfile(command=...)`. The command can include an MPI launcher, for example `mpirun -np 4 /path/to/abacus`; ABACUS rank 0 opens the socket connection and broadcasts the i-PI data to the other ranks internally. The ASE interface can be installed from this repository with: ```bash cd interfaces/ASE_interface @@ -153,7 +153,14 @@ with abacus as calc: BFGS(atoms).run(fmax=0.05) ``` -`AbacusSocketIO` sets `socket_driver=1` and `cal_force=1` automatically. The socket endpoint is selected by the calculator arguments and passed to ABACUS through `ABACUS_SOCKET_ADDRESS`. `unixsocket="abacus_si"` means that ASE listens on the local Unix-domain socket `/tmp/ipi_abacus_si`; abacuslite launches ABACUS with `ABACUS_SOCKET_ADDRESS=/tmp/ipi_abacus_si:UNIX`. For TCP sockets, use the `port` argument instead, which maps to an address such as `localhost:31415`. Calling `atoms.get_potential_energy()` is supported, but the ABACUS client still computes forces because the i-PI `GETFORCE` exchange returns energy, forces, and virial as one response. +`AbacusSocketIO` sets `socket_driver=1` and `cal_force=1` automatically. It also selects the socket endpoint and passes it to ABACUS through `ABACUS_SOCKET_ADDRESS`, so users normally do not set this environment variable by hand when using abacuslite. + +There are two endpoint styles: + +- `unixsocket="abacus_si"` uses a local Unix-domain socket. ASE creates and listens on `/tmp/ipi_abacus_si`; abacuslite launches ABACUS with `ABACUS_SOCKET_ADDRESS=/tmp/ipi_abacus_si:UNIX`. The `:UNIX` suffix is part of ABACUS' address syntax and means that `/tmp/ipi_abacus_si` is a filesystem socket path, not a TCP host. This is usually the best choice when ASE and ABACUS run on the same node because it avoids TCP port conflicts. +- `port=31415` uses a TCP socket. abacuslite launches ABACUS with `ABACUS_SOCKET_ADDRESS=localhost:31415`, meaning host `localhost` and TCP port `31415`. Use this style when the socket server should listen on a TCP port. If ABACUS is launched manually instead of through `AbacusSocketIO`, set `ABACUS_SOCKET_ADDRESS` yourself to the same `host:port` or `path:UNIX` endpoint. + +Calling `atoms.get_potential_energy()` is supported, but the ABACUS client still computes forces because the i-PI `GETFORCE` exchange returns energy, forces, and virial as one response. A socket calculator owns one ABACUS process initialized from one fixed `INPUT`/`STRU` setup. Reuse the same `AbacusSocketIO` instance only for position updates under the same electronic-structure settings and the same cell. Do not change `kpts`, `kspacing`, `nspin`, `basis_type`, `basissets`, pseudopotentials, species, atom count, cell, or other core `INPUT`/`STRU` parameters through an existing socket calculator; create a new `AbacusSocketIO` instance and a new ABACUS client process for those changes. `AbacusSocketIO` rejects cell changes before sending them to ABACUS, and the ABACUS socket driver also checks incoming POSDATA cells against the initial `STRU` cell and exits if they differ. diff --git a/docs/parameters.yaml b/docs/parameters.yaml index 74127d1b841..6722db689bb 100644 --- a/docs/parameters.yaml +++ b/docs/parameters.yaml @@ -44,7 +44,10 @@ parameters: description: | If set to True, ABACUS keeps the calculation type as scf and receives atomic positions from an external driver through the i-PI socket protocol. - [NOTE] Use calculation = scf with socket_driver = True. Set ABACUS_SOCKET_ADDRESS to host:port or path:UNIX to choose the socket endpoint; if unset, localhost:31415 is used. + [NOTE] Use calculation = scf with socket_driver = True. ABACUS connects to the external i-PI server selected by ABACUS_SOCKET_ADDRESS. If ABACUS_SOCKET_ADDRESS is unset, ABACUS uses localhost:31415. The value can use one of two forms: + * host:port, for example localhost:31415 or 127.0.0.1:31415, opens a TCP connection to that host and port. Use this when the i-PI server listens on a TCP port. + * path:UNIX, for example /tmp/ipi_abacus_si:UNIX, opens a Unix-domain socket at the given filesystem path. The :UNIX suffix tells ABACUS that the preceding value is a local socket path rather than a TCP host name. This form only works on the same machine. + When using the ASE AbacusSocketIO interface, this environment variable is set automatically from the port or unixsocket calculator argument. default_value: "False" unit: "" availability: "" diff --git a/source/source_io/module_parameter/read_input_item_system.cpp b/source/source_io/module_parameter/read_input_item_system.cpp index bc2d1c94d4f..10444a1b4bd 100644 --- a/source/source_io/module_parameter/read_input_item_system.cpp +++ b/source/source_io/module_parameter/read_input_item_system.cpp @@ -143,7 +143,10 @@ void ReadInput::item_system() item.type = "Boolean"; item.description = R"(If set to True, ABACUS keeps the calculation type as scf and receives atomic positions from an external driver through the i-PI socket protocol. -[NOTE] Use calculation = scf with socket_driver = True. Set ABACUS_SOCKET_ADDRESS to host:port or path:UNIX to choose the socket endpoint; if unset, localhost:31415 is used.)"; +[NOTE] Use calculation = scf with socket_driver = True. ABACUS connects to the external i-PI server selected by ABACUS_SOCKET_ADDRESS. If ABACUS_SOCKET_ADDRESS is unset, ABACUS uses localhost:31415. The value can use one of two forms: +* host:port, for example localhost:31415 or 127.0.0.1:31415, opens a TCP connection to that host and port. Use this when the i-PI server listens on a TCP port. +* path:UNIX, for example /tmp/ipi_abacus_si:UNIX, opens a Unix-domain socket at the given filesystem path. The :UNIX suffix tells ABACUS that the preceding value is a local socket path rather than a TCP host name. This form only works on the same machine. +When using the ASE AbacusSocketIO interface, this environment variable is set automatically from the port or unixsocket calculator argument.)"; item.default_value = "False"; read_sync_bool(input.socket_driver); item.check_value = [](const Input_Item& item, const Parameter& para) { From 7c5c0c609f7de64de2481ca061d84272a6ecfd99 Mon Sep 17 00:00:00 2001 From: jianrui geng Date: Sun, 19 Jul 2026 14:57:32 +0800 Subject: [PATCH 8/9] Fix socket docs generation and launcher parsing --- docs/advanced/input_files/input-main.md | 2 -- docs/advanced/interface/ase.md | 6 +++++- interfaces/ASE_interface/abacuslite/core.py | 21 ++++++++++++++++----- 3 files changed, 21 insertions(+), 8 deletions(-) diff --git a/docs/advanced/input_files/input-main.md b/docs/advanced/input_files/input-main.md index dee59db41d1..a6f71db14cf 100644 --- a/docs/advanced/input_files/input-main.md +++ b/docs/advanced/input_files/input-main.md @@ -596,14 +596,12 @@ ### socket_driver - **Type**: Boolean -- **Availability**: *calculation==scf* - **Description**: If set to True, ABACUS keeps the calculation type as scf and receives atomic positions from an external driver through the i-PI socket protocol. > Note: Use calculation = scf with socket_driver = True. ABACUS connects to the external i-PI server selected by ABACUS_SOCKET_ADDRESS. If ABACUS_SOCKET_ADDRESS is unset, ABACUS uses localhost:31415. The value can use one of two forms: - host:port, for example localhost:31415 or 127.0.0.1:31415, opens a TCP connection to that host and port. Use this when the i-PI server listens on a TCP port. - path:UNIX, for example /tmp/ipi_abacus_si:UNIX, opens a Unix-domain socket at the given filesystem path. The :UNIX suffix tells ABACUS that the preceding value is a local socket path rather than a TCP host name. This form only works on the same machine. - When using the ASE AbacusSocketIO interface, this environment variable is set automatically from the port or unixsocket calculator argument. - **Default**: False diff --git a/docs/advanced/interface/ase.md b/docs/advanced/interface/ase.md index 7e4ed1a6af5..d69fdd33d5f 100644 --- a/docs/advanced/interface/ase.md +++ b/docs/advanced/interface/ase.md @@ -119,7 +119,11 @@ cmake -S . -B build-lcao -DENABLE_MPI=ON -DENABLE_LCAO=ON cmake --build build-lcao --target abacus_basic_para -j ``` -With the ABACUS toolchain workflow, build the normal ABACUS executable with LCAO support when `basis_type=lcao` is needed, then pass that executable to `AbacusProfile(command=...)`. The command can include an MPI launcher, for example `mpirun -np 4 /path/to/abacus`; ABACUS rank 0 opens the socket connection and broadcasts the i-PI data to the other ranks internally. The ASE interface can be installed from this repository with: +With the ABACUS toolchain workflow, build the normal ABACUS executable with LCAO support when `basis_type=lcao` is needed, then pass that executable to `AbacusProfile(command=...)`. The command can include an MPI launcher, for example `mpirun -np 4 /path/to/abacus`; ABACUS rank 0 opens the socket connection and broadcasts the i-PI data to the other ranks internally. On managed clusters, keep scheduler-specific launch options outside the calculator when possible and test the exact launcher command on a compute node. + +For PW calculations on CUDA/ROCm with multiple MPI ranks, use a k-point layout compatible with ABACUS' GPU parallelization. In practice, make sure each k-point pool contains one MPI rank; for example, a 4-rank PW GPU socket calculation should use at least four k-points so the default GPU `kpar` adjustment can assign one rank per pool. A one-k-point PW GPU job with several MPI ranks can fail in the PW GPU transform path; reduce the rank count or use a denser k-point mesh such as a smaller `kspacing`. + +The ASE interface can be installed from this repository with: ```bash cd interfaces/ASE_interface diff --git a/interfaces/ASE_interface/abacuslite/core.py b/interfaces/ASE_interface/abacuslite/core.py index e2c00d4c889..6285b426e5e 100644 --- a/interfaces/ASE_interface/abacuslite/core.py +++ b/interfaces/ASE_interface/abacuslite/core.py @@ -125,11 +125,14 @@ def __init__(self, @staticmethod def parse_version(stdout) -> str: - # up to the ABACUS version v3.9.0.17, the run of command - # `abacus --version` would returns the information organized - # in the following way: - # ABACUS version v3.9.0.17 - return re.match(r'ABACUS version (\S+)', stdout).group(1) + # MPI launchers may add informational lines before ABACUS output. + match = re.search(r'ABACUS version (\S+)', stdout or '') + if match is None: + raise RuntimeError( + 'Could not parse ABACUS version from command output. ' + 'Expected a line like "ABACUS version vX.Y.Z".' + ) + return match.group(1) def get_calculator_command(self, inputfile) -> List[str]: # because ABACUS run in the folder where there are INPUT files, so the @@ -778,6 +781,14 @@ def test_socketio_rejects_cell_changes(self): with self.assertRaisesRegex(PropertyNotImplementedError, 'fixed-cell'): calc._check_fixed_cell(changed) + def test_parse_version_allows_launcher_noise(self): + stdout = 'launcher info\nABACUS version v3.11.0-beta6\n' + self.assertEqual(AbacusProfile.parse_version(stdout), 'v3.11.0-beta6') + + def test_parse_version_rejects_missing_version(self): + with self.assertRaisesRegex(RuntimeError, 'ABACUS version'): + AbacusProfile.parse_version('launcher failed before abacus started') + def test_calculator_results(self): from ase.build.bulk import bulk silicon = bulk('Si', crystalstructure='diamond', a=5.43) From 81a3885eb732c54840634a1177b51b8887573cc7 Mon Sep 17 00:00:00 2001 From: jianrui geng Date: Mon, 20 Jul 2026 13:21:21 +0800 Subject: [PATCH 9/9] Document socket running log behavior --- docs/advanced/interface/ase.md | 2 ++ 1 file changed, 2 insertions(+) diff --git a/docs/advanced/interface/ase.md b/docs/advanced/interface/ase.md index d69fdd33d5f..b9735279d1c 100644 --- a/docs/advanced/interface/ase.md +++ b/docs/advanced/interface/ase.md @@ -168,6 +168,8 @@ Calling `atoms.get_potential_energy()` is supported, but the ABACUS client still A socket calculator owns one ABACUS process initialized from one fixed `INPUT`/`STRU` setup. Reuse the same `AbacusSocketIO` instance only for position updates under the same electronic-structure settings and the same cell. Do not change `kpts`, `kspacing`, `nspin`, `basis_type`, `basissets`, pseudopotentials, species, atom count, cell, or other core `INPUT`/`STRU` parameters through an existing socket calculator; create a new `AbacusSocketIO` instance and a new ABACUS client process for those changes. `AbacusSocketIO` rejects cell changes before sending them to ABACUS, and the ABACUS socket driver also checks incoming POSDATA cells against the initial `STRU` cell and exits if they differ. +In socket mode, ABACUS keeps one client process alive. All SCF evaluations produced by the same `AbacusSocketIO` instance are appended to the same `OUT.ABACUS/running_scf.log`, because the ABACUS calculation type remains `scf`. The authoritative per-step energy and force results are returned through the i-PI socket to ASE. Use ASE trajectory and optimizer log files, such as `BFGS(atoms, trajectory="opt.traj", logfile="opt.log")`, when each optimizer or MD step should be saved separately. Treat `running_scf.log` mainly as the ABACUS diagnostic log for the socket client, not as one independent FileIO result per structure. + The i-PI protocol does not transmit element symbols. `AbacusSocketIO` therefore sorts the internal socket atoms with the same first-occurrence species grouping used when writing `STRU`, and maps returned forces back to the original ASE `Atoms` order. This avoids silent force/atom mismatches when structures are read from CIF, extxyz, POSCAR, or other formats whose atom order is not already grouped for ABACUS. Users should not manually reorder atoms for socket I/O; pass the physical ASE `Atoms` object directly to the calculator. A complete fixed-cell validation and benchmark example is available in `interfaces/ASE_interface/examples/socketio.py`.