blackjax.mcmc.coupled_hmc#
Two HMC marginals coupled only through their random inputs.
The pair consists of ordinary HMC transitions with one shared normal innovation and one shared uniform variate. Each marginal keeps its own transition, acceptance probability, decision, and cached state. This module does not claim product-target invariance, meeting, unbiased estimators, or an efficiency improvement.
The contract#
Given the same (state, standard_normal, uniform) triple and per-marginal
parameters, each marginal performs the same transition as ordinary HMC. A
shared uniform is compared against each marginal’s own acceptance probability;
it is not a shared decision.
Mathematical semantics#
Writing \(p = A z\) for the momentum and \(A A^\top = M\) for its covariance, both couplings draw one \(z \sim N(0,I)\) and hand each marginal a transformed copy:
synchronousBoth marginals receive the same \(z\).
reflectionThe first marginal receives \(z\); the second receives
\[z' = z - 2 e (e^\top z)\]For unit
eindependent ofz, this orthogonal reflection preserves the normal law exactly in mathematics (floating-point normalization is accurate to rounding).ecomes fromdirection_fn, fixed at build time and applied to the incoming states and first metric only; it receives no innovations and must be pure in those arguments. The default direction is the difference of the two positions whitened by the first marginal’s metric, \(e \propto A^{-1}(x_1 - x_2)\). That particular choice is a heuristic. With different metrics it remains only a first-metric-based direction with no claimed property for the pair, and a zero direction is the identity.
Relation to blackjax.hmc#
Each marginal accepts with probability \(\min(1,e^{\Delta})\), comparing the supplied uniform directly. This is distributionally equivalent to ordinary HMC’s Bernoulli implementation, but not draw-for-draw equivalent.
Usage#
There is no top-level blackjax.coupled_hmc; use
blackjax.mcmc.coupled_hmc. Every per-marginal argument is an explicit
(first, second) pair, including shared values as (value, value);
nothing is broadcast or inferred.
import blackjax.mcmc.coupled_hmc as coupled_hmc
algorithm = coupled_hmc.as_top_level_api(
(logdensity_fn, logdensity_fn),
step_size=(0.1, 0.1),
inverse_mass_matrix=(inverse_mass_matrix, inverse_mass_matrix),
num_integration_steps=(10, 10),
coupling="reflection",
)
state = algorithm.init((position_one, position_two))
state, info = algorithm.step(rng_key, state)
Classes#
Pair of HMC states, each with its own caches. |
|
Per-marginal transition information and the shared random inputs. |
Functions#
|
Eagerly validate one marginal's concrete metric and settings. |
|
Initialize a coupled pair from explicit position and log-density pairs. |
|
Build a coupled HMC kernel. |
|
Build a |
Module Contents#
- class CoupledHMCState[source]#
Pair of HMC states, each with its own caches.
Each marginal retains its own position and cached
logdensity/logdensity_grad.- first
State of the first marginal chain.
- second
State of the second marginal chain.
- class CoupledHMCInfo[source]#
Per-marginal transition information and the shared random inputs.
firstandsecondretain complete, separate HMC diagnostics.- first
Transition information for the first marginal.
- second
Transition information for the second marginal.
- common_normal
Standard normal handed to the first marginal; the second receives its synchronous or reflected image.
- reflection_unit
Reflection direction in the same coordinates, or exactly zero for synchronous/identity coupling.
- uniform
Uniform shared by both tests; each marginal compares it with its own acceptance probability.
- validate_marginal_inputs(inverse_mass_matrix, step_size, num_integration_steps)[source]#
Eagerly validate one marginal’s concrete metric and settings.
Checks use NumPy on concrete host values and are absent from the traced kernel. Passing them validates only the values supplied now; callers using traced values through
build_kernel()retain responsibility for them.- Raises:
TypeError, ValueError – If the metric is not a supported kind, is not finite, is not positive definite, or the integration settings are not positive and finite.
- init(position: blackjax.base.Position, logdensity_fn: Sequence[Callable]) CoupledHMCState[source]#
Initialize a coupled pair from explicit position and log-density pairs.
- Parameters:
position – An explicit
(first_position, second_position)pair. The two positions must share a pytree structure, leaf shapes and one floating dtype.logdensity_fn – An explicit
(first_logdensity_fn, second_logdensity_fn)pair. The two may target different distributions.
- build_kernel(integrator: Callable = integrators.velocity_verlet, divergence_threshold: float = 1000, *, coupling: str = 'synchronous', direction_fn: Callable | None = None)[source]#
Build a coupled HMC kernel.
couplinganddirection_fnare fixed here, when the kernel is built, and not at call time.direction_fnis applied to the pair of incoming states and the first metric, and is supplied no innovations.The kernel draws
zand the uniform before the transition evaluatesdirection_fn; the contract concerns its arguments, not ordering. A callable that closes over innovations or draws randomness breaks reflection marginal correctness and cannot be detected here, so callers must keep it pure in the supplied arguments.- Parameters:
integrator – Symplectic integrator used by both marginals.
divergence_threshold – Energy difference above which a marginal transition is flagged divergent. Applied to each marginal separately.
coupling –
"synchronous", which gives both marginals the same standard normal, or"reflection", which gives the second marginal a reflected copy.direction_fn – Only for
"reflection". A callable(first_state, second_state, first_metric) -> Arrayreturning a flat direction; defaults towhitened_difference(). Passing one under synchronous coupling is an error, since it would have no effect.
- Returns:
A kernel ``(rng_key, state, logdensity_fn, step_size,
inverse_mass_matrix, num_integration_steps) -> (CoupledHMCState,
CoupledHMCInfo)`` in which every per-marginal parameter is an explicit
(first, second)pair.
- as_top_level_api(logdensity_fn: Sequence[Callable], step_size: Sequence[float], inverse_mass_matrix: Sequence[blackjax.mcmc.metrics.MetricTypes], num_integration_steps: Sequence[int], *, coupling: str = 'synchronous', direction_fn: Callable | None = None, integrator: Callable = integrators.velocity_verlet, divergence_threshold: float = 1000) blackjax.base.SamplingAlgorithm[source]#
Build a
SamplingAlgorithmfor a coupled HMC pair.There is deliberately no top-level
blackjax.coupled_hmc; reach this throughblackjax.mcmc.coupled_hmc.as_top_level_api.Every per-marginal argument is an explicit
(first, second)pair, so a shared value is written twice. The two marginals may target different distributions and use different fixed Gaussian Euclidean metrics, step sizes and integration counts, but their positions must share a pytree structure, leaf shapes and one floating dtype.- Parameters:
logdensity_fn – Pair of log-density functions.
step_size – Pair of integration step sizes.
inverse_mass_matrix – Pair of inverse mass matrices. Each is a diagonal array, a dense array, or a
LowRankInverseMassMatrix; callable (Riemannian) metrics and pre-builtMetricobjects are not supported.num_integration_steps – Pair of integration step counts.
coupling –
"synchronous"or"reflection"; seebuild_kernel().direction_fn – Reflection direction policy; see
build_kernel().integrator – Symplectic integrator used by both marginals.
divergence_threshold – Per-marginal divergence threshold.
Notes
Each marginal’s concrete metric and integration settings are checked here by
validate_marginal_inputs(). Those checks read concrete values on the host, so they cover the arguments given here and say nothing about traced values appearing later underjitorvmap. Callers who need to supply traced metrics usebuild_kernel()directly, which performs no eager numerical validation.- Return type:
A
SamplingAlgorithmwhoseinittakes a pair of positions.