DoruC/Grounded-Segment-Anything
0
1# coding=utf-82# Copyright 2021 VinAI Research and the HuggingFace Inc. team.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License15""" Tokenization classes for BARTpho-syllable model."""16 17 18import os19from shutil import copyfile20from typing import Any, Dict, List, Optional, Tuple21 22import sentencepiece as spm23 24from ...tokenization_utils import AddedToken, PreTrainedTokenizer25from ...utils import logging26 27 28logger = logging.get_logger(__name__)29 30SPIECE_UNDERLINE = "โ"31 32VOCAB_FILES_NAMES = {"vocab_file": "sentencepiece.bpe.model", "monolingual_vocab_file": "dict.txt"}33 34PRETRAINED_VOCAB_FILES_MAP = {35 "vocab_file": {36 "vinai/bartpho-syllable": "https://huggingface.co/vinai/bartpho-syllable/resolve/main/sentencepiece.bpe.model",37 },38 "monolingual_vocab_file": {39 "vinai/bartpho-syllable": "https://huggingface.co/vinai/bartpho-syllable/resolve/main/dict.txt",40 },41}42 43PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES = {"vinai/bartpho-syllable": 1024}44 45 46class BartphoTokenizer(PreTrainedTokenizer):47 """48 Adapted from [`XLMRobertaTokenizer`]. Based on [SentencePiece](https://github.com/google/sentencepiece).49 50 This tokenizer inherits from [`PreTrainedTokenizer`] which contains most of the main methods. Users should refer to51 this superclass for more information regarding those methods.52 53 Args:54 vocab_file (`str`):55 Path to the vocabulary file. This vocabulary is the pre-trained SentencePiece model available from the56 multilingual XLM-RoBERTa, also used in mBART, consisting of 250K types.57 monolingual_vocab_file (`str`):58 Path to the monolingual vocabulary file. This monolingual vocabulary consists of Vietnamese-specialized59 types extracted from the multilingual vocabulary vocab_file of 250K types.60 bos_token (`str`, *optional*, defaults to `"<s>"`):61 The beginning of sequence token that was used during pretraining. Can be used a sequence classifier token.62 63 <Tip>64 65 When building a sequence using special tokens, this is not the token that is used for the beginning of66 sequence. The token used is the `cls_token`.67 68 </Tip>69 70 eos_token (`str`, *optional*, defaults to `"</s>"`):71 The end of sequence token.72 73 <Tip>74 75 When building a sequence using special tokens, this is not the token that is used for the end of sequence.76 The token used is the `sep_token`.77 78 </Tip>79 80 sep_token (`str`, *optional*, defaults to `"</s>"`):81 The separator token, which is used when building a sequence from multiple sequences, e.g. two sequences for82 sequence classification or for a text and a question for question answering. It is also used as the last83 token of a sequence built with special tokens.84 cls_token (`str`, *optional*, defaults to `"<s>"`):85 The classifier token which is used when doing sequence classification (classification of the whole sequence86 instead of per-token classification). It is the first token of the sequence when built with special tokens.87 unk_token (`str`, *optional*, defaults to `"<unk>"`):88 The unknown token. A token that is not in the vocabulary cannot be converted to an ID and is set to be this89 token instead.90 pad_token (`str`, *optional*, defaults to `"<pad>"`):91 The token used for padding, for example when batching sequences of different lengths.92 mask_token (`str`, *optional*, defaults to `"<mask>"`):93 The token used for masking values. This is the token used when training this model with masked language94 modeling. This is the token which the model will try to predict.95 sp_model_kwargs (`dict`, *optional*):96 Will be passed to the `SentencePieceProcessor.__init__()` method. The [Python wrapper for97 SentencePiece](https://github.com/google/sentencepiece/tree/master/python) can be used, among other things,98 to set:99 100 - `enable_sampling`: Enable subword regularization.101 - `nbest_size`: Sampling parameters for unigram. Invalid for BPE-Dropout.102 103 - `nbest_size = {0,1}`: No sampling is performed.104 - `nbest_size > 1`: samples from the nbest_size results.105 - `nbest_size < 0`: assuming that nbest_size is infinite and samples from the all hypothesis (lattice)106 using forward-filtering-and-backward-sampling algorithm.107 108 - `alpha`: Smoothing parameter for unigram sampling, and dropout probability of merge operations for109 BPE-dropout.110 111 Attributes:112 sp_model (`SentencePieceProcessor`):113 The *SentencePiece* processor that is used for every conversion (string, tokens and IDs).114 """115 116 vocab_files_names = VOCAB_FILES_NAMES117 pretrained_vocab_files_map = PRETRAINED_VOCAB_FILES_MAP118 max_model_input_sizes = PRETRAINED_POSITIONAL_EMBEDDINGS_SIZES119 model_input_names = ["input_ids", "attention_mask"]120 121 def __init__(122 self,123 vocab_file,124 monolingual_vocab_file,125 bos_token="<s>",126 eos_token="</s>",127 sep_token="</s>",128 cls_token="<s>",129 unk_token="<unk>",130 pad_token="<pad>",131 mask_token="<mask>",132 sp_model_kwargs: Optional[Dict[str, Any]] = None,133 **kwargs,134 ) -> None:135 # Mask token behave like a normal word, i.e. include the space before it136 mask_token = AddedToken(mask_token, lstrip=True, rstrip=False) if isinstance(mask_token, str) else mask_token137 138 self.sp_model_kwargs = {} if sp_model_kwargs is None else sp_model_kwargs139 140 self.vocab_file = vocab_file141 self.monolingual_vocab_file = monolingual_vocab_file142 self.sp_model = spm.SentencePieceProcessor(**self.sp_model_kwargs)143 self.sp_model.Load(str(vocab_file))144 145 # Load the reduced vocab146 147 # Keep order of special tokens for backward compatibility148 self.fairseq_tokens_to_ids = {}149 cnt = 0150 for token in [bos_token, pad_token, eos_token, unk_token, sep_token, cls_token]:151 if str(token) not in self.fairseq_tokens_to_ids:152 self.fairseq_tokens_to_ids[str(token)] = cnt153 cnt += 1154 with open(monolingual_vocab_file, "r", encoding="utf-8") as f:155 for line in f.readlines():156 token = line.strip().split()[0]157 self.fairseq_tokens_to_ids[token] = len(self.fairseq_tokens_to_ids)158 if str(mask_token) not in self.fairseq_tokens_to_ids:159 self.fairseq_tokens_to_ids[str(mask_token)] = len(self.fairseq_tokens_to_ids)160 161 self.fairseq_ids_to_tokens = {v: k for k, v in self.fairseq_tokens_to_ids.items()}162 163 super().__init__(164 bos_token=bos_token,165 eos_token=eos_token,166 unk_token=unk_token,167 sep_token=sep_token,168 cls_token=cls_token,169 pad_token=pad_token,170 mask_token=mask_token,171 sp_model_kwargs=self.sp_model_kwargs,172 **kwargs,173 )174 175 def __getstate__(self):176 state = self.__dict__.copy()177 state["sp_model"] = None178 state["sp_model_proto"] = self.sp_model.serialized_model_proto()179 return state180 181 def __setstate__(self, d):182 self.__dict__ = d183 184 # for backward compatibility185 if not hasattr(self, "sp_model_kwargs"):186 self.sp_model_kwargs = {}187 188 self.sp_model = spm.SentencePieceProcessor(**self.sp_model_kwargs)189 self.sp_model.LoadFromSerializedProto(self.sp_model_proto)190 191 def build_inputs_with_special_tokens(192 self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None193 ) -> List[int]:194 """195 Build model inputs from a sequence or a pair of sequence for sequence classification tasks by concatenating and196 adding special tokens. An BARTPho sequence has the following format:197 198 - single sequence: `<s> X </s>`199 - pair of sequences: `<s> A </s></s> B </s>`200 201 Args:202 token_ids_0 (`List[int]`):203 List of IDs to which the special tokens will be added.204 token_ids_1 (`List[int]`, *optional*):205 Optional second list of IDs for sequence pairs.206 207 Returns:208 `List[int]`: List of [input IDs](../glossary#input-ids) with the appropriate special tokens.209 """210 211 if token_ids_1 is None:212 return [self.cls_token_id] + token_ids_0 + [self.sep_token_id]213 cls = [self.cls_token_id]214 sep = [self.sep_token_id]215 return cls + token_ids_0 + sep + sep + token_ids_1 + sep216 217 def get_special_tokens_mask(218 self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None, already_has_special_tokens: bool = False219 ) -> List[int]:220 """221 Retrieve sequence ids from a token list that has no special tokens added. This method is called when adding222 special tokens using the tokenizer `prepare_for_model` method.223 224 Args:225 token_ids_0 (`List[int]`):226 List of IDs.227 token_ids_1 (`List[int]`, *optional*):228 Optional second list of IDs for sequence pairs.229 already_has_special_tokens (`bool`, *optional*, defaults to `False`):230 Whether or not the token list is already formatted with special tokens for the model.231 232 Returns:233 `List[int]`: A list of integers in the range [0, 1]: 1 for a special token, 0 for a sequence token.234 """235 236 if already_has_special_tokens:237 return super().get_special_tokens_mask(238 token_ids_0=token_ids_0, token_ids_1=token_ids_1, already_has_special_tokens=True239 )240 241 if token_ids_1 is None:242 return [1] + ([0] * len(token_ids_0)) + [1]243 return [1] + ([0] * len(token_ids_0)) + [1, 1] + ([0] * len(token_ids_1)) + [1]244 245 def create_token_type_ids_from_sequences(246 self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None247 ) -> List[int]:248 """249 Create a mask from the two sequences passed to be used in a sequence-pair classification task. BARTPho does not250 make use of token type ids, therefore a list of zeros is returned.251 252 Args:253 token_ids_0 (`List[int]`):254 List of IDs.255 token_ids_1 (`List[int]`, *optional*):256 Optional second list of IDs for sequence pairs.257 258 Returns:259 `List[int]`: List of zeros.260 261 """262 263 sep = [self.sep_token_id]264 cls = [self.cls_token_id]265 266 if token_ids_1 is None:267 return len(cls + token_ids_0 + sep) * [0]268 return len(cls + token_ids_0 + sep + sep + token_ids_1 + sep) * [0]269 270 @property271 def vocab_size(self):272 return len(self.fairseq_ids_to_tokens)273 274 def get_vocab(self):275 vocab = {self.convert_ids_to_tokens(i): i for i in range(self.vocab_size)}276 vocab.update(self.added_tokens_encoder)277 return vocab278 279 def _tokenize(self, text: str) -> List[str]:280 return self.sp_model.encode(text, out_type=str)281 282 def _convert_token_to_id(self, token):283 """Converts a token (str) in an id using the vocab."""284 if token in self.fairseq_tokens_to_ids:285 return self.fairseq_tokens_to_ids[token]286 else:287 return self.unk_token_id288 289 def _convert_id_to_token(self, index):290 """Converts an index (integer) in a token (str) using the vocab."""291 return self.fairseq_ids_to_tokens[index]292 293 def convert_tokens_to_string(self, tokens):294 """Converts a sequence of tokens (strings for sub-words) in a single string."""295 out_string = "".join(tokens).replace(SPIECE_UNDERLINE, " ").strip()296 return out_string297 298 def save_vocabulary(self, save_directory: str, filename_prefix: Optional[str] = None) -> Tuple[str]:299 if not os.path.isdir(save_directory):300 logger.error(f"Vocabulary path ({save_directory}) should be a directory")301 return302 out_vocab_file = os.path.join(303 save_directory, (filename_prefix + "-" if filename_prefix else "") + VOCAB_FILES_NAMES["vocab_file"]304 )305 out_monolingual_vocab_file = os.path.join(306 save_directory,307 (filename_prefix + "-" if filename_prefix else "") + VOCAB_FILES_NAMES["monolingual_vocab_file"],308 )309 310 if os.path.abspath(self.vocab_file) != os.path.abspath(out_vocab_file) and os.path.isfile(self.vocab_file):311 copyfile(self.vocab_file, out_vocab_file)312 elif not os.path.isfile(self.vocab_file):313 with open(out_vocab_file, "wb") as fi:314 content_spiece_model = self.sp_model.serialized_model_proto()315 fi.write(content_spiece_model)316 317 if os.path.abspath(self.monolingual_vocab_file) != os.path.abspath(318 out_monolingual_vocab_file319 ) and os.path.isfile(self.monolingual_vocab_file):320 copyfile(self.monolingual_vocab_file, out_monolingual_vocab_file)321 elif not os.path.isfile(self.monolingual_vocab_file):322 with open(out_monolingual_vocab_file, "w", encoding="utf-8") as fp:323 for token in self.fairseq_tokens_to_ids:324 if token not in self.all_special_tokens:325 fp.write(f"{str(token)} \n")326 327 return out_vocab_file, out_monolingual_vocab_file328 