diff --git a/CHANGELOG.md b/CHANGELOG.md index 911fd3019..b66e08f6c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,8 +8,36 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- **Nonstationary and space–time kernels + ([#812](https://github.com/QuantClimate/GPJax/issues/812)).** + - `VaryingAmplitude(base_kernel, amplitude)` multiplies a base kernel by + $\sigma(x)\sigma(y)$, so that variability changes with location. + - `Gibbs(base_kernel, lengthscale)` lets the lengthscale of an isotropic base + kernel (RBF, the Matérn kernels, RationalQuadratic or PoweredExponential) + change with location. It is the Paciorek–Schervish construction, and it + reuses the existing kernels and their ARD lengthscales. + - `Gneiting(space_dims, time_dim)` is the nonseparable space–time kernel of + Gneiting (2002), with an interaction parameter $\beta$ (0 is separable). + - The new `gpjax.kernels.location_functions` module gives the functions of + location that the first two kernels use: `Constant`, `Linear` (log-linear in + selected covariate columns, starting at zero) and an abstract base class for + custom functions. A location function is not a mean function; see + `docs/adr/0001-location-functions.md`. + - Stationary kernels have a new class flag, `isotropic_radial`, which marks + the kernels that are valid Gibbs base kernels. + - A new example, *Nonstationary Kernels over Complex Terrain*, fits these + kernels to Colorado precipitation normals with elevation as the covariate. + - A new example, *Space–Time Modelling of Winter Temperature*, fits the + Gneiting kernel to daily NCEP-NCAR reanalysis temperature anomalies over + Europe and compares it with its separable version. + ### Changed +- `RFF`, and so pathwise sampling with `sample_approx`, now names the kernel it + cannot approximate and tells you to sample from the predictive distribution. + - Allow JAX and JAXlib 0.11 in downstream environments by removing the `<0.11` dependency bounds ([#801](https://github.com/QuantClimate/GPJax/issues/801)). diff --git a/GLOSSARY.md b/GLOSSARY.md new file mode 100644 index 000000000..faee846e3 --- /dev/null +++ b/GLOSSARY.md @@ -0,0 +1,43 @@ +# GPJax + +GPJax is a Gaussian process library built on JAX. This glossary gives the +canonical terms for concepts that are specific to GPJax. + +## Nonstationary kernels + +**Location function**: +A kernel parameter that has a value at each input location, for example a +standard deviation or a lengthscale that changes over space. It is not a mean +function: a mean function describes the process, but a location function +describes the covariance. +_Avoid_: mean function, warping function, parameter field + +**Base kernel**: +The stationary kernel that a nonstationary kernel modifies with a location +function. +_Avoid_: inner kernel, wrapped kernel + +**Varying amplitude kernel**: +A base kernel whose standard deviation changes with location, so that +variability is larger in some regions than in others. +_Avoid_: heteroscedastic kernel (GPJax uses "heteroscedastic" for noise), +scaled kernel + +**Gibbs kernel**: +A base kernel whose lengthscale changes with location, so that correlation +decays faster in some regions than in others. The literature also calls it the +Paciorek–Schervish kernel. +_Avoid_: nonstationary Matérn, Paciorek–Schervish kernel (as a name) + +## Space–time kernels + +**Gneiting kernel**: +A nonseparable space–time kernel. The spatial correlation changes with the +time lag, so that it decays more slowly at longer lags. +_Avoid_: space–time Matérn, separable space–time kernel + +**Space columns**: +The input columns over which a space–time kernel measures spatial distance. + +**Time columns**: +The input columns over which a space–time kernel measures the time lag. diff --git a/docs/adr/0001-location-functions.md b/docs/adr/0001-location-functions.md new file mode 100644 index 000000000..a3749b437 --- /dev/null +++ b/docs/adr/0001-location-functions.md @@ -0,0 +1,33 @@ +# Location functions are a separate type from mean functions + +Nonstationary kernels (#812) need a kernel parameter that changes with input +location, such as a standard deviation or a lengthscale. We decided to add a +separate abstract type for these location functions in +`gpjax/kernels/location_functions.py`, and not to reuse +`gpjax.mean_functions`. A mean function describes the process, while a +location function describes the covariance. They also have different +contracts: a location function evaluates one point, selects its own input +columns with `active_dims`, and returns a value on the log scale, to which the +kernel applies `exp`. + +## Considered options + +- **Reuse `AbstractMeanFunction`.** Rejected. It works on `(N, D)` batches, + has no column selection, and would mix two different concepts in one type. +- **Accept any `eqx.Module` callable.** Rejected. It gives no contract that + beartype can check, and a plain Python callable silently puts its parameters + in a static field, so `fit` does not train them. + +## Consequences + +- The user documentation must explain why a location function is not a mean + function, because a reader will expect the two to be the same. +- A location function that has an intercept overlaps with the scale parameters + of the base kernel. For this reason, `location_functions.Linear` has no + intercept by default. +- Kernels evaluate their location functions inside `__call__`, and no special + compute engine is used. This is not wasteful: `vmap` batches only the + operations that depend on the mapped input, so the dense engine evaluates + a location function once for each row ($N + M$ times for an $N \times M$ + matrix), also inside sum and product kernels. A dedicated engine would add + code and give no saving. diff --git a/docs/examples/data/_pull_reference_datasets.py b/docs/examples/data/_pull_reference_datasets.py index 6811045b3..3bc114fc6 100644 --- a/docs/examples/data/_pull_reference_datasets.py +++ b/docs/examples/data/_pull_reference_datasets.py @@ -32,6 +32,25 @@ 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``. +- Colorado precipitation normals (``colorado_precipitation_normals.nc``): the + 1991-2020 annual precipitation normal at each Colorado COOP and first-order + station, with the station location and elevation. Source: NOAA NCEI U.S. + Climate Normals, one CSV per station at + https://www.ncei.noaa.gov/data/normals-annualseasonal/1991-2020/access/ , + with Colorado stations found from + https://www.ncei.noaa.gov/pub/data/ghcn/daily/ghcnd-stations.txt . + Work of the US Government, so public domain. Cite Palecki et al. (2021), + U.S. Climate Normals 1991-2020, NOAA NCEI. +- European winter temperature anomalies (``ncep_europe_winter_2019.nc``): daily + near-surface (sigma 0.995) air temperature anomalies over Europe and the + north-east Atlantic, 1 January to 1 March 2019, on a 7.5 x 10 degree subset of + the 2.5 degree grid. The anomaly is the daily value minus the 1991-2020 daily + long-term mean. Source: NCEP-NCAR Reanalysis 1 (Kalnay et al., 1996), + https://psl.noaa.gov/thredds/fileServer/Datasets/ncep.reanalysis/Dailies/surface/air.sig995.2019.nc + and .../ncep.reanalysis.derived/surface/air.sig995.day.ltm.1991-2020.nc . + Public domain; "NCEP-NCAR Reanalysis 1 data provided by the NOAA PSL, + Boulder, Colorado, USA". The source files are netCDF4, so this pull needs + ``uv run --extra docs --with h5netcdf --with h5py python ...``. """ from __future__ import annotations @@ -104,8 +123,200 @@ def pull_auto_mpg() -> None: _save(pd.concat([features, targets], axis=1), "auto_mpg.csv") +# --------------------------------------------------------------------------- # +# nonstationary_terrain — Colorado annual precipitation normals, 1991-2020. # +# --------------------------------------------------------------------------- # +NORMALS_URL = "https://www.ncei.noaa.gov/data/normals-annualseasonal/1991-2020/access/" +STATIONS_URL = "https://www.ncei.noaa.gov/pub/data/ghcn/daily/ghcnd-stations.txt" + + +def _get(url: str) -> requests.Response: + last_err = None + for attempt in range(4): + try: + resp = requests.get(url, timeout=60) + if resp.status_code == 404: + return resp + resp.raise_for_status() + return resp + except requests.RequestException as err: + last_err = err + time.sleep(2 * (attempt + 1)) + raise RuntimeError(f"Failed to fetch {url}: {last_err}") + + +def pull_colorado_precipitation() -> None: + print("nonstationary_terrain: Colorado precipitation normals 1991-2020") + import io + import re + + import xarray as xr + + stations = _get(STATIONS_URL).text.splitlines() + colorado = sorted( + line[:11] + for line in stations + if line[38:40] == "CO" and line[:3] in ("USC", "USW") + ) + listing = _get(NORMALS_URL).text + available = set(re.findall(r'href="(US[CW]\d+)\.csv"', listing)) + rows = [] + for station in (s for s in colorado if s in available): + frame = pd.read_csv(io.StringIO(_get(f"{NORMALS_URL}{station}.csv").text)) + if "ANN-PRCP-NORMAL" not in frame: + continue + precipitation = pd.to_numeric(frame["ANN-PRCP-NORMAL"], errors="coerce") + if not precipitation.iloc[0] > 0: + continue + rows.append( + { + "station": station, + "name": str(frame["NAME"].iloc[0]).strip(), + "lat": float(frame["LATITUDE"].iloc[0]), + "lon": float(frame["LONGITUDE"].iloc[0]), + "elevation": float(frame["ELEVATION"].iloc[0]), + # The normals are in inches. + "precipitation": 25.4 * float(precipitation.iloc[0]), + } + ) + table = pd.DataFrame(rows) + ds = xr.Dataset( + { + "precipitation": ( + "station", + table["precipitation"].to_numpy(), + { + "long_name": "Annual precipitation normal, 1991-2020", + "units": "mm", + }, + ), + "elevation": ( + "station", + table["elevation"].to_numpy(), + {"long_name": "Station elevation", "units": "m"}, + ), + }, + coords={ + "station": ("station", table["station"].to_numpy()), + "name": ("station", table["name"].to_numpy()), + "lat": ("station", table["lat"].to_numpy(), {"units": "degrees_north"}), + "lon": ("station", table["lon"].to_numpy(), {"units": "degrees_east"}), + }, + attrs={ + "title": "Colorado annual precipitation normals, 1991-2020", + "source": "NOAA NCEI U.S. Climate Normals 1991-2020 (annual/seasonal)", + "references": "Palecki et al. (2021), U.S. Climate Normals 1991-2020, " + "NOAA National Centers for Environmental Information.", + "license": "Public domain (work of the US Government).", + "Conventions": "CF-1.8", + }, + ) + path = HERE / "colorado_precipitation_normals.nc" + ds.to_netcdf(path, engine="scipy", format="NETCDF3_64BIT") + print(f" wrote {path.name}: {ds.sizes['station']} stations") + + +# --------------------------------------------------------------------------- # +# spacetime_temperature — daily NCEP-NCAR R1 temperature anomalies, Europe. # +# --------------------------------------------------------------------------- # +PSL_THREDDS = "https://psl.noaa.gov/thredds/fileServer/Datasets/" + + +def _download_resumable(url: str, path: Path) -> None: + """Download a large file. The PSL server drops long transfers, so resume.""" + for _ in range(20): + done = path.stat().st_size if path.exists() else 0 + headers = {"Range": f"bytes={done}-"} if done else {} + try: + with requests.get(url, headers=headers, stream=True, timeout=120) as resp: + if resp.status_code == 416: + return + resp.raise_for_status() + if done and resp.status_code != 206: + done = 0 + total = done + int(resp.headers.get("Content-Length", 0)) + with path.open("ab" if done else "wb") as handle: + for chunk in resp.iter_content(1 << 20): + handle.write(chunk) + if path.stat().st_size >= total: + return + except requests.RequestException: + time.sleep(3) + raise RuntimeError(f"Failed to fetch {url}") + + +def pull_europe_winter_temperature() -> None: + print("spacetime_temperature: NCEP-NCAR R1 daily temperature anomalies") + import tempfile + + import numpy as np + import xarray as xr + + with tempfile.TemporaryDirectory() as tmp: + daily_path = Path(tmp) / "air.sig995.2019.nc" + normal_path = Path(tmp) / "air.sig995.day.ltm.1991-2020.nc" + _download_resumable( + PSL_THREDDS + "ncep.reanalysis/Dailies/surface/air.sig995.2019.nc", + daily_path, + ) + _download_resumable( + PSL_THREDDS + + "ncep.reanalysis.derived/surface/air.sig995.day.ltm.1991-2020.nc", + normal_path, + ) + daily = xr.open_dataset(daily_path, engine="h5netcdf")["air"] + normal = xr.open_dataset(normal_path, engine="h5netcdf", decode_times=False)[ + "air" + ] + daily = daily.sel(time=slice("2019-01-01", "2019-03-01")) + day_of_year = np.minimum(daily.time.dt.dayofyear.to_numpy(), 365) - 1 + anomaly = ( + daily - normal.isel(time=xr.DataArray(day_of_year, dims="time")).values + ) + # Longitudes 340E-40E, as -20 to 40 degrees east. + anomaly = anomaly.assign_coords(lon=((anomaly.lon + 180) % 360) - 180) + anomaly = anomaly.sortby("lon").sel(lat=slice(70, 40), lon=slice(-20, 40)) + anomaly = anomaly.isel(lat=slice(None, None, 3), lon=slice(None, None, 4)) + anomaly = anomaly.load() + + ds = xr.Dataset( + { + "tas_anomaly": ( + ("time", "lat", "lon"), + anomaly.to_numpy().astype("float32"), + { + "long_name": "Daily near-surface air temperature anomaly " + "(sigma 0.995) against the 1991-2020 daily mean", + "units": "K", + }, + ) + }, + coords={ + "time": anomaly.time.to_numpy(), + "lat": ("lat", anomaly.lat.to_numpy(), {"units": "degrees_north"}), + "lon": ("lon", anomaly.lon.to_numpy(), {"units": "degrees_east"}), + }, + attrs={ + "title": "Daily temperature anomalies over Europe, winter 2019", + "source": "NCEP-NCAR Reanalysis 1, NOAA PSL", + "references": "Kalnay et al. (1996), The NCEP/NCAR 40-year " + "reanalysis project, Bull. Amer. Meteor. Soc. 77, 437-471.", + "acknowledgment": "NCEP-NCAR Reanalysis 1 data provided by the NOAA " + "PSL, Boulder, Colorado, USA, from their website at " + "https://psl.noaa.gov", + "license": "Public domain.", + "Conventions": "CF-1.8", + }, + ) + path = HERE / "ncep_europe_winter_2019.nc" + ds.to_netcdf(path, engine="scipy", format="NETCDF3_64BIT") + print(f" wrote {path.name}: {dict(ds.sizes)}") + + if __name__ == "__main__": pull_mauna_loa_co2() pull_gulf_velocities() pull_auto_mpg() + pull_colorado_precipitation() + pull_europe_winter_temperature() print("\nDone.") diff --git a/docs/examples/data/colorado_precipitation_normals.nc b/docs/examples/data/colorado_precipitation_normals.nc new file mode 100644 index 000000000..74060ef8b Binary files /dev/null and b/docs/examples/data/colorado_precipitation_normals.nc differ diff --git a/docs/examples/data/ncep_europe_winter_2019.nc b/docs/examples/data/ncep_europe_winter_2019.nc new file mode 100644 index 000000000..f02f9eeee Binary files /dev/null and b/docs/examples/data/ncep_europe_winter_2019.nc differ diff --git a/docs/examples/nonstationary_terrain.py b/docs/examples/nonstationary_terrain.py new file mode 100644 index 000000000..7b719472e --- /dev/null +++ b/docs/examples/nonstationary_terrain.py @@ -0,0 +1,359 @@ +# --- +# 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] +# # Nonstationary Kernels over Complex Terrain +# +# Download this notebook: {nb-download}`nonstationary_terrain.ipynb` +# +# A stationary kernel uses one lengthscale and one variance everywhere. Over +# complex terrain this is a poor assumption. In Colorado, annual precipitation +# changes over a few kilometres in the Rocky Mountains, but only slowly over +# the Great Plains to the east. A stationary kernel must use one compromise +# lengthscale for both regions. Its intervals are then too narrow in the +# mountains and too wide on the plains. +# +# Paciorek & Schervish (2006) used this example to motivate nonstationary +# kernels. In this notebook we +# +# 1. load the 1991–2020 annual precipitation normals at 247 Colorado stations, +# 2. let the variance and the lengthscale change with elevation, with the +# [`VaryingAmplitude`](#gpjax.kernels.VaryingAmplitude) and +# [`Gibbs`](#gpjax.kernels.Gibbs) kernels, +# 3. compare these models with a stationary kernel on the marginal likelihood +# and on held-out stations, separately for the mountains and the plains, and +# 4. map the fitted lengthscale and amplitude. + +# %% 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 +from jaxtyping import install_import_hook +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +import xarray as xr + +config.update("jax_enable_x64", True) + +with install_import_hook("gpjax", "beartype.beartype"): + import gpjax as gpx + from gpjax.kernels import location_functions + from gpjax.parameters import val + +gpx.plotting.use_style() + +# %% [markdown] +# ## The data +# +# The NOAA U.S. Climate Normals give the mean annual precipitation for +# 1991–2020 at each station, with its location and elevation. The file in this +# repository is a small netCDF subset for Colorado; see +# `docs/examples/data/_pull_reference_datasets.py` for how it was made. +# +# Precipitation is positive and skewed, so we model its logarithm, as Paciorek +# & Schervish did. + +# %% +DATA = Path("data") if Path("data").exists() else Path("docs/examples/data") +stations = xr.open_dataset(DATA / "colorado_precipitation_normals.nc", engine="scipy") +stations + +# %% +fig, (precip_ax, elev_ax) = plt.subplots(1, 2, figsize=(10, 3.6), sharey=True) +for ax, values, label, cmap in [ + (precip_ax, stations.precipitation, "Annual precipitation (mm)", "viridis"), + (elev_ax, stations.elevation, "Elevation (m)", "cividis"), +]: + points = ax.scatter(stations.lon, stations.lat, c=values, s=14, cmap=cmap) + fig.colorbar(points, ax=ax, label=label) + ax.set_xlabel("Longitude") + ax.set_aspect(1 / np.cos(np.deg2rad(39.0))) +precip_ax.set_ylabel("Latitude") + +# %% [markdown] +# The wettest stations are high in the mountains in the west. The plains in the +# east are dry, and they change slowly. Elevation is the obvious covariate for +# the structure of the field. +# +# We project the station locations to kilometres on a local plane, and we use +# units of 100 km so that the lengthscales start at a sensible value. We also +# standardise elevation, so that one unit of its weight is one standard +# deviation of elevation. The inputs have three columns: east, north and +# standardised elevation. + +# %% +lat = stations.lat.to_numpy() +lon = stations.lon.to_numpy() +elevation = stations.elevation.to_numpy() + +east = (lon - lon.mean()) * 111.32 * np.cos(np.deg2rad(lat.mean())) / 100.0 +north = (lat - lat.mean()) * 110.57 / 100.0 +elevation_scale = elevation.std() +standardised_elevation = (elevation - elevation.mean()) / elevation_scale + +log_precipitation = np.log(stations.precipitation.to_numpy()) +y_mean, y_scale = log_precipitation.mean(), log_precipitation.std() + +X = np.column_stack([east, north, standardised_elevation]) +y = ((log_precipitation - y_mean) / y_scale)[:, None] +data = gpx.Dataset(X=jnp.asarray(X), y=jnp.asarray(y)) + +mountains = elevation > 2000.0 +print(f"{data.n} stations: {mountains.sum()} above 2000 m, {(~mountains).sum()} below") + +# %% [markdown] +# ## Location functions +# +# A nonstationary kernel needs a parameter that has a value at each location. +# GPJax calls this a *location function*. The module +# [`gpjax.kernels.location_functions`](../reference/kernels.md) gives a +# `Constant` and a log-linear `Linear` function, and you can subclass +# `AbstractLocationFunction` for anything else. +# +# A location function is not a mean function. A mean function describes the +# Gaussian process itself, but a location function describes its covariance. +# The two also have different contracts: +# +# - A location function evaluates one input point, and it selects its own +# columns with `active_dims`. Here the base kernel measures distance over the +# east and north columns, while the location function reads only the +# elevation column. +# - A location function returns a value on the log scale, and the kernel applies +# `exp`. So the parameter is always positive, and a weight $w$ multiplies the +# parameter by $e^{w}$ for each unit of its column. +# +# `Linear` starts with zero weights and no intercept. A model with these +# functions therefore starts as exactly its stationary base kernel, and the base +# kernel keeps the overall scale. +# +# We compare four models. All of them use a Matérn-3/2 base kernel over east and +# north, a constant mean and Gaussian noise. +# +# - **Stationary**: the base kernel alone. +# - **Varying amplitude**: +# $k(x, y) = \sigma(x)\,\sigma(y)\,k_0(x, y)$ with +# $\log\sigma(x) = w_\sigma\,z(x)$, where $z$ is standardised elevation. +# - **Gibbs**: the lengthscale is $\ell_0\,\ell(x)$ with +# $\log\ell(x) = w_\ell\,z(x)$. The marginal variance stays the same +# everywhere. +# - **Both**: a Gibbs kernel inside a varying-amplitude kernel. + + +# %% +def build_model(kind: str): + base = gpx.kernels.Matern32(active_dims=[0, 1]) + + def elevation_fn(): + return location_functions.Linear(active_dims=[2]) + + if kind == "Stationary": + kernel = base + elif kind == "Varying amplitude": + kernel = gpx.kernels.VaryingAmplitude(base, amplitude=elevation_fn()) + elif kind == "Gibbs": + kernel = gpx.kernels.Gibbs(base, lengthscale=elevation_fn()) + elif kind == "Both": + kernel = gpx.kernels.VaryingAmplitude( + gpx.kernels.Gibbs(base, lengthscale=elevation_fn()), + amplitude=elevation_fn(), + ) + prior = gpx.gps.Prior(mean_function=gpx.mean_functions.Constant(), kernel=kernel) + return prior * gpx.likelihoods.Gaussian(obs_stddev=0.3) + + +def negative_mll(model, data): + return -gpx.objectives.conjugate_mll(model, data) + + +def fit(model, data): + return gpx.fit_scipy( + model=model, objective=negative_mll, train_data=data, verbose=False + ) + + +KINDS = ["Stationary", "Varying amplitude", "Gibbs", "Both"] +fitted = {} +log_marginal_likelihood = {} +for kind in KINDS: + fitted[kind], history = fit(build_model(kind), data) + log_marginal_likelihood[kind] = -float(history[-1]) + +# %% [markdown] +# ## What the fitted weights mean +# +# Because of the log link, each weight converts directly into a factor per +# 1000 m of elevation. + + +# %% +def weights_of(kernel): + """The elevation weights of the amplitude and lengthscale functions.""" + weights = {} + if isinstance(kernel, gpx.kernels.VaryingAmplitude): + weights["amplitude"] = float(val(kernel.amplitude.weights)[0]) + kernel = kernel.base_kernel + if isinstance(kernel, gpx.kernels.Gibbs): + weights["lengthscale"] = float(val(kernel.lengthscale.weights)[0]) + return weights + + +per_km = 1000.0 / elevation_scale +for kind in KINDS[1:]: + for name, weight in weights_of(fitted[kind].prior.kernel).items(): + print( + f"{kind:18s} {name:12s} weight {weight:+.2f}: " + f"x{np.exp(weight * per_km):.2f} per 1000 m" + ) + +# %% [markdown] +# The Gibbs kernel learns a lengthscale that becomes much shorter at higher +# elevation, as we expected: precipitation changes over short distances in the +# mountains and over long distances on the plains. The amplitude alone learns a +# larger variance in the mountains. When both are present, most of the effect +# goes to the lengthscale. +# +# ## Marginal likelihood and held-out stations +# +# The marginal likelihood is the first test. To test the predictions, we also +# use 5-fold cross-validation: we refit each model with one fifth of the +# stations held out, and predict those stations. We report the negative log +# predictive density (NLPD, lower is better) and the coverage of the 90% +# predictive intervals, separately for stations above and below 2000 m. A good +# model has coverage close to 0.90 in both regions. + +# %% +folds = np.random.default_rng(0).permutation(data.n) % 5 +z90 = 1.6449 + + +def cross_validate(kind: str) -> dict: + nlpd = np.zeros(data.n) + covered = np.zeros(data.n, dtype=bool) + error = np.zeros(data.n) + for fold in range(5): + train, test = folds != fold, folds == fold + train_data = gpx.Dataset(X=data.X[train], y=data.y[train]) + model, _ = fit(build_model(kind), train_data) + latent = model.predict(data.X[test], train_data, covariance="diagonal") + mean = np.asarray(latent.mean) + variance = ( + np.asarray(latent.variance) + float(val(model.likelihood.obs_stddev)) ** 2 + ) + residual = y[test, 0] - mean + nlpd[test] = 0.5 * np.log(2 * np.pi * variance) + 0.5 * residual**2 / variance + covered[test] = np.abs(residual) < z90 * np.sqrt(variance) + error[test] = residual * y_scale + return { + "Log marginal likelihood": log_marginal_likelihood[kind], + "NLPD": nlpd.mean(), + "NLPD, mountains": nlpd[mountains].mean(), + "NLPD, plains": nlpd[~mountains].mean(), + "90% coverage, mountains": covered[mountains].mean(), + "90% coverage, plains": covered[~mountains].mean(), + "RMSE (log mm)": np.sqrt(np.mean(error**2)), + } + + +scores = pd.DataFrame({kind: cross_validate(kind) for kind in KINDS}).T +scores.round(3) + +# %% [markdown] +# The results agree with Paciorek & Schervish (2006): +# +# - **The marginal likelihood improves strongly.** The Gibbs kernel is much +# better than the stationary kernel, with only one more parameter. +# - **The uncertainty becomes calibrated by region.** The stationary kernel +# uses one compromise lengthscale. Its intervals are too narrow in the +# mountains and much too wide on the plains. The Gibbs kernel moves the +# mountain coverage towards 0.90. On the plains, the coverage changes from +# much too wide to slightly too narrow. The held-out NLPD falls in both +# regions, which shows that the predictive distributions are better overall. +# - **The error falls less.** The RMSE improves, but by less than the NLPD. +# Paciorek & Schervish also found that the main gain is in the likelihood +# and the uncertainty, not in the point predictions. +# +# The varying amplitude alone helps less than the Gibbs kernel. Here the main +# nonstationarity is in the correlation length, not in the variance. +# +# ## Maps of the fitted lengthscale and amplitude +# +# The lengthscale of the Gibbs kernel at each station is $\ell_0 e^{w_\ell +# z(x)}$. We plot it in kilometres, together with the amplitude +# $\sigma(x) = e^{w_\sigma z(x)}$ from the model with both functions. + +# %% +both = fitted["Both"].prior.kernel +gibbs = both.base_kernel +lengthscale_km = ( + 100.0 + * float(val(gibbs.base_kernel.lengthscale)) + * np.exp(np.asarray([gibbs.lengthscale(x) for x in data.X])) +) +amplitude = np.exp(np.asarray([both.amplitude(x) for x in data.X])) + +fig, (ls_ax, amp_ax) = plt.subplots(1, 2, figsize=(10, 3.6), sharey=True) +for ax, values, label in [ + (ls_ax, lengthscale_km, "Lengthscale (km)"), + (amp_ax, amplitude, "Amplitude multiplier"), +]: + points = ax.scatter(stations.lon, stations.lat, c=values, s=14, cmap="magma") + fig.colorbar(points, ax=ax, label=label) + ax.set_xlabel("Longitude") + ax.set_aspect(1 / np.cos(np.deg2rad(39.0))) +ls_ax.set_ylabel("Latitude") + +# %% [markdown] +# The lengthscale is a few tens of kilometres along the high ranges and a few +# hundred kilometres on the eastern plains. +# +# ## Notes +# +# - **Valid base kernels.** The Gibbs construction is positive definite only +# for an isotropic radial base kernel: RBF, the Matérn kernels, +# RationalQuadratic or PoweredExponential. `Gibbs` raises a `TypeError` for +# other kernels. `VaryingAmplitude` accepts any base kernel that evaluates +# one pair of points at a time. +# - **Extrapolation.** Because of the log link, the lengthscale and the +# amplitude change exponentially with the covariate. Outside the range of +# elevations in the data, they can become very large or very small. Keep the +# predictions inside the range of the covariate. +# - **Elevation as an input.** You can also give elevation to the base kernel +# as a third distance column. This is a different idea: it makes stations at +# different elevations less correlated. It combines with the location +# functions, and on these data the combination is slightly better again. +# - **Cost.** The location function runs once for each row of a kernel matrix, +# so a nonstationary kernel costs about the same as its base kernel. +# - **Pathwise sampling.** These kernels have no spectral density, so +# `sample_approx` does not support them. Sample from the predictive +# distribution instead. +# +# ## References +# +# - Gibbs, M. N. (1997). *Bayesian Gaussian processes for regression and +# classification*. PhD thesis, University of Cambridge. +# - Paciorek, C. J. and Schervish, M. J. (2006). Spatial modelling using a new +# class of nonstationary covariance functions. *Environmetrics* 17, 483–506. +# - Palecki, M. et al. (2021). *U.S. Climate Normals 1991–2020*. NOAA National +# Centers for Environmental Information. diff --git a/docs/examples/spacetime_temperature.py b/docs/examples/spacetime_temperature.py new file mode 100644 index 000000000..5432aaa07 --- /dev/null +++ b/docs/examples/spacetime_temperature.py @@ -0,0 +1,275 @@ +# --- +# 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] +# # Space–Time Modelling of Winter Temperature +# +# Download this notebook: {nb-download}`spacetime_temperature.ipynb` +# +# A simple space–time kernel is *separable*: it is the product of a kernel in +# space and a kernel in time. Then the spatial correlation has the same shape at +# every time lag. Weather does not behave like this. A large anomaly, for +# example a blocking high over Scandinavia, persists for many days, while a +# small anomaly is gone after one or two days. So the spatial correlation +# between two days that are far apart comes mostly from the large anomalies, +# and it is *broader* than the spatial correlation on the same day. +# +# The [`Gneiting`](#gpjax.kernels.Gneiting) kernel (Gneiting, 2002) models this +# with one interaction parameter. In this notebook we fit it to daily +# temperature anomalies over Europe in the winter of 2019, and compare it with +# its separable version on the marginal likelihood and on held-out data. + +# %% 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) + +# %% +import os +from pathlib import Path + +from jax import config +import jax.numpy as jnp +from jaxtyping import install_import_hook +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +import paramax +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 + +gpx.plotting.use_style() + +# Smoke-render flag: set GPJAX_DOCS_CI=1 to shrink the optimiser for fast CI builds. +ci = os.environ.get("GPJAX_DOCS_CI") == "1" +max_iters = 25 if ci else 500 + +# %% [markdown] +# ## The data +# +# We use daily near-surface air temperature from the NCEP-NCAR Reanalysis 1 +# (Kalnay et al., 1996), from 1 January to 1 March 2019. The anomaly is the +# daily value minus the 1991–2020 mean for that day of the year. The file in +# this repository is a coarse subset: 35 grid cells, 7.5° apart in latitude and +# 10° apart in longitude, over Europe and the north-east Atlantic. See +# `docs/examples/data/_pull_reference_datasets.py` for how it was made. + +# %% +DATA = Path("data") if Path("data").exists() else Path("docs/examples/data") +field = xr.open_dataset(DATA / "ncep_europe_winter_2019.nc", engine="scipy")[ + "tas_anomaly" +] +field + +# %% +limit = float(np.abs(field).max()) +fig, axes = plt.subplots(1, 4, figsize=(12, 2.8), sharey=True) +for ax, index in zip(axes, [0, 7, 14, 21], strict=True): + image = field.isel(time=index).plot( + ax=ax, cmap="RdBu_r", vmin=-limit, vmax=limit, add_colorbar=False + ) + ax.set_title(str(field.time.values[index])[:10]) + ax.set_ylabel("Latitude" if index == 0 else "") + ax.set_xlabel("Longitude") +fig.colorbar(image, ax=axes, label="Anomaly (K)") + +# %% [markdown] +# The anomalies are up to several thousand kilometres wide, and the pattern +# changes from one week to the next. +# +# ## Inputs and a held-out gap +# +# The inputs have three columns: the day number, and east and north positions +# in units of 1000 km on a local plane. We hold out 20% of the grid cells for +# the middle third of the period. This is like a gap in the record of some +# stations, which a model must fill from their neighbours in space and in time. + +# %% +lat, lon = np.meshgrid(field.lat, field.lon, indexing="ij") +east = np.deg2rad(lon - lon.mean()) * 6371 * np.cos(np.deg2rad(lat.mean())) / 1000 +north = np.deg2rad(lat - lat.mean()) * 6371 / 1000 + +n_days, n_cells = field.sizes["time"], lat.size +day = np.repeat(np.arange(n_days), n_cells).astype(float) +cell = np.tile(np.arange(n_cells), n_days) +X = np.column_stack([day, east.ravel()[cell], north.ravel()[cell]]) +y = field.to_numpy().astype(np.float64).reshape(-1) + +held_cells = np.random.default_rng(0).choice(n_cells, n_cells // 5, replace=False) +in_gap = (day >= n_days // 3) & (day < 2 * n_days // 3) +test = np.isin(cell, held_cells) & in_gap +train = ~test + +y_mean, y_scale = y[train].mean(), y[train].std() +y = (y - y_mean) / y_scale +train_data = gpx.Dataset(X=jnp.asarray(X[train]), y=jnp.asarray(y[train, None])) +print(f"{train.sum()} training points, {test.sum()} held-out points") + +# %% [markdown] +# ## The Gneiting kernel +# +# For a spatial separation $h$ and a time lag $u$, the kernel is +# +# $$ +# k(h, u) = \frac{\sigma^2}{\psi(u)^{d/2}} +# \exp\!\left(-\frac{(\lVert h\rVert/\ell_s)^{2\gamma}}{\psi(u)^{\beta\gamma}}\right), +# \qquad \psi(u) = \left(\frac{\lvert u\rvert}{\ell_t}\right)^{2\alpha} + 1, +# $$ +# +# where $d = 2$ is the number of space columns. The spatial lengthscale is +# effectively $\ell_s\,\psi(u)^{\beta/2}$, so it grows with the time lag when +# $\beta > 0$. With $\beta = 0$ the kernel is separable: a powered exponential +# kernel in space times a generalised Cauchy kernel in time. So the separable +# model is the same kernel with one parameter fixed, and the comparison is fair. +# +# The kernel reads its columns with `space_dims` and `time_dim`. To fix $\beta$ +# at 0, we pass a non-trainable value. + + +# %% +def build_model(beta): + kernel = gpx.kernels.Gneiting(space_dims=[1, 2], time_dim=0, beta=beta) + prior = gpx.gps.Prior(mean_function=gpx.mean_functions.Zero(), kernel=kernel) + return prior * gpx.likelihoods.Gaussian(obs_stddev=0.3) + + +def negative_mll(model, data): + return -gpx.objectives.conjugate_mll(model, data) + + +candidates = { + "Separable (β = 0)": build_model(paramax.non_trainable(jnp.array(0.0))), + "Nonseparable (β fitted)": build_model(0.5), +} +fitted, log_marginal_likelihood = {}, {} +for name, model in candidates.items(): + fitted[name], history = gpx.fit_scipy( + model=model, + objective=negative_mll, + train_data=train_data, + max_iters=max_iters, + verbose=False, + ) + log_marginal_likelihood[name] = -float(history[-1]) + +# %% [markdown] +# ## Results +# +# We compare the log marginal likelihood on the training data, and three scores +# on the held-out gap: the negative log predictive density (NLPD, lower is +# better), the root-mean-square error in kelvin, and the coverage of the 90% +# predictive intervals. + +# %% +z90 = 1.6449 + + +def score(name): + model = fitted[name] + latent = model.predict(jnp.asarray(X[test]), train_data, covariance="diagonal") + mean = np.asarray(latent.mean) + variance = ( + np.asarray(latent.variance) + float(val(model.likelihood.obs_stddev)) ** 2 + ) + residual = y[test] - mean + kernel = model.prior.kernel + return { + "β": float(val(kernel.beta)), + "Log marginal likelihood": log_marginal_likelihood[name], + "Held-out NLPD": np.mean( + 0.5 * np.log(2 * np.pi * variance) + 0.5 * residual**2 / variance + ), + "Held-out RMSE (K)": np.sqrt(np.mean(residual**2)) * y_scale, + "90% coverage": np.mean(np.abs(residual) < z90 * np.sqrt(variance)), + } + + +pd.DataFrame({name: score(name) for name in fitted}).T.round(3) + +# %% [markdown] +# The interaction is strong. The fitted $\beta$ is at its upper bound of 1, the +# log marginal likelihood is much higher than for the separable model, and the +# held-out NLPD and RMSE are lower. The coverage of both models is close to +# 0.90. The improvement is in the shape of the covariance, and the separable +# model compensates for it with more observation noise. +# +# ## What the interaction looks like +# +# We plot the spatial correlation as a function of distance, at time lags of 0 +# to 3 days, for each fitted kernel. For each lag, we divide by the covariance +# at zero distance, so that only the shape of the spatial correlation remains. + +# %% +distance = np.linspace(0.0, 4.0, 200) +fig, axes = plt.subplots(1, 2, figsize=(10, 3.6), sharey=True) +for ax, name in zip(axes, fitted, strict=True): + kernel = fitted[name].prior.kernel + for lag in range(4): + origin = jnp.array([0.0, 0.0, 0.0]) + points = jnp.column_stack( + [jnp.full_like(distance, lag), distance, jnp.zeros_like(distance)] + ) + covariance = np.asarray(kernel.cross_covariance(origin[None, :], points))[0] + ax.plot(1000 * distance, covariance / covariance[0], label=f"lag {lag} d") + ax.set_title(name) + ax.set_xlabel("Distance (km)") +axes[0].set_ylabel("Spatial correlation at the lag") +axes[1].legend() + +# %% [markdown] +# For the separable kernel, the curves are on top of each other: the shape of +# the spatial correlation does not change with the lag. For the nonseparable +# kernel, the spatial correlation becomes broader as the lag grows, because +# only the large anomalies persist from one day to the next. +# +# ## Notes +# +# - **Bounds.** The trainable $\alpha$, $\beta$ and $\gamma$ are in the open +# interval $(0, 1)$. To fix one of them at a bound, pass +# `paramax.non_trainable(jnp.array(value))`, as we did for $\beta = 0$. +# - **Symmetry.** The Gneiting kernel is symmetric in space and in time. It +# cannot represent advection, where anomalies move in one main direction. +# For daily wind speeds in Ireland (Haslett & Raftery, 1989), which westerly +# winds carry from west to east, we found no benefit over the separable +# kernel. Gneiting, Genton & Guttorp (2007) discuss asymmetric models. +# - **Monthly data.** Monthly anomalies over Europe showed no interaction: the +# temporal correlation is shorter than one month, so the time lags carry +# little information about it. +# - **Pathwise sampling.** The kernel has no closed-form spectral density, so +# `sample_approx` does not support it. Sample from the predictive +# distribution instead. +# +# ## References +# +# - Gneiting, T. (2002). Nonseparable, stationary covariance functions for +# space–time data. *Journal of the American Statistical Association* 97, +# 590–600. +# - Gneiting, T., Genton, M. G. and Guttorp, P. (2007). Geostatistical +# space–time models, stationarity, separability and full symmetry. In +# *Statistical Methods for Spatio-Temporal Systems*, 151–175. Chapman & +# Hall/CRC. +# - Haslett, J. and Raftery, A. E. (1989). Space–time modelling with +# long-memory dependence: assessing Ireland's wind power resource. *Applied +# Statistics* 38, 1–50. +# - Kalnay, E. et al. (1996). The NCEP/NCAR 40-year reanalysis project. +# *Bulletin of the American Meteorological Society* 77, 437–471. diff --git a/docs/index.md b/docs/index.md index c75c7e81c..b28eb0fc4 100644 --- a/docs/index.md +++ b/docs/index.md @@ -177,8 +177,10 @@ examples/barycentres examples/graph_kernels examples/heteroscedastic_inference examples/multioutput +examples/nonstationary_terrain examples/oak examples/oceanmodelling +examples/spacetime_temperature examples/spatial_linear_gp examples/xarray_workflow examples/yacht diff --git a/docs/reference/kernels.md b/docs/reference/kernels.md index e3bee1f7e..95dd14b12 100644 --- a/docs/reference/kernels.md +++ b/docs/reference/kernels.md @@ -17,6 +17,8 @@ DenseKernelComputation DiagonalKernelComputation EigenKernelComputation + Gibbs + Gneiting GraphKernel ICMKernel LCMKernel @@ -33,5 +35,27 @@ ProductKernel RationalQuadratic SumKernel + VaryingAmplitude White ``` + +## Location functions + +A location function gives a kernel parameter that changes with input location, +such as the standard deviation of {class}`~gpjax.kernels.VaryingAmplitude` or +the lengthscale of {class}`~gpjax.kernels.Gibbs`. It is not a mean function: a +mean function describes the Gaussian process, but a location function describes +its covariance. A location function evaluates one input point, selects its own +columns with `active_dims`, and returns a value on the log scale. + +```{eval-rst} +.. currentmodule:: gpjax.kernels.location_functions + +.. autosummary:: + :toctree: generated/ + :nosignatures: + + AbstractLocationFunction + Constant + Linear +``` diff --git a/gpjax/kernels/__init__.py b/gpjax/kernels/__init__.py index 5ad342b86..40cd52398 100644 --- a/gpjax/kernels/__init__.py +++ b/gpjax/kernels/__init__.py @@ -15,7 +15,10 @@ """JaxKern.""" -from gpjax.kernels import stationary +from gpjax.kernels import ( + location_functions, + stationary, +) from gpjax.kernels.additive import ( OrthogonalAdditiveKernel, ) @@ -42,11 +45,14 @@ from gpjax.kernels.non_euclidean import GraphKernel from gpjax.kernels.nonstationary import ( ArcCosine, + Gibbs, Linear, Polynomial, + VaryingAmplitude, ) from gpjax.kernels.stationary import ( RBF, + Gneiting, Matern12, Matern32, Matern52, @@ -67,6 +73,8 @@ "DenseKernelComputation", "DiagonalKernelComputation", "EigenKernelComputation", + "Gibbs", + "Gneiting", "GraphKernel", "ICMKernel", "LCMKernel", @@ -83,6 +91,8 @@ "ProductKernel", "RationalQuadratic", "SumKernel", + "VaryingAmplitude", "White", + "location_functions", "stationary", ] diff --git a/gpjax/kernels/approximations/rff.py b/gpjax/kernels/approximations/rff.py index d9c9595f4..e36c70a56 100644 --- a/gpjax/kernels/approximations/rff.py +++ b/gpjax/kernels/approximations/rff.py @@ -36,7 +36,7 @@ class RFF(AbstractKernel): def __init__( self, - base_kernel: StationaryKernel, + base_kernel: AbstractKernel, num_basis_fns: int = 50, frequencies: tp.Union[Float[Array, "M D"], None] = None, compute_engine: BasisFunctionComputation = BasisFunctionComputation(), @@ -88,7 +88,12 @@ def _check_valid_base_kernel(kernel: AbstractKernel): kernel (AbstractKernel): The kernel to be checked. """ if not isinstance(kernel, StationaryKernel): - raise TypeError("RFF can only be applied to stationary kernels.") + raise TypeError( + "RFF needs a stationary kernel with a spectral density, but got " + f"{type(kernel).__name__}. Pathwise sampling (`sample_approx`) " + "uses RFF, so it does not support this kernel. Sample from the " + "predictive distribution instead." + ) # check that the kernel has a spectral density _ = kernel.spectral_density diff --git a/gpjax/kernels/location_functions.py b/gpjax/kernels/location_functions.py new file mode 100644 index 000000000..3b1672582 --- /dev/null +++ b/gpjax/kernels/location_functions.py @@ -0,0 +1,195 @@ +# 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"""Location functions: kernel parameters that change with input location. + +A location function gives one value at each input location, for example the +standard deviation of a :class:`~gpjax.kernels.VaryingAmplitude` kernel or the +lengthscale of a :class:`~gpjax.kernels.Gibbs` kernel. + +A location function is not a mean function. A mean function describes the +Gaussian process itself, but a location function describes its covariance. +The two types also have different contracts: + +- A location function evaluates one input point of shape `(D,)` and returns a + scalar. JAX's `vmap` then evaluates it once for each row of a kernel matrix. +- A location function selects its own input columns with `active_dims`. This + lets a kernel measure distance over spatial columns while its location + function reads covariate columns, such as elevation. +- A location function returns a value on the log scale. The kernel applies + `exp`, so the parameter it controls is always positive and a coefficient + multiplies that parameter: a weight $\beta$ multiplies the parameter by + $e^{\beta}$ for each unit of its column. +""" + +import abc + +import beartype.typing as tp +import equinox as eqx +import jax.numpy as jnp +from jaxtyping import ( + Float, + Num, +) +from paramax import AbstractUnwrappable + +from gpjax.parameters import ( + Real, + val, +) +from gpjax.summary import _SummaryMixin +from gpjax.typing import ( + Array, + ScalarFloat, +) + + +class AbstractLocationFunction(_SummaryMixin, eqx.Module): + r"""Base class for location functions. + + A subclass implements `__call__`, which maps one input point to a scalar on + the log scale. Use `slice_input` to read only the columns in + `active_dims`. + """ + + active_dims: tp.Union[list[int], slice] = eqx.field( + static=True, default_factory=lambda: slice(None) + ) + + @abc.abstractmethod + def __call__(self, x: Num[Array, " D"]) -> ScalarFloat: + r"""Evaluate the location function at one input point. + + Args: + x: one input point, with all columns of the data. + + Returns: + The log of the parameter value at `x`. + """ + ... + + def slice_input(self, x: Num[Array, " D"]) -> Num[Array, " Q"]: + r"""Select the columns in `active_dims` from one input point. + + Args: + x: one input point, with all columns of the data. + + Returns: + The selected columns. + """ + return x[..., self.active_dims] + + +class Constant(AbstractLocationFunction): + r"""A location function with the same value at all locations. + + $$\log g(x) = c$$ + + With this function, a nonstationary kernel is equal to its stationary base + kernel with its scale multiplied by $e^{c}$. The default $c = 0$ gives the + base kernel exactly. + """ + + value: tp.Any + + def __init__(self, value: tp.Union[ScalarFloat, AbstractUnwrappable] = 0.0): + """Initialise the location function. + + Args: + value: the log value $c$. A float is wrapped as a trainable `Real`. + """ + if isinstance(value, AbstractUnwrappable): + self.value = value + else: + self.value = Real(jnp.asarray(value, dtype=float)) + self.active_dims = slice(None) + + def __call__(self, x: Num[Array, " D"]) -> ScalarFloat: + return jnp.asarray(val(self.value)).squeeze() + + +class Linear(AbstractLocationFunction): + r"""A location function that is log-linear in selected input columns. + + $$\log g(x) = \beta^{\top} x_{\mathcal{A}} + b$$ + + Here $x_{\mathcal{A}}$ are the columns in `active_dims`. The weights + $\beta$ start at zero, so a model starts as its stationary base kernel and + moves away from it only as far as the data supports. + + By default the function has no intercept ($b = 0$): the base kernel holds + the overall scale, and the location function holds only the change over + location. An intercept would duplicate the variance or lengthscale of the + base kernel, and the data could not separate the two. + + Standardise the covariate columns before fitting, so that the weights have + similar scales. + """ + + weights: tp.Any + bias: tp.Any + + def __init__( + self, + active_dims: list[int], + weights: tp.Union[ + Float[Array, " Q"], list[float], AbstractUnwrappable, None + ] = None, + intercept: bool = False, + ): + r"""Initialise the location function. + + Args: + active_dims: the indices of the covariate columns. This argument is + required, so that the function cannot read the spatial columns + by mistake. + weights: the initial weights $\beta$, one for each column in + `active_dims`. Defaults to zeros. + intercept: whether to add a trainable intercept $b$, which starts + at zero. + """ + if not isinstance(active_dims, list) or not active_dims: + raise TypeError( + "Expected `active_dims` to be a non-empty list of column indices. " + f"Got {active_dims!r}." + ) + + if weights is None: + weights = jnp.zeros(len(active_dims)) + if not isinstance(weights, AbstractUnwrappable): + weights = jnp.asarray(weights, dtype=float) + if weights.shape != (len(active_dims),): + raise ValueError( + f"Expected one weight for each of the {len(active_dims)} " + f"columns in `active_dims`. Got weights of shape " + f"{weights.shape}." + ) + weights = Real(weights) + + self.active_dims = active_dims + self.weights = weights + self.bias = Real(jnp.array(0.0)) if intercept else None + + def __call__(self, x: Num[Array, " D"]) -> ScalarFloat: + value = jnp.dot(self.slice_input(x), val(self.weights)) + if self.bias is not None: + value = value + val(self.bias) + return value.squeeze() + + +__all__ = [ + "AbstractLocationFunction", + "Constant", + "Linear", +] diff --git a/gpjax/kernels/nonstationary/__init__.py b/gpjax/kernels/nonstationary/__init__.py index 11e2a179f..fc6025879 100644 --- a/gpjax/kernels/nonstationary/__init__.py +++ b/gpjax/kernels/nonstationary/__init__.py @@ -14,7 +14,9 @@ # ============================================================================== from gpjax.kernels.nonstationary.arccosine import ArcCosine +from gpjax.kernels.nonstationary.gibbs import Gibbs from gpjax.kernels.nonstationary.linear import Linear from gpjax.kernels.nonstationary.polynomial import Polynomial +from gpjax.kernels.nonstationary.varying_amplitude import VaryingAmplitude -__all__ = ["ArcCosine", "Linear", "Polynomial"] +__all__ = ["ArcCosine", "Gibbs", "Linear", "Polynomial", "VaryingAmplitude"] diff --git a/gpjax/kernels/nonstationary/gibbs.py b/gpjax/kernels/nonstationary/gibbs.py new file mode 100644 index 000000000..91677b2e6 --- /dev/null +++ b/gpjax/kernels/nonstationary/gibbs.py @@ -0,0 +1,111 @@ +# 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. +# ============================================================================== + +from typing import ClassVar + +import jax.numpy as jnp +from jaxtyping import Float + +from gpjax.kernels.base import AbstractKernel +from gpjax.kernels.computations import ( + AbstractKernelComputation, + DenseKernelComputation, +) +from gpjax.kernels.location_functions import AbstractLocationFunction +from gpjax.typing import ( + Array, + ScalarFloat, +) + + +class Gibbs(AbstractKernel): + r"""A base kernel whose lengthscale changes with location. + + The Gibbs kernel (Gibbs, 1997), also known as the Paciorek–Schervish + kernel (Paciorek & Schervish, 2006). A location function $g$ gives + $\log\ell(x)$, and $\ell(x)$ multiplies the lengthscale of an isotropic + base kernel $k_0$ with correlation $\rho$ and variance $\sigma^2$: + $$ + k(x, y) = \sigma^2 + \left(\frac{2\,\ell(x)\,\ell(y)}{\ell(x)^2 + \ell(y)^2}\right)^{d/2} + \rho\!\left(\sqrt{\frac{2}{\ell(x)^2 + \ell(y)^2}}\, + \lVert x - y\rVert\right), + \qquad \ell(x) = \exp g(x). + $$ + + Here $d$ is the number of columns over which the base kernel measures + distance, and $\lVert\cdot\rVert$ uses the base kernel's lengthscales, so + an ARD base kernel keeps its shape and $\ell(x)$ scales it. Correlation + decays faster where $\ell(x)$ is small, for example over mountains, and + more slowly where it is large. The marginal variance is $\sigma^2$ at every + location. + + The kernel is positive definite when $\rho$ is positive definite in every + dimension. Only base kernels with `isotropic_radial = True` meet this + condition: RBF, the Matérn kernels, RationalQuadratic and + PoweredExponential. + + The base kernel selects the columns over which it measures distance with + its own `active_dims`, and the location function selects its covariate + columns. The wrapper itself always receives every column. + """ + + name: ClassVar[str] = "Gibbs" + base_kernel: AbstractKernel + lengthscale: AbstractLocationFunction + + def __init__( + self, + base_kernel: AbstractKernel, + lengthscale: AbstractLocationFunction, + compute_engine: AbstractKernelComputation = DenseKernelComputation(), + ): + r"""Initialise the kernel. + + Args: + base_kernel: the isotropic kernel $k_0$ whose lengthscale changes. + lengthscale: the location function that gives $\log\ell(x)$. + compute_engine: the computation engine that the kernel uses to + compute its covariance matrices. + + Raises: + TypeError: if `base_kernel` is not an isotropic radial kernel. + """ + if not getattr(base_kernel, "isotropic_radial", False): + raise TypeError( + "Gibbs needs an isotropic radial base kernel: RBF, Matern12, " + "Matern32, Matern52, RationalQuadratic or PoweredExponential. " + f"Got {type(base_kernel).__name__}, for which the Gibbs " + "construction is not guaranteed to be positive definite." + ) + self.base_kernel = base_kernel + self.lengthscale = lengthscale + super().__init__(compute_engine=compute_engine) + + def __call__(self, x: Float[Array, " D"], y: Float[Array, " D"]) -> ScalarFloat: + log_lx = self.lengthscale(x) + log_ly = self.lengthscale(y) + # log(ℓ(x)² + ℓ(y)²), computed stably for large or small lengthscales. + log_sum = jnp.logaddexp(2.0 * log_lx, 2.0 * log_ly) + log_ratio = jnp.log(2.0) + log_lx + log_ly - log_sum + dims = self.base_kernel.slice_input(x).shape[-1] + prefactor = jnp.exp(0.5 * dims * log_ratio) + # Scaling both inputs by the same factor scales their distance, so the + # base kernel evaluates ρ at the Gibbs distance with its own variance. + scale = jnp.exp(0.5 * (jnp.log(2.0) - log_sum)) + return (prefactor * self.base_kernel(scale * x, scale * y)).squeeze() + + +__all__ = ["Gibbs"] diff --git a/gpjax/kernels/nonstationary/varying_amplitude.py b/gpjax/kernels/nonstationary/varying_amplitude.py new file mode 100644 index 000000000..a4ef38162 --- /dev/null +++ b/gpjax/kernels/nonstationary/varying_amplitude.py @@ -0,0 +1,93 @@ +# 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. +# ============================================================================== + +from typing import ClassVar + +import jax.numpy as jnp +from jaxtyping import Float + +from gpjax.kernels.approximations import RFF +from gpjax.kernels.base import AbstractKernel +from gpjax.kernels.computations import ( + AbstractKernelComputation, + DenseKernelComputation, +) +from gpjax.kernels.location_functions import AbstractLocationFunction +from gpjax.typing import ( + Array, + ScalarFloat, +) + + +class VaryingAmplitude(AbstractKernel): + r"""A base kernel whose standard deviation changes with location. + + Computes the covariance for a pair of inputs $(x, y)$ from a base kernel + $k_0$ and a location function $g$ that gives $\log\sigma(x)$: + $$ + k(x, y) = \sigma(x)\,\sigma(y)\,k_0(x, y), \qquad \sigma(x) = \exp g(x). + $$ + + The kernel is positive definite for every base kernel that is positive + definite. The marginal variance at $x$ is $\sigma(x)^2 k_0(x, x)$, so the + base kernel's variance holds the overall scale and $\sigma(x)$ holds the + change over location. Use it when variability is larger in some regions + than in others, for example over land than over sea. + + The base kernel selects the columns over which it measures distance with + its own `active_dims`, and the location function selects its covariate + columns. The wrapper itself always receives every column. + """ + + name: ClassVar[str] = "Varying amplitude" + base_kernel: AbstractKernel + amplitude: AbstractLocationFunction + + def __init__( + self, + base_kernel: AbstractKernel, + amplitude: AbstractLocationFunction, + compute_engine: AbstractKernelComputation = DenseKernelComputation(), + ): + r"""Initialise the kernel. + + Args: + base_kernel: the kernel $k_0$ whose amplitude changes. It must + evaluate one pair of points at a time. + amplitude: the location function that gives $\log\sigma(x)$. + compute_engine: the computation engine that the kernel uses to + compute its covariance matrices. + """ + _check_pointwise(base_kernel, type(self).__name__) + self.base_kernel = base_kernel + self.amplitude = amplitude + super().__init__(compute_engine=compute_engine) + + def __call__(self, x: Float[Array, " D"], y: Float[Array, " D"]) -> ScalarFloat: + scale = jnp.exp(self.amplitude(x) + self.amplitude(y)) + return (scale * self.base_kernel(x, y)).squeeze() + + +def _check_pointwise(base_kernel: AbstractKernel, wrapper: str) -> None: + # RFF returns None from __call__: it builds its matrices from features. + if isinstance(base_kernel, RFF): + raise TypeError( + f"{wrapper} needs a base kernel that evaluates one pair of points at " + "a time, but RFF builds its matrices from random features. Use the " + "kernel that RFF approximates as the base kernel." + ) + + +__all__ = ["VaryingAmplitude"] diff --git a/gpjax/kernels/stationary/__init__.py b/gpjax/kernels/stationary/__init__.py index 6ebac3fd6..03d985b3d 100644 --- a/gpjax/kernels/stationary/__init__.py +++ b/gpjax/kernels/stationary/__init__.py @@ -14,6 +14,7 @@ # ============================================================================== from gpjax.kernels.stationary.base import StationaryKernel +from gpjax.kernels.stationary.gneiting import Gneiting from gpjax.kernels.stationary.matern12 import Matern12 from gpjax.kernels.stationary.matern32 import Matern32 from gpjax.kernels.stationary.matern52 import Matern52 @@ -25,6 +26,7 @@ __all__ = [ "RBF", + "Gneiting", "Matern12", "Matern32", "Matern52", diff --git a/gpjax/kernels/stationary/base.py b/gpjax/kernels/stationary/base.py index 722a5080f..28f0850bb 100644 --- a/gpjax/kernels/stationary/base.py +++ b/gpjax/kernels/stationary/base.py @@ -14,6 +14,8 @@ # ============================================================================== +from typing import ClassVar + import beartype.typing as tp import equinox as eqx import jax.numpy as jnp @@ -49,6 +51,10 @@ class StationaryKernel(AbstractKernel): for each input dimension. """ + # True when the kernel is a function of the lengthscale-scaled Euclidean + # distance only and is positive definite in every dimension. Gibbs accepts + # only such base kernels. + isotropic_radial: ClassVar[bool] = False lengthscale: AbstractUnwrappable = eqx.field( default_factory=lambda: PositiveReal(1.0) ) diff --git a/gpjax/kernels/stationary/gneiting.py b/gpjax/kernels/stationary/gneiting.py new file mode 100644 index 000000000..56372f86d --- /dev/null +++ b/gpjax/kernels/stationary/gneiting.py @@ -0,0 +1,187 @@ +# 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. +# ============================================================================== + +from typing import ClassVar + +import beartype.typing as tp +import equinox as eqx +import jax.numpy as jnp +from jaxtyping import Float +from paramax import AbstractUnwrappable + +from gpjax.kernels.base import ( + AbstractKernel, + val, +) +from gpjax.kernels.computations import ( + AbstractKernelComputation, + DenseKernelComputation, +) +from gpjax.parameters import ( + NonNegativeReal, + PositiveReal, + SigmoidBounded, +) +from gpjax.typing import ( + Array, + ScalarFloat, +) + + +class Gneiting(AbstractKernel): + r"""The Gneiting nonseparable space–time kernel. + + Computes the covariance for a pair of inputs with spatial separation $h$ + over the space columns and time lag $u$ over the time column (Gneiting, + 2002, eq. 14): + $$ + k(h, u) = \frac{\sigma^2}{\psi(u)^{d/2}} + \exp\!\left(-\frac{(\lVert h\rVert/\ell_s)^{2\gamma}} + {\psi(u)^{\beta\gamma}}\right), + \qquad \psi(u) = \left(\frac{\lvert u\rvert}{\ell_t}\right)^{2\alpha} + 1, + $$ + where $d$ is the number of space columns. + + The kernel is stationary, but it is not separable: as the time lag grows, + $\psi(u)$ grows and the spatial correlation decays more slowly. The + interaction parameter $\beta \in [0, 1]$ controls this effect, and + $\beta = 0$ gives the separable product of a powered exponential kernel in + space and a generalised Cauchy kernel in time. $\alpha \in (0, 1]$ and + $\gamma \in (0, 1]$ set the smoothness in time and in space. + + The trainable parameters $\alpha$, $\beta$ and $\gamma$ are bounded to the + open interval $(0, 1)$. To fix one of them at a bound, pass a + non-trainable value, for example `paramax.non_trainable(jnp.array(1.0))`. + + The kernel has two lengthscales and no closed-form spectral density, so it + is not a :class:`StationaryKernel` subclass and it does not support random + Fourier features. + """ + + name: ClassVar[str] = "Gneiting" + space_dims: list[int] = eqx.field(static=True) + time_dim: int = eqx.field(static=True) + variance: tp.Any + space_lengthscale: tp.Any + time_lengthscale: tp.Any + alpha: tp.Any + beta: tp.Any + gamma: tp.Any + + def __init__( + self, + space_dims: list[int], + time_dim: int, + variance: tp.Union[ScalarFloat, AbstractUnwrappable] = 1.0, + space_lengthscale: tp.Union[ScalarFloat, AbstractUnwrappable] = 1.0, + time_lengthscale: tp.Union[ScalarFloat, AbstractUnwrappable] = 1.0, + alpha: tp.Union[ScalarFloat, AbstractUnwrappable] = 0.5, + beta: tp.Union[ScalarFloat, AbstractUnwrappable] = 0.5, + gamma: tp.Union[ScalarFloat, AbstractUnwrappable] = 0.5, + compute_engine: AbstractKernelComputation = DenseKernelComputation(), + ): + r"""Initialise the kernel. + + Args: + space_dims: the indices of the space columns. + time_dim: the index of the time column. + variance: the variance $\sigma^2$. + space_lengthscale: the spatial lengthscale $\ell_s$. + time_lengthscale: the temporal lengthscale $\ell_t$. + alpha: the smoothness in time, $\alpha \in (0, 1]$. + beta: the space–time interaction, $\beta \in [0, 1]$. + gamma: the smoothness in space, $\gamma \in (0, 1]$. + compute_engine: the computation engine that the kernel uses to + compute its covariance matrices. + + Raises: + ValueError: if the columns are not valid, or if a float value of + `alpha`, `beta` or `gamma` is not in the open interval + $(0, 1)$. + """ + _check_columns(space_dims, time_dim) + self.space_dims = list(space_dims) + self.time_dim = time_dim + self.variance = _wrap(variance, NonNegativeReal) + self.space_lengthscale = _wrap(space_lengthscale, PositiveReal) + self.time_lengthscale = _wrap(time_lengthscale, PositiveReal) + self.alpha = _wrap_unit(alpha, "alpha") + self.beta = _wrap_unit(beta, "beta") + self.gamma = _wrap_unit(gamma, "gamma") + super().__init__(compute_engine=compute_engine) + + def __call__(self, x: Float[Array, " D"], y: Float[Array, " D"]) -> ScalarFloat: + h = (x[..., self.space_dims] - y[..., self.space_dims]) / val( + self.space_lengthscale + ) + u = (x[..., self.time_dim] - y[..., self.time_dim]) / val(self.time_lengthscale) + gamma = val(self.gamma) + psi = _power(u**2, val(self.alpha)) + 1.0 + space_term = _power(jnp.sum(h**2), gamma) / psi ** (val(self.beta) * gamma) + dims = len(self.space_dims) + K = val(self.variance) * psi ** (-0.5 * dims) * jnp.exp(-space_term) + return K.squeeze() + + +def _power(t: Float[Array, ""], p: Float[Array, ""]) -> Float[Array, ""]: + # t**p with t >= 0. The gradient of 0**p with respect to p, or of t**p at + # t = 0 for p < 1, is not finite, so the zero case takes a separate branch. + positive = t > 0 + safe = jnp.where(positive, t, 1.0) + return jnp.where(positive, safe**p, 0.0) + + +def _check_columns(space_dims: tp.Any, time_dim: tp.Any) -> None: + if ( + not isinstance(space_dims, (list, tuple)) + or not space_dims + or not all(isinstance(i, int) for i in space_dims) + ): + raise ValueError( + "Expected `space_dims` to be a non-empty list of column indices. " + f"Got {space_dims!r}." + ) + if len(set(space_dims)) != len(space_dims): + raise ValueError(f"`space_dims` has repeated columns: {space_dims!r}.") + if not isinstance(time_dim, int): + raise ValueError( + f"Expected `time_dim` to be one column index. Got {time_dim!r}." + ) + if time_dim in space_dims: + raise ValueError( + f"Column {time_dim} is both a space column and the time column." + ) + + +def _wrap(value: tp.Any, parameter: type) -> tp.Any: + if isinstance(value, AbstractUnwrappable): + return value + return parameter(jnp.asarray(value, dtype=float)) + + +def _wrap_unit(value: tp.Any, label: str) -> tp.Any: + if isinstance(value, AbstractUnwrappable): + return value + value = jnp.asarray(value, dtype=float) + if not 0.0 < float(value) < 1.0: + raise ValueError( + f"Expected `{label}` in the open interval (0, 1), so that it can be " + f"trained. Got {float(value)}. To fix it at a bound, pass " + "`paramax.non_trainable(jnp.array(value))`." + ) + return SigmoidBounded(value, low=0.0, high=1.0) + + +__all__ = ["Gneiting"] diff --git a/gpjax/kernels/stationary/matern12.py b/gpjax/kernels/stationary/matern12.py index a10dd4b75..c17989a4b 100644 --- a/gpjax/kernels/stationary/matern12.py +++ b/gpjax/kernels/stationary/matern12.py @@ -43,6 +43,7 @@ class Matern12(StationaryKernel): """ name: ClassVar[str] = "Matérn12" + isotropic_radial: ClassVar[bool] = True def __call__(self, x: Float[Array, " D"], y: Float[Array, " D"]) -> ScalarFloat: x = self.slice_input(x) / val(self.lengthscale) diff --git a/gpjax/kernels/stationary/matern32.py b/gpjax/kernels/stationary/matern32.py index de101d130..51b7f74f8 100644 --- a/gpjax/kernels/stationary/matern32.py +++ b/gpjax/kernels/stationary/matern32.py @@ -40,6 +40,7 @@ class Matern32(StationaryKernel): """ name: ClassVar[str] = "Matérn32" + isotropic_radial: ClassVar[bool] = True def __call__( self, diff --git a/gpjax/kernels/stationary/matern52.py b/gpjax/kernels/stationary/matern52.py index 3c721f806..67d832e2d 100644 --- a/gpjax/kernels/stationary/matern52.py +++ b/gpjax/kernels/stationary/matern52.py @@ -41,6 +41,7 @@ class Matern52(StationaryKernel): """ name: ClassVar[str] = "Matérn52" + isotropic_radial: ClassVar[bool] = True def __call__( self, x: Float[Array, " D"], y: Float[Array, " D"] diff --git a/gpjax/kernels/stationary/powered_exponential.py b/gpjax/kernels/stationary/powered_exponential.py index f2d72d1fe..9c0149207 100644 --- a/gpjax/kernels/stationary/powered_exponential.py +++ b/gpjax/kernels/stationary/powered_exponential.py @@ -53,6 +53,7 @@ class PoweredExponential(StationaryKernel): """ name: ClassVar[str] = "Powered Exponential" + isotropic_radial: ClassVar[bool] = True power: tp.Any def __init__( diff --git a/gpjax/kernels/stationary/rational_quadratic.py b/gpjax/kernels/stationary/rational_quadratic.py index 3bfda9362..860a80afd 100644 --- a/gpjax/kernels/stationary/rational_quadratic.py +++ b/gpjax/kernels/stationary/rational_quadratic.py @@ -51,6 +51,7 @@ class RationalQuadratic(StationaryKernel): """ name: ClassVar[str] = "Rational Quadratic" + isotropic_radial: ClassVar[bool] = True alpha: tp.Any def __init__( diff --git a/gpjax/kernels/stationary/rbf.py b/gpjax/kernels/stationary/rbf.py index c6263c3be..8bd8837f1 100644 --- a/gpjax/kernels/stationary/rbf.py +++ b/gpjax/kernels/stationary/rbf.py @@ -39,6 +39,7 @@ class RBF(StationaryKernel): """ name: ClassVar[str] = "RBF" + isotropic_radial: ClassVar[bool] = True def __call__(self, x: Float[Array, " D"], y: Float[Array, " D"]) -> ScalarFloat: x = self.slice_input(x) / val(self.lengthscale) diff --git a/tests/test_kernels/test_location_functions.py b/tests/test_kernels/test_location_functions.py new file mode 100644 index 000000000..b001d1634 --- /dev/null +++ b/tests/test_kernels/test_location_functions.py @@ -0,0 +1,105 @@ +# 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. +# ============================================================================== + +from gpjax.kernels.location_functions import ( + AbstractLocationFunction, + Constant, + Linear, +) +from gpjax.parameters import Real +import jax +from jax import config +import jax.numpy as jnp +import paramax +import pytest + +config.update("jax_enable_x64", True) + +X = jnp.array([0.3, -1.2, 2.0, 0.5]) + + +def test_constant_returns_its_value(): + assert jnp.allclose(Constant(0.7)(X), 0.7) + + +def test_constant_defaults_to_zero(): + assert jnp.allclose(Constant()(X), 0.0) + + +def test_constant_accepts_a_parameter(): + frozen = paramax.non_trainable(Real(jnp.array(1.5))) + assert jnp.allclose(Constant(frozen)(X), 1.5) + + +def test_linear_starts_at_zero(): + fn = Linear(active_dims=[1, 3]) + assert jnp.allclose(fn(X), 0.0) + assert fn.bias is None + + +def test_linear_reads_only_its_columns(): + fn = Linear(active_dims=[1, 3], weights=[2.0, -1.0]) + assert jnp.allclose(fn(X), 2.0 * X[1] - 1.0 * X[3]) + + +def test_linear_intercept_starts_at_zero_and_is_trainable(): + fn = Linear(active_dims=[0], weights=[1.0], intercept=True) + assert jnp.allclose(fn(X), X[0]) + shifted = jax.tree_util.tree_map(lambda leaf: leaf + 0.5, fn) + assert jnp.allclose(shifted(X), 1.5 * X[0] + 0.5) + + +def test_linear_requires_active_dims(): + with pytest.raises(TypeError): + Linear() # type: ignore[call-arg] + + +def test_linear_rejects_empty_active_dims(): + with pytest.raises(TypeError, match="non-empty list"): + Linear(active_dims=[]) + + +@pytest.mark.parametrize("active_dims", [slice(None), (0, 1)]) +def test_linear_rejects_active_dims_that_are_not_a_list(active_dims): + with pytest.raises(TypeError): + Linear(active_dims=active_dims) + + +def test_linear_rejects_wrong_number_of_weights(): + with pytest.raises(ValueError, match="one weight for each"): + Linear(active_dims=[0, 1], weights=[1.0]) + + +def test_linear_gradient_is_the_selected_columns(): + fn = Linear(active_dims=[1, 2], weights=[0.1, 0.2]) + grads = jax.grad(lambda f: f(X))(fn) + (gradient,) = jax.tree_util.tree_leaves(grads.weights) + assert jnp.allclose(gradient, X[jnp.array([1, 2])]) + + +def test_location_functions_work_under_jit_and_vmap(): + fn = Linear(active_dims=[0], weights=[2.0]) + rows = jnp.stack([X, 2 * X]) + values = jax.jit(jax.vmap(fn))(rows) + assert jnp.allclose(values, 2.0 * rows[:, 0]) + + +def test_custom_location_function(): + class Quadratic(AbstractLocationFunction): + def __call__(self, x): + return jnp.sum(self.slice_input(x) ** 2) + + fn = Quadratic(active_dims=[0, 1]) + assert jnp.allclose(fn(X), X[0] ** 2 + X[1] ** 2) diff --git a/tests/test_kernels/test_nonstationary.py b/tests/test_kernels/test_nonstationary.py index 2b92a35ea..6457243c4 100644 --- a/tests/test_kernels/test_nonstationary.py +++ b/tests/test_kernels/test_nonstationary.py @@ -17,19 +17,45 @@ from typing import Any import equinox as eqx +from gpjax.dataset import Dataset +from gpjax.fit import fit +from gpjax.gps import Prior +from gpjax.kernels import location_functions +from gpjax.kernels.approximations import RFF from gpjax.kernels.base import AbstractKernel from gpjax.kernels.computations import AbstractKernelComputation from gpjax.kernels.nonstationary import ( ArcCosine, + Gibbs, Linear, Polynomial, + VaryingAmplitude, ) +from gpjax.kernels.stationary import ( + RBF, + Matern12, + Matern32, + Matern52, + Periodic, + PoweredExponential, + RationalQuadratic, + White, +) +from gpjax.likelihoods import Gaussian +from gpjax.mean_functions import Zero +from gpjax.objectives import conjugate_mll from gpjax.parameters import NonNegativeReal, val +from gpjax.summary import _collect +from hypothesis import ( + given, + strategies as st, +) import jax from jax import config import jax.numpy as jnp import jax.random as jr import lineax as lx +import optax as ox from paramax import AbstractUnwrappable import pytest @@ -231,3 +257,195 @@ def loss(p): assert val(k_new.weight_variance) > 0.0 gram = k_new.gram(x).as_matrix() assert jnp.all(jnp.isfinite(gram)) + + +# --------------------------------------------------------------------------- +# Location-function kernels: VaryingAmplitude and Gibbs. +# +# Inputs have two space columns (0, 1) and one covariate column (2). +# --------------------------------------------------------------------------- + +INPUTS = jr.uniform(jr.key(0), (40, 3), minval=-2.0, maxval=2.0) +RADIAL_KERNELS = [ + RBF, + Matern12, + Matern32, + Matern52, + RationalQuadratic, + PoweredExponential, +] + + +def _min_eigenvalue(kernel: AbstractKernel) -> float: + return float(jnp.linalg.eigvalsh(kernel.gram(INPUTS).as_matrix()).min()) + + +def _covariate_fn(weight: float) -> location_functions.Linear: + return location_functions.Linear(active_dims=[2], weights=[weight]) + + +@pytest.mark.parametrize("base", RADIAL_KERNELS) +@pytest.mark.parametrize("wrapper", ["amplitude", "gibbs"]) +def test_zero_location_function_gives_the_base_kernel(base, wrapper): + base_kernel = base(active_dims=[0, 1], lengthscale=0.8, variance=1.7) + fn = location_functions.Constant() + kernel = ( + VaryingAmplitude(base_kernel, amplitude=fn) + if wrapper == "amplitude" + else Gibbs(base_kernel, lengthscale=fn) + ) + assert jnp.allclose( + kernel.gram(INPUTS).as_matrix(), base_kernel.gram(INPUTS).as_matrix() + ) + + +def test_varying_amplitude_scales_the_base_kernel(): + base_kernel = Matern32(active_dims=[0, 1]) + kernel = VaryingAmplitude(base_kernel, amplitude=_covariate_fn(0.6)) + sigma = jnp.exp(0.6 * INPUTS[:, 2]) + expected = sigma[:, None] * sigma[None, :] * base_kernel.gram(INPUTS).as_matrix() + assert jnp.allclose(kernel.gram(INPUTS).as_matrix(), expected) + + +def test_varying_amplitude_diagonal_is_the_local_variance(): + base_kernel = RBF(active_dims=[0, 1], variance=2.0) + kernel = VaryingAmplitude(base_kernel, amplitude=_covariate_fn(-0.4)) + diagonal = kernel.diagonal(INPUTS).as_matrix().diagonal() + assert jnp.allclose(diagonal, 2.0 * jnp.exp(2 * -0.4 * INPUTS[:, 2])) + + +def test_varying_amplitude_accepts_a_nonstationary_base_kernel(): + kernel = VaryingAmplitude(Linear(active_dims=[0]), amplitude=_covariate_fn(0.3)) + assert _min_eigenvalue(kernel) > -1e-8 + + +def test_varying_amplitude_rejects_rff(): + rff = RFF(base_kernel=RBF(n_dims=3), num_basis_fns=5) + with pytest.raises(TypeError, match="one pair of points"): + VaryingAmplitude(rff, amplitude=location_functions.Constant()) + + +def test_wrappers_do_not_take_active_dims(): + with pytest.raises(TypeError): + Gibbs(RBF(), lengthscale=location_functions.Constant(), active_dims=[0]) + + +def _paciorek_schervish(x, y, ell, base_lengthscale, variance): + """Direct Paciorek-Schervish Matern-3/2 with Sigma(x) = ell(x)^2 diag(l0^2).""" + sigma_x = jnp.diag((ell(x) * base_lengthscale) ** 2) + sigma_y = jnp.diag((ell(y) * base_lengthscale) ** 2) + sigma = (sigma_x + sigma_y) / 2 + prefactor = ( + jnp.linalg.det(sigma_x) ** 0.25 + * jnp.linalg.det(sigma_y) ** 0.25 + / jnp.sqrt(jnp.linalg.det(sigma)) + ) + h = x[:2] - y[:2] + r = jnp.sqrt(h @ jnp.linalg.solve(sigma, h) + 1e-36) + return variance * prefactor * (1 + jnp.sqrt(3.0) * r) * jnp.exp(-jnp.sqrt(3.0) * r) + + +def test_gibbs_matches_the_paciorek_schervish_formula(): + base_lengthscale = jnp.array([0.7, 1.3]) + fn = _covariate_fn(0.9) + kernel = Gibbs( + Matern32(active_dims=[0, 1], lengthscale=base_lengthscale, variance=2.0), + lengthscale=fn, + ) + ell = lambda x: jnp.exp(fn(x)) + expected = jax.vmap( + lambda a: jax.vmap( + lambda b: _paciorek_schervish(a, b, ell, base_lengthscale, 2.0) + )(INPUTS) + )(INPUTS) + assert jnp.allclose(kernel.gram(INPUTS).as_matrix(), expected) + + +@pytest.mark.parametrize("base", RADIAL_KERNELS) +def test_gibbs_keeps_the_base_variance(base): + kernel = Gibbs( + base(active_dims=[0, 1], variance=1.7), lengthscale=_covariate_fn(1.2) + ) + diagonal = kernel.diagonal(INPUTS).as_matrix().diagonal() + assert jnp.allclose(diagonal, 1.7) + + +def test_gibbs_correlation_is_shorter_where_the_lengthscale_is_smaller(): + kernel = Gibbs(Matern32(active_dims=[0, 1]), lengthscale=_covariate_fn(1.0)) + low = jnp.array([[0.0, 0.0, -1.0], [0.5, 0.0, -1.0]]) + high = low.at[:, 2].set(1.0) + assert kernel(low[0], low[1]) < kernel(high[0], high[1]) + + +@pytest.mark.parametrize( + "base", + [ + Periodic(), + White(), + RBF() + Matern32(), + RBF() * Matern32(), + Linear(), + ], + ids=["periodic", "white", "sum", "product", "linear"], +) +def test_gibbs_rejects_bases_that_are_not_isotropic_radial(base): + with pytest.raises(TypeError, match="isotropic radial"): + Gibbs(base, lengthscale=location_functions.Constant()) + + +@given( + weights=st.lists(st.floats(min_value=-1.5, max_value=1.5), min_size=2, max_size=2), + base_index=st.integers(min_value=0, max_value=len(RADIAL_KERNELS) - 1), +) +def test_location_function_kernels_are_positive_definite(weights, base_index): + base_kernel = RADIAL_KERNELS[base_index](active_dims=[0, 1], lengthscale=0.5) + gibbs = Gibbs(base_kernel, lengthscale=_covariate_fn(weights[0])) + kernel = VaryingAmplitude(gibbs, amplitude=_covariate_fn(weights[1])) + gram = kernel.gram(INPUTS).as_matrix() + assert jnp.allclose(gram, gram.T) + assert _min_eigenvalue(kernel) > -1e-8 * jnp.max(jnp.diag(gram)) + + +def test_location_function_kernels_fit_and_have_finite_gradients(): + y = jnp.sin(INPUTS[:, :1] * 2.0) * jnp.exp(0.5 * INPUTS[:, 2:3]) + data = Dataset(X=INPUTS, y=y) + kernel = VaryingAmplitude( + Gibbs(Matern52(active_dims=[0, 1]), lengthscale=_covariate_fn(0.0)), + amplitude=_covariate_fn(0.0), + ) + model = Prior(mean_function=Zero(), kernel=kernel) * Gaussian() + + objective = lambda m, d: -conjugate_mll(m, d) + grads = jax.jit(jax.grad(objective))(model, data) + assert all(jnp.all(jnp.isfinite(g)) for g in jax.tree_util.tree_leaves(grads)) + + fitted, history = fit( + model=model, + objective=objective, + train_data=data, + optim=ox.adam(0.05), + num_iters=30, + verbose=False, + ) + assert history[-1] < history[0] + weight = val(fitted.prior.kernel.amplitude.weights) + assert not jnp.allclose(weight, 0.0) + + +def test_summary_shows_location_function_parameters(): + kernel = Gibbs(Matern32(active_dims=[0, 1]), lengthscale=_covariate_fn(0.1)) + names = {r.name for r in _collect(Prior(mean_function=Zero(), kernel=kernel))} + assert "kernel.lengthscale.weights" in names + assert "kernel.base_kernel.lengthscale" in names + + +@pytest.mark.parametrize("wrapper", ["amplitude", "gibbs"]) +def test_rff_names_the_kernel_it_cannot_approximate(wrapper): + fn = location_functions.Constant() + kernel = ( + VaryingAmplitude(RBF(), amplitude=fn) + if wrapper == "amplitude" + else Gibbs(RBF(), lengthscale=fn) + ) + with pytest.raises(TypeError, match=f"{type(kernel).__name__}.*sample_approx"): + RFF(base_kernel=kernel) diff --git a/tests/test_kernels/test_stationary.py b/tests/test_kernels/test_stationary.py index c36c41ea9..9e0811600 100644 --- a/tests/test_kernels/test_stationary.py +++ b/tests/test_kernels/test_stationary.py @@ -20,6 +20,7 @@ from gpjax.kernels.computations import AbstractKernelComputation from gpjax.kernels.stationary import ( RBF, + Gneiting, Matern12, Matern32, Matern52, @@ -33,10 +34,16 @@ NonNegativeReal, PositiveReal, ) +from hypothesis import ( + given, + strategies as st, +) import jax from jax import config import jax.numpy as jnp +import jax.random as jr import lineax as lx +import paramax from paramax import AbstractUnwrappable import pytest @@ -254,3 +261,133 @@ def test_name_is_not_a_constructor_argument(kernel: type[StationaryKernel]): assert "name" not in [f.name for f in dataclasses.fields(kernel)] with pytest.raises(TypeError): kernel(name="renamed") + + +# --------------------------------------------------------------------------- +# Gneiting space-time kernel. Columns 0 and 1 are space, column 2 is time. +# --------------------------------------------------------------------------- + +SPACE_TIME = jnp.concatenate( + [ + jr.uniform(jr.key(1), (30, 2), minval=-2.0, maxval=2.0), + jr.uniform(jr.key(2), (30, 1), minval=0.0, maxval=5.0), + ], + axis=1, +) + + +def _fixed(value: float): + return paramax.non_trainable(jnp.array(value)) + + +def test_gneiting_with_zero_interaction_is_separable(): + kernel = Gneiting( + space_dims=[0, 1], + time_dim=2, + variance=1.5, + space_lengthscale=0.8, + time_lengthscale=2.0, + alpha=0.7, + gamma=0.6, + beta=_fixed(0.0), + ) + x, y = SPACE_TIME[0], SPACE_TIME[1] + h = jnp.linalg.norm(x[:2] - y[:2]) / 0.8 + psi = (jnp.abs(x[2] - y[2]) / 2.0) ** (2 * 0.7) + 1.0 + space = jnp.exp(-(h ** (2 * 0.6))) + time = psi ** (-2 / 2) + assert jnp.allclose(kernel(x, y), 1.5 * space * time) + + +def test_gneiting_at_zero_lag_is_a_powered_exponential_in_space(): + kernel = Gneiting( + space_dims=[0, 1], time_dim=2, variance=1.3, space_lengthscale=0.9, gamma=0.4 + ) + same_time = SPACE_TIME.at[:, 2].set(1.0) + powered = PoweredExponential( + active_dims=[0, 1], lengthscale=0.9, variance=1.3, power=2 * 0.4 + ) + assert jnp.allclose( + kernel.gram(same_time).as_matrix(), powered.gram(same_time).as_matrix() + ) + + +def test_gneiting_space_correlation_decays_more_slowly_at_longer_lags(): + kernel = Gneiting(space_dims=[0, 1], time_dim=2, beta=0.9) + origin = jnp.array([0.0, 0.0, 0.0]) + + def correlation(distance, lag): + far = jnp.array([distance, 0.0, lag]) + same = jnp.array([0.0, 0.0, lag]) + return kernel(origin, far) / kernel(origin, same) + + assert correlation(1.0, 0.0) < correlation(1.0, 3.0) + + +def test_gneiting_diagonal_is_the_variance(): + kernel = Gneiting(space_dims=[0, 1], time_dim=2, variance=2.5) + diagonal = kernel.diagonal(SPACE_TIME).as_matrix().diagonal() + assert jnp.allclose(diagonal, 2.5) + + +@given( + alpha=st.floats(min_value=0.05, max_value=0.95), + beta=st.floats(min_value=0.05, max_value=0.95), + gamma=st.floats(min_value=0.05, max_value=0.95), + space_lengthscale=st.floats(min_value=0.2, max_value=3.0), + time_lengthscale=st.floats(min_value=0.2, max_value=3.0), +) +def test_gneiting_is_positive_definite( + alpha, beta, gamma, space_lengthscale, time_lengthscale +): + kernel = Gneiting( + space_dims=[0, 1], + time_dim=2, + alpha=alpha, + beta=beta, + gamma=gamma, + space_lengthscale=space_lengthscale, + time_lengthscale=time_lengthscale, + ) + gram = kernel.gram(SPACE_TIME).as_matrix() + assert jnp.allclose(gram, gram.T) + assert float(jnp.linalg.eigvalsh(gram).min()) > -1e-8 + + +def test_gneiting_accepts_fixed_values_at_the_bounds(): + kernel = Gneiting( + space_dims=[0, 1], + time_dim=2, + alpha=_fixed(1.0), + beta=_fixed(1.0), + gamma=_fixed(1.0), + ) + gram = kernel.gram(SPACE_TIME).as_matrix() + assert float(jnp.linalg.eigvalsh(gram).min()) > -1e-8 + + +def test_gneiting_gradients_are_finite_on_the_diagonal(): + kernel = Gneiting(space_dims=[0, 1], time_dim=2, alpha=0.3, gamma=0.3) + grads = jax.grad(lambda k: jnp.sum(k.gram(SPACE_TIME).as_matrix()))(kernel) + assert all(jnp.all(jnp.isfinite(g)) for g in jax.tree_util.tree_leaves(grads)) + + +@pytest.mark.parametrize( + ("space_dims", "time_dim", "message"), + [ + ([0, 1], 1, "both a space column and the time column"), + ([0, 0], 2, "repeated columns"), + ([], 2, "non-empty list"), + ([0, 1], [2], "time_dim"), + ], +) +def test_gneiting_rejects_invalid_columns(space_dims, time_dim, message): + with pytest.raises((ValueError, TypeError), match=message): + Gneiting(space_dims=space_dims, time_dim=time_dim) + + +@pytest.mark.parametrize("name", ["alpha", "beta", "gamma"]) +@pytest.mark.parametrize("value", [0.0, 1.0, 1.5]) +def test_gneiting_rejects_trainable_values_outside_the_open_interval(name, value): + with pytest.raises(ValueError, match="non_trainable"): + Gneiting(space_dims=[0, 1], time_dim=2, **{name: value})