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,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()