Skip to content

Closed-loop control

Dynestyx can combine simulation, observation, state estimation, and closed-loop (or online) control for a single discrete-time trajectory. Which control an observation sees is set by the model's observation_control_alignment (the same field that governs open-loop simulation). Unlike open-loop control, closed-loop control defaults to "previous_transition" and raises a warning if the alignment is not specified. "same_time" with state estimation is currently not supported; its support is tracked in Issue #372.

Both conventions (with and without state estimation) are described below.

In the following, \(s_0\) is simulate's initial_policy_state: None by default, for a stateless policy. It must be initialized and passed explicitly if the policy requires it. By default, the loop runs with state estimation; pass use_true_state=True to run it on the true state instead.

previous_transition convention (default for closed-loop control)

In this convention, the control \(u_k\) drives the transition into the next state \(x_{k+1}\) and generates the observation \(y_{k+1}\). Hence, \(y_0\) never exists.

With the true state

On \(\text{Times} = [t_0, \dots, t_N]\):

\[ \begin{aligned} &x_0 \sim p_0, \quad s_0 \text{ given} && \text{Initialization step} \\ &\text{for } k = 0, \dots, N-1: \\ &\quad u_k, s_{k+1} = \pi(x_k, t_k, t_{k+1}, s_k) && \text{Select the control} \\ &\quad x_{k+1} \sim p(x_{k+1} \mid x_k, u_k, t_k, t_{k+1}) && \text{State transition} \\ &\quad y_{k+1} \sim p(y_{k+1} \mid x_{k+1}, u_k, t_{k+1}) && \text{Emit observation} \end{aligned} \]

At each step, the true state \(x_k\) is passed to the policy as a Delta distribution (its .mean is the state itself).

With an estimated state

Notation:

\[ \begin{aligned} \tilde{p}_k &\approx p(x_k \mid y_1, \dots, y_{k-1},\ u_0, \dots, u_{k-1}) && \text{is the predicted distribution,} \\ \hat{p}_k &\approx p(x_k \mid y_1, \dots, y_k,\ u_0, \dots, u_{k-1}) && \text{is the filtered distribution.} \end{aligned} \]

The loop is:

\[ \begin{aligned} & \text{Times} = [t_0, \dots, t_{N}]\\ &x_0 \sim p_0, \quad \hat{p}_0 = p_0, \quad s_0 \text{ given} \quad \text{Initialization step} \\ &\text{for } k = 0,\dots N-1:\\ &\quad u_{k}, s_{k+1} =\pi(\hat{p}_{k}, t_{k}, t_{k+1}, s_k) \quad \text{Select the control}\\ &\quad x_{k+1} \sim p(x_{k+1} \mid x_k, u_k, t_k, t_{k+1}) \quad \text{State transition} \\ &\quad y_{k+1} \sim p(y_{k+1} \mid x_{k+1}, u_k, t_{k+1}) \quad \text{Emit observation}\\ &\quad \hat{p}_{k+1} = \text{FilterUpdate}(y_{k+1}, \hat{p}_k, u_k) \quad \text{Update the filtering distribution using the observation} \end{aligned} \]

At each step, an estimate of the state, the filtered distribution \(\hat{p}_k\), is passed to the policy as a Distribution.

In the previous_transition convention:

  1. There are \(N+1\) states \(x_0, x_1, \dots, x_{N}\) on the time grid \([t_0, \dots, t_N]\) (reported in the results as times).
  2. There are \(N\) controls \(u_0, u_1, \dots, u_{N-1}\) on the time grid \([t_0, \dots, t_{N-1}]\) (reported in the results as ctrl_times).
  3. There are \(N\) observations \(y_1, \dots, y_{N}\) on the time grid \([t_1, \dots, t_{N}]\) (reported in the results as obs_times).

same_time convention

In this convention, the control \(u_k\) drives the transition into the next state \(x_{k+1}\) and generates the observation \(y_k\). Hence, the last observation \(y_{N}\) never exists.

With the true state

On \(\text{Times} = [t_0, \dots, t_N]\):

\[ \begin{aligned} &x_0 \sim p_0, \quad s_0 \text{ given} && \text{Initialization step} \\ &\text{for } k = 0, \dots, N-1: \\ &\quad u_k, s_{k+1} = \pi(x_k, t_k, t_{k+1}, s_k) && \text{Select the control} \\ &\quad y_k \sim p(y_k \mid x_k, u_k, t_k) && \text{Emit observation} \\ &\quad x_{k+1} \sim p(x_{k+1} \mid x_k, u_k, t_k, t_{k+1}) && \text{State transition} \end{aligned} \]

At each step, the true state \(x_k\) is passed to the policy as a Delta distribution (its .mean is the state itself).

With an estimated state (not implemented)

Notation:

\[ \begin{aligned} \tilde{p}_k &\approx p(x_k \mid y_0, \dots, y_{k-1},\ u_0, \dots, u_{k-1}) && \text{is the predicted distribution,} \\ \hat{p}_k &\approx p(x_k \mid y_0, \dots, y_k,\ u_0, \dots, u_k) && \text{is the filtered distribution.} \end{aligned} \]

The loop is:

\[ \begin{aligned} & \text{Times} = [t_0, \dots, t_{N}]\\ &x_0 \sim p_0, \quad \tilde{p}_0 = p_0, \quad s_0 \text{ given} \quad \text{Initialization step} \\ &\text{for } k = 0,\dots N-1:\\ &\quad u_{k}, s_{k+1} =\pi(\tilde{p}_{k}, t_{k}, t_{k+1}, s_k) \quad \text{Select the control}\\ &\quad y_{k} \sim p(y_k \mid x_k, u_k, t_k) \quad \text{Emit observation}\\ &\quad \hat{p}_k = \text{FilterAnalysis}(y_k, \tilde{p}_k, u_k) \quad \text{Update the filtering distribution using the observation} \\ &\quad x_{k+1} \sim p(x_{k+1} \mid x_k, u_k, t_k, t_{k+1}) \quad \text{State transition} \\ &\quad \tilde{p}_{k+1} = \text{PredictionUpdate}(\hat{p}_k, u_k) \quad \text{Predict the filtering distribution} \end{aligned} \]

At each step, an estimate of the state, the predicted distribution \(\tilde{p}_k\), is passed to the policy as a Distribution.

In the same_time convention:

  1. There are \(N+1\) states \(x_0, x_1, \dots, x_{N}\) on the time grid \([t_0, \dots, t_N]\) (reported in the results as times).
  2. There are \(N\) controls \(u_0, u_1, \dots, u_{N-1}\) on the time grid \([t_0, \dots, t_{N-1}]\) (reported in the results as ctrl_times).
  3. There are \(N\) observations \(y_0, y_1, \dots, y_{N-1}\) on the time grid \([t_0, \dots, t_{N-1}]\) (reported in the results as obs_times).

Telling the two apart

Both conventions return \(N+1\) states on \([t_0, \dots, t_N]\) and \(N\) controls on \([t_0, \dots, t_{N-1}]\). They differ in exactly one place:

states ctrl_times obs_times
"same_time" \([t_0 \dots t_N]\) \([t_0 \dots t_{N-1}]\) \([t_0 \dots t_{N-1}]\)
"previous_transition" \([t_0 \dots t_N]\) \([t_0 \dots t_{N-1}]\) \([t_1 \dots t_N]\)

Array lengths are therefore identical and cannot identify which convention produced a result. Read obs_times and ctrl_times off the result rather than inferring alignment from shapes. With state estimation, filtered_states_mean also differs: \(N+1\) beliefs under "previous_transition" (one per state), against \(N\) under "same_time". On the true state, it is None.

Simulator and policy protocol

dynestyx.control.discrete_controller_simulators.DiscreteControlLoopSimulator

Bases: BaseSimulator

Closed-loop simulator: simulate, observe, filter, and decide controls online.

