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,324 @@
1
+ """Describes a DiscreteTimeStateTransitionModel."""
2
+
3
+ import tensorflow as tf
4
+ import tensorflow_probability as tfp
5
+ from tensorflow_probability.python.internal import (
6
+ dtype_util,
7
+ reparameterization,
8
+ samplers,
9
+ )
10
+
11
+ from gemlib.distributions.discrete_markov import (
12
+ compute_state,
13
+ discrete_markov_log_prob,
14
+ discrete_markov_simulation,
15
+ )
16
+ from gemlib.util import batch_gather, transition_coords
17
+
18
+ Tensor = tf.Tensor
19
+ tla = tf.linalg
20
+ tfd = tfp.distributions
21
+
22
+
23
+ class DiscreteTimeStateTransitionModel(tfd.Distribution):
24
+ """Discrete-time state transition model
25
+
26
+ A discrete-time state transition model assumes a population of
27
+ individuals is divided into a number of mutually exclusive states,
28
+ where transitions between states occur according to a Markov process.
29
+ Such models are commonly found in epidemiological and ecological
30
+ applications, where rapid implementation and modification is necessary.
31
+
32
+ This class provides a programmable implementation of the discrete-time
33
+ state transition model, compatible with TensorFlow Probability.
34
+ """
35
+
36
+ def __init__(
37
+ self,
38
+ transition_rates,
39
+ incidence_matrix,
40
+ initial_state,
41
+ initial_step,
42
+ time_delta,
43
+ num_steps,
44
+ validate_args=False,
45
+ allow_nan_stats=True,
46
+ name="DiscreteTimeStateTransitionModel",
47
+ ):
48
+ """A discrete-time Markov jump process for a state transition model.
49
+
50
+ Args:
51
+ ----
52
+ transition_rates: Python callable of the form `fn(t, state)` taking
53
+ the current time `t` (Python float) and state tensor `state`. This
54
+ function returns a tensor which broadcasts to the first dimension of
55
+ `incidence_matrix`.
56
+ incidence_matrix: `Tensor` representing the stochiometry matrix for
57
+ the state transition model where rows represent the transitions and
58
+ columns states.
59
+ initial_state: `Tensor` representing an initial state of counts per
60
+ compartment. The inner dimension is equal to the first dimension
61
+ of `incidence_matrix`.
62
+ initial_step: Python float representing an offset giving the time `t`
63
+ of the first time step in the model.
64
+ time_delta: Python float representing the size of the time step to be
65
+ used.
66
+ num_steps: Python integer representing the number of time steps across
67
+ which the model runs.
68
+
69
+
70
+ Example:
71
+ -------
72
+ A homogeneously mixing SIR model implementation::
73
+
74
+ import tensorflow as tf
75
+ from gemlib.distributions.discrete_time_state_transition_model \
76
+ import (
77
+ DiscreteTimeStateTransitionModel
78
+ )
79
+ from gemlib.util import compute_state
80
+
81
+ dtype = tf.float32
82
+
83
+ # Initial state, counts per compartment (S, I, R), for one
84
+ # population
85
+ initial_state = tf.constant([[99, 1, 0]], dtype)
86
+
87
+ # incidence_matrix S->I, I->R,
88
+ incidence_matrix = tf.constant([[-1, 0], # S
89
+ [1,-1], # I
90
+ [0,1]], # R
91
+
92
+ dtype)
93
+
94
+ # time parameters
95
+ initial_step, time_delta, num_steps = 0.0, 1.0, 100
96
+
97
+ def txrates(t, state):
98
+ # Transition rate per individual corresponding to each row of
99
+ # the incidence matrix.
100
+ #
101
+ # state: `Tensor` representing the current state (count of
102
+ # individuals in each compartment).
103
+ # t: Python float representing the current time. For example
104
+ # seasonality in the S->I
105
+ # transition could be driven by tensors of the following form
106
+ # seasonality = tf.math.sin(2 * 3.14159 * t / 20) + 1
107
+ # si = seasonality * beta * state[:, 1] / tf.reduce_sum(
108
+ # state)
109
+ #
110
+ # Returns: List of `Tensor`(s) each of which corresponds to a
111
+ # transition.
112
+
113
+ beta, gamma = 0.28, 0.14 # note R0=beta/gamma
114
+ si = beta * state[:, 1] / tf.reduce_sum(state)
115
+ ir = tf.constant([gamma], dtype)
116
+ return [si, ir]
117
+
118
+ # Instantiate model
119
+ sir = DiscreteTimeStateTransitionModel(
120
+ transition_rates=txrates,
121
+ incidence_matrix=incidence_matrix,
122
+ initial_state=initial_state,
123
+ initial_step=initial_step,
124
+ time_delta=time_delta,
125
+ num_steps=num_steps,
126
+ )
127
+
128
+ # One realisation of the epidemic process
129
+ @tf.function
130
+ def simulate_one(elems):
131
+ return sir.sample()
132
+
133
+ nsim = 15 # Number of realisations of the epidemic process
134
+ eventlist = tf.map_fn(simulate_one,
135
+ tf.ones([nsim, incidence_matrix.shape[0]]),
136
+ fn_output_signature=dtype)
137
+
138
+ # Events for each transition with shape (simulation, population,
139
+ # time, transition)
140
+ print('I->R events:', eventlist[0, 0, :, 1])
141
+
142
+ # Log prob of observing the eventlist, of first simulation, given
143
+ # the model
144
+ print('Log prob:', sir.log_prob(eventlist[0, ...]))
145
+
146
+ # Timeseries of counts per state with shape (simulation, population,
147
+ # time, state)
148
+ state_timeseries = compute_state(
149
+ initial_state,
150
+ eventlist,
151
+ incidence_matrix,
152
+ )
153
+ print('Susceptible state counts for first simulation:',
154
+ state_timeseries[0, 0, :, 0])
155
+
156
+ Note:
157
+ ----
158
+ See http://gitlab.com/gem-epidemics/gemlib/distributions/discrete_time_state_transition_model_examples.py
159
+ for further examples.
160
+
161
+ """
162
+ parameters = dict(locals())
163
+ with tf.name_scope(name) as name:
164
+ self._transition_rates = transition_rates
165
+ self._incidence_matrix = tf.convert_to_tensor(
166
+ incidence_matrix, dtype=initial_state.dtype
167
+ )
168
+ self._source_states = _compute_source_states(incidence_matrix)
169
+ self._initial_state = initial_state
170
+ self._initial_step = initial_step
171
+ self._time_delta = time_delta
172
+ self._num_steps = num_steps
173
+
174
+ super().__init__(
175
+ dtype=initial_state.dtype,
176
+ reparameterization_type=reparameterization.FULLY_REPARAMETERIZED,
177
+ validate_args=validate_args,
178
+ allow_nan_stats=allow_nan_stats,
179
+ parameters=parameters,
180
+ name=name,
181
+ )
182
+
183
+ self.dtype = initial_state.dtype
184
+
185
+ @property
186
+ def transition_rates(self):
187
+ return self._transition_rates
188
+
189
+ @property
190
+ def incidence_matrix(self):
191
+ return self._incidence_matrix
192
+
193
+ @property
194
+ def initial_state(self):
195
+ return self._initial_state
196
+
197
+ @property
198
+ def initial_step(self):
199
+ return self._initial_step
200
+
201
+ @property
202
+ def source_states(self):
203
+ return self._source_states
204
+
205
+ @property
206
+ def time_delta(self):
207
+ return self._time_delta
208
+
209
+ @property
210
+ def num_steps(self):
211
+ return self._num_steps
212
+
213
+ def _batch_shape(self):
214
+ return tf.TensorShape([])
215
+
216
+ def _event_shape(self):
217
+ shape = tf.TensorShape(
218
+ [
219
+ self.initial_state.shape[0],
220
+ tf.get_static_value(self._num_steps),
221
+ self._incidence_matrix.shape[1],
222
+ ]
223
+ )
224
+ return shape
225
+
226
+ def compute_state(
227
+ self, events: Tensor, include_final_state: bool = False
228
+ ) -> Tensor:
229
+ """Computes a state timeseries given a transition events
230
+
231
+ Args
232
+ ----
233
+ events: a `[self.num_steps, self.num_steps, self.num_events]`
234
+ shaped tensor of events
235
+ include_final_state: should the result include the final state? If
236
+ `False` (default), then `result.shape[1] ==
237
+ events.shape[1]`. If `True`, then
238
+ `results.shape[1] == events.shape[1] + 1`.
239
+
240
+ Returns
241
+ -------
242
+ A tensor of shape `[self.num_meta, self.num_steps, self.num_states]`
243
+ giving the number of individuals in each state at each time point in
244
+ each unit.
245
+ """
246
+ return compute_state(
247
+ incidence_matrix=self.incidence_matrix,
248
+ initial_state=self.initial_state,
249
+ events=events,
250
+ closed=include_final_state,
251
+ )
252
+
253
+ def _sample_n(self, n, seed=None):
254
+ """Runs a simulation from the epidemic model
255
+
256
+ :param param: a dictionary of model parameters
257
+ :param state_init: the initial state
258
+ :returns: a tuple of times and simulated states.
259
+ """
260
+ seed = samplers.sanitize_seed(
261
+ seed, salt="DiscreteTimeStateTransitionModel"
262
+ )
263
+ t, sim = discrete_markov_simulation(
264
+ hazard_fn=self.transition_rates,
265
+ state=self.initial_state,
266
+ start=self.initial_step,
267
+ end=self.initial_step + self.num_steps * self.time_delta,
268
+ time_step=self.time_delta,
269
+ incidence_matrix=self.incidence_matrix,
270
+ seed=seed,
271
+ )
272
+
273
+ # `sim` is `[T, M, S, S]`, and we need to pick out
274
+ # elements `[..., i, j]` for all our relevant transitions
275
+ # `i->j`. `batch_gather` computes these coordinates and
276
+ # invokes tf.gather.
277
+ indices = transition_coords(self.incidence_matrix)
278
+ sim = batch_gather(sim, indices)
279
+
280
+ # `sim` is now `[T, M, R]` structure for T times,
281
+ # M population units, and R transitions.
282
+ sim = tf.transpose(sim, perm=(1, 0, 2))
283
+ return tf.expand_dims(sim, 0)
284
+
285
+ def _log_prob(self, y, **kwargs):
286
+ dtype = dtype_util.common_dtype(
287
+ [y, self.initial_state], dtype_hint=self.dtype
288
+ )
289
+ y = tf.convert_to_tensor(y, dtype)
290
+
291
+ hazard = self.transition_rates
292
+ return discrete_markov_log_prob(
293
+ events=y,
294
+ init_state=self.initial_state,
295
+ init_step=self.initial_step,
296
+ time_delta=self.time_delta,
297
+ hazard_fn=hazard,
298
+ incidence_matrix=self.incidence_matrix,
299
+ )
300
+
301
+
302
+ def _compute_source_states(incidence_matrix, dtype=tf.int32):
303
+ """Computes the indices of the source states for each
304
+ transition in a state transition model.
305
+
306
+ :param incidence_matrix: incidence matrix in `[S, R]` orientation
307
+ for `S` states and `R` transitions.
308
+ :returns: a tensor of shape `(R,)` containing source state indices.
309
+ """
310
+ incidence_matrix = tf.transpose(incidence_matrix)
311
+
312
+ source_states = tf.reduce_sum(
313
+ tf.cumsum(
314
+ tf.clip_by_value(
315
+ -incidence_matrix, clip_value_min=0, clip_value_max=1
316
+ ),
317
+ axis=-1,
318
+ reverse=True,
319
+ exclusive=True,
320
+ ),
321
+ axis=-1,
322
+ )
323
+
324
+ return tf.cast(source_states, dtype)