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,279 @@
1
+ """Functions for chain binomial simulation."""
2
+
3
+ import numpy as np
4
+ import tensorflow as tf
5
+ import tensorflow_probability as tfp
6
+ from tensorflow_probability.python.internal import prefer_static as ps
7
+ from tensorflow_probability.python.internal import samplers
8
+ from tensorflow_probability.python.mcmc.internal import util as mcmc_util
9
+
10
+ from gemlib.util import transition_coords
11
+
12
+ tfd = tfp.distributions
13
+
14
+
15
+ def _gen_index(state_shape, trm_coords):
16
+ """Generate indices for broadcasting transition rates."""
17
+ trm_coords = tf.convert_to_tensor(trm_coords)
18
+
19
+ i_shp = state_shape[:-1] + [trm_coords.shape[0]] + [len(state_shape) + 1]
20
+
21
+ b_idx = np.array(list(np.ndindex(*i_shp[:-1])))[:, :-1]
22
+ m_idx = tf.tile(trm_coords, [tf.reduce_prod(i_shp[:-2]), 1])
23
+
24
+ idx = tf.concat([b_idx, m_idx], axis=-1)
25
+ return tf.reshape(idx, i_shp)
26
+
27
+
28
+ def _make_transition_matrix(rates, rate_coords, state_shape):
29
+ """Create a transition rate matrix.
30
+
31
+ Args
32
+ rates: batched transition rate tensors [b1, b2, n_rates] or a list of
33
+ length n_rates of batched tensors [b1, b2]
34
+ rate_coords: coordinates of rates in resulting transition matrix
35
+ state_shape: the shape of the state tensor with ns states
36
+ Returns
37
+ a tensor of shape [..., ns, ns]
38
+ """
39
+ indices = _gen_index(state_shape, rate_coords)
40
+ if mcmc_util.is_list_like(rates):
41
+ rates = tf.stack(rates, axis=-1)
42
+ output_shape = state_shape + [state_shape[-1]]
43
+ rate_tensor = tf.scatter_nd(
44
+ indices=indices,
45
+ updates=rates,
46
+ shape=output_shape,
47
+ name="build_markov_matrix",
48
+ )
49
+ return rate_tensor
50
+
51
+
52
+ def compute_state(initial_state, events, incidence_matrix, closed=False):
53
+ """Compute a state tensor from initial state and event tensor.
54
+
55
+ Args
56
+ ----
57
+ initial_state: a tensor of shape [M, S]
58
+ events: a tensor of shape [M, T, R]
59
+ incidence_matrix: a incidence_matrix matrix of shape [S,R] describing
60
+ how transitions update the state.
61
+ closed: if `True`, return state in close interval [0, T], otherwise
62
+ [0, T)
63
+
64
+ Returns
65
+ -------
66
+ a tensor of shape [M, T, S] if `closed=False` or [M, T+1, S] if
67
+ `closed=True` describing the state of the system for each batch
68
+ M at time T.
69
+ """
70
+ if isinstance(incidence_matrix, tf.Tensor):
71
+ incidence_matrix = ps.cast(incidence_matrix, dtype=events.dtype)
72
+ else:
73
+ incidence_matrix = tf.convert_to_tensor(
74
+ incidence_matrix, dtype=events.dtype
75
+ )
76
+ increments = tf.einsum("...tr,sr->...ts", events, incidence_matrix)
77
+
78
+ if closed is False:
79
+ cum_increments = tf.cumsum(increments, axis=-2, exclusive=True)
80
+ else:
81
+ cum_increments = tf.cumsum(increments, axis=-2, exclusive=False)
82
+ cum_increments = tf.concat(
83
+ [tf.zeros_like(cum_increments[..., 0:1, :]), cum_increments],
84
+ axis=-2,
85
+ )
86
+ state = cum_increments + tf.expand_dims(initial_state, axis=-2)
87
+ return state
88
+
89
+
90
+ def approx_expm(rates):
91
+ """Approximates a full Markov transition matrix
92
+ :param rates: un-normalised rate matrix (i.e. diagonal zero)
93
+ :returns: approximation to Markov transition matrix
94
+ """
95
+ total_rates = tf.reduce_sum(rates, axis=-1, keepdims=True)
96
+ prob = 1.0 - tf.math.exp(-tf.reduce_sum(rates, axis=-1, keepdims=True))
97
+ mt1 = tf.math.multiply_no_nan(rates / total_rates, prob)
98
+ return tf.linalg.set_diag(mt1, 1.0 - tf.reduce_sum(mt1, axis=-1))
99
+
100
+
101
+ def chain_binomial_propagate(h, time_step, incidence_matrix):
102
+ """Propagates the state of a population according to discrete time dynamics.
103
+
104
+ :param h: a hazard rate function returning the non-row-normalised Markov
105
+ transition rate matrix. This function should return a list of
106
+ length R equal to the number of transitions, with each element a
107
+ tensor of shape `[M]` where `M` is the number of population units.
108
+ :param time_step: the time step
109
+ :param incidence_matrix: a `[S, R]` tensor giving the state transition graph
110
+ :returns : a function that propagate `state[t]` -> `state[t+time_step]`
111
+ """
112
+
113
+ def propagate_fn(t, state, seed):
114
+ rates = h(t, state)
115
+
116
+ # `rate_matrix` needs to be a tensor of shape
117
+ # `[M, S, S]` where `M` is the number of population units,
118
+ # and `S` is the number of states. Then, `rate_matrix[m, i, j]`
119
+ # gives the transition rate for transitioning from state `i` to
120
+ # state `j` in unit `m`.
121
+ rate_matrix = _make_transition_matrix(
122
+ rates, transition_coords(incidence_matrix), state.shape
123
+ )
124
+ # Set diagonal to be the negative of the sum of other elements in
125
+ # each row
126
+ markov_transition = approx_expm(rate_matrix * time_step)
127
+ num_states = markov_transition.shape[-1]
128
+ prev_probs = tf.zeros_like(markov_transition[..., :, 0])
129
+ counts = tf.zeros(
130
+ markov_transition.shape[:-1].as_list() + [0],
131
+ dtype=markov_transition.dtype,
132
+ )
133
+ total_count = state
134
+ # This for loop is ok because there are (currently) only 4 states (SEIR)
135
+ # and we're only actually creating work for 3 of them. Even for as many
136
+ # as a ~10 states it should probably be fine, just increasing the size
137
+ # of the graph a bit.
138
+ seeds = samplers.split_seed(seed, n=num_states - 1, salt="propagate_fn")
139
+ for i in range(num_states - 1):
140
+ probs = markov_transition[..., :, i]
141
+ binom = tfd.Binomial(
142
+ total_count=total_count,
143
+ probs=tf.clip_by_value(probs / (1.0 - prev_probs), 0.0, 1.0),
144
+ )
145
+ sample = binom.sample(seed=seeds[i])
146
+ counts = tf.concat([counts, sample[..., tf.newaxis]], axis=-1)
147
+ total_count -= sample
148
+ prev_probs += probs
149
+
150
+ counts = tf.concat([counts, total_count[..., tf.newaxis]], axis=-1)
151
+
152
+ # Counts is a `[M, S, S]` tensor, where each inner dimension represents
153
+ # a draw from a Multinomial random variable. Each element
154
+ # `counts[m, i, j]` gives the number of transitions from state `i` to
155
+ # state `j` in each unit `m`. We now sum over the `i` axis to get the
156
+ # new state.
157
+ new_state = tf.reduce_sum(counts, axis=-2)
158
+
159
+ # `new_state` is of shape `[M, S]`
160
+ return counts, new_state
161
+
162
+ return propagate_fn
163
+
164
+
165
+ def discrete_markov_simulation(
166
+ hazard_fn, state, start, end, time_step, incidence_matrix, seed=None
167
+ ):
168
+ """Simulates from a discrete time Markov state transition model using
169
+ multinomial sampling across rows of the transition matrix"""
170
+ state = tf.convert_to_tensor(state)
171
+
172
+ propagate = chain_binomial_propagate(hazard_fn, time_step, incidence_matrix)
173
+
174
+ times = tf.range(start, end, time_step, dtype=state.dtype)
175
+ state = tf.convert_to_tensor(state)
176
+
177
+ output = tf.TensorArray(state.dtype, size=times.shape[0])
178
+
179
+ def cond(i, *_):
180
+ return i < times.shape[0]
181
+
182
+ def body(i, state, output, seed):
183
+ seed, next_seed = samplers.split_seed(seed)
184
+ event_counts, state = propagate(times[i], state, seed)
185
+ output = output.write(i, event_counts)
186
+ return i + 1, state, output, next_seed
187
+
188
+ _, state, output, _ = tf.while_loop(
189
+ cond, body, loop_vars=(0, state, output, seed)
190
+ )
191
+
192
+ # `output.stack()` returns a `[T, M, S, S]` tensor of event numbers.
193
+ return times, output.stack()
194
+
195
+
196
+ def discrete_markov_log_prob(
197
+ events, init_state, init_step, time_delta, hazard_fn, incidence_matrix
198
+ ):
199
+ """Calculates an unnormalised log_prob function for a discrete time epidemic
200
+ model.
201
+
202
+ :param events: a `[M, T, X]` batch of transition events for metapopulation
203
+ `M` times `T`, and transitions `X`.
204
+ :param init_state: a vector of shape `[M, S]` the initial state of the
205
+ epidemic for `M` metapopulations and `S` states
206
+ :param init_step: the initial time step, as an offset to
207
+ `range(events.shape[-2])`
208
+ :param time_delta: the size of the time step.
209
+ :param hazard_fn: a function that takes a state and returns a matrix of
210
+ transition rates.
211
+ :param incidence_matrix: a `[S, R]` matrix describing the state update for
212
+ each transition.
213
+ :return: a scalar log probability for the epidemic.
214
+ """
215
+ num_meta = events.shape[-3]
216
+ num_times = events.shape[-2]
217
+ num_states = incidence_matrix.shape[-2]
218
+
219
+ state_timeseries = compute_state(
220
+ init_state, events, incidence_matrix
221
+ ) # MxTxS
222
+
223
+ tms_timeseries = tf.transpose(state_timeseries, perm=(1, 0, 2))
224
+
225
+ def fn(elems):
226
+ return hazard_fn(*elems)
227
+
228
+ tx_coords = transition_coords(incidence_matrix)
229
+ rates = tf.vectorized_map(
230
+ fn=fn,
231
+ elems=[
232
+ tf.range(
233
+ init_step, time_delta * num_times + init_step, delta=time_delta
234
+ ),
235
+ tms_timeseries,
236
+ ],
237
+ )
238
+ rate_matrix = _make_transition_matrix(
239
+ rates, tx_coords, tms_timeseries.shape
240
+ )
241
+ probs = approx_expm(rate_matrix * time_delta)
242
+
243
+ # [T, M, S, S] to [M, T, S, S]
244
+ probs = tf.transpose(probs, perm=(1, 0, 2, 3))
245
+ event_matrix = _make_transition_matrix(
246
+ events, tx_coords, [num_meta, num_times, num_states]
247
+ )
248
+ event_matrix = tf.linalg.set_diag(
249
+ event_matrix, state_timeseries - tf.reduce_sum(event_matrix, axis=-1)
250
+ )
251
+
252
+ logp = tfd.Multinomial(
253
+ total_count=state_timeseries,
254
+ # logits=logits,
255
+ probs=probs + 1.0e-9,
256
+ name="log_prob",
257
+ ).log_prob(event_matrix)
258
+ return tf.reduce_sum(logp)
259
+
260
+
261
+ def events_to_full_transitions(events, initial_state):
262
+ """Creates a state tensor given matrices of transition events
263
+ and the initial state
264
+
265
+ :param events: a tensor of shape [t, c, s, s] for t timepoints, c
266
+ metapopulations and s states.
267
+ :param initial_state: the initial state matrix of shape [c, s]
268
+ """
269
+
270
+ def f(state, events):
271
+ survived = tf.reduce_sum(state, axis=-2) - tf.reduce_sum(
272
+ events, axis=-1
273
+ )
274
+ new_state = tf.linalg.set_diag(events, survived)
275
+ return new_state
276
+
277
+ return tf.scan(
278
+ fn=f, elems=events, initializer=tf.linalg.diag(initial_state)
279
+ )
@@ -0,0 +1,149 @@
1
+ # Copyright 2020 The TensorFlow Probability Authors.
2
+ #
3
+ # Licensed under the Apache License, Version 2.0 (the "License");
4
+ # you may not use this file except in compliance with the License.
5
+ # You may obtain a copy of the License at
6
+ #
7
+ # http://www.apache.org/licenses/LICENSE-2.0
8
+ #
9
+ # Unless required by applicable law or agreed to in writing, software
10
+ # distributed under the License is distributed on an "AS IS" BASIS,
11
+ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12
+ # See the License for the specific language governing permissions and
13
+ # limitations under the License.
14
+ # ============================================================================
15
+ """Batched discrete rejection samplers."""
16
+
17
+ import numpy as np
18
+ import tensorflow.compat.v2 as tf
19
+ from tensorflow_probability.python import random as tfp_random
20
+ from tensorflow_probability.python.distributions import exponential
21
+ from tensorflow_probability.python.internal import (
22
+ batched_rejection_sampler as brs,
23
+ )
24
+ from tensorflow_probability.python.internal import prefer_static as ps
25
+ from tensorflow_probability.python.internal import samplers
26
+
27
+ __all__ = [
28
+ "log_concave_rejection_sampler",
29
+ ]
30
+
31
+
32
+ def log_concave_rejection_sampler(
33
+ mode,
34
+ prob_fn,
35
+ dtype,
36
+ sample_shape=(),
37
+ distribution_minimum=None,
38
+ distribution_maximum=None,
39
+ seed=None,
40
+ ):
41
+ """Utility for rejection sampling from log-concave discrete distributions.
42
+
43
+ This utility constructs an easy-to-sample-from upper bound for a discrete
44
+ univariate log-concave distribution (for discrete univariate distributions,
45
+ a necessary and sufficient condition is p_k^2 >= p_{k-1} p_{k+1} for all k).
46
+ The method requires that the mode of the distribution is known. While a
47
+ better method can likely be derived for any given distribution, this method
48
+ is general and easy to implement. The expected number of iterations is
49
+ bounded by 4+m, where m is the probability of the mode. For details, see
50
+ [(Devroye, 1979)][1].
51
+
52
+ Args:
53
+ ----
54
+ mode: Tensor, the mode[s] of the [batch of] distribution[s].
55
+ prob_fn: Python callable, counts -> prob(counts).
56
+ dtype: DType of the generated samples.
57
+ sample_shape: 0D or 1D `int32` `Tensor`. Shape of the generated samples.
58
+ distribution_minimum: Tensor of type `dtype`. The minimum value
59
+ taken by the distribution. The `prob` method will only be called on
60
+ values greater than equal to the specified minimum. The shape must
61
+ broadcast with the batch shape of the distribution. If unspecified, the
62
+ domain is treated as unbounded below.
63
+ distribution_maximum: Tensor of type `dtype`. The maximum value
64
+ taken by the distribution. See `distribution_minimum` for details.
65
+ seed: PRNG seed; see `tfp.random.sanitize_seed` for details.
66
+
67
+ Returns:
68
+ -------
69
+ samples: a `Tensor` with prepended dimensions `sample_shape`.
70
+
71
+ #### References
72
+
73
+ [1] Luc Devroye. A Simple Generator for Discrete Log-Concave
74
+ Distributions. Computing, 1987.
75
+
76
+ """
77
+ mode = tf.broadcast_to(
78
+ mode, ps.concat([sample_shape, ps.shape(mode)], axis=0)
79
+ )
80
+
81
+ mode_height = prob_fn(mode)
82
+ mode_shape = ps.shape(mode)
83
+
84
+ top_width = 1.0 + mode_height / 2.0 # w in ref [1].
85
+ top_fraction = top_width / (1 + top_width)
86
+ exponential_distribution = exponential.Exponential(
87
+ rate=tf.ones([], dtype=dtype)
88
+ ) # E in ref [1].
89
+
90
+ if distribution_minimum is None:
91
+ distribution_minimum = tf.constant(-np.inf, dtype)
92
+ if distribution_maximum is None:
93
+ distribution_maximum = tf.constant(np.inf, dtype)
94
+
95
+ def proposal(seed):
96
+ """Proposal for log-concave rejection sampler."""
97
+ (
98
+ top_lobe_fractions_seed,
99
+ exponential_samples_seed,
100
+ top_selector_seed,
101
+ rademacher_seed,
102
+ ) = samplers.split_seed(seed, n=4)
103
+
104
+ top_lobe_fractions = samplers.uniform(
105
+ mode_shape, seed=top_lobe_fractions_seed, dtype=dtype
106
+ ) # V in ref [1].
107
+ top_offsets = top_lobe_fractions * top_width / mode_height
108
+
109
+ exponential_samples = exponential_distribution.sample(
110
+ mode_shape, seed=exponential_samples_seed
111
+ ) # E in ref [1].
112
+ exponential_height = (
113
+ exponential_distribution.prob(exponential_samples) * mode_height
114
+ )
115
+ exponential_offsets = (top_width + exponential_samples) / mode_height
116
+
117
+ top_selector = samplers.uniform(
118
+ mode_shape, seed=top_selector_seed, dtype=dtype
119
+ ) # U in ref [1].
120
+ on_top_mask = top_selector <= top_fraction
121
+
122
+ unsigned_offsets = tf.where(
123
+ on_top_mask, top_offsets, exponential_offsets
124
+ )
125
+ offsets = tf.round(
126
+ tfp_random.rademacher(mode_shape, seed=rademacher_seed, dtype=dtype)
127
+ * unsigned_offsets
128
+ )
129
+
130
+ potential_samples = mode + offsets
131
+ envelope_height = tf.where(on_top_mask, mode_height, exponential_height)
132
+
133
+ return potential_samples, envelope_height
134
+
135
+ def target(values):
136
+ # Check for out of bounds rather than in bounds to avoid accidentally
137
+ # masking a `nan` value.
138
+ out_of_bounds_mask = (values < distribution_minimum) | (
139
+ values > distribution_maximum
140
+ )
141
+ in_bounds_values = tf.where(
142
+ out_of_bounds_mask, tf.constant(0.0, dtype=values.dtype), values
143
+ )
144
+ probs = prob_fn(in_bounds_values)
145
+ return tf.where(out_of_bounds_mask, tf.zeros([], probs.dtype), probs)
146
+
147
+ return tf.stop_gradient(
148
+ brs.batched_rejection_sampler(proposal, target, seed, dtype=dtype)[0]
149
+ ) # Discard `num_iters`.