From 3a298fdfeaa5a19173c2ffd72b4396977612efd8 Mon Sep 17 00:00:00 2001 From: Rakesh Pai <41351936+developer-rpai@users.noreply.github.com> Date: Sat, 19 Sep 2026 22:26:18 -0700 Subject: [PATCH] Forward n_timepoints and first_day_dow to ascertainment models (#890) MultiSignalModel now passes n_timepoints (shared model-axis length) and first_day_dow (axis-origin day-of-week) when sampling registered ascertainment models, so ascertainment models with temporal processes can produce full-axis, calendar-aligned trajectories. Existing scalar ascertainment implementations accept **kwargs and are unaffected. Adds AscertainmentModel.requires_calendar_anchor() (default False, mirroring the latent-process API) and includes ascertainment models in MultiSignalModel's obs_start_date entry check. Tests: time-varying day-of-week ascertainment test double verifies the forwarded axis context and calendar alignment; anchor checks fire from both sample() and validate_data(); scalar JointAscertainment output is unchanged with and without the new kwargs; requires_calendar_anchor() defaults to False for scalar implementations. --- pyrenew/ascertainment/base.py | 24 +++- pyrenew/model/multisignal_model.py | 23 +++- test/test_ascertainment.py | 25 ++++ test/test_pyrenew_builder.py | 186 ++++++++++++++++++++++++++++- 4 files changed, 252 insertions(+), 6 deletions(-) diff --git a/pyrenew/ascertainment/base.py b/pyrenew/ascertainment/base.py index 9c0c823f2..ceb08ca47 100644 --- a/pyrenew/ascertainment/base.py +++ b/pyrenew/ascertainment/base.py @@ -134,6 +134,25 @@ def __init__( self.name = name self.signals = signals + def requires_calendar_anchor(self) -> bool: + """ + Report whether this ascertainment model needs a calendar anchor + at sample time. + + The default implementation returns ``False`` for scalar, + time-constant ascertainment rates. Subclasses that sample a + calendar-aligned temporal process (for example a day-of-week + effect) should override this to return ``True``. + + Returns + ------- + bool + ``True`` if the caller of :meth:`sample` must supply a + ``first_day_dow`` (derived from ``obs_start_date`` at the + model entry point); ``False`` otherwise. + """ + return False + def for_signal(self, signal_name: str) -> AscertainmentSignal: """ Return an observation-process accessor for one signal. @@ -177,7 +196,10 @@ def sample(self, **kwargs: object) -> Mapping[str, ArrayLike]: ---------- **kwargs Additional model-context arguments supplied by ``MultiSignalModel``. - Subclasses may ignore unused values. + Currently ``n_timepoints`` (the shared model-axis length) and + ``first_day_dow`` (day-of-week of element 0 of the shared axis, + or ``None`` when no calendar anchor was supplied). Subclasses + may ignore unused values. Returns ------- diff --git a/pyrenew/model/multisignal_model.py b/pyrenew/model/multisignal_model.py index cd8efa144..08818cc8b 100644 --- a/pyrenew/model/multisignal_model.py +++ b/pyrenew/model/multisignal_model.py @@ -45,7 +45,11 @@ class MultiSignalModel(Model): ascertainment_models Optional dictionary mapping names to ascertainment model instances. Each ascertainment model is sampled once per model execution before - observation processes run. + observation processes run. The model forwards the shared model-axis + length (``n_timepoints``) and the axis-origin day-of-week + (``first_day_dow``) to each ascertainment model's ``sample()`` so + time-varying ascertainment can build full-axis, calendar-aligned + trajectories; scalar ascertainment models ignore these arguments. Notes ----- @@ -224,8 +228,9 @@ def _check_obs_start_date( Raises ------ ValueError - If ``obs_start_date`` is ``None`` and any observation or - the latent process requires a calendar anchor. + If ``obs_start_date`` is ``None`` and any observation, the + latent process, or any ascertainment model requires a + calendar anchor. """ if obs_start_date is not None: return @@ -245,6 +250,13 @@ def _check_obs_start_date( "obs_start_date is required when the latent process uses a " "calendar-aligned temporal process." ) + for name, ascertainment_model in self.ascertainment_models.items(): + if ascertainment_model.requires_calendar_anchor(): + raise ValueError( + f"obs_start_date is required when any ascertainment model " + f"uses a calendar-aligned temporal process; " + f"ascertainment model '{name}' does." + ) def shift_times(self, times: jnp.ndarray) -> jnp.ndarray: """ @@ -413,7 +425,10 @@ def sample( } ascertainment_values = { - name: ascertainment_model.sample() + name: ascertainment_model.sample( + n_timepoints=self.latent.n_initialization_points + n_days_post_init, + first_day_dow=first_day_dow, + ) for name, ascertainment_model in self.ascertainment_models.items() } diff --git a/test/test_ascertainment.py b/test/test_ascertainment.py index 6cf77cd2b..68a791e5f 100644 --- a/test/test_ascertainment.py +++ b/test/test_ascertainment.py @@ -481,3 +481,28 @@ def test_context_clears_after_exception(self): with pytest.raises(RuntimeError, match="before ascertainment values"): get_ascertainment_value("he_ascertainment", "hospital") + + +class TestRequiresCalendarAnchor: + """Default calendar-anchor requirements for ascertainment models.""" + + def test_joint_ascertainment_does_not_require_calendar_anchor(self): + """Scalar joint ascertainment rates need no calendar anchor.""" + ascertainment = JointAscertainment( + name="he_ascertainment", + signals=("hospital", "ed"), + baseline_rates=jnp.full(2, 0.5), + scale_tril=jnp.eye(2), + ) + assert ascertainment.requires_calendar_anchor() is False + + def test_ratio_linked_ascertainment_does_not_require_calendar_anchor(self): + """Scalar ratio-linked ascertainment rates need no calendar anchor.""" + ascertainment = RatioLinkedAscertainment( + name="he_ascertainment", + base_signal="hospital", + linked_signal="ed", + base_rate_rv=DeterministicVariable("base_rate", 0.01), + ratio_rv=DeterministicVariable("ratio", 1.5), + ) + assert ascertainment.requires_calendar_anchor() is False diff --git a/test/test_pyrenew_builder.py b/test/test_pyrenew_builder.py index a03af53ea..a7c266812 100644 --- a/test/test_pyrenew_builder.py +++ b/test/test_pyrenew_builder.py @@ -9,7 +9,11 @@ import numpyro.distributions as dist import pytest -from pyrenew.ascertainment import JointAscertainment, RatioLinkedAscertainment +from pyrenew.ascertainment import ( + AscertainmentModel, + JointAscertainment, + RatioLinkedAscertainment, +) from pyrenew.deterministic import DeterministicPMF, DeterministicVariable from pyrenew.latent import ( InfectionsWithFeedback, @@ -986,6 +990,186 @@ def _daily_ed_counts(name="ed"): ) +class _TimeVaryingAscertainment(AscertainmentModel): + """ + Test ascertainment model with a day-of-week effect over the model axis. + + Requires the ``n_timepoints`` and ``first_day_dow`` model-context + arguments at sample time to build a full-axis, calendar-aligned + trajectory, and requires a calendar anchor for ``obs_start_date``. + """ + + def __init__(self, name, signals, baseline_rate, dow_effect): + """ + Initialize the test ascertainment model. + + Parameters + ---------- + name + Name of the ascertainment model. + signals + Unique signal names produced by this model. + baseline_rate + Scalar baseline ascertainment rate. + dow_effect + Multiplicative day-of-week effect of length 7. + """ + super().__init__(name=name, signals=signals) + self.baseline_rate = baseline_rate + self.dow_effect = jnp.asarray(dow_effect) + self.seen_context = {} + + def requires_calendar_anchor(self): + """This model samples a calendar-aligned temporal process.""" + return True + + def sample(self, **kwargs): + """ + Sample a full-axis, calendar-aligned ascertainment trajectory. + + Requires ``n_timepoints`` and ``first_day_dow`` in ``kwargs``. + """ + n_timepoints = kwargs["n_timepoints"] + first_day_dow = kwargs["first_day_dow"] + self.seen_context = { + "n_timepoints": n_timepoints, + "first_day_dow": first_day_dow, + } + dow_indices = (jnp.arange(n_timepoints) + first_day_dow) % 7 + trajectory = self.baseline_rate * self.dow_effect[dow_indices] + numpyro.deterministic(f"{self.name}_trajectory", trajectory) + return {signal: trajectory for signal in self.signals} + + +class TestTimeVaryingAscertainment: + """MultiSignalModel forwards model-axis context to ascertainment models.""" + + DOW_EFFECT = jnp.array([1.0, 0.9, 0.8, 1.0, 1.1, 1.2, 1.0]) + + def _build(self, ascertainment): + """Build a daily ED model wired to the given ascertainment model.""" + latent = PopulationInfections( + name="PopulationInfections", + gen_int_rv=DeterministicPMF("gen_int", jnp.array([0.2, 0.5, 0.3])), + I0_rv=DeterministicVariable("I0", 0.001), + log_rt_time_0_rv=DeterministicVariable("initial_log_rt", 0.0), + single_rt_process=fixed_ar1(autoreg=0.9, innovation_sd=0.05), + n_initialization_points=3, + ) + obs = PopulationCounts( + name="ed", + ascertainment_rate_rv=ascertainment.for_signal("ed"), + delay_distribution_rv=DeterministicPMF("ed_delay", jnp.array([1.0])), + noise=PoissonNoise(), + ) + return MultiSignalModel( + latent, + {"ed": obs}, + ascertainment_models={ascertainment.name: ascertainment}, + ) + + def test_ascertainment_receives_model_axis_context(self): + """The model forwards n_timepoints and first_day_dow to ascertainment.""" + ascertainment = _TimeVaryingAscertainment( + name="tv_asc", + signals=("ed",), + baseline_rate=0.01, + dow_effect=self.DOW_EFFECT, + ) + model = self._build(ascertainment) + n_days = 10 + n_total = model.latent.n_initialization_points + n_days + obs_start_date = _obs_date_for_dow(target_first_day_dow=3, n_init=3) + + with numpyro.handlers.seed(rng_seed=42): + with numpyro.handlers.trace() as trace: + model.sample( + n_days_post_init=n_days, + population_size=1_000_000, + obs_start_date=obs_start_date, + ed={"obs": None}, + ) + + expected_dow = model._resolve_first_day_dow(obs_start_date) + assert ascertainment.seen_context["n_timepoints"] == n_total + assert ascertainment.seen_context["first_day_dow"] == expected_dow + + trajectory = trace["tv_asc_trajectory"]["value"] + assert trajectory.shape == (n_total,) + expected = 0.01 * self.DOW_EFFECT[(jnp.arange(n_total) + expected_dow) % 7] + assert jnp.allclose(trajectory, expected) + + def test_missing_obs_start_date_for_calendar_aligned_ascertainment_raises( + self, + ): + """Calendar-aligned ascertainment models trigger the anchor check.""" + ascertainment = _TimeVaryingAscertainment( + name="tv_asc", + signals=("ed",), + baseline_rate=0.01, + dow_effect=self.DOW_EFFECT, + ) + model = self._build(ascertainment) + + with numpyro.handlers.seed(rng_seed=42): + with pytest.raises(ValueError, match="obs_start_date is required"): + model.sample( + n_days_post_init=10, + population_size=1_000_000, + ed={"obs": None}, + ) + + def test_validate_data_missing_obs_start_date_for_ascertainment_raises(self): + """The anchor check also fires on the validate_data path.""" + ascertainment = _TimeVaryingAscertainment( + name="tv_asc", + signals=("ed",), + baseline_rate=0.01, + dow_effect=self.DOW_EFFECT, + ) + model = self._build(ascertainment) + n_total = model.latent.n_initialization_points + 10 + + with pytest.raises(ValueError, match="ascertainment model 'tv_asc'"): + model.validate_data( + n_days_post_init=10, + ed={"obs": jnp.full(n_total, jnp.nan)}, + ) + + def test_scalar_ascertainment_ignores_new_context_kwargs(self): + """Scalar ascertainment models keep working with the new kwargs.""" + ascertainment = JointAscertainment( + name="he_ascertainment", + signals=("ed",), + baseline_rates=jnp.array([0.5]), + scale_tril=jnp.eye(1), + ) + model = self._build(ascertainment) + obs_start_date = _obs_date_for_dow(target_first_day_dow=3, n_init=3) + + with numpyro.handlers.seed(rng_seed=42): + with numpyro.handlers.trace() as trace_without_anchor: + model.sample( + n_days_post_init=10, + population_size=1_000_000, + ed={"obs": None}, + ) + + with numpyro.handlers.seed(rng_seed=42): + with numpyro.handlers.trace() as trace_with_anchor: + model.sample( + n_days_post_init=10, + population_size=1_000_000, + obs_start_date=obs_start_date, + ed={"obs": None}, + ) + + assert jnp.allclose( + trace_without_anchor["he_ascertainment_ed"]["value"], + trace_with_anchor["he_ascertainment_ed"]["value"], + ) + + class TestBuilderConfigurations: """PyrenewBuilder.build() accepts varied R(t) and observation cadences."""