Gridded Data with xarray#

Download this notebook: xarray_workflow.ipynb

Climate and environmental data rarely arrive as a tidy matrix. They arrive as labelled xarray objects: a temperature field over latitude and longitude, a covariate such as elevation, and gaps where a sensor failed or a cloud covered the scene. A GP in GPJax, on the other hand, consumes a Dataset of flat inputs \(\mathbf{X} \in \mathbb{R}^{N \times D}\) and outputs \(\mathbf{y} \in \mathbb{R}^{N \times 1}\).

The gpjax.xarray module converts between the two at the edges of a workflow. In this notebook we

  1. flatten a gappy, labelled field into a Dataset with from_xarray,

  2. fit a GP exactly as we would on any other Dataset,

  3. build inputs for a finer prediction grid with GridSpec.inputs_for, and

  4. map the predictions, and joint posterior samples, back onto that grid with GridSpec.to_xarray.

The module needs the optional extra: pip install "gpjax[xarray]".

from jax import config
import jax.numpy as jnp
import jax.random as jr
from jaxtyping import install_import_hook
import matplotlib.pyplot as plt
import numpy as np
from utils import use_mpl_style
import xarray as xr

config.update("jax_enable_x64", True)

with install_import_hook("gpjax", "beartype.beartype"):
    import gpjax as gpx
    from gpjax.xarray import from_xarray

key = jr.key(42)
use_mpl_style()

A synthetic temperature field#

We simulate near-surface temperature on a regional latitude-longitude grid. It cools towards the pole and with elevation (a lapse rate of roughly 6.5 K per kilometre), with a smooth large-scale anomaly on top. Elevation is a separate variable over the same grid. The data are synthetic, so the notebook needs no download and every number in it can be checked against the truth.

def elevation_at(lat, lon):
    """A single mountain range, in metres."""
    return 2500.0 * np.exp(-(((lon - 12.0) / 4.0) ** 2) - ((lat - 47.0) / 3.0) ** 2)


def temperature_at(lat, lon, elevation):
    """Temperature in kelvin: latitude gradient, lapse rate and an anomaly."""
    anomaly = 1.5 * np.sin(lon / 3.0) * np.cos(lat / 4.0)
    return 290.0 - 0.6 * (lat - 40.0) - 0.0065 * elevation + anomaly


def regional_field(lats, lons) -> xr.Dataset:
    lat_grid, lon_grid = np.meshgrid(lats, lons, indexing="ij")
    elevation = elevation_at(lat_grid, lon_grid)
    return xr.Dataset(
        {
            "t2m": (
                ("lat", "lon"),
                temperature_at(lat_grid, lon_grid, elevation),
                {"units": "K", "long_name": "2 m air temperature"},
            ),
            "elevation": (("lat", "lon"), elevation, {"units": "m"}),
        },
        coords={
            "lat": ("lat", lats, {"units": "degrees_north"}),
            "lon": ("lon", lons, {"units": "degrees_east"}),
        },
    )


coarse = regional_field(np.linspace(40.0, 54.0, 12), np.linspace(2.0, 22.0, 16))

Real observations have holes. We knock out a block of cells, as a cloud would, and add a little measurement noise to the rest.

key, noise_key = jr.split(key)
noise = 0.2 * np.asarray(jr.normal(noise_key, coarse["t2m"].shape))
observed = coarse.copy(deep=True)
observed["t2m"] = observed["t2m"] + noise
observed["t2m"].attrs = coarse["t2m"].attrs
observed["t2m"][4:7, 9:13] = np.nan

observed["t2m"].plot(cmap="coolwarm")
plt.title("Observed temperature (gaps in white)")
plt.show()
../_images/d97953d0a447bb8ead66741d0ce7907b943fdd7132ffabd83bdd7ebea27447e6.png

From labelled data to a Dataset#

from_xarray takes the target variable and the inputs we want the GP to depend on. Inputs can be coordinates (lat, lon) or other data variables (elevation), and the columns of \(\mathbf{X}\) follow the order we list them in. Cells where the target or any input is NaN are dropped by default, and the returned GridSpec records which ones.

inputs = ["lat", "lon", "elevation"]
data, spec = from_xarray(observed, target="t2m", inputs=inputs)

print(data)
print(spec)
Dataset(Number of observations: 180 - Input dimension: 3)
GridSpec(target='t2m', inputs=('lat', 'lon', 'elevation'), grid={'lat': 12, 'lon': 16}, kept=180, dropped=12)

data is an ordinary Dataset, so nothing downstream knows it came from xarray. The GridSpec stays with us, outside the model, until we want labelled output again.

Fitting the model#

Temperature varies over hundreds of kilometres in latitude and longitude but over hundreds of metres in elevation, so we give the RBF kernel one lengthscale per input. A constant mean absorbs the ~285 K offset.

prior = gpx.gps.Prior(
    mean_function=gpx.mean_functions.Constant(jnp.array([285.0])),
    kernel=gpx.kernels.RBF(lengthscale=jnp.array([3.0, 3.0, 1000.0]), variance=25.0),
)
model = prior * gpx.likelihoods.Gaussian(obs_stddev=jnp.array(0.5))

model, history = gpx.fit_scipy(
    model=model,
    objective=lambda candidate, train_data: -gpx.objectives.conjugate_mll(
        candidate, train_data
    ),
    train_data=data,
    verbose=False,
)

