codelion commited on
Commit
bed5ea1
·
verified ·
1 Parent(s): 06285fe

Use perplexity from mlxbench-2 only (BOS in every window)

Browse files
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-1"
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
- n_windows = min(len(ids), max_tokens) // window
 
 
 
 
 
 
 
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 * window:(w + 1) * window])[None]
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