dasLOL's picture
classify evalcl.py
7f7399f verified
Raw History Blame Contribute Delete
1.64 kB
"""exact-label accuracy of an HF checkpoint on a classify jsonl (greedy). usage: evalcl.py <model> <jsonl> [N]"""
import json, os, re, sys, time
from pathlib import Path
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
H = Path(os.path.expanduser("~/mt92"))
MODEL = str(H/sys.argv[1]); SET = H/sys.argv[2]; N = int(sys.argv[3]) if len(sys.argv) > 3 else 10**9
tok = AutoTokenizer.from_pretrained(MODEL); tok.padding_side = "left"
if tok.pad_token_id is None: tok.pad_token_id = 151643
model = AutoModelForCausalLM.from_pretrained(MODEL, dtype=torch.bfloat16).cuda().eval()
rows = [json.loads(l) for l in SET.read_text(encoding="utf-8").splitlines() if l.strip()][:N]
order = sorted(range(len(rows)), key=lambda i: len(rows[i]["prompt"]))
norm = lambda s: re.sub(r"\s+", " ", str(s).strip().lower())
correct = 0; ntok = []; t0 = time.time(); B = 128
for s in range(0, len(order), B):
idx = order[s:s+B]; chunk = [rows[i] for i in idx]
enc = tok([r["prompt"] for r in chunk], return_tensors="pt", padding=True, add_special_tokens=False).to("cuda")
with torch.no_grad():
o = model.generate(**enc, max_new_tokens=16, do_sample=False, pad_token_id=151643, eos_token_id=151643)
L = enc["input_ids"].shape[1]
for j, i in enumerate(idx):
gen = tok.decode(o[j][L:], skip_special_tokens=True); ntok.append(int((o[j][L:] != 151643).sum()))
correct += norm(gen) == norm(rows[i]["completion"])
print(f"=== {sys.argv[1]} on {SET.name} ({len(rows)} docs, {time.time()-t0:.0f}s) ===")
print(f" ACCURACY {correct/len(rows):.4f} | out tok med {sorted(ntok)[len(ntok)//2]} max {max(ntok)}")