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
9 changes: 3 additions & 6 deletions pyrenew/latent/subpopulation_infections.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down
57 changes: 36 additions & 21 deletions pyrenew/math.py
100755 → 100644
Original file line number Diff line number Diff line change
Expand Up @@ -30,24 +30,26 @@ 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`.

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
-----
Expand All @@ -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`
Expand All @@ -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`
Expand All @@ -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
-----
Expand Down Expand Up @@ -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,
)

Expand Down
79 changes: 79 additions & 0 deletions test/test_math.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down