ObservationModel¶
Bases: Module
Observation or emission model for state-space systems.
Defines the conditional distribution of observations given the latent state, control, and time:
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 |
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 |
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, ...)