test: cover evaluator skip branches

This commit is contained in:
Joseph Magly
2026-08-14 19:03:36 -04:00
parent 78caa5eab2
commit 9b93e20276
+38
View File
@@ -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"],