VI Inference#
Variational-inference utilities for blayers.
Provides Batched_Trace_ELBO, a drop-in Trace_ELBO replacement
that handles minibatching without requiring the model to use numpyro.plate,
and svi_run_batched(), an svi.run-style helper that drives it.
Use Batched_Trace_ELBO + svi_run_batched for plate-free batched VI;
fall back to standard Trace_ELBO if your model already uses plates.
- class blayers.vi_infer.Batched_Trace_ELBO(num_obs, num_particles=1, batch_size=None)[source]#
Bases:
ELBOELBO estimator for minibatched VI without
numpyro.plate.Behaves like
Trace_ELBObut rescales the per-batch log-likelihood bynum_obs / batch_sizeso the gradient is an unbiased estimate of the full-dataset ELBO. Drive it withsvi_run_batched().Assumes all latent variables are global. The whole observed log-likelihood is scaled by
num_obs / batch_sizeand the KL over latents is not rescaled, which is only correct when every latent is shared across observations (the usual case for BLayers: coefficients, scales, embeddings). Models with per-observation (local) latents — e.g. a latent variable sampled once per row — are not supported here; usenumpyro.platewith the standardTrace_ELBOinstead.- Parameters:
num_obs (int) – Total number of observations in the full training set.
num_particles (int) – Number of Monte Carlo samples per gradient step.
batch_size (int | None) – Fallback batch size for calls without array inputs. Otherwise the actual leading dimension of row-aligned positional and keyword arrays is used, including short remainder batches. Bind non-row arrays (e.g. knots) into the model with a closure.
NumPyro scale and mask handlers are honored for model and guide sites. All observed sites, including
numpyro.factorterms, are treated as row-wise likelihood contributions and receive N/B scaling. Global factors and local latent variables are unsupported in minibatched mode.Warning
Does not mix with
numpyro.plate. AValueErroris raised if a plate is detected in the model trace — thenum_obs / batch_sizerescaling double-counts plate-subsampled sites, so the ELBO would be silently wrong. Use the standardTrace_ELBOwith plates instead.- __init__(num_obs, num_particles=1, batch_size=None)[source]#
- Parameters:
num_obs (int)
num_particles (int)
batch_size (int | None)
- loss(rng_key, param_map, model, guide, *args, **kwargs)[source]#
Evaluates the ELBO with an estimator that uses num_particles many samples/particles.
- Parameters:
rng_key (jax.random.key) – random number generator seed.
param_map (dict) – dictionary of current parameter values keyed by site name.
model (Callable[[...], Any]) – Python callable with NumPyro primitives for the model.
guide (Callable[[...], Any]) – Python callable with NumPyro primitives for the guide.
args (Any) – arguments to the model / guide (these can possibly vary during the course of fitting).
kwargs (Any) – keyword arguments to the model / guide (these can possibly vary during the course of fitting).
- Returns:
negative of the Evidence Lower Bound (ELBO) to be minimized.
- Return type:
Array
- blayers.vi_infer.svi_run_batched(svi, rng_key, batch_size, num_steps=None, num_epochs=None, shuffle=True, progress_bar=False, **data)[source]#
Drive batched VI with compiled minibatch and epoch loops.
- Parameters:
shuffle (bool) – Re-permute rows each epoch (default). False retains the existing contiguous batch order, including its statistical tradeoff.
progress_bar (bool) – Report progress once per epoch. Defaults to False so the complete update sequence runs without Python dispatch per epoch.
svi (SVI)
rng_key (Array)
batch_size (int)
num_steps (int | None)
num_epochs (int | None)
data (Array)
- Return type:
SVIRunResult
Returns one loss per update. Short final batches retain their actual row count and N/B likelihood scaling; no observations are padded or dropped. Batch order and random-key splitting match the Python batch generator.