Skip to main content

jaxbo: Bayesian optimization in JAX

CI PyPI License

A maintained, modern-JAX Bayesian optimization library: a small stable core (exact GP regression, 5 kernels, the classic acquisition functions plus a batched score_candidates helper, input priors, L-BFGS training) with optional research extras (MCMC inference, multifidelity and manifold GPs, weighted sampling) that install on demand and are never paid for otherwise.

This is a maintained fork of PredictiveIntelligenceLab/JAX-BO; see the acknowledgment below.

Install

pip install jaxbo

# The quickstart below needs 0.2.2 (score_candidates, extras). Until it is
# on PyPI, install from the repo:
# pip install "jaxbo @ git+https://github.com/ricardogr07/JAX-BO"

The core installs with exactly 4 dependencies: jax, jaxlib, numpy, scipy. Everything else lives behind optional extras:

Install Adds Extra dependencies
pip install jaxbo core: jaxbo.gp, jaxbo.kernels, jaxbo.acquisitions, jaxbo.optimizers, jaxbo.priors, jaxbo.test_functions none
pip install jaxbo[mcmc] jaxbo.mcmc: NUTS-based GP models numpyro
pip install jaxbo[multifidelity] jaxbo.multifidelity: multifidelity, manifold, gradient, and multi-output GPs none
pip install jaxbo[weighted] jaxbo.weights: GMM/KDE weighted acquisition machinery scikit-learn, KDEpy
pip install jaxbo[all] all of the above union

import jaxbo never imports an extra's dependencies; importing an extra without them raises an ImportError naming the pip install jaxbo[extra] fix.

Supported versions

The package pins jax>=0.6,<0.11 (matching jaxlib) and requires Python 3.10 or newer. CI tests every lane below at the jax floor and the newest jax the pin allows:

Python jax tested Status
3.10 0.6.0 to 0.6.2 supported (0.6.2 is the last jax with 3.10 wheels)
3.11 0.6.0 to 0.10.2 supported
3.12 0.6.0 to 0.10.2 supported
3.13 0.6.0 to 0.10.2 tested, advisory until the CI lanes hold a green streak
3.14 0.7.2 to 0.10.2 tested, advisory until the CI lanes hold a green streak

Quickstart: 60 seconds to the next point

Fit a GP to observations of an objective, then score a batch of candidates with expected improvement in one vmapped pass:

import jax.numpy as jnp
import numpy as np
from jax import random

from jaxbo.acquisitions import score_candidates
from jaxbo.gp import GP
from jaxbo.priors import uniform_prior
from jaxbo.utils import normalize


def f(x):
    return ((x - 1.5) ** 2).ravel()


# Domain and observations (raw domain)
lb, ub = jnp.array([-2.0]), jnp.array([3.0])
bounds = {"lb": lb, "ub": ub}
X = jnp.linspace(-2.0, 3.0, 8)[:, None]
y = f(X)

# GP expects an already normalized training batch
batch, norm_const = normalize(X, y, bounds)

gp = GP({"kernel": "RBF", "input_prior": uniform_prior(lb, ub), "criterion": "EI"})
params = gp.train(batch, random.PRNGKey(0), num_restarts=5)

# Score 256 raw-domain candidates in one vmapped pass and pick the best
X_cand = np.linspace(-2.0, 3.0, 256)[:, None]
scores = score_candidates(
    gp, X_cand, params=params, batch=batch, bounds=bounds,
    best=float(np.min(batch["y"])),
)
x_next = X_cand[np.argmin(scores)]
print("next point to evaluate:", x_next)  # close to the true minimum at 1.5

Two contract points worth knowing before you swap in a real objective:

  • Normalization: train consumes batch["X"] exactly as given, so pass the already normalized batch (utils.normalize). predict and score_candidates normalize raw-domain inputs internally against bounds. Normalized batch in, raw candidates in; mixing these up fails silently.
  • Scores: lower is better for every acquisition in jaxbo.acquisitions (EI is returned negated), so the next point is X_cand[np.argmin(scores)].

Prefer a notebook? Launch the interactive tutorial on Colab: Open Demo in Colab

Docker

