From 4e1eaa109e62a447b67c681bd19bf64d7f34c9d5 Mon Sep 17 00:00:00 2001 From: Matthew Levine Date: Fri, 31 Jul 2026 11:24:50 -0400 Subject: [PATCH] 1-shot un-edited codex update from PR #279 for SLDS support --- .../inference/configs/filter_configs.md | 2 + .../specialized/mixed_state_distribution.md | 7 + .../switching_linear_gaussian_observation.md | 11 + ...itching_linear_gaussian_state_evolution.md | 11 + dynestyx/__init__.py | 6 + dynestyx/distributions.py | 315 +++++++++++++++ dynestyx/inference/configs/filter.py | 64 ++- dynestyx/inference/filters.py | 15 +- .../cd_dynamax/discrete_filter.py | 288 +++++++++++++- .../inference/utils/distribution_utils.py | 21 + dynestyx/inference/utils/numpyro_sites.py | 48 +++ dynestyx/models/__init__.py | 6 + dynestyx/models/observations.py | 75 ++++ dynestyx/models/state_evolution.py | 91 +++++ dynestyx/utils.py | 16 +- mkdocs.yml | 5 + pyproject.toml | 4 +- tests/test_slds_rbpf.py | 371 ++++++++++++++++++ 18 files changed, 1330 insertions(+), 26 deletions(-) create mode 100644 docs/api_reference/public/models/specialized/mixed_state_distribution.md create mode 100644 docs/api_reference/public/models/specialized/switching_linear_gaussian_observation.md create mode 100644 docs/api_reference/public/models/specialized/switching_linear_gaussian_state_evolution.md create mode 100644 dynestyx/distributions.py create mode 100644 tests/test_slds_rbpf.py diff --git a/docs/api_reference/public/inference/configs/filter_configs.md b/docs/api_reference/public/inference/configs/filter_configs.md index df5d650d..188b3f66 100644 --- a/docs/api_reference/public/inference/configs/filter_configs.md +++ b/docs/api_reference/public/inference/configs/filter_configs.md @@ -11,6 +11,7 @@ The single `Filter()` handler is directed to the appropriate filtering algorithm | `EKFConfig` | Discrete | Nonlinear, differentiable Gaussian dynamics, nonlinear (and with `cuthbert`, non-Gaussian) but differentiable observations (approximate). | | `UKFConfig` | Discrete | Nonlinear, differentiable Gaussian dynamics, nonlinear but differentiable Gaussian observations (approximate). Generally more accurate, but slower than `EKFConfig`. | | `PFConfig` | Discrete | Applicable for arbitrary state-space models, but quite expensive and noisy estimates (asymptotically exact in the limit of infinite particles, approximate in practice). | +| `RBPFConfig` | Discrete (SLDS) | Switching linear-Gaussian models; samples regimes while marginalizing the continuous state with Kalman updates. | | `HMMConfig` | Discrete (HMM) | Finite discrete latent state space (exact & optimal). | | `ContinuousTimeKFConfig` | Continuous-discrete | Linear-Gaussian SDE + linear-Gaussian observations (exact and optimal). | | `ContinuousTimeEKFConfig` | Continuous-discrete | Mildly nonlinear SDE with differentiable drift and difussion terms; Gaussian observations (approximate). | @@ -28,6 +29,7 @@ The single `Filter()` handler is directed to the appropriate filtering algorithm - EKFConfig - UKFConfig - PFConfig + - RBPFConfig - EnKFConfig ## Continuous Time Configuration Classes diff --git a/docs/api_reference/public/models/specialized/mixed_state_distribution.md b/docs/api_reference/public/models/specialized/mixed_state_distribution.md new file mode 100644 index 00000000..a6ec3c48 --- /dev/null +++ b/docs/api_reference/public/models/specialized/mixed_state_distribution.md @@ -0,0 +1,7 @@ +# MixedStateDistribution + +::: dynestyx.distributions.MixedStateDistribution + options: + show_root_heading: false + show_root_toc_entry: false + diff --git a/docs/api_reference/public/models/specialized/switching_linear_gaussian_observation.md b/docs/api_reference/public/models/specialized/switching_linear_gaussian_observation.md new file mode 100644 index 00000000..502359ac --- /dev/null +++ b/docs/api_reference/public/models/specialized/switching_linear_gaussian_observation.md @@ -0,0 +1,11 @@ +# SwitchingLinearGaussianObservation + +::: dynestyx.models.observations.SwitchingLinearGaussianObservation + options: + show_root_heading: false + show_root_toc_entry: false + +See the [switching linear dynamical systems +guide](../../../../tutorials/state_space_models/slds_rbpf.md) for a complete +simulation and filtering example. + diff --git a/docs/api_reference/public/models/specialized/switching_linear_gaussian_state_evolution.md b/docs/api_reference/public/models/specialized/switching_linear_gaussian_state_evolution.md new file mode 100644 index 00000000..1a71c267 --- /dev/null +++ b/docs/api_reference/public/models/specialized/switching_linear_gaussian_state_evolution.md @@ -0,0 +1,11 @@ +# SwitchingLinearGaussianStateEvolution + +::: dynestyx.models.state_evolution.SwitchingLinearGaussianStateEvolution + options: + show_root_heading: false + show_root_toc_entry: false + +See the [switching linear dynamical systems +guide](../../../../tutorials/state_space_models/slds_rbpf.md) for a complete +simulation and filtering example. + diff --git a/dynestyx/__init__.py b/dynestyx/__init__.py index d4b546a7..93c6cb72 100644 --- a/dynestyx/__init__.py +++ b/dynestyx/__init__.py @@ -33,9 +33,12 @@ LinearGaussianStateEvolution, LTI_continuous, LTI_discrete, + MixedStateDistribution, ObservationModel, ScalarDiffusion, StochasticContinuousTimeStateEvolution, + SwitchingLinearGaussianObservation, + SwitchingLinearGaussianStateEvolution, ) from dynestyx.observation_missingness import ( MissingObservationMetadata, @@ -66,6 +69,7 @@ "LTI_discrete", "LinearGaussianParams", "LinearGaussianStateEvolution", + "MixedStateDistribution", "GaussianStateEvolution", "Discretizer", "ObservationModel", @@ -85,6 +89,8 @@ "DiracIdentityObservation", "LinearGaussianObservation", "LinearGaussianObservationParams", + "SwitchingLinearGaussianObservation", + "SwitchingLinearGaussianStateEvolution", "GaussianObservation", "ODESimulatorConfig", "SDESimulatorConfig", diff --git a/dynestyx/distributions.py b/dynestyx/distributions.py new file mode 100644 index 00000000..8c62959e --- /dev/null +++ b/dynestyx/distributions.py @@ -0,0 +1,315 @@ +"""Probability distributions used by specialized dynestyx models.""" + +from __future__ import annotations + +import jax +import jax.numpy as jnp +import jax.random as jr +import numpyro.distributions as dist +from jax import lax +from jaxtyping import Array, Float, Int, PRNGKeyArray, Real +from numpyro.distributions import constraints + + +class MixedStateDistribution(dist.Distribution): + r"""Joint distribution for the discrete and continuous state of an SLDS. + + A switching linear dynamical system (SLDS) has a regime \(z\) and a + continuous state \(x\): + + \[ + z \sim \operatorname{Categorical}(\pi), \qquad + x \mid z \sim \mathcal{N}(\mu_z, \Sigma_z). + \] + + Dynestyx represents the joint state as the homogeneous JAX vector + ``[z, *x]``. The regime is therefore encoded in the first floating-point + entry, although it is sampled and scored as a categorical integer. + + Leading batch dimensions are supported. The regime axis must be the final + batch axis of ``continuous_locs`` and ``continuous_covariances``. + + Args: + categorical_probs: Regime probabilities with shape + ``(*batch, num_regimes)``. + continuous_locs: Conditional means with shape + ``(*batch, num_regimes, state_dim)``. + continuous_covariances: Conditional covariance matrices with shape + ``(*batch, num_regimes, state_dim, state_dim)``. + validate_args: Whether NumPyro should validate samples and parameters. + + Attributes: + categorical_probs: Regime probabilities. + continuous_locs: Regime-conditional continuous-state means. + continuous_covariances: Regime-conditional continuous-state + covariances. + """ + + arg_constraints = { + "categorical_probs": constraints.simplex, + "continuous_locs": constraints.real, + "continuous_covariances": constraints.positive_definite, + } + support = constraints.real_vector + pytree_data_fields = ( + "categorical_probs", + "continuous_locs", + "continuous_covariances", + ) + pytree_aux_fields = ("_batch_shape", "_event_shape") + + categorical_probs: Float[Array, "*batch num_regimes"] + continuous_locs: Float[Array, "*batch num_regimes state_dim"] + continuous_covariances: Float[Array, "*batch num_regimes state_dim state_dim"] + + def __init__( + self, + categorical_probs: Float[Array, "*batch num_regimes"], + continuous_locs: Float[Array, "*batch num_regimes state_dim"], + continuous_covariances: Float[Array, "*batch num_regimes state_dim state_dim"], + *, + validate_args: bool | None = None, + ) -> None: + if continuous_locs.ndim < 2: + raise ValueError( + "continuous_locs must have shape (*batch, num_regimes, state_dim)." + ) + if continuous_covariances.ndim < 3: + raise ValueError( + "continuous_covariances must have shape " + "(*batch, num_regimes, state_dim, state_dim)." + ) + + num_regimes = categorical_probs.shape[-1] + state_dim = continuous_locs.shape[-1] + if continuous_locs.shape[-2] != num_regimes: + raise ValueError( + "categorical_probs and continuous_locs disagree on " + f"num_regimes: {num_regimes} != {continuous_locs.shape[-2]}." + ) + if continuous_covariances.shape[-3:] != ( + num_regimes, + state_dim, + state_dim, + ): + raise ValueError( + "continuous_covariances must end in " + f"({num_regimes}, {state_dim}, {state_dim}); got " + f"{continuous_covariances.shape[-3:]}." + ) + + batch_shape = lax.broadcast_shapes( + categorical_probs.shape[:-1], + continuous_locs.shape[:-2], + continuous_covariances.shape[:-3], + ) + probs = jnp.broadcast_to(categorical_probs, batch_shape + (num_regimes,)) + locs = jnp.broadcast_to(continuous_locs, batch_shape + (num_regimes, state_dim)) + covariances = jnp.broadcast_to( + continuous_covariances, + batch_shape + (num_regimes, state_dim, state_dim), + ) + self.categorical_probs = probs + self.continuous_locs = locs + self.continuous_covariances = covariances + super().__init__( + batch_shape=batch_shape, + event_shape=(state_dim + 1,), + validate_args=validate_args, + ) + + @property + def num_regimes(self) -> int: + """Number of discrete regimes.""" + return int(self.categorical_probs.shape[-1]) + + @property + def continuous_state_dim(self) -> int: + """Dimension of the continuous part of the state.""" + return int(self.continuous_locs.shape[-1]) + + def sample( + self, + key: PRNGKeyArray, + sample_shape: tuple[int, ...] = (), + ) -> Real[Array, "*sample *batch joint_state_dim"]: + """Draw joint regime and continuous-state samples.""" + regime_key, state_key = jr.split(key) + regimes = dist.Categorical(probs=self.categorical_probs).sample( + regime_key, sample_shape + ) + component_samples = dist.MultivariateNormal( + self.continuous_locs, + covariance_matrix=self.continuous_covariances, + ).sample(state_key, sample_shape) + state_indices = jnp.broadcast_to( + regimes[..., None, None], + component_samples.shape[:-2] + (1, self.continuous_state_dim), + ) + states = jnp.take_along_axis(component_samples, state_indices, axis=-2)[ + ..., 0, : + ] + return jnp.concatenate( + (regimes[..., None].astype(states.dtype), states), + axis=-1, + ) + + def log_prob( + self, + value: Real[Array, "*sample *batch joint_state_dim"], + ) -> Float[Array, "*sample *batch"]: + """Evaluate the joint log density of ``[z, *x]``.""" + regimes = jnp.rint(value[..., 0]).astype(jnp.int32) + continuous_state = value[..., 1:] + regime_log_prob = dist.Categorical(probs=self.categorical_probs).log_prob( + regimes + ) + component_log_probs = dist.MultivariateNormal( + self.continuous_locs, + covariance_matrix=self.continuous_covariances, + ).log_prob(continuous_state[..., None, :]) + continuous_log_prob = jnp.take_along_axis( + component_log_probs, regimes[..., None], axis=-1 + )[..., 0] + return regime_log_prob + continuous_log_prob + + +class RaoBlackwellizedParticleDistribution(dist.Distribution): + r"""Gaussian-mixture posterior represented by Rao-Blackwellized particles. + + Each particle stores a discrete regime and a Gaussian conditional + distribution for the continuous state. This distribution is used for + filter-to-simulator posterior rollout; unlike a point-particle + approximation, sampling retains the within-particle Gaussian uncertainty. + + Args: + log_weights: Normalized or unnormalized particle log weights with shape + ``(*batch, num_particles)``. + regimes: Integer regime labels with shape + ``(*batch, num_particles)``. + continuous_locs: Particle-conditional means with shape + ``(*batch, num_particles, state_dim)``. + continuous_covariances: Particle-conditional covariances with shape + ``(*batch, num_particles, state_dim, state_dim)``. + validate_args: Whether NumPyro should validate samples and parameters. + """ + + arg_constraints: dict = {} + support = constraints.real_vector + pytree_data_fields = ( + "log_weights", + "regimes", + "continuous_locs", + "continuous_covariances", + ) + pytree_aux_fields = ("_batch_shape", "_event_shape") + + log_weights: Float[Array, "*batch num_particles"] + regimes: Int[Array, "*batch num_particles"] + continuous_locs: Float[Array, "*batch num_particles state_dim"] + continuous_covariances: Float[Array, "*batch num_particles state_dim state_dim"] + + def __init__( + self, + log_weights: Float[Array, "*batch num_particles"], + regimes: Int[Array, "*batch num_particles"], + continuous_locs: Float[Array, "*batch num_particles state_dim"], + continuous_covariances: Float[ + Array, "*batch num_particles state_dim state_dim" + ], + *, + validate_args: bool | None = None, + ) -> None: + state_dim = continuous_locs.shape[-1] + num_particles = continuous_locs.shape[-2] + if regimes.shape[-1] != num_particles or log_weights.shape[-1] != num_particles: + raise ValueError("All RBPF inputs must agree on num_particles.") + if continuous_covariances.shape[-3:] != ( + num_particles, + state_dim, + state_dim, + ): + raise ValueError( + "continuous_covariances must end in " + f"({num_particles}, {state_dim}, {state_dim})." + ) + batch_shape = lax.broadcast_shapes( + log_weights.shape[:-1], + regimes.shape[:-1], + continuous_locs.shape[:-2], + continuous_covariances.shape[:-3], + ) + weights = jnp.broadcast_to(log_weights, batch_shape + (num_particles,)) + regimes = jnp.broadcast_to(regimes, batch_shape + (num_particles,)) + locs = jnp.broadcast_to( + continuous_locs, batch_shape + (num_particles, state_dim) + ) + covariances = jnp.broadcast_to( + continuous_covariances, + batch_shape + (num_particles, state_dim, state_dim), + ) + self.log_weights = jax.nn.log_softmax(weights, axis=-1) + self.regimes = regimes + self.continuous_locs = locs + self.continuous_covariances = covariances + super().__init__( + batch_shape=batch_shape, + event_shape=(state_dim + 1,), + validate_args=validate_args, + ) + + def sample( + self, + key: PRNGKeyArray, + sample_shape: tuple[int, ...] = (), + ) -> Real[Array, "*sample *batch joint_state_dim"]: + """Draw a particle, then draw its conditional continuous state.""" + particle_key, state_key = jr.split(key) + particle_indices = dist.Categorical(logits=self.log_weights).sample( + particle_key, sample_shape + ) + component_samples = dist.MultivariateNormal( + self.continuous_locs, + covariance_matrix=self.continuous_covariances, + ).sample(state_key, sample_shape) + state_indices = jnp.broadcast_to( + particle_indices[..., None, None], + component_samples.shape[:-2] + (1, self.event_shape[0] - 1), + ) + states = jnp.take_along_axis(component_samples, state_indices, axis=-2)[ + ..., 0, : + ] + regimes = jnp.take_along_axis( + jnp.broadcast_to( + self.regimes, + sample_shape + self.regimes.shape, + ), + particle_indices[..., None], + axis=-1, + )[..., 0] + return jnp.concatenate( + (regimes[..., None].astype(states.dtype), states), + axis=-1, + ) + + def log_prob( + self, + value: Real[Array, "*sample *batch joint_state_dim"], + ) -> Float[Array, "*sample *batch"]: + """Evaluate the particle-mixture density of ``[z, *x]``.""" + regime = jnp.rint(value[..., 0]).astype(self.regimes.dtype) + continuous_state = value[..., 1:] + component_log_probs = dist.MultivariateNormal( + self.continuous_locs, + covariance_matrix=self.continuous_covariances, + ).log_prob(continuous_state[..., None, :]) + regime_matches = regime[..., None] == self.regimes + joint_component_log_probs = jnp.where( + regime_matches, + self.log_weights + component_log_probs, + -jnp.inf, + ) + return jax.scipy.special.logsumexp(joint_component_log_probs, axis=-1) + + +__all__ = ["MixedStateDistribution"] diff --git a/dynestyx/inference/configs/filter.py b/dynestyx/inference/configs/filter.py index 1590d935..9fa99ad5 100644 --- a/dynestyx/inference/configs/filter.py +++ b/dynestyx/inference/configs/filter.py @@ -12,6 +12,7 @@ ResamplingDifferentiableMethod = Literal["stop_gradient", "straight_through", "soft"] FilterEmissionOrder = Literal["zeroth", "first", "second"] FilterStateOrder = Literal["zeroth", "first", "second"] +RBPFProposal = Literal["prior", "optimal"] CuthbertOnlyFilterSource = Literal["cuthbert"] CDDynamaxOnlyFilterSource = Literal["cd_dynamax"] @@ -285,6 +286,44 @@ class PFConfig(BaseFilterConfig): filter_source: CuthbertOnlyFilterSource = "cuthbert" +@dataclasses.dataclass +class RBPFConfig(BaseFilterConfig): + r"""Rao-Blackwellized particle filter (RBPF) for an SLDS. + + The filter samples the discrete regime path while analytically + marginalizing the conditionally linear-Gaussian continuous state with + Kalman updates. It is therefore usually much more efficient than a + bootstrap particle filter over the full joint state. + + Use this config with a model composed of + `MixedStateDistribution`, + `SwitchingLinearGaussianStateEvolution`, and + `SwitchingLinearGaussianObservation`. + + Does not support missing observations (data cannot contain NaNs). + + Attributes: + n_particles: Number of particles over discrete regime paths. Defaults + to `1_000`. + proposal: Discrete proposal strategy. `"prior"` samples from the + regime transition prior and resamples below the configured + effective-sample-size threshold. `"optimal"` uses cd-dynamax's + observation-adapted optimal path. Defaults to `"optimal"`. + ess_threshold_ratio: Resample the prior-proposal filter when its + effective sample size falls below this fraction of + `n_particles`. Ignored by the optimal path. Defaults to `0.5`. + record_filtered_regime_probs: Save + \(p(z_t \mid y_{0:t})\) as a deterministic NumPyro site. + filter_source: Backend. Always `"cd_dynamax"`. + """ + + n_particles: int = 1_000 + proposal: RBPFProposal = "optimal" + ess_threshold_ratio: float = 0.5 + record_filtered_regime_probs: bool | None = None + filter_source: CDDynamaxOnlyFilterSource = "cd_dynamax" + + @dataclasses.dataclass class EKFConfig(BaseFilterConfig): r"""Extended Kalman Filter (EKF) for discrete-time models. @@ -675,6 +714,7 @@ class ContinuousTimeUKFConfig(UKFConfig, ContinuousTimeConfig): DiscreteTimeConfigs: tuple[type, ...] = ( EnKFConfig, PFConfig, + RBPFConfig, EKFConfig, KFConfig, UKFConfig, @@ -744,13 +784,17 @@ def _config_to_record_kwargs(config: BaseFilterConfig) -> dict: "record_log_filtered": config.record_log_filtered, "record_max_elems": config.record_max_elems, } - else: - return { - "record_filtered_states_mean": config.record_filtered_states_mean, - "record_filtered_states_cov": config.record_filtered_states_cov, - "record_filtered_states_cov_diag": config.record_filtered_states_cov_diag, - "record_filtered_particles": config.record_filtered_particles, - "record_filtered_log_weights": config.record_filtered_log_weights, - "record_filtered_states_chol_cov": config.record_filtered_states_chol_cov, - "record_max_elems": config.record_max_elems, - } + record_kwargs = { + "record_filtered_states_mean": config.record_filtered_states_mean, + "record_filtered_states_cov": config.record_filtered_states_cov, + "record_filtered_states_cov_diag": config.record_filtered_states_cov_diag, + "record_filtered_particles": config.record_filtered_particles, + "record_filtered_log_weights": config.record_filtered_log_weights, + "record_filtered_states_chol_cov": config.record_filtered_states_chol_cov, + "record_max_elems": config.record_max_elems, + } + if isinstance(config, RBPFConfig): + record_kwargs["record_filtered_regime_probs"] = ( + config.record_filtered_regime_probs + ) + return record_kwargs diff --git a/dynestyx/inference/filters.py b/dynestyx/inference/filters.py index ae84dc95..3ee891ea 100644 --- a/dynestyx/inference/filters.py +++ b/dynestyx/inference/filters.py @@ -34,6 +34,7 @@ KFConfig, PFConfig, PFResamplingConfig, + RBPFConfig, UKFConfig, ) from dynestyx.inference.hmm_filters import _filter_hmm, compute_hmm_filter @@ -58,6 +59,7 @@ _categorical_log_probs_to_dists, _cholesky_state_sequence_to_dists, _posterior_sequence_to_dists, + _rbpf_sequence_to_dists, ) from dynestyx.inference.utils.numpyro_sites import ( register_filter_sites, @@ -486,6 +488,7 @@ def compute_output(dyn, ot, ov, ovf, om, ct, cv, k): return compute_cd_dynamax_discrete_filter( dyn, config, + key=k, obs_times=ot, obs_values=ov, ctrl_times=ct, @@ -637,6 +640,11 @@ def compute_output_member(dyn, ot, ov, ovf, om, ct, cv, k, *idxs): ), ) if output_kind == "cd_dynamax_discrete": + if isinstance(config, RBPFConfig): + return _rbpf_sequence_to_dists( + outputs, + plate_shapes=plate_shapes, + ) return _posterior_sequence_to_dists( outputs, means_attr="filtered_means", @@ -678,8 +686,9 @@ def _filter_discrete_time( ) -> tuple[jax.Array | None, object | None, list[numpyro.distributions.Distribution]]: """Discrete-time marginal likelihood via cuthbert or cd-dynamax. - Filter type inferred from config class: KFConfig, EKFConfig, UKFConfig - (cd-dynamax) or KFConfig, EKFConfig, EnKFConfig, PFConfig (cuthbert). + Filter type inferred from config class: KFConfig, EKFConfig, UKFConfig, + RBPFConfig (cd-dynamax) or KFConfig, EKFConfig, EnKFConfig, PFConfig + (cuthbert). Args: name: Name of the factor. @@ -696,6 +705,7 @@ def _filter_discrete_time( name, dynamics, filter_config, + key=key, obs_times=obs_times, obs_values=obs_values, ctrl_times=ctrl_times, @@ -772,5 +782,6 @@ def _filter_continuous_time( "KFConfig", "PFConfig", "PFResamplingConfig", + "RBPFConfig", "UKFConfig", ] diff --git a/dynestyx/inference/integrations/cd_dynamax/discrete_filter.py b/dynestyx/inference/integrations/cd_dynamax/discrete_filter.py index 089c3321..f064a7fa 100644 --- a/dynestyx/inference/integrations/cd_dynamax/discrete_filter.py +++ b/dynestyx/inference/integrations/cd_dynamax/discrete_filter.py @@ -1,4 +1,6 @@ -"""Discrete-time filters via cd-dynamax (dynamax): KF, EKF, UKF.""" +"""Discrete-time filters via cd-dynamax: KF, EKF, UKF, and SLDS RBPF.""" + +from typing import NamedTuple import jax import jax.numpy as jnp @@ -14,11 +16,20 @@ UKFHyperParams, unscented_kalman_filter, ) +from cd_dynamax.dynamax.slds.inference import ( + DiscreteParamsSLDS, + LGParamsSLDS, + ParamsSLDS, + rbpfilter, + rbpfilter_optimal, +) +from jaxtyping import Array, Float, Int, PRNGKeyArray from dynestyx.inference.configs.filter import ( BaseFilterConfig, EKFConfig, KFConfig, + RBPFConfig, UKFConfig, ) from dynestyx.inference.integrations.cd_dynamax.utils import ( @@ -26,14 +37,251 @@ gaussian_to_nlgssm_params, ) from dynestyx.inference.integrations.utils import squeeze_leading_singletons -from dynestyx.inference.utils.distribution_utils import _posterior_sequence_to_dists +from dynestyx.inference.utils.distribution_utils import ( + _posterior_sequence_to_dists, + _rbpf_sequence_to_dists, +) from dynestyx.models import ( DynamicalModel, LinearGaussianObservation, LinearGaussianStateEvolution, + MixedStateDistribution, + SwitchingLinearGaussianObservation, + SwitchingLinearGaussianStateEvolution, ) +class RBPFPosterior(NamedTuple): + """Normalized dynestyx view of a cd-dynamax SLDS RBPF result. + + Attributes: + marginal_loglik: Particle estimate of the total marginal log + likelihood. + weights: Normalized particle weights with shape + ``(time, num_particles)``. + regimes: Discrete regime particles with shape + ``(time, num_particles)``. + means: Conditional continuous-state means with shape + ``(time, num_particles, state_dim)``. + covariances: Conditional continuous-state covariances with shape + ``(time, num_particles, state_dim, state_dim)``. + filtered_means: Mixture mean of the continuous state. + filtered_covariances: Mixture covariance of the continuous state, + including both within-particle and between-particle uncertainty. + filtered_regime_probs: Filtered categorical probabilities over + regimes. + """ + + marginal_loglik: Float[Array, ""] + weights: Float[Array, "time num_particles"] + regimes: Int[Array, "time num_particles"] + means: Float[Array, "time num_particles state_dim"] + covariances: Float[Array, "time num_particles state_dim state_dim"] + filtered_means: Float[Array, "time state_dim"] + filtered_covariances: Float[Array, "time state_dim state_dim"] + filtered_regime_probs: Float[Array, "time num_regimes"] + + @property + def particles(self) -> Float[Array, "time num_particles joint_state_dim"]: + """Joint particles encoded as ``[regime, *conditional_mean]``.""" + return jnp.concatenate( + (self.regimes[..., None].astype(self.means.dtype), self.means), + axis=-1, + ) + + @property + def log_weights(self) -> Float[Array, "time num_particles"]: + """Logarithms of normalized particle weights.""" + return jnp.log(self.weights) + + +def _slds_to_params(dynamics: DynamicalModel) -> ParamsSLDS: + """Translate a structured dynestyx SLDS to cd-dynamax parameters.""" + if not ( + isinstance(dynamics.state_evolution, SwitchingLinearGaussianStateEvolution) + and isinstance(dynamics.observation_model, SwitchingLinearGaussianObservation) + and isinstance(dynamics.initial_condition, MixedStateDistribution) + ): + raise TypeError( + "RBPFConfig requires a DynamicalModel with a " + "MixedStateDistribution initial condition, " + "SwitchingLinearGaussianStateEvolution, and " + "SwitchingLinearGaussianObservation." + ) + + evolution = dynamics.state_evolution + observation = dynamics.observation_model + initial = dynamics.initial_condition + num_regimes = evolution.num_regimes + state_dim = evolution.continuous_state_dim + observation_dim = dynamics.observation_dim + control_dim = dynamics.control_dim + + if observation.num_regimes != num_regimes: + raise ValueError( + "State evolution and observation model disagree on num_regimes: " + f"{num_regimes} != {observation.num_regimes}." + ) + if initial.num_regimes != num_regimes: + raise ValueError( + "Initial condition and state evolution disagree on num_regimes: " + f"{initial.num_regimes} != {num_regimes}." + ) + if initial.continuous_state_dim != state_dim: + raise ValueError( + "Initial condition and state evolution disagree on continuous " + f"state dimension: {initial.continuous_state_dim} != {state_dim}." + ) + + return ParamsSLDS( + discrete=DiscreteParamsSLDS( + initial_distribution=initial.categorical_probs, + transition_matrix=evolution.transition_matrix, + # Dynestyx intentionally exposes only the bootstrap proposal here. + # An arbitrary proposal requires an importance-ratio correction + # that is not part of the current cd-dynamax RBPF contract. + proposal_transition_matrix=evolution.transition_matrix, + ), + linear_gaussian=LGParamsSLDS( + initial_mean=initial.continuous_locs, + initial_cov=initial.continuous_covariances, + dynamics_weights=evolution.A, + dynamics_cov=evolution.cov, + dynamics_bias=( + jnp.zeros((num_regimes, state_dim)) + if evolution.bias is None + else evolution.bias + ), + dynamics_input_weights=( + jnp.zeros((num_regimes, state_dim, control_dim)) + if evolution.B is None + else evolution.B + ), + emission_weights=observation.H, + emission_cov=observation.R, + emission_bias=( + jnp.zeros((num_regimes, observation_dim)) + if observation.bias is None + else observation.bias + ), + emission_input_weights=( + jnp.zeros((num_regimes, observation_dim, control_dim)) + if observation.D is None + else observation.D + ), + initialized=True, + ), + ) + + +def _prepare_slds_inputs( + dynamics: DynamicalModel, + obs_values: jax.Array, + obs_times: jax.Array, + ctrl_times: jax.Array | None, + ctrl_values: jax.Array | None, +) -> tuple[jax.Array, jax.Array]: + """Prepare observation and control arrays for the cd-dynamax SLDS API.""" + observations = obs_values[:, None] if obs_values.ndim == 1 else obs_values + _, inputs = _prepare_inputs( + dynamics, + observations, + obs_times, + ctrl_times, + ctrl_values, + ) + return observations, inputs + + +def _normalize_rbpf_output( + output, + *, + num_regimes: int, +) -> RBPFPosterior: + """Validate and normalize the cd-dynamax RBPF output contract.""" + marginal_loglik = getattr(output, "marginal_loglik", None) + if marginal_loglik is None: + raise RuntimeError( + "The installed cd-dynamax version does not expose an SLDS RBPF " + "`marginal_loglik`. Install the commit pinned by dynestyx while " + "the upstream RBPF likelihood release is pending." + ) + weights = output.weights + regimes = output.states + means = output.means + covariances = output.covariances + if any(value is None for value in (weights, regimes, means, covariances)): + raise RuntimeError("cd-dynamax returned an incomplete SLDS RBPF posterior.") + + filtered_means = jnp.einsum("...tp,...tpd->...td", weights, means) + centered = means - filtered_means[..., :, None, :] + filtered_covariances = jnp.einsum( + "...tp,...tpij->...tij", + weights, + covariances + centered[..., :, :, None] * centered[..., :, None, :], + ) + filtered_regime_probs = jnp.einsum( + "...tp,...tpk->...tk", + weights, + jax.nn.one_hot(regimes, num_regimes), + ) + return RBPFPosterior( + marginal_loglik=marginal_loglik, + weights=weights, + regimes=regimes, + means=means, + covariances=covariances, + filtered_means=filtered_means, + filtered_covariances=filtered_covariances, + filtered_regime_probs=filtered_regime_probs, + ) + + +def _compute_slds_rbpf( + dynamics: DynamicalModel, + filter_config: RBPFConfig, + key: PRNGKeyArray | None, + *, + obs_times: jax.Array, + obs_values: jax.Array, + ctrl_times: jax.Array | None, + ctrl_values: jax.Array | None, +) -> RBPFPosterior: + """Run cd-dynamax's SLDS RBPF and normalize its posterior.""" + if key is None: + raise ValueError( + "RBPFConfig requires a PRNG key. Set `crn_seed` or run the model " + "inside a NumPyro seed handler." + ) + params = _slds_to_params(dynamics) + observations, inputs = _prepare_slds_inputs( + dynamics, obs_values, obs_times, ctrl_times, ctrl_values + ) + if filter_config.proposal == "prior": + output = rbpfilter( + filter_config.n_particles, + params, + observations, + key, + inputs=inputs, + ess_threshold=filter_config.ess_threshold_ratio, + ) + elif filter_config.proposal == "optimal": + output = rbpfilter_optimal( + filter_config.n_particles, + params, + observations, + key, + inputs=inputs, + ) + else: + raise ValueError(f"Unknown RBPF proposal: {filter_config.proposal!r}.") + return _normalize_rbpf_output( + output, + num_regimes=int(params.discrete.transition_matrix.shape[-1]), + ) + + def _lti_to_lgssm_params(dynamics: DynamicalModel): """Build dynamax ParamsLGSSM from LinearGaussianSSM.initialize for an LTI model.""" state_dim = dynamics.state_dim @@ -101,6 +349,7 @@ def _prepare_inputs(dynamics, obs_values, obs_times, ctrl_times, ctrl_values): def compute_cd_dynamax_discrete_filter( dynamics: DynamicalModel, filter_config: BaseFilterConfig, + key: PRNGKeyArray | None = None, *, obs_times: jax.Array, obs_values: jax.Array, @@ -108,6 +357,17 @@ def compute_cd_dynamax_discrete_filter( ctrl_values=None, ): """Pure-JAX cd-dynamax discrete filter computation (no numpyro side-effects).""" + if isinstance(filter_config, RBPFConfig): + return _compute_slds_rbpf( + dynamics, + filter_config, + key, + obs_times=obs_times, + obs_values=obs_values, + ctrl_times=ctrl_times, + ctrl_values=ctrl_values, + ) + emissions, inputs = _prepare_inputs( dynamics, obs_values, obs_times, ctrl_times, ctrl_values ) @@ -132,7 +392,7 @@ def compute_cd_dynamax_discrete_filter( ) raise ValueError( f"Unsupported cd-dynamax discrete config: {type(filter_config).__name__}. " - "Expected KFConfig, EKFConfig, or UKFConfig." + "Expected KFConfig, EKFConfig, UKFConfig, or RBPFConfig." ) @@ -140,6 +400,7 @@ def run_discrete_filter( name: str, dynamics: DynamicalModel, filter_config: BaseFilterConfig, + key: PRNGKeyArray | None = None, *, obs_times: jax.Array, obs_values: jax.Array, @@ -147,7 +408,7 @@ def run_discrete_filter( ctrl_values=None, **kwargs, ) -> tuple[jax.Array, object, list[dist.Distribution]]: - """Run discrete-time filter via cd-dynamax (KF, EKF, UKF). + """Run a discrete-time cd-dynamax filter (KF, EKF, UKF, or SLDS RBPF). Pure computation — no numpyro side-effects. Callers are responsible for registering numpyro.factor / numpyro.deterministic if needed. @@ -163,19 +424,23 @@ def run_discrete_filter( posterior = compute_cd_dynamax_discrete_filter( dynamics, filter_config, + key=key, obs_times=obs_times, obs_values=obs_values, ctrl_times=ctrl_times, ctrl_values=ctrl_values, ) - filtered_dists = _posterior_sequence_to_dists( - posterior, - means_attr="filtered_means", - covariances_attr="filtered_covariances", - particle_mode=False, - missing="empty", - ) + if isinstance(filter_config, RBPFConfig): + filtered_dists = _rbpf_sequence_to_dists(posterior) + else: + filtered_dists = _posterior_sequence_to_dists( + posterior, + means_attr="filtered_means", + covariances_attr="filtered_covariances", + particle_mode=False, + missing="empty", + ) return posterior.marginal_loglik, posterior, filtered_dists @@ -184,4 +449,5 @@ def run_discrete_filter( "run_discrete_filter", "_lti_to_lgssm_params", "_prepare_inputs", + "RBPFPosterior", ] diff --git a/dynestyx/inference/utils/distribution_utils.py b/dynestyx/inference/utils/distribution_utils.py index 0b1c3616..ec78b05a 100644 --- a/dynestyx/inference/utils/distribution_utils.py +++ b/dynestyx/inference/utils/distribution_utils.py @@ -8,6 +8,7 @@ from jaxtyping import Array, Float, PRNGKeyArray, Real, Shaped from numpyro.distributions import constraints +from dynestyx.distributions import RaoBlackwellizedParticleDistribution from dynestyx.inference.integrations.utils import ( WeightedParticles, covariance_from_cholesky, @@ -169,6 +170,26 @@ def _posterior_sequence_to_dists( ) +def _rbpf_sequence_to_dists( + posterior, + *, + plate_shapes: tuple[int, ...] = (), +) -> list[dist.Distribution]: + """Convert an RBPF Gaussian-mixture sequence to joint-state distributions.""" + t_len = _time_len_from_array(posterior.weights, plate_shapes) + return [ + RaoBlackwellizedParticleDistribution( + log_weights=jnp.log(_slice_time_axis(posterior.weights, t, plate_shapes)), + regimes=_slice_time_axis(posterior.regimes, t, plate_shapes), + continuous_locs=_slice_time_axis(posterior.means, t, plate_shapes), + continuous_covariances=_slice_time_axis( + posterior.covariances, t, plate_shapes + ), + ) + for t in range(t_len) + ] + + def _cholesky_state_sequence_to_dists( states, *, diff --git a/dynestyx/inference/utils/numpyro_sites.py b/dynestyx/inference/utils/numpyro_sites.py index 86b5c6d5..451d68db 100644 --- a/dynestyx/inference/utils/numpyro_sites.py +++ b/dynestyx/inference/utils/numpyro_sites.py @@ -10,6 +10,7 @@ ContinuousTimeConfigs, HMMConfig, PFConfig, + RBPFConfig, _config_to_record_kwargs, ) from dynestyx.inference.configs.smoother import ( @@ -54,6 +55,8 @@ def register_filter_sites( if isinstance(filter_config, tuple(ContinuousTimeConfigs)): _add_continuous_filter_sites(name, states, record_kwargs) + elif isinstance(filter_config, RBPFConfig): + _add_rbpf_sites(name, states, record_kwargs) elif isinstance(filter_config, PFConfig): _add_cuthbert_pf_sites(name, states, record_kwargs) else: @@ -241,6 +244,51 @@ def _add_cuthbert_pf_sites(name: str, states, record_kwargs: dict) -> None: numpyro.deterministic(f"{name}_filtered_states_cov_diag", diag_cov) +def _add_rbpf_sites(name: str, states, record_kwargs: dict) -> None: + """Register requested continuous-state and regime summaries from an RBPF.""" + max_elems = record_kwargs["record_max_elems"] + fields = ( + ( + "record_filtered_states_mean", + "filtered_states_mean", + states.filtered_means, + ), + ( + "record_filtered_states_cov", + "filtered_states_cov", + states.filtered_covariances, + ), + ( + "record_filtered_particles", + "filtered_particles", + states.particles, + ), + ( + "record_filtered_log_weights", + "filtered_log_weights", + states.log_weights, + ), + ( + "record_filtered_regime_probs", + "filtered_regime_probs", + states.filtered_regime_probs, + ), + ) + for config_field, site_suffix, value in fields: + if _should_record_field( + record_kwargs.get(config_field), value.shape, max_elems + ): + numpyro.deterministic(f"{name}_{site_suffix}", value) + + covariance_diag = jnp.diagonal(states.filtered_covariances, axis1=-2, axis2=-1) + if _should_record_field( + record_kwargs["record_filtered_states_cov_diag"], + covariance_diag.shape, + max_elems, + ): + numpyro.deterministic(f"{name}_filtered_states_cov_diag", covariance_diag) + + def _add_gaussian_filter_sites( name: str, states, filter_config: BaseFilterConfig, record_kwargs: dict ) -> None: diff --git a/dynestyx/models/__init__.py b/dynestyx/models/__init__.py index 28015a25..4df8e610 100644 --- a/dynestyx/models/__init__.py +++ b/dynestyx/models/__init__.py @@ -3,6 +3,7 @@ Structure anticipates future extension to LTI factories, Neural SDEs, etc. """ +from dynestyx.distributions import MixedStateDistribution from dynestyx.models.core import ( ContinuousTimeStateEvolution, DeterministicContinuousTimeStateEvolution, @@ -24,12 +25,14 @@ GaussianObservation, LinearGaussianObservation, LinearGaussianObservationParams, + SwitchingLinearGaussianObservation, ) from dynestyx.models.state_evolution import ( AffineDrift, GaussianStateEvolution, LinearGaussianParams, LinearGaussianStateEvolution, + SwitchingLinearGaussianStateEvolution, ) __all__ = [ @@ -49,7 +52,10 @@ "LinearGaussianObservationParams", "LinearGaussianParams", "LinearGaussianStateEvolution", + "MixedStateDistribution", "ObservationModel", + "SwitchingLinearGaussianObservation", + "SwitchingLinearGaussianStateEvolution", "StochasticContinuousTimeStateEvolution", "LTI_continuous", "LTI_discrete", diff --git a/dynestyx/models/observations.py b/dynestyx/models/observations.py index 352cd68c..6efa9ad2 100644 --- a/dynestyx/models/observations.py +++ b/dynestyx/models/observations.py @@ -173,6 +173,81 @@ def __call__(self, x, u, t): return dist.MultivariateNormal(loc=loc, covariance_matrix=R) +class SwitchingLinearGaussianObservation(ObservationModel): + r"""Regime-switching linear-Gaussian observation model for an SLDS. + + Given the joint state ``[z_t, *x_t]``, observations follow + + \[ + y_t \mid x_t, z_t + \sim \mathcal{N}(H_{z_t}x_t + D_{z_t}u_t + b_{z_t}, R_{z_t}). + \] + + Args: + H: Regime-specific observation matrices with shape + ``(num_regimes, observation_dim, state_dim)``. + R: Regime-specific observation covariance matrices with shape + ``(num_regimes, observation_dim, observation_dim)``. + D: Optional regime-specific control matrices with shape + ``(num_regimes, observation_dim, control_dim)``. + bias: Optional regime-specific observation biases with shape + ``(num_regimes, observation_dim)``. + + Attributes: + H: Regime-specific observation matrices. + R: Regime-specific observation covariance matrices. + D: Optional regime-specific control matrices. + bias: Optional regime-specific observation biases. + """ + + H: Float[Array, "num_regimes observation_dim state_dim"] + R: Float[Array, "num_regimes observation_dim observation_dim"] + D: Float[Array, "num_regimes observation_dim control_dim"] | None = None + bias: Float[Array, "num_regimes observation_dim"] | None = None + + def __init__( + self, + H: Float[Array, "num_regimes observation_dim state_dim"], + R: Float[Array, "num_regimes observation_dim observation_dim"], + D: Float[Array, "num_regimes observation_dim control_dim"] | None = None, + bias: Float[Array, "num_regimes observation_dim"] | None = None, + ) -> None: + self.H = H + self.R = R + self.D = D + self.bias = bias + + @property + def num_regimes(self) -> int: + """Number of discrete regimes.""" + return int(self.H.shape[-3]) + + @property + def continuous_state_dim(self) -> int: + """Dimension of the continuous part of the state.""" + return int(self.H.shape[-1]) + + def __call__( + self, + x: Real[Array, " joint_state_dim"], + u: Real[Array, " control_dim"] | Real[Array, ""] | None, + t: float | int | Real[Array, ""], + ) -> dist.MultivariateNormal: + """Return the observation distribution for ``[z_t, *x_t]``.""" + del t + regime = jnp.rint(x[0]).astype(jnp.int32) + continuous_state = x[1:] + loc = jnp.dot(self.H[regime], continuous_state) + if self.D is not None and u is not None: + loc = loc + jnp.dot(self.D[regime], u) + if self.bias is not None: + loc = loc + self.bias[regime] + return dist.MultivariateNormal( + loc=loc, + covariance_matrix=self.R[regime], + ) + + class GaussianObservation(ObservationModel): """ Nonlinear Gaussian observation model. diff --git a/dynestyx/models/state_evolution.py b/dynestyx/models/state_evolution.py index 2e91fd81..7b145d27 100644 --- a/dynestyx/models/state_evolution.py +++ b/dynestyx/models/state_evolution.py @@ -12,6 +12,7 @@ import numpyro.distributions as dist from jaxtyping import Array, Float, Real +from dynestyx.distributions import MixedStateDistribution from dynestyx.models.core import DiscreteTimeStateEvolution @@ -187,6 +188,96 @@ def __call__(self, x, u, t_now, t_next): return dist.MultivariateNormal(loc=loc, covariance_matrix=cov) +class SwitchingLinearGaussianStateEvolution(DiscreteTimeStateEvolution): + r"""Regime-switching linear-Gaussian transition for an SLDS. + + Given a joint state ``[z_t, *x_t]``, the transition is + + \[ + z_{t+1} \mid z_t + \sim \operatorname{Categorical}(P_{z_t}), \qquad + x_{t+1} \mid x_t, z_{t+1} + \sim \mathcal{N}(A_{z_{t+1}}x_t + + B_{z_{t+1}}u_t + b_{z_{t+1}}, Q_{z_{t+1}}). + \] + + The next regime selects the continuous transition parameters. This + convention matches the Rao-Blackwellized particle-filter backend and keeps + simulation and filtering under one model definition. + + Args: + transition_matrix: Row-stochastic regime transition matrix with shape + ``(num_regimes, num_regimes)``. + A: Regime-specific state matrices with shape + ``(num_regimes, state_dim, state_dim)``. + cov: Regime-specific process covariance matrices with shape + ``(num_regimes, state_dim, state_dim)``. + B: Optional regime-specific control matrices with shape + ``(num_regimes, state_dim, control_dim)``. + bias: Optional regime-specific transition biases with shape + ``(num_regimes, state_dim)``. + + Attributes: + transition_matrix: Regime transition probabilities. + A: Regime-specific state matrices. + cov: Regime-specific process covariance matrices. + B: Optional regime-specific control matrices. + bias: Optional regime-specific transition biases. + """ + + transition_matrix: Float[Array, "num_regimes num_regimes"] + A: Float[Array, "num_regimes state_dim state_dim"] + cov: Float[Array, "num_regimes state_dim state_dim"] + B: Float[Array, "num_regimes state_dim control_dim"] | None = None + bias: Float[Array, "num_regimes state_dim"] | None = None + + def __init__( + self, + transition_matrix: Float[Array, "num_regimes num_regimes"], + A: Float[Array, "num_regimes state_dim state_dim"], + cov: Float[Array, "num_regimes state_dim state_dim"], + B: Float[Array, "num_regimes state_dim control_dim"] | None = None, + bias: Float[Array, "num_regimes state_dim"] | None = None, + ) -> None: + self.transition_matrix = transition_matrix + self.A = A + self.cov = cov + self.B = B + self.bias = bias + + @property + def num_regimes(self) -> int: + """Number of discrete regimes.""" + return int(self.transition_matrix.shape[-1]) + + @property + def continuous_state_dim(self) -> int: + """Dimension of the continuous part of the state.""" + return int(self.A.shape[-1]) + + def __call__( + self, + x: Real[Array, " joint_state_dim"], + u: Real[Array, " control_dim"] | Real[Array, ""] | None, + t_now: float | int | Real[Array, ""], + t_next: float | int | Real[Array, ""], + ) -> MixedStateDistribution: + """Return the joint transition distribution from ``[z_t, *x_t]``.""" + del t_now, t_next + regime = jnp.rint(x[0]).astype(jnp.int32) + continuous_state = x[1:] + locs = jnp.einsum("...kij,...j->...ki", self.A, continuous_state) + if self.B is not None and u is not None: + locs = locs + jnp.einsum("...kij,...j->...ki", self.B, u) + if self.bias is not None: + locs = locs + self.bias + return MixedStateDistribution( + categorical_probs=self.transition_matrix[regime], + continuous_locs=locs, + continuous_covariances=self.cov, + ) + + class GaussianStateEvolution(DiscreteTimeStateEvolution): """ Nonlinear Gaussian discrete-time state transition. diff --git a/dynestyx/utils.py b/dynestyx/utils.py index 97a8ee4b..7ee79635 100644 --- a/dynestyx/utils.py +++ b/dynestyx/utils.py @@ -13,7 +13,12 @@ from jax import Array, lax from jaxtyping import Real, Shaped -from dynestyx.models import Diffusion, DynamicalModel +from dynestyx.models import ( + Diffusion, + DynamicalModel, + SwitchingLinearGaussianObservation, + SwitchingLinearGaussianStateEvolution, +) def flatten_draws(arr: Shaped[Array, "..."]) -> Shaped[Array, "..."]: @@ -212,7 +217,14 @@ def _is_opaque_plate_leaf(node) -> bool: """ if isinstance(node, Diffusion): return not callable(node.coefficient) - return isinstance(node, numpyro.distributions.Distribution) + return isinstance( + node, + ( + numpyro.distributions.Distribution, + SwitchingLinearGaussianObservation, + SwitchingLinearGaussianStateEvolution, + ), + ) def _dist_has_plate_batch_dims(dist_obj, plate_shapes: tuple[int, ...]) -> bool: diff --git a/mkdocs.yml b/mkdocs.yml index 5ec9b191..4d2182a3 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -45,6 +45,8 @@ nav: - Tracking an object with a Kalman filter: tutorials/state_space_models/kf_tracking.ipynb - Online linear regression with a Kalman filter: tutorials/state_space_models/kf_linreg.ipynb - Parameter estimation for an LGSSM (SGD and MCMC): tutorials/state_space_models/lgssm_learning.ipynb + - Switching linear dynamical systems: + - Rao-Blackwellized particle filtering: tutorials/state_space_models/slds_rbpf.md - Nonlinear Gaussian SSMs: - Tracking a spiraling object with an EKF: tutorials/state_space_models/ekf_spiral.ipynb - Online learning of an MLP with an EKF: tutorials/state_space_models/ekf_mlp.ipynb @@ -77,6 +79,9 @@ nav: - LinearGaussianObservation: api_reference/public/models/specialized/linear_gaussian_observation.md - GaussianObservation: api_reference/public/models/specialized/gaussian_observation.md - LinearGaussianStateEvolution: api_reference/public/models/specialized/linear_gaussian_state_evolution.md + - MixedStateDistribution: api_reference/public/models/specialized/mixed_state_distribution.md + - SwitchingLinearGaussianStateEvolution: api_reference/public/models/specialized/switching_linear_gaussian_state_evolution.md + - SwitchingLinearGaussianObservation: api_reference/public/models/specialized/switching_linear_gaussian_observation.md - GaussianStateEvolution: api_reference/public/models/specialized/gaussian_state_evolution.md - AffineDrift: api_reference/public/models/specialized/affine_drift.md - LTI_continuous: api_reference/public/models/specialized/lti_continuous.md diff --git a/pyproject.toml b/pyproject.toml index ed4d7242..54d38235 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -33,7 +33,9 @@ dependencies = [ "effectful>=0.4.0", "cuthbert>=0.0.10", "cuthbertlib>=0.0.10", - "cd-dynamax>=0.3.3", + # Temporary commit pin while the SLDS RBPF marginal-likelihood change is + # awaiting an upstream cd-dynamax release. + "cd-dynamax @ git+https://github.com/anushri10/cd_dynamax.git@6295031bf19162735564c3b3c7cad697ea9176d2", "matplotlib>=3.10.7", "numpyro>=0.19.0", "pytest>=9.0.1", diff --git a/tests/test_slds_rbpf.py b/tests/test_slds_rbpf.py new file mode 100644 index 00000000..bfa9779b --- /dev/null +++ b/tests/test_slds_rbpf.py @@ -0,0 +1,371 @@ +import jax.numpy as jnp +import jax.random as jr +import numpyro +import numpyro.distributions as dist +import pytest +from numpyro.handlers import seed, trace + +import dynestyx as dsx +from dynestyx.distributions import RaoBlackwellizedParticleDistribution +from dynestyx.inference.configs.mcmc import NUTSConfig +from dynestyx.inference.filters import KFConfig, RBPFConfig +from dynestyx.inference.integrations.cd_dynamax.discrete_filter import ( + RBPFPosterior, + compute_cd_dynamax_discrete_filter, +) +from dynestyx.inference.mcmc import MCMCInference + + +def _make_slds_dynamics(*, bias_shift=0.0): + num_regimes = 2 + state_dim = 2 + observation_dim = 1 + transition_matrix = jnp.array([[0.95, 0.05], [0.10, 0.90]]) + dynamics_matrices = jnp.array( + [ + [[0.95, 0.10], [-0.05, 0.90]], + [[0.55, -0.25], [0.20, 0.65]], + ] + ) + dynamics_covariances = jnp.tile(0.05**2 * jnp.eye(state_dim), (num_regimes, 1, 1)) + dynamics_biases = jnp.array([[0.0, 0.0], [0.75 + bias_shift, -0.25]]) + observation_matrices = jnp.tile(jnp.array([[1.0, 0.0]]), (num_regimes, 1, 1)) + observation_covariances = jnp.tile( + 0.10**2 * jnp.eye(observation_dim), (num_regimes, 1, 1) + ) + + return dsx.DynamicalModel( + initial_condition=dsx.MixedStateDistribution( + categorical_probs=jnp.array([0.80, 0.20]), + continuous_locs=jnp.zeros((num_regimes, state_dim)), + continuous_covariances=jnp.tile(jnp.eye(state_dim), (num_regimes, 1, 1)), + ), + state_evolution=dsx.SwitchingLinearGaussianStateEvolution( + transition_matrix=transition_matrix, + A=dynamics_matrices, + cov=dynamics_covariances, + bias=dynamics_biases, + ), + observation_model=dsx.SwitchingLinearGaussianObservation( + H=observation_matrices, + R=observation_covariances, + ), + ) + + +def _make_observations(): + times = jnp.arange(8.0) + simulation = dsx.simulate( + _make_slds_dynamics(), + rng_key=jr.PRNGKey(0), + predict_times=times, + ) + assert simulation.observations is not None + return times, simulation.observations[0] + + +def _slds_model( + obs_times=None, + obs_values=None, + *, + sample_bias_shift=False, +): + bias_shift = ( + numpyro.sample("bias_shift", dist.Normal(0.0, 0.1)) + if sample_bias_shift + else 0.0 + ) + return dsx.sample( + "f", + _make_slds_dynamics(bias_shift=bias_shift), + obs_times=obs_times, + obs_values=obs_values, + ) + + +def test_mixed_state_distribution_samples_and_scores_batched_states(): + probs = jnp.array([[0.25, 0.75], [0.60, 0.40]]) + locs = jnp.array( + [ + [[0.0, 1.0], [2.0, -1.0]], + [[-1.0, 0.5], [0.25, 1.5]], + ] + ) + covariances = jnp.tile(jnp.eye(2)[None, None], (2, 2, 1, 1)) + mixed = dsx.MixedStateDistribution(probs, locs, covariances) + + samples = mixed.sample(jr.PRNGKey(0), sample_shape=(5,)) + value = jnp.array([[1.0, 1.7, -0.4], [0.0, -0.8, 0.1]]) + expected = dist.Categorical(probs=probs).log_prob( + jnp.array([1, 0]) + ) + dist.MultivariateNormal( + jnp.array([locs[0, 1], locs[1, 0]]), + covariance_matrix=jnp.eye(2), + ).log_prob(value[:, 1:]) + + assert mixed.batch_shape == (2,) + assert mixed.event_shape == (3,) + assert samples.shape == (5, 2, 3) + assert jnp.allclose(mixed.log_prob(value), expected) + + +def test_rbpf_distribution_retains_conditional_gaussian_uncertainty(): + posterior = RaoBlackwellizedParticleDistribution( + log_weights=jnp.array([0.0]), + regimes=jnp.array([1]), + continuous_locs=jnp.zeros((1, 1)), + continuous_covariances=jnp.ones((1, 1, 1)), + ) + + samples = posterior.sample(jr.PRNGKey(0), sample_shape=(256,)) + + assert jnp.all(samples[:, 0] == 1) + assert jnp.var(samples[:, 1]) > 0.2 + assert jnp.isfinite(posterior.log_prob(jnp.array([1.0, 0.0]))) + + +@pytest.mark.parametrize("proposal", ["prior", "optimal"]) +def test_slds_rbpf_returns_finite_typed_posterior(proposal): + obs_times, obs_values = _make_observations() + config = RBPFConfig(n_particles=64, proposal=proposal) + posterior = compute_cd_dynamax_discrete_filter( + _make_slds_dynamics(), + config, + key=jr.PRNGKey(1), + obs_times=obs_times, + obs_values=obs_values, + ) + + assert isinstance(posterior, RBPFPosterior) + assert jnp.isfinite(posterior.marginal_loglik) + assert posterior.weights.shape == (len(obs_times), config.n_particles) + assert posterior.means.shape == (len(obs_times), config.n_particles, 2) + assert posterior.filtered_means.shape == (len(obs_times), 2) + assert posterior.filtered_covariances.shape == (len(obs_times), 2, 2) + assert posterior.filtered_regime_probs.shape == (len(obs_times), 2) + assert jnp.allclose(posterior.weights.sum(axis=-1), 1.0) + assert jnp.allclose(posterior.filtered_regime_probs.sum(axis=-1), 1.0) + + +def test_slds_rbpf_registers_requested_trace_sites(): + obs_times, obs_values = _make_observations() + config = RBPFConfig( + n_particles=32, + proposal="optimal", + record_filtered_states_mean=True, + record_filtered_states_cov=True, + record_filtered_regime_probs=True, + crn_seed=jr.PRNGKey(1), + ) + + with trace() as tr, seed(rng_seed=jr.PRNGKey(2)), dsx.Filter(config): + _slds_model(obs_times, obs_values) + + assert jnp.isfinite(tr["f_marginal_loglik"]["value"]) + assert tr["f_filtered_states_mean"]["value"].shape == (len(obs_times), 2) + assert tr["f_filtered_states_cov"]["value"].shape == (len(obs_times), 2, 2) + assert tr["f_filtered_regime_probs"]["value"].shape == (len(obs_times), 2) + + +def test_slds_rbpf_supports_shared_model_in_plate(): + obs_times, obs_values = _make_observations() + batched_observations = jnp.stack((obs_values, obs_values)) + + def plated_model(): + with dsx.plate("trajectories", 2): + _slds_model(obs_times, batched_observations) + + config = RBPFConfig( + n_particles=32, + proposal="optimal", + crn_seed=jr.PRNGKey(1), + ) + with trace() as tr, seed(rng_seed=jr.PRNGKey(2)), dsx.Filter(config): + plated_model() + + marginal_loglik = tr["f_marginal_loglik"]["value"] + assert marginal_loglik.shape == (2,) + assert jnp.isfinite(marginal_loglik).all() + + +def test_slds_rbpf_supports_controls(): + obs_times, obs_values = _make_observations() + base = _make_slds_dynamics() + evolution = base.state_evolution + observation = base.observation_model + controlled = dsx.DynamicalModel( + control_dim=1, + initial_condition=base.initial_condition, + state_evolution=dsx.SwitchingLinearGaussianStateEvolution( + transition_matrix=evolution.transition_matrix, + A=evolution.A, + cov=evolution.cov, + B=jnp.ones((evolution.num_regimes, 2, 1)) * 0.05, + bias=evolution.bias, + ), + observation_model=dsx.SwitchingLinearGaussianObservation( + H=observation.H, + R=observation.R, + D=jnp.ones((observation.num_regimes, 1, 1)) * 0.01, + bias=observation.bias, + ), + ) + + posterior = compute_cd_dynamax_discrete_filter( + controlled, + RBPFConfig(n_particles=32), + key=jr.PRNGKey(1), + obs_times=obs_times, + obs_values=obs_values, + ctrl_times=obs_times, + ctrl_values=jnp.ones((len(obs_times), 1)), + ) + + assert jnp.isfinite(posterior.marginal_loglik) + + +def _make_degenerate_slds_and_lgssm(*, compensate_backend_initialization): + num_regimes = 2 + state_dim = 2 + initial_mean = jnp.array([0.2, -0.1]) + initial_cov = 0.5 * jnp.eye(state_dim) + A = jnp.array([[0.75, 0.1], [-0.05, 0.8]]) + Q = 0.05 * jnp.eye(state_dim) + H = jnp.array([[1.0, 0.3], [-0.2, 0.7]]) + R = 0.1 * jnp.eye(2) + transition_matrix = jnp.full((num_regimes, num_regimes), 1.0 / num_regimes) + + slds = dsx.DynamicalModel( + initial_condition=dsx.MixedStateDistribution( + jnp.full((num_regimes,), 1.0 / num_regimes), + jnp.tile(initial_mean[None], (num_regimes, 1)), + jnp.tile(initial_cov[None], (num_regimes, 1, 1)), + ), + state_evolution=dsx.SwitchingLinearGaussianStateEvolution( + transition_matrix, + jnp.tile(A[None], (num_regimes, 1, 1)), + jnp.tile(Q[None], (num_regimes, 1, 1)), + ), + observation_model=dsx.SwitchingLinearGaussianObservation( + jnp.tile(H[None], (num_regimes, 1, 1)), + jnp.tile(R[None], (num_regimes, 1, 1)), + ), + ) + + if compensate_backend_initialization: + # The pinned backend samples an initial Gaussian mean and then performs + # one transition before its first observation. These are the exact + # LGSSM moments implied by that temporary backend contract. + initial_mean = A @ initial_mean + initial_cov = 2.0 * A @ initial_cov @ A.T + Q + lgssm = dsx.LTI_discrete( + A=A, + Q=Q, + H=H, + R=R, + initial_mean=initial_mean, + initial_cov=initial_cov, + ) + return slds, lgssm + + +@pytest.mark.parametrize("proposal", ["prior", "optimal"]) +def test_degenerate_slds_marginal_loglik_matches_backend_equivalent_kf(proposal): + """Identical SLDS regimes match the KF under the pinned backend contract.""" + times = jnp.arange(6.0) + observations = jnp.array( + [[0.2, -0.1], [0.4, 0.0], [0.1, 0.3], [-0.2, 0.1], [0.0, -0.2], [0.3, 0.2]] + ) + slds, lgssm = _make_degenerate_slds_and_lgssm( + compensate_backend_initialization=True + ) + + rbpf_posterior = compute_cd_dynamax_discrete_filter( + slds, + RBPFConfig(n_particles=2_000, proposal=proposal), + key=jr.PRNGKey(3), + obs_times=times, + obs_values=observations, + ) + kf_posterior = compute_cd_dynamax_discrete_filter( + lgssm, + KFConfig(), + obs_times=times, + obs_values=observations, + ) + + assert jnp.allclose( + rbpf_posterior.marginal_loglik, + kf_posterior.marginal_loglik, + atol=0.1, + ) + + +@pytest.mark.xfail( + strict=True, + reason=( + "The pinned CD-Dynamax RBPF advances the continuous state before y[0]; " + "remove this marker when the upstream release uses the model's x[0] prior." + ), +) +def test_degenerate_slds_marginal_loglik_matches_same_model_kf(): + """Identical regimes should reduce to the same-model Kalman filter.""" + times = jnp.arange(6.0) + observations = jnp.array( + [[0.2, -0.1], [0.4, 0.0], [0.1, 0.3], [-0.2, 0.1], [0.0, -0.2], [0.3, 0.2]] + ) + slds, lgssm = _make_degenerate_slds_and_lgssm( + compensate_backend_initialization=False + ) + + rbpf_posterior = compute_cd_dynamax_discrete_filter( + slds, + RBPFConfig(n_particles=4_000, proposal="optimal"), + key=jr.PRNGKey(3), + obs_times=times, + obs_values=observations, + ) + kf_posterior = compute_cd_dynamax_discrete_filter( + lgssm, + KFConfig(), + obs_times=times, + obs_values=observations, + ) + + assert jnp.allclose( + rbpf_posterior.marginal_loglik, + kf_posterior.marginal_loglik, + atol=0.1, + ) + + +def test_slds_rbpf_numpyro_nuts_one_iteration_smoke(): + obs_times, obs_values = _make_observations() + + with dsx.Filter( + RBPFConfig( + n_particles=16, + proposal="optimal", + crn_seed=jr.PRNGKey(1), + ) + ): + inference = MCMCInference( + mcmc_config=NUTSConfig( + num_samples=1, + num_warmup=1, + num_chains=1, + mcmc_source="numpyro", + ), + model=lambda obs_times, obs_values, **_: _slds_model( + obs_times, + obs_values, + sample_bias_shift=True, + ), + ) + posterior_samples = inference.run( + jr.PRNGKey(3), + obs_times, + obs_values, + ) + + assert jnp.isfinite(posterior_samples["bias_shift"]).all()