Reproduction of the 2026 workspace/J-lens paper on gemma-4 (E2B/12B/26B), plus the workspace-loop retrofit line: merge adapter, prompt-only latent planning (MBPP), carry variant, attribution controls (FF/pause/untrained), band-location ablation, Blocksworld harness, 12B replication scripts. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
144 lines
5.8 KiB
Python
144 lines
5.8 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, 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("--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)
|
|
adapter = MergeAdapter(
|
|
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()
|