Repository navigation
Upgrade uv.lock to latest compatible versions - #265
Conversation
`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>
|
Pushed 9e7f921, a fix for the Cause: Fix:
Evidence:
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 |
|
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)
Trace structure (round 2)
Pre-existing issue found along the way (separate PR):
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 — 🤖 Blackjax-devs AI TL |
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
|
Bumped blackjax to 1.7.0 in — 🤖 Blackjax-devs AI TL |
|
The blackjax 1.7.0 bump surfaces one
tuningfork's gate already classifies NaN as FAIL, but — 🤖 Blackjax-devs AI TL |
|
Full CI picture with blackjax 1.7.0: three failures, all from 1.7.0's corrected diagnostics on degenerate or tied chains.
— 🤖 Blackjax-devs AI TL |
…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
|
Pushed the two tuningfork-side fixes for blackjax 1.7.0:
Locally: the fast suite gives 2089 passed and 5 skipped,
— 🤖 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
Summary
Upgrades
uv.lockto the latest compatible versions, and fixes what that upgrade surfaced. Every fix is forward: no package is held back.Lock
uv lock --upgrade.lax.scanis 2–10× slower on CPU from jax 0.10.2) and guards MCLMC stage-3Lagainst zero ESS (blackjax#1046).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 buildingNormal(loc=NaN).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_verdictnow maps a non-finiterhat_max,min_bulk_essormax_abs_mean_ztoNone, lists it innonfinite_metricsand forces the verdict to FAIL.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
test-fast,test-slowandtest-e2eall green on blackjax 1.7.1.🤖 Generated with Claude Code
https://claude.ai/code/session_013EfAQsGd5umbz3Yk97TMgx