Spaces:
Running
Running
Use perplexity from mlxbench-2 only (BOS in every window)
Browse files- app/recommend.py +3 -0
- app/stats.py +3 -2
- bench/mlx_explorer_bench.py +10 -3
- tests/test_events.py +11 -0
app/recommend.py
CHANGED
|
@@ -70,6 +70,9 @@ class CommunitySignal:
|
|
| 70 |
|
| 71 |
|
| 72 |
EVAL_DATASET = "wikitext2-test-v1"
|
|
|
|
|
|
|
|
|
|
| 73 |
EVAL_LABEL = "wikitext-2"
|
| 74 |
# Perplexity is only comparable between quantizations of the same base model (same tokenizer).
|
| 75 |
# Across bit widths the priority-based preference already applies, so measurements only move a
|
|
|
|
| 70 |
|
| 71 |
|
| 72 |
EVAL_DATASET = "wikitext2-test-v1"
|
| 73 |
+
# Benchmark versions whose perplexity is comparable. mlxbench-1 left BOS off most windows,
|
| 74 |
+
# which inflated perplexity for models that have a BOS token (Gemma by 100x).
|
| 75 |
+
PPL_BENCHMARK_VERSIONS = ("mlxbench-2",)
|
| 76 |
EVAL_LABEL = "wikitext-2"
|
| 77 |
# Perplexity is only comparable between quantizations of the same base model (same tokenizer).
|
| 78 |
# Across bit widths the priority-based preference already applies, so measurements only move a
|
app/stats.py
CHANGED
|
@@ -14,7 +14,7 @@ from collections import Counter, defaultdict
|
|
| 14 |
from typing import Callable, Iterable
|
| 15 |
|
| 16 |
from .events import latest_sessions
|
| 17 |
-
from .recommend import EVAL_DATASET, CommunitySignal
|
| 18 |
|
| 19 |
K_MIN = 5
|
| 20 |
BENCH_TYPES = ("mlx_benchmark", "mlx_benchmark_submission")
|
|
@@ -91,7 +91,8 @@ def community_signals(rows: Iterable[dict]) -> dict[str, CommunitySignal]:
|
|
| 91 |
elif r.get("event_type") in BENCH_TYPES:
|
| 92 |
if r.get("generation_tps"):
|
| 93 |
bench[m].append(float(r["generation_tps"]))
|
| 94 |
-
if r.get("perplexity") and r.get("eval_dataset") == EVAL_DATASET
|
|
|
|
| 95 |
ppl[m].append(float(r["perplexity"]))
|
| 96 |
out = {}
|
| 97 |
for m in set(fb) | set(bench) | set(ppl):
|
|
|
|
| 14 |
from typing import Callable, Iterable
|
| 15 |
|
| 16 |
from .events import latest_sessions
|
| 17 |
+
from .recommend import EVAL_DATASET, PPL_BENCHMARK_VERSIONS, CommunitySignal
|
| 18 |
|
| 19 |
K_MIN = 5
|
| 20 |
BENCH_TYPES = ("mlx_benchmark", "mlx_benchmark_submission")
|
|
|
|
| 91 |
elif r.get("event_type") in BENCH_TYPES:
|
| 92 |
if r.get("generation_tps"):
|
| 93 |
bench[m].append(float(r["generation_tps"]))
|
| 94 |
+
if (r.get("perplexity") and r.get("eval_dataset") == EVAL_DATASET
|
| 95 |
+
and r.get("benchmark_version") in PPL_BENCHMARK_VERSIONS):
|
| 96 |
ppl[m].append(float(r["perplexity"]))
|
| 97 |
out = {}
|
| 98 |
for m in set(fb) | set(bench) | set(ppl):
|
bench/mlx_explorer_bench.py
CHANGED
|
@@ -29,7 +29,7 @@ import urllib.error
|
|
| 29 |
import urllib.request
|
| 30 |
from pathlib import Path
|
| 31 |
|
| 32 |
-
BENCHMARK_VERSION = "mlxbench-
|
| 33 |
|
| 34 |
# Quality: perplexity on a fixed slice of the wikitext-2 test set (CC BY-SA 3.0).
|
| 35 |
# Perplexity depends on the tokenizer, so results are only compared between
|
|
@@ -92,12 +92,19 @@ def perplexity(model, tokenizer, max_tokens: int, window: int) -> tuple[float, f
|
|
| 92 |
import mlx.nn as nn
|
| 93 |
|
| 94 |
ids = tokenizer.encode(load_eval_text())
|
| 95 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
if n_windows < 1:
|
| 97 |
raise SystemExit("evaluation text is shorter than one window")
|
| 98 |
losses = []
|
| 99 |
for w in range(n_windows):
|
| 100 |
-
chunk = mx.array(ids[w *
|
| 101 |
logits = model(chunk[:, :-1]).astype(mx.float32)
|
| 102 |
loss = nn.losses.cross_entropy(logits, chunk[:, 1:], reduction="none")
|
| 103 |
mx.eval(loss)
|
|
|
|
| 29 |
import urllib.request
|
| 30 |
from pathlib import Path
|
| 31 |
|
| 32 |
+
BENCHMARK_VERSION = "mlxbench-2" # 2: every perplexity window starts with BOS
|
| 33 |
|
| 34 |
# Quality: perplexity on a fixed slice of the wikitext-2 test set (CC BY-SA 3.0).
|
| 35 |
# Perplexity depends on the tokenizer, so results are only compared between
|
|
|
|
| 92 |
import mlx.nn as nn
|
| 93 |
|
| 94 |
ids = tokenizer.encode(load_eval_text())
|
| 95 |
+
# Models with a BOS token (Gemma, LFM, MiniCPM, ...) expect it at the start of
|
| 96 |
+
# every sequence; without it Gemma's perplexity runs into the thousands.
|
| 97 |
+
# Some tokenizers add it on encode, some don't, so strip it and add it per window.
|
| 98 |
+
bos = getattr(tokenizer, "bos_token_id", None)
|
| 99 |
+
if bos is not None and ids and ids[0] == bos:
|
| 100 |
+
ids = ids[1:]
|
| 101 |
+
step = window - 1 if bos is not None else window
|
| 102 |
+
n_windows = min(len(ids), max_tokens) // step
|
| 103 |
if n_windows < 1:
|
| 104 |
raise SystemExit("evaluation text is shorter than one window")
|
| 105 |
losses = []
|
| 106 |
for w in range(n_windows):
|
| 107 |
+
chunk = mx.array(([bos] if bos is not None else []) + ids[w * step:(w + 1) * step])[None]
|
| 108 |
logits = model(chunk[:, :-1]).astype(mx.float32)
|
| 109 |
loss = nn.losses.cross_entropy(logits, chunk[:, 1:], reduction="none")
|
| 110 |
mx.eval(loss)
|
tests/test_events.py
CHANGED
|
@@ -155,3 +155,14 @@ def test_quality_submission_validation():
|
|
| 155 |
ev(event_type="mlx_benchmark", perplexity=0.5)
|
| 156 |
with pytest.raises(ValidationError):
|
| 157 |
ev(event_type="mlx_benchmark", perplexity=9.8, eval_dataset="my-own-text")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 155 |
ev(event_type="mlx_benchmark", perplexity=0.5)
|
| 156 |
with pytest.raises(ValidationError):
|
| 157 |
ev(event_type="mlx_benchmark", perplexity=9.8, eval_dataset="my-own-text")
|
| 158 |
+
|
| 159 |
+
|
| 160 |
+
def test_perplexity_from_old_bench_version_is_ignored():
|
| 161 |
+
from app.stats import community_signals
|
| 162 |
+
|
| 163 |
+
row = dict(event_type="mlx_benchmark", selected_model="mlx-community/gemma-4-e2b-it-4bit", perplexity=5268.9,
|
| 164 |
+
eval_dataset="wikitext2-test-v1", suspicious_flags=[])
|
| 165 |
+
sig = community_signals([{**row, "benchmark_version": "mlxbench-1"},
|
| 166 |
+
{**row, "benchmark_version": "mlxbench-2", "perplexity": 21.4}])
|
| 167 |
+
assert sig[row["selected_model"]].median_perplexity == 21.4
|
| 168 |
+
assert sig[row["selected_model"]].perplexity_count == 1
|