Voice / benchmarking /shim.py
Wiself's picture
benchmarking: refresh run logs and shim
285bb31
Raw History Blame Contribute Delete
6.37 kB
#!/usr/bin/env python3
"""Token-shape shim between lm-eval and llama-server. Stdlib only.
llama-server and lm-eval disagree on three shapes; this translates and
reverse-proxies everything else untouched:
- GET /tokenizer_info (missing on llama-server): served from the server's
own /props (bos/eos strings, canonical chat template).
- POST /tokenize {"prompt"} -> upstream {"content"} -> {"tokens"}.
- POST /detokenize {"tokens"} -> upstream -> {"prompt"}.
- /v1/completions logprobs: server returns chat-style
content[{token, logprob, top_logprobs}]; harness expects legacy
{token_logprobs[], top_logprobs[]}. Translated field-for-field.
- /v1/chat/completions: system messages are folded into the following user
turn, but only for families without a system role (Gemma). See families.py.
- Sidecar: every completions request + raw response (incl. reasoning_content)
is appended to thinking-log.jsonl for later inspection. Scoring untouched.
Usage: python3 shim.py [upstream] [port]
"""
import json
import os
import sys
import time
import urllib.request
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
UPSTREAM = sys.argv[1] if len(sys.argv) > 1 else "http://host.docker.internal:8080"
PORT = int(sys.argv[2]) if len(sys.argv) > 2 else 8081
THINK_LOG = os.path.join(os.path.dirname(os.path.abspath(__file__)), "thinking-log.jsonl")
def _up(method, path, body=None):
data = json.dumps(body).encode() if body is not None else None
req = urllib.request.Request(
UPSTREAM + path, data=data, method=method,
headers={"Content-Type": "application/json"},
)
with urllib.request.urlopen(req, timeout=300) as r:
return r.status, r.read()
def _log(path, req_body, resp_raw):
try:
resp = json.loads(resp_raw)
label = (req_body or {}).get("model", "model")
for ch in resp.get("choices", []):
msg = ch.get("message", {})
# server echoes its local file path as the model id: personal,
# meaningless to others. Log our run label instead.
if isinstance(msg, dict):
msg.pop("model", None)
if "model" in resp:
resp["model"] = label
with open(THINK_LOG, "a") as f:
f.write(json.dumps({
"ts": time.time(),
"path": path,
"request": req_body,
"response": resp,
}) + "\n")
except Exception:
pass
def _info():
try:
_, raw = _up("GET", "/props")
p = json.loads(raw)
eos = p.get("eos_token", "<eos>")
bos = p.get("bos_token", "<bos>")
tmpl = p.get("chat_template") or ""
except Exception:
eos, bos, tmpl = "<eos>", "<bos>", ""
return {"eos_token": eos, "bos_token": bos, "pad_token": eos,
"chat_template": tmpl}
def _family():
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import families
try:
_, raw = _up("GET", "/props")
tmpl = json.loads(raw).get("chat_template", "")
except Exception:
tmpl = ""
name = families.detect(tmpl, sys.argv[3] if len(sys.argv) > 3 else "")
sys.stderr.write(f"shim: family={name}\n")
return families.profile(name)
def _fix_logprobs(raw):
try:
obj = json.loads(raw)
except Exception:
return raw
for ch in obj.get("choices", []):
lp = (ch.get("logprobs") or {}).get("content")
if isinstance(lp, list):
ch["logprobs"] = {
"token_logprobs": [e.get("logprob", 0.0) for e in lp],
"top_logprobs": [
{t.get("token", ""): t.get("logprob", 0.0)
for t in (e.get("top_logprobs") or [])}
or {e.get("token", ""): e.get("logprob", 0.0)}
for e in lp
],
}
return json.dumps(obj).encode()
def _fold_system(body):
try:
msgs = body.get("messages")
if isinstance(msgs, list) and msgs and msgs[0].get("role") == "system":
rest = msgs[1:]
if rest and rest[0].get("role") == "user":
rest[0] = dict(rest[0], content=msgs[0].get("content", "")
+ "\n" + rest[0].get("content", ""))
return dict(body, messages=rest)
except Exception:
pass
return body
INFO = _info()
FAMILY = _family()
class H(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def log_message(self, *a):
pass
def _send(self, code, obj=None, raw=None):
body = raw if raw is not None else json.dumps(obj).encode()
self.send_response(code)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def _read_json(self):
try:
n = int(self.headers.get("Content-Length") or 0)
except ValueError:
n = 0
if not n:
return {}
return json.loads(self.rfile.read(n) or b"{}")
def do_GET(self):
if self.path == "/tokenizer_info":
return self._send(200, INFO)
code, raw = _up("GET", self.path)
self._send(code, None, raw)
def do_POST(self):
body = self._read_json()
if self.path == "/tokenize":
code, raw = _up("POST", "/tokenize",
{"content": body.get("prompt", "")})
toks = json.loads(raw).get("tokens", [])
return self._send(code, {"tokens": toks})
if self.path == "/detokenize":
code, raw = _up("POST", "/detokenize",
{"tokens": body.get("tokens", [])})
text = json.loads(raw).get("content", "")
return self._send(code, {"prompt": text})
if self.path == "/v1/chat/completions" and FAMILY["fold_system"]:
body = _fold_system(body)
try:
code, raw = _up("POST", self.path, body)
except Exception as e:
return self._send(502, {"error": f"upstream: {e}"})
if "completions" in self.path:
_log(self.path, body, raw)
self._send(code, None, _fix_logprobs(raw))
if __name__ == "__main__":
ThreadingHTTPServer(("127.0.0.1", PORT), H).serve_forever()