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,270 @@
1
+ """Implementation of Metropolis-within-Gibbs framework"""
2
+
3
+ from collections import ChainMap, namedtuple
4
+ from typing import AnyStr, Callable, List, Optional, Tuple
5
+
6
+ import tensorflow_probability as tfp
7
+
8
+ from .mcmc_base import (
9
+ ChainState,
10
+ KernelInfo,
11
+ KernelState,
12
+ Position,
13
+ SamplingAlgorithm,
14
+ )
15
+
16
+ split_seed = tfp.random.split_seed
17
+
18
+ __all__ = ["Step"]
19
+
20
+
21
+ def _maybe_list(x):
22
+ if isinstance(x, list):
23
+ return x
24
+ return [x]
25
+
26
+
27
+ def _project_position(
28
+ position: Position, varnames: List[AnyStr]
29
+ ) -> Tuple[Position, Position]:
30
+ """Splits `position` into `position[varnames]` and
31
+ `position[~varnames]`
32
+ """
33
+ if varnames is None:
34
+ return (position, ())
35
+
36
+ target = {k: v for k, v in position._asdict().items() if k in varnames}
37
+ target_compl = {
38
+ k: v for k, v in position._asdict().items() if k not in varnames
39
+ }
40
+
41
+ return (
42
+ namedtuple("Target", target.keys())(**target),
43
+ namedtuple("TargetCompl", target_compl.keys())(**target_compl),
44
+ )
45
+
46
+
47
+ def _join_dicts(a: dict, b: dict):
48
+ """Joins two dictionaries `a` and `b`"""
49
+ return dict(ChainMap(a, b))
50
+
51
+
52
+ def _maybe_flatten(x: List):
53
+ """Flatten a list if `len(x) <= 1`"""
54
+ if len(x) == 0:
55
+ return None
56
+ if len(x) == 1:
57
+ return x[0]
58
+ return x
59
+
60
+
61
+ class KernelInitMonad:
62
+ """KernelInitMonad is a Writer monad allowing us to build an initial
63
+ state tuple for a Metropolis-within-Gibbs algorithm
64
+ """
65
+
66
+ def __init__(self, fn: SamplingAlgorithm):
67
+ """The monad 'unit' function"""
68
+ self.__fn = fn
69
+
70
+ def __call__(self, *args, **kwargs):
71
+ """Monad ``run'' function"""
72
+ return self.__fn(*args, **kwargs)
73
+
74
+ def fish(self, next_kernel):
75
+ """Monad combination, i.e. Haskell fish operator"""
76
+
77
+ @KernelInitMonad
78
+ def compound_init_fn(
79
+ target_log_prob_fn: Callable[[Position], float],
80
+ initial_position: ChainState,
81
+ ):
82
+ _, self_kernel_state = self(target_log_prob_fn, initial_position)
83
+ next_chain_state, next_kernel_state = next_kernel(
84
+ target_log_prob_fn, initial_position
85
+ )
86
+
87
+ return (
88
+ next_chain_state,
89
+ _maybe_list(self_kernel_state) + _maybe_list(next_kernel_state),
90
+ )
91
+
92
+ return compound_init_fn
93
+
94
+ def __rshift__(self, next_kernel):
95
+ return self.fish(next_kernel)
96
+
97
+
98
+ class KernelStepMonad:
99
+ """StepMonad is a state monad that allows us to chain MCMC kernels
100
+ together.
101
+ """
102
+
103
+ def __init__(self, fn: SamplingAlgorithm):
104
+ """The monad 'unit' function"""
105
+ self.__fn = fn # Make private
106
+
107
+ def __call__(self, *args, **kwargs):
108
+ """Apply the state transformer computation to a state."""
109
+ return self.__fn(*args, **kwargs)
110
+
111
+ def __rshift__(self, next_kernel_fn):
112
+ """The monad 'bind' operator which allows chaining.
113
+ ma >> f :: ma -> (a -> mb) -> mb
114
+ """
115
+
116
+ def compound_step_kernel(
117
+ target_log_prob_fn: Callable[[Position], float],
118
+ chain_and_kernel_state: Tuple[ChainState, KernelState],
119
+ seed,
120
+ ):
121
+ first_seed, second_seed = split_seed(seed)
122
+ # Pre-order recursive descent
123
+ chain_state, kernel_state = chain_and_kernel_state
124
+
125
+ self_kernel_state = kernel_state[:-1]
126
+ next_kernel_state = kernel_state[-1]
127
+
128
+ # Execute self kernel step
129
+ (chain_state, self_kernel_state), self_info = self.__fn(
130
+ target_log_prob_fn,
131
+ (chain_state, _maybe_flatten(self_kernel_state)),
132
+ seed=first_seed,
133
+ )
134
+
135
+ # Descend right hand branch
136
+ (chain_state, next_kernel_state), next_info = next_kernel_fn(
137
+ target_log_prob_fn,
138
+ (chain_state, next_kernel_state),
139
+ seed=second_seed,
140
+ )
141
+
142
+ return (
143
+ (
144
+ chain_state,
145
+ _maybe_list(self_kernel_state)
146
+ + _maybe_list(next_kernel_state),
147
+ ),
148
+ self_info + next_info,
149
+ )
150
+
151
+ return KernelStepMonad(compound_step_kernel)
152
+
153
+
154
+ class ComposableSamplingAlgorithm:
155
+ """ComposableSamplingAlgorithm implements a Metropolis-within-Gibbs step"""
156
+
157
+ def __init__(self, init: KernelInitMonad, step: KernelStepMonad):
158
+ self.__initialize = init
159
+ self.__step = step
160
+
161
+ @property
162
+ def init(self):
163
+ """Initialize an MCMC chain"""
164
+ return self.__initialize
165
+
166
+ @property
167
+ def step(self):
168
+ """Function to invoke the MCMC kernel"""
169
+ return self.__step
170
+
171
+ def __rshift__(self, rhs):
172
+ """Combinator"""
173
+ return ComposableSamplingAlgorithm(
174
+ init=self.init >> rhs.init, step=self.step >> rhs.step
175
+ )
176
+
177
+
178
+ class Step: # pylint: disable=too-few-public-methods
179
+ """A Metropolis-within-Gibbs step"""
180
+
181
+ def __new__(
182
+ cls,
183
+ sampling_algorithm: SamplingAlgorithm,
184
+ target_names: Optional[List[str]] = None,
185
+ step_kwargs_fn: Callable[[Position], dict] = lambda _: {},
186
+ ):
187
+ """Transforms a base kernel to operate on a substate of a Markov chain.
188
+
189
+ Args:
190
+ ----
191
+ sampling_algorithm: a named tuple containing the generic kernel `init`
192
+ and `step` function.
193
+ target_log_prob_fn: a list of variable names on which the
194
+ Metropolis-within-Gibbs step is to operate
195
+ step_kwargs_fn: a callable taking the chain position as an argument,
196
+ and returning a dictionary of extra kwargs to
197
+ `sampling_algorithm.step`.
198
+
199
+ Returns:
200
+ -------
201
+ A monad of type StepMonad (State -> (State, Info))
202
+
203
+ """
204
+
205
+ @KernelInitMonad
206
+ def init(
207
+ target_log_prob_fn: Callable[[Position], float],
208
+ initial_position: Position,
209
+ ):
210
+ target, target_compl = _project_position(
211
+ initial_position, target_names
212
+ )
213
+
214
+ def conditional_tlp(*args):
215
+ state = _join_dicts(
216
+ dict(zip(target._fields, args)),
217
+ target_compl._asdict(),
218
+ )
219
+ return target_log_prob_fn(**state)
220
+
221
+ kernel_state = sampling_algorithm.init(conditional_tlp, target)
222
+
223
+ chain_state = ChainState(
224
+ position=initial_position,
225
+ log_density=kernel_state[0].log_density,
226
+ log_density_grad=kernel_state[0].log_density_grad,
227
+ )
228
+
229
+ return chain_state, kernel_state[1]
230
+
231
+ @KernelStepMonad
232
+ def step(
233
+ target_log_prob_fn: Callable[[Position], float],
234
+ chain_and_kernel_state: Tuple[ChainState, KernelState],
235
+ seed,
236
+ ) -> Tuple[Tuple[ChainState, KernelState], KernelInfo]:
237
+ chain_state, kernel_state = chain_and_kernel_state
238
+
239
+ # Split global state and generate conditional density
240
+ target, target_compl = _project_position(
241
+ chain_state.position, target_names
242
+ )
243
+
244
+ # Calculate the conditional log density
245
+ def conditional_tlp(*args):
246
+ state = _join_dicts(
247
+ dict(zip(target._fields, args)),
248
+ target_compl._asdict(),
249
+ )
250
+ return target_log_prob_fn(**state)
251
+
252
+ # Invoke the kernel on the target state
253
+ (new_target, new_kernel_state), info = sampling_algorithm.step(
254
+ conditional_tlp,
255
+ (chain_state._replace(position=target), kernel_state),
256
+ seed,
257
+ **step_kwargs_fn(chain_state.position),
258
+ )
259
+
260
+ # Stitch the global position back together
261
+ new_global_position = chain_state.position.__class__(
262
+ **new_target.position._asdict(), **target_compl._asdict()
263
+ )
264
+ new_global_state = new_target._replace(
265
+ position=new_global_position,
266
+ )
267
+
268
+ return (new_global_state, new_kernel_state), info
269
+
270
+ return ComposableSamplingAlgorithm(init, step)
@@ -0,0 +1 @@
1
+ """DiscreteTimeStateTransitionModel-related MCMC samplers"""
@@ -0,0 +1,149 @@
1
+ """Left-censored events MCMC kernel for DiscreteTimeStateTransitionModel"""
2
+
3
+ from functools import partial
4
+ from typing import NamedTuple, Optional, Tuple
5
+
6
+ import tensorflow as tf
7
+ import tensorflow_probability as tfp
8
+
9
+ from gemlib.mcmc.discrete_time import UncalibratedLeftCensoredEventTimesUpdate
10
+ from gemlib.mcmc.experimental.mcmc_base import ChainState, SamplingAlgorithm
11
+
12
+
13
+ class LeftCensoredEventsState(NamedTuple):
14
+ max_timepoint: int
15
+ max_events: int
16
+
17
+
18
+ class LeftCensoredEventsInfo(NamedTuple):
19
+ is_accepted: bool
20
+ log_acceptance_correction: float
21
+ target_log_prob: float
22
+ unit: int
23
+ timepoint: int
24
+ direction: int
25
+ num_events: int
26
+ seed: Tuple[int, int]
27
+
28
+
29
+ class LeftCensoredEventsPosition(NamedTuple):
30
+ initial_conditions: tf.Tensor
31
+ events: tf.Tensor
32
+
33
+
34
+ def _get_state_tuple(
35
+ initial_conditions_varname: str,
36
+ events_varname: str,
37
+ position: NamedTuple,
38
+ ):
39
+ return LeftCensoredEventsPosition(
40
+ getattr(position, initial_conditions_varname),
41
+ getattr(position, events_varname),
42
+ )
43
+
44
+
45
+ def _repack_state_tuple(
46
+ initial_conditions_varname: str,
47
+ events_varname: str,
48
+ new_structure: LeftCensoredEventsPosition,
49
+ original_structure: NamedTuple,
50
+ ):
51
+ return original_structure.__class__(
52
+ **{
53
+ initial_conditions_varname: new_structure.initial_conditions,
54
+ events_varname: new_structure.events,
55
+ }
56
+ )
57
+
58
+
59
+ def left_censored_events_mh(
60
+ transition_index: int,
61
+ max_timepoint: int,
62
+ max_events: int,
63
+ stoichiometry: tf.Tensor,
64
+ events_varname: str = "events",
65
+ initial_conditions_varname: str = "initial_conditions",
66
+ name: Optional[str] = None,
67
+ ):
68
+ """Left censored event times Metropolis Hastings
69
+
70
+ This MCMC kernel updates left-censored event times in a
71
+ DiscreteTimeStateTransitionModel realisation.
72
+
73
+ Args
74
+ ----
75
+ transition_index: the index of the transition to update
76
+ max_timepoint: max timepoint up to which to propose moves
77
+ max_events: max number of events per unit/timepoint to move
78
+
79
+ Returns
80
+ -------
81
+ A instance of SamplingAlgorithm
82
+ """
83
+
84
+ flatten_state = partial(
85
+ _get_state_tuple, initial_conditions_varname, events_varname
86
+ )
87
+ repack_state = partial(
88
+ _repack_state_tuple, initial_conditions_varname, events_varname
89
+ )
90
+
91
+ def _build_kernel(target_log_prob_fn):
92
+ return tfp.mcmc.MetropolisHastings(
93
+ inner_kernel=UncalibratedLeftCensoredEventTimesUpdate(
94
+ target_log_prob_fn,
95
+ transition_index,
96
+ stoichiometry,
97
+ max_timepoint,
98
+ max_events,
99
+ name,
100
+ )
101
+ )
102
+
103
+ def init_fn(target_log_prob_fn, target_state):
104
+ kernel = _build_kernel(target_log_prob_fn)
105
+ results = kernel.bootstrap_results(flatten_state(target_state))
106
+ chain_state = ChainState(
107
+ position=target_state,
108
+ log_density=results.accepted_results.target_log_prob,
109
+ )
110
+ kernel_state = LeftCensoredEventsState(
111
+ max_timepoint=max_timepoint, max_events=max_events
112
+ )
113
+
114
+ return chain_state, kernel_state
115
+
116
+ def step_fn(target_log_prob_fn, target_and_kernel_state, seed):
117
+ kernel = _build_kernel(target_log_prob_fn)
118
+ target_chain_state, kernel_state = target_and_kernel_state
119
+
120
+ new_target_position, results = kernel.one_step(
121
+ flatten_state(target_chain_state.position),
122
+ kernel.bootstrap_results(
123
+ flatten_state(target_chain_state.position)
124
+ ),
125
+ seed=seed,
126
+ )
127
+
128
+ new_chain_and_kernel_state = (
129
+ ChainState(
130
+ position=repack_state(
131
+ new_target_position, target_chain_state.position
132
+ ),
133
+ log_density=results.accepted_results.target_log_prob,
134
+ ),
135
+ kernel_state,
136
+ )
137
+
138
+ return new_chain_and_kernel_state, LeftCensoredEventsInfo(
139
+ is_accepted=results.is_accepted,
140
+ log_acceptance_correction=results.proposed_results.log_acceptance_correction,
141
+ target_log_prob=results.proposed_results.target_log_prob,
142
+ unit=results.proposed_results.unit,
143
+ timepoint=results.proposed_results.timepoint,
144
+ direction=results.proposed_results.direction,
145
+ num_events=results.proposed_results.num_events,
146
+ seed=seed,
147
+ )
148
+
149
+ return SamplingAlgorithm(init_fn, step_fn)
@@ -0,0 +1,71 @@
1
+ """Move partially-censored events in DiscreteTimeStateTransitionModel"""
2
+
3
+ from typing import NamedTuple
4
+
5
+ import tensorflow_probability as tfp
6
+
7
+ from gemlib.mcmc.discrete_time_state_transition_model import (
8
+ TransitionTopology,
9
+ UncalibratedEventTimesUpdate,
10
+ )
11
+ from gemlib.mcmc.experimental.mcmc_base import ChainState, SamplingAlgorithm
12
+
13
+
14
+ class MoveEventsState(NamedTuple):
15
+ dmax: int
16
+ mmax: int
17
+ nmax: int
18
+
19
+
20
+ class MoveEventsInfo(NamedTuple):
21
+ is_accepted: bool
22
+
23
+
24
+ def move_events(topology: TransitionTopology, dmax, mmax, nmax):
25
+ def _build_kernel(target_log_prob_fn, initial_conditions):
26
+ return tfp.mcmc.MetropolisHastings(
27
+ inner_kernel=UncalibratedEventTimesUpdate(
28
+ target_log_prob_fn,
29
+ topology,
30
+ initial_conditions,
31
+ dmax,
32
+ mmax,
33
+ nmax,
34
+ )
35
+ )
36
+
37
+ def init_fn(target_log_prob_fn, target_state, initial_conditions):
38
+ kernel = _build_kernel(target_log_prob_fn, initial_conditions)
39
+ results = kernel.bootstrap_results(target_state)
40
+ chain_state = ChainState(
41
+ position=target_state,
42
+ log_density=results.accepted_results.target_log_prob,
43
+ )
44
+ kernel_state = MoveEventsState(dmax=dmax, mmax=mmax, nmax=nmax)
45
+
46
+ return chain_state, kernel_state
47
+
48
+ def step_fn(
49
+ target_log_prob_fn, target_and_kernel_state, seed, initial_conditions
50
+ ):
51
+ kernel = _build_kernel(target_log_prob_fn, initial_conditions)
52
+
53
+ target_chain_state, kernel_state = target_and_kernel_state
54
+
55
+ new_target_position, results = kernel.one_step(
56
+ target_chain_state.position,
57
+ kernel.bootstrap_results(target_chain_state.position),
58
+ seed=seed,
59
+ )
60
+
61
+ new_chain_and_kernel_state = (
62
+ ChainState(
63
+ position=new_target_position,
64
+ log_density=results.accepted_results.target_log_prob,
65
+ ),
66
+ kernel_state,
67
+ )
68
+
69
+ return new_chain_and_kernel_state, MoveEventsInfo(results.is_accepted)
70
+
71
+ return SamplingAlgorithm(init_fn, step_fn)
@@ -0,0 +1,48 @@
1
+ """Test partially censored events move for DiscreteTimeStateTransitionModel"""
2
+
3
+ import pytest
4
+ import tensorflow as tf
5
+
6
+ from gemlib.mcmc.discrete_time_state_transition_model import TransitionTopology
7
+
8
+ from .move_events import move_events
9
+
10
+
11
+ @pytest.fixture
12
+ def random_events():
13
+ """SEIR model with prescribed starting conditions"""
14
+ events = tf.random.uniform(
15
+ [10, 10, 3], minval=0, maxval=100, dtype=tf.float64, seed=0
16
+ )
17
+ return events
18
+
19
+
20
+ @pytest.fixture
21
+ def initial_state():
22
+ popsize = tf.fill([10], tf.constant(100.0, tf.float64))
23
+ initial_state = tf.stack(
24
+ [
25
+ popsize,
26
+ tf.ones_like(popsize),
27
+ tf.zeros_like(popsize),
28
+ tf.zeros_like(popsize),
29
+ ],
30
+ axis=-1,
31
+ )
32
+ return initial_state
33
+
34
+
35
+ def test_move_events(random_events, initial_state):
36
+ def tlp(x):
37
+ return tf.constant(0.0, dtype=tf.float64)
38
+
39
+ kernel = move_events(
40
+ topology=TransitionTopology(0, 1, 2),
41
+ dmax=5,
42
+ mmax=1,
43
+ nmax=10,
44
+ )
45
+
46
+ state = kernel.init(tlp, random_events, initial_state)
47
+ seed = [0, 0]
48
+ new_state, results = kernel.step(tlp, state, seed, initial_state)
@@ -0,0 +1,89 @@
1
+ """Right-censored events MCMC kernel for DiscreteTimeStateTransitionModel"""
2
+
3
+ from typing import NamedTuple, Optional, Tuple
4
+
5
+ import tensorflow as tf
6
+ import tensorflow_probability as tfp
7
+
8
+ from gemlib.mcmc.discrete_time_state_transition_model import (
9
+ TransitionTopology,
10
+ UncalibratedOccultUpdate,
11
+ )
12
+ from gemlib.mcmc.experimental.mcmc_base import ChainState, SamplingAlgorithm
13
+
14
+
15
+ class RightCensoredEventsState(NamedTuple):
16
+ nmax: int
17
+ t_range: Tuple[int, int]
18
+
19
+
20
+ class RightCensoredEventsInfo(NamedTuple):
21
+ is_accepted: bool
22
+ m: int
23
+ t: int
24
+ delta: int
25
+ x_star: int
26
+ proposed_state: tf.Tensor
27
+ seed: Tuple[int, int]
28
+
29
+
30
+ def right_censored_events_mh(
31
+ topology: TransitionTopology,
32
+ nmax: int = 1,
33
+ t_range: Optional[Tuple[int, int]] = None,
34
+ name: Optional[str] = None,
35
+ ):
36
+ def _build_kernel(target_log_prob_fn, initial_conditions):
37
+ return tfp.mcmc.MetropolisHastings(
38
+ inner_kernel=UncalibratedOccultUpdate(
39
+ target_log_prob_fn,
40
+ topology,
41
+ initial_conditions,
42
+ nmax,
43
+ t_range,
44
+ name,
45
+ ),
46
+ )
47
+
48
+ def init_fn(target_log_prob_fn, target_state, initial_conditions):
49
+ kernel = _build_kernel(target_log_prob_fn, initial_conditions)
50
+ results = kernel.bootstrap_results(target_state)
51
+ chain_state = ChainState(
52
+ position=target_state,
53
+ log_density=results.accepted_results.target_log_prob,
54
+ )
55
+ kernel_state = RightCensoredEventsState(nmax=nmax, t_range=t_range)
56
+
57
+ return chain_state, kernel_state
58
+
59
+ def step_fn(
60
+ target_log_prob_fn, target_and_kernel_state, seed, initial_conditions
61
+ ):
62
+ kernel = _build_kernel(target_log_prob_fn, initial_conditions)
63
+ target_chain_state, kernel_state = target_and_kernel_state
64
+
65
+ new_target_position, results = kernel.one_step(
66
+ target_chain_state.position,
67
+ kernel.bootstrap_results(target_chain_state.position),
68
+ seed=seed,
69
+ )
70
+
71
+ new_chain_and_kernel_state = (
72
+ ChainState(
73
+ position=new_target_position,
74
+ log_density=results.accepted_results.target_log_prob,
75
+ ),
76
+ kernel_state,
77
+ )
78
+
79
+ return new_chain_and_kernel_state, RightCensoredEventsInfo(
80
+ results.is_accepted,
81
+ m=results.proposed_results.m,
82
+ t=results.proposed_results.t,
83
+ delta=results.proposed_results.delta_t,
84
+ x_star=results.proposed_results.x_star,
85
+ proposed_state=results.proposed_state,
86
+ seed=seed,
87
+ )
88
+
89
+ return SamplingAlgorithm(init_fn, step_fn)
@@ -0,0 +1,28 @@
1
+ """Right-censored events MCMC test"""
2
+
3
+ import pytest
4
+ import tensorflow as tf
5
+
6
+ from gemlib.mcmc.discrete_time_state_transition_model import TransitionTopology
7
+
8
+ from .right_censored_events_mh import right_censored_events_mh
9
+
10
+
11
+ @pytest.mark.skip("Work in progress")
12
+ def test_right_censored_events_mh(random_events, initial_state):
13
+ def tlp(_):
14
+ return tf.constant(0.0, tf.float64)
15
+
16
+ kernel = right_censored_events_mh(
17
+ topology=TransitionTopology(0, 1, 2),
18
+ nmax=10,
19
+ t_range=(6, 10),
20
+ )
21
+
22
+ state = kernel.init(tlp, random_events, initial_state)
23
+ seed = [0, 0]
24
+ new_state, info = kernel.step(tlp, state, seed, initial_state)
25
+
26
+ print(info)
27
+
28
+ raise AssertionError()