{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"},{"sourceId":8462195,"sourceType":"datasetVersion","datasetId":5044537}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"**paper**  \n- 'Direct Molecular Conformation Generation'- Jinhua Zhu, arvix 2022  \n- https://github.com/DirectMolecularConfGen/DMCG\n\nA demo to show how to use open souce conformer generator to generate xyz coords for a given SMILES.  \nI have modified some of the code to make it run in kaggle. \nBut i have not check if the modifications would affects other function (e.g. training, etc).  \nIt may be better for you to use the orginal files from the github repo.  ","metadata":{}},{"cell_type":"code","source":"#!pip install rdkit\n#!pip install torch_geometric\n#!pip install torch-scatter\n#!pip install py3Dmol\n#!pip install torch-sparse","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-20T00:58:20.545879Z","iopub.execute_input":"2024-05-20T00:58:20.546458Z","iopub.status.idle":"2024-05-20T01:06:41.688397Z","shell.execute_reply.started":"2024-05-20T00:58:20.546411Z","shell.execute_reply":"2024-05-20T01:06:41.686421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from rdkit import Chem \nfrom rdkit.Chem import AllChem  \nfrom rdkit.Chem import rdMolAlign\nfrom rdkit.Chem import rdDistGeom\n\nimport sys\nsys.path.append('/kaggle/input/dcgm-confomer-generator-demo')\nfrom confgen.molecule.graph import rdk2graph\nfrom confgen.model.gnn import GNN\n\nimport torch\nfrom torch_geometric.data import Batch, Data\nfrom torch_sparse import SparseTensor\n\nimport re\nimport pickle\nimport py3Dmol\n\nprint('import ok!!!')","metadata":{"execution":{"iopub.status.busy":"2024-05-20T01:13:19.981753Z","iopub.execute_input":"2024-05-20T01:13:19.982349Z","iopub.status.idle":"2024-05-20T01:13:20.047229Z","shell.execute_reply.started":"2024-05-20T01:13:19.982312Z","shell.execute_reply":"2024-05-20T01:13:20.045907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#load model\ncheckpoint_file = '/kaggle/input/dcgm-confomer-generator-demo/checkpoint_94.pt'\nmodel_param = {\n\t'mlp_hidden_size': 1024,\n\t'mlp_layers': 2,\n\t'latent_size': 256,\n\t'use_layer_norm': False,\n\t'num_message_passing_steps': 6,\n\t'global_reducer': 'sum',\n\t'node_reducer': 'sum',\n\t'dropedge_rate': 0.1,\n\t'dropnode_rate': 0.1,\n\t'dropout': 0.1,\n\t'layernorm_before': False,\n\t'encoder_dropout': 0.0,\n\t'use_bn': True,\n\t'vae_beta': 1.0,\n\t'decoder_layers': None,\n\t'reuse_prior': True,\n\t'cycle': 1,\n\t'pred_pos_residual': True,\n\t'node_attn': True,\n\t'global_attn': False,\n\t'shared_decoder': False,\n\t'sample_beta': 1.2,\n\t'shared_output': True\n}\nmodel = GNN(**model_param)\n#print(model)\n\ncheckpoint = torch.load(checkpoint_file, map_location=lambda storage, loc: storage)['model_state_dict']\nprint(model.load_state_dict(checkpoint,strict=False))\nprint('model ok!!!')","metadata":{"execution":{"iopub.status.busy":"2024-05-20T01:08:57.269826Z","iopub.execute_input":"2024-05-20T01:08:57.271195Z","iopub.status.idle":"2024-05-20T01:09:14.279887Z","shell.execute_reply.started":"2024-05-20T01:08:57.271152Z","shell.execute_reply":"2024-05-20T01:09:14.278186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#create graph\nclass CustomData(Data):\n    def __cat_dim__(self, key, value, *args, **kwargs):\n        if isinstance(value, SparseTensor):\n            return (0, 1)\n        elif bool(re.search(\"(index|face|nei_tgt_mask)\", key)):\n            return -1\n        return 0\n\n\n#input : [Dy] has been replaced by C\nkaggle_smiles=[\n\t 'C#CCOc1ccc(CNc2nc(NCC3CCCN3c3cccnn3)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1',\n\t 'C#CCOc1ccc(CNc2nc(NCc3cccc(Br)n3)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1',\n\t 'C#CCOc1ccc(CNc2nc(Nc3cc(C(C)(C)C)[nH]n3)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1',\n\t 'C#CCOc1ccc(CNc2nc(Nc3nc(Cl)nc4cc(OC)c(OC)cc34)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1',\n\t 'C#CCOc1cccc(CNc2nc(NCc3nc4ccccc4s3)nc(N[C@@H](CC#C)CC(=O)NC)n2)c1',\n\t 'C#CC[C@@H](CC(=O)NC)Nc1nc(NCC23CCC(C)(CO2)C3)nc(Nc2cccc(S(=O)(=O)NC(C)(C)C)c2)n1',\n\t 'C#CC[C@@H](CC(=O)NC)Nc1nc(NCc2cccnc2N(C)C)nc(Nc2cnc(Cl)nc2OC)n1',\n\t 'C#CC[C@@](C)(Nc1nc(NCc2ccccc2N2CCOCC2)nc(NCc2ncccc2N(C)C)n1)C(=O)NC',\n\t 'CNC(=O)c1cc(OC)c(Nc2nc(NCCC3CC3)nc(NCc3nnc4c(=O)[nH]ccn34)n2)cc1N',\n\t 'CNC(=O)[C@H](CCCN=[N+]=[N-])Nc1nc(Nc2noc3ccc(F)cc23)nc(Nc2noc3ccc(F)cc23)n1'\n]\n# make graph\nkaggle_mol = [\n\tChem.MolFromSmiles(s) for s in kaggle_smiles\n]\ngraph = []\nfor m in kaggle_mol:\n\tg = rdk2graph(m)\n\tassert len(g[\"edge_attr\"]) == g[\"edge_index\"].shape[1]\n\tassert len(g[\"node_feat\"]) == g[\"num_nodes\"]\n\n\tdata = CustomData()\n\tdata.edge_index = torch.from_numpy(g[\"edge_index\"]).to(torch.int64)\n\tdata.edge_attr = torch.from_numpy(g[\"edge_attr\"]).to(torch.int64)\n\tdata.x = torch.from_numpy(g[\"node_feat\"]).to(torch.int64)\n\tdata.n_nodes = g[\"n_nodes\"]\n\tdata.n_edges = g[\"n_edges\"]\n\tdata.nei_src_index = torch.from_numpy(g[\"nei_src_index\"]).to(torch.int64)\n\tdata.nei_tgt_index = torch.from_numpy(g[\"nei_tgt_index\"]).to(torch.int64)\n\tdata.nei_tgt_mask = torch.from_numpy(g[\"nei_tgt_mask\"]).to(torch.bool)\n\t# data.pos = torch.from_numpy(mol.GetConformer(0).GetPositions()).to(torch.float)\n\t# data.isomorphisms = isomorphic_core(mol)\n\tgraph.append(data)\n\n# collate to a batch and send to network\nbatch_graph = Batch.from_data_list(graph)\nprint(batch_graph)\nprint('data ok!!!')","metadata":{"execution":{"iopub.status.busy":"2024-05-20T01:09:18.578910Z","iopub.execute_input":"2024-05-20T01:09:18.579416Z","iopub.status.idle":"2024-05-20T01:09:18.677947Z","shell.execute_reply.started":"2024-05-20T01:09:18.579383Z","shell.execute_reply":"2024-05-20T01:09:18.676323Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#start conformer generator here !!!!!!!!!!!\n\nnum_mol  = len(kaggle_mol)\nnum_conf = 3 #we generator 3 conformers per mol\ndevice ='cpu' #'cuda:0'\n\nmodel = model.eval()\nif device!='cpu':\n\tmodel.cuda()\n\nbatch_graph = batch_graph.to(device)\n\ndef split_prediction(pred,batch_graph):\n    split = []\n    batch_size = batch_graph.num_graphs\n    n_nodes = batch_graph.n_nodes.tolist()\n    n = 0\n    for i in range(batch_size):\n        p = pred[n: n + n_nodes[i]]\n        n += n_nodes[i]\n        split.append(p)\n    return split\n\n\npredict = [[] for j in range(num_mol)]\nfor i in range(num_conf):\n\twith torch.no_grad():\n\t\tout, _ = model(batch_graph, sample=True)\n\n\tpos = out[-1]  # 3d xyz position\n\tpos = split_prediction(pos, batch_graph)\n\n\tfor j in range(num_mol):\n\t\tpredict[j].append(pos[j])\n\nprint('completed!')\nprint('smiles[0]', kaggle_smiles[0])\nprint('estimated xyz pos:')\nprint('\\tconformer 0:')\nprint(predict[0][0][:5],'...\\n')\nprint('\\tconformer 1:')\nprint(predict[0][1][:5],'...\\n')\nprint('\\tconformer 2:')\nprint(predict[0][2][:5],'...\\n')\n\n#put conformers back into rdkit format\nfor j in range(num_mol):\n\t# make a dummy conformer\n\t# rdDistGeom.EmbedMolecule(mol[j])\n\trdDistGeom.EmbedMultipleConfs(kaggle_mol[j], numConfs=num_conf)\n\n\tfor i in range(num_conf):\n\t\tpos = predict[j][i]\n\t\tfor k in range(len(pos)):\n\t\t\tkaggle_mol[j].GetConformer(i).SetAtomPosition(k, pos[k].tolist())\n\n\npickle_file = 'kaggle_mol.pickle'\nwith open(pickle_file, 'wb') as f:\n\tpickle.dump(kaggle_mol, f, pickle.HIGHEST_PROTOCOL)\n\nprint('infer ok!')","metadata":{"execution":{"iopub.status.busy":"2024-05-20T01:09:23.444332Z","iopub.execute_input":"2024-05-20T01:09:23.444790Z","iopub.status.idle":"2024-05-20T01:09:28.239657Z","shell.execute_reply.started":"2024-05-20T01:09:23.444758Z","shell.execute_reply":"2024-05-20T01:09:28.237870Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#visualization\npickle_file = 'kaggle_mol.pickle'\nwith open(pickle_file,'rb') as f:\n    kaggle_mol = pickle.load(f)\n\n    \nmol = kaggle_mol[7]  \nprint(kaggle_smiles[7])\n\np = py3Dmol.view(width=1200, height=400, viewergrid=(1,3))\n#\nfor j in range(3):\n    p.removeAllModels(viewer=(0,j))\n    p.addModel(Chem.MolToMolBlock(mol, confId=j), 'sdf', viewer=(0,j))\n    p.setStyle({'stick':{}}, viewer=(0,j))\np.zoomTo()\np.show()\n    \n#example of alignment\nrms = rdMolAlign.GetBestRMS(mol,mol,0,1)\nprint(f'rms of conformer0,1, is {rms}')\n        \nAllChem.AlignMol(mol,mol,0,1)\n\np = py3Dmol.view(width=400,height=400)\np.removeAllModels()\np.addModel(Chem.MolToMolBlock(mol,confId=0),'sdf')\np.addModel(Chem.MolToMolBlock(mol,confId=1),'sdf') \np.setStyle({'stick':{}}) \np.zoomTo()\np.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-05-20T01:17:59.046704Z","iopub.execute_input":"2024-05-20T01:17:59.048394Z","iopub.status.idle":"2024-05-20T01:17:59.069646Z","shell.execute_reply.started":"2024-05-20T01:17:59.048336Z","shell.execute_reply":"2024-05-20T01:17:59.068101Z"},"trusted":true},"execution_count":null,"outputs":[]}]}