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 usesstate_space_mllas 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