Skip to content

Commit 0dce0f2

Browse files
authored
Merge pull request #38 from naist-nlp/typed-registry
Improve type annotation of registry with generics
2 parents ead6663 + c4b31ce commit 0dce0f2

14 files changed

Lines changed: 174 additions & 80 deletions

File tree

.github/workflows/ci.yaml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -34,8 +34,8 @@ jobs:
3434
- name: Set up Python ${{ matrix.python-version }}
3535
run: uv python install
3636
- name: Install the project
37-
run: uv sync --all-extras --dev
37+
run: uv sync --all-extras --dev --index torch=https://download.pytorch.org/whl/cpu --index-strategy unsafe-best-match
3838
- name: Test with pytest
3939
run: |
40-
uv run huggingface-cli login --token ${{ secrets.HUGGINGFACE_TOKEN }}
40+
uv run hf auth login --token ${{ secrets.HUGGINGFACE_TOKEN }}
4141
uv run pytest ${{ matrix.pytest_marker && format('-m {0}', matrix.pytest_marker) || '' }}

mbrs/args_test.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -27,7 +27,7 @@ def test_plugin_load(self, tmp_path: pathlib.Path):
2727
f.writelines(["tests", "a test"])
2828

2929
cmd_args = ["--config_path", str(config_path)]
30-
with pytest.raises(NotImplementedError):
30+
with pytest.raises(KeyError):
3131
parser = get_argparser(cmd_args)
3232

3333
plugin_dir = os.path.join(

mbrs/cli/decode.py

Lines changed: 13 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -18,10 +18,8 @@
1818
)
1919
logger = logging.getLogger(__name__)
2020

21-
import simple_parsing
2221
import torch
2322
from simple_parsing import choice, field, flag
24-
from simple_parsing.wrappers import dataclass_wrapper
2523
from tabulate import tabulate, tabulate_formats
2624
from tqdm import tqdm
2725

@@ -33,7 +31,7 @@
3331
DecoderReferenceless,
3432
get_decoder,
3533
)
36-
from mbrs.metrics import Metric, MetricEnum, get_metric
34+
from mbrs.metrics import Metric, MetricEnum, MetricReferenceless, get_metric
3735
from mbrs.selectors import Selector, get_selector
3836

3937

@@ -64,15 +62,21 @@ class CommonArguments:
6462
num_references: int | None = field(default=None)
6563
# Type of the decoder.
6664
decoder: str = field(
67-
default="mbr", metadata={"choices": registry.get_registry("decoder")}
65+
default="mbr",
66+
metadata={
67+
"choices": registry.get_registry(
68+
DecoderReferenceBased | DecoderReferenceless
69+
)
70+
},
6871
)
6972
# Type of the metric.
7073
metric: str = field(
71-
default="bleu", metadata={"choices": registry.get_registry("metric")}
74+
default="bleu",
75+
metadata={"choices": registry.get_registry(Metric | MetricReferenceless)},
7276
)
7377
# Type of the selector.
7478
selector: str = field(
75-
default="nbest", metadata={"choices": registry.get_registry("selector")}
79+
default="nbest", metadata={"choices": registry.get_registry(Selector)}
7680
)
7781
# Return the n-best hypotheses.
7882
nbest: int = field(default=1)
@@ -177,11 +181,9 @@ def main(args: Namespace) -> None:
177181
reference_lprobs = f.readlines()
178182
assert len(references) == len(reference_lprobs)
179183

180-
metric: Metric = get_metric(args.common.metric)(args.metric)
181-
selector: Selector = get_selector(args.common.selector)(args.selector)
182-
decoder: DecoderReferenceBased | DecoderReferenceless = get_decoder(
183-
args.common.decoder
184-
)(args.decoder, metric, selector)
184+
metric = get_metric(args.common.metric)(args.metric)
185+
selector = get_selector(args.common.selector)(args.selector)
186+
decoder = get_decoder(args.common.decoder)(args.decoder, metric, selector)
185187

186188
num_cands = args.common.num_candidates
187189
num_refs = args.common.num_references or num_cands

