Releases: QuantClimate/GPJax
Release list
GPJax v1.0.0
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
GPJax v0.18.0
Release v0.18.0
🐛 Bug Fixes
- require the noise latent in heteroscedastic link_function (#706)
- parameterise the spectral measure by lengthscale and dimension (#702)
📝 Other Changes
- chore(release): prepare v0.18.0 (#711)
- perf(variational): compute the prior KL in closed form (#708)
- perf: reuse Cholesky factors in Gaussian KL divergence (#707)
- build: drop the unused tensorstore runtime dependency (#705)
- docs(kernels): fix RationalQuadratic and PoweredExponential formulas (#703)
- ci(ruff): pin ruff to the lockfile and stop CI auto-fixing (#704)
- ci(deps): bump actions/setup-node from 6 to 7 (#700)
- deps(deps): bump tqdm from 4.68.4 to 4.69.0 (#701)
- Fix runtime dependency metadata (#699)
- Add real data examples (#696)
- deps(deps): bump tqdm from 4.68.3 to 4.68.4 (#697)
- Fix Matern kernel docstring formulas (#690)
- deps(deps): bump absl-py from 2.4.0 to 2.5.0 (#694)
- deps(deps): bump typing-extensions from 4.15.0 to 4.16.0 (#693)
- docs: clarify collapsed elbo batching (#686)
- Consolidate the duplicated paramax-unwrap _val helper into gpjax.parameters (#692)
- Merge pull request #691 from thomaspinder/fix/skip-pages-deploy-for-forks
- ci: skip docs preview deploy for fork PRs
📊 Performance & ML Improvements
- See performance regression tests in this release
- Model validation results available in CI artifacts
🔍 What's Changed
- Full diff: v0.17.0...v0.18.0
What's Changed
- ci: skip docs preview deploy for fork PRs by @thomaspinder in #691
- Consolidate the duplicated paramax-unwrap _val helper into gpjax.parameters by @businessarshgoyal in #692
- docs: clarify collapsed elbo batching by @vku2018 in #686
- deps(deps): bump typing-extensions from 4.15.0 to 4.16.0 by @dependabot[bot] in #693
- deps(deps): bump absl-py from 2.4.0 to 2.5.0 by @dependabot[bot] in #694
- Fix Matérn kernel docstring formulas by @jordansilly77-stack in #690
- deps(deps): bump tqdm from 4.68.3 to 4.68.4 by @dependabot[bot] in #697
- Add real data examples by @thomaspinder in #696
- Fix runtime dependency declarations by @ShreyanshGoyal in #699
- deps(deps): bump tqdm from 4.68.4 to 4.69.0 by @dependabot[bot] in #701
- ci(deps): bump actions/setup-node from 6 to 7 by @dependabot[bot] in #700
- ci(ruff): pin ruff to the lockfile and stop CI auto-fixing by @thomaspinder in #704
- docs(kernels): fix RationalQuadratic and PoweredExponential formulas by @thomaspinder in #703
- fix(kernels): parameterise the spectral measure by lengthscale and dimension by @thomaspinder in #702
- build: drop the unused tensorstore runtime dependency by @thomaspinder in #705
- fix(likelihoods): require the noise latent in heteroscedastic link_function by @thomaspinder in #706
- perf: reuse Cholesky factors in Gaussian KL divergence by @thomaspinder in #707
- perf: use the closed-form KL in variational families by @thomaspinder in #708
- chore(release): prepare v0.18.0 by @thomaspinder in #711
New Contributors
- @vku2018 made their first contribution in #686
- @jordansilly77-stack made their first contribution in #690
- @ShreyanshGoyal made their first contribution in #699
Full Changelog: v0.17.0...v0.18.0
GPJax v0.17.0
Release v0.17.0
✨ New Features
- add gpx.summarise() model summary tables (#661)
🐛 Bug Fixes
- OILMM PCA data-init + diagonal predict operator (0.17.0)
- correctness fixes — sample_approx, Matern RFF, multi-output LOOCV, ArcCosine (0.17.0)
📝 Other Changes
- perf: diagonal predict returns DiagonalLinearOperator (0.17.0)
- Remove unreachable gpjax/linalg/_compat.py deprecation shim (#684)
- docs: update GraphKernel spectral docstring (#685)
- deps(deps): bump tqdm from 4.67.3 to 4.68.3 (#660)
- deps(deps): bump the jax-ecosystem group across 1 directory with 3 updates (#659)
- ci(deps): bump actions/checkout from 6 to 7 (#658)
- deps(deps): bump chex from 0.1.91 to 0.1.92 (#657)
- ci(deps): bump codecov/codecov-action from 6 to 7 (#653)
📊 Performance & ML Improvements
- See performance regression tests in this release
- Model validation results available in CI artifacts
🔍 What's Changed
- Full diff: v0.15.0...v0.17.0
What's Changed
- feat: add gpx.summarise() model summary tables by @thomaspinder in #661
- ci(deps): bump codecov/codecov-action from 6 to 7 by @dependabot[bot] in #653
- deps(deps): bump chex from 0.1.91 to 0.1.92 by @dependabot[bot] in #657
- ci(deps): bump actions/checkout from 6 to 7 by @dependabot[bot] in #658
- deps(deps): bump the jax-ecosystem group across 1 directory with 3 updates by @dependabot[bot] in #659
- deps(deps): bump tqdm from 4.67.3 to 4.68.3 by @dependabot[bot] in #660
- docs: update GraphKernel spectral docstring by @milekv in #685
- Remove unreachable gpjax/linalg/_compat.py deprecation shim by @businessarshgoyal in #684
- Correctness fixes: sample_approx, Matern RFF, multi-output LOOCV, ArcCosine (#662) by @thomaspinder in #687
- OILMM: PCA data-init + diagonal predict operator (0.17.0) by @thomaspinder in #688
- Perf: diagonal predict returns DiagonalLinearOperator (0.17.0) by @thomaspinder in #689
New Contributors
- @milekv made their first contribution in #685
- @businessarshgoyal made their first contribution in #684
Full Changelog: v0.15.0...v0.17.0
GPJax v0.15.0
Release v0.15.0
📝 Other Changes
- Add v0 state space work (#633)
- deps(deps): bump the ml-tools group with 2 updates (#650)
- deps(deps): bump the jax-ecosystem group with 3 updates (#649)
- deps(deps): bump the jax-ecosystem group across 1 directory with 5 updates (#641)
- deps(deps): bump tensorstore from 0.1.83 to 0.1.84 (#648)
- deps(deps): bump tqdm from 4.67.1 to 4.67.3 (#643)
- deps(deps): bump numpyro from 0.19.0 to 0.21.0 (#644)
- deps(deps): bump absl-py from 2.3.1 to 2.4.0 (#645)
- deps(deps): bump tensorstore from 0.1.80 to 0.1.83 (#646)
- deps(deps): bump lineax from 0.1.0 to 0.1.1 (#647)
- Extend benchmarks (#640)
- Reorder index (#639)
- Add ASV Benchmarking (#637)
- Kronecker mv fix (#634)
- ci(deps): bump conda-incubator/setup-miniconda from 3 to 4 (#631)
- deps(deps): bump the jax-ecosystem group with 2 updates (#632)
- ci(deps): bump actions/upload-pages-artifact from 4 to 5 (#625)
- Add call to action (#630)
- Drop explicit float64 type (#629)
📊 Performance & ML Improvements
- See performance regression tests in this release
- Model validation results available in CI artifacts
🔍 What's Changed
- Full diff: v0.14.0...v0.15.0
What's Changed
- Drop explicit float64 type by @thomaspinder in #629
- Add call to action by @thomaspinder in #630
- ci(deps): bump actions/upload-pages-artifact from 4 to 5 by @dependabot[bot] in #625
- deps(deps): bump the jax-ecosystem group with 2 updates by @dependabot[bot] in #632
- ci(deps): bump conda-incubator/setup-miniconda from 3 to 4 by @dependabot[bot] in #631
- Kronecker mv fix by @lugin100 in #634
- Add ASV Benchmarking by @thomaspinder in #637
- Reorder index by @thomaspinder in #639
- Extend benchmarks by @thomaspinder in #640
- deps(deps): bump lineax from 0.1.0 to 0.1.1 by @dependabot[bot] in #647
- deps(deps): bump tensorstore from 0.1.80 to 0.1.83 by @dependabot[bot] in #646
- deps(deps): bump absl-py from 2.3.1 to 2.4.0 by @dependabot[bot] in #645
- deps(deps): bump numpyro from 0.19.0 to 0.21.0 by @dependabot[bot] in #644
- deps(deps): bump tqdm from 4.67.1 to 4.67.3 by @dependabot[bot] in #643
- deps(deps): bump tensorstore from 0.1.83 to 0.1.84 by @dependabot[bot] in #648
- deps(deps): bump the jax-ecosystem group across 1 directory with 5 updates by @dependabot[bot] in #641
- deps(deps): bump the jax-ecosystem group with 3 updates by @dependabot[bot] in #649
- deps(deps): bump the ml-tools group with 2 updates by @dependabot[bot] in #650
- Add v0 state space work by @thomaspinder in #633
New Contributors
Full Changelog: v0.14.0...v0.15.0
GPJax v0.14.0
Release v0.14.0
🐛 Bug Fixes
- upper-bound jax and jaxlib to <0.10 (#627)
📝 Other Changes
📊 Performance & ML Improvements
- See performance regression tests in this release
- Model validation results available in CI artifacts
🔍 What's Changed
- Full diff: v0.14.0-rc1...v0.14.0
What's Changed
- ci(deps): bump actions/upload-pages-artifact from 3 to 4 by @dependabot[bot] in #603
- ci(deps): bump actions/upload-artifact from 4 to 6 by @dependabot[bot] in #602
- thomaspinder/multi output by @thomaspinder in #604
- Add LCM Kernel by @thomaspinder in #605
- Add OILMM update by @thomaspinder in #606
- Add OAK by @thomaspinder in #607
- ci(deps): bump actions/download-artifact from 7 to 8 by @dependabot[bot] in #611
- ci(deps): bump actions/upload-artifact from 6 to 7 by @dependabot[bot] in #610
- Add Variable subclass section by @thomaspinder in #613
- Fix docs by @thomaspinder in #616
- Fix docs by @thomaspinder in #617
- ci(deps): bump codecov/codecov-action from 5 to 6 by @dependabot[bot] in #619
- ci(deps): bump actions/deploy-pages from 4 to 5 by @dependabot[bot] in #620
- ci(deps): bump actions/github-script from 8 to 9 by @dependabot[bot] in #623
- ci(deps): bump softprops/action-gh-release from 2 to 3 by @dependabot[bot] in #622
- Test Equinox backend by @thomaspinder in #614
- Fix migration doc linting by @thomaspinder in #624
- Release v0.14.0 by @thomaspinder in #626
- fix: upper-bound jax and jaxlib to <0.10 by @thomaspinder in #627
Full Changelog: v0.13.6...v0.14.0
GPJax v0.14.0-rc1
Release v0.14.0-rc1
📝 Other Changes
- Add Equinox backend (#614)
- ci(deps): bump softprops/action-gh-release from 2 to 3 (#622)
- ci(deps): bump actions/github-script from 8 to 9 (#623)
- ci(deps): bump actions/deploy-pages from 4 to 5 (#620)
- ci(deps): bump codecov/codecov-action from 5 to 6 (#619)
- Fix docs (#617)
- Fix docs (#616)
- Add Variable subclass section (#613)
- ci(deps): bump actions/upload-artifact from 6 to 7 (#610)
- ci(deps): bump actions/download-artifact from 7 to 8 (#611)
- Add OAK (#607)
- Add OILMM update (#606)
- Add LCM Kernel (#605)
- thomaspinder/ICM Multi-Output Kernel
- ci(deps): bump actions/upload-artifact from 4 to 6 (#602)
- ci(deps): bump actions/upload-pages-artifact from 3 to 4 (#603)
📊 Performance & ML Improvements
- See performance regression tests in this release
- Model validation results available in CI artifacts
🔍 What's Changed
- Full diff: v0.13.6...v0.14.0-rc1
What's Changed
- ci(deps): bump actions/upload-pages-artifact from 3 to 4 by @dependabot[bot] in #603
- ci(deps): bump actions/upload-artifact from 4 to 6 by @dependabot[bot] in #602
- thomaspinder/multi output by @thomaspinder in #604
- Add LCM Kernel by @thomaspinder in #605
- Add OILMM update by @thomaspinder in #606
- Add OAK by @thomaspinder in #607
- ci(deps): bump actions/download-artifact from 7 to 8 by @dependabot[bot] in #611
- ci(deps): bump actions/upload-artifact from 6 to 7 by @dependabot[bot] in #610
- Add Variable subclass section by @thomaspinder in #613
- Fix docs by @thomaspinder in #616
- Fix docs by @thomaspinder in #617
- ci(deps): bump codecov/codecov-action from 5 to 6 by @dependabot[bot] in #619
- ci(deps): bump actions/deploy-pages from 4 to 5 by @dependabot[bot] in #620
- ci(deps): bump actions/github-script from 8 to 9 by @dependabot[bot] in #623
- ci(deps): bump softprops/action-gh-release from 2 to 3 by @dependabot[bot] in #622
- Test Equinox backend by @thomaspinder in #614
Full Changelog: v0.13.6...v0.14.0-rc1
GPJax v0.13.6
Release v0.13.6
📝 Other Changes
- bump version (#601)
- Feb 26' Audit (#600)
- NumPyro Integration (#583)
- bug: Fix combination kernel flattening (#599)
📊 Performance & ML Improvements
- See performance regression tests in this release
- Model validation results available in CI artifacts
🔍 What's Changed
- Full diff: v0.13.5...v0.13.6
What's Changed
- bug: Fix combination kernel flattening by @thomaspinder in #599
- NumPyro Integration by @thomaspinder in #583
- Feb 26' Audit by @thomaspinder in #600
- bump version by @thomaspinder in #601
Full Changelog: v0.13.5...v0.13.6
GPJax v0.13.5
Release v0.13.5
🐛 Bug Fixes
- adds marker to avoid problematic tensorstore in mac builds (#595)
📝 Other Changes
- Bump version (#598)
- Resolve parameter annotation (#597)
- build: Add 7-day lag on latest package versions (#596)
- ci(deps): bump JamesIves/github-pages-deploy-action from 4.7.6 to 4.8.0 (#590)
- ci(deps): bump actions/upload-artifact from 5 to 6 (#585)
- ci(deps): bump actions/download-artifact from 6 to 7 (#586)
- ci(deps): bump JamesIves/github-pages-deploy-action from 4.7.4 to 4.7.6 (#587)
- dev-deps(deps-dev): update mkdocstrings[python] requirement (#584)
- Merge pull request #582 from thomaspinder/dependabot/github_actions/actions/checkout-6
- ci(deps): bump actions/checkout from 5 to 6
- Merge pull request #581 from thomaspinder/thomaspinder/temp
- Improve sparse plotting
- Merge pull request #580 from thomaspinder:thomaspinder/docs-header
- Correct title
📊 Performance & ML Improvements
- See performance regression tests in this release
- Model validation results available in CI artifacts
🔍 What's Changed
- Full diff: v0.13.4...v0.13.5
What's Changed
- Correct title by @thomaspinder in #580
- Improve sparse plotting by @thomaspinder in #581
- ci(deps): bump actions/checkout from 5 to 6 by @dependabot[bot] in #582
- dev-deps(deps-dev): update mkdocstrings[python] requirement from <0.31.0 to <1.1.0 by @dependabot[bot] in #584
- ci(deps): bump JamesIves/github-pages-deploy-action from 4.7.4 to 4.7.6 by @dependabot[bot] in #587
- ci(deps): bump actions/download-artifact from 6 to 7 by @dependabot[bot] in #586
- ci(deps): bump actions/upload-artifact from 5 to 6 by @dependabot[bot] in #585
- ci(deps): bump JamesIves/github-pages-deploy-action from 4.7.6 to 4.8.0 by @dependabot[bot] in #590
- fix: adds marker to avoid problematic
tensorstorein mac builds by @miguelgondu in #595 - build: Add 7-day lag on latest package versions by @thomaspinder in #596
- Resolve parameter annotation by @thomaspinder in #597
- Bump version by @thomaspinder in #598
Full Changelog: v0.13.4...v0.13.5
GPJax v0.13.4
Release v0.13.4
📝 Other Changes
- Add Heteroscedastic likelihood (#579)
- Align parameter tagging with Flax conventions (#578)
- Limit dtype promotion in Constant mean (#573)
🔍 What's Changed
- Full diff: v0.13.3...v0.13.4
What's Changed
- Limit dtype promotion in Constant mean by @sethaxen in #573
- Align parameter tagging with Flax conventions by @thomaspinder in #578
- Add Heteroscedastic likelihood by @thomaspinder in #579
New Contributors
Full Changelog: v0.13.3...v0.13.4
GPJax v0.13.3
Release v0.13.3
📝 Other Changes
- Bump version (#572)
- Ability to return just diag on
predictcall (#567) - ci(deps): bump JamesIves/github-pages-deploy-action from 4.7.3 to 4.7.4 (#570)
- Create FUNDING.yml (#569)
- Update Gardeners (#568)
📊 Performance & ML Improvements
- See performance regression tests in this release
- Model validation results available in CI artifacts
🔍 What's Changed
- Full diff: v0.13.2...v0.13.3
What's Changed
- Update Gardeners by @thomaspinder in #568
- Create FUNDING.yml by @thomaspinder in #569
- ci(deps): bump JamesIves/github-pages-deploy-action from 4.7.3 to 4.7.4 by @dependabot[bot] in #570
- Ability to return just diag on
predictcall by @mathDR in #567 - Bump version by @thomaspinder in #572
New Contributors
Full Changelog: v0.13.2...v0.13.3