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
12 changes: 12 additions & 0 deletions pyrenew/convolve.py
100755 → 100644
Original file line number Diff line number Diff line change
Expand Up @@ -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])
10 changes: 3 additions & 7 deletions pyrenew/observation/count_observations.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down
68 changes: 68 additions & 0 deletions test/test_convolve.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:])
12 changes: 9 additions & 3 deletions test/test_observation_counts.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -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])
Expand Down