State-level shim routing ML prognostic ghost exchange through the O1 comm primitives.
This module provides a single call site for the outer-step ghost fill of
the ocean multilayer C-grid prognostic fields. The routing follows D0
(MPI-agnostic solver code): on a single-rank non-periodic build the calls
resolve to no-ops; on a single-rank periodic build they resolve to local
wrap copies (handled inside rdb_ocean_halo); on a multi-rank build they
resolve to messages. Corner ghosts are valid after the call because the
O1 primitives perform two-pass (E/W then N/S) exchanges internally (D2).
The per-tracer loop is OUTSIDE any do concurrent region (outer-shim
pattern: array-of-derived-types cannot be dereferenced on-device).
!! State-level shim routing ML prognostic ghost exchange through the O1 !! comm primitives. !! !! This module provides a single call site for the outer-step ghost fill of !! the ocean multilayer C-grid prognostic fields. The routing follows D0 !! (MPI-agnostic solver code): on a single-rank non-periodic build the calls !! resolve to no-ops; on a single-rank periodic build they resolve to local !! wrap copies (handled inside `rdb_ocean_halo`); on a multi-rank build they !! resolve to messages. Corner ghosts are valid after the call because the !! O1 primitives perform two-pass (E/W then N/S) exchanges internally (D2). !! !! The per-tracer loop is OUTSIDE any `do concurrent` region (outer-shim !! pattern: array-of-derived-types cannot be dereferenced on-device). module rdb_ocean_halo_state use rdb_ocean_halo, only: ocean_halo_centre, & ocean_halo_face_x, & ocean_halo_face_y, & ocean_halo_is_decomposed_x, & ocean_halo_is_decomposed_y use rdb_ocean_halo_counters, only: oh_count_ml_state, & oh_count_suppress_on, oh_count_suppress_off use rdb_profiler, only: profiler_start, profiler_stop use rdb_multilayer_state, only: multilayer_state_t use rdb_grid, only: hgrid_t use rdb_ocean_boundary_types, only: ocean_bc_state_t use rdb_ocean_periodic, only: ocean_periodic_wrap_face_x_2d, & ocean_periodic_wrap_face_y_2d use rdb_ocean_fold_apply, only: ocean_fold_wrap_stress, ocean_fold_wrap_centre_flat use rdb_ocean_surface_stress, only: ocean_surface_stress_t, & ocean_surface_stress_set_derived use rdb_ice_state, only: ocean_sea_ice_t use rdb_constants, only: wp implicit none private public :: ocean_halo_exchange_ml_state public :: ocean_seam_refresh_surface_stress public :: ocean_halo_exchange_ice_state public :: ocean_halo_exchange_ice_fluxes public :: ocean_halo_exchange_ice_transport contains subroutine ocean_halo_exchange_ml_state(ms, device_resident) !! Exchange ghost cells for the four multilayer prognostic field kinds !! via the O1 halo primitives: !! !! * `h_layer` — cell-centred layer thickness (centre_3d) !! * `u_face_x_layer` — east-face layer velocity (face_x_3d) !! * `v_face_y_layer` — north-face layer velocity (face_y_3d) !! * `tracers(it)%hTr` — per-tracer thickness-weighted scalar !! (centre_3d, outer-shim loop) !! !! Unconditional (D0): single-rank + non-periodic ⇒ no-op; !! single-rank + periodic ⇒ local wrap; multi-rank ⇒ messages. !! Corner ghosts valid on return (D2 two-pass inside the primitives). type(multilayer_state_t), intent(inout) :: ms !! Multilayer C-grid state whose ghost bands are to be filled. logical, intent(in), optional :: device_resident !! Forwarded to every primitive call. Pass .false. for init-time !! host-side exchanges that occur before ocean_state_enter_data. !! Default (.true.) is the normal device-resident path. integer :: it call profiler_start("ocean_comms_ml") call oh_count_ml_state() call oh_count_suppress_on() call ocean_halo_centre(ms%h_layer, ms%nz_ml, device_resident) call ocean_halo_face_x(ms%u_face_x_layer, ms%nz_ml, device_resident) call ocean_halo_face_y(ms%v_face_y_layer, ms%nz_ml, device_resident) ! Per-tracer loop outside DC — outer-shim pattern: array-of-DT ! cannot be dereferenced inside a do concurrent on-device. if (allocated(ms%tracers)) then do it = 1, size(ms%tracers) if (.not. allocated(ms%tracers(it)%hTr)) cycle call ocean_halo_centre(ms%tracers(it)%hTr, ms%nz_ml, device_resident) end do end if call oh_count_suppress_off() call profiler_stop("ocean_comms_ml") end subroutine ocean_halo_exchange_ml_state subroutine ocean_seam_refresh_surface_stress(ss, grid, bc, device_resident) !! Make the surface-stress pair valid in every ghost cell, then !! re-derive `stress_mag` from it. !! !! **Why this exists.** `tau_x`/`tau_y` are C-grid face fields that !! several kernels read ONE CELL BEYOND the cell they write: !! !! * `ocean_surfstress_derived_impl` averages `tau_x(i)`+`tau_x(i+1)` !! into the cell-centred `stress_mag`, which feeds KPP/EPBL `u_*`. !! * `mle_face_ustar_x/y` (Fox-Kemper / Bodner) take a 4-point !! corner average reaching `tau_y(i-1, ·)` / `tau_x(·, j-1)`. !! !! At an MPI seam those reads land in ghost cells that belong to the !! neighbour rank, so they MUST come from an exchange. Nothing may !! extrapolate them: a zero-gradient / edge-copy fill silently !! substitutes this rank's edge value for the neighbour's real data, !! which is decomposition-dependent and therefore invisible to any !! single-rank test. (This mirrors the convention in MOM6, which !! halo-exchanges the stress pair and never extrapolates forcing.) !! !! **Order is load-bearing** and matches the prognostic-state path in !! `ocean_dyn_step_split`: exchange, THEN periodic wrap on any axis !! the exchange did not own, THEN the north fold. The periodic !! kernels are skipped on a decomposed axis because the halo already !! filled those ghosts — running both would overwrite correct !! neighbour data with a local wrap (see `rdb_ocean_periodic`). !! !! `stress_mag` needs no wrap or fold of its own: it is recomputed !! LAST, over the full array, from a `tau` pair whose ghosts are by !! then already valid, so its ghosts come out right for free. !! !! Safe to call before `ocean_state_enter_data` with !! `device_resident = .false.` (the configure-time seed path). type(ocean_surface_stress_t), intent(inout) :: ss !! Stress slot whose `tau_x`/`tau_y` ghosts are to be filled and !! whose `stress_mag` is then refreshed. type(hgrid_t), intent(in) :: grid type(ocean_bc_state_t), intent(in) :: bc !! Supplies `periodic_x`/`periodic_y`/`north_fold`. logical, intent(in), optional :: device_resident !! Forwarded to the halo primitives; `.false.` for host-side !! configure-time calls. integer :: nxt, nyt, nxp, nyp, ng logical :: wrap_x, wrap_y if (.not. allocated(ss%tau_x) .or. .not. allocated(ss%tau_y)) return nxt = grid%nx_total nyt = grid%ny_total nxp = grid%nx_phys nyp = grid%ny_phys ng = grid%nghost call profiler_start("ocean_comms_stress") call oh_count_suppress_on() call ocean_halo_face_x(ss%tau_x, device_resident) call ocean_halo_face_y(ss%tau_y, device_resident) call oh_count_suppress_off() call profiler_stop("ocean_comms_stress") ! Local periodic wrap only on an axis the halo did NOT own. wrap_x = bc%periodic_x .and. (.not. ocean_halo_is_decomposed_x()) wrap_y = bc%periodic_y .and. (.not. ocean_halo_is_decomposed_y()) if (wrap_x .or. wrap_y) then call ocean_periodic_wrap_face_x_2d(ss%tau_x, nxt + 1, nyt, & nxp, nyp, ng, wrap_x, wrap_y) call ocean_periodic_wrap_face_y_2d(ss%tau_y, nxt, nyt + 1, & nxp, nyp, ng, wrap_x, wrap_y) end if ! Tripolar north fold. `tau_x`/`tau_y` are TRUE VECTOR components, ! so the sign-flipping u/v-face variants are the correct ones (the ! scalar-copy duplicates in `rdb_ocean_metrics` exist precisely ! because those are NOT vectors). px = 1: the local kernels; px > 1: ! one owner-routed exchange group (`ocean_fold_wrap_stress`). if (bc%north_fold) call ocean_fold_wrap_stress(grid, bc, ss%tau_x, ss%tau_y, device_resident) call ocean_surface_stress_set_derived(grid, ss) end subroutine ocean_seam_refresh_surface_stress subroutine ocean_halo_exchange_ice_state(ice, grid, bc, device_resident) !! Make the sea-ice CATEGORY state valid in every ghost cell (X1 of !! the sea-ice MPI plan): `part_size`, `m_ice`, `m_snow`, !! `enth_ice`, `sal_ice`, `enth_snow`, one two-pass centre exchange !! each, all categories (and ice layers) in one message per !! direction, THEN the tripolar north fold of each (added with the !! fold-seam fix below). !! !! **Why.** Nothing inside the ice step refreshes these ghosts — the !! column thermodynamics, ITD and (PR 4b) transport compress write !! PHYSICAL cells only — yet three consumers read them one cell into !! the halo: the EVP's category gather (`mis`/`mice`/`ci` over the !! full array, then strength, face mass and corner ratios), and the !! stress coupler's face concentration `a_u = (ci(i-1)+ci(i))/2` at !! the west/south-most owned face. On one rank with a periodic axis !! the primitives' local wrap closes the seam the same way (the !! pre-existing "D7" stale-ghost note in `rdb_ice_evp`). !! !! **The fold.** On a `north = 'tripolar_fold'` grid the MPI/periodic !! exchange above fills every ghost EXCEPT the north cap: the fold !! seam needs its own 180-degree-rotated mirror (`rdb_ocean_fold`), !! which `ocean_halo_centre` knows nothing about. Before this fix !! the north-fold ghost band of every one of these fields was stale !! (uninitialised / previous-step), so the ITD/transport readers one !! cell into it on the fold-seam row effectively saw an unrelated !! cell — the root cause of unbounded ice growth on the seam row !! (a mass source with no physical origin). Every field here is a !! per-category SCALAR (mass, enthalpy, salinity, fractional area), !! never a vector, so the fold is always the `negate=.false.` copy !! contract (`ocean_fold_wrap_centre_flat`). !! !! **When.** At the end of every thermo block (the category state !! changes only there) and once at cold-start configure, host-side, !! BEFORE `enter_data` and NEVER on a warm restart (the checkpoint !! carries the writer's ghosts; re-deriving them resumed a different !! state, `8e1931f20`). Single-rank non-periodic: a no-op. Requires !! `ocean_halo_init`; the caller skips it otherwise. The fold itself !! is gated independently on `bc%north_fold` and no-ops on any !! non-tripolar grid. !! !! The rank-4 `enth_ice`/`sal_ice`/`enth_snow` and the `0:ncat` !! `part_size` are contiguous and go out as one flat `nz` each, !! through an explicit-shape seam (`ice_halo_centre_flat`), never !! the aggregate `ice` (CLAUDE.md: component arrays only). type(ocean_sea_ice_t), intent(inout) :: ice !! Live sea-ice slot (`ice%is_init`); a no-op otherwise. type(hgrid_t), intent(in) :: grid type(ocean_bc_state_t), intent(in) :: bc !! Supplies `north_fold` + the grid metrics the fold kernels need. logical, intent(in), optional :: device_resident !! Forwarded to the halo primitives; `.false.` for the host-side !! configure-time call. integer :: nxt, nyt if (.not. ice%is_init) return nxt = ice%nx_total nyt = ice%ny_total call profiler_start("ice_comms_state") call oh_count_suppress_on() call ice_halo_centre_flat(ice%part_size, nxt, nyt, ice%ncat + 1, device_resident) call ice_halo_centre_flat(ice%m_ice, nxt, nyt, ice%ncat, device_resident) call ice_halo_centre_flat(ice%m_snow, nxt, nyt, ice%ncat, device_resident) call ice_halo_centre_flat(ice%enth_ice, nxt, nyt, ice%ncat*ice%nk_ice, device_resident) call ice_halo_centre_flat(ice%sal_ice, nxt, nyt, ice%ncat*ice%nk_ice, device_resident) call ice_halo_centre_flat(ice%enth_snow, nxt, nyt, ice%ncat, device_resident) call oh_count_suppress_off() call profiler_stop("ice_comms_state") ! Tripolar north fold. Every field above is a per-category ! SCALAR (mass / enthalpy / salinity / fractional area), never a ! vector, so each fold is the plain copy (`negate=.false.`) ! contract — see `ocean_fold_wrap_centre_flat`. No-op off a ! tripolar grid (`bc%north_fold = .false.`). if (bc%north_fold) then call ocean_fold_wrap_centre_flat(grid, bc, ice%part_size, nxt, nyt, & ice%ncat + 1, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%m_ice, nxt, nyt, & ice%ncat, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%m_snow, nxt, nyt, & ice%ncat, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%enth_ice, nxt, nyt, & ice%ncat*ice%nk_ice, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%sal_ice, nxt, nyt, & ice%ncat*ice%nk_ice, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%enth_snow, nxt, nyt, & ice%ncat, device_resident) end if end subroutine ocean_halo_exchange_ice_state subroutine ocean_halo_exchange_ice_fluxes(ice, grid, bc, device_resident) !! Seam ghosts of the three per-cell ice->ocean flux diagnostics the !! couplers hand to the ocean — `salt_flux_diag`, `heat_flux_diag`, !! `sw_thru_diag` — one two-pass centre exchange each. !! !! **Why.** The column driver, frazil uptake and snowfall share write !! them on PHYSICAL cells only, but the brine / heat / shortwave !! couplers copy them over the FULL array into `Q_salt` / `Q_heat` / !! `q_sw` (or their components), and the ocean's surface-flux !! application reads a seam ghost before its next exchange. A stale !! ghost there is the neighbour's flux replaced by this tile's old !! one (measured: the first decomposed run diverged from the serial !! one in the outer step after the first thermo block). Call after !! the last contributor and before the couplers. Single-rank !! non-periodic: a no-op; requires `ocean_halo_init`. Also carries !! the tripolar north fold of the same three fields (plain-copy !! scalar contract) — the fold-seam fix's third exchange site. type(ocean_sea_ice_t), intent(inout) :: ice !! Live sea-ice slot (`ice%is_init`); a no-op otherwise. type(hgrid_t), intent(in) :: grid type(ocean_bc_state_t), intent(in) :: bc !! Supplies `north_fold` + the grid metrics the fold kernels need. logical, intent(in), optional :: device_resident !! Forwarded to the halo primitives. if (.not. ice%is_init) return call profiler_start("ice_comms_fluxes") call oh_count_suppress_on() call ocean_halo_centre(ice%salt_flux_diag, device_resident) call ocean_halo_centre(ice%heat_flux_diag, device_resident) call ocean_halo_centre(ice%sw_thru_diag, device_resident) call oh_count_suppress_off() call profiler_stop("ice_comms_fluxes") if (bc%north_fold) then call ocean_fold_wrap_centre_flat(grid, bc, ice%salt_flux_diag, & size(ice%salt_flux_diag, 1), & size(ice%salt_flux_diag, 2), 1, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%heat_flux_diag, & size(ice%heat_flux_diag, 1), & size(ice%heat_flux_diag, 2), 1, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%sw_thru_diag, & size(ice%sw_thru_diag, 1), & size(ice%sw_thru_diag, 2), 1, device_resident) end if end subroutine ocean_halo_exchange_ice_fluxes subroutine ocean_halo_exchange_ice_transport(ice, grid, bc, device_resident) !! X4 of the sea-ice MPI plan: the seam ghosts every advective !! substep of `ice_transport_step` reads — the cell-averaged !! category masses `mca_ice`/`mca_snow` (the PPM donors, 5-point !! stencil) and the riding intensive tracers `m_ice`, `enth_ice`, !! `sal_ice`, `enth_snow` (the PCM donors). `mca_*` ghosts are !! zeroed by the IST->CAS conversion and the ride/mass updates leave !! the ghost band one substep old, so this runs at the top of EVERY !! substep. On one rank with a periodic axis the primitives wrap. !! Also carries the tripolar north fold of the same six fields !! (plain-copy scalar contract) — without it the fold-seam row's !! advective stencil read a stale/unrelated mirror cell every !! substep, which is how ice piled up without bound on that row. type(ocean_sea_ice_t), intent(inout) :: ice !! Live sea-ice slot (`ice%is_init`); a no-op otherwise. type(hgrid_t), intent(in) :: grid type(ocean_bc_state_t), intent(in) :: bc !! Supplies `north_fold` + the grid metrics the fold kernels need. logical, intent(in), optional :: device_resident !! Forwarded to the halo primitives. integer :: nxt, nyt if (.not. ice%is_init) return nxt = ice%nx_total nyt = ice%ny_total call profiler_start("ice_comms_transport") call oh_count_suppress_on() call ice_halo_centre_flat(ice%mca_ice, nxt, nyt, ice%ncat, device_resident) call ice_halo_centre_flat(ice%mca_snow, nxt, nyt, ice%ncat, device_resident) call ice_halo_centre_flat(ice%m_ice, nxt, nyt, ice%ncat, device_resident) call ice_halo_centre_flat(ice%enth_ice, nxt, nyt, ice%ncat*ice%nk_ice, device_resident) call ice_halo_centre_flat(ice%sal_ice, nxt, nyt, ice%ncat*ice%nk_ice, device_resident) call ice_halo_centre_flat(ice%enth_snow, nxt, nyt, ice%ncat, device_resident) call oh_count_suppress_off() call profiler_stop("ice_comms_transport") if (bc%north_fold) then call ocean_fold_wrap_centre_flat(grid, bc, ice%mca_ice, nxt, nyt, & ice%ncat, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%mca_snow, nxt, nyt, & ice%ncat, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%m_ice, nxt, nyt, & ice%ncat, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%enth_ice, nxt, nyt, & ice%ncat*ice%nk_ice, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%sal_ice, nxt, nyt, & ice%ncat*ice%nk_ice, device_resident) call ocean_fold_wrap_centre_flat(grid, bc, ice%enth_snow, nxt, nyt, & ice%ncat, device_resident) end if end subroutine ocean_halo_exchange_ice_transport subroutine ice_halo_centre_flat(fld, nxt, nyt, nz, device_resident) !! Explicit-shape seam: a contiguous ice array of any rank (`0:ncat` !! third bound, rank-4 category x layer) is handed in by sequence !! association and exchanged as one `(nxt, nyt, nz)` centre field. !! The generic `ocean_halo_centre` resolves on the DUMMY's rank, so !! the rank-4 actuals cannot call it directly. integer, intent(in) :: nxt, nyt, nz real(wp), intent(inout) :: fld(nxt, nyt, nz) logical, intent(in), optional :: device_resident call ocean_halo_centre(fld, nz, device_resident) end subroutine ice_halo_centre_flat end module rdb_ocean_halo_state