LCOV - code coverage report
Current view: top level - src - manybody_e3nn.F (source / functions) Coverage Total Hit
Test: CP2K Regtests (git:71c3ab0) Lines: 97.2 % 282 274
Test Date: 2026-07-25 06:35:44 Functions: 100.0 % 14 14

            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
        

Generated by: LCOV version 2.0-1