Two-pass X-then-Y halo exchange for a cell-centred 3D field. Batched: packs ALL nz layers into one buffer per direction and posts ONE isend+irecv pair per needed direction (not per layer). Buffer index formula: ((L-1)rows + (row-1))width + k_within_strip where L=layer, rows=nyt, width=ng (centre X-pass).
| Type | Intent | Optional | Attributes | Name | ||
|---|---|---|---|---|---|---|
| real(kind=wp), | intent(inout) | :: | fld(oh_nx_total,oh_ny_total,nz) |
Cell-centred 3D field (nx_total, ny_total, nz) |
||
| integer, | intent(in) | :: | nz |
Number of vertical layers |
||
| logical, | intent(in), | optional | :: | device_resident |
If .false., operate on host memory (no OpenACC directives). Default (.true.) is the normal device-resident path. |
| Type | Visibility | Attributes | Name | Initial | |||
|---|---|---|---|---|---|---|---|
| integer, | private | :: | L | ||||
| type(comm_t), | private | :: | comm | ||||
| integer, | private | :: | j | ||||
| integer, | private | :: | k | ||||
| logical, | private | :: | need_e | ||||
| logical, | private | :: | need_n | ||||
| logical, | private | :: | need_s | ||||
| logical, | private | :: | need_w | ||||
| integer, | private | :: | ng | ||||
| integer, | private | :: | nreq | ||||
| integer, | private | :: | nxl | ||||
| integer, | private | :: | nxt | ||||
| integer, | private | :: | nyl | ||||
| integer, | private | :: | nyt | ||||
| logical, | private | :: | on_device | ||||
| type(request_t), | private | :: | reqs(MAX_REQS) | ||||
| integer, | private | :: | rk_e | ||||
| integer, | private | :: | rk_n | ||||
| integer, | private | :: | rk_s | ||||
| integer, | private | :: | rk_w | ||||
| type(MPI_Status), | private | :: | stats(MAX_REQS) | ||||
| integer, | private | :: | strip_ew_3d | ||||
| integer, | private | :: | strip_ns_3d |
subroutine ocean_halo_centre_3d(fld, nz, device_resident) !! Two-pass X-then-Y halo exchange for a cell-centred 3D field. !! Batched: packs ALL nz layers into one buffer per direction and !! posts ONE isend+irecv pair per needed direction (not per layer). !! Buffer index formula: ((L-1)*rows + (row-1))*width + k_within_strip !! where L=layer, rows=nyt, width=ng (centre X-pass). integer, intent(in) :: nz !! Number of vertical layers real(wp), intent(inout) :: fld(oh_nx_total, oh_ny_total, nz) !! Cell-centred 3D field (nx_total, ny_total, nz) logical, intent(in), optional :: device_resident !! If .false., operate on host memory (no OpenACC directives). !! Default (.true.) is the normal device-resident path. integer :: ng, nxl, nyl, nxt, nyt, j, k, L integer :: rk_e, rk_w, rk_n, rk_s integer :: nreq type(comm_t) :: comm type(request_t) :: reqs(MAX_REQS) type(MPI_Status) :: stats(MAX_REQS) logical :: need_e, need_w, need_n, need_s logical :: on_device integer :: strip_ew_3d, strip_ns_3d call oh_count_centre_3d() on_device = .true. if (present(device_resident)) on_device = device_resident ng = oh_nghost nxl = oh_nx_local nyl = oh_ny_local nxt = oh_nx_total nyt = oh_ny_total if (oh_decomp%px > 1 .or. oh_decomp%py > 1) comm = comm_env_compute_comm() call needs_flags(need_w, need_e, need_s, need_n) ! Ensure buffers are large enough for nz layers call ocean_halo_buffers_ensure_nz(nz) ! Batched strip sizes: ng elements per row per layer strip_ew_3d = ng*nyt*nz ! centre X-pass: ng columns per j-row per layer strip_ns_3d = nxt*ng*nz ! centre Y-pass: ng rows per i-col per layer ! ---- X pass ---- if (oh_decomp%px == 1) then if (oh_periodic_x) then call ocean_periodic_wrap_centre_3d(fld, nxt, nyt, nz, nxl, nyl, ng, .true., .false.) end if else if (need_e) rk_e = ew_rank_east() if (need_w) rk_w = ew_rank_west() ! Pack: buffer index = ((L-1)*nyt + (j-1))*ng + k if (on_device) then if (need_e) then !$acc parallel loop collapse(3) present(oh_buf_send_east, fld) do L = 1, nz do j = 1, nyt do k = 1, ng oh_buf_send_east(((L - 1)*nyt + (j - 1))*ng + k) = fld(nxl + k, j, L) end do end do end do end if if (need_w) then !$acc parallel loop collapse(3) present(oh_buf_send_west, fld) do L = 1, nz do j = 1, nyt do k = 1, ng oh_buf_send_west(((L - 1)*nyt + (j - 1))*ng + k) = fld(ng + k, j, L) end do end do end do end if else if (need_e) then do L = 1, nz do j = 1, nyt do k = 1, ng oh_buf_send_east(((L - 1)*nyt + (j - 1))*ng + k) = fld(nxl + k, j, L) end do end do end do end if if (need_w) then do L = 1, nz do j = 1, nyt do k = 1, ng oh_buf_send_west(((L - 1)*nyt + (j - 1))*ng + k) = fld(ng + k, j, L) end do end do end do end if end if nreq = 0 if (on_device) then if (need_e) then !$acc host_data use_device(oh_buf_send_east, oh_buf_recv_east) nreq = nreq + 1 call HALO_ISEND_N(comm, oh_buf_send_east, strip_ew_3d, rk_e, TAG_OC_W_TO_E, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_east, strip_ew_3d, rk_e, TAG_OC_E_TO_W, reqs(nreq)) !$acc end host_data end if if (need_w) then !$acc host_data use_device(oh_buf_send_west, oh_buf_recv_west) nreq = nreq + 1 call HALO_ISEND_N(comm, oh_buf_send_west, strip_ew_3d, rk_w, TAG_OC_E_TO_W, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_west, strip_ew_3d, rk_w, TAG_OC_W_TO_E, reqs(nreq)) !$acc end host_data end if else if (need_e) then nreq = nreq + 1 call HALO_ISEND_N(comm, oh_buf_send_east, strip_ew_3d, rk_e, TAG_OC_W_TO_E, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_east, strip_ew_3d, rk_e, TAG_OC_E_TO_W, reqs(nreq)) end if if (need_w) then nreq = nreq + 1 call HALO_ISEND_N(comm, oh_buf_send_west, strip_ew_3d, rk_w, TAG_OC_E_TO_W, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_west, strip_ew_3d, rk_w, TAG_OC_W_TO_E, reqs(nreq)) end if end if ! nreq = 2*(need_e + need_w); isends = nreq/2 (paired isend+irecv) if (nreq > 0) then call oh_count_msgs(nreq/2) call waitall(reqs(1:nreq), stats(1:nreq)) end if ! Unpack: X pass fills west ghosts (i=1..ng) and east ghosts (i=ng+nxl+1..nxt) if (on_device) then if (need_w) then !$acc parallel loop collapse(3) present(oh_buf_recv_west, fld) do L = 1, nz do j = 1, nyt do k = 1, ng fld(k, j, L) = oh_buf_recv_west(((L - 1)*nyt + (j - 1))*ng + k) end do end do end do end if if (need_e) then !$acc parallel loop collapse(3) present(oh_buf_recv_east, fld) do L = 1, nz do j = 1, nyt do k = 1, ng fld(ng + nxl + k, j, L) = oh_buf_recv_east(((L - 1)*nyt + (j - 1))*ng + k) end do end do end do end if else if (need_w) then do L = 1, nz do j = 1, nyt do k = 1, ng fld(k, j, L) = oh_buf_recv_west(((L - 1)*nyt + (j - 1))*ng + k) end do end do end do end if if (need_e) then do L = 1, nz do j = 1, nyt do k = 1, ng fld(ng + nxl + k, j, L) = oh_buf_recv_east(((L - 1)*nyt + (j - 1))*ng + k) end do end do end do end if end if end if ! ---- Y pass (spans full i=1..nxt including x-filled ghosts) ---- if (oh_decomp%py == 1) then if (oh_periodic_y) then call ocean_periodic_wrap_centre_3d(fld, nxt, nyt, nz, nxl, nyl, ng, .false., .true.) end if else if (need_n) rk_n = ns_rank_north() if (need_s) rk_s = ns_rank_south() ! Pack: buffer index = ((L-1)*ng + (k-1))*nxt + j if (on_device) then if (need_n) then !$acc parallel loop collapse(3) present(oh_buf_send_north, fld) do L = 1, nz do k = 1, ng do j = 1, nxt oh_buf_send_north(((L - 1)*ng + (k - 1))*nxt + j) = fld(j, nyl + k, L) end do end do end do end if if (need_s) then !$acc parallel loop collapse(3) present(oh_buf_send_south, fld) do L = 1, nz do k = 1, ng do j = 1, nxt oh_buf_send_south(((L - 1)*ng + (k - 1))*nxt + j) = fld(j, ng + k, L) end do end do end do end if else if (need_n) then do L = 1, nz do k = 1, ng do j = 1, nxt oh_buf_send_north(((L - 1)*ng + (k - 1))*nxt + j) = fld(j, nyl + k, L) end do end do end do end if if (need_s) then do L = 1, nz do k = 1, ng do j = 1, nxt oh_buf_send_south(((L - 1)*ng + (k - 1))*nxt + j) = fld(j, ng + k, L) end do end do end do end if end if nreq = 0 if (on_device) then if (need_n) then !$acc host_data use_device(oh_buf_send_north, oh_buf_recv_north) nreq = nreq + 1 call HALO_ISEND_N(comm, oh_buf_send_north, strip_ns_3d, rk_n, TAG_OC_S_TO_N, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_north, strip_ns_3d, rk_n, TAG_OC_N_TO_S, reqs(nreq)) !$acc end host_data end if if (need_s) then !$acc host_data use_device(oh_buf_send_south, oh_buf_recv_south) nreq = nreq + 1 call HALO_ISEND_N(comm, oh_buf_send_south, strip_ns_3d, rk_s, TAG_OC_N_TO_S, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_south, strip_ns_3d, rk_s, TAG_OC_S_TO_N, reqs(nreq)) !$acc end host_data end if else if (need_n) then nreq = nreq + 1 call HALO_ISEND_N(comm, oh_buf_send_north, strip_ns_3d, rk_n, TAG_OC_S_TO_N, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_north, strip_ns_3d, rk_n, TAG_OC_N_TO_S, reqs(nreq)) end if if (need_s) then nreq = nreq + 1 call HALO_ISEND_N(comm, oh_buf_send_south, strip_ns_3d, rk_s, TAG_OC_N_TO_S, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_south, strip_ns_3d, rk_s, TAG_OC_S_TO_N, reqs(nreq)) end if end if ! nreq = 2*(need_n + need_s); isends = nreq/2 (paired isend+irecv) if (nreq > 0) then call oh_count_msgs(nreq/2) call waitall(reqs(1:nreq), stats(1:nreq)) end if ! Unpack: Y pass fills south ghosts (j=1..ng) and north ghosts (j=ng+nyl+1..nyt) if (on_device) then if (need_s) then !$acc parallel loop collapse(3) present(oh_buf_recv_south, fld) do L = 1, nz do k = 1, ng do j = 1, nxt fld(j, k, L) = oh_buf_recv_south(((L - 1)*ng + (k - 1))*nxt + j) end do end do end do end if if (need_n) then !$acc parallel loop collapse(3) present(oh_buf_recv_north, fld) do L = 1, nz do k = 1, ng do j = 1, nxt fld(j, ng + nyl + k, L) = oh_buf_recv_north(((L - 1)*ng + (k - 1))*nxt + j) end do end do end do end if else if (need_s) then do L = 1, nz do k = 1, ng do j = 1, nxt fld(j, k, L) = oh_buf_recv_south(((L - 1)*ng + (k - 1))*nxt + j) end do end do end do end if if (need_n) then do L = 1, nz do k = 1, ng do j = 1, nxt fld(j, ng + nyl + k, L) = oh_buf_recv_north(((L - 1)*ng + (k - 1))*nxt + j) end do end do end do end if end if end if end subroutine ocean_halo_centre_3d