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