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