Skip to content

Test Equinox backend - #614

Merged
thomaspinder merged 9 commits into
mainfrom
thomaspinder/equinox-backend
Apr 13, 2026
Merged

thomaspinder merged 9 commits into
mainfrom
thomaspinder/equinox-backend

Conversation

@thomaspinder

@thomaspinder thomaspinder commented Mar 14, 2026 •

Copy link
Copy Markdown
Collaborator

Breaking changes

LowerTriangular parameter now requires a valid Cholesky factor

gpjax.parameters.LowerTriangular previously accepted any lower-triangular matrix (its diagonal was unconstrained). It now parameterises a Cholesky factor via numpyro.distributions.constraints.softplus_lower_cholesky, so the diagonal must be strictly positive.

  • Passing a matrix with zero or negative diagonal entries will produce 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. This change removes a latent footgun — under the old parameterisation, zero/negative diagonals led to singular or sign-ambiguous variational covariances.

Comment thread docs/sharp_bits.md Outdated
Comment thread examples/backend.py
Comment thread gpjax/parameters.py Outdated
Comment thread gpjax/numpyro_extras.py Outdated
Comment thread gpjax/linalg/custom_operators.py
@thomaspinder

Copy link
Copy Markdown
Collaborator Author

One thing that's concerning with Equinox is that things run quite a bit slower. This is not a detailed analysis, but the CI for docs building jumps to 21mins from 15mins in this PR. Like I say, required more thorough testing, but if the degradation is true, then Equinox is a no go. cc @theorashid

@theorashid

theorashid commented Mar 24, 2026 •

Copy link
Copy Markdown
Contributor

That's really surprising. I wonder if Claude can do some analysis on that at the jaxpr level. If anything, the Equinox authors would probably be interested.

@theorashid

Copy link
Copy Markdown
Contributor

Hey, so I got Claude to do some benchmarking on both branches.

Here are the results:

The core GP operations (fit, fit_scipy, predict) are NOT slower with
Equinox.
Benchmarks on the same machine show Equinox is equal or faster
than NNX for standard conjugate/non-conjugate regression. The 6-minute
docs build regression is NOT caused by the JAX/XLA computation itself.

The slowdown must come from something else in the docs build pipeline:
notebook execution overhead, plotting, specific complex examples
(deep kernels, multi-output, OILMM), or CI infrastructure variance.

Benchmark Results

Same machine, clean venvs, CPU-only, JAX 0.9.0, Python 3.12.

Light workloads (N=100)

                               NNX (master)   Equinox (branch)
fit_scipy compile              0.6532s        0.3517s   << eqx faster
fit_scipy run (avg of 3)       0.3490s        0.1595s   << eqx 2x faster
fit compile                    0.3260s        0.3365s      ~same
fit run (avg of 3)             0.2668s        0.2372s      ~same
fit_nc compile                 0.3260s        0.2905s      ~same
fit_nc run (avg of 3)          0.2516s        0.2477s      ~same
predict compile                0.0208s        0.0256s      ~same
predict run                    ~0s            ~0s          ~same

Heavy workloads (N=200, 500 iters; N=500 for large_conj)

                               NNX (master)   Equinox (branch)
collapsed_vi compile           0.6729s        0.6681s      ~same
collapsed_vi run (avg of 2)    0.4481s        0.4912s      ~same
uncollapsed_vi compile         1.5940s        1.0846s   << eqx faster
uncollapsed_vi run (avg of 2)  0.8758s        0.8916s      ~same
large_conj_scipy compile       0.7595s        0.5094s   << eqx faster
large_conj_scipy run (avg of 2)0.5115s        0.3732s   << eqx faster

Conclusion: Equinox is equal or faster on every single benchmark.
The docs build slowdown is NOT caused by slower GP computation. Carrying
a full eqx.Module through jax.lax.scan does NOT cause a performance hit.


I'll attach the benchmark scripts at the end. Claude was pretty confident the internals were sound.

The more likely issue would be CI instance variance. Maybe there was a cold runner or a busy cloud at the time. That was my instinct too.

Other things it suggested:

  1. Import time -- Equinox + Lineax + Paramax may have different cold
    import costs than Flax NNX. Each notebook is a fresh Python process, so
    import overhead multiplies by ~21.

3. Specific example differences -- Some examples may have changed
behaviour (different iteration counts, dataset sizes, or added plots).

  1. The UserWarning: A JAX array is being set as static on
    NonConjugatePosterior
    -- key: tp.Any = eqx.field(static=True) means
    the PRNG key is baked into the treedef. If different keys are used across
    calls this would cause retracing. Unlikely to matter for docs though since
    each notebook runs once.

