CoolFace
Modelpublic

swgoo/pmnet

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes8downloads
tokenization_pmnet.py92 linesDownload Raw Back to root
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