reparam.time_reparam()
Reparameterize every in-sample time latent of model along the time axis.
Usage
reparam.time_reparam(
model,
transform,
)Port of the time_reparam argument of Pyro’s Forecaster and HMCForecaster (forecaster.py). The returned model is numpyro.handlers.reparam(model, config) with a config that targets every non-observed continuous sample site under the time plate opened by innovations(), exactly as Pyro’s time_reparam_haar targets every site inside its time plate. Each targeted site name becomes a deterministic site of the same shape, computed from a new sample site f"{name}_haar" or f"{name}_dct" whose event shape owns the time axis and every plate axis to its right (plate axes to the left of time stay batch axes). The transform is orthonormal, so the model’s log density is unchanged; only the geometry inference sees changes.
Parameters
model: ForecastModel-
A forecasting model
(covariates, data=None) -> None. transform: TimeTransform-
"haar"fornumpyro.distributions.transforms.HaarTransform(a multi-resolution average/difference basis, suited to blocky or multi-scale dynamics) or"dct"fornumpyro.distributions.transforms.DiscreteCosineTransform(a cosine frequency basis, close to the Karhunen-Loeve basis of first-order Markov processes). Note that Pyro’s own string mapping is swapped: its"haar"runs a DCT and its"dct"runs a Haar transform.
Returns
ForecastModel-
The wrapped model. It is a
numpyro.handlers.reparamhandler and is the single object to hand to the guide,SVI/MCMC/ blackjax, forecast(), predict_in_sample(), to_datatree() and themodel_fnof backtest().
Raises
ValueError-
If
modelis already the result of time_reparam(). Nesting is not supported: the inner handler would turn the site into a deterministic before the outer one sees it, so the outer transform would silently be a no-op.
Notes
- Create the wrapped model once and reuse it: the drivers jit-compile with the model as a static argument, so wrapping again inside a loop recompiles.
- Apply it innermost.
scope(time_reparam(model), prefix="a")yieldsa/drift_haar;time_reparam(scope(model, "a"))yieldsa/a/drift_haarbecause the auxiliary sample passes throughscopea second time (generic NumPyro reparam-under-scope behavior). - It composes with the per-block
reparam=hook: afterinnovations(..., reparam=LocScaleReparam(0))the site under the plate isdrift_decentered, so the auxiliary site isdrift_decentered_haarand bothdrift_decenteredanddriftbecome deterministic. This matches Pyro, whose config applies to every site in the plate. - Untouched sites: the
_futuresuffix sites (they stay prior-drawn undertime_future, so forecast() is unaffected), the scan sites of markov_series(), the error sites of ssoe(), observed sites and discrete sites. - Posterior dictionaries from draw_posterior() and
mcmc.get_samples()contain bothdrift(deterministic) anddrift_haar;Predictivesubstitutes only the latter and recomputes the former.init_to_valuemust therefore targetdrift_haar. - Measured on random-walk level models: mean-field
AutoNormalreaches a better ELBO for both transforms (DCT slightly ahead of Haar), but an optimizer schedule tuned for the original coordinates does not transfer as is. The rotated posterior is better conditioned, so it tolerates and may need a larger learning rate to converge within the same step budget; a schedule that is too small stalls with part of the intercept still held by the level. NUTS trajectories become cheaper (fewer leapfrog steps per iteration) while the effective sample size per draw is model dependent.
Examples
model_dct = time_reparam(seasonal_model, "dct")
guide = AutoNormal(model_dct)
svi = SVI(model_dct, guide, Adam(0.01), Trace_ELBO())
svi_result = svi.run(key_fit, 1_500, covariates[:t_obs], data, progress_bar=False)
posterior = draw_posterior(key_post, guide, svi_result.params, num_samples=100)
samples = forecast(key_pred, model_dct, posterior, data, covariates)