From 34d1e17cc8723beb5de2e4bf022a1188893cec1e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ole=20Sch=C3=BCtt?= Date: Sat, 25 Jan 2025 15:32:33 +0100 Subject: [PATCH] Torch: Rename torch_model_eval -> torch_model_forward --- src/manybody_allegro.F | 4 ++-- src/manybody_nequip.F | 4 ++-- src/nequip_unittest.F | 4 ++-- src/pao_model.F | 4 ++-- src/torch_api.F | 18 +++++++++--------- src/torch_c_api.cpp | 5 ++--- 6 files changed, 19 insertions(+), 20 deletions(-) diff --git a/src/manybody_allegro.F b/src/manybody_allegro.F index ec49bfa30b..69334f7b60 100644 --- a/src/manybody_allegro.F +++ b/src/manybody_allegro.F @@ -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) diff --git a/src/manybody_nequip.F b/src/manybody_nequip.F index 0ab5487f00..63e71fac79 100644 --- a/src/manybody_nequip.F +++ b/src/manybody_nequip.F @@ -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) diff --git a/src/nequip_unittest.F b/src/nequip_unittest.F index cc2700031a..4fb943fb4e 100644 --- a/src/nequip_unittest.F +++ b/src/nequip_unittest.F @@ -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) diff --git a/src/pao_model.F b/src/pao_model.F index f87389cd4a..e3102297ba 100644 --- a/src/pao_model.F +++ b/src/pao_model.F @@ -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) diff --git a/src/torch_api.F b/src/torch_api.F index 1e407ac9c9..3ead41b295 100644 --- a/src/torch_api.F +++ b/src/torch_api.F @@ -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. diff --git a/src/torch_c_api.cpp b/src/torch_c_api.cpp index 7763594fe5..19a13f7499 100644 --- a/src/torch_c_api.cpp +++ b/src/torch_c_api.cpp @@ -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();