Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
102 changes: 54 additions & 48 deletions astropy_xarray/accessors.py

Large diffs are not rendered by default.

48 changes: 27 additions & 21 deletions astropy_xarray/conversion.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,7 +70,7 @@ def array_attach_unit(data, unit) -> astropy.units.Quantity:
return astropy.units.Quantity(data, unit)


def array_convert_unit(data, unit) -> astropy.units.Quantity:
def array_convert_unit(data, unit, equivalencies) -> astropy.units.Quantity:
"""convert the unit of an array

This is roughly the same as ``data.to(unit)``.
Expand All @@ -80,8 +80,14 @@ def array_convert_unit(data, unit) -> astropy.units.Quantity:
data : quantity or array-like
The data to convert. If it is not a quantity, it is assumed to be
dimensionless.
unit : str or astropy.units.Unit
unit : str or astropy.units.UnitBase
The unit to convert to. If a string ``data`` has to be a quantity.
equivalencies : list | None
A list of equivalence pairs to try if the units are not
directly convertible. See :py:doc:`astropy:units/equivalencies`.
This list is in addition to possible global defaults set by,
e.g., :py:func:`astropy.units.set_enabled_equivalencies`.
Use None to turn off all equivalencies.

Returns
-------
Expand All @@ -95,7 +101,7 @@ def array_convert_unit(data, unit) -> astropy.units.Quantity:
raise ValueError(f"cannot convert a non-quantity using {unit!r} as unit")

data = (
data.to(unit)
data.to(unit, equivalencies)
if isinstance(data, astropy.units.Quantity)
else astropy.units.Quantity(data, unit)
)
Expand Down Expand Up @@ -239,32 +245,32 @@ def attach_unit_attributes(obj, units, attr="units"):
return new_obj


def convert_units_variable(variable, units):
def convert_unit_variable(variable, unit, equivalencies):
if isinstance(variable, IndexVariable):
if variable.level_names:
# don't try to convert MultiIndexes
return variable

if units is not None:
if unit is not None:
quantity = array_attach_unit(
variable.data, variable.attrs.get(unit_attribute_name)
)
converted = array_convert_unit(quantity, units)
converted = array_convert_unit(quantity, unit, equivalencies)
new_obj = variable.copy(data=array_strip_unit(converted))

new_obj.attrs[unit_attribute_name] = array_extract_unit(converted)
else:
new_obj = variable
elif isinstance(variable, Variable):
converted = array_convert_unit(variable.data, units)
converted = array_convert_unit(variable.data, unit, equivalencies)
new_obj = variable.copy(data=converted)
else:
raise ValueError(f"unknown type: {variable}")

return new_obj


def convert_units_index(index, index_vars, units):
def convert_units_index(index, index_vars, units, equivalencies):
if not isinstance(index, AstropyIndex):
raise ValueError("cannot convert non-quantified index")

Expand All @@ -273,7 +279,7 @@ def convert_units_index(index, index_vars, units):
for name, var in index_vars.items():
unit = units.get(name)
try:
converted = convert_units_variable(var, unit)
converted = convert_unit_variable(var, unit, equivalencies)
converted_vars[name] = strip_units_variable(converted)
except (ValueError, astropy.units.core.UnitConversionError) as e:
failed[name] = e
Expand All @@ -287,7 +293,7 @@ def convert_units_index(index, index_vars, units):
return AstropyIndex(index=converted_index, units=units)


def convert_units_dataset(obj, units):
def convert_units_dataset(obj, units, equivalencies):
converted = {}
failed = {}
indexed_variables = obj.xindexes.variables
Expand All @@ -297,7 +303,7 @@ def convert_units_dataset(obj, units):

unit = units.get(name)
try:
converted[name] = convert_units_variable(var, unit)
converted[name] = convert_unit_variable(var, unit, equivalencies)
except (ValueError, astropy.units.core.UnitConversionError) as e:
failed[name] = e

Expand All @@ -308,7 +314,7 @@ def convert_units_dataset(obj, units):
continue

try:
converted_index = convert_units_index(idx, idx_vars, idx_units)
converted_index = convert_units_index(idx, idx_vars, idx_units, equivalencies)
indexes.update({k: converted_index for k in idx_vars})
index_vars.update(converted_index.create_variables())
except (ValueError, astropy.units.core.UnitConversionError) as e:
Expand All @@ -324,7 +330,7 @@ def convert_units_dataset(obj, units):
return dataset_from_variables(reordered, obj._coord_names, indexes, obj.attrs)


def convert_units(obj, units):
def convert_units(obj, units, equivalencies):
if not isinstance(obj, (DataArray, Dataset)):
raise ValueError(f"cannot convert object: {obj!r}: unknown type")

Expand All @@ -335,7 +341,7 @@ def convert_units(obj, units):

try:
new_obj = call_on_dataset(
convert_units_dataset, obj, name=temporary_name, units=units
convert_units_dataset, obj, name=temporary_name, units=units, equivalencies=equivalencies
)
except ValueError as e:
(failed,) = e.args
Expand Down Expand Up @@ -464,27 +470,27 @@ def slice_extract_units(indexer):
return astropy.units.Quantity(1, units_).si.unit


def convert_units_slice(indexer, units):
def convert_units_slice(indexer, units, equivalencies):
attrs = {name: getattr(indexer, name) for name in slice_attributes}
converted = {
name: array_convert_unit(value, units) if value is not None else None
name: array_convert_unit(value, units, equivalencies) if value is not None else None
for name, value in attrs.items()
}
args = [converted[name] for name in slice_attributes]

return slice(*args)


def convert_indexer_units(indexers, units):
def convert_indexer_units(indexers, units, equivalencies):
def convert(indexer, units):
if isinstance(indexer, slice):
return convert_units_slice(indexer, units)
return convert_units_slice(indexer, units, equivalencies)
elif isinstance(indexer, DataArray):
return convert_units(indexer, {None: units})
return convert_units(indexer, {None: units}, equivalencies)
elif isinstance(indexer, Variable):
return convert_units_variable(indexer, units)
return convert_unit_variable(indexer, units, equivalencies)
else:
return array_convert_unit(indexer, units)
return array_convert_unit(indexer, units, equivalencies)

converted = {}
invalid = {}
Expand Down
4 changes: 2 additions & 2 deletions astropy_xarray/index.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,8 +67,8 @@ def stack(cls, variables, dim):
def unstack(self):
raise NotImplementedError()

def sel(self, labels, **options):
converted_labels = conversion.convert_indexer_units(labels, self.units)
def sel(self, labels, equivalencies=None, **options):
converted_labels = conversion.convert_indexer_units(labels, self.units, equivalencies)
stripped_labels = conversion.strip_indexer_units(converted_labels)

return self.index.sel(stripped_labels, **options)
Expand Down
Loading
Loading