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.