diff --git a/CHANGELOG.md b/CHANGELOG.md index 5fa8d0cd1..b426bcba4 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -419,6 +419,7 @@ guidance; the itemized changes follow. ### Fixed - **`GraphManager.step`, `run` and `run_adaptive` no longer store a step that reshapes a state leaf** (`MADD-ANO-220`, in 0.1.0 to 0.3.1): a list given for a scalar constant (`BallNode(initial_velocity=[1.0, 2.0])`), or an external input of another shape at a later step, broadcast a scalar leaf, and the state's checkpoint then did not load after `reset_state()`. The step is refused by name (compared once per trace, host-side: no compiled program changes), as `run_scan` always did; the FMU sidecar and bridge refuse it too. +- **`coupling_diagnostics()` describes the step that ran when the state is written afterwards** (MADD-ANO-231, never released): the float floor under `spectral_error_bound`, `precision_limited` and `spectral_usable` was measured on the live state, so `set_node_state` after a step moved that step's report (a bound of 0.0, flag set, with every member written to zero). The graph now keeps what the step left. A checkpoint saved after a member was written says so (an optional archive member), and the graph that loads it reports that group's bound as NaN and its `*_usable` flags False with a `not_usable_reason` until it steps; `_reports` is now a reserved node name. Action: none. - **`coupling_diagnostics()`: three numbers that read wrong with their flag set, and one found beside them** (MADD-ANO-222, 225, 226, 229, never released): `gradient_relative_error_bound` at the float floor took the change along one stand-in direction (5.1e-6 for 3.0e-5); `rho_spectral` read settled past eight interface scalars (0.273 for 0.219), from sampled rounding on a float32 group far from normal (0.250 for 0.206), and lost a loop below a field's rounding on a sub-cycled group (1e-12 for 1.25e-4). The gradient bound's undirected distance now takes an operator norm; a Krylov space still growing at the cap is never settled (`spectral_usable=False` above seven interface scalars, eight where they are the whole state); a certificate over every perturbation of the measured size replaces the samples; a sub-cycled boundary value is differentiated as `(1 - alpha) a + alpha b` (values unchanged, gradients of sub-cycled groups move at rounding level). Action: none; expect `spectral_usable=False` on more float32 and 16-bit groups whose Jacobian is far from normal (usable fraction on the search's draws 0.979 to 0.976). @@ -803,6 +804,7 @@ guidance; the itemized changes follow. exactly representable ### Verification +- **Tests for verdicts taken over a span of steps**: a `mask_unconverged` window whose early steps hit the cap and whose last step converged, the profiler's fractions on such a window, and the batched `strict_convergence` firing gate; the slow normal-contraction property asserts the bound where `spectral_usable` and holds the flag's share of its draws to a floor (it asserted the flag on every draw, which no claim promises). `docs/user_guide/inspection.md` states what `diagnostics=True` costs (a GPU measurement in `benchmarks/results/gpu_eigvals_probe/`; MADD-ANO-232 records a not-usable value that differs in kind between backends). - **The step-program gate lowers each graph from emptied jax caches** (`scripts/capture_step_programs.py`, `tests/core/test_step_program_digests.py`): eleven gate graphs failed in a slow-lane process with no program changed. The lowered text also says how many identical copies of its own helpers (`isinf`, `frexp`, `_where`) jax wrote out, which depends on what its bounded trace caches still hold. `graph_digests` now calls `jax.clear_caches()` before it builds a graph; the capture of 396eb59a is unchanged and still holds for all 24 graphs on jax 0.10.2, 0.11.0 and 0.11.2. - **REST write-sequence oracle: a run that fails part-way is held to a replay of the steps it counted** (`tests/property/injected_failures.py`): the rule "an empty graph's clock stays" was wrong for a graph emptied over the API, whose steps the streams go on counting, and failed the slow hundred-request machine; the server was consistent. A per-push example now runs an emptied and a never-filled graph with a later slice failing - **MADD-VER-002's acceptance band no longer admits the defect it measured**: diff --git a/benchmarks/results/gpu_eigvals_probe/RESULT.md b/benchmarks/results/gpu_eigvals_probe/RESULT.md new file mode 100644 index 000000000..461d5dbcb --- /dev/null +++ b/benchmarks/results/gpu_eigvals_probe/RESULT.md @@ -0,0 +1,51 @@ +# GPU probe: the non-symmetric eigenvalue solve in coupling diagnostics + +One-off local run, maintainer-approved, 2026-10-07. Tree: `release/0.4.0` at `1f66801c`. +Machine: NVIDIA RTX A2000 8 GB Laptop GPU. jax/jaxlib 0.11.0 with `jax-cuda12-plugin` and +`jax-cuda12-pjrt` 0.11.0 installed into a scratch `--target` directory (the shared venv has no +CUDA plugin and was not modified). Scripts: `probe.py`, `cost_split.py`; raw results +`probe_{cpu,gpu}{,_x64}.json`, logs `run_*.log`. + +## Question + +PR #248 put `jnp.linalg.eigvals` of a small non-symmetric matrix (at most 9 × 9) inside the +diagnostics branch of the coupled step. It had only ever run on CPU. + +## Answers + +1. **It runs on the GPU backend**, eagerly and under `jit`, inside `lax.scan`, `vmap` and + `lax.cond`, in float32 and float64. Values equal the CPU's to rounding (worst relative + difference 6e-7 in float32, 2e-15 in float64) and NumPy's float64 reference. +2. **A coupled pair with `diagnostics=True` compiles and runs on GPU** through `step` and + `run_scan`. With diagnostics off, states equal the CPU's to 1e-7 (float32) and 4e-16 (float64). +3. **Reports agree where they are usable.** In float64 (flags True) every reported number equals + the CPU's to rounding. In float32 this pair sits at its float floor: `spectral_usable` and + `gradient_bound_usable` are False on both backends, and the not-usable `rho_spectral` differs + between them (1.0e-5 vs 1.0e-5 after `step`; 2.3e-5 vs 3.9e-5 after `run_scan`), as noise does. +4. **One difference worth a line:** after 40 `step` calls in float32, with the residual exactly 0, + `gradient_relative_error_bound` reads 4.3e-6 on CPU and **NaN on GPU**, with + `gradient_bound_usable=False` on both. Inside the contract (the number is flagged not usable), + but a NaN where the other backend has a number. +5. **Cost on GPU, this tiny pair, float32, per step:** diagnostics off 0.3 ms; diagnostics on + 6 to 8 ms with `eigvals`; 2 to 3.4 ms with the repeated-squaring fallback forced. So the + eigenvalue solve (LAPACK on the host, a device round trip each call) is about 4 ms of a + diagnostics-on step on GPU, and the rest of the diagnostics program about 2 to 3 ms. On CPU a + diagnostics-on step of the same pair is about 0.25 to 0.5 ms. + +## An environment trap, not a library defect + +The first GPU pass failed in every `eigvals` call (and in `jnp.linalg.solve`) with +`Error loading CUDA libraries. GPU will not be used.` from jaxlib's GPU solver kernels, while +`eigh`, `svd`, `qr` and the whole solve path worked. Cause: the NVIDIA runtime libraries +installed as wheels in the venv were not on the loader path for those kernels. With +`LD_LIBRARY_PATH` pointing at `site-packages/nvidia/*/lib` everything runs. **On a pod, check +this before trusting a diagnostics run:** `python -c "import jax.numpy as jnp; print(jnp.linalg.eigvals(jnp.eye(3)))"` +on the GPU backend must not raise. + +## What follows + +- `eigvals` stays the GPU path (repeated squaring was the wrong number it replaced). The cost is + opt-in (`diagnostics=True`) and belongs in the user guide's diagnostics section. +- The NaN in a not-usable `gradient_relative_error_bound` on GPU is a small item for the next + coupling batch (a not-usable number should be the same kind of value on both backends). +- The one-line `eigvals` check joins the RunPod section-0 checklist. diff --git a/benchmarks/results/gpu_eigvals_probe/cost_split.py b/benchmarks/results/gpu_eigvals_probe/cost_split.py new file mode 100644 index 000000000..ca5c84570 --- /dev/null +++ b/benchmarks/results/gpu_eigvals_probe/cost_split.py @@ -0,0 +1,24 @@ +"""Where does the diagnostics-on step time go on GPU: the eigenvalue solve +(a host round trip) or the rest of the diagnostics program? Same pair, +once with eigvals and once with the squaring fallback forced.""" +import os, sys, time +os.environ["JAX_PLATFORMS"] = "cuda,cpu" +import jax, numpy as np +from maddening.core.coupling import acceleration as acc +if sys.argv[1] == "squaring": + acc._EIGVALS_BACKENDS = ("cpu",) +from maddening import GraphManager +from maddening.nodes import SpringDamperNode +gm = GraphManager() +gm.add_node(SpringDamperNode("left", 0.01, stiffness=30.0, damping=2.0, initial_position=1.5)) +gm.add_node(SpringDamperNode("right", 0.01, stiffness=30.0, damping=2.0)) +gm.add_edge("left", "right", "position", "anchor_position") +gm.add_edge("right", "left", "position", "anchor_position") +gm.add_coupling_group(["left", "right"], max_iterations=8, tolerance=1e-6, solver="ift", diagnostics=True) +gm.compile(); gm.step() +t = time.perf_counter() +for _ in range(100): gm.step() +step = (time.perf_counter() - t) / 100 +t = time.perf_counter(); gm.run_scan(2000); scan = (time.perf_counter() - t) +t = time.perf_counter(); gm.run_scan(2000); scan2 = (time.perf_counter() - t) / 2000 +print(sys.argv[1], jax.default_backend(), f"step {step*1e3:.3f} ms; run_scan second call {scan2*1e3:.4f} ms per step") diff --git a/benchmarks/results/gpu_eigvals_probe/probe.py b/benchmarks/results/gpu_eigvals_probe/probe.py new file mode 100644 index 000000000..3874c7b74 --- /dev/null +++ b/benchmarks/results/gpu_eigvals_probe/probe.py @@ -0,0 +1,99 @@ +"""One-off local GPU probe (maintainer-approved, 2026-10-07): does the +non-symmetric eigenvalue solve that coupling diagnostics now use run on a +GPU, inside jit / scan / vmap, and does a coupled group with diagnostics on +give the same report on GPU as on CPU? + +Run twice: PROBE_PLATFORM=cpu and PROBE_PLATFORM=gpu (sets JAX_PLATFORMS). +Each run writes probe_.json; compare.py diffs them. +""" +import json, os, sys, time +plat = os.environ["PROBE_PLATFORM"] +os.environ["JAX_PLATFORMS"] = "cuda,cpu" if plat == "gpu" else "cpu" +x64 = os.environ.get("PROBE_X64") == "1" +import jax +if x64: + jax.config.update("jax_enable_x64", True) +import jax.numpy as jnp +import numpy as np + +out = {"platform": plat, "x64": x64, "jax": jax.__version__, + "default_backend": jax.default_backend(), + "devices": [str(d) for d in jax.devices()]} +print(out, flush=True) + +def record(name, fn): + t = time.perf_counter() + try: + v = fn() + out[name] = {"ok": True, "value": v, "seconds": round(time.perf_counter() - t, 3)} + except Exception as e: # the probe's point is to see what raises + out[name] = {"ok": False, "error": f"{type(e).__name__}: {str(e)[:300]}", + "seconds": round(time.perf_counter() - t, 3)} + print(name, json.dumps(out[name])[:400], flush=True) + +rng = np.random.default_rng(0) +dt = jnp.float64 if x64 else jnp.float32 +mats = {n: jnp.asarray(rng.standard_normal((n, n)) * 0.4, dt) for n in (2, 5, 9)} +# a non-normal one: graded upper Hessenberg, the shape the estimator compresses to +H = np.triu(rng.standard_normal((9, 9)), -1) * (0.7 ** np.arange(9))[:, None] +mats["hess9"] = jnp.asarray(H, dt) + +def radius(M): + return jnp.max(jnp.abs(jnp.linalg.eigvals(M))) + +for k, M in mats.items(): + record(f"eigvals_eager_{k}", lambda M=M: float(radius(M))) + record(f"eigvals_jit_{k}", lambda M=M: float(jax.jit(radius)(M))) +ref = {k: float(np.max(np.abs(np.linalg.eigvals(np.asarray(M, np.float64))))) for k, M in mats.items()} +out["numpy_float64_reference"] = ref + +M9 = mats["hess9"] +record("eigvals_in_scan", lambda: [float(x) for x in jax.jit( + lambda M: jax.lax.scan(lambda c, s: (c, radius(M * s)), 0.0, jnp.asarray([0.5, 1.0, 1.5], dt))[1])(M9)]) +record("eigvals_in_vmap", lambda: [float(x) for x in jax.jit(jax.vmap(radius))( + jnp.stack([M9 * s for s in (0.5, 1.0, 1.5)]))]) +record("eigvals_in_cond", lambda: float(jax.jit( + lambda M: jax.lax.cond(M[0, 0] > -1e9, radius, lambda m: jnp.zeros((), m.dtype), M))(M9))) + +# the estimator's own function +from maddening.core.coupling import acceleration as acc +record("estimator_spectral_radius", lambda: float(jax.jit(acc._spectral_radius)(M9))) +record("estimator_squaring_fallback", lambda: float(jax.jit(acc._spectral_radius_small)(M9))) + +# a stock coupled pair with diagnostics on and off +from maddening import GraphManager +from maddening.nodes import SpringDamperNode + +def pair(diagnostics): + gm = GraphManager() + gm.add_node(SpringDamperNode("left", 0.01, stiffness=30.0, damping=2.0, initial_position=1.5)) + gm.add_node(SpringDamperNode("right", 0.01, stiffness=30.0, damping=2.0)) + gm.add_edge("left", "right", "position", "anchor_position") + gm.add_edge("right", "left", "position", "anchor_position") + gm.add_coupling_group(["left", "right"], max_iterations=8, tolerance=1e-6, + solver="ift", diagnostics=diagnostics) + gm.compile() + return gm + +def run_pair(diagnostics, n=40): + gm = pair(diagnostics) + t = time.perf_counter(); gm.step(); first = time.perf_counter() - t + t = time.perf_counter() + for _ in range(n - 1): + gm.step() + per = (time.perf_counter() - t) / (n - 1) + st = {k: {f: np.asarray(v).tolist() for f, v in d.items()} for k, d in gm._state.items() + if not k.startswith("_")} + rep = {} + if diagnostics: + for g, r in gm.coupling_diagnostics().items(): + rep[str(g)] = {k: (np.asarray(v).tolist() if not isinstance(v, (str, bool, type(None))) else v) + for k, v in r.items()} + return {"first_step_s": round(first, 3), "per_step_ms": round(per * 1e3, 3), "state": st, "report": rep} + +record("pair_diagnostics_off", lambda: run_pair(False)) +record("pair_diagnostics_on", lambda: run_pair(True)) +record("pair_diagnostics_on_run_scan", lambda: (lambda gm: (gm.run_scan(20), {str(g): {k: (np.asarray(v).tolist() if not isinstance(v, (str, bool, type(None))) else v) for k, v in r.items()} for g, r in gm.coupling_diagnostics().items()})[1])(pair(True))) + +json.dump(out, open(f"probe_{plat}{'_x64' if x64 else ''}.json", "w"), indent=1, default=str) +print("written", flush=True) diff --git a/benchmarks/results/gpu_eigvals_probe/probe_cpu.json b/benchmarks/results/gpu_eigvals_probe/probe_cpu.json new file mode 100644 index 000000000..88f9d64de --- /dev/null +++ b/benchmarks/results/gpu_eigvals_probe/probe_cpu.json @@ -0,0 +1,165 @@ +{ + "platform": "cpu", + "x64": false, + "jax": "0.11.0", + "default_backend": "cpu", + "devices": [ + "cpu:0" + ], + "eigvals_eager_2": { + "ok": true, + "value": 0.12508688867092133, + "seconds": 0.225 + }, + "eigvals_jit_2": { + "ok": true, + "value": 0.12508688867092133, + "seconds": 0.041 + }, + "eigvals_eager_5": { + "ok": true, + "value": 0.7798304557800293, + "seconds": 0.096 + }, + "eigvals_jit_5": { + "ok": true, + "value": 0.7798304557800293, + "seconds": 0.043 + }, + "eigvals_eager_9": { + "ok": true, + "value": 1.2581210136413574, + "seconds": 0.112 + }, + "eigvals_jit_9": { + "ok": true, + "value": 1.2581210136413574, + "seconds": 0.047 + }, + "eigvals_eager_hess9": { + "ok": true, + "value": 0.5837985873222351, + "seconds": 0.0 + }, + "eigvals_jit_hess9": { + "ok": true, + "value": 0.5837985873222351, + "seconds": 0.0 + }, + "numpy_float64_reference": { + "2": 0.12508688842884425, + "5": 0.7798302029153219, + "9": 1.2581206133678078, + "hess9": 0.5837984961035871 + }, + "eigvals_in_scan": { + "ok": true, + "value": [ + 0.29189929366111755, + 0.5837985873222351, + 0.8756980299949646 + ], + "seconds": 0.118 + }, + "eigvals_in_vmap": { + "ok": true, + "value": [ + 0.29189929366111755, + 0.5837985873222351, + 0.8756980299949646 + ], + "seconds": 0.146 + }, + "eigvals_in_cond": { + "ok": true, + "value": 0.5837985873222351, + "seconds": 0.055 + }, + "estimator_spectral_radius": { + "ok": true, + "value": 0.5837985873222351, + "seconds": 0.114 + }, + "estimator_squaring_fallback": { + "ok": true, + "value": 0.5837984681129456, + "seconds": 0.122 + }, + "pair_diagnostics_off": { + "ok": true, + "value": { + "first_step_s": 0.353, + "per_step_ms": 0.091, + "state": { + "left": { + "position": 2.2816643714904785, + "velocity": 8.540358543395996 + }, + "right": { + "position": 3.21620512008667, + "velocity": 9.014348983764648 + } + }, + "report": {} + }, + "seconds": 0.396 + }, + "pair_diagnostics_on": { + "ok": true, + "value": { + "first_step_s": 5.365, + "per_step_ms": 0.255, + "state": { + "left": { + "position": 2.2816643714904785, + "velocity": 8.540358543395996 + }, + "right": { + "position": 3.21620512008667, + "velocity": 9.014348983764648 + } + }, + "report": { + "left+right": { + "iterations": 2, + "total_iterations": 2, + "residual": 0.0, + "amplification": 1.0, + "error_estimate": 0.0, + "ratio_usable": true, + "gradient_error_estimate": 0.0, + "converged": true, + "rho_spectral": 1.0359331099607516e-05, + "spectral_error_bound": 1.9366693777556065e-06, + "spectral_usable": false, + "gradient_relative_error_bound": 4.262304173607845e-06, + "gradient_bound_usable": false, + "precision_limited": true + } + } + }, + "seconds": 6.397 + }, + "pair_diagnostics_on_run_scan": { + "ok": true, + "value": { + "left+right": { + "iterations": 2, + "total_iterations": 2, + "residual": 5.226464168117673e-07, + "amplification": 1.0010093450546265, + "error_estimate": 5.231739237387956e-07, + "ratio_usable": true, + "gradient_error_estimate": 5.231739237387956e-07, + "converged": true, + "rho_spectral": 2.2723110305378214e-05, + "spectral_error_bound": 3.5370385376154445e-06, + "spectral_usable": false, + "gradient_relative_error_bound": 5.320030595612479e-06, + "gradient_bound_usable": false, + "precision_limited": true + } + }, + "seconds": 5.213 + } +} \ No newline at end of file diff --git a/benchmarks/results/gpu_eigvals_probe/probe_cpu_x64.json b/benchmarks/results/gpu_eigvals_probe/probe_cpu_x64.json new file mode 100644 index 000000000..b865e06aa --- /dev/null +++ b/benchmarks/results/gpu_eigvals_probe/probe_cpu_x64.json @@ -0,0 +1,165 @@ +{ + "platform": "cpu", + "x64": true, + "jax": "0.11.0", + "default_backend": "cpu", + "devices": [ + "cpu:0" + ], + "eigvals_eager_2": { + "ok": true, + "value": 0.12508688923006284, + "seconds": 0.251 + }, + "eigvals_jit_2": { + "ok": true, + "value": 0.12508688923006284, + "seconds": 0.042 + }, + "eigvals_eager_5": { + "ok": true, + "value": 0.7798302084302585, + "seconds": 0.096 + }, + "eigvals_jit_5": { + "ok": true, + "value": 0.7798302084302585, + "seconds": 0.042 + }, + "eigvals_eager_9": { + "ok": true, + "value": 1.2581206637447568, + "seconds": 0.121 + }, + "eigvals_jit_9": { + "ok": true, + "value": 1.2581206637447568, + "seconds": 0.068 + }, + "eigvals_eager_hess9": { + "ok": true, + "value": 0.5837984972178895, + "seconds": 0.0 + }, + "eigvals_jit_hess9": { + "ok": true, + "value": 0.5837984972178895, + "seconds": 0.0 + }, + "numpy_float64_reference": { + "2": 0.12508688923006284, + "5": 0.7798302084302585, + "9": 1.2581206637447568, + "hess9": 0.5837984972178895 + }, + "eigvals_in_scan": { + "ok": true, + "value": [ + 0.29189924860894473, + 0.5837984972178895, + 0.8756977458268352 + ], + "seconds": 0.136 + }, + "eigvals_in_vmap": { + "ok": true, + "value": [ + 0.29189924860894473, + 0.5837984972178895, + 0.8756977458268352 + ], + "seconds": 0.149 + }, + "eigvals_in_cond": { + "ok": true, + "value": 0.5837984972178895, + "seconds": 0.078 + }, + "estimator_spectral_radius": { + "ok": true, + "value": 0.5837984972178895, + "seconds": 0.11 + }, + "estimator_squaring_fallback": { + "ok": true, + "value": 0.5837985476527052, + "seconds": 0.128 + }, + "pair_diagnostics_off": { + "ok": true, + "value": { + "first_step_s": 0.341, + "per_step_ms": 0.086, + "state": { + "left": { + "position": 2.2816641889740357, + "velocity": 8.54036092626129 + }, + "right": { + "position": 3.2162056478290686, + "velocity": 9.014346578747606 + } + }, + "report": {} + }, + "seconds": 0.383 + }, + "pair_diagnostics_on": { + "ok": true, + "value": { + "first_step_s": 5.569, + "per_step_ms": 0.201, + "state": { + "left": { + "position": 2.2816641889740357, + "velocity": 8.54036092626129 + }, + "right": { + "position": 3.2162056478290686, + "velocity": 9.014346578747606 + } + }, + "report": { + "left+right": { + "iterations": 2, + "total_iterations": 2, + "residual": 2.8518300604636335e-08, + "amplification": 1.0006420029605732, + "error_estimate": 2.8536609438055027e-08, + "ratio_usable": true, + "gradient_error_estimate": 2.8536609438055027e-08, + "converged": true, + "rho_spectral": 8.999999947148507e-06, + "spectral_error_bound": 3.0168976714134447e-08, + "spectral_usable": true, + "gradient_relative_error_bound": 5.226361320304525e-12, + "gradient_bound_usable": true, + "precision_limited": false + } + } + }, + "seconds": 6.602 + }, + "pair_diagnostics_on_run_scan": { + "ok": true, + "value": { + "left+right": { + "iterations": 2, + "total_iterations": 2, + "residual": 5.809858401802299e-07, + "amplification": 1.0010642761087856, + "error_estimate": 5.816041695294764e-07, + "ratio_usable": true, + "gradient_error_estimate": 5.816041695294764e-07, + "converged": true, + "rho_spectral": 8.999998349523697e-06, + "spectral_error_bound": 8.589090417530644e-07, + "spectral_usable": true, + "gradient_relative_error_bound": 1.4343104899068651e-11, + "gradient_bound_usable": true, + "precision_limited": false + } + }, + "seconds": 5.529 + } +} \ No newline at end of file diff --git a/benchmarks/results/gpu_eigvals_probe/probe_gpu.json b/benchmarks/results/gpu_eigvals_probe/probe_gpu.json new file mode 100644 index 000000000..61ac053ab --- /dev/null +++ b/benchmarks/results/gpu_eigvals_probe/probe_gpu.json @@ -0,0 +1,165 @@ +{ + "platform": "gpu", + "x64": false, + "jax": "0.11.0", + "default_backend": "gpu", + "devices": [ + "cuda:0" + ], + "eigvals_eager_2": { + "ok": true, + "value": 0.12508688867092133, + "seconds": 0.358 + }, + "eigvals_jit_2": { + "ok": true, + "value": 0.12508688867092133, + "seconds": 0.061 + }, + "eigvals_eager_5": { + "ok": true, + "value": 0.7798306345939636, + "seconds": 0.252 + }, + "eigvals_jit_5": { + "ok": true, + "value": 0.7798306345939636, + "seconds": 0.07 + }, + "eigvals_eager_9": { + "ok": true, + "value": 1.2581210136413574, + "seconds": 0.237 + }, + "eigvals_jit_9": { + "ok": true, + "value": 1.2581210136413574, + "seconds": 0.074 + }, + "eigvals_eager_hess9": { + "ok": true, + "value": 0.5837985873222351, + "seconds": 0.001 + }, + "eigvals_jit_hess9": { + "ok": true, + "value": 0.5837985873222351, + "seconds": 0.001 + }, + "numpy_float64_reference": { + "2": 0.12508688842884425, + "5": 0.7798302029153219, + "9": 1.2581206133678078, + "hess9": 0.5837984961035871 + }, + "eigvals_in_scan": { + "ok": true, + "value": [ + 0.29189929366111755, + 0.5837985873222351, + 0.8756974935531616 + ], + "seconds": 0.17 + }, + "eigvals_in_vmap": { + "ok": true, + "value": [ + 0.29189929366111755, + 0.5837985873222351, + 0.8756974935531616 + ], + "seconds": 0.236 + }, + "eigvals_in_cond": { + "ok": true, + "value": 0.5837985873222351, + "seconds": 0.084 + }, + "estimator_spectral_radius": { + "ok": true, + "value": 0.5837985873222351, + "seconds": 0.171 + }, + "estimator_squaring_fallback": { + "ok": true, + "value": 0.5838320255279541, + "seconds": 0.316 + }, + "pair_diagnostics_off": { + "ok": true, + "value": { + "first_step_s": 0.646, + "per_step_ms": 0.295, + "state": { + "left": { + "position": 2.2816643714904785, + "velocity": 8.540358543395996 + }, + "right": { + "position": 3.21620512008667, + "velocity": 9.014348030090332 + } + }, + "report": {} + }, + "seconds": 0.739 + }, + "pair_diagnostics_on": { + "ok": true, + "value": { + "first_step_s": 10.844, + "per_step_ms": 9.231, + "state": { + "left": { + "position": 2.2816643714904785, + "velocity": 8.540358543395996 + }, + "right": { + "position": 3.21620512008667, + "velocity": 9.014348030090332 + } + }, + "report": { + "left+right": { + "iterations": 2, + "total_iterations": 2, + "residual": 0.0, + "amplification": 1.0, + "error_estimate": 0.0, + "ratio_usable": true, + "gradient_error_estimate": 0.0, + "converged": true, + "rho_spectral": 1.0357935025240295e-05, + "spectral_error_bound": 1.9366693777556065e-06, + "spectral_usable": false, + "gradient_relative_error_bound": NaN, + "gradient_bound_usable": false, + "precision_limited": true + } + } + }, + "seconds": 12.906 + }, + "pair_diagnostics_on_run_scan": { + "ok": true, + "value": { + "left+right": { + "iterations": 2, + "total_iterations": 2, + "residual": 5.226464736551861e-07, + "amplification": 1.0010093450546265, + "error_estimate": 5.231739805822144e-07, + "ratio_usable": true, + "gradient_error_estimate": 5.231739805822144e-07, + "converged": true, + "rho_spectral": 3.888137507601641e-05, + "spectral_error_bound": 3.5370376281207427e-06, + "spectral_usable": false, + "gradient_relative_error_bound": 5.320030595612479e-06, + "gradient_bound_usable": false, + "precision_limited": true + } + }, + "seconds": 9.882 + } +} \ No newline at end of file diff --git a/benchmarks/results/gpu_eigvals_probe/probe_gpu_x64.json b/benchmarks/results/gpu_eigvals_probe/probe_gpu_x64.json new file mode 100644 index 000000000..e8f4f99ee --- /dev/null +++ b/benchmarks/results/gpu_eigvals_probe/probe_gpu_x64.json @@ -0,0 +1,165 @@ +{ + "platform": "gpu", + "x64": true, + "jax": "0.11.0", + "default_backend": "gpu", + "devices": [ + "cuda:0" + ], + "eigvals_eager_2": { + "ok": true, + "value": 0.12508688923006284, + "seconds": 0.342 + }, + "eigvals_jit_2": { + "ok": true, + "value": 0.12508688923006284, + "seconds": 0.074 + }, + "eigvals_eager_5": { + "ok": true, + "value": 0.7798302084302591, + "seconds": 0.264 + }, + "eigvals_jit_5": { + "ok": true, + "value": 0.7798302084302591, + "seconds": 0.1 + }, + "eigvals_eager_9": { + "ok": true, + "value": 1.2581206637447575, + "seconds": 0.297 + }, + "eigvals_jit_9": { + "ok": true, + "value": 1.2581206637447575, + "seconds": 0.13 + }, + "eigvals_eager_hess9": { + "ok": true, + "value": 0.5837984972178905, + "seconds": 0.002 + }, + "eigvals_jit_hess9": { + "ok": true, + "value": 0.5837984972178905, + "seconds": 0.002 + }, + "numpy_float64_reference": { + "2": 0.12508688923006284, + "5": 0.7798302084302585, + "9": 1.2581206637447568, + "hess9": 0.5837984972178895 + }, + "eigvals_in_scan": { + "ok": true, + "value": [ + 0.2918992486089452, + 0.5837984972178905, + 0.8756977458268366 + ], + "seconds": 0.219 + }, + "eigvals_in_vmap": { + "ok": true, + "value": [ + 0.2918992486089452, + 0.5837984972178905, + 0.8756977458268366 + ], + "seconds": 0.283 + }, + "eigvals_in_cond": { + "ok": true, + "value": 0.5837984972178905, + "seconds": 0.139 + }, + "estimator_spectral_radius": { + "ok": true, + "value": 0.5837984972178905, + "seconds": 0.189 + }, + "estimator_squaring_fallback": { + "ok": true, + "value": 0.5837985476527054, + "seconds": 0.314 + }, + "pair_diagnostics_off": { + "ok": true, + "value": { + "first_step_s": 0.642, + "per_step_ms": 0.322, + "state": { + "left": { + "position": 2.2816641889740366, + "velocity": 8.540360926261291 + }, + "right": { + "position": 3.216205647829069, + "velocity": 9.014346578747604 + } + }, + "report": {} + }, + "seconds": 0.706 + }, + "pair_diagnostics_on": { + "ok": true, + "value": { + "first_step_s": 10.537, + "per_step_ms": 11.237, + "state": { + "left": { + "position": 2.2816641889740366, + "velocity": 8.540360926261291 + }, + "right": { + "position": 3.216205647829069, + "velocity": 9.014346578747604 + } + }, + "report": { + "left+right": { + "iterations": 2, + "total_iterations": 2, + "residual": 2.8518300396787208e-08, + "amplification": 1.0006420029582321, + "error_estimate": 2.85366092300057e-08, + "ratio_usable": true, + "gradient_error_estimate": 2.85366092300057e-08, + "converged": true, + "rho_spectral": 9.00000001989587e-06, + "spectral_error_bound": 3.01685302491266e-08, + "spectral_usable": true, + "gradient_relative_error_bound": 5.2265488282509716e-12, + "gradient_bound_usable": true, + "precision_limited": false + } + } + }, + "seconds": 12.671 + }, + "pair_diagnostics_on_run_scan": { + "ok": true, + "value": { + "left+right": { + "iterations": 2, + "total_iterations": 2, + "residual": 5.809858403019144e-07, + "amplification": 1.001064276108897, + "error_estimate": 5.816041696513553e-07, + "ratio_usable": true, + "gradient_error_estimate": 5.816041696513553e-07, + "converged": true, + "rho_spectral": 8.999999394909697e-06, + "spectral_error_bound": 8.61603479694276e-07, + "spectral_usable": true, + "gradient_relative_error_bound": 1.4387405688167306e-11, + "gradient_bound_usable": true, + "precision_limited": false + } + }, + "seconds": 9.626 + } +} \ No newline at end of file diff --git a/docs/release_notes/v0.4.0.md b/docs/release_notes/v0.4.0.md index f6502d0b3..12a4816cf 100644 --- a/docs/release_notes/v0.4.0.md +++ b/docs/release_notes/v0.4.0.md @@ -7158,6 +7158,10 @@ takes for rounding can carry the dominant mode of a float32 group with a field below a hundredth of what it drives (0.157 for 0.349, settled); run such a group in float64 +`MADD-ANO-232` — *A `gradient_relative_error_bound` whose flag is `False` can +read NaN on a GPU backend where the CPU backend reads a number* — **open**. Do +not compare not-usable values between backends. + **Resolved.** `MADD-ANO-006`, `019`, `020`, `024`, `025`, `026`, `028`, `029`, `030`, `049`, `062` and `063` above, `064`, `065`, `066` and `067` in "[What the differential sharding harness @@ -7376,6 +7380,8 @@ its nodes on different clocks; now a `ValueError` at `compile()` and a 400 at `POST /graph/nodes`. `228` (resolved; carried by 0.1.0 to 0.3.1): REST read a boolean or numeric text as a node's timestep and clamped `/sim/profile`'s counts silently; now a 422. +`231` (never released): a coupling report was measured on the live state and +moved when the state was written after the step. ## Verification evidence diff --git a/docs/user_guide/inspection.md b/docs/user_guide/inspection.md index 21355f95e..8a8f439f4 100644 --- a/docs/user_guide/inspection.md +++ b/docs/user_guide/inspection.md @@ -361,6 +361,32 @@ relay, `N` for a loop of `N` sub-steps), and read a float32 group with a field under a hundredth of what drives it should be run in float64 before its flag is relied on (MADD-ANO-230). +**What `diagnostics=True` costs.** It is opt-in per group, and the work is done in every step of +that group. Beside the solve, the step runs 9 Jacobian-vector products for the spectrum (18 under +`convergence_norm="interface"` with a mapping or a transform on an internal edge, or a field more +than one internal edge reads) and `11 + 4 k + 5 n_p + 2 k n_p` more for the gradient bound +(`k <= 8`, `n_p` the probed constants), then small dense factorisations (QR, linear solves, an +SVD) and one non-symmetric eigenvalue solve of a matrix of at most 9 x 9. The eigenvalue solve +runs in LAPACK on the host; on a GPU backend it is a device round trip in every step. Measured +once on a two-spring pair in float32 (an RTX A2000 laptop GPU, jax 0.11.0, +`benchmarks/results/gpu_eigvals_probe/RESULT.md`): a step takes 0.3 ms with diagnostics off and +6 to 8 ms with them on, about 4 ms of it the eigenvalue solve; on CPU the same diagnostics-on step +takes 0.25 to 0.5 ms. + +**A report describes the step that ran.** The bound's float floor is measured on the state the +step returned. Writing the state afterwards (`set_node_state`, `PUT /graph/state/{node}`) does not +change the entry: it describes that step until the group steps again. A checkpoint is a copy of the +state, so one saved after a member's state was written holds the written state and not the returned +one. It says so, and the graph that loads it reports that group's `spectral_error_bound` as NaN and +its `spectral_usable`, `gradient_bound_usable` and `precision_limited` as `False`, with a +`not_usable_reason`, until the group steps. + +**A number whose flag is `False` is not a number to compare.** Where `gradient_bound_usable` is +`False` the value beside it can be finite, `inf` or NaN, and at the float floor it can differ in +kind between backends: on the pair above after 40 float32 steps (residual exactly 0), +`gradient_relative_error_bound` read 4.3e-6 on CPU and NaN on the GPU, with the flag `False` on +both (MADD-ANO-232). + ## Graphs that are not ready None of these methods compiles a graph or puts it back after a diff --git a/docs/validation/coupling_claims.yaml b/docs/validation/coupling_claims.yaml index 6c26c2690..f4179bfa4 100644 --- a/docs/validation/coupling_claims.yaml +++ b/docs/validation/coupling_claims.yaml @@ -3398,6 +3398,7 @@ claims: - tests/core/test_coupling_spectral_bound_where_edges_share_a_field.py::test_the_interface_norm_counts_a_field_once_for_every_edge_that_reads_it - tests/core/test_coupling_spectral_bound_where_edges_share_a_field.py::test_a_usable_bound_covers_the_distance_where_a_node_reads_a_field_twice - tests/property/test_coupling_targeted_search.py::test_a_usable_error_bound_is_never_below_the_distance_per_push + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_write_after_the_step_does_not_move_its_report status: verified notes: >- The round-6 audit found the interface norm's transformed reading @@ -3638,6 +3639,7 @@ claims: - tests/property/test_coupling_targeted_search.py::test_a_defect_the_search_reached_stays_fixed - tests/core/test_coupling_precision_floor.py::test_an_undeclared_sub_stepped_node_is_not_usable_at_the_floor - tests/property/test_differential_coupling_interactions.py::test_a_usable_spectral_bound_bounds_the_distance_in_every_row + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_write_after_the_step_does_not_move_its_report status: verified domains: f32: tests/core/test_coupling_report_flags_read_their_conditions.py::test_spectral_usable_is_false_where_the_krylov_space_did_not_settle @@ -3882,6 +3884,7 @@ claims: - tests/core/test_coupling_precision_floor.py::test_precision_limited_is_the_residual_at_or_below_the_floor - tests/core/test_coupling_precision_floor.py::test_a_residual_above_the_floor_is_not_precision_limited - tests/core/test_coupling_precision_floor.py::test_a_sub_stepped_node_is_floored_per_evaluation + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_write_after_the_step_does_not_move_its_report status: verified domains: f32: tests/core/test_coupling_precision_floor.py::test_precision_limited_is_the_residual_at_or_below_the_floor @@ -5042,14 +5045,21 @@ claims: - tests/core/test_sysid_mask_cuts_unconverged_windows.py::test_the_verdict_reads_every_coupling_group - tests/core/test_sysid_mask_cuts_unconverged_windows.py::test_the_mask_drops_the_window_the_report_calls_unconverged - tests/core/test_sysid_mask_reads_the_last_waveform_sweep.py::test_a_window_whose_first_sweep_hit_its_cap_is_kept_by_the_mask + - 'tests/core/test_sysid_mask_reads_every_step_of_the_window.py::test_a_window_whose_early_steps_hit_the_cap_is_dropped_though_its_last_step_converged' + - 'tests/core/test_sysid_mask_reads_every_step_of_the_window.py::test_the_same_window_is_kept_when_the_cap_lets_every_step_converge' + - tests/core/test_sysid_mask_reads_every_step_of_the_window.py::test_multiple_shooting_drops_the_recovering_window_and_keeps_the_converged_one status: verified domains: - f32: tests/core/test_sysid_mask_cuts_unconverged_windows.py::test_a_masked_window_contributes_nothing_whatever_its_samples_hold + f32: + - tests/core/test_sysid_mask_cuts_unconverged_windows.py::test_a_masked_window_contributes_nothing_whatever_its_samples_hold + - 'tests/core/test_sysid_mask_reads_every_step_of_the_window.py::test_a_window_whose_early_steps_hit_the_cap_is_dropped_though_its_last_step_converged[30-1]' f64: narrowed mixed_dtype: narrowed 16bit: narrowed jit: tests/core/test_sysid_mask_cuts_unconverged_windows.py::test_a_diverged_window_leaves_a_finite_gradient_of_the_other_windows - grad: tests/core/test_sysid_mask_cuts_unconverged_windows.py::test_a_diverged_window_leaves_a_finite_gradient_of_the_other_windows + grad: + - tests/core/test_sysid_mask_cuts_unconverged_windows.py::test_a_diverged_window_leaves_a_finite_gradient_of_the_other_windows + - tests/core/test_sysid_mask_reads_every_step_of_the_window.py::test_multiple_shooting_drops_the_recovering_window_and_keeps_the_converged_one vmap: narrowed multi_rate: tests/core/test_multirate_group_meta_follows_firing.py::test_the_sysid_mask_keeps_a_window_whose_applied_solve_converged sub_cycled: tests/core/test_sysid_mask_reads_the_last_waveform_sweep.py::test_a_window_whose_first_sweep_hit_its_cap_is_kept_by_the_mask @@ -5076,9 +5086,12 @@ claims: - tests/core/test_sysid_mask_refuses_unrecorded_group.py::test_every_group_that_records_a_verdict_is_masked - tests/core/test_sysid_degenerate_inputs.py::test_windowed_loss_refuses_a_mask_unconverged_that_is_not_a_bool - tests/core/test_multirate_group_meta_follows_firing.py::test_the_sysid_mask_keeps_a_window_whose_applied_solve_converged + - tests/core/test_sysid_mask_reads_every_step_of_the_window.py::test_the_window_recovers_its_last_step_converges_and_its_early_ones_do_not status: verified domains: - f32: tests/core/test_sysid_mask_refuses_unrecorded_group.py::test_every_group_that_records_a_verdict_is_masked + f32: + - tests/core/test_sysid_mask_refuses_unrecorded_group.py::test_every_group_that_records_a_verdict_is_masked + - tests/core/test_sysid_mask_reads_every_step_of_the_window.py::test_multiple_shooting_drops_the_recovering_window_and_keeps_the_converged_one f64: n/a mixed_dtype: n/a 16bit: n/a @@ -5623,8 +5636,8 @@ claims: source field is floating: only that group owns the `_meta` slot coupling__reading_floor (the floor per evaluation, NaN until a step has written it, restarted with the group's other report slots); - every other group's floor is taken from the returned state alone, as it - was. A state whose slot was never written (a checkpoint from before + every other group's floor is taken from the returned state alone, which + the graph keeps across a later state write (CPL-189). A state whose slot was never written (a checkpoint from before the slot existed) is read with the graph's own weights. oracle: >- residual_precision_floor of the returned state under the step's @@ -5657,3 +5670,66 @@ claims: adaptive: 'tests/core/test_coupling_mapped_edges_in_every_domain.py::test_the_reports_floor_is_the_one_the_step_measured_with_its_weights[adaptive]' checkpoint_restart: 'tests/core/test_coupling_mapped_edges_in_every_domain.py::test_the_reports_floor_is_the_one_the_step_measured_with_its_weights[checkpoint_restart]' sharded: 'tests/cloud/multigpu/test_coupling_domain_claims_on_a_sharded_graph.py::test_the_claim_holds_with_a_sharded_member[CPL-188]' + +- id: CPL-189 + area: report + claim: >- + "The floor is measured on the state the step returned, and stays so when + the state is written afterwards (set_node_state, PUT /graph/state, a node + replaced): the entry describes the step that ran until the group steps + again, the state is reset or a member is removed. A checkpoint is + a copy of the state, _meta included, so one saved after such a write to + a member holds the written state and not the returned one; it carries a + marker saying so, and the graph that loads it reports the group with + spectral_error_bound NaN, spectral_usable, gradient_bound_usable and + precision_limited False and a not_usable_reason, until the group + steps. The entry's other numbers are the slots' own." + sources: + - src/maddening/core/graph_manager.py::GraphManager.coupling_diagnostics (precision_limited) + - src/maddening/core/graph_manager.py::GraphManager._keep_state_for_reports + - src/maddening/core/graph_manager.py::GraphManager._groups_written_after_their_step + - src/maddening/core/graph_manager.py::GraphManager._loaded_after_a_write + - src/maddening/core/simulation/checkpoint.py::_load_from_archive + - src/maddening/core/simulation/checkpoint.py::save_state + conditions: >- + as stated: writes through set_node_state (which the REST state route and + load_state use), add_node and remove_node. An assignment into the + private gm._state is not seen. A group that stores its floor in the + step (CPL-188) does not read the kept state and is never marked. The + marker is the archive member _reports//written_after_step, + present only for such a group; an archive without it loads as before, + and a marker for a group the graph does not have is ignored. + oracle: the report read before the write, key for key. + tests: + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_write_after_the_step_does_not_move_its_report + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_the_report_of_every_stepping_entry_point_survives_a_write + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_the_next_step_is_reported_from_the_state_it_left + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_loaded_checkpoint_is_reported_as_the_step_it_saved + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_checkpoint_is_a_copy_of_the_state_and_says_when_it_was_written_after_the_step + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_checkpoint_keeps_the_report_where_no_member_changed + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_an_archive_without_the_marker_loads_as_it_always_did + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_node_cannot_take_the_markers_prefix + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_failed_load_leaves_the_report_of_the_step_before_it + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_recompile_keeps_the_report_and_what_it_was_measured_on + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_reset_state_and_a_removed_member_still_end_the_report + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_transform_that_stepped_the_graph_leaves_the_report_of_the_step_before_it + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_traced_write_keeps_no_tracer_for_the_report + - tests/core/test_coupling_report_describes_the_step_that_ran.py::test_the_profilers_measurements_hand_the_report_back_with_the_state + - tests/api/test_a_coupling_report_survives_a_rest_state_write.py::test_a_state_put_after_the_step_does_not_move_its_report + - tests/api/test_a_coupling_report_survives_a_rest_state_write.py::test_a_checkpoint_saved_over_rest_after_a_put_reloads_the_state_and_withholds_the_bound + - tests/property/test_coupling_targeted_search.py::test_a_report_on_a_seed_shape_does_not_move_when_the_state_is_written_afterwards + status: verified + domains: + f32: tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_write_after_the_step_does_not_move_its_report + f64: 'tests/property/test_coupling_targeted_search.py::test_a_report_on_a_seed_shape_does_not_move_when_the_state_is_written_afterwards[gain-0.98-in-other-units]' + mixed_dtype: n/a + 16bit: n/a + jit: 'tests/core/test_coupling_report_describes_the_step_that_ran.py::test_the_report_of_every_stepping_entry_point_survives_a_write[run_scan]' + grad: n/a + vmap: n/a + multi_rate: n/a + sub_cycled: n/a + predictors_warm_starts: n/a + adaptive: 'tests/core/test_coupling_report_describes_the_step_that_ran.py::test_the_report_of_every_stepping_entry_point_survives_a_write[run_adaptive]' + checkpoint_restart: tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_checkpoint_is_a_copy_of_the_state_and_says_when_it_was_written_after_the_step + sharded: n/a diff --git a/docs/validation/known_anomalies.yaml b/docs/validation/known_anomalies.yaml index 2a57916b0..0d9cc664e 100644 --- a/docs/validation/known_anomalies.yaml +++ b/docs/validation/known_anomalies.yaml @@ -15533,3 +15533,106 @@ anomalies: - "tests/property/test_rest_requests_generated_from_the_schema.py::test_every_integer_query_declares_its_range" - "tests/api/test_profile_endpoints.py::TestSimProfile::test_a_count_outside_its_range_is_a_422_and_nothing_is_profiled" github_issue: null + + - anomaly_id: "MADD-ANO-231" + title: "coupling_diagnostics() measured the residual's float floor on the live state: after set_node_state the same step read another spectral_error_bound (0.0 with spectral_usable=True where every member was written to zero)" + description: > + The `_meta` slots of a report are the step's own record, but the + float floor that `spectral_error_bound` adds, `precision_limited` + reads and `spectral_usable` rests on was computed at report time + from the node states the graph held then. A state write after the + step (`set_node_state`, `PUT /graph/state/{node}`) therefore gave + the same step another report. Measured before the fix (CPU, + jaxlib 0.11.0): float32 Gauss-Seidel pairs with gains 0.999 and + 0.9999 at tolerance 1e-9 stall at `residual == 0.0`, where the + floor is the whole bound; with both members written to zero (a + field at exactly zero leaves the norm) the report read + `spectral_error_bound = 0.0`, `spectral_usable=True`, + `precision_limited=False` for a state 9.0e-5 and 7.6e-4 from its + fixed point. Writing one member to zero shrank the bound 1.37 + times; at tolerance 1e-6 the bound shrank 12.7 times and stayed + above the distance. Only a group whose interface norm reads a + mapped edge stored its floor in the step. No release carries the + floor or the spectral bound. + severity: "major" + safety_relevance: "context_dependent" + safety_relevance_rationale: > + A bound of 0.0 flagged usable for a state that is not at its fixed + point is a silent wrong result, never `minor`. Not `critical`: it + is reported, not applied, and it needs a state write between the + step and the read. `context_dependent` as for the other bound + entries. + affected_components: + - "maddening.core.graph_manager.GraphManager.coupling_diagnostics" + - "maddening.core.graph_manager.GraphManager.set_node_state" + - "maddening.core.simulation.checkpoint.save_state" + affected_versions: "none" + workaround: > + None needed in 0.4.0. + resolution_status: "resolved" + resolution_version: "0.4.0" + residual_risk: > + The graph keeps a shallow copy of the state (references, no data) + from the first write after a step, beside the report slots it goes + with, and the report is measured on it while those slots are still + the ones in the graph; a step drops it. The doors that keep it are + `set_node_state` (so `PUT /graph/state/{node}` and `load_state`), + `add_node` and `remove_node`. An assignment into the private + `gm._state` dict is not seen. A checkpoint is a copy of the + state, `_meta` included, and holds the written state, not the one + the step returned. One saved after a member was written to other + values carries the member `_reports//written_after_step`, + and the graph that loads it reports that group's + `spectral_error_bound` as NaN, `spectral_usable`, + `gradient_bound_usable` and `precision_limited` as False, with a + `not_usable_reason`, until the group steps: the bound cannot be + recovered from such a checkpoint. An archive without the member + loads as before. `_reports` joins the node names the graph + refuses. The replaced arrays stay referenced until the next step. + verification: + - "tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_write_after_the_step_does_not_move_its_report" + - "tests/core/test_coupling_report_describes_the_step_that_ran.py::test_the_report_of_every_stepping_entry_point_survives_a_write" + - "tests/core/test_coupling_report_describes_the_step_that_ran.py::test_a_checkpoint_is_a_copy_of_the_state_and_says_when_it_was_written_after_the_step" + - "tests/property/test_coupling_targeted_search.py::test_a_report_on_a_seed_shape_does_not_move_when_the_state_is_written_afterwards" + - "tests/api/test_a_coupling_report_survives_a_rest_state_write.py::test_a_state_put_after_the_step_does_not_move_its_report" + github_issue: null + + - anomaly_id: "MADD-ANO-232" + title: "A gradient_relative_error_bound whose flag is False can read NaN on a GPU backend where the CPU backend reads a number" + description: > + One local GPU run (jax 0.11.0, float32, + `benchmarks/results/gpu_eigvals_probe/RESULT.md`): the stock spring + pair under `solver="ift"`, `diagnostics=True`, after 40 steps with + the residual exactly 0, reported + `gradient_relative_error_bound` 4.3e-6 on CPU and NaN on the GPU, + `gradient_bound_usable=False` on both (the pair is at its float + floor and stock nodes do not declare their evaluation count). The + other reported numbers agreed to rounding. Inside the contract: + the flag says the number is not to be used, and NaN is a documented + reading of a bound that was not computed. The cause is not + established. Read on CPU: every quotient in the bound is guarded, + and the eight-vector range basis captures this pair's Jacobian + (four state entries) with or without one-ulp noise on its + products, so the documented NaN paths (nothing responds, range not + captured, a non-finite tangent) are not reached there; forcing + "not captured" on CPU gives the GPU's reading. + severity: "minor" + safety_relevance: "not_safety_relevant" + safety_relevance_rationale: > + Not a wrong result: the value is flagged not usable on both + backends, and a NaN cannot be taken for a bound. + affected_components: + - "maddening.core.graph_manager.GraphManager.coupling_diagnostics" + affected_versions: ">=0.4.0.dev0" + workaround: > + Read `gradient_bound_usable` before the value, as documented; do + not compare a not-usable value between backends. + resolution_status: "open" + resolution_version: null + residual_risk: > + A usable bound has not been seen to differ in kind between + backends: in float64, where the flags were True, every reported + number equalled the CPU's to rounding. Finding the intermediate + needs a GPU run with the bound's parts printed. + verification: [] + github_issue: null diff --git a/docs/validation/soup_package.md b/docs/validation/soup_package.md index 225a57ceb..5c54387c4 100644 --- a/docs/validation/soup_package.md +++ b/docs/validation/soup_package.md @@ -293,8 +293,10 @@ stale copy fails CI rather than shipping. | MADD-ANO-230 | A direction the Arnoldi breakdown test takes for rounding can carry the dominant mode of a Jacobian far from normal: float32 fan-out hubs with a field 1e-4 to 1e-2 of the rest read rho_spectral 0.157 for 0.349 and 0.503 for 0.397, settled | `major` | `context_dependent` | `open` | >=0.4.0.dev0 | | MADD-ANO-227 | Nodes whose timesteps are more than about 1e9 apart, or have no common step the float GCD finds, ran on different clocks with no error | `major` | `context_dependent` | `resolved` (in 0.4.0) | >=0.1.0, <0.4.0 | | MADD-ANO-228 | REST read a boolean or numeric text as a timestep (`true` added a node stepping at 1.0 s), and clamped the counts of POST /sim/profile without saying so | `major` | `context_dependent` | `resolved` (in 0.4.0) | >=0.1.0, <0.4.0 | +| MADD-ANO-231 | coupling_diagnostics() measured the residual's float floor on the live state: after set_node_state the same step read another spectral_error_bound (0.0 with spectral_usable=True where every member was written to zero) | `major` | `context_dependent` | `resolved` (in 0.4.0) | none | +| MADD-ANO-232 | A gradient_relative_error_bound whose flag is False can read NaN on a GPU backend where the CPU backend reads a number | `minor` | `not_safety_relevant` | `open` | >=0.4.0.dev0 | -*230 anomalies registered. 37 have a defect reachable in this version — every entry whose `resolution_status` is not `resolved` or `duplicate`, which is 28 `open` plus 9 `partially_resolved` whose residual risk is still live. The Affected Versions column is a PEP 440 specifier set read against this document's version; `none` marks a defect introduced and fixed within one development cycle, which no release carried. The convention, and the gate that holds every range to it, are in the header of `known_anomalies.yaml`. Rationale, workaround, affected components and verification evidence for each: `known_anomalies.yaml`.* +*232 anomalies registered. 38 have a defect reachable in this version — every entry whose `resolution_status` is not `resolved` or `duplicate`, which is 29 `open` plus 9 `partially_resolved` whose residual risk is still live. The Affected Versions column is a PEP 440 specifier set read against this document's version; `none` marks a defect introduced and fixed within one development cycle, which no release carried. The convention, and the gate that holds every range to it, are in the header of `known_anomalies.yaml`. Rationale, workaround, affected components and verification evidence for each: `known_anomalies.yaml`.* ## 4. Verification Evidence diff --git a/src/maddening/core/_graph_specs.py b/src/maddening/core/_graph_specs.py index c10837161..53d15c5e7 100644 --- a/src/maddening/core/_graph_specs.py +++ b/src/maddening/core/_graph_specs.py @@ -565,7 +565,7 @@ def _holds_every_value(have, dtype) -> bool: #: State and checkpoint keys a node may not be named: the graph's own state #: (``_meta``) and a checkpoint's params prefixes #: (``maddening.core.simulation.checkpoint``). -_RESERVED_STATE_KEYS = frozenset({_META_KEY, "_params", "_params_mappings"}) +_RESERVED_STATE_KEYS = frozenset({_META_KEY, "_params", "_params_mappings", "_reports"}) def _uncarriable_characters(text: str) -> list[str]: @@ -639,7 +639,8 @@ def _node_name_refusal(name: Any) -> Optional[str]: if name in _RESERVED_STATE_KEYS: # The graph's own state lives under ``_meta`` (coupling and # multirate carries) and a checkpoint keeps the params under - # ``_params`` and ``_params_mappings``. A node named for one of + # ``_params`` and ``_params_mappings`` (and a per-group marker + # under ``_reports``). A node named for one of # them was taken (POST /graph/nodes answered 201), the next # compile dropped its state, every step was a KeyError and a # checkpoint save was refused until the node was deleted. diff --git a/src/maddening/core/coupling/_group_layout.py b/src/maddening/core/coupling/_group_layout.py index 019ca9b12..2bc39e34e 100644 --- a/src/maddening/core/coupling/_group_layout.py +++ b/src/maddening/core/coupling/_group_layout.py @@ -307,6 +307,15 @@ def _group_geometry_edges(group, edges) -> list: ) +_WRITTEN_BEFORE_SAVE_REASON = ( + "this report was loaded from a checkpoint saved after the group's state had been " + "written (set_node_state) since its last step; the state that step returned, which " + "the float floor is measured on, is not in the checkpoint, so spectral_error_bound, " + "precision_limited and the *_usable flags are not reported until the group steps. " + "iterations, residual, converged and the estimates are the step's own." +) + + def _geometry_edge_coupling_errors(group, edges) -> list[str]: """``ERROR:`` issues for a group setting a geometry-dependent mapping cannot serve (experimental; empty for every other group). diff --git a/src/maddening/core/graph_manager.py b/src/maddening/core/graph_manager.py index bf83d2a1e..11d0e5a8e 100644 --- a/src/maddening/core/graph_manager.py +++ b/src/maddening/core/graph_manager.py @@ -260,6 +260,12 @@ def __init__(self) -> None: # Escaped-tracer bookkeeping; see ``_recover_from_escaped_tracers``. self._state_traced = False self._state_before_trace: Optional[dict] = None + # What the last step left, kept from the first write made to the + # state after it; see ``_keep_state_for_reports``. + self._state_as_reported: Optional[tuple[dict, dict]] = None + # Groups whose report was loaded from a checkpoint saved after + # their state was written; see ``_note_loaded_after_a_write``. + self._reports_loaded_after_a_write: dict[str, tuple] = {} # The rate dividers of the step that is actually compiled, as # opposed to ``_rate_dividers``, which ``compile`` overwrites on # its way through and leaves behind if it raises. This is what @@ -1664,6 +1670,7 @@ def add_node(self, node: SimulationNode) -> None: # must leave the graph exactly as it was. state = node.initial_state() self._nodes[node.name] = spec + self._keep_state_for_reports() self._state[node.name] = state self._dirty = True self._notify(EVENT_NODE_ADDED, node.name) @@ -2044,6 +2051,7 @@ def _remove_node(self, name: str, *, replacing: bool) -> list[str]: # ``pop`` rather than ``del``: a graph whose state entry is missing # must still be removable, so the removal cannot itself fail # half-way and leave ``_nodes`` and ``_state`` disagreeing. + self._keep_state_for_reports() self._state.pop(name, None) self._edges = [ e for e in self._edges @@ -3879,6 +3887,10 @@ def _store_state(self, new_state: dict) -> None: else: self._state_traced = False self._state_before_trace = None + # The state being stored is what its own step left. (A traced + # one is put back to the state before it, whose reports the + # kept copy still describes.) + self._state_as_reported = None self._state = new_state if self._underflow_check_pending and not self._state_traced: self._underflow_check_pending = False @@ -3966,6 +3978,142 @@ def _recover_from_escaped_tracers(self, *, warn: bool = True) -> None: stacklevel=3, ) + # ------------------------------------------------------------------ + # A report describes the step that ran + # ------------------------------------------------------------------ + + def _keep_state_for_reports(self) -> None: + """Keep what the last step left, before a write that is not a step's. + + ``coupling_diagnostics()`` measures the residual's float floor on + the state the step returned. The ``_meta`` slots are the step's + own record, but the node states are live, so the same step's + ``spectral_error_bound``, ``precision_limited`` and + ``spectral_usable`` moved when the state was written afterwards: + a float32 pair stalled at ``residual == 0.0`` read a bound of + ``0.0``, ``spectral_usable=True``, after ``set_node_state`` to + zeros, 9e-5 from its fixed point. + + Called by every door that writes node states in place + (``set_node_state`` -- so ``PUT /graph/state`` and + ``load_state`` -- ``add_node``, ``remove_node``) *before* its + write. The first such write after a step keeps a shallow copy of + the state (references to the arrays, no data) beside the report + slots it goes with; later ones find it there. A step drops it + (``_store_state``), so it holds the replaced arrays only until + the next step. The reader (``_state_a_report_describes``) uses + it only while the group's slots are still the very objects kept + here, so a door that brings its own slots -- a step, a loaded + checkpoint, ``reset_state``, a recompile that restarts a + replaced group -- needs no bookkeeping: the report then reads + the live state, which is what came with those slots. + + Nothing is kept while the state holds tracers: the graph is put + back to the state before the trace, which the copy already kept + (or the live state) describes. + """ + if self._state_traced or not self._committed_floor_inputs: + return + meta = self._state.get(_graph_specs._META_KEY) + if not meta: + return + kept = self._state_as_reported + if (kept is not None and len(kept[1]) == len(meta) + and all(kept[1].get(slot) is value for slot, value in meta.items())): + return + self._state_as_reported = ( + {name: fields for name, fields in self._state.items() + if name != _graph_specs._META_KEY}, + dict(meta), + ) + + def _state_a_report_describes(self, slots: Sequence[str], nodes) -> dict: + """The node states the step that wrote ``slots`` left. + + The copy kept before the first later write + (:meth:`_keep_state_for_reports`) while ``slots`` are still the + objects it was kept with; otherwise the live state. + """ + kept = self._state_as_reported + if kept is None: + return self._state + state, kept_meta = kept + meta = self._state.get(_graph_specs._META_KEY, {}) + if (all(slot in meta and kept_meta.get(slot) is meta[slot] for slot in slots) + and all(name in state for name in nodes)): + return state + return self._state + + def _groups_written_after_their_step(self) -> list[str]: + """The groups a checkpoint marks "state written after the last step". + + A checkpoint is a copy of the state, ``_meta`` included, and holds + the *written* state: what the step returned, which the report's + float floor is measured on, is not in it. So the archive carries + the fact beside the state (``checkpoint.save_state``), and the + graph that loads it reports the group's floor-dependent entries + as not usable (:meth:`_loaded_after_a_write`). + + A group is marked where its report is measured on a kept copy + because a member has since been written to other values, and + where this graph itself loaded it marked and has not stepped it + since. Not a group that stores its floor in the step (CPL-188). + """ + meta = self._state.get(_graph_specs._META_KEY, {}) + if not meta: + return [] + marked = [] + for key in self._committed_floor_inputs: + nodes = key.split("+") + slots = (f"coupling_{key}_iterations", f"coupling_{key}_residual") + if any(slot not in meta for slot in slots) or int(meta[slots[0]]) <= 0: + continue # no report to qualify + stored = meta.get(f"coupling_{key}_reading_floor") + if stored is not None and np.isfinite(np.asarray(stored)): + continue # the step stored this group's floor itself + if self._loaded_after_a_write(key): + marked.append(key) + continue + state = self._state_a_report_describes(slots, nodes) + if state is self._state: + continue + for name in nodes: + live = self._state.get(name) + if live is None or state[name] is live: + continue + if (set(live) != set(state[name]) or any( + np.asarray(live[f]).tobytes() != np.asarray(state[name][f]).tobytes() + or np.asarray(live[f]).dtype != np.asarray(state[name][f]).dtype + for f in live)): + marked.append(key) + break + return marked + + def _note_loaded_after_a_write(self, keys) -> None: + """Record that a checkpoint just loaded marks *keys* as written + after their last step (``checkpoint.load_state``, on success). + + Kept beside the report slots the load installed: the record holds + while they are the ones in the graph, so the group's next step, a + reset or another load ends it with no bookkeeping. + """ + meta = self._state.get(_graph_specs._META_KEY, {}) + self._reports_loaded_after_a_write = { + key: (meta.get(f"coupling_{key}_iterations"), meta.get(f"coupling_{key}_residual")) + for key in keys if key in self._committed_floor_inputs + } + + def _loaded_after_a_write(self, key: str) -> bool: + """Whether *key*'s report came from a checkpoint saved after its + members were written (see :meth:`_note_loaded_after_a_write`).""" + noted = self._reports_loaded_after_a_write.get(key) + if noted is None: + return False + meta = self._state.get(_graph_specs._META_KEY, {}) + return (noted[0] is not None + and meta.get(f"coupling_{key}_iterations") is noted[0] + and meta.get(f"coupling_{key}_residual") is noted[1]) + # ------------------------------------------------------------------ # Internal helpers for _meta stripping # ------------------------------------------------------------------ @@ -4582,6 +4730,19 @@ def coupling_diagnostics(self) -> dict[str, dict]: count declared the bound with its floor *is* a bound, and a group converged to float32 -- the best a float32 group can do -- stays usable. + The floor is measured on the state the step returned, + and stays so when the state is written afterwards + (``set_node_state``, ``PUT /graph/state``, a node replaced): + the entry describes the step that ran until the group steps + again, the state is reset or a member is removed. A + checkpoint is a copy of the state, ``_meta`` included, so + one saved after such a write to a member holds the written + state and not the returned one; it carries a marker saying + so, and the graph that loads it reports the group with + ``"spectral_error_bound"`` NaN, ``"spectral_usable"``, + ``"gradient_bound_usable"`` and ``"precision_limited"`` + ``False`` and a ``"not_usable_reason"``, until the group + steps. The entry's other numbers are the slots' own. ``converged=True`` is a statement about the state this step returned: both solvers stop on the iterate whose residual @@ -4813,8 +4974,11 @@ def coupling_diagnostics(self) -> dict[str, dict]: unit = np.asarray(measured_floor) floor = float(unit * np.asarray(evaluations, unit.dtype)) else: + # On the state the step left, not on whatever has been + # written to it since (``_keep_state_for_reports``). floor = float(residual_precision_floor( - self._state, sorted(group.nodes), group.convergence_norm, + self._state_a_report_describes((iter_key, res_key), group.nodes), + sorted(group.nodes), group.convergence_norm, group.atol, group.rtol, list(internal_edges), evaluations=evaluations, mappings=(self._params or {}).get("mappings"), @@ -4862,6 +5026,21 @@ def coupling_diagnostics(self) -> dict[str, dict]: ), "precision_limited": precision_limited, }) + if self._loaded_after_a_write(key) and not ( + measured_floor is not None + and np.isfinite(np.asarray(measured_floor))): + # The checkpoint this report was loaded from was saved + # after the group's state had been written: the state + # the step returned, which the floor is measured on, + # is not in it. Everything built on the floor is + # withheld; the slots' own numbers stand. + result[key].update({ + "spectral_error_bound": float("nan"), + "spectral_usable": False, + "gradient_bound_usable": False, + "precision_limited": False, + "not_usable_reason": _group_layout._WRITTEN_BEFORE_SAVE_REASON, + }) geometry_keys = self._committed_geometry_edges.get(key, ()) if geometry_keys: # Experimental, 0.4.0: the diagnostics do not read a @@ -6115,6 +6294,7 @@ def set_node_state(self, name: str, state: dict) -> None: if name not in self._nodes: raise KeyError(f"No node named '{name}'.") state = _graph_specs._strong_typed(state) + self._keep_state_for_reports() if _graph_specs._holds_tracer({name: state}) and not self._state_traced: # Snapshot before the write, per node, so the recovery has # something untraced to go back to. Only on the tracer path, @@ -7031,6 +7211,9 @@ def coupling_report(self) -> InspectionTable: * in place of the three above, ``not_usable_reason`` for a group that resolves a geometry-dependent mapping (experimental): its bounds, estimates and ``*_usable`` flags are withheld; + * ``not_usable_reason`` for a group loaded from a checkpoint saved + after its state was written: the bound and the flags that rest + on the float floor are withheld; * why a group has no report (``solver="fori"`` without ``diagnostics``, no step since ``compile()`` / ``reset_state()``, added since the last compile). diff --git a/src/maddening/core/inspection.py b/src/maddening/core/inspection.py index 1b73d1e1c..a2c96220a 100644 --- a/src/maddening/core/inspection.py +++ b/src/maddening/core/inspection.py @@ -1344,6 +1344,17 @@ def _coupling_flags(group: Any, d: Mapping[str, Any], whole: tuple = ()) -> list + ("; under solver='ift' the gradient through this step is unreliable" if group.solver == "ift" else "")) reason = d.get("not_usable_reason") + estimate = d.get("error_estimate") + if reason and isinstance(estimate, float) and not math.isnan(estimate): + # Only what rests on the float floor is withheld (a checkpoint + # saved after the state was written): the estimates are there, + # and their caveats below still apply. + flags.append(f"no bound reported: {reason}") + if not d["ratio_usable"]: + flags.append("ratio_usable=False: the contraction ratio was unusable, so the " + "criterion fell back to the raw residual test; converged reports that " + "test, not a distance estimate (MADD-ANO-005)") + return flags if reason: # The report withholds every bound, estimate and ``*_usable`` flag # of this group, so the caveats below (which read them) would be diff --git a/src/maddening/core/simulation/checkpoint.py b/src/maddening/core/simulation/checkpoint.py index e48a0bcc3..71759f7d2 100644 --- a/src/maddening/core/simulation/checkpoint.py +++ b/src/maddening/core/simulation/checkpoint.py @@ -43,6 +43,11 @@ #: Python identifier, so it can never be taken for a weight: a weight's #: name is one (``register_mapping``'s contract for ``params_pytree()``). _STRUCTURE_MEMBER = "structure.sha256" +#: ``_reports//written_after_step``: present (a uint8 one) for a +#: coupling group whose state was written after its last step. Optional; +#: an archive without it loads as it always did. +_REPORTS_KEY = "_reports" +_WRITTEN_MEMBER = "written_after_step" _DIGEST_BYTES = hashlib.sha256().digest_size # Schema version for the integrity manifest. @@ -96,6 +101,17 @@ def save_state(graph_manager: "GraphManager", path: str | Path) -> Path: ------- Path The resolved path of the written file (always ends in ``.npz``). + + Notes + ----- + The state and ``_meta`` are saved as they are. Where a coupling + group's members were written to other values after its last step + (``set_node_state``), the archive also carries + ``_reports//written_after_step``: the state that step returned, + which the report's float floor is measured on, is not in the archive, + and the graph that loads it reports that group's bound and + ``*_usable`` flags as not usable, with a ``not_usable_reason``, + until the group steps. """ path = Path(path) @@ -115,6 +131,15 @@ def save_state(graph_manager: "GraphManager", path: str | Path) -> Path: for field_name, value in raw_state[_META_KEY].items(): key = f"{_META_KEY}/{field_name}" arrays[key] = np.asarray(value) + # The state and ``_meta`` are saved as they are. What the archive + # cannot hold is the state a coupling group's last step returned + # where a member has been written since, and the report's float + # floor is measured on that: the fact goes beside the state, one + # member per such group, and the load reports the group's + # floor-dependent entries as not usable. No member where no + # group is in that case, so every other archive is as it was. + for group_key in graph_manager._groups_written_after_their_step(): # noqa: SLF001 + arrays[f"{_REPORTS_KEY}/{group_key}/{_WRITTEN_MEMBER}"] = np.ones((), np.uint8) # Differentiable graph parameters (node constants), so a calibrated # graph restores with the values it was calibrated to. @@ -321,6 +346,7 @@ def read(key: str, want: Any, what: str) -> np.ndarray: node_keys: dict[str, dict[str, str]] = {} param_keys: dict[str, dict[str, str]] = {} mapping_keys: dict[str, dict[str, str]] = {} + written_groups: list[str] = [] for flat_key in archive: parts = flat_key.split("/", 1) @@ -338,6 +364,12 @@ def read(key: str, want: Any, what: str) -> np.ndarray: elif prefix == _MAPPINGS_KEY: edge_key, wname = field.rsplit("/", 1) mapping_keys.setdefault(edge_key, {})[wname] = flat_key + elif prefix == _REPORTS_KEY: + # Read by name only: presence is the content. A group this + # graph does not have is ignored, as its ``_meta`` slots are. + group_key, member = field.rsplit("/", 1) if "/" in field else (field, "") + if member == _WRITTEN_MEMBER: + written_groups.append(group_key) else: node_keys.setdefault(prefix, {})[field] = flat_key @@ -490,6 +522,9 @@ def _stage_params(section: str, saved_tree: dict) -> list: except BaseException: _restore_state_and_params(graph_manager, undo) raise + # Only once the load has succeeded: the reports it installed, and + # which of them describe a state written after its step. + graph_manager._note_loaded_after_a_write(written_groups) # noqa: SLF001 class CheckpointFormatError(ValueError): diff --git a/src/maddening/core/simulation/profiler.py b/src/maddening/core/simulation/profiler.py index f93e65b97..2441296fe 100644 --- a/src/maddening/core/simulation/profiler.py +++ b/src/maddening/core/simulation/profiler.py @@ -536,6 +536,11 @@ def compile_counts( if scan_steps > 0: counts.scan_steps = int(scan_steps) saved_state = jax.tree.map(lambda x: x, gm._state) + # With the state goes what its last step left where the state has + # been written since (``GraphManager._keep_state_for_reports``): + # the run below drops it, and the caller's report would then be + # measured on the written state. + saved_as_reported = gm._state_as_reported try: # Populates ``_scan_cache``; ``scan_trace_count`` counts the # Python traces, one per XLA compile of a scan program. It @@ -555,6 +560,7 @@ def compile_counts( ) finally: gm._state = saved_state + gm._state_as_reported = saved_as_reported return counts @@ -659,6 +665,7 @@ def _one_iteration_variant(gm): def _cm(): saved_groups = list(gm._coupling_groups) saved_state = jax.tree.map(lambda x: x, gm._state) + saved_as_reported = gm._state_as_reported saved_params = gm.params saved_step = gm._compiled_step try: @@ -693,6 +700,7 @@ def _cm(): with quiet_warnings(): gm.compile() gm._state = saved_state + gm._state_as_reported = saved_as_reported gm.params = saved_params # ``compile`` rebuilt the step; the original object is fine # to keep for callers holding a reference (same graph). diff --git a/tests/api/test_a_coupling_report_survives_a_rest_state_write.py b/tests/api/test_a_coupling_report_survives_a_rest_state_write.py new file mode 100644 index 000000000..4a4a41b06 --- /dev/null +++ b/tests/api/test_a_coupling_report_survives_a_rest_state_write.py @@ -0,0 +1,110 @@ +"""A coupling report describes the step that ran, through the REST write doors too. + +``PUT /graph/state/{node}`` writes through ``set_node_state`` and +``POST /checkpoint/save`` / ``load`` through ``save_state`` / +``load_state``; the report is read in process (no route returns it). See +``tests/core/test_coupling_report_describes_the_step_that_ran.py`` for the +invariant and the in-process doors. +""" + +from __future__ import annotations + +import os + +os.environ.setdefault("JAX_PLATFORMS", "cpu") + +import jax.numpy as jnp +import pytest + +from maddening.api.server import SimulationServer +from maddening.core.graph_manager import GraphManager +from maddening.core.node import BoundaryInputSpec, SimulationNode +from tests._loopback_client import LoopbackTestClient as TestClient + +N = 3 +KEY = "a+b" + + +class Lin(SimulationNode): + def __init__(self, name, b): + super().__init__(name, 0.01, g=jnp.float32(0.99), b=jnp.full(N, b, jnp.float32)) + + def initial_state(self): + return {"x": jnp.zeros(N, jnp.float32)} + + def boundary_input_spec(self): + return {"u": BoundaryInputSpec(shape=(N,), dtype=jnp.float32, + default=jnp.zeros(N, jnp.float32))} + + def update(self, state, boundary_inputs, dt, *, params=None): + p = self.params if params is None else {**self.params, **params} + return {"x": p["g"] * boundary_inputs["u"] + p["b"]} + + def update_evaluations(self): + return 1 + + +@pytest.fixture(scope="module") +def served(tmp_path_factory): + gm = GraphManager() + gm.add_node(Lin("a", 1.0)) + gm.add_node(Lin("b", 2.0)) + gm.add_edge("a", "b", "x", "u") + gm.add_edge("b", "a", "x", "u") + gm.add_coupling_group(["a", "b"], max_iterations=400000, tolerance=1e-9, + diagnostics=True) + gm.compile() + server = SimulationServer({"Lin": Lin}, graph_manager=gm, + checkpoint_root=str(tmp_path_factory.mktemp("checkpoints"))) + return gm, TestClient(server.create_app(), raise_server_exceptions=False) + + +def _report(gm): + d = gm.coupling_diagnostics().get(KEY) + if d is None: + return None + return {k: ("nan" if isinstance(v, float) and v != v else v) for k, v in d.items()} + + +def _zero(client): + for name in ("a", "b"): + reply = client.put(f"/graph/state/{name}", json={"state": {"x": [0.0] * N}}) + assert reply.status_code == 200, reply.text + + +def test_a_state_put_after_the_step_does_not_move_its_report(served): + gm, client = served + assert client.post("/sim/reset").status_code == 200 + assert client.post("/sim/step").status_code == 200 + first = _report(gm) + assert first["spectral_usable"] and first["precision_limited"], first + assert first["spectral_error_bound"] > 0.0 + _zero(client) + assert client.get("/graph/state/a").json()["x"] == [0.0] * N + assert _report(gm) == first + + +def test_a_checkpoint_saved_over_rest_after_a_put_reloads_the_state_and_withholds_the_bound(served): + gm, client = served + assert client.post("/sim/reset").status_code == 200 + assert client.post("/sim/step").status_code == 200 + first = _report(gm) + assert client.post("/checkpoint/save", params={"path": "stepped.npz"}).status_code == 200 + _zero(client) + assert client.post("/checkpoint/save", params={"path": "written.npz"}).status_code == 200 + assert _report(gm) == first + before = client.get("/graph/state").json() + + reply = client.post("/checkpoint/load", params={"path": "written.npz"}) + assert reply.status_code == 200, reply.text + # A checkpoint is a copy of the state, ``_meta`` included. + assert client.get("/graph/state").json() == before + loaded = _report(gm) + assert loaded["spectral_usable"] is False and loaded["precision_limited"] is False + assert "written" in loaded["not_usable_reason"] + assert loaded["iterations"] == first["iterations"] + reply = client.post("/checkpoint/load", params={"path": "stepped.npz"}) + assert reply.status_code == 200, reply.text + assert _report(gm) == first + _zero(client) + assert _report(gm) == first diff --git a/tests/api/test_new_nodes_need_a_usable_name_and_timestep.py b/tests/api/test_new_nodes_need_a_usable_name_and_timestep.py index 0b83c97ec..627b13983 100644 --- a/tests/api/test_new_nodes_need_a_usable_name_and_timestep.py +++ b/tests/api/test_new_nodes_need_a_usable_name_and_timestep.py @@ -31,7 +31,7 @@ from tests._loopback_client import LoopbackTestClient as TestClient REGISTRY = {"BallNode": BallNode} -RESERVED = ["_meta", "_params", "_params_mappings"] +RESERVED = ["_meta", "_params", "_params_mappings", "_reports"] def _served(*, nodes=True, root=None): diff --git a/tests/compliance/test_soup_evidence.py b/tests/compliance/test_soup_evidence.py index 92f21d245..da65005a7 100644 --- a/tests/compliance/test_soup_evidence.py +++ b/tests/compliance/test_soup_evidence.py @@ -97,7 +97,7 @@ def test_the_committed_soup_tables_match_a_fresh_generation(): # anomaly is closed -- resolved, or a duplicate of the entry that carries it # -- and stays in the registry; only a closed entry, or a number no commit # ever recorded, can be retired. The gate reads the last status from git. -_HIGHEST_ANOMALY_ID = 230 +_HIGHEST_ANOMALY_ID = 232 _RETIRED_ANOMALY_IDS: dict = {} _HIGHEST_BENCHMARK_ID = 16 diff --git a/tests/core/test_coupling_report_describes_the_step_that_ran.py b/tests/core/test_coupling_report_describes_the_step_that_ran.py new file mode 100644 index 000000000..ab529484e --- /dev/null +++ b/tests/core/test_coupling_report_describes_the_step_that_ran.py @@ -0,0 +1,470 @@ +"""``coupling_diagnostics()`` describes the step that ran, whatever is written afterwards. + +The residual's float floor -- which ``spectral_error_bound`` adds, +``precision_limited`` reads and ``spectral_usable`` rests on -- is measured +on the state the step returned. The report took that state from the live +graph at report time, so a ``set_node_state`` after the step gave the same +step another report: a float32 pair stalled at ``residual == 0.0`` read +``spectral_error_bound == 0.0`` with ``spectral_usable=True`` once both +members were written to zero (a field at exactly zero leaves the norm), +for a state 1e-5 from its fixed point. + +The graph now keeps what the step left from the first later write +(``GraphManager._keep_state_for_reports``). The invariant, for every +stepping entry point and every door that writes the state without +stepping: **the report read after the write equals the report read before +it**, key for key; and a door that brings its own report (a step, a loaded +checkpoint) is described by the state it brought. +""" + +import os + +os.environ.setdefault("JAX_PLATFORMS", "cpu") + +import jax +import jax.numpy as jnp +import numpy as np +import pytest + +from maddening.core.graph_manager import GraphManager +from maddening.core.node import BoundaryInputSpec, SimulationNode +from maddening.core.simulation.profiler import _one_iteration_variant, compile_counts + +N = 3 +KEY = "a+b" +GAIN = 0.99 +ZERO = {"x": jnp.zeros(N, jnp.float32)} +ADAPTIVE = dict(dt_initial=0.01, dt_min=0.01, dt_max=0.01, atol=1e3, rtol=1.0) + + +class Lin(SimulationNode): + """``x <- g * u + b``, one evaluation per pass, declared.""" + + def __init__(self, name, b): + super().__init__(name, 0.01, g=jnp.float32(GAIN), b=jnp.full(N, b, jnp.float32)) + + def initial_state(self): + return {"x": jnp.zeros(N, jnp.float32)} + + def boundary_input_spec(self): + return {"u": BoundaryInputSpec(shape=(N,), dtype=jnp.float32, + default=jnp.zeros(N, jnp.float32))} + + def update(self, state, boundary_inputs, dt, *, params=None): + p = self.params if params is None else {**self.params, **params} + return {"x": p["g"] * boundary_inputs["u"] + p["b"]} + + def update_evaluations(self): + return 1 + + +class Idle(SimulationNode): + """A node outside the group.""" + + def __init__(self, name): + super().__init__(name, 0.01) + + def initial_state(self): + return {"y": jnp.ones(N, jnp.float32)} + + def update(self, state, boundary_inputs, dt): + return state + + +def _pair(diagnostics=True): + """A float32 Gauss-Seidel pair that stalls at ``residual == 0.0`` under + the default tolerance, 1e-5 from its float64 fixed point: the floor is + the whole of its bound.""" + gm = GraphManager() + gm.add_node(Lin("a", 1.0)) + gm.add_node(Lin("b", 2.0)) + gm.add_node(Idle("c")) + gm.add_edge("a", "b", "x", "u") + gm.add_edge("b", "a", "x", "u") + gm.add_coupling_group(["a", "b"], max_iterations=400000, tolerance=1e-9, + diagnostics=diagnostics) + gm.compile() + return gm + + +def _report(gm): + d = gm.coupling_diagnostics().get(KEY) + if d is None: + return None + return {k: ("nan" if isinstance(v, float) and v != v else v) for k, v in d.items()} + + +def _true_distance(gm): + """The returned state's distance to the float64 fixed point, in the + group's norm (each field over its largest entry, root-sum-square).""" + g = float(np.float32(GAIN)) + a_star = (1.0 + g * 2.0) / (1.0 - g * g) + b_star = 2.0 + g * a_star + xa = np.asarray(gm.get_node_state("a")["x"], np.float64) + xb = np.asarray(gm.get_node_state("b")["x"], np.float64) + return float(np.sqrt(np.sum(((xa - a_star) / np.max(np.abs(xa))) ** 2) + + np.sum(((xb - b_star) / np.max(np.abs(xb))) ** 2))) + + +STEPPERS = { + "step": lambda gm: gm.step(), + "run": lambda gm: gm.run(2), + "run_scan": lambda gm: gm.run_scan(2), + "run_scan_with_history": lambda gm: gm.run_scan_with_history(2), + "run_adaptive": lambda gm: gm.run_adaptive(0.01, **ADAPTIVE), + "run_adaptive_scan": lambda gm: gm.run_adaptive_scan(0.01, max_steps=1, **ADAPTIVE), +} + + +def _scaled(gm, name, factor): + return {f: v * jnp.asarray(factor, v.dtype) for f, v in gm.get_node_state(name).items()} + + +#: Writes that are not a step, each as ``write(gm)``. +WRITES = { + "one member to zero": lambda gm: gm.set_node_state("b", ZERO), + "both members to zero": lambda gm: (gm.set_node_state("a", ZERO), + gm.set_node_state("b", ZERO)), + "both members a million times larger": lambda gm: ( + gm.set_node_state("a", _scaled(gm, "a", 1e6)), + gm.set_node_state("b", _scaled(gm, "b", 1e6))), + "a member to NaN": lambda gm: gm.set_node_state( + "a", {"x": jnp.full(N, jnp.nan, jnp.float32)}), + "a node outside the group": lambda gm: gm.set_node_state( + "c", {"y": jnp.zeros(N, jnp.float32)}), + "the same member twice": lambda gm: (gm.set_node_state("a", _scaled(gm, "a", 2.0)), + gm.set_node_state("a", ZERO)), +} + + +@pytest.fixture(scope="module") +def stepped(): + """One stepped pair and its report, for the tests that only write: + ``(gm, report, returned state)``. No test steps this graph.""" + gm = _pair() + gm.step() + return gm, _report(gm), {n: gm.get_node_state(n) for n in ("a", "b", "c")} + + +@pytest.fixture(scope="module") +def _second(): + return _pair() + + +@pytest.fixture +def graph(_second): + """A compiled pair for the tests that step, reset before each.""" + _second.reset_state() + return _second + + +def _restore(gm, returned): + for name, fields in returned.items(): + gm.set_node_state(name, fields) + + +def test_the_fixture_rests_its_bound_on_the_floor(stepped): + """The case can express the defect: the floor is the whole bound, and + the bound covers the true distance.""" + gm, first, returned = stepped + _restore(gm, returned) + assert first["residual"] == 0.0 and first["precision_limited"], first + assert first["spectral_usable"] and first["spectral_error_bound"] > 0.0, first + true = _true_distance(gm) + assert 0.0 < true <= first["spectral_error_bound"], (true, first) + + +@pytest.mark.parametrize("write", sorted(WRITES)) +def test_a_write_after_the_step_does_not_move_its_report(stepped, write): + gm, first, returned = stepped + _restore(gm, returned) + WRITES[write](gm) + assert _report(gm) == first + # ...and the state really was written: the report is not a refusal to write. + _restore(gm, returned) + assert _report(gm) == first + + +@pytest.mark.parametrize("stepper", sorted(STEPPERS)) +def test_the_report_of_every_stepping_entry_point_survives_a_write(graph, stepper): + STEPPERS[stepper](graph) + first = _report(graph) + assert first is not None and first["spectral_usable"], first + WRITES["both members to zero"](graph) + assert _report(graph) == first + + +def test_the_next_step_is_reported_from_the_state_it_left(graph): + """A step after the write brings its own report, and nothing is held + back from the step before it.""" + graph.step() + first = _report(graph) + graph.set_node_state("a", _scaled(graph, "a", 0.5)) + assert graph._state_as_reported is not None + graph.step() + assert graph._state_as_reported is None + second = _report(graph) + assert second["iterations"] != first["iterations"], (first, second) + true = _true_distance(graph) + assert second["spectral_usable"] and second["spectral_error_bound"] >= true > 0.0 + # Measured on the state this step left: writing the same values back + # (which keeps that state) changes nothing. + for name in ("a", "b"): + graph.set_node_state(name, graph.get_node_state(name)) + assert _report(graph) == second + + +def test_a_loaded_checkpoint_is_reported_as_the_step_it_saved(stepped, graph, tmp_path): + """The decision for ``load_state``: the slots a checkpoint carries + describe the step that left the state beside them, so the report read + right after a load is the saved step's -- in a graph that has not + stepped and over one that has stepped and been written elsewhere -- and + it survives a later write like any other.""" + gm, first, returned = stepped + _restore(gm, returned) + path = gm.save_state(tmp_path / "stepped.npz") + + graph.load_state(path) + assert _report(graph) == first + for slot, value in gm._state["_meta"].items(): + got = graph._state["_meta"][slot] + assert np.asarray(got).tobytes() == np.asarray(value).tobytes(), slot + assert np.asarray(got).dtype == np.asarray(value).dtype, slot + WRITES["both members to zero"](graph) + assert _report(graph) == first + + # Over a graph somewhere else, with a write of its own pending. + graph.reset_state() + graph.step() + graph.step() + elsewhere = _report(graph) + assert elsewhere != first + WRITES["both members to zero"](graph) + assert _report(graph) == elsewhere + graph.load_state(path) + assert _report(graph) == first + WRITES["one member to zero"](graph) + assert _report(graph) == first + + +#: What the report of a group loaded from such a checkpoint withholds. +WITHHELD = {"spectral_usable": False, "gradient_bound_usable": False, + "precision_limited": False, "spectral_error_bound": "nan"} + + +def _archive_members(path): + with np.load(path) as archive: + return {name: archive[name] for name in archive.files} + + +def test_a_checkpoint_is_a_copy_of_the_state_and_says_when_it_was_written_after_the_step( + graph, tmp_path): + """Saved after a write, the archive holds the written state and + ``_meta`` exactly as they are, plus one member saying the group's state + was written after its last step. The graph that loads it has the same + state, ``_meta`` included, and reports what rests on the float floor + as not usable, with the reason, until the group steps; the restart + steps bit for bit like the graph that saved it.""" + graph.step() + first = _report(graph) + WRITES["both members to zero"](graph) + path = graph.save_state(tmp_path / "written.npz") + assert _report(graph) == first # saving moved nothing + live_meta = {k: np.asarray(v) for k, v in graph._state["_meta"].items()} + members = _archive_members(path) + for slot, value in live_meta.items(): + assert members[f"_meta/{slot}"].tobytes() == value.tobytes(), slot + assert members[f"_meta/{slot}"].dtype == value.dtype, slot + assert [m for m in members if m.startswith("_reports/")] == [ + f"_reports/{KEY}/written_after_step"] + graph.step() + want = ({n: np.asarray(graph.get_node_state(n)["x"]) for n in ("a", "b")}, _report(graph)) + + graph.reset_state() + graph.load_state(path) + for slot, value in live_meta.items(): + got = np.asarray(graph._state["_meta"][slot]) + assert got.tobytes() == value.tobytes() and got.dtype == value.dtype, slot + np.testing.assert_array_equal(np.asarray(graph.get_node_state("a")["x"]), 0.0) + loaded = _report(graph) + assert "written" in loaded["not_usable_reason"] + for key, value in loaded.items(): + if key == "not_usable_reason": + continue + assert value == WITHHELD.get(key, first[key]), key + # It stays so under a further write, and a save of it says so again. + WRITES["both members a million times larger"](graph) + assert _report(graph) == loaded + graph.load_state(path) + again = graph.save_state(tmp_path / "again.npz") + assert f"_reports/{KEY}/written_after_step" in _archive_members(again) + table = graph.coupling_report() + flags = " ".join(next(iter(table))["flags"]) + assert "no bound reported" in flags and "saved after" in flags, flags + # The next step brings a whole report. + graph.step() + for name in ("a", "b"): + assert np.asarray(graph.get_node_state(name)["x"]).tobytes() == want[0][name].tobytes() + assert _report(graph) == want[1] + assert "not_usable_reason" not in want[1] + + +def test_an_archive_without_the_marker_loads_as_it_always_did(graph, tmp_path): + """Both directions: the marker stripped from an archive (what a build + without it wrote) loads, with the report measured on the loaded state; + and a checkpoint saved straight after a step has no marker at all.""" + graph.step() + first = _report(graph) + clean = graph.save_state(tmp_path / "stepped.npz") + assert not [m for m in _archive_members(clean) if m.startswith("_reports/")] + WRITES["one member to zero"](graph) + marked = _archive_members(graph.save_state(tmp_path / "written.npz")) + stripped = {k: v for k, v in marked.items() if not k.startswith("_reports/")} + assert len(stripped) == len(marked) - 1 + np.savez(tmp_path / "stripped.npz", **stripped) + graph.reset_state() + graph.load_state(tmp_path / "stripped.npz") + old = _report(graph) + assert "not_usable_reason" not in old and old["iterations"] == first["iterations"] + # A marker for a group this graph does not have is ignored. + foreign = dict(stripped) + foreign["_reports/x+y/written_after_step"] = np.ones((), np.uint8) + np.savez(tmp_path / "foreign.npz", **foreign) + graph.load_state(tmp_path / "foreign.npz") + assert _report(graph) == old + # A marked load that fails leaves no mark behind. + graph.load_state(clean) + broken = dict(marked) + broken["a/x"] = np.zeros(N + 1, np.float32) + np.savez(tmp_path / "broken.npz", **broken) + with pytest.raises(ValueError): + graph.load_state(tmp_path / "broken.npz") + assert _report(graph) == first + + +def test_a_node_cannot_take_the_markers_prefix(): + gm = GraphManager() + with pytest.raises(ValueError, match="reserves for its own state and checkpoints"): + gm.add_node(Idle("_reports")) + + +@pytest.mark.parametrize("write", ["a node outside the group", "the values it already holds"]) +def test_a_checkpoint_keeps_the_report_where_no_member_changed(graph, tmp_path, write): + graph.step() + first = _report(graph) + if write == "a node outside the group": + WRITES[write](graph) + else: + for name in ("a", "b"): + graph.set_node_state(name, {"x": jnp.asarray(np.asarray( + graph.get_node_state(name)["x"]))}) + path = graph.save_state(tmp_path / "kept.npz") + graph.reset_state() + graph.load_state(path) + assert _report(graph) == first + + +def test_a_failed_load_leaves_the_report_of_the_step_before_it(stepped, tmp_path): + gm, first, returned = stepped + _restore(gm, returned) + WRITES["both members to zero"](gm) + bad = tmp_path / "not_a_checkpoint.npz" + bad.write_bytes(b"not an archive") + with pytest.raises(Exception): + gm.load_state(bad) + assert _report(gm) == first + _restore(gm, returned) + + +def test_a_recompile_keeps_the_report_and_what_it_was_measured_on(graph): + graph.step() + first = _report(graph) + WRITES["both members to zero"](graph) + graph.compile() + assert _report(graph) == first + WRITES["both members a million times larger"](graph) + assert _report(graph) == first + + +def test_reset_state_and_a_removed_member_still_end_the_report(): + gm = _pair(diagnostics=False) + gm.step() + assert _report(gm) is not None + WRITES["both members to zero"](gm) + gm.reset_state() + assert _report(gm) is None + gm.step() + first = _report(gm) + WRITES["one member to zero"](gm) + assert _report(gm) == first + with pytest.warns(UserWarning, match="removed the coupling group"): + gm.remove_node("b") + assert _report(gm) is None + + +def test_a_transform_that_stepped_the_graph_leaves_the_report_of_the_step_before_it(graph): + """A transform of a function that calls ``run_scan`` stores tracers and + the graph is put back to the state before it: the written one, whose + report is still the earlier step's.""" + graph.step() + first = _report(graph) + WRITES["both members to zero"](graph) + + def loss(p): + graph.run_scan(1, params=p) + # A write to the traced state: the copy kept before the transform + # is the one the graph comes back to, and stays. + graph.set_node_state("c", {"y": graph.get_node_state("a")["x"]}) + return jnp.sum(graph.get_node_state("a")["x"]) + + jax.make_jaxpr(loss)(graph.params) + assert graph._state_traced + with pytest.warns(RuntimeWarning, match="put back"): + assert _report(graph) == first + np.testing.assert_array_equal(np.asarray(graph.get_node_state("a")["x"]), 0.0) + + +def test_a_traced_write_keeps_no_tracer_for_the_report(graph): + """A differentiable initial condition written inside a transform: the + report read afterwards is the earlier step's and reads no tracer.""" + graph.step() + first = _report(graph) + + def loss(scale): + graph.set_node_state("a", {"x": scale * jnp.ones(N, jnp.float32)}) + graph.set_node_state("b", {"x": scale * jnp.ones(N, jnp.float32)}) + graph.run_scan(1) + # A write after a traced step: nothing of the trace may be kept. + graph.set_node_state("b", {"x": scale * jnp.ones(N, jnp.float32)}) + return jnp.sum(graph.get_node_state("a")["x"]) + + jax.make_jaxpr(loss)(jnp.float32(0.0)) + assert graph._state_traced + with pytest.warns(RuntimeWarning, match="put back"): + assert _report(graph) == first + WRITES["both members to zero"](graph) + assert _report(graph) == first + + +def test_the_profilers_measurements_hand_the_report_back_with_the_state(graph): + """``compile_counts`` and the one-iteration variant step the graph and + put the state back: the report comes back with it.""" + graph.step() + first = _report(graph) + WRITES["both members to zero"](graph) + compile_counts(graph, warmup_steps=0, scan_steps=1) + assert _report(graph) == first + with _one_iteration_variant(graph): + graph.step() + assert _report(graph) == first + + +def test_the_kept_state_holds_references_not_copies(graph): + """Keeping what the step left costs no array: the kept fields are the + step's own objects.""" + graph.step() + before = {n: graph._state[n] for n in ("a", "b", "c")} + graph.set_node_state("a", ZERO) + kept_state, _kept_meta = graph._state_as_reported + for name, fields in before.items(): + assert kept_state[name] is fields diff --git a/tests/core/test_multirate_group_meta_follows_firing.py b/tests/core/test_multirate_group_meta_follows_firing.py index 7d85f0f2d..6e98a4b31 100644 --- a/tests/core/test_multirate_group_meta_follows_firing.py +++ b/tests/core/test_multirate_group_meta_follows_firing.py @@ -157,6 +157,66 @@ def test_strict_convergence_still_raises_about_an_applied_solve(): gm.step() +def _capped_strict_graph(): + """The graph whose applied solve (phase 0) cannot meet its tolerance + within its cap, and whose discarded solves (phases 3 and 5, gains 2 and + 3) do not contract at all.""" + gm = GraphManager() + gm.add_node(_Phase("clock", 0.001)) + gm.add_node(_Gained("a", 0.01, bias=1.0, phased=True)) + gm.add_node(_Gained("b", 0.01, bias=0.0, phased=False)) + gm.add_edge("clock", "a", "phase", "phase") + gm.add_edge("b", "a", "x", "u") + gm.add_edge("a", "b", "x", "u") + gm.add_coupling_group(["a", "b"], max_iterations=2, tolerance=1e-7, + strict_convergence=True) + with warnings.catch_warnings(): + warnings.simplefilter("ignore") + gm.compile() + return gm + + +def _at_base_step(state, count): + """``state`` as the compiled step sees it on base step ``count``.""" + meta = dict(state["_meta"]) + meta["step_count"] = jnp.asarray(count, meta["step_count"].dtype) + clock = {"phase": jnp.float32((count - 1) % 10)} + return {**state, "clock": clock, "_meta": meta} + + +def _batched_step(gm, counts): + step = gm._build_step_fn() + ext = gm._default_external_inputs() + states = [_at_base_step(gm._state, c) for c in counts] + batch = jax.tree.map(lambda *leaves: jnp.stack(leaves), *states) + return states, jax.jit(jax.vmap(step, in_axes=(0, None)))(batch, ext) + + +def test_a_batched_strict_check_is_silent_about_solves_no_element_applies(): + """The strict check's firing gate, on its own. + + One at a time a non-firing base step never runs the solve (the + ``cond`` skips it), so nothing there can tell a gated check from an + ungated one. Under ``vmap`` both branches run: every element here is + between firings, each discarded solve fails its tolerance, and the + step must neither raise nor move the coupled members. + """ + gm = _capped_strict_graph() + states, out = _batched_step(gm, [3, 5]) + for i, state in enumerate(states): + for name in ("a", "b"): + assert (np.asarray(out[name]["x"][i]).tobytes() + == np.asarray(state[name]["x"]).tobytes()), (i, name) + + +def test_a_batched_strict_check_raises_when_one_element_applies_an_unconverged_solve(): + """...and the gate lets the applied solve through: one firing element + whose solve exits at its cap raises, whatever the others do.""" + gm = _capped_strict_graph() + with pytest.raises(Exception, match="without converging"): + _batched_step(gm, [5, 10]) + + def test_the_linear_predictor_learns_only_from_applied_solves(): """Two applied solves in, the extrapolated guess is still near the answer. diff --git a/tests/core/test_profiler.py b/tests/core/test_profiler.py index 4bc084560..178657d43 100644 --- a/tests/core/test_profiler.py +++ b/tests/core/test_profiler.py @@ -565,3 +565,57 @@ def test_per_node_timing_feeds_each_node_its_declared_inputs(): rep = profile_graph(gm, n_steps=2, n_warmup=1, counts=False) assert set(rep.node_times_ms) == {"a", "b"} and all( t > 0 for t in rep.node_times_ms.values()), rep.node_times_ms + + +class _Lin(SimulationNode): + """``x <- 0.9 * u + b``: a Gauss-Seidel pair of these contracts 0.81 per pass.""" + + def __init__(self, name, b): + super().__init__(name, 0.01, b=jnp.float32(b)) + + def initial_state(self): + return {"x": jnp.zeros((), jnp.float32)} + + def boundary_input_spec(self): + return {"u": BoundaryInputSpec(shape=(), dtype=jnp.float32, + default=jnp.zeros((), jnp.float32))} + + def update(self, state, boundary_inputs, dt): + # ``.get``: the profiler times each node with no inputs at all. + u = boundary_inputs.get("u", jnp.zeros((), jnp.float32)) + return {"x": jnp.float32(0.9) * u + self.params["b"]} + + +def test_the_fractions_count_every_step_of_a_window_that_recovers(): + """``at_cap_fraction`` and ``converged_fraction`` are shares of the + window's steps, not the verdict of its last one. + + Every other window here is uniform (all at the cap, or none), where a + statistic read from the last step alone gives the same number. This + pair starts 15 from its fixed point with a cap of 4 passes: its first + steps exit at the cap unconverged and, each step starting from the + previous iterate, the later ones converge. + """ + gm = GraphManager() + gm.add_node(_Lin("a", 1.0)) + gm.add_node(_Lin("b", 2.0)) + gm.add_edge("a", "b", "x", "u") + gm.add_edge("b", "a", "x", "u") + gm.add_coupling_group(["a", "b"], max_iterations=4, tolerance=1e-6) + gm.compile() + n_stat = 30 + rep = profile_graph(gm, n_steps=5, n_warmup=0, n_stat_steps=n_stat, + measure_coupling=False, counts=False) + st = rep.coupling_iter_stats["a+b"] + + gm.reset_state() + at_cap, converged = [], [] + for _ in range(n_stat): + gm.step() + d = gm.coupling_diagnostics()["a+b"] + at_cap.append(d["iterations"] >= 4) + converged.append(d["converged"]) + assert not converged[0] and converged[-1] and not at_cap[-1], (at_cap, converged) + assert st["n"] == n_stat + assert 0.0 < st["converged_fraction"] == float(np.mean(converged)) < 1.0, st + assert 0.0 < st["at_cap_fraction"] == float(np.mean(at_cap)) < 1.0, st diff --git a/tests/core/test_sysid_mask_reads_every_step_of_the_window.py b/tests/core/test_sysid_mask_reads_every_step_of_the_window.py new file mode 100644 index 000000000..e1e07eafa --- /dev/null +++ b/tests/core/test_sysid_mask_reads_every_step_of_the_window.py @@ -0,0 +1,183 @@ +"""``windowed_loss(mask_unconverged=True)`` reads every base step of a window. + +The mask drops a window "when any coupling group exited at +``max_iterations`` unconverged *during* it". Every other mask test uses a +window that is unconverged at its END (a diverged window stays diverged, a +one-step window has one verdict), so a mask that read only the window's +last step passed all of them. + +Here the window recovers: a linear Gauss-Seidel pair (contraction 0.81 per +pass, cap 4) restarted far from its fixed point exits its first steps at the +cap, and, because each step starts from the previous step's iterate, the +later steps converge. The last step's verdict alone says "keep"; the +window's says "drop". +""" + +import os + +os.environ.setdefault("JAX_PLATFORMS", "cpu") + +import jax +import jax.numpy as jnp +import numpy as np +import pytest + +from maddening import sysid +from maddening.core.graph_manager import GraphManager +from maddening.core.node import BoundaryInputSpec, SimulationNode + +STEPS = 30 +FAR = (100.0, -50.0) +# The pair's fixed point: a = 0.9 b + 1, b = 0.9 a + 2. +FIXED = (2.8 / 0.19, 2.0 + 0.9 * 2.8 / 0.19) + + +class Lin(SimulationNode): + """``x <- g * u + b``; its own state is only the coupling's first guess.""" + + def __init__(self, name, g, b): + super().__init__(name, 0.01, g=jnp.float32(g), b=jnp.float32(b)) + + def initial_state(self): + return {"x": jnp.zeros((), jnp.float32)} + + def boundary_input_spec(self): + return {"u": BoundaryInputSpec(shape=(), dtype=jnp.float32, + default=jnp.zeros((), jnp.float32))} + + def update(self, state, boundary_inputs, dt, *, params=None): + p = self.params if params is None else {**self.params, **params} + return {"x": p["g"] * boundary_inputs["u"] + p["b"]} + + +def _pair(cap): + gm = GraphManager() + gm.add_node(Lin("a", 0.9, 1.0)) + gm.add_node(Lin("b", 0.9, 2.0)) + gm.add_edge("a", "b", "x", "u") + gm.add_edge("b", "a", "x", "u") + gm.add_coupling_group(["a", "b"], max_iterations=cap, tolerance=1e-6) + gm.compile() + return gm + + +def _verdicts(gm, start, steps): + gm.reset_state() + gm.set_node_state("a", {"x": jnp.float32(start[0])}) + gm.set_node_state("b", {"x": jnp.float32(start[1])}) + out = [] + for _ in range(steps): + gm.step() + out.append(bool(gm.coupling_diagnostics()["a+b"]["converged"])) + gm.reset_state() + return out + + +def _obs(starts, n_samples): + """Observations whose window starts are ``starts`` and whose other + samples are zero (so a kept window has a loss far from zero).""" + per = n_samples + a = np.zeros(1 + per * len(starts), np.float32) + b = np.zeros_like(a) + for w, (sa, sb) in enumerate(starts): + a[w * per], b[w * per] = sa, sb + return {"a": {"x": jnp.asarray(a)}, "b": {"x": jnp.asarray(b)}} + + +@pytest.fixture(scope="module") +def capped(): + return _pair(4) + + +@pytest.fixture(scope="module") +def uncapped(): + return _pair(400) + + +def test_the_window_recovers_its_last_step_converges_and_its_early_ones_do_not(capped): + """The fixture expresses the case: early steps at the cap, the last one fine.""" + v = _verdicts(capped, FAR, STEPS) + assert not v[0] and not v[1], v + assert v[-1], v + # Monotone: once the iterate is close enough, every later step converges. + first_ok = v.index(True) + assert all(v[first_ok:]) and 2 <= first_ok < STEPS - 1, v + + +# (window, sample_every): the verdict is and-ed over the base steps of one +# sample and over the samples of one window; each shape puts the recovery +# across a different one of the two reductions. +SHAPES = [(STEPS, 1), (1, STEPS), (6, 5)] + + +@pytest.mark.parametrize("window,sample_every", SHAPES) +def test_a_window_whose_early_steps_hit_the_cap_is_dropped_though_its_last_step_converged( + capped, window, sample_every): + obs = _obs([FAR], window) + + def loss(p, mask): + return sysid.windowed_loss(capped, p, obs, obs_fn=lambda s: s, window=window, + sample_every=sample_every, mask_unconverged=mask) + + value, grad = jax.value_and_grad(loss)(capped.params, True) + assert float(value) == 0.0 + for leaf in jax.tree.leaves(grad): + np.testing.assert_array_equal(np.asarray(leaf), 0.0) + # The window is not trivially zero: unmasked, it has a loss and a gradient. + kept, kept_grad = jax.value_and_grad(loss)(capped.params, False) + assert float(kept) > 1.0 + assert float(abs(kept_grad["nodes"]["a"]["g"])) > 0.0 + + +@pytest.mark.parametrize("window,sample_every", SHAPES) +def test_the_same_window_is_kept_when_the_cap_lets_every_step_converge( + uncapped, window, sample_every): + assert all(_verdicts(uncapped, FAR, STEPS)) + obs = _obs([FAR], window) + + def loss(p, mask): + return sysid.windowed_loss(uncapped, p, obs, obs_fn=lambda s: s, window=window, + sample_every=sample_every, mask_unconverged=mask) + + value, grad = jax.value_and_grad(loss)(uncapped.params, True) + plain, plain_grad = jax.value_and_grad(loss)(uncapped.params, False) + assert float(value) == float(plain) > 1.0 + for a, b in zip(jax.tree.leaves(grad), jax.tree.leaves(plain_grad)): + np.testing.assert_array_equal(np.asarray(a), np.asarray(b)) + + +def test_multiple_shooting_drops_the_recovering_window_and_keeps_the_converged_one(capped): + """Two free-start windows: the first recovers (dropped, with its + continuity term), the second starts at the fixed point (kept).""" + assert all(_verdicts(capped, FIXED, STEPS)) + obs = _obs([FAR, FIXED], STEPS) + ws = sysid.init_window_states(obs, STEPS) + + def loss(p, w, mask=True): + return sysid.windowed_loss(capped, p, obs, obs_fn=lambda s: s, window=STEPS, + mask_unconverged=mask, window_states=w, + continuity_weight=1.0) + + value, (gp, gw) = jax.value_and_grad(loss, argnums=(0, 1))(capped.params, ws) + + # What is left is exactly the second window alone (it is the last + # window, so it has no continuity term of its own). + tail = jax.tree.map(lambda x: x[STEPS:], obs) + tail_ws = sysid.init_window_states(tail, STEPS) + + def alone(p, w): + return sysid.windowed_loss(capped, p, tail, obs_fn=lambda s: s, window=STEPS, + mask_unconverged=True, window_states=w, + continuity_weight=1.0) + + ref, (rp, rw) = jax.value_and_grad(alone, argnums=(0, 1))(capped.params, tail_ws) + assert float(value) == float(ref) > 1.0 + for a, b in zip(jax.tree.leaves(gp), jax.tree.leaves(rp)): + np.testing.assert_array_equal(np.asarray(a), np.asarray(b)) + # The dropped window's free start has no gradient; the kept one's has + # the gradient it has alone. + for leaf, ref_leaf in zip(jax.tree.leaves(gw), jax.tree.leaves(rw)): + np.testing.assert_array_equal(np.asarray(leaf[0]), 0.0) + np.testing.assert_array_equal(np.asarray(leaf[1]), np.asarray(ref_leaf[0])) + # Unmasked, the first window and its continuity term are in the loss. + assert float(loss(capped.params, ws, mask=False)) > float(value) diff --git a/tests/property/test_coupling_error_bound.py b/tests/property/test_coupling_error_bound.py index f52037dde..591f9f2b3 100644 --- a/tests/property/test_coupling_error_bound.py +++ b/tests/property/test_coupling_error_bound.py @@ -74,6 +74,7 @@ coupling_residual_l2, coupling_residual_mixed, relaxation_step_scale, + spectral_rate_settled, ) from maddening.core.node import BoundaryInputSpec, SimulationNode @@ -1072,6 +1073,29 @@ def _normal_contractions(draw): return A, c, acceleration, relaxation +#: A normal contraction whose flag 0.4.0 withdraws (eigenvalues 0 and 0.5 +#: five times; float32): see the test below for why. +_GROWING_AT_THE_CAP = ( + np.array([ + [0.499134, -0.0089817, 0.01601435, -0.00283523, -0.00622292, 0.00695405], + [-0.0089817, 0.40684713, 0.16609148, -0.02940534, -0.06454052, 0.07212334], + [0.01601435, 0.16609148, 0.20385906, 0.05242969, 0.11507568, -0.12859584], + [-0.00283523, -0.02940534, 0.05242969, 0.49071769, -0.02037335, 0.022767], + [-0.00622292, -0.06454052, 0.11507568, -0.02037335, 0.45528341, 0.04997031], + [0.00695405, 0.07212334, -0.12859584, 0.022767, 0.04997031, 0.44415872], + ]), + np.array([-0.05665856, 1.55795134, 1.73617406, -0.56881921, 0.28611932, -0.71252244]), + "none", + 1.0, +) + +#: The least share of drawn normal contractions ``spectral_usable`` is set +#: on. Measured over this test's own draws (CPU, jaxlib 0.11.0): 300 of +#: 300 over fifteen seeded runs of twenty, and the one draw pinned above in +#: one unseeded run. The floor leaves two of twenty. +_NORMAL_USABLE_FLOOR = 0.9 # units: dimensionless -- a share of the drawn examples + + # Slow-marked (still run by slow-tests.yml): the draw varies the mode count # (a shape) and, under ``"fixed"``, the static relaxation factor, so most # examples are a compile of their own (31-37 s on the CI runner). The @@ -1081,43 +1105,78 @@ def _normal_contractions(draw): # Per push: tests/property/test_coupling_error_bound.py::test_the_spectral_bound_is_never_smaller_than_the_distance_it_bounds and # tests/core/test_coupling_error_bound.py::test_the_spectral_bound_does_not_depend_on_the_relaxation_factor @pytest.mark.slow -@settings(max_examples=EXAMPLES_COSTLY, deadline=None) -@given(case=_normal_contractions()) -def test_the_spectral_bound_holds_on_random_normal_contractions(case): - """``spectral_error_bound >= ||x - x*||`` for any normal ``A``. +def test_the_spectral_bound_holds_on_random_normal_contractions(): + """Where ``spectral_usable``, ``spectral_error_bound >= ||x - x*||`` and + ``rho_spectral`` is the radius, for any normal ``A``; and the flag is + set on at least :data:`_NORMAL_USABLE_FLOOR` of the draws. Under ``"none"`` and ``"fixed"`` at any relaxation, with negative eigenvalues, alternation and a spectral radius up to 0.98, and whether or not the group met its criterion within the cap -- the - bound is about the returned iterate, not about convergence. The - spectrum is exact here (``n <= 6 < SPECTRAL_KRYLOV_STEPS``), so - ``spectral_usable`` is asserted True rather than assumed: a False - would mean the Arnoldi breakdown handling stopped resolving a - resolvable spectrum. + bound is about the returned iterate, not about convergence. + + **Why the flag is not asserted on every draw.** This test used to + assert ``spectral_usable is True`` on each, reasoning that the spectrum + is exact (``n <= 6 < SPECTRAL_KRYLOV_STEPS``). No claim says so: the + bound is claimed where the flag is set (CPL-088), and the flag is + False "for a group whose Krylov space is still growing at the cap" + (CPL-092), a rule 0.4.0 holds without exception since a Ritz value of + a space that is not invariant is within no computed distance of the + radius. :data:`_GROWING_AT_THE_CAP` is a normal contraction it + withdraws the flag on: five equal eigenvalues, so the Krylov space of + the start vector closes after two steps and continues from the + residual, whose part outside the one eigenspace is the float32 + iterate's rounding (1e-4 of a residual of 2.6e-4, far above the + breakdown test's eight ulps of a product). That is carried as a + direction, each later step grows by 1e-5 to 5e-4, and the space is + still growing at step eight. The radius and the bound it reported + were right; the flag declines to say so. The example is pinned, and + the share of draws the flag is set on is held to a floor instead, so + a change that withdraws it broadly still fails here. """ - A, c, acceleration, relaxation = case - kw = dict(acceleration=acceleration) - if acceleration == "fixed": - kw["relaxation"] = relaxation - gm = _linear_cycle(A, c, max_iterations=400, - tolerance=_ANALYTIC_TOLERANCE, **kw) - gm.step() - d = gm.coupling_diagnostics()["a+b"] - # The fixed point of the float32 map the graph evaluates (see - # ``_exact_distance``), taken in float64. - A32 = np.asarray(A, np.float32).astype(np.float64) - c32 = np.asarray(c, np.float32).astype(np.float64) - x_star = np.linalg.solve(np.eye(len(c)) - A32, c32) - distance = _distance_to(gm, x_star) - rho = float(np.max(np.abs(np.linalg.eigvalsh(A)))) - note(f"rho={rho} {acceleration} omega={relaxation} distance={distance} {d}") - assert d["spectral_usable"] is True, d - assert d["rho_spectral"] == pytest.approx(rho, abs=1e-4), ( - f"rho_spectral={d['rho_spectral']} for a spectral radius of {rho}" - ) - # Bare: the key carries its own float resolution (see the two-mode - # property above). - assert d["spectral_error_bound"] >= distance, ( - f"{acceleration} omega={relaxation}: bound {d['spectral_error_bound']:.4e} " - f"below the true distance {distance:.4e} at rho={rho:.4f}" - ) + usable = [] + + @settings(max_examples=EXAMPLES_COSTLY, deadline=None) + @example(case=_GROWING_AT_THE_CAP) + @given(case=_normal_contractions()) + def check(case): + A, c, acceleration, relaxation = case + kw = dict(acceleration=acceleration) + if acceleration == "fixed": + kw["relaxation"] = relaxation + gm = _linear_cycle(A, c, max_iterations=400, + tolerance=_ANALYTIC_TOLERANCE, **kw) + gm.step() + d = gm.coupling_diagnostics()["a+b"] + # The fixed point of the float32 map the graph evaluates (see + # ``_exact_distance``), taken in float64. + A32 = np.asarray(A, np.float32).astype(np.float64) + c32 = np.asarray(c, np.float32).astype(np.float64) + x_star = np.linalg.solve(np.eye(len(c)) - A32, c32) + distance = _distance_to(gm, x_star) + rho = float(np.max(np.abs(np.linalg.eigvalsh(A)))) + note(f"rho={rho} {acceleration} omega={relaxation} distance={distance} {d}") + usable.append(bool(d["spectral_usable"])) + if not d["spectral_usable"]: + # Withdrawn for the reason the flag documents, not for none: + # the stored Arnoldi residual is over the settle margin. + stored = float(gm._state["_meta"]["coupling_a+b_spectral_residual"]) + assert not bool(spectral_rate_settled(d["rho_spectral"], stored)), d + return + assert d["rho_spectral"] == pytest.approx(rho, abs=1e-4), ( + f"rho_spectral={d['rho_spectral']} for a spectral radius of {rho}" + ) + # Bare: the key carries its own float resolution (see the two-mode + # property above). + assert d["spectral_error_bound"] >= distance, ( + f"{acceleration} omega={relaxation}: bound {d['spectral_error_bound']:.4e} " + f"below the true distance {distance:.4e} at rho={rho:.4f}" + ) + + check() + assert not usable[0], "the pinned draw: its Krylov space grows at the cap" + drawn = usable[1:] + assert drawn, "no example was drawn" + print(f"spectral_usable on {sum(drawn)} of {len(drawn)} drawn normal contractions") + assert float(np.mean(drawn)) >= _NORMAL_USABLE_FLOOR, ( + f"spectral_usable on {sum(drawn)} of {len(drawn)} normal contractions") diff --git a/tests/property/test_coupling_targeted_search.py b/tests/property/test_coupling_targeted_search.py index c31e50695..6a208f2f1 100644 --- a/tests/property/test_coupling_targeted_search.py +++ b/tests/property/test_coupling_targeted_search.py @@ -86,6 +86,7 @@ os.environ.setdefault("JAX_PLATFORMS", "cpu") +import jax.numpy as jnp import numpy as np import pytest from hypothesis import strategies as st @@ -750,6 +751,46 @@ def test_every_score_holds_on_the_seed_shapes(seed): assert not over, f"{seed}: {over} ({seen['report']})" +def _same_report(a: dict, b: dict) -> bool: + """Equal key for key, a NaN (a number not computed) equal to itself.""" + def norm(d): + return {k: ("nan" if isinstance(v, float) and v != v else v) for k, v in d.items()} + return norm(a) == norm(b) + + +@pytest.mark.parametrize("seed", sorted(SEEDS)) +def test_a_report_on_a_seed_shape_does_not_move_when_the_state_is_written_afterwards(seed): + """Step, write the state, read the report: it is still the step's. + + The scores above are of the report read right after the step. The + floor under the bound is measured on the returned state at report + time, so a ``set_node_state`` in between used to give the same step + another bound (to 0.0 with its flag set, where every member was + written to zero: a field at exactly zero leaves the norm). + """ + case = SEEDS[seed] + cell = CELLS[case.cell] + with precision(cell.dtype == "float64"): + built = _built(case.cell) + (step,) = ct.run(built, values_of(case), 1) + gm, key = built.gm, cell.topo.group_key(0) + first = dict(gm.coupling_diagnostics()[key]) + assert _same_report(first, step.reports[0]) + members = key.split("+") + returned = {name: gm.get_node_state(name) for name in members} + edits = { + "one member to zero": {members[0]: 0.0}, + "every member to zero": dict.fromkeys(members, 0.0), + "every member a thousand times larger": dict.fromkeys(members, 1e3), + "put back": dict.fromkeys(members, 1.0), + } + for label, factors in edits.items(): + for name, factor in factors.items(): + gm.set_node_state(name, {f: v * jnp.asarray(factor, v.dtype) + for f, v in returned[name].items()}) + assert _same_report(dict(gm.coupling_diagnostics()[key]), first), (seed, label) + + def _known(case: Case, score: str, reason: str): return pytest.param(case, score, marks=pytest.mark.xfail(strict=True, reason=reason))