{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":10338,"databundleVersionId":862042,"sourceType":"competition"}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"raw","source":"!pip install pydicom","metadata":{"execution":{"iopub.status.busy":"2023-12-02T15:47:49.781246Z","iopub.execute_input":"2023-12-02T15:47:49.781515Z","iopub.status.idle":"2023-12-02T15:48:03.245874Z","shell.execute_reply.started":"2023-12-02T15:47:49.781490Z","shell.execute_reply":"2023-12-02T15:48:03.244785Z"}}},{"cell_type":"code","source":"!pip install numpy","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:07.681043Z","iopub.execute_input":"2023-12-03T00:09:07.681433Z","iopub.status.idle":"2023-12-03T00:09:19.628562Z","shell.execute_reply.started":"2023-12-03T00:09:07.681401Z","shell.execute_reply":"2023-12-03T00:09:19.627147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pydicom \nimport os \nimport pandas as pd \nimport numpy as np \n\nimport cv2\nimport matplotlib.pyplot as plt \n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:19.631180Z","iopub.execute_input":"2023-12-03T00:09:19.631524Z","iopub.status.idle":"2023-12-03T00:09:19.636576Z","shell.execute_reply.started":"2023-12-03T00:09:19.631495Z","shell.execute_reply":"2023-12-03T00:09:19.635677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Load the metadata from the corresponding CSV files and store them in Pandas DataFrames","metadata":{}},{"cell_type":"code","source":"class_data = pd.read_csv(\"/kaggle/input/rsna-pneumonia-detection-challenge/stage_2_detailed_class_info.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:19.637786Z","iopub.execute_input":"2023-12-03T00:09:19.638034Z","iopub.status.idle":"2023-12-03T00:09:19.686216Z","shell.execute_reply.started":"2023-12-03T00:09:19.638012Z","shell.execute_reply":"2023-12-03T00:09:19.685511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels = pd.read_csv(\"/kaggle/input/rsna-pneumonia-detection-challenge/stage_2_train_labels.csv\")","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:19.688185Z","iopub.execute_input":"2023-12-03T00:09:19.688460Z","iopub.status.idle":"2023-12-03T00:09:19.726575Z","shell.execute_reply.started":"2023-12-03T00:09:19.688436Z","shell.execute_reply":"2023-12-03T00:09:19.725817Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Define the directory where the DICOM images are stored","metadata":{}},{"cell_type":"code","source":"data_dir = \"/kaggle/input/rsna-pneumonia-detection-challenge/stage_2_train_images\"","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:19.727550Z","iopub.execute_input":"2023-12-03T00:09:19.727821Z","iopub.status.idle":"2023-12-03T00:09:19.731510Z","shell.execute_reply.started":"2023-12-03T00:09:19.727799Z","shell.execute_reply":"2023-12-03T00:09:19.730678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"List the content of the directory to calculate the total number of DICOM images\n","metadata":{}},{"cell_type":"code","source":"num_dir = os.listdir(data_dir)\nprint(len(num_dir))","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:19.732569Z","iopub.execute_input":"2023-12-03T00:09:19.732872Z","iopub.status.idle":"2023-12-03T00:09:19.756452Z","shell.execute_reply.started":"2023-12-03T00:09:19.732821Z","shell.execute_reply":"2023-12-03T00:09:19.755720Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Combine metadata from class_data and train_labels by joining them on every column except ‘patientId’\n","metadata":{}},{"cell_type":"code","source":"dataset = pd.concat([class_data.drop(columns = 'patientId'), train_labels], axis = 1)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:19.757429Z","iopub.execute_input":"2023-12-03T00:09:19.757704Z","iopub.status.idle":"2023-12-03T00:09:19.765596Z","shell.execute_reply.started":"2023-12-03T00:09:19.757681Z","shell.execute_reply":"2023-12-03T00:09:19.764479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Display the first few rows of the combined dataset to verify correct concatenation\n","metadata":{}},{"cell_type":"code","source":"dataset.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:19.766831Z","iopub.execute_input":"2023-12-03T00:09:19.767495Z","iopub.status.idle":"2023-12-03T00:09:19.786433Z","shell.execute_reply.started":"2023-12-03T00:09:19.767461Z","shell.execute_reply":"2023-12-03T00:09:19.785552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Create a new column ‘bbox’ representing the combined information of bounding box columns\n","metadata":{}},{"cell_type":"code","source":"dataset['bbox'] = dataset[['x', 'y', 'height', 'width']].apply(lambda x: '-'.join(str(i) for i in x), axis=1)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:19.787572Z","iopub.execute_input":"2023-12-03T00:09:19.787886Z","iopub.status.idle":"2023-12-03T00:09:20.112453Z","shell.execute_reply.started":"2023-12-03T00:09:19.787863Z","shell.execute_reply":"2023-12-03T00:09:20.111465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Drop the original bounding box columns as they are no longer needed due to the creation of ‘bbox’ column","metadata":{}},{"cell_type":"code","source":"dataset.drop(columns = [\"x\", \"y\", \"width\", \"height\"])","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.117844Z","iopub.execute_input":"2023-12-03T00:09:20.118169Z","iopub.status.idle":"2023-12-03T00:09:20.134007Z","shell.execute_reply.started":"2023-12-03T00:09:20.118141Z","shell.execute_reply":"2023-12-03T00:09:20.132979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Count the number of cases where the ‘Target’ column is 1 indicating the presence of pneumonia\n","metadata":{}},{"cell_type":"code","source":"print(len(dataset[dataset[\"Target\"] == 1]))","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.135132Z","iopub.execute_input":"2023-12-03T00:09:20.135456Z","iopub.status.idle":"2023-12-03T00:09:20.147443Z","shell.execute_reply.started":"2023-12-03T00:09:20.135429Z","shell.execute_reply":"2023-12-03T00:09:20.146530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Image View","metadata":{}},{"cell_type":"markdown","source":"Concatenate the file path and filename to form the complete image path for each image\n","metadata":{}},{"cell_type":"code","source":"image_directory = '/kaggle/input/rsna-pneumonia-detection-challenge/stage_2_train_images/'\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.148762Z","iopub.execute_input":"2023-12-03T00:09:20.149109Z","iopub.status.idle":"2023-12-03T00:09:20.156663Z","shell.execute_reply.started":"2023-12-03T00:09:20.149077Z","shell.execute_reply":"2023-12-03T00:09:20.155874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset['image_path'] = image_directory + dataset['patientId'] + '.dcm'\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.157715Z","iopub.execute_input":"2023-12-03T00:09:20.158047Z","iopub.status.idle":"2023-12-03T00:09:20.179253Z","shell.execute_reply.started":"2023-12-03T00:09:20.158023Z","shell.execute_reply":"2023-12-03T00:09:20.178571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Display the updated dataset with the ‘image_path’ column added\n","metadata":{}},{"cell_type":"code","source":"dataset.head()","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.180222Z","iopub.execute_input":"2023-12-03T00:09:20.180530Z","iopub.status.idle":"2023-12-03T00:09:20.200592Z","shell.execute_reply.started":"2023-12-03T00:09:20.180506Z","shell.execute_reply":"2023-12-03T00:09:20.199690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Function to read DICOM images, resize them to 224x224 pixels, convert to RGB, and store them along with their targets\n","metadata":{}},{"cell_type":"code","source":"def read_and_resize_images(pneumonia, num_samples=None):\n    resized_images = []\n    boxes = []\n    \n    if num_samples:\n        pneumonia = pneumonia[:num_samples]\n    \n    for _, row in pneumonia.iterrows():\n        image_path = row['image_path']\n        target = row['Target']\n        \n        dicom_data = pydicom.read_file(image_path)\n        img = dicom_data.pixel_array\n        \n        img = cv2.resize(img, (224, 224))\n        img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)\n        \n        resized_images.append(img)\n        boxes.append(np.array(target, dtype=np.float32))\n    \n    return np.array(resized_images), np.array(boxes)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.201618Z","iopub.execute_input":"2023-12-03T00:09:20.201981Z","iopub.status.idle":"2023-12-03T00:09:20.212614Z","shell.execute_reply.started":"2023-12-03T00:09:20.201951Z","shell.execute_reply":"2023-12-03T00:09:20.211675Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Choose a random patient’s image for visualization\n","metadata":{}},{"cell_type":"code","source":"import random","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.213744Z","iopub.execute_input":"2023-12-03T00:09:20.214038Z","iopub.status.idle":"2023-12-03T00:09:20.228969Z","shell.execute_reply.started":"2023-12-03T00:09:20.214016Z","shell.execute_reply":"2023-12-03T00:09:20.227975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"random_patient_id = random.choice(dataset['patientId'])\n\nimage_path = dataset[dataset['patientId'] == random_patient_id]['image_path'].values[0]\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.230185Z","iopub.execute_input":"2023-12-03T00:09:20.230541Z","iopub.status.idle":"2023-12-03T00:09:20.246736Z","shell.execute_reply.started":"2023-12-03T00:09:20.230508Z","shell.execute_reply":"2023-12-03T00:09:20.245945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dicom_data = pydicom.read_file(image_path)\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.248472Z","iopub.execute_input":"2023-12-03T00:09:20.248866Z","iopub.status.idle":"2023-12-03T00:09:20.269188Z","shell.execute_reply.started":"2023-12-03T00:09:20.248833Z","shell.execute_reply":"2023-12-03T00:09:20.268491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"display a DICOM image","metadata":{}},{"cell_type":"code","source":"def visualize_dicom_image(dicom_data):\n   \n    img = dicom_data.pixel_array\n\n    plt.figure(figsize=(6, 6))\n    plt.imshow(img, cmap='gray')\n    plt.title('DICOM Image')\n    plt.axis('off') \n    plt.show()\n    \nvisualize_dicom_image(dicom_data)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.270170Z","iopub.execute_input":"2023-12-03T00:09:20.270455Z","iopub.status.idle":"2023-12-03T00:09:20.575513Z","shell.execute_reply.started":"2023-12-03T00:09:20.270430Z","shell.execute_reply":"2023-12-03T00:09:20.574545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CNN - Simple CNN","metadata":{}},{"cell_type":"markdown","source":"## CNN Model - transfer learning with resnet101\n\n","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom torchvision import transforms","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.576797Z","iopub.execute_input":"2023-12-03T00:09:20.577129Z","iopub.status.idle":"2023-12-03T00:09:20.581436Z","shell.execute_reply.started":"2023-12-03T00:09:20.577100Z","shell.execute_reply":"2023-12-03T00:09:20.580499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Read and process a subset (1000 samples) of the images and targets for modeling","metadata":{}},{"cell_type":"code","source":"X, y = read_and_resize_images(dataset[:1000])","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:20.582791Z","iopub.execute_input":"2023-12-03T00:09:20.583153Z","iopub.status.idle":"2023-12-03T00:09:27.568523Z","shell.execute_reply.started":"2023-12-03T00:09:20.583118Z","shell.execute_reply":"2023-12-03T00:09:27.567672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Split the data into training and testing sets, and then transform them","metadata":{}},{"cell_type":"code","source":"X_train, X_test, y_train, y_test = train_test_split(X, y, test_size = 0.2)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:27.569805Z","iopub.execute_input":"2023-12-03T00:09:27.570093Z","iopub.status.idle":"2023-12-03T00:09:27.615719Z","shell.execute_reply.started":"2023-12-03T00:09:27.570068Z","shell.execute_reply":"2023-12-03T00:09:27.614899Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_size = 224","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:27.616941Z","iopub.execute_input":"2023-12-03T00:09:27.617253Z","iopub.status.idle":"2023-12-03T00:09:27.621391Z","shell.execute_reply.started":"2023-12-03T00:09:27.617227Z","shell.execute_reply":"2023-12-03T00:09:27.620542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.ToPILImage(), \n    transforms.Resize((int(img_size), int(img_size))), \n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n])\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:27.622494Z","iopub.execute_input":"2023-12-03T00:09:27.622763Z","iopub.status.idle":"2023-12-03T00:09:27.633985Z","shell.execute_reply.started":"2023-12-03T00:09:27.622739Z","shell.execute_reply":"2023-12-03T00:09:27.633339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torchvision.transforms import ToTensor\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nimport torchvision.models as models\nimport torch.nn as nn\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:27.635146Z","iopub.execute_input":"2023-12-03T00:09:27.635405Z","iopub.status.idle":"2023-12-03T00:09:27.647038Z","shell.execute_reply.started":"2023-12-03T00:09:27.635382Z","shell.execute_reply":"2023-12-03T00:09:27.646133Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"PyTorch Dataset class that encapsulates the preprocessed images and corresponding labels","metadata":{}},{"cell_type":"code","source":"class PneumoniaDataset(Dataset):\n    def __init__(self, images, labels, transform=None):\n        self.images = images\n        self.labels = labels.astype(int)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n        image = self.images[idx]\n        label = self.labels[idx]\n\n        if self.transform:\n            image = self.transform(image)\n\n        label = torch.tensor(label, dtype=torch.long)\n\n        return image, label\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:27.648141Z","iopub.execute_input":"2023-12-03T00:09:27.648715Z","iopub.status.idle":"2023-12-03T00:09:27.658476Z","shell.execute_reply.started":"2023-12-03T00:09:27.648681Z","shell.execute_reply":"2023-12-03T00:09:27.657804Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Create dataset and dataloader instances for training and testing\n","metadata":{}},{"cell_type":"code","source":"train_dataset = PneumoniaDataset(X_train, y_train, transform=transform)\ntest_dataset = PneumoniaDataset(X_test, y_test, transform=transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)\ntest_loader = DataLoader(test_dataset, batch_size=32, shuffle=False)\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:27.659543Z","iopub.execute_input":"2023-12-03T00:09:27.659913Z","iopub.status.idle":"2023-12-03T00:09:27.675318Z","shell.execute_reply.started":"2023-12-03T00:09:27.659889Z","shell.execute_reply":"2023-12-03T00:09:27.674303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Prepare the Pretrained ResNet101 model for fine-tuning, and set up the loss function and optimizer\n","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nresnet101 = models.resnet101(pretrained=True)\nnum_classes = 2  \nresnet101.fc = nn.Linear(resnet101.fc.in_features, num_classes)\nresnet101 = resnet101.to(device)\n\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(resnet101.parameters(), lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:27.681237Z","iopub.execute_input":"2023-12-03T00:09:27.681514Z","iopub.status.idle":"2023-12-03T00:09:28.617091Z","shell.execute_reply.started":"2023-12-03T00:09:27.681490Z","shell.execute_reply":"2023-12-03T00:09:28.616128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Training loop for the ResNet101 model with a set number of epochs","metadata":{}},{"cell_type":"code","source":"num_epochs = 10\n\ntrain_losses = []\ntrain_accuracies = []\n\nfor epoch in range(num_epochs):\n    resnet101.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    for images, labels in train_loader:\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n\n        outputs = resnet101(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n        _, predicted = torch.max(outputs.data, 1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n\n    epoch_loss = running_loss / len(train_loader)\n    epoch_accuracy = 100 * correct / total\n    train_losses.append(epoch_loss)\n    train_accuracies.append(epoch_accuracy)\n\n    print(f\"Epoch {epoch+1}/{num_epochs}, Loss: {epoch_loss}, Accuracy: {epoch_accuracy}%\")\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:09:28.618274Z","iopub.execute_input":"2023-12-03T00:09:28.618577Z","iopub.status.idle":"2023-12-03T00:11:54.337351Z","shell.execute_reply.started":"2023-12-03T00:09:28.618551Z","shell.execute_reply":"2023-12-03T00:11:54.336356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Plot the training loss and accuracy for visualization of training progress\n","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(12, 5))\n\nplt.subplot(1, 2, 1)\nplt.plot(train_losses, label='Training Loss')\nplt.title('Training Loss over Epochs')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\n\nplt.subplot(1, 2, 2)\nplt.plot(train_accuracies, label='Training Accuracy')\nplt.title('Training Accuracy over Epochs')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.legend()\n\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:11:54.338714Z","iopub.execute_input":"2023-12-03T00:11:54.339071Z","iopub.status.idle":"2023-12-03T00:11:54.769523Z","shell.execute_reply.started":"2023-12-03T00:11:54.339043Z","shell.execute_reply":"2023-12-03T00:11:54.768594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Model evaluation to calculate the accuracy on the test set\n","metadata":{}},{"cell_type":"code","source":"resnet101 = resnet101.to(device)\nresnet101.eval()  \n\ncorrect = 0\ntotal = 0\nwith torch.no_grad():  \n    for images, labels in test_loader:\n        images = images.to(device)\n        labels = labels.to(device)\n\n        outputs = resnet101(images)\n        _, predicted = torch.max(outputs.data, 1)\n        total += labels.size(0)\n        correct += (predicted == labels).sum().item()\n\naccuracy = 100 * correct / total\nprint(f'Accuracy on the test set: {accuracy:.2f}%')","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:11:54.770890Z","iopub.execute_input":"2023-12-03T00:11:54.771608Z","iopub.status.idle":"2023-12-03T00:11:55.927159Z","shell.execute_reply.started":"2023-12-03T00:11:54.771572Z","shell.execute_reply":"2023-12-03T00:11:55.926178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Custom CNN\n","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:11:55.928560Z","iopub.execute_input":"2023-12-03T00:11:55.928974Z","iopub.status.idle":"2023-12-03T00:11:55.933651Z","shell.execute_reply.started":"2023-12-03T00:11:55.928937Z","shell.execute_reply":"2023-12-03T00:11:55.932693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Define Simple CNN architecture for model building\n","metadata":{}},{"cell_type":"code","source":"class SimpleCNN(nn.Module):\n    def __init__(self):\n        super(SimpleCNN, self).__init__()\n        \n        self.conv1 = nn.Conv2d(3, 32, kernel_size=3, padding=1)\n        self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1)\n        self.fc1 = nn.Linear(64 * 56 * 56, 128)\n        self.fc2 = nn.Linear(128, 2)\n\n    def forward(self, x):\n        x = F.relu(self.conv1(x))\n        x = F.max_pool2d(x, 2)\n        x = F.relu(self.conv2(x))\n        x = F.max_pool2d(x, 2)\n        x = x.view(x.size(0), -1)  \n        x = F.relu(self.fc1(x))\n        x = self.fc2(x)\n        return x\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:11:55.935253Z","iopub.execute_input":"2023-12-03T00:11:55.935629Z","iopub.status.idle":"2023-12-03T00:11:55.947474Z","shell.execute_reply.started":"2023-12-03T00:11:55.935595Z","shell.execute_reply":"2023-12-03T00:11:55.946564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Instantiate the SimpleCNN model and move it to the appropriate device (CPU or GPU)","metadata":{}},{"cell_type":"code","source":"model = SimpleCNN()\nmodel = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:11:55.948840Z","iopub.execute_input":"2023-12-03T00:11:55.949132Z","iopub.status.idle":"2023-12-03T00:11:56.139485Z","shell.execute_reply.started":"2023-12-03T00:11:55.949107Z","shell.execute_reply":"2023-12-03T00:11:56.138694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Define loss function criterion and optimizer for the SimpleCNN","metadata":{}},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:11:56.140451Z","iopub.execute_input":"2023-12-03T00:11:56.140715Z","iopub.status.idle":"2023-12-03T00:11:56.147415Z","shell.execute_reply.started":"2023-12-03T00:11:56.140692Z","shell.execute_reply":"2023-12-03T00:11:56.146566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Function to train the Custom CNN model with a training loop for the specified number of epochs","metadata":{}},{"cell_type":"code","source":"def train_model(model, train_loader, criterion, optimizer, num_epochs=10):\n    model.train()\n    train_losses = []\n    train_accuracies = []\n\n    for epoch in range(num_epochs):\n        running_loss = 0.0\n        correct = 0\n        total = 0\n\n        for images, labels in train_loader:\n            images, labels = images.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n\n            running_loss += loss.item()\n            _, predicted = torch.max(outputs.data, 1)\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n\n        epoch_loss = running_loss / len(train_loader)\n        epoch_accuracy = 100 * correct / total\n        train_losses.append(epoch_loss)\n        train_accuracies.append(epoch_accuracy)\n\n        print(f\"Epoch {epoch+1}/{num_epochs}, Loss: {epoch_loss:.4f}, Accuracy: {epoch_accuracy:.2f}%\")\n\n    return train_losses, train_accuracies\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:11:56.148911Z","iopub.execute_input":"2023-12-03T00:11:56.149252Z","iopub.status.idle":"2023-12-03T00:11:56.158590Z","shell.execute_reply.started":"2023-12-03T00:11:56.149219Z","shell.execute_reply":"2023-12-03T00:11:56.157793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 10 ","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:11:56.160013Z","iopub.execute_input":"2023-12-03T00:11:56.160911Z","iopub.status.idle":"2023-12-03T00:11:56.173458Z","shell.execute_reply.started":"2023-12-03T00:11:56.160859Z","shell.execute_reply":"2023-12-03T00:11:56.172579Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_losses, train_accuracies = train_model(model, train_loader, criterion, optimizer, num_epochs)\n","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:11:56.174477Z","iopub.execute_input":"2023-12-03T00:11:56.174806Z","iopub.status.idle":"2023-12-03T00:12:14.344898Z","shell.execute_reply.started":"2023-12-03T00:11:56.174775Z","shell.execute_reply":"2023-12-03T00:12:14.344007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Plot training losses and accuracies","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(12, 5))\n\nplt.subplot(1, 2, 1)\nplt.plot(train_losses, label='Training Loss')\nplt.title('Training Loss over Epochs')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\n\nplt.subplot(1, 2, 2)\nplt.plot(train_accuracies, label='Training Accuracy')\nplt.title('Training Accuracy over Epochs')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy (%)')\nplt.legend()\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:12:14.346092Z","iopub.execute_input":"2023-12-03T00:12:14.346384Z","iopub.status.idle":"2023-12-03T00:12:14.827380Z","shell.execute_reply.started":"2023-12-03T00:12:14.346358Z","shell.execute_reply":"2023-12-03T00:12:14.826403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nimport seaborn as sns","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:12:14.828611Z","iopub.execute_input":"2023-12-03T00:12:14.828909Z","iopub.status.idle":"2023-12-03T00:12:14.833434Z","shell.execute_reply.started":"2023-12-03T00:12:14.828883Z","shell.execute_reply":"2023-12-03T00:12:14.832454Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Evaluate the Custom CNN model and calculate the test set accuracy\n","metadata":{}},{"cell_type":"code","source":"def evaluate_model(model, test_loader):\n    model.eval()\n    y_true = []\n    y_pred = []\n\n    with torch.no_grad():\n        for images, labels in test_loader:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            _, predicted = torch.max(outputs, 1)\n            y_true.extend(labels.cpu().numpy())\n            y_pred.extend(predicted.cpu().numpy())\n\n    accuracy = 100 * sum([1 for true, pred in zip(y_true, y_pred) if true == pred]) / len(y_true)\n    print(f'Accuracy on test set: {accuracy:.2f}%')\n\n    \nevaluate_model(model, test_loader)","metadata":{"execution":{"iopub.status.busy":"2023-12-03T00:12:14.834563Z","iopub.execute_input":"2023-12-03T00:12:14.834848Z","iopub.status.idle":"2023-12-03T00:12:15.106271Z","shell.execute_reply.started":"2023-12-03T00:12:14.834824Z","shell.execute_reply":"2023-12-03T00:12:15.105311Z"},"trusted":true},"execution_count":null,"outputs":[]}]}