GridSpec#
- class gpjax.xarray.GridSpec(target, target_attrs, inputs, dims, coords, time_origins, mask, n_dropped, transforms=())[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, and the encoding of the inputs as the columns ofX, is all this object records, and allinputs_for(),predict()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:
- transforms#
The input transforms, fitted on the training cells, that turn
inputsinto the columns ofX.- 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, 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 fromto_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 theGridSpecof 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]
- predict(predict_fn, obj, *, chunk_size=4096)[source]#
Predict the mean and variance on a new grid, in chunks.
The grid is built from
objas ininputs_for(), but the cells go topredict_fnat mostchunk_sizeat a time, so the memory use does not grow with the size of the grid. The last chunk is padded tochunk_size, sopredict_fnis compiled once, withjax.jit.If
objholds Dask arrays, the result is lazy: each Dask block is predicted when it is computed, for example by.compute()orto_netcdf. Chunkobjalong 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 withmeanandvarianceof shape(M,), for examplelambda x: posterior(x, covariance="diagonal"). Use the diagonal covariance: a dense one costschunk_size**2memory 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_fnat once.
- Returns:
An
xr.Datasetholding{target}_meanand{target}_varianceover the grid. Cells where an input is NaN are NaN.- Raises:
ValueError – If
chunk_sizeis not positive, an input is missing fromobj, orpredict_fnreturns 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 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