Repository navigation
Test Equinox backend - #614
Conversation
|
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 |
|
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. |
|
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 The slowdown must come from something else in the docs build pipeline: Benchmark ResultsSame machine, clean venvs, CPU-only, JAX 0.9.0, Python 3.12. Light workloads (N=100)Heavy workloads (N=200, 500 iters; N=500 for large_conj)Conclusion: Equinox is equal or faster on every single benchmark. 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:
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() |
|
@claude Review this |
|
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
Heavy Benchmarks
Scaling Benchmarks
End-to-End Workflows
|
|
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. |
…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)
|
btw, |
Yes - I think this fine! |
|
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! |
`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>
Breaking changes
LowerTriangularparameter now requires a valid Cholesky factorgpjax.parameters.LowerTriangularpreviously accepted any lower-triangular matrix (its diagonal was unconstrained). It now parameterises a Cholesky factor vianumpyro.distributions.constraints.softplus_lower_cholesky, so the diagonal must be strictly positive.NaNduring construction (inv_softplusof a non-positive number).VariationalGaussian.variational_root_covariance, which is initialised to the identity by default.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.