{"metadata":{"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":75176,"databundleVersionId":8252256,"sourceType":"competition"},{"sourceId":69985,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":58407,"modelId":80655}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3","language":"python","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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# AMIA Challenge: Explorative Data Analysis, Loading, and Training with MONAI (classification)","metadata":{"execution":{"iopub.execute_input":"2024-04-23T13:50:04.495349Z","iopub.status.busy":"2024-04-23T13:50:04.494867Z","iopub.status.idle":"2024-04-23T13:50:04.5Z","shell.execute_reply":"2024-04-23T13:50:04.498974Z","shell.execute_reply.started":"2024-04-23T13:50:04.495319Z"}}},{"cell_type":"code","source":"# imports, helper functions and globals\nfrom pathlib import Path\nfrom typing import List, Tuple, Optional\n\n# import 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 tqdm.notebook import tqdm, trange\nfrom torchinfo import summary\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":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2024-08-15T11:23:58.491004Z","iopub.execute_input":"2024-08-15T11:23:58.491325Z","iopub.status.idle":"2024-08-15T11:24:05.495480Z","shell.execute_reply.started":"2024-08-15T11:23:58.491298Z","shell.execute_reply":"2024-08-15T11:24:05.494606Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Run this if monai not found\n!python -c \"import monai\" || pip install -q \"monai-weekly[pillow, tqdm]\"","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:24:05.496574Z","iopub.execute_input":"2024-08-15T11:24:05.496986Z","iopub.status.idle":"2024-08-15T11:24:20.568448Z","shell.execute_reply.started":"2024-08-15T11:24:05.496961Z","shell.execute_reply":"2024-08-15T11:24:20.567420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport shutil\nimport tempfile\nimport matplotlib.pyplot as plt\nimport PIL\nimport torch\nfrom torch.utils.tensorboard import SummaryWriter\nimport numpy as np\nfrom sklearn.metrics import classification_report\nfrom torchmetrics.classification import MultilabelAccuracy\nfrom torchinfo import summary\nimport time\nimport copy\n\nfrom monai.apps import download_and_extract\nfrom monai.config import print_config\nfrom monai.data import decollate_batch, DataLoader\nfrom monai.metrics import ROCAUCMetric\nfrom monai.networks.nets import DenseNet121\nfrom monai.transforms import (\n    Activations,\n    EnsureChannelFirst,\n    AsDiscrete,\n    Compose,\n    LoadImage,\n    RandFlip,\n    RandRotate,\n    RandZoom,\n    ScaleIntensity,\n    Resize\n)\nfrom monai.utils import set_determinism\n\nprint_config()","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:24:20.569867Z","iopub.execute_input":"2024-08-15T11:24:20.570195Z","iopub.status.idle":"2024-08-15T11:25:06.274776Z","shell.execute_reply.started":"2024-08-15T11:24:20.570165Z","shell.execute_reply":"2024-08-15T11:25:06.273890Z"},"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-15T11:25:06.277948Z","iopub.execute_input":"2024-08-15T11:25:06.279096Z","iopub.status.idle":"2024-08-15T11:25:06.286173Z","shell.execute_reply.started":"2024-08-15T11:25:06.279042Z","shell.execute_reply":"2024-08-15T11:25:06.285281Z"},"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-15T11:25:06.287621Z","iopub.execute_input":"2024-08-15T11:25:06.288035Z","iopub.status.idle":"2024-08-15T11:25:06.751395Z","shell.execute_reply.started":"2024-08-15T11:25:06.288004Z","shell.execute_reply":"2024-08-15T11:25:06.750518Z"},"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-15T11:25:06.752449Z","iopub.execute_input":"2024-08-15T11:25:06.752711Z","iopub.status.idle":"2024-08-15T11:25:07.122937Z","shell.execute_reply.started":"2024-08-15T11:25:06.752690Z","shell.execute_reply":"2024-08-15T11:25:07.122037Z"},"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-15T11:25:07.124378Z","iopub.execute_input":"2024-08-15T11:25:07.125008Z","iopub.status.idle":"2024-08-15T11:25:07.175787Z","shell.execute_reply.started":"2024-08-15T11:25:07.124973Z","shell.execute_reply":"2024-08-15T11:25:07.174890Z"},"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-15T11:25:07.176856Z","iopub.execute_input":"2024-08-15T11:25:07.177160Z","iopub.status.idle":"2024-08-15T11:25:14.390429Z","shell.execute_reply.started":"2024-08-15T11:25:07.177136Z","shell.execute_reply":"2024-08-15T11:25:14.389503Z"},"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-15T11:25:14.391715Z","iopub.execute_input":"2024-08-15T11:25:14.392565Z","iopub.status.idle":"2024-08-15T11:25:15.242335Z","shell.execute_reply.started":"2024-08-15T11:25:14.392537Z","shell.execute_reply":"2024-08-15T11:25:15.241435Z"},"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 == \"train\":\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\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-15T11:25:15.243537Z","iopub.execute_input":"2024-08-15T11:25:15.243886Z","iopub.status.idle":"2024-08-15T11:25:15.251862Z","shell.execute_reply.started":"2024-08-15T11:25:15.243855Z","shell.execute_reply":"2024-08-15T11:25:15.250904Z"},"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-15T11:25:15.253004Z","iopub.execute_input":"2024-08-15T11:25:15.253331Z","iopub.status.idle":"2024-08-15T11:25:15.506135Z","shell.execute_reply.started":"2024-08-15T11:25:15.253301Z","shell.execute_reply":"2024-08-15T11:25:15.505107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:25:15.507431Z","iopub.execute_input":"2024-08-15T11:25:15.507736Z","iopub.status.idle":"2024-08-15T11:25:15.530883Z","shell.execute_reply.started":"2024-08-15T11:25:15.507710Z","shell.execute_reply":"2024-08-15T11:25:15.530005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:25:15.535473Z","iopub.execute_input":"2024-08-15T11:25:15.535806Z","iopub.status.idle":"2024-08-15T11:25:15.547003Z","shell.execute_reply.started":"2024-08-15T11:25:15.535781Z","shell.execute_reply":"2024-08-15T11:25:15.545971Z"},"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-15T11:25:15.548291Z","iopub.execute_input":"2024-08-15T11:25:15.548576Z","iopub.status.idle":"2024-08-15T11:25:15.578985Z","shell.execute_reply.started":"2024-08-15T11:25:15.548552Z","shell.execute_reply":"2024-08-15T11:25:15.578188Z"},"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-15T11:25:15.580117Z","iopub.execute_input":"2024-08-15T11:25:15.580451Z","iopub.status.idle":"2024-08-15T11:25:15.588137Z","shell.execute_reply.started":"2024-08-15T11:25:15.580420Z","shell.execute_reply":"2024-08-15T11:25:15.587288Z"},"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-15T11:25:15.589264Z","iopub.execute_input":"2024-08-15T11:25:15.589520Z","iopub.status.idle":"2024-08-15T11:25:15.607856Z","shell.execute_reply.started":"2024-08-15T11:25:15.589499Z","shell.execute_reply":"2024-08-15T11:25:15.607096Z"},"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-15T11:25:15.608779Z","iopub.execute_input":"2024-08-15T11:25:15.609041Z","iopub.status.idle":"2024-08-15T11:25:15.939288Z","shell.execute_reply.started":"2024-08-15T11:25:15.609019Z","shell.execute_reply":"2024-08-15T11:25:15.938442Z"},"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-15T11:25:15.940503Z","iopub.execute_input":"2024-08-15T11:25:15.940789Z","iopub.status.idle":"2024-08-15T11:25:21.027889Z","shell.execute_reply.started":"2024-08-15T11:25:15.940764Z","shell.execute_reply":"2024-08-15T11:25:21.026965Z"},"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-15T11:25:21.029231Z","iopub.execute_input":"2024-08-15T11:25:21.029565Z","iopub.status.idle":"2024-08-15T11:25:21.041469Z","shell.execute_reply.started":"2024-08-15T11:25:21.029537Z","shell.execute_reply":"2024-08-15T11:25:21.040378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AMIADatasetOneHot:\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        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        image = torch.unsqueeze(torch.tensor(image), 0) # add channel dim\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)\n            image = transformed\n\n        return image, labels","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:25:21.042822Z","iopub.execute_input":"2024-08-15T11:25:21.043272Z","iopub.status.idle":"2024-08-15T11:25:21.055494Z","shell.execute_reply.started":"2024-08-15T11:25:21.043241Z","shell.execute_reply":"2024-08-15T11:25:21.054561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = AMIADatasetOneHot(iou_threshold=0.7)\nimage, ohc_labels = dataset[1]\n\nprint(image.size())\nprint(ohc_labels)\n\ndel dataset","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:25:21.056798Z","iopub.execute_input":"2024-08-15T11:25:21.057600Z","iopub.status.idle":"2024-08-15T11:25:21.305225Z","shell.execute_reply.started":"2024-08-15T11:25:21.057567Z","shell.execute_reply":"2024-08-15T11:25:21.304235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Adding augmentations using *monai transforms*.\nFor more info check:\n- [Overview of augmentations](https://docs.monai.io/en/latest/transforms.html)","metadata":{"execution":{"iopub.execute_input":"2024-04-23T21:29:09.04779Z","iopub.status.busy":"2024-04-23T21:29:09.047218Z","iopub.status.idle":"2024-04-23T21:29:09.056893Z","shell.execute_reply":"2024-04-23T21:29:09.054798Z","shell.execute_reply.started":"2024-04-23T21:29:09.047754Z"}}},{"cell_type":"code","source":"train_transforms = Compose(\n    [\n        ScaleIntensity(),\n        Resize(spatial_size=(128,128), anti_aliasing=True),\n        RandRotate(range_x=np.pi / 12, prob=0.5, keep_size=True),\n        RandFlip(spatial_axis=0, prob=0.5),\n        RandZoom(min_zoom=0.9, max_zoom=1.1, prob=0.5),\n    ]\n)\n\ndataset = AMIADatasetOneHot(transform=train_transforms, iou_threshold=0.7)\n\n# print(dataset)\n\ntotal_size = len(dataset)\ntrain_size = int(0.8 * total_size)  # 80% for training\nval_size = total_size - train_size  # 20% for validation\n\n# Split the dataset\ntrain_dataset, val_dataset = torch.utils.data.random_split(dataset, [train_size, val_size])\n\ntrain_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=4)\nval_loader = DataLoader(val_dataset, batch_size=64, shuffle=True, num_workers=4)\n\nprint(f\"Train dataset size: {len(train_dataset)}\")\nprint(f\"Validation dataset size: {len(val_dataset)}\")\n\ny_pred_trans = Compose([Activations(softmax=True)])\ny_trans = Compose([AsDiscrete(to_onehot=15)])\n\ntest_dataset = AMIADatasetOneHot(partition=\"test\", transform=None, iou_threshold=0.7)\ntest_loader = DataLoader(test_dataset, batch_size=64, num_workers=4)\nprint(f\"Test dataset size: {len(test_dataset)}\")","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:25:21.306275Z","iopub.execute_input":"2024-08-15T11:25:21.306540Z","iopub.status.idle":"2024-08-15T11:25:21.962811Z","shell.execute_reply.started":"2024-08-15T11:25:21.306517Z","shell.execute_reply":"2024-08-15T11:25:21.961897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Classification training using *monai DenseNet121*.","metadata":{}},{"cell_type":"code","source":"num_class = dataset.num_classes\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = DenseNet121(spatial_dims=2, in_channels=1, out_channels=num_class).to(device)\nloss_function = torch.nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), 1e-5)\n\nauc_metric = ROCAUCMetric()","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:25:21.963896Z","iopub.execute_input":"2024-08-15T11:25:21.964196Z","iopub.status.idle":"2024-08-15T11:25:22.452135Z","shell.execute_reply.started":"2024-08-15T11:25:21.964170Z","shell.execute_reply":"2024-08-15T11:25:22.451006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(model)","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:25:22.453339Z","iopub.execute_input":"2024-08-15T11:25:22.453616Z","iopub.status.idle":"2024-08-15T11:25:22.510283Z","shell.execute_reply.started":"2024-08-15T11:25:22.453593Z","shell.execute_reply":"2024-08-15T11:25:22.509356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class EarlyStopping:\n    \"\"\"Early stops the training if validation loss doesn't improve after a given patience.\"\"\"\n    def __init__(self, patience=5, verbose=False, delta=0.01, path='checkpoint.pt', trace_func=print):\n        \"\"\"\n        Args:\n            patience (int): How long to wait after last time validation loss improved.\n                            Default: 5\n            verbose (bool): If True, prints a message for each validation loss improvement. \n                            Default: False\n            delta (float): Minimum change in the monitored quantity to qualify as an improvement.\n                            Default: 0.01\n            path (str): Path for the checkpoint to be saved to.\n                            Default: 'checkpoint.pt'\n            trace_func (function): trace print function.\n                            Default: print\n        \"\"\"\n        self.patience = patience\n        self.verbose = verbose\n        self.counter = 0\n        self.best_score = None\n        self.early_stop = False\n        self.val_loss_min = np.inf\n        self.delta = delta\n        self.path = path\n        self.trace_func = trace_func\n\n    def __call__(self, val_loss, model):\n        score = -val_loss\n\n        if self.best_score is None:\n            self.best_score = score\n            self.save_checkpoint(val_loss, model)\n        elif score < self.best_score + self.delta:\n            self.counter += 1\n            self.trace_func(f'EarlyStopping counter: {self.counter} out of {self.patience}')\n            if self.counter >= self.patience:\n                self.early_stop = True\n        else:\n            self.best_score = score\n            self.save_checkpoint(val_loss, model)\n            self.counter = 0\n\n    def save_checkpoint(self, val_loss, model):\n        '''Saves model when validation loss decreases.'''\n        if self.verbose:\n            self.trace_func(f'Validation loss decreased ({self.val_loss_min:.6f} --> {val_loss:.6f}).  Saving model ...')\n        torch.save(model.state_dict(), self.path)\n        self.val_loss_min = val_loss\n\ndef train_model(model, train_loader, val_loader, optimizer, loss_function, device, max_epochs=4, val_interval=1, early_stopping=None):\n    train_losses = []\n    val_losses = [] \n    train_accs = []\n    val_accs = []\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_loss = float('inf')\n    best_epoch = -1\n    print(f'Training using {device}')\n    start_time = time.time()\n\n    for epoch in range(max_epochs):\n        print(f\"\\n---------------- Epoch {epoch + 1}/{max_epochs} ----------------\")\n        model.train()\n        running_loss = 0.0\n        running_acc = 0.0\n        step = 0\n        now = time.time()\n\n        for inputs, labels in train_loader:\n            step += 1\n            inputs = inputs.to(device)\n            labels = labels.to(device)\n\n            optimizer.zero_grad()\n\n            outputs = model(inputs)\n            loss = loss_function(outputs, labels)\n            acc = metric(outputs, labels)\n            loss.backward()\n            optimizer.step()\n\n            running_loss += loss.item()\n            running_acc += acc.item()\n            \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 = running_loss / step\n        epoch_acc = running_acc / step\n        \n        train_losses.append(epoch_loss)\n        train_accs.append(epoch_acc)\n        \n        writer.add_scalar(\"epoch_accuracy\", epoch_acc, epoch)\n        writer.add_scalar(\"epoch_loss\", epoch_loss, epoch)\n        writer.add_scalar(\"running_time\", running_time, epoch)\n        \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            with torch.no_grad():\n                for inputs, labels in val_loader:\n                    inputs = inputs.to(device)\n                    labels = labels.to(device)\n\n                    outputs = model(inputs.float())\n                    loss = loss_function(outputs, labels)\n                    acc = metric(outputs, labels)\n\n                    val_loss += loss.item()\n                    val_acc += acc.item()\n                    step += 1\n                    \n            val_loss /= step\n            val_acc /= step\n\n            val_losses.append(val_loss)\n            val_accs.append(val_acc)\n            writer.add_scalar(\"val_accuracy\", val_acc, epoch)\n            writer.add_scalar(\"val_loss\", val_loss, epoch)\n            print(f'val loss: {val_loss:.4f}, val accuracy: {val_acc:.4f}')\n            \n            if early_stopping:\n                early_stopping(val_loss, model)\n                if early_stopping.early_stop:\n                    print(\"\\nEarly stopping\")\n                    break\n\n            if val_loss < best_loss:\n                best_loss = val_loss\n                best_epoch = epoch+1\n                best_model_wts = copy.deepcopy(model.state_dict())\n\n    end_time = time.time()\n    model.load_state_dict(best_model_wts)\n    print(f\"\\nTraining Completed! Total time taken: {(end_time - start_time)/60:.2f} mins\")\n    print(f\"Best Model Checkpoint Saved!! with loss: {best_loss} at epoch:{best_epoch}\\n\")\n    return model, train_losses, val_losses, train_accs, val_accs","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:25:22.511631Z","iopub.execute_input":"2024-08-15T11:25:22.512410Z","iopub.status.idle":"2024-08-15T11:25:22.537077Z","shell.execute_reply.started":"2024-08-15T11:25:22.512375Z","shell.execute_reply":"2024-08-15T11:25:22.536117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metric = MultilabelAccuracy(num_labels=15).to(device)\nwriter = SummaryWriter()\nmax_epochs = 7\nval_interval = 2\n\n# Initialize the early stopping\nearly_stopping = EarlyStopping(patience=3, verbose=True)\n\n# Train the model with early stopping\ntrained_model, train_losses, val_losses, train_accs, val_accs = train_model(model, train_loader, val_loader, optimizer, loss_function, device, max_epochs, val_interval, early_stopping=early_stopping)","metadata":{"execution":{"iopub.status.busy":"2024-08-15T11:25:22.538361Z","iopub.execute_input":"2024-08-15T11:25:22.538718Z","iopub.status.idle":"2024-08-15T12:46:29.051543Z","shell.execute_reply.started":"2024-08-15T11:25:22.538682Z","shell.execute_reply":"2024-08-15T12:46:29.050456Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax1 = plt.subplots()\n\nepochs = list(range(1, len(train_losses)+1))\nax1.set_xlabel('epoch')\nax1.set_ylabel('Training Loss', color='tab:red')\nax1.plot(epochs, train_losses, color='tab:red', marker='o', label='loss')\nax1.tick_params(axis='y', labelcolor='tab:red')\n\nax2 = ax1.twinx()\nax2.set_ylabel('Training Accuracy', color='tab:blue')\nax2.plot(epochs, train_accs, color='tab:blue', marker='o', label='accuracy')\nax2.tick_params(axis='y', labelcolor='tab:blue')\n\nplt.title('Training Loss and Accuracy over Epochs')\n\nfig.tight_layout() \nax1.legend(loc='upper left')\nax2.legend(loc='upper right')\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-15T12:46:29.053618Z","iopub.execute_input":"2024-08-15T12:46:29.053931Z","iopub.status.idle":"2024-08-15T12:46:30.111894Z","shell.execute_reply.started":"2024-08-15T12:46:29.053902Z","shell.execute_reply":"2024-08-15T12:46:30.110984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, ax1 = plt.subplots()\n\nepochs = list(range(1, len(val_losses)+1))\nax1.set_xlabel('epoch')\nax1.set_ylabel('Validation Loss', color='tab:red')\nax1.plot(epochs, val_losses, color='tab:red', marker='o', label='loss')\nax1.tick_params(axis='y', labelcolor='tab:red')\n\nax2 = ax1.twinx()\nax2.set_ylabel('Validation Accuracy', color='tab:blue')\nax2.plot(epochs, val_accs, color='tab:blue', marker='o', label='accuracy')\nax2.tick_params(axis='y', labelcolor='tab:blue')\n\nplt.title('Validation Loss and Accuracy over Epochs')\n\nfig.tight_layout() \nax1.legend(loc='upper left')\nax2.legend(loc='upper right')\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-15T13:09:13.857070Z","iopub.execute_input":"2024-08-15T13:09:13.857476Z","iopub.status.idle":"2024-08-15T13:09:14.420798Z","shell.execute_reply.started":"2024-08-15T13:09:13.857448Z","shell.execute_reply":"2024-08-15T13:09:14.419970Z"},"trusted":true},"execution_count":null,"outputs":[]}]}