Use the engine#
Use an EngineBuilder when you need an explicit schedule
or want to add kernels yourself. For routine runs, start with
MCMC sampling.
import jax.numpy as jnp
import numpy as np
import tensorflow_probability.substrates.jax.bijectors as tfb
import tensorflow_probability.substrates.jax.distributions as tfd
import liesel.goose as gs
import liesel.model as lsl
Prepare the example#
These examples require the regression model, including its transformation
and inference specifications, from Sample your first posterior.
Build a schedule#
For the model from Sample your first posterior:
builder = gs.LieselMCMC(model).get_engine_builder(seed=2026, num_chains=4)
builder.positions_included = ["sigma_sq"]
builder.add_adaptation(1000)
builder.add_burnin(200)
builder.add_posterior(2000)
builder.show_progress = False
engine = builder.build()
liesel.goose.engine - INFO - Initializing kernels...
liesel.goose.engine - INFO - Done
engine.sample_all_epochs()
results = engine.get_results()
liesel.goose.engine - INFO - Finished warmup
get_engine_builder supplies the model interface, initial state, kernels,
and configured jitter. The remaining calls choose the schedule, construct
the engine, run it, and retrieve its results.
Customize adaptation#
Use a new builder to change the fast and slow adaptation windows:
builder = gs.LieselMCMC(model).get_engine_builder(seed=2026, num_chains=4)
builder.add_adaptation(
1000,
init=100,
term=100,
base=50,
)
builder.add_posterior(1000)
builder.show_progress = False
engine = builder.build()
liesel.goose.engine - INFO - Initializing kernels...
liesel.goose.engine - INFO - Done
engine.sample_all_epochs()
results = engine.get_results()
liesel.goose.engine - INFO - Finished warmup
Integer init and term values specify numbers of iterations; floats
specify fractions of the adaptation duration. base sets the first slow
window. See add_adaptation() for the schedule.
Add adaptation before burnin and posterior sampling.
Supply a log density#
Goose can also sample a log density represented without a Liesel graph. This independent example targets a standard normal variable:
def log_prob(state):
return -0.5 * jnp.square(state["x"]).sum()
builder = gs.EngineBuilder(seed=2026, num_chains=4)
builder.set_model(gs.DictInterface(log_prob))
builder.set_initial_values({"x": jnp.array(0.0)})
builder.add_kernel(gs.NUTSKernel(["x"]))
builder.add_adaptation(1000)
builder.add_posterior(1000)
builder.show_progress = False
engine = builder.build()
liesel.goose.builder - WARNING - No jitter functions provided. The initial values won't be jittered
liesel.goose.engine - INFO - Initializing kernels...
liesel.goose.engine - INFO - Done
engine.sample_all_epochs()
normal_results = engine.get_results()
liesel.goose.engine - INFO - Finished warmup
Here set_initial_values broadcasts one state to all chains. For different
starts, pass state leaves with a leading chain axis and
multiple_chains=True. The interface provides log-density evaluation and
position updates. See ModelInterface for the protocol.
The same engine can use other representations through
ModelInterface. For custom update logic, start with
an MH proposal before implementing a
complete kernel class.