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:
Nils
2026-07-14 00:54:12 +02:00
co-authored by Claude Fable 5
commit ef9c08966c
42 changed files with 5068 additions and 0 deletions
+90
View File
@@ -0,0 +1,90 @@
"""Eval design C (prefill k + carry through p pauses) on GSM8K test.
Usage: eval_carry.py --adapter ../results-loop/adapter_carry_e600.pt \
--tag carry --grid "0:0,2:0,2:2,2:6"
Grid entries are k:p pairs; k=0,p=0 is the plain baseline (same harness).
"""
import argparse
import json
import sys
import time
from pathlib import Path
import torch
from carry_common import generate_carry_c
from loop_common import (DIRECT_SUFFIX, BandLooper, MergeAdapter, chat_prompt,
last_number, num_eq)
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from jlens.core import load_model # noqa: E402
OUT = Path(__file__).resolve().parent.parent / "results-loop"
@torch.no_grad()
def acc_at(looper, adapter, tok, items, k, p, batch=16, feedforward=False):
hits, per_label = 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")
gen = generate_carry_c(looper, adapter, tok, enc["input_ids"],
enc["attention_mask"], k=k, p=p,
max_new_tokens=10, feedforward=feedforward)
n0 = enc["input_ids"].shape[1] + p
for j, it in enumerate(chunk):
txt = tok.decode(gen[j, n0:], skip_special_tokens=True)
ok = num_eq(last_number(txt), it["gold"])
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()}
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--adapter", default=None)
ap.add_argument("--tag", default="carry")
ap.add_argument("--grid", default="0:0,2:0,2:2,2:6")
ap.add_argument("--n", type=int, default=0)
ap.add_argument("--feedforward", action="store_true")
args = ap.parse_args()
model, tok = load_model(dtype=torch.bfloat16)
tok.padding_side = "left"
looper = BandLooper(model)
adapter = MergeAdapter().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}] GSM8K carry eval on {len(items)} items, "
f"grid={args.grid}", flush=True)
res = {"tag": args.tag, "grid": {}, "n": len(items)}
for kp in args.grid.split(","):
k, p = (int(x) for x in kp.split(":"))
t0 = time.time()
acc, by_label = acc_at(looper, adapter, tok, items, k, p,
feedforward=args.feedforward)
res["grid"][kp] = {"acc": acc, "by_label": by_label}
print(f"k={k} p={p}: acc={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_{args.tag}.json", "w") as f:
json.dump(res, f, indent=1)
print("wrote", OUT / f"eval_{args.tag}.json")
if __name__ == "__main__":
main()