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,21 @@
1
+ """DiscreteTimeStateTransitionModel-related MCMC samplers"""
2
+
3
+ from gemlib.mcmc.discrete_time_state_transition_model.left_censored_events_mh import ( # noqa: E501
4
+ UncalibratedLeftCensoredEventTimesUpdate,
5
+ )
6
+ from gemlib.mcmc.discrete_time_state_transition_model.move_events import (
7
+ UncalibratedEventTimesUpdate,
8
+ )
9
+ from gemlib.mcmc.discrete_time_state_transition_model.right_censored_events_mh import ( # noqa: E501
10
+ UncalibratedOccultUpdate,
11
+ )
12
+ from gemlib.mcmc.discrete_time_state_transition_model.util import (
13
+ TransitionTopology,
14
+ )
15
+
16
+ __all__ = [
17
+ "TransitionTopology",
18
+ "UncalibratedEventTimesUpdate",
19
+ "UncalibratedLeftCensoredEventTimesUpdate",
20
+ "UncalibratedOccultUpdate",
21
+ ]
@@ -0,0 +1,275 @@
1
+ """Mechanism for proposing event times to move"""
2
+
3
+ from warnings import warn
4
+
5
+ import numpy as np
6
+ import tensorflow as tf
7
+ import tensorflow_probability as tfp
8
+ from tensorflow_probability.python.mcmc.internal import util as mcmc_util
9
+
10
+ from gemlib.distributions import (
11
+ Categorical2,
12
+ UniformInteger,
13
+ UniformKCategorical,
14
+ )
15
+
16
+ tfd = tfp.distributions
17
+
18
+
19
+ def _events_or_inf(events, transition_id):
20
+ if transition_id is None:
21
+ return tf.fill(
22
+ events.shape[:-1], tf.constant(np.inf, dtype=events.dtype)
23
+ )
24
+ return tf.gather(events, transition_id, axis=-1)
25
+
26
+
27
+ def _abscumdiff(
28
+ events, initial_state, topology, t, delta_t, bound_times, int_dtype=tf.int32
29
+ ):
30
+ """Returns the number of free events to move in target_events
31
+ bounded by max([N_{target_id}(t)-N_{bound_id}(t)]_{bound_t}).
32
+
33
+ :param events: a [(M), T, X] tensor of transition events
34
+ :param initial_state: a [M, X] tensor of the constraining initial state
35
+ :param target_id: the Xth index of the target event
36
+ :param bound_t: the times to compute the constraints
37
+ :param bound_id: the Xth index of the bounding event, -1 implies no bound
38
+
39
+ :returns: a tensor of shape [M] + bound_t.shape[0] + of max free events,
40
+ dtype=target_events.dtype
41
+ """
42
+ with tf.name_scope("_abscumdiff"):
43
+ # This line prevents negative indices. However, we must have
44
+ # a contract that the output of the algorithm is invalid!
45
+ bound_times = tf.clip_by_value(
46
+ bound_times, clip_value_min=0, clip_value_max=events.shape[-2] - 1
47
+ )
48
+
49
+ # Maybe replace with pad to avoid unstack/stack
50
+ prev_events = _events_or_inf(events, topology.prev)
51
+ target_events = tf.gather(events, topology.target, axis=-1)
52
+ next_events = _events_or_inf(events, topology.next)
53
+ event_tensor = tf.stack(
54
+ [prev_events, target_events, next_events], axis=-1
55
+ )
56
+
57
+ # Compute the absolute cumulative difference between event times
58
+ diff = event_tensor[..., 1:] - event_tensor[..., :-1] # [m, T, 2]
59
+ cumdiff = tf.abs(tf.cumsum(diff, axis=-2)) # cumsum along time axis
60
+
61
+ # Create indices into cumdiff [m, d_max, 2]. Last dimension selects
62
+ # the bound for either the previous or next event.
63
+ indices = tf.stack(
64
+ [
65
+ tf.repeat(
66
+ tf.range(events.shape[0], dtype=int_dtype),
67
+ [bound_times.shape[1]],
68
+ ),
69
+ tf.reshape(bound_times, [-1]),
70
+ tf.repeat(tf.where(delta_t < 0, 0, 1), [bound_times.shape[1]]),
71
+ ],
72
+ axis=-1,
73
+ )
74
+ indices = tf.reshape(
75
+ indices, [events.shape[-3], bound_times.shape[1], 3]
76
+ )
77
+ free_events = tf.gather_nd(cumdiff, indices)
78
+
79
+ # Add on initial state
80
+ indices = tf.stack(
81
+ [
82
+ tf.range(events.shape[0]),
83
+ tf.where(
84
+ delta_t[:, 0] < 0, topology.target, topology.target + 1
85
+ ),
86
+ ],
87
+ axis=-1,
88
+ )
89
+ bound_init_state = tf.gather_nd(initial_state, indices)
90
+ free_events += bound_init_state[..., tf.newaxis]
91
+
92
+ return free_events
93
+
94
+
95
+ class Deterministic2(tfd.Deterministic):
96
+ def __init__(
97
+ self,
98
+ loc,
99
+ atol=None,
100
+ rtol=None,
101
+ validate_args=False,
102
+ allow_nan_stats=True,
103
+ log_prob_dtype=tf.float32,
104
+ name="Deterministic",
105
+ ):
106
+ parameters = dict(locals())
107
+ super().__init__(
108
+ loc,
109
+ atol=atol,
110
+ rtol=rtol,
111
+ validate_args=validate_args,
112
+ allow_nan_stats=allow_nan_stats,
113
+ name=name,
114
+ )
115
+ self.log_prob_dtype = log_prob_dtype
116
+
117
+ def _prob(self, x):
118
+ return tf.constant(1, dtype=self.log_prob_dtype)
119
+
120
+
121
+ def event_time_proposal(
122
+ events, initial_state, topology, d_max, n_max, dtype=tf.int32, name=None
123
+ ):
124
+ """Draws an event time move proposal.
125
+ :param events: a [M, T, K] tensor of event times (M number of
126
+ metapopulations, T number of full_timepoints, K number of
127
+ transitions)
128
+ :param initial_state: a [M, S] tensor of initial metapopulation x state
129
+ counts
130
+ :param topology: a 3-element tuple of (previous_transition,
131
+ target_transition, next_transition), eg "(s->e, e->i,
132
+ i->r)" (assuming we are interested presently in e->i,
133
+ `None` for boundaries)
134
+ :param d_max: the maximum distance over which to move (in time)
135
+ :param n_max: the maximum number of events to move
136
+ """
137
+ target_events = tf.gather(events, topology.target, axis=-1)
138
+ time_interval = tf.range(d_max, dtype=dtype)
139
+
140
+ def t():
141
+ with tf.name_scope("t"):
142
+ # Waiting for fixed tf.nn.sparse_softmax_cross_entropy_with_logits
143
+ x = tf.cast(target_events > 0, dtype=events.dtype) # [M, T]
144
+ return Categorical2(logits=tf.math.log(x), name="event_coords")
145
+
146
+ def delta_t(t):
147
+ with tf.name_scope("delta_t"):
148
+ d_max_bcast = tf.broadcast_to(d_max, [events.shape[-3]])
149
+ low = -tf.clip_by_value(
150
+ d_max_bcast, clip_value_min=0, clip_value_max=t
151
+ )
152
+ high = tf.clip_by_value(
153
+ d_max_bcast,
154
+ clip_value_min=0,
155
+ clip_value_max=events.shape[-2] - t - 1,
156
+ )
157
+ return UniformInteger(
158
+ low=low, high=high + 1, float_dtype=events.dtype
159
+ )
160
+
161
+ def x_star(t, delta_t):
162
+ with tf.name_scope("x_star"):
163
+ # Compute bounds
164
+ # The limitations of XLA mean that we must calculate bounds for
165
+ # intervals [t, t+delta_t) if delta_t > 0, and [t+delta_t, t) if
166
+ # delta_t is < 0.
167
+ t = t[..., tf.newaxis]
168
+ delta_t = delta_t[..., tf.newaxis]
169
+ bound_times = tf.where(
170
+ delta_t < 0,
171
+ t - time_interval - 1,
172
+ t + time_interval, # [t+delta_t, t)
173
+ ) # [t, t+delta_t)
174
+ free_events = _abscumdiff(
175
+ events=events,
176
+ initial_state=initial_state,
177
+ topology=topology,
178
+ t=t,
179
+ delta_t=delta_t,
180
+ bound_times=bound_times,
181
+ int_dtype=dtype,
182
+ )
183
+
184
+ # Mask out bits of the interval we don't need for our delta_t
185
+ inf_mask = tf.cumsum(
186
+ tf.one_hot(
187
+ tf.math.abs(delta_t[:, 0]),
188
+ d_max,
189
+ on_value=tf.constant(np.inf, events.dtype),
190
+ dtype=events.dtype,
191
+ )
192
+ )
193
+ free_events = tf.maximum(inf_mask, free_events)
194
+ free_events = tf.reduce_min(free_events, axis=-1)
195
+
196
+ indices = tf.stack(
197
+ [tf.range(events.shape[0], dtype=dtype), t[:, 0]], axis=-1
198
+ )
199
+ available_events = tf.gather_nd(target_events, indices)
200
+ max_events = tf.minimum(free_events, available_events)
201
+ max_events = tf.clip_by_value(
202
+ max_events, clip_value_min=0, clip_value_max=n_max
203
+ )
204
+ # Draw x_star
205
+ return UniformInteger(
206
+ low=1, high=max_events + 1, float_dtype=events.dtype
207
+ )
208
+
209
+ return tfd.JointDistributionNamed(
210
+ {"t": t, "delta_t": delta_t, "x_star": x_star}, name=name
211
+ )
212
+
213
+
214
+ def filtered_event_time_proposal( # pylint: disable-invalid-name
215
+ events,
216
+ initial_state,
217
+ topology,
218
+ m_max,
219
+ d_max,
220
+ n_max,
221
+ dtype=tf.int32,
222
+ name=None,
223
+ ):
224
+ """FilteredEventTimeProposal allows us to choose a subset of indices
225
+ in `range(events.shape[0])` for which to propose an update. The
226
+ results are then broadcast back to `events.shape[0]`.
227
+
228
+ :param events: a [M, T, X] event tensor
229
+ :param initial_state: a [M, S] initial state tensor
230
+ :param topology: a TransitionTopology named tuple describing the ordering
231
+ of events
232
+ :param m: the number of metapopulations to move
233
+ :param d_max: maximum distance in time to move
234
+ :param n_max: maximum number of events to move (user defined)
235
+ :return: an instance of a JointDistributionNamed
236
+ """
237
+ if mcmc_util.is_list_like(events):
238
+ warn(
239
+ "Batched FilteredEventTimeProposals are not yet supported",
240
+ stacklevel=1,
241
+ )
242
+ events = events[0]
243
+
244
+ target_events = tf.gather(events, topology.target, axis=-1)
245
+
246
+ def m():
247
+ with tf.name_scope("m"):
248
+ hot_meta = tf.math.count_nonzero(target_events, axis=1) > 0
249
+ X = UniformKCategorical(
250
+ m_max, hot_meta, float_dtype=events.dtype, name="m"
251
+ )
252
+ return X
253
+
254
+ def move(m):
255
+ """We select out meta-population `m` from the first
256
+ dimension of `events`.
257
+ :param m: a 1-D tensor of indices of meta-populations
258
+ :return: a random variable of type `EventTimeProposal`
259
+ """
260
+ with tf.name_scope("move"):
261
+ select_meta = tf.gather(events, m, axis=0)
262
+ select_init = tf.gather(initial_state, m, axis=0)
263
+ return event_time_proposal(
264
+ select_meta,
265
+ select_init,
266
+ topology,
267
+ d_max,
268
+ n_max,
269
+ dtype=dtype,
270
+ name=name,
271
+ )
272
+
273
+ return tfd.JointDistributionNamed(
274
+ {"m": m, "move": move}, name="FilteredEventTimeProposal"
275
+ )
@@ -0,0 +1,56 @@
1
+ """Test fixtures"""
2
+
3
+ import numpy as np
4
+ import pytest
5
+
6
+
7
+ @pytest.fixture(scope="module")
8
+ def sir_metapop_example():
9
+ """Outcome of a simulation from a 3-metapopulation model
10
+ with mixing, implemented in https://colab.research.google.com/drive/1Q1PUcOnYlvCGHzRUBUAp4CxYhZ8RJzg8?usp=sharing
11
+ """
12
+
13
+ incidence_matrix = np.array([[-1, 0], [1, -1], [0, 1]], dtype=np.float32)
14
+
15
+ initial_conditions = np.array(
16
+ [[999, 50, 0], [500, 20, 0], [250, 10, 0]], dtype=np.float32
17
+ )
18
+
19
+ events = np.array(
20
+ [
21
+ [
22
+ [11.0, 11.0],
23
+ [11.0, 6.0],
24
+ [2.0, 6.0],
25
+ [10.0, 6.0],
26
+ [12.0, 7.0],
27
+ [11.0, 5.0],
28
+ [13.0, 10.0],
29
+ ],
30
+ [
31
+ [5.0, 4.0],
32
+ [5.0, 2.0],
33
+ [11.0, 1.0],
34
+ [8.0, 4.0],
35
+ [7.0, 4.0],
36
+ [5.0, 5.0],
37
+ [12.0, 6.0],
38
+ ],
39
+ [
40
+ [2.0, 2.0],
41
+ [4.0, 2.0],
42
+ [1.0, 0.0],
43
+ [2.0, 2.0],
44
+ [6.0, 1.0],
45
+ [4.0, 1.0],
46
+ [6.0, 1.0],
47
+ ],
48
+ ],
49
+ dtype=np.float32,
50
+ )
51
+
52
+ return {
53
+ "initial_conditions": initial_conditions,
54
+ "events": events,
55
+ "incidence_matrix": incidence_matrix,
56
+ }
@@ -0,0 +1,258 @@
1
+ """Metropolis Hastings implementation for left-censored events
2
+ in a discrete-time metapopulation epidemic model
3
+ """
4
+
5
+ from typing import NamedTuple, Tuple
6
+
7
+ import tensorflow as tf
8
+ import tensorflow_probability as tfp
9
+ from tensorflow_probability.python.internal import prefer_static as ps
10
+ from tensorflow_probability.python.internal import samplers
11
+ from tensorflow_probability.python.mcmc.internal import util as mcmc_util
12
+
13
+ from .left_censored_events_proposal import (
14
+ left_censored_event_time_proposal,
15
+ )
16
+
17
+ tfd = tfp.distributions
18
+
19
+ __all__ = ["UncalibratedLeftCensoredEventTimesUpdate"]
20
+
21
+ LEFT_CENSORED_EVENT_TIMES_UPDATE_TUPLE_LEN = 2
22
+
23
+
24
+ def _update_state(update, current_state, transition_idx, incidence_matrix):
25
+ current_initial_conditions, current_events = current_state
26
+ update = {k: tf.convert_to_tensor(v) for k, v in update.items()}
27
+
28
+ # -1 is moving events forward in time
29
+ sign = tf.gather(
30
+ [-1, 1],
31
+ update["direction"],
32
+ )
33
+ events_delta = update["num_events"] * sign
34
+
35
+ # Update initial conditions
36
+ indices = tf.stack(
37
+ [
38
+ tf.broadcast_to(
39
+ update["unit"], [current_initial_conditions.shape[-1]]
40
+ ),
41
+ ps.range(current_initial_conditions.shape[-1]),
42
+ ],
43
+ axis=-1,
44
+ )
45
+ new_initial_conditions = tf.tensor_scatter_nd_add(
46
+ current_initial_conditions,
47
+ indices=indices,
48
+ updates=tf.cast(
49
+ tf.cast(
50
+ events_delta, incidence_matrix.dtype
51
+ ) # TODO sort this out: casts are usually a code smell!
52
+ * ps.gather(incidence_matrix, transition_idx, axis=-1),
53
+ current_initial_conditions.dtype,
54
+ ),
55
+ )
56
+
57
+ # Update events
58
+ indices = tf.stack(
59
+ [
60
+ update["unit"],
61
+ update["timepoint"],
62
+ ps.broadcast_to(transition_idx, update["unit"].shape),
63
+ ],
64
+ axis=-1,
65
+ )
66
+ new_events = tf.tensor_scatter_nd_sub(
67
+ current_events,
68
+ indices=indices,
69
+ updates=tf.cast(events_delta, dtype=current_events.dtype),
70
+ )
71
+
72
+ return new_initial_conditions, new_events
73
+
74
+
75
+ def _reverse_update(update):
76
+ direction = (update["direction"] + 1) % 2
77
+ return {
78
+ "unit": update["unit"],
79
+ "timepoint": update["timepoint"],
80
+ "direction": direction,
81
+ "num_events": update["num_events"],
82
+ }
83
+
84
+
85
+ class LeftCensoredEventTimeResults(NamedTuple):
86
+ log_acceptance_correction: float
87
+ target_log_prob: float
88
+ unit: int
89
+ timepoint: int
90
+ direction: int
91
+ num_events: int
92
+ seed: Tuple[int, int]
93
+
94
+
95
+ class UncalibratedLeftCensoredEventTimesUpdate(tfp.mcmc.TransitionKernel):
96
+ """UncalibratedLeftCensoredEventTimesUpdate"""
97
+
98
+ def __init__(
99
+ self,
100
+ target_log_prob_fn,
101
+ transition_index,
102
+ incidence_matrix,
103
+ max_timepoint,
104
+ max_events,
105
+ name=None,
106
+ ):
107
+ """An uncalibrated random walk for initial conditions.
108
+ :param target_log_prob_fn: the log density of the target distribution
109
+ :param transition_index: the index of the transition to adjust
110
+ :param incidence_matrix: the `[S,R]` incidence matrix
111
+ :param max_timepoint: max timepoint up to which to move events
112
+ :param max_events: max number of events per unit/timepoint to move
113
+ """
114
+ self._name = name
115
+ self._parameters = {
116
+ "target_log_prob_fn": target_log_prob_fn,
117
+ "transition_index": transition_index,
118
+ "incidence_matrix": incidence_matrix,
119
+ "max_timepoint": max_timepoint,
120
+ "max_events": max_events,
121
+ "name": name,
122
+ }
123
+
124
+ @property
125
+ def target_log_prob_fn(self):
126
+ return self._parameters["target_log_prob_fn"]
127
+
128
+ @property
129
+ def transition_index(self):
130
+ return self._parameters["transition_index"]
131
+
132
+ @property
133
+ def incidence_matrix(self):
134
+ return self._parameters["incidence_matrix"]
135
+
136
+ @property
137
+ def max_timepoint(self):
138
+ return self._parameters["max_timepoint"]
139
+
140
+ @property
141
+ def max_events(self):
142
+ return self._parameters["max_events"]
143
+
144
+ @property
145
+ def name(self):
146
+ return self._parameters["name"]
147
+
148
+ @property
149
+ def parameters(self):
150
+ return self._parameters
151
+
152
+ @property
153
+ def is_calibrated(self):
154
+ return False
155
+
156
+ def one_step(self, current_state, previous_kernel_results, seed=None):
157
+ """Update the initial conditions
158
+
159
+ :param current_state: a tuple of `(current_initial_conditions,
160
+ current_events)`
161
+ :param previous_kernel_results: previous kernel results tuple
162
+ :param seed: optional seed tuple `(int32, int32)`
163
+ :returns: new state tuple `(next_initial_conditions, next_events)`
164
+ """
165
+ with tf.name_scope("uncalibrated_left_censored_events_mh/one_step"):
166
+ seed = samplers.sanitize_seed(
167
+ seed, salt="uncalibrated_left_censored_events_mh"
168
+ )
169
+
170
+ if (not mcmc_util.is_list_like(current_state)) and (
171
+ len(current_state) == LEFT_CENSORED_EVENT_TIMES_UPDATE_TUPLE_LEN
172
+ ):
173
+ raise ValueError(
174
+ f"State for LeftCensoredEventTimesUpdate must be a\
175
+ list/tuple of length {LEFT_CENSORED_EVENT_TIMES_UPDATE_TUPLE_LEN}"
176
+ )
177
+
178
+ proposal = left_censored_event_time_proposal(
179
+ events=current_state[1],
180
+ initial_state=current_state[0],
181
+ transition=self.transition_index,
182
+ incidence_matrix=self.incidence_matrix,
183
+ num_units=1,
184
+ max_timepoint=self.max_timepoint,
185
+ max_events=self.max_events,
186
+ name=f"{self.name}/fwd_proposal",
187
+ )
188
+ fwd_update = proposal.sample(seed=seed)
189
+ fwd_proposal_log_prob = proposal.log_prob(
190
+ fwd_update, name="fwd_proposal_log_prob"
191
+ )
192
+
193
+ next_state = _update_state(
194
+ fwd_update,
195
+ current_state,
196
+ self.transition_index,
197
+ self.incidence_matrix,
198
+ )
199
+ next_target_log_prob = self.target_log_prob_fn(*next_state)
200
+
201
+ rev_update = _reverse_update(fwd_update)
202
+ rev_proposal = left_censored_event_time_proposal(
203
+ events=next_state[1],
204
+ initial_state=next_state[0],
205
+ transition=self.transition_index,
206
+ incidence_matrix=self.incidence_matrix,
207
+ num_units=1,
208
+ max_timepoint=self.max_timepoint,
209
+ max_events=self.max_events,
210
+ name=f"{self.name}/rev_proposal",
211
+ )
212
+ rev_proposal_log_prob = rev_proposal.log_prob(rev_update)
213
+ log_acceptance_correction = tf.reduce_sum(
214
+ rev_proposal_log_prob - fwd_proposal_log_prob
215
+ )
216
+ results = (
217
+ next_state,
218
+ LeftCensoredEventTimeResults(
219
+ log_acceptance_correction=log_acceptance_correction,
220
+ target_log_prob=next_target_log_prob,
221
+ unit=fwd_update["unit"],
222
+ timepoint=fwd_update["timepoint"],
223
+ direction=fwd_update["direction"],
224
+ num_events=fwd_update["num_events"],
225
+ seed=seed,
226
+ ),
227
+ )
228
+
229
+ return results
230
+
231
+ def bootstrap_results(self, init_state):
232
+ with tf.name_scope(
233
+ "uncalibrated_left_censored_events_mh/boostrap_results"
234
+ ):
235
+ if (not mcmc_util.is_list_like(init_state)) and (
236
+ len(init_state) == 2 # noqa: PLR2004
237
+ ):
238
+ raise ValueError(
239
+ f"State for LeftCensoredEventTimesUpdate must be a \
240
+ list/tuple of length {LEFT_CENSORED_EVENT_TIMES_UPDATE_TUPLE_LEN}"
241
+ )
242
+
243
+ initial_conditions = tf.convert_to_tensor(init_state[0])
244
+ events = tf.convert_to_tensor(init_state[1])
245
+ init_target_log_prob = self.target_log_prob_fn(
246
+ initial_conditions, events
247
+ )
248
+ return LeftCensoredEventTimeResults(
249
+ log_acceptance_correction=tf.constant(
250
+ 0.0, dtype=init_target_log_prob.dtype
251
+ ),
252
+ target_log_prob=init_target_log_prob,
253
+ unit=tf.zeros([1], dtype=tf.int32),
254
+ timepoint=tf.zeros([1], dtype=tf.int32),
255
+ direction=tf.constant(0, dtype=tf.int32),
256
+ num_events=tf.ones([1], dtype=tf.int32),
257
+ seed=samplers.zeros_seed(),
258
+ )