Perplexity
Be able to compute perplexity and explain what it measures — and does not measure.
Prerequisites
Intuition
Perplexity is the cross-entropy, packaged so that it can be interpreted:
The interpretation: how many alternatives the model is effectively choosing between at each token.
| Perplexity | Means |
|---|---|
| 1 | the model always knows exactly what the next token is |
| 10 | it is as uncertain as at a choice between ten equivalent options |
| 50 000 | it is guessing entirely at random from the vocabulary |
An improvement from a loss of 2.3 to 2.0 sounds marginal. In perplexity it is 10.0 → 7.4 — a quarter off the effective uncertainty. That is why perplexity is reported: it makes small changes in the loss intelligible.
Formal
Perplexity can only be compared under strict conditions. This is often missed, and makes published comparisons misleading.
| Has to be identical | Why |
|---|---|
| The tokenizer | more tokens per word gives a lower perplexity per token — the same model, a different number |
| The test set | different text has a different difficulty |
| The context length | more context makes every token easier to guess |
| The unit of normalisation | per token, per word or per character gives entirely different numbers |
The first is the sneakiest: a model with a tokenizer that splits the text into more, shorter tokens automatically gets a lower perplexity per token, without being better. Perplexity per character (bits per character) is therefore the measure used when the tokenizers differ.
What perplexity does not measure:
| Property | Captured by perplexity? |
|---|---|
| Next-token uncertainty | yes — that is precisely what it measures |
| Factual correctness | no |
| Instruction following | no |
| Reasoning ability | no |
| Helpfulness and tone | no |
| Safety | no |
An instruction-tuned model often has a higher perplexity on raw text than the base model it comes from — it has learnt to answer instead of continuing the text, and that is an improvement the metric punishes.
Where perplexity is still useful:
- During pretraining — as a continuous measure that things are moving forwards.
- When comparing the same model before and after a change, with everything else equal.
- For detecting a domain shift — the perplexity rises when the model meets text unlike the training data.
- For quality-filtering training data — an extremely high or extremely low perplexity indicates rubbish and repetition respectively.
The fourth is a real and underrated use: running a small model over a corpus and filtering out the documents in the tails is a cheap and effective clean-up.
Contamination is what makes published perplexity numbers unreliable: if the test set happens to be in the training data the perplexity becomes artificially low. Always check for overlap between test and training before you believe a figure.
Code
import math, torch
def perplexity(model, tokens, context=1024, stride=512):
"""A sliding window with overlap; only the new tokens count."""
model.eval()
nll, count = 0.0, 0
for start in range(0, len(tokens) - 1, stride):
end = min(start + context, len(tokens))
chunk = tokens[start:end]
target = chunk.clone()
target[:-stride] = -100 # count only the new tokens
with torch.no_grad():
out = model(chunk.unsqueeze(0), labels=target.unsqueeze(0))
n = int((target != -100).sum())
nll += float(out.loss) * n
count += n
if end == len(tokens):
break
return math.exp(nll / count)
# The loss → the perplexity
for loss in (1.6, 2.0, 2.3, 3.0, 4.6):
print(f" a loss of {loss:.1f} → a perplexity of {math.exp(loss):>8.1f}")
# a loss of 1.6 → a perplexity of 5.0
# a loss of 2.3 → a perplexity of 10.0
# a loss of 4.6 → a perplexity of 99.5
# Bits per character — comparable ACROSS tokenizers
def bpc(model, text, tokenizer):
tokens = tokenizer(text, return_tensors="pt").input_ids[0]
with torch.no_grad():
loss = float(model(tokens.unsqueeze(0), labels=tokens.unsqueeze(0)).loss)
return loss * len(tokens) / (len(text) * math.log(2))
# A domain shift shows directly in the perplexity
def domain_check(model, tokenizer, texts: dict):
for name, t in texts.items():
ids = tokenizer(t, return_tensors="pt").input_ids
with torch.no_grad():
p = math.exp(float(model(ids, labels=ids).loss))
print(f" {name:<22} PPL {p:>8.1f}")
# news text PPL 18.3
# medical records PPL 142.7 ← far outside the training domain
# Quality filtering: throw the tails away
def filter_corpus(model, tokenizer, documents, low=10, high=1000):
keep = []
for d in documents:
ids = tokenizer(d, return_tensors="pt", truncation=True, max_length=512).input_ids
with torch.no_grad():
p = math.exp(float(model(ids, labels=ids).loss))
if low <= p <= high:
keep.append(d)
return keep
# too LOW a perplexity = repetition, boilerplate; too HIGH = rubbish, code sludge
Mastery means
- Computes perplexity from the cross-entropy
- Interprets the value
- Knows when perplexity cannot be compared
Sign in to do the exercises and build your mastery up.
Sources
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0
- Hugging Face — dokumentation (Apache-2.0) — Apache-2.0
- arXiv — The Curious Case of Neural Text Degeneration — arXiv (open access; licence per article)