| """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}") |
|
|