From ce52988bcfce61197d8c955d3ece9db972e6d069 Mon Sep 17 00:00:00 2001 From: Rakesh Pai <41351936+developer-rpai@users.noreply.github.com> Date: Wed, 23 Sep 2026 23:06:58 -0700 Subject: [PATCH 1/2] Vectorize r_approx_from_R over array-valued R Generalize r_approx_from_R, neg_MGF, and neg_MGF_del_r so the reproduction number (and rate) may be arrays of any shape. The MGF sums are taken over the weights axis only, giving each entry an independent Newton solve with output shaped like the input. Scalar inputs behave exactly as before. Replace the jax.vmap workaround in SubpopulationInfections with a direct vectorized call. Add five regression tests to test_math.py. Closes #486. --- pyrenew/latent/subpopulation_infections.py | 9 +-- pyrenew/math.py | 57 ++++++++++------ test/test_math.py | 79 ++++++++++++++++++++++ 3 files changed, 118 insertions(+), 27 deletions(-) mode change 100755 => 100644 pyrenew/math.py diff --git a/pyrenew/latent/subpopulation_infections.py b/pyrenew/latent/subpopulation_infections.py index 0182966fe..5e757566b 100644 --- a/pyrenew/latent/subpopulation_infections.py +++ b/pyrenew/latent/subpopulation_infections.py @@ -4,9 +4,6 @@ from __future__ import annotations -from functools import partial - -import jax import jax.numpy as jnp import numpyro from jax.typing import ArrayLike @@ -252,9 +249,9 @@ def sample( I0_subpop = self._validate_and_prepare_I0(jnp.asarray(self.I0_rv()), pop) - initial_r_subpop = jax.vmap( - partial(r_approx_from_R, g=gen_int, n_newton_steps=4) - )(rt_subpop[0, :]) + initial_r_subpop = r_approx_from_R( + R=rt_subpop[0, :], g=gen_int, n_newton_steps=4 + ) time_indices = jnp.arange(self.n_initialization_points) I0_all = I0_subpop[jnp.newaxis, :] * jnp.exp( diff --git a/pyrenew/math.py b/pyrenew/math.py old mode 100755 new mode 100644 index a77cc3ecf..e42b9e3de --- a/pyrenew/math.py +++ b/pyrenew/math.py @@ -30,7 +30,7 @@ def _positive_ints_like(vec: ArrayLike) -> jnp.ndarray: return jnp.arange(1, jnp.size(jnp.asarray(vec)) + 1) -def neg_MGF(r: float, w: ArrayLike) -> float: +def neg_MGF(r: ArrayLike, w: ArrayLike) -> ArrayLike: """ Compute the negative moment generating function (MGF) for a given rate `r` and weights `w`. @@ -38,16 +38,18 @@ def neg_MGF(r: float, w: ArrayLike) -> float: Parameters ---------- r - The rate parameter. + The rate parameter. May be a scalar or an array, in which + case the MGF is evaluated independently for each entry. w - An array of weights. + A 1D array of weights. Returns ------- - float + ArrayLike The value of the negative MGF evaluated at `r` - and `w`. + and `w`. Has the same shape as `r` + (a scalar when `r` is a scalar). Notes ----- @@ -63,12 +65,17 @@ def neg_MGF(r: float, w: ArrayLike) -> float: ```math M_-(r) = \\sum_{t = 1}^{n} w_i \\exp(-rt) ``` - """ - return jnp.sum(w * jnp.exp(-r * _positive_ints_like(w))) + The sum is always over the weights (last) axis, so each + entry of an array-valued `r` gets its own weighted sum. + """ + w_arr = jnp.asarray(w) + t_vec = _positive_ints_like(w_arr) + r_col = jnp.asarray(r)[..., jnp.newaxis] + return jnp.sum(w_arr * jnp.exp(-r_col * t_vec), axis=-1) -def neg_MGF_del_r(r: float, w: ArrayLike) -> float: +def neg_MGF_del_r(r: ArrayLike, w: ArrayLike) -> ArrayLike: """ Compute the value of the partial deriative of [`pyrenew.math.neg_MGF`][] with respect to `r` @@ -77,22 +84,26 @@ def neg_MGF_del_r(r: float, w: ArrayLike) -> float: Parameters ---------- r - The rate parameter. + The rate parameter. May be a scalar or an array, in which + case the derivative is evaluated independently for each entry. w - An array of weights. + A 1D array of weights. Returns ------- - float + ArrayLike The value of the partial derivative evaluated at `r` - and `w`. + and `w`. Has the same shape as `r` + (a scalar when `r` is a scalar). """ - t_vec = _positive_ints_like(w) - return -jnp.sum(w * t_vec * jnp.exp(-r * t_vec)) + w_arr = jnp.asarray(w) + t_vec = _positive_ints_like(w_arr) + r_col = jnp.asarray(r)[..., jnp.newaxis] + return -jnp.sum(w_arr * t_vec * jnp.exp(-r_col * t_vec), axis=-1) -def r_approx_from_R(R: float, g: ArrayLike, n_newton_steps: int) -> ArrayLike: +def r_approx_from_R(R: ArrayLike, g: ArrayLike, n_newton_steps: int) -> ArrayLike: """ Get the approximate asymptotic geometric growth rate `r` for a renewal process with a fixed reproduction number `R` @@ -103,19 +114,22 @@ def r_approx_from_R(R: float, g: ArrayLike, n_newton_steps: int) -> ArrayLike: Parameters ---------- R - The reproduction number + The reproduction number. May be a scalar or an array of + any shape, in which case the growth rate is computed + independently for each entry. g The probability mass function of the generation - interval. + interval, as a 1D array. n_newton_steps Number of steps to take when performing Newton's method. Returns ------- - float - The approximate value of `r`. + ArrayLike + The approximate value(s) of `r`, with the same shape + as `R` (a scalar when `R` is a scalar). Notes ----- @@ -143,14 +157,15 @@ def r_approx_from_R(R: float, g: ArrayLike, n_newton_steps: int) -> ArrayLike: We then refine this approximation by applying Newton's method for a fixed number of steps. """ + R_arr = jnp.asarray(R) mean_gi = jnp.dot(g, _positive_ints_like(g)) - init_r = (R - 1) / (R * mean_gi) + init_r = (R_arr - 1) / (R_arr * mean_gi) def _r_next( r: ArrayLike, _: None ) -> tuple[ArrayLike, None]: # numpydoc ignore=GL08 return ( - r - ((R * neg_MGF(r, g) - 1) / (R * neg_MGF_del_r(r, g))), + r - ((R_arr * neg_MGF(r, g) - 1) / (R_arr * neg_MGF_del_r(r, g))), None, ) diff --git a/test/test_math.py b/test/test_math.py index fe2e96609..13a4b1c48 100644 --- a/test/test_math.py +++ b/test/test_math.py @@ -60,6 +60,85 @@ def test_r_approx(R, G): assert_almost_equal(e_val, 1, decimal=5) +def test_r_approx_vectorized_matches_scalar(): + """ + Test that r_approx_from_R with an array of R values + gives the same answers as calling it once per scalar + R value. + """ + vec_rng = RandomState(11) + G = vec_rng.dirichlet(np.ones(8)) + R_vec = jnp.array([0.5, 0.99, 1.0, 1.01, 1.5, 3.0]) + r_vec = pmath.r_approx_from_R(R_vec, G, n_newton_steps=8) + r_expected = jnp.array( + [pmath.r_approx_from_R(float(R), G, n_newton_steps=8) for R in R_vec] + ) + assert r_vec.shape == R_vec.shape + assert_array_almost_equal(r_vec, r_expected) + + +def test_r_approx_vectorized_satisfies_defining_equation(): + """ + Test that each entry of a vectorized r_approx_from_R + result satisfies the defining equation + R * M_-(r) - 1 == 0 for its own R value. + """ + vec_rng = RandomState(12) + G = vec_rng.dirichlet(np.ones(6)) + R_vec = jnp.array([0.7, 1.2, 2.5]) + r_vec = pmath.r_approx_from_R(R_vec, G, n_newton_steps=8) + residuals = R_vec * pmath.neg_MGF(r_vec, G) - 1 + assert residuals.shape == R_vec.shape + assert_array_almost_equal(residuals, jnp.zeros_like(residuals), decimal=5) + + +def test_r_approx_vectorized_multidimensional(): + """ + Test that r_approx_from_R preserves the shape of a + multidimensional R array and matches scalar calls + entry by entry. + """ + vec_rng = RandomState(13) + G = vec_rng.dirichlet(np.ones(5)) + R_mat = jnp.array([[0.8, 1.0, 1.4], [2.0, 0.9, 1.1]]) + r_mat = pmath.r_approx_from_R(R_mat, G, n_newton_steps=8) + assert r_mat.shape == R_mat.shape + for i in range(R_mat.shape[0]): + for j in range(R_mat.shape[1]): + r_scalar = pmath.r_approx_from_R(float(R_mat[i, j]), G, n_newton_steps=8) + assert_almost_equal(float(r_mat[i, j]), float(r_scalar)) + + +def test_neg_MGF_batched(): + """ + Test that neg_MGF and neg_MGF_del_r evaluate + independently for each entry of an array-valued r. + """ + vec_rng = RandomState(14) + w = vec_rng.dirichlet(np.ones(7)) + r_vec = jnp.array([-0.1, 0.0, 0.2]) + mgf_vec = pmath.neg_MGF(r_vec, w) + dmgf_vec = pmath.neg_MGF_del_r(r_vec, w) + assert mgf_vec.shape == r_vec.shape + assert dmgf_vec.shape == r_vec.shape + for k in range(len(r_vec)): + assert_almost_equal(float(mgf_vec[k]), float(pmath.neg_MGF(float(r_vec[k]), w))) + assert_almost_equal( + float(dmgf_vec[k]), float(pmath.neg_MGF_del_r(float(r_vec[k]), w)) + ) + + +def test_r_approx_scalar_still_scalar(): + """ + Test backward compatibility: a scalar R still yields + a scalar r. + """ + G = np.array([0.2, 0.1, 0.2, 0.15, 0.05, 0.025, 0.025, 0.25]) + r_val = pmath.r_approx_from_R(1.2, G, n_newton_steps=5) + assert jnp.asarray(r_val).shape == () + assert_almost_equal(float(1.2 * pmath.neg_MGF(r_val, G) - 1), 0.0, decimal=5) + + def test_asymptotic_properties(): """ Check that the calculated From 08680d3152f802c7abe2a4a7adb0e986bcca94cb Mon Sep 17 00:00:00 2001 From: "Dylan H. Morris" Date: Mon, 28 Sep 2026 10:59:52 -0400 Subject: [PATCH 2/2] Update test/test_math.py --- test/test_math.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/test/test_math.py b/test/test_math.py index 13a4b1c48..f41a879f0 100644 --- a/test/test_math.py +++ b/test/test_math.py @@ -134,9 +134,10 @@ def test_r_approx_scalar_still_scalar(): a scalar r. """ G = np.array([0.2, 0.1, 0.2, 0.15, 0.05, 0.025, 0.025, 0.25]) - r_val = pmath.r_approx_from_R(1.2, G, n_newton_steps=5) + R = 1.2 + r_val = pmath.r_approx_from_R(R, G, n_newton_steps=5) assert jnp.asarray(r_val).shape == () - assert_almost_equal(float(1.2 * pmath.neg_MGF(r_val, G) - 1), 0.0, decimal=5) + assert_almost_equal(float(R * pmath.neg_MGF(r_val, G) - 1), 0.0, decimal=5) def test_asymptotic_properties():