GridSpec#
- class gpjax.xarray.GridSpec(target, target_attrs, inputs, dims, coords, time_origins, mask, n_dropped)[source]#
Bases:
objectThe labelled grid behind a flattened
Dataset.Row
iof the flattened data is thei-th kept cell of the grid in C order overdims. That correspondence is all this object records, and allinputs_for()andto_xarray()need.- Parameters:
- target_attrs#
The target’s attributes (units, long_name, …), carried onto predictions.
- 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:
- mask#
Boolean array over the full grid;
Truemarks a cell that has a row in the flattened data.- Type:
- 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 fromto_xarray().- Parameters:
obj (Dataset | DataArray) – Labelled data holding every input named by this spec.
- Returns:
The
(M, D)prediction inputs and theGridSpecof 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]
- 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 byinputs_for().num_samples (int | None) – If given, return this many joint draws from
distinstead 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.Datasetholding{target}_meanand{target}_varianceover the grid. Withnum_samples, one variable{target}over("sample", *dims). Dropped cells are NaN either way.- Raises:
ValueError – If
distdoes not match the number of kept cells, ifnum_samplesis given withoutkey, or if samples are requested from a distribution holding only marginal variances.- Return type:
Dataset