| from transformers import PreTrainedTokenizer |
| import os |
| import re |
|
|
| def replace_halogen(smiles): |
| return smiles.replace("Cl", "L").replace("Br", "R") |
|
|
| def restore_halogen(smiles): |
| return smiles.replace("L", "Cl").replace("R", "Br") |
|
|
|
|
| class SMILESTokenizer(PreTrainedTokenizer): |
| def __init__(self, vocab_file=None, max_length=140, **kwargs): |
| self.special_tokens = ['[MASK]'] |
| self.additional_chars = set() |
| self.max_length = max_length |
|
|
| self.chars = self.special_tokens |
| self.vocab = {} |
| self.ids_to_tokens = {} |
|
|
| if vocab_file is not None: |
| with open(vocab_file, 'r') as f: |
| chars = f.read().split() |
| self.add_characters(chars) |
|
|
| super().__init__(unk_token='[MASK]', mask_token='[MASK]', **kwargs) |
|
|
| def tokenize(self, smiles): |
| smiles = replace_halogen(smiles) |
| regex = '(\[[^\[\]]{1,6}\])' |
| char_list = re.split(regex, smiles) |
| tokens = [] |
| for group in char_list: |
| if group.startswith('['): |
| tokens.append(group) |
| else: |
| tokens.extend(list(group)) |
| return tokens |
|
|
| def _convert_token_to_id(self, token): |
| return self.vocab.get(token, self.vocab['[MASK]']) |
|
|
| def _convert_id_to_token(self, index): |
| return self.ids_to_tokens.get(index, '[MASK]') |
|
|
| def convert_tokens_to_string(self, tokens): |
| smiles = ''.join(tokens) |
| return restore_halogen(smiles) |
|
|
| def build_inputs_with_special_tokens(self, token_ids_0, token_ids_1=None): |
| if token_ids_1 is not None: |
| |
| return token_ids_0 + token_ids_1 |
| return token_ids_0 |
|
|
| def get_vocab(self): |
| return self.vocab |
|
|
| def add_characters(self, chars): |
| self.additional_chars.update(chars) |
| all_chars = sorted(list(self.additional_chars)) + self.special_tokens |
| self.vocab = dict(zip(all_chars, range(len(all_chars)))) |
| self.ids_to_tokens = {v: k for k, v in self.vocab.items()} |
|
|
| def save_vocabulary(self, save_directory): |
| path = os.path.join(save_directory, "vocab.txt") |
| with open(path, "w") as f: |
| for token in sorted(self.vocab, key=lambda x: self.vocab[x]): |
| f.write(token + "\n") |
| return (path,) |
| |
| @property |
| def vocab_size(self): |
| return len(self.vocab) |
| |
| def save_vocabulary(self, save_directory, filename_prefix=None): |
| vocab_file = os.path.join(save_directory, (filename_prefix + "-" if filename_prefix else "") + "vocab.txt") |
| with open(vocab_file, "w") as f: |
| for token in sorted(self.vocab, key=self.vocab.get): |
| f.write(token + "\n") |
| return (vocab_file,) |
| |
| @classmethod |
| def from_pretrained(cls, pretrained_model_name_or_path, *args, **kwargs): |
| vocab_file = os.path.join(pretrained_model_name_or_path, "vocab.txt") |
| return cls(vocab_file=vocab_file, *args, **kwargs) |