Files
jspace/scripts/eval_mc_panel.py
T

121 lines
4.6 KiB
Python

"""General-capability panel: MC-likelihood scoring with the prompt-loop.
The safety question for the implant: what does looping do to general
abilities? (k=0 is bit-exact base by construction; k>0 is measured here.)
Benchmarks match Lys et al. / McLeish et al.: ARC-Challenge, WinoGrande,
HellaSwag, MMLU (capped per-benchmark for Spark runtime). Scoring:
length-normalized option log-likelihood, loop applied to the context span.
"""
import argparse
import json
import os
import random
import sys
import time
from pathlib import Path
import torch
import torch.nn.functional as F
from datasets import load_dataset
from loop_common import AdaptiveMergeAdapter, BandLooper, MergeAdapter
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
CAP = int(os.environ.get("MC_CAP", "800"))
def get_items(name):
rng = random.Random(0)
if name == "arc":
ds = load_dataset("allenai/ai2_arc", "ARC-Challenge", split="test")
items = [(r["question"], r["choices"]["text"],
r["choices"]["label"].index(r["answerKey"]))
for r in ds if r["answerKey"] in r["choices"]["label"]]
elif name == "winogrande":
ds = load_dataset("winogrande", "winogrande_xl", split="validation",
trust_remote_code=True)
items = [(r["sentence"], [r["option1"], r["option2"]],
int(r["answer"]) - 1) for r in ds]
elif name == "hellaswag":
ds = load_dataset("hellaswag", split="validation",
trust_remote_code=True)
items = [(r["ctx"], r["endings"], int(r["label"])) for r in ds]
elif name == "mmlu":
ds = load_dataset("cais/mmlu", "all", split="test")
items = [(r["question"], r["choices"], r["answer"]) for r in ds]
else:
raise ValueError(name)
if len(items) > CAP:
items = rng.sample(items, CAP)
return items
@torch.no_grad()
def score_item(looper, adapter, tok, q, options, k, feedforward):
prompt = q if q.endswith((" ", "\n")) else q + " "
p_ids = tok(prompt, add_special_tokens=True)["input_ids"]
seqs = [p_ids + tok(o, add_special_tokens=False)["input_ids"]
for o in options]
T = max(len(s) for s in seqs)
pad = tok.pad_token_id or 0
ids = torch.full((len(seqs), T), pad, dtype=torch.long)
msk = torch.zeros((len(seqs), T), dtype=torch.long)
lmask = torch.zeros((len(seqs), T), dtype=torch.bool)
for i, s in enumerate(seqs):
ids[i, : len(s)] = torch.tensor(s)
msk[i, : len(s)] = 1
lmask[i, : len(p_ids)] = True
ids, msk, lmask = ids.cuda(), msk.cuda(), lmask.cuda()
logits = looper.loop_logits(adapter, ids, k, attention_mask=msk,
loop_mask=lmask, feedforward=feedforward)
lp = F.log_softmax(logits.float(), -1)
scores = []
for i, s in enumerate(seqs):
n = len(s) - len(p_ids)
tokp = lp[i, len(p_ids) - 1: len(s) - 1].gather(
-1, ids[i, len(p_ids): len(s)].unsqueeze(-1)).sum().item()
scores.append(tokp / max(1, n))
return max(range(len(scores)), key=lambda i: scores[i])
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--adapter", default=None)
ap.add_argument("--tag", default="panel")
ap.add_argument("--k", type=int, default=2)
ap.add_argument("--feedforward", action="store_true")
ap.add_argument("--adaptive", action="store_true")
ap.add_argument("--benchmarks", default="arc,winogrande,hellaswag,mmlu")
args = ap.parse_args()
model, tok = load_model(dtype=torch.bfloat16)
looper = BandLooper(model)
cls = AdaptiveMergeAdapter if args.adaptive else MergeAdapter
adapter = cls(d=model.config.get_text_config().hidden_size).cuda()
if args.adapter:
adapter.load_state_dict(torch.load(args.adapter, map_location="cuda"))
adapter.eval()
res = {"tag": args.tag, "k": args.k}
for b in args.benchmarks.split(","):
items = get_items(b)
t0 = time.time()
ok = 0
for q, opts, gold in items:
ok += score_item(looper, adapter, tok, q, opts, args.k,
args.feedforward) == gold
res[b] = ok / len(items)
print(f"[{args.tag}] {b}: {res[b]:.4f} (n={len(items)}, "
f"{time.time()-t0:.0f}s)", flush=True)
json.dump(res, open(OUT / f"eval_panel_{args.tag}.json", "w"), indent=1)
print("wrote", OUT / f"eval_panel_{args.tag}.json")
if __name__ == "__main__":
main()