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