#!/usr/bin/env python3 """Drive MMSpec (github.com/killthefullmoon/MMSpec) against mlx-vlm in-process. Port of mmspec_llamacpp.py to the MLX stack: same data, same prompting protocol (MMSpec's verbatim system prompt, greedy, image attached to the first user turn only, multi-turn fed sequentially with assistant replies accumulated), same output jsonl schema, same tau conventions: tau = predicted_n / (predicted_n - draft_n_accepted) No prompt cache: every generate() call pays its full prompt, so prompt_ms is a physically meaningful TTFT (matches the llama.cpp --no-prompt-cache runs). Known protocol delta vs the llama.cpp driver: mlx-vlm's chat formatting places the token before the first turn's text (llama.cpp sent text first). Usage: python mmspec_mlx.py --model [--drafter ] \ --data ~/Downloads/MMSpec/dataset/MMSpec/testmini --out results.jsonl """ import argparse import json import os import sys import time SYSTEM_PROMPT = ("A chat between a curious human and an artificial intelligence assistant. " "The assistant gives helpful, detailed, and polite answers to the human's questions.") TOPICS = { "chart understanding": "CharXiv", "complex reasoning pro": "MMMU-Pro", "general vqa": "GQA", "image captioning": "COCO", "multi-turn conversation": "multi-turn", "text vqa": "TextVQA", } def subset_of(sample): return TOPICS.get(sample.get("topic"), sample.get("topic", "unknown")) def run_sample(model, processor, config, drafter, draft_kind, data_dir, sample, max_tokens, draft_block_size=None): from mlx_vlm.generate import generate from mlx_vlm.prompt_utils import apply_chat_template img_names = sample.get("images") or [sample["image"]] img_paths = [os.path.join(data_dir, "images", n) for n in img_names] messages = [{"role": "system", "content": SYSTEM_PROMPT}] turns_out = [] for i, turn in enumerate(sample["turns"]): if i == 0: content = [{"type": "text", "text": turn}] + [ {"type": "image"} for _ in img_paths ] else: content = turn messages.append({"role": "user", "content": content}) formatted = apply_chat_template( processor, config, messages, num_images=len(img_paths) ) gen_kwargs = {} if drafter is not None: gen_kwargs = {"draft_model": drafter, "draft_kind": draft_kind} if draft_block_size: gen_kwargs["draft_block_size"] = draft_block_size result = generate( model, processor, formatted, image=img_paths, max_tokens=max_tokens, temperature=0.0, **gen_kwargs, ) text = result.text messages.append({"role": "assistant", "content": text}) accepts = list(getattr(drafter, "accept_lens", []) or []) if drafter else [] drafts = list(getattr(drafter, "draft_lens", []) or []) if drafter else [] pred_n = result.generation_tokens turns_out.append({ "text": text, "predicted_n": pred_n, "predicted_ms": pred_n / result.generation_tps * 1000 if result.generation_tps else 0.0, "prompt_n": result.prompt_tokens, "prompt_ms": result.prompt_tokens / result.prompt_tps * 1000 if result.prompt_tps else 0.0, "draft_n": sum(drafts), "draft_n_accepted": sum(accepts), }) return turns_out def main(): ap = argparse.ArgumentParser() ap.add_argument("--model", required=True) ap.add_argument("--drafter", default="", help="HF drafter dir; empty = baseline arm") ap.add_argument("--draft-block-size", type=int, default=0, help="override drafter block width (0 = mlx-vlm default)") ap.add_argument("--data", required=True, help="split dir containing mmspec.jsonl + images/") ap.add_argument("--limit", type=int, default=0, help="first N samples only (0 = all)") ap.add_argument("--subsets", default="", help="comma-separated subset filter (e.g. GQA,COCO)") ap.add_argument("--max-tokens", type=int, default=2048) ap.add_argument("--out", default="") args = ap.parse_args() from mlx_vlm import load from mlx_vlm.speculative.drafters import load_drafter from mlx_vlm.utils import load_config model, processor = load(os.path.expanduser(args.model)) config = load_config(os.path.expanduser(args.model)) drafter, draft_kind = (None, None) if args.drafter: drafter, draft_kind = load_drafter(os.path.expanduser(args.drafter)) data_dir = os.path.expanduser(args.data) samples = [json.loads(l) for l in open(os.path.join(data_dir, "mmspec.jsonl"))] if args.subsets: want = set(args.subsets.split(",")) samples = [s for s in samples if subset_of(s) in want] if args.limit: samples = samples[:args.limit] # warmup: Metal pipeline compile + first-touch allocations, not measured warm = {"id": "warmup", "topic": samples[0].get("topic"), "image": (samples[0].get("images") or [samples[0]["image"]])[0] if not samples[0].get("images") else None, "images": samples[0].get("images"), "turns": [samples[0]["turns"][0]]} run_sample(model, processor, config, drafter, draft_kind, data_dir, warm, max_tokens=32, draft_block_size=args.draft_block_size or None) print("warmup done", flush=True) out_f = open(args.out, "w") if args.out else None agg = {} # subset -> [pred_n, cycles, pred_ms] taus = {} # subset -> per-turn tau list (request-mean convention) t_start = time.time() for n_done, s in enumerate(samples, 1): try: turns = run_sample(model, processor, config, drafter, draft_kind, data_dir, s, args.max_tokens, draft_block_size=args.draft_block_size or None) except Exception as e: print(f"{s['id']}: ERROR {e}", file=sys.stderr, flush=True) continue sub = subset_of(s) a = agg.setdefault(sub, [0, 0, 0.0]) for t in turns: cyc = t["predicted_n"] - t["draft_n_accepted"] a[0] += t["predicted_n"] a[1] += cyc a[2] += t["predicted_ms"] if cyc > 0: taus.setdefault(sub, []).append(t["predicted_n"] / cyc) if out_f: out_f.write(json.dumps({"id": s["id"], "subset": sub, "turns": turns}) + "\n") out_f.flush() tau = a[0] / a[1] if a[1] else 0.0 print(f"[{n_done}/{len(samples)}] {s['id']} {sub}: subset tau={tau:.2f} " f"({time.time()-t_start:.0f}s elapsed)", flush=True) print("\n=== per subset ===") tot = [0, 0, 0.0] all_taus = [] for sub in sorted(agg): p, c, ms = agg[sub] for i, v in enumerate((p, c, ms)): tot[i] += v ts = taus.get(sub, []) all_taus += ts rmean = sum(ts) / len(ts) if ts else 0.0 print(f"{sub:12s} tau={p/c if c else 0:.2f} tau_rmean={rmean:.2f} tps={p/ms*1000 if ms else 0:.1f} n={p}") p, c, ms = tot if c: rmean = sum(all_taus) / len(all_taus) if all_taus else 0.0 print(f"{'overall':12s} tau={p/c:.2f} tau_rmean={rmean:.2f} tps={p/ms*1000:.1f} n={p}") if __name__ == "__main__": main()