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,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()
@@ -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
+ ]