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
|