CoolFace
Modelpublic

GenerTeam/GENERator-v2-eukaryote-3b-base

sourceHugging Facemitupdated 3mo agoView on Hugging Face
1likes700downloads
tokenizer.py176 linesDownload Raw Back to root
1import itertools2import os3import json4import re5from typing import List, Optional, Tuple6from transformers import PreTrainedTokenizer7 8class DNAKmerTokenizer(PreTrainedTokenizer):9    def __init__(self, k, add_bos_token=True, add_eos_token=False, **kwargs):10        self.k = k11        self.add_bos_token = add_bos_token12        self.add_eos_token = add_eos_token13        self.special_tokens = [14            "<oov>",15            "<s>",16            "</s>",17            "<pad>",18            "<mask>",19            "<bog>",20            "<eog>",21            "<bok>",22            "<eok>",23            "<+>",24            "<->",25            "<cds>",26            "<pseudo>",27            "<tRNA>",28            "<rRNA>",29            "<ncRNA>",30            "<miscRNA>",31            "<mam>",32            "<vrt>",33            "<inv>",34            "<pln>",35            "<fng>",36            "<prt>",37            "<arc>",38            "<bct>",39            "<mit>",40            "<plt>",41            "<plm>",42            "<vir>",43            "<sp0>",44            "<sp1>",45            "<sp2>",46        ]47        self.kmers = [48            "".join(kmer) for kmer in itertools.product("ATCG", repeat=self.k)49        ]50        self.vocab = {51            token: i for i, token in enumerate(self.special_tokens + self.kmers)52        }53        self.ids_to_tokens = {v: k for k, v in self.vocab.items()}54        self.special_token_pattern = re.compile(55            "|".join(re.escape(token) for token in self.special_tokens)56        )57        self.dna_pattern = re.compile(f"[A-Z]{{{self.k}}}|[A-Z]+")58        kwargs.setdefault("unk_token", "<oov>")59        kwargs.setdefault("bos_token", "<s>")60        kwargs.setdefault("eos_token", "</s>")61        kwargs.setdefault("pad_token", "<pad>")62        kwargs.setdefault("mask_token", "<mask>")63        super().__init__(**kwargs)64 65    @property66    def vocab_size(self):67        return len(self.vocab)68 69    def get_vocab(self):70        return dict(self.vocab)71 72    def _tokenize(self, text, **kwargs) -> List[str]:73        tokens = []74        pos = 075        while pos < len(text):76            special_match = self.special_token_pattern.match(text, pos)77            if special_match:78                tokens.append(special_match.group())79                pos = special_match.end()80            else:81                dna_match = self.dna_pattern.match(text, pos)82                if dna_match:83                    dna_seq = dna_match.group()84                    tokens.append(dna_seq)85                    pos = dna_match.end()86                else:87                    tokens.append(text[pos])88                    pos += 189        return tokens90 91    def _convert_token_to_id(self, token: str) -> int:92        return self.vocab.get(token, self.vocab["<oov>"])93 94    def _convert_id_to_token(self, index: int) -> str:95        return self.ids_to_tokens.get(index, "<oov>")96 97    def convert_tokens_to_string(self, tokens: List[str]) -> str:98        return "".join(tokens)99 100    def build_inputs_with_special_tokens(self, token_ids_0, token_ids_1=None):101        bos = [self.bos_token_id] if self.add_bos_token else []102        eos = [self.eos_token_id] if self.add_eos_token else []103 104        if token_ids_1 is None:105            return bos + token_ids_0 + eos106        # Dual sequence case: bos + token_ids_0 + bos + token_ids_1 + eos107        return bos + token_ids_0 + bos + token_ids_1 + eos108 109    def get_special_tokens_mask(110            self, token_ids_0, token_ids_1=None, already_has_special_tokens=False111    ):112        if already_has_special_tokens:113            return super().get_special_tokens_mask(114                token_ids_0, token_ids_1, already_has_special_tokens=True115            )116 117        bos_mask = [1] if self.add_bos_token else []118        eos_mask = [1] if self.add_eos_token else []119 120        if token_ids_1 is None:121            return bos_mask + ([0] * len(token_ids_0)) + eos_mask122        # Dual sequence case: bos + token_ids_0 + bos + token_ids_1 + eos123        return bos_mask + ([0] * len(token_ids_0)) + bos_mask + ([0] * len(token_ids_1)) + eos_mask124 125    def prepare_for_model(self, *args, **kwargs):126        encoding = super().prepare_for_model(*args, **kwargs)127        if "token_type_ids" in encoding:128            del encoding["token_type_ids"]129        return encoding130 131    def save_vocabulary(132            self, save_directory: str, filename_prefix: Optional[str] = None133    ) -> Tuple[str]:134        if not os.path.exists(save_directory):135            os.makedirs(save_directory)136 137        vocab_file = os.path.join(138            save_directory,139            (filename_prefix + "-" if filename_prefix else "") + "vocab.txt",140        )141 142        with open(vocab_file, "w", encoding="utf-8") as f:143            for token, idx in sorted(self.vocab.items(), key=lambda x: x[1]):144                f.write(f"{token} {idx}\n")145        return (vocab_file,)146    147    def save_pretrained(self, save_directory: str, **kwargs):148        vocab_files = super().save_pretrained(save_directory, **kwargs)149        tokenizer_config_path = os.path.join(save_directory, "tokenizer_config.json")150 151        # Read existing config or create new one152        if os.path.exists(tokenizer_config_path):153            with open(tokenizer_config_path, "r", encoding="utf-8") as f:154                config = json.load(f)155        else:156            config = {}157 158        # Add auto_map configuration159        config.update({160            "auto_map": {161                "AutoTokenizer": [162                    "tokenizer.DNAKmerTokenizer",163                    None164                ]165            },166            "k": self.k,167            "add_bos_token": self.add_bos_token,168            "add_eos_token": self.add_eos_token,169        })170 171        # Save config172        with open(tokenizer_config_path, "w", encoding="utf-8") as f:173            json.dump(config, f, ensure_ascii=False, indent=2)174 175        return vocab_files176