Two-pass X-then-Y halo exchange for a 3D x-face field. Batched: packs ALL nz layers into one buffer per direction. Face-x ownership: west rank sends (ng+1)nytnz eastward, east rank sends ngnytnz westward (D1 asymmetry preserved). Y-pass strip is nxt1ngnz (full face-x i extent × ng rows × nz).
| Type | Intent | Optional | Attributes | Name | ||
|---|---|---|---|---|---|---|
| real(kind=wp), | intent(inout) | :: | fld(oh_nx_total+1,oh_ny_total,nz) |
x-face field (nx_total+1, 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 | :: | nxt1 | ||||
| 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) |
subroutine ocean_halo_face_x_3d(fld, nz, device_resident) !! Two-pass X-then-Y halo exchange for a 3D x-face field. !! Batched: packs ALL nz layers into one buffer per direction. !! Face-x ownership: west rank sends (ng+1)*nyt*nz eastward, !! east rank sends ng*nyt*nz westward (D1 asymmetry preserved). !! Y-pass strip is nxt1*ng*nz (full face-x i extent × ng rows × nz). integer, intent(in) :: nz !! Number of vertical layers real(wp), intent(inout) :: fld(oh_nx_total + 1, oh_ny_total, nz) !! x-face field (nx_total+1, 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, nxt1, 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 call oh_count_face_x_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 nxt1 = nxt + 1 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) ! ---- X pass ---- if (oh_decomp%px == 1) then if (oh_periodic_x) then call ocean_periodic_wrap_face_x_3d(fld, nxt1, 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() ! Send east: ng+1 per j-row per layer — my LAST ng+1 physical faces ! Pack: buffer index = ((L-1)*nyt + (j-1))*(ng+1) + 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 + 1 oh_buf_send_east(((L - 1)*nyt + (j - 1))*(ng + 1) + k) = fld(nxl + k, j, L) end do end do end do end if ! Send west: ng per j-row per layer (first ng interior faces after west-seam copy) 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 + 1 + 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 + 1 oh_buf_send_east(((L - 1)*nyt + (j - 1))*(ng + 1) + 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 + 1 + 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, (ng + 1)*nyt*nz, rk_e, TAG_OC_W_TO_E, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_east, ng*nyt*nz, 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, ng*nyt*nz, rk_w, TAG_OC_E_TO_W, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_west, (ng + 1)*nyt*nz, 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, (ng + 1)*nyt*nz, rk_e, TAG_OC_W_TO_E, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_east, ng*nyt*nz, 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, ng*nyt*nz, rk_w, TAG_OC_E_TO_W, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_west, (ng + 1)*nyt*nz, 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 if (on_device) then ! Unpack recv-from-west: fills i=1..ng+1 (west ghosts + seam-copy) 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 + 1 fld(k, j, L) = oh_buf_recv_west(((L - 1)*nyt + (j - 1))*(ng + 1) + k) end do end do end do end if ! Unpack recv-from-east: fills east ghosts i=ng+nxl+2..ng+nxl+ng+1 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 + 1 + 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 + 1 fld(k, j, L) = oh_buf_recv_west(((L - 1)*nyt + (j - 1))*(ng + 1) + 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 + 1 + 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..nxt1 including x-ghosts) ---- if (oh_decomp%py == 1) then if (oh_periodic_y) then call ocean_periodic_wrap_face_x_3d(fld, nxt1, 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 Y-pass: buffer index = ((L-1)*ng + (k-1))*nxt1 + 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, nxt1 oh_buf_send_north(((L - 1)*ng + (k - 1))*nxt1 + 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, nxt1 oh_buf_send_south(((L - 1)*ng + (k - 1))*nxt1 + 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, nxt1 oh_buf_send_north(((L - 1)*ng + (k - 1))*nxt1 + 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, nxt1 oh_buf_send_south(((L - 1)*ng + (k - 1))*nxt1 + 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, nxt1*ng*nz, rk_n, TAG_OC_S_TO_N, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_north, nxt1*ng*nz, 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, nxt1*ng*nz, rk_s, TAG_OC_N_TO_S, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_south, nxt1*ng*nz, 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, nxt1*ng*nz, rk_n, TAG_OC_S_TO_N, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_north, nxt1*ng*nz, 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, nxt1*ng*nz, rk_s, TAG_OC_N_TO_S, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_south, nxt1*ng*nz, 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 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, nxt1 fld(j, k, L) = oh_buf_recv_south(((L - 1)*ng + (k - 1))*nxt1 + 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, nxt1 fld(j, ng + nyl + k, L) = oh_buf_recv_north(((L - 1)*ng + (k - 1))*nxt1 + 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, nxt1 fld(j, k, L) = oh_buf_recv_south(((L - 1)*ng + (k - 1))*nxt1 + j) end do end do end do end if if (need_n) then do L = 1, nz do k = 1, ng do j = 1, nxt1 fld(j, ng + nyl + k, L) = oh_buf_recv_north(((L - 1)*ng + (k - 1))*nxt1 + j) end do end do end do end if end if end if end subroutine ocean_halo_face_x_3d