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,156 @@
1
+ """Gibbs sampling kernel"""
2
+
3
+ from collections import namedtuple
4
+
5
+ import tensorflow_probability as tfp
6
+ from tensorflow_probability.python.internal import (
7
+ samplers,
8
+ structural_tuple,
9
+ unnest,
10
+ )
11
+ from tensorflow_probability.python.mcmc.internal import util as mcmc_util
12
+
13
+ __all__ = [
14
+ "CompoundKernel",
15
+ "split_unpinned_by_name",
16
+ "unpinned_and_conditional_model",
17
+ ]
18
+
19
+
20
+ class CompoundKernelResults(
21
+ mcmc_util.PrettyNamedTupleMixin,
22
+ namedtuple(
23
+ "CompoundKernelResults",
24
+ ["inner_results", "seed"],
25
+ ),
26
+ ):
27
+ __slots__ = ()
28
+
29
+
30
+ def _make_namedtuple(input_dict):
31
+ return structural_tuple.structtuple(input_dict.keys())(**input_dict)
32
+
33
+
34
+ def split_unpinned_by_name(full_struct_tuple, unpinned_names):
35
+ """Splits a StructTuple of variables into `unpinned` and `pinned`
36
+
37
+ :param full_struct_tuple: StructTuple to split
38
+ :param unpinned_names: names of required unpinned vars
39
+ :returns: a tuple `(unpinned: dict, pinned: dict)`
40
+ """
41
+ full_dict = full_struct_tuple._asdict()
42
+ unpinned = _make_namedtuple(
43
+ {k: v for k, v in full_dict.items() if k in unpinned_names}
44
+ )
45
+ pinned = _make_namedtuple(
46
+ {k: v for k, v in full_dict.items() if k not in unpinned_names}
47
+ )
48
+ return unpinned, pinned
49
+
50
+
51
+ def unpinned_and_conditional_model(varnames, vars, joint_model):
52
+ unpinned, pins = split_unpinned_by_name(vars, varnames)
53
+ return unpinned, joint_model.experimental_pin(pins)
54
+
55
+
56
+ def _replace_tlp(current_results, other_results):
57
+ """Replaces tlp in `current_results` with that in `other_results`"""
58
+ other_results_wrapped = unnest.UnnestingWrapper(other_results)
59
+
60
+ return unnest.replace_innermost(
61
+ current_results, target_log_prob=other_results_wrapped.target_log_prob
62
+ )
63
+
64
+
65
+ def _maybe_replace_grads(current_results, other_results):
66
+ """Replaces grads in `current_results` with that in `other_results`"""
67
+ other_results_wrapped = unnest.UnnestingWrapper(other_results)
68
+ if hasattr(other_results_wrapped, "grads_target_log_prob"):
69
+ return unnest.replace_innermost(
70
+ current_results,
71
+ grads_target_log_prob=other_results_wrapped.grads_target_log_prob,
72
+ )
73
+ return current_results
74
+
75
+
76
+ class CompoundKernel(tfp.mcmc.TransitionKernel):
77
+ class Step(namedtuple("Step", ["varnames", "build"])):
78
+ """Represents a Step within the CompoundKernel"""
79
+
80
+ __slots__ = ()
81
+
82
+ def __init__(self, joint_model, kernels, name=None):
83
+ self._parameters = locals()
84
+
85
+ @property
86
+ def is_calibrated(self):
87
+ return True
88
+
89
+ @property
90
+ def joint_model(self):
91
+ return self._parameters["joint_model"]
92
+
93
+ @property
94
+ def kernels(self):
95
+ return self._parameters["kernels"]
96
+
97
+ @property
98
+ def name(self):
99
+ return self._parameters["name"]
100
+
101
+ def one_step(self, current_state, previous_results, seed=None):
102
+ seeds = samplers.split_seed(
103
+ seed, n=len(self.kernels), salt="CompoundKernel.one_step"
104
+ )
105
+
106
+ next_results = []
107
+ for kernel_tuple, results, seed in zip(
108
+ self.kernels, previous_results.inner_results, seeds
109
+ ):
110
+ # Create the sub-state and conditioned model to
111
+ # then present to the kernel builder function.
112
+ unpinned, conditional_model = unpinned_and_conditional_model(
113
+ kernel_tuple.varnames, current_state, self.joint_model
114
+ )
115
+ kernel = kernel_tuple.build(conditional_model, current_state)
116
+
117
+ # In a Gibbs scheme, we have to re-calculate the current value
118
+ # (and grads) of the conditional log prob, pushing it back into
119
+ # the previous kernel results.
120
+ #
121
+ # A nested CompoundKernel does not have a target_log_prob or grads,
122
+ # so we can't do any replacement. Fortunately this doesn't matter,
123
+ # as one_step gets called recursively anyway.
124
+ if not isinstance(results, CompoundKernelResults):
125
+ pre_results = kernel.bootstrap_results(unpinned)
126
+ step_results = _replace_tlp(results, pre_results)
127
+ step_results = _maybe_replace_grads(step_results, pre_results)
128
+
129
+ new_unpinned, next_kernel_results = kernel.one_step(
130
+ unpinned, step_results, seed=seed
131
+ )
132
+
133
+ # Update the global state
134
+ current_state = current_state._replace(**new_unpinned._asdict())
135
+ next_results.append(next_kernel_results)
136
+
137
+ return (
138
+ current_state,
139
+ CompoundKernelResults(inner_results=tuple(next_results), seed=seed),
140
+ )
141
+
142
+ def bootstrap_results(self, current_state):
143
+ results = []
144
+ for kernel_tuple in self.kernels:
145
+ unpinned, pins = split_unpinned_by_name(
146
+ current_state, kernel_tuple.varnames
147
+ )
148
+ unpinned, conditional_model = unpinned_and_conditional_model(
149
+ kernel_tuple.varnames, current_state, self.joint_model
150
+ )
151
+ kernel = kernel_tuple.build(conditional_model, current_state)
152
+ results.append(kernel.bootstrap_results(unpinned))
153
+
154
+ return CompoundKernelResults(
155
+ inner_results=tuple(results), seed=samplers.zeros_seed()
156
+ )
@@ -0,0 +1,4 @@
1
+ """Pytest config"""
2
+
3
+
4
+ # Fixtures