from __future__ import annotations import json from pathlib import Path from typing import Sequence, List import sentencepiece as spm def train_sentencepiece( data_files: Sequence[str], model_prefix: str = 'tokenizer', vocab_size: int = 32000, model_type: str = 'bpe', character_coverage: float = 1.0, byte_fallback: bool = True, pad_id: int = 1, unk_id: int = 0, bos_id: int = 2, eos_id: int = 3, ) -> str: data_files = [str(Path(p)) for p in data_files] if not data_files: raise ValueError('data_files is empty') args = [ f"--input={','.join(data_files)}", f'--model_prefix={model_prefix}', f'--vocab_size={int(vocab_size)}', f'--model_type={model_type}', f'--character_coverage={character_coverage}', f'--pad_id={pad_id}', f'--unk_id={unk_id}', f'--bos_id={bos_id}', f'--eos_id={eos_id}', '--hard_vocab_limit=false', '--normalization_rule_name=nmt_nfkc', ] if byte_fallback: args.append('--byte_fallback=true') spm.SentencePieceTrainer.train(' '.join(args)) return f'{model_prefix}.model' class TokenizerWrapper: def __init__(self, model_path: str): model_path = str(Path(model_path)) if not Path(model_path).exists(): raise FileNotFoundError(model_path) self.sp = spm.SentencePieceProcessor(model_file=model_path) self.vocab_size = int(self.sp.vocab_size()) self.pad_id = self.sp.pad_id() self.unk_id = self.sp.unk_id() self.bos_id = self.sp.bos_id() self.eos_id = self.sp.eos_id() for name, val in [('pad', self.pad_id), ('unk', self.unk_id), ('bos', self.bos_id), ('eos', self.eos_id)]: if val < 0: raise ValueError(f'SentencePiece model missing <{name}>') def encode(self, text: str, add_bos: bool = True, add_eos: bool = False) -> List[int]: ids = list(self.sp.encode(text, out_type=int)) if add_bos: ids = [self.bos_id] + ids if add_eos: ids = ids + [self.eos_id] return ids def encode_batch(self, texts: Sequence[str], add_bos: bool = True, add_eos: bool = False) -> List[List[int]]: return [self.encode(t, add_bos=add_bos, add_eos=add_eos) for t in texts] def decode(self, ids: Sequence[int]) -> str: return self.sp.decode([int(i) for i in ids if int(i) != self.pad_id]) def save_config(self, path: str) -> None: Path(path).write_text(json.dumps({ 'vocab_size': self.vocab_size, 'pad_id': self.pad_id, 'unk_id': self.unk_id, 'bos_id': self.bos_id, 'eos_id': self.eos_id, }, indent=2), encoding='utf-8')