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,144 @@
|
|
|
1
|
+
# Dependency imports
|
|
2
|
+
import numpy as np
|
|
3
|
+
import tensorflow as tf
|
|
4
|
+
import tensorflow_probability as tfp
|
|
5
|
+
from tensorflow_probability.python import distributions as tfd
|
|
6
|
+
from tensorflow_probability.python.internal import test_util
|
|
7
|
+
|
|
8
|
+
from gemlib.mcmc.adaptive_random_walk_metropolis import (
|
|
9
|
+
AdaptiveRandomWalkMetropolis,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@test_util.test_all_tf_execution_regimes
|
|
14
|
+
class TestAdaptiveRandomWalkMetropolis(test_util.TestCase):
|
|
15
|
+
def test_1d_normal(self):
|
|
16
|
+
"""Sample from Standard Normal Distribution."""
|
|
17
|
+
dtype = np.float32
|
|
18
|
+
|
|
19
|
+
target = tfd.Normal(loc=dtype(0), scale=dtype(1))
|
|
20
|
+
|
|
21
|
+
kernel = AdaptiveRandomWalkMetropolis(
|
|
22
|
+
target_log_prob_fn=target.log_prob,
|
|
23
|
+
target_accept_ratio=0.44,
|
|
24
|
+
initial_covariance=dtype(0.001),
|
|
25
|
+
)
|
|
26
|
+
samples = tfp.mcmc.sample_chain(
|
|
27
|
+
num_results=2000,
|
|
28
|
+
current_state=dtype([0.1]),
|
|
29
|
+
kernel=kernel,
|
|
30
|
+
num_burnin_steps=500,
|
|
31
|
+
trace_fn=None,
|
|
32
|
+
)
|
|
33
|
+
|
|
34
|
+
sample_mean = tf.math.reduce_mean(samples, axis=0)
|
|
35
|
+
sample_std = tf.math.reduce_std(samples, axis=0)
|
|
36
|
+
[sample_mean_, sample_std_] = self.evaluate([sample_mean, sample_std])
|
|
37
|
+
|
|
38
|
+
self.assertAllClose([0.0], sample_mean_, atol=0.17, rtol=0.0)
|
|
39
|
+
self.assertAllClose([1.0], sample_std_, atol=0.2, rtol=0.0)
|
|
40
|
+
|
|
41
|
+
def test_3d_mvn(self):
|
|
42
|
+
"""Sample from 3-variate Gaussian Distribution."""
|
|
43
|
+
dtype = np.float32
|
|
44
|
+
|
|
45
|
+
true_mean = dtype([1.0, 2.0, 3.0])
|
|
46
|
+
true_cov = dtype(
|
|
47
|
+
[[0.36, 0.12, 0.06], [0.12, 0.29, -0.13], [0.06, -0.13, 0.26]]
|
|
48
|
+
)
|
|
49
|
+
target = tfd.MultivariateNormalTriL(
|
|
50
|
+
loc=true_mean, scale_tril=tf.linalg.cholesky(true_cov)
|
|
51
|
+
)
|
|
52
|
+
kernel = AdaptiveRandomWalkMetropolis(
|
|
53
|
+
target_log_prob_fn=target.log_prob,
|
|
54
|
+
initial_covariance=dtype(0.001) * np.eye(3, dtype=dtype),
|
|
55
|
+
)
|
|
56
|
+
samples = tfp.mcmc.sample_chain(
|
|
57
|
+
num_results=2000,
|
|
58
|
+
current_state=dtype([0.1, 0.1, 0.1]),
|
|
59
|
+
kernel=kernel,
|
|
60
|
+
num_burnin_steps=500,
|
|
61
|
+
trace_fn=None,
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
sample_mean = tf.math.reduce_mean(samples, axis=0)
|
|
65
|
+
[sample_mean_] = self.evaluate([sample_mean])
|
|
66
|
+
self.assertAllClose(sample_mean_, true_mean, atol=0.1, rtol=0.1)
|
|
67
|
+
|
|
68
|
+
sample_cov = tfp.stats.covariance(samples)
|
|
69
|
+
sample_cov_ = self.evaluate(sample_cov)
|
|
70
|
+
self.assertAllClose(sample_cov_, true_cov, atol=0.1, rtol=0.1)
|
|
71
|
+
|
|
72
|
+
def test_float64(self):
|
|
73
|
+
"""Sample with dtype float64."""
|
|
74
|
+
dtype = np.float64
|
|
75
|
+
|
|
76
|
+
target = tfd.Normal(loc=dtype(0), scale=dtype(1))
|
|
77
|
+
|
|
78
|
+
kernel = AdaptiveRandomWalkMetropolis(
|
|
79
|
+
target_log_prob_fn=target.log_prob,
|
|
80
|
+
target_accept_ratio=0.44,
|
|
81
|
+
initial_covariance=dtype(0.001),
|
|
82
|
+
)
|
|
83
|
+
tfp.mcmc.sample_chain(
|
|
84
|
+
num_results=20,
|
|
85
|
+
current_state=dtype([0.1]),
|
|
86
|
+
kernel=kernel,
|
|
87
|
+
trace_fn=None,
|
|
88
|
+
)
|
|
89
|
+
|
|
90
|
+
def test_bijector(self):
|
|
91
|
+
"""Employ bijector when sampling."""
|
|
92
|
+
dtype = np.float32
|
|
93
|
+
|
|
94
|
+
target = tfd.Normal(loc=dtype(0), scale=dtype(1))
|
|
95
|
+
|
|
96
|
+
kernel = AdaptiveRandomWalkMetropolis(
|
|
97
|
+
target_log_prob_fn=target.log_prob,
|
|
98
|
+
target_accept_ratio=0.44,
|
|
99
|
+
initial_covariance=dtype(0.001),
|
|
100
|
+
)
|
|
101
|
+
kernel = tfp.mcmc.TransformedTransitionKernel(
|
|
102
|
+
inner_kernel=kernel, bijector=tfp.bijectors.Exp()
|
|
103
|
+
)
|
|
104
|
+
tfp.mcmc.sample_chain(
|
|
105
|
+
num_results=20,
|
|
106
|
+
current_state=dtype([0.1]),
|
|
107
|
+
kernel=kernel,
|
|
108
|
+
trace_fn=None,
|
|
109
|
+
)
|
|
110
|
+
|
|
111
|
+
def test_is_calibrated(self):
|
|
112
|
+
dtype = np.float32
|
|
113
|
+
kernel = AdaptiveRandomWalkMetropolis(
|
|
114
|
+
target_log_prob_fn=lambda x: -tf.square(x) / 2.0,
|
|
115
|
+
initial_covariance=dtype(0.001),
|
|
116
|
+
)
|
|
117
|
+
self.assertTrue(kernel.is_calibrated)
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
# *** RUN TESTS ***
|
|
121
|
+
|
|
122
|
+
# ** Method 1 **
|
|
123
|
+
# When using tf.test.main() from within jupyter notebook use
|
|
124
|
+
# import sys
|
|
125
|
+
# sys.argv = sys.argv[:1]
|
|
126
|
+
# to eliminate the ERROR "unknown command line flag 'f"
|
|
127
|
+
# if __name__ == '__main__':
|
|
128
|
+
# try:
|
|
129
|
+
# tf.test.main()
|
|
130
|
+
# except SystemExit as inst:
|
|
131
|
+
# if inst.args[0] is True: # raised by sys.exit(True) if tests fail
|
|
132
|
+
# raise
|
|
133
|
+
|
|
134
|
+
# ** Method 2 **
|
|
135
|
+
# Althernatively use unittest.main() instead of tf.test.main()
|
|
136
|
+
# when running tests within jupyter notebook:
|
|
137
|
+
# import unittest
|
|
138
|
+
# unittest.main(argv=['first-arg-is-ignored'], exit=False)
|
|
139
|
+
|
|
140
|
+
# ** Method 3 **
|
|
141
|
+
# If not using jupyter notebook (or similar environment) the
|
|
142
|
+
# following should work:
|
|
143
|
+
if __name__ == "__main__":
|
|
144
|
+
tf.test.main()
|
|
Binary file
|
|
@@ -0,0 +1,291 @@
|
|
|
1
|
+
"""A Brownian Bridge kernel is intended to operate on
|
|
2
|
+
a timeseries
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
# ruff: noqa: B023
|
|
6
|
+
|
|
7
|
+
import tensorflow as tf
|
|
8
|
+
import tensorflow_probability as tfp
|
|
9
|
+
from tensorflow_probability.python.internal import samplers
|
|
10
|
+
from tensorflow_probability.python.mcmc.internal import util as mcmc_util
|
|
11
|
+
from tensorflow_probability.python.mcmc.random_walk_metropolis import (
|
|
12
|
+
UncalibratedRandomWalkResults,
|
|
13
|
+
)
|
|
14
|
+
|
|
15
|
+
from gemlib.distributions import BrownianBridge, BrownianMotion, UniformInteger
|
|
16
|
+
|
|
17
|
+
tfd = tfp.distributions
|
|
18
|
+
mcmc = tfp.mcmc
|
|
19
|
+
|
|
20
|
+
MIN_SPAN = 3
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def _slide_left(x, shift):
|
|
24
|
+
x_right = x[..., -1:]
|
|
25
|
+
y = tf.roll(x, -shift, axis=-1)
|
|
26
|
+
mask = tf.range(x.shape[-1]) >= (x.shape[-1] - shift)
|
|
27
|
+
pad = x_right * tf.cast(mask, x_right.dtype)
|
|
28
|
+
return y * tf.cast(~mask, y.dtype) + pad
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def _slide_right(x, shift):
|
|
32
|
+
x_left = x[..., 0:]
|
|
33
|
+
y = tf.roll(x, shift, axis=-1)
|
|
34
|
+
mask = tf.range(x.shape[-1]) < shift
|
|
35
|
+
pad = x_left * tf.cast(mask, x_left.dtype)
|
|
36
|
+
return y * tf.cast(~mask, x_left.dtype) + pad
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class UncalibratedBrownianBridgeKernel(mcmc.TransitionKernel):
|
|
40
|
+
def __init__(
|
|
41
|
+
self,
|
|
42
|
+
target_log_prob_fn,
|
|
43
|
+
index_points,
|
|
44
|
+
span=3,
|
|
45
|
+
scale=0.1,
|
|
46
|
+
left=True,
|
|
47
|
+
right=True,
|
|
48
|
+
name="UncalibratedBrownianBridgeKernel",
|
|
49
|
+
):
|
|
50
|
+
with tf.name_scope(
|
|
51
|
+
mcmc_util.make_name(
|
|
52
|
+
name, "UncalibratedBrownianBridgeKernel", "__init__"
|
|
53
|
+
)
|
|
54
|
+
) as name:
|
|
55
|
+
if span < MIN_SPAN:
|
|
56
|
+
raise ValueError(
|
|
57
|
+
f"`span` must be at least {MIN_SPAN} timepoints"
|
|
58
|
+
)
|
|
59
|
+
if scale <= 0.0:
|
|
60
|
+
raise ValueError("`scale` must be positive")
|
|
61
|
+
|
|
62
|
+
span_parts = list(span) if mcmc_util.is_list_like(span) else [span]
|
|
63
|
+
self.span_parts = [
|
|
64
|
+
tf.convert_to_tensor(s, name="span") for s in span_parts
|
|
65
|
+
]
|
|
66
|
+
|
|
67
|
+
if mcmc_util.is_list_like(scale):
|
|
68
|
+
scale_parts = list(scale)
|
|
69
|
+
else:
|
|
70
|
+
scale_parts = [scale]
|
|
71
|
+
self.scale_parts = [
|
|
72
|
+
tf.convert_to_tensor(s, name="scale") for s in scale_parts
|
|
73
|
+
]
|
|
74
|
+
|
|
75
|
+
self._index_points = tf.convert_to_tensor(index_points)
|
|
76
|
+
self._left = tf.cast(left, tf.int32)
|
|
77
|
+
self._right = tf.cast(right, tf.int32)
|
|
78
|
+
|
|
79
|
+
self.dtype = self.scale_parts[0].dtype
|
|
80
|
+
cls_name = mcmc_util.make_name(
|
|
81
|
+
name, "UncalibratedBrownianBridgeKernel", ""
|
|
82
|
+
)
|
|
83
|
+
|
|
84
|
+
self._parameters = {
|
|
85
|
+
"target_log_prob_fn": target_log_prob_fn,
|
|
86
|
+
"index_points": index_points,
|
|
87
|
+
"span": span,
|
|
88
|
+
"scale": scale,
|
|
89
|
+
"left": left,
|
|
90
|
+
"right": right,
|
|
91
|
+
"name": cls_name,
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
@property
|
|
95
|
+
def is_calibrated(self):
|
|
96
|
+
return False
|
|
97
|
+
|
|
98
|
+
@property
|
|
99
|
+
def name(self):
|
|
100
|
+
return self._parameters["name"]
|
|
101
|
+
|
|
102
|
+
@property
|
|
103
|
+
def span(self):
|
|
104
|
+
return self._parameters["span"]
|
|
105
|
+
|
|
106
|
+
@property
|
|
107
|
+
def scale(self):
|
|
108
|
+
return self._parameters["scale"]
|
|
109
|
+
|
|
110
|
+
@property
|
|
111
|
+
def jitter(self):
|
|
112
|
+
return self._parameters["jitter"]
|
|
113
|
+
|
|
114
|
+
@property
|
|
115
|
+
def target_log_prob_fn(self):
|
|
116
|
+
return self._parameters["target_log_prob_fn"]
|
|
117
|
+
|
|
118
|
+
def one_step(self, current_state, previous_results, seed=None):
|
|
119
|
+
with tf.name_scope(mcmc_util.make_name(self.name, "bbmh", "one_step")):
|
|
120
|
+
with tf.name_scope("initialize"):
|
|
121
|
+
if mcmc_util.is_list_like(current_state):
|
|
122
|
+
current_state_parts = list(current_state)
|
|
123
|
+
else:
|
|
124
|
+
current_state_parts = [current_state]
|
|
125
|
+
current_state_parts = [
|
|
126
|
+
tf.convert_to_tensor(s, name="current_state")
|
|
127
|
+
for s in current_state_parts
|
|
128
|
+
]
|
|
129
|
+
seed = samplers.sanitize_seed(
|
|
130
|
+
seed, salt="UncalibratedBrownianBridgeKernel"
|
|
131
|
+
)
|
|
132
|
+
|
|
133
|
+
new_state_parts = []
|
|
134
|
+
log_acceptance_correction_parts = []
|
|
135
|
+
for current_state_part, span_part, scale_part in zip(
|
|
136
|
+
current_state_parts,
|
|
137
|
+
self.span_parts,
|
|
138
|
+
self.scale_parts,
|
|
139
|
+
):
|
|
140
|
+
t_low_seed, bridge_seed = samplers.split_seed(seed)
|
|
141
|
+
|
|
142
|
+
# Evaluate bridge limits
|
|
143
|
+
t_low = UniformInteger(
|
|
144
|
+
0 - self._left,
|
|
145
|
+
current_state_part.shape[-1] - span_part + self._right,
|
|
146
|
+
).sample(seed=seed)
|
|
147
|
+
|
|
148
|
+
# We have 3 cases:
|
|
149
|
+
# 0. If t_low > 0 and t_high < (current_state.shape[-1]-1):
|
|
150
|
+
# Brownian Bridge
|
|
151
|
+
# 1. If t_low == 0: reverse Brownian motion
|
|
152
|
+
# 2. If t_high >= (current_state.shape[-1]-1): Brownian motion
|
|
153
|
+
def brownian_bridge_proposal():
|
|
154
|
+
with tf.name_scope("brownian_bridge_proposal"):
|
|
155
|
+
indices = t_low + tf.range(span_part)
|
|
156
|
+
current_bridge = tf.gather(
|
|
157
|
+
current_state_part,
|
|
158
|
+
indices=indices,
|
|
159
|
+
)
|
|
160
|
+
bridge = BrownianBridge(
|
|
161
|
+
index_points=tf.gather(self._index_points, indices),
|
|
162
|
+
x0=current_bridge[..., 0],
|
|
163
|
+
x1=current_bridge[..., -1],
|
|
164
|
+
scale=scale_part,
|
|
165
|
+
)
|
|
166
|
+
new_bridge = bridge.sample(seed=bridge_seed)
|
|
167
|
+
log_acceptance_correction = bridge.log_prob(
|
|
168
|
+
current_bridge[1:-1]
|
|
169
|
+
) - bridge.log_prob(new_bridge)
|
|
170
|
+
|
|
171
|
+
new_state = tf.tensor_scatter_nd_update(
|
|
172
|
+
current_state_part,
|
|
173
|
+
indices=indices[1:-1][:, tf.newaxis],
|
|
174
|
+
updates=new_bridge,
|
|
175
|
+
name="update_new_state",
|
|
176
|
+
)
|
|
177
|
+
return new_state, log_acceptance_correction
|
|
178
|
+
|
|
179
|
+
def brownian_motion_right_proposal():
|
|
180
|
+
with tf.name_scope("brownian_motion_proposal"):
|
|
181
|
+
indices = tf.range(
|
|
182
|
+
current_state_part.shape[0] - span_part,
|
|
183
|
+
current_state_part.shape[0],
|
|
184
|
+
) # Index into current_state_part
|
|
185
|
+
current_bridge = tf.gather(
|
|
186
|
+
current_state_part,
|
|
187
|
+
indices=indices,
|
|
188
|
+
name="current_state_slice",
|
|
189
|
+
)
|
|
190
|
+
bridge = BrownianMotion(
|
|
191
|
+
index_points=tf.gather(self._index_points, indices),
|
|
192
|
+
x0=current_bridge[..., 0],
|
|
193
|
+
scale=scale_part,
|
|
194
|
+
)
|
|
195
|
+
new_bridge = bridge.sample(seed=bridge_seed)
|
|
196
|
+
log_acceptance_correction = bridge.log_prob(
|
|
197
|
+
current_bridge[1:]
|
|
198
|
+
) - bridge.log_prob(new_bridge)
|
|
199
|
+
new_state = tf.tensor_scatter_nd_update(
|
|
200
|
+
current_state_part,
|
|
201
|
+
indices=tf.expand_dims(indices[1:], -1),
|
|
202
|
+
updates=new_bridge,
|
|
203
|
+
name="update_new_state",
|
|
204
|
+
)
|
|
205
|
+
return new_state, log_acceptance_correction
|
|
206
|
+
|
|
207
|
+
def brownian_motion_left_proposal():
|
|
208
|
+
with tf.name_scope("brownian_motion_left_proposal"):
|
|
209
|
+
indices = tf.range(span_part)
|
|
210
|
+
current_bridge = tf.gather(
|
|
211
|
+
current_state_part,
|
|
212
|
+
indices=indices,
|
|
213
|
+
name="current_state_slice",
|
|
214
|
+
)
|
|
215
|
+
bridge = BrownianMotion(
|
|
216
|
+
index_points=tf.gather(self._index_points, indices),
|
|
217
|
+
x0=current_bridge[..., -1],
|
|
218
|
+
scale=scale_part,
|
|
219
|
+
)
|
|
220
|
+
new_bridge = bridge.sample(seed=bridge_seed)
|
|
221
|
+
log_acceptance_correction = bridge.log_prob(
|
|
222
|
+
tf.reverse(
|
|
223
|
+
current_bridge[:-1],
|
|
224
|
+
axis=[-1],
|
|
225
|
+
name="reverse_current_bridge",
|
|
226
|
+
),
|
|
227
|
+
) - bridge.log_prob(new_bridge)
|
|
228
|
+
new_state = tf.tensor_scatter_nd_update(
|
|
229
|
+
current_state_part,
|
|
230
|
+
indices=tf.expand_dims(indices[:-1], -1),
|
|
231
|
+
updates=tf.reverse(
|
|
232
|
+
new_bridge, axis=[-1], name="reverse_new_bridge"
|
|
233
|
+
),
|
|
234
|
+
name="update_new_state",
|
|
235
|
+
)
|
|
236
|
+
return new_state, log_acceptance_correction
|
|
237
|
+
|
|
238
|
+
case_enum = ( # 0=bridge, 1=right, 2=left
|
|
239
|
+
tf.cast(
|
|
240
|
+
t_low == (current_state_part.shape[0] - span_part),
|
|
241
|
+
tf.int32,
|
|
242
|
+
)
|
|
243
|
+
+ tf.cast(t_low == -1, tf.int32) * 2
|
|
244
|
+
)
|
|
245
|
+
|
|
246
|
+
(
|
|
247
|
+
new_state,
|
|
248
|
+
log_acceptance_correction_part,
|
|
249
|
+
) = tf.switch_case(
|
|
250
|
+
case_enum,
|
|
251
|
+
[
|
|
252
|
+
brownian_bridge_proposal,
|
|
253
|
+
brownian_motion_right_proposal,
|
|
254
|
+
brownian_motion_left_proposal,
|
|
255
|
+
],
|
|
256
|
+
)
|
|
257
|
+
|
|
258
|
+
new_state_parts.append(new_state)
|
|
259
|
+
log_acceptance_correction_parts.append(
|
|
260
|
+
log_acceptance_correction_part
|
|
261
|
+
)
|
|
262
|
+
|
|
263
|
+
target_log_prob = self.target_log_prob_fn(*new_state_parts)
|
|
264
|
+
|
|
265
|
+
def maybe_flatten(x):
|
|
266
|
+
return x if mcmc_util.is_list_like(current_state) else x[0]
|
|
267
|
+
|
|
268
|
+
return [
|
|
269
|
+
maybe_flatten(new_state_parts),
|
|
270
|
+
UncalibratedRandomWalkResults(
|
|
271
|
+
log_acceptance_correction=maybe_flatten(
|
|
272
|
+
log_acceptance_correction_parts
|
|
273
|
+
),
|
|
274
|
+
target_log_prob=target_log_prob,
|
|
275
|
+
seed=seed,
|
|
276
|
+
),
|
|
277
|
+
]
|
|
278
|
+
|
|
279
|
+
def bootstrap_results(self, current_state):
|
|
280
|
+
if mcmc_util.is_list_like(current_state):
|
|
281
|
+
current_state_parts = list(current_state)
|
|
282
|
+
else:
|
|
283
|
+
current_state_parts = [current_state]
|
|
284
|
+
|
|
285
|
+
init_target_log_prob = self.target_log_prob_fn(*current_state_parts)
|
|
286
|
+
|
|
287
|
+
return UncalibratedRandomWalkResults(
|
|
288
|
+
log_acceptance_correction=tf.zeros_like(init_target_log_prob),
|
|
289
|
+
target_log_prob=init_target_log_prob,
|
|
290
|
+
seed=samplers.zeros_seed(),
|
|
291
|
+
)
|
|
@@ -0,0 +1,164 @@
|
|
|
1
|
+
"""Tests Brownian Bridge kernel"""
|
|
2
|
+
|
|
3
|
+
import os
|
|
4
|
+
import pickle as pkl
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
import tensorflow as tf
|
|
8
|
+
import tensorflow_probability as tfp
|
|
9
|
+
from tensorflow_probability.python.internal import test_util
|
|
10
|
+
|
|
11
|
+
from gemlib.distributions import BrownianMotion
|
|
12
|
+
from gemlib.mcmc.brownian_bridge_kernel import UncalibratedBrownianBridgeKernel
|
|
13
|
+
|
|
14
|
+
tfd = tfp.distributions
|
|
15
|
+
|
|
16
|
+
DTYPE = tf.float64
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def model_fixture():
|
|
20
|
+
"""Fixture from model below"""
|
|
21
|
+
dir_path = os.path.dirname(os.path.realpath(__file__))
|
|
22
|
+
with open(os.path.join(dir_path, "bb_fixture.pkl"), "rb") as f:
|
|
23
|
+
return pkl.load(f)
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
class TestBrownianBridgeKernel(test_util.TestCase):
|
|
27
|
+
def test_simple_brownian_motion(self):
|
|
28
|
+
x = tf.range(0.0, 10.0, 0.1, dtype=DTYPE)
|
|
29
|
+
Y = BrownianMotion(x)
|
|
30
|
+
y = Y.sample()
|
|
31
|
+
|
|
32
|
+
kernel = tfp.mcmc.MetropolisHastings(
|
|
33
|
+
inner_kernel=UncalibratedBrownianBridgeKernel(
|
|
34
|
+
Y.log_prob,
|
|
35
|
+
index_points=x,
|
|
36
|
+
span=90,
|
|
37
|
+
scale=tf.constant(1.0, DTYPE),
|
|
38
|
+
)
|
|
39
|
+
)
|
|
40
|
+
|
|
41
|
+
# kernel = tfp.mcmc.DualAveragingStepSizeAdaptation(
|
|
42
|
+
# inner_kernel=tfp.mcmc.HamiltonianMonteCarlo(
|
|
43
|
+
# target_log_prob_fn=Y.log_prob,
|
|
44
|
+
# num_leapfrog_steps=3,
|
|
45
|
+
# step_size=0.1,
|
|
46
|
+
# ),
|
|
47
|
+
# num_adaptation_steps=500,
|
|
48
|
+
# )
|
|
49
|
+
|
|
50
|
+
samples, results = tf.function(
|
|
51
|
+
lambda: tfp.mcmc.sample_chain(
|
|
52
|
+
num_results=10000, kernel=kernel, current_state=y
|
|
53
|
+
)
|
|
54
|
+
)()
|
|
55
|
+
|
|
56
|
+
print(
|
|
57
|
+
"Acceptance rate:",
|
|
58
|
+
tf.reduce_mean(tf.cast(results.is_accepted, tf.float32)),
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
# fig, ax = plt.subplots(1, 2)
|
|
62
|
+
# ax[0].plot(
|
|
63
|
+
# x[1:],
|
|
64
|
+
# samples.numpy().T,
|
|
65
|
+
# color="lightblue",
|
|
66
|
+
# alpha=0.3,
|
|
67
|
+
# )
|
|
68
|
+
# ax[0].plot(
|
|
69
|
+
# x[1:],
|
|
70
|
+
# y,
|
|
71
|
+
# "o",
|
|
72
|
+
# color="black",
|
|
73
|
+
# )
|
|
74
|
+
# ax[1].plot(samples[:, 75])
|
|
75
|
+
# plt.show()
|
|
76
|
+
|
|
77
|
+
self.assertAllClose(0.0, np.mean(samples[:, 0]), atol=1.0, rtol=0.1)
|
|
78
|
+
self.assertAllClose(10.0, np.var(samples[:, -1]), atol=1.0, rtol=0.1)
|
|
79
|
+
|
|
80
|
+
def test_poisson_with_brownian_mean(self):
|
|
81
|
+
x = tf.range(0.0, 10.0, 0.1, dtype=DTYPE)
|
|
82
|
+
|
|
83
|
+
def model():
|
|
84
|
+
mu0 = tfd.Normal(
|
|
85
|
+
loc=tf.constant(1.0, DTYPE),
|
|
86
|
+
scale=tf.constant(1.0, DTYPE),
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
def mu(mu0):
|
|
90
|
+
return BrownianMotion(x, x0=mu0)
|
|
91
|
+
|
|
92
|
+
def y(mu):
|
|
93
|
+
rate = tf.concat([[0.0], mu], axis=-1)
|
|
94
|
+
return tfd.Independent(
|
|
95
|
+
tfd.Poisson(rate=tf.math.exp(rate)),
|
|
96
|
+
reinterpreted_batch_ndims=1,
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
return tfd.JointDistributionNamed({"mu0": mu0, "mu": mu, "y": y})
|
|
100
|
+
|
|
101
|
+
model = model()
|
|
102
|
+
trial = model_fixture()
|
|
103
|
+
|
|
104
|
+
def logp(mu):
|
|
105
|
+
return model.log_prob(
|
|
106
|
+
{"mu0": trial["mu0"], "mu": mu, "y": trial["y"]}
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
mcmc_kernel = tfp.mcmc.MetropolisHastings(
|
|
110
|
+
inner_kernel=UncalibratedBrownianBridgeKernel(
|
|
111
|
+
logp,
|
|
112
|
+
index_points=x,
|
|
113
|
+
span=5,
|
|
114
|
+
scale=tf.constant(1.0, DTYPE),
|
|
115
|
+
left=True,
|
|
116
|
+
right=True,
|
|
117
|
+
)
|
|
118
|
+
)
|
|
119
|
+
|
|
120
|
+
# mcmc_kernel = tfp.mcmc.DualAveragingStepSizeAdaptation(
|
|
121
|
+
# inner_kernel=tfp.mcmc.HamiltonianMonteCarlo(
|
|
122
|
+
# target_log_prob_fn=logp, num_leapfrog_steps=3, step_size=1.0
|
|
123
|
+
# ),
|
|
124
|
+
# num_adaptation_steps=500,
|
|
125
|
+
# )
|
|
126
|
+
|
|
127
|
+
samples, results = tf.function(
|
|
128
|
+
lambda: tfp.mcmc.sample_chain(
|
|
129
|
+
num_results=5000,
|
|
130
|
+
kernel=mcmc_kernel,
|
|
131
|
+
current_state=tf.fill(
|
|
132
|
+
trial["mu"].shape, tf.constant(3.0, DTYPE)
|
|
133
|
+
),
|
|
134
|
+
)
|
|
135
|
+
)()
|
|
136
|
+
print(
|
|
137
|
+
"Acceptance rate:",
|
|
138
|
+
tf.reduce_mean(tf.cast(results.is_accepted, tf.float32)),
|
|
139
|
+
)
|
|
140
|
+
|
|
141
|
+
# fig, ax = plt.subplots(1, 2)
|
|
142
|
+
# ax[0].plot(
|
|
143
|
+
# np.exp(samples[500:].numpy().T),
|
|
144
|
+
# color="lightblue",
|
|
145
|
+
# alpha=0.3,
|
|
146
|
+
# )
|
|
147
|
+
# ax[0].plot(
|
|
148
|
+
# np.exp(trial["mu"]),
|
|
149
|
+
# "o",
|
|
150
|
+
# color="black",
|
|
151
|
+
# )
|
|
152
|
+
# ax[1].plot(np.exp(samples[:, 75]))
|
|
153
|
+
# plt.show()
|
|
154
|
+
|
|
155
|
+
self.assertAllClose(
|
|
156
|
+
0.0,
|
|
157
|
+
0.0,
|
|
158
|
+
rtol=1.5,
|
|
159
|
+
atol=2.0,
|
|
160
|
+
)
|
|
161
|
+
|
|
162
|
+
|
|
163
|
+
if __name__ == "__main__":
|
|
164
|
+
TestBrownianBridgeKernel().test_poisson_with_brownian_mean()
|