From fda6850c4fefc0f89e5d28689694a861bb72e394 Mon Sep 17 00:00:00 2001 From: Thomas Pinder Date: Sun, 4 Oct 2026 12:27:50 +0000 Subject: [PATCH 1/3] feat(xarray): add input transforms and chunked GridSpec.predict from_xarray takes transforms (Standardise, UnitSphere, Cyclic) that turn the named inputs into the columns of X. The GridSpec keeps them, fitted on the training cells, and applies them to every new grid. GridSpec.predict gives the predictive mean and variance on a grid in chunks of fixed size, compiled once, and keeps a Dask-backed grid lazy. Dask is a dev dependency for the lazy-grid test only. Co-Authored-By: Claude Opus 5.5 --- docs/reference/xarray.md | 19 ++ gpjax/xarray.py | 469 +++++++++++++++++++++++++++++++++++---- pyproject.toml | 2 + tests/test_xarray.py | 243 +++++++++++++++++++- uv.lock | 70 ++++++ 5 files changed, 753 insertions(+), 50 deletions(-) diff --git a/docs/reference/xarray.md b/docs/reference/xarray.md index 243686f4f..f0526ba0d 100644 --- a/docs/reference/xarray.md +++ b/docs/reference/xarray.md @@ -14,3 +14,22 @@ predictions back onto the grid. Requires the optional extra: from_xarray GridSpec ``` + +## Input transforms + +Transforms turn the named inputs into the columns of `X`. Pass them to +{func}`~gpjax.xarray.from_xarray`; the {class}`~gpjax.xarray.GridSpec` applies the +same fitted transforms to every new grid. + +```{eval-rst} +.. currentmodule:: gpjax.xarray + +.. autosummary:: + :toctree: generated/ + :nosignatures: + + Standardise + UnitSphere + Cyclic + InputTransform +``` diff --git a/gpjax/xarray.py b/gpjax/xarray.py index 1b89569cf..e9dd80207 100644 --- a/gpjax/xarray.py +++ b/gpjax/xarray.py @@ -30,12 +30,24 @@ rebuild the grid lives on the :class:`GridSpec`, which never enters ``fit``, ``condition`` or a traced computation. +Input transforms (:class:`Standardise`, :class:`UnitSphere`, :class:`Cyclic`) +turn the named inputs into the columns of ``X``. The spec records them, fitted on +the training data, and applies the same ones to every new grid. For large grids, +:meth:`GridSpec.predict` gives the predictive mean and variance in chunks of +bounded size, and keeps a Dask-backed grid lazy. + Requires the optional ``xarray`` extra: ``pip install "gpjax[xarray]"``. """ -from dataclasses import dataclass +import abc +from dataclasses import ( + dataclass, + field, + replace, +) import beartype.typing as tp +import jax import jax.numpy as jnp from jaxtyping import Float import lineax as lx @@ -55,26 +67,194 @@ "gpjax.xarray requires xarray; install it with pip install 'gpjax[xarray]'" ) from error +Columns = dict[str, np.ndarray] + + +class InputTransform(abc.ABC): + r"""A map from named input columns to the named columns of ``X``. + + A transform receives the columns in order, keyed by name, and returns a new + ordered mapping. :func:`from_xarray` calls :meth:`fit` on the training + columns, and the :class:`GridSpec` then applies the fitted transform to the + training data and to every new grid, so both see the same encoding. + Datetime inputs are already float days since their time origin. + """ + + def fit(self, columns: Columns) -> "InputTransform": + r"""Return this transform fitted to the training ``columns``. + + The default returns the transform unchanged; override it when the + transform learns a state from the training data. + """ + return self + + @abc.abstractmethod + def __call__(self, columns: Columns) -> Columns: + r"""Transform ``columns`` and return the new ordered columns.""" + + +@dataclass(frozen=True) +class Standardise(InputTransform): + r"""Centre and scale inputs to zero mean and unit standard deviation. + + The mean and standard deviation come from the training cells. A new grid + is scaled with the training values, so one lengthscale means the same + distance in training and in prediction. + + Attributes: + names: The columns to standardise. ``None`` standardises every column + present when the transform runs. + loc: The fitted training mean of each column; ``None`` before fitting. + scale: The fitted training standard deviation of each column; ``None`` + before fitting. + """ + + names: tp.Optional[tp.Sequence[str]] = None + loc: tp.Optional[dict[str, float]] = field(default=None, compare=False) + scale: tp.Optional[dict[str, float]] = field(default=None, compare=False) + + def __post_init__(self) -> None: + if isinstance(self.names, str): + raise TypeError( + f"names must be a list of names, e.g. [{self.names!r}], not a string" + ) + if self.names is not None: + object.__setattr__(self, "names", tuple(self.names)) + + def fit(self, columns: Columns) -> "Standardise": + r"""Record the mean and standard deviation of each named column. + + Raises: + ValueError: If a name is not a column, or a column is constant. + """ + names = self.names if self.names is not None else tuple(columns) + _require(columns, names, "Standardise") + loc = {name: float(np.mean(columns[name])) for name in names} + scale = {name: float(np.std(columns[name])) for name in names} + constant = [name for name in names if not scale[name] > 0.0] + if constant: + raise ValueError( + f"cannot standardise {constant}: the column is constant over the " + "training cells" + ) + return replace(self, names=names, loc=loc, scale=scale) + + def __call__(self, columns: Columns) -> Columns: + r"""Apply the training mean and standard deviation.""" + if self.loc is None or self.scale is None: + raise RuntimeError("Standardise must be fitted before it is applied") + _require(columns, tuple(self.loc), "Standardise") + return { + name: (values - self.loc[name]) / self.scale[name] + if name in self.loc + else values + for name, values in columns.items() + } + + +@dataclass(frozen=True) +class UnitSphere(InputTransform): + r"""Map latitude and longitude in degrees to a point on the unit sphere. + + The two columns are replaced, at the position of ``lat``, by + ``{name}_x = cos(lat) cos(lon)``, ``{name}_y = cos(lat) sin(lon)`` and + ``{name}_z = sin(lat)``. A stationary kernel on these columns uses the + chord distance through the sphere, which is valid on the whole globe: it + has no seam at the antimeridian, no singularity at the poles, and it does + not stretch distances at high latitudes. A lengthscale is then in Earth + radii (one radius is about 6371 km). + + Attributes: + lat: Name of the latitude column, in degrees north. + lon: Name of the longitude column, in degrees east. + name: Prefix of the three new columns. + """ + + lat: str = "lat" + lon: str = "lon" + name: str = "sphere" + + def __call__(self, columns: Columns) -> Columns: + r"""Replace ``lat`` and ``lon`` with the three unit-sphere columns. + + Raises: + ValueError: If a name is not a column, a new column name is already + taken, or a latitude is outside [-90, 90] degrees. + """ + _require(columns, (self.lat, self.lon), "UnitSphere") + lat_degrees = columns[self.lat] + if np.any(np.abs(lat_degrees[np.isfinite(lat_degrees)]) > 90.0): + raise ValueError( + f"{self.lat!r} has values outside [-90, 90]; UnitSphere expects " + "latitude in degrees north" + ) + lat = np.deg2rad(lat_degrees) + lon = np.deg2rad(columns[self.lon]) + sphere = { + f"{self.name}_x": np.cos(lat) * np.cos(lon), + f"{self.name}_y": np.cos(lat) * np.sin(lon), + f"{self.name}_z": np.sin(lat), + } + return _substitute(columns, (self.lat, self.lon), sphere) + + +@dataclass(frozen=True) +class Cyclic(InputTransform): + r"""Encode a periodic input as a point on a circle. + + The column is replaced by ``{name}_sin`` and ``{name}_cos`` of + ``2 pi value / period``, so values one period apart get the same encoding. + Datetime inputs are days since their time origin: ``Cyclic("time", 365.25)`` + encodes the seasonal cycle, and ``Cyclic("lon", 360.0)`` removes the seam + in longitude. + + Attributes: + name: The column to encode. + period: The period, in the units of the column. + """ + + name: str + period: float + + def __post_init__(self) -> None: + if not self.period > 0.0: + raise ValueError(f"period must be positive, not {self.period}") + + def __call__(self, columns: Columns) -> Columns: + r"""Replace the column with its sine and cosine. + + Raises: + ValueError: If the name is not a column, or a new column name is + already taken. + """ + _require(columns, (self.name,), "Cyclic") + phase = 2.0 * np.pi * columns[self.name] / self.period + circle = {f"{self.name}_sin": np.sin(phase), f"{self.name}_cos": np.cos(phase)} + return _substitute(columns, (self.name,), circle) + @dataclass(frozen=True, repr=False) class GridSpec: r"""The labelled grid behind a flattened :class:`~gpjax.dataset.Dataset`. Row ``i`` of the flattened data is the ``i``-th kept cell of the grid in C - order over ``dims``. That correspondence is all this object records, and all - :meth:`inputs_for` and :meth:`to_xarray` need. + order over ``dims``. That correspondence, and the encoding of the inputs as + the columns of ``X``, is all this object records, and all + :meth:`inputs_for`, :meth:`predict` and :meth:`to_xarray` need. Attributes: target: Name of the modelled variable. target_attrs: The target's attributes (units, long_name, ...), carried onto predictions. - inputs: Input names, in the column order of ``X``. + inputs: Input names, in the order they are read from the data. dims: The grid's dims, in the target's order. coords: The grid's coordinates, used to rebuild labelled output. time_origins: For each datetime input, the timestamp encoded as day 0. mask: Boolean array over the full grid; ``True`` marks a cell that has a row in the flattened data. n_dropped: Number of grid cells dropped for containing NaN. + transforms: The input transforms, fitted on the training cells, that + turn ``inputs`` into the columns of ``X``. """ target: str @@ -85,18 +265,26 @@ class GridSpec: time_origins: dict[str, np.datetime64] mask: np.ndarray n_dropped: int + transforms: tuple[InputTransform, ...] = () @property def n_kept(self) -> int: r"""Number of grid cells with a row in the flattened data.""" return int(self.mask.sum()) + @property + def columns(self) -> tuple[str, ...]: + r"""Names of the columns of ``X``, after the input transforms.""" + empty = {name: np.empty(0) for name in self.inputs} + return tuple(_apply(self.transforms, empty)) + def __repr__(self) -> str: r"""Summarise the grid without printing its coordinates or mask.""" grid = dict(zip(self.dims, self.mask.shape, strict=True)) + columns = "" if self.columns == self.inputs else f", columns={self.columns}" return ( - f"GridSpec(target={self.target!r}, inputs={self.inputs}, grid={grid}, " - f"kept={self.n_kept}, dropped={self.n_dropped})" + f"GridSpec(target={self.target!r}, inputs={self.inputs}{columns}, " + f"grid={grid}, kept={self.n_kept}, dropped={self.n_dropped})" ) def inputs_for( @@ -107,42 +295,108 @@ def inputs_for( The new grid is the broadcast of this spec's inputs as found in ``obj``; no target is needed. Its dims follow the training grid's order, with any new dims after them. Datetime inputs reuse the training time origins, so - a date maps to the same number here as it did in training. Cells where an - input is NaN get no row, and come back as NaN from :meth:`to_xarray`. + a date maps to the same number here as it did in training, and the + fitted input transforms are applied as they were in training. Cells + where an input is NaN get no row, and come back as NaN from + :meth:`to_xarray`. + + This builds every row at once. For a large grid, :meth:`predict` gives + the mean and variance in chunks instead. Args: obj: Labelled data holding every input named by this spec. Returns: The ``(M, D)`` prediction inputs and the ``GridSpec`` of the new grid, - which carries this spec's target name and attributes. + which carries this spec's target name, attributes and transforms. Raises: ValueError: If an input is missing from ``obj``. TypeError: If an input is non-numeric, or is a datetime now but was not in training (or the reverse). """ - dataset = _as_dataset(obj) - variables = _resolve(dataset, self.inputs) - input_dims = dict.fromkeys(dim for var in variables for dim in var.dims) - dims = tuple( - [dim for dim in self.dims if dim in input_dims] - + [dim for dim in input_dims if dim not in self.dims] - ) + dataset, variables, dims = self._grid_of(obj) grid = xr.broadcast(*variables)[0].transpose(*dims) input_matrix = _input_matrix(variables, grid, dims, self.time_origins) keep = np.isfinite(input_matrix).all(axis=1) - test_spec = GridSpec( - target=self.target, + test_spec = replace( + self, target_attrs=dict(self.target_attrs), - inputs=self.inputs, dims=dims, coords=_grid_coords(dataset, dims), - time_origins=self.time_origins, mask=keep.reshape(grid.shape), n_dropped=int(keep.size - keep.sum()), ) - return jnp.asarray(input_matrix[keep]), test_spec + return jnp.asarray(self._encode(input_matrix[keep])), test_spec + + def predict( + self, + predict_fn: tp.Callable[[Float[Array, "M D"]], tp.Any], + obj: tp.Union[xr.Dataset, xr.DataArray], + *, + chunk_size: int = 4096, + ) -> xr.Dataset: + r"""Predict the mean and variance on a new grid, in chunks. + + The grid is built from ``obj`` as in :meth:`inputs_for`, but the cells + go to ``predict_fn`` at most ``chunk_size`` at a time, so the memory use + does not grow with the size of the grid. The last chunk is padded to + ``chunk_size``, so ``predict_fn`` is compiled once, with ``jax.jit``. + + If ``obj`` holds Dask arrays, the result is lazy: each Dask block is + predicted when it is computed, for example by ``.compute()`` or + ``to_netcdf``. Chunk ``obj`` along the grid's dims to control the size + of a block. + + Args: + predict_fn: Maps ``(M, D)`` inputs to a distribution with ``mean`` + and ``variance`` of shape ``(M,)``, for example + ``lambda x: posterior(x, covariance="diagonal")``. Use the + diagonal covariance: a dense one costs ``chunk_size**2`` memory + and gives the same marginals. + obj: Labelled data holding every input named by this spec. + chunk_size: The number of cells given to ``predict_fn`` at once. + + Returns: + An ``xr.Dataset`` holding ``{target}_mean`` and ``{target}_variance`` + over the grid. Cells where an input is NaN are NaN. + + Raises: + ValueError: If ``chunk_size`` is not positive, an input is missing + from ``obj``, or ``predict_fn`` returns the wrong shape. + TypeError: If an input is non-numeric, or is a datetime now but was + not in training (or the reverse). + """ + if chunk_size < 1: + raise ValueError(f"chunk_size must be positive, not {chunk_size}") + dataset, variables, dims = self._grid_of(obj) + grid_inputs = [var.transpose(*dims) for var in xr.broadcast(*variables)] + moments = jax.jit(lambda inputs: _moments(predict_fn(inputs))) + + def predict_block(*blocks: np.ndarray) -> tuple[np.ndarray, np.ndarray]: + columns = [ + _encode(name, block.ravel(), self.time_origins) + for name, block in zip(self.inputs, blocks, strict=True) + ] + input_matrix = np.stack(columns, axis=1) + keep = np.isfinite(input_matrix).all(axis=1) + mean = np.full(keep.size, np.nan) + variance = np.full(keep.size, np.nan) + if keep.any(): + inputs = self._encode(input_matrix[keep]) + mean[keep], variance[keep] = _in_chunks(moments, inputs, chunk_size) + shape = blocks[0].shape + return mean.reshape(shape), variance.reshape(shape) + + mean, variance = xr.apply_ufunc( + predict_block, + *grid_inputs, + output_core_dims=[[], []], + dask="parallelized", + output_dtypes=[float, float], + ) + prediction = self._moments_dataset(mean, variance) + return prediction.assign_coords(_grid_coords(dataset, dims)) def to_xarray( self, @@ -180,14 +434,42 @@ def to_xarray( ) if num_samples is not None: return self._samples(dist, num_samples, key) + return self._moments_dataset( + self._scatter(dist.mean, {}), self._scatter(dist.variance, {}) + ) + + def _grid_of( + self, obj: tp.Union[xr.Dataset, xr.DataArray] + ) -> tuple[xr.Dataset, list[xr.DataArray], tuple[str, ...]]: + r"""The inputs of this spec in ``obj``, and the dims of their grid. + + The dims follow the training grid's order, with any new dims after them. + """ + dataset = _as_dataset(obj) + variables = _resolve(dataset, self.inputs) + input_dims = dict.fromkeys(dim for var in variables for dim in var.dims) + dims = tuple( + [dim for dim in self.dims if dim in input_dims] + + [dim for dim in input_dims if dim not in self.dims] + ) + return dataset, variables, dims + + def _encode(self, input_matrix: np.ndarray) -> np.ndarray: + r"""Apply the fitted transforms to rows of raw inputs.""" + columns = dict(zip(self.inputs, input_matrix.T, strict=True)) + return np.stack(list(_apply(self.transforms, columns).values()), axis=1) + + def _moments_dataset( + self, mean: xr.DataArray, variance: xr.DataArray + ) -> xr.Dataset: + r"""Name ``mean`` and ``variance`` after the target, with its attributes.""" variance_attrs = dict(self.target_attrs) if "units" in variance_attrs: variance_attrs["units"] = f"({variance_attrs['units']})^2" + mean, variance = mean.copy(deep=False), variance.copy(deep=False) + mean.attrs, variance.attrs = dict(self.target_attrs), variance_attrs return xr.Dataset( - { - f"{self.target}_mean": self._scatter(dist.mean, self.target_attrs), - f"{self.target}_variance": self._scatter(dist.variance, variance_attrs), - } + {f"{self.target}_mean": mean, f"{self.target}_variance": variance} ) def _samples( @@ -239,6 +521,7 @@ def from_xarray( target: str, inputs: tp.Sequence[str], *, + transforms: tp.Sequence[InputTransform] = (), dropna: bool = True, ) -> tuple[Dataset, GridSpec]: r"""Flatten labelled xarray data into a :class:`~gpjax.dataset.Dataset`. @@ -246,14 +529,17 @@ def from_xarray( Every input is broadcast onto the target's grid, so an input on fewer dims (e.g. ``elevation(lat, lon)`` for a ``(time, lat, lon)`` target) repeats along the rest. Datetime inputs become float days since their earliest - timestamp. + timestamp. The ``transforms`` then run in order, fitted on the kept cells, + to give the columns of ``X``; :attr:`GridSpec.columns` names them. Args: obj: The labelled data. A ``DataArray`` must be named, and that name is the target. target: Name of the data variable to model. Its dims define the grid. inputs: Coordinates and/or data variables to use as inputs, in the order - of the columns of ``X``. + of the columns of ``X`` before any transforms. + transforms: Input transforms, such as :class:`Standardise`, + :class:`UnitSphere` or :class:`Cyclic`, applied in order. dropna: Drop grid cells where the target or any input is NaN. When ``False``, such cells raise instead. @@ -264,8 +550,9 @@ def from_xarray( Raises: ValueError: If a name is missing, ``inputs`` is empty, repeats a name or includes the target, an input has a dim the target lacks, a - ``DataArray`` is unnamed, or no cells survive NaN handling (or any - NaN is present when ``dropna=False``). + ``DataArray`` is unnamed, no cells survive NaN handling (or any + NaN is present when ``dropna=False``), or a transform rejects its + columns. TypeError: If ``inputs`` is a single string, or the target or an input is not numeric (inputs may also be ``datetime64``). """ @@ -300,7 +587,16 @@ def from_xarray( if not keep.any(): raise ValueError("no cells are left once NaN cells are dropped") - data = Dataset(X=jnp.asarray(input_matrix[keep]), y=jnp.asarray(outputs[keep])) + columns = dict(zip(inputs, input_matrix[keep].T, strict=True)) + fitted = [] + for transform in transforms: + fitted.append(transform.fit(columns)) + columns = fitted[-1](columns) + + data = Dataset( + X=jnp.asarray(np.stack(list(columns.values()), axis=1)), + y=jnp.asarray(outputs[keep]), + ) spec = GridSpec( target=target, target_attrs=dict(target_values.attrs), @@ -310,10 +606,78 @@ def from_xarray( time_origins=time_origins, mask=keep.reshape(target_values.shape), n_dropped=n_dropped, + transforms=tuple(fitted), ) return data, spec +def _apply(transforms: tp.Sequence[InputTransform], columns: Columns) -> Columns: + r"""Run fitted ``transforms`` over ``columns`` in order.""" + for transform in transforms: + columns = transform(columns) + return columns + + +def _require(columns: Columns, names: tp.Sequence[str], transform: str) -> None: + r"""Raise if a name that ``transform`` reads is not a column.""" + missing = [name for name in names if name not in columns] + if missing: + raise ValueError( + f"{transform} needs columns {missing}, but the columns are {list(columns)}" + ) + + +def _substitute(columns: Columns, replaced: tuple[str, ...], new: Columns) -> Columns: + r"""Swap the ``replaced`` columns for ``new``, at the first one's position.""" + taken = [name for name in new if name in columns and name not in replaced] + if taken: + raise ValueError(f"cannot add columns {taken}: the names are already taken") + out = {} + for name, values in columns.items(): + if name == replaced[0]: + out.update(new) + elif name not in replaced: + out[name] = values + return out + + +def _moments(dist: tp.Any) -> tuple[Float[Array, " M"], Float[Array, " M"]]: + r"""The mean and marginal variance of a predictive distribution.""" + return dist.mean, dist.variance + + +def _in_chunks( + moments: tp.Callable[[Float[Array, "M D"]], tuple[Array, Array]], + inputs: np.ndarray, + chunk_size: int, +) -> tuple[np.ndarray, np.ndarray]: + r"""Evaluate ``moments`` over ``inputs``, ``chunk_size`` rows at a time. + + The last chunk is padded with copies of its first row, so every call has + the same shape and a jitted ``moments`` compiles once. + """ + n_rows = inputs.shape[0] + means, variances = [], [] + for start in range(0, n_rows, chunk_size): + chunk = inputs[start : start + chunk_size] + n_valid = chunk.shape[0] + if n_valid < chunk_size: + padding = np.repeat(chunk[:1], chunk_size - n_valid, axis=0) + chunk = np.concatenate([chunk, padding]) + mean, variance = moments(jnp.asarray(chunk)) + mean = np.asarray(mean).reshape(-1) + variance = np.asarray(variance).reshape(-1) + if mean.shape != (chunk_size,) or variance.shape != (chunk_size,): + raise ValueError( + f"predict_fn returned a mean of shape {mean.shape} and a variance " + f"of shape {variance.shape} for {chunk_size} inputs; it must " + "return one value per input" + ) + means.append(mean[:n_valid]) + variances.append(variance[:n_valid]) + return np.concatenate(means), np.concatenate(variances) + + def _as_dataset(obj: tp.Union[xr.Dataset, xr.DataArray]) -> xr.Dataset: r"""Promote a named ``DataArray`` to a one-variable ``Dataset``.""" if isinstance(obj, xr.Dataset): @@ -360,29 +724,34 @@ def _input_matrix( dims: tuple[str, ...], time_origins: dict[str, np.datetime64], ) -> np.ndarray: - r"""Stack ``variables`` into an ``(N, D)`` float matrix over ``grid``. + r"""Stack ``variables`` into an ``(N, D)`` float matrix over ``grid``.""" + columns = [ + _encode(variable.name, _column(variable, grid, dims), time_origins) + for variable in variables + ] + return np.stack(columns, axis=1) + + +def _encode( + name: tp.Hashable, column: np.ndarray, time_origins: dict[str, np.datetime64] +) -> np.ndarray: + r"""Encode one flat input column as floats. A datetime input must have a time origin, and an input with a time origin must still be a datetime: otherwise the same number would mean different things in training and prediction. """ - columns = [] - for variable in variables: - name = variable.name - column = _column(variable, grid, dims) - is_datetime = np.issubdtype(column.dtype, np.datetime64) - if is_datetime != (name in time_origins): - then, now = ("a", "not a") if name in time_origins else ("not a", "a") - raise TypeError( - f"input {name!r} was {then} datetime when the spec was built but " - f"is {now} datetime now; encode it the same way in both" - ) - if is_datetime: - days = (column - time_origins[name]) / np.timedelta64(1, "D") - columns.append(np.where(np.isnat(column), np.nan, days)) - else: - columns.append(_as_float(column, f"input {name!r}")) - return np.stack(columns, axis=1) + is_datetime = np.issubdtype(column.dtype, np.datetime64) + if is_datetime != (name in time_origins): + then, now = ("a", "not a") if name in time_origins else ("not a", "a") + raise TypeError( + f"input {name!r} was {then} datetime when the spec was built but " + f"is {now} datetime now; encode it the same way in both" + ) + if is_datetime: + days = (column - time_origins[name]) / np.timedelta64(1, "D") + return np.where(np.isnat(column), np.nan, days) + return _as_float(column, f"input {name!r}") def _column( @@ -413,6 +782,10 @@ def _as_float(values: np.ndarray, description: str) -> np.ndarray: __all__ = [ + "Cyclic", "GridSpec", + "InputTransform", + "Standardise", + "UnitSphere", "from_xarray", ] diff --git a/pyproject.toml b/pyproject.toml index 588106caa..aec26f042 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -126,6 +126,8 @@ dev = [ # The `xarray` extra, repeated here because CI syncs without extras and the # `gpjax.xarray` tests must still run. "xarray>=2024.1", + # Lazy, Dask-backed grids in `GridSpec.predict`; tests only. + "dask>=2024.1", # Workflow security audit (`poe lint-actions`, .github/workflows/zizmor.yml). # Pinned through uv.lock rather than `uvx zizmor@latest` so a new release # cannot turn every open PR red without a commit. diff --git a/tests/test_xarray.py b/tests/test_xarray.py index ee94426d0..9cf7b1b81 100644 --- a/tests/test_xarray.py +++ b/tests/test_xarray.py @@ -19,7 +19,12 @@ from gpjax.dataset import Dataset from gpjax.distributions import GaussianDistribution -from gpjax.xarray import from_xarray +from gpjax.xarray import ( + Cyclic, + Standardise, + UnitSphere, + from_xarray, +) import jax import jax.numpy as jnp import jax.random as jr @@ -327,12 +332,20 @@ def gridded(lats, lons) -> xr.Dataset: verbose=False, ) + posterior = model.condition(data) test_inputs, test_spec = spec.inputs_for(fine.drop_vars("t2m")) - out = test_spec.to_xarray(model.condition(data)(test_inputs)) + out = test_spec.to_xarray(posterior(test_inputs)) assert out["t2m_mean"].sizes == {"lat": 11, "lon": 13} np.testing.assert_allclose(out["t2m_mean"], fine["t2m"], atol=0.1) + chunked = spec.predict( + lambda x: posterior(x, covariance="diagonal"), + fine.drop_vars("t2m"), + chunk_size=50, + ) + xr.testing.assert_allclose(chunked, out) + def test_samples_from_a_tagged_diagonal_covariance_are_refused(two_cell_field): _, spec = from_xarray(two_cell_field, target="t2m", inputs=["site"]) @@ -401,3 +414,229 @@ def test_inputs_of_any_numeric_dtype_become_floats(dtype): np.testing.assert_array_equal(data.X[:, 0], [0.0, 1.0]) assert data.X.dtype == jnp.float64 + + +# --- Input transforms ------------------------------------------------------- + + +def test_standardise_gives_columns_with_zero_mean_and_unit_std(field): + data, spec = from_xarray( + field, target="t2m", inputs=["lat", "elevation"], transforms=[Standardise()] + ) + + np.testing.assert_allclose(data.X.mean(axis=0), 0.0, atol=1e-12) + np.testing.assert_allclose(data.X.std(axis=0), 1.0) + (standardise,) = spec.transforms + assert standardise.loc == {"lat": 15.0, "elevation": 350.0} + + +def test_standardise_scales_a_new_grid_with_the_training_statistics(field): + data, spec = from_xarray( + field, target="t2m", inputs=["elevation"], transforms=[Standardise()] + ) + (standardise,) = spec.transforms + grid = xr.Dataset( + {"elevation": ("site", [350.0, 350.0 + standardise.scale["elevation"]])} + ) + + test_inputs, _ = spec.inputs_for(grid) + + np.testing.assert_allclose(test_inputs[:, 0], [0.0, 1.0]) + # elevation(lat, lon) alone spans 6 cells: the first time step of training. + training_inputs, _ = spec.inputs_for(field.drop_vars("t2m")) + np.testing.assert_allclose(training_inputs, data.X[:6]) + + +def test_standardise_changes_only_the_named_columns(field): + data, _ = from_xarray( + field, + target="t2m", + inputs=["lat", "elevation"], + transforms=[Standardise(["elevation"])], + ) + + np.testing.assert_array_equal(np.unique(data.X[:, 0]), [10.0, 20.0]) + np.testing.assert_allclose(data.X[:, 1].std(), 1.0) + + +def test_standardise_rejects_a_constant_input(field): + field["flat"] = (("lat", "lon"), np.ones((2, 3))) + + with pytest.raises(ValueError, match=r"cannot standardise \['flat'\]"): + from_xarray(field, target="t2m", inputs=["flat"], transforms=[Standardise()]) + + +def test_standardise_rejects_a_bare_string(): + with pytest.raises(TypeError, match=r"\['lat'\]"): + Standardise("lat") + + +def test_unit_sphere_replaces_lat_and_lon_with_unit_vectors(): + lons = np.array([0.0, 90.0, 360.0]) + field = xr.Dataset( + {"t2m": (("lat", "lon"), np.zeros((2, 3)))}, + coords={"lat": [0.0, 90.0], "lon": lons, "level": 1.0}, + ) + field["height"] = (("lat", "lon"), np.arange(6.0).reshape(2, 3)) + + data, spec = from_xarray( + field, target="t2m", inputs=["height", "lat", "lon"], transforms=[UnitSphere()] + ) + + assert spec.columns == ("height", "sphere_x", "sphere_y", "sphere_z") + xyz = np.asarray(data.X[:, 1:]) + np.testing.assert_allclose(np.linalg.norm(xyz, axis=1), 1.0) + # On the equator: lon 0 -> +x, lon 90 -> +y, and lon 360 is lon 0 again. + np.testing.assert_allclose(xyz[:3], [[1, 0, 0], [0, 1, 0], [1, 0, 0]], atol=1e-12) + # At the pole every longitude is the same point. + np.testing.assert_allclose(xyz[3:], [[0, 0, 1]] * 3, atol=1e-12) + + +def test_unit_sphere_rejects_latitude_outside_the_degree_range(field): + field = field.assign_coords(lat=[10.0, 100.0]) + + with pytest.raises(ValueError, match=r"outside \[-90, 90\]"): + from_xarray( + field, target="t2m", inputs=["lat", "lon"], transforms=[UnitSphere()] + ) + + +def test_cyclic_maps_values_one_period_apart_to_the_same_point(field): + data, spec = from_xarray( + field, target="t2m", inputs=["lon"], transforms=[Cyclic("lon", period=2.0)] + ) + + assert spec.columns == ("lon_sin", "lon_cos") + # lon is 0, 1, 2 along each row: 0 and 2 are one period apart. + np.testing.assert_allclose(data.X[0], data.X[2], atol=1e-12) + np.testing.assert_allclose(data.X[1], [0.0, -1.0], atol=1e-12) + + +def test_cyclic_on_a_datetime_input_uses_days_since_the_origin(field): + data, _ = from_xarray( + field, target="t2m", inputs=["time"], transforms=[Cyclic("time", period=8.0)] + ) + + # The second timestamp is 2 days, a quarter period, after the first. + np.testing.assert_allclose(data.X[0], [0.0, 1.0], atol=1e-12) + np.testing.assert_allclose(data.X[-1], [1.0, 0.0], atol=1e-12) + + +@pytest.mark.parametrize("period", [0.0, -1.0]) +def test_cyclic_rejects_a_non_positive_period(period): + with pytest.raises(ValueError, match="period must be positive"): + Cyclic("time", period=period) + + +def test_a_transform_reports_a_missing_column(field): + with pytest.raises(ValueError, match=r"UnitSphere needs columns \['lon'\]"): + from_xarray(field, target="t2m", inputs=["lat"], transforms=[UnitSphere()]) + + +def test_a_transform_may_not_overwrite_another_column(field): + field["lon_sin"] = (("lat", "lon"), np.zeros((2, 3))) + + with pytest.raises(ValueError, match=r"\['lon_sin'\].*already taken"): + from_xarray( + field, + target="t2m", + inputs=["lon", "lon_sin"], + transforms=[Cyclic("lon", period=360.0)], + ) + + +def test_transforms_run_in_order_and_the_spec_names_the_columns(field): + data, spec = from_xarray( + field, + target="t2m", + inputs=["lat", "lon", "elevation"], + transforms=[UnitSphere(), Standardise(["elevation"])], + ) + + assert data.X.shape == (12, 4) + assert spec.columns == ("sphere_x", "sphere_y", "sphere_z", "elevation") + assert "columns=('sphere_x', 'sphere_y', 'sphere_z', 'elevation')" in repr(spec) + + +# --- Chunked prediction ------------------------------------------------------ + + +def _toy_predict_fn(inputs): + """A deterministic stand-in for a posterior: one mean and variance per row.""" + return _diagonal_gaussian(inputs.sum(axis=1), inputs[:, 0] ** 2 + 1.0) + + +@pytest.fixture +def gappy_grid(field) -> xr.Dataset: + """The field's inputs, with one NaN covariate cell and lat attrs.""" + grid = field.drop_vars("t2m") + grid["elevation"][1, 2] = np.nan + grid["lat"].attrs = {"units": "degrees_north"} + return grid + + +@pytest.mark.parametrize("chunk_size", [1, 5, 12, 100]) +def test_predict_matches_inputs_for_then_to_xarray(field, gappy_grid, chunk_size): + _, spec = from_xarray( + field, + target="t2m", + inputs=["time", "lat", "elevation"], + transforms=[Standardise()], + ) + test_inputs, test_spec = spec.inputs_for(gappy_grid) + expected = test_spec.to_xarray(_toy_predict_fn(test_inputs)) + + out = spec.predict(_toy_predict_fn, gappy_grid, chunk_size=chunk_size) + + xr.testing.assert_allclose(out, expected) + assert np.isnan(out["t2m_mean"][:, 1, 2]).all() + assert out["t2m_mean"].attrs == {"units": "K"} + assert out["t2m_variance"].attrs == {"units": "(K)^2"} + + +def test_predict_compiles_once_for_every_chunk(field): + _, spec = from_xarray(field, target="t2m", inputs=["lat", "lon"]) + traces = [] + + def predict_fn(inputs): + traces.append(inputs.shape) + return _toy_predict_fn(inputs) + + spec.predict(predict_fn, field.drop_vars("t2m"), chunk_size=5) + + assert traces == [(5, 2)] + + +def test_predict_keeps_a_dask_grid_lazy(field, gappy_grid): + pytest.importorskip("dask") + _, spec = from_xarray(field, target="t2m", inputs=["time", "lat", "elevation"]) + calls = [] + + def predict_fn(inputs): + calls.append(inputs.shape) + return _toy_predict_fn(inputs) + + lazy = spec.predict(predict_fn, gappy_grid.chunk({"lat": 1}), chunk_size=4) + + assert lazy["t2m_mean"].chunks is not None + assert calls == [] + eager = spec.predict(_toy_predict_fn, gappy_grid, chunk_size=4) + xr.testing.assert_allclose(lazy.compute(), eager) + + +@pytest.mark.parametrize("chunk_size", [0, -3]) +def test_predict_rejects_a_non_positive_chunk_size(field, chunk_size): + _, spec = from_xarray(field, target="t2m", inputs=["lat"]) + + with pytest.raises(ValueError, match="chunk_size must be positive"): + spec.predict(_toy_predict_fn, field, chunk_size=chunk_size) + + +def test_predict_rejects_a_predict_fn_of_the_wrong_shape(field): + _, spec = from_xarray(field, target="t2m", inputs=["lat"]) + + def one_value(inputs): + return _diagonal_gaussian(inputs.sum(keepdims=True)[0], jnp.ones(1)) + + with pytest.raises(ValueError, match="one value per input"): + spec.predict(one_value, field, chunk_size=4) diff --git a/uv.lock b/uv.lock index f62a5a867..dd43f1fdd 100644 --- a/uv.lock +++ b/uv.lock @@ -374,6 +374,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/98/78/01c019cdb5d6498122777c1a43056ebb3ebfeef2076d9d026bfe15583b2b/click-8.3.1-py3-none-any.whl", hash = "sha256:981153a64e25f12d547d3426c367a4857371575ee7ad18df2a6183ab0545b2a6", size = 108274, upload-time = "2025-11-15T20:45:41.139Z" }, ] +[[package]] +name = "cloudpickle" +version = "3.1.2" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/27/fb/576f067976d320f5f0114a8d9fa1215425441bb35627b1993e5afd8111e5/cloudpickle-3.1.2.tar.gz", hash = "sha256:7fda9eb655c9c230dab534f1983763de5835249750e85fbcef43aaa30a9a2414", size = 22330, upload-time = "2025-11-03T09:25:26.604Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/88/39/799be3f2f0f38cc727ee3b4f1445fe6d5e4133064ec2e4115069418a5bb6/cloudpickle-3.1.2-py3-none-any.whl", hash = "sha256:9acb47f6afd73f60dc1df93bb801b472f05ff42fa6c84167d25cb206be1fbf4a", size = 22228, upload-time = "2025-11-03T09:25:25.534Z" }, +] + [[package]] name = "codespell" version = "2.4.3" @@ -653,6 +662,25 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/e7/05/c19819d5e3d95294a6f5947fb9b9629efb316b96de511b418c53d245aae6/cycler-0.12.1-py3-none-any.whl", hash = "sha256:85cef7cff222d8644161529808465972e51340599459b8ac3ccbac5a854e0d30", size = 8321, upload-time = "2023-10-07T05:32:16.783Z" }, ] +[[package]] +name = "dask" +version = "2026.8.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "click" }, + { name = "cloudpickle" }, + { name = "fsspec" }, + { name = "importlib-metadata", marker = "python_full_version < '3.12'" }, + { name = "packaging" }, + { name = "partd" }, + { name = "pyyaml" }, + { name = "toolz" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/33/a7/6b3c7ac32b642fbbe0821111654e0bd8cfbe88f68560bcf23cc78ab35c71/dask-2026.8.0.tar.gz", hash = "sha256:8a94c37b5de6d869343340dc26c3c3acca7ec48a3abdabe00ea3abb1125884d5", size = 11561752, upload-time = "2026-08-24T19:21:25.906Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/f8/3a/4fc99e788bcfa1b3b3f21abf57da45898d807d007e7f6fd1c7300904eb70/dask-2026.8.0-py3-none-any.whl", hash = "sha256:ccc0c83a189b0398602435189771d28dad7b5773b6089bb8dce14ae732dd782c", size = 1492182, upload-time = "2026-08-24T19:21:23.997Z" }, +] + [[package]] name = "debugpy" version = "1.8.19" @@ -797,6 +825,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/c7/4e/ce75a57ff3aebf6fc1f4e9d508b8e5810618a33d900ad6c19eb30b290b97/fonttools-4.61.1-py3-none-any.whl", hash = "sha256:17d2bf5d541add43822bcf0c43d7d847b160c9bb01d15d5007d84e2217aaa371", size = 1148996, upload-time = "2025-12-12T17:31:21.03Z" }, ] +[[package]] +name = "fsspec" +version = "2026.9.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/77/cd/9be253869fc42e764de7f3dedd6969af7d44ff9c3375214a3442a6f3fc08/fsspec-2026.9.0.tar.gz", hash = "sha256:0f08147951c8cb31d844c3547d631053b127863b60be04cf06e121333ee0e2fe", size = 333545, upload-time = "2026-09-18T17:50:42.825Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/6c/c0/a98505f18594f1bce828bb159cec0fcf9860562f1a2c85913409fc8f3d9e/fsspec-2026.9.0-py3-none-any.whl", hash = "sha256:8dd6e646e99ea382bd85f97a45e6b526a442d79423a7dc673f1e2756d05fcb5f", size = 221738, upload-time = "2026-09-18T17:50:41.341Z" }, +] + [[package]] name = "gpjax" source = { editable = "." } @@ -858,6 +895,7 @@ dev = [ { name = "asv" }, { name = "codespell" }, { name = "coverage" }, + { name = "dask" }, { name = "hypothesis" }, { name = "interrogate" }, { name = "jupytext" }, @@ -929,6 +967,7 @@ dev = [ { name = "asv", specifier = ">=0.6" }, { name = "codespell", specifier = ">=2.2.4" }, { name = "coverage", specifier = ">=7.2.2" }, + { name = "dask", specifier = ">=2024.1" }, { name = "hypothesis", specifier = ">=6.148.2" }, { name = "interrogate", specifier = ">=1.5.0" }, { name = "jupytext" }, @@ -1734,6 +1773,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/8d/f8/0d32243c6b8e5dbee9097cf0c95bbdf8681ba4463c927c2e3445f3775814/lineax-0.1.1-py3-none-any.whl", hash = "sha256:2e399f1674773ab2ba54d76175a618977a554f47abd0a345198d53d92c07beb2", size = 77567, upload-time = "2026-05-01T15:59:05.517Z" }, ] +[[package]] +name = "locket" +version = "1.0.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/2f/83/97b29fe05cb6ae28d2dbd30b81e2e402a3eed5f460c26e9eaa5895ceacf5/locket-1.0.0.tar.gz", hash = "sha256:5c0d4c052a8bbbf750e056a8e65ccd309086f4f0f18a2eac306a8dfa4112a632", size = 4350, upload-time = "2022-04-20T22:04:44.312Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/db/bc/83e112abc66cd466c6b83f99118035867cecd41802f8d044638aa78a106e/locket-1.0.0-py2.py3-none-any.whl", hash = "sha256:b6c819a722f7b6bd955b80781788e4a66a55628b858d347536b7e81325a3a5e3", size = 4398, upload-time = "2022-04-20T22:04:42.23Z" }, +] + [[package]] name = "markdown-it-py" version = "4.0.0" @@ -2357,6 +2405,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/16/32/f8e3c85d1d5250232a5d3477a2a28cc291968ff175caeadaf3cc19ce0e4a/parso-0.8.5-py2.py3-none-any.whl", hash = "sha256:646204b5ee239c396d040b90f9e272e9a8017c630092bf59980beb62fd033887", size = 106668, upload-time = "2025-08-23T15:15:25.663Z" }, ] +[[package]] +name = "partd" +version = "1.4.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "locket" }, + { name = "toolz" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/b2/3a/3f06f34820a31257ddcabdfafc2672c5816be79c7e353b02c1f318daa7d4/partd-1.4.2.tar.gz", hash = "sha256:d022c33afbdc8405c226621b015e8067888173d85f7f5ecebb3cafed9a20f02c", size = 21029, upload-time = "2024-05-06T19:51:41.945Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/71/e7/40fb618334dcdf7c5a316c0e7343c5cd82d3d866edc100d98e29bc945ecd/partd-1.4.2-py3-none-any.whl", hash = "sha256:978e4ac767ec4ba5b86c6eaa52e5a2a3bc748a2ca839e8cc798f1cc6ce6efb0f", size = 18905, upload-time = "2024-05-06T19:51:39.271Z" }, +] + [[package]] name = "pastel" version = "0.2.1" @@ -3774,6 +3835,15 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/23/d1/136eb2cb77520a31e1f64cbae9d33ec6df0d78bdf4160398e86eec8a8754/tomli-2.4.0-py3-none-any.whl", hash = "sha256:1f776e7d669ebceb01dee46484485f43a4048746235e683bcdffacdf1fb4785a", size = 14477, upload-time = "2026-01-11T11:22:37.446Z" }, ] +[[package]] +name = "toolz" +version = "1.1.0" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/11/d6/114b492226588d6ff54579d95847662fc69196bdeec318eb45393b24c192/toolz-1.1.0.tar.gz", hash = "sha256:27a5c770d068c110d9ed9323f24f1543e83b2f300a687b7891c1a6d56b697b5b", size = 52613, upload-time = "2025-10-17T04:03:21.661Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/fb/12/5911ae3eeec47800503a238d971e51722ccea5feb8569b735184d5fcdbc0/toolz-1.1.0-py3-none-any.whl", hash = "sha256:15ccc861ac51c53696de0a5d6d4607f99c210739caf987b5d2054f3efed429d8", size = 58093, upload-time = "2025-10-17T04:03:20.435Z" }, +] + [[package]] name = "tornado" version = "6.5.4" From 247c9a02f1674c2261227b9382f9910a2b770d64 Mon Sep 17 00:00:00 2001 From: Thomas Pinder Date: Sun, 4 Oct 2026 12:27:50 +0000 Subject: [PATCH 2/3] docs(xarray): infill a real reanalysis field in the xarray example The example now reads a netCDF file of the 2024 NCEP-NCAR Reanalysis 1 temperature anomaly, which is complete, removes cells, and scores the GP infill against the truth. Random removal gives a global mean whose interval holds the truth; value-dependent removal shows the bias of a GP when data are missing not at random. The pull script records how the 44 KB file was made, so docs builds stay offline. Co-Authored-By: Claude Opus 5.5 --- CHANGELOG.md | 21 + .../examples/data/_pull_reference_datasets.py | 90 +++- docs/examples/data/ncep_air_anomaly_2024.nc | Bin 0 -> 44008 bytes docs/examples/xarray_workflow.py | 453 ++++++++++++------ 4 files changed, 408 insertions(+), 156 deletions(-) create mode 100644 docs/examples/data/ncep_air_anomaly_2024.nc diff --git a/CHANGELOG.md b/CHANGELOG.md index 911fd3019..d6dbee5ad 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,27 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Added + +- **`gpjax.xarray` input transforms.** `from_xarray(..., transforms=[...])` + turns the named inputs into the columns of `X`, and the `GridSpec` applies the + same fitted transforms to every new grid. `Standardise` scales inputs with the + training mean and standard deviation. `UnitSphere` maps latitude and longitude + to a point on the unit sphere, so a stationary kernel is valid on the whole + globe. `Cyclic` encodes a periodic input, such as the seasonal cycle, as a + point on a circle. `GridSpec.columns` names the resulting columns. +- **`GridSpec.predict`: chunked prediction on large grids.** + `spec.predict(lambda x: posterior(x, covariance="diagonal"), grid)` gives the + predictive mean and variance as an `xr.Dataset`, with at most `chunk_size` + cells in memory at once and one compilation. A Dask-backed grid gives a lazy + result that is predicted block by block. +- **The *Gridded Data with xarray* example now uses real data.** It reads a + netCDF file of the 2024 temperature anomaly from the NCEP-NCAR Reanalysis 1, + which is complete, and removes cells to test the infill against the truth. + Cells removed at random are filled well, and joint samples give a global mean + whose interval holds the truth. Cells removed because they are warm show how a + GP is biased, and overconfident, when data are missing not at random. + ### Changed - Allow JAX and JAXlib 0.11 in downstream environments by removing the diff --git a/docs/examples/data/_pull_reference_datasets.py b/docs/examples/data/_pull_reference_datasets.py index 6811045b3..efff6cc04 100644 --- a/docs/examples/data/_pull_reference_datasets.py +++ b/docs/examples/data/_pull_reference_datasets.py @@ -10,8 +10,9 @@ uv run --extra docs python docs/examples/data/_pull_reference_datasets.py The ``--extra docs`` is only needed for the UCI Auto MPG pull, which uses -``ucimlrepo``; the other three pulls need nothing beyond ``pandas`` and -``requests``. +``ucimlrepo``. The NCEP reanalysis pull reads netCDF4, so it also needs +``--with h5netcdf --with h5py``. The other pulls need nothing beyond ``pandas`` +and ``requests``. Data sources ------------ @@ -32,11 +33,23 @@ https://archive.ics.uci.edu/dataset/9/auto-mpg (fetched via ``ucimlrepo``). Licence: CC BY 4.0. The features and the target are concatenated into a single frame so the notebook can split them back out without ``ucimlrepo``. +- NCEP-NCAR Reanalysis 1 temperature anomaly (``ncep_air_anomaly_2024.nc``): + the 2024 annual-mean anomaly of near-surface (0.995 sigma) air temperature, + relative to the 1991-2020 mean of each grid cell, on the native 2.5-degree + grid. A reanalysis has a value in every cell, so the field is complete. + Source: https://psl.noaa.gov/data/gridded/data.ncep.reanalysis.html (monthly + means, ``air.mon.mean.nc``). Provider: NOAA Physical Sciences Laboratory + (Kalnay et al., 1996, https://doi.org/10.1175/1520-0477(1996)077<0437:TNYRP>2.0.CO;2). + Work of the US Government, so public domain; PSL ask for the acknowledgement + "NCEP-NCAR Reanalysis 1 data provided by the NOAA PSL, Boulder, Colorado, USA, + from their website at https://psl.noaa.gov". Written as netCDF3, which xarray + reads with SciPy, so the notebook needs no netCDF4 library. """ from __future__ import annotations from pathlib import Path +import tempfile import time import pandas as pd @@ -104,8 +117,81 @@ def pull_auto_mpg() -> None: _save(pd.concat([features, targets], axis=1), "auto_mpg.csv") +# --------------------------------------------------------------------------- # +# xarray_workflow — NCEP-NCAR Reanalysis 1 temperature anomaly for 2024. # +# --------------------------------------------------------------------------- # +NCEP_URL = ( + "https://psl.noaa.gov/thredds/fileServer/Datasets/ncep.reanalysis/" + "Monthlies/surface/air.mon.mean.nc" +) + + +def _download_resumable(url: str, path: Path, attempts: int = 8) -> None: + """Download ``url`` to ``path``, resuming when the server cuts it short.""" + for _ in range(attempts): + start = path.stat().st_size if path.exists() else 0 + headers = {"User-Agent": "gpjax-docs", "Range": f"bytes={start}-"} + try: + with requests.get(url, headers=headers, stream=True, timeout=120) as resp: + if resp.status_code == 416: # nothing left to fetch + return + resp.raise_for_status() + total = start + int(resp.headers["Content-Length"]) + mode = "ab" if resp.status_code == 206 else "wb" + with path.open(mode) as file: + for chunk in resp.iter_content(1 << 20): + file.write(chunk) + if path.stat().st_size >= total: + return + except requests.RequestException as err: + print(f" retrying after: {err}") + time.sleep(2) + raise RuntimeError(f"Failed to fetch {url} in {attempts} attempts") + + +def pull_ncep_anomaly(nc_path: Path | None = None) -> None: + """Save the 2024 NCEP anomaly; pass ``nc_path`` to reuse a download.""" + print("xarray_workflow: NCEP-NCAR Reanalysis 1 temperature anomaly, 2024") + import xarray as xr + + with tempfile.TemporaryDirectory() as tmp: + if nc_path is None: + nc_path = Path(tmp) / "air.mon.mean.nc" + _download_resumable(NCEP_URL, nc_path) + with xr.open_dataset(nc_path, engine="h5netcdf") as monthly: + annual = monthly["air"].resample(time="YS").mean().load() + climatology = annual.sel(time=slice("1991", "2020")).mean("time") + anomaly = annual.sel(time="2024").squeeze("time", drop=True) - climatology + # Longitude from 0..357.5 to -180..177.5, so maps are centred on Greenwich. + anomaly = anomaly.assign_coords(lon=(anomaly["lon"] + 180.0) % 360.0 - 180.0) + anomaly = anomaly.sortby(["lat", "lon"]).astype("float32") + anomaly.attrs = { + "long_name": "Near-surface air temperature anomaly", + "units": "K", + "cell_methods": "time: mean (2024, anomaly relative to 1991-2020)", + } + anomaly["lat"].attrs = {"standard_name": "latitude", "units": "degrees_north"} + anomaly["lon"].attrs = {"standard_name": "longitude", "units": "degrees_east"} + dataset = anomaly.to_dataset(name="tas_anomaly") + dataset.attrs = { + "title": "2024 annual-mean near-surface air temperature anomaly", + "source": "NCEP-NCAR Reanalysis 1, monthly air.sig995 (air.mon.mean.nc)", + "references": "Kalnay et al. (1996), doi:10.1175/1520-0477(1996)077<0437:TNYRP>2.0.CO;2", + "acknowledgement": ( + "NCEP-NCAR Reanalysis 1 data provided by the NOAA PSL, Boulder, " + "Colorado, USA, from their website at https://psl.noaa.gov" + ), + "history": "Annual means of the monthly means; 2024 minus the 1991-2020 mean.", + "Conventions": "CF-1.8", + } + path = HERE / "ncep_air_anomaly_2024.nc" + dataset.to_netcdf(path, engine="scipy", format="NETCDF3_64BIT") + print(f" wrote {path.name}: {dict(dataset.sizes)}") + + if __name__ == "__main__": pull_mauna_loa_co2() pull_gulf_velocities() pull_auto_mpg() + pull_ncep_anomaly() print("\nDone.") diff --git a/docs/examples/data/ncep_air_anomaly_2024.nc b/docs/examples/data/ncep_air_anomaly_2024.nc new file mode 100644 index 0000000000000000000000000000000000000000..01ffe6d465bc2f7f29a61f3292f378397528071c GIT binary patch literal 44008 zcmeFYcXSlT^Do+f5CQ}&S>z#$BvwQr?NqfSEJR+CO*A8eK!AWHgN>od!lV^ZZ(swS5?P0jmydY z^T|iex9eQV3B$kLY5wo^;6dN6-9*7efke57_Zyy^NEh_^_I$&7_IeGJgh7KwBqY}v zn3ym~8I+hXw9c>*L;EE3N>md14ONCG4s<0BO&C66DE&WZ@W6!RG5`LY1r8fLVrZ}b zboQ+;N1NF1oORm7Mz>emQ&05Vuztgo@Y>42!Gne;C6D>mTiCFEeVa6CsDysIru$*v z-WfKiR}K1ozM+YI5{D)Z>XkU`-xqxAGrna)@}PtGCV0!Y11}3T5)!0Ra=P} zJR-Sw;?UYk?BL|VLlb%ruB~+H5M5j8Gj#C4Z;uk|MkW3pOTPy*p1*zhT1%5B-;;-^LOh{qIsJ-`Fxt8QkaJcm5au{=FHc{JUfW z`wbc~?B8#RCE<1Gp8ek+g;6K;$sap-&`6qNzrln4W#+fO17qXrgoj1^=X?5_>l-Ei zM@Dkc^*8z^C*STD7&bg%Q167Hy?YKy82CTQOI~`uZ@=LqdjF4mxkn7@H~c@#^{+ey z_fG6PG%;~l&%}gb-x!>)XI#JJzs40n*Z;o%zm2I%oButL{qJMV(emH!S#Yn!P^zEJhgYEx)9tHl3=Ys|h9iH@`{{9zy&r#(+k8M11xB+8_!v#BBm)hFY)}*#N zwN0n`%r4t*v1Nc5J-PxOT7A<-S8 zn?zTL&J&#>I!<($ND}QM+C{XDXd}@&BAsX{(L$nBqM1a~h$a&KOynZ!L)4w9JyARn z5YfC2(OeGE_z%&T579Ue(bx{scn;DS4$^xM(z^~){|BkhgVfJK`t3n_-lS(tx^L3; zw+~l{bRxPRP1n)%Tr@o&O}~q#`_a@-H1!uv{YKL>G4!4odS48^H-?^%p>f2}cw%T= zG4$IQ8fOfRH-^R?L;b|iJYr}*F*L6j>NAGs8AJ1pp?Swp|1rde7~)3^@g;`d6+?WA zA%4XW-(u*!F~r9h;%5x;HHO9zLwt@Qe#a2sW4?`t)**)0BZk%`mc|wvAXsc3YH8hK zX`HbIsipOdrFD&^F~^2bOY0p=>mEzvkF7*)RcfnKOLK{>O)bewEXfVNd#wLo`oBE* zKb8UaM27RKnz&$(iD~6cT;gQ7?7D_4?l4>xC2&n!1=ocuxV{j_^ou%foTuYv)NsoN z6St`hJzE6sXuxsjDg}4F)p5^wff+dk?h9kM|CGSYQWknsbu?;NXzthXP(u?BZ{_$$ zPl3mxG(5p|Jk>?RGYd>Sx6#D&H3eSGHu1_~4X^LA@K${t@Ax!)kZj`Le{y{GRmWFz zbbK#N{4`kO?192_4R`avrh?}!!+8E|ix>Qx^TLNMUTlNKOQ!jFX_wB+IZdtv`gtX< zpI7bc*Tfbx+5$ezS89NZ}@n_uP)yBjl!E;ws@4}Tpi=&7{WQ9U z)rr(+N5&&wS-epT8?Qgf&h2wJuQSEuHSHF!cE{wE2e^3!uZNcnv3N*(jTf6JxNRll z!5iE>DAnQtqfE|{J(xXQLw`LDv&xxx+h^gG3mnhwCZ1jvXjauQ!_aa23Jo{T;<&bs zhRZK8OiR{qUMCZ0?-2OgWDTd~HgPgzIAOEE@#_Wt+C$(keFXlz#Kd1ZFdX}viQ}?3 zit9Q~oMYkSfdZ$Vw=m^}f-_vyPYw&;HK8@Asq=u;2p-fml;L( zi!@H>MNTnZv|tD?TFT-@ZYS{~6G!vHeh)9)!Om^RJ98WHs?ZG13)0#YtYYT{`h@b} zG@a)^>Eii3DLilgEFKi>;enBM9?&R==QtaLUn;pVJFgpGPGb1jWygC_e!P)x$BV@k zJpI^-M>$8s&vCz7!976=?htPD%yr|o1r~1eo49?o8+T21W5zrWdbcn{oM_D_-q!cyi?s^rmJdIlH+jw+9v*PV!}FZ8c)sskJb1mE7kudBgdo12%v&3Pqm@v2|_yvA1_4-2*PdL#plvN(?_(&JZ_VX$IwL7rwWgl?&9@kD7@AOH?JC}^YV>3FImyz zwz?$y9sE4!0}p-}!0_2_hIe%nukO+CTr9(*Q#i^9;zNp#TMu$vUyS3*-xORl)xx93=2&+F|-B8ri%ss)W*aKVFFXroVe<~hC6cE@!%>K`dYg1mS)Gt z&pdd3@o~IT%ZWE;T6lY*!29j}n03^`XTJ!1MPq!QD)7@P8)qqQp7R9f0gID(pgNl8 zS!&}!q@Q`+CeNn`l9M!^KM(1DkIp>*{^LBjCec!z2fwp;fue#JsCShIU&`jeK|Y>u zE9ufclBZn6vo%-obGi@T7V_Z}&5d`uy79_GC!THZ!artc=sm3C-o_*^dtJC?pn@A$ zn3$g7!So^w)Bm*L#tsT@eo3;`&4#-hlf3P5qq)_C$JU#8_OlBwucW!gdGNtP3!lC) z@%0NEejMZD>^qa^>Y2g=r_tJNqj6q!bK5k=i}e>gB;C%-)%5X-4+XEb-_2{SP`JIL z;0-Pkc09MZW4_?H+~m;@bsjrX;Z0+lylEQ`k6q~JF*)6|z&dyIBzdXe;q@*FUVA0U z;dg?UtD^H_g={?EIKkOeA3l9TGV%|}$9x@kwd`@d9aOORMliHd%a+8Vu#+lgT zlfaIH1-2Y5Fy;cs$ma~}cj8zl%Eam}hLwkLRF(@YPxs56;#jUA$Fe5{mI)>3EGMw& z7?LBBjogzMe0ap*aU_GAeHdKq&EP_D0XJp{cz&F|_j4>|Y8bvm!`7h+_B+RMa*!M6 zc2aQVUw+(NnBktIq}K~ZBoV4WB6x){T}VG5Sd%drvZ zbUT{I;O}*u`bNQp^J$(RN%t=*cxbN;&(&4%+Gf(19WMO)k&E=*iLZA!@txO)A97jv zvA2o;w6^e{Jp%vHb^JI-$LuHMcU5Efs)Q4tRq^AaB7U?|1zwv?_UpV2&s6Z^@tr94Yy5G>^k=Z7<&v)^vWo*1u3lA@FmgMoQ z3*Q?iKK+g3!wm{PXsO_%Kn)*n*74aL4PW$Q_`E*Dzo%$;r=y9NHwZi#$I-i-bh{|& z#Y`PnP1kXGl8#HeFjKCt1I?i#K+=*C4=+?Y~><6y$T)>#&YrE_nq6y}|K zlX+nYbHw_Xx_dK&I=cn@(vHFMoF?oU%;Ds84mT4Rym+c%zzq|NXECh#)x_8g1%EUd z{#=#gZxu{j*ujRYGz~Y+bm2~~4fiLw(fGlGha$=6X<^61Q7$|bYvI8l1&xfr`dU2H1JEL4IEEvy(eU+Ih@gACJ@RJpQwrw>YWrR`o63ntYa4*BNiV zO6M_E{k&0O8?U-p;e}}(J`A<+e3FSff7MX0r{Uro3|C%axGi1akqZp3uUGKtQI6RY zHT?8I;D;z3U*yrz+Nq<@No!Z1;f`MjQ&(EJdVz`h1rt}-(s1o03xBW6aBZB4%M;07 zJrOwSx`|EQ9P<`6VO|5F){9^=;vJLnr7&{dVn*@W!YFdjlyBQHxu}rH{E)?9QFnpW zqsSgrVz}XfhA01IcqdEX^Ya2<&*S*28u^xmO?+@k!J7~CfBn!SVRqb6QNg9hd^myW z*lfJO;9?qV%*UX-#vG?wF=L{`X4D(4%dJJ2a$j@yl7B5qo^U&{1UN48AH;A-&?I9*FeVTl$edK#BB44$cg;Sgg z4(`El_+^eMI|xhcB-7~}_q5}nMw{DQE{zw^<{sI+wH!xu;FwCXoc4!;^MVCVsl%~X z3x*Bi7`!japhs`!J<);5d##z#qM2pPtSpT549(bcyAND=Eh7bekJ9keL>Hb+6L_Y2I^kdto}I4X zF|QvF9Jq>G)~DdA`Z_L3)zMu<;Iv789Cpx-E&kH6*n1D`du~FbAxxe6oT>K4gm>o{ z)SSy8*=>TDWE1tK4GxQ#<1vpjxCNe>|cOk@-v2=XOjOC zW@5$T94jz^jXp8#8%RFtGujtaA&kCF*l@_iL&Wc`4O}?qsvm!7&M_#!gf>%|49dym z<1NgiWxJ{b7OvlFcx-KjI>bGSaFFPn`&#ufoRJ(eqA%p7SW9(8H%y-m1RsK zUWJZka$#i=HC8Z)yvATK@xOi&Q>%PpGSJK9st?S_)ms>j9?Tfui5ct0>BcFK$9VGG zXM8#BF}@tO8GrvO4D+_j_$|L3n{-le+EW{DzJ44}xZ?3glv+&-At9WTx1Rm*}fqTM| zaO3rC)D=6X#$ClRq>If0{aCo93k9l(OLZJq-&*(H`NH4y9bC-1*Ej`Y^w-c~qKlCPgtWPCss(SnV_R zDl})!V=QdXjA=_kjX{>rXk}VP*l5it6{i^iFa7c`SG<2#CampB*mN=-Z@o2n+5VXv zo@R2Vw>EF%oyj8yf8r$%cf{+iES$8{1y&_1k_#p$E)a6>SWWJKtjV3bnKU{I`Q#oq zO6Jv#Mx>uHn-n8FuSb545VAGlXZ{WX9(zn|d`!n)kvfiSN4T_%_BJC;j4w?-a9e?; zcad(5)-a@tj+H9=u~s1s>ykd#y+U$1m9RUQe9=1!<{55+T*aY$Rp#9i#N=WplTYuF zt@?o(HTIcC(=m$C`y@9eHYd4CWyY*AnlYu1%NVp&G1{goM!g-1QMwK@a{odym@edB zaZF~lU`C$#rctf7Zp8n=jUS=_KHEoQ5T&h0Pz4AM;ShxwU+_daNC{YN5C{$&Z;MYUBQ&GjVJ8(U_JILiV5- zYLBmCf{4e7jw-nP-bLJVz5yOS+Yt|T(QxNF@&PMqINN<6#~5zxG}Db?L((u$!)!SI z!GgKBbZFMXRFCZtndcjk-R`W*ZS&m*E0SSU9qc!n&rLA~mo$y(quj>)^B0ZOMoEUN z$!Mcgx+xO}D2}@oSkz`e$=PGgNP6Km4%O9+Px+3Ud9P)dxes+VUS#`?bAv*R?Y*;& zh2wq3&;4yit1OpMtB&6&Gfy{46;+H9&4rPpm?9TFG}T(PA9+7T!O?XYX3*L`d6~_F ze)aRhw*@c!tHOf98Wa#%LcqPu)yd zHHp5j~7(ZPz zjl}u7F$yhX+6lXn`p$1GT-n)}ca0mr-4@1}=DN{wkI$%a%r4)dDHlvKy$>cawSNu< zi8g-V_iUzmd<=>lwm{urLUNP_X?g8%axceRMdGnslTfVjGy^N$$-v+P>G0x)4g3|I zSn6d*{C-wvoS9S)*Y07s<+tIuzQhVlYw`n5eK-~ed^Z}Kl}^VBgUiCxEt#9zv zYWFZ>WO3an`Q9(f9d|;-hc<{2+MlW)tLGUIPc#*Hx5II=Ydn^D}vu;U`y6CSPP z$ERWQd9LeQd4agqyjb8uUaG3(<<|D+6+aZ@)w=w|YjqCe;iJm)h=647I6av=l$*Rk z)yKR_b}?RfE8stQi{a&qLAbB}SX>=?6=$FJ;i#lEjQ=$r3u-1TOcjpfOPK7{naPXI zO(QVYZRCBR$rqZCJAWoDMlSP2GVgfCGXMC(yip~Yj0|M*)jZA6^12PZt7&97n;Kgt zk2aZhQ^pz_ld2f&GUAPuhsn;Z){J>a zEn`9MBx6MdW^7vA+1S%qGxk=eXMfEymLM~x8k#XU(K5P|?6hy_G2$a^MugjAl#@a7 zWieeYo#>YZ^O`Uw{v=E;7!93z{{saNpH|Ob2~wX-9t{oC<6-jc6xcAq14f7dv&t^m zmt!qtE_cB72|*Y%=Ok9`7Kx3^mdD0N2Vl+oiCCiZL(Eyf8{GeS9~^%*A2t{23aRQj z7_s&pbUnNpoUiUe#QOjUeC1O0%1*UuqNPSpQPfIJU76AAOz**u9$7ANvaHa3toQ4P zbnj!gB1^}H$?s3ilzrFVl;!i>as1d)K!-0Ou)bC#%+NPG=D%n1=u2UAo23}*ONJVk z{PT@BozjgjyWPgS&|u?kt?tIDx*lV9`%q(9Xai%+e#>YO;F6C=dE|_l+$$S0sPxpr zz-MVVv|0+{^I6^FaN>=($-+Mau#3w?-NW>d@x>h!o(>d@audXIz4kijb8;+dul>#c{|_}x9VSKsR4aO zWOfpBTv;#FT6dY__qrl#r(kN!-KP49u(3|64I-C#py?0KNP2z!{d-ME5x^i2~m%Om| zEzkdD7-w?|Vb+ixcsXPv9$yuK2de#!o5noBv`;T_(ufcE!{_rDzI_~491(-n_8q|* z-_1c=rD5PXQ$;=9#pRum%kn;Xpn10jYTj-onRodgimcnnC70J5E$?0=ObW{|a;_OI z&6Ytjf&OMQ{9@#UpW1we zgf(-ZQ_w-EarRGj+w8919Q{19aZ-}p?0F`i+iXUGRJ&1j{THKJ^cf@c#0#Us$NEO0 zQa9zNT=nF`t6${x#!KbNCfnrZ_K)P00rzE-CzkilW5pZm=234$n3_0^t4}Bn85YN( zV;KgWLriFq%LCDMIJBC;$cnatE8?$FepPn z|N0tCim}6(3_H|5$kZo`IM@PRP&toVomf(MS9jpvy&o-ATP2`enyyYx<=!*%O|QR- zsj7_x#O1R=<(GWl6Vly^E1Ys+>r5kOct@k#rJoJmd}16sve~%Vx{7hvIAPohbsHy3 z?=V*JxkiVS0C}opk~j6P9f~Yc*P^IxF%bTd4( zqdgio>Z5UXA09dR6CQau68Fr?#ASDaaZ;^8*y~h2Y&9n*Hheo3D@?78`S~(!zxEV>cDOq5V|wI?znItTG+)-4(o*8;qB5*iQ&}nOhz$JzvSo|w@XtC=B=d0o@9=_$v%g@giuSJXX?=v_aXelPWaX6 z3G=)iVEV3fXzRbPewZlKuD>ux@&ple_#<=t@Q`ITE6lv-X8UFGb4%`CV9CGJd`7_% zE~8>xXQSfA&PK4;BX1{emV0k^mg$d^p}QKp9$F3ogZ7DD}(K)Ud0ZJ+G6|#1?z8Lf)%`7vG~vxSiI2y4EcE&7Txd; zbIlt7H}`gj-S3XVtg^oY#Am2hf=f+nCsfC5o_Tbs@V+kXlzFi;?WXb|Lo>4!X)raq^&KFl@M(^6N5JNdrDCi5-{XR=&A zr<~I)NxreC8I?v*93+oZ-sta!7^8|*bF}^ zxvpZ}oY@$$A_DCJ(S3TuW2Gt`Y~D`G(5^&9jLhcQMr5xL+3L0FU0jaICcQN|@0}&*Z+6M#ZHnws!zJfNd1Pkg zAbGx3n!Iu^MP8^mR_>XXCZ`(?8Mn|b!wPX(bxXR;U)1JZJW+A%8cMO7nr>Ks^C{*k z^$shp=U6p=HkP?H4$IaojkPxVF{)=8Iu1Lq{`>Y==co&7eCUYfs$a&u^NwTgp;NI) zXdsq|y@@%iJK=&o8@6PJg4k+aw1*wINAb8;%bCob;+7qPJaS5QhFr6c$;=Q< zn$ObY(X~PHMYP{2wm#XYpzSb<#m+WzjVUevRpXE|#VE?beIR*xqG9zjBo|Q{+TYl* z-e-YP|JX79H$TRo5(sZy7%2rt?zdsRAF{D}D-VVSrC_a7PAvaQ$9%;@F`#ZJ27T|t z97j9DgJ3sYoa~3=6a5qr^_XcncUJaMDA-5 zA4;P>s^7<1c$^%^;`?oo!0HZE*oQ?Nrl z%Hb_?;`f^^Y#J&syn%(4b}3k*q8$sb@?q(oKCHCFiDhCGEYQx6`B&Jn$ZJ0qsOy9$ ziUMBUg5{M3{B%P=z>i$L(nzR(77&i@v`<~%PsqY9AtydIaxZ=;Z5mdX8OWe=JbDD3|$$54OB#+8;=p)(lPE)1jhUmf{}%;Vl?GuqbZkHm1kk_!m^BLk@ESNJohWO>;7++L(dut5CNuWKlC{@78S+?cYdmidG|e8O>qvb;^?^ zTCn4m1qVlIkU3w0Kk8WUE?i)tRUQoRX<$jpJw#qF79VvZ9QycZf=Euz7pTLSi2Cd~g)K+DS}lzVQf?~e+{$NkJZr>gKi zJP0HTzOg;Q;OTo;+a0PM#BfjnU-!NPdld!*}}Zj z1~5mQhw@@i1Wfl@a3$S@)1)tR7IJXf8PqDw)yuVnquouGIkpNr5TG#c(aB79@iAHK zlgO-Fn?(x24n)yDFXey$TTR+aR1z@hlMZR0{cxw4h1Cvd7!%;bmN_YIT>B~}Z}MZ0 zmu}QD1r8hN!hWwAwrK0aCNF&0c!~|1RC1xChmJMF-N?F*hCd1S?>{3RMVj#PumY!c za7Z3yIa+jNvbssWS~Vs=EK!VnoBZ<8Ph1{b!R3uWMLzth$)neqoL5+rO*fienaHBD z(wVxk5QBUp7&QOHA*GWBE5jLVkJaJCk0zXa%V5nr9THlaPfY3swdiQ&wPQS*~k!Kl{xk9x;XkqzZEQAh35jxf{qhocS8>=33VvP-UEZ^9U zc`Xaem9$~+AUn=E9fZG?@uBOi3tLkSMh~**zh$Q2mWjvl$hIIn)J;dvtrT27#f?+D zP|Tv12YY4_<~VhXZtKMQ-+8doI}`KXW$z%ul{eVF8=!3s^ft zhr|dILcCo4OtQZ>r%<5@b9{^>J576n8xxtd6*i^X!<5Y~3Yi*2x!hAs{t-y=RGTnt z@0{{_3)7g@Cflqq@vK=t?1FK&aF&thjV{F)Q&tHu52zjn&F%^KP?uw=B~EygLxFQP z(j#5K4+*CF#V^z`1z6_r6c$y{Vh3JjGwPa8V4q_ae*XvMYct(Appp%fuDh@s@vBp1s?nL{ z!q`ZS^2>HqvRrU{ivWFt2^ZR!aJD94&oQzqNgP-!p5bR?$8s~d>{mrTX1e^mm~uDW z6?yhAO}>sdWpeWQxHhz zyAHKuyW(VbrZbHAo&1%bIo91mJP0J}?ZFl=s19*R7A|=~bq$R(aO$NHoH!&DXP#AX zb>(#2-83DI;0)Y*OT+cQ*fH%;5>6X)6-QPyadc^c1Gekfx|0nf7nv9`-w)>t)Bdf2 z4s%Lz`0=cOu6Y=6;!n8{0by@xPd<=29{QN$f`esFPNKc*IpN)|32$m2=5?1Qe{u-z zDO)jFW+|8XTXWfKHSseq*{><4JUUbug$LN=Kd*%$%4V3s7pj?sLmWnzJi2<0*4!ZZ zXwbw2Gv0(J$4snQ!;X>570kcXg@Ki9@V1r)$rS}uDr`c;MpI>fvdq+0Vm=;$!~ST9mw*b*C-*=B+D(S)t{O~~v|vA+~Wy*iS4 zSp$mo>}9glcqT{OW^(y14{?c{MpZmnqv_puEOZQ{qr2^Ys_| z>4XPVc|+qoWWtJ`ggYNN94W66PICBiTZgxeEV$i72mcU(xz^Y)=N=9prqR8_8XWDW zz=cUV+;v-UBTRt@ElB6GG|WjhsCa&XwF^@1gI_`Ao(s!9bz;a=hQ(HL40+|j$_3pR zs%lteoEvNH5ZEbD$B8K;1yy`i}Kl!<%m2;AbaM&_h>NYoe9C8xtbl%R5#UHbRNyrk6}!0P=k5Z z7c8>~VKuA3<>F{vPA_4}ZfRWNWbVxjX5Oa@gzUT6luZ_xvg|FAox@x%ZL7!+lyj(; z%H+#(mQiWEX^uA z(+q|sFvw>kJ4W&|x0+CcJ2Q301(uoFlzDsaWscEKrbdirP_``LXI~BG-xWYPZCJIK z^nz@F*;<1ucrF^OEYy1lAm&tgP5(r)&ve0%5iurTldtlHVs-}<@Yk~7pWzx7nXF@(D>@dKV`71K0;?8ah#g2LCRo@$gkxA2 z8`^ejXv?8sani+ng9JP)q`~E99DFlOI8vI!U9v$DBQ%_P^Ei&JpyPxFR5M9_&+=v# zrvK`~O@uMJ;=<`mTqw}Oku!B{|H8zGFbgYtH4G*U4<>BeJ6eov;ZSkIu+XQrMUM1Cn}nLX{4E4?b@lJB{kpqaAwVCu`&%S77Q1$wRBDno>?69 zz9yEuMYWy@9DDE45N?_X#QS=e6nOEKY+|@LV6&K1l7#FRPr1I6LQeW*$|3#8_em7q zmDO027IgSADTH~Qgo~NAn78FI=D5|Lsjq46;$H|z{!Fzkn*{vkH(~8aT1(*YrlSJy zxDGEP8C<_*f^Uq4ITzZoTsKDfdLPC<(XriqKZcGo;aRE)N1L0F5y@dQ$=|+LCOo~! zFyyX|_Tm=CZC7xRQ^D~)J=lYM*LralhSPV=q=pemCg#tf!=1zQ`(P8kk7m%T9o4$l zG_lJ%4O=l2`{dJc&}4?c?BF;h*2I+J0;lD5;gHuHorxS9?G_j`pX@8)-S#OAHm>Jj zRA8`W0*78?j|T|>wJ0w7z{S)OhuDEDF6M2rm8njXsn>fjDA-=8MT;?6w4CtnBN-{X zi+L*!WoqtpcHpT_K3p-vj4Zv|Py4`IXuJWrC>-4v}nE zW>8`bgF0P>dWK{wdNbYsl|f#GL9E3f@}+>HhbU(GmI2iq!_=h$7E?RBGT9$#f{{#i z!^Pldn*WBo3|{5r7|~V7Nx$20Tvgi7`b=y%Ov3@RCyiZg!inh|!pkwoH-Ej_5g|uD zWOB@GCR^SU-sR(&+H5R?(oF?STub;D$6@?Z$_GpoQ2e%lc4Ntpk64tpkbG}RF^b@aN?Z_m)B8UzzE7o zxlH(p=D2PP;btVmK}iaVqYMXhFtK~KhCk-v*rSrbAG&kwKhMIc1N`{wA13XUG_1Nu z$3Vic$)7kh{F%YY7Xm9C;#iBu)abT?VSka_4OTFgFeWBc!)B`(Mz=RHg#4}-siezK z$d3;peRxbY1@#0J{>aqZCt2p-I`h^fc?lt3=Ia?I>wIK#>=z*?(EhZTD!k;+sOy%p zsI7h`^Zq3q3uwK{yWi@MyYAK zq0|ua)$TavJV+7trFWz)&K{<+pLxblMu z+usXlRaHRk&vfSCGKY;FIjpb2VPOQ-L=fIgUqN+WHv~wJfW2i**nkXXD&%vn;&5H1 zy4{ij{+z(z(0C4i^)un!eR{4w-RsJ*`9B=Tl(lfU!@|^a3Z{i*V8>dv#d!u0fB6U311I!csrbY`xOk{KBRjsOw4(MWAJRM6MCf~|C979kB-w0 zQ@!OdTGtc{4v*!q;;IP?Nsi8I926K;Pz};v9m^Kc5T%a2 zpK3VyJL=<}j)Ox?{Jyxr?_LV5Jc?t%34~qc7%Gzpo31k~(1v5~fi#YSWIvN>4X8e) z+*jJG%^@G-t^&udn~?T5gNS)d{oIwQ{WO+Y(aAEmQ4Hav&SaHWq)&v$Q>rqVvPH;s ze+#*K4VNwJGFijMyj|uqZ^P0|R@%Y5EAKIH+EQDGx1#eB7AghuxUrxJ!IL&3Mo23wi2}{B&|w6I`ljzzLbUy%6oyNPtp5^a`=h1O_Ta-xEzH4O_N7FeW! zj=9bVxE0Fa^#WR}Rvhl^qPST-22az7FN+ndm}Fx0*&3!*)$!^pKgw7O7xnky#>P(E zN4#7(HVqq>qV;ufEccpYT$sT6J&5P+G&pvFLA^=LQ7VIF(iu^;iJw7{6b8S(p#Ad& zs>SV1{BFar1ScGMrlDhzhK;ss6kFq1^rVIXi5k2wLVodQh6O_jhk7$CV<=ed2NRq8 zY+^f_d*5>k&K>E&jDZ~2_wixJx`fd&CI(Khz)!wyjqU0Qso~uf7CzhN#>1-> z95bJ=y@Cc07jP^#LSP*I-)obFBb*#3tnlNA_hhdskv&V-u;_O><}O8gagoLxC*byG z2Jae?J(xs3{A;R@zfN*A#KasCa6)0wgfq2@n`O@4#WJ6{nKyMP<)=tCQUsH^ zW-`gk3OOpAd_lXAd;!(w=4Ia5gimj5%yBY>sW-!!8qkJ0F5hSB-hm7z?V&hdehyW8 z2=yQ8HzSMcQa3Q~+4mF+Ss>(}2e_PHl^LN&x$G7ujJhSb5r3C){3(mC{#1omIu?LS zsJ_{nV8csC7+$R=@NRAoZY8~IL43{xs_A)dVwEu(CT1x(>4Ogi?HfnZwJVC_kdgw2 z*VnM?cJhC(Q_TH1`NMG(kI2h0FqeY)kCXp0K){uYBtx$mmV0Dk>q{h;g>=k2LPKwF zjpqzwcz&lH58QI&@q-q6j+;0u+m4M2Q|zc4`P_DnrJFMh_{l`u0t<_sv|!b26ABm% zDv)h$U=c?3+|nuF%RI_b3Dgk!5jyyvvBlnx(?^qOlMg70L8E- z33xD?!{@UcPSi5NJBP!0;?>{JOpO;qIgH51(&@VzB{7D#PN?G2`HJO~Btw>Kc z!%VN67dW*FU*)&)kc-(o|3N4ImCVuq$c-2FbKLtQoiCW9;V<_T?BM0tZ@h&I6^_3r zGh7g&A2=EM}T=H>&LGvsO zA4L9S6^h}krW(=~9Ou3gc%d+zFH7{}%l8WUH&KoJQX39E$1t31=ccFRFBKASxCevu z&J1%6Rxs$Z2~+Nq-4!Le;QCgvB| z&`WvhXADb%iG@#c3_7J?!Lka)x(Hk5Dp)LxaA*SMvl~%eQ{~vd)sh$n2vJ?nz;F$ z^0eTYzb#{qiM010GLK?&#Wh?x)1{_n2;wiDw<(^E&(>ArtL?fb94MvirpuT-z!jG>~NBFB2C|5qPznjpzBv!)^Jp@e{=ocfC|_ zXa|9IlVZRW$N0RFaJx5yZiN}V7)JK9rU}+u(%Xj&J9H%sJYwO1Stbsg&2Y$Afg^Ks z{I!yaza(qeZ>NSGlXQepCRTh%xY)}gKa|EdoNQ@5hJim@nCA|M*GCw99w)HS1d_4u zh`P|TLrkpX(s6u4hHLJcIJc4=TL)51vpmiFF3~s>`#m%9*M}6tnc~9{_XRetAh2*6 zjq5B!_ZZT#Ms$YcoWP!hMb5e$yLDIaC(S~>U&ksxDVT?N#@Y!?jN{m~j)n~v3#`6c z!y=_ExY>xq$wUGBr*K#@nL%6|rmm?*er`_ze~@pmex-mJ#NXMEO}H_X_B-hUUS@L0 z-a&Fwnn6|p;ydX=uI?J<+rZ)DUK4J%HerW{LCQuB(1ocT0$JoO!nEtnC_Yw8$QoCf z6jw|+Q&KGC4doX$F|W8peo->>UP@)^zJm-;c68x!X2bl?Jh+qi@U#W-e5r{;q>c-j z9eamr*ti|V^X_t7m1N@H6Lc=1rjF^C$+so^nKImk$>k{*(MZF3Cv>bZ&cb{JE%=np z;D~O*nm;Mm{1b=n&uEVlVPdWwIx2}|CrSt`){4%A45PSTO|k>S8D86~`tQ5q(H5jcIc zg=vi$uAq3$Qj(MT%#YJ3mp$sHi3y61-|bScYGD)egi{WLWbbw$pOWu$hG-hCHR(Wj7m~HMggIX| zEI)_#jV^&bCMcNbrQf!u9Qj>Am}_DEbRFy3bd1WQVH=HO^FspL);IBoxfXUVqhRt@ zj@nNqHtTO<>F))+%Aj~=6!CMF30Wy7=J=81JW#;BR}{w_O?Xwp!ctopmR&|vh;l=v z=se0T3+;u8H$K9$UK$pdNjN)0!Q9hLcw1P2Pcvbcn{u1WX|G>hK(UuhO*_mSyKXW^ zo-F3r!6~MekIAw_gm>nCCRNGTvW0saTHmQ)0bwr4hEU;pPfmc^7+^?R98`Ft38k(53&OpnW^x>1-$EW1g zml@ylYZU2Y$}c;qVZTzS*Gr*t7Dqk~RR~z^qFY7>b7!Pr-B+G>Op3s^G!vV?(4IaL zM{#c>4;kn-+Clw%E{Yzukg=PPZdM_GPEhik67oh_As^-;=J-<~e;oPyDTTt#P1K&^ zqTOo)?WeiGE`@3ab-Vcng1QsqlP63P$ZKlU@=;HI=U{%WG>q-&pk8J^ z>s^hI_X3TcvbTacDmYP~8P~ky2?OQ~7o##V2b^%wegd&Jf2PSI#@9CFzB4TOPqLPO z4pAr+WubBb1rPteMJ~p_qs*gI`A+o(dh}o&XugZy+_xUr9JFX-q2@A$3Xc>D4iU)O zN66=+8OI{1^W7y@ZcDu+C-YZPC9Psg?kYeIeM8G8YgwZrj>_CjocK`_qc>w6m44nZ zQF!|h7meUM4qL2oD-+{Pu~=_xCi?VJ{(dKQ%`z^}laf0RR9I9n4aX7;9B-V6BP-Le zEin;M-H0JJYK(g!&^NnAW1oS-hCoKckyku?7EmW*zS_gMdZd?)Dikv6lMW1|Q&(E#x0JjxdHP$fWTukp{Vo{L7%Jvr4bB(v0O2ZB}0Ba@_ z?}w$J-vkFOo02PUHs#CaLRvGm>_H#QG=drtvBPZgL~ll7)pZ)#FFDB7!bPr;4swrk z;D1}l$32xi!g!e4N}$M0K6?xDfU>ly{={GRTm+8PsNU3pI-oI#cJxDC?(sqD(UTms z%AwF?k^w0-FvCFe8YVjJ)#&A~(fu?1@v(z8c}#Q*Rv1K^9avYP(-;els~c$anArcl zLImUM^!)#VftxfKMWq=WvOX@3pLQELd4E5Y|1!?=1v!~Y@qZ#C)} zB^a;YQ%jjEEV^6`><09li59$G1 zxHi8iTn$aamZ^#IOPGTq%0g&+fyrqG4tj!ctZ@pKaxXUZbZ~4|1TGE?!ZUvZC;4ts z`3#i(kv9BNFCodD876XK*=rSS8?bayzh%xYI_9!f7T{(-1vX@zh8=>Tx1R)tklaF1v!<~8omO8VhT2feB@kHt8ccG<*)S=938JLHLJi26-H z=W#Ih1$}&{g-R{RnV&0p;DjTGz7?|hP$379n=B8|asz#FX)hty2k~8FlzbJd<(q1Z zOGEf>dmQ;Wqrx|n9TX2XQE5K&RW}!T{0x+`n4g|#v@s3L7_Oj$4fKxE81`PHUw0E7 zhOjQvk{Ge5Lc^=Hivt496E!-ww9sLsBa(--#1H1Q@(4* z{n|z@(~_8|BXb{ZRs<2-MN)G*YspuPzrU4t3fCrXTI9a{+wmM32zqm^Z3WJ zt~i5!W700RMWEx~LZDvYc$dPaVb)@g2Blq+ZGI^SiUdHoZ#}Rw(;Lly5 z9`VdU=X49hrU`WG%hq@< zFI^mLXd!70=Znl#NEETKVp6bpvWY3r1$vKV&M(2d6-WN~7wzJmkRSFlM{g8zFLBSB z0?ZN1v|N{_<+g%K9=f6A_MMJ&?hy~scA~2**=v)QIV$OpC#)3)R~DY`{v!B;Ba-t4 z2+!#W0`s(3hO zVN3|ak)SkGJnYJw6_tG5Nys-l6kC?G{Gu5Lm%6CBQKMR4fiCF^10(pkCu6bYlE&5_ z1rm;GY{B{{Z#tJa4mmcNe%wL_i3dOyvxLzrv~=UNyD-|iRjOG@%=|~-6|Fu zO=jE~OuhD=fl;*t2G9mO_frVZMEsSR^^VibXQ{M@Mi%WvpztgO{}l%EW+1O=M?YJ~ zSU#6_dxYotuQ3O=RN#0abbdt}*rL$ku!S~xT{PTDyvgS+#QiK>iJF5W0thy@U0G8FM~! z*3LN2#c&dF>am5h(FVNVaqdoA3k&arVR$)(=E>9>wlg1HRPy(ETE6H(46)mh+deB9 zmCxZBKOwu`5VFB=?YU2T`1+HOV!HPH<`tgTF~ak(vGDAuD?C>p3eTaEBBVmB@VJZ< zeYOf&oVNICyMx72Bkpzp+V5dZog)x4$;IkzDcJslfy53xCz~k%M@ne?*+#1Zsc9jF%qgd_tFZ&GpNkdiViEC-&&hbP zbYu{MOA6Gub>Du-!Q zxJPFmXRM zgH@)yI#tPCu8?oHQP+~R@j4btJ}{6!AGKWiVaeK5ij4LLv6gAWZt6zCDB(a+8HYfy66^ z9sKZ!H39x?RIxD3{f-(+8y7W-YSfD-c74O|t?>9P;w9XSQF|QqYdQC!5Cee zy6I5r=hc}H-?DzwiM6KuT3%nr|5u23B|ClY2M3XSt{HcTomSEJ%BEonap3Yf5m=rS zizR~{EK9Vov@H8LE}GC~3=FGb;k!;Qn*GN3HBF%QTm!|j6YF0w<%_q}X8CMCh6}lF zg_cnRgzVQ|%g7)t*B^A`?q`nNyw;Hi$UUz=7xMBB{%k{XTz|%fjv57h2C5t(Pt9c^ zUkLqW3-RR#VtV@d)Z2`a1r0225r)l;W3l}&11tDUu|^nn)E77#C~&1D>piPjXV}M_ z8Ksbcwos~6$JE-_0HUA2%1H=)R4kst-Z=qW!#_*c_&bbPG zo|8Y0r(V8H%O#DJT(*>ck*MMEF}~D}z{Y5W)nzpjDl*sIbg|_y*X*gr%mDgbbBpty zT#Q&wTe_<;Z;-%(QRG^D9(h5d#wmePo2f~L^IdBZBjs^X|CvCO@-A8k7p*PE41Q*v z-HeT8iJjV7=%2tk!ZwBJEj6O4?Je)>VrhH?CT+D4b5mnAZEosfVzA-lON<}=Cul^p zG!R^Yaj2p|UHVX-`C9(5-H~5v3S{3SkcI1FrUd{{Vj6@T2O0ag+#)7%r!FNOHuiGuOp3i5yrBUj+M!qA&eYb@?C&&k$J91S8 z_TFSra{FQ}cLvdxU5y++J1F2NWNYXm`$_}Z?~%hCrv9*o`&*Lxy3Rn|Lk8+43vw0% zZMe?irJ1vDD@>|tVqtO`mak63;-QS`OBB{*4}g2r#RJBogTJL=U)EU6H4U_AOH7-a z&pOdXCy#|zS&0=2Q(L&fn7+?Ny-pTt{w5Gwgg8(M)O*4BGK2ZV&%}aGvBcqth^K==w5eU=DxTFT!q zb@6>p3!ZL*{nQq|FUVXGWMDS6iCMiPFrul#lw%6(%g18rZt|d`0<1fu{#t=L1DLb8 z4$V4SsFh@*`f(F^k`*#FGLR!&%glX0VekGS_p-=NS+7Fy~2yG7U88hNBLPj6&9W4_cmr6;`4jp;@ipw>Tzw0Mba-Txc}FH zvg~E-$z;d}jQc0~tXW@@pS89Aubn!dnJ30La&0#ydtVfu&p!)Uj5(`x9@eACgVrW8 z9u#F==RIw2tAX6z6mp-`C^3|=^fC8^YyZ8!fgv4S3_c{#d(40SE)Yh)>Xe5*yiQy;68H45f!a~jR;xP1&IW3E43vvC(Xk(YuAzy}Y5cCT4w}!S&1K>9jj_<6 zD`UnSffn~P!ct7MET}R0q=}8iQV{z&mYmDQ+$O9g{TzVkmRmPJ&B*4Tg$37$!S7ZAG#)jzfi{=)=zj2vS;?sCPLQzPRpU) zgxuQ5ktYoyuX$-pnYh3EER^ctpb5W|eFUhvlJ@JfP>=6WMKhLg9ZP#O0xvkI6YC(P zH1lyu7ga8ps2uE~a()Mu+7N%9a*?$r^JsB$&R;F*7Glrw5G}9qx$bl(W`7{$vuZ-# z+pFc*(*ph-P5JRW^H*zy^3Mc{)#dl}Gf?Qfl8=@Czy4k(#X{Xo7UI4oPCUh&zgJ*k zl!KL3E$m|)aN1iq_&05ZIrg|t!?~lexSJ~gFQ*xJ_An9Gn+4!(Sp)lWreLif=dsLn zFt=9h|7XB;kI;}WP1JK4gC;qs_sT)-0uGEEChCL>?w>%VC&ZLnT^W_=cuVvWc5U|E z&Iu5JOcQMAe7YhAR>X`=ZQ=Eo%tDtt?B;xp0u4&&ZL7c(yf zVQ#Y&^y2#spg(kJ$2k4MLWSQHbvX;=KM5FR$lq^TD16XFsb(f}Khg-^;-JJfN6xLp zI>iO9O&Lp`+OFi*%1R!eNWL*&q1qSrcFiWn{@cPba)_mIK|E{Y$VdDym3wg{l6dBs zkVQTT*|dk2o)jU=MhjWgPk4H65+UzGg*ONLw{j)0H{H?pH?4&C_h8{&@B_~jtQDT> z)%kO*eI9)-WUF`~=jRpjh}JT7n369BC}dwt4ehmwf?fj^ue+$anERQ97@K>WrTBmE zW1(U_fzmTA(qGwCvMN%ibBaT*#PtE}J9!@_nM8l2g}rQEm_Y=M;UvAY(xSacOP? z$@IltHOS8&XgGhQaPDRzP8JrpzA6^4hezPy5sk|&9Ne-w`*)>@*Kr#6Z@9SNXX1Ey z3%geYVPkF+s|py1>Fr`P_qlhCa`fp+3YV{S@jBQ3#Ns-7HX)am)^hEAAyMNIq@|&w(iZ%nmJI#vn{VWf|>aUDNFB36yxq~)2h&|pYJgjQsQ4D*xW;;ldVK{X- z1s6((;XL=m{x=MV&qQEvlz|;360zCO!LonCFzbbb3HPY!mSX+ph{AB?p!F(Z%^9rM zkEPDlky>ds_B?18)0Yx2opI!jU5?E1kC06sX*uWyxxq^%T@U$)BV-bD!GdwZn`evg z<-cY6x>ZrWMstMyvW+9-N-21wU#+3 z7DAn>xsr#Dvj$<2?_^}{?5iX9lyX8Y?-q8AJleatzmP5HZ&nt@_F?>YIRzFk2;%J4 zF#M8At?C9fxuXh3F~KwV3Katkr|F5=_w$MK=JN97qk$SFZyGb(Zkbh|J?7E`S8b^3mH>VaemNBS3Cn3$%98=4XMHpO{zMyJ+ypK#?<+OwXj`CyVx+!I4jfFn%vna^7@u@~%Rz zVhmYK`-vT@-1KRmq9WudRAn!Q= z>y3fmtwA_fMPu*^VxL_OVmE6n>mPwtIcN}M>K!t*Zdo2db9PpzQs@Lk$&Jzv<}J8HX1Y2kf-MtDa%B6zjQ=gBQ( zLB_N8j6?rvc6(+j`DC0UpWhS6-G(vzlt#IQE-G~;RxD+p+DMc7nkj#6tK>~npxhf1 z#Y3#{WYrG`|1a!P(W^r)7yabAy`tbd9-O%XQ?wD{l+LHciB$Tn-kWWnD3!M!iwg z+47pmpW>osl!bOZSbzVhVGSYHIbdN+^F*v@#Gel}u>3+8*7zCNP@MIMloV{3orvXS z49rssGqRKSb|h}OpwVqOG0z-wwM=0u>%wSNVM(gWH;uEv7b=xsn=0T5pD)8RJj3zos&-Pt@MyC)qpsLVGWArtGN=LWU32^37f9Z2b%@ zS(t_edt=eAnu&_Ey`}>NCUoH%{3hh5CrWZnf?FE&_j=l{b5`4h$7yfF-6D9ZA-uJC z*64G2A$L5bzE(&;6eVB0Lp{P&D9(LpRfzWQmxvY@%>U~qYZ+@!Gv}q|V6Rys`OHrO zl?z(PGo5{1r5U?!JMwr6C>CTNL$J$ybI?D?`ih=U#=*;(>TYh z4*%bY0Fj(CQG`5B5}s>9$Og0lT~^4w#M37Z3;DH`K#s-aOv5$G)Zv=8bkHh;gV4qX zN^jJ}i}aJLF7n2?^5{`5r@Rqzz!D+XUFNeyQQP@jBgb18#r_cRFDzu*ed?Fr2zj`j zl9ieYyIr*Q6&fk*kBrSLhqBhRo%=}NK0GK4zb$8NPI#OmLKaKave6JN zNA#lIOcrv}M+XZ6Sg-kzhz@@;Z*uML-q2{DY@l;H@}c$u6$_-G9zS#5uZ}EoR>&!| zx9DfIxopI%r5!YUuBZudy_Rdv6>(rpp)W=o_~txuDEZZ))=JK=Ej+30ncUx2c&?Cx zb+d$LJkQRYVNEijsqhp`7M=mb)Nksuw)9@e4wtlC$XN9F5I=LALPh#w`3ENIwH0W) z*+u6h`e#n^*k=v`ue->;$&klKDLL~KF^;We*Ft>8hgu%HNL~AdMrMC%h$pBM9ChT| z-0ZLZUCE?i;qev`_Q)li^LLauuCIXnl{!Hcfu1#)pMFY2V)aDU+gvQ7-7Ff*eDgjH zr$-x@8^hS^BTp*jpxsnzYG=r4vlyJoPt4ZNl>Zt^ey*z^E>XW8Nt>6=0t&Jks4L@YxYUG zP@x)mjyv*?%^G?V$}9PhzVjc?YtqNW4+UxKZ4DGBuJHd^ zdpx^^+&_x8m!MF%qJh4l4%WXiP=G#I{F^{u+}8+S(`&I3Q-u6nOnJ`I4!68FP;j3m zyROxqm9#-g?mg!h#)IEuF^utVE&X)W83$vyru{}yhxy>5oifm6pNVGi#9?d62l;Hb zD-h3?Q*xRqWaM-qXB5_Q*&0VK40dGK#l*LNbKl2l+4CWDVKqm3ⓈSqCY=0QF#Py zWh*~3fx2K%#@22IkV@xz|Oujx$qtG8NvYlfsyPJ|HMmch8ZXx#-rT+}n@?=>h zpS#>s=Gwd5)2sQE+*d_;tzE)C%5!=z*9dqDQdb&bqU~_TmzDx2k7}fB3&Lej5Oy>V zLesMX^(!(S7jZFskAYmn4SbiVvF&^S4wNM3S(=6xZyfpSQ!S&fDLJ<~G3E;`@3mpg z*;4Yaj9Q+n;>i8=xE7NHdtAsT!DmmR7JH4c>_-bjY=OwVh8u1y@u=^xybe%1H|3D-60uyU17eiS4`hRB>NVQxaymjRNEIn z{+%!G_(;1e>m!H&FBlz4x< zgADbk2e-CRBE&%1;s!$RGq)5M82!M-kd_+ZMFr0TIOrX~--i-wPPS01iHU-)D?iV7 z$|jY`dbcAvgk?{s}dX3Q zE@DsG(V&_JIu*0fvAls6!yVKq;bPvzAl#W9fV4RR_qPY&Rex7~GbLU5T?u4vbMfd5 zZF#nZH9xY)ab*y8F9^W;cLvsA{#cITq8C6buerW%<1o5qYTDTwUn!aU)k=kLrl z7cI2l9@e_Re^W?t77XK-aOB-EEf2>Cd2tx^$ux4-{p2Hm5I0w}WT!mJbLEx@A*YZ< zT06)XZXkG-iMs28P-Rps!oOqh#drRAnk^RB&uSdGC~)E%7m3Bwuqaz1#_UT&+cFOF z6gK7V)k9%-m{WMqB{e{%6zGm?OOog7;gfn?Yf2P1?n?Andm_ zE}Rv3^Q~4HYa2-SDO{UpVzZ&Kv=i<9zKcCIsV{U^*jzLK=T9k|ESrL;KRDNAp@pq8 zs6R}iU%KSwJC$6@eDJ25CI6WyWT`UR7L|2~4iuhwqqTkEFT>X?CEOR&_pXzzaY-&{ zEhpcYAm6bbj_=HR({~~y(l^iw^7VbIeQ~*>d^6hn`+Da*@2ggEt@l;CG`Y90Mi$yh zm8~vHHV;GIanvl=S=ewT4QI!P;?<`_d^)FbcMf%p73}q&C9tSm1FYN4_ZX@WcO(`Q z!!#ljHHJ4bi4Psj-9!z)ZW>~fW3jeHEH->{v8H?)*6iZHr=86*HIA-Q*#0z2S% z4Vfn|g(LQQIHGI07Kkkbi?Suaj3GKB0>jW}^I@%QQ!c`1+KdE16D_6y@^jRyFuPoT=U zAOz1n+?yX9>^#>1^AjVn!N%+7mbFo9y7;8_hBwZ7I)ZJAAzYEEzH_tV0NxB zbg!-dUqkKNUZcl*>W2*sEK3W(g|9At>cCn_0~d#4VsRv=i=885u{zem@@g)&o@d>E zhJ)+V({OQG5N_U8Dg*n^KYbI5ebcEq=Mv<9e3#oP^1(URQthx<#U>h@}&f&_@c5l_O9BcVLWH;utpk|I>g&ww_|);Y9WKuaad{m^gapC z-V4R4BltW2c3e*g!fD#o;T)0JpE(K6Ab*@#$GKs30@RIJZIL_SGQ#Ji~&|%Nlk?_9jF|V?m0-sCJxn6KA0~ zwVXe*x-zl1C8zIDvhi_8_PML%-kIFXfz%xQf>5-31hNe>Yd5+0p-P1OqmU7DCZk9`SXsynS%qx4U&wgc@NZ>=Oy5C1 zu*gK6DvWnS7{Bki=<|!hxLwSf$1IHJz7OvvVCFE;|dEM!<({#=BCC_^gMBRDfokDS&X{E1fxOBd>!<96*$iLzlA*rFHJ*5u zX=37lSggt!jRM*aS(6RIEa^$OjcxLZJ3j9aA-u|eW~c~@O$B%~(t z%tJaV)k*Xmy(qjhN>dA)L5S*}hy^`UuxMrg zq6(()oTY<7PW|rUkBpAWa5_R|IUb>M_8g{i4OHrz{LB?Kx8s%IS+)2i`8k8@ zt1#E}wR#xgeZ5c1KQ5anQC7<*TZOC|C}iGMN^YAk;uB8!HZ1(fwwWK=WT%hs8H5$f;}L5`V$IliY$#S3%i0GaCQoN9{oDY{ z1IJ?g_9VzJeyEj}j?OtWmVeU>XYUNbO}7m$ZwTkSbAjYM0y`3@(N7J;!9C$P^mz-8 zCT+*@m(*WhIoR08#g>n>y(J1soujd9NIEv{55u;v&yZLm90yLH$MmYP7#qray+)?t z`=27vCdfgRJ+a7UCdx~594Sqm>{jRO6`lue8LRE@*DCv4<=7RMEluXm6t6hpnJsjG_mTPj!kB$a*RTKs!JV|5GO7cxM}^W(@r_X`Zz zVxVbxYA@jta$^@$ZZlo^XIb{@ay>+=2(+Y5IdCU6i7DFqPX}%Pbv(^ie3Q1F_lA72 z(c&GQp?I<~0%r<-K%BQ9#`UO*Ni}j{N^uvxHiY2Ym(l1pH3`%Dn^@TNCZb2T!^ls> zv)3%F8an_9%X48-t9bNo?}x@qBhmC^7M3yM*9)v-3F8 z_!17k;taJr0oZ%|3B1-Y9Bm$p<&T0eE^hz^w2Vdf{1)mDG?2fAiL5cym%@XPnRCsr z{c6gyMJ@Ryl=Y*3s4;|Dn0qi41DT_<3&)pzqps}5(C4}e*gi*{ao#SsGpyaq#nL69?Wa?5&=NglX*M{PjFckBKHt zdEefIGz45p1m}~-`F{*;&%;`C#}WGCbO{mKyNJDHW(|v@x%I} z?J)n#7F6gNAlq+@^i21$p81OUO}#ddyE&scyt91M!jZcRGUx1xkO%(a{Klf}9r`Cl zekntZ_-Q2aE~l>bmbm5?xzJWi{(RGrD{}~$V1~(a>jLD)Zb6>oLCSmMlw()gpzSgl zw4L{ZX=mSL+Qoiz?69`lj#?aM$B)wX{=ULKw%D<|jdwhk_!%2hlpJm86kxq^ig#s_I7GeY2>!%9B(=WtnnKWg0w2iE)wm&2I{>zyAqt(|_>1 z-(6h&D+vc1hhqD{a@bUec96>--_8s}h36V&MH=#QF8hntj=WdQmG7&&^7&Ci?y2p_ zjsDb(e8h99VVoPTSYJ_yZXbcUJ$X*)A@#)*@mP0__p1zTfJ43gagA8|Rk=jGJCm+5 z-q@}38?98a63HrOhd_L|@D}f1X4N;2&8heSXZeKaNiE#DIE2hgd{4nu#JYiFVx6MBj|gmwn@{cst;^K+|%5 zSQ2p*b5^auxNZifg-2j~{blHKZ7SL~q_4d1iIK>U?u-1Y?fOvfIy&h)O^p*<_pl+5tE<=MBz39i61poQ9M zPs6oZ$`7s#GmVg?uCRyaH*8<^So>TrWpC&k;C)hp7`eP7JFVBp;~Zg+IH&FCw?X#o zdbMn~V}h^5{x$Zk7iBz0bGXQvU!#K+gZYPdbJn64HwN#*@1Of3z3DBL;g*3~oQO*#rQ<^_CV-_+5<@>u2?a@YXY zth%Lox0F+MsfUiINii*-&9m@r>2)w``k`kkdHWxO&|{hpt^e+ZkQslY-t65dy){%G z>X^g3>!o9#(*gF5W#_%O=Q?)N!&`i{zklzG9g@)(ndhQ?;Y_r*!tcrs?8>__W6IfP zdXOythmcM;LvG({$hcOCJey5@@oqZrj#Hk`XH4%OW!dMFMcdz<4D*(VCno(UWW^LS zWMMLAqvlq@ai)+n_c`)N&oo(RxVAlaO#9f~SbO;p$GdW!kmc$)GPb{v_1AHp;U?FP zKljNV(cq4^%?OozqMWdUhFbQNeEivXSI+zQI*eP*F|x%&tb8>ehhL1w&1|>vPliP* zL-Bv{`NdM@clmEsqCi2_D6F3fEmlE=hVgFoc}-Q9FF&eIkuw!arm9LQrIaybxhhht zzsj^|2QK$}gB1%tpnFsr^3P3|`_?JR18Xu<4aakTkn$|uuRV3k3D30W!t?h)EgP&B za@ZavXKy#}vQVXN5`x#dC|BhzYSbEv34O0&P0zB3 z`TjL$es%UFTN&*RC8E98RvB`4KnS|${0WO*9>VPWN73Qic+~qf453L8=tKR(e3^(Q zNzahwLaZ!(-n17kw|xF>9Q$SS6nk6ILB4zqX8G3L?&-VL`myiejhVipf&SiyYfO9F znKr%-o!9#YFZtlhSnW>ojRrz?idJ&g+rn~Ft#o<1ZM@v}JV0*ncO=h5dRB#N?@pNB zk9{3)*d=}J9%o~$eCFEi&hRcXUzqJzASn6BLy?@De?ROF?`)o^8@-F5&Ga(lqb1Z9 z?>|HH5<8J6Z=CED80($i+qLVJa_zyr8`xX#KeMk~{%Efom*%PYy9NIaHm2Hhuxn0N zT;8}BFLQRq*E>U$_-DP!aAmj($hBToX%ej(|By*F-Q?N_&_i-)*rF zl$3@pWrEPgr%-!_K*POZ7Qe@+{>2wJbhA=i6*(CELt)`uL@+*A-5bKXW!`BG@#bS-pXVQkr02dA3%#_2HjuO!{Z z>QVk!@wFM&KKO(+ZYmZx@WZ%D5eS}?B5(a-cuo+{wY{$FFHyW#&-L?dy>Y;ocs{Rh z_~<0wdFlhp2-bI&&vuI ziCHow$>w7(dFC{kk-WS^15bvwt_zcA+hO@*y@%>( zPlk%hlW@%RJi3-BCl&C@FWrx!;k+2m7AS|zdrdj}VL6#+R!`Y|!Z3OEKo|<%+Kl$& zS0noSSJ<(#8qTDS$Gz>t@a}#ym1PL=CxwiHIIl{i{i&+{`bo)v?y7mIEh_kZbyazY zQpLMJQTdA&SK0UNRpM|nyztEH|W7HRTMG=}$-PH7jOhw_KYv2*lI2kHdKy6dJE)&P%nkD3;6(Ux7skor)L=4b z#2k*txL-mrrt(^z^Dd5|m#Fct%YnAxqtR-8B(QBPDoh%NoWm|6V>ETy_KlF=t&93O zCtz^(28arOf_^z$$s1yTU8SJ1KlLtbuUir;V-}~PcdJ2Iu41q(>2EB0R0iuNk44gk zpKyLu2b{L+z}vPfocTdGd(uVn$st(NKL)+8M?p zIPtQf@Prjn_VnDe32E51c&A9YW}&vIsggI>SF+};X!+Boy~x_~14OMPxuI~V=WKuq zUc#;JjZ~{5WF@PGJkKbU1`}1&66T>5-~j^fXl!TUF)#>s7@qEmXy`GgaW+R8{r_ zYv8MP;{BO)__pOhT(CdqQe4M^(Su+$YX<9eJ48N>lb_# zgzzpsr#-{FYB{HdCEZS}ha4t08j^^rMHM7-M$6o3_?B4q+lA-xeNPKfSNmc8`xxwU zOeAF9jTM7#VDZ~Sh(DH&ZM$QTSYa)8MCHJ;F_|#5?Oaqoc2oX4sfB#dx{myiEd{yC zSK%3`+bEcOGb%NXN4tMAV|eIY3`jnW#v^i|+^k`!UF0tM&A5)awu$-qIg@5pBt~7J zjz1_rmQn}b5O)^|6SiPak7vxS4i5XLW7)UoQGIQsteVrY+g^3-U*-n+vgJ+p1>SRf zwZ=sFvJD8cJDwDtuX(f_I7o$jw~Vv@+ZsMJH+@5bc{l1wQM~3EnACjvzwuD zU^q&*Hj!gQ59Hht$lj%7lsbJESwhq0?q-E0dnG-iiA!@D)bq*+lwD_F;P`mNTx)~1 zJ0p?sTNHLhK0)H_VK`n_;Fo#+cz0^8D$sbns*$O>lI4?B$aSTv>$|GN*-t8eWDb=t z%~JU$2C6J4x+8V(A=Y7d|JZ~?WJ=&&II9I_Y~h*M4b3n&?+sMXnl5L~ALQ*lPS~cG z@wc|Nucay9H^Jw9*@~{Rf3FyCZ}AVZpFP$CbjqgdCZDr|0$cd6!4|;-8K4 zbq(&}%XzPXed0&cZU?JTV{(r%KMLtiG?4QG>z_HFp?Hn}RGBCc`s+@NxDkSR%W`Ao zq!_Fl#PdVzDqvZYtyuJVK9-iwgw1ObajeA}obP)B=MKNZ!3_-&cfAvOgdRkxFISPV zP*G$)7>F`SV}T=)Xw%1_{+W!B+gZ_7-Fmd{B%(y-t`avA;xCun7pDmQ! zd0zgzpZu{#7_v-rQE^rSgys?GchCL*{WsGq@Q&lKcod%$E@SGOAu3&XkEgQ_Hly~= z_)B;bKk@gsG;@xS#|Bua(UEn|*%mtV4nmQ_5i;SC>(S@f%eqh450;v~inAl;-(*7rjj!`JAx>oXQ;U(j;bH?+qq-^!~8ecMkh^sRk+**E%8 z6JLoW(>_o*!se_VZ_`%V^YjS$Y8De2#+^s8Olc_m?=xh*?2oJ;{ZX^xcIc)XvA>dy z+np4(o+)_Velo7zZi*A%Wx&p6AF!GCf^BR_4b#5H8DJg^o*%|@Hm^}+QxoL$3q^sG zNvN@?Gr|r8@N91qChis(S3DK{1?P}94skdVAG0ttc?oCgxm(Qd<;k)GD!FeEeP-}x_{p1J0Iq|@hHf5 z<7~L^LWv~b@gkwVJ)6>ePpU`yj`qFoo7&W}zfMo{hRxM-Sv}6$xe<$EZ%ou|!k)gp z0@J2CSQ`<6VGgCrV zXbGWp&ywAC7dyKfuhwSW*ZvFNpYL9;_p9@oIp^~^=gc|R=X!Ww*X6w`3Zo|#BFJck zVyZvJudW2n7GmNg6DGOq@W~AWRH5k@cRd}W=d0nPYJ!_vI|jNmqC5NJe|2c0x)4y& zH5+A4NT;m@TFN|bA>&>4v-h@8Wu1-Au2#_gMeHN}PQmdDop?~+5d3X@pox zn3qgzQu%j_*V0O*K+7C-6noX57K%2S|5q62RJW39j){hTqZ6N}v7XN|UfbI!9`gcq z3zAc~+D7Sx%sZ0Tk#V4&eX=Inu#00*qq#@hVt)$1;!jKZs%g`mI@;->r5$5*wArzp z*8WRIYqx7?OLihZzXxpl6|(k0f#(x->k z(&v4#rH8l1@~DO?A!1-W1_8btyC%3 z({D3e(AOsh9tg$9(FTNe)IojN1#^^6STsX}1n*=d)t4YS)CbAFZCG9HhLutW62C0L z7aC7|{^SLg{Lq30{o4^SLCGY$B#ijJ1N}=4upi<^zohxo-Oj~S<*lb<5hlv*(oWfL z*~WQmqoQBsRQ@oEZr*R8d&^~Xqq~kyPRbzd+Z6hM?aI{2nkxG(+z&g0>raKa|Aj;UKpq2ePXg04mlmC+VdEp$4)Cc6ZY!A8ZUr6%v>;pD&FFN+qOjTPb zwy}Yd?)%f)4Qfi7&i?ttc$#~V=$z~06>ujV@Fn?w<_cC#Bpb^m;3Sqr| z@go!G*ekeyi3=Go)KP9kI&J58M4UoR!K*A&FMm@-x(3y;RDZO2K1Y5qxZ%#jtgt3&Q3#RDQpi9=6c*wTqhyRM8!|^bbLMY zac(EmDOVenj4@GR5h!PIGHLJ1>C<)GYxAN^G`!OW7e)(Jt(k(fd96h%%5$!JgEm29Oy!a{cFkCm_Yfr`*1#N zU&^PGbVz%UjD73r(8z^kxaLK75)66)Yl5qfW=z2w~SJ&2g~8l6U_r(+RV;d zVD8Mmnx;s(WGpsHa}Pse-*MINTe-e>mR1UAgS2*$MLOcCk#7HKmHyDdY+r5_cOQdT zf0+Hk0lXG|siCB94WwUUp)w~=YGQq=eUkz`?prW;t`=UW$}sllVoX}SW zHE|9|-6O}=zgw}>dj~S@k09&PBj|Q)g08Ryx_A#{mDC_}eI_!txnlc+41rdEUw)x5(mZq}t+Oy54+`ZXcEC z-o^{`IIj32f7GNApRlg+*iJ4szJhpTn^CM@%C<{3`xM_<#II8! z=6ocGoB8(*cMyU*Y6NAzGt1o+L0Y4QRQZ!e`s1C&+-s)J+;6AO+-sv;dUMe#)dX0j z%{h=9f`!T}oJX|s59S^h{o^Ckii7$JqRW1Z=n@Lii~rC1N}w}-?NnW3qq~bO^xThg z5OU<`)SyI9SuC7-yI{!BcDRlA#K=CO2-?vTAx=FprD6zX$J=9Z*LtiTQ-;*0IBf6R z0G*37a;|hoUS|g!u879bnp~804Z`Vp{y3AHi_+u>oUuCMbgqCCh8z?uVP*$cF7sgd+7<4*M8T+IpobmMjz>usL& zpF2is31%MTp?GTf-b&9-s_E%(0=KAipw@?BWJ__Oj*JF$e&&kat(kB+XoTw*ZSdai zhS3S_80%(*@~c27Z@$5}-cJ$e7Xj~>7mzzlgyX3i*iX{aFOKPS%g9`=Q0DM8DJbhG zzXLWaZFp><^>dA+?J!Z=Z06}KFQ#4mj$~#SX_u{xwii0krXmF;KC#iv3^jfHlaZWe zTg7IMMPqTK~;r4z@f?G{1Otbn8|w@St-kc>)=ls(HLtzDp%LN;opF28B3if;;4 zK5;^&I!qYfUo9xJq6Ot2oa6E7AtCrsU-pA#a9)lJL~S37cr=IeJi^&-kXdO!shmbO z@%m{3Ma!A5`X0D@2~baAUkuT^?DMU`Oizzz&s;VSf(LpY@!hM7SUV$d~UbbM=WIX~oUnB5< z8lxUK!S|sCo(+K*Hn0sYUH#GPc^f*dY@**D#nY2C)(1+pRP|dtl@8a^u?7VdP7b32 zqm_(oUllCk+HXl5a~se^UyqEZf>EuMx7>mDvp$m*rRBPla;{Top=cZPmJAxMF{UBM zQxMpIh~f#kRGfb ztxqJc``l0RyorKi)ii01jAk84r1%y!ebtai>6glA?`6&hZYU;Ge+5auaJ~cd^n`81 zm&PbK3Nj2B%W|f;1+LQoPemYnrvib;tr#EN0M$@nQa~fr*NuqEu|adP0ZZ6dv*L*r z>&{dn^=crtA_+Urdm^iQ3HE&Ig#D9%gPkuRf43Y3SAoJkJ8-m@3di~<;P|aOC>a!p zlRcFZ`k_Mdgt3%29WSrRX2*>Bwp!m-i99=LEMa3q3U7;# zInmxt|end=W1do2cR3&en)VQ`-6g#Kv_aAJa^ zV@)Lb25?@G*oI!7a`bT2!@-(_E*D#2zo`j!{QL=O<|m&9wHMi_?Vg5gty+3%*G{dh zW4HWh;<#-wHAgc~_PLDi))?uwK}R@Om)=kKT37yk-56?*LdaU2eG3B-YS<;cBx z0SB{_U02SgO{eQ|8Sz=L<2ep9wkZkK58-IU#SF-w})G#mgVaUxVfj}qrW@g#`C+m^LJx;ds!vi@>%}2xWJudrTd-~cnk(SvLiA4 zqY?}+WqGl*5+lYm!E0zMyxw@iyR{J`ozmf}83I3p8UaqB7`4d>ffFYD*YuM&*w zp@#C4LWsslEa&6?hyUG!|K~GMb50OyPBV@(iWo-5K}HT^HzR|wjj@Tbj*-ZSW5h5b z7-0;N5yfSSE`&c1;c-HEybvBYg!fVLJSv_~#q+9oe-)pHiqA*I=cVFfReYW* uK3^4|w~EJ6@pY*9dQ^N}DjrkC*Qw&`Rq=JJczhM#2NmBB72g*X&-Gtk@nA6k literal 0 HcmV?d00001 diff --git a/docs/examples/xarray_workflow.py b/docs/examples/xarray_workflow.py index 0801884ec..9f06dc030 100644 --- a/docs/examples/xarray_workflow.py +++ b/docs/examples/xarray_workflow.py @@ -14,28 +14,35 @@ # name: python3 # --- +# %% [markdown] # %% [markdown] # # Gridded Data with xarray # # Download this notebook: {nb-download}`xarray_workflow.ipynb` # -# Climate and environmental data rarely arrive as a tidy matrix. They arrive as -# labelled [xarray](https://docs.xarray.dev/) objects: a temperature field over -# latitude and longitude, a covariate such as elevation, and gaps where a sensor -# failed or a cloud covered the scene. A GP in GPJax, on the other hand, consumes a -# [`Dataset`](#gpjax.dataset.Dataset) of flat inputs $\mathbf{X} \in \mathbb{R}^{N -# \times D}$ and outputs $\mathbf{y} \in \mathbb{R}^{N \times 1}$. +# Climate data usually arrive as labelled [xarray](https://docs.xarray.dev/) +# objects, read from netCDF files: a temperature field over latitude and +# longitude, often with gaps where there were no observations. A GP in GPJax, on +# the other hand, consumes a [`Dataset`](#gpjax.dataset.Dataset) of flat inputs +# $\mathbf{X} \in \mathbb{R}^{N \times D}$ and outputs $\mathbf{y} \in +# \mathbb{R}^{N \times 1}$. The [`gpjax.xarray`](../reference/xarray.md) module +# converts between the two at the edges of a workflow. # -# The [`gpjax.xarray`](../reference/xarray.md) module converts between the two at the -# edges of a workflow. In this notebook we +# In this notebook we fill gaps in a global temperature field. To know how good +# the filled values are, we start from a field that is complete, remove cells +# ourselves, and compare the GP with the cells that we removed. We # -# 1. flatten a gappy, labelled field into a `Dataset` with -# [`from_xarray`](#gpjax.xarray.from_xarray), +# 1. flatten a gappy, labelled global field into a `Dataset` with +# [`from_xarray`](#gpjax.xarray.from_xarray), and put the inputs on the sphere +# with the [`UnitSphere`](#gpjax.xarray.UnitSphere) transform, # 2. fit a GP exactly as we would on any other `Dataset`, -# 3. build inputs for a finer prediction grid with -# [`GridSpec.inputs_for`](#gpjax.xarray.GridSpec.inputs_for), and -# 4. map the predictions, and joint posterior samples, back onto that grid with -# [`GridSpec.to_xarray`](#gpjax.xarray.GridSpec.to_xarray). +# 3. predict the mean and variance on a finer grid with +# [`GridSpec.predict`](#gpjax.xarray.GridSpec.predict), which works in chunks, +# 4. draw joint posterior samples with +# [`GridSpec.inputs_for`](#gpjax.xarray.GridSpec.inputs_for) and +# [`GridSpec.to_xarray`](#gpjax.xarray.GridSpec.to_xarray), to get the global +# mean temperature with its uncertainty, and +# 5. show what goes wrong when cells are missing *because of* their values. # # The module needs the optional extra: `pip install "gpjax[xarray]"`. @@ -47,6 +54,8 @@ logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR) # %% +from pathlib import Path + from jax import config import jax.numpy as jnp import jax.random as jr @@ -59,83 +68,98 @@ with install_import_hook("gpjax", "beartype.beartype"): import gpjax as gpx - from gpjax.xarray import from_xarray + from gpjax.parameters import val + from gpjax.xarray import ( + UnitSphere, + from_xarray, + ) key = jr.key(42) +rng = np.random.default_rng(0) gpx.plotting.use_style() # %% [markdown] -# ## A synthetic temperature field +# ## A complete temperature field # -# We simulate near-surface temperature on a regional latitude-longitude grid. It -# cools towards the pole and with elevation (a lapse rate of roughly 6.5 K per -# kilometre), with a smooth large-scale anomaly on top. Elevation is a separate -# variable over the same grid. The data are synthetic, so the notebook needs no -# download and every number in it can be checked against the truth. - +# We use the [NCEP-NCAR Reanalysis 1](https://psl.noaa.gov/data/gridded/data.ncep.reanalysis.html) +# (Kalnay et al., 1996). A reanalysis combines a weather model with the +# observations, so it has a value in every grid cell. The file holds the 2024 +# annual-mean near-surface air temperature as an *anomaly*: the difference from +# the 1991–2020 mean of the same cell. Anomalies remove the large, fixed +# differences between the equator and the poles, so what remains is the +# signal of interest, and it varies smoothly over large distances. +# +# NCEP-NCAR Reanalysis 1 data provided by the NOAA PSL, Boulder, Colorado, USA, +# from their website at . # %% -def elevation_at(lat, lon): - """A single mountain range, in metres.""" - return 2500.0 * np.exp(-(((lon - 12.0) / 4.0) ** 2) - ((lat - 47.0) / 3.0) ** 2) - - -def temperature_at(lat, lon, elevation): - """Temperature in kelvin: latitude gradient, lapse rate and an anomaly.""" - anomaly = 1.5 * np.sin(lon / 3.0) * np.cos(lat / 4.0) - return 290.0 - 0.6 * (lat - 40.0) - 0.0065 * elevation + anomaly - - -def regional_field(lats, lons) -> xr.Dataset: - lat_grid, lon_grid = np.meshgrid(lats, lons, indexing="ij") - elevation = elevation_at(lat_grid, lon_grid) - return xr.Dataset( - { - "t2m": ( - ("lat", "lon"), - temperature_at(lat_grid, lon_grid, elevation), - {"units": "K", "long_name": "2 m air temperature"}, - ), - "elevation": (("lat", "lon"), elevation, {"units": "m"}), - }, - coords={ - "lat": ("lat", lats, {"units": "degrees_north"}), - "lon": ("lon", lons, {"units": "degrees_east"}), - }, - ) +nc_candidates = [ + Path("docs/examples/data/ncep_air_anomaly_2024.nc"), + Path("data/ncep_air_anomaly_2024.nc"), +] +nc_path = next(path for path in nc_candidates if path.exists()) +reanalysis = xr.open_dataset(nc_path) +reanalysis + +# %% [markdown] +# The native grid is 2.5°. We fit on a 5° grid, which is every second grid point, +# and predict back on the 2.5° grid, so the predictions are also tested at points +# the model never saw. +# %% +truth_fine = reanalysis["tas_anomaly"].astype(float) +truth = truth_fine.isel(lat=slice(1, None, 2), lon=slice(1, None, 2)) +print(f"5° grid: {dict(truth.sizes)}, 2.5° grid: {dict(truth_fine.sizes)}") -coarse = regional_field(np.linspace(40.0, 54.0, 12), np.linspace(2.0, 22.0, 16)) +anomaly_style = dict(cmap="RdBu_r", vmin=-4.0, vmax=4.0) +error_style = dict( + cmap="PuOr_r", vmin=-2.0, vmax=2.0, cbar_kwargs={"label": "Error [K]"} +) +truth_fine.plot(figsize=(7, 3.5), **anomaly_style) +plt.title("2024 anomaly, NCEP-NCAR Reanalysis 1") +plt.show() # %% [markdown] -# Real observations have holes. We knock out a block of cells, as a cloud would, -# and add a little measurement noise to the rest. +# A 5° cell near a pole is much smaller than one at the equator, so a global mean +# weights each cell by the cosine of its latitude. + # %% -key, noise_key = jr.split(key) -noise = 0.2 * np.asarray(jr.normal(noise_key, coarse["t2m"].shape)) -observed = coarse.copy(deep=True) -observed["t2m"] = observed["t2m"] + noise -observed["t2m"].attrs = coarse["t2m"].attrs -observed["t2m"][4:7, 9:13] = np.nan - -observed["t2m"].plot(cmap="coolwarm") -plt.title("Observed temperature (gaps in white)") -plt.show() +def global_mean(field: xr.DataArray) -> xr.DataArray: + """Area-weighted mean over the cells that have a value.""" + weights = np.cos(np.deg2rad(field["lat"])) + return field.weighted(weights).mean(["lat", "lon"]) + + +print(f"True global mean anomaly: {float(global_mean(truth)):.2f} K") # %% [markdown] -# ## From labelled data to a `Dataset` +# ## Cells missing at random # -# `from_xarray` takes the target variable and the inputs we want the GP to depend -# on. Inputs can be coordinates (`lat`, `lon`) or other data variables -# (`elevation`), and the columns of $\mathbf{X}$ follow the order we list them in. -# Cells where the target or any input is NaN are dropped by default, and the -# returned `GridSpec` records which ones. +# First we keep a random 25% of the 5° cells. Which cells are missing has nothing +# to do with their values. Statisticians call this *missing completely at +# random*. # %% -inputs = ["lat", "lon", "elevation"] -data, spec = from_xarray(observed, target="t2m", inputs=inputs) +kept_at_random = rng.uniform(size=truth.shape) < 0.25 +observed = truth.where(kept_at_random).rename("tas").to_dataset() +print(f"{int(kept_at_random.sum())} of {truth.size} cells kept") + +# %% [markdown] +# ### Inputs on the sphere +# +# Latitude and longitude in degrees are not good GP inputs for a global field. A +# 5° cell is about 555 km wide at the equator but only about 24 km wide next to a +# pole, and longitude 177.5°E is next to 177.5°W. The +# [`UnitSphere`](#gpjax.xarray.UnitSphere) transform replaces `lat` and `lon` with +# the three coordinates of a point on the unit sphere. A stationary kernel on these +# coordinates uses the chord distance through the Earth, so it has no seam at the +# antimeridian and no distortion at the poles. Its lengthscale is in Earth radii. +# %% +data, spec = from_xarray( + observed, target="tas", inputs=["lat", "lon"], transforms=[UnitSphere()] +) print(data) print(spec) @@ -144,119 +168,240 @@ def regional_field(lats, lons) -> xr.Dataset: # The `GridSpec` stays with us, outside the model, until we want labelled output # again. # -# ## Fitting the model +# ### Fitting the model # -# Temperature varies over hundreds of kilometres in latitude and longitude but -# over hundreds of metres in elevation, so we give the RBF kernel one lengthscale -# per input. A constant mean absorbs the ~285 K offset. +# We use a Matérn-3/2 kernel, which gives rougher fields than an RBF kernel, as +# temperature anomalies are. A constant mean absorbs the global warming signal. + # %% -prior = gpx.gps.Prior( - mean_function=gpx.mean_functions.Constant(jnp.array([285.0])), - kernel=gpx.kernels.RBF(lengthscale=jnp.array([3.0, 3.0, 1000.0]), variance=25.0), -) -model = prior * gpx.likelihoods.Gaussian(obs_stddev=jnp.array(0.5)) - -model, history = gpx.fit_scipy( - model=model, - objective=lambda candidate, train_data: -gpx.objectives.conjugate_mll( - candidate, train_data - ), - train_data=data, - verbose=False, -) +def fit_model(train_data: gpx.Dataset): + prior = gpx.gps.Prior( + mean_function=gpx.mean_functions.Constant(jnp.array([0.5])), + kernel=gpx.kernels.Matern32(lengthscale=jnp.array(0.3), variance=1.0), + ) + model = prior * gpx.likelihoods.Gaussian(obs_stddev=jnp.array(0.1)) + model, _ = gpx.fit_scipy( + model=model, + objective=lambda candidate, train_data: ( + -gpx.objectives.conjugate_mll(candidate, train_data) + ), + train_data=train_data, + verbose=False, + ) + return model + + +earth_radius_km = 6371.0 +model = fit_model(data) +lengthscale_km = float(val(model.prior.kernel.lengthscale)) * earth_radius_km +print(f"Lengthscale: {lengthscale_km:.0f} km") # %% [markdown] -# ## Predicting on a finer grid +# ### Predicting on the finer grid +# +# `spec.predict` gives the predictive mean and variance on any grid that holds the +# same inputs, encoded exactly as in training. Here the grid is the 2.5° grid of +# the reanalysis. The function we pass maps a block of inputs to a distribution; +# we use a diagonal covariance, because we need only the variance of each cell. # -# `spec.inputs_for` builds the prediction inputs for any grid that holds the same -# input variables, encoded exactly as in training. Here we predict on a grid four -# times finer in each direction, including the cells that were missing from the -# observations. We pass the likelihood's predictive distribution, so the variance -# includes observation noise. +# `spec.predict` sends the cells to this function in chunks of `chunk_size`, so +# the memory use stays the same for a finer or larger grid. If the grid holds +# Dask arrays, the result is lazy, and each Dask block is predicted only when it is +# computed or written with `to_netcdf`. # %% -fine = regional_field(np.linspace(40.0, 54.0, 45), np.linspace(2.0, 22.0, 61)) -test_inputs, test_spec = spec.inputs_for(fine[["elevation"]]) - posterior = model.condition(data) -predictive = model.likelihood(posterior(test_inputs)) -prediction = test_spec.to_xarray(predictive) +prediction = spec.predict( + lambda x: model.likelihood(posterior(x, covariance="diagonal")), + truth_fine.to_dataset(), + chunk_size=2048, +) prediction # %% [markdown] -# The result is a labelled `xr.Dataset` on the fine grid, with the target's -# attributes carried over. The variance is in $\mathrm{K}^2$. Everything xarray -# offers, from plotting to `to_netcdf`, works on it directly. +# Because the field is complete, we can score every prediction. We compare with a +# simple baseline, the mean of the kept cells, and check the uncertainty: about +# 95% of the true values should be within two predictive standard deviations. + # %% -fig, (mean_ax, std_ax, error_ax) = plt.subplots(1, 3, figsize=(15, 4)) -prediction["t2m_mean"].plot(ax=mean_ax, cmap="coolwarm") -mean_ax.set_title("Predictive mean") -np.sqrt(prediction["t2m_variance"]).plot(ax=std_ax, cmap="viridis") -std_ax.set_title("Predictive standard deviation") -(prediction["t2m_mean"] - fine["t2m"]).plot(ax=error_ax, cmap="RdBu_r", center=0.0) -error_ax.set_title("Error against the true field") -for ax in (mean_ax, std_ax, error_ax): - ax.add_patch( - plt.Rectangle( - (observed.lon[9], observed.lat[4]), - float(observed.lon[12] - observed.lon[9]), - float(observed.lat[6] - observed.lat[4]), - fill=False, - linestyle="--", - ) +def score(prediction: xr.Dataset, observed: xr.Dataset) -> None: + error = prediction["tas_mean"] - truth_fine + baseline_error = float(observed["tas"].mean()) - truth_fine + within = abs(error) < 2 * np.sqrt(prediction["tas_variance"]) + print(f"RMSE, GP: {float(np.sqrt((error**2).mean())):.2f} K") + print( + f"RMSE, mean of kept cells: {float(np.sqrt((baseline_error**2).mean())):.2f} K" ) -plt.show() + print(f"Within 2 sd: {float(within.mean()):.0%}") + + +def plot_infill(prediction: xr.Dataset, observed: xr.Dataset) -> None: + fig, axes = plt.subplots(1, 3, figsize=(15, 3.5), sharey=True) + observed["tas"].plot(ax=axes[0], **anomaly_style) + axes[0].set_title("Kept cells") + prediction["tas_mean"].plot(ax=axes[1], **anomaly_style) + axes[1].set_title("Predictive mean") + (prediction["tas_mean"] - truth_fine).plot(ax=axes[2], **error_style) + axes[2].set_title("Predictive mean minus truth") + for ax in axes[1:]: + ax.set_ylabel("") + plt.show() + + +score(prediction, observed) +plot_infill(prediction, observed) # %% [markdown] -# The standard deviation grows inside the dashed box where observations were -# missing, and the error stays small across the mountain range because elevation is -# an input. +# From a quarter of the cells, the GP recovers the large-scale pattern, including +# the strong warmth over the Arctic. The errors are largest where the field +# changes over short distances, and the uncertainty is about right. # -# ## Joint samples and regional averages +# ### The global mean, with joint samples # -# The mean and variance describe each cell on its own. Many questions are about -# several cells together, such as the average temperature over the Alpine box -# $\mathcal{R}$. Its variance depends on the covariance between the cells, +# The global mean anomaly fills each missing cell and averages. Its variance +# depends on the covariance between the filled cells, # # $$ -# \operatorname{Var}\Big[\tfrac{1}{|\mathcal{R}|} \sum_{i \in \mathcal{R}} f_i\Big] -# = \tfrac{1}{|\mathcal{R}|^2} \sum_{i, j \in \mathcal{R}} \operatorname{Cov}[f_i, f_j], -# $$ (eq-xarray-regional-variance) +# \operatorname{Var}\Big[\sum_{i} w_i f_i\Big] +# = \sum_{i, j} w_i w_j \operatorname{Cov}[f_i, f_j], +# $$ (eq-xarray-global-variance) # -# which the per-cell variances alone cannot give. Passing `num_samples` to -# `to_xarray` draws from the joint predictive distribution instead, and returns the -# draws with a leading `sample` dimension. Averaging each draw over the region gives -# samples of the regional mean. +# which the per-cell variances alone cannot give. For this we need the full +# covariance, so we build all the prediction inputs on the 5° grid at once with +# `spec.inputs_for`. Passing `num_samples` to `to_xarray` then draws from the joint +# distribution of the field, and returns the draws with a leading `sample` +# dimension. `combine_first` keeps the kept cells and takes each missing cell from +# a draw. + # %% -latent = posterior(test_inputs) # the field itself, without observation noise +def global_mean_samples(posterior, spec, observed: xr.Dataset, key) -> xr.DataArray: + test_inputs, test_spec = spec.inputs_for(truth.to_dataset()) + samples = test_spec.to_xarray(posterior(test_inputs), num_samples=500, key=key) + return global_mean(observed["tas"].combine_first(samples["tas"])) + + +def report_global_mean(means: xr.DataArray, observed: xr.Dataset) -> None: + print(f"Truth: {float(global_mean(truth)):.2f} K") + print(f"Mean of kept cells: {float(global_mean(observed['tas'])):.2f} K") + print( + f"Gaps filled by GP: {float(means.mean()):.2f}" + f" ± {2 * float(means.std('sample')):.2f} K (2 sd)" + ) + + key, sample_key = jr.split(key) -samples = test_spec.to_xarray(latent, num_samples=500, key=sample_key) +means = global_mean_samples(posterior, spec, observed, sample_key) +report_global_mean(means, observed) -alps = dict(lat=slice(45.0, 49.0), lon=slice(8.0, 16.0)) -regional_mean = samples["t2m"].sel(**alps).mean(["lat", "lon"]) -true_regional_mean = float(fine["t2m"].sel(**alps).mean()) +# %% [markdown] +# With random gaps, even the mean of the kept cells is close to the truth, and the +# GP gives the global mean with an interval that holds it. +# +# ## Cells missing because of their values +# +# Real gaps are rarely random. A satellite cannot see the surface through cloud, +# a station fails in extreme weather, and the polar regions, which warm fastest, +# have the fewest observations. When the chance that a cell is missing depends on +# the value that is missing, the data are *missing not at random*. +# +# We simulate this. We again aim to keep a quarter of the cells, but now the +# warmer a cell's anomaly, the less likely it is to be kept. -# The same latent distribution, but treating the cells as independent. -latent_variance = test_spec.to_xarray(latent)["t2m_variance"].sel(**alps) -joint_std = float(regional_mean.std("sample")) -naive_std = float(np.sqrt(latent_variance.sum()) / latent_variance.size) +# %% +standardised = ((truth - truth.mean()) / truth.std()).values +odds = np.exp(-1.5 * standardised) +keep_probability = np.clip(0.25 * odds / odds.mean(), 0.0, 1.0) +kept_selectively = rng.uniform(size=truth.shape) < keep_probability +selective = truth.where(kept_selectively).rename("tas").to_dataset() +print(f"{int(kept_selectively.sum())} of {truth.size} cells kept") -print(f"Regional mean: {float(regional_mean.mean()):.2f} K (truth {true_regional_mean:.2f} K)") -print(f"Standard deviation from joint samples: {joint_std:.3f} K") -print(f"Standard deviation if cells were independent: {naive_std:.3f} K") +# %% [markdown] +# The steps are the same as before. + +# %% +data_selective, spec_selective = from_xarray( + selective, target="tas", inputs=["lat", "lon"], transforms=[UnitSphere()] +) +model_selective = fit_model(data_selective) +posterior_selective = model_selective.condition(data_selective) +prediction_selective = spec_selective.predict( + lambda x: model_selective.likelihood(posterior_selective(x, covariance="diagonal")), + truth_fine.to_dataset(), + chunk_size=2048, +) + +score(prediction_selective, selective) +plot_infill(prediction_selective, selective) + +key, sample_key = jr.split(key) +means_selective = global_mean_samples( + posterior_selective, spec_selective, selective, sample_key +) +report_global_mean(means_selective, selective) # %% [markdown] -# Treating the cells as independent understates the uncertainty in the regional -# average by roughly an order of magnitude, because neighbouring cells tend to be -# wrong in the same direction. Against the joint standard deviation the true -# regional mean is a plausible outcome; against the independent one it would look -# like a many-sigma surprise. Joint -# samples keep that correlation, which is why `to_xarray` refuses to draw samples -# from a distribution that only holds marginal variances (as returned by -# `posterior(test_inputs, covariance="diagonal")`). +# The kept cells are mostly the cold ones, so their mean is far too cold. The GP +# removes part of this bias, because it fills each gap from its neighbours and +# the warm regions still have some kept cells. But the result is still too cold, +# and the interval does not hold the truth: the GP is confidently wrong. +# +# The reason is that a GP conditions only on the values it sees. Its prior has one +# constant mean, which it learns from the kept cells, so the cold kept cells pull +# that mean down. It also has no way to know that the missing cells are warm, +# because nothing in its inputs says so. Its uncertainty describes the spread of +# values that are *consistent with the kept cells*, not the error of the +# selection. +# +# In real data we cannot see this bias, because we do not have the missing values. +# Useful steps are: +# +# - Add inputs that explain why cells are missing, such as cloud fraction or a +# covariate that is observed everywhere. If the chance of a gap depends only on +# the inputs, the gaps are *missing at random* given those inputs, and the GP +# can correct for them. +# - Model the observation process together with the field. +# - Test how sensitive the result is to different assumptions about the missing +# values, as we did here with a complete field. +# +# ## Other inputs and seasonal cycles +# +# The same workflow takes more inputs. Other transforms encode them: +# [`Cyclic`](#gpjax.xarray.Cyclic) encodes a periodic input as a point on a +# circle, and [`Standardise`](#gpjax.xarray.Standardise) scales an input to zero +# mean and unit standard deviation over the training cells. Datetime inputs are +# days since their first timestamp, so a period of 365.25 gives the seasonal +# cycle of a monthly record: +# +# ```python +# data, spec = from_xarray( +# monthly, +# target="tas", +# inputs=["lat", "lon", "time", "cloud_fraction"], +# transforms=[ +# UnitSphere(), +# Cyclic("time", 365.25), +# Standardise(["cloud_fraction"]), +# ], +# ) +# spec.columns +# # ('sphere_x', 'sphere_y', 'sphere_z', 'time_sin', 'time_cos', 'cloud_fraction') +# ``` +# +# The transforms run in order, and `spec.columns` names the columns of +# $\mathbf{X}$ that they produce. The spec applies the same fitted transforms to +# every grid that we predict on. +# +# ## References +# +# Kalnay, E., Kanamitsu, M., Kistler, R., Collins, W., Deaven, D., Gandin, L., +# Iredell, M., Saha, S., White, G., Woollen, J., Zhu, Y., Chelliah, M., Ebisuzaki, +# W., Higgins, W., Janowiak, J., Mo, K. C., Ropelewski, C., Wang, J., Leetmaa, A., +# Reynolds, R., Jenne, R. and Joseph, D. (1996). The NCEP/NCAR 40-year reanalysis +# project. *Bulletin of the American Meteorological Society*, 77(3), 437–471. +# [doi:10.1175/1520-0477(1996)077<0437:TNYRP>2.0.CO;2](https://doi.org/10.1175/1520-0477(1996)077%3C0437:TNYRP%3E2.0.CO;2) # # ## System configuration From 6744b6c3c8303eeea997d99a001d18e04b9e5a46 Mon Sep 17 00:00:00 2001 From: Thomas Pinder Date: Sun, 4 Oct 2026 12:53:22 +0000 Subject: [PATCH 3/3] docs(examples): split the xarray introduction from the infilling example Working with Gridded Data (examples/xarray_workflow, same URL) is the synthetic introduction to gpjax.xarray again, now under Getting started. It uses Standardise and GridSpec.predict, and shows UnitSphere and Cyclic. Infilling Global Surface Temperature (examples/infilling_surface_temperature) holds the reanalysis study, under Applied modelling. It also drops an empty markdown cell at the top of the notebook. Co-Authored-By: Claude Opus 5.5 --- CHANGELOG.md | 15 +- .../examples/infilling_surface_temperature.py | 379 +++++++++++++++ docs/examples/xarray_workflow.py | 458 +++++++----------- docs/index.md | 3 +- 4 files changed, 574 insertions(+), 281 deletions(-) create mode 100644 docs/examples/infilling_surface_temperature.py diff --git a/CHANGELOG.md b/CHANGELOG.md index d6dbee5ad..ce8746280 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -22,12 +22,15 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 predictive mean and variance as an `xr.Dataset`, with at most `chunk_size` cells in memory at once and one compilation. A Dask-backed grid gives a lazy result that is predicted block by block. -- **The *Gridded Data with xarray* example now uses real data.** It reads a - netCDF file of the 2024 temperature anomaly from the NCEP-NCAR Reanalysis 1, - which is complete, and removes cells to test the infill against the truth. - Cells removed at random are filled well, and joint samples give a global mean - whose interval holds the truth. Cells removed because they are warm show how a - GP is biased, and overconfident, when data are missing not at random. +- **New example: *Infilling Global Surface Temperature*.** It reads a netCDF + file of the 2024 temperature anomaly from the NCEP-NCAR Reanalysis 1, which is + complete, and removes cells to test the infill against the truth. Cells removed + at random are filled well, and joint samples give a global mean whose interval + holds the truth. Cells removed because they are warm show how a GP is biased, + and overconfident, when data are missing not at random. +- **The xarray introduction is now *Working with Gridded Data*, under Getting + started.** It uses `Standardise` and `GridSpec.predict`, and shows + `UnitSphere` and `Cyclic`. Its URL is unchanged. ### Changed diff --git a/docs/examples/infilling_surface_temperature.py b/docs/examples/infilling_surface_temperature.py new file mode 100644 index 000000000..683371677 --- /dev/null +++ b/docs/examples/infilling_surface_temperature.py @@ -0,0 +1,379 @@ +# --- +# jupyter: +# jupytext: +# cell_metadata_filter: -all +# custom_cell_magics: kql +# text_representation: +# extension: .py +# format_name: percent +# format_version: '1.3' +# jupytext_version: 1.19.1 +# kernelspec: +# display_name: Python 3 +# language: python +# name: python3 +# --- + +# %% [markdown] +# # Infilling Global Surface Temperature +# +# Download this notebook: {nb-download}`infilling_surface_temperature.ipynb` +# +# Temperature records have gaps: over the poles, over parts of Africa and the +# Southern Ocean, and, for satellites, wherever there is cloud. Climate +# scientists fill these gaps to map the field and to compute the global mean +# temperature. In this notebook we fill gaps with a GP, and we test the result. +# +# To know how good the filled values are, we start from a field that is complete, +# remove cells ourselves, and compare the GP with the cells that we removed. We +# +# 1. fill gaps that are random, map the field, and get the global mean +# temperature with its uncertainty from joint posterior samples, and +# 2. show what goes wrong when cells are missing *because of* their values. +# +# The notebook uses the [`gpjax.xarray`](../reference/xarray.md) module, which +# [Working with Gridded Data](xarray_workflow.py) introduces. It needs the +# optional extra: `pip install "gpjax[xarray]"`. + +# %% tags=["remove-cell"] +import logging + +# The build host may not have the ledger fonts (Public Sans, Spectral); hide +# matplotlib's font fallback messages. +logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR) + +# %% +from pathlib import Path + +from jax import config +import jax.numpy as jnp +import jax.random as jr +from jaxtyping import install_import_hook +import matplotlib.pyplot as plt +import numpy as np +import xarray as xr + +config.update("jax_enable_x64", True) + +with install_import_hook("gpjax", "beartype.beartype"): + import gpjax as gpx + from gpjax.parameters import val + from gpjax.xarray import ( + UnitSphere, + from_xarray, + ) + +key = jr.key(42) +rng = np.random.default_rng(0) +gpx.plotting.use_style() + +# %% [markdown] +# ## A complete temperature field +# +# We use the [NCEP-NCAR Reanalysis 1](https://psl.noaa.gov/data/gridded/data.ncep.reanalysis.html) +# (Kalnay et al., 1996). A reanalysis combines a weather model with the +# observations, so it has a value in every grid cell. The file holds the 2024 +# annual-mean near-surface air temperature as an *anomaly*: the difference from +# the 1991–2020 mean of the same cell. Anomalies remove the large, fixed +# differences between the equator and the poles, so what remains is the +# signal of interest, and it varies smoothly over large distances. +# +# NCEP-NCAR Reanalysis 1 data provided by the NOAA PSL, Boulder, Colorado, USA, +# from their website at . + +# %% +nc_candidates = [ + Path("docs/examples/data/ncep_air_anomaly_2024.nc"), + Path("data/ncep_air_anomaly_2024.nc"), +] +nc_path = next(path for path in nc_candidates if path.exists()) +reanalysis = xr.open_dataset(nc_path) +reanalysis + +# %% [markdown] +# The native grid is 2.5°. We fit on a 5° grid, which is every second grid point, +# and predict back on the 2.5° grid, so the predictions are also tested at points +# the model never saw. + +# %% +truth_fine = reanalysis["tas_anomaly"].astype(float) +truth = truth_fine.isel(lat=slice(1, None, 2), lon=slice(1, None, 2)) +print(f"5° grid: {dict(truth.sizes)}, 2.5° grid: {dict(truth_fine.sizes)}") + +anomaly_style = dict(cmap="RdBu_r", vmin=-4.0, vmax=4.0) +error_style = dict( + cmap="PuOr_r", vmin=-2.0, vmax=2.0, cbar_kwargs={"label": "Error [K]"} +) +truth_fine.plot(figsize=(7, 3.5), **anomaly_style) +plt.title("2024 anomaly, NCEP-NCAR Reanalysis 1") +plt.show() + +# %% [markdown] +# A 5° cell near a pole is much smaller than one at the equator, so a global mean +# weights each cell by the cosine of its latitude. + + +# %% +def global_mean(field: xr.DataArray) -> xr.DataArray: + """Area-weighted mean over the cells that have a value.""" + weights = np.cos(np.deg2rad(field["lat"])) + return field.weighted(weights).mean(["lat", "lon"]) + + +print(f"True global mean anomaly: {float(global_mean(truth)):.2f} K") + +# %% [markdown] +# ## Cells missing at random +# +# First we keep a random 25% of the 5° cells. Which cells are missing has nothing +# to do with their values. Statisticians call this *missing completely at +# random*. + +# %% +kept_at_random = rng.uniform(size=truth.shape) < 0.25 +observed = truth.where(kept_at_random).rename("tas").to_dataset() +print(f"{int(kept_at_random.sum())} of {truth.size} cells kept") + +# %% [markdown] +# ### Inputs on the sphere +# +# Latitude and longitude in degrees are not good GP inputs for a global field. A +# 5° cell is about 555 km wide at the equator but only about 24 km wide next to a +# pole, and longitude 177.5°E is next to 177.5°W. The +# [`UnitSphere`](#gpjax.xarray.UnitSphere) transform replaces `lat` and `lon` with +# the three coordinates of a point on the unit sphere. A stationary kernel on these +# coordinates uses the chord distance through the Earth, so it has no seam at the +# antimeridian and no distortion at the poles. Its lengthscale is in Earth radii. + +# %% +data, spec = from_xarray( + observed, target="tas", inputs=["lat", "lon"], transforms=[UnitSphere()] +) +print(data) +print(spec) + +# %% [markdown] +# ### Fitting the model +# +# We use a Matérn-3/2 kernel, which gives rougher fields than an RBF kernel, as +# temperature anomalies are. A constant mean absorbs the global warming signal. + + +# %% +def fit_model(train_data: gpx.Dataset): + prior = gpx.gps.Prior( + mean_function=gpx.mean_functions.Constant(jnp.array([0.5])), + kernel=gpx.kernels.Matern32(lengthscale=jnp.array(0.3), variance=1.0), + ) + model = prior * gpx.likelihoods.Gaussian(obs_stddev=jnp.array(0.1)) + model, _ = gpx.fit_scipy( + model=model, + objective=lambda candidate, train_data: ( + -gpx.objectives.conjugate_mll(candidate, train_data) + ), + train_data=train_data, + verbose=False, + ) + return model + + +earth_radius_km = 6371.0 +model = fit_model(data) +lengthscale_km = float(val(model.prior.kernel.lengthscale)) * earth_radius_km +print(f"Lengthscale: {lengthscale_km:.0f} km") + +# %% [markdown] +# ### Predicting on the finer grid +# +# `spec.predict` gives the predictive mean and variance on the 2.5° grid of the +# reanalysis, in chunks of `chunk_size` cells. We use a diagonal covariance, +# because we need only the variance of each cell. + +# %% +posterior = model.condition(data) +prediction = spec.predict( + lambda x: model.likelihood(posterior(x, covariance="diagonal")), + truth_fine.to_dataset(), + chunk_size=2048, +) +prediction + +# %% [markdown] +# Because the field is complete, we can score every prediction. We compare with a +# simple baseline, the mean of the kept cells, and check the uncertainty: about +# 95% of the true values should be within two predictive standard deviations. + + +# %% +def score(prediction: xr.Dataset, observed: xr.Dataset) -> None: + error = prediction["tas_mean"] - truth_fine + baseline_error = float(observed["tas"].mean()) - truth_fine + within = abs(error) < 2 * np.sqrt(prediction["tas_variance"]) + print(f"RMSE, GP: {float(np.sqrt((error**2).mean())):.2f} K") + print( + f"RMSE, mean of kept cells: {float(np.sqrt((baseline_error**2).mean())):.2f} K" + ) + print(f"Within 2 sd: {float(within.mean()):.0%}") + + +def plot_infill(prediction: xr.Dataset, observed: xr.Dataset) -> None: + fig, axes = plt.subplots(1, 3, figsize=(15, 3.5), sharey=True) + observed["tas"].plot(ax=axes[0], **anomaly_style) + axes[0].set_title("Kept cells") + prediction["tas_mean"].plot(ax=axes[1], **anomaly_style) + axes[1].set_title("Predictive mean") + (prediction["tas_mean"] - truth_fine).plot(ax=axes[2], **error_style) + axes[2].set_title("Predictive mean minus truth") + for ax in axes[1:]: + ax.set_ylabel("") + plt.show() + + +score(prediction, observed) +plot_infill(prediction, observed) + +# %% [markdown] +# From a quarter of the cells, the GP recovers the large-scale pattern, including +# the strong warmth over the Arctic. The errors are largest where the field +# changes over short distances, and the uncertainty is about right. +# +# ### The global mean, with joint samples +# +# The global mean anomaly fills each missing cell and averages. Its variance +# depends on the covariance between the filled cells, +# +# $$ +# \operatorname{Var}\Big[\sum_{i} w_i f_i\Big] +# = \sum_{i, j} w_i w_j \operatorname{Cov}[f_i, f_j], +# $$ (eq-xarray-global-variance) +# +# which the per-cell variances alone cannot give. For this we need the full +# covariance, so we build all the prediction inputs on the 5° grid at once with +# `spec.inputs_for`. Passing `num_samples` to `to_xarray` then draws from the joint +# distribution of the field, and returns the draws with a leading `sample` +# dimension. `combine_first` keeps the kept cells and takes each missing cell from +# a draw. + + +# %% +def global_mean_samples(posterior, spec, observed: xr.Dataset, key) -> xr.DataArray: + test_inputs, test_spec = spec.inputs_for(truth.to_dataset()) + samples = test_spec.to_xarray(posterior(test_inputs), num_samples=500, key=key) + return global_mean(observed["tas"].combine_first(samples["tas"])) + + +def report_global_mean(means: xr.DataArray, observed: xr.Dataset) -> None: + print(f"Truth: {float(global_mean(truth)):.2f} K") + print(f"Mean of kept cells: {float(global_mean(observed['tas'])):.2f} K") + print( + f"Gaps filled by GP: {float(means.mean()):.2f}" + f" ± {2 * float(means.std('sample')):.2f} K (2 sd)" + ) + + +key, sample_key = jr.split(key) +means = global_mean_samples(posterior, spec, observed, sample_key) +report_global_mean(means, observed) + +# %% [markdown] +# With random gaps, even the mean of the kept cells is close to the truth, and the +# GP gives the global mean with an interval that holds it. +# +# ## Cells missing because of their values +# +# Real gaps are rarely random. A satellite cannot see the surface through cloud, +# a station fails in extreme weather, and the polar regions, which warm fastest, +# have the fewest observations. When the chance that a cell is missing depends on +# the value that is missing, the data are *missing not at random*. +# +# We simulate this. We again aim to keep a quarter of the cells, but now the +# warmer a cell's anomaly, the less likely it is to be kept. + +# %% +standardised = ((truth - truth.mean()) / truth.std()).values +odds = np.exp(-1.5 * standardised) +keep_probability = np.clip(0.25 * odds / odds.mean(), 0.0, 1.0) +kept_selectively = rng.uniform(size=truth.shape) < keep_probability +selective = truth.where(kept_selectively).rename("tas").to_dataset() +print(f"{int(kept_selectively.sum())} of {truth.size} cells kept") + +# %% [markdown] +# The steps are the same as before. + +# %% +data_selective, spec_selective = from_xarray( + selective, target="tas", inputs=["lat", "lon"], transforms=[UnitSphere()] +) +model_selective = fit_model(data_selective) +posterior_selective = model_selective.condition(data_selective) +prediction_selective = spec_selective.predict( + lambda x: model_selective.likelihood(posterior_selective(x, covariance="diagonal")), + truth_fine.to_dataset(), + chunk_size=2048, +) + +score(prediction_selective, selective) +plot_infill(prediction_selective, selective) + +key, sample_key = jr.split(key) +means_selective = global_mean_samples( + posterior_selective, spec_selective, selective, sample_key +) +report_global_mean(means_selective, selective) + +# %% [markdown] +# The kept cells are mostly the cold ones, so their mean is far too cold. The GP +# removes part of this bias, because it fills each gap from its neighbours and +# the warm regions still have some kept cells. But the result is still too cold, +# and the interval does not hold the truth: the GP is confidently wrong. +# +# The reason is that a GP conditions only on the values it sees. Its prior has one +# constant mean, which it learns from the kept cells, so the cold kept cells pull +# that mean down. It also has no way to know that the missing cells are warm, +# because nothing in its inputs says so. Its uncertainty describes the spread of +# values that are *consistent with the kept cells*, not the error of the +# selection. +# +# In real data we cannot see this bias, because we do not have the missing values. +# Useful steps are: +# +# - Add inputs that explain why cells are missing, such as cloud fraction or a +# covariate that is observed everywhere. If the chance of a gap depends only on +# the inputs, the gaps are *missing at random* given those inputs, and the GP +# can correct for them. +# - Model the observation process together with the field. +# - Test how sensitive the result is to different assumptions about the missing +# values, as we did here with a complete field. +# +# ## Adding an input that explains the gaps +# +# If a variable that explains the gaps is observed everywhere, add it as an input. +# [`Standardise`](#gpjax.xarray.Standardise) puts it on the same scale as the +# sphere coordinates: +# +# ```python +# data, spec = from_xarray( +# observed, +# target="tas", +# inputs=["lat", "lon", "cloud_fraction"], +# transforms=[UnitSphere(), Standardise(["cloud_fraction"])], +# ) +# ``` +# +# For a monthly record, [`Cyclic`](#gpjax.xarray.Cyclic) adds the seasonal cycle; +# see [Working with Gridded Data](xarray_workflow.py). +# +# ## References +# +# Kalnay, E., Kanamitsu, M., Kistler, R., Collins, W., Deaven, D., Gandin, L., +# Iredell, M., Saha, S., White, G., Woollen, J., Zhu, Y., Chelliah, M., Ebisuzaki, +# W., Higgins, W., Janowiak, J., Mo, K. C., Ropelewski, C., Wang, J., Leetmaa, A., +# Reynolds, R., Jenne, R. and Joseph, D. (1996). The NCEP/NCAR 40-year reanalysis +# project. *Bulletin of the American Meteorological Society*, 77(3), 437–471. +# [doi:10.1175/1520-0477(1996)077<0437:TNYRP>2.0.CO;2](https://doi.org/10.1175/1520-0477(1996)077%3C0437:TNYRP%3E2.0.CO;2) +# +# ## System configuration + +# %% +# %reload_ext watermark +# %watermark -n -u -v -iv -w -a 'Thomas Pinder' diff --git a/docs/examples/xarray_workflow.py b/docs/examples/xarray_workflow.py index 9f06dc030..99fadc9e8 100644 --- a/docs/examples/xarray_workflow.py +++ b/docs/examples/xarray_workflow.py @@ -15,34 +15,31 @@ # --- # %% [markdown] -# %% [markdown] -# # Gridded Data with xarray +# # Working with Gridded Data # # Download this notebook: {nb-download}`xarray_workflow.ipynb` # -# Climate data usually arrive as labelled [xarray](https://docs.xarray.dev/) -# objects, read from netCDF files: a temperature field over latitude and -# longitude, often with gaps where there were no observations. A GP in GPJax, on -# the other hand, consumes a [`Dataset`](#gpjax.dataset.Dataset) of flat inputs -# $\mathbf{X} \in \mathbb{R}^{N \times D}$ and outputs $\mathbf{y} \in -# \mathbb{R}^{N \times 1}$. The [`gpjax.xarray`](../reference/xarray.md) module -# converts between the two at the edges of a workflow. +# Climate and environmental data rarely arrive as a tidy matrix. They arrive as +# labelled [xarray](https://docs.xarray.dev/) objects: a temperature field over +# latitude and longitude, a covariate such as elevation, and gaps where a sensor +# failed or a cloud covered the scene. A GP in GPJax, on the other hand, consumes a +# [`Dataset`](#gpjax.dataset.Dataset) of flat inputs $\mathbf{X} \in \mathbb{R}^{N +# \times D}$ and outputs $\mathbf{y} \in \mathbb{R}^{N \times 1}$. # -# In this notebook we fill gaps in a global temperature field. To know how good -# the filled values are, we start from a field that is complete, remove cells -# ourselves, and compare the GP with the cells that we removed. We +# The [`gpjax.xarray`](../reference/xarray.md) module converts between the two at the +# edges of a workflow. In this notebook we # -# 1. flatten a gappy, labelled global field into a `Dataset` with -# [`from_xarray`](#gpjax.xarray.from_xarray), and put the inputs on the sphere -# with the [`UnitSphere`](#gpjax.xarray.UnitSphere) transform, +# 1. flatten a gappy, labelled field into a `Dataset` with +# [`from_xarray`](#gpjax.xarray.from_xarray), # 2. fit a GP exactly as we would on any other `Dataset`, # 3. predict the mean and variance on a finer grid with -# [`GridSpec.predict`](#gpjax.xarray.GridSpec.predict), which works in chunks, -# 4. draw joint posterior samples with +# [`GridSpec.predict`](#gpjax.xarray.GridSpec.predict), and +# 4. draw joint posterior samples on that grid with # [`GridSpec.inputs_for`](#gpjax.xarray.GridSpec.inputs_for) and -# [`GridSpec.to_xarray`](#gpjax.xarray.GridSpec.to_xarray), to get the global -# mean temperature with its uncertainty, and -# 5. show what goes wrong when cells are missing *because of* their values. +# [`GridSpec.to_xarray`](#gpjax.xarray.GridSpec.to_xarray). +# +# For the same workflow on real data, see +# [Infilling Global Surface Temperature](infilling_surface_temperature.py). # # The module needs the optional extra: `pip install "gpjax[xarray]"`. @@ -54,8 +51,6 @@ logging.getLogger("matplotlib.font_manager").setLevel(logging.ERROR) # %% -from pathlib import Path - from jax import config import jax.numpy as jnp import jax.random as jr @@ -68,98 +63,95 @@ with install_import_hook("gpjax", "beartype.beartype"): import gpjax as gpx - from gpjax.parameters import val from gpjax.xarray import ( - UnitSphere, + Standardise, from_xarray, ) key = jr.key(42) -rng = np.random.default_rng(0) gpx.plotting.use_style() # %% [markdown] -# ## A complete temperature field -# -# We use the [NCEP-NCAR Reanalysis 1](https://psl.noaa.gov/data/gridded/data.ncep.reanalysis.html) -# (Kalnay et al., 1996). A reanalysis combines a weather model with the -# observations, so it has a value in every grid cell. The file holds the 2024 -# annual-mean near-surface air temperature as an *anomaly*: the difference from -# the 1991–2020 mean of the same cell. Anomalies remove the large, fixed -# differences between the equator and the poles, so what remains is the -# signal of interest, and it varies smoothly over large distances. +# ## A synthetic temperature field # -# NCEP-NCAR Reanalysis 1 data provided by the NOAA PSL, Boulder, Colorado, USA, -# from their website at . - -# %% -nc_candidates = [ - Path("docs/examples/data/ncep_air_anomaly_2024.nc"), - Path("data/ncep_air_anomaly_2024.nc"), -] -nc_path = next(path for path in nc_candidates if path.exists()) -reanalysis = xr.open_dataset(nc_path) -reanalysis - -# %% [markdown] -# The native grid is 2.5°. We fit on a 5° grid, which is every second grid point, -# and predict back on the 2.5° grid, so the predictions are also tested at points -# the model never saw. - -# %% -truth_fine = reanalysis["tas_anomaly"].astype(float) -truth = truth_fine.isel(lat=slice(1, None, 2), lon=slice(1, None, 2)) -print(f"5° grid: {dict(truth.sizes)}, 2.5° grid: {dict(truth_fine.sizes)}") - -anomaly_style = dict(cmap="RdBu_r", vmin=-4.0, vmax=4.0) -error_style = dict( - cmap="PuOr_r", vmin=-2.0, vmax=2.0, cbar_kwargs={"label": "Error [K]"} -) -truth_fine.plot(figsize=(7, 3.5), **anomaly_style) -plt.title("2024 anomaly, NCEP-NCAR Reanalysis 1") -plt.show() - -# %% [markdown] -# A 5° cell near a pole is much smaller than one at the equator, so a global mean -# weights each cell by the cosine of its latitude. +# We simulate near-surface temperature on a regional latitude-longitude grid. It +# cools towards the pole and with elevation (a lapse rate of roughly 6.5 K per +# kilometre), with a smooth large-scale anomaly on top. Elevation is a separate +# variable over the same grid. The data are synthetic, so the notebook needs no +# download and every number in it can be checked against the truth. # %% -def global_mean(field: xr.DataArray) -> xr.DataArray: - """Area-weighted mean over the cells that have a value.""" - weights = np.cos(np.deg2rad(field["lat"])) - return field.weighted(weights).mean(["lat", "lon"]) +def elevation_at(lat, lon): + """A single mountain range, in metres.""" + return 2500.0 * np.exp(-(((lon - 12.0) / 4.0) ** 2) - ((lat - 47.0) / 3.0) ** 2) + + +def temperature_at(lat, lon, elevation): + """Temperature in kelvin: latitude gradient, lapse rate and an anomaly.""" + anomaly = 1.5 * np.sin(lon / 3.0) * np.cos(lat / 4.0) + return 290.0 - 0.6 * (lat - 40.0) - 0.0065 * elevation + anomaly + + +def regional_field(lats, lons) -> xr.Dataset: + lat_grid, lon_grid = np.meshgrid(lats, lons, indexing="ij") + elevation = elevation_at(lat_grid, lon_grid) + return xr.Dataset( + { + "t2m": ( + ("lat", "lon"), + temperature_at(lat_grid, lon_grid, elevation), + {"units": "K", "long_name": "2 m air temperature"}, + ), + "elevation": (("lat", "lon"), elevation, {"units": "m"}), + }, + coords={ + "lat": ("lat", lats, {"units": "degrees_north"}), + "lon": ("lon", lons, {"units": "degrees_east"}), + }, + ) -print(f"True global mean anomaly: {float(global_mean(truth)):.2f} K") +coarse = regional_field(np.linspace(40.0, 54.0, 12), np.linspace(2.0, 22.0, 16)) # %% [markdown] -# ## Cells missing at random -# -# First we keep a random 25% of the 5° cells. Which cells are missing has nothing -# to do with their values. Statisticians call this *missing completely at -# random*. +# Real observations have holes. We knock out a block of cells, as a cloud would, +# and add a little measurement noise to the rest. # %% -kept_at_random = rng.uniform(size=truth.shape) < 0.25 -observed = truth.where(kept_at_random).rename("tas").to_dataset() -print(f"{int(kept_at_random.sum())} of {truth.size} cells kept") +key, noise_key = jr.split(key) +noise = 0.2 * np.asarray(jr.normal(noise_key, coarse["t2m"].shape)) +observed = coarse.copy(deep=True) +observed["t2m"] = observed["t2m"] + noise +observed["t2m"].attrs = coarse["t2m"].attrs +observed["t2m"][4:7, 9:13] = np.nan + +observed["t2m"].plot(cmap="coolwarm") +plt.title("Observed temperature (gaps in white)") +plt.show() # %% [markdown] -# ### Inputs on the sphere +# ## From labelled data to a `Dataset` # -# Latitude and longitude in degrees are not good GP inputs for a global field. A -# 5° cell is about 555 km wide at the equator but only about 24 km wide next to a -# pole, and longitude 177.5°E is next to 177.5°W. The -# [`UnitSphere`](#gpjax.xarray.UnitSphere) transform replaces `lat` and `lon` with -# the three coordinates of a point on the unit sphere. A stationary kernel on these -# coordinates uses the chord distance through the Earth, so it has no seam at the -# antimeridian and no distortion at the poles. Its lengthscale is in Earth radii. +# `from_xarray` takes the target variable and the inputs we want the GP to depend +# on. Inputs can be coordinates (`lat`, `lon`) or other data variables +# (`elevation`), and the columns of $\mathbf{X}$ follow the order we list them in. +# Cells where the target or any input is NaN are dropped by default, and the +# returned `GridSpec` records which ones. +# +# Temperature varies over degrees of latitude and longitude, but over hundreds of +# metres of elevation. The [`Standardise`](#gpjax.xarray.Standardise) transform +# scales each input to zero mean and unit standard deviation over the training +# cells, so one starting lengthscale suits every input. The spec keeps the +# training mean and standard deviation, and applies them to every grid that we +# predict on. # %% +inputs = ["lat", "lon", "elevation"] data, spec = from_xarray( - observed, target="tas", inputs=["lat", "lon"], transforms=[UnitSphere()] + observed, target="t2m", inputs=inputs, transforms=[Standardise()] ) + print(data) print(spec) @@ -168,240 +160,158 @@ def global_mean(field: xr.DataArray) -> xr.DataArray: # The `GridSpec` stays with us, outside the model, until we want labelled output # again. # -# ### Fitting the model +# ## Fitting the model # -# We use a Matérn-3/2 kernel, which gives rougher fields than an RBF kernel, as -# temperature anomalies are. A constant mean absorbs the global warming signal. - +# We give the RBF kernel one lengthscale per input, in units of that input's +# standard deviation. A constant mean absorbs the ~285 K offset. # %% -def fit_model(train_data: gpx.Dataset): - prior = gpx.gps.Prior( - mean_function=gpx.mean_functions.Constant(jnp.array([0.5])), - kernel=gpx.kernels.Matern32(lengthscale=jnp.array(0.3), variance=1.0), - ) - model = prior * gpx.likelihoods.Gaussian(obs_stddev=jnp.array(0.1)) - model, _ = gpx.fit_scipy( - model=model, - objective=lambda candidate, train_data: ( - -gpx.objectives.conjugate_mll(candidate, train_data) - ), - train_data=train_data, - verbose=False, - ) - return model - - -earth_radius_km = 6371.0 -model = fit_model(data) -lengthscale_km = float(val(model.prior.kernel.lengthscale)) * earth_radius_km -print(f"Lengthscale: {lengthscale_km:.0f} km") +prior = gpx.gps.Prior( + mean_function=gpx.mean_functions.Constant(jnp.array([285.0])), + kernel=gpx.kernels.RBF(lengthscale=jnp.ones(3), variance=25.0), +) +model = prior * gpx.likelihoods.Gaussian(obs_stddev=jnp.array(0.5)) + +model, history = gpx.fit_scipy( + model=model, + objective=lambda candidate, train_data: -gpx.objectives.conjugate_mll( + candidate, train_data + ), + train_data=data, + verbose=False, +) # %% [markdown] -# ### Predicting on the finer grid +# ## Predicting on a finer grid # # `spec.predict` gives the predictive mean and variance on any grid that holds the -# same inputs, encoded exactly as in training. Here the grid is the 2.5° grid of -# the reanalysis. The function we pass maps a block of inputs to a distribution; -# we use a diagonal covariance, because we need only the variance of each cell. +# same input variables, encoded exactly as in training. Here we predict on a grid +# four times finer in each direction, including the cells that were missing from +# the observations. The function we pass maps a block of inputs to a +# distribution. We use the likelihood's predictive distribution, so the variance +# includes observation noise, and a diagonal covariance, because we need only the +# variance of each cell. # # `spec.predict` sends the cells to this function in chunks of `chunk_size`, so -# the memory use stays the same for a finer or larger grid. If the grid holds -# Dask arrays, the result is lazy, and each Dask block is predicted only when it is +# the memory use stays the same for a global grid. If the grid holds Dask +# arrays, the result is lazy, and each Dask block is predicted only when it is # computed or written with `to_netcdf`. # %% +fine = regional_field(np.linspace(40.0, 54.0, 45), np.linspace(2.0, 22.0, 61)) posterior = model.condition(data) + prediction = spec.predict( lambda x: model.likelihood(posterior(x, covariance="diagonal")), - truth_fine.to_dataset(), - chunk_size=2048, + fine[["elevation"]], + chunk_size=1024, ) prediction # %% [markdown] -# Because the field is complete, we can score every prediction. We compare with a -# simple baseline, the mean of the kept cells, and check the uncertainty: about -# 95% of the true values should be within two predictive standard deviations. - +# The result is a labelled `xr.Dataset` on the fine grid, with the target's +# attributes carried over. The variance is in $\mathrm{K}^2$. Everything xarray +# offers, from plotting to `to_netcdf`, works on it directly. # %% -def score(prediction: xr.Dataset, observed: xr.Dataset) -> None: - error = prediction["tas_mean"] - truth_fine - baseline_error = float(observed["tas"].mean()) - truth_fine - within = abs(error) < 2 * np.sqrt(prediction["tas_variance"]) - print(f"RMSE, GP: {float(np.sqrt((error**2).mean())):.2f} K") - print( - f"RMSE, mean of kept cells: {float(np.sqrt((baseline_error**2).mean())):.2f} K" +fig, (mean_ax, std_ax, error_ax) = plt.subplots(1, 3, figsize=(15, 4)) +prediction["t2m_mean"].plot(ax=mean_ax, cmap="coolwarm") +mean_ax.set_title("Predictive mean") +np.sqrt(prediction["t2m_variance"]).plot(ax=std_ax, cmap="viridis") +std_ax.set_title("Predictive standard deviation") +(prediction["t2m_mean"] - fine["t2m"]).plot(ax=error_ax, cmap="RdBu_r", center=0.0) +error_ax.set_title("Error against the true field") +for ax in (mean_ax, std_ax, error_ax): + ax.add_patch( + plt.Rectangle( + (observed.lon[9], observed.lat[4]), + float(observed.lon[12] - observed.lon[9]), + float(observed.lat[6] - observed.lat[4]), + fill=False, + linestyle="--", + ) ) - print(f"Within 2 sd: {float(within.mean()):.0%}") - - -def plot_infill(prediction: xr.Dataset, observed: xr.Dataset) -> None: - fig, axes = plt.subplots(1, 3, figsize=(15, 3.5), sharey=True) - observed["tas"].plot(ax=axes[0], **anomaly_style) - axes[0].set_title("Kept cells") - prediction["tas_mean"].plot(ax=axes[1], **anomaly_style) - axes[1].set_title("Predictive mean") - (prediction["tas_mean"] - truth_fine).plot(ax=axes[2], **error_style) - axes[2].set_title("Predictive mean minus truth") - for ax in axes[1:]: - ax.set_ylabel("") - plt.show() - - -score(prediction, observed) -plot_infill(prediction, observed) +plt.show() # %% [markdown] -# From a quarter of the cells, the GP recovers the large-scale pattern, including -# the strong warmth over the Arctic. The errors are largest where the field -# changes over short distances, and the uncertainty is about right. +# The standard deviation grows inside the dashed box where observations were +# missing, and the error stays small across the mountain range because elevation is +# an input. # -# ### The global mean, with joint samples +# ## Joint samples and regional averages # -# The global mean anomaly fills each missing cell and averages. Its variance -# depends on the covariance between the filled cells, +# The mean and variance describe each cell on its own. Many questions are about +# several cells together, such as the average temperature over the Alpine box +# $\mathcal{R}$. Its variance depends on the covariance between the cells, # # $$ -# \operatorname{Var}\Big[\sum_{i} w_i f_i\Big] -# = \sum_{i, j} w_i w_j \operatorname{Cov}[f_i, f_j], -# $$ (eq-xarray-global-variance) +# \operatorname{Var}\Big[\tfrac{1}{|\mathcal{R}|} \sum_{i \in \mathcal{R}} f_i\Big] +# = \tfrac{1}{|\mathcal{R}|^2} \sum_{i, j \in \mathcal{R}} \operatorname{Cov}[f_i, f_j], +# $$ (eq-xarray-regional-variance) # # which the per-cell variances alone cannot give. For this we need the full -# covariance, so we build all the prediction inputs on the 5° grid at once with +# covariance, so we build all the prediction inputs at once with # `spec.inputs_for`. Passing `num_samples` to `to_xarray` then draws from the joint -# distribution of the field, and returns the draws with a leading `sample` -# dimension. `combine_first` keeps the kept cells and takes each missing cell from -# a draw. - +# predictive distribution, and returns the draws with a leading `sample` +# dimension. Averaging each draw over the region gives samples of the regional +# mean. # %% -def global_mean_samples(posterior, spec, observed: xr.Dataset, key) -> xr.DataArray: - test_inputs, test_spec = spec.inputs_for(truth.to_dataset()) - samples = test_spec.to_xarray(posterior(test_inputs), num_samples=500, key=key) - return global_mean(observed["tas"].combine_first(samples["tas"])) - - -def report_global_mean(means: xr.DataArray, observed: xr.Dataset) -> None: - print(f"Truth: {float(global_mean(truth)):.2f} K") - print(f"Mean of kept cells: {float(global_mean(observed['tas'])):.2f} K") - print( - f"Gaps filled by GP: {float(means.mean()):.2f}" - f" ± {2 * float(means.std('sample')):.2f} K (2 sd)" - ) - - +test_inputs, test_spec = spec.inputs_for(fine[["elevation"]]) +latent = posterior(test_inputs) # the field itself, without observation noise key, sample_key = jr.split(key) -means = global_mean_samples(posterior, spec, observed, sample_key) -report_global_mean(means, observed) - -# %% [markdown] -# With random gaps, even the mean of the kept cells is close to the truth, and the -# GP gives the global mean with an interval that holds it. -# -# ## Cells missing because of their values -# -# Real gaps are rarely random. A satellite cannot see the surface through cloud, -# a station fails in extreme weather, and the polar regions, which warm fastest, -# have the fewest observations. When the chance that a cell is missing depends on -# the value that is missing, the data are *missing not at random*. -# -# We simulate this. We again aim to keep a quarter of the cells, but now the -# warmer a cell's anomaly, the less likely it is to be kept. - -# %% -standardised = ((truth - truth.mean()) / truth.std()).values -odds = np.exp(-1.5 * standardised) -keep_probability = np.clip(0.25 * odds / odds.mean(), 0.0, 1.0) -kept_selectively = rng.uniform(size=truth.shape) < keep_probability -selective = truth.where(kept_selectively).rename("tas").to_dataset() -print(f"{int(kept_selectively.sum())} of {truth.size} cells kept") - -# %% [markdown] -# The steps are the same as before. +samples = test_spec.to_xarray(latent, num_samples=500, key=sample_key) -# %% -data_selective, spec_selective = from_xarray( - selective, target="tas", inputs=["lat", "lon"], transforms=[UnitSphere()] -) -model_selective = fit_model(data_selective) -posterior_selective = model_selective.condition(data_selective) -prediction_selective = spec_selective.predict( - lambda x: model_selective.likelihood(posterior_selective(x, covariance="diagonal")), - truth_fine.to_dataset(), - chunk_size=2048, -) +alps = dict(lat=slice(45.0, 49.0), lon=slice(8.0, 16.0)) +regional_mean = samples["t2m"].sel(**alps).mean(["lat", "lon"]) +true_regional_mean = float(fine["t2m"].sel(**alps).mean()) -score(prediction_selective, selective) -plot_infill(prediction_selective, selective) +# The same latent distribution, but treating the cells as independent. +latent_variance = test_spec.to_xarray(latent)["t2m_variance"].sel(**alps) +joint_std = float(regional_mean.std("sample")) +naive_std = float(np.sqrt(latent_variance.sum()) / latent_variance.size) -key, sample_key = jr.split(key) -means_selective = global_mean_samples( - posterior_selective, spec_selective, selective, sample_key -) -report_global_mean(means_selective, selective) +print(f"Regional mean: {float(regional_mean.mean()):.2f} K (truth {true_regional_mean:.2f} K)") +print(f"Standard deviation from joint samples: {joint_std:.3f} K") +print(f"Standard deviation if cells were independent: {naive_std:.3f} K") # %% [markdown] -# The kept cells are mostly the cold ones, so their mean is far too cold. The GP -# removes part of this bias, because it fills each gap from its neighbours and -# the warm regions still have some kept cells. But the result is still too cold, -# and the interval does not hold the truth: the GP is confidently wrong. +# Treating the cells as independent understates the uncertainty in the regional +# average by roughly an order of magnitude, because neighbouring cells tend to be +# wrong in the same direction. Against the joint standard deviation the true +# regional mean is a plausible outcome; against the independent one it would look +# like a many-sigma surprise. Joint +# samples keep that correlation, which is why `to_xarray` refuses to draw samples +# from a distribution that only holds marginal variances (as returned by +# `posterior(test_inputs, covariance="diagonal")`). # -# The reason is that a GP conditions only on the values it sees. Its prior has one -# constant mean, which it learns from the kept cells, so the cold kept cells pull -# that mean down. It also has no way to know that the missing cells are warm, -# because nothing in its inputs says so. Its uncertainty describes the spread of -# values that are *consistent with the kept cells*, not the error of the -# selection. +# ## Global grids and seasonal cycles # -# In real data we cannot see this bias, because we do not have the missing values. -# Useful steps are: -# -# - Add inputs that explain why cells are missing, such as cloud fraction or a -# covariate that is observed everywhere. If the chance of a gap depends only on -# the inputs, the gaps are *missing at random* given those inputs, and the GP -# can correct for them. -# - Model the observation process together with the field. -# - Test how sensitive the result is to different assumptions about the missing -# values, as we did here with a complete field. -# -# ## Other inputs and seasonal cycles -# -# The same workflow takes more inputs. Other transforms encode them: -# [`Cyclic`](#gpjax.xarray.Cyclic) encodes a periodic input as a point on a -# circle, and [`Standardise`](#gpjax.xarray.Standardise) scales an input to zero -# mean and unit standard deviation over the training cells. Datetime inputs are -# days since their first timestamp, so a period of 365.25 gives the seasonal -# cycle of a monthly record: +# Latitude and longitude in degrees are not good inputs for a global field. A +# degree of longitude is about 111 km at the equator but only 38 km at 70°N, and +# longitude 359° is next to 0°. The +# [`UnitSphere`](#gpjax.xarray.UnitSphere) transform replaces `lat` and `lon` with +# the three coordinates of a point on the unit sphere. A stationary kernel on these +# coordinates uses the chord distance through the Earth, which has no seam and no +# distortion at the poles. In the same way, [`Cyclic`](#gpjax.xarray.Cyclic) +# encodes a periodic input as a point on a circle. Datetime inputs are days since +# their first timestamp, so a period of 365.25 gives the seasonal cycle: # # ```python # data, spec = from_xarray( -# monthly, -# target="tas", -# inputs=["lat", "lon", "time", "cloud_fraction"], -# transforms=[ -# UnitSphere(), -# Cyclic("time", 365.25), -# Standardise(["cloud_fraction"]), -# ], +# ds, +# target="t2m", +# inputs=["lat", "lon", "time", "elevation"], +# transforms=[UnitSphere(), Cyclic("time", 365.25), Standardise(["elevation"])], # ) # spec.columns -# # ('sphere_x', 'sphere_y', 'sphere_z', 'time_sin', 'time_cos', 'cloud_fraction') +# # ('sphere_x', 'sphere_y', 'sphere_z', 'time_sin', 'time_cos', 'elevation') # ``` # # The transforms run in order, and `spec.columns` names the columns of -# $\mathbf{X}$ that they produce. The spec applies the same fitted transforms to -# every grid that we predict on. -# -# ## References -# -# Kalnay, E., Kanamitsu, M., Kistler, R., Collins, W., Deaven, D., Gandin, L., -# Iredell, M., Saha, S., White, G., Woollen, J., Zhu, Y., Chelliah, M., Ebisuzaki, -# W., Higgins, W., Janowiak, J., Mo, K. C., Ropelewski, C., Wang, J., Leetmaa, A., -# Reynolds, R., Jenne, R. and Joseph, D. (1996). The NCEP/NCAR 40-year reanalysis -# project. *Bulletin of the American Meteorological Society*, 77(3), 437–471. -# [doi:10.1175/1520-0477(1996)077<0437:TNYRP>2.0.CO;2](https://doi.org/10.1175/1520-0477(1996)077%3C0437:TNYRP%3E2.0.CO;2) +# $\mathbf{X}$ that they produce, one for each kernel lengthscale. +# [Infilling Global Surface Temperature](infilling_surface_temperature.py) uses +# `UnitSphere` on a real global field. # # ## System configuration diff --git a/docs/index.md b/docs/index.md index c75c7e81c..2af488138 100644 --- a/docs/index.md +++ b/docs/index.md @@ -152,6 +152,7 @@ examples/intro_to_kernels examples/regression examples/classification examples/poisson +examples/xarray_workflow examples/natural_gradients ``` @@ -176,11 +177,11 @@ examples/oilmm examples/barycentres examples/graph_kernels examples/heteroscedastic_inference +examples/infilling_surface_temperature examples/multioutput examples/oak examples/oceanmodelling examples/spatial_linear_gp -examples/xarray_workflow examples/yacht ```