subroutine ocean_sponge_apply(grid, bc, ms, dt)
!! Apply momentum relaxation in any sponge-tagged edge band.
!! No-op when no edge is OBC_SPONGE.
type(hgrid_t), intent(in) :: grid
type(ocean_bc_state_t), intent(in) :: bc
type(multilayer_state_t), intent(inout) :: ms
real(wp), intent(in) :: dt
integer :: i, j, k, nz, d
integer :: i0, i1, j0, j1
integer :: wall_face, band
integer :: nx_u, ny_u, nx_v, ny_v
real(wp) :: strength, tau, decay, alpha
real(wp), parameter :: PI = acos(-1.0_wp)
nz = ms%nz_ml
i0 = grid%nghost + 1
i1 = grid%nghost + grid%nx_phys
j0 = grid%nghost + 1
j1 = grid%nghost + grid%ny_phys
! Hoisted array extents for the in-band bounds guards (don't call
! size() on a mapped array inside a do concurrent body).
nx_u = size(ms%u_face_x_layer, 1)
ny_u = size(ms%u_face_x_layer, 2)
nx_v = size(ms%v_face_y_layer, 1)
ny_v = size(ms%v_face_y_layer, 2)
! GPU note: `do concurrent` on the device-resident layer velocities.
! Within one edge block every (d, ·, k) tuple writes a distinct
! element, so the DC is race-free; overlapping corner faces between
! edge blocks get both decays applied sequentially.
! ---- West edge ----
! has_* gate: a subdomain seam never carries a sponge band (O0).
if (bc%west%bc_type == OBC_SPONGE .and. bc%west%sponge_width > 0 .and. bc%has_west) then
band = bc%west%sponge_width
strength = bc%west%sponge_strength
wall_face = grid%nghost + 1
do concurrent(k=1:nz, j=j0:j1, d=0:band - 1) local(alpha, tau, decay, i)
! Cosine ramp: τ peaks at d=0 (outer), tapers to 0 at the
! interior edge of the band.
alpha = 0.5_wp*(1.0_wp + cos(PI*real(d, wp)/real(band, wp)))
tau = strength*alpha
decay = exp(-tau*dt)
! u-face just east of the west wall sits at (wall_face + d + 1)
i = wall_face + d + 1
if (i >= 1 .and. i <= nx_u) then
ms%u_face_x_layer(i, j, k) = decay*ms%u_face_x_layer(i, j, k)
end if
end do
do concurrent(k=1:nz, j=j0:j1 + 1, d=0:band - 1) local(alpha, tau, decay, i)
alpha = 0.5_wp*(1.0_wp + cos(PI*real(d, wp)/real(band, wp)))
tau = strength*alpha
decay = exp(-tau*dt)
i = wall_face + d
if (i >= 1 .and. i <= nx_v) then
ms%v_face_y_layer(i, j, k) = decay*ms%v_face_y_layer(i, j, k)
end if
end do
end if
! ---- East edge ----
if (bc%east%bc_type == OBC_SPONGE .and. bc%east%sponge_width > 0 .and. bc%has_east) then
band = bc%east%sponge_width
strength = bc%east%sponge_strength
wall_face = grid%nghost + grid%nx_phys + 1
do concurrent(k=1:nz, j=j0:j1, d=0:band - 1) local(alpha, tau, decay, i)
alpha = 0.5_wp*(1.0_wp + cos(PI*real(d, wp)/real(band, wp)))
tau = strength*alpha
decay = exp(-tau*dt)
i = wall_face - d - 1
if (i >= 1 .and. i <= nx_u) then
ms%u_face_x_layer(i, j, k) = decay*ms%u_face_x_layer(i, j, k)
end if
end do
do concurrent(k=1:nz, j=j0:j1 + 1, d=0:band - 1) local(alpha, tau, decay, i)
alpha = 0.5_wp*(1.0_wp + cos(PI*real(d, wp)/real(band, wp)))
tau = strength*alpha
decay = exp(-tau*dt)
i = wall_face - d - 1
if (i >= 1 .and. i <= nx_v) then
ms%v_face_y_layer(i, j, k) = decay*ms%v_face_y_layer(i, j, k)
end if
end do
end if
! ---- South edge ----
if (bc%south%bc_type == OBC_SPONGE .and. bc%south%sponge_width > 0 .and. bc%has_south) then
band = bc%south%sponge_width
strength = bc%south%sponge_strength
wall_face = grid%nghost + 1
do concurrent(k=1:nz, i=i0:i1 + 1, d=0:band - 1) local(alpha, tau, decay, j)
alpha = 0.5_wp*(1.0_wp + cos(PI*real(d, wp)/real(band, wp)))
tau = strength*alpha
decay = exp(-tau*dt)
j = wall_face + d
if (j >= 1 .and. j <= ny_u) then
ms%u_face_x_layer(i, j, k) = decay*ms%u_face_x_layer(i, j, k)
end if
end do
do concurrent(k=1:nz, i=i0:i1, d=0:band - 1) local(alpha, tau, decay, j)
alpha = 0.5_wp*(1.0_wp + cos(PI*real(d, wp)/real(band, wp)))
tau = strength*alpha
decay = exp(-tau*dt)
j = wall_face + d + 1
if (j >= 1 .and. j <= ny_v) then
ms%v_face_y_layer(i, j, k) = decay*ms%v_face_y_layer(i, j, k)
end if
end do
end if
! ---- North edge ----
if (bc%north%bc_type == OBC_SPONGE .and. bc%north%sponge_width > 0 .and. bc%has_north) then
band = bc%north%sponge_width
strength = bc%north%sponge_strength
wall_face = grid%nghost + grid%ny_phys + 1
do concurrent(k=1:nz, i=i0:i1 + 1, d=0:band - 1) local(alpha, tau, decay, j)
alpha = 0.5_wp*(1.0_wp + cos(PI*real(d, wp)/real(band, wp)))
tau = strength*alpha
decay = exp(-tau*dt)
j = wall_face - d - 1
if (j >= 1 .and. j <= ny_u) then
ms%u_face_x_layer(i, j, k) = decay*ms%u_face_x_layer(i, j, k)
end if
end do
do concurrent(k=1:nz, i=i0:i1, d=0:band - 1) local(alpha, tau, decay, j)
alpha = 0.5_wp*(1.0_wp + cos(PI*real(d, wp)/real(band, wp)))
tau = strength*alpha
decay = exp(-tau*dt)
j = wall_face - d - 1
if (j >= 1 .and. j <= ny_v) then
ms%v_face_y_layer(i, j, k) = decay*ms%v_face_y_layer(i, j, k)
end if
end do
end if
end subroutine ocean_sponge_apply