Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 18 additions & 0 deletions dynestyx/inference/checkers.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,24 @@ def _leading_dims(
return tuple(int(d) for d in arr.shape[:n])


def _ensure_trailing_event_axis(
values: Real[Array, "..."],
*,
plate_shapes: tuple[int, ...],
) -> Real[Array, "..."]:
"""Lift scalar time series to ``(*plate, time, 1)`` for numerical inference."""
n_plate_dims = len(plate_shapes)
has_plate_axes = (
n_plate_dims > 0
and values.ndim > n_plate_dims
and tuple(values.shape[:n_plate_dims]) == plate_shapes
)
scalar_series_ndim = n_plate_dims + 1 if has_plate_axes else 1
if values.ndim == scalar_series_ndim:
return values[..., None]
return values


def _summarize_dynamics_leading_dims(
dynamics: DynamicalModel, n_dims: int, max_items: int = 6
) -> str:
Expand Down
10 changes: 5 additions & 5 deletions dynestyx/inference/configs/filter.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,8 @@
import math
from typing import Literal

import jax
import jax.random as jr
from jaxtyping import PRNGKeyArray

ResamplingBaseMethod = Literal["systematic", "multinomial", "stratified"]
ResamplingDifferentiableMethod = Literal["stop_gradient", "straight_through", "soft"]
Expand Down Expand Up @@ -55,7 +55,7 @@ class BaseFilterConfig(abc.ABC):
this factor before the update. Values slightly above `1.0`
implement covariance inflation, which can improve robustness when
the model is misspecified. `None` disables rescaling.
crn_seed (jax.Array | None): Fix the PRNG key for stochastic filters
crn_seed (PRNGKeyArray | None): Fix the PRNG key for stochastic filters
(EnKF, PF). Useful when differentiating through the filter:
a fixed key makes the randomness a deterministic function of model
parameters. `None` draws a fresh key each call.
Expand All @@ -78,7 +78,7 @@ class BaseFilterConfig(abc.ABC):
record_max_elems: int = 100_000
filter_source: FilterSource | None = None
cov_rescaling: float | None = None
crn_seed: jax.Array | None = None
crn_seed: PRNGKeyArray | None = None


@dataclasses.dataclass
Expand Down Expand Up @@ -107,7 +107,7 @@ class EnKFConfig(BaseFilterConfig):
n_particles (int): Number of ensemble members. More members give a
better covariance estimate at higher compute cost. Defaults to
`30`.
crn_seed (jax.Array | None): Fixed PRNG key for the ensemble. Defaults
crn_seed (PRNGKeyArray | None): Fixed PRNG key for the ensemble. Defaults
to `jr.PRNGKey(0)`, i.e., common random numbers are used. This
can reduce variance in gradient-based learning, but introduces
further bias.
Expand Down Expand Up @@ -163,7 +163,7 @@ class EnKFConfig(BaseFilterConfig):
"""

n_particles: int = 30
crn_seed: jax.Array | None = dataclasses.field(
crn_seed: PRNGKeyArray | None = dataclasses.field(
default_factory=lambda: jr.PRNGKey(0)
)
perturb_measurements: bool | None = None
Expand Down
43 changes: 29 additions & 14 deletions dynestyx/inference/filters.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

from dynestyx.handlers import HandlesSelf, _condition_intp
from dynestyx.inference.checkers import (
_ensure_trailing_event_axis,
_validate_batched_plate_alignment,
_validate_inference_supported_model_classes,
_validate_missing_observation_support,
Expand Down Expand Up @@ -88,7 +89,7 @@ def _sample_ds(
name: str,
dynamics: DynamicalModel,
*,
plate_shapes=(),
plate_shapes: tuple[int, ...] = (),
obs_times: Real[Array, "*obs_time_plate obs_time"] | None = None,
obs_values: Real[Array, "*obs_value_plate obs_time observation_dim"]
| Real[Array, "*obs_value_plate obs_time"]
Expand Down Expand Up @@ -143,7 +144,7 @@ def _add_log_factors(
name: str,
dynamics: DynamicalModel,
*,
plate_shapes=(),
plate_shapes: tuple[int, ...] = (),
obs_times: Real[Array, "*obs_time_plate obs_time"] | None = None,
obs_values: Real[Array, "*obs_value_plate obs_time observation_dim"]
| Real[Array, "*obs_value_plate obs_time"]
Expand All @@ -161,7 +162,7 @@ def _build_infer_result(
) -> ConditionedResult: ...


def _default_filter_config(dynamics: DynamicalModel):
def _default_filter_config(dynamics: DynamicalModel) -> BaseFilterConfig:
"""Return appropriate default filter config when none specified."""
if dynamics.continuous_time:
return ContinuousTimeEnKFConfig()
Expand Down Expand Up @@ -237,7 +238,7 @@ def _add_log_factors(
name: str,
dynamics: DynamicalModel,
*,
plate_shapes=(),
plate_shapes: tuple[int, ...] = (),
obs_times: Real[Array, "*obs_time_plate obs_time"] | None = None,
obs_values: Real[Array, "*obs_value_plate obs_time observation_dim"]
| Real[Array, "*obs_value_plate obs_time"]
Expand Down Expand Up @@ -275,6 +276,16 @@ def _add_log_factors(
obs_values=obs_values,
mode="filter",
)
if not isinstance(config, HMMConfigs):
obs_values = _ensure_trailing_event_axis(
obs_values,
plate_shapes=plate_shapes,
)
if ctrl_values is not None:
ctrl_values = _ensure_trailing_event_axis(
ctrl_values,
plate_shapes=plate_shapes,
)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Codex points out (and I agree) that this can be problematic to do before plating; for plate size M and time series length T, if M = T, then shared scalar controls (T,) get promoted to plated controls (T, 1), which gets interpreted as plated. We should either do this later, or raise an error in _ensure_trailing_event_axis for this (rather rare) edge case.


# Resolve PRNG key: use explicit seed from config, fall back to numpyro
# context (inside a seeded model), or None (deterministic filters don't need one).
Expand Down Expand Up @@ -669,13 +680,15 @@ def _filter_discrete_time(
key: PRNGKeyArray | None = None,
*,
obs_times: Real[Array, " obs_time"],
obs_values: Real[Array, "obs_time observation_dim"] | Real[Array, " obs_time"],
obs_values: Real[Array, "obs_time observation_dim"],
ctrl_times: Real[Array, " ctrl_time"] | None = None,
ctrl_values: Real[Array, "ctrl_time control_dim"]
| Real[Array, " ctrl_time"]
| None = None,
ctrl_values: Real[Array, "ctrl_time control_dim"] | None = None,
**kwargs,
) -> tuple[jax.Array | None, object | None, list[numpyro.distributions.Distribution]]:
) -> tuple[
Real[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
Expand Down Expand Up @@ -725,13 +738,15 @@ def _filter_continuous_time(
key: PRNGKeyArray | None = None,
*,
obs_times: Real[Array, " obs_time"],
obs_values: Real[Array, "obs_time observation_dim"] | Real[Array, " obs_time"],
obs_values: Real[Array, "obs_time observation_dim"],
ctrl_times: Real[Array, " ctrl_time"] | None = None,
ctrl_values: Real[Array, "ctrl_time control_dim"]
| Real[Array, " ctrl_time"]
| None = None,
ctrl_values: Real[Array, "ctrl_time control_dim"] | None = None,
**kwargs,
) -> tuple[jax.Array, object, list[numpyro.distributions.Distribution]]:
) -> tuple[
Real[Array, ""],
object,
list[numpyro.distributions.Distribution],
]:
"""Continuous-time marginal likelihood via CD-Dynamax.

Supports: EnKF, DPF, EKF, UKF (inferred from config type).
Expand Down
45 changes: 24 additions & 21 deletions dynestyx/inference/integrations/cd_dynamax/continuous_filter.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
"""Continuous-time filters via CD-Dynamax: KF, EnKF, DPF, EKF, UKF."""

from typing import Any

import jax
import jax.numpy as jnp
import numpyro.distributions as dist
Expand All @@ -11,6 +13,7 @@
from cd_dynamax.src.continuous_discrete_linear_gaussian_ssm.models import (
PosteriorGSSMFiltered,
)
from jaxtyping import Array, PRNGKeyArray, Real

from dynestyx.inference.configs.filter import (
ContinuousTimeDPFConfig,
Expand Down Expand Up @@ -39,12 +42,12 @@

def _config_to_cd_dynamax_filter_kwargs(
config: ContinuousTimeFilterConfig,
params,
obs_values,
obs_times,
ctrl_values,
key,
) -> dict:
params: Any,
obs_values: Real[Array, "obs_time observation_dim"],
obs_times: Real[Array, "obs_time one"],
ctrl_values: Real[Array, "ctrl_time control_dim"],
key: PRNGKeyArray | None,
) -> dict[str, Any]:
"""Build the filter_kwargs dict passed to cd_dynamax_model.filter()."""

# cd-dynamax uses the legacy PRNG key interface, but newer numpyro uses typed keys.
Expand Down Expand Up @@ -111,9 +114,9 @@ def _config_to_cd_dynamax_filter_kwargs(

def _run_linear_kf(
dynamics: DynamicalModel,
obs_times,
obs_values,
ctrl_values,
obs_times: Real[Array, "obs_time one"],
obs_values: Real[Array, "obs_time observation_dim"],
ctrl_values: Real[Array, "ctrl_time control_dim"],
filter_config: ContinuousTimeKFConfig,
) -> PosteriorGSSMFiltered:
"""Run exact continuous-discrete KF (AffineLinearDrift + constant diffusion + LinearGaussianObservation)."""
Expand All @@ -136,13 +139,13 @@ def _run_linear_kf(
def compute_continuous_filter(
dynamics: DynamicalModel,
filter_config: ContinuousTimeFilterConfig,
key: jax.Array | None = None,
key: PRNGKeyArray | None = None,
*,
obs_times: jax.Array,
obs_values: jax.Array,
ctrl_times=None,
ctrl_values=None,
):
obs_times: Real[Array, " obs_time"],
obs_values: Real[Array, "obs_time observation_dim"],
ctrl_times: Real[Array, " ctrl_time"] | None = None,
ctrl_values: Real[Array, "ctrl_time control_dim"] | None = None,
) -> Any:
"""Pure-JAX continuous-time filter computation (no numpyro side-effects)."""
obs_times_arr = jnp.asarray(obs_times)
if obs_times_arr.ndim == 1:
Expand Down Expand Up @@ -196,14 +199,14 @@ def run_continuous_filter(
name: str,
dynamics: DynamicalModel,
filter_config: ContinuousTimeFilterConfig,
key: jax.Array | None = None,
key: PRNGKeyArray | None = None,
*,
obs_times: jax.Array,
obs_values: jax.Array,
ctrl_times=None,
ctrl_values=None,
obs_times: Real[Array, " obs_time"],
obs_values: Real[Array, "obs_time observation_dim"],
ctrl_times: Real[Array, " ctrl_time"] | None = None,
ctrl_values: Real[Array, "ctrl_time control_dim"] | None = None,
**kwargs,
) -> tuple[jax.Array, object, list[dist.Distribution]]:
) -> tuple[Real[Array, ""], object, list[dist.Distribution]]:
"""Run continuous-time filter via CD-Dynamax.

Pure computation — no numpyro side-effects. Callers are responsible for
Expand Down
27 changes: 15 additions & 12 deletions dynestyx/inference/integrations/cd_dynamax/continuous_smoother.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
"""Continuous-time smoothers via CD-Dynamax: KF, EKF."""

from typing import Any

import jax
import jax.numpy as jnp
import numpyro.distributions as dist
Expand All @@ -10,6 +12,7 @@
cdlgssm_smoother,
cdnlgssm_smoother,
)
from jaxtyping import Array, PRNGKeyArray, Real

from dynestyx.inference.configs.smoother import (
ContinuousTimeEKFSmootherConfig,
Expand All @@ -30,13 +33,13 @@
def compute_continuous_smoother(
dynamics: DynamicalModel,
smoother_config: ContinuousTimeSmootherConfig,
key: jax.Array | None = None,
key: PRNGKeyArray | None = None,
*,
obs_times: jax.Array,
obs_values: jax.Array,
ctrl_times=None,
ctrl_values=None,
):
obs_times: Real[Array, " obs_time"],
obs_values: Real[Array, "obs_time observation_dim"],
ctrl_times: Real[Array, " ctrl_time"] | None = None,
ctrl_values: Real[Array, "ctrl_time control_dim"] | None = None,
) -> Any:
"""Pure-JAX continuous-time smoother computation (no numpyro side-effects)."""
obs_times_arr = jnp.asarray(obs_times)
if obs_times_arr.ndim == 1:
Expand Down Expand Up @@ -118,14 +121,14 @@ def run_continuous_smoother(
name: str,
dynamics: DynamicalModel,
smoother_config: ContinuousTimeSmootherConfig,
key: jax.Array | None = None,
key: PRNGKeyArray | None = None,
*,
obs_times: jax.Array,
obs_values: jax.Array,
ctrl_times=None,
ctrl_values=None,
obs_times: Real[Array, " obs_time"],
obs_values: Real[Array, "obs_time observation_dim"],
ctrl_times: Real[Array, " ctrl_time"] | None = None,
ctrl_values: Real[Array, "ctrl_time control_dim"] | None = None,
**kwargs,
) -> tuple[jax.Array, object, list[dist.Distribution]]:
) -> tuple[Real[Array, ""], object, list[dist.Distribution]]:
"""Run continuous-time smoother via CD-Dynamax.

Pure computation — no numpyro side-effects. Callers are responsible for
Expand Down
38 changes: 25 additions & 13 deletions dynestyx/inference/integrations/cd_dynamax/discrete_filter.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
"""Discrete-time filters via cd-dynamax (dynamax): KF, EKF, UKF."""

import jax
from typing import Any, cast

import jax.numpy as jnp
import numpyro.distributions as dist
from cd_dynamax.dynamax.linear_gaussian_ssm.inference import (
Expand All @@ -14,6 +15,7 @@
UKFHyperParams,
unscented_kalman_filter,
)
from jaxtyping import Array, Real

from dynestyx.inference.configs.filter import (
BaseFilterConfig,
Expand Down Expand Up @@ -83,15 +85,25 @@ def _lti_to_lgssm_params(dynamics: DynamicalModel):
)


def _prepare_inputs(dynamics, obs_values, obs_times, ctrl_times, ctrl_values):
def _prepare_inputs(
dynamics: DynamicalModel,
obs_values: Real[Array, "obs_time observation_dim"],
obs_times: Real[Array, " obs_time"],
ctrl_times: Real[Array, " ctrl_time"] | None,
ctrl_values: Real[Array, "ctrl_time control_dim"] | None,
) -> tuple[
Real[Array, "obs_time observation_dim"],
Real[Array, "obs_time control_dim"],
]:
"""Prepare emissions and inputs arrays for cd-dynamax discrete filters."""
emissions = obs_values
t1 = emissions.shape[0]
control_dim = dynamics.control_dim
if ctrl_values is None:
inputs = jnp.zeros((t1, control_dim))
elif ctrl_values.shape[0] > t1:
inds = jnp.searchsorted(ctrl_times, obs_times, side="left")
aligned_ctrl_times = cast(Real[Array, " ctrl_time"], ctrl_times)
inds = jnp.searchsorted(aligned_ctrl_times, obs_times, side="left")
inputs = ctrl_values[inds]
else:
inputs = ctrl_values
Expand All @@ -102,11 +114,11 @@ def compute_cd_dynamax_discrete_filter(
dynamics: DynamicalModel,
filter_config: BaseFilterConfig,
*,
obs_times: jax.Array,
obs_values: jax.Array,
ctrl_times=None,
ctrl_values=None,
):
obs_times: Real[Array, " obs_time"],
obs_values: Real[Array, "obs_time observation_dim"],
ctrl_times: Real[Array, " ctrl_time"] | None = None,
ctrl_values: Real[Array, "ctrl_time control_dim"] | None = None,
) -> Any:
"""Pure-JAX cd-dynamax discrete filter computation (no numpyro side-effects)."""
emissions, inputs = _prepare_inputs(
dynamics, obs_values, obs_times, ctrl_times, ctrl_values
Expand Down Expand Up @@ -141,12 +153,12 @@ def run_discrete_filter(
dynamics: DynamicalModel,
filter_config: BaseFilterConfig,
*,
obs_times: jax.Array,
obs_values: jax.Array,
ctrl_times=None,
ctrl_values=None,
obs_times: Real[Array, " obs_time"],
obs_values: Real[Array, "obs_time observation_dim"],
ctrl_times: Real[Array, " ctrl_time"] | None = None,
ctrl_values: Real[Array, "ctrl_time control_dim"] | None = None,
**kwargs,
) -> tuple[jax.Array, object, list[dist.Distribution]]:
) -> tuple[Real[Array, ""], object, list[dist.Distribution]]:
"""Run discrete-time filter via cd-dynamax (KF, EKF, UKF).

Pure computation — no numpyro side-effects. Callers are responsible for
Expand Down
Loading
Loading