models.ssoe()
Run a single-source-of-error recursion over the full horizon.
Usage
models.ssoe(
h,
name,
y,
init_carry,
mean,
update,
noise_dist,
xs=None,
)The building block for innovations state-space models (ARMA, exponential smoothing, Croston/TSB levels, censored autoregressions): a deterministic filter whose state is driven by the one-step-ahead error eps_t = y_t - mu_t. In-sample it runs mean and update in a raw jax.lax.scan over the observed series y (no sample sites inside); when forecasting it draws iid future errors at the site f"{name}_future" from noise_dist under a plate("time_future", h.future) and runs a second scan from the final in-sample carry with y_t = mu_t + eps_t fed back through update. The guide never sees the future site, because fitting always happens with future == 0. Linear-Gaussian members (ARMA, additive exponential smoothing) can be marginalized exactly by a Kalman filter; the error-feedback form is the one that also covers the nonlinear members.
The block registers nothing but the error site. The caller writes the likelihood against r.mu and registers numpyro.deterministic("forecast", r.y_future) when h.future > 0 (an unconditional registration is harmless to forecast(), but lands a size-0 variable in every posterior). Driver contract: predict_in_sample() and to_datatree() call the model with data=None and read "obs", so y must come from covariates or be computed in the model, never from h.data; forecast() reads "forecast".
Frozen gates. Route an update gate (Croston’s demand indicator, an availability mask) through an xs leaf frozen over the horizon with pad_future(), and read it from x_t, never from y_t (over the horizon y_t = mu_t + eps_t is nonzero). With the gate off, update returns the carry unchanged and the forecast is the last level plus iid errors. backtest() and forecast() hand the model real future covariate rows, so a gate sliced from the full covariates keeps updating on sampled values and leaks the test window; scenario inputs such as a future availability mask are the only thing to read from those rows.
Shapes. Rows are (*batch, obs): a scalar state emits mu[None] and starts from init[None] (the block refuses a scalar or a wider mean because either would silently broadcast the likelihood into a (t, t) log-prob); a tuple carry with scalar leaves reads eps_t[0] (the ETS idiom). A panel puts the series on the observation axis: y is (t_obs, series), the carry (series,), and a noise sampled under plate("series") has exactly the batch shape (series,) the block needs. Batch dims to the left of time ((B, t_obs, obs)) take a (B, obs) carry and a (B, 1, obs) noise batch. Errors correlated across the observation axis are a multivariate noise_dist with event shape (obs,), which is how a vector autoregression enters the block (see var_step()). Inputs are jax Arrays (the import hook rejects NumPy). With obs == 1 a noise batch shape (future, 1) is indistinguishable from time and is consumed as such: per-step error scales, if that is what you meant.
Composition. Two channels are two calls sharing the same Horizon; each opens its own time_future plate. Scoping a helper that contains the call is fine (handlers.scope prefixes the error site and the plate; use name="eps" inside a scoped helper so the site reads z_eps_future); register "obs" and "forecast" outside any scope and build the forecast from the channels’ y_future.
Parameters
h: Horizon-
The horizon for the current model call (see Horizon).
name: str-
Base name of the error site; the future errors are drawn at
f"{name}_future". y: Array | None-
The driving series over the observed window, shape
(*batch, t_obs, obs)with time at axis-2(integer counts are fine as long as the carry, hence the mean, is floating; the error promotes). Sliced fromcovariatesor computed in the model;None(the value ofh.dataunderdata=None) raises. init_carry: Carry-
Initial carry, any PyTree, already broadcast to the
(*batch, obs)rows: a scalar level isinit[None], a panel level(series,). Every leaf must keep its shape and dtype throughupdate. mean: SSOEMean[Carry]-
(carry, x_t) -> mu_t(see SSOEMean): the mean for the current row (shape(*batch, obs), so a scalar state emitsmu[None]). update: SSOEUpdate[Carry]-
(carry, y_t, eps_t, x_t) -> carry(see SSOEUpdate): the next carry from the row’s value and error. Over the horizon it receives the drawneps_t(not a recomputedy_t - mu_t, which can differ by an ulp); when the update needs the mean, callmean(carry, x_t)inside it. noise_dist: dist.Distribution-
Zero-centered per-step error distribution, either an elementwise family (event rank 0) or a multivariate family over the observation axis (event rank 1, e.g.
dist.MultivariateNormal(jnp.zeros(obs), scale_tril=L)for shocks correlated across series, the VAR case). For a(t_obs, obs)series the batch shape is(obs,)(()is fine whenobs == 1) for the elementwise form and()for the multivariate form; a batched(B, t_obs, obs)series takes(B, 1, obs)and(B, 1)respectively. Either way the draw under the time plate (dim=-2for event rank 0,dim=-1for event rank 1) is exactly(*batch, future, obs)with the dtype of the means; event rank 2 or higher is rejected. xs: PyTree[Array] | None = None-
Optional exogenous inputs over the full horizon: a PyTree of arrays with time at axis
-2anddurationrows (a single array, a tuple, a dict, …), split ath.t_obsand handed to mean andupdaterow by row asx_t;Nonefor autonomous dynamics.
Returns
SSOEResult-
mu(in-sample means),mu_futureandy_future(forecast means and sampled values; size-0 time axis while training).
Raises
ValueError-
If
yisNone, lacks the time or observation axis, or does not cover exactlyh.t_obsrows; if anxsleaf lacks the axes or does not spanh.durationrows; if mean returns a value without the observation axis orupdatea carry with a different tree structure, shape or dtype; if mean orupdatecallsnumpyro.sample; ifnoise_disthas event rank 2 or higher, or has event rank 1 inside an enclosing plate atdim=-1; or if it draws errors of the wrong shape or dtype.
Examples
ARMA(1,1) (y is the observed series routed through covariates):
>>> def mean(carry, _):
... y_prev, eps_prev = carry
... return mu + phi * y_prev + theta * eps_prev
>>> def update(carry, y_t, eps_t, _):
... return y_t, eps_t
>>> r = ssoe(h, "eps", y, (mu[None], jnp.zeros((1,))), mean, update, dist.Normal(0.0, sigma))
>>> numpyro.sample("obs", dist.Normal(r.mu, sigma), obs=h.data)
>>> if h.future > 0:
... numpyro.deterministic("forecast", r.y_future)A gated level (Croston, TSB) with the gate frozen over the horizon:
>>> def mean(level, _):
... return level
>>> def update(level, y_t, _, gate_t):
... return jnp.where(gate_t, alpha * y_t + (1 - alpha) * level, level)
>>> gate_full = pad_future(gate, h.future)
>>> r = ssoe(h, "eps", y, init[None], mean, update, dist.Normal(0.0, noise), xs=gate_full)