fit_scipy#

gpjax.state_space.fit_scipy(*, model, train_data, observation_mask=None, max_iters=500, verbose=True, safe=True)[source]#

Fit a state-space posterior with SciPy’s L-BFGS-B.

Thin wrapper around gpx.fit_scipy. Validates data, sorts if necessary, and uses state_space_mll as the objective.

Example

>>> import jax.numpy as jnp
>>> import gpjax as gpx
>>> from gpjax.state_space import StateSpacePrior, fit_scipy
>>> 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_scipy(
...     model=posterior,
...     train_data=gpx.Dataset(X=X, y=y),
...     max_iters=2,
...     verbose=False,
... )
>>> bool(jnp.all(jnp.isfinite(history)))
True

Expand for references to gpjax.state_space.fit_scipy

fit_scipy

State-Space (Markovian) Gaussian Processes

Parameters: