diff --git a/pyrenew/latent/subpopulation_infections.py b/pyrenew/latent/subpopulation_infections.py index 0182966f..5e757566 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 a77cc3ec..e42b9e3d --- 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 fe2e9660..f41a879f 100644 --- a/test/test_math.py +++ b/test/test_math.py @@ -60,6 +60,86 @@ 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 = 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(R * pmath.neg_MGF(r_val, G) - 1), 0.0, decimal=5) + + def test_asymptotic_properties(): """ Check that the calculated