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 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 singlecovariance=keyword. - Variational families condition through the same machinery.
q.condition(D)returns aPosterior, exactly likemodel.condition(D). - Pathwise sampling (
sample_approx) now also works on sparse posteriors. - One jitter setting:
Prior.jitteris applied exactly once, inside conditioning, via the newgpjax.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 VariationalGaussian and WhitenedVariationalGaussian. On a conjugate full-batch problem, one step with natgrad_lr=1.0 reaches the optimal
Dual parameterisation (t-SVGP): DualVariationalGaussian and dual_elbo
The dual parameterisation of Adam et al. (2021) stores Gaussian sites instead of the moments of 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"onStateSpacePrior.predict,StateSpaceConjugateModel.predictand the smoothedStateSpacePosteriorreturns 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érnproduct kernels now have an SDE representation (quasi-periodic models).rts_smootheris 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, andWhiteno longer carries an unused trainable lengthscale (#695).- Random Fourier feature kernels no longer compute the gram matrix twice.
BlockDiagandKroneckeroperators correctly report unit diagonals to Lineax.- Noise-free Gaussian likelihoods (
obs_stddev=0.0) work with numpyro 0.22 (#785).
📚 Documentation
- The docs moved from MkDocs to Sphinx and now live at gpjax.quantclimate.com. Old MkDocs-era URLs still resolve.
- New notebooks: Natural Gradients, Natural Gradients in Practice, Dual Parameterisation of Sparse GPs (t-SVGP) and Gridded Data with xarray.
- A migration guide, and a Sharp Bits section explaining how the parameter system works.
🏗️ 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