swgoo/pmnet
08
1from transformers import PreTrainedTokenizer2from typing import Dict, List, Optional, Any3 4 5class ByteTokenizer(PreTrainedTokenizer):6 model_input_names = ["input_ids", "attention_mask"]7 8 def __init__(9 self,10 bos_token="<|bos|>",11 eos_token="<|eos|>",12 pad_token="<|pad|>",13 vocab_size=384,14 **kwargs,15 ):16 self.pad_idx = 017 self.bos_idx = 25418 self.eos_idx = 25519 self._vocab_size = vocab_size20 21 self.byte_to_token = [f"<byte_{i}>" for i in range(256)]22 self.token_to_byte = {t: i for i, t in enumerate(self.byte_to_token)}23 24 super().__init__(25 bos_token=bos_token,26 eos_token=eos_token,27 pad_token=pad_token,28 **kwargs,29 )30 31 @property32 def vocab_size(self) -> int:33 return self._vocab_size34 35 def get_vocab(self) -> Dict[str, int]:36 vocab = {t: i for i, t in enumerate(self.byte_to_token)}37 vocab.update(38 {39 self.bos_token: self.bos_idx,40 self.eos_token: self.eos_idx,41 self.pad_token: self.pad_idx,42 }43 )44 return vocab45 46 def _tokenize(self, text, **kwargs):47 return [self.byte_to_token[b] for b in text.encode("utf-8")]48 49 def _convert_token_to_id(self, token):50 if token == self.bos_token:51 return self.bos_idx52 if token == self.eos_token:53 return self.eos_idx54 if token == self.pad_token:55 return self.pad_idx56 return self.token_to_byte.get(token, self.pad_idx)57 58 def _convert_id_to_token(self, index):59 if index == self.bos_idx:60 return self.bos_token61 if index == self.eos_idx:62 return self.eos_token63 if index == self.pad_idx:64 return self.pad_token65 if 0 <= index < 256:66 return self.byte_to_token[index]67 return f"<unk_{index}>"68 69 def build_inputs_with_special_tokens(self, token_ids_0, token_ids_1=None):70 return [self.bos_idx] + token_ids_0 + [self.eos_idx]71 72 def _decode(73 self, token_ids: List[int], skip_special_tokens: bool = False, **kwargs74 ) -> str:75 clean_ids = []76 for i in token_ids:77 if skip_special_tokens and i in [self.bos_idx, self.eos_idx, self.pad_idx]:78 continue79 if 0 <= i < 256:80 clean_ids.append(i)81 return bytes(clean_ids).decode("utf-8", errors="ignore")82 83 def save_vocabulary(84 self, save_directory: str, filename_prefix: Optional[str] = None85 ) -> tuple:86 return ()87 88 89__all__ = [90 "ByteTokenizer",91]92 