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,505 @@
1
+ """Gibbs sampling kernel"""
2
+ # ruff: noqa: B023
3
+
4
+ import logging
5
+ from collections import namedtuple
6
+ from typing import Callable, List, Tuple
7
+
8
+ import tensorflow as tf
9
+ import tensorflow_probability as tfp
10
+ from tensorflow_probability.python.internal import samplers, unnest
11
+ from tensorflow_probability.python.mcmc.internal import util as mcmc_util
12
+
13
+ tfd = tfp.distributions # pylint: disable=no-member
14
+ tfb = tfp.bijectors # pylint: disable=no-member
15
+ mcmc = tfp.mcmc # pylint: disable=no-member
16
+
17
+ logging.basicConfig(format="[%(asctime)s] %(levelname)s %(message)s")
18
+ logger = logging.getLogger("gemlib.gibbs_kernel")
19
+ logger.setLevel(logging.DEBUG)
20
+
21
+
22
+ class GibbsKernelResults(
23
+ mcmc_util.PrettyNamedTupleMixin,
24
+ namedtuple(
25
+ "GibbsKernelResults",
26
+ ["target_log_prob", "inner_results", "seed"],
27
+ ),
28
+ ):
29
+ """Represents kernel results"""
30
+
31
+ __slots__ = ()
32
+
33
+
34
+ class GibbsStep(
35
+ mcmc_util.PrettyNamedTupleMixin,
36
+ namedtuple(
37
+ "GibbsStep",
38
+ ["state_parts", "kernel_fn"],
39
+ ),
40
+ ):
41
+ """Represents a Gibbs step"""
42
+
43
+ __slots__ = ()
44
+
45
+
46
+ def is_namedtuple(x):
47
+ """Return `True` if `x` looks like a `namedtuple`."""
48
+ return hasattr(x, "_fields")
49
+
50
+
51
+ def _flatten_results(results):
52
+ """Results structures from nested Gibbs samplers sometimes
53
+ need flattening for writing out purposes.
54
+ """
55
+
56
+ def recurse(r):
57
+ for i in iter(r):
58
+ if isinstance(i, list):
59
+ yield from _flatten_results(i)
60
+ else:
61
+ yield i
62
+
63
+ return list(recurse(results))
64
+
65
+
66
+ def _has_gradients(results):
67
+ return unnest.has_nested(results, "grads_target_log_prob")
68
+
69
+
70
+ def _get_target_log_prob(results):
71
+ """Fetches a target log prob from a results structure"""
72
+ return unnest.get_innermost(results, "target_log_prob")
73
+
74
+
75
+ def _update_target_log_prob(results, target_log_prob):
76
+ """Puts a target log prob into a results structure"""
77
+ if isinstance(results, GibbsKernelResults):
78
+ replace_fn = unnest.replace_outermost
79
+ else:
80
+ replace_fn = unnest.replace_innermost
81
+ return replace_fn(results, target_log_prob=target_log_prob)
82
+
83
+
84
+ def _maybe_transform_value(tlp, state, kernel, direction):
85
+ if not isinstance(kernel, tfp.mcmc.TransformedTransitionKernel):
86
+ return tlp
87
+
88
+ jacobian_parts = [
89
+ b.inverse_log_det_jacobian(x)
90
+ for b, x in zip(
91
+ tf.nest.flatten(kernel.bijector), tf.nest.flatten(state)
92
+ )
93
+ ]
94
+ jacobian = tf.math.add_n(jacobian_parts)
95
+
96
+ if direction == "forward":
97
+ return tlp + jacobian
98
+ if direction == "inverse":
99
+ return tlp - jacobian
100
+
101
+ raise AttributeError("`direction` must be `forward` or `inverse`")
102
+
103
+
104
+ def _make_namedtuple(input_dict):
105
+ return namedtuple("NamedTuple", input_dict.keys())(**input_dict)
106
+
107
+
108
+ def _split_namedtuple(full_namedtuple, subset_names):
109
+ """Splits a StructTuple of variables into `subset` and `compl`
110
+
111
+ :param full_struct_tuple: `namedtuple` to split
112
+ :param subset_names: names of required subset vars
113
+ :returns: a tuple `(subset: namedtuple, compl: namedtuple)`
114
+ """
115
+ full_dict = full_namedtuple._asdict()
116
+ subset = _make_namedtuple(
117
+ {k: v for k, v in full_dict.items() if k in subset_names}
118
+ )
119
+ compl = _make_namedtuple(
120
+ {k: v for k, v in full_dict.items() if k not in subset_names}
121
+ )
122
+ return subset, compl
123
+
124
+
125
+ def _split_state(global_state, indices):
126
+ """Split a global state into subset and complement
127
+
128
+ Args:
129
+ ----
130
+ global_state: a tuple or namedtuple representing the global state
131
+ subset: a tuple of indices or names (if `global_state` is a namedtuple)
132
+ representing a subset.
133
+
134
+ Returns:
135
+ -------
136
+ a tuple of `(subset, complement)`
137
+
138
+ """
139
+ if is_namedtuple(global_state):
140
+ return _split_namedtuple(global_state, indices)
141
+
142
+ return (
143
+ [s for i, s in enumerate(global_state) if i in indices],
144
+ [s for i, s in enumerate(global_state) if i not in indices],
145
+ )
146
+
147
+
148
+ def _scatter_state(global_state, subset, indices=()):
149
+ """Scatters `subset` into `global_state`.
150
+
151
+ Args:
152
+ ----
153
+ global_state: a tuple or namedtuple representing the global state
154
+ subset: a tuple or namedtuple containing values to be scattered into
155
+ `global_state`.
156
+ indices: if `subset` is a tuple, `indices` is a tuple of corresponding
157
+ indices into `global_state`.
158
+
159
+ Returns:
160
+ -------
161
+ a tuple or namedtuple of the same structure as `global_state`.
162
+
163
+ """
164
+ if is_namedtuple(global_state) and is_namedtuple(subset):
165
+ return global_state._replace(**subset._asdict())
166
+
167
+ for i, state_part in zip(indices, subset):
168
+ global_state[i] = state_part
169
+
170
+ return global_state
171
+
172
+
173
+ class GibbsKernel(mcmc.TransitionKernel):
174
+ """Component-wise MCMC sampling.
175
+
176
+ ``GibbsKernel`` is designed to fit within TensorFlow Probability's MCMC
177
+ framework, essentially acting as a "meta-kernel" that aggregates a
178
+ sequence of component-wise kernels.
179
+
180
+ Example:
181
+ -------
182
+ Sample from the posterior of a linear model::
183
+
184
+ import numpy as np
185
+ import tensorflow as tf
186
+ import tensorflow_probability as tfp
187
+ from gemlib.mcmc.gibbs_kernel import GibbsKernel
188
+
189
+ tfd = tfp.distributions
190
+
191
+ dtype = np.float32
192
+
193
+ # data
194
+ x = dtype([2.9, 4.2, 8.3, 1.9, 2.6, 1.0, 8.4, 8.6, 7.9, 4.3])
195
+ y = dtype([6.2, 7.8, 8.1, 2.7, 4.8, 2.4, 10.7, 9.0, 9.6, 5.7])
196
+
197
+
198
+ # define linear regression model
199
+ def Model(x):
200
+ def alpha():
201
+ return tfd.Normal(loc=dtype(0.0), scale=dtype(1000.0))
202
+
203
+ def beta():
204
+ return tfd.Normal(loc=dtype(0.0), scale=dtype(100.0))
205
+
206
+ def sigma():
207
+ return tfd.Gamma(concentration=dtype(0.1), rate=dtype(0.1))
208
+
209
+ def y(alpha, beta, sigma):
210
+ mu = alpha + beta * x
211
+ return tfd.Normal(mu, scale=sigma)
212
+
213
+ return tfd.JointDistributionNamed(
214
+ dict(alpha=alpha, beta=beta, sigma=sigma, y=y)
215
+ )
216
+
217
+
218
+ # target log probability of linear model
219
+ def log_prob(alpha, beta, sigma):
220
+ lp = model.log_prob(
221
+ {"alpha": alpha, "beta": beta, "sigma": sigma, "y": y}
222
+ )
223
+ return tf.reduce_sum(lp)
224
+
225
+
226
+ # random walk Markov chain function
227
+ def kernel_make_fn(target_log_prob_fn, state):
228
+ return tfp.mcmc.RandomWalkMetropolis(
229
+ target_log_prob_fn=target_log_prob_fn
230
+ )
231
+
232
+
233
+ # posterior distribution MCMC chain
234
+ @tf.function
235
+ def posterior(iterations, burnin, thinning, initial_state):
236
+ kernel_list = [
237
+ (
238
+ 0,
239
+ kernel_make_fn,
240
+ ), # conditional probability for zeroth parmeter alpha
241
+ (
242
+ 1,
243
+ kernel_make_fn,
244
+ ), # conditional probability for first parameter beta
245
+ (2, kernel_make_fn),
246
+ ] # conditional probability for second parameter sigma
247
+ kernel = GibbsKernel(
248
+ target_log_prob_fn=log_prob, kernel_list=kernel_list
249
+ )
250
+ return tfp.mcmc.sample_chain(
251
+ num_results=iterations,
252
+ current_state=initial_state,
253
+ kernel=kernel,
254
+ num_burnin_steps=burnin,
255
+ num_steps_between_results=thinning,
256
+ parallel_iterations=1,
257
+ trace_fn=None,
258
+ )
259
+
260
+
261
+ # initialize model
262
+ model = Model(x)
263
+ initial_state = [
264
+ dtype(0.1),
265
+ dtype(0.1),
266
+ dtype(0.1),
267
+ ] # start chain at alpha=0.1, beta=0.1, sigma=0.1
268
+
269
+ # estimate posterior distribution
270
+ samples = posterior(
271
+ iterations=10000,
272
+ burnin=1000,
273
+ thinning=0,
274
+ initial_state=initial_state,
275
+ )
276
+
277
+ tf.print("alpha samples:", samples[0])
278
+ tf.print("beta samples:", samples[1])
279
+ tf.print("sigma samples:", samples[2])
280
+ tf.print(
281
+ "sample means: [alpha, beta, sigma] =",
282
+ tf.math.reduce_mean(samples, axis=1),
283
+ )
284
+
285
+ """
286
+
287
+ def __init__(
288
+ self,
289
+ target_log_prob_fn: Callable[[float], float],
290
+ kernel_list: List[Tuple[Tuple[int, ...], Callable]],
291
+ name: str = None,
292
+ ):
293
+ """Build a Gibbs sampling scheme from component kernels.
294
+
295
+ Args:
296
+ ----
297
+ target_log_prob_fn: a function that takes `state` arguments
298
+ and returns the target log probability
299
+ density.
300
+ kernel_list: a list of tuples `(state_part_idx, kernel_make_fn)`.
301
+ `state_part_idx` denotes the index (relative to
302
+ positional args in `target_log_prob_fn`) of the
303
+ state the kernel updates. `kernel_make_fn` takes
304
+ arguments `target_log_prob_fn` and `state`,
305
+ returning a `tfp.mcmc.TransitionKernel`.
306
+
307
+ Returns:
308
+ -------
309
+ an instance of `GibbsKernel`
310
+
311
+ """
312
+ # Require to check if all kernel.is_calibrated is True
313
+ self._parameters = {
314
+ "target_log_prob_fn": target_log_prob_fn,
315
+ "kernel_list": kernel_list,
316
+ "name": name,
317
+ }
318
+
319
+ @property
320
+ def is_calibrated(self):
321
+ return True
322
+
323
+ @property
324
+ def target_log_prob_fn(self):
325
+ """Target log probability function."""
326
+ return self._parameters["target_log_prob_fn"]
327
+
328
+ @property
329
+ def kernel_list(self):
330
+ """List of kernel-build functions."""
331
+ return self._parameters["kernel_list"]
332
+
333
+ @property
334
+ def name(self):
335
+ """Name of the kernel."""
336
+ return self._parameters["name"]
337
+
338
+ def one_step(self, current_state, previous_results, seed=None):
339
+ """Iterate over the state elements, calling each kernel in turn.
340
+
341
+ The ``target_log_prob`` is forwarded to the next ``previous_results``
342
+ such that each kernel has a current ``target_log_prob`` value.
343
+ Transformations are automatically performed if the kernel is of
344
+ type ``tfp.mcmc.TransformedTransitionKernel``.
345
+
346
+ In graph and XLA modes, the for loop should be unrolled.
347
+
348
+ Args:
349
+ ----
350
+ current_state: the current chain state
351
+ previous_results: a ``GibbsKernelResults`` instance
352
+ seed: an optional list of two scalar ``int`` tensors.
353
+
354
+ Returns:
355
+ -------
356
+ a tuple of ``(next_state, results)``.
357
+
358
+ """
359
+ seed = samplers.sanitize_seed(seed, salt="GibbsKernel")
360
+
361
+ global_state_parts = current_state
362
+
363
+ next_results = []
364
+ untransformed_target_log_prob = previous_results.target_log_prob
365
+ seeds = samplers.split_seed(seed, n=len(self.kernel_list))
366
+
367
+ for (state_part_indices, kernel_fn), previous_step_results, seed in zip(
368
+ self.kernel_list, previous_results.inner_results, seeds
369
+ ):
370
+ if not mcmc_util.is_list_like(state_part_indices):
371
+ state_part_indices = [state_part_indices] # noqa: PLW2901
372
+
373
+ # Extract state parts required for step
374
+ step_state_parts, _ = _split_state(
375
+ global_state_parts, state_part_indices
376
+ )
377
+
378
+ def target_log_prob_fn(*kernel_state):
379
+ if is_namedtuple(step_state_parts):
380
+ kernel_state = step_state_parts.__class__(*kernel_state)
381
+ state_parts = _scatter_state(
382
+ global_state_parts,
383
+ kernel_state,
384
+ state_part_indices,
385
+ )
386
+ return self.target_log_prob_fn(*state_parts)
387
+
388
+ # Build kernel function
389
+ kernel = kernel_fn(target_log_prob_fn, global_state_parts)
390
+
391
+ # Forward the current tlp to the kernel. If the kernel is
392
+ # gradient-based, we need to calculate fresh gradients,
393
+ # as these cannot easily be forwarded
394
+ # from the previous Gibbs step.
395
+ if _has_gradients(previous_step_results):
396
+ # TODO would be better to avoid re-calculating the whole of
397
+ # `bootstrap_results` when we just need to calculate gradients.
398
+ fresh_previous_results = unnest.UnnestingWrapper(
399
+ kernel.bootstrap_results(step_state_parts)
400
+ )
401
+ previous_step_results = unnest.replace_innermost( # noqa: PLW2901
402
+ previous_step_results,
403
+ target_log_prob=fresh_previous_results.target_log_prob,
404
+ grads_target_log_prob=fresh_previous_results.grads_target_log_prob,
405
+ )
406
+
407
+ else:
408
+ previous_step_results = _update_target_log_prob( # noqa: PLW2901
409
+ previous_step_results,
410
+ _maybe_transform_value(
411
+ tlp=untransformed_target_log_prob,
412
+ state=step_state_parts,
413
+ kernel=kernel,
414
+ direction="inverse",
415
+ ),
416
+ )
417
+
418
+ new_step_state_parts, next_kernel_results = kernel.one_step(
419
+ step_state_parts, previous_step_results, seed
420
+ )
421
+ if is_namedtuple(step_state_parts):
422
+ new_step_state_parts = step_state_parts.__class__(
423
+ *new_step_state_parts,
424
+ )
425
+
426
+ next_results.append(next_kernel_results)
427
+
428
+ # Cache the new tlp for use in the next Gibbs step
429
+ untransformed_target_log_prob = _maybe_transform_value(
430
+ tlp=_get_target_log_prob(next_kernel_results),
431
+ state=new_step_state_parts,
432
+ kernel=kernel,
433
+ direction="forward",
434
+ )
435
+
436
+ global_state_parts = _scatter_state(
437
+ global_state_parts, new_step_state_parts, state_part_indices
438
+ )
439
+
440
+ if is_namedtuple(current_state):
441
+ global_state_parts = current_state.__class__(
442
+ *global_state_parts,
443
+ )
444
+
445
+ return (
446
+ global_state_parts
447
+ if mcmc_util.is_list_like(current_state)
448
+ else global_state_parts[0],
449
+ GibbsKernelResults(
450
+ target_log_prob=untransformed_target_log_prob,
451
+ inner_results=next_results,
452
+ seed=seeds[-1],
453
+ ),
454
+ )
455
+
456
+ def bootstrap_results(self, current_state):
457
+ """Set up the results tuple.
458
+
459
+ Args:
460
+ ----
461
+ current_state: a list of state parts representing the Markov chain
462
+ state
463
+ Returns:
464
+ an instance of `GibbsKernelResults`
465
+
466
+ """
467
+ global_state_parts = current_state
468
+ inner_results = []
469
+ untransformed_target_log_prob = 0.0
470
+
471
+ for state_part_indices, kernel_fn in self.kernel_list:
472
+ if not mcmc_util.is_list_like(state_part_indices):
473
+ state_part_indices = [state_part_indices] # noqa: PLW2901
474
+
475
+ step_state_parts, _ = _split_state(
476
+ global_state_parts, state_part_indices
477
+ )
478
+
479
+ def tlp_fn(*kernel_state):
480
+ if is_namedtuple(step_state_parts):
481
+ kernel_state = step_state_parts.__class__(*kernel_state)
482
+ state_parts = _scatter_state(
483
+ global_state_parts,
484
+ kernel_state,
485
+ state_part_indices,
486
+ )
487
+
488
+ return self.target_log_prob_fn(*state_parts)
489
+
490
+ kernel = kernel_fn(tlp_fn, global_state_parts)
491
+ kernel_results = kernel.bootstrap_results(step_state_parts)
492
+
493
+ inner_results.append(kernel_results)
494
+ untransformed_target_log_prob = _maybe_transform_value(
495
+ tlp=_get_target_log_prob(kernel_results),
496
+ state=step_state_parts,
497
+ kernel=kernel,
498
+ direction="forward",
499
+ )
500
+
501
+ return GibbsKernelResults(
502
+ target_log_prob=untransformed_target_log_prob,
503
+ inner_results=inner_results,
504
+ seed=samplers.zeros_seed(),
505
+ )