diff --git a/CHANGELOG.md b/CHANGELOG.md index 911fd3019..ce8746280 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,30 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- **`gpjax.xarray` input transforms.** `from_xarray(..., transforms=[...])` + turns the named inputs into the columns of `X`, and the `GridSpec` applies the + same fitted transforms to every new grid. `Standardise` scales inputs with the + training mean and standard deviation. `UnitSphere` maps latitude and longitude + to a point on the unit sphere, so a stationary kernel is valid on the whole + globe. `Cyclic` encodes a periodic input, such as the seasonal cycle, as a + point on a circle. `GridSpec.columns` names the resulting columns. +- **`GridSpec.predict`: chunked prediction on large grids.** + `spec.predict(lambda x: posterior(x, covariance="diagonal"), grid)` gives the + predictive mean and variance as an `xr.Dataset`, with at most `chunk_size` + cells in memory at once and one compilation. A Dask-backed grid gives a lazy + result that is predicted block by block. +- **New example: *Infilling Global Surface Temperature*.** It reads a netCDF + file of the 2024 temperature anomaly from the NCEP-NCAR Reanalysis 1, which is + complete, and removes cells to test the infill against the truth. Cells removed + at random are filled well, and joint samples give a global mean whose interval + holds the truth. Cells removed because they are warm show how a GP is biased, + and overconfident, when data are missing not at random. +- **The xarray introduction is now *Working with Gridded Data*, under Getting + started.** It uses `Standardise` and `GridSpec.predict`, and shows + `UnitSphere` and `Cyclic`. Its URL is unchanged. + ### Changed - Allow JAX and JAXlib 0.11 in downstream environments by removing the diff --git a/docs/examples/data/_pull_reference_datasets.py b/docs/examples/data/_pull_reference_datasets.py index 6811045b3..efff6cc04 100644 --- a/docs/examples/data/_pull_reference_datasets.py +++ b/docs/examples/data/_pull_reference_datasets.py @@ -10,8 +10,9 @@ uv run --extra docs python docs/examples/data/_pull_reference_datasets.py The ``--extra docs`` is only needed for the UCI Auto MPG pull, which uses -``ucimlrepo``; the other three pulls need nothing beyond ``pandas`` and -``requests``. +``ucimlrepo``. The NCEP reanalysis pull reads netCDF4, so it also needs +``--with h5netcdf --with h5py``. The other pulls need nothing beyond ``pandas`` +and ``requests``. Data sources ------------ @@ -32,11 +33,23 @@ https://archive.ics.uci.edu/dataset/9/auto-mpg (fetched via ``ucimlrepo``). Licence: CC BY 4.0. The features and the target are concatenated into a single frame so the notebook can split them back out without ``ucimlrepo``. +- NCEP-NCAR Reanalysis 1 temperature anomaly (``ncep_air_anomaly_2024.nc``): + the 2024 annual-mean anomaly of near-surface (0.995 sigma) air temperature, + relative to the 1991-2020 mean of each grid cell, on the native 2.5-degree + grid. A reanalysis has a value in every cell, so the field is complete. + Source: https://psl.noaa.gov/data/gridded/data.ncep.reanalysis.html (monthly + means, ``air.mon.mean.nc``). Provider: NOAA Physical Sciences Laboratory + (Kalnay et al., 1996, https://doi.org/10.1175/1520-0477(1996)077<0437:TNYRP>2.0.CO;2). + Work of the US Government, so public domain; PSL ask for the acknowledgement + "NCEP-NCAR Reanalysis 1 data provided by the NOAA PSL, Boulder, Colorado, USA, + from their website at https://psl.noaa.gov". Written as netCDF3, which xarray + reads with SciPy, so the notebook needs no netCDF4 library. """ from __future__ import annotations from pathlib import Path +import tempfile import time import pandas as pd @@ -104,8 +117,81 @@ def pull_auto_mpg() -> None: _save(pd.concat([features, targets], axis=1), "auto_mpg.csv") +# --------------------------------------------------------------------------- # +# xarray_workflow — NCEP-NCAR Reanalysis 1 temperature anomaly for 2024. # +# --------------------------------------------------------------------------- # +NCEP_URL = ( + "https://psl.noaa.gov/thredds/fileServer/Datasets/ncep.reanalysis/" + "Monthlies/surface/air.mon.mean.nc" +) + + +def _download_resumable(url: str, path: Path, attempts: int = 8) -> None: + """Download ``url`` to ``path``, resuming when the server cuts it short.""" + for _ in range(attempts): + start = path.stat().st_size if path.exists() else 0 + headers = {"User-Agent": "gpjax-docs", "Range": f"bytes={start}-"} + try: + with requests.get(url, headers=headers, stream=True, timeout=120) as resp: + if resp.status_code == 416: # nothing left to fetch + return + resp.raise_for_status() + total = start + int(resp.headers["Content-Length"]) + mode = "ab" if resp.status_code == 206 else "wb" + with path.open(mode) as file: + for chunk in resp.iter_content(1 << 20): + file.write(chunk) + if path.stat().st_size >= total: + return + except requests.RequestException as err: + print(f" retrying after: {err}") + time.sleep(2) + raise RuntimeError(f"Failed to fetch {url} in {attempts} attempts") + + +def pull_ncep_anomaly(nc_path: Path | None = None) -> None: + """Save the 2024 NCEP anomaly; pass ``nc_path`` to reuse a download.""" + print("xarray_workflow: NCEP-NCAR Reanalysis 1 temperature anomaly, 2024") + import xarray as xr + + with tempfile.TemporaryDirectory() as tmp: + if nc_path is None: + nc_path = Path(tmp) / "air.mon.mean.nc" + _download_resumable(NCEP_URL, nc_path) + with xr.open_dataset(nc_path, engine="h5netcdf") as monthly: + annual = monthly["air"].resample(time="YS").mean().load() + climatology = annual.sel(time=slice("1991", "2020")).mean("time") + anomaly = annual.sel(time="2024").squeeze("time", drop=True) - climatology + # Longitude from 0..357.5 to -180..177.5, so maps are centred on Greenwich. + anomaly = anomaly.assign_coords(lon=(anomaly["lon"] + 180.0) % 360.0 - 180.0) + anomaly = anomaly.sortby(["lat", "lon"]).astype("float32") + anomaly.attrs = { + "long_name": "Near-surface air temperature anomaly", + "units": "K", + "cell_methods": "time: mean (2024, anomaly relative to 1991-2020)", + } + anomaly["lat"].attrs = {"standard_name": "latitude", "units": "degrees_north"} + anomaly["lon"].attrs = {"standard_name": "longitude", "units": "degrees_east"} + dataset = anomaly.to_dataset(name="tas_anomaly") + dataset.attrs = { + "title": "2024 annual-mean near-surface air temperature anomaly", + "source": "NCEP-NCAR Reanalysis 1, monthly air.sig995 (air.mon.mean.nc)", + "references": "Kalnay et al. (1996), doi:10.1175/1520-0477(1996)077<0437:TNYRP>2.0.CO;2", + "acknowledgement": ( + "NCEP-NCAR Reanalysis 1 data provided by the NOAA PSL, Boulder, " + "Colorado, USA, from their website at https://psl.noaa.gov" + ), + "history": "Annual means of the monthly means; 2024 minus the 1991-2020 mean.", + "Conventions": "CF-1.8", + } + path = HERE / "ncep_air_anomaly_2024.nc" + dataset.to_netcdf(path, engine="scipy", format="NETCDF3_64BIT") + print(f" wrote {path.name}: {dict(dataset.sizes)}") + + if __name__ == "__main__": pull_mauna_loa_co2() pull_gulf_velocities() pull_auto_mpg() + pull_ncep_anomaly() print("\nDone.") diff --git a/docs/examples/data/ncep_air_anomaly_2024.nc b/docs/examples/data/ncep_air_anomaly_2024.nc new file mode 100644 index 000000000..01ffe6d46 Binary files /dev/null and b/docs/examples/data/ncep_air_anomaly_2024.nc differ diff --git a/docs/examples/infilling_surface_temperature.py b/docs/examples/infilling_surface_temperature.py new file mode 100644 index 000000000..683371677 --- /dev/null +++ b/docs/examples/infilling_surface_temperature.py @@ -0,0 +1,379 @@ +# --- +# jupyter: +# jupytext: +# cell_metadata_filter: -all +# custom_cell_magics: kql +# text_representation: +# extension: .py +# format_name: percent +# format_version: '1.3' +# jupytext_version: 1.19.1 +# kernelspec: +# display_name: Python 3 +# language: python +# name: python3 +# --- + +# %% [markdown] +# # Infilling Global Surface Temperature +# +# Download this notebook: {nb-download}`infilling_surface_temperature.ipynb` +# +# Temperature records have gaps: over the poles, over parts of Africa and the +# Southern Ocean, and, for satellites, wherever there is cloud. Climate +# scientists fill these gaps to map the field and to compute the global mean +# temperature. In this notebook we fill gaps with a GP, and we test the result. +# +# To know how good the filled values are, we start from a field that is complete, +# remove cells ourselves, and compare the GP with the cells that we removed. We +# +# 1. fill gaps that are random, map the field, and get the global mean +# temperature with its uncertainty from joint posterior samples, and +# 2. show what goes wrong when cells are missing *because of* their values. +# +# The notebook uses the [`gpjax.xarray`](../reference/xarray.md) module, which +# [Working with Gridded Data](xarray_workflow.py) introduces. It needs the +# optional extra: `pip install "gpjax[xarray]"`. + +# %% tags=["remove-cell"] +import logging + +# The build host may not have the ledger fonts (Public Sans, Spectral); hide +# matplotlib's font fallback messages. +logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR) + +# %% +from pathlib import Path + +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 +import xarray as xr + +config.update("jax_enable_x64", True) + +with install_import_hook("gpjax", "beartype.beartype"): + import gpjax as gpx + from gpjax.parameters import val + from gpjax.xarray import ( + UnitSphere, + from_xarray, + ) + +key = jr.key(42) +rng = np.random.default_rng(0) +gpx.plotting.use_style() + +# %% [markdown] +# ## A complete temperature field +# +# We use the [NCEP-NCAR Reanalysis 1](https://psl.noaa.gov/data/gridded/data.ncep.reanalysis.html) +# (Kalnay et al., 1996). A reanalysis combines a weather model with the +# observations, so it has a value in every grid cell. The file holds the 2024 +# annual-mean near-surface air temperature as an *anomaly*: the difference from +# the 1991–2020 mean of the same cell. Anomalies remove the large, fixed +# differences between the equator and the poles, so what remains is the +# signal of interest, and it varies smoothly over large distances. +# +# NCEP-NCAR Reanalysis 1 data provided by the NOAA PSL, Boulder, Colorado, USA, +# from their website at . + +# %% +nc_candidates = [ + Path("docs/examples/data/ncep_air_anomaly_2024.nc"), + Path("data/ncep_air_anomaly_2024.nc"), +] +nc_path = next(path for path in nc_candidates if path.exists()) +reanalysis = xr.open_dataset(nc_path) +reanalysis + +# %% [markdown] +# The native grid is 2.5°. We fit on a 5° grid, which is every second grid point, +# and predict back on the 2.5° grid, so the predictions are also tested at points +# the model never saw. + +# %% +truth_fine = reanalysis["tas_anomaly"].astype(float) +truth = truth_fine.isel(lat=slice(1, None, 2), lon=slice(1, None, 2)) +print(f"5° grid: {dict(truth.sizes)}, 2.5° grid: {dict(truth_fine.sizes)}") + +anomaly_style = dict(cmap="RdBu_r", vmin=-4.0, vmax=4.0) +error_style = dict( + cmap="PuOr_r", vmin=-2.0, vmax=2.0, cbar_kwargs={"label": "Error [K]"} +) +truth_fine.plot(figsize=(7, 3.5), **anomaly_style) +plt.title("2024 anomaly, NCEP-NCAR Reanalysis 1") +plt.show() + +# %% [markdown] +# A 5° cell near a pole is much smaller than one at the equator, so a global mean +# weights each cell by the cosine of its latitude. + + +# %% +def global_mean(field: xr.DataArray) -> xr.DataArray: + """Area-weighted mean over the cells that have a value.""" + weights = np.cos(np.deg2rad(field["lat"])) + return field.weighted(weights).mean(["lat", "lon"]) + + +print(f"True global mean anomaly: {float(global_mean(truth)):.2f} K") + +# %% [markdown] +# ## Cells missing at random +# +# First we keep a random 25% of the 5° cells. Which cells are missing has nothing +# to do with their values. Statisticians call this *missing completely at +# random*. + +# %% +kept_at_random = rng.uniform(size=truth.shape) < 0.25 +observed = truth.where(kept_at_random).rename("tas").to_dataset() +print(f"{int(kept_at_random.sum())} of {truth.size} cells kept") + +# %% [markdown] +# ### Inputs on the sphere +# +# Latitude and longitude in degrees are not good GP inputs for a global field. A +# 5° cell is about 555 km wide at the equator but only about 24 km wide next to a +# pole, and longitude 177.5°E is next to 177.5°W. The +# [`UnitSphere`](#gpjax.xarray.UnitSphere) transform replaces `lat` and `lon` with +# the three coordinates of a point on the unit sphere. A stationary kernel on these +# coordinates uses the chord distance through the Earth, so it has no seam at the +# antimeridian and no distortion at the poles. Its lengthscale is in Earth radii. + +# %% +data, spec = from_xarray( + observed, target="tas", inputs=["lat", "lon"], transforms=[UnitSphere()] +) +print(data) +print(spec) + +# %% [markdown] +# ### Fitting the model +# +# We use a Matérn-3/2 kernel, which gives rougher fields than an RBF kernel, as +# temperature anomalies are. A constant mean absorbs the global warming signal. + + +# %% +def fit_model(train_data: gpx.Dataset): + prior = gpx.gps.Prior( + mean_function=gpx.mean_functions.Constant(jnp.array([0.5])), + kernel=gpx.kernels.Matern32(lengthscale=jnp.array(0.3), variance=1.0), + ) + model = prior * gpx.likelihoods.Gaussian(obs_stddev=jnp.array(0.1)) + model, _ = gpx.fit_scipy( + model=model, + objective=lambda candidate, train_data: ( + -gpx.objectives.conjugate_mll(candidate, train_data) + ), + train_data=train_data, + verbose=False, + ) + return model + + +earth_radius_km = 6371.0 +model = fit_model(data) +lengthscale_km = float(val(model.prior.kernel.lengthscale)) * earth_radius_km +print(f"Lengthscale: {lengthscale_km:.0f} km") + +# %% [markdown] +# ### Predicting on the finer grid +# +# `spec.predict` gives the predictive mean and variance on the 2.5° grid of the +# reanalysis, in chunks of `chunk_size` cells. We use a diagonal covariance, +# because we need only the variance of each cell. + +# %% +posterior = model.condition(data) +prediction = spec.predict( + lambda x: model.likelihood(posterior(x, covariance="diagonal")), + truth_fine.to_dataset(), + chunk_size=2048, +) +prediction + +# %% [markdown] +# Because the field is complete, we can score every prediction. We compare with a +# simple baseline, the mean of the kept cells, and check the uncertainty: about +# 95% of the true values should be within two predictive standard deviations. + + +# %% +def score(prediction: xr.Dataset, observed: xr.Dataset) -> None: + error = prediction["tas_mean"] - truth_fine + baseline_error = float(observed["tas"].mean()) - truth_fine + within = abs(error) < 2 * np.sqrt(prediction["tas_variance"]) + print(f"RMSE, GP: {float(np.sqrt((error**2).mean())):.2f} K") + print( + f"RMSE, mean of kept cells: {float(np.sqrt((baseline_error**2).mean())):.2f} K" + ) + print(f"Within 2 sd: {float(within.mean()):.0%}") + + +def plot_infill(prediction: xr.Dataset, observed: xr.Dataset) -> None: + fig, axes = plt.subplots(1, 3, figsize=(15, 3.5), sharey=True) + observed["tas"].plot(ax=axes[0], **anomaly_style) + axes[0].set_title("Kept cells") + prediction["tas_mean"].plot(ax=axes[1], **anomaly_style) + axes[1].set_title("Predictive mean") + (prediction["tas_mean"] - truth_fine).plot(ax=axes[2], **error_style) + axes[2].set_title("Predictive mean minus truth") + for ax in axes[1:]: + ax.set_ylabel("") + plt.show() + + +score(prediction, observed) +plot_infill(prediction, observed) + +# %% [markdown] +# From a quarter of the cells, the GP recovers the large-scale pattern, including +# the strong warmth over the Arctic. The errors are largest where the field +# changes over short distances, and the uncertainty is about right. +# +# ### The global mean, with joint samples +# +# The global mean anomaly fills each missing cell and averages. Its variance +# depends on the covariance between the filled cells, +# +# $$ +# \operatorname{Var}\Big[\sum_{i} w_i f_i\Big] +# = \sum_{i, j} w_i w_j \operatorname{Cov}[f_i, f_j], +# $$ (eq-xarray-global-variance) +# +# which the per-cell variances alone cannot give. For this we need the full +# covariance, so we build all the prediction inputs on the 5° grid at once with +# `spec.inputs_for`. Passing `num_samples` to `to_xarray` then draws from the joint +# distribution of the field, and returns the draws with a leading `sample` +# dimension. `combine_first` keeps the kept cells and takes each missing cell from +# a draw. + + +# %% +def global_mean_samples(posterior, spec, observed: xr.Dataset, key) -> xr.DataArray: + test_inputs, test_spec = spec.inputs_for(truth.to_dataset()) + samples = test_spec.to_xarray(posterior(test_inputs), num_samples=500, key=key) + return global_mean(observed["tas"].combine_first(samples["tas"])) + + +def report_global_mean(means: xr.DataArray, observed: xr.Dataset) -> None: + print(f"Truth: {float(global_mean(truth)):.2f} K") + print(f"Mean of kept cells: {float(global_mean(observed['tas'])):.2f} K") + print( + f"Gaps filled by GP: {float(means.mean()):.2f}" + f" ± {2 * float(means.std('sample')):.2f} K (2 sd)" + ) + + +key, sample_key = jr.split(key) +means = global_mean_samples(posterior, spec, observed, sample_key) +report_global_mean(means, observed) + +# %% [markdown] +# With random gaps, even the mean of the kept cells is close to the truth, and the +# GP gives the global mean with an interval that holds it. +# +# ## Cells missing because of their values +# +# Real gaps are rarely random. A satellite cannot see the surface through cloud, +# a station fails in extreme weather, and the polar regions, which warm fastest, +# have the fewest observations. When the chance that a cell is missing depends on +# the value that is missing, the data are *missing not at random*. +# +# We simulate this. We again aim to keep a quarter of the cells, but now the +# warmer a cell's anomaly, the less likely it is to be kept. + +# %% +standardised = ((truth - truth.mean()) / truth.std()).values +odds = np.exp(-1.5 * standardised) +keep_probability = np.clip(0.25 * odds / odds.mean(), 0.0, 1.0) +kept_selectively = rng.uniform(size=truth.shape) < keep_probability +selective = truth.where(kept_selectively).rename("tas").to_dataset() +print(f"{int(kept_selectively.sum())} of {truth.size} cells kept") + +# %% [markdown] +# The steps are the same as before. + +# %% +data_selective, spec_selective = from_xarray( + selective, target="tas", inputs=["lat", "lon"], transforms=[UnitSphere()] +) +model_selective = fit_model(data_selective) +posterior_selective = model_selective.condition(data_selective) +prediction_selective = spec_selective.predict( + lambda x: model_selective.likelihood(posterior_selective(x, covariance="diagonal")), + truth_fine.to_dataset(), + chunk_size=2048, +) + +score(prediction_selective, selective) +plot_infill(prediction_selective, selective) + +key, sample_key = jr.split(key) +means_selective = global_mean_samples( + posterior_selective, spec_selective, selective, sample_key +) +report_global_mean(means_selective, selective) + +# %% [markdown] +# The kept cells are mostly the cold ones, so their mean is far too cold. The GP +# removes part of this bias, because it fills each gap from its neighbours and +# the warm regions still have some kept cells. But the result is still too cold, +# and the interval does not hold the truth: the GP is confidently wrong. +# +# The reason is that a GP conditions only on the values it sees. Its prior has one +# constant mean, which it learns from the kept cells, so the cold kept cells pull +# that mean down. It also has no way to know that the missing cells are warm, +# because nothing in its inputs says so. Its uncertainty describes the spread of +# values that are *consistent with the kept cells*, not the error of the +# selection. +# +# In real data we cannot see this bias, because we do not have the missing values. +# Useful steps are: +# +# - Add inputs that explain why cells are missing, such as cloud fraction or a +# covariate that is observed everywhere. If the chance of a gap depends only on +# the inputs, the gaps are *missing at random* given those inputs, and the GP +# can correct for them. +# - Model the observation process together with the field. +# - Test how sensitive the result is to different assumptions about the missing +# values, as we did here with a complete field. +# +# ## Adding an input that explains the gaps +# +# If a variable that explains the gaps is observed everywhere, add it as an input. +# [`Standardise`](#gpjax.xarray.Standardise) puts it on the same scale as the +# sphere coordinates: +# +# ```python +# data, spec = from_xarray( +# observed, +# target="tas", +# inputs=["lat", "lon", "cloud_fraction"], +# transforms=[UnitSphere(), Standardise(["cloud_fraction"])], +# ) +# ``` +# +# For a monthly record, [`Cyclic`](#gpjax.xarray.Cyclic) adds the seasonal cycle; +# see [Working with Gridded Data](xarray_workflow.py). +# +# ## References +# +# Kalnay, E., Kanamitsu, M., Kistler, R., Collins, W., Deaven, D., Gandin, L., +# Iredell, M., Saha, S., White, G., Woollen, J., Zhu, Y., Chelliah, M., Ebisuzaki, +# W., Higgins, W., Janowiak, J., Mo, K. C., Ropelewski, C., Wang, J., Leetmaa, A., +# Reynolds, R., Jenne, R. and Joseph, D. (1996). The NCEP/NCAR 40-year reanalysis +# project. *Bulletin of the American Meteorological Society*, 77(3), 437–471. +# [doi:10.1175/1520-0477(1996)077<0437:TNYRP>2.0.CO;2](https://doi.org/10.1175/1520-0477(1996)077%3C0437:TNYRP%3E2.0.CO;2) +# +# ## System configuration + +# %% +# %reload_ext watermark +# %watermark -n -u -v -iv -w -a 'Thomas Pinder' diff --git a/docs/examples/xarray_workflow.py b/docs/examples/xarray_workflow.py index 0801884ec..99fadc9e8 100644 --- a/docs/examples/xarray_workflow.py +++ b/docs/examples/xarray_workflow.py @@ -15,7 +15,7 @@ # --- # %% [markdown] -# # Gridded Data with xarray +# # Working with Gridded Data # # Download this notebook: {nb-download}`xarray_workflow.ipynb` # @@ -32,11 +32,15 @@ # 1. flatten a gappy, labelled field into a `Dataset` with # [`from_xarray`](#gpjax.xarray.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`](#gpjax.xarray.GridSpec.inputs_for), and -# 4. map the predictions, and joint posterior samples, back onto that grid with +# 3. predict the mean and variance on a finer grid with +# [`GridSpec.predict`](#gpjax.xarray.GridSpec.predict), and +# 4. draw joint posterior samples on that grid with +# [`GridSpec.inputs_for`](#gpjax.xarray.GridSpec.inputs_for) and # [`GridSpec.to_xarray`](#gpjax.xarray.GridSpec.to_xarray). # +# For the same workflow on real data, see +# [Infilling Global Surface Temperature](infilling_surface_temperature.py). +# # The module needs the optional extra: `pip install "gpjax[xarray]"`. # %% tags=["remove-cell"] @@ -59,7 +63,10 @@ with install_import_hook("gpjax", "beartype.beartype"): import gpjax as gpx - from gpjax.xarray import from_xarray + from gpjax.xarray import ( + Standardise, + from_xarray, + ) key = jr.key(42) gpx.plotting.use_style() @@ -131,10 +138,19 @@ def regional_field(lats, lons) -> xr.Dataset: # (`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. +# +# Temperature varies over degrees of latitude and longitude, but over hundreds of +# metres of elevation. The [`Standardise`](#gpjax.xarray.Standardise) transform +# scales each input to zero mean and unit standard deviation over the training +# cells, so one starting lengthscale suits every input. The spec keeps the +# training mean and standard deviation, and applies them to every grid that we +# predict on. # %% inputs = ["lat", "lon", "elevation"] -data, spec = from_xarray(observed, target="t2m", inputs=inputs) +data, spec = from_xarray( + observed, target="t2m", inputs=inputs, transforms=[Standardise()] +) print(data) print(spec) @@ -146,14 +162,13 @@ def regional_field(lats, lons) -> xr.Dataset: # # ## 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. +# We give the RBF kernel one lengthscale per input, in units of that input's +# standard deviation. 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), + kernel=gpx.kernels.RBF(lengthscale=jnp.ones(3), variance=25.0), ) model = prior * gpx.likelihoods.Gaussian(obs_stddev=jnp.array(0.5)) @@ -169,19 +184,28 @@ def regional_field(lats, lons) -> xr.Dataset: # %% [markdown] # ## 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. +# `spec.predict` gives the predictive mean and variance on 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. The function we pass maps a block of inputs to a +# distribution. We use the likelihood's predictive distribution, so the variance +# includes observation noise, and a diagonal covariance, because we need only the +# variance of each cell. +# +# `spec.predict` sends the cells to this function in chunks of `chunk_size`, so +# the memory use stays the same for a global grid. If the grid holds Dask +# arrays, the result is lazy, and each Dask block is predicted only when it is +# computed or written with `to_netcdf`. # %% 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 = spec.predict( + lambda x: model.likelihood(posterior(x, covariance="diagonal")), + fine[["elevation"]], + chunk_size=1024, +) prediction # %% [markdown] @@ -225,12 +249,15 @@ def regional_field(lats, lons) -> xr.Dataset: # = \tfrac{1}{|\mathcal{R}|^2} \sum_{i, j \in \mathcal{R}} \operatorname{Cov}[f_i, f_j], # $$ (eq-xarray-regional-variance) # -# 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. +# which the per-cell variances alone cannot give. For this we need the full +# covariance, so we build all the prediction inputs at once with +# `spec.inputs_for`. Passing `num_samples` to `to_xarray` then draws from the joint +# predictive distribution, and returns the draws with a leading `sample` +# dimension. Averaging each draw over the region gives samples of the regional +# mean. # %% +test_inputs, test_spec = spec.inputs_for(fine[["elevation"]]) 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) @@ -258,6 +285,34 @@ def regional_field(lats, lons) -> xr.Dataset: # from a distribution that only holds marginal variances (as returned by # `posterior(test_inputs, covariance="diagonal")`). # +# ## Global grids and seasonal cycles +# +# Latitude and longitude in degrees are not good inputs for a global field. A +# degree of longitude is about 111 km at the equator but only 38 km at 70°N, and +# longitude 359° is next to 0°. The +# [`UnitSphere`](#gpjax.xarray.UnitSphere) transform replaces `lat` and `lon` with +# the three coordinates of a point on the unit sphere. A stationary kernel on these +# coordinates uses the chord distance through the Earth, which has no seam and no +# distortion at the poles. In the same way, [`Cyclic`](#gpjax.xarray.Cyclic) +# encodes a periodic input as a point on a circle. Datetime inputs are days since +# their first timestamp, so a period of 365.25 gives the seasonal cycle: +# +# ```python +# data, spec = from_xarray( +# ds, +# target="t2m", +# inputs=["lat", "lon", "time", "elevation"], +# transforms=[UnitSphere(), Cyclic("time", 365.25), Standardise(["elevation"])], +# ) +# spec.columns +# # ('sphere_x', 'sphere_y', 'sphere_z', 'time_sin', 'time_cos', 'elevation') +# ``` +# +# The transforms run in order, and `spec.columns` names the columns of +# $\mathbf{X}$ that they produce, one for each kernel lengthscale. +# [Infilling Global Surface Temperature](infilling_surface_temperature.py) uses +# `UnitSphere` on a real global field. +# # ## System configuration # %% diff --git a/docs/index.md b/docs/index.md index c75c7e81c..2af488138 100644 --- a/docs/index.md +++ b/docs/index.md @@ -152,6 +152,7 @@ examples/intro_to_kernels examples/regression examples/classification examples/poisson +examples/xarray_workflow examples/natural_gradients ``` @@ -176,11 +177,11 @@ examples/oilmm examples/barycentres examples/graph_kernels examples/heteroscedastic_inference +examples/infilling_surface_temperature examples/multioutput examples/oak examples/oceanmodelling examples/spatial_linear_gp -examples/xarray_workflow examples/yacht ``` diff --git a/docs/reference/xarray.md b/docs/reference/xarray.md index 243686f4f..f0526ba0d 100644 --- a/docs/reference/xarray.md +++ b/docs/reference/xarray.md @@ -14,3 +14,22 @@ predictions back onto the grid. Requires the optional extra: from_xarray GridSpec ``` + +## Input transforms + +Transforms turn the named inputs into the columns of `X`. Pass them to +{func}`~gpjax.xarray.from_xarray`; the {class}`~gpjax.xarray.GridSpec` applies the +same fitted transforms to every new grid. + +```{eval-rst} +.. currentmodule:: gpjax.xarray + +.. autosummary:: + :toctree: generated/ + :nosignatures: + + Standardise + UnitSphere + Cyclic + InputTransform +``` diff --git a/gpjax/xarray.py b/gpjax/xarray.py index 1b89569cf..e9dd80207 100644 --- a/gpjax/xarray.py +++ b/gpjax/xarray.py @@ -30,12 +30,24 @@ rebuild the grid lives on the :class:`GridSpec`, which never enters ``fit``, ``condition`` or a traced computation. +Input transforms (:class:`Standardise`, :class:`UnitSphere`, :class:`Cyclic`) +turn the named inputs into the columns of ``X``. The spec records them, fitted on +the training data, and applies the same ones to every new grid. For large grids, +:meth:`GridSpec.predict` gives the predictive mean and variance in chunks of +bounded size, and keeps a Dask-backed grid lazy. + Requires the optional ``xarray`` extra: ``pip install "gpjax[xarray]"``. """ -from dataclasses import dataclass +import abc +from dataclasses import ( + dataclass, + field, + replace, +) import beartype.typing as tp +import jax import jax.numpy as jnp from jaxtyping import Float import lineax as lx @@ -55,26 +67,194 @@ "gpjax.xarray requires xarray; install it with pip install 'gpjax[xarray]'" ) from error +Columns = dict[str, np.ndarray] + + +class InputTransform(abc.ABC): + r"""A map from named input columns to the named columns of ``X``. + + A transform receives the columns in order, keyed by name, and returns a new + ordered mapping. :func:`from_xarray` calls :meth:`fit` on the training + columns, and the :class:`GridSpec` then applies the fitted transform to the + training data and to every new grid, so both see the same encoding. + Datetime inputs are already float days since their time origin. + """ + + def fit(self, columns: Columns) -> "InputTransform": + r"""Return this transform fitted to the training ``columns``. + + The default returns the transform unchanged; override it when the + transform learns a state from the training data. + """ + return self + + @abc.abstractmethod + def __call__(self, columns: Columns) -> Columns: + r"""Transform ``columns`` and return the new ordered columns.""" + + +@dataclass(frozen=True) +class Standardise(InputTransform): + r"""Centre and scale inputs to zero mean and unit standard deviation. + + The mean and standard deviation come from the training cells. A new grid + is scaled with the training values, so one lengthscale means the same + distance in training and in prediction. + + Attributes: + names: The columns to standardise. ``None`` standardises every column + present when the transform runs. + loc: The fitted training mean of each column; ``None`` before fitting. + scale: The fitted training standard deviation of each column; ``None`` + before fitting. + """ + + names: tp.Optional[tp.Sequence[str]] = None + loc: tp.Optional[dict[str, float]] = field(default=None, compare=False) + scale: tp.Optional[dict[str, float]] = field(default=None, compare=False) + + def __post_init__(self) -> None: + if isinstance(self.names, str): + raise TypeError( + f"names must be a list of names, e.g. [{self.names!r}], not a string" + ) + if self.names is not None: + object.__setattr__(self, "names", tuple(self.names)) + + def fit(self, columns: Columns) -> "Standardise": + r"""Record the mean and standard deviation of each named column. + + Raises: + ValueError: If a name is not a column, or a column is constant. + """ + names = self.names if self.names is not None else tuple(columns) + _require(columns, names, "Standardise") + loc = {name: float(np.mean(columns[name])) for name in names} + scale = {name: float(np.std(columns[name])) for name in names} + constant = [name for name in names if not scale[name] > 0.0] + if constant: + raise ValueError( + f"cannot standardise {constant}: the column is constant over the " + "training cells" + ) + return replace(self, names=names, loc=loc, scale=scale) + + def __call__(self, columns: Columns) -> Columns: + r"""Apply the training mean and standard deviation.""" + if self.loc is None or self.scale is None: + raise RuntimeError("Standardise must be fitted before it is applied") + _require(columns, tuple(self.loc), "Standardise") + return { + name: (values - self.loc[name]) / self.scale[name] + if name in self.loc + else values + for name, values in columns.items() + } + + +@dataclass(frozen=True) +class UnitSphere(InputTransform): + r"""Map latitude and longitude in degrees to a point on the unit sphere. + + The two columns are replaced, at the position of ``lat``, by + ``{name}_x = cos(lat) cos(lon)``, ``{name}_y = cos(lat) sin(lon)`` and + ``{name}_z = sin(lat)``. A stationary kernel on these columns uses the + chord distance through the sphere, which is valid on the whole globe: it + has no seam at the antimeridian, no singularity at the poles, and it does + not stretch distances at high latitudes. A lengthscale is then in Earth + radii (one radius is about 6371 km). + + Attributes: + lat: Name of the latitude column, in degrees north. + lon: Name of the longitude column, in degrees east. + name: Prefix of the three new columns. + """ + + lat: str = "lat" + lon: str = "lon" + name: str = "sphere" + + def __call__(self, columns: Columns) -> Columns: + r"""Replace ``lat`` and ``lon`` with the three unit-sphere columns. + + Raises: + ValueError: If a name is not a column, a new column name is already + taken, or a latitude is outside [-90, 90] degrees. + """ + _require(columns, (self.lat, self.lon), "UnitSphere") + lat_degrees = columns[self.lat] + if np.any(np.abs(lat_degrees[np.isfinite(lat_degrees)]) > 90.0): + raise ValueError( + f"{self.lat!r} has values outside [-90, 90]; UnitSphere expects " + "latitude in degrees north" + ) + lat = np.deg2rad(lat_degrees) + lon = np.deg2rad(columns[self.lon]) + sphere = { + f"{self.name}_x": np.cos(lat) * np.cos(lon), + f"{self.name}_y": np.cos(lat) * np.sin(lon), + f"{self.name}_z": np.sin(lat), + } + return _substitute(columns, (self.lat, self.lon), sphere) + + +@dataclass(frozen=True) +class Cyclic(InputTransform): + r"""Encode a periodic input as a point on a circle. + + The column is replaced by ``{name}_sin`` and ``{name}_cos`` of + ``2 pi value / period``, so values one period apart get the same encoding. + Datetime inputs are days since their time origin: ``Cyclic("time", 365.25)`` + encodes the seasonal cycle, and ``Cyclic("lon", 360.0)`` removes the seam + in longitude. + + Attributes: + name: The column to encode. + period: The period, in the units of the column. + """ + + name: str + period: float + + def __post_init__(self) -> None: + if not self.period > 0.0: + raise ValueError(f"period must be positive, not {self.period}") + + def __call__(self, columns: Columns) -> Columns: + r"""Replace the column with its sine and cosine. + + Raises: + ValueError: If the name is not a column, or a new column name is + already taken. + """ + _require(columns, (self.name,), "Cyclic") + phase = 2.0 * np.pi * columns[self.name] / self.period + circle = {f"{self.name}_sin": np.sin(phase), f"{self.name}_cos": np.cos(phase)} + return _substitute(columns, (self.name,), circle) + @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. + order over ``dims``. That correspondence, and the encoding of the inputs as + the columns of ``X``, is all this object records, and all + :meth:`inputs_for`, :meth:`predict` 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``. + inputs: Input names, in the order they are read from the data. 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. + transforms: The input transforms, fitted on the training cells, that + turn ``inputs`` into the columns of ``X``. """ target: str @@ -85,18 +265,26 @@ class GridSpec: time_origins: dict[str, np.datetime64] mask: np.ndarray n_dropped: int + transforms: tuple[InputTransform, ...] = () @property def n_kept(self) -> int: r"""Number of grid cells with a row in the flattened data.""" return int(self.mask.sum()) + @property + def columns(self) -> tuple[str, ...]: + r"""Names of the columns of ``X``, after the input transforms.""" + empty = {name: np.empty(0) for name in self.inputs} + return tuple(_apply(self.transforms, empty)) + def __repr__(self) -> str: r"""Summarise the grid without printing its coordinates or mask.""" grid = dict(zip(self.dims, self.mask.shape, strict=True)) + columns = "" if self.columns == self.inputs else f", columns={self.columns}" return ( - f"GridSpec(target={self.target!r}, inputs={self.inputs}, grid={grid}, " - f"kept={self.n_kept}, dropped={self.n_dropped})" + f"GridSpec(target={self.target!r}, inputs={self.inputs}{columns}, " + f"grid={grid}, kept={self.n_kept}, dropped={self.n_dropped})" ) def inputs_for( @@ -107,42 +295,108 @@ def inputs_for( The new grid is the broadcast of this spec's inputs as found in ``obj``; no target is needed. Its dims follow the training grid's order, with any new dims after them. Datetime inputs reuse the training time origins, so - a date maps to the same number here as it did in training. Cells where an - input is NaN get no row, and come back as NaN from :meth:`to_xarray`. + a date maps to the same number here as it did in training, and the + fitted input transforms are applied as they were in training. Cells + where an input is NaN get no row, and come back as NaN from + :meth:`to_xarray`. + + This builds every row at once. For a large grid, :meth:`predict` gives + the mean and variance in chunks instead. Args: obj: Labelled data holding every input named by this spec. Returns: The ``(M, D)`` prediction inputs and the ``GridSpec`` of the new grid, - which carries this spec's target name and attributes. + which carries this spec's target name, attributes and transforms. Raises: ValueError: If an input is missing from ``obj``. TypeError: If an input is non-numeric, or is a datetime now but was not in training (or the reverse). """ - dataset = _as_dataset(obj) - variables = _resolve(dataset, self.inputs) - input_dims = dict.fromkeys(dim for var in variables for dim in var.dims) - dims = tuple( - [dim for dim in self.dims if dim in input_dims] - + [dim for dim in input_dims if dim not in self.dims] - ) + dataset, variables, dims = self._grid_of(obj) grid = xr.broadcast(*variables)[0].transpose(*dims) input_matrix = _input_matrix(variables, grid, dims, self.time_origins) keep = np.isfinite(input_matrix).all(axis=1) - test_spec = GridSpec( - target=self.target, + test_spec = replace( + self, target_attrs=dict(self.target_attrs), - inputs=self.inputs, dims=dims, coords=_grid_coords(dataset, dims), - time_origins=self.time_origins, mask=keep.reshape(grid.shape), n_dropped=int(keep.size - keep.sum()), ) - return jnp.asarray(input_matrix[keep]), test_spec + return jnp.asarray(self._encode(input_matrix[keep])), test_spec + + def predict( + self, + predict_fn: tp.Callable[[Float[Array, "M D"]], tp.Any], + obj: tp.Union[xr.Dataset, xr.DataArray], + *, + chunk_size: int = 4096, + ) -> xr.Dataset: + r"""Predict the mean and variance on a new grid, in chunks. + + The grid is built from ``obj`` as in :meth:`inputs_for`, but the cells + go to ``predict_fn`` at most ``chunk_size`` at a time, so the memory use + does not grow with the size of the grid. The last chunk is padded to + ``chunk_size``, so ``predict_fn`` is compiled once, with ``jax.jit``. + + If ``obj`` holds Dask arrays, the result is lazy: each Dask block is + predicted when it is computed, for example by ``.compute()`` or + ``to_netcdf``. Chunk ``obj`` along the grid's dims to control the size + of a block. + + Args: + predict_fn: Maps ``(M, D)`` inputs to a distribution with ``mean`` + and ``variance`` of shape ``(M,)``, for example + ``lambda x: posterior(x, covariance="diagonal")``. Use the + diagonal covariance: a dense one costs ``chunk_size**2`` memory + and gives the same marginals. + obj: Labelled data holding every input named by this spec. + chunk_size: The number of cells given to ``predict_fn`` at once. + + Returns: + An ``xr.Dataset`` holding ``{target}_mean`` and ``{target}_variance`` + over the grid. Cells where an input is NaN are NaN. + + Raises: + ValueError: If ``chunk_size`` is not positive, an input is missing + from ``obj``, or ``predict_fn`` returns the wrong shape. + TypeError: If an input is non-numeric, or is a datetime now but was + not in training (or the reverse). + """ + if chunk_size < 1: + raise ValueError(f"chunk_size must be positive, not {chunk_size}") + dataset, variables, dims = self._grid_of(obj) + grid_inputs = [var.transpose(*dims) for var in xr.broadcast(*variables)] + moments = jax.jit(lambda inputs: _moments(predict_fn(inputs))) + + def predict_block(*blocks: np.ndarray) -> tuple[np.ndarray, np.ndarray]: + columns = [ + _encode(name, block.ravel(), self.time_origins) + for name, block in zip(self.inputs, blocks, strict=True) + ] + input_matrix = np.stack(columns, axis=1) + keep = np.isfinite(input_matrix).all(axis=1) + mean = np.full(keep.size, np.nan) + variance = np.full(keep.size, np.nan) + if keep.any(): + inputs = self._encode(input_matrix[keep]) + mean[keep], variance[keep] = _in_chunks(moments, inputs, chunk_size) + shape = blocks[0].shape + return mean.reshape(shape), variance.reshape(shape) + + mean, variance = xr.apply_ufunc( + predict_block, + *grid_inputs, + output_core_dims=[[], []], + dask="parallelized", + output_dtypes=[float, float], + ) + prediction = self._moments_dataset(mean, variance) + return prediction.assign_coords(_grid_coords(dataset, dims)) def to_xarray( self, @@ -180,14 +434,42 @@ def to_xarray( ) if num_samples is not None: return self._samples(dist, num_samples, key) + return self._moments_dataset( + self._scatter(dist.mean, {}), self._scatter(dist.variance, {}) + ) + + def _grid_of( + self, obj: tp.Union[xr.Dataset, xr.DataArray] + ) -> tuple[xr.Dataset, list[xr.DataArray], tuple[str, ...]]: + r"""The inputs of this spec in ``obj``, and the dims of their grid. + + The dims follow the training grid's order, with any new dims after them. + """ + dataset = _as_dataset(obj) + variables = _resolve(dataset, self.inputs) + input_dims = dict.fromkeys(dim for var in variables for dim in var.dims) + dims = tuple( + [dim for dim in self.dims if dim in input_dims] + + [dim for dim in input_dims if dim not in self.dims] + ) + return dataset, variables, dims + + def _encode(self, input_matrix: np.ndarray) -> np.ndarray: + r"""Apply the fitted transforms to rows of raw inputs.""" + columns = dict(zip(self.inputs, input_matrix.T, strict=True)) + return np.stack(list(_apply(self.transforms, columns).values()), axis=1) + + def _moments_dataset( + self, mean: xr.DataArray, variance: xr.DataArray + ) -> xr.Dataset: + r"""Name ``mean`` and ``variance`` after the target, with its attributes.""" variance_attrs = dict(self.target_attrs) if "units" in variance_attrs: variance_attrs["units"] = f"({variance_attrs['units']})^2" + mean, variance = mean.copy(deep=False), variance.copy(deep=False) + mean.attrs, variance.attrs = dict(self.target_attrs), variance_attrs return xr.Dataset( - { - f"{self.target}_mean": self._scatter(dist.mean, self.target_attrs), - f"{self.target}_variance": self._scatter(dist.variance, variance_attrs), - } + {f"{self.target}_mean": mean, f"{self.target}_variance": variance} ) def _samples( @@ -239,6 +521,7 @@ def from_xarray( target: str, inputs: tp.Sequence[str], *, + transforms: tp.Sequence[InputTransform] = (), dropna: bool = True, ) -> tuple[Dataset, GridSpec]: r"""Flatten labelled xarray data into a :class:`~gpjax.dataset.Dataset`. @@ -246,14 +529,17 @@ def from_xarray( 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. + timestamp. The ``transforms`` then run in order, fitted on the kept cells, + to give the columns of ``X``; :attr:`GridSpec.columns` names them. 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``. + of the columns of ``X`` before any transforms. + transforms: Input transforms, such as :class:`Standardise`, + :class:`UnitSphere` or :class:`Cyclic`, applied in order. dropna: Drop grid cells where the target or any input is NaN. When ``False``, such cells raise instead. @@ -264,8 +550,9 @@ def from_xarray( 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``). + ``DataArray`` is unnamed, no cells survive NaN handling (or any + NaN is present when ``dropna=False``), or a transform rejects its + columns. TypeError: If ``inputs`` is a single string, or the target or an input is not numeric (inputs may also be ``datetime64``). """ @@ -300,7 +587,16 @@ def from_xarray( 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])) + columns = dict(zip(inputs, input_matrix[keep].T, strict=True)) + fitted = [] + for transform in transforms: + fitted.append(transform.fit(columns)) + columns = fitted[-1](columns) + + data = Dataset( + X=jnp.asarray(np.stack(list(columns.values()), axis=1)), + y=jnp.asarray(outputs[keep]), + ) spec = GridSpec( target=target, target_attrs=dict(target_values.attrs), @@ -310,10 +606,78 @@ def from_xarray( time_origins=time_origins, mask=keep.reshape(target_values.shape), n_dropped=n_dropped, + transforms=tuple(fitted), ) return data, spec +def _apply(transforms: tp.Sequence[InputTransform], columns: Columns) -> Columns: + r"""Run fitted ``transforms`` over ``columns`` in order.""" + for transform in transforms: + columns = transform(columns) + return columns + + +def _require(columns: Columns, names: tp.Sequence[str], transform: str) -> None: + r"""Raise if a name that ``transform`` reads is not a column.""" + missing = [name for name in names if name not in columns] + if missing: + raise ValueError( + f"{transform} needs columns {missing}, but the columns are {list(columns)}" + ) + + +def _substitute(columns: Columns, replaced: tuple[str, ...], new: Columns) -> Columns: + r"""Swap the ``replaced`` columns for ``new``, at the first one's position.""" + taken = [name for name in new if name in columns and name not in replaced] + if taken: + raise ValueError(f"cannot add columns {taken}: the names are already taken") + out = {} + for name, values in columns.items(): + if name == replaced[0]: + out.update(new) + elif name not in replaced: + out[name] = values + return out + + +def _moments(dist: tp.Any) -> tuple[Float[Array, " M"], Float[Array, " M"]]: + r"""The mean and marginal variance of a predictive distribution.""" + return dist.mean, dist.variance + + +def _in_chunks( + moments: tp.Callable[[Float[Array, "M D"]], tuple[Array, Array]], + inputs: np.ndarray, + chunk_size: int, +) -> tuple[np.ndarray, np.ndarray]: + r"""Evaluate ``moments`` over ``inputs``, ``chunk_size`` rows at a time. + + The last chunk is padded with copies of its first row, so every call has + the same shape and a jitted ``moments`` compiles once. + """ + n_rows = inputs.shape[0] + means, variances = [], [] + for start in range(0, n_rows, chunk_size): + chunk = inputs[start : start + chunk_size] + n_valid = chunk.shape[0] + if n_valid < chunk_size: + padding = np.repeat(chunk[:1], chunk_size - n_valid, axis=0) + chunk = np.concatenate([chunk, padding]) + mean, variance = moments(jnp.asarray(chunk)) + mean = np.asarray(mean).reshape(-1) + variance = np.asarray(variance).reshape(-1) + if mean.shape != (chunk_size,) or variance.shape != (chunk_size,): + raise ValueError( + f"predict_fn returned a mean of shape {mean.shape} and a variance " + f"of shape {variance.shape} for {chunk_size} inputs; it must " + "return one value per input" + ) + means.append(mean[:n_valid]) + variances.append(variance[:n_valid]) + return np.concatenate(means), np.concatenate(variances) + + 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): @@ -360,29 +724,34 @@ def _input_matrix( dims: tuple[str, ...], time_origins: dict[str, np.datetime64], ) -> np.ndarray: - r"""Stack ``variables`` into an ``(N, D)`` float matrix over ``grid``. + r"""Stack ``variables`` into an ``(N, D)`` float matrix over ``grid``.""" + columns = [ + _encode(variable.name, _column(variable, grid, dims), time_origins) + for variable in variables + ] + return np.stack(columns, axis=1) + + +def _encode( + name: tp.Hashable, column: np.ndarray, time_origins: dict[str, np.datetime64] +) -> np.ndarray: + r"""Encode one flat input column as floats. 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) + 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") + return np.where(np.isnat(column), np.nan, days) + return _as_float(column, f"input {name!r}") def _column( @@ -413,6 +782,10 @@ def _as_float(values: np.ndarray, description: str) -> np.ndarray: __all__ = [ + "Cyclic", "GridSpec", + "InputTransform", + "Standardise", + "UnitSphere", "from_xarray", ] diff --git a/pyproject.toml b/pyproject.toml index 588106caa..aec26f042 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -126,6 +126,8 @@ dev = [ # The `xarray` extra, repeated here because CI syncs without extras and the # `gpjax.xarray` tests must still run. "xarray>=2024.1", + # Lazy, Dask-backed grids in `GridSpec.predict`; tests only. + "dask>=2024.1", # Workflow security audit (`poe lint-actions`, .github/workflows/zizmor.yml). # Pinned through uv.lock rather than `uvx zizmor@latest` so a new release # cannot turn every open PR red without a commit. diff --git a/tests/test_xarray.py b/tests/test_xarray.py index ee94426d0..9cf7b1b81 100644 --- a/tests/test_xarray.py +++ b/tests/test_xarray.py @@ -19,7 +19,12 @@ from gpjax.dataset import Dataset from gpjax.distributions import GaussianDistribution -from gpjax.xarray import from_xarray +from gpjax.xarray import ( + Cyclic, + Standardise, + UnitSphere, + from_xarray, +) import jax import jax.numpy as jnp import jax.random as jr @@ -327,12 +332,20 @@ def gridded(lats, lons) -> xr.Dataset: verbose=False, ) + posterior = model.condition(data) test_inputs, test_spec = spec.inputs_for(fine.drop_vars("t2m")) - out = test_spec.to_xarray(model.condition(data)(test_inputs)) + out = test_spec.to_xarray(posterior(test_inputs)) assert out["t2m_mean"].sizes == {"lat": 11, "lon": 13} np.testing.assert_allclose(out["t2m_mean"], fine["t2m"], atol=0.1) + chunked = spec.predict( + lambda x: posterior(x, covariance="diagonal"), + fine.drop_vars("t2m"), + chunk_size=50, + ) + xr.testing.assert_allclose(chunked, out) + def test_samples_from_a_tagged_diagonal_covariance_are_refused(two_cell_field): _, spec = from_xarray(two_cell_field, target="t2m", inputs=["site"]) @@ -401,3 +414,229 @@ def test_inputs_of_any_numeric_dtype_become_floats(dtype): np.testing.assert_array_equal(data.X[:, 0], [0.0, 1.0]) assert data.X.dtype == jnp.float64 + + +# --- Input transforms ------------------------------------------------------- + + +def test_standardise_gives_columns_with_zero_mean_and_unit_std(field): + data, spec = from_xarray( + field, target="t2m", inputs=["lat", "elevation"], transforms=[Standardise()] + ) + + np.testing.assert_allclose(data.X.mean(axis=0), 0.0, atol=1e-12) + np.testing.assert_allclose(data.X.std(axis=0), 1.0) + (standardise,) = spec.transforms + assert standardise.loc == {"lat": 15.0, "elevation": 350.0} + + +def test_standardise_scales_a_new_grid_with_the_training_statistics(field): + data, spec = from_xarray( + field, target="t2m", inputs=["elevation"], transforms=[Standardise()] + ) + (standardise,) = spec.transforms + grid = xr.Dataset( + {"elevation": ("site", [350.0, 350.0 + standardise.scale["elevation"]])} + ) + + test_inputs, _ = spec.inputs_for(grid) + + np.testing.assert_allclose(test_inputs[:, 0], [0.0, 1.0]) + # elevation(lat, lon) alone spans 6 cells: the first time step of training. + training_inputs, _ = spec.inputs_for(field.drop_vars("t2m")) + np.testing.assert_allclose(training_inputs, data.X[:6]) + + +def test_standardise_changes_only_the_named_columns(field): + data, _ = from_xarray( + field, + target="t2m", + inputs=["lat", "elevation"], + transforms=[Standardise(["elevation"])], + ) + + np.testing.assert_array_equal(np.unique(data.X[:, 0]), [10.0, 20.0]) + np.testing.assert_allclose(data.X[:, 1].std(), 1.0) + + +def test_standardise_rejects_a_constant_input(field): + field["flat"] = (("lat", "lon"), np.ones((2, 3))) + + with pytest.raises(ValueError, match=r"cannot standardise \['flat'\]"): + from_xarray(field, target="t2m", inputs=["flat"], transforms=[Standardise()]) + + +def test_standardise_rejects_a_bare_string(): + with pytest.raises(TypeError, match=r"\['lat'\]"): + Standardise("lat") + + +def test_unit_sphere_replaces_lat_and_lon_with_unit_vectors(): + lons = np.array([0.0, 90.0, 360.0]) + field = xr.Dataset( + {"t2m": (("lat", "lon"), np.zeros((2, 3)))}, + coords={"lat": [0.0, 90.0], "lon": lons, "level": 1.0}, + ) + field["height"] = (("lat", "lon"), np.arange(6.0).reshape(2, 3)) + + data, spec = from_xarray( + field, target="t2m", inputs=["height", "lat", "lon"], transforms=[UnitSphere()] + ) + + assert spec.columns == ("height", "sphere_x", "sphere_y", "sphere_z") + xyz = np.asarray(data.X[:, 1:]) + np.testing.assert_allclose(np.linalg.norm(xyz, axis=1), 1.0) + # On the equator: lon 0 -> +x, lon 90 -> +y, and lon 360 is lon 0 again. + np.testing.assert_allclose(xyz[:3], [[1, 0, 0], [0, 1, 0], [1, 0, 0]], atol=1e-12) + # At the pole every longitude is the same point. + np.testing.assert_allclose(xyz[3:], [[0, 0, 1]] * 3, atol=1e-12) + + +def test_unit_sphere_rejects_latitude_outside_the_degree_range(field): + field = field.assign_coords(lat=[10.0, 100.0]) + + with pytest.raises(ValueError, match=r"outside \[-90, 90\]"): + from_xarray( + field, target="t2m", inputs=["lat", "lon"], transforms=[UnitSphere()] + ) + + +def test_cyclic_maps_values_one_period_apart_to_the_same_point(field): + data, spec = from_xarray( + field, target="t2m", inputs=["lon"], transforms=[Cyclic("lon", period=2.0)] + ) + + assert spec.columns == ("lon_sin", "lon_cos") + # lon is 0, 1, 2 along each row: 0 and 2 are one period apart. + np.testing.assert_allclose(data.X[0], data.X[2], atol=1e-12) + np.testing.assert_allclose(data.X[1], [0.0, -1.0], atol=1e-12) + + +def test_cyclic_on_a_datetime_input_uses_days_since_the_origin(field): + data, _ = from_xarray( + field, target="t2m", inputs=["time"], transforms=[Cyclic("time", period=8.0)] + ) + + # The second timestamp is 2 days, a quarter period, after the first. + np.testing.assert_allclose(data.X[0], [0.0, 1.0], atol=1e-12) + np.testing.assert_allclose(data.X[-1], [1.0, 0.0], atol=1e-12) + + +@pytest.mark.parametrize("period", [0.0, -1.0]) +def test_cyclic_rejects_a_non_positive_period(period): + with pytest.raises(ValueError, match="period must be positive"): + Cyclic("time", period=period) + + +def test_a_transform_reports_a_missing_column(field): + with pytest.raises(ValueError, match=r"UnitSphere needs columns \['lon'\]"): + from_xarray(field, target="t2m", inputs=["lat"], transforms=[UnitSphere()]) + + +def test_a_transform_may_not_overwrite_another_column(field): + field["lon_sin"] = (("lat", "lon"), np.zeros((2, 3))) + + with pytest.raises(ValueError, match=r"\['lon_sin'\].*already taken"): + from_xarray( + field, + target="t2m", + inputs=["lon", "lon_sin"], + transforms=[Cyclic("lon", period=360.0)], + ) + + +def test_transforms_run_in_order_and_the_spec_names_the_columns(field): + data, spec = from_xarray( + field, + target="t2m", + inputs=["lat", "lon", "elevation"], + transforms=[UnitSphere(), Standardise(["elevation"])], + ) + + assert data.X.shape == (12, 4) + assert spec.columns == ("sphere_x", "sphere_y", "sphere_z", "elevation") + assert "columns=('sphere_x', 'sphere_y', 'sphere_z', 'elevation')" in repr(spec) + + +# --- Chunked prediction ------------------------------------------------------ + + +def _toy_predict_fn(inputs): + """A deterministic stand-in for a posterior: one mean and variance per row.""" + return _diagonal_gaussian(inputs.sum(axis=1), inputs[:, 0] ** 2 + 1.0) + + +@pytest.fixture +def gappy_grid(field) -> xr.Dataset: + """The field's inputs, with one NaN covariate cell and lat attrs.""" + grid = field.drop_vars("t2m") + grid["elevation"][1, 2] = np.nan + grid["lat"].attrs = {"units": "degrees_north"} + return grid + + +@pytest.mark.parametrize("chunk_size", [1, 5, 12, 100]) +def test_predict_matches_inputs_for_then_to_xarray(field, gappy_grid, chunk_size): + _, spec = from_xarray( + field, + target="t2m", + inputs=["time", "lat", "elevation"], + transforms=[Standardise()], + ) + test_inputs, test_spec = spec.inputs_for(gappy_grid) + expected = test_spec.to_xarray(_toy_predict_fn(test_inputs)) + + out = spec.predict(_toy_predict_fn, gappy_grid, chunk_size=chunk_size) + + xr.testing.assert_allclose(out, expected) + assert np.isnan(out["t2m_mean"][:, 1, 2]).all() + assert out["t2m_mean"].attrs == {"units": "K"} + assert out["t2m_variance"].attrs == {"units": "(K)^2"} + + +def test_predict_compiles_once_for_every_chunk(field): + _, spec = from_xarray(field, target="t2m", inputs=["lat", "lon"]) + traces = [] + + def predict_fn(inputs): + traces.append(inputs.shape) + return _toy_predict_fn(inputs) + + spec.predict(predict_fn, field.drop_vars("t2m"), chunk_size=5) + + assert traces == [(5, 2)] + + +def test_predict_keeps_a_dask_grid_lazy(field, gappy_grid): + pytest.importorskip("dask") + _, spec = from_xarray(field, target="t2m", inputs=["time", "lat", "elevation"]) + calls = [] + + def predict_fn(inputs): + calls.append(inputs.shape) + return _toy_predict_fn(inputs) + + lazy = spec.predict(predict_fn, gappy_grid.chunk({"lat": 1}), chunk_size=4) + + assert lazy["t2m_mean"].chunks is not None + assert calls == [] + eager = spec.predict(_toy_predict_fn, gappy_grid, chunk_size=4) + xr.testing.assert_allclose(lazy.compute(), eager) + + +@pytest.mark.parametrize("chunk_size", [0, -3]) +def test_predict_rejects_a_non_positive_chunk_size(field, chunk_size): + _, spec = from_xarray(field, target="t2m", inputs=["lat"]) + + with pytest.raises(ValueError, match="chunk_size must be positive"): + spec.predict(_toy_predict_fn, field, chunk_size=chunk_size) + + +def test_predict_rejects_a_predict_fn_of_the_wrong_shape(field): + _, spec = from_xarray(field, target="t2m", inputs=["lat"]) + + def one_value(inputs): + return _diagonal_gaussian(inputs.sum(keepdims=True)[0], jnp.ones(1)) + + with pytest.raises(ValueError, match="one value per input"): + spec.predict(one_value, field, chunk_size=4) diff --git a/uv.lock b/uv.lock index f62a5a867..dd43f1fdd 100644 --- a/uv.lock +++ b/uv.lock @@ -374,6 +374,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/98/78/01c019cdb5d6498122777c1a43056ebb3ebfeef2076d9d026bfe15583b2b/click-8.3.1-py3-none-any.whl", hash = "sha256:981153a64e25f12d547d3426c367a4857371575ee7ad18df2a6183ab0545b2a6", size = 108274, upload-time = "2025-11-15T20:45:41.139Z" }, ] +[[package]] +name = "cloudpickle" +version = "3.1.2" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/27/fb/576f067976d320f5f0114a8d9fa1215425441bb35627b1993e5afd8111e5/cloudpickle-3.1.2.tar.gz", hash = "sha256:7fda9eb655c9c230dab534f1983763de5835249750e85fbcef43aaa30a9a2414", size = 22330, upload-time = "2025-11-03T09:25:26.604Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/88/39/799be3f2f0f38cc727ee3b4f1445fe6d5e4133064ec2e4115069418a5bb6/cloudpickle-3.1.2-py3-none-any.whl", hash = "sha256:9acb47f6afd73f60dc1df93bb801b472f05ff42fa6c84167d25cb206be1fbf4a", size = 22228, upload-time = "2025-11-03T09:25:25.534Z" }, +] + [[package]] name = "codespell" version = "2.4.3" @@ -653,6 +662,25 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/e7/05/c19819d5e3d95294a6f5947fb9b9629efb316b96de511b418c53d245aae6/cycler-0.12.1-py3-none-any.whl", hash = "sha256:85cef7cff222d8644161529808465972e51340599459b8ac3ccbac5a854e0d30", size = 8321, upload-time = "2023-10-07T05:32:16.783Z" }, ] +[[package]] +name = "dask" +version = "2026.8.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "click" }, + { name = "cloudpickle" }, + { name = "fsspec" }, + { name = "importlib-metadata", marker = "python_full_version < '3.12'" }, + { name = "packaging" }, + { name = "partd" }, + { name = "pyyaml" }, + { name = "toolz" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/33/a7/6b3c7ac32b642fbbe0821111654e0bd8cfbe88f68560bcf23cc78ab35c71/dask-2026.8.0.tar.gz", hash = "sha256:8a94c37b5de6d869343340dc26c3c3acca7ec48a3abdabe00ea3abb1125884d5", size = 11561752, upload-time = "2026-08-24T19:21:25.906Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/f8/3a/4fc99e788bcfa1b3b3f21abf57da45898d807d007e7f6fd1c7300904eb70/dask-2026.8.0-py3-none-any.whl", hash = "sha256:ccc0c83a189b0398602435189771d28dad7b5773b6089bb8dce14ae732dd782c", size = 1492182, upload-time = "2026-08-24T19:21:23.997Z" }, +] + [[package]] name = "debugpy" version = "1.8.19" @@ -797,6 +825,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/c7/4e/ce75a57ff3aebf6fc1f4e9d508b8e5810618a33d900ad6c19eb30b290b97/fonttools-4.61.1-py3-none-any.whl", hash = "sha256:17d2bf5d541add43822bcf0c43d7d847b160c9bb01d15d5007d84e2217aaa371", size = 1148996, upload-time = "2025-12-12T17:31:21.03Z" }, ] +[[package]] +name = "fsspec" +version = "2026.9.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/77/cd/9be253869fc42e764de7f3dedd6969af7d44ff9c3375214a3442a6f3fc08/fsspec-2026.9.0.tar.gz", hash = "sha256:0f08147951c8cb31d844c3547d631053b127863b60be04cf06e121333ee0e2fe", size = 333545, upload-time = "2026-09-18T17:50:42.825Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/6c/c0/a98505f18594f1bce828bb159cec0fcf9860562f1a2c85913409fc8f3d9e/fsspec-2026.9.0-py3-none-any.whl", hash = "sha256:8dd6e646e99ea382bd85f97a45e6b526a442d79423a7dc673f1e2756d05fcb5f", size = 221738, upload-time = "2026-09-18T17:50:41.341Z" }, +] + [[package]] name = "gpjax" source = { editable = "." } @@ -858,6 +895,7 @@ dev = [ { name = "asv" }, { name = "codespell" }, { name = "coverage" }, + { name = "dask" }, { name = "hypothesis" }, { name = "interrogate" }, { name = "jupytext" }, @@ -929,6 +967,7 @@ dev = [ { name = "asv", specifier = ">=0.6" }, { name = "codespell", specifier = ">=2.2.4" }, { name = "coverage", specifier = ">=7.2.2" }, + { name = "dask", specifier = ">=2024.1" }, { name = "hypothesis", specifier = ">=6.148.2" }, { name = "interrogate", specifier = ">=1.5.0" }, { name = "jupytext" }, @@ -1734,6 +1773,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/8d/f8/0d32243c6b8e5dbee9097cf0c95bbdf8681ba4463c927c2e3445f3775814/lineax-0.1.1-py3-none-any.whl", hash = "sha256:2e399f1674773ab2ba54d76175a618977a554f47abd0a345198d53d92c07beb2", size = 77567, upload-time = "2026-05-01T15:59:05.517Z" }, ] +[[package]] +name = "locket" +version = "1.0.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/2f/83/97b29fe05cb6ae28d2dbd30b81e2e402a3eed5f460c26e9eaa5895ceacf5/locket-1.0.0.tar.gz", hash = "sha256:5c0d4c052a8bbbf750e056a8e65ccd309086f4f0f18a2eac306a8dfa4112a632", size = 4350, upload-time = "2022-04-20T22:04:44.312Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/db/bc/83e112abc66cd466c6b83f99118035867cecd41802f8d044638aa78a106e/locket-1.0.0-py2.py3-none-any.whl", hash = "sha256:b6c819a722f7b6bd955b80781788e4a66a55628b858d347536b7e81325a3a5e3", size = 4398, upload-time = "2022-04-20T22:04:42.23Z" }, +] + [[package]] name = "markdown-it-py" version = "4.0.0" @@ -2357,6 +2405,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/16/32/f8e3c85d1d5250232a5d3477a2a28cc291968ff175caeadaf3cc19ce0e4a/parso-0.8.5-py2.py3-none-any.whl", hash = "sha256:646204b5ee239c396d040b90f9e272e9a8017c630092bf59980beb62fd033887", size = 106668, upload-time = "2025-08-23T15:15:25.663Z" }, ] +[[package]] +name = "partd" +version = "1.4.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "locket" }, + { name = "toolz" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/b2/3a/3f06f34820a31257ddcabdfafc2672c5816be79c7e353b02c1f318daa7d4/partd-1.4.2.tar.gz", hash = "sha256:d022c33afbdc8405c226621b015e8067888173d85f7f5ecebb3cafed9a20f02c", size = 21029, upload-time = "2024-05-06T19:51:41.945Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/71/e7/40fb618334dcdf7c5a316c0e7343c5cd82d3d866edc100d98e29bc945ecd/partd-1.4.2-py3-none-any.whl", hash = "sha256:978e4ac767ec4ba5b86c6eaa52e5a2a3bc748a2ca839e8cc798f1cc6ce6efb0f", size = 18905, upload-time = "2024-05-06T19:51:39.271Z" }, +] + [[package]] name = "pastel" version = "0.2.1" @@ -3774,6 +3835,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/23/d1/136eb2cb77520a31e1f64cbae9d33ec6df0d78bdf4160398e86eec8a8754/tomli-2.4.0-py3-none-any.whl", hash = "sha256:1f776e7d669ebceb01dee46484485f43a4048746235e683bcdffacdf1fb4785a", size = 14477, upload-time = "2026-01-11T11:22:37.446Z" }, ] +[[package]] +name = "toolz" +version = "1.1.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/11/d6/114b492226588d6ff54579d95847662fc69196bdeec318eb45393b24c192/toolz-1.1.0.tar.gz", hash = "sha256:27a5c770d068c110d9ed9323f24f1543e83b2f300a687b7891c1a6d56b697b5b", size = 52613, upload-time = "2025-10-17T04:03:21.661Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/fb/12/5911ae3eeec47800503a238d971e51722ccea5feb8569b735184d5fcdbc0/toolz-1.1.0-py3-none-any.whl", hash = "sha256:15ccc861ac51c53696de0a5d6d4607f99c210739caf987b5d2054f3efed429d8", size = 58093, upload-time = "2025-10-17T04:03:20.435Z" }, +] + [[package]] name = "tornado" version = "6.5.4"