Skip to content

Commit 40e3140

Browse files
committed
Add vectorized indexing
1 parent 9f3ef9c commit 40e3140

1 file changed

Lines changed: 189 additions & 7 deletions

File tree

tests/test_indexing.py

Lines changed: 189 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -81,7 +81,7 @@ def pandas_da(raster_da):
8181
return da
8282

8383

84-
def pos_to_label_indexer(idx: pd.Index, idxr: int | slice | np.ndarray) -> Any:
84+
def pos_to_label_indexer(idx: pd.Index, idxr: int | slice | np.ndarray, *, use_scalar: bool = True) -> Any:
8585
if isinstance(idxr, slice):
8686
return slice(
8787
None if idxr.start is None else idx[idxr.start],
@@ -93,7 +93,7 @@ def pos_to_label_indexer(idx: pd.Index, idxr: int | slice | np.ndarray) -> Any:
9393
return idx[idxr].values
9494
else:
9595
val = idx[idxr]
96-
if st.booleans():
96+
if use_scalar:
9797
try:
9898
# pass python scalars occasionally
9999
val = val.item()
@@ -109,7 +109,7 @@ def basic_indexers(
109109
/,
110110
*,
111111
sizes: dict[Hashable, int],
112-
min_dims: int = 0,
112+
min_dims: int = 1,
113113
max_dims: int | None = None,
114114
) -> dict[Hashable, int | slice]:
115115
"""Generate basic indexers using hypothesis.extra.numpy.basic_indices.
@@ -121,9 +121,9 @@ def basic_indexers(
121121
sizes : dict[Hashable, int]
122122
Dictionary mapping dimension names to their sizes.
123123
min_dims : int, optional
124-
Minimum dimensionality of the generated index. Default is 0.
124+
Minimum dimensionality of the generated index.
125125
max_dims : int or None, optional
126-
Maximum dimensionality of the generated index. Default is None (no limit).
126+
Maximum dimensionality of the generated index.
127127
128128
Returns
129129
-------
@@ -187,7 +187,10 @@ def basic_label_indexers(draw, /, *, indexes: Indexes) -> dict[Hashable, float |
187187
pos_indexer = draw(basic_indexers(sizes=sizes))
188188
pdindexes = indexes.to_pandas_indexes()
189189

190-
label_indexer = {dim: pos_to_label_indexer(pdindexes[dim], idx) for dim, idx in pos_indexer.items()}
190+
label_indexer = {
191+
dim: pos_to_label_indexer(pdindexes[dim], idx, use_scalar=draw(st.booleans()))
192+
for dim, idx in pos_indexer.items()
193+
}
191194
return label_indexer
192195

193196

@@ -274,7 +277,155 @@ def outer_array_label_indexers(draw, /, *, indexes: Indexes) -> dict[Hashable, n
274277
pos_indexer = draw(outer_array_indexers(sizes=sizes))
275278
pdindexes = indexes.to_pandas_indexes()
276279

277-
label_indexer = {dim: pos_to_label_indexer(pdindexes[dim], idx) for dim, idx in pos_indexer.items()}
280+
label_indexer = {
281+
dim: pos_to_label_indexer(pdindexes[dim], idx, use_scalar=False) for dim, idx in pos_indexer.items()
282+
}
283+
return label_indexer
284+
285+
286+
@st.composite
287+
def vectorized_indexers(
288+
draw,
289+
/,
290+
*,
291+
sizes: dict[Hashable, int],
292+
min_dims: int = 2,
293+
max_dims: int | None = None,
294+
min_ndim: int = 1,
295+
max_ndim: int = 3,
296+
min_size: int = 1,
297+
max_size: int = 5,
298+
) -> dict[Hashable, xr.DataArray]:
299+
"""Generate vectorized (fancy) indexers where all arrays are broadcastable.
300+
301+
In vectorized indexing, all array indexers must have compatible shapes
302+
that can be broadcast together, and the result shape is determined by
303+
broadcasting the indexer arrays.
304+
305+
Parameters
306+
----------
307+
draw : callable
308+
The Hypothesis draw function (automatically provided by @st.composite).
309+
sizes : dict[Hashable, int]
310+
Dictionary mapping dimension names to their sizes.
311+
min_dims : int, optional
312+
Minimum number of dimensions to index. Default is 2, so that we always have a "trajectory".
313+
Use ``outer_array_indexers`` for the ``min_dims==1`` case.
314+
max_dims : int or None, optional
315+
Maximum number of dimensions to index. Default is None (no limit).
316+
min_ndim : int, optional
317+
Minimum number of dimensions for the result arrays. Default is 1.
318+
max_ndim : int, optional
319+
Maximum number of dimensions for the result arrays. Default is 3.
320+
min_size : int, optional
321+
Minimum size for each dimension in the result arrays. Default is 1.
322+
max_size : int, optional
323+
Maximum size for each dimension in the result arrays. Default is 5.
324+
325+
Returns
326+
-------
327+
dict[Hashable, xr.DataArray]
328+
Indexers as a dict with keys randomly selected from sizes.keys().
329+
Values are DataArrays of integer indices that are all broadcastable
330+
to a common shape.
331+
"""
332+
# Get all dimension names
333+
all_dims = list(sizes.keys())
334+
335+
# Determine how many dimensions to index
336+
num_dims = draw(st.integers(min_value=min_dims, max_value=min(max_dims or len(all_dims), len(all_dims))))
337+
338+
# Randomly select which dimensions to index
339+
selected_dim_names = draw(st.permutations(all_dims).map(lambda x: x[:num_dims]))
340+
selected_dims = {dim: sizes[dim] for dim in selected_dim_names}
341+
342+
# Require at least one dimension to be indexed to avoid edge cases
343+
if num_dims == 0:
344+
return {}
345+
346+
# Generate a common broadcast shape for all arrays
347+
# Use min_ndim to max_ndim dimensions for the result shape
348+
result_ndim = draw(st.integers(min_value=min_ndim, max_value=max_ndim))
349+
result_shape = tuple(
350+
draw(st.integers(min_value=min_size, max_value=max_size)) for _ in range(result_ndim)
351+
)
352+
353+
# Create dimension names for the vectorized result
354+
vec_dims = tuple(f"vec_{i}" for i in range(result_ndim))
355+
356+
# Generate array indexers for each selected dimension
357+
# All arrays must be broadcastable to the same result_shape
358+
# To ensure proper broadcasting, use the same decision for all dimensions
359+
# (i.e., if vec_1 should be size 1, it's 1 for all dimensions)
360+
broadcast_mask = tuple(draw(st.booleans()) for _ in result_shape)
361+
362+
# If min_size > 1, prevent all dimensions from being broadcast to 1
363+
# This ensures the resulting arrays have at least min_size total elements
364+
if min_size > 1 and all(broadcast_mask):
365+
idx_to_keep = draw(st.integers(min_value=0, max_value=result_ndim - 1))
366+
broadcast_mask = tuple(False if i == idx_to_keep else True for i in range(result_ndim))
367+
368+
result = {}
369+
for dim, size in selected_dims.items():
370+
# Apply the same broadcast mask to all arrays
371+
array_shape = tuple(
372+
1 if use_one else s for use_one, s in zip(broadcast_mask, result_shape, strict=True)
373+
)
374+
375+
# Generate array of valid indices for this dimension
376+
indices = draw(
377+
npst.arrays(
378+
dtype=np.int64,
379+
shape=array_shape,
380+
elements=st.integers(min_value=0, max_value=size - 1),
381+
)
382+
)
383+
384+
# Wrap in DataArray with named dimensions for vectorized indexing
385+
result[dim] = xr.DataArray(indices, dims=vec_dims)
386+
387+
return result
388+
389+
390+
@st.composite
391+
def vectorized_label_indexers(draw, /, *, indexes: Indexes, **kwargs) -> dict[Hashable, xr.DataArray]:
392+
"""Generate label-based vectorized indexers by converting position indexers to labels.
393+
394+
This works in label space by using the coordinate Index values.
395+
396+
Parameters
397+
----------
398+
draw : callable
399+
The Hypothesis draw function (automatically provided by @st.composite).
400+
indexes : Indexes
401+
Dictionary mapping dimension names to their associated indexes
402+
**kwargs : dict
403+
Additional keyword arguments to pass to vectorized_indexers
404+
405+
Returns
406+
-------
407+
dict[Hashable, xr.DataArray]
408+
Label-based indexers as a dict with keys from indexes.
409+
Values are DataArrays of label values for each dimension.
410+
"""
411+
idxs = indexes.get_unique()
412+
assert all(isinstance(idx, xr.indexes.PandasIndex) for idx in idxs)
413+
414+
# FIXME: this should be indexes.sizes!
415+
sizes = indexes.dims
416+
417+
pos_indexer = draw(vectorized_indexers(sizes=sizes, **kwargs))
418+
pdindexes = indexes.to_pandas_indexes()
419+
420+
label_indexer = {}
421+
for dim, idx_array in pos_indexer.items():
422+
# Convert each position in the array to its corresponding label
423+
# Flatten, index, then reshape back to original shape
424+
flat_indices = idx_array.values.ravel()
425+
flat_labels = pdindexes[dim][flat_indices].values
426+
label_values = flat_labels.reshape(idx_array.shape)
427+
label_indexer[dim] = xr.DataArray(label_values, dims=idx_array.dims)
428+
278429
return label_indexer
279430

280431

@@ -360,3 +511,34 @@ def test_outer_array_label_indexing(data, raster_da, pandas_da):
360511
result_pandas = pandas_da.sel(indexers, method="nearest")
361512

362513
xr.testing.assert_identical(result_raster, result_pandas)
514+
515+
516+
@given(data=st.data())
517+
@settings(max_examples=200, suppress_health_check=[HealthCheck.function_scoped_fixture])
518+
def test_vectorized_indexing(data, raster_da, pandas_da):
519+
"""Test that vectorized indexing produces identical results for RasterIndex and PandasIndex."""
520+
sizes = dict(raster_da.sizes)
521+
indexers = data.draw(vectorized_indexers(sizes=sizes))
522+
523+
result_raster = raster_da.isel(indexers)
524+
result_pandas = pandas_da.isel(indexers)
525+
526+
xr.testing.assert_identical(result_raster, result_pandas)
527+
528+
529+
@given(data=st.data())
530+
@settings(
531+
deadline=None,
532+
max_examples=200,
533+
suppress_health_check=[HealthCheck.function_scoped_fixture],
534+
)
535+
def test_vectorized_label_indexing(data, raster_da, pandas_da):
536+
"""Test that vectorized label indexing produces identical results for RasterIndex and PandasIndex."""
537+
# RasterIndex has a bug with size-1 arrays in vectorized indexing
538+
# Use min_size=2 to avoid creating arrays with only 1 element total
539+
indexers = data.draw(vectorized_label_indexers(indexes=pandas_da.xindexes))
540+
541+
result_raster = raster_da.sel(indexers, method="nearest")
542+
result_pandas = pandas_da.sel(indexers, method="nearest")
543+
544+
xr.testing.assert_identical(result_raster, result_pandas)

0 commit comments

Comments
 (0)