GenerTeam/GENERator-v2-eukaryote-3b-base
1700
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 