LCOV - code coverage report
Current view: top level - src - pao_model.F (source / functions) Coverage Total Hit
Test: CP2K Regtests (git:71c3ab0) Lines: 95.3 % 169 161
Test Date: 2026-07-25 06:35:44 Functions: 100.0 % 4 4

            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
        

Generated by: LCOV version 2.0-1