Files
esh-pfi-infrastructure/services/parakeet-ab-2026-09-30/code/bench.py
T
vh a6c1d3c454 docs(parakeet): seat A/B vs parakeet-unified-en-0.6b - latency is the int8-on-CPU runtime; unified wins WER
A/B of the live STT seat (fv-ml1 GPU 0, sherpa-onnx int8 v3) against
nvidia/parakeet-unified-en-0.6b, measured on GPU 3 with the seat's own image,
k2-fsa's published unified int8 export, fp32/fp16 exports made with k2-fsa's
recipe, v2 int8, and NeMo 3.0.0 (fp32, bf16 autocast, bf16 weights).

- Seat int8 graph runs on one CPU thread (cpu/wall 1.00, GPU 2-9%).
- unified-en under NeMo: -121/-234/-530 ms vs the seat at 1-3/3-8/8-20 s
  (paired, n=120/bin; floor <=6 ms; +50 ms positive control reads +52-54).
- unified-en WER lower in every runtime: -0.7 pp clean, -1.5 pp other,
  -3.2 to -4.4 pp AMI (paired CIs exclude 0).
- Seat defects found: hard 400 s input ceiling (HTTP 500), truncation after
  a quiet 1.5 s pause, and severe long-window dropouts (int8 v3 only).
- B-bf16w needs +0.8 to +1.5 GB over the seat's 1,690 MiB on GPU 0.

Raw requests, hypotheses, manifests and the full harness under
services/parakeet-ab-2026-09-30/. No deploy; live seat untouched apart
from 240 light test requests.
2026-09-30 18:51:44 -07:00

187 lines
7.6 KiB
Python
Executable File

