var.var_mean()

Compute the VAR conditional mean of the next row from a lag window.

Usage

Source

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); None for 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)