Torch: Rename torch_model_eval -> torch_model_forward

This commit is contained in:
Ole Schütt 2025-01-25 15:32:33 +01:00 committed by Ole Schütt
parent a26eb0bcc9
commit 34d1e17cc8
6 changed files with 19 additions and 20 deletions

View file

@ -32,7 +32,7 @@ MODULE manybody_allegro
USE particle_types, ONLY: particle_type
USE torch_api, ONLY: &
torch_dict_create, torch_dict_get, torch_dict_insert, torch_dict_release, torch_dict_type, &
torch_model_eval, torch_model_freeze, torch_model_load, torch_tensor_data_ptr, &
torch_model_forward, torch_model_freeze, torch_model_load, torch_tensor_data_ptr, &
torch_tensor_from_array, torch_tensor_release, torch_tensor_type
USE util, ONLY: sort
#include "./base/base_uses.f90"
@ -450,7 +450,7 @@ CONTAINS
CALL torch_tensor_release(atom_types_tensor)
CALL torch_dict_create(outputs)
CALL torch_model_eval(allegro_data%model, inputs, outputs)
CALL torch_model_forward(allegro_data%model, inputs, outputs)
pot_allegro = 0.0_dp
CALL torch_dict_get(outputs, "atomic_energy", atomic_energy_tensor)

View file

@ -32,7 +32,7 @@ MODULE manybody_nequip
USE particle_types, ONLY: particle_type
USE torch_api, ONLY: &
torch_dict_create, torch_dict_get, torch_dict_insert, torch_dict_release, torch_dict_type, &
torch_model_eval, torch_model_freeze, torch_model_load, torch_tensor_data_ptr, &
torch_model_forward, torch_model_freeze, torch_model_load, torch_tensor_data_ptr, &
torch_tensor_from_array, torch_tensor_release, torch_tensor_type
USE util, ONLY: sort
#include "./base/base_uses.f90"
@ -450,7 +450,7 @@ CONTAINS
CALL torch_tensor_release(atom_types_tensor)
CALL torch_dict_create(outputs)
CALL torch_model_eval(nequip_data%model, inputs, outputs)
CALL torch_model_forward(nequip_data%model, inputs, outputs)
CALL torch_dict_get(outputs, "total_energy", total_energy_tensor)
CALL torch_dict_get(outputs, "atomic_energy", atomic_energy_tensor)

View file

@ -17,7 +17,7 @@ PROGRAM nequip_unittest
evolt
USE torch_api, ONLY: &
torch_cuda_is_available, torch_dict_create, torch_dict_get, torch_dict_insert, &
torch_dict_release, torch_dict_type, torch_model_eval, torch_model_load, &
torch_dict_release, torch_dict_type, torch_model_forward, torch_model_load, &
torch_model_read_metadata, torch_model_release, torch_model_type, torch_tensor_data_ptr, &
torch_tensor_from_array, torch_tensor_release, torch_tensor_type
#include "./base/base_uses.f90"
@ -197,7 +197,7 @@ PROGRAM nequip_unittest
CALL torch_dict_insert(inputs, "atom_types", atom_types_tensor)
CALL torch_tensor_release(atom_types_tensor)
CALL torch_model_eval(model, inputs, outputs)
CALL torch_model_forward(model, inputs, outputs)
CALL torch_dict_get(outputs, "total_energy", total_energy_tensor)
CALL torch_tensor_data_ptr(total_energy_tensor, total_energy)

View file