#!/usr/bin/env python3
"""A/B client. Stdlib only. One NEW HTTP connection per request (what talk's httpx.AsyncClient-per-call does).
bench.py lat URL ARM MANIFEST OUT --rounds R [--seed S] [--pause SEC] [--auth KEY] [--bins b1_3,b3_8]
R rounds; each round sends every clip once, in a seeded shuffled order. Single stream.
bench.py conc URL ARM MANIFEST OUT --per-bin N [--workers 4] [--seed S]
Per bin: N requests drawn from that bin's clips, W concurrent workers, back to back.
bench.py acc URL ARM MANIFEST OUT
Every manifest row once, in manifest order; stores the text.
Rows (jsonl): arm, mode, round, id, bin, dur, e2e_ms, server_ms (x-ab-decode-ms header; null on the live seat),
status, text (acc/lat), t_wall (unix), gpu_snapshot (lat: per round, nvidia-smi compute apps).
"""
import argparse
import sys
import http.client
import json
import random
import subprocess
import threading
import time
import urllib.parse
import uuid
def multipart(wav: bytes):
b = uuid.uuid4().hex
body = (f"--{b}\r\nContent-Disposition: form-data; name=\"model\"\r\n\r\next-stt\r\n"
f"--{b}\r\nContent-Disposition: form-data; name=\"file\"; filename=\"turn.wav\"\r\n"
f"Content-Type: audio/wav\r\n\r\n").encode() + wav + f"\r\n--{b}--\r\n".encode()
return body, f"multipart/form-data; boundary={b}"
def post(url, wav, auth=None, timeout=3600):
u = urllib.parse.urlparse(url)
body, ctype = multipart(wav)
hdr = {"Content-Type": ctype, "Content-Length": str(len(body))}
if auth:
hdr["Authorization"] = f"Bearer {auth}"
t0 = time.perf_counter()
c = http.client.HTTPConnection(u.hostname, u.port or 80, timeout=timeout)
try:
c.request("POST", u.path, body=body, headers=hdr)
r = c.getresponse()
data = r.read()
t1 = time.perf_counter()
sm = r.getheader("x-ab-decode-ms")
status = r.status
finally:
c.close()
text = None
try:
text = json.loads(data).get("text")
except Exception:
text = None
return dict(e2e_ms=round((t1 - t0) * 1000, 3), server_ms=float(sm) if sm else None, status=status, text=text,
err=None if status == 200 else data[:300].decode("utf-8", "replace"))
def gpu_snapshot():
try:
out = subprocess.run(["nvidia-smi", "--query-compute-apps=gpu_uuid,pid,used_memory", "--format=csv,noheader,nounits"],
capture_output=True, text=True, timeout=20).stdout
return [l.strip() for l in out.splitlines() if l.strip()]
except Exception as e: # noqa: BLE001
return [f"err {e}"]
def trimmed(wav_bytes, trim_ms):
"""Same clip, `trim_ms` shorter at the tail: a never-seen input length (production utterances are all
unique lengths, and ORT / NeMo cache per shape). Built BEFORE timing starts."""
import io, wave
if not trim_ms:
return wav_bytes
with wave.open(io.BytesIO(wav_bytes)) as w:
p = w.getparams()
fr = w.readframes(p.nframes)
keep = max(1, p.nframes - int(p.framerate * trim_ms / 1000)) * p.sampwidth * p.nchannels
bio = io.BytesIO()
with wave.open(bio, "wb") as o:
o.setnchannels(p.nchannels); o.setsampwidth(p.sampwidth); o.setframerate(p.framerate)
o.writeframes(fr[:keep])
return bio.getvalue()
def jitter_ms(seed, k, uid, max_ms):
return 0 if not max_ms else random.Random(f"{seed}-{k}-{uid}").randint(0, max_ms // 10) * 10
def hostpath(p):
"""Manifests carry the container path (/ab/...); the client may run on the host."""
import os
return p if os.path.exists(p) else p.replace("/ab/", "/tank/spikes/parakeet-ab/", 1)
def load(manifest, bins=None):
rows = [json.loads(l) for l in open(manifest)]
if bins:
rows = [r for r in rows if r.get("bin") in bins]
for r in rows:
r["_wav"] = open(hostpath(r["wav"]), "rb").read()
return rows
def main():
ap = argparse.ArgumentParser()
ap.add_argument("mode")
ap.add_argument("url")
ap.add_argument("arm")
ap.add_argument("manifest")
ap.add_argument("out")
ap.add_argument("--rounds", type=int, default=1)
ap.add_argument("--round-offset", type=int, default=0)
ap.add_argument("--seed", type=int, default=20260930)
ap.add_argument("--pause", type=float, default=0.0)
ap.add_argument("--auth")
ap.add_argument("--bins")
ap.add_argument("--per-bin", type=int, default=40)
ap.add_argument("--workers", type=int, default=4)
ap.add_argument("--stop-on-error", action="store_true", help="abort at the first non-200 (used on the live seat)")
ap.add_argument("--jitter-ms", type=int, default=0, help="trim a seeded 0..N ms (10 ms steps) off each request's tail")
a = ap.parse_args()
rows = load(a.manifest, a.bins.split(",") if a.bins else None)
out = open(a.out, "a")
def emit(d):
out.write(json.dumps(d) + "\n")
out.flush()
if a.mode == "acc":
for r in rows:
res = post(a.url, r["_wav"], a.auth)
emit(dict(arm=a.arm, mode="acc", id=r["id"], dur=r["dur"], t_wall=time.time(), **res))
elif a.mode == "lat":
for k in range(a.round_offset, a.round_offset + a.rounds):
order = list(range(len(rows)))
random.Random(a.seed * 1000 + k).shuffle(order)
jit = {i: jitter_ms(a.seed, k, rows[i]["id"], a.jitter_ms) for i in order}
body = {i: trimmed(rows[i]["_wav"], jit[i]) for i in order}
snap0 = gpu_snapshot()
for i in order:
r = rows[i]
res = post(a.url, body[i], a.auth)
res["trim_ms"] = jit[i]
res.pop("err", None) if res["status"] == 200 else None
emit(dict(arm=a.arm, mode="lat", round=k, id=r["id"], bin=r.get("bin"), dur=r["dur"], t_wall=time.time(), **res))
if a.stop_on_error and res["status"] != 200:
print("STOP: non-200 from", a.url, res.get("err"), flush=True)
sys.exit(3)
if a.pause:
time.sleep(a.pause)
emit(dict(arm=a.arm, mode="lat-round", round=k, gpu_start=snap0, gpu_end=gpu_snapshot(), t_wall=time.time()))
elif a.mode == "conc":
for tag in sorted({r["bin"] for r in rows}, key=lambda t: int(t[1:].split("_")[0])):
pool = [r for r in rows if r["bin"] == tag]
rng = random.Random(a.seed)
queue = []
for q in range(a.per_bin):
r = rng.choice(pool)
t = jitter_ms(a.seed, 10_000 + q, r["id"], a.jitter_ms)
queue.append(dict(r, _wav=trimmed(r["_wav"], t), _trim=t))
lock = threading.Lock()
snap0 = gpu_snapshot()
def worker(w):
while True:
with lock:
if not queue:
return
r = queue.pop()
res = post(a.url, r["_wav"], a.auth)
res.pop("text", None)
with lock:
emit(dict(arm=a.arm, mode="conc", workers=a.workers, worker=w, id=r["id"], bin=tag, dur=r["dur"],
trim_ms=r["_trim"], t_wall=time.time(), **res))
t0 = time.perf_counter()
ts = [threading.Thread(target=worker, args=(w,)) for w in range(a.workers)]
[t.start() for t in ts]
[t.join() for t in ts]
emit(dict(arm=a.arm, mode="conc-bin", bin=tag, workers=a.workers, n=a.per_bin,
wall_s=round(time.perf_counter() - t0, 3), gpu_start=snap0, gpu_end=gpu_snapshot()))
out.close()
if __name__ == "__main__":
main()