Unlike DiscreteTimeSimulator, which requires the entire control trajectory as a pre-supplied ctrl_values array, DiscreteControlLoopSimulator computes each \(u_k\) online from the current belief via control_policy.

Which control an observation sees is set by the model's dynamics.observation_control_alignment -- the same field that governs open-loop simulation.

With convention "previous_transition", the control \(u_k\) is chosen after seeing the observation \(y_k\) (which uses the previous control). The policy therefore sees the filtered belief \(\hat p_k\). This is the only convention the filtered closed loop supports for now, and it is what an unspecified (None) field resolves to, with a warning.

With convention "same_time", the control \(u_k\) would be chosen before seeing the observation \(y_k\) it drives, so the policy would see the predicted belief \(\tilde p_k\). That needs separate prediction and analysis filter steps, which cuthbert does not expose; an explicit "same_time" raises NotImplementedError until they are built.

Both loops are written out in full on the Closed-loop control page.

With use_true_state=True the loop skips filtering altogether and hands the policy the true state \(x_k\). Both conventions run.

The one-step filter update runs on the cuthbert backend (filter_source="cuthbert"). See the filters page for the available filters, and build_cuthbert_filter in discrete_filter.py for which of them the online update supports. Plated controlled simulation is not yet supported; see Issue #318.

Attributes:

Name Type Description
control_policy

Control policy \(\pi\); see PolicyCallable. Its initial state \(s_0\) is exactly simulate's initial_policy_state argument (default None, for a stateless policy) -- control_policy is never introspected for an initial_state() method; a stateful policy's initial state must always be passed explicitly.

filter_config

Selects the filtering algorithm; any config build_cuthbert_filter accepts (see above). Defaults to _default_filter_config(dynamics) when None. The online one-step update currently requires filter_source="cuthbert". Its record_filtered_states_mean/record_max_elems fields gate whether the filtered_states_mean output is recorded, exactly as they do for Filter (see dynestyx.utils._should_record_field). Must not be given together with use_true_state=True, which filters nothing.

use_true_state

Give the policy the true state \(x_k\) instead of a filtered belief, and run no filter at all. Defaults to False. Observations are still emitted and returned, but nothing consumes them, and filtered_states_mean is always None.

n_simulations int

Currently only 1 is supported.

simulate(dynamics: DynamicalModel, *, rng_key: PRNGKeyArray, ctrl_times: Real[Array, ' ctrl_time'] | None = None, ctrl_values: Real[Array, 'ctrl_time control_dim'] | Real[Array, ' ctrl_time'] | None = None, predict_times: Real[Array, ' predict_time'] | None = None, initial_policy_state: PyTree | None = None, **kwargs: Any) -> ControlledSimulatedResult

Simulate one online controlled trajectory.

Parameters:

Name Type Description Default
dynamics DynamicalModel

Discrete-time dynamical model.

required
rng_key PRNGKeyArray

Root key for environment and fallback filter randomness.

required
ctrl_times Real[Array, ' ctrl_time'] | None

Unsupported because controls are selected online.

None
ctrl_values Real[Array, 'ctrl_time control_dim'] | Real[Array, ' ctrl_time'] | None

Unsupported because controls are selected online.

None
predict_times Real[Array, ' predict_time'] | None

Strictly increasing simulation times.

None
initial_policy_state PyTree | None

Initial state passed to control_policy.

None
**kwargs Any

Additional shared simulator-handler metadata, ignored here.

{}

Returns:

Type Description
ControlledSimulatedResult

States, observations, controls, filtered beliefs and policy states

ControlledSimulatedResult

on all predict_times.

ControlledSimulatedResult

Under the "same_time" convention, the states are of length len(predict_times), but the controls and observations are of length len(predict_times) - 1.

ControlledSimulatedResult

Under the previous_transition convention, the states are of length len(predict_times), but the controls and observations are of length len(predict_times) - 1.

Raises:

Type Description
ValueError

If inputs are incompatible with online discrete control.

NotImplementedError

