Skip to content

Commit d1fbb9c

Browse files
authored
Support dictionary-style CF indexing (#661)
1 parent 4b5dc74 commit d1fbb9c

2 files changed

Lines changed: 80 additions & 4 deletions

File tree

cf_xarray/accessor.py

Lines changed: 44 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1406,6 +1406,31 @@ def __iter__(self):
14061406
return iter(self.wrapped)
14071407

14081408

1409+
class _CFWrappedLocIndexer:
1410+
"""Wrap xarray's label-based indexer with CF name translation."""
1411+
1412+
def __init__(self, accessor: CFAccessor):
1413+
self.accessor = accessor
1414+
1415+
def _translate_indexers(self, key):
1416+
if not isinstance(key, Mapping):
1417+
return key
1418+
1419+
_, arguments = self.accessor._process_signature(
1420+
self.accessor._obj.sel,
1421+
(key,),
1422+
{},
1423+
dict(_DEFAULT_KEY_MAPPERS),
1424+
)
1425+
return arguments["indexers"]
1426+
1427+
def __getitem__(self, key):
1428+
return self.accessor._obj.loc[self._translate_indexers(key)]
1429+
1430+
def __setitem__(self, key, value):
1431+
self.accessor._obj.loc[self._translate_indexers(key)] = value
1432+
1433+
14091434
class _CFWrappedPlotMethods:
14101435
"""
14111436
This class wraps DataArray.plot
@@ -1942,6 +1967,11 @@ def __contains__(self, item: str) -> bool:
19421967
"""
19431968
return item in self.keys()
19441969

1970+
@property
1971+
def loc(self):
1972+
"""Label-based indexer that understands CF names."""
1973+
return _CFWrappedLocIndexer(self)
1974+
19451975
@property
19461976
def plot(self):
19471977
"""
@@ -2809,13 +2839,15 @@ def grid_mappings(self) -> tuple[GridMapping, ...]:
28092839

28102840
@xr.register_dataset_accessor("cf")
28112841
class CFDatasetAccessor(CFAccessor):
2812-
def __getitem__(self, key: Hashable | Iterable[Hashable]) -> DataArray | Dataset:
2842+
def __getitem__(
2843+
self, key: Hashable | Iterable[Hashable] | Mapping[Hashable, Any]
2844+
) -> DataArray | Dataset:
28132845
"""
28142846
Index into a Dataset making use of CF attributes.
28152847
28162848
Parameters
28172849
----------
2818-
key : str, Iterable[str], optional
2850+
key : str, Iterable[str], or Mapping, optional
28192851
One of
28202852
- axes names: "X", "Y", "Z", "T"
28212853
- coordinate names: "longitude", "latitude", "vertical", "time"
@@ -2841,6 +2873,9 @@ def __getitem__(self, key: Hashable | Iterable[Hashable]) -> DataArray | Dataset
28412873
28422874
Add additional keys by specifying "custom criteria". See :ref:`custom_criteria` for more.
28432875
"""
2876+
if isinstance(key, Mapping):
2877+
return self.isel(key)
2878+
28442879
return _getitem(self, key)
28452880

28462881
@property
@@ -3474,13 +3509,15 @@ def grid_mapping_name(self) -> str:
34743509
# Return the single grid mapping name
34753510
return next(iter(grid_mapping_names.keys()))
34763511

3477-
def __getitem__(self, key: Hashable | Iterable[Hashable]) -> DataArray:
3512+
def __getitem__(
3513+
self, key: Hashable | Iterable[Hashable] | Mapping[Hashable, Any]
3514+
) -> DataArray:
34783515
"""
34793516
Index into a DataArray making use of CF attributes.
34803517
34813518
Parameters
34823519
----------
3483-
key : str, Iterable[str], optional
3520+
key : str, Iterable[str], or Mapping, optional
34843521
One of
34853522
- axes names: "X", "Y", "Z", "T"
34863523
- coordinate names: "longitude", "latitude", "vertical", "time"
@@ -3509,6 +3546,9 @@ def __getitem__(self, key: Hashable | Iterable[Hashable]) -> DataArray:
35093546
Add additional keys by specifying "custom criteria". See :ref:`custom_criteria` for more.
35103547
"""
35113548

3549+
if isinstance(key, Mapping):
3550+
return self.isel(key)
3551+
35123552
if not isinstance(key, Hashable):
35133553
raise KeyError(
35143554
f"Cannot use an Iterable of keys with DataArrays. Expected a single string. Received {key!r} instead."

cf_xarray/tests/test_accessor.py

Lines changed: 36 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -517,6 +517,42 @@ def test_kwargs_methods(obj):
517517
assert_identical(expected, actual)
518518

519519

520+
@pytest.mark.parametrize("obj", objects)
521+
def test_dictionary_indexing(obj):
522+
with raise_if_dask_computes():
523+
expected = obj.isel(time=0)
524+
actual = obj.cf[{"T": 0}]
525+
assert_identical(expected, actual)
526+
527+
528+
@pytest.mark.parametrize("obj", objects)
529+
def test_loc_indexing(obj):
530+
label = obj.time.values[0]
531+
with raise_if_dask_computes():
532+
expected = obj.loc[{"time": label}]
533+
actual = obj.cf.loc[{"T": label}]
534+
assert_identical(expected, actual)
535+
536+
537+
@pytest.mark.parametrize("obj", objects)
538+
def test_loc_assignment(obj):
539+
expected = obj.copy()
540+
actual = obj.copy()
541+
label = obj.time.values[0]
542+
543+
expected.loc[{"time": label}] = -1
544+
actual.cf.loc[{"T": label}] = -1
545+
546+
assert_identical(expected, actual)
547+
548+
549+
def test_dictionary_indexing_expands_multiple_dimensions():
550+
expected = multiple.isel(x1=5, x2=5)
551+
actual = multiple.cf[{"X": 5}]
552+
553+
assert_identical(expected, actual)
554+
555+
520556
def test_pos_args_methods() -> None:
521557
expected = airds.transpose("lon", "time", "lat")
522558
actual = airds.cf.transpose("longitude", "T", "latitude")

0 commit comments

Comments
 (0)