GridSpec#

class gpjax.xarray.GridSpec(target, target_attrs, inputs, dims, coords, time_origins, mask, n_dropped, transforms=())[source]#

Bases: object

The labelled grid behind a flattened Dataset.

Row i of the flattened data is the i-th kept cell of the grid in C order over dims. That correspondence, and the encoding of the inputs as the columns of X, is all this object records, and all inputs_for(), predict() and to_xarray() need.

Parameters:
target#

Name of the modelled variable.

Type:

str

target_attrs#

The target’s attributes (units, long_name, …), carried onto predictions.

Type:

dict[str, Any]

inputs#

Input names, in the order they are read from the data.

Type:

tuple[str, …]

dims#

The grid’s dims, in the target’s order.

Type:

tuple[str, …]

coords#

The grid’s coordinates, used to rebuild labelled output.

Type:

collections.abc.Mapping[collections.abc.Hashable, xarray.core.dataarray.DataArray]

time_origins#

For each datetime input, the timestamp encoded as day 0.

Type:

dict[str, numpy.datetime64]

mask#

Boolean array over the full grid; True marks a cell that has a row in the flattened data.

Type:

numpy.ndarray

n_dropped#

Number of grid cells dropped for containing NaN.

Type:

int

transforms#

The input transforms, fitted on the training cells, that turn inputs into the columns of X.

Type:

tuple[gpjax.xarray.InputTransform, …]

property columns: tuple[str, ...]#

Names of the columns of X, after the input transforms.

inputs_for(obj)[source]#

Build prediction inputs on a new grid, encoded as in training.

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, 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 to_xarray().

This builds every row at once. For a large grid, predict() gives the mean and variance in chunks instead.

Parameters:

obj (Dataset | DataArray) – 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, 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).

Return type:

tuple[Float[jaxlib._jax.Array, ‘M D’] | Float[ndarray, ‘M D’], GridSpec]

property n_kept: int#

Number of grid cells with a row in the flattened data.

predict(predict_fn, obj, *, chunk_size=4096)[source]#

Predict the mean and variance on a new grid, in chunks.

The grid is built from obj as in 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.

Parameters:
  • predict_fn (Callable[[Float[jaxlib._jax.Array, 'M D'] | Float[ndarray, 'M D']], Any]) – 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 (Dataset | DataArray) – Labelled data holding every input named by this spec.

  • chunk_size (int) – 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).

Return type:

Dataset

to_xarray(dist, *, num_samples=None, key=None)[source]#

Map a predictive distribution back onto the labelled grid.

Parameters:
  • dist (GaussianDistribution) – A distribution over exactly the cells this spec kept, in flattened order – e.g. posterior(test_inputs) for inputs built by inputs_for().

  • num_samples (int | None) – If given, return this many joint draws from dist instead of its mean and variance.

  • key (UInt32[jaxlib._jax.Array, '2'] | Key[jaxlib._jax.Array, ''] | None) – PRNG key for the draws; required with num_samples.

Returns:

By default, an xr.Dataset holding {target}_mean and {target}_variance over the grid. With num_samples, one variable {target} over ("sample", *dims). Dropped cells are NaN either way.

Raises:

ValueError – If dist does not match the number of kept cells, if num_samples is given without key, or if samples are requested from a distribution holding only marginal variances.

Return type:

Dataset