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