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,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)