[ADD] Added s/z/c version of custom gather/scatter in halo exchange for new comm routines

This commit is contained in:
Stack-1
2026-08-07 15:41:47 +02:00
parent 87eca22895
commit f9f72f35e9
3 changed files with 126 additions and 6 deletions
+42 -2
View File
@@ -77,6 +77,7 @@ module psb_c_cuda_vect_mod
procedure, pass(x) :: set_sync => c_cuda_set_sync
procedure, pass(x) :: set_scal => c_cuda_set_scal
!!$ procedure, pass(x) :: set_vect => c_cuda_set_vect
procedure, pass(x) :: gthzv => c_cuda_gthzv
procedure, pass(x) :: gthzv_x => c_cuda_gthzv_x
procedure, pass(y) :: sctb => c_cuda_sctb
procedure, pass(y) :: sctb_x => c_cuda_sctb_x
@@ -231,6 +232,45 @@ contains
end select
end subroutine c_cuda_check_addr
subroutine c_cuda_gthzv(n,idx,x,y)
! GPU override of the plain-array-index gather y(1:n) = x(idx(1:n)).
! Without it the generic gth(n,idx(:),buf) -- used ONLY by the RMA swap path
! (psi_c_swapdata, rma_pull/push) -- fell back to the host c_base_gthzv, which
! calls x%sync() and copies the WHOLE vector device->host on EVERY swap. Here
! we gather on the device straight from x%deviceVect (mirrors the class-default
! branch of c_cuda_gthzv_x): only the n boundary indices move H2D, no full D2H.
use psb_cuda_env_mod
use psi_serial_mod
implicit none
integer(psb_mpk_) :: n
integer(psb_ipk_) :: idx(:)
complex(psb_spk_) :: y(:)
class(psb_c_vect_cuda) :: x
integer :: info, ni
info = 0
if (x%is_host()) call x%sync() ! ensure device copy is current (no D2H of the whole vector)
ni = size(idx)
if (x%i_buf_sz < ni) then
if (c_associated(x%i_buf)) then
call freeInt(x%i_buf); x%i_buf = c_null_ptr
end if
info = allocateInt(x%i_buf,ni); x%i_buf_sz = ni
end if
if (x%dt_buf_sz < n) then
if (c_associated(x%dt_buf)) then
call freeFloatComplex(x%dt_buf); x%dt_buf = c_null_ptr
end if
info = allocateFloatComplex(x%dt_buf,n); x%dt_buf_sz = n
end if
if (info == 0) info = writeInt(x%i_buf,idx,ni)
if (info == 0) info = igathMultiVecDeviceFloatComplex(x%deviceVect, 0, n, 1, x%i_buf, 1, x%dt_buf, 1)
if (info == 0) info = readFloatComplex(x%dt_buf,y,n)
end subroutine c_cuda_gthzv
subroutine c_cuda_gthzv_x(i,n,idx,x,y)
use psb_cuda_env_mod
use psi_serial_mod
@@ -1459,7 +1499,7 @@ contains
!!$ allocate(x%buffer(n),stat=info)
!!$ if (info == 0) info = inner_register(x%buffer,x%dt_buf)
!!$ endif
!!$ info = igathMultiVecDeviceDouble(x%deviceVect,&
!!$ info = igathMultiVecDeviceFloatComplex(x%deviceVect,&
!!$ & 0, i, n, ii%deviceVect, x%dt_buf, 1)
!!$ call psb_cudaSync()
!!$ y(1:n) = x%buffer(1:n)
@@ -1514,7 +1554,7 @@ contains
!!$ if (info == 0) info = inner_register(y%buffer,y%dt_buf)
!!$ endif
!!$ y%buffer(1:n) = x(1:n)
!!$ info = iscatMultiVecDeviceDouble(y%deviceVect,&
!!$ info = iscatMultiVecDeviceFloatComplex(y%deviceVect,&
!!$ & 0, i, n, ii%deviceVect, y%dt_buf, 1,beta)
!!$
!!$ call y%set_dev()
+42 -2
View File
@@ -77,6 +77,7 @@ module psb_s_cuda_vect_mod
procedure, pass(x) :: set_sync => s_cuda_set_sync
procedure, pass(x) :: set_scal => s_cuda_set_scal
!!$ procedure, pass(x) :: set_vect => s_cuda_set_vect
procedure, pass(x) :: gthzv => s_cuda_gthzv
procedure, pass(x) :: gthzv_x => s_cuda_gthzv_x
procedure, pass(y) :: sctb => s_cuda_sctb
procedure, pass(y) :: sctb_x => s_cuda_sctb_x
@@ -231,6 +232,45 @@ contains
end select
end subroutine s_cuda_check_addr
subroutine s_cuda_gthzv(n,idx,x,y)
! GPU override of the plain-array-index gather y(1:n) = x(idx(1:n)).
! Without it the generic gth(n,idx(:),buf) -- used ONLY by the RMA swap path
! (psi_s_swapdata, rma_pull/push) -- fell back to the host s_base_gthzv, which
! calls x%sync() and copies the WHOLE vector device->host on EVERY swap. Here
! we gather on the device straight from x%deviceVect (mirrors the class-default
! branch of s_cuda_gthzv_x): only the n boundary indices move H2D, no full D2H.
use psb_cuda_env_mod
use psi_serial_mod
implicit none
integer(psb_mpk_) :: n
integer(psb_ipk_) :: idx(:)
real(psb_spk_) :: y(:)
class(psb_s_vect_cuda) :: x
integer :: info, ni
info = 0
if (x%is_host()) call x%sync() ! ensure device copy is current (no D2H of the whole vector)
ni = size(idx)
if (x%i_buf_sz < ni) then
if (c_associated(x%i_buf)) then
call freeInt(x%i_buf); x%i_buf = c_null_ptr
end if
info = allocateInt(x%i_buf,ni); x%i_buf_sz = ni
end if
if (x%dt_buf_sz < n) then
if (c_associated(x%dt_buf)) then
call freeFloat(x%dt_buf); x%dt_buf = c_null_ptr
end if
info = allocateFloat(x%dt_buf,n); x%dt_buf_sz = n
end if
if (info == 0) info = writeInt(x%i_buf,idx,ni)
if (info == 0) info = igathMultiVecDeviceFloat(x%deviceVect, 0, n, 1, x%i_buf, 1, x%dt_buf, 1)
if (info == 0) info = readFloat(x%dt_buf,y,n)
end subroutine s_cuda_gthzv
subroutine s_cuda_gthzv_x(i,n,idx,x,y)
use psb_cuda_env_mod
use psi_serial_mod
@@ -1459,7 +1499,7 @@ contains
!!$ allocate(x%buffer(n),stat=info)
!!$ if (info == 0) info = inner_register(x%buffer,x%dt_buf)
!!$ endif
!!$ info = igathMultiVecDeviceDouble(x%deviceVect,&
!!$ info = igathMultiVecDeviceFloat(x%deviceVect,&
!!$ & 0, i, n, ii%deviceVect, x%dt_buf, 1)
!!$ call psb_cudaSync()
!!$ y(1:n) = x%buffer(1:n)
@@ -1514,7 +1554,7 @@ contains
!!$ if (info == 0) info = inner_register(y%buffer,y%dt_buf)
!!$ endif
!!$ y%buffer(1:n) = x(1:n)
!!$ info = iscatMultiVecDeviceDouble(y%deviceVect,&
!!$ info = iscatMultiVecDeviceFloat(y%deviceVect,&
!!$ & 0, i, n, ii%deviceVect, y%dt_buf, 1,beta)
!!$
!!$ call y%set_dev()
+42 -2
View File
@@ -77,6 +77,7 @@ module psb_z_cuda_vect_mod
procedure, pass(x) :: set_sync => z_cuda_set_sync
procedure, pass(x) :: set_scal => z_cuda_set_scal
!!$ procedure, pass(x) :: set_vect => z_cuda_set_vect
procedure, pass(x) :: gthzv => z_cuda_gthzv
procedure, pass(x) :: gthzv_x => z_cuda_gthzv_x
procedure, pass(y) :: sctb => z_cuda_sctb
procedure, pass(y) :: sctb_x => z_cuda_sctb_x
@@ -231,6 +232,45 @@ contains
end select
end subroutine z_cuda_check_addr
subroutine z_cuda_gthzv(n,idx,x,y)
! GPU override of the plain-array-index gather y(1:n) = x(idx(1:n)).
! Without it the generic gth(n,idx(:),buf) -- used ONLY by the RMA swap path
! (psi_z_swapdata, rma_pull/push) -- fell back to the host z_base_gthzv, which
! calls x%sync() and copies the WHOLE vector device->host on EVERY swap. Here
! we gather on the device straight from x%deviceVect (mirrors the class-default
! branch of z_cuda_gthzv_x): only the n boundary indices move H2D, no full D2H.
use psb_cuda_env_mod
use psi_serial_mod
implicit none
integer(psb_mpk_) :: n
integer(psb_ipk_) :: idx(:)
complex(psb_dpk_) :: y(:)
class(psb_z_vect_cuda) :: x
integer :: info, ni
info = 0
if (x%is_host()) call x%sync() ! ensure device copy is current (no D2H of the whole vector)
ni = size(idx)
if (x%i_buf_sz < ni) then
if (c_associated(x%i_buf)) then
call freeInt(x%i_buf); x%i_buf = c_null_ptr
end if
info = allocateInt(x%i_buf,ni); x%i_buf_sz = ni
end if
if (x%dt_buf_sz < n) then
if (c_associated(x%dt_buf)) then
call freeDoubleComplex(x%dt_buf); x%dt_buf = c_null_ptr
end if
info = allocateDoubleComplex(x%dt_buf,n); x%dt_buf_sz = n
end if
if (info == 0) info = writeInt(x%i_buf,idx,ni)
if (info == 0) info = igathMultiVecDeviceDoubleComplex(x%deviceVect, 0, n, 1, x%i_buf, 1, x%dt_buf, 1)
if (info == 0) info = readDoubleComplex(x%dt_buf,y,n)
end subroutine z_cuda_gthzv
subroutine z_cuda_gthzv_x(i,n,idx,x,y)
use psb_cuda_env_mod
use psi_serial_mod
@@ -1459,7 +1499,7 @@ contains
!!$ allocate(x%buffer(n),stat=info)
!!$ if (info == 0) info = inner_register(x%buffer,x%dt_buf)
!!$ endif
!!$ info = igathMultiVecDeviceDouble(x%deviceVect,&
!!$ info = igathMultiVecDeviceDoubleComplex(x%deviceVect,&
!!$ & 0, i, n, ii%deviceVect, x%dt_buf, 1)
!!$ call psb_cudaSync()
!!$ y(1:n) = x%buffer(1:n)
@@ -1514,7 +1554,7 @@ contains
!!$ if (info == 0) info = inner_register(y%buffer,y%dt_buf)
!!$ endif
!!$ y%buffer(1:n) = x(1:n)
!!$ info = iscatMultiVecDeviceDouble(y%deviceVect,&
!!$ info = iscatMultiVecDeviceDoubleComplex(y%deviceVect,&
!!$ & 0, i, n, ii%deviceVect, y%dt_buf, 1,beta)
!!$
!!$ call y%set_dev()