halo_exchange_2d_device Subroutine

public subroutine halo_exchange_2d_device(fld, decomp, nghost, nx_local, ny_local)

GPU-direct halo exchange via CUDA-aware MPI

Pack/unpack buffers live on the device. MPI operates on device pointers via !$acc host_data use_device, eliminating the full-array GPU<->host copies required by the host-staged path.

Uses persistent module-level send/recv buffers (hs_buf_*) so back-to-back halo calls do not hit the cudaMalloc/cudaFree allocator on every call. Lazy-allocated by halo_sync_buffers_ensure; the solver normally calls that from solver_enter_data, but a guard here keeps the routine self-contained for callers that haven’t been migrated yet.

Arguments

Type IntentOptional Attributes Name
real(kind=wp), intent(inout) :: fld(:,:)

2D field with ghost cells (present on device)

type(decomp_t), intent(in) :: decomp
integer, intent(in) :: nghost
integer, intent(in) :: nx_local
integer, intent(in) :: ny_local

Calls

proc~~halo_exchange_2d_device~~CallsGraph proc~halo_exchange_2d_device halo_exchange_2d_device comm_irecv_real_sp_array_n comm_irecv_real_sp_array_n proc~halo_exchange_2d_device->comm_irecv_real_sp_array_n comm_isend_real_sp_array_n comm_isend_real_sp_array_n proc~halo_exchange_2d_device->comm_isend_real_sp_array_n proc~comm_env_compute_comm comm_env_compute_comm proc~halo_exchange_2d_device->proc~comm_env_compute_comm proc~decomp_rank_from_coords decomp_rank_from_coords proc~halo_exchange_2d_device->proc~decomp_rank_from_coords proc~halo_sync_buffers_ensure halo_sync_buffers_ensure proc~halo_exchange_2d_device->proc~halo_sync_buffers_ensure waitall waitall proc~halo_exchange_2d_device->waitall comm_world comm_world proc~comm_env_compute_comm->comm_world proc~halo_sync_buffers_cleanup halo_sync_buffers_cleanup proc~halo_sync_buffers_ensure->proc~halo_sync_buffers_cleanup

Variables

Type Visibility Attributes Name Initial
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_2d_device(fld, decomp, nghost, nx_local, ny_local)
      !! GPU-direct halo exchange via CUDA-aware MPI
      !!
      !! Pack/unpack buffers live on the device. MPI operates on device
      !! pointers via !$acc host_data use_device, eliminating the
      !! full-array GPU<->host copies required by the host-staged path.
      !!
      !! Uses persistent module-level send/recv buffers
      !! (`hs_buf_*`) so back-to-back halo calls do not hit the
      !! cudaMalloc/cudaFree allocator on every call.  Lazy-allocated by
      !! `halo_sync_buffers_ensure`; the solver normally calls that from
      !! `solver_enter_data`, but a guard here keeps the routine
      !! self-contained for callers that haven't been migrated yet.
      real(wp), intent(inout) :: fld(:, :)
         !! 2D field with ghost cells (present on device)
      type(decomp_t), intent(in) :: decomp
      integer, intent(in) :: nghost
      integer, intent(in) :: nx_local
      integer, intent(in) :: ny_local

      type(comm_t) :: comm
      integer :: nx_total, ny_total
      integer :: i, j, k
      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
      strip_sn = nx_total*nghost

      call halo_sync_buffers_ensure(nghost, nx_total, ny_total)

      ! --- Pack on device ---
      if (.not. decomp%has_west) then
         rank_west = decomp_rank_from_coords(decomp%px, decomp%rx - 1, decomp%ry)
         !$acc parallel loop collapse(2) present(hs_buf_send_west, fld)
         do j = 1, ny_total
            do k = 1, nghost
               hs_buf_send_west((j - 1)*nghost + k) = fld(nghost + k, j)
            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(2) present(hs_buf_send_east, fld)
         do j = 1, ny_total
            do k = 1, nghost
               hs_buf_send_east((j - 1)*nghost + k) = fld(nghost + nx_local - nghost + k, j)
            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(2) present(hs_buf_send_south, fld)
         do k = 1, nghost
            do i = 1, nx_total
               hs_buf_send_south((k - 1)*nx_total + i) = fld(i, nghost + k)
            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(2) present(hs_buf_send_north, fld)
         do k = 1, nghost
            do i = 1, nx_total
               hs_buf_send_north((k - 1)*nx_total + i) = fld(i, nghost + ny_local - nghost + k)
            end do
         end do
      end if

      ! --- MPI with device pointers ---
      nreq = 0

      if (.not. decomp%has_west) then
         !$acc host_data use_device(hs_buf_send_west, hs_buf_recv_west)
         nreq = nreq + 1
         call HALO_ISEND_N(comm, hs_buf_send_west, strip_ew, rank_west, 1, reqs(nreq))
         nreq = nreq + 1
         call HALO_IRECV_N(comm, hs_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(hs_buf_send_east, hs_buf_recv_east)
         nreq = nreq + 1
         call HALO_ISEND_N(comm, hs_buf_send_east, strip_ew, rank_east, 2, reqs(nreq))
         nreq = nreq + 1
         call HALO_IRECV_N(comm, hs_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(hs_buf_send_south, hs_buf_recv_south)
         nreq = nreq + 1
         call HALO_ISEND_N(comm, hs_buf_send_south, strip_sn, rank_south, 3, reqs(nreq))
         nreq = nreq + 1
         call HALO_IRECV_N(comm, hs_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(hs_buf_send_north, hs_buf_recv_north)
         nreq = nreq + 1
         call HALO_ISEND_N(comm, hs_buf_send_north, strip_sn, rank_north, 4, reqs(nreq))
         nreq = nreq + 1
         call HALO_IRECV_N(comm, hs_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 ---
      if (.not. decomp%has_west) then
         !$acc parallel loop collapse(2) present(hs_buf_recv_west, fld)
         do j = 1, ny_total
            do k = 1, nghost
               fld(k, j) = hs_buf_recv_west((j - 1)*nghost + k)
            end do
         end do
      end if

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

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

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

   end subroutine halo_exchange_2d_device