"""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()