{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport time\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-21T23:09:07.169076Z","iopub.execute_input":"2024-05-21T23:09:07.169411Z","iopub.status.idle":"2024-05-21T23:09:08.910273Z","shell.execute_reply.started":"2024-05-21T23:09:07.169382Z","shell.execute_reply":"2024-05-21T23:09:08.909350Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### install stuff\n!pip install rdkit\n!pip install duckdb\n!pip install deepchem","metadata":{"execution":{"iopub.status.busy":"2024-05-21T23:09:08.912068Z","iopub.execute_input":"2024-05-21T23:09:08.912624Z","iopub.status.idle":"2024-05-21T23:09:54.910304Z","shell.execute_reply.started":"2024-05-21T23:09:08.912584Z","shell.execute_reply":"2024-05-21T23:09:54.909312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### import and install stuff\n\nimport os\nimport torch\nimport numpy as np\nimport duckdb\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport deepchem\nimport rdkit.Chem as Chem\nfrom rdkit.Chem import Draw\nfrom deepchem.utils.typing import RDKitMol\nfrom deepchem.feat.base_classes import MolecularFeaturizer\nfrom deepchem.feat import MolGraphConvFeaturizer as MGCF\nfrom sklearn.model_selection import train_test_split\n\n\nos.environ['TORCH'] = torch.__version__\nprint(torch.__version__)\n\n!pip install -q torch-scatter -f https://data.pyg.org/whl/torch-${TORCH}.html\n!pip install -q torch-sparse -f https://data.pyg.org/whl/torch-${TORCH}.html\n!pip install -q git+https://github.com/pyg-team/pytorch_geometric.git\n    \n!pip install torch-geometric\nimport torch_geometric\nimport torch_geometric.nn as geom_nn\nimport torch_geometric.data as geom_data","metadata":{"execution":{"iopub.status.busy":"2024-05-21T23:09:54.911771Z","iopub.execute_input":"2024-05-21T23:09:54.912073Z","iopub.status.idle":"2024-05-21T23:37:38.296390Z","shell.execute_reply.started":"2024-05-21T23:09:54.912044Z","shell.execute_reply":"2024-05-21T23:37:38.295561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path = '/kaggle/input/leash-BELKA/train.parquet'\n\ntest_path = '/kaggle/input/leash-BELKA/test.parquet'\n\ncon = duckdb.connect()\n\ndf_train = con.query(f\"\"\"(SELECT *\n                        FROM parquet_scan('{train_path}')\n                        WHERE binds = 0\n                        ORDER BY random()\n                        LIMIT 10000)\n                        UNION ALL\n                        (SELECT *\n                        FROM parquet_scan('{train_path}')\n                        WHERE binds = 1\n                        ORDER BY random()\n                        LIMIT 10000)\"\"\").df()\n\ndf_test = con.query(f\"\"\"(SELECT *\n                        FROM parquet_scan('{test_path}')\n                        ORDER BY random()\n                        LIMIT 5000)\n                        UNION ALL\n                        (SELECT *\n                        FROM parquet_scan('{test_path}')\n                        ORDER BY random()\n                        LIMIT 5000)\"\"\").df()\n\ncon.close()","metadata":{"execution":{"iopub.status.busy":"2024-05-21T23:37:38.298549Z","iopub.execute_input":"2024-05-21T23:37:38.299067Z","iopub.status.idle":"2024-05-21T23:38:22.649173Z","shell.execute_reply.started":"2024-05-21T23:37:38.299041Z","shell.execute_reply":"2024-05-21T23:38:22.648331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df_train = pd.read_csv('/kaggle/input/leash-BELKA/train.csv') #too slow","metadata":{"execution":{"iopub.status.busy":"2024-05-21T23:38:22.650761Z","iopub.execute_input":"2024-05-21T23:38:22.651120Z","iopub.status.idle":"2024-05-21T23:38:22.655014Z","shell.execute_reply.started":"2024-05-21T23:38:22.651088Z","shell.execute_reply":"2024-05-21T23:38:22.654109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = df_train[df_train['protein_name'] == 'BRD4']\ndf_train, df_val = train_test_split(df_train, test_size=0.2, random_state=42)\ndisplay(df_train.head())\ndisplay(df_val.head())","metadata":{"execution":{"iopub.status.busy":"2024-05-21T23:38:22.656549Z","iopub.execute_input":"2024-05-21T23:38:22.656964Z","iopub.status.idle":"2024-05-21T23:38:22.747353Z","shell.execute_reply.started":"2024-05-21T23:38:22.656914Z","shell.execute_reply":"2024-05-21T23:38:22.746384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display(df_test.head())","metadata":{"execution":{"iopub.status.busy":"2024-05-21T23:38:22.748403Z","iopub.execute_input":"2024-05-21T23:38:22.748657Z","iopub.status.idle":"2024-05-21T23:38:22.760233Z","shell.execute_reply.started":"2024-05-21T23:38:22.748634Z","shell.execute_reply":"2024-05-21T23:38:22.759152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### demo for using deepchem featurizer\ndemo_smiles = df_train['molecule_smiles'].values[1]\n\nfeat_graph = MGCF(use_edges=True).featurize([demo_smiles])[0] # node_features=[21, 30], edge_index=[2, 46], edge_features=[46, 11], pos=[0]\n\nx = torch.tensor(feat_graph.node_features, dtype=torch.float32) # node input features from deepchem\nedge_index = torch.tensor(feat_graph.edge_index, dtype=torch.long)\nedge_attr = torch.tensor(feat_graph.edge_features, dtype=torch.long)\n# pos = feat_graph.pos\n# y = logD.values[0] # logD is our label\nprint(x)\nprint(edge_index)\nprint(edge_attr)\nprint(feat_graph)","metadata":{"execution":{"iopub.status.busy":"2024-05-21T23:38:22.761732Z","iopub.execute_input":"2024-05-21T23:38:22.762105Z","iopub.status.idle":"2024-05-21T23:38:22.948267Z","shell.execute_reply.started":"2024-05-21T23:38:22.762071Z","shell.execute_reply":"2024-05-21T23:38:22.947310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### dataset\n\n\nclass deepchem_Dataset(torch_geometric.data.Dataset):\n    '''\n    Pytorch gemmetric Dataset, takes pandas dataframe and turn smiles into graph as input and bind as target\n    '''\n    def __init__(self, transform=None, train=True, val=False, test=False):\n        super().__init__()\n        self.transform = transform\n        self.train = train\n        self.test = test\n        if train:\n          self.df = df_train\n        else:\n          self.df = df_val\n\n\n    def len(self):\n        self.filelength = len(self.df)\n        return self.filelength\n\n    def get(self, index):\n\n        smiles = self.df['molecule_smiles'].iloc[[index]]\n        binds = self.df['binds'].iloc[[index]]\n\n\n\n        feat_graph = MGCF(use_edges=True).featurize([smiles.values[0]])[0] # node_features=[21, 30], edge_index=[2, 46], edge_features=[46, 11], pos=[0]\n\n        x = torch.tensor(feat_graph.node_features, dtype=torch.float32) # node input features from deepchem\n        edge_index = torch.tensor(feat_graph.edge_index, dtype=torch.long)\n        edge_attr = torch.tensor(feat_graph.edge_features, dtype=torch.long)\n        y = binds.values[0] # binds is our label\n        y = torch.tensor(y)\n\n\n        data = Data(x=x, edge_index=edge_index, edge_attr=edge_attr, y=y) # create pyg data\n#         print(smiles.values[0], binds.values[0])\n        return data #, (smiles.values[0], binds.values[0])","metadata":{"execution":{"iopub.status.busy":"2024-05-21T23:47:10.872868Z","iopub.execute_input":"2024-05-21T23:47:10.873626Z","iopub.status.idle":"2024-05-21T23:47:10.883874Z","shell.execute_reply.started":"2024-05-21T23:47:10.873589Z","shell.execute_reply":"2024-05-21T23:47:10.882886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn.functional as F\nfrom torch_geometric.nn import GATConv\nfrom torch_geometric.datasets import QM9\nfrom torch_geometric.data import Data\nfrom torch_geometric.loader import DataLoader\nfrom torch_geometric.nn import global_mean_pool\n\nglobal_batch_size = 32\n# train_loader = DataLoader(train_dataset, batch_size=global_batch_size, shuffle=True)\n# val_loader = DataLoader(val_dataset, batch_size=global_batch_size, shuffle=False)\n# test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)\n\n\n# train_loader = torch.utils.data.DataLoader(Pyg_Dataset(train=True),batch_size = 64, shuffle = True)\n# val_loader = DataLoader(Pyg_Dataset(train=False),batch_size = 64, shuffle = True)\n# test_loader = DataLoader(Pyg_Dataset(train=False),batch_size = 64, shuffle = True)\n\n\n\n# train_loader = DataLoader(Pyg_Dataset(train=True), batch_size = global_batch_size, shuffle = True)\n# val_loader = DataLoader(Pyg_Dataset(train=False), batch_size = global_batch_size, shuffle = True)\n# test_loader = DataLoader(Pyg_Dataset(train=False), batch_size = global_batch_size, shuffle = True)\n\ntrain_loader = DataLoader(deepchem_Dataset(train=True), batch_size = global_batch_size, shuffle = True)\nval_loader = DataLoader(deepchem_Dataset(train=False), batch_size = global_batch_size, shuffle = True)\ntest_loader = DataLoader(deepchem_Dataset(train=False), batch_size = global_batch_size, shuffle = True)\n\n\n\n# Step 2: Define the GAT model\nclass GATNet(torch.nn.Module):\n    def __init__(self):\n        super(GATNet, self).__init__()\n        self.input_dim = 30 #dataset.num_features\n        self.hidden_dim = 32\n        self.head1 = 12\n        self.head2 = 12\n        self.dropout_rate = 0.1\n\n        self.conv1 = GATConv(self.input_dim, self.hidden_dim, heads=self.head1)\n        self.conv2 = GATConv(self.hidden_dim * self.head1, self.hidden_dim, heads=self.head2)\n        self.fc = torch.nn.Linear(self.hidden_dim * self.head2, 1)  # Predicting a single property\n\n\n    def forward(self, data):\n        x, edge_index = data.x, data.edge_index\n        batch = data.batch if hasattr(data, 'batch') else None\n#         print('x', x.shape)\n\n        x = F.dropout(x, p=self.dropout_rate, training=self.training)\n        x = F.elu(self.conv1(x, edge_index))\n        x = F.dropout(x, p=self.dropout_rate, training=self.training)\n        x = F.elu(self.conv2(x, edge_index))\n\n        # x = torch.mean(x, dim=0)  # Global average pooling\n        x = global_mean_pool(x, batch)  # [batch_size, hidden_channels]\n        x = self.fc(x)\n        return x\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = GATNet().to(device)\noptimizer = torch.optim.Adam(model.parameters(), lr=0.004)\n# criterion = torch.nn.MSELoss()\ncriterion = torch.nn.BCEWithLogitsLoss()\n\n# Step 3: Train the model\n# Step 3: Train the model\ndef train():\n    model.train()\n\n    total_loss = 0\n    for data in train_loader:\n        data = data.to(device)\n#         print(data)\n        optimizer.zero_grad()\n        pred = model(data)\n        loss = criterion(pred, data.y.unsqueeze(1).float())\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item() * data.num_graphs\n    return total_loss / len(train_loader.dataset)\n\n# Step 4: Evaluate the model\ndef evaluate(loader):\n    model.eval()\n\n    total_loss = 0\n    with torch.no_grad():\n        for data in loader:\n            data = data.to(device)\n            pred = model(data)\n            loss = criterion(pred, data.y.unsqueeze(1).float())\n            total_loss += loss.item() * data.num_graphs\n    return total_loss / len(loader.dataset)\n\n# Training loop\nt1 = time.time()\nfor epoch in range(1, 10):\n    loss = train()\n    val_loss = evaluate(val_loader)\n    print(f'Epoch: {epoch:03d}, Loss: {loss:.4f}, Val Loss: {val_loss:.4f}')\nt2 = time.time()\nprint('training time', t2 - t1)\n# Test performance\ntest_loss = evaluate(test_loader)\nprint(f'Test Loss: {test_loss:.4f}')","metadata":{"execution":{"iopub.status.busy":"2024-05-21T23:47:28.077238Z","iopub.execute_input":"2024-05-21T23:47:28.077606Z","iopub.status.idle":"2024-05-22T00:08:40.901144Z","shell.execute_reply.started":"2024-05-21T23:47:28.077577Z","shell.execute_reply":"2024-05-22T00:08:40.900215Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### visualize the result and analysis (as first step)\n\nfrom sklearn.metrics import r2_score, mean_squared_error\nimport matplotlib.pyplot as plt\ntest_loader = DataLoader(deepchem_Dataset(train=False), batch_size = 1, shuffle = True)\n\ndef get_evaluate(loader):\n    target_list = []\n    pred_list = []\n    model.eval()\n    for data in loader:\n        data = data.to(device)\n        pred = model(data)\n        loss = criterion(pred, data.y.unsqueeze(1).float())\n        pred = F.sigmoid(pred)\n        pred_list.append(pred.cpu().detach().numpy()[0][0])\n        target_list.append(data.y.cpu().detach().numpy()[0])\n        # total_loss += loss.item() * data.num_graphs\n    return pred_list, target_list\n\npred_list, target_list = get_evaluate(test_loader)\n\n\nscale= np.linspace(-2, 5 ,100)\nplt.figure(figsize=(8, 8))\nplt.plot(scale, scale, \"r\")\nplt.plot(target_list, pred_list, 'o')\nplt.xlabel('target')\nplt.ylabel('pred')\nplt.show()\n\n\nr2_score = r2_score(target_list, pred_list)\nrmse = mean_squared_error(target_list, pred_list)\nprint(\"r2 score\", r2_score)\nprint(\"rmse\", rmse) ### the model has only been trained for 20 epochs at the time, so just checking whether the the training script and hyperparameters work\n\n\n###\ntorch.save(model, 'GAT_v0.mdl') ### save the model first, and will further train this one with more epochs","metadata":{"execution":{"iopub.status.busy":"2024-05-22T00:13:15.688528Z","iopub.execute_input":"2024-05-22T00:13:15.688898Z","iopub.status.idle":"2024-05-22T00:13:48.127506Z","shell.execute_reply.started":"2024-05-22T00:13:15.688869Z","shell.execute_reply":"2024-05-22T00:13:48.126390Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def binder(val):\n    if val >= 0.5:\n        return \"bind\"\n    else:\n        return \"no_bind\"","metadata":{"execution":{"iopub.status.busy":"2024-05-22T00:19:40.481281Z","iopub.execute_input":"2024-05-22T00:19:40.482020Z","iopub.status.idle":"2024-05-22T00:19:40.486659Z","shell.execute_reply.started":"2024-05-22T00:19:40.481984Z","shell.execute_reply":"2024-05-22T00:19:40.485793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for idx, data in enumerate(test_loader):\n    if idx > 10:\n        break\n    data = data.to(device)\n    pred = model(data)\n    loss = criterion(pred, data.y.unsqueeze(1).float())\n    pred = F.sigmoid(pred)\n#     print(smile, bind, pred)\n#     pred_list.append(pred.cpu().detach().numpy()[0][0])\n#     target_list.append(data.y.cpu().detach().numpy()[0])","metadata":{"execution":{"iopub.status.busy":"2024-05-22T00:19:56.832366Z","iopub.execute_input":"2024-05-22T00:19:56.832773Z","iopub.status.idle":"2024-05-22T00:19:57.156035Z","shell.execute_reply.started":"2024-05-22T00:19:56.832733Z","shell.execute_reply":"2024-05-22T00:19:57.155161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn import svm, datasets\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix\nfrom sklearn.utils.multiclass import unique_labels\n\ndef plot_confusion_matrix(y_true, y_pred, classes,\n                          normalize=False,\n                          title=None,\n                          cmap=plt.cm.Blues):\n    \"\"\"\n    This function prints and plots the confusion matrix.\n    Normalization can be applied by setting `normalize=True`.\n    \"\"\"\n    if not title:\n        if normalize:\n            title = 'Normalized confusion matrix'\n        else:\n            title = 'Confusion matrix, without normalization'\n\n    # Compute confusion matrix\n    cm = confusion_matrix(y_true, y_pred)\n    # Only use the labels that appear in the data\n    classes = classes[unique_labels(y_true, y_pred)]\n    print(classes)\n    if normalize:\n        cm = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]\n        print(\"Normalized confusion matrix\")\n    else:\n        print('Confusion matrix, without normalization')\n\n    #print(cm)\n\n    fig, ax = plt.subplots()\n    im = ax.imshow(cm, interpolation='nearest', cmap=cmap)\n    ax.figure.colorbar(im, ax=ax)\n    # We want to show all ticks...\n    print(np.arange(0, cm.shape[1]))\n    ax.set(xticks=np.arange(0, cm.shape[1]),\n           yticks=np.arange(0, cm.shape[0]),\n           # ... and label them with the respective list entries\n           xticklabels=classes, yticklabels=classes,\n           title=title,\n           ylabel='True label',\n           xlabel='Predicted label')\n\n    # Rotate the tick labels and set their alignment.\n    plt.setp(ax.get_xticklabels(), rotation=45, ha=\"right\",\n             rotation_mode=\"anchor\")\n\n    # Loop over data dimensions and create text annotations.\n    fmt = '.2f' if normalize else 'd'\n    thresh = cm.max() / 2.\n    for i in range(cm.shape[0]):\n        for j in range(cm.shape[1]):\n            ax.text(j, i, format(cm[i, j], fmt),\n                    ha=\"center\", va=\"center\",\n                    color=\"white\" if cm[i, j] > thresh else \"black\")\n    fig.tight_layout()\n    \n    return ax\n\n\nnp.set_printoptions(precision=2)\n\n\n\nmad_sane_classes = np.array(['bind', 'no-bind'])\ntarget_int_list = np.array(list(map(round, target_list)), dtype=np.int8)\npred_int_list = np.array(list(map(round, pred_list)), dtype=np.int8)\n\n# print(target_int_list)\nplot_confusion_matrix(target_int_list, pred_int_list, classes=mad_sane_classes,\n                      title='graph neural network baseline, bind/no-bind prediction from validation BRD4)')","metadata":{"execution":{"iopub.status.busy":"2024-05-22T00:20:01.364475Z","iopub.execute_input":"2024-05-22T00:20:01.364851Z","iopub.status.idle":"2024-05-22T00:20:02.089877Z","shell.execute_reply.started":"2024-05-22T00:20:01.364820Z","shell.execute_reply":"2024-05-22T00:20:02.088938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class submission_Dataset(torch_geometric.data.Dataset):\n    '''\n    Pytorch gemmetric Dataset, takes pandas dataframe and turn smiles into graph as input and bind as target\n    '''\n    def __init__(self, transform=None, train=True, val=False, test=False):\n        super().__init__()\n        self.transform = transform\n        self.train = train\n        self.test = test\n        if train:\n          self.df = df_train\n        else:\n          self.df = df_test\n\n\n    def len(self):\n        self.filelength = len(self.df)\n        return self.filelength\n\n    def get(self, index):\n\n        smiles = self.df['molecule_smiles'].iloc[[index]]\n#         binds = self.df['binds'].iloc[[index]]\n\n\n\n        feat_graph = MGCF(use_edges=True).featurize([smiles.values[0]])[0] # node_features=[21, 30], edge_index=[2, 46], edge_features=[46, 11], pos=[0]\n\n        x = torch.tensor(feat_graph.node_features, dtype=torch.float32) # node input features from deepchem\n        edge_index = torch.tensor(feat_graph.edge_index, dtype=torch.long)\n        edge_attr = torch.tensor(feat_graph.edge_features, dtype=torch.long)\n#         y = binds.values[0] # binds is our label\n#         y = torch.tensor(y)\n\n\n        data = Data(x=x, edge_index=edge_index, edge_attr=edge_attr) # create pyg data\n#         print(smiles.values[0], binds.values[0])\n        return data #, (smiles.values[0], binds.values[0])","metadata":{"execution":{"iopub.status.busy":"2024-05-22T00:22:58.339212Z","iopub.execute_input":"2024-05-22T00:22:58.340147Z","iopub.status.idle":"2024-05-22T00:22:58.349452Z","shell.execute_reply.started":"2024-05-22T00:22:58.340097Z","shell.execute_reply":"2024-05-22T00:22:58.348581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_loader = DataLoader(submission_Dataset(train=False), batch_size = 1, shuffle = True)\npred_list = []\nfor idx, data in enumerate(test_loader):\n#     if idx > 10:\n#         break\n    data = data.to(device)\n    pred = model(data)\n#     loss = criterion(pred, data.y.unsqueeze(1).float())\n    pred = F.sigmoid(pred)\n    print(pred)\n    pred_list.append(pred[0][0].detach().cpu().numpy())\n","metadata":{"execution":{"iopub.status.busy":"2024-05-22T00:25:38.421422Z","iopub.execute_input":"2024-05-22T00:25:38.421808Z","iopub.status.idle":"2024-05-22T00:29:57.331140Z","shell.execute_reply.started":"2024-05-22T00:25:38.421778Z","shell.execute_reply":"2024-05-22T00:29:57.330138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test['pred'] = pred_list\ndf_test.to_csv('test_prediction.csv')","metadata":{"execution":{"iopub.status.busy":"2024-05-22T00:29:57.333393Z","iopub.execute_input":"2024-05-22T00:29:57.334074Z","iopub.status.idle":"2024-05-22T00:29:57.487282Z","shell.execute_reply.started":"2024-05-22T00:29:57.334032Z","shell.execute_reply":"2024-05-22T00:29:57.486535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test","metadata":{"execution":{"iopub.status.busy":"2024-05-22T00:33:51.032127Z","iopub.execute_input":"2024-05-22T00:33:51.032965Z","iopub.status.idle":"2024-05-22T00:33:51.057138Z","shell.execute_reply.started":"2024-05-22T00:33:51.032931Z","shell.execute_reply":"2024-05-22T00:33:51.056176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}