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.
| Type | Intent | Optional | 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 |
| 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 |
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