Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 28 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)).

Expand Down
43 changes: 43 additions & 0 deletions GLOSSARY.md
Original file line number Diff line number Diff line change
@@ -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.
33 changes: 33 additions & 0 deletions docs/adr/0001-location-functions.md
Original file line number Diff line number Diff line change
@@ -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.
211 changes: 211 additions & 0 deletions docs/examples/data/_pull_reference_datasets.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.")
Binary file not shown.
Binary file added docs/examples/data/ncep_europe_winter_2019.nc
Binary file not shown.
Loading
Loading