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:
1 parent
789549ed2f
commit
7c980c35aa
6 files changed
+2805
No files matched your search
@@ -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
|
||||
Reference in new issue
Block a user