Skip to content

DiscreteTimeStateEvolution

Bases: Module

Discrete-time state evolution via Markov transition distributions.

The next state is drawn from a conditional distribution given the current state, control, and time indices:

\[ x_{t_{k+1}} \sim p\left(x_{t_{k+1}} \mid x_{t_k}, u_{t_k}, t_k, t_{k+1}\right) \]

Implementations must return a NumPyro-compatible distribution (e.g., numpyro.distributions.Distribution). Most transitions provide sampling, moments, and log_prob; explicitly sample-only transitions are also valid for simulators, ensemble filters, and bootstrap particle filters that never evaluate the transition density. Such transitions should raise a targeted error when an unavailable moment or density is requested.

Parameters:

Name Type Description Default
x State

Current state \(x \in \mathbb{R}^{d_x}\).

required
u Control | None

Current control input or None.

required
t_now Time

Current time index \(t_k\).

required
t_next Time

Next time index \(t_{k+1}\) (for non-uniform sampling or continuous-time embeddings).

required

Returns:

Type Description

numpyro.distributions.Distribution: Distribution over the next state \(x_{t_{k+1}}\). In practice this should be a numpyro.distributions.Distribution instance.

All-pairs transition scores

Use vmap to score every pair of previous and current states:

import jax


def score_from(previous_state):
    transition = dynamics.state_evolution(
        x=previous_state, u=control, t_now=t_now, t_next=t_next
    )
    return jax.vmap(transition.log_prob, out_axes=-1)(states)


pairwise_log_prob = jax.vmap(score_from, out_axes=-2)(previous_states)

The result has shape (*log_prob_batch, n_previous, n_current), preserving distribution batch axes. Dense evaluation requires work and output storage proportional to n_previous * n_current per batch member.

For joint state-path and observation scoring, pass the complete discrete path to log_prob as state_path_params.