LCOV - code coverage report
Current view: top level - src - skala_torch_api.F (source / functions) Coverage Total Hit
Test: CP2K Regtests (git:24d69ee) Lines: 61.9 % 63 39
Test Date: 2026-09-03 07:32:15 Functions: 33.3 % 9 3

            Line data    Source code
       1              : !--------------------------------------------------------------------------------------------------!
       2              : !   CP2K: A general program to perform molecular dynamics simulations                              !
       3              : !   Copyright 2000-2026 CP2K developers group <https://cp2k.org>                                   !
       4              : !                                                                                                  !
       5              : !   SPDX-License-Identifier: GPL-2.0-or-later                                                      !
       6              : !--------------------------------------------------------------------------------------------------!
       7              : 
       8              : ! **************************************************************************************************
       9              : !> \brief Small CP2K wrapper around the SKALA TorchScript functional protocol.
      10              : ! **************************************************************************************************
      11              : MODULE skala_torch_api
      12              : #if defined (__HAS_IEEE_EXCEPTIONS)
      13              :    USE ieee_exceptions, ONLY: ieee_all, &
      14              :                               ieee_get_halting_mode, &
      15              :                               ieee_set_halting_mode
      16              : #endif
      17              :    USE kinds, ONLY: default_string_length, &
      18              :                     dp
      19              :    USE string_utilities, ONLY: uppercase
      20              :    USE torch_api, ONLY: &
      21              :       torch_dict_type, torch_model_disable_parameter_gradients, torch_model_forward_mol_tensor, &
      22              :       torch_model_freeze_preserving_method, torch_model_load_with_metadata, torch_model_release, &
      23              :       torch_model_remap_device_constants, torch_model_type, &
      24              :       torch_tensor_item_double, torch_tensor_release, torch_tensor_type, &
      25              :       torch_tensor_weighted_sum
      26              : #include "./base/base_uses.f90"
      27              : 
      28              :    IMPLICIT NONE
      29              : 
      30              :    PRIVATE
      31              : 
      32              :    CHARACTER(len=*), PARAMETER, PRIVATE :: moduleN = 'skala_torch_api'
      33              : 
      34              :    PUBLIC :: skala_torch_model_type, skala_torch_model_load, skala_torch_model_release
      35              :    PUBLIC :: skala_torch_model_get_exc, skala_torch_model_get_exc_density
      36              :    PUBLIC :: skala_torch_model_needs_feature, skala_torch_model_protocol_version
      37              : 
      38              :    TYPE skala_torch_model_type
      39              :       PRIVATE
      40              :       INTEGER                                            :: protocol_version = -1
      41              :       CHARACTER(len=default_string_length), ALLOCATABLE, &
      42              :          DIMENSION(:)                                    :: features
      43              :       TYPE(torch_model_type)                             :: torch_model
      44              :    END TYPE skala_torch_model_type
      45              : 
      46              : CONTAINS
      47              : 
      48              : ! **************************************************************************************************
      49              : !> \brief Load a SKALA TorchScript model and its feature metadata.
      50              : !> \param model ...
      51              : !> \param filename ...
      52              : ! **************************************************************************************************
      53           91 :    SUBROUTINE skala_torch_model_load(model, filename)
      54              :       TYPE(skala_torch_model_type), INTENT(INOUT)        :: model
      55              :       CHARACTER(len=*), INTENT(IN)                       :: filename
      56              : 
      57           91 :       CHARACTER(:), ALLOCATABLE                          :: features_json, protocol_string
      58              :       INTEGER                                            :: ios
      59              : 
      60              :       CALL torch_model_load_with_metadata(model%torch_model, filename, &
      61              :                                           "protocol_version", protocol_string, &
      62           91 :                                           "features", features_json)
      63           91 :       CALL torch_model_remap_device_constants(model%torch_model)
      64           91 :       CALL torch_model_disable_parameter_gradients(model%torch_model)
      65           91 :       READ (protocol_string, *, IOSTAT=ios) model%protocol_version
      66           91 :       IF (ios /= 0) CPABORT("Could not parse SKALA TorchScript protocol_version metadata")
      67           91 :       IF (model%protocol_version /= 2) THEN
      68            0 :          CPABORT("Unsupported SKALA TorchScript protocol version")
      69              :       END IF
      70              : 
      71           91 :       CALL parse_feature_list(features_json, model%features)
      72              :       ! Preserve the exported SKALA entry point while folding constant model state.
      73           91 :       CALL torch_model_freeze_preserving_method(model%torch_model, "get_exc_density")
      74              : 
      75           91 :    END SUBROUTINE skala_torch_model_load
      76              : 
      77              : ! **************************************************************************************************
      78              : !> \brief Release a loaded SKALA TorchScript model.
      79              : !> \param model ...
      80              : ! **************************************************************************************************
      81            0 :    SUBROUTINE skala_torch_model_release(model)
      82              :       TYPE(skala_torch_model_type), INTENT(INOUT)        :: model
      83              : 
      84            0 :       CALL torch_model_release(model%torch_model)
      85            0 :       IF (ALLOCATED(model%features)) DEALLOCATE (model%features)
      86            0 :       model%protocol_version = -1
      87              : 
      88            0 :    END SUBROUTINE skala_torch_model_release
      89              : 
      90              : ! **************************************************************************************************
      91              : !> \brief Check whether a loaded SKALA model requests a feature.
      92              : !> \param model ...
      93              : !> \param feature ...
      94              : !> \return ...
      95              : ! **************************************************************************************************
      96            0 :    FUNCTION skala_torch_model_needs_feature(model, feature) RESULT(needs_feature)
      97              :       TYPE(skala_torch_model_type), INTENT(IN)           :: model
      98              :       CHARACTER(len=*), INTENT(IN)                       :: feature
      99              :       LOGICAL                                            :: needs_feature
     100              : 
     101              :       CHARACTER(len=default_string_length)               :: feature_key, model_feature
     102              :       INTEGER                                            :: i
     103              : 
     104            0 :       feature_key = ADJUSTL(feature)
     105            0 :       CALL uppercase(feature_key)
     106              : 
     107            0 :       needs_feature = .FALSE.
     108            0 :       IF (.NOT. ALLOCATED(model%features)) RETURN
     109              : 
     110            0 :       DO i = 1, SIZE(model%features)
     111            0 :          model_feature = ADJUSTL(model%features(i))
     112            0 :          CALL uppercase(model_feature)
     113            0 :          IF (TRIM(model_feature) == TRIM(feature_key)) THEN
     114            0 :             needs_feature = .TRUE.
     115              :             RETURN
     116              :          END IF
     117              :       END DO
     118              : 
     119            0 :    END FUNCTION skala_torch_model_needs_feature
     120              : 
     121              : ! **************************************************************************************************
     122              : !> \brief Return the loaded SKALA TorchScript protocol version.
     123              : !> \param model ...
     124              : !> \return ...
     125              : ! **************************************************************************************************
     126            0 :    FUNCTION skala_torch_model_protocol_version(model) RESULT(protocol_version)
     127              :       TYPE(skala_torch_model_type), INTENT(IN)           :: model
     128              :       INTEGER                                            :: protocol_version
     129              : 
     130            0 :       protocol_version = model%protocol_version
     131              : 
     132            0 :    END FUNCTION skala_torch_model_protocol_version
     133              : 
     134              : ! **************************************************************************************************
     135              : !> \brief Evaluate the SKALA exchange-correlation energy density.
     136              : !> \param model ...
     137              : !> \param inputs ...
     138              : !> \param exc_density ...
     139              : ! **************************************************************************************************
     140            0 :    SUBROUTINE skala_torch_model_get_exc_density(model, inputs, exc_density)
     141              :       TYPE(skala_torch_model_type), INTENT(INOUT)        :: model
     142              :       TYPE(torch_dict_type), INTENT(IN)                  :: inputs
     143              :       TYPE(torch_tensor_type), INTENT(INOUT)             :: exc_density
     144              : 
     145              : #if defined (__HAS_IEEE_EXCEPTIONS)
     146              :       LOGICAL, DIMENSION(5)                              :: ieee_halt
     147              : 
     148              :       CALL ieee_get_halting_mode(IEEE_ALL, ieee_halt)
     149              :       CALL ieee_set_halting_mode(IEEE_ALL, .FALSE.)
     150              : #endif
     151            0 :       CALL torch_model_forward_mol_tensor(model%torch_model, "get_exc_density", inputs, exc_density)
     152              : #if defined (__HAS_IEEE_EXCEPTIONS)
     153              :       CALL ieee_set_halting_mode(IEEE_ALL, ieee_halt)
     154              : #endif
     155              : 
     156            0 :    END SUBROUTINE skala_torch_model_get_exc_density
     157              : 
     158              : ! **************************************************************************************************
     159              : !> \brief Evaluate the weighted SKALA exchange-correlation energy.
     160              : !> \param model ...
     161              : !> \param inputs ...
     162              : !> \param grid_weights ...
     163              : !> \param exc_tensor ...
     164              : !> \param exc ...
     165              : ! **************************************************************************************************
     166          296 :    SUBROUTINE skala_torch_model_get_exc(model, inputs, grid_weights, exc_tensor, exc)
     167              :       TYPE(skala_torch_model_type), INTENT(INOUT)        :: model
     168              :       TYPE(torch_dict_type), INTENT(IN)                  :: inputs
     169              :       TYPE(torch_tensor_type), INTENT(IN)                :: grid_weights
     170              :       TYPE(torch_tensor_type), INTENT(INOUT)             :: exc_tensor
     171              :       REAL(KIND=dp), INTENT(OUT)                         :: exc
     172              : 
     173              :       TYPE(torch_tensor_type)                            :: exc_density
     174              : 
     175              : #if defined (__HAS_IEEE_EXCEPTIONS)
     176              :       LOGICAL, DIMENSION(5)                              :: ieee_halt
     177              : 
     178              :       CALL ieee_get_halting_mode(IEEE_ALL, ieee_halt)
     179              :       CALL ieee_set_halting_mode(IEEE_ALL, .FALSE.)
     180              : #endif
     181          296 :       CALL torch_model_forward_mol_tensor(model%torch_model, "get_exc_density", inputs, exc_density)
     182          296 :       CALL torch_tensor_weighted_sum(exc_density, grid_weights, exc_tensor)
     183          296 :       CALL torch_tensor_release(exc_density)
     184          296 :       exc = torch_tensor_item_double(exc_tensor)
     185              : #if defined (__HAS_IEEE_EXCEPTIONS)
     186              :       CALL ieee_set_halting_mode(IEEE_ALL, ieee_halt)
     187              : #endif
     188              : 
     189          296 :    END SUBROUTINE skala_torch_model_get_exc
     190              : 
     191              : ! **************************************************************************************************
     192              : !> \brief Parse a TorchScript extra_files JSON list of feature names.
     193              : !> \param features_json ...
     194              : !> \param features ...
     195              : ! **************************************************************************************************
     196           91 :    SUBROUTINE parse_feature_list(features_json, features)
     197              :       CHARACTER(len=*), INTENT(IN)                       :: features_json
     198              :       CHARACTER(len=default_string_length), &
     199              :          ALLOCATABLE, DIMENSION(:), INTENT(OUT)          :: features
     200              : 
     201              :       INTEGER                                            :: end_pos, feature_count, i, pos, quote1, &
     202              :                                                             quote2, start_pos
     203              : 
     204           91 :       feature_count = 0
     205           91 :       pos = 1
     206          819 :       DO
     207          910 :          quote1 = INDEX(features_json(pos:), '"')
     208          910 :          IF (quote1 == 0) EXIT
     209          819 :          start_pos = pos + quote1
     210          819 :          quote2 = INDEX(features_json(start_pos:), '"')
     211          819 :          IF (quote2 == 0) EXIT
     212          819 :          feature_count = feature_count + 1
     213          819 :          pos = start_pos + quote2
     214              :       END DO
     215              : 
     216           91 :       IF (feature_count == 0) CPABORT("SKALA TorchScript model does not list any features")
     217          273 :       ALLOCATE (features(feature_count))
     218          910 :       features = ""
     219              : 
     220              :       pos = 1
     221          910 :       DO i = 1, feature_count
     222          819 :          quote1 = INDEX(features_json(pos:), '"')
     223          819 :          start_pos = pos + quote1
     224          819 :          quote2 = INDEX(features_json(start_pos:), '"')
     225          819 :          end_pos = start_pos + quote2 - 2
     226          819 :          features(i) = features_json(start_pos:end_pos)
     227          910 :          pos = start_pos + quote2
     228              :       END DO
     229              : 
     230           91 :    END SUBROUTINE parse_feature_list
     231              : 
     232            0 : END MODULE skala_torch_api
        

Generated by: LCOV version 2.0-1