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