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.
355 {
356
357
359
360 eonc::FPEHandler fpeh;
362
363 if (!atomicNrs) {
364 throw std::runtime_error(
365 "[MetatomicPotential] `atomicNrs` must be provided.");
366 }
367
368 const bool use_rotation =
370
371 const long n_passes =
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) {
389 }
390
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
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)
418
419 auto torch_cell = torch::from_blob(cell_buf.data(), {3, 3}, f64_options)
422
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>()};
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
435 cell_buf.data(), periodic);
436 metatomic_torch::register_autograd_neighbors(system, neighbors,
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},
447 });
448 auto dict_output = ivalue_output.toGenericDict();
449 auto output_map = dict_output.at(this->
energy_key_)
450 .toCustomClass<metatensor_torch::TensorMapHolder>();
451
453 try {
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,
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 =
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(
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 {}: {}",
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
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
551
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
567}