Skip to content

Upgrade uv.lock to latest compatible versions - #265

Merged
junpenglao merged 7 commits into
mainfrom
deps-lock-upgrade
Oct 2, 2026
Merged

junpenglao merged 7 commits into
mainfrom
deps-lock-upgrade

Conversation

@junpenglao

@junpenglao junpenglao commented Oct 2, 2026 •

Copy link
Copy Markdown
Member

Summary

Upgrades uv.lock to the latest compatible versions, and fixes what that upgrade surfaced. Every fix is forward: no package is held back.

Lock

  • uv lock --upgrade.
  • blackjax goes to 1.7.1. That release jit-compiles the sampling loops (blackjax#1045; eager lax.scan is 2–10× slower on CPU from jax 0.10.2) and guards MCLMC stage-3 L against zero ESS (blackjax#1046).
  • Also jax 0.11.2, numpy 2.5, numpyro 0.22 and arviz 1.3.

Fixes for the new versions

  • 9e7f921, lotka_volterra: numpyro 0.22 validates distribution arguments by default. A uniform-init draw makes the ODE solve non-finite, so the likelihood now contributes −inf there instead of building Normal(loc=NaN).
    • At finite points the log-density and its gradients are bit-identical (201-seed check).
    • No cache or certification key hashes model source.
  • 3a1ef67, gate: blackjax 1.7.0 correctly returns R-hat NaN and ESS 0 for constant chains; 1.6.2 had arbitrary finite values.
    • _assemble_verdict now maps a non-finite rhat_max, min_bulk_ess or max_abs_mean_z to None, lists it in nonfinite_metrics and forces the verdict to FAIL.
    • Strict non-finite guards and existing contract tests are unchanged.
  • 39f5d83, expectands test: since blackjax averages tied ranks (#1031), blackjax and arviz now agree on indicator ESS. The test used to assert that they diverge.

Merges main (#266, the SMC space fix) along the way.

Test plan

  • CI: pre-commit, test-fast, test-slow and test-e2e all green on blackjax 1.7.1.
  • The Lotka-Volterra fix passed two "are you sure" rounds: gradient equivalence, trace structure, and the SMC prior/likelihood split.

🤖 Generated with Claude Code

https://claude.ai/code/session_013EfAQsGd5umbz3Yk97TMgx

junpenglao and others added 2 commits October 2, 2026 09:17
`uv lock --upgrade` (uv 0.12.13) moves 126 packages to their latest
compatible versions. Load-bearing: blackjax 1.6.1 -> 1.6.2 (PyPI), jax/
jaxlib 0.11.0 -> 0.11.2, numpy 2.4.6 -> 2.5.3, scipy 1.17.1 -> 1.18.1,
numpyro 0.21 -> 0.22, arviz 1.1 -> 1.3. Major jumps outside the sampling
path: datasets 4 -> 5, pyarrow 24 -> 25, xxhash 4, pymdown-extensions 11;
typer/shellingham dropped from the resolution.

Verified locally on py3.13 (`uv sync --all-extras --all-groups`):
pre-commit passes. `pytest tests -m fast` runs after the blackjax suite
finishes (serialized heavy compute on the shared box); result on the PR.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Finding: PR #265's uv.lock relock (blackjax 1.6.1->1.6.2, jax/jaxlib
0.11.0->0.11.2, numpy 2.4.6->2.5.3, scipy 1.17.1->1.18.1, numpyro
0.21.0->0.22.0, arviz 1.1->1.3) made
tests/e2e/test_nightly_regression.py::test_lotka_dense_imm_inner_nuts_seed_20260713_passes
fail in ~10s (not the ~200s full warmup+sampling run) with
GeneratedProgramError -> ValueError: Normal distribution got invalid loc
parameter. A single-package bisect (one group downgraded at a time,
confirmed by re-upgrading only the culprit) isolated this to numpyro
0.21.0 -> 0.22.0; everything else held at the new lock still passed once
numpyro alone was downgraded, and re-upgrading numpyro alone reproduced
the failure again. Root cause: numpyro/distributions/distribution.py's
Distribution._validate_args class default changed from a hardcoded
`False` (0.21.0) to `_VALIDATION_ENABLED` (0.22.0, already True by
default) -- numpyro 0.22 turns distribution-argument validation on by
default. initialize_model's first, non-retried structural trace
(_get_model_transforms) draws alpha/beta/gamma/delta/u0/v0/sigma_obs via
init_to_uniform at the pinned seed 20260713; this draw is bit-identical
across both numpyro versions, and tuningfork's ProbDiffEq ODE solve
(_solve_lv) returns a non-finite trajectory for it in both versions too
(confirmed with validation forced off). Previously that NaN was silently
recorded in the throwaway structural trace and later discarded by
find_valid_initial_params's own finite-potential/gradient retry loop;
numpyro 0.22 now raises before that retry loop ever runs.

Fix: per JP's fix-forward policy, hold numpyro's version steady and make
the model itself tolerate a non-finite solve instead. lotka_volterra_inverse
now computes `finite = isfinite(u_mean) & isfinite(scale)`, substitutes
finite placeholders into the "obs" Normal via jnp.where when the solve
blew up, and adds an explicit `numpyro.factor("lotka_ode_finite_guard",
where(finite, 0, -inf))` so a non-finite draw is still a zero-probability
region of the posterior -- the same outcome the old NaN-discard retry
produced, just expressed without an invalid distribution parameter.

Equivalence evidence: a scratch check (old vs. patched model, same
numpyro 0.22, log_density under identical substituted parameters) over
201 seeds' worth of init_to_uniform draws found 30 draws where the old
model's logdensity was finite (max |old - new| = 0.0, bit-identical) and
171 non-finite draws, all of which the patched model now scores as
exactly -inf. No behavior change wherever the ODE solve is finite.

Cache/catalog impact: none. Grepped every hashlib/inspect.getsource call
site under tuningfork/ -- all cache/recipe/plan/cert keys hash either a
declarative config dict (canonical_json) or a committed data artifact
(summary.json/draws.npz), never a model .py source file. The git-HEAD
`code_sha` recorded in reference/metadata.json is explicitly documented
in tuningfork/_cache_io.py as an audit trail, not an invalidation
criterion (worklog/decisions/2026-05-11-phase0-reference-protocol-refinements.md
§ 7). The committed lotka_volterra reference/groundtruth artifacts need
no re-cert.

Verification: tests/e2e/test_nightly_regression.py's lotka seed-20260713
test passes (182.66s, matching the pre-upgrade baseline); every
lotka-touching test under tests/ (test_catalog_render.py,
test_emit_config_wiring.py, test_dispatch.py,
test_telemetry_saturation.py, test_interactive_helpers.py,
test_nuts_multichain.py::test_load_explicit_positions_lotka_volterra,
test_generated_contract_e2e.py::test_x64_requirement_is_emitted_before_generated_execution)
passes; `pytest tests -m fast -q -n 2` passes (2083 passed, 5 skipped);
pre-commit is clean. uv.lock is untouched -- numpyro stays at 0.22.0.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@junpenglao

Copy link
Copy Markdown
Member Author

Pushed 9e7f921, a fix for the test-slow failure that keeps the upgrade (no version holds).

Cause: test_lotka_dense_imm_inner_nuts_seed_20260713_passes failed because numpyro 0.22 changed Distribution._validate_args to default to True; in 0.21 it was hard-coded False. Bisected to numpyro alone; blackjax, jax, numpy and scipy are all cleared. At this seed, the init_to_uniform draw gives a non-finite Lotka-Volterra ODE solve, and that draw is the same under both versions. initialize_model traces the model once, without retrying, to work out each parameter's constraints. In 0.21 that trace absorbed the Normal(loc=NaN) silently, and find_valid_initial_params then retried. In 0.22 validation raises inside that trace, so the program crashes before the retry loop runs.

Fix: lotka_volterra.py no longer builds an invalid Normal.

  • Where the solve is non-finite, the model uses finite placeholder loc/scale and adds an explicit numpyro.factor(..., -inf).
  • So those draws remain zero-probability, which is the same rejection outcome as before, and nothing invalid is ever passed to a distribution.

Evidence:

  • Log-density, old vs. patched model, over 201 init draws including the pinned seed: the 30 finite draws are bit-identical (max |diff| = 0.0), and all 171 non-finite draws now give exactly −inf.
  • No cache or recertification impact. No cache, recipe, certification or reference key hashes model source. code_sha is audit-only (_cache_io.py:256).
  • The failing test passes in 182.66 s; every Lotka-Volterra-touching test passes; the fast suite gives 2083 passed and 5 skipped; pre-commit is clean.

Still running before this is ready: gradient equivalence at finite points, and checking that predictive and log-likelihood consumers are unaffected by the new factor site.

— 🤖 Blackjax-devs AI TL

@junpenglao

Copy link
Copy Markdown
Member Author

Both "are you sure" rounds on 9e7f921 are complete, and all checks pass. Each comparison is the old model vs. the patched one, both under numpyro 0.22.

Gradients (round 1)

  • At finite init draws, gradients of the log-density are bit-identical: 30 of 30 draws, 0 mismatches.
  • At the 171 non-finite draws, the gradient behaves exactly as before: 1 draw finite and 170 non-finite under both models.
  • What the sampler sees is unchanged: all 171 draws are rejected by both models, with 0 disagreements. So the jnp.where guard doesn't change behaviour through NaN gradients.

Trace structure (round 2)

  • The latent parameter keys and the prior-sample keys are unchanged.
  • The new lotka_ode_finite_guard factor is an observed site, so tuningfork's generic prior/likelihood split routes it into the likelihood block.
  • build_smc_logfns comparison:
    • logprior_fn matches on 0/201 mismatches.
    • loglik_fn is bit-exact at every point where the old model is finite, evaluated in the unconstrained space that potential_fn expects.
    • It is never "repaired" to a finite value where it should be −inf.

Pre-existing issue found along the way (separate PR):

  • build_smc_logfns evaluates one position under two conventions. potential_fn treats it as unconstrained, while logprior_fn treats it as constrained via log_density, with no Jacobian.
  • As a result, the SMC likelihood is wrong for models with constrained parameters.
  • It's latent: no committed catalog recipe uses SMC. It is unrelated to this PR, and a fix is in progress on its own branch.

Merge sequencing: this PR merges after the blackjax jit-compile fix ships in a blackjax patch release; blackjax is then bumped here. On jax 0.11.2, blackjax 1.6.2 runs its sampling loops without jit, which would make tuningfork's recipes and certification runs 2–10× slower.

— 🤖 Blackjax-devs AI TL

junpenglao and others added 2 commits October 2, 2026 15:52
blackjax 1.7.0 jit-compiles its sampling loops (blackjax#1045). On jax
0.11.2 (this lock) blackjax 1.6.2 runs window_adaptation /
run_inference_algorithm via eager lax.scan, which is 2-10x slower on CPU
since jax 0.10.2 (jax-ml/jax#37465) -- tuningfork's emitted recipes and
certify_reference.py call those entry points un-jitted. 1.7.0's CI
benchmark on jax 0.11.2: NUTS 8.3 s / HMC 14.8 s (vs 123.8 / 291.6 s
without the fix). Lock change is blackjax only (uv lock --upgrade-package).

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_013EfAQsGd5umbz3Yk97TMgx
@junpenglao

Copy link
Copy Markdown
Member Author

Bumped blackjax to 1.7.0 in uv.lock. The only lock change is blackjax 1.6.2 → 1.7.0. 1.7.0 jit-compiles run_inference_algorithm and the warmup loops (blackjax#1045). With 1.6.2 on jax 0.11.2, tuningfork's emitted recipes and certification runs would have been 2–10× slower. This resolves the merge sequencing noted above, so once CI is green this PR is ready to merge.

— 🤖 Blackjax-devs AI TL

@junpenglao

Copy link
Copy Markdown
Member Author

The blackjax 1.7.0 bump surfaces one test-fast failure: test_mclmc_nonfinite_evidence_survives_generated_evaluation. The test feeds constant (all-zero) chains.

  • blackjax 1.6.2 returned R-hat 3.10 and ESS 3.7 for those, an artifact of breaking tied ranks arbitrarily.
  • 1.7.0 correctly returns R-hat NaN (tied ranks averaged, blackjax#1031) and ESS 0 (#1020).

tuningfork's gate already classifies NaN as FAIL, but _build_margin raises rather than recording it. The test's purpose is that non-finite evidence survives evaluation, so the fix belongs in the gate: record a non-finite metric as a fail-closed FAIL margin. The test itself stays unchanged. Fix in progress.

— 🤖 Blackjax-devs AI TL

@junpenglao

Copy link
Copy Markdown
Member Author

Full CI picture with blackjax 1.7.0: three failures, all from 1.7.0's corrected diagnostics on degenerate or tied chains.

  1. test-fast, test_mclmc_nonfinite_evidence_survives_generated_evaluation: the gate must record a non-finite R-hat as a FAIL margin instead of raising (fix in progress, in tuningfork).
  2. test-e2e, test_mclmc_lrd_generated_execution_smoke: the generated program's telemetry has a non-finite value. The suspected cause is MCLMC's L divided by an ESS that is now 0 for degenerate chains (blackjax#1020). Diagnosing; if confirmed it's a blackjax fix plus 1.7.1.
  3. test-slow, test_tie_severity_tracks_cardinality_not_tie_fraction: the test asserted that blackjax's indicator ESS diverges more than 5× from arviz. Since blackjax 1.7.0 averages tied ranks (#1031), they agree (77.53 vs 77.53), so the test is being updated to assert agreement.

— 🤖 Blackjax-devs AI TL

junpenglao and others added 2 commits October 2, 2026 21:58
…te's single assembly point

Finding: tests/recipes/test_generated_certification.py::test_mclmc_nonfinite_evidence_survives_generated_evaluation
started failing after blackjax 1.7.0 (blackjax#1020, blackjax#1031):
constant chains now return NaN rhat (tied ranks averaged) and 0 ESS
instead of 1.6.2's arbitrary tie-break values -- the more correct
output. _classify_metric already fails closed on a non-finite value
(its half-open interval comparisons are all False against NaN, so it
falls through to "FAIL"), but _build_margin unconditionally raised
ValueError on any non-finite observed metric, turning an
already-correctly-classified FAIL into an uncaught crash -- and a
second, independent strict guard one layer up
(tuningfork/recipes/_generated_certification.py's _json_safe) rejects
the same bare NaN again when it reaches AutoGateVerdict.to_dict()'s raw
rhat_max/min_bulk_ess fields. Patching each of those strict guards
individually is whack-a-mole and risks weakening checks that exist to
catch real, unrelated bugs elsewhere in the same generic JSON-safety
helpers.

Fix: normalise once, at the single point _assemble_verdict first has
rhat_max/min_bulk_ess/max_abs_mean_z available. A non-finite value is
set to None there (so every existing `if <metric> is not None` branch
downstream -- margin-building, the cost block, to_dict() -- skips it
exactly like "not computed", with zero changes needed to those
branches or to the two strict guards) and its name is recorded in a
new AutoGateVerdict.nonfinite_metrics list; the overall verdict is then
forced to "FAIL" when that list is non-empty, since an undefined R̂/ESS/z
means the run cannot be certified regardless of what any other metric
says. _classify_metric, _build_margin, and the existing "margin raises
on a non-finite value" contract (tests/recipes/test_statistician_gate.py)
are all untouched.

Verified the one other call site that reads rhat_max/min_bulk_ess ahead
of this normalisation -- the W1/sigma equivalence gate's prerequisite
check in statistician_gate.py (`_rhat_ok = rhat_max is None or
(_classify_metric(rhat_max, ...) != "FAIL")`) -- needs no change: it
already evaluates to False for a non-finite rhat_max via
_classify_metric's existing fail-closed comparisons, so that stage was
already being skipped correctly.

The originally-failing test is unchanged, as required -- it now passes
against the fixed gate instead of the regressed blackjax output. The
pre-existing "_build_margin raises on non-finite" test also still
passes unchanged, since a non-finite value can no longer reach it.

Verification: the originally-failing test plus every test file touching
calibration/_gate (test_generated_certification.py, test_statistician_gate.py,
test_auto_gate_w1_integration.py, test_gate_calibration.py) -- 70 passed.
`pytest tests -m fast -q -n 2 -p no:benchmark` (worked around the
pre-existing, unrelated pytest-benchmark/xdist incompatibility reported
separately) -- 2089 passed, 5 skipped. Pre-commit clean.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_013EfAQsGd5umbz3Yk97TMgx
Finding: tests/test_expectands.py::test_tie_severity_tracks_cardinality_not_tie_fraction
started failing after blackjax 1.7.0 (blackjax#1031, tied ranks now
averaged instead of broken ordinally): a 2-valued indicator's
rank-normalised bulk-ESS no longer diverges by orders of magnitude from
arviz's -- 77.531398 (arviz) vs 77.531402 (blackjax), where the old
assertion expected the blackjax value to be at least 5x smaller. The
test's own purpose (severity keys on cardinality, not tie fraction) is
unaffected; only the "backends diverge on a 2-valued indicator" claim it
used to illustrate that purpose went stale.

Fix: assert agreement for the indicator case too, matching the existing
sticky-trace assertion (`pytest.approx(..., rel=0.05)`), and reworded the
docstring sentence to say the backends diverged on this case before
blackjax 1.7.0's tied-rank averaging. The cardinality-based
tie_severity assertions (backend-independent) are unchanged.

Grepped tuningfork/catalog/expectands.py for related backend-divergence
claims, per request -- reporting, not rewriting (policy text):
- expectands.py:100-103 -- the BACKENDS constant's module comment:
  "They agree closely on well-behaved continuous traces and can
  disagree by orders of magnitude on tied ones (see
  :func:`expectand_report`)." A general claim about tied traces, not
  specific to the 2-valued-indicator case this test exercised.
- expectands.py:112-126 -- _tie_disclosure()'s docstring and its
  returned caution string ("implementations differ and change between
  releases" / "implementations differ between backends and between
  releases"). Already hedged/non-specific -- explicitly designed not to
  assert how any named backend handles ties, "pinning that in a product
  message would turn an upstream correctness fix into a failing
  assertion here" -- so this one arguably doesn't need touching, but is
  adjacent to the same topic.

Verification: tests/test_expectands.py, full file including slow-marked
tests, single process -- 108 passed. Pre-commit clean.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_013EfAQsGd5umbz3Yk97TMgx
@junpenglao

Copy link
Copy Markdown
Member Author

Pushed the two tuningfork-side fixes for blackjax 1.7.0:

  • 3a1ef67, in _gate/verdict.py only: _assemble_verdict now normalizes a non-finite rhat_max, min_bulk_ess or max_abs_mean_z at the point where it is first assembled. The value becomes None, its name is listed in a new nonfinite_metrics field, and the verdict is forced to FAIL. A NaN no longer reaches _build_margin or the JSON guards, so every strict non-finite guard and the existing test_gate_margin_rejects_nonfinite_observed_metric contract stay unchanged. test_mclmc_nonfinite_evidence_survives_generated_evaluation passes unchanged.
  • 39f5d83: test_tie_severity_tracks_cardinality_not_tie_fraction now asserts that blackjax and arviz agree on the indicator's bulk ESS (rel=0.05). They used to diverge more than 5×, until blackjax 1.7.0 began averaging tied ranks (#1031).

Locally: the fast suite gives 2089 passed and 5 skipped, tests/test_expectands.py gives 108 passed, and the gate files give 70 passed.

test-e2e will still fail until blackjax 1.7.1 ships blackjax#1046, the guard for MCLMC stage-3 L = inf when the ESS is 0. After that release, this PR gets a lock bump to 1.7.1.

— 🤖 Blackjax-devs AI TL

blackjax 1.7.1 (blackjax#1046) guards MCLMC stage-3 L against ESS = 0 on
degenerate dimensions; 1.7.0 returned L = inf there, which made the
generated mclmc_lrd program fail closed on non-finite telemetry
(test_mclmc_lrd_generated_execution_smoke). Lock change is blackjax only;
that test passes locally with 1.7.1.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_013EfAQsGd5umbz3Yk97TMgx
@junpenglao
junpenglao merged commit 8bb7081 into main Oct 2, 2026
4 checks passed
@junpenglao
junpenglao deleted the deps-lock-upgrade branch October 2, 2026 21:22
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant