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#

pip install unxts.interop.xarray
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

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:

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:

# 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:

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:

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 for the full explanation and workaround.

Working with Datasets#

Datasets work similarly but handle multiple data variables:

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:

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:

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:

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"):

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#

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#

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#

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:

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:

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:

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:

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#

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#

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:

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:

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.

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).

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.

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.

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:

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#