If the requested simulation mode is unsupported, including an explicit observation_control_alignment="same_time" while filtering (use_true_state=False).

Warns:

Type Description
UserWarning

If dynamics.observation_control_alignment is unspecified (None); "previous_transition" is used.

dynestyx.control.discrete_controller_simulators.ControlledSimulatedResult

Bases: SimulatedResult

SimulatedResult extended with the control loop's extra outputs.

Registered as deterministic sites the same generic way as SimulatedResult's own fields.

Both conventions return \(N+1\) states and \(N\) controls, so the array lengths alone do not say which one ran: consult obs_times, which is \([t_0, \dots, t_{N-1}]\) under "same_time" and \([t_1, \dots, t_N]\) under "previous_transition".

dynestyx.control.discrete_controller_simulators.PolicyCallable

Bases: Protocol

Structural protocol for a control policy \(\pi\).

\[u_k, s_{k+1} = \pi(\tilde x_k, t_k, t_{k+1}, s_k)\]

\(\tilde x_k\) is the loop's current state or state estimate. Which one it is depends on use_true_state and on the convention.

When using use_true_state=True, \(\tilde x_k\) is a Delta distribution centered on the true state \(x_k\) (including at \(k=0\)).

When using use_true_state=False (default), this is a distribution depending on the filter configuration:

  • under same_time, \(\tilde x_k\) is the predicted state \(\hat{x}_{k|k-1}\)
  • under previous_transition, \(\tilde x_k\) is the filtered state \(\hat{x}_{k|k}\)

where \(\hat{x}_{k|j}\) is the state estimate at time \(t_k\) given observations up to time \(t_j\). At \(k=0\), this is the model's initial-state distribution \(p_0\).

With use_true_state=True there is no filtering: \(\tilde x_k\) is a Delta centered on the true state \(x_k\).

It is always a NumPyro Distribution. Use x_hat.mean for a family-agnostic point estimate. t_now/t_next are the current and next time points. Any plain callable matching this signature works, including an equinox.Module with a matching __call__ or a plain Python function. control_policy must return a concrete value.

s is an internal state for the policy (such as a PRNG key), which is passed back to it on the next call. A stateless policy can ignore it and return None for the next state.

MPPI policy

dynestyx.control.mppi.MPPI

Bases: Module

Model Predictive Path Integral (MPPI) controller.

At each call: sample n_samples candidate control sequences of length horizon as Gaussian perturbations around a nominal sequence (the policy state s, warm-started from the previous call), roll each one forward horizon steps through dynamics.state_evolution (built internally -- the caller only ever supplies the one-step dynamics, never a hand-written rollout), score the resulting trajectories with loss_fn, and combine them via the standard MPPI weighting

\[w_i \propto \exp(-\mathrm{loss}_i / \lambda), \qquad u_{0:H-1} = \sum_i w_i\, u^{(i)}_{0:H-1}\]

i.e. a softmax over the (negated, temperature-scaled) per-sample losses. Only the first control of that weighted-mean sequence is applied this step (receding horizon); the remainder becomes next step's nominal sequence, shifted left by one with the last entry repeated.

Each rollout is run under the "previous_transition" observation/ control convention, so a candidate's \(u_k\) influences \(x_{k+1}\) and \(y_{k+1}\). dynamics is copied (via equinox.tree_at).

Attributes:

Name Type Description
dynamics DynamicalModel

a DynamicalModel (the same model used for the real simulation or some approximate). Each candidate rollout is computed by calling dsx.simulate. If dynamics holds trainable parameters you're also fitting via the outer simulation, they remain in the differentiable pytree so gradients through planning are tracked too.

loss_fn MPPILossFn

MPPILossFn, i.e. (result: SimulatedResult) -> scalar, called once per sample (vmapped) on that candidate's full rollout. Every field carries a leading n_simulations axis -- e.g. result.states.shape == (n_simulations, horizon, state_dim), so (1, horizon, state_dim) by default. times/states/observations/controls all have length horizon and are index-aligned: at index k, states[k] is \(x_{k+1}\), observations[k] is \(y_{k+1}\), and controls[k] is \(u_k\) -- the control that produced that state. The starting state \(x_0\) is not in states (no control produced it); it is available separately as result.x_0, shape (1, state_dim).

