kelu01's picture
Update tokenizer.py
d7a1630 verified
Raw
History Blame Contribute Delete
3.04 kB
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)