Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -75,3 +75,16 @@ sacrebleu = ["py.typed"]

[tool.setuptools_scm]
version_file = "sacrebleu/version.py"

[tool.ruff.lint]
ignore = [
"B019",
"B020",
"PLR1704",
"PLW1508",
"RUF012",
"SIM115",
"SIM117",
"TRY002",
"TRY203",
]
28 changes: 13 additions & 15 deletions sacrebleu/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-

# Copyright 2017--2018 Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
Expand Down Expand Up @@ -42,25 +40,25 @@
from .version import __version__

__all__ = [
"smart_open",
"SACREBLEU_DIR",
"download_test_set",
"get_source_file",
"get_reference_files",
"get_available_testsets",
"get_langpairs_for_testset",
"extract_word_ngrams",
"extract_char_ngrams",
"DATASETS",
"BLEU",
"CHRF",
"DATASETS",
"SACREBLEU_DIR",
"TER",
"__version__",
"corpus_bleu",
"corpus_chrf",
"corpus_ter",
"download_test_set",
"extract_char_ngrams",
"extract_word_ngrams",
"get_available_testsets",
"get_langpairs_for_testset",
"get_reference_files",
"get_source_file",
"raw_corpus_bleu",
"sentence_bleu",
"corpus_chrf",
"sentence_chrf",
"corpus_ter",
"sentence_ter",
"__version__",
"smart_open",
]
2 changes: 0 additions & 2 deletions sacrebleu/__main__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,3 @@
#!/usr/bin/env python3
# -*- coding: utf-8 -*-

# Copyright 2017--2018 Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
Expand Down
8 changes: 5 additions & 3 deletions sacrebleu/compat.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
from typing import Sequence, Optional
from __future__ import annotations

from collections.abc import Sequence

from .metrics import BLEU, CHRF, TER, BLEUScore, CHRFScore, TERScore

