mirror of
https://github.com/elder-plinius/OBLITERATUS.git
synced 2026-08-17 16:37:30 +02:00
test: cover causal evaluation edge cases
This commit is contained in:
@@ -2,7 +2,6 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
|
||||
@@ -56,14 +55,20 @@ class Evaluator:
|
||||
|
||||
ds = self.dataset
|
||||
|
||||
# WikiText and similar corpora contain empty separator rows.
|
||||
# Filter them before max_samples so max_samples means usable texts.
|
||||
raw_texts = ds[self.text_column]
|
||||
valid_indices = [
|
||||
index
|
||||
for index, text in enumerate(raw_texts)
|
||||
if isinstance(text, str) and text.strip()
|
||||
]
|
||||
if self.max_samples is not None and self.max_samples <= 0:
|
||||
raise ValueError("max_samples must be a positive integer or None.")
|
||||
|
||||
# WikiText and similar corpora contain empty separator rows. Walk the
|
||||
# map-style dataset row by row so a bounded evaluation stops as soon as
|
||||
# it has enough usable texts instead of materializing the whole column.
|
||||
valid_indices = []
|
||||
for index in range(len(ds)):
|
||||
text = ds[index][self.text_column]
|
||||
if not isinstance(text, str) or not text.strip():
|
||||
continue
|
||||
valid_indices.append(index)
|
||||
if self.max_samples is not None and len(valid_indices) >= self.max_samples:
|
||||
break
|
||||
|
||||
if not valid_indices:
|
||||
raise ValueError(
|
||||
@@ -73,9 +78,6 @@ class Evaluator:
|
||||
|
||||
ds = ds.select(valid_indices)
|
||||
|
||||
if self.max_samples is not None:
|
||||
ds = ds.select(range(min(self.max_samples, len(ds))))
|
||||
|
||||
total_loss = 0.0
|
||||
total_tokens = 0
|
||||
skipped_batches = 0
|
||||
|
||||
@@ -0,0 +1,228 @@
|
||||
"""Focused regression tests for causal-LM evaluation and report output."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, Mock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from obliteratus.evaluation.evaluator import Evaluator
|
||||
|
||||
|
||||
class _MapDataset:
|
||||
def __init__(self, texts: list[object]):
|
||||
self.texts = texts
|
||||
self.read_indices: list[int] = []
|
||||
self.selected_indices: list[int] | None = None
|
||||
|
||||
def __len__(self):
|
||||
return len(self.texts)
|
||||
|
||||
def __getitem__(self, key):
|
||||
if isinstance(key, str):
|
||||
raise TypeError("the evaluator must not materialize a full dataset column")
|
||||
if isinstance(key, slice):
|
||||
return {"text": self.texts[key]}
|
||||
self.read_indices.append(key)
|
||||
return {"text": self.texts[key]}
|
||||
|
||||
def select(self, indices):
|
||||
self.selected_indices = list(indices)
|
||||
return _MapDataset([self.texts[index] for index in self.selected_indices])
|
||||
|
||||
|
||||
class _Encoding(dict):
|
||||
def to(self, _device):
|
||||
return self
|
||||
|
||||
|
||||
class _ScriptedTokenizer:
|
||||
def __init__(self, encodings: list[_Encoding]):
|
||||
self.encodings = list(encodings)
|
||||
self.calls: list[list[str]] = []
|
||||
|
||||
def __call__(self, texts, **_kwargs):
|
||||
self.calls.append(list(texts))
|
||||
return self.encodings.pop(0)
|
||||
|
||||
|
||||
class _LossModel(nn.Module):
|
||||
def __init__(self, losses: list[float | None]):
|
||||
super().__init__()
|
||||
self.anchor = nn.Parameter(torch.zeros(()))
|
||||
self.losses = list(losses)
|
||||
self.labels: list[torch.Tensor] = []
|
||||
|
||||
def forward(self, *, input_ids, attention_mask, labels):
|
||||
del input_ids, attention_mask
|
||||
self.labels.append(labels.detach().clone())
|
||||
loss = self.losses.pop(0)
|
||||
return SimpleNamespace(loss=None if loss is None else torch.tensor(loss))
|
||||
|
||||
|
||||
def _encoding(input_ids, attention_mask) -> _Encoding:
|
||||
return _Encoding(
|
||||
input_ids=torch.as_tensor(input_ids, dtype=torch.long),
|
||||
attention_mask=torch.as_tensor(attention_mask, dtype=torch.long),
|
||||
)
|
||||
|
||||
|
||||
def _evaluator(
|
||||
texts: list[object],
|
||||
encodings: list[_Encoding],
|
||||
losses: list[float | None],
|
||||
*,
|
||||
batch_size: int = 8,
|
||||
max_samples: int | None = None,
|
||||
):
|
||||
dataset = _MapDataset(texts)
|
||||
tokenizer = _ScriptedTokenizer(encodings)
|
||||
model = _LossModel(losses)
|
||||
handle = SimpleNamespace(model=model, tokenizer=tokenizer, task="causal_lm")
|
||||
evaluator = Evaluator(
|
||||
handle=handle,
|
||||
dataset=dataset,
|
||||
batch_size=batch_size,
|
||||
max_samples=max_samples,
|
||||
)
|
||||
return evaluator, dataset, tokenizer, model
|
||||
|
||||
|
||||
def test_causal_lm_skips_blank_rows_bounds_scan_and_masks_padding():
|
||||
evaluator, dataset, tokenizer, model = _evaluator(
|
||||
["", " \t", "alpha", "beta", "must not be read"],
|
||||
[_encoding([[1, 2, 0], [3, 4, 5]], [[1, 1, 0], [1, 1, 1]])],
|
||||
[math.log(4.0)],
|
||||
batch_size=2,
|
||||
max_samples=2,
|
||||
)
|
||||
|
||||
result = evaluator.evaluate()
|
||||
|
||||
assert result["perplexity"] == pytest.approx(4.0)
|
||||
assert dataset.read_indices == [0, 1, 2, 3]
|
||||
assert dataset.selected_indices == [2, 3]
|
||||
assert tokenizer.calls == [["alpha", "beta"]]
|
||||
assert model.labels[0].tolist() == [[1, 2, -100], [3, 4, 5]]
|
||||
|
||||
|
||||
def test_causal_lm_rejects_all_empty_dataset():
|
||||
evaluator, dataset, tokenizer, model = _evaluator(["", " \n", None], [], [])
|
||||
|
||||
with pytest.raises(ValueError, match="no non-empty text"):
|
||||
evaluator.evaluate()
|
||||
|
||||
assert dataset.read_indices == [0, 1, 2]
|
||||
assert tokenizer.calls == []
|
||||
assert model.labels == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize("max_samples", [0, -1])
|
||||
def test_causal_lm_requires_positive_max_samples(max_samples):
|
||||
evaluator, dataset, _tokenizer, _model = _evaluator(
|
||||
["alpha"],
|
||||
[],
|
||||
[],
|
||||
max_samples=max_samples,
|
||||
)
|
||||
|
||||
with pytest.raises(ValueError, match="max_samples must be a positive integer"):
|
||||
evaluator.evaluate()
|
||||
|
||||
assert dataset.read_indices == []
|
||||
|
||||
|
||||
def test_causal_lm_rejects_tokenless_and_one_token_sequences():
|
||||
evaluator, _dataset, _tokenizer, model = _evaluator(
|
||||
["tokenless", "single"],
|
||||
[
|
||||
_encoding(torch.empty((1, 0)), torch.empty((1, 0))),
|
||||
_encoding([[1]], [[1]]),
|
||||
],
|
||||
[],
|
||||
batch_size=1,
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="zero valid prediction tokens"):
|
||||
evaluator.evaluate()
|
||||
|
||||
assert model.labels == []
|
||||
|
||||
|
||||
def test_causal_lm_rejects_missing_loss():
|
||||
evaluator, _dataset, _tokenizer, _model = _evaluator(
|
||||
["alpha"],
|
||||
[_encoding([[1, 2]], [[1, 1]])],
|
||||
[None],
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="returned no loss"):
|
||||
evaluator.evaluate()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("loss", [float("nan"), float("inf"), float("-inf")])
|
||||
def test_causal_lm_rejects_non_finite_loss(loss):
|
||||
evaluator, _dataset, _tokenizer, _model = _evaluator(
|
||||
["alpha"],
|
||||
[_encoding([[1, 2]], [[1, 1]])],
|
||||
[loss],
|
||||
)
|
||||
|
||||
with pytest.raises(RuntimeError, match="Non-finite evaluation loss"):
|
||||
evaluator.evaluate()
|
||||
|
||||
|
||||
def test_causal_lm_perplexity_is_weighted_by_prediction_tokens():
|
||||
evaluator, _dataset, _tokenizer, _model = _evaluator(
|
||||
["three tokens", "two tokens"],
|
||||
[
|
||||
_encoding([[1, 2, 3]], [[1, 1, 1]]),
|
||||
_encoding([[4, 5]], [[1, 1]]),
|
||||
],
|
||||
[math.log(2.0), math.log(8.0)],
|
||||
batch_size=1,
|
||||
)
|
||||
|
||||
result = evaluator.evaluate()
|
||||
|
||||
expected = math.exp((2 * math.log(2.0) + math.log(8.0)) / 3)
|
||||
assert result["perplexity"] == pytest.approx(expected)
|
||||
|
||||
|
||||
def test_report_creates_nested_output_directory(monkeypatch, tmp_path):
|
||||
import obliteratus.reporting.report
|
||||
from obliteratus import cli
|
||||
|
||||
results_path = tmp_path / "results.json"
|
||||
results_path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"model_name": "fixture",
|
||||
"baseline_metrics": {"perplexity": 4.0},
|
||||
"results": [],
|
||||
}
|
||||
)
|
||||
)
|
||||
report = MagicMock()
|
||||
monkeypatch.setattr(
|
||||
obliteratus.reporting.report,
|
||||
"AblationReport",
|
||||
Mock(return_value=report),
|
||||
)
|
||||
output_dir = tmp_path / "nested" / "plots"
|
||||
|
||||
cli._cmd_report(
|
||||
SimpleNamespace(results_json=str(results_path), output_dir=str(output_dir))
|
||||
)
|
||||
|
||||
assert output_dir.is_dir()
|
||||
report.plot_impact.assert_called_once_with(
|
||||
metric="perplexity",
|
||||
output_path=output_dir / "impact.png",
|
||||
)
|
||||
report.plot_heatmap.assert_called_once_with(output_path=output_dir / "heatmap.png")
|
||||
Reference in New Issue
Block a user