2727import logging
2828from collections import Counter
2929from multiprocessing import Pool
30+ from typing import Literal
3031
3132from lighteval .metrics .imports .data_stats_utils import Fragments
3233from lighteval .utils .imports import NO_SPACY_ERROR_MSG , is_spacy_available
3334
3435
3536logger = logging .getLogger (__name__ )
3637
37-
38- _en = None
38+ LANGUAGE_TO_SPACY_MODEL_MAP = {
39+ "en" : "en_core_web_sm" ,
40+ "de" : "de_core_news_sm" ,
41+ "fr" : "fr_core_news_sm" ,
42+ "it" : "it_core_news_sm" ,
43+ }
3944
4045
4146class Metric :
@@ -51,8 +56,16 @@ def find_ngrams(input_list, n):
5156
5257
5358class DataStatsMetric (Metric ):
54- def __init__ (self , n_gram = 3 , n_workers = 24 , case = False , tokenize = True ):
55- """Data Statistics metric
59+ def __init__ (
60+ self ,
61+ n_gram : int = 3 ,
62+ n_workers : int = 24 ,
63+ case : bool = False ,
64+ tokenize : bool = True ,
65+ language : Literal ["en" , "de" , "fr" , "it" ] = "en" ,
66+ ):
67+ """
68+ Data Statistics metric
5669 Makes use of Newsroom code: \
5770 https://github.com/lil-lab/newsroom/blob/master/newsroom/analyze/fragments.py
5871 Calculates extractive statistics such as coverage, density, compression as
@@ -69,6 +82,9 @@ def __init__(self, n_gram=3, n_workers=24, case=False, tokenize=True):
6982 case (bool): whether to lowercase input before calculating statistics.
7083 tokenize (bool): whether to tokenize the input; otherwise assumes that the input
7184 is a string of space-separated tokens.
85+ language (Literal["en", "de", "fr", "it"]): the language of the input text. This
86+ determines the spaCy model used for tokenization. Currently supports English,
87+ German, French, and Italian.
7288 """
7389 if not is_spacy_available ():
7490 raise ImportError (NO_SPACY_ERROR_MSG )
@@ -78,22 +94,24 @@ def __init__(self, n_gram=3, n_workers=24, case=False, tokenize=True):
7894 self .n_workers = n_workers
7995 self .case = case
8096 self .tokenize = tokenize
97+ self .language = language
98+ self .nlp = None
8199
82- global _en
100+ spacy_model = LANGUAGE_TO_SPACY_MODEL_MAP . get ( self . language , "en_core_web_sm" )
83101 try :
84- _en = spacy .load ("en_core_web_sm" )
102+ self . nlp = spacy .load (spacy_model )
85103 except OSError :
86- logger .info ("Downloading the spacy en_core_web_sm model\n (don't worry, this will only happen once)" )
104+ logger .info (f "Downloading the spacy { spacy_model } model\n (don't worry, this will only happen once)" )
87105 from spacy .cli import download
88106
89- download ("en_core_web_sm" )
90- _en = spacy .load ("en_core_web_sm" )
107+ download (spacy_model )
108+ self . nlp = spacy .load (spacy_model )
91109
92110 def evaluate_example (self , summary , input_text ):
93111 if self .tokenize :
94- input_text = _en (input_text , disable = ["tagger" , "parser" , "ner" , "textcat" ])
112+ input_text = self . nlp (input_text , disable = ["tagger" , "parser" , "ner" , "textcat" ])
95113 input_text = [tok .text for tok in input_text ]
96- summary = _en (summary , disable = ["tagger" , "parser" , "ner" , "textcat" ])
114+ summary = self . nlp (summary , disable = ["tagger" , "parser" , "ner" , "textcat" ])
97115 summary = [tok .text for tok in summary ]
98116 fragments = Fragments (summary , input_text , case = self .case )
99117 coverage = fragments .coverage ()
0 commit comments