Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
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
14 changes: 10 additions & 4 deletions dynestyx/evaluation/handlers.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,7 +78,10 @@ def _sample_ds(
scoring_config=self.observation_scoring_config,
plate_shapes=plate_shapes,
)
filtered_result.evaluation_result = evaluation_result
filtered_result = dataclasses.replace(
filtered_result,
evaluation_result=evaluation_result,
)

forwarded_result = fwd(
name,
Expand All @@ -95,9 +98,12 @@ def _sample_ds(
evaluation_result=evaluation_result,
**kwargs,
)
evaluation_result._register_numpyro_sites = chain_numpyro_site_registrations(
evaluation_result._register_numpyro_sites,
getattr(forwarded_result, "_register_numpyro_sites", None),
evaluation_result = dataclasses.replace(
evaluation_result,
_register_numpyro_sites=chain_numpyro_site_registrations(
evaluation_result._register_numpyro_sites,
getattr(forwarded_result, "_register_numpyro_sites", None),
),
)
return evaluation_result

Expand Down
15 changes: 12 additions & 3 deletions dynestyx/inference/filters.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,7 @@
from dynestyx.models import DynamicalModel
from dynestyx.types import (
ConditionedResult,
EvaluationResult,
FunctionOfTime,
chain_numpyro_site_registrations,
)
Expand Down Expand Up @@ -144,9 +145,17 @@ def _sample_ds(
)

forwarded_register = getattr(forwarded_result, "_register_numpyro_sites", None)
result._register_numpyro_sites = chain_numpyro_site_registrations(
result._register_numpyro_sites,
forwarded_register,
if isinstance(forwarded_result, EvaluationResult):
result = dataclasses.replace(
result,
evaluation_result=forwarded_result,
)
result = dataclasses.replace(
result,
_register_numpyro_sites=chain_numpyro_site_registrations(
result._register_numpyro_sites,
forwarded_register,
),
)

return result
Expand Down
5 changes: 2 additions & 3 deletions dynestyx/inference/observation_predictions.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,9 +6,9 @@

from __future__ import annotations

import dataclasses
from typing import Any

import equinox as eqx
import jax
import jax.numpy as jnp
import numpyro
Expand Down Expand Up @@ -39,8 +39,7 @@
)


@dataclasses.dataclass(frozen=True)
class PredictedObservationOutputs:
class PredictedObservationOutputs(eqx.Module):
"""Canonical predicted-observation outputs for Dynestyx filters."""

mean: Float[Array, "*plate time observation_dim"] | None = None
Expand Down
9 changes: 6 additions & 3 deletions dynestyx/inference/smoothers.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,9 +180,12 @@ def _sample_ds(
)

forwarded_register = getattr(forwarded_result, "_register_numpyro_sites", None)
result._register_numpyro_sites = chain_numpyro_site_registrations(
result._register_numpyro_sites,
forwarded_register,
result = dataclasses.replace(
result,
_register_numpyro_sites=chain_numpyro_site_registrations(
result._register_numpyro_sites,
forwarded_register,
),
)

return result
Expand Down
13 changes: 6 additions & 7 deletions dynestyx/observation_missingness.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,10 +2,10 @@

from __future__ import annotations

import dataclasses
from collections.abc import Callable
from typing import Literal

import equinox as eqx
import jax.numpy as jnp
import jax.scipy as jsp
import numpy as np
Expand All @@ -27,8 +27,7 @@
MissingObservationStrategy = Literal["auto", "marginalize", "augment", "error"]


@dataclasses.dataclass
class MissingObservationMetadata:
class MissingObservationMetadata(eqx.Module):
"""Describe the missing entries in one observation array.

Flattened indices list missing entries by time and then by component.
Expand All @@ -54,10 +53,10 @@ class MissingObservationMetadata:
missing_obs_times: Real[Array, " n_missing_obs"]
missing_obs_coordinate_indices: Int[Array, " n_missing_obs"] | None
missing_flat_indices: Int[Array, " n_missing_obs"]
observation_shape: tuple[int, ...]
has_missing: bool
has_partial_missing: bool
has_fully_missing_rows: bool
observation_shape: tuple[int, ...] = eqx.field(static=True)
has_missing: bool = eqx.field(static=True)
has_partial_missing: bool = eqx.field(static=True)
has_fully_missing_rows: bool = eqx.field(static=True)


def _concrete_observation_mask(
Expand Down
17 changes: 12 additions & 5 deletions dynestyx/simulation/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
)
from dynestyx.types import (
ConditionedResult,
EvaluationResult,
SimulatedResult,
chain_numpyro_site_registrations,
)
Expand Down Expand Up @@ -504,13 +505,19 @@ def _register_self(site_name: str) -> None:
downstream_register = getattr(
downstream_result, "_register_numpyro_sites", None
)
combined_register = chain_numpyro_site_registrations(
_register_self,
results._register_numpyro_sites,
downstream_register,
)
if isinstance(downstream_result, EvaluationResult):
return dataclasses.replace(
downstream_result,
_register_numpyro_sites=combined_register,
)
return dataclasses.replace(
results,
_register_numpyro_sites=chain_numpyro_site_registrations(
_register_self,
results._register_numpyro_sites,
downstream_register,
),
_register_numpyro_sites=combined_register,
)

def simulate(
Expand Down
22 changes: 8 additions & 14 deletions dynestyx/types.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,5 @@
"""Shared typing helpers for dynamical systems."""

import dataclasses
from collections.abc import Callable
from typing import Protocol, runtime_checkable

Expand All @@ -17,25 +16,21 @@ def __call__(
raise NotImplementedError()


@dataclasses.dataclass
class EvaluationResult:
class EvaluationResult(eqx.Module):
"""Outputs computed by an evaluation handler.

Evaluation handlers attach this object to the ``ConditionedResult`` they
consume. NumPyro registration remains deferred so ``dsx.condition`` stays
side-effect free while ``dsx.sample`` can register the same outputs later.
"""

observation_scores: dict[str, Real[Array, "..."]] = dataclasses.field(
default_factory=dict
)
_register_numpyro_sites: Callable[[str], None] | None = dataclasses.field(
default=None, repr=False
observation_scores: dict[str, Real[Array, "..."]] = eqx.field(default_factory=dict)
_register_numpyro_sites: Callable[[str], None] | None = eqx.field(
default=None, repr=False, static=True
)


@dataclasses.dataclass
class ConditionedResult:
class ConditionedResult(eqx.Module):
"""Common base for results from the NumPyro-free conditioning primitive.

``dsx.condition`` returns this type under both ``Filter`` and ``Smoother``.
Expand All @@ -55,8 +50,8 @@ class ConditionedResult:
dists: list | None = None
predicted_observations: object = None
evaluation_result: EvaluationResult | None = None
_register_numpyro_sites: Callable[[str], None] | None = dataclasses.field(
default=None, repr=False
_register_numpyro_sites: Callable[[str], None] | None = eqx.field(
default=None, repr=False, static=True
)

def __call__(
Expand All @@ -68,8 +63,7 @@ def __call__(
)


@dataclasses.dataclass
class LatentStateResult:
class LatentStateResult(eqx.Module):
"""Result of latent-state construction / scoring without NumPyro side effects.

Let ``z = state_path_params`` denote the free variables used to
Expand Down
Loading