bench.py

"""Benchmark script for GPJax -- works on both NNX and Equinox backends."""

import time
import jax
import jax.numpy as jnp
import jax.random as jr
import optax as ox

jax.config.update("jax_enable_x64", True)

import gpjax as gpx

SEED = jr.key(0)
N_TRAIN = 100
N_TEST = 50


def make_regression_data(key):
    x = jnp.linspace(0.0, 10.0, N_TRAIN).reshape(-1, 1)
    y = jnp.sin(x) + 0.1 * jr.normal(key, x.shape)
    return gpx.Dataset(X=x, y=y)


def make_classification_data(key):
    x = jnp.sort(jr.uniform(key, shape=(N_TRAIN, 1), minval=-3, maxval=3), axis=0)
    y = (jnp.sin(x) > 0).astype(jnp.float64)
    return gpx.Dataset(X=x, y=y)


def build_conjugate(data):
    meanf = gpx.mean_functions.Constant()
    kernel = gpx.kernels.RBF()
    likelihood = gpx.likelihoods.Gaussian(num_datapoints=data.n)
    prior = gpx.gps.Prior(mean_function=meanf, kernel=kernel)
    return prior * likelihood


def build_nonconjugate(data):
    meanf = gpx.mean_functions.Constant()
    kernel = gpx.kernels.RBF()
    likelihood = gpx.likelihoods.Bernoulli(num_datapoints=data.n)
    prior = gpx.gps.Prior(mean_function=meanf, kernel=kernel)
    return prior * likelihood


def block(result):
    jax.tree_util.tree_map(
        lambda x: x.block_until_ready() if hasattr(x, "block_until_ready") else x,
        result,
    )


def time_fn(fn, label, warmup=1, repeats=3):
    for _ in range(warmup):
        block(fn())
    times = []
    for _ in range(repeats):
        t0 = time.perf_counter()
        block(fn())
        elapsed = time.perf_counter() - t0
        times.append(elapsed)
    avg = sum(times) / len(times)
    print(f"  {label}: {avg:.4f}s  ({[f'{t:.4f}' for t in times]})")
    return avg


def time_first(fn, label):
    t0 = time.perf_counter()
    block(fn())
    elapsed = time.perf_counter() - t0
    print(f"  {label} (compile): {elapsed:.4f}s")
    return elapsed


def main():
    print(f"JAX {jax.__version__}, GPJax {gpx.__version__}")
    try:
        import equinox
        print(f"Backend: Equinox {equinox.__version__}")
    except ImportError:
        print("Backend: Flax NNX")
    print(f"Device: {jax.devices()[0]}")
    print()

    k1, k2 = jr.split(SEED)
    reg_data = make_regression_data(k1)
    cls_data = make_classification_data(k2)
    nmll = lambda p, d: -gpx.objectives.conjugate_mll(p, d)
    xtest = jnp.linspace(0.0, 10.0, N_TEST).reshape(-1, 1)
    R = {}

    # 1. fit_scipy
    print("1. fit_scipy (conjugate, L-BFGS-B, 500 iters)")
    R["fit_scipy_compile"] = time_first(
        lambda: gpx.fit_scipy(model=build_conjugate(reg_data), objective=nmll,
                              train_data=reg_data, verbose=False),
        "fit_scipy")
    R["fit_scipy_run"] = time_fn(
        lambda: gpx.fit_scipy(model=build_conjugate(reg_data), objective=nmll,
                              train_data=reg_data, verbose=False),
        "fit_scipy", warmup=0, repeats=3)

    # 2. fit (conjugate, Adam, 200 iters)
    print("\n2. fit (conjugate, Adam, 200 iters)")
    R["fit_compile"] = time_first(
        lambda: gpx.fit(model=build_conjugate(reg_data), objective=nmll,
                        train_data=reg_data, optim=ox.adam(0.01),
                        num_iters=200, verbose=False),
        "fit")
    R["fit_run"] = time_fn(
        lambda: gpx.fit(model=build_conjugate(reg_data), objective=nmll,
                        train_data=reg_data, optim=ox.adam(0.01),
                        num_iters=200, verbose=False),
        "fit", warmup=0, repeats=3)

    # 3. fit (non-conjugate, Adam, 200 iters)
    print("\n3. fit (non-conjugate, Adam, 200 iters)")
    nc_obj = lambda p, d: -gpx.objectives.log_posterior_density(p, d)
    R["fit_nc_compile"] = time_first(
        lambda: gpx.fit(model=build_nonconjugate(cls_data), objective=nc_obj,
                        train_data=cls_data, optim=ox.adam(0.01),
                        num_iters=200, verbose=False),
        "fit_nc")
    R["fit_nc_run"] = time_fn(
        lambda: gpx.fit(model=build_nonconjugate(cls_data), objective=nc_obj,
                        train_data=cls_data, optim=ox.adam(0.01),
                        num_iters=200, verbose=False),
        "fit_nc", warmup=0, repeats=3)

    # 4. Predict
    print("\n4. Posterior prediction (conjugate)")
    opt_post, _ = gpx.fit_scipy(model=build_conjugate(reg_data), objective=nmll,
                                train_data=reg_data, verbose=False)
    predict_jit = jax.jit(lambda m, t, d: m(t, d))
    R["predict_compile"] = time_first(
        lambda: predict_jit(opt_post, xtest, reg_data), "predict")
    R["predict_run"] = time_fn(
        lambda: predict_jit(opt_post, xtest, reg_data),
        "predict", warmup=0, repeats=5)

    # Summary
    print("\n" + "=" * 55)
    print("SUMMARY")
    print("=" * 55)
    for k, v in R.items():
        print(f"  {k:30s} {v:.4f}s")


