File size: 5,047 Bytes
a484e22
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
"""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()