File size: 6,771 Bytes
a2d6c00
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
"""Tokenizer wrappers and BPE training helpers."""

from __future__ import annotations

from pathlib import Path
from typing import Iterable, Mapping, Sequence


class TokenizerWrapper:
    """Small adapter that gives HF tokenizers and tokenizers.Tokenizer one API."""

    def __init__(
        self,
        tokenizer,
        pad_token: str = "<pad>",
        unk_token: str = "<unk>",
        bos_token: str = "<s>",
        eos_token: str = "</s>",
    ):
        self.tokenizer = tokenizer
        self.pad_token = pad_token
        self.unk_token = unk_token
        self.bos_token = bos_token
        self.eos_token = eos_token

    @property
    def pad_token_id(self) -> int:
        return self.token_to_id(self.pad_token)

    @property
    def unk_token_id(self) -> int:
        return self.token_to_id(self.unk_token)

    @property
    def bos_token_id(self) -> int:
        return self.token_to_id(self.bos_token)

    @property
    def eos_token_id(self) -> int:
        return self.token_to_id(self.eos_token)

    @property
    def vocab_size(self) -> int:
        if hasattr(self.tokenizer, "get_vocab_size"):
            return int(self.tokenizer.get_vocab_size())
        return int(len(self.tokenizer))

    def token_to_id(self, token: str) -> int:
        if hasattr(self.tokenizer, "token_to_id"):
            idx = self.tokenizer.token_to_id(token)
        elif hasattr(self.tokenizer, "convert_tokens_to_ids"):
            idx = self.tokenizer.convert_tokens_to_ids(token)
        else:
            raise TypeError("Unsupported tokenizer type")

        if idx is None:
            raise ValueError(f"Token {token!r} is not in the tokenizer vocabulary")
        return int(idx)

    def encode(self, text: str, add_special_tokens: bool = False, max_length: int | None = None) -> list[int]:
        if hasattr(self.tokenizer, "encode") and self.tokenizer.__class__.__module__.startswith("tokenizers"):
            ids = self.tokenizer.encode(text).ids
        else:
            ids = self.tokenizer.encode(text, add_special_tokens=add_special_tokens)
            add_special_tokens = False

        if add_special_tokens:
            ids = [self.bos_token_id] + list(ids) + [self.eos_token_id]

        if max_length is not None:
            ids = list(ids)[:max_length]

        return list(ids)

    def decode(self, ids: Sequence[int], skip_special_tokens: bool = True) -> str:
        if hasattr(self.tokenizer, "decode"):
            try:
                return self.tokenizer.decode(list(ids), skip_special_tokens=skip_special_tokens)
            except TypeError:
                return self.tokenizer.decode(list(ids))
        raise TypeError("Unsupported tokenizer type")

    def save(self, path: str | Path) -> None:
        path = Path(path)
        path.parent.mkdir(parents=True, exist_ok=True)

        if hasattr(self.tokenizer, "save"):
            self.tokenizer.save(str(path))
            return
        if hasattr(self.tokenizer, "save_pretrained"):
            self.tokenizer.save_pretrained(str(path))
            return
        raise TypeError("Unsupported tokenizer type")


def _special_tokens(config: Mapping | None = None) -> dict[str, str]:
    tokens = {
        "pad": "<pad>",
        "unk": "<unk>",
        "bos": "<s>",
        "eos": "</s>",
    }
    if config:
        tokens.update(dict(config))
    return tokens


def train_bpe_tokenizer(
    texts: Iterable[str],
    vocab_size: int = 32000,
    min_frequency: int = 2,
    special_tokens: Mapping[str, str] | None = None,
    save_path: str | Path | None = None,
) -> TokenizerWrapper:
    """Train a byte-level BPE tokenizer on source and target training text."""
    from tokenizers import Tokenizer
    from tokenizers.decoders import ByteLevel as ByteLevelDecoder
    from tokenizers.models import BPE
    from tokenizers.normalizers import NFKC, Sequence as NormalizerSequence
    from tokenizers.pre_tokenizers import ByteLevel
    from tokenizers.trainers import BpeTrainer

    tokens = _special_tokens(special_tokens)
    ordered_specials = [tokens["pad"], tokens["unk"], tokens["bos"], tokens["eos"]]

    tokenizer = Tokenizer(BPE(unk_token=tokens["unk"]))
    tokenizer.normalizer = NormalizerSequence([NFKC()])
    tokenizer.pre_tokenizer = ByteLevel(add_prefix_space=False)
    tokenizer.decoder = ByteLevelDecoder()

    trainer = BpeTrainer(
        vocab_size=vocab_size,
        min_frequency=min_frequency,
        special_tokens=ordered_specials,
        show_progress=True,
    )
    tokenizer.train_from_iterator((text for text in texts if text), trainer=trainer)

    wrapper = TokenizerWrapper(
        tokenizer,
        pad_token=tokens["pad"],
        unk_token=tokens["unk"],
        bos_token=tokens["bos"],
        eos_token=tokens["eos"],
    )

    if save_path is not None:
        wrapper.save(save_path)

    return wrapper


def build_tokenizer(config: Mapping, train_texts: Iterable[str] | None = None) -> TokenizerWrapper:
    """Build a tokenizer from project config."""
    tokenizer_type = config.get("type", "bpe")
    tokens = _special_tokens(config.get("special_tokens"))

    if tokenizer_type == "pretrained":
        from transformers import AutoTokenizer

        model_name = config.get("model_name") or config.get("pretrained_model_name")
        if not model_name:
            raise ValueError("pretrained tokenizer requires config['model_name']")
        tokenizer = AutoTokenizer.from_pretrained(model_name)
        return TokenizerWrapper(
            tokenizer,
            pad_token=tokenizer.pad_token or tokens["pad"],
            unk_token=tokenizer.unk_token or tokens["unk"],
            bos_token=tokenizer.bos_token or tokens["bos"],
            eos_token=tokenizer.eos_token or tokens["eos"],
        )

    if tokenizer_type in {"bpe", "sentencepiece"}:
        tokenizer_path = config.get("path") or config.get("tokenizer_path")
        if tokenizer_path and Path(tokenizer_path).exists():
            from tokenizers import Tokenizer

            return TokenizerWrapper(
                Tokenizer.from_file(str(tokenizer_path)),
                pad_token=tokens["pad"],
                unk_token=tokens["unk"],
                bos_token=tokens["bos"],
                eos_token=tokens["eos"],
            )

        if train_texts is None:
            raise ValueError("BPE tokenizer requires train_texts when no tokenizer path is provided")

        return train_bpe_tokenizer(
            train_texts,
            vocab_size=int(config.get("vocab_size", 32000)),
            min_frequency=int(config.get("min_frequency", 2)),
            special_tokens=tokens,
            save_path=tokenizer_path,
        )

    raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")