guychuk commited on
Commit
f902f39
·
verified ·
1 Parent(s): 46b0675

Serving-path speedup: state-KV LRU cache + suffix diet (see inference/bench_serve.py)

Browse files
Files changed (1) hide show
  1. 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}")