mirror of
https://github.com/cp2k/cp2k.git
synced 2026-07-27 13:45:19 -04:00
PAO: Add forces for equivariant PyTorch models
This commit is contained in:
parent
b6963b94bc
commit
5912d788db
6 changed files with 184 additions and 35 deletions
|
|
@ -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
|
||||
|
|
|
|||
132
src/pao_model.F
132
src/pao_model.F
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
72
tests/QS/regtest-pao-5/2H2O_equivar_ML_checkforces.inp
Normal file
72
tests/QS/regtest-pao-5/2H2O_equivar_ML_checkforces.inp
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue