Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
199 changes: 93 additions & 106 deletions src/parcels/interpolators/_xinterpolators.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,61 @@
from parcels._core.xgrid import XGrid


_CORNER_AXES: tuple[ptyping.XgridAxis, ...] = ("T", "Z", "Y", "X")


def _gather_corners(
data: np.ndarray | xr.DataArray,
axis_dim: dict[ptyping.ptyping.XgridAxis, str],
levels: dict[ptyping.XgridAxis, tuple[np.ndarray, ...]],
npart: int,
) -> np.ndarray:
"""Gather field data at the corners bracketing each particle.

The gather is the outer product over ``_CORNER_AXES``, with each axis
contributing one or two levels and the particle index innermost. The local
variable ``shape`` states that layout, and is used both to build the index
arrays and to reshape the result.

Parameters
----------
data : np.ndarray or xr.DataArray
Field data, with dimensions ordered ``(time, Z, Y, X)``.
axis_dim : dict
Maps ``"X"``, ``"Y"`` and ``"Z"`` to dimension names. The ``"T"``
dimension is always named ``"time"``.
levels : dict
Maps an axis to the per-particle index arrays of its corner levels, e.g.
``{"T": (ti, ti + 1), "Y": (yi, yi + 1)}``. An axis absent from
``levels`` contributes a single level. An axis the field has no
dimension for still shapes the result, but is not indexed, so its
corners repeat.
npart : int
Number of particles.

Returns
-------
np.ndarray
Array of shape ``(*counts, npart)``, where ``counts[i]`` is the number
of levels gathered on ``_CORNER_AXES[i]``.
"""
dims = {axis: ("time" if axis == "T" else axis_dim.get(axis)) for axis in _CORNER_AXES}
counts = tuple(len(levels[axis]) if axis in levels else 1 for axis in _CORNER_AXES)
shape = (*counts, npart)

selection_dict = {}
for i, axis in enumerate(_CORNER_AXES):
if axis in levels and dims[axis] is not None and dims[axis] in data.dims:
stacked = np.stack(np.broadcast_arrays(*levels[axis])) # (n_levels, npart)
# Put the level axis in slot i of the corner grid, spread it over the
# other slots, and flatten. This is the inverse of the reshape below.
other_slots = tuple(j for j in range(len(_CORNER_AXES)) if j != i)
in_slot_i = np.expand_dims(stacked, other_slots)
selection_dict[dims[axis]] = xr.DataArray(np.broadcast_to(in_slot_i, shape).reshape(-1), dims="points")

return data.isel(selection_dict).data.reshape(shape)


