Skip to content

Commit 2c22664

Browse files
authored
Merge pull request #39 from naist-nlp/fix-constructor
Fix constructor
2 parents 0dce0f2 + 377dc6a commit 2c22664

9 files changed

Lines changed: 17 additions & 12 deletions

File tree

mbrs/decoders/base.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88

99
from mbrs import registry
1010
from mbrs.metrics.base import Metric, MetricBase, MetricReferenceless
11-
from mbrs.selectors import Selector, SelectorNbest
11+
from mbrs.selectors import SELECTOR_NBEST, Selector
1212

1313

1414
class 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

mbrs/decoders/centroid_mbr.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
from mbrs import functional, timer
1010
from mbrs.metrics import MetricAggregatableCache
1111
from mbrs.modules.kmeans import Kmeans
12-
from mbrs.selectors import Selector, SelectorNbest
12+
from mbrs.selectors import SELECTOR_NBEST, Selector
1313

1414
from . import register
1515
from .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)

mbrs/decoders/pruning_mbr.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88

99
from mbrs import functional, timer
1010
from mbrs.metrics import Metric, MetricCacheable
11-
from mbrs.selectors import Selector, SelectorNbest
11+
from mbrs.selectors import SELECTOR_NBEST, Selector, SelectorNbest
1212

1313
from . import register
1414
from .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(

mbrs/metrics/bleurt.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff 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:

mbrs/metrics/comet.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff 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():

mbrs/metrics/cometkiwi.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff 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():

mbrs/metrics/metricx.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff 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(

mbrs/metrics/xcomet.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff 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:

mbrs/selectors/__init__.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,10 +5,14 @@
55
from .diverse import SelectorDiverse
66
from .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
]

0 commit comments

Comments
 (0)