ocean_halo_centre_3d Subroutine

private 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).

Arguments

Type IntentOptional 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.


Calls

proc~~ocean_halo_centre_3d~~CallsGraph proc~ocean_halo_centre_3d ocean_halo_centre_3d comm_irecv_real_sp_array_n comm_irecv_real_sp_array_n proc~ocean_halo_centre_3d->comm_irecv_real_sp_array_n comm_isend_real_sp_array_n comm_isend_real_sp_array_n proc~ocean_halo_centre_3d->comm_isend_real_sp_array_n proc~comm_env_compute_comm comm_env_compute_comm proc~ocean_halo_centre_3d->proc~comm_env_compute_comm proc~ew_rank_east ew_rank_east proc~ocean_halo_centre_3d->proc~ew_rank_east proc~ew_rank_west ew_rank_west proc~ocean_halo_centre_3d->proc~ew_rank_west proc~needs_flags needs_flags proc~ocean_halo_centre_3d->proc~needs_flags proc~ns_rank_north ns_rank_north proc~ocean_halo_centre_3d->proc~ns_rank_north proc~ns_rank_south ns_rank_south proc~ocean_halo_centre_3d->proc~ns_rank_south proc~ocean_halo_buffers_ensure_nz ocean_halo_buffers_ensure_nz proc~ocean_halo_centre_3d->proc~ocean_halo_buffers_ensure_nz proc~ocean_periodic_wrap_centre_3d ocean_periodic_wrap_centre_3d proc~ocean_halo_centre_3d->proc~ocean_periodic_wrap_centre_3d proc~oh_count_centre_3d oh_count_centre_3d proc~ocean_halo_centre_3d->proc~oh_count_centre_3d proc~oh_count_msgs oh_count_msgs proc~ocean_halo_centre_3d->proc~oh_count_msgs waitall waitall proc~ocean_halo_centre_3d->waitall comm_world comm_world proc~comm_env_compute_comm->comm_world proc~decomp_rank_from_coords decomp_rank_from_coords proc~ew_rank_east->proc~decomp_rank_from_coords proc~ew_rank_west->proc~decomp_rank_from_coords proc~ns_rank_north->proc~decomp_rank_from_coords proc~ns_rank_south->proc~decomp_rank_from_coords to_string to_string proc~ocean_halo_buffers_ensure_nz->to_string warning warning proc~ocean_halo_buffers_ensure_nz->warning

Called by

proc~~ocean_halo_centre_3d~~CalledByGraph proc~ocean_halo_centre_3d ocean_halo_centre_3d interface~ocean_halo_centre ocean_halo_centre interface~ocean_halo_centre->proc~ocean_halo_centre_3d proc~configure_ocean_land_mask configure_ocean_land_mask proc~configure_ocean_land_mask->interface~ocean_halo_centre proc~continuity_gm_apply continuity_gm_apply proc~continuity_gm_apply->interface~ocean_halo_centre proc~continuity_tracer_step_split continuity_tracer_step_split proc~continuity_tracer_step_split->interface~ocean_halo_centre proc~engine_setup engine_setup proc~engine_setup->interface~ocean_halo_centre proc~engine_setup->proc~configure_ocean_land_mask proc~ocean_halo_exchange_ml_state ocean_halo_exchange_ml_state proc~engine_setup->proc~ocean_halo_exchange_ml_state proc~ocean_halo_exchange_ice_state ocean_halo_exchange_ice_state proc~engine_setup->proc~ocean_halo_exchange_ice_state proc~ice_halo_centre_flat ice_halo_centre_flat proc~ice_halo_centre_flat->interface~ocean_halo_centre proc~ocean_dyn_step_split ocean_dyn_step_split proc~ocean_dyn_step_split->interface~ocean_halo_centre proc~run_stage_split run_stage_split proc~ocean_dyn_step_split->proc~run_stage_split proc~run_gm_step run_gm_step proc~ocean_dyn_step_split->proc~run_gm_step proc~ocean_halo_exchange_ice_fluxes ocean_halo_exchange_ice_fluxes proc~ocean_halo_exchange_ice_fluxes->interface~ocean_halo_centre proc~ocean_halo_exchange_ml_state->interface~ocean_halo_centre proc~refresh_tracer_ghosts refresh_tracer_ghosts proc~refresh_tracer_ghosts->interface~ocean_halo_centre proc~run_continuity_chain run_continuity_chain proc~run_continuity_chain->interface~ocean_halo_centre proc~run_continuity_chain->proc~continuity_tracer_step_split proc~run_continuity_chain->proc~ocean_halo_exchange_ml_state proc~run_continuity_chain->proc~refresh_tracer_ghosts proc~run_stage_split->interface~ocean_halo_centre proc~run_stage_split->proc~ocean_halo_exchange_ml_state proc~run_stage_split->proc~refresh_tracer_ghosts proc~run_stage_split->proc~run_continuity_chain proc~complete_ocean_create complete_ocean_create proc~complete_ocean_create->proc~engine_setup proc~engine_enter_data engine_enter_data proc~complete_ocean_create->proc~engine_enter_data proc~driver_run_ocean driver_run_ocean proc~driver_run_ocean->proc~engine_setup proc~driver_run_ocean->proc~engine_enter_data proc~engine_step engine_step proc~driver_run_ocean->proc~engine_step proc~engine_step_ice engine_step_ice proc~driver_run_ocean->proc~engine_step_ice proc~driver_validate driver_validate proc~driver_validate->proc~engine_setup proc~engine_enter_data->proc~ocean_halo_exchange_ml_state proc~engine_step->proc~ocean_dyn_step_split proc~ocean_dyn_step ocean_dyn_step proc~engine_step->proc~ocean_dyn_step proc~engine_step_ice->proc~ocean_halo_exchange_ice_fluxes proc~engine_step_ice->proc~ocean_halo_exchange_ice_state proc~ice_transport_step ice_transport_step proc~engine_step_ice->proc~ice_transport_step proc~ocean_halo_exchange_ice_state->proc~ice_halo_centre_flat proc~ocean_halo_exchange_ice_transport ocean_halo_exchange_ice_transport proc~ocean_halo_exchange_ice_transport->proc~ice_halo_centre_flat proc~run_gm_step->proc~continuity_gm_apply proc~run_gm_step->proc~ocean_halo_exchange_ml_state proc~run_stage run_stage proc~run_stage->proc~continuity_tracer_step_split proc~driver_run driver_run proc~driver_run->proc~driver_run_ocean proc~ice_transport_step->proc~ocean_halo_exchange_ice_transport proc~ocean_dyn_step->proc~run_stage proc~rdb_ocean_create_finalize rdb_ocean_create_finalize proc~rdb_ocean_create_finalize->proc~complete_ocean_create proc~rdb_ocean_create_from_string rdb_ocean_create_from_string proc~rdb_ocean_create_from_string->proc~complete_ocean_create proc~rdb_ocean_step rdb_ocean_step proc~rdb_ocean_step->proc~engine_step proc~rdb_ocean_step->proc~engine_step_ice

Variables

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

Source Code

   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