import collections import json import os from typing import List, Optional, Tuple from transformers import PreTrainedTokenizer VOCAB_FILES_NAMES = {"vocab_file": "vocab.txt"} VOCAB_SIZE_TO_KMER = {69: 3, 261: 4, 1029: 5, 4101: 6} def load_vocab(vocab_file): vocab = collections.OrderedDict() with open(vocab_file, "r", encoding="utf-8") as f: for index, line in enumerate(f): token = line.rstrip("\n") vocab[token] = index return vocab class UTRBertTokenizer(PreTrainedTokenizer): vocab_files_names = VOCAB_FILES_NAMES model_input_names = ["input_ids", "attention_mask"] def __init__( self, vocab_file, unk_token="[UNK]", sep_token="[SEP]", pad_token="[PAD]", cls_token="[CLS]", mask_token="[MASK]", **kwargs, ): self._vocab = load_vocab(vocab_file) self._ids_to_tokens = {v: k for k, v in self._vocab.items()} vocab_size = len(self._vocab) if vocab_size not in VOCAB_SIZE_TO_KMER: raise ValueError(f"Unrecognised vocab size {vocab_size}; expected one of {list(VOCAB_SIZE_TO_KMER)}") self.kmer = VOCAB_SIZE_TO_KMER[vocab_size] super().__init__( unk_token=unk_token, sep_token=sep_token, pad_token=pad_token, cls_token=cls_token, mask_token=mask_token, **kwargs, ) @property def vocab_size(self): return len(self._vocab) def get_vocab(self): return dict(self._vocab) def _tokenize(self, text: str) -> List[str]: seq = text.upper().replace("T", "U").replace(" ", "") k = self.kmer return [seq[i : i + k] for i in range(len(seq) + 1 - k)] def _convert_token_to_id(self, token: str) -> int: return self._vocab.get(token, self._vocab.get(self.unk_token, 0)) def _convert_id_to_token(self, index: int) -> str: return self._ids_to_tokens.get(index, self.unk_token) def convert_tokens_to_string(self, tokens: List[str]) -> str: return " ".join(tokens) def build_inputs_with_special_tokens(self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None) -> List[int]: cls = [self.cls_token_id] sep = [self.sep_token_id] if token_ids_1 is None: return cls + token_ids_0 + sep return cls + token_ids_0 + sep + token_ids_1 + sep def get_special_tokens_mask(self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None, already_has_special_tokens: bool = False) -> List[int]: if already_has_special_tokens: return super().get_special_tokens_mask(token_ids_0, token_ids_1, already_has_special_tokens=True) if token_ids_1 is None: return [1] + [0] * len(token_ids_0) + [1] return [1] + [0] * len(token_ids_0) + [1] + [0] * len(token_ids_1) + [1] def create_token_type_ids_from_sequences(self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None) -> List[int]: sep = [self.sep_token_id] cls = [self.cls_token_id] if token_ids_1 is None: return [0] * len(cls + token_ids_0 + sep) return [0] * len(cls + token_ids_0 + sep) + [1] * len(token_ids_1 + sep) def save_vocabulary(self, save_directory: str, filename_prefix: Optional[str] = None) -> Tuple[str]: os.makedirs(save_directory, exist_ok=True) fname = (filename_prefix + "-" if filename_prefix else "") + "vocab.txt" path = os.path.join(save_directory, fname) with open(path, "w", encoding="utf-8") as f: for token, _ in sorted(self._vocab.items(), key=lambda kv: kv[1]): f.write(token + "\n") return (path,)