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,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)
|
gemlib/mcmc/experimental/discrete_time_state_transition_model/right_censored_events_mh_test.py
ADDED
|
@@ -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()
|