halo_exchange_3d_device Subroutine

public subroutine halo_exchange_3d_device(fld, decomp, nghost, nx_local, ny_local, nz)

GPU-direct halo exchange for a 3D field, batched across layers.

Packs all nz layers into one persistent device buffer per direction, fires one MPI Isend/Irecv per direction, and unpacks all layers in one kernel. This replaces the previous do kk = 1, nz; call halo_exchange_2d_device(fld(:,:,kk), ...) loop, which (a) issued nz×4 MPI messages per call and (b) tripped NVHPC’s per-call array-section descriptor push for each 2D slice – both visible as gaps in the timeline.

fld is declared explicit-shape so that NVHPC passes a raw pointer + bounds rather than a descriptor that has to be re-attached to the device on each call.

Arguments

Type IntentOptional Attributes Name
real(kind=wp), intent(inout) :: fld(nx_local+2*nghost,ny_local+2*nghost,nz)
type(decomp_t), intent(in) :: decomp
integer, intent(in) :: nghost
integer, intent(in) :: nx_local
integer, intent(in) :: ny_local
integer, intent(in) :: nz

Calls

proc~~halo_exchange_3d_device~~CallsGraph proc~halo_exchange_3d_device halo_exchange_3d_device comm_irecv_real_sp_array_n comm_irecv_real_sp_array_n proc~halo_exchange_3d_device->comm_irecv_real_sp_array_n comm_isend_real_sp_array_n comm_isend_real_sp_array_n proc~halo_exchange_3d_device->comm_isend_real_sp_array_n proc~comm_env_compute_comm comm_env_compute_comm proc~halo_exchange_3d_device->proc~comm_env_compute_comm proc~decomp_rank_from_coords decomp_rank_from_coords proc~halo_exchange_3d_device->proc~decomp_rank_from_coords proc~halo_sync_buffers_ensure_3d halo_sync_buffers_ensure_3d proc~halo_exchange_3d_device->proc~halo_sync_buffers_ensure_3d waitall waitall proc~halo_exchange_3d_device->waitall comm_world comm_world proc~comm_env_compute_comm->comm_world proc~halo_sync_buffers_cleanup_3d halo_sync_buffers_cleanup_3d proc~halo_sync_buffers_ensure_3d->proc~halo_sync_buffers_cleanup_3d

Variables

Type Visibility Attributes Name Initial
integer, private :: L
type(comm_t), private :: comm
integer, private :: i
integer, private :: j
integer, private :: k
integer, private :: nreq
integer, private :: nx_total
integer, private :: ny_total
integer, private :: rank_east
integer, private :: rank_north
integer, private :: rank_south
integer, private :: rank_west
type(request_t), private :: reqs(MAX_REQS)
type(MPI_Status), private :: stats(MAX_REQS)
integer, private :: strip_ew
integer, private :: strip_sn

