From c53a51da81811627bba14224c0e17511fc585488 Mon Sep 17 00:00:00 2001 From: Erik van Sebille Date: Thu, 6 Aug 2026 19:49:28 +0200 Subject: [PATCH 1/5] Fixing Nearest Neighbour interpolation --- src/parcels/interpolators/_xinterpolators.py | 37 ++++++-------------- 1 file changed, 11 insertions(+), 26 deletions(-) diff --git a/src/parcels/interpolators/_xinterpolators.py b/src/parcels/interpolators/_xinterpolators.py index f80f2241a..390a655c5 100644 --- a/src/parcels/interpolators/_xinterpolators.py +++ b/src/parcels/interpolators/_xinterpolators.py @@ -530,41 +530,26 @@ def interp( # Spatial coordinates: left if barycentric < 0.5, otherwise right zi_1 = np.clip(zi + 1, 0, data.shape[1] - 1) - zi_full = np.where(zeta < 0.5, zi, zi_1) + zi_full = np.where(zeta <= 0.5, zi, zi_1) yi_1 = np.clip(yi + 1, 0, data.shape[2] - 1) - yi_full = np.where(eta < 0.5, yi, yi_1) + yi_full = np.where(eta <= 0.5, yi, yi_1) xi_1 = np.clip(xi + 1, 0, data.shape[3] - 1) - xi_full = np.where(xsi < 0.5, xi, xi_1) + xi_full = np.where(xsi <= 0.5, xi, xi_1) - # Time coordinates: 1 point at ti, then 1 point at ti+1 - if lenT == 1: - ti_full = ti - else: - ti_1 = np.clip(ti + 1, 0, data.shape[0] - 1) - ti_full = np.concatenate([ti, ti_1]) - xi_full = np.repeat(xi_full, 2) - yi_full = np.repeat(yi_full, 2) - zi_full = np.repeat(zi_full, 2) - - # Create DataArrays for indexing - selection_dict = { - axis_dim["X"]: xr.DataArray(xi_full, dims=("points")), - axis_dim["Y"]: xr.DataArray(yi_full, dims=("points")), + levels: dict[ptyping.XgcmAxisDirection, tuple[np.ndarray, ...]] = { + "T": (ti,) if lenT == 1 else (ti, np.clip(ti + 1, 0, data.shape[0] - 1)), + "Z": (zi_full,), + "Y": (yi_full,), + "X": (xi_full,), } - if "Z" in axis_dim: - selection_dict[axis_dim["Z"]] = xr.DataArray(zi_full, dims=("points")) - if "time" in data.dims: - selection_dict["time"] = xr.DataArray(ti_full, dims=("points")) - - corner_data = data.isel(selection_dict).data.reshape(lenT, len(xsi)) + corner_data = _gather_corners(data, axis_dim, levels, len(xsi))[:, 0, 0, 0] if lenT == 2: - value = corner_data[0, :] * (1 - tau) + corner_data[1, :] * tau + value = corner_data[0] * (1 - tau) + corner_data[1] * tau else: - value = corner_data[0, :] - + value = corner_data[0] return value.compute() if is_dask_collection(value) else value From ef81fe4f0b29a1f3359293d0c430acf314af70fb Mon Sep 17 00:00:00 2001 From: Erik van Sebille Date: Thu, 6 Aug 2026 19:50:06 +0200 Subject: [PATCH 2/5] Testing all interpolators --- tests/test_interpolation.py | 36 ++++++++++++++++++++++++++++-------- 1 file changed, 28 insertions(+), 8 deletions(-) diff --git a/tests/test_interpolation.py b/tests/test_interpolation.py index 37f766c40..b06ccf9a3 100644 --- a/tests/test_interpolation.py +++ b/tests/test_interpolation.py @@ -17,12 +17,14 @@ from parcels._core.mesh import get_mesh from parcels._datasets.structured.generated import simple_UV_dataset from parcels.interpolators import ( + CGrid_Velocity, XFreeslip, XLinear, XLinearInvdistLandTracer, XNearest, XPartialslip, ) +from parcels.interpolators._base import VectorInterpolator from parcels.interpolators._xinterpolators import _get_corner_data_Agrid from parcels.kernels import AdvectionRK4_3D from tests.utils import TEST_DATA @@ -275,12 +277,32 @@ def test_corner_gather_keeps_axes_missing_from_the_mapping(): assert out[0, 0, 0, 0, p] == data.values[ti[p], 0, yi[p], xi[p]] -interp_methods = { - "linear": XLinear, -} +class XNearest_Velocity(VectorInterpolator): # noqa: N801 + """Nearest-Neighbour interpolation on a regular grid for VectorFields of velocity.""" + def interp( + self, + particle_positions: dict[str, float | np.ndarray], + grid_positions, + vectorfield: VectorField, + ): + """Nearest-Neighbour interpolation on a regular grid for VectorFields of velocity.""" + _xnearest = XNearest() + u = _xnearest.interp(particle_positions, grid_positions, vectorfield.U) + v = _xnearest.interp(particle_positions, grid_positions, vectorfield.V) + w = _xnearest.interp(particle_positions, grid_positions, vectorfield.W) + return u, v, w -@pytest.mark.parametrize(("interp_name", "interp_method"), [("linear", XLinear)]) + +@pytest.mark.parametrize( + ("interp_name", "interp_method"), + [ + ("linear", XLinear), + ("freeslip", XFreeslip), + ("nearest", XNearest_Velocity), + ("cgrid_velocity", CGrid_Velocity), + ], +) def test_interp_regression_v3(interp_name, interp_method): """Test that the v4 versions of the interpolation are the same as the v3 versions.""" ds_input = xr.open_dataset(str(TEST_DATA / f"test_interpolation_data_random_{interp_name}.nc")) @@ -320,9 +342,8 @@ def test_interp_regression_v3(interp_name, interp_method): ) fieldset = FieldSet.from_sgrid_conventions(ds, mesh="flat") - assert isinstance(fieldset.U.interp_method, interp_method) - assert isinstance(fieldset.V.interp_method, interp_method) - assert isinstance(fieldset.W.interp_method, interp_method) + if interp_name in ["cgrid_velocity", "freeslip", "nearest"]: + fieldset.UVW.interp_method = interp_method() x, y, z = np.meshgrid(np.linspace(0, 1, 7), np.linspace(0, 1, 13), np.linspace(0, 1, 5)) @@ -341,7 +362,6 @@ def DeleteParticle(particles, fieldset): output_file=outfile, ) - print(str(TEST_DATA / f"test_interpolation_jit_{interp_name}.zarr")) ds_v3 = xr.open_zarr(str(TEST_DATA / f"test_interpolation_jit_{interp_name}.zarr")) ds_v4 = read_particlefile(f"test_interpolation_v4_{interp_name}.parquet") From 5c01b3b11e002fbec5a4a3d7c58cb00bc3f92a13 Mon Sep 17 00:00:00 2001 From: Erik van Sebille Date: Wed, 12 Aug 2026 13:37:23 +0200 Subject: [PATCH 3/5] Matching index search to match v3 behaviour This matters for C-grid interpolation, which is non-continuous --- src/parcels/_core/index_search.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/parcels/_core/index_search.py b/src/parcels/_core/index_search.py index 77e43aff3..10cf50803 100644 --- a/src/parcels/_core/index_search.py +++ b/src/parcels/_core/index_search.py @@ -44,7 +44,7 @@ def _search_1d_array( # TODO v4: We probably rework this to deal with 0D arrays before this point (as we already know field dimensionality) if len(arr) < 2: return np.zeros(shape=x.shape, dtype=np.int32), np.zeros_like(x) - index = np.clip(np.searchsorted(arr, x, side="right") - 1, 0, len(arr) - 2) + index = np.clip(np.searchsorted(arr, x, side="left") - 1, 0, len(arr) - 2) # Use broadcasting to avoid repeated array access arr_index = arr[index] arr_next = arr[np.clip(index + 1, 1, len(arr) - 1)] # Ensure we don't go out of bounds From 69a020072e297b49c979617ed96b468b243454a8 Mon Sep 17 00:00:00 2001 From: Erik van Sebille Date: Wed, 12 Aug 2026 13:37:43 +0200 Subject: [PATCH 4/5] Update regression test to fix Cgrid interpolation --- tests/test_interpolation.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/tests/test_interpolation.py b/tests/test_interpolation.py index b06ccf9a3..e0ccc18e5 100644 --- a/tests/test_interpolation.py +++ b/tests/test_interpolation.py @@ -310,6 +310,10 @@ def test_interp_regression_v3(interp_name, interp_method): xdim = ds_input["U"].shape[3] time = [np.timedelta64(int(t), "s") for t in ds_input["time"].values] + # Convert the coordinates to float32 to match v3 behavior. This makes a difference for Cgrid velocity interpolation + for dim in ["lon", "lat", "depth"]: + ds_input[dim] = ds_input[dim].astype(np.float32) + ds = xr.Dataset( { "U": (["time", "depth", "YG", "XG"], ds_input["U"].values), @@ -333,8 +337,8 @@ def test_interp_regression_v3(interp_name, interp_method): topology_dimension=2, node_dimensions=("XG", "YG"), face_dimensions=( - sgrid.FaceNodePadding("XC", "XG", sgrid.Padding.HIGH), - sgrid.FaceNodePadding("YC", "YG", sgrid.Padding.HIGH), + sgrid.FaceNodePadding("XC", "XG", sgrid.Padding.LOW), + sgrid.FaceNodePadding("YC", "YG", sgrid.Padding.LOW), ), node_coordinates=("lon", "lat"), vertical_dimensions=(sgrid.FaceNodePadding("ZC", "depth", sgrid.Padding.HIGH),), From ecda1cea9c7ebbdf895cfa34e902d3ed2143ed13 Mon Sep 17 00:00:00 2001 From: Erik van Sebille Date: Thu, 13 Aug 2026 09:04:16 +0200 Subject: [PATCH 5/5] Using particlefile_to_v3_zarr for regression comparison --- tests/conftest.py | 5 ++++ tests/test_interpolation.py | 55 ++++++++----------------------------- 2 files changed, 17 insertions(+), 43 deletions(-) diff --git a/tests/conftest.py b/tests/conftest.py index 71a3740e7..99d041446 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -22,6 +22,11 @@ def tmp_parquet(tmp_path): return tmp_path / "tmp.parquet" +@pytest.fixture +def tmp_zarr(tmp_path): + return tmp_path / "tmp.zarr" + + @pytest.fixture def fieldset() -> FieldSet: """FieldSet with U and V""" diff --git a/tests/test_interpolation.py b/tests/test_interpolation.py index e0ccc18e5..cfb690916 100644 --- a/tests/test_interpolation.py +++ b/tests/test_interpolation.py @@ -11,7 +11,7 @@ StatusCode, Variable, VectorField, - read_particlefile, + particlefile_to_v3_zarr, ) from parcels._core.index_search import _search_time_index from parcels._core.mesh import get_mesh @@ -303,7 +303,7 @@ def interp( ("cgrid_velocity", CGrid_Velocity), ], ) -def test_interp_regression_v3(interp_name, interp_method): +def test_interp_regression_v3(interp_name, interp_method, tmp_zarr, tmp_parquet): """Test that the v4 versions of the interpolation are the same as the v3 versions.""" ds_input = xr.open_dataset(str(TEST_DATA / f"test_interpolation_data_random_{interp_name}.nc")) ydim = ds_input["U"].shape[2] @@ -358,52 +358,21 @@ def DeleteParticle(particles, fieldset): any_error = particles.state >= 50 # This captures all Errors particles[any_error].state = StatusCode.Delete - outfile = ParticleFile(f"test_interpolation_v4_{interp_name}.parquet", outputdt=np.timedelta64(1, "s"), mode="w") + outfile = ParticleFile(tmp_parquet, outputdt=np.timedelta64(1, "s"), mode="w") pset.execute( [AdvectionRK4_3D, DeleteParticle], runtime=np.timedelta64(4, "s"), dt=np.timedelta64(1, "s"), output_file=outfile, ) + particlefile_to_v3_zarr(tmp_parquet, tmp_zarr) + ds_v4 = xr.open_zarr(tmp_zarr) ds_v3 = xr.open_zarr(str(TEST_DATA / f"test_interpolation_jit_{interp_name}.zarr")) - ds_v4 = read_particlefile(f"test_interpolation_v4_{interp_name}.parquet") - - v3_starts = np.column_stack([ds_v3.lon[:, 0].values, ds_v3.lat[:, 0].values, ds_v3.z[:, 0].values]) - unique_starts_v3, inverse_indices = np.unique(v3_starts, axis=0, return_inverse=True) - - v4_pid_to_data = {pid: ds_v4.filter(ds_v4["particle_id"] == pid) for pid in ds_v4["particle_id"].unique()} - - for start_lon, start_lat, start_z in unique_starts_v3: - # Find particles in v3 with this starting position - ind_v3 = np.where( - inverse_indices - == np.where( - (unique_starts_v3[:, 0] == start_lon) - & (unique_starts_v3[:, 1] == start_lat) - & (unique_starts_v3[:, 2] == start_z) - )[0][0] - )[0][0] - - # Find particles in v4 with this starting position using vectorized filter - v4_mask = (ds_v4["x"] == start_lon) & (ds_v4["y"] == start_lat) & (ds_v4["z"] == start_z) - ind_v4 = ds_v4.filter(v4_mask)["particle_id"].unique().to_numpy() - - v3_lon = ds_v3.lon[ind_v3, :].values - v3_lat = ds_v3.lat[ind_v3, :].values - v3_z = ds_v3.z[ind_v3, :].values - - # Use cached v4 data - v4_data = v4_pid_to_data[ind_v4[0]] - v4_lon = v4_data["x"].to_numpy()[:-1] - v4_lat = v4_data["y"].to_numpy()[:-1] - v4_z = v4_data["z"].to_numpy()[:-1] - - # Skip if all NaN - if np.all(np.isnan(v3_lon)) or np.all(np.isnan(v4_lon)): - continue - - tol = 1e-6 - np.testing.assert_allclose(v3_lon, v4_lon, atol=tol) - np.testing.assert_allclose(v3_lat, v4_lat, atol=tol) - np.testing.assert_allclose(v3_z, v4_z, atol=tol) + # v3 zarr is not sorted by particle_id, so we sort it here to match the v4 output + ds_v3 = ds_v3.sortby("trajectory") + + # v4 also writes last timestep, so has one more observation than v3. We ignore the last timestep to match v3. + np.testing.assert_allclose(ds_v3["lon"].values, ds_v4["lon"].values[:, :-1], atol=1e-6, equal_nan=True) + np.testing.assert_allclose(ds_v3["lat"].values, ds_v4["lat"].values[:, :-1], atol=1e-6, equal_nan=True) + np.testing.assert_allclose(ds_v3["z"].values, ds_v4["z"].values[:, :-1], atol=1e-6, equal_nan=True)