if __name__ == "__main__":
    main()

bench_heavy.py

"""Heavier benchmarks: variational inference and larger problems."""

import time
import jax
import jax.numpy as jnp
import jax.random as jr
import optax as ox

jax.config.update("jax_enable_x64", True)

import gpjax as gpx


def block(result):
    jax.tree_util.tree_map(
        lambda x: x.block_until_ready() if hasattr(x, "block_until_ready") else x,
        result,
    )


def timeit(fn, label, repeats=3):
    times = []
    for i in range(repeats):
        t0 = time.perf_counter()
        block(fn())
        elapsed = time.perf_counter() - t0
        times.append(elapsed)
        tag = " (compile)" if i == 0 else ""
        print(f"  {label} run {i}{tag}: {elapsed:.4f}s")
    print(f"  {label} avg (excl compile): {sum(times[1:])/max(len(times)-1,1):.4f}s")
    return times


def main():
    print(f"JAX {jax.__version__}, GPJax {gpx.__version__}")
    try:
        import equinox
        print(f"Backend: Equinox {equinox.__version__}")
    except ImportError:
        print("Backend: Flax NNX")
    print(f"Device: {jax.devices()[0]}")
    print()

    key = jr.key(42)
    N = 200

    # --- Collapsed VI (like collapsed_vi.py example) ---
    print("1. Collapsed VI (N=200, M=20, 500 Adam iters)")
    x = jnp.linspace(0.0, 10.0, N).reshape(-1, 1)
    y = jnp.sin(x) + 0.2 * jr.normal(key, x.shape)
    D = gpx.Dataset(X=x, y=y)

    meanf = gpx.mean_functions.Constant()
    kernel = gpx.kernels.Matern52()
    likelihood = gpx.likelihoods.Gaussian(num_datapoints=D.n)
    prior = gpx.gps.Prior(mean_function=meanf, kernel=kernel)
    posterior = prior * likelihood

    z = jnp.linspace(0.0, 10.0, 20).reshape(-1, 1)
    q = gpx.variational_families.CollapsedVariationalGaussian(
        posterior=posterior, inducing_inputs=z
    )

    def run_collapsed_vi():
        q_fresh = gpx.variational_families.CollapsedVariationalGaussian(
            posterior=posterior, inducing_inputs=z
        )
        return gpx.fit(
            model=q_fresh,
            objective=lambda p, d: -gpx.objectives.collapsed_elbo(p, d),
            train_data=D,
            optim=ox.adamw(learning_rate=1e-2),
            num_iters=500,
            verbose=False,
        )

    timeit(run_collapsed_vi, "collapsed_vi")

    # --- Uncollapsed VI (like uncollapsed_vi.py example) ---
    print("\n2. Uncollapsed VI (N=200, M=20, 500 Adam iters)")
    meanf2 = gpx.mean_functions.Constant()
    kernel2 = gpx.kernels.Matern52()
    likelihood2 = gpx.likelihoods.Bernoulli(num_datapoints=N)
    prior2 = gpx.gps.Prior(mean_function=meanf2, kernel=kernel2)
    posterior2 = prior2 * likelihood2

    q2 = gpx.variational_families.VariationalGaussian(
        posterior=posterior2, inducing_inputs=z
    )

    def run_uncollapsed_vi():
        q2_fresh = gpx.variational_families.VariationalGaussian(
            posterior=posterior2, inducing_inputs=z
        )
        return gpx.fit(
            model=q2_fresh,
            objective=lambda p, d: -gpx.objectives.elbo(p, d),
            train_data=D,
            optim=ox.adam(learning_rate=1e-2),
            num_iters=500,
            verbose=False,
        )

    timeit(run_uncollapsed_vi, "uncollapsed_vi")

    # --- Larger conjugate regression (N=500) ---
    print("\n3. Conjugate regression (N=500, fit_scipy)")
    x3 = jnp.linspace(0.0, 10.0, 500).reshape(-1, 1)
    y3 = jnp.sin(x3) + 0.1 * jr.normal(key, x3.shape)
    D3 = gpx.Dataset(X=x3, y=y3)

    def run_large_conj():
        p = gpx.gps.Prior(
            mean_function=gpx.mean_functions.Constant(),
            kernel=gpx.kernels.RBF(),
        ) * gpx.likelihoods.Gaussian(num_datapoints=500)
        return gpx.fit_scipy(
            model=p,
            objective=lambda m, d: -gpx.objectives.conjugate_mll(m, d),
            train_data=D3,
            verbose=False,
        )

    timeit(run_large_conj, "large_conj_scipy")


