| from unimol.data.dictionary import DecoderDictionary |
| import selfies as sf |
| from rdkit import Chem |
| from rdkit.Chem import AllChem |
| from rdkit.Chem.Crippen import MolLogP |
| from rdkit.Chem import MolFromSmiles |
|
|
|
|
| def one_hot_to_selfies(hot, dict1:DecoderDictionary): |
| '''> 3 means to get rid of special tokens in the molecule representation.''' |
| selfies_list = [] |
| |
| for idx in hot.transpose(0, 1).argmax(1): |
| if idx.item() == dict1.index('[SEP]') or idx.item() == dict1.index('[PAD]'): |
| break |
| elif idx.item() == dict1.index('[UNK]') or idx.item() == dict1.index('[CLS]'): |
| selfies_list.append('[nop]') |
| else: |
| selfies_list.append(dict1.index2symbol(idx.item())) |
| |
| |
| |
| return ''.join(selfies_list).replace(' ', '') |
|
|
|
|
| def one_hot_to_smiles(hot, dict_): |
| '''Return both the smile repre. and the selfies rep.''' |
| selfies = one_hot_to_selfies(hot, dict_) |
| |
| |
| return sf.decoder(selfies) |
|
|
|
|
| def label_smiles(smiles:list): |
| """Label a batch of smiles to in the form of Unimol compatible dataset""" |
|
|
| selfies = [list(sf.split_selfies(sf.encoder(smile))) for smile in smiles] |
| new_data_list = [] |
| |
| for idx, smile in enumerate(smiles): |
| data_dict = dict() |
| try: |
| m = Chem.MolFromSmiles(smile) |
| m3d = Chem.AddHs(m) |
| except: |
| |
| continue |
|
|
| atom_list = [] |
| for atom in m3d.GetAtoms(): |
| atom_list.append(atom.GetSymbol()) |
| |
| selfie = selfies[idx] |
|
|
| |
|
|
| |
| |
|
|
|
|
|
|
|
|
| data_dict['atoms'] = atom_list |
| |
| |
| |
| |
| |
| |
| |
| data_dict['coordinates'] = [] |
| |
| data_dict['smi'] = smile |
| data_dict['scaffold'] = '' |
| data_dict['ori_index'] = -1 |
| data_dict['selfies'] = selfies[idx] |
| data_dict['target'] = MolLogP(MolFromSmiles(smile)) |
|
|
| new_data_list.append(data_dict) |
| |
| return new_data_list |