diff --git a/pyproject.toml b/pyproject.toml index fbcde65b..27223d1b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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", +] diff --git a/sacrebleu/__init__.py b/sacrebleu/__init__.py index 111f284d..aa3c2157 100644 --- a/sacrebleu/__init__.py +++ b/sacrebleu/__init__.py @@ -1,5 +1,3 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- # Copyright 2017--2018 Amazon.com, Inc. or its affiliates. All Rights Reserved. # @@ -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", ] diff --git a/sacrebleu/__main__.py b/sacrebleu/__main__.py index 3833741e..d5a72e4c 100644 --- a/sacrebleu/__main__.py +++ b/sacrebleu/__main__.py @@ -1,5 +1,3 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- # Copyright 2017--2018 Amazon.com, Inc. or its affiliates. All Rights Reserved. # diff --git a/sacrebleu/compat.py b/sacrebleu/compat.py index 57359603..f912181c 100644 --- a/sacrebleu/compat.py +++ b/sacrebleu/compat.py @@ -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 @@ -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`. @@ -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: diff --git a/sacrebleu/dataset/__init__.py b/sacrebleu/dataset/__init__.py index 68c35cef..4eb7788e 100644 --- a/sacrebleu/dataset/__init__.py +++ b/sacrebleu/dataset/__init__.py @@ -1,5 +1,3 @@ -#!/usr/bin/env python -# -*- coding: utf-8 -*- # Copyright 2017--2018 Amazon.com, Inc. or its affiliates. All Rights Reserved. # @@ -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 diff --git a/sacrebleu/dataset/__main__.py b/sacrebleu/dataset/__main__.py index 5b13d59a..175dd779 100644 --- a/sacrebleu/dataset/__main__.py +++ b/sacrebleu/dataset/__main__.py @@ -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) diff --git a/sacrebleu/dataset/base.py b/sacrebleu/dataset/base.py index cf3c092f..16edc392 100644 --- a/sacrebleu/dataset/base.py +++ b/sacrebleu/dataset/base.py @@ -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 @@ -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, ): """ @@ -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. @@ -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 @@ -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): """ @@ -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) diff --git a/sacrebleu/dataset/fake_sgml.py b/sacrebleu/dataset/fake_sgml.py index d1f63812..d6e387fc 100644 --- a/sacrebleu/dataset/fake_sgml.py +++ b/sacrebleu/dataset/fake_sgml.py @@ -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 diff --git a/sacrebleu/dataset/iwslt_xml.py b/sacrebleu/dataset/iwslt_xml.py index 4381271d..99df3d05 100644 --- a/sacrebleu/dataset/iwslt_xml.py +++ b/sacrebleu/dataset/iwslt_xml.py @@ -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 diff --git a/sacrebleu/dataset/wmt_xml.py b/sacrebleu/dataset/wmt_xml.py index 1aedf1d3..00850ce1 100644 --- a/sacrebleu/dataset/wmt_xml.py +++ b/sacrebleu/dataset/wmt_xml.py @@ -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: @@ -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"]) diff --git a/sacrebleu/metrics/__init__.py b/sacrebleu/metrics/__init__.py index a18c2277..465595ea 100644 --- a/sacrebleu/metrics/__init__.py +++ b/sacrebleu/metrics/__init__.py @@ -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, diff --git a/sacrebleu/metrics/base.py b/sacrebleu/metrics/base.py index 5315eb4d..8aa693af 100644 --- a/sacrebleu/metrics/base.py +++ b/sacrebleu/metrics/base.py @@ -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__ @@ -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. @@ -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. @@ -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: @@ -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. @@ -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. @@ -371,7 +368,7 @@ 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 @@ -379,7 +376,7 @@ def _cache_references(self, references: Sequence[Sequence[str]]) -> List[Any]: 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. @@ -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). diff --git a/sacrebleu/metrics/bleu.py b/sacrebleu/metrics/bleu.py index 14f31f5a..50beddae 100644 --- a/sacrebleu/metrics/bleu.py +++ b/sacrebleu/metrics/bleu.py @@ -1,13 +1,14 @@ """The implementation of the BLEU metric (Papineni et al., 2002).""" +from __future__ import annotations -import math import logging +import math +from collections.abc import Sequence from importlib import import_module -from typing import List, Sequence, Optional, Dict, Any +from typing import Any from ..utils import my_log, sum_of_lists - -from .base import Score, Signature, Metric +from .base import Metric, Score, Signature from .helpers import extract_all_word_ngrams sacrelogger = logging.getLogger('sacrebleu') @@ -88,8 +89,8 @@ class BLEUScore(Score): :param sys_len: The cumulative system length. :param ref_len: The cumulative reference length. """ - def __init__(self, score: float, counts: List[int], totals: List[int], - precisions: List[float], bp: float, + def __init__(self, score: float, counts: list[int], totals: list[int], + precisions: list[float], bp: float, sys_len: int, ref_len: int): """`BLEUScore` initializer.""" super().__init__('BLEU', score) @@ -127,7 +128,7 @@ class BLEU(Metric): across many systems. """ - SMOOTH_DEFAULTS: Dict[str, Optional[float]] = { + SMOOTH_DEFAULTS: dict[str, float | None] = { # The defaults for `floor` and `add-k` are obtained from the following paper # A Systematic Comparison of Smoothing Techniques for Sentence-Level BLEU # Boxing Chen and Colin Cherry @@ -155,13 +156,13 @@ class BLEU(Metric): def __init__(self, lowercase: bool = False, force: bool = False, - tokenize: Optional[str] = None, + tokenize: str | None = None, smooth_method: str = 'exp', - smooth_value: Optional[float] = None, + smooth_value: float | None = None, max_ngram_order: int = MAX_NGRAM_ORDER, effective_order: bool = False, trg_lang: str = '', - references: Optional[Sequence[Sequence[str]]] = None): + references: Sequence[Sequence[str]] | None = None): """`BLEU` initializer.""" super().__init__() @@ -174,7 +175,7 @@ def __init__(self, lowercase: bool = False, self.effective_order = effective_order # Sanity check - assert self.smooth_method in self.SMOOTH_DEFAULTS.keys(), \ + assert self.smooth_method in self.SMOOTH_DEFAULTS, \ "Unknown smooth_method {self.smooth_method!r}" # If the tokenizer wasn't specified, choose it according to the @@ -210,8 +211,8 @@ def __init__(self, lowercase: bool = False, self._ref_cache = self._cache_references(references) @staticmethod - def compute_bleu(correct: List[int], - total: List[int], + def compute_bleu(correct: list[int], + total: list[int], sys_len: int, ref_len: int, smooth_method: str = 'none', @@ -239,7 +240,7 @@ def compute_bleu(correct: List[int], :param max_ngram_order: If given, it overrides the maximum n-gram order (default: 4) when computing precisions. :return: A `BLEUScore` instance. """ - assert smooth_method in BLEU.SMOOTH_DEFAULTS.keys(), \ + assert smooth_method in BLEU.SMOOTH_DEFAULTS, \ "Unknown smooth_method {smooth_method!r}" # Fetch the default value for floor and add-k @@ -301,7 +302,7 @@ def _preprocess_segment(self, sent: str) -> str: sent = sent.lower() return self.tokenizer(sent.rstrip()) - def _compute_score_from_stats(self, stats: List[int]) -> BLEUScore: + def _compute_score_from_stats(self, stats: list[int]) -> BLEUScore: """Computes the final score from already aggregated statistics. :param stats: A list or numpy array of segment-level statistics. @@ -316,7 +317,7 @@ def _compute_score_from_stats(self, stats: List[int]) -> BLEUScore: max_ngram_order=self.max_ngram_order ) - def _aggregate_and_compute(self, stats: List[List[int]]) -> BLEUScore: + def _aggregate_and_compute(self, stats: list[list[int]]) -> BLEUScore: """Computes the final BLEU score given the pre-computed corpus statistics. :param stats: A list of segment-level statistics @@ -324,7 +325,7 @@ def _aggregate_and_compute(self, stats: List[List[int]]) -> BLEUScore: """ return self._compute_score_from_stats(sum_of_lists(stats)) - def _get_closest_ref_len(self, hyp_len: int, ref_lens: List[int]) -> int: + def _get_closest_ref_len(self, hyp_len: int, ref_lens: list[int]) -> int: """Given a hypothesis length and a list of reference lengths, returns the closest reference length to be used by BLEU. @@ -344,7 +345,7 @@ def _get_closest_ref_len(self, hyp_len: int, ref_lens: List[int]) -> int: return closest_len - 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 n-grams and reference lengths. The latter will be useful when comparing hypothesis and reference lengths for BLEU. @@ -372,7 +373,7 @@ def _extract_reference_info(self, refs: Sequence[str]) -> Dict[str, Any]: return {'ref_ngrams': ngrams, 'ref_lens': ref_lens} def _compute_segment_statistics(self, hypothesis: str, - ref_kwargs: Dict) -> List[int]: + ref_kwargs: dict) -> list[int]: """Given a (pre-processed) hypothesis sentence and already computed reference n-grams & lengths, returns the best match statistics across the references. diff --git a/sacrebleu/metrics/chrf.py b/sacrebleu/metrics/chrf.py index f7d4f685..04274dc0 100644 --- a/sacrebleu/metrics/chrf.py +++ b/sacrebleu/metrics/chrf.py @@ -1,10 +1,11 @@ """The implementation of chrF (Popović 2015) and chrF++ (Popović 2017) metrics.""" +from __future__ import annotations -from typing import List, Sequence, Optional, Dict from collections import Counter +from collections.abc import Sequence from ..utils import sum_of_lists -from .base import Score, Signature, Metric +from .base import Metric, Score, Signature from .helpers import extract_all_char_ngrams, extract_word_ngrams @@ -89,7 +90,7 @@ def __init__(self, char_order: int = CHAR_ORDER, lowercase: bool = False, whitespace: bool = False, eps_smoothing: bool = False, - references: Optional[Sequence[Sequence[str]]] = None): + references: Sequence[Sequence[str]] | None = None): """`CHRF` initializer.""" super().__init__() @@ -106,7 +107,7 @@ def __init__(self, char_order: int = CHAR_ORDER, self._ref_cache = self._cache_references(references) @staticmethod - def _get_match_statistics(hyp_ngrams: Counter, ref_ngrams: Counter) -> List[int]: + def _get_match_statistics(hyp_ngrams: Counter, ref_ngrams: Counter) -> list[int]: """Computes the match statistics between hypothesis and reference n-grams. :param hyp_ngrams: A `Counter` holding hypothesis n-grams. @@ -128,7 +129,7 @@ def _get_match_statistics(hyp_ngrams: Counter, ref_ngrams: Counter) -> List[int] match_count, ] - def _remove_punctuation(self, sent: str) -> List[str]: + def _remove_punctuation(self, sent: str) -> list[str]: """Separates out punctuations from beginning and end of words for chrF. Adapted from https://github.com/m-popovic/chrF @@ -157,7 +158,7 @@ def _preprocess_segment(self, sent: str) -> str: """ return sent.lower() if self.lowercase else sent - def _compute_f_score(self, statistics: List[int]) -> float: + def _compute_f_score(self, statistics: list[int]) -> float: """Compute the chrF score given the n-gram match statistics. :param statistics: A flattened list of 3 * (`char_order` + `word_order`) @@ -202,7 +203,7 @@ def _compute_f_score(self, statistics: List[int]) -> float: else: return 0.0 - def _compute_score_from_stats(self, stats: List[int]) -> CHRFScore: + def _compute_score_from_stats(self, stats: list[int]) -> CHRFScore: """Computes the final score from already aggregated statistics. :param stats: A list or numpy array of segment-level statistics. @@ -212,7 +213,7 @@ def _compute_score_from_stats(self, stats: List[int]) -> CHRFScore: self._compute_f_score(stats), self.char_order, self.word_order, self.beta) - def _aggregate_and_compute(self, stats: List[List[int]]) -> CHRFScore: + def _aggregate_and_compute(self, stats: list[list[int]]) -> CHRFScore: """Computes the final score given the pre-computed corpus statistics. :param stats: A list of segment-level statistics @@ -220,7 +221,7 @@ def _aggregate_and_compute(self, stats: List[List[int]]) -> CHRFScore: """ return self._compute_score_from_stats(sum_of_lists(stats)) - def _extract_reference_info(self, refs: Sequence[str]) -> Dict[str, List[List[Counter]]]: + def _extract_reference_info(self, refs: Sequence[str]) -> dict[str, list[list[Counter]]]: """Given a list of reference segments, extract the character and word n-grams. :param refs: A sequence of reference segments. @@ -244,7 +245,7 @@ def _extract_reference_info(self, refs: Sequence[str]) -> Dict[str, List[List[Co return {'ref_ngrams': ngrams} def _compute_segment_statistics( - self, hypothesis: str, ref_kwargs: Dict) -> List[int]: + self, hypothesis: str, ref_kwargs: dict) -> list[int]: """Given a (pre-processed) hypothesis sentence and already computed reference n-grams, returns the best match statistics across the references. diff --git a/sacrebleu/metrics/helpers.py b/sacrebleu/metrics/helpers.py index 72ec1446..ffddd984 100644 --- a/sacrebleu/metrics/helpers.py +++ b/sacrebleu/metrics/helpers.py @@ -1,10 +1,11 @@ """Various utility functions for word and character n-gram extraction.""" +from __future__ import annotations + from collections import Counter -from typing import List, Tuple -def extract_all_word_ngrams(line: str, min_order: int, max_order: int) -> Tuple[Counter, int]: +def extract_all_word_ngrams(line: str, min_order: int, max_order: int) -> tuple[Counter, int]: """Extracts all ngrams (min_order <= n <= max_order) from a sentence. :param line: A string sentence. @@ -17,13 +18,13 @@ def extract_all_word_ngrams(line: str, min_order: int, max_order: int) -> Tuple[ tokens = line.split() for n in range(min_order, max_order + 1): - for i in range(0, len(tokens) - n + 1): + for i in range(len(tokens) - n + 1): ngrams.append(tuple(tokens[i: i + n])) return Counter(ngrams), len(tokens) -def extract_word_ngrams(tokens: List[str], n: int) -> Counter: +def extract_word_ngrams(tokens: list[str], n: int) -> Counter: """Extracts n-grams with order `n` from a list of tokens. :param tokens: A list of tokens. @@ -48,7 +49,7 @@ def extract_char_ngrams(line: str, n: int, include_whitespace: bool = False) -> def extract_all_char_ngrams( - line: str, max_order: int, include_whitespace: bool = False) -> List[Counter]: + line: str, max_order: int, include_whitespace: bool = False) -> list[Counter]: """Extracts all character n-grams at once for convenience. :param line: A segment containing a sequence of words. diff --git a/sacrebleu/metrics/lib_ter.py b/sacrebleu/metrics/lib_ter.py index 2d2de494..5e5d09b9 100644 --- a/sacrebleu/metrics/lib_ter.py +++ b/sacrebleu/metrics/lib_ter.py @@ -16,8 +16,6 @@ import math -from typing import List, Tuple, Dict - _COST_INS = 1 _COST_DEL = 1 @@ -42,7 +40,7 @@ _FLIP_OPS = str.maketrans(_OP_INS + _OP_DEL, _OP_DEL + _OP_INS) -def translation_edit_rate(words_hyp: List[str], words_ref: List[str]) -> Tuple[int, int]: +def translation_edit_rate(words_hyp: list[str], words_ref: list[str]) -> tuple[int, int]: """Calculate the translation edit rate. :param words_hyp: Tokenized translation hypothesis. @@ -52,8 +50,6 @@ def translation_edit_rate(words_hyp: List[str], words_ref: List[str]) -> Tuple[i n_words_ref = len(words_ref) n_words_hyp = len(words_hyp) if n_words_ref == 0: - # FIXME: This trace here is not used? - trace = _OP_DEL * n_words_hyp # special treatment of empty refs return n_words_hyp, 0 @@ -75,14 +71,14 @@ def translation_edit_rate(words_hyp: List[str], words_ref: List[str]) -> Tuple[i shifts += 1 input_words = new_input_words - edit_distance, trace = cached_ed(input_words) + edit_distance, _trace = cached_ed(input_words) total_edits = shifts + edit_distance return total_edits, n_words_ref -def _shift(words_h: List[str], words_r: List[str], cached_ed, - checked_candidates: int) -> Tuple[int, List[str], int]: +def _shift(words_h: list[str], words_r: list[str], cached_ed, + checked_candidates: int) -> tuple[int, list[str], int]: """Attempt to shift words in hypothesis to match reference. Returns the shift that reduces the edit distance the most. @@ -166,7 +162,7 @@ def _shift(words_h: List[str], words_r: List[str], cached_ed, return best_score, shifted_words, checked_candidates -def _perform_shift(words: List[str], start: int, length: int, target: int) -> List[str]: +def _perform_shift(words: list[str], start: int, length: int, target: int) -> list[str]: """Perform a shift in `words` from `start` to `target`. :param words: Words to shift. @@ -189,7 +185,7 @@ def _perform_shift(words: List[str], start: int, length: int, target: int) -> Li + words[start: start + length] + words[length + target:] -def _find_shifted_pairs(words_h: List[str], words_r: List[str]): +def _find_shifted_pairs(words_h: list[str], words_r: list[str]): """Find matching word sub-sequences in two lists of words. Ignores sub-sequences starting at the same position. @@ -229,7 +225,7 @@ def _flip_trace(trace): return trace.translate(_FLIP_OPS) -def trace_to_alignment(trace: str) -> Tuple[Dict, List, List]: +def trace_to_alignment(trace: str) -> tuple[dict, list, list]: """Transform trace of edit operations into an alignment of the sequences. :param trace: Trace of edit operations (' '=no change or 's'/'i'/'d'). @@ -294,7 +290,7 @@ class BeamEditDistance: :param words_ref: A list of reference tokens. """ - def __init__(self, words_ref: List[str]): + def __init__(self, words_ref: list[str]): """`BeamEditDistance` initializer.""" self._words_ref = words_ref self._n_words_ref = len(self._words_ref) @@ -304,14 +300,14 @@ def __init__(self, words_ref: List[str]): self._initial_row = [(i * _COST_INS, _OP_INS) for i in range(self._n_words_ref + 1)] - self._cache = {} # type: Dict[str, Tuple] + self._cache: dict[str, tuple] = {} self._cache_size = 0 # Precomputed empty matrix row. Contains infinities so that beam search # avoids using the uninitialized cells. self._empty_row = [(_INT_INFINITY, _OP_UNDEF)] * (self._n_words_ref + 1) - def __call__(self, words_hyp: List[str]) -> Tuple[int, str]: + def __call__(self, words_hyp: list[str]) -> tuple[int, str]: """Calculate edit distance between self._words_ref and the hypothesis. Uses cache to skip some of the computation. @@ -333,8 +329,8 @@ def __call__(self, words_hyp: List[str]) -> Tuple[int, str]: return edit_distance, trace - def _edit_distance(self, words_h: List[str], start_h: int, - cache: List[List[Tuple[int, str]]]) -> Tuple[int, List, str]: + def _edit_distance(self, words_h: list[str], start_h: int, + cache: list[list[tuple[int, str]]]) -> tuple[int, list, str]: """Actual edit distance calculation. Can be initialized with the last cached row and a start position in @@ -422,7 +418,7 @@ def _edit_distance(self, words_h: List[str], start_h: int, return dist[-1][-1][0], dist[len(cache):], trace - def _add_cache(self, words_hyp: List[str], mat: List[List[Tuple]]): + def _add_cache(self, words_hyp: list[str], mat: list[list[tuple]]): """Add newly computed rows to cache. Since edit distance is only calculated on the hypothesis suffix that @@ -456,7 +452,7 @@ def _add_cache(self, words_hyp: List[str], mat: List[List[Tuple]]): value = node[word] node = value[0] - def _find_cache(self, words_hyp: List[str]) -> Tuple[int, List[List]]: + def _find_cache(self, words_hyp: list[str]) -> tuple[int, list[list]]: """Find the already computed rows of the edit distance matrix in cache. Returns a partially computed edit distance matrix. diff --git a/sacrebleu/metrics/ter.py b/sacrebleu/metrics/ter.py index 40f82218..abb1fbd0 100644 --- a/sacrebleu/metrics/ter.py +++ b/sacrebleu/metrics/ter.py @@ -15,11 +15,14 @@ # limitations under the License. -from typing import List, Dict, Sequence, Optional, Any +from __future__ import annotations + +from collections.abc import Sequence +from typing import Any from ..tokenizers.tokenizer_ter import TercomTokenizer from ..utils import sum_of_lists -from .base import Score, Signature, Metric +from .base import Metric, Score, Signature from .lib_ter import translation_edit_rate @@ -97,7 +100,7 @@ def __init__(self, normalized: bool = False, no_punct: bool = False, asian_support: bool = False, case_sensitive: bool = False, - references: Optional[Sequence[Sequence[str]]] = None): + references: Sequence[Sequence[str]] | None = None): """`TER` initializer.""" super().__init__() @@ -125,7 +128,7 @@ def _preprocess_segment(self, sent: str) -> str: """ return self.tokenizer(sent.rstrip()) - def _compute_score_from_stats(self, stats: List[float]) -> TERScore: + def _compute_score_from_stats(self, stats: list[float]) -> TERScore: """Computes the final score from already aggregated statistics. :param stats: A list or numpy array of segment-level statistics. @@ -142,7 +145,7 @@ def _compute_score_from_stats(self, stats: List[float]) -> TERScore: return TERScore(100 * score, total_edits, sum_ref_lengths) - def _aggregate_and_compute(self, stats: List[List[float]]) -> TERScore: + def _aggregate_and_compute(self, stats: list[list[float]]) -> TERScore: """Computes the final TER score given the pre-computed corpus statistics. :param stats: A list of segment-level statistics @@ -151,7 +154,7 @@ def _aggregate_and_compute(self, stats: List[List[float]]) -> TERScore: return self._compute_score_from_stats(sum_of_lists(stats)) def _compute_segment_statistics( - self, hypothesis: str, ref_kwargs: Dict) -> List[float]: + self, hypothesis: str, ref_kwargs: dict) -> list[float]: """Given a (pre-processed) hypothesis sentence and already computed reference words, returns the segment statistics required to compute the full TER score. @@ -173,13 +176,12 @@ def _compute_segment_statistics( for words_ref in ref_words: num_edits, ref_len = translation_edit_rate(words_hyp, words_ref) ref_lengths += ref_len - if num_edits < best_num_edits: - best_num_edits = num_edits + best_num_edits = min(best_num_edits, num_edits) avg_ref_len = ref_lengths / len(ref_words) return [best_num_edits, avg_ref_len] - 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, applies pre-processing & tokenization and returns list of tokens for each reference. diff --git a/sacrebleu/sacrebleu.py b/sacrebleu/sacrebleu.py index 6b7cd9e7..1122be9f 100755 --- a/sacrebleu/sacrebleu.py +++ b/sacrebleu/sacrebleu.py @@ -21,15 +21,14 @@ See the [README.md] file for more information. """ +import argparse import io -import os -import sys import logging +import os import pathlib -import argparse +import sys from collections import defaultdict - # Allows calling the script as a standalone utility # See: https://github.com/mjpost/sacrebleu/issues/86 if __package__ is None and __name__ == '__main__': @@ -37,23 +36,36 @@ sys.path.insert(0, str(parent)) __package__ = 'sacrebleu' +from . import __version__ as VERSION from .dataset import DATASETS from .metrics import METRICS -from .utils import smart_open, filter_subset, get_langpairs_for_testset, get_available_testsets -from .utils import print_test_set, print_subset_results, get_reference_files, download_test_set -from .utils import args_to_dict, sanity_check_lengths, print_results_table, print_single_results -from .utils import get_available_testsets_for_langpair, Color - -from . import __version__ as VERSION +from .utils import ( + Color, + args_to_dict, + download_test_set, + filter_subset, + get_available_testsets, + get_available_testsets_for_langpair, + get_langpairs_for_testset, + get_reference_files, + print_results_table, + print_single_results, + print_subset_results, + print_test_set, + sanity_check_lengths, + smart_open, +) sacrelogger = logging.getLogger('sacrebleu') try: # SIGPIPE is not available on Windows machines, throwing an exception. - from signal import SIGPIPE # type: ignore - # If SIGPIPE is available, change behaviour to default instead of ignore. - from signal import signal, SIG_DFL + from signal import ( # type: ignore[attr-defined] + SIG_DFL, + SIGPIPE, + signal, + ) signal(SIGPIPE, SIG_DFL) except ImportError: pass @@ -202,7 +214,7 @@ def parse_args(): '`json` and `text` apply to single-system mode only. This flag is overridden if the ' 'SACREBLEU_FORMAT environment variable is set to one of the valid choices (Default: %(default)s).') - arg_parser.add_argument('--version', '-V', action='version', version='%(prog)s {}'.format(VERSION)) + arg_parser.add_argument('--version', '-V', action='version', version=f'%(prog)s {VERSION}') args = arg_parser.parse_args() @@ -495,7 +507,7 @@ def main(): # Handle sentence level and quit if args.sentence_level: # one metric and one system in use for sentence-level - metric, system = list(metrics.values())[0], systems[0] + metric, system = next(iter(metrics.values())), systems[0] for hypothesis, *references in zip(system, *refs): score = metric.sentence_score(hypothesis, references) diff --git a/sacrebleu/significance.py b/sacrebleu/significance.py index a9c71d0a..214d0f76 100644 --- a/sacrebleu/significance.py +++ b/sacrebleu/significance.py @@ -1,7 +1,10 @@ -import os +from __future__ import annotations + import logging import multiprocessing as mp -from typing import Sequence, Dict, Optional, Tuple, List, Union, Any, Mapping +import os +from collections.abc import Mapping, Sequence +from typing import Any import numpy as np @@ -24,18 +27,18 @@ class Result: :param ci: When paired bootstrap test is applied, this represents the 95% confidence interval around the true mean score `sys_mean`. """ - def __init__(self, score: float, p_value: Optional[float] = None, - mean: Optional[float] = None, ci: Optional[float] = None): + def __init__(self, score: float, p_value: float | None = None, + mean: float | None = None, ci: float | None = None): self.score = score self.p_value = p_value self.mean = mean self.ci = ci def __repr__(self): - return ','.join([f'{k}={str(v)}' for k, v in self.__dict__.items()]) + return ','.join([f'{k}={v!s}' for k, v in self.__dict__.items()]) -def estimate_ci(scores: np.ndarray) -> Tuple[float, float]: +def estimate_ci(scores: np.ndarray) -> tuple[float, float]: """Takes a list of scores and returns mean and 95% confidence interval around the mean. @@ -54,8 +57,8 @@ def estimate_ci(scores: np.ndarray) -> Tuple[float, float]: return (scores.mean(), ci) -def _bootstrap_resample(stats: List[List[Union[int, float]]], - metric: Metric, n_samples: int = 1000) -> Tuple[str, List[Score]]: +def _bootstrap_resample(stats: list[list[int | float]], + metric: Metric, n_samples: int = 1000) -> tuple[str, list[Score]]: """Performs bootstrap resampling for a single system to estimate a confidence interval around the true mean. :param stats: A list of statistics extracted from the system's hypotheses. @@ -109,14 +112,14 @@ def _compute_p_value(stats: np.ndarray, real_difference: float) -> float: return p -def _paired_ar_test(baseline_info: Dict[str, Tuple[np.ndarray, Result]], +def _paired_ar_test(baseline_info: dict[str, tuple[np.ndarray, Result]], sys_name: str, hypotheses: Sequence[str], - references: Optional[Sequence[Sequence[str]]], - metrics: Dict[str, Metric], + references: Sequence[Sequence[str]] | None, + metrics: dict[str, Metric], n_samples: int = 10000, n_ar_confidence: int = -1, - seed: Optional[int] = None) -> Tuple[str, Dict[str, Result]]: + seed: int | None = None) -> tuple[str, dict[str, Result]]: """Paired two-sided approximate randomization (AR) test for MT evaluation. :param baseline_info: A dictionary with `Metric` instances as the keys, @@ -197,14 +200,14 @@ def _paired_ar_test(baseline_info: Dict[str, Tuple[np.ndarray, Result]], return sys_name, results -def _paired_bs_test(baseline_info: Dict[str, Tuple[np.ndarray, Result]], +def _paired_bs_test(baseline_info: dict[str, tuple[np.ndarray, Result]], sys_name: str, hypotheses: Sequence[str], - references: Optional[Sequence[Sequence[str]]], - metrics: Dict[str, Metric], + references: Sequence[Sequence[str]] | None, + metrics: dict[str, Metric], n_samples: int = 1000, n_ar_confidence: int = -1, - seed: Optional[int] = None) -> Tuple[str, Dict[str, Result]]: + seed: int | None = None) -> tuple[str, dict[str, Result]]: """Paired bootstrap resampling test for MT evaluation. This function replicates the behavior of the Moses script called `bootstrap-hypothesis-difference-significance.pl`. @@ -300,9 +303,9 @@ class PairedTest: 'bs': 1000, } - def __init__(self, named_systems: List[Tuple[str, Sequence[str]]], + def __init__(self, named_systems: list[tuple[str, Sequence[str]]], metrics: Mapping[str, Metric], - references: Optional[Sequence[Sequence[str]]], + references: Sequence[Sequence[str]] | None, test_type: str = 'ar', n_samples: int = 0, n_ar_confidence: int = -1, @@ -349,8 +352,8 @@ def __init__(self, named_systems: List[Tuple[str, Sequence[str]]], # Don't use more workers than the number of CPUs self.n_jobs = min(n_max_jobs, self.n_systems) - self._signatures: Dict[str, Signature] = {} - self._baseline_info: Dict[str, Tuple[Any, Result]] = {} + self._signatures: dict[str, Signature] = {} + self._baseline_info: dict[str, tuple[Any, Result]] = {} ################################################## # Pre-compute and cache baseline system statistics @@ -389,10 +392,10 @@ def __init__(self, named_systems: List[Tuple[str, Sequence[str]]], sig.update('bs', self.n_ar_confidence) self._signatures[bl_score.name] = sig - def __call__(self) -> Tuple[Dict[str, Signature], Dict[str, List[Union[str, Result]]]]: + def __call__(self) -> tuple[dict[str, Signature], dict[str, list[str | Result]]]: """Runs the paired test either on single or multiple worker processes.""" tasks = [] - scores: Dict[str, List[Union[str, Result]]] = {} + scores: dict[str, list[str | Result]] = {} # Add the name column scores['System'] = [ns[0] for ns in self.named_systems] diff --git a/sacrebleu/tokenizers/__init__.py b/sacrebleu/tokenizers/__init__.py index d658a1ba..0059ebbb 100644 --- a/sacrebleu/tokenizers/__init__.py +++ b/sacrebleu/tokenizers/__init__.py @@ -1,2 +1,2 @@ # Base tokenizer to derive from -from .tokenizer_base import BaseTokenizer # noqa: F401 +from .tokenizer_base import BaseTokenizer # noqa: F401 diff --git a/sacrebleu/tokenizers/tokenizer_13a.py b/sacrebleu/tokenizers/tokenizer_13a.py index 6441a762..f0d660e9 100644 --- a/sacrebleu/tokenizers/tokenizer_13a.py +++ b/sacrebleu/tokenizers/tokenizer_13a.py @@ -1,4 +1,5 @@ from functools import lru_cache + from .tokenizer_base import BaseTokenizer from .tokenizer_re import TokenizerRegexp diff --git a/sacrebleu/tokenizers/tokenizer_char.py b/sacrebleu/tokenizers/tokenizer_char.py index 8b8f8c5d..c70a4f79 100644 --- a/sacrebleu/tokenizers/tokenizer_char.py +++ b/sacrebleu/tokenizers/tokenizer_char.py @@ -1,4 +1,5 @@ from functools import lru_cache + from .tokenizer_base import BaseTokenizer @@ -16,4 +17,4 @@ def __call__(self, line): :param line: a segment to tokenize :return: the tokenized line """ - return " ".join((char for char in line)) + return " ".join(char for char in line) diff --git a/sacrebleu/tokenizers/tokenizer_ja_mecab.py b/sacrebleu/tokenizers/tokenizer_ja_mecab.py index 2844c5fc..65cced35 100644 --- a/sacrebleu/tokenizers/tokenizer_ja_mecab.py +++ b/sacrebleu/tokenizers/tokenizer_ja_mecab.py @@ -1,8 +1,8 @@ from functools import lru_cache try: - import MeCab import ipadic + import MeCab except ImportError: # Don't fail until the tokenizer is actually used MeCab = None diff --git a/sacrebleu/tokenizers/tokenizer_none.py b/sacrebleu/tokenizers/tokenizer_none.py index a204c000..9cc28232 100644 --- a/sacrebleu/tokenizers/tokenizer_none.py +++ b/sacrebleu/tokenizers/tokenizer_none.py @@ -1,5 +1,6 @@ from .tokenizer_base import BaseTokenizer + class NoneTokenizer(BaseTokenizer): """Don't apply any tokenization. Not recommended!.""" diff --git a/sacrebleu/tokenizers/tokenizer_re.py b/sacrebleu/tokenizers/tokenizer_re.py index 7eb67eb5..e73b7eb1 100644 --- a/sacrebleu/tokenizers/tokenizer_re.py +++ b/sacrebleu/tokenizers/tokenizer_re.py @@ -1,5 +1,5 @@ -from functools import lru_cache import re +from functools import lru_cache from .tokenizer_base import BaseTokenizer diff --git a/sacrebleu/tokenizers/tokenizer_spm.py b/sacrebleu/tokenizers/tokenizer_spm.py index 6cbc34d3..0b4356b6 100644 --- a/sacrebleu/tokenizers/tokenizer_spm.py +++ b/sacrebleu/tokenizers/tokenizer_spm.py @@ -1,9 +1,8 @@ -# -*- coding: utf-8 -*- -import os import logging - +import os from functools import lru_cache + from ..utils import SACREBLEU_DIR, download_file from .tokenizer_base import BaseTokenizer @@ -39,7 +38,7 @@ def __init__(self, key="spm"): self.name = SPM_MODELS[key]["signature"] if key == "spm": - sacrelogger.warn("Tokenizer 'spm' has been changed to 'flores101', and may be removed in the future.") + sacrelogger.warning("Tokenizer 'spm' has been changed to 'flores101', and may be removed in the future.") try: import sentencepiece as spm diff --git a/sacrebleu/tokenizers/tokenizer_zh.py b/sacrebleu/tokenizers/tokenizer_zh.py index 8ec831aa..e4666b31 100644 --- a/sacrebleu/tokenizers/tokenizer_zh.py +++ b/sacrebleu/tokenizers/tokenizer_zh.py @@ -43,29 +43,29 @@ from .tokenizer_re import TokenizerRegexp _UCODE_RANGES = [ - (u'\u3400', u'\u4db5'), # CJK Unified Ideographs Extension A, release 3.0 - (u'\u4e00', u'\u9fa5'), # CJK Unified Ideographs, release 1.1 - (u'\u9fa6', u'\u9fbb'), # CJK Unified Ideographs, release 4.1 - (u'\uf900', u'\ufa2d'), # CJK Compatibility Ideographs, release 1.1 - (u'\ufa30', u'\ufa6a'), # CJK Compatibility Ideographs, release 3.2 - (u'\ufa70', u'\ufad9'), # CJK Compatibility Ideographs, release 4.1 - (u'\u20000', u'\u2a6d6'), # (UTF16) CJK Unified Ideographs Extension B, release 3.1 - (u'\u2f800', u'\u2fa1d'), # (UTF16) CJK Compatibility Supplement, release 3.1 - (u'\uff00', u'\uffef'), # Full width ASCII, full width of English punctuation, + ('\u3400', '\u4db5'), # CJK Unified Ideographs Extension A, release 3.0 + ('\u4e00', '\u9fa5'), # CJK Unified Ideographs, release 1.1 + ('\u9fa6', '\u9fbb'), # CJK Unified Ideographs, release 4.1 + ('\uf900', '\ufa2d'), # CJK Compatibility Ideographs, release 1.1 + ('\ufa30', '\ufa6a'), # CJK Compatibility Ideographs, release 3.2 + ('\ufa70', '\ufad9'), # CJK Compatibility Ideographs, release 4.1 + ('\u20000', '\u2a6d6'), # (UTF16) CJK Unified Ideographs Extension B, release 3.1 + ('\u2f800', '\u2fa1d'), # (UTF16) CJK Compatibility Supplement, release 3.1 + ('\uff00', '\uffef'), # Full width ASCII, full width of English punctuation, # half width Katakana, half wide half width kana, Korean alphabet - (u'\u2e80', u'\u2eff'), # CJK Radicals Supplement - (u'\u3000', u'\u303f'), # CJK punctuation mark - (u'\u31c0', u'\u31ef'), # CJK stroke - (u'\u2f00', u'\u2fdf'), # Kangxi Radicals - (u'\u2ff0', u'\u2fff'), # Chinese character structure - (u'\u3100', u'\u312f'), # Phonetic symbols - (u'\u31a0', u'\u31bf'), # Phonetic symbols (Taiwanese and Hakka expansion) - (u'\ufe10', u'\ufe1f'), - (u'\ufe30', u'\ufe4f'), - (u'\u2600', u'\u26ff'), - (u'\u2700', u'\u27bf'), - (u'\u3200', u'\u32ff'), - (u'\u3300', u'\u33ff'), + ('\u2e80', '\u2eff'), # CJK Radicals Supplement + ('\u3000', '\u303f'), # CJK punctuation mark + ('\u31c0', '\u31ef'), # CJK stroke + ('\u2f00', '\u2fdf'), # Kangxi Radicals + ('\u2ff0', '\u2fff'), # Chinese character structure + ('\u3100', '\u312f'), # Phonetic symbols + ('\u31a0', '\u31bf'), # Phonetic symbols (Taiwanese and Hakka expansion) + ('\ufe10', '\ufe1f'), + ('\ufe30', '\ufe4f'), + ('\u2600', '\u26ff'), + ('\u2700', '\u27bf'), + ('\u3200', '\u32ff'), + ('\u3300', '\u33ff'), ] diff --git a/sacrebleu/utils.py b/sacrebleu/utils.py index 34d286c4..6aae33b1 100644 --- a/sacrebleu/utils.py +++ b/sacrebleu/utils.py @@ -1,20 +1,21 @@ +from __future__ import annotations + +import gzip +import hashlib import itertools import json +import logging +import math import os import re import sys -import gzip -import math -import hashlib -import logging -import portalocker -from collections import defaultdict -from typing import List, Optional, Sequence, Dict from argparse import Namespace +from collections import defaultdict +from collections.abc import Sequence -from tabulate import tabulate import colorama - +import portalocker +from tabulate import tabulate # Where to store downloaded test sets. # Define the environment variable $SACREBLEU, or use the default of ~/.sacrebleu. @@ -50,7 +51,7 @@ def format(msg: str, color: str) -> str: def _format_score_lines(scores: dict, width: int = 2, - multiline: bool = True) -> Dict[str, List[str]]: + multiline: bool = True) -> dict[str, list[str]]: """Formats the scores prior to tabulating them.""" new_scores = {'System': scores.pop('System')} p_val_break_char = '\n' if multiline else ' ' @@ -126,7 +127,7 @@ def print_results_table(results: dict, signatures: dict, args: Namespace): # Color the column names and the baseline system name and scores has_baseline = False baseline_name = '' - for name in results.keys(): + for name in results: val = results[name] if val[0].startswith('Baseline:') or has_baseline: if val[0].startswith('Baseline:'): @@ -188,7 +189,7 @@ def print_results_table(results: dict, signatures: dict, args: Namespace): print(f' - {name:<10} {sig}') -def print_single_results(results: List[str], args: Namespace): +def print_single_results(results: list[str], args: Namespace): """Re-process metric strings to align them nicely.""" if args.format == 'json': if len(results) > 1: @@ -228,7 +229,7 @@ def print_single_results(results: List[str], args: Namespace): def sanity_check_lengths(system: Sequence[str], refs: Sequence[Sequence[str]], - test_set: Optional[str] = None): + test_set: str | None = None): n_hyps = len(system) if any(len(ref_stream) != n_hyps for ref_stream in refs): sacrelogger.error("System and reference streams have different lengths.") @@ -333,7 +334,7 @@ def print_test_set(test_set, langpair, requested_fields, origlang=None, subset=N streams = [smart_open(file) for file in files] streams = filter_subset(streams, test_set, langpair, origlang, subset) for lines in zip(*streams): - print('\t'.join(map(lambda x: x.rstrip(), lines))) + print('\t'.join(x.rstrip() for x in lines)) def get_source_file(test_set: str, langpair: str) -> str: @@ -351,7 +352,7 @@ def get_source_file(test_set: str, langpair: str) -> str: return DATASETS[test_set].get_source_file(langpair) -def get_reference_files(test_set: str, langpair: str) -> List[str]: +def get_reference_files(test_set: str, langpair: str) -> list[str]: """ Returns a list of one or more reference file paths for the given testset/langpair. Downloads the references first if they are not already local. @@ -365,7 +366,7 @@ def get_reference_files(test_set: str, langpair: str) -> List[str]: return DATASETS[test_set].get_reference_files(langpair) -def get_files(test_set, langpair) -> List[str]: +def get_files(test_set, langpair) -> list[str]: """ Returns the path of the source file and all reference files for the provided test set / language pair. @@ -383,7 +384,7 @@ def get_files(test_set, langpair) -> List[str]: def extract_tarball(filepath, destdir): sacrelogger.info(f'Extracting {filepath} to {destdir}') - if filepath.endswith('.tar.gz') or filepath.endswith('.tgz'): + if filepath.endswith(('.tar.gz', '.tgz')): import tarfile with tarfile.open(filepath) as tar: tar.extractall(path=destdir) @@ -413,8 +414,8 @@ def download_file(source_path, dest_path, extract_to=None, expected_md5=None): :param expected_md5: the MD5 sum :return: the set of processed file names """ - import urllib.request import ssl + import urllib.request outdir = os.path.dirname(dest_path) os.makedirs(outdir, exist_ok=True) @@ -462,18 +463,18 @@ def download_test_set(test_set, langpair=None): return file_paths -def get_langpairs_for_testset(testset: str) -> List[str]: +def get_langpairs_for_testset(testset: str) -> list[str]: """Return a list of language pairs for a given test set.""" if testset not in DATASETS: return [] return list(DATASETS[testset].langpairs.keys()) -def get_available_testsets() -> List[str]: +def get_available_testsets() -> list[str]: """Return a list of available test sets.""" return sorted(DATASETS.keys(), reverse=True) -def get_available_testsets_for_langpair(langpair: str) -> List[str]: +def get_available_testsets_for_langpair(langpair: str) -> list[str]: """Return a list of available test sets for a given language pair""" parts = langpair.split('-') srclang = parts[0] @@ -488,7 +489,7 @@ def get_available_testsets_for_langpair(langpair: str) -> List[str]: return testsets -def get_available_origlangs(test_sets, langpair) -> List[str]: +def get_available_origlangs(test_sets, langpair) -> list[str]: """Return a list of origlang values according to the raw XML/SGM files.""" if test_sets is None: return [] @@ -507,10 +508,10 @@ def get_available_origlangs(test_sets, langpair) -> List[str]: if line.startswith(' List[str]: +def get_available_subsets(test_sets, langpair) -> list[str]: """Return a list of domain values according to the raw XML files and domain/country values from the SGM files.""" if test_sets is None: return [] @@ -525,9 +526,9 @@ def get_available_subsets(test_sets, langpair) -> List[str]: if 'domain' in fields: subsets |= set(fields['domain']) elif test_set in SUBSETS: - subsets |= set("country:" + v.split("-")[0] for v in SUBSETS[test_set].values()) - subsets |= set(v.split("-")[1] for v in SUBSETS[test_set].values()) - return sorted(list(subsets)) + subsets |= {"country:" + v.split("-")[0] for v in SUBSETS[test_set].values()} + subsets |= {v.split("-")[1] for v in SUBSETS[test_set].values()} + return sorted(subsets) def filter_subset(systems, test_sets, langpair, origlang, subset=None): """Filter sentences with a given origlang (or subset) according to the raw SGM files.""" @@ -628,12 +629,12 @@ def print_subset_results(metrics, full_system, full_refs, args): score = metric.corpus_score(system, refs) results[key].append((len(system), score)) - max_left_width = max([len(k) for k in results.keys()]) + 1 - max_metric_width = max([len(val[1].name) for val in list(results.values())[0]]) + max_left_width = max([len(k) for k in results]) + 1 + max_metric_width = max([len(val[1].name) for val in next(iter(results.values()))]) for key, scores in results.items(): key = Color.format(f'{key:<{max_left_width}}', 'yellow') for n_system, score in scores: print(f'{key}: sentences={n_system:<6} {score.name:<{max_metric_width}} = {score.score:.{w}f}') # import at the end to avoid circular import -from .dataset import DATASETS, SUBSETS # noqa: E402 +from .dataset import DATASETS, SUBSETS diff --git a/scripts/perf_test.py b/scripts/perf_test.py old mode 100644 new mode 100755 index f2812db5..b3302a81 --- a/scripts/perf_test.py +++ b/scripts/perf_test.py @@ -1,13 +1,12 @@ #!/usr/bin/env python +import statistics import sys import time -import statistics sys.path.insert(0, '.') -import sacrebleu # noqa: E402 -from sacrebleu.metrics import BLEU, CHRF # noqa: E402 - +import sacrebleu +from sacrebleu.metrics import BLEU, CHRF N_REPEATS = 5 diff --git a/test/test_api.py b/test/test_api.py index 02ac9f03..f53de228 100644 --- a/test/test_api.py +++ b/test/test_api.py @@ -13,9 +13,14 @@ import pytest -from sacrebleu.utils import get_available_testsets, get_available_testsets_for_langpair, get_langpairs_for_testset -from sacrebleu.utils import get_source_file, get_reference_files from sacrebleu.dataset import DATASETS +from sacrebleu.utils import ( + get_available_testsets, + get_available_testsets_for_langpair, + get_langpairs_for_testset, + get_reference_files, + get_source_file, +) test_api_get_data = [ ("wmt19", "de-en", 1, "Schöne Münchnerin 2018: Schöne Münchnerin 2018 in Hvar: Neun Dates", "The Beauty of Munich 2018: the Beauty of Munich 2018 in Hvar: Nine dates"), @@ -48,7 +53,7 @@ def test_api_get_available_testsets(): assert "wmt19" in available assert "wmt05" not in available - for testset in DATASETS.keys(): + for testset in DATASETS: assert testset in available assert "slashdot_" + testset not in available @@ -75,10 +80,10 @@ def test_api_get_langpairs_for_testset(): Loop over the datasets directly, and ensure the API function returns each language pair in each test set. """ - for testset in DATASETS.keys(): + for testset in DATASETS: available = get_langpairs_for_testset(testset) assert isinstance(available, list) - for langpair in DATASETS[testset].langpairs.keys(): + for langpair in DATASETS[testset].langpairs: # skip non-language keys if "-" not in langpair: assert langpair not in available diff --git a/test/test_bleu.py b/test/test_bleu.py index d8312de6..20936bf0 100644 --- a/test/test_bleu.py +++ b/test/test_bleu.py @@ -12,13 +12,12 @@ # permissions and limitations under the License. from collections import namedtuple + import pytest import sacrebleu - from sacrebleu.metrics import BLEU - EPSILON = 1e-8 Statistics = namedtuple('Statistics', ['common', 'total']) diff --git a/test/test_chrf.py b/test/test_chrf.py index 24df6d52..56b6e7c9 100644 --- a/test/test_chrf.py +++ b/test/test_chrf.py @@ -12,6 +12,7 @@ # permissions and limitations under the License. import pytest + import sacrebleu EPSILON = 1e-4 diff --git a/test/test_dataset.py b/test/test_dataset.py index d3ece8cb..0c06aad6 100644 --- a/test/test_dataset.py +++ b/test/test_dataset.py @@ -1,8 +1,8 @@ import os -import shutil import random +import shutil -import sacrebleu.dataset as dataset +from sacrebleu import dataset from sacrebleu.utils import smart_open diff --git a/test/test_sentence_bleu.py b/test/test_sentence_bleu.py index 88afa2eb..6cb46c02 100644 --- a/test/test_sentence_bleu.py +++ b/test/test_sentence_bleu.py @@ -1,4 +1,5 @@ import pytest + import sacrebleu EPSILON = 1e-3 diff --git a/test/test_significance.py b/test/test_significance.py index f7098328..edee928e 100644 --- a/test/test_significance.py +++ b/test/test_significance.py @@ -1,13 +1,11 @@ import os - from collections import defaultdict -from typing import DefaultDict + +import pytest from sacrebleu.metrics import BLEU from sacrebleu.significance import PairedTest, Result -import pytest - def _read_pickle_file(): import bz2 @@ -58,8 +56,8 @@ def _read_pickle_file(): } -SACREBLEU_BS_P_VALS: DefaultDict[str, float] = defaultdict(float) -SACREBLEU_AR_P_VALS: DefaultDict[str, float] = defaultdict(float) +SACREBLEU_BS_P_VALS: defaultdict[str, float] = defaultdict(float) +SACREBLEU_AR_P_VALS: defaultdict[str, float] = defaultdict(float) # Load data from pickled file to not bother with WMT17 downloading named_systems = _read_pickle_file() diff --git a/test/test_ter.py b/test/test_ter.py index df650a93..b32211d3 100644 --- a/test/test_ter.py +++ b/test/test_ter.py @@ -1,4 +1,5 @@ import pytest + import sacrebleu EPSILON = 1e-3 diff --git a/test/test_tokenizer_ter.py b/test/test_tokenizer_ter.py index e76a03bf..143f0aae 100644 --- a/test/test_tokenizer_ter.py +++ b/test/test_tokenizer_ter.py @@ -2,7 +2,6 @@ from sacrebleu.tokenizers.tokenizer_ter import TercomTokenizer - test_cases_default = [ ("a b c d", "a b c d"), ("", ""),