Skip to content

Instantly share code, notes, and snippets.

@greg76
Last active April 19, 2026 21:01
Show Gist options
  • Select an option

  • Save greg76/30e58c4b93446064d153471c9c4b771e to your computer and use it in GitHub Desktop.

Select an option

Save greg76/30e58c4b93446064d153471c9c4b771e to your computer and use it in GitHub Desktop.
Benchmark: mlx-lm vs. Ollama. Comparing Time to First Token (TTFT) and decode speeds for on-device AI creative writing.

Built a local LLM benchmark to pick the right inference stack for my "sleepstory" weekend project (fully on-device, AI-narrated bedtime stories).

Should I use mlx-lm (Python, direct MLX inference) or Ollama (new MLX backend)?

Output Length mlx-lm direct Ollama (nvfp4/MLX)
Short (~7 tokens) 70.8 tok/s 58.1 tok/s
Mid (~313 tokens) 69.5 tok/s 57.6 tok/s
Long (~826 tokens) 68.5 tok/s 56.7 tok/s

Hardware: M4 Pro 24GB, Model: Gemma 4 E4B

  • ~21% faster decode across all output lengths. The gap holds whether you're generating 7 tokens or 826 — neither stack degrades meaningfully with longer output at this model size.
  • Time to First Token (TTFT): Ollama is ~7x faster (~26ms vs ~185ms). This isn't a runtime advantage, but an architectural one: Ollama's persistent server eliminates per-call setup costs.
  • Worth noting, using GGUF or NVFP4 on Ollama has almost identical performance for my use case. On an M4 Pro with <500 token prompts, the reported MLX/NVFP4 speedups over standard GGUF/Metal don't trigger. The prefill window is too short.

Match the stack to your bottleneck: If you're building interactive chat, Ollama's ~26ms TTFT feels instantly responsive. If you're running batch generations where total completion time is the only metric, the raw mlx-lm library saves you ~21% in compute.

Keep in mind: I used different weight formats, one model, one chip — your mileage will vary.

#!/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()
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment