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
8 changes: 8 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -248,6 +248,14 @@ cython_debug/
# Hatch-VCS
_version.py

# LLMs
.claude

/scratch
/data
notebooks
docs/_static/logo*

# Benchmark smoke runs verify the harness; only real runs are committed.
benchmarks/results/smoke.json
benchmarks/smoke/
28 changes: 28 additions & 0 deletions benchmarks/README.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,28 @@
# benchmarks

This directory contains the code and results for scaling benchmarks of `harv` functionality.
Currently, this is mainly for `harv.samplers.RejectionSampler`, and we compare
benchmarks for different model parameterizations, epoch counts, prior library sizes, and
`batch_size`.

- `grid.py` defines the grid of benchmark cells and builds the data, priors, and
models for each one.
- `test_rejection_scaling.py` is the benchmark itself, one `pytest-benchmark`
test parametrized over that grid.
- `report.py` merges `results/*.json` into `docs/benchmarks.md` and its figures.
- `results/` holds the committed JSON from the runs that page is built from.

These are deliberately not part of the test suite: they live outside
`testpaths`, require `--bench` to run, and depend on the `bench` group that
is not installed in CI.
The results page is committed rather than rebuilt so we can compare CPU and GPU
performance.

For the measurements themselves see `docs/benchmarks.md`, for how to reproduce
them see `docs/running-benchmarks.md`, and for how to use the numbers when
running over a survey see `docs/at-scale.md`.

Most of the code in this directory was written by Claude Opus 5.

TODO: generalize the benchmark code so we can benchmark other functionality, like the
periodogram code.
251 changes: 251 additions & 0 deletions benchmarks/conftest.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,251 @@
"""Fixtures and options for the rejection-sampler benchmarks.

Nothing here runs unless ``--bench`` is passed. See ``docs/running-benchmarks.md``.
"""

# x64 FIRST, before anything imports harv (and therefore JAX). harv deliberately
# does not enable it -- docs/sharp-bits.md makes it the user's job -- and every
# tutorial turns it on. float32 changes the sampler's *arithmetic*, not just its
# precision, so a float32 timing would not describe how anyone runs harv.
# Safe here: the root conftest.py only installs import hooks and creates no arrays.
import jax

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

import functools # noqa: E402
import os # noqa: E402
from pathlib import Path # noqa: E402
from typing import Any # noqa: E402

import jax.random as jr # noqa: E402
import pytest # noqa: E402
from grid import Cell, build_prior_and_model, enumerate_cells # noqa: E402

REPO_ROOT = Path(__file__).resolve().parent.parent


def pytest_addoption(parser: pytest.Parser) -> None:
group = parser.getgroup("harv benchmarks")
group.addoption(
"--bench",
action="store_true",
default=False,
help="Run the rejection-sampler benchmarks. Without this they all skip, "
"so a stray `pytest benchmarks/` cannot start a multi-hour run.",
)
group.addoption(
"--bench-full",
action="store_true",
default=False,
help="Full cartesian product instead of the default star design (~10x longer).",
)
group.addoption(
"--bench-smoke",
action="store_true",
default=False,
help="Two tiny cells that exercise the whole pipeline in under a minute.",
)
group.addoption(
"--bench-rounds",
type=int,
default=5,
help="Timed rounds per cell, after one warmup round (default: 5).",
)
group.addoption(
"--bench-expect",
choices=("cpu", "gpu"),
default=None,
help="Fail immediately unless JAX is on this backend. Guards the CPU/GPU "
"pair: a filename says nothing about which device actually ran, so without "
"this a silent fallback costs the whole run.",
)
group.addoption(
"--bench-cache-dir",
default=None,
help="Directory for HDF5 prior caches. Default: a pytest temp dir, "
"rebuilt each run. Point it somewhere persistent to reuse them.",
)


def pytest_configure(config: pytest.Config) -> None:
"""Make the --benchmark-json destination usable before any measuring starts.

pytest-benchmark writes that file from ``pytest_sessionfinish`` -- after every
benchmark has run. So a missing directory or a typo'd path does not fail fast,
it discards the entire session, which on a real run is hours of compute. Create
the directory and prove it is writable now instead.
"""
expect = config.getoption("--bench-expect")
if expect is not None:
# jax.default_backend() is the documented answer to "what am I running on";
# Device.platform has spelled CUDA both "gpu" and "cuda" across versions.
backend = jax.default_backend()
actual = "cpu" if backend == "cpu" else "gpu"
if actual != expect:
msg = (
f"--bench-expect={expect} but JAX is on {backend!r} "
f"({jax.devices()[0].device_kind}). "
+ (
"JAX falls back to CPU when CUDA-enabled jaxlib or the driver is "
"missing; see docs/running-benchmarks.md, 'Installing with GPU "
"support'."
if expect == "gpu"
else "Pass JAX_PLATFORMS=cpu to force the CPU backend."
)
)
raise pytest.UsageError(msg)

