PAO: Add forces for equivariant PyTorch models

This commit is contained in:
Ole Schütt 2025-01-25 15:32:33 +01:00 committed by Ole Schütt
parent b6963b94bc
commit 5912d788db
6 changed files with 184 additions and 35 deletions

View file

@ -44,7 +44,8 @@ MODULE pao_methods
diamat_all
USE message_passing, ONLY: mp_para_env_type
USE pao_ml, ONLY: pao_ml_forces
USE pao_model, ONLY: pao_model_load
USE pao_model, ONLY: pao_model_forces,&
pao_model_load
USE pao_param, ONLY: pao_calc_AB,&
pao_param_count
USE pao_types, ONLY: pao_env_type
@ -999,7 +1000,7 @@ CONTAINS
CALL pao_ml_forces(pao, qs_env, pao%matrix_G, forces)
IF (ALLOCATED(pao%models)) &
CPABORT("PAO forces for PyTorch models are not yet implemented.")
CALL pao_model_forces(pao, qs_env, pao%matrix_G, forces)
CALL para_env%sum(forces)
DO iatom = 1, natoms

View file

@ -10,16 +10,19 @@
!> \author Ole Schuett
! **************************************************************************************************
MODULE pao_model
USE OMP_LIB, ONLY: omp_init_lock
USE atomic_kind_types, ONLY: atomic_kind_type,&
get_atomic_kind
USE basis_set_types, ONLY: gto_basis_set_type
USE cell_types, ONLY: cell_type,&
pbc
USE cp_dbcsr_api, ONLY: dbcsr_iterator_blocks_left,&
USE cp_dbcsr_api, ONLY: dbcsr_get_info,&
dbcsr_iterator_blocks_left,&
dbcsr_iterator_next_block,&
dbcsr_iterator_start,&
dbcsr_iterator_stop,&
dbcsr_iterator_type
dbcsr_iterator_type,&
dbcsr_type
USE kinds, ONLY: default_path_length,&
default_string_length,&
dp,&
@ -35,8 +38,9 @@ MODULE pao_model
qs_kind_type
USE torch_api, ONLY: &
torch_dict_create, torch_dict_get, torch_dict_insert, torch_dict_release, torch_dict_type, &
torch_model_forward, torch_model_get_attr, torch_model_load, torch_tensor_data_ptr, &
torch_tensor_from_array, torch_tensor_release, torch_tensor_type
torch_model_forward, torch_model_get_attr, torch_model_load, torch_tensor_backward, &
torch_tensor_data_ptr, torch_tensor_from_array, torch_tensor_grad, torch_tensor_release, &
torch_tensor_type
USE util, ONLY: sort
#include "./base/base_uses.f90"
@ -44,7 +48,7 @@ MODULE pao_model
PRIVATE
PUBLIC :: pao_model_load, pao_model_predict, pao_model_type
PUBLIC :: pao_model_load, pao_model_predict, pao_model_forces, pao_model_type
CONTAINS
@ -136,6 +140,7 @@ CONTAINS
IF (model%pao_basis_size /= pao_basis_size) &
CPABORT("PAO basis size does not match.")
CALL omp_init_lock(model%lock)
CALL timestop(handle)
END SUBROUTINE pao_model_load
@ -163,7 +168,7 @@ CONTAINS
CALL dbcsr_iterator_next_block(iter, arow, acol, block_X)
IF (SIZE(block_X) == 0) CYCLE ! pao disabled for iatom
iatom = arow; CPASSERT(arow == acol)
CALL predict_single_atom(pao, qs_env, iatom, block_X)
CALL predict_single_atom(pao, qs_env, iatom, block_X=block_X)
END DO
CALL dbcsr_iterator_stop(iter)
!$OMP END PARALLEL
@ -172,25 +177,65 @@ CONTAINS
END SUBROUTINE pao_model_predict
! **************************************************************************************************
!> \brief Calculate forces contributed by machine learning
!> \param pao ...
!> \param qs_env ...
!> \param matrix_G ...
!> \param forces ...
! **************************************************************************************************
SUBROUTINE pao_model_forces(pao, qs_env, matrix_G, forces)
TYPE(pao_env_type), POINTER :: pao
TYPE(qs_environment_type), POINTER :: qs_env
TYPE(dbcsr_type) :: matrix_G
REAL(dp), DIMENSION(:, :), INTENT(INOUT) :: forces
CHARACTER(len=*), PARAMETER :: routineN = 'pao_model_forces'
INTEGER :: acol, arow, handle, iatom
REAL(dp), DIMENSION(:, :), POINTER :: block_G
TYPE(dbcsr_iterator_type) :: iter
CALL timeset(routineN, handle)
!$OMP PARALLEL DEFAULT(NONE) SHARED(pao,qs_env,matrix_G,forces) PRIVATE(iter,arow,acol,iatom,block_G)
CALL dbcsr_iterator_start(iter, matrix_G)
DO WHILE (dbcsr_iterator_blocks_left(iter))
CALL dbcsr_iterator_next_block(iter, arow, acol, block_G)
iatom = arow; CPASSERT(arow == acol)
IF (SIZE(block_G) == 0) CYCLE ! pao disabled for iatom
CALL predict_single_atom(pao, qs_env, iatom, block_G=block_G, forces=forces)
END DO
CALL dbcsr_iterator_stop(iter)
!$OMP END PARALLEL
CALL timestop(handle)
END SUBROUTINE pao_model_forces
! **************************************************************************************************
!> \brief Predicts a single block_X.
!> \param pao ...
!> \param qs_env ...
!> \param iatom ...
!> \param block_X ...
!> \param block_G ...
!> \param forces ...
! **************************************************************************************************
SUBROUTINE predict_single_atom(pao, qs_env, iatom, block_X)
SUBROUTINE predict_single_atom(pao, qs_env, iatom, block_X, block_G, forces)
TYPE(pao_env_type), INTENT(IN), POINTER :: pao
TYPE(qs_environment_type), INTENT(IN), POINTER :: qs_env
INTEGER, INTENT(IN) :: iatom
REAL(dp), DIMENSION(:, :), INTENT(OUT) :: block_X
REAL(dp), DIMENSION(:, :), OPTIONAL :: block_X, block_G, forces
INTEGER :: ikind, jatom, jkind, jneighbor, natoms
INTEGER :: ikind, jatom, jkind, jneighbor, m, n, &
natoms
INTEGER, ALLOCATABLE, DIMENSION(:) :: neighbors_index
INTEGER, DIMENSION(:), POINTER :: blk_sizes_pao, blk_sizes_pri
REAL(dp), DIMENSION(3) :: Ri, Rij, Rj
REAL(KIND=dp), ALLOCATABLE, DIMENSION(:) :: neighbors_distance
REAL(sp), ALLOCATABLE, DIMENSION(:, :) :: neighbors_features, neighbors_relpos
REAL(sp), DIMENSION(:, :), POINTER :: predicted_xblock
REAL(sp), ALLOCATABLE, DIMENSION(:, :) :: features, outer_grad, relpos
REAL(sp), DIMENSION(:, :), POINTER :: predicted_xblock, relpos_grad
TYPE(atomic_kind_type), DIMENSION(:), POINTER :: atomic_kind_set
TYPE(cell_type), POINTER :: cell
TYPE(mp_para_env_type), POINTER :: para_env
@ -198,9 +243,13 @@ CONTAINS
TYPE(particle_type), DIMENSION(:), POINTER :: particle_set
TYPE(qs_kind_type), DIMENSION(:), POINTER :: qs_kind_set
TYPE(torch_dict_type) :: model_inputs, model_outputs
TYPE(torch_tensor_type) :: neighbors_features_tensor, &
neighbors_relpos_tensor, &
predicted_xblock_tensor
TYPE(torch_tensor_type) :: features_tensor, outer_grad_tensor, &
predicted_xblock_tensor, &
relpos_grad_tensor, relpos_tensor
CALL dbcsr_get_info(pao%matrix_Y, row_blk_size=blk_sizes_pri, col_blk_size=blk_sizes_pao)
n = blk_sizes_pri(iatom) ! size of primary basis
m = blk_sizes_pao(iatom) ! size of pao basis
CALL get_qs_env(qs_env, &
para_env=para_env, &
@ -213,6 +262,7 @@ CONTAINS
CALL get_atomic_kind(particle_set(iatom)%atomic_kind, kind_number=ikind)
model => pao%models(ikind)
CPASSERT(model%version > 0)
CALL omp_set_lock(model%lock) ! TODO: might not be needed for inference.
! Find neighbors.
! TODO: this is a quadratic algorithm, use a neighbor-list instead
@ -227,34 +277,31 @@ CONTAINS
CPASSERT(neighbors_index(1) == iatom) ! central atom should be closesd to itself
! Compute neighbors relative positions.
ALLOCATE (neighbors_relpos(3, model%num_neighbors))
neighbors_relpos(:, :) = 0.0_sp
ALLOCATE (relpos(3, model%num_neighbors))
relpos(:, :) = 0.0_sp
DO jneighbor = 1, MIN(model%num_neighbors, natoms - 1)
jatom = neighbors_index(jneighbor + 1) ! skipping central atom
Rj = particle_set(jatom)%r
Rij = pbc(Ri, Rj, cell)
neighbors_relpos(:, jneighbor) = REAL(angstrom*Rij, kind=sp)
relpos(:, jneighbor) = REAL(angstrom*Rij, kind=sp)
END DO
! Compute neighbors features.
ALLOCATE (neighbors_features(SIZE(model%feature_kinds), model%num_neighbors))
neighbors_features(:, :) = 0.0_sp
ALLOCATE (features(SIZE(model%feature_kinds), model%num_neighbors))
features(:, :) = 0.0_sp
DO jneighbor = 1, MIN(model%num_neighbors, natoms - 1)
jatom = neighbors_index(jneighbor + 1) ! skipping central atom
jkind = particle_set(jatom)%atomic_kind%kind_number
WHERE (model%feature_kinds == jkind) neighbors_features(:, jneighbor) = 1.0_sp
WHERE (model%feature_kinds == jkind) features(:, jneighbor) = 1.0_sp
END DO
! Inference.
CALL torch_dict_create(model_inputs)
CALL torch_tensor_from_array(neighbors_relpos_tensor, neighbors_relpos)
CALL torch_dict_insert(model_inputs, "neighbors_relpos", neighbors_relpos_tensor)
CALL torch_tensor_release(neighbors_relpos_tensor)
CALL torch_tensor_from_array(neighbors_features_tensor, neighbors_features)
CALL torch_dict_insert(model_inputs, "neighbors_features", neighbors_features_tensor)
CALL torch_tensor_release(neighbors_features_tensor)
CALL torch_tensor_from_array(relpos_tensor, relpos, requires_grad=PRESENT(block_G))
CALL torch_dict_insert(model_inputs, "neighbors_relpos", relpos_tensor)
CALL torch_tensor_from_array(features_tensor, features)
CALL torch_dict_insert(model_inputs, "neighbors_features", features_tensor)
CALL torch_dict_create(model_outputs)
CALL torch_model_forward(model%torch_model, model_inputs, model_outputs)
@ -263,13 +310,38 @@ CONTAINS
NULLIFY (predicted_xblock)
CALL torch_dict_get(model_outputs, "xblock", predicted_xblock_tensor)
CALL torch_tensor_data_ptr(predicted_xblock_tensor, predicted_xblock)
block_X = RESHAPE(predicted_xblock, (/SIZE(block_X), 1/))
CALL torch_tensor_release(predicted_xblock_tensor)
CPASSERT(SIZE(predicted_xblock, 1) == n .AND. SIZE(predicted_xblock, 2) == m)
IF (PRESENT(block_X)) THEN
block_X = RESHAPE(predicted_xblock, [n*m, 1])
END IF
! TURNING POINT (if calc forces) ------------------------------------------
IF (PRESENT(block_G)) THEN
ALLOCATE (outer_grad(n, m))
outer_grad(:, :) = REAL(RESHAPE(block_G, [n, m]), kind=sp)
CALL torch_tensor_from_array(outer_grad_tensor, outer_grad)
CALL torch_tensor_backward(predicted_xblock_tensor, outer_grad_tensor)
CALL torch_tensor_grad(relpos_tensor, relpos_grad_tensor)
NULLIFY (relpos_grad)
CALL torch_tensor_data_ptr(relpos_grad_tensor, relpos_grad)
CPASSERT(SIZE(relpos_grad, 1) == 3 .AND. SIZE(relpos_grad, 2) == model%num_neighbors)
DO jneighbor = 1, MIN(model%num_neighbors, natoms - 1)
jatom = neighbors_index(jneighbor + 1) ! skipping central atom
forces(iatom, :) = forces(iatom, :) + relpos_grad(:, jneighbor)*angstrom
forces(jatom, :) = forces(jatom, :) - relpos_grad(:, jneighbor)*angstrom
END DO
CALL torch_tensor_release(outer_grad_tensor)
CALL torch_tensor_release(relpos_grad_tensor)
END IF
! Clean up.
CALL torch_tensor_release(relpos_tensor)
CALL torch_tensor_release(features_tensor)
CALL torch_tensor_release(predicted_xblock_tensor)
CALL torch_dict_release(model_inputs)
CALL torch_dict_release(model_outputs)
DEALLOCATE (neighbors_distance, neighbors_index, neighbors_relpos, neighbors_features)
DEALLOCATE (neighbors_distance, neighbors_index, relpos, features)
CALL omp_unset_lock(model%lock)
END SUBROUTINE predict_single_atom

View file

@ -10,6 +10,8 @@
!> \author Ole Schuett
! **************************************************************************************************
MODULE pao_types
USE OMP_LIB, ONLY: omp_destroy_lock,&
omp_lock_kind
USE cp_dbcsr_api, ONLY: dbcsr_distribution_release,&
dbcsr_distribution_type,&
dbcsr_release,&
@ -66,6 +68,7 @@ MODULE pao_types
INTEGER :: num_neighbors = -1
REAL(dp) :: cutoff = 0.0_dp
INTEGER, DIMENSION(:), ALLOCATABLE :: feature_kinds
INTEGER(kind=omp_lock_kind) :: lock = -1
END TYPE pao_model_type
! **************************************************************************************************
@ -257,6 +260,7 @@ CONTAINS
IF (pao%models(i)%version > 0) THEN
CALL torch_model_release(pao%models(i)%torch_model)
END IF
CALL omp_destroy_lock(pao%models(i)%lock)
END DO
DEALLOCATE (pao%models)
END IF

View file

@ -10,8 +10,7 @@
BASIS_SET_FILE_NAME BASIS_MOLOPT
POTENTIAL_FILE_NAME GTH_POTENTIALS
&LS_SCF
EPS_FILTER 1.0E-8
EPS_SCF 1.0E-6
EPS_SCF 1.0E-8
EXTRAPOLATION_ORDER 1
MAX_SCF 25
PURIFICATION_METHOD TRS4

View file

@ -0,0 +1,72 @@
&GLOBAL
PROJECT 2H2O_equivar_ML_checkforces
RUN_TYPE DEBUG
&END GLOBAL
&DEBUG
CHECK_ATOM_FORCE 1 XYZ
DEBUG_FORCES YES
MAX_RELATIVE_ERROR 2.0
STOP_ON_MISMATCH YES
&END DEBUG
&FORCE_EVAL
METHOD Quickstep
&DFT
BASIS_SET_FILE_NAME BASIS_MOLOPT
POTENTIAL_FILE_NAME GTH_POTENTIALS
&LS_SCF
EPS_SCF 1.0E-8
EXTRAPOLATION_ORDER 1
MAX_SCF 25
PURIFICATION_METHOD TRS4
REPORT_ALL_SPARSITIES OFF
S_PRECONDITIONER NONE
&PAO
MAX_PAO 0
PARAMETERIZATION EQUIVARIANT
&END PAO
&END LS_SCF
&MGRID
CUTOFF 100
REL_CUTOFF 40
&END MGRID
&QS
LS_SCF
&END QS
&XC
&XC_FUNCTIONAL PBE
&END XC_FUNCTIONAL
&END XC
&END DFT
&SUBSYS
&CELL
ABC 8.0 8.0 8.0
&END CELL
! From 2H2O_rotations/phi_18/coords.xyz
&COORD
O -2.89210201 -3.94312785 4.35000000
H -2.44914329 -4.51011689 3.71000000
H -2.47994647 -3.07859282 4.25600000
O -5.79767781 -3.72981531 3.63500000
H -6.35104168 -4.02138593 4.36500000
H -4.88507656 -3.83102770 3.94700000
&END COORD
&KIND H
BASIS_SET DZVP-MOLOPT-GTH
PAO_BASIS_SIZE 4
PAO_MODEL_FILE DZVP-MOLOPT-GTH-PAO4-H.pt
POTENTIAL GTH-PBE
&END KIND
&KIND O
BASIS_SET DZVP-MOLOPT-GTH
PAO_BASIS_SIZE 4
PAO_MODEL_FILE DZVP-MOLOPT-GTH-PAO4-O.pt
POTENTIAL GTH-PBE
&END KIND
&TOPOLOGY
&CENTER_COORDINATES
&END CENTER_COORDINATES
&END TOPOLOGY
&END SUBSYS
&END FORCE_EVAL

View file

@ -1,3 +1,4 @@
2H2O_equivar_ML.inp 11 1e-9 -34.467980282273217
2H2O_equivar_ML.inp 11 1e-9 -34.467985160768627
2H2O_equivar_ML_checkforces.inp 0
#EOF