Benchmark: Compare quadratic entry precision and arithmetic cost

Generate plain long-double and experimental FMA variants from the production entry kernel. Compare exact-rational references, geometric validation, fallback counts and repeated timings with self-contained fixtures.

Validate cached build dependencies and reject stale timing output. Count all unconfirmed entry outcomes independently of reference classification.
This commit is contained in:
wyj committed 2026-10-09 00:15:02 -04:00
1 parent 789549ed2f
commit 7c980c35aa
6 files changed
+2805

No files matched your search

+571
View File
@@ -0,0 +1,571 @@
#!/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