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,352 @@
|
|
|
1
|
+
"""Describes a State Transition Model with marginalised baseline
|
|
2
|
+
hazard rates.
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
import tensorflow as tf
|
|
6
|
+
import tensorflow_probability as tfp
|
|
7
|
+
from tensorflow_probability.python.internal import (
|
|
8
|
+
dtype_util,
|
|
9
|
+
reparameterization,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
from gemlib.distributions.discrete_markov import discrete_markov_simulation
|
|
13
|
+
from gemlib.util import batch_gather, compute_state, transition_coords
|
|
14
|
+
|
|
15
|
+
tla = tf.linalg
|
|
16
|
+
tfd = tfp.distributions
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class StateTransitionMarginalModel(tfd.Distribution):
|
|
20
|
+
def __init__(
|
|
21
|
+
self,
|
|
22
|
+
transition_rates,
|
|
23
|
+
baseline_hazard_rate_priors,
|
|
24
|
+
stoichiometry,
|
|
25
|
+
initial_state,
|
|
26
|
+
initial_step,
|
|
27
|
+
time_delta,
|
|
28
|
+
num_steps,
|
|
29
|
+
validate_args=False,
|
|
30
|
+
allow_nan_stats=True,
|
|
31
|
+
name="StateTransitionMarginalModel",
|
|
32
|
+
):
|
|
33
|
+
"""Implements a discrete-time Markov jump process for a state transition
|
|
34
|
+
model.
|
|
35
|
+
|
|
36
|
+
:param transition_rates: a function of the form `fn(t, state)` taking
|
|
37
|
+
the current time `t` and state tensor `state`.
|
|
38
|
+
This function returns a tensor which broadcasts
|
|
39
|
+
to the first dimension of `stoichiometry`.
|
|
40
|
+
Transition rates are assumed to be risk ratios,
|
|
41
|
+
with the baseline hazard rate marginalised out
|
|
42
|
+
from the model.
|
|
43
|
+
:param baseline_hazard_rate_priors: a dictionary of `concentration` and
|
|
44
|
+
`rate` hyperparameters for implicit
|
|
45
|
+
Gamma priors on (marginalised)
|
|
46
|
+
baseline hazard rates. Both
|
|
47
|
+
`concentration` and `rate` should
|
|
48
|
+
broadcast with the number of rows in
|
|
49
|
+
`stoichiometry`.
|
|
50
|
+
:param stoichiometry: the stochiometry matrix for the state transition
|
|
51
|
+
model with rows representing transitions and
|
|
52
|
+
columns representing states.
|
|
53
|
+
:param initial_state: an initial state tensor with inner dimension equal
|
|
54
|
+
to the first dimension of `stoichiometry`.
|
|
55
|
+
:param initial_step: an offset giving the time `t` of the first timestep
|
|
56
|
+
in the model.
|
|
57
|
+
:param time_delta: the size of the time step to be used.
|
|
58
|
+
:param num_steps: the number of time steps across which the model runs.
|
|
59
|
+
"""
|
|
60
|
+
parameters = dict(locals())
|
|
61
|
+
with tf.name_scope(name) as name:
|
|
62
|
+
self._transition_rates = transition_rates
|
|
63
|
+
self._stoichiometry = tf.convert_to_tensor(
|
|
64
|
+
stoichiometry,
|
|
65
|
+
dtype=initial_state.dtype,
|
|
66
|
+
)
|
|
67
|
+
self._initial_state = initial_state
|
|
68
|
+
self._initial_step = initial_step
|
|
69
|
+
self._time_delta = time_delta
|
|
70
|
+
self._num_steps = num_steps
|
|
71
|
+
self._baseline_hazard_rate_priors = baseline_hazard_rate_priors
|
|
72
|
+
|
|
73
|
+
super().__init__(
|
|
74
|
+
dtype=initial_state.dtype,
|
|
75
|
+
reparameterization_type=reparameterization.FULLY_REPARAMETERIZED,
|
|
76
|
+
validate_args=validate_args,
|
|
77
|
+
allow_nan_stats=allow_nan_stats,
|
|
78
|
+
parameters=parameters,
|
|
79
|
+
name=name,
|
|
80
|
+
)
|
|
81
|
+
|
|
82
|
+
self.dtype = initial_state.dtype
|
|
83
|
+
|
|
84
|
+
@property
|
|
85
|
+
def transition_rates(self):
|
|
86
|
+
return self._transition_rates
|
|
87
|
+
|
|
88
|
+
@property
|
|
89
|
+
def baseline_hazard_rate_priors(self):
|
|
90
|
+
return self._baseline_hazard_rate_priors
|
|
91
|
+
|
|
92
|
+
@property
|
|
93
|
+
def stoichiometry(self):
|
|
94
|
+
return self._stoichiometry
|
|
95
|
+
|
|
96
|
+
@property
|
|
97
|
+
def initial_state(self):
|
|
98
|
+
return self._initial_state
|
|
99
|
+
|
|
100
|
+
@property
|
|
101
|
+
def initial_step(self):
|
|
102
|
+
return self._initial_step
|
|
103
|
+
|
|
104
|
+
@property
|
|
105
|
+
def time_delta(self):
|
|
106
|
+
return self._time_delta
|
|
107
|
+
|
|
108
|
+
@property
|
|
109
|
+
def num_steps(self):
|
|
110
|
+
return self._num_steps
|
|
111
|
+
|
|
112
|
+
def _batch_shape(self):
|
|
113
|
+
return tf.TensorShape([])
|
|
114
|
+
|
|
115
|
+
def _event_shape(self):
|
|
116
|
+
shape = tf.TensorShape(
|
|
117
|
+
[
|
|
118
|
+
self.initial_state.shape[0],
|
|
119
|
+
tf.get_static_value(self._num_steps),
|
|
120
|
+
self._stoichiometry.shape[0],
|
|
121
|
+
]
|
|
122
|
+
)
|
|
123
|
+
return shape
|
|
124
|
+
|
|
125
|
+
def _sample_n(self, n, seed=None):
|
|
126
|
+
"""Runs a simulation from the epidemic model
|
|
127
|
+
|
|
128
|
+
:param param: a dictionary of model parameters
|
|
129
|
+
:param state_init: the initial state
|
|
130
|
+
:returns: a tuple of times and simulated states.
|
|
131
|
+
"""
|
|
132
|
+
with tf.name_scope("DiscreteTimeStateTransitionModel.log_prob"):
|
|
133
|
+
|
|
134
|
+
def hazard_fn(t, state):
|
|
135
|
+
return self.transition_rates(t, state)
|
|
136
|
+
|
|
137
|
+
t, sim = discrete_markov_simulation(
|
|
138
|
+
hazard_fn=hazard_fn,
|
|
139
|
+
state=self.initial_state,
|
|
140
|
+
start=self.initial_step,
|
|
141
|
+
end=self.initial_step + self.num_steps * self.time_delta,
|
|
142
|
+
time_step=self.time_delta,
|
|
143
|
+
stoichiometry=self.stoichiometry,
|
|
144
|
+
seed=seed,
|
|
145
|
+
)
|
|
146
|
+
indices = transition_coords(self.stoichiometry)
|
|
147
|
+
sim = batch_gather(sim, indices)
|
|
148
|
+
sim = tf.transpose(sim, perm=(1, 0, 2))
|
|
149
|
+
return tf.expand_dims(sim, 0)
|
|
150
|
+
|
|
151
|
+
def _log_prob(self, y, **kwargs):
|
|
152
|
+
"""Calculates the log probability of observing epidemic events y
|
|
153
|
+
:param y: a list of tensors. The first is of shape [n_times] containing
|
|
154
|
+
times, the second is of shape [n_times, n_states, n_states]
|
|
155
|
+
containing event matrices.
|
|
156
|
+
:param param: a list of parameters
|
|
157
|
+
:returns: a scalar giving the log probability of the epidemic
|
|
158
|
+
"""
|
|
159
|
+
dtype = dtype_util.common_dtype(
|
|
160
|
+
[y, self.initial_state], dtype_hint=self.dtype
|
|
161
|
+
)
|
|
162
|
+
events = tf.convert_to_tensor(y, dtype)
|
|
163
|
+
with tf.name_scope("StateTransitionMarginalModel.log_prob"):
|
|
164
|
+
state_timeseries = compute_state(
|
|
165
|
+
initial_state=self.initial_state,
|
|
166
|
+
events=events,
|
|
167
|
+
stoichiometry=self.stoichiometry,
|
|
168
|
+
closed=True,
|
|
169
|
+
)
|
|
170
|
+
|
|
171
|
+
tms_timeseries = tf.transpose(state_timeseries, perm=(1, 0, 2))
|
|
172
|
+
tmr_events = tf.transpose(events, perm=(1, 0, 2))
|
|
173
|
+
|
|
174
|
+
def fn(elems):
|
|
175
|
+
return tf.stack(self.transition_rates(*elems), axis=-1)
|
|
176
|
+
|
|
177
|
+
rates = tf.vectorized_map(
|
|
178
|
+
fn=fn,
|
|
179
|
+
elems=(
|
|
180
|
+
self._initial_step + tf.range(tms_timeseries.shape[0]),
|
|
181
|
+
tms_timeseries,
|
|
182
|
+
),
|
|
183
|
+
)
|
|
184
|
+
|
|
185
|
+
def integrated_rate_fn():
|
|
186
|
+
"""Use mid-point integration to estimate the constant rate
|
|
187
|
+
over time.
|
|
188
|
+
"""
|
|
189
|
+
integrated_rates = tms_timeseries[..., :-1] * rates
|
|
190
|
+
return (
|
|
191
|
+
integrated_rates[:-1, ...] + integrated_rates[1:, ...]
|
|
192
|
+
) / 2.0
|
|
193
|
+
|
|
194
|
+
integrated_rates = integrated_rate_fn()
|
|
195
|
+
|
|
196
|
+
log_norm_constant = tf.reduce_sum(
|
|
197
|
+
tf.math.multiply_no_nan(
|
|
198
|
+
tf.math.log(integrated_rates), tmr_events
|
|
199
|
+
)
|
|
200
|
+
- tf.math.lgamma(tmr_events + 1.0),
|
|
201
|
+
axis=(0, 1),
|
|
202
|
+
)
|
|
203
|
+
pi_concentration = (
|
|
204
|
+
tf.reduce_sum(tmr_events, axis=(0, 1))
|
|
205
|
+
+ self.baseline_hazard_rate_priors["concentration"]
|
|
206
|
+
)
|
|
207
|
+
pi_rate = (
|
|
208
|
+
tf.reduce_sum(integrated_rates * self.time_delta, axis=(0, 1))
|
|
209
|
+
+ self.baseline_hazard_rate_priors["rate"]
|
|
210
|
+
)
|
|
211
|
+
|
|
212
|
+
log_prob = (
|
|
213
|
+
log_norm_constant
|
|
214
|
+
+ tf.math.lgamma(pi_concentration)
|
|
215
|
+
- (pi_concentration) * tf.math.log(pi_rate)
|
|
216
|
+
)
|
|
217
|
+
|
|
218
|
+
return tf.reduce_sum(log_prob)
|
|
219
|
+
|
|
220
|
+
|
|
221
|
+
class BaselineHazardRateMarginal(tfd.Distribution):
|
|
222
|
+
def __init__(
|
|
223
|
+
self,
|
|
224
|
+
events,
|
|
225
|
+
transition_rate_fn,
|
|
226
|
+
baseline_hazard_rate_priors,
|
|
227
|
+
stoichiometry,
|
|
228
|
+
initial_state,
|
|
229
|
+
initial_step,
|
|
230
|
+
time_delta,
|
|
231
|
+
num_steps,
|
|
232
|
+
validate_args=False,
|
|
233
|
+
allow_nan_stats=True,
|
|
234
|
+
name="StateTransitionMarginalModel",
|
|
235
|
+
):
|
|
236
|
+
"""Implements a discrete-time Markov jump process for a state transition
|
|
237
|
+
model.
|
|
238
|
+
|
|
239
|
+
:param events: a [M, T, R] event tensor
|
|
240
|
+
:param transition_rates: a function of the form `fn(t, state)` taking
|
|
241
|
+
the current time `t` and state tensor `state`.
|
|
242
|
+
This function returns a tensor which broadcasts
|
|
243
|
+
to the first dimension of `stoichiometry`.
|
|
244
|
+
Transition rates are assumed to be risk ratios,
|
|
245
|
+
with the baseline hazard rate marginalised out
|
|
246
|
+
from the model.
|
|
247
|
+
:param baseline_hazard_rate_priors: a dictionary of `concentration` and
|
|
248
|
+
`rate` hyperparameters for implicit
|
|
249
|
+
Gamma priors on (marginalised)
|
|
250
|
+
baseline hazard rates. Both
|
|
251
|
+
`concentration` and `rate` should
|
|
252
|
+
broadcast with the number of rows in
|
|
253
|
+
`stoichiometry`.
|
|
254
|
+
:param stoichiometry: the stochiometry matrix for the state transition
|
|
255
|
+
model with rows representing transitions and
|
|
256
|
+
columns representing states.
|
|
257
|
+
:param initial_state: an initial state tensor with inner dimension equal
|
|
258
|
+
to the first dimension of `stoichiometry`.
|
|
259
|
+
:param initial_step: an offset giving the time `t` of the first timestep
|
|
260
|
+
in the model.
|
|
261
|
+
:param time_delta: the size of the time step to be used.
|
|
262
|
+
:param num_steps: the number of time steps across which the model runs.
|
|
263
|
+
"""
|
|
264
|
+
parameters = dict(locals())
|
|
265
|
+
with tf.name_scope(name) as name:
|
|
266
|
+
self._events = events
|
|
267
|
+
self._transition_rate_fn = transition_rate_fn
|
|
268
|
+
self._stoichiometry = tf.convert_to_tensor(
|
|
269
|
+
stoichiometry,
|
|
270
|
+
dtype=initial_state.dtype,
|
|
271
|
+
)
|
|
272
|
+
self._initial_state = initial_state
|
|
273
|
+
self._initial_step = initial_step
|
|
274
|
+
self._time_delta = time_delta
|
|
275
|
+
self._num_steps = num_steps
|
|
276
|
+
self._baseline_hazard_rate_priors = baseline_hazard_rate_priors
|
|
277
|
+
|
|
278
|
+
super().__init__(
|
|
279
|
+
dtype=initial_state.dtype,
|
|
280
|
+
reparameterization_type=reparameterization.FULLY_REPARAMETERIZED,
|
|
281
|
+
validate_args=validate_args,
|
|
282
|
+
allow_nan_stats=allow_nan_stats,
|
|
283
|
+
parameters=parameters,
|
|
284
|
+
name=name,
|
|
285
|
+
)
|
|
286
|
+
|
|
287
|
+
self.dtype = initial_state.dtype
|
|
288
|
+
|
|
289
|
+
@property
|
|
290
|
+
def baseline_hazard_rate_priors(self):
|
|
291
|
+
return self._baseline_hazard_rate_priors
|
|
292
|
+
|
|
293
|
+
def _batch_shape(self):
|
|
294
|
+
return tf.TensorShape(())
|
|
295
|
+
|
|
296
|
+
def _event_shape(self):
|
|
297
|
+
shape = tf.TensorShape(self._events.shape[-1])
|
|
298
|
+
return shape
|
|
299
|
+
|
|
300
|
+
def concentration(self):
|
|
301
|
+
"""Calculates the concentration parameter"""
|
|
302
|
+
return (
|
|
303
|
+
tf.reduce_sum(self._events, axis=(0, 1))
|
|
304
|
+
+ self.baseline_hazard_rate_priors["concentration"]
|
|
305
|
+
)
|
|
306
|
+
|
|
307
|
+
def rate(self):
|
|
308
|
+
"""Calculates the rate parameter"""
|
|
309
|
+
state = compute_state(
|
|
310
|
+
self._initial_state, self._events, self._stoichiometry, closed=True
|
|
311
|
+
)
|
|
312
|
+
tms_state = tf.transpose(state, perm=(1, 0, 2))
|
|
313
|
+
|
|
314
|
+
def fn(elems):
|
|
315
|
+
return tf.stack(self._transition_rate_fn(*elems), axis=-1)
|
|
316
|
+
|
|
317
|
+
rates = tf.vectorized_map(
|
|
318
|
+
fn=fn,
|
|
319
|
+
elems=(
|
|
320
|
+
self._initial_step + tf.range(tms_state.shape[0]),
|
|
321
|
+
tms_state,
|
|
322
|
+
),
|
|
323
|
+
)
|
|
324
|
+
|
|
325
|
+
def integrated_rate_fn():
|
|
326
|
+
"""Use mid-point integration to estimate the constant rate
|
|
327
|
+
over time.
|
|
328
|
+
"""
|
|
329
|
+
integrated_rates = tms_state[..., :-1] * rates
|
|
330
|
+
return (
|
|
331
|
+
integrated_rates[:-1, ...] + integrated_rates[1:, ...]
|
|
332
|
+
) / 2.0
|
|
333
|
+
|
|
334
|
+
return (
|
|
335
|
+
tf.reduce_sum(integrated_rate_fn(), axis=(-3, -2))
|
|
336
|
+
+ self._baseline_hazard_rate_priors["rate"]
|
|
337
|
+
)
|
|
338
|
+
|
|
339
|
+
def _sample_n(self, n, seed=None):
|
|
340
|
+
tf.print("Concentration:", self.concentration())
|
|
341
|
+
tf.print("Rate:", self.rate())
|
|
342
|
+
rv = tfd.Gamma(concentration=self.concentration(), rate=self.rate())
|
|
343
|
+
return rv.sample(n, seed=seed)
|
|
344
|
+
|
|
345
|
+
def _log_prob(self, y, **kwargs):
|
|
346
|
+
"""Calculates the log prob"""
|
|
347
|
+
rv = tfd.Gamma(concentration=self.concentration(), rate=self.rate())
|
|
348
|
+
return tf.reduce_sum(rv.log_prob(y))
|
|
349
|
+
|
|
350
|
+
def _mean(self):
|
|
351
|
+
rv = tfd.Gamma(concentration=self.concentration(), rate=self.rate())
|
|
352
|
+
return rv.mean()
|
|
@@ -0,0 +1,144 @@
|
|
|
1
|
+
"""Hypergeometric random variable"""
|
|
2
|
+
|
|
3
|
+
# ruff: noqa: N803, N802
|
|
4
|
+
|
|
5
|
+
import tensorflow as tf
|
|
6
|
+
import tensorflow_probability as tfp
|
|
7
|
+
from tensorflow_probability.python.internal import (
|
|
8
|
+
dtype_util,
|
|
9
|
+
parameter_properties,
|
|
10
|
+
reparameterization,
|
|
11
|
+
)
|
|
12
|
+
|
|
13
|
+
from gemlib.distributions.hypergeometric_sampler import sample_hypergeometric
|
|
14
|
+
|
|
15
|
+
tfd = tfp.distributions
|
|
16
|
+
|
|
17
|
+
__all__ = ["Hypergeometric"]
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def _log_factorial(x):
|
|
21
|
+
"""Computes x!"""
|
|
22
|
+
return tf.math.lgamma(x + 1.0)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def _log_choose(n, k):
|
|
26
|
+
"""Computes nCk"""
|
|
27
|
+
return _log_factorial(n) - _log_factorial(k) - _log_factorial(n - k)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
class Hypergeometric(tfd.Distribution):
|
|
31
|
+
def __init__(
|
|
32
|
+
self,
|
|
33
|
+
N,
|
|
34
|
+
K,
|
|
35
|
+
n,
|
|
36
|
+
validate_args=False,
|
|
37
|
+
allow_nan_stats=True,
|
|
38
|
+
name="Hypergeometric",
|
|
39
|
+
):
|
|
40
|
+
"""Hypergeometric distribution
|
|
41
|
+
|
|
42
|
+
Args:
|
|
43
|
+
----
|
|
44
|
+
N: Population size
|
|
45
|
+
K: number of units of interest in the population
|
|
46
|
+
n: size of sample drawn from the population
|
|
47
|
+
validate_args: should arguments be validated for correctness
|
|
48
|
+
allow_nan_stats: allow NaN to be returned for mode, mean, variance...
|
|
49
|
+
|
|
50
|
+
"""
|
|
51
|
+
parameters = dict(locals())
|
|
52
|
+
with tf.name_scope(name) as name:
|
|
53
|
+
dtype = dtype_util.common_dtype([N, K, n], tf.float32)
|
|
54
|
+
self._N = tf.cast(N, dtype=dtype)
|
|
55
|
+
self._K = tf.cast(K, dtype=dtype)
|
|
56
|
+
self._n = tf.cast(n, dtype=dtype)
|
|
57
|
+
self._n_positive_mask = N > 0.0
|
|
58
|
+
super().__init__(
|
|
59
|
+
dtype=dtype,
|
|
60
|
+
reparameterization_type=reparameterization.FULLY_REPARAMETERIZED,
|
|
61
|
+
validate_args=validate_args,
|
|
62
|
+
allow_nan_stats=allow_nan_stats,
|
|
63
|
+
parameters=parameters,
|
|
64
|
+
name=name,
|
|
65
|
+
)
|
|
66
|
+
if validate_args is True:
|
|
67
|
+
tf.debugging.assert_non_negative(
|
|
68
|
+
N, message="N must be non-negative"
|
|
69
|
+
)
|
|
70
|
+
tf.debugging.assert_less_equal(K, N, message="K must be <= N")
|
|
71
|
+
tf.debugging.assert_less_equal(n, N, message="n must be <= N")
|
|
72
|
+
|
|
73
|
+
@classmethod
|
|
74
|
+
def _parameter_properties(cls, dtype, num_classes=None):
|
|
75
|
+
return {
|
|
76
|
+
"N": parameter_properties.ParameterProperties(
|
|
77
|
+
default_constraining_bijector_fn=parameter_properties.BIJECTOR_NOT_IMPLEMENTED
|
|
78
|
+
),
|
|
79
|
+
"K": parameter_properties.ParameterProperties(),
|
|
80
|
+
"n": parameter_properties.ParameterProperties(),
|
|
81
|
+
}
|
|
82
|
+
|
|
83
|
+
@staticmethod
|
|
84
|
+
def _param_shapes(sample_shape):
|
|
85
|
+
return dict(
|
|
86
|
+
zip(
|
|
87
|
+
("N", "K", "n"),
|
|
88
|
+
([tf.convert_to_tensor(sample_shape, dtype=tf.int32)] * 3),
|
|
89
|
+
)
|
|
90
|
+
)
|
|
91
|
+
|
|
92
|
+
@classmethod
|
|
93
|
+
def _params_event_ndims(cls):
|
|
94
|
+
return {"N": 0, "K": 0, "n": 0}
|
|
95
|
+
|
|
96
|
+
@property
|
|
97
|
+
def N(self):
|
|
98
|
+
"""Population size"""
|
|
99
|
+
return self._parameters["N"]
|
|
100
|
+
|
|
101
|
+
@property
|
|
102
|
+
def K(self):
|
|
103
|
+
"""Number of units of interest in population"""
|
|
104
|
+
return self._parameters["K"]
|
|
105
|
+
|
|
106
|
+
@property
|
|
107
|
+
def n(self):
|
|
108
|
+
"""Sample size"""
|
|
109
|
+
return self._parameters["n"]
|
|
110
|
+
|
|
111
|
+
def _default_event_space_bijector(self):
|
|
112
|
+
return
|
|
113
|
+
|
|
114
|
+
def _event_shape_tensor(self):
|
|
115
|
+
return tf.constant([], dtype=tf.int32)
|
|
116
|
+
|
|
117
|
+
def _event_shape(self):
|
|
118
|
+
return tf.TensorShape([])
|
|
119
|
+
|
|
120
|
+
def _sample_n(self, n, seed=None):
|
|
121
|
+
with tf.name_scope(self.name + "/sample_n"):
|
|
122
|
+
sample = sample_hypergeometric(n, self.N, self.K, self.n, seed=seed)
|
|
123
|
+
sample = tf.where(self._n_positive_mask, sample, 0.0)
|
|
124
|
+
return sample
|
|
125
|
+
|
|
126
|
+
def _log_prob(self, x):
|
|
127
|
+
numerator = _log_choose(self._K, x) + _log_choose(
|
|
128
|
+
self._N - self._K, self._n - x
|
|
129
|
+
)
|
|
130
|
+
denominator = _log_choose(self._N, self._n)
|
|
131
|
+
return numerator - denominator
|
|
132
|
+
|
|
133
|
+
def _mode(self):
|
|
134
|
+
return tf.math.floor((self._n + 1) * (self._K + 1) / (self._N + 2))
|
|
135
|
+
|
|
136
|
+
def _mean(self):
|
|
137
|
+
return self._n * self._K / self._N
|
|
138
|
+
|
|
139
|
+
def _variance(self):
|
|
140
|
+
n = self._n
|
|
141
|
+
N = self._N
|
|
142
|
+
K = self._K
|
|
143
|
+
|
|
144
|
+
return n * (K / N) * (N - K) / N * (N - n) / (N - 1)
|
|
@@ -0,0 +1,103 @@
|
|
|
1
|
+
"""Hypergeometric sampling algorithm"""
|
|
2
|
+
|
|
3
|
+
# ruff: noqa: N803
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import tensorflow as tf
|
|
7
|
+
from tensorflow_probability.python.internal import (
|
|
8
|
+
batched_rejection_sampler as brs,
|
|
9
|
+
)
|
|
10
|
+
from tensorflow_probability.python.internal import (
|
|
11
|
+
dtype_util,
|
|
12
|
+
samplers,
|
|
13
|
+
tensor_util,
|
|
14
|
+
)
|
|
15
|
+
from tensorflow_probability.python.internal import prefer_static as ps
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def sample_hypergeometric(num_samples, N, K, n, seed=None):
|
|
19
|
+
dtype = dtype_util.common_dtype([N, K, n], tf.float32)
|
|
20
|
+
N = tensor_util.convert_nonref_to_tensor(N, dtype, name="N")
|
|
21
|
+
K = tensor_util.convert_nonref_to_tensor(K, dtype, name="K")
|
|
22
|
+
n = tensor_util.convert_nonref_to_tensor(n, dtype, name="n")
|
|
23
|
+
good_params_mask = (N >= 1.0) & (N >= K) & (n <= N)
|
|
24
|
+
N = tf.where(good_params_mask, N, 100.0)
|
|
25
|
+
K = tf.where(good_params_mask, K, 50.0)
|
|
26
|
+
n = tf.where(good_params_mask, n, 50.0)
|
|
27
|
+
sample_shape = ps.concat(
|
|
28
|
+
[
|
|
29
|
+
[num_samples],
|
|
30
|
+
ps.broadcast_shape(
|
|
31
|
+
ps.broadcast_shape(ps.shape(N), ps.shape(K)), ps.shape(n)
|
|
32
|
+
),
|
|
33
|
+
],
|
|
34
|
+
axis=0,
|
|
35
|
+
)
|
|
36
|
+
|
|
37
|
+
# First Transform N, K, n such that
|
|
38
|
+
# N / 2 >= K, N / 2 >= n
|
|
39
|
+
is_k_small = 0.5 * N >= K
|
|
40
|
+
is_n_small = n <= 0.5 * N
|
|
41
|
+
previous_K = K
|
|
42
|
+
previous_n = n
|
|
43
|
+
K = tf.where(is_k_small, K, N - K)
|
|
44
|
+
n = tf.where(is_n_small, n, N - n)
|
|
45
|
+
|
|
46
|
+
# TODO: Can we write this in a more numerically stable way?
|
|
47
|
+
def _log_hypergeometric_coeff(x):
|
|
48
|
+
return (
|
|
49
|
+
tf.math.lgamma(x + 1.0)
|
|
50
|
+
+ tf.math.lgamma(K - x + 1.0)
|
|
51
|
+
+ tf.math.lgamma(n - x + 1.0)
|
|
52
|
+
+ tf.math.lgamma(N - K - n + x + 1.0)
|
|
53
|
+
)
|
|
54
|
+
|
|
55
|
+
p = K / N
|
|
56
|
+
q = 1 - p
|
|
57
|
+
a = n * p + 0.5
|
|
58
|
+
c = tf.math.sqrt(2.0 * a * q * (1.0 - n / N))
|
|
59
|
+
k = tf.math.floor((n + 1) * (K + 1) / (N + 2))
|
|
60
|
+
g = _log_hypergeometric_coeff(k)
|
|
61
|
+
diff = tf.math.floor(a - c)
|
|
62
|
+
x = (a - diff - 1) / (a - diff)
|
|
63
|
+
diff = tf.where(
|
|
64
|
+
(n - diff) * (p - diff / N) * tf.math.square(x)
|
|
65
|
+
> (diff + 1.0) * (q - (n - diff - 1) / N),
|
|
66
|
+
diff + 1.0,
|
|
67
|
+
diff,
|
|
68
|
+
)
|
|
69
|
+
# TODO: Can we write this difference of lgammas more numerically stably?
|
|
70
|
+
h = (a - diff) * tf.math.exp(
|
|
71
|
+
0.5 * (g - _log_hypergeometric_coeff(diff)) + np.log(2.0)
|
|
72
|
+
)
|
|
73
|
+
b = tf.math.minimum(tf.math.minimum(n, K) + 1, tf.math.floor(a + 5 * c))
|
|
74
|
+
|
|
75
|
+
def generate_and_test_samples(seed):
|
|
76
|
+
v_seed, u_seed = samplers.split_seed(seed)
|
|
77
|
+
U = samplers.uniform(sample_shape, dtype=dtype, seed=u_seed)
|
|
78
|
+
V = samplers.uniform(sample_shape, dtype=dtype, seed=v_seed)
|
|
79
|
+
# Guard against 0.
|
|
80
|
+
|
|
81
|
+
X = a + h * (V - 0.5) / (1.0 - U)
|
|
82
|
+
samples = tf.math.floor(X)
|
|
83
|
+
good_sample_mask = (samples >= 0.0) & (samples < b)
|
|
84
|
+
T = g - _log_hypergeometric_coeff(samples)
|
|
85
|
+
# Uses slow pass since we are trying to do this in a vectorized way.
|
|
86
|
+
good_sample_mask = good_sample_mask & (2 * tf.math.log1p(-U) <= T)
|
|
87
|
+
return samples, good_sample_mask
|
|
88
|
+
|
|
89
|
+
samples = brs.batched_las_vegas_algorithm(
|
|
90
|
+
generate_and_test_samples, seed=seed
|
|
91
|
+
)[0]
|
|
92
|
+
samples = tf.where(good_params_mask, samples, np.nan)
|
|
93
|
+
# Now transform the samples depending on if we constrained N and / or k
|
|
94
|
+
samples = tf.where(
|
|
95
|
+
~is_k_small & ~is_n_small,
|
|
96
|
+
samples + previous_K + previous_n - N,
|
|
97
|
+
tf.where(
|
|
98
|
+
~is_k_small,
|
|
99
|
+
previous_n - samples,
|
|
100
|
+
tf.where(~is_n_small, previous_K - samples, samples),
|
|
101
|
+
),
|
|
102
|
+
)
|
|
103
|
+
return samples
|
|
@@ -0,0 +1,46 @@
|
|
|
1
|
+
"""Test the Hypergeometric random vaiable"""
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import tensorflow as tf
|
|
5
|
+
import tensorflow_probability as tfp
|
|
6
|
+
from tensorflow_probability.python.internal import test_util
|
|
7
|
+
|
|
8
|
+
from gemlib.distributions.hypergeometric import Hypergeometric
|
|
9
|
+
|
|
10
|
+
tfd = tfp.distributions
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@test_util.test_all_tf_execution_regimes
|
|
14
|
+
class TestHypergeometric(test_util.TestCase):
|
|
15
|
+
def setUp(self):
|
|
16
|
+
self._rng = np.random.RandomState(5)
|
|
17
|
+
super().setUp()
|
|
18
|
+
|
|
19
|
+
def test_neg_args(self):
|
|
20
|
+
"""Test for invalid arguments"""
|
|
21
|
+
with self.assertRaisesRegex(
|
|
22
|
+
tf.errors.InvalidArgumentError, "N must be non-negative"
|
|
23
|
+
):
|
|
24
|
+
self.evaluate(Hypergeometric(N=-3, K=1, n=2, validate_args=True))
|
|
25
|
+
|
|
26
|
+
def test_sample_n_float32(self):
|
|
27
|
+
"""Sample returning float32 args"""
|
|
28
|
+
|
|
29
|
+
X = Hypergeometric(345.0, 35.0, 100.0)
|
|
30
|
+
x = X.sample([1000, 1000], seed=1)
|
|
31
|
+
|
|
32
|
+
x = self.evaluate(x)
|
|
33
|
+
self.assertDTypeEqual(x, np.float32)
|
|
34
|
+
self.assertAllClose(tf.reduce_mean(x), X.mean(), atol=1e-3, rtol=1e-3)
|
|
35
|
+
|
|
36
|
+
def test_sample_n_float64(self):
|
|
37
|
+
"""Sample returning float32 args"""
|
|
38
|
+
|
|
39
|
+
X = Hypergeometric(
|
|
40
|
+
np.float64(345.0), np.float64(35.0), np.float64(100.0)
|
|
41
|
+
)
|
|
42
|
+
x = X.sample([1000, 1000], seed=1)
|
|
43
|
+
|
|
44
|
+
x = self.evaluate(x)
|
|
45
|
+
self.assertDTypeEqual(x, np.float64)
|
|
46
|
+
self.assertAllClose(tf.reduce_mean(x), X.mean(), atol=1e-3, rtol=1e-3)
|