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,254 @@
1
+ """Sampler for discrete-space occult events"""
2
+
3
+ from typing import NamedTuple, Tuple
4
+ from warnings import warn
5
+
6
+ import tensorflow as tf
7
+ import tensorflow_probability as tfp
8
+ from tensorflow_probability.python.internal import samplers
9
+ from tensorflow_probability.python.mcmc.internal import util as mcmc_util
10
+
11
+ from gemlib.mcmc.discrete_time_state_transition_model.right_censored_events_proposal import ( # noqa:E501
12
+ add_occult_proposal,
13
+ del_occult_proposal,
14
+ )
15
+
16
+ tfd = tfp.distributions
17
+
18
+ __all__ = ["UncalibratedOccultUpdate"]
19
+
20
+
21
+ PROB_DIRECTION = 0.5
22
+
23
+
24
+ class OccultKernelResults(NamedTuple):
25
+ log_acceptance_correction: float
26
+ target_log_prob: float
27
+ m: int
28
+ t: int
29
+ delta_t: int
30
+ x_star: int
31
+ seed: Tuple[int, int]
32
+
33
+
34
+ def _nonzero_rows(m):
35
+ return tf.cast(tf.reduce_sum(m, axis=-1) > 0.0, m.dtype)
36
+
37
+
38
+ def _maybe_expand_dims(x):
39
+ """If x is a scalar, give it at least 1 dimension"""
40
+ x = tf.convert_to_tensor(x)
41
+ if x.shape == ():
42
+ return tf.expand_dims(x, axis=0)
43
+ return x
44
+
45
+
46
+ def _add_events(events, m, t, x, x_star):
47
+ """Adds `x_star` events to metapopulation `m`,
48
+ time `t`, transition `x` in `events`.
49
+ """
50
+ x = _maybe_expand_dims(x)
51
+ indices = tf.stack([m, t, x], axis=-1)
52
+ return tf.tensor_scatter_nd_add(events, indices, x_star)
53
+
54
+
55
+ class UncalibratedOccultUpdate(tfp.mcmc.TransitionKernel):
56
+ """UncalibratedOccultUpdate"""
57
+
58
+ def __init__(
59
+ self,
60
+ target_log_prob_fn,
61
+ topology,
62
+ cumulative_event_offset,
63
+ nmax,
64
+ t_range=None,
65
+ name=None,
66
+ ):
67
+ """An uncalibrated random walk for event times.
68
+ :param target_log_prob_fn: the log density of the target distribution
69
+ :param target_event_id: the position in the last dimension of the events
70
+ tensor that we wish to move
71
+ :param t_range: a tuple containing earliest and latest times between
72
+ which to update occults.
73
+ :param seed: a random seed
74
+ :param name: the name of the update step
75
+ """
76
+ self._name = name or "uncalibrated_occult_update"
77
+ self._parameters = {
78
+ "target_log_prob_fn": target_log_prob_fn,
79
+ "topology": topology,
80
+ "cumulative_event_offset": cumulative_event_offset,
81
+ "nmax": nmax,
82
+ "t_range": t_range,
83
+ "name": name,
84
+ }
85
+ self.tx_topology = topology
86
+ self.initial_state = cumulative_event_offset
87
+ self._dtype = self.initial_state.dtype
88
+
89
+ @property
90
+ def target_log_prob_fn(self):
91
+ return self._parameters["target_log_prob_fn"]
92
+
93
+ @property
94
+ def target_event_id(self):
95
+ return self._parameters["topology"]["target_transition"]
96
+
97
+ @property
98
+ def name(self):
99
+ return self._parameters["name"]
100
+
101
+ @property
102
+ def parameters(self):
103
+ """Return `dict` of ``__init__`` arguments and their values."""
104
+ return self._parameters
105
+
106
+ @property
107
+ def is_calibrated(self):
108
+ return False
109
+
110
+ def one_step(self, current_events, previous_kernel_results, seed=None):
111
+ """One update of event times.
112
+ :param current_events: a [M, T, X] tensor containing number of events
113
+ per time t, metapopulation m,
114
+ and transition x.
115
+ :param previous_kernel_results: an object of type
116
+ UncalibratedRandomWalkResults.
117
+ :returns: a tuple containing new_state and UncalibratedRandomWalkResults
118
+ """
119
+ with tf.name_scope("occult_rw/onestep"):
120
+ seed = samplers.sanitize_seed(seed, salt="occult_rw")
121
+ proposal_seed, add_del_seed = samplers.split_seed(seed)
122
+
123
+ if mcmc_util.is_list_like(current_events):
124
+ step_events = current_events[0]
125
+ warn(
126
+ "Batched updating of occults is not supported.",
127
+ stacklevel=2,
128
+ )
129
+ else:
130
+ step_events = current_events
131
+
132
+ def add_occult_fn():
133
+ with tf.name_scope("true_fn"):
134
+ proposal = add_occult_proposal(
135
+ events=step_events,
136
+ topology=self.tx_topology,
137
+ initial_state=self.initial_state,
138
+ n_max=self.parameters["nmax"],
139
+ t_range=self.parameters["t_range"],
140
+ name=self.name,
141
+ )
142
+ update = proposal.sample(seed=proposal_seed)
143
+ next_state = _add_events(
144
+ events=step_events,
145
+ m=update["m"],
146
+ t=update["t"],
147
+ x=self.tx_topology.target,
148
+ x_star=tf.cast(update["x_star"], step_events.dtype),
149
+ )
150
+ reverse = del_occult_proposal(
151
+ events=next_state,
152
+ topology=self.tx_topology,
153
+ initial_state=self.initial_state,
154
+ t_range=self.parameters["t_range"],
155
+ n_max=self.parameters["nmax"],
156
+ )
157
+ q_fwd = tf.reduce_sum(proposal.log_prob(update))
158
+ q_rev = tf.reduce_sum(reverse.log_prob(update))
159
+ log_acceptance_correction = q_rev - q_fwd
160
+
161
+ return (
162
+ update,
163
+ next_state,
164
+ log_acceptance_correction,
165
+ tf.ones(1, dtype=tf.int32),
166
+ )
167
+
168
+ def del_occult_fn():
169
+ with tf.name_scope("false_fn"):
170
+ proposal = del_occult_proposal(
171
+ events=step_events,
172
+ topology=self.tx_topology,
173
+ initial_state=self.initial_state,
174
+ t_range=self.parameters["t_range"],
175
+ n_max=self.parameters["nmax"],
176
+ )
177
+ update = proposal.sample(seed=proposal_seed)
178
+ next_state = _add_events(
179
+ events=step_events,
180
+ m=update["m"],
181
+ t=update["t"],
182
+ x=[self.tx_topology.target],
183
+ x_star=tf.cast(-update["x_star"], step_events.dtype),
184
+ )
185
+ reverse = add_occult_proposal(
186
+ events=next_state,
187
+ topology=self.tx_topology,
188
+ initial_state=self.initial_state,
189
+ n_max=self.parameters["nmax"],
190
+ t_range=self.parameters["t_range"],
191
+ name=f"{self.name}rev",
192
+ )
193
+ q_fwd = tf.reduce_sum(proposal.log_prob(update))
194
+ q_rev = tf.reduce_sum(reverse.log_prob(update))
195
+ log_acceptance_correction = q_rev - q_fwd
196
+
197
+ return (
198
+ update,
199
+ next_state,
200
+ log_acceptance_correction,
201
+ -tf.ones(1, dtype=tf.int32),
202
+ )
203
+
204
+ u = tfd.Uniform().sample(seed=add_del_seed)
205
+ delta, next_state, log_acceptance_correction, direction = tf.cond(
206
+ (u < PROB_DIRECTION)
207
+ & (
208
+ tf.math.count_nonzero(
209
+ step_events[..., self.tx_topology.target]
210
+ )
211
+ > 0
212
+ ),
213
+ del_occult_fn,
214
+ add_occult_fn,
215
+ )
216
+
217
+ next_target_log_prob = self.target_log_prob_fn(next_state)
218
+
219
+ if mcmc_util.is_list_like(current_events):
220
+ next_state = [next_state]
221
+
222
+ return [
223
+ next_state,
224
+ OccultKernelResults(
225
+ log_acceptance_correction=log_acceptance_correction,
226
+ target_log_prob=next_target_log_prob,
227
+ m=delta["m"],
228
+ t=delta["t"],
229
+ delta_t=direction,
230
+ x_star=delta["x_star"],
231
+ seed=add_del_seed,
232
+ ),
233
+ ]
234
+
235
+ def bootstrap_results(self, init_state):
236
+ with tf.name_scope("uncalibrated_event_times_rw/bootstrap_results"):
237
+ if not mcmc_util.is_list_like(init_state):
238
+ init_state = [init_state]
239
+
240
+ init_state = [
241
+ tf.convert_to_tensor(x, dtype=self._dtype) for x in init_state
242
+ ]
243
+ init_target_log_prob = self.target_log_prob_fn(*init_state)
244
+ return OccultKernelResults(
245
+ log_acceptance_correction=tf.constant(
246
+ 0.0, dtype=init_target_log_prob.dtype
247
+ ),
248
+ target_log_prob=init_target_log_prob,
249
+ m=tf.zeros([1], dtype=tf.int32),
250
+ t=tf.zeros([1], dtype=tf.int32),
251
+ delta_t=tf.zeros([1], dtype=tf.int32),
252
+ x_star=tf.zeros([1], dtype=tf.int32),
253
+ seed=samplers.zeros_seed(),
254
+ )
@@ -0,0 +1,28 @@
1
+ """Test event time samplers"""
2
+
3
+ import pytest
4
+ import tensorflow as tf
5
+
6
+
7
+ @pytest.fixture
8
+ def random_events():
9
+ """SEIR model with prescribed starting conditions"""
10
+ events = tf.random.uniform(
11
+ [10, 10, 3], minval=0, maxval=100, dtype=tf.float64, seed=0
12
+ )
13
+ return events
14
+
15
+
16
+ @pytest.fixture
17
+ def initial_state():
18
+ popsize = tf.fill([10], tf.constant(100.0, tf.float64))
19
+ initial_state = tf.stack(
20
+ [
21
+ popsize,
22
+ tf.ones_like(popsize),
23
+ tf.zeros_like(popsize),
24
+ tf.zeros_like(popsize),
25
+ ],
26
+ axis=-1,
27
+ )
28
+ return initial_state
@@ -0,0 +1,150 @@
1
+ import tensorflow as tf
2
+ import tensorflow_probability as tfp
3
+
4
+ from gemlib.distributions import Categorical2, UniformInteger
5
+
6
+ tfd = tfp.distributions
7
+
8
+
9
+ def add_occult_proposal(
10
+ events,
11
+ topology,
12
+ initial_state,
13
+ n_max,
14
+ t_range=None,
15
+ dtype=tf.int32,
16
+ name=None,
17
+ ):
18
+ if t_range is None:
19
+ t_range = [0, events.shape[-2]]
20
+
21
+ def m():
22
+ """Select a metapopulation"""
23
+ with tf.name_scope("m"):
24
+ return UniformInteger(
25
+ low=[0],
26
+ high=[events.shape[0]],
27
+ dtype=dtype,
28
+ float_dtype=events.dtype,
29
+ )
30
+
31
+ def t():
32
+ """Select a timepoint"""
33
+ with tf.name_scope("t"):
34
+ return UniformInteger(
35
+ low=[t_range[0]],
36
+ high=[t_range[1]],
37
+ dtype=dtype,
38
+ float_dtype=events.dtype,
39
+ )
40
+
41
+ def x_star(m, t):
42
+ """Draw num to add bounded by counting process contraint"""
43
+ if topology.prev is not None:
44
+ mask = ( # Mask out times prior to t
45
+ tf.cast(tf.range(events.shape[-2]) < t[0], events.dtype)
46
+ * events.dtype.max
47
+ )
48
+ m_events = tf.gather(events, m, axis=-3)
49
+ m_inits = tf.gather(initial_state, m, axis=-2)
50
+ diff = m_events[..., topology.prev] - m_events[..., topology.target]
51
+ diff = tf.gather(m_inits, topology.target, axis=-1) + tf.cumsum(
52
+ diff, axis=-1
53
+ )
54
+ diff = diff + mask
55
+ bound = tf.cast(tf.reduce_min(diff, axis=-1), dtype=tf.int32)
56
+ # bound = tf.maximum(0, bound)
57
+ bound = tf.minimum(n_max, bound)
58
+ else:
59
+ bound = tf.broadcast_to(n_max, m.shape)
60
+
61
+ return UniformInteger(
62
+ low=1, high=bound + 1, dtype=dtype, float_dtype=events.dtype
63
+ )
64
+
65
+ return tfd.JointDistributionNamed(
66
+ {"m": m, "t": t, "x_star": x_star}, name=name
67
+ )
68
+
69
+
70
+ def del_occult_proposal(
71
+ events,
72
+ topology,
73
+ initial_state,
74
+ n_max,
75
+ t_range=None,
76
+ dtype=tf.int32,
77
+ name=None,
78
+ ):
79
+ if t_range is None:
80
+ t_range = [0, events.shape[-2]]
81
+
82
+ def m():
83
+ """Select a metapopulation"""
84
+ with tf.name_scope("m"):
85
+ hot_meta = (
86
+ tf.math.count_nonzero(
87
+ events[..., slice(*t_range), topology.target],
88
+ axis=1,
89
+ keepdims=True,
90
+ )
91
+ > 0
92
+ )
93
+ hot_meta = tf.cast(tf.transpose(hot_meta), dtype=events.dtype)
94
+ logits = tf.math.log(hot_meta)
95
+ X = Categorical2(
96
+ logits=tf.cast(logits, tf.float32), dtype=dtype, name="m"
97
+ )
98
+ return X
99
+
100
+ def t(m):
101
+ """Draw timepoint"""
102
+ with tf.name_scope("t"):
103
+ metapops = tf.gather(events, m)
104
+ hot_times = (
105
+ (metapops[..., topology.target] > 0)
106
+ & (t_range[0] <= tf.range(events.shape[-2]))
107
+ & (tf.range(events.shape[-2]) < t_range[1])
108
+ )
109
+ hot_times = tf.cast(hot_times, dtype=events.dtype)
110
+ logits = tf.math.log(hot_times)
111
+ return Categorical2(
112
+ logits=tf.cast(logits, tf.float32), dtype=dtype, name="t"
113
+ )
114
+
115
+ def x_star(m, t):
116
+ """Draw num to delete"""
117
+ with tf.name_scope("x_star"):
118
+ if topology.next is not None:
119
+ mask = ( # Mask out times prior to t
120
+ tf.cast(tf.range(events.shape[-2]) < t[0], events.dtype)
121
+ * events.dtype.max
122
+ )
123
+ m_events = tf.gather(events, m, axis=-3)
124
+ m_inits = tf.gather(initial_state, m, axis=-2)
125
+ # calc offset[target] + N_{target}(t) - N_{next} bound
126
+ diff = (
127
+ m_events[..., topology.target]
128
+ - m_events[..., topology.next]
129
+ )
130
+ diff = tf.gather(m_inits, topology.next, axis=-1) + tf.cumsum(
131
+ diff, axis=-1
132
+ )
133
+ diff = diff + mask
134
+ bound = tf.cast(tf.reduce_min(diff, axis=-1), dtype=tf.int32)
135
+ # bound = tf.maximum(0, bound)
136
+ bound = tf.minimum(n_max, bound)
137
+ else:
138
+ bound = tf.broadcast_to(n_max, m.shape)
139
+
140
+ return UniformInteger(
141
+ low=1,
142
+ high=bound + 1,
143
+ dtype=dtype,
144
+ float_dtype=events.dtype,
145
+ name="x_star",
146
+ )
147
+
148
+ return tfd.JointDistributionNamed(
149
+ {"m": m, "t": t, "x_star": x_star}, name=name
150
+ )
@@ -0,0 +1,9 @@
1
+ """Utilities for DiscreteTimeStateTransitionModel MCMC kernels"""
2
+
3
+ from typing import NamedTuple
4
+
5
+
6
+ class TransitionTopology(NamedTuple):
7
+ prev: int
8
+ target: int
9
+ next: int
File without changes