{"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"}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"for details, refer to https://www.kaggle.com/competitions/leash-BELKA/discussion/498858","metadata":{}},{"cell_type":"code","source":"#!pip install rdkit\n#!pip install torch_geometric\n#!pip install torch-scatter","metadata":{"execution":{"iopub.status.busy":"2024-05-11T16:08:01.411876Z","iopub.execute_input":"2024-05-11T16:08:01.4124Z","iopub.status.idle":"2024-05-11T16:08:01.418209Z","shell.execute_reply.started":"2024-05-11T16:08:01.41236Z","shell.execute_reply":"2024-05-11T16:08:01.416626Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\n\nimport rdkit\nfrom rdkit import Chem\n\nimport torch\nfrom torch_geometric.data import Data\nfrom torch_geometric.loader import DataLoader\n\nfrom torch_geometric.nn import MessagePassing, global_mean_pool\nfrom torch_scatter import scatter\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nprint('import ok!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-11T16:08:01.425787Z","iopub.execute_input":"2024-05-11T16:08:01.427288Z","iopub.status.idle":"2024-05-11T16:08:01.435928Z","shell.execute_reply.started":"2024-05-11T16:08:01.427226Z","shell.execute_reply":"2024-05-11T16:08:01.43484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# helper\n# torch version of np unpackbits\n#https://gist.github.com/vadimkantorov/30ea6d278bc492abf6ad328c6965613a\n\ndef tensor_dim_slice(tensor, dim, dim_slice):\n\treturn tensor[(dim if dim >= 0 else dim + tensor.dim()) * (slice(None),) + (dim_slice,)]\n\n# @torch.jit.script\ndef packshape(shape, dim: int = -1, mask: int = 0b00000001, dtype=torch.uint8, pack=True):\n\tdim = dim if dim >= 0 else dim + len(shape)\n\tbits, nibble = (\n\t\t8 if dtype is torch.uint8 else 16 if dtype is torch.int16 else 32 if dtype is torch.int32 else 64 if dtype is torch.int64 else 0), (\n\t\t1 if mask == 0b00000001 else 2 if mask == 0b00000011 else 4 if mask == 0b00001111 else 8 if mask == 0b11111111 else 0)\n\t# bits = torch.iinfo(dtype).bits # does not JIT compile\n\tassert nibble <= bits and bits % nibble == 0\n\tnibbles = bits // nibble\n\tshape = (shape[:dim] + (int(math.ceil(shape[dim] / nibbles)),) + shape[1 + dim:]) if pack else (\n\t\t\t\tshape[:dim] + (shape[dim] * nibbles,) + shape[1 + dim:])\n\treturn shape, nibbles, nibble\n\n# @torch.jit.script\ndef F_unpackbits(tensor, dim: int = -1, mask: int = 0b00000001, shape=None, out=None, dtype=torch.uint8):\n\tdim = dim if dim >= 0 else dim + tensor.dim()\n\tshape_, nibbles, nibble = packshape(tensor.shape, dim=dim, mask=mask, dtype=tensor.dtype, pack=False)\n\tshape = shape if shape is not None else shape_\n\tout = out if out is not None else torch.empty(shape, device=tensor.device, dtype=dtype)\n\tassert out.shape == shape\n\n\tif shape[dim] % nibbles == 0:\n\t\tshift = torch.arange((nibbles - 1) * nibble, -1, -nibble, dtype=torch.uint8, device=tensor.device)\n\t\tshift = shift.view(nibbles, *((1,) * (tensor.dim() - dim - 1)))\n\t\treturn torch.bitwise_and((tensor.unsqueeze(1 + dim) >> shift).view_as(out), mask, out=out)\n\n\telse:\n\t\tfor i in range(nibbles):\n\t\t\tshift = nibble * i\n\t\t\tsliced_output = tensor_dim_slice(out, dim, slice(i, None, nibbles))\n\t\t\tsliced_input = tensor.narrow(dim, 0, sliced_output.shape[dim])\n\t\t\ttorch.bitwise_and(sliced_input >> shift, mask, out=sliced_output)\n\treturn out\n\nclass dotdict(dict):\n\t__setattr__ = dict.__setitem__\n\t__delattr__ = dict.__delitem__\n\t\n\tdef __getattr__(self, name):\n\t\ttry:\n\t\t\treturn self[name]\n\t\texcept KeyError:\n\t\t\traise AttributeError(name)\n\n            \nprint('helper ok!')","metadata":{"execution":{"iopub.status.busy":"2024-05-11T16:08:01.437878Z","iopub.execute_input":"2024-05-11T16:08:01.438686Z","iopub.status.idle":"2024-05-11T16:08:01.467314Z","shell.execute_reply.started":"2024-05-11T16:08:01.438647Z","shell.execute_reply":"2024-05-11T16:08:01.464879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# mol to graph adopted from\n# from https://github.com/LiZhang30/GPCNDTA/blob/main/utils/DrugGraph.py\n\nPACK_NODE_DIM=9\nPACK_EDGE_DIM=1\nNODE_DIM=PACK_NODE_DIM*8\nEDGE_DIM=PACK_EDGE_DIM*8\n\ndef one_of_k_encoding(x, allowable_set, allow_unk=False):\n\tif x not in allowable_set:\n\t\tif allow_unk:\n\t\t\tx = allowable_set[-1]\n\t\telse:\n\t\t\traise Exception(f'input {x} not in allowable set{allowable_set}!!!')\n\treturn list(map(lambda s: x == s, allowable_set))\n\n\n#Get features of an atom (one-hot encoding:)\n'''\n\t1.atom element: 44+1 dimensions    \n\t2.the atom's hybridization: 5 dimensions\n\t3.degree of atom: 6 dimensions                        \n\t4.total number of H bound to atom: 6 dimensions\n\t5.number of implicit H bound to atom: 6 dimensions    \n\t6.whether the atom is on ring: 1 dimension\n\t7.whether the atom is aromatic: 1 dimension           \n\tTotal: 70 dimensions\n'''\n\nATOM_SYMBOL = [\n\t'C', 'N', 'O', 'S', 'F', 'Si', 'P', 'Cl', 'Br', 'Mg',\n\t'Na', 'Ca', 'Fe', 'As', 'Al', 'I', 'B', 'V', 'K', 'Tl',\n\t'Yb', 'Sb', 'Sn', 'Ag', 'Pd', 'Co', 'Se', 'Ti', 'Zn', 'H',\n\t'Li', 'Ge', 'Cu', 'Au', 'Ni', 'Cd', 'In', 'Mn', 'Zr', 'Cr',\n\t'Pt', 'Hg', 'Pb', 'Dy',\n\t#'Unknown'\n]\n#print('ATOM_SYMBOL', len(ATOM_SYMBOL))44\nHYBRIDIZATION_TYPE = [\n\tChem.rdchem.HybridizationType.S,\n\tChem.rdchem.HybridizationType.SP,\n\tChem.rdchem.HybridizationType.SP2,\n\tChem.rdchem.HybridizationType.SP3,\n\tChem.rdchem.HybridizationType.SP3D\n]\n\ndef get_atom_feature(atom):\n\tfeature = (\n\t\t one_of_k_encoding(atom.GetSymbol(), ATOM_SYMBOL)\n\t   + one_of_k_encoding(atom.GetHybridization(), HYBRIDIZATION_TYPE)\n\t   + one_of_k_encoding(atom.GetDegree(), [0, 1, 2, 3, 4, 5])\n\t   + one_of_k_encoding(atom.GetTotalNumHs(), [0, 1, 2, 3, 4, 5])\n\t   + one_of_k_encoding(atom.GetImplicitValence(), [0, 1, 2, 3, 4, 5])\n\t   + [atom.IsInRing()]\n\t   + [atom.GetIsAromatic()]\n\t)\n\t#feature = np.array(feature, dtype=np.uint8)\n\tfeature = np.packbits(feature)\n\treturn feature\n\n\n#Get features of an edge (one-hot encoding)\n'''\n\t1.single/double/triple/aromatic: 4 dimensions       \n\t2.the atom's hybridization: 1 dimensions\n\t3.whether the bond is on ring: 1 dimension          \n\tTotal: 6 dimensions\n'''\n\ndef get_bond_feature(bond):\n\tbond_type = bond.GetBondType()\n\tfeature = [\n\t\tbond_type == Chem.rdchem.BondType.SINGLE,\n\t\tbond_type == Chem.rdchem.BondType.DOUBLE,\n\t\tbond_type == Chem.rdchem.BondType.TRIPLE,\n\t\tbond_type == Chem.rdchem.BondType.AROMATIC,\n\t\tbond.GetIsConjugated(),\n\t\tbond.IsInRing()\n\t]\n\t#feature = np.array(feature, dtype=np.uint8)\n\tfeature = np.packbits(feature)\n\treturn feature\n\n\ndef smile_to_graph(smiles):\n\tmol = Chem.MolFromSmiles(smiles)\n\tN = mol.GetNumAtoms()\n\tnode_feature = []\n\tedge_feature = []\n\tedge = []\n\tfor i in range(mol.GetNumAtoms()):\n\t\tatom_i = mol.GetAtomWithIdx(i)\n\t\tatom_i_features = get_atom_feature(atom_i)\n\t\tnode_feature.append(atom_i_features)\n\n\t\tfor j in range(mol.GetNumAtoms()):\n\t\t\tbond_ij = mol.GetBondBetweenAtoms(i, j)\n\t\t\tif bond_ij is not None:\n\t\t\t\tedge.append([i, j])\n\t\t\t\tbond_features_ij = get_bond_feature(bond_ij)\n\t\t\t\tedge_feature.append(bond_features_ij)\n\tnode_feature=np.stack(node_feature)\n\tedge_feature=np.stack(edge_feature)\n\tedge = np.array(edge,dtype=np.uint8)\n\treturn N,edge,node_feature,edge_feature\n\ndef to_pyg_format(N,edge,node_feature,edge_feature):\n\tgraph = Data(\n\t\tidx=-1,\n\t\tedge_index = torch.from_numpy(edge.T).int(),\n\t\tx          = torch.from_numpy(node_feature).byte(),\n\t\tedge_attr  = torch.from_numpy(edge_feature).byte(),\n\t)\n\treturn graph\n\n#debug one example\ng = to_pyg_format(*smile_to_graph(smiles=\"C#CCOc1ccc(CNc2nc(NCc3cccc(Br)n3)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1\"))\nprint(g)\nprint('[Dy] is replaced by C !!')\nprint('smile_to_graph() ok!')","metadata":{"execution":{"iopub.status.busy":"2024-05-11T16:08:01.568872Z","iopub.execute_input":"2024-05-11T16:08:01.569366Z","iopub.status.idle":"2024-05-11T16:08:01.604085Z","shell.execute_reply.started":"2024-05-11T16:08:01.56933Z","shell.execute_reply":"2024-05-11T16:08:01.60246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#MODEL: simple MPNNModel\n#from https://github.com/chaitjo/geometric-gnn-dojo/blob/main/geometric_gnn_101.ipynb\n\nDEVICE='cpu'\n\n# i have removed all comments here to jepp it clean. refer to orginal link for code comments\n# of MPNNModel\nclass MPNNLayer(MessagePassing):\n\tdef __init__(self, emb_dim=64, edge_dim=4, aggr='add'):\n\t\tsuper().__init__(aggr=aggr)\n\n\t\tself.emb_dim = emb_dim\n\t\tself.edge_dim = edge_dim\n\t\tself.mlp_msg = nn.Sequential(\n\t\t\tnn.Linear(2 * emb_dim + edge_dim, emb_dim), nn.BatchNorm1d(emb_dim), nn.ReLU(),\n\t\t\tnn.Linear(emb_dim, emb_dim), nn.BatchNorm1d(emb_dim), nn.ReLU()\n\t\t)\n\t\tself.mlp_upd = nn.Sequential(\n\t\t\tnn.Linear(2 * emb_dim, emb_dim), nn.BatchNorm1d(emb_dim), nn.ReLU(),\n\t\t\tnn.Linear(emb_dim, emb_dim), nn.BatchNorm1d(emb_dim), nn.ReLU()\n\t\t)\n\n\tdef forward(self, h, edge_index, edge_attr):\n\t\tout = self.propagate(edge_index, h=h, edge_attr=edge_attr)\n\t\treturn out\n\n\tdef message(self, h_i, h_j, edge_attr):\n\t\tmsg = torch.cat([h_i, h_j, edge_attr], dim=-1)\n\t\treturn self.mlp_msg(msg)\n\n\tdef aggregate(self, inputs, index):\n\t\treturn scatter(inputs, index, dim=self.node_dim, reduce=self.aggr)\n\n\tdef update(self, aggr_out, h):\n\t\tupd_out = torch.cat([h, aggr_out], dim=-1)\n\t\treturn self.mlp_upd(upd_out)\n\n\tdef __repr__(self) -> str:\n\t\treturn (f'{self.__class__.__name__}(emb_dim={self.emb_dim}, aggr={self.aggr})')\n\n\nclass MPNNModel(nn.Module):\n\tdef __init__(self, num_layers=4, emb_dim=64, in_dim=11, edge_dim=4, out_dim=1):\n\t\tsuper().__init__()\n\n\t\tself.lin_in = nn.Linear(in_dim, emb_dim)\n\n\t\t# Stack of MPNN layers\n\t\tself.convs = torch.nn.ModuleList()\n\t\tfor layer in range(num_layers):\n\t\t\tself.convs.append(MPNNLayer(emb_dim, edge_dim, aggr='add'))\n\n\t\tself.pool = global_mean_pool\n\n\tdef forward(self, data): #PyG.Data - batch of PyG graphs\n\n\t\th = self.lin_in(F_unpackbits(data.x,-1).float())  \n\n\t\tfor conv in self.convs:\n\t\t\th = h + conv(h, data.edge_index.long(), F_unpackbits(data.edge_attr,-1).float())  # (n, d) -> (n, d)\n\n\t\th_graph = self.pool(h, data.batch)  \n\t\treturn h_graph\n\n# our prediction model here !!!!\nclass Net(nn.Module):\n\tdef __init__(self, ):\n\t\tsuper().__init__()\n\n\t\tself.output_type = ['infer', 'loss']\n\n\t\tgraph_dim=96\n\t\tself.smile_encoder = MPNNModel(\n\t\t\t in_dim=NODE_DIM, edge_dim=EDGE_DIM, emb_dim=graph_dim, num_layers=4,\n\t\t)\n\t\tself.bind = nn.Sequential(\n\t\t\tnn.Linear(graph_dim, 1024),\n\t\t\t#nn.BatchNorm1d(1024),\n\t\t\tnn.ReLU(inplace=True),\n\t\t\tnn.Dropout(0.1),\n\t\t\tnn.Linear(1024, 1024),\n\t\t\t#nn.BatchNorm1d(1024),\n\t\t\tnn.ReLU(inplace=True),\n\t\t\tnn.Dropout(0.1),\n\t\t\tnn.Linear(1024, 512),\n\t\t\t#nn.BatchNorm1d(512),\n\t\t\tnn.ReLU(inplace=True),\n\t\t\tnn.Dropout(0.1),\n\t\t\tnn.Linear(512, 3),\n\t\t)\n\n\tdef forward(self, batch):\n\t\tgraph = batch['graph']\n\t\tx = self.smile_encoder(graph) \n\t\tbind = self.bind(x)\n\n\t\t# --------------------------\n\t\toutput = {}\n\t\tif 'loss' in self.output_type:\n\t\t\ttarget = batch['bind']\n\t\t\toutput['bce_loss'] = F.binary_cross_entropy_with_logits(bind.float(), target.float())\n\t\tif 'infer' in self.output_type:\n\t\t\toutput['bind'] = torch.sigmoid(bind)\n\n\t\treturn output\n    \n#debug: make some dummy data and run\n\ndef run_check_net():\n\tbatch_size = 3\n\tnode_dim=NODE_DIM\n\tedge_dim=EDGE_DIM\n\n\tdata = []\n\tfor b in range(batch_size):\n\t\tN = np.random.randint(5,10)\n\t\tE = np.random.randint(3,N*(N-1))\n\t\tedge_index = np.stack([\n\t\t\tnp.random.choice(N, E, replace=True),\n\t\t\tnp.random.choice(N, E, replace=True),\n\t\t]).T\n\t\tedge_index = np.sort(edge_index)\n\t\tedge_index = edge_index[edge_index[:, 0].argsort()]\n\t\tedge_index[0] = [0,1] #default\n\t\tedge_index = edge_index[edge_index[:,0]!=edge_index[:,1]]\n\t\tedge_index = np.unique(edge_index, axis=0)\n\n\t\tE = len(edge_index)\n\t\tedge_index = np.ascontiguousarray(edge_index.T)\n\n\t\td = Data(\n\t\t\tidx        = b,\n\t\t\tedge_index = torch.from_numpy(edge_index).int(),\n\t\t\tx          = torch.from_numpy(np.packbits(np.random.choice(2, (N, node_dim)),-1)).byte(),\n\t\t\tedge_attr  = torch.from_numpy(np.packbits(np.random.choice(2, (E, edge_dim)),-1)).byte(),\n\t\t)\n\t\tdata.append(d)\n\n\t#from my_mol2graph import make_dummy_data\n\t#data = make_dummy_data()\n\n\tloader = DataLoader(data, batch_size=batch_size)\n\tgraph = next(iter(loader))\n\tidx = graph.idx.tolist()  #use to index bind array\n\tbatch = dotdict( \n\t\tgraph = graph.to(DEVICE),\n\t\tbind  = torch.from_numpy(np.random.choice(2, (batch_size, 3))).float().to(DEVICE),\n\t)\n\tzz=0\n \n\tnet = Net().to(DEVICE)\n\t#print(net)\n\n\twith torch.no_grad():\n\t\twith torch.cuda.amp.autocast(enabled=True): # dtype=torch.float16):\n\t\t\toutput = net(batch)\n\t\t\t#print(output['bind'])\n\n\t# ---\n\tprint('batch')\n\tfor k, v in batch.items():\n\t\tif k=='idx':\n\t\t\tprint(f'{k:>32} : {len(v)} ')\n\t\telif k=='graph':\n\t\t\tprint(f'{k:>32} : {graph} ')\n\t\telse:\n\t\t\tprint(f'{k:>32} : {v.shape} ')\n\n\tprint('output')\n\tfor k, v in output.items():\n\t\tif 'loss' not in k:\n\t\t\tprint(f'{k:>32} : {v.shape} ')\n\tprint('loss')\n\tfor k, v in output.items():\n\t\tif 'loss' in k:\n\t\t\tprint(f'{k:>32} : {v.item()} ')\n\n            \nrun_check_net()\nprint('model ok!')","metadata":{"execution":{"iopub.status.busy":"2024-05-11T16:08:01.606222Z","iopub.execute_input":"2024-05-11T16:08:01.606965Z","iopub.status.idle":"2024-05-11T16:08:01.712101Z","shell.execute_reply.started":"2024-05-11T16:08:01.6069Z","shell.execute_reply":"2024-05-11T16:08:01.710764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#example of parallel conversion of smiles to graph\n\nfrom multiprocessing import Pool\nfrom tqdm import tqdm\nimport gc\nfrom torch_geometric.loader import DataLoader as PyGDataLoader\n\ndef to_pyg_list(graph):\n\tL = len(graph)\n\tfor i in tqdm(range(L)):\n\t\tN, edge, node_feature, edge_feature = graph[i]\n\t\tgraph[i] = Data(\n\t\t\tidx=i,\n\t\t\tedge_index=torch.from_numpy(edge.T).int(),\n\t\t\tx=torch.from_numpy(node_feature).byte(),\n\t\t\tedge_attr=torch.from_numpy(edge_feature).byte(),\n\t\t)\n\treturn graph\n\n\ntrain_smiles=[ #replace [Dy] with C\n    \"C#CCOc1ccc(CNc2nc(NCc3cccc(Br)n3)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1\",\n    \"C#CCOc1ccc(CNc2nc(NCc3cccc(Br)n3)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1\",\n    \"C#CCOc1ccc(CNc2nc(NCc3cccc(Br)n3)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1\",\n    \"C#CCOc1ccc(CNc2nc(NCc3cccc(Br)n3)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1\",\n    \"C#CCOc1ccc(CNc2nc(NCc3cccc(Br)n3)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1\",\n    \"C#CCOc1ccc(CNc2nc(NCc3cccc(Br)n3)nc(N[C@@H](CC#C)CC(=O)NC)n2)cc1\",\n]\ntrain_bind =np.array([\n    [0,0,0],[1,0,0],[0,1,0],[0,0,1],[1,1,0],[0,0,0],\n])\nnum_train= len(train_smiles)\nwith Pool(processes=64) as pool:\n    train_graph = list(tqdm(pool.imap(smile_to_graph, train_smiles), total=num_train))\n\ntrain_graph = to_pyg_list(train_graph)\ntrain_loader = PyGDataLoader(train_graph, batch_size=3, shuffle=True)\n\n## example training loop\nscaler = torch.cuda.amp.GradScaler(enabled=True)\nnet = Net()\nnet.to(DEVICE)\n\noptimizer =\\\n\ttorch.optim.AdamW(filter(lambda p: p.requires_grad, net.parameters()), lr=0.001)\n\nnum_epoch=10\nepoch=0\niteration=0\nwhile epoch<num_epoch: \n\tfor t, graph_batch in enumerate(train_loader): \n\t\tindex = graph_batch.idx.tolist()\n\t\tB = len(index)\n\t\tbatch = dotdict(\n\t\t\tgraph  = graph_batch.to(DEVICE),\n\t\t\tbind   = torch.from_numpy(train_bind[index]).to(DEVICE),\n\t\t)\n\n\t\tnet.train()\n\t\tnet.output_type = ['loss', 'infer']\n\t\twith torch.cuda.amp.autocast(enabled=True):\n\t\t\toutput = net(batch)  #data_parallel(net,batch) #\n\t\t\tbce_loss = output['bce_loss']\n\n\t\toptimizer.zero_grad() \n\t\tscaler.scale(bce_loss).backward() \n\t\tscaler.step(optimizer)\n\t\tscaler.update()\n\t\t \n\t\ttorch.clear_autocast_cache()\n\t\tprint(epoch,iteration,bce_loss.item())\n\t\titeration +=  1\n        \n\tepoch += 1","metadata":{"execution":{"iopub.status.busy":"2024-05-11T16:08:01.716159Z","iopub.execute_input":"2024-05-11T16:08:01.717077Z","iopub.status.idle":"2024-05-11T16:08:04.32686Z","shell.execute_reply.started":"2024-05-11T16:08:01.717028Z","shell.execute_reply":"2024-05-11T16:08:04.325228Z"},"trusted":true},"execution_count":null,"outputs":[]}]}