xiexinyuan341's picture
Implement data pipeline module
a2d6c00 verified
Raw
History Blame
6.77 kB
"""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}")