Loading...
Searching...
No Matches
MetatomicPotential Class Reference

A potential class that uses a metatomic model for energy and force calculations. More...

#include <MetatomicPotential.h>

Inheritance diagram for MetatomicPotential:

Classes

struct  CloneTag

Public Member Functions

 MetatomicPotential (const eonc::Parameters &params)
 Constructor for the MetatomicPotential.
 ~MetatomicPotential () override=default
 Destructor.
std::shared_ptr< eonc::Potential > clonePotential () const override
 Independent instance that does not reload from disk.
void force (long nAtoms, const double *positions, const int *atomicNrs, double *forces, double *energy, double *variance, const double *box) override
 Calculates the energy and forces for a given atomic configuration.
bool isThreadSafe () const noexcept override
 Single shared instance, serialized via mutex.
bool needsPerImageInstance () const noexcept override
 Whether NEB should create separate Potential instances per image for true parallel force evaluation.
bool supportsBatchEvaluation () const noexcept override
 Batched evaluation via single shared instance.
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
Public Member Functions inherited from eonc::Potential
 Potential (PotType a_ptype)
 Production default: construction-scope registry, else PotRegistry::get().
 Potential (PotType a_ptype, IPotRegistry &registry)
 Test seam: injected registry, no process-default get() counters.
 Potential (PotType a_ptype, const Parameters &p)
 Potential (const Parameters &a_params)
virtual ~Potential ()
void force (std::span< const double > positions, std::span< const int > atomicNrs, std::span< double > forces, double *energy, double *variance, std::span< const double > box)
 C++ call site: size-checked view over the raw FFI force().
virtual void setFixedMask (long nAtoms, const double *isFixed)
 Optional frozen-atom mask (nAtoms*3, 1.0 = fixed).
std::tuple< double, AtomMatrix > get_ef (const AtomMatrix &pos, const VectorXi &atmnrs, const Matrix3d &box)
PotType getType () const
virtual double finiteCutoff () const noexcept
 Finite interaction range in position length units.
virtual bool isSurrogate () const noexcept
 Whether this is a surrogate (GP) potential.
virtual bool requiresIsolatedMoleculeLayout () const noexcept
 True for molecular QM / non-PBC backends (NWChem socket, ASE ORCA/NWChem, …).
virtual unsigned layoutFlags () const noexcept
virtual bool isSharedInstanceThreadSafe () const noexcept
 Conservative gate for sharing one Potential instance across threads.
virtual bool computesStress () const noexcept
 True when force() leaves a Cauchy stress that cauchyStress() can read until the next force() on this instance.
virtual Matrix3d cauchyStress () const
 Cauchy stress in eV/Angstrom^3.
virtual void forceBatchOwned (long nSystems, long nAtoms, const double *const *positions, const int *const *atomicNrs, double *const *forces, double *energies, double *variances, const double *const *boxes, const long *owners)
 Evaluate forces for N systems in a single call.

Private Member Functions

metatensor_torch::TensorBlock computeNeighbors (metatomic_torch::NeighborListOptions request, long nAtoms, const double *positions, const double *box, const bool periodic[3])
 Computes neighbor list using the vesin library.
 MetatomicPotential (const MetatomicPotential &src, CloneTag)
void forceBatchNative (long nSystems, long nAtoms, const double *const *positions, const int *const *atomicNrs, double *const *forces, double *energies, double *variances, const double *const *boxes)

Private Attributes

eonc::log::Scoped m_log
eonc::Parameters::metatomic_options_t m_metatomic_opts
metatensor_torch::Module model_
metatomic_torch::ModelCapabilities capabilities_
std::vector< metatomic_torch::NeighborListOptions > nl_requests_
metatomic_torch::ModelEvaluationOptions evaluations_options_
torch::ScalarType dtype_
c10::DeviceType device_type_
torch::Device device_
bool check_consistency_
std::string energy_key_
std::string energy_uncertainty_key_
std::string nc_forces_key_
bool non_conservative_ {false}
bool random_rotation_ {false}
long n_symmetry_rotations_ {0}
double uncertainty_threshold_ {-1.0}
std::mutex inference_mutex_

Additional Inherited Members

Public Types inherited from eonc::Potential
enum class  PotLayout : unsigned { InProcess = 1u << 0 , NeedsWorkingDirectory = 1u << 1 , Subprocess = 1u << 2 }
 How the pot is executed. Combine with bitwise or. More...
Public Attributes inherited from eonc::Potential
std::atomic< size_t > forceCallCounter
Protected Attributes inherited from eonc::Potential
PotType ptype

Detailed Description

A potential class that uses a metatomic model for energy and force calculations.

This class loads a pre-trained atomistic model (via metatomic/PyTorch) and uses it to compute the potential energy and corresponding forces on atoms. It uses the vesin library to compute neighbor lists required by the model.

Definition at line 54 of file MetatomicPotential.h.

Constructor & Destructor Documentation

◆ MetatomicPotential() [1/2]

MetatomicPotential::MetatomicPotential ( const eonc::Parameters & params)

Constructor for the MetatomicPotential.

Parameters
paramsA shared pointer to the simulation parameters object. This object should contain the settings needed for the metatomic model.

Definition at line 56 of file MetatomicPotential.cpp.

57 : eonc::Potential(eonc::PotType::METATOMIC),
59 model_(torch::jit::Module()),
60 device_type_(c10::DeviceType::CPU),
61 device_(torch::Device(device_type_)) {
62
63 // Determinism knobs (see
64 // https://rgoswami.me/snippets/pytorch-deterministic-regression/): JIT
65 // profiling specializes graphs after the first few forwards, so a fresh model
66 // instance can differ at ULP level from a warm one — bad for parallel NEB.
67 // cuBLAS / index_add_ on CUDA are likewise nondeterministic unless forced.
68 // [Metatomic] deterministic=true (default) applies the safe defaults;
69 // deterministic_strict=true requires CUBLAS_WORKSPACE_CONFIG (e.g. :4096:8)
70 // and fails on nondeterministic ops instead of warning.
71 if (m_metatomic_opts.deterministic) {
72 torch::jit::getProfilingMode() = false;
73 const bool strict = m_metatomic_opts.deterministic_strict ||
74 (std::getenv("CUBLAS_WORKSPACE_CONFIG") != nullptr);
75 at::globalContext().setDeterministicAlgorithms(true,
76 /*warn_only=*/!strict);
77 at::globalContext().setBenchmarkCuDNN(false);
78 QUILL_LOG_INFO(m_log,
79 "[MetatomicPotential] Deterministic algorithms enabled "
80 "(strict={})",
81 strict);
82 } else {
83 QUILL_LOG_INFO(m_log,
84 "[MetatomicPotential] Deterministic algorithms disabled "
85 "(faster, may diverge across runs / NEB images)");
86 }
87
88 eonc::FPEHandler fpeh;
89 fpeh.eat_fpe();
90
91 QUILL_LOG_INFO(m_log, "[MetatomicPotential] Initializing...");
92
93 // 1. Load the model from the path specified in parameters
94 torch::optional<std::string> extensions_directory = torch::nullopt;
95 if (!m_metatomic_opts.extensions_directory.empty()) {
96 extensions_directory = m_metatomic_opts.extensions_directory;
97 }
98
99 try {
100 QUILL_LOG_INFO(m_log, "[MetatomicPotential] Loading model from '{}'",
101 m_metatomic_opts.model_path);
102 this->model_ = metatomic_torch::load_atomistic_model(
103 m_metatomic_opts.model_path, extensions_directory);
104 } catch (const std::exception &e) {
105 QUILL_LOG_ERROR(m_log, "[MetatomicPotential] Failed to load model: {}",
106 e.what());
107 throw;
108 }
109
110 // 2. Extract capabilities and neighbor list requests from the model
111 this->capabilities_ =
112 this->model_.run_method("capabilities")
113 .toCustomClass<metatomic_torch::ModelCapabilitiesHolder>();
114 auto requests_ivalue = this->model_.run_method("requested_neighbor_lists");
115 for (const auto &request_ivalue : requests_ivalue.toList()) {
116 auto request =
117 request_ivalue.get()
118 .toCustomClass<metatomic_torch::NeighborListOptionsHolder>();
119 this->nl_requests_.push_back(request);
120 }
121
122 // 3. Determine and set up the device (CPU/CUDA/MPS)
123 torch::optional<std::string> desired = torch::nullopt;
124 if (!m_metatomic_opts.device.empty()) {
125 desired = m_metatomic_opts.device;
126 }
127 device_type_ = metatomic_torch::pick_device(
128 this->capabilities_->supported_devices, desired);
129
130 device_ = torch::Device(device_type_);
131 QUILL_LOG_INFO(m_log, "[MetatomicPotential] Using device: {}", device_.str());
132
133 move_atomistic_model(this->model_, this->device_);
134
135 // 4. Set data type (float32/float64) based on model capabilities
136 if (this->capabilities_->dtype() == "float64") {
137 this->dtype_ = torch::kFloat64;
138 } else if (this->capabilities_->dtype() == "float32") {
139 this->dtype_ = torch::kFloat32;
140 } else {
141 throw std::runtime_error("Unsupported dtype: " +
142 this->capabilities_->dtype());
143 }
144 QUILL_LOG_INFO(m_log, "[MetatomicPotential] Using dtype: {}",
145 this->capabilities_->dtype().c_str());
146
147 // 5. Resolve energy / force output keys: explicit keys (#215) or variants
148 // (#296 for non_conservative_force)
149 auto outputs = this->capabilities_->outputs();
150
151 auto v_base = normalize_variant(m_metatomic_opts.variant.base);
152 auto v_energy = m_metatomic_opts.variant.energy.empty()
153 ? v_base
154 : normalize_variant(m_metatomic_opts.variant.energy);
155 auto v_energy_uq =
156 m_metatomic_opts.variant.energy_uncertainty.empty()
157 ? v_energy
158 : normalize_variant(m_metatomic_opts.variant.energy_uncertainty);
159 auto v_force = m_metatomic_opts.variant.force.empty()
160 ? v_energy
161 : normalize_variant(m_metatomic_opts.variant.force);
162
163 if (!m_metatomic_opts.energy_output.empty()) {
164 this->energy_key_ = m_metatomic_opts.energy_output;
165 } else {
166 this->energy_key_ =
167 metatomic_torch::pick_output("energy", outputs, v_energy);
168 }
169
170 this->non_conservative_ = m_metatomic_opts.non_conservative;
171 this->random_rotation_ = m_metatomic_opts.random_rotation;
172 this->n_symmetry_rotations_ = m_metatomic_opts.n_symmetry_rotations;
173 if (this->non_conservative_) {
174 if (!m_metatomic_opts.force_output.empty()) {
175 this->nc_forces_key_ = m_metatomic_opts.force_output;
176 if (!outputs.contains(this->nc_forces_key_)) {
177 throw std::runtime_error(
178 "Missing explicit force_output in metatomic model: " +
179 this->nc_forces_key_);
180 }
181 } else {
182 this->nc_forces_key_ = metatomic_torch::pick_output(
183 "non_conservative_force", outputs, v_force);
184 }
185 QUILL_LOG_INFO(m_log,
186 "[MetatomicPotential] Non-conservative forces from '{}'",
187 this->nc_forces_key_);
188 }
189 if (this->n_symmetry_rotations_ > 0) {
190 QUILL_LOG_INFO(m_log,
191 "[MetatomicPotential] Symmetry averaging over {} rotations",
193 } else if (this->random_rotation_) {
194 QUILL_LOG_INFO(
195 m_log, "[MetatomicPotential] Per-call random SO(3) rotation enabled");
196 }
197 if ((this->random_rotation_ || this->n_symmetry_rotations_ > 0) &&
198 params.main_options().randomSeed > 0) {
199 torch::manual_seed(static_cast<uint64_t>(params.main_options().randomSeed));
200 QUILL_LOG_INFO(m_log,
201 "[MetatomicPotential] torch RNG seeded from "
202 "main.randomSeed={}",
203 params.main_options().randomSeed);
204 }
205
206 if (!outputs.contains(this->energy_key_)) {
207 QUILL_LOG_ERROR(
208 m_log,
209 "[MetatomicPotential] The model does not provide an '{}' output.",
210 this->energy_key_);
211 throw std::runtime_error("Missing energy output in metatomic model");
212 }
213
214 // 6. Set up evaluation options to request total energy
216 torch::make_intrusive<metatomic_torch::ModelEvaluationOptionsHolder>();
217 evaluations_options_->set_length_unit(m_metatomic_opts.length_unit);
218
219 auto model_output = outputs.at(this->energy_key_);
220 auto requested_output =
221 torch::make_intrusive<metatomic_torch::ModelOutputHolder>();
222
223 // Per-atom granularity is sample_kind == "atom" (get/set_per_atom removed).
224 requested_output->set_sample_kind(model_output->sample_kind());
225 requested_output->explicit_gradients = {};
226 requested_output->set_unit("eV");
227 evaluations_options_->outputs.insert(this->energy_key_, requested_output);
228
229 // Request non-conservative forces when enabled (#296)
230 if (this->non_conservative_ && !this->nc_forces_key_.empty()) {
231 auto nc_info = outputs.at(this->nc_forces_key_);
232 auto requested_nc =
233 torch::make_intrusive<metatomic_torch::ModelOutputHolder>();
234 requested_nc->set_sample_kind(nc_info->sample_kind());
235 requested_nc->explicit_gradients = {};
236 requested_nc->set_unit("eV/Angstrom");
237 evaluations_options_->outputs.insert(this->nc_forces_key_, requested_nc);
238 }
239
240 // 7. Optionally request energy uncertainty if threshold is positive
241 if (m_metatomic_opts.uncertainty_threshold > 0) {
242 this->uncertainty_threshold_ = m_metatomic_opts.uncertainty_threshold;
243 const bool explicit_uq_key =
244 !m_metatomic_opts.energy_uncertainty_output.empty();
245 if (explicit_uq_key) {
246 // User-specified key: hard fail if missing (not soft-disabled).
248 m_metatomic_opts.energy_uncertainty_output;
249 if (!outputs.contains(this->energy_uncertainty_key_)) {
250 QUILL_LOG_ERROR(m_log,
251 "[MetatomicPotential] energy_uncertainty_output '{}' "
252 "is not provided by the model.",
254 throw std::runtime_error(
255 "Missing explicit energy_uncertainty_output in metatomic model: " +
257 }
258 } else {
259 try {
260 this->energy_uncertainty_key_ = metatomic_torch::pick_output(
261 "energy_uncertainty", outputs, v_energy_uq);
262 } catch (const std::exception &e) {
263 QUILL_LOG_DEBUG(
264 m_log, "[MetatomicPotential] No uncertainty output available: {}",
265 e.what());
266 this->uncertainty_threshold_ = -1.0;
267 }
268 }
269
270 if (this->uncertainty_threshold_ > 0) {
271 auto uncertainty_info = outputs.at(this->energy_uncertainty_key_);
272 if (uncertainty_info->sample_kind() == "atom") {
273 auto requested_uncertainty =
274 torch::make_intrusive<metatomic_torch::ModelOutputHolder>();
275 requested_uncertainty->set_sample_kind("atom");
276 requested_uncertainty->explicit_gradients = {};
277 requested_uncertainty->set_unit("eV");
278 evaluations_options_->outputs.insert(this->energy_uncertainty_key_,
279 requested_uncertainty);
280 QUILL_LOG_INFO(m_log,
281 "[MetatomicPotential] Requested per-atom "
282 "'{}' from model (threshold = {})",
285 } else {
286 QUILL_LOG_DEBUG(m_log,
287 "[MetatomicPotential] Model provides '{}' "
288 "but sample_kind is not \"atom\"; skipping uncertainty "
289 "checks.",
291 this->uncertainty_threshold_ = -1.0;
292 }
293 }
294 }
295
296 this->check_consistency_ = m_metatomic_opts.check_consistency;
297 QUILL_LOG_INFO(m_log, "[MetatomicPotential] Initialization complete.");
298
299 fpeh.restore_fpe();
300}
static torch::optional< std::string > normalize_variant(const std::string &s)
std::vector< metatomic_torch::NeighborListOptions > nl_requests_
metatomic_torch::ModelEvaluationOptions evaluations_options_
eonc::log::Scoped m_log
metatomic_torch::ModelCapabilities capabilities_
std::string energy_uncertainty_key_
eonc::Parameters::metatomic_options_t m_metatomic_opts
c10::DeviceType device_type_
metatensor_torch::Module model_
torch::ScalarType dtype_
const metatomic_options_t & metatomic_options() const
const main_options_t & main_options() const

◆ ~MetatomicPotential()

MetatomicPotential::~MetatomicPotential ( )
overridedefault

Destructor.

◆ MetatomicPotential() [2/2]

MetatomicPotential::MetatomicPotential ( const MetatomicPotential & src,
CloneTag  )
private

Definition at line 302 of file MetatomicPotential.cpp.

Member Function Documentation

◆ clonePotential()

std::shared_ptr< eonc::Potential > MetatomicPotential::clonePotential ( ) const
nodiscardoverridevirtual

Independent instance that does not reload from disk.

nullptr means the caller should use makePotential().

Reimplemented from eonc::Potential.

Definition at line 323 of file MetatomicPotential.cpp.

323 {
324 return std::shared_ptr<eonc::Potential>(
325 new MetatomicPotential(*this, CloneTag{}));
326}
MetatomicPotential(const eonc::Parameters &params)
Constructor for the MetatomicPotential.

◆ computeNeighbors()

metatensor_torch::TensorBlock MetatomicPotential::computeNeighbors ( metatomic_torch::NeighborListOptions request,
long nAtoms,
const double * positions,
const double * box,
const bool periodic[3] )
private

Computes neighbor list using the vesin library.

This function calls vesin_neighbors to build a neighbor list based on the model's requirements (cutoff, full/half list) and converts it into a metatensor::TensorBlock suitable for metatomic.

Parameters
requestThe neighbor list options requested by the model.
nAtomsThe number of atoms in the system.
positionsPointer to the atomic positions array.
boxPointer to the simulation box matrix.
Returns
A metatensor_torch::TensorBlock containing the neighbor list.

Definition at line 571 of file MetatomicPotential.cpp.

573 {
574
575 auto cutoff = request->engine_cutoff(m_metatomic_opts.length_unit);
576
577 // Zero-init so vesin 0.6 skin/n_threads stay 0 when compiling against 0.6
578 // headers (must match linked libvesin — pin vesin>=0.6 for metatomic builds).
579 VesinOptions options{};
580 options.cutoff = cutoff;
581 options.full = request->full_list();
582 options.sorted = false;
583 options.return_shifts = true;
584 options.return_distances = false; // we don't need distances
585 options.return_vectors = true; // metatomic uses vectors for autograd
586
587 VesinNeighborList *vesin_neighbor_list = new VesinNeighborList();
588
589 VesinDevice cpu{VesinCPU, 0};
590 const char *error_message = nullptr;
591 int status = vesin_neighbors(reinterpret_cast<const double (*)[3]>(positions),
592 static_cast<size_t>(nAtoms),
593 reinterpret_cast<const double (*)[3]>(box),
594 const_cast<bool *>(periodic), cpu, options,
595 vesin_neighbor_list, &error_message);
596
597 if (status != EXIT_SUCCESS) {
598 std::string err_str = "vesin_neighbors failed";
599 if (error_message != nullptr) {
600 err_str += ": " + std::string(error_message);
601 } else {
602 err_str += " (no message; vesin header/lib ABI mismatch? need vesin>=0.6 "
603 "with matching engine)";
604 }
605 delete vesin_neighbor_list;
606 throw std::runtime_error(err_str);
607 }
608
609 // Convert from vesin to metatomic format
610 auto n_pairs = static_cast<int64_t>(vesin_neighbor_list->length);
611 auto labels_options_cpu =
612 torch::TensorOptions().dtype(torch::kInt32).device(torch::kCPU);
613
614 auto pair_samples_values = torch::empty({n_pairs, 5}, labels_options_cpu);
615 auto pair_samples_values_ptr = pair_samples_values.accessor<int32_t, 2>();
616 for (int64_t i = 0; i < n_pairs; i++) {
617 pair_samples_values_ptr[i][0] =
618 static_cast<int32_t>(vesin_neighbor_list->pairs[i][0]);
619 pair_samples_values_ptr[i][1] =
620 static_cast<int32_t>(vesin_neighbor_list->pairs[i][1]);
621 pair_samples_values_ptr[i][2] = vesin_neighbor_list->shifts[i][0];
622 pair_samples_values_ptr[i][3] = vesin_neighbor_list->shifts[i][1];
623 pair_samples_values_ptr[i][4] = vesin_neighbor_list->shifts[i][2];
624 }
625
626 // Custom deleter to free vesin's memory when the torch tensor is destroyed
627 auto deleter = [=](void *) {
628 vesin_free(vesin_neighbor_list);
629 delete vesin_neighbor_list;
630 };
631
632 auto pair_vectors = torch::from_blob(
633 vesin_neighbor_list->vectors, {n_pairs, 3, 1}, deleter,
634 torch::TensorOptions().dtype(torch::kFloat64).device(torch::kCPU));
635
636 auto neighbor_samples = torch::make_intrusive<metatensor_torch::LabelsHolder>(
637 std::vector<std::string>{"first_atom", "second_atom", "cell_shift_a",
638 "cell_shift_b", "cell_shift_c"},
639 pair_samples_values.to(this->device_));
640
641 auto labels_options_dev =
642 torch::TensorOptions().dtype(torch::kInt32).device(this->device_);
643 auto neighbor_component =
644 torch::make_intrusive<metatensor_torch::LabelsHolder>(
645 "xyz", torch::tensor({0, 1, 2}, labels_options_dev).reshape({3, 1}));
646 auto neighbor_properties =
647 torch::make_intrusive<metatensor_torch::LabelsHolder>(
648 "distance", torch::zeros({1, 1}, labels_options_dev));
649
650 return torch::make_intrusive<metatensor_torch::TensorBlockHolder>(
651 pair_vectors.to(this->dtype_).to(this->device_), neighbor_samples,
652 std::vector<metatensor_torch::Labels>{neighbor_component},
653 neighbor_properties);
654}

◆ force()

void MetatomicPotential::force ( long nAtoms,
const double * positions,
const int * atomicNrs,
double * forces,
double * energy,
double * variance,
const double * box )
overridevirtual

Calculates the energy and forces for a given atomic configuration.

This is the core method of the potential. It takes the current atomic positions, builds the necessary data structures for metatomic, executes the model to get the potential energy, and uses PyTorch's autograd to compute forces.

Parameters
nAtomsNumber of atoms.
positionsFlat array of atomic positions (size nAtoms * 3).
atomicNrsFlat array of atomic numbers (size nAtoms), used for consistency checks.
forcesFlat array to store the calculated forces (size nAtoms * 3).
energyPointer to a double to store the calculated potential energy.
variancePointer to a double to store the variance of the energy (currently unused, set to NULL).
boxThe simulation box vectors (3x3 matrix).

Implements eonc::Potential.

Definition at line 352 of file MetatomicPotential.cpp.

355 {
356 // Serialize concurrent calls -- PyTorch model inference on the same
357 // Module instance is not thread-safe
358 std::lock_guard<std::mutex> lock(inference_mutex_);
359
360 eonc::FPEHandler fpeh;
361 fpeh.eat_fpe();
362
363 if (!atomicNrs) {
364 throw std::runtime_error(
365 "[MetatomicPotential] `atomicNrs` must be provided.");
366 }
367
368 const bool use_rotation =
369 this->random_rotation_ || this->n_symmetry_rotations_ > 0;
370 // n_symmetry_rotations averages; random_rotation alone is one rotated eval
371 const long n_passes =
372 this->n_symmetry_rotations_ > 0 ? this->n_symmetry_rotations_ : 1;
373
374 auto f64_options =
375 torch::TensorOptions().dtype(torch::kFloat64).device(torch::kCPU);
376 std::vector<int32_t> types_vec(atomicNrs, atomicNrs + nAtoms);
377 auto atomic_types_cpu =
378 torch::tensor(types_vec, torch::TensorOptions().dtype(torch::kInt32));
379
380 double energy_acc = 0.0;
381 auto forces_acc = torch::zeros({nAtoms, 3}, f64_options);
382 bool variance_set = false;
383
384 for (long i_pass = 0; i_pass < n_passes; ++i_pass) {
385 torch::Tensor R = torch::eye(
386 3, torch::TensorOptions().dtype(this->dtype_).device(this->device_));
387 if (use_rotation) {
388 R = random_so3(this->device_, this->dtype_);
389 }
390 // R is applied to row vectors: pos' = pos @ R^T (equiv. R @ pos for cols)
391 auto R_cpu = R.to(torch::kCPU).to(torch::kFloat64);
392 auto R_T = R.transpose(0, 1);
393
394 auto pos_cpu = torch::from_blob(const_cast<double *>(positions),
395 {nAtoms, 3}, f64_options)
396 .clone();
397 auto cell_cpu =
398 torch::from_blob(const_cast<double *>(box), {3, 3}, f64_options)
399 .clone();
400 if (use_rotation) {
401 pos_cpu = pos_cpu.matmul(R_cpu.transpose(0, 1));
402 // Rotate cell vectors (rows) the same way
403 cell_cpu = cell_cpu.matmul(R_cpu.transpose(0, 1));
404 }
405
406 std::vector<double> pos_buf(static_cast<size_t>(nAtoms) * 3);
407 std::vector<double> cell_buf(9);
408 std::memcpy(pos_buf.data(), pos_cpu.contiguous().data_ptr<double>(),
409 pos_buf.size() * sizeof(double));
410 std::memcpy(cell_buf.data(), cell_cpu.contiguous().data_ptr<double>(),
411 9 * sizeof(double));
412
413 auto torch_positions =
414 torch::from_blob(pos_buf.data(), {nAtoms, 3}, f64_options)
415 .to(this->dtype_)
416 .to(this->device_)
417 .set_requires_grad(!this->non_conservative_);
418
419 auto torch_cell = torch::from_blob(cell_buf.data(), {3, 3}, f64_options)
420 .to(this->dtype_)
421 .to(this->device_);
422
423 auto cell_norms = torch::norm(torch_cell, 2, /*dim=*/1);
424 auto torch_pbc = cell_norms.abs() > 1e-9;
425 bool periodic[3] = {torch_pbc[0].item<bool>(), torch_pbc[1].item<bool>(),
426 torch_pbc[2].item<bool>()};
427
428 auto atomic_types = atomic_types_cpu.to(this->device_);
429
430 auto system = torch::make_intrusive<metatomic_torch::SystemHolder>(
431 atomic_types, torch_positions, torch_cell, torch_pbc);
432
433 for (const auto &request : this->nl_requests_) {
434 auto neighbors = this->computeNeighbors(request, nAtoms, pos_buf.data(),
435 cell_buf.data(), periodic);
436 metatomic_torch::register_autograd_neighbors(system, neighbors,
437 this->check_consistency_);
438 system->add_neighbor_list(request, neighbors);
439 }
440
441 torch::Tensor forces_tensor;
442 try {
443 auto ivalue_output = this->model_.forward({
444 std::vector<metatomic_torch::System>{system},
446 this->check_consistency_,
447 });
448 auto dict_output = ivalue_output.toGenericDict();
449 auto output_map = dict_output.at(this->energy_key_)
450 .toCustomClass<metatensor_torch::TensorMapHolder>();
451
452 if (this->uncertainty_threshold_ > 0 && i_pass == 0) {
453 try {
454 if (dict_output.contains(this->energy_uncertainty_key_)) {
455 auto uncertainty_map =
456 dict_output.at(this->energy_uncertainty_key_)
457 .toCustomClass<metatensor_torch::TensorMapHolder>();
458 auto uncertainty_block =
459 metatensor_torch::TensorMapHolder::block_by_id(uncertainty_map,
460 0);
461 auto flat_uncertainty =
462 uncertainty_block->values().reshape({-1}).to(torch::kCPU);
463 if (variance != nullptr && flat_uncertainty.numel() > 0) {
464 try {
465 *variance =
466 flat_uncertainty.to(torch::kFloat64).mean().item<double>();
467 variance_set = true;
468 } catch (...) {
469 QUILL_LOG_DEBUG(m_log,
470 "[MetatomicPotential] Failed to compute mean "
471 "uncertainty for variance.");
472 }
473 }
474 auto atoms_above_threshold =
475 flat_uncertainty > this->uncertainty_threshold_;
476 if (torch::any(atoms_above_threshold).item<bool>()) {
477 auto samples = uncertainty_block->samples();
478 auto atom_indices_all = samples->column("atom").to(torch::kCPU);
479 auto atom_indices_above =
480 atom_indices_all.index({atoms_above_threshold});
481 std::ostringstream ss;
482 ss << "atoms at index [";
483 auto n_report = std::min<int64_t>(10, atom_indices_above.size(0));
484 for (int64_t i = 0; i < n_report; ++i) {
485 if (i > 0)
486 ss << ", ";
487 ss << atom_indices_above[i].item<int32_t>();
488 }
489 ss << "]";
490 if (atom_indices_above.size(0) > n_report) {
491 ss << " and " << (atom_indices_above.size(0) - n_report)
492 << " more";
493 }
494 QUILL_LOG_WARNING(
495 m_log,
496 "[MetatomicPotential] The uncertainty on atomic energies for "
497 "{} are larger than the threshold of {}. (Key: {}) Be "
498 "careful "
499 "when analyzing the results, and consider retraining the "
500 "model to better describe these configurations.",
501 ss.str(), this->uncertainty_threshold_,
502 this->energy_uncertainty_key_);
503 }
504 }
505 } catch (const std::exception &e) {
506 QUILL_LOG_WARNING(m_log,
507 "[MetatomicPotential] Failed to check {}: {}",
508 this->energy_uncertainty_key_, e.what());
509 }
510 }
511
512 auto energy_block =
513 metatensor_torch::TensorMapHolder::block_by_id(output_map, 0);
514 auto energy_tensor = energy_block->values();
515 energy_acc += energy_tensor.sum().item<double>();
516
517 if (this->non_conservative_ && !this->nc_forces_key_.empty()) {
518 auto nc_map = dict_output.at(this->nc_forces_key_)
519 .toCustomClass<metatensor_torch::TensorMapHolder>();
520 auto nc_block =
521 metatensor_torch::TensorMapHolder::block_by_id(nc_map, 0);
522 auto nc_vals = nc_block->values();
523 if (nc_vals.numel() != nAtoms * 3) {
524 throw std::runtime_error("[MetatomicPotential] NC force block has " +
525 std::to_string(nc_vals.numel()) +
526 " values, expected " +
527 std::to_string(nAtoms * 3));
528 }
529 auto nc_samples = nc_block->samples();
530 if (nc_samples->size() > 0 && nc_samples->names().size() > 1) {
531 auto atom_col = nc_samples->column("atom").to(torch::kCPU);
532 if (atom_col.size(0) != nAtoms) {
533 throw std::runtime_error(
534 "[MetatomicPotential] NC force samples atom count mismatch");
535 }
536 }
537 forces_tensor =
538 nc_vals.reshape({nAtoms, 3}).to(torch::kCPU).to(torch::kFloat64);
539 } else {
540 energy_tensor.backward(torch::ones_like(energy_tensor));
541 auto positions_grad = system->positions().grad();
542 forces_tensor = (-positions_grad).to(torch::kCPU).to(torch::kFloat64);
543 }
544 } catch (const std::exception &e) {
545 QUILL_LOG_ERROR(m_log, "[MetatomicPotential] Model evaluation failed: {}",
546 e.what());
547 throw;
548 }
549
550 // Rotate forces back to original frame: F = F' @ R (since pos' = pos @
551 // R^T)
552 if (use_rotation) {
553 forces_tensor = forces_tensor.matmul(R_cpu);
554 }
555 forces_acc += forces_tensor;
556 }
557
558 const double inv_n = 1.0 / static_cast<double>(n_passes);
559 *energy = energy_acc * inv_n;
560 forces_acc = forces_acc * inv_n;
561 (void)variance_set;
562
563 std::memcpy(forces, forces_acc.contiguous().data_ptr<double>(),
564 nAtoms * 3 * sizeof(double));
565
566 fpeh.restore_fpe();
567}
metatensor_torch::TensorBlock computeNeighbors(metatomic_torch::NeighborListOptions request, long nAtoms, const double *positions, const double *box, const bool periodic[3])
Computes neighbor list using the vesin library.

◆ forceBatch()

void MetatomicPotential::forceBatch ( long nSystems,
long nAtoms,
const double *const * positions,
const int *const * atomicNrs,
double *const * forces,
double * energies,
double * variances,
const double *const * boxes )
overridevirtual

Reimplemented from eonc::Potential.

Definition at line 662 of file MetatomicPotential.cpp.

667 {
668 if (nSystems > 1) {
669 try {
670 forceBatchNative(nSystems, nAtoms, positions, atomicNrs, forces, energies,
671 variances, boxes);
672 forceCallCounter += nSystems;
674 return;
675 } catch (const std::exception &e) {
676 QUILL_LOG_WARNING(m_log,
677 "[MetatomicPotential] batched forward failed ({}); "
678 "falling back to sequential force()",
679 e.what());
680 }
681 }
682 for (long s = 0; s < nSystems; s++) {
683 double var = 0;
684 force(nAtoms, positions[s], atomicNrs[s], forces[s], &energies[s], &var,
685 boxes[s]);
686 if (variances)
687 variances[s] = var;
690 }
691}
void force(long nAtoms, const double *positions, const int *atomicNrs, double *forces, double *energy, double *variance, const double *box) override
Calculates the energy and forces for a given atomic configuration.
void forceBatchNative(long nSystems, long nAtoms, const double *const *positions, const int *const *atomicNrs, double *const *forces, double *energies, double *variances, const double *const *boxes)
void on_force_call(PotType t) noexcept override
static PotRegistry & get() noexcept
Process-lifetime singleton.
std::atomic< size_t > forceCallCounter
Definition Potential.h:53
PotType ptype
Definition Potential.h:44

◆ forceBatchNative()

void MetatomicPotential::forceBatchNative ( long nSystems,
long nAtoms,
const double *const * positions,
const int *const * atomicNrs,
double *const * forces,
double * energies,
double * variances,
const double *const * boxes )
private

Definition at line 693 of file MetatomicPotential.cpp.

698 {
699 std::lock_guard<std::mutex> lock(inference_mutex_);
700
701 eonc::FPEHandler fpeh;
702 fpeh.eat_fpe();
703
704 auto f64_options =
705 torch::TensorOptions().dtype(torch::kFloat64).device(torch::kCPU);
706
707 std::vector<metatomic_torch::System> systems;
708 std::vector<torch::Tensor> pos_tensors;
709 systems.reserve(static_cast<size_t>(nSystems));
710 pos_tensors.reserve(static_cast<size_t>(nSystems));
711
712 for (long s = 0; s < nSystems; s++) {
713 auto torch_positions = torch::from_blob(const_cast<double *>(positions[s]),
714 {nAtoms, 3}, f64_options)
715 .to(this->dtype_)
716 .to(this->device_)
717 .set_requires_grad(true);
718 pos_tensors.push_back(torch_positions);
719
720 auto torch_cell =
721 torch::from_blob(const_cast<double *>(boxes[s]), {3, 3}, f64_options)
722 .to(this->dtype_)
723 .to(this->device_);
724
725 auto cell_norms = torch::norm(torch_cell, 2, /*dim=*/1);
726 auto torch_pbc = cell_norms.abs() > 1e-9;
727 bool periodic[3] = {torch_pbc[0].item<bool>(), torch_pbc[1].item<bool>(),
728 torch_pbc[2].item<bool>()};
729
730 if (!atomicNrs[s]) {
731 throw std::runtime_error(
732 "[MetatomicPotential] `atomicNrs` must be provided.");
733 }
734 std::vector<int32_t> types_vec(atomicNrs[s], atomicNrs[s] + nAtoms);
735 auto atomic_types =
736 torch::tensor(types_vec, torch::TensorOptions().dtype(torch::kInt32))
737 .to(this->device_);
738
739 auto system = torch::make_intrusive<metatomic_torch::SystemHolder>(
740 atomic_types, torch_positions, torch_cell, torch_pbc);
741
742 // Compute and register neighbor lists for this system
743 for (const auto &request : this->nl_requests_) {
744 auto neighbors = this->computeNeighbors(request, nAtoms, positions[s],
745 boxes[s], periodic);
746 metatomic_torch::register_autograd_neighbors(system, neighbors,
747 this->check_consistency_);
748 system->add_neighbor_list(request, neighbors);
749 }
750
751 systems.push_back(system);
752 }
753
754 // Single batched forward pass
755 metatensor_torch::TensorMap output_map;
756 try {
757 auto ivalue_output = this->model_.forward({
758 systems,
760 this->check_consistency_,
761 });
762 auto dict_output = ivalue_output.toGenericDict();
763 output_map = dict_output.at(this->energy_key_)
764 .toCustomClass<metatensor_torch::TensorMapHolder>();
765 } catch (const std::exception &e) {
766 QUILL_LOG_ERROR(m_log,
767 "[MetatomicPotential] Batched model evaluation failed: {}",
768 e.what());
769 throw;
770 }
771
772 // Extract per-system energies from the output TensorMap.
773 // For per-atom output: samples have ["system", "atom"] dimensions.
774 // For system-level output: samples have ["system"] dimension.
775 // In both cases, sum over all non-system dimensions to get per-system energy.
776 auto energy_block =
777 metatensor_torch::TensorMapHolder::block_by_id(output_map, 0);
778 auto energy_values = energy_block->values();
779 auto samples = energy_block->samples();
780
781 // Check if this is per-atom or per-system output
782 bool per_atom = samples->size() > 0 && samples->names().size() > 1;
783
784 if (per_atom) {
785 // Per-atom output: sum energies by system index
786 auto system_col = samples->column("system").to(torch::kCPU);
787 auto flat_energies = energy_values.reshape({-1}).to(torch::kCPU);
788
789 // Sum per-system energies for output
790 for (long s = 0; s < nSystems; s++) {
791 auto mask = (system_col == s);
792 energies[s] =
793 flat_energies.index({mask}).sum().to(torch::kFloat64).item<double>();
794 }
795
796 // Backward: sum ALL energies, single backward call.
797 // Each system's positions.grad() gets only its own contribution
798 // because energy_i depends only on positions_i.
799 energy_values.sum().backward();
800 } else {
801 // System-level output: values shape is (nSystems, 1) or similar
802 auto cpu_energies = energy_values.to(torch::kCPU).to(torch::kFloat64);
803 for (long s = 0; s < nSystems; s++) {
804 energies[s] = cpu_energies[s].sum().item<double>();
805 }
806 energy_values.backward(torch::ones_like(energy_values));
807 }
808
809 // Extract per-system forces from position gradients
810 for (long s = 0; s < nSystems; s++) {
811 auto positions_grad = pos_tensors[s].grad();
812 auto forces_tensor = -positions_grad.to(torch::kCPU).to(torch::kFloat64);
813 std::memcpy(forces[s], forces_tensor.contiguous().data_ptr<double>(),
814 nAtoms * 3 * sizeof(double));
815 }
816
817 // Variances: not yet supported in batched path
818 if (variances) {
819 for (long s = 0; s < nSystems; s++) {
820 variances[s] = 0.0;
821 }
822 }
823
824 fpeh.restore_fpe();
825}

◆ isThreadSafe()

bool MetatomicPotential::isThreadSafe ( ) const
inlinenodiscardoverridevirtualnoexcept

Single shared instance, serialized via mutex.

Sequential evaluation through computePotential() ensures correct force counting, removeNetForce, and PotRegistry bookkeeping. JIT profiling is disabled so all calls on the same instance produce deterministic results.

Reimplemented from eonc::Potential.

Definition at line 140 of file MetatomicPotential.h.

140{ return false; }

◆ needsPerImageInstance()

bool MetatomicPotential::needsPerImageInstance ( ) const
inlinenodiscardoverridevirtualnoexcept

Whether NEB should create separate Potential instances per image for true parallel force evaluation.

When true, NEB calls makePotential() once per image instead of sharing one instance. Override in potentials that use internal mutexes (e.g. MetatomicPotential).

Reimplemented from eonc::Potential.

Definition at line 141 of file MetatomicPotential.h.

141 {
142 return false;
143 }

◆ supportsBatchEvaluation()

bool MetatomicPotential::supportsBatchEvaluation ( ) const
inlinenodiscardoverridevirtualnoexcept

Batched evaluation via single shared instance.

Processes each system sequentially through the same model (identical to N force() calls but bypasses Matter::computePotential overhead). Forces, energies, and forceCallCounter are handled correctly.

Reimplemented from eonc::Potential.

Definition at line 149 of file MetatomicPotential.h.

149 {
150 return true;
151 }

Member Data Documentation

◆ capabilities_

metatomic_torch::ModelCapabilities MetatomicPotential::capabilities_
private

Definition at line 60 of file MetatomicPotential.h.

◆ check_consistency_

bool MetatomicPotential::check_consistency_
private

Definition at line 67 of file MetatomicPotential.h.

◆ device_

torch::Device MetatomicPotential::device_
private

Definition at line 66 of file MetatomicPotential.h.

◆ device_type_

c10::DeviceType MetatomicPotential::device_type_
private

Definition at line 65 of file MetatomicPotential.h.

◆ dtype_

torch::ScalarType MetatomicPotential::dtype_
private

Definition at line 64 of file MetatomicPotential.h.

◆ energy_key_

std::string MetatomicPotential::energy_key_
private

Definition at line 69 of file MetatomicPotential.h.

◆ energy_uncertainty_key_

std::string MetatomicPotential::energy_uncertainty_key_
private

Definition at line 70 of file MetatomicPotential.h.

◆ evaluations_options_

metatomic_torch::ModelEvaluationOptions MetatomicPotential::evaluations_options_
private

Definition at line 62 of file MetatomicPotential.h.

◆ inference_mutex_

std::mutex MetatomicPotential::inference_mutex_
mutableprivate

Definition at line 165 of file MetatomicPotential.h.

◆ m_log

eonc::log::Scoped MetatomicPotential::m_log
private

Definition at line 56 of file MetatomicPotential.h.

◆ m_metatomic_opts

eonc::Parameters::metatomic_options_t MetatomicPotential::m_metatomic_opts
private

Definition at line 57 of file MetatomicPotential.h.

◆ model_

metatensor_torch::Module MetatomicPotential::model_
private

Definition at line 59 of file MetatomicPotential.h.

◆ n_symmetry_rotations_

long MetatomicPotential::n_symmetry_rotations_ {0}
private

Definition at line 74 of file MetatomicPotential.h.

74{0};

◆ nc_forces_key_

std::string MetatomicPotential::nc_forces_key_
private

Definition at line 71 of file MetatomicPotential.h.

◆ nl_requests_

std::vector<metatomic_torch::NeighborListOptions> MetatomicPotential::nl_requests_
private

Definition at line 61 of file MetatomicPotential.h.

◆ non_conservative_

bool MetatomicPotential::non_conservative_ {false}
private

Definition at line 72 of file MetatomicPotential.h.

72{false};

◆ random_rotation_

bool MetatomicPotential::random_rotation_ {false}
private

Definition at line 73 of file MetatomicPotential.h.

73{false};

◆ uncertainty_threshold_

double MetatomicPotential::uncertainty_threshold_ {-1.0}
private

Definition at line 77 of file MetatomicPotential.h.

77{-1.0};

The documentation for this class was generated from the following files: