diff --git a/sacrebleu/metrics/base.py b/sacrebleu/metrics/base.py index 8aa693a..6a50631 100644 --- a/sacrebleu/metrics/base.py +++ b/sacrebleu/metrics/base.py @@ -396,6 +396,12 @@ def _extract_corpus_statistics( else: raise RuntimeError("No references provided and the cache is empty.") + if len(hypotheses) != len(ref_cache): + raise TypeError( + f"{self.__class__.__name__}: The reference cache must have the " + "same length as `hypotheses`." + ) + stats = [] tok_count = 0 diff --git a/test/test_metrics_base.py b/test/test_metrics_base.py index 679ec71..713d76b 100644 --- a/test/test_metrics_base.py +++ b/test/test_metrics_base.py @@ -2,7 +2,9 @@ import pytest +from sacrebleu.metrics import BLEU, CHRF, TER from sacrebleu.metrics.base import Metric +from sacrebleu.significance import PairedTest class ConcreteTestMetric(Metric): @@ -353,3 +355,56 @@ def test_check_corpus_score_args(hyps, refs, expected_context): metric = ConcreteTestMetric() with expected_context: metric._check_corpus_score_args(hyps=hyps, refs=refs) + + +@pytest.mark.parametrize("metric_type", [BLEU, CHRF, TER]) +@pytest.mark.parametrize("hyp_count", [1, 3]) +def test_cached_corpus_rejects_length_mismatch(metric_type, hyp_count): + metric = metric_type(references=[["one two three four"] * 2]) + with pytest.raises(TypeError, match="same length"): + metric.corpus_score(["one two three four"] * hyp_count, None) + + +@pytest.mark.parametrize("metric_type", [BLEU, CHRF, TER]) +def test_cached_corpus_matches_explicit_references(metric_type): + refs = [ + ["one two three four", "five six seven eight"], + ["one two three five", None], + ] + hyps = ["one two three four", "five six seven nine"] + metric = metric_type(references=refs) + expected = metric_type().corpus_score(hyps, refs).score + assert metric.corpus_score(hyps, None).score == pytest.approx(expected) + + # Explicit references take precedence over a differently sized cache. + assert metric.corpus_score(hyps[:1], [refs[0][:1]]).score == pytest.approx( + metric_type().corpus_score(hyps[:1], [refs[0][:1]]).score + ) + assert metric.sentence_score(hyps[0], [refs[0][0]]).score == pytest.approx( + metric_type().sentence_score(hyps[0], [refs[0][0]]).score + ) + assert metric.corpus_score(hyps, None).score == pytest.approx(expected) + + +@pytest.mark.parametrize("test_type", ["bs", "ar"]) +@pytest.mark.parametrize("mismatched_system", ["baseline", "candidate"]) +@pytest.mark.parametrize("n_jobs", [1, 2]) +def test_paired_test_rejects_cached_length_mismatch( + test_type, mismatched_system, n_jobs +): + refs = [["one two three four", "five six seven eight"]] + systems = [("baseline", refs[0].copy()), ("candidate", refs[0].copy())] + for name, hypotheses in systems: + if name == mismatched_system: + hypotheses.append("this extra sentence must not be discarded") + + with pytest.raises(TypeError, match="same length"): + PairedTest( + systems, + {"BLEU": BLEU(references=refs)}, + references=None, + test_type=test_type, + n_samples=10, + n_ar_confidence=-1, + n_jobs=n_jobs, + )()