if __name__ == "__main__":
    main()

@theorashid theorashid mentioned this pull request Apr 5, 2026
2 of 3 tasks
@thomaspinder

Copy link
Copy Markdown
Collaborator Author

@claude Review this

@thomaspinder

Copy link
Copy Markdown
Collaborator Author

Hey @theorashid. Thanks for pushing this. I ran some more benchmarks locally. Sharing the full results below, but overall I think you're right - the Equinox branch is on the whole faster...

Core Benchmarks

Benchmark Flax NNX 0.12.2 Equinox 0.13.5 Verdict
fit_scipy_compile 0.5266s 0.2369s 2.2x faster
fit_scipy_run 0.2349s 0.1123s 2.1x faster
fit_lbfgs_compile 0.4081s 0.3376s 1.2x faster
fit_lbfgs_run 0.3495s 0.3115s 1.1x faster
fit_compile 0.2140s 0.2000s 1.1x faster
fit_run 0.1656s 0.1554s 1.1x faster
fit_nc_compile 0.1980s 0.2490s 1.3x slower
fit_nc_run 0.1707s 0.1557s 1.1x faster
predict_compile 0.0189s 0.0220s 1.2x slower
predict_run 0.0000s 0.0000s 1.2x faster

Heavy Benchmarks

Benchmark Flax NNX 0.12.2 Equinox 0.13.5 Verdict
collapsed_vi_compile 0.4717s 0.3850s 1.2x faster
collapsed_vi_run 0.2894s 0.2708s 1.1x faster
uncollapsed_vi_compile 1.2418s 1.0413s 1.2x faster
uncollapsed_vi_run 0.5422s 0.4974s 1.1x faster
large_conj_scipy_compile 0.4133s 0.2716s 1.5x faster
large_conj_scipy_run 0.3312s 0.1969s 1.7x faster

Scaling Benchmarks

Benchmark Flax NNX 0.12.2 Equinox 0.13.5 Verdict
N=100 fit_adam 0.1649s 0.1529s 1.1x faster
N=100 fit_lbfgs 0.3302s 0.3104s 1.1x faster
N=100 fit_nc 0.1638s 0.1546s 1.1x faster
N=100 fit_scipy 0.2287s 0.1008s 2.3x faster
N=100 predict 0.0000s 0.0000s 1.1x faster
N=300 fit_adam 0.3526s 0.3664s ~same
N=300 fit_lbfgs 0.3529s 0.3425s ~same
N=300 fit_nc 0.3431s 0.3448s ~same
N=300 fit_scipy 0.2552s 0.1364s 1.9x faster
N=300 predict 0.0000s 0.0000s 1.1x faster
N=500 fit_adam 0.6078s 0.7252s 1.2x slower
N=500 fit_lbfgs 0.3887s 0.4332s 1.1x slower
N=500 fit_nc 0.6965s 0.6988s ~same
N=500 fit_scipy 0.3085s 0.2270s 1.4x faster
N=500 predict 0.0000s 0.0000s 1.1x faster
N=1000 fit_adam 2.2133s 2.8214s 1.3x slower
N=1000 fit_lbfgs 0.7622s 1.0100s 1.3x slower
N=1000 fit_nc 3.0305s 2.6835s 1.1x faster
N=1000 fit_scipy 0.5317s 1.2481s 2.3x slower
N=1000 predict 0.0000s 0.0000s ~same

End-to-End Workflows

