From 9882d4b535c0e9204b65448a466559ec0511e544 Mon Sep 17 00:00:00 2001 From: Rakesh Pai <41351936+developer-rpai@users.noreply.github.com> Date: Wed, 23 Sep 2026 22:43:53 -0700 Subject: [PATCH] Support observation windows shorter than the right-truncation PMF Generalizes compute_prop_already_reported so that a prediction vector shorter than the reporting-delay PMF support returns the trailing slice of the reported-proportion tail instead of failing. _apply_right_truncation no longer requires the observation window to cover the full delay support. Closes #714. --- pyrenew/convolve.py | 12 ++++ pyrenew/observation/count_observations.py | 10 +--- test/test_convolve.py | 68 +++++++++++++++++++++++ test/test_observation_counts.py | 12 +++- 4 files changed, 92 insertions(+), 10 deletions(-) mode change 100755 => 100644 pyrenew/convolve.py diff --git a/pyrenew/convolve.py b/pyrenew/convolve.py old mode 100755 new mode 100644 index 0a9129f93..a068fafcb --- a/pyrenew/convolve.py +++ b/pyrenew/convolve.py @@ -275,8 +275,20 @@ def compute_prop_already_reported( proportion of events already reported at that timepoint. Earlier timepoints are 1.0 (fully reported); recent timepoints approach reporting_delay_pmf[0] (minimally reported). + + Notes + ----- + When ``n_timepoints`` is shorter than the delay distribution's + support (minus ``right_truncation_offset``), the window covers + only the most recent timepoints, so the trailing slice of the + reported-proportion tail is returned rather than padding with + ones. This equals the last ``n_timepoints`` entries of the + full-support result. """ cdf = jnp.cumsum(reporting_delay_pmf) tail = jnp.flip(cdf[right_truncation_offset:]) + if n_timepoints <= tail.shape[0]: + # Short observation window: keep the most recent timepoints. + return tail[tail.shape[0] - n_timepoints :] n_pad = n_timepoints - tail.shape[0] return jnp.concatenate([jnp.ones(n_pad), tail]) diff --git a/pyrenew/observation/count_observations.py b/pyrenew/observation/count_observations.py index 3aaa7a342..935a0e377 100644 --- a/pyrenew/observation/count_observations.py +++ b/pyrenew/observation/count_observations.py @@ -254,16 +254,12 @@ def _apply_right_truncation( ----- Assumes a single truncation PMF shared across all subpopulations. The 1D proportion array is broadcast to match 2D predicted counts. + Observation windows shorter than the truncation PMF support are + supported: they are scaled by the trailing (most recent) + entries of the reported-proportion tail. """ trunc_pmf = self.right_truncation_rv() n_timepoints = predicted.shape[0] - delay_support = trunc_pmf.shape[0] - right_truncation_offset - if n_timepoints < delay_support: - raise ValueError( - f"Observation window length ({n_timepoints}) must be >= " - f"delay distribution support minus right_truncation_offset " - f"({delay_support})." - ) prop = compute_prop_already_reported( trunc_pmf, n_timepoints, right_truncation_offset ) diff --git a/test/test_convolve.py b/test/test_convolve.py index a9211f85f..6fa5b84ed 100644 --- a/test/test_convolve.py +++ b/test/test_convolve.py @@ -456,3 +456,71 @@ def test_compute_prop_already_reported( ) assert result.shape == (n_timepoints,) assert_array_equal(result, expected) + + +@pytest.mark.parametrize( + ["reporting_delay_pmf", "n_timepoints", "right_truncation_offset", "expected"], + [ + # PMF [0.2, 0.3, 0.5] has CDF [0.2, 0.5, 1.0]. + # offset=0: tail = flip(CDF[0:]) = [1.0, 0.5, 0.2] + # Short windows keep the trailing (most recent) entries. + [ + jnp.array([0.2, 0.3, 0.5]), + 2, + 0, + jnp.array([0.5, 0.2]), + ], + [ + jnp.array([0.2, 0.3, 0.5]), + 1, + 0, + jnp.array([0.2]), + ], + # offset=1: tail = flip(CDF[1:]) = [1.0, 0.5] + [ + jnp.array([0.2, 0.3, 0.5]), + 1, + 1, + jnp.array([0.5]), + ], + # Boundary: window exactly as long as the tail. + [ + jnp.array([0.2, 0.3, 0.5]), + 3, + 0, + jnp.array([1.0, 0.5, 0.2]), + ], + ], +) +def test_compute_prop_already_reported_short_window( + reporting_delay_pmf, + n_timepoints, + right_truncation_offset, + expected, +): + """ + Short observation windows (n_timepoints below the delay PMF + support) return the trailing slice of the reported-proportion + tail instead of raising. + """ + result = pc.compute_prop_already_reported( + reporting_delay_pmf, n_timepoints, right_truncation_offset + ) + assert result.shape == (n_timepoints,) + assert_array_equal(result, expected) + + +@pytest.mark.parametrize("n_timepoints", [1, 2, 3, 4, 5]) +@pytest.mark.parametrize("right_truncation_offset", [0, 1]) +def test_compute_prop_already_reported_short_window_matches_full( + n_timepoints, right_truncation_offset +): + """ + A short window's result must equal the last n_timepoints entries + of the full-support result: truncation only affects recent + timepoints, and a short window covers exactly those. + """ + pmf = jnp.array([0.2, 0.3, 0.5]) + full = pc.compute_prop_already_reported(pmf, 5, right_truncation_offset) + short = pc.compute_prop_already_reported(pmf, n_timepoints, right_truncation_offset) + assert_array_equal(short, full[-n_timepoints:]) diff --git a/test/test_observation_counts.py b/test/test_observation_counts.py index 5fd5b6abb..9a2b54d9d 100644 --- a/test/test_observation_counts.py +++ b/test/test_observation_counts.py @@ -524,8 +524,8 @@ def test_validate_catches_invalid_rt_pmf(self, simple_delay_pmf): with pytest.raises(ValueError, match="must sum to 1.0"): process.validate() - def test_short_observation_window_raises(self, simple_delay_pmf): - """Test that observation window shorter than delay support raises.""" + def test_short_observation_window_supported(self, simple_delay_pmf): + """Short observation windows use the trailing tail slice, not an error.""" rt_pmf = jnp.array([0.2, 0.3, 0.5]) process = PopulationCounts( name="test", @@ -537,13 +537,19 @@ def test_short_observation_window_raises(self, simple_delay_pmf): infections = jnp.ones(2) * 100 with numpyro.handlers.seed(rng_seed=42): - with pytest.raises(ValueError, match="Observation window length"): + with numpyro.handlers.trace() as trace: process.sample( infections=infections, obs=None, right_truncation_offset=0, ) + # PMF [0.2, 0.3, 0.5] has CDF [0.2, 0.5, 1.0]; a 2-day window + # keeps the trailing (most recent) tail entries [0.5, 0.2]. + prop = trace["test_prop_already_reported"]["value"] + assert prop.shape == (2,) + assert jnp.all(jnp.isclose(prop, jnp.array([0.5, 0.2]))) + def test_counts_by_subpop_2d_broadcasting(self): """Test right-truncation with SubpopulationCounts 2D infections.""" rt_pmf = jnp.array([0.2, 0.3, 0.5])