Predicting on a finer grid#

spec.inputs_for builds the prediction inputs for any grid that holds the same input variables, encoded exactly as in training. Here we predict on a grid four times finer in each direction, including the cells that were missing from the observations. We pass the likelihood’s predictive distribution, so the variance includes observation noise.

fine = regional_field(np.linspace(40.0, 54.0, 45), np.linspace(2.0, 22.0, 61))
test_inputs, test_spec = spec.inputs_for(fine[["elevation"]])

posterior = model.condition(data)
predictive = model.likelihood(posterior(test_inputs))
prediction = test_spec.to_xarray(predictive)
prediction
<xarray.Dataset> Size: 45kB
Dimensions:       (lat: 45, lon: 61)
Coordinates:
  * lat           (lat) float64 360B 40.0 40.32 40.64 40.95 ... 53.36 53.68 54.0
  * lon           (lon) float64 488B 2.0 2.333 2.667 3.0 ... 21.33 21.67 22.0
Data variables:
    t2m_mean      (lat, lon) float64 22kB 289.2 289.1 289.0 ... 282.4 282.4
    t2m_variance  (lat, lon) float64 22kB 0.06723 0.06189 ... 0.06195 0.06734

The result is a labelled xr.Dataset on the fine grid, with the target’s attributes carried over. The variance is in \(\mathrm{K}^2\). Everything xarray offers, from plotting to to_netcdf, works on it directly.

fig, (mean_ax, std_ax, error_ax) = plt.subplots(1, 3, figsize=(15, 4))
prediction["t2m_mean"].plot(ax=mean_ax, cmap="coolwarm")
mean_ax.set_title("Predictive mean")
np.sqrt(prediction["t2m_variance"]).plot(ax=std_ax, cmap="viridis")
std_ax.set_title("Predictive standard deviation")
(prediction["t2m_mean"] - fine["t2m"]).plot(ax=error_ax, cmap="RdBu_r", center=0.0)
error_ax.set_title("Error against the true field")
for ax in (mean_ax, std_ax, error_ax):
    ax.add_patch(
        plt.Rectangle(
            (observed.lon[9], observed.lat[4]),
            float(observed.lon[12] - observed.lon[9]),
            float(observed.lat[6] - observed.lat[4]),
            fill=False,
            linestyle="--",
        )
    )
plt.show()
../_images/e8b00ab90b0df13129a0d210a28318b321cc0216e7193abfeb5f505f402745ed.png

The standard deviation grows inside the dashed box where observations were missing, and the error stays small across the mountain range because elevation is an input.

Joint samples and regional averages#

The mean and variance describe each cell on its own. Many questions are about several cells together, such as the average temperature over the Alpine box \(\mathcal{R}\). Its variance depends on the covariance between the cells,

(1)#\[ \operatorname{Var}\Big[\tfrac{1}{|\mathcal{R}|} \sum_{i \in \mathcal{R}} f_i\Big] = \tfrac{1}{|\mathcal{R}|^2} \sum_{i, j \in \mathcal{R}} \operatorname{Cov}[f_i, f_j], \]

which the per-cell variances alone cannot give. Passing num_samples to to_xarray draws from the joint predictive distribution instead, and returns the draws with a leading sample dimension. Averaging each draw over the region gives samples of the regional mean.

latent = posterior(test_inputs)  # the field itself, without observation noise
key, sample_key = jr.split(key)
samples = test_spec.to_xarray(latent, num_samples=500, key=sample_key)

alps = dict(lat=slice(45.0, 49.0), lon=slice(8.0, 16.0))
regional_mean = samples["t2m"].sel(**alps).mean(["lat", "lon"])
true_regional_mean = float(fine["t2m"].sel(**alps).mean())

# The same latent distribution, but treating the cells as independent.
latent_variance = test_spec.to_xarray(latent)["t2m_variance"].sel(**alps)
joint_std = float(regional_mean.std("sample"))
naive_std = float(np.sqrt(latent_variance.sum()) / latent_variance.size)

print(f"Regional mean: {float(regional_mean.mean()):.2f} K (truth {true_regional_mean:.2f} K)")
print(f"Standard deviation from joint samples:        {joint_std:.3f} K")
print(f"Standard deviation if cells were independent: {naive_std:.3f} K")
Regional mean: 275.11 K (truth 275.03 K)
Standard deviation from joint samples:        0.065 K
Standard deviation if cells were independent: 0.007 K

Treating the cells as independent understates the uncertainty in the regional average by roughly an order of magnitude, because neighbouring cells tend to be wrong in the same direction. Against the joint standard deviation the true regional mean is a plausible outcome; against the independent one it would look like a many-sigma surprise. Joint samples keep that correlation, which is why to_xarray refuses to draw samples from a distribution that only holds marginal variances (as returned by posterior(test_inputs, covariance="diagonal")).

System configuration#

%reload_ext watermark
%watermark -n -u -v -iv -w -a 'Thomas Pinder'
Author: Thomas Pinder

Last updated: Mon, 28 Sep 2026

Python implementation: CPython
Python version       : 3.11.16
IPython version      : 9.17.1

gpjax     : 1.0.0
jax       : 0.10.2
jaxtyping : 0.3.11
matplotlib: 3.11.2
numpy     : 2.4.6
xarray    : 2026.7.0

Watermark: 2.6.0