{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59094,"databundleVersionId":7010844,"sourceType":"competition"}],"dockerImageVersionId":30558,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"I have constructed a model that converts molecules into graph structures and makes predictions using Graph Attention Networks (GAT).  \nThe cell type is one-hot encoded and combined using a linear layer, but there is also a method of combining them using an embedding layer.  \nThere are still areas for improvement in this code. If you have any good ideas, please let me know.   \nFinally, if you found this code even slightly helpful, I would appreciate it if you could vote for it.","metadata":{}},{"cell_type":"code","source":"!pip install deepchem torch torchvision torch-geometric rdkit lightning","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:28:58.079446Z","iopub.execute_input":"2023-11-16T02:28:58.080225Z","iopub.status.idle":"2023-11-16T02:29:17.345078Z","shell.execute_reply.started":"2023-11-16T02:28:58.080190Z","shell.execute_reply":"2023-11-16T02:29:17.343876Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport tensorflow as tf\nimport numpy as np\nimport torch\nimport matplotlib.pyplot as plt\nimport deepchem as dc\nimport torch_geometric\nfrom torch_geometric.data import Data, DataLoader\nfrom torch_geometric.nn import GCNConv, GATConv, BatchNorm,GATv2Conv\nfrom torch_geometric.utils import to_networkx\nimport torch.nn.functional as F\nfrom torch.utils.data import  random_split\nimport pytorch_lightning as pl\nfrom sklearn.preprocessing import OneHotEncoder\nimport networkx as nx\nimport matplotlib.pyplot as plt\nfrom rdkit import Chem\nfrom rdkit.Chem import Draw\nimport PIL\nimport lightning as L","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:29:17.347430Z","iopub.execute_input":"2023-11-16T02:29:17.347759Z","iopub.status.idle":"2023-11-16T02:29:17.355057Z","shell.execute_reply.started":"2023-11-16T02:29:17.347729Z","shell.execute_reply":"2023-11-16T02:29:17.354161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"SEED = 126\ndef seed_everything(seed):\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True # Fix the network according to random seed\n    print('Finish seeding with seed {}'.format(seed))\n\nseed_everything(SEED)","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:29:17.356265Z","iopub.execute_input":"2023-11-16T02:29:17.356547Z","iopub.status.idle":"2023-11-16T02:29:17.371260Z","shell.execute_reply.started":"2023-11-16T02:29:17.356523Z","shell.execute_reply":"2023-11-16T02:29:17.370365Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Read data ","metadata":{}},{"cell_type":"code","source":"de_train =   pd.read_parquet(\"/kaggle/input/open-problems-single-cell-perturbations/de_train.parquet\")\nid_map = pd.read_csv(\"/kaggle/input/open-problems-single-cell-perturbations/id_map.csv\")\nsample_sub = pd.read_csv(\"/kaggle/input/open-problems-single-cell-perturbations/sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:29:17.372362Z","iopub.execute_input":"2023-11-16T02:29:17.372708Z","iopub.status.idle":"2023-11-16T02:29:22.429797Z","shell.execute_reply.started":"2023-11-16T02:29:17.372675Z","shell.execute_reply":"2023-11-16T02:29:22.428918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocess data ","metadata":{}},{"cell_type":"code","source":"id_map_merge = pd.merge(id_map,de_train[['sm_name','SMILES']].drop_duplicates(['sm_name']),on='sm_name',how='left')\nid_map_merge","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:29:22.433046Z","iopub.execute_input":"2023-11-16T02:29:22.433759Z","iopub.status.idle":"2023-11-16T02:29:22.457184Z","shell.execute_reply.started":"2023-11-16T02:29:22.433728Z","shell.execute_reply":"2023-11-16T02:29:22.456160Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"featurizer = dc.feat.MolGraphConvFeaturizer(use_edges=True)\ngraphs = featurizer.featurize(de_train['SMILES'])\ngraphs_test = featurizer.featurize(id_map_merge['SMILES'])","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:29:22.458485Z","iopub.execute_input":"2023-11-16T02:29:22.458870Z","iopub.status.idle":"2023-11-16T02:29:36.826682Z","shell.execute_reply.started":"2023-11-16T02:29:22.458836Z","shell.execute_reply":"2023-11-16T02:29:36.825864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Function to create a PyTorch Geometric data object from deepchem graph data\ndef convert_to_pyg_graph(dc_graph):\n    node_features = torch.tensor(dc_graph.node_features, dtype=torch.float)\n    edge_index = torch.tensor(dc_graph.edge_index, dtype=torch.long)\n    \n    if 'edge_features' in dc_graph.__dict__:\n        edge_features = torch.tensor(dc_graph.edge_features, dtype=torch.float)\n    else:\n        edge_features = None\n\n    data = Data(x=node_features, edge_index=edge_index, edge_attr=edge_features)\n    return data\n\npyg_graphs = [convert_to_pyg_graph(graph) for graph in graphs]\npyg_graphs_test = [convert_to_pyg_graph(graph) for graph in graphs_test]","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:29:36.828013Z","iopub.execute_input":"2023-11-16T02:29:36.828308Z","iopub.status.idle":"2023-11-16T02:29:36.894904Z","shell.execute_reply.started":"2023-11-16T02:29:36.828286Z","shell.execute_reply":"2023-11-16T02:29:36.894143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder = OneHotEncoder(sparse=False)\nencoder.fit(de_train[['cell_type']])\ncell_type_encoded = encoder.transform(de_train[['cell_type']])\ncell_type_encoded_test = encoder.transform(id_map_merge[['cell_type']])\ncell_type_encoded","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:29:36.896032Z","iopub.execute_input":"2023-11-16T02:29:36.896316Z","iopub.status.idle":"2023-11-16T02:29:36.910783Z","shell.execute_reply.started":"2023-11-16T02:29:36.896291Z","shell.execute_reply":"2023-11-16T02:29:36.909813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Dataset","metadata":{}},{"cell_type":"code","source":"class GeneExpressionDataset(torch.utils.data.Dataset):\n    def __init__(self, graphs, cell_types, expressions = None):\n        self.graphs = graphs\n        self.cell_types = cell_types\n        self.expressions = expressions\n\n    def __len__(self):\n        return len(self.graphs)\n\n    def __getitem__(self, idx):\n        if self.expressions is None:\n            return self.graphs[idx], torch.tensor(self.cell_types[idx], dtype=torch.float)\n        else:\n            return self.graphs[idx], torch.tensor(self.cell_types[idx], dtype=torch.float), torch.tensor(self.expressions[idx], dtype=torch.float)\n\nbatch_size = 32\nuse_val = False\n\ndataset = GeneExpressionDataset(pyg_graphs, cell_type_encoded, de_train.iloc[:, 5:].values)\ndataset_test = GeneExpressionDataset(pyg_graphs_test, cell_type_encoded)\n\ntrain_loader = DataLoader(dataset, batch_size=32, shuffle=True,)\ntest_loader = DataLoader(dataset_test, batch_size=len(dataset_test))\n\nif use_val:\n    total_size = len(dataset)\n    train_size = int(total_size * 0.9)\n    val_size = total_size - train_size\n    train_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n    print(f'train_size:{len(train_dataset)}\\nval_size:{len(val_dataset)}')\n\n   \n    train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=False,num_workers=4)\n    val_loader = DataLoader(val_dataset, batch_size=batch_size,num_workers=4,shuffle=False)\nelse:\n    train_loader = DataLoader(dataset, batch_size=32, shuffle=True,num_workers=4)\n    test_loader = DataLoader(dataset_test, batch_size=len(dataset_test),num_workers=4)\n","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:29:36.911959Z","iopub.execute_input":"2023-11-16T02:29:36.912245Z","iopub.status.idle":"2023-11-16T02:29:36.959429Z","shell.execute_reply.started":"2023-11-16T02:29:36.912221Z","shell.execute_reply":"2023-11-16T02:29:36.958423Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Draw the graph structure","metadata":{}},{"cell_type":"code","source":"struc_num = 3\ndata_list = [(dataset[i][0], de_train['SMILES'][i]) for i in np.random.randint(0,len(dataset),struc_num)]\n\n\nfig, axes = plt.subplots(struc_num, 2, figsize=(10, struc_num * 3))\n\nfor idx, (data, smiles) in enumerate(data_list):\n    G = nx.Graph()\n    for i in range(data.num_nodes):\n        G.add_node(i)\n    for i, j in data.edge_index.t().tolist():\n        G.add_edge(i, j)\n    pos = nx.spring_layout(G)  \n    nx.draw(G, pos, ax=axes[idx, 0], with_labels=True, node_color='lightblue', edge_color='gray')\n    \n    mol = Chem.MolFromSmiles(smiles)\n    img = Draw.MolToImage(mol, size=(300, 300))\n    axes[idx, 1].imshow(PIL.ImageOps.expand(img, border=10, fill='white'))\n    axes[idx, 1].axis('off') \n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:29:36.960872Z","iopub.execute_input":"2023-11-16T02:29:36.961239Z","iopub.status.idle":"2023-11-16T02:29:37.702395Z","shell.execute_reply.started":"2023-11-16T02:29:36.961205Z","shell.execute_reply":"2023-11-16T02:29:37.701355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"from pytorch_lightning.loggers import TensorBoardLogger\nclass GATModel(L.LightningModule):\n    def __init__(self, num_node_features, num_classes, num_cell_types):\n        super().__init__()\n        self.gat1 = GATConv(num_node_features, 128)\n        self.bn1 = BatchNorm(128)\n        self.gat2 = GATConv(128, 256)\n        self.bn2 = BatchNorm(256)\n        self.fc1 = torch.nn.Linear(256 + num_cell_types, 1024)\n        self.fc2 = torch.nn.Linear(1024,4096)\n        self.fc3 = torch.nn.Linear(4096, num_classes)\n        \n        self.dropout = torch.nn.Dropout(0.1)\n        self.val_step_outputs = []\n        self.val_step_labels = []\n\n    def forward(self, data, cell_type):\n        x, edge_index, batch = data.x, data.edge_index, data.batch\n        x = self.gat1(x, edge_index)\n        x = self.bn1(x)\n        x = torch.relu(x)\n        x = self.dropout(x)\n        x = self.gat2(x, edge_index)\n        x = self.bn2(x)\n        x = torch.relu(x)\n        x = torch_geometric.nn.global_mean_pool(x, batch)\n        x = torch.cat((x, cell_type), dim=1)\n        x = self.fc1(x)\n        x = torch.relu(x)\n        x = self.fc2(x)\n        x = torch.relu(x)\n        x = self.fc3(x)\n        return x\n\n    def training_step(self, batch, batch_idx):\n        graphs, cell_types, expressions = batch\n        preds = self(graphs, cell_types)\n        loss = torch.sqrt(torch.nn.functional.mse_loss(preds, expressions))\n        batch_size = graphs.num_graphs\n        self.log('train_loss', loss, on_step=False, on_epoch=True, prog_bar=True,batch_size=batch_size)\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        graphs, cell_types, expressions = batch\n        preds = self(graphs, cell_types)\n        val_loss = torch.sqrt(torch.nn.functional.mse_loss(preds, expressions))\n        batch_size = graphs.num_graphs\n        self.log('val_loss', val_loss, on_step=False, on_epoch=True, prog_bar=True, batch_size=batch_size)\n        self.val_step_outputs.append(preds)\n        self.val_step_labels.append(expressions)\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=0.001)\n        return optimizer\n    \n    def cal_mrrmse(self, labels, preds):\n        labels = labels.numpy()\n        preds = preds.numpy()\n        score = labels - preds\n        score = np.mean(np.sqrt(np.mean(score ** 2, axis=1)))\n        return score\n    \n    def on_validation_epoch_end(self):\n        all_preds = torch.cat(self.val_step_outputs)\n        all_labels = torch.cat(self.val_step_labels)\n        score = self.cal_mrrmse(all_labels, all_preds)\n        print(f'val_MRRMSE:{score}')\n        self.val_step_labels.clear()\n        self.val_step_outputs.clear()\n\n\nnum_node_features = dataset[0][0].num_node_features\nnum_classes = de_train.iloc[:, 5:].shape[1]\nnum_cell_types = cell_type_encoded.shape[1]\nmodel = GATModel(num_node_features, num_classes, num_cell_types)\nlogger = TensorBoardLogger(\"tb_logs\", name=\"my_model\")\n\n\ntrainer = L.Trainer(max_epochs=15,logger=logger,log_every_n_steps=4)\nif not use_val:\n    trainer.fit(model, train_loader)\nelse:\n    trainer.fit(model,train_loader,val_dataloaders=val_loader,)","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:29:37.703892Z","iopub.execute_input":"2023-11-16T02:29:37.704180Z","iopub.status.idle":"2023-11-16T02:30:34.202329Z","shell.execute_reply.started":"2023-11-16T02:29:37.704155Z","shell.execute_reply":"2023-11-16T02:30:34.201391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"# Prediction","metadata":{}},{"cell_type":"code","source":"model.eval()\npredictions = []\nwith torch.no_grad():\n    for batch in test_loader:\n        preds = model(*batch)\n        predictions.append(preds)","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:30:34.203957Z","iopub.execute_input":"2023-11-16T02:30:34.204328Z","iopub.status.idle":"2023-11-16T02:30:35.118948Z","shell.execute_reply.started":"2023-11-16T02:30:34.204282Z","shell.execute_reply":"2023-11-16T02:30:35.117671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predict = predictions[0]\npredict","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:30:35.120854Z","iopub.execute_input":"2023-11-16T02:30:35.121987Z","iopub.status.idle":"2023-11-16T02:30:35.131695Z","shell.execute_reply.started":"2023-11-16T02:30:35.121940Z","shell.execute_reply":"2023-11-16T02:30:35.130384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if predict.is_cuda:\n    predict = predict.cpu()\nsample_sub.iloc[:,1:] = predict.numpy()","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:30:35.136323Z","iopub.execute_input":"2023-11-16T02:30:35.136670Z","iopub.status.idle":"2023-11-16T02:30:35.512079Z","shell.execute_reply.started":"2023-11-16T02:30:35.136621Z","shell.execute_reply":"2023-11-16T02:30:35.511008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_sub","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:30:35.513357Z","iopub.execute_input":"2023-11-16T02:30:35.513754Z","iopub.status.idle":"2023-11-16T02:30:35.548772Z","shell.execute_reply.started":"2023-11-16T02:30:35.513727Z","shell.execute_reply":"2023-11-16T02:30:35.547859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_sub.to_csv('submission1.csv',index=False)","metadata":{"execution":{"iopub.status.busy":"2023-11-16T02:30:35.550017Z","iopub.execute_input":"2023-11-16T02:30:35.550588Z","iopub.status.idle":"2023-11-16T02:30:46.862238Z","shell.execute_reply.started":"2023-11-16T02:30:35.550562Z","shell.execute_reply":"2023-11-16T02:30:46.861327Z"},"trusted":true},"execution_count":null,"outputs":[]}]}