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.
- {gemlib-0.9.2 → gemlib-0.9.4}/PKG-INFO +6 -3
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/continuous_markov.py +15 -14
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/continuous_time_state_transition_model.py +41 -20
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/continuous_time_state_transition_model_test.py +169 -8
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_time_state_transition_model.py +8 -8
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_time_state_transition_model_examples.py +5 -5
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_time_state_transition_model_test.py +2 -2
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/chain_binomial_rippler.py +2 -2
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/chain_binomial_rippler_test.py +2 -2
- {gemlib-0.9.2 → gemlib-0.9.4}/pyproject.toml +6 -2
- {gemlib-0.9.2 → gemlib-0.9.4}/LICENSE +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/__init__.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/__init__.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/brownian.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/categorical2.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_markov.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_rejection_sampling.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/experimental/__init__.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/experimental/discrete_approx_cont_state_transition_model.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/experimental/state_transition_marginal_model.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/hypergeometric.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/hypergeometric_sampler.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/hypergeometric_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/kcategorical.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/kcategorical_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/uniform_integer.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/uniform_integer_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/__init__.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/adaptive_random_walk_metropolis.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/adaptive_random_walk_metropolis_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/bb_fixture.pkl +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/brownian_bridge_kernel.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/brownian_bridge_kernel_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/compound_kernel.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/conftest.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/damped_chain_binomial_rippler.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/__init__.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/event_time_proposal.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/fixtures.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/move_events.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/move_events_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_proposal.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/util.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/__init__.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/composable_kernel.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/__init__.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/left_censored_events_mh.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/hmc.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/hmc_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/mcmc_base.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/mcmc_sampler.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/mcmc_sampler_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/multi_scan.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/multi_scan_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/random_walk_metropolis.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/random_walk_metropolis_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/experimental/test_util.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/gibbs_kernel.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/gibbs_kernel_test.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/h5_posterior.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/multi_scan_kernel.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/zarr_posterior.py +0 -0
- {gemlib-0.9.2 → gemlib-0.9.4}/gemlib/util.py +0 -0
- {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.
|
|
3
|
+
Version: 0.9.4
|
|
4
4
|
Summary: GEMlib scientific compute library for epidemic modelling
|
|
5
|
-
Home-page:
|
|
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,
|
|
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
|
|
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:
|
|
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
|
-
) ->
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
) ->
|
|
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
|
-
|
|
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 =
|
|
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 =
|
|
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 =
|
|
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 =
|
|
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:
|
|
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 (
|
|
317
|
+
event (EventList): The event data containing the times
|
|
317
318
|
and states.
|
|
318
319
|
|
|
319
320
|
Returns:
|
{gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/continuous_time_state_transition_model.py
RENAMED
|
@@ -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
|
-
|
|
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(
|
|
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=
|
|
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:
|
|
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) ->
|
|
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) ->
|
|
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
|
-
|
|
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
|
|
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 (
|
|
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) ->
|
|
176
|
-
return
|
|
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) ->
|
|
179
|
-
return
|
|
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) ->
|
|
182
|
-
return
|
|
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) ->
|
|
185
|
-
return
|
|
201
|
+
def _batch_shape(self) -> EventList:
|
|
202
|
+
return EventList(
|
|
203
|
+
time=tf.TensorShape([]),
|
|
204
|
+
transition=tf.TensorShape([]),
|
|
205
|
+
individual=tf.TensorShape([]),
|
|
206
|
+
)
|
{gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/continuous_time_state_transition_model_test.py
RENAMED
|
@@ -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
|
-
|
|
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":
|
|
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(),
|
|
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,
|
|
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(),
|
|
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,
|
|
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,
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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.
|
|
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
|
|
187
|
-
return self.
|
|
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.
|
|
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.
|
|
291
|
+
hazard = self.transition_rate_fn
|
|
292
292
|
return discrete_markov_log_prob(
|
|
293
293
|
events=y,
|
|
294
294
|
init_state=self.initial_state,
|
{gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_time_state_transition_model_examples.py
RENAMED
|
@@ -56,7 +56,7 @@ def txrates(t, state):
|
|
|
56
56
|
|
|
57
57
|
# Instantiate model
|
|
58
58
|
sir = DiscreteTimeStateTransitionModel(
|
|
59
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
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
|
-
|
|
412
|
+
transition_rate_fn=txrates,
|
|
413
413
|
stoichiometry=stoichiometry,
|
|
414
414
|
initial_state=initial_state,
|
|
415
415
|
initial_step=initial_step,
|
{gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/discrete_time_state_transition_model_test.py
RENAMED
|
@@ -38,7 +38,7 @@ class TestDiscreteTimeStateTransitionModel(test_util.TestCase):
|
|
|
38
38
|
return [si, ir]
|
|
39
39
|
|
|
40
40
|
return DiscreteTimeStateTransitionModel(
|
|
41
|
-
|
|
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
|
-
|
|
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.
|
|
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.
|
|
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
|
-
|
|
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
|
-
|
|
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.
|
|
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
|
-
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{gemlib-0.9.2 → gemlib-0.9.4}/gemlib/distributions/experimental/state_transition_marginal_model.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/move_events.py
RENAMED
|
File without changes
|
{gemlib-0.9.2 → gemlib-0.9.4}/gemlib/mcmc/discrete_time_state_transition_model/move_events_test.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|