File size: 3,041 Bytes
9d2fb14
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
d7a1630
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
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:
            # If token_ids_1 is provided, combine both sequences
            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)