#!/usr/bin/env python3 """High-precision reference for the quadratic entry benchmark. Inputs are exact IEEE-754 values, so they are represented exactly as ``fractions.Fraction``. The relative-distance quadratic coefficients, the discriminant and the polynomial residuals are therefore *exact* rationals; only the square root needs ``decimal`` (160 significant digits by default). This avoids ``float.fromhex`` (which would silently drop a long-double coefficient to 53 bits) and keeps the sign of a near-zero discriminant exact. Reference classification mirrors the production contract: * c < 0 -> camera already INSIDE (never an entry failure) * c == 0 -> boundary slope b decides; b < 0 enters at once, b == 0 with a < 0 enters later, otherwise no crossing * c > 0 -> smallest positive root with inward slope; a missing/tangent root is MISS. ``a == 0`` is the exact linear branch. """ from __future__ import annotations import csv import json import math from decimal import Decimal, localcontext from fractions import Fraction DEFAULT_PREC = 160 def parse_hex(s: str) -> Fraction: """Exact rational value of a C ``%a`` / ``%La`` hex float literal.""" s = s.strip() low = s.lower() if "nan" in low: raise ValueError("NaN input") if "inf" in low: return Fraction(0) # only used for flags; never present in fixtures neg = False if s and s[0] in "+-": neg = s[0] == "-" s = s[1:] if s[:2].lower() == "0x": s = s[2:] mant, _, exp_s = s.partition("p") if not exp_s: mant, _, exp_s = s.partition("P") exp = int(exp_s) if exp_s else 0 ip, _, fp = mant.partition(".") digits = (ip + fp) or "0" val = Fraction(int(digits, 16), 1) shift = exp - 4 * len(fp) if shift >= 0: val *= Fraction(2) ** shift else: val /= Fraction(2) ** (-shift) return -val if neg else val def frac_to_dec(fr: Fraction, prec: int = DEFAULT_PREC) -> Decimal: with localcontext() as ctx: ctx.prec = prec return Decimal(fr.numerator) / Decimal(fr.denominator) def hex_to_dec(s: str, prec: int = DEFAULT_PREC) -> Decimal: return frac_to_dec(parse_hex(s), prec) def classify(a: Fraction, b: Fraction, c: Fraction, prec: int, R0: Fraction | None = None, rr: Fraction | None = None): """Return (status, root_decimal_or_None, info). ``R0``/``rr`` (radius at the segment start and dR/dt) enforce the physical positive-radius domain R(s) = R0 - rr*s > 0 along the past parameter. A positive root outside that domain is not a physical entry: it is reported as MISS with ``info['domain_clipped']`` set (shrinking worldtubes). """ info = {"a_zero": a == 0, "disc_sign": 0, "domain_clipped": False} def domain_ok(root: Fraction) -> bool: if R0 is None: return True return (R0 - (rr if rr is not None else Fraction(0)) * root) > 0 if c < 0: return "INSIDE", None, info if c == 0: if b < 0: return "ENTER", Decimal(0), info if b == 0: if a < 0: return "ENTER", Decimal(0), info return "MISS", None, info if a < 0: root = -b / a if root > 0 and domain_ok(root): return "ENTER", frac_to_dec(root, prec), info if root > 0: info["domain_clipped"] = True return "MISS", None, info return "MISS", None, info if a == 0: if b < 0: root = -c / b if root > 0: if domain_ok(root): return "ENTER", frac_to_dec(root, prec), info info["domain_clipped"] = True return "MISS", None, info disc = b * b - 4 * a * c info["disc_sign"] = (disc > 0) - (disc < 0) if disc < 0: return "MISS", None, info if disc == 0: return "MISS", None, info # tangency is not a crossing with localcontext() as ctx: ctx.prec = prec sd = frac_to_dec(disc, prec).sqrt() ad = frac_to_dec(a, prec) bd = frac_to_dec(b, prec) r1 = (-bd - sd) / (2 * ad) r2 = (-bd + sd) / (2 * ad) pos = sorted(r for r in (r1, r2) if r > 0) for r in pos: if 2 * ad * r + bd < 0: # inward (outside -> inside) slope rfr = _dec_to_frac_snapshot(r) if domain_ok(rfr): return "ENTER", r, info info["domain_clipped"] = True return "MISS", None, info def _dec_to_frac_snapshot(d: Decimal) -> Fraction: return Fraction(d) def coeffs(x, c, w, v, R0, rr): d = [x[i] - c[i] for i in range(3)] q = [w[i] + v[i] for i in range(3)] qq = sum(q[i] * q[i] for i in range(3)) dq = sum(d[i] * q[i] for i in range(3)) dd = sum(d[i] * d[i] for i in range(3)) a = qq - rr * rr b = 2 * (dq + R0 * rr) cq = dd - R0 * R0 return a, b, cq, qq def load_cases(path): data = json.loads(open(path).read()) kernel = {} for c in data["kernel"]: kernel[c["id"]] = { "category": c["category"], "x": [parse_hex(z) for z in c["x"]], "c": [parse_hex(z) for z in c["c"]], "w": [parse_hex(z) for z in c["w"]], "v": [parse_hex(z) for z in c["v"]], "R0": parse_hex(c["R0"]), "rr": parse_hex(c["rr"]), } route = {c["id"]: c for c in data["route"]} return kernel, route def kernel_reference(case, prec): a, b, cq, qq = coeffs(case["x"], case["c"], case["w"], case["v"], case["R0"], case["rr"]) status, root, info = classify(a, b, cq, prec, case["R0"], case["rr"]) dd = sum(xi * xi for xi in (case["x"][i] - case["c"][i] for i in range(3))) return { "status": status, "root": root, "a": a, "b": b, "c": cq, "qq": qq, "dd": dd, "R0": case["R0"], "rr": case["rr"], "a_zero": info["a_zero"], "disc_sign": info["disc_sign"], "domain_clipped": info["domain_clipped"], } def ulp_dec(cand: float) -> Decimal: return frac_to_dec(Fraction(math.ulp(cand))) def dec_of_float(f: float) -> Decimal: return frac_to_dec(Fraction(f), 300) def poly_resid(a: Fraction, b: Fraction, c: Fraction, s: Fraction) -> Fraction: return a * s * s + b * s + c def poly_scale(a: Fraction, b: Fraction, c: Fraction, s: Fraction) -> Fraction: return abs(a) * s * s + abs(b) * s + abs(c) def analyze_kernel_variant(rows, refs, prec): """Return per-case records for one variant's kernel CSV.""" out = [] for r in rows: cid = int(r["id"]) ref = refs[cid] status = int(r["status"]) cand = float.fromhex(r["sigma"]) if r["sigma"] not in ("", "nan") else float("nan") cand_fr = parse_hex(r["sigma"]) if r["sigma"] not in ("", "nan") else None ak = parse_hex(r["a"]) bk = parse_hex(r["b"]) ck = parse_hex(r["c"]) rec = { "id": cid, "category": r["category"], "variant_status": status, "ref_status": ref["status"], "a_zero_ref": ref["a_zero"], "a_zero_kernel": ak == 0, "a_sign_ref": _sign(ref["a"]), "a_sign_kernel": _sign(ak), "disc_sign_ref": ref["disc_sign"], "late": False, "near_linear": False, "near_boundary": False, "ill_conditioned": False, "domain_clipped": ref["domain_clipped"], "root_err_ulps": None, "root_rel_err": None, "resid_kernel_rel": None, "resid_ideal_rel": None, "baseline_resid_rel": None, } # conditioning flags if ref["status"] == "ENTER" and ref["root"] is not None: s = ref["root"] scale = max(abs(ref["R0"]), Fraction(1)) rec["late"] = bool(s > frac_to_dec(scale, prec) * (10 ** 6)) rec["near_linear"] = abs(ref["a"]) <= Fraction(1, 10**10) * max( ref["qq"], ref["rr"] * ref["rr"], Fraction(1) ) r0sq = ref["R0"] * ref["R0"] dd = ref["dd"] rec["near_boundary"] = bool( dd + r0sq != 0 and abs(ref["c"]) <= Fraction(1, 10**8) * (dd + r0sq) ) rec["ill_conditioned"] = bool( rec["near_linear"] or rec["late"] or rec["near_boundary"] or r["category"].startswith("growing_out") or r["category"] == "shrinking" ) # root-level comparison if status == 1 and ref["status"] == "ENTER" and ref["root"] is not None and cand_fr is not None: err = abs(frac_to_dec(cand_fr, prec) - ref["root"]) u = ulp_dec(cand) if u > 0: rec["root_err_ulps"] = float(err / u) if ref["root"] != 0: rec["root_rel_err"] = float(err / abs(ref["root"])) rec["resid_kernel_rel"] = _rel( poly_resid(ak, bk, ck, cand_fr), ak, bk, ck, cand_fr ) rec["resid_ideal_rel"] = _rel( poly_resid(ref["a"], ref["b"], ref["c"], cand_fr), ref["a"], ref["b"], ref["c"], cand_fr, ) try: nearest = float(ref["root"]) nfr = Fraction(nearest) rec["baseline_resid_rel"] = _rel( poly_resid(ref["a"], ref["b"], ref["c"], nfr), ref["a"], ref["b"], ref["c"], nfr, ) except (OverflowError, ValueError): rec["baseline_resid_rel"] = None out.append(rec) return out def _sign(fr: Fraction) -> int: return (fr > 0) - (fr < 0) def _rel(resid: Fraction, a, b, c, s) -> float: scale = poly_scale(a, b, c, s) if scale == 0: return 0.0 return float(abs(resid) / scale) def route_reference(row, prec): x = [parse_hex(row[f"canonx{i}"]) for i in range(3)] w = [parse_hex(row[f"canonw{i}"]) for i in range(3)] v = [parse_hex(row[f"v{i}"]) for i in range(3)] center = [parse_hex(row[f"ct0{i}"]) for i in range(3)] rr = parse_hex(row["rr"]) R0 = parse_hex(row["rt0"]) if row["rt0"] not in ("", "nan") else parse_hex(row["R0"]) a, b, cq, qq = coeffs(x, center, w, v, R0, rr) status, root, info = classify(a, b, cq, prec, R0, rr) return status, root, a, b, cq, qq, info["domain_clipped"] STATUS_NAME = {-1: "INVALID", 0: "MISS", 1: "ENTRY", 2: "UNCERTAIN"} KIND_NAME = { 0: "INSIDE", 1: "ENTRY", 2: "ESCAPED", 3: "TIME_RANGE_EXHAUSTED", 4: "INVALID", } def analyze_route_variant(rows, prec, dbl_eps=2.220446049250313e-16): out = [] for r in rows: (ref_status, ref_root, a, b, cq, qq, ref_domain_clipped) = route_reference(r, prec) status = int(r["status"]) kind = int(r["kind"]) F = _maybe_dec(r["F"]) tol = _maybe_dec(r["tol"]) kern_ok = r.get("kern_ok", "0") == "1" kern_status = int(r["kern_status"]) if kern_ok else None rec = { "id": int(r["id"]), "category": r["category"], "ref_status": ref_status, "ref_domain_clipped": ref_domain_clipped, "status": status, "kind": kind, "kind_name": KIND_NAME.get(kind, "?"), "failure_reason": int(r["failure_reason"]), "fallback": int(r["fallback_evals"]), "pi_match": int(r["pi_match"]), "lcam_match": int(r["lcam_match"]), "kern_status": kern_status if kern_ok else "", "kern_sigma": r.get("kern_sigma", ""), "F_over_tol": None, "inside_mismatch": False, "false_miss": False, "false_candidate": False, "unconfirmed": kind == 4 and r["reason_name"] == "ENTRY_UNCONFIRMED", "kernel_false_miss": False, "kernel_false_miss_escaped": False, } if ref_status == "INSIDE": if kind != 0: rec["inside_mismatch"] = True elif ref_status == "ENTER": if kind == 2: rec["false_miss"] = True elif ref_status == "MISS": if kind == 1: rec["false_candidate"] = True if ref_status == "ENTER" and kern_ok and kern_status == 0: rec["kernel_false_miss"] = True if kind == 2: rec["kernel_false_miss_escaped"] = True if F is not None and tol is not None and tol > 0: rec["F_over_tol"] = float(F / tol) out.append(rec) return out def _maybe_dec(s): if s in ("", "nan"): return None low = s.lower() if "inf" in low: return None return hex_to_dec(s, 300) def summarize_kernel(records_by_variant, refs): """Aggregate counts, ULP distributions, and common-ENTRY intersections.""" variants = list(records_by_variant.keys()) by_id = {v: {} for v in variants} for v in variants: for rec in records_by_variant[v]: by_id[v][rec["id"]] = rec ids = sorted(refs.keys()) summary = {"variants": {}} for v in variants: recs = by_id[v] counts = { "enter": 0, "miss": 0, "uncertain": 0, "false_miss": 0, "false_uncertain": 0, "false_candidate": 0, "a_zero_ref": 0, "a_zero_kernel": 0, "a_sign_mismatch": 0, "a_zero_collapse": 0, "a_spurious_nonzero": 0, "domain_clipped_candidate": 0, "ill_conditioned_enter": 0, "near_boundary_enter": 0, "domain_clipped_cases": 0, "late_ref_enter": 0, } ulps = [] rels = [] for cid in ids: rec = recs[cid] st = rec["variant_status"] if st == 1: counts["enter"] += 1 elif st == 0: counts["miss"] += 1 elif st == 2: counts["uncertain"] += 1 if rec["ref_status"] == "ENTER" and st == 0: counts["false_miss"] += 1 if rec["ref_status"] == "ENTER" and st == 2: counts["false_uncertain"] += 1 if rec["ref_status"] == "MISS" and st == 1: if rec["domain_clipped"]: counts["domain_clipped_candidate"] += 1 else: counts["false_candidate"] += 1 if rec["a_zero_ref"]: counts["a_zero_ref"] += 1 if rec["a_zero_kernel"]: counts["a_zero_kernel"] += 1 if not rec["a_zero_ref"] and rec["a_zero_kernel"]: counts["a_zero_collapse"] += 1 if rec["a_zero_ref"] and not rec["a_zero_kernel"]: counts["a_spurious_nonzero"] += 1 if rec["ill_conditioned"] and rec["ref_status"] == "ENTER": counts["ill_conditioned_enter"] += 1 if rec["near_boundary"] and rec["ref_status"] == "ENTER": counts["near_boundary_enter"] += 1 if rec["domain_clipped"]: counts["domain_clipped_cases"] += 1 if (rec["a_sign_ref"] != 0 and rec["a_sign_kernel"] != 0 and rec["a_sign_ref"] != rec["a_sign_kernel"]): counts["a_sign_mismatch"] += 1 if rec["late"]: counts["late_ref_enter"] += 1 if rec["root_err_ulps"] is not None: ulps.append(rec["root_err_ulps"]) if rec["root_rel_err"] is not None: rels.append(rec["root_rel_err"]) summary["variants"][v] = { "counts": counts, "root_ulp_all_ref_enter": _dist(ulps), "root_rel_all_ref_enter": _dist(rels), } # Common reference-ENTER subset that every variant solved as ENTRY. common = [] for cid in ids: if refs[cid]["status"] != "ENTER": continue if all(by_id[v][cid]["variant_status"] == 1 for v in variants): common.append(cid) summary["common_enter_ids"] = len(common) for v in variants: vals = [] vals_nonlate = [] rels_nonlate = [] for cid in common: rec = by_id[v][cid] e = rec["root_err_ulps"] if e is not None: vals.append(e) if not rec["ill_conditioned"]: vals_nonlate.append(e) if rec["root_rel_err"] is not None: rels_nonlate.append(rec["root_rel_err"]) summary["variants"][v]["root_ulp_common_enter"] = _dist(vals) summary["variants"][v]["root_ulp_common_enter_wellcond"] = _dist(vals_nonlate) summary["variants"][v]["root_rel_common_enter_wellcond"] = _dist(rels_nonlate) return summary def _dist(vals): if not vals: return {"n": 0} vs = sorted(vals) n = len(vs) def pct(p): idx = min(n - 1, max(0, int(math.ceil(p * n)) - 1)) return vs[idx] return { "n": n, "min": vs[0], "median": pct(0.5), "p95": pct(0.95), "max": vs[-1], } def summarize_route(records_by_variant): summary = {"variants": {}} for v, recs in records_by_variant.items(): counts = { "inside_ok": 0, "inside_mismatch": 0, "false_miss": 0, "false_candidate": 0, "unconfirmed": 0, "fallback_used": 0, "accepted_entry": 0, "F_over_tol_gt1": 0, "pi_mismatch": 0, "lcam_mismatch": 0, "kernel_false_miss": 0, "kernel_false_miss_escaped": 0, "ref_domain_clipped": 0, } fmax = 0.0 by_cat = {} for rec in recs: cat = by_cat.setdefault(rec["category"], {"cases": 0, "fallback": 0, "entry": 0, "escaped": 0}) cat["cases"] += 1 if rec["kind"] == 0 and rec["ref_status"] == "INSIDE": counts["inside_ok"] += 1 if rec["inside_mismatch"]: counts["inside_mismatch"] += 1 if rec["false_miss"]: counts["false_miss"] += 1 if rec["false_candidate"]: counts["false_candidate"] += 1 if rec["unconfirmed"]: counts["unconfirmed"] += 1 if rec["kernel_false_miss"]: counts["kernel_false_miss"] += 1 if rec["kernel_false_miss_escaped"]: counts["kernel_false_miss_escaped"] += 1 if rec["ref_domain_clipped"]: counts["ref_domain_clipped"] += 1 if rec["fallback"] > 0: counts["fallback_used"] += 1 cat["fallback"] += 1 if rec["kind"] == 1: counts["accepted_entry"] += 1 cat["entry"] += 1 if rec["F_over_tol"] is not None: fmax = max(fmax, rec["F_over_tol"]) if rec["F_over_tol"] > 1.0: counts["F_over_tol_gt1"] += 1 if rec["kind"] == 2: cat["escaped"] += 1 if not rec["pi_match"]: counts["pi_mismatch"] += 1 if not rec["lcam_match"]: counts["lcam_mismatch"] += 1 summary["variants"][v] = {"counts": counts, "max_F_over_tol": fmax, "by_category": by_cat} return summary def precision_consistency(cases, ids, prec_a, prec_b): """Compare reference classification and nearest-double root at two Decimal precisions. Coefficients/discriminant are exact rationals, so only the square-root precision can differ. Returns a list of mismatches.""" mismatches = [] for cid in sorted(ids): ra = kernel_reference(cases[cid], prec_a) rb = kernel_reference(cases[cid], prec_b) reason = None if ra["status"] != rb["status"]: reason = "status" elif ra["disc_sign"] != rb["disc_sign"]: reason = "disc_sign" elif ra["status"] == "ENTER": try: fa = float(ra["root"]) fb = float(rb["root"]) except (OverflowError, ValueError): reason = "root_unrepresentable" else: if not (fa == fb or (math.isnan(fa) and math.isnan(fb))): reason = "nearest_double_root" if reason is not None: mismatches.append({"id": cid, "reason": reason, "status_a": ra["status"], "status_b": rb["status"]}) return mismatches