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,336 @@
|
|
|
1
|
+
# Dependency imports
|
|
2
|
+
import numpy as np
|
|
3
|
+
import tensorflow as tf
|
|
4
|
+
import tensorflow_probability as tfp
|
|
5
|
+
from tensorflow_probability.python.internal import test_util
|
|
6
|
+
|
|
7
|
+
from gemlib.distributions.discrete_markov import compute_state
|
|
8
|
+
from gemlib.distributions.discrete_time_state_transition_model import (
|
|
9
|
+
DiscreteTimeStateTransitionModel,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@test_util.test_all_tf_execution_regimes
|
|
14
|
+
class TestDiscreteTimeStateTransitionModel(test_util.TestCase):
|
|
15
|
+
def setUp(self):
|
|
16
|
+
self.dtype = tf.float32
|
|
17
|
+
self.incidence_matrix = [[-1, 0], [1, -1], [0, 1]]
|
|
18
|
+
self.initial_state_A = [[99, 1, 0]]
|
|
19
|
+
self.initial_state_B = [[8000, 2000, 0]]
|
|
20
|
+
self.beta = 0.28
|
|
21
|
+
self.gamma = 0.14
|
|
22
|
+
self.nsim = 50
|
|
23
|
+
|
|
24
|
+
def init_model(
|
|
25
|
+
self,
|
|
26
|
+
beta,
|
|
27
|
+
gamma,
|
|
28
|
+
incidence_matrix,
|
|
29
|
+
initial_state,
|
|
30
|
+
initial_step=0.0,
|
|
31
|
+
time_delta=1.0,
|
|
32
|
+
num_steps=100,
|
|
33
|
+
dtype=tf.float32,
|
|
34
|
+
):
|
|
35
|
+
def txrates(t, state):
|
|
36
|
+
si = beta * state[:, 1] / tf.reduce_sum(state)
|
|
37
|
+
ir = tf.constant([gamma], dtype)
|
|
38
|
+
return [si, ir]
|
|
39
|
+
|
|
40
|
+
return DiscreteTimeStateTransitionModel(
|
|
41
|
+
transition_rates=txrates,
|
|
42
|
+
incidence_matrix=incidence_matrix,
|
|
43
|
+
initial_state=initial_state,
|
|
44
|
+
initial_step=initial_step,
|
|
45
|
+
time_delta=time_delta,
|
|
46
|
+
num_steps=num_steps,
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
def test_float32(self):
|
|
50
|
+
incidence_matrix = tf.constant(self.incidence_matrix, self.dtype)
|
|
51
|
+
initial_state = tf.constant(self.initial_state_A, self.dtype)
|
|
52
|
+
|
|
53
|
+
sir = self.init_model(
|
|
54
|
+
self.beta, self.gamma, incidence_matrix, initial_state, num_steps=5
|
|
55
|
+
)
|
|
56
|
+
|
|
57
|
+
eventlist = sir.sample()
|
|
58
|
+
eventlist_ = self.evaluate(eventlist)
|
|
59
|
+
self.assertDTypeEqual(eventlist_, np.float32)
|
|
60
|
+
|
|
61
|
+
lp = sir.log_prob(eventlist)
|
|
62
|
+
lp_ = self.evaluate(lp)
|
|
63
|
+
self.assertDTypeEqual(lp_, np.float32)
|
|
64
|
+
|
|
65
|
+
def test_float64(self):
|
|
66
|
+
dtype = tf.float64
|
|
67
|
+
incidence_matrix = tf.constant(self.incidence_matrix, dtype)
|
|
68
|
+
initial_state = tf.constant(self.initial_state_A, dtype)
|
|
69
|
+
|
|
70
|
+
sir = self.init_model(
|
|
71
|
+
self.beta,
|
|
72
|
+
self.gamma,
|
|
73
|
+
incidence_matrix,
|
|
74
|
+
initial_state,
|
|
75
|
+
num_steps=5,
|
|
76
|
+
dtype=dtype,
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
eventlist = sir.sample()
|
|
80
|
+
eventlist_ = self.evaluate(eventlist)
|
|
81
|
+
self.assertDTypeEqual(eventlist_, np.float64)
|
|
82
|
+
|
|
83
|
+
lp = sir.log_prob(eventlist)
|
|
84
|
+
lp_ = self.evaluate(lp)
|
|
85
|
+
self.assertDTypeEqual(lp_, np.float64)
|
|
86
|
+
|
|
87
|
+
def test_non_integer_time_steps(self):
|
|
88
|
+
incidence_matrix = tf.constant(self.incidence_matrix, self.dtype)
|
|
89
|
+
initial_state = tf.constant(self.initial_state_A, self.dtype)
|
|
90
|
+
|
|
91
|
+
sir = self.init_model(
|
|
92
|
+
self.beta,
|
|
93
|
+
self.gamma,
|
|
94
|
+
incidence_matrix,
|
|
95
|
+
initial_state,
|
|
96
|
+
initial_step=1.5,
|
|
97
|
+
time_delta=0.5,
|
|
98
|
+
num_steps=100,
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
eventlist = sir.sample()
|
|
102
|
+
self.assertShapeEqual(np.ndarray(shape=(1, 100, 2)), eventlist)
|
|
103
|
+
|
|
104
|
+
lp = sir.log_prob(eventlist)
|
|
105
|
+
self.assertShapeEqual(np.ndarray(shape=()), lp)
|
|
106
|
+
|
|
107
|
+
def test_log_prob_over_simuations(self):
|
|
108
|
+
incidence_matrix = tf.constant(self.incidence_matrix, self.dtype)
|
|
109
|
+
initial_state = tf.constant(self.initial_state_B, self.dtype)
|
|
110
|
+
|
|
111
|
+
sir = self.init_model(
|
|
112
|
+
self.beta, self.gamma, incidence_matrix, initial_state, num_steps=60
|
|
113
|
+
)
|
|
114
|
+
|
|
115
|
+
def simulate_one(elems):
|
|
116
|
+
return sir.sample()
|
|
117
|
+
|
|
118
|
+
eventlist = tf.map_fn(
|
|
119
|
+
simulate_one,
|
|
120
|
+
tf.ones([self.nsim, incidence_matrix.shape[1]]),
|
|
121
|
+
fn_output_signature=self.dtype,
|
|
122
|
+
)
|
|
123
|
+
|
|
124
|
+
lp = tf.vectorized_map(
|
|
125
|
+
fn=lambda i: sir.log_prob(eventlist[i, ...]),
|
|
126
|
+
elems=tf.range(self.nsim),
|
|
127
|
+
)
|
|
128
|
+
lp_mean = tf.math.reduce_mean(lp)
|
|
129
|
+
lp_mean_ = self.evaluate(lp_mean)
|
|
130
|
+
actual_mean = (
|
|
131
|
+
-395
|
|
132
|
+
) # sample_mean ~= -395 derived from 1000 simulations of this model
|
|
133
|
+
self.assertAllClose(
|
|
134
|
+
lp_mean_, actual_mean, rtol=1e-06, atol=8.1
|
|
135
|
+
) # sample_variance ~= 65
|
|
136
|
+
|
|
137
|
+
def test_model_constraints(self):
|
|
138
|
+
incidence_matrix = tf.constant(self.incidence_matrix, self.dtype)
|
|
139
|
+
initial_state = tf.constant(self.initial_state_A, self.dtype)
|
|
140
|
+
time_delta = 1.0
|
|
141
|
+
num_steps = 100
|
|
142
|
+
|
|
143
|
+
sir = self.init_model(
|
|
144
|
+
self.beta,
|
|
145
|
+
self.gamma,
|
|
146
|
+
incidence_matrix,
|
|
147
|
+
initial_state,
|
|
148
|
+
time_delta=time_delta,
|
|
149
|
+
num_steps=num_steps,
|
|
150
|
+
)
|
|
151
|
+
|
|
152
|
+
def simulate_one(elems):
|
|
153
|
+
return sir.sample()
|
|
154
|
+
|
|
155
|
+
eventlist = tf.map_fn(
|
|
156
|
+
simulate_one,
|
|
157
|
+
tf.ones([self.nsim, incidence_matrix.shape[1]]),
|
|
158
|
+
fn_output_signature=self.dtype,
|
|
159
|
+
)
|
|
160
|
+
ts = compute_state(initial_state, eventlist, incidence_matrix)
|
|
161
|
+
|
|
162
|
+
# Crude check that some simulations have nontrivial dynamics
|
|
163
|
+
# i.e. some units arrived in recovered compartments for some simulations
|
|
164
|
+
sum_at_tmax = tf.reduce_sum(ts[:, :, num_steps - 1, 2])
|
|
165
|
+
test_sum_at_tmax = (
|
|
166
|
+
tf.cast(self.nsim * num_steps, self.dtype) / 4
|
|
167
|
+
) # factor 4 is a choice
|
|
168
|
+
self.assertGreater(
|
|
169
|
+
self.evaluate(sum_at_tmax), self.evaluate(test_sum_at_tmax)
|
|
170
|
+
)
|
|
171
|
+
|
|
172
|
+
# Check N is conserved at each time point
|
|
173
|
+
# Note dS/dt + dI/dt + dR/dt = 0 then integrating over dt leads to
|
|
174
|
+
# N = S + I + R
|
|
175
|
+
sums_at_t = tf.vectorized_map(
|
|
176
|
+
fn=lambda i: tf.reduce_sum(ts[:, :, i, :]),
|
|
177
|
+
elems=tf.range(num_steps),
|
|
178
|
+
)
|
|
179
|
+
expected_sums = tf.broadcast_to(
|
|
180
|
+
tf.cast(self.nsim * num_steps, self.dtype), [num_steps]
|
|
181
|
+
)
|
|
182
|
+
self.assertAllClose(sums_at_t, expected_sums, rtol=1e-06, atol=1e-06)
|
|
183
|
+
|
|
184
|
+
# Check dS/dt + dI/dt + dR/dt = 0 at each time point
|
|
185
|
+
def forward_difference(i):
|
|
186
|
+
"""Numerical differentiation of states wrt time using forward
|
|
187
|
+
difference.
|
|
188
|
+
"""
|
|
189
|
+
x1 = ts[i, 0, :, :]
|
|
190
|
+
x2 = tf.roll(x1, shift=-1, axis=0)
|
|
191
|
+
diffs = (
|
|
192
|
+
tf.math.subtract(
|
|
193
|
+
x2[0 : ts.shape[-2] - 1, :], x1[0 : ts.shape[-2] - 1, :]
|
|
194
|
+
)
|
|
195
|
+
/ time_delta
|
|
196
|
+
)
|
|
197
|
+
return tf.math.reduce_sum(diffs, axis=-1)
|
|
198
|
+
|
|
199
|
+
finite_diffs = tf.vectorized_map(
|
|
200
|
+
fn=forward_difference, elems=tf.range(self.nsim)
|
|
201
|
+
)
|
|
202
|
+
expected_diffs = tf.zeros_like(finite_diffs, self.dtype)
|
|
203
|
+
self.assertAllClose(
|
|
204
|
+
finite_diffs, expected_diffs, rtol=1e-06, atol=1e-06
|
|
205
|
+
)
|
|
206
|
+
|
|
207
|
+
def test_model_dynamics(self):
|
|
208
|
+
"""Check simulation adheres to the SIR system of ODEs.
|
|
209
|
+
|
|
210
|
+
This check is performed without being in the thermodynamic limit
|
|
211
|
+
(N->inf, t->inf).
|
|
212
|
+
|
|
213
|
+
Let dS/dt=-bIS/N, dI/dt=bIS/N-gI and dR/dt=gI.
|
|
214
|
+
Dividing first equation by third gives dS/dR=-b/g.S/N.
|
|
215
|
+
Separating variables and integrating wrt dR yields
|
|
216
|
+
int(1/S, dS)=-b/g.1/N.int(1, dR).
|
|
217
|
+
Let the integrals have limits S(0), S(t), R(0), R(t).
|
|
218
|
+
The solution to this integral is the transcendental equation
|
|
219
|
+
S(t)=S(0)exp(-b/g.(R(t)-R(0))/N). Due the stochastic nature of the
|
|
220
|
+
chain binomial algorithm naively checking the simulated right
|
|
221
|
+
hand side of this solution equals (with a given tolerence) the simulated
|
|
222
|
+
left hand side is fraught with difficulty. However rearranging the
|
|
223
|
+
solution in terms of the time invariant factor
|
|
224
|
+
b/g=-N.ln(S(t)/S(0))/(R(t)-R(0)) makes it possible to check the
|
|
225
|
+
simulated dynamics adhere to the dynamics given by the SIR system of
|
|
226
|
+
ODEs (except when R(t)=R0).
|
|
227
|
+
|
|
228
|
+
"""
|
|
229
|
+
|
|
230
|
+
incidence_matrix = tf.constant(self.incidence_matrix, self.dtype)
|
|
231
|
+
initial_state = tf.constant(self.initial_state_B, self.dtype) * 10
|
|
232
|
+
num_steps = 200
|
|
233
|
+
buffer = 20 # number if initial steps to be omitted
|
|
234
|
+
|
|
235
|
+
sir = self.init_model(
|
|
236
|
+
self.beta,
|
|
237
|
+
self.gamma,
|
|
238
|
+
incidence_matrix,
|
|
239
|
+
initial_state,
|
|
240
|
+
time_delta=0.25,
|
|
241
|
+
num_steps=num_steps,
|
|
242
|
+
)
|
|
243
|
+
|
|
244
|
+
def simulate_one(elems):
|
|
245
|
+
return sir.sample()
|
|
246
|
+
|
|
247
|
+
eventlist = tf.map_fn(
|
|
248
|
+
simulate_one,
|
|
249
|
+
tf.ones([self.nsim, incidence_matrix.shape[1]]),
|
|
250
|
+
fn_output_signature=self.dtype,
|
|
251
|
+
)
|
|
252
|
+
ts = compute_state(initial_state, eventlist, incidence_matrix)
|
|
253
|
+
|
|
254
|
+
S0 = initial_state[0, -3]
|
|
255
|
+
R0 = initial_state[0, -1]
|
|
256
|
+
St = ts[:, 0, :, -3]
|
|
257
|
+
Rt = ts[:, 0, :, -1]
|
|
258
|
+
N = tf.reduce_sum(initial_state)
|
|
259
|
+
r0_sim = -N * tf.math.log(St / S0) / (Rt - R0) # r0=beta/gamma
|
|
260
|
+
|
|
261
|
+
# Crude check that some simulations have nontrivial dynamics
|
|
262
|
+
sum_at_tmax = tf.reduce_sum(ts[:, :, num_steps - 1, 2])
|
|
263
|
+
test_sum_at_tmax = (
|
|
264
|
+
tf.cast(self.nsim * num_steps, self.dtype) / 4
|
|
265
|
+
) # factor 4 is a choice
|
|
266
|
+
self.assertGreater(
|
|
267
|
+
self.evaluate(sum_at_tmax), self.evaluate(test_sum_at_tmax)
|
|
268
|
+
)
|
|
269
|
+
|
|
270
|
+
# First time step must be omitted as it will be undefined due to
|
|
271
|
+
# division by zero n.b. R(t=0)=R0
|
|
272
|
+
# Soft test - summarise the mean of each simulation
|
|
273
|
+
mean_r0_sim = tf.reduce_mean(r0_sim[:, buffer:num_steps], axis=1)
|
|
274
|
+
r0_actual = tf.broadcast_to(self.beta / self.gamma, [self.nsim])
|
|
275
|
+
self.assertAllClose(
|
|
276
|
+
mean_r0_sim, r0_actual, rtol=1e-06, atol=0.11
|
|
277
|
+
) # atol scales inversely with the size of N
|
|
278
|
+
|
|
279
|
+
# Hard test - check all times apart from a few initial steps when R(t)
|
|
280
|
+
# may equal R0
|
|
281
|
+
r0_all_actual = tf.broadcast_to(
|
|
282
|
+
self.beta / self.gamma, [self.nsim, num_steps - buffer]
|
|
283
|
+
)
|
|
284
|
+
self.assertAllClose(
|
|
285
|
+
r0_sim[:, buffer:num_steps], r0_all_actual, rtol=1e-06, atol=0.14
|
|
286
|
+
)
|
|
287
|
+
|
|
288
|
+
|
|
289
|
+
@test_util.test_all_tf_execution_regimes
|
|
290
|
+
class TestDiscreteTimeStateTransitionModelLogProbMaxima(test_util.TestCase):
|
|
291
|
+
def init_model(self, params, incidence_matrix, initial_state):
|
|
292
|
+
def txrates(t, state):
|
|
293
|
+
si = 1e-9 + beta * state[:, 1] / tf.reduce_sum(state)
|
|
294
|
+
ir = tf.expand_dims(gamma, axis=0)
|
|
295
|
+
return [si, ir]
|
|
296
|
+
|
|
297
|
+
beta, gamma = tf.unstack(params)
|
|
298
|
+
return DiscreteTimeStateTransitionModel(
|
|
299
|
+
transition_rates=txrates,
|
|
300
|
+
incidence_matrix=incidence_matrix,
|
|
301
|
+
initial_state=initial_state,
|
|
302
|
+
initial_step=0.0,
|
|
303
|
+
time_delta=1.0,
|
|
304
|
+
num_steps=100,
|
|
305
|
+
)
|
|
306
|
+
|
|
307
|
+
def test_log_prob_mle(self):
|
|
308
|
+
"""Test maximum likelihood estimation"""
|
|
309
|
+
|
|
310
|
+
dtype = np.float32
|
|
311
|
+
incidence_matrix = tf.constant([[-1, 0], [1, -1], [0, 1]], dtype)
|
|
312
|
+
initial_state = tf.constant([[8000, 2000, 0]], dtype)
|
|
313
|
+
pars = tf.constant([0.5, 0.3], dtype)
|
|
314
|
+
|
|
315
|
+
# Simulate a dataset
|
|
316
|
+
sir_orig = self.init_model(pars, incidence_matrix, initial_state)
|
|
317
|
+
events = self.evaluate(sir_orig.sample(seed=(0, 0)))
|
|
318
|
+
print(np.sum(events, axis=-2))
|
|
319
|
+
|
|
320
|
+
def logp(pars):
|
|
321
|
+
return -self.init_model(
|
|
322
|
+
pars, incidence_matrix, initial_state
|
|
323
|
+
).log_prob(events)
|
|
324
|
+
|
|
325
|
+
optim_results = tfp.optimizer.nelder_mead_minimize(
|
|
326
|
+
logp,
|
|
327
|
+
initial_vertex=tf.zeros_like(pars),
|
|
328
|
+
)
|
|
329
|
+
print(self.evaluate(optim_results))
|
|
330
|
+
|
|
331
|
+
self.assertAllTrue(optim_results.converged)
|
|
332
|
+
self.assertAllClose(optim_results.position, pars, rtol=0.01, atol=0.005)
|
|
333
|
+
|
|
334
|
+
|
|
335
|
+
if __name__ == "__main__":
|
|
336
|
+
tf.test.main()
|
|
@@ -0,0 +1,194 @@
|
|
|
1
|
+
"""Describes a continuous time State Transition Model with discrete event time
|
|
2
|
+
approximation.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
import tensorflow as tf
|
|
6
|
+
import tensorflow_probability as tfp
|
|
7
|
+
from tensorflow_probability.python.internal import (
|
|
8
|
+
dtype_util,
|
|
9
|
+
reparameterization,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
from gemlib.distributions.discrete_markov import (
|
|
13
|
+
_transition_coords,
|
|
14
|
+
compute_state,
|
|
15
|
+
discrete_markov_simulation,
|
|
16
|
+
)
|
|
17
|
+
from gemlib.util import batch_gather
|
|
18
|
+
|
|
19
|
+
tla = tf.linalg
|
|
20
|
+
tfd = tfp.distributions
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class DiscreteApproxContStateTransitionModel(tfd.Distribution):
|
|
24
|
+
def __init__(
|
|
25
|
+
self,
|
|
26
|
+
transition_rates,
|
|
27
|
+
stoichiometry,
|
|
28
|
+
initial_state,
|
|
29
|
+
initial_step,
|
|
30
|
+
time_delta,
|
|
31
|
+
num_steps,
|
|
32
|
+
validate_args=False,
|
|
33
|
+
allow_nan_stats=True,
|
|
34
|
+
name="StateTransitionMarginalModel",
|
|
35
|
+
):
|
|
36
|
+
"""Implements a discrete-time Markov jump process for a state transition
|
|
37
|
+
model.
|
|
38
|
+
|
|
39
|
+
:param transition_rates: a function of the form `fn(t, state)` taking
|
|
40
|
+
the current time `t` and state tensor `state`.
|
|
41
|
+
This function returns a tensor which broadcasts
|
|
42
|
+
to the first dimension of `stoichiometry`.
|
|
43
|
+
Transition rates are assumed to be risk ratios,
|
|
44
|
+
with the baseline hazard rate marginalised out
|
|
45
|
+
from the model.
|
|
46
|
+
:param stoichiometry: the stochiometry matrix for the state transition
|
|
47
|
+
model with rows representing transitions and
|
|
48
|
+
columns representing states.
|
|
49
|
+
:param initial_state: an initial state tensor with inner dimension equal
|
|
50
|
+
to the first dimension of `stoichiometry`.
|
|
51
|
+
:param initial_step: an offset giving the time `t` of the first timestep
|
|
52
|
+
in the model.
|
|
53
|
+
:param time_delta: the size of the time step to be used.
|
|
54
|
+
:param num_steps: the number of time steps across which the model runs.
|
|
55
|
+
"""
|
|
56
|
+
parameters = dict(locals())
|
|
57
|
+
with tf.name_scope(name) as name:
|
|
58
|
+
self._transition_rates = transition_rates
|
|
59
|
+
self._stoichiometry = tf.convert_to_tensor(
|
|
60
|
+
stoichiometry, dtype=initial_state.dtype
|
|
61
|
+
)
|
|
62
|
+
self._initial_state = initial_state
|
|
63
|
+
self._initial_step = initial_step
|
|
64
|
+
self._time_delta = time_delta
|
|
65
|
+
self._num_steps = num_steps
|
|
66
|
+
|
|
67
|
+
super().__init__(
|
|
68
|
+
dtype=initial_state.dtype,
|
|
69
|
+
reparameterization_type=reparameterization.FULLY_REPARAMETERIZED,
|
|
70
|
+
validate_args=validate_args,
|
|
71
|
+
allow_nan_stats=allow_nan_stats,
|
|
72
|
+
parameters=parameters,
|
|
73
|
+
name=name,
|
|
74
|
+
)
|
|
75
|
+
|
|
76
|
+
self.dtype = initial_state.dtype
|
|
77
|
+
|
|
78
|
+
@property
|
|
79
|
+
def transition_rates(self):
|
|
80
|
+
return self._transition_rates
|
|
81
|
+
|
|
82
|
+
@property
|
|
83
|
+
def baseline_hazard_rate_priors(self):
|
|
84
|
+
return self._baseline_hazard_rate_priors
|
|
85
|
+
|
|
86
|
+
@property
|
|
87
|
+
def stoichiometry(self):
|
|
88
|
+
return self._stoichiometry
|
|
89
|
+
|
|
90
|
+
@property
|
|
91
|
+
def initial_state(self):
|
|
92
|
+
return self._initial_state
|
|
93
|
+
|
|
94
|
+
@property
|
|
95
|
+
def initial_step(self):
|
|
96
|
+
return self._initial_step
|
|
97
|
+
|
|
98
|
+
@property
|
|
99
|
+
def time_delta(self):
|
|
100
|
+
return self._time_delta
|
|
101
|
+
|
|
102
|
+
@property
|
|
103
|
+
def num_steps(self):
|
|
104
|
+
return self._num_steps
|
|
105
|
+
|
|
106
|
+
def _batch_shape(self):
|
|
107
|
+
return tf.TensorShape([])
|
|
108
|
+
|
|
109
|
+
def _event_shape(self):
|
|
110
|
+
shape = tf.TensorShape(
|
|
111
|
+
[
|
|
112
|
+
self.initial_state.shape[0],
|
|
113
|
+
tf.get_static_value(self._num_steps),
|
|
114
|
+
self._stoichiometry.shape[0],
|
|
115
|
+
]
|
|
116
|
+
)
|
|
117
|
+
return shape
|
|
118
|
+
|
|
119
|
+
def _sample_n(self, n, seed=None):
|
|
120
|
+
"""Runs a simulation from the epidemic model
|
|
121
|
+
|
|
122
|
+
:param param: a dictionary of model parameters
|
|
123
|
+
:param state_init: the initial state
|
|
124
|
+
:returns: a tuple of times and simulated states.
|
|
125
|
+
"""
|
|
126
|
+
with tf.name_scope("DiscreteTimeStateTransitionModel.log_prob"):
|
|
127
|
+
|
|
128
|
+
def hazard_fn(t, state):
|
|
129
|
+
return self.transition_rates(t, state)
|
|
130
|
+
|
|
131
|
+
t, sim = discrete_markov_simulation(
|
|
132
|
+
hazard_fn=hazard_fn,
|
|
133
|
+
state=self.initial_state,
|
|
134
|
+
start=self.initial_step,
|
|
135
|
+
end=self.initial_step + self.num_steps * self.time_delta,
|
|
136
|
+
time_step=self.time_delta,
|
|
137
|
+
stoichiometry=self.stoichiometry,
|
|
138
|
+
seed=seed,
|
|
139
|
+
)
|
|
140
|
+
indices = _transition_coords(self.stoichiometry)
|
|
141
|
+
sim = batch_gather(sim, indices)
|
|
142
|
+
sim = tf.transpose(sim, perm=(1, 0, 2))
|
|
143
|
+
return tf.expand_dims(sim, 0)
|
|
144
|
+
|
|
145
|
+
def _log_prob(self, y, **kwargs):
|
|
146
|
+
"""Calculates the log probability of observing epidemic events y
|
|
147
|
+
:param y: a list of tensors. The first is of shape [n_times] containing
|
|
148
|
+
times, the second is of shape [n_times, n_states, n_states]
|
|
149
|
+
containing event matrices.
|
|
150
|
+
:param param: a list of parameters
|
|
151
|
+
:returns: a scalar giving the log probability of the epidemic
|
|
152
|
+
"""
|
|
153
|
+
dtype = dtype_util.common_dtype(
|
|
154
|
+
[y, self.initial_state], dtype_hint=self.dtype
|
|
155
|
+
)
|
|
156
|
+
events = tf.convert_to_tensor(y, dtype)
|
|
157
|
+
with tf.name_scope("DiscreteApproxContStateTransitionModel.log_prob"):
|
|
158
|
+
state_timeseries = compute_state(
|
|
159
|
+
self.initial_state,
|
|
160
|
+
events,
|
|
161
|
+
self.stoichiometry,
|
|
162
|
+
closed=True,
|
|
163
|
+
)
|
|
164
|
+
|
|
165
|
+
tms_timeseries = tf.transpose(state_timeseries, perm=(1, 0, 2))
|
|
166
|
+
tmr_events = tf.transpose(events, perm=(1, 0, 2))
|
|
167
|
+
|
|
168
|
+
def fn(elems):
|
|
169
|
+
return tf.stack(self.transition_rates(*elems), axis=-1)
|
|
170
|
+
|
|
171
|
+
rates = tf.vectorized_map(
|
|
172
|
+
fn=fn,
|
|
173
|
+
elems=(
|
|
174
|
+
self.initial_step + tf.range(tms_timeseries.shape[0]),
|
|
175
|
+
tms_timeseries,
|
|
176
|
+
),
|
|
177
|
+
)
|
|
178
|
+
|
|
179
|
+
def integrated_rates():
|
|
180
|
+
"""Use mid-point integration to estimate the constant rate
|
|
181
|
+
over time
|
|
182
|
+
"""
|
|
183
|
+
integrated_rates = tms_timeseries[..., :-1] * rates
|
|
184
|
+
return (
|
|
185
|
+
integrated_rates[:-1, ...] + integrated_rates[1:, ...]
|
|
186
|
+
) / 2.0
|
|
187
|
+
|
|
188
|
+
log_hazard_rate = tf.reduce_sum(
|
|
189
|
+
tmr_events * tf.math.log(integrated_rates())
|
|
190
|
+
)
|
|
191
|
+
log_survival = tf.reduce_sum(integrated_rates()) * self.time_delta
|
|
192
|
+
log_denom = tf.reduce_sum(tf.math.lgamma(tmr_events + 1.0))
|
|
193
|
+
|
|
194
|
+
return log_hazard_rate - log_survival - log_denom
|