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
2 changes: 1 addition & 1 deletion .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -78,7 +78,7 @@ jobs:

- name: install dependencies
run: |
python -m pip install -r ci/requirements.txt
python -m pip install ".[test]"

- name: install astropy-xarray
run: python -m pip install --no-deps .
Expand Down
2 changes: 1 addition & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -104,7 +104,7 @@ celerybeat.pid

# Environments
.env
.venv
.venv*/
env/
venv/
ENV/
Expand Down
2 changes: 1 addition & 1 deletion .pre-commit-config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,7 @@ repos:
- repo: https://github.com/astral-sh/ruff-pre-commit
rev: v0.12.3
hooks:
- id: ruff
- id: ruff-check
args: [--fix]
- repo: https://github.com/kynan/nbstripout
rev: 0.8.1
Expand Down
1 change: 1 addition & 0 deletions astropy_xarray/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@

import astropy.units

import astropy_xarray.time_compat # noqa: F401
from astropy_xarray import accessors, formatting, testing # noqa: F401
from astropy_xarray.index import AstropyIndex

Expand Down
57 changes: 54 additions & 3 deletions astropy_xarray/accessors.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,11 +18,16 @@

import astropy
import astropy.units
from xarray import register_dataarray_accessor, register_dataset_accessor
from xarray import (
register_dataarray_accessor,
register_dataset_accessor,
register_datatree_accessor,
)
from xarray.core.dtypes import NA

from astropy_xarray import conversion
from astropy_xarray.conversion import no_unit_values
from astropy_xarray.conversion import AstropyUnitType, no_unit_values
from astropy_xarray.coordinates.sky_coord import dataset_to_skycoord
from astropy_xarray.errors import format_error_message

# sentinel to fallback to attribute unit, then value container type
Expand Down Expand Up @@ -64,7 +69,7 @@ def units_to_str_or_none(mapping, unit_format):
formatter = str if not unit_format else lambda v: unit_format.format(v)

return {
key: formatter(value) if isinstance(value, astropy.units.UnitBase) else value
key: formatter(value) if isinstance(value, AstropyUnitType) else value
for key, value in mapping.items()
}

Expand Down Expand Up @@ -96,6 +101,8 @@ def _decide_unit(unit, unit_attribute):
elif unit is _default:
if unit_attribute in no_unit_values:
return unit_attribute
if isinstance(unit_attribute, dict):
return unit_attribute
if isinstance(unit_attribute, astropy.units.UnitBase):
unit = unit_attribute
else:
Expand Down Expand Up @@ -1718,3 +1725,47 @@ def interpolate_na(
)

return conversion.attach_units(interpolated, units)

def to_sky_coord(self):
"""TODO"""
return dataset_to_skycoord(self.ds)


@register_datatree_accessor("astropy")
class AstropyDataTreeAccessor:
"""
Access methods for DataTree with units using Astropy.

Methods and attributes can be accessed through the `.astropy` attribute.
"""

def __init__(self, dt):
self.dt = dt

def quantify(self, units=_default, **unit_kwargs):
from xarray import DataTree

def breadth_first_quantify(dt: DataTree):
quantified_ds = (
dt.dataset.astropy.quantify(units=_default, **unit_kwargs)
if dt.dataset is not None
else None
)
children = {k: breadth_first_quantify(v) for k, v in dt.children.items()}
return DataTree(dataset=quantified_ds, children=children)

return breadth_first_quantify(self.dt)

def dequantify(self, format=None):
from xarray import DataTree

def breadth_first_dequantify(dt: DataTree):
dequantified_ds = (
dt.dataset.astropy.dequantify(format=format)
if dt.dataset is not None
else None
)
children = {k: breadth_first_dequantify(v) for k, v in dt.children.items()}
return DataTree(dataset=dequantified_ds, children=children)

return breadth_first_dequantify(self.dt)
70 changes: 63 additions & 7 deletions astropy_xarray/conversion.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@
import re

import astropy
import astropy.coordinates
import astropy.time
import astropy.units
from xarray import Coordinates, DataArray, Dataset, IndexVariable, Variable

Expand All @@ -34,11 +36,15 @@
datetime_units_re = re.compile(rf"{time_units_re} since {datetime_re}")


AstropyType = astropy.units.Quantity | astropy.time.TimeBase
AstropyUnitType = astropy.units.UnitBase | astropy.units.FunctionUnitBase


def is_datetime_unit(unit):
return isinstance(unit, str) and datetime_units_re.match(unit) is not None


def array_attach_unit(data, unit) -> astropy.units.Quantity:
def array_attach_unit(data, unit) -> AstropyType:
"""attach a unit to the data

Parameters
Expand All @@ -55,7 +61,7 @@ def array_attach_unit(data, unit) -> astropy.units.Quantity:
if unit in no_unit_values:
return data

if not isinstance(unit, astropy.units.UnitBase):
if not isinstance(unit, AstropyUnitType | dict):
raise ValueError(f"cannot use {unit!r} as a unit")

if isinstance(data, astropy.units.Quantity):
Expand All @@ -67,7 +73,35 @@ def array_attach_unit(data, unit) -> astropy.units.Quantity:
f"already has units {data.unit}"
)

return astropy.units.Quantity(data, unit)
if isinstance(unit, dict):
match unit["class"].lower():
case "time":
return astropy.time.Time(
data,
format=unit["format"],
scale=unit["scale"],
precision=unit["precision"],
)
case "timedelta":
return astropy.time.TimeDelta(
data,
format=unit["format"],
scale=unit["scale"],
precision=unit["precision"],
)
case "angle":
return astropy.coordinates.Angle(data, unit=unit["unit"])
case "longitude":
return astropy.coordinates.Longitude(data, unit=unit["unit"])
case "latitude":
return astropy.coordinates.Latitude(data, unit=unit["unit"])
case "distance":
return astropy.coordinates.Distance(data, unit=unit["unit"])

if isinstance(unit, astropy.units.LogUnit):
return astropy.units.LogQuantity(data, unit)
else:
return astropy.units.Quantity(data, unit)


def array_convert_unit(data, unit, equivalencies) -> astropy.units.Quantity:
Expand Down Expand Up @@ -111,21 +145,43 @@ def array_convert_unit(data, unit, equivalencies) -> astropy.units.Quantity:
return data


def array_extract_unit(data):
def array_extract_unit(data) -> dict | astropy.units.Unit | astropy.units.LogUnit:
"""extract the unit of an array

If ``data`` is not a quantity, the units are ``None``
"""
try:
return data.unit
if isinstance(
data,
(
astropy.coordinates.Longitude,
astropy.coordinates.Latitude,
astropy.coordinates.Distance,
),
):
return {"class": data.__class__.__name__.lower(), "unit": str(data.unit)}
elif isinstance(data, astropy.time.TimeBase):
return {
"class": data.__class__.__name__.lower(),
"format": data.format,
"scale": data.scale,
"precision": data.precision,
}
elif isinstance(data, astropy.units.Quantity):
return data.unit
else:
return None
except AttributeError:
return None


def array_strip_unit(data):
"""strip the unit of a quantity"""
try:
return data.value
if isinstance(data, (astropy.units.Quantity, astropy.time.TimeBase)):
return data.value
else:
return data
except AttributeError:
return data

Expand Down Expand Up @@ -405,7 +461,7 @@ def extract_unit_attributes(obj, attr="units"):


def strip_units_variable(var):
if not isinstance(var.data, astropy.units.Quantity):
if not isinstance(var.data, AstropyType):
return var

data = array_strip_unit(var.data)
Expand Down
12 changes: 12 additions & 0 deletions astropy_xarray/coordinates/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,12 @@
from astropy_xarray.coordinates.frame import load_frame, load_representation
from astropy_xarray.coordinates.sky_coord import (
dataset_to_skycoord,
skycoord_to_dataset,
)

__all__ = [
"dataset_to_skycoord",
"skycoord_to_dataset",
"load_frame",
"load_representation",
]
26 changes: 26 additions & 0 deletions astropy_xarray/coordinates/core.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,26 @@
from typing import TypeVar

import astropy.units as u
from astropy.time import Time


# Time
def dump_time(time: Time):
return {
"val": time.value,
"format": time.format,
"precision": time.precision,
"scale": time.scale,
}


# Quantity
def dump_quantity(q: u.Quantity):
return {"value": float(q.value), "unit": str(q.unit)}


_T = TypeVar("_T", bound=u.Quantity | Time)


def load_optional_object(cls: type[_T], kwargs: dict | None) -> _T | None:
return cls(**kwargs) if kwargs is not None else None
Loading
Loading