horizon int

Planning horizon length H -- the number of internal one-step dynamics calls per rollout. Defaults to 10.

noise_std Real[Array, ''] | Real[Array, ' control_dim']

Standard deviation of the Gaussian perturbations added to the nominal sequence, scalar or shape (control_dim,). Defaults to 1.0.

n_samples int

Number of sampled control sequences per call. Defaults to 20.

n_simulations int

Number of independent rollouts drawn per candidate control sequence, forwarded to dsx.simulate. Defaults to 1.

dt float

Fixed planning step size. Defaults to 1.0.

temperature float

MPPI's \(\lambda\); higher values flatten the weights toward a uniform average, lower values concentrate weight on the lowest-loss samples.

batched bool

Whether the n_samples candidate rollouts are computed with jax.vmap (default, fast, requires dynamics.state_evolution to be vmap-compatible) or jax.lax.map (a sequential loop -- slower, but works for a dynamics.state_evolution that isn't vmap-compatible, e.g. wraps an external simulator via jax.pure_callback).

seed int

Seeds MPPI's own PRNG key, carried inside the policy state s (as (nominal_sequence, key)) and split internally on every call.

initial_state() -> tuple[Real[Array, 'horizon control_dim'], PRNGKeyArray]

Zero nominal control sequence plus MPPI's own seeded PRNG key. Pass this call's result as initial_policy_state to simulate/ dsx.simulate -- DiscreteControlLoopSimulator never calls this automatically, so it must be supplied explicitly.

plan_step(x_hat: Distribution, t_now: Real[Array, ''], s: tuple[Real[Array, 'horizon control_dim'], PRNGKeyArray]) -> tuple[Real[Array, ' control_dim'], tuple[Real[Array, 'horizon control_dim'], PRNGKeyArray], SimulatedResult]

Do MPPI's full planning step and also return the batch of every candidate rollout considered (n_samples-wide SimulatedResult). -- useful for debugging/plotting what MPPI weighed, or diagnosing a loss_fn.

__call__ (used by DiscreteControlLoopSimulator) is a thin wrapper around this that drops the rollout batch, since PolicyCallable's return signature can't carry a third value.

Every field is shaped (n_samples, n_simulations, horizon, ...). predicted_* are always None (not meaningful for a planning rollout).

Policy helpers

dynestyx.control.discrete_controller_simulators.filter_state_mean(state: Any) -> Real[Array, ...]

Point-estimate summary of a cuthbert filter state, any family.

Kalman-family states (KFConfig, EKFConfig, EnKFConfig) expose a .mean property directly. PFConfig states (ParticleFilterState) have no such property -- they represent the belief as a weighted particle cloud (.particles, .log_weights), so the point estimate is the weighted mean instead. Broadcasts over any leading batch/time axis, so it works on both a single belief and a whole scanned-out sequence of them.

dynestyx.control.discrete_controller_simulators.filter_state_dist(state: Any, filter_config: BaseFilterConfig) -> Distribution

Full-belief NumPyro distribution for a filter state.

Only filter_source="cuthbert" is supported today, matching DiscreteControlLoopSimulator's own restriction. Converts the filter state to a NumPyro Distribution in the same way that ConditionedResult.dists does for the recorded filtered states.

The shared conversion is time-indexed, so the state is given a leading axis of length one and the single distribution unwrapped.

Parameters:

Name Type Description Default
state Any

A single, unbatched filter state produced by filter_config.

required
filter_config BaseFilterConfig

The config that produced state. Selects the backend and carries recorded_filtered_states_cov_jitter for ensemble states.

required

Returns:

Type Description
Distribution

The belief as a NumPyro Distribution.

Raises:

Type Description
ValueError

If filter_config.filter_source is not "cuthbert".