# Copyright 2026 The thomaspinder Contributors. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
r"""Bridge between labelled xarray data and GPJax.
A GP in GPJax takes a :class:`~gpjax.dataset.Dataset` of flat ``(N, D)`` inputs
and ``(N, 1)`` outputs. Gridded data arrives as labelled xarray objects instead.
This module converts between the two at the edges of a workflow and nowhere
else:
.. code-block:: python
data, spec = from_xarray(ds, target="t2m", inputs=["lat", "lon", "elevation"])
posterior = model.condition(data)
test_inputs, test_spec = spec.inputs_for(grid)
out = test_spec.to_xarray(posterior(test_inputs))
``data`` is an ordinary :class:`~gpjax.dataset.Dataset`. Everything needed to
rebuild the grid lives on the :class:`GridSpec`, which never enters ``fit``,
``condition`` or a traced computation.
Requires the optional ``xarray`` extra: ``pip install "gpjax[xarray]"``.
"""
from dataclasses import dataclass
import beartype.typing as tp
import jax.numpy as jnp
from jaxtyping import Float
import lineax as lx
import numpy as np
from gpjax.dataset import Dataset
from gpjax.distributions import GaussianDistribution
from gpjax.typing import (
Array,
KeyArray,
)
try:
import xarray as xr
except ImportError as error: # pragma: no cover - exercised without the extra
raise ImportError(
"gpjax.xarray requires xarray; install it with pip install 'gpjax[xarray]'"
) from error
[docs]
@dataclass(frozen=True, repr=False)
class GridSpec:
r"""The labelled grid behind a flattened :class:`~gpjax.dataset.Dataset`.
Row ``i`` of the flattened data is the ``i``-th kept cell of the grid in C
order over ``dims``. That correspondence is all this object records, and all
:meth:`inputs_for` and :meth:`to_xarray` need.
Attributes:
target: Name of the modelled variable.
target_attrs: The target's attributes (units, long_name, ...), carried
onto predictions.
inputs: Input names, in the column order of ``X``.
dims: The grid's dims, in the target's order.
coords: The grid's coordinates, used to rebuild labelled output.
time_origins: For each datetime input, the timestamp encoded as day 0.
mask: Boolean array over the full grid; ``True`` marks a cell that has a
row in the flattened data.
n_dropped: Number of grid cells dropped for containing NaN.
"""
target: str
target_attrs: dict[str, tp.Any]
inputs: tuple[str, ...]
dims: tuple[str, ...]
coords: tp.Mapping[tp.Hashable, xr.DataArray]
time_origins: dict[str, np.datetime64]
mask: np.ndarray
n_dropped: int
@property
def n_kept(self) -> int:
r"""Number of grid cells with a row in the flattened data."""
return int(self.mask.sum())
def __repr__(self) -> str:
r"""Summarise the grid without printing its coordinates or mask."""
grid = dict(zip(self.dims, self.mask.shape, strict=True))
return (
f"GridSpec(target={self.target!r}, inputs={self.inputs}, grid={grid}, "
f"kept={self.n_kept}, dropped={self.n_dropped})"
)
[docs]
def to_xarray(
self,
dist: GaussianDistribution,
*,
num_samples: tp.Optional[int] = None,
key: tp.Optional[KeyArray] = None,
) -> xr.Dataset:
r"""Map a predictive distribution back onto the labelled grid.
Args:
dist: A distribution over exactly the cells this spec kept, in
flattened order -- e.g. ``posterior(test_inputs)`` for inputs
built by :meth:`inputs_for`.
num_samples: If given, return this many joint draws from ``dist``
instead of its mean and variance.
key: PRNG key for the draws; required with ``num_samples``.
Returns:
By default, an ``xr.Dataset`` holding ``{target}_mean`` and
``{target}_variance`` over the grid. With ``num_samples``, one
variable ``{target}`` over ``("sample", *dims)``. Dropped cells are
NaN either way.
Raises:
ValueError: If ``dist`` does not match the number of kept cells, if
``num_samples`` is given without ``key``, or if samples are
requested from a distribution holding only marginal variances.
"""
n_points = dist.mean.shape[0]
if n_points != self.n_kept:
raise ValueError(
f"distribution has {n_points} points but this spec expects "
f"{self.n_kept}; was it built from a different grid?"
)
if num_samples is not None:
return self._samples(dist, num_samples, key)
variance_attrs = dict(self.target_attrs)
if "units" in variance_attrs:
variance_attrs["units"] = f"({variance_attrs['units']})^2"
return xr.Dataset(
{
f"{self.target}_mean": self._scatter(dist.mean, self.target_attrs),
f"{self.target}_variance": self._scatter(dist.variance, variance_attrs),
}
)
def _samples(
self,
dist: GaussianDistribution,
num_samples: int,
key: tp.Optional[KeyArray],
) -> xr.Dataset:
r"""Joint draws from ``dist``, scattered onto the grid per sample."""
if key is None:
raise ValueError("num_samples requires a PRNG key, e.g. key=jr.key(0)")
if lx.is_diagonal(dist.scale):
# Independent draws would ignore the spatial correlation, so any
# aggregate over cells (a region, a portfolio) would be overconfident.
raise ValueError(
"this distribution holds only marginal variances, so it cannot "
'give joint samples; predict with covariance="dense" to keep '
"the correlation between cells"
)
attrs = {
**self.target_attrs,
"description": "joint posterior predictive draws",
}
draws = self._scatter(dist.sample(key, (num_samples,)), attrs, ("sample",))
return xr.Dataset({self.target: draws})
def _scatter(
self,
values: Float[Array, "... N"],
attrs: dict[str, tp.Any],
leading_dims: tuple[str, ...] = (),
) -> xr.DataArray:
r"""Place per-kept-cell ``values`` into a NaN-filled labelled grid.
Any leading axes of ``values`` (e.g. samples) are kept as
``leading_dims`` in front of the grid's dims.
"""
values = np.asarray(values)
leading_shape = values.shape[:-1]
grid = np.full((*leading_shape, *self.mask.shape), np.nan)
grid[..., self.mask] = values
return xr.DataArray(
grid, coords=self.coords, dims=(*leading_dims, *self.dims), attrs=attrs
)
[docs]
def from_xarray(
obj: tp.Union[xr.Dataset, xr.DataArray],
target: str,
inputs: tp.Sequence[str],
*,
dropna: bool = True,
) -> tuple[Dataset, GridSpec]:
r"""Flatten labelled xarray data into a :class:`~gpjax.dataset.Dataset`.
Every input is broadcast onto the target's grid, so an input on fewer dims
(e.g. ``elevation(lat, lon)`` for a ``(time, lat, lon)`` target) repeats
along the rest. Datetime inputs become float days since their earliest
timestamp.
Args:
obj: The labelled data. A ``DataArray`` must be named, and that name is
the target.
target: Name of the data variable to model. Its dims define the grid.
inputs: Coordinates and/or data variables to use as inputs, in the order
of the columns of ``X``.
dropna: Drop grid cells where the target or any input is NaN. When
``False``, such cells raise instead.
Returns:
The flattened ``Dataset`` and the ``GridSpec`` needed to map
predictions back onto the grid.
Raises:
ValueError: If a name is missing, ``inputs`` is empty, repeats a name or
includes the target, an input has a dim the target lacks, a
``DataArray`` is unnamed, or no cells survive NaN handling (or any
NaN is present when ``dropna=False``).
TypeError: If ``inputs`` is a single string, or the target or an input
is not numeric (inputs may also be ``datetime64``).
"""
if isinstance(inputs, str):
raise TypeError(
f"inputs must be a list of names, e.g. [{inputs!r}], not a string"
)
if not inputs:
raise ValueError("inputs must name at least one coordinate or variable")
if len(set(inputs)) != len(inputs):
raise ValueError(f"inputs {list(inputs)} contain duplicate names")
if target in inputs:
raise ValueError(f"the target {target!r} cannot also be an input")
dataset = _as_dataset(obj)
target_values, *variables = _resolve(dataset, [target, *inputs])
dims = tuple(target_values.dims)
time_origins = {
variable.name: _earliest(variable.values)
for variable in variables
if np.issubdtype(variable.dtype, np.datetime64)
}
input_matrix = _input_matrix(variables, target_values, dims, time_origins)
outputs = _as_float(target_values.values.reshape(-1, 1), f"target {target!r}")
keep = np.isfinite(input_matrix).all(axis=1) & np.isfinite(outputs[:, 0])
n_dropped = int(keep.size - keep.sum())
if n_dropped and not dropna:
raise ValueError(
f"{n_dropped} cell(s) contain NaN in the target or an input; pass "
"dropna=True to drop them"
)
if not keep.any():
raise ValueError("no cells are left once NaN cells are dropped")
data = Dataset(X=jnp.asarray(input_matrix[keep]), y=jnp.asarray(outputs[keep]))
spec = GridSpec(
target=target,
target_attrs=dict(target_values.attrs),
inputs=tuple(inputs),
dims=dims,
coords=target_values.coords,
time_origins=time_origins,
mask=keep.reshape(target_values.shape),
n_dropped=n_dropped,
)
return data, spec
def _as_dataset(obj: tp.Union[xr.Dataset, xr.DataArray]) -> xr.Dataset:
r"""Promote a named ``DataArray`` to a one-variable ``Dataset``."""
if isinstance(obj, xr.Dataset):
return obj
if obj.name is None:
raise ValueError(
"a DataArray must have a name to be used as the target; set one with "
".rename('name') or pass an xr.Dataset"
)
return obj.to_dataset()
def _resolve(dataset: xr.Dataset, names: tp.Sequence[str]) -> list[xr.DataArray]:
r"""Look up coordinates or data variables by name, reporting every miss."""
missing = [name for name in names if name not in dataset.variables]
if missing:
raise ValueError(
f"{missing} not found; available coordinates and variables are "
f"{sorted(map(str, dataset.variables))}"
)
return [dataset[name] for name in names]
def _grid_coords(
dataset: xr.Dataset, dims: tuple[str, ...]
) -> dict[tp.Hashable, xr.DataArray]:
r"""Every coordinate of ``dataset`` that lives on the grid's dims."""
return {
name: coord
for name, coord in dataset.coords.items()
if set(coord.dims) <= set(dims)
}
def _earliest(timestamps: np.ndarray) -> np.datetime64:
r"""The earliest non-NaT timestamp, or NaT if there is none."""
valid = timestamps[~np.isnat(timestamps)]
return valid.min() if valid.size else np.datetime64("NaT")
def _input_matrix(
variables: list[xr.DataArray],
grid: xr.DataArray,
dims: tuple[str, ...],
time_origins: dict[str, np.datetime64],
) -> np.ndarray:
r"""Stack ``variables`` into an ``(N, D)`` float matrix over ``grid``.
A datetime input must have a time origin, and an input with a time origin
must still be a datetime: otherwise the same number would mean different
things in training and prediction.
"""
columns = []
for variable in variables:
name = variable.name
column = _column(variable, grid, dims)
is_datetime = np.issubdtype(column.dtype, np.datetime64)
if is_datetime != (name in time_origins):
then, now = ("a", "not a") if name in time_origins else ("not a", "a")
raise TypeError(
f"input {name!r} was {then} datetime when the spec was built but "
f"is {now} datetime now; encode it the same way in both"
)
if is_datetime:
days = (column - time_origins[name]) / np.timedelta64(1, "D")
columns.append(np.where(np.isnat(column), np.nan, days))
else:
columns.append(_as_float(column, f"input {name!r}"))
return np.stack(columns, axis=1)
def _column(
variable: xr.DataArray, grid: xr.DataArray, dims: tuple[str, ...]
) -> np.ndarray:
r"""Broadcast ``variable`` onto ``grid`` and flatten it in C order over ``dims``.
A variable on fewer dims than the grid repeats along the missing ones. A
variable on a dim the grid lacks has no cell to land in, so it is rejected.
"""
extra_dims = [dim for dim in variable.dims if dim not in dims]
if extra_dims:
raise ValueError(
f"input {variable.name!r} has dims {extra_dims} that the grid "
f"{dims} does not have"
)
return xr.broadcast(grid, variable)[1].transpose(*dims).values.ravel()
def _as_float(values: np.ndarray, description: str) -> np.ndarray:
r"""Cast numeric or boolean ``values`` to float; reject anything else."""
if np.issubdtype(values.dtype, np.number) or np.issubdtype(values.dtype, np.bool_):
return values.astype(float)
raise TypeError(
f"{description} has dtype {values.dtype}, which is not numeric; convert "
"it before calling from_xarray"
)
__all__ = [
"GridSpec",
"from_xarray",
]