Files
jspace/scripts/eval_loop_code.py

152 lines
6.3 KiB
Python

"""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 (AdaptiveMergeAdapter, BandLooper, MergeAdapter,
ParcaeAdapter, RecurrentAdapter)
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, halt=False):
codes = []
k_convs = []
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)
cv = {} if halt else None
gen = looper.generate_frozen_prompt(adapter, tok, enc["input_ids"], k,
max_new_tokens=max_new,
attention_mask=enc["attention_mask"],
feedforward=feedforward,
conv_out=cv)
if halt:
k_convs.extend(cv.get("k_conv", []))
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)]
if halt and k_convs:
for r, kc in zip(per_item, k_convs):
r["k_conv"] = kc
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")
ap.add_argument("--alpha", type=float, default=0.3)
ap.add_argument("--adaptive", action="store_true")
ap.add_argument("--rec", action="store_true",
help="RecurrentAdapter (noise h0, learned A/B)")
ap.add_argument("--parcae", action="store_true",
help="ParcaeAdapter (rho(A)<1 by construction)")
ap.add_argument("--perdepth", action="store_true",
help="PerDepthAdapter (Bae-style per-iteration merges)")
ap.add_argument("--tiedalpha", action="store_true",
help="TiedAlphaAdapter (learned per-dim alpha, tied B)")
ap.add_argument("--noises0", action="store_true")
ap.add_argument("--hidden", type=int, default=512)
ap.add_argument("--halt", action="store_true",
help="record per-item convergence depth (free-ACT probe)")
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)
from loop_common import (NoisyMergeAdapter, PerDepthAdapter,
TiedAlphaAdapter)
cls = (NoisyMergeAdapter if args.noises0
else TiedAlphaAdapter if args.tiedalpha
else PerDepthAdapter if args.perdepth
else ParcaeAdapter if args.parcae
else RecurrentAdapter if args.rec
else AdaptiveMergeAdapter if args.adaptive else MergeAdapter)
kw = ({"alpha0": args.alpha} if args.adaptive else {"alpha": args.alpha})
if not args.adaptive:
kw["hidden"] = args.hidden
adapter = cls(d=model.config.get_text_config().hidden_size, **kw).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, halt=args.halt)
if args.halt:
by_lbl_k = {}
for it, r in zip(items, per_item):
by_lbl_k.setdefault(it["label"], []).append(r.get("k_conv", k))
print(" mean k_conv:", {l: round(sum(v)/len(v), 2)
for l, v in by_lbl_k.items()}, flush=True)
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()