ASHQ1 / scripts /run_humaneval.py
wepiqx's picture
v7: whole-head MTP, Q5_K embd pins, CAN_Q3 types, top-down mode, transit fixes, scripts cleanup
491cce7
Raw History Blame Contribute Delete
2.44 kB
#!/usr/bin/env python3
"""HumanEval runner using /v1/chat/completions with template."""
import json
import os
import time
import requests
from human_eval.data import read_problems, write_jsonl
from human_eval.evaluation import evaluate_functional_correctness
SERVER_URL = "http://127.0.0.1:28082"
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
OUTPUT_FILE = os.environ.get(
"HUMANEVAL_OUT",
os.path.join(REPO_ROOT, "eval_results", "humaneval_opts.jsonl"),
)
def generate(problem, max_tokens=512, temp=0.0):
resp = requests.post(
f"{SERVER_URL}/v1/chat/completions",
json={
"messages": [
{
"role": "system",
"content": "You are an expert Python programmer. Complete the following function. Return ONLY Python code inside a fenced block ```python...```",
},
{"role": "user", "content": problem["prompt"]},
],
"max_tokens": max_tokens,
"temperature": temp,
"chat_template_kwargs": {
"add_generation_prompt": True,
"enable_thinking": False,
},
},
timeout=120,
)
resp.raise_for_status()
return resp.json()["choices"][0]["message"]["content"]
def extract_code(text):
text = text.strip()
if "```python" in text:
text = text.split("```python")[1]
if "```" in text:
text = text.split("```")[0]
return text.strip()
def main():
os.makedirs(os.path.dirname(OUTPUT_FILE), exist_ok=True)
problems = read_problems()
results = []
total = len(problems)
for i, (task_id, problem) in enumerate(sorted(problems.items())):
print(f"[{i+1}/{total}] {task_id} ... ", end="", flush=True)
try:
raw = generate(problem)
code = extract_code(raw) or raw
results.append({"task_id": task_id, "completion": code})
print("OK")
except Exception as e:
print(f"ERROR: {e}")
results.append({"task_id": task_id, "completion": ""})
time.sleep(0.1)
write_jsonl(OUTPUT_FILE, results)
print(f"\nSaved {len(results)} to {OUTPUT_FILE}")
print("\n--- Evaluating pass@1 ---")
r = evaluate_functional_correctness(sample_file=OUTPUT_FILE, k=[1], n_workers=4)
print(f"pass@1: {r}")
if __name__ == "__main__":
main()