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,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,7 @@
1
+ """Experimental gemlib distributions"""
2
+
3
+ from gemlib.distributions.experimental.discrete_approx_cont_state_transition_model import ( # noqa: E501
4
+ DiscreteApproxContStateTransitionModel,
5
+ )
6
+
7
+ __all__ = ["DiscreteApproxContStateTransitionModel"]
@@ -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