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