Source Code

   subroutine halo_exchange_3d_device(fld, decomp, nghost, nx_local, ny_local, nz)
      !! GPU-direct halo exchange for a 3D field, batched across layers.
      !!
      !! Packs all `nz` layers into one persistent device buffer per
      !! direction, fires one MPI Isend/Irecv per direction, and unpacks
      !! all layers in one kernel.  This replaces the previous
      !! `do kk = 1, nz; call halo_exchange_2d_device(fld(:,:,kk), ...)`
      !! loop, which (a) issued nz×4 MPI messages per call and (b)
      !! tripped NVHPC's per-call array-section descriptor push for each
      !! 2D slice -- both visible as gaps in the timeline.
      !!
      !! `fld` is declared explicit-shape so that NVHPC passes a raw
      !! pointer + bounds rather than a descriptor that has to be
      !! re-attached to the device on each call.
      integer, intent(in) :: nghost, nx_local, ny_local, nz
      real(wp), intent(inout) :: fld(nx_local + 2*nghost, ny_local + 2*nghost, nz)
      type(decomp_t), intent(in) :: decomp

      type(comm_t) :: comm
      integer :: nx_total, ny_total
      integer :: i, j, k, L
      integer :: rank_west, rank_east, rank_south, rank_north
      type(request_t) :: reqs(MAX_REQS)
      type(MPI_Status) :: stats(MAX_REQS)
      integer :: nreq
      integer :: strip_ew, strip_sn

      comm = comm_env_compute_comm()
      nx_total = nx_local + 2*nghost
      ny_total = ny_local + 2*nghost
      strip_ew = nghost*ny_total*nz
      strip_sn = nx_total*nghost*nz

      call halo_sync_buffers_ensure_3d(nghost, nx_total, ny_total, nz)

      ! --- Pack on device (all layers in a single kernel per direction) ---
      if (.not. decomp%has_west) then
         rank_west = decomp_rank_from_coords(decomp%px, decomp%rx - 1, decomp%ry)
         !$acc parallel loop collapse(3) present(hs3_buf_send_west, fld)
         do L = 1, nz
            do j = 1, ny_total
               do k = 1, nghost
                  hs3_buf_send_west(((L - 1)*ny_total + (j - 1))*nghost + k) = fld(nghost + k, j, L)
               end do
            end do
         end do
      end if

      if (.not. decomp%has_east) then
         rank_east = decomp_rank_from_coords(decomp%px, decomp%rx + 1, decomp%ry)
         !$acc parallel loop collapse(3) present(hs3_buf_send_east, fld)
         do L = 1, nz
            do j = 1, ny_total
               do k = 1, nghost
                  hs3_buf_send_east(((L - 1)*ny_total + (j - 1))*nghost + k) = fld(nghost + nx_local - nghost + k, j, L)
               end do
            end do
         end do
      end if

      if (.not. decomp%has_south) then
         rank_south = decomp_rank_from_coords(decomp%px, decomp%rx, decomp%ry - 1)
         !$acc parallel loop collapse(3) present(hs3_buf_send_south, fld)
         do L = 1, nz
            do k = 1, nghost
               do i = 1, nx_total
                  hs3_buf_send_south(((L - 1)*nghost + (k - 1))*nx_total + i) = fld(i, nghost + k, L)
               end do
            end do
         end do
      end if

      if (.not. decomp%has_north) then
         rank_north = decomp_rank_from_coords(decomp%px, decomp%rx, decomp%ry + 1)
         !$acc parallel loop collapse(3) present(hs3_buf_send_north, fld)
         do L = 1, nz
            do k = 1, nghost
               do i = 1, nx_total
                  hs3_buf_send_north(((L - 1)*nghost + (k - 1))*nx_total + i) = fld(i, nghost + ny_local - nghost + k, L)
               end do
            end do
         end do
      end if

      ! --- One MPI exchange per direction, all layers in one message ---
      nreq = 0

      if (.not. decomp%has_west) then
         !$acc host_data use_device(hs3_buf_send_west, hs3_buf_recv_west)
         nreq = nreq + 1
         call HALO_ISEND_N(comm, hs3_buf_send_west, strip_ew, rank_west, 1, reqs(nreq))
         nreq = nreq + 1
         call HALO_IRECV_N(comm, hs3_buf_recv_west, strip_ew, rank_west, 2, reqs(nreq))
         !$acc end host_data
      end if

      if (.not. decomp%has_east) then
         !$acc host_data use_device(hs3_buf_send_east, hs3_buf_recv_east)
         nreq = nreq + 1
         call HALO_ISEND_N(comm, hs3_buf_send_east, strip_ew, rank_east, 2, reqs(nreq))
         nreq = nreq + 1
         call HALO_IRECV_N(comm, hs3_buf_recv_east, strip_ew, rank_east, 1, reqs(nreq))
         !$acc end host_data
      end if

      if (.not. decomp%has_south) then
         !$acc host_data use_device(hs3_buf_send_south, hs3_buf_recv_south)
         nreq = nreq + 1
         call HALO_ISEND_N(comm, hs3_buf_send_south, strip_sn, rank_south, 3, reqs(nreq))
         nreq = nreq + 1
         call HALO_IRECV_N(comm, hs3_buf_recv_south, strip_sn, rank_south, 4, reqs(nreq))
         !$acc end host_data
      end if

      if (.not. decomp%has_north) then
         !$acc host_data use_device(hs3_buf_send_north, hs3_buf_recv_north)
         nreq = nreq + 1
         call HALO_ISEND_N(comm, hs3_buf_send_north, strip_sn, rank_north, 4, reqs(nreq))
         nreq = nreq + 1
         call HALO_IRECV_N(comm, hs3_buf_recv_north, strip_sn, rank_north, 3, reqs(nreq))
         !$acc end host_data
      end if

      if (nreq > 0) call waitall(reqs(1:nreq), stats(1:nreq))

      ! --- Unpack on device (all layers in a single kernel per direction) ---
      if (.not. decomp%has_west) then
         !$acc parallel loop collapse(3) present(hs3_buf_recv_west, fld)
         do L = 1, nz
            do j = 1, ny_total
               do k = 1, nghost
                  fld(k, j, L) = hs3_buf_recv_west(((L - 1)*ny_total + (j - 1))*nghost + k)
               end do
            end do
         end do
      end if

      if (.not. decomp%has_east) then
         !$acc parallel loop collapse(3) present(hs3_buf_recv_east, fld)
         do L = 1, nz
            do j = 1, ny_total
               do k = 1, nghost
                  fld(nghost + nx_local + k, j, L) = hs3_buf_recv_east(((L - 1)*ny_total + (j - 1))*nghost + k)
               end do
            end do
         end do
      end if

      if (.not. decomp%has_south) then
         !$acc parallel loop collapse(3) present(hs3_buf_recv_south, fld)
         do L = 1, nz
            do k = 1, nghost
               do i = 1, nx_total
                  fld(i, k, L) = hs3_buf_recv_south(((L - 1)*nghost + (k - 1))*nx_total + i)
               end do
            end do
         end do
      end if

      if (.not. decomp%has_north) then
         !$acc parallel loop collapse(3) present(hs3_buf_recv_north, fld)
         do L = 1, nz
            do k = 1, nghost
               do i = 1, nx_total
                  fld(i, nghost + ny_local + k, L) = hs3_buf_recv_north(((L - 1)*nghost + (k - 1))*nx_total + i)
               end do
            end do
         end do
      end if

   end subroutine halo_exchange_3d_device