"""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("--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 PerDepthAdapter cls = (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}) 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()