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.
Files changed (75) hide show
  1. gemlib/__init__.py +9 -0
  2. gemlib/distributions/__init__.py +26 -0
  3. gemlib/distributions/brownian.py +141 -0
  4. gemlib/distributions/categorical2.py +33 -0
  5. gemlib/distributions/continuous_markov.py +371 -0
  6. gemlib/distributions/continuous_time_state_transition_model.py +185 -0
  7. gemlib/distributions/continuous_time_state_transition_model_test.py +289 -0
  8. gemlib/distributions/discrete_markov.py +279 -0
  9. gemlib/distributions/discrete_rejection_sampling.py +149 -0
  10. gemlib/distributions/discrete_time_state_transition_model.py +324 -0
  11. gemlib/distributions/discrete_time_state_transition_model_examples.py +453 -0
  12. gemlib/distributions/discrete_time_state_transition_model_test.py +336 -0
  13. gemlib/distributions/experimental/__init__.py +7 -0
  14. gemlib/distributions/experimental/discrete_approx_cont_state_transition_model.py +194 -0
  15. gemlib/distributions/experimental/state_transition_marginal_model.py +352 -0
  16. gemlib/distributions/hypergeometric.py +144 -0
  17. gemlib/distributions/hypergeometric_sampler.py +103 -0
  18. gemlib/distributions/hypergeometric_test.py +46 -0
  19. gemlib/distributions/kcategorical.py +113 -0
  20. gemlib/distributions/kcategorical_test.py +52 -0
  21. gemlib/distributions/uniform_integer.py +169 -0
  22. gemlib/distributions/uniform_integer_test.py +55 -0
  23. gemlib/mcmc/__init__.py +23 -0
  24. gemlib/mcmc/adaptive_random_walk_metropolis.py +859 -0
  25. gemlib/mcmc/adaptive_random_walk_metropolis_test.py +144 -0
  26. gemlib/mcmc/bb_fixture.pkl +0 -0
  27. gemlib/mcmc/brownian_bridge_kernel.py +291 -0
  28. gemlib/mcmc/brownian_bridge_kernel_test.py +164 -0
  29. gemlib/mcmc/chain_binomial_rippler.py +524 -0
  30. gemlib/mcmc/chain_binomial_rippler_test.py +120 -0
  31. gemlib/mcmc/compound_kernel.py +156 -0
  32. gemlib/mcmc/conftest.py +4 -0
  33. gemlib/mcmc/damped_chain_binomial_rippler.py +840 -0
  34. gemlib/mcmc/discrete_time_state_transition_model/__init__.py +21 -0
  35. gemlib/mcmc/discrete_time_state_transition_model/event_time_proposal.py +275 -0
  36. gemlib/mcmc/discrete_time_state_transition_model/fixtures.py +56 -0
  37. gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh.py +258 -0
  38. gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_mh_test.py +108 -0
  39. gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal.py +188 -0
  40. gemlib/mcmc/discrete_time_state_transition_model/left_censored_events_proposal_test.py +33 -0
  41. gemlib/mcmc/discrete_time_state_transition_model/move_events.py +239 -0
  42. gemlib/mcmc/discrete_time_state_transition_model/move_events_test.py +63 -0
  43. gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh.py +254 -0
  44. gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_mh_test.py +28 -0
  45. gemlib/mcmc/discrete_time_state_transition_model/right_censored_events_proposal.py +150 -0
  46. gemlib/mcmc/discrete_time_state_transition_model/util.py +9 -0
  47. gemlib/mcmc/experimental/__init__.py +0 -0
  48. gemlib/mcmc/experimental/composable_kernel.py +270 -0
  49. gemlib/mcmc/experimental/discrete_time_state_transition_model/__init__.py +1 -0
  50. gemlib/mcmc/experimental/discrete_time_state_transition_model/left_censored_events_mh.py +149 -0
  51. gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events.py +71 -0
  52. gemlib/mcmc/experimental/discrete_time_state_transition_model/move_events_test.py +48 -0
  53. gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh.py +89 -0
  54. gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh_test.py +28 -0
  55. gemlib/mcmc/experimental/hmc.py +111 -0
  56. gemlib/mcmc/experimental/hmc_test.py +61 -0
  57. gemlib/mcmc/experimental/mcmc_base.py +33 -0
  58. gemlib/mcmc/experimental/mcmc_sampler.py +94 -0
  59. gemlib/mcmc/experimental/mcmc_sampler_test.py +44 -0
  60. gemlib/mcmc/experimental/multi_scan.py +53 -0
  61. gemlib/mcmc/experimental/multi_scan_test.py +71 -0
  62. gemlib/mcmc/experimental/random_walk_metropolis.py +107 -0
  63. gemlib/mcmc/experimental/random_walk_metropolis_test.py +214 -0
  64. gemlib/mcmc/experimental/test_util.py +49 -0
  65. gemlib/mcmc/gibbs_kernel.py +505 -0
  66. gemlib/mcmc/gibbs_kernel_test.py +212 -0
  67. gemlib/mcmc/h5_posterior.py +77 -0
  68. gemlib/mcmc/multi_scan_kernel.py +59 -0
  69. gemlib/mcmc/zarr_posterior.py +132 -0
  70. gemlib/util.py +117 -0
  71. gemlib/util_test.py +75 -0
  72. gemlib-0.9.2.dist-info/LICENSE +21 -0
  73. gemlib-0.9.2.dist-info/METADATA +19 -0
  74. gemlib-0.9.2.dist-info/RECORD +75 -0
  75. 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)