@ -35,7 +35,7 @@ MODULE pao_model
qs_kind_type
USE torch_api, ONLY: &
torch_dict_create, torch_dict_get, torch_dict_insert, torch_dict_release, torch_dict_type, &
torch_model_eval, torch_model_get_attr, torch_model_load, torch_tensor_data_ptr, &
torch_model_forward, torch_model_get_attr, torch_model_load, torch_tensor_data_ptr, &
torch_tensor_from_array, torch_tensor_release, torch_tensor_type
USE util, ONLY: sort
#include "./base/base_uses.f90"
@ -257,7 +257,7 @@ CONTAINS
CALL torch_tensor_release(neighbors_features_tensor)
CALL torch_dict_create(model_outputs)
CALL torch_model_eval(model%torch_model, model_inputs, model_outputs)
CALL torch_model_forward(model%torch_model, model_inputs, model_outputs)
! Copy predicted XBlock.
NULLIFY (predicted_xblock)

View file

@ -67,7 +67,7 @@ MODULE torch_api
PUBLIC :: torch_tensor_type, torch_tensor_from_array, torch_tensor_data_ptr, torch_tensor_release
PUBLIC :: torch_dict_type, torch_dict_create, torch_dict_insert, torch_dict_get, torch_dict_release
PUBLIC :: torch_model_type, torch_model_load, torch_model_eval, torch_model_release
PUBLIC :: torch_model_type, torch_model_load, torch_model_forward, torch_model_release
PUBLIC :: torch_model_get_attr, torch_model_read_metadata
PUBLIC :: torch_cuda_is_available, torch_allow_tf32, torch_model_freeze
@ -340,37 +340,37 @@ CONTAINS
END SUBROUTINE torch_model_load
! **************************************************************************************************
!> \brief Evaluates the given Torch model. (In Torch lingo this operation is called forward())
!> \brief Evaluates the given Torch model.
!> \author Ole Schuett
! **************************************************************************************************
SUBROUTINE torch_model_eval(model, inputs, outputs)
SUBROUTINE torch_model_forward(model, inputs, outputs)
TYPE(torch_model_type), INTENT(INOUT) :: model
TYPE(torch_dict_type), INTENT(IN) :: inputs
TYPE(torch_dict_type), INTENT(INOUT) :: outputs
#if defined(__LIBTORCH)
INTERFACE
SUBROUTINE torch_c_model_eval(model, inputs, outputs) BIND(C, name="torch_c_model_eval")
SUBROUTINE torch_c_model_forward(model, inputs, outputs) BIND(C, name="torch_c_model_forward")
IMPORT :: C_PTR
TYPE(C_PTR), VALUE :: model
TYPE(C_PTR), VALUE :: inputs
TYPE(C_PTR), VALUE :: outputs
END SUBROUTINE torch_c_model_eval
END SUBROUTINE torch_c_model_forward
END INTERFACE
CPASSERT(C_ASSOCIATED(model%c_ptr))
CPASSERT(C_ASSOCIATED(inputs%c_ptr))
CPASSERT(C_ASSOCIATED(outputs%c_ptr))
CALL torch_c_model_eval(model=model%c_ptr, &
inputs=inputs%c_ptr, &
outputs=outputs%c_ptr)
CALL torch_c_model_forward(model=model%c_ptr, &
inputs=inputs%c_ptr, &
outputs=outputs%c_ptr)
#else
CPABORT("CP2K was compiled without Torch library.")
MARK_USED(model)
MARK_USED(inputs)
MARK_USED(outputs)
#endif
END SUBROUTINE torch_model_eval
END SUBROUTINE torch_model_forward
! **************************************************************************************************
!> \brief Releases a Torch model and all its ressources.

View file

@ -175,11 +175,10 @@ void torch_c_model_load(torch_c_model_t **model_out, const char *filename) {
/*******************************************************************************
* \brief Evaluates the given Torch model.
* In Torch lingo this operation is called forward().
* \author Ole Schuett
******************************************************************************/
void torch_c_model_eval(torch_c_model_t *model, const torch_c_dict_t *inputs,
torch_c_dict_t *outputs) {
void torch_c_model_forward(torch_c_model_t *model, const torch_c_dict_t *inputs,
torch_c_dict_t *outputs) {
auto untyped_output = model->forward({*inputs}).toGenericDict();
outputs->clear();