var.var_mean()
Compute the VAR conditional mean of the next row from a lag window.
Usage
var.var_mean(
phi,
lags,
intercept=None,
)\mu_t = c + \sum_{l=1}^{p} \Phi_l \, y_{t-l},
with phi[..., l - 1, :, :] holding \Phi_l and lags[..., -l, :] holding y_{t-l} (most recent row last).
Parameters
phi: Float[Array, " *#batch lags obs obs"]-
Coefficient tensor
(*batch, lags, obs, obs); see the module docstring for the index convention. lags: Float[Array, " *#batch lags obs"]-
Lag window
(*batch, lags, obs)in natural time order. intercept: Float[Array, " *#batch obs"] | None = None-
Optional intercept c of shape
(*batch, obs);Nonefor a zero-mean recursion (the impulse response case).
Returns
Float[Array, "*batch obs"]- The conditional mean row, broadcast over the batch axes of the inputs.
Examples
A latent VAR under markov_series(): the transition returns the next-row distribution and advance shifts the lag window, with phi and the Cholesky factor scale_tril sampled outside:
def transition(window, _):
return dist.MultivariateNormal(var_mean(phi, window), scale_tril=scale_tril)
def advance(window, z_t, _):
return jnp.concatenate([window[..., 1:, :], z_t[..., None, :]], axis=-2)
z = markov_series(h, "z", jnp.zeros((n_lags, n_obs)), transition, advance=advance)