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