gemlib 0.9.2__tar.gz → 0.9.4__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (74) hide show
  1. {gemlib-0.9.2 → gemlib-0.9.4}/PKG-INFO +6 -3
  2. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/continuous_markov.py +15 -14
  3. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/continuous_time_state_transition_model.py +41 -20
  4. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/continuous_time_state_transition_model_test.py +169 -8
  5. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_time_state_transition_model.py +8 -8
  6. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_time_state_transition_model_examples.py +5 -5
  7. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_time_state_transition_model_test.py +2 -2
  8. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/chain_binomial_rippler.py +2 -2
  9. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/chain_binomial_rippler_test.py +2 -2
  10. {gemlib-0.9.2 → gemlib-0.9.4}/pyproject.toml +6 -2
  11. {gemlib-0.9.2 → gemlib-0.9.4}/LICENSE +0 -0
  12. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/__init__.py +0 -0
  13. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/__init__.py +0 -0
  14. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/brownian.py +0 -0
  15. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/categorical2.py +0 -0
  16. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_markov.py +0 -0
  17. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_rejection_sampling.py +0 -0
  18. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/experimental/__init__.py +0 -0
  19. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/experimental/discrete_approx_cont_state_transition_model.py +0 -0
  20. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/experimental/state_transition_marginal_model.py +0 -0
  21. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/hypergeometric.py +0 -0
  22. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/hypergeometric_sampler.py +0 -0
  23. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/hypergeometric_test.py +0 -0
  24. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/kcategorical.py +0 -0
  25. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/kcategorical_test.py +0 -0
  26. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/uniform_integer.py +0 -0
  27. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/uniform_integer_test.py +0 -0
  28. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/__init__.py +0 -0
  29. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/adaptive_random_walk_metropolis.py +0 -0
  30. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/adaptive_random_walk_metropolis_test.py +0 -0
  31. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/bb_fixture.pkl +0 -0
  32. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/brownian_bridge_kernel.py +0 -0
  33. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/brownian_bridge_kernel_test.py +0 -0
  34. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/compound_kernel.py +0 -0
  35. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/conftest.py +0 -0
  36. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/damped_chain_binomial_rippler.py +0 -0
  37. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/__init__.py +0 -0
  38. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/event_time_proposal.py +0 -0
  39. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/fixtures.py +0 -0
  40. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh.py +0 -0
  41. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh_test.py +0 -0
  42. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal.py +0 -0
  43. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal_test.py +0 -0
  44. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/move_events.py +0 -0
  45. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/move_events_test.py +0 -0
  46. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh.py +0 -0
  47. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh_test.py +0 -0
  48. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_proposal.py +0 -0
  49. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/util.py +0 -0
  50. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/__init__.py +0 -0
  51. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/composable_kernel.py +0 -0
  52. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/__init__.py +0 -0
  53. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/left_censored_events_mh.py +0 -0
  54. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events.py +0 -0
  55. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events_test.py +0 -0
  56. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh.py +0 -0
  57. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh_test.py +0 -0
  58. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/hmc.py +0 -0
  59. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/hmc_test.py +0 -0
  60. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/mcmc_base.py +0 -0
  61. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/mcmc_sampler.py +0 -0
  62. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/mcmc_sampler_test.py +0 -0
  63. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/multi_scan.py +0 -0
  64. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/multi_scan_test.py +0 -0
  65. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/random_walk_metropolis.py +0 -0
  66. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/random_walk_metropolis_test.py +0 -0
  67. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/test_util.py +0 -0
  68. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/gibbs_kernel.py +0 -0
  69. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/gibbs_kernel_test.py +0 -0
  70. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/h5_posterior.py +0 -0
  71. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/multi_scan_kernel.py +0 -0
  72. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/zarr_posterior.py +0 -0
  73. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/util.py +0 -0
  74. {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/util_test.py +0 -0
@@ -1,10 +1,13 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: gemlib
3
- Version: 0.9.2
3
+ Version: 0.9.4
4
4
  Summary: GEMlib scientific compute library for epidemic modelling
5
- Home-page: http://fhm-chicas-code.lancs.ac.uk/GEM/gemlib
5
+ Home-page: https://gem-epidemics.gitlab.io/gemlib
6
+ Keywords: epidemic,Bayesian,inference,infectious disease model,probabilistic programming
6
7
  Author: Chris Jewell
7
8
  Author-email: c.jewell@lancaster.ac.uk
9
+ Maintainer: Jessica Bridgen
10
+ Maintainer-email: j.bridgen@lancaster.ac.uk
8
11
  Requires-Python: >=3.9.0,<3.12.0
9
12
  Classifier: Programming Language :: Python :: 3
10
13
  Classifier: Programming Language :: Python :: 3.9
@@ -16,4 +19,4 @@ Requires-Dist: tensorflow (>=2.15.0,<2.16.0) ; sys_platform == "linux"
16
19
  Requires-Dist: tensorflow-cpu (>=2.15.0,<2.16.0) ; sys_platform == "darwin"
17
20
  Requires-Dist: tensorflow-intel (>=2.15.0,<2.16.0) ; sys_platform == "win32"
18
21
  Requires-Dist: tensorflow-probability (>=0.23.0,<0.24.0)
19
- Project-URL: Repository, http://fhm-chicas-code.lancs.ac.uk/GEM/gemlib
22
+ Project-URL: Repository, https://gitlab.com/gem-epidemics/gemlib
@@ -14,7 +14,7 @@ Tensor = tf.Tensor
14
14
  DTYPE = tf.float32
15
15
 
16
16
 
17
- class EpidemicEvent(NamedTuple):
17
+ class EventList(NamedTuple):
18
18
  """Tracker of an event in an epidemic simulation
19
19
 
20
20
  Attributes:
@@ -65,6 +65,7 @@ def _total_flux(transition_rates, state, incidence_matrix):
65
65
  A [R,N] tensor of total flux along each transition, taking into account the
66
66
  availability of individuals in the source state.
67
67
  """
68
+
68
69
  source_state_idx = transition_coords(incidence_matrix)[:, 0]
69
70
  source_states = batch_gather(state, indices=source_state_idx[:, tf.newaxis])
70
71
  transition_rates = tf.stack(transition_rates, axis=-1)
@@ -75,7 +76,7 @@ def _total_flux(transition_rates, state, incidence_matrix):
75
76
  def compute_state(
76
77
  incidence_matrix: Tensor,
77
78
  initial_state: Tensor,
78
- event_list: EpidemicEvent,
79
+ event_list: EventList,
79
80
  include_final_state: bool = False,
80
81
  ):
81
82
  """Given an event list `event_list`, compute a timeseries
@@ -140,7 +141,7 @@ def compute_state(
140
141
 
141
142
  def exponential_propogate(
142
143
  transition_rate_fn: Callable, incidence_matrix: Tensor
143
- ) -> EpidemicEvent:
144
+ ) -> EventList:
144
145
  """Generates a function for propogating an epidemic forward in time
145
146
 
146
147
  Closure over the transition rate function and the incidence matrix
@@ -156,7 +157,7 @@ def exponential_propogate(
156
157
  `R` transitions and the columns correspond to the `S` states.
157
158
 
158
159
  Returns:
159
- EpidemicEvent: A NamedTuple that describes the next event in the
160
+ EventList: A NamedTuple that describes the next event in the
160
161
  epidemic.
161
162
  """
162
163
  tr_incidence_matrix = tf.transpose(incidence_matrix)
@@ -170,7 +171,7 @@ def exponential_propogate(
170
171
  state (tensor): `[N,S]` representing the current state.
171
172
 
172
173
  Returns:
173
- EpidemicEvent: The next event in the epidemic.
174
+ EventList: The next event in the epidemic.
174
175
  """
175
176
  seed_exp, seed_cat = tfp.random.split_seed(seed, n=2)
176
177
  num_units = state.shape[-2]
@@ -204,7 +205,7 @@ def exponential_propogate(
204
205
  return (
205
206
  time + t_next,
206
207
  new_state,
207
- EpidemicEvent(time + t_next, transition_idx, unit_idx),
208
+ EventList(time + t_next, transition_idx, unit_idx),
208
209
  )
209
210
 
210
211
  return propogate_fn
@@ -217,7 +218,7 @@ def continuous_markov_simulation(
217
218
  num_markov_jumps: int,
218
219
  initial_time: float = 0.0,
219
220
  seed=None,
220
- ) -> EpidemicEvent:
221
+ ) -> EventList:
221
222
  """
222
223
  Simulates a continuous-time Markov process
223
224
 
@@ -231,17 +232,17 @@ def continuous_markov_simulation(
231
232
  state transition model with S states and R transitions.
232
233
  seed (Optional[List(int,int)): The random seed.
233
234
  Returns:
234
- EpidemicEvent: An object containing the simulated epidemic events.
235
+ EventList: An object containing the simulated epidemic events.
235
236
 
236
237
  """
237
238
  initial_state = tf.convert_to_tensor(initial_state)
238
239
  incidence_matrix = tf.convert_to_tensor(incidence_matrix)
239
- dtype = initial_state.dtype
240
+ dtype = tf.float32
240
241
  seed = tfp.random.sanitize_seed(seed, salt="continuous_markov_simulation")
241
242
 
242
243
  propagate_fn = exponential_propogate(transition_rate_fn, incidence_matrix)
243
244
 
244
- accum = EpidemicEvent(
245
+ accum = EventList(
245
246
  time=tf.TensorArray(dtype, size=num_markov_jumps, dynamic_size=False),
246
247
  transition=tf.TensorArray(
247
248
  tf.int32, size=num_markov_jumps, dynamic_size=False
@@ -261,7 +262,7 @@ def continuous_markov_simulation(
261
262
  def body(i, time, state, seed, accum):
262
263
  next_seed, this_seed = tfp.random.split_seed(seed, salt="body")
263
264
  next_time, next_state, event = propagate_fn(time, state, this_seed)
264
- accum = EpidemicEvent(*[x.write(i, y) for x, y in zip(accum, event)])
265
+ accum = EventList(*[x.write(i, y) for x, y in zip(accum, event)])
265
266
  return i + 1, next_time, next_state, next_seed, accum
266
267
 
267
268
  actual_markov_jumps, _, _, _, accum = tf.while_loop(
@@ -273,7 +274,7 @@ def continuous_markov_simulation(
273
274
  indices = tf.range(actual_markov_jumps, num_markov_jumps)
274
275
  fills = tf.fill([num_markov_jumps - actual_markov_jumps], np.inf)
275
276
 
276
- output = EpidemicEvent(
277
+ output = EventList(
277
278
  time=accum.time.scatter(
278
279
  indices,
279
280
  fills,
@@ -299,7 +300,7 @@ def continuous_time_log_likelihood(
299
300
  initial_state: Tensor,
300
301
  initial_time: float,
301
302
  num_jumps: int,
302
- event_list: EpidemicEvent,
303
+ event_list: EventList,
303
304
  ) -> float:
304
305
  """
305
306
  Computes the log-likelihood of a continuous-time Markov process
@@ -313,7 +314,7 @@ def continuous_time_log_likelihood(
313
314
  the connections between states in `[S,R]` format.
314
315
  initial_state: The initial state of the process as a `[N,R]`.
315
316
  num_jumps (int): The number of jumps to simulate.
316
- event (EpidemicEvent): The event data containing the times
317
+ event (EventList): The event data containing the times
317
318
  and states.
318
319
 
319
320
  Returns:
@@ -7,7 +7,7 @@ import tensorflow_probability as tfp
7
7
  from tensorflow_probability.python.internal import reparameterization
8
8
 
9
9
  from gemlib.distributions.continuous_markov import (
10
- EpidemicEvent,
10
+ EventList,
11
11
  compute_state,
12
12
  continuous_markov_simulation,
13
13
  continuous_time_log_likelihood,
@@ -55,10 +55,18 @@ class ContinuousTimeStateTransitionModel(tfd.Distribution):
55
55
 
56
56
  self._incidence_matrix = tf.convert_to_tensor(incidence_matrix)
57
57
  self._initial_state = tf.convert_to_tensor(initial_state)
58
- self._initial_time = tf.convert_to_tensor(initial_time)
58
+ self._initial_time = tf.convert_to_tensor(
59
+ initial_time, dtype=self._initial_state.dtype
60
+ )
61
+
62
+ dtype = EventList(
63
+ time=self._initial_time.dtype,
64
+ transition=tf.int32,
65
+ individual=tf.int32,
66
+ )
59
67
 
60
68
  super().__init__(
61
- dtype=self._incidence_matrix.dtype,
69
+ dtype=dtype,
62
70
  reparameterization_type=reparameterization.FULLY_REPARAMETERIZED,
63
71
  validate_args=validate_args,
64
72
  allow_nan_stats=allow_nan_stats,
@@ -92,7 +100,7 @@ class ContinuousTimeStateTransitionModel(tfd.Distribution):
92
100
  return self._parameters["initial_time"]
93
101
 
94
102
  def compute_state(
95
- self, event_list: EpidemicEvent, include_final_state: bool = False
103
+ self, event_list: EventList, include_final_state: bool = False
96
104
  ) -> Tensor:
97
105
  """Given an event list `event_list`, compute a timeseries
98
106
  of state given the model.
@@ -118,10 +126,10 @@ class ContinuousTimeStateTransitionModel(tfd.Distribution):
118
126
  )
119
127
 
120
128
  # Bypass the reshaping that tfd.Distribution._call_sample_n does
121
- def _call_sample_n(self, sample_shape, seed) -> EpidemicEvent:
129
+ def _call_sample_n(self, sample_shape, seed) -> EventList:
122
130
  return self._sample_n(sample_shape, seed)
123
131
 
124
- def _sample_n(self, sample_shape: int, seed=None) -> EpidemicEvent:
132
+ def _sample_n(self, sample_shape: int, seed=None) -> EventList:
125
133
  """
126
134
  Samples n outcomes from the continuous time state transition model.
127
135
 
@@ -132,7 +140,7 @@ class ContinuousTimeStateTransitionModel(tfd.Distribution):
132
140
  Defaults to None.
133
141
 
134
142
  Returns:
135
- EpidemicEvent: A list of n outcomes sampled from the continuous time
143
+ EventList: A list of n outcomes sampled from the continuous time
136
144
  state transition model.
137
145
  """
138
146
 
@@ -147,15 +155,12 @@ class ContinuousTimeStateTransitionModel(tfd.Distribution):
147
155
 
148
156
  return outcome
149
157
 
150
- def log_prob(self, value: EpidemicEvent) -> float:
151
- return self._log_prob(value)
152
-
153
- def _log_prob(self, value: EpidemicEvent) -> float:
158
+ def _log_prob(self, value: EventList) -> float:
154
159
  """
155
160
  Computes the log probability of the given outcomes.
156
161
 
157
162
  Args:
158
- value (EpidemicEvent): an EpidemicEvent object representing the
163
+ value (EventList): an EventList object representing the
159
164
  outcomes.
160
165
 
161
166
  Returns:
@@ -172,14 +177,30 @@ class ContinuousTimeStateTransitionModel(tfd.Distribution):
172
177
 
173
178
  return log_lik
174
179
 
175
- def _event_shape_tensor(self) -> Tensor:
176
- return tf.constant([self.num_events], dtype=tf.int32)
180
+ def _event_shape_tensor(self) -> EventList:
181
+ return EventList(
182
+ time=tf.constant([self.num_events], dtype=tf.int32),
183
+ transition=tf.constant([self.num_events], dtype=tf.int32),
184
+ individual=tf.constant([self.num_events], dtype=tf.int32),
185
+ )
177
186
 
178
- def _event_shape(self) -> tf.TensorShape:
179
- return tf.TensorShape([self.num_events])
187
+ def _event_shape(self) -> EventList:
188
+ return EventList(
189
+ time=tf.TensorShape([self.num_events]),
190
+ transition=tf.TensorShape([self.num_events]),
191
+ individual=tf.TensorShape([self.num_events]),
192
+ )
180
193
 
181
- def _batch_shape_tensor(self) -> Tensor:
182
- return tf.constant([])
194
+ def _batch_shape_tensor(self) -> EventList:
195
+ return EventList(
196
+ time=tf.constant([]),
197
+ transition=tf.constant([]),
198
+ individual=tf.constant([]),
199
+ )
183
200
 
184
- def _batch_shape(self) -> tf.TensorShape:
185
- return tf.TensorShape([])
201
+ def _batch_shape(self) -> EventList:
202
+ return EventList(
203
+ time=tf.TensorShape([]),
204
+ transition=tf.TensorShape([]),
205
+ individual=tf.TensorShape([]),
206
+ )
@@ -1,17 +1,21 @@
1
1
  """Test ContinuousTimeStateTransitionModel"""
2
2
 
3
+ from collections import namedtuple
4
+
3
5
  import numpy as np
4
6
  import pytest
5
7
  import tensorflow as tf
8
+ import tensorflow_probability as tfp
6
9
  from scipy.optimize import minimize
7
10
 
8
11
  from gemlib.distributions.continuous_time_state_transition_model import (
9
12
  ContinuousTimeStateTransitionModel,
10
- EpidemicEvent,
13
+ EventList,
11
14
  compute_state,
12
15
  )
13
16
 
14
17
  NUM_EVENTS = 1999
18
+ tfd = tfp.distributions
15
19
 
16
20
 
17
21
  @pytest.fixture
@@ -21,7 +25,7 @@ def example_ilm():
21
25
  "incidence_matrix": np.array(
22
26
  [[-1, 0], [1, -1], [0, 1]], dtype=np.float32
23
27
  ),
24
- "event_list": EpidemicEvent(
28
+ "event_list": EventList(
25
29
  time=np.array(
26
30
  [0.4, 1.3, 1.5, 1.9, 2.3, np.inf, np.inf], dtype=np.float32
27
31
  ),
@@ -58,18 +62,163 @@ def simple_sir_model():
58
62
  )
59
63
 
60
64
 
65
+ @pytest.fixture
66
+ def bayesian_sir_model():
67
+ DTYPE = np.float32
68
+
69
+ @tfd.JointDistributionCoroutine
70
+ def model():
71
+ # Priors
72
+ beta = yield tfd.Gamma(
73
+ concentration=DTYPE(0.1),
74
+ rate=DTYPE(0.1),
75
+ name="beta",
76
+ )
77
+ gamma = yield tfd.Gamma(
78
+ concentration=DTYPE(2.0), rate=DTYPE(8.0), name="gamma"
79
+ )
80
+
81
+ # Epidemic model
82
+ incidence_matrix = np.array(
83
+ [ # SI IR
84
+ [-1, 0], # S
85
+ [1, -1], # I
86
+ [0, 1], # R
87
+ ],
88
+ dtype=DTYPE,
89
+ )
90
+
91
+ initial_state = np.array([[99, 1, 0]]).astype(DTYPE)
92
+
93
+ def transition_rates(t, state):
94
+ si_rate = beta * state[:, 1] / tf.reduce_sum(state, axis=-1)
95
+ ir_rate = tf.fill((state.shape[0],), gamma)
96
+ return si_rate, ir_rate
97
+
98
+ NUM_EVENTS = 200
99
+
100
+ yield ContinuousTimeStateTransitionModel(
101
+ transition_rate_fn=transition_rates,
102
+ incidence_matrix=incidence_matrix,
103
+ initial_state=initial_state,
104
+ num_events=NUM_EVENTS,
105
+ name="sir",
106
+ )
107
+
108
+ ModelType = namedtuple("StructTuple", ["beta", "gamma", "sir"])
109
+
110
+ example = ModelType(
111
+ beta=0.5,
112
+ gamma=0.14,
113
+ sir=EventList(
114
+ time=np.array(
115
+ [ 3.8722684, 4.0912385, 4.5234103, 4.722512 , 4.9462056,
116
+ 5.7793403, 5.7914076, 6.028009 , 6.1961355, 6.8580866,
117
+ 7.6919003, 7.8862953, 8.106527 , 8.489728 , 8.565479 ,
118
+ 8.710205 , 8.7175045, 8.807663 , 8.824789 , 8.863979 ,
119
+ 9.035543 , 9.428677 , 9.470956 , 9.492353 , 9.529175 ,
120
+ 9.57752 , 9.618181 , 9.834693 , 9.88009 , 9.963752 ,
121
+ 10.158042 , 10.233447 , 10.283622 , 10.464923 , 10.472176 ,
122
+ 10.509909 , 10.713008 , 10.932794 , 10.937864 , 11.025746 ,
123
+ 11.3029 , 11.452712 , 11.45593 , 11.801975 , 11.944026 ,
124
+ 12.169918 , 12.2015085, 12.289305 , 12.379549 , 12.658843 ,
125
+ 12.673543 , 12.675543 , 12.758116 , 12.839079 , 12.887654 ,
126
+ 13.002403 , 13.063832 , 13.06984 , 13.204261 , 13.279876 ,
127
+ 13.368993 , 13.465008 , 13.545585 , 13.6250305, 13.652922 ,
128
+ 13.6884165, 13.722584 , 13.7678585, 13.773789 , 13.814655 ,
129
+ 13.861033 , 14.0029125, 14.039422 , 14.068574 , 14.082346 ,
130
+ 14.352536 , 14.35859 , 14.422744 , 14.575989 , 14.657471 ,
131
+ 14.69332 , 14.700054 , 14.848383 , 14.882924 , 15.0288925,
132
+ 15.093655 , 15.108234 , 15.236706 , 15.322691 , 15.328567 ,
133
+ 15.404174 , 15.421979 , 15.657191 , 15.922981 , 15.9314375,
134
+ 16.213312 , 16.311094 , 16.331137 , 16.50571 , 16.542715 ,
135
+ 16.634466 , 16.743319 , 16.77621 , 16.864067 , 17.02836 ,
136
+ 17.149601 , 17.375605 , 17.418259 , 17.424673 , 17.46273 ,
137
+ 17.503548 , 17.648726 , 17.659864 , 17.879429 , 18.013472 ,
138
+ 18.117981 , 18.781914 , 18.822414 , 18.923925 , 18.97823 ,
139
+ 19.034103 , 19.124441 , 19.150362 , 19.70805 , 19.848843 ,
140
+ 19.925968 , 19.967867 , 20.104042 , 20.159836 , 20.608583 ,
141
+ 21.046917 , 21.156258 , 21.233568 , 21.242342 , 21.574625 ,
142
+ 21.72764 , 21.865955 , 22.06923 , 22.242723 , 23.101812 ,
143
+ 23.451588 , 23.622063 , 23.676891 , 23.758053 , 24.117605 ,
144
+ 24.329521 , 25.265572 , 25.2762 , 25.319914 , 25.386583 ,
145
+ 25.476671 , 25.518967 , 25.865902 , 26.031075 , 26.163708 ,
146
+ 26.169014 , 26.238436 , 26.445032 , 26.973305 , 27.094973 ,
147
+ 27.237108 , 27.26869 , 27.58751 , 27.857409 , 28.816969 ,
148
+ 29.008768 , 29.6066 , 30.18358 , 31.610064 , 32.21432 ,
149
+ 32.693485 , 33.329193 , 33.356907 , 33.576965 , 33.786686 ,
150
+ 34.548702 , 35.99837 , 36.461132 , 36.63711 , 36.97347 ,
151
+ 37.104862 , 38.354027 , 38.914967 , 39.609924 , 39.66717 ,
152
+ 42.54664 , 42.610905 , 42.873867 , 43.515656 , 43.913628 ,
153
+ 45.195133 , 45.952778 , 46.862324 , 47.602505 , 50.113663 ,
154
+ 54.11174 , 56.933784 , np.inf, np.inf, np.inf,
155
+ ],
156
+ dtype=np.float32,
157
+ ),
158
+ transition=np.array(
159
+ [0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0,
160
+ 1, 0, 0, 0, 0, 1, 0, 1, 0, 0, 1, 1, 1, 0, 0, 1, 0, 1, 0, 0, 1, 1,
161
+ 1, 0, 1, 0, 0, 0, 0, 0, 0, 1, 0, 0, 0, 1, 0, 0, 0, 0, 1, 1, 1, 1,
162
+ 0, 0, 1, 0, 0, 1, 0, 0, 1, 0, 1, 0, 1, 0, 0, 0, 1, 1, 1, 0, 0, 1,
163
+ 0, 0, 1, 0, 1, 0, 1, 0, 0, 0, 0, 1, 1, 1, 0, 0, 1, 1, 1, 0, 1, 1,
164
+ 1, 0, 1, 1, 1, 0, 1, 0, 0, 1, 1, 1, 1, 0, 0, 1, 1, 0, 0, 0, 0, 1,
165
+ 0, 1, 1, 1, 1, 0, 1, 0, 1, 1, 1, 1, 1, 0, 0, 1, 1, 0, 1, 0, 1, 0,
166
+ 0, 1, 1, 0, 1, 0, 1, 1, 1, 1, 1, 0, 1, 1, 0, 1, 0, 1, 1, 0, 1, 1,
167
+ 1, 0, 1, 1, 1, 1, 1, 1, 0, 0, 1, 1, 1, 0, 1, 1, 1, 1, 1, 1, 1, 2,
168
+ 2, 2,
169
+ ],
170
+ dtype=np.int32,
171
+ ),
172
+ individual=np.array(
173
+ [0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
174
+ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
175
+ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
176
+ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
177
+ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
178
+ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
179
+ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
180
+ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
181
+ 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0,
182
+ 0, 0,
183
+ ],
184
+ dtype=np.int32,
185
+ ),
186
+ )
187
+ ) # fmt: skip
188
+
189
+ return {"model": model, "example": example}
190
+
191
+
61
192
  def test_simple_sir_shapes(simple_sir_model):
62
193
  """Test expected output shape"""
63
194
  tf.debugging.assert_equal(
64
- simple_sir_model.event_shape_tensor(), tf.constant(NUM_EVENTS)
195
+ simple_sir_model.event_shape_tensor(),
196
+ EventList(
197
+ tf.constant(NUM_EVENTS),
198
+ tf.constant(NUM_EVENTS),
199
+ tf.constant(NUM_EVENTS),
200
+ ),
65
201
  )
66
202
  tf.debugging.assert_equal(
67
- simple_sir_model.event_shape, tf.TensorShape([NUM_EVENTS])
203
+ simple_sir_model.event_shape,
204
+ EventList(
205
+ tf.TensorShape([NUM_EVENTS]),
206
+ tf.TensorShape([NUM_EVENTS]),
207
+ tf.TensorShape([NUM_EVENTS]),
208
+ ),
68
209
  )
69
210
  tf.debugging.assert_equal(
70
- simple_sir_model.batch_shape_tensor(), tf.constant([], tf.int32)
211
+ simple_sir_model.batch_shape_tensor(),
212
+ EventList(
213
+ tf.constant([], tf.int32),
214
+ tf.constant([], tf.int32),
215
+ tf.constant([], tf.int32),
216
+ ),
217
+ )
218
+ tf.debugging.assert_equal(
219
+ simple_sir_model.batch_shape,
220
+ EventList(tf.TensorShape([]), tf.TensorShape([]), tf.TensorShape([])),
71
221
  )
72
- tf.debugging.assert_equal(simple_sir_model.batch_shape, tf.TensorShape([]))
73
222
 
74
223
 
75
224
  def test_simple_sir_eager(simple_sir_model):
@@ -77,7 +226,7 @@ def test_simple_sir_eager(simple_sir_model):
77
226
 
78
227
  sample = simple_sir_model.sample(seed=[0, 0])
79
228
 
80
- assert isinstance(sample, EpidemicEvent)
229
+ assert isinstance(sample, EventList)
81
230
 
82
231
  state = simple_sir_model.compute_state(sample)
83
232
  tf.debugging.assert_non_negative(state)
@@ -92,7 +241,7 @@ def test_simple_sir_graph(simple_sir_model):
92
241
 
93
242
  sample = fn()
94
243
 
95
- assert isinstance(sample, EpidemicEvent)
244
+ assert isinstance(sample, EventList)
96
245
 
97
246
  state = simple_sir_model.compute_state(sample)
98
247
  tf.debugging.assert_non_negative(state)
@@ -287,3 +436,15 @@ def test_simple_sir_workflow(simple_sir_model):
287
436
 
288
437
  assert opt.success
289
438
  assert np.all((lower_ci < actuals) & (actuals < upper_ci))
439
+
440
+
441
+ def test_tfp_jd_integration(bayesian_sir_model):
442
+ model = bayesian_sir_model["model"]
443
+ example = bayesian_sir_model["example"]
444
+
445
+ model.sample(seed=[20240714, 1139])
446
+
447
+ conditioned_model = model.experimental_pin(sir=example.sir)
448
+ lp = conditioned_model.log_prob(beta=0.5, gamma=0.14)
449
+
450
+ np.testing.assert_approx_equal(lp, 10.0128927)
@@ -35,7 +35,7 @@ class DiscreteTimeStateTransitionModel(tfd.Distribution):
35
35
 
36
36
  def __init__(
37
37
  self,
38
- transition_rates,
38
+ transition_rate_fn,
39
39
  incidence_matrix,
40
40
  initial_state,
41
41
  initial_step,
@@ -49,7 +49,7 @@ class DiscreteTimeStateTransitionModel(tfd.Distribution):
49
49
 
50
50
  Args:
51
51
  ----
52
- transition_rates: Python callable of the form `fn(t, state)` taking
52
+ transition_rate_fn: Python callable of the form `fn(t, state)` taking
53
53
  the current time `t` (Python float) and state tensor `state`. This
54
54
  function returns a tensor which broadcasts to the first dimension of
55
55
  `incidence_matrix`.
@@ -117,7 +117,7 @@ class DiscreteTimeStateTransitionModel(tfd.Distribution):
117
117
 
118
118
  # Instantiate model
119
119
  sir = DiscreteTimeStateTransitionModel(
120
- transition_rates=txrates,
120
+ transition_rate_fn=txrates,
121
121
  incidence_matrix=incidence_matrix,
122
122
  initial_state=initial_state,
123
123
  initial_step=initial_step,
@@ -161,7 +161,7 @@ class DiscreteTimeStateTransitionModel(tfd.Distribution):
161
161
  """
162
162
  parameters = dict(locals())
163
163
  with tf.name_scope(name) as name:
164
- self._transition_rates = transition_rates
164
+ self._transition_rate_fn = transition_rate_fn
165
165
  self._incidence_matrix = tf.convert_to_tensor(
166
166
  incidence_matrix, dtype=initial_state.dtype
167
167
  )
@@ -183,8 +183,8 @@ class DiscreteTimeStateTransitionModel(tfd.Distribution):
183
183
  self.dtype = initial_state.dtype
184
184
 
185
185
  @property
186
- def transition_rates(self):
187
- return self._transition_rates
186
+ def transition_rate_fn(self):
187
+ return self._transition_rate_fn
188
188
 
189
189
  @property
190
190
  def incidence_matrix(self):
@@ -261,7 +261,7 @@ class DiscreteTimeStateTransitionModel(tfd.Distribution):
261
261
  seed, salt="DiscreteTimeStateTransitionModel"
262
262
  )
263
263
  t, sim = discrete_markov_simulation(
264
- hazard_fn=self.transition_rates,
264
+ hazard_fn=self.transition_rate_fn,
265
265
  state=self.initial_state,
266
266
  start=self.initial_step,
267
267
  end=self.initial_step + self.num_steps * self.time_delta,
@@ -288,7 +288,7 @@ class DiscreteTimeStateTransitionModel(tfd.Distribution):
288
288
  )
289
289
  y = tf.convert_to_tensor(y, dtype)
290
290
 
291
- hazard = self.transition_rates
291
+ hazard = self.transition_rate_fn
292
292
  return discrete_markov_log_prob(
293
293
  events=y,
294
294
  init_state=self.initial_state,
@@ -56,7 +56,7 @@ def txrates(t, state):
56
56
 
57
57
  # Instantiate model
58
58
  sir = DiscreteTimeStateTransitionModel(
59
- transition_rates=txrates,
59
+ transition_rate_fn=txrates,
60
60
  stoichiometry=stoichiometry,
61
61
  initial_state=initial_state,
62
62
  initial_step=initial_step,
@@ -166,7 +166,7 @@ def txrates(t, state):
166
166
 
167
167
  # Instantiate model
168
168
  sirs = DiscreteTimeStateTransitionModel(
169
- transition_rates=txrates,
169
+ transition_rate_fn=txrates,
170
170
  stoichiometry=stoichiometry,
171
171
  initial_state=initial_state,
172
172
  initial_step=initial_step,
@@ -252,7 +252,7 @@ def txrates(t, state):
252
252
 
253
253
  # Instantiate model
254
254
  seir = DiscreteTimeStateTransitionModel(
255
- transition_rates=txrates,
255
+ transition_rate_fn=txrates,
256
256
  stoichiometry=stoichiometry,
257
257
  initial_state=initial_state,
258
258
  initial_step=initial_step,
@@ -336,7 +336,7 @@ def txrates(t, state):
336
336
  initial_step, time_delta, num_steps = 0.0, 1.0, 100
337
337
 
338
338
  sirc = DiscreteTimeStateTransitionModel(
339
- transition_rates=txrates,
339
+ transition_rate_fn=txrates,
340
340
  stoichiometry=stoichiometry,
341
341
  initial_state=initial_state,
342
342
  initial_step=initial_step,
@@ -409,7 +409,7 @@ def txrates(t, state):
409
409
 
410
410
  # Instantiate model
411
411
  sivr = DiscreteTimeStateTransitionModel(
412
- transition_rates=txrates,
412
+ transition_rate_fn=txrates,
413
413
  stoichiometry=stoichiometry,
414
414
  initial_state=initial_state,
415
415
  initial_step=initial_step,
@@ -38,7 +38,7 @@ class TestDiscreteTimeStateTransitionModel(test_util.TestCase):
38
38
  return [si, ir]
39
39
 
40
40
  return DiscreteTimeStateTransitionModel(
41
- transition_rates=txrates,
41
+ transition_rate_fn=txrates,
42
42
  incidence_matrix=incidence_matrix,
43
43
  initial_state=initial_state,
44
44
  initial_step=initial_step,
@@ -296,7 +296,7 @@ class TestDiscreteTimeStateTransitionModelLogProbMaxima(test_util.TestCase):
296
296
 
297
297
  beta, gamma = tf.unstack(params)
298
298
  return DiscreteTimeStateTransitionModel(
299
- transition_rates=txrates,
299
+ transition_rate_fn=txrates,
300
300
  incidence_matrix=incidence_matrix,
301
301
  initial_state=initial_state,
302
302
  initial_step=0.0,
@@ -209,7 +209,7 @@ def default_initial_ripple(model, current_events, current_state, seed):
209
209
 
210
210
  # Choose new infection events at time t
211
211
  proposed_transition_rates = tf.stack(
212
- model.transition_rates(
212
+ model.transition_rate_fn(
213
213
  proposed_time_idx, tf.transpose(current_state_t)
214
214
  ),
215
215
  axis=0,
@@ -267,7 +267,7 @@ def chain_binomial_rippler(model, current_events, initial_ripple_fn, seed=None):
267
267
  # Calculate transition rates for current and new states
268
268
  def transition_probs(time, state):
269
269
  rates = tf.stack(
270
- model.transition_rates(time, tf.transpose(state)), axis=-2
270
+ model.transition_rate_fn(time, tf.transpose(state)), axis=-2
271
271
  )
272
272
  return 1.0 - tf.math.exp(-rates * model.time_delta)
273
273
 
@@ -33,7 +33,7 @@ def _make_model():
33
33
  return [si, ir]
34
34
 
35
35
  model = DiscreteTimeStateTransitionModel(
36
- transition_rates=hazard_fn,
36
+ transition_rate_fn=hazard_fn,
37
37
  initial_state=init_state,
38
38
  initial_step=0,
39
39
  time_delta=1.0,
@@ -87,7 +87,7 @@ class CBRSIRTest(test_util.TestCase):
87
87
  return [si, ir]
88
88
 
89
89
  model = DiscreteTimeStateTransitionModel(
90
- transition_rates=hazard_fn,
90
+ transition_rate_fn=hazard_fn,
91
91
  initial_state=init_state,
92
92
  initial_step=0,
93
93
  time_delta=1.0,
@@ -1,10 +1,14 @@
1
1
  [tool.poetry]
2
2
  name = "gemlib"
3
- version = "0.9.2"
3
+ version = "0.9.4"
4
4
  description = "GEMlib scientific compute library for epidemic modelling"
5
5
  authors = ["Chris Jewell <c.jewell@lancaster.ac.uk>",
6
6
  "Alison Hale <haleac@lancaster.ac.uk>"]
7
- repository = "http://fhm-chicas-code.lancs.ac.uk/GEM/gemlib"
7
+ maintainers = ["Jessica Bridgen <j.bridgen@lancaster.ac.uk>",
8
+ "Alin Morariu <a.morariu@lancaster.ac.uk"]
9
+ repository = "https://gitlab.com/gem-epidemics/gemlib"
10
+ homepage = "https://gem-epidemics.gitlab.io/gemlib"
11
+ keywords = ["epidemic", "Bayesian", "inference", "infectious disease model", "probabilistic programming"]
8
12
 
9
13
  [tool.poetry.dependencies]
10
14
  python = ">=3.9.0, <3.12.0"
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes