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 Equivariant parametrization
10 : !> \author Ole Schuett
11 : ! **************************************************************************************************
12 : MODULE pao_param_equi
13 : USE basis_set_types, ONLY: gto_basis_set_type
14 : USE cp_dbcsr_api, ONLY: &
15 : dbcsr_complete_redistribute, dbcsr_create, dbcsr_distribution_type, dbcsr_get_block_p, &
16 : dbcsr_get_info, dbcsr_iterator_blocks_left, dbcsr_iterator_next_block, &
17 : dbcsr_iterator_start, dbcsr_iterator_stop, dbcsr_iterator_type, dbcsr_p_type, &
18 : dbcsr_release, dbcsr_type
19 : USE cp_dbcsr_contrib, ONLY: dbcsr_reserve_diag_blocks
20 : USE dm_ls_scf_types, ONLY: ls_mstruct_type,&
21 : ls_scf_env_type
22 : USE kinds, ONLY: dp
23 : USE mathlib, ONLY: diamat_all
24 : USE message_passing, ONLY: mp_comm_type
25 : USE pao_param_methods, ONLY: pao_calc_grad_lnv_wrt_AB
26 : USE pao_potentials, ONLY: pao_guess_initial_potential
27 : USE pao_types, ONLY: pao_env_type
28 : USE qs_environment_types, ONLY: get_qs_env,&
29 : qs_environment_type
30 : USE qs_kind_types, ONLY: get_qs_kind,&
31 : qs_kind_type
32 : #include "./base/base_uses.f90"
33 :
34 : IMPLICIT NONE
35 :
36 : PRIVATE
37 :
38 : CHARACTER(len=*), PARAMETER, PRIVATE :: moduleN = 'pao_param_equi'
39 :
40 : PUBLIC :: pao_param_init_equi, pao_param_finalize_equi, pao_calc_AB_equi
41 : PUBLIC :: pao_param_count_equi, pao_param_initguess_equi
42 :
43 : CONTAINS
44 :
45 : ! **************************************************************************************************
46 : !> \brief Initialize equivariant parametrization
47 : !> \param pao ...
48 : ! **************************************************************************************************
49 26 : SUBROUTINE pao_param_init_equi(pao)
50 : TYPE(pao_env_type), POINTER :: pao
51 :
52 26 : IF (pao%precondition) THEN
53 0 : CPABORT("PAO preconditioning not supported for selected parametrization.")
54 : END IF
55 :
56 26 : END SUBROUTINE pao_param_init_equi
57 :
58 : ! **************************************************************************************************
59 : !> \brief Finalize equivariant parametrization
60 : ! **************************************************************************************************
61 26 : SUBROUTINE pao_param_finalize_equi()
62 :
63 : ! Nothing to do.
64 :
65 26 : END SUBROUTINE pao_param_finalize_equi
66 :
67 : ! **************************************************************************************************
68 : !> \brief Returns the number of parameters for given atomic kind
69 : !> \param qs_env ...
70 : !> \param ikind ...
71 : !> \param nparams ...
72 : ! **************************************************************************************************
73 112 : SUBROUTINE pao_param_count_equi(qs_env, ikind, nparams)
74 : TYPE(qs_environment_type), POINTER :: qs_env
75 : INTEGER, INTENT(IN) :: ikind
76 : INTEGER, INTENT(OUT) :: nparams
77 :
78 : INTEGER :: pao_basis_size, pri_basis_size
79 : TYPE(gto_basis_set_type), POINTER :: basis_set
80 56 : TYPE(qs_kind_type), DIMENSION(:), POINTER :: qs_kind_set
81 :
82 56 : CALL get_qs_env(qs_env, qs_kind_set=qs_kind_set)
83 : CALL get_qs_kind(qs_kind_set(ikind), &
84 : basis_set=basis_set, &
85 56 : pao_basis_size=pao_basis_size)
86 56 : pri_basis_size = basis_set%nsgf
87 :
88 56 : nparams = pao_basis_size*pri_basis_size
89 :
90 56 : END SUBROUTINE pao_param_count_equi
91 :
92 : ! **************************************************************************************************
93 : !> \brief Fills matrix_X with an initial guess
94 : !> \param pao ...
95 : !> \param qs_env ...
96 : ! **************************************************************************************************
97 10 : SUBROUTINE pao_param_initguess_equi(pao, qs_env)
98 : TYPE(pao_env_type), POINTER :: pao
99 : TYPE(qs_environment_type), POINTER :: qs_env
100 :
101 : CHARACTER(len=*), PARAMETER :: routineN = 'pao_param_initguess_equi'
102 :
103 : INTEGER :: acol, arow, handle, i, iatom, m, n
104 10 : INTEGER, DIMENSION(:), POINTER :: blk_sizes_pao, blk_sizes_pri
105 : LOGICAL :: found
106 10 : REAL(dp), DIMENSION(:), POINTER :: H_evals
107 10 : REAL(dp), DIMENSION(:, :), POINTER :: A, block_H0, block_N, block_N_inv, &
108 10 : block_X, H, H_evecs, V0
109 : TYPE(dbcsr_iterator_type) :: iter
110 :
111 10 : CALL timeset(routineN, handle)
112 :
113 10 : CALL dbcsr_get_info(pao%matrix_Y, row_blk_size=blk_sizes_pri, col_blk_size=blk_sizes_pao)
114 :
115 : !$OMP PARALLEL DEFAULT(NONE) SHARED(pao,qs_env,blk_sizes_pri,blk_sizes_pao) &
116 : !$OMP PRIVATE(iter,arow,acol,iatom,n,m,i,found) &
117 10 : !$OMP PRIVATE(block_X,block_H0,block_N,block_N_inv,A,H,H_evecs,H_evals,V0)
118 : CALL dbcsr_iterator_start(iter, pao%matrix_X)
119 : DO WHILE (dbcsr_iterator_blocks_left(iter))
120 : CALL dbcsr_iterator_next_block(iter, arow, acol, block_X)
121 : iatom = arow; CPASSERT(arow == acol)
122 :
123 : CALL dbcsr_get_block_p(matrix=pao%matrix_H0, row=iatom, col=iatom, block=block_H0, found=found)
124 : CALL dbcsr_get_block_p(matrix=pao%matrix_N_diag, row=iatom, col=iatom, block=block_N, found=found)
125 : CALL dbcsr_get_block_p(matrix=pao%matrix_N_inv_diag, row=iatom, col=iatom, block=block_N_inv, found=found)
126 : CPASSERT(ASSOCIATED(block_H0) .AND. ASSOCIATED(block_N) .AND. ASSOCIATED(block_N_inv))
127 :
128 : n = blk_sizes_pri(iatom) ! size of primary basis
129 : m = blk_sizes_pao(iatom) ! size of pao basis
130 :
131 : ALLOCATE (V0(n, n))
132 : CALL pao_guess_initial_potential(qs_env, iatom, V0)
133 :
134 : ! construct H
135 : ALLOCATE (H(n, n))
136 : H = MATMUL(MATMUL(block_N, block_H0 + V0), block_N) ! transform into orthonormal basis
137 :
138 : ! diagonalize H
139 : ALLOCATE (H_evecs(n, n), H_evals(n))
140 : H_evecs = H
141 : CALL diamat_all(H_evecs, H_evals)
142 :
143 : ! use first m eigenvectors as initial guess
144 : ALLOCATE (A(n, m))
145 : A = MATMUL(block_N_inv, H_evecs(:, 1:m))
146 :
147 : ! normalize vectors
148 : DO i = 1, m
149 : A(:, i) = A(:, i)/NORM2(A(:, i))
150 : END DO
151 :
152 : block_X = RESHAPE(A, [n*m, 1])
153 : DEALLOCATE (H, V0, A, H_evecs, H_evals)
154 :
155 : END DO
156 : CALL dbcsr_iterator_stop(iter)
157 : !$OMP END PARALLEL
158 :
159 10 : CALL timestop(handle)
160 :
161 10 : END SUBROUTINE pao_param_initguess_equi
162 :
163 : ! **************************************************************************************************
164 : !> \brief Takes current matrix_X and calculates the matrices A and B.
165 : !> \param pao ...
166 : !> \param qs_env ...
167 : !> \param ls_scf_env ...
168 : !> \param gradient ...
169 : !> \param penalty ...
170 : ! **************************************************************************************************
171 3412 : SUBROUTINE pao_calc_AB_equi(pao, qs_env, ls_scf_env, gradient, penalty)
172 : TYPE(pao_env_type), POINTER :: pao
173 : TYPE(qs_environment_type), POINTER :: qs_env
174 : TYPE(ls_scf_env_type), TARGET :: ls_scf_env
175 : LOGICAL, INTENT(IN) :: gradient
176 : REAL(dp), INTENT(INOUT), OPTIONAL :: penalty
177 :
178 : CHARACTER(len=*), PARAMETER :: routineN = 'pao_calc_AB_equi'
179 :
180 : INTEGER :: acol, arow, handle, i, iatom, j, k, m, n
181 : LOGICAL :: found
182 : REAL(dp) :: denom, penalty_sum, w
183 1706 : REAL(dp), DIMENSION(:), POINTER :: ANNA_evals
184 1706 : REAL(dp), DIMENSION(:, :), POINTER :: ANNA, ANNA_evecs, ANNA_inv, block_A, &
185 1706 : block_B, block_G, block_Ma, block_Mb, &
186 1706 : block_N, block_X, D, G, M1, M2, M3, &
187 1706 : M4, M5, NN
188 : TYPE(dbcsr_distribution_type) :: main_dist
189 : TYPE(dbcsr_iterator_type) :: iter
190 1706 : TYPE(dbcsr_p_type), DIMENSION(:), POINTER :: matrix_s
191 : TYPE(dbcsr_type) :: matrix_G_nondiag, matrix_Ma, matrix_Mb, &
192 : matrix_X_nondiag
193 : TYPE(ls_mstruct_type), POINTER :: ls_mstruct
194 : TYPE(mp_comm_type) :: group
195 :
196 1706 : CALL timeset(routineN, handle)
197 1706 : ls_mstruct => ls_scf_env%ls_mstruct
198 :
199 1706 : IF (gradient) THEN
200 234 : CALL pao_calc_grad_lnv_wrt_AB(qs_env, ls_scf_env, matrix_Ma, matrix_Mb)
201 : END IF
202 :
203 : ! Redistribute matrix_X from diag_distribution to distribution of matrix_s.
204 1706 : CALL get_qs_env(qs_env, matrix_s=matrix_s)
205 1706 : CALL dbcsr_get_info(matrix=matrix_s(1)%matrix, distribution=main_dist)
206 : CALL dbcsr_create(matrix_X_nondiag, &
207 : name="PAO matrix_X_nondiag", &
208 : dist=main_dist, &
209 1706 : template=pao%matrix_X)
210 1706 : CALL dbcsr_reserve_diag_blocks(matrix_X_nondiag)
211 1706 : CALL dbcsr_complete_redistribute(pao%matrix_X, matrix_X_nondiag)
212 :
213 : ! Compuation of matrix_G uses distr. of matrix_s, afterwards we redistribute to diag_distribution.
214 1706 : IF (gradient) THEN
215 : CALL dbcsr_create(matrix_G_nondiag, &
216 : name="PAO matrix_G_nondiag", &
217 : dist=main_dist, &
218 234 : template=pao%matrix_G)
219 234 : CALL dbcsr_reserve_diag_blocks(matrix_G_nondiag)
220 : END IF
221 :
222 : penalty_sum = 0.0_dp
223 :
224 : !$OMP PARALLEL DEFAULT(NONE) &
225 : !$OMP SHARED(pao,ls_mstruct,matrix_X_nondiag,matrix_G_nondiag,matrix_Ma,matrix_Mb,gradient,penalty) &
226 : !$OMP PRIVATE(iter,arow,acol,iatom,found,n,m,w,i,j,k,denom) &
227 : !$OMP PRIVATE(NN,ANNA,ANNA_evals,ANNA_evecs,ANNA_inv,D,G,M1,M2,M3,M4,M5) &
228 : !$OMP PRIVATE(block_X,block_A,block_B,block_N,block_Ma, block_Mb, block_G) &
229 1706 : !$OMP REDUCTION(+:penalty_sum)
230 : CALL dbcsr_iterator_start(iter, matrix_X_nondiag)
231 : DO WHILE (dbcsr_iterator_blocks_left(iter))
232 : CALL dbcsr_iterator_next_block(iter, arow, acol, block_X)
233 : iatom = arow; CPASSERT(arow == acol)
234 : CALL dbcsr_get_block_p(matrix=ls_mstruct%matrix_A, row=iatom, col=iatom, block=block_A, found=found)
235 : CPASSERT(ASSOCIATED(block_A))
236 : CALL dbcsr_get_block_p(matrix=ls_mstruct%matrix_B, row=iatom, col=iatom, block=block_B, found=found)
237 : CPASSERT(ASSOCIATED(block_B))
238 : CALL dbcsr_get_block_p(matrix=pao%matrix_N, row=iatom, col=iatom, block=block_N, found=found)
239 : CPASSERT(ASSOCIATED(block_N))
240 :
241 : n = SIZE(block_A, 1) ! size of primary basis
242 : m = SIZE(block_A, 2) ! size of pao basis
243 : block_A = RESHAPE(block_X, [n, m])
244 :
245 : ! restrain pao basis vectors to unit norm
246 : IF (PRESENT(penalty)) THEN
247 : DO i = 1, m
248 : w = 1.0_dp - SUM(block_A(:, i)**2)
249 : penalty_sum = penalty_sum + pao%penalty_strength*w**2
250 : END DO
251 : END IF
252 :
253 : ALLOCATE (NN(n, n), ANNA(m, m))
254 : NN = MATMUL(block_N, block_N) ! it's actually S^{-1}
255 : ANNA = MATMUL(MATMUL(TRANSPOSE(block_A), NN), block_A)
256 :
257 : ! diagonalize ANNA
258 : ALLOCATE (ANNA_evecs(m, m), ANNA_evals(m))
259 : ANNA_evecs(:, :) = ANNA
260 : CALL diamat_all(ANNA_evecs, ANNA_evals)
261 : IF (MINVAL(ABS(ANNA_evals)) < 1e-10_dp) CPABORT("PAO basis singualar.")
262 :
263 : ! build ANNA_inv
264 : ALLOCATE (ANNA_inv(m, m))
265 : ANNA_inv(:, :) = 0.0_dp
266 : DO k = 1, m
267 : w = 1.0_dp/ANNA_evals(k)
268 : DO i = 1, m
269 : DO j = 1, m
270 : ANNA_inv(i, j) = ANNA_inv(i, j) + w*ANNA_evecs(i, k)*ANNA_evecs(j, k)
271 : END DO
272 : END DO
273 : END DO
274 :
275 : !B = 1/S * A * 1/(A^T 1/S A)
276 : block_B = MATMUL(MATMUL(NN, block_A), ANNA_inv)
277 :
278 : ! TURNING POINT (if calc grad) ------------------------------------------
279 : IF (gradient) THEN
280 : CALL dbcsr_get_block_p(matrix=matrix_G_nondiag, row=iatom, col=iatom, block=block_G, found=found)
281 : CPASSERT(ASSOCIATED(block_G))
282 : CALL dbcsr_get_block_p(matrix=matrix_Ma, row=iatom, col=iatom, block=block_Ma, found=found)
283 : CALL dbcsr_get_block_p(matrix=matrix_Mb, row=iatom, col=iatom, block=block_Mb, found=found)
284 : ! don't check ASSOCIATED(block_M), it might have been filtered out.
285 :
286 : ALLOCATE (G(n, m))
287 : G(:, :) = 0.0_dp
288 :
289 : IF (PRESENT(penalty)) THEN
290 : DO i = 1, m
291 : w = 1.0_dp - SUM(block_A(:, i)**2)
292 : G(:, i) = -4.0_dp*pao%penalty_strength*w*block_A(:, i)
293 : END DO
294 : END IF
295 :
296 : IF (ASSOCIATED(block_Ma)) THEN
297 : G = G + block_Ma
298 : END IF
299 :
300 : IF (ASSOCIATED(block_Mb)) THEN
301 : G = G + MATMUL(MATMUL(NN, block_Mb), ANNA_inv)
302 :
303 : ! calculate derivatives dAA_inv/ dAA
304 : ALLOCATE (D(m, m), M1(m, m), M2(m, m), M3(m, m), M4(m, m), M5(m, m))
305 :
306 : DO i = 1, m
307 : DO j = 1, m
308 : denom = ANNA_evals(i) - ANNA_evals(j)
309 : IF (i == j) THEN
310 : D(i, i) = -1.0_dp/ANNA_evals(i)**2 ! diagonal elements
311 : ELSE IF (ABS(denom) > 1e-10_dp) THEN
312 : D(i, j) = (1.0_dp/ANNA_evals(i) - 1.0_dp/ANNA_evals(j))/denom
313 : ELSE
314 : D(i, j) = -1.0_dp ! limit according to L'Hospital's rule
315 : END IF
316 : END DO
317 : END DO
318 :
319 : M1 = MATMUL(MATMUL(TRANSPOSE(block_A), NN), block_Mb)
320 : M2 = MATMUL(MATMUL(TRANSPOSE(ANNA_evecs), M1), ANNA_evecs)
321 : M3 = M2*D ! Hadamard product
322 : M4 = MATMUL(MATMUL(ANNA_evecs, M3), TRANSPOSE(ANNA_evecs))
323 : M5 = 0.5_dp*(M4 + TRANSPOSE(M4))
324 : G = G + 2.0_dp*MATMUL(MATMUL(NN, block_A), M5)
325 :
326 : DEALLOCATE (D, M1, M2, M3, M4, M5)
327 : END IF
328 :
329 : block_G = RESHAPE(G, [n*m, 1])
330 : DEALLOCATE (G)
331 : END IF
332 :
333 : DEALLOCATE (NN, ANNA, ANNA_evecs, ANNA_evals, ANNA_inv)
334 : END DO
335 : CALL dbcsr_iterator_stop(iter)
336 : !$OMP END PARALLEL
337 :
338 : ! sum penalty energies across ranks
339 1706 : IF (PRESENT(penalty)) THEN
340 1678 : CALL dbcsr_get_info(pao%matrix_X, group=group)
341 1678 : CALL group%sum(penalty_sum)
342 1678 : penalty = penalty_sum
343 : END IF
344 :
345 1706 : CALL dbcsr_release(matrix_X_nondiag)
346 :
347 1706 : IF (gradient) THEN
348 234 : CALL dbcsr_complete_redistribute(matrix_G_nondiag, pao%matrix_G)
349 234 : CALL dbcsr_release(matrix_G_nondiag)
350 234 : CALL dbcsr_release(matrix_Ma)
351 234 : CALL dbcsr_release(matrix_Mb)
352 : END IF
353 :
354 1706 : CALL timestop(handle)
355 :
356 1706 : END SUBROUTINE pao_calc_AB_equi
357 :
358 : END MODULE pao_param_equi
|