J-lens workspace reproduction + loop retrofit: lens, band looping, adapters, controls, multi-task evals
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>
This commit is contained in:
@@ -0,0 +1,112 @@
|
||||
"""Eval MBPP pass@1 vs loop depth k with prompt-only ("latent planning") loops.
|
||||
|
||||
Usage:
|
||||
eval_loop_code.py --tag untrained
|
||||
eval_loop_code.py --adapter ../results-loop/adapter_code.pt --tag trained
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from loop_common import BandLooper, MergeAdapter
|
||||
from prep_mbpp import DIRECT_SUFFIX, extract_code, mbpp_prompt, run_tests
|
||||
|
||||
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"))
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def pass1_at_k(looper, adapter, tok, items, k, batch=8, max_new=220,
|
||||
feedforward=False, pause=0):
|
||||
codes = []
|
||||
for i in range(0, len(items), batch):
|
||||
torch.cuda.empty_cache()
|
||||
chunk = items[i : i + batch]
|
||||
enc = tok([mbpp_prompt(tok, it, DIRECT_SUFFIX) for it in chunk],
|
||||
return_tensors="pt", padding=True,
|
||||
add_special_tokens=False).to("cuda")
|
||||
if pause:
|
||||
B = enc["input_ids"].shape[0]
|
||||
pcol = torch.full((B, pause), 6, dtype=torch.long, device="cuda")
|
||||
enc["input_ids"] = torch.cat([enc["input_ids"], pcol], 1)
|
||||
enc["attention_mask"] = torch.cat(
|
||||
[enc["attention_mask"], torch.ones_like(pcol)], 1)
|
||||
gen = looper.generate_frozen_prompt(adapter, tok, enc["input_ids"], k,
|
||||
max_new_tokens=max_new,
|
||||
attention_mask=enc["attention_mask"],
|
||||
feedforward=feedforward)
|
||||
for j in range(len(chunk)):
|
||||
txt = tok.decode(gen[j, enc["input_ids"].shape[1]:],
|
||||
skip_special_tokens=True)
|
||||
codes.append(extract_code(txt))
|
||||
with ThreadPoolExecutor(8) as ex:
|
||||
oks = list(ex.map(lambda ci: run_tests(ci[0], ci[1]),
|
||||
zip(codes, items)))
|
||||
hits, per_label = 0, {}
|
||||
for it, ok in zip(items, oks):
|
||||
hits += ok
|
||||
d = per_label.setdefault(it["label"], [0, 0])
|
||||
d[0] += ok
|
||||
d[1] += 1
|
||||
per_item = [{"task_id": it["task_id"], "ok": bool(ok)}
|
||||
for it, ok in zip(items, oks)]
|
||||
return (hits / len(items),
|
||||
{l: c / n for l, (c, n) in per_label.items()}, per_item)
|
||||
|
||||
|
||||
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")
|
||||
ap.add_argument("--n", type=int, default=250)
|
||||
ap.add_argument("--feedforward", action="store_true",
|
||||
help="no-recurrence control arm: adapter(e,e) once")
|
||||
ap.add_argument("--pause", type=int, default=0,
|
||||
help="append p pause tokens to each prompt")
|
||||
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 / "mbpp_data.json"))
|
||||
if it["split"] == "test"][: args.n]
|
||||
print(f"[{args.tag}] MBPP 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 = pass1_at_k(looper, adapter, tok, items, k,
|
||||
feedforward=args.feedforward,
|
||||
pause=args.pause)
|
||||
res["ks"][k] = {"acc": acc, "by_label": by_label,
|
||||
"per_item": per_item}
|
||||
print(f"k={k}: pass@1={acc:.3f} "
|
||||
f"by_label={ {l: round(v,3) for l,v in by_label.items()} }"
|
||||
f" ({time.time()-t0:.0f}s)", flush=True)
|
||||
|
||||
with open(OUT / f"eval_code_{args.tag}.json", "w") as f:
|
||||
json.dump(res, f, indent=1)
|
||||
print("wrote", OUT / f"eval_code_{args.tag}.json")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user