cuthbert API Reference
- General filtering API
- General smoothing API
- Discrete hidden Markov models
- Ensemble Kalman
- Factorial state-space models
- Gaussian filters and smoothers
- Sequential Monte Carlo
Unified Interface
All inference methods are implemented with the following unified interface:
from jax import tree
import cuthbert
# Define model_inputs
init_model_inputs = ...
filter_model_inputs = ...
# Load inference method
kalman_filter = cuthbert.gaussian.kalman.build_filter(
get_init_params=get_init_params,
get_dynamics_params=get_dynamics_params,
get_observation_params=get_observation_params,
) # build_filter function takes all inference-specific arguments, swap this out for different inference methods.
# Online inference
state = kalman_filter.init_prepare(init_model_inputs)
for t in range(T):
model_inputs_t = tree.map(lambda x: x[t], filter_model_inputs)
prepare_state = kalman_filter.filter_prepare(model_inputs_t)
state = kalman_filter.filter_combine(state, prepare_state)
Or for offline inference:
kalman_smoother = cuthbert.gaussian.kalman.build_smoother(get_dynamics_params)
init_state = kalman_filter.init_prepare(init_model_inputs)
filter_states = cuthbert.filter(kalman_filter, filter_model_inputs, init_state)
smoother_states = cuthbert.smoother(kalman_smoother, filter_states, filter_model_inputs)