File size: 2,272 Bytes
414fc14 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 | #!/usr/bin/env python
import argparse
import json
import math
import statistics
from pathlib import Path
TRACKS = ("iid_single", "iid_multi", "ood_single", "ood_multi")
METHODS = ("sft", "random", "greedy_verifier", "greedy_oracle")
METRICS = ("success", "regret", "direction")
def ci(xs):
sd = statistics.stdev(xs)
return {"mean": statistics.fmean(xs), "sd": sd, "ci95": 2.776 * sd / math.sqrt(len(xs))}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--input", type=Path, required=True)
ap.add_argument("--output", type=Path, required=True)
ap.add_argument("--markdown", type=Path, required=True)
args = ap.parse_args()
rows = [json.loads(p.read_text()) for p in sorted(args.input.glob("seed_*.json"))]
if len(rows) != 5 or len({r["rule_split_hash"] for r in rows}) != 1:
raise SystemExit("expected five seeds with one rule split")
summary = {track: {method: {metric: ci([r["results"][track][method][metric] for r in rows])
for metric in METRICS} for method in METHODS} for track in TRACKS}
drops = {}
for scale in ("single", "multi"):
drops[scale] = {method: {metric: ci([
r["results"][f"ood_{scale}"][method][metric]
- r["results"][f"iid_{scale}"][method][metric] for r in rows])
for metric in METRICS} for method in METHODS}
output = {"seeds": [r["seed"] for r in rows], "rule_split_hash": rows[0]["rule_split_hash"],
"summary": summary, "ood_minus_iid": drops}
args.output.write_text(json.dumps(output, indent=2) + "\n")
lines = ["# PolyEdit real edit-rule OOD", "",
"Values are five-seed means with 95% t confidence intervals.", ""]
for track in TRACKS:
lines += [f"## {track}", "", "| Method | Success | Regret | Direction |",
"|---|---:|---:|---:|"]
for method in METHODS:
row = summary[track][method]
cell = lambda metric: f"{row[metric]['mean']:.3f} ± {row[metric]['ci95']:.3f}"
lines.append(f"| {method} | {cell('success')} | {cell('regret')} | {cell('direction')} |")
lines.append("")
args.markdown.write_text("\n".join(lines) + "\n")
if __name__ == "__main__":
main()
|