fit_lbfgs#
- gpjax.state_space.fit_lbfgs(*, model, train_data, observation_mask=None, max_iters=500, safe=True)[source]#
Fit a state-space posterior with Optax’s L-BFGS (
while_loopdriver).Thin wrapper around
gpx.fit_lbfgs.Example
>>> import jax.numpy as jnp >>> import gpjax as gpx >>> from gpjax.state_space import StateSpacePrior, fit_lbfgs >>> X = jnp.linspace(0.0, 5.0, 20).reshape(-1, 1) >>> y = jnp.sin(X) >>> prior = StateSpacePrior( ... mean_function=gpx.mean_functions.Zero(), ... kernel=gpx.kernels.Matern32(lengthscale=1.0, variance=1.0), ... ) >>> likelihood = gpx.likelihoods.Gaussian(obs_stddev=0.1) >>> posterior = prior * likelihood >>> fitted, history = fit_lbfgs( ... model=posterior, ... train_data=gpx.Dataset(X=X, y=y), ... max_iters=2, ... ) >>> fitted is not None True
Expand for references to
gpjax.state_space.fit_lbfgs