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,505 @@
|
|
|
1
|
+
"""Gibbs sampling kernel"""
|
|
2
|
+
# ruff: noqa: B023
|
|
3
|
+
|
|
4
|
+
import logging
|
|
5
|
+
from collections import namedtuple
|
|
6
|
+
from typing import Callable, List, Tuple
|
|
7
|
+
|
|
8
|
+
import tensorflow as tf
|
|
9
|
+
import tensorflow_probability as tfp
|
|
10
|
+
from tensorflow_probability.python.internal import samplers, unnest
|
|
11
|
+
from tensorflow_probability.python.mcmc.internal import util as mcmc_util
|
|
12
|
+
|
|
13
|
+
tfd = tfp.distributions # pylint: disable=no-member
|
|
14
|
+
tfb = tfp.bijectors # pylint: disable=no-member
|
|
15
|
+
mcmc = tfp.mcmc # pylint: disable=no-member
|
|
16
|
+
|
|
17
|
+
logging.basicConfig(format="[%(asctime)s] %(levelname)s %(message)s")
|
|
18
|
+
logger = logging.getLogger("gemlib.gibbs_kernel")
|
|
19
|
+
logger.setLevel(logging.DEBUG)
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
class GibbsKernelResults(
|
|
23
|
+
mcmc_util.PrettyNamedTupleMixin,
|
|
24
|
+
namedtuple(
|
|
25
|
+
"GibbsKernelResults",
|
|
26
|
+
["target_log_prob", "inner_results", "seed"],
|
|
27
|
+
),
|
|
28
|
+
):
|
|
29
|
+
"""Represents kernel results"""
|
|
30
|
+
|
|
31
|
+
__slots__ = ()
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
class GibbsStep(
|
|
35
|
+
mcmc_util.PrettyNamedTupleMixin,
|
|
36
|
+
namedtuple(
|
|
37
|
+
"GibbsStep",
|
|
38
|
+
["state_parts", "kernel_fn"],
|
|
39
|
+
),
|
|
40
|
+
):
|
|
41
|
+
"""Represents a Gibbs step"""
|
|
42
|
+
|
|
43
|
+
__slots__ = ()
|
|
44
|
+
|
|
45
|
+
|
|
46
|
+
def is_namedtuple(x):
|
|
47
|
+
"""Return `True` if `x` looks like a `namedtuple`."""
|
|
48
|
+
return hasattr(x, "_fields")
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _flatten_results(results):
|
|
52
|
+
"""Results structures from nested Gibbs samplers sometimes
|
|
53
|
+
need flattening for writing out purposes.
|
|
54
|
+
"""
|
|
55
|
+
|
|
56
|
+
def recurse(r):
|
|
57
|
+
for i in iter(r):
|
|
58
|
+
if isinstance(i, list):
|
|
59
|
+
yield from _flatten_results(i)
|
|
60
|
+
else:
|
|
61
|
+
yield i
|
|
62
|
+
|
|
63
|
+
return list(recurse(results))
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
def _has_gradients(results):
|
|
67
|
+
return unnest.has_nested(results, "grads_target_log_prob")
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def _get_target_log_prob(results):
|
|
71
|
+
"""Fetches a target log prob from a results structure"""
|
|
72
|
+
return unnest.get_innermost(results, "target_log_prob")
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def _update_target_log_prob(results, target_log_prob):
|
|
76
|
+
"""Puts a target log prob into a results structure"""
|
|
77
|
+
if isinstance(results, GibbsKernelResults):
|
|
78
|
+
replace_fn = unnest.replace_outermost
|
|
79
|
+
else:
|
|
80
|
+
replace_fn = unnest.replace_innermost
|
|
81
|
+
return replace_fn(results, target_log_prob=target_log_prob)
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def _maybe_transform_value(tlp, state, kernel, direction):
|
|
85
|
+
if not isinstance(kernel, tfp.mcmc.TransformedTransitionKernel):
|
|
86
|
+
return tlp
|
|
87
|
+
|
|
88
|
+
jacobian_parts = [
|
|
89
|
+
b.inverse_log_det_jacobian(x)
|
|
90
|
+
for b, x in zip(
|
|
91
|
+
tf.nest.flatten(kernel.bijector), tf.nest.flatten(state)
|
|
92
|
+
)
|
|
93
|
+
]
|
|
94
|
+
jacobian = tf.math.add_n(jacobian_parts)
|
|
95
|
+
|
|
96
|
+
if direction == "forward":
|
|
97
|
+
return tlp + jacobian
|
|
98
|
+
if direction == "inverse":
|
|
99
|
+
return tlp - jacobian
|
|
100
|
+
|
|
101
|
+
raise AttributeError("`direction` must be `forward` or `inverse`")
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def _make_namedtuple(input_dict):
|
|
105
|
+
return namedtuple("NamedTuple", input_dict.keys())(**input_dict)
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def _split_namedtuple(full_namedtuple, subset_names):
|
|
109
|
+
"""Splits a StructTuple of variables into `subset` and `compl`
|
|
110
|
+
|
|
111
|
+
:param full_struct_tuple: `namedtuple` to split
|
|
112
|
+
:param subset_names: names of required subset vars
|
|
113
|
+
:returns: a tuple `(subset: namedtuple, compl: namedtuple)`
|
|
114
|
+
"""
|
|
115
|
+
full_dict = full_namedtuple._asdict()
|
|
116
|
+
subset = _make_namedtuple(
|
|
117
|
+
{k: v for k, v in full_dict.items() if k in subset_names}
|
|
118
|
+
)
|
|
119
|
+
compl = _make_namedtuple(
|
|
120
|
+
{k: v for k, v in full_dict.items() if k not in subset_names}
|
|
121
|
+
)
|
|
122
|
+
return subset, compl
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
def _split_state(global_state, indices):
|
|
126
|
+
"""Split a global state into subset and complement
|
|
127
|
+
|
|
128
|
+
Args:
|
|
129
|
+
----
|
|
130
|
+
global_state: a tuple or namedtuple representing the global state
|
|
131
|
+
subset: a tuple of indices or names (if `global_state` is a namedtuple)
|
|
132
|
+
representing a subset.
|
|
133
|
+
|
|
134
|
+
Returns:
|
|
135
|
+
-------
|
|
136
|
+
a tuple of `(subset, complement)`
|
|
137
|
+
|
|
138
|
+
"""
|
|
139
|
+
if is_namedtuple(global_state):
|
|
140
|
+
return _split_namedtuple(global_state, indices)
|
|
141
|
+
|
|
142
|
+
return (
|
|
143
|
+
[s for i, s in enumerate(global_state) if i in indices],
|
|
144
|
+
[s for i, s in enumerate(global_state) if i not in indices],
|
|
145
|
+
)
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
def _scatter_state(global_state, subset, indices=()):
|
|
149
|
+
"""Scatters `subset` into `global_state`.
|
|
150
|
+
|
|
151
|
+
Args:
|
|
152
|
+
----
|
|
153
|
+
global_state: a tuple or namedtuple representing the global state
|
|
154
|
+
subset: a tuple or namedtuple containing values to be scattered into
|
|
155
|
+
`global_state`.
|
|
156
|
+
indices: if `subset` is a tuple, `indices` is a tuple of corresponding
|
|
157
|
+
indices into `global_state`.
|
|
158
|
+
|
|
159
|
+
Returns:
|
|
160
|
+
-------
|
|
161
|
+
a tuple or namedtuple of the same structure as `global_state`.
|
|
162
|
+
|
|
163
|
+
"""
|
|
164
|
+
if is_namedtuple(global_state) and is_namedtuple(subset):
|
|
165
|
+
return global_state._replace(**subset._asdict())
|
|
166
|
+
|
|
167
|
+
for i, state_part in zip(indices, subset):
|
|
168
|
+
global_state[i] = state_part
|
|
169
|
+
|
|
170
|
+
return global_state
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
class GibbsKernel(mcmc.TransitionKernel):
|
|
174
|
+
"""Component-wise MCMC sampling.
|
|
175
|
+
|
|
176
|
+
``GibbsKernel`` is designed to fit within TensorFlow Probability's MCMC
|
|
177
|
+
framework, essentially acting as a "meta-kernel" that aggregates a
|
|
178
|
+
sequence of component-wise kernels.
|
|
179
|
+
|
|
180
|
+
Example:
|
|
181
|
+
-------
|
|
182
|
+
Sample from the posterior of a linear model::
|
|
183
|
+
|
|
184
|
+
import numpy as np
|
|
185
|
+
import tensorflow as tf
|
|
186
|
+
import tensorflow_probability as tfp
|
|
187
|
+
from gemlib.mcmc.gibbs_kernel import GibbsKernel
|
|
188
|
+
|
|
189
|
+
tfd = tfp.distributions
|
|
190
|
+
|
|
191
|
+
dtype = np.float32
|
|
192
|
+
|
|
193
|
+
# data
|
|
194
|
+
x = dtype([2.9, 4.2, 8.3, 1.9, 2.6, 1.0, 8.4, 8.6, 7.9, 4.3])
|
|
195
|
+
y = dtype([6.2, 7.8, 8.1, 2.7, 4.8, 2.4, 10.7, 9.0, 9.6, 5.7])
|
|
196
|
+
|
|
197
|
+
|
|
198
|
+
# define linear regression model
|
|
199
|
+
def Model(x):
|
|
200
|
+
def alpha():
|
|
201
|
+
return tfd.Normal(loc=dtype(0.0), scale=dtype(1000.0))
|
|
202
|
+
|
|
203
|
+
def beta():
|
|
204
|
+
return tfd.Normal(loc=dtype(0.0), scale=dtype(100.0))
|
|
205
|
+
|
|
206
|
+
def sigma():
|
|
207
|
+
return tfd.Gamma(concentration=dtype(0.1), rate=dtype(0.1))
|
|
208
|
+
|
|
209
|
+
def y(alpha, beta, sigma):
|
|
210
|
+
mu = alpha + beta * x
|
|
211
|
+
return tfd.Normal(mu, scale=sigma)
|
|
212
|
+
|
|
213
|
+
return tfd.JointDistributionNamed(
|
|
214
|
+
dict(alpha=alpha, beta=beta, sigma=sigma, y=y)
|
|
215
|
+
)
|
|
216
|
+
|
|
217
|
+
|
|
218
|
+
# target log probability of linear model
|
|
219
|
+
def log_prob(alpha, beta, sigma):
|
|
220
|
+
lp = model.log_prob(
|
|
221
|
+
{"alpha": alpha, "beta": beta, "sigma": sigma, "y": y}
|
|
222
|
+
)
|
|
223
|
+
return tf.reduce_sum(lp)
|
|
224
|
+
|
|
225
|
+
|
|
226
|
+
# random walk Markov chain function
|
|
227
|
+
def kernel_make_fn(target_log_prob_fn, state):
|
|
228
|
+
return tfp.mcmc.RandomWalkMetropolis(
|
|
229
|
+
target_log_prob_fn=target_log_prob_fn
|
|
230
|
+
)
|
|
231
|
+
|
|
232
|
+
|
|
233
|
+
# posterior distribution MCMC chain
|
|
234
|
+
@tf.function
|
|
235
|
+
def posterior(iterations, burnin, thinning, initial_state):
|
|
236
|
+
kernel_list = [
|
|
237
|
+
(
|
|
238
|
+
0,
|
|
239
|
+
kernel_make_fn,
|
|
240
|
+
), # conditional probability for zeroth parmeter alpha
|
|
241
|
+
(
|
|
242
|
+
1,
|
|
243
|
+
kernel_make_fn,
|
|
244
|
+
), # conditional probability for first parameter beta
|
|
245
|
+
(2, kernel_make_fn),
|
|
246
|
+
] # conditional probability for second parameter sigma
|
|
247
|
+
kernel = GibbsKernel(
|
|
248
|
+
target_log_prob_fn=log_prob, kernel_list=kernel_list
|
|
249
|
+
)
|
|
250
|
+
return tfp.mcmc.sample_chain(
|
|
251
|
+
num_results=iterations,
|
|
252
|
+
current_state=initial_state,
|
|
253
|
+
kernel=kernel,
|
|
254
|
+
num_burnin_steps=burnin,
|
|
255
|
+
num_steps_between_results=thinning,
|
|
256
|
+
parallel_iterations=1,
|
|
257
|
+
trace_fn=None,
|
|
258
|
+
)
|
|
259
|
+
|
|
260
|
+
|
|
261
|
+
# initialize model
|
|
262
|
+
model = Model(x)
|
|
263
|
+
initial_state = [
|
|
264
|
+
dtype(0.1),
|
|
265
|
+
dtype(0.1),
|
|
266
|
+
dtype(0.1),
|
|
267
|
+
] # start chain at alpha=0.1, beta=0.1, sigma=0.1
|
|
268
|
+
|
|
269
|
+
# estimate posterior distribution
|
|
270
|
+
samples = posterior(
|
|
271
|
+
iterations=10000,
|
|
272
|
+
burnin=1000,
|
|
273
|
+
thinning=0,
|
|
274
|
+
initial_state=initial_state,
|
|
275
|
+
)
|
|
276
|
+
|
|
277
|
+
tf.print("alpha samples:", samples[0])
|
|
278
|
+
tf.print("beta samples:", samples[1])
|
|
279
|
+
tf.print("sigma samples:", samples[2])
|
|
280
|
+
tf.print(
|
|
281
|
+
"sample means: [alpha, beta, sigma] =",
|
|
282
|
+
tf.math.reduce_mean(samples, axis=1),
|
|
283
|
+
)
|
|
284
|
+
|
|
285
|
+
"""
|
|
286
|
+
|
|
287
|
+
def __init__(
|
|
288
|
+
self,
|
|
289
|
+
target_log_prob_fn: Callable[[float], float],
|
|
290
|
+
kernel_list: List[Tuple[Tuple[int, ...], Callable]],
|
|
291
|
+
name: str = None,
|
|
292
|
+
):
|
|
293
|
+
"""Build a Gibbs sampling scheme from component kernels.
|
|
294
|
+
|
|
295
|
+
Args:
|
|
296
|
+
----
|
|
297
|
+
target_log_prob_fn: a function that takes `state` arguments
|
|
298
|
+
and returns the target log probability
|
|
299
|
+
density.
|
|
300
|
+
kernel_list: a list of tuples `(state_part_idx, kernel_make_fn)`.
|
|
301
|
+
`state_part_idx` denotes the index (relative to
|
|
302
|
+
positional args in `target_log_prob_fn`) of the
|
|
303
|
+
state the kernel updates. `kernel_make_fn` takes
|
|
304
|
+
arguments `target_log_prob_fn` and `state`,
|
|
305
|
+
returning a `tfp.mcmc.TransitionKernel`.
|
|
306
|
+
|
|
307
|
+
Returns:
|
|
308
|
+
-------
|
|
309
|
+
an instance of `GibbsKernel`
|
|
310
|
+
|
|
311
|
+
"""
|
|
312
|
+
# Require to check if all kernel.is_calibrated is True
|
|
313
|
+
self._parameters = {
|
|
314
|
+
"target_log_prob_fn": target_log_prob_fn,
|
|
315
|
+
"kernel_list": kernel_list,
|
|
316
|
+
"name": name,
|
|
317
|
+
}
|
|
318
|
+
|
|
319
|
+
@property
|
|
320
|
+
def is_calibrated(self):
|
|
321
|
+
return True
|
|
322
|
+
|
|
323
|
+
@property
|
|
324
|
+
def target_log_prob_fn(self):
|
|
325
|
+
"""Target log probability function."""
|
|
326
|
+
return self._parameters["target_log_prob_fn"]
|
|
327
|
+
|
|
328
|
+
@property
|
|
329
|
+
def kernel_list(self):
|
|
330
|
+
"""List of kernel-build functions."""
|
|
331
|
+
return self._parameters["kernel_list"]
|
|
332
|
+
|
|
333
|
+
@property
|
|
334
|
+
def name(self):
|
|
335
|
+
"""Name of the kernel."""
|
|
336
|
+
return self._parameters["name"]
|
|
337
|
+
|
|
338
|
+
def one_step(self, current_state, previous_results, seed=None):
|
|
339
|
+
"""Iterate over the state elements, calling each kernel in turn.
|
|
340
|
+
|
|
341
|
+
The ``target_log_prob`` is forwarded to the next ``previous_results``
|
|
342
|
+
such that each kernel has a current ``target_log_prob`` value.
|
|
343
|
+
Transformations are automatically performed if the kernel is of
|
|
344
|
+
type ``tfp.mcmc.TransformedTransitionKernel``.
|
|
345
|
+
|
|
346
|
+
In graph and XLA modes, the for loop should be unrolled.
|
|
347
|
+
|
|
348
|
+
Args:
|
|
349
|
+
----
|
|
350
|
+
current_state: the current chain state
|
|
351
|
+
previous_results: a ``GibbsKernelResults`` instance
|
|
352
|
+
seed: an optional list of two scalar ``int`` tensors.
|
|
353
|
+
|
|
354
|
+
Returns:
|
|
355
|
+
-------
|
|
356
|
+
a tuple of ``(next_state, results)``.
|
|
357
|
+
|
|
358
|
+
"""
|
|
359
|
+
seed = samplers.sanitize_seed(seed, salt="GibbsKernel")
|
|
360
|
+
|
|
361
|
+
global_state_parts = current_state
|
|
362
|
+
|
|
363
|
+
next_results = []
|
|
364
|
+
untransformed_target_log_prob = previous_results.target_log_prob
|
|
365
|
+
seeds = samplers.split_seed(seed, n=len(self.kernel_list))
|
|
366
|
+
|
|
367
|
+
for (state_part_indices, kernel_fn), previous_step_results, seed in zip(
|
|
368
|
+
self.kernel_list, previous_results.inner_results, seeds
|
|
369
|
+
):
|
|
370
|
+
if not mcmc_util.is_list_like(state_part_indices):
|
|
371
|
+
state_part_indices = [state_part_indices] # noqa: PLW2901
|
|
372
|
+
|
|
373
|
+
# Extract state parts required for step
|
|
374
|
+
step_state_parts, _ = _split_state(
|
|
375
|
+
global_state_parts, state_part_indices
|
|
376
|
+
)
|
|
377
|
+
|
|
378
|
+
def target_log_prob_fn(*kernel_state):
|
|
379
|
+
if is_namedtuple(step_state_parts):
|
|
380
|
+
kernel_state = step_state_parts.__class__(*kernel_state)
|
|
381
|
+
state_parts = _scatter_state(
|
|
382
|
+
global_state_parts,
|
|
383
|
+
kernel_state,
|
|
384
|
+
state_part_indices,
|
|
385
|
+
)
|
|
386
|
+
return self.target_log_prob_fn(*state_parts)
|
|
387
|
+
|
|
388
|
+
# Build kernel function
|
|
389
|
+
kernel = kernel_fn(target_log_prob_fn, global_state_parts)
|
|
390
|
+
|
|
391
|
+
# Forward the current tlp to the kernel. If the kernel is
|
|
392
|
+
# gradient-based, we need to calculate fresh gradients,
|
|
393
|
+
# as these cannot easily be forwarded
|
|
394
|
+
# from the previous Gibbs step.
|
|
395
|
+
if _has_gradients(previous_step_results):
|
|
396
|
+
# TODO would be better to avoid re-calculating the whole of
|
|
397
|
+
# `bootstrap_results` when we just need to calculate gradients.
|
|
398
|
+
fresh_previous_results = unnest.UnnestingWrapper(
|
|
399
|
+
kernel.bootstrap_results(step_state_parts)
|
|
400
|
+
)
|
|
401
|
+
previous_step_results = unnest.replace_innermost( # noqa: PLW2901
|
|
402
|
+
previous_step_results,
|
|
403
|
+
target_log_prob=fresh_previous_results.target_log_prob,
|
|
404
|
+
grads_target_log_prob=fresh_previous_results.grads_target_log_prob,
|
|
405
|
+
)
|
|
406
|
+
|
|
407
|
+
else:
|
|
408
|
+
previous_step_results = _update_target_log_prob( # noqa: PLW2901
|
|
409
|
+
previous_step_results,
|
|
410
|
+
_maybe_transform_value(
|
|
411
|
+
tlp=untransformed_target_log_prob,
|
|
412
|
+
state=step_state_parts,
|
|
413
|
+
kernel=kernel,
|
|
414
|
+
direction="inverse",
|
|
415
|
+
),
|
|
416
|
+
)
|
|
417
|
+
|
|
418
|
+
new_step_state_parts, next_kernel_results = kernel.one_step(
|
|
419
|
+
step_state_parts, previous_step_results, seed
|
|
420
|
+
)
|
|
421
|
+
if is_namedtuple(step_state_parts):
|
|
422
|
+
new_step_state_parts = step_state_parts.__class__(
|
|
423
|
+
*new_step_state_parts,
|
|
424
|
+
)
|
|
425
|
+
|
|
426
|
+
next_results.append(next_kernel_results)
|
|
427
|
+
|
|
428
|
+
# Cache the new tlp for use in the next Gibbs step
|
|
429
|
+
untransformed_target_log_prob = _maybe_transform_value(
|
|
430
|
+
tlp=_get_target_log_prob(next_kernel_results),
|
|
431
|
+
state=new_step_state_parts,
|
|
432
|
+
kernel=kernel,
|
|
433
|
+
direction="forward",
|
|
434
|
+
)
|
|
435
|
+
|
|
436
|
+
global_state_parts = _scatter_state(
|
|
437
|
+
global_state_parts, new_step_state_parts, state_part_indices
|
|
438
|
+
)
|
|
439
|
+
|
|
440
|
+
if is_namedtuple(current_state):
|
|
441
|
+
global_state_parts = current_state.__class__(
|
|
442
|
+
*global_state_parts,
|
|
443
|
+
)
|
|
444
|
+
|
|
445
|
+
return (
|
|
446
|
+
global_state_parts
|
|
447
|
+
if mcmc_util.is_list_like(current_state)
|
|
448
|
+
else global_state_parts[0],
|
|
449
|
+
GibbsKernelResults(
|
|
450
|
+
target_log_prob=untransformed_target_log_prob,
|
|
451
|
+
inner_results=next_results,
|
|
452
|
+
seed=seeds[-1],
|
|
453
|
+
),
|
|
454
|
+
)
|
|
455
|
+
|
|
456
|
+
def bootstrap_results(self, current_state):
|
|
457
|
+
"""Set up the results tuple.
|
|
458
|
+
|
|
459
|
+
Args:
|
|
460
|
+
----
|
|
461
|
+
current_state: a list of state parts representing the Markov chain
|
|
462
|
+
state
|
|
463
|
+
Returns:
|
|
464
|
+
an instance of `GibbsKernelResults`
|
|
465
|
+
|
|
466
|
+
"""
|
|
467
|
+
global_state_parts = current_state
|
|
468
|
+
inner_results = []
|
|
469
|
+
untransformed_target_log_prob = 0.0
|
|
470
|
+
|
|
471
|
+
for state_part_indices, kernel_fn in self.kernel_list:
|
|
472
|
+
if not mcmc_util.is_list_like(state_part_indices):
|
|
473
|
+
state_part_indices = [state_part_indices] # noqa: PLW2901
|
|
474
|
+
|
|
475
|
+
step_state_parts, _ = _split_state(
|
|
476
|
+
global_state_parts, state_part_indices
|
|
477
|
+
)
|
|
478
|
+
|
|
479
|
+
def tlp_fn(*kernel_state):
|
|
480
|
+
if is_namedtuple(step_state_parts):
|
|
481
|
+
kernel_state = step_state_parts.__class__(*kernel_state)
|
|
482
|
+
state_parts = _scatter_state(
|
|
483
|
+
global_state_parts,
|
|
484
|
+
kernel_state,
|
|
485
|
+
state_part_indices,
|
|
486
|
+
)
|
|
487
|
+
|
|
488
|
+
return self.target_log_prob_fn(*state_parts)
|
|
489
|
+
|
|
490
|
+
kernel = kernel_fn(tlp_fn, global_state_parts)
|
|
491
|
+
kernel_results = kernel.bootstrap_results(step_state_parts)
|
|
492
|
+
|
|
493
|
+
inner_results.append(kernel_results)
|
|
494
|
+
untransformed_target_log_prob = _maybe_transform_value(
|
|
495
|
+
tlp=_get_target_log_prob(kernel_results),
|
|
496
|
+
state=step_state_parts,
|
|
497
|
+
kernel=kernel,
|
|
498
|
+
direction="forward",
|
|
499
|
+
)
|
|
500
|
+
|
|
501
|
+
return GibbsKernelResults(
|
|
502
|
+
target_log_prob=untransformed_target_log_prob,
|
|
503
|
+
inner_results=inner_results,
|
|
504
|
+
seed=samplers.zeros_seed(),
|
|
505
|
+
)
|