brandonmusic's picture
Add independent eval harness: rerun_errors.py
b88d516 verified
Raw History Blame Contribute Delete
3.44 kB
#!/usr/bin/env python3
"""Re-run samples that died with exhausted retries (vLLM crash window) and merge
corrected records + summary. Works for both mathbench and gpqa samples files."""
import argparse
import asyncio
import json
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).parent))
async def main():
ap = argparse.ArgumentParser()
ap.add_argument("--samples", required=True, help="path to *_samples.jsonl")
ap.add_argument("--kind", choices=["math", "gpqa"], required=True)
ap.add_argument("--dataset", help="HF dataset name (math kind)", default=None)
ap.add_argument("--concurrency", type=int, default=6)
args = ap.parse_args()
path = Path(args.samples)
records = [json.loads(l) for l in path.open()]
bad = [r for r in records if str(r.get("finish_reason", "")).startswith("error")]
print(f"{path.name}: {len(records)} records, {len(bad)} error records to re-run")
if not bad:
return
import time
from openai import AsyncOpenAI
client = AsyncOpenAI(base_url="http://localhost:8000/v1", api_key="dummy", timeout=10800.0, max_retries=0)
sem = asyncio.Semaphore(args.concurrency)
results = []
t0 = time.monotonic()
if args.kind == "math":
from datasets import load_dataset
import mathbench as mb
ds = load_dataset(args.dataset)
items = {str(i["problem_idx"]): i for i in ds[list(ds.keys())[0]]}
total = len(bad)
await asyncio.gather(*[
mb.run_one(client, sem, "GLM-5.2", items[str(r["problem_idx"])], r["repeat"], 163840, results, t0, total)
for r in bad
])
key = lambda r: (str(r["problem_idx"]), r["repeat"])
else:
from datasets import load_dataset
import gpqa_bench as gb
ds = load_dataset("Idavidrein/gpqa", "gpqa_diamond")
items = list(ds[list(ds.keys())[0]])
total = len(bad)
await asyncio.gather(*[
gb.run_one(client, sem, "GLM-5.2", items[r["idx"]], r["idx"], r["repeat"], 131072, results, t0, total)
for r in bad
])
key = lambda r: (r["idx"], r["repeat"])
fixed = {key(r): r for r in results}
merged = [fixed.get(key(r), r) for r in records]
still_bad = sum(1 for r in merged if str(r.get("finish_reason", "")).startswith("error"))
backup = path.with_suffix(".jsonl.pre-fixup")
path.rename(backup)
with path.open("w") as f:
for r in merged:
f.write(json.dumps(r) + "\n")
n = len(merged)
acc = sum(r["correct"] for r in merged) / n
print(f"MERGED: {n} records, accuracy_pass_at_1={acc:.4f}, still_errored={still_bad}")
# patch summary file if present
sp = path.parent / path.name.replace("_samples.jsonl", "_summary.json")
if sp.exists():
s = json.loads(sp.read_text())
s["accuracy_pass_at_1"] = round(acc, 4)
s["errors"] = still_bad
s["crash_fixup"] = f"re-ran {len(bad)} samples killed by vLLM crash"
per_q = {}
idx_field = "problem_idx" if args.kind == "math" else "idx"
for r in merged:
per_q.setdefault(str(r[idx_field]), []).append(r["correct"])
s["per_question_correct_rate"] = {k: round(sum(v) / len(v), 3) for k, v in sorted(per_q.items())}
sp.write_text(json.dumps(s, indent=2))
print(f"summary updated: {sp}")
if __name__ == "__main__":
asyncio.run(main())