predictive.predict_in_sample()
Sample the in-sample posterior predictive of the obs site.
Usage
predictive.predict_in_sample(
rng_key,
model,
posterior,
covariates,
*,
batch_size=None,
parallel=True,
device=None
)Runs Predictive with the in-sample covariates and the supplied posterior latent draws. Unlike forecast() there is no forecast horizon: covariates span only the observed window, so the model’s obs site is sampled at every step. The number of predictive samples equals the leading (sample) axis of posterior (see draw_posterior()).
Parameters
rng_key: Array-
PRNG key.
model: ForecastModel-
The forecasting model callable (the same one that produced
posterior). posterior: Mapping[str, ArrayLike]-
Posterior samples of the latent sites, sample axis leading. The output of a
device="host"stage (CPU-committed jax leaves or NumPy leaves) is accepted directly. covariates: Array-
Covariates with time at axis
-2spanning the observed window. Its time length must match the data theposteriorwas fit on, since the in-sample latent sites are sized to that window. batch_size: int | None = None-
Optional chunk size for sampling (caps peak memory).
parallel: bool = True-
Whether
Predictivevectorizes over the sample axis withvmap(True, faster, higher peak memory) or maps it serially withlax.map(False). See forecast() for how this interacts withbatch_size. device: jax.Device | str | None = None-
Where each chunk of draws is placed as soon as it is drawn and where the stitched result lives; the same placement contract as the
deviceargument of draw_posterior() ("host"for pageable host memory as a CPU-committedjax.Array, or a NumPy array when no CPU backend is initialized;"numpy";"pinned_host"; ajax.Deviceor platform name;Nonefor the default device), including its mixing rules for host-committed results. Withbatch_sizeset on an accelerator, any host target bounds accelerator memory by a single chunk instead of the full(sample, time, obs)array; the draw values are unchanged, only where the result lives. The bound requiresbatch_sizestrictly below the sample count: at or above it, the single-shot path runs and the full array is materialized on the default device before the one transfer. The result feeds straight into to_datatree() and thebatch_size-chunked evaluation metrics inevaluate, which accept host-resident draws.
Returns
Num[Array, " sample *batch time obs"]-
In-sample posterior-predictive draws of the
obssite (withdevice="host"committed to the CPU device, or a NumPy array when no CPU backend is initialized).
Raises
HostMemoryKindError-
If
device="pinned_host"is requested on a device that exposes no host memory kind (see_host_memory_kind()). DevicePlatformError-
If
devicenames a platform whose backend is not initialized (see_resolve_device()).
Warns
UserWarning-
If
device="cpu"is requested and the JAX CPU backend is not initialized, so the draws take the NumPy path of"host"instead.