blackjax.adaptation.laps_burn_in#

Classes#

Functions#

no_nans(a)

nan_reject(nonans, old, new)

Equivalent to

build_kernel(logdensity_fn, ndims[, microcanonical])

MCLMC kernel (with nan rejection)

initialize(rng_key, logdensity_fn, microcanonical, ...)

initialize the chains based on the equipartition of the initial condition.

update_history(new_vals, history)

update_history_scalar(new_val, history)

contract_history(theta, weights)

equipartition_diagonal(state)

Ei = E_ensemble (- grad log p_i x_i ). Ei is 1 if we have converged.

equipartition_fullrank(state, rng_key)

loss = Tr[(1 - E)^T (1 - E)] / d^2

equipartition_diagonal_loss(Eii)

equipartition_fullrank_loss(delta_z)

Module Contents#

no_nans(a)[source]#
nan_reject(nonans, old, new)[source]#

Equivalent to return new if nonans else old

build_kernel(logdensity_fn, ndims, microcanonical=True)[source]#

MCLMC kernel (with nan rejection)

initialize(rng_key, logdensity_fn, microcanonical, sample_init, num_chains, mesh, superchain_size)[source]#

initialize the chains based on the equipartition of the initial condition. We initialize the velocity along grad log p if E_ii > 1 and along -grad log p if E_ii < 1.

update_history(new_vals, history)[source]#
update_history_scalar(new_val, history)[source]#
contract_history(theta, weights)[source]#
class History[source]#
observables: blackjax.types.Array[source]#
stopping: blackjax.types.Array[source]#
weights: blackjax.types.Array[source]#
class AdaptationState[source]#
L: float[source]#
inverse_mass_matrix: Any[source]#
step_size: float[source]#
step_count: int[source]#
EEVPD: float[source]#
EEVPD_wanted: float[source]#
history: Any[source]#
equipartition_diagonal(state)[source]#

Ei = E_ensemble (- grad log p_i x_i ). Ei is 1 if we have converged. equipartition_loss = average over parameters (Ei)

equipartition_fullrank(state, rng_key)[source]#

loss = Tr[(1 - E)^T (1 - E)] / d^2 where Eij = <xi gj> is the equipartition patrix. Loss is computed with the Hutchinson’s trick.

equipartition_diagonal_loss(Eii)[source]#
equipartition_fullrank_loss(delta_z)[source]#
class Adaptation(ndims, microcanonical, alpha=1.0, C=0.1, r_end=0.01, bias_type=0, save_num=10, observables=lambda x: ..., observables_for_bias=lambda x: ..., contract=lambda x: ...)[source]#
ndims[source]#
alpha = 1.0[source]#
C = 0.1[source]#
r_end = 0.01[source]#
observables[source]#
observables_for_bias[source]#
contract[source]#
bias_type = 0[source]#
save_num = 10[source]#
norm_factor[source]#
initial_state[source]#
summary_statistics_fn(state, info, rng_key)[source]#
update(adaptation_state, Etheta)[source]#
while_cond(info, counter)[source]#

determine if we want to switch to adjustment