Files
wyj 7c980c35aa 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.
2026-10-09 00:15:02 -04:00

459 lines
20 KiB
Python

#!/usr/bin/env python3
"""End-to-end driver for the quadratic precision/performance benchmark.
Self-contained: reads the current working-tree ``src/asymptotic.c``, generates
and builds four variants, runs the accuracy probe, evaluates the exact-decimal
reference, and times the kernel and the public pre-route. All raw artifacts
are written under ``--output-dir`` (default /tmp/opencode/quadratic-comparison).
No production source, Makefile or git state is modified.
"""
from __future__ import annotations
import argparse
import csv
import json
import os
import re
import subprocess
import sys
# Do not leave __pycache__ inside the repository benchmark directory: the
# imported local modules (build/cases/oracle) must not be cached here.
sys.dont_write_bytecode = True
from pathlib import Path
HERE = Path(__file__).resolve().parent
sys.path.insert(0, str(HERE))
import build as buildmod # noqa: E402
import cases as casesmod # noqa: E402
import oracle # noqa: E402
VARIANTS = list(buildmod.VARIANTS)
def run_cmd(cmd, log_path: Path, env=None):
with open(log_path, "a") as fh:
fh.write("$ " + " ".join(str(c) for c in cmd) + "\n")
proc = subprocess.run(cmd, capture_output=True, text=True, check=False, env=env)
with open(log_path, "a") as fh:
if proc.stdout:
fh.write(proc.stdout)
if proc.stderr:
fh.write(proc.stderr)
fh.write(f"exit={proc.returncode}\n")
return proc
def read_csv(path: Path):
with open(path, newline="") as fh:
return list(csv.DictReader(fh))
def run_timing(cmd, outp: Path, log_path: Path):
# Never ingest output left by a previous invocation, even after a failure.
outp.unlink(missing_ok=True)
proc = run_cmd(cmd, log_path)
if proc.returncode != 0 or not outp.is_file():
raise RuntimeError(f"timing failed or produced no fresh output: {log_path}")
return json.loads(outp.read_text())
def write_metrics_csv(path: Path, records, columns):
with open(path, "w", newline="") as fh:
w = csv.DictWriter(fh, fieldnames=columns, extrasaction="ignore")
w.writeheader()
for r in records:
w.writerow(r)
def write_curated(path: Path, rows, columns):
with open(path, "w", newline="") as fh:
w = csv.DictWriter(fh, fieldnames=columns, extrasaction="ignore")
w.writeheader()
for r in rows:
w.writerow(r)
def build_curated_kernel(kernel_records):
ref_ids = [r["id"] for r in kernel_records[VARIANTS[0]] if r["id"] < 1000]
by = {v: {r["id"]: r for r in kernel_records[v]} for v in VARIANTS}
rows = []
for cid in sorted(ref_ids):
row = {"id": cid, "category": by[VARIANTS[0]][cid]["category"],
"ref": by[VARIANTS[0]][cid]["ref_status"]}
for v in VARIANTS:
row[f"{v}_status"] = by[v][cid]["variant_status"]
row[f"{v}_ulp"] = by[v][cid]["root_err_ulps"]
row[f"{v}_rel"] = by[v][cid]["root_rel_err"]
rows.append(row)
return rows
def build_curated_route(route_records):
ref_ids = [r["id"] for r in route_records[VARIANTS[0]] if r["id"] < 100]
by = {v: {r["id"]: r for r in route_records[v]} for v in VARIANTS}
rows = []
for cid in sorted(ref_ids):
row = {"id": cid, "category": by[VARIANTS[0]][cid]["category"],
"ref": by[VARIANTS[0]][cid]["ref_status"]}
for v in VARIANTS:
row[f"{v}_kind"] = by[v][cid]["kind"]
row[f"{v}_fallback"] = by[v][cid]["fallback"]
row[f"{v}_F_over_tol"] = by[v][cid]["F_over_tol"]
rows.append(row)
return rows
def main():
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--repo", default=str(HERE.parents[1]))
ap.add_argument("--output-dir", default="/tmp/opencode/quadratic-comparison")
ap.add_argument("--cc", default=os.environ.get("CC", "cc"))
ap.add_argument("--threads", type=int, default=4)
ap.add_argument("--precision", type=int, default=160)
ap.add_argument("--rounds", type=int, default=3)
ap.add_argument("--kernel-target-calls", type=int, default=10_000_000)
ap.add_argument("--route-target-calls", type=int, default=1_000_000)
ap.add_argument("--skip-build", action="store_true")
args = ap.parse_args()
repo = Path(args.repo).resolve()
out = Path(args.output_dir).resolve()
(out / "raw").mkdir(parents=True, exist_ok=True)
(out / "cases").mkdir(parents=True, exist_ok=True)
(out / "logs").mkdir(parents=True, exist_ok=True)
import time as _time
t_start = _time.time()
phase_times = {}
# ---------------------------------------------------------------- cases
t0 = _time.time()
kernel_cases, route_cases = casesmod.build()
casesmod.write_header(kernel_cases, route_cases, out / "cases" / "quadratic_cases.h")
casesmod.write_json(kernel_cases, route_cases, out / "cases" / "cases.json")
phase_times["case_gen"] = _time.time() - t0
print(f"cases: kernel={len(kernel_cases)} route={len(route_cases)}")
# ------------------------------------------------------------ environment
env_info = buildmod.environment(repo, args.cc, args.threads)
env_info["python_version"] = sys.version.split()[0]
env_info["invocation"] = " ".join([sys.executable, *sys.argv])
(out / "environment.json").write_text(json.dumps(env_info, indent=2) + "\n")
(out / "raw" / "invocation.txt").write_text(
" ".join([sys.executable, *sys.argv]) + "\n"
)
print(f"source sha256: {env_info['source_sha256'][:16]} cc: {env_info['cc_version']}")
# ---------------------------------------------------------------- build
t0 = _time.time()
if not args.skip_build:
info = buildmod.build_all(repo, out, args.cc, buildmod.BASE_FLAGS, args.threads)
exes = {k: Path(v) for k, v in info["executables"].items()}
else:
manifest, reasons = buildmod.verify_manifest(repo, out, buildmod.BASE_FLAGS, args.cc)
if reasons:
print("refusing --skip-build: existing build does not match the "
f"current source/flags/fixtures ({', '.join(reasons)}); "
"rerun without --skip-build", file=sys.stderr)
sys.exit(1)
exes = {k: Path(v) for k, v in manifest["executables"].items()}
phase_times["build"] = _time.time() - t0
for name, exe in exes.items():
if not exe.exists():
print(f"missing executable for {name}: {exe}", file=sys.stderr)
sys.exit(1)
# -------------------------------------------------------------- accuracy
t0 = _time.time()
env_log = out / "logs" / "probe_environment.log"
env_log.write_text("")
for name in VARIANTS:
kout = out / "raw" / f"kernel_{name}.csv"
rout = out / "raw" / f"route_{name}.csv"
proc = run_cmd(
[str(exes[name]), "accuracy", "--kernel-out", str(kout),
"--route-out", str(rout)],
out / "logs" / f"accuracy_{name}.log",
)
with open(env_log, "a") as fh:
fh.write(proc.stdout)
if proc.returncode != 0:
print(f"accuracy probe failed for {name}", file=sys.stderr)
sys.exit(1)
phase_times["accuracy"] = _time.time() - t0
# ---------------------------------------------------------------- oracle
t0 = _time.time()
kernel_ref_cases, route_cases_json = oracle.load_cases(out / "cases" / "cases.json")
refs = {
cid: oracle.kernel_reference(case, args.precision)
for cid, case in kernel_ref_cases.items()
}
kernel_records = {}
route_records = {}
for name in VARIANTS:
krows = read_csv(out / "raw" / f"kernel_{name}.csv")
kernel_records[name] = oracle.analyze_kernel_variant(krows, refs, args.precision)
rrows = read_csv(out / "raw" / f"route_{name}.csv")
route_records[name] = oracle.analyze_route_variant(rrows, args.precision)
kcols = ["id", "category", "ref_status", "variant_status", "root_err_ulps",
"root_rel_err",
"resid_kernel_rel", "resid_ideal_rel", "baseline_resid_rel",
"a_zero_ref", "a_zero_kernel", "a_sign_ref", "a_sign_kernel",
"disc_sign_ref", "near_linear", "near_boundary",
"ill_conditioned", "domain_clipped", "late"]
for name in VARIANTS:
write_metrics_csv(out / "raw" / f"kernel_metrics_{name}.csv",
kernel_records[name], kcols)
rcols = ["id", "category", "ref_status", "ref_domain_clipped", "status",
"kind", "kind_name", "failure_reason", "fallback", "pi_match",
"lcam_match", "kern_status", "kern_sigma", "F_over_tol",
"inside_mismatch", "false_miss", "false_candidate", "unconfirmed",
"kernel_false_miss", "kernel_false_miss_escaped"]
for name in VARIANTS:
write_metrics_csv(out / "raw" / f"route_metrics_{name}.csv",
route_records[name], rcols)
ksummary = oracle.summarize_kernel(kernel_records, refs)
rsummary = oracle.summarize_route(route_records)
consistency_ids = {
rec["id"] for rec in kernel_records[VARIANTS[0]]
if rec["id"] < 1000 or rec["near_linear"] or rec["late"]
or rec["near_boundary"] or rec["domain_clipped"]
}
precision_mismatches = oracle.precision_consistency(
kernel_ref_cases, consistency_ids, args.precision, 240
)
curated_kernel = build_curated_kernel(kernel_records)
curated_route = build_curated_route(route_records)
write_curated(out / "raw" / "curated_kernel.csv", curated_kernel,
["id", "category", "ref"] + [f"{v}_{f}" for v in VARIANTS
for f in ("status", "ulp", "rel")])
write_curated(out / "raw" / "curated_route.csv", curated_route,
["id", "category", "ref"] + [f"{v}_{f}" for v in VARIANTS
for f in ("kind", "fallback",
"F_over_tol")])
# ------------------------------------------------------------- hardware FMA
phase_times["oracle"] = _time.time() - t0
fma_info = verify_fma(out, args.threads)
# ---------------------------------------------------------------- timing
t0 = _time.time()
timings = timed_runs(exes, out, args, kernel_target_calls=args.kernel_target_calls,
route_target_calls=args.route_target_calls)
lin_target = max(100_000, args.kernel_target_calls // 20)
timings["linearity"] = linearity_check(exes, out, lin_target)
phase_times["timing"] = _time.time() - t0
# ---------------------------------------------------------------- summary
summary = {
"source_sha256": env_info["source_sha256"],
"flags": buildmod.BASE_FLAGS,
"precision": args.precision,
"precision_consistency": {
"ids_checked": len(consistency_ids),
"mismatches": precision_mismatches,
},
"threads_timing": args.threads,
"cases": {"kernel": len(kernel_cases), "route": len(route_cases)},
"kernel_accuracy": ksummary,
"route_accuracy": rsummary,
"curated_kernel": curated_kernel,
"fma": fma_info,
"timing": timings,
"phase_seconds": phase_times,
"total_seconds": _time.time() - t_start,
}
(out / "summary.json").write_text(json.dumps(summary, indent=2) + "\n")
print_summary(summary)
print(f"\nartifacts under {out}")
def verify_fma(out: Path, threads: int):
info = {"objects": {}}
for name in VARIANTS:
obj = out / "build" / f"probe_{name}.o"
if not obj.exists():
continue
proc = subprocess.run(["objdump", "-d", str(obj)], capture_output=True,
text=True, check=False)
n = sum(1 for line in proc.stdout.splitlines()
if "vfmadd" in line or "vfmsub" in line or "fmadd" in line)
nm = subprocess.run(["nm", "-u", str(obj)], capture_output=True,
text=True, check=False)
undef = sorted({tok for line in nm.stdout.splitlines()
for tok in line.split()
if tok.startswith("fma")})
kd = subprocess.run(["objdump", "-dr", str(obj)], capture_output=True,
text=True, check=False)
kb_lines = []
inside = False
for line in kd.stdout.splitlines():
if re.match(r"^[0-9a-f]+ <entry_solve", line):
inside = True
kb_lines.append(line)
continue
if inside:
if re.match(r"^[0-9a-f]+ <", line):
break
kb_lines.append(line)
kb = "\n".join(kb_lines)
x87 = sum(1 for line in kb.splitlines()
if re.search(r"\b(fld|fstp|fmul|fadd|fsub|fdiv|fcom)\b", line))
vfma = sum(1 for line in kb.splitlines()
if "vfmadd" in line or "vfmsub" in line)
calls_fmal = sum(1 for line in kb.splitlines() if "fmal" in line)
info["objects"][name] = {
"fma_instructions": n,
"lowered_hardware": n > 0,
"undefined_fma_symbols": undef,
"entry_solve_x87": x87,
"entry_solve_vfmadd": vfma,
"entry_solve_fmal_calls": calls_fmal,
}
print(f"arith verify {name}: fused={n} undef={undef} "
f"entry_solve[x87={x87} vfmadd={vfma} fmal_calls={calls_fmal}]")
(out / "raw" / "fma_verify.json").write_text(json.dumps(info, indent=2) + "\n")
return info
def linearity_check(exes, out: Path, target_calls: int):
"""Confirm the microbench time scales with the call count (no hoisting)."""
res = {}
for name in ("ld_fma", "double_fma"):
secs = []
for mult in (1, 2):
outp = out / "raw" / f"linearity_{name}_{mult}.json"
d = run_timing([str(exes[name]), "microbench", "--out", str(outp),
"--target-calls", str(target_calls * mult)],
outp, out / "logs" / f"linearity_{name}_{mult}.log")
secs.append(d["seconds"])
res[name] = {"t1": secs[0], "t2": secs[1],
"ratio": secs[1] / secs[0] if secs[0] > 0 else None}
return res
def timed_runs(exes, out: Path, args, kernel_target_calls, route_target_calls):
schedules = []
fwd = list(VARIANTS)
schedules.append(fwd)
schedules.append(list(reversed(fwd)))
for i in range(2, args.rounds):
schedules.append(fwd if i % 2 == 0 else list(reversed(fwd)))
schedules = schedules[: args.rounds]
results = {"microbench": [], "routebench": []}
def order(seq):
return [exes[n] for n in seq]
for rnd, seq in enumerate(schedules):
for name in seq:
outp = out / "raw" / f"microbench_{name}_r{rnd}.json"
d = run_timing([str(exes[name]), "microbench", "--out", str(outp),
"--target-calls", str(kernel_target_calls)],
outp, out / "logs" / f"microbench_{name}_r{rnd}.log")
d["round"] = rnd
results["microbench"].append(d)
for rnd, seq in enumerate(schedules):
for name in seq:
outp = out / "raw" / f"routebench_{name}_r{rnd}.json"
d = run_timing([str(exes[name]), "routebench", "--out", str(outp),
"--target-calls", str(route_target_calls),
"--threads", str(args.threads), "--max-seconds", "20"],
outp, out / "logs" / f"routebench_{name}_r{rnd}.log")
d["round"] = rnd
results["routebench"].append(d)
return results
def print_summary(summary):
print("\n=== kernel accuracy (per variant) ===")
print(f"{'variant':<18}{'ENTER':>7}{'MISS':>7}{'UNC':>6}{'falseMiss':>10}"
f"{'falseCand':>10}{'domClipCand':>12}{'a0ref':>7}{'a0k':>6}{'asign':>6}"
f"{'illcond':>8}{'nearBnd':>8}{'medULP':>10}{'p95ULP':>11}{'maxULP':>11}")
for v, s in summary["kernel_accuracy"]["variants"].items():
c = s["counts"]
d = s["root_ulp_all_ref_enter"]
print(f"{v:<18}{c['enter']:>7}{c['miss']:>7}{c['uncertain']:>6}"
f"{c['false_miss']:>10}{c['false_candidate']:>10}"
f"{c['domain_clipped_candidate']:>12}"
f"{c['a_zero_ref']:>7}{c['a_zero_kernel']:>6}"
f"{c['a_sign_mismatch']:>6}{c['ill_conditioned_enter']:>8}"
f"{c['near_boundary_enter']:>8}"
f"{d.get('median', float('nan')):>10.4g}{d.get('p95', float('nan')):>11.4g}"
f"{d.get('max', float('nan')):>11.4g}")
print("well-conditioned common-ENTER ULP (excludes late / near-linear /"
" growing-out / shrinking):")
for v, s in summary["kernel_accuracy"]["variants"].items():
d = s["root_ulp_common_enter_wellcond"]
r = s["root_rel_common_enter_wellcond"]
print(f" {v:<18}n={d.get('n',0):>5} median={d.get('median', float('nan')):>9.3g}"
f" p95={d.get('p95', float('nan')):>9.3g}"
f" max={d.get('max', float('nan')):>9.3g}"
f" |rel err median={r.get('median', float('nan')):>9.3g}"
f" p95={r.get('p95', float('nan')):>9.3g}")
print(f"common reference-ENTER ids solved ENTRY by all variants: "
f"{summary['kernel_accuracy']['common_enter_ids']}")
print("\n=== curated kernel cases (status/ULP) ===")
print(f"{'id':>4} {'category':<20} {'ref':<6} "
+ " ".join(f"{v:>16}" for v in VARIANTS))
for row in summary.get("curated_kernel", []):
cells = []
for v in VARIANTS:
st = row.get(f"{v}_status")
u = row.get(f"{v}_ulp")
cells.append(f"{st}/{float(u):.3g}" if u not in (None, "") else f"{st}/-")
print(f"{row['id']:>4} {row['category']:<20} {row.get('ref',''):<6} "
+ " ".join(f"{c:>16}" for c in cells))
print("\n=== route accuracy (per variant) ===")
print(f"{'variant':<18}{'insideOK':>9}{'insideBad':>10}{'falseMiss':>10}"
f"{'falseCand':>10}{'unconf':>8}{'kernFM':>8}{'kernFMesc':>10}"
f"{'fallback':>9}{'F>tol':>7}{'maxF/tol':>10}")
for v, s in summary["route_accuracy"]["variants"].items():
c = s["counts"]
print(f"{v:<18}{c['inside_ok']:>9}{c['inside_mismatch']:>10}"
f"{c['false_miss']:>10}{c['false_candidate']:>10}"
f"{c['unconfirmed']:>8}{c['kernel_false_miss']:>8}"
f"{c['kernel_false_miss_escaped']:>10}"
f"{c['fallback_used']:>9}"
f"{c['F_over_tol_gt1']:>7}{s['max_F_over_tol']:>10.3g}")
pc = summary.get("precision_consistency", {})
print(f"reference precision {summary['precision']} vs 240: checked "
f"{pc.get('ids_checked')} curated/near-linear/late ids, "
f"mismatches={len(pc.get('mismatches', []))}")
print("\nroute fallback / entry / escaped by category "
"(ld_plain | ld_fma | double_fma):")
ra = summary["route_accuracy"]["variants"]
cats = sorted({c for s in ra.values() for c in s.get("by_category", {})})
for cat in cats:
cells = []
for v in ("ld_plain", "ld_fma", "double_fma"):
d = ra[v].get("by_category", {}).get(cat, {})
cells.append(f"fb={d.get('fallback',0)} e={d.get('entry',0)} "
f"x={d.get('escaped',0)}")
print(f" {cat:<20} " + " | ".join(cells))
print("\n=== timing (median ns/call over rounds) ===")
for mode in ("microbench", "routebench"):
per = {}
for rec in summary["timing"][mode]:
per.setdefault(rec["variant"], []).append(rec["ns_per_call"])
for v, vals in per.items():
vals.sort()
med = vals[len(vals) // 2]
print(f"{mode:<12}{v:<18}median={med:>10.2f} ns/call "
f"rounds={len(vals)}")
lin = summary["timing"].get("linearity", {})
for v, d in lin.items():
print(f"linearity {v:<18}t1={d['t1']:.4f}s t2={d['t2']:.4f}s "
f"ratio={d['ratio']:.3f} (expect ~2)")
if __name__ == "__main__":
main()