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,21 @@
|
|
|
1
|
+
"""DiscreteTimeStateTransitionModel-related MCMC samplers"""
|
|
2
|
+
|
|
3
|
+
from gemlib.mcmc.discrete_time_state_transition_model.left_censored_events_mh import ( # noqa: E501
|
|
4
|
+
UncalibratedLeftCensoredEventTimesUpdate,
|
|
5
|
+
)
|
|
6
|
+
from gemlib.mcmc.discrete_time_state_transition_model.move_events import (
|
|
7
|
+
UncalibratedEventTimesUpdate,
|
|
8
|
+
)
|
|
9
|
+
from gemlib.mcmc.discrete_time_state_transition_model.right_censored_events_mh import ( # noqa: E501
|
|
10
|
+
UncalibratedOccultUpdate,
|
|
11
|
+
)
|
|
12
|
+
from gemlib.mcmc.discrete_time_state_transition_model.util import (
|
|
13
|
+
TransitionTopology,
|
|
14
|
+
)
|
|
15
|
+
|
|
16
|
+
__all__ = [
|
|
17
|
+
"TransitionTopology",
|
|
18
|
+
"UncalibratedEventTimesUpdate",
|
|
19
|
+
"UncalibratedLeftCensoredEventTimesUpdate",
|
|
20
|
+
"UncalibratedOccultUpdate",
|
|
21
|
+
]
|
|
@@ -0,0 +1,275 @@
|
|
|
1
|
+
"""Mechanism for proposing event times to move"""
|
|
2
|
+
|
|
3
|
+
from warnings import warn
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import tensorflow as tf
|
|
7
|
+
import tensorflow_probability as tfp
|
|
8
|
+
from tensorflow_probability.python.mcmc.internal import util as mcmc_util
|
|
9
|
+
|
|
10
|
+
from gemlib.distributions import (
|
|
11
|
+
Categorical2,
|
|
12
|
+
UniformInteger,
|
|
13
|
+
UniformKCategorical,
|
|
14
|
+
)
|
|
15
|
+
|
|
16
|
+
tfd = tfp.distributions
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def _events_or_inf(events, transition_id):
|
|
20
|
+
if transition_id is None:
|
|
21
|
+
return tf.fill(
|
|
22
|
+
events.shape[:-1], tf.constant(np.inf, dtype=events.dtype)
|
|
23
|
+
)
|
|
24
|
+
return tf.gather(events, transition_id, axis=-1)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def _abscumdiff(
|
|
28
|
+
events, initial_state, topology, t, delta_t, bound_times, int_dtype=tf.int32
|
|
29
|
+
):
|
|
30
|
+
"""Returns the number of free events to move in target_events
|
|
31
|
+
bounded by max([N_{target_id}(t)-N_{bound_id}(t)]_{bound_t}).
|
|
32
|
+
|
|
33
|
+
:param events: a [(M), T, X] tensor of transition events
|
|
34
|
+
:param initial_state: a [M, X] tensor of the constraining initial state
|
|
35
|
+
:param target_id: the Xth index of the target event
|
|
36
|
+
:param bound_t: the times to compute the constraints
|
|
37
|
+
:param bound_id: the Xth index of the bounding event, -1 implies no bound
|
|
38
|
+
|
|
39
|
+
:returns: a tensor of shape [M] + bound_t.shape[0] + of max free events,
|
|
40
|
+
dtype=target_events.dtype
|
|
41
|
+
"""
|
|
42
|
+
with tf.name_scope("_abscumdiff"):
|
|
43
|
+
# This line prevents negative indices. However, we must have
|
|
44
|
+
# a contract that the output of the algorithm is invalid!
|
|
45
|
+
bound_times = tf.clip_by_value(
|
|
46
|
+
bound_times, clip_value_min=0, clip_value_max=events.shape[-2] - 1
|
|
47
|
+
)
|
|
48
|
+
|
|
49
|
+
# Maybe replace with pad to avoid unstack/stack
|
|
50
|
+
prev_events = _events_or_inf(events, topology.prev)
|
|
51
|
+
target_events = tf.gather(events, topology.target, axis=-1)
|
|
52
|
+
next_events = _events_or_inf(events, topology.next)
|
|
53
|
+
event_tensor = tf.stack(
|
|
54
|
+
[prev_events, target_events, next_events], axis=-1
|
|
55
|
+
)
|
|
56
|
+
|
|
57
|
+
# Compute the absolute cumulative difference between event times
|
|
58
|
+
diff = event_tensor[..., 1:] - event_tensor[..., :-1] # [m, T, 2]
|
|
59
|
+
cumdiff = tf.abs(tf.cumsum(diff, axis=-2)) # cumsum along time axis
|
|
60
|
+
|
|
61
|
+
# Create indices into cumdiff [m, d_max, 2]. Last dimension selects
|
|
62
|
+
# the bound for either the previous or next event.
|
|
63
|
+
indices = tf.stack(
|
|
64
|
+
[
|
|
65
|
+
tf.repeat(
|
|
66
|
+
tf.range(events.shape[0], dtype=int_dtype),
|
|
67
|
+
[bound_times.shape[1]],
|
|
68
|
+
),
|
|
69
|
+
tf.reshape(bound_times, [-1]),
|
|
70
|
+
tf.repeat(tf.where(delta_t < 0, 0, 1), [bound_times.shape[1]]),
|
|
71
|
+
],
|
|
72
|
+
axis=-1,
|
|
73
|
+
)
|
|
74
|
+
indices = tf.reshape(
|
|
75
|
+
indices, [events.shape[-3], bound_times.shape[1], 3]
|
|
76
|
+
)
|
|
77
|
+
free_events = tf.gather_nd(cumdiff, indices)
|
|
78
|
+
|
|
79
|
+
# Add on initial state
|
|
80
|
+
indices = tf.stack(
|
|
81
|
+
[
|
|
82
|
+
tf.range(events.shape[0]),
|
|
83
|
+
tf.where(
|
|
84
|
+
delta_t[:, 0] < 0, topology.target, topology.target + 1
|
|
85
|
+
),
|
|
86
|
+
],
|
|
87
|
+
axis=-1,
|
|
88
|
+
)
|
|
89
|
+
bound_init_state = tf.gather_nd(initial_state, indices)
|
|
90
|
+
free_events += bound_init_state[..., tf.newaxis]
|
|
91
|
+
|
|
92
|
+
return free_events
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
class Deterministic2(tfd.Deterministic):
|
|
96
|
+
def __init__(
|
|
97
|
+
self,
|
|
98
|
+
loc,
|
|
99
|
+
atol=None,
|
|
100
|
+
rtol=None,
|
|
101
|
+
validate_args=False,
|
|
102
|
+
allow_nan_stats=True,
|
|
103
|
+
log_prob_dtype=tf.float32,
|
|
104
|
+
name="Deterministic",
|
|
105
|
+
):
|
|
106
|
+
parameters = dict(locals())
|
|
107
|
+
super().__init__(
|
|
108
|
+
loc,
|
|
109
|
+
atol=atol,
|
|
110
|
+
rtol=rtol,
|
|
111
|
+
validate_args=validate_args,
|
|
112
|
+
allow_nan_stats=allow_nan_stats,
|
|
113
|
+
name=name,
|
|
114
|
+
)
|
|
115
|
+
self.log_prob_dtype = log_prob_dtype
|
|
116
|
+
|
|
117
|
+
def _prob(self, x):
|
|
118
|
+
return tf.constant(1, dtype=self.log_prob_dtype)
|
|
119
|
+
|
|
120
|
+
|
|
121
|
+
def event_time_proposal(
|
|
122
|
+
events, initial_state, topology, d_max, n_max, dtype=tf.int32, name=None
|
|
123
|
+
):
|
|
124
|
+
"""Draws an event time move proposal.
|
|
125
|
+
:param events: a [M, T, K] tensor of event times (M number of
|
|
126
|
+
metapopulations, T number of full_timepoints, K number of
|
|
127
|
+
transitions)
|
|
128
|
+
:param initial_state: a [M, S] tensor of initial metapopulation x state
|
|
129
|
+
counts
|
|
130
|
+
:param topology: a 3-element tuple of (previous_transition,
|
|
131
|
+
target_transition, next_transition), eg "(s->e, e->i,
|
|
132
|
+
i->r)" (assuming we are interested presently in e->i,
|
|
133
|
+
`None` for boundaries)
|
|
134
|
+
:param d_max: the maximum distance over which to move (in time)
|
|
135
|
+
:param n_max: the maximum number of events to move
|
|
136
|
+
"""
|
|
137
|
+
target_events = tf.gather(events, topology.target, axis=-1)
|
|
138
|
+
time_interval = tf.range(d_max, dtype=dtype)
|
|
139
|
+
|
|
140
|
+
def t():
|
|
141
|
+
with tf.name_scope("t"):
|
|
142
|
+
# Waiting for fixed tf.nn.sparse_softmax_cross_entropy_with_logits
|
|
143
|
+
x = tf.cast(target_events > 0, dtype=events.dtype) # [M, T]
|
|
144
|
+
return Categorical2(logits=tf.math.log(x), name="event_coords")
|
|
145
|
+
|
|
146
|
+
def delta_t(t):
|
|
147
|
+
with tf.name_scope("delta_t"):
|
|
148
|
+
d_max_bcast = tf.broadcast_to(d_max, [events.shape[-3]])
|
|
149
|
+
low = -tf.clip_by_value(
|
|
150
|
+
d_max_bcast, clip_value_min=0, clip_value_max=t
|
|
151
|
+
)
|
|
152
|
+
high = tf.clip_by_value(
|
|
153
|
+
d_max_bcast,
|
|
154
|
+
clip_value_min=0,
|
|
155
|
+
clip_value_max=events.shape[-2] - t - 1,
|
|
156
|
+
)
|
|
157
|
+
return UniformInteger(
|
|
158
|
+
low=low, high=high + 1, float_dtype=events.dtype
|
|
159
|
+
)
|
|
160
|
+
|
|
161
|
+
def x_star(t, delta_t):
|
|
162
|
+
with tf.name_scope("x_star"):
|
|
163
|
+
# Compute bounds
|
|
164
|
+
# The limitations of XLA mean that we must calculate bounds for
|
|
165
|
+
# intervals [t, t+delta_t) if delta_t > 0, and [t+delta_t, t) if
|
|
166
|
+
# delta_t is < 0.
|
|
167
|
+
t = t[..., tf.newaxis]
|
|
168
|
+
delta_t = delta_t[..., tf.newaxis]
|
|
169
|
+
bound_times = tf.where(
|
|
170
|
+
delta_t < 0,
|
|
171
|
+
t - time_interval - 1,
|
|
172
|
+
t + time_interval, # [t+delta_t, t)
|
|
173
|
+
) # [t, t+delta_t)
|
|
174
|
+
free_events = _abscumdiff(
|
|
175
|
+
events=events,
|
|
176
|
+
initial_state=initial_state,
|
|
177
|
+
topology=topology,
|
|
178
|
+
t=t,
|
|
179
|
+
delta_t=delta_t,
|
|
180
|
+
bound_times=bound_times,
|
|
181
|
+
int_dtype=dtype,
|
|
182
|
+
)
|
|
183
|
+
|
|
184
|
+
# Mask out bits of the interval we don't need for our delta_t
|
|
185
|
+
inf_mask = tf.cumsum(
|
|
186
|
+
tf.one_hot(
|
|
187
|
+
tf.math.abs(delta_t[:, 0]),
|
|
188
|
+
d_max,
|
|
189
|
+
on_value=tf.constant(np.inf, events.dtype),
|
|
190
|
+
dtype=events.dtype,
|
|
191
|
+
)
|
|
192
|
+
)
|
|
193
|
+
free_events = tf.maximum(inf_mask, free_events)
|
|
194
|
+
free_events = tf.reduce_min(free_events, axis=-1)
|
|
195
|
+
|
|
196
|
+
indices = tf.stack(
|
|
197
|
+
[tf.range(events.shape[0], dtype=dtype), t[:, 0]], axis=-1
|
|
198
|
+
)
|
|
199
|
+
available_events = tf.gather_nd(target_events, indices)
|
|
200
|
+
max_events = tf.minimum(free_events, available_events)
|
|
201
|
+
max_events = tf.clip_by_value(
|
|
202
|
+
max_events, clip_value_min=0, clip_value_max=n_max
|
|
203
|
+
)
|
|
204
|
+
# Draw x_star
|
|
205
|
+
return UniformInteger(
|
|
206
|
+
low=1, high=max_events + 1, float_dtype=events.dtype
|
|
207
|
+
)
|
|
208
|
+
|
|
209
|
+
return tfd.JointDistributionNamed(
|
|
210
|
+
{"t": t, "delta_t": delta_t, "x_star": x_star}, name=name
|
|
211
|
+
)
|
|
212
|
+
|
|
213
|
+
|
|
214
|
+
def filtered_event_time_proposal( # pylint: disable-invalid-name
|
|
215
|
+
events,
|
|
216
|
+
initial_state,
|
|
217
|
+
topology,
|
|
218
|
+
m_max,
|
|
219
|
+
d_max,
|
|
220
|
+
n_max,
|
|
221
|
+
dtype=tf.int32,
|
|
222
|
+
name=None,
|
|
223
|
+
):
|
|
224
|
+
"""FilteredEventTimeProposal allows us to choose a subset of indices
|
|
225
|
+
in `range(events.shape[0])` for which to propose an update. The
|
|
226
|
+
results are then broadcast back to `events.shape[0]`.
|
|
227
|
+
|
|
228
|
+
:param events: a [M, T, X] event tensor
|
|
229
|
+
:param initial_state: a [M, S] initial state tensor
|
|
230
|
+
:param topology: a TransitionTopology named tuple describing the ordering
|
|
231
|
+
of events
|
|
232
|
+
:param m: the number of metapopulations to move
|
|
233
|
+
:param d_max: maximum distance in time to move
|
|
234
|
+
:param n_max: maximum number of events to move (user defined)
|
|
235
|
+
:return: an instance of a JointDistributionNamed
|
|
236
|
+
"""
|
|
237
|
+
if mcmc_util.is_list_like(events):
|
|
238
|
+
warn(
|
|
239
|
+
"Batched FilteredEventTimeProposals are not yet supported",
|
|
240
|
+
stacklevel=1,
|
|
241
|
+
)
|
|
242
|
+
events = events[0]
|
|
243
|
+
|
|
244
|
+
target_events = tf.gather(events, topology.target, axis=-1)
|
|
245
|
+
|
|
246
|
+
def m():
|
|
247
|
+
with tf.name_scope("m"):
|
|
248
|
+
hot_meta = tf.math.count_nonzero(target_events, axis=1) > 0
|
|
249
|
+
X = UniformKCategorical(
|
|
250
|
+
m_max, hot_meta, float_dtype=events.dtype, name="m"
|
|
251
|
+
)
|
|
252
|
+
return X
|
|
253
|
+
|
|
254
|
+
def move(m):
|
|
255
|
+
"""We select out meta-population `m` from the first
|
|
256
|
+
dimension of `events`.
|
|
257
|
+
:param m: a 1-D tensor of indices of meta-populations
|
|
258
|
+
:return: a random variable of type `EventTimeProposal`
|
|
259
|
+
"""
|
|
260
|
+
with tf.name_scope("move"):
|
|
261
|
+
select_meta = tf.gather(events, m, axis=0)
|
|
262
|
+
select_init = tf.gather(initial_state, m, axis=0)
|
|
263
|
+
return event_time_proposal(
|
|
264
|
+
select_meta,
|
|
265
|
+
select_init,
|
|
266
|
+
topology,
|
|
267
|
+
d_max,
|
|
268
|
+
n_max,
|
|
269
|
+
dtype=dtype,
|
|
270
|
+
name=name,
|
|
271
|
+
)
|
|
272
|
+
|
|
273
|
+
return tfd.JointDistributionNamed(
|
|
274
|
+
{"m": m, "move": move}, name="FilteredEventTimeProposal"
|
|
275
|
+
)
|
|
@@ -0,0 +1,56 @@
|
|
|
1
|
+
"""Test fixtures"""
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import pytest
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
@pytest.fixture(scope="module")
|
|
8
|
+
def sir_metapop_example():
|
|
9
|
+
"""Outcome of a simulation from a 3-metapopulation model
|
|
10
|
+
with mixing, implemented in https://colab.research.google.com/drive/1Q1PUcOnYlvCGHzRUBUAp4CxYhZ8RJzg8?usp=sharing
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
incidence_matrix = np.array([[-1, 0], [1, -1], [0, 1]], dtype=np.float32)
|
|
14
|
+
|
|
15
|
+
initial_conditions = np.array(
|
|
16
|
+
[[999, 50, 0], [500, 20, 0], [250, 10, 0]], dtype=np.float32
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
events = np.array(
|
|
20
|
+
[
|
|
21
|
+
[
|
|
22
|
+
[11.0, 11.0],
|
|
23
|
+
[11.0, 6.0],
|
|
24
|
+
[2.0, 6.0],
|
|
25
|
+
[10.0, 6.0],
|
|
26
|
+
[12.0, 7.0],
|
|
27
|
+
[11.0, 5.0],
|
|
28
|
+
[13.0, 10.0],
|
|
29
|
+
],
|
|
30
|
+
[
|
|
31
|
+
[5.0, 4.0],
|
|
32
|
+
[5.0, 2.0],
|
|
33
|
+
[11.0, 1.0],
|
|
34
|
+
[8.0, 4.0],
|
|
35
|
+
[7.0, 4.0],
|
|
36
|
+
[5.0, 5.0],
|
|
37
|
+
[12.0, 6.0],
|
|
38
|
+
],
|
|
39
|
+
[
|
|
40
|
+
[2.0, 2.0],
|
|
41
|
+
[4.0, 2.0],
|
|
42
|
+
[1.0, 0.0],
|
|
43
|
+
[2.0, 2.0],
|
|
44
|
+
[6.0, 1.0],
|
|
45
|
+
[4.0, 1.0],
|
|
46
|
+
[6.0, 1.0],
|
|
47
|
+
],
|
|
48
|
+
],
|
|
49
|
+
dtype=np.float32,
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
return {
|
|
53
|
+
"initial_conditions": initial_conditions,
|
|
54
|
+
"events": events,
|
|
55
|
+
"incidence_matrix": incidence_matrix,
|
|
56
|
+
}
|
|
@@ -0,0 +1,258 @@
|
|
|
1
|
+
"""Metropolis Hastings implementation for left-censored events
|
|
2
|
+
in a discrete-time metapopulation epidemic model
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from typing import NamedTuple, Tuple
|
|
6
|
+
|
|
7
|
+
import tensorflow as tf
|
|
8
|
+
import tensorflow_probability as tfp
|
|
9
|
+
from tensorflow_probability.python.internal import prefer_static as ps
|
|
10
|
+
from tensorflow_probability.python.internal import samplers
|
|
11
|
+
from tensorflow_probability.python.mcmc.internal import util as mcmc_util
|
|
12
|
+
|
|
13
|
+
from .left_censored_events_proposal import (
|
|
14
|
+
left_censored_event_time_proposal,
|
|
15
|
+
)
|
|
16
|
+
|
|
17
|
+
tfd = tfp.distributions
|
|
18
|
+
|
|
19
|
+
__all__ = ["UncalibratedLeftCensoredEventTimesUpdate"]
|
|
20
|
+
|
|
21
|
+
LEFT_CENSORED_EVENT_TIMES_UPDATE_TUPLE_LEN = 2
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def _update_state(update, current_state, transition_idx, incidence_matrix):
|
|
25
|
+
current_initial_conditions, current_events = current_state
|
|
26
|
+
update = {k: tf.convert_to_tensor(v) for k, v in update.items()}
|
|
27
|
+
|
|
28
|
+
# -1 is moving events forward in time
|
|
29
|
+
sign = tf.gather(
|
|
30
|
+
[-1, 1],
|
|
31
|
+
update["direction"],
|
|
32
|
+
)
|
|
33
|
+
events_delta = update["num_events"] * sign
|
|
34
|
+
|
|
35
|
+
# Update initial conditions
|
|
36
|
+
indices = tf.stack(
|
|
37
|
+
[
|
|
38
|
+
tf.broadcast_to(
|
|
39
|
+
update["unit"], [current_initial_conditions.shape[-1]]
|
|
40
|
+
),
|
|
41
|
+
ps.range(current_initial_conditions.shape[-1]),
|
|
42
|
+
],
|
|
43
|
+
axis=-1,
|
|
44
|
+
)
|
|
45
|
+
new_initial_conditions = tf.tensor_scatter_nd_add(
|
|
46
|
+
current_initial_conditions,
|
|
47
|
+
indices=indices,
|
|
48
|
+
updates=tf.cast(
|
|
49
|
+
tf.cast(
|
|
50
|
+
events_delta, incidence_matrix.dtype
|
|
51
|
+
) # TODO sort this out: casts are usually a code smell!
|
|
52
|
+
* ps.gather(incidence_matrix, transition_idx, axis=-1),
|
|
53
|
+
current_initial_conditions.dtype,
|
|
54
|
+
),
|
|
55
|
+
)
|
|
56
|
+
|
|
57
|
+
# Update events
|
|
58
|
+
indices = tf.stack(
|
|
59
|
+
[
|
|
60
|
+
update["unit"],
|
|
61
|
+
update["timepoint"],
|
|
62
|
+
ps.broadcast_to(transition_idx, update["unit"].shape),
|
|
63
|
+
],
|
|
64
|
+
axis=-1,
|
|
65
|
+
)
|
|
66
|
+
new_events = tf.tensor_scatter_nd_sub(
|
|
67
|
+
current_events,
|
|
68
|
+
indices=indices,
|
|
69
|
+
updates=tf.cast(events_delta, dtype=current_events.dtype),
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
return new_initial_conditions, new_events
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def _reverse_update(update):
|
|
76
|
+
direction = (update["direction"] + 1) % 2
|
|
77
|
+
return {
|
|
78
|
+
"unit": update["unit"],
|
|
79
|
+
"timepoint": update["timepoint"],
|
|
80
|
+
"direction": direction,
|
|
81
|
+
"num_events": update["num_events"],
|
|
82
|
+
}
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
class LeftCensoredEventTimeResults(NamedTuple):
|
|
86
|
+
log_acceptance_correction: float
|
|
87
|
+
target_log_prob: float
|
|
88
|
+
unit: int
|
|
89
|
+
timepoint: int
|
|
90
|
+
direction: int
|
|
91
|
+
num_events: int
|
|
92
|
+
seed: Tuple[int, int]
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
class UncalibratedLeftCensoredEventTimesUpdate(tfp.mcmc.TransitionKernel):
|
|
96
|
+
"""UncalibratedLeftCensoredEventTimesUpdate"""
|
|
97
|
+
|
|
98
|
+
def __init__(
|
|
99
|
+
self,
|
|
100
|
+
target_log_prob_fn,
|
|
101
|
+
transition_index,
|
|
102
|
+
incidence_matrix,
|
|
103
|
+
max_timepoint,
|
|
104
|
+
max_events,
|
|
105
|
+
name=None,
|
|
106
|
+
):
|
|
107
|
+
"""An uncalibrated random walk for initial conditions.
|
|
108
|
+
:param target_log_prob_fn: the log density of the target distribution
|
|
109
|
+
:param transition_index: the index of the transition to adjust
|
|
110
|
+
:param incidence_matrix: the `[S,R]` incidence matrix
|
|
111
|
+
:param max_timepoint: max timepoint up to which to move events
|
|
112
|
+
:param max_events: max number of events per unit/timepoint to move
|
|
113
|
+
"""
|
|
114
|
+
self._name = name
|
|
115
|
+
self._parameters = {
|
|
116
|
+
"target_log_prob_fn": target_log_prob_fn,
|
|
117
|
+
"transition_index": transition_index,
|
|
118
|
+
"incidence_matrix": incidence_matrix,
|
|
119
|
+
"max_timepoint": max_timepoint,
|
|
120
|
+
"max_events": max_events,
|
|
121
|
+
"name": name,
|
|
122
|
+
}
|
|
123
|
+
|
|
124
|
+
@property
|
|
125
|
+
def target_log_prob_fn(self):
|
|
126
|
+
return self._parameters["target_log_prob_fn"]
|
|
127
|
+
|
|
128
|
+
@property
|
|
129
|
+
def transition_index(self):
|
|
130
|
+
return self._parameters["transition_index"]
|
|
131
|
+
|
|
132
|
+
@property
|
|
133
|
+
def incidence_matrix(self):
|
|
134
|
+
return self._parameters["incidence_matrix"]
|
|
135
|
+
|
|
136
|
+
@property
|
|
137
|
+
def max_timepoint(self):
|
|
138
|
+
return self._parameters["max_timepoint"]
|
|
139
|
+
|
|
140
|
+
@property
|
|
141
|
+
def max_events(self):
|
|
142
|
+
return self._parameters["max_events"]
|
|
143
|
+
|
|
144
|
+
@property
|
|
145
|
+
def name(self):
|
|
146
|
+
return self._parameters["name"]
|
|
147
|
+
|
|
148
|
+
@property
|
|
149
|
+
def parameters(self):
|
|
150
|
+
return self._parameters
|
|
151
|
+
|
|
152
|
+
@property
|
|
153
|
+
def is_calibrated(self):
|
|
154
|
+
return False
|
|
155
|
+
|
|
156
|
+
def one_step(self, current_state, previous_kernel_results, seed=None):
|
|
157
|
+
"""Update the initial conditions
|
|
158
|
+
|
|
159
|
+
:param current_state: a tuple of `(current_initial_conditions,
|
|
160
|
+
current_events)`
|
|
161
|
+
:param previous_kernel_results: previous kernel results tuple
|
|
162
|
+
:param seed: optional seed tuple `(int32, int32)`
|
|
163
|
+
:returns: new state tuple `(next_initial_conditions, next_events)`
|
|
164
|
+
"""
|
|
165
|
+
with tf.name_scope("uncalibrated_left_censored_events_mh/one_step"):
|
|
166
|
+
seed = samplers.sanitize_seed(
|
|
167
|
+
seed, salt="uncalibrated_left_censored_events_mh"
|
|
168
|
+
)
|
|
169
|
+
|
|
170
|
+
if (not mcmc_util.is_list_like(current_state)) and (
|
|
171
|
+
len(current_state) == LEFT_CENSORED_EVENT_TIMES_UPDATE_TUPLE_LEN
|
|
172
|
+
):
|
|
173
|
+
raise ValueError(
|
|
174
|
+
f"State for LeftCensoredEventTimesUpdate must be a\
|
|
175
|
+
list/tuple of length {LEFT_CENSORED_EVENT_TIMES_UPDATE_TUPLE_LEN}"
|
|
176
|
+
)
|
|
177
|
+
|
|
178
|
+
proposal = left_censored_event_time_proposal(
|
|
179
|
+
events=current_state[1],
|
|
180
|
+
initial_state=current_state[0],
|
|
181
|
+
transition=self.transition_index,
|
|
182
|
+
incidence_matrix=self.incidence_matrix,
|
|
183
|
+
num_units=1,
|
|
184
|
+
max_timepoint=self.max_timepoint,
|
|
185
|
+
max_events=self.max_events,
|
|
186
|
+
name=f"{self.name}/fwd_proposal",
|
|
187
|
+
)
|
|
188
|
+
fwd_update = proposal.sample(seed=seed)
|
|
189
|
+
fwd_proposal_log_prob = proposal.log_prob(
|
|
190
|
+
fwd_update, name="fwd_proposal_log_prob"
|
|
191
|
+
)
|
|
192
|
+
|
|
193
|
+
next_state = _update_state(
|
|
194
|
+
fwd_update,
|
|
195
|
+
current_state,
|
|
196
|
+
self.transition_index,
|
|
197
|
+
self.incidence_matrix,
|
|
198
|
+
)
|
|
199
|
+
next_target_log_prob = self.target_log_prob_fn(*next_state)
|
|
200
|
+
|
|
201
|
+
rev_update = _reverse_update(fwd_update)
|
|
202
|
+
rev_proposal = left_censored_event_time_proposal(
|
|
203
|
+
events=next_state[1],
|
|
204
|
+
initial_state=next_state[0],
|
|
205
|
+
transition=self.transition_index,
|
|
206
|
+
incidence_matrix=self.incidence_matrix,
|
|
207
|
+
num_units=1,
|
|
208
|
+
max_timepoint=self.max_timepoint,
|
|
209
|
+
max_events=self.max_events,
|
|
210
|
+
name=f"{self.name}/rev_proposal",
|
|
211
|
+
)
|
|
212
|
+
rev_proposal_log_prob = rev_proposal.log_prob(rev_update)
|
|
213
|
+
log_acceptance_correction = tf.reduce_sum(
|
|
214
|
+
rev_proposal_log_prob - fwd_proposal_log_prob
|
|
215
|
+
)
|
|
216
|
+
results = (
|
|
217
|
+
next_state,
|
|
218
|
+
LeftCensoredEventTimeResults(
|
|
219
|
+
log_acceptance_correction=log_acceptance_correction,
|
|
220
|
+
target_log_prob=next_target_log_prob,
|
|
221
|
+
unit=fwd_update["unit"],
|
|
222
|
+
timepoint=fwd_update["timepoint"],
|
|
223
|
+
direction=fwd_update["direction"],
|
|
224
|
+
num_events=fwd_update["num_events"],
|
|
225
|
+
seed=seed,
|
|
226
|
+
),
|
|
227
|
+
)
|
|
228
|
+
|
|
229
|
+
return results
|
|
230
|
+
|
|
231
|
+
def bootstrap_results(self, init_state):
|
|
232
|
+
with tf.name_scope(
|
|
233
|
+
"uncalibrated_left_censored_events_mh/boostrap_results"
|
|
234
|
+
):
|
|
235
|
+
if (not mcmc_util.is_list_like(init_state)) and (
|
|
236
|
+
len(init_state) == 2 # noqa: PLR2004
|
|
237
|
+
):
|
|
238
|
+
raise ValueError(
|
|
239
|
+
f"State for LeftCensoredEventTimesUpdate must be a \
|
|
240
|
+
list/tuple of length {LEFT_CENSORED_EVENT_TIMES_UPDATE_TUPLE_LEN}"
|
|
241
|
+
)
|
|
242
|
+
|
|
243
|
+
initial_conditions = tf.convert_to_tensor(init_state[0])
|
|
244
|
+
events = tf.convert_to_tensor(init_state[1])
|
|
245
|
+
init_target_log_prob = self.target_log_prob_fn(
|
|
246
|
+
initial_conditions, events
|
|
247
|
+
)
|
|
248
|
+
return LeftCensoredEventTimeResults(
|
|
249
|
+
log_acceptance_correction=tf.constant(
|
|
250
|
+
0.0, dtype=init_target_log_prob.dtype
|
|
251
|
+
),
|
|
252
|
+
target_log_prob=init_target_log_prob,
|
|
253
|
+
unit=tf.zeros([1], dtype=tf.int32),
|
|
254
|
+
timepoint=tf.zeros([1], dtype=tf.int32),
|
|
255
|
+
direction=tf.constant(0, dtype=tf.int32),
|
|
256
|
+
num_events=tf.ones([1], dtype=tf.int32),
|
|
257
|
+
seed=samplers.zeros_seed(),
|
|
258
|
+
)
|