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,111 @@
|
|
|
1
|
+
"""Base Hamiltonian Monte Carlo"""
|
|
2
|
+
|
|
3
|
+
from typing import Iterable, NamedTuple, Optional
|
|
4
|
+
|
|
5
|
+
import tensorflow as tf
|
|
6
|
+
import tensorflow_probability as tfp
|
|
7
|
+
import tensorflow_probability.python.experimental.mcmc.preconditioning_utils as pu # noqa: E501
|
|
8
|
+
|
|
9
|
+
from .mcmc_base import ChainState, SamplingAlgorithm
|
|
10
|
+
|
|
11
|
+
tfd = tfp.distributions
|
|
12
|
+
tfde = tfp.experimental.distributions
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class HmcKernelState(NamedTuple):
|
|
16
|
+
step_size: float
|
|
17
|
+
num_leapfrog_steps: float
|
|
18
|
+
mass_matrix: Iterable
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def hmc(
|
|
22
|
+
step_size: float = 0.1,
|
|
23
|
+
num_leapfrog_steps: int = 16,
|
|
24
|
+
mass_matrix: Optional[Iterable] = None,
|
|
25
|
+
):
|
|
26
|
+
"""Hamiltonian Monte Carlo
|
|
27
|
+
|
|
28
|
+
Args
|
|
29
|
+
----
|
|
30
|
+
step_size: the step size to take
|
|
31
|
+
num_leapfrog_steps: number of leapfrog steps to take
|
|
32
|
+
mass_matrix: a mass matrix (defaults to diag(1) if None)
|
|
33
|
+
|
|
34
|
+
Returns
|
|
35
|
+
-------
|
|
36
|
+
A SamplingAlgorithm
|
|
37
|
+
"""
|
|
38
|
+
|
|
39
|
+
step_size = tf.convert_to_tensor(step_size)
|
|
40
|
+
mass_matrix = (
|
|
41
|
+
tf.convert_to_tensor(mass_matrix) if mass_matrix is not None else None
|
|
42
|
+
)
|
|
43
|
+
|
|
44
|
+
def _make_momentum_distribution(position):
|
|
45
|
+
# return None
|
|
46
|
+
if mass_matrix is None:
|
|
47
|
+
return pu.make_momentum_distribution(
|
|
48
|
+
position, tf.constant([], dtype=tf.int32)
|
|
49
|
+
)
|
|
50
|
+
|
|
51
|
+
else:
|
|
52
|
+
return tfde.MultivariateNormalPrecisionFactorLinearOperator(
|
|
53
|
+
precision_factor=tf.linalg.LinearOperatorFullMatrix(
|
|
54
|
+
mass_matrix
|
|
55
|
+
),
|
|
56
|
+
precision=tf.linalg.LinearOperatorFullMatrix(mass_matrix),
|
|
57
|
+
)
|
|
58
|
+
|
|
59
|
+
def _build_kernel(target_log_prob_fn, momentum_distribution):
|
|
60
|
+
return tfp.experimental.mcmc.PreconditionedHamiltonianMonteCarlo(
|
|
61
|
+
target_log_prob_fn=target_log_prob_fn,
|
|
62
|
+
step_size=step_size,
|
|
63
|
+
num_leapfrog_steps=num_leapfrog_steps,
|
|
64
|
+
momentum_distribution=momentum_distribution,
|
|
65
|
+
store_parameters_in_results=True,
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
def init_fn(target_log_prob_fn, initial_position):
|
|
69
|
+
kernel = _build_kernel(
|
|
70
|
+
target_log_prob_fn, _make_momentum_distribution(initial_position)
|
|
71
|
+
)
|
|
72
|
+
results = kernel.bootstrap_results(initial_position)
|
|
73
|
+
|
|
74
|
+
# Repack the results data structure into our own
|
|
75
|
+
chain_state = ChainState(
|
|
76
|
+
position=initial_position,
|
|
77
|
+
log_density=results.accepted_results.target_log_prob,
|
|
78
|
+
log_density_grad=results.accepted_results.grads_target_log_prob,
|
|
79
|
+
)
|
|
80
|
+
kernel_state = results
|
|
81
|
+
|
|
82
|
+
return chain_state, kernel_state
|
|
83
|
+
|
|
84
|
+
def step_fn(target_log_prob_fn, chain_and_kernel_state, seed):
|
|
85
|
+
chain_state, kernel_state = chain_and_kernel_state
|
|
86
|
+
|
|
87
|
+
# Pack kernel state into results here
|
|
88
|
+
seed = tfp.random.sanitize_seed(seed)
|
|
89
|
+
|
|
90
|
+
kernel = _build_kernel(
|
|
91
|
+
target_log_prob_fn,
|
|
92
|
+
_make_momentum_distribution(chain_state.position),
|
|
93
|
+
)
|
|
94
|
+
|
|
95
|
+
new_position, results = kernel.one_step(
|
|
96
|
+
chain_state.position,
|
|
97
|
+
kernel.bootstrap_results(chain_state.position),
|
|
98
|
+
seed,
|
|
99
|
+
)
|
|
100
|
+
|
|
101
|
+
info = results
|
|
102
|
+
|
|
103
|
+
chain_state = ChainState(
|
|
104
|
+
position=new_position,
|
|
105
|
+
log_density=results.accepted_results.target_log_prob,
|
|
106
|
+
log_density_grad=results.accepted_results.grads_target_log_prob,
|
|
107
|
+
)
|
|
108
|
+
|
|
109
|
+
return (chain_state, results), info
|
|
110
|
+
|
|
111
|
+
return SamplingAlgorithm(init_fn, step_fn)
|
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
"""Tests for HMC sampler"""
|
|
2
|
+
|
|
3
|
+
import numpy as np
|
|
4
|
+
import pytest
|
|
5
|
+
import tensorflow as tf
|
|
6
|
+
import tensorflow_probability as tfp
|
|
7
|
+
|
|
8
|
+
from .hmc import hmc
|
|
9
|
+
from .mcmc_sampler import mcmc
|
|
10
|
+
|
|
11
|
+
tfd = tfp.distributions
|
|
12
|
+
|
|
13
|
+
NUM_SAMPLES = 100000
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
@pytest.fixture
|
|
17
|
+
def simple_model():
|
|
18
|
+
@tfp.distributions.JointDistributionCoroutine
|
|
19
|
+
def model():
|
|
20
|
+
yield tfp.distributions.Normal(loc=0.0, scale=1.0, name="foo")
|
|
21
|
+
yield tfp.distributions.Normal(loc=1.0, scale=1.0, name="bar")
|
|
22
|
+
yield tfp.distributions.Normal(loc=2.0, scale=1.0, name="baz")
|
|
23
|
+
|
|
24
|
+
return model
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
def test_hmc(simple_model):
|
|
28
|
+
mcmc_init_state = simple_model.sample(seed=[0, 0])
|
|
29
|
+
algorithm = hmc(step_size=0.1, num_leapfrog_steps=16)
|
|
30
|
+
|
|
31
|
+
state = algorithm.init(simple_model.log_prob, mcmc_init_state)
|
|
32
|
+
new_state, info = algorithm.step(simple_model.log_prob, state, [0, 0])
|
|
33
|
+
|
|
34
|
+
tf.nest.assert_same_structure(state, new_state)
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def test_many_hmc(simple_model):
|
|
38
|
+
mcmc_init_state = simple_model.sample(seed=[0, 0])
|
|
39
|
+
algorithm = hmc(step_size=1.2, num_leapfrog_steps=16)
|
|
40
|
+
|
|
41
|
+
samples, info = tf.function(
|
|
42
|
+
lambda: mcmc(
|
|
43
|
+
NUM_SAMPLES,
|
|
44
|
+
sampling_algorithm=algorithm,
|
|
45
|
+
target_density_fn=simple_model.log_prob,
|
|
46
|
+
initial_position=mcmc_init_state,
|
|
47
|
+
seed=[0, 0],
|
|
48
|
+
),
|
|
49
|
+
jit_compile=True,
|
|
50
|
+
)()
|
|
51
|
+
|
|
52
|
+
np.testing.assert_allclose(
|
|
53
|
+
np.array([np.mean(x) for x in samples]),
|
|
54
|
+
np.array([0.0, 1.0, 2.0]),
|
|
55
|
+
atol=1e-2,
|
|
56
|
+
)
|
|
57
|
+
np.testing.assert_allclose(
|
|
58
|
+
np.array([np.var(x) for x in samples]),
|
|
59
|
+
np.array([1.0, 1.0, 1.0]),
|
|
60
|
+
atol=1e-2,
|
|
61
|
+
)
|
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
"""Base MCMC datatypes"""
|
|
2
|
+
|
|
3
|
+
from typing import Callable, NamedTuple, Optional, Tuple
|
|
4
|
+
|
|
5
|
+
Position = NamedTuple
|
|
6
|
+
KernelInfo = NamedTuple
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class ChainState(NamedTuple):
|
|
10
|
+
"""Represent the state of an MCMC probability space"""
|
|
11
|
+
|
|
12
|
+
position: Position
|
|
13
|
+
log_density: float
|
|
14
|
+
log_density_grad: Optional[float] = None
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
class KernelState(NamedTuple):
|
|
18
|
+
"""Represent the state of a stateful MCMC kernel"""
|
|
19
|
+
|
|
20
|
+
pass
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class SamplingAlgorithm(NamedTuple):
|
|
24
|
+
"""Represent a sampling algorithm"""
|
|
25
|
+
|
|
26
|
+
init: Callable[[NamedTuple], Tuple[ChainState, KernelState]]
|
|
27
|
+
step: Callable[
|
|
28
|
+
[
|
|
29
|
+
ChainState,
|
|
30
|
+
Callable[[NamedTuple], float],
|
|
31
|
+
],
|
|
32
|
+
Callable[[ChainState], Tuple[ChainState, KernelInfo]],
|
|
33
|
+
]
|
|
@@ -0,0 +1,94 @@
|
|
|
1
|
+
"""Higher-order functions to run MCMC"""
|
|
2
|
+
|
|
3
|
+
from functools import partial
|
|
4
|
+
from typing import Any, Callable, Iterable, Tuple
|
|
5
|
+
|
|
6
|
+
import tensorflow as tf
|
|
7
|
+
import tensorflow_probability as tfp
|
|
8
|
+
|
|
9
|
+
from .mcmc_base import SamplingAlgorithm
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def split_seed(seed, n):
|
|
13
|
+
n = tf.convert_to_tensor(n)
|
|
14
|
+
return tfp.random.split_seed(seed, n=n)
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _tensor_array_from_element(elem, size):
|
|
18
|
+
return tf.TensorArray(
|
|
19
|
+
elem.dtype,
|
|
20
|
+
size=size,
|
|
21
|
+
element_shape=elem.shape,
|
|
22
|
+
)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def scan(fn, init, xs):
|
|
26
|
+
"""Scan
|
|
27
|
+
|
|
28
|
+
This function is equivalent to
|
|
29
|
+
|
|
30
|
+
```
|
|
31
|
+
scan :: (c -> a -> (c, b)) -> c -> [a] -> (c, [b])
|
|
32
|
+
```
|
|
33
|
+
"""
|
|
34
|
+
# Set up results accumulator
|
|
35
|
+
_, initial_trace = fn(init, xs[0])
|
|
36
|
+
|
|
37
|
+
flat_initial_trace = tf.nest.flatten(initial_trace, expand_composites=True)
|
|
38
|
+
trace_arrays = []
|
|
39
|
+
for trace_elem in flat_initial_trace:
|
|
40
|
+
trace_arrays.append(
|
|
41
|
+
_tensor_array_from_element(trace_elem, size=xs.shape[0])
|
|
42
|
+
)
|
|
43
|
+
|
|
44
|
+
def trace_one_step(i, trace_arrays, trace):
|
|
45
|
+
return [
|
|
46
|
+
ta.write(i, x)
|
|
47
|
+
for ta, x in zip(
|
|
48
|
+
trace_arrays, tf.nest.flatten(trace, expand_composites=True)
|
|
49
|
+
)
|
|
50
|
+
]
|
|
51
|
+
|
|
52
|
+
def cond(i, carry, accum):
|
|
53
|
+
return i < xs.shape[0]
|
|
54
|
+
|
|
55
|
+
def body(i, carry, accum):
|
|
56
|
+
new_carry, result = fn(carry, xs[i])
|
|
57
|
+
new_accum = trace_one_step(i, accum, result)
|
|
58
|
+
return i + 1, new_carry, new_accum
|
|
59
|
+
|
|
60
|
+
_, final_state, trace_arrays = tf.while_loop(
|
|
61
|
+
cond=cond, body=body, loop_vars=(0, init, trace_arrays)
|
|
62
|
+
)
|
|
63
|
+
|
|
64
|
+
stacked_trace = tf.nest.pack_sequence_as(
|
|
65
|
+
initial_trace,
|
|
66
|
+
[ta.stack() for ta in trace_arrays],
|
|
67
|
+
expand_composites=True,
|
|
68
|
+
)
|
|
69
|
+
|
|
70
|
+
return final_state, stacked_trace
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def mcmc(
|
|
74
|
+
num_samples: int,
|
|
75
|
+
sampling_algorithm: SamplingAlgorithm,
|
|
76
|
+
target_density_fn: Callable[[Any, ...], float],
|
|
77
|
+
initial_position: Iterable,
|
|
78
|
+
seed: Tuple[int, int],
|
|
79
|
+
):
|
|
80
|
+
initial_position = tf.nest.map_structure(
|
|
81
|
+
lambda x: tf.convert_to_tensor(x), initial_position
|
|
82
|
+
)
|
|
83
|
+
initial_state = sampling_algorithm.init(target_density_fn, initial_position)
|
|
84
|
+
kernel_step_fn = partial(sampling_algorithm.step, target_density_fn)
|
|
85
|
+
|
|
86
|
+
def one_step(state, rng_key):
|
|
87
|
+
new_state, info = kernel_step_fn(state, rng_key)
|
|
88
|
+
return new_state, (new_state[0].position, info)
|
|
89
|
+
|
|
90
|
+
keys = split_seed(seed, num_samples)
|
|
91
|
+
|
|
92
|
+
_, trace = scan(one_step, initial_state, keys)
|
|
93
|
+
|
|
94
|
+
return trace
|
|
@@ -0,0 +1,44 @@
|
|
|
1
|
+
"""Test mcmc_sampler"""
|
|
2
|
+
|
|
3
|
+
from typing import NamedTuple
|
|
4
|
+
|
|
5
|
+
import tensorflow as tf
|
|
6
|
+
|
|
7
|
+
from .mcmc_sampler import mcmc
|
|
8
|
+
from .test_util import CountingKernelInfo, counting_kernel
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
class TestPosition(NamedTuple):
|
|
12
|
+
x: float
|
|
13
|
+
y: float
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
NUM_SAMPLES = 100
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def test_mcmc():
|
|
20
|
+
sampling_algorithm = counting_kernel()
|
|
21
|
+
|
|
22
|
+
initial_position = TestPosition(0.0, -100.0)
|
|
23
|
+
|
|
24
|
+
def tlp(x, y):
|
|
25
|
+
return tf.constant(0.0)
|
|
26
|
+
|
|
27
|
+
samples, info = mcmc(
|
|
28
|
+
num_samples=NUM_SAMPLES,
|
|
29
|
+
sampling_algorithm=sampling_algorithm,
|
|
30
|
+
target_density_fn=tlp,
|
|
31
|
+
initial_position=initial_position,
|
|
32
|
+
seed=[0, 0],
|
|
33
|
+
)
|
|
34
|
+
print(samples)
|
|
35
|
+
tf.debugging.assert_equal(
|
|
36
|
+
samples,
|
|
37
|
+
TestPosition(
|
|
38
|
+
x=tf.range(1.0, 1.0 + NUM_SAMPLES, delta=1.0),
|
|
39
|
+
y=tf.range(-99.0, -99.0 + NUM_SAMPLES, delta=1.0),
|
|
40
|
+
),
|
|
41
|
+
)
|
|
42
|
+
tf.debugging.assert_equal(
|
|
43
|
+
info, CountingKernelInfo(tf.fill(NUM_SAMPLES, True))
|
|
44
|
+
)
|
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
"""MultiScanKernel calls one_step a number of times on an inner kernel"""
|
|
2
|
+
|
|
3
|
+
from functools import partial
|
|
4
|
+
|
|
5
|
+
import tensorflow as tf
|
|
6
|
+
from tensorflow_probability.python.internal import samplers
|
|
7
|
+
|
|
8
|
+
from .mcmc_base import SamplingAlgorithm
|
|
9
|
+
|
|
10
|
+
__all__ = ["multi_scan"]
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def multi_scan(
|
|
14
|
+
num_updates: int, sampling_algorithm: SamplingAlgorithm
|
|
15
|
+
) -> SamplingAlgorithm:
|
|
16
|
+
"""Performs multiple applications of a kernel
|
|
17
|
+
|
|
18
|
+
`sampling_algorithm` is invoked `num_updates` times
|
|
19
|
+
returning the state and info after the last step.
|
|
20
|
+
|
|
21
|
+
Args
|
|
22
|
+
----
|
|
23
|
+
num_updates: integer giving the number of updates
|
|
24
|
+
sampling_algorithm: an instance of `SamplingAlgorithm`
|
|
25
|
+
|
|
26
|
+
Return
|
|
27
|
+
------
|
|
28
|
+
An instance of `SamplingAlgorithm`
|
|
29
|
+
"""
|
|
30
|
+
|
|
31
|
+
num_updates_ = tf.convert_to_tensor(num_updates)
|
|
32
|
+
init_fn = sampling_algorithm.init
|
|
33
|
+
|
|
34
|
+
def step_fn(target_log_prob_fn, current_state, seed=None):
|
|
35
|
+
seed = samplers.sanitize_seed(seed, salt="multi_scan_kernel")
|
|
36
|
+
seeds = samplers.split_seed(seed, n=num_updates)
|
|
37
|
+
kernel = partial(sampling_algorithm.step, target_log_prob_fn)
|
|
38
|
+
|
|
39
|
+
def body(i, state, _):
|
|
40
|
+
state, info = kernel(state, seeds[i])
|
|
41
|
+
return i + 1, state, info
|
|
42
|
+
|
|
43
|
+
def cond(i, *_):
|
|
44
|
+
return i < num_updates_
|
|
45
|
+
|
|
46
|
+
init_state, init_info = kernel(current_state, seed) # unrolled first it
|
|
47
|
+
|
|
48
|
+
_, last_state, last_info = tf.while_loop(
|
|
49
|
+
cond, body, loop_vars=(1, init_state, init_info)
|
|
50
|
+
)
|
|
51
|
+
return last_state, last_info
|
|
52
|
+
|
|
53
|
+
return SamplingAlgorithm(init_fn, step_fn)
|
|
@@ -0,0 +1,71 @@
|
|
|
1
|
+
"""Test modules for multiscan kernel"""
|
|
2
|
+
|
|
3
|
+
from typing import NamedTuple
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import tensorflow as tf
|
|
7
|
+
|
|
8
|
+
from .mcmc_sampler import mcmc
|
|
9
|
+
from .multi_scan import multi_scan
|
|
10
|
+
from .test_util import CountingKernelInfo, counting_kernel
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
class TestPosition(NamedTuple):
|
|
14
|
+
x: float
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def test_one_multi_scan():
|
|
18
|
+
multi_scan_iterations = 100
|
|
19
|
+
|
|
20
|
+
sampler = multi_scan(multi_scan_iterations, counting_kernel())
|
|
21
|
+
|
|
22
|
+
def tlp(x):
|
|
23
|
+
return x
|
|
24
|
+
|
|
25
|
+
initial_position = TestPosition(0.0)
|
|
26
|
+
|
|
27
|
+
state = sampler.init(tlp, initial_position)
|
|
28
|
+
(chain_state, kernel_state), info = sampler.step(
|
|
29
|
+
target_log_prob_fn=tlp, current_state=state, seed=[0, 0]
|
|
30
|
+
)
|
|
31
|
+
|
|
32
|
+
tf.debugging.assert_equal(
|
|
33
|
+
chain_state.position, np.float32(multi_scan_iterations)
|
|
34
|
+
)
|
|
35
|
+
tf.debugging.assert_equal(
|
|
36
|
+
chain_state.log_density, np.float32(multi_scan_iterations)
|
|
37
|
+
)
|
|
38
|
+
tf.debugging.assert_equal(kernel_state.invocation, multi_scan_iterations)
|
|
39
|
+
tf.debugging.assert_equal(info, CountingKernelInfo(True))
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def test_many_multi_scan():
|
|
43
|
+
num_samples = 5
|
|
44
|
+
multi_scan_iterations = 100
|
|
45
|
+
|
|
46
|
+
sampler = multi_scan(multi_scan_iterations, counting_kernel())
|
|
47
|
+
|
|
48
|
+
def tlp(x):
|
|
49
|
+
return x
|
|
50
|
+
|
|
51
|
+
initial_position = TestPosition(0.0)
|
|
52
|
+
|
|
53
|
+
samples, info = mcmc(
|
|
54
|
+
num_samples=num_samples,
|
|
55
|
+
sampling_algorithm=sampler,
|
|
56
|
+
target_density_fn=tlp,
|
|
57
|
+
initial_position=initial_position,
|
|
58
|
+
seed=[0, 0],
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
tf.debugging.assert_equal(
|
|
62
|
+
samples,
|
|
63
|
+
tf.range(
|
|
64
|
+
100.0,
|
|
65
|
+
100.0 + (num_samples * multi_scan_iterations),
|
|
66
|
+
delta=multi_scan_iterations,
|
|
67
|
+
),
|
|
68
|
+
)
|
|
69
|
+
tf.debugging.assert_equal(
|
|
70
|
+
info, CountingKernelInfo(tf.fill(num_samples, True))
|
|
71
|
+
)
|
|
@@ -0,0 +1,107 @@
|
|
|
1
|
+
"""Implementation of the Random Walk Metropolis algorithm"""
|
|
2
|
+
|
|
3
|
+
from typing import Callable, NamedTuple, Tuple
|
|
4
|
+
|
|
5
|
+
import tensorflow_probability as tfp
|
|
6
|
+
|
|
7
|
+
from .mcmc_base import ChainState, SamplingAlgorithm
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class RwmhInfo(NamedTuple):
|
|
11
|
+
"""Represents the information about a random walk Metropolis-Hastings (RWMH)
|
|
12
|
+
step.
|
|
13
|
+
This can be expanded to include more information in the future (as needed
|
|
14
|
+
for a specific kernel).
|
|
15
|
+
|
|
16
|
+
Attributes
|
|
17
|
+
----------
|
|
18
|
+
is_accepted (bool): Indicates whether the proposal was accepted or not.
|
|
19
|
+
|
|
20
|
+
"""
|
|
21
|
+
|
|
22
|
+
is_accepted: bool
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class RwmhKernelState(NamedTuple):
|
|
26
|
+
"""Represents the kernel state of an arbitrary MCMC kernel.
|
|
27
|
+
|
|
28
|
+
Attributes
|
|
29
|
+
----------
|
|
30
|
+
scale (float): The scale parameter for the RWMH kernel.
|
|
31
|
+
|
|
32
|
+
"""
|
|
33
|
+
|
|
34
|
+
scale: float
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def rwmh(scale=1.0):
|
|
38
|
+
def _build_kernel(log_prob_fn):
|
|
39
|
+
"""Partial"""
|
|
40
|
+
return tfp.mcmc.RandomWalkMetropolis(
|
|
41
|
+
target_log_prob_fn=log_prob_fn,
|
|
42
|
+
new_state_fn=tfp.mcmc.random_walk_normal_fn(scale=scale),
|
|
43
|
+
)
|
|
44
|
+
|
|
45
|
+
def init_fn(target_log_prob_fn, target_state):
|
|
46
|
+
kernel = _build_kernel(target_log_prob_fn)
|
|
47
|
+
results = kernel.bootstrap_results(target_state)
|
|
48
|
+
|
|
49
|
+
chain_state = ChainState(
|
|
50
|
+
position=target_state,
|
|
51
|
+
log_density=results.accepted_results.target_log_prob,
|
|
52
|
+
log_density_grad=(),
|
|
53
|
+
)
|
|
54
|
+
kernel_state = RwmhKernelState(scale=scale)
|
|
55
|
+
|
|
56
|
+
return chain_state, kernel_state
|
|
57
|
+
|
|
58
|
+
def step_fn(
|
|
59
|
+
target_log_prob_fn: Callable[[NamedTuple], float],
|
|
60
|
+
target_and_kernel_state: Tuple[ChainState, RwmhKernelState],
|
|
61
|
+
seed,
|
|
62
|
+
) -> Callable[[ChainState], Tuple[ChainState, RwmhInfo]]:
|
|
63
|
+
"""Computation that calls a kernel.
|
|
64
|
+
|
|
65
|
+
Args:
|
|
66
|
+
----
|
|
67
|
+
target_log_prob_fn: the conditional log target density/mass function
|
|
68
|
+
conditioned_state: Parts of the global state that are not updated by
|
|
69
|
+
the kernel, but may be needed to instantiate it.
|
|
70
|
+
chain_and_kernel_state: a tuple containing a ChainState object and
|
|
71
|
+
kernel-specific state.
|
|
72
|
+
target_state: the sub-state on which the kernel operates.
|
|
73
|
+
|
|
74
|
+
Returns:
|
|
75
|
+
-------
|
|
76
|
+
a tuple of the new target sub-state and information about the sampler
|
|
77
|
+
|
|
78
|
+
"""
|
|
79
|
+
# This could be replaced with BlackJAX easily
|
|
80
|
+
kernel = _build_kernel(target_log_prob_fn)
|
|
81
|
+
|
|
82
|
+
target_chain_state, kernel_state = target_and_kernel_state
|
|
83
|
+
|
|
84
|
+
new_target_position, results = kernel.one_step(
|
|
85
|
+
target_chain_state.position,
|
|
86
|
+
kernel.bootstrap_results(target_chain_state.position),
|
|
87
|
+
seed=seed,
|
|
88
|
+
)
|
|
89
|
+
|
|
90
|
+
new_chain_and_kernel_state = (
|
|
91
|
+
ChainState(
|
|
92
|
+
position=new_target_position,
|
|
93
|
+
log_density=results.accepted_results.target_log_prob,
|
|
94
|
+
log_density_grad=(),
|
|
95
|
+
),
|
|
96
|
+
kernel_state,
|
|
97
|
+
)
|
|
98
|
+
|
|
99
|
+
info = (
|
|
100
|
+
RwmhInfo(
|
|
101
|
+
results.is_accepted,
|
|
102
|
+
),
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
return new_chain_and_kernel_state, info
|
|
106
|
+
|
|
107
|
+
return SamplingAlgorithm(init_fn, step_fn)
|