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,840 @@
|
|
|
1
|
+
"""Chain binomial process rippler algorithm"""
|
|
2
|
+
|
|
3
|
+
import warnings
|
|
4
|
+
from collections import namedtuple
|
|
5
|
+
|
|
6
|
+
import tensorflow as tf
|
|
7
|
+
import tensorflow_probability as tfp
|
|
8
|
+
from tensorflow_probability.python.internal import prefer_static
|
|
9
|
+
from tensorflow_probability.python.mcmc.internal import util as mcmc_util
|
|
10
|
+
|
|
11
|
+
from gemlib.distributions import Hypergeometric, UniformInteger
|
|
12
|
+
|
|
13
|
+
tfd = tfp.distributions
|
|
14
|
+
samplers = tfp.random
|
|
15
|
+
|
|
16
|
+
__all__ = [
|
|
17
|
+
"default_initial_ripple",
|
|
18
|
+
"damped_initial_ripple_fn",
|
|
19
|
+
"DampedCBRKernel",
|
|
20
|
+
]
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def _compute_state(initial_state, events, stoichiometry, closed=False):
|
|
24
|
+
"""Computes a state tensor from initial state and event tensor
|
|
25
|
+
|
|
26
|
+
:param initial_state: a tensor of shape [S, M]
|
|
27
|
+
:param events: a tensor of shape [T, R, M]
|
|
28
|
+
:param stoichiometry: a stoichiometry matrix of shape [R, S] describing
|
|
29
|
+
how transitions update the state.
|
|
30
|
+
:param closed: if `True`, return state in close interval [0, T], otherwise
|
|
31
|
+
[0, T)
|
|
32
|
+
:return: a tensor of shape [T, S, M] if `closed=False` or [T+1, S, M] if
|
|
33
|
+
`closed=True`
|
|
34
|
+
describing the state of the
|
|
35
|
+
system for each batch M at time T.
|
|
36
|
+
"""
|
|
37
|
+
if isinstance(stoichiometry, tf.Tensor):
|
|
38
|
+
stoichiometry = prefer_static.cast(stoichiometry, dtype=events.dtype)
|
|
39
|
+
else:
|
|
40
|
+
stoichiometry = tf.convert_to_tensor(stoichiometry, dtype=events.dtype)
|
|
41
|
+
|
|
42
|
+
increments = tf.einsum("...trm,rs->...tsm", events, stoichiometry)
|
|
43
|
+
|
|
44
|
+
if closed is False:
|
|
45
|
+
cum_increments = tf.cumsum(increments, axis=-3, exclusive=True)
|
|
46
|
+
else:
|
|
47
|
+
cum_increments = tf.cumsum(increments, axis=-3, exclusive=False)
|
|
48
|
+
cum_increments = tf.concat(
|
|
49
|
+
[tf.zeros_like(cum_increments[..., 0:1, :, :]), cum_increments],
|
|
50
|
+
axis=-2,
|
|
51
|
+
)
|
|
52
|
+
state = cum_increments + tf.expand_dims(initial_state, axis=-3)
|
|
53
|
+
return state
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
class DampingFunction:
|
|
57
|
+
def __init__(self, p, upper_bound, gamma):
|
|
58
|
+
self._parameters = locals()
|
|
59
|
+
self._dtype = self.p.dtype
|
|
60
|
+
|
|
61
|
+
@property
|
|
62
|
+
def p(self):
|
|
63
|
+
return tf.convert_to_tensor(self._parameters["p"])
|
|
64
|
+
|
|
65
|
+
@property
|
|
66
|
+
def upper_bound(self):
|
|
67
|
+
return tf.convert_to_tensor(self._parameters["upper_bound"])
|
|
68
|
+
|
|
69
|
+
@property
|
|
70
|
+
def gamma(self):
|
|
71
|
+
return tf.convert_to_tensor(self._parameters["gamma"])
|
|
72
|
+
|
|
73
|
+
def __call__(self, u):
|
|
74
|
+
u = tf.convert_to_tensor(u, dtype=self._dtype)
|
|
75
|
+
|
|
76
|
+
r_transformed = self.p + tf.math.pow(
|
|
77
|
+
self.upper_bound - self.p, 1 - self.gamma
|
|
78
|
+
) * tf.math.pow(u - self.p, self.gamma)
|
|
79
|
+
return tf.where(
|
|
80
|
+
(self.p < u) & (u <= self.upper_bound), r_transformed, u
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
def forward(self, u):
|
|
84
|
+
return self.__call__(u)
|
|
85
|
+
|
|
86
|
+
def inverse(self, u):
|
|
87
|
+
u = tf.convert_to_tensor(u, dtype=self._dtype)
|
|
88
|
+
|
|
89
|
+
r_transformed = self.p + tf.math.pow(
|
|
90
|
+
u - self.p, 1 / self.gamma
|
|
91
|
+
) * tf.math.pow(self.upper_bound - self.p, 1 - 1 / self.gamma)
|
|
92
|
+
return tf.where(
|
|
93
|
+
(self.p < u) & (u <= self.upper_bound), r_transformed, u
|
|
94
|
+
)
|
|
95
|
+
|
|
96
|
+
def log_inverse_jacobian(self, u):
|
|
97
|
+
"""N.B. only implemented for p < u <= upper_bound"""
|
|
98
|
+
u = tf.convert_to_tensor(u, dtype=self._dtype)
|
|
99
|
+
|
|
100
|
+
deriv = (
|
|
101
|
+
(1.0 - 1.0 / self.gamma) * tf.math.log(self.upper_bound - self.p)
|
|
102
|
+
+ (1.0 / self.gamma - 1.0) * tf.math.log(u - self.p)
|
|
103
|
+
- tf.math.log(self.gamma)
|
|
104
|
+
)
|
|
105
|
+
|
|
106
|
+
return tf.where((self.p < u) & (u <= self.upper_bound), deriv, 0.0)
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
def _reduce_first_n(values, n):
|
|
110
|
+
"""Reduces the first `n` elements of `values` over the first dimension
|
|
111
|
+
of `values` in a vectorized way.
|
|
112
|
+
"""
|
|
113
|
+
values_shape = tf.shape(values)
|
|
114
|
+
seq = tf.range(values_shape[0], dtype=n.dtype)
|
|
115
|
+
seq = tf.reshape(seq, shape=[values_shape[0]] + [1] * len(values.shape[1:]))
|
|
116
|
+
mask = tf.cast(seq < n, values.dtype)
|
|
117
|
+
return tf.reduce_sum(values * mask, axis=0)
|
|
118
|
+
|
|
119
|
+
|
|
120
|
+
def _reduce_uniforms(
|
|
121
|
+
p_low,
|
|
122
|
+
p_high,
|
|
123
|
+
size,
|
|
124
|
+
u_transform_fn,
|
|
125
|
+
seed,
|
|
126
|
+
chunksize=10,
|
|
127
|
+
):
|
|
128
|
+
"""Generates `size` U(`p_low`, `p_high`) variates, and reduces
|
|
129
|
+
through the transform `u_transform_fn`.
|
|
130
|
+
"""
|
|
131
|
+
seed = samplers.sanitize_seed(seed, salt="binomial_log_jacobian")
|
|
132
|
+
size = tf.cast(size, tf.int32)
|
|
133
|
+
|
|
134
|
+
def cond(i, _):
|
|
135
|
+
return tf.reduce_sum(size - i * chunksize) > 0
|
|
136
|
+
|
|
137
|
+
def body(i, accum):
|
|
138
|
+
u = tfd.Uniform(low=p_low, high=p_high).sample(
|
|
139
|
+
sample_shape=chunksize, seed=seed
|
|
140
|
+
)
|
|
141
|
+
size_local = tf.clip_by_value(
|
|
142
|
+
size - i * chunksize, clip_value_min=0, clip_value_max=chunksize
|
|
143
|
+
)
|
|
144
|
+
log_jacobian = _reduce_first_n(u_transform_fn(u), size_local)
|
|
145
|
+
return i + 1, accum + log_jacobian
|
|
146
|
+
|
|
147
|
+
_, log_jacobian = tf.while_loop(
|
|
148
|
+
cond, body, loop_vars=(0, tf.zeros_like(p_low))
|
|
149
|
+
)
|
|
150
|
+
|
|
151
|
+
return log_jacobian
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def _p_step_ps_gt_p(z, x, p, ps, gamma, seed):
|
|
155
|
+
r"""Run the damped p-step where ps > p
|
|
156
|
+
|
|
157
|
+
:param z: a [R, M] tensor of events for R transitions and M units
|
|
158
|
+
:param x: a [S, M] tensor of state values for S states and M units
|
|
159
|
+
:param p: a [R, M] tensor of transition rates for the current time series
|
|
160
|
+
:param ps: a [R, M] tensor of transition rates for the rippled time series
|
|
161
|
+
:param gamma: a scalar ($\gamma \geq 1$) giving the damping.
|
|
162
|
+
:param seed: a random seed.
|
|
163
|
+
:returns: a tuple of ([R, M] tensor of new event numbers,
|
|
164
|
+
log_acceptance_correction)
|
|
165
|
+
"""
|
|
166
|
+
seeds = samplers.split_seed(seed, n=4, salt="_p_step_ps_gt_p")
|
|
167
|
+
|
|
168
|
+
upper_bound = 2 * ps - p
|
|
169
|
+
|
|
170
|
+
damp = DampingFunction(p, upper_bound, gamma)
|
|
171
|
+
|
|
172
|
+
p_v = (upper_bound - p) / (1.0 - p)
|
|
173
|
+
p_w = (damp(ps) - p) / (upper_bound - p)
|
|
174
|
+
|
|
175
|
+
v = tfd.Binomial(total_count=x - z, probs=p_v).sample(seed=seeds[0])
|
|
176
|
+
w = tfd.Binomial(total_count=v, probs=p_w).sample(seed=seeds[1])
|
|
177
|
+
|
|
178
|
+
# Chunk up the sampling domain and use masking here
|
|
179
|
+
def jacobian(x):
|
|
180
|
+
return damp.log_inverse_jacobian(x)
|
|
181
|
+
|
|
182
|
+
neg_log_inverse_jacobian = _reduce_uniforms(
|
|
183
|
+
p, ps, w, jacobian, seeds[2]
|
|
184
|
+
) + _reduce_uniforms(ps, upper_bound, v - w, jacobian, seeds[3])
|
|
185
|
+
|
|
186
|
+
return (
|
|
187
|
+
z + w,
|
|
188
|
+
neg_log_inverse_jacobian,
|
|
189
|
+
) # log_acceptance_correction
|
|
190
|
+
|
|
191
|
+
|
|
192
|
+
def _p_step_ps_leq_p(z, x, p, ps, gamma, seed):
|
|
193
|
+
r"""Run the damped p-step for ps <= p
|
|
194
|
+
|
|
195
|
+
:param z: a [R, M] tensor of events for R transitions and M units
|
|
196
|
+
:param x: a [S, M] tensor of state values for S states and M units
|
|
197
|
+
:param p: a [R, M] tensor of transition rates for the current time series
|
|
198
|
+
:param ps: a [R, M] tensor of transition rates for the rippled time series
|
|
199
|
+
:param gamma: a scalar ($\gamma \geq 1$) giving the damping.
|
|
200
|
+
:param seed: a random seed.
|
|
201
|
+
:returns: a tuple of ([R, M] tensor of new event numbers,
|
|
202
|
+
log_acceptance_correction)
|
|
203
|
+
"""
|
|
204
|
+
seeds = samplers.split_seed(seed, n=4, salt="_p_step_ps_leq_p")
|
|
205
|
+
|
|
206
|
+
# Sample z_new
|
|
207
|
+
z_new = tfd.Binomial(total_count=z, probs=ps / p).sample(seed=seeds[0])
|
|
208
|
+
|
|
209
|
+
# Jacobian
|
|
210
|
+
upper_bound = 2 * p - ps
|
|
211
|
+
damping_fn = DampingFunction(ps, upper_bound, gamma)
|
|
212
|
+
|
|
213
|
+
w = z - z_new
|
|
214
|
+
w_prime = tfd.Binomial(
|
|
215
|
+
total_count=x - z,
|
|
216
|
+
probs=(upper_bound - p) / (1 - p),
|
|
217
|
+
).sample(seed=seeds[1])
|
|
218
|
+
|
|
219
|
+
def jacobian(x):
|
|
220
|
+
return damping_fn.log_inverse_jacobian(damping_fn(x))
|
|
221
|
+
|
|
222
|
+
log_jacobian_fwd = _reduce_uniforms(
|
|
223
|
+
ps,
|
|
224
|
+
p,
|
|
225
|
+
w,
|
|
226
|
+
jacobian,
|
|
227
|
+
seeds[2],
|
|
228
|
+
) + _reduce_uniforms(p, upper_bound, w_prime, jacobian, seed=seeds[3])
|
|
229
|
+
|
|
230
|
+
return (
|
|
231
|
+
z_new,
|
|
232
|
+
-log_jacobian_fwd,
|
|
233
|
+
) # log_acceptance_correction
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
def _pstep(z, x, p, ps, gamma, seed):
|
|
237
|
+
"""Compute the p-step of Rippler.
|
|
238
|
+
|
|
239
|
+
Since there are two possible distributions to draw from,
|
|
240
|
+
but both are Binomial, we compute `offset`, `total_count`, and
|
|
241
|
+
`prob` parameters for both branches, and select which we need
|
|
242
|
+
based on p <= ps.
|
|
243
|
+
|
|
244
|
+
:param z: current $z$
|
|
245
|
+
:param x: current $x$
|
|
246
|
+
:param p: current probability
|
|
247
|
+
:param ps: new probability
|
|
248
|
+
"""
|
|
249
|
+
with tf.name_scope("_pstep"):
|
|
250
|
+
seed1, seed2 = samplers.split_seed(seed, salt="_pstep")
|
|
251
|
+
ps_leq_p = _p_step_ps_leq_p(z, x, p, ps, gamma, seed1)
|
|
252
|
+
ps_gt_p = _p_step_ps_gt_p(z, x, p, ps, gamma, seed2)
|
|
253
|
+
|
|
254
|
+
z_prime = tf.where(ps <= p, ps_leq_p[0], ps_gt_p[0])
|
|
255
|
+
log_acceptance_correction = tf.where(ps <= p, ps_leq_p[1], ps_gt_p[1])
|
|
256
|
+
|
|
257
|
+
return z_prime, log_acceptance_correction
|
|
258
|
+
|
|
259
|
+
|
|
260
|
+
def _xstep(z_prime, x, xs, ps, seed, validate_args=False):
|
|
261
|
+
"""Computes the x-step of the Rippler algorithm.
|
|
262
|
+
|
|
263
|
+
Both xs >= x and xs < x are sampled and results selected.
|
|
264
|
+
"""
|
|
265
|
+
with tf.name_scope("_xstep"):
|
|
266
|
+
seeds = samplers.split_seed(seed, salt="_xstep")
|
|
267
|
+
|
|
268
|
+
# xs >= x
|
|
269
|
+
# Switch off `validate_args` because `xs-x` may be -ve.
|
|
270
|
+
z_new_geq = z_prime + tfd.Binomial(
|
|
271
|
+
xs - x,
|
|
272
|
+
probs=ps,
|
|
273
|
+
validate_args=False,
|
|
274
|
+
name="_xstep_Binomial",
|
|
275
|
+
).sample(seed=seeds[0])
|
|
276
|
+
|
|
277
|
+
# xs < x - explicitly vectorize
|
|
278
|
+
def safe_hypergeom(N, K, n): # noqa: N803
|
|
279
|
+
# xs is clipped to min(x, xs) to avoid errors in the Hypergeometric
|
|
280
|
+
# sampler these values won't be selected anyway due to the
|
|
281
|
+
# xs >= x condition below.
|
|
282
|
+
return Hypergeometric(
|
|
283
|
+
N=N,
|
|
284
|
+
K=K,
|
|
285
|
+
n=tf.math.minimum(N, n),
|
|
286
|
+
validate_args=False,
|
|
287
|
+
name="_xstep_Hypergeom",
|
|
288
|
+
)
|
|
289
|
+
|
|
290
|
+
z_new_lt = safe_hypergeom(x, z_prime, xs).sample(seed=seeds[1])
|
|
291
|
+
|
|
292
|
+
return tf.where(xs >= x, z_new_geq, z_new_lt)
|
|
293
|
+
|
|
294
|
+
|
|
295
|
+
def _dispatch_update(z, x, p, xs, ps, gamma, seed, validate_args=False):
|
|
296
|
+
r"""Dispatches update function based on values of
|
|
297
|
+
parameters.
|
|
298
|
+
|
|
299
|
+
:param z: current $z$
|
|
300
|
+
:param x: current $x$
|
|
301
|
+
:param p: $p$ current probability
|
|
302
|
+
:param xs: $x^\star$ new state
|
|
303
|
+
:param ps: $p_star$ new probability
|
|
304
|
+
|
|
305
|
+
:returns: an updated number of events
|
|
306
|
+
"""
|
|
307
|
+
with tf.name_scope("dispatch_update"):
|
|
308
|
+
p = tf.convert_to_tensor(p)
|
|
309
|
+
ps = tf.convert_to_tensor(ps)
|
|
310
|
+
z = tf.cast(z, p.dtype)
|
|
311
|
+
x = tf.cast(x, p.dtype)
|
|
312
|
+
xs = tf.cast(xs, p.dtype)
|
|
313
|
+
|
|
314
|
+
seeds = samplers.split_seed(seed, salt="_dispatch_update")
|
|
315
|
+
|
|
316
|
+
z_prime, log_acceptance_correction = _pstep(
|
|
317
|
+
z, x, p, ps, gamma, seed=seeds[0]
|
|
318
|
+
)
|
|
319
|
+
z_new = _xstep(
|
|
320
|
+
z_prime, x, xs, ps, seed=seeds[1], validate_args=validate_args
|
|
321
|
+
)
|
|
322
|
+
|
|
323
|
+
return z_new, log_acceptance_correction, ps > p
|
|
324
|
+
|
|
325
|
+
|
|
326
|
+
# Tests
|
|
327
|
+
def test_dispatch():
|
|
328
|
+
tf.debugging.assert_scalar(
|
|
329
|
+
_dispatch_update(z=10, x=100, p=0.1, xs=100, ps=0.1)
|
|
330
|
+
)
|
|
331
|
+
tf.debugging.assert_scalar(
|
|
332
|
+
_dispatch_update(z=10, x=100, p=0.1, xs=100, ps=0.05)
|
|
333
|
+
)
|
|
334
|
+
tf.debugging.assert_scalar(
|
|
335
|
+
_dispatch_update(z=10, x=100, p=0.1, xs=100, ps=0.2)
|
|
336
|
+
)
|
|
337
|
+
tf.debugging.assert_scalar(
|
|
338
|
+
_dispatch_update(z=10, x=100, p=0.1, xs=50, ps=0.1)
|
|
339
|
+
)
|
|
340
|
+
tf.debugging.assert_scalar(
|
|
341
|
+
_dispatch_update(z=10, x=100, p=0.1, xs=50, ps=0.05)
|
|
342
|
+
)
|
|
343
|
+
tf.debugging.assert_scalar(
|
|
344
|
+
_dispatch_update(z=10, x=100, p=0.1, xs=50, ps=0.2)
|
|
345
|
+
)
|
|
346
|
+
tf.debugging.assert_scalar(
|
|
347
|
+
_dispatch_update(z=10, x=100, p=0.1, xs=200, ps=0.1)
|
|
348
|
+
)
|
|
349
|
+
tf.debugging.assert_scalar(
|
|
350
|
+
_dispatch_update(z=10, x=100, p=0.1, xs=200, ps=0.05)
|
|
351
|
+
)
|
|
352
|
+
tf.debugging.assert_scalar(
|
|
353
|
+
_dispatch_update(z=10, x=100, p=0.1, xs=200, ps=0.2)
|
|
354
|
+
)
|
|
355
|
+
|
|
356
|
+
|
|
357
|
+
def default_initial_ripple(model, current_events, current_state, seed):
|
|
358
|
+
"""Produces the initial ripple.
|
|
359
|
+
|
|
360
|
+
:param model: an instance of `DiscreteTimeStateTransitionModel`
|
|
361
|
+
:param current_events: a tensor of events in [T, R, M] order
|
|
362
|
+
:param current_state: a tensor of state in [T, S, M] order
|
|
363
|
+
:param seed: the seed to initialise the ripple
|
|
364
|
+
|
|
365
|
+
:returns: a tuple of `(proposed_time_idx, new_events_t, current_state_t)`
|
|
366
|
+
"""
|
|
367
|
+
init_time_seed, init_pop_seed, init_events_seed = samplers.split_seed(
|
|
368
|
+
seed, n=3, salt="_initial_ripple"
|
|
369
|
+
)
|
|
370
|
+
|
|
371
|
+
# Choose timepoint, t
|
|
372
|
+
proposed_time_idx = UniformInteger(low=0, high=model.num_steps).sample(
|
|
373
|
+
seed=init_time_seed
|
|
374
|
+
)
|
|
375
|
+
current_state_t = tf.gather(current_state, proposed_time_idx, axis=-3)
|
|
376
|
+
|
|
377
|
+
# Choose subpopulation - KCategorical?
|
|
378
|
+
proposed_pop_idx = UniformInteger(
|
|
379
|
+
low=0, high=current_events.shape[-1]
|
|
380
|
+
).sample(seed=init_pop_seed)
|
|
381
|
+
|
|
382
|
+
# Choose new infection events at time t
|
|
383
|
+
proposed_transition_rates = tf.stack(
|
|
384
|
+
model.transition_rates(
|
|
385
|
+
proposed_time_idx, tf.transpose(current_state_t)
|
|
386
|
+
),
|
|
387
|
+
axis=0,
|
|
388
|
+
)
|
|
389
|
+
prob_t = 1.0 - tf.math.exp(
|
|
390
|
+
-tf.gather(proposed_transition_rates[0], proposed_pop_idx, axis=-1)
|
|
391
|
+
* model.time_delta,
|
|
392
|
+
) # First event to perturb.
|
|
393
|
+
|
|
394
|
+
required_state = tf.gather(current_state_t[0], proposed_pop_idx, axis=-1)
|
|
395
|
+
new_si_events_t = tfd.Binomial(
|
|
396
|
+
total_count=required_state,
|
|
397
|
+
probs=prob_t, # Perturb SI events here
|
|
398
|
+
).sample(seed=init_events_seed)
|
|
399
|
+
|
|
400
|
+
new_events_t = tf.tensor_scatter_nd_update(
|
|
401
|
+
current_events[proposed_time_idx],
|
|
402
|
+
[[0, proposed_pop_idx]],
|
|
403
|
+
[new_si_events_t],
|
|
404
|
+
)
|
|
405
|
+
|
|
406
|
+
return proposed_time_idx, new_events_t, current_state_t
|
|
407
|
+
|
|
408
|
+
|
|
409
|
+
def damped_initial_ripple_fn(resampling_fraction=1.0):
|
|
410
|
+
"""Construct a damped initial ripple function
|
|
411
|
+
|
|
412
|
+
Args
|
|
413
|
+
----
|
|
414
|
+
sampling_fraction: the fraction of the initial state to resample
|
|
415
|
+
|
|
416
|
+
Returns
|
|
417
|
+
-------
|
|
418
|
+
a callable fn(model, current_events, current_state, seed)
|
|
419
|
+
"""
|
|
420
|
+
|
|
421
|
+
def fn(model, current_events, current_state, seed):
|
|
422
|
+
"""Produces a damped initial ripple.
|
|
423
|
+
|
|
424
|
+
:param model: an instance of `DiscreteTimeStateTransitionModel`
|
|
425
|
+
:param current_events: a tensor of events in [T, R, M] order
|
|
426
|
+
:param current_state: a tensor of state in [T, S, M] order
|
|
427
|
+
:param resampling_fraction: a value between 0 and 1 where 0
|
|
428
|
+
is fully damped, and 1 is no damping
|
|
429
|
+
:param seed: the seed to initialise the ripple
|
|
430
|
+
|
|
431
|
+
:returns: a tuple of `(proposed_time_idx, new_events_t,
|
|
432
|
+
current_state_t)`
|
|
433
|
+
"""
|
|
434
|
+
|
|
435
|
+
init_time_seed, init_pop_seed, hypergeom_seed, binom_seed = (
|
|
436
|
+
samplers.split_seed(seed, n=4, salt="_initial_ripple")
|
|
437
|
+
)
|
|
438
|
+
|
|
439
|
+
# Choose timepoint, t
|
|
440
|
+
proposed_time_idx = UniformInteger(low=0, high=model.num_steps).sample(
|
|
441
|
+
seed=init_time_seed
|
|
442
|
+
)
|
|
443
|
+
current_state_t = tf.gather(current_state, proposed_time_idx, axis=-3)
|
|
444
|
+
|
|
445
|
+
# Choose subpopulation - KCategorical?
|
|
446
|
+
proposed_pop_idx = UniformInteger(
|
|
447
|
+
low=0, high=current_events.shape[-1]
|
|
448
|
+
).sample(seed=init_pop_seed)
|
|
449
|
+
|
|
450
|
+
# Choose new infection events at time t
|
|
451
|
+
proposed_transition_rates = tf.stack(
|
|
452
|
+
model.transition_rates(
|
|
453
|
+
proposed_time_idx, tf.transpose(current_state_t)
|
|
454
|
+
),
|
|
455
|
+
axis=0,
|
|
456
|
+
)
|
|
457
|
+
prob_t = 1.0 - tf.math.exp(
|
|
458
|
+
-tf.gather(proposed_transition_rates[0], proposed_pop_idx, axis=-1)
|
|
459
|
+
* model.time_delta,
|
|
460
|
+
) # First event to perturb.
|
|
461
|
+
|
|
462
|
+
required_state = tf.gather(
|
|
463
|
+
current_state_t[0], proposed_pop_idx, axis=-1
|
|
464
|
+
)
|
|
465
|
+
required_events = current_events[proposed_time_idx, 0, proposed_pop_idx]
|
|
466
|
+
sample_size = tf.math.floor(required_state * resampling_fraction)
|
|
467
|
+
|
|
468
|
+
new_si_events_t = (
|
|
469
|
+
required_events
|
|
470
|
+
- Hypergeometric(
|
|
471
|
+
required_state, required_events, sample_size
|
|
472
|
+
).sample(seed=hypergeom_seed)
|
|
473
|
+
+ tfd.Binomial(total_count=sample_size, probs=prob_t).sample(
|
|
474
|
+
seed=binom_seed
|
|
475
|
+
)
|
|
476
|
+
)
|
|
477
|
+
|
|
478
|
+
new_events_t = tf.tensor_scatter_nd_update(
|
|
479
|
+
current_events[proposed_time_idx],
|
|
480
|
+
[[0, proposed_pop_idx]],
|
|
481
|
+
[new_si_events_t],
|
|
482
|
+
)
|
|
483
|
+
|
|
484
|
+
return proposed_time_idx, new_events_t, current_state_t
|
|
485
|
+
|
|
486
|
+
return fn
|
|
487
|
+
|
|
488
|
+
|
|
489
|
+
def chain_binomial_rippler(
|
|
490
|
+
model,
|
|
491
|
+
current_events,
|
|
492
|
+
initial_ripple_fn,
|
|
493
|
+
ripple_damping_constant,
|
|
494
|
+
seed,
|
|
495
|
+
):
|
|
496
|
+
init_seed, ripple_seed = samplers.split_seed(
|
|
497
|
+
seed, salt="chain_binomial_rippler"
|
|
498
|
+
)
|
|
499
|
+
src_states = model.source_states
|
|
500
|
+
|
|
501
|
+
# Transpose to [T, S/R, M]
|
|
502
|
+
current_events = tf.transpose(current_events, perm=(1, 2, 0))
|
|
503
|
+
|
|
504
|
+
# Calculate current state
|
|
505
|
+
current_state = _compute_state(
|
|
506
|
+
initial_state=tf.transpose(model.initial_state),
|
|
507
|
+
events=current_events,
|
|
508
|
+
stoichiometry=model.stoichiometry,
|
|
509
|
+
)
|
|
510
|
+
|
|
511
|
+
# Begin the ripple by sampling a time point, and perturbing the
|
|
512
|
+
# events at that timepoint
|
|
513
|
+
(
|
|
514
|
+
proposed_time_idx,
|
|
515
|
+
new_events_t,
|
|
516
|
+
current_state_t,
|
|
517
|
+
) = initial_ripple_fn(model, current_events, current_state, init_seed)
|
|
518
|
+
new_events = tf.tensor_scatter_nd_update(
|
|
519
|
+
current_events, indices=[[proposed_time_idx]], updates=[new_events_t]
|
|
520
|
+
)
|
|
521
|
+
|
|
522
|
+
# Propagate from t+1 up to end of the timeseries
|
|
523
|
+
def draw_events(time, new_state_t, current_events_t, current_state_t, seed):
|
|
524
|
+
with tf.name_scope("draw_events"):
|
|
525
|
+
# Calculate transition rates for current and new states
|
|
526
|
+
def transition_probs(time, state):
|
|
527
|
+
rates = tf.stack(
|
|
528
|
+
model.transition_rates(time, tf.transpose(state)), axis=-2
|
|
529
|
+
)
|
|
530
|
+
return 1.0 - tf.math.exp(-rates * model.time_delta)
|
|
531
|
+
|
|
532
|
+
current_p = transition_probs(time, current_state_t)
|
|
533
|
+
new_p = transition_probs(time, new_state_t)
|
|
534
|
+
|
|
535
|
+
new_events, log_acceptance_correction, ps_gt_p = _dispatch_update(
|
|
536
|
+
z=current_events_t,
|
|
537
|
+
x=prefer_static.gather(current_state_t, indices=src_states),
|
|
538
|
+
p=current_p,
|
|
539
|
+
xs=prefer_static.gather(new_state_t, indices=src_states),
|
|
540
|
+
ps=new_p,
|
|
541
|
+
gamma=ripple_damping_constant,
|
|
542
|
+
seed=seed,
|
|
543
|
+
)
|
|
544
|
+
tf.debugging.assert_non_negative(new_events)
|
|
545
|
+
|
|
546
|
+
return new_events, log_acceptance_correction, ps_gt_p
|
|
547
|
+
|
|
548
|
+
def time_loop_body(
|
|
549
|
+
t,
|
|
550
|
+
new_events_t,
|
|
551
|
+
new_state_t,
|
|
552
|
+
new_events_buffer,
|
|
553
|
+
log_acceptance_correction_accum,
|
|
554
|
+
ps_gt_p_accum,
|
|
555
|
+
seed,
|
|
556
|
+
):
|
|
557
|
+
sample_seed, next_seed = samplers.split_seed(
|
|
558
|
+
seed, salt="time_loop_body"
|
|
559
|
+
)
|
|
560
|
+
|
|
561
|
+
# Propagate new_state[t] to new_state[t+1]
|
|
562
|
+
new_state_t1 = new_state_t + tf.einsum(
|
|
563
|
+
"...ik,ij->...jk", new_events_t, model.stoichiometry
|
|
564
|
+
)
|
|
565
|
+
# tf.debugging.assert_non_negative(new_state_t1, summarize=100)
|
|
566
|
+
|
|
567
|
+
# Gather current states and events, and draw new events
|
|
568
|
+
new_events_t1, log_acceptance_correction, ps_gt_p = draw_events(
|
|
569
|
+
t + 1,
|
|
570
|
+
new_state_t1,
|
|
571
|
+
current_events[t + 1],
|
|
572
|
+
current_state[t + 1],
|
|
573
|
+
sample_seed,
|
|
574
|
+
)
|
|
575
|
+
|
|
576
|
+
# Update new_events_buffer
|
|
577
|
+
new_events_buffer = tf.tensor_scatter_nd_update(
|
|
578
|
+
new_events_buffer, indices=[[t + 1]], updates=[new_events_t1]
|
|
579
|
+
)
|
|
580
|
+
|
|
581
|
+
return (
|
|
582
|
+
t + 1,
|
|
583
|
+
new_events_t1,
|
|
584
|
+
new_state_t1,
|
|
585
|
+
new_events_buffer,
|
|
586
|
+
log_acceptance_correction_accum.write(t, log_acceptance_correction),
|
|
587
|
+
ps_gt_p_accum.write(t, ps_gt_p),
|
|
588
|
+
next_seed,
|
|
589
|
+
)
|
|
590
|
+
|
|
591
|
+
def time_loop_cond(t, _1, _2, new_events_buffer, *_3):
|
|
592
|
+
t_stop = t < (model.num_steps - 1)
|
|
593
|
+
delta_stop = tf.reduce_any(new_events_buffer != current_events)
|
|
594
|
+
return t_stop & delta_stop
|
|
595
|
+
|
|
596
|
+
log_acceptance_correction_accum = tf.TensorArray(
|
|
597
|
+
current_state.dtype, size=model.num_steps
|
|
598
|
+
)
|
|
599
|
+
ps_gt_p_accum = tf.TensorArray(tf.bool, size=model.num_steps)
|
|
600
|
+
(
|
|
601
|
+
_,
|
|
602
|
+
_,
|
|
603
|
+
_,
|
|
604
|
+
new_events,
|
|
605
|
+
log_acceptance_correction,
|
|
606
|
+
ps_gt_p_accum,
|
|
607
|
+
_,
|
|
608
|
+
) = tf.while_loop(
|
|
609
|
+
time_loop_cond,
|
|
610
|
+
time_loop_body,
|
|
611
|
+
loop_vars=(
|
|
612
|
+
proposed_time_idx,
|
|
613
|
+
new_events_t,
|
|
614
|
+
current_state_t,
|
|
615
|
+
new_events,
|
|
616
|
+
log_acceptance_correction_accum,
|
|
617
|
+
ps_gt_p_accum,
|
|
618
|
+
ripple_seed,
|
|
619
|
+
),
|
|
620
|
+
) # new_events.shape = [T, R, M]
|
|
621
|
+
|
|
622
|
+
new_events = tf.transpose(new_events, perm=(2, 0, 1))
|
|
623
|
+
|
|
624
|
+
return (
|
|
625
|
+
new_events,
|
|
626
|
+
{
|
|
627
|
+
"log_acceptance_correction": log_acceptance_correction.stack(),
|
|
628
|
+
"is_ps_gt_p": ps_gt_p_accum.stack(),
|
|
629
|
+
"delta": tf.transpose(
|
|
630
|
+
new_events_t - current_events[proposed_time_idx]
|
|
631
|
+
),
|
|
632
|
+
"timepoint": proposed_time_idx,
|
|
633
|
+
"initial_ripple": new_events_t,
|
|
634
|
+
"current_state_t": tf.transpose(current_state_t),
|
|
635
|
+
},
|
|
636
|
+
)
|
|
637
|
+
|
|
638
|
+
|
|
639
|
+
# The Chain Binomial Rippler kernel
|
|
640
|
+
CBRResults = namedtuple(
|
|
641
|
+
"CBRResults",
|
|
642
|
+
[
|
|
643
|
+
"target_log_prob",
|
|
644
|
+
"is_accepted",
|
|
645
|
+
"delta",
|
|
646
|
+
"current_state_t",
|
|
647
|
+
"initial_ripple",
|
|
648
|
+
"timepoint",
|
|
649
|
+
"proposed_state",
|
|
650
|
+
"proposed_target_log_prob",
|
|
651
|
+
"log_acceptance_correction",
|
|
652
|
+
"is_ps_gt_p",
|
|
653
|
+
"seed",
|
|
654
|
+
],
|
|
655
|
+
)
|
|
656
|
+
|
|
657
|
+
|
|
658
|
+
class DampedCBRKernel(tfp.mcmc.TransitionKernel):
|
|
659
|
+
def __init__(
|
|
660
|
+
self,
|
|
661
|
+
target_log_prob_fn,
|
|
662
|
+
model,
|
|
663
|
+
initial_ripple_fn=default_initial_ripple,
|
|
664
|
+
ripple_damping_constant=1.0,
|
|
665
|
+
name=None,
|
|
666
|
+
):
|
|
667
|
+
self._target_log_prob_fn = target_log_prob_fn
|
|
668
|
+
self._model = model
|
|
669
|
+
|
|
670
|
+
name = mcmc_util.make_name(name, "CBRKernel", "")
|
|
671
|
+
|
|
672
|
+
self._parameters = {
|
|
673
|
+
"target_log_prob_fn": target_log_prob_fn,
|
|
674
|
+
"model": model,
|
|
675
|
+
"initial_ripple_fn": initial_ripple_fn,
|
|
676
|
+
"ripple_damping_constant": ripple_damping_constant,
|
|
677
|
+
"name": name,
|
|
678
|
+
}
|
|
679
|
+
|
|
680
|
+
@property
|
|
681
|
+
def is_calibrated(self):
|
|
682
|
+
return True
|
|
683
|
+
|
|
684
|
+
@property
|
|
685
|
+
def target_log_prob(self):
|
|
686
|
+
return self._target_log_prob_fn
|
|
687
|
+
|
|
688
|
+
@property
|
|
689
|
+
def model(self):
|
|
690
|
+
return self._model
|
|
691
|
+
|
|
692
|
+
@property
|
|
693
|
+
def name(self):
|
|
694
|
+
return self._parameters["name"]
|
|
695
|
+
|
|
696
|
+
@property
|
|
697
|
+
def initial_ripple_fn(self):
|
|
698
|
+
return self._parameters["initial_ripple_fn"]
|
|
699
|
+
|
|
700
|
+
@property
|
|
701
|
+
def ripple_damping_constant(self):
|
|
702
|
+
return self._parameters["ripple_damping_constant"]
|
|
703
|
+
|
|
704
|
+
def one_step(self, current_state, previous_results, seed=None):
|
|
705
|
+
with tf.name_scope("CBRKernel/one_step"):
|
|
706
|
+
seed_rippler, seed_u, seed_results = samplers.split_seed(
|
|
707
|
+
seed, n=3, salt="cbr_kernel"
|
|
708
|
+
)
|
|
709
|
+
|
|
710
|
+
if mcmc_util.is_list_like(current_state):
|
|
711
|
+
current_state_parts = list(current_state)
|
|
712
|
+
else:
|
|
713
|
+
current_state_parts = [current_state]
|
|
714
|
+
|
|
715
|
+
if len(current_state_parts) > 1:
|
|
716
|
+
warnings.warn(
|
|
717
|
+
"CBRKernel.boostrap_results: multiple state parts detected,\
|
|
718
|
+
but only the first will be used",
|
|
719
|
+
stacklevel=2,
|
|
720
|
+
)
|
|
721
|
+
|
|
722
|
+
current_state_part = tf.convert_to_tensor(
|
|
723
|
+
current_state_parts[0], name="current_state"
|
|
724
|
+
)
|
|
725
|
+
|
|
726
|
+
proposed_state, proposal_results = chain_binomial_rippler(
|
|
727
|
+
self.model,
|
|
728
|
+
current_state_part,
|
|
729
|
+
initial_ripple_fn=self.initial_ripple_fn,
|
|
730
|
+
ripple_damping_constant=self.ripple_damping_constant,
|
|
731
|
+
seed=seed_rippler,
|
|
732
|
+
)
|
|
733
|
+
|
|
734
|
+
proposed_target_log_prob = self.target_log_prob(proposed_state)
|
|
735
|
+
|
|
736
|
+
delta_logp = (
|
|
737
|
+
proposed_target_log_prob
|
|
738
|
+
- previous_results.target_log_prob
|
|
739
|
+
+ tf.reduce_sum(proposal_results["log_acceptance_correction"])
|
|
740
|
+
)
|
|
741
|
+
|
|
742
|
+
def accept():
|
|
743
|
+
return (
|
|
744
|
+
proposed_state,
|
|
745
|
+
CBRResults(
|
|
746
|
+
target_log_prob=proposed_target_log_prob,
|
|
747
|
+
is_accepted=tf.constant(True),
|
|
748
|
+
delta=proposal_results["delta"],
|
|
749
|
+
current_state_t=proposal_results["current_state_t"],
|
|
750
|
+
initial_ripple=proposal_results["initial_ripple"],
|
|
751
|
+
timepoint=proposal_results["timepoint"],
|
|
752
|
+
proposed_state=proposed_state,
|
|
753
|
+
proposed_target_log_prob=proposed_target_log_prob,
|
|
754
|
+
log_acceptance_correction=proposal_results[
|
|
755
|
+
"log_acceptance_correction"
|
|
756
|
+
],
|
|
757
|
+
is_ps_gt_p=proposal_results["is_ps_gt_p"],
|
|
758
|
+
seed=seed_results,
|
|
759
|
+
),
|
|
760
|
+
)
|
|
761
|
+
|
|
762
|
+
def reject():
|
|
763
|
+
return (
|
|
764
|
+
current_state_part,
|
|
765
|
+
CBRResults(
|
|
766
|
+
target_log_prob=previous_results.target_log_prob,
|
|
767
|
+
is_accepted=tf.constant(False),
|
|
768
|
+
delta=proposal_results["delta"],
|
|
769
|
+
current_state_t=proposal_results["current_state_t"],
|
|
770
|
+
initial_ripple=proposal_results["initial_ripple"],
|
|
771
|
+
timepoint=proposal_results["timepoint"],
|
|
772
|
+
proposed_state=proposed_state,
|
|
773
|
+
proposed_target_log_prob=proposed_target_log_prob,
|
|
774
|
+
log_acceptance_correction=proposal_results[
|
|
775
|
+
"log_acceptance_correction"
|
|
776
|
+
],
|
|
777
|
+
is_ps_gt_p=proposal_results["is_ps_gt_p"],
|
|
778
|
+
seed=seed_results,
|
|
779
|
+
),
|
|
780
|
+
)
|
|
781
|
+
|
|
782
|
+
u = tf.math.log(
|
|
783
|
+
tfd.Uniform(low=tf.zeros(1, dtype=delta_logp.dtype)).sample(
|
|
784
|
+
seed=seed_u
|
|
785
|
+
)
|
|
786
|
+
)
|
|
787
|
+
new_state, results = tf.cond(u < delta_logp, accept, reject)
|
|
788
|
+
|
|
789
|
+
def maybe_flatten(x):
|
|
790
|
+
if mcmc_util.is_list_like(current_state):
|
|
791
|
+
return type(current_state)(new_state)
|
|
792
|
+
return x
|
|
793
|
+
|
|
794
|
+
new_state = maybe_flatten(new_state)
|
|
795
|
+
return new_state, results
|
|
796
|
+
|
|
797
|
+
def bootstrap_results(self, current_state):
|
|
798
|
+
with tf.name_scope("CBRKernel/bootstrap_results"):
|
|
799
|
+
if mcmc_util.is_list_like(current_state):
|
|
800
|
+
current_state_parts = list(current_state)
|
|
801
|
+
else:
|
|
802
|
+
current_state_parts = [current_state]
|
|
803
|
+
|
|
804
|
+
if len(current_state_parts) > 1:
|
|
805
|
+
warnings.warn(
|
|
806
|
+
"CBRKernel.boostrap_results: multiple state parts detected,\
|
|
807
|
+
but only the first will be used",
|
|
808
|
+
stacklevel=2,
|
|
809
|
+
)
|
|
810
|
+
state_part = current_state_parts[0]
|
|
811
|
+
|
|
812
|
+
num_times = state_part.shape[-2]
|
|
813
|
+
num_pop = state_part.shape[-3]
|
|
814
|
+
num_transitions = state_part.shape[-1]
|
|
815
|
+
num_states = self.model.stoichiometry.shape[-1]
|
|
816
|
+
|
|
817
|
+
target_log_prob = self.target_log_prob(state_part)
|
|
818
|
+
|
|
819
|
+
return CBRResults(
|
|
820
|
+
target_log_prob=target_log_prob,
|
|
821
|
+
is_accepted=tf.constant(False),
|
|
822
|
+
delta=tf.zeros((num_pop, num_transitions), state_part.dtype),
|
|
823
|
+
current_state_t=tf.zeros(
|
|
824
|
+
[num_pop, num_states], state_part.dtype
|
|
825
|
+
),
|
|
826
|
+
initial_ripple=tf.zeros(
|
|
827
|
+
(num_transitions, num_pop), state_part.dtype
|
|
828
|
+
),
|
|
829
|
+
timepoint=tf.constant(0, dtype=tf.int32),
|
|
830
|
+
proposed_state=state_part,
|
|
831
|
+
proposed_target_log_prob=target_log_prob,
|
|
832
|
+
log_acceptance_correction=tf.zeros(
|
|
833
|
+
(num_times, num_transitions, num_pop),
|
|
834
|
+
current_state[0].dtype,
|
|
835
|
+
),
|
|
836
|
+
is_ps_gt_p=tf.fill(
|
|
837
|
+
(num_times, num_transitions, num_pop), False
|
|
838
|
+
),
|
|
839
|
+
seed=samplers.sanitize_seed(0),
|
|
840
|
+
)
|