File size: 14,021 Bytes
11ecc5b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
#!/usr/bin/env python3
"""
Make bge-m3 fast enough to serve from a CPU -- and PROVE it still retrieves.

THE MEASURED PROBLEM
    P50 143.3 ms · P70 152.1 ms · P100 229.9 ms   (laptop CPU, batch of 1)
P100 alone is 30 ms over the 200 ms budget before retrieval, the reader or TTS
run at all. Requirement 4 asks for P100 by name, so the tail is going in the
submission whether or not it flatters us.

WHY max_len IS THE WRONG KNOB
Queries are ~10-15 tokens. With batch=1 the tokenizer pads to the longest item
in the batch -- which is the query itself -- so the sequence is ALREADY short.
The 144 ms is 568M parameters being multiplied, not a long sequence. Dropping
max_len 192 -> 128 changes almost nothing, because nothing was reaching 192.

WHAT ACTUALLY HELPS, IN ORDER OF SAFETY

  1 THREADS      torch often defaults to fewer threads than the machine has.
                 Free, exact, zero risk to output. Try it first.

                 MEASURED IN A SEPARATE PROCESS PER CONFIG. torch.set_num_threads()
                 only takes effect reliably BEFORE the intra-op thread pool is
                 built. Calling it in a loop, after the model has already run,
                 does not re-partition the pool -- it produces numbers like
                 "threads=2 is slower than threads=1", which is physically
                 impossible and is exactly what the first version of this script
                 reported. One process per config, or the measurement is fiction.
  2 INT8 DYNAMIC Quantises Linear weights to int8, activations on the fly.
                 Typically 2-3x on CPU transformers. Changes the numbers
                 slightly -- which is why step 3 exists.
  3 AGREEMENT    THE STEP NOBODY RUNS. A quantised encoder that returns
                 different neighbours has not been optimised, it has been
                 broken, and a cosine similarity of 0.999 does not prove the
                 TOP-K ORDER survived. This measures top-k overlap and rank
                 correlation against the fp32 model on real queries, and refuses
                 to recommend int8 if agreement drops.

The index vectors stay fp32/fp16 bge-m3 -- only the QUERY encoder is quantised.
That asymmetry is fine if and only if agreement holds, which is exactly what
gets measured here rather than assumed.

  python scripts/optimize_cpu.py --n 60
"""
from __future__ import annotations

import argparse
import json
import os
import statistics as st
import sys
import time
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from scripts.bench_latency import QUERIES, pct  # noqa: E402
from src.schema_utils import default_root  # noqa: E402


# One config per interpreter. See the module docstring for why.
_CHILD = r"""
import json, sys, time, os
nt, model, max_len, n, reps = json.loads(sys.argv[1])
import torch
torch.set_num_threads(nt)          # BEFORE any tensor work
sys.path.insert(0, os.getcwd())
from src.evaluate_retrieval import Embedder
from scripts.bench_latency import QUERIES, pct
texts = [q for lg in QUERIES for q in QUERIES[lg]]
emb = Embedder(model, "cpu", 1, max_len)
for _ in range(5):
    emb.encode([texts[0]])
best = None
for _rep in range(reps):           # repeat the whole sweep; keep the cleanest
    ms = []
    for i in range(n):
        t0 = time.perf_counter(); emb.encode([texts[i % len(texts)]])
        ms.append((time.perf_counter() - t0) * 1000)
    if best is None or pct(ms, .5) < best[0]:
        best = (pct(ms, .5), pct(ms, .7), pct(ms, 1.0), sum(ms) / len(ms))
print(json.dumps({"p50": round(best[0], 1), "p70": round(best[1], 1),
                  "p100": round(best[2], 1), "mean": round(best[3], 1),
                  "n": n, "reps": reps}))
"""


def _run_isolated(nt, model, max_len, texts, n, reps):
    import subprocess
    import sys as _s
    try:
        out = subprocess.run(
            [_s.executable, "-c", _CHILD,
             json.dumps([nt, model, max_len, n, reps])],
            capture_output=True, text=True, timeout=900, cwd=str(Path.cwd()))
        line = [l for l in out.stdout.strip().splitlines() if l.startswith("{")]
        if not line:
            return None
        d = json.loads(line[-1])
        d["config"] = f"threads={nt}"
        return d
    except Exception:
        return None


def timeit(fn, texts, n, warmup=5):
    for _ in range(warmup):
        fn([texts[0]])
    out = []
    for i in range(n):
        t0 = time.perf_counter()
        fn([texts[i % len(texts)]])
        out.append((time.perf_counter() - t0) * 1000)
    return out


