kenosistron-lora / scripts /bench_accept.py
disinfozone's picture
Add files using upload-large-folder tool
4335e83 verified
Raw
History Blame Contribute Delete
4.73 kB
"""Measure real MTP acceptance for a model by parsing the server's MTP[ lines.
Controlled A/B: identical prompts, identical sampler, one model then the other.
Prompts are deliberately NOT drawn from the training corpus -- the head was
fitted to that distribution, so measuring on it would flatter the trained head.
These are fresh held-out prompts in the model's normal serving register.
Usage: python3 bench_accept.py <model_id> <label> [temperature]
"""
import json
import os
import re
import subprocess
import sys
import time
import urllib.request
BASE = "http://localhost:8000/v1"
KEY = json.load(open(os.path.expanduser("~/.omlx/settings.json")))["auth"]["api_key"]
LOG = "/Users/david/.omlx/logs/server.log"
PROMPTS = [
"Describe the sound a room makes after everyone has left it.",
"What is the difference between silence and refusal to speak?",
"Write a short scene where two people fail to say goodbye.",
"Explain why a mirror is not a window, without using the word reflection.",
"A man inherits a house he has never seen. Describe his first hour inside.",
"What does it mean to be emptied rather than filled?",
"Write instructions for forgetting something on purpose.",
"Describe hunger to someone who has never eaten.",
"Why do people apologize to objects they bump into?",
"Write a paragraph that begins in a kitchen and ends in grief.",
"What is the smallest unit of betrayal?",
"Describe a city that exists only while someone is remembering it.",
"Explain the appeal of doors that lead nowhere.",
"Write a letter from a body to the person living in it.",
"What would it mean for a word to die?",
"Describe the moment just before a decision becomes irreversible.",
]
def post(model, prompt, temperature, max_tokens=300):
body = {
"model": model,
"messages": [{"role": "user", "content": prompt}],
"max_tokens": max_tokens,
"stream": False,
}
if temperature is not None:
body["temperature"] = temperature
if temperature == 0:
# Deterministic decode. t=0 alone is NOT enough here: the server's
# file-level xtc_probability=0.4 / min_p / top_p still inject
# randomness, which swamped the t1.3 A/B (paired t=1.08 on a +4pt
# effect, sd 17.3). Zero them so paired deltas reflect the model.
body["xtc_probability"] = 0.0
body["min_p"] = 0.0
body["top_p"] = 1.0
req = urllib.request.Request(
f"{BASE}/chat/completions",
data=json.dumps(body).encode(),
headers={"Authorization": f"Bearer {KEY}",
"Content-Type": "application/json"},
)
with urllib.request.urlopen(req, timeout=600) as r:
return json.load(r)
def log_lines():
return int(subprocess.run(["wc", "-l", LOG], capture_output=True,
text=True).stdout.split()[0])
def main():
model, label = sys.argv[1], sys.argv[2]
temp = float(sys.argv[3]) if len(sys.argv) > 3 else None
start = log_lines()
t0 = time.time()
completion_tokens = 0
errors = 0
for i, p in enumerate(PROMPTS):
try:
r = post(model, p, temp)
completion_tokens += r["usage"]["completion_tokens"]
except Exception as e:
errors += 1
print(f" [{i}] ERROR {e}", flush=True)
print(f" [{i + 1}/{len(PROMPTS)}] done", flush=True)
dt = time.time() - t0
# Only MTP lines emitted after our first request.
new = subprocess.run(["tail", "-n", f"+{start + 1}", LOG],
capture_output=True, text=True).stdout
A = D = 0
rates = []
for line in new.splitlines():
if "MTP[" not in line:
continue
m = re.search(r"accept=(\d+)/(\d+)", line)
if not m:
continue
a, d = int(m.group(1)), int(m.group(2))
if d == 0:
continue
A += a
D += d
rates.append(a / d * 100)
rates.sort()
out = {
"label": label, "model": model, "temperature": temp,
"requests": len(PROMPTS), "errors": errors,
"mtp_lines": len(rates),
"pooled_accept_pct": round(A / D * 100, 2) if D else None,
"accepts": A, "drafted": D,
"median_pct": round(rates[len(rates) // 2], 2) if rates else None,
"completion_tokens": completion_tokens,
"wall_s": round(dt, 1),
"tok_per_s": round(completion_tokens / dt, 2) if dt else None,
}
print(json.dumps(out, indent=1))
with open(f"/Users/david/AI/mtp_training/bench_{label}.json", "w") as f:
json.dump(out, f, indent=1)
if __name__ == "__main__":
main()