# `xarray` Integration Guide

This guide shows how to use `unxts.interop.xarray` to integrate JAX-based physical quantities with `xarray`'s labeled multi-dimensional arrays.

## Overview

`unxts.interop.xarray` provides seamless integration between:

- **`unxt`**: JAX-based physical quantities with dimension checking
- **`xarray`**: N-dimensional labeled arrays for scientific computing

The integration enables you to:

- Attach physical units to `xarray` DataArrays and Datasets
- Preserve units through `xarray` operations
- Convert between unit-aware (Quantity) and plain arrays with metadata
- Use JAX transformations (jit, vmap, grad) on unit-aware `xarray` objects

## Installation

::::{tab-set}

:::{tab-item} pip

```bash
pip install unxts.interop.xarray
```

:::

:::{tab-item} uv

```bash
uv add unxts.interop.xarray
```

:::

::::

## Basic Usage

### The `.unxt` Accessor

After importing `unxts.interop.xarray`, all DataArrays and Datasets gain a `.unxt` accessor with two main methods:

- `quantify()`: Convert attrs to Quantities
- `dequantify()`: Convert Quantities back to plain arrays with attrs

```python
import xarray as xr
import unxt as u
import unxts.interop.xarray  # Registers the .unxt accessor

# Create a DataArray with unit metadata
da = xr.DataArray(
    [1.0, 2.0, 3.0],
    dims=["x"],
    attrs={"units": "m"},
)

# Convert to Quantities
quantified = da.unxt.quantify()
print(quantified.data)
# Quantity(Array([1., 2., 3.], dtype=float32), unit='m')

# Convert back
dequantified = quantified.unxt.dequantify()
print(dequantified.attrs["units"])
# 'm'
```

## Working with DataArrays

### Quantifying from Attributes

The most common workflow starts with `xarray` objects that have unit information stored in their `attrs`:

```python
import jax.numpy as jnp
import xarray as xr
import unxt as u
import unxts.interop.xarray

# Temperature data with metadata
temp = xr.DataArray(
    jnp.array([20.0, 25.0, 30.0]),
    dims=["time"],
    coords={"time": [0, 1, 2]},
    attrs={"units": "K", "description": "Temperature measurements"},
)

# Convert to Quantities - units become part of the data
q_temp = temp.unxt.quantify()
print(q_temp.data)
# Quantity(Array([20., 25., 30.], dtype=float32), unit='K')

# Other attributes are preserved
print(q_temp.attrs["description"])
# 'Temperature measurements'
```

### Explicit Unit Specification

You can override or set units explicitly:

```python
# Override the units attribute
da = xr.DataArray([100.0, 200.0], dims=["x"], attrs={"units": "cm"})
quantified = da.unxt.quantify("m")
print(quantified.data)
# Quantity(Array([100., 200.], dtype=float32), unit='m')
```

Pass a string or `AbstractUnit` directly to apply it to the DataArray's data.

### Inspecting Discovered Units

You can inspect unit metadata discovered by the accessor before quantifying:

```python
import xarray as xr
import unxts.interop.xarray

da = xr.DataArray([1.0, 2.0], dims=["x"], attrs={"units": "m"})
print(da.unxt.units)
# {None: Unit("m")}
```

For DataArrays, the `None` key refers to the DataArray's own data.

### Coordinates with Units

Coordinates can also have units:

```python
import xarray as xr
import unxt as u
import unxts.interop.xarray

# Use a non-dimension coordinate to preserve units after quantify().
# (Dimension coordinates are coerced to plain arrays by xarray's indexing.)
da = xr.DataArray(
    [10.0, 20.0, 30.0],
    dims=["i"],
    coords={"i": [0, 1, 2], "time": ("i", [0.0, 1.0, 2.0], {"units": "s"})},
    attrs={"units": "m"},
)

# Quantify both data and the non-dimension coordinate
quantified = da.unxt.quantify()
print(quantified.data)
# Quantity(Array([10., 20., 30.], dtype=float32), unit='m')
print(quantified.coords["time"].data)
# Quantity(Array([0., 1., 2.], dtype=float32), unit='s')
```

