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,143 @@
|
||||
"""HumanEval transfer eval: does the MBPP-trained loop adapter generalize?
|
||||
|
||||
No HumanEval training exists (no train split) — this is pure distribution
|
||||
transfer. Labels the 164 items with the frozen model (direct vs terse-plan,
|
||||
greedy; descriptive buckets only), then evaluates the loop adapter at
|
||||
k grid + the untrained control.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from datasets import load_dataset
|
||||
|
||||
from loop_common import BandLooper, MergeAdapter
|
||||
from prep_mbpp import batch_generate, extract_code
|
||||
|
||||
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"))
|
||||
|
||||
DIRECT = ("Complete the following Python function. Return the COMPLETE "
|
||||
"function (signature included) in a ```python code block. "
|
||||
"No explanation.\n\n```python\n{prompt}```")
|
||||
PLAN = ("First write a very brief plan: at most 4 short bullet lines. Then "
|
||||
"return the COMPLETE function (signature included) in a ```python "
|
||||
"code block.\n\n```python\n{prompt}```")
|
||||
|
||||
|
||||
def he_prompt(tok, item, tmpl):
|
||||
return tok.apply_chat_template(
|
||||
[{"role": "user", "content": tmpl.format(prompt=item["prompt"])}],
|
||||
tokenize=False, add_generation_prompt=True)
|
||||
|
||||
|
||||
def run_he_tests(code, item, timeout=10):
|
||||
if not code:
|
||||
return False
|
||||
script = (code + "\n\n" + item["test"] +
|
||||
f"\ncheck({item['entry_point']})\n")
|
||||
try:
|
||||
with tempfile.TemporaryDirectory() as td:
|
||||
r = subprocess.run([sys.executable, "-c", script], cwd=td,
|
||||
capture_output=True, timeout=timeout)
|
||||
return r.returncode == 0
|
||||
except (subprocess.TimeoutExpired, OSError):
|
||||
return False
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def loop_eval(looper, adapter, tok, items, k, batch=8, max_new=380):
|
||||
codes = []
|
||||
for i in range(0, len(items), batch):
|
||||
torch.cuda.empty_cache()
|
||||
chunk = items[i : i + batch]
|
||||
enc = tok([he_prompt(tok, it, DIRECT) for it in chunk],
|
||||
return_tensors="pt", padding=True,
|
||||
add_special_tokens=False).to("cuda")
|
||||
gen = looper.generate_frozen_prompt(adapter, tok, enc["input_ids"], k,
|
||||
max_new_tokens=max_new,
|
||||
attention_mask=enc["attention_mask"])
|
||||
for j in range(len(chunk)):
|
||||
codes.append(extract_code(
|
||||
tok.decode(gen[j, enc["input_ids"].shape[1]:],
|
||||
skip_special_tokens=True)))
|
||||
with ThreadPoolExecutor(8) as ex:
|
||||
oks = list(ex.map(lambda ci: run_he_tests(ci[0], ci[1]),
|
||||
zip(codes, items)))
|
||||
return oks
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("--adapter", default=None)
|
||||
ap.add_argument("--tag", default="he")
|
||||
ap.add_argument("--ks", default="0,2,4")
|
||||
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 = list(load_dataset("openai/openai_humaneval")["test"])
|
||||
print(f"[{args.tag}] HumanEval: {len(items)} items", flush=True)
|
||||
|
||||
# labeling pass (descriptive buckets; greedy — outcome-selection caveat)
|
||||
lab_path = OUT / "humaneval_labels.json"
|
||||
if lab_path.exists():
|
||||
labels = json.load(open(lab_path))
|
||||
else:
|
||||
plans = batch_generate(model, tok,
|
||||
[he_prompt(tok, it, PLAN) for it in items],
|
||||
max_new_tokens=800, batch_size=16)
|
||||
with ThreadPoolExecutor(8) as ex:
|
||||
plan_ok = list(ex.map(
|
||||
lambda gi: run_he_tests(extract_code(gi[0]), gi[1]),
|
||||
zip(plans, items)))
|
||||
labels = {it["task_id"]: bool(ok) for it, ok in zip(items, plan_ok)}
|
||||
json.dump(labels, open(lab_path, "w"), indent=1)
|
||||
print(f"plan-reachable: {sum(labels.values())}/{len(items)}", flush=True)
|
||||
|
||||
res = {"tag": args.tag, "ks": {}, "n": len(items)}
|
||||
k0_ok = None
|
||||
for k in ks:
|
||||
t0 = time.time()
|
||||
oks = loop_eval(looper, adapter, tok, items, k)
|
||||
if k == 0:
|
||||
k0_ok = oks
|
||||
hard = [i for i, it in enumerate(items)
|
||||
if k0_ok and not k0_ok[i] and labels[it["task_id"]]]
|
||||
acc = sum(oks) / len(items)
|
||||
hard_acc = (sum(oks[i] for i in hard) / len(hard)) if hard else None
|
||||
res["ks"][k] = {"acc": acc, "hard_n": len(hard),
|
||||
"hard_acc": hard_acc,
|
||||
"per_item": [{"task_id": it["task_id"],
|
||||
"ok": bool(o)}
|
||||
for it, o in zip(items, oks)]}
|
||||
print(f"k={k}: pass@1={acc:.3f} hard({len(hard)})="
|
||||
f"{hard_acc if hard_acc is None else round(hard_acc,3)} "
|
||||
f"({time.time()-t0:.0f}s)", flush=True)
|
||||
|
||||
json.dump(res, open(OUT / f"eval_humaneval_{args.tag}.json", "w"),
|
||||
indent=1)
|
||||
print("wrote", OUT / f"eval_humaneval_{args.tag}.json")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user