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