financial-rag / src /evaluation /perplexity.py
tolivert's picture
deploy: financial_rag streamlit app
4e316d6
Raw
History Blame Contribute Delete
3.74 kB
"""
Perplexity evaluation on WikiText-2.
Perplexity measures how well a language model predicts a held-out
corpus — lower is better. It is defined as exp(average NLL) over
all tokens. For sequences longer than the model's context window
a sliding-window approach is used: the window advances by a stride
and only the newly exposed tokens contribute to the loss, avoiding
double-counting while still conditioning on full context.
This module downloads WikiText-2 from HuggingFace Datasets and
evaluates any model + tokenizer pair with configurable stride.
"""
import math
import torch
import torch.nn.functional as F
from datasets import load_dataset
def evaluate_perplexity(model, tokenizer, max_length=1024, stride=512, device=None, dataset_text=None):
"""
Compute perplexity on WikiText-2 test set using a sliding window.
The window is `max_length` tokens wide and advances by `stride` tokens.
Only the last `stride` tokens in each window are scored (except the
first window, which scores all tokens). Earlier tokens in the window
provide conditioning context but don't contribute to the loss, so
every token is scored exactly once with maximal left context.
Args:
model: Language model with forward(input_ids) -> logits
tokenizer: Tokenizer with encode(text) -> list[int]
max_length: Context window size for each forward pass
stride: How far the window advances each step (stride <= max_length)
device: Torch device
dataset_text: Override text instead of WikiText-2 (for testing)
Returns:
Perplexity as a float
"""
model.eval()
if device is None:
device = next(model.parameters()).device
# Load and tokenize
if dataset_text is None:
dataset = load_dataset("wikitext", "wikitext-2-raw-v1", split="test")
text = "\n\n".join(dataset["text"])
else:
text = dataset_text
encodings = tokenizer.encode(text)
token_ids = torch.tensor(encodings, dtype=torch.long)
seq_len = token_ids.size(0)
if seq_len < 2:
raise ValueError(f"Need at least 2 tokens to compute perplexity, got {seq_len}")
total_loss = 0.0
total_tokens = 0
prev_end = 0
with torch.no_grad():
for begin in range(0, seq_len - 1, stride):
end = min(begin + max_length, seq_len)
input_ids = token_ids[begin:end].unsqueeze(0).to(device)
logits = model(input_ids) # (1, end-begin, vocab_size)
# logits[:, i, :] predicts token at absolute position begin+i+1.
# Targets: token_ids[begin+1 .. min(end+1, seq_len)-1]
target_end = min(end + 1, seq_len)
num_preds = target_end - begin - 1
targets = token_ids[begin + 1:target_end].unsqueeze(0).to(device)
logits = logits[:, :num_preds, :]
# Only score newly exposed tokens to avoid double-counting.
# Positions [begin+1, prev_end] were already scored by the
# previous window; skip them (they served as context here).
score_from = max(prev_end - begin, 0)
scored_logits = logits[:, score_from:, :]
scored_targets = targets[:, score_from:]
if scored_targets.numel() == 0:
continue
loss = F.cross_entropy(
scored_logits.reshape(-1, scored_logits.size(-1)),
scored_targets.reshape(-1),
reduction="sum",
)
total_loss += loss.item()
total_tokens += scored_targets.numel()
prev_end = end
if end == seq_len:
break
avg_loss = total_loss / total_tokens
return math.exp(avg_loss)