File tree Expand file tree Collapse file tree
Expand file tree Collapse file tree Original file line number Diff line number Diff line change 88
99from mbrs import registry
1010from mbrs .metrics .base import Metric , MetricBase , MetricReferenceless
11- from mbrs .selectors import Selector , SelectorNbest
11+ from mbrs .selectors import SELECTOR_NBEST , Selector
1212
1313
1414class DecoderBase (abc .ABC ):
@@ -18,7 +18,7 @@ def __init__(
1818 self ,
1919 cfg : DecoderBase .Config ,
2020 metric : MetricBase ,
21- selector : Selector = SelectorNbest ( SelectorNbest . Config ()) ,
21+ selector : Selector = SELECTOR_NBEST ,
2222 ) -> None :
2323 self .cfg = cfg
2424 self .metric = metric
Original file line number Diff line number Diff line change 99from mbrs import functional , timer
1010from mbrs .metrics import MetricAggregatableCache
1111from mbrs .modules .kmeans import Kmeans
12- from mbrs .selectors import Selector , SelectorNbest
12+ from mbrs .selectors import SELECTOR_NBEST , Selector
1313
1414from . import register
1515from .mbr import DecoderMBR
@@ -34,7 +34,7 @@ def __init__(
3434 self ,
3535 cfg : DecoderCentroidMBR .Config ,
3636 metric : MetricAggregatableCache ,
37- selector : Selector = SelectorNbest ( SelectorNbest . Config ()) ,
37+ selector : Selector = SELECTOR_NBEST ,
3838 ) -> None :
3939 super ().__init__ (cfg , metric , selector = selector )
4040 self .kmeans = Kmeans (cfg .kmeans )
Original file line number Diff line number Diff line change 88
99from mbrs import functional , timer
1010from mbrs .metrics import Metric , MetricCacheable
11- from mbrs .selectors import Selector , SelectorNbest
11+ from mbrs .selectors import SELECTOR_NBEST , Selector , SelectorNbest
1212
1313from . import register
1414from .mbr import DecoderMBR
@@ -28,7 +28,7 @@ def __init__(
2828 self ,
2929 cfg : DecoderPruningMBR .Config ,
3030 metric : Metric ,
31- selector : Selector = SelectorNbest ( SelectorNbest . Config ()) ,
31+ selector : Selector = SELECTOR_NBEST ,
3232 ) -> None :
3333 if not isinstance (selector , SelectorNbest ):
3434 raise ValueError (
Original file line number Diff line number Diff line change @@ -62,7 +62,7 @@ class Config(Metric.Config):
6262 cpu : bool = False
6363
6464 def __init__ (self , cfg : MetricBLEURT .Config ):
65- self . cfg = cfg
65+ super (). __init__ ( cfg )
6666 self .scorer = BleurtForSequenceClassification .from_pretrained (cfg .model )
6767 self .tokenizer = BleurtTokenizer .from_pretrained (cfg .model )
6868 self .max_length = self .tokenizer .max_model_input_sizes [
@@ -197,7 +197,8 @@ def pairwise_scores(
197197 return torch .cat (scores ).view (len (references ), len (hypotheses )).transpose (0 , 1 )
198198
199199 def corpus_score (
200- self , hypotheses : list [str ],
200+ self ,
201+ hypotheses : list [str ],
201202 references_lists : list [list [str ]],
202203 sources : Optional [list [str ]] = None ,
203204 ) -> float :
Original file line number Diff line number Diff line change @@ -34,7 +34,7 @@ class Config(MetricAggregatableCache.Config):
3434 cpu : bool = False
3535
3636 def __init__ (self , cfg : MetricCOMET .Config ):
37- self . cfg = cfg
37+ super (). __init__ ( cfg )
3838 self .scorer = load_from_checkpoint (download_model (cfg .model ))
3939 self .scorer .eval ()
4040 for param in self .scorer .parameters ():
Original file line number Diff line number Diff line change @@ -32,7 +32,7 @@ class Config(MetricReferenceless.Config):
3232 cpu : bool = False
3333
3434 def __init__ (self , cfg : MetricCOMETkiwi .Config ):
35- self . cfg = cfg
35+ super (). __init__ ( cfg )
3636 self .scorer = load_from_checkpoint (download_model (cfg .model ))
3737 self .scorer .eval ()
3838 for param in self .scorer .parameters ():
Original file line number Diff line number Diff line change @@ -272,7 +272,7 @@ class InputPrefix:
272272 }
273273
274274 def __init__ (self , cfg : MetricMetricX .Config ):
275- self . cfg = cfg
275+ super (). __init__ ( cfg )
276276 self .scorer = MT5ForRegression .from_pretrained (cfg .model )
277277 os .environ ["PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION" ] = "python"
278278 self .tokenizer = AutoTokenizer .from_pretrained (
Original file line number Diff line number Diff line change @@ -166,7 +166,7 @@ class Config(Metric.Config):
166166 cpu : bool = False
167167
168168 def __init__ (self , cfg : MetricXCOMET .Config ):
169- self . cfg = cfg
169+ super (). __init__ ( cfg )
170170 if cfg .model == "myyycroft/XCOMET-lite" :
171171 self .scorer = XCOMETLiteMetric .from_pretrained (cfg .model )
172172 else :
Original file line number Diff line number Diff line change 55from .diverse import SelectorDiverse
66from .nbest import SelectorNbest
77
8+ # Singleton of default selector
9+ SELECTOR_NBEST = SelectorNbest (SelectorNbest .Config ())
10+
811__all__ = [
912 "Selector" ,
1013 "SelectorNbest" ,
1114 "SelectorDiverse" ,
1215 "register" ,
1316 "get_selector" ,
17+ "SELECTOR_NBEST" ,
1418]
You can’t perform that action at this time.
0 commit comments