mirror of
https://github.com/cp2k/cp2k.git
synced 2026-07-28 22:25:32 -04:00
Torch: Rename torch_model_eval -> torch_model_forward
This commit is contained in:
parent
a26eb0bcc9
commit
34d1e17cc8
6 changed files with 19 additions and 20 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue