Two-pass X-then-Y halo exchange for a 2D x-face field. Core body shared by ocean_halo_face_x_2d (ng=oh_nghost) and ocean_halo_face_x_2d_wide (ng=ng_wide).
| Type | Intent | Optional | Attributes | Name | ||
|---|---|---|---|---|---|---|
| real(kind=wp), | intent(inout) | :: | fld(oh_nx_local+2*ng+1,oh_ny_local+2*ng) |
x-face field (nx_local + 2ng + 1, ny_local + 2ng) |
||
| integer, | intent(in) | :: | ng |
Ghost cell width for this exchange |
||
| logical, | intent(in), | optional | :: | device_resident |
If .false., operate on host memory. Default = .true. |
| Type | Visibility | Attributes | Name | Initial | |||
|---|---|---|---|---|---|---|---|
| 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 | :: | 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_2d_impl(fld, ng, device_resident) !! Two-pass X-then-Y halo exchange for a 2D x-face field. !! Core body shared by ocean_halo_face_x_2d (ng=oh_nghost) and !! ocean_halo_face_x_2d_wide (ng=ng_wide). integer, intent(in) :: ng !! Ghost cell width for this exchange real(wp), intent(inout) :: fld(oh_nx_local + 2*ng + 1, oh_ny_local + 2*ng) !! x-face field (nx_local + 2*ng + 1, ny_local + 2*ng) logical, intent(in), optional :: device_resident !! If .false., operate on host memory. Default = .true. integer :: nxl, nyl, nxt, nxt1, nyt, j, k 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 on_device = .true. if (present(device_resident)) on_device = device_resident nxl = oh_nx_local nyl = oh_ny_local nxt = nxl + 2*ng nxt1 = nxt + 1 nyt = nyl + 2*ng 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) ! ---- X pass ---- if (oh_decomp%px == 1) then if (oh_periodic_x) then call ocean_periodic_wrap_face_x_2d(fld, nxt1, nyt, nxl, nyl, ng, .true., .false.) end if else if (need_e) rk_e = ew_rank_east() if (need_w) rk_w = ew_rank_west() if (on_device) then if (need_e) then !$acc parallel loop collapse(2) present(oh_buf_send_east, fld) do j = 1, nyt do k = 1, ng + 1 oh_buf_send_east((j - 1)*(ng + 1) + k) = fld(nxl + k, j) end do end do end if if (need_w) then !$acc parallel loop collapse(2) present(oh_buf_send_west, fld) do j = 1, nyt do k = 1, ng oh_buf_send_west((j - 1)*ng + k) = fld(ng + 1 + k, j) end do end do end if else if (need_e) then do j = 1, nyt do k = 1, ng + 1 oh_buf_send_east((j - 1)*(ng + 1) + k) = fld(nxl + k, j) end do end do end if if (need_w) then do j = 1, nyt do k = 1, ng oh_buf_send_west((j - 1)*ng + k) = fld(ng + 1 + k, j) 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, rk_e, TAG_OC_W_TO_E, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_east, ng*nyt, 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, rk_w, TAG_OC_E_TO_W, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_west, (ng + 1)*nyt, 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, rk_e, TAG_OC_W_TO_E, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_east, ng*nyt, 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, rk_w, TAG_OC_E_TO_W, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_west, (ng + 1)*nyt, rk_w, TAG_OC_W_TO_E, reqs(nreq)) end if end if 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(2) present(oh_buf_recv_west, fld) do j = 1, nyt do k = 1, ng + 1 fld(k, j) = oh_buf_recv_west((j - 1)*(ng + 1) + k) end do end do end if if (need_e) then !$acc parallel loop collapse(2) present(oh_buf_recv_east, fld) do j = 1, nyt do k = 1, ng fld(ng + nxl + 1 + k, j) = oh_buf_recv_east((j - 1)*ng + k) end do end do end if else if (need_w) then do j = 1, nyt do k = 1, ng + 1 fld(k, j) = oh_buf_recv_west((j - 1)*(ng + 1) + k) end do end do end if if (need_e) then do j = 1, nyt do k = 1, ng fld(ng + nxl + 1 + k, j) = oh_buf_recv_east((j - 1)*ng + k) 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_2d(fld, nxt1, nyt, nxl, nyl, ng, .false., .true.) end if else if (need_n) rk_n = ns_rank_north() if (need_s) rk_s = ns_rank_south() if (on_device) then if (need_n) then !$acc parallel loop collapse(2) present(oh_buf_send_north, fld) do k = 1, ng do j = 1, nxt1 oh_buf_send_north((k - 1)*nxt1 + j) = fld(j, nyl + k) end do end do end if if (need_s) then !$acc parallel loop collapse(2) present(oh_buf_send_south, fld) do k = 1, ng do j = 1, nxt1 oh_buf_send_south((k - 1)*nxt1 + j) = fld(j, ng + k) end do end do end if else if (need_n) then do k = 1, ng do j = 1, nxt1 oh_buf_send_north((k - 1)*nxt1 + j) = fld(j, nyl + k) end do end do end if if (need_s) then do k = 1, ng do j = 1, nxt1 oh_buf_send_south((k - 1)*nxt1 + j) = fld(j, ng + k) 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, rk_n, TAG_OC_S_TO_N, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_north, nxt1*ng, 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, rk_s, TAG_OC_N_TO_S, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_south, nxt1*ng, 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, rk_n, TAG_OC_S_TO_N, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_north, nxt1*ng, 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, rk_s, TAG_OC_N_TO_S, reqs(nreq)) nreq = nreq + 1 call HALO_IRECV_N(comm, oh_buf_recv_south, nxt1*ng, rk_s, TAG_OC_S_TO_N, reqs(nreq)) end if end if 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(2) present(oh_buf_recv_south, fld) do k = 1, ng do j = 1, nxt1 fld(j, k) = oh_buf_recv_south((k - 1)*nxt1 + j) end do end do end if if (need_n) then !$acc parallel loop collapse(2) present(oh_buf_recv_north, fld) do k = 1, ng do j = 1, nxt1 fld(j, ng + nyl + k) = oh_buf_recv_north((k - 1)*nxt1 + j) end do end do end if else if (need_s) then do k = 1, ng do j = 1, nxt1 fld(j, k) = oh_buf_recv_south((k - 1)*nxt1 + j) end do end do end if if (need_n) then do k = 1, ng do j = 1, nxt1 fld(j, ng + nyl + k) = oh_buf_recv_north((k - 1)*nxt1 + j) end do end do end if end if end if end subroutine ocean_halo_face_x_2d_impl