E1 halting-gate machinery (adapter, trainer, gated eval) + pre-registration item 18
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,114 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user