def _get_corner_data_Agrid(
data: np.ndarray | xr.DataArray,
ti: int,
Expand All @@ -29,40 +84,13 @@ def _get_corner_data_Agrid(
axis_dim: dict[ptyping.ptyping.XgridAxis, str],
) -> np.ndarray:
"""Helper function to get the corner data for a given A-grid field and position."""
# Time coordinates: 8 points at ti, then 8 points at ti+1
if lenT == 1:
ti = np.repeat(ti, lenZ * 4)
else:
ti_1 = np.clip(ti + 1, 0, data.shape[0] - 1)
ti = np.concatenate([np.repeat(ti, lenZ * 4), np.repeat(ti_1, lenZ * 4)])

# Z coordinates: 4 points at zi, 4 at zi+1, repeated for both time levels
if lenZ == 1:
zi = np.repeat(zi, lenT * 4)
else:
zi_1 = np.clip(zi + 1, 0, data.shape[1] - 1)
zi = np.tile(np.array([zi, zi, zi, zi, zi_1, zi_1, zi_1, zi_1]).flatten(), lenT)

# Y coordinates: [yi, yi, yi+1, yi+1] for each spatial point, repeated for time/z
yi_1 = np.clip(yi + 1, 0, data.shape[2] - 1)
yi = np.tile(np.array([yi, yi, yi_1, yi_1]).flatten(), lenT * lenZ)

# X coordinates: [xi, xi+1, xi, xi+1] for each spatial point, repeated for time/z
xi_1 = np.clip(xi + 1, 0, data.shape[3] - 1)
xi = np.tile(np.array([xi, xi_1]).flatten(), lenT * lenZ * 2)

# Create DataArrays for indexing
selection_dict = {}
if "X" in axis_dim:
selection_dict[axis_dim["X"]] = xr.DataArray(xi, dims=("points"))
if "Y" in axis_dim:
selection_dict[axis_dim["Y"]] = xr.DataArray(yi, dims=("points"))
if "Z" in axis_dim:
selection_dict[axis_dim["Z"]] = xr.DataArray(zi, dims=("points"))
if "time" in data.dims:
selection_dict["time"] = xr.DataArray(ti, dims=("points"))

return data.isel(selection_dict).data.reshape(lenT, lenZ, 2, 2, npart)
levels = {
"T": (ti,) if lenT == 1 else (ti, np.clip(ti + 1, 0, data.shape[0] - 1)),
"Z": (zi,) if lenZ == 1 else (zi, np.clip(zi + 1, 0, data.shape[1] - 1)),
"Y": (yi, np.clip(yi + 1, 0, data.shape[2] - 1)),
"X": (xi, np.clip(xi + 1, 0, data.shape[3] - 1)),
}
return _gather_corners(data, axis_dim, levels, npart)


def _get_offsets_dictionary(grid: XGrid) -> dict[ptyping.CfAxisSpatial, Literal[1, 0]]:
Expand Down Expand Up @@ -213,39 +241,23 @@ def interp(
py[3], py[0], px[3], px[0], grid._mesh, np.einsum("ij,ji->i", i_u.phi2D_lin(eta, 0.0), py), grid.deg2m
)

def _create_selection_dict(dims, zdir=False):
"""Helper function to create DataArrays for indexing."""
axis_dim = grid.get_axis_dim_mapping(dims)
selection_dict = {
axis_dim["X"]: xr.DataArray(xi_full, dims=("points")),
axis_dim["Y"]: xr.DataArray(yi_full, dims=("points")),
}
npart = len(xsi)
t_levels = (ti,) if lenT == 1 else (ti, np.clip(ti + 1, 0, tdim - 1))

def _compute_corner_data(data, y_levels, x_levels, z_levels=None) -> np.ndarray:
"""Gather the two bracketing face values and reduce over time if needed.

# Time coordinates: 2 points at ti, then 2 points at ti+1
if "time" in dims:
if lenT == 1:
ti_full = np.repeat(ti, 2)
else:
ti_1 = np.clip(ti + 1, 0, tdim - 1)
ti_full = np.concatenate([np.repeat(ti, 2), np.repeat(ti_1, 2)])
selection_dict["time"] = xr.DataArray(ti_full, dims=("points"))

if "Z" in axis_dim:
if zdir:
# Z coordinates: 1 point at zi and 1 point at zi+1 repeated for lenT time levels
zi_0 = np.clip(zi + offsets["Z"], 0, zdim - 1)
zi_1 = np.clip(zi + offsets["Z"] + 1, 0, zdim - 1)
zi_full = np.tile(np.array([zi_0, zi_1]).flatten(), lenT)
else:
# Z coordinates: 2 points at zi, repeated for lenT time levels
zi_full = np.repeat(zi, lenT * 2)
selection_dict[axis_dim["Z"]] = xr.DataArray(zi_full, dims=("points"))

return selection_dict

def _compute_corner_data(data, selection_dict) -> np.ndarray:
"""Helper function to load and reduce corner data over time dimension if needed."""
corner_data = data.isel(selection_dict).data.reshape(lenT, 2, len(xsi))
Exactly one of the Z, Y and X axes contributes the two corners. The
other two contribute a single level each.
"""
levels = {
"T": t_levels,
"Z": z_levels if z_levels is not None else (zi,),
"Y": y_levels,
"X": x_levels,
}
axis_dim = grid.get_axis_dim_mapping(data.dims)
corner_data = _gather_corners(data, axis_dim, levels, npart).reshape(lenT, 2, npart)

if lenT == 2:
tau_full = tau[np.newaxis, :]
Expand All @@ -254,29 +266,19 @@ def _compute_corner_data(data, selection_dict) -> np.ndarray:
corner_data = corner_data[0, :]
return corner_data

# Compute U velocity
# Compute U velocity: the two corners are the X faces
yi_o = np.clip(yi + offsets["Y"], 0, ydim - 1)
yi_full = np.tile(np.array([yi_o, yi_o]).flatten(), lenT)

xi_1 = np.clip(xi + 1, 0, xdim - 1)
xi_full = np.tile(np.array([xi, xi_1]).flatten(), lenT)

selection_dict = _create_selection_dict(U.dims)
corner_data = _compute_corner_data(U, selection_dict)
corner_data = _compute_corner_data(U, y_levels=(yi_o,), x_levels=(xi, xi_1))

U0 = corner_data[0, :] * c4
U1 = corner_data[1, :] * c2
Uvel = (1 - xsi) * U0 + xsi * U1

# Compute V velocity
# Compute V velocity: the two corners are the Y faces
yi_1 = np.clip(yi + 1, 0, ydim - 1)
yi_full = np.tile(np.array([yi, yi_1]).flatten(), lenT)

xi_o = np.clip(xi + offsets["X"], 0, xdim - 1)
xi_full = np.tile(np.array([xi_o, xi_o]).flatten(), lenT)

selection_dict = _create_selection_dict(V.dims)
corner_data = _compute_corner_data(V, selection_dict)
corner_data = _compute_corner_data(V, y_levels=(yi, yi_1), x_levels=(xi_o,))

V0 = corner_data[0, :] * c1
V1 = corner_data[1, :] * c3
Expand Down Expand Up @@ -311,16 +313,12 @@ def _compute_corner_data(data, selection_dict) -> np.ndarray:
if vectorfield.W:
W = vectorfield.W.data

# Y coordinates: yi+offset for each spatial point, repeated for time
# Compute W velocity: the two corners are the Z faces
yi_o = np.clip(yi + offsets["Y"], 0, ydim - 1)
yi_full = np.tile(yi_o, (lenT) * 2)

# X coordinates: xi+offset for each spatial point, repeated for time
xi_o = np.clip(xi + offsets["X"], 0, xdim - 1)
xi_full = np.tile(xi_o, (lenT) * 2)

selection_dict = _create_selection_dict(W.dims, zdir=True)
corner_data = _compute_corner_data(W, selection_dict)
zi_0 = np.clip(zi + offsets["Z"], 0, zdim - 1)
zi_1 = np.clip(zi + offsets["Z"] + 1, 0, zdim - 1)
corner_data = _compute_corner_data(W, y_levels=(yi_o,), x_levels=(xi_o,), z_levels=(zi_0, zi_1))

w = corner_data[0, :] * (1 - zeta) + corner_data[1, :] * zeta
if is_dask_collection(w):
Expand Down Expand Up @@ -364,28 +362,17 @@ def interp(
xi = np.clip(xi + offsets["X"], 0, data.shape[3] - 1)

lenT = 2 if np.any(tau > 0) else 1
npart = len(xi)

if lenT == 2:
ti_1 = np.clip(ti + 1, 0, data.shape[0] - 1)
ti = np.concatenate([np.repeat(ti), np.repeat(ti_1)])
zi = np.tile(zi, (lenT) * 2)
yi = np.tile(yi, (lenT) * 2)
xi = np.tile(xi, (lenT) * 2)

# Create DataArrays for indexing
selection_dict = {
axis_dim["X"]: xr.DataArray(xi, dims=("points")),
axis_dim["Y"]: xr.DataArray(yi, dims=("points")),
levels = {
"T": (ti,) if lenT == 1 else (ti, np.clip(ti + 1, 0, data.shape[0] - 1)),
"Z": (zi,),
"Y": (yi,),
"X": (xi,),
}
if "Z" in axis_dim:
selection_dict[axis_dim["Z"]] = xr.DataArray(zi, dims=("points"))
if "time" in field.data.dims:
selection_dict["time"] = xr.DataArray(ti, dims=("points"))

value = data.isel(selection_dict).data.reshape(lenT, len(xi))
value = _gather_corners(data, axis_dim, levels, npart).reshape(lenT, npart)

if lenT == 2:
tau = tau[:, np.newaxis]
value = value[0, :] * (1 - tau) + value[1, :] * tau
else:
value = value[0, :]
Expand Down
6 changes: 5 additions & 1 deletion tests/test_advection.py
Original file line number Diff line number Diff line change
Expand Up @@ -466,9 +466,13 @@ def test_nemo_3D_curvilinear_fieldset(kernel):
np.testing.assert_allclose([p.z for p in pset], z_initial)
elif kernel == AdvectionRK4_3D:
# TODO check why decimals needs to be so low in RK4_3D (compare to v3)
# These particles sit at different depth levels, so the C-grid gather used to
# mix their zi indices; the previous expectations recorded that. Each value
# below is what the same particle gets when advected on its own.
np.testing.assert_allclose(
[p.z for p in pset],
[0.666162, 0.8667131, 0.92150104, 0.9605109, 0.9577529, 1.0041442, 1.0284728, 1.0033542, 1.2949713, 1.3928112],
[0.66616201, 0.86671311, 0.92108649, 0.95940739, 0.95945358, 1.00413370, 1.02847278, 1.00335419, 1.27260256, 1.38021827],
rtol=1e-6,
) # fmt:skip


Expand Down
73 changes: 73 additions & 0 deletions tests/test_interpolation.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
XNearest,
XPartialslip,
)
from parcels.interpolators._xinterpolators import _get_corner_data_Agrid
from parcels.kernels import AdvectionRK4_3D
from tests.utils import TEST_DATA

Expand Down Expand Up @@ -201,6 +202,78 @@ def test_interpolation_mesh_type(mesh, npart=10):
assert v == 0.0


@pytest.fixture
def corner_gather_data() -> xr.DataArray:
rng = np.random.default_rng(0)
return xr.DataArray(rng.random((6, 5, 7, 8)), dims=("time", "depth", "lat", "lon"))


CORNER_GATHER_AXIS_DIM = {"X": "lon", "Y": "lat", "Z": "depth"}


def _agrid_corners(data, ti, zi, yi, xi, lenT, lenZ, npart): # noqa: N803
return _get_corner_data_Agrid(data, ti, zi, yi, xi, lenT, lenZ, npart, CORNER_GATHER_AXIS_DIM)


@pytest.mark.parametrize("lenT", [1, 2])
@pytest.mark.parametrize("lenZ", [1, 2])
@pytest.mark.parametrize("uniform_clock", [True, False])
def test_corner_gather_batch_matches_single_particle(corner_gather_data, lenT, lenZ, uniform_clock): # noqa: N803
"""Gathering a batch must give each particle what it would get on its own.

Interpolation is per-particle, so batching cannot change a result. The index
arrays and the final reshape must therefore agree on where each particle sits
in the flat gather.
"""
rng = np.random.default_rng(1)
npart = 5
if uniform_clock:
ti, zi = np.full(npart, 2), np.full(npart, 1)
else:
ti, zi = rng.integers(0, 4, npart), rng.integers(0, 3, npart)
yi, xi = rng.integers(0, 5, npart), rng.integers(0, 6, npart)

batch = _agrid_corners(corner_gather_data, ti, zi, yi, xi, lenT, lenZ, npart)

for p in range(npart):
single = _agrid_corners(
corner_gather_data, ti[p : p + 1], zi[p : p + 1], yi[p : p + 1], xi[p : p + 1], lenT, lenZ, 1
)
np.testing.assert_array_equal(batch[..., p], single[..., 0])


def test_corner_gather_axes_are_ordered_t_z_y_x(corner_gather_data):
"""The returned axes must mean (T, Z, Y, X, particle), as the interpolators assume."""
rng = np.random.default_rng(2)
npart = 4
ti, zi = rng.integers(0, 4, npart), rng.integers(0, 3, npart)
yi, xi = rng.integers(0, 5, npart), rng.integers(0, 6, npart)

out = _agrid_corners(corner_gather_data, ti, zi, yi, xi, 2, 2, npart)
raw = corner_gather_data.values

assert out.shape == (2, 2, 2, 2, npart)
for p in range(npart):
for it, iz, iy, ix in np.ndindex(2, 2, 2, 2):
assert out[it, iz, iy, ix, p] == raw[ti[p] + it, zi[p] + iz, yi[p] + iy, xi[p] + ix]


def test_corner_gather_keeps_axes_missing_from_the_mapping():
"""An axis absent from ``axis_dim`` is not indexed, but still shapes the result."""
rng = np.random.default_rng(3)
data = xr.DataArray(rng.random((6, 1, 7, 8)), dims=("time", "depth", "lat", "lon"))
npart = 3
ti = np.array([0, 2, 3])
zi = np.zeros(npart, dtype=int)
yi, xi = rng.integers(0, 5, npart), rng.integers(0, 6, npart)

out = _get_corner_data_Agrid(data, ti, zi, yi, xi, 2, 1, npart, {"X": "lon", "Y": "lat"})

assert out.shape == (2, 1, 2, 2, npart)
for p in range(npart):
assert out[0, 0, 0, 0, p] == data.values[ti[p], 0, yi[p], xi[p]]


interp_methods = {
"linear": XLinear,
}
Expand Down
Loading