def summarise(label, ms, budget):
    d = {"config": label, "n": len(ms), "mean": round(st.mean(ms), 1),
         "p50": round(pct(ms, .5), 1), "p70": round(pct(ms, .7), 1),
         "p100": round(pct(ms, 1.0), 1)}
    d["p100_within_budget"] = d["p100"] <= budget
    print(f"  {label:28s} P50 {d['p50']:7.1f}  P70 {d['p70']:7.1f}  "
          f"P100 {d['p100']:7.1f}  {'OK' if d['p100_within_budget'] else 'OVER'}")
    return d


def main() -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--root", type=Path, default=None)
    ap.add_argument("--n", type=int, default=60)
    ap.add_argument("--budget-ms", type=float, default=200.0)
    ap.add_argument("--reserve-ms", type=float, default=15.0,
                    help="ms to leave for retrieval + reader + overhead")
    ap.add_argument("--min-topk-agreement", type=float, default=0.95)
    ap.add_argument("--reps", type=int, default=3,
                    help="repeat each config N times and keep the cleanest run — "
                         "a laptop under load produces a fake tail")
    args = ap.parse_args()

    import torch
    if torch.cuda.is_available():
        print("!! CUDA is available — this script is for CPU serving. Nothing to do.")
        return 0

    root = args.root.expanduser().resolve() if args.root else default_root()
    man_p = root / "index" / "manifest.json"
    man = json.loads(man_p.read_text()) if man_p.exists() else {}
    model = man.get("model")
    if model is None:
        hits = list((root / "hf_cache" / "hub").glob("models--BAAI--bge-m3/snapshots/*"))
        model = str(hits[0]) if hits else "BAAI/bge-m3"
    max_len = man.get("max_len", 192)
    texts = [q for lg in QUERIES for q in QUERIES[lg]]
    target = args.budget_ms - args.reserve_ms

    print(f"==> model {model}\n==> cpus {os.cpu_count()}  torch threads "
          f"{torch.get_num_threads()}")
    print(f"==> P100 must land under {target:.0f} ms "
          f"({args.budget_ms:.0f} budget − {args.reserve_ms:.0f} reserved)\n")

    results = []

    # ---- 1. threads -------------------------------------------------
    print(f"{'='*72}\n1. THREAD COUNT — one SUBPROCESS per config\n{'='*72}")
    print("  Each config runs in a fresh interpreter: set_num_threads() must be")
    print("  called before the thread pool exists, or the result is meaningless.\n")

    cand = sorted({1, 2, 4, 8, os.cpu_count() or 4, max(1, (os.cpu_count() or 4) // 2)})
    cand = [c for c in cand if c <= (os.cpu_count() or 4)]
    best_threads, best_p100, best_p50 = None, None, None
    for nt in cand:
        d = _run_isolated(nt, model, max_len, texts, args.n, args.reps)
        if d is None:
            print(f"  threads={nt:<3d} FAILED to measure")
            continue
        d["threads"] = nt
        results.append(d)
        stable = d["p100"] / max(1e-9, d["p50"])
        flag = "  <-- NOISY, machine not idle" if stable > 2.5 else ""
        print(f"  threads={nt:<3d} P50 {d['p50']:7.1f}  P70 {d['p70']:7.1f}  "
              f"P100 {d['p100']:7.1f}  spread {stable:4.2f}x{flag}")
        if best_p100 is None or d["p100"] < best_p100:
            best_threads, best_p100, best_p50 = nt, d["p100"], d["p50"]

    if best_threads is None:
        print("\n  could not measure any config — is torch importable in a subprocess?")
        return 1

    noisy = [d for d in results if d["p100"] / max(1e-9, d["p50"]) > 2.5]
    if noisy:
        print(f"\n  !! {len(noisy)} config(s) had a P100/P50 spread over 2.5x.")
        print("     A fixed ~15-token forward pass has no inherent tail that big, so")
        print("     something else was using the CPU. Close other work and re-run;")
        print("     do NOT make the pod-vs-laptop decision on these numbers.")

    print(f"\n  best: threads={best_threads}  P50 {best_p50:.1f}  P100 {best_p100:.1f} ms")
    torch.set_num_threads(best_threads)

    if best_p100 <= target:
        print(f"\n  >> DONE. Threads alone bring P100 inside {target:.0f} ms.")
        print(f"     Set torch.set_num_threads({best_threads}) at startup. No")
        print("     quantisation needed, so retrieval is bit-identical.")
        _save(root, results, best_threads, None, None)
        return 0

    # ---- 2. int8 ----------------------------------------------------
    print(f"\n{'='*72}\n2. INT8 DYNAMIC QUANTISATION\n{'='*72}")
    from src.evaluate_retrieval import Embedder
    base = Embedder(model, "cpu", 1, max_len)
    fp32 = timeit(base.encode, texts, args.n)
    print(f"  fp32 re-measured in THIS process for a like-for-like comparison:")
    d32 = summarise("fp32 (same process)", fp32, target)
    try:
        qmodel = torch.quantization.quantize_dynamic(
            base.model, {torch.nn.Linear}, dtype=torch.qint8)
    except Exception as exc:
        print(f"  quantisation unavailable: {exc}")
        _save(root, results, best_threads, None, None)
        return 1

    class QEmb:
        def __init__(self, tok, m, dev, ml):
            self.tok, self.model, self.device, self.max_len = tok, m, dev, ml

        def encode(self, ts):
            with torch.inference_mode():
                enc = self.tok(ts, padding=True, truncation=True,
                               max_length=self.max_len, return_tensors="pt")
                h = self.model(**enc).last_hidden_state[:, 0]
                return torch.nn.functional.normalize(h, dim=-1)

    q = QEmb(base.tok, qmodel, "cpu", max_len)
    d = summarise("int8 dynamic", timeit(q.encode, texts, args.n), target)
    d["threads"], d["quantised"] = best_threads, True
    results.append(d)
    # Compare against fp32 measured in the SAME process, not against a
    # subprocess number -- otherwise the "speedup" mixes two environments.
    speedup = d32["p50"] / max(1e-9, d["p50"])
    print(f"\n  speedup vs fp32 in this same process, at P50: {speedup:.2f}x")
    if speedup < 1.0:
        print("  int8 came out SLOWER. That happens when the CPU lacks the int8")
        print("  kernels torch expects, or when threads are oversubscribed. It is")
        print("  a real result: do not ship it.")

    # ---- 3. agreement ----------------------------------------------
    print(f"\n{'='*72}\n3. DOES IT STILL RETRIEVE THE SAME THINGS?\n{'='*72}")
    print("  Cosine between the two query vectors is NOT the test — what matters")
    print("  is whether the same chunks come back in the same order.\n")

    import numpy as np
    lang = next((lg for lg in QUERIES if (root / "index" / f"{lg}__FW.vecs.npy").exists()),
                None)
    if lang is None:
        print("  !! no index found — cannot verify agreement.")
        print("     DO NOT ship int8 on the strength of the speedup alone.")
        _save(root, results, best_threads, d, None)
        return 2

    vecs = np.load(root / "index" / f"{lang}__FW.vecs.npy").astype("float32")
    qs = QUERIES[lang]
    K = 5
    overlaps, top1, cosines = [], 0, []
    for text in qs:
        a = base.encode([text]).cpu().numpy().astype("float32")[0]
        b = q.encode([text]).cpu().numpy().astype("float32")[0]
        cosines.append(float(a @ b))
        ra = np.argsort(-(vecs @ a))[:K]
        rb = np.argsort(-(vecs @ b))[:K]
        overlaps.append(len(set(ra.tolist()) & set(rb.tolist())) / K)
        top1 += int(ra[0] == rb[0])

    agree = sum(overlaps) / len(overlaps)
    t1 = top1 / len(qs)
    print(f"  query-vector cosine   {st.mean(cosines):.5f}   <- looks perfect, proves little")
    print(f"  top-{K} set overlap     {agree:.3f}")
    print(f"  top-1 identical       {t1:.3f}   ({top1}/{len(qs)} queries)")

    ok = agree >= args.min_topk_agreement and d["p100_within_budget"]
    print(f"\n{'='*72}\nVERDICT\n{'='*72}")
    if ok:
        print(f"  USE INT8. P100 {d['p100']:.1f} ms is inside {target:.0f} ms and top-{K}")
        print(f"  agreement is {agree:.3f} (>= {args.min_topk_agreement}).")
        print("  Quantise the QUERY encoder only; leave the index as built.")
        print("  State it in the writeup — a quantised query encoder is a real")
        print("  engineering decision, not something to hide.")
    elif not d["p100_within_budget"]:
        print(f"  STILL OVER. P100 {d['p100']:.1f} ms > {target:.0f} ms even quantised.")
        print("  Serve from the GPU pod, or accept and REPORT a P100 over 200 ms.")
        print("  Reporting an honest miss beats a number nobody can reproduce.")
    else:
        print(f"  DO NOT USE INT8. Top-{K} agreement {agree:.3f} < "
              f"{args.min_topk_agreement}.")
        print("  It is faster and it retrieves different passages — that is a")
        print("  regression wearing a speedup's clothes.")

    _save(root, results, best_threads, d, {"topk_overlap": round(agree, 4),
                                           "top1_identical": round(t1, 4),
                                           "mean_cosine": round(st.mean(cosines), 5),
                                           "recommend_int8": bool(ok)})
    return 0 if ok else 3


def _save(root, results, threads, int8, agreement):
    out = root / "results" / "cpu_optimisation.json"
    out.parent.mkdir(parents=True, exist_ok=True)
    out.write_text(json.dumps({"configs": results, "best_threads": threads,
                               "int8": int8, "agreement": agreement}, indent=2))
    print(f"\n==> wrote {out}")


if __name__ == "__main__":
    raise SystemExit(main())