Files
jspace/scripts/eval_humaneval.py
T
2026-07-14 12:34:57 +02:00

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()