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,840 @@
1
+ """Chain binomial process rippler algorithm"""
2
+
3
+ import warnings
4
+ from collections import namedtuple
5
+
6
+ import tensorflow as tf
7
+ import tensorflow_probability as tfp
8
+ from tensorflow_probability.python.internal import prefer_static
9
+ from tensorflow_probability.python.mcmc.internal import util as mcmc_util
10
+
11
+ from gemlib.distributions import Hypergeometric, UniformInteger
12
+
13
+ tfd = tfp.distributions
14
+ samplers = tfp.random
15
+
16
+ __all__ = [
17
+ "default_initial_ripple",
18
+ "damped_initial_ripple_fn",
19
+ "DampedCBRKernel",
20
+ ]
21
+
22
+
23
+ def _compute_state(initial_state, events, stoichiometry, closed=False):
24
+ """Computes a state tensor from initial state and event tensor
25
+
26
+ :param initial_state: a tensor of shape [S, M]
27
+ :param events: a tensor of shape [T, R, M]
28
+ :param stoichiometry: a stoichiometry matrix of shape [R, S] describing
29
+ how transitions update the state.
30
+ :param closed: if `True`, return state in close interval [0, T], otherwise
31
+ [0, T)
32
+ :return: a tensor of shape [T, S, M] if `closed=False` or [T+1, S, M] if
33
+ `closed=True`
34
+ describing the state of the
35
+ system for each batch M at time T.
36
+ """
37
+ if isinstance(stoichiometry, tf.Tensor):
38
+ stoichiometry = prefer_static.cast(stoichiometry, dtype=events.dtype)
39
+ else:
40
+ stoichiometry = tf.convert_to_tensor(stoichiometry, dtype=events.dtype)
41
+
42
+ increments = tf.einsum("...trm,rs->...tsm", events, stoichiometry)
43
+
44
+ if closed is False:
45
+ cum_increments = tf.cumsum(increments, axis=-3, exclusive=True)
46
+ else:
47
+ cum_increments = tf.cumsum(increments, axis=-3, exclusive=False)
48
+ cum_increments = tf.concat(
49
+ [tf.zeros_like(cum_increments[..., 0:1, :, :]), cum_increments],
50
+ axis=-2,
51
+ )
52
+ state = cum_increments + tf.expand_dims(initial_state, axis=-3)
53
+ return state
54
+
55
+
56
+ class DampingFunction:
57
+ def __init__(self, p, upper_bound, gamma):
58
+ self._parameters = locals()
59
+ self._dtype = self.p.dtype
60
+
61
+ @property
62
+ def p(self):
63
+ return tf.convert_to_tensor(self._parameters["p"])
64
+
65
+ @property
66
+ def upper_bound(self):
67
+ return tf.convert_to_tensor(self._parameters["upper_bound"])
68
+
69
+ @property
70
+ def gamma(self):
71
+ return tf.convert_to_tensor(self._parameters["gamma"])
72
+
73
+ def __call__(self, u):
74
+ u = tf.convert_to_tensor(u, dtype=self._dtype)
75
+
76
+ r_transformed = self.p + tf.math.pow(
77
+ self.upper_bound - self.p, 1 - self.gamma
78
+ ) * tf.math.pow(u - self.p, self.gamma)
79
+ return tf.where(
80
+ (self.p < u) & (u <= self.upper_bound), r_transformed, u
81
+ )
82
+
83
+ def forward(self, u):
84
+ return self.__call__(u)
85
+
86
+ def inverse(self, u):
87
+ u = tf.convert_to_tensor(u, dtype=self._dtype)
88
+
89
+ r_transformed = self.p + tf.math.pow(
90
+ u - self.p, 1 / self.gamma
91
+ ) * tf.math.pow(self.upper_bound - self.p, 1 - 1 / self.gamma)
92
+ return tf.where(
93
+ (self.p < u) & (u <= self.upper_bound), r_transformed, u
94
+ )
95
+
96
+ def log_inverse_jacobian(self, u):
97
+ """N.B. only implemented for p < u <= upper_bound"""
98
+ u = tf.convert_to_tensor(u, dtype=self._dtype)
99
+
100
+ deriv = (
101
+ (1.0 - 1.0 / self.gamma) * tf.math.log(self.upper_bound - self.p)
102
+ + (1.0 / self.gamma - 1.0) * tf.math.log(u - self.p)
103
+ - tf.math.log(self.gamma)
104
+ )
105
+
106
+ return tf.where((self.p < u) & (u <= self.upper_bound), deriv, 0.0)
107
+
108
+
109
+ def _reduce_first_n(values, n):
110
+ """Reduces the first `n` elements of `values` over the first dimension
111
+ of `values` in a vectorized way.
112
+ """
113
+ values_shape = tf.shape(values)
114
+ seq = tf.range(values_shape[0], dtype=n.dtype)
115
+ seq = tf.reshape(seq, shape=[values_shape[0]] + [1] * len(values.shape[1:]))
116
+ mask = tf.cast(seq < n, values.dtype)
117
+ return tf.reduce_sum(values * mask, axis=0)
118
+
119
+
120
+ def _reduce_uniforms(
121
+ p_low,
122
+ p_high,
123
+ size,
124
+ u_transform_fn,
125
+ seed,
126
+ chunksize=10,
127
+ ):
128
+ """Generates `size` U(`p_low`, `p_high`) variates, and reduces
129
+ through the transform `u_transform_fn`.
130
+ """
131
+ seed = samplers.sanitize_seed(seed, salt="binomial_log_jacobian")
132
+ size = tf.cast(size, tf.int32)
133
+
134
+ def cond(i, _):
135
+ return tf.reduce_sum(size - i * chunksize) > 0
136
+
137
+ def body(i, accum):
138
+ u = tfd.Uniform(low=p_low, high=p_high).sample(
139
+ sample_shape=chunksize, seed=seed
140
+ )
141
+ size_local = tf.clip_by_value(
142
+ size - i * chunksize, clip_value_min=0, clip_value_max=chunksize
143
+ )
144
+ log_jacobian = _reduce_first_n(u_transform_fn(u), size_local)
145
+ return i + 1, accum + log_jacobian
146
+
147
+ _, log_jacobian = tf.while_loop(
148
+ cond, body, loop_vars=(0, tf.zeros_like(p_low))
149
+ )
150
+
151
+ return log_jacobian
152
+
153
+
154
+ def _p_step_ps_gt_p(z, x, p, ps, gamma, seed):
155
+ r"""Run the damped p-step where ps > p
156
+
157
+ :param z: a [R, M] tensor of events for R transitions and M units
158
+ :param x: a [S, M] tensor of state values for S states and M units
159
+ :param p: a [R, M] tensor of transition rates for the current time series
160
+ :param ps: a [R, M] tensor of transition rates for the rippled time series
161
+ :param gamma: a scalar ($\gamma \geq 1$) giving the damping.
162
+ :param seed: a random seed.
163
+ :returns: a tuple of ([R, M] tensor of new event numbers,
164
+ log_acceptance_correction)
165
+ """
166
+ seeds = samplers.split_seed(seed, n=4, salt="_p_step_ps_gt_p")
167
+
168
+ upper_bound = 2 * ps - p
169
+
170
+ damp = DampingFunction(p, upper_bound, gamma)
171
+
172
+ p_v = (upper_bound - p) / (1.0 - p)
173
+ p_w = (damp(ps) - p) / (upper_bound - p)
174
+
175
+ v = tfd.Binomial(total_count=x - z, probs=p_v).sample(seed=seeds[0])
176
+ w = tfd.Binomial(total_count=v, probs=p_w).sample(seed=seeds[1])
177
+
178
+ # Chunk up the sampling domain and use masking here
179
+ def jacobian(x):
180
+ return damp.log_inverse_jacobian(x)
181
+
182
+ neg_log_inverse_jacobian = _reduce_uniforms(
183
+ p, ps, w, jacobian, seeds[2]
184
+ ) + _reduce_uniforms(ps, upper_bound, v - w, jacobian, seeds[3])
185
+
186
+ return (
187
+ z + w,
188
+ neg_log_inverse_jacobian,
189
+ ) # log_acceptance_correction
190
+
191
+
192
+ def _p_step_ps_leq_p(z, x, p, ps, gamma, seed):
193
+ r"""Run the damped p-step for ps <= p
194
+
195
+ :param z: a [R, M] tensor of events for R transitions and M units
196
+ :param x: a [S, M] tensor of state values for S states and M units
197
+ :param p: a [R, M] tensor of transition rates for the current time series
198
+ :param ps: a [R, M] tensor of transition rates for the rippled time series
199
+ :param gamma: a scalar ($\gamma \geq 1$) giving the damping.
200
+ :param seed: a random seed.
201
+ :returns: a tuple of ([R, M] tensor of new event numbers,
202
+ log_acceptance_correction)
203
+ """
204
+ seeds = samplers.split_seed(seed, n=4, salt="_p_step_ps_leq_p")
205
+
206
+ # Sample z_new
207
+ z_new = tfd.Binomial(total_count=z, probs=ps / p).sample(seed=seeds[0])
208
+
209
+ # Jacobian
210
+ upper_bound = 2 * p - ps
211
+ damping_fn = DampingFunction(ps, upper_bound, gamma)
212
+
213
+ w = z - z_new
214
+ w_prime = tfd.Binomial(
215
+ total_count=x - z,
216
+ probs=(upper_bound - p) / (1 - p),
217
+ ).sample(seed=seeds[1])
218
+
219
+ def jacobian(x):
220
+ return damping_fn.log_inverse_jacobian(damping_fn(x))
221
+
222
+ log_jacobian_fwd = _reduce_uniforms(
223
+ ps,
224
+ p,
225
+ w,
226
+ jacobian,
227
+ seeds[2],
228
+ ) + _reduce_uniforms(p, upper_bound, w_prime, jacobian, seed=seeds[3])
229
+
230
+ return (
231
+ z_new,
232
+ -log_jacobian_fwd,
233
+ ) # log_acceptance_correction
234
+
235
+
236
+ def _pstep(z, x, p, ps, gamma, seed):
237
+ """Compute the p-step of Rippler.
238
+
239
+ Since there are two possible distributions to draw from,
240
+ but both are Binomial, we compute `offset`, `total_count`, and
241
+ `prob` parameters for both branches, and select which we need
242
+ based on p <= ps.
243
+
244
+ :param z: current $z$
245
+ :param x: current $x$
246
+ :param p: current probability
247
+ :param ps: new probability
248
+ """
249
+ with tf.name_scope("_pstep"):
250
+ seed1, seed2 = samplers.split_seed(seed, salt="_pstep")
251
+ ps_leq_p = _p_step_ps_leq_p(z, x, p, ps, gamma, seed1)
252
+ ps_gt_p = _p_step_ps_gt_p(z, x, p, ps, gamma, seed2)
253
+
254
+ z_prime = tf.where(ps <= p, ps_leq_p[0], ps_gt_p[0])
255
+ log_acceptance_correction = tf.where(ps <= p, ps_leq_p[1], ps_gt_p[1])
256
+
257
+ return z_prime, log_acceptance_correction
258
+
259
+
260
+ def _xstep(z_prime, x, xs, ps, seed, validate_args=False):
261
+ """Computes the x-step of the Rippler algorithm.
262
+
263
+ Both xs >= x and xs < x are sampled and results selected.
264
+ """
265
+ with tf.name_scope("_xstep"):
266
+ seeds = samplers.split_seed(seed, salt="_xstep")
267
+
268
+ # xs >= x
269
+ # Switch off `validate_args` because `xs-x` may be -ve.
270
+ z_new_geq = z_prime + tfd.Binomial(
271
+ xs - x,
272
+ probs=ps,
273
+ validate_args=False,
274
+ name="_xstep_Binomial",
275
+ ).sample(seed=seeds[0])
276
+
277
+ # xs < x - explicitly vectorize
278
+ def safe_hypergeom(N, K, n): # noqa: N803
279
+ # xs is clipped to min(x, xs) to avoid errors in the Hypergeometric
280
+ # sampler these values won't be selected anyway due to the
281
+ # xs >= x condition below.
282
+ return Hypergeometric(
283
+ N=N,
284
+ K=K,
285
+ n=tf.math.minimum(N, n),
286
+ validate_args=False,
287
+ name="_xstep_Hypergeom",
288
+ )
289
+
290
+ z_new_lt = safe_hypergeom(x, z_prime, xs).sample(seed=seeds[1])
291
+
292
+ return tf.where(xs >= x, z_new_geq, z_new_lt)
293
+
294
+
295
+ def _dispatch_update(z, x, p, xs, ps, gamma, seed, validate_args=False):
296
+ r"""Dispatches update function based on values of
297
+ parameters.
298
+
299
+ :param z: current $z$
300
+ :param x: current $x$
301
+ :param p: $p$ current probability
302
+ :param xs: $x^\star$ new state
303
+ :param ps: $p_star$ new probability
304
+
305
+ :returns: an updated number of events
306
+ """
307
+ with tf.name_scope("dispatch_update"):
308
+ p = tf.convert_to_tensor(p)
309
+ ps = tf.convert_to_tensor(ps)
310
+ z = tf.cast(z, p.dtype)
311
+ x = tf.cast(x, p.dtype)
312
+ xs = tf.cast(xs, p.dtype)
313
+
314
+ seeds = samplers.split_seed(seed, salt="_dispatch_update")
315
+
316
+ z_prime, log_acceptance_correction = _pstep(
317
+ z, x, p, ps, gamma, seed=seeds[0]
318
+ )
319
+ z_new = _xstep(
320
+ z_prime, x, xs, ps, seed=seeds[1], validate_args=validate_args
321
+ )
322
+
323
+ return z_new, log_acceptance_correction, ps > p
324
+
325
+
326
+ # Tests
327
+ def test_dispatch():
328
+ tf.debugging.assert_scalar(
329
+ _dispatch_update(z=10, x=100, p=0.1, xs=100, ps=0.1)
330
+ )
331
+ tf.debugging.assert_scalar(
332
+ _dispatch_update(z=10, x=100, p=0.1, xs=100, ps=0.05)
333
+ )
334
+ tf.debugging.assert_scalar(
335
+ _dispatch_update(z=10, x=100, p=0.1, xs=100, ps=0.2)
336
+ )
337
+ tf.debugging.assert_scalar(
338
+ _dispatch_update(z=10, x=100, p=0.1, xs=50, ps=0.1)
339
+ )
340
+ tf.debugging.assert_scalar(
341
+ _dispatch_update(z=10, x=100, p=0.1, xs=50, ps=0.05)
342
+ )
343
+ tf.debugging.assert_scalar(
344
+ _dispatch_update(z=10, x=100, p=0.1, xs=50, ps=0.2)
345
+ )
346
+ tf.debugging.assert_scalar(
347
+ _dispatch_update(z=10, x=100, p=0.1, xs=200, ps=0.1)
348
+ )
349
+ tf.debugging.assert_scalar(
350
+ _dispatch_update(z=10, x=100, p=0.1, xs=200, ps=0.05)
351
+ )
352
+ tf.debugging.assert_scalar(
353
+ _dispatch_update(z=10, x=100, p=0.1, xs=200, ps=0.2)
354
+ )
355
+
356
+
357
+ def default_initial_ripple(model, current_events, current_state, seed):
358
+ """Produces the initial ripple.
359
+
360
+ :param model: an instance of `DiscreteTimeStateTransitionModel`
361
+ :param current_events: a tensor of events in [T, R, M] order
362
+ :param current_state: a tensor of state in [T, S, M] order
363
+ :param seed: the seed to initialise the ripple
364
+
365
+ :returns: a tuple of `(proposed_time_idx, new_events_t, current_state_t)`
366
+ """
367
+ init_time_seed, init_pop_seed, init_events_seed = samplers.split_seed(
368
+ seed, n=3, salt="_initial_ripple"
369
+ )
370
+
371
+ # Choose timepoint, t
372
+ proposed_time_idx = UniformInteger(low=0, high=model.num_steps).sample(
373
+ seed=init_time_seed
374
+ )
375
+ current_state_t = tf.gather(current_state, proposed_time_idx, axis=-3)
376
+
377
+ # Choose subpopulation - KCategorical?
378
+ proposed_pop_idx = UniformInteger(
379
+ low=0, high=current_events.shape[-1]
380
+ ).sample(seed=init_pop_seed)
381
+
382
+ # Choose new infection events at time t
383
+ proposed_transition_rates = tf.stack(
384
+ model.transition_rates(
385
+ proposed_time_idx, tf.transpose(current_state_t)
386
+ ),
387
+ axis=0,
388
+ )
389
+ prob_t = 1.0 - tf.math.exp(
390
+ -tf.gather(proposed_transition_rates[0], proposed_pop_idx, axis=-1)
391
+ * model.time_delta,
392
+ ) # First event to perturb.
393
+
394
+ required_state = tf.gather(current_state_t[0], proposed_pop_idx, axis=-1)
395
+ new_si_events_t = tfd.Binomial(
396
+ total_count=required_state,
397
+ probs=prob_t, # Perturb SI events here
398
+ ).sample(seed=init_events_seed)
399
+
400
+ new_events_t = tf.tensor_scatter_nd_update(
401
+ current_events[proposed_time_idx],
402
+ [[0, proposed_pop_idx]],
403
+ [new_si_events_t],
404
+ )
405
+
406
+ return proposed_time_idx, new_events_t, current_state_t
407
+
408
+
409
+ def damped_initial_ripple_fn(resampling_fraction=1.0):
410
+ """Construct a damped initial ripple function
411
+
412
+ Args
413
+ ----
414
+ sampling_fraction: the fraction of the initial state to resample
415
+
416
+ Returns
417
+ -------
418
+ a callable fn(model, current_events, current_state, seed)
419
+ """
420
+
421
+ def fn(model, current_events, current_state, seed):
422
+ """Produces a damped initial ripple.
423
+
424
+ :param model: an instance of `DiscreteTimeStateTransitionModel`
425
+ :param current_events: a tensor of events in [T, R, M] order
426
+ :param current_state: a tensor of state in [T, S, M] order
427
+ :param resampling_fraction: a value between 0 and 1 where 0
428
+ is fully damped, and 1 is no damping
429
+ :param seed: the seed to initialise the ripple
430
+
431
+ :returns: a tuple of `(proposed_time_idx, new_events_t,
432
+ current_state_t)`
433
+ """
434
+
435
+ init_time_seed, init_pop_seed, hypergeom_seed, binom_seed = (
436
+ samplers.split_seed(seed, n=4, salt="_initial_ripple")
437
+ )
438
+
439
+ # Choose timepoint, t
440
+ proposed_time_idx = UniformInteger(low=0, high=model.num_steps).sample(
441
+ seed=init_time_seed
442
+ )
443
+ current_state_t = tf.gather(current_state, proposed_time_idx, axis=-3)
444
+
445
+ # Choose subpopulation - KCategorical?
446
+ proposed_pop_idx = UniformInteger(
447
+ low=0, high=current_events.shape[-1]
448
+ ).sample(seed=init_pop_seed)
449
+
450
+ # Choose new infection events at time t
451
+ proposed_transition_rates = tf.stack(
452
+ model.transition_rates(
453
+ proposed_time_idx, tf.transpose(current_state_t)
454
+ ),
455
+ axis=0,
456
+ )
457
+ prob_t = 1.0 - tf.math.exp(
458
+ -tf.gather(proposed_transition_rates[0], proposed_pop_idx, axis=-1)
459
+ * model.time_delta,
460
+ ) # First event to perturb.
461
+
462
+ required_state = tf.gather(
463
+ current_state_t[0], proposed_pop_idx, axis=-1
464
+ )
465
+ required_events = current_events[proposed_time_idx, 0, proposed_pop_idx]
466
+ sample_size = tf.math.floor(required_state * resampling_fraction)
467
+
468
+ new_si_events_t = (
469
+ required_events
470
+ - Hypergeometric(
471
+ required_state, required_events, sample_size
472
+ ).sample(seed=hypergeom_seed)
473
+ + tfd.Binomial(total_count=sample_size, probs=prob_t).sample(
474
+ seed=binom_seed
475
+ )
476
+ )
477
+
478
+ new_events_t = tf.tensor_scatter_nd_update(
479
+ current_events[proposed_time_idx],
480
+ [[0, proposed_pop_idx]],
481
+ [new_si_events_t],
482
+ )
483
+
484
+ return proposed_time_idx, new_events_t, current_state_t
485
+
486
+ return fn
487
+
488
+
489
+ def chain_binomial_rippler(
490
+ model,
491
+ current_events,
492
+ initial_ripple_fn,
493
+ ripple_damping_constant,
494
+ seed,
495
+ ):
496
+ init_seed, ripple_seed = samplers.split_seed(
497
+ seed, salt="chain_binomial_rippler"
498
+ )
499
+ src_states = model.source_states
500
+
501
+ # Transpose to [T, S/R, M]
502
+ current_events = tf.transpose(current_events, perm=(1, 2, 0))
503
+
504
+ # Calculate current state
505
+ current_state = _compute_state(
506
+ initial_state=tf.transpose(model.initial_state),
507
+ events=current_events,
508
+ stoichiometry=model.stoichiometry,
509
+ )
510
+
511
+ # Begin the ripple by sampling a time point, and perturbing the
512
+ # events at that timepoint
513
+ (
514
+ proposed_time_idx,
515
+ new_events_t,
516
+ current_state_t,
517
+ ) = initial_ripple_fn(model, current_events, current_state, init_seed)
518
+ new_events = tf.tensor_scatter_nd_update(
519
+ current_events, indices=[[proposed_time_idx]], updates=[new_events_t]
520
+ )
521
+
522
+ # Propagate from t+1 up to end of the timeseries
523
+ def draw_events(time, new_state_t, current_events_t, current_state_t, seed):
524
+ with tf.name_scope("draw_events"):
525
+ # Calculate transition rates for current and new states
526
+ def transition_probs(time, state):
527
+ rates = tf.stack(
528
+ model.transition_rates(time, tf.transpose(state)), axis=-2
529
+ )
530
+ return 1.0 - tf.math.exp(-rates * model.time_delta)
531
+
532
+ current_p = transition_probs(time, current_state_t)
533
+ new_p = transition_probs(time, new_state_t)
534
+
535
+ new_events, log_acceptance_correction, ps_gt_p = _dispatch_update(
536
+ z=current_events_t,
537
+ x=prefer_static.gather(current_state_t, indices=src_states),
538
+ p=current_p,
539
+ xs=prefer_static.gather(new_state_t, indices=src_states),
540
+ ps=new_p,
541
+ gamma=ripple_damping_constant,
542
+ seed=seed,
543
+ )
544
+ tf.debugging.assert_non_negative(new_events)
545
+
546
+ return new_events, log_acceptance_correction, ps_gt_p
547
+
548
+ def time_loop_body(
549
+ t,
550
+ new_events_t,
551
+ new_state_t,
552
+ new_events_buffer,
553
+ log_acceptance_correction_accum,
554
+ ps_gt_p_accum,
555
+ seed,
556
+ ):
557
+ sample_seed, next_seed = samplers.split_seed(
558
+ seed, salt="time_loop_body"
559
+ )
560
+
561
+ # Propagate new_state[t] to new_state[t+1]
562
+ new_state_t1 = new_state_t + tf.einsum(
563
+ "...ik,ij->...jk", new_events_t, model.stoichiometry
564
+ )
565
+ # tf.debugging.assert_non_negative(new_state_t1, summarize=100)
566
+
567
+ # Gather current states and events, and draw new events
568
+ new_events_t1, log_acceptance_correction, ps_gt_p = draw_events(
569
+ t + 1,
570
+ new_state_t1,
571
+ current_events[t + 1],
572
+ current_state[t + 1],
573
+ sample_seed,
574
+ )
575
+
576
+ # Update new_events_buffer
577
+ new_events_buffer = tf.tensor_scatter_nd_update(
578
+ new_events_buffer, indices=[[t + 1]], updates=[new_events_t1]
579
+ )
580
+
581
+ return (
582
+ t + 1,
583
+ new_events_t1,
584
+ new_state_t1,
585
+ new_events_buffer,
586
+ log_acceptance_correction_accum.write(t, log_acceptance_correction),
587
+ ps_gt_p_accum.write(t, ps_gt_p),
588
+ next_seed,
589
+ )
590
+
591
+ def time_loop_cond(t, _1, _2, new_events_buffer, *_3):
592
+ t_stop = t < (model.num_steps - 1)
593
+ delta_stop = tf.reduce_any(new_events_buffer != current_events)
594
+ return t_stop & delta_stop
595
+
596
+ log_acceptance_correction_accum = tf.TensorArray(
597
+ current_state.dtype, size=model.num_steps
598
+ )
599
+ ps_gt_p_accum = tf.TensorArray(tf.bool, size=model.num_steps)
600
+ (
601
+ _,
602
+ _,
603
+ _,
604
+ new_events,
605
+ log_acceptance_correction,
606
+ ps_gt_p_accum,
607
+ _,
608
+ ) = tf.while_loop(
609
+ time_loop_cond,
610
+ time_loop_body,
611
+ loop_vars=(
612
+ proposed_time_idx,
613
+ new_events_t,
614
+ current_state_t,
615
+ new_events,
616
+ log_acceptance_correction_accum,
617
+ ps_gt_p_accum,
618
+ ripple_seed,
619
+ ),
620
+ ) # new_events.shape = [T, R, M]
621
+
622
+ new_events = tf.transpose(new_events, perm=(2, 0, 1))
623
+
624
+ return (
625
+ new_events,
626
+ {
627
+ "log_acceptance_correction": log_acceptance_correction.stack(),
628
+ "is_ps_gt_p": ps_gt_p_accum.stack(),
629
+ "delta": tf.transpose(
630
+ new_events_t - current_events[proposed_time_idx]
631
+ ),
632
+ "timepoint": proposed_time_idx,
633
+ "initial_ripple": new_events_t,
634
+ "current_state_t": tf.transpose(current_state_t),
635
+ },
636
+ )
637
+
638
+
639
+ # The Chain Binomial Rippler kernel
640
+ CBRResults = namedtuple(
641
+ "CBRResults",
642
+ [
643
+ "target_log_prob",
644
+ "is_accepted",
645
+ "delta",
646
+ "current_state_t",
647
+ "initial_ripple",
648
+ "timepoint",
649
+ "proposed_state",
650
+ "proposed_target_log_prob",
651
+ "log_acceptance_correction",
652
+ "is_ps_gt_p",
653
+ "seed",
654
+ ],
655
+ )
656
+
657
+
658
+ class DampedCBRKernel(tfp.mcmc.TransitionKernel):
659
+ def __init__(
660
+ self,
661
+ target_log_prob_fn,
662
+ model,
663
+ initial_ripple_fn=default_initial_ripple,
664
+ ripple_damping_constant=1.0,
665
+ name=None,
666
+ ):
667
+ self._target_log_prob_fn = target_log_prob_fn
668
+ self._model = model
669
+
670
+ name = mcmc_util.make_name(name, "CBRKernel", "")
671
+
672
+ self._parameters = {
673
+ "target_log_prob_fn": target_log_prob_fn,
674
+ "model": model,
675
+ "initial_ripple_fn": initial_ripple_fn,
676
+ "ripple_damping_constant": ripple_damping_constant,
677
+ "name": name,
678
+ }
679
+
680
+ @property
681
+ def is_calibrated(self):
682
+ return True
683
+
684
+ @property
685
+ def target_log_prob(self):
686
+ return self._target_log_prob_fn
687
+
688
+ @property
689
+ def model(self):
690
+ return self._model
691
+
692
+ @property
693
+ def name(self):
694
+ return self._parameters["name"]
695
+
696
+ @property
697
+ def initial_ripple_fn(self):
698
+ return self._parameters["initial_ripple_fn"]
699
+
700
+ @property
701
+ def ripple_damping_constant(self):
702
+ return self._parameters["ripple_damping_constant"]
703
+
704
+ def one_step(self, current_state, previous_results, seed=None):
705
+ with tf.name_scope("CBRKernel/one_step"):
706
+ seed_rippler, seed_u, seed_results = samplers.split_seed(
707
+ seed, n=3, salt="cbr_kernel"
708
+ )
709
+
710
+ if mcmc_util.is_list_like(current_state):
711
+ current_state_parts = list(current_state)
712
+ else:
713
+ current_state_parts = [current_state]
714
+
715
+ if len(current_state_parts) > 1:
716
+ warnings.warn(
717
+ "CBRKernel.boostrap_results: multiple state parts detected,\
718
+ but only the first will be used",
719
+ stacklevel=2,
720
+ )
721
+
722
+ current_state_part = tf.convert_to_tensor(
723
+ current_state_parts[0], name="current_state"
724
+ )
725
+
726
+ proposed_state, proposal_results = chain_binomial_rippler(
727
+ self.model,
728
+ current_state_part,
729
+ initial_ripple_fn=self.initial_ripple_fn,
730
+ ripple_damping_constant=self.ripple_damping_constant,
731
+ seed=seed_rippler,
732
+ )
733
+
734
+ proposed_target_log_prob = self.target_log_prob(proposed_state)
735
+
736
+ delta_logp = (
737
+ proposed_target_log_prob
738
+ - previous_results.target_log_prob
739
+ + tf.reduce_sum(proposal_results["log_acceptance_correction"])
740
+ )
741
+
742
+ def accept():
743
+ return (
744
+ proposed_state,
745
+ CBRResults(
746
+ target_log_prob=proposed_target_log_prob,
747
+ is_accepted=tf.constant(True),
748
+ delta=proposal_results["delta"],
749
+ current_state_t=proposal_results["current_state_t"],
750
+ initial_ripple=proposal_results["initial_ripple"],
751
+ timepoint=proposal_results["timepoint"],
752
+ proposed_state=proposed_state,
753
+ proposed_target_log_prob=proposed_target_log_prob,
754
+ log_acceptance_correction=proposal_results[
755
+ "log_acceptance_correction"
756
+ ],
757
+ is_ps_gt_p=proposal_results["is_ps_gt_p"],
758
+ seed=seed_results,
759
+ ),
760
+ )
761
+
762
+ def reject():
763
+ return (
764
+ current_state_part,
765
+ CBRResults(
766
+ target_log_prob=previous_results.target_log_prob,
767
+ is_accepted=tf.constant(False),
768
+ delta=proposal_results["delta"],
769
+ current_state_t=proposal_results["current_state_t"],
770
+ initial_ripple=proposal_results["initial_ripple"],
771
+ timepoint=proposal_results["timepoint"],
772
+ proposed_state=proposed_state,
773
+ proposed_target_log_prob=proposed_target_log_prob,
774
+ log_acceptance_correction=proposal_results[
775
+ "log_acceptance_correction"
776
+ ],
777
+ is_ps_gt_p=proposal_results["is_ps_gt_p"],
778
+ seed=seed_results,
779
+ ),
780
+ )
781
+
782
+ u = tf.math.log(
783
+ tfd.Uniform(low=tf.zeros(1, dtype=delta_logp.dtype)).sample(
784
+ seed=seed_u
785
+ )
786
+ )
787
+ new_state, results = tf.cond(u < delta_logp, accept, reject)
788
+
789
+ def maybe_flatten(x):
790
+ if mcmc_util.is_list_like(current_state):
791
+ return type(current_state)(new_state)
792
+ return x
793
+
794
+ new_state = maybe_flatten(new_state)
795
+ return new_state, results
796
+
797
+ def bootstrap_results(self, current_state):
798
+ with tf.name_scope("CBRKernel/bootstrap_results"):
799
+ if mcmc_util.is_list_like(current_state):
800
+ current_state_parts = list(current_state)
801
+ else:
802
+ current_state_parts = [current_state]
803
+
804
+ if len(current_state_parts) > 1:
805
+ warnings.warn(
806
+ "CBRKernel.boostrap_results: multiple state parts detected,\
807
+ but only the first will be used",
808
+ stacklevel=2,
809
+ )
810
+ state_part = current_state_parts[0]
811
+
812
+ num_times = state_part.shape[-2]
813
+ num_pop = state_part.shape[-3]
814
+ num_transitions = state_part.shape[-1]
815
+ num_states = self.model.stoichiometry.shape[-1]
816
+
817
+ target_log_prob = self.target_log_prob(state_part)
818
+
819
+ return CBRResults(
820
+ target_log_prob=target_log_prob,
821
+ is_accepted=tf.constant(False),
822
+ delta=tf.zeros((num_pop, num_transitions), state_part.dtype),
823
+ current_state_t=tf.zeros(
824
+ [num_pop, num_states], state_part.dtype
825
+ ),
826
+ initial_ripple=tf.zeros(
827
+ (num_transitions, num_pop), state_part.dtype
828
+ ),
829
+ timepoint=tf.constant(0, dtype=tf.int32),
830
+ proposed_state=state_part,
831
+ proposed_target_log_prob=target_log_prob,
832
+ log_acceptance_correction=tf.zeros(
833
+ (num_times, num_transitions, num_pop),
834
+ current_state[0].dtype,
835
+ ),
836
+ is_ps_gt_p=tf.fill(
837
+ (num_times, num_transitions, num_pop), False
838
+ ),
839
+ seed=samplers.sanitize_seed(0),
840
+ )