json_path = config.getoption("benchmark_json", None)
if not json_path:
return
parent = Path(json_path).parent
try:
parent.mkdir(parents=True, exist_ok=True)
probe = parent / f".write-probe-{os.getpid()}"
probe.touch()
probe.unlink()
except OSError as exc:
msg = f"--benchmark-json path is not writable: {json_path} ({exc})"
raise pytest.UsageError(msg) from exc


def pytest_collection_modifyitems(
config: pytest.Config, items: list[pytest.Item]
) -> None:
if config.getoption("--bench"):
return
skip = pytest.mark.skip(
reason="benchmarks require --bench (see docs/running-benchmarks.md)"
)
for item in items:
if "benchmarks/" in str(item.fspath).replace(os.sep, "/"):
item.add_marker(skip)


def pytest_generate_tests(metafunc: pytest.Metafunc) -> None:
"""Parametrize over the grid, chosen by the --bench-* flags."""
if "cell" not in metafunc.fixturenames:
return
cells = enumerate_cells(
full=metafunc.config.getoption("--bench-full"),
smoke=metafunc.config.getoption("--bench-smoke"),
)
metafunc.parametrize("cell", cells, ids=[c.ident for c in cells])


@pytest.fixture(scope="session")
def rounds(request: pytest.FixtureRequest) -> int:
n: int = request.config.getoption("--bench-rounds")
return 1 if request.config.getoption("--bench-smoke") else n


@pytest.fixture(scope="session")
def grid_mode(request: pytest.FixtureRequest) -> str:
"""Which grid produced this run.

Recorded into every result so report.py can reconstruct curve membership
exactly. Inferring it from the curve names in the data does not work: curves
share baseline cells, and a smoke run legitimately reuses a real curve name.
"""
if request.config.getoption("--bench-smoke"):
return "smoke"
if request.config.getoption("--bench-full"):
return "full"
return "star"


@pytest.fixture(scope="session")
def device_info() -> dict[str, Any]:
"""Device and version metadata.

pytest-benchmark's ``machine_info`` records the CPU and Python build but knows
nothing about accelerators, so the GPU identity has to come from here or the
CPU and GPU result files are indistinguishable in the report.
"""
import harv

device = jax.devices()[0]
return {
"device_platform": device.platform,
"device_kind": device.device_kind,
"device_count": jax.device_count(),
"jax_version": jax.__version__,
"harv_version": getattr(harv, "__version__", "unknown"),
"x64": bool(jax.config.read("jax_enable_x64")),
"typecheck_hooks": not os.environ.get("HARV_NO_TYPECHECK"),
}


@pytest.fixture(scope="session")
def cache_dir(request: pytest.FixtureRequest, tmp_path_factory: Any) -> Path:
opt = request.config.getoption("--bench-cache-dir")
if opt:
path = Path(opt).expanduser().resolve()
path.mkdir(parents=True, exist_ok=True)
return path
return tmp_path_factory.mktemp("prior-caches")


@pytest.fixture(scope="session")
def max_samples_by_parameterization(request: pytest.FixtureRequest) -> dict[str, int]:
"""Largest library each parameterization needs, so none is built oversized."""
cells = enumerate_cells(
full=request.config.getoption("--bench-full"),
smoke=request.config.getoption("--bench-smoke"),
)
out: dict[str, int] = {}
for cell in cells:
out[cell.parameterization] = max(
out.get(cell.parameterization, 0), cell.n_prior_samples
)
return out


# `maxsize=1` plus the parameterization-sorted cell list means exactly one prior
# library is resident at a time. At M = 1e7 a Gaia library is ~500 MB, so holding
# all six would be gigabytes for no reason.
@functools.lru_cache(maxsize=1)
def _build_memory_cache(parameterization: str, n_samples: int) -> Any:
prior, model = build_prior_and_model(parameterization)
return prior.sample(jr.key(0), n_samples, model=model)


@functools.cache
def _build_hdf5_cache(parameterization: str, n_samples: int, out_dir: str) -> str:
from harv.samplers import make_prior_cache

path = Path(out_dir) / f"{parameterization}-{n_samples}.h5"
if path.exists():
return str(path)
prior, model = build_prior_and_model(parameterization)
make_prior_cache(
prior,
model,
n_samples,
path,
key=jr.key(0),
batch_size=min(100_000, n_samples),
)
return str(path)


@pytest.fixture
def prior_cache(
cell: Cell,
cache_dir: Path,
max_samples_by_parameterization: dict[str, int],
) -> Any:
"""The prior library for this cell, in whichever backend the cell names.

The in-memory library is built once per parameterization at its largest
required size and sliced down -- ``Samples`` slices all arrays along the
leading axis (docs/spec.md, "Dict-style and index access"), so a slice is a
view-shaped copy rather than a fresh round of prior sampling.
"""
if cell.backend == "memory":
n_max = max_samples_by_parameterization[cell.parameterization]
full = _build_memory_cache(cell.parameterization, n_max)
return full if cell.n_prior_samples == n_max else full[: cell.n_prior_samples]
return _build_hdf5_cache(
cell.parameterization, cell.n_prior_samples, str(cache_dir)
)
Loading
Loading