SpliNet / tokenization_splinet.py
Angshul's picture
Upload repaired SpliNet 2B-token pretrained model
4ec5e47 verified
Raw History Blame Contribute Delete
3.56 kB
import os
import shutil
import unicodedata
import sentencepiece as spm
from transformers import PreTrainedTokenizer
class SpliNetTokenizer(PreTrainedTokenizer):
vocab_files_names = {
"vocab_file": "spiece.model"
}
model_input_names = [
"input_ids",
"token_type_ids",
"attention_mask",
]
def __init__(
self,
vocab_file,
do_lower_case=True,
**kwargs,
):
self.vocab_file = vocab_file
self.do_lower_case = bool(do_lower_case)
self.sp_model = spm.SentencePieceProcessor(
model_file=vocab_file
)
# tokenizer_config.json may already provide these.
# setdefault prevents passing any keyword twice.
kwargs.setdefault("unk_token", "<unk>")
kwargs.setdefault("bos_token", "<s>")
kwargs.setdefault("eos_token", "</s>")
kwargs.setdefault("pad_token", "<pad>")
kwargs.setdefault("cls_token", "<cls>")
kwargs.setdefault("sep_token", "<sep>")
kwargs.setdefault("mask_token", "<mask>")
super().__init__(**kwargs)
@property
def vocab_size(self):
return int(
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 _normalize(self, text):
text = text or ""
if self.do_lower_case:
text = unicodedata.normalize(
"NFKC",
text,
).lower()
return " ".join(text.split())
def _tokenize(self, text):
return self.sp_model.encode(
self._normalize(text),
out_type=str,
)
def _convert_token_to_id(self, token):
return int(
self.sp_model.piece_to_id(token)
)
def _convert_id_to_token(self, index):
return self.sp_model.id_to_piece(
int(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 None:
return (
[self.cls_token_id]
+ list(token_ids_0)
+ [self.sep_token_id]
)
return (
[self.cls_token_id]
+ list(token_ids_0)
+ [self.sep_token_id]
+ list(token_ids_1)
+ [self.sep_token_id]
)
def create_token_type_ids_from_sequences(
self,
token_ids_0,
token_ids_1=None,
):
if token_ids_1 is None:
return [0] * (
len(token_ids_0) + 2
)
return (
[0] * (len(token_ids_0) + 2)
+ [1] * (len(token_ids_1) + 1)
)
def save_vocabulary(
self,
save_directory,
filename_prefix=None,
):
os.makedirs(
save_directory,
exist_ok=True,
)
prefix = (
filename_prefix + "-"
if filename_prefix
else ""
)
destination = os.path.join(
save_directory,
prefix + "spiece.model",
)
if (
os.path.abspath(self.vocab_file)
!= os.path.abspath(destination)
):
shutil.copy2(
self.vocab_file,
destination,
)
return (destination,)