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 Module for equivariant PAO-ML based on PyTorch.
10 : !> \author Ole Schuett
11 : ! **************************************************************************************************
12 : MODULE pao_model
13 : USE OMP_LIB, ONLY: omp_init_lock,&
14 : omp_set_lock,&
15 : omp_unset_lock
16 : USE atomic_kind_types, ONLY: atomic_kind_type,&
17 : get_atomic_kind
18 : USE basis_set_types, ONLY: gto_basis_set_type
19 : USE cell_types, ONLY: cell_type
20 : USE cp_dbcsr_api, ONLY: dbcsr_get_info,&
21 : dbcsr_iterator_blocks_left,&
22 : dbcsr_iterator_next_block,&
23 : dbcsr_iterator_start,&
24 : dbcsr_iterator_stop,&
25 : dbcsr_iterator_type,&
26 : dbcsr_type
27 : USE kinds, ONLY: default_path_length,&
28 : default_string_length,&
29 : dp,&
30 : int_8,&
31 : sp
32 : USE message_passing, ONLY: mp_para_env_type
33 : USE pao_types, ONLY: pao_env_type,&
34 : pao_model_type
35 : USE particle_types, ONLY: particle_type
36 : USE physcon, ONLY: angstrom
37 : USE qs_environment_types, ONLY: get_qs_env,&
38 : qs_environment_type
39 : USE qs_kind_types, ONLY: get_qs_kind,&
40 : qs_kind_type
41 : USE torch_api, ONLY: &
42 : torch_dict_create, torch_dict_get, torch_dict_insert, torch_dict_release, torch_dict_type, &
43 : torch_model_forward, torch_model_get_attr, torch_model_load, torch_tensor_backward, &
44 : torch_tensor_data_ptr, torch_tensor_from_array, torch_tensor_grad, torch_tensor_release, &
45 : torch_tensor_type
46 : #include "./base/base_uses.f90"
47 :
48 : IMPLICIT NONE
49 :
50 : PRIVATE
51 :
52 : PUBLIC :: pao_model_load, pao_model_predict, pao_model_forces, pao_model_type
53 :
54 : CONTAINS
55 :
56 : ! **************************************************************************************************
57 : !> \brief Loads a PAO-ML model.
58 : !> \param pao ...
59 : !> \param qs_env ...
60 : !> \param ikind ...
61 : !> \param pao_model_file ...
62 : !> \param model ...
63 : ! **************************************************************************************************
64 0 : SUBROUTINE pao_model_load(pao, qs_env, ikind, pao_model_file, model)
65 : TYPE(pao_env_type), INTENT(IN) :: pao
66 : TYPE(qs_environment_type), INTENT(IN) :: qs_env
67 : INTEGER, INTENT(IN) :: ikind
68 : CHARACTER(LEN=default_path_length), INTENT(IN) :: pao_model_file
69 : TYPE(pao_model_type), INTENT(OUT) :: model
70 :
71 : CHARACTER(len=*), PARAMETER :: routineN = 'pao_model_load'
72 :
73 : CHARACTER(LEN=default_string_length) :: kind_name
74 : CHARACTER(LEN=default_string_length), &
75 8 : ALLOCATABLE, DIMENSION(:) :: model_kind_names
76 : INTEGER :: handle, jkind, kkind, pao_basis_size, z
77 : REAL(dp) :: cutoff_angstrom
78 8 : TYPE(atomic_kind_type), DIMENSION(:), POINTER :: atomic_kind_set
79 : TYPE(gto_basis_set_type), POINTER :: basis_set
80 8 : TYPE(qs_kind_type), DIMENSION(:), POINTER :: qs_kind_set
81 :
82 8 : CALL timeset(routineN, handle)
83 8 : CALL get_qs_env(qs_env, qs_kind_set=qs_kind_set, atomic_kind_set=atomic_kind_set)
84 :
85 8 : IF (pao%iw > 0) WRITE (pao%iw, '(A)') " PAO| Loading PyTorch model from: "//TRIM(pao_model_file)
86 8 : CALL torch_model_load(model%torch_model, pao_model_file)
87 :
88 : ! Read model attributes.
89 8 : CALL torch_model_get_attr(model%torch_model, "pao_model_version", model%version)
90 8 : CALL torch_model_get_attr(model%torch_model, "kind_name", model%kind_name)
91 8 : CALL torch_model_get_attr(model%torch_model, "atomic_number", model%atomic_number)
92 8 : CALL torch_model_get_attr(model%torch_model, "prim_basis_name", model%prim_basis_name)
93 8 : CALL torch_model_get_attr(model%torch_model, "prim_basis_size", model%prim_basis_size)
94 8 : CALL torch_model_get_attr(model%torch_model, "pao_basis_size", model%pao_basis_size)
95 8 : CALL torch_model_get_attr(model%torch_model, "num_layers", model%num_layers)
96 8 : CALL torch_model_get_attr(model%torch_model, "cutoff", cutoff_angstrom)
97 8 : CALL torch_model_get_attr(model%torch_model, "all_kind_names", model_kind_names)
98 8 : model%cutoff = cutoff_angstrom/angstrom
99 :
100 : ! Freeze model after all attributes have been read.
101 : ! TODO Re-enable once the memory leaks of torch::jit::freeze() are fixed.
102 : ! https://github.com/pytorch/pytorch/issues/96726
103 : ! CALL torch_model_freeze(model%torch_model)
104 :
105 : ! For each of the model's kind names lookup the corresponding atomic kind index.
106 24 : ALLOCATE (model%kinds_mapping(SIZE(atomic_kind_set)))
107 24 : model%kinds_mapping(:) = -1
108 24 : DO jkind = 1, SIZE(atomic_kind_set)
109 24 : DO kkind = 1, SIZE(model_kind_names)
110 24 : IF (TRIM(atomic_kind_set(jkind)%name) == TRIM(model_kind_names(kkind))) THEN
111 16 : model%kinds_mapping(jkind) = kkind - 1
112 16 : EXIT
113 : END IF
114 : END DO
115 24 : IF (model%kinds_mapping(jkind) < 0) THEN
116 0 : CALL cp_abort(__LOCATION__, "PAO-ML model lacks kind '"//TRIM(atomic_kind_set(jkind)%name)//"' .")
117 : END IF
118 : END DO
119 :
120 : ! Check compatibility
121 8 : CALL get_qs_kind(qs_kind_set(ikind), basis_set=basis_set, pao_basis_size=pao_basis_size)
122 8 : CALL get_atomic_kind(atomic_kind_set(ikind), name=kind_name, z=z)
123 8 : IF (model%version /= 2) THEN
124 0 : CPABORT("Model version not supported.")
125 : END IF
126 8 : IF (TRIM(model%kind_name) /= TRIM(kind_name)) THEN
127 0 : CPABORT("Kind name does not match.")
128 : END IF
129 8 : IF (model%atomic_number /= z) THEN
130 0 : CPABORT("Atomic number does not match.")
131 : END IF
132 8 : IF (TRIM(model%prim_basis_name) /= TRIM(basis_set%name)) THEN
133 0 : CPABORT("Primary basis set name does not match.")
134 : END IF
135 8 : IF (model%prim_basis_size /= basis_set%nsgf) THEN
136 0 : CPABORT("Primary basis set size does not match.")
137 : END IF
138 8 : IF (model%pao_basis_size /= pao_basis_size) THEN
139 0 : CPABORT("PAO basis size does not match.")
140 : END IF
141 :
142 8 : CALL omp_init_lock(model%lock)
143 8 : CALL timestop(handle)
144 :
145 32 : END SUBROUTINE pao_model_load
146 :
147 : ! **************************************************************************************************
148 : !> \brief Fills pao%matrix_X based on machine learning predictions
149 : !> \param pao ...
150 : !> \param qs_env ...
151 : ! **************************************************************************************************
152 16 : SUBROUTINE pao_model_predict(pao, qs_env)
153 : TYPE(pao_env_type), POINTER :: pao
154 : TYPE(qs_environment_type), POINTER :: qs_env
155 :
156 : CHARACTER(len=*), PARAMETER :: routineN = 'pao_model_predict'
157 :
158 : INTEGER :: acol, arow, handle, iatom
159 16 : REAL(dp), DIMENSION(:, :), POINTER :: block_X
160 : TYPE(dbcsr_iterator_type) :: iter
161 :
162 16 : CALL timeset(routineN, handle)
163 :
164 16 : !$OMP PARALLEL DEFAULT(NONE) SHARED(pao,qs_env) PRIVATE(iter,arow,acol,iatom,block_X)
165 : CALL dbcsr_iterator_start(iter, pao%matrix_X)
166 : DO WHILE (dbcsr_iterator_blocks_left(iter))
167 : CALL dbcsr_iterator_next_block(iter, arow, acol, block_X)
168 : IF (SIZE(block_X) == 0) CYCLE ! pao disabled for iatom
169 : iatom = arow; CPASSERT(arow == acol)
170 : CALL predict_single_atom(pao, qs_env, iatom, block_X=block_X)
171 : END DO
172 : CALL dbcsr_iterator_stop(iter)
173 : !$OMP END PARALLEL
174 :
175 16 : CALL timestop(handle)
176 :
177 16 : END SUBROUTINE pao_model_predict
178 :
179 : ! **************************************************************************************************
180 : !> \brief Calculate forces contributed by machine learning
181 : !> \param pao ...
182 : !> \param qs_env ...
183 : !> \param matrix_G ...
184 : !> \param forces ...
185 : ! **************************************************************************************************
186 2 : SUBROUTINE pao_model_forces(pao, qs_env, matrix_G, forces)
187 : TYPE(pao_env_type), POINTER :: pao
188 : TYPE(qs_environment_type), POINTER :: qs_env
189 : TYPE(dbcsr_type) :: matrix_G
190 : REAL(dp), DIMENSION(:, :), INTENT(INOUT) :: forces
191 :
192 : CHARACTER(len=*), PARAMETER :: routineN = 'pao_model_forces'
193 :
194 : INTEGER :: acol, arow, handle, iatom
195 2 : REAL(dp), DIMENSION(:, :), POINTER :: block_G
196 : TYPE(dbcsr_iterator_type) :: iter
197 :
198 2 : CALL timeset(routineN, handle)
199 :
200 2 : !$OMP PARALLEL DEFAULT(NONE) SHARED(pao,qs_env,matrix_G,forces) PRIVATE(iter,arow,acol,iatom,block_G)
201 : CALL dbcsr_iterator_start(iter, matrix_G)
202 : DO WHILE (dbcsr_iterator_blocks_left(iter))
203 : CALL dbcsr_iterator_next_block(iter, arow, acol, block_G)
204 : iatom = arow; CPASSERT(arow == acol)
205 : IF (SIZE(block_G) == 0) CYCLE ! pao disabled for iatom
206 : CALL predict_single_atom(pao, qs_env, iatom, block_G=block_G, forces=forces)
207 : END DO
208 : CALL dbcsr_iterator_stop(iter)
209 : !$OMP END PARALLEL
210 :
211 2 : CALL timestop(handle)
212 :
213 2 : END SUBROUTINE pao_model_forces
214 :
215 : ! **************************************************************************************************
216 : !> \brief Predicts a single block_X.
217 : !> \param pao ...
218 : !> \param qs_env ...
219 : !> \param iatom ...
220 : !> \param block_X ...
221 : !> \param block_G ...
222 : !> \param forces ...
223 : ! **************************************************************************************************
224 54 : SUBROUTINE predict_single_atom(pao, qs_env, iatom, block_X, block_G, forces)
225 : TYPE(pao_env_type), INTENT(IN), POINTER :: pao
226 : TYPE(qs_environment_type), INTENT(IN), POINTER :: qs_env
227 : INTEGER, INTENT(IN) :: iatom
228 : REAL(dp), DIMENSION(:, :), OPTIONAL :: block_X, block_G, forces
229 :
230 : INTEGER :: i, iedge, ikind, j, jatom, jcell, jkind, &
231 : jneighbor, k, katom, kneighbor, m, n, &
232 : natoms, num_edges, num_neighbors
233 : INTEGER(kind=int_8), ALLOCATABLE, DIMENSION(:) :: neighbor_atom_types
234 54 : INTEGER(kind=int_8), ALLOCATABLE, DIMENSION(:, :) :: central_edge_index, edge_index
235 54 : INTEGER, ALLOCATABLE, DIMENSION(:) :: neighbor_atom_index
236 54 : INTEGER, DIMENSION(:), POINTER :: blk_sizes_pao, blk_sizes_pri
237 : REAL(dp), DIMENSION(3) :: Ri, Rj, Rjk
238 54 : REAL(KIND=dp), ALLOCATABLE, DIMENSION(:, :) :: cell_shifts, neighbor_pos
239 54 : REAL(sp), ALLOCATABLE, DIMENSION(:, :) :: edge_vectors
240 54 : REAL(sp), ALLOCATABLE, DIMENSION(:, :, :) :: outer_grad
241 54 : REAL(sp), DIMENSION(:, :), POINTER :: edge_vectors_grad
242 54 : REAL(sp), DIMENSION(:, :, :), POINTER :: predicted_xblock
243 54 : TYPE(atomic_kind_type), DIMENSION(:), POINTER :: atomic_kind_set
244 : TYPE(cell_type), POINTER :: cell
245 : TYPE(mp_para_env_type), POINTER :: para_env
246 : TYPE(pao_model_type), POINTER :: model
247 54 : TYPE(particle_type), DIMENSION(:), POINTER :: particle_set
248 54 : TYPE(qs_kind_type), DIMENSION(:), POINTER :: qs_kind_set
249 : TYPE(torch_dict_type) :: model_inputs, model_outputs
250 : TYPE(torch_tensor_type) :: atom_types_tensor, central_edge_index_tensor, edge_index_tensor, &
251 : edge_vectors_grad_tensor, edge_vectors_tensor, outer_grad_tensor, predicted_xblock_tensor
252 :
253 54 : CALL dbcsr_get_info(pao%matrix_Y, row_blk_size=blk_sizes_pri, col_blk_size=blk_sizes_pao)
254 54 : n = blk_sizes_pri(iatom) ! size of primary basis
255 54 : m = blk_sizes_pao(iatom) ! size of pao basis
256 :
257 : CALL get_qs_env(qs_env, &
258 : para_env=para_env, &
259 : cell=cell, &
260 : particle_set=particle_set, &
261 : atomic_kind_set=atomic_kind_set, &
262 : qs_kind_set=qs_kind_set, &
263 54 : natom=natoms)
264 :
265 54 : CALL get_atomic_kind(particle_set(iatom)%atomic_kind, kind_number=ikind)
266 216 : Ri = particle_set(iatom)%r
267 54 : model => pao%models(ikind)
268 54 : CPASSERT(model%version > 0)
269 54 : CALL omp_set_lock(model%lock) ! TODO: might not be needed for inference.
270 :
271 : ! TODO: this is a quadratic algorithm, use a neighbor-list instead.
272 :
273 : ! Enumerate all neighboring images. TODO: should be all images within num_layers*cutoff.
274 54 : ALLOCATE (cell_shifts(27, 3))
275 216 : jcell = 0
276 216 : DO i = -1, +1
277 702 : DO j = -1, +1
278 2106 : DO k = -1, +1
279 1458 : jcell = jcell + 1
280 6318 : cell_shifts(jcell, :) = i*cell%hmat(:, 1) + j*cell%hmat(:, 2) + k*cell%hmat(:, 3)
281 : END DO
282 : END DO
283 : END DO
284 :
285 : ! Find neighbors, ie. atoms that are reachable within num_layers*cutoff.
286 : ! 1st pass to count neighbors.
287 54 : num_neighbors = 1 ! first neighbor is always the central atom
288 378 : DO jatom = 1, natoms
289 9126 : DO jcell = 1, 27
290 34992 : Rj = particle_set(jatom)%r + cell_shifts(jcell, :)
291 36018 : IF (NORM2(Rj - Ri) < model%num_layers*model%cutoff .AND. ANY(Rj /= Ri)) THEN
292 180 : num_neighbors = num_neighbors + 1
293 : END IF
294 : END DO
295 : END DO
296 :
297 : ! 2nd pass to collect neighbors.
298 378 : ALLOCATE (neighbor_pos(num_neighbors, 3), neighbor_atom_types(num_neighbors), neighbor_atom_index(num_neighbors))
299 54 : num_neighbors = 1 ! first neighbor is always the central atom
300 216 : neighbor_pos(1, :) = Ri
301 54 : neighbor_atom_types(1) = model%kinds_mapping(ikind)
302 54 : neighbor_atom_index(1) = iatom
303 378 : DO jatom = 1, natoms
304 9126 : DO jcell = 1, 27
305 34992 : Rj = particle_set(jatom)%r + cell_shifts(jcell, :)
306 8748 : jkind = particle_set(jatom)%atomic_kind%kind_number
307 36018 : IF (NORM2(Rj - Ri) < model%num_layers*model%cutoff .AND. ANY(Rj /= Ri)) THEN
308 180 : num_neighbors = num_neighbors + 1
309 720 : neighbor_pos(num_neighbors, :) = Rj
310 180 : neighbor_atom_types(num_neighbors) = model%kinds_mapping(jkind)
311 180 : neighbor_atom_index(num_neighbors) = jatom
312 : END IF
313 : END DO
314 : END DO
315 :
316 : ! Build connectivity graph of neighbors.
317 : ! 1st pass to count edges.
318 : num_edges = 0
319 288 : DO jneighbor = 1, num_neighbors
320 1350 : DO kneighbor = 1, num_neighbors
321 4248 : Rjk = neighbor_pos(kneighbor, :) - neighbor_pos(jneighbor, :)
322 4482 : IF (NORM2(Rjk) < model%cutoff .AND. jneighbor /= kneighbor) THEN
323 684 : num_edges = num_edges + 1
324 : END IF
325 : END DO
326 : END DO
327 :
328 : ! 2nd pass to collect edges.
329 270 : ALLOCATE (edge_index(num_edges, 2), edge_vectors(3, num_edges)) ! edge_index is transposed
330 54 : num_edges = 0
331 288 : DO jneighbor = 1, num_neighbors
332 1350 : DO kneighbor = 1, num_neighbors
333 4248 : Rjk = neighbor_pos(kneighbor, :) - neighbor_pos(jneighbor, :)
334 4482 : IF (NORM2(Rjk) < model%cutoff .AND. jneighbor /= kneighbor) THEN
335 684 : num_edges = num_edges + 1
336 2052 : edge_index(num_edges, :) = [jneighbor - 1, kneighbor - 1]
337 2736 : edge_vectors(:, num_edges) = REAL(Rjk*angstrom, kind=sp)
338 : END IF
339 : END DO
340 : END DO
341 :
342 54 : ALLOCATE (central_edge_index(1, 2))
343 54 : central_edge_index(:, :) = 0
344 :
345 : ! Inference.
346 54 : CALL torch_dict_create(model_inputs)
347 :
348 54 : CALL torch_tensor_from_array(atom_types_tensor, neighbor_atom_types)
349 54 : CALL torch_dict_insert(model_inputs, "atom_types", atom_types_tensor)
350 :
351 54 : CALL torch_tensor_from_array(edge_index_tensor, edge_index)
352 54 : CALL torch_dict_insert(model_inputs, "edge_index", edge_index_tensor)
353 :
354 54 : CALL torch_tensor_from_array(edge_vectors_tensor, edge_vectors, requires_grad=PRESENT(block_G))
355 54 : CALL torch_dict_insert(model_inputs, "edge_vectors", edge_vectors_tensor)
356 :
357 54 : CALL torch_tensor_from_array(central_edge_index_tensor, central_edge_index)
358 54 : CALL torch_dict_insert(model_inputs, "central_edge_index", central_edge_index_tensor)
359 :
360 54 : CALL torch_dict_create(model_outputs)
361 54 : CALL torch_model_forward(model%torch_model, model_inputs, model_outputs)
362 :
363 : ! Copy predicted XBlock.
364 54 : NULLIFY (predicted_xblock)
365 54 : CALL torch_dict_get(model_outputs, "xblock", predicted_xblock_tensor)
366 54 : CALL torch_tensor_data_ptr(predicted_xblock_tensor, predicted_xblock)
367 54 : CPASSERT(SIZE(predicted_xblock, 1) == n)
368 54 : CPASSERT(SIZE(predicted_xblock, 2) == m)
369 54 : CPASSERT(SIZE(predicted_xblock, 3) == 1)
370 1980 : CPASSERT(ALL(predicted_xblock == predicted_xblock)) ! checking for NaNs
371 54 : IF (PRESENT(block_X)) THEN
372 1664 : block_X = RESHAPE(predicted_xblock, [n*m, 1])
373 : END IF
374 :
375 : ! TURNING POINT (if calc forces) ------------------------------------------
376 54 : IF (PRESENT(block_G)) THEN
377 24 : ALLOCATE (outer_grad(n, m, 1))
378 238 : outer_grad(:, :, :) = REAL(RESHAPE(block_G, [n, m, 1]), kind=sp)
379 6 : CALL torch_tensor_from_array(outer_grad_tensor, outer_grad)
380 6 : CALL torch_tensor_backward(predicted_xblock_tensor, outer_grad_tensor)
381 6 : CALL torch_tensor_grad(edge_vectors_tensor, edge_vectors_grad_tensor)
382 6 : NULLIFY (edge_vectors_grad)
383 6 : CALL torch_tensor_data_ptr(edge_vectors_grad_tensor, edge_vectors_grad)
384 6 : IF (ASSOCIATED(edge_vectors_grad)) THEN ! Torch may return NULL pointer as gradient.
385 6 : CPASSERT(SIZE(edge_vectors_grad, 1) == 3 .AND. SIZE(edge_vectors_grad, 2) == num_edges)
386 310 : CPASSERT(ALL(edge_vectors_grad == edge_vectors_grad)) ! checking for NaNs
387 82 : DO iedge = 1, num_edges
388 76 : jneighbor = INT(edge_index(iedge, 1) + 1)
389 76 : kneighbor = INT(edge_index(iedge, 2) + 1)
390 76 : jatom = neighbor_atom_index(jneighbor)
391 76 : katom = neighbor_atom_index(kneighbor)
392 304 : forces(jatom, :) = forces(jatom, :) + edge_vectors_grad(:, iedge)*angstrom
393 310 : forces(katom, :) = forces(katom, :) - edge_vectors_grad(:, iedge)*angstrom
394 : END DO
395 : END IF
396 6 : CALL torch_tensor_release(outer_grad_tensor)
397 6 : CALL torch_tensor_release(edge_vectors_grad_tensor)
398 : END IF
399 :
400 : ! Clean up.
401 54 : CALL torch_tensor_release(atom_types_tensor)
402 54 : CALL torch_tensor_release(edge_index_tensor)
403 54 : CALL torch_tensor_release(edge_vectors_tensor)
404 54 : CALL torch_tensor_release(central_edge_index_tensor)
405 54 : CALL torch_tensor_release(predicted_xblock_tensor)
406 54 : CALL torch_dict_release(model_inputs)
407 54 : CALL torch_dict_release(model_outputs)
408 54 : CALL omp_unset_lock(model%lock)
409 :
410 162 : END SUBROUTINE predict_single_atom
411 :
412 : END MODULE pao_model
|