Serving-path speedup: state-KV LRU cache + suffix diet (see inference/bench_serve.py)
Browse files- inference/bench_serve.py +161 -0
inference/bench_serve.py
ADDED
|
@@ -0,0 +1,161 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Profile a single Typical decision on the serving path, and produce the before/after
|
| 2 |
+
latency table (p50/p95 ms) across models x state length x K x {cold, warm}.
|
| 3 |
+
|
| 4 |
+
"cold" = the state text is different on every call (state-cache miss every time, i.e. the
|
| 5 |
+
pre-optimisation baseline path: full re-encode of the state prefix). "warm" = the same
|
| 6 |
+
state text is reused across calls with a new question each time (state-cache hit after the
|
| 7 |
+
first call) -- the path Typical.choice/score/noul/decide take automatically via the LRU in
|
| 8 |
+
core.py.
|
| 9 |
+
|
| 10 |
+
Usage:
|
| 11 |
+
uv run --no-sync python inference/bench_serve.py --models typical-small --out_dir runs/serve_bench
|
| 12 |
+
uv run --no-sync python inference/bench_serve.py --models typical-small --profile # chrome trace + key_averages table
|
| 13 |
+
"""
|
| 14 |
+
import argparse
|
| 15 |
+
import json
|
| 16 |
+
import os
|
| 17 |
+
import statistics
|
| 18 |
+
import sys
|
| 19 |
+
import time
|
| 20 |
+
|
| 21 |
+
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
| 22 |
+
import torch
|
| 23 |
+
from typical import Typical
|
| 24 |
+
|
| 25 |
+
MODELS = {
|
| 26 |
+
"typical-small": ("guychuk/pcdm-runs", "typical-small/best.pt"), # 1.7B, Qwen3
|
| 27 |
+
"typical-medium": ("guychuk/pcdm-runs", "typical-medium/best.pt"), # 4B, Qwen3
|
| 28 |
+
"tm2": ("guychuk/pcdm-runs", "tm2/best.pt"), # 4B, Qwen3.5
|
| 29 |
+
}
|
| 30 |
+
STATE_LENS = (256, 1024)
|
| 31 |
+
KS = (2, 4, 10, 32)
|
| 32 |
+
FILLER = ("Ticket log entry: the customer reported an issue with their recent order and "
|
| 33 |
+
"followed up twice asking for a status update on the resolution. ")
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def make_state(tok, n_tokens: int) -> str:
|
| 37 |
+
"""Filler text truncated to exactly n_tokens tokens (no special tokens)."""
|
| 38 |
+
reps = n_tokens // len(tok(FILLER, add_special_tokens=False)["input_ids"]) + 3
|
| 39 |
+
ids = tok(FILLER * reps, add_special_tokens=False)["input_ids"][:n_tokens]
|
| 40 |
+
return tok.decode(ids)
|
| 41 |
+
|
| 42 |
+
|
| 43 |
+
def make_labels(k: int) -> list[str]:
|
| 44 |
+
return [f"candidate option {j} short description text" for j in range(k)]
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def make_items(tok, n: int = 20) -> list[dict]:
|
| 48 |
+
"""20 JevBench-style items spanning the (state_len, K) grid."""
|
| 49 |
+
grid = [(sl, k) for sl in STATE_LENS for k in KS]
|
| 50 |
+
items = []
|
| 51 |
+
for i in range(n):
|
| 52 |
+
sl, k = grid[i % len(grid)]
|
| 53 |
+
items.append({"state_len": sl, "K": k, "state": make_state(tok, sl),
|
| 54 |
+
"question": f"Question {i}: which option best applies?",
|
| 55 |
+
"labels": make_labels(k)})
|
| 56 |
+
return items
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def _sync():
|
| 60 |
+
if torch.cuda.is_available():
|
| 61 |
+
torch.cuda.synchronize()
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def timeit_ms(fn, n: int = 12) -> dict:
|
| 65 |
+
times = []
|
| 66 |
+
for _ in range(n):
|
| 67 |
+
_sync()
|
| 68 |
+
t0 = time.perf_counter()
|
| 69 |
+
fn()
|
| 70 |
+
_sync()
|
| 71 |
+
times.append((time.perf_counter() - t0) * 1000)
|
| 72 |
+
times.sort()
|
| 73 |
+
return {"p50_ms": statistics.median(times), "p95_ms": times[min(len(times) - 1, int(0.95 * len(times)))],
|
| 74 |
+
"n": n}
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
def bench_model(name: str, repo_id: str, filename: str) -> dict:
|
| 78 |
+
m = Typical.from_pretrained(repo_id, filename=filename)
|
| 79 |
+
tok = m.tok
|
| 80 |
+
configs = [(sl, k) for sl in STATE_LENS for k in KS]
|
| 81 |
+
results = {"model": name, "backbone": m.head.backbone.model.config.name_or_path
|
| 82 |
+
if hasattr(m.head.backbone.model.config, "name_or_path") else None, "configs": []}
|
| 83 |
+
for sl, k in configs:
|
| 84 |
+
state = make_state(tok, sl)
|
| 85 |
+
question, labels = "Which option best applies?", make_labels(k)
|
| 86 |
+
nonce = [0]
|
| 87 |
+
|
| 88 |
+
def cold_call():
|
| 89 |
+
nonce[0] += 1
|
| 90 |
+
m.choice(f"[{nonce[0]}] " + state, question, labels)
|
| 91 |
+
|
| 92 |
+
def warm_call():
|
| 93 |
+
m.choice(state, question, labels)
|
| 94 |
+
|
| 95 |
+
warm_call() # warm up CUDA kernels/allocator, populate state cache for warm_call
|
| 96 |
+
cold_call() # warm up the cold path's own kernels too (same shapes, different state text)
|
| 97 |
+
cold = timeit_ms(cold_call)
|
| 98 |
+
warm = timeit_ms(warm_call)
|
| 99 |
+
print(f" [{name}] state={sl:5d} K={k:3d} cold p50={cold['p50_ms']:7.2f}ms p95={cold['p95_ms']:7.2f}ms"
|
| 100 |
+
f" warm p50={warm['p50_ms']:7.2f}ms p95={warm['p95_ms']:7.2f}ms"
|
| 101 |
+
f" speedup={cold['p50_ms']/max(warm['p50_ms'], 1e-6):.2f}x")
|
| 102 |
+
results["configs"].append({"state_tokens": sl, "K": k, "cold": cold, "warm": warm})
|
| 103 |
+
if torch.cuda.is_available():
|
| 104 |
+
results["peak_mem_mb"] = torch.cuda.max_memory_allocated() / 2**20
|
| 105 |
+
return results
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
def profile_one(name: str, repo_id: str, filename: str, out_dir: str, state_len: int = 256, k: int = 10):
|
| 109 |
+
"""PyTorch profiler over a single warm choice() call -- reports the split between
|
| 110 |
+
state-prefix forward, suffix forward, tokenisation, mask building, and cpu syncs (see
|
| 111 |
+
the record_function blocks in typical/native.py)."""
|
| 112 |
+
from torch.profiler import ProfilerActivity, profile
|
| 113 |
+
|
| 114 |
+
m = Typical.from_pretrained(repo_id, filename=filename)
|
| 115 |
+
tok = m.tok
|
| 116 |
+
state, question, labels = make_state(tok, state_len), "Which option best applies?", make_labels(k)
|
| 117 |
+
for _ in range(3):
|
| 118 |
+
m.choice(state, question, labels) # warmup (state cache hit after call 1)
|
| 119 |
+
activities = [ProfilerActivity.CPU] + ([ProfilerActivity.CUDA] if torch.cuda.is_available() else [])
|
| 120 |
+
with profile(activities=activities) as prof:
|
| 121 |
+
m.choice(state, question, labels)
|
| 122 |
+
sort_by = "self_cuda_time_total" if torch.cuda.is_available() else "self_cpu_time_total"
|
| 123 |
+
table = prof.key_averages().table(sort_by=sort_by, row_limit=20)
|
| 124 |
+
print(f"\n=== profile: {name} warm choice() state={state_len} K={k} ===\n{table}")
|
| 125 |
+
os.makedirs(out_dir, exist_ok=True)
|
| 126 |
+
trace_path = os.path.join(out_dir, f"{name}_trace.json")
|
| 127 |
+
prof.export_chrome_trace(trace_path)
|
| 128 |
+
print(f"trace written to {trace_path}")
|
| 129 |
+
|
| 130 |
+
# cold, for contrast
|
| 131 |
+
m._state_kv.clear()
|
| 132 |
+
with profile(activities=activities) as prof_cold:
|
| 133 |
+
m.choice(state + " [cold]", question, labels)
|
| 134 |
+
table_cold = prof_cold.key_averages().table(sort_by=sort_by, row_limit=20)
|
| 135 |
+
print(f"\n=== profile: {name} cold choice() state={state_len} K={k} ===\n{table_cold}")
|
| 136 |
+
|
| 137 |
+
|
| 138 |
+
if __name__ == "__main__":
|
| 139 |
+
ap = argparse.ArgumentParser()
|
| 140 |
+
ap.add_argument("--models", nargs="+", default=list(MODELS.keys()), choices=list(MODELS.keys()))
|
| 141 |
+
ap.add_argument("--out_dir", default="runs/serve_bench")
|
| 142 |
+
ap.add_argument("--profile", action="store_true", help="also run torch.profiler over one warm+cold call")
|
| 143 |
+
ap.add_argument("--tag", default="results", help="output json filename stem")
|
| 144 |
+
args = ap.parse_args()
|
| 145 |
+
os.makedirs(args.out_dir, exist_ok=True)
|
| 146 |
+
|
| 147 |
+
all_results = []
|
| 148 |
+
for name in args.models:
|
| 149 |
+
repo_id, filename = MODELS[name]
|
| 150 |
+
print(f"=== {name} ({repo_id}/{filename}) ===")
|
| 151 |
+
if torch.cuda.is_available():
|
| 152 |
+
torch.cuda.reset_peak_memory_stats()
|
| 153 |
+
res = bench_model(name, repo_id, filename)
|
| 154 |
+
all_results.append(res)
|
| 155 |
+
if args.profile:
|
| 156 |
+
profile_one(name, repo_id, filename, args.out_dir)
|
| 157 |
+
|
| 158 |
+
out_path = os.path.join(args.out_dir, f"{args.tag}.json")
|
| 159 |
+
with open(out_path, "w") as f:
|
| 160 |
+
json.dump(all_results, f, indent=2)
|
| 161 |
+
print(f"\nwrote {out_path}")
|