{"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":30698,"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\nCLASS_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-14T17:35:58.932858Z","iopub.execute_input":"2024-08-14T17:35:58.933129Z","iopub.status.idle":"2024-08-14T17:36:07.547059Z","shell.execute_reply.started":"2024-08-14T17:35:58.933107Z","shell.execute_reply":"2024-08-14T17:36:07.546184Z"},"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-14T17:36:07.552080Z","iopub.execute_input":"2024-08-14T17:36:07.552353Z","iopub.status.idle":"2024-08-14T17:36:07.560522Z","shell.execute_reply.started":"2024-08-14T17:36:07.552329Z","shell.execute_reply":"2024-08-14T17:36:07.559727Z"},"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-14T17:36:07.561589Z","iopub.execute_input":"2024-08-14T17:36:07.561836Z","iopub.status.idle":"2024-08-14T17:36:07.592979Z","shell.execute_reply.started":"2024-08-14T17:36:07.561814Z","shell.execute_reply":"2024-08-14T17:36:07.592180Z"},"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-14T17:36:07.595171Z","iopub.execute_input":"2024-08-14T17:36:07.595430Z","iopub.status.idle":"2024-08-14T17:36:07.881580Z","shell.execute_reply.started":"2024-08-14T17:36:07.595409Z","shell.execute_reply":"2024-08-14T17:36:07.880547Z"},"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-14T17:36:07.882964Z","iopub.execute_input":"2024-08-14T17:36:07.883294Z","iopub.status.idle":"2024-08-14T17:36:07.925842Z","shell.execute_reply.started":"2024-08-14T17:36:07.883267Z","shell.execute_reply":"2024-08-14T17:36:07.924907Z"},"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-14T17:36:07.927065Z","iopub.execute_input":"2024-08-14T17:36:07.927406Z","iopub.status.idle":"2024-08-14T17:36:15.203320Z","shell.execute_reply.started":"2024-08-14T17:36:07.927374Z","shell.execute_reply":"2024-08-14T17:36:15.202349Z"},"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-14T17:36:15.204650Z","iopub.execute_input":"2024-08-14T17:36:15.204957Z","iopub.status.idle":"2024-08-14T17:36:16.124846Z","shell.execute_reply.started":"2024-08-14T17:36:15.204931Z","shell.execute_reply":"2024-08-14T17:36:16.123941Z"},"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":"def read_and_proccess_annotations(partition):\n    assert partition in [\"train\", \"test\"]\n    \n    # This dataframe contains the original image sizes.\n    # They are used to normalize the bounding box coordinates.\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    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        # Normalize coordinates accoring to height and width of the original images\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    return df","metadata":{"execution":{"iopub.status.busy":"2024-08-14T17:36:16.126131Z","iopub.execute_input":"2024-08-14T17:36:16.126441Z","iopub.status.idle":"2024-08-14T17:36:16.136977Z","shell.execute_reply.started":"2024-08-14T17:36:16.126417Z","shell.execute_reply":"2024-08-14T17:36:16.133398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = read_and_proccess_annotations(\"train\")\ntest_df = read_and_proccess_annotations(\"test\")","metadata":{"execution":{"iopub.status.busy":"2024-08-14T17:36:16.138266Z","iopub.execute_input":"2024-08-14T17:36:16.138740Z","iopub.status.idle":"2024-08-14T17:36:16.318859Z","shell.execute_reply.started":"2024-08-14T17:36:16.138701Z","shell.execute_reply":"2024-08-14T17:36:16.318088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-08-14T17:36:16.320094Z","iopub.execute_input":"2024-08-14T17:36:16.320427Z","iopub.status.idle":"2024-08-14T17:36:16.341462Z","shell.execute_reply.started":"2024-08-14T17:36:16.320402Z","shell.execute_reply":"2024-08-14T17:36:16.340384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-08-14T17:36:16.342660Z","iopub.execute_input":"2024-08-14T17:36:16.342918Z","iopub.status.idle":"2024-08-14T17:36:16.351360Z","shell.execute_reply.started":"2024-08-14T17:36:16.342896Z","shell.execute_reply":"2024-08-14T17:36:16.350523Z"},"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-14T17:36:16.352644Z","iopub.execute_input":"2024-08-14T17:36:16.353173Z","iopub.status.idle":"2024-08-14T17:36:16.388453Z","shell.execute_reply.started":"2024-08-14T17:36:16.353149Z","shell.execute_reply":"2024-08-14T17:36:16.387643Z"},"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-14T17:36:16.393747Z","iopub.execute_input":"2024-08-14T17:36:16.394004Z","iopub.status.idle":"2024-08-14T17:36:16.400521Z","shell.execute_reply.started":"2024-08-14T17:36:16.393983Z","shell.execute_reply":"2024-08-14T17:36:16.399546Z"},"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-14T17:36:16.401457Z","iopub.execute_input":"2024-08-14T17:36:16.402058Z","iopub.status.idle":"2024-08-14T17:36:16.425366Z","shell.execute_reply.started":"2024-08-14T17:36:16.402027Z","shell.execute_reply":"2024-08-14T17:36:16.424543Z"},"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-14T17:36:16.426346Z","iopub.execute_input":"2024-08-14T17:36:16.426687Z","iopub.status.idle":"2024-08-14T17:36:16.774393Z","shell.execute_reply.started":"2024-08-14T17:36:16.426656Z","shell.execute_reply":"2024-08-14T17:36:16.773400Z"},"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-14T17:36:16.775787Z","iopub.execute_input":"2024-08-14T17:36:16.776243Z","iopub.status.idle":"2024-08-14T17:36:21.934960Z","shell.execute_reply.started":"2024-08-14T17:36:16.776207Z","shell.execute_reply":"2024-08-14T17:36:21.934052Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset class","metadata":{}},{"cell_type":"code","source":"def to_one_hot_encoded(class_indeces, num_classes=15) -> torch.Tensor:\n    one_hot_encoded = torch.zeros(num_classes)\n    one_hot_encoded[torch.unique(torch.tensor(class_indeces))] = 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_IDS_NAMES[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-14T17:36:21.936303Z","iopub.execute_input":"2024-08-14T17:36:21.936645Z","iopub.status.idle":"2024-08-14T17:36:21.948230Z","shell.execute_reply.started":"2024-08-14T17:36:21.936616Z","shell.execute_reply":"2024-08-14T17:36:21.947085Z"},"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":"class AMIADatasetOneHot:\n    def __init__(self, partition=\"train\", transform=None):\n          \n        self.num_classes = 15        \n\n        self.data = read_and_proccess_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.class_id.unique())\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-14T17:36:21.949593Z","iopub.execute_input":"2024-08-14T17:36:21.950228Z","iopub.status.idle":"2024-08-14T17:36:21.964536Z","shell.execute_reply.started":"2024-08-14T17:36:21.950194Z","shell.execute_reply":"2024-08-14T17:36:21.963634Z"},"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-14T17:36:21.965581Z","iopub.execute_input":"2024-08-14T17:36:21.966078Z","iopub.status.idle":"2024-08-14T17:36:22.226085Z","shell.execute_reply.started":"2024-08-14T17:36:21.966054Z","shell.execute_reply":"2024-08-14T17:36:22.225118Z"},"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":"train_transforms = A.Compose([\n    A.RandomCrop(width=256, height=256), # 900 x 900 leads to OutOfMemoryError\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\ntest_dataset = AMIADatasetOneHot(partition=\"test\", transform=None) # used only for the challenge, not here\ntest_loader = DataLoader(test_dataset, batch_size=64, num_workers=4)\nprint(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-14T17:36:22.227142Z","iopub.execute_input":"2024-08-14T17:36:22.227418Z","iopub.status.idle":"2024-08-14T17:36:23.005356Z","shell.execute_reply.started":"2024-08-14T17:36:22.227394Z","shell.execute_reply":"2024-08-14T17:36:23.004169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Classification training using *ResNet-50*.","metadata":{}},{"cell_type":"code","source":"def train(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=15).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                    optimizer.zero_grad()\n                    outputs = model(inputs.float())\n                    loss = loss_function(outputs, labels)\n                    acc = metric(outputs, labels)\n                    optimizer.step()\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        \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-14T17:36:23.007094Z","iopub.execute_input":"2024-08-14T17:36:23.007793Z","iopub.status.idle":"2024-08-14T17:36:23.026770Z","shell.execute_reply.started":"2024-08-14T17:36:23.007759Z","shell.execute_reply":"2024-08-14T17:36:23.025755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def 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-14T17:36:23.028131Z","iopub.execute_input":"2024-08-14T17:36:23.028421Z","iopub.status.idle":"2024-08-14T17:36:23.041500Z","shell.execute_reply.started":"2024-08-14T17:36:23.028398Z","shell.execute_reply":"2024-08-14T17:36:23.040772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes = 15\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-14T17:36:23.042461Z","iopub.execute_input":"2024-08-14T17:36:23.042743Z","iopub.status.idle":"2024-08-14T17:36:24.719471Z","shell.execute_reply.started":"2024-08-14T17:36:23.042721Z","shell.execute_reply":"2024-08-14T17:36:24.718559Z"},"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":"model.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-14T17:36:24.720714Z","iopub.execute_input":"2024-08-14T17:36:24.721015Z","iopub.status.idle":"2024-08-14T17:36:24.730859Z","shell.execute_reply.started":"2024-08-14T17:36:24.720991Z","shell.execute_reply":"2024-08-14T17:36:24.729891Z"},"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-14T17:36:24.732349Z","iopub.execute_input":"2024-08-14T17:36:24.732640Z","iopub.status.idle":"2024-08-14T17:36:24.739208Z","shell.execute_reply.started":"2024-08-14T17:36:24.732614Z","shell.execute_reply":"2024-08-14T17:36:24.738496Z"},"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)\nmax_epochs = 10\nval_interval = 2\n\nmodel, epoch_loss_values, epoch_acc_values, val_loss_values, val_acc_values = train(model, max_epochs, val_interval)","metadata":{"execution":{"iopub.status.busy":"2024-08-14T17:36:24.740404Z","iopub.execute_input":"2024-08-14T17:36:24.741011Z","iopub.status.idle":"2024-08-14T17:51:55.140682Z","shell.execute_reply.started":"2024-08-14T17:36:24.740979Z","shell.execute_reply":"2024-08-14T17:51:55.139515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plott(epoch_loss_values, epoch_acc_values)","metadata":{"execution":{"iopub.status.busy":"2024-08-14T17:51:55.142200Z","iopub.execute_input":"2024-08-14T17:51:55.142528Z","iopub.status.idle":"2024-08-14T17:51:55.617676Z","shell.execute_reply.started":"2024-08-14T17:51:55.142477Z","shell.execute_reply":"2024-08-14T17:51:55.616695Z"},"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-14T17:51:55.618860Z","iopub.execute_input":"2024-08-14T17:51:55.619126Z","iopub.status.idle":"2024-08-14T17:51:56.154084Z","shell.execute_reply.started":"2024-08-14T17:51:55.619103Z","shell.execute_reply":"2024-08-14T17:51:56.153177Z"},"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(model, max_epochs, val_interval, save=True)","metadata":{"execution":{"iopub.status.busy":"2024-08-14T17:51:56.155154Z","iopub.execute_input":"2024-08-14T17:51:56.155437Z","iopub.status.idle":"2024-08-14T18:08:07.232877Z","shell.execute_reply.started":"2024-08-14T17:51:56.155414Z","shell.execute_reply":"2024-08-14T18:08:07.231572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plott(epoch_loss_values, epoch_acc_values)","metadata":{"execution":{"iopub.status.busy":"2024-08-14T18:08:07.234917Z","iopub.execute_input":"2024-08-14T18:08:07.235823Z","iopub.status.idle":"2024-08-14T18:08:07.770316Z","shell.execute_reply.started":"2024-08-14T18:08:07.235778Z","shell.execute_reply":"2024-08-14T18:08:07.769526Z"},"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-14T18:08:07.771375Z","iopub.execute_input":"2024-08-14T18:08:07.771700Z","iopub.status.idle":"2024-08-14T18:08:08.360778Z","shell.execute_reply.started":"2024-08-14T18:08:07.771674Z","shell.execute_reply":"2024-08-14T18:08:08.359832Z"},"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":"class AMIADataset:\n    def __init__(self, partition=\"train\", iou_threshold=0.7, transform=None):\n          \n        self.num_classes = 15        \n        self.iou_threshold = iou_threshold\n\n        self.data = read_and_proccess_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 Supression 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.class_id.unique())\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 bobxes\n        class_labels = sample.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-14T18:08:08.362520Z","iopub.execute_input":"2024-08-14T18:08:08.362898Z","iopub.status.idle":"2024-08-14T18:08:08.377960Z","shell.execute_reply.started":"2024-08-14T18:08:08.362866Z","shell.execute_reply":"2024-08-14T18:08:08.376988Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Remove the previous datasets and dataloaders\ndel model\ndel train_dataset\ndel val_dataset\ndel test_dataset\ndel train_loader\ndel val_loader\ndel test_loader\n\ntrain_transforms = A.Compose([\n    A.Resize(width=256, height=256), # affects accuracy but higher dimensions take too much time to train\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\n# Use the former Dataset class to get new train, val, and test\ntrain_dataset = AMIADataset(transform=train_transforms, iou_threshold=0.7)\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=1, shuffle=True, num_workers=4)\nprint(f\"Train dataset size: {len(train_dataset)}\")\n\nval_loader = DataLoader(val_dataset, batch_size=1, shuffle=True, num_workers=4)\nprint(f\"Validation dataset size: {len(val_dataset)}\")","metadata":{"execution":{"iopub.status.busy":"2024-08-14T18:08:08.379128Z","iopub.execute_input":"2024-08-14T18:08:08.379392Z","iopub.status.idle":"2024-08-14T18:08:08.533176Z","shell.execute_reply.started":"2024-08-14T18:08:08.379369Z","shell.execute_reply":"2024-08-14T18:08:08.532135Z"},"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-14T18:08:08.534307Z","iopub.execute_input":"2024-08-14T18:08:08.534593Z","iopub.status.idle":"2024-08-14T18:08:08.587424Z","shell.execute_reply.started":"2024-08-14T18:08:08.534569Z","shell.execute_reply":"2024-08-14T18:08:08.586538Z"},"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-14T18:08:08.588450Z","iopub.execute_input":"2024-08-14T18:08:08.588761Z","iopub.status.idle":"2024-08-14T18:08:08.841191Z","shell.execute_reply.started":"2024-08-14T18:08:08.588738Z","shell.execute_reply":"2024-08-14T18:08:08.840304Z"},"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\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\nimport torch.nn.functional as F\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# Define the number of classes\nnum_classes = 15\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-14T18:08:08.842456Z","iopub.execute_input":"2024-08-14T18:08:08.842795Z","iopub.status.idle":"2024-08-14T18:08:10.958591Z","shell.execute_reply.started":"2024-08-14T18:08:08.842769Z","shell.execute_reply":"2024-08-14T18:08:10.957721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(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            images = list(image.to(device, dtype=torch.float).unsqueeze(0) for image in images) # add channel dim (C, H, W)\n            targets = [{k: v.to(device) for k, v in t.items()} for t in [targets]]\n            targets[0]['boxes'] = targets[0]['boxes'].squeeze(0) # remove batch dim\n            targets[0]['labels'] = targets[0]['labels'].squeeze(0)\n            targets[0]['boxes'] *= 256 # absolute coords from normalized\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        \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                    images = list(image.to(device, dtype=torch.float).unsqueeze(0) for image in images) # add channel dim (C, H, W)\n                    targets = [{k: v.to(device) for k, v in t.items()} for t in [targets]]\n                    targets[0]['boxes'] = targets[0]['boxes'].squeeze(0) # remove batch dim\n                    targets[0]['labels'] = targets[0]['labels'].squeeze(0)\n                    targets[0]['boxes'] *= 256 # absolute coords from normalized\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                        \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(best_model_wts, \"/kaggle/working/trained_frcnn_bbox.pth\")\n        \n    return model, epoch_losses","metadata":{"execution":{"iopub.status.busy":"2024-08-14T18:34:05.465922Z","iopub.execute_input":"2024-08-14T18:34:05.466346Z","iopub.status.idle":"2024-08-14T18:34:05.489263Z","shell.execute_reply.started":"2024-08-14T18:34:05.466316Z","shell.execute_reply":"2024-08-14T18:34:05.487913Z"},"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\noptimizer = torch.optim.Adam(model.parameters(), 1e-8)\n\nlr_scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.1)\n\nmax_epochs = 5\nval_interval = 1","metadata":{"execution":{"iopub.status.busy":"2024-08-14T18:34:08.353020Z","iopub.execute_input":"2024-08-14T18:34:08.353385Z","iopub.status.idle":"2024-08-14T18:34:08.366335Z","shell.execute_reply.started":"2024-08-14T18:34:08.353354Z","shell.execute_reply":"2024-08-14T18:34:08.365316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trained_model, epoch_losses = train(model, max_epochs, val_interval)","metadata":{"execution":{"iopub.status.busy":"2024-08-14T18:34:09.557503Z","iopub.execute_input":"2024-08-14T18:34:09.558370Z"},"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":{}}]}