GridSpec#

class gpjax.xarray.GridSpec(target, target_attrs, inputs, dims, coords, time_origins, mask, n_dropped)[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 is all this object records, and all inputs_for() 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 column order of X.

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

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. Cells where an input is NaN get no row, and come back as NaN from to_xarray().

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 and attributes.

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.

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