diff --git a/tests/test_evaluator.py b/tests/test_evaluator.py index 43487cd..7f834d2 100644 --- a/tests/test_evaluator.py +++ b/tests/test_evaluator.py @@ -36,6 +36,14 @@ class _MapDataset: return _MapDataset([self.texts[index] for index in self.selected_indices]) +class _ChangingDataset(_MapDataset): + """Simulate a transformed dataset changing after index selection.""" + + def select(self, indices): + self.selected_indices = list(indices) + return _MapDataset([""]) + + class _Encoding(dict): def to(self, _device): return self @@ -154,6 +162,36 @@ def test_causal_lm_rejects_tokenless_and_one_token_sequences(): assert model.labels == [] +def test_causal_lm_skips_batch_emptied_by_dataset_transform(): + dataset = _ChangingDataset(["alpha"]) + tokenizer = _ScriptedTokenizer([]) + model = _LossModel([]) + handle = SimpleNamespace(model=model, tokenizer=tokenizer, task="causal_lm") + + with pytest.raises(RuntimeError, match="zero valid prediction tokens"): + Evaluator(handle=handle, dataset=dataset).evaluate() + + assert tokenizer.calls == [] + assert model.labels == [] + + +def test_causal_lm_skips_zero_target_batch_then_reports_valid_perplexity(): + evaluator, _dataset, _tokenizer, model = _evaluator( + ["padded target", "valid target"], + [ + _encoding([[1, 0]], [[1, 0]]), + _encoding([[2, 3]], [[1, 1]]), + ], + [math.log(3.0)], + batch_size=1, + ) + + result = evaluator.evaluate() + + assert result["perplexity"] == pytest.approx(3.0) + assert len(model.labels) == 1 + + def test_causal_lm_rejects_missing_loss(): evaluator, _dataset, _tokenizer, _model = _evaluator( ["alpha"],