Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions .github/labels.yml
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,9 @@
#
# The repository labels will be automatically configured using this file and
# the GitHub Action https://github.com/marketplace/actions/github-labeler.
- name: 0.14-migration
description: Issues related to upgrading from 0.13.x to 0.14.x
color: f9d0c4
- name: automated
description: Automated PRs
color: 000000
Expand Down
65 changes: 65 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
# Changelog

All notable changes to GPJax are documented here.

The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.1.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).

## [0.14.0rc1] — 2026-04-13

See [`docs/migration.md`](docs/migration.md) for full upgrade instructions.

### Changed (breaking)

- **Module backend**: `flax.nnx.Module` has been replaced with
`equinox.Module`. Custom kernels, mean functions, likelihoods, and
variational families that subclassed `nnx.Module` must now subclass
`eqx.Module` and declare class-level type annotations for every field.
- **Parameter system**: `PositiveReal`, `NonNegativeReal`, `Real`,
`SigmoidBounded`, and `LowerTriangular` now inherit from
`paramax.AbstractUnwrappable`. Read a parameter value with
`param.unwrap()`, or resolve an entire model tree with
`paramax.unwrap(model)`. `param.value` no longer exists.
- **Fitting API**: `gpx.fit`, `gpx.fit_scipy`, and `gpx.fit_lbfgs` no
longer accept `params_bijection=` or `trainable=`. Bijection handling
is automatic; freeze parameters by wrapping them with
`paramax.non_trainable` via `eqx.tree_at`.
- **`LowerTriangular`** now requires a strictly positive diagonal
(parameterised by `softplus_lower_cholesky`). Passing a matrix with
zero or negative diagonal entries produces `NaN` during construction.
- **Linear algebra**: `gpjax.linalg` has been rewritten on top of
[Lineax](https://docs.kidger.site/lineax/). Kernel `gram()` returns a
`lineax.AbstractLinearOperator` instead of a `cola.LinearOperator`.
`PSD`, `psd`, `Dense`, `Diagonal`, `Identity`, `Triangular`,
`LinearOperator`, `diag`, `solve`, and `lower_cholesky` are removed;
see the migration guide for Lineax equivalents.
- **Bijectors**: the custom bijector stack has been removed.
`Parameter` (the old generic base class), `DEFAULT_BIJECTION`,
`transform(...)`, and `FillTriangularTransform` are gone. Custom
parameter classes should use `numpyro.distributions.biject_to` with a
`numpyro.distributions.constraints` object.
- **`register_parameters` removed**. With the Equinox backend, GPJax
identifies parameter classes through `isinstance` checks on
`AbstractUnwrappable`, so the decorator is no longer needed.

### Added

- `gpjax.linalg.cholesky_factor` — single-dispatch Cholesky factoriser
that returns a lower-triangular Lineax operator.
- `gpjax.linalg.add_jitter` — helper for adding a jitter term to a
covariance operator.
- [`docs/migration.md`](docs/migration.md) — upgrade guide from 0.13.x.

### Removed

- Flax is no longer a runtime dependency. The package `__description__`
has been updated from "Gaussian processes in JAX and Flax" to
"Gaussian processes in JAX".
- `cola-ml` is no longer a runtime dependency.

### Dependencies

- Added: `equinox>=0.11`, `paramax>=0.0.5`, `lineax`.
- Removed: `flax`, `cola-ml`.

[0.14.0rc1]: https://github.com/thomaspinder/GPJax/releases/tag/v0.14.0rc1
24 changes: 11 additions & 13 deletions CLAUDE.md
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ This file provides guidance to Claude Code (claude.ai/code) when working with co

## What is GPJax?

GPJax is a Gaussian process library built on JAX. The API mirrors GP mathematics: you compose a `Prior` (kernel + mean function), multiply by a `Likelihood` to get a `Posterior`, then optimise an objective. All modules are Flax NNX `nnx.Module` subclasses, making them JAX-pytree-compatible for `jit`, `vmap`, and `grad`.
GPJax is a Gaussian process library built on JAX. The API mirrors GP mathematics: you compose a `Prior` (kernel + mean function), multiply by a `Likelihood` to get a `Posterior`, then optimise an objective. All modules are Equinox `eqx.Module` subclasses, making them JAX-pytree-compatible for `jit`, `vmap`, and `grad`.

## Commands

Expand Down Expand Up @@ -47,19 +47,17 @@ Prior(kernel, mean_function) * Likelihood --> Posterior

### Parameter system (`gpjax/parameters.py`)

Parameters are `nnx.Variable` subclasses with a `tag` string that selects the bijection for unconstrained optimisation:
Parameters are `paramax.AbstractUnwrappable` subclasses. Each class stores its value in an unconstrained internal field and implements `unwrap()` to apply the constraining bijection:

| Class | Tag | Bijection |
| Class | Bijection | Internal storage |
|---|---|---|
| `Real` | `"real"` | Identity |
| `PositiveReal` | `"positive"` | Softplus |
| `NonNegativeReal` | `"non_negative"` | Softplus |
| `SigmoidBounded` | `"sigmoid"` | Sigmoid |
| `LowerTriangular` | `"lower_triangular"` | FillTriangular |
| `Real` | Identity | `value` (unchanged) |
| `PositiveReal` | Softplus | `_unconstrained` via `inv_softplus` |
| `NonNegativeReal` | Softplus | `_unconstrained` via `inv_softplus` |
| `SigmoidBounded` | Sigmoid scaled to `[low, high]` | `_unconstrained` via `logit` |
| `LowerTriangular` | Fill-triangular | `_flat` vector |

Access the raw value with `param[...]` (the `__getitem__` ellipsis convention from Flax NNX). The `transform()` function maps entire `nnx.State` trees between constrained and unconstrained spaces using `DEFAULT_BIJECTION`.

**Flax NNX Variable metadata**: `Variable` uses `__slots__`, so metadata must be set via `self.set_metadata(key=value)` and read via `self.get_metadata("key", default)`.
Access the constrained value with `param.unwrap()`. To unwrap an entire model tree, use `paramax.unwrap(model)` which recursively resolves all `AbstractUnwrappable` leaves. To freeze parameters, wrap them with `paramax.non_trainable(param)`.

### Kernel system (`gpjax/kernels/`)

Expand All @@ -69,7 +67,7 @@ Kernel categories: `stationary/` (RBF, Matern12/32/52, Periodic, etc.), `nonstat

### Linear algebra (`gpjax/linalg/`)

Custom `LinearOperator` hierarchy: `Dense`, `Diagonal`, `Triangular`, `Identity`, `BlockDiag`, `Kronecker`. The `psd()` wrapper marks an operator as PSD. Key operations: `lower_cholesky()`, `solve()`, `logdet()`, `diag()`. These wrap JAX linalg but provide custom gradient-friendly implementations.
Built on [Lineax](https://docs.kidger.site/lineax/). Kernel `gram()` returns `lx.AbstractLinearOperator` (typically `lx.MatrixLinearOperator`). Use `.as_matrix()` to materialise. Custom operators: `BlockDiag`, `Kronecker`. Key utilities: `cholesky_factor()` (singledispatch, returns lower-triangular operator), `logdet()`, `add_jitter()`. Linear solves use `lx.linear_solve()`.

### Objectives (`gpjax/objectives.py`)

Expand All @@ -83,7 +81,7 @@ Optimise by negating: `nmll = lambda p, d: -conjugate_mll(p, d)`

### Fitting (`gpjax/fit.py`)

Three optimisers: `fit()` (Optax gradient descent with scan), `fit_scipy()` (SciPy L-BFGS-B), `fit_lbfgs()` (Optax L-BFGS with `while_loop`). All handle the constrained/unconstrained bijection automatically via `nnx.split`/`nnx.merge`.
Three optimisers: `fit()` (Optax gradient descent with scan), `fit_scipy()` (SciPy L-BFGS-B), `fit_lbfgs()` (Optax L-BFGS with `while_loop`). All handle the constrained/unconstrained bijection automatically: `paramax.unwrap(model)` is called inside the loss function, and `eqx.partition`/`eqx.combine` with `eqx.is_array` manage trainable vs static parts.

### Variational inference (`gpjax/variational_families.py`)

Expand Down
2 changes: 1 addition & 1 deletion docs/design.md
Original file line number Diff line number Diff line change
Expand Up @@ -35,4 +35,4 @@ Prior to building GPJax, the developers of GPJax have benefited greatly from the
[GPyTorch](https://github.com/cornellius-gp/gpytorch) packages. As such, many of
the design principles in GPJax are inspired by the excellent precursory
packages. Documentation designs have been greatly inspired by the exceptional
[Flax docs](https://flax.readthedocs.io/en/latest/index.html).
[Equinox docs](https://docs.kidger.site/equinox/).
234 changes: 234 additions & 0 deletions docs/migration.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,234 @@
# Migration guide: `0.13.x` → `0.14.0`

GPJax `0.14` replaces the Flax NNX backend with
[Equinox](https://docs.kidger.site/equinox/) +
[paramax](https://danielward27.github.io/paramax/), and introduces a linear-algebra layer via
[Lineax](https://docs.kidger.site/lineax/). It also removes the custom bijector
stack in favour of
[numpyro constraints](https://num.pyro.ai/en/stable/distributions.html#constraints).

The changes are mostly internal. They surface in three places:

1. How you **define custom modules and custom parameter classes**.
2. How you **read a parameter value** (`param.unwrap()` / `paramax.unwrap(model)`
instead of `param.value`).
3. How you **freeze parameters** (`paramax.non_trainable(...)` instead of the
`trainable=` filter argument on `fit`).

If you only use the high-level API (`gpx.Prior`, `gpx.Posterior`, `gpx.fit`,
etc.) most code keeps working once you update the two or three call sites
below.

## Installation

```bash
pip install "gpjax==0.14.0rc1"
# or
uv add "gpjax==0.14.0rc1"
```

New dependencies (pulled in automatically): `equinox>=0.11`, `paramax>=0.0.5`.
Flax is no longer a runtime dependency.

## Breaking changes

### 1. Backend: `flax.nnx.Module` → `equinox.Module`

If you subclassed `nnx.Module` to build a custom model, kernel, mean function,
likelihood, or variational family, change the base class:

```py
# Before (0.13.x)
from flax import nnx

class MyKernel(nnx.Module):
def __init__(self, lengthscale):
self.lengthscale = gpx.parameters.PositiveReal(lengthscale)
```

```py
# After (0.14.0)
import equinox as eqx

class MyKernel(eqx.Module):
lengthscale: gpx.parameters.PositiveReal

def __init__(self, lengthscale):
self.lengthscale = gpx.parameters.PositiveReal(lengthscale)
```

Equinox requires **class-level field annotations** for every attribute, and
static configuration fields should be marked with `eqx.field(static=True)`.

### 2. Parameter classes are now `paramax.AbstractUnwrappable`

`PositiveReal`, `NonNegativeReal`, `Real`, `SigmoidBounded`, and
`LowerTriangular` all live in `gpjax.parameters` with the same names, but they
now inherit from `paramax.AbstractUnwrappable` and store their value in an
**unconstrained** internal field. The constraining bijection is applied at
read time through `unwrap()`.

```py
# Before (0.13.x) — nnx.Variable-based
length = gpx.parameters.PositiveReal(0.5)
length.value # -> 0.5
length.value = 1.0 # in-place mutation (nnx)
```

```py
# After (0.14.0) — paramax.AbstractUnwrappable
length = gpx.parameters.PositiveReal(0.5)
length.unwrap() # -> 0.5 (applies softplus to the stored unconstrained value)

# Unwrap an entire model tree in one call:
import paramax
model_resolved = paramax.unwrap(model)
```

`Parameter` (the old generic base class), `DEFAULT_BIJECTION`, the
`transform(...)` function, and `FillTriangularTransform` have been
**removed**. `numpyro.distributions.biject_to` now handles every
constraint → bijection mapping.

### 3. `LowerTriangular` now requires a valid Cholesky factor

Previously `LowerTriangular` accepted any lower-triangular matrix (the
diagonal was unconstrained). It is now parameterised via
`numpyro.distributions.constraints.softplus_lower_cholesky`, so the diagonal
**must be strictly positive**.

- Passing a matrix with zero or negative diagonal entries produces `NaN`
during construction (`inv_softplus` of a non-positive number).
- In-library usage is unaffected: the only consumer is
`VariationalGaussian.variational_root_covariance`, which is initialised to
the identity by default.
- If you previously supplied a custom `variational_root_covariance`, ensure
its diagonal is strictly positive. Under the old parameterisation, zero or
negative diagonals produced singular or sign-ambiguous variational
covariances.

### 4. `gpx.fit` / `fit_scipy` / `fit_lbfgs`: removed `params_bijection` and `trainable`

Bijection handling is now automatic via `paramax.unwrap` inside the loss
function, and freezing parameters is expressed by wrapping them in
`paramax.non_trainable`:

```py
# Before (0.13.x)
opt_model, history = gpx.fit(
model=posterior,
objective=gpx.objectives.conjugate_mll,
train_data=D,
optim=ox.adam(1e-2),
params_bijection=gpx.parameters.DEFAULT_BIJECTION,
trainable=gpx.parameters.Parameter, # filter-based trainability
)
```

```py
# After (0.14.0)
import paramax

# Freeze specific parameters up-front by wrapping them:
posterior = eqx.tree_at(
lambda m: m.prior.kernel.lengthscale,
posterior,
replace_fn=paramax.non_trainable,
)

opt_model, history = gpx.fit(
model=posterior,
objective=gpx.objectives.conjugate_mll,
train_data=D,
optim=ox.adam(1e-2),
)
```

Internally, `fit` now splits the model with `eqx.partition(model, eqx.is_array)`
so only concrete JAX arrays participate in the gradient update; everything
wrapped in `paramax.non_trainable` is held constant.

### 5. `register_parameters` removed

The `gpx.parameters.register_parameters` decorator (added in 0.13.x to mark
NNX variables as GPJax parameters) is gone. With Equinox, GPJax identifies
parameter classes through `isinstance` checks on `AbstractUnwrappable`, so
registration is unnecessary.

### 6. `gpjax.linalg` rewrite: cola → Lineax

Kernel `gram()` now returns a `lineax.AbstractLinearOperator`
(typically `lineax.MatrixLinearOperator`) instead of a `cola.LinearOperator`.
Materialise with `.as_matrix()`.

The following names have been **removed** from `gpjax.linalg`:

| Removed | Replacement |
| --------------------------- | -------------------------------------------------- |
| `PSD`, `psd` | Not needed — Lineax operators carry tags directly. |
| `Dense`, `Diagonal`, `Identity`, `Triangular` | `lineax.MatrixLinearOperator`, `lineax.DiagonalLinearOperator`, `lineax.IdentityLinearOperator`, `lineax.TriangularLinearOperator` |
| `LinearOperator` | `lineax.AbstractLinearOperator` |
| `diag`, `solve` | `operator.diagonal()`, `lineax.linear_solve(...)` |
| `lower_cholesky` | `gpjax.linalg.cholesky_factor` (singledispatch, returns a lower-triangular operator) |

`BlockDiag`, `Kronecker`, and `logdet` are unchanged. Use
`gpjax.linalg.add_jitter` to add a jitter term to a covariance operator.

### 7. Custom bijectors replaced with numpyro constraints

If you had a custom `Parameter` subclass that declared a bijection, replace
the bijection with a numpyro constraint and use `biject_to`:

```py
# Before (0.13.x)
class MyParam(gpx.parameters.Parameter):
# Custom bijection registered via DEFAULT_BIJECTION
...
```

```py
# After (0.14.0)
from numpyro.distributions import biject_to, constraints
from paramax import AbstractUnwrappable
import jax

class MyParam(AbstractUnwrappable):
_constraint = constraints.positive
_unconstrained: jax.Array

def __init__(self, value):
self._unconstrained = biject_to(self._constraint).inv(value)

def unwrap(self):
return biject_to(self._constraint)(self._unconstrained)
```

## Non-breaking cleanup

- `__description__` changed from `"Gaussian processes in JAX and Flax"` to
`"Gaussian processes in JAX"`, since Flax is no longer a dependency.
- Many kernel `compute_engine` internals moved; the public `kernel(x, y)`,
`kernel.gram(x)`, `kernel.cross_covariance(x, y)`, and `kernel.diagonal(x)`
methods are unchanged.

## Upgrade checklist

- [ ] Replace `nnx.Module` base classes with `eqx.Module`, and add class-level
type annotations for every field.
- [ ] Replace `param.value` reads with `param.unwrap()`, or call
`paramax.unwrap(model)` once at the top of your loss / prediction
function.
- [ ] Drop any `params_bijection=` / `trainable=` arguments passed to
`gpx.fit`. To freeze parameters, wrap them with `paramax.non_trainable`
using `eqx.tree_at`.
- [ ] Remove any `gpx.parameters.register_parameters` decorator calls.
- [ ] If you construct `LowerTriangular` from a custom matrix, verify the
diagonal is strictly positive.
- [ ] If you used `gpjax.linalg` operators directly, switch to the Lineax
equivalents listed above.

## Reporting issues

This is a pre-release (`0.14.0rc1`). Please file migration issues at
<https://github.com/thomaspinder/GPJax/issues> with the `0.14-migration`
label so they can be triaged before the stable `0.14.0` release.
Loading
Loading