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
4 changes: 2 additions & 2 deletions mbrs/decoders/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

from mbrs import registry
from mbrs.metrics.base import Metric, MetricBase, MetricReferenceless
from mbrs.selectors import Selector, SelectorNbest
from mbrs.selectors import SELECTOR_NBEST, Selector


class DecoderBase(abc.ABC):
Expand All @@ -18,7 +18,7 @@ def __init__(
self,
cfg: DecoderBase.Config,
metric: MetricBase,
selector: Selector = SelectorNbest(SelectorNbest.Config()),
selector: Selector = SELECTOR_NBEST,
) -> None:
self.cfg = cfg
self.metric = metric
Expand Down
4 changes: 2 additions & 2 deletions mbrs/decoders/centroid_mbr.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
from mbrs import functional, timer
from mbrs.metrics import MetricAggregatableCache
from mbrs.modules.kmeans import Kmeans
from mbrs.selectors import Selector, SelectorNbest
from mbrs.selectors import SELECTOR_NBEST, Selector

from . import register
from .mbr import DecoderMBR
Expand All @@ -34,7 +34,7 @@ def __init__(
self,
cfg: DecoderCentroidMBR.Config,
metric: MetricAggregatableCache,
selector: Selector = SelectorNbest(SelectorNbest.Config()),
selector: Selector = SELECTOR_NBEST,
) -> None:
super().__init__(cfg, metric, selector=selector)
self.kmeans = Kmeans(cfg.kmeans)
Expand Down
4 changes: 2 additions & 2 deletions mbrs/decoders/pruning_mbr.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@

from mbrs import functional, timer
from mbrs.metrics import Metric, MetricCacheable
from mbrs.selectors import Selector, SelectorNbest
from mbrs.selectors import SELECTOR_NBEST, Selector, SelectorNbest

from . import register
from .mbr import DecoderMBR
Expand All @@ -28,7 +28,7 @@ def __init__(
self,
cfg: DecoderPruningMBR.Config,
metric: Metric,
selector: Selector = SelectorNbest(SelectorNbest.Config()),
selector: Selector = SELECTOR_NBEST,
) -> None:
if not isinstance(selector, SelectorNbest):
raise ValueError(
Expand Down
5 changes: 3 additions & 2 deletions mbrs/metrics/bleurt.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@ class Config(Metric.Config):
cpu: bool = False

def __init__(self, cfg: MetricBLEURT.Config):
self.cfg = cfg
super().__init__(cfg)
self.scorer = BleurtForSequenceClassification.from_pretrained(cfg.model)
self.tokenizer = BleurtTokenizer.from_pretrained(cfg.model)
self.max_length = self.tokenizer.max_model_input_sizes[
Expand Down Expand Up @@ -197,7 +197,8 @@ def pairwise_scores(
return torch.cat(scores).view(len(references), len(hypotheses)).transpose(0, 1)

def corpus_score(
self, hypotheses: list[str],
self,
hypotheses: list[str],
references_lists: list[list[str]],
sources: Optional[list[str]] = None,
) -> float:
Expand Down
2 changes: 1 addition & 1 deletion mbrs/metrics/comet.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@ class Config(MetricAggregatableCache.Config):
cpu: bool = False

def __init__(self, cfg: MetricCOMET.Config):
self.cfg = cfg
super().__init__(cfg)
self.scorer = load_from_checkpoint(download_model(cfg.model))
self.scorer.eval()
for param in self.scorer.parameters():
Expand Down
2 changes: 1 addition & 1 deletion mbrs/metrics/cometkiwi.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ class Config(MetricReferenceless.Config):
cpu: bool = False

def __init__(self, cfg: MetricCOMETkiwi.Config):
self.cfg = cfg
super().__init__(cfg)
self.scorer = load_from_checkpoint(download_model(cfg.model))
self.scorer.eval()
for param in self.scorer.parameters():
Expand Down
2 changes: 1 addition & 1 deletion mbrs/metrics/metricx.py
Original file line number Diff line number Diff line change
Expand Up @@ -272,7 +272,7 @@ class InputPrefix:
}

def __init__(self, cfg: MetricMetricX.Config):
self.cfg = cfg
super().__init__(cfg)
self.scorer = MT5ForRegression.from_pretrained(cfg.model)
os.environ["PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION"] = "python"
self.tokenizer = AutoTokenizer.from_pretrained(
Expand Down
2 changes: 1 addition & 1 deletion mbrs/metrics/xcomet.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,7 +166,7 @@ class Config(Metric.Config):
cpu: bool = False

def __init__(self, cfg: MetricXCOMET.Config):
self.cfg = cfg
super().__init__(cfg)
if cfg.model == "myyycroft/XCOMET-lite":
self.scorer = XCOMETLiteMetric.from_pretrained(cfg.model)
else:
Expand Down
4 changes: 4 additions & 0 deletions mbrs/selectors/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,10 +5,14 @@
from .diverse import SelectorDiverse
from .nbest import SelectorNbest

# Singleton of default selector
SELECTOR_NBEST = SelectorNbest(SelectorNbest.Config())

__all__ = [
"Selector",
"SelectorNbest",
"SelectorDiverse",
"register",
"get_selector",
"SELECTOR_NBEST",
]