"""E1 eval (PLAN_SELFPACED): halting-gate MBPP eval. Per item: deploy-time halting picks k* (1..kmax); generation then uses the frozen-prompt path at that k (items grouped by k* for batching). Reports pass@1 by label, the k* distribution by label, mean compute, and the gate-difficulty point-biserial correlation. """ import argparse import json import os import sys import time from concurrent.futures import ThreadPoolExecutor from pathlib import Path import torch from halting_common import HaltingMergeAdapter, halted_k_per_item from loop_common import BandLooper 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")) def main(): ap = argparse.ArgumentParser() ap.add_argument("--adapter", required=True) ap.add_argument("--tag", required=True) ap.add_argument("--kmax", type=int, default=4) ap.add_argument("--n", type=int, default=250) ap.add_argument("--batch", type=int, default=8) args = ap.parse_args() model, tok = load_model(dtype=torch.bfloat16) tok.padding_side = "left" looper = BandLooper(model) adapter = HaltingMergeAdapter( d=model.config.get_text_config().hidden_size).cuda() 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}] gated MBPP eval on {len(items)}, kmax={args.kmax}", flush=True) # phase 1: per-item k* kstars = [] with torch.no_grad(): for i in range(0, len(items), args.batch): chunk = items[i : i + args.batch] enc = tok([mbpp_prompt(tok, it, DIRECT_SUFFIX) for it in chunk], return_tensors="pt", padding=True, add_special_tokens=False).to("cuda") pl = enc["attention_mask"].sum(-1) - 1 # left-pad: last position pl = torch.full_like(pl, enc["input_ids"].shape[1] - 1) ks = halted_k_per_item(looper, adapter, enc["input_ids"], args.kmax, enc["attention_mask"], pl) kstars.extend(ks.tolist()) t0 = time.time() # phase 2: generate grouped by k* codes = [None] * len(items) for kval in sorted(set(kstars)): idxs = [i for i, kk in enumerate(kstars) if kk == kval] for j in range(0, len(idxs), args.batch): grp = idxs[j : j + args.batch] enc = tok([mbpp_prompt(tok, items[i], DIRECT_SUFFIX) for i in grp], return_tensors="pt", padding=True, add_special_tokens=False).to("cuda") gen = looper.generate_frozen_prompt( adapter, tok, enc["input_ids"], kval, max_new_tokens=220, attention_mask=enc["attention_mask"]) for gi, i in enumerate(grp): txt = tok.decode(gen[gi, enc["input_ids"].shape[1]:], skip_special_tokens=True) codes[i] = extract_code(txt) with ThreadPoolExecutor(8) as ex: oks = list(ex.map(lambda ci: run_tests(ci[0], ci[1]), zip(codes, items))) per_label, kdist = {}, {} for it, ok, kk in zip(items, oks, kstars): d = per_label.setdefault(it["label"], [0, 0, 0.0]) d[0] += ok; d[1] += 1; d[2] += kk kdist.setdefault(it["label"], []).append(kk) acc = sum(oks) / len(items) by_label = {l: c / n for l, (c, n, _) in per_label.items()} mean_k = {l: sum(v) / len(v) for l, v in kdist.items()} hard = torch.tensor([it["label"] == "hard" for it in items], dtype=torch.float) kk = torch.tensor(kstars, dtype=torch.float) r = ((kk - kk.mean()) * (hard - hard.mean())).mean() / (kk.std() * hard.std() + 1e-9) print(f"pass@1={acc:.3f} by_label={ {l: round(v,3) for l,v in by_label.items()} }") print(f"mean k* by label: { {l: round(v,2) for l,v in mean_k.items()} } " f"overall E[k]={sum(kstars)/len(kstars):.2f} " f"gate-difficulty r={r:.3f} ({time.time()-t0:.0f}s)", flush=True) json.dump({"tag": args.tag, "acc": acc, "by_label": by_label, "mean_kstar": mean_k, "corr_hard": r.item(), "kstars": kstars, "per_item": [{"task_id": it["task_id"], "ok": bool(o), "kstar": kk_} for it, o, kk_ in zip(items, oks, kstars)]}, open(OUT / f"eval_code_{args.tag}.json", "w"), indent=1) print("wrote", OUT / f"eval_code_{args.tag}.json") if __name__ == "__main__": main()