Skip to content

GPJax v1.0.0

Latest

Choose a tag to compare

@github-actions github-actions released this 28 Sep 08:39
e44bb7a

GPJax 1.0 is the first stable release. It rebuilds the core API around conditioning, so the code reads like the maths: a joint model is conditioned on data to give a posterior, and the posterior is queried directly. On top of that come natural-gradient training, the dual (t-SVGP) parameterisation of sparse GPs, xarray support for gridded data, a robust Student's t likelihood, and dense joint covariances for state-space GPs.

GPJax now lives at QuantClimate/GPJax, with documentation at gpjax.quantclimate.com. Old links to docs.jaxgaussianprocesses.com redirect permanently.

Important

1.0 contains breaking changes. Work through the 0.18 → 1.0 migration guide when upgrading. Most prediction code keeps working unchanged.

✨ Highlights

Conditioning is the core API

prior * likelihood now returns a JointModel, the joint $p(f, y)$ that gpx.fit trains. Conditioning it on data gives a Posterior that caches its factorisation and is queried directly:

model = prior * gpx.likelihoods.Gaussian()
posterior = model.condition(D)          # or: model | D
predictive = posterior(xtest)           # condition once, predict many times
marginals = posterior(xtest, covariance="diagonal")
evidence = posterior.log_marginal_likelihood
  • One contract for every process: exact (ExactPosterior), non-conjugate (LatentPosterior), sparse variational (SparsePosterior), collapsed (CollapsedPosterior), state-space (Kalman) and OILMM posteriors all condition and predict the same way, with a single covariance= keyword.
  • Variational families condition through the same machinery. q.condition(D) returns a Posterior, exactly like model.condition(D).
  • Pathwise sampling (sample_approx) now also works on sparse posteriors.
  • One jitter setting: Prior.jitter is applied exactly once, inside conditioning, via the new gpjax.linalg.stabilised_cholesky. Prediction and the marginal likelihood can no longer factorise different matrices.
  • Likelihoods are pure conditional distributions and no longer take num_datapoints. The minibatch ELBO scale comes from the dataset itself (Dataset.n_total), so it can't be set wrong.
  • The design is recorded in ADR-0001, and the vocabulary in CONTEXT.md.

Natural gradients: gpx.fit_natgrads

Trains a sparse variational GP by alternating a natural-gradient step on $q(u)$ with an Optax step on the hyperparameters and inducing inputs, following Salimbeni et al. (2018). It supports VariationalGaussian and WhitenedVariationalGaussian. On a conjugate full-batch problem, one step with natgrad_lr=1.0 reaches the optimal $q$.

Dual parameterisation (t-SVGP): DualVariationalGaussian and dual_elbo

The dual parameterisation of Adam et al. (2021) stores Gaussian sites instead of the moments of $q(u)$. A natural-gradient step then becomes a closed-form convex combination, and the M-step gets t-SVGP's better-behaved hyperparameter gradients. fit_natgrads dispatches on the family automatically.

Gridded data with xarray: gpjax.xarray

from_xarray(ds, target=..., inputs=[...]) flattens a labelled xr.Dataset into an ordinary Dataset. Inputs can be coordinates or data variables, datetimes are converted to days, and NaN cells are dropped. It also returns a GridSpec, which builds inputs for any new grid and maps predictions back onto it: means and variances, or joint posterior samples with a sample dimension. Install with pip install "gpjax[xarray]"; import gpjax never imports xarray.

Robust regression: gpx.likelihoods.StudentT

A heavy-tailed Student's t likelihood, so outliers pull the posterior mean less strongly. Its expected log-likelihood is computed by Gauss–Hermite quadrature.

State-space GPs

  • covariance="dense" on StateSpacePrior.predict, StateSpaceConjugateModel.predict and the smoothed StateSpacePosterior returns the full joint covariance across test points. It is built from the RTS smoother's cross-covariance recursion, so the cost stays linear in the number of training points.
  • TruncatedPeriodic × Matérn product kernels now have an SDE representation (quasi-periodic models).
  • rts_smoother is now a genuine QR-based square-root smoother.

Parameters: gpjax.parameters.val

val(param) is now public and the single way to read a parameter's constrained value. Plain arrays pass through unchanged, and frozen (paramax.non_trainable) parameters are read like any other. Models are no longer unwrapped behind the scenes, so custom kernels, means, likelihoods and objectives see the same model whether you call them directly or through fit.

🚨 Breaking changes

See the migration guide for before-and-after code. In brief:

0.18 1.0
ConjugatePosterior / NonConjugatePosterior ConjugateModel / NonConjugateModel
HeteroscedasticPosterior / ChainedPosterior HeteroscedasticModel (takes noise_prior= directly)
construct_posterior construct_model
AbstractPosterior split into JointModel and Posterior
StateSpaceConjugatePosterior StateSpaceConjugateModel
posterior.predict(x, D) model.condition(D)(x) (the old call still works)
return_covariance_type= covariance=
Gaussian(num_datapoints=n, ...) Gaussian(...), and likewise for every likelihood
HeteroscedasticGaussian(noise_prior=...) the noise prior moves onto HeteroscedasticModel
variational family posterior=, jitter= model=; set jitter on the Prior
NaturalVariationalGaussian, ExpectationVariationalGaussian removed; use VariationalGaussian with fit_natgrads
OILMMModel.condition_on_observations(D) model.condition(D) (the old name is deprecated, not removed)
OILMMPosterior default covariance "dense" "diagonal"
custom code reading kernel.lengthscale directly wrap each read in val(...)

NonConjugateModel also sizes its latent vector lazily on first fit; call model.init_latent(n) to use it earlier.

🐛 Fixes

  • The default Zero() mean no longer drifts away from zero during fitting. This was a regression introduced by the Equinox migration (#712).
  • gpx.kernels.RBF() and friends now type-check with no arguments, and White no longer carries an unused trainable lengthscale (#695).
  • Random Fourier feature kernels no longer compute the gram matrix twice.
  • BlockDiag and Kronecker operators correctly report unit diagonals to Lineax.
  • Noise-free Gaussian likelihoods (obs_stddev=0.0) work with numpyro 0.22 (#785).

📚 Documentation

🏗️ Project

  • The repository moved to the QuantClimate organisation.
  • Releases are published to PyPI with trusted publishing, so no long-lived tokens are stored.
  • Development dependencies moved to PEP 735 dependency groups, and GitHub Actions are pinned and audited with zizmor.

🙌 Contributors

Thanks to everyone who contributed to 1.0: @thomaspinder, @Thomas-Christie and @stephen-huan.

Full changelog: CHANGELOG.md · v0.18.0...v1.0.0