59 model_(torch::jit::Module()),
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,
77 at::globalContext().setBenchmarkCuDNN(false);
79 "[MetatomicPotential] Deterministic algorithms enabled "
84 "[MetatomicPotential] Deterministic algorithms disabled "
85 "(faster, may diverge across runs / NEB images)");
91 QUILL_LOG_INFO(
m_log,
"[MetatomicPotential] Initializing...");
94 torch::optional<std::string> extensions_directory = torch::nullopt;
96 extensions_directory = m_metatomic_opts.extensions_directory;
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: {}",
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()) {
118 .toCustomClass<metatomic_torch::NeighborListOptionsHolder>();
119 this->nl_requests_.push_back(request);
123 torch::optional<std::string> desired = torch::nullopt;
125 desired = m_metatomic_opts.device;
128 this->capabilities_->supported_devices, desired);
131 QUILL_LOG_INFO(
m_log,
"[MetatomicPotential] Using device: {}",
device_.str());
137 this->dtype_ = torch::kFloat64;
139 this->dtype_ = torch::kFloat32;
141 throw std::runtime_error(
"Unsupported dtype: " +
142 this->capabilities_->dtype());
144 QUILL_LOG_INFO(
m_log,
"[MetatomicPotential] Using dtype: {}",
149 auto outputs = this->capabilities_->outputs();
164 this->energy_key_ = m_metatomic_opts.energy_output;
167 metatomic_torch::pick_output(
"energy", outputs, v_energy);
176 if (!outputs.contains(this->nc_forces_key_)) {
177 throw std::runtime_error(
178 "Missing explicit force_output in metatomic model: " +
183 "non_conservative_force", outputs, v_force);
185 QUILL_LOG_INFO(
m_log,
186 "[MetatomicPotential] Non-conservative forces from '{}'",
190 QUILL_LOG_INFO(
m_log,
191 "[MetatomicPotential] Symmetry averaging over {} rotations",
195 m_log,
"[MetatomicPotential] Per-call random SO(3) rotation enabled");
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);
206 if (!outputs.contains(this->energy_key_)) {
209 "[MetatomicPotential] The model does not provide an '{}' output.",
211 throw std::runtime_error(
"Missing energy output in metatomic model");
215 this->evaluations_options_ =
216 torch::make_intrusive<metatomic_torch::ModelEvaluationOptionsHolder>();
219 auto model_output = outputs.at(this->energy_key_);
220 auto requested_output =
221 torch::make_intrusive<metatomic_torch::ModelOutputHolder>();
224 requested_output->set_sample_kind(model_output->sample_kind());
225 requested_output->explicit_gradients = {};
226 requested_output->set_unit(
"eV");
231 auto nc_info = outputs.at(this->nc_forces_key_);
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");
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) {
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_);
260 this->energy_uncertainty_key_ = metatomic_torch::pick_output(
261 "energy_uncertainty", outputs, v_energy_uq);
262 } catch (
const std::exception &e) {
264 m_log,
"[MetatomicPotential] No uncertainty output available: {}",
266 this->uncertainty_threshold_ = -1.0;
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_);
286 QUILL_LOG_DEBUG(m_log,
287 "[MetatomicPotential] Model provides '{}' "
288 "but sample_kind is not \"atom\"; skipping uncertainty "
290 this->energy_uncertainty_key_);
291 this->uncertainty_threshold_ = -1.0;
296 this->check_consistency_ = m_metatomic_opts.check_consistency;
297 QUILL_LOG_INFO(m_log,
"[MetatomicPotential] Initialization complete.");
353 const int *atomicNrs,
double *forces,
354 double *energy,
double *variance,
364 throw std::runtime_error(
365 "[MetatomicPotential] `atomicNrs` must be provided.");
368 const bool use_rotation =
371 const long n_passes =
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));
380 double energy_acc = 0.0;
381 auto forces_acc = torch::zeros({nAtoms, 3}, f64_options);
382 bool variance_set =
false;
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_));
391 auto R_cpu = R.to(torch::kCPU).to(torch::kFloat64);
392 auto R_T = R.transpose(0, 1);
394 auto pos_cpu = torch::from_blob(
const_cast<double *
>(positions),
395 {nAtoms, 3}, f64_options)
398 torch::from_blob(
const_cast<double *
>(box), {3, 3}, f64_options)
401 pos_cpu = pos_cpu.matmul(R_cpu.transpose(0, 1));
403 cell_cpu = cell_cpu.matmul(R_cpu.transpose(0, 1));
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>(),
413 auto torch_positions =
414 torch::from_blob(pos_buf.data(), {nAtoms, 3}, f64_options)
419 auto torch_cell = torch::from_blob(cell_buf.data(), {3, 3}, f64_options)
423 auto cell_norms = torch::norm(torch_cell, 2, 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>()};
428 auto atomic_types = atomic_types_cpu.to(this->
device_);
430 auto system = torch::make_intrusive<metatomic_torch::SystemHolder>(
431 atomic_types, torch_positions, torch_cell, torch_pbc);
435 cell_buf.data(), periodic);
436 metatomic_torch::register_autograd_neighbors(system, neighbors,
438 system->add_neighbor_list(request, neighbors);
441 torch::Tensor forces_tensor;
443 auto ivalue_output = this->
model_.forward({
444 std::vector<metatomic_torch::System>{system},
448 auto dict_output = ivalue_output.toGenericDict();
449 auto output_map = dict_output.at(this->
energy_key_)
450 .toCustomClass<metatensor_torch::TensorMapHolder>();
454 if (dict_output.contains(this->energy_uncertainty_key_)) {
455 auto uncertainty_map =
457 .toCustomClass<metatensor_torch::TensorMapHolder>();
458 auto uncertainty_block =
459 metatensor_torch::TensorMapHolder::block_by_id(uncertainty_map,
461 auto flat_uncertainty =
462 uncertainty_block->values().reshape({-1}).to(torch::kCPU);
463 if (variance !=
nullptr && flat_uncertainty.numel() > 0) {
466 flat_uncertainty.to(torch::kFloat64).mean().item<
double>();
469 QUILL_LOG_DEBUG(
m_log,
470 "[MetatomicPotential] Failed to compute mean "
471 "uncertainty for variance.");
474 auto atoms_above_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) {
487 ss << atom_indices_above[i].item<int32_t>();
490 if (atom_indices_above.size(0) > n_report) {
491 ss <<
" and " << (atom_indices_above.size(0) - n_report)
496 "[MetatomicPotential] The uncertainty on atomic energies for "
497 "{} are larger than the threshold of {}. (Key: {}) Be "
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_);
505 }
catch (
const std::exception &e) {
506 QUILL_LOG_WARNING(
m_log,
507 "[MetatomicPotential] Failed to check {}: {}",
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>();
519 .toCustomClass<metatensor_torch::TensorMapHolder>();
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));
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");
538 nc_vals.reshape({nAtoms, 3}).to(torch::kCPU).to(torch::kFloat64);
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);
544 }
catch (
const std::exception &e) {
545 QUILL_LOG_ERROR(
m_log,
"[MetatomicPotential] Model evaluation failed: {}",
553 forces_tensor = forces_tensor.matmul(R_cpu);
555 forces_acc += forces_tensor;
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;
563 std::memcpy(forces, forces_acc.contiguous().data_ptr<
double>(),
564 nAtoms * 3 *
sizeof(
double));
572 metatomic_torch::NeighborListOptions request,
long nAtoms,
573 const double *positions,
const double *box,
const bool periodic[3]) {
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;
585 options.return_vectors =
true;
587 VesinNeighborList *vesin_neighbor_list =
new VesinNeighborList();
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);
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);
602 err_str +=
" (no message; vesin header/lib ABI mismatch? need vesin>=0.6 "
603 "with matching engine)";
605 delete vesin_neighbor_list;
606 throw std::runtime_error(err_str);
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);
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];
627 auto deleter = [=](
void *) {
628 vesin_free(vesin_neighbor_list);
629 delete vesin_neighbor_list;
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));
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_));
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));
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);
694 const double *
const *positions,
695 const int *
const *atomicNrs,
696 double *
const *forces,
697 double *energies,
double *variances,
698 const double *
const *boxes) {
705 torch::TensorOptions().dtype(torch::kFloat64).device(torch::kCPU);
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));
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)
717 .set_requires_grad(
true);
718 pos_tensors.push_back(torch_positions);
721 torch::from_blob(
const_cast<double *
>(boxes[s]), {3, 3}, f64_options)
725 auto cell_norms = torch::norm(torch_cell, 2, 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>()};
731 throw std::runtime_error(
732 "[MetatomicPotential] `atomicNrs` must be provided.");
734 std::vector<int32_t> types_vec(atomicNrs[s], atomicNrs[s] + nAtoms);
736 torch::tensor(types_vec, torch::TensorOptions().dtype(torch::kInt32))
739 auto system = torch::make_intrusive<metatomic_torch::SystemHolder>(
740 atomic_types, torch_positions, torch_cell, torch_pbc);
746 metatomic_torch::register_autograd_neighbors(system, neighbors,
748 system->add_neighbor_list(request, neighbors);
751 systems.push_back(system);
755 metatensor_torch::TensorMap output_map;
757 auto ivalue_output = this->
model_.forward({
762 auto dict_output = ivalue_output.toGenericDict();
764 .toCustomClass<metatensor_torch::TensorMapHolder>();
765 }
catch (
const std::exception &e) {
766 QUILL_LOG_ERROR(
m_log,
767 "[MetatomicPotential] Batched model evaluation failed: {}",
777 metatensor_torch::TensorMapHolder::block_by_id(output_map, 0);
778 auto energy_values = energy_block->values();
779 auto samples = energy_block->samples();
782 bool per_atom = samples->size() > 0 && samples->names().size() > 1;
786 auto system_col = samples->column(
"system").to(torch::kCPU);
787 auto flat_energies = energy_values.reshape({-1}).to(torch::kCPU);
790 for (
long s = 0; s < nSystems; s++) {
791 auto mask = (system_col == s);
793 flat_energies.index({mask}).sum().to(torch::kFloat64).item<
double>();
799 energy_values.sum().backward();
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>();
806 energy_values.backward(torch::ones_like(energy_values));
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));
819 for (
long s = 0; s < nSystems; s++) {