From 5912d788db39b623e4c07ba720ff657d0b0ebfc3 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Ole=20Sch=C3=BCtt?= Date: Sat, 25 Jan 2025 15:32:33 +0100 Subject: [PATCH] PAO: Add forces for equivariant PyTorch models --- src/pao_methods.F | 5 +- src/pao_model.F | 132 ++++++++++++++---- src/pao_types.F | 4 + tests/QS/regtest-pao-5/2H2O_equivar_ML.inp | 3 +- .../2H2O_equivar_ML_checkforces.inp | 72 ++++++++++ tests/QS/regtest-pao-5/TEST_FILES | 3 +- 6 files changed, 184 insertions(+), 35 deletions(-) create mode 100644 tests/QS/regtest-pao-5/2H2O_equivar_ML_checkforces.inp diff --git a/src/pao_methods.F b/src/pao_methods.F index ba5600f19d..d1fc6d2247 100644 --- a/src/pao_methods.F +++ b/src/pao_methods.F @@ -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 diff --git a/src/pao_model.F b/src/pao_model.F index e3102297ba..e05443b91b 100644 --- a/src/pao_model.F +++ b/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 diff --git a/src/pao_types.F b/src/pao_types.F index 73f84a8c0a..72ca92cc73 100644 --- a/src/pao_types.F +++ b/src/pao_types.F @@ -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 diff --git a/tests/QS/regtest-pao-5/2H2O_equivar_ML.inp b/tests/QS/regtest-pao-5/2H2O_equivar_ML.inp index 659dbf105a..b0c79eccd4 100644 --- a/tests/QS/regtest-pao-5/2H2O_equivar_ML.inp +++ b/tests/QS/regtest-pao-5/2H2O_equivar_ML.inp @@ -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 diff --git a/tests/QS/regtest-pao-5/2H2O_equivar_ML_checkforces.inp b/tests/QS/regtest-pao-5/2H2O_equivar_ML_checkforces.inp new file mode 100644 index 0000000000..dfaf6ac255 --- /dev/null +++ b/tests/QS/regtest-pao-5/2H2O_equivar_ML_checkforces.inp @@ -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 diff --git a/tests/QS/regtest-pao-5/TEST_FILES b/tests/QS/regtest-pao-5/TEST_FILES index ec66d8d9d6..e9aadfbc5d 100644 --- a/tests/QS/regtest-pao-5/TEST_FILES +++ b/tests/QS/regtest-pao-5/TEST_FILES @@ -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