Pragya / tokenizer.py
Arush kumar
Upload tokenizer.py
bf55e4e
Raw History Blame
2.86 kB
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')