Loading...
Searching...
No Matches
ASE.cpp
Go to the documentation of this file.
1/*
2** This file is part of eOn.
3**
4** SPDX-License-Identifier: BSD-3-Clause
5**
6** Copyright (c) 2010--present, eOn Development Team
7** All rights reserved.
8**
9** Repo:
10** https://github.com/TheochemUI/eOn
11*/
12
14#include "eon/Eigen.h"
15#include "eon/EonLogger.h"
16#include "eon/Parameters.h"
17#include "eon/PyGuard.h"
18#include "eon/fpe_handler.h"
19#include <pybind11/eigen.h>
20#include <pybind11/embed.h>
21#include <pybind11/numpy.h>
22#include <pybind11/pybind11.h>
23#include <stdexcept>
24#include <string>
25#include <tuple>
26#include <vector>
27
28namespace py = pybind11;
29
30ASE::ASE(const eonc::Parameters &a_params)
31 : eonc::Potential(eonc::PotType::ASE_POT, a_params) {
33 counter = 1;
34 std::string py_file = a_params.potential_options().extPotPath;
35
36 // import
37 try {
38 // must briefly disable FPE because Python packages like Numpy causes it
39 // during import
41 fpeh.eat_fpe();
42
43 // Create a Python script to use importlib.util to load the module
44 py::exec(R"(
45 import sys
46 import importlib.util
47
48 def load_module_from_path(module_name, file_path):
49 spec = importlib.util.spec_from_file_location(module_name, file_path)
50 module = importlib.util.module_from_spec(spec)
51 sys.modules[module_name] = module
52 spec.loader.exec_module(module)
53 return module
54 )");
55
56 // Prepare the module name and file path
57 std::string module_name = "ase_eon";
58 py::object load_module = py::globals()["load_module_from_path"];
59 py_module = load_module(module_name, py_file);
60
61 fpeh.restore_fpe();
62
63 calculator = py_module.attr("ase_calc")();
64 _calculate = py_module.attr("_calculate");
65 has_batch_ = py::hasattr(py_module, "batch_calculate");
66 if (has_batch_) {
67 batch_calculate_ = py_module.attr("batch_calculate");
68 }
70 } catch (const std::exception &e) {
71 EONC_LOG_ERROR("ASE calculator import failed for {}: {}", py_file,
72 e.what());
73 throw std::runtime_error(std::string("ASE calculator import failed: ") +
74 e.what());
75 }
76 return;
77}
78
79void ASE::force(long nAtoms, const double *R, const int *atomicNrs, double *F,
80 double *U, double *variance, const double *box) {
81 if (variance != nullptr) {
82 *variance = 0.0;
83 }
84 py::gil_scoped_acquire gil;
85 try {
86 const Eigen::Map<const AtomMatrix> positions(R, nAtoms, 3);
87 const Eigen::Map<const RotationMatrix> boxx(box);
88 const Eigen::Map<const Eigen::VectorXi> atmnmrs(atomicNrs, nAtoms);
89
90 std::tuple<double, py::array_t<double>> py_result =
91 _calculate(positions, atmnmrs, boxx, calculator)
92 .cast<std::tuple<double, py::array_t<double>>>();
93
94 *U = std::get<0>(py_result);
95 py::array_t<double> forces = std::get<1>(py_result);
96 auto buffer = forces.request();
97 if (buffer.size < nAtoms * 3) {
98 throw std::runtime_error(
99 "ASE _calculate returned forces of the wrong size");
100 }
101 Eigen::Map<AtomMatrix>(F, nAtoms, 3) = Eigen::Map<const AtomMatrix>(
102 static_cast<const double *>(buffer.ptr), nAtoms, 3);
103
104 } catch (py::error_already_set &e) {
105 EONC_LOG_ERROR("ASE calculator Python error: {}", e.what());
106 throw std::runtime_error(std::string("ASE calculator Python error: ") +
107 e.what());
108 } catch (const std::exception &e) {
109 EONC_LOG_ERROR("ASE calculator C++ exception: {}", e.what());
110 throw std::runtime_error(std::string("ASE calculator C++ exception: ") +
111 e.what());
112 }
113
114 counter++;
115 return;
116}
117
118void ASE::forceBatch(long nSystems, long nAtoms, const double *const *positions,
119 const int *const *atomicNrs, double *const *forces,
120 double *energies, double *variances,
121 const double *const *boxes) {
122 if (!has_batch_) {
123 eonc::Potential::forceBatch(nSystems, nAtoms, positions, atomicNrs, forces,
124 energies, variances, boxes);
125 return;
126 }
127 py::gil_scoped_acquire gil;
128 try {
129 const auto ns = static_cast<size_t>(nSystems);
130 const auto na = static_cast<size_t>(nAtoms);
131 std::vector<double> R_data(ns * na * 3);
132 std::vector<int> Z_data(ns * na);
133 std::vector<double> box_data(ns * 9);
134 for (long s = 0; s < nSystems; ++s) {
135 std::copy(positions[s], positions[s] + static_cast<long>(na) * 3,
136 R_data.data() + static_cast<size_t>(s) * na * 3);
137 std::copy(atomicNrs[s], atomicNrs[s] + nAtoms,
138 Z_data.data() + static_cast<size_t>(s) * na);
139 std::copy(boxes[s], boxes[s] + 9,
140 box_data.data() + static_cast<size_t>(s) * 9);
141 }
142 py::array_t<double> R_np({ns, na, size_t{3}}, R_data.data());
143 py::array_t<int> Z_np({ns, na}, Z_data.data());
144 py::array_t<double> box_np({ns, size_t{3}, size_t{3}}, box_data.data());
145 auto py_result =
146 batch_calculate_(R_np, Z_np, box_np, calculator)
147 .cast<std::tuple<py::array_t<double>, py::array_t<double>>>();
148 py::array_t<double> E = std::get<0>(py_result);
149 py::array_t<double> F = std::get<1>(py_result);
150 auto bufE = E.request();
151 auto bufF = F.request();
152 if (bufE.size < nSystems || bufF.size < nSystems * nAtoms * 3) {
153 throw std::runtime_error(
154 "ASE batch_calculate returned energies/forces of the wrong size");
155 }
156 auto *ePtr = static_cast<double *>(bufE.ptr);
157 auto *fPtr = static_cast<double *>(bufF.ptr);
158 std::copy(ePtr, ePtr + nSystems, energies);
159 for (long s = 0; s < nSystems; ++s) {
160 std::copy(fPtr + s * nAtoms * 3, fPtr + (s + 1) * nAtoms * 3, forces[s]);
161 if (variances) {
162 variances[s] = 0.0;
163 }
164 }
165 } catch (py::error_already_set &e) {
166 EONC_LOG_ERROR("ASE calculator Python error: {}", e.what());
167 throw std::runtime_error(std::string("ASE calculator Python error: ") +
168 e.what());
169 } catch (const std::exception &e) {
170 EONC_LOG_ERROR("ASE calculator C++ exception: {}", e.what());
171 throw std::runtime_error(std::string("ASE calculator C++ exception: ") +
172 e.what());
173 }
174 counter += static_cast<size_t>(nSystems);
175}
#define EONC_LOG_ERROR(...)
Definition EonLogger.h:261
py::object batch_calculate_
Definition ASE.h:27
bool has_batch_
Definition ASE.h:28
py::object calculator
Definition ASE.h:24
void force(long nAtoms, const double *R, const int *atomicNrs, double *F, double *U, double *variance, const double *box) override
Definition ASE.cpp:69
py::module_ py_module
Definition ASE.h:23
void forceBatch(long nSystems, long nAtoms, const double *const *positions, const int *const *atomicNrs, double *const *forces, double *energies, double *variances, const double *const *boxes) override
Definition ASE.cpp:108
size_t counter
Definition ASE.h:22
py::object _calculate
Definition ASE.h:25
ASE(const eonc::Parameters &a_params)
Definition ASE.cpp:30
const potential_options_t & potential_options() const
virtual void forceBatch(long nSystems, long nAtoms, const double *const *positions, const int *const *atomicNrs, double *const *forces, double *energies, double *variances, const double *const *boxes)
Definition Potential.h:198
Potential(PotType a_ptype)
Production default: construction-scope registry, else PotRegistry::get().
RAII resource manager for the ARTn C library with global synchronization.
void ensure_interpreter()
Definition NbGuard.h:21