|
#!/usr/bin/env python3 |
|
""" |
|
Benchmark: Ollama (GGUF Q4_K_M & NVFP4) vs mlx-lm (4-bit) |
|
Models: |
|
gemma4:e4b |
|
gemma4:e4b-nvfp4 |
|
mlx-community/gemma-4-E4B-it-4bit |
|
|
|
Usage: |
|
pip install ollama mlx-lm |
|
ollama pull gemma4:e4b |
|
python bench.py |
|
|
|
Metrics: |
|
* TTFT (Time To First Token): How long the model takes to process your prompt and start writing. This is the "thinking before speaking" phase — all input tokens are processed in parallel before any output appears. Shorter = more responsive feel. |
|
* Decode (Generation time): How long it takes to write the full response, one token, at a time, after the first token appears. This is the dominant cost for long outputs like stories. |
|
* WALL (Total elapsed time): What you actually wait for: TTFT + decode + any overhead from the framework (HTTP, streaming, scheduling). Wall - TTFT - decode = framework overhead. |
|
* TOK/S (Tokens per second): Decode speed. A token is roughly ¾ of a word on average, so 70 tok/s ≈ 50 words/second. Human reading speed is ~200-250 words/min (~3-4 words/sec), so anything above ~5 tok/s feels instant in practice. |
|
|
|
""" |
|
|
|
import json |
|
import statistics |
|
import sys |
|
import threading |
|
import time |
|
|
|
import ollama |
|
from mlx_lm import generate, load, stream_generate |
|
from mlx_lm.sample_utils import make_sampler |
|
|
|
# ── Config ─────────────────────────────────────────────────────────────────── |
|
|
|
OLLAMA_MODELS = {"GGUF": "gemma4:e4b", "MLX": "gemma4:e4b-nvfp4"} |
|
MLX_MODEL = "mlx-community/gemma-4-E4B-it-4bit" |
|
N_RUNS = 5 |
|
CONTEXT_SIZE = 8192 |
|
|
|
PROMPTS = { |
|
"short_factual": { |
|
"text": "What is the capital of Japan? Answer in one sentence.", |
|
"max_tokens": 50, |
|
}, |
|
"creative_mid": { |
|
"text": ( |
|
"Write the opening paragraph of a cyberpunk short story set in 2087 Tokyo. " |
|
"A rain-soaked courier discovers a data chip that shouldn't exist. " |
|
"Atmospheric, noir tone, ~150 words." |
|
), |
|
"max_tokens": 400, |
|
}, |
|
"creative_long": { |
|
"max_tokens": 1500, |
|
"text": """ |
|
# System prompt |
|
|
|
* You are a sleep podcasts writer! Your style is "Cyberpunk ASMR" — heavily atmospheric & sensory. |
|
* Use markdown formatting. Start with the title right ahead. |
|
* Do not write any instructions, all the text is read out by the narrator. |
|
* Write a long, immersive cyberpunk sleep story of approximately 500 words. |
|
* Break the story into many atmospheric paragraphs. Maintain a steady, rhythmic pace. |
|
* Crucial: Ensure the story has a clear, atmospheric conclusion and ends with 'The end.' |
|
* Do not wander; keep the pacing tight so you finish within the word count. |
|
|
|
# Task |
|
|
|
Write a "sleep story" set in a rain-slicked, neon-lit cyberpunk city. |
|
|
|
# The Vibe |
|
|
|
Inspired by Neuromancer, The Snowcrash and Ready Player One. Focus on the "low-life, high-tech" aesthetic but through a lens of calm, late-night solitude. |
|
|
|
# Story constraints |
|
|
|
1) No Conflict: The "story arc" should be a flat plateau of calm activity. |
|
2) Sensory Details: Focus on the hum of cooling fans, the rhythmic drip & smell of rain, the soft glow of emerald terminal text. |
|
3) Pacing: Use long, flowing sentences, use hypnotic descriptions. Use commas and ellipses to create a slow, rhythmic pace for the narrator. |
|
4) The "Decking" Segment: Describe a slow, peaceful transition into a "private server" virtual reality that looks like a calm, digital Zen garden or a low-poly ocean. |
|
|
|
# Structure |
|
|
|
This is just for the overall arch for the actual theme, take some creative liberty. |
|
|
|
* 0-150 words: Setting the scene in a small, cozy apartment or hideout at night. The sound of the city outside. |
|
* 150-350 words: The process of "booting up" / "jacking in". The tactile feel of the cyberdeck, the soft click of switches & cables, the slow crawl of data on the screen. |
|
* 350-500 words: A drift into a virtual void. Ending with pulsing neon lines that slowly fade to black. |
|
""", |
|
}, |
|
} |
|
|
|
# ── Spinner ─────────────────────────────────────────────────────────────────── |
|
|
|
|
|
class Spinner: |
|
"""Simple CLI spinner for blocking operations (model load, warmup).""" |
|
|
|
FRAMES = "⠋⠙⠹⠸⠼⠴⠦⠧⠇⠏" |
|
|
|
def __init__(self, label: str): |
|
self.label = label |
|
self._stop = threading.Event() |
|
self._thread = threading.Thread(target=self._spin, daemon=True) |
|
|
|
def _spin(self): |
|
i = 0 |
|
while not self._stop.is_set(): |
|
frame = self.FRAMES[i % len(self.FRAMES)] |
|
sys.stdout.write(f"\r {frame} {self.label} ") |
|
sys.stdout.flush() |
|
time.sleep(0.08) |
|
i += 1 |
|
|
|
def __enter__(self): |
|
self._thread.start() |
|
return self |
|
|
|
def __exit__(self, *_): |
|
self._stop.set() |
|
self._thread.join() |
|
sys.stdout.write("\r" + " " * (len(self.label) + 8) + "\r") |
|
sys.stdout.flush() |
|
|
|
|
|
# ── Ollama benchmark ────────────────────────────────────────────────────────── |
|
def bench_ollama(model: str, prompt: str, max_tokens: int, run_index: int) -> dict: |
|
first_token_time = None |
|
token_count = 0 |
|
collected = [] |
|
|
|
print(" ", end="", flush=True) |
|
|
|
t_start = time.perf_counter() |
|
|
|
stream = ollama.chat( |
|
model=model, |
|
messages=[{"role": "user", "content": prompt}], |
|
options={ |
|
"num_predict": max_tokens, |
|
"num_ctx": CONTEXT_SIZE, |
|
"temperature": 0, |
|
"keep_alive": "30m", |
|
}, |
|
stream=True, |
|
) |
|
|
|
final_chunk = None |
|
for chunk in stream: |
|
if first_token_time is None: |
|
first_token_time = time.perf_counter() |
|
# Print TTFT immediately so you can see prefill completing |
|
ttft_so_far = first_token_time - t_start |
|
print(f"[TTFT {ttft_so_far:.2f}s] ", end="", flush=True) |
|
|
|
token_text = chunk.message.content or "" |
|
collected.append(token_text) |
|
token_count += 1 |
|
|
|
# Print a dot every 10 tokens to show decode progress |
|
if token_count % 10 == 0: |
|
print("·", end="", flush=True) |
|
|
|
final_chunk = chunk |
|
|
|
t_end = time.perf_counter() |
|
print() # newline after dots |
|
|
|
# Pull timing from the final streaming chunk (same fields as non-streaming) |
|
prompt_tokens = getattr(final_chunk, "prompt_eval_count", 0) or 0 |
|
output_tokens = getattr(final_chunk, "eval_count", 0) or token_count |
|
ttft_s = (getattr(final_chunk, "prompt_eval_duration", 0) or 0) / 1e9 |
|
decode_s = (getattr(final_chunk, "eval_duration", 0) or 0) / 1e9 |
|
total_s = t_end - t_start |
|
tok_per_s = output_tokens / decode_s if decode_s > 0 else 0 |
|
|
|
# Fallback: if the final chunk didn't carry duration fields, use wall time |
|
if ttft_s == 0 and first_token_time: |
|
ttft_s = first_token_time - t_start |
|
decode_s = t_end - first_token_time |
|
|
|
return { |
|
"run": run_index, |
|
"prompt_tokens": prompt_tokens, |
|
"output_tokens": output_tokens, |
|
"ttft_s": round(ttft_s, 3), |
|
"decode_s": round(decode_s, 3), |
|
"total_wall_s": round(total_s, 3), |
|
"tok_per_s": round(tok_per_s, 1), |
|
"text": "".join(collected), |
|
} |
|
|
|
|
|
# ── mlx-lm benchmark ───────────────────────────────────────────────────────── |
|
|
|
# Module-level cache so the model is only loaded once across all prompts |
|
_mlx_model_cache = {} |
|
|
|
|
|
def get_mlx_model(): |
|
if "model" not in _mlx_model_cache: |
|
with Spinner(f"Loading {MLX_MODEL}"): |
|
model, tokenizer = load(MLX_MODEL) |
|
_mlx_model_cache["model"] = model |
|
_mlx_model_cache["tokenizer"] = tokenizer |
|
return _mlx_model_cache["model"], _mlx_model_cache["tokenizer"] |
|
|
|
|
|
def warmup_mlx(): |
|
with Spinner(" [Warming up MLX JIT...]"): |
|
model, tokenizer = get_mlx_model() |
|
# Create a sampler with temp=0.0 for greedy decoding |
|
sampler = make_sampler(temp=0.0) |
|
generate(model, tokenizer, prompt="Hi", max_tokens=1, sampler=sampler) |
|
|
|
|
|
def bench_mlx(prompt: str, max_tokens: int, run_index: int) -> dict: |
|
model, tokenizer = get_mlx_model() |
|
|
|
first_token_time = None |
|
token_count = 0 |
|
collected = [] |
|
|
|
print(" ", end="", flush=True) |
|
|
|
messages = [{"role": "user", "content": prompt}] |
|
formatted = tokenizer.apply_chat_template( |
|
messages, tokenize=False, add_generation_prompt=True |
|
) |
|
|
|
t_start = time.perf_counter() |
|
sampler = make_sampler(temp=0.0) |
|
|
|
for response in stream_generate( |
|
model, |
|
tokenizer, |
|
prompt=formatted, |
|
max_kv_size=CONTEXT_SIZE, |
|
max_tokens=max_tokens, |
|
sampler=sampler, |
|
): |
|
if first_token_time is None: |
|
first_token_time = time.perf_counter() |
|
ttft_so_far = first_token_time - t_start |
|
print(f"[TTFT {ttft_so_far:.2f}s] ", end="", flush=True) |
|
|
|
token_count += 1 |
|
# response is a GenerationResponse-like object with a .text attribute |
|
collected.append(response.text if hasattr(response, "text") else str(response)) |
|
|
|
if token_count % 10 == 0: |
|
print("·", end="", flush=True) |
|
|
|
t_end = time.perf_counter() |
|
print() |
|
|
|
ttft_s = round(first_token_time - t_start, 3) if first_token_time else None |
|
decode_s = round(t_end - (first_token_time or t_start), 3) |
|
total_s = round(t_end - t_start, 3) |
|
tok_per_s = round(token_count / decode_s, 1) if decode_s > 0 else 0 |
|
|
|
return { |
|
"run": run_index, |
|
"output_tokens": token_count, |
|
"ttft_s": ttft_s, |
|
"decode_s": decode_s, |
|
"total_wall_s": total_s, |
|
"tok_per_s": tok_per_s, |
|
"text": "".join(collected), |
|
} |
|
|
|
|
|
# ── Reporting ───────────────────────────────────────────────────────────────── |
|
def summarize(runs): |
|
def med(key): |
|
vals = sorted(r[key] for r in runs if r.get(key) is not None) |
|
return statistics.median(vals) |
|
|
|
def stdev(key): |
|
vals = [r[key] for r in runs if r.get(key) is not None] |
|
return statistics.stdev(vals) if len(vals) > 1 else 0.0 |
|
|
|
return { |
|
"median_ttft_s": med("ttft_s"), |
|
"median_decode_s": med("decode_s"), |
|
"median_total_wall_s": med("total_wall_s"), |
|
"median_tok_per_s": med("tok_per_s"), |
|
"stdev_tok_per_s": stdev("tok_per_s"), # tells you how stable the runs are |
|
} |
|
|
|
|
|
def print_run_result(r: dict): |
|
print( |
|
f" run {r['run']}: " |
|
f"TTFT={r['ttft_s']}s " |
|
f"decode={r['decode_s']}s " |
|
f"wall={r['total_wall_s']}s " |
|
f"{r['tok_per_s']} tok/s " |
|
f"({r.get('output_tokens', '?')} tokens)" |
|
) |
|
|
|
|
|
def print_summary(label: str, prompt_name: str, runs: list[dict]): |
|
s = summarize(runs) |
|
print(f"\n ── {label} / {prompt_name} averages over {len(runs)} runs:") |
|
for k, v in s.items(): |
|
if v is None: |
|
formatted = " n/a" |
|
else: |
|
formatted = f"{v:>10.3f}" |
|
print(f" {k:<25} {formatted}") |
|
|
|
|
|
def prompt_header(prompt_name: str) -> tuple[str, str, int]: |
|
p = PROMPTS[prompt_name] |
|
first_line = p["text"].splitlines()[0] |
|
short = first_line if len(first_line) < 80 else f"{first_line[:80]}..." |
|
return p["text"], short, p["max_tokens"] |
|
|
|
|
|
# ── Main ────────────────────────────────────────────────────────────────────── |
|
|
|
|
|
def main(): |
|
all_results = {} |
|
|
|
# ── Ollama runs — all prompts first ── |
|
for stack, model in OLLAMA_MODELS.items(): |
|
print("\n" + "═" * 62) |
|
print(f" STACK: OLLAMA {stack}") |
|
print("═" * 62) |
|
for prompt_name in PROMPTS: |
|
prompt_text, prompt_short, max_tokens = prompt_header(prompt_name) |
|
print(f"\n PROMPT: {prompt_name}") |
|
print(f" {prompt_short}") |
|
|
|
ollama_runs = [] |
|
for i in range(1, N_RUNS + 1): |
|
print(f" run {i}/{N_RUNS}:") |
|
result = bench_ollama(model, prompt_text, max_tokens, i) |
|
print_run_result(result) |
|
ollama_runs.append(result) |
|
print_summary(f"OLLAMA {stack}", prompt_name, ollama_runs) |
|
all_results[f"OLLAMA_{stack}_{prompt_name}"] = ollama_runs |
|
|
|
# ── mlx-lm runs — all prompts second ── |
|
warmup_mlx() |
|
|
|
print("\n" + "═" * 62) |
|
print(" STACK: mlx-lm") |
|
print("═" * 62) |
|
for prompt_name in PROMPTS: |
|
prompt_text, prompt_short, max_tokens = prompt_header(prompt_name) |
|
print(f"\n PROMPT: {prompt_name}") |
|
print(f" {prompt_short}") |
|
|
|
mlx_runs = [] |
|
for i in range(1, N_RUNS + 1): |
|
print(f" run {i}/{N_RUNS}:") |
|
result = bench_mlx(prompt_text, max_tokens, i) |
|
print_run_result(result) |
|
mlx_runs.append(result) |
|
print_summary("MLX-LM", prompt_name, mlx_runs) |
|
all_results[f"mlx_{prompt_name}"] = mlx_runs |
|
|
|
# Save raw data |
|
with open("bench_results.json", "w") as f: |
|
json.dump( |
|
{ |
|
k: [{kk: vv for kk, vv in r.items() if kk != "text"} for r in v] |
|
for k, v in all_results.items() |
|
}, |
|
f, |
|
indent=2, |
|
) |
|
print("\n\nRaw results (without generated text) saved to bench_results.json") |
|
|
|
|
|
if __name__ == "__main__": |
|
main() |