12#if defined (__HAS_IEEE_EXCEPTIONS)
13 USE ieee_exceptions,
ONLY: ieee_all, &
14 ieee_get_halting_mode, &
26#include "./base/base_uses.f90"
32 CHARACTER(len=*),
PARAMETER,
PRIVATE :: moduleN =
'skala_torch_api'
40 INTEGER :: protocol_version = -1
41 CHARACTER(len=default_string_length),
ALLOCATABLE, &
42 DIMENSION(:) :: features
55 CHARACTER(len=*),
INTENT(IN) :: filename
57 CHARACTER(:),
ALLOCATABLE :: features_json, protocol_string
61 "protocol_version", protocol_string, &
62 "features", features_json)
65 READ (protocol_string, *, iostat=ios) model%protocol_version
66 IF (ios /= 0) cpabort(
"Could not parse SKALA TorchScript protocol_version metadata")
67 IF (model%protocol_version /= 2)
THEN
68 cpabort(
"Unsupported SKALA TorchScript protocol version")
71 CALL parse_feature_list(features_json, model%features)
85 IF (
ALLOCATED(model%features))
DEALLOCATE (model%features)
86 model%protocol_version = -1
98 CHARACTER(len=*),
INTENT(IN) :: feature
99 LOGICAL :: needs_feature
101 CHARACTER(len=default_string_length) :: feature_key, model_feature
104 feature_key = adjustl(feature)
107 needs_feature = .false.
108 IF (.NOT.
ALLOCATED(model%features))
RETURN
110 DO i = 1,
SIZE(model%features)
111 model_feature = adjustl(model%features(i))
113 IF (trim(model_feature) == trim(feature_key))
THEN
114 needs_feature = .true.
128 INTEGER :: protocol_version
130 protocol_version = model%protocol_version
145#if defined (__HAS_IEEE_EXCEPTIONS)
146 LOGICAL,
DIMENSION(5) :: ieee_halt
148 CALL ieee_get_halting_mode(ieee_all, ieee_halt)
149 CALL ieee_set_halting_mode(ieee_all, .false.)
152#if defined (__HAS_IEEE_EXCEPTIONS)
153 CALL ieee_set_halting_mode(ieee_all, ieee_halt)
171 REAL(kind=
dp),
INTENT(OUT) :: exc
175#if defined (__HAS_IEEE_EXCEPTIONS)
176 LOGICAL,
DIMENSION(5) :: ieee_halt
178 CALL ieee_get_halting_mode(ieee_all, ieee_halt)
179 CALL ieee_set_halting_mode(ieee_all, .false.)
185#if defined (__HAS_IEEE_EXCEPTIONS)
186 CALL ieee_set_halting_mode(ieee_all, ieee_halt)
196 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
201 INTEGER :: end_pos, feature_count, i, pos, quote1, &
207 quote1 = index(features_json(pos:),
'"')
208 IF (quote1 == 0)
EXIT
209 start_pos = pos + quote1
210 quote2 = index(features_json(start_pos:),
'"')
211 IF (quote2 == 0)
EXIT
212 feature_count = feature_count + 1
213 pos = start_pos + quote2
216 IF (feature_count == 0) cpabort(
"SKALA TorchScript model does not list any features")
217 ALLOCATE (features(feature_count))
221 DO i = 1, feature_count
222 quote1 = index(features_json(pos:),
'"')
223 start_pos = pos + quote1
224 quote2 = index(features_json(start_pos:),
'"')
225 end_pos = start_pos + quote2 - 2
226 features(i) = features_json(start_pos:end_pos)
227 pos = start_pos + quote2
230 END SUBROUTINE parse_feature_list
Defines the basic variable types.
integer, parameter, public dp
integer, parameter, public default_string_length
Small CP2K wrapper around the SKALA TorchScript functional protocol.
subroutine, public skala_torch_model_release(model)
Release a loaded SKALA TorchScript model.
logical function, public skala_torch_model_needs_feature(model, feature)
Check whether a loaded SKALA model requests a feature.
subroutine, public skala_torch_model_get_exc(model, inputs, grid_weights, exc_tensor, exc)
Evaluate the weighted SKALA exchange-correlation energy.
integer function, public skala_torch_model_protocol_version(model)
Return the loaded SKALA TorchScript protocol version.
subroutine, public skala_torch_model_get_exc_density(model, inputs, exc_density)
Evaluate the SKALA exchange-correlation energy density.
subroutine, public skala_torch_model_load(model, filename)
Load a SKALA TorchScript model and its feature metadata.
Utilities for string manipulations.
elemental subroutine, public uppercase(string)
Convert all lower case characters in a string to upper case.
subroutine, public torch_model_freeze_preserving_method(model, method_name)
Freeze a Torch model while preserving one exported method.
real(kind=dp) function, public torch_tensor_item_double(tensor)
Returns a scalar double value from a Torch tensor.
subroutine, public torch_model_load_with_metadata(model, filename, key1, value1, key2, value2)
Loads a Torch model and reads two metadata entries in the same archive pass.
subroutine, public torch_model_remap_device_constants(model)
Maps serialized TorchScript device constants to the active Torch device.
subroutine, public torch_model_forward_mol_tensor(model, method_name, inputs, output)
Evaluates a TorchScript model method expecting keyword argument "mol".
subroutine, public torch_model_disable_parameter_gradients(model)
Disable gradients for inference-only model parameters.
subroutine, public torch_model_release(model)
Releases a Torch model and all its ressources.
subroutine, public torch_tensor_weighted_sum(values, weights, result)
Returns the weighted sum of two Torch tensors.
subroutine, public torch_tensor_release(tensor)
Releases a Torch tensor and all its ressources.