File size: 6,093 Bytes
fdc6474 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 | #!/usr/bin/env python3
"""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")
# ------------------------------------------------------------ data loading
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
# ------------------------------------------------------------ server call
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"]
# ------------------------------------------------------------ HumanEval eval
def _run_candidate(program, q):
g = {"__name__": "__main__"}
try:
exec(program, g) # noqa: S102 - sandboxed in subprocess
q.put(True)
except Exception: # noqa: BLE001
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: # noqa: BLE001
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)
# ------------------------------------------------------------ GSM8K eval
_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)
# ------------------------------------------------------------ self test
def self_test():
probs = load_humaneval(20)
assert len(probs) == 20 and all(
{"prompt", "test", "entry_point"} <= set(p) for p in probs)
# exercise the subprocess sandbox with the canonical solution (should pass)
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()
|