nur-dev's picture
Add files using upload-large-folder tool
e69b72a verified
Raw
History Blame Contribute Delete
2.83 kB
"""Held-out perplexity over a packed token corpus."""
from __future__ import annotations
import math
from dataclasses import dataclass
from pathlib import Path
import torch
from strata.training.lm_data import PackedLMDataset
@dataclass(slots=True)
class PerplexityResult:
corpus: str
perplexity: float
loss: float
tokens: int
windows: int
def to_dict(self) -> dict[str, object]:
return {
"corpus": self.corpus,
"perplexity": self.perplexity,
"loss": self.loss,
"tokens": self.tokens,
"windows": self.windows,
}
@torch.no_grad()
def evaluate_perplexity(
model,
corpus_dir: Path,
*,
seq_len: int,
device: torch.device,
batch_size: int = 8,
max_windows: int | None = None,
precision: str = "bf16",
predicate_memory_intervention: str = "none",
predicate_memory_residual_scale: float | None = None,
graph_object_residual_scale: float | None = None,
) -> PerplexityResult:
"""Token-weighted mean NLL / perplexity over sequential corpus windows.
Evaluation is deterministic (no shuffle). Windows are scored in order; pass
``max_windows`` to cap the number evaluated for a quick estimate.
"""
dataset = PackedLMDataset(corpus_dir, seq_len=seq_len)
n_windows = len(dataset) if max_windows is None else min(len(dataset), max_windows)
if n_windows == 0:
raise ValueError(f"corpus {corpus_dir} has no windows at seq_len {seq_len}")
device_type = device.type
autocast = (
torch.autocast(device_type=device_type, enabled=False)
if (precision == "fp32" or device_type == "cpu")
else torch.autocast(device_type=device_type, dtype=torch.bfloat16 if precision == "bf16" else torch.float16)
)
model.eval()
loss_sum = 0.0
token_count = 0
for start in range(0, n_windows, batch_size):
indices = range(start, min(start + batch_size, n_windows))
batch = torch.stack([dataset[i] for i in indices]).to(device)
with autocast:
output = model(
batch,
labels=batch,
predicate_memory_intervention=predicate_memory_intervention,
predicate_memory_residual_scale=predicate_memory_residual_scale,
graph_object_residual_scale=graph_object_residual_scale,
)
# output.loss is the mean CE over B*(S-1) predicted tokens.
predicted = batch.shape[0] * (batch.shape[1] - 1)
loss_sum += float(output.loss) * predicted
token_count += predicted
mean_loss = loss_sum / token_count
return PerplexityResult(
corpus=corpus_dir.name,
perplexity=math.exp(mean_loss),
loss=mean_loss,
tokens=token_count,
windows=n_windows,
)