A quickstart image with jaxbo[all], JupyterLab, and the examples/ notebooks is published to GHCR with each release, starting with 0.2.0:

docker run -p 8888:8888 ghcr.io/ricardogr07/jaxbo

Then open the printed http://127.0.0.1:8888/lab?token=... URL.

Project documentation

Final benchmark evidence

The 0.2.2 benchmark summary records four clean runs from the existing .venv, using the command and environment listed in the evidence file. Historical ratios remain descriptive only because the cross-session noise bands are wide.

Bench 0.2.2 median vs 2026-07-28 vs v0.2.0
Train warm, n=32 16.53 ms 0.19x 0.31x
Train warm, n=128 110.84 ms 0.62x 0.80x
Train warm, n=512 715.70 ms 0.29x 0.40x
Train fresh, n=128 101.37 ms 0.16x 0.20x
Predict, 256 points 1.35 ms 1.26x 1.36x
EI consumer, 256 candidates 116.83 ms 0.85x 0.96x
EI fused, 256 candidates 65.63 ms 0.86x 0.92x
score_candidates, 256 candidates 1.38 ms N/A 0.86x

Acknowledgment

This project is a fork of JAX-BO by the Predictive Intelligence Lab at the University of Pennsylvania (Paris Perdikaris, Yibo Yang, Mohamed Aziz Bhouri). The core model structure, kernels, and acquisition functions originate there; this fork modernizes the library (current jax/python support, core/extras packaging, tests, benchmarks, CI) and is maintained independently by Ricardo García Ramírez. It is not affiliated with the original authors.

If you use this library in your research, please cite the original work (see also CITATION.cff):

@software{jaxbo2020github,
  author = {Paris Perdikaris, Yibo Yang, Mohamed Aziz Bhouri},
  title = {{JAX-BO}: A Bayesian optimization library in {JAX}},
  url = {https://github.com/PredictiveIntelligenceLab/JAX-BO},
  version = {0.2},
  year = {2020},
}

License

Apache License 2.0. See LICENSE.

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

jaxbo-0.3.0.tar.gz (98.8 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

jaxbo-0.3.0-py3-none-any.whl (98.8 kB view details)

Uploaded Python 3

File details

Details for the file jaxbo-0.3.0.tar.gz.

File metadata

  • Download URL: jaxbo-0.3.0.tar.gz
  • Upload date:
  • Size: 98.8 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for jaxbo-0.3.0.tar.gz
Algorithm Hash digest
SHA256 ddd8a330cd45a05822cf82dae5d1452d86fdb3d46b8bf056094b64a34c0ed1df
MD5 e100b9df387f437c09b2fc667450a10d
BLAKE2b-256 7865b0655dff1bc9bd6e94d4d0891f7f0e378915661298c38da6ff1b43d6f3f5

See more details on using hashes here.

Provenance

The following attestation bundles were made for jaxbo-0.3.0.tar.gz:

Publisher: release.yml on ricardogr07/JAX-BO

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

File details

Details for the file jaxbo-0.3.0-py3-none-any.whl.

File metadata

  • Download URL: jaxbo-0.3.0-py3-none-any.whl
  • Upload date:
  • Size: 98.8 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? Yes
  • Uploaded via: twine/7.0.0 CPython/3.13.14

File hashes

Hashes for jaxbo-0.3.0-py3-none-any.whl
Algorithm Hash digest
SHA256 ce571134697c7242d8d99840a1b0189cd18abfdcf1c26bcfceb8d7d498e80c03
MD5 740286df0d7ef74e370573d5a8e36ac2
BLAKE2b-256 843b6544e11d8bb9a4dcd8662143965f112ce8f77339d86ff7caa016d5a7bb2d

See more details on using hashes here.

Provenance

The following attestation bundles were made for jaxbo-0.3.0-py3-none-any.whl:

Publisher: release.yml on ricardogr07/JAX-BO

Attestations: Values shown here reflect the state when the release was signed and may no longer be current.

Release history Release notifications | RSS feed

This release

0.3.0 This release

2 files

0.2.2

2 files

0.2.1

2 files

0.2.0

2 files

0.1.2

2 files

0.1.1

2 files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page