pns-bind-25m / eval /verify_checkpoints.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
1.93 kB
#!/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()