#!/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]+ 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()