Ines-1 / scripts /bench_jev.py
Endikavi's picture
Ines-1 RC1 (private staging; release commit b7f5644)
61b6fb9 verified
Raw History Blame Contribute Delete
5.26 kB
"""Accuracy and load test for serve_jev.py (stdlib only).
python scripts/bench_jev.py accuracy URL CASES.jsonl # same cases as the offline test
python scripts/bench_jev.py load URL [--cases CASES.jsonl] [--levels 1,8,32,128,256] [--n 2000] [--procs 8]
load: closed loop, `c` clients spread over --procs generator processes (one Python process is one
core, it would measure itself). Without --cases a request is the small 2-question one used for
the other Jev servers (choice of 3 + noul); with --cases, each request is one real case with all
its questions.
"""
import argparse
import collections
import json
import multiprocessing as mp
import sys
import time
import urllib.error
import urllib.request
from concurrent.futures import ThreadPoolExecutor
SMALL = {"model": "jev-latest", "state": "Me habéis cobrado dos veces la cuota de septiembre, devolvedme una ya.",
"questions": {"dep": {"type": "choice", "instructions": "¿Qué departamento?",
"criteria": {"administracion": "recibos", "mantenimiento": "averías", "juridico": "morosos"}},
"urg": {"type": "noul", "instructions": "¿Es urgente?"}}}
def post(url, body, timeout=300):
req = urllib.request.Request(url, json.dumps(body).encode(), {"Content-Type": "application/json"})
t = time.perf_counter()
try:
with urllib.request.urlopen(req, timeout=timeout) as r:
return r.status, json.loads(r.read()), time.perf_counter() - t
except urllib.error.HTTPError as e:
return e.code, None, time.perf_counter() - t
except Exception as e: # noqa: BLE001
return type(e).__name__, None, time.perf_counter() - t
def body_of(c):
return {"state": c["state"], "questions": c["questions"], "lang": c.get("lang", "es")}
def accuracy(url, path):
cases = [json.loads(x) for x in open(path, encoding="utf-8") if x.strip()]
ok = tot = 0
with ThreadPoolExecutor(32) as ex:
for c, (st, out, _) in zip(cases, ex.map(lambda c: post(url, body_of(c)), cases)):
for qid, q in c["questions"].items():
a = out["answers"][qid]
if q["type"] == "noul":
pred = "true" if a["noul"] >= 0.5 else "false"
else:
pred = max(a["probabilities"], key=a["probabilities"].get)
ok += pred in c["gold"][qid]
tot += 1
print("accuracy %d/%d (%.1f %%)" % (ok, tot, 100 * ok / tot))
def _worker(args):
url, bodies, n, threads, warm = args
def one(i):
st, out, dt = post(url, bodies[i % len(bodies)])
return (st, dt, (sum(1 for _ in out["answers"]) if out else 0),
((out.get("usage") or {}).get("input_tokens", 0) if out else 0))
with ThreadPoolExecutor(threads) as ex:
list(ex.map(one, range(min(warm, n))))
return list(ex.map(one, range(n)))
def load(url, bodies, levels, n, procs):
health = url.rsplit("/v1/", 1)[0] + "/health"
for c in levels:
p = min(procs, c)
per = [c // p + (i < c % p) for i in range(p)]
reqs = [max(1, n * k // c) for k in per]
def stats():
try:
with urllib.request.urlopen(health, timeout=10) as r:
return json.loads(r.read())["batching"]
except Exception: # noqa: BLE001
return {}
s0 = stats()
t0 = time.time()
with mp.Pool(p) as pool:
res = [x for part in pool.map(_worker, [(url, bodies[i::p] or bodies, reqs[i], per[i], per[i]) for i in range(p)])
for x in part]
dur = time.time() - t0
s1 = stats()
ms = sorted(x[1] * 1000 for x in res)
q = sum(x[2] for x in res)
tok_q = sum(x[3] for x in res) / max(1, q)
k = len(ms)
rows = (s1.get("rows", 0) - s0.get("rows", 0)) / max(1, s1.get("calls", 0) - s0.get("calls", 0))
busy = (s1.get("gpu_s", s1.get("busy_s", 0)) - s0.get("gpu_s", s0.get("busy_s", 0))) / dur
print(" c=%-4d n=%-5d %7.0f req/s %7.0f questions/s p50 %6.1f ms p90 %6.1f p95 %6.1f p99 %6.1f "
"input tok/question %6.1f rows/forward %5.1f in forward %3.0f %% %s"
% (c, k, k / dur, q / dur, ms[k // 2], ms[int(k * .9)], ms[int(k * .95)], ms[min(k - 1, int(k * .99))],
tok_q, rows, 100 * busy, dict(collections.Counter(x[0] for x in res))), flush=True)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("mode", choices=["accuracy", "load"])
ap.add_argument("url")
ap.add_argument("cases", nargs="?")
ap.add_argument("--cases", dest="cases_opt")
ap.add_argument("--levels", default="1,8,32,128,256")
ap.add_argument("--n", type=int, default=2000)
ap.add_argument("--procs", type=int, default=8)
a = ap.parse_args()
url = a.url.rstrip("/") + "/v1/systemone"
if a.mode == "accuracy":
return accuracy(url, a.cases)
src = a.cases_opt or a.cases
bodies = [body_of(json.loads(x)) for x in open(src, encoding="utf-8") if x.strip()] if src else [SMALL]
load(url, bodies, [int(x) for x in a.levels.split(",")], a.n, a.procs)
if __name__ == "__main__":
sys.exit(main())