From 6a954bd3e41220d697e0ae8de7f5d17a8e6b1a8c Mon Sep 17 00:00:00 2001 From: Nils Date: Tue, 14 Jul 2026 13:58:57 +0200 Subject: [PATCH] MC capability panel: ARC/WinoGrande/HellaSwag/MMLU with prompt-loop Co-Authored-By: Claude Fable 5 --- scripts/eval_mc_panel.py | 120 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 120 insertions(+) create mode 100644 scripts/eval_mc_panel.py diff --git a/scripts/eval_mc_panel.py b/scripts/eval_mc_panel.py new file mode 100644 index 0000000..b175b95 --- /dev/null +++ b/scripts/eval_mc_panel.py @@ -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()