{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-07T08:27:53.158291Z","iopub.execute_input":"2023-09-07T08:27:53.158719Z","iopub.status.idle":"2023-09-07T08:27:53.501818Z","shell.execute_reply.started":"2023-09-07T08:27:53.158678Z","shell.execute_reply":"2023-09-07T08:27:53.500845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import zipfile\n\npath_to_zip_file = \"/kaggle/input/carvana-image-masking-challenge/train.zip\"\ndirectory_to_extract_to = \"/kaggle/working/\"\nwith zipfile.ZipFile(path_to_zip_file, 'r') as zip_ref:\n    zip_ref.extractall(directory_to_extract_to)\n    \npath_to_zip_file = \"/kaggle/input/carvana-image-masking-challenge/train_masks.zip\"\ndirectory_to_extract_to = \"/kaggle/working/\"\nwith zipfile.ZipFile(path_to_zip_file, 'r') as zip_ref:\n    zip_ref.extractall(directory_to_extract_to)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:27:53.503840Z","iopub.execute_input":"2023-09-07T08:27:53.504334Z","iopub.status.idle":"2023-09-07T08:28:03.671174Z","shell.execute_reply.started":"2023-09-07T08:27:53.504297Z","shell.execute_reply":"2023-09-07T08:28:03.670069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:03.672726Z","iopub.execute_input":"2023-09-07T08:28:03.673059Z","iopub.status.idle":"2023-09-07T08:28:07.018015Z","shell.execute_reply.started":"2023-09-07T08:28:03.673026Z","shell.execute_reply":"2023-09-07T08:28:07.017062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\" Parts of the U-Net model \"\"\"\n\nclass DoubleConv(nn.Module):\n    \"\"\"(convolution => [BN] => ReLU) * 2\"\"\"\n\n    def __init__(self, in_channels, out_channels, mid_channels=None):\n        super().__init__()\n        if not mid_channels:\n            mid_channels = out_channels\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(mid_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)\n\n\nclass Down(nn.Module):\n    \"\"\"Downscaling with maxpool then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool2d(2),\n            DoubleConv(in_channels, out_channels)\n        )\n\n    def forward(self, x):\n        return self.maxpool_conv(x)\n\n\nclass Up(nn.Module):\n    \"\"\"Upscaling then double conv\"\"\"\n\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super().__init__()\n\n        # if bilinear, use the normal convolutions to reduce the number of channels\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n            self.conv = DoubleConv(in_channels, out_channels, in_channels // 2)\n        else:\n            self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)\n            self.conv = DoubleConv(in_channels, out_channels)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        # input is CHW\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n\n        x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,\n                        diffY // 2, diffY - diffY // 2])\n        # if you have padding issues, see\n        # https://github.com/HaiyongJiang/U-Net-Pytorch-Unstructured-Buggy/commit/0e854509c2cea854e247a9c615f175f76fbb2e3a\n        # https://github.com/xiaopeng-liao/Pytorch-UNet/commit/8ebac70e633bac59fc22bb5195e513d5832fb3bd\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\n\nclass OutConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(OutConv, self).__init__()\n        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=1)\n\n    def forward(self, x):\n        return self.conv(x)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:07.020733Z","iopub.execute_input":"2023-09-07T08:28:07.021223Z","iopub.status.idle":"2023-09-07T08:28:07.038972Z","shell.execute_reply.started":"2023-09-07T08:28:07.021187Z","shell.execute_reply":"2023-09-07T08:28:07.038040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\" Full assembly of the parts to form the complete network \"\"\"\n\nclass UNet(nn.Module):\n    def __init__(self, n_channels, n_classes, bilinear=False):\n        super(UNet, self).__init__()\n        self.n_channels = n_channels\n        self.n_classes = n_classes\n        self.bilinear = bilinear\n\n        self.inc = (DoubleConv(n_channels, 64))\n        self.down1 = (Down(64, 128))\n        self.down2 = (Down(128, 256))\n        self.down3 = (Down(256, 512))\n        factor = 2 if bilinear else 1\n        self.down4 = (Down(512, 1024 // factor))\n        self.up1 = (Up(1024, 512 // factor, bilinear))\n        self.up2 = (Up(512, 256 // factor, bilinear))\n        self.up3 = (Up(256, 128 // factor, bilinear))\n        self.up4 = (Up(128, 64, bilinear))\n        self.outc = (OutConv(64, n_classes))\n\n    def forward(self, x):\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        x = self.up1(x5, x4)\n        x = self.up2(x, x3)\n        x = self.up3(x, x2)\n        x = self.up4(x, x1)\n        logits = F.softmax(self.outc(x))\n        return logits\n\n    def use_checkpointing(self):\n        self.inc = torch.utils.checkpoint(self.inc)\n        self.down1 = torch.utils.checkpoint(self.down1)\n        self.down2 = torch.utils.checkpoint(self.down2)\n        self.down3 = torch.utils.checkpoint(self.down3)\n        self.down4 = torch.utils.checkpoint(self.down4)\n        self.up1 = torch.utils.checkpoint(self.up1)\n        self.up2 = torch.utils.checkpoint(self.up2)\n        self.up3 = torch.utils.checkpoint(self.up3)\n        self.up4 = torch.utils.checkpoint(self.up4)\n        self.outc = torch.utils.checkpoint(self.outc)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:07.040365Z","iopub.execute_input":"2023-09-07T08:28:07.040690Z","iopub.status.idle":"2023-09-07T08:28:07.055819Z","shell.execute_reply.started":"2023-09-07T08:28:07.040658Z","shell.execute_reply":"2023-09-07T08:28:07.054840Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset\nfrom torchvision import transforms\nfrom PIL import Image","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:07.057227Z","iopub.execute_input":"2023-09-07T08:28:07.057739Z","iopub.status.idle":"2023-09-07T08:28:07.366343Z","shell.execute_reply.started":"2023-09-07T08:28:07.057708Z","shell.execute_reply":"2023-09-07T08:28:07.365371Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, image_paths, mask_paths, transform=None):\n        self.image_paths = image_paths\n        self.mask_paths = mask_paths\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        image = Image.open('/kaggle/working/train/'+self.image_paths[idx])\n        mask = Image.open('/kaggle/working/train_masks/'+self.mask_paths[idx])\n\n        if self.transform:\n            transformed = self.transform(image = np.array(image),\n                                         mask = np.array(mask))\n            image = transformed['image']\n            mask = transformed['mask']\n\n        return image, mask","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:07.367623Z","iopub.execute_input":"2023-09-07T08:28:07.368250Z","iopub.status.idle":"2023-09-07T08:28:07.376840Z","shell.execute_reply.started":"2023-09-07T08:28:07.368217Z","shell.execute_reply":"2023-09-07T08:28:07.375792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\nimage_path = '/kaggle/working/train/'\nmask_path = '/kaggle/working/train_masks/'\n\nimage_list = []\nmask_list = []\n\nfor (directory, subdirectories, files) in os.walk(image_path):\n    image_list.extend(files)\n\nfor file_name in image_list:\n    image_name = file_name.split(\".\")[0]\n    mask_name = [image_name + \"_mask.gif\"]\n    mask_list.extend(mask_name)\n","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:07.378414Z","iopub.execute_input":"2023-09-07T08:28:07.378828Z","iopub.status.idle":"2023-09-07T08:28:07.404895Z","shell.execute_reply.started":"2023-09-07T08:28:07.378797Z","shell.execute_reply":"2023-09-07T08:28:07.403895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Show images data\n# print(image_list)\n# print(mask_list)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:07.406228Z","iopub.execute_input":"2023-09-07T08:28:07.406538Z","iopub.status.idle":"2023-09-07T08:28:07.411165Z","shell.execute_reply.started":"2023-09-07T08:28:07.406510Z","shell.execute_reply":"2023-09-07T08:28:07.410190Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n'''\ntrain_transform = transforms.Compose([\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(degrees=30),\n    transforms.Resize((256, 256)),\n    transforms.ToTensor(),\n])\n'''\ntrain_transform = A.Compose([\n    A.Resize(256, 256),\n    A.HorizontalFlip(p=0.3),\n    A.VerticalFlip(p=0.2),\n    A.Rotate(p=0.3),\n    ToTensorV2(),\n])\n\ndataset = CustomDataset(image_paths=image_list,\n                        mask_paths=mask_list,\n                        transform=train_transform)\n\n","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:07.414927Z","iopub.execute_input":"2023-09-07T08:28:07.415816Z","iopub.status.idle":"2023-09-07T08:28:09.055721Z","shell.execute_reply.started":"2023-09-07T08:28:07.415792Z","shell.execute_reply":"2023-09-07T08:28:09.054769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\nbatch_size = 5\ndataloader = DataLoader(dataset,\n                        batch_size=batch_size,\n                        shuffle=True)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:09.058062Z","iopub.execute_input":"2023-09-07T08:28:09.058955Z","iopub.status.idle":"2023-09-07T08:28:09.064889Z","shell.execute_reply.started":"2023-09-07T08:28:09.058919Z","shell.execute_reply":"2023-09-07T08:28:09.063319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Device configuration\nprint(torch.cuda.is_available())\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:09.066087Z","iopub.execute_input":"2023-09-07T08:28:09.066388Z","iopub.status.idle":"2023-09-07T08:28:09.103698Z","shell.execute_reply.started":"2023-09-07T08:28:09.066363Z","shell.execute_reply":"2023-09-07T08:28:09.102892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.optim as optim\nfrom tqdm import tqdm\n\nn_epoch = 15\nlearning_rate = 0.001\nearly_stopping_patience = 10\ncounter = 0\nmin_val_loss = float('inf')\ncriterion = nn.CrossEntropyLoss()\n\n# Unet\nmodel = UNet(n_channels=3, n_classes=2).to(device)\n\noptimizer = optim.Adam(model.parameters(), lr=learning_rate)","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:09.104955Z","iopub.execute_input":"2023-09-07T08:28:09.105495Z","iopub.status.idle":"2023-09-07T08:28:12.417466Z","shell.execute_reply.started":"2023-09-07T08:28:09.105455Z","shell.execute_reply":"2023-09-07T08:28:12.416403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Define a function to visualize images, masks, and predictions\ndef visualize_images_with_masks(images, masks_true, masks_pred, num_rows=5):\n    fig, axes = plt.subplots(num_rows, 3, figsize=(10, 15))\n    for i in range(num_rows):\n\n\n        ax = axes[i]\n        idx = np.random.randint(len(images))\n        \n        ax[0].imshow(images[idx][:])\n        ax[0].set_title('Image')\n        ax[0].axis('off')\n        \n        ax[1].imshow(masks_true[idx][0], cmap='gray')\n        ax[1].set_title('Ground Truth Mask')\n        ax[1].axis('off')\n        \n        ax[2].imshow(masks_pred[idx][0], cmap='gray')\n        ax[2].set_title('Predicted Mask')\n        ax[2].axis('off')\n\n    plt.tight_layout()\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:12.418911Z","iopub.execute_input":"2023-09-07T08:28:12.419279Z","iopub.status.idle":"2023-09-07T08:28:12.429202Z","shell.execute_reply.started":"2023-09-07T08:28:12.419247Z","shell.execute_reply":"2023-09-07T08:28:12.427328Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train loop\nfor epoch in range(n_epoch):\n    model.train()\n    train_loss = 0.0\n\n    for batch_idx, (img, msk) in enumerate(dataloader):\n        img, msk = img.to(torch.float32).to(device), msk.to(torch.int64).to(device)\n        img = img.float() / 255.0\n        optimizer.zero_grad()\n        \n        outputs = model(img)\n        msk = torch.nn.functional.one_hot(msk, 2)\n        msk = torch.permute(msk, (0,3,1,2))\n        msk = msk.float()\n        \n        loss = criterion(outputs, msk)\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item()\n        \n        image = img.cpu().numpy().transpose(0, 2, 3, 1)\n        predicted_mask = outputs.cpu().detach().numpy()\n        ground_truth_mask = msk.cpu().detach().numpy()\n\n        predicted_mask = np.expand_dims(np.argmax(predicted_mask,1), axis=1)\n        ground_truth_mask = np.expand_dims(np.argmax(ground_truth_mask,1), axis=1)\n\n        predicted_mask = (predicted_mask > 0.5).astype(np.uint8)\n\n        if batch_idx%500 == 0:\n            visualize_images_with_masks(image, ground_truth_mask, predicted_mask, num_rows=5)\n        \n        \n\n    train_loss /= len(dataloader)\n    \n\n    print(f\"Epoch [{epoch+1}/{n_epoch}] Train Loss: {train_loss:.4f}\")","metadata":{"execution":{"iopub.status.busy":"2023-09-07T08:28:12.430634Z","iopub.execute_input":"2023-09-07T08:28:12.431692Z","iopub.status.idle":"2023-09-07T09:35:17.834815Z","shell.execute_reply.started":"2023-09-07T08:28:12.431656Z","shell.execute_reply":"2023-09-07T09:35:17.832930Z"},"trusted":true},"execution_count":null,"outputs":[]}],"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}}