bashkir-roberta / tokenization_bashkir_roberta.py
failed09's picture
Initial BashkirRoBERTa release
2b3226f verified
Raw
History Blame Contribute Delete
2.45 kB
"""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": "<s>", "eos_token": "</s>", "unk_token": "<unk>",
"pad_token": "<pad>", "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),)