#!/usr/bin/env python3 """Load every published checkpoint and check it against its recorded config. python3 eval/verify_checkpoints.py # all runs python3 eval/verify_checkpoints.py --sha # also verify SHA-256 Exits non-zero if any checkpoint fails to load or disagrees with its record. """ import argparse import json import sys from pathlib import Path ROOT = Path(__file__).resolve().parent.parent sys.path.insert(0, str(ROOT / "src")) from pns.checkpoint import verify # noqa: E402 from pns.common import ckpt_root, sha256_file # noqa: E402 def main(): ap = argparse.ArgumentParser() ap.add_argument("--sha", action="store_true", help="also verify SHA-256") ap.add_argument("--device", default="cpu") args = ap.parse_args() index = json.loads((ckpt_root() / "index.json").read_text()) sums = {} p = ROOT / "results" / "artifact_hashes.csv" if args.sha and p.exists(): import csv for row in csv.DictReader(p.open()): sums[row["path"]] = row["sha256"] bad = [] for run in sorted(index): r = verify(run, args.device) ok = r["loaded"] and r["tensor_total_match"] and r["params_match"] status = "ok" if ok else "MISMATCH" line = f" {run:18s} {r['kind']:8s} {r['params']:>12,} params {status}" if args.sha: rel = f"checkpoints/{run}/model.safetensors" want = sums.get(rel) got = sha256_file(ckpt_root() / run / "model.safetensors") shaok = want is None or want == got ok = ok and shaok line += f" sha {'ok' if shaok else 'MISMATCH'}" print(line, flush=True) if not ok: bad.append(run) print(f"\n{len(index) - len(bad)}/{len(index)} checkpoints verified") if bad: print("FAILED:", ", ".join(bad), file=sys.stderr) sys.exit(1) if __name__ == "__main__": main()