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_loop driver).

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(num_datapoints=20, 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

fit_lbfgs

Parameters: