"""exact-label accuracy of an HF checkpoint on a classify jsonl (greedy). usage: evalcl.py [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)}")