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,214 @@
|
|
|
1
|
+
"""Test the random walk metropolis kernel"""
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import tensorflow as tf
|
|
5
|
+
import tensorflow_probability as tfp
|
|
6
|
+
|
|
7
|
+
from .composable_kernel import Step
|
|
8
|
+
from .mcmc_sampler import mcmc
|
|
9
|
+
from .random_walk_metropolis import RwmhInfo, rwmh
|
|
10
|
+
|
|
11
|
+
NUM_SAMPLES = 100000
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def split_seed(seed, n):
|
|
15
|
+
n = tf.convert_to_tensor(n)
|
|
16
|
+
return tfp.random.split_seed(seed, n=n)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def tree_map(fn, *args):
|
|
20
|
+
return tf.nest.map_structure(fn, *args)
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def tree_flatten(tree):
|
|
24
|
+
return tf.nest.flatten(tree)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def get_seed():
|
|
28
|
+
# jax.random.PRNGKey(42)
|
|
29
|
+
return [0, 0]
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
@tfp.distributions.JointDistributionCoroutine
|
|
33
|
+
def simple_model():
|
|
34
|
+
yield tfp.distributions.Normal(loc=0.0, scale=1.0, name="foo")
|
|
35
|
+
yield tfp.distributions.Normal(loc=1.0, scale=1.0, name="bar")
|
|
36
|
+
yield tfp.distributions.Normal(loc=2.0, scale=1.0, name="baz")
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def test_rwmh_1kernel():
|
|
40
|
+
seed = get_seed()
|
|
41
|
+
|
|
42
|
+
initial_position = simple_model.sample(seed=seed)
|
|
43
|
+
|
|
44
|
+
kernel = rwmh(scale=0.3)
|
|
45
|
+
|
|
46
|
+
state = kernel.init(simple_model.log_prob, initial_position)
|
|
47
|
+
new_state, results = kernel.step(simple_model.log_prob, state, seed)
|
|
48
|
+
|
|
49
|
+
assert tree_map(lambda x, y: None, new_state, state)
|
|
50
|
+
|
|
51
|
+
expected_results = (RwmhInfo(is_accepted=True),)
|
|
52
|
+
assert tree_map(lambda x, y: x == y, results, expected_results)
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def test_rwmh_2kernel():
|
|
56
|
+
seed = get_seed()
|
|
57
|
+
|
|
58
|
+
initial_position = simple_model.sample(seed=seed)
|
|
59
|
+
|
|
60
|
+
kernel = Step(rwmh(scale=0.3), ["foo"]) >> Step(
|
|
61
|
+
rwmh(scale=0.1), ["bar", "baz"]
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
state = kernel.init(simple_model.log_prob, initial_position)
|
|
65
|
+
new_state, results = kernel.step(simple_model.log_prob, state, seed)
|
|
66
|
+
|
|
67
|
+
assert tree_map(lambda x, y: None, new_state, state)
|
|
68
|
+
|
|
69
|
+
expected_results = (RwmhInfo(is_accepted=True), RwmhInfo(is_accepted=True))
|
|
70
|
+
assert tree_map(lambda x, y: x == y, results, expected_results)
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def test_rwmh_3kernel():
|
|
74
|
+
seed = get_seed()
|
|
75
|
+
|
|
76
|
+
initial_position = simple_model.sample(seed=seed)
|
|
77
|
+
|
|
78
|
+
kernel = (
|
|
79
|
+
Step(rwmh(scale=0.3), ["foo"])
|
|
80
|
+
>> Step(rwmh(scale=0.1), ["bar"])
|
|
81
|
+
>> Step(rwmh(scale=0.2), ["baz"])
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
state = kernel.init(simple_model.log_prob, initial_position)
|
|
85
|
+
new_state, results = kernel.step(simple_model.log_prob, state, seed)
|
|
86
|
+
|
|
87
|
+
assert tree_map(lambda x, y: None, new_state, state)
|
|
88
|
+
|
|
89
|
+
expected_results = (
|
|
90
|
+
RwmhInfo(is_accepted=True),
|
|
91
|
+
RwmhInfo(is_accepted=True),
|
|
92
|
+
RwmhInfo(is_accepted=True),
|
|
93
|
+
)
|
|
94
|
+
assert tree_map(lambda x, y: x == y, results, expected_results)
|
|
95
|
+
|
|
96
|
+
|
|
97
|
+
def test_rwmh_1kernel_mcmc():
|
|
98
|
+
seed = get_seed()
|
|
99
|
+
|
|
100
|
+
initial_position = simple_model.sample(seed=seed)
|
|
101
|
+
|
|
102
|
+
kernel = rwmh(scale=1.8)
|
|
103
|
+
|
|
104
|
+
posterior, info = tf.function(
|
|
105
|
+
lambda: mcmc(
|
|
106
|
+
NUM_SAMPLES,
|
|
107
|
+
sampling_algorithm=kernel,
|
|
108
|
+
target_density_fn=simple_model.log_prob,
|
|
109
|
+
initial_position=initial_position,
|
|
110
|
+
seed=get_seed(),
|
|
111
|
+
),
|
|
112
|
+
jit_compile=True,
|
|
113
|
+
)()
|
|
114
|
+
|
|
115
|
+
# Test results
|
|
116
|
+
np.testing.assert_approx_equal(
|
|
117
|
+
np.mean(info[0].is_accepted), 0.23, significant=1
|
|
118
|
+
)
|
|
119
|
+
np.testing.assert_allclose(
|
|
120
|
+
tree_map(lambda x: np.mean(x), posterior),
|
|
121
|
+
[0.0, 1.0, 2.0],
|
|
122
|
+
rtol=0.01,
|
|
123
|
+
atol=0.05,
|
|
124
|
+
)
|
|
125
|
+
np.testing.assert_allclose(
|
|
126
|
+
tree_map(lambda x: np.var(x), posterior),
|
|
127
|
+
[1.0, 1.0, 1.0],
|
|
128
|
+
rtol=0.01,
|
|
129
|
+
atol=0.05,
|
|
130
|
+
)
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
def test_rwmh_2kernel_mcmc():
|
|
134
|
+
seed = get_seed()
|
|
135
|
+
|
|
136
|
+
initial_position = simple_model.sample(seed=seed)
|
|
137
|
+
|
|
138
|
+
kernel = Step(rwmh(scale=2.3), ["foo"]) >> Step(
|
|
139
|
+
rwmh(scale=1.8), ["bar", "baz"]
|
|
140
|
+
)
|
|
141
|
+
|
|
142
|
+
posterior, info = tf.function(
|
|
143
|
+
lambda: mcmc(
|
|
144
|
+
NUM_SAMPLES,
|
|
145
|
+
sampling_algorithm=kernel,
|
|
146
|
+
target_density_fn=simple_model.log_prob,
|
|
147
|
+
initial_position=initial_position,
|
|
148
|
+
seed=get_seed(),
|
|
149
|
+
),
|
|
150
|
+
jit_compile=True,
|
|
151
|
+
)()
|
|
152
|
+
|
|
153
|
+
# Test results
|
|
154
|
+
np.testing.assert_allclose(
|
|
155
|
+
tree_flatten(tree_map(lambda x: np.mean(x), info)),
|
|
156
|
+
[0.45, 0.33],
|
|
157
|
+
atol=0.01,
|
|
158
|
+
rtol=0.05,
|
|
159
|
+
)
|
|
160
|
+
np.testing.assert_allclose(
|
|
161
|
+
tree_map(lambda x: np.mean(x), posterior),
|
|
162
|
+
[0.0, 1.0, 2.0],
|
|
163
|
+
rtol=0.01,
|
|
164
|
+
atol=0.05,
|
|
165
|
+
)
|
|
166
|
+
np.testing.assert_allclose(
|
|
167
|
+
tree_map(lambda x: np.var(x), posterior),
|
|
168
|
+
[1.0, 1.0, 1.0],
|
|
169
|
+
rtol=0.01,
|
|
170
|
+
atol=0.05,
|
|
171
|
+
)
|
|
172
|
+
|
|
173
|
+
|
|
174
|
+
def test_rwmh_3kernel_mcmc():
|
|
175
|
+
seed = get_seed()
|
|
176
|
+
|
|
177
|
+
initial_position = simple_model.sample(seed=seed)
|
|
178
|
+
|
|
179
|
+
kernel = (
|
|
180
|
+
Step(rwmh(scale=2.3), ["foo"])
|
|
181
|
+
>> Step(rwmh(scale=2.3), ["bar"])
|
|
182
|
+
>> Step(rwmh(scale=2.3), ["baz"])
|
|
183
|
+
)
|
|
184
|
+
|
|
185
|
+
posterior, info = tf.function(
|
|
186
|
+
lambda: mcmc(
|
|
187
|
+
NUM_SAMPLES,
|
|
188
|
+
sampling_algorithm=kernel,
|
|
189
|
+
target_density_fn=simple_model.log_prob,
|
|
190
|
+
initial_position=initial_position,
|
|
191
|
+
seed=get_seed(),
|
|
192
|
+
),
|
|
193
|
+
jit_compile=True,
|
|
194
|
+
)()
|
|
195
|
+
|
|
196
|
+
# Test results
|
|
197
|
+
np.testing.assert_allclose(
|
|
198
|
+
tree_flatten(tree_map(lambda x: np.mean(x), info))[0],
|
|
199
|
+
[0.45, 0.45, 0.45],
|
|
200
|
+
atol=0.01,
|
|
201
|
+
rtol=0.05,
|
|
202
|
+
)
|
|
203
|
+
np.testing.assert_allclose(
|
|
204
|
+
tree_map(lambda x: np.mean(x), posterior),
|
|
205
|
+
[0.0, 1.0, 2.0],
|
|
206
|
+
rtol=0.01,
|
|
207
|
+
atol=0.05,
|
|
208
|
+
)
|
|
209
|
+
np.testing.assert_allclose(
|
|
210
|
+
tree_map(lambda x: np.var(x), posterior),
|
|
211
|
+
[1.0, 1.0, 1.0],
|
|
212
|
+
rtol=0.01,
|
|
213
|
+
atol=0.05,
|
|
214
|
+
)
|
|
@@ -0,0 +1,49 @@
|
|
|
1
|
+
"""Simple counting kernel for testing"""
|
|
2
|
+
|
|
3
|
+
from typing import NamedTuple
|
|
4
|
+
|
|
5
|
+
import tensorflow as tf
|
|
6
|
+
|
|
7
|
+
from .mcmc_base import ChainState, SamplingAlgorithm
|
|
8
|
+
|
|
9
|
+
__all__ = ["CountingKernelInfo", "CountingKernelState", "counting_kernel"]
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
class CountingKernelState(NamedTuple):
|
|
13
|
+
invocation: int
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class CountingKernelInfo(NamedTuple):
|
|
17
|
+
is_accepted: bool
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def counting_kernel():
|
|
21
|
+
def init_fn(target_log_prob_fn, position):
|
|
22
|
+
chain_state = ChainState(
|
|
23
|
+
position=position,
|
|
24
|
+
log_density=target_log_prob_fn(**position._asdict()),
|
|
25
|
+
log_density_grad=(),
|
|
26
|
+
)
|
|
27
|
+
kernel_state = CountingKernelState(tf.constant(0))
|
|
28
|
+
|
|
29
|
+
return chain_state, kernel_state
|
|
30
|
+
|
|
31
|
+
def step_fn(target_log_prob_fn, chain_and_kernel_state, seed):
|
|
32
|
+
chain_state, kernel_state = chain_and_kernel_state
|
|
33
|
+
|
|
34
|
+
new_position = chain_state.position.__class__(
|
|
35
|
+
**{k: v + 1.0 for k, v in chain_state.position._asdict().items()}
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
new_chain_state = ChainState(
|
|
39
|
+
position=new_position,
|
|
40
|
+
log_density=target_log_prob_fn(**new_position._asdict()),
|
|
41
|
+
log_density_grad=(),
|
|
42
|
+
)
|
|
43
|
+
new_kernel_state = CountingKernelState(kernel_state.invocation + 1)
|
|
44
|
+
|
|
45
|
+
return (new_chain_state, new_kernel_state), CountingKernelInfo(
|
|
46
|
+
tf.constant(True)
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
return SamplingAlgorithm(init_fn, step_fn)
|