MC capability panel: ARC/WinoGrande/HellaSwag/MMLU with prompt-loop
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,120 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user