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,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)
|