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.
Files changed (75) hide show
  1. gemlib/__init__.py +9 -0
  2. gemlib/distributions/__init__.py +26 -0
  3. gemlib/distributions/brownian.py +141 -0
  4. gemlib/distributions/categorical2.py +33 -0
  5. gemlib/distributions/continuous_markov.py +371 -0
  6. gemlib/distributions/continuous_time_state_transition_model.py +185 -0
  7. gemlib/distributions/continuous_time_state_transition_model_test.py +289 -0
  8. gemlib/distributions/discrete_markov.py +279 -0
  9. gemlib/distributions/discrete_rejection_sampling.py +149 -0
  10. gemlib/distributions/discrete_time_state_transition_model.py +324 -0
  11. gemlib/distributions/discrete_time_state_transition_model_examples.py +453 -0
  12. gemlib/distributions/discrete_time_state_transition_model_test.py +336 -0
  13. gemlib/distributions/experimental/__init__.py +7 -0
  14. gemlib/distributions/experimental/discrete_approx_cont_state_transition_model.py +194 -0
  15. gemlib/distributions/experimental/state_transition_marginal_model.py +352 -0
  16. gemlib/distributions/hypergeometric.py +144 -0
  17. gemlib/distributions/hypergeometric_sampler.py +103 -0
  18. gemlib/distributions/hypergeometric_test.py +46 -0
  19. gemlib/distributions/kcategorical.py +113 -0
  20. gemlib/distributions/kcategorical_test.py +52 -0
  21. gemlib/distributions/uniform_integer.py +169 -0
  22. gemlib/distributions/uniform_integer_test.py +55 -0
  23. gemlib/mcmc/__init__.py +23 -0
  24. gemlib/mcmc/adaptive_random_walk_metropolis.py +859 -0
  25. gemlib/mcmc/adaptive_random_walk_metropolis_test.py +144 -0
  26. gemlib/mcmc/bb_fixture.pkl +0 -0
  27. gemlib/mcmc/brownian_bridge_kernel.py +291 -0
  28. gemlib/mcmc/brownian_bridge_kernel_test.py +164 -0
  29. gemlib/mcmc/chain_binomial_rippler.py +524 -0
  30. gemlib/mcmc/chain_binomial_rippler_test.py +120 -0
  31. gemlib/mcmc/compound_kernel.py +156 -0
  32. gemlib/mcmc/conftest.py +4 -0
  33. gemlib/mcmc/damped_chain_binomial_rippler.py +840 -0
  34. gemlib/mcmc/discrete_time_state_transition_model/__init__.py +21 -0
  35. gemlib/mcmc/discrete_time_state_transition_model/event_time_proposal.py +275 -0
  36. gemlib/mcmc/discrete_time_state_transition_model/fixtures.py +56 -0
  37. gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh.py +258 -0
  38. gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh_test.py +108 -0
  39. gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal.py +188 -0
  40. gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal_test.py +33 -0
  41. gemlib/mcmc/discrete_time_state_transition_model/move_events.py +239 -0
  42. gemlib/mcmc/discrete_time_state_transition_model/move_events_test.py +63 -0
  43. gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh.py +254 -0
  44. gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh_test.py +28 -0
  45. gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_proposal.py +150 -0
  46. gemlib/mcmc/discrete_time_state_transition_model/util.py +9 -0
  47. gemlib/mcmc/experimental/__init__.py +0 -0
  48. gemlib/mcmc/experimental/composable_kernel.py +270 -0
  49. gemlib/mcmc/experimental/discrete_time_state_transition_model/__init__.py +1 -0
  50. gemlib/mcmc/experimental/discrete_time_state_transition_model/left_censored_events_mh.py +149 -0
  51. gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events.py +71 -0
  52. gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events_test.py +48 -0
  53. gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh.py +89 -0
  54. gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh_test.py +28 -0
  55. gemlib/mcmc/experimental/hmc.py +111 -0
  56. gemlib/mcmc/experimental/hmc_test.py +61 -0
  57. gemlib/mcmc/experimental/mcmc_base.py +33 -0
  58. gemlib/mcmc/experimental/mcmc_sampler.py +94 -0
  59. gemlib/mcmc/experimental/mcmc_sampler_test.py +44 -0
  60. gemlib/mcmc/experimental/multi_scan.py +53 -0
  61. gemlib/mcmc/experimental/multi_scan_test.py +71 -0
  62. gemlib/mcmc/experimental/random_walk_metropolis.py +107 -0
  63. gemlib/mcmc/experimental/random_walk_metropolis_test.py +214 -0
  64. gemlib/mcmc/experimental/test_util.py +49 -0
  65. gemlib/mcmc/gibbs_kernel.py +505 -0
  66. gemlib/mcmc/gibbs_kernel_test.py +212 -0
  67. gemlib/mcmc/h5_posterior.py +77 -0
  68. gemlib/mcmc/multi_scan_kernel.py +59 -0
  69. gemlib/mcmc/zarr_posterior.py +132 -0
  70. gemlib/util.py +117 -0
  71. gemlib/util_test.py +75 -0
  72. gemlib-0.9.2.dist-info/LICENSE +21 -0
  73. gemlib-0.9.2.dist-info/METADATA +19 -0
  74. gemlib-0.9.2.dist-info/RECORD +75 -0
  75. 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()