From 4b6ca3a95ad45e8be2f89ed7ed24bf8279ab48af Mon Sep 17 00:00:00 2001 From: RanaPriyansh Date: Mon, 28 Sep 2026 22:51:13 +0530 Subject: [PATCH] fix: keep chunk-cached DataArray.data lazy --- src/parcels/_chunk_cached_array/core.py | 2 +- tests/test_chunk_cached_array.py | 50 +++++++++++++++++++++++++ tests/test_interpolation.py | 12 ++++++ 3 files changed, 63 insertions(+), 1 deletion(-) create mode 100644 tests/test_chunk_cached_array.py diff --git a/src/parcels/_chunk_cached_array/core.py b/src/parcels/_chunk_cached_array/core.py index bded6d1a5..520f48452 100644 --- a/src/parcels/_chunk_cached_array/core.py +++ b/src/parcels/_chunk_cached_array/core.py @@ -74,7 +74,7 @@ def __init__(self, dask_array: dask.array.Array, max_cache_bytes: int) -> None: self._boundaries.append(np.concatenate(([0], np.cumsum(dim_chunks)))) def get_duck_array(self): - return self.array.compute() + return self.array def _raw_vindex(self, *indices: np.ndarray) -> np.ndarray: """Vectorized indexing with chunk caching. diff --git a/tests/test_chunk_cached_array.py b/tests/test_chunk_cached_array.py new file mode 100644 index 000000000..935c7e3e0 --- /dev/null +++ b/tests/test_chunk_cached_array.py @@ -0,0 +1,50 @@ +import dask +import dask.array as da +import numpy as np +import xarray as xr + +from parcels._chunk_cached_array import ChunkCachedArray, wrap_dataset + + +def test_chunk_cached_data_stays_lazy_until_explicit_materialization(): + loaded = [] + + @dask.delayed + def chunk(index): + loaded.append(index) + return np.arange(16 * index, 16 * index + 16).reshape(4, 4) + + array = da.concatenate([da.from_delayed(chunk(i), shape=(4, 4), dtype=int) for i in range(2)]) + dataset = wrap_dataset(xr.Dataset({"value": (("x", "y"), array)}), max_cache_bytes=1024) + + with dask.config.set(scheduler="synchronous"): + data = dataset.value.data + assert isinstance(data, da.Array) + assert loaded == [] + + np.testing.assert_array_equal(dataset.value.values, np.arange(32).reshape(8, 4)) + assert sorted(loaded) == [0, 1] + + loaded.clear() + np.testing.assert_array_equal(np.asarray(dataset.value), np.arange(32).reshape(8, 4)) + assert sorted(loaded) == [0, 1] + + +def test_chunk_cached_vectorized_selection_reuses_cached_chunk(): + loaded = [] + + @dask.delayed + def chunk(index): + loaded.append(index) + return np.arange(16 * index, 16 * index + 16).reshape(4, 4) + + array = da.concatenate([da.from_delayed(chunk(i), shape=(4, 4), dtype=int) for i in range(2)]) + dataset = wrap_dataset(xr.Dataset({"value": (("x", "y"), array)}), max_cache_bytes=1024) + assert isinstance(dataset.value.variable._data, ChunkCachedArray) + indices = {"x": xr.DataArray([0, 1], dims="points"), "y": xr.DataArray([1, 2], dims="points")} + + with dask.config.set(scheduler="synchronous"): + np.testing.assert_array_equal(dataset.value.isel(indices).data, [1, 6]) + assert loaded == [0] + np.testing.assert_array_equal(dataset.value.isel(indices).data, [1, 6]) + assert loaded == [0] diff --git a/tests/test_interpolation.py b/tests/test_interpolation.py index cfb690916..5773ee65f 100644 --- a/tests/test_interpolation.py +++ b/tests/test_interpolation.py @@ -13,6 +13,7 @@ VectorField, particlefile_to_v3_zarr, ) +from parcels._chunk_cached_array import ChunkCachedArray from parcels._core.index_search import _search_time_index from parcels._core.mesh import get_mesh from parcels._datasets.structured.generated import simple_UV_dataset @@ -118,6 +119,17 @@ def test_raw_2d_interpolation(field, interpolator, t, z, y, x, expected): np.testing.assert_equal(value, expected) +def test_linear_interpolation_with_chunk_cached_data(field): + field.model.data = field.model.data.chunk({"time": 1, "depth": 1, "lat": 2, "lon": 2}) + field.model.to_chunk_cached_arrays(max_cache_bytes=1024) + assert isinstance(field.data.variable._data, ChunkCachedArray) + particle_positions = {"time": [0, 1], "z": [0, 0], "lat": [0.49, 0.49], "lon": [0.51, 0.51]} + grid_positions = field.grid.search(particle_positions["z"], particle_positions["lat"], particle_positions["lon"]) + grid_positions.update(_search_time_index(field, particle_positions["time"])) + + np.testing.assert_allclose(field.interp_method.interp(particle_positions, grid_positions, field), [1.49, 6.49]) + + @pytest.mark.parametrize("mesh", ["flat", "spherical"]) @pytest.mark.parametrize( "func, t, z, y, x, expected",