mbrs/cli/score.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -45,7 +45,8 @@ class CommonArguments:
4545
format: Format = choice(Format, default=Format.json)
4646
# Type of the metric.
4747
metric: str = field(
48-
default="bleu", metadata={"choices": registry.get_registry("metric")}
48+
default="bleu",
49+
metadata={"choices": registry.get_registry(Metric | MetricReferenceless)},
4950
)
5051
# No verbose information and report.
5152
quiet: bool = flag(default=False)
@@ -94,7 +95,7 @@ def main(args: Namespace) -> None:
9495
assert num_sents == len(references)
9596
references_lists.append(references)
9697

97-
metric: Metric | MetricReferenceless = get_metric(args.common.metric)(args.metric)
98+
metric = get_metric(args.common.metric)(args.metric)
9899

99100
if isinstance(metric, MetricReferenceless):
100101
assert sources is not None

mbrs/decoders/__init__.py

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -2,9 +2,14 @@
22

33
from mbrs import registry
44

5-
from .base import DecoderBase, DecoderReferenceBased, DecoderReferenceless
5+
from .base import (
6+
DecoderBase,
7+
DecoderReferenceBased,
8+
DecoderReferenceless,
9+
register,
10+
get_decoder,
11+
)
612

7-
register, get_decoder = registry.setup("decoder")
813

914
from .aggregate_mbr import DecoderAggregateMBR
1015
from .centroid_mbr import DecoderCentroidMBR
@@ -17,12 +22,12 @@
1722
"DecoderBase",
1823
"DecoderReferenceBased",
1924
"DecoderReferenceless",
25+
"register",
26+
"get_decoder",
2027
"DecoderMBR",
2128
"DecoderAggregateMBR",
2229
"DecoderCentroidMBR",
2330
"DecoderProbabilisticMBR",
2431
"DecoderPruningMBR",
2532
"DecoderRerank",
2633
]
27-
28-
Decoders = enum.Enum("Decoders", registry.get_registry("decoder"))

mbrs/decoders/base.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66

77
from torch import Tensor
88

9+
from mbrs import registry
910
from mbrs.metrics.base import Metric, MetricBase, MetricReferenceless
1011
from mbrs.selectors import Selector, SelectorNbest
1112

