mirror of
https://github.com/elder-plinius/OBLITERATUS.git
synced 2026-08-17 16:37:30 +02:00
test: cover evaluator skip branches
This commit is contained in:
@@ -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"],
|
||||
|
||||
Reference in New Issue
Block a user