Benchmark Flax NNX 0.12.2 Equinox 0.13.5 Verdict
N=100 conjugate_regression 0.6770s 0.3916s 1.7x faster
N=100 poisson_regression 0.2819s 0.2705s ~same
N=100 sparse_variational 0.4306s 0.5180s 1.2x slower
N=300 conjugate_regression 0.4051s 0.2499s 1.6x faster
N=300 poisson_regression 0.4649s 0.4794s ~same
N=300 sparse_variational 0.4578s 0.5162s 1.1x slower
N=500 conjugate_regression 0.4021s 0.4111s ~same
N=500 poisson_regression 0.9973s 1.0490s 1.1x slower
N=500 sparse_variational 0.6621s 0.6915s ~same
N=1000 conjugate_regression 1.2659s 0.7946s 1.6x faster
N=1000 poisson_regression 3.9737s 4.0111s ~same
N=1000 sparse_variational 1.3904s 1.1782s 1.2x faster

@theorashid

Copy link
Copy Markdown
Contributor

Nice, glad you got to the same conclusion with your benchmarks. The #621 wouldn't affect performance, but I think it streamlines some of the code, particularly the integration with numpyro.

I reviewed the current PR myself whilst making that suggestions PR, using Claude too. Could probably do with a final fresh Claude review. But I'm happier with the Equinox backend.

theorashid and others added 2 commits April 13, 2026 20:57
…parameters (#621)

- Replace custom FillTriangularTransform and inv_softplus with numpyro's
  biject_to + constraints (softplus_positive, interval, softplus_lower_cholesky)
- Remove prior field from all Parameter classes (PositiveReal, NonNegativeReal,
  Real, SigmoidBounded, LowerTriangular)
- Delete numpyro_extras.py (register_parameters, resolve_prior, tree_path_to_name)
- Update numpyro examples to sample hyperparameters directly with numpyro.sample
  and pass them to GPJax constructors as raw JAX arrays
- Remove unused self.key static field from NonConjugatePosterior (was causing
  'JAX array is being set as static' warning)
@theorashid

theorashid commented Apr 13, 2026 •

Copy link
Copy Markdown
Contributor

btw, FillTriangularTransform is gone and LowerTriangular now uses constraints.softplus_lower_cholesky, which is Lower-triangular matrix parameter with positive diagonal (Cholesky factor). Is that okay?

@thomaspinder

Copy link
Copy Markdown
Collaborator Author

btw, FillTriangularTransform is gone and LowerTriangular now uses constraints.softplus_lower_cholesky, which is Lower-triangular matrix parameter with positive diagonal (Cholesky factor). Is that okay?

Yes - I think this fine!

@thomaspinder
thomaspinder merged commit 85ab347 into main Apr 13, 2026
19 checks passed
@theorashid

Copy link
Copy Markdown
Contributor

Sweet. Let's update the numpyro docs notebook when you next do a release

@thomaspinder

Copy link
Copy Markdown
Collaborator Author

Sweet. Let's update the numpyro docs notebook when you next do a release

Just triggered the release workflow. Released as an rc, given the magnitude. If there's no issues by end of next week, then I'll formally release and update numpyro!

@thomaspinder
thomaspinder deleted the thomaspinder/equinox-backend branch April 13, 2026 21:15
thomaspinder added a commit that referenced this pull request Jul 26, 2026
`Zero()` was trainable and drifted towards the data mean during `fit`
(0.0 -> 5.09 on a dataset with mean 5), contradicting its own docstring and
silently changing the posterior mean of every model using the default mean
function. Regression of #330, fixed once in #500.

The cause is a changed trainability contract rather than a lost line. Under
nnx, `fit` optimised only `Parameter` instances, so `Zero`'s bare array was
inert by construction -- which is why #530 could drop the `Static` wrapper
and stay correct. Under Equinox, `fit` partitions on `eqx.is_array`, making
every array leaf trainable, and `Zero` was still relying on the old meaning
of a bare array.

Wrap the constant in `paramax.non_trainable` so the invariant holds by
construction rather than by the ambient filter semantics.

The guard for this was weakened rather than removed: #614 replaced
`test_zero_mean_remains_zero` with a test of the initial value only, and
`test_zero_mean_function_uses_raw_value` asserted the defect outright.
Restore the end-to-end fit assertion and invert the unit test.

Fixes #712


Claude-Session: https://claude.ai/code/session_01Bj9k5fnAZ8JzD4Rg3HMDMj

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>

This branch was previously deployed

1 inactive deployment
docs-preview — 23c101e4 Deployed Apr 13, 2026 by thomaspinder via deploy-docs-preview #734
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants