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,212 @@
|
|
|
1
|
+
"""Test GibbsKernel"""
|
|
2
|
+
|
|
3
|
+
# Dependency imports
|
|
4
|
+
|
|
5
|
+
from collections import namedtuple
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import tensorflow as tf
|
|
9
|
+
import tensorflow_probability as tfp
|
|
10
|
+
from tensorflow_probability.python import distributions as tfd
|
|
11
|
+
from tensorflow_probability.python.internal import test_util
|
|
12
|
+
|
|
13
|
+
from gemlib.mcmc.gibbs_kernel import GibbsKernel, GibbsStep
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
@test_util.test_all_tf_execution_regimes
|
|
17
|
+
class TestGibbsKernel(test_util.TestCase):
|
|
18
|
+
def test_2d_mvn(self):
|
|
19
|
+
"""Sample from 2-variate MVN Distribution."""
|
|
20
|
+
dtype = np.float32
|
|
21
|
+
true_mean = dtype([1, 1])
|
|
22
|
+
true_cov = dtype([[1, 0.5], [0.5, 1]])
|
|
23
|
+
target = tfd.MultivariateNormalTriL(
|
|
24
|
+
loc=true_mean, scale_tril=tf.linalg.cholesky(true_cov)
|
|
25
|
+
)
|
|
26
|
+
|
|
27
|
+
def logp(x1, x2):
|
|
28
|
+
return target.log_prob([x1, x2])
|
|
29
|
+
|
|
30
|
+
def kernel_make_fn(target_log_prob_fn, state):
|
|
31
|
+
return tfp.mcmc.RandomWalkMetropolis(
|
|
32
|
+
target_log_prob_fn=target_log_prob_fn
|
|
33
|
+
)
|
|
34
|
+
|
|
35
|
+
kernel_list = [(0, kernel_make_fn), (1, kernel_make_fn)]
|
|
36
|
+
kernel = GibbsKernel(target_log_prob_fn=logp, kernel_list=kernel_list)
|
|
37
|
+
samples = tfp.mcmc.sample_chain(
|
|
38
|
+
num_results=2000,
|
|
39
|
+
current_state=[dtype(1), dtype(1)],
|
|
40
|
+
kernel=kernel,
|
|
41
|
+
num_burnin_steps=500,
|
|
42
|
+
trace_fn=None,
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
sample_mean = tf.math.reduce_mean(samples, axis=1)
|
|
46
|
+
[sample_mean_] = self.evaluate([sample_mean])
|
|
47
|
+
self.assertAllClose(sample_mean_, true_mean, atol=0.2, rtol=0.2)
|
|
48
|
+
|
|
49
|
+
sample_cov = tfp.stats.covariance(tf.transpose(samples))
|
|
50
|
+
sample_cov_ = self.evaluate(sample_cov)
|
|
51
|
+
self.assertAllClose(sample_cov_, true_cov, atol=0.1, rtol=0.1)
|
|
52
|
+
|
|
53
|
+
def test_2d_mvn_namedtuple(self):
|
|
54
|
+
"""Sample from 2-variate MVN Distribution."""
|
|
55
|
+
dtype = np.float32
|
|
56
|
+
true_mean = dtype([1, 1])
|
|
57
|
+
true_cov = dtype([[1, 0.5], [0.5, 1]])
|
|
58
|
+
target = tfd.MultivariateNormalTriL(
|
|
59
|
+
loc=true_mean, scale_tril=tf.linalg.cholesky(true_cov)
|
|
60
|
+
)
|
|
61
|
+
|
|
62
|
+
def logp(x, y):
|
|
63
|
+
return target.log_prob([x, y])
|
|
64
|
+
|
|
65
|
+
def kernel_make_fn(target_log_prob_fn, state):
|
|
66
|
+
return tfp.mcmc.RandomWalkMetropolis(
|
|
67
|
+
target_log_prob_fn=target_log_prob_fn
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
kernel_list = [
|
|
71
|
+
GibbsStep("x", kernel_make_fn),
|
|
72
|
+
GibbsStep("y", kernel_make_fn),
|
|
73
|
+
]
|
|
74
|
+
kernel = GibbsKernel(target_log_prob_fn=logp, kernel_list=kernel_list)
|
|
75
|
+
StateTuple = namedtuple("StateTuple", ["x", "y"])
|
|
76
|
+
current_state = StateTuple(dtype(1), dtype(1))
|
|
77
|
+
|
|
78
|
+
samples = tfp.mcmc.sample_chain(
|
|
79
|
+
num_results=2000,
|
|
80
|
+
current_state=current_state,
|
|
81
|
+
kernel=kernel,
|
|
82
|
+
num_burnin_steps=500,
|
|
83
|
+
trace_fn=None,
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
sample_mean = tf.math.reduce_mean(samples, axis=1)
|
|
87
|
+
[sample_mean_] = self.evaluate([sample_mean])
|
|
88
|
+
self.assertAllClose(sample_mean_, true_mean, atol=0.2, rtol=0.2)
|
|
89
|
+
|
|
90
|
+
sample_cov = tfp.stats.covariance(tf.transpose(samples))
|
|
91
|
+
sample_cov_ = self.evaluate(sample_cov)
|
|
92
|
+
self.assertAllClose(sample_cov_, true_cov, atol=0.1, rtol=0.1)
|
|
93
|
+
|
|
94
|
+
def test_float64(self):
|
|
95
|
+
"""Sample with dtype float64."""
|
|
96
|
+
dtype = np.float64
|
|
97
|
+
true_mean = dtype([1, 1])
|
|
98
|
+
true_cov = dtype([[1, 0.5], [0.5, 1]])
|
|
99
|
+
target = tfd.MultivariateNormalTriL(
|
|
100
|
+
loc=true_mean, scale_tril=tf.linalg.cholesky(true_cov)
|
|
101
|
+
)
|
|
102
|
+
|
|
103
|
+
def logp(x1, x2):
|
|
104
|
+
return target.log_prob([x1, x2])
|
|
105
|
+
|
|
106
|
+
def kernel_make_fn(target_log_prob_fn, state):
|
|
107
|
+
return tfp.mcmc.RandomWalkMetropolis(
|
|
108
|
+
target_log_prob_fn=target_log_prob_fn
|
|
109
|
+
)
|
|
110
|
+
|
|
111
|
+
kernel_list = [(0, kernel_make_fn), (1, kernel_make_fn)]
|
|
112
|
+
kernel = GibbsKernel(target_log_prob_fn=logp, kernel_list=kernel_list)
|
|
113
|
+
tfp.mcmc.sample_chain(
|
|
114
|
+
num_results=20,
|
|
115
|
+
current_state=[dtype(1), dtype(1)],
|
|
116
|
+
kernel=kernel,
|
|
117
|
+
trace_fn=None,
|
|
118
|
+
)
|
|
119
|
+
|
|
120
|
+
def test_bijector(self):
|
|
121
|
+
"""Employ bijector when sampling."""
|
|
122
|
+
dtype = np.float32
|
|
123
|
+
true_mean = dtype([1, 1])
|
|
124
|
+
true_cov = dtype([[1, 0.5], [0.5, 1]])
|
|
125
|
+
target = tfd.MultivariateNormalTriL(
|
|
126
|
+
loc=true_mean, scale_tril=tf.linalg.cholesky(true_cov)
|
|
127
|
+
)
|
|
128
|
+
|
|
129
|
+
def logp(x1, x2):
|
|
130
|
+
return target.log_prob([x1, x2])
|
|
131
|
+
|
|
132
|
+
def kernel_make_fn(target_log_prob_fn, state):
|
|
133
|
+
inner_kernel = tfp.mcmc.RandomWalkMetropolis(
|
|
134
|
+
target_log_prob_fn=target_log_prob_fn
|
|
135
|
+
)
|
|
136
|
+
return tfp.mcmc.TransformedTransitionKernel(
|
|
137
|
+
inner_kernel=inner_kernel, bijector=tfp.bijectors.Exp()
|
|
138
|
+
)
|
|
139
|
+
|
|
140
|
+
kernel_list = [(0, kernel_make_fn), (1, kernel_make_fn)]
|
|
141
|
+
kernel = GibbsKernel(target_log_prob_fn=logp, kernel_list=kernel_list)
|
|
142
|
+
tfp.mcmc.sample_chain(
|
|
143
|
+
num_results=20,
|
|
144
|
+
current_state=[dtype(1), dtype(1)],
|
|
145
|
+
kernel=kernel,
|
|
146
|
+
trace_fn=None,
|
|
147
|
+
)
|
|
148
|
+
|
|
149
|
+
def test_gradient_based_sampler(self):
|
|
150
|
+
"""Make sure Gibbs kernel is compatible with gradient-based
|
|
151
|
+
samplers
|
|
152
|
+
"""
|
|
153
|
+
dtype = np.float32
|
|
154
|
+
true_mean = dtype([1, 1])
|
|
155
|
+
true_cov = dtype([[1, 0.5], [0.5, 1]])
|
|
156
|
+
target = tfd.MultivariateNormalTriL(
|
|
157
|
+
loc=true_mean, scale_tril=tf.linalg.cholesky(true_cov)
|
|
158
|
+
)
|
|
159
|
+
|
|
160
|
+
def logp(x1, x2):
|
|
161
|
+
return target.log_prob([x1, x2])
|
|
162
|
+
|
|
163
|
+
def kernel_make_rwm_fn(target_log_prob_fn, state):
|
|
164
|
+
inner_kernel = tfp.mcmc.RandomWalkMetropolis(
|
|
165
|
+
target_log_prob_fn=target_log_prob_fn
|
|
166
|
+
)
|
|
167
|
+
return tfp.mcmc.TransformedTransitionKernel(
|
|
168
|
+
inner_kernel=inner_kernel, bijector=tfp.bijectors.Exp()
|
|
169
|
+
)
|
|
170
|
+
|
|
171
|
+
def kernel_make_hmc_fn(target_log_prob_fn, state):
|
|
172
|
+
inner_kernel = tfp.mcmc.HamiltonianMonteCarlo(
|
|
173
|
+
target_log_prob_fn=target_log_prob_fn,
|
|
174
|
+
step_size=0.1,
|
|
175
|
+
num_leapfrog_steps=3,
|
|
176
|
+
)
|
|
177
|
+
return tfp.mcmc.TransformedTransitionKernel(
|
|
178
|
+
inner_kernel=inner_kernel, bijector=tfp.bijectors.Exp()
|
|
179
|
+
)
|
|
180
|
+
|
|
181
|
+
kernel_list = [(0, kernel_make_rwm_fn), (1, kernel_make_hmc_fn)]
|
|
182
|
+
kernel = GibbsKernel(target_log_prob_fn=logp, kernel_list=kernel_list)
|
|
183
|
+
tfp.mcmc.sample_chain(
|
|
184
|
+
num_results=20,
|
|
185
|
+
current_state=[dtype(1), dtype(1)],
|
|
186
|
+
kernel=kernel,
|
|
187
|
+
trace_fn=None,
|
|
188
|
+
)
|
|
189
|
+
|
|
190
|
+
def test_is_calibrated(self):
|
|
191
|
+
dtype = np.float32
|
|
192
|
+
true_mean = dtype([1, 1])
|
|
193
|
+
true_cov = dtype([[1, 0.5], [0.5, 1]])
|
|
194
|
+
target = tfd.MultivariateNormalTriL(
|
|
195
|
+
loc=true_mean, scale_tril=tf.linalg.cholesky(true_cov)
|
|
196
|
+
)
|
|
197
|
+
|
|
198
|
+
def logp(x1, x2):
|
|
199
|
+
return target.log_prob([x1, x2])
|
|
200
|
+
|
|
201
|
+
def kernel_make_fn(target_log_prob_fn, state):
|
|
202
|
+
return tfp.mcmc.RandomWalkMetropolis(
|
|
203
|
+
target_log_prob_fn=target_log_prob_fn
|
|
204
|
+
)
|
|
205
|
+
|
|
206
|
+
kernel_list = [(0, kernel_make_fn), (1, kernel_make_fn)]
|
|
207
|
+
kernel = GibbsKernel(target_log_prob_fn=logp, kernel_list=kernel_list)
|
|
208
|
+
self.assertTrue(kernel.is_calibrated)
|
|
209
|
+
|
|
210
|
+
|
|
211
|
+
if __name__ == "__main__":
|
|
212
|
+
tf.test.main()
|
|
@@ -0,0 +1,77 @@
|
|
|
1
|
+
"""Class for writing posterior samples"""
|
|
2
|
+
|
|
3
|
+
import h5py
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def _maybe_tf_dtype(dtype):
|
|
7
|
+
if hasattr(dtype, "as_numpy_dtype"):
|
|
8
|
+
return dtype.as_numpy_dtype
|
|
9
|
+
return dtype
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def _maybe_to_numpy(val):
|
|
13
|
+
if hasattr(val, "numpy"):
|
|
14
|
+
return val.numpy()
|
|
15
|
+
return val
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
class Posterior:
|
|
19
|
+
def __init__(self, filename, sample_dict, results_dict, num_samples):
|
|
20
|
+
"""Constructs a posterior output object
|
|
21
|
+
|
|
22
|
+
:param filename: the name of the backend HDF5 file
|
|
23
|
+
:param sample_dict: a dictionary containing `key`:`shape_tuple`
|
|
24
|
+
:param results_dict: a dictionary containing `key`:`shape_tuple`
|
|
25
|
+
:param num_samples: total number of samples
|
|
26
|
+
"""
|
|
27
|
+
self._num_samples = num_samples
|
|
28
|
+
self._file = h5py.File(
|
|
29
|
+
filename,
|
|
30
|
+
"w",
|
|
31
|
+
rdcc_nbytes=1024**2 * 400,
|
|
32
|
+
rdcc_nslots=100000,
|
|
33
|
+
libver="latest",
|
|
34
|
+
)
|
|
35
|
+
self._file.swmr_mode = True
|
|
36
|
+
|
|
37
|
+
self._sample_group = self._file.create_group("samples")
|
|
38
|
+
self._create_data_tree(sample_dict, self._sample_group)
|
|
39
|
+
|
|
40
|
+
self._results_group = self._file.create_group("results")
|
|
41
|
+
self._create_data_tree(results_dict, self._results_group)
|
|
42
|
+
|
|
43
|
+
def __del__(self):
|
|
44
|
+
self._file.close()
|
|
45
|
+
|
|
46
|
+
def __getitem__(self, path):
|
|
47
|
+
return self._file[path]
|
|
48
|
+
|
|
49
|
+
def _create_data_tree(self, data_dict, h5dataset):
|
|
50
|
+
for k, v in data_dict.items():
|
|
51
|
+
if isinstance(v, dict):
|
|
52
|
+
h5group = h5dataset.create_group(k)
|
|
53
|
+
self._create_data_tree(v, h5group)
|
|
54
|
+
|
|
55
|
+
else:
|
|
56
|
+
h5dataset.create_dataset(
|
|
57
|
+
k,
|
|
58
|
+
(self._num_samples,) + v.shape[1:],
|
|
59
|
+
dtype=_maybe_tf_dtype(v.dtype),
|
|
60
|
+
compression="gzip",
|
|
61
|
+
)
|
|
62
|
+
|
|
63
|
+
def _write(self, sample_dict, dset, first_dim_offset=0):
|
|
64
|
+
for k, v in sample_dict.items():
|
|
65
|
+
if isinstance(v, dict):
|
|
66
|
+
self._write(v, dset[k], first_dim_offset)
|
|
67
|
+
else:
|
|
68
|
+
s = slice(first_dim_offset, first_dim_offset + v.shape[0])
|
|
69
|
+
dset[k][s, ...] = _maybe_to_numpy(v)
|
|
70
|
+
|
|
71
|
+
def write_samples(self, samples_dict, first_dim_offset=0):
|
|
72
|
+
self._write(samples_dict, self._sample_group, first_dim_offset)
|
|
73
|
+
self._file.flush()
|
|
74
|
+
|
|
75
|
+
def write_results(self, results_dict, first_dim_offset=0):
|
|
76
|
+
self._write(results_dict, self._results_group, first_dim_offset)
|
|
77
|
+
self._file.flush()
|
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
"""MultiScanKernel calls one_step a number of times on an inner kernel"""
|
|
2
|
+
|
|
3
|
+
import tensorflow as tf
|
|
4
|
+
import tensorflow_probability as tfp
|
|
5
|
+
from tensorflow_probability.python.internal import samplers
|
|
6
|
+
|
|
7
|
+
mcmc = tfp.mcmc
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class MultiScanKernel(mcmc.TransitionKernel):
|
|
11
|
+
def __init__(self, num_updates, inner_kernel, name=None):
|
|
12
|
+
"""Performs multiple steps of an inner kernel
|
|
13
|
+
returning the state and results after the last step.
|
|
14
|
+
|
|
15
|
+
:param num_updates: integer giving the number of updates
|
|
16
|
+
:param inner_kernel: an instance of a `tfp.mcmc.TransitionKernel`
|
|
17
|
+
"""
|
|
18
|
+
self._parameters = {
|
|
19
|
+
"num_updates": num_updates,
|
|
20
|
+
"inner_kernel": inner_kernel,
|
|
21
|
+
"name": name,
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
@property
|
|
25
|
+
def is_calibrated(self):
|
|
26
|
+
return True
|
|
27
|
+
|
|
28
|
+
@property
|
|
29
|
+
def num_updates(self):
|
|
30
|
+
return self._parameters["num_updates"]
|
|
31
|
+
|
|
32
|
+
@property
|
|
33
|
+
def inner_kernel(self):
|
|
34
|
+
return self._parameters["inner_kernel"]
|
|
35
|
+
|
|
36
|
+
@property
|
|
37
|
+
def name(self):
|
|
38
|
+
return self._parameters["name"]
|
|
39
|
+
|
|
40
|
+
def one_step(self, current_state, prev_results, seed=None):
|
|
41
|
+
seed = samplers.sanitize_seed(seed, salt="MultiScanKernel")
|
|
42
|
+
|
|
43
|
+
def body(i, state, results, seed):
|
|
44
|
+
this_seed, next_seed = samplers.split_seed(seed)
|
|
45
|
+
state, results = self.inner_kernel.one_step(
|
|
46
|
+
state, results, this_seed
|
|
47
|
+
)
|
|
48
|
+
return i + 1, state, results, next_seed
|
|
49
|
+
|
|
50
|
+
def cond(i, *_):
|
|
51
|
+
return i < self.num_updates
|
|
52
|
+
|
|
53
|
+
_, next_state, next_results, _ = tf.while_loop(
|
|
54
|
+
cond, body, (0, current_state, prev_results, seed)
|
|
55
|
+
)
|
|
56
|
+
return next_state, next_results
|
|
57
|
+
|
|
58
|
+
def bootstrap_results(self, current_state):
|
|
59
|
+
return self.inner_kernel.bootstrap_results(current_state)
|
|
@@ -0,0 +1,132 @@
|
|
|
1
|
+
"""Class for writing posterior samples"""
|
|
2
|
+
|
|
3
|
+
import math
|
|
4
|
+
from datetime import datetime
|
|
5
|
+
from typing import Tuple
|
|
6
|
+
|
|
7
|
+
import numpy as np
|
|
8
|
+
import zarr
|
|
9
|
+
|
|
10
|
+
import gemlib
|
|
11
|
+
|
|
12
|
+
__all__ = ["ZarrPosterior"]
|
|
13
|
+
|
|
14
|
+
CHUNK_BASE = 100 * 1024 * 1024 # Multiplier by which chunks are adjusted
|
|
15
|
+
CHUNK_MIN = 128 * 1024 # Soft lower limit (128k)
|
|
16
|
+
CHUNK_MAX = 300 * 1024 * 1024 # Hard upper limit
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def _guess_chunks(shape: Tuple[int, ...], typesize: int) -> Tuple[int, ...]:
|
|
20
|
+
"""Guess an appropriate chunk layout for an array, given its shape and
|
|
21
|
+
the size of each element in bytes. Will allocate chunks only as large
|
|
22
|
+
as MAX_SIZE. Chunks are generally close to some power-of-2 fraction of
|
|
23
|
+
each axis, slightly favoring bigger values for the last index.
|
|
24
|
+
Undocumented and subject to change without warning.
|
|
25
|
+
"""
|
|
26
|
+
ndims = len(shape)
|
|
27
|
+
# require chunks to have non-zero length for all dimensions
|
|
28
|
+
chunks = np.maximum(np.array(shape, dtype="=f8"), 1)
|
|
29
|
+
|
|
30
|
+
# Determine the optimal chunk size in bytes using a PyTables expression.
|
|
31
|
+
# This is kept as a float.
|
|
32
|
+
dset_size = np.product(chunks) * typesize
|
|
33
|
+
target_size = CHUNK_BASE * (2 ** np.log10(dset_size / (1024.0 * 1024)))
|
|
34
|
+
|
|
35
|
+
if target_size > CHUNK_MAX:
|
|
36
|
+
target_size = CHUNK_MAX
|
|
37
|
+
elif target_size < CHUNK_MIN:
|
|
38
|
+
target_size = CHUNK_MIN
|
|
39
|
+
|
|
40
|
+
idx = 0
|
|
41
|
+
while True:
|
|
42
|
+
# Repeatedly loop over the axes, dividing them by 2. Stop when:
|
|
43
|
+
# 1a. We're smaller than the target chunk size, OR
|
|
44
|
+
# 1b. We're within 50% of the target chunk size, AND
|
|
45
|
+
# 2. The chunk is smaller than the maximum chunk size
|
|
46
|
+
chunk_bytes = np.product(chunks) * typesize
|
|
47
|
+
if (
|
|
48
|
+
chunk_bytes < target_size
|
|
49
|
+
or abs(chunk_bytes - target_size) / target_size < 0.5 # noqa: PLR2004
|
|
50
|
+
) and chunk_bytes < CHUNK_MAX:
|
|
51
|
+
break
|
|
52
|
+
|
|
53
|
+
if np.product(chunks) == 1:
|
|
54
|
+
break # Element size larger than CHUNK_MAX
|
|
55
|
+
|
|
56
|
+
chunks[idx % ndims] = math.ceil(chunks[idx % ndims] / 2.0)
|
|
57
|
+
idx += 1
|
|
58
|
+
|
|
59
|
+
return tuple(int(x) for x in chunks)
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _maybe_tf_dtype(dtype):
|
|
63
|
+
if hasattr(dtype, "as_numpy_dtype"):
|
|
64
|
+
return dtype.as_numpy_dtype
|
|
65
|
+
return dtype
|
|
66
|
+
|
|
67
|
+
|
|
68
|
+
def _maybe_to_numpy(val):
|
|
69
|
+
if hasattr(val, "numpy"):
|
|
70
|
+
return val.numpy()
|
|
71
|
+
return val
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
class ZarrPosterior:
|
|
75
|
+
def __init__(self, filename, sample_dict, results_dict, num_samples):
|
|
76
|
+
"""Constructs a posterior output object
|
|
77
|
+
|
|
78
|
+
:param filename: the name of the backend HDF5 file
|
|
79
|
+
:param sample_dict: a dictionary containing `key`:`shape_tuple`
|
|
80
|
+
:param results_dict: a dictionary containing `key`:`shape_tuple`
|
|
81
|
+
:param num_samples: total number of samples
|
|
82
|
+
"""
|
|
83
|
+
self._num_samples = num_samples
|
|
84
|
+
self._archive = zarr.open(
|
|
85
|
+
filename,
|
|
86
|
+
"w",
|
|
87
|
+
)
|
|
88
|
+
|
|
89
|
+
self._sample_group = self._archive.create_group("samples")
|
|
90
|
+
self._create_data_tree(sample_dict, self._sample_group)
|
|
91
|
+
|
|
92
|
+
self._results_group = self._archive.create_group("results")
|
|
93
|
+
self._create_data_tree(results_dict, self._results_group)
|
|
94
|
+
|
|
95
|
+
self._archive.attrs["created_at"] = str(datetime.now())
|
|
96
|
+
self._archive.attrs["inference_library"] = "gemlib"
|
|
97
|
+
self._archive.attrs["inference_library_version"] = gemlib.__version__
|
|
98
|
+
|
|
99
|
+
def __getitem__(self, path):
|
|
100
|
+
return self._archive[path]
|
|
101
|
+
|
|
102
|
+
def _create_data_tree(self, data_dict, group):
|
|
103
|
+
for k, v in data_dict.items():
|
|
104
|
+
if isinstance(v, dict):
|
|
105
|
+
subgroup = group.create_group(k)
|
|
106
|
+
self._create_data_tree(v, subgroup)
|
|
107
|
+
|
|
108
|
+
else:
|
|
109
|
+
dtype = _maybe_tf_dtype(v.dtype)
|
|
110
|
+
chunks = _guess_chunks(
|
|
111
|
+
shape=(self._num_samples,) + v.shape,
|
|
112
|
+
typesize=np.dtype(dtype).itemsize,
|
|
113
|
+
)
|
|
114
|
+
dset_shape = (0,) + v.shape
|
|
115
|
+
|
|
116
|
+
group.create_dataset(
|
|
117
|
+
k,
|
|
118
|
+
shape=dset_shape,
|
|
119
|
+
chunks=chunks,
|
|
120
|
+
dtype=dtype,
|
|
121
|
+
)
|
|
122
|
+
|
|
123
|
+
def _append(self, sample_dict, dset):
|
|
124
|
+
for k, v in sample_dict.items():
|
|
125
|
+
if isinstance(v, dict):
|
|
126
|
+
self._append(v, dset[k])
|
|
127
|
+
else:
|
|
128
|
+
dset[k].append(_maybe_to_numpy(v))
|
|
129
|
+
|
|
130
|
+
def append(self, samples_dict, results_dict):
|
|
131
|
+
self._append(samples_dict, self._sample_group)
|
|
132
|
+
self._append(results_dict, self._results_group)
|
gemlib/util.py
ADDED
|
@@ -0,0 +1,117 @@
|
|
|
1
|
+
"""Utility functions for model implementation code."""
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import tensorflow as tf
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def which(predicate):
|
|
8
|
+
"""Return the indices of True elements of `predicate`."""
|
|
9
|
+
with tf.name_scope("which"):
|
|
10
|
+
x = tf.cast(predicate, dtype=tf.int32)
|
|
11
|
+
index_range = tf.range(x.shape[0])
|
|
12
|
+
indices = tf.cumsum(x) * x
|
|
13
|
+
indices = tf.scatter_nd(indices[:, None], index_range, x.shape)
|
|
14
|
+
return indices[1:]
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def batch_gather(arr, indices):
|
|
18
|
+
"""Gather `indices` from the right-most dimensions of `arr`.
|
|
19
|
+
|
|
20
|
+
This function gathers elements on the right-most `indices` of `tensor`
|
|
21
|
+
|
|
22
|
+
Args
|
|
23
|
+
----
|
|
24
|
+
arr: an N-dimensional tensor
|
|
25
|
+
indices: an iterable of N-dimensional coordinates into the rightmost
|
|
26
|
+
`indices.shape[-1]` dimensions of `arr`
|
|
27
|
+
|
|
28
|
+
Returns
|
|
29
|
+
-------
|
|
30
|
+
A tensor of dimension `rank(arr) - indices.shape[-1]` of gathered values in
|
|
31
|
+
`arr`.
|
|
32
|
+
"""
|
|
33
|
+
|
|
34
|
+
arr = tf.convert_to_tensor(arr)
|
|
35
|
+
# TF shapes and indices are 32 bit
|
|
36
|
+
indices = tf.cast(tf.convert_to_tensor(indices), tf.int32)
|
|
37
|
+
|
|
38
|
+
index_dims = indices.shape[-1]
|
|
39
|
+
|
|
40
|
+
# Flatten the dims which we are indexing - this is cheap, as no data needs
|
|
41
|
+
# to be copied. `flat_arr` is just a "view" of `arr`.
|
|
42
|
+
flat_shape = arr.shape[:-index_dims].as_list() + [
|
|
43
|
+
np.prod(arr.shape[-index_dims:])
|
|
44
|
+
]
|
|
45
|
+
flat_arr = tf.reshape(arr, shape=flat_shape)
|
|
46
|
+
|
|
47
|
+
# Compute the stride for each dim in the indices
|
|
48
|
+
flat_coord_stride = tf.math.cumprod(
|
|
49
|
+
tf.concat(
|
|
50
|
+
[arr.shape[arr.shape.rank - (index_dims - 1) :], [1]], axis=0
|
|
51
|
+
),
|
|
52
|
+
axis=0,
|
|
53
|
+
reverse=True,
|
|
54
|
+
)
|
|
55
|
+
flat_indices = tf.linalg.matvec(indices, flat_coord_stride)
|
|
56
|
+
|
|
57
|
+
return tf.gather(flat_arr, flat_indices, axis=-1)
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def transition_coords(incidence_matrix):
|
|
61
|
+
"""Compute coordinates of transitions in a Markov transition matrix
|
|
62
|
+
|
|
63
|
+
Args
|
|
64
|
+
----
|
|
65
|
+
incidence_matrix: a (batch of) `[S, R]` matrix describing R
|
|
66
|
+
transitions between S states.
|
|
67
|
+
|
|
68
|
+
Returns
|
|
69
|
+
-------
|
|
70
|
+
a [..., R, 2] tensor of coordinates of the transitions in a square
|
|
71
|
+
transition matrix.
|
|
72
|
+
"""
|
|
73
|
+
|
|
74
|
+
incidence_matrix = tf.convert_to_tensor(incidence_matrix)
|
|
75
|
+
|
|
76
|
+
is_src_dest = tf.stack(
|
|
77
|
+
[incidence_matrix < 0, incidence_matrix > 0], axis=-1
|
|
78
|
+
)
|
|
79
|
+
|
|
80
|
+
coords = tf.reduce_sum(
|
|
81
|
+
tf.cumsum(
|
|
82
|
+
tf.cast(is_src_dest, tf.int64),
|
|
83
|
+
exclusive=True,
|
|
84
|
+
reverse=True,
|
|
85
|
+
axis=-3,
|
|
86
|
+
),
|
|
87
|
+
axis=-3,
|
|
88
|
+
)
|
|
89
|
+
|
|
90
|
+
return coords
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def states_from_transition_idx(
|
|
94
|
+
transition_index, incidence_matrix, output_type=tf.int32
|
|
95
|
+
):
|
|
96
|
+
"""Return source and destination state indices given a transition index.
|
|
97
|
+
|
|
98
|
+
Given the index of a transition in `stoichiometry`, return
|
|
99
|
+
the indices of the source and destination states.
|
|
100
|
+
|
|
101
|
+
Note: this algorithm depends on the stoichiometry matrix
|
|
102
|
+
describing a state transition model and taking the values `[-1, 0, 1]`.
|
|
103
|
+
|
|
104
|
+
Args:
|
|
105
|
+
----
|
|
106
|
+
event_index: the index (row id) of the event in `stoichiometry`
|
|
107
|
+
incidence_matrix: a `[S, R]` matrix relating transitions to states
|
|
108
|
+
|
|
109
|
+
Returns:
|
|
110
|
+
-------
|
|
111
|
+
a tuple of integers denoting indices of `(src, dest)`.
|
|
112
|
+
|
|
113
|
+
"""
|
|
114
|
+
transition_index = tf.convert_to_tensor(transition_index, tf.int32)
|
|
115
|
+
coords = transition_coords(incidence_matrix)[..., transition_index, :]
|
|
116
|
+
|
|
117
|
+
return coords[..., 0], coords[..., 1]
|
gemlib/util_test.py
ADDED
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
"""Test `gemlib` utility functions."""
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import pytest
|
|
5
|
+
|
|
6
|
+
from gemlib.util import (
|
|
7
|
+
batch_gather,
|
|
8
|
+
states_from_transition_idx,
|
|
9
|
+
transition_coords,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@pytest.fixture
|
|
14
|
+
def svir_incidence():
|
|
15
|
+
"""Fixture for SVIR incidence matrix."""
|
|
16
|
+
|
|
17
|
+
return np.array(
|
|
18
|
+
[ # SI SV VI IR
|
|
19
|
+
[-1, -1, 0, 0], # S
|
|
20
|
+
[0, 1, -1, 0], # V
|
|
21
|
+
[1, 0, 1, -1], # I
|
|
22
|
+
[0, 0, 0, 1], # R
|
|
23
|
+
]
|
|
24
|
+
)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@pytest.fixture
|
|
28
|
+
def sirs_incidence():
|
|
29
|
+
"""Fixture for SIRS incidence matrix."""
|
|
30
|
+
|
|
31
|
+
return np.array(
|
|
32
|
+
[ # SI IR RS
|
|
33
|
+
[-1, 0, 1], # S
|
|
34
|
+
[1, -1, 0], # I
|
|
35
|
+
[0, 1, -1], # R
|
|
36
|
+
],
|
|
37
|
+
)
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def test_transition_coords_svir(svir_incidence):
|
|
41
|
+
# Test svir
|
|
42
|
+
coords = transition_coords(svir_incidence)
|
|
43
|
+
expected = np.array([[0, 2], [0, 1], [1, 2], [2, 3]])
|
|
44
|
+
np.testing.assert_equal(coords, expected)
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
def test_transition_coords_sirs(sirs_incidence):
|
|
48
|
+
# Test SIRS
|
|
49
|
+
coords = transition_coords(sirs_incidence)
|
|
50
|
+
expected = np.array([[0, 1], [1, 2], [2, 0]])
|
|
51
|
+
np.testing.assert_equal(coords, expected)
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def test_states_from_transition_idx(svir_incidence):
|
|
55
|
+
"""Ensure source and destination enums are correct."""
|
|
56
|
+
# S->I
|
|
57
|
+
assert states_from_transition_idx(0, svir_incidence) == (0, 2)
|
|
58
|
+
# S->V
|
|
59
|
+
assert states_from_transition_idx(1, svir_incidence) == (0, 1)
|
|
60
|
+
# V->I
|
|
61
|
+
assert states_from_transition_idx(2, svir_incidence) == (1, 2)
|
|
62
|
+
# I->R
|
|
63
|
+
assert states_from_transition_idx(3, svir_incidence) == (2, 3)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def test_batch_gather():
|
|
67
|
+
arr = np.random.uniform(low=0, high=1, size=[20, 100, 30, 20])
|
|
68
|
+
|
|
69
|
+
indices = np.array([[2, 3], [4, 7], [15, 10]])
|
|
70
|
+
|
|
71
|
+
slice_arr = batch_gather(arr, indices)
|
|
72
|
+
|
|
73
|
+
np.testing.assert_array_equal(
|
|
74
|
+
slice_arr, arr[..., indices[:, 0], indices[:, 1]]
|
|
75
|
+
)
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
MIT License
|
|
2
|
+
|
|
3
|
+
Copyright (c) 2020 The GEM Authors. All rights reserved.
|
|
4
|
+
|
|
5
|
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
|
6
|
+
of this software and associated documentation files (the "Software"), to deal
|
|
7
|
+
in the Software without restriction, including without limitation the rights
|
|
8
|
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
|
9
|
+
copies of the Software, and to permit persons to whom the Software is
|
|
10
|
+
furnished to do so, subject to the following conditions:
|
|
11
|
+
|
|
12
|
+
The above copyright notice and this permission notice shall be included in all
|
|
13
|
+
copies or substantial portions of the Software.
|
|
14
|
+
|
|
15
|
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
|
16
|
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
|
17
|
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
|
18
|
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
|
19
|
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
|
20
|
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
|
21
|
+
SOFTWARE.
|