{"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":8776139,"sourceType":"datasetVersion","datasetId":5274724}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 2d Segmentation of Sagittal Lumbar Spine MRI\n\n- Training data Spider dataset (https://doi.org/10.5281/zenodo.10159290)\n- Very simple model using segmentation_models_pytorch\n- Trained model attached\n- Images resized to 256x256","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install segmentation-models-pytorch\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-01T07:54:51.666438Z","iopub.execute_input":"2024-12-01T07:54:51.667026Z","iopub.status.idle":"2024-12-01T07:55:07.364314Z","shell.execute_reply.started":"2024-12-01T07:54:51.666984Z","shell.execute_reply":"2024-12-01T07:55:07.363316Z"}},"outputs":[{"name":"stdout","text":"Collecting segmentation-models-pytorch\n  Downloading segmentation_models_pytorch-0.3.4-py3-none-any.whl.metadata (30 kB)\nCollecting efficientnet-pytorch==0.7.1 (from segmentation-models-pytorch)\n  Downloading efficientnet_pytorch-0.7.1.tar.gz (21 kB)\n  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hCollecting huggingface-hub>=0.24.6 (from segmentation-models-pytorch)\n  Downloading huggingface_hub-0.26.3-py3-none-any.whl.metadata (13 kB)\nRequirement already satisfied: pillow in /opt/conda/lib/python3.10/site-packages (from segmentation-models-pytorch) (9.5.0)\nCollecting pretrainedmodels==0.7.4 (from segmentation-models-pytorch)\n  Downloading pretrainedmodels-0.7.4.tar.gz (58 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m58.8/58.8 kB\u001b[0m \u001b[31m3.5 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25h  Preparing metadata (setup.py) ... \u001b[?25ldone\n\u001b[?25hRequirement already satisfied: six in /opt/conda/lib/python3.10/site-packages (from segmentation-models-pytorch) (1.16.0)\nCollecting timm==0.9.7 (from segmentation-models-pytorch)\n  Downloading timm-0.9.7-py3-none-any.whl.metadata (58 kB)\n\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m58.8/58.8 kB\u001b[0m \u001b[31m5.1 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hRequirement already satisfied: torchvision>=0.5.0 in /opt/conda/lib/python3.10/site-packages (from segmentation-models-pytorch) (0.16.2)\nRequirement already satisfied: tqdm in /opt/conda/lib/python3.10/site-packages (from segmentation-models-pytorch) (4.66.4)\nRequirement already satisfied: torch in /opt/conda/lib/python3.10/site-packages (from efficientnet-pytorch==0.7.1->segmentation-models-pytorch) (2.1.2)\nCollecting munch (from pretrainedmodels==0.7.4->segmentation-models-pytorch)\n  Downloading munch-4.0.0-py2.py3-none-any.whl.metadata (5.9 kB)\nRequirement already satisfied: pyyaml in /opt/conda/lib/python3.10/site-packages (from timm==0.9.7->segmentation-models-pytorch) (6.0.1)\nRequirement already satisfied: safetensors in /opt/conda/lib/python3.10/site-packages (from timm==0.9.7->segmentation-models-pytorch) (0.4.3)\nRequirement already satisfied: filelock in /opt/conda/lib/python3.10/site-packages (from huggingface-hub>=0.24.6->segmentation-models-pytorch) (3.13.1)\nRequirement already satisfied: fsspec>=2023.5.0 in /opt/conda/lib/python3.10/site-packages (from huggingface-hub>=0.24.6->segmentation-models-pytorch) (2024.3.1)\nRequirement already satisfied: packaging>=20.9 in /opt/conda/lib/python3.10/site-packages (from huggingface-hub>=0.24.6->segmentation-models-pytorch) (21.3)\nRequirement already satisfied: requests in /opt/conda/lib/python3.10/site-packages (from huggingface-hub>=0.24.6->segmentation-models-pytorch) (2.32.3)\nRequirement already satisfied: typing-extensions>=3.7.4.3 in /opt/conda/lib/python3.10/site-packages (from huggingface-hub>=0.24.6->segmentation-models-pytorch) (4.9.0)\nRequirement already satisfied: numpy in /opt/conda/lib/python3.10/site-packages (from torchvision>=0.5.0->segmentation-models-pytorch) (1.26.4)\nRequirement already satisfied: pyparsing!=3.0.5,>=2.0.2 in /opt/conda/lib/python3.10/site-packages (from packaging>=20.9->huggingface-hub>=0.24.6->segmentation-models-pytorch) (3.1.1)\nRequirement already satisfied: sympy in /opt/conda/lib/python3.10/site-packages (from torch->efficientnet-pytorch==0.7.1->segmentation-models-pytorch) (1.12.1)\nRequirement already satisfied: networkx in /opt/conda/lib/python3.10/site-packages (from torch->efficientnet-pytorch==0.7.1->segmentation-models-pytorch) (3.2.1)\nRequirement already satisfied: jinja2 in /opt/conda/lib/python3.10/site-packages (from torch->efficientnet-pytorch==0.7.1->segmentation-models-pytorch) (3.1.2)\nRequirement already satisfied: charset-normalizer<4,>=2 in /opt/conda/lib/python3.10/site-packages (from requests->huggingface-hub>=0.24.6->segmentation-models-pytorch) (3.3.2)\nRequirement already satisfied: idna<4,>=2.5 in /opt/conda/lib/python3.10/site-packages (from requests->huggingface-hub>=0.24.6->segmentation-models-pytorch) (3.6)\nRequirement already satisfied: urllib3<3,>=1.21.1 in /opt/conda/lib/python3.10/site-packages (from requests->huggingface-hub>=0.24.6->segmentation-models-pytorch) (1.26.18)\nRequirement already satisfied: certifi>=2017.4.17 in /opt/conda/lib/python3.10/site-packages (from requests->huggingface-hub>=0.24.6->segmentation-models-pytorch) (2024.2.2)\nRequirement already satisfied: MarkupSafe>=2.0 in /opt/conda/lib/python3.10/site-packages (from jinja2->torch->efficientnet-pytorch==0.7.1->segmentation-models-pytorch) (2.1.3)\nRequirement already satisfied: mpmath<1.4.0,>=1.1.0 in /opt/conda/lib/python3.10/site-packages (from sympy->torch->efficientnet-pytorch==0.7.1->segmentation-models-pytorch) (1.3.0)\nDownloading segmentation_models_pytorch-0.3.4-py3-none-any.whl (109 kB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m109.5/109.5 kB\u001b[0m \u001b[31m7.0 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hDownloading timm-0.9.7-py3-none-any.whl (2.2 MB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m2.2/2.2 MB\u001b[0m \u001b[31m39.7 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0ma \u001b[36m0:00:01\u001b[0m\n\u001b[?25hDownloading huggingface_hub-0.26.3-py3-none-any.whl (447 kB)\n\u001b[2K   \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m447.6/447.6 kB\u001b[0m \u001b[31m35.0 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m\n\u001b[?25hDownloading munch-4.0.0-py2.py3-none-any.whl (9.9 kB)\nBuilding wheels for collected packages: efficientnet-pytorch, pretrainedmodels\n  Building wheel for efficientnet-pytorch (setup.py) ... \u001b[?25ldone\n\u001b[?25h  Created wheel for efficientnet-pytorch: filename=efficientnet_pytorch-0.7.1-py3-none-any.whl size=16428 sha256=0164d3d76b6d2cca281fedf6b387835172a8d2859a39f56e9f146450c77fc4cc\n  Stored in directory: /root/.cache/pip/wheels/03/3f/e9/911b1bc46869644912bda90a56bcf7b960f20b5187feea3baf\n  Building wheel for pretrainedmodels (setup.py) ... \u001b[?25ldone\n\u001b[?25h  Created wheel for pretrainedmodels: filename=pretrainedmodels-0.7.4-py3-none-any.whl size=60945 sha256=91a83ebaccdb8d2a61f503add0083725c70c4377dab5e3bbf28b3a9baa8afaf3\n  Stored in directory: /root/.cache/pip/wheels/35/cb/a5/8f534c60142835bfc889f9a482e4a67e0b817032d9c6883b64\nSuccessfully built efficientnet-pytorch pretrainedmodels\nInstalling collected packages: munch, huggingface-hub, efficientnet-pytorch, timm, pretrainedmodels, segmentation-models-pytorch\n  Attempting uninstall: huggingface-hub\n    Found existing installation: huggingface-hub 0.23.2\n    Uninstalling huggingface-hub-0.23.2:\n      Successfully uninstalled huggingface-hub-0.23.2\n  Attempting uninstall: timm\n    Found existing installation: timm 1.0.3\n    Uninstalling timm-1.0.3:\n      Successfully uninstalled timm-1.0.3\nSuccessfully installed efficientnet-pytorch-0.7.1 huggingface-hub-0.26.3 munch-4.0.0 pretrainedmodels-0.7.4 segmentation-models-pytorch-0.3.4 timm-0.9.7\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd \nimport os\nfrom pathlib import Path\nfrom PIL import Image\nfrom sklearn.model_selection import KFold\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset\nimport segmentation_models_pytorch as sm\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-12-01T07:55:07.366209Z","iopub.execute_input":"2024-12-01T07:55:07.366588Z","iopub.status.idle":"2024-12-01T07:55:14.485817Z","shell.execute_reply.started":"2024-12-01T07:55:07.366554Z","shell.execute_reply":"2024-12-01T07:55:14.484846Z"},"trusted":true},"outputs":[],"execution_count":2},{"cell_type":"code","source":"#transforms\nnewsize = (256, 256)\n#dataset\nfold = 1\n#dataloader\nbatch_size = 64\nnum_workers = 4\n#model\nnum_classes = 20\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")#run\nepochs = 100\nlearning_rate = 1e-3\n\nTRAIN = True #or False for inference only","metadata":{"execution":{"iopub.status.busy":"2024-11-30T15:15:41.047866Z","iopub.execute_input":"2024-11-30T15:15:41.048571Z","iopub.status.idle":"2024-11-30T15:15:41.111519Z","shell.execute_reply.started":"2024-11-30T15:15:41.04853Z","shell.execute_reply":"2024-11-30T15:15:41.110646Z"},"trusted":true},"outputs":[],"execution_count":6},{"cell_type":"markdown","source":"### Model","metadata":{}},{"cell_type":"code","source":"model1 = sm.Unet('resnet34', encoder_weights='imagenet', classes=num_classes, activation='softmax')\nmodel2 = sm.Unet('efficientnet-b6', encoder_weights='imagenet', classes=num_classes, activation='softmax')\nmodel3 = sm.Unet('vgg16', encoder_weights='imagenet', classes=num_classes, activation='softmax')\n","metadata":{"execution":{"iopub.status.busy":"2024-11-30T15:15:44.057273Z","iopub.execute_input":"2024-11-30T15:15:44.057613Z","iopub.status.idle":"2024-11-30T15:15:52.152694Z","shell.execute_reply.started":"2024-11-30T15:15:44.057587Z","shell.execute_reply":"2024-11-30T15:15:52.151806Z"},"trusted":true},"outputs":[{"name":"stderr","text":"Downloading: \"https://download.pytorch.org/models/resnet34-333f7ec4.pth\" to /root/.cache/torch/hub/checkpoints/resnet34-333f7ec4.pth\n100%|██████████| 83.3M/83.3M [00:00<00:00, 226MB/s]\nDownloading: \"https://github.com/lukemelas/EfficientNet-PyTorch/releases/download/1.0/efficientnet-b6-c76e70fd.pth\" to /root/.cache/torch/hub/checkpoints/efficientnet-b6-c76e70fd.pth\n100%|██████████| 165M/165M [00:00<00:00, 206MB/s]  \nDownloading: \"https://download.pytorch.org/models/vgg16-397923af.pth\" to /root/.cache/torch/hub/checkpoints/vgg16-397923af.pth\n100%|██████████| 528M/528M [00:02<00:00, 239MB/s]  \n","output_type":"stream"}],"execution_count":7},{"cell_type":"markdown","source":"### Create folds","metadata":{}},{"cell_type":"code","source":"output_dir = \"/kaggle/input/spider-mri-spine-t2-png/data\"\nim_dir = os.path.join(output_dir, \"images\")\nmask_dir = os.path.join(output_dir, \"masks\")\n\n# get list of data\nitems = list(Path(im_dir).glob(\"*.png\"))\nimage_names = [o.name for o in items]\nimages = list(set([o.split('_')[0] for o in image_names]))\n\nfold_df = pd.DataFrame({\"image_name\": images})\n# Seed for reproducibility\nnp.random.seed(42)\n\n# Split the DataFrame into 5 folds\nkf = KFold(n_splits=5, shuffle=True, random_state=42)\nfor i, (_, v_ind) in enumerate(kf.split(fold_df)):\n    fold_df.loc[v_ind, 'fold'] = i+1\n\n# Create df with image_names and their respective folds\ndef get_fold(fn, df):\n    image_name = fn.name.split(\"_\")[0] \n    return df.loc[df.image_name==image_name, 'fold'].values[0]\n\nfolds = [get_fold(o, fold_df) for o in items]\ndf = pd.DataFrame({\"image\": image_names, \"fold\": folds})\n\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2024-11-30T15:16:10.523709Z","iopub.execute_input":"2024-11-30T15:16:10.52449Z","iopub.status.idle":"2024-11-30T15:16:11.697555Z","shell.execute_reply.started":"2024-11-30T15:16:10.524462Z","shell.execute_reply":"2024-11-30T15:16:11.696705Z"},"trusted":true},"outputs":[{"execution_count":9,"output_type":"execute_result","data":{"text/plain":"        image  fold\n0   10_11.png   1.0\n1  191_01.png   4.0\n2  210_06.png   2.0\n3  239_03.png   5.0\n4   61_01.png   1.0","text/html":"<div>\n<style scoped>\n    .dataframe tbody tr th:only-of-type {\n        vertical-align: middle;\n    }\n\n    .dataframe tbody tr th {\n        vertical-align: top;\n    }\n\n    .dataframe thead th {\n        text-align: right;\n    }\n</style>\n<table border=\"1\" class=\"dataframe\">\n  <thead>\n    <tr style=\"text-align: right;\">\n      <th></th>\n      <th>image</th>\n      <th>fold</th>\n    </tr>\n  </thead>\n  <tbody>\n    <tr>\n      <th>0</th>\n      <td>10_11.png</td>\n      <td>1.0</td>\n    </tr>\n    <tr>\n      <th>1</th>\n      <td>191_01.png</td>\n      <td>4.0</td>\n    </tr>\n    <tr>\n      <th>2</th>\n      <td>210_06.png</td>\n      <td>2.0</td>\n    </tr>\n    <tr>\n      <th>3</th>\n      <td>239_03.png</td>\n      <td>5.0</td>\n    </tr>\n    <tr>\n      <th>4</th>\n      <td>61_01.png</td>\n      <td>1.0</td>\n    </tr>\n  </tbody>\n</table>\n</div>"},"metadata":{}}],"execution_count":9},{"cell_type":"markdown","source":"### Dataset class","metadata":{}},{"cell_type":"code","source":"class SEGDataset(Dataset):\n    def __init__(self, df, mode, transforms=None):\n        self.df = df.reset_index()\n        self.mode = mode\n        self.transforms = transforms\n\n    def __len__(self):\n        return self.df.shape[0]\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n\n        image_path = os.path.join(im_dir, row.image)\n        mask_path = os.path.join(mask_dir, row.image)\n\n        # Open image\n        image = Image.open(image_path)\n        if image.mode != 'RGB':  # Ensure image is RGB\n            image = image.convert('RGB')\n        image = np.asarray(image)\n        if (image > 1).any():  # Normalize if pixel values are between 0-255\n            image = image / 255.0\n\n        # Open mask\n        mask = Image.open(mask_path)\n        mask = np.asarray(mask)\n        assert mask.max() < num_classes, f\"Mask value {mask.max()} exceeds number of classes {num_classes}\"\n\n        # Apply transformations\n        if self.transforms is not None:\n            transformed = self.transforms(image=image, mask=mask)\n            image = transformed[\"image\"]\n            mask = transformed[\"mask\"]\n        \n        # Create one layer for each label\n        mask = torch.as_tensor(mask).long()\n        mask = torch.nn.functional.one_hot(mask, num_classes=num_classes).permute(2,0,1).float()\n        #mask = torch.nn.functional.one_hot(mask, num_classes=num_classes).permute(0,3,1,2).squeeze(0).float()\n\n        # Convert image to tensor\n        image = torch.as_tensor(image).float()\n\n        return image, mask          ","metadata":{"execution":{"iopub.status.busy":"2024-11-30T15:16:14.074889Z","iopub.execute_input":"2024-11-30T15:16:14.075636Z","iopub.status.idle":"2024-11-30T15:16:14.083004Z","shell.execute_reply.started":"2024-11-30T15:16:14.075611Z","shell.execute_reply":"2024-11-30T15:16:14.082005Z"},"trusted":true},"outputs":[],"execution_count":10},{"cell_type":"markdown","source":"### Transforms","metadata":{}},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\ntransforms_train = A.Compose([\n    A.Resize(newsize[0], newsize[1]),\n    A.HorizontalFlip(),\n    A.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n    ),\n    ToTensorV2()\n])\n\ntransforms_valid = A.Compose([\n    A.Resize(newsize[0], newsize[1]),\n    A.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225],\n    ),\n    ToTensorV2()\n])","metadata":{"execution":{"iopub.status.busy":"2024-11-30T15:16:16.574896Z","iopub.execute_input":"2024-11-30T15:16:16.575232Z","iopub.status.idle":"2024-11-30T15:16:16.901652Z","shell.execute_reply.started":"2024-11-30T15:16:16.575207Z","shell.execute_reply":"2024-11-30T15:16:16.900988Z"},"trusted":true},"outputs":[],"execution_count":11},{"cell_type":"markdown","source":"### Loss","metadata":{}},{"cell_type":"code","source":"class CombinedLoss(nn.Module):\n    def __init__(self, weight_ce=1.0, weight_iou=1.0):\n        super(CombinedLoss, self).__init__()\n        self.weight_ce = weight_ce\n        self.weight_iou = weight_iou\n        self.cross_entropy_loss = nn.CrossEntropyLoss()\n\n    def forward(self, inputs, targets):\n        # Cross-Entropy Loss\n        ce_loss = self.cross_entropy_loss(inputs, targets)\n\n        # IoU Loss\n        # Apply softmax to the inputs to get probabilities\n        probs = F.softmax(inputs, dim=1)\n\n        intersection = torch.sum(probs * targets, dim=(2, 3))\n        union = torch.sum(probs + targets, dim=(2, 3)) - intersection\n        iou = (intersection + 1e-6) / (union + 1e-6)\n        iou_loss = 1 - iou.mean()\n\n        # Combine losses\n        loss = self.weight_ce * ce_loss + self.weight_iou * iou_loss\n        return loss","metadata":{"execution":{"iopub.status.busy":"2024-11-30T13:48:16.30904Z","iopub.execute_input":"2024-11-30T13:48:16.309779Z","iopub.status.idle":"2024-11-30T13:48:16.316115Z","shell.execute_reply.started":"2024-11-30T13:48:16.309745Z","shell.execute_reply":"2024-11-30T13:48:16.315337Z"},"trusted":true},"outputs":[],"execution_count":8},{"cell_type":"markdown","source":"### Create datasets and dataloaders","metadata":{}},{"cell_type":"code","source":"train_ = df[df['fold'] != fold].reset_index(drop=True)\nvalid_ = df[df['fold'] == fold].reset_index(drop=True)\n\ndataset_train = SEGDataset(train_, 'train',  transforms_train)\ndataset_valid = SEGDataset(valid_, 'valid',  transforms_valid)\n\ntrain_loader = torch.utils.data.DataLoader(dataset_train, batch_size=batch_size, shuffle=True, num_workers=num_workers)\nval_loader = torch.utils.data.DataLoader(dataset_valid, batch_size=batch_size, shuffle=False, num_workers=num_workers)","metadata":{"execution":{"iopub.status.busy":"2024-11-30T15:16:19.057312Z","iopub.execute_input":"2024-11-30T15:16:19.057855Z","iopub.status.idle":"2024-11-30T15:16:19.069339Z","shell.execute_reply.started":"2024-11-30T15:16:19.057826Z","shell.execute_reply":"2024-11-30T15:16:19.068567Z"},"trusted":true},"outputs":[],"execution_count":12},{"cell_type":"markdown","source":"### Run function","metadata":{}},{"cell_type":"code","source":"from torch import optim\nfrom torch.nn import BCEWithLogitsLoss\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\n\ndef run(train_loader, val_loader, model, learning_rate, criterion, epochs, device):\n  \"\"\"\n  Trains a U-net model for multi-label segmentation.\n\n  Args:\n      train_loader: DataLoader for training data.\n      val_loader: DataLoader for validation data.\n      model: U-net model instance.\n      learning_rate: Learning rate for optimizer.\n      epochs: Number of epochs to train.\n      device: Device to use for training (CPU or GPU).\n  \"\"\"\n  # Define loss function and optimizer\n  optimizer = optim.Adam(model.parameters(), lr=learning_rate)\n\n  # Define a learning rate scheduler\n  scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=5, verbose=True)\n\n\n  # Training loop\n  for epoch in range(epochs):\n    model.train()\n    train_loss = 0.0\n    for images, masks in train_loader:\n      images, masks = images.to(device), masks.to(device)\n\n      # Forward pass and calculate loss\n      outputs = model(images)\n      loss = criterion(outputs, masks)\n\n      # Backward pass and update weights\n      optimizer.zero_grad()\n      loss.backward()\n      optimizer.step()\n\n      train_loss += loss.item()\n\n    train_loss /= len(train_loader)\n\n    # Validation step (optional)\n    model.eval()\n    with torch.no_grad():\n      val_loss = 0.0\n      for images, masks in val_loader:\n        images, masks = images.to(device), masks.to(device)\n        outputs = model(images)\n        val_loss += criterion(outputs, masks).item()\n\n    val_loss /= len(val_loader)\n    \n    # Step the scheduler\n    scheduler.step(val_loss)\n\n    # Print training and validation loss\n    print(f\"Epoch: {epoch+1}/{epochs} | Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")","metadata":{"execution":{"iopub.status.busy":"2024-11-30T15:16:21.194983Z","iopub.execute_input":"2024-11-30T15:16:21.19538Z","iopub.status.idle":"2024-11-30T15:16:21.20309Z","shell.execute_reply.started":"2024-11-30T15:16:21.195355Z","shell.execute_reply":"2024-11-30T15:16:21.202252Z"},"trusted":true},"outputs":[],"execution_count":13},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef inference(model, dataloader, device, num_samples=16):\n    model.eval()\n    images_batch = []\n    preds_batch = []\n    \n    with torch.no_grad():\n        for images, _ in dataloader:\n            images = images.to(device)\n            outputs = model(images)\n            preds = torch.argmax(outputs, dim=1)\n            \n            images_batch.append(images.cpu())\n            preds_batch.append(preds.cpu())\n            \n            if len(images_batch) * images.size(0) >= num_samples:\n                break\n\n    images_batch = torch.cat(images_batch)[:num_samples]\n    preds_batch = torch.cat(preds_batch)[:num_samples]\n    \n    return images_batch, preds_batch\n\n\n# Define a color map with fixed colors for each label\ndef get_label_colors(num_classes):\n    colors = plt.cm.tab20(np.linspace(0, 1, num_classes))\n    return colors\n\ndef visualize_predictions(images, masks, num_classes=20, num_samples=16):\n    num_samples = min(num_samples, len(images))\n    plt.figure(figsize=(20, 20))\n    \n    label_colors = get_label_colors(num_classes)\n    \n    for i in range(num_samples):\n        plt.subplot(4, 8, i * 2 + 1)\n        im = images[i].numpy()\n        im = np.transpose(im, (1, 2, 0))\n        #denormalize\n        im = ((im * [0.229, 0.224, 0.225]) + [0.485, 0.456, 0.406]) * 255\n        plt.imshow(im)\n        plt.title(\"Input Image\")\n        plt.axis('off')\n        \n        plt.subplot(4, 8, i * 2 + 2)\n        mask = masks[i].numpy()\n\n        color_mask = np.zeros((mask.shape[0], mask.shape[1], 3))\n        for label in range(num_classes):\n            color_mask[mask == label] = label_colors[label][:3] * 255\n        \n        plt.imshow(color_mask.astype(np.uint8))\n        plt.title(\"Predicted Mask\")\n        plt.axis('off')\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2024-11-30T15:16:24.225279Z","iopub.execute_input":"2024-11-30T15:16:24.2256Z","iopub.status.idle":"2024-11-30T15:16:24.234932Z","shell.execute_reply.started":"2024-11-30T15:16:24.225574Z","shell.execute_reply":"2024-11-30T15:16:24.234212Z"},"trusted":true},"outputs":[],"execution_count":14},{"cell_type":"markdown","source":"### Train","metadata":{}},{"cell_type":"code","source":"criterion = CombinedLoss()\nmodels = [model2, model3]  # List of models\nmodel_paths = ['./model2.pth','./model3.pth']  # Paths to save models\n\n# Send models to the device\nfor model in models:\n    model.to(device)\n\nif TRAIN:\n    # Train each model individually\n    for i, model in enumerate(models):\n        print(f\"Training model {i + 1}...\")\n        run(train_loader, val_loader, model, learning_rate, criterion, epochs, device)\n        # Save the trained model\n        torch.save(model.state_dict(), model_paths[i])\nelse:\n    # Load pre-trained weights for each model\n    for i, model in enumerate(models):\n        model.load_state_dict(torch.load(model_paths[i]))\n        model.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-30T15:17:11.485352Z","iopub.execute_input":"2024-11-30T15:17:11.48572Z","iopub.status.idle":"2024-11-30T15:17:11.515639Z","shell.execute_reply.started":"2024-11-30T15:17:11.485692Z","shell.execute_reply":"2024-11-30T15:17:11.514419Z"}},"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mNameError\u001b[0m                                 Traceback (most recent call last)","Cell \u001b[0;32mIn[18], line 1\u001b[0m\n\u001b[0;32m----> 1\u001b[0m criterion \u001b[38;5;241m=\u001b[39m \u001b[43mCombinedLoss\u001b[49m()\n\u001b[1;32m      2\u001b[0m models \u001b[38;5;241m=\u001b[39m [model2, model3]  \u001b[38;5;66;03m# List of models\u001b[39;00m\n\u001b[1;32m      3\u001b[0m model_paths \u001b[38;5;241m=\u001b[39m [\u001b[38;5;124m'\u001b[39m\u001b[38;5;124m./model2.pth\u001b[39m\u001b[38;5;124m'\u001b[39m,\u001b[38;5;124m'\u001b[39m\u001b[38;5;124m./model3.pth\u001b[39m\u001b[38;5;124m'\u001b[39m]  \u001b[38;5;66;03m# Paths to save models\u001b[39;00m\n","\u001b[0;31mNameError\u001b[0m: name 'CombinedLoss' is not defined"],"ename":"NameError","evalue":"name 'CombinedLoss' is not defined","output_type":"error"}],"execution_count":18},{"cell_type":"code","source":"from keras.models import load_model\n\n#Set compile=False as we are not loading it for training, only for prediction.\nmodel1 = load_model('saved_models/res34_backbone_50epochs.hdf5', compile=False)\nmodel2 = load_model('saved_models/inceptionv3_backbone_50epochs.hdf5', compile=False)\nmodel3 = load_model('saved_models/vgg19_backbone_50epochs.hdf5', compile=False)\n\n#Weighted average ensemble\nmodels = [model1, model2, model3]\n#preds = [model.predict(X_test) for model in models]\n\npred1 = model1.predict(X_test1)\npred2 = model2.predict(X_test2)\npred3 = model3.predict(X_test3)\n\npreds=np.array([pred1, pred2, pred3])\n\n#preds=np.array(preds)\nweights = [0.3, 0.5, 0.2]\n\n#Use tensordot to sum the products of all elements over specified axes.\nweighted_preds = np.tensordot(preds, weights, axes=((0),(0)))\nweighted_ensemble_prediction = np.argmax(weighted_preds, axis=3)\n\ny_pred1_argmax=np.argmax(pred1, axis=3)\ny_pred2_argmax=np.argmax(pred2, axis=3)\ny_pred3_argmax=np.argmax(pred3, axis=3)\n\n\n#Using built in keras function\nn_classes = 4\nIOU1 = MeanIoU(num_classes=n_classes)  \nIOU2 = MeanIoU(num_classes=n_classes)  \nIOU3 = MeanIoU(num_classes=n_classes)  \nIOU_weighted = MeanIoU(num_classes=n_classes)  \n\nIOU1.update_state(y_test[:,:,:,0], y_pred1_argmax)\nIOU2.update_state(y_test[:,:,:,0], y_pred2_argmax)\nIOU3.update_state(y_test[:,:,:,0], y_pred3_argmax)\nIOU_weighted.update_state(y_test[:,:,:,0], weighted_ensemble_prediction)\n\n\nprint('IOU Score for model1 = ', IOU1.result().numpy())\nprint('IOU Score for model2 = ', IOU2.result().numpy())\nprint('IOU Score for model3 = ', IOU3.result().numpy())\nprint('IOU Score for weighted average ensemble = ', IOU_weighted.result().numpy())\n###########################################\n#Grid search for the best combination of w1, w2, w3 that gives maximum acuracy\n\nimport pandas as pd\ndf = pd.DataFrame([])\n\nfor w1 in range(0, 4):\n    for w2 in range(0,4):\n        for w3 in range(0,4):\n            wts = [w1/10.,w2/10.,w3/10.]\n            \n            IOU_wted = MeanIoU(num_classes=n_classes) \n            wted_preds = np.tensordot(preds, wts, axes=((0),(0)))\n            wted_ensemble_pred = np.argmax(wted_preds, axis=3)\n            IOU_wted.update_state(y_test[:,:,:,0], wted_ensemble_pred)\n            print(\"Now predciting for weights :\", w1/10., w2/10., w3/10., \" : IOU = \", IOU_wted.result().numpy())\n            df = df.append(pd.DataFrame({'wt1':wts[0],'wt2':wts[1], \n                                         'wt3':wts[2], 'IOU': IOU_wted.result().numpy()}, index=[0]), ignore_index=True)\n            \nmax_iou_row = df.iloc[df['IOU'].idxmax()]\nprint(\"Max IOU of \", max_iou_row[3], \" obained with w1=\", max_iou_row[0],\n      \" w2=\", max_iou_row[1], \" and w3=\", max_iou_row[2])         \n\n\n#############################################################\nopt_weights = [max_iou_row[0], max_iou_row[1], max_iou_row[2]]\n\n#Use tensordot to sum the products of all elements over specified axes.\nopt_weighted_preds = np.tensordot(preds, opt_weights, axes=((0),(0)))\nopt_weighted_ensemble_prediction = np.argmax(opt_weighted_preds, axis=3)\n#######################################################\n#Predict on a few images\n\nimport random\ntest_img_number = random.randint(0, len(X_test))\ntest_img = X_test[test_img_number]\nground_truth=y_test[test_img_number]\ntest_img_norm=test_img[:,:,:]\ntest_img_input=np.expand_dims(test_img_norm, 0)\n\n#Weighted average ensemble\nmodels = [model1, model2, model3]\n\ntest_img_input1 = preprocess_input1(test_img_input)\ntest_img_input2 = preprocess_input2(test_img_input)\ntest_img_input3 = preprocess_input3(test_img_input)\n\ntest_pred1 = model1.predict(test_img_input1)\ntest_pred2 = model2.predict(test_img_input2)\ntest_pred3 = model3.predict(test_img_input3)\n\ntest_preds=np.array([test_pred1, test_pred2, test_pred3])\n\n#Use tensordot to sum the products of all elements over specified axes.\nweighted_test_preds = np.tensordot(test_preds, opt_weights, axes=((0),(0)))\nweighted_ensemble_test_prediction = np.argmax(weighted_test_preds, axis=3)[0,:,:]\n\n\nplt.figure(figsize=(12, 8))\nplt.subplot(231)\nplt.title('Testing Image')\nplt.imshow(test_img[:,:,0], cmap='gray')\nplt.subplot(232)\nplt.title('Testing Label')\nplt.imshow(ground_truth[:,:,0], cmap='jet')\nplt.subplot(233)\nplt.title('Prediction on test image')\nplt.imshow(weighted_ensemble_test_prediction, cmap='jet')\nplt.show()\n\n#####################################################################\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}