-
-
Notifications
You must be signed in to change notification settings - Fork 5.9k
Expand file tree
/
Copy pathbenchmark-local.py
More file actions
179 lines (148 loc) · 6.57 KB
/
Copy pathbenchmark-local.py
File metadata and controls
179 lines (148 loc) · 6.57 KB
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
"""
Ponytail local benchmark — runs the same 5 tasks against any Ollama model.
No promptfoo required. Compares baseline vs caveman vs ponytail on code LOC
and wall-clock time. Results are printed as a table and saved to a JSON file.
Usage:
python benchmarks/benchmark-local.py
python benchmarks/benchmark-local.py --model llama3.2 --repeat 3
Prerequisites: Ollama running locally (https://ollama.com), model pulled.
"""
import argparse
import json
import re
import time
import urllib.request
import urllib.parse
import urllib.error
from pathlib import Path
ROOT = Path(__file__).parent.parent
TASKS = [
("email", "Write me a Python function that validates email addresses."),
("debounce", "Add debounce to a search input in vanilla JavaScript. It currently fires an API call on every keystroke."),
("csv-sum", "Write Python code that reads sales.csv and sums the 'amount' column."),
("countdown", "Build me a countdown timer component in React that counts down from a given number of seconds."),
("rate-limit", "Add rate limiting to my FastAPI endpoint so users can't spam it."),
]
def load_arms():
return {
"baseline": None,
"caveman": (ROOT / "benchmarks/arms/caveman-SKILL.md").read_text(encoding="utf-8"),
"ponytail": (ROOT / "skills/ponytail/SKILL.md").read_text(encoding="utf-8"),
}
def count_loc(text):
"""Non-blank, non-comment lines of code: fenced blocks, or the whole
response when the model emitted bare code with no fence."""
blocks = re.findall(r"```[a-zA-Z0-9_+\-]*\n([\s\S]*?)```", text)
lines = ("\n".join(blocks) if blocks else text).splitlines()
return sum(
1 for l in lines
if l.strip()
and not l.strip().startswith("//")
and not l.strip().startswith("#")
and l.strip() not in ("*/",)
and not l.strip().startswith("/*")
and not l.strip().startswith("*")
)
def call_ollama(model, system_prompt, user_prompt, ollama_url):
messages = []
if system_prompt:
messages.append({"role": "system", "content": system_prompt})
messages.append({"role": "user", "content": user_prompt})
payload = json.dumps({
"model": model,
"messages": messages,
"stream": False,
"options": {"temperature": 0.7},
}).encode()
req = urllib.request.Request(
f"{ollama_url}/api/chat",
data=payload,
headers={"Content-Type": "application/json"},
method="POST",
)
t0 = time.time()
try:
with urllib.request.urlopen(req, timeout=180) as resp:
data = json.loads(resp.read())
except urllib.error.URLError as e:
raise RuntimeError(
f"Could not reach Ollama at {ollama_url} (model '{model}'): {e}. "
"Is Ollama running and is the model pulled?"
) from e
except json.JSONDecodeError as e:
raise RuntimeError(f"Ollama returned invalid JSON for model '{model}': {e}") from e
elapsed = time.time() - t0
try:
return data["message"]["content"], round(elapsed, 1)
except KeyError as e:
raise RuntimeError(
f"Unexpected response shape from Ollama for model '{model}': {data}"
) from e
def run(model, repeat, ollama_url):
arms = load_arms()
task_ids = [t[0] for t in TASKS]
# results[arm][task_id] = list of {loc, time}
results = {arm: {t: [] for t in task_ids} for arm in arms}
total = len(arms) * len(TASKS) * repeat
done = 0
for r in range(repeat):
for arm, system in arms.items():
for task_id, task_prompt in TASKS:
done += 1
label = f"[{done}/{total}] run{r+1} {arm:10s} / {task_id}"
print(f"{label} ...", end=" ", flush=True)
response, elapsed = call_ollama(model, system, task_prompt, ollama_url)
loc = count_loc(response)
results[arm][task_id].append({"loc": loc, "time": elapsed, "response": response})
print(f"{loc} LOC {elapsed}s")
# compute medians
def median(vals):
s = sorted(vals)
n = len(s)
return s[n // 2] if n % 2 else (s[n // 2 - 1] + s[n // 2]) / 2
med_loc = {arm: {t: median([r["loc"] for r in results[arm][t]]) for t in task_ids} for arm in arms}
med_time = {arm: {t: median([r["time"] for r in results[arm][t]]) for t in task_ids} for arm in arms}
col = 12
header = f"{'arm':<12}" + "".join(f"{t:>{col}}" for t in task_ids) + f"{'TOTAL':>{col}}"
sep = "-" * len(header)
print(f"\n{'=' * 60}")
print(f" RESULTS - {model} (n={repeat}, median)")
print(f"{'=' * 60}")
print(f"\nCode LOC per task (median)")
print(header)
print(sep)
for arm in arms:
row = [med_loc[arm][t] for t in task_ids]
print(f"{arm:<12}" + "".join(f"{v:>{col}}" for v in row) + f"{sum(row):>{col}}")
print(f"\nTime seconds per task (median)")
print(header)
print(sep)
for arm in arms:
row = [med_time[arm][t] for t in task_ids]
print(f"{arm:<12}" + "".join(f"{v:>{col}.1f}" for v in row) + f"{sum(row):>{col}.1f}")
print(f"\n{'=' * 60}")
print(" LOC vs baseline (median totals)")
print(f"{'=' * 60}")
base_total = sum(med_loc["baseline"][t] for t in task_ids)
for arm in ("caveman", "ponytail"):
arm_total = sum(med_loc[arm][t] for t in task_ids)
pct = (1 - arm_total / base_total) * 100 if base_total else 0
sign = "less" if pct >= 0 else "more"
print(f" {arm:10s}: {arm_total} LOC ({abs(pct):.0f}% {sign} than baseline)")
out = Path(__file__).parent / "benchmark-local-results.json"
out.write_text(json.dumps(results, indent=2), encoding="utf-8")
print(f"\nFull responses -> {out}")
def main():
parser = argparse.ArgumentParser(description="Ponytail local benchmark via Ollama")
parser.add_argument("--model", default="llama3.2", help="Ollama model name (default: llama3.2)")
parser.add_argument("--repeat", type=int, default=1, help="Runs per cell; median reported (default: 1)")
parser.add_argument("--ollama-url", default="http://localhost:11434", help="Ollama base URL")
args = parser.parse_args()
parsed_url = urllib.parse.urlparse(args.ollama_url)
if parsed_url.scheme not in ("http", "https"):
parser.error(f"Invalid --ollama-url scheme: '{parsed_url.scheme}'. Only 'http' and 'https' are supported.")
if not parsed_url.netloc:
parser.error(f"--ollama-url must include a host, e.g. http://localhost:11434 (got '{args.ollama_url}').")
run(args.model, args.repeat, args.ollama_url)
if __name__ == "__main__":
main()