Neko 1.99.9
A portable framework for high-order spectral element flow simulations
Loading...
Searching...
No Matches
fusedcg_cpld_device.F90
Go to the documentation of this file.
1! Copyright (c) 2021-2026, The Neko Authors
2! All rights reserved.
3!
4! Redistribution and use in source and binary forms, with or without
5! modification, are permitted provided that the following conditions
6! are met:
7!
8! * Redistributions of source code must retain the above copyright
9! notice, this list of conditions and the following disclaimer.
10!
11! * Redistributions in binary form must reproduce the above
12! copyright notice, this list of conditions and the following
13! disclaimer in the documentation and/or other materials provided
14! with the distribution.
15!
16! * Neither the name of the authors nor the names of its
17! contributors may be used to endorse or promote products derived
18! from this software without specific prior written permission.
19!
20! THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS
21! "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT
22! LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS
23! FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE
24! COPYRIGHT OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT,
25! INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING,
26! BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES;
27! LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
28! CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT
29! LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN
30! ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE
31! POSSIBILITY OF SUCH DAMAGE.
32!
36 use precon, only : pc_t
37 use ax_product, only : ax_t
38 use num_types, only : rp, c_rp
39 use field, only : field_t
40 use coefs, only : coef_t
41 use gather_scatter, only : gs_t, gs_op_add
45 use math, only : glsc3, rzero, copy, abscmp
50 use utils, only : neko_error
52 use mpi_f08, only : mpi_in_place, mpi_allreduce, &
53 mpi_sum
54 use operators, only : rotate_cyc
55 use, intrinsic :: iso_c_binding, only : c_ptr, c_null_ptr, &
56 c_associated, c_size_t, c_sizeof, c_int, c_loc
57 implicit none
58 private
59
60 integer, parameter :: device_fusedcg_cpld_p_space = 10
61
63 type, public, extends(ksp_t) :: fusedcg_cpld_device_t
64 real(kind=rp), allocatable :: w1(:)
65 real(kind=rp), allocatable :: w2(:)
66 real(kind=rp), allocatable :: w3(:)
67 real(kind=rp), allocatable :: r1(:)
68 real(kind=rp), allocatable :: r2(:)
69 real(kind=rp), allocatable :: r3(:)
70 real(kind=rp), allocatable :: z1(:)
71 real(kind=rp), allocatable :: z2(:)
72 real(kind=rp), allocatable :: z3(:)
73 real(kind=rp), allocatable :: tmp(:)
74 real(kind=rp), allocatable :: p1(:,:)
75 real(kind=rp), allocatable :: p2(:,:)
76 real(kind=rp), allocatable :: p3(:,:)
77 real(kind=rp), allocatable :: alpha(:)
78 type(c_ptr) :: w1_d = c_null_ptr
79 type(c_ptr) :: w2_d = c_null_ptr
80 type(c_ptr) :: w3_d = c_null_ptr
81 type(c_ptr) :: r1_d = c_null_ptr
82 type(c_ptr) :: r2_d = c_null_ptr
83 type(c_ptr) :: r3_d = c_null_ptr
84 type(c_ptr) :: z1_d = c_null_ptr
85 type(c_ptr) :: z2_d = c_null_ptr
86 type(c_ptr) :: z3_d = c_null_ptr
87 type(c_ptr) :: alpha_d = c_null_ptr
88 type(c_ptr) :: p1_d_d = c_null_ptr
89 type(c_ptr) :: p2_d_d = c_null_ptr
90 type(c_ptr) :: p3_d_d = c_null_ptr
91 type(c_ptr) :: tmp_d = c_null_ptr
92 type(c_ptr), allocatable :: p1_d(:)
93 type(c_ptr), allocatable :: p2_d(:)
94 type(c_ptr), allocatable :: p3_d(:)
95 type(c_ptr) :: gs_event1 = c_null_ptr
96 type(c_ptr) :: gs_event2 = c_null_ptr
97 type(c_ptr) :: gs_event3 = c_null_ptr
98 contains
99 procedure, pass(this) :: init => fusedcg_cpld_device_init
100 procedure, pass(this) :: free => fusedcg_cpld_device_free
101 procedure, pass(this) :: solve => fusedcg_cpld_device_solve
102 procedure, pass(this) :: solve_coupled => fusedcg_cpld_device_solve_coupled
103 end type fusedcg_cpld_device_t
104
105#ifdef HAVE_CUDA
106 interface
107 subroutine cuda_fusedcg_cpld_part1(a1_d, a2_d, a3_d, &
108 b1_d, b2_d, b3_d, tmp_d, n) bind(c, name = 'cuda_fusedcg_cpld_part1')
109 use, intrinsic :: iso_c_binding
110 import c_rp
111 implicit none
112 type(c_ptr), value :: a1_d, a2_d, a3_d, b1_d, b2_d, b3_d, tmp_d
113 integer(c_int) :: n
114 end subroutine cuda_fusedcg_cpld_part1
115 end interface
116
117 interface
118 subroutine cuda_fusedcg_cpld_update_p(p1_d, p2_d, p3_d, z1_d, z2_d, z3_d, &
119 po1_d, po2_d, po3_d, beta, n) &
120 bind(c, name = 'cuda_fusedcg_cpld_update_p')
121 use, intrinsic :: iso_c_binding
122 import c_rp
123 implicit none
124 type(c_ptr), value :: p1_d, p2_d, p3_d, z1_d, z2_d, z3_d
125 type(c_ptr), value :: po1_d, po2_d, po3_d
126 real(c_rp) :: beta
127 integer(c_int) :: n
128 end subroutine cuda_fusedcg_cpld_update_p
129 end interface
130
131 interface
132 subroutine cuda_fusedcg_cpld_update_x(x1_d, x2_d, x3_d, p1_d, p2_d, p3_d, &
133 alpha, p_cur, n) bind(c, name = 'cuda_fusedcg_cpld_update_x')
134 use, intrinsic :: iso_c_binding
135 implicit none
136 type(c_ptr), value :: x1_d, x2_d, x3_d, p1_d, p2_d, p3_d, alpha
137 integer(c_int) :: p_cur, n
138 end subroutine cuda_fusedcg_cpld_update_x
139 end interface
140
141 interface
142 real(c_rp) function cuda_fusedcg_cpld_part2(a1_d, a2_d, a3_d, b_d, &
143 c1_d, c2_d, c3_d, alpha_d, alpha, p_cur, n) &
144 bind(c, name = 'cuda_fusedcg_cpld_part2')
145 use, intrinsic :: iso_c_binding
146 import c_rp
147 implicit none
148 type(c_ptr), value :: a1_d, a2_d, a3_d, b_d
149 type(c_ptr), value :: c1_d, c2_d, c3_d, alpha_d
150 real(c_rp) :: alpha
151 integer(c_int) :: n, p_cur
152 end function cuda_fusedcg_cpld_part2
153 end interface
154#elif HAVE_HIP
155 interface
156 subroutine hip_fusedcg_cpld_part1(a1_d, a2_d, a3_d, &
157 b1_d, b2_d, b3_d, tmp_d, n) &
158 bind(c, name = 'hip_fusedcg_cpld_part1')
159 use, intrinsic :: iso_c_binding
160 import c_rp
161 implicit none
162 type(c_ptr), value :: a1_d, a2_d, a3_d, b1_d, b2_d, b3_d, tmp_d
163 integer(c_int) :: n
164 end subroutine hip_fusedcg_cpld_part1
165 end interface
166
167 interface
168 subroutine hip_fusedcg_cpld_update_p(p1_d, p2_d, p3_d, z1_d, z2_d, z3_d, &
169 po1_d, po2_d, po3_d, beta, n) &
170 bind(c, name = 'hip_fusedcg_cpld_update_p')
171 use, intrinsic :: iso_c_binding
172 import c_rp
173 implicit none
174 type(c_ptr), value :: p1_d, p2_d, p3_d, z1_d, z2_d, z3_d
175 type(c_ptr), value :: po1_d, po2_d, po3_d
176 real(c_rp) :: beta
177 integer(c_int) :: n
178 end subroutine hip_fusedcg_cpld_update_p
179 end interface
180
181 interface
182 subroutine hip_fusedcg_cpld_update_x(x1_d, x2_d, x3_d, p1_d, p2_d, p3_d, &
183 alpha, p_cur, n) bind(c, name = 'hip_fusedcg_cpld_update_x')
184 use, intrinsic :: iso_c_binding
185 implicit none
186 type(c_ptr), value :: x1_d, x2_d, x3_d, p1_d, p2_d, p3_d, alpha
187 integer(c_int) :: p_cur, n
188 end subroutine hip_fusedcg_cpld_update_x
189 end interface
190
191 interface
192 real(c_rp) function hip_fusedcg_cpld_part2(a1_d, a2_d, a3_d, b_d, &
193 c1_d, c2_d, c3_d, alpha_d, alpha, p_cur, n) &
194 bind(c, name = 'hip_fusedcg_cpld_part2')
195 use, intrinsic :: iso_c_binding
196 import c_rp
197 implicit none
198 type(c_ptr), value :: a1_d, a2_d, a3_d, b_d
199 type(c_ptr), value :: c1_d, c2_d, c3_d, alpha_d
200 real(c_rp) :: alpha
201 integer(c_int) :: n, p_cur
202 end function hip_fusedcg_cpld_part2
203 end interface
204#endif
205
206contains
207
208 subroutine device_fusedcg_cpld_part1(a1_d, a2_d, a3_d, &
209 b1_d, b2_d, b3_d, tmp_d, n)
210 type(c_ptr), value :: a1_d, a2_d, a3_d, b1_d, b2_d, b3_d
211 type(c_ptr), value :: tmp_d
212 integer(c_int) :: n
213#ifdef HAVE_HIP
214 call hip_fusedcg_cpld_part1(a1_d, a2_d, a3_d, b1_d, b2_d, b3_d, tmp_d, n)
215#elif HAVE_CUDA
216 call cuda_fusedcg_cpld_part1(a1_d, a2_d, a3_d, b1_d, b2_d, b3_d, tmp_d, n)
217#else
218 call neko_error('No device backend configured')
219#endif
220 end subroutine device_fusedcg_cpld_part1
221
222 subroutine device_fusedcg_cpld_update_p(p1_d, p2_d, p3_d, z1_d, z2_d, z3_d, &
223 po1_d, po2_d, po3_d, beta, n)
224 type(c_ptr), value :: p1_d, p2_d, p3_d, z1_d, z2_d, z3_d
225 type(c_ptr), value :: po1_d, po2_d, po3_d
226 real(c_rp) :: beta
227 integer(c_int) :: n
228#ifdef HAVE_HIP
229 call hip_fusedcg_cpld_update_p(p1_d, p2_d, p3_d, z1_d, z2_d, z3_d, &
230 po1_d, po2_d, po3_d, beta, n)
231#elif HAVE_CUDA
232 call cuda_fusedcg_cpld_update_p(p1_d, p2_d, p3_d, z1_d, z2_d, z3_d, &
233 po1_d, po2_d, po3_d, beta, n)
234#else
235 call neko_error('No device backend configured')
236#endif
237 end subroutine device_fusedcg_cpld_update_p
238
239 subroutine device_fusedcg_cpld_update_x(x1_d, x2_d, x3_d, &
240 p1_d, p2_d, p3_d, alpha, p_cur, n)
241 type(c_ptr), value :: x1_d, x2_d, x3_d, p1_d, p2_d, p3_d, alpha
242 integer(c_int) :: p_cur, n
243#ifdef HAVE_HIP
244 call hip_fusedcg_cpld_update_x(x1_d, x2_d, x3_d, &
245 p1_d, p2_d, p3_d, alpha, p_cur, n)
246#elif HAVE_CUDA
247 call cuda_fusedcg_cpld_update_x(x1_d, x2_d, x3_d, &
248 p1_d, p2_d, p3_d, alpha, p_cur, n)
249#else
250 call neko_error('No device backend configured')
251#endif
252 end subroutine device_fusedcg_cpld_update_x
253
254 function device_fusedcg_cpld_part2(a1_d, a2_d, a3_d, b_d, &
255 c1_d, c2_d, c3_d, alpha_d, alpha, p_cur, n) result(res)
256 type(c_ptr), value :: a1_d, a2_d, a3_d, b_d
257 type(c_ptr), value :: c1_d, c2_d, c3_d, alpha_d
258 real(c_rp) :: alpha
259 integer :: n, p_cur
260 real(kind=rp) :: res
261 integer :: ierr
262#ifdef HAVE_HIP
263 res = hip_fusedcg_cpld_part2(a1_d, a2_d, a3_d, b_d, &
264 c1_d, c2_d, c3_d, alpha_d, alpha, p_cur, n)
265#elif HAVE_CUDA
266 res = cuda_fusedcg_cpld_part2(a1_d, a2_d, a3_d, b_d, &
267 c1_d, c2_d, c3_d, alpha_d, alpha, p_cur, n)
268#else
269 call neko_error('No device backend configured')
270#endif
271
272#ifndef HAVE_DEVICE_MPI
273 if (pe_size .gt. 1) then
274 call mpi_allreduce(mpi_in_place, res, 1, &
275 mpi_real_precision, mpi_sum, neko_comm, ierr)
276 end if
277#endif
278
279 end function device_fusedcg_cpld_part2
280
282 subroutine fusedcg_cpld_device_init(this, n, max_iter, M, &
283 rel_tol, abs_tol, monitor)
284 class(fusedcg_cpld_device_t), target, intent(inout) :: this
285 class(pc_t), optional, intent(in), target :: M
286 integer, intent(in) :: n
287 integer, intent(in) :: max_iter
288 real(kind=rp), optional, intent(in) :: rel_tol
289 real(kind=rp), optional, intent(in) :: abs_tol
290 logical, optional, intent(in) :: monitor
291 type(c_ptr) :: ptr
292 integer(c_size_t) :: p_size
293 integer :: i
294
295 call this%free()
296
297 allocate(this%w1(n))
298 allocate(this%w2(n))
299 allocate(this%w3(n))
300 allocate(this%r1(n))
301 allocate(this%r2(n))
302 allocate(this%r3(n))
303 allocate(this%z1(n))
304 allocate(this%z2(n))
305 allocate(this%z3(n))
306 allocate(this%tmp(n))
307 allocate(this%p1(n, device_fusedcg_cpld_p_space))
308 allocate(this%p2(n, device_fusedcg_cpld_p_space))
309 allocate(this%p3(n, device_fusedcg_cpld_p_space))
310 allocate(this%p1_d(device_fusedcg_cpld_p_space))
311 allocate(this%p2_d(device_fusedcg_cpld_p_space))
312 allocate(this%p3_d(device_fusedcg_cpld_p_space))
313 allocate(this%alpha(device_fusedcg_cpld_p_space))
314
315 if (present(m)) then
316 this%M => m
317 end if
318
319 call device_map(this%w1, this%w1_d, n)
320 call device_map(this%w2, this%w2_d, n)
321 call device_map(this%w3, this%w3_d, n)
322 call device_map(this%r1, this%r1_d, n)
323 call device_map(this%r2, this%r2_d, n)
324 call device_map(this%r3, this%r3_d, n)
325 call device_map(this%z1, this%z1_d, n)
326 call device_map(this%z2, this%z2_d, n)
327 call device_map(this%z3, this%z3_d, n)
328 call device_map(this%tmp, this%tmp_d, n)
329 call device_map(this%alpha, this%alpha_d, device_fusedcg_cpld_p_space)
331 this%p1_d(i) = c_null_ptr
332 call device_map(this%p1(:,i), this%p1_d(i), n)
333
334 this%p2_d(i) = c_null_ptr
335 call device_map(this%p2(:,i), this%p2_d(i), n)
336
337 this%p3_d(i) = c_null_ptr
338 call device_map(this%p3(:,i), this%p3_d(i), n)
339 end do
340
341 p_size = c_sizeof(c_null_ptr) * (device_fusedcg_cpld_p_space)
342 call device_alloc(this%p1_d_d, p_size)
343 call device_alloc(this%p2_d_d, p_size)
344 call device_alloc(this%p3_d_d, p_size)
345 ptr = c_loc(this%p1_d)
346 call device_memcpy(ptr, this%p1_d_d, p_size, &
347 host_to_device, sync=.false.)
348 ptr = c_loc(this%p2_d)
349 call device_memcpy(ptr, this%p2_d_d, p_size, &
350 host_to_device, sync=.false.)
351 ptr = c_loc(this%p3_d)
352 call device_memcpy(ptr, this%p3_d_d, p_size, &
353 host_to_device, sync=.false.)
354 if (present(rel_tol) .and. present(abs_tol) .and. present(monitor)) then
355 call this%ksp_init(max_iter, rel_tol, abs_tol, monitor = monitor)
356 else if (present(rel_tol) .and. present(abs_tol)) then
357 call this%ksp_init(max_iter, rel_tol, abs_tol)
358 else if (present(monitor) .and. present(abs_tol)) then
359 call this%ksp_init(max_iter, abs_tol = abs_tol, monitor = monitor)
360 else if (present(rel_tol) .and. present(monitor)) then
361 call this%ksp_init(max_iter, rel_tol, monitor = monitor)
362 else if (present(rel_tol)) then
363 call this%ksp_init(max_iter, rel_tol = rel_tol)
364 else if (present(abs_tol)) then
365 call this%ksp_init(max_iter, abs_tol = abs_tol)
366 else if (present(monitor)) then
367 call this%ksp_init(max_iter, monitor = monitor)
368 else
369 call this%ksp_init(max_iter)
370 end if
371
372 call device_event_create(this%gs_event1, 2)
373 call device_event_create(this%gs_event2, 2)
374 call device_event_create(this%gs_event3, 2)
375
376 end subroutine fusedcg_cpld_device_init
377
380 class(fusedcg_cpld_device_t), intent(inout) :: this
381 integer :: i
382
383 call this%ksp_free()
384
385 if (allocated(this%w1)) then
386 if (c_associated(this%w1_d)) then
387 call device_unmap(this%w1, this%w1_d)
388 end if
389 deallocate(this%w1)
390 end if
391
392 if (allocated(this%w2)) then
393 if (c_associated(this%w2_d)) then
394 call device_unmap(this%w2, this%w2_d)
395 end if
396 deallocate(this%w2)
397 end if
398
399 if (allocated(this%w3)) then
400 if (c_associated(this%w3_d)) then
401 call device_unmap(this%w3, this%w3_d)
402 end if
403 deallocate(this%w3)
404 end if
405
406 if (allocated(this%r1)) then
407 if (c_associated(this%r1_d)) then
408 call device_unmap(this%r1, this%r1_d)
409 end if
410 deallocate(this%r1)
411 end if
412
413 if (allocated(this%r2)) then
414 if (c_associated(this%r2_d)) then
415 call device_unmap(this%r2, this%r2_d)
416 end if
417 deallocate(this%r2)
418 end if
419
420 if (allocated(this%r3)) then
421 if (c_associated(this%r3_d)) then
422 call device_unmap(this%r3, this%r3_d)
423 end if
424 deallocate(this%r3)
425 end if
426
427 if (allocated(this%z1)) then
428 if (c_associated(this%z1_d)) then
429 call device_unmap(this%z1, this%z1_d)
430 end if
431 deallocate(this%z1)
432 end if
433
434 if (allocated(this%z2)) then
435 if (c_associated(this%z2_d)) then
436 call device_unmap(this%z2, this%z2_d)
437 end if
438 deallocate(this%z2)
439 end if
440
441 if (allocated(this%z3)) then
442 if (c_associated(this%z3_d)) then
443 call device_unmap(this%z3, this%z3_d)
444 end if
445 deallocate(this%z3)
446 end if
447
448 if (allocated(this%tmp)) then
449 if (c_associated(this%tmp_d)) then
450 call device_unmap(this%tmp, this%tmp_d)
451 end if
452 deallocate(this%tmp)
453 end if
454
455 if (allocated(this%alpha)) then
456 if (c_associated(this%alpha_d)) then
457 call device_unmap(this%alpha, this%alpha_d)
458 end if
459 deallocate(this%alpha)
460 end if
461
462 if (allocated(this%p1)) then
463 if (allocated(this%p1_d)) then
465 if (c_associated(this%p1_d(i))) then
466 call device_unmap(this%p1(:,i), this%p1_d(i))
467 end if
468 end do
469 end if
470 deallocate(this%p1)
471 end if
472
473 if (allocated(this%p2)) then
474 if (allocated(this%p2_d)) then
476 if (c_associated(this%p2_d(i))) then
477 call device_unmap(this%p2(:,i), this%p2_d(i))
478 end if
479 end do
480 end if
481 deallocate(this%p2)
482 end if
483
484 if (allocated(this%p3)) then
485 if (allocated(this%p3_d)) then
487 if (c_associated(this%p3_d(i))) then
488 call device_unmap(this%p3(:,i), this%p3_d(i))
489 end if
490 end do
491 end if
492 deallocate(this%p3)
493 end if
494
495 if (allocated(this%p1_d)) then
496 deallocate(this%p1_d)
497 end if
498
499 if (allocated(this%p2_d)) then
500 deallocate(this%p2_d)
501 end if
502
503 if (allocated(this%p3_d)) then
504 deallocate(this%p3_d)
505 end if
506
507 if (c_associated(this%p1_d_d)) then
508 call device_free(this%p1_d_d)
509 end if
510
511 if (c_associated(this%p2_d_d)) then
512 call device_free(this%p2_d_d)
513 end if
514
515 if (c_associated(this%p3_d_d)) then
516 call device_free(this%p3_d_d)
517 end if
518
519 nullify(this%M)
520
521 if (c_associated(this%gs_event1)) then
522 call device_event_destroy(this%gs_event1)
523 end if
524
525 if (c_associated(this%gs_event2)) then
526 call device_event_destroy(this%gs_event2)
527 end if
528
529 if (c_associated(this%gs_event3)) then
530 call device_event_destroy(this%gs_event3)
531 end if
532
533 end subroutine fusedcg_cpld_device_free
534
536 function fusedcg_cpld_device_solve_coupled(this, Ax, x, y, z, fx, fy, fz, &
537 n, coef, bc_projector, gs_h, niter) result(ksp_results)
538 class(fusedcg_cpld_device_t), intent(inout) :: this
539 class(ax_t), intent(in) :: ax
540 type(field_t), intent(inout) :: x
541 type(field_t), intent(inout) :: y
542 type(field_t), intent(inout) :: z
543 integer, intent(in) :: n
544 real(kind=rp), dimension(n), intent(in) :: fx
545 real(kind=rp), dimension(n), intent(in) :: fy
546 real(kind=rp), dimension(n), intent(in) :: fz
547 type(coef_t), intent(inout) :: coef
548 class(vector_bc_projector_t), intent(inout) :: bc_projector
549 type(gs_t), intent(inout) :: gs_h
550 type(ksp_monitor_t), dimension(3) :: ksp_results
551 integer, optional, intent(in) :: niter
552 integer :: iter, max_iter, ierr, i, p_cur, p_prev
553 real(kind=rp) :: rnorm, rtr, norm_fac, rtz1, rtz2
554 real(kind=rp) :: pap, beta
555 type(c_ptr) :: fx_d
556 type(c_ptr) :: fy_d
557 type(c_ptr) :: fz_d
558
559 fx_d = device_get_ptr(fx)
560 fy_d = device_get_ptr(fy)
561 fz_d = device_get_ptr(fz)
562
563 if (present(niter)) then
564 max_iter = niter
565 else
566 max_iter = ksp_max_iter
567 end if
568 norm_fac = 1.0_rp / sqrt(coef%volume)
569
570 associate(w1 => this%w1, w2 => this%w2, w3 => this%w3, r1 => this%r1, &
571 r2 => this%r2, r3 => this%r3, p1 => this%p1, p2 => this%p2, &
572 p3 => this%p3, z1 => this%z1, z2 => this%z2, z3 => this%z3, &
573 tmp_d => this%tmp_d, alpha => this%alpha, alpha_d => this%alpha_d, &
574 w1_d => this%w1_d, w2_d => this%w2_d, w3_d => this%w3_d, &
575 r1_d => this%r1_d, r2_d => this%r2_d, r3_d => this%r3_d, &
576 z1_d => this%z1_d, z2_d => this%z2_d, z3_d => this%z3_d, &
577 p1_d => this%p1_d, p2_d => this%p2_d, p3_d => this%p3_d, &
578 p1_d_d => this%p1_d_d, p2_d_d => this%p2_d_d, p3_d_d => this%p3_d_d)
579
580 rtz1 = 1.0_rp
582 p_cur = 1
583
584
585 call device_rzero(x%x_d, n)
586 call device_rzero(y%x_d, n)
587 call device_rzero(z%x_d, n)
588 call device_rzero(p1_d(1), n)
589 call device_rzero(p2_d(1), n)
590 call device_rzero(p3_d(1), n)
591 call device_copy(r1_d, fx_d, n)
592 call device_copy(r2_d, fy_d, n)
593 call device_copy(r3_d, fz_d, n)
594
595 call device_fusedcg_cpld_part1(r1_d, r2_d, r3_d, r1_d, &
596 r2_d, r3_d, tmp_d, n)
597
598 rtr = device_glsc2(tmp_d, coef%mult_d, n)
599
600 rnorm = sqrt(rtr)*norm_fac
601 ksp_results%res_start = rnorm
602 ksp_results%res_final = rnorm
603 ksp_results(1)%iter = 0
604 ksp_results(2:3)%iter = -1
605 if (abscmp(rnorm, 0.0_rp)) then
606 ksp_results%converged = .true.
607 return
608 end if
609 call this%monitor_start('fcpldCG')
610 do iter = 1, max_iter
611 call this%M%solve(z1, r1, n)
612 call this%M%solve(z2, r2, n)
613 call this%M%solve(z3, r3, n)
614 rtz2 = rtz1
615 call device_fusedcg_cpld_part1(z1_d, z2_d, z3_d, &
616 r1_d, r2_d, r3_d, tmp_d, n)
617 rtz1 = device_glsc2(tmp_d, coef%mult_d, n)
618
619 beta = rtz1 / rtz2
620 if (iter .eq. 1) beta = 0.0_rp
621
622 call device_fusedcg_cpld_update_p(p1_d(p_cur), p2_d(p_cur), p3_d(p_cur), &
623 z1_d, z2_d, z3_d, p1_d(p_prev), p2_d(p_prev), p3_d(p_prev), beta, n)
624
625 call ax%compute_vector(w1, w2, w3, &
626 p1(1, p_cur), p2(1, p_cur), p3(1, p_cur), coef, x%msh, x%Xh)
627
628 call rotate_cyc(w1_d, w2_d, w3_d, 1, coef)
629 call gs_h%op(w1, w2, w3, n, gs_op_add, this%gs_event1)
630 call device_event_sync(this%gs_event1)
631 call bc_projector%apply(w1, w2, w3, n)
632 call rotate_cyc(w1_d, w2_d, w3_d, 0, coef)
633
634 call device_fusedcg_cpld_part1(w1_d, w2_d, w3_d, p1_d(p_cur), &
635 p2_d(p_cur), p3_d(p_cur), tmp_d, n)
636
637 pap = device_glsc2(tmp_d, coef%mult_d, n)
638
639 alpha(p_cur) = rtz1 / pap
640 rtr = device_fusedcg_cpld_part2(r1_d, r2_d, r3_d, coef%mult_d, &
641 w1_d, w2_d, w3_d, alpha_d, alpha(p_cur), p_cur, n)
642 rnorm = sqrt(rtr)*norm_fac
643 call this%monitor_iter(iter, rnorm)
644 if ((p_cur .eq. device_fusedcg_cpld_p_space) .or. &
645 (rnorm .lt. this%abs_tol) .or. iter .eq. max_iter) then
646 call device_fusedcg_cpld_update_x(x%x_d, y%x_d, z%x_d, &
647 p1_d_d, p2_d_d, p3_d_d, alpha_d, p_cur, n)
648 p_prev = p_cur
649 p_cur = 1
650 if (rnorm .lt. this%abs_tol) exit
651 else
652 p_prev = p_cur
653 p_cur = p_cur + 1
654 end if
655 end do
656 call this%monitor_stop()
657 ksp_results%res_final = rnorm
658 ksp_results%iter = iter
659 ksp_results%converged = this%is_converged(iter, rnorm)
660
661 end associate
662
664
666 function fusedcg_cpld_device_solve(this, Ax, x, f, n, coef, bc_projector, &
667 gs_h, niter) result(ksp_results)
668 class(fusedcg_cpld_device_t), intent(inout) :: this
669 class(ax_t), intent(in) :: ax
670 type(field_t), intent(inout) :: x
671 integer, intent(in) :: n
672 real(kind=rp), dimension(n), intent(in) :: f
673 type(coef_t), intent(inout) :: coef
674 class(scalar_bc_projector_t), intent(inout) :: bc_projector
675 type(gs_t), intent(inout) :: gs_h
676 type(ksp_monitor_t) :: ksp_results
677 integer, optional, intent(in) :: niter
678
679 ! Throw and error
680 call neko_error('The cpldcg solver is only defined for coupled solves')
681
682 ksp_results%res_final = 0.0
683 ksp_results%iter = 0
684 ksp_results%converged = .false.
685
686 end function fusedcg_cpld_device_solve
687
688end module fusedcg_cpld_device
__device__ T solve(const T u, const T y, const T guess, const T nu, const T kappa, const T B)
void hip_fusedcg_cpld_update_x(void *x1, void *x2, void *x3, void *p1, void *p2, void *p3, void *alpha, int *p_cur, int *n)
void hip_fusedcg_cpld_update_p(void *p1, void *p2, void *p3, void *z1, void *z2, void *z3, void *po1, void *po2, void *po3, real *beta, int *n)
real hip_fusedcg_cpld_part2(void *a1, void *a2, void *a3, void *b, void *c1, void *c2, void *c3, void *alpha_d, real *alpha, int *p_cur, int *n)
void hip_fusedcg_cpld_part1(void *a1, void *a2, void *a3, void *b1, void *b2, void *b3, void *tmp, int *n)
Return the device pointer for an associated Fortran array.
Definition device.F90:113
Map a Fortran array to a device (allocate and associate)
Definition device.F90:83
Copy data between host and device (or device and device)
Definition device.F90:72
Unmap a Fortran array from a device (deassociate and free)
Definition device.F90:89
Apply cyclic boundary condition to a vector field.
Defines a Matrix-vector product.
Definition ax.f90:34
Coefficients.
Definition coef.f90:34
Definition comm.F90:1
type(mpi_datatype), public mpi_real_precision
MPI type for working precision of REAL types.
Definition comm.F90:54
integer, public pe_size
MPI size of communicator.
Definition comm.F90:62
type(mpi_comm), public neko_comm
MPI communicator.
Definition comm.F90:46
subroutine, public device_rzero(a_d, n, strm)
Zero a real vector.
subroutine, public device_copy(a_d, b_d, n, strm)
Copy a vector .
real(kind=rp) function, public device_glsc2(a_d, b_d, n, strm)
Weighted inner product .
Device abstraction, common interface for various accelerators.
Definition device.F90:34
subroutine, public device_event_sync(event)
Synchronize an event.
Definition device.F90:1667
integer, parameter, public host_to_device
Definition device.F90:48
subroutine, public device_free(x_d)
Deallocate memory on the device.
Definition device.F90:243
subroutine, public device_event_destroy(event)
Destroy a device event.
Definition device.F90:1623
subroutine, public device_alloc(x_d, s)
Allocate memory on the device.
Definition device.F90:212
subroutine, public device_event_create(event, flags)
Create a device event queue.
Definition device.F90:1589
Defines a field.
Definition field.f90:34
Defines a fused Conjugate Gradient method for accelerators.
subroutine device_fusedcg_cpld_update_x(x1_d, x2_d, x3_d, p1_d, p2_d, p3_d, alpha, p_cur, n)
type(ksp_monitor_t) function, dimension(3) fusedcg_cpld_device_solve_coupled(this, ax, x, y, z, fx, fy, fz, n, coef, bc_projector, gs_h, niter)
Pipelined PCG solve coupled solve.
type(ksp_monitor_t) function fusedcg_cpld_device_solve(this, ax, x, f, n, coef, bc_projector, gs_h, niter)
Pipelined PCG solve.
subroutine fusedcg_cpld_device_free(this)
Deallocate a pipelined PCG solver.
subroutine device_fusedcg_cpld_update_p(p1_d, p2_d, p3_d, z1_d, z2_d, z3_d, po1_d, po2_d, po3_d, beta, n)
subroutine fusedcg_cpld_device_init(this, n, max_iter, m, rel_tol, abs_tol, monitor)
Initialise a fused PCG solver.
real(kind=rp) function device_fusedcg_cpld_part2(a1_d, a2_d, a3_d, b_d, c1_d, c2_d, c3_d, alpha_d, alpha, p_cur, n)
integer, parameter device_fusedcg_cpld_p_space
subroutine device_fusedcg_cpld_part1(a1_d, a2_d, a3_d, b1_d, b2_d, b3_d, tmp_d, n)
Gather-scatter.
Implements the base abstract type for Krylov solvers plus helper types.
Definition krylov.f90:34
integer, parameter, public ksp_max_iter
Maximum number of iters.
Definition krylov.f90:52
Definition math.f90:60
real(kind=rp) function, public glsc3(a, b, c, n)
Weighted inner product .
Definition math.f90:1290
subroutine, public copy(a, b, n)
Copy a vector .
Definition math.f90:294
subroutine, public rzero(a, n)
Zero a real vector.
Definition math.f90:238
integer, parameter, public c_rp
Definition num_types.f90:15
integer, parameter, public rp
Global precision used in computations.
Definition num_types.f90:14
Operators.
Definition operators.f90:34
Krylov preconditioner.
Definition precon.f90:34
Implements scalar_projector_t.
Utilities.
Definition utils.f90:35
Implements boundary condition projectors for vector fields. Two types concrete types are provided: se...
subroutine, public vector_bc_projector_components(this, x, y, z)
Access the component scalar projectors from a segregated vector projector.
Base type for a matrix-vector product providing .
Definition ax.f90:43
Coefficients defined on a given (mesh, ) tuple. Arrays use indices (i,j,k,e): element e,...
Definition coef.f90:93
Fused preconditioned conjugate gradient method.
Gather-scatter kernel.
Type for storing initial and final residuals in a Krylov solver.
Definition krylov.f90:57
Base abstract type for a canonical Krylov method, solving .
Definition krylov.f90:74
Defines a canonical Krylov preconditioner.
Definition precon.f90:40
Projector for scalar boundary conditions.
Abstract type for resolving vector boundary conditions.