@@ -184,3 +185,8 @@ def decode(
184185
Returns:
185186
Decoder.Output: The n-best hypotheses.
186187
"""
188+
189+
190+
register, get_decoder = registry.Registry(
191+
DecoderReferenceBased | DecoderReferenceless
192+
).get_closure()

mbrs/metrics/__init__.py

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -11,10 +11,9 @@
1111
MetricBase,
1212
MetricCacheable,
1313
MetricReferenceless,
14+
get_metric,
15+
register,
1416
)
15-
16-
register, get_metric = registry.setup("metric")
17-
1817
from .bertscore import MetricBERTScore
1918
from .bleu import MetricBLEU
2019
from .bleurt import MetricBLEURT
@@ -32,6 +31,8 @@
3231
"MetricAggregatableCache",
3332
"MetricCacheable",
3433
"MetricReferenceless",
34+
"get_metric",
35+
"register",
3536
"MetricBERTScore",
3637
"MetricBLEU",
3738
"MetricChrF",
@@ -47,4 +48,7 @@
4748
class MetricEnum(str, enum.Enum): ...
4849

4950

50-
Metrics = MetricEnum("Metrics", {k: k for k in registry.get_registry("metric").keys()})
51+
Metrics = MetricEnum(
52+
"Metrics",
53+
{k: k for k in registry.get_registry(Metric | MetricReferenceless).keys()},
54+
)

mbrs/metrics/base.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
import torch
88
from torch import Tensor
99

10-
from mbrs import functional, timer
10+
from mbrs import functional, registry, timer
1111
from mbrs.modules.kmeans import Kmeans
1212

1313

@@ -489,3 +489,6 @@ def corpus_score(self, hypotheses: list[str], sources: list[str]) -> float:
489489
float: The corpus score.
490490
"""
491491
return self.scores(hypotheses, sources=sources).mean().cpu().float().item()
492+
493+
494+
register, get_metric = registry.Registry(Metric | MetricReferenceless).get_closure()

mbrs/metrics/metricx.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -55,7 +55,7 @@ def __init__(self, config: MT5Config):
5555

5656
decoder_config = copy.deepcopy(config)
5757
decoder_config.is_decoder = True
58-
decoder_config.is_encoder_decoder = False
58+
decoder_config.is_encoder_decoder = True
5959
decoder_config.num_layers = config.num_decoder_layers
6060
self.decoder = MT5Stack(decoder_config, self.shared)
6161

mbrs/registry.py

Lines changed: 51 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -1,58 +1,76 @@
1-
from typing import Any, Callable, Dict, Type, TypeVar
1+
from typing import Callable, TypeVar
22

33
T = TypeVar("T")
44

5-
REGISTRIES = {}
65

6+
class Registry(dict[str, type[T]]):
7+
"""Registry that maps a name to its corresponding type."""
78

8-
def setup(registry_name: str):
9-
"""Setup a registry.
9+
def __init__(self, base_type: type[T]):
10+
super().__init__()
11+
REGISTRIES[base_type] = self
12+
self._base_type = base_type
1013

11-
Args:
12-
registry_name (str): Registry name for grouping classes.
14+
def register(self, name: str) -> Callable[[type[T]], type[T]]:
15+
"""Register a type as the given name.
1316
14-
Returns:
15-
Tuple of the two functions:
16-
- register: Register a class as the given name.
17-
- get_cls: Return the registered class of the given name.
18-
"""
19-
REGISTRY = {}
20-
REGISTRIES[registry_name] = REGISTRY
17+
Args:
18+
name (str): The name of a type.
2119
22-
def register(name: str) -> Callable[[Type[T]], Type[T]]:
23-
"""Register a class as the given name.
20+
Returns:
21+
Callable[[type[T]], type[T]]: Register decorator function.
2422
25-
Args:
26-
name (str): The name of a class.
23+
Raises:
24+
ValueError: The type is already registered.
2725
"""
2826

29-
def _register(cls: Type[T]):
30-
if name in REGISTRY:
27+
def _register(cls: type[T]) -> type[T]:
28+
if not issubclass(cls, self._base_type):
29+
raise ValueError(f"`{cls.__name__}` must inherit `{self._base_type}`.")
30+
31+
if (registered := self.get(name)) is not None:
3132
raise ValueError(
32-
f"{name} already registered as {REGISTRY[name].__name__}. ({cls.__name__})"
33+
f"{cls.__name__}: `{name}` already registered as `{registered.__name__}`."
3334
)
34-
REGISTRY[name] = cls
35+
self[name] = cls
3536
return cls
3637

3738
return _register
3839

39-
def get_cls(name: str):
40-
if name not in REGISTRY:
41-
raise NotImplementedError(
42-
f"`{name}` is not registered in `{registry_name}`."
43-
)
44-
return REGISTRY[name]
40+
def get_cls(self, name: str) -> type[T]:
41+
"""Get a class type.
42+
43+
Args:
44+
name: A registered name.
45+
46+
Returns:
47+
type[T]: Class type.
48+
"""
49+
return self.__getitem__(name)
50+
51+
def get_closure(
52+
self,
53+
) -> tuple[Callable[[str], Callable[[type[T]], type[T]]], Callable[[str], type[T]]]:
54+
"""Get closure functions: `register()` and `get_cls()`.
55+
56+
Returns:
57+
tuple:
58+
- Callable[[str], Callable[[type[T]], type[T]]]: `register()` function.
59+
- Callable[[str], type[T]]: `get_cls()` function.
60+
"""
61+
return (self.register, self.get_cls)
62+
4563

46-
return register, get_cls
64+
REGISTRIES: dict[type, Registry] = {}
4765

4866

49-
def get_registry(registry_name: str) -> Dict[str, Type[Any]]:
50-
"""Get registry of the given name.
67+
def get_registry(base_type: type[T]) -> Registry[T]:
68+
"""Get registry of the given base class type.
5169
5270
Args:
53-
registry_name (str): Registry name to be returned.
71+
base_type (type[T]): Base class type that associated with the registry to be returned.
5472
5573
Returns:
56-
Dict[str, Type[Any]]: Class mapper from registered name to its corresponding class.
74+
Registry[T]: Class mapper from registered name to its corresponding class.
5775
"""
58-
return REGISTRIES[registry_name]
76+
return REGISTRIES[base_type]

0 commit comments

Comments
 (0)