148 lines
5.6 KiB
Python
148 lines
5.6 KiB
Python
"""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,
|
|
feedforward=False):
|
|
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"],
|
|
feedforward=feedforward)
|
|
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")
|
|
ap.add_argument("--feedforward", action="store_true")
|
|
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,
|
|
feedforward=args.feedforward)
|
|
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()
|