gemlib 0.9.2__py3-none-any.whl
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/__init__.py +9 -0
- gemlib/distributions/__init__.py +26 -0
- gemlib/distributions/brownian.py +141 -0
- gemlib/distributions/categorical2.py +33 -0
- gemlib/distributions/continuous_markov.py +371 -0
- gemlib/distributions/continuous_time_state_transition_model.py +185 -0
- gemlib/distributions/continuous_time_state_transition_model_test.py +289 -0
- gemlib/distributions/discrete_markov.py +279 -0
- gemlib/distributions/discrete_rejection_sampling.py +149 -0
- gemlib/distributions/discrete_time_state_transition_model.py +324 -0
- gemlib/distributions/discrete_time_state_transition_model_examples.py +453 -0
- gemlib/distributions/discrete_time_state_transition_model_test.py +336 -0
- gemlib/distributions/experimental/__init__.py +7 -0
- gemlib/distributions/experimental/discrete_approx_cont_state_transition_model.py +194 -0
- gemlib/distributions/experimental/state_transition_marginal_model.py +352 -0
- gemlib/distributions/hypergeometric.py +144 -0
- gemlib/distributions/hypergeometric_sampler.py +103 -0
- gemlib/distributions/hypergeometric_test.py +46 -0
- gemlib/distributions/kcategorical.py +113 -0
- gemlib/distributions/kcategorical_test.py +52 -0
- gemlib/distributions/uniform_integer.py +169 -0
- gemlib/distributions/uniform_integer_test.py +55 -0
- gemlib/mcmc/__init__.py +23 -0
- gemlib/mcmc/adaptive_random_walk_metropolis.py +859 -0
- gemlib/mcmc/adaptive_random_walk_metropolis_test.py +144 -0
- gemlib/mcmc/bb_fixture.pkl +0 -0
- gemlib/mcmc/brownian_bridge_kernel.py +291 -0
- gemlib/mcmc/brownian_bridge_kernel_test.py +164 -0
- gemlib/mcmc/chain_binomial_rippler.py +524 -0
- gemlib/mcmc/chain_binomial_rippler_test.py +120 -0
- gemlib/mcmc/compound_kernel.py +156 -0
- gemlib/mcmc/conftest.py +4 -0
- gemlib/mcmc/damped_chain_binomial_rippler.py +840 -0
- gemlib/mcmc/discrete_time_state_transition_model/__init__.py +21 -0
- gemlib/mcmc/discrete_time_state_transition_model/event_time_proposal.py +275 -0
- gemlib/mcmc/discrete_time_state_transition_model/fixtures.py +56 -0
- gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh.py +258 -0
- gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh_test.py +108 -0
- gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal.py +188 -0
- gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal_test.py +33 -0
- gemlib/mcmc/discrete_time_state_transition_model/move_events.py +239 -0
- gemlib/mcmc/discrete_time_state_transition_model/move_events_test.py +63 -0
- gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh.py +254 -0
- gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh_test.py +28 -0
- gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_proposal.py +150 -0
- gemlib/mcmc/discrete_time_state_transition_model/util.py +9 -0
- gemlib/mcmc/experimental/__init__.py +0 -0
- gemlib/mcmc/experimental/composable_kernel.py +270 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/__init__.py +1 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/left_censored_events_mh.py +149 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events.py +71 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events_test.py +48 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh.py +89 -0
- gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh_test.py +28 -0
- gemlib/mcmc/experimental/hmc.py +111 -0
- gemlib/mcmc/experimental/hmc_test.py +61 -0
- gemlib/mcmc/experimental/mcmc_base.py +33 -0
- gemlib/mcmc/experimental/mcmc_sampler.py +94 -0
- gemlib/mcmc/experimental/mcmc_sampler_test.py +44 -0
- gemlib/mcmc/experimental/multi_scan.py +53 -0
- gemlib/mcmc/experimental/multi_scan_test.py +71 -0
- gemlib/mcmc/experimental/random_walk_metropolis.py +107 -0
- gemlib/mcmc/experimental/random_walk_metropolis_test.py +214 -0
- gemlib/mcmc/experimental/test_util.py +49 -0
- gemlib/mcmc/gibbs_kernel.py +505 -0
- gemlib/mcmc/gibbs_kernel_test.py +212 -0
- gemlib/mcmc/h5_posterior.py +77 -0
- gemlib/mcmc/multi_scan_kernel.py +59 -0
- gemlib/mcmc/zarr_posterior.py +132 -0
- gemlib/util.py +117 -0
- gemlib/util_test.py +75 -0
- gemlib-0.9.2.dist-info/LICENSE +21 -0
- gemlib-0.9.2.dist-info/METADATA +19 -0
- gemlib-0.9.2.dist-info/RECORD +75 -0
- gemlib-0.9.2.dist-info/WHEEL +4 -0
|
@@ -0,0 +1,453 @@
|
|
|
1
|
+
"""DiscreteTimeStateTransitionModel examples"""
|
|
2
|
+
|
|
3
|
+
import matplotlib.pyplot as plt
|
|
4
|
+
import tensorflow as tf
|
|
5
|
+
from matplotlib.lines import Line2D
|
|
6
|
+
|
|
7
|
+
from gemlib.distributions.discrete_time_state_transition_model import (
|
|
8
|
+
DiscreteTimeStateTransitionModel,
|
|
9
|
+
)
|
|
10
|
+
from gemlib.util import compute_state
|
|
11
|
+
|
|
12
|
+
# Example 1: SIR model for a single population
|
|
13
|
+
|
|
14
|
+
dtype = tf.float32
|
|
15
|
+
|
|
16
|
+
# Initial state, counts per compartment (S, I, R), for one population
|
|
17
|
+
initial_state = tf.constant([[99, 1, 0]], dtype)
|
|
18
|
+
|
|
19
|
+
# Stoichiometry matrix
|
|
20
|
+
stoichiometry = tf.constant(
|
|
21
|
+
[ # S I R
|
|
22
|
+
[-1, 1, 0], # S->I
|
|
23
|
+
[0, -1, 1], # I->R
|
|
24
|
+
],
|
|
25
|
+
dtype,
|
|
26
|
+
)
|
|
27
|
+
|
|
28
|
+
# time parameters
|
|
29
|
+
initial_step, time_delta, num_steps = 0.0, 1.0, 100
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def txrates(t, state):
|
|
33
|
+
"""Transition rate per individual corresponding to each row of the
|
|
34
|
+
stoichiometry matrix.
|
|
35
|
+
|
|
36
|
+
Args:
|
|
37
|
+
----
|
|
38
|
+
state: `Tensor` representing the current state (count of individuals in
|
|
39
|
+
each compartment).
|
|
40
|
+
t: Python float representing the current time. For example seasonality
|
|
41
|
+
in the S->I transition could be driven by tensors of the following
|
|
42
|
+
form:
|
|
43
|
+
seasonality = tf.math.sin(2 * 3.14159 * t / 20) + 1
|
|
44
|
+
si = seasonality * beta * state[:, 1] / tf.reduce_sum(state)
|
|
45
|
+
|
|
46
|
+
Returns:
|
|
47
|
+
-------
|
|
48
|
+
List of `Tensor`(s) each of which corresponds to a transition.
|
|
49
|
+
|
|
50
|
+
"""
|
|
51
|
+
beta, gamma = 0.28, 0.14 # note R0=beta/gamma
|
|
52
|
+
si = beta * state[:, 1] / tf.reduce_sum(state) # S->I transition rate
|
|
53
|
+
ir = tf.constant([gamma], dtype) # I->R transition rate
|
|
54
|
+
return [si, ir]
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
# Instantiate model
|
|
58
|
+
sir = DiscreteTimeStateTransitionModel(
|
|
59
|
+
transition_rates=txrates,
|
|
60
|
+
stoichiometry=stoichiometry,
|
|
61
|
+
initial_state=initial_state,
|
|
62
|
+
initial_step=initial_step,
|
|
63
|
+
time_delta=time_delta,
|
|
64
|
+
num_steps=num_steps,
|
|
65
|
+
)
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
@tf.function
|
|
69
|
+
def simulate_one(elems):
|
|
70
|
+
"""One realisation of the epidemic process."""
|
|
71
|
+
return sir.sample()
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
nsim = 15 # Number of realisations of the epidemic process
|
|
75
|
+
eventlist = tf.map_fn(
|
|
76
|
+
simulate_one,
|
|
77
|
+
tf.ones([nsim, stoichiometry.shape[0]]),
|
|
78
|
+
fn_output_signature=dtype,
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
# Events for each transition with shape (simulation, population, time,
|
|
82
|
+
# transition)
|
|
83
|
+
print("I->R events:", eventlist[0, 0, :, 1])
|
|
84
|
+
|
|
85
|
+
# Log prob of observing the eventlist, of first simulation, given the model
|
|
86
|
+
print("Log prob:", sir.log_prob(eventlist[0, ...]))
|
|
87
|
+
print(
|
|
88
|
+
"log prob of each simulation:",
|
|
89
|
+
tf.vectorized_map(
|
|
90
|
+
fn=lambda i: sir.log_prob(eventlist[i, ...]), elems=tf.range(nsim)
|
|
91
|
+
),
|
|
92
|
+
)
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
# Plot timeseries of counts in each compartment from eventlist
|
|
96
|
+
def plot_timeseries(
|
|
97
|
+
initial_state,
|
|
98
|
+
eventlist,
|
|
99
|
+
stoichiometry,
|
|
100
|
+
initial_step,
|
|
101
|
+
time_delta,
|
|
102
|
+
num_steps,
|
|
103
|
+
nsim,
|
|
104
|
+
col,
|
|
105
|
+
legend,
|
|
106
|
+
):
|
|
107
|
+
"""Plot timeseries of counts in each compartment."""
|
|
108
|
+
state_timeseries = compute_state(initial_state, eventlist, stoichiometry)
|
|
109
|
+
x = tf.range(
|
|
110
|
+
initial_step, initial_step + time_delta * num_steps, delta=time_delta
|
|
111
|
+
)
|
|
112
|
+
for i in range(0, initial_state.shape[0]): # populations
|
|
113
|
+
plt.subplots_adjust(hspace=0.8)
|
|
114
|
+
plt.subplot(initial_state.shape[0], 1, i + 1)
|
|
115
|
+
for j in range(0, nsim): # simulations
|
|
116
|
+
for k in range(0, initial_state.shape[1]): # compartments
|
|
117
|
+
plt.step(x, state_timeseries[j, i, :, k], col[k], lw=0.5)
|
|
118
|
+
lines = [Line2D([0], [0], color=c, lw=0.5) for c in col]
|
|
119
|
+
plt.legend(lines, legend)
|
|
120
|
+
plt.title(str(nsim) + " simulations of population " + str(i))
|
|
121
|
+
plt.xlabel("time")
|
|
122
|
+
plt.ylabel("count")
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
plot_timeseries(
|
|
126
|
+
initial_state,
|
|
127
|
+
eventlist,
|
|
128
|
+
stoichiometry,
|
|
129
|
+
initial_step,
|
|
130
|
+
time_delta,
|
|
131
|
+
num_steps,
|
|
132
|
+
nsim,
|
|
133
|
+
["k", "r", "b"],
|
|
134
|
+
["S", "I", "R"],
|
|
135
|
+
)
|
|
136
|
+
plt.show()
|
|
137
|
+
|
|
138
|
+
# Example 2: SIR feedback model where recovered become susceptible
|
|
139
|
+
|
|
140
|
+
dtype = tf.float32
|
|
141
|
+
|
|
142
|
+
# Initial state, counts per compartment (S, I, R), for one population
|
|
143
|
+
initial_state = tf.constant([[999, 1, 0]], dtype)
|
|
144
|
+
|
|
145
|
+
# Stoichiometry matrix # S, I, R
|
|
146
|
+
stoichiometry = tf.constant(
|
|
147
|
+
[
|
|
148
|
+
[-1, 1, 0], # S->I
|
|
149
|
+
[0, -1, 1], # I->R
|
|
150
|
+
[1, 0, -1],
|
|
151
|
+
], # R->S
|
|
152
|
+
dtype,
|
|
153
|
+
)
|
|
154
|
+
# time parameters
|
|
155
|
+
initial_step, time_delta, num_steps = 0.0, 1.0, 200
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
def txrates(t, state):
|
|
159
|
+
"""Transition rate per individual with feedback from R to S."""
|
|
160
|
+
beta, gamma, eta = 0.6, 0.4, 0.05
|
|
161
|
+
si = beta * state[:, 1] / tf.reduce_sum(state) # S->I transition rate
|
|
162
|
+
ir = tf.constant([gamma], dtype) # I->R transition rate
|
|
163
|
+
rs = tf.constant([eta], dtype) # R->S transition rate
|
|
164
|
+
return [si, ir, rs]
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
# Instantiate model
|
|
168
|
+
sirs = DiscreteTimeStateTransitionModel(
|
|
169
|
+
transition_rates=txrates,
|
|
170
|
+
stoichiometry=stoichiometry,
|
|
171
|
+
initial_state=initial_state,
|
|
172
|
+
initial_step=initial_step,
|
|
173
|
+
time_delta=time_delta,
|
|
174
|
+
num_steps=num_steps,
|
|
175
|
+
)
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
@tf.function
|
|
179
|
+
def simulate_one(elems):
|
|
180
|
+
"""One realisation of the epidemic process."""
|
|
181
|
+
return sirs.sample()
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
nsim = 12
|
|
185
|
+
eventlist = tf.map_fn(
|
|
186
|
+
simulate_one,
|
|
187
|
+
tf.ones([nsim, stoichiometry.shape[0]]),
|
|
188
|
+
fn_output_signature=dtype,
|
|
189
|
+
)
|
|
190
|
+
|
|
191
|
+
# Second simulation eventlist for transition I->R
|
|
192
|
+
print("I->R events:", eventlist[1, 0, :, 1])
|
|
193
|
+
|
|
194
|
+
# Log prob of observing the eventlist, of second simulation, given the model
|
|
195
|
+
print("Log prob:", sirs.log_prob(eventlist[1, ...]))
|
|
196
|
+
|
|
197
|
+
# Plot timeseries using plot_timeseries() from first example
|
|
198
|
+
plot_timeseries(
|
|
199
|
+
initial_state,
|
|
200
|
+
eventlist,
|
|
201
|
+
stoichiometry,
|
|
202
|
+
initial_step,
|
|
203
|
+
time_delta,
|
|
204
|
+
num_steps,
|
|
205
|
+
nsim,
|
|
206
|
+
["k", "r", "b"],
|
|
207
|
+
["S", "I", "R"],
|
|
208
|
+
)
|
|
209
|
+
plt.show()
|
|
210
|
+
|
|
211
|
+
# Example 3: SEIR model for two populations
|
|
212
|
+
|
|
213
|
+
dtype = tf.float32
|
|
214
|
+
|
|
215
|
+
# Initial state, counts per compartment (S, E, I, R), for two populations
|
|
216
|
+
initial_state = tf.constant(
|
|
217
|
+
[
|
|
218
|
+
[450, 5, 5, 0], # first population
|
|
219
|
+
[990, 10, 0, 0],
|
|
220
|
+
], # second population
|
|
221
|
+
dtype,
|
|
222
|
+
)
|
|
223
|
+
|
|
224
|
+
# Stoichiometry matrix # S, E, I, R
|
|
225
|
+
stoichiometry = tf.constant(
|
|
226
|
+
[
|
|
227
|
+
[-1, 1, 0, 0], # S->E
|
|
228
|
+
[0, -1, 1, 0], # E->I
|
|
229
|
+
[0, 0, -1, 1],
|
|
230
|
+
], # I->R
|
|
231
|
+
dtype,
|
|
232
|
+
)
|
|
233
|
+
# time parameters
|
|
234
|
+
initial_step, time_delta, num_steps = 0.0, 1.0, 150
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def txrates(t, state):
|
|
238
|
+
"""Transition rate per individual for 2 populations [population_0,
|
|
239
|
+
population_1].
|
|
240
|
+
"""
|
|
241
|
+
beta, delta, gamma = 0.25, 0.15, 0.05
|
|
242
|
+
se = tf.stack(
|
|
243
|
+
[
|
|
244
|
+
beta * state[0, 1] / tf.reduce_sum(state[0, :]),
|
|
245
|
+
beta * state[1, 1] / tf.reduce_sum(state[1, :]),
|
|
246
|
+
]
|
|
247
|
+
) # S->E transition rate
|
|
248
|
+
ei = tf.constant([delta, delta], dtype) # E->I transition rate
|
|
249
|
+
ir = tf.constant([gamma, gamma], dtype) # I->R transition rate
|
|
250
|
+
return [se, ei, ir]
|
|
251
|
+
|
|
252
|
+
|
|
253
|
+
# Instantiate model
|
|
254
|
+
seir = DiscreteTimeStateTransitionModel(
|
|
255
|
+
transition_rates=txrates,
|
|
256
|
+
stoichiometry=stoichiometry,
|
|
257
|
+
initial_state=initial_state,
|
|
258
|
+
initial_step=initial_step,
|
|
259
|
+
time_delta=time_delta,
|
|
260
|
+
num_steps=num_steps,
|
|
261
|
+
)
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
@tf.function
|
|
265
|
+
def simulate_one(elems):
|
|
266
|
+
"""One realisation of the epidemic process."""
|
|
267
|
+
return seir.sample()
|
|
268
|
+
|
|
269
|
+
|
|
270
|
+
nsim = 10
|
|
271
|
+
eventlist = tf.map_fn(
|
|
272
|
+
simulate_one,
|
|
273
|
+
tf.ones([nsim, stoichiometry.shape[0]]),
|
|
274
|
+
fn_output_signature=dtype,
|
|
275
|
+
)
|
|
276
|
+
|
|
277
|
+
# Sixth simulation, second population, transition E->I
|
|
278
|
+
print("E->I events:", eventlist[5, 1, :, 1].numpy())
|
|
279
|
+
|
|
280
|
+
# Log prob of observing the eventlist, of sixth simulation, given the model
|
|
281
|
+
print("Log prob:", seir.log_prob(eventlist[5, ...]))
|
|
282
|
+
|
|
283
|
+
# Plot timeseries using plot_timeseries() from first example
|
|
284
|
+
plot_timeseries(
|
|
285
|
+
initial_state,
|
|
286
|
+
eventlist,
|
|
287
|
+
stoichiometry,
|
|
288
|
+
initial_step,
|
|
289
|
+
time_delta,
|
|
290
|
+
num_steps,
|
|
291
|
+
nsim,
|
|
292
|
+
["k", "g", "r", "b"],
|
|
293
|
+
["S", "E", "I", "R"],
|
|
294
|
+
)
|
|
295
|
+
plt.show()
|
|
296
|
+
|
|
297
|
+
|
|
298
|
+
# Example 4: SIR model with mixing between the two populations
|
|
299
|
+
|
|
300
|
+
dtype = tf.float32
|
|
301
|
+
|
|
302
|
+
# Initial state (counts). Infections in population 1 will be driven by
|
|
303
|
+
# population 0
|
|
304
|
+
initial_state = tf.constant(
|
|
305
|
+
[
|
|
306
|
+
[700, 300, 0], # population 0
|
|
307
|
+
[1000, 0, 0],
|
|
308
|
+
], # population 1
|
|
309
|
+
dtype,
|
|
310
|
+
)
|
|
311
|
+
|
|
312
|
+
# Stoichiometry matrix # S, I, R
|
|
313
|
+
stoichiometry = tf.constant(
|
|
314
|
+
[
|
|
315
|
+
[-1, 1, 0], # S->I
|
|
316
|
+
[0, -1, 1],
|
|
317
|
+
], # I->R
|
|
318
|
+
dtype,
|
|
319
|
+
)
|
|
320
|
+
|
|
321
|
+
# Mixing matrix describes 25% mixing between 2 populations
|
|
322
|
+
C = tf.constant([[0.75, 0.25], [0.25, 0.75]], dtype)
|
|
323
|
+
|
|
324
|
+
|
|
325
|
+
def txrates(t, state):
|
|
326
|
+
"""Transition rate per individual with mixing between populations."""
|
|
327
|
+
beta, gamma = 0.4, 0.1
|
|
328
|
+
si = (
|
|
329
|
+
beta * tf.linalg.matvec(C, state[:, 1]) / tf.reduce_sum(state)
|
|
330
|
+
) # S->I transition
|
|
331
|
+
ir = tf.constant([gamma, gamma], dtype) # I->R transition rate
|
|
332
|
+
return [si, ir]
|
|
333
|
+
|
|
334
|
+
|
|
335
|
+
# time parameters
|
|
336
|
+
initial_step, time_delta, num_steps = 0.0, 1.0, 100
|
|
337
|
+
|
|
338
|
+
sirc = DiscreteTimeStateTransitionModel(
|
|
339
|
+
transition_rates=txrates,
|
|
340
|
+
stoichiometry=stoichiometry,
|
|
341
|
+
initial_state=initial_state,
|
|
342
|
+
initial_step=initial_step,
|
|
343
|
+
time_delta=time_delta,
|
|
344
|
+
num_steps=num_steps,
|
|
345
|
+
)
|
|
346
|
+
|
|
347
|
+
|
|
348
|
+
@tf.function
|
|
349
|
+
def simulate_one(elems):
|
|
350
|
+
"""One realisation of the epidemic process."""
|
|
351
|
+
return sirc.sample()
|
|
352
|
+
|
|
353
|
+
|
|
354
|
+
nsim = 15
|
|
355
|
+
eventlist = tf.map_fn(
|
|
356
|
+
simulate_one,
|
|
357
|
+
tf.ones([nsim, stoichiometry.shape[0]]),
|
|
358
|
+
fn_output_signature=dtype,
|
|
359
|
+
)
|
|
360
|
+
|
|
361
|
+
# Log prob of observing the eventlist, of second simulation, given the model
|
|
362
|
+
print("Log prob:", sirc.log_prob(eventlist[1, ...]))
|
|
363
|
+
|
|
364
|
+
# Plot timeseries using plot_timeseries() from first example
|
|
365
|
+
plot_timeseries(
|
|
366
|
+
initial_state,
|
|
367
|
+
eventlist,
|
|
368
|
+
stoichiometry,
|
|
369
|
+
initial_step,
|
|
370
|
+
time_delta,
|
|
371
|
+
num_steps,
|
|
372
|
+
nsim,
|
|
373
|
+
["k", "r", "b"],
|
|
374
|
+
["S", "I", "R"],
|
|
375
|
+
)
|
|
376
|
+
plt.show()
|
|
377
|
+
|
|
378
|
+
# Example 5: SIVR model for a single population (V=vaccinated)
|
|
379
|
+
dtype = tf.float32
|
|
380
|
+
|
|
381
|
+
# Initial state, counts per compartment (S, I, V, R), for one population
|
|
382
|
+
initial_state = tf.constant([[998, 1, 1, 0]], dtype)
|
|
383
|
+
|
|
384
|
+
# Stoichiometry matrix # S, I, V, R
|
|
385
|
+
stoichiometry = tf.constant(
|
|
386
|
+
[
|
|
387
|
+
[-1, 1, 0, 0], # S->I
|
|
388
|
+
[0, -1, 0, 1], # I->R
|
|
389
|
+
[-1, 0, 1, 0], # S->V
|
|
390
|
+
[0, 0, -1, 1],
|
|
391
|
+
], # V->R
|
|
392
|
+
dtype,
|
|
393
|
+
)
|
|
394
|
+
|
|
395
|
+
# time parameters
|
|
396
|
+
initial_step, time_delta, num_steps = 0.0, 1.0, 100
|
|
397
|
+
|
|
398
|
+
|
|
399
|
+
def txrates(t, state):
|
|
400
|
+
"""Transition rate per individual"""
|
|
401
|
+
beta, gamma = 0.3, 0.2
|
|
402
|
+
xi, eta = 0.6, 0.1
|
|
403
|
+
si = beta * state[:, 1] / tf.reduce_sum(state) # S->I transition rate
|
|
404
|
+
ir = tf.constant([gamma], dtype) # I->R transition rate
|
|
405
|
+
sv = xi * state[:, 1] / tf.reduce_sum(state) # S->V transition rate
|
|
406
|
+
vr = tf.constant([eta], dtype) # V->R transition rate
|
|
407
|
+
return [si, ir, sv, vr]
|
|
408
|
+
|
|
409
|
+
|
|
410
|
+
# Instantiate model
|
|
411
|
+
sivr = DiscreteTimeStateTransitionModel(
|
|
412
|
+
transition_rates=txrates,
|
|
413
|
+
stoichiometry=stoichiometry,
|
|
414
|
+
initial_state=initial_state,
|
|
415
|
+
initial_step=initial_step,
|
|
416
|
+
time_delta=time_delta,
|
|
417
|
+
num_steps=num_steps,
|
|
418
|
+
)
|
|
419
|
+
|
|
420
|
+
|
|
421
|
+
@tf.function
|
|
422
|
+
def simulate_one(elems):
|
|
423
|
+
"""One realisation of the epidemic process."""
|
|
424
|
+
return sivr.sample()
|
|
425
|
+
|
|
426
|
+
|
|
427
|
+
nsim = 30 # Number of realisations of the epidemic process
|
|
428
|
+
eventlist = tf.map_fn(
|
|
429
|
+
simulate_one,
|
|
430
|
+
tf.ones([nsim, stoichiometry.shape[0]]),
|
|
431
|
+
fn_output_signature=dtype,
|
|
432
|
+
)
|
|
433
|
+
|
|
434
|
+
# Events for each transition with shape (simulation, population, time,
|
|
435
|
+
# transition)
|
|
436
|
+
print("S->V events:", eventlist[0, 0, :, 2])
|
|
437
|
+
|
|
438
|
+
# Log prob of observing the eventlist, of third simulation, given the model
|
|
439
|
+
print("Log prob:", sivr.log_prob(eventlist[2, ...]))
|
|
440
|
+
|
|
441
|
+
# Plot timeseries using plot_timeseries() from first example
|
|
442
|
+
plot_timeseries(
|
|
443
|
+
initial_state,
|
|
444
|
+
eventlist,
|
|
445
|
+
stoichiometry,
|
|
446
|
+
initial_step,
|
|
447
|
+
time_delta,
|
|
448
|
+
num_steps,
|
|
449
|
+
nsim,
|
|
450
|
+
["k", "r", "g", "b"],
|
|
451
|
+
["S", "I", "V", "R"],
|
|
452
|
+
)
|
|
453
|
+
plt.show()
|