**Important**: Use non-dimension coordinates (coordinates not marked with `*` in `xarray` output) to preserve Quantity objects. Dimension coordinates are automatically converted to plain arrays by `xarray`.

:::{warning}

**Dimension coordinates cannot hold Quantities.** xarray automatically coerces dimension coordinates (those named after their dimension, shown with `*` in the repr) to plain NumPy arrays via `pandas.Index`. This silently drops the Quantity wrapper and its unit — there is no error or warning.

The workaround is to use _non-dimension_ coordinates: give the coordinate a name different from its dimension (e.g., `"time"` attached to dimension `"i"`). See [Limitations: Dimension Coordinates Cannot Hold Quantities](#dimension-coordinates-cannot-hold-quantities) for the full explanation and workaround.

:::

## Working with Datasets

Datasets work similarly but handle multiple data variables:

```python
import xarray as xr
import unxt as u
import unxts.interop.xarray

# Create Dataset with multiple variables
ds = xr.Dataset(
    {
        "temperature": (["time"], [273.0, 293.0, 313.0], {"units": "K"}),
        "pressure": (["time"], [101325.0, 102000.0, 103000.0], {"units": "Pa"}),
    },
    coords={"time": [0, 1, 2]},
)

# Quantify all variables at once
q_ds = ds.unxt.quantify()
print(q_ds["temperature"].data)
# Quantity(Array([273., 293., 313.], dtype=float32), unit='K')
print(q_ds["pressure"].data)
# Quantity(Array([101325., 102000., 103000.], dtype=float32), unit='Pa')
```

### Per-Variable Units

You can specify units for specific variables:

```python
ds = xr.Dataset(
    {
        "distance": (["x"], [1.0, 2.0, 3.0]),
        "velocity": (["x"], [10.0, 20.0, 30.0]),
    }
)

q_ds = ds.unxt.quantify(
    units={
        "distance": "m",
        "velocity": "m/s",
    }
)
```

## Operations Preserve Units

Because the `Quantity` is stored as the DataArray's underlying (duck) array, units propagate through `xarray` operations directly — you operate on the labeled object, not on `.data`:

```python
import jax.numpy as jnp
import xarray as xr
import unxt as u
import unxts.interop.xarray

da = xr.DataArray(u.Quantity(jnp.array([1.0, 2.0, 3.0]), "m"), dims=["x"])

# Arithmetic combines units dimensionally
print((da * da).data)
# Quantity(Array([1., 4., 9.], dtype=float32), unit='m2')

# Reductions keep the unit
print(da.sum().data)
# Quantity(Array(6., dtype=float32), unit='m')
print(da.prod().data)
# Quantity(Array(6., dtype=float32), unit='m3')

# Masking is unit-aware
masked = da.where(da > u.Quantity(1.0, "m"))
print(masked.fillna(u.Quantity(0.0, "m")).data)
# Quantity(Array([0., 2., 3.], dtype=float32), unit='m')
```

`mean`, `std`, `min`, `max`, `median`, `quantile`, `dot`, `clip`, `cumsum`, `diff`, `integrate`, `concat`, `groupby`, and `weighted` likewise preserve units. This works through `unxt`'s Array-API-conformant `quaxed.numpy` namespace — `unxts.interop.xarray` does **not** monkeypatch `xarray` or its array namespace machinery.

## Dequantification

Converting back to plain arrays with unit metadata:

```python
import xarray as xr
import unxt as u
import unxts.interop.xarray

# Start with quantified data
q = u.Quantity([1.0, 2.0, 3.0], "m")
da = xr.DataArray(q, dims=["x"])

# Convert to plain arrays with unit attributes
plain = da.unxt.dequantify()
print(plain.data)
# Array([1., 2., 3.], dtype=float32)
print(plain.attrs["units"])
# 'm'
```

The `unit_attribute` parameter controls the attribute name (default: `"units"`):

```python
plain = da.unxt.dequantify(unit_attribute="unit_str")
print(plain.attrs["unit_str"])
# 'm'
```

## JAX Integration

Since `unxt` uses JAX arrays, all JAX transformations work seamlessly:

### JIT Compilation

```python
import jax
import jax.numpy as jnp
import xarray as xr
import unxt as u
import unxts.interop.xarray


@jax.jit
def process_data(data):
    """JIT-compiled function operating on the underlying Quantity."""
    return data * 2.0


# Create quantified DataArray
q = u.Quantity([1.0, 2.0, 3.0], "m")
da = xr.DataArray(q, dims=["x"])

# JIT works with the underlying data (a Quantity)
result = process_data(da.data)
print(result)
# Quantity(Array([2., 4., 6.], dtype=float32), unit='m')
```

### Vectorization

```python
import jax
import xarray as xr
import unxt as u


@jax.vmap
def square(x):
    return x**2


q = u.Quantity([[1.0, 2.0], [3.0, 4.0]], "m")
da = xr.DataArray(q, dims=["x", "y"])

# vmap over the data
squared = square(da.data)
print(squared)
# Quantity(Array([[ 1.,  4.],
#        [ 9., 16.]], dtype=float32), unit='m2')
```

### Auto-differentiation

```python
import jax
import xarray as xr
import unxt as u


def kinetic_energy(v):
    """Kinetic energy: KE = 0.5 * m * v^2."""
    m = u.Quantity(2.0, "kg")
    return 0.5 * m * v**2


v = u.Quantity([1.0, 2.0, 3.0], "m/s")
da = xr.DataArray(v, dims=["time"])

# Gradient with respect to velocity
grad_fn = jax.grad(lambda v_val: jax.numpy.sum(kinetic_energy(v_val).value))
dKE_dv = grad_fn(da.data)
```

## Roundtrip Conversions

The quantify/dequantify operations are designed to roundtrip:

```python
import jax.numpy as jnp
import xarray as xr
import unxt as u
import unxts.interop.xarray

# Start with attrs
original = xr.DataArray([1.0, 2.0], dims=["x"], attrs={"units": "m"})

# Roundtrip: attrs → Quantity → attrs
roundtrip = original.unxt.quantify().unxt.dequantify()

assert roundtrip.attrs["units"] == original.attrs["units"]
assert jnp.allclose(roundtrip.data, original.data)
```

## Best Practices

### 1. Import unxts.interop.xarray Early

Always import `unxts.interop.xarray` before using the `.unxt` accessor:

```python
import unxts.interop.xarray  # Registers the accessor
```

This registers the accessor on `xarray`'s DataArray and Dataset classes.

### 2. Use Non-Dimension Coordinates for Units

When working with coordinates that need to preserve Quantities:

```python
import unxt as u

# ✓ Good: non-dimension coordinate
coords = {"i": [0, 1], "x": ("i", u.Quantity([1.0, 2.0], "m"))}

# ✗ Bad: dimension coordinate (xarray will extract values)
coords = {"x": u.Quantity([1.0, 2.0], "m")}  # x is marked as dimension
```

### 3. Consistent Unit Attributes

Use consistent attribute names throughout your workflow. The default `"units"` is standard in many scientific data formats (CF conventions, NetCDF, etc.).

### 4. Preserve Other Metadata

The `quantify()` and `dequantify()` methods preserve all other attributes:

```python
da = xr.DataArray(
    [1.0, 2.0],
    dims=["x"],
    attrs={
        "units": "m",
        "long_name": "Distance",
        "standard_name": "distance",
    },
)

quantified = da.unxt.quantify()
# All non-unit attrs are preserved
assert quantified.attrs["long_name"] == "Distance"
```

## Common Patterns

### Loading from NetCDF

```python
import xarray as xr
import unxts.interop.xarray
from pathlib import Path

# Load dataset with unit metadata. The path to the bundled sample data is
# relative to the repository root (the directory the docs are built and run
# from); adjust it to wherever your own NetCDF file lives.
docs_dir = Path("packages/unxts.interop.xarray/docs")
data_path = docs_dir / "_data" / "sample_data.nc"

ds = xr.open_dataset(data_path)
print(ds)
# <xarray.Dataset> Size: 144B
# Dimensions:      (time: 2, location: 3)
# Coordinates:
#   * time         (time) float64 16B 0.0 3.6e+03
#   * location     (location) int64 24B 0 1 2
# Data variables:
#     temperature  (time, location) float64 48B 273.1 293.1 313.1 275.0 295.0 315.0
#     pressure     (time, location) float64 48B 1.013e+05 1.02e+05 ... 1.032e+05
#     distance     (location) float64 24B 0.0 100.0 200.0
# Attributes: (12/13)
#     ...

# Variables have unit metadata
print(ds["temperature"].attrs["units"])
# 'K'

# Convert all variables with units to Quantities
q_ds = ds.unxt.quantify()
print(q_ds["temperature"].data)
# Quantity(Array([[273.15, 293.15, 313.15],
#                 [275.  , 295.  , 315.  ]], dtype=float64), unit='K')
print(q_ds["pressure"].data)
# Quantity(Array([[101325., 102000., 103000.],
#                 [101500., 102200., 103200.]], dtype=float64), unit='Pa')
```

### Saving to NetCDF

```python
import xarray as xr
import unxt as u
import unxts.interop.xarray

# Create a quantified dataset
q_ds = xr.Dataset(
    {
        "distance": (["x"], u.Quantity([1.0, 2.0, 3.0], "m")),
        "velocity": (["x"], u.Quantity([10.0, 20.0, 30.0], "m/s")),
    }
)

# Dequantify before saving
plain_ds = q_ds.unxt.dequantify()
print(plain_ds["distance"].attrs["units"])
# 'm'

# Save to file
# plain_ds.to_netcdf("output.nc")
```

### Unit Conversion

Use `u.uconvert()` on the underlying `Quantity`, then wrap the result back into a DataArray:

```python
import xarray as xr
import unxt as u
import unxts.interop.xarray

# Quantified DataArray in metres
q = u.Quantity([1.0, 2.0, 3.0], "m")
da = xr.DataArray(q, dims=["x"])

# Convert to centimetres
da_cm = xr.DataArray(u.uconvert(u.unit("cm"), da.data), dims=da.dims, coords=da.coords)
print(da_cm.data)
# Quantity(Array([100., 200., 300.], dtype=float32), unit='cm')
```

For DataArrays that start with unit attrs, quantify first, convert, then dequantify:

```python
import xarray as xr
import unxt as u
import unxts.interop.xarray

plain = xr.DataArray([1.0, 2.0, 3.0], dims=["x"], attrs={"units": "m"})
quantified = plain.unxt.quantify()

converted = u.uconvert(u.unit("km"), quantified.data)
da_km = xr.DataArray(converted, dims=plain.dims).unxt.dequantify()
print(da_km.attrs["units"])
# 'km'
```

## Lower-Level API

The `.unxt` accessor covers most workflows, but the four underlying functions are also exported for use in pipelines, custom integrations, or cases where you need direct control.

### `extract_unit_attributes`

Reads `"units"` attrs from each variable and coordinate — without converting anything to a Quantity. Use this to inspect declared units before committing to a conversion.

```python
import xarray as xr
from unxts.interop.xarray import extract_unit_attributes

ds = xr.Dataset(
    {
        "temperature": ("time", [273.0, 293.0], {"units": "K"}),
        "pressure": ("time", [101325.0, 102000.0]),
    }
)
print(extract_unit_attributes(ds))
# {'temperature': Unit("K")}
```

### `attach_units`

Attaches units to a DataArray or Dataset, converting plain array data into Quantities. Use `None` as the key for a DataArray's own data (as opposed to a named coordinate).

```python
import xarray as xr
from unxts.interop.xarray import attach_units

da = xr.DataArray([1.0, 2.0, 3.0], dims=["x"])
quantified = attach_units(da, {None: "m"})
print(quantified.data)
# Quantity(Array([1., 2., 3.], dtype=float32), unit='m')
```

Use `attach_units` directly when you already have a units mapping (e.g., from a file header or a prior `extract_unit_attributes` call) and want to skip the attribute-reading step.

### `extract_units`

Reads the units from **existing Quantities** in a DataArray or Dataset. This is the inverse of `attach_units` — use it when you need the units for computation before stripping them.

```python
import xarray as xr
import unxt as u
from unxts.interop.xarray import extract_units

q = u.Quantity([1.0, 2.0], "m")
da = xr.DataArray(q, dims=["x"])
print(extract_units(da))
# {None: Unit("m")}
```

### `strip_units`

Removes Quantity wrappers, returning plain arrays. The unit information is discarded unless you capture it with `extract_units` first.

```python
import xarray as xr
import unxt as u
from unxts.interop.xarray import strip_units

q = u.Quantity([1.0, 2.0], "m")
da = xr.DataArray(q, dims=["x"])
stripped = strip_units(da)
print(stripped.data)
# Array([1., 2.], dtype=float32)
```

### When to use the low-level API

| Task | Use |
| --- | --- |
| Interactive quantify/dequantify | `.unxt.quantify()` / `.unxt.dequantify()` |
| Inspect declared units without converting | `extract_unit_attributes` |
| Attach a pre-built units mapping | `attach_units` |
| Read units from already-quantified data | `extract_units` |
| Strip Quantities to plain arrays | `strip_units` |
| Build a custom quantify/dequantify pipeline | All four, composed manually |

## Limitations

### Dimension Coordinates Cannot Hold Quantities

`xarray` backs every _dimension coordinate_ (one named like its dimension, shown with a `*` in the repr) with a `pandas.Index`. Building that index coerces the data to a plain `numpy` array, so a `Quantity` assigned to a dimension coordinate is silently unwrapped — its unit is lost. This is inherent to `xarray`'s indexing model, not something `unxts.interop.xarray` can override, and it affects every duck-array unit library (including `pint-xarray`) the same way.

**Workaround**: store the unitful values on a _non-dimension_ coordinate, keeping a plain index on the dimension itself:

```python
import unxt as u
import xarray as xr

data = [10.0, 20.0, 30.0]
quantities = u.Quantity([1.0, 2.0, 3.0], "m")

# Dimension coordinate: ``x`` is unwrapped to a plain array, unit lost
da = xr.DataArray(data, dims=["x"], coords={"x": quantities})
print(type(da.coords["x"].data).__name__)
# ndarray

# Non-dimension coordinate: the Quantity (and its unit) is preserved
da = xr.DataArray(data, dims=["i"], coords={"i": [0, 1, 2], "x": ("i", quantities)})
print(da.coords["x"].data)
# Quantity(Array([1., 2., 3.], dtype=float32), unit='m')
```

### Operations That Drop Units

A few `xarray` operations route through code paths that cannot preserve a `Quantity`:

- **`rolling` / sliding-window reductions** use `numpy.lib.stride_tricks`, which has no Array API (or `jax.numpy`) equivalent, so they are unsupported on JAX-backed data generally — not specific to units.
- **`interp`** delegates to `scipy`/`numpy` interpolation internally and returns a plain array (the same behavior as `pint-xarray`).

For these, `dequantify`, operate, then re-`quantify`, or work on `.data` with `unxt`/`quaxed` directly.

## See Also

- [unxt documentation](https://unxt.readthedocs.io/) - Core unitful quantities
- [xarray documentation](https://docs.xarray.dev/) - Labeled arrays
- [JAX documentation](https://jax.readthedocs.io/) - Composable transformations
- [Astropy units](https://docs.astropy.org/en/stable/units/) - Unit definitions
