SpecMem / harness /phase4_throughput.py
inweriok's picture
Initial release: SpecMem harness (code only, credentials-free)
a484e22 verified
Raw
History Blame Contribute Delete
5.05 kB
"""Phase 4.2 — concurrency / throughput under real load.
Drives CONCURRENT (not sequential) requests against the serving stack and
measures how aggregate throughput scales with concurrency, per arm, to answer:
does the memory arm's decode-step advantage hold, shrink, or reverse under
contention?
Method: reuse the faithful external spec-decode accounting (Phase 4.1) — each
arm's per-request work is decoding ``T - L`` tokens (accept length ``L`` from
the real served targets). For each arm and concurrency level ``C`` we fire the
sampled requests through a ``C``-worker pool against the live sglang server
(real batched serving, ``ignore_eos`` forces exact token counts) and record
aggregate tokens/s and per-request latency. The no-memory arm decodes ~``T``
tokens/request; ours decodes ~``T-L`` with ``L`` large, so under batching it
should sustain higher throughput.
Run from the repo root: python -m harness.phase4_throughput --url http://localhost:30000/v1
"""
from __future__ import annotations
import argparse
import json
import random
import time
from collections import defaultdict
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import requests
from . import metrics
from .phase4_wallclock import _collect_points, MODEL_PATH
ROOT = Path(__file__).resolve().parent.parent
RESULTS = ROOT / "results"
ARMS = ["no_memory", "static_global", "personal_memory"]
def _gen(chat, prompt, ntok, T):
t0 = time.perf_counter()
r = requests.post(chat, json={"model": "gpt-oss-120b", "temperature": 0.0,
"max_tokens": ntok, "ignore_eos": True,
"messages": [{"role": "user",
"content": prompt}]}, timeout=300)
dt = time.perf_counter() - t0
r.raise_for_status()
return dt, T
def _run_arm(chat, jobs, concurrency):
"""Fire (prompt, decode_ntok, target_T) jobs through a C-worker pool.
Each request PRODUCES its full target (``T`` tokens) whether by decoding or
by accepting the draft's prefix; the memory arm decodes fewer (``T-L``) but
still produces ``T``, so the correct throughput counts effective OUTPUT
tokens (``T``), not decode work. We also report requests/s."""
lat, out_toks = [], 0
t0 = time.perf_counter()
with ThreadPoolExecutor(max_workers=concurrency) as ex:
for dt, T in ex.map(lambda j: _gen(chat, j[0], j[1], j[2]), jobs):
lat.append(dt * 1000)
out_toks += T
wall = time.perf_counter() - t0
lat.sort()
return {"eff_output_tokens_per_s": round(out_toks / wall, 1),
"requests_per_s": round(len(jobs) / wall, 2),
"wall_s": round(wall, 2),
"req_p50_ms": round(lat[len(lat) // 2], 1),
"req_p95_ms": round(lat[min(len(lat) - 1, int(0.95 * len(lat)))], 1)}
def main():
p = argparse.ArgumentParser()
p.add_argument("--url", default="http://localhost:30000/v1")
p.add_argument("--domains", nargs="+",
default=["airline", "retail", "telecom"])
p.add_argument("--concurrency", nargs="+", type=int, default=[1, 8, 32])
p.add_argument("--sample", type=int, default=96)
p.add_argument("--seed", type=int, default=0)
args = p.parse_args()
chat = args.url.rstrip("/") + "/chat/completions"
metrics.get_tokenizer(MODEL_PATH)
pts = _collect_points(args.domains)
rng = random.Random(args.seed)
rng.shuffle(pts)
pts = [x for x in pts if x["T"] >= 2][:args.sample]
print(f"[throughput] {len(pts)} sampled requests", flush=True)
# precompute per-arm job list: (prompt, decode_ntok=T-L, target_T)
jobs = {a: [(x["query"][:1500], max(1, x["T"] - x[a]), x["T"]) for x in pts]
for a in ARMS}
mean_tok = {a: round(sum(n for _, n, _ in jobs[a]) / len(jobs[a]), 1)
for a in ARMS}
out = {"concurrency_levels": args.concurrency, "arms": ARMS,
"n_requests": len(pts), "mean_decode_tokens_per_req": mean_tok,
"results": defaultdict(dict)}
for C in args.concurrency:
for a in ARMS:
r = _run_arm(chat, jobs[a], C)
out["results"][str(C)][a] = r
print(f" C={C:2d} {a:16s} out={r['eff_output_tokens_per_s']:8.1f} tok/s "
f"req/s={r['requests_per_s']:.2f} p50={r['req_p50_ms']:.0f}ms "
f"p95={r['req_p95_ms']:.0f}ms", flush=True)
# personal over no_memory effective-output throughput at each C
out["personal_over_vanilla_throughput"] = {
str(C): round(out["results"][str(C)]["personal_memory"]["eff_output_tokens_per_s"] /
out["results"][str(C)]["no_memory"]["eff_output_tokens_per_s"], 3)
for C in args.concurrency}
out["results"] = dict(out["results"])
(RESULTS / "phase4_throughput.json").write_text(json.dumps(out, indent=2))
print("\npersonal/vanilla throughput ratio by concurrency:",
out["personal_over_vanilla_throughput"])
if __name__ == "__main__":
main()