Files
SystemSimulationApp/tests/manual/compare_jacobian_trajectories.py
T

302 lines
19 KiB
Python

r"""Diagnose complete native trajectories without changing acceptance tolerances.
.venv/bin/python tests/manual/compare_jacobian_trajectories.py \
--baseline old/result.json --candidate new/result.json \
--manifest new/cache/KEY/manifest.json --output test/jacobian/comparison.json
Optional --reference tight/result.json compares each trajectory to that supplied
reference. Its precision/convergence must be established separately. There is no
numerical pass threshold here: successful analysis is not accuracy acceptance.
The ordinary grid defaults to 0..10 s at .01 s. Only exactly shared grid times
are compared; all off-grid saved points are listed and examined separately.
"""
from __future__ import annotations
import argparse
from hashlib import sha256
import json
import math
from pathlib import Path
import sys
ROOT = Path(__file__).resolve().parents[2]
sys.path.insert(0, str(ROOT))
from app.simulation.native_codegen.tolerances import state_absolute_tolerance
RTOL = 1e-8
PAYLOAD = {"series", "final", "finalState"}
def read_json(path: Path):
def unique(items):
value = {}
for key, item in items:
if key in value:
raise ValueError(f"Duplicate JSON key: {key}")
value[key] = item
return value
def reject(token):
raise ValueError(f"Nonfinite JSON token: {token}")
def parsed_float(token):
value = float(token)
if not math.isfinite(value):
raise ValueError(f"Nonfinite JSON number: {token}")
return value
raw = path.read_bytes()
value = json.loads(raw, parse_float=parsed_float, parse_int=lambda token: -0.0 if token == "-0" else int(token),
parse_constant=reject, object_pairs_hook=unique)
return value, {"path": str(path.resolve()), "sha256": sha256(raw).hexdigest(), "bytes": len(raw)}
def finite(value):
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValueError("Expected a finite numeric payload")
number = float(value)
if not math.isfinite(number) or (isinstance(value, int) and int(number) != value):
raise ValueError("Payload is nonfinite or not exactly representable as binary64")
return number
def rms(values):
return math.hypot(*values) / math.sqrt(len(values)) if values else None
def peak(values, times):
i = max(range(len(values)), key=lambda i: abs(values[i]))
return {"absolute": abs(values[i]), "value": values[i], "time": times[i]}
def validate(data, metadata, states, start, stop):
if not isinstance(data, dict) or data.get("success") is not True or data.get("status") != "completed":
raise ValueError("A completed native result is required")
if not isinstance(data.get("series"), dict) or not isinstance(data.get("final"), dict) or not isinstance(data.get("finalState"), list):
raise ValueError("Missing native series/final/finalState structure")
if set(data["series"]) != set(metadata) | {"time"} or set(data["final"]) != set(metadata):
raise ValueError("Series/final key sets do not match the supplied manifest")
if len(data["finalState"]) != len(states) or not set(states) <= set(metadata):
raise ValueError("State keys/order cannot be mapped from the supplied manifest")
times = data["series"]["time"]
if not isinstance(times, list) or not times:
raise ValueError("No saved samples")
for key, values in data["series"].items():
if not isinstance(values, list) or len(values) != len(times):
raise ValueError(f"Invalid sample length: {key}")
for value in values:
finite(value)
for value in (*data["final"].values(), *data["finalState"]):
finite(value)
if any(a >= b for a, b in zip(times, times[1:])):
raise ValueError("Sample times are not strictly increasing")
if times[0] != start or times[-1] != stop or data.get("simulatedUntil") != stop:
raise ValueError("Result does not span the requested complete interval")
def input_observations(data, metadata, states, grid):
series, times = data["series"], data["series"]["time"]
fixed = set(grid)
mapping = {t: i for i, t in enumerate(times)}
extras = [i for i, t in enumerate(times) if t not in fixed]
mechanical = [key for key, item in metadata.items() if item.get("scope") == "component"
and item.get("quantity") in {"force", "length", "velocity", "acceleration"}]
events = []
for i in extras:
indices = list(range(max(0, i - 1), min(len(times), i + 2)))
neighborhood = {}
for key, item in metadata.items():
group = (item["quantity"], item["unit"])
best = max(indices, key=lambda j: abs(series[key][j]))
current = neighborhood.get(group)
if current is None or abs(series[key][best]) > current["absolute"]:
neighborhood[group] = {"quantity": group[0], "unit": group[1], "key": key,
"absolute": abs(series[key][best]), "value": series[key][best], "time": times[best]}
events.append({"sampleIndex": i, "time": times[i], "neighborSampleTimes": [times[j] for j in indices],
"stateSamples": [{"time": times[j], "values": {key: series[key][j] for key in states}} for j in indices],
"mechanicalSamples": [{"key": key, "unit": metadata[key]["unit"],
"values": [series[key][j] for j in indices],
"neighborhoodPeak": peak([series[key][j] for j in indices], [times[j] for j in indices])}
for key in mechanical],
"neighborhoodQuantityPeaks": list(neighborhood.values())})
mass_keys = [key for key in states if key.rsplit(".", 1)[-1] in {"m", "m1", "m2"}
and metadata[key]["quantity"] == "mass" and metadata[key]["unit"] == "kg"]
conservation = None
if mass_keys:
totals = [math.fsum(series[key][i] for key in mass_keys) for i in range(len(times))]
drift = [abs(value - totals[0]) for value in totals]
worst = max(range(len(times)), key=drift.__getitem__)
conservation = {"stateKeys": mass_keys, "uniqueMassStateCount": len(mass_keys), "initialKg": totals[0],
"finalKg": totals[-1], "maxAbsoluteDriftKg": drift[worst], "worstTime": times[worst],
"sampleCount": len(times), "includesOffGridSamples": True,
"method": "math.fsum of unique manifest gas mass states; output aliases are not accumulated"}
return {"metadata": {k: v for k, v in data.items() if k not in PAYLOAD},
"sampleCount": len(times), "seriesVariableCount": len(metadata), "stateCount": len(states),
"missingFixedGridTimes": [t for t in grid if t not in mapping],
"excludedFromFixedGrid": [{"index": i, "time": times[i]} for i in extras],
"extraSampleCount": len(extras), "reportedStateTransitions": data.get("stateTransitions"),
"eventInterpretation": "Off-grid saved points are event candidates, not guaranteed event identities. An event on the fixed grid is not distinguishable from ordinary samples. Neighbors are nearest saved samples, not the true pre/post impact limits; a non-grid final endpoint can also be extra.",
"events": events, "massConservation": conservation,
"final": data["final"], "finalStateByKey": dict(zip(states, data["finalState"], strict=True)),
"finalVsLastSeriesUnequalKeys": [key for key in metadata if data["final"][key] != series[key][-1]],
"finalStateVsLastSeriesUnequalKeys": [key for key, value in zip(states, data["finalState"], strict=True) if value != series[key][-1]]}
def aggregate(rows):
groups = {}
for row in rows:
group = (row["quantity"], row["unit"])
if group not in groups:
groups[group] = {"quantity": group[0], "unit": group[1], "variableCount": 0,
"sampleValues": 0, "maxAbsoluteError": -1., "norm": 0.}
total = groups[group]
total["variableCount"] += 1
total["sampleValues"] += row["sampleCount"]
total["norm"] = math.hypot(total["norm"], row["rmsError"] * math.sqrt(row["sampleCount"]))
if row["maxAbsoluteError"] > total["maxAbsoluteError"]:
total.update(maxAbsoluteError=row["maxAbsoluteError"], worstKey=row["key"], worstTime=row["worstTime"])
for total in groups.values():
total["rmsError"] = total.pop("norm") / math.sqrt(total["sampleValues"])
return list(groups.values())
def trajectory_comparison(left, right, metadata, states, grid, label):
a_times, b_times = left["series"]["time"], right["series"]["time"]
a_index, b_index = {t: i for i, t in enumerate(a_times)}, {t: i for i, t in enumerate(b_times)}
common = [t for t in grid if t in a_index and t in b_index]
if not common:
raise ValueError(f"No exactly shared fixed-grid times: {label}")
ai, bi = [a_index[t] for t in common], [b_index[t] for t in common]
curves = []
state_rows = []
state_norms = [0.] * len(common)
state_set = set(states)
for key, item in metadata.items():
av = [left["series"][key][i] for i in ai]
bv = [right["series"][key][i] for i in bi]
errors = [b - a for a, b in zip(av, bv, strict=True)]
if not all(math.isfinite(e) for e in errors):
raise ValueError(f"Difference exceeds binary64 range: {key}")
worst = max(range(len(common)), key=lambda i: abs(errors[i]))
row = {"key": key, "quantity": item["quantity"], "unit": item["unit"], "sampleCount": len(common),
"maxAbsoluteError": abs(errors[worst]), "rmsError": rms(errors), "worstTime": common[worst],
"leftValueAtWorst": av[worst], "rightValueAtWorst": bv[worst],
"leftAllSavedPeak": peak(left["series"][key], a_times),
"rightAllSavedPeak": peak(right["series"][key], b_times)}
curves.append(row)
if key in state_set:
atol = float(state_absolute_tolerance(key))
z = [abs(e) / (atol + RTOL * max(abs(a), abs(b))) for e, a, b in zip(errors, av, bv, strict=True)]
at = max(range(len(z)), key=z.__getitem__)
state_rows.append({"key": key, "unit": item["unit"], "atol": atol,
"maxWeightedError": z[at], "rmsWeightedError": rms(z), "worstTime": common[at],
"maxAbsoluteError": row["maxAbsoluteError"], "rmsError": row["rmsError"]})
for i, value in enumerate(z):
state_norms[i] = math.hypot(state_norms[i], value)
wrms = [value / math.sqrt(len(states)) for value in state_norms]
weighted_worst = max(state_rows, key=lambda row: row["maxWeightedError"])
wrms_at = max(range(len(wrms)), key=wrms.__getitem__)
final_rows = []
for key, item in metadata.items():
a, b = left["final"][key], right["final"][key]
error = abs(b - a)
final_rows.append({"key": key, "quantity": item["quantity"], "unit": item["unit"],
"sampleCount": 1, "maxAbsoluteError": error, "rmsError": error,
"worstTime": right["simulatedUntil"], "left": a, "right": b})
terminal = []
for key, a, b in zip(states, left["finalState"], right["finalState"], strict=True):
atol = float(state_absolute_tolerance(key))
terminal.append({"key": key, "unit": metadata[key]["unit"], "left": a, "right": b,
"absoluteError": abs(b - a), "weightedError": abs(b - a) / (atol + RTOL * max(abs(a), abs(b)))})
fixed = set(grid)
a_extra, b_extra = [t for t in a_times if t not in fixed], [t for t in b_times if t not in fixed]
return {"label": label, "comparedFixedGridTimes": common, "comparedSampleCount": len(common),
"expectedFixedGridCount": len(grid), "allFixedGridTimesCompared": len(common) == len(grid),
"missingFromLeft": [t for t in grid if t not in a_index], "missingFromRight": [t for t in grid if t not in b_index],
"curves": curves, "quantityGroups": aggregate(curves),
"stateErrors": {"formula": "abs(right-left)/(state_atol + 1e-8*max(abs(left),abs(right))), independently at each state/time",
"interpretation": "Diagnostic normalization only. Local integration tolerances are not global trajectory acceptance thresholds.",
"rows": state_rows, "maxWeightedError": weighted_worst["maxWeightedError"],
"worstKey": weighted_worst["key"], "worstTime": weighted_worst["worstTime"],
"wrmsAtEachComparedTime": wrms, "maxWrms": wrms[wrms_at], "maxWrmsTime": common[wrms_at]},
"final": {"rows": final_rows, "quantityGroups": aggregate(final_rows)},
"finalState": {"rows": terminal, "maxWeightedError": max(row["weightedError"] for row in terminal),
"wrms": rms([row["weightedError"] for row in terminal])},
"eventTimeDiagnostics": {"leftExtraTimes": a_extra, "rightExtraTimes": b_extra,
"leftCount": len(a_extra), "rightCount": len(b_extra),
"reportedTransitions": [left.get("stateTransitions"), right.get("stateTransitions")],
"ordinalTimeDifferences": [{"ordinal": i + 1, "leftTime": a, "rightTime": b, "rightMinusLeftSeconds": b - a}
for i, (a, b) in enumerate(zip(a_extra, b_extra, strict=True))] if len(a_extra) == len(b_extra) else None,
"interpretation": "Equal-count ordinal differences are observations only, not verified physical event matching. Different counts are not paired. Event and neighbor values/peaks are in input observations; no time shifting or interpolation."}}
def main():
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--baseline", required=True, type=Path)
parser.add_argument("--candidate", required=True, type=Path)
parser.add_argument("--manifest", required=True, type=Path)
parser.add_argument("--reference", type=Path)
parser.add_argument("--output", required=True, type=Path)
parser.add_argument("--start", type=float, default=0.)
parser.add_argument("--stop", type=float, default=10.)
parser.add_argument("--sample-step", type=float, default=.01)
args = parser.parse_args()
if not all(math.isfinite(v) for v in (args.start, args.stop, args.sample_step)) or not (args.stop > args.start and args.sample_step > 0):
parser.error("Require finite increasing interval and positive sample step")
if (args.stop - args.start) / args.sample_step > 1000000 or args.start + args.sample_step == args.start:
parser.error("Invalid or excessive sampling grid")
output = args.output.resolve()
sources = [args.baseline, args.candidate, args.manifest] + ([args.reference] if args.reference else [])
if output in {path.resolve() for path in sources}:
parser.error("Output must not overwrite an input")
report = {"schemaVersion": 1, "complete": False, "errors": [], "inputs": {}, "comparisons": {},
"numericalAcceptance": "Not assessed: diagnostic errors only; no tolerance relaxation or automatic pass threshold.",
"definitions": {"fixedGrid": "start + integer index * sampleStep; exact floating-point time membership, no interpolation",
"quantityRms": "Pooled RMS of every compared value in that quantity/unit group; output aliases are included. Per-curve RMS is also reported.",
"peaks": "All saved points including off-grid events. A sampled peak need not be the continuous-time peak.",
"reference": "User-supplied reference; native JSON alone does not establish tighter effective tolerances or convergence. Both comparisons retain normalization rtol=1e-8.",
"sharedManifest": "Caller must establish identical state order and physical output mapping for every input; one shared manifest is checked against all key sets and lengths."},
"scriptSha256": sha256(Path(__file__).read_bytes()).hexdigest(),
"stateToleranceSourceSha256": sha256((ROOT / 'app/simulation/native_codegen/tolerances.py').read_bytes()).hexdigest(),
"normalizationRtol": RTOL, "grid": {"start": args.start, "stop": args.stop, "sampleStep": args.sample_step}}
output.parent.mkdir(parents=True, exist_ok=True)
try:
manifest, identity = read_json(args.manifest)
report["manifest"] = identity
states = manifest["stateKeys"]
variables = manifest["variables"]
metadata = {row["key"]: row for row in variables}
if not states or len(states) != len(set(states)) or len(metadata) != len(variables):
raise ValueError("Manifest has empty/duplicate state keys or duplicate output keys")
if not all(isinstance(row.get("quantity"), str) and isinstance(row.get("unit"), str) for row in variables):
raise ValueError("Every manifest output requires quantity and unit metadata")
report["stateKeys"] = states
report["stateAbsoluteTolerances"] = {key: float(state_absolute_tolerance(key)) for key in states}
grid = []
i = 0
while (t := args.start + i * args.sample_step) <= args.stop:
grid.append(t)
i += 1
data = {}
for label, path in [("baseline", args.baseline), ("candidate", args.candidate)] + ([("reference", args.reference)] if args.reference else []):
data[label], identity = read_json(path)
report["inputs"][label] = identity
validate(data[label], metadata, states, args.start, args.stop)
report["inputs"][label]["observations"] = input_observations(data[label], metadata, states, grid)
report["comparisons"]["candidateVsBaseline"] = trajectory_comparison(data["baseline"], data["candidate"], metadata, states, grid, "candidate minus baseline")
if args.reference:
for label in ("baseline", "candidate"):
report["comparisons"][label + "VsReference"] = trajectory_comparison(data["reference"], data[label], metadata, states, grid, label + " minus supplied reference")
report["complete"] = True
except (OSError, ValueError, KeyError, TypeError, OverflowError) as exc:
report["errors"].append(f"{type(exc).__name__}: {exc}")
output.write_text(json.dumps(report, ensure_ascii=False, indent=2, allow_nan=False) + "\n", encoding="utf-8")
print(json.dumps({"complete": report["complete"], "errors": report["errors"], "output": str(output),
"comparedSamples": {key: value["comparedSampleCount"] for key, value in report["comparisons"].items()},
"numericalAcceptance": report["numericalAcceptance"]}, ensure_ascii=False), flush=True)
return 0 if report["complete"] else 2
if __name__ == "__main__":
raise SystemExit(main())