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,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
|
+
)
|
|
File without changes
|