Skip to content

Releases: QuantClimate/GPJax

GPJax v1.0.0

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

GPJax v0.18.0

Choose a tag to compare

@github-actions github-actions released this 26 Jul 19:46
177dabb

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

What's Changed

New Contributors

Full Changelog: v0.17.0...v0.18.0

GPJax v0.17.0

Choose a tag to compare

@github-actions github-actions released this 04 Jul 12:58
e4572d7

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

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

Full Changelog: v0.15.0...v0.17.0

GPJax v0.15.0

Choose a tag to compare

@github-actions github-actions released this 03 Jun 17:12
7d68b58

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

What's Changed

New Contributors

Full Changelog: v0.14.0...v0.15.0

GPJax v0.14.0

Choose a tag to compare

@github-actions github-actions released this 21 Apr 07:42
ed4bfd3

Release v0.14.0

🐛 Bug Fixes

  • upper-bound jax and jaxlib to <0.10 (#627)

📝 Other Changes

  • Release v0.14.0 (#626)
  • Fix linting (#624)

📊 Performance & ML Improvements

  • See performance regression tests in this release
  • Model validation results available in CI artifacts

🔍 What's Changed

What's Changed

Full Changelog: v0.13.6...v0.14.0

GPJax v0.14.0-rc1

GPJax v0.14.0-rc1 Pre-release
Pre-release

Choose a tag to compare

@github-actions github-actions released this 13 Apr 21:23
85ab347

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

What's Changed

Full Changelog: v0.13.6...v0.14.0-rc1

GPJax v0.13.6

Choose a tag to compare

@github-actions github-actions released this 08 Feb 22:53
2349ce9

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

What's Changed

Full Changelog: v0.13.5...v0.13.6

GPJax v0.13.5

Choose a tag to compare

@github-actions github-actions released this 05 Feb 21:40
77abc37

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

What's Changed

Full Changelog: v0.13.4...v0.13.5

GPJax v0.13.4

Choose a tag to compare

@github-actions github-actions released this 23 Nov 21:37
7b33936

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

What's Changed

New Contributors

Full Changelog: v0.13.3...v0.13.4

GPJax v0.13.3

Choose a tag to compare

@github-actions github-actions released this 11 Nov 08:14
c5e1fd7

Release v0.13.3

📝 Other Changes

  • Bump version (#572)
  • Ability to return just diag on predict call (#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

What's Changed

New Contributors

Full Changelog: v0.13.2...v0.13.3