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,113 @@
|
|
|
1
|
+
import tensorflow as tf
|
|
2
|
+
import tensorflow_probability as tfp
|
|
3
|
+
from tensorflow_probability.python.internal import (
|
|
4
|
+
reparameterization,
|
|
5
|
+
tensorshape_util,
|
|
6
|
+
)
|
|
7
|
+
from tensorflow_probability.python.internal.tensor_util import (
|
|
8
|
+
convert_nonref_to_tensor,
|
|
9
|
+
)
|
|
10
|
+
|
|
11
|
+
tfd = tfp.distributions
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def _log_choose(N, k): # noqa: N803
|
|
15
|
+
return (
|
|
16
|
+
tf.math.lgamma(N + 1.0)
|
|
17
|
+
- tf.math.lgamma(k + 1.0)
|
|
18
|
+
- tf.math.lgamma(N - k + 1.0)
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class UniformKCategorical(tfd.Distribution):
|
|
23
|
+
def __init__(
|
|
24
|
+
self,
|
|
25
|
+
k,
|
|
26
|
+
mask,
|
|
27
|
+
float_dtype=tf.float32,
|
|
28
|
+
validate_args=False,
|
|
29
|
+
allow_nan_stats=True,
|
|
30
|
+
name="UniformKCategorical",
|
|
31
|
+
):
|
|
32
|
+
"""Uniform K-Categorical distribution.
|
|
33
|
+
|
|
34
|
+
Given a set of items indexed $1,...,n$ and a boolean mask of the same
|
|
35
|
+
shape sample $k$ indices without replacement.
|
|
36
|
+
|
|
37
|
+
:param k: the number of indices to sample
|
|
38
|
+
:param mask: a boolean mask with `True` where an element is valid,
|
|
39
|
+
otherwise `False`
|
|
40
|
+
:param validate_args: Whether to validate args
|
|
41
|
+
:param allow_nan_stats: allow nan stats
|
|
42
|
+
:param name: name of the distribution
|
|
43
|
+
|
|
44
|
+
Example 1: Generate 4 samples of size k given a mask
|
|
45
|
+
import numpy as np
|
|
46
|
+
import tensorflow as tf
|
|
47
|
+
import tensorflow_probability as tfp
|
|
48
|
+
from gemlib.distributions.kcategorical import UniformKCategorical
|
|
49
|
+
|
|
50
|
+
# Mask determines which indices are valid and returned by sample().
|
|
51
|
+
# Below combinations of the indices 0, 3, 4, and 6 will be realised
|
|
52
|
+
# when sampling.
|
|
53
|
+
mask = [True, False, False, True, True, False, True]
|
|
54
|
+
X = UniformKCategorical(k=3, mask=mask)
|
|
55
|
+
x = X.sample(4)
|
|
56
|
+
tf.print(x)
|
|
57
|
+
|
|
58
|
+
Example 2: Probability of a given sample
|
|
59
|
+
import numpy as np
|
|
60
|
+
import tensorflow as tf
|
|
61
|
+
import tensorflow_probability as tfp
|
|
62
|
+
from gemlib.distributions.kcategorical import UniformKCategorical
|
|
63
|
+
|
|
64
|
+
# Probability of drawing an unordered sample of
|
|
65
|
+
# size k from N items where N equals the number
|
|
66
|
+
# of True states in the mask determines N.
|
|
67
|
+
mask = [True, False, False, True, True, False, True]
|
|
68
|
+
sample = tf.convert_to_tensor([4, 3, 6])
|
|
69
|
+
X = UniformKCategorical(k=sample.shape[-1], mask=mask)
|
|
70
|
+
lp = X.log_prob(sample)
|
|
71
|
+
tf.print('prob:', tf.exp(lp))
|
|
72
|
+
|
|
73
|
+
"""
|
|
74
|
+
parameters = dict(locals())
|
|
75
|
+
self._mask = convert_nonref_to_tensor(mask, dtype_hint=tf.bool)
|
|
76
|
+
self._k = convert_nonref_to_tensor(k, dtype_hint=tf.int32)
|
|
77
|
+
self._float_dtype = float_dtype
|
|
78
|
+
dtype = self._k.dtype
|
|
79
|
+
|
|
80
|
+
with tf.name_scope(name) as name:
|
|
81
|
+
super().__init__(
|
|
82
|
+
dtype=dtype,
|
|
83
|
+
reparameterization_type=reparameterization.FULLY_REPARAMETERIZED,
|
|
84
|
+
validate_args=validate_args,
|
|
85
|
+
allow_nan_stats=allow_nan_stats,
|
|
86
|
+
parameters=parameters,
|
|
87
|
+
name=name,
|
|
88
|
+
)
|
|
89
|
+
|
|
90
|
+
def _batch_shape(self):
|
|
91
|
+
return tf.TensorShape(self._mask.shape[:-1])
|
|
92
|
+
|
|
93
|
+
def _event_shape(self):
|
|
94
|
+
return tensorshape_util.constant_value_as_shape(
|
|
95
|
+
tf.expand_dims(self._k, axis=0)
|
|
96
|
+
)
|
|
97
|
+
|
|
98
|
+
def _sample_n(self, n, seed=None):
|
|
99
|
+
seed = tfp.random.sanitize_seed(seed, salt="KCategorical._sample_n")
|
|
100
|
+
u = tfd.Uniform(
|
|
101
|
+
low=tf.zeros(self._mask.shape, dtype=tf.float32),
|
|
102
|
+
high=tf.ones(self._mask.shape, dtype=tf.float32),
|
|
103
|
+
).sample(n, seed=seed)
|
|
104
|
+
u = u * tf.cast(self._mask, u.dtype)
|
|
105
|
+
_, x = tf.math.top_k(u, k=self._k, sorted=True)
|
|
106
|
+
return x
|
|
107
|
+
|
|
108
|
+
def _log_prob(self, x):
|
|
109
|
+
N = tf.math.count_nonzero(self._mask, axis=-1)
|
|
110
|
+
return -_log_choose(
|
|
111
|
+
tf.cast(N, dtype=self._float_dtype),
|
|
112
|
+
tf.cast(self._k, dtype=self._float_dtype),
|
|
113
|
+
)
|
|
@@ -0,0 +1,52 @@
|
|
|
1
|
+
# Dependency imports
|
|
2
|
+
import numpy as np
|
|
3
|
+
import tensorflow as tf
|
|
4
|
+
from tensorflow_probability.python.internal import test_util
|
|
5
|
+
|
|
6
|
+
from gemlib.distributions.kcategorical import UniformKCategorical
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@test_util.test_all_tf_execution_regimes
|
|
10
|
+
class TestUniformInteger(test_util.TestCase):
|
|
11
|
+
def setUp(self):
|
|
12
|
+
self.seed = 10402302
|
|
13
|
+
self.mask = [True, False, False, True, True, False, True]
|
|
14
|
+
|
|
15
|
+
def test_sample(self):
|
|
16
|
+
"""Sample draws one sample with shape (1,3) ."""
|
|
17
|
+
tf.random.set_seed(self.seed)
|
|
18
|
+
target = tf.convert_to_tensor([[0, 6, 3]])
|
|
19
|
+
X = UniformKCategorical(k=target.shape[-1], mask=self.mask)
|
|
20
|
+
x = X.sample(1, seed=1234)
|
|
21
|
+
x_ = self.evaluate(x)
|
|
22
|
+
target_ = self.evaluate(target)
|
|
23
|
+
self.assertAllEqual(target_, x_)
|
|
24
|
+
self.assertDTypeEqual(x_, np.int32)
|
|
25
|
+
|
|
26
|
+
def test_log_prob_float32(self):
|
|
27
|
+
"""Log probability of 1 realisations using float32."""
|
|
28
|
+
target = tf.convert_to_tensor([[0, 6, 3]])
|
|
29
|
+
X = UniformKCategorical(
|
|
30
|
+
k=target.shape[-1], mask=self.mask, float_dtype=tf.float32
|
|
31
|
+
)
|
|
32
|
+
lp = X.log_prob(target)
|
|
33
|
+
lp_ = self.evaluate(lp)
|
|
34
|
+
print("32", lp)
|
|
35
|
+
self.assertAlmostEqual(lp_, -1.3862944, places=5)
|
|
36
|
+
self.assertDTypeEqual(lp_, np.float32)
|
|
37
|
+
|
|
38
|
+
def test_log_prob_float64(self):
|
|
39
|
+
"""Log probability of 1 realisations using float32."""
|
|
40
|
+
target = tf.convert_to_tensor([[0, 6, 3]])
|
|
41
|
+
X = UniformKCategorical(
|
|
42
|
+
k=target.shape[-1], mask=self.mask, float_dtype=tf.float64
|
|
43
|
+
)
|
|
44
|
+
lp = X.log_prob(target)
|
|
45
|
+
lp_ = self.evaluate(lp)
|
|
46
|
+
print("64", lp)
|
|
47
|
+
self.assertAlmostEqual(lp_, -1.3862944, places=5)
|
|
48
|
+
self.assertDTypeEqual(lp_, np.float64)
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
if __name__ == "__main__":
|
|
52
|
+
tf.test.main()
|
|
@@ -0,0 +1,169 @@
|
|
|
1
|
+
"""The UniformInteger distribution class"""
|
|
2
|
+
|
|
3
|
+
import tensorflow as tf
|
|
4
|
+
import tensorflow_probability as tfp
|
|
5
|
+
from tensorflow_probability.python.internal import (
|
|
6
|
+
reparameterization,
|
|
7
|
+
samplers,
|
|
8
|
+
)
|
|
9
|
+
|
|
10
|
+
tfd = tfp.distributions
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class UniformInteger(tfd.Distribution):
|
|
14
|
+
def __init__(
|
|
15
|
+
self,
|
|
16
|
+
low=0,
|
|
17
|
+
high=1,
|
|
18
|
+
validate_args=False,
|
|
19
|
+
allow_nan_stats=True,
|
|
20
|
+
dtype=tf.int32,
|
|
21
|
+
float_dtype=tf.float32,
|
|
22
|
+
name="UniformInteger",
|
|
23
|
+
):
|
|
24
|
+
"""Integer uniform distribution.
|
|
25
|
+
|
|
26
|
+
Args:
|
|
27
|
+
----
|
|
28
|
+
low: Integer tensor, lower boundary of the output interval. Must have
|
|
29
|
+
`low <= high`.
|
|
30
|
+
high: Integer tensor, _inclusive_ upper boundary of the output
|
|
31
|
+
interval. Must have `low <= high`.
|
|
32
|
+
validate_args: Python `bool`, default `False`. When `True`
|
|
33
|
+
distribution parameters are checked for validity despite possibly
|
|
34
|
+
degrading runtime performance. When `False` invalid inputs may
|
|
35
|
+
silently render incorrect outputs.
|
|
36
|
+
allow_nan_stats: Python `bool`, default `True`. When `True`,
|
|
37
|
+
statistics (e.g., mean, mode, variance) use the value "`NaN`" to
|
|
38
|
+
indicate the result is undefined. When `False`, an exception is
|
|
39
|
+
raised if one or more of the statistic's batch members are undefined.
|
|
40
|
+
dtype: returned integer dtype when sampling.
|
|
41
|
+
float_dtype: returned float dtype of log probability.
|
|
42
|
+
name: Python `str` name prefixed to Ops created by this class.
|
|
43
|
+
|
|
44
|
+
Example 1: sampling
|
|
45
|
+
```python
|
|
46
|
+
import tensorflow as tf
|
|
47
|
+
from gemlib.distributions.uniform_integer import UniformInteger
|
|
48
|
+
|
|
49
|
+
tf.random.set_seed(10402302)
|
|
50
|
+
X = UniformInteger(0, 10, dtype=tf.int32)
|
|
51
|
+
x = X.sample([3, 3], seed=1)
|
|
52
|
+
tf.print("samples:", x, "=", [[8, 4, 8], [2, 7, 9], [6, 0, 9]])
|
|
53
|
+
```
|
|
54
|
+
|
|
55
|
+
Example 2: log probability
|
|
56
|
+
```python
|
|
57
|
+
import tensorflow as tf
|
|
58
|
+
from gemlib.distributions.uniform_integer import UniformInteger
|
|
59
|
+
|
|
60
|
+
X = UniformInteger(0, 10, float_dtype=tf.float32)
|
|
61
|
+
lp = X.log_prob([[8, 4, 8], [2, 7, 9], [6, 0, 9]])
|
|
62
|
+
total_lp = tf.math.round(tf.math.reduce_sum(lp) * 1e5) / 1e5
|
|
63
|
+
tf.print("total lp:", total_lp, "= -20.72327")
|
|
64
|
+
```
|
|
65
|
+
|
|
66
|
+
Raises:
|
|
67
|
+
------
|
|
68
|
+
InvalidArgument if `low > high` and `validate_args=False`.
|
|
69
|
+
|
|
70
|
+
"""
|
|
71
|
+
parameters = dict(locals())
|
|
72
|
+
with tf.name_scope(name) as name:
|
|
73
|
+
self._low = tf.cast(low, name="low", dtype=dtype)
|
|
74
|
+
self._high = tf.cast(high, name="high", dtype=dtype)
|
|
75
|
+
super().__init__(
|
|
76
|
+
dtype=dtype,
|
|
77
|
+
reparameterization_type=reparameterization.FULLY_REPARAMETERIZED,
|
|
78
|
+
validate_args=validate_args,
|
|
79
|
+
allow_nan_stats=allow_nan_stats,
|
|
80
|
+
parameters=parameters,
|
|
81
|
+
name=name,
|
|
82
|
+
)
|
|
83
|
+
self.float_dtype = float_dtype
|
|
84
|
+
if validate_args is True:
|
|
85
|
+
tf.assert_greater(
|
|
86
|
+
self._high, self._low, "Condition low < high failed"
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
@staticmethod
|
|
90
|
+
def _param_shapes(sample_shape):
|
|
91
|
+
return dict(
|
|
92
|
+
zip(
|
|
93
|
+
("low", "high"),
|
|
94
|
+
([tf.convert_to_tensor(sample_shape, dtype=tf.int32)] * 2),
|
|
95
|
+
)
|
|
96
|
+
)
|
|
97
|
+
|
|
98
|
+
@classmethod
|
|
99
|
+
def _params_event_ndims(cls):
|
|
100
|
+
return {"low": 0, "high": 0}
|
|
101
|
+
|
|
102
|
+
@property
|
|
103
|
+
def low(self):
|
|
104
|
+
"""Lower boundary of the output interval."""
|
|
105
|
+
return self._low
|
|
106
|
+
|
|
107
|
+
@property
|
|
108
|
+
def high(self):
|
|
109
|
+
"""Upper boundary of the output interval."""
|
|
110
|
+
return self._high
|
|
111
|
+
|
|
112
|
+
def range(self, name="range"):
|
|
113
|
+
"""`high - low`."""
|
|
114
|
+
with self._name_and_control_scope(name):
|
|
115
|
+
return self._range()
|
|
116
|
+
|
|
117
|
+
def _range(self, low=None, high=None):
|
|
118
|
+
low = self.low if low is None else low
|
|
119
|
+
high = self.high if high is None else high
|
|
120
|
+
return high - low
|
|
121
|
+
|
|
122
|
+
def _batch_shape_tensor(self, low=None, high=None):
|
|
123
|
+
return tf.broadcast_dynamic_shape(
|
|
124
|
+
tf.shape(self.low if low is None else low),
|
|
125
|
+
tf.shape(self.high if high is None else high),
|
|
126
|
+
)
|
|
127
|
+
|
|
128
|
+
def _batch_shape(self):
|
|
129
|
+
return tf.broadcast_static_shape(self.low.shape, self.high.shape)
|
|
130
|
+
|
|
131
|
+
def _event_shape_tensor(self):
|
|
132
|
+
return tf.constant([], dtype=tf.int32)
|
|
133
|
+
|
|
134
|
+
def _event_shape(self):
|
|
135
|
+
return tf.TensorShape([])
|
|
136
|
+
|
|
137
|
+
def _sample_n(self, n, seed=None):
|
|
138
|
+
with tf.name_scope("sample_n"):
|
|
139
|
+
low = tf.convert_to_tensor(self.low)
|
|
140
|
+
high = tf.convert_to_tensor(self.high)
|
|
141
|
+
shape = tf.concat(
|
|
142
|
+
[[n], self._batch_shape_tensor(low=low, high=high)], 0
|
|
143
|
+
)
|
|
144
|
+
samples = samplers.uniform(shape=shape, dtype=tf.float32, seed=seed)
|
|
145
|
+
return low + tf.cast(
|
|
146
|
+
tf.cast(self._range(low=low, high=high), tf.float32) * samples,
|
|
147
|
+
self.dtype,
|
|
148
|
+
)
|
|
149
|
+
|
|
150
|
+
def _prob(self, x):
|
|
151
|
+
with tf.name_scope("prob"):
|
|
152
|
+
low = tf.cast(self.low, self.float_dtype)
|
|
153
|
+
high = tf.cast(self.high, self.float_dtype)
|
|
154
|
+
x = tf.cast(x, dtype=self.float_dtype)
|
|
155
|
+
|
|
156
|
+
return tf.where(
|
|
157
|
+
tf.math.is_nan(x),
|
|
158
|
+
x,
|
|
159
|
+
tf.where(
|
|
160
|
+
(x < low) | (x >= high),
|
|
161
|
+
tf.zeros_like(x),
|
|
162
|
+
tf.ones_like(x) / self._range(low=low, high=high),
|
|
163
|
+
),
|
|
164
|
+
)
|
|
165
|
+
|
|
166
|
+
def _log_prob(self, x):
|
|
167
|
+
with tf.name_scope("log_prob"):
|
|
168
|
+
res = tf.math.log(self._prob(x))
|
|
169
|
+
return res
|
|
@@ -0,0 +1,55 @@
|
|
|
1
|
+
# Dependency imports
|
|
2
|
+
import numpy as np
|
|
3
|
+
import tensorflow as tf
|
|
4
|
+
from tensorflow_probability.python.internal import test_util
|
|
5
|
+
|
|
6
|
+
from gemlib.distributions.uniform_integer import UniformInteger
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@test_util.test_all_tf_execution_regimes
|
|
10
|
+
class TestUniformInteger(test_util.TestCase):
|
|
11
|
+
def setUp(self):
|
|
12
|
+
self.seed = 10402302
|
|
13
|
+
self.fixture = [[8, 4, 8], [2, 7, 9], [6, 0, 9]]
|
|
14
|
+
|
|
15
|
+
def test_sample_n_int32(self):
|
|
16
|
+
"""Sample returning dtype int32."""
|
|
17
|
+
tf.random.set_seed(self.seed)
|
|
18
|
+
X = UniformInteger(0, 10)
|
|
19
|
+
x = X.sample([3, 3], seed=1)
|
|
20
|
+
x_ = self.evaluate(x)
|
|
21
|
+
self.fixture_ = self.evaluate(tf.convert_to_tensor(self.fixture))
|
|
22
|
+
self.assertAllEqual(self.fixture, x_)
|
|
23
|
+
self.assertDTypeEqual(x_, np.int32)
|
|
24
|
+
|
|
25
|
+
def test_sample_n_int64(self):
|
|
26
|
+
"""Sample returning int64."""
|
|
27
|
+
tf.random.set_seed(self.seed)
|
|
28
|
+
X = UniformInteger(0, 10, dtype=tf.int64)
|
|
29
|
+
x = X.sample([3, 3], seed=1)
|
|
30
|
+
x_ = self.evaluate(x)
|
|
31
|
+
self.fixture_ = self.evaluate(tf.convert_to_tensor(self.fixture))
|
|
32
|
+
self.assertAllEqual(self.fixture_, x_)
|
|
33
|
+
self.assertDTypeEqual(x_, np.int64)
|
|
34
|
+
|
|
35
|
+
def test_log_prob_float32(self):
|
|
36
|
+
"""log_prob returning float32."""
|
|
37
|
+
X = UniformInteger(0, 10)
|
|
38
|
+
lp = X.log_prob(self.fixture)
|
|
39
|
+
self.assertSequenceEqual(lp.shape, [3, 3])
|
|
40
|
+
lp_ = self.evaluate(lp)
|
|
41
|
+
self.assertAlmostEqual(np.sum(lp_), -20.723265, places=5)
|
|
42
|
+
self.assertDTypeEqual(lp_, np.float32)
|
|
43
|
+
|
|
44
|
+
def test_log_prob_float64(self):
|
|
45
|
+
"""log_prob returning float64."""
|
|
46
|
+
X = UniformInteger(0, 10, float_dtype=tf.float64)
|
|
47
|
+
lp = X.log_prob(self.fixture)
|
|
48
|
+
self.assertSequenceEqual(lp.shape, [3, 3])
|
|
49
|
+
lp_ = self.evaluate(lp)
|
|
50
|
+
self.assertAlmostEqual(np.sum(lp_), -20.723265, places=5)
|
|
51
|
+
self.assertDTypeEqual(lp_, np.float64)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
if __name__ == "__main__":
|
|
55
|
+
tf.test.main()
|
gemlib/mcmc/__init__.py
ADDED
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
"""MCMC kernel addons"""
|
|
2
|
+
|
|
3
|
+
import gemlib.mcmc.discrete_time_state_transition_model as discrete_time
|
|
4
|
+
from gemlib.mcmc.adaptive_random_walk_metropolis import (
|
|
5
|
+
AdaptiveRandomWalkMetropolis,
|
|
6
|
+
)
|
|
7
|
+
from gemlib.mcmc.chain_binomial_rippler import CBRKernel
|
|
8
|
+
from gemlib.mcmc.compound_kernel import CompoundKernel
|
|
9
|
+
from gemlib.mcmc.damped_chain_binomial_rippler import DampedCBRKernel
|
|
10
|
+
from gemlib.mcmc.gibbs_kernel import GibbsKernel
|
|
11
|
+
from gemlib.mcmc.h5_posterior import Posterior
|
|
12
|
+
from gemlib.mcmc.multi_scan_kernel import MultiScanKernel
|
|
13
|
+
|
|
14
|
+
__all__ = [
|
|
15
|
+
"AdaptiveRandomWalkMetropolis",
|
|
16
|
+
"CBRKernel",
|
|
17
|
+
"CompoundKernel",
|
|
18
|
+
"DampedCBRKernel",
|
|
19
|
+
"GibbsKernel",
|
|
20
|
+
"MultiScanKernel",
|
|
21
|
+
"Posterior",
|
|
22
|
+
"discrete_time",
|
|
23
|
+
]
|