|
1 |
+#!/usr/bin/env python3 |
|
2 |
+# ============================================================================= |
|
3 |
+# Project : localvm-research |
|
4 |
+# File : experiments/candidate_01/benchmark_scale.py |
|
5 |
+# Purpose : Scale run — 32B model whose q8 does NOT fit beside the resident |
|
6 |
+# base: q4 resident + layer-streamed q8 verification sweeps |
|
7 |
+# Author : Simon-Pierre Boucher |
|
8 |
+# Contact : contact@spboucher.ai |
|
9 |
+# Created : 2026-08-12 |
|
10 |
+# Modified : 2026-08-12 |
|
11 |
+# Platform : macOS / Apple Silicon (arm64) — MLX / Metal |
|
12 |
+# License : All rights reserved (research code) |
|
13 |
+# ============================================================================= |
|
14 |
+"""Candidate-01 scale benchmark (the regime the architecture exists for). |
|
15 |
+ |
|
16 |
+Qwen3-32B on a 48 GB Mac: q4 (17.5 GB) resident; q8 (34.8 GB) cannot be |
|
17 |
+co-resident — sweeps stream it layer-by-layer from SSD (StreamingVerifier). |
|
18 |
+Baseline: pure q4 (the only real alternative on this machine). Quality judged |
|
19 |
+by Qwen3-8B-bf16 (independent judge; the 32B bf16 obviously cannot run). |
|
20 |
+ |
|
21 |
+Usage: |
|
22 |
+ .venv/bin/python benchmark_scale.py [--per-domain 2] [--max-tokens 96] |
|
23 |
+""" |
|
24 |
+ |
|
25 |
+from __future__ import annotations |
|
26 |
+ |
|
27 |
+import argparse |
|
28 |
+import gc |
|
29 |
+import json |
|
30 |
+import sys |
|
31 |
+import time |
|
32 |
+from datetime import datetime, timezone |
|
33 |
+from pathlib import Path |
|
34 |
+ |
|
35 |
+import mlx.core as mx |
|
36 |
+from huggingface_hub import snapshot_download |
|
37 |
+from mlx_lm import load |
|
38 |
+ |
|
39 |
+REPO_ROOT = Path(__file__).resolve().parents[2] |
|
40 |
+sys.path.insert(0, str(REPO_ROOT / "benchmarks")) |
|
41 |
+sys.path.insert(0, str(Path(__file__).parent / "implementation")) |
|
42 |
+from hardware_manifest import collect_manifest # noqa: E402 |
|
43 |
+from runtime import generate_deferred # noqa: E402 |
|
44 |
+from streaming_verifier import StreamingVerifier # noqa: E402 |
|
45 |
+ |
|
46 |
+Q4_REPO = "mlx-community/Qwen3-32B-4bit" |
|
47 |
+Q8_REPO = "mlx-community/Qwen3-32B-8bit" |
|
48 |
+JUDGE_REPO = "mlx-community/Qwen3-8B-bf16" |
|
49 |
+ |
|
50 |
+ |
|
51 |
+def greedy_baseline(model, tokenizer, prompt_ids, max_tokens): |
|
52 |
+ from mlx_lm.models.cache import make_prompt_cache |
|
53 |
+ |
|
54 |
+ cache = make_prompt_cache(model) |
|
55 |
+ tokens = [] |
|
56 |
+ inp = mx.array(list(prompt_ids))[None] |
|
57 |
+ t0 = time.perf_counter() |
|
58 |
+ for _ in range(max_tokens): |
|
59 |
+ nxt = int(mx.argmax(model(inp, cache=cache)[0, -1]).item()) |
|
60 |
+ if nxt == tokenizer.eos_token_id: |
|
61 |
+ break |
|
62 |
+ tokens.append(nxt) |
|
63 |
+ inp = mx.array([[nxt]]) |
|
64 |
+ return tokens, time.perf_counter() - t0 |
|
65 |
+ |
|
66 |
+ |
|
67 |
+def main() -> None: |
|
68 |
+ ap = argparse.ArgumentParser() |
|
69 |
+ ap.add_argument("--per-domain", type=int, default=2) |
|
70 |
+ ap.add_argument("--max-tokens", type=int, default=96) |
|
71 |
+ ap.add_argument("--window", type=int, default=32) |
|
72 |
+ ap.add_argument("--taus", default="2.0") |
|
73 |
+ ap.add_argument("--modes", default="margin,verify-all") |
|
74 |
+ args = ap.parse_args() |
|
75 |
+ |
|
76 |
+ q4_path = snapshot_download(Q4_REPO) |
|
77 |
+ q8_path = snapshot_download(Q8_REPO) |
|
78 |
+ domains = json.loads((REPO_ROOT / "benchmarks/datasets/eval_prompts.json").read_text())["domains"] |
|
79 |
+ |
|
80 |
+ print("loading q4 resident …", flush=True) |
|
81 |
+ base_model, tokenizer = load(q4_path) |
|
82 |
+ verifier = StreamingVerifier(q8_path) |
|
83 |
+ q8_bytes = verifier.weight_bytes |
|
84 |
+ print(f"q8 checkpoint (streamed): {q8_bytes/1e9:.1f} GB", flush=True) |
|
85 |
+ |
|
86 |
+ prompts = [] |
|
87 |
+ for domain, plist in domains.items(): |
|
88 |
+ for prompt in plist[: args.per_domain]: |
|
89 |
+ ids = tokenizer.apply_chat_template( |
|
90 |
+ [{"role": "user", "content": prompt}], add_generation_prompt=True) |
|
91 |
+ prompts.append({"domain": domain, "ids": list(ids)}) |
|
92 |
+ |
|
93 |
+ print("baseline: pure q4 …", flush=True) |
|
94 |
+ q4_out, q4_times = [], [] |
|
95 |
+ for k, p in enumerate(prompts): |
|
96 |
+ toks, dt = greedy_baseline(base_model, tokenizer, p["ids"], args.max_tokens) |
|
97 |
+ q4_out.append(toks); q4_times.append((len(toks), dt)) |
|
98 |
+ print(f" {k+1}/{len(prompts)} ({len(toks)} tok, {len(toks)/dt:.1f} tok/s)", flush=True) |
|
99 |
+ |
|
100 |
+ outputs = {"pure_q4": q4_out} |
|
101 |
+ runs = [] |
|
102 |
+ for mode in args.modes.split(","): |
|
103 |
+ for tau in ([float(x) for x in args.taus.split(",")] if mode == "margin" else [2.0]): |
|
104 |
+ print(f"runtime: mode={mode} tau={tau} W={args.window} …", flush=True) |
|
105 |
+ outs, agg = [], {"tokens": 0, "deferred": 0, "sweeps": 0, "rollbacks": 0, |
|
106 |
+ "sweep_s": 0.0, "gen_s": 0.0, "logical_bytes": 0, "io_s": []} |
|
107 |
+ for k, p in enumerate(prompts): |
|
108 |
+ toks, st = generate_deferred( |
|
109 |
+ base_model, verifier, tokenizer, p["ids"], |
|
110 |
+ args.max_tokens, tau, args.window, mode, q8_bytes) |
|
111 |
+ outs.append(toks) |
|
112 |
+ agg["tokens"] += st.tokens_out; agg["deferred"] += st.deferred |
|
113 |
+ agg["sweeps"] += st.sweeps; agg["rollbacks"] += st.rollbacks |
|
114 |
+ agg["sweep_s"] += st.sweep_time_s; agg["gen_s"] += st.gen_time_s |
|
115 |
+ agg["logical_bytes"] += st.sweep_logical_bytes |
|
116 |
+ print(f" {k+1}/{len(prompts)} ({st.tokens_out} tok, {st.sweeps} sweeps, " |
|
117 |
+ f"{st.rollbacks} rollbacks, last sweep io {verifier.last_sweep_io_s:.1f}s)", |
|
118 |
+ flush=True) |
|
119 |
+ n = max(agg["tokens"], 1) |
|
120 |
+ runs.append({ |
|
121 |
+ "mode": mode, "tau": tau, "window": args.window, |
|
122 |
+ "tokens_per_s": n / (agg["gen_s"] + agg["sweep_s"]), |
|
123 |
+ "deferral_rate": agg["deferred"] / n, |
|
124 |
+ "rollback_rate": agg["rollbacks"] / n, |
|
125 |
+ "sweep_latency_s_mean": agg["sweep_s"] / max(agg["sweeps"], 1), |
|
126 |
+ "logical_verify_bytes_per_token": agg["logical_bytes"] / n, |
|
127 |
+ "raw": {k: v for k, v in agg.items() if k != "io_s"}, |
|
128 |
+ }) |
|
129 |
+ outputs[f"{mode}_tau{tau}"] = outs |
|
130 |
+ r = runs[-1] |
|
131 |
+ print(f" tok/s={r['tokens_per_s']:.2f} sweepLat={r['sweep_latency_s_mean']:.1f}s " |
|
132 |
+ f"GB/token(logical)={r['logical_verify_bytes_per_token']/1e9:.2f}", flush=True) |
|
133 |
+ |
|
134 |
+ print("freeing 32B models; loading 8B bf16 judge …", flush=True) |
|
135 |
+ del base_model, verifier |
|
136 |
+ gc.collect(); mx.clear_cache() |
|
137 |
+ judge, _ = load(JUDGE_REPO) |
|
138 |
+ quality = {} |
|
139 |
+ for name, outs in outputs.items(): |
|
140 |
+ vals = [] |
|
141 |
+ for p, toks in zip(prompts, outs): |
|
142 |
+ if len(toks) < 2: |
|
143 |
+ continue |
|
144 |
+ full = p["ids"] + list(toks) |
|
145 |
+ logits = judge(mx.array(full)[None])[0] |
|
146 |
+ sel = logits[len(p["ids"]) - 1 : len(full) - 1].astype(mx.float32) |
|
147 |
+ lp = sel - mx.logsumexp(sel, axis=-1, keepdims=True) |
|
148 |
+ tok_lp = mx.take_along_axis(lp, mx.array(toks)[:, None], axis=-1) |
|
149 |
+ mx.eval(tok_lp) |
|
150 |
+ vals.append(float(mx.mean(tok_lp).item())) |
|
151 |
+ quality[name] = {"mean_logprob_8b_judge": sum(vals) / len(vals), "n": len(vals)} |
|
152 |
+ print(f" {name:>18}: {quality[name]['mean_logprob_8b_judge']:.4f}", flush=True) |
|
153 |
+ |
|
154 |
+ ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ") |
|
155 |
+ out_dir = REPO_ROOT / "results" / "candidate_01_scale32b" / ts |
|
156 |
+ out_dir.mkdir(parents=True) |
|
157 |
+ (out_dir / "results.json").write_text(json.dumps({ |
|
158 |
+ "experiment": "candidate_01_scale32b", |
|
159 |
+ "author": "Simon-Pierre Boucher", |
|
160 |
+ "contact": "contact@spboucher.ai", |
|
161 |
+ "manifest": collect_manifest(), |
|
162 |
+ "config": vars(args), |
|
163 |
+ "models": {"base": Q4_REPO, "verify": Q8_REPO, "judge": JUDGE_REPO}, |
|
164 |
+ "q8_streamed_bytes": q8_bytes, |
|
165 |
+ "baseline_pure_q4_tokens_per_s": |
|
166 |
+ sum(t for t, _ in q4_times) / max(sum(d for _, d in q4_times), 1e-9), |
|
167 |
+ "runs": runs, |
|
168 |
+ "quality_8b_judge": quality, |
|
169 |
+ }, indent=2)) |
|
170 |
+ print(f"\nwrote {out_dir / 'results.json'}") |
|
171 |
+ |
|
172 |
+ |
|
173 |
+if __name__ == "__main__": |
|
174 |
+ main() |