Files
2026-07-14 12:33:03 +02:00

147 lines
5.9 KiB
Python

"""Eval the looped-band model: accuracy vs k, and J-lens concept sharpening.
Usage:
eval_loop.py # untrained adapter (alpha-merge only)
eval_loop.py --adapter results-loop/adapter.pt --tag trained
"""
import argparse
import json
import os
import time
from pathlib import Path
import torch
from loop_common import (DIRECT_SUFFIX, AdaptiveMergeAdapter, BandLooper,
MergeAdapter, chat_prompt,
last_number, num_eq)
import sys
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import JLens, load_model # noqa: E402
OUT = Path(os.environ.get("LOOP_OUT",
Path(__file__).resolve().parent.parent / "results-loop"))
ROOT = OUT.parent
@torch.no_grad()
def accuracy_at_k(looper, adapter, tok, items, k, batch=16, prompt_only=False,
feedforward=False):
hits, per_label, per_item = 0, {}, []
for i in range(0, len(items), batch):
torch.cuda.empty_cache()
chunk = items[i : i + batch]
enc = tok([chat_prompt(tok, it["question"], DIRECT_SUFFIX) for it in chunk],
return_tensors="pt", padding=True,
add_special_tokens=False).to("cuda")
if prompt_only:
gen = looper.generate_frozen_prompt(
adapter, tok, enc["input_ids"], k, max_new_tokens=10,
attention_mask=enc["attention_mask"], feedforward=feedforward)
else:
gen = looper.loop_generate(adapter, tok, enc["input_ids"], k,
max_new_tokens=10,
attention_mask=enc["attention_mask"])
for j, it in enumerate(chunk):
txt = tok.decode(gen[j, enc["input_ids"].shape[1]:],
skip_special_tokens=True)
ok = num_eq(last_number(txt), it["gold"])
per_item.append({"idx": it["idx"], "ok": bool(ok)})
hits += ok
d = per_label.setdefault(it["label"], [0, 0])
d[0] += ok
d[1] += 1
return (hits / len(items),
{l: c / n for l, (c, n) in per_label.items()}, per_item)
@torch.no_grad()
def spider_sharpening(looper, adapter, model, tok, ks):
"""P('spider') under the J-lens at L30, and P('8') at the output, vs k."""
jbar = torch.load(ROOT / "results" / "jbar.pt", map_location="cuda")
Jbar = (jbar["Jbar"] if isinstance(jbar, dict) else jbar).float()
lens = JLens(model, tok, Jbar)
q = ("The animal that spins webs has how many legs? "
"Answer with just the number.")
ids = tok.apply_chat_template([{"role": "user", "content": q}],
return_tensors="pt", add_generation_prompt=True,
return_dict=True)["input_ids"].cuda()
spider = tok.encode(" spider", add_special_tokens=False)[0]
eight = [tok.encode(t, add_special_tokens=False)[0] for t in (" 8", "8")]
calls, _ = looper.capture(ids)
e = looper._hin[looper.l0]
s = looper.band(e, calls)
rows = []
kmax = max(ks)
for k in range(0, kmax + 1):
if k > 0:
s = looper.band(adapter(e, s), calls)
if k in ks:
_, probs = lens.read(s[0], looper.l1, topk=1)
logits = lens._readout((s[0].float() @ Jbar[looper.l1].T))
p_spider = torch.softmax(logits.float(), -1)[:, spider].max().item()
out = looper.suffix_logits(s, calls)
p8 = torch.softmax(out[0, -1].float(), -1)[eight].max().item()
rows.append({"k": k, "P_spider_lens": p_spider, "P_8_out": p8})
return rows
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--adapter", default=None)
ap.add_argument("--tag", default="untrained")
ap.add_argument("--ks", default="0,1,2,4,8")
ap.add_argument("--n", type=int, default=0, help="cap test items (0=all)")
ap.add_argument("--prompt-only", action="store_true",
help="loop the prompt span only (unified regime, fast path)")
ap.add_argument("--no-spider", action="store_true")
ap.add_argument("--adaptive", action="store_true")
ap.add_argument("--feedforward", action="store_true",
help="no-recurrence control arm: adapter(e,e) once")
args = ap.parse_args()
ks = [int(x) for x in args.ks.split(",")]
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
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()
items = [it for it in json.load(open(OUT / "star_data.json"))
if it["split"] == "test"]
if args.n:
items = items[: args.n]
print(f"[{args.tag}] eval on {len(items)} test items, ks={ks}", flush=True)
res = {"tag": args.tag, "ks": {}, "n": len(items)}
for k in ks:
t0 = time.time()
acc, by_label, per_item = accuracy_at_k(looper, adapter, tok, items, k,
prompt_only=args.prompt_only,
feedforward=args.feedforward)
res["ks"][k] = {"acc": acc, "by_label": by_label,
"per_item": per_item}
print(f"k={k}: acc={acc:.3f} by_label={ {l: round(v,3) for l,v in by_label.items()} }"
f" ({time.time()-t0:.0f}s)", flush=True)
res["spider"] = ([] if args.no_spider else
spider_sharpening(looper, adapter, model, tok, set(ks)))
for r in res["spider"]:
print(f"spider k={r['k']}: P_lens={r['P_spider_lens']:.3f} "
f"P(8)={r['P_8_out']:.3f}", flush=True)
OUT.mkdir(exist_ok=True)
with open(OUT / f"eval_{args.tag}.json", "w") as f:
json.dump(res, f, indent=1)
print("wrote", OUT / f"eval_{args.tag}.json")
if __name__ == "__main__":
main()