diff --git a/.github/workflows/pre-commit.yml b/.github/workflows/pre-commit.yml index 4b434e27..9175970c 100644 --- a/.github/workflows/pre-commit.yml +++ b/.github/workflows/pre-commit.yml @@ -1,6 +1,6 @@ # Pre-commit / lint check. # -# Runs `pre-commit run --all-files` (black, isort, flake8, pyupgrade, mypy, +# Runs `pre-commit run --all-files` (ruff-check, ruff-format, mypy, # end-of-file-fixer, trailing-whitespace, check-merge-conflict, debug-statements). # ~30 s wall. Runs on every PR + push to main. diff --git a/.github/workflows/test-fast.yml b/.github/workflows/test-fast.yml index d25c2b57..26a1b959 100644 --- a/.github/workflows/test-fast.yml +++ b/.github/workflows/test-fast.yml @@ -1,6 +1,6 @@ # Fast tests (PR gate). # -# `pytest -m fast -n 2` on Python 3.11 + 3.13 matrix. ~1.5 min wall per Python. +# `pytest -m fast -n 2` on Python 3.13 + 3.14 matrix. ~1.5 min wall per Python. # 676 tests in the current `fast` set (no JAX trace, < 100 ms each). # Runs on every PR + push to main. @@ -23,8 +23,14 @@ jobs: strategy: fail-fast: false matrix: - # pyproject declares ``requires-python = ">=3.13"``. Matrix matches. - python: ['3.13'] + # pyproject declares ``requires-python = ">=3.13"``. Matrix covers + # both supported Pythons (3.13 and 3.14). + python: ['3.13', '3.14'] + env: + # Pin every uv command to the matrix Python. Without this, `uv run` + # falls back to .python-version (3.13) and rebuilds the env without + # the bench group, so the 3.14 leg fails with "Failed to spawn: pytest". + UV_PYTHON: ${{ matrix.python }} steps: - name: Check out repo uses: actions/checkout@v4 diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 6ad47d20..024ee16c 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -1,50 +1,34 @@ -# Lint stack for tuningfork. Versions bumped from sampling-book's 2022-era pins -# so they install cleanly on Python 3.13. Tools and ignore lists kept identical -# in spirit to sampling-book/blackjax for cross-repo consistency. +# Lint stack for tuningfork. ruff (ruff-check + ruff-format) replaces +# isort/pyupgrade/flake8/black/nbQA — one tool, one config, native Jupyter +# notebook support if any .ipynb ever lands. Tools kept identical in spirit +# to sampling-book/blackjax for cross-repo consistency. repos: - repo: https://github.com/pre-commit/pre-commit-hooks - rev: v5.0.0 + rev: v6.0.0 hooks: - id: check-merge-conflict - id: debug-statements - id: end-of-file-fixer - id: trailing-whitespace -- repo: https://github.com/PyCQA/isort - rev: 5.13.2 - hooks: - - id: isort - args: [--profile, black] -- repo: https://github.com/asottile/pyupgrade - rev: v3.19.0 - hooks: - - id: pyupgrade - args: [--py311-plus] -- repo: https://github.com/PyCQA/flake8 - rev: 7.1.1 - hooks: - - id: flake8 - args: ['--ignore=E501,E203,E731,W503'] - exclude: '^scripts/' -- repo: https://github.com/psf/black - rev: 24.10.0 - hooks: - - id: black +- repo: https://github.com/astral-sh/ruff-pre-commit + rev: v0.16.10 + hooks: + - id: ruff-check + args: [--fix] + - id: ruff-format + # ruff-format's hook default also formats fenced code blocks in + # markdown (types_or includes "markdown"); the old stack + # (black/isort/flake8/nbQA) never touched .md prose docs, only .py + # and .ipynb. Keep that scope so curated doc formatting (aligned + # inline comments in catalog READMEs/schemas) isn't rewritten as a + # side effect of this migration. + types_or: [python, pyi, jupyter] - repo: https://github.com/pre-commit/mirrors-mypy - rev: v1.13.0 + rev: v2.4.0 hooks: - id: mypy args: [--ignore-missing-imports] exclude: '^scripts/' -- repo: https://github.com/nbQA-dev/nbQA - rev: 1.9.0 - hooks: - - id: nbqa-black - - id: nbqa-pyupgrade - args: [--py311-plus] - - id: nbqa-isort - args: ['--profile=black'] - - id: nbqa-flake8 - args: ['--ignore=E501,E203,E302,E402,E731,W503'] - repo: https://github.com/jorisroovers/gitlint rev: v0.19.1 hooks: diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 3c3d4be5..0ec47cad 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -126,7 +126,7 @@ All `test-*` targets automatically run `make clean-orphans` first (except `test- assert logp == 0.0 ``` -5. **Run `make lint` before committing** to pass pre-commit hooks (black, isort, flake8, mypy). +5. **Run `make lint` before committing** to pass pre-commit hooks (ruff, mypy). ## Pre-Commit and Commit Messages diff --git a/benchmarks/_benchmark_helpers.py b/benchmarks/_benchmark_helpers.py index fb2df59e..ba61a547 100644 --- a/benchmarks/_benchmark_helpers.py +++ b/benchmarks/_benchmark_helpers.py @@ -15,6 +15,7 @@ Used by both test_fast_recipes.py and test_e2e_recipes.py. """ + from __future__ import annotations import json diff --git a/benchmarks/config.py b/benchmarks/config.py index 3e393226..ea007776 100644 --- a/benchmarks/config.py +++ b/benchmarks/config.py @@ -27,6 +27,7 @@ The split keeps the fast suite quick enough for local smoke + targeted CI triggers, while the slow e2e cells run nightly only. """ + from __future__ import annotations # --------------------------------------------------------------------------- diff --git a/benchmarks/test_e2e_recipes.py b/benchmarks/test_e2e_recipes.py index e11d51b9..171fa35c 100644 --- a/benchmarks/test_e2e_recipes.py +++ b/benchmarks/test_e2e_recipes.py @@ -24,6 +24,7 @@ DO NOT add these to ``make benchmark-fast`` or per-PR triggers — the wall time makes them unsuitable for quick local checks. """ + from __future__ import annotations from typing import Any diff --git a/benchmarks/test_fast_recipes.py b/benchmarks/test_fast_recipes.py index 2c7ed061..9d202cc3 100644 --- a/benchmarks/test_fast_recipes.py +++ b/benchmarks/test_fast_recipes.py @@ -25,6 +25,7 @@ Cell selection: see ``benchmarks/config.py`` (FAST_CELLS, ≤60s/cell in CI). Slow e2e cells (>60s): see ``benchmarks/test_e2e_recipes.py``. """ + from __future__ import annotations from typing import Any diff --git a/benchmarks/test_speed_lite.py b/benchmarks/test_speed_lite.py index c0ed4cc6..46a67041 100644 --- a/benchmarks/test_speed_lite.py +++ b/benchmarks/test_speed_lite.py @@ -41,6 +41,7 @@ Trend persisted in ``actions/cache`` (not gh-pages). Alert at 200% with ``fail-on-alert: true`` (nightly red = regression signal). """ + from __future__ import annotations from pathlib import Path diff --git a/pyproject.toml b/pyproject.toml index 6d5fb533..667f682b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -126,8 +126,24 @@ exclude = [ [tool.setuptools_scm] write_to = "tuningfork/_version.py" -[tool.isort] -profile = "black" +[tool.ruff] +target-version = "py313" +line-length = 88 + +[tool.ruff.lint] +select = ["E", "F", "W", "I", "UP"] # = the old flake8 (pycodestyle+pyflakes) + isort + pyupgrade +ignore = [ + "E501", # line length: ruff-format owns it (was ignored before) + "E731", # lambda assignment (was ignored before) + "UP040", "UP046", "UP047", # PEP 695 `type`/generic rewrites: change runtime semantics, out of scope + "UP042", # str+Enum -> StrEnum changes str()/f-string output ("Effort.LOW" -> + # "low"); test_schema.py:61 asserts the current str() form and + # _resolve_execution_plan.py builds a file path from an f-strung + # effort value, so this rewrite is a behavior change, not a style fix +] + +[tool.ruff.lint.isort] +known-first-party = ["tuningfork"] [tool.pytest.ini_options] filterwarnings = [ diff --git a/tests/e2e/test_nightly_regression.py b/tests/e2e/test_nightly_regression.py index e84fdf61..94abaffe 100644 --- a/tests/e2e/test_nightly_regression.py +++ b/tests/e2e/test_nightly_regression.py @@ -22,6 +22,7 @@ make test-slow JAX_PLATFORM_NAME=cpu uv run pytest tests/e2e/test_nightly_regression.py -v """ + from __future__ import annotations import pytest @@ -81,9 +82,9 @@ def test_lotka_dense_imm_inner_nuts_seed_20260713_passes() -> None: catalog_root=catalog, ) z = compute_max_abs_mean_z(idata, "lotka_volterra") - assert ( - z is not None - ), "compute_max_abs_mean_z returned None — reference/summary.json missing?" + assert z is not None, ( + "compute_max_abs_mean_z returned None — reference/summary.json missing?" + ) assert z < 4.0, ( f"Seed 20260713 should pass z < 4.0 (it is a clean seed), got z={z:.3f}. " "If this seed now fails, the step-collapse has worsened — see issue #232." diff --git a/tests/e2e/test_phase1.py b/tests/e2e/test_phase1.py index 9dbae73e..82e1b2c5 100644 --- a/tests/e2e/test_phase1.py +++ b/tests/e2e/test_phase1.py @@ -75,9 +75,9 @@ def test_mvn_cache_hit_second_run(self, tmp_path: Path) -> None: second_meta = json.load(fh) second_ts = second_meta["timestamp_utc"] - assert ( - first_ts == second_ts - ), f"Cache hit expected but timestamp changed: {first_ts!r} → {second_ts!r}" + assert first_ts == second_ts, ( + f"Cache hit expected but timestamp changed: {first_ts!r} → {second_ts!r}" + ) def test_mvn_force_regenerates(self, tmp_path: Path) -> None: """--force must update the timestamp (regeneration happened).""" @@ -94,9 +94,9 @@ def test_mvn_force_regenerates(self, tmp_path: Path) -> None: second_meta = json.load(fh) second_ts = second_meta["timestamp_utc"] - assert ( - first_ts != second_ts - ), f"--force expected regeneration but timestamp did not change: {first_ts!r}" + assert first_ts != second_ts, ( + f"--force expected regeneration but timestamp did not change: {first_ts!r}" + ) def test_mvn_output_contains_summary(self, tmp_path: Path) -> None: """CLI must print a summary table with expected fields.""" diff --git a/tests/e2e/test_phase3.py b/tests/e2e/test_phase3.py index 797a94de..7e74ab17 100644 --- a/tests/e2e/test_phase3.py +++ b/tests/e2e/test_phase3.py @@ -47,32 +47,32 @@ class TestLeaderboardCLI: def test_leaderboard_mvn_10_markdown(self, tmp_path: Path) -> None: """tuningfork leaderboard mvn_10 must exit 0 and print markdown table.""" result = _run_leaderboard(["mvn_10"], tmp_path) - assert ( - result.returncode == 0 - ), f"Expected exit 0.\nstdout:\n{result.stdout}\nstderr:\n{result.stderr}" + assert result.returncode == 0, ( + f"Expected exit 0.\nstdout:\n{result.stdout}\nstderr:\n{result.stderr}" + ) stdout = result.stdout # Table header and title must be present - assert ( - "Leaderboard for mvn_10" in stdout - ), f"'Leaderboard for mvn_10' missing from stdout:\n{stdout}" - assert ( - "| effort |" in stdout - ), f"Markdown table header missing from stdout:\n{stdout}" + assert "Leaderboard for mvn_10" in stdout, ( + f"'Leaderboard for mvn_10' missing from stdout:\n{stdout}" + ) + assert "| effort |" in stdout, ( + f"Markdown table header missing from stdout:\n{stdout}" + ) assert "|" in stdout, f"Table pipes missing from stdout:\n{stdout}" # At least a few rows expected lines = stdout.strip().split("\n") table_lines = [line for line in lines if line.startswith("|")] # header + separator + at least 2 data rows - assert ( - len(table_lines) >= 4 - ), f"Expected at least 2 data rows in markdown table, got {len(table_lines)} table lines:\n{stdout}" + assert len(table_lines) >= 4, ( + f"Expected at least 2 data rows in markdown table, got {len(table_lines)} table lines:\n{stdout}" + ) def test_leaderboard_mvn_10_effort_medium(self, tmp_path: Path) -> None: """tuningfork leaderboard mvn_10 --effort medium filters to MEDIUM rows.""" result = _run_leaderboard(["mvn_10", "--effort", "medium"], tmp_path) - assert ( - result.returncode == 0 - ), f"Expected exit 0.\nstdout:\n{result.stdout}\nstderr:\n{result.stderr}" + assert result.returncode == 0, ( + f"Expected exit 0.\nstdout:\n{result.stdout}\nstderr:\n{result.stderr}" + ) stdout = result.stdout assert "Leaderboard for mvn_10" in stdout # Count MEDIUM-effort rows (ignoring header/separator) @@ -82,29 +82,29 @@ def test_leaderboard_mvn_10_effort_medium(self, tmp_path: Path) -> None: for line in lines if line.startswith("|") and "effort" not in line and "-----" not in line ] - assert ( - len(data_lines) >= 1 - ), f"Expected at least 1 MEDIUM-effort row, got {len(data_lines)} rows:\n{stdout}" + assert len(data_lines) >= 1, ( + f"Expected at least 1 MEDIUM-effort row, got {len(data_lines)} rows:\n{stdout}" + ) # All data rows should have "medium" in the effort column. for line in data_lines: parts = line.split("|") effort_col = parts[1].strip() - assert ( - effort_col == "medium" - ), f"Expected 'medium' in effort column, got '{effort_col}' in line:\n{line}" + assert effort_col == "medium", ( + f"Expected 'medium' in effort column, got '{effort_col}' in line:\n{line}" + ) def test_leaderboard_mvn_10_json_format(self, tmp_path: Path) -> None: """tuningfork leaderboard mvn_10 --format json must output JSON list.""" result = _run_leaderboard(["mvn_10", "--format", "json"], tmp_path) - assert ( - result.returncode == 0 - ), f"Expected exit 0.\nstdout:\n{result.stdout}\nstderr:\n{result.stderr}" + assert result.returncode == 0, ( + f"Expected exit 0.\nstdout:\n{result.stdout}\nstderr:\n{result.stderr}" + ) # Parse JSON output data = json.loads(result.stdout) assert isinstance(data, list), f"Expected JSON list, got {type(data)}" - assert ( - len(data) >= 1 - ), f"Expected at least 1 element in JSON list, got {len(data)}" + assert len(data) >= 1, ( + f"Expected at least 1 element in JSON list, got {len(data)}" + ) # Verify schema for item in data: assert "effort" in item @@ -116,11 +116,11 @@ def test_leaderboard_mvn_10_json_format(self, tmp_path: Path) -> None: def test_leaderboard_bad_model_exits_2(self, tmp_path: Path) -> None: """tuningfork leaderboard does_not_exist must exit 2 and mention model in stderr.""" result = _run_leaderboard(["does_not_exist"], tmp_path) - assert ( - result.returncode == 2 - ), f"Expected exit 2 for bad model.\nstdout:\n{result.stdout}\nstderr:\n{result.stderr}" + assert result.returncode == 2, ( + f"Expected exit 2 for bad model.\nstdout:\n{result.stdout}\nstderr:\n{result.stderr}" + ) # Stderr should mention model or list valid models stderr_lower = result.stderr.lower() - assert ( - "does_not_exist" in result.stderr or "model" in stderr_lower - ), f"Expected bad-model name or 'model' in stderr:\n{result.stderr}" + assert "does_not_exist" in result.stderr or "model" in stderr_lower, ( + f"Expected bad-model name or 'model' in stderr:\n{result.stderr}" + ) diff --git a/tests/groundtruth/test_analytic_iid.py b/tests/groundtruth/test_analytic_iid.py index f891353d..e0800b69 100644 --- a/tests/groundtruth/test_analytic_iid.py +++ b/tests/groundtruth/test_analytic_iid.py @@ -84,12 +84,12 @@ def test_analytic_iid_smoke(model_name: str, tmp_path: Path) -> None: with np.load(str(draws_path), allow_pickle=True) as draws: for site in draws.files: shape = draws[site].shape - assert ( - shape[0] == _SMOKE_N_CHAINS - ), f"{model_name}.{site}: expected n_chains={_SMOKE_N_CHAINS}, got {shape[0]}" - assert ( - shape[1] == _SMOKE_N_DRAWS - ), f"{model_name}.{site}: expected n_draws={_SMOKE_N_DRAWS}, got {shape[1]}" + assert shape[0] == _SMOKE_N_CHAINS, ( + f"{model_name}.{site}: expected n_chains={_SMOKE_N_CHAINS}, got {shape[0]}" + ) + assert shape[1] == _SMOKE_N_DRAWS, ( + f"{model_name}.{site}: expected n_draws={_SMOKE_N_DRAWS}, got {shape[1]}" + ) # --- crude mean coherence vs committed GT --- # At smoke scale (2 chains × 50 draws = 100 samples), the new mean can @@ -123,9 +123,9 @@ def test_analytic_iid_smoke(model_name: str, tmp_path: Path) -> None: "wrong distribution" ) # Also check means are finite (catches NaN/Inf in sampler output) - assert np.all( - np.isfinite(new_mean) - ), f"{model_name}.{site}: new mean contains non-finite values" + assert np.all(np.isfinite(new_mean)), ( + f"{model_name}.{site}: new mean contains non-finite values" + ) def test_analytic_iid_verify_integration(tmp_path: Path) -> None: @@ -197,6 +197,6 @@ def test_analytic_iid_custom_seed_differs(tmp_path: Path) -> None: np.load(str(out2 / "draws.npz"), allow_pickle=True) as d2, ): site = d1.files[0] - assert not np.allclose( - d1[site], d2[site] - ), "Different seeds produced identical draws — RNG not working correctly" + assert not np.allclose(d1[site], d2[site]), ( + "Different seeds produced identical draws — RNG not working correctly" + ) diff --git a/tests/groundtruth/test_dispatch.py b/tests/groundtruth/test_dispatch.py index de17e38f..bab7d069 100644 --- a/tests/groundtruth/test_dispatch.py +++ b/tests/groundtruth/test_dispatch.py @@ -70,9 +70,9 @@ def test_committed_gt_dir_returns_valid_path() -> None: for model_name in _EXPECTED: gt_dir = committed_gt_dir(model_name) assert gt_dir.is_dir(), f"GT dir missing for {model_name!r}: {gt_dir}" - assert ( - gt_dir / "summary_v2.json" - ).exists(), f"summary_v2.json missing for {model_name!r}" + assert (gt_dir / "summary_v2.json").exists(), ( + f"summary_v2.json missing for {model_name!r}" + ) assert (gt_dir / "draws.npz").exists(), f"draws.npz missing for {model_name!r}" @@ -80,9 +80,9 @@ def test_load_committed_summary_schema_version() -> None: """All committed summaries have schema_version='gt_v2_multichain'.""" for model_name in _EXPECTED: summary = load_committed_summary(model_name) - assert ( - summary["schema_version"] == "gt_v2_multichain" - ), f"Model {model_name!r}: unexpected schema_version {summary['schema_version']!r}" + assert summary["schema_version"] == "gt_v2_multichain", ( + f"Model {model_name!r}: unexpected schema_version {summary['schema_version']!r}" + ) def test_resolve_gt_method_unknown_generator_raises() -> None: diff --git a/tests/groundtruth/test_gp_marginal.py b/tests/groundtruth/test_gp_marginal.py index d479c825..a2c23492 100644 --- a/tests/groundtruth/test_gp_marginal.py +++ b/tests/groundtruth/test_gp_marginal.py @@ -71,9 +71,9 @@ def test_gp_marginal_committed_has_init_positions() -> None: sites = list(pos.keys()) # Should contain the three hyperparameter sites expected_sites = {"log_lengthscale", "log_kernel_scale", "log_noise_scale"} - assert expected_sites.issubset( - set(sites) - ), f"Missing sites: {expected_sites - set(sites)}" + assert expected_sites.issubset(set(sites)), ( + f"Missing sites: {expected_sites - set(sites)}" + ) n_chains = len(pos[sites[0]]) assert n_chains > 0 @@ -150,7 +150,9 @@ def test_gp_marginal_smoke(tmp_path: Path) -> None: _SMOKE_N_CHAINS, _SMOKE_N_DRAWS, 200, - ), f"f_raw: expected ({_SMOKE_N_CHAINS}, {_SMOKE_N_DRAWS}, 200), got {f_raw_shape}" + ), ( + f"f_raw: expected ({_SMOKE_N_CHAINS}, {_SMOKE_N_DRAWS}, 200), got {f_raw_shape}" + ) # f_raw must be finite assert np.all(np.isfinite(draws["f_raw"])), "f_raw contains non-finite values" @@ -187,9 +189,9 @@ def test_gp_marginal_smoke(tmp_path: Path) -> None: # The committed GT has 10 chains × 10k draws; at 2×50 smoke scale the # per-dim std of f_raw is roughly the posterior std, so 10× that is a wide # but meaningful sanity bound. - assert ( - "f_raw" in committed_draws.files - ), "committed draws.npz missing f_raw site" + assert "f_raw" in committed_draws.files, ( + "committed draws.npz missing f_raw site" + ) gen_f_raw = draws["f_raw"] # shape (nc, nd, 200) com_f_raw = committed_draws["f_raw"] # shape (nc_com, nd_com, 200) diff --git a/tests/groundtruth/test_nuts_multichain.py b/tests/groundtruth/test_nuts_multichain.py index 4316ec0a..deae159b 100644 --- a/tests/groundtruth/test_nuts_multichain.py +++ b/tests/groundtruth/test_nuts_multichain.py @@ -77,7 +77,7 @@ def test_load_explicit_positions_lotka_volterra() -> None: assert n_chains > 0, "must have at least one chain" for site in sites: assert len(pos_dict[site]) == n_chains, ( - f"site {site!r}: expected {n_chains} chains, " f"got {len(pos_dict[site])}" + f"site {site!r}: expected {n_chains} chains, got {len(pos_dict[site])}" ) @@ -116,9 +116,9 @@ def test_nuts_multichain_smoke_radon(tmp_path: Path) -> None: assert "max_rhat" in gate assert "min_bulk_ess" in gate # R̂ < 2.0 at 2×100 is a very loose sanity check (catches sampler failures) - assert ( - gate["max_rhat"] < 2.0 - ), f"{model_name}: max_rhat={gate['max_rhat']:.4f} >= 2.0" + assert gate["max_rhat"] < 2.0, ( + f"{model_name}: max_rhat={gate['max_rhat']:.4f} >= 2.0" + ) # --- diagnostics present --- dpc = result["diagnostics_per_chain"] @@ -135,12 +135,12 @@ def test_nuts_multichain_smoke_radon(tmp_path: Path) -> None: with np.load(str(draws_path), allow_pickle=True) as draws: for site in draws.files: shape = draws[site].shape - assert ( - shape[0] == _SMOKE_N_CHAINS - ), f"{model_name}.{site}: expected n_chains={_SMOKE_N_CHAINS}, got {shape[0]}" - assert ( - shape[1] == _SMOKE_N_DRAWS - ), f"{model_name}.{site}: expected n_draws={_SMOKE_N_DRAWS}, got {shape[1]}" + assert shape[0] == _SMOKE_N_CHAINS, ( + f"{model_name}.{site}: expected n_chains={_SMOKE_N_CHAINS}, got {shape[0]}" + ) + assert shape[1] == _SMOKE_N_DRAWS, ( + f"{model_name}.{site}: expected n_draws={_SMOKE_N_DRAWS}, got {shape[1]}" + ) # --- crude mean coherence vs committed GT --- # At 2×100 draws, use 10× posterior std as the tolerance. @@ -165,9 +165,9 @@ def test_nuts_multichain_smoke_radon(tmp_path: Path) -> None: f"{model_name}.{site}: new mean deviates " f"{max_dev:.2f}× committed posterior std" ) - assert np.all( - np.isfinite(new_mean) - ), f"{model_name}.{site}: new mean contains non-finite values" + assert np.all(np.isfinite(new_mean)), ( + f"{model_name}.{site}: new mean contains non-finite values" + ) @pytest.mark.slow @@ -232,9 +232,9 @@ def test_nuts_multichain_custom_seed_differs(tmp_path: Path) -> None: np.load(str(out2 / "draws.npz"), allow_pickle=True) as d2, ): site = d1.files[0] - assert not np.allclose( - d1[site], d2[site] - ), "Different seeds produced identical draws — RNG not working correctly" + assert not np.allclose(d1[site], d2[site]), ( + "Different seeds produced identical draws — RNG not working correctly" + ) # --------------------------------------------------------------------------- # @@ -314,9 +314,9 @@ def test_nuts_multichain_precision_emitted_in_sampler_config(tmp_path: Path) -> """ committed = load_committed_summary("radon") result = generate_nuts_multichain("radon", committed, tmp_path, smoke=True) - assert ( - "precision" in result["sampler_config"] - ), "sampler_config must include 'precision' key for artifact self-description" + assert "precision" in result["sampler_config"], ( + "sampler_config must include 'precision' key for artifact self-description" + ) assert result["sampler_config"]["precision"] == "float32" @@ -349,9 +349,9 @@ def test_nuts_multichain_horseshoe_x64_path(tmp_path: Path) -> None: with np.load(str(draws_path), allow_pickle=True) as draws: for site in draws.files: arr = draws[site] - assert ( - arr.dtype == np.float64 - ), f"horseshoe.{site}: expected float64 draws, got {arr.dtype}" + assert arr.dtype == np.float64, ( + f"horseshoe.{site}: expected float64 draws, got {arr.dtype}" + ) # (c) x64 flag must be restored to the value it had before generation restored_x64 = jax.config.read("jax_enable_x64") diff --git a/tests/groundtruth/test_verify.py b/tests/groundtruth/test_verify.py index face0b94..a1ab91b8 100644 --- a/tests/groundtruth/test_verify.py +++ b/tests/groundtruth/test_verify.py @@ -370,9 +370,9 @@ def test_coherence_materiality_review_pass() -> None: com = _coh_summary({"x": _per_site_summary([0.0], [se], std=[std_c])}) passed, results, meta = _check_coherence(gen, com) - assert ( - passed - ), "Large z but immaterial (mat << TAU_SCI) should be REVIEW (counts as pass)" + assert passed, ( + "Large z but immaterial (mat << TAU_SCI) should be REVIEW (counts as pass)" + ) assert results[0]["verdict"] == "REVIEW" assert len(results[0]["review_dims"]) > 0 assert len(results[0]["hard_fail_dims"]) == 0 @@ -444,9 +444,9 @@ def _make(D: int, n_chains: int = 10) -> dict: assert meta2["D_total"] == D2 expected_z2 = float(scipy_stats.t.ppf(1.0 - alpha / (2.0 * D2), 18)) assert meta2["z_crit"] == pytest.approx(expected_z2, rel=1e-5) - assert meta2["z_crit"] == pytest.approx( - 3.6281, rel=1e-3 - ), f"z_crit for D={D2}: got {meta2['z_crit']:.4f}, expected ≈3.6281." + assert meta2["z_crit"] == pytest.approx(3.6281, rel=1e-3), ( + f"z_crit for D={D2}: got {meta2['z_crit']:.4f}, expected ≈3.6281." + ) def test_coherence_shape_shrink_fails() -> None: @@ -496,9 +496,9 @@ def test_coherence_materiality_boundary_strict_gt() -> None: gen_at = _coh_summary({"x": _per_site_summary([delta_at], [se])}) com_at = _coh_summary({"x": _per_site_summary([0.0], [se], std=[std_c])}) passed_at, results_at, _ = _check_coherence(gen_at, com_at) - assert ( - passed_at - ), f"mat exactly at boundary ({_TAU_SCI}) should be REVIEW (strict >), not FAIL." + assert passed_at, ( + f"mat exactly at boundary ({_TAU_SCI}) should be REVIEW (strict >), not FAIL." + ) assert results_at[0]["verdict"] == "REVIEW" # Just above the boundary: mat = TAU_SCI + epsilon → FAIL (strict > is True) @@ -506,9 +506,9 @@ def test_coherence_materiality_boundary_strict_gt() -> None: gen_above = _coh_summary({"x": _per_site_summary([delta_above], [se])}) com_above = _coh_summary({"x": _per_site_summary([0.0], [se], std=[std_c])}) passed_above, results_above, _ = _check_coherence(gen_above, com_above) - assert ( - not passed_above - ), f"mat just above boundary ({_TAU_SCI}+eps) should hard-FAIL." + assert not passed_above, ( + f"mat just above boundary ({_TAU_SCI}+eps) should hard-FAIL." + ) assert results_above[0]["verdict"] == "FAIL" diff --git a/tests/inference/warmup/test_vi_warmup_hp_space.py b/tests/inference/warmup/test_vi_warmup_hp_space.py index b66b9f6e..5198e331 100644 --- a/tests/inference/warmup/test_vi_warmup_hp_space.py +++ b/tests/inference/warmup/test_vi_warmup_hp_space.py @@ -83,9 +83,9 @@ def test_other_warmups_have_empty_hp_space() -> None: if name in _known_non_empty_hp_space: continue hp_space = getattr(entry, "default_hp_space", ()) - assert ( - hp_space == () or len(hp_space) == 0 - ), f"Non-VI warmup {name!r} unexpectedly has default_hp_space: {hp_space}" + assert hp_space == () or len(hp_space) == 0, ( + f"Non-VI warmup {name!r} unexpectedly has default_hp_space: {hp_space}" + ) def test_low_rank_warmup_has_max_rank_hp_space() -> None: @@ -153,6 +153,6 @@ def test_warmup_hp_space_override_roundtrip() -> None: if any(k == s.name for s in getattr(mfwu, "default_hp_space", ())) } ) - assert ( - merged["num_optimization_steps"] == 10_000 - ), "warmup_kwargs_override must override the HP default" + assert merged["num_optimization_steps"] == 10_000, ( + "warmup_kwargs_override must override the HP default" + ) diff --git a/tests/metrics/test_grad_counter.py b/tests/metrics/test_grad_counter.py index 68ab2ceb..7be52dc6 100644 --- a/tests/metrics/test_grad_counter.py +++ b/tests/metrics/test_grad_counter.py @@ -112,9 +112,9 @@ def test_mala_returns_python_int(self) -> None: """total_grad_evals must return a plain Python int, not a JAX Array.""" infos = FakeConstantInfo(accepted=jnp.ones((10,), dtype=jnp.bool_)) result = total_grad_evals(infos, lambda i: 1) - assert isinstance( - result, int - ), f"Expected Python int, got {type(result).__name__}" + assert isinstance(result, int), ( + f"Expected Python int, got {type(result).__name__}" + ) # --------------------------------------------------------------------------- diff --git a/tests/metrics/test_headline.py b/tests/metrics/test_headline.py index f500c749..458f21cd 100644 --- a/tests/metrics/test_headline.py +++ b/tests/metrics/test_headline.py @@ -89,9 +89,9 @@ def test_iid_returns_float(self) -> None: """Return type must be a Python float, not a JAX Array.""" samples = {"x": jax.random.normal(_key(1), (2, 100, 3))} headline = min_bulk_ess_per_grad(samples, n_grad_evals=200) - assert isinstance( - headline, float - ), f"Expected Python float, got {type(headline).__name__}" + assert isinstance(headline, float), ( + f"Expected Python float, got {type(headline).__name__}" + ) def test_iid_positive(self) -> None: """Headline must be strictly positive for valid i.i.d. samples.""" @@ -141,9 +141,9 @@ def test_worst_site_governs(self) -> None: headline_good_only = min_bulk_ess_per_grad({"good": good_site}, n_grad_evals) headline_both = min_bulk_ess_per_grad(sites, n_grad_evals) - assert ( - headline_both < headline_good_only - ), "Adding a bad site should lower the headline metric" + assert headline_both < headline_good_only, ( + "Adding a bad site should lower the headline metric" + ) # --------------------------------------------------------------------------- diff --git a/tests/metrics/test_reference_compare.py b/tests/metrics/test_reference_compare.py index aea894cb..bed09e40 100644 --- a/tests/metrics/test_reference_compare.py +++ b/tests/metrics/test_reference_compare.py @@ -119,8 +119,7 @@ def test_dict_and_array_inputs_agree(self) -> None: for key in result_dict: assert abs(result_dict[key] - result_arr[key]) < 1e-10, ( - f"Mismatch on {key!r}: dict={result_dict[key]}, " - f"array={result_arr[key]}" + f"Mismatch on {key!r}: dict={result_dict[key]}, array={result_arr[key]}" ) @@ -302,9 +301,9 @@ def test_mean_shift_recoverable(self) -> None: } result = compute_sample_quality(draws, ref) # The mean is shifted by k, so mae_vs_reference ≈ k. - assert ( - abs(result["mae_vs_reference"] - k) < 0.5 - ), f"Expected mae_vs_reference ≈ {k}, got {result['mae_vs_reference']:.4f}" + assert abs(result["mae_vs_reference"] - k) < 0.5, ( + f"Expected mae_vs_reference ≈ {k}, got {result['mae_vs_reference']:.4f}" + ) def test_reference_std_normalization_not_empirical(self) -> None: """Draws with doubled std → std_ratio_max_dev ≈ 1.0, not 0.0. @@ -398,9 +397,9 @@ def test_perfect_scalar_draws_near_zero(self) -> None: draws = {"x": rng.standard_normal((4, 4000, 1))} ref = {"x": {"mean": 0.0, "std": 1.0, "q05": -1.645, "q95": 1.645}} sq = compute_sample_quality(draws, ref) - assert ( - sq["std_ratio_max_dev"] < 0.05 - ), f"scalar std_ratio_max_dev={sq['std_ratio_max_dev']:.4f}" + assert sq["std_ratio_max_dev"] < 0.05, ( + f"scalar std_ratio_max_dev={sq['std_ratio_max_dev']:.4f}" + ) assert sq["q05_error"] < 0.1, f"scalar q05_error={sq['q05_error']:.4f}" assert sq["q95_error"] < 0.1, f"scalar q95_error={sq['q95_error']:.4f}" @@ -459,6 +458,6 @@ def test_per_element_ref_used_not_grand_mean(self) -> None: } draws_dict = {"x": draws_2d} sq = compute_sample_quality(draws_dict, ref_dict) - assert ( - sq["std_ratio_max_dev"] < 0.1 - ), f"Per-element-std test failed: std_ratio_max_dev={sq['std_ratio_max_dev']:.4f}" + assert sq["std_ratio_max_dev"] < 0.1, ( + f"Per-element-std test failed: std_ratio_max_dev={sq['std_ratio_max_dev']:.4f}" + ) diff --git a/tests/models/test_models_parametrized.py b/tests/models/test_models_parametrized.py index d5689ddd..e9de5797 100644 --- a/tests/models/test_models_parametrized.py +++ b/tests/models/test_models_parametrized.py @@ -36,9 +36,9 @@ @pytest.mark.parametrize("model_name", _MODELS) def test_model_registered(model_name: str) -> None: """Model is registered in MODELS.""" - assert ( - model_name in MODELS - ), f"Model '{model_name}' not found in MODELS; registered: {sorted(MODELS)}" + assert model_name in MODELS, ( + f"Model '{model_name}' not found in MODELS; registered: {sorted(MODELS)}" + ) @pytest.mark.parametrize("model_name", _MODELS) @@ -46,6 +46,6 @@ def test_model_has_dim_and_class(model_name: str) -> None: """ENTRY has dim (positive int) and class_ (non-empty string label).""" entry = MODELS[model_name] assert isinstance(entry.dim, int) and entry.dim > 0 - assert ( - isinstance(entry.class_, str) and entry.class_ - ), f"Model '{model_name}' must have a non-empty string class_; got {entry.class_!r}" + assert isinstance(entry.class_, str) and entry.class_, ( + f"Model '{model_name}' must have a non-empty string class_; got {entry.class_!r}" + ) diff --git a/tests/models/test_posterior_entry.py b/tests/models/test_posterior_entry.py index 46c55b1e..72fd2b53 100644 --- a/tests/models/test_posterior_entry.py +++ b/tests/models/test_posterior_entry.py @@ -163,6 +163,6 @@ def test_all_other_models_use_default_acceptance(self): for name, entry in MODELS.items() if name != "lgcp" and entry.reference_target_acceptance != 0.80 } - assert ( - deviations == {} - ), f"Unexpected non-default reference_target_acceptance: {deviations}" + assert deviations == {}, ( + f"Unexpected non-default reference_target_acceptance: {deviations}" + ) diff --git a/tests/notebooks/test_inspect.py b/tests/notebooks/test_inspect.py index 0f14222d..01f2559f 100644 --- a/tests/notebooks/test_inspect.py +++ b/tests/notebooks/test_inspect.py @@ -253,17 +253,17 @@ def test_summarize_recipe_sample_budget_rows_low_recipe() -> None: props = dict(zip(df["Property"].tolist(), df["Value"].tolist())) assert "num_chains" in props, "summarize_recipe must include 'num_chains' row" - assert ( - props["num_chains"] == "4" - ), f"Expected num_chains='4' for LOW recipe, got {props['num_chains']!r}" + assert props["num_chains"] == "4", ( + f"Expected num_chains='4' for LOW recipe, got {props['num_chains']!r}" + ) assert "n_warmup" in props, "summarize_recipe must include 'n_warmup' row" - assert ( - props["n_warmup"] == "1000" - ), f"Expected n_warmup='1000' for LOW recipe, got {props['n_warmup']!r}" + assert props["n_warmup"] == "1000", ( + f"Expected n_warmup='1000' for LOW recipe, got {props['n_warmup']!r}" + ) assert "n_samples" in props, "summarize_recipe must include 'n_samples' row" - assert ( - props["n_samples"] == "1000" - ), f"Expected n_samples='1000' for LOW recipe, got {props['n_samples']!r}" + assert props["n_samples"] == "1000", ( + f"Expected n_samples='1000' for LOW recipe, got {props['n_samples']!r}" + ) def test_summarize_recipe_sample_budget_rows_legacy_groundtruth( @@ -282,17 +282,17 @@ def test_summarize_recipe_sample_budget_rows_legacy_groundtruth( props = dict(zip(df["Property"].tolist(), df["Value"].tolist())) # num_chains absent from both warmup_params and calibration_budget - assert ( - props.get("num_chains") == "N/A" - ), f"Expected num_chains='N/A' for legacy recipe, got {props.get('num_chains')!r}" + assert props.get("num_chains") == "N/A", ( + f"Expected num_chains='N/A' for legacy recipe, got {props.get('num_chains')!r}" + ) # n_warmup present in warmup_params (n_warmup=1000 in the fixture) - assert ( - props.get("n_warmup") == "1000" - ), f"Expected n_warmup='1000', got {props.get('n_warmup')!r}" + assert props.get("n_warmup") == "1000", ( + f"Expected n_warmup='1000', got {props.get('n_warmup')!r}" + ) # n_samples absent - assert ( - props.get("n_samples") == "N/A" - ), f"Expected n_samples='N/A' for legacy recipe, got {props.get('n_samples')!r}" + assert props.get("n_samples") == "N/A", ( + f"Expected n_samples='N/A' for legacy recipe, got {props.get('n_samples')!r}" + ) # --------------------------------------------------------------------------- @@ -335,9 +335,9 @@ def test_summarize_recipe_warmup_inner_kernel_shown_when_set() -> None: "Schema extension: summarize_recipe must include 'warmup_inner_kernel' row " "when recipe.warmup_inner_kernel is explicitly set." ) - assert ( - props["warmup_inner_kernel"] == "nuts" - ), f"Expected warmup_inner_kernel='nuts', got {props['warmup_inner_kernel']!r}" + assert props["warmup_inner_kernel"] == "nuts", ( + f"Expected warmup_inner_kernel='nuts', got {props['warmup_inner_kernel']!r}" + ) def test_summarize_recipe_warmup_inner_kernel_absent_when_none() -> None: diff --git a/tests/notebooks/test_interactive_helpers.py b/tests/notebooks/test_interactive_helpers.py index 05c62fb4..6d855b3a 100644 --- a/tests/notebooks/test_interactive_helpers.py +++ b/tests/notebooks/test_interactive_helpers.py @@ -66,9 +66,9 @@ def test_all_14_models_have_headline_fields() -> None: assert isinstance(hc, dict), f"{name} headline_coords not dict: {hc!r}" for k, v in hc.items(): assert isinstance(k, str), f"{name} headline_coords key not str: {k!r}" - assert isinstance(v, list) and all( - isinstance(i, int) for i in v - ), f"{name} headline_coords value not list[int]: {v!r}" + assert isinstance(v, list) and all(isinstance(i, int) for i in v), ( + f"{name} headline_coords value not list[int]: {v!r}" + ) def test_headline_params_per_decision_doc() -> None: @@ -102,9 +102,9 @@ def test_headline_coords_per_decision_doc() -> None: for name in MODELS: if name == "german_credit": continue - assert ( - MODELS[name].headline_coords is None - ), f"{name} expected None headline_coords, got {MODELS[name].headline_coords!r}" + assert MODELS[name].headline_coords is None, ( + f"{name} expected None headline_coords, got {MODELS[name].headline_coords!r}" + ) # --------------------------------------------------------------------------- diff --git a/tests/notebooks/test_rerun_inference.py b/tests/notebooks/test_rerun_inference.py index ea720213..865fc27e 100644 --- a/tests/notebooks/test_rerun_inference.py +++ b/tests/notebooks/test_rerun_inference.py @@ -29,11 +29,14 @@ from dataclasses import replace from pathlib import Path from types import SimpleNamespace -from typing import Any, cast +from typing import TYPE_CHECKING, Any, cast import numpy as np import pytest +if TYPE_CHECKING: + from tuningfork.recipes._base import Recipe + pytestmark = pytest.mark.fast @@ -135,7 +138,7 @@ def test_load_from_cache_rejects_stat_name_collisions(tmp_path: Path) -> None: # --------------------------------------------------------------------------- -def _failed_recipe(): +def _failed_recipe() -> Recipe: from tuningfork.recipes._base import Effort, Recipe return Recipe( diff --git a/tests/numpyro/test_helper.py b/tests/numpyro/test_helper.py index 4f458c4c..6f6bf642 100644 --- a/tests/numpyro/test_helper.py +++ b/tests/numpyro/test_helper.py @@ -80,9 +80,9 @@ def test_logdensity_at_mode_is_maximum(self): # logdensity at mode (0,0,0) should exceed logdensity at (2,2,2) at_mode = logdensity_fn({"x": jnp.zeros(3)}) away_from_mode = logdensity_fn({"x": jnp.array([2.0, 2.0, 2.0])}) - assert ( - at_mode > away_from_mode - ), f"Expected logdensity at mode > away: {at_mode:.3f} vs {away_from_mode:.3f}" + assert at_mode > away_from_mode, ( + f"Expected logdensity at mode > away: {at_mode:.3f} vs {away_from_mode:.3f}" + ) def test_logdensity_fn_is_differentiable(self): key = jax.random.key(3) @@ -142,6 +142,6 @@ def test_negative_potential(self): # Use the same init_pos as above ld = logdensity_fn(init_pos) pot = potential_fn(init_pos) - assert jnp.isclose( - ld, -pot, atol=1e-5 - ), f"logdensity_fn != -potential_fn: {ld} vs {-pot}" + assert jnp.isclose(ld, -pot, atol=1e-5), ( + f"logdensity_fn != -potential_fn: {ld} vs {-pot}" + ) diff --git a/tests/numpyro/test_smc_logfns_space.py b/tests/numpyro/test_smc_logfns_space.py index 95aad518..6d048601 100644 --- a/tests/numpyro/test_smc_logfns_space.py +++ b/tests/numpyro/test_smc_logfns_space.py @@ -88,9 +88,9 @@ def test_logprior_fn_matches_analytic_prior(self) -> None: for z in _Z_GRID: got = float(logprior_fn({"sigma": jnp.asarray(z)})) want = _analytic_logprior(z) - assert got == pytest.approx( - want, abs=1e-5 - ), f"logprior_fn({z}) = {got}, expected {want}" + assert got == pytest.approx(want, abs=1e-5), ( + f"logprior_fn({z}) = {got}, expected {want}" + ) @pytest.mark.fast def test_loglikelihood_fn_matches_analytic_likelihood(self) -> None: @@ -99,9 +99,9 @@ def test_loglikelihood_fn_matches_analytic_likelihood(self) -> None: for z in _Z_GRID: got = float(loglik_fn({"sigma": jnp.asarray(z)})) want = _analytic_loglik(z) - assert got == pytest.approx( - want, abs=1e-5 - ), f"loglikelihood_fn({z}) = {got}, expected {want}" + assert got == pytest.approx(want, abs=1e-5), ( + f"loglikelihood_fn({z}) = {got}, expected {want}" + ) @pytest.mark.fast def test_logprior_plus_loglikelihood_equals_negative_potential(self) -> None: @@ -126,9 +126,9 @@ def test_logprior_plus_loglikelihood_equals_negative_potential(self) -> None: position = {"sigma": jnp.asarray(z)} total = float(logprior_fn(position)) + float(loglik_fn(position)) expected = float(-potential_fn(position)) - assert total == pytest.approx( - expected, abs=1e-5 - ), f"logprior+loglik={total}, -potential_fn={expected} at z={z}" + assert total == pytest.approx(expected, abs=1e-5), ( + f"logprior+loglik={total}, -potential_fn={expected} at z={z}" + ) @pytest.mark.fast def test_blocked_model_latent_sites_match_joint_model(self) -> None: diff --git a/tests/recipes/test_auto_gate_w1_integration.py b/tests/recipes/test_auto_gate_w1_integration.py index ca4063bc..1b2718b3 100644 --- a/tests/recipes/test_auto_gate_w1_integration.py +++ b/tests/recipes/test_auto_gate_w1_integration.py @@ -168,12 +168,12 @@ def test_w1_realm_skipped_when_rhat_fails(): assert verdict.rhat_max is not None and verdict.rhat_max >= 1.05 # W1 realm must have been skipped - assert ( - verdict.w1_realm_result is None - ), "Expected W1 realm to be skipped when R̂ FAILs" - assert ( - "w1_realm" not in verdict.margins - ), "Expected no 'w1_realm' key in margins when W1 was skipped" + assert verdict.w1_realm_result is None, ( + "Expected W1 realm to be skipped when R̂ FAILs" + ) + assert "w1_realm" not in verdict.margins, ( + "Expected no 'w1_realm' key in margins when W1 was skipped" + ) assert verdict.verdict == "FAIL" @@ -254,9 +254,9 @@ def test_w1_null_case_does_not_add_top_level_to_dict_key(): "verdict", "margins", } - assert ( - set(d.keys()) == expected_top_level - ), f"Unexpected top-level keys in to_dict(): {set(d.keys()) - expected_top_level}" + assert set(d.keys()) == expected_top_level, ( + f"Unexpected top-level keys in to_dict(): {set(d.keys()) - expected_top_level}" + ) # W1 lives under margins assert "w1_realm" in d["margins"] @@ -284,10 +284,10 @@ def test_w1_verdict_propagates_to_overall(): assert verdict.w1_realm_result is not None w1_m = verdict.margins["w1_realm"] - assert ( - w1_m["max_prong_verdict"] == "FAIL" - ), f"Expected max prong FAIL with +5σ shift, got {w1_m['max_prong_verdict']}" + assert w1_m["max_prong_verdict"] == "FAIL", ( + f"Expected max prong FAIL with +5σ shift, got {w1_m['max_prong_verdict']}" + ) # Overall verdict must propagate the W1 FAIL - assert ( - verdict.verdict == "FAIL" - ), f"Expected overall FAIL from W1 prong, got {verdict.verdict}" + assert verdict.verdict == "FAIL", ( + f"Expected overall FAIL from W1 prong, got {verdict.verdict}" + ) diff --git a/tests/recipes/test_benchmark_regression.py b/tests/recipes/test_benchmark_regression.py index 2fd20a5f..fea15a00 100644 --- a/tests/recipes/test_benchmark_regression.py +++ b/tests/recipes/test_benchmark_regression.py @@ -734,9 +734,9 @@ def call_fn(fn): ) # All 3 seeds must be present in the returned dict - assert ( - set(per_seed.keys()) == expected_seeds - ), f"Expected seeds {expected_seeds}, got {set(per_seed.keys())}" + assert set(per_seed.keys()) == expected_seeds, ( + f"Expected seeds {expected_seeds}, got {set(per_seed.keys())}" + ) # per_seed_metrics must be stored in extra_info assert "per_seed_metrics" in mock_benchmark.extra_info @@ -1055,14 +1055,14 @@ def call_fn(fn): # Each seed's ESS should be mean(1600, 2000) = 1800 for s, m in per_seed.items(): - assert m["min_bulk_ess"] == pytest.approx( - 1800.0 - ), f"seed {s}: expected mean ESS 1800.0, got {m['min_bulk_ess']}" + assert m["min_bulk_ess"] == pytest.approx(1800.0), ( + f"seed {s}: expected mean ESS 1800.0, got {m['min_bulk_ess']}" + ) # runtime is summed: 1.0 + 1.2 = 2.2 for s, m in per_seed.items(): - assert m["runtime_warmup_s"] == pytest.approx( - 2.2 - ), f"seed {s}: expected summed runtime 2.2, got {m['runtime_warmup_s']}" + assert m["runtime_warmup_s"] == pytest.approx(2.2), ( + f"seed {s}: expected summed runtime 2.2, got {m['runtime_warmup_s']}" + ) # --------------------------------------------------------------------------- diff --git a/tests/recipes/test_catalog_render.py b/tests/recipes/test_catalog_render.py index 12dcb5c3..5019071f 100644 --- a/tests/recipes/test_catalog_render.py +++ b/tests/recipes/test_catalog_render.py @@ -17,6 +17,7 @@ treated as single-chain draws, garbling the posterior group via the cert-protocol reshape path. """ + from __future__ import annotations import json diff --git a/tests/recipes/test_codegen_boundary.py b/tests/recipes/test_codegen_boundary.py index 5c8180db..e8c75445 100644 --- a/tests/recipes/test_codegen_boundary.py +++ b/tests/recipes/test_codegen_boundary.py @@ -112,13 +112,12 @@ def test_sampling_constructor_alias_is_reported(tmp_path: Path) -> None: ), ( "import_module.py", - "import blackjax.mcmc.hmc as internal_hmc\n" - "internal_hmc.build_kernel()\n", + "import blackjax.mcmc.hmc as internal_hmc\ninternal_hmc.build_kernel()\n", "blackjax.mcmc.hmc.build_kernel", ), ( "import_module_unaliased.py", - "import blackjax.mcmc.hmc\n" "blackjax.mcmc.hmc.build_kernel()\n", + "import blackjax.mcmc.hmc\nblackjax.mcmc.hmc.build_kernel()\n", "blackjax.mcmc.hmc.build_kernel", ), ( @@ -129,7 +128,7 @@ def test_sampling_constructor_alias_is_reported(tmp_path: Path) -> None: ), ( "from_function.py", - "from blackjax.mcmc.random_walk import build_rmh\n" "build_rmh()\n", + "from blackjax.mcmc.random_walk import build_rmh\nbuild_rmh()\n", "blackjax.mcmc.random_walk.build_rmh", ), ( @@ -139,7 +138,7 @@ def test_sampling_constructor_alias_is_reported(tmp_path: Path) -> None: ), ( "top_level_api.py", - "from blackjax.mcmc.hmc import as_top_level_api\n" "as_top_level_api()\n", + "from blackjax.mcmc.hmc import as_top_level_api\nas_top_level_api()\n", "blackjax.mcmc.hmc.as_top_level_api", ), ( @@ -168,8 +167,7 @@ def test_low_level_sampling_builder_import_forms_are_reported( ), ( "getattr_builder.py", - "import blackjax.mcmc.hmc\n" - "getattr(blackjax.mcmc.hmc, 'build_kernel')()\n", + "import blackjax.mcmc.hmc\ngetattr(blackjax.mcmc.hmc, 'build_kernel')()\n", "blackjax.mcmc.hmc.build_kernel", ), ( @@ -196,7 +194,7 @@ def test_constant_dynamic_sampling_forms_are_reported( def test_unrelated_build_function_is_ignored(tmp_path: Path) -> None: (tmp_path / "helper.py").write_text( - "from project.helpers import build_sampler\n" "build_sampler(logdensity_fn)\n" + "from project.helpers import build_sampler\nbuild_sampler(logdensity_fn)\n" ) assert not boundary.scan_source(tmp_path) @@ -261,7 +259,7 @@ def test_chained_assignment_alias_is_reported(tmp_path: Path) -> None: def test_unknown_assignment_alias_is_ignored(tmp_path: Path) -> None: (tmp_path / "unknown.py").write_text( - "def sample():\n" " helper = object()\n" " return helper(logdensity_fn)\n" + "def sample():\n helper = object()\n return helper(logdensity_fn)\n" ) assert not boundary.scan_source(tmp_path) diff --git a/tests/recipes/test_emit.py b/tests/recipes/test_emit.py index 10bec7b9..13adc328 100644 --- a/tests/recipes/test_emit.py +++ b/tests/recipes/test_emit.py @@ -121,9 +121,9 @@ def test_catalog_headline_basis_reproduces_headline_metric() -> None: f"total_grad_evals={tge}). headline_basis must store the HEADLINE " f"ESS (blackjax effective_sample_size), not the gate ESS (ess_bulk)." ) - assert ( - not failures - ), "headline_basis does not reproduce headline_metric:\n" + "\n".join(failures) + assert not failures, ( + "headline_basis does not reproduce headline_metric:\n" + "\n".join(failures) + ) @pytest.mark.fast diff --git a/tests/recipes/test_emit_script.py b/tests/recipes/test_emit_script.py index 0e6eb091..5fa09aff 100644 --- a/tests/recipes/test_emit_script.py +++ b/tests/recipes/test_emit_script.py @@ -116,15 +116,15 @@ def test_emit_script_rmhmc_generated_valid_python() -> None: ) # 3. No stray blackjax.hmc calls. - assert ( - "blackjax.hmc(" not in script - ), "generated RMHMC program must not contain blackjax.hmc() calls" + assert "blackjax.hmc(" not in script, ( + "generated RMHMC program must not contain blackjax.hmc() calls" + ) # 4. IMM→mass_matrix helper present (required because upstream rmhmc takes # mass_matrix not inverse_mass_matrix) - assert ( - "_imm_to_mass_matrix" in script - ), "generated RMHMC program must define _imm_to_mass_matrix" + assert "_imm_to_mass_matrix" in script, ( + "generated RMHMC program must define _imm_to_mass_matrix" + ) @pytest.mark.fast @@ -170,9 +170,9 @@ def test_emit_script_num_samples_defaults_to_calibration_budget() -> None: # (re-stamped to the 1000×4 production config 2026-05-27 — assert dynamically # against whatever the recipe currently declares, not a hardcoded value). n_samples = recipe.calibration_budget.get("n_samples") - assert isinstance( - n_samples, int - ), f"Unexpected n_samples in calibration_budget: {recipe.calibration_budget}" + assert isinstance(n_samples, int), ( + f"Unexpected n_samples in calibration_budget: {recipe.calibration_budget}" + ) # emit_script with no num_samples arg → must use calibration_budget value script = emit_script(recipe) @@ -759,12 +759,12 @@ def test_emit_script_warmup_execution_contract_mhmc_dense(tmp_path: Path) -> Non emitted = json.load(f) # Verify structure: single-chain warmup produces scalar step_size (ndim=0). - assert emitted[ - "step_size_finite" - ], f"Emitted _adapted_params['step_size'] is not finite: {emitted['step_size']}" - assert emitted[ - "imm_finite" - ], "Emitted _adapted_params['inverse_mass_matrix'] has non-finite values" + assert emitted["step_size_finite"], ( + f"Emitted _adapted_params['step_size'] is not finite: {emitted['step_size']}" + ) + assert emitted["imm_finite"], ( + "Emitted _adapted_params['inverse_mass_matrix'] has non-finite values" + ) # Dense IMM for eight_schools_ncp (9 free params): shape should be (9, 9). assert emitted["imm_ndim"] == 2, ( f"Expected dense IMM (ndim=2) for window_adaptation_dense_imm, " @@ -801,12 +801,12 @@ def test_emit_script_laplace_high_recipe_multiphase_warmup_structure() -> None: script = emit_script(recipe, num_samples=10, num_chains=2) # Both phases must appear in the warmup section. - assert ( - "_warmup_p1" in script - ), "Phase 1 warmup (`_warmup_p1`) not found in emitted script" - assert ( - "_warmup_p2" in script - ), "Phase 2 warmup (`_warmup_p2`) not found in emitted script" + assert "_warmup_p1" in script, ( + "Phase 1 warmup (`_warmup_p1`) not found in emitted script" + ) + assert "_warmup_p2" in script, ( + "Phase 2 warmup (`_warmup_p2`) not found in emitted script" + ) # Phase 1 uses maxiter=100, Phase 2 uses maxiter=400 (from the recipe JSON). assert "maxiter=100" in script, "Expected maxiter=100 (Phase 1) in emitted script" @@ -880,9 +880,9 @@ def test_emit_script_laplace_optimizer_kwargs_persisted() -> None: assert "400" in script, "maxiter=400 value not in emitted script" assert "20" in script, "maxcor=20 value not in emitted script" # LaplaceMarginal factory call must include the kwargs - assert ( - "_lmf(log_joint_fn, theta_init" in script - ), "LaplaceMarginal factory call missing from emitted script" + assert "_lmf(log_joint_fn, theta_init" in script, ( + "LaplaceMarginal factory call missing from emitted script" + ) # --------------------------------------------------------------------------- @@ -951,19 +951,19 @@ def test_emit_script_num_warmup_multiphase_list() -> None: # Phase 1 should use 100 steps; Phase 2 should use 10. # The run() call is multi-line so check the comment header which has n_warmup=. # e.g. "# Phase 1: window_adaptation_diag_imm (n_warmup=100, ...)" - assert ( - "n_warmup=100" in script - ), "Expected 'n_warmup=100' (Phase 1 comment) in emitted script with num_warmup=[100, 10]." - assert ( - "n_warmup=10," in script - ), "Expected 'n_warmup=10,' (Phase 2 comment) in emitted script with num_warmup=[100, 10]." + assert "n_warmup=100" in script, ( + "Expected 'n_warmup=100' (Phase 1 comment) in emitted script with num_warmup=[100, 10]." + ) + assert "n_warmup=10," in script, ( + "Expected 'n_warmup=10,' (Phase 2 comment) in emitted script with num_warmup=[100, 10]." + ) # The recipe's original values (500/200) should NOT appear in the warmup comments. - assert ( - "n_warmup=500" not in script - ), "Recipe's Phase 1 n_warmup=500 should be overridden by num_warmup=[100, 10]." - assert ( - "n_warmup=200" not in script - ), "Recipe's Phase 2 n_warmup=200 should be overridden by num_warmup=[100, 10]." + assert "n_warmup=500" not in script, ( + "Recipe's Phase 1 n_warmup=500 should be overridden by num_warmup=[100, 10]." + ) + assert "n_warmup=200" not in script, ( + "Recipe's Phase 2 n_warmup=200 should be overridden by num_warmup=[100, 10]." + ) @pytest.mark.fast diff --git a/tests/recipes/test_emit_warmup_progress_bar.py b/tests/recipes/test_emit_warmup_progress_bar.py index 96de76e8..da318e4d 100644 --- a/tests/recipes/test_emit_warmup_progress_bar.py +++ b/tests/recipes/test_emit_warmup_progress_bar.py @@ -33,6 +33,7 @@ _samples leading axis = num_chains (topology is unaffected by the flag). - No vmap/io_callback errors in either mode. """ + from __future__ import annotations import dataclasses @@ -218,9 +219,9 @@ def test_progress_bar_none_default_same_as_false(warmup_name: str) -> None: recipe, num_samples=10, num_warmup=10, progress_bar=False ) - assert ( - script_none == script_false - ), f"[{warmup_name}] progress_bar=None must produce the same script as False.\n" + assert script_none == script_false, ( + f"[{warmup_name}] progress_bar=None must produce the same script as False.\n" + ) # --------------------------------------------------------------------------- diff --git a/tests/recipes/test_emit_x64_timing.py b/tests/recipes/test_emit_x64_timing.py index ea225915..2272bc2e 100644 --- a/tests/recipes/test_emit_x64_timing.py +++ b/tests/recipes/test_emit_x64_timing.py @@ -20,6 +20,7 @@ - Emitted script contains warmup_wall_seconds and sampling_wall_seconds prints. - Warmup timing fence (_warmup_t0, _warmup_t1, _warmup_wall) present in emitted script. """ + from __future__ import annotations import ast @@ -151,25 +152,31 @@ def test_x64_line_precedes_model_computation() -> None: # Find line numbers for anchor tokens. import_jax_line = next( - (i for i, l in enumerate(lines) if l.strip() == "import jax"), None + (i for i, line in enumerate(lines) if line.strip() == "import jax"), None + ) + x64_line = next( + (i for i, line in enumerate(lines) if "jax_enable_x64" in line), None ) - x64_line = next((i for i, l in enumerate(lines) if "jax_enable_x64" in l), None) model_import_line = next( - (i for i, l in enumerate(lines) if "from tuningfork.model import MODELS" in l), + ( + i + for i, line in enumerate(lines) + if "from tuningfork.model import MODELS" in line + ), None, ) assert import_jax_line is not None, "Could not find 'import jax' in emitted script" - assert ( - x64_line is not None - ), "gp_regression emitted script is missing 'jax_enable_x64' config line" - assert ( - model_import_line is not None - ), "Could not find 'from tuningfork.model import MODELS' in emitted script" - - assert ( - x64_line > import_jax_line - ), f"x64 config line (line {x64_line}) must come AFTER 'import jax' (line {import_jax_line})" + assert x64_line is not None, ( + "gp_regression emitted script is missing 'jax_enable_x64' config line" + ) + assert model_import_line is not None, ( + "Could not find 'from tuningfork.model import MODELS' in emitted script" + ) + + assert x64_line > import_jax_line, ( + f"x64 config line (line {x64_line}) must come AFTER 'import jax' (line {import_jax_line})" + ) assert x64_line < model_import_line, ( f"x64 config line (line {x64_line}) must come BEFORE model import " f"(line {model_import_line}) — JAX locks precision on first use.\n" @@ -344,6 +351,6 @@ def test_emitted_dense_window_adaptation_default_is_multichain() -> None: "window_adaptation_dense_imm in emitted script must NOT forward a " "progress_bar= kwarg to blackjax (removed upstream in blackjax #964)." ) - assert ( - "_warmup_is_perchain = True" in script - ), "window_adaptation_dense_imm in emitted script should default to multichain." + assert "_warmup_is_perchain = True" in script, ( + "window_adaptation_dense_imm in emitted script should default to multichain." + ) diff --git a/tests/recipes/test_emitted_scripts_golden.py b/tests/recipes/test_emitted_scripts_golden.py index 60ebaad0..c2683c8b 100644 --- a/tests/recipes/test_emitted_scripts_golden.py +++ b/tests/recipes/test_emitted_scripts_golden.py @@ -128,8 +128,7 @@ "failed__laplace_dhmc__window_adaptation_low_rank_imm.json", "neals_funnel/recipes/" "failed__laplace_dmhmc__window_adaptation_low_rank_imm.json", - "neals_funnel/recipes/" - "failed__laplace_hmc__window_adaptation_low_rank_imm.json", + "neals_funnel/recipes/failed__laplace_hmc__window_adaptation_low_rank_imm.json", "neals_funnel/recipes/" "failed__laplace_mhmc__window_adaptation_low_rank_imm.json", "radon/recipes/failed__nuts__fullrank_vi.json", diff --git a/tests/recipes/test_execute_recipe.py b/tests/recipes/test_execute_recipe.py index 0add0ba9..a37d6be5 100644 --- a/tests/recipes/test_execute_recipe.py +++ b/tests/recipes/test_execute_recipe.py @@ -161,7 +161,7 @@ def test_execute_recipe_recipe_evidence_preserves_negative_and_unknown_fields( monkeypatch.setattr( emit_module, "launch_generated_program", - lambda *args, **kwargs: (seen.append(kwargs) or object()), + lambda *args, **kwargs: seen.append(kwargs) or object(), ) emit_module.execute_recipe(recipe, Path("runs")) @@ -203,7 +203,7 @@ def test_execute_recipe_real_recipe_encodes_nonfinite_gate_evidence(monkeypatch) monkeypatch.setattr( emit_module, "launch_generated_program", - lambda *args, **kwargs: (calls.append(kwargs) or object()), + lambda *args, **kwargs: calls.append(kwargs) or object(), ) emit_module.execute_recipe(recipe, Path("runs")) @@ -225,7 +225,7 @@ def test_canonical_recipe_snapshot_matches_receipt_for_structured_values(monkeyp monkeypatch.setattr( emit_module, "launch_generated_program", - lambda *args, **kwargs: (seen.append(kwargs) or object()), + lambda *args, **kwargs: seen.append(kwargs) or object(), ) emit_module.execute_recipe(recipe, Path("runs")) @@ -247,7 +247,7 @@ def test_execute_recipe_recipe_evidence_is_immutable_and_preserves_caller_identi monkeypatch.setattr( emit_module, "launch_generated_program", - lambda *args, **kwargs: (seen.append(kwargs) or object()), + lambda *args, **kwargs: seen.append(kwargs) or object(), ) emit_module.execute_recipe(recipe, Path("runs"), reference_identity=caller) @@ -282,8 +282,9 @@ def test_execute_recipe_supports_smc_generated_program(monkeypatch, tmp_path): monkeypatch.setattr( emit_module, "launch_generated_program", - lambda source, run_root, **kwargs: calls.append((source, run_root, kwargs)) - or object(), + lambda source, run_root, **kwargs: ( + calls.append((source, run_root, kwargs)) or object() + ), ) recipe = SMCRecipe( model_name="gmm_25", @@ -357,7 +358,7 @@ def test_execute_recipe_diagnostics_sets_child_environment(monkeypatch): monkeypatch.setattr( emit_module, "launch_generated_program", - lambda *args, **kwargs: (calls.append(kwargs) or object()), + lambda *args, **kwargs: calls.append(kwargs) or object(), ) emit_module.execute_recipe(_Recipe(), Path("runs"), diagnostics=True) diff --git a/tests/recipes/test_gate_calibration.py b/tests/recipes/test_gate_calibration.py index 1daa1d70..222261f4 100644 --- a/tests/recipes/test_gate_calibration.py +++ b/tests/recipes/test_gate_calibration.py @@ -181,9 +181,9 @@ def test_benchmark_tau_sci_review_at_seed18_mat() -> None: f"n_review={result.calibrated_n_review}" ) assert result.calibrated_n_fail == 0 - assert ( - result.calibrated_n_review is not None and result.calibrated_n_review >= 1 - ), "hot dim (mat approx 0.085) should be REVIEW under TAU_SCI_BENCHMARK=0.15" + assert result.calibrated_n_review is not None and result.calibrated_n_review >= 1, ( + "hot dim (mat approx 0.085) should be REVIEW under TAU_SCI_BENCHMARK=0.15" + ) def test_benchmark_tau_sci_hard_fail_at_genuine_bias() -> None: diff --git a/tests/recipes/test_generate_groundtruth.py b/tests/recipes/test_generate_groundtruth.py index c7207157..1d699ed3 100644 --- a/tests/recipes/test_generate_groundtruth.py +++ b/tests/recipes/test_generate_groundtruth.py @@ -62,12 +62,12 @@ def test_generate_groundtruth_analytic_returns_none_and_populates_cache( # Cache files must be populated (per-model layout post cleanup-and-simplify) assert (tmp_path / "mvn_10" / "_cache" / "draws.npz").exists(), "draws npz missing" - assert ( - tmp_path / "mvn_10" / "reference" / "summary.json" - ).exists(), "summary json missing" - assert ( - tmp_path / "mvn_10" / "reference" / "metadata.json" - ).exists(), "metadata json missing" + assert (tmp_path / "mvn_10" / "reference" / "summary.json").exists(), ( + "summary json missing" + ) + assert (tmp_path / "mvn_10" / "reference" / "metadata.json").exists(), ( + "metadata json missing" + ) # No adaptation file for analytic models assert not (tmp_path / "mvn_10" / "reference" / "adaptation.json").exists() diff --git a/tests/recipes/test_gt_compare_shape_alignment.py b/tests/recipes/test_gt_compare_shape_alignment.py index a226f9b3..f7202fb0 100644 --- a/tests/recipes/test_gt_compare_shape_alignment.py +++ b/tests/recipes/test_gt_compare_shape_alignment.py @@ -94,9 +94,9 @@ def test_gt_compare_structured_event_shape_no_broadcast_error() -> None: assert result.max_abs_mean_z is not None assert np.isfinite(result.max_abs_mean_z) - assert ( - result.n_dims == D - ), f"every grid cell must contribute one z-score; got n_dims={result.n_dims}" + assert result.n_dims == D, ( + f"every grid cell must contribute one z-score; got n_dims={result.n_dims}" + ) assert result.calibrated_D_total == D @@ -133,9 +133,9 @@ def test_gt_compare_structured_event_shape_preserves_dim_order() -> None: _flat_gt(multichain=True, mean_offset=offset, offset_flat_idx=k), min_bulk_ess=None, ) - assert ( - misaligned.calibrated_pass is False - ), "a one-cell misalignment must be detected as a hard FAIL" + assert misaligned.calibrated_pass is False, ( + "a one-cell misalignment must be detected as a hard FAIL" + ) def test_gt_compare_structured_event_shape_legacy_single_chain_path() -> None: diff --git a/tests/recipes/test_laplace_splits_table.py b/tests/recipes/test_laplace_splits_table.py index 89962993..5d160115 100644 --- a/tests/recipes/test_laplace_splits_table.py +++ b/tests/recipes/test_laplace_splits_table.py @@ -25,9 +25,9 @@ @pytest.mark.fast def test_gp_regression_in_splits_table() -> None: """gp_regression is registered in _LAPLACE_PHI_THETA_SPLITS.""" - assert ( - "gp_regression" in LAPLACE_PHI_THETA_SPLITS - ), "_LAPLACE_PHI_THETA_SPLITS missing 'gp_regression' entry" + assert "gp_regression" in LAPLACE_PHI_THETA_SPLITS, ( + "_LAPLACE_PHI_THETA_SPLITS missing 'gp_regression' entry" + ) @pytest.mark.fast @@ -35,15 +35,15 @@ def test_gp_regression_phi_sites() -> None: """gp_regression phi sites are the 3 log-scale hyperparameters.""" phi_sites, _ = LAPLACE_PHI_THETA_SPLITS["gp_regression"] expected = {"log_lengthscale", "log_kernel_scale", "log_noise_scale"} - assert ( - set(phi_sites) == expected - ), f"phi sites mismatch: expected {expected}, got {set(phi_sites)}" + assert set(phi_sites) == expected, ( + f"phi sites mismatch: expected {expected}, got {set(phi_sites)}" + ) @pytest.mark.fast def test_gp_regression_theta_sites() -> None: """gp_regression theta sites is ('f_raw',) — the NCP base variable.""" _, theta_sites = LAPLACE_PHI_THETA_SPLITS["gp_regression"] - assert theta_sites == ( - "f_raw", - ), f"theta sites mismatch: expected ('f_raw',), got {theta_sites}" + assert theta_sites == ("f_raw",), ( + f"theta sites mismatch: expected ('f_raw',), got {theta_sites}" + ) diff --git a/tests/recipes/test_launcher.py b/tests/recipes/test_launcher.py index 3a61a271..1af449c2 100644 --- a/tests/recipes/test_launcher.py +++ b/tests/recipes/test_launcher.py @@ -490,9 +490,7 @@ def test_child_fixed_receipt_temp_symlink_cannot_overwrite_target(tmp_path): def test_successful_parent_with_live_process_group_descendant_is_rejected(tmp_path): - descendant = ( - "import time; time.sleep(0.5); " "open('late.bin', 'wb').write(b'late')" - ) + descendant = "import time; time.sleep(0.5); open('late.bin', 'wb').write(b'late')" source = _source( _manifest(), "import subprocess, sys\n" @@ -512,7 +510,7 @@ def test_successful_parent_with_live_process_group_descendant_is_rejected(tmp_pa def test_timeout_terminates_descendants_before_receipt_is_written(tmp_path): manifest = _manifest() descendant = ( - "import time; time.sleep(0.5); " "open('late.draws.npz', 'wb').write(b'late')" + "import time; time.sleep(0.5); open('late.draws.npz', 'wb').write(b'late')" ) source = ( f"EXECUTION_MANIFEST_JSON = {manifest.to_json()!r}\n" diff --git a/tests/recipes/test_multichain_gt_schema.py b/tests/recipes/test_multichain_gt_schema.py index 43221a16..4e942435 100644 --- a/tests/recipes/test_multichain_gt_schema.py +++ b/tests/recipes/test_multichain_gt_schema.py @@ -20,6 +20,7 @@ - list_recipes still enumerates groundtruth.json for migrated models - load_idata returns correctly-shaped multichain posterior on new schema (mock-based) """ + from __future__ import annotations from pathlib import Path @@ -127,9 +128,9 @@ def test_radon_gt_in_list_recipes(self) -> None: paths = list_recipes("radon") filenames = [p.name for p in paths] - assert ( - "groundtruth.json" in filenames - ), f"groundtruth.json missing from list_recipes('radon'). Got: {filenames}" + assert "groundtruth.json" in filenames, ( + f"groundtruth.json missing from list_recipes('radon'). Got: {filenames}" + ) def test_gp_regression_gt_in_list_recipes(self) -> None: """list_recipes('gp_regression') includes groundtruth.json (legacy model).""" @@ -137,9 +138,9 @@ def test_gp_regression_gt_in_list_recipes(self) -> None: paths = list_recipes("gp_regression") filenames = [p.name for p in paths] - assert ( - "groundtruth.json" in filenames - ), f"groundtruth.json missing from list_recipes('gp_regression'). Got: {filenames}" + assert "groundtruth.json" in filenames, ( + f"groundtruth.json missing from list_recipes('gp_regression'). Got: {filenames}" + ) class TestLoadIdataMultichainNewSchema: diff --git a/tests/recipes/test_scan_vmap_progress_bar.py b/tests/recipes/test_scan_vmap_progress_bar.py index 96d77bcd..825b07f8 100644 --- a/tests/recipes/test_scan_vmap_progress_bar.py +++ b/tests/recipes/test_scan_vmap_progress_bar.py @@ -40,6 +40,7 @@ they assert the emitted script executes (structure correct, no vmap/io_callback errors), NOT inference quality. This keeps the e2e gate fast and memory-safe. """ + from __future__ import annotations import dataclasses @@ -118,6 +119,6 @@ def test_multichain_progress_bar_no_vmap_of_cond_error(tmp_path: Path) -> None: f"Emitted multi-chain progress_bar script exited with code {result.returncode}.\n" f"stdout:\n{result.stdout}\nstderr:\n{result.stderr}" ) - assert ( - "DONE" in result.stdout - ), f"Expected 'DONE' in stdout.\nstdout:\n{result.stdout}" + assert "DONE" in result.stdout, ( + f"Expected 'DONE' in stdout.\nstdout:\n{result.stdout}" + ) diff --git a/tests/recipes/test_schema.py b/tests/recipes/test_schema.py index dd3979e9..2141f18f 100644 --- a/tests/recipes/test_schema.py +++ b/tests/recipes/test_schema.py @@ -185,9 +185,9 @@ def test_save_load_roundtrip(tmp_path: Path) -> None: # calibration_budget: original keys must survive round-trip; the backward-compat # backfill in Recipe.load adds timing keys (None) that the original dict may lack. for k, v in recipe.calibration_budget.items(): - assert ( - loaded.calibration_budget[k] == v - ), f"calibration_budget[{k!r}]: {loaded.calibration_budget[k]!r} != {v!r}" + assert loaded.calibration_budget[k] == v, ( + f"calibration_budget[{k!r}]: {loaded.calibration_budget[k]!r} != {v!r}" + ) assert loaded.difficulty == recipe.difficulty assert loaded.instructions == recipe.instructions assert loaded.notes == recipe.notes @@ -1734,10 +1734,9 @@ def test_every_catalog_recipe_round_trips_through_load() -> None: except Exception as exc: failures.append((Path(p).name, type(exc).__name__, str(exc)[:120])) - assert ( - not failures - ), f"{len(failures)} catalog recipes fail load_recipe() round-trip:\n" + "\n".join( - f" {name}: {etype}: {emsg}" for name, etype, emsg in failures + assert not failures, ( + f"{len(failures)} catalog recipes fail load_recipe() round-trip:\n" + + "\n".join(f" {name}: {etype}: {emsg}" for name, etype, emsg in failures) ) @@ -1771,10 +1770,9 @@ def test_every_catalog_recipe_has_required_fields_on_disk() -> None: if absent: missing.append((Path(p).name, absent)) - assert ( - not missing - ), f"{len(missing)} recipes missing required keys on disk:\n" + "\n".join( - f" {name}: missing {keys}" for name, keys in missing + assert not missing, ( + f"{len(missing)} recipes missing required keys on disk:\n" + + "\n".join(f" {name}: missing {keys}" for name, keys in missing) ) diff --git a/tests/recipes/test_staged_adaptation_auto_codegen.py b/tests/recipes/test_staged_adaptation_auto_codegen.py index b8d5e663..58cdcd5b 100644 --- a/tests/recipes/test_staged_adaptation_auto_codegen.py +++ b/tests/recipes/test_staged_adaptation_auto_codegen.py @@ -194,9 +194,9 @@ def test_max_grad_budget_is_a_material_plan_field() -> None: bumped = resolve_execution_plan(_recipe(max_grad_budget=40_000)) assert base.plan_hash != bumped.plan_hash assert base.executable_config_hash != bumped.executable_config_hash - assert ( - base.config.warmup_stages[0].params["max_grad_budget"] == 20_000 - ), "max_grad_budget must survive into the normalized plan, not just the recipe" + assert base.config.warmup_stages[0].params["max_grad_budget"] == 20_000, ( + "max_grad_budget must survive into the normalized plan, not just the recipe" + ) # --------------------------------------------------------------------------- diff --git a/tests/recipes/test_statistician_gate.py b/tests/recipes/test_statistician_gate.py index ebeef8a2..4f3091b8 100644 --- a/tests/recipes/test_statistician_gate.py +++ b/tests/recipes/test_statistician_gate.py @@ -743,15 +743,15 @@ def test_dimension_aware_gate_still_fails_genuine_bias(): assert verdict.min_bulk_ess is not None if verdict.min_bulk_ess > Z_VERDICT_ESS_CEILING: # Large realm: z-driven FAIL demotes to REVIEW - assert ( - verdict.verdict == "REVIEW" - ), f"d={d}, ESS={verdict.min_bulk_ess:.0f}: expected REVIEW with z_advisory" + assert verdict.verdict == "REVIEW", ( + f"d={d}, ESS={verdict.min_bulk_ess:.0f}: expected REVIEW with z_advisory" + ) assert verdict.margins["max_abs_mean_z"].get("z_advisory") is True else: # Small realm: verdict is FAIL as before - assert ( - verdict.verdict == "FAIL" - ), f"d={d} did not FAIL at z={verdict.max_abs_mean_z}" + assert verdict.verdict == "FAIL", ( + f"d={d} did not FAIL at z={verdict.max_abs_mean_z}" + ) # --------------------------------------------------------------------------- @@ -1077,9 +1077,9 @@ def test_z_advisory_bias_sigma_numerically_correct(): bias_sigma_max_z4 = margins_z["bias_sigma_max_at_z4"] # Should be close to 0.005 (from dim 0), NOT 0.048 (from dim 1). # This pins the semantic: max among z>=4 failing dims. - assert ( - 0.003 < bias_sigma_max_z4 < 0.010 - ), f"bias_sigma_max_at_z4={bias_sigma_max_z4}, expected ≈0.005 (from failing dim z≥4)" + assert 0.003 < bias_sigma_max_z4 < 0.010, ( + f"bias_sigma_max_at_z4={bias_sigma_max_z4}, expected ≈0.005 (from failing dim z≥4)" + ) def test_z_advisory_cost_block_optional(): diff --git a/tests/recipes/test_step_policy_medium_emit.py b/tests/recipes/test_step_policy_medium_emit.py index 509ff45d..bac66554 100644 --- a/tests/recipes/test_step_policy_medium_emit.py +++ b/tests/recipes/test_step_policy_medium_emit.py @@ -109,9 +109,9 @@ def test_medium_with_policy_tag_mvn10(tmp_path): / "recipes" / f"medium__dynamic_hmc__window_adaptation_diag_imm__{policy_tag}.json" ) - assert ( - expected_path.exists() - ), f"Expected recipe at {expected_path}; got result.recipe_path={result.recipe_path}" + assert expected_path.exists(), ( + f"Expected recipe at {expected_path}; got result.recipe_path={result.recipe_path}" + ) # Load and verify the recipe from tuningfork.recipes._base import Recipe @@ -135,12 +135,12 @@ def test_medium_with_policy_tag_mvn10(tmp_path): "MEDIUM recipe headline_basis is null — grad_count_convention + is_lower_bound " "must be populated from BaseMethod.grad_count_convention." ) - assert recipe.headline_basis.get( - "grad_count_convention" - ), f"headline_basis.grad_count_convention is empty: {recipe.headline_basis}" - assert isinstance( - recipe.headline_basis.get("is_lower_bound"), bool - ), f"headline_basis.is_lower_bound must be bool, got: {recipe.headline_basis}" + assert recipe.headline_basis.get("grad_count_convention"), ( + f"headline_basis.grad_count_convention is empty: {recipe.headline_basis}" + ) + assert isinstance(recipe.headline_basis.get("is_lower_bound"), bool), ( + f"headline_basis.is_lower_bound must be bool, got: {recipe.headline_basis}" + ) def test_medium_with_policy_tag_none_preserves_low(tmp_path): diff --git a/tests/recipes/test_w1_realm.py b/tests/recipes/test_w1_realm.py index bb94f2ed..d697c087 100644 --- a/tests/recipes/test_w1_realm.py +++ b/tests/recipes/test_w1_realm.py @@ -326,9 +326,9 @@ def test_w1_unequal_n_agrees_with_equal_n(): # The two measures should be in the same ballpark; not exact because # different samples are drawn, but both estimate E[W1] ≈ 1/(2√n). assert w1_unequal > 0.0, "W1 should be positive for unequal-n samples" - assert ( - w1_unequal < 0.15 - ), f"W1 suspiciously large for same distribution: {w1_unequal}" + assert w1_unequal < 0.15, ( + f"W1 suspiciously large for same distribution: {w1_unequal}" + ) @pytest.mark.fast @@ -569,14 +569,14 @@ def test_eight_schools_injected_max_prong_fail(): ) max_w1 = float(np.max(computed_w1)) - assert ( - abs(max_w1 - _ENS_INJECTED_MAX_W1_SIGMA) < 1e-8 - ), f"INJECTED max W1/σ={max_w1:.8f} ≠ pinned {_ENS_INJECTED_MAX_W1_SIGMA:.8f}" + assert abs(max_w1 - _ENS_INJECTED_MAX_W1_SIGMA) < 1e-8, ( + f"INJECTED max W1/σ={max_w1:.8f} ≠ pinned {_ENS_INJECTED_MAX_W1_SIGMA:.8f}" + ) # MAX-prong FAIL - assert ( - max_w1 > _ENS_FLOOR_OF_MAX - ), f"INJECTED case should FAIL: max W1/σ={max_w1:.6f} ≤ floor={_ENS_FLOOR_OF_MAX:.8f}" + assert max_w1 > _ENS_FLOOR_OF_MAX, ( + f"INJECTED case should FAIL: max W1/σ={max_w1:.6f} ≤ floor={_ENS_FLOOR_OF_MAX:.8f}" + ) @pytest.mark.slow @@ -640,9 +640,9 @@ def test_eight_schools_loo_conservatism_guard(): ) # k_crit=1 for n_chains=10 (P(Bin(10,0.05) >= 2) = 0.086 ≤ 0.10) - assert ( - result["k_crit"] == 1 - ), f"Expected k_crit=1 for n_chains=10, got {result['k_crit']}" + assert result["k_crit"] == 1, ( + f"Expected k_crit=1 for n_chains=10, got {result['k_crit']}" + ) # Count rule: exactly 1/10 violation (leave-one-out gives the honest null). assert result["violation_count"] <= result["k_crit"], ( f"violation_count={result['violation_count']} > k_crit={result['k_crit']}: " @@ -652,9 +652,9 @@ def test_eight_schools_loo_conservatism_guard(): ) # Severity rule: max LOO null ≤ floor * 1.05 (1.9% overshoot < 5% limit). max_loo = max(result["loo_max_w1_sigma"]) - assert ( - max_loo <= _ENS_FLOOR_OF_MAX * 1.05 - ), f"Severity rule FAILED: max_loo={max_loo:.5f} > floor*1.05={_ENS_FLOOR_OF_MAX * 1.05:.5f}" + assert max_loo <= _ENS_FLOOR_OF_MAX * 1.05, ( + f"Severity rule FAILED: max_loo={max_loo:.5f} > floor*1.05={_ENS_FLOOR_OF_MAX * 1.05:.5f}" + ) # is_conservative must be True (both count and severity rules pass). assert result["is_conservative"] is True, ( f"is_conservative=False despite valid count ({result['violation_count']}) and " @@ -728,9 +728,9 @@ def test_eight_schools_floor_of_max_e2e_pinned(): ) # tau_frac sanity (D=10, may be coarse-grained at B=5000) - assert ( - 0.5 <= result.tau_frac <= 0.9 - ), f"tau_frac={result.tau_frac:.4f} out of expected [0.5, 0.9] range for D=10" + assert 0.5 <= result.tau_frac <= 0.9, ( + f"tau_frac={result.tau_frac:.4f} out of expected [0.5, 0.9] range for D=10" + ) # NULL should PASS (gen = chain0[:1000] from same distribution) assert result.verdict == "PASS", ( @@ -739,9 +739,9 @@ def test_eight_schools_floor_of_max_e2e_pinned(): ) # No heavy-tail dims (all khat < 0 for eight_schools_ncp NCP) - assert ( - result.n_heavy_tail_dims == 0 - ), f"Expected 0 heavy-tail dims (GPD k-hat fix), got {result.n_heavy_tail_dims}" + assert result.n_heavy_tail_dims == 0, ( + f"Expected 0 heavy-tail dims (GPD k-hat fix), got {result.n_heavy_tail_dims}" + ) # =========================================================================== @@ -814,9 +814,9 @@ def test_radon_frac_prong_tau_frac_pinned(): # Floor must be in the real-ESS regime (>0.20); if ESS regresses to raw 1000 # the floor would drop to ~0.128. - assert ( - result.floor_of_max > 0.20 - ), f"Radon floor_of_max={result.floor_of_max:.6f} < 0.20 — likely ESS regression" + assert result.floor_of_max > 0.20, ( + f"Radon floor_of_max={result.floor_of_max:.6f} < 0.20 — likely ESS regression" + ) # Pin tau_frac to 5% relative tolerance (B=2000, well-resolved) assert abs(result.tau_frac - _RADON_TAU_FRAC) / _RADON_TAU_FRAC < 0.05, ( @@ -844,9 +844,9 @@ def test_radon_frac_prong_tau_frac_pinned(): ) # No heavy-tail dims (GPD k-hat fix: radon posteriors are not heavy-tailed) - assert ( - result.n_heavy_tail_dims == 0 - ), f"Expected 0 heavy-tail dims after GPD k-hat fix, got {result.n_heavy_tail_dims}" + assert result.n_heavy_tail_dims == 0, ( + f"Expected 0 heavy-tail dims after GPD k-hat fix, got {result.n_heavy_tail_dims}" + ) @pytest.mark.slow @@ -878,9 +878,9 @@ def test_radon_null_max_w1_sigma_deterministic(): w1 = _w1_1d_local(gt_d, gen_d) / float(sig[dim_i]) w1_max = max(w1_max, w1) - assert ( - abs(w1_max - _RADON_NULL_MAX_W1_SIGMA) < 1e-4 - ), f"Radon NULL max W1/σ={w1_max:.6f} ≠ pinned {_RADON_NULL_MAX_W1_SIGMA:.6f}" + assert abs(w1_max - _RADON_NULL_MAX_W1_SIGMA) < 1e-4, ( + f"Radon NULL max W1/σ={w1_max:.6f} ≠ pinned {_RADON_NULL_MAX_W1_SIGMA:.6f}" + ) @pytest.mark.slow @@ -939,9 +939,9 @@ def test_radon_loo_conservatism_guard(): ) # k_crit=1 for n_chains=10 - assert ( - result["k_crit"] == 1 - ), f"Expected k_crit=1 for n_chains=10, got {result['k_crit']}" + assert result["k_crit"] == 1, ( + f"Expected k_crit=1 for n_chains=10, got {result['k_crit']}" + ) # Count rule: 1/10 violation is expected (and within k_crit=1). assert result["violation_count"] <= result["k_crit"], ( f"violation_count={result['violation_count']} > k_crit={result['k_crit']}: " @@ -1061,9 +1061,9 @@ def test_w1_realm_result_skip_counts_as_pass_in_overall(): # Overall verdict should be PASS (SKIP is not a failure) if result.frac_prong_verdict == "SKIP": if result.max_prong_verdict == "PASS": - assert ( - result.verdict == "PASS" - ), f"SKIP frac + PASS max should give overall PASS, got {result.verdict}" + assert result.verdict == "PASS", ( + f"SKIP frac + PASS max should give overall PASS, got {result.verdict}" + ) # =========================================================================== @@ -1165,18 +1165,18 @@ def test_w1_realm_skip_fold_no_crash(): from tuningfork.calibration._gate.bands import _worst # _worst("PASS", "SKIP") must not raise KeyError - assert ( - _worst("PASS", "SKIP") == "PASS" - ), "_worst('PASS', 'SKIP') should return 'PASS' (SKIP ≡ PASS rank)" - assert ( - _worst("SKIP", "PASS") == "PASS" - ), "_worst('SKIP', 'PASS') should return 'PASS'" - assert ( - _worst("SKIP", "SKIP") == "PASS" - ), "_worst('SKIP', 'SKIP') should return 'PASS'" - assert ( - _worst("SKIP", "FAIL") == "FAIL" - ), "_worst('SKIP', 'FAIL') should return 'FAIL'" + assert _worst("PASS", "SKIP") == "PASS", ( + "_worst('PASS', 'SKIP') should return 'PASS' (SKIP ≡ PASS rank)" + ) + assert _worst("SKIP", "PASS") == "PASS", ( + "_worst('SKIP', 'PASS') should return 'PASS'" + ) + assert _worst("SKIP", "SKIP") == "PASS", ( + "_worst('SKIP', 'SKIP') should return 'PASS'" + ) + assert _worst("SKIP", "FAIL") == "FAIL", ( + "_worst('SKIP', 'FAIL') should return 'FAIL'" + ) # compute_w1_realm with non-overlapping sites must return SKIP, not crash rng = np.random.default_rng(9) diff --git a/tests/recipes/test_warmup_inner_kernel_emit.py b/tests/recipes/test_warmup_inner_kernel_emit.py index 7b3780cd..17386811 100644 --- a/tests/recipes/test_warmup_inner_kernel_emit.py +++ b/tests/recipes/test_warmup_inner_kernel_emit.py @@ -99,18 +99,18 @@ def test_inner_nuts_hmc_emit_mvn10(tmp_path): # (2026-05-27) — a textbook noisy-MC-assertion flake. # # Hard structural gates (these should NEVER fail on a working pipeline): - assert ( - result.gate_n_div == 0 - ), f"n_div must be 0 for well-behaved MVN-10 (structural); got {result.gate_n_div}" + assert result.gate_n_div == 0, ( + f"n_div must be 0 for well-behaved MVN-10 (structural); got {result.gate_n_div}" + ) import math assert result.gate_rhat_max is not None and math.isfinite(result.gate_rhat_max), ( f"rhat_max must be finite (pipeline produced draws); " f"got {result.gate_rhat_max!r}" ) - assert ( - result.gate_min_ess is not None and result.gate_min_ess > 0 - ), f"min_ess must be positive (chain mixed at all); got {result.gate_min_ess!r}" + assert result.gate_min_ess is not None and result.gate_min_ess > 0, ( + f"min_ess must be positive (chain mixed at all); got {result.gate_min_ess!r}" + ) assert result.verdict in {"PASS", "REVIEW", "FAIL"} assert result.recipe_path is not None @@ -123,8 +123,7 @@ def test_inner_nuts_hmc_emit_mvn10(tmp_path): / "low__hmc__window_adaptation_diag_imm__inner_nuts.json" ) assert expected_path.exists(), ( - f"Expected recipe at {expected_path}; " - f"result.recipe_path={result.recipe_path}" + f"Expected recipe at {expected_path}; result.recipe_path={result.recipe_path}" ) # --- Load and verify recipe fields --- @@ -135,14 +134,14 @@ def test_inner_nuts_hmc_emit_mvn10(tmp_path): assert recipe.warmup_name == "window_adaptation_diag_imm" # Schema extension: warmup_inner_kernel persisted correctly. - assert ( - recipe.warmup_inner_kernel == "nuts" - ), f"Expected warmup_inner_kernel='nuts', got {recipe.warmup_inner_kernel!r}" + assert recipe.warmup_inner_kernel == "nuts", ( + f"Expected warmup_inner_kernel='nuts', got {recipe.warmup_inner_kernel!r}" + ) # Schema extension: warmups list populated (not just flat fields). - assert ( - recipe.warmups - ), "recipe.warmups must be non-empty after schema-extension save/load" + assert recipe.warmups, ( + "recipe.warmups must be non-empty after schema-extension save/load" + ) assert recipe.warmups[0]["name"] == "window_adaptation_diag_imm", ( f"Expected warmups[0].name='window_adaptation_diag_imm', " f"got {recipe.warmups[0]['name']!r}" @@ -156,9 +155,9 @@ def test_inner_nuts_hmc_emit_mvn10(tmp_path): f"Actual base_method_params keys: {list(recipe.base_method_params.keys())}" ) nis_value = recipe.base_method_params["num_integration_steps"] - assert ( - isinstance(nis_value, int) and nis_value >= 1 - ), f"num_integration_steps must be a positive int, got {nis_value!r}" + assert isinstance(nis_value, int) and nis_value >= 1, ( + f"num_integration_steps must be a positive int, got {nis_value!r}" + ) # Schema extension: the recipe must NOT write legacy flat warmup fields to JSON. import json diff --git a/tests/reference/test_analytic.py b/tests/reference/test_analytic.py index 9a244f26..74fb4e05 100644 --- a/tests/reference/test_analytic.py +++ b/tests/reference/test_analytic.py @@ -144,9 +144,9 @@ def test_v_mean_near_zero(self) -> None: v = np.asarray(draws["v"]) # (N,) # std(v) = 3, so MC SE = 3 / sqrt(N) tol = 4.0 * 3.0 / np.sqrt(N) - assert ( - abs(v.mean()) < tol - ), f"Funnel v mean={v.mean():.6f} not within 4-sigma tol={tol:.6f}" + assert abs(v.mean()) < tol, ( + f"Funnel v mean={v.mean():.6f} not within 4-sigma tol={tol:.6f}" + ) def test_v_std_near_three(self) -> None: key = jax.random.key(101) @@ -154,9 +154,9 @@ def test_v_std_near_three(self) -> None: v = np.asarray(draws["v"]) # (N,) # MC SE of std estimator ≈ 3 / sqrt(2N) tol = 4.0 * 3.0 / np.sqrt(2 * N) - assert ( - abs(v.std() - 3.0) < tol - ), f"Funnel v std={v.std():.6f} not within 4-sigma of 3 (tol={tol:.6f})" + assert abs(v.std() - 3.0) < tol, ( + f"Funnel v std={v.std():.6f} not within 4-sigma of 3 (tol={tol:.6f})" + ) def test_theta_mean_near_zero(self) -> None: # Generous tolerance: theta marginal std is large (heavy tails). @@ -180,9 +180,9 @@ def test_summaries_v_mean(self) -> None: _, summaries = certify_reference_analytic(FUNNEL_ENTRY, N, key) v_mean = float(jnp.asarray(summaries.mean["v"])) tol = 4.0 * 3.0 / np.sqrt(N) - assert ( - abs(v_mean) < tol - ), f"Summaries v mean={v_mean:.6f} not within tolerance {tol:.6f}" + assert abs(v_mean) < tol, ( + f"Summaries v mean={v_mean:.6f} not within tolerance {tol:.6f}" + ) def test_draws_shape_funnel(self) -> None: key = jax.random.key(104) diff --git a/tests/reference/test_nuts.py b/tests/reference/test_nuts.py index bdd397ef..14b102b6 100644 --- a/tests/reference/test_nuts.py +++ b/tests/reference/test_nuts.py @@ -154,9 +154,9 @@ def test_draws_keys_contain_theta_raw(self, smoke_cert_result) -> None: def test_draws_sample_axis(self, smoke_cert_result) -> None: draws, _, _, _, _ = smoke_cert_result for site, arr in draws.items(): - assert ( - arr.shape[0] == SMOKE_N_SAMPLES - ), f"Site {site!r}: expected shape[0]={SMOKE_N_SAMPLES}, got {arr.shape[0]}" + assert arr.shape[0] == SMOKE_N_SAMPLES, ( + f"Site {site!r}: expected shape[0]={SMOKE_N_SAMPLES}, got {arr.shape[0]}" + ) def test_summaries_instance(self, smoke_cert_result) -> None: _, summaries, _, _, _ = smoke_cert_result @@ -206,9 +206,9 @@ def test_chain_stats_has_required_fields(self, smoke_cert_result) -> None: def test_chain_stats_field_shapes(self, smoke_cert_result) -> None: _, _, _, _, chain_stats = smoke_cert_result for field_name, arr in chain_stats.items(): - assert ( - arr.shape[0] == SMOKE_N_SAMPLES - ), f"Field {field_name!r}: expected shape[0]={SMOKE_N_SAMPLES}, got {arr.shape[0]}" + assert arr.shape[0] == SMOKE_N_SAMPLES, ( + f"Field {field_name!r}: expected shape[0]={SMOKE_N_SAMPLES}, got {arr.shape[0]}" + ) # Default-clean baseline used by gate-logic tests. Mutated per-test by diff --git a/tests/reference/test_posteriordb_xcheck.py b/tests/reference/test_posteriordb_xcheck.py index 3d1c2b1e..a7c352f9 100644 --- a/tests/reference/test_posteriordb_xcheck.py +++ b/tests/reference/test_posteriordb_xcheck.py @@ -275,9 +275,9 @@ def test_vector_param_passes(self) -> None: n_samples_ours=n, ) - assert ( - result.passed is True - ), f"failed_dims={result.failed_dims}, max_z={result.max_abs_mean_z:.3f}" + assert result.passed is True, ( + f"failed_dims={result.failed_dims}, max_z={result.max_abs_mean_z:.3f}" + ) assert result.n_dims_compared == 3 @@ -327,9 +327,9 @@ def test_mean_shift_fails(self) -> None: n_samples_ours=n_ours, ) - assert ( - result.passed is False - ), f"Expected passed=False; max_z={result.max_abs_mean_z:.3f}" + assert result.passed is False, ( + f"Expected passed=False; max_z={result.max_abs_mean_z:.3f}" + ) assert len(result.failed_dims) > 0 assert result.max_abs_mean_z >= 2.0 @@ -369,9 +369,9 @@ def test_std_ratio_fails(self) -> None: ) assert result.passed is False - assert ( - result.max_std_ratio_dev >= 0.05 - ), f"Expected max_std_ratio_dev≥0.05; got {result.max_std_ratio_dev:.4f}" + assert result.max_std_ratio_dev >= 0.05, ( + f"Expected max_std_ratio_dev≥0.05; got {result.max_std_ratio_dev:.4f}" + ) def test_failed_dims_contains_param_name(self) -> None: """When 'mu' fails, 'mu' (or 'mu[0]') appears in failed_dims.""" @@ -406,9 +406,9 @@ def test_failed_dims_contains_param_name(self) -> None: assert result.passed is False # "mu" should appear in some form in failed_dims (could be "mu" or "mu[0]") - assert any( - "mu" in d for d in result.failed_dims - ), f"Expected 'mu' in failed_dims; got {result.failed_dims}" + assert any("mu" in d for d in result.failed_dims), ( + f"Expected 'mu' in failed_dims; got {result.failed_dims}" + ) # --------------------------------------------------------------------------- diff --git a/tests/reference/test_revalidation.py b/tests/reference/test_revalidation.py index f953a4f8..31b1d8f4 100644 --- a/tests/reference/test_revalidation.py +++ b/tests/reference/test_revalidation.py @@ -160,9 +160,9 @@ def test_thresholds_imported_not_hardcoded(self): # One below the FAIL boundary (= review hi - 1) must be REVIEW or PASS result_not_fail = compute_stage1_verdict(draws, n_divergences=ndiv_fail - 1) - assert ( - result_fail["stage1_verdict"] == "FAIL" - ), f"n_div={ndiv_fail} (= DEFAULT_THRESHOLDS FAIL boundary) must → FAIL" + assert result_fail["stage1_verdict"] == "FAIL", ( + f"n_div={ndiv_fail} (= DEFAULT_THRESHOLDS FAIL boundary) must → FAIL" + ) assert result_not_fail["stage1_verdict"] in ( "PASS", "REVIEW", diff --git a/tests/smc/test_smc_parametrized.py b/tests/smc/test_smc_parametrized.py index 0b8c1f4a..5d239220 100644 --- a/tests/smc/test_smc_parametrized.py +++ b/tests/smc/test_smc_parametrized.py @@ -51,9 +51,9 @@ @pytest.mark.parametrize("smc_name", _SMC_METHODS) def test_smc_method_registered(smc_name: str) -> None: """SMC method is registered in SMC_METHODS.""" - assert ( - smc_name in SMC_METHODS - ), f"SMC_METHODS must contain '{smc_name}'; registered: {sorted(SMC_METHODS)}" + assert smc_name in SMC_METHODS, ( + f"SMC_METHODS must contain '{smc_name}'; registered: {sorted(SMC_METHODS)}" + ) @pytest.mark.parametrize("smc_name", _SMC_METHODS) @@ -74,33 +74,33 @@ def test_smc_method_family_is_smc(smc_name: str) -> None: def test_smc_method_has_default_inner_method(smc_name: str) -> None: """ENTRY.default_inner_method is defined.""" entry = SMC_METHODS[smc_name] - assert ( - entry.default_inner_method is not None - ), f"{smc_name}.default_inner_method must be defined" - assert isinstance( - entry.default_inner_method, str - ), f"{smc_name}.default_inner_method must be str" + assert entry.default_inner_method is not None, ( + f"{smc_name}.default_inner_method must be defined" + ) + assert isinstance(entry.default_inner_method, str), ( + f"{smc_name}.default_inner_method must be str" + ) @pytest.mark.parametrize("smc_name", _SMC_METHODS) def test_smc_method_num_particles_default_positive(smc_name: str) -> None: """ENTRY.num_particles_default is a positive integer.""" entry = SMC_METHODS[smc_name] - assert isinstance( - entry.num_particles_default, int - ), f"{smc_name}.num_particles_default must be int" - assert ( - entry.num_particles_default > 0 - ), f"{smc_name}.num_particles_default must be positive" + assert isinstance(entry.num_particles_default, int), ( + f"{smc_name}.num_particles_default must be int" + ) + assert entry.num_particles_default > 0, ( + f"{smc_name}.num_particles_default must be positive" + ) @pytest.mark.parametrize("smc_name", _SMC_METHODS) def test_smc_method_hp_space_non_empty(smc_name: str) -> None: """ENTRY.default_hp_space is non-empty.""" entry = SMC_METHODS[smc_name] - assert ( - len(entry.default_hp_space) > 0 - ), f"{smc_name}.default_hp_space must be non-empty" + assert len(entry.default_hp_space) > 0, ( + f"{smc_name}.default_hp_space must be non-empty" + ) @pytest.mark.parametrize("smc_name", _SMC_METHODS) @@ -114,18 +114,18 @@ def test_smc_method_notes_non_empty(smc_name: str) -> None: def test_smc_method_compatible_inner_non_empty(smc_name: str) -> None: """ENTRY.compatible_inner_methods is non-empty.""" entry = SMC_METHODS[smc_name] - assert ( - len(entry.compatible_inner_methods) > 0 - ), f"{smc_name}.compatible_inner_methods must be non-empty" + assert len(entry.compatible_inner_methods) > 0, ( + f"{smc_name}.compatible_inner_methods must be non-empty" + ) @pytest.mark.parametrize("smc_name", _SMC_METHODS) def test_smc_method_default_inner_in_compatible(smc_name: str) -> None: """ENTRY.default_inner_method is in compatible_inner_methods.""" entry = SMC_METHODS[smc_name] - assert ( - entry.default_inner_method in entry.compatible_inner_methods - ), f"{smc_name}: default_inner_method={entry.default_inner_method} not in compatible_inner_methods" + assert entry.default_inner_method in entry.compatible_inner_methods, ( + f"{smc_name}: default_inner_method={entry.default_inner_method} not in compatible_inner_methods" + ) # =========================================================================== @@ -137,24 +137,24 @@ def test_smc_method_default_inner_in_compatible(smc_name: str) -> None: def test_smc_method_mclmc_excluded(smc_name: str) -> None: """mclmc is excluded from compatible_inner_methods (microcanonical invariance violated by tempering).""" entry = SMC_METHODS[smc_name] - assert ( - "mclmc" not in entry.compatible_inner_methods - ), f"{smc_name}: mclmc must be excluded from compatible_inner_methods" + assert "mclmc" not in entry.compatible_inner_methods, ( + f"{smc_name}: mclmc must be excluded from compatible_inner_methods" + ) @pytest.mark.parametrize("smc_name", _SMC_METHODS) def test_smc_method_adjusted_mclmc_excluded(smc_name: str) -> None: """adjusted_mclmc is excluded from compatible_inner_methods.""" entry = SMC_METHODS[smc_name] - assert ( - "adjusted_mclmc" not in entry.compatible_inner_methods - ), f"{smc_name}: adjusted_mclmc must be excluded from compatible_inner_methods" + assert "adjusted_mclmc" not in entry.compatible_inner_methods, ( + f"{smc_name}: adjusted_mclmc must be excluded from compatible_inner_methods" + ) @pytest.mark.parametrize("smc_name", _SMC_METHODS) def test_smc_method_adjusted_mclmc_dynamic_excluded(smc_name: str) -> None: """adjusted_mclmc_dynamic is excluded from compatible_inner_methods.""" entry = SMC_METHODS[smc_name] - assert ( - "adjusted_mclmc_dynamic" not in entry.compatible_inner_methods - ), f"{smc_name}: adjusted_mclmc_dynamic must be excluded from compatible_inner_methods" + assert "adjusted_mclmc_dynamic" not in entry.compatible_inner_methods, ( + f"{smc_name}: adjusted_mclmc_dynamic must be excluded from compatible_inner_methods" + ) diff --git a/tests/test_api_pins_smc.py b/tests/test_api_pins_smc.py index 01fcf8fd..2e435bb2 100644 --- a/tests/test_api_pins_smc.py +++ b/tests/test_api_pins_smc.py @@ -30,9 +30,9 @@ def _assert_parameters( f"current parameters: {list(parameters)}" ) if var_keyword is not None: - assert ( - parameters[var_keyword].kind is inspect.Parameter.VAR_KEYWORD - ), f"{callable_obj!r} must retain **{var_keyword} for generated SMC options" + assert parameters[var_keyword].kind is inspect.Parameter.VAR_KEYWORD, ( + f"{callable_obj!r} must retain **{var_keyword} for generated SMC options" + ) def test_adaptive_tempered_constructor_parameters() -> None: diff --git a/tests/test_cli_leaderboard.py b/tests/test_cli_leaderboard.py index 0b8cccca..4fc9c249 100644 --- a/tests/test_cli_leaderboard.py +++ b/tests/test_cli_leaderboard.py @@ -114,11 +114,10 @@ def test_leaderboard_smc_note_printed_in_markdown(capsys) -> None: _cmd_leaderboard(args) out = capsys.readouterr().out assert "SMC" in out, ( - f"Expected SMC note in leaderboard output for {model_name!r}.\n" f"Got:\n{out}" + f"Expected SMC note in leaderboard output for {model_name!r}.\nGot:\n{out}" ) assert "not ranked" in out, ( - f"Expected 'not ranked' note for SMC in output for {model_name!r}.\n" - f"Got:\n{out}" + f"Expected 'not ranked' note for SMC in output for {model_name!r}.\nGot:\n{out}" ) @@ -223,6 +222,6 @@ def test_leaderboard_json_format_mcmc_only_on_smc_model() -> None: assert "effort" in entry, f"Entry missing 'effort': {entry}" assert "base_method_name" in entry, f"Entry missing 'base_method_name': {entry}" # SMC entries would have 'smc_method_name' instead — assert absent - assert ( - "smc_method_name" not in entry - ), f"SMC entry leaked into JSON output: {entry}" + assert "smc_method_name" not in entry, ( + f"SMC entry leaked into JSON output: {entry}" + ) diff --git a/tests/test_expectands.py b/tests/test_expectands.py index c9574c6d..13eb2951 100644 --- a/tests/test_expectands.py +++ b/tests/test_expectands.py @@ -1169,9 +1169,9 @@ def test_every_sampler_declaring_an_inexact_count_is_disclosed_as_such(): derivation = sampling_grad_evals_from_chain_stats(stats, name) if derivation.count is None: continue # refused for an unrelated reason; nothing to disclose - assert any( - "INEXACT" in item for item in derivation.excluded - ), f"{name} counts inexactly without disclosing it" + assert any("INEXACT" in item for item in derivation.excluded), ( + f"{name} counts inexactly without disclosing it" + ) @pytest.mark.fast diff --git a/tests/test_registry.py b/tests/test_registry.py index be47460d..630ffa25 100644 --- a/tests/test_registry.py +++ b/tests/test_registry.py @@ -149,8 +149,8 @@ def test_inference_namespace_imports(): assert isinstance(BASE_METHODS, dict) # Core 6 must still be present; more entries may be added later. core_six = {"hmc", "nuts", "mala", "barker", "rwm", "mclmc"} - assert core_six <= set( - BASE_METHODS.keys() - ), f"missing core base methods: {core_six - set(BASE_METHODS.keys())}" + assert core_six <= set(BASE_METHODS.keys()), ( + f"missing core base methods: {core_six - set(BASE_METHODS.keys())}" + ) assert isinstance(WARMUPS, dict) # may be empty assert Warmup is not None diff --git a/tests/transport/test_chart_mutants.py b/tests/transport/test_chart_mutants.py index 673b57a4..e0a9afce 100644 --- a/tests/transport/test_chart_mutants.py +++ b/tests/transport/test_chart_mutants.py @@ -159,7 +159,7 @@ def test_omitted_normalisation_is_rejected(): jnp.asarray(rng.normal(size=DIM)), jnp.asarray(np.exp(0.2 * rng.normal(size=DIM))), ) - assert ( - float(jnp.abs(jnp.linalg.norm(broken.h) - 1.0)) > 1e-3 - ), "must be un-normalised" + assert float(jnp.abs(jnp.linalg.norm(broken.h) - 1.0)) > 1e-3, ( + "must be un-normalised" + ) assert _run(broken) != (True, True) diff --git a/tests/transport/test_chart_refuters.py b/tests/transport/test_chart_refuters.py index e3f4178c..5fa03ef4 100644 --- a/tests/transport/test_chart_refuters.py +++ b/tests/transport/test_chart_refuters.py @@ -219,9 +219,7 @@ def test_exact_funnel_chart_factorises_the_clock(): """ h = jnp.zeros(FUNNEL_DIM).at[-1].set(1.0) chart = make_chart(h, -0.5 * h, h, 0.5, jnp.zeros(FUNNEL_DIM), jnp.ones(FUNNEL_DIM)) - reference = lambda y: _funnel_logdensity(chart.forward(y)) + chart.log_det( - y - ) # noqa: E731 + reference = lambda y: _funnel_logdensity(chart.forward(y)) + chart.log_det(y) # noqa: E731 clock_row = jax.jacfwd(lambda y: jax.grad(reference)(y)[-1]) rng = np.random.default_rng(13) @@ -230,9 +228,9 @@ def test_exact_funnel_chart_factorises_the_clock(): np.concatenate([rng.normal(size=FUNNEL_DIM - 1), [rng.normal() * 3]]) ) row = clock_row(y) - assert ( - float(jnp.max(jnp.abs(row[:-1]))) < 1e-10 - ), "clock score depends on section" + assert float(jnp.max(jnp.abs(row[:-1]))) < 1e-10, ( + "clock score depends on section" + ) assert abs(float(row[-1]) + 1.0 / 9.0) < 1e-10, "clock curvature not constant" @@ -285,9 +283,9 @@ def test_finite_but_overflowing_basis_is_refused(): assert bool(jnp.all(jnp.isfinite(basis))), "every entry must be finite" gram = basis.T @ basis - assert not bool( - jnp.all(jnp.isfinite(gram)) - ), "premise: derived Gram must be non-finite" + assert not bool(jnp.all(jnp.isfinite(gram))), ( + "premise: derived Gram must be non-finite" + ) h = jnp.zeros(d).at[-1].set(1.0) with pytest.raises(ValueError, match="orthonormal"): diff --git a/tuningfork/_posteriordb_xcheck.py b/tuningfork/_posteriordb_xcheck.py index 4c235605..2927293e 100644 --- a/tuningfork/_posteriordb_xcheck.py +++ b/tuningfork/_posteriordb_xcheck.py @@ -266,7 +266,9 @@ def cross_check_against_posteriordb( pdb_path = ( str(posteriordb_root) if posteriordb_root is not None - else env_path if env_path else None + else env_path + if env_path + else None ) # Priority: local PosteriorDatabase → PosteriorDatabaseGithub fallback. diff --git a/tuningfork/calibration/_gate/verdict.py b/tuningfork/calibration/_gate/verdict.py index dffd89e1..5d5b5520 100644 --- a/tuningfork/calibration/_gate/verdict.py +++ b/tuningfork/calibration/_gate/verdict.py @@ -293,15 +293,15 @@ def _assemble_verdict( # Add the new REPORTED margins (never verdict; bias effect sizes in GT-σ units) if _bias_sigma_at_argmax_z is not None: - margins["max_abs_mean_z"][ - "bias_sigma_at_argmax_z" - ] = _bias_sigma_at_argmax_z + margins["max_abs_mean_z"]["bias_sigma_at_argmax_z"] = ( + _bias_sigma_at_argmax_z + ) if _bias_sigma_max_at_z4 is not None: margins["max_abs_mean_z"]["bias_sigma_max_at_z4"] = _bias_sigma_max_at_z4 if _achieved_bias_bound_sigma is not None: - margins["max_abs_mean_z"][ - "achieved_bias_bound_sigma" - ] = _achieved_bias_bound_sigma + margins["max_abs_mean_z"]["achieved_bias_bound_sigma"] = ( + _achieved_bias_bound_sigma + ) overall_verdict = _worst(overall_verdict, band) diff --git a/tuningfork/catalog/_timing.py b/tuningfork/catalog/_timing.py index 5f4501cb..1cf1a7da 100644 --- a/tuningfork/catalog/_timing.py +++ b/tuningfork/catalog/_timing.py @@ -104,7 +104,6 @@ def format_timing_context(recipe: Recipe) -> dict[str, str]: """ budget = recipe.calibration_budget or {} - n_warmup = budget.get("n_warmup") n_samples = budget.get("n_samples") num_chains = budget.get("num_chains") warmup_wall = budget.get("warmup_wall_seconds") diff --git a/tuningfork/catalog/expectands.py b/tuningfork/catalog/expectands.py index 62647179..e06937a7 100644 --- a/tuningfork/catalog/expectands.py +++ b/tuningfork/catalog/expectands.py @@ -1607,7 +1607,7 @@ def compare_reports( cost_blocked: tuple[str, ...] = () if not is_rate: cost_blocked += ( - f"{statistic} is not a rate; cost normalisation does not " "apply", + f"{statistic} is not a rate; cost normalisation does not apply", ) elif base_val is None or cand_val is None: # The statistic itself is undefined on one side, so no cost diff --git a/tuningfork/catalog/notebooks/catalog_explorer.py b/tuningfork/catalog/notebooks/catalog_explorer.py index 2fcafe7a..0fb44771 100644 --- a/tuningfork/catalog/notebooks/catalog_explorer.py +++ b/tuningfork/catalog/notebooks/catalog_explorer.py @@ -33,7 +33,6 @@ def _(): import arviz as az import marimo as mo - import matplotlib.pyplot as plt from tuningfork.catalog import ( cached_idata_for_recipe, diff --git a/tuningfork/groundtruth/_emit.py b/tuningfork/groundtruth/_emit.py index 842009b1..e8e046eb 100644 --- a/tuningfork/groundtruth/_emit.py +++ b/tuningfork/groundtruth/_emit.py @@ -23,7 +23,7 @@ import json import platform import subprocess -from datetime import datetime, timezone +from datetime import UTC, datetime from pathlib import Path from typing import Any @@ -207,7 +207,7 @@ def write_gt_artifacts( "x64_enabled": bool(jax.config.read("jax_enable_x64")), "device": jax.devices()[0].platform, "platform": platform.platform(), - "timestamp_utc": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"), + "timestamp_utc": datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%SZ"), "code_sha": _git_sha(), } if reproduced_from is not None: diff --git a/tuningfork/model/irt_2pl.py b/tuningfork/model/irt_2pl.py index 4f2d3979..06f3b015 100644 --- a/tuningfork/model/irt_2pl.py +++ b/tuningfork/model/irt_2pl.py @@ -80,7 +80,9 @@ assert set(np.unique(_responses_flat).tolist()) <= { 0.0, 1.0, -}, f"Response values must be binary {{0, 1}}, got {set(np.unique(_responses_flat).tolist())}" +}, ( + f"Response values must be binary {{0, 1}}, got {set(np.unique(_responses_flat).tolist())}" +) # --------------------------------------------------------------------------- # NumPyro model (NCP IRT 2PL) diff --git a/tuningfork/model/lotka_volterra.py b/tuningfork/model/lotka_volterra.py index 7333e18b..1bd4928c 100644 --- a/tuningfork/model/lotka_volterra.py +++ b/tuningfork/model/lotka_volterra.py @@ -91,9 +91,9 @@ T_OBS, 2, ), f"Expected OBSERVATIONS shape ({T_OBS}, 2), got {OBSERVATIONS.shape}" -assert OBSERVATION_TIMES.shape == ( - T_OBS, -), f"Expected OBSERVATION_TIMES shape ({T_OBS},), got {OBSERVATION_TIMES.shape}" +assert OBSERVATION_TIMES.shape == (T_OBS,), ( + f"Expected OBSERVATION_TIMES shape ({T_OBS},), got {OBSERVATION_TIMES.shape}" +) #: Ground-truth parameters used to generate synthetic data MU_TRUE: dict[str, float] = { diff --git a/tuningfork/model/radon.py b/tuningfork/model/radon.py index 8588a356..5ea3e93b 100644 --- a/tuningfork/model/radon.py +++ b/tuningfork/model/radon.py @@ -63,9 +63,9 @@ _county_min = int(COUNTY_IDX.min()) _county_max = int(COUNTY_IDX.max()) assert _county_min == 0, f"Expected county_idx min=0, got {_county_min}" -assert ( - _county_max == N_COUNTIES - 1 -), f"Expected county_idx max={N_COUNTIES - 1}, got {_county_max}" +assert _county_max == N_COUNTIES - 1, ( + f"Expected county_idx max={N_COUNTIES - 1}, got {_county_max}" +) # --------------------------------------------------------------------------- # NumPyro model (NCP varying-intercept hierarchical radon) diff --git a/tuningfork/model/stoch_vol.py b/tuningfork/model/stoch_vol.py index c237f5f0..c21033e6 100644 --- a/tuningfork/model/stoch_vol.py +++ b/tuningfork/model/stoch_vol.py @@ -74,9 +74,9 @@ def _load_returns(path: Path) -> np.ndarray: RETURNS: jnp.ndarray = jnp.array(_load_returns(_CSV_PATH), dtype=jnp.float32) # Validate shape -assert RETURNS.shape == ( - T_LENGTH, -), f"Expected RETURNS shape ({T_LENGTH},), got {RETURNS.shape}" +assert RETURNS.shape == (T_LENGTH,), ( + f"Expected RETURNS shape ({T_LENGTH},), got {RETURNS.shape}" +) # --------------------------------------------------------------------------- # NumPyro model (NCP recursive AR(1) stochastic volatility) diff --git a/tuningfork/recipes/_attempt_evidence.py b/tuningfork/recipes/_attempt_evidence.py index f4f4d4b6..35ae8288 100644 --- a/tuningfork/recipes/_attempt_evidence.py +++ b/tuningfork/recipes/_attempt_evidence.py @@ -25,7 +25,7 @@ import math from collections.abc import Mapping from dataclasses import replace -from datetime import datetime, timezone +from datetime import UTC, datetime from typing import Any from tuningfork.recipes._base import Recipe @@ -85,7 +85,7 @@ def _canonical_json(value: Any) -> bytes: def _timestamp(value: str | None) -> str: if value is None: - value = datetime.now(timezone.utc).isoformat().replace("+00:00", "Z") + value = datetime.now(UTC).isoformat().replace("+00:00", "Z") if not isinstance(value, str): raise TypeError("recorded_at must be an ISO-8601 string") try: @@ -94,7 +94,7 @@ def _timestamp(value: str | None) -> str: raise ValueError("recorded_at must be an ISO-8601 timestamp") from exc if parsed.tzinfo is None or parsed.utcoffset() is None: raise ValueError("recorded_at must include timezone information") - return parsed.astimezone(timezone.utc).isoformat().replace("+00:00", "Z") + return parsed.astimezone(UTC).isoformat().replace("+00:00", "Z") def build_attempt_record( diff --git a/tuningfork/recipes/_base.py b/tuningfork/recipes/_base.py index 09affdd0..f0e1ceb5 100644 --- a/tuningfork/recipes/_base.py +++ b/tuningfork/recipes/_base.py @@ -447,8 +447,7 @@ def validate_warmup_num_chains(v: list[int] | None, n_phases: int) -> None: return if not isinstance(v, list): raise ValueError( - f"warmup_num_chains must be a list[int] or None; " - f"got {type(v).__name__!r}" + f"warmup_num_chains must be a list[int] or None; got {type(v).__name__!r}" ) if len(v) != n_phases: raise ValueError( diff --git a/tuningfork/recipes/_certification_intent.py b/tuningfork/recipes/_certification_intent.py index 365889ff..503d4cf9 100644 --- a/tuningfork/recipes/_certification_intent.py +++ b/tuningfork/recipes/_certification_intent.py @@ -223,7 +223,7 @@ def build_certification_intent( init_strategy=init_strategy, variant_label=variant_label, tuning_seed=seed, - timestamp_utc=_datetime.datetime.now(_datetime.timezone.utc).isoformat(), + timestamp_utc=_datetime.datetime.now(_datetime.UTC).isoformat(), tuningfork_version=tuningfork.__version__, blackjax_version=blackjax.__version__, jax_version=jax.__version__, diff --git a/tuningfork/recipes/_certification_record.py b/tuningfork/recipes/_certification_record.py index ae080ddf..349cb6ec 100644 --- a/tuningfork/recipes/_certification_record.py +++ b/tuningfork/recipes/_certification_record.py @@ -140,7 +140,9 @@ def import_legacy_current_view( lifecycle = ( "CURATED" if verdict in {"PASS", "FAIL"} - else "EVALUATED" if verdict == "REVIEW" else "DRAFT" + else "EVALUATED" + if verdict == "REVIEW" + else "DRAFT" ) diagnosis = recipe.failure_diagnosis failure = { diff --git a/tuningfork/recipes/_certification_runner.py b/tuningfork/recipes/_certification_runner.py index e4e18d6c..0f019b0b 100644 --- a/tuningfork/recipes/_certification_runner.py +++ b/tuningfork/recipes/_certification_runner.py @@ -251,7 +251,9 @@ def _record_error( lifecycle_stage=( "SAMPLED" if result is not None and result.receipt.status == "success" - else "GENERATED" if result is not None else "DRAFT" + else "GENERATED" + if result is not None + else "DRAFT" ), automatic_verdict="ERROR", rationale=rationale, @@ -360,7 +362,7 @@ def emit_low_recipe_for_cell( ) except Exception as error: # noqa: BLE001 wall_seconds = time.perf_counter() - started_at - note = f"ERROR during intent validation: " f"{type(error).__name__}: {error}" + note = f"ERROR during intent validation: {type(error).__name__}: {error}" _append_outcome(outcomes, model_name, warmup_name, sampler_name, note) return _result( model_name=model_name, @@ -510,7 +512,7 @@ def emit_low_recipe_for_cell( if sampler_name in LAPLACE_METHOD_NAMES: if model_name not in LAPLACE_PHI_THETA_SPLITS: raise ValueError( - f"Laplace model {model_name!r} has no declared " "phi/theta split" + f"Laplace model {model_name!r} has no declared phi/theta split" ) allowed_sites = LAPLACE_PHI_THETA_SPLITS[model_name][0] evaluation = evaluate_generated_run( diff --git a/tuningfork/recipes/_emit/_init_strategy.py b/tuningfork/recipes/_emit/_init_strategy.py index d26a23dd..dc6bf376 100644 --- a/tuningfork/recipes/_emit/_init_strategy.py +++ b/tuningfork/recipes/_emit/_init_strategy.py @@ -109,8 +109,7 @@ def emit_init_strategy( "_init_leaves, _init_treedef = jax.tree_util.tree_flatten(init_position)", "_init_keys = jax.random.split(_init_strategy_key, len(_init_leaves))", "init_position = _init_treedef.unflatten([", - " jax.random.uniform(k, x.shape, dtype=x.dtype, minval=%r, maxval=%r)" - % (low, high), + f" jax.random.uniform(k, x.shape, dtype=x.dtype, minval={low!r}, maxval={high!r})", " for k, x in zip(_init_keys, _init_leaves)", "])", "_init_position_is_prebatched = False", @@ -146,7 +145,7 @@ def emit_init_strategy( ) draw = ( "jax.random.uniform(k, (1,) + _init_leaf.shape, " - "dtype=_init_leaf.dtype, minval=%r, maxval=%r)" % (low, high) + f"dtype=_init_leaf.dtype, minval={low!r}, maxval={high!r})" ) else: jitter = ( @@ -155,8 +154,8 @@ def emit_init_strategy( else 0.5 ) draw = ( - "(%r * jax.random.normal(k, (1,) + _init_leaf.shape, " - "dtype=_init_leaf.dtype))" % jitter + f"({jitter!r} * jax.random.normal(k, (1,) + _init_leaf.shape, " + "dtype=_init_leaf.dtype))" ) lines += [ "_init_leaves, _init_treedef = jax.tree_util.tree_flatten(init_position)", @@ -166,7 +165,7 @@ def emit_init_strategy( "for _init_leaf_idx, _init_leaf in enumerate(_init_leaves):", " _init_leaf_keys = _init_keys[:, _init_leaf_idx]", " _init_new_leaves.append(", - " jnp.concatenate([%s for k in _init_leaf_keys], axis=0)" % draw, + f" jnp.concatenate([{draw} for k in _init_leaf_keys], axis=0)", " )", "init_position = _init_treedef.unflatten(_init_new_leaves)", "_init_position_is_prebatched = True", diff --git a/tuningfork/recipes/_emit/_warmup.py b/tuningfork/recipes/_emit/_warmup.py index 9de877da..e8b71013 100644 --- a/tuningfork/recipes/_emit/_warmup.py +++ b/tuningfork/recipes/_emit/_warmup.py @@ -1360,10 +1360,7 @@ def _emit_adjusted_mclmc_trajectory_tuning(ctx: dict[str, Any]) -> str: ) a("_traj_warmup_keys = jax.random.split(_traj_warmup_key, num_chains)") a("_traj_init_positions = jax.tree.map(") - a( - " lambda x: jnp.broadcast_to(x[None], (num_chains,) + x.shape), " - "init_position" - ) + a(" lambda x: jnp.broadcast_to(x[None], (num_chains,) + x.shape), init_position") a(")") a("") a("") @@ -1392,10 +1389,7 @@ def _emit_adjusted_mclmc_trajectory_tuning(ctx: dict[str, Any]) -> str: a("_traj_warm_states, _traj_adaptation, _traj_warm_total = _traj_tune_one(") a(" _traj_warmup_keys, _traj_init_states") a(")") - a( - "jax.block_until_ready(" - "(_traj_warm_states, _traj_adaptation, _traj_warm_total))" - ) + a("jax.block_until_ready((_traj_warm_states, _traj_adaptation, _traj_warm_total))") a("_traj_step_sizes = _traj_adaptation.step_size") a("_traj_inverse_mass_matrices = _traj_adaptation.inverse_mass_matrix") a("") @@ -1435,10 +1429,7 @@ def _emit_adjusted_mclmc_trajectory_tuning(ctx: dict[str, Any]) -> str: a(" inverse_mass_matrix=inverse_mass_matrix,") a(" integration_steps_params=(avg,),") a(" )") - a( - " return new_state, " - "(new_state.position, info.num_integration_steps)" - ) + a(" return new_state, (new_state.position, info.num_integration_steps)") a("") a(" _, (positions, integration_steps) = jax.lax.scan(") a(" step, state, scan_keys") @@ -1485,7 +1476,7 @@ def _emit_adjusted_mclmc_trajectory_tuning(ctx: dict[str, Any]) -> str: a("_traj_avg_star = float(max(_traj_scores, key=lambda avg: _traj_scores[avg]))") a("# Preserve the historical upper-bound estimate based on candidate averages.") a("_traj_pilot_grad_estimate = sum(") - a(" int(2 * avg * _traj_n_pilot * num_chains) " "for avg in _traj_avg_grid") + a(" int(2 * avg * _traj_n_pilot * num_chains) for avg in _traj_avg_grid") a(")") a("_adapted_params = {") a(' "L": _traj_avg_star * _traj_step_sizes,') @@ -1494,7 +1485,7 @@ def _emit_adjusted_mclmc_trajectory_tuning(ctx: dict[str, Any]) -> str: a(' "_avg_star": _traj_avg_star,') a(' "_avg_search_ess_per_grad": _traj_scores,') a(' "_total_tuning_steps": (') - a(" int(jnp.asarray(_traj_warm_total)[0]) " "+ _traj_pilot_grad_estimate") + a(" int(jnp.asarray(_traj_warm_total)[0]) + _traj_pilot_grad_estimate") a(" ),") a("}") a("_state_post_warmup = _traj_dynamic_states") diff --git a/tuningfork/recipes/_emit_script.py b/tuningfork/recipes/_emit_script.py index 70c7e22d..4118be72 100644 --- a/tuningfork/recipes/_emit_script.py +++ b/tuningfork/recipes/_emit_script.py @@ -280,10 +280,7 @@ def _build_inference_loop( "# Re-init per-chain state (dynamic_hmc / dmhmc / ghmc: different state" " type than warmup)." ) - a( - f"_reinit_keys = jax.random.split(jax.random.key({reinit_seed})," - f" num_chains)" - ) + a(f"_reinit_keys = jax.random.split(jax.random.key({reinit_seed}), num_chains)") if warmup_is_perchain and not warmup_init_is_prebatched: a('_batched_step_size = _adapted_params["step_size"]') a('_batched_imm = _adapted_params["inverse_mass_matrix"]') diff --git a/tuningfork/recipes/_execution_receipt.py b/tuningfork/recipes/_execution_receipt.py index 2a1ce4ae..55a5e300 100644 --- a/tuningfork/recipes/_execution_receipt.py +++ b/tuningfork/recipes/_execution_receipt.py @@ -20,7 +20,7 @@ import re from collections.abc import Mapping, Sequence from dataclasses import dataclass -from datetime import datetime, timezone +from datetime import UTC, datetime from pathlib import PurePosixPath, PureWindowsPath from typing import Any @@ -82,7 +82,7 @@ def _timestamp(name: str, value: Any) -> str: raise ValueError(f"{name} must be ISO-8601") from exc if parsed.tzinfo is None or parsed.utcoffset() is None: raise ValueError(f"{name} must include a timezone") - return parsed.astimezone(timezone.utc).isoformat().replace("+00:00", "Z") + return parsed.astimezone(UTC).isoformat().replace("+00:00", "Z") @dataclass(frozen=True) diff --git a/tuningfork/recipes/_execution_telemetry.py b/tuningfork/recipes/_execution_telemetry.py index 00137f52..fb7efeac 100644 --- a/tuningfork/recipes/_execution_telemetry.py +++ b/tuningfork/recipes/_execution_telemetry.py @@ -126,8 +126,10 @@ def _validate_low_rank( raise ValueError( "batched low-rank marker chain count does not match manifest" ) - for s, u, l in zip(sigma, U, lam): - _validate_low_rank({"type": marker["type"], "sigma": s, "U": u, "lam": l}) + for s, u, lam_i in zip(sigma, U, lam): + _validate_low_rank( + {"type": marker["type"], "sigma": s, "U": u, "lam": lam_i} + ) return if not all( isinstance(x, (int, float)) diff --git a/tuningfork/recipes/_launcher.py b/tuningfork/recipes/_launcher.py index 97f9f480..1e3581a2 100644 --- a/tuningfork/recipes/_launcher.py +++ b/tuningfork/recipes/_launcher.py @@ -29,7 +29,7 @@ import uuid from collections.abc import Mapping from dataclasses import dataclass, field -from datetime import datetime, timezone +from datetime import UTC, datetime from pathlib import Path from typing import Any @@ -80,7 +80,7 @@ def _parse_timings(stdout: bytes) -> ExecutionTimings | None: required_fields = {"warmup_seconds", "sampling_seconds", "total_seconds"} if set(payload) != required_fields: raise ValueError( - "timing sentinel fields must be exactly " f"{sorted(required_fields)!r}" + f"timing sentinel fields must be exactly {sorted(required_fields)!r}" ) values: list[float] = [] for name in ("warmup_seconds", "sampling_seconds", "total_seconds"): @@ -439,7 +439,7 @@ def build_receipt(receipt_error: str | None) -> ExecutionReceipt: status="success" if receipt_error is None else "failed", run_id=run_dir.name, started_at=started_at, - finished_at=datetime.now(timezone.utc).isoformat(), + finished_at=datetime.now(UTC).isoformat(), manifest=manifest, source_sha256=source_sha256, program_path=_PROGRAM_FILENAME, @@ -586,7 +586,7 @@ def launch_generated_program( python_executable, tuple(sorted(() if env is None else env)) ) environment["child_working_directory"] = _WORK_DIRECTORY - started_at = datetime.now(timezone.utc).isoformat() + started_at = datetime.now(UTC).isoformat() child_env = os.environ.copy() if env is not None: diff --git a/tuningfork/recipes/_smc_certification_runner.py b/tuningfork/recipes/_smc_certification_runner.py index ed68cc71..4d0a0712 100644 --- a/tuningfork/recipes/_smc_certification_runner.py +++ b/tuningfork/recipes/_smc_certification_runner.py @@ -251,9 +251,7 @@ def _try_persist_attempt( updated = replace(updated, calibration_budget=budget) return updated, updated.save(root), attempt_id, None except Exception as error: # noqa: BLE001 - note = ( - "attempt recording/persistence failed: " f"{type(error).__name__}: {error}" - ) + note = f"attempt recording/persistence failed: {type(error).__name__}: {error}" return base, None, attempt_id, note