diff --git a/src/parcels/_chunk_cached_array/core.py b/src/parcels/_chunk_cached_array/core.py index bded6d1a5..67f5f4a6e 100644 --- a/src/parcels/_chunk_cached_array/core.py +++ b/src/parcels/_chunk_cached_array/core.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, Any import numpy as np from xarray.core.indexing import BasicIndexer, ExplicitlyIndexedNDArrayMixin, OuterIndexer, VectorizedIndexer @@ -48,6 +48,42 @@ def wrap_dataset(ds: xr.Dataset, max_cache_bytes: int) -> xr.Dataset: return ds +def _is_equal_length_1d(key: tuple[Any, ...]) -> bool: + """Return whether every indexer is a 1D integer array of the same length.""" + if not key: + return False + length: int | None = None + for item in key: + if not isinstance(item, np.ndarray) or item.ndim != 1: + return False + item_length = int(item.shape[0]) + if length is None: + length = item_length + elif item_length != length: + return False + return True + + +def _raise_if_out_of_bounds(key: tuple[Any, ...], shape: tuple[int, ...]) -> None: + """Raise IndexError if an integer indexer points outside its axis. + + Parameters + ---------- + key : tuple + Per-axis indexer. Slices are ignored; integer arrays are checked. + shape : tuple of int + Length of each axis of the array being indexed. + """ + for axis, item in enumerate(key): + if not isinstance(item, np.ndarray): + continue + size = shape[axis] + out_of_bounds = (item < -size) | (item >= size) + if item.size and np.any(out_of_bounds): + bad = int(np.ravel(item)[int(np.flatnonzero(out_of_bounds)[0])]) + raise IndexError(f"index {bad} is out of bounds for axis {axis} with size {size}") + + class ChunkCachedArray(ExplicitlyIndexedNDArrayMixin): """Chunk-level LRU cache on top of a dask array for vectorized indexing. @@ -58,9 +94,13 @@ class ChunkCachedArray(ExplicitlyIndexedNDArrayMixin): underlying dask array. On each vectorized index: - 1. Maps global indices -> (chunk_coord, local_index) per dimension. - 2. Fetches missing chunks via dask_array.blocks[...].compute(). - 3. Assembles the result from cached numpy arrays. + 1. Expands mixed slice and array keys into flat per-dimension indices. + Integer indexers are broadcast, and slice axes follow those broadcast + axes (xarray's vectorized-indexing order). Equal-length 1D integer + keys skip this step. + 2. Maps global indices -> (chunk_coord, local_index) per dimension. + 3. Fetches missing chunks via dask_array.blocks[...].compute(). + 4. Assembles the result from cached numpy arrays, reshaping mixed keys. """ def __init__(self, dask_array: dask.array.Array, max_cache_bytes: int) -> None: @@ -87,11 +127,21 @@ def _raw_vindex(self, *indices: np.ndarray) -> np.ndarray: Returns ------- np.ndarray - 1D array of length N with the selected values. + 1D array of length N with the selected values. An empty selection + returns an empty array without fetching a chunk. + + Raises + ------ + IndexError + If an index is smaller than ``-size`` or greater than or equal to + ``size`` on its axis. """ ndim = len(self.array.chunks) assert len(indices) == ndim + _raise_if_out_of_bounds(indices, tuple(int(size) for size in self.array.shape)) n_points = len(indices[0]) + if n_points == 0: + return np.empty(0, dtype=self.array.dtype) # Step 1: Map global indices to chunk coords and local indices. # Normalize negative indices (e.g. -1 → last element) to positive, @@ -139,8 +189,90 @@ def _raw_vindex(self, *indices: np.ndarray) -> np.ndarray: # --- ExplicitlyIndexed protocol --- def _vindex_get(self, indexer: VectorizedIndexer): + """Gather values for an xarray vectorized indexer. + + Parameters + ---------- + indexer : VectorizedIndexer + One slice or integer array per dimension. Equal-length 1D integer + keys are gathered directly. Mixed keys are broadcast, each slice is + expanded with ``slice.indices`` onto an axis after those broadcast + axes, and the gathered values are reshaped to that result. + + Returns + ------- + np.ndarray + Selected values in xarray's vectorized-indexing order. + """ key = indexer.tuple - return self._raw_vindex(*key) + if _is_equal_length_1d(key): + return self._raw_vindex(*key) + + # Reject before expanding slices or touching a chunk. Valid negative + # indices stay negative here; ``_raw_vindex`` normalizes them. + _raise_if_out_of_bounds(key, tuple(int(size) for size in self.array.shape)) + flat_indices, result_shape = self._flatten_mixed_vindex(key) + if 0 in result_shape: + return np.empty(result_shape, dtype=self.array.dtype) + return self._raw_vindex(*flat_indices).reshape(result_shape) + + def _flatten_mixed_vindex(self, key: tuple[Any, ...]) -> tuple[tuple[np.ndarray, ...], tuple[int, ...]]: + """Broadcast integer indexers and expand slices to flat coordinates. + + Parameters + ---------- + key : tuple + One slice or integer array per dimension. + + Returns + ------- + tuple + ``(flat_indices, result_shape)``. ``flat_indices`` is empty when + ``result_shape`` contains a zero-size axis. Otherwise it holds one + 1D coordinate array per dimension, raveled in ``result_shape`` order: + the broadcast integer-index shape, then one axis per slice. + """ + shape = tuple(int(size) for size in self.array.shape) + if len(key) != len(shape): + raise ValueError(f"vectorized indexer length {len(key)} does not match ndim {len(shape)}") + + array_shapes: list[tuple[int, ...]] = [] + slice_lengths: list[int] = [] + normalized_slices: list[tuple[int, int, int]] = [] + for axis, item in enumerate(key): + if isinstance(item, np.ndarray): + array_shapes.append(tuple(int(size) for size in item.shape)) + elif isinstance(item, slice): + start, stop, step = item.indices(shape[axis]) + normalized_slices.append((start, stop, step)) + slice_lengths.append(len(range(start, stop, step))) + else: + raise TypeError(f"unsupported vectorized indexer type: {type(item)!r}") + + broadcast_shape = np.broadcast_shapes(*array_shapes) if array_shapes else () + result_shape = broadcast_shape + tuple(slice_lengths) + if 0 in result_shape: + return (), result_shape + + # Place each slice on its own axis after the broadcast index axes. + n_slices = len(slice_lengths) + expanded: list[np.ndarray] = [] + slice_axis = 0 + for item in key: + if isinstance(item, np.ndarray): + reshaped = np.reshape(item, tuple(int(size) for size in item.shape) + (1,) * n_slices) + expanded.append(np.broadcast_to(reshaped, broadcast_shape + (1,) * n_slices)) + continue + start, stop, step = normalized_slices[slice_axis] + positions = np.arange(start, stop, step) + slice_shape = ( + (1,) * (len(broadcast_shape) + slice_axis) + (int(positions.size),) + (1,) * (n_slices - slice_axis - 1) + ) + expanded.append(positions.reshape(slice_shape)) + slice_axis += 1 + + flat_indices = tuple(np.reshape(idx, -1) for idx in np.broadcast_arrays(*expanded)) + return flat_indices, result_shape def _oindex_get(self, indexer: OuterIndexer): # Delegate to dask for orthogonal indexing diff --git a/tests/test_chunk_cached_array.py b/tests/test_chunk_cached_array.py new file mode 100644 index 000000000..16ccee38b --- /dev/null +++ b/tests/test_chunk_cached_array.py @@ -0,0 +1,326 @@ +"""Tests for ChunkCachedArray vectorized indexing, including mixed slices.""" + +from __future__ import annotations + +from collections.abc import Iterator +from contextlib import contextmanager + +import dask +import dask.array as da +import numpy as np +import pytest +import xarray as xr +from dask import is_dask_collection +from dask.callbacks import Callback + +from parcels._chunk_cached_array import ChunkCachedArray, wrap_dataset + +# Dask drops a trailing "-" token when it names fused block tasks, so the +# callback key contains this whole string only when the name has no such suffix. +SOURCE_NAME = "parcelschunksource" +DIMS = ("time", "depth", "lat", "lon") +SHAPE = (5, 4, 7, 6) +CHUNKS = (2, 3, 3, 4) + + +def _points(*values: int, dims: str | tuple[str, ...] = "points") -> xr.DataArray: + return xr.DataArray(np.array(values, dtype=np.int64), dims=dims) + + +def _dataset( + shape: tuple[int, ...] = SHAPE, + chunks: tuple[int, ...] = CHUNKS, + dims: tuple[str, ...] = DIMS, +) -> tuple[xr.Dataset, xr.Dataset, ChunkCachedArray]: + values = np.arange(int(np.prod(shape)), dtype=np.float64).reshape(shape) + darr = da.from_array(values, chunks=chunks, name=SOURCE_NAME) + wrapped = wrap_dataset(xr.Dataset({"data": (dims, darr)}), max_cache_bytes=int(values.nbytes)) + reference = xr.Dataset({"data": (dims, values)}) + cached = wrapped["data"].variable._data + assert isinstance(cached, ChunkCachedArray) + return wrapped, reference, cached + + +def _assert_same_isel(wrapped: xr.Dataset, reference: xr.Dataset, indexers: dict) -> xr.DataArray: + got = wrapped["data"].isel(indexers) + ref = reference["data"].isel(indexers) + assert tuple(got.dims) == tuple(ref.dims) + assert not is_dask_collection(got.data) + np.testing.assert_array_equal(np.asarray(got.data), np.asarray(ref.data)) + return got + + +class _ChunkComputeCounter(Callback): + """Count dask tasks whose key names the source array.""" + + def __init__(self, array_name: str) -> None: + super().__init__() + self.array_name = array_name + self.all_keys: list[object] = [] + + def _pretask(self, key, dsk, state) -> None: + self.all_keys.append(key) + + @property + def chunk_computes(self) -> int: + return sum(self.array_name in (key if isinstance(key, str) else str(key)) for key in self.all_keys) + + +@contextmanager +def _counting() -> Iterator[_ChunkComputeCounter]: + counter = _ChunkComputeCounter(SOURCE_NAME) + with dask.config.set(scheduler="sync"), counter: + yield counter + + +@pytest.mark.parametrize( + "indexers", + [ + pytest.param( + { + "time": slice(None), + "depth": slice(1, 4), + "lat": _points(0, 6, 3, 1), + "lon": _points(0, 5, 2, 4), + }, + id="slices-before-full-and-bounded", + ), + pytest.param( + { + "time": _points(0, 4, 2, 1), + "depth": slice(0, 4, 2), + "lat": slice(None, None, -1), + "lon": _points(5, 0, 3, 1), + }, + id="slices-between-stepped-and-reverse", + ), + pytest.param( + { + "time": _points(4, 0, -1, 2), + "depth": _points(3, 0, 1, -4), + "lat": slice(1, 7, 2), + "lon": slice(None, None, -1), + }, + id="slices-after-bounded-and-reverse", + ), + ], +) +def test_mixed_isel_slice_position_matches_numpy(indexers): + """Full, bounded, stepped, and reverse slices keep xarray's axis order.""" + wrapped, reference, _ = _dataset() + _assert_same_isel(wrapped, reference, indexers) + + +def test_isel_omits_nonsingleton_dimension(): + """An omitted axis longer than one is a full slice in vectorized order.""" + wrapped, reference, _ = _dataset() + indexers = { + "time": _points(0, 4, 2, -1), + "depth": _points(0, 3, 1, 2), + "lon": _points(5, 0, -2, 4), + } + got = _assert_same_isel(wrapped, reference, indexers) + assert "lat" in got.dims + assert got.sizes["lat"] == SHAPE[DIMS.index("lat")] + + +@pytest.mark.parametrize("explicit_full_slice", [False, True], ids=["omitted", "explicit-full-slice"]) +def test_isel_singleton_dimension_matches_numpy(explicit_full_slice): + """A size-1 axis omitted from isel is the reported mixed-slice key.""" + dims = ("time", "mockZ", "N", "M") + wrapped, reference, _ = _dataset((4, 1, 5, 6), (2, 1, 2, 3), dims) + indexers = { + "time": _points(0, 3, -1, 1), + "N": _points(0, 4, 2, -2), + "M": _points(5, 0, 3, 1), + } + if explicit_full_slice: + indexers["mockZ"] = slice(None) + got = _assert_same_isel(wrapped, reference, indexers) + assert got.sizes["mockZ"] == 1 + + +def test_broadcast_multidimensional_indexers_with_leading_slice(): + """Size-1 index arrays broadcast, and a leading slice stays at the end.""" + wrapped, reference, _ = _dataset() + indexers = { + "time": slice(None, None, -1), + "depth": xr.DataArray(np.array([[0, 3], [1, 2], [3, 0]], dtype=np.int64), dims=("i", "j")), + "lat": xr.DataArray(np.array([0, 6], dtype=np.int64), dims="j"), + "lon": xr.DataArray(np.array([[5, 1], [0, 4], [2, 3]], dtype=np.int64), dims=("i", "j")), + } + got = _assert_same_isel(wrapped, reference, indexers) + assert tuple(got.dims) == ("time", "i", "j") + assert got.sizes["time"] == SHAPE[0] + + +def test_all_array_vindex_preserves_duplicates_order_and_negative_indices(): + """Equal-length 1D keys keep point order, duplicates, and valid negatives.""" + wrapped, reference, _ = _dataset() + indexers = { + "time": _points(4, 0, -1, 2, 4, -5), + "depth": _points(3, 0, -4, 1, 3, 2), + "lat": _points(6, 1, 0, -3, 6, 2), + "lon": _points(5, 0, -2, 4, 5, 1), + } + got = _assert_same_isel(wrapped, reference, indexers) + assert tuple(got.dims) == ("points",) + assert got.shape == (6,) + + +def test_mixed_vindex_duplicates_negative_indices_and_uneven_chunks(): + """Mixed keys select duplicate, unordered, and final short-chunk points.""" + wrapped, reference, _ = _dataset() + indexers = { + "time": _points(4, 0, -1, 2, 4), + "depth": slice(None, None, -1), + "lat": _points(6, 0, -7, 3, 1), + "lon": slice(0, 6, 2), + } + _assert_same_isel(wrapped, reference, indexers) + + +@pytest.mark.parametrize( + "indexers", + [ + pytest.param( + { + "time": _points(), + "depth": _points(), + "lat": _points(), + "lon": _points(), + }, + id="empty-point-arrays", + ), + pytest.param( + { + "time": _points(), + "depth": slice(None), + "lat": slice(1, 4), + "lon": _points(), + }, + id="empty-points-with-slices", + ), + pytest.param( + { + "time": _points(0, 2, 4), + "depth": slice(2, 2), + "lat": slice(None), + "lon": _points(1, 5, 0), + }, + id="zero-length-slice", + ), + pytest.param( + { + "time": _points(), + "depth": slice(1, 1), + "lat": slice(0, 0), + "lon": _points(), + }, + id="empty-points-and-slices", + ), + ], +) +def test_empty_selections_match_numpy_without_loading_chunks(indexers): + """Zero-length slices and empty point arrays do not fetch a chunk.""" + wrapped, reference, cached = _dataset() + with _counting() as counter: + got = _assert_same_isel(wrapped, reference, indexers) + assert got.size == 0 + assert counter.chunk_computes == 0, counter.all_keys + assert cached.cache.current_bytes == 0 + + +@pytest.mark.parametrize( + "indexers", + [ + pytest.param( + { + "time": _points(0), + "depth": _points(0), + "lat": _points(0), + "lon": _points(6), + }, + id="all-array-too-large", + ), + pytest.param( + { + "time": _points(-6), + "depth": _points(0), + "lat": _points(0), + "lon": _points(0), + }, + id="all-array-too-negative", + ), + pytest.param( + { + "time": _points(5), + "depth": slice(None), + "lat": _points(0), + "lon": _points(0), + }, + id="mixed-too-large", + ), + pytest.param( + { + "time": _points(0), + "depth": _points(-5), + "lat": slice(0, 3), + "lon": _points(1), + }, + id="mixed-too-negative", + ), + ], +) +def test_out_of_bounds_indices_raise_before_chunk_lookup(indexers): + """Indices outside ``[-size, size)`` raise IndexError and load nothing.""" + wrapped, _, cached = _dataset() + with _counting() as counter, pytest.raises(IndexError, match="out of bounds"): + wrapped["data"].isel(indexers) + assert counter.chunk_computes == 0, counter.all_keys + assert cached.cache.current_bytes == 0 + + +def test_zero_slice_step_raises_without_loading_chunks(): + """A slice step of zero is invalid and must not fetch a chunk.""" + wrapped, _, cached = _dataset() + indexers = { + "time": _points(0, 1), + "depth": slice(0, 4, 0), + "lat": _points(0, 2), + "lon": slice(None), + } + with _counting() as counter, pytest.raises(ValueError): + wrapped["data"].isel(indexers) + assert counter.chunk_computes == 0, counter.all_keys + assert cached.cache.current_bytes == 0 + + +def test_repeated_mixed_selection_loads_only_referenced_chunks(): + """The first mixed selection loads referenced chunks; the repeat loads none. + + ``depth`` chunks are ``(2, 2, 1)`` on a length-5 axis, so slice ``0:3`` + touches chunk 0 (indices 0, 1) and chunk 1 (index 2). Paired with time + indices 0 and 2 (chunks 0 and 1) and lat indices inside chunk 0, the + selection is exactly four chunks of the eighteen in the array. + """ + shape = (6, 5, 4) + chunks = (2, 2, 2) + wrapped, reference, cached = _dataset(shape, chunks, ("time", "depth", "lat")) + total_chunks = 3 * 3 * 2 + indexers = { + "time": _points(0, 2), + "depth": slice(0, 3), + "lat": _points(0, 1), + } + with _counting() as first: + _assert_same_isel(wrapped, reference, indexers) + assert first.chunk_computes == 4, first.all_keys + assert first.chunk_computes < total_chunks + assert cached.cache.current_bytes > 0 + + cached_bytes = cached.cache.current_bytes + with _counting() as second: + _assert_same_isel(wrapped, reference, indexers) + assert second.chunk_computes == 0, second.all_keys + assert cached.cache.current_bytes == cached_bytes diff --git a/tests/test_interpolation.py b/tests/test_interpolation.py index cfb690916..1fa8c8fa7 100644 --- a/tests/test_interpolation.py +++ b/tests/test_interpolation.py @@ -1,3 +1,4 @@ +import dask.array as da import numpy as np import pytest import xarray as xr @@ -13,6 +14,7 @@ VectorField, particlefile_to_v3_zarr, ) +from parcels._chunk_cached_array import ChunkCachedArray, wrap_dataset 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 @@ -277,6 +279,39 @@ 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]] +def test_get_corner_data_Agrid_chunk_cached_singleton_mockz(): + """Singleton unindexed depth must gather through the chunk cache. + + A Delft3D surface field has shape ``(time, mockZ, N, M)`` with ``mockZ`` + size 1 and ``axis_dim`` mapping only ``Y`` and ``X``. Corner gathering then + leaves ``mockZ`` as ``slice(None)`` inside a vectorized index. Before the + cache indexed only integer arrays, that slice raised ``TypeError``. The + cached corners must match NumPy-backed data with shape + ``(lenT, lenZ, 2, 2, npart)``. + """ + shape = (6, 1, 8, 7) + values = np.arange(np.prod(shape), dtype=np.float64).reshape(shape) + dims = ("time", "mockZ", "N", "M") + lazy = xr.DataArray(da.from_array(values, chunks=(2, 1, 3, 3)), dims=dims, name="U") + wrapped = wrap_dataset(lazy.to_dataset(), max_cache_bytes=int(values.nbytes)) + cached = wrapped["U"] + reference = xr.DataArray(values, dims=dims, name="U") + assert isinstance(cached.variable._data, ChunkCachedArray) + + npart = 5 + lenZ = 1 + ti = np.array([0, 2, 4, 3, 1], dtype=np.int32) + zi = np.zeros(npart, dtype=np.int32) + yi = np.array([0, 7, -3, 4, 6], dtype=np.int32) + xi = np.array([6, 0, -1, 3, 2], dtype=np.int32) + axis_dim = {"Y": "N", "X": "M"} + for lenT in (1, 2): + got = _get_corner_data_Agrid(cached, ti, zi, yi, xi, lenT, lenZ, npart, axis_dim) + ref = _get_corner_data_Agrid(reference, ti, zi, yi, xi, lenT, lenZ, npart, axis_dim) + assert got.shape == (lenT, lenZ, 2, 2, npart) + np.testing.assert_array_equal(got, ref) + + class XNearest_Velocity(VectorInterpolator): # noqa: N801 """Nearest-Neighbour interpolation on a regular grid for VectorFields of velocity."""