Two-pass X-then-Y halo exchange for a 3D y-face field. Batched: packs ALL nz layers into one buffer per direction. Face-y ownership (D1): south rank sends (ng+1)nxtnz northward, north rank sends ngnxtnz southward. X-pass uses standard centre-style (ng per column per layer over nyt1 rows).
| Type | Intent | Optional | Attributes | Name | ||
|---|---|---|---|---|---|---|
| real(kind=wp), | intent(inout) | :: | fld(oh_nx_total,oh_ny_total+1,nz) |
y-face field (nx_total, ny_total+1, 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 | :: | i | ||||
| 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 | ||||
| integer, | private | :: | nyt1 | ||||
| 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_y_3d(fld, nz, device_resident) !! Two-pass X-then-Y halo exchange for a 3D y-face field. !! Batched: packs ALL nz layers into one buffer per direction. !! Face-y ownership (D1): south rank sends (ng+1)*nxt*nz northward, !! north rank sends ng*nxt*nz southward. !! X-pass uses standard centre-style (ng per column per layer over nyt1 rows). integer, intent(in) :: nz !! Number of vertical layers real(wp), intent(inout) :: fld(oh_nx_total, oh_ny_total + 1, nz) !! y-face field (nx_total, ny_total+1, 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, nyt1, i, 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_y_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 nyt1 = nyt + 1 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: standard centre-style, full y-face extent 1..nyt1 ---- if (oh_decomp%px == 1) then if (oh_periodic_x) then call ocean_periodic_wrap_face_y_3d(fld, nxt, nyt1, 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 X-pass: buffer index = ((L-1)*ng + (k-1))*nyt1 + i if (on_device) then if (need_e) then !$acc parallel loop collapse(3) present(oh_buf_send_east, fld) do L = 1, nz do k = 1, ng do i = 1, nyt1 oh_buf_send_east(((L - 1)*ng + (k - 1))*nyt1 + i) = fld(nxl + k, i, 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 k = 1, ng do i = 1, nyt1 oh_buf_send_west(((L - 1)*ng + (k - 1))*nyt1 + i) = fld(ng + k, i, L) end do end do end do end if else if (need_e) then do L = 1, nz do k = 1, ng do i = 1, nyt1 oh_buf_send_east(((L - 1)*ng + (k - 1))*nyt1 + i) = fld(nxl + k, i, L) end do end do end do end if if (need_w) then do L = 1, nz do k = 1, ng do i = 1, nyt1 oh_buf_send_west(((L - 1)*ng + (k - 1))*nyt1 + i) = fld(ng + k, i, 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*nyt1*nz, rk_e, TAG_OC_W_TO_E, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_east, ng*nyt1*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*nyt1*nz, rk_w, TAG_OC_E_TO_W, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_west, ng*nyt1*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*nyt1*nz, rk_e, TAG_OC_W_TO_E, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_east, ng*nyt1*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*nyt1*nz, rk_w, TAG_OC_E_TO_W, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_west, ng*nyt1*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 if (need_w) then !$acc parallel loop collapse(3) present(oh_buf_recv_west, fld) do L = 1, nz do k = 1, ng do i = 1, nyt1 fld(k, i, L) = oh_buf_recv_west(((L - 1)*ng + (k - 1))*nyt1 + i) 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 k = 1, ng do i = 1, nyt1 fld(ng + nxl + k, i, L) = oh_buf_recv_east(((L - 1)*ng + (k - 1))*nyt1 + i) end do end do end do end if else if (need_w) then do L = 1, nz do k = 1, ng do i = 1, nyt1 fld(k, i, L) = oh_buf_recv_west(((L - 1)*ng + (k - 1))*nyt1 + i) end do end do end do end if if (need_e) then do L = 1, nz do k = 1, ng do i = 1, nyt1 fld(ng + nxl + k, i, L) = oh_buf_recv_east(((L - 1)*ng + (k - 1))*nyt1 + i) end do end do end do end if end if end if ! ---- Y pass: ownership — south rank owns seam face ---- if (oh_decomp%py == 1) then if (oh_periodic_y) then call ocean_periodic_wrap_face_y_3d(fld, nxt, nyt1, 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)*nxt + (i-1))*(ng+1) + k (send-north, ng+1 per col) ! and ((L-1)*nxt + (i-1))*ng + k (send-south, ng per col) if (on_device) then ! Send north: ng+1 per i-col per layer if (need_n) then !$acc parallel loop collapse(3) present(oh_buf_send_north, fld) do L = 1, nz do i = 1, nxt do k = 1, ng + 1 oh_buf_send_north(((L - 1)*nxt + (i - 1))*(ng + 1) + k) = fld(i, nyl + k, L) end do end do end do end if ! Send south: ng per i-col per layer if (need_s) then !$acc parallel loop collapse(3) present(oh_buf_send_south, fld) do L = 1, nz do i = 1, nxt do k = 1, ng oh_buf_send_south(((L - 1)*nxt + (i - 1))*ng + k) = fld(i, ng + 1 + k, L) end do end do end do end if else if (need_n) then do L = 1, nz do i = 1, nxt do k = 1, ng + 1 oh_buf_send_north(((L - 1)*nxt + (i - 1))*(ng + 1) + k) = fld(i, nyl + k, L) end do end do end do end if if (need_s) then do L = 1, nz do i = 1, nxt do k = 1, ng oh_buf_send_south(((L - 1)*nxt + (i - 1))*ng + k) = fld(i, ng + 1 + 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, (ng + 1)*nxt*nz, rk_n, TAG_OC_S_TO_N, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_north, ng*nxt*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, ng*nxt*nz, rk_s, TAG_OC_N_TO_S, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_south, (ng + 1)*nxt*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, (ng + 1)*nxt*nz, rk_n, TAG_OC_S_TO_N, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_north, ng*nxt*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, ng*nxt*nz, rk_s, TAG_OC_N_TO_S, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_south, (ng + 1)*nxt*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 ! Recv from south (ng+1 per i-col per layer) → fills south ghosts + seam-copy j=1..ng+1 if (need_s) then !$acc parallel loop collapse(3) present(oh_buf_recv_south, fld) do L = 1, nz do i = 1, nxt do k = 1, ng + 1 fld(i, k, L) = oh_buf_recv_south(((L - 1)*nxt + (i - 1))*(ng + 1) + k) end do end do end do end if ! Recv from north (ng per i-col per layer) → fills north ghosts j=ng+nyl+2..ng+nyl+ng+1 if (need_n) then !$acc parallel loop collapse(3) present(oh_buf_recv_north, fld) do L = 1, nz do i = 1, nxt do k = 1, ng fld(i, ng + nyl + 1 + k, L) = oh_buf_recv_north(((L - 1)*nxt + (i - 1))*ng + k) end do end do end do end if else if (need_s) then do L = 1, nz do i = 1, nxt do k = 1, ng + 1 fld(i, k, L) = oh_buf_recv_south(((L - 1)*nxt + (i - 1))*(ng + 1) + k) end do end do end do end if if (need_n) then do L = 1, nz do i = 1, nxt do k = 1, ng fld(i, ng + nyl + 1 + k, L) = oh_buf_recv_north(((L - 1)*nxt + (i - 1))*ng + k) end do end do end do end if end if end if end subroutine ocean_halo_face_y_3d