andy88836's picture
Deploy MOFScreen-Agent FastAPI backend
4d0da28 verified
Raw
History Blame Contribute Delete
4.66 kB
import torch
import numpy as np
from rdkit import Chem
from rdkit.Chem import rdMolDescriptors
from rdkit.Chem.Descriptors import ExactMolWt
from rdkit.ML.Descriptors.CompoundDescriptors import GetAllDescriptorNames
from torch_geometric.data import Data as GraphData
from rdkit.ML.Descriptors.MoleculeDescriptors import MolecularDescriptorCalculator
calculator = {}
def smiles_to_graph(smi):
mol = Chem.MolFromSmiles(smi)
if mol is None:
return None
if not calculator:
chosen_descriptors = ['BertzCT',
'Chi0',
'Chi0n',
'Chi0v',
'Chi1',
'Chi1n',
'Chi1v',
'Chi2n',
'Chi2v',
'Chi3n',
'Chi3v',
'Chi4n',
'Chi4v',
'EState_VSA1',
'EState_VSA10',
'EState_VSA11',
'EState_VSA2',
'EState_VSA3',
'EState_VSA4',
'EState_VSA5',
'EState_VSA6',
'EState_VSA7',
'EState_VSA8',
'EState_VSA9',
'ExactMolWt',
'FpDensityMorgan1',
'FpDensityMorgan2',
'FpDensityMorgan3',
'FractionCSP3',
'HallKierAlpha',
'HeavyAtomCount',
'HeavyAtomMolWt',
'Kappa1',
'Kappa2',
'Kappa3',
'LabuteASA',
'MaxAbsEStateIndex',
'MaxEStateIndex',
'MinAbsEStateIndex',
'MinAbsPartialCharge',
'MinEStateIndex',
'MinPartialCharge',
'MolLogP',
'MolMR',
'MolWt',
'NHOHCount',
'NOCount',
'NumAliphaticCarbocycles',
'NumAliphaticHeterocycles',
'NumAliphaticRings',
'NumAromaticCarbocycles',
'NumAromaticHeterocycles',
'NumAromaticRings',
'NumHAcceptors',
'NumHDonors',
'NumHeteroatoms',
'NumRotatableBonds',
'NumSaturatedCarbocycles',
'NumSaturatedHeterocycles',
'NumSaturatedRings',
'NumValenceElectrons',
'PEOE_VSA1',
'PEOE_VSA10',
'PEOE_VSA11',
'PEOE_VSA12',
'PEOE_VSA13',
'PEOE_VSA14',
'PEOE_VSA2',
'PEOE_VSA3',
'PEOE_VSA4',
'PEOE_VSA5',
'PEOE_VSA6',
'PEOE_VSA7',
'PEOE_VSA8',
'PEOE_VSA9',
'RingCount',
'SMR_VSA1',
'SMR_VSA10',
'SMR_VSA2',
'SMR_VSA3',
'SMR_VSA4',
'SMR_VSA5',
'SMR_VSA6',
'SMR_VSA7',
'SMR_VSA9',
'SlogP_VSA1',
'SlogP_VSA10',
'SlogP_VSA11',
'SlogP_VSA12',
'SlogP_VSA2',
'SlogP_VSA3',
'SlogP_VSA4',
'SlogP_VSA5',
'SlogP_VSA6',
'SlogP_VSA7',
'SlogP_VSA8',
'TPSA',
'VSA_EState1',
'VSA_EState10',
'VSA_EState2',
'VSA_EState3',
'VSA_EState4',
'VSA_EState5',
'VSA_EState6',
'VSA_EState7',
'VSA_EState8',
'VSA_EState9',
'fr_Al_COO',
'fr_Al_OH',
'fr_Al_OH_noTert',
'fr_ArN',
'fr_Ar_COO',
'fr_Ar_N',
'fr_Ar_NH',
'fr_Ar_OH',
'fr_COO',
'fr_COO2',
'fr_C_O',
'fr_C_O_noCOO',
'fr_C_S',
'fr_HOCCN',
'fr_Imine',
'fr_NH0',
'fr_NH1',
'fr_NH2',
'fr_N_O',
'fr_Ndealkylation1',
'fr_Ndealkylation2',
'fr_Nhpyrrole',
'fr_SH',
'fr_aldehyde',
'fr_alkyl_carbamate',
'fr_alkyl_halide',
'fr_allylic_oxid',
'fr_amide',
'fr_amidine',
'fr_aniline',
'fr_aryl_methyl',
'fr_azide',
'fr_azo',
'fr_barbitur',
'fr_benzene',
'fr_benzodiazepine',
'fr_bicyclic',
'fr_dihydropyridine',
'fr_epoxide',
'fr_ester',
'fr_ether',
'fr_furan',
'fr_guanido',
'fr_halogen',
'fr_hdrzine',
'fr_hdrzone',
'fr_imidazole',
'fr_imide',
'fr_isothiocyan',
'fr_ketone',
'fr_ketone_Topliss',
'fr_lactam',
'fr_lactone',
'fr_methoxy',
'fr_morpholine',
'fr_nitrile',
'fr_nitro',
'fr_nitro_arom',
'fr_nitro_arom_nonortho',
'fr_nitroso',
'fr_oxazole',
'fr_oxime',
'fr_para_hydroxylation',
'fr_phenol',
'fr_phenol_noOrthoHbond',
'fr_phos_acid',
'fr_phos_ester',
'fr_piperdine',
'fr_piperzine',
'fr_priamide',
'fr_pyridine',
'fr_quatN',
'fr_sulfide',
'fr_sulfonamd',
'fr_sulfone',
'fr_term_acetylene',
'fr_tetrazole',
'fr_thiazole',
'fr_thiocyan',
'fr_thiophene',
'fr_unbrch_alkane',
'fr_urea',
'qed']
calculator["calculator"] = MolecularDescriptorCalculator(chosen_descriptors)
x = np.array(calculator["calculator"].CalcDescriptors(mol))
x = x.reshape((1,-1))
x = torch.FloatTensor(x)
if torch.isnan(x).any():
return None
graph = GraphData(x=x, edge_index=torch.zeros(2,0, dtype=torch.long), edge_attr=torch.zeros(0,1))
return graph
def test(smi = "C1=CC=C(C=C1)CC(C(=O)O)N"):
mol = Chem.MolFromSmiles(smi)
# print(mol)
graph = smiles_to_graph(smi)
print(graph.keys)
print(graph.x)
print(graph.edge_index)
print( graph.edge_attr )
print( graph.x.shape, graph.edge_index.shape, graph.edge_attr.shape)
print(graph.x.max(), graph.x.min())
for idx in range(mol.GetNumAtoms()):
mol.GetAtomWithIdx(idx).SetProp('molAtomMapNumber', str(mol.GetAtomWithIdx(idx).GetIdx()))
from matplotlib import pyplot as plt; from rdkit.Chem import Draw; plt.imshow(Draw.MolsToGridImage([mol], molsPerRow=2)); plt.show()
if __name__ == "__main__":
smis_list = ["[C@H](C)1CCCO1", "O[C@@H](N)C", "C1=CC=C(C=C1)CC(C(=O)O)N"]
for smi in smis_list:
test(smi)