"""SentencePiece tokenizer preserving BashkirRoBERTa's original token IDs.""" from pathlib import Path import sentencepiece as spm from transformers import PreTrainedTokenizer class BashkirRobertaTokenizer(PreTrainedTokenizer): vocab_files_names = {"vocab_file": "spm_bashkir_bert_16k.model"} model_input_names = ["input_ids", "attention_mask"] def __init__(self, vocab_file, **kwargs): self.vocab_file = vocab_file self.sp_model = spm.SentencePieceProcessor(model_file=vocab_file) defaults = { "bos_token": "", "eos_token": "", "unk_token": "", "pad_token": "", "cls_token": "[CLS]", "sep_token": "[SEP]", "mask_token": "[MASK]", } for key, value in defaults.items(): kwargs.setdefault(key, value) super().__init__(**kwargs) @property def vocab_size(self): return self.sp_model.get_piece_size() def get_vocab(self): return {self.sp_model.id_to_piece(i): i for i in range(self.vocab_size)} def _tokenize(self, text): return self.sp_model.encode(text, out_type=str) def _convert_token_to_id(self, token): return self.sp_model.piece_to_id(token) def _convert_id_to_token(self, index): return self.sp_model.id_to_piece(index) def convert_tokens_to_string(self, tokens): return self.sp_model.decode(tokens) def build_inputs_with_special_tokens(self, token_ids_0, token_ids_1=None): if token_ids_1 is not None: return [self.bos_token_id] + token_ids_0 + [self.sep_token_id] + token_ids_1 + [self.eos_token_id] return [self.bos_token_id] + token_ids_0 + [self.eos_token_id] def get_special_tokens_mask(self, token_ids_0, token_ids_1=None, already_has_special_tokens=False): if already_has_special_tokens: return super().get_special_tokens_mask(token_ids_0, token_ids_1, True) if token_ids_1 is None: return [1] + [0] * len(token_ids_0) + [1] return [1] + [0] * len(token_ids_0) + [1] + [0] * len(token_ids_1) + [1] def save_vocabulary(self, save_directory, filename_prefix=None): source = Path(self.vocab_file) name = ((filename_prefix + "-") if filename_prefix else "") + self.vocab_files_names["vocab_file"] target = Path(save_directory) / name target.write_bytes(source.read_bytes()) return (str(target),)