{"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":75176,"databundleVersionId":8252256,"sourceType":"competition"}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# AMIA Challenge: Explorative Data Analysis, Loading, and Training with ResNet-50 (classification) and Faster R-CNN (bounding box regression)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T13:50:04.494867Z","iopub.execute_input":"2024-04-23T13:50:04.495349Z","iopub.status.idle":"2024-04-23T13:50:04.5Z","shell.execute_reply.started":"2024-04-23T13:50:04.495319Z","shell.execute_reply":"2024-04-23T13:50:04.498974Z"}}},{"cell_type":"code","source":"# imports, helper functions and globals\nfrom pathlib import Path\nfrom typing import List, Tuple, Optional\nimport time\nimport copy\n\nimport albumentations as A\nimport matplotlib as mpl\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport torch\nfrom PIL import Image\nfrom torchvision.ops import nms\nfrom torchvision import datasets, models, transforms\nfrom tqdm.notebook import tqdm, trange\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchinfo import summary\nfrom torch.utils.data import DataLoader\nfrom torch.utils.tensorboard import SummaryWriter\nfrom torchmetrics.classification import MultilabelAccuracy\nfrom sklearn.model_selection import train_test_split\n\ndef ls(path: Path) -> List[Path]:\n    return list(path.iterdir())\n\nROOT = Path(\"/kaggle/input/amia-public-challenge-2024\")\n\n# Mapping of the new classes\nCLASS_MAPPING = {\n    0: 'Normal',\n    1: 'Aortic enlargement',\n    2: 'Other abnormalities'\n}\n\n# Original CLASS_IDS_NAMES from dataset\nORIGINAL_CLASS_IDS_NAMES = {\n    0: 'Aortic enlargement',\n    1: 'Atelectasis',\n    2: 'Calcification',\n    3: 'Cardiomegaly',\n    4: 'Consolidation',\n    5: 'ILD',\n    6: 'Infiltration',\n    7: 'Lung Opacity',\n    8: 'Nodule/Mass',\n    9: 'Other lesion',\n    10: 'Pleural effusion',\n    11: 'Pleural thickening',\n    12: 'Pneumothorax',\n    13: 'Pulmonary fibrosis',\n    14: 'No finding'\n}","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-08-28T22:44:02.456694Z","iopub.execute_input":"2024-08-28T22:44:02.457469Z","iopub.status.idle":"2024-08-28T22:44:23.810687Z","shell.execute_reply.started":"2024-08-28T22:44:02.457431Z","shell.execute_reply":"2024-08-28T22:44:23.809764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Function to map original classes to new classes\ndef map_classes(original_class_ids):\n    # Define the mapping logic\n    mapped_classes = []\n    for class_id in original_class_ids:\n        if class_id == 0:  # Aortic enlargement\n            mapped_classes.append(1)\n        elif class_id == 14:  # No finding\n            mapped_classes.append(0)\n        else:  # All other classes\n            mapped_classes.append(2)\n    return mapped_classes","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:44:23.812461Z","iopub.execute_input":"2024-08-28T22:44:23.813058Z","iopub.status.idle":"2024-08-28T22:44:23.818976Z","shell.execute_reply.started":"2024-08-28T22:44:23.813029Z","shell.execute_reply":"2024-08-28T22:44:23.817844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Folder structure","metadata":{}},{"cell_type":"code","source":"ls(ROOT)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:44:23.823526Z","iopub.execute_input":"2024-08-28T22:44:23.823880Z","iopub.status.idle":"2024-08-28T22:44:23.997591Z","shell.execute_reply.started":"2024-08-28T22:44:23.823850Z","shell.execute_reply":"2024-08-28T22:44:23.996559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# List of train image paths\nimages = ls(ROOT / \"train/train\")\nimages[:2]","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:44:24.000949Z","iopub.execute_input":"2024-08-28T22:44:24.001396Z","iopub.status.idle":"2024-08-28T22:44:24.693183Z","shell.execute_reply.started":"2024-08-28T22:44:24.001364Z","shell.execute_reply":"2024-08-28T22:44:24.692196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# List of test image paths\ntest_images = ls(ROOT / \"train/train\")\ntest_images[:2]","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:44:24.694465Z","iopub.execute_input":"2024-08-28T22:44:24.694857Z","iopub.status.idle":"2024-08-28T22:44:24.718969Z","shell.execute_reply.started":"2024-08-28T22:44:24.694825Z","shell.execute_reply":"2024-08-28T22:44:24.718058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Explore images","metadata":{}},{"cell_type":"code","source":"# Reading a single image\nImage.open(images[0]).resize((256, 256))","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:44:24.720026Z","iopub.execute_input":"2024-08-28T22:44:24.720309Z","iopub.status.idle":"2024-08-28T22:44:24.775330Z","shell.execute_reply.started":"2024-08-28T22:44:24.720285Z","shell.execute_reply":"2024-08-28T22:44:24.774404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Some more examples\n\nnrows=3\nncols=10\n\nfig, axs = plt.subplots(nrows=nrows, ncols=ncols, sharey=True, sharex=True)\nfig.set_size_inches(ncols * 2, nrows * 2)\ni = 0\nfor row in range(nrows):\n    for col in range(ncols):\n        image = np.array(Image.open(images[i]))\n        axs[row, col].imshow(image, cmap=\"gray\")\n        i += 1\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:44:24.776553Z","iopub.execute_input":"2024-08-28T22:44:24.776946Z","iopub.status.idle":"2024-08-28T22:44:32.656552Z","shell.execute_reply.started":"2024-08-28T22:44:24.776918Z","shell.execute_reply":"2024-08-28T22:44:32.655337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Lets look at the pixel distributions\nfor i in range(20):\n    pixels = np.array(Image.open(images[i])).flatten()\n    count, _ = np.histogram(pixels, bins=256, range=(0, 255))\n    plt.plot(count, color=\"black\", alpha=0.5)\n    \nplt.ylabel(\"Count\")\n# plt.yscale(\"log\")\nplt.xlabel(\"Pixel value\")\nplt.grid(axis=\"y\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:44:32.657696Z","iopub.execute_input":"2024-08-28T22:44:32.658072Z","iopub.status.idle":"2024-08-28T22:44:33.576996Z","shell.execute_reply.started":"2024-08-28T22:44:32.658045Z","shell.execute_reply":"2024-08-28T22:44:33.575787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Bounding Box Annotations\n\nAll images are rescaled to have shape (1024x1024). The bounding box coordinates still have to be adapted to this new image shape. image_size.csv contains the original image size and is used here to rescale the boudning box coordinates.  \nCheck explanations [here](https://www.kaggle.com/competitions/amia-public-challenge-2024/data).","metadata":{}},{"cell_type":"code","source":"# Updating read_and_process_annotations function\ndef read_and_process_annotations(partition):\n    assert partition in [\"train\", \"test\"]\n    \n    # This dataframe contains the original image sizes.\n    image_sizes = pd.read_csv(ROOT / \"img_size.csv\")\n    image_sizes.rename({\"dim0\": \"original_image_height\", \"dim1\": \"original_image_width\"},\n                       axis=1, inplace=True)\n    \n    # Read dataframe and merge with image_size information\n    df = pd.read_csv(ROOT / f\"{partition}.csv\")\n    \n    df.head(3)\n    \n    # Check if 'class_id' exists in DataFrame\n    if 'class_id' not in df.columns:\n        raise KeyError(f\"Column 'class_id' not found in the dataset. Available columns are: {df.columns}\")\n    \n    df = df.merge(image_sizes)\n    \n    # Normalize the bounding box coordinates to the resized images of shape (1024, 1024)\n    if partition in [\"train\", \"val\"]:\n        df[\"x_min_norm\"] = (df.x_min / df.original_image_width) * 1024\n        df[\"x_max_norm\"] = (df.x_max / df.original_image_width) * 1024\n        df[\"y_min_norm\"] = (df.y_min / df.original_image_height) * 1024\n        df[\"y_max_norm\"] = (df.y_max / df.original_image_height) * 1024\n\n        # Compute the bounding box width and height\n        df[\"width\"] = df.x_max_norm - df.x_min_norm \n        df[\"height\"] = df.y_max_norm - df.y_min_norm\n    \n    # Map original classes to new classes\n    df[\"new_class_id\"] = map_classes(df[\"class_id\"])\n    \n    return df","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:51:06.425361Z","iopub.execute_input":"2024-08-28T22:51:06.426055Z","iopub.status.idle":"2024-08-28T22:51:06.434951Z","shell.execute_reply.started":"2024-08-28T22:51:06.426021Z","shell.execute_reply":"2024-08-28T22:51:06.434012Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = read_and_process_annotations(\"train\")\n# test_df = read_and_process_annotations(\"test\")","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:51:43.224011Z","iopub.execute_input":"2024-08-28T22:51:43.224691Z","iopub.status.idle":"2024-08-28T22:51:43.372821Z","shell.execute_reply.started":"2024-08-28T22:51:43.224661Z","shell.execute_reply":"2024-08-28T22:51:43.372000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:51:46.484730Z","iopub.execute_input":"2024-08-28T22:51:46.485358Z","iopub.status.idle":"2024-08-28T22:51:46.508311Z","shell.execute_reply.started":"2024-08-28T22:51:46.485329Z","shell.execute_reply":"2024-08-28T22:51:46.506863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:44:36.014997Z","iopub.status.idle":"2024-08-28T22:44:36.015337Z","shell.execute_reply.started":"2024-08-28T22:44:36.015170Z","shell.execute_reply":"2024-08-28T22:44:36.015184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Class Annotations","metadata":{}},{"cell_type":"code","source":"# For every sample we have multiple labels. This is a multi-label classification task!\n# Additionally we have labels from multiple radiologists (rad_id). \n# Every entry corresponds to a single bounding box.\n\ntrain_df[train_df.image_id == \"0FDQVdLgDKI1sRnPL94LzVh9EvXDVM9m\"]","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:51:59.012641Z","iopub.execute_input":"2024-08-28T22:51:59.013353Z","iopub.status.idle":"2024-08-28T22:51:59.048050Z","shell.execute_reply.started":"2024-08-28T22:51:59.013318Z","shell.execute_reply":"2024-08-28T22:51:59.046861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# We have 15 unique classes!\nsorted(train_df.class_id.unique())","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:51:59.050194Z","iopub.execute_input":"2024-08-28T22:51:59.050786Z","iopub.status.idle":"2024-08-28T22:51:59.058829Z","shell.execute_reply.started":"2024-08-28T22:51:59.050754Z","shell.execute_reply":"2024-08-28T22:51:59.057832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Mapping of class ids to class names\ntrain_df.groupby(\"class_id\").class_name.first().to_dict()","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:51:59.060156Z","iopub.execute_input":"2024-08-28T22:51:59.060486Z","iopub.status.idle":"2024-08-28T22:51:59.078465Z","shell.execute_reply.started":"2024-08-28T22:51:59.060457Z","shell.execute_reply":"2024-08-28T22:51:59.077625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Exploration\n\nIdeas for classification labels:\n- Why are the more counts for some classes then we have samples?\n- Do we have class imbalance?\n- What is the frequency of every class?\n- Can we see classes accuring often together?\n\nIdeas for bounding boxes:\n- What are the size of the boxes for different classes?\n- Are the boxes equally distributed across the images? Any differences per class?\n- How do the outliers look?","metadata":{}},{"cell_type":"code","source":"counts = train_df['class_name'].value_counts()\n\nsns.barplot(y=counts.index, x=counts.values, color=\"gray\")\nplt.xlabel(\"Count\")\nplt.ylabel(\"\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:51:59.079850Z","iopub.execute_input":"2024-08-28T22:51:59.080425Z","iopub.status.idle":"2024-08-28T22:51:59.362842Z","shell.execute_reply.started":"2024-08-28T22:51:59.080394Z","shell.execute_reply":"2024-08-28T22:51:59.361939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"g = sns.FacetGrid(train_df, col=\"class_name\", col_wrap=5)\ng.map(sns.scatterplot, \"width\", \"height\", alpha=0.2);","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:51:59.365867Z","iopub.execute_input":"2024-08-28T22:51:59.366263Z","iopub.status.idle":"2024-08-28T22:52:04.683020Z","shell.execute_reply.started":"2024-08-28T22:51:59.366233Z","shell.execute_reply":"2024-08-28T22:52:04.682138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset class","metadata":{}},{"cell_type":"code","source":"# Update one hot encoding function\ndef to_one_hot_encoded(class_indices, num_classes=3) -> torch.Tensor:\n    one_hot_encoded = torch.zeros(num_classes)\n    one_hot_encoded[torch.unique(torch.tensor(class_indices))] = 1.\n    return one_hot_encoded\n\ndef plot_bounding_boxes(image, bounding_boxes, class_ids=None, ax=None):\n    \n    if ax is None:\n        plt.imshow(image, cmap=\"gray\")\n        ax = plt.gca()\n    else:\n        ax.imshow(image, cmap=\"gray\")\n    \n    if class_ids is not None:\n        colors = [mpl.colormaps[\"tab20b\"](i) for i in class_ids]\n    \n    for i, bbox in enumerate(bounding_boxes):\n        xmin, ymin, xmax, ymax = bbox\n        width = xmax - xmin\n        height = ymax - ymin\n        rect = plt.Rectangle((xmin, ymin), width, height, fill=False, edgecolor=colors[i], linewidth=2)\n        ax.add_patch(rect)\n        if class_ids is not None:\n            class_id = class_ids[i]\n            ax.text(xmin, ymin - 10, f'{CLASS_MAPPING[class_id]}', bbox=dict(facecolor=colors[i], alpha=0.2), fontsize=12, color='white')\n    \n    return ax","metadata":{"execution":{"iopub.status.busy":"2024-08-28T23:01:46.227592Z","iopub.execute_input":"2024-08-28T23:01:46.228533Z","iopub.status.idle":"2024-08-28T23:01:46.239162Z","shell.execute_reply.started":"2024-08-28T23:01:46.228496Z","shell.execute_reply":"2024-08-28T23:01:46.238195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We create a simplified version of the Dataset Class that only returns the one-hot encoded labels.\n\nAs bounding boxes have no use currently in classification...","metadata":{}},{"cell_type":"code","source":"# Update the dataset class for one-hot encoding\nclass AMIADatasetOneHot:\n    def __init__(self, partition=\"train\", transform=None):\n          \n        self.num_classes = 3  # Only three classes now\n\n        self.data = read_and_process_annotations(partition)\n        \n        self.images = ls(ROOT / f\"{partition}/{partition}\")\n        \n        self.transform = transform\n    \n    def __len__(self) -> int:\n        return len(self.images)\n    \n    def __getitem__(self, idx: int) -> Tuple[np.ndarray, List[int], Optional[torch.Tensor]]:\n        \n        # Get sample by image_id\n        image_path = self.images[idx]\n        image_id = image_path.stem\n        image = np.array(Image.open(image_path))\n\n        sample = self.data[self.data.image_id == image_id]\n\n        # Generate one hot encode target vector\n        labels = to_one_hot_encoded(sample.new_class_id.unique(), num_classes=3)\n        \n        if self.transform is not None:\n            transformed = self.transform(image=image)\n            image = transformed['image']\n            \n        image = torch.unsqueeze(torch.tensor(image), 0)  # add channel dim\n\n        return image, labels","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:52:04.696198Z","iopub.execute_input":"2024-08-28T22:52:04.696577Z","iopub.status.idle":"2024-08-28T22:52:04.707494Z","shell.execute_reply.started":"2024-08-28T22:52:04.696545Z","shell.execute_reply":"2024-08-28T22:52:04.706391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = AMIADatasetOneHot()\nimage, ohc_labels = dataset[1]\n\nprint(image.size())\nprint(ohc_labels)\n\ndel dataset","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:52:04.708639Z","iopub.execute_input":"2024-08-28T22:52:04.708938Z","iopub.status.idle":"2024-08-28T22:52:04.985843Z","shell.execute_reply.started":"2024-08-28T22:52:04.708914Z","shell.execute_reply":"2024-08-28T22:52:04.984836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Adding augmentations using *albumentation*.\nFor more info check:\n- [Overview of augmentations](https://albumentations.ai/docs/getting_started/transforms_and_targets/)\n- [Bounding box tutorial](https://albumentations.ai/docs/getting_started/bounding_boxes_augmentation/)","metadata":{"execution":{"iopub.status.busy":"2024-04-23T21:29:09.047218Z","iopub.execute_input":"2024-04-23T21:29:09.04779Z","iopub.status.idle":"2024-04-23T21:29:09.056893Z","shell.execute_reply.started":"2024-04-23T21:29:09.047754Z","shell.execute_reply":"2024-04-23T21:29:09.054798Z"}}},{"cell_type":"code","source":"# Prepare dataset for training\ntrain_transforms = A.Compose([\n    A.RandomCrop(width=256, height=256), \n    A.HorizontalFlip(p=0.5),\n    A.RandomBrightnessContrast(p=0.2),\n])\n\ntrain_dataset = AMIADatasetOneHot(transform=train_transforms)\n\n# As we do not have a separate val folder, use some of the train data for validation\ntotal_size = len(train_dataset)\ntrain_size = int(0.9 * total_size)  # 90% for training\nval_size = total_size - train_size  # 10% for validation\n\n# Split the dataset\ntrain_dataset, val_dataset = torch.utils.data.random_split(train_dataset, [train_size, val_size])\n\ntrain_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=4)\nprint(f\"Train dataset size: {len(train_dataset)}\")\n\nval_loader = DataLoader(val_dataset, batch_size=64, shuffle=True, num_workers=4)\nprint(f\"Validation dataset size: {len(val_dataset)}\")\n\n# test_dataset = AMIADatasetOneHot(partition=\"test\", transform=None) # used only for the challenge, not here\n# test_loader = DataLoader(test_dataset, batch_size=64, num_workers=4)\n# print(f\"Test dataset size: {len(test_dataset)}\")\n\nimage, class_ids = train_dataset[1]\nprint(\"\\nSample\")\nprint(image.shape)\nprint(class_ids)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:52:52.552532Z","iopub.execute_input":"2024-08-28T22:52:52.552915Z","iopub.status.idle":"2024-08-28T22:52:52.788839Z","shell.execute_reply.started":"2024-08-28T22:52:52.552885Z","shell.execute_reply":"2024-08-28T22:52:52.787849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Classification training using *ResNet-50*.","metadata":{}},{"cell_type":"code","source":"# Define train function for classification\ndef train_classification(model, max_epochs=10, val_interval=2, save=False):\n    epoch_loss_values = []\n    epoch_acc_values = []\n    val_loss_values = []\n    val_acc_values = []\n    writer = SummaryWriter()\n    metric = MultilabelAccuracy(num_labels=3).to(device)\n    best_val_loss = float('inf')\n    best_epoch = -1\n\n    print(f\"Training using: {device}\")\n    start_time = time.time()\n    for epoch in range(max_epochs):\n        print(f\"\\n---------------- Epoch {epoch + 1}/{max_epochs} ----------------\")\n        model.train()\n        epoch_loss = 0\n        epoch_acc = 0\n        step = 0\n        now = time.time()\n        for batch_data in train_loader:\n            step += 1\n            inputs, labels = batch_data[0].to(device), batch_data[1].to(device)\n            optimizer.zero_grad()\n            outputs = model(inputs.float())\n            loss = loss_function(outputs, labels)\n            loss.backward()\n            optimizer.step()\n            epoch_loss += loss.item()\n            acc = metric(outputs, labels)\n            epoch_acc += acc.item()\n            if (step % 10 == 0) or (step == 1):\n                print(f\"{step}/{len(train_dataset) // train_loader.batch_size}, loss: {loss.item():.4f}, accuracy: {acc.item():.4f}\")\n            \n        running_time = time.time() - now\n        epoch_loss /= step\n        epoch_loss_values.append(epoch_loss)\n        epoch_acc /= step\n        epoch_acc_values.append(epoch_acc)\n        writer.add_scalar(\"epoch_accuracy\", epoch_acc, epoch)\n        writer.add_scalar(\"epoch_loss\", epoch_loss, epoch)\n        print(f\"epoch {epoch + 1}, average loss: {epoch_loss:.4f}, average accuracy: {epoch_acc:.4f}, time taken: {running_time/60:.2f} mins\")\n\n        if (epoch+1) % val_interval == 0:\n            print(f\"=========== Validation Epoch ===========\")\n            model.eval()\n            val_loss = 0.0\n            val_acc = 0.0\n            step = 0\n            now = time.time()\n            with torch.no_grad():\n                for batch_data in val_loader:\n                    step += 1\n                    inputs, labels = batch_data[0].to(device), batch_data[1].to(device)\n                    outputs = model(inputs.float())\n                    loss = loss_function(outputs, labels)\n                    acc = metric(outputs, labels)\n                    val_loss += loss.item()\n                    val_acc += acc.item()\n                    if (step % 10 == 0) or (step == 1):\n                        print(f\"{step}/{len(val_dataset) // val_loader.batch_size}, loss: {loss.item():.4f}, accuracy: {acc.item():.4f}\")\n\n            running_time = time.time() - now\n            val_loss /= step\n            val_loss_values.append(val_loss)\n            val_acc /= step\n            val_acc_values.append(val_acc)\n            writer.add_scalar(\"val_accuracy\", val_acc, epoch)\n            writer.add_scalar(\"val_loss\", val_loss, epoch)\n            print(f\"Validation epoch {epoch + 1}, val_loss: {val_loss:.4f}, val_accuracy: {val_acc:.4f}, time taken: {running_time/60:.2f} mins\")\n\n            if val_loss < best_val_loss:\n                best_val_loss = val_loss\n                best_model_wts = copy.deepcopy(model.state_dict())\n                best_epoch = epoch+1\n        \n    end_time = time.time()  \n    total_time = end_time - start_time\n    print(f\"\\nClassification Training Completed! Total time taken: {total_time/60:.2f} mins\\n\")\n    writer.close()\n    \n    if save:\n        print(f\"Best Model Saved!! with loss: {best_val_loss} at epoch:{best_epoch}\")\n        torch.save(best_model_wts, \"/kaggle/working/best_resnet50.pth\")\n        \n    return model, epoch_loss_values, epoch_acc_values, val_loss_values, val_acc_values","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:52:52.916024Z","iopub.execute_input":"2024-08-28T22:52:52.916861Z","iopub.status.idle":"2024-08-28T22:52:52.934725Z","shell.execute_reply.started":"2024-08-28T22:52:52.916824Z","shell.execute_reply":"2024-08-28T22:52:52.933611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plotting function\ndef plott(loss_values, acc_values, mode='Train'):\n    fig, ax1 = plt.subplots()\n\n    epochs = range(1, len(loss_values)+1)\n    ax1.set_xlabel('epoch')\n    ax1.set_ylabel(f'{mode} Loss', color='tab:red')\n    ax1.plot(epochs, loss_values, color='tab:red', marker='o', label=f'{mode} loss')\n    ax1.tick_params(axis='y', labelcolor='tab:red')\n\n    ax2 = ax1.twinx()\n    ax2.set_ylabel(f'{mode} Accuracy', color='tab:blue')\n    ax2.plot(epochs, acc_values, color='tab:blue', marker='o', label=f'{mode} accuracy')\n    ax2.tick_params(axis='y', labelcolor='tab:blue')\n\n    plt.title(f'{mode} Loss and Accuracy over Epochs')\n\n    fig.tight_layout() \n    ax1.legend(loc='upper left')\n    ax2.legend(loc='upper right')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:52:53.093848Z","iopub.execute_input":"2024-08-28T22:52:53.094143Z","iopub.status.idle":"2024-08-28T22:52:53.101621Z","shell.execute_reply.started":"2024-08-28T22:52:53.094121Z","shell.execute_reply":"2024-08-28T22:52:53.100570Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Modifying ResNet-50 output layer for 3 classes\nnum_classes = 3\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = models.resnet50(weights=models.ResNet50_Weights.DEFAULT).to(device)\n\nsummary(model)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:52:53.103172Z","iopub.execute_input":"2024-08-28T22:52:53.103452Z","iopub.status.idle":"2024-08-28T22:52:54.806163Z","shell.execute_reply.started":"2024-08-28T22:52:53.103423Z","shell.execute_reply":"2024-08-28T22:52:54.805230Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We now change the first conv layer to accept inputs having only 1 channel dim instead of 3, and change the last fully connected layer to the no. of output classes (15 in our case).","metadata":{}},{"cell_type":"code","source":"# Modify the first conv layer and the last fully connected layer\nmodel.conv1 = nn.Conv2d(1, model.conv1.out_channels, kernel_size=model.conv1.kernel_size, stride=model.conv1.stride, padding=model.conv1.padding, bias=False)\nmodel.fc = nn.Linear(model.fc.in_features, num_classes)\nmodel = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:52:54.807974Z","iopub.execute_input":"2024-08-28T22:52:54.808276Z","iopub.status.idle":"2024-08-28T22:52:54.818555Z","shell.execute_reply.started":"2024-08-28T22:52:54.808251Z","shell.execute_reply":"2024-08-28T22:52:54.817751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We now freeze these layers","metadata":{}},{"cell_type":"code","source":"for name, param in model.named_parameters():\n    if 'conv1' in name or 'bn1' in name or 'fc' in name:\n        param.requires_grad = False","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:52:54.819848Z","iopub.execute_input":"2024-08-28T22:52:54.820136Z","iopub.status.idle":"2024-08-28T22:52:54.825055Z","shell.execute_reply.started":"2024-08-28T22:52:54.820113Z","shell.execute_reply":"2024-08-28T22:52:54.824150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Start Training","metadata":{}},{"cell_type":"code","source":"loss_function = torch.nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), 1e-5)\n\n# Train the ResNet-50 model\nmodel, epoch_loss_values, epoch_acc_values, val_loss_values, val_acc_values = train_classification(model, max_epochs=5, val_interval=2)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:52:54.827232Z","iopub.execute_input":"2024-08-28T22:52:54.827494Z","iopub.status.idle":"2024-08-28T22:58:00.293750Z","shell.execute_reply.started":"2024-08-28T22:52:54.827473Z","shell.execute_reply":"2024-08-28T22:58:00.291367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot training and validation results\nplott(epoch_loss_values, epoch_acc_values)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:58:10.826314Z","iopub.execute_input":"2024-08-28T22:58:10.826706Z","iopub.status.idle":"2024-08-28T22:58:10.872452Z","shell.execute_reply.started":"2024-08-28T22:58:10.826674Z","shell.execute_reply":"2024-08-28T22:58:10.870594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plott(val_loss_values, val_acc_values, mode=\"Validation\")","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:58:10.873169Z","iopub.status.idle":"2024-08-28T22:58:10.873566Z","shell.execute_reply.started":"2024-08-28T22:58:10.873355Z","shell.execute_reply":"2024-08-28T22:58:10.873369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Unfreeze them and train again\n","metadata":{}},{"cell_type":"code","source":"for param in model.parameters():\n    param.requires_grad = True\n    \nmodel, epoch_loss_values, epoch_acc_values, val_loss_values, val_acc_values = train_classification(model, max_epochs=5, val_interval=2)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:58:10.875542Z","iopub.status.idle":"2024-08-28T22:58:10.876054Z","shell.execute_reply.started":"2024-08-28T22:58:10.875783Z","shell.execute_reply":"2024-08-28T22:58:10.875820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plott(epoch_loss_values, epoch_acc_values)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:58:10.877239Z","iopub.status.idle":"2024-08-28T22:58:10.877890Z","shell.execute_reply.started":"2024-08-28T22:58:10.877602Z","shell.execute_reply":"2024-08-28T22:58:10.877624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plott(val_loss_values, val_acc_values, mode=\"Validation\")","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:58:10.879410Z","iopub.status.idle":"2024-08-28T22:58:10.879907Z","shell.execute_reply.started":"2024-08-28T22:58:10.879636Z","shell.execute_reply":"2024-08-28T22:58:10.879656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Bounding Box Regression","metadata":{}},{"cell_type":"markdown","source":"We use the `AMIADataset` class to get our train, and test datasets and split our train to get val","metadata":{}},{"cell_type":"code","source":"# Update the AMIADataset class for bounding box regression\nclass AMIADataset:\n    def __init__(self, partition=\"train\", iou_threshold=0.7, transform=None):\n          \n        self.num_classes = 3  # Update to 3 classes        \n        self.iou_threshold = iou_threshold\n\n        self.data = read_and_process_annotations(partition)\n    \n        self.images = ls(ROOT / f\"{partition}/{partition}\")\n        \n        self.transform = transform\n    \n    def __len__(self) -> int:\n        return len(self.images)\n    \n    def _read_bounding_box(self, sample: pd.Series) -> Tuple[torch.Tensor, List]:\n        \"\"\"Process bounding boxes for a given sample.\"\"\"\n\n        bboxes = sample.loc[:, [\"x_min_norm\", \"y_min_norm\", \"x_max_norm\", \"y_max_norm\"]]\n        bboxes = torch.tensor(bboxes.values).float() / 1024 # Normalize to 0 and 1.\n        \n        # Apply Non-maximum Suppression to remove highly overlapping bounding boxes.\n        boxes_to_keep = nms(bboxes, torch.ones(len(bboxes)), self.iou_threshold)\n        bboxes = bboxes[boxes_to_keep]\n        \n        return bboxes, boxes_to_keep.tolist()\n    \n    def __getitem__(self, idx: int) -> Tuple[np.ndarray, List[int], Optional[torch.Tensor]]:\n        \n        # Get sample by image_id\n        image_path = self.images[idx]\n        image_id = image_path.stem\n        image = np.array(Image.open(image_path))\n\n        sample = self.data[self.data.image_id == image_id]\n\n        # Generate one hot encode target vector\n        labels = to_one_hot_encoded(sample.new_class_id.unique(), num_classes=3)\n\n        # Generate bounding box list           \n        bboxes, boxes_to_keep = self._read_bounding_box(sample)\n\n        # Only keep class labels of kept bounding boxes\n        class_labels = sample.new_class_id.iloc[boxes_to_keep].tolist()\n\n        has_bbox = torch.isnan(bboxes).sum() == 0\n\n        # Generate fake bbox for transformation\n        if not has_bbox:\n            bboxes = torch.zeros(bboxes.shape[0], 4)\n            bboxes[:, 2:] += 0.1\n\n        if self.transform is not None:\n            transformed = self.transform(image=image, bboxes=bboxes, class_labels=class_labels)\n            image = transformed['image']\n            bboxes = torch.tensor(transformed['bboxes'])\n\n        # For no findings, create a class_labels of size 1 and bbox of size (1, 4) containing zeros\n        if not has_bbox:\n            class_labels = [class_labels[0]]\n            bboxes = torch.zeros(1, 4) \n            bboxes[:, 2:] += 0.01\n            \n        # Prepare target dictionary for Fast R-CNN\n        target = {}\n        target['boxes'] = bboxes  # (N, 4)\n        target['labels'] = torch.tensor(class_labels, dtype=torch.int64)  # (N,)\n            \n        return image, target","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:58:43.388473Z","iopub.execute_input":"2024-08-28T22:58:43.389465Z","iopub.status.idle":"2024-08-28T22:58:43.405459Z","shell.execute_reply.started":"2024-08-28T22:58:43.389426Z","shell.execute_reply":"2024-08-28T22:58:43.404393Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def custom_collate_fn(batch):\n    \"\"\"\n    Custom collate function to handle batches with varying numbers of bounding boxes.\n    \"\"\"\n    images = [item[0] for item in batch]\n    targets = [item[1] for item in batch]\n    return images, targets","metadata":{"execution":{"iopub.status.busy":"2024-08-28T23:06:37.972127Z","iopub.execute_input":"2024-08-28T23:06:37.972844Z","iopub.status.idle":"2024-08-28T23:06:37.978015Z","shell.execute_reply.started":"2024-08-28T23:06:37.972806Z","shell.execute_reply":"2024-08-28T23:06:37.977047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Remove the previous datasets and dataloaders\n# del model\n# del train_dataset\n# del val_dataset\n# del test_dataset\n# del train_loader\n# del val_loader\n\n# Prepare dataset for bounding box training\ntrain_transforms = A.Compose([\n    A.Resize(width=256, height=256), \n    A.HorizontalFlip(p=0.5),\n    A.RandomBrightnessContrast(p=0.2),\n], bbox_params=A.BboxParams(format='albumentations', label_fields=['class_labels'], min_visibility=0.2))\n\ntrain_dataset = AMIADataset(transform=train_transforms, iou_threshold=0.7)\n\n# Split the dataset for bounding box regression\ntotal_size = len(train_dataset)\ntrain_size = int(0.9 * total_size)  # 90% for training\nval_size = total_size - train_size  # 10% for validation\n\ntrain_dataset, val_dataset = torch.utils.data.random_split(train_dataset, [train_size, val_size])\n\n# Update DataLoader to use the custom collate function\ntrain_loader = DataLoader(train_dataset, batch_size=2, shuffle=True, num_workers=4, collate_fn=custom_collate_fn)\nval_loader = DataLoader(val_dataset, batch_size=2, shuffle=True, num_workers=4, collate_fn=custom_collate_fn)\n\nprint(f\"Train dataset size: {len(train_dataset)}\")\nprint(f\"Validation dataset size: {len(val_dataset)}\")","metadata":{"execution":{"iopub.status.busy":"2024-08-28T23:06:39.860529Z","iopub.execute_input":"2024-08-28T23:06:39.860936Z","iopub.status.idle":"2024-08-28T23:06:40.050620Z","shell.execute_reply.started":"2024-08-28T23:06:39.860906Z","shell.execute_reply":"2024-08-28T23:06:40.049278Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image, target = train_dataset[1]\n\nprint(target['boxes'])\nprint(target['labels'])","metadata":{"execution":{"iopub.status.busy":"2024-08-28T22:59:47.484394Z","iopub.execute_input":"2024-08-28T22:59:47.484698Z","iopub.status.idle":"2024-08-28T22:59:47.535534Z","shell.execute_reply.started":"2024-08-28T22:59:47.484667Z","shell.execute_reply":"2024-08-28T22:59:47.534589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_bounding_boxes(image, target['boxes'] * 256, target['labels'].numpy());","metadata":{"execution":{"iopub.status.busy":"2024-08-28T23:01:56.686236Z","iopub.execute_input":"2024-08-28T23:01:56.686904Z","iopub.status.idle":"2024-08-28T23:01:57.384249Z","shell.execute_reply.started":"2024-08-28T23:01:56.686874Z","shell.execute_reply":"2024-08-28T23:01:57.383217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As each image here has a different number of bounding boxes, training with a simple ResNet50 would be very complicated. One way of circumventing this would be to cosider multiple samples of the same image, each sample having only one `class_id` and `bounding_box` associated with it\n\nAnother way would be to train using Faster R-CNN as it already handles variable number of bounding boxes.","metadata":{}},{"cell_type":"code","source":"# Go with the Faster R-CNN approach\n# Update the model to have 3 classes for Faster R-CNN\nimport torchvision\nfrom torchvision.models.detection.faster_rcnn import FastRCNNPredictor\nfrom torchvision.models.detection import FasterRCNN_ResNet50_FPN_Weights\nimport torch.nn.utils as utils\n\n# Load a pre-trained model for classification and return only the features\nmodel = torchvision.models.detection.fasterrcnn_resnet50_fpn(weights=FasterRCNN_ResNet50_FPN_Weights.DEFAULT)\n\n# Get the number of input features for the classifier\nin_features = model.roi_heads.box_predictor.cls_score.in_features\n\n# Update the number of classes\nnum_classes = 3  # Update to 3 classes\n\n# Replace the pre-trained head with a new one\nmodel.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)\n\ndevice = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T23:02:52.783577Z","iopub.execute_input":"2024-08-28T23:02:52.784466Z","iopub.status.idle":"2024-08-28T23:02:53.642404Z","shell.execute_reply.started":"2024-08-28T23:02:52.784432Z","shell.execute_reply":"2024-08-28T23:02:53.641511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_bounding_box(model, max_epochs=10, val_interval=2):\n    train_losses = [] \n    writer = SummaryWriter()\n\n    print(f\"Training using: {device}\")\n    start_time = time.time()\n    \n    for epoch in range(max_epochs):\n        # Training phase\n        print(f\"\\n---------------- Epoch {epoch + 1}/{max_epochs} ----------------\")\n        model.train()\n        epoch_loss = 0\n        step = 0\n        now = time.time()\n        for images, targets in train_loader:\n            # Convert images from numpy to PyTorch tensor and add channel dimension\n            images = [torch.tensor(image, dtype=torch.float).unsqueeze(0).to(device) for image in images]\n            targets = [{k: v.to(device) for k, v in t.items()} for t in targets]\n\n            optimizer.zero_grad()\n\n            # Forward pass\n            loss_dict = model(images, targets)\n\n            # Compute the total loss\n            losses = sum(loss for loss in loss_dict.values())\n            epoch_loss += losses.item()\n            \n            losses.backward()\n            utils.clip_grad_norm_(model.parameters(), max_norm=2.0) # prevent gradient explosion\n            optimizer.step()\n            step += 1\n            \n            if (step % 500 == 0) or (step == 1):\n                print(f\"{step}/{len(train_dataset) // 2}, loss: {losses.item():.4f}\")\n              \n            # Training takes too much time (approx 30mins/epoch for the full dataset)\n            # Hence, only train till half\n            if step >= (len(train_dataset) // 2):\n                print(f\"{step}/{len(train_dataset) // 2}, loss: {losses.item():.4f}\")\n                break\n            \n        # Update learning rate\n        lr_scheduler.step()\n            \n        running_time = time.time() - now\n        epoch_loss /= step\n        train_losses.append(epoch_loss)\n        writer.add_scalar(\"epoch_loss\", epoch_loss, epoch)\n        print(f\"epoch {epoch + 1}, average loss: {epoch_loss:.4f}, time taken: {running_time/60:.2f} mins\")\n            \n        # Validation phase\n        if (epoch+1) % val_interval == 0:\n            print(f\"=========== Validation Epoch ===========\")\n            model.eval()\n            step = 0\n            \n            # Get 5 images from the val set to show during validation\n            indices = [int(i * len(val_dataset) / 5) + 1 for i in range(5)]\n\n            with torch.no_grad():\n                for images, targets in val_loader:\n                    # Convert images from numpy to PyTorch tensor and add channel dimension\n                    images = [torch.tensor(image, dtype=torch.float).unsqueeze(0).to(device) for image in images]\n                    targets = [{k: v.to(device) for k, v in t.items()} for t in targets]\n\n                    # Get predictions\n                    predictions = model(images)\n\n                    step += 1\n                    \n                    if step in indices:\n                        fig, axs = plt.subplots(1, 2, figsize=(12, 6))\n\n                        # Plot true bounding boxes\n                        axs[0] = plot_bounding_boxes(images[0].squeeze().cpu(), targets[0]['boxes'].cpu() * 256, class_ids=targets[0]['labels'].cpu().numpy(), ax=axs[0])\n                        axs[0].set_title(\"True\")\n\n                        # Plot predicted bounding boxes\n                        axs[1] = plot_bounding_boxes(images[0].squeeze().cpu(), predictions[0]['boxes'].cpu(), class_ids=predictions[0]['labels'].cpu().numpy(), ax=axs[1])\n                        axs[1].set_title(\"Predicted\")\n\n                        plt.suptitle(f\"Image {step}\\n\")\n                        plt.show()\n                        \n    end_time = time.time()  \n    total_time = end_time - start_time\n    print(f\"\\nBounding Box Training Completed! Total time taken: {total_time/60:.2f} mins\\n\")\n    writer.close()\n    end_wts = copy.deepcopy(model.state_dict())\n    \n    print(f\"Model Saved!!\")\n    torch.save(end_wts, \"/kaggle/working/trained_frcnn_bbox.pth\")\n        \n    return model, train_losses","metadata":{"execution":{"iopub.status.busy":"2024-08-28T23:09:18.636488Z","iopub.execute_input":"2024-08-28T23:09:18.636886Z","iopub.status.idle":"2024-08-28T23:09:18.656700Z","shell.execute_reply.started":"2024-08-28T23:09:18.636854Z","shell.execute_reply":"2024-08-28T23:09:18.655640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')\nmodel.to(device)\n\n# Update optimizer and learning rate scheduler\noptimizer = torch.optim.Adam(model.parameters(), 1e-5)\nlr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T23:09:18.658458Z","iopub.execute_input":"2024-08-28T23:09:18.658729Z","iopub.status.idle":"2024-08-28T23:09:18.673645Z","shell.execute_reply.started":"2024-08-28T23:09:18.658707Z","shell.execute_reply":"2024-08-28T23:09:18.672844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train the Faster R-CNN model\ntrained_model, epoch_losses = train_bounding_box(model, max_epochs=2, val_interval=1)","metadata":{"execution":{"iopub.status.busy":"2024-08-28T23:09:18.674713Z","iopub.execute_input":"2024-08-28T23:09:18.675033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As one can see, the model predicts way too many bounding boxes and that also poorly.\n\nThe reason for this might be: `(bbox_pred): Linear(in_features=1024, out_features=60, bias=True)` in the `box_predictor` ROI Head. This value is automatically calculated as *num_classes x 4* (4 coords for each class) and changing it did not make any sense...","metadata":{}},{"cell_type":"markdown","source":"The runtime eventually ran out of memory but we do get atleast one nice image that also matches the true output :/ ","metadata":{}}]}