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,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
|
+
)
|
gemlib/mcmc/conftest.py
ADDED