Expand Down Expand Up @@ -39,7 +41,7 @@ def corpus_bleu(hypotheses: Sequence[str],

def raw_corpus_bleu(hypotheses: Sequence[str],
references: Sequence[Sequence[str]],
smooth_value: Optional[float] = BLEU.SMOOTH_DEFAULTS['floor']) -> BLEUScore:
smooth_value: float | None = BLEU.SMOOTH_DEFAULTS['floor']) -> BLEUScore:
"""Computes BLEU for a corpus against a single (or multiple) reference(s).
This convenience function assumes a particular set of arguments i.e.
it disables tokenization and applies a `floor` smoothing with value `0.1`.
Expand All @@ -64,7 +66,7 @@ def raw_corpus_bleu(hypotheses: Sequence[str],
def sentence_bleu(hypothesis: str,
references: Sequence[str],
smooth_method: str = 'exp',
smooth_value: Optional[float] = None,
smooth_value: float | None = None,
lowercase: bool = False,
tokenize=BLEU.TOKENIZER_DEFAULT,
use_effective_order: bool = True) -> BLEUScore:
Expand Down
6 changes: 2 additions & 4 deletions sacrebleu/dataset/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,3 @@
#!/usr/bin/env python
# -*- coding: utf-8 -*-

# Copyright 2017--2018 Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
Expand Down Expand Up @@ -69,8 +67,8 @@
k: {d.split("=")[0]: d.split("=")[1] for d in v.split()}
for (k, v) in _SUBSETS.items()
}
COUNTRIES = sorted(list({v.split("-")[0] for v in SUBSETS["wmt19"].values()}))
DOMAINS = sorted(list({v.split("-")[1] for v in SUBSETS["wmt19"].values()}))
COUNTRIES = sorted({v.split("-")[0] for v in SUBSETS["wmt19"].values()})
DOMAINS = sorted({v.split("-")[1] for v in SUBSETS["wmt19"].values()})

DATASETS = {
# wmt
Expand Down
4 changes: 2 additions & 2 deletions sacrebleu/dataset/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,8 +28,8 @@
print("Downloading ", url)
with urllib.request.urlopen(url) as f:
data = f.read()
except Exception as exc:
raise (exc)
except Exception:
raise

if hashlib.md5(data).hexdigest() != md5_hash:
print("MD5 check failed for", url)
Expand Down
24 changes: 11 additions & 13 deletions sacrebleu/dataset/base.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,11 @@
"""
The base class for all types of datasets.
"""
from __future__ import annotations

import os
import re
from abc import ABCMeta, abstractmethod
from typing import Dict, List, Optional

from ..utils import SACREBLEU_DIR, download_file, smart_open

Expand All @@ -13,11 +14,11 @@ class Dataset(metaclass=ABCMeta):
def __init__(
self,
name: str,
data: Optional[List[str]] = None,
description: Optional[str] = None,
citation: Optional[str] = None,
md5: Optional[List[str]] = None,
langpairs=Dict[str, List[str]],
data: list[str] | None = None,
description: str | None = None,
citation: str | None = None,
md5: list[str] | None = None,
langpairs=None,
**kwargs,
):
"""
Expand All @@ -35,7 +36,7 @@ def __init__(
self.description = description
self.citation = citation
self.md5 = md5
self.langpairs = langpairs
self.langpairs = {} if langpairs is None else langpairs
self.kwargs = kwargs

# Don't do any downloading or further processing now.
Expand Down Expand Up @@ -118,9 +119,8 @@ def process_to_text(self, langpair=None) -> None:

:param langpair: The language pair to process. e.g. "en-de". If None, all files will be processed.
"""
pass

def fieldnames(self, langpair) -> List[str]:
def fieldnames(self, langpair) -> list[str]:
"""
Return a list of all the field names. For most source, this is just
the source and the reference. For others, it might include the document
Expand All @@ -141,8 +141,7 @@ def __iter__(self, langpair):
all_files = self.get_files(langpair)
all_fins = [smart_open(f) for f in all_files]

for item in zip(*all_fins):
yield item
yield from zip(*all_fins)

def source(self, langpair):
"""
Expand All @@ -160,8 +159,7 @@ def references(self, langpair):
ref_files = self.get_reference_files(langpair)
ref_fins = [smart_open(f) for f in ref_files]

for item in zip(*ref_fins):
yield item
yield from zip(*ref_fins)

def get_source_file(self, langpair):
all_files = self.get_files(langpair)
Expand Down
2 changes: 1 addition & 1 deletion sacrebleu/dataset/fake_sgml.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@ def process_to_text(self, langpair=None):
origin_file = os.path.join(self._rawdir, origin_file)
output_file = self._get_txt_file_path(langpair, field)

if field.startswith("src") or field.startswith("ref"):
if field.startswith(("src", "ref")):
self._convert_format(origin_file, output_file)
else:
# document metadata keys
Expand Down
1 change: 0 additions & 1 deletion sacrebleu/dataset/iwslt_xml.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,4 +5,3 @@ class IWSLTXMLDataset(FakeSGMLDataset):
"""IWSLT dataset format. Can be parsed with the lxml parser."""

# Same as FakeSGMLDataset. Nothing to do here.
pass
7 changes: 3 additions & 4 deletions sacrebleu/dataset/wmt_xml.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,11 @@
import os
from collections import defaultdict

import lxml.etree as ET

from ..utils import smart_open
from .base import Dataset

from collections import defaultdict


def _get_field_by_translator(translator):
if not translator:
Expand Down Expand Up @@ -92,14 +91,14 @@ def get_sents(doc):

for seg_id in sorted(src_sents.keys()):
# no ref translation is available for this segment
if not any([value.get(seg_id, "") for value in trans_to_ref.values()]):
if not any(value.get(seg_id, "") for value in trans_to_ref.values()):
continue
for translator in translators:
refs[_get_field_by_translator(translator)].append(
trans_to_ref.get(translator, {translator: {}}).get(seg_id, "")
)
src.append(src_sents[seg_id])
for system_name in hyps.keys():
for system_name in hyps:
systems[system_name].append(hyps[system_name][seg_id])
docids.append(doc.attrib["id"])
orig_langs.append(doc.attrib["origlang"])
Expand Down
6 changes: 3 additions & 3 deletions sacrebleu/metrics/__init__.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,8 @@
"""The implementation of various metrics."""

from .bleu import BLEU, BLEUScore # noqa: F401
from .chrf import CHRF, CHRFScore # noqa: F401
from .ter import TER, TERScore # noqa: F401
from .bleu import BLEU, BLEUScore # noqa: F401
from .chrf import CHRF, CHRFScore # noqa: F401
from .ter import TER, TERScore # noqa: F401

METRICS = {
'BLEU': BLEU,
Expand Down
31 changes: 14 additions & 17 deletions sacrebleu/metrics/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,14 @@
of abstract methods. This way, a correctly implemented metric will work
seamlessly with the rest of the codebase.
"""
from __future__ import annotations

import json
import logging
import statistics
from abc import ABCMeta, abstractmethod
from typing import Any, Dict, List, Optional, Sequence
from collections.abc import Sequence
from typing import Any

from ..version import __version__

Expand Down Expand Up @@ -88,7 +90,7 @@ def format(

return full_score

def estimate_ci(self, scores: List["Score"]):
def estimate_ci(self, scores: list[Score]):
"""Takes a list of scores and stores mean, stdev and 95% confidence
interval around the mean.

Expand Down Expand Up @@ -237,7 +239,7 @@ def _check_sentence_score_args(self, hyp: str, refs: Sequence[str]):
raise TypeError(f"{prefix}: {err_msg}")

def _check_corpus_score_args(
self, hyps: Sequence[str], refs: Optional[Sequence[Sequence[str]]]
self, hyps: Sequence[str], refs: Sequence[Sequence[str]] | None
):
"""Performs sanity checks on `corpus_score` method's arguments.

Expand Down Expand Up @@ -289,22 +291,20 @@ def _check_corpus_score_args(
raise TypeError(f"{prefix}: {err_msg}")

@abstractmethod
def _aggregate_and_compute(self, stats: List[List[Any]]) -> Any:
def _aggregate_and_compute(self, stats: list[list[Any]]) -> Any:
"""Computes the final score given the pre-computed match statistics.

:param stats: A list of segment-level statistics.
:return: A `Score` instance.
"""
pass

@abstractmethod
def _compute_score_from_stats(self, stats: List[Any]) -> Any:
def _compute_score_from_stats(self, stats: list[Any]) -> Any:
"""Computes the final score from already aggregated statistics.

:param stats: A list or numpy array of segment-level statistics.
:return: A `Score` object.
"""
pass

@abstractmethod
def _preprocess_segment(self, sent: str) -> str:
Expand All @@ -314,22 +314,20 @@ def _preprocess_segment(self, sent: str) -> str:
:param sent: The input sentence.
:return: The pre-processed output sentence.
"""
pass

@abstractmethod
def _extract_reference_info(self, refs: Sequence[str]) -> Dict[str, Any]:
def _extract_reference_info(self, refs: Sequence[str]) -> dict[str, Any]:
"""Given a list of reference segments, extract the required
information (such as n-grams for BLEU and chrF). This should be implemented
for the generic `_cache_references()` to work across all metrics.

:param refs: A sequence of strings.
"""
pass

@abstractmethod
def _compute_segment_statistics(
self, hypothesis: str, ref_kwargs: Dict
) -> List[Any]:
self, hypothesis: str, ref_kwargs: dict
) -> list[Any]:
"""Given a (pre-processed) hypothesis sentence and already computed
reference info, returns the best match statistics across the
references. The return type is usually a List of ints or floats.
Expand All @@ -339,9 +337,8 @@ def _compute_segment_statistics(
within. This is formulated as a dictionary as different metrics may
require different information regarding a reference segment.
"""
pass

def _cache_references(self, references: Sequence[Sequence[str]]) -> List[Any]:
def _cache_references(self, references: Sequence[Sequence[str]]) -> list[Any]:
"""Given the full set of document references, extract segment n-grams
(or other necessary information) for caching purposes.

Expand Down Expand Up @@ -371,15 +368,15 @@ def _cache_references(self, references: Sequence[Sequence[str]]) -> List[Any]:
ref_cache.append(self._extract_reference_info(lines))

if len(num_refs) == 1:
self.num_refs = list(num_refs)[0]
self.num_refs = next(iter(num_refs))
else:
# A variable number of refs exist
self.num_refs = -1

return ref_cache

def _extract_corpus_statistics(
self, hypotheses: Sequence[str], references: Optional[Sequence[Sequence[str]]]
self, hypotheses: Sequence[str], references: Sequence[Sequence[str]] | None
) -> Any:
"""Reads the corpus and returns sentence-level match statistics for
faster re-computations esp. during statistical tests.
Expand Down Expand Up @@ -440,7 +437,7 @@ def sentence_score(self, hypothesis: str, references: Sequence[str]) -> Any:
def corpus_score(
self,
hypotheses: Sequence[str],
references: Optional[Sequence[Sequence[str]]],
references: Sequence[Sequence[str]] | None,
n_bootstrap: int = 1,
) -> Any:
"""Compute the metric for a corpus against a single (or multiple) reference(s).
Expand Down
Loading
Loading