#!/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()