Skip to content

ObservationModel

Bases: Module

Observation or emission model for state-space systems.

Defines the conditional distribution of observations given the latent state, control, and time:

\[ y_t \sim p(y_t \mid x_t, u_t, t) \]

Subclasses implement __call__ to return a NumPyro-compatible distribution. The base class provides log_prob and sample for convenience. Subclasses may add parameters (e.g., observation noise scale) as module attributes.

Methods:

Name Description
__call__

Return the observation distribution (a NumPyro distribution; see the NumPyro distributions API) for \(p(y_t \mid x_t, u_t, t)\).

log_prob

Compute \(\log p(y_t \mid x_t, u_t, t)\).

sample

Sample \(y_t \sim p(y_t \mid x_t, u_t, t)\).

masked_observation_log_prob

Score the observed marginal of a distribution returned by an observation model:

import dynestyx as dsx

observation_dist = dynamics.observation_model(state, control, time)
log_likelihood = dsx.masked_observation_log_prob(
    observation_dist, y=observation, obs_mask=observed
)

This also works with plain callable observation models.

Evaluate the marginal log density of observed components.

Partial observations require MultivariateNormal or factorizable Independent(..., 1) distributions. Scalar observations and fully observed or fully missing vectors support other distribution families.

Parameters:

Name Type Description Default
obs_dist Distribution

Scalar or vector observation distribution.

required
y Real[Array, ' observation_dim'] | Real[Array, '']

One observation, broadcast over distribution batch axes. Values at missing components are ignored and may be NaN.

required
obs_mask Bool[Array, ' observation_dim'] | Bool[Array, '']

Boolean mask matching y, with True at observed components.

required

Returns:

Name Type Description
Array Real[Array, '*log_prob_batch']

Observed-marginal log density, retaining distribution batch axes. A fully missing observation contributes zero.

Raises:

Type Description
ValueError

If y and obs_mask do not match the scalar or vector event shape, or partial marginalization is unsupported.

RuntimeError

If unsupported partial marginalization is detected during JIT execution.

Example

Negative Binomial observation model
import jax
import jax.numpy as jnp
from numpyro import distributions as dist
from dynestyx import ObservationModel


class NegativeBinomialObservation(ObservationModel):
    def __init__(self, W: jnp.ndarray, alpha: float = 10.0):
        self.W = W
        self.alpha = alpha  # concentration/over-dispersion parameter

    def __call__(self, x, u, t):
        # log link: mean rate must stay positive
        mean = jnp.exp(self.W @ x)
        return dist.NegativeBinomial2(mean=mean, concentration=self.alpha)


obs_model = NegativeBinomialObservation(
    W=jnp.array([[1.0, -0.5, 0.25]]),
    alpha=8.0,
)

dynamics = DynamicalModel(observation_model=obs_model, ...)