Skip to content
Merged
2 changes: 2 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Expand Down Expand Up @@ -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**:
Expand Down
51 changes: 51 additions & 0 deletions benchmarks/results/gpu_eigvals_probe/RESULT.md
Original file line number Diff line number Diff line change
@@ -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.
24 changes: 24 additions & 0 deletions benchmarks/results/gpu_eigvals_probe/cost_split.py
Original file line number Diff line number Diff line change
@@ -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")
99 changes: 99 additions & 0 deletions benchmarks/results/gpu_eigvals_probe/probe.py
Original file line number Diff line number Diff line change
@@ -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_<platform>.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)
165 changes: 165 additions & 0 deletions benchmarks/results/gpu_eigvals_probe/probe_cpu.json
Original file line number Diff line number Diff line change
@@ -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
}
}
Loading
Loading