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

490 lines
18 KiB
Python

#!/usr/bin/env python3
"""Generate the four precision variants of the asymptotic entry kernel and build
one self-contained probe executable per variant.
This module never edits the production tree. It reads the *current working
tree* ``src/asymptotic.c``, extracts the "entry kernel" region delimited by two
lexical anchors, emits four generated full-module copies under the scratch
directory, and compiles a probe that textually includes exactly one copy so the
private static ``entry_solve``/``entry_quadratic_coeffs`` are exercised as the
real generated code (not an independent toy copy).
Variants (all share the same stable-entry algebra, scaling, uncertainty policy,
validation and fallback; only the arithmetic under test changes):
ld_plain long double coefficients/solver, ordinary b*b-4ac and slope
ld_fma reconstructed experimental compensated fma disc/slope
double_fma double type/solver with compensated fma disc/slope,
ordinary coefficient accumulation (type-precision control)
double_fma_coeff double like double_fma plus fma coefficient accumulation
Only the kernel region is rewritten. Every other production dispatch/fallback
path is copied verbatim.
"""
from __future__ import annotations
import hashlib
import json
import os
import re
import shutil
import subprocess
import sys
from pathlib import Path
# ---------------------------------------------------------------------------
# Fixed compiler configuration. Identical for all four variants so the
# comparison is not confounded by flags. ``-ffp-contract=off`` disables the
# implicit FMA contraction GCC enables at -O2 for this target, so the ld_plain
# arm cannot silently gain extra precision; explicit fma()/fmal() calls still
# lower to hardware FMA.
# ---------------------------------------------------------------------------
BASE_FLAGS = [
"-std=c11",
"-march=native",
"-O2",
"-DNDEBUG",
"-ffp-contract=off",
"-fopenmp",
]
VARIANTS = ("ld_plain", "ld_fma", "double_fma", "double_fma_coeff")
# Common production translation units needed by the probe. Chosen as the
# minimal dependency closure of the entry kernel plus the public pre-route:
# geodesic initialisation, the numerical entry localizer, the Schwarzschild
# exterior referenced by the full module, the dispatch wrapper and the
# Minkowski provider. main/asymptotic/frame/catalog/PSF/FFTW are excluded.
COMMON_SOURCES = (
"src/geodesic.c",
"src/asymptotic_entry.c",
"src/asymptotic_schwarzschild.c",
"src/spacetime_common.c",
"src/spacetime_minkowski.c",
)
START_ANCHOR = "/* Long-double coefficients of the relative-distance quadratic"
FWD_PREFIX = "static void minkowski_route_entry("
END_ANCHOR = (
"static void minkowski_route_entry(const SpacetimeAsymptoticEnd *end, "
"double t0,"
)
class BuildError(RuntimeError):
pass
def _replace_once(text: str, old: str, new: str, what: str) -> str:
count = text.count(old)
if count != 1:
raise BuildError(f"expected exactly one occurrence of {what}, found {count}")
return text.replace(old, new, 1)
def locate_kernel(lines: list[str]) -> tuple[int, int]:
"""Return [start, end) line indices of the rewritten kernel region.
start is the 'Long-double coefficients' comment; end is the forward
declaration of ``minkowski_route_entry`` (preserved verbatim).
"""
starts = [i for i, l in enumerate(lines) if l.startswith(START_ANCHOR)]
if len(starts) != 1:
raise BuildError(
f"kernel start anchor must appear exactly once, found {len(starts)}"
)
start = starts[0]
ends = [
i
for i, l in enumerate(lines)
if i > start and l.startswith(FWD_PREFIX)
]
if not ends:
raise BuildError("kernel end anchor (minkowski_route_entry) not found")
end = ends[0]
if not lines[end].startswith(END_ANCHOR):
raise BuildError(
"kernel end anchor does not match the expected signature; "
"production source changed"
)
region = "".join(lines[start:end])
for needle in (
"typedef struct {",
"static EntryQuadratic entry_quadratic_coeffs(",
"static EntrySolveResult entry_solve(",
"static long double entry_discriminant(",
"EntryQuadratic;",
):
if region.count(needle) != 1:
raise BuildError(
f"kernel region must contain exactly one '{needle}', "
f"found {region.count(needle)}"
)
# The region must not contain the forward declaration or any later route
# code: those must be byte-for-byte preserved.
if "minkowski_preroute" in region or "AsymptoticStatus asymptotic_route" in region:
raise BuildError("kernel region overran into preserved route code")
return start, end
def transform_ld_plain(kernel: str) -> str:
"""Convert a compensated experimental kernel back to plain long double."""
old_disc = (
" const long double four_a = 4.0L * k->a;\n"
" const long double ac = four_a * k->c;\n"
" const long double ac_error = fmal(four_a, k->c, -ac);\n"
" *scale = k->b * k->b + fabsl(ac);\n"
" return fmal(k->b, k->b, -ac) - ac_error;\n"
)
new_disc = (
" const long double four_a = 4.0L * k->a;\n"
" const long double ac = four_a * k->c;\n"
" *scale = k->b * k->b + fabsl(ac);\n"
" return k->b * k->b - ac;\n"
)
out = _replace_once(kernel, old_disc, new_disc, "ld_plain discriminant body")
old_slope = "const long double slope = fmal(2.0L * k->a, s, k->b);"
new_slope = "const long double slope = 2.0L * k->a * s + k->b;"
out = _replace_once(out, old_slope, new_slope, "ld_plain slope")
return out
def _assert_double_clean(t: str) -> None:
"""The rewritten double kernel must contain no long-double promotion."""
forbidden = ("long double", "fmal(", "fabsl(", "fmaxl(", "frexpl(",
"scalbnl(", "sqrtl(", "copysignl(", "LDBL_")
for tok in forbidden:
if tok in t:
raise BuildError(f"double kernel still contains '{tok}'")
if re.search(r"(?<=[0-9])L(?![A-Za-z0-9_])", t):
raise BuildError("double kernel still contains an L-suffixed literal")
def transform_to_double(kernel: str) -> str:
"""Convert the kernel from long double to double arithmetic.
Explicit fmal->fma, *l->*, LDBL_*->DBL_*, long double->double and strips the
L suffix from the (few) decimal literals. Explicit fma calls are retained
(double_fma keeps compensated discriminant/slope).
"""
t = kernel
for old, new in (
("fmal(", "fma("),
("fabsl(", "fabs("),
("fmaxl(", "fmax("),
("frexpl(", "frexp("),
("scalbnl(", "scalbn("),
("sqrtl(", "sqrt("),
("copysignl(", "copysign("),
):
t = t.replace(old, new)
t = t.replace("LDBL_MIN", "DBL_MIN")
t = t.replace("LDBL_EPSILON", "DBL_EPSILON")
t = t.replace("long double", "double")
# Strip L suffixes from decimal literals; otherwise 4.0L/2.0L would
# re-promote the double expression and confound type accuracy/timing.
t = re.sub(r"(?<=[0-9])L(?![A-Za-z0-9_])", "", t)
_assert_double_clean(t)
return t
def transform_double_fma_coeff(kernel: str) -> str:
"""double variants plus fma coefficient accumulation and products."""
dbl = transform_to_double(kernel)
old_block = (
" double qq = 0.0, dot_dq = 0.0, dot_dd = 0.0;\n"
" for (int i = 0; i < 3; ++i) {\n"
" qq += q[i] * q[i];\n"
" dot_dq += d[i] * q[i];\n"
" dot_dd += d[i] * d[i];\n"
" }\n"
" const double R0_ld = (double)R0;\n"
" const double rr_ld = (double)rr;\n"
" EntryQuadratic k;\n"
" k.a = qq - rr_ld * rr_ld;\n"
" k.b = 2.0 * (dot_dq + R0_ld * rr_ld);\n"
" k.c = dot_dd - R0_ld * R0_ld;\n"
)
new_block = (
" double qq = 0.0, dot_dq = 0.0, dot_dd = 0.0;\n"
" for (int i = 0; i < 3; ++i) {\n"
" qq = fma(q[i], q[i], qq);\n"
" dot_dq = fma(d[i], q[i], dot_dq);\n"
" dot_dd = fma(d[i], d[i], dot_dd);\n"
" }\n"
" const double R0_ld = (double)R0;\n"
" const double rr_ld = (double)rr;\n"
" EntryQuadratic k;\n"
" k.a = fma(-rr_ld, rr_ld, qq);\n"
" k.b = 2.0 * fma(R0_ld, rr_ld, dot_dq);\n"
" k.c = fma(-R0_ld, R0_ld, dot_dd);\n"
)
return _replace_once(dbl, old_block, new_block, "double_fma_coeff block")
def generate_variants(repo: Path, out_dir: Path) -> dict:
"""Read current src/asymptotic.c and emit the four variant modules."""
src_path = repo / "src" / "asymptotic.c"
text = src_path.read_text()
lines = text.splitlines(keepends=True)
start, end = locate_kernel(lines)
kernel = "".join(lines[start:end])
prefix = "".join(lines[:start])
suffix = "".join(lines[end:])
sha = hashlib.sha256(text.encode()).hexdigest()
# Production now uses plain long-double arithmetic. Reconstruct the former
# compensated arm so the original four-way comparison remains reproducible.
plain_disc = (
" const long double four_a = 4.0L * k->a;\n"
" const long double ac = four_a * k->c;\n"
" *scale = k->b * k->b + fabsl(ac);\n"
" return k->b * k->b - ac;\n"
)
fused_disc = plain_disc.replace(
" *scale =", " const long double ac_error = fmal(four_a, k->c, -ac);\n *scale ="
).replace("return k->b * k->b - ac;", "return fmal(k->b, k->b, -ac) - ac_error;")
fused = _replace_once(kernel, plain_disc, fused_disc, "ld_fma discriminant body")
fused = _replace_once(
fused, "const long double slope = 2.0L * k->a * s + k->b;",
"const long double slope = fmal(2.0L * k->a, s, k->b);", "ld_fma slope"
)
kernels = {
"ld_plain": kernel,
"ld_fma": fused,
"double_fma": transform_to_double(fused),
"double_fma_coeff": transform_double_fma_coeff(fused),
}
gen_dir = out_dir / "generated"
gen_dir.mkdir(parents=True, exist_ok=True)
paths = {}
metas = {}
for name in VARIANTS:
body = kernels[name]
banner = (
f"/* GENERATED FILE - do not edit.\n"
f" * variant: {name}\n"
f" * source: src/asymptotic.c sha256={sha}\n"
f" * region lines [{start + 1}, {end}] rewritten; all other code verbatim.\n"
f" */\n"
)
out = banner + prefix + body + suffix
path = gen_dir / f"{name}.c"
path.write_text(out)
paths[name] = path
metas[name] = {
"path": str(path),
"sha256": hashlib.sha256(out.encode()).hexdigest(),
"kernel_sha256": hashlib.sha256(body.encode()).hexdigest(),
}
return {
"source": str(src_path),
"source_sha256": sha,
"kernel_line_range": [start + 1, end],
"variants": metas,
}
def cc_version(cc: str) -> str:
try:
out = subprocess.run(
[cc, "--version"], capture_output=True, text=True, check=False
)
return (out.stdout or out.stderr).splitlines()[0] if out.stdout or out.stderr else ""
except OSError:
return ""
def environment(repo: Path, cc: str, threads: int) -> dict:
src_path = repo / "src" / "asymptotic.c"
env = {
"repo": str(repo),
"cc": cc,
"cc_version": cc_version(cc),
"base_flags": list(BASE_FLAGS),
"threads": threads,
"source_sha256": hashlib.sha256(src_path.read_bytes()).hexdigest(),
"git_head": _git(repo, "rev-parse", "HEAD"),
"git_status_short": _git(repo, "status", "--short"),
"uname": os.uname().sysname + " " + os.uname().release + " " + os.uname().machine,
"cpu_model": _cpu_model(),
}
return env
def _git(repo: Path, *args: str) -> str:
try:
out = subprocess.run(
["git", "-C", str(repo), *args], capture_output=True, text=True, check=False
)
return out.stdout.strip()
except OSError:
return ""
def _cpu_model() -> str:
try:
for line in Path("/proc/cpuinfo").read_text().splitlines():
if line.startswith("model name"):
return line.split(":", 1)[1].strip()
except OSError:
pass
return "unknown"
def compile_common(repo: Path, out_dir: Path, cc: str, flags: list[str]) -> list[Path]:
obj_dir = out_dir / "build" / "common"
obj_dir.mkdir(parents=True, exist_ok=True)
objs = []
for rel in COMMON_SOURCES:
src = repo / rel
obj = obj_dir / (Path(rel).stem + ".o")
cmd = [cc, *flags, f"-I{repo / 'src'}", "-c", str(src), "-o", str(obj)]
_run(cmd, out_dir, f"compile common {rel}")
objs.append(obj)
return objs
def compile_probe(
repo: Path,
out_dir: Path,
cc: str,
flags: list[str],
variant: str,
probe_src: Path,
cases_dir: Path,
common_objs: list[Path],
) -> Path:
"""Compile the probe (which #includes the generated variant) and link."""
exe_dir = out_dir / "build" / "bin"
exe_dir.mkdir(parents=True, exist_ok=True)
obj = out_dir / "build" / f"probe_{variant}.o"
variant_src = out_dir / "generated" / f"{variant}.c"
cmd = [
cc,
*flags,
f"-I{repo / 'src'}",
f"-I{cases_dir}",
f'-DPROBE_VARIANT_SOURCE="{variant_src}"',
f'-DPROBE_VARIANT_NAME="{variant}"',
"-c",
str(probe_src),
"-o",
str(obj),
]
_run(cmd, out_dir, f"compile probe {variant}")
exe = exe_dir / f"probe_{variant}"
link = [cc, *flags, str(obj), *[str(o) for o in common_objs], "-lm", "-o", str(exe)]
_run(link, out_dir, f"link probe {variant}")
return exe
def _run(cmd: list[str], out_dir: Path, what: str) -> None:
log_dir = out_dir / "logs"
log_dir.mkdir(parents=True, exist_ok=True)
with open(log_dir / "build_commands.log", "a") as fh:
fh.write(what + "\n$ " + " ".join(cmd) + "\n")
try:
proc = subprocess.run(cmd, capture_output=True, text=True, check=False)
except OSError as exc:
raise BuildError(f"{what}: {exc}") from exc
with open(log_dir / "build_commands.log", "a") as fh:
fh.write(f"exit={proc.returncode}\n")
if proc.stdout:
fh.write(proc.stdout)
if proc.stderr:
fh.write(proc.stderr)
if proc.returncode != 0:
raise BuildError(
f"{what} failed (exit {proc.returncode}); see logs/build_commands.log\n"
+ (proc.stderr or "")[-4000:]
)
def build_all(repo: Path, out_dir: Path, cc: str, flags: list[str], threads: int) -> dict:
out_dir.mkdir(parents=True, exist_ok=True)
(out_dir / "logs").mkdir(parents=True, exist_ok=True)
variant_info = generate_variants(repo, out_dir)
common = compile_common(repo, out_dir, cc, flags)
probe_src = repo / "benchmarks" / "quadratic_precision" / "fixtures" / "probe.c"
cases_dir = out_dir / "cases"
exes = {}
for name in VARIANTS:
exes[name] = compile_probe(
repo, out_dir, cc, flags, name, probe_src, cases_dir, common
)
result = {"variants": variant_info["variants"],
"source_sha256": variant_info["source_sha256"],
"base_flags": list(flags),
"dependencies": dependency_hashes(repo),
"compiler": {"command": cc, "version": cc_version(cc)},
"probe_sha256": _file_sha(probe_src),
"cases_header_sha256": _file_sha(cases_dir / "quadratic_cases.h")}
result["executables"] = {k: str(v) for k, v in exes.items()}
result["common_objects"] = [str(o) for o in common]
(out_dir / "build_manifest.json").write_text(json.dumps(result, indent=2) + "\n")
return result
def _file_sha(path: Path):
if not Path(path).exists():
return None
return hashlib.sha256(Path(path).read_bytes()).hexdigest()
def dependency_hashes(repo: Path):
# Include all project headers conservatively, including transitive includes.
paths = {repo / rel for rel in COMMON_SOURCES}
paths.update((repo / "src").rglob("*.h"))
return {str(p.relative_to(repo)): _file_sha(p) for p in sorted(paths)}
def verify_manifest(repo: Path, out_dir: Path, flags: list[str], cc: str):
"""Check that an existing build matches the current source, flags and
fixtures. Returns (manifest, list_of_mismatch_reasons). Regenerates the
variant sources as a side effect so their hashes can be compared."""
manifest = json.loads((out_dir / "build_manifest.json").read_text())
reasons = []
if manifest.get("dependencies") != dependency_hashes(repo):
reasons.append("dependencies")
if manifest.get("compiler") != {"command": cc, "version": cc_version(cc)}:
reasons.append("compiler")
current = generate_variants(repo, out_dir)
if manifest.get("source_sha256") != current["source_sha256"]:
reasons.append("source_sha256")
if list(manifest.get("base_flags", [])) != list(flags):
reasons.append("base_flags")
for name, meta in current["variants"].items():
if manifest.get("variants", {}).get(name, {}).get("sha256") != meta["sha256"]:
reasons.append(f"variant:{name}")
probe_src = repo / "benchmarks" / "quadratic_precision" / "fixtures" / "probe.c"
if manifest.get("probe_sha256") != _file_sha(probe_src):
reasons.append("probe_sha256")
cases_header = out_dir / "cases" / "quadratic_cases.h"
if cases_header.exists() and manifest.get("cases_header_sha256") != _file_sha(cases_header):
reasons.append("cases_header_sha256")
for name, exe in manifest.get("executables", {}).items():
if not Path(exe).exists():
reasons.append(f"missing_exe:{name}")
return manifest, reasons
if __name__ == "__main__":
import argparse
ap = argparse.ArgumentParser(description=__doc__)
ap.add_argument("--repo", default=str(Path(__file__).resolve().parents[2]))
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)
args = ap.parse_args()
repo = Path(args.repo).resolve()
out = Path(args.output_dir).resolve()
try:
info = build_all(repo, out, args.cc, BASE_FLAGS, args.threads)
except BuildError as exc:
print(f"build failed: {exc}", file=sys.stderr)
sys.exit(1)
print(json.dumps(info, indent=2))