| |
| """bench_small: quick HumanEval(20) pass@1 + GSM8K(50) exact-match against a |
| running vLLM OpenAI-compatible server. |
| |
| python tools/sanity/bench_small.py --port 8199 |
| |
| HumanEval: first 20 problems from openai/human-eval (GitHub raw jsonl.gz), |
| greedy, 512 max tokens, completions executed in a subprocess sandbox with a |
| 5 s timeout; reports pass@1. |
| GSM8K: 50 problems from the HF `gsm8k` (main) test split (streaming), 3-shot, |
| final-number extraction, exact match. |
| |
| Use --self-test to only validate that imports and dataset downloads work |
| (no server required); this is what CI/staging runs when no server is up. |
| """ |
| import argparse |
| import gzip |
| import io |
| import json |
| import multiprocessing as mp |
| import re |
| import urllib.request |
|
|
| HUMANEVAL_URL = ("https://github.com/openai/human-eval/raw/master/data/" |
| "HumanEval.jsonl.gz") |
|
|
|
|
| |
| def load_humaneval(n=20): |
| raw = urllib.request.urlopen(HUMANEVAL_URL, timeout=120).read() |
| text = gzip.GzipFile(fileobj=io.BytesIO(raw)).read().decode() |
| probs = [json.loads(line) for line in text.splitlines() if line.strip()] |
| return probs[:n] |
|
|
|
|
| def load_gsm8k(n_test=50, n_shot=3): |
| from datasets import load_dataset |
| train = load_dataset("openai/gsm8k", "main", split="train", streaming=True) |
| test = load_dataset("openai/gsm8k", "main", split="test", streaming=True) |
| shots = [] |
| for row in train: |
| shots.append((row["question"], row["answer"])) |
| if len(shots) >= n_shot: |
| break |
| tests = [] |
| for row in test: |
| tests.append((row["question"], row["answer"])) |
| if len(tests) >= n_test: |
| break |
| return shots, tests |
|
|
|
|
| |
| def complete(port, prompt, max_tokens, stop=None): |
| model = json.load(urllib.request.urlopen( |
| f"http://localhost:{port}/v1/models", timeout=30))["data"][0]["id"] |
| body = {"model": model, "prompt": prompt, "max_tokens": max_tokens, |
| "temperature": 0.0} |
| if stop: |
| body["stop"] = stop |
| req = urllib.request.Request( |
| f"http://localhost:{port}/v1/completions", |
| data=json.dumps(body).encode(), |
| headers={"Content-Type": "application/json"}) |
| with urllib.request.urlopen(req, timeout=600) as r: |
| return json.load(r)["choices"][0]["text"] |
|
|
|
|
| |
| def _run_candidate(program, q): |
| g = {"__name__": "__main__"} |
| try: |
| exec(program, g) |
| q.put(True) |
| except Exception: |
| q.put(False) |
|
|
|
|
| def check_humaneval(problem, completion, timeout=5): |
| program = (problem["prompt"] + completion + "\n" |
| + problem["test"] + "\n" |
| + f"check({problem['entry_point']})\n") |
| ctx = mp.get_context("fork") |
| q = ctx.Queue() |
| p = ctx.Process(target=_run_candidate, args=(program, q)) |
| p.start() |
| p.join(timeout) |
| if p.is_alive(): |
| p.terminate() |
| p.join() |
| return False |
| try: |
| return q.get_nowait() |
| except Exception: |
| return False |
|
|
|
|
| def eval_humaneval(port, n=20): |
| probs = load_humaneval(n) |
| passed = 0 |
| for pr in probs: |
| comp = complete(port, pr["prompt"], 512, |
| stop=["\ndef ", "\nclass ", "\nif __name__", "\nprint("]) |
| if check_humaneval(pr, comp): |
| passed += 1 |
| print(f"HumanEval pass@1: {passed}/{len(probs)} = {passed/len(probs):.3f}") |
| return passed / len(probs) |
|
|
|
|
| |
| _NUM = re.compile(r"-?\d[\d,]*(?:\.\d+)?") |
|
|
|
|
| def extract_answer(text): |
| if "####" in text: |
| text = text.split("####")[-1] |
| nums = _NUM.findall(text) |
| if not nums: |
| return None |
| return nums[-1].replace(",", "").rstrip(".") |
|
|
|
|
| def build_prompt(shots, question): |
| parts = [] |
| for q, a in shots: |
| parts.append(f"Question: {q}\nAnswer: {a}") |
| parts.append(f"Question: {question}\nAnswer:") |
| return "\n\n".join(parts) |
|
|
|
|
| def eval_gsm8k(port, n=50, n_shot=3): |
| shots, tests = load_gsm8k(n, n_shot) |
| correct = 0 |
| for q, gold_ans in tests: |
| gold = extract_answer(gold_ans) |
| prompt = build_prompt(shots, q) |
| out = complete(port, prompt, 512, stop=["\n\nQuestion:", "\nQuestion:"]) |
| if extract_answer(out) == gold: |
| correct += 1 |
| print(f"GSM8K exact-match: {correct}/{len(tests)} = {correct/len(tests):.3f}") |
| return correct / len(tests) |
|
|
|
|
| |
| def self_test(): |
| probs = load_humaneval(20) |
| assert len(probs) == 20 and all( |
| {"prompt", "test", "entry_point"} <= set(p) for p in probs) |
| |
| p0 = probs[0] |
| ok = check_humaneval(p0, p0["canonical_solution"]) |
| assert ok is True, "sandbox failed to verify canonical HumanEval solution" |
| shots, tests = load_gsm8k(50, 3) |
| assert len(shots) == 3 and len(tests) == 50 |
| assert extract_answer("The result is #### 42") == "42" |
| assert extract_answer("so the answer is 1,024 apples.") == "1024" |
| print(f"self-test OK: HumanEval={len(probs)} problems " |
| f"(entry_point[0]={probs[0]['entry_point']}), " |
| f"GSM8K shots={len(shots)} tests={len(tests)}, " |
| f"gold[0]={extract_answer(tests[0][1])}, sandbox_canonical_pass={ok}") |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--port", type=int) |
| ap.add_argument("--self-test", action="store_true") |
| a = ap.parse_args() |
| if a.self_test: |
| self_test() |
| return |
| if a.port is None: |
| ap.error("--port is required unless --self-test") |
| he = eval_humaneval(a.port, 20) |
| gs = eval_gsm8k(a.port, 50, 3) |
| print(f"bench_small: HumanEval={he:.3f} GSM8K={gs:.3f}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|