Loading...
Searching...
No Matches
MetatomicPotential.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*/
13#include "eon/Parameters.h"
14#include "eon/fpe_handler.h"
15#include "vesin.h"
16
17#include <torch/csrc/jit/runtime/graph_executor.h>
18
19#include <cstdint>
20#include <random>
21#include <sstream>
22#include <string>
23#include <vector>
24
25namespace {
26
27// metatensor-torch 0.10.3 Module::to() walks every attribute when
28// `_mts_buffer_names` is missing. Exported PET-MAD stores mixed dicts as
29// ordinary attrs; empty containers count as non-metatensor and throw.
30// Weights already moved. Swallow only that mixed-dict error. Do not
31// register `_mts_buffer_names` on scripted modules (JIT slot assert).
32bool is_mixed_mts_to_error(const c10::Error &e) {
33 const std::string w = e.what_without_backtrace();
34 return w.find("metatensor and non-metatensor") != std::string::npos;
35}
36
37void move_atomistic_model(metatensor_torch::Module &model,
38 torch::Device device) {
39 try {
40 model.to(device);
41 } catch (const c10::Error &e) {
42 if (!is_mixed_mts_to_error(e)) {
43 throw;
44 }
45 }
46}
47
48} // namespace
49
50static torch::optional<std::string> normalize_variant(const std::string &s) {
51 if (s.empty() || s == "off")
52 return torch::nullopt;
53 return s;
54}
55
57 : eonc::Potential(eonc::PotType::METATOMIC),
58 m_metatomic_opts{params.metatomic_options()},
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
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
215 this->evaluations_options_ =
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).
247 this->energy_uncertainty_key_ =
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.",
253 this->energy_uncertainty_key_);
254 throw std::runtime_error(
255 "Missing explicit energy_uncertainty_output in metatomic model: " +
256 this->energy_uncertainty_key_);
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 = {})",
283 this->energy_uncertainty_key_,
284 this->uncertainty_threshold_);
285 } else {
286 QUILL_LOG_DEBUG(m_log,
287 "[MetatomicPotential] Model provides '{}' "
288 "but sample_kind is not \"atom\"; skipping uncertainty "
289 "checks.",
290 this->energy_uncertainty_key_);
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}
301
322
323std::shared_ptr<eonc::Potential> MetatomicPotential::clonePotential() const {
324 return std::shared_ptr<eonc::Potential>(
325 new MetatomicPotential(*this, CloneTag{}));
326}
327
328// --- helpers for random / symmetry rotations (#287, #292) ---
329
330namespace {
331
332// Uniform random rotation in SO(3) via QR of a Gaussian matrix with positive
333// determinant (Arvo / Shoemake style, sufficient for stochastic averaging).
334torch::Tensor random_so3(torch::Device device, torch::ScalarType dtype) {
335 auto A =
336 torch::randn({3, 3}, torch::TensorOptions().dtype(dtype).device(device));
337 auto qr = torch::linalg_qr(A);
338 auto Q = std::get<0>(qr);
339 auto R = std::get<1>(qr);
340 auto d = torch::sign(torch::diagonal(R));
341 Q = Q * d.unsqueeze(0);
342 if (torch::det(Q).item<double>() < 0) {
343 Q.select(1, 0).mul_(-1);
344 }
345 return Q;
346}
347
348} // namespace
349
350// --- MetatomicPotential::force ---
351
352void MetatomicPotential::force(long nAtoms, const double *positions,
353 const int *atomicNrs, double *forces,
354 double *energy, double *variance,
355 const double *box) {
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}
568
569// --- MetatomicPotential::computeNeighbors (helper) ---
570
571metatensor_torch::TensorBlock MetatomicPotential::computeNeighbors(
572 metatomic_torch::NeighborListOptions request, long nAtoms,
573 const double *positions, const double *box, const bool periodic[3]) {
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}
655
656// --- MetatomicPotential::forceBatch ---
657// Processes N systems sequentially through the same model instance.
658// Numerically identical to N individual force() calls. The single-instance
659// design avoids JIT profiling divergence and N model copies. True batched
660// model.forward({sys0..sysN}) is a future optimization (see #if 0 block below).
661
662void MetatomicPotential::forceBatch(long nSystems, long nAtoms,
663 const double *const *positions,
664 const int *const *atomicNrs,
665 double *const *forces, double *energies,
666 double *variances,
667 const double *const *boxes) {
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}
692
693void MetatomicPotential::forceBatchNative(long nSystems, long nAtoms,
694 const double *const *positions,
695 const int *const *atomicNrs,
696 double *const *forces,
697 double *energies, double *variances,
698 const double *const *boxes) {
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}
static torch::optional< std::string > normalize_variant(const std::string &s)
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)
std::vector< metatomic_torch::NeighborListOptions > nl_requests_
metatomic_torch::ModelEvaluationOptions evaluations_options_
eonc::log::Scoped m_log
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
metatomic_torch::ModelCapabilities capabilities_
MetatomicPotential(const eonc::Parameters &params)
Constructor for the MetatomicPotential.
std::string energy_uncertainty_key_
eonc::Parameters::metatomic_options_t m_metatomic_opts
c10::DeviceType device_type_
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.
std::shared_ptr< eonc::Potential > clonePotential() const override
Independent instance that does not reload from disk.
metatensor_torch::Module model_
torch::ScalarType dtype_
const main_options_t & main_options() const
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
Potential(PotType a_ptype)
Production default: construction-scope registry, else PotRegistry::get().
RAII resource manager for the ARTn C library with global synchronization.