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%
9.0 KB · 232 lines python
Raw Blame History
1#!/usr/bin/env python32# =============================================================================3#  Project   : localvm-research4#  File      : experiments/candidate_01/benchmark.py5#  Purpose   : End-to-end evaluation of the margin-gated deferred-refinement6#              runtime vs pure-q4 / pure-q8 baselines7#  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 benchmark.1516Build once (downloads + quantizes):17    .venv/bin/python benchmark.py --build18Run:19    .venv/bin/python benchmark.py [--per-domain 4] [--max-tokens 128]20        [--window 32] [--taus 1.0,2.0]21"""2223from __future__ import annotations2425import argparse26import difflib27import json28import sys29import time30from datetime import datetime, timezone31from pathlib import Path3233import mlx.core as mx34from mlx_lm import load3536REPO_ROOT = Path(__file__).resolve().parents[2]37sys.path.insert(0, str(REPO_ROOT / "benchmarks"))38sys.path.insert(0, str(Path(__file__).parent / "implementation"))39from hardware_manifest import collect_manifest  # noqa: E40240from runtime import generate_deferred  # noqa: E4024142MODELS_DIR = Path(__file__).parent / "implementation" / "models"43HF_MODEL = "mlx-community/Qwen3-1.7B-bf16"444546def build() -> None:47    from mlx_lm import convert4849    for bits in (4, 8):50        out = MODELS_DIR / f"q{bits}"51        if out.exists():52            print(f"{out} exists, skipping")53            continue54        print(f"converting {HF_MODEL} → q{bits} …", flush=True)55        convert(HF_MODEL, mlx_path=str(out), quantize=True, q_bits=bits, q_group_size=64)56    print("build done")575859def dir_weight_bytes(d: Path) -> int:60    return sum(f.stat().st_size for f in d.glob("*.safetensors"))616263def greedy_baseline(model, tokenizer, prompt_ids, max_tokens):64    from mlx_lm.models.cache import make_prompt_cache6566    cache = make_prompt_cache(model)67    tokens = []68    inp = mx.array(list(prompt_ids))[None]69    t0 = time.perf_counter()70    for _ in range(max_tokens):71        logits = model(inp, cache=cache)72        nxt = int(mx.argmax(logits[0, -1]).item())73        if nxt == tokenizer.eos_token_id:74            break75        tokens.append(nxt)76        inp = mx.array([[nxt]])77    return tokens, time.perf_counter() - t0787980def judge_outputs(outputs_by_config: dict, prompts: list[dict]) -> dict:81    """Quality-level metric: mean per-token logprob of each config's generated82    continuation under the bf16 reference model (higher = better). Token-exact83    fidelity is incoherent on Metal (1.56%/token prefill/decode flips), so the84    judge scores usefulness of the text the system actually produced."""85    import gc8687    gc.collect(); mx.clear_cache()88    judge, _ = load(HF_MODEL)89    scores = {}90    for name, outs in outputs_by_config.items():91        vals = []92        for p, toks in zip(prompts, outs):93            if len(toks) < 2:94                continue95            full = p["ids"] + list(toks)96            logits = judge(mx.array(full)[None])[0]97            sel = logits[len(p["ids"]) - 1 : len(full) - 1].astype(mx.float32)98            logprobs = sel - mx.logsumexp(sel, axis=-1, keepdims=True)99            idx = mx.array(toks)100            tok_lp = mx.take_along_axis(logprobs, idx[:, None], axis=-1)101            mx.eval(tok_lp)102            vals.append(float(mx.mean(tok_lp).item()))103        scores[name] = {"mean_logprob_bf16": sum(vals) / len(vals), "n": len(vals)}104    del judge105    gc.collect(); mx.clear_cache()106    return scores107108109def fidelity(a: list[int], b: list[int]) -> float:110    """Similarity of two token sequences (difflib ratio — robust to length111    drift after divergence)."""112    if not a and not b:113        return 1.0114    return difflib.SequenceMatcher(None, a, b).ratio()115116117def main() -> None:118    ap = argparse.ArgumentParser()119    ap.add_argument("--build", action="store_true")120    ap.add_argument("--per-domain", type=int, default=4)121    ap.add_argument("--max-tokens", type=int, default=128)122    ap.add_argument("--window", type=int, default=32)123    ap.add_argument("--taus", default="1.0,2.0")124    args = ap.parse_args()125    if args.build:126        build()127        return128129    domains = json.loads((REPO_ROOT / "benchmarks/datasets/eval_prompts.json").read_text())["domains"]130    q4_dir, q8_dir = MODELS_DIR / "q4", MODELS_DIR / "q8"131    q8_bytes = dir_weight_bytes(q8_dir)132    q4_bytes = dir_weight_bytes(q4_dir)133    print(f"resident q4: {q4_bytes/1e9:.2f} GB · streamed q8: {q8_bytes/1e9:.2f} GB", flush=True)134135    base_model, tokenizer = load(str(q4_dir))136    verify_model, _ = load(str(q8_dir))137138    prompts = []139    for domain, plist in domains.items():140        for prompt in plist[: args.per_domain]:141            ids = tokenizer.apply_chat_template(142                [{"role": "user", "content": prompt}], add_generation_prompt=True)143            prompts.append({"domain": domain, "ids": list(ids)})144145    # baselines146    print("baseline: pure q8 greedy …", flush=True)147    q8_out, q8_times = [], []148    for p in prompts:149        toks, dt = greedy_baseline(verify_model, tokenizer, p["ids"], args.max_tokens)150        q8_out.append(toks); q8_times.append((len(toks), dt))151    print("baseline: pure q4 greedy …", flush=True)152    q4_out, q4_times = [], []153    for p in prompts:154        toks, dt = greedy_baseline(base_model, tokenizer, p["ids"], args.max_tokens)155        q4_out.append(toks); q4_times.append((len(toks), dt))156157    def toks_per_s(times):158        n = sum(t for t, _ in times); s = sum(d for _, d in times)159        return n / s if s else 0.0160161    configs = []162    for mode in ("margin", "verify-all"):163        for tau in ([float(x) for x in args.taus.split(",")] if mode == "margin" else [2.0]):164            configs.append({"mode": mode, "tau": tau})165166    outputs_by_config = {"pure_q4": q4_out, "pure_q8": q8_out}167    results = []168    for cfg in configs:169        print(f"runtime: mode={cfg['mode']} tau={cfg['tau']} W={args.window} …", flush=True)170        fid, agg = [], {"tokens": 0, "deferred": 0, "sweeps": 0, "rollbacks": 0,171                        "sweep_s": 0.0, "gen_s": 0.0, "logical_bytes": 0}172        cfg_outputs = []173        for p, ref in zip(prompts, q8_out):174            toks, st = generate_deferred(175                base_model, verify_model, tokenizer, p["ids"],176                args.max_tokens, cfg["tau"], args.window, cfg["mode"], q8_bytes)177            cfg_outputs.append(toks)178            fid.append(fidelity(toks, ref))179            agg["tokens"] += st.tokens_out; agg["deferred"] += st.deferred180            agg["sweeps"] += st.sweeps; agg["rollbacks"] += st.rollbacks181            agg["sweep_s"] += st.sweep_time_s; agg["gen_s"] += st.gen_time_s182            agg["logical_bytes"] += st.sweep_logical_bytes183        n = max(agg["tokens"], 1)184        results.append({185            **cfg, "window": args.window,186            "fidelity_vs_q8_mean": sum(fid) / len(fid),187            "tokens_per_s": n / (agg["gen_s"] + agg["sweep_s"]),188            "deferral_rate": agg["deferred"] / n,189            "rollback_rate": agg["rollbacks"] / n,190            "sweeps_per_100tok": 100 * agg["sweeps"] / n,191            "sweep_latency_s_mean": agg["sweep_s"] / max(agg["sweeps"], 1),192            "logical_verify_bytes_per_token": agg["logical_bytes"] / n,193            "raw": agg,194        })195        outputs_by_config[f"{cfg['mode']}_tau{cfg['tau']}"] = cfg_outputs196        r = results[-1]197        print(f"  fidelity={r['fidelity_vs_q8_mean']:.4f} tok/s={r['tokens_per_s']:.1f} "198              f"defer={r['deferral_rate']:.2f} rollback={r['rollback_rate']:.3f} "199              f"MB/token(logical)={r['logical_verify_bytes_per_token']/1e6:.0f}", flush=True)200201    print("judging outputs with bf16 reference …", flush=True)202    quality = judge_outputs(outputs_by_config, prompts)203    for name, s in quality.items():204        print(f"  {name:>18}: mean logprob (bf16 judge) = {s['mean_logprob_bf16']:.4f}", flush=True)205206    payload = {207        "experiment": "candidate_01_deferred_refinement",208        "author": "Simon-Pierre Boucher",209        "contact": "contact@spboucher.ai",210        "manifest": collect_manifest(),211        "config": vars(args),212        "model": HF_MODEL,213        "q4_resident_bytes": q4_bytes, "q8_stream_bytes": q8_bytes,214        "baselines": {215            "pure_q4": {"tokens_per_s": toks_per_s(q4_times),216                        "fidelity_vs_q8_mean": sum(fidelity(a, b) for a, b in zip(q4_out, q8_out)) / len(q8_out)},217            "pure_q8": {"tokens_per_s": toks_per_s(q8_times), "fidelity_vs_q8_mean": 1.0},218        },219        "runs": results,220        "quality_bf16_judge": quality,221    }222    ts = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")223    out_dir = REPO_ROOT / "results" / "candidate_01" / ts224    out_dir.mkdir(parents=True)225    (out_dir / "results.json").write_text(json.dumps(payload, indent=2))226    print(f"\nwrote {out_dir / 'results.json'}")227    print("baselines:", json.dumps(payload["baselines"], indent=1))228229230if __name__ == "__main__":231    main()232