SPB Git

spb/localvm-research Public License

Running LLMs larger than memory on a consumer Mac — falsification-driven research: margin-gated deferred refinement, out-of-core verification on Apple Silicon. TR-01 published.

Python 63.2% JavaScript 23.5% CSS 11.8% Shell 0.9% Makefile 0.5%
7.4 KB · 175 lines python
Raw Blame History
1#!/usr/bin/env python32# =============================================================================3#  Project   : localvm-research4#  File      : experiments/candidate_01/benchmark_scale.py5#  Purpose   : Scale run — 32B model whose q8 does NOT fit beside the resident6#              base: q4 resident + layer-streamed q8 verification sweeps7#  Author    : Simon-Pierre Boucher8#  Contact   : contact@spboucher.ai9#  Created   : 2026-08-1210#  Modified  : 2026-08-1211#  Platform  : macOS / Apple Silicon (arm64) — MLX / Metal12#  License   : All rights reserved (research code)13# =============================================================================14"""Candidate-01 scale benchmark (the regime the architecture exists for).1516Qwen3-32B on a 48 GB Mac: q4 (17.5 GB) resident; q8 (34.8 GB) cannot be17co-resident — sweeps stream it layer-by-layer from SSD (StreamingVerifier).18Baseline: pure q4 (the only real alternative on this machine). Quality judged19by Qwen3-8B-bf16 (independent judge; the 32B bf16 obviously cannot run).2021Usage:22    .venv/bin/python benchmark_scale.py [--per-domain 2] [--max-tokens 96]23"""2425from __future__ import annotations2627import argparse28import gc29import json30import sys31import time32from datetime import datetime, timezone33from pathlib import Path3435import mlx.core as mx36from huggingface_hub import snapshot_download37from mlx_lm import load3839REPO_ROOT = Path(__file__).resolve().parents[2]40sys.path.insert(0, str(REPO_ROOT / "benchmarks"))41sys.path.insert(0, str(Path(__file__).parent / "implementation"))42from hardware_manifest import collect_manifest  # noqa: E40243from runtime import generate_deferred  # noqa: E40244from streaming_verifier import StreamingVerifier  # noqa: E4024546Q4_REPO = "mlx-community/Qwen3-32B-4bit"47Q8_REPO = "mlx-community/Qwen3-32B-8bit"48JUDGE_REPO = "mlx-community/Qwen3-8B-bf16"495051def greedy_baseline(model, tokenizer, prompt_ids, max_tokens):52    from mlx_lm.models.cache import make_prompt_cache5354    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            break62        tokens.append(nxt)63        inp = mx.array([[nxt]])64    return tokens, time.perf_counter() - t0656667def 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()7576    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"]7980    print("loading q4 resident …", flush=True)81    base_model, tokenizer = load(q4_path)82    verifier = StreamingVerifier(q8_path)83    q8_bytes = verifier.weight_bytes84    print(f"q8 checkpoint (streamed): {q8_bytes/1e9:.1f} GB", flush=True)8586    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)})9293    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)99100    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.deferred113                agg["sweeps"] += st.sweeps; agg["rollbacks"] += st.rollbacks114                agg["sweep_s"] += st.sweep_time_s; agg["gen_s"] += st.gen_time_s115                agg["logical_bytes"] += st.sweep_logical_bytes116                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}"] = outs130            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)133134    print("freeing 32B models; loading 8B bf16 judge …", flush=True)135    del base_model, verifier136    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                continue144            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)153154    ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")155    out_dir = REPO_ROOT / "results" / "candidate_01_scale32b" / ts156    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'}")171172173if __name__ == "__main__":174    main()175