aimz.ImpactModel.predict#
- ImpactModel.predict(X, *, intervention=None, rng_key=None, in_sample=True, return_sites=None, shard_axis='obs', batch_size=None, output_dir=None, progress=True, **kwargs)[source]#
Predict the output based on the fitted model.
This method performs posterior predictive sampling to generate model-based predictions. It is optimized for batch processing of large input data and is not recommended for use in loops that process only a few inputs at a time. Results are written to disk in the Zarr format, with sampling and file writing decoupled and executed concurrently.
- Parameters:
X (ArrayLike | ArrayLoader) – Input data. If array-like, the leading axis is
Alternatively (the observation axis.)
array-like (a data loader that holds all)
internally. (objects and handles batching)
intervention (dict | None) – A dictionary mapping sample sites to their corresponding intervention values. Interventions enable counterfactual analysis by modifying the specified sample sites during prediction (posterior predictive sampling).
rng_key (Array | None) – A pseudo-random number generator key. By default, an internal key is used and split as needed.
in_sample (bool) – Specifies the group where posterior predictive samples are stored in the returned output. If
True, samples are stored in theposterior_predictivegroup, indicating they were generated based on data used during model fitting. IfFalse, samples are stored in thepredictionsgroup, indicating they were generated based on out-of-sample data.return_sites (str | Iterable[str] | None) – Names of variables (sites) to return. If
None, samplesparam_outputand deterministic sites.shard_axis (Literal['obs', 'draw']) – Multi-device sharding strategy; no effect on a single device.
"obs"(default) shards the input across devices and replicates the posterior."draw"shards the posterior across devices and replicates the input, which must be an array, not a data loader. If the model has no posterior samples, the data path is used regardless ofshard_axis.batch_size (int | None) – Size of each batch, taken from the input under
shard_axis="obs"and from the draws undershard_axis="draw". Also used as the chunk size when storing results. IfNone, it is determined automatically from the input size and number of samples. Ignored ifXis a data loader, in which case the data loader is expected to handle batching internally.output_dir (str | Path | None) – The directory where the outputs will be saved. If the specified directory does not exist, it will be created automatically. If
None, a model-owned temporary directory is used. A subdirectory is generated within this directory to store the outputs. The temporary directory is removed bycleanup(),cleanup_models(), or when the model is garbage-collected. Pass an explicitoutput_dirto keep results beyond the model’s lifetime.progress (bool) – Whether to display a progress bar.
**kwargs (object) – Additional arguments passed to the model.
- Returns:
Posterior predictive samples. Posterior samples are included if available.
- Raises:
TypeError – If
param_outputis passed as an argument, orshard_axis="draw"is used with a data loaderX.ValueError – If
shard_axisis not"obs"or"draw".NotImplementedError – If a return site’s axis-1 size does not match the input batch size (
shard_axis="obs"only).
- Return type:
xr.DataTree
See also
cleanup()to remove the temporary directory if created.