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 Shared TorchScript evaluation path for e3nn-based equivariant message-passing
10 : !> potentials (NequIP, Allegro and MACE).
11 : !> \par History
12 : !> Implementation of NequIP and Allegro potentials - [gtocci] 2022
13 : !> Index mapping of atoms from .xyz to Allegro config.yaml file - [mbilichenko] 2024
14 : !> Refactoring and update to NequIP version >= v0.7.0 - [gtocci] 2026
15 : !> Renamed manybody_nequip -> manybody_e3nn as it now also serves MACE - [xysun] 2026
16 : !> \author Gabriele Tocci
17 : ! **************************************************************************************************
18 : MODULE manybody_e3nn
19 :
20 : USE atomic_kind_types, ONLY: atomic_kind_type
21 : USE cell_types, ONLY: cell_type
22 : USE distribution_1d_types, ONLY: distribution_1d_type
23 : USE fist_neighbor_list_types, ONLY: fist_neighbor_type,&
24 : neighbor_kind_pairs_type
25 : USE fist_nonbond_env_types, ONLY: fist_nonbond_env_get,&
26 : fist_nonbond_env_set,&
27 : fist_nonbond_env_type,&
28 : nequip_data_type,&
29 : pos_type
30 : USE kinds, ONLY: default_string_length,&
31 : dp,&
32 : int_8
33 : USE message_passing, ONLY: mp_para_env_type
34 : USE pair_potential_types, ONLY: mace_type,&
35 : nequip_pot_type,&
36 : nequip_type,&
37 : pair_potential_pp_type,&
38 : pair_potential_single_type
39 : USE particle_types, ONLY: particle_type
40 : USE string_utilities, ONLY: uppercase
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_freeze, torch_model_load, torch_tensor_data_ptr, &
44 : torch_tensor_from_array, torch_tensor_release, torch_tensor_type
45 : #include "./base/base_uses.f90"
46 :
47 : IMPLICIT NONE
48 :
49 : PRIVATE
50 : PUBLIC :: e3nn_energy_store_force_virial, &
51 : e3nn_add_force_virial
52 :
53 : CHARACTER(len=*), PARAMETER, PRIVATE :: moduleN = 'manybody_e3nn'
54 :
55 : TYPE, PRIVATE :: nequip_work_type
56 : INTEGER :: target_pot_type
57 : INTEGER :: n_atoms_use
58 : LOGICAL :: use_virial
59 :
60 : TYPE(cell_type), POINTER :: cell => NULL()
61 : TYPE(pos_type), DIMENSION(:), POINTER :: r_pbc => NULL()
62 : TYPE(distribution_1d_type), POINTER :: local_particles => NULL()
63 : TYPE(particle_type), POINTER :: particle_set(:) => NULL()
64 : TYPE(mp_para_env_type), POINTER :: para_env => NULL()
65 :
66 : LOGICAL, ALLOCATABLE :: use_atom(:)
67 : INTEGER(kind=int_8), ALLOCATABLE :: local_edges(:, :)
68 : REAL(kind=dp), ALLOCATABLE :: local_shifts(:, :)
69 : INTEGER(kind=int_8), ALLOCATABLE :: final_edges(:, :)
70 : REAL(kind=dp), ALLOCATABLE :: final_shifts(:, :)
71 : INTEGER, DIMENSION(:), ALLOCATABLE :: kind_mapper
72 : LOGICAL, ALLOCATABLE :: sum_energy(:)
73 : END TYPE nequip_work_type
74 :
75 : CONTAINS
76 :
77 : ! **************************************************************************************************
78 : !> \brief ...
79 : !> \param nonbonded ...
80 : !> \param particle_set ...
81 : !> \param local_particles ...
82 : !> \param cell ...
83 : !> \param atomic_kind_set ...
84 : !> \param potparm ...
85 : !> \param r_last_update_pbc ...
86 : !> \param pot_total ...
87 : !> \param fist_nonbond_env ...
88 : !> \param para_env ...
89 : !> \param use_virial ...
90 : !> \param target_pot_type ...
91 : !> \par History
92 : !> Implementation of the nequip potential - [gtocci] 2022
93 : !> Refactoring and unifying NequIP and Allegro - [gtocci] 2026
94 : !> \author Gabriele Tocci - University of Zurich
95 : ! **************************************************************************************************
96 6 : SUBROUTINE e3nn_energy_store_force_virial(nonbonded, particle_set, local_particles, cell, &
97 : atomic_kind_set, potparm, r_last_update_pbc, &
98 : pot_total, fist_nonbond_env, para_env, use_virial, &
99 : target_pot_type)
100 :
101 : TYPE(fist_neighbor_type), POINTER :: nonbonded
102 : TYPE(particle_type), POINTER :: particle_set(:)
103 : TYPE(distribution_1d_type), POINTER :: local_particles
104 : TYPE(cell_type), POINTER :: cell
105 : TYPE(atomic_kind_type), POINTER :: atomic_kind_set(:)
106 : TYPE(pair_potential_pp_type), POINTER :: potparm
107 : TYPE(pos_type), DIMENSION(:), POINTER :: r_last_update_pbc
108 : REAL(kind=dp) :: pot_total
109 : TYPE(fist_nonbond_env_type), POINTER :: fist_nonbond_env
110 : TYPE(mp_para_env_type), POINTER :: para_env
111 : LOGICAL, INTENT(IN) :: use_virial
112 : INTEGER, INTENT(IN) :: target_pot_type
113 :
114 : CHARACTER(LEN=*), PARAMETER :: routineN = 'e3nn_energy_store_force_virial'
115 :
116 : INTEGER :: handle
117 : TYPE(nequip_data_type), POINTER :: neq_data
118 : TYPE(nequip_pot_type), POINTER :: neq_pot
119 6 : TYPE(nequip_work_type) :: nequip_work
120 : TYPE(torch_dict_type) :: outputs
121 :
122 6 : CALL timeset(routineN, handle)
123 :
124 : CALL nequip_work_create(nequip_work, atomic_kind_set, particle_set, local_particles, cell, &
125 : r_last_update_pbc, para_env, potparm, target_pot_type, use_virial, &
126 6 : neq_pot)
127 :
128 6 : IF (.NOT. ASSOCIATED(neq_pot)) THEN
129 0 : CALL timestop(handle)
130 0 : RETURN
131 : END IF
132 :
133 6 : CALL build_local_edges_shifts(nonbonded, potparm, nequip_work)
134 :
135 6 : CALL build_torch_edge_indexes(nequip_work)
136 :
137 6 : CALL setup_neq_data(fist_nonbond_env, neq_data, neq_pot, nequip_work)
138 :
139 6 : IF (nequip_work%target_pot_type == nequip_type .OR. &
140 : nequip_work%target_pot_type == mace_type) THEN
141 4 : CALL prepare_edges_shifts_nequip(nequip_work)
142 : ELSE
143 2 : CALL prepare_edges_shifts_allegro(nequip_work)
144 : END IF
145 :
146 6 : CALL run_torch_model(neq_data, neq_pot, nequip_work, outputs)
147 :
148 6 : CALL process_outputs(outputs, neq_data, neq_pot, pot_total, nequip_work)
149 :
150 6 : CALL torch_dict_release(outputs)
151 6 : CALL release_nequip_work(nequip_work)
152 :
153 6 : CALL timestop(handle)
154 6 : END SUBROUTINE e3nn_energy_store_force_virial
155 :
156 : ! **************************************************************************************************
157 : !> \brief ...
158 : !> \param nequip_work ...
159 : !> \param atomic_kind_set ...
160 : !> \param particle_set ...
161 : !> \param local_particles ...
162 : !> \param cell ...
163 : !> \param r_pbc ...
164 : !> \param para_env ...
165 : !> \param potparm ...
166 : !> \param target_pot_type ...
167 : !> \param use_virial ...
168 : !> \param neq_pot ...
169 : !> \author Gabriele Tocci - University of Zurich
170 : ! **************************************************************************************************
171 6 : SUBROUTINE nequip_work_create(nequip_work, atomic_kind_set, particle_set, local_particles, cell, &
172 : r_pbc, para_env, potparm, target_pot_type, use_virial, neq_pot)
173 : TYPE(nequip_work_type), INTENT(OUT) :: nequip_work
174 : TYPE(atomic_kind_type), POINTER :: atomic_kind_set(:)
175 : TYPE(particle_type), POINTER :: particle_set(:)
176 : TYPE(distribution_1d_type), POINTER :: local_particles
177 : TYPE(cell_type), POINTER :: cell
178 : TYPE(pos_type), DIMENSION(:), POINTER :: r_pbc
179 : TYPE(mp_para_env_type), POINTER :: para_env
180 : TYPE(pair_potential_pp_type), POINTER :: potparm
181 : INTEGER, INTENT(IN) :: target_pot_type
182 : LOGICAL, INTENT(IN) :: use_virial
183 : TYPE(nequip_pot_type), INTENT(OUT), POINTER :: neq_pot
184 :
185 6 : nequip_work%target_pot_type = target_pot_type
186 6 : nequip_work%use_virial = use_virial
187 6 : nequip_work%cell => cell
188 6 : nequip_work%r_pbc => r_pbc
189 6 : nequip_work%particle_set => particle_set
190 6 : nequip_work%para_env => para_env
191 6 : nequip_work%local_particles => local_particles
192 :
193 6 : CALL get_potential_config(atomic_kind_set, potparm, target_pot_type, neq_pot)
194 :
195 6 : IF (.NOT. ASSOCIATED(neq_pot)) THEN
196 : RETURN
197 : END IF
198 :
199 6 : CALL build_kind_mapper(atomic_kind_set, neq_pot, nequip_work)
200 :
201 6 : CALL init_atom_masks(nequip_work)
202 :
203 : END SUBROUTINE nequip_work_create
204 :
205 : ! **************************************************************************************************
206 : !> \brief ...
207 : !> \param nequip_work ...
208 : !> \author Gabriele Tocci - University of Zurich
209 : ! **************************************************************************************************
210 6 : SUBROUTINE release_nequip_work(nequip_work)
211 : TYPE(nequip_work_type), INTENT(INOUT) :: nequip_work
212 :
213 6 : IF (ALLOCATED(nequip_work%final_edges)) DEALLOCATE (nequip_work%final_edges)
214 6 : IF (ALLOCATED(nequip_work%final_shifts)) DEALLOCATE (nequip_work%final_shifts)
215 6 : IF (ALLOCATED(nequip_work%local_edges)) DEALLOCATE (nequip_work%local_edges)
216 6 : IF (ALLOCATED(nequip_work%local_shifts)) DEALLOCATE (nequip_work%local_shifts)
217 6 : IF (ALLOCATED(nequip_work%use_atom)) DEALLOCATE (nequip_work%use_atom)
218 6 : IF (ALLOCATED(nequip_work%kind_mapper)) DEALLOCATE (nequip_work%kind_mapper)
219 6 : IF (ALLOCATED(nequip_work%sum_energy)) DEALLOCATE (nequip_work%sum_energy)
220 6 : NULLIFY (nequip_work%cell, nequip_work%r_pbc, nequip_work%particle_set, nequip_work%para_env, &
221 6 : nequip_work%local_particles)
222 :
223 6 : END SUBROUTINE release_nequip_work
224 :
225 : ! **************************************************************************************************
226 : !> \brief ...
227 : !> \param nonbonded ...
228 : !> \param potparm ...
229 : !> \param nequip_work ...
230 : !> \par History
231 : !> Build edges and cell shifts for the GNN - [gtocci] 2026
232 : !> \author Gabriele Tocci - University of Zurich
233 : ! **************************************************************************************************
234 6 : SUBROUTINE build_local_edges_shifts(nonbonded, potparm, nequip_work)
235 : TYPE(fist_neighbor_type), POINTER :: nonbonded
236 : TYPE(pair_potential_pp_type), POINTER :: potparm
237 : TYPE(nequip_work_type), INTENT(INOUT) :: nequip_work
238 :
239 : INTEGER :: atom_a, atom_b, i, idx_i, idx_j, iend, &
240 : igrp, ikind, ilist, ipair, istart, &
241 : jkind, n_max_edges, nedges, npairs
242 6 : INTEGER, DIMENSION(:, :), POINTER :: list
243 : LOGICAL :: do_nequip_allegro
244 : REAL(kind=dp) :: cutsq_ij, drij, rij(3)
245 : REAL(kind=dp), DIMENSION(3) :: cell_v, cvi
246 : TYPE(neighbor_kind_pairs_type), POINTER :: neighbor_kind_pair
247 : TYPE(pair_potential_single_type), POINTER :: pot
248 :
249 6 : n_max_edges = 0
250 168 : DO ilist = 1, nonbonded%nlists
251 162 : neighbor_kind_pair => nonbonded%neighbor_kind_pairs(ilist)
252 168 : n_max_edges = n_max_edges + neighbor_kind_pair%npairs
253 : END DO
254 :
255 30 : ALLOCATE (nequip_work%local_edges(2, n_max_edges), nequip_work%local_shifts(3, n_max_edges))
256 6 : nedges = 0
257 :
258 168 : DO ilist = 1, nonbonded%nlists
259 162 : neighbor_kind_pair => nonbonded%neighbor_kind_pairs(ilist)
260 162 : npairs = neighbor_kind_pair%npairs
261 162 : IF (npairs == 0) CYCLE
262 :
263 470 : Kind_Loop: DO igrp = 1, neighbor_kind_pair%ngrp_kind
264 316 : istart = neighbor_kind_pair%grp_kind_start(igrp)
265 316 : iend = neighbor_kind_pair%grp_kind_end(igrp)
266 316 : ikind = neighbor_kind_pair%ij_kind(1, igrp)
267 316 : jkind = neighbor_kind_pair%ij_kind(2, igrp)
268 :
269 316 : idx_i = nequip_work%kind_mapper(ikind)
270 316 : idx_j = nequip_work%kind_mapper(jkind)
271 :
272 316 : IF (idx_i < 1 .OR. idx_j < 1) THEN
273 : ! pair involving atom not defined in the NequIP model, skipping..
274 : CYCLE Kind_Loop
275 : END IF
276 316 : pot => potparm%pot(ikind, jkind)%pot
277 316 : do_nequip_allegro = .FALSE.
278 316 : DO i = 1, SIZE(pot%type)
279 316 : IF (pot%type(i) == nequip_work%target_pot_type) THEN
280 : do_nequip_allegro = .TRUE.
281 : EXIT
282 : END IF
283 : END DO
284 :
285 316 : IF (.NOT. do_nequip_allegro) CYCLE Kind_Loop
286 :
287 316 : cutsq_ij = pot%set(i)%nequip%cutoff_matrix(idx_i, idx_j)
288 316 : list => neighbor_kind_pair%list
289 1264 : cvi = neighbor_kind_pair%cell_vector
290 316 : pot => potparm%pot(ikind, jkind)%pot
291 4108 : cell_v = MATMUL(nequip_work%cell%hmat, cvi)
292 :
293 20058 : DO ipair = istart, iend
294 19580 : atom_a = neighbor_kind_pair%list(1, ipair)
295 19580 : atom_b = neighbor_kind_pair%list(2, ipair)
296 :
297 78320 : rij(:) = nequip_work%r_pbc(atom_b)%r(:) - nequip_work%r_pbc(atom_a)%r(:) + cell_v
298 78320 : drij = DOT_PRODUCT(rij, rij)
299 :
300 19896 : IF (drij <= cutsq_ij) THEN
301 11406 : nedges = nedges + 1
302 34218 : nequip_work%local_edges(:, nedges) = [atom_a, atom_b]
303 45624 : nequip_work%local_shifts(:, nedges) = cvi
304 : END IF
305 : END DO
306 : END DO Kind_Loop
307 : END DO
308 :
309 6 : IF (nedges < n_max_edges) THEN
310 : BLOCK
311 6 : INTEGER(kind=int_8), ALLOCATABLE :: tmp_idx(:, :)
312 6 : REAL(kind=dp), ALLOCATABLE :: tmp_sft(:, :)
313 :
314 30 : ALLOCATE (tmp_idx(2, nedges), tmp_sft(3, nedges))
315 :
316 34224 : tmp_idx(:, :) = nequip_work%local_edges(:, 1:nedges)
317 45630 : tmp_sft(:, :) = nequip_work%local_shifts(:, 1:nedges)
318 :
319 6 : CALL MOVE_ALLOC(tmp_idx, nequip_work%local_edges)
320 6 : CALL MOVE_ALLOC(tmp_sft, nequip_work%local_shifts)
321 : END BLOCK
322 : END IF
323 :
324 6 : END SUBROUTINE build_local_edges_shifts
325 :
326 : ! **************************************************************************************************
327 : !> \brief ...
328 : !> \param atomic_kind_set ...
329 : !> \param potparm ...
330 : !> \param target_pot_type ...
331 : !> \param neq_pot ...
332 : !> \par History
333 : !> Get the NequIP or Allegro potential - [gtocci] 2026
334 : !> \author Gabriele Tocci - University of Zurich
335 : ! **************************************************************************************************
336 6 : SUBROUTINE get_potential_config(atomic_kind_set, potparm, target_pot_type, neq_pot)
337 : TYPE(atomic_kind_type), POINTER :: atomic_kind_set(:)
338 : TYPE(pair_potential_pp_type), POINTER :: potparm
339 : INTEGER, INTENT(IN) :: target_pot_type
340 : TYPE(nequip_pot_type), INTENT(OUT), POINTER :: neq_pot
341 :
342 : INTEGER :: i, ikind, jkind
343 : TYPE(pair_potential_single_type), POINTER :: pot
344 :
345 6 : NULLIFY (neq_pot)
346 6 : OuterLoop: DO ikind = 1, SIZE(atomic_kind_set)
347 6 : DO jkind = ikind, SIZE(atomic_kind_set)
348 6 : pot => potparm%pot(ikind, jkind)%pot
349 6 : DO i = 1, SIZE(pot%type)
350 6 : IF (pot%type(i) == target_pot_type) THEN
351 6 : neq_pot => pot%set(i)%nequip
352 6 : EXIT OuterLoop
353 : END IF
354 : END DO
355 : END DO
356 : END DO OuterLoop
357 6 : END SUBROUTINE get_potential_config
358 :
359 : ! **************************************************************************************************
360 : !> \brief ...
361 : !> \param nequip_work ...
362 : !> \par History
363 : !> Inits masks for torch evaluation (use_atom) and MPI summation (sum_energy) - [gtocci] 2026
364 : !> \author Gabriele Tocci - University of Zurich
365 : ! **************************************************************************************************
366 6 : SUBROUTINE init_atom_masks(nequip_work)
367 : TYPE(nequip_work_type), INTENT(INOUT) :: nequip_work
368 :
369 : INTEGER :: iat, ikind, ilocal, n_atoms, n_local
370 :
371 6 : IF (.NOT. ALLOCATED(nequip_work%kind_mapper)) THEN
372 0 : CPABORT("kind_mapper not initialized before init_atom_masks")
373 : END IF
374 :
375 6 : n_atoms = SIZE(nequip_work%particle_set)
376 :
377 6 : IF (ALLOCATED(nequip_work%use_atom)) DEALLOCATE (nequip_work%use_atom)
378 18 : ALLOCATE (nequip_work%use_atom(n_atoms))
379 454 : nequip_work%use_atom = .FALSE.
380 :
381 454 : DO iat = 1, n_atoms
382 448 : ikind = nequip_work%particle_set(iat)%atomic_kind%kind_number
383 454 : IF (nequip_work%kind_mapper(ikind) > 0) THEN
384 448 : nequip_work%use_atom(iat) = .TRUE.
385 : END IF
386 : END DO
387 454 : nequip_work%n_atoms_use = COUNT(nequip_work%use_atom)
388 :
389 6 : IF (ALLOCATED(nequip_work%sum_energy)) DEALLOCATE (nequip_work%sum_energy)
390 12 : ALLOCATE (nequip_work%sum_energy(n_atoms))
391 454 : nequip_work%sum_energy = .FALSE.
392 :
393 6 : IF (ASSOCIATED(nequip_work%local_particles)) THEN
394 16 : DO ikind = 1, SIZE(nequip_work%local_particles%n_el)
395 16 : IF (nequip_work%kind_mapper(ikind) > 0) THEN
396 10 : n_local = nequip_work%local_particles%n_el(ikind)
397 234 : DO ilocal = 1, n_local
398 224 : iat = nequip_work%local_particles%list(ikind)%array(ilocal)
399 234 : nequip_work%sum_energy(iat) = .TRUE.
400 : END DO
401 : END IF
402 : END DO
403 : ELSE
404 0 : nequip_work%sum_energy(:) = nequip_work%use_atom(:)
405 : END IF
406 :
407 6 : END SUBROUTINE init_atom_masks
408 :
409 : ! **************************************************************************************************
410 : !> \brief ...
411 : !> \param atomic_kind_set ...
412 : !> \param neq_pot ...
413 : !> \param nequip_work ...
414 : !> \author Gabriele Tocci - University of Zurich
415 : ! **************************************************************************************************
416 6 : SUBROUTINE build_kind_mapper(atomic_kind_set, neq_pot, nequip_work)
417 : TYPE(atomic_kind_type), POINTER :: atomic_kind_set(:)
418 : TYPE(nequip_pot_type), POINTER :: neq_pot
419 : TYPE(nequip_work_type), INTENT(INOUT) :: nequip_work
420 :
421 : CHARACTER(LEN=100) :: model_sym
422 : CHARACTER(LEN=default_string_length) :: kind_sym
423 : INTEGER :: i, ikind, n_kinds
424 :
425 6 : n_kinds = SIZE(atomic_kind_set)
426 :
427 6 : IF (ALLOCATED(nequip_work%kind_mapper)) DEALLOCATE (nequip_work%kind_mapper)
428 18 : ALLOCATE (nequip_work%kind_mapper(n_kinds))
429 16 : nequip_work%kind_mapper = -1
430 :
431 16 : DO ikind = 1, n_kinds
432 10 : kind_sym = atomic_kind_set(ikind)%element_symbol
433 10 : CALL uppercase(kind_sym)
434 :
435 30 : DO i = 1, neq_pot%num_types
436 24 : model_sym = neq_pot%type_names_torch(i)
437 24 : CALL uppercase(model_sym)
438 24 : IF (TRIM(kind_sym) == TRIM(model_sym)) THEN
439 10 : nequip_work%kind_mapper(ikind) = i
440 10 : EXIT
441 : END IF
442 : END DO
443 : END DO
444 6 : END SUBROUTINE build_kind_mapper
445 :
446 : ! **************************************************************************************************
447 : !> \brief ...
448 : !> \param fist_nonbond_env ...
449 : !> \param neq_data ...
450 : !> \param pot ...
451 : !> \param nequip_work ...
452 : !> \par History
453 : !> load the NequIP/Allegro model, initialize forces, positions - [gtocci] 2026
454 : !> \author Gabriele Tocci - University of Zurich
455 : ! **************************************************************************************************
456 6 : SUBROUTINE setup_neq_data(fist_nonbond_env, neq_data, pot, nequip_work)
457 : TYPE(fist_nonbond_env_type), POINTER :: fist_nonbond_env
458 : TYPE(nequip_data_type), POINTER :: neq_data
459 : TYPE(nequip_pot_type), POINTER :: pot
460 : TYPE(nequip_work_type), INTENT(IN) :: nequip_work
461 :
462 : INTEGER :: iat, iat_use, n_atoms
463 :
464 6 : CALL fist_nonbond_env_get(fist_nonbond_env, nequip_data=neq_data)
465 :
466 6 : IF (.NOT. ASSOCIATED(neq_data)) THEN
467 84 : ALLOCATE (neq_data)
468 6 : CALL fist_nonbond_env_set(fist_nonbond_env, nequip_data=neq_data)
469 6 : NULLIFY (neq_data%use_indices, neq_data%force)
470 :
471 6 : CALL torch_model_load(neq_data%model, pot%pot_file_name)
472 6 : CALL torch_model_freeze(neq_data%model)
473 : END IF
474 :
475 6 : IF (ASSOCIATED(neq_data%force)) THEN
476 0 : IF (SIZE(neq_data%force, 2) /= nequip_work%n_atoms_use) THEN
477 0 : DEALLOCATE (neq_data%force, neq_data%use_indices)
478 : END IF
479 : END IF
480 :
481 6 : IF (.NOT. ASSOCIATED(neq_data%force)) THEN
482 18 : ALLOCATE (neq_data%force(3, nequip_work%n_atoms_use))
483 18 : ALLOCATE (neq_data%use_indices(nequip_work%n_atoms_use))
484 : END IF
485 :
486 6 : n_atoms = SIZE(nequip_work%use_atom)
487 6 : iat_use = 0
488 454 : DO iat = 1, n_atoms
489 454 : IF (nequip_work%use_atom(iat)) THEN
490 448 : iat_use = iat_use + 1
491 448 : neq_data%use_indices(iat_use) = iat
492 : END IF
493 : END DO
494 6 : END SUBROUTINE setup_neq_data
495 :
496 : ! **************************************************************************************************
497 : !> \brief ...
498 : !> \param nequip_work ...
499 : !> \par History
500 : !> Prepare edges and cell shifts for NequIP - [gtocci] 2026
501 : !> \author Gabriele Tocci - University of Zurich
502 : ! **************************************************************************************************
503 4 : SUBROUTINE prepare_edges_shifts_nequip(nequip_work)
504 : TYPE(nequip_work_type), INTENT(INOUT) :: nequip_work
505 :
506 : INTEGER :: ipair, nedges, nedges_tot
507 : INTEGER(kind=int_8), ALLOCATABLE :: temp_edge_index(:, :)
508 : INTEGER, ALLOCATABLE :: displ(:), displ_cell(:), edge_count(:), &
509 : edge_count_cell(:)
510 :
511 4 : nedges = SIZE(nequip_work%local_edges, 2)
512 :
513 16 : ALLOCATE (edge_count(nequip_work%para_env%num_pe), edge_count_cell(nequip_work%para_env%num_pe))
514 12 : ALLOCATE (displ_cell(nequip_work%para_env%num_pe), displ(nequip_work%para_env%num_pe))
515 :
516 4 : CALL nequip_work%para_env%allgather(nedges, edge_count)
517 12 : nedges_tot = SUM(edge_count)
518 :
519 12 : ALLOCATE (temp_edge_index(2, nedges_tot))
520 12 : ALLOCATE (nequip_work%final_shifts(3, nedges_tot))
521 :
522 12 : edge_count_cell(:) = edge_count*3
523 12 : edge_count = edge_count*2
524 4 : displ(1) = 0
525 4 : displ_cell(1) = 0
526 8 : DO ipair = 2, nequip_work%para_env%num_pe
527 4 : displ(ipair) = displ(ipair - 1) + edge_count(ipair - 1)
528 8 : displ_cell(ipair) = displ_cell(ipair - 1) + edge_count_cell(ipair - 1)
529 : END DO
530 :
531 4 : CALL nequip_work%para_env%allgatherv(nequip_work%local_shifts, nequip_work%final_shifts, edge_count_cell, displ_cell)
532 4 : CALL nequip_work%para_env%allgatherv(nequip_work%local_edges, temp_edge_index, edge_count, displ)
533 :
534 8 : ALLOCATE (nequip_work%final_edges(nedges_tot, 2))
535 25740 : nequip_work%final_edges(:, :) = TRANSPOSE(temp_edge_index)
536 :
537 4 : DEALLOCATE (edge_count, edge_count_cell, displ, displ_cell, temp_edge_index)
538 :
539 4 : END SUBROUTINE prepare_edges_shifts_nequip
540 :
541 : ! **************************************************************************************************
542 : !> \brief ...
543 : !> \param nequip_work ...
544 : !> \par History
545 : !> Prepare edges and cell shifts for Allegro - [gtocci] 2026
546 : !> \author Gabriele Tocci - University of Zurich
547 : ! **************************************************************************************************
548 2 : SUBROUTINE prepare_edges_shifts_allegro(nequip_work)
549 : TYPE(nequip_work_type), INTENT(INOUT) :: nequip_work
550 :
551 19904 : ALLOCATE (nequip_work%final_shifts, SOURCE=nequip_work%local_shifts)
552 6 : ALLOCATE (nequip_work%final_edges(SIZE(nequip_work%local_edges, 2), 2))
553 19908 : nequip_work%final_edges(:, :) = TRANSPOSE(nequip_work%local_edges)
554 2 : END SUBROUTINE prepare_edges_shifts_allegro
555 :
556 : ! **************************************************************************************************
557 : !> \brief ...
558 : !> \param nequip_work ...
559 : !> \par History
560 : !> Build edges from cp2k global neigh lists to local/packed ones for torch - [gtocci] 2026
561 : !> \author Gabriele Tocci - University of Zurich
562 : ! **************************************************************************************************
563 6 : SUBROUTINE build_torch_edge_indexes(nequip_work)
564 : TYPE(nequip_work_type), INTENT(INOUT) :: nequip_work
565 :
566 : INTEGER :: atom_a, atom_b, i, iat, iat_use, n_atoms
567 6 : INTEGER, ALLOCATABLE :: global_to_packed(:)
568 :
569 6 : n_atoms = SIZE(nequip_work%particle_set)
570 :
571 : ! for allegro ensure ghost atoms are included in the evaluation
572 6 : IF (nequip_work%target_pot_type /= nequip_type .AND. &
573 : nequip_work%target_pot_type /= mace_type) THEN
574 : ! label atoms in the local edges
575 4976 : DO i = 1, SIZE(nequip_work%local_edges, 2)
576 4974 : atom_a = INT(nequip_work%local_edges(1, i))
577 4974 : atom_b = INT(nequip_work%local_edges(2, i))
578 4974 : nequip_work%use_atom(atom_a) = .TRUE.
579 4976 : nequip_work%use_atom(atom_b) = .TRUE.
580 : END DO
581 194 : nequip_work%n_atoms_use = COUNT(nequip_work%use_atom)
582 : END IF
583 :
584 : ! mapping from global CP2K index to packed/local Torch index
585 18 : ALLOCATE (global_to_packed(n_atoms))
586 6 : global_to_packed = 0
587 6 : iat_use = 0
588 454 : DO iat = 1, n_atoms
589 454 : IF (nequip_work%use_atom(iat)) THEN
590 448 : iat_use = iat_use + 1
591 448 : global_to_packed(iat) = iat_use
592 : END IF
593 : END DO
594 :
595 : ! remap local_edges to use 0-based dense indices for torch
596 11412 : DO i = 1, SIZE(nequip_work%local_edges, 2)
597 11406 : atom_a = INT(nequip_work%local_edges(1, i))
598 11406 : atom_b = INT(nequip_work%local_edges(2, i))
599 :
600 11406 : nequip_work%local_edges(1, i) = INT(global_to_packed(atom_a) - 1, kind=int_8)
601 11412 : nequip_work%local_edges(2, i) = INT(global_to_packed(atom_b) - 1, kind=int_8)
602 : END DO
603 :
604 6 : DEALLOCATE (global_to_packed)
605 :
606 6 : END SUBROUTINE build_torch_edge_indexes
607 :
608 : ! **************************************************************************************************
609 : !> \brief ...
610 : !> \param neq_data ...
611 : !> \param pot ...
612 : !> \param nequip_work ...
613 : !> \param outputs ...
614 : !> \par History
615 : !> Run forward pass using torch api - [gtocci] 2026
616 : !> \author Gabriele Tocci - University of Zurich
617 : ! **************************************************************************************************
618 6 : SUBROUTINE run_torch_model(neq_data, pot, nequip_work, outputs)
619 : TYPE(nequip_data_type), POINTER :: neq_data
620 : TYPE(nequip_pot_type), POINTER :: pot
621 : TYPE(nequip_work_type), INTENT(IN) :: nequip_work
622 : TYPE(torch_dict_type), INTENT(OUT) :: outputs
623 :
624 : INTEGER :: iat, iat_use, ikind
625 : INTEGER(kind=int_8), ALLOCATABLE :: atom_types(:)
626 : REAL(kind=dp), ALLOCATABLE :: lattice(:, :), pos(:, :)
627 : TYPE(torch_dict_type) :: inputs
628 : TYPE(torch_tensor_type) :: cell_t, idx_t, pos_t, shift_t, types_t
629 :
630 0 : ALLOCATE (lattice(3, 3))
631 78 : lattice(:, :) = nequip_work%cell%hmat/pot%unit_length_val
632 :
633 30 : ALLOCATE (pos(3, nequip_work%n_atoms_use), atom_types(nequip_work%n_atoms_use))
634 6 : iat_use = 0
635 454 : DO iat = 1, SIZE(nequip_work%particle_set)
636 448 : IF (.NOT. nequip_work%use_atom(iat)) CYCLE
637 448 : iat_use = iat_use + 1
638 :
639 448 : ikind = nequip_work%particle_set(iat)%atomic_kind%kind_number
640 448 : IF (nequip_work%kind_mapper(ikind) < 1) THEN
641 0 : CALL cp_abort(__LOCATION__, "Atom symbol not found in NequIP model!")
642 : END IF
643 :
644 : ! Convert 1-based Fortran index to 0-based PyTorch index
645 448 : atom_types(iat_use) = nequip_work%kind_mapper(ikind) - 1
646 1798 : pos(:, iat_use) = nequip_work%r_pbc(iat)%r(:)/pot%unit_length_val
647 : END DO
648 :
649 6 : CALL torch_dict_create(inputs)
650 :
651 6 : CALL torch_tensor_from_array(pos_t, pos)
652 6 : CALL torch_tensor_from_array(shift_t, nequip_work%final_shifts)
653 6 : CALL torch_tensor_from_array(cell_t, lattice)
654 :
655 6 : CALL torch_dict_insert(inputs, "pos", pos_t)
656 6 : CALL torch_dict_insert(inputs, "edge_cell_shift", shift_t)
657 6 : CALL torch_dict_insert(inputs, "cell", cell_t)
658 6 : CALL torch_tensor_release(pos_t)
659 6 : CALL torch_tensor_release(shift_t)
660 6 : CALL torch_tensor_release(cell_t)
661 :
662 6 : CALL torch_tensor_from_array(idx_t, nequip_work%final_edges)
663 6 : CALL torch_dict_insert(inputs, "edge_index", idx_t)
664 6 : CALL torch_tensor_release(idx_t)
665 :
666 6 : CALL torch_tensor_from_array(types_t, atom_types)
667 6 : CALL torch_dict_insert(inputs, "atom_types", types_t)
668 6 : CALL torch_tensor_release(types_t)
669 :
670 6 : CALL torch_dict_create(outputs)
671 6 : CALL torch_model_forward(neq_data%model, inputs, outputs)
672 :
673 6 : CALL torch_dict_release(inputs)
674 :
675 6 : IF (ALLOCATED(pos)) DEALLOCATE (pos)
676 6 : IF (ALLOCATED(lattice)) DEALLOCATE (lattice)
677 6 : IF (ALLOCATED(atom_types)) DEALLOCATE (atom_types)
678 :
679 12 : END SUBROUTINE run_torch_model
680 :
681 : ! **************************************************************************************************
682 : !> \brief ...
683 : !> \param outputs ...
684 : !> \param neq_data ...
685 : !> \param pot ...
686 : !> \param pot_total ...
687 : !> \param nequip_work ...
688 : !> \par History
689 : !> Collect potential, forces, virial - [gtocci] 2026
690 : !> \author Gabriele Tocci - University of Zurich
691 : ! **************************************************************************************************
692 6 : SUBROUTINE process_outputs(outputs, neq_data, pot, pot_total, nequip_work)
693 : TYPE(torch_dict_type), INTENT(IN) :: outputs
694 : TYPE(nequip_data_type), POINTER :: neq_data
695 : TYPE(nequip_pot_type), POINTER :: pot
696 : REAL(kind=dp), INTENT(OUT) :: pot_total
697 : TYPE(nequip_work_type), INTENT(IN) :: nequip_work
698 :
699 : INTEGER :: iat, iat_use
700 6 : REAL(kind=dp), POINTER :: e_ptr(:, :), f_ptr(:, :), v_ptr(:, :, :)
701 : TYPE(torch_tensor_type) :: t_energy, t_forces, t_virial
702 :
703 6 : NULLIFY (f_ptr, e_ptr, v_ptr)
704 :
705 6 : CALL torch_dict_get(outputs, "forces", t_forces)
706 6 : CALL torch_tensor_data_ptr(t_forces, f_ptr)
707 :
708 3596 : neq_data%force = f_ptr*pot%unit_forces_val
709 6 : CALL torch_tensor_release(t_forces)
710 6 : CALL torch_dict_get(outputs, "atomic_energy", t_energy)
711 6 : CALL torch_tensor_data_ptr(t_energy, e_ptr)
712 :
713 6 : pot_total = 0.0_dp
714 454 : DO iat_use = 1, SIZE(neq_data%use_indices)
715 448 : iat = neq_data%use_indices(iat_use)
716 : ! Only apply the local mask for Allegro models
717 448 : IF (nequip_work%target_pot_type /= nequip_type .AND. &
718 : nequip_work%target_pot_type /= mace_type) THEN
719 192 : IF (.NOT. nequip_work%sum_energy(iat)) CYCLE
720 : END IF
721 :
722 454 : pot_total = pot_total + e_ptr(1, iat_use)
723 : END DO
724 6 : CALL torch_tensor_release(t_energy)
725 6 : pot_total = pot_total*pot%unit_energy_val
726 :
727 6 : IF (nequip_work%target_pot_type == nequip_type .OR. &
728 : nequip_work%target_pot_type == mace_type) THEN
729 1028 : neq_data%force = neq_data%force/REAL(nequip_work%para_env%num_pe, dp)
730 4 : pot_total = pot_total/REAL(nequip_work%para_env%num_pe, dp)
731 : END IF
732 :
733 6 : IF (nequip_work%use_virial) THEN
734 4 : CALL torch_dict_get(outputs, "virial", t_virial)
735 4 : CALL torch_tensor_data_ptr(t_virial, v_ptr)
736 :
737 52 : neq_data%virial(:, :) = RESHAPE(v_ptr, [3, 3])*pot%unit_energy_val
738 4 : CALL torch_tensor_release(t_virial)
739 4 : IF (nequip_work%target_pot_type == nequip_type .OR. &
740 : nequip_work%target_pot_type == mace_type) THEN
741 26 : neq_data%virial = neq_data%virial/REAL(nequip_work%para_env%num_pe, dp)
742 : END IF
743 : END IF
744 :
745 6 : END SUBROUTINE process_outputs
746 :
747 : ! **************************************************************************************************
748 : !> \brief ...
749 : !> \param fist_nonbond_env ...
750 : !> \param f_nonbond ...
751 : !> \param pv_nonbond ...
752 : !> \param use_virial ...
753 : !> \par History
754 : !> Sum forces, virial to nonbond - [gtocci] 2026
755 : !> \author Gabriele Tocci - University of Zurich
756 : ! **************************************************************************************************
757 6 : SUBROUTINE e3nn_add_force_virial(fist_nonbond_env, f_nonbond, pv_nonbond, use_virial)
758 : TYPE(fist_nonbond_env_type), POINTER :: fist_nonbond_env
759 : REAL(KIND=dp), DIMENSION(:, :), INTENT(INOUT) :: f_nonbond, pv_nonbond
760 : LOGICAL, INTENT(IN) :: use_virial
761 :
762 : INTEGER :: iat, iat_use
763 : TYPE(nequip_data_type), POINTER :: neq_data
764 :
765 6 : CALL fist_nonbond_env_get(fist_nonbond_env, nequip_data=neq_data)
766 :
767 6 : IF (use_virial) THEN
768 52 : pv_nonbond = pv_nonbond + neq_data%virial
769 : END IF
770 :
771 454 : DO iat_use = 1, SIZE(neq_data%use_indices)
772 448 : iat = neq_data%use_indices(iat_use)
773 1798 : f_nonbond(1:3, iat) = f_nonbond(1:3, iat) + neq_data%force(1:3, iat_use)
774 : END DO
775 :
776 6 : END SUBROUTINE e3nn_add_force_virial
777 :
778 : END MODULE manybody_e3nn
|