{"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":"gpu","dataSources":[{"sourceId":67356,"databundleVersionId":8006601,"sourceType":"competition"}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"pip install rdkit","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:02:01.145062Z","iopub.execute_input":"2024-07-10T04:02:01.145657Z","iopub.status.idle":"2024-07-10T04:02:14.498393Z","shell.execute_reply.started":"2024-07-10T04:02:01.145629Z","shell.execute_reply":"2024-07-10T04:02:14.497185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install duckdb","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:02:14.500687Z","iopub.execute_input":"2024-07-10T04:02:14.501014Z","iopub.status.idle":"2024-07-10T04:02:26.849288Z","shell.execute_reply.started":"2024-07-10T04:02:14.500986Z","shell.execute_reply":"2024-07-10T04:02:26.848105Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nfrom sklearn.model_selection import train_test_split\n\nimport duckdb\nimport rdkit\nfrom rdkit import Chem\nfrom rdkit.Chem import rdFingerprintGenerator\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset,DataLoader\nfrom torchvision import datasets, transforms\nimport torch.nn.functional as F \nimport torchvision\nfrom torchvision import models\nimport torchvision.transforms as T\n\n\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:02:26.850733Z","iopub.execute_input":"2024-07-10T04:02:26.851035Z","iopub.status.idle":"2024-07-10T04:02:26.858553Z","shell.execute_reply.started":"2024-07-10T04:02:26.850993Z","shell.execute_reply":"2024-07-10T04:02:26.857664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path ='/kaggle/input/leash-BELKA/train.parquet'\ntest_path = '/kaggle/input/leash-BELKA/test.parquet'","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:02:26.859871Z","iopub.execute_input":"2024-07-10T04:02:26.860200Z","iopub.status.idle":"2024-07-10T04:02:26.872986Z","shell.execute_reply.started":"2024-07-10T04:02:26.860166Z","shell.execute_reply":"2024-07-10T04:02:26.872121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"con = duckdb.connect()\ndf = con.query(f\"\"\"\n    (SELECT \n        molecule_smiles,\n        MAX(CASE WHEN protein_name = 'BRD4' THEN binds END) AS binds_BRD4,\n        MAX(CASE WHEN protein_name = 'sEH' THEN binds END) AS binds_sEH,\n        MAX(CASE WHEN protein_name = 'HSA' THEN binds END) AS binds_HSA\n     FROM parquet_scan('{train_path}')\n     WHERE binds = 0\n     GROUP BY molecule_smiles\n     HAVING \n        MAX(CASE WHEN protein_name = 'BRD4' THEN binds END) IS NOT NULL AND\n        MAX(CASE WHEN protein_name = 'sEH' THEN binds END) IS NOT NULL AND\n        MAX(CASE WHEN protein_name = 'HSA' THEN binds END) IS NOT NULL\n     ORDER BY RANDOM()\n     LIMIT 30000)\n    UNION ALL\n    (SELECT \n        molecule_smiles,\n        MAX(CASE WHEN protein_name = 'BRD4' THEN binds END) AS binds_BRD4,\n        MAX(CASE WHEN protein_name = 'sEH' THEN binds END) AS binds_sEH,\n        MAX(CASE WHEN protein_name = 'HSA' THEN binds END) AS binds_HSA\n     FROM parquet_scan('{train_path}')\n     WHERE binds = 1\n     GROUP BY molecule_smiles\n     ORDER BY RANDOM()\n     LIMIT 30000)\n    UNION ALL \n    (SELECT \n        molecule_smiles,\n        MAX(CASE WHEN protein_name = 'BRD4' THEN binds END) AS binds_BRD4,\n        MAX(CASE WHEN protein_name = 'sEH' THEN binds END) AS binds_sEH,\n        MAX(CASE WHEN protein_name = 'HSA' THEN binds END) AS binds_HSA\n     FROM parquet_scan('{train_path}')\n     GROUP BY molecule_smiles\n     HAVING \n        MAX(CASE WHEN protein_name = 'BRD4' THEN binds END) IS NOT NULL AND\n        MAX(CASE WHEN protein_name = 'sEH' THEN binds END) IS NOT NULL AND\n        MAX(CASE WHEN protein_name = 'HSA' THEN binds END) IS NOT NULL AND\n        (MAX(CASE WHEN protein_name = 'BRD4' THEN binds END) != MAX(CASE WHEN protein_name = 'sEH' THEN binds END) OR\n         MAX(CASE WHEN protein_name = 'BRD4' THEN binds END) != MAX(CASE WHEN protein_name = 'HSA' THEN binds END) OR\n         MAX(CASE WHEN protein_name = 'sEH' THEN binds END) != MAX(CASE WHEN protein_name = 'HSA' THEN binds END))\n     LIMIT 30000)\n\"\"\").df()\n\n# Close the connection\ncon.close()","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:02:26.875618Z","iopub.execute_input":"2024-07-10T04:02:26.876109Z","iopub.status.idle":"2024-07-10T04:04:49.630454Z","shell.execute_reply.started":"2024-07-10T04:02:26.876079Z","shell.execute_reply":"2024-07-10T04:04:49.628375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.shape","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:04:49.635911Z","iopub.execute_input":"2024-07-10T04:04:49.636332Z","iopub.status.idle":"2024-07-10T04:04:49.645919Z","shell.execute_reply.started":"2024-07-10T04:04:49.636295Z","shell.execute_reply":"2024-07-10T04:04:49.645078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.info()","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:04:49.647033Z","iopub.execute_input":"2024-07-10T04:04:49.647299Z","iopub.status.idle":"2024-07-10T04:04:49.730625Z","shell.execute_reply.started":"2024-07-10T04:04:49.647276Z","shell.execute_reply":"2024-07-10T04:04:49.729748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head(-10)","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:04:49.731713Z","iopub.execute_input":"2024-07-10T04:04:49.731972Z","iopub.status.idle":"2024-07-10T04:04:49.752103Z","shell.execute_reply.started":"2024-07-10T04:04:49.731949Z","shell.execute_reply":"2024-07-10T04:04:49.751014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_heatmap(df):\n    # Extract the relevant columns for the heatmap\n    heatmap_data = df[['binds_BRD4', 'binds_sEH', 'binds_HSA']]\n    # Plot the heatmap\n    plt.figure(figsize=(10, 8))\n    sns.heatmap(heatmap_data, cmap='viridis', fmt='.2f')\n    plt.title('Heatmap of Protein Binding Values')\n    plt.xlabel('Proteins')\n    plt.ylabel('Proteins')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:04:49.753256Z","iopub.execute_input":"2024-07-10T04:04:49.753520Z","iopub.status.idle":"2024-07-10T04:04:49.761636Z","shell.execute_reply.started":"2024-07-10T04:04:49.753497Z","shell.execute_reply":"2024-07-10T04:04:49.760877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_heatmap(df)","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:04:49.762663Z","iopub.execute_input":"2024-07-10T04:04:49.762914Z","iopub.status.idle":"2024-07-10T04:04:50.684603Z","shell.execute_reply.started":"2024-07-10T04:04:49.762892Z","shell.execute_reply":"2024-07-10T04:04:50.683677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Fill NaN value with 1 on three protein colums\n#to increase the binds =1 values\ncolumns_to_fill = ['binds_BRD4','binds_sEH','binds_HSA']\ndf[columns_to_fill] =df[columns_to_fill].fillna(1)\n\nplot_heatmap(df)","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:04:50.685765Z","iopub.execute_input":"2024-07-10T04:04:50.686105Z","iopub.status.idle":"2024-07-10T04:04:51.532508Z","shell.execute_reply.started":"2024-07-10T04:04:50.686078Z","shell.execute_reply":"2024-07-10T04:04:51.531144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.info()","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:04:51.533811Z","iopub.execute_input":"2024-07-10T04:04:51.534117Z","iopub.status.idle":"2024-07-10T04:04:51.555040Z","shell.execute_reply.started":"2024-07-10T04:04:51.534090Z","shell.execute_reply":"2024-07-10T04:04:51.554036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Generate ECFP \ndf['molecule'] = df['molecule_smiles'].apply(Chem.MolFromSmiles)\n\nmfpgen = rdFingerprintGenerator.GetMorganGenerator(radius=2, fpSize=1024)\ndf['ecfp']=df['molecule'].apply(mfpgen.GetFingerprint).apply(lambda bitvec:list(bitvec))","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:04:51.556556Z","iopub.execute_input":"2024-07-10T04:04:51.557149Z","iopub.status.idle":"2024-07-10T04:07:41.810805Z","shell.execute_reply.started":"2024-07-10T04:04:51.557107Z","shell.execute_reply":"2024-07-10T04:07:41.809973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#0-29999,30000-59999,60000-89999\n\n#Split for 0-29999\ndf_part1_train,df_part1_temp=train_test_split(df.iloc[0:30000],train_size=0.7,random_state=872)\ndf_part1_val,df_part1_test=train_test_split(df_part1_temp,train_size=0.7,random_state=872)\n\n#Split for 30000-59999\ndf_part2_train,df_part2_temp=train_test_split(df.iloc[30000:60000],train_size=0.7,random_state=872)\ndf_part2_val,df_part2_test=train_test_split(df_part2_temp,train_size=0.7,random_state=872)\n\n#Split for 60000-90000\ndf_part3_train,df_part3_temp=train_test_split(df.iloc[60000:90000],train_size=0.7,random_state=872)\ndf_part3_val,df_part3_test=train_test_split(df_part2_temp,train_size=0.7,random_state=872)\n\n#Concatenate train, val and test sets\ndf_train = pd.concat([df_part1_train,df_part2_train,df_part3_train]).reset_index(drop=True)\ndf_val = pd.concat([df_part1_val,df_part2_val,df_part3_val])\ndf_test = pd.concat([df_part1_test,df_part2_test,df_part3_test])\n\nprint(f\"Train set size: {len(df_train)}\")\nprint(f\"Validation set size: {len(df_val)}\")\nprint(f\"Test set size: {len(df_test)}\")","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:07:41.815105Z","iopub.execute_input":"2024-07-10T04:07:41.815399Z","iopub.status.idle":"2024-07-10T04:07:47.905527Z","shell.execute_reply.started":"2024-07-10T04:07:41.815374Z","shell.execute_reply":"2024-07-10T04:07:47.904495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.info()","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:07:47.906554Z","iopub.execute_input":"2024-07-10T04:07:47.906806Z","iopub.status.idle":"2024-07-10T04:07:47.943500Z","shell.execute_reply.started":"2024-07-10T04:07:47.906783Z","shell.execute_reply":"2024-07-10T04:07:47.942527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import DataLoader, Dataset","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:07:47.944648Z","iopub.execute_input":"2024-07-10T04:07:47.945005Z","iopub.status.idle":"2024-07-10T04:07:47.954493Z","shell.execute_reply.started":"2024-07-10T04:07:47.944974Z","shell.execute_reply":"2024-07-10T04:07:47.953646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MoleculeDataset(Dataset):\n    def __init__(self, df, feature_column, label_columns):\n        self.features = df[feature_column].tolist()\n        self.labels = df[label_columns].values\n        \n    def __len__(self):\n        return len(self.features)\n    \n    def __getitem__(self, idx):\n        features = torch.tensor(self.features[idx], dtype=torch.float32)\n        labels = torch.tensor(self.labels[idx], dtype=torch.float32)\n        return features, labels","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:07:47.955595Z","iopub.execute_input":"2024-07-10T04:07:47.955913Z","iopub.status.idle":"2024-07-10T04:07:47.965980Z","shell.execute_reply.started":"2024-07-10T04:07:47.955884Z","shell.execute_reply":"2024-07-10T04:07:47.965229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create datasets\ntrain_dataset = MoleculeDataset(df_train, 'ecfp', ['binds_BRD4', 'binds_sEH', 'binds_HSA'])\nval_dataset = MoleculeDataset(df_val, 'ecfp', ['binds_BRD4', 'binds_sEH', 'binds_HSA'])\ntest_dataset = MoleculeDataset(df_test, 'ecfp', ['binds_BRD4', 'binds_sEH', 'binds_HSA'])\n\n# Create dataloaders\ntrain_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=64, shuffle=False)\ntest_loader = DataLoader(test_dataset, batch_size=64, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:07:47.967177Z","iopub.execute_input":"2024-07-10T04:07:47.967627Z","iopub.status.idle":"2024-07-10T04:07:48.069837Z","shell.execute_reply.started":"2024-07-10T04:07:47.967596Z","shell.execute_reply":"2024-07-10T04:07:48.069041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nfrom sklearn.metrics import accuracy_score, classification_report,roc_auc_score\nimport time\n\ndef get_class_labels():\n    return ['binds_BRD4','binds_sEH','binds_HSA']\n\ndef train_and_evaluate(model,train_loader,test_loader,optimizer,scheduler,model_name,num_epochs=10,p=1):\n    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    model.to(device)\n    print(device)\n    \n    criterion = nn.CrossEntropyLoss(reduction='mean') \n    train_loss_set,test_loss_set=[],[]\n    acc_train_set,acc_test_set=[],[]\n    best_test_loss = float('inf')\n    \n    train_start= time.time()\n    for epoch in range(num_epochs):\n        model.train()\n        total_loss = 0.0\n        train_acc_sum=0.0\n        all_train_labels,all_train_outputs=[],[]\n        \n        for inputs,labels in train_loader:\n            inputs,labels= inputs.to(device) , labels.to(device)\n            optimizer.zero_grad()\n            outputs = model(inputs)\n            #print(labels)\n            #print(outputs)\n            loss = criterion(outputs,labels)\n            loss.backward()\n            optimizer.step()\n            total_loss += loss.item()\n            all_train_labels.extend(labels.cpu().detach().numpy())\n            #print(outputs)\n            preds = (outputs > 0.5).float()\n            #print(preds)\n            #print(labels)\n            all_train_outputs.extend(preds.cpu().detach().numpy())\n        \n        average_train_loss = total_loss / len(train_loader)\n        train_accuracy=accuracy_score(all_train_labels,all_train_outputs)\n        train_loss_set.append(average_train_loss)\n        acc_train_set.append(train_accuracy) \n        scheduler.step()\n        \n        model.eval()\n        all_test_labels,all_test_preds=[],[]\n        with torch.no_grad():\n            total_test_loss=0.0\n            for inputs ,labels in test_loader:\n                inputs ,labels = inputs.to(device) ,labels.to(device)\n                outputs = model(inputs)\n                \n                preds = (outputs > 0.5).float()\n                all_test_labels.extend(labels.cpu().numpy())\n                all_test_preds.extend(preds.cpu().numpy())\n                total_test_loss += criterion(outputs,labels).item()\n                \n        average_test_loss = total_test_loss/len(test_loader)\n        test_loss_set.append(average_test_loss)\n        \n        test_accuracy = accuracy_score(all_test_labels,all_test_preds)\n        acc_test_set.append(test_accuracy)\n        current_lr = optimizer.param_groups[0]['lr']\n        \n        if average_test_loss < best_test_loss:\n            # Save the model weights when the test loss is the lowest\n            best_test_loss = average_test_loss\n            torch.save(model.state_dict(), f'model_weights_{model_name}.pth')\n            \n        if epoch%p==0:\n            print(f'Epoch {epoch+1}/{num_epochs}: \\t Lr:{current_lr},\\t train loss:{average_train_loss:.4f},\\t train acc:{train_accuracy:.4f},\\t test loss:{average_test_loss:.4f},\\t test acc:{test_accuracy:.4f} ')\n            print('-'* 80 )\n            \n    train_end = time.time()\n    time_use = train_end - train_start\n    print(f'Time used for Training:{time_use} sec ')\n    print('-'* 80) \n    \n    report_train = classification_report(all_train_labels,all_train_outputs,target_names=get_class_labels())\n    report_test = classification_report(all_test_labels,all_test_preds,target_names=get_class_labels())\n    print(f'Training Classification Report :\\n {report_train}')\n    print(f'Test Classification Report :\\n {report_test}')\n        \n    return model, train_loss_set, test_loss_set, acc_train_set, acc_test_set, time_use","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:07:48.073087Z","iopub.execute_input":"2024-07-10T04:07:48.073349Z","iopub.status.idle":"2024-07-10T04:07:48.094509Z","shell.execute_reply.started":"2024-07-10T04:07:48.073327Z","shell.execute_reply":"2024-07-10T04:07:48.093592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nclass FeatureExtractor(nn.Module):\n    def __init__(self, input_dim):\n        super(FeatureExtractor, self).__init__()\n        self.conv1 = nn.Conv1d(in_channels=1, out_channels=16, kernel_size=3, padding=1)\n        self.conv2 = nn.Conv1d(in_channels=16, out_channels=32, kernel_size=3, padding=1)\n        self.bn = nn.BatchNorm1d(32)  # Batch normalization after convolution\n        self.relu = nn.ReLU()\n        self.dropout = nn.Dropout(0.5)\n        self.conv3 = nn.Conv1d(in_channels=32, out_channels=64, kernel_size=3, padding=1)\n        self.global_avg_pool = nn.AdaptiveAvgPool1d(1)  # Global average pooling across the entire sequence\n        self.fc1 = nn.Linear(64, 32)\n        self.fc2 = nn.Linear(32, 16)\n\n    def forward(self, x):\n        # Input x should have shape (batch_size, input_dim)\n        x = x.unsqueeze(1)  # Add a channel dimension for Conv1d: (batch_size, 1, input_dim)\n        x = self.conv1(x)\n        x = self.relu(x)\n        x = self.conv2(x)\n        x = self.bn(x)\n        x = self.relu(x)\n        x = self.dropout(x)\n        x = self.conv3(x)\n        x = self.relu(x)\n        x = self.global_avg_pool(x).squeeze(-1)  # Global average pooling: (batch_size, 64)\n        x = self.relu(self.fc1(x))\n        x = self.relu(self.fc2(x))\n        return x\n\nclass Classifier(nn.Module):\n    def __init__(self, input_dim, output_dim):\n        super(Classifier, self).__init__()\n        self.fc3 = nn.Linear(input_dim, output_dim)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        x = self.fc3(x)\n        x = self.sigmoid(x)\n        return x\n\nclass MultiLabelNN(nn.Module):\n    def __init__(self, input_dim, output_dim):\n        super(MultiLabelNN, self).__init__()\n        self.feature_extractor = FeatureExtractor(input_dim)\n        self.classifier = Classifier(16, output_dim)\n\n    def forward(self, x):\n        x = self.feature_extractor(x)\n        x = self.classifier(x)\n        return x\n","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:07:48.095747Z","iopub.execute_input":"2024-07-10T04:07:48.096109Z","iopub.status.idle":"2024-07-10T04:07:48.111106Z","shell.execute_reply.started":"2024-07-10T04:07:48.096077Z","shell.execute_reply":"2024-07-10T04:07:48.110153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_name='NeuralLinear'\ninput_size = len(df_train.iloc[0]['ecfp'])  # Length of the ECFP vector\nclass_num = 3  # Number of labels (binds_BRD4, binds_sEH, binds_HSA)\nnumber_epochs =60\np=1\n\nmodel = MultiLabelNN(input_dim=input_size, output_dim=class_num)\n\noptimizer = optim.Adam(model.parameters(),lr=0.01)\n# Define learning rate scheduler\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1)  # Reduce learning rate after every 10 epochs\n\nmodel, train_loss, test_loss, acc_train, acc_test, time_use = train_and_evaluate(model,train_loader,test_loader,optimizer=optimizer,scheduler=scheduler, model_name=model_name,num_epochs=number_epochs,p=p)","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:07:48.112305Z","iopub.execute_input":"2024-07-10T04:07:48.112577Z","iopub.status.idle":"2024-07-10T04:24:33.808987Z","shell.execute_reply.started":"2024-07-10T04:07:48.112554Z","shell.execute_reply":"2024-07-10T04:24:33.807954Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_name='NeuralLinear'\ninput_size = len(df_train.iloc[0]['ecfp'])  # Length of the ECFP vector\nclass_num = 3  # Number of labels (binds_BRD4, binds_sEH, binds_HSA)\nnumber_epochs =60\np=1\n\nmodel = MultiLabelNN(input_dim=input_size, output_dim=class_num)\n\noptimizer = optim.Adam(model.parameters(),lr=0.05)\n# Define learning rate scheduler\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.1)  # Reduce learning rate after every 10 epochs\n\nmodel, train_loss, test_loss, acc_train, acc_test, time_use = train_and_evaluate(model,train_loader,test_loader,optimizer=optimizer,scheduler=scheduler, model_name=model_name,num_epochs=number_epochs,p=p)","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:24:33.810586Z","iopub.execute_input":"2024-07-10T04:24:33.810895Z","iopub.status.idle":"2024-07-10T04:27:24.439107Z","shell.execute_reply.started":"2024-07-10T04:24:33.810868Z","shell.execute_reply":"2024-07-10T04:27:24.437632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def test_model(model,test_loader,model_name):\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    model.to(device)\n    print(device)\n    \n    model.eval\n    all_test_labels,all_test_outputs=[],[]\n    time_start=time.time()\n    for inputs,labels in test_loader:\n        inputs,labels=inputs.to(device), labels.to(device)\n        outputs = model(inputs)\n        #print(outputs)\n        _, preds = torch.max(outputs, 1) # pick the highest\n        all_test_labels.extend(labels.cpu().numpy())\n        all_test_outputs.extend(preds.cpu().numpy())\n        \n    # Save labels and predictions to a DataFrame\n    results_df = pd.DataFrame({'True Labels': all_test_labels, 'Predicted Labels': all_test_outputs})\n    results_df.to_csv(f'/kaggle/working/test_results_{model_name}.csv', index=False)\n    \n    time_end=time.time()\n    time_use=time_end-time_start\n    report_test = classification_report(all_test_labels,all_test_outputs,target_names=get_class_labels())\n    conf_matrix = confusion_matrix(all_test_labels, all_test_outputs)\n    print(f'Time use : {time_use} sec')\n    print(f'Test Classification Report :\\n {report_test}')\n    print(f\"Confusion Matrix:\\n{conf_matrix}\")\n    # Plot confusion matrix\n    plt.figure(figsize=(8, 6))\n    sns.heatmap(conf_matrix, annot=True, fmt='d', cmap='Blues', \n                xticklabels=[0, 1], yticklabels=[0, 1])\n    plt.xlabel('Predicted')\n    plt.ylabel('True')\n    plt.title('Confusion Matrix')\n    plt.savefig(f'/kaggle/working/Confusion Matrix on {model_name}.jpg')\n    plt.show()\n    ","metadata":{"execution":{"iopub.status.busy":"2024-07-10T04:27:24.440228Z","iopub.status.idle":"2024-07-10T04:27:24.440664Z","shell.execute_reply.started":"2024-07-10T04:27:24.440477Z","shell.execute_reply":"2024-07-10T04:27:24.440494Z"},"trusted":true},"execution_count":null,"outputs":[]}]}