{"metadata":{"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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **Import Libraries**","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport zipfile\n\nimport numpy as np\n\nimport torch\nimport torchvision.transforms as T\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom torchvision.utils import make_grid\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport cv2\nfrom PIL import Image\nimport matplotlib\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:50:50.724851Z","iopub.execute_input":"2023-09-22T13:50:50.725515Z","iopub.status.idle":"2023-09-22T13:50:53.679586Z","shell.execute_reply.started":"2023-09-22T13:50:50.725474Z","shell.execute_reply":"2023-09-22T13:50:53.678563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\nseed_everything(42)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:50:53.684945Z","iopub.execute_input":"2023-09-22T13:50:53.686008Z","iopub.status.idle":"2023-09-22T13:50:53.695601Z","shell.execute_reply.started":"2023-09-22T13:50:53.685964Z","shell.execute_reply":"2023-09-22T13:50:53.694552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Unzip Files\nwith zipfile.ZipFile('/kaggle/input/carvana-image-masking-challenge/train.zip', 'r') as zip_ref:\n    zip_ref.extractall('/kaggle/working/')\n\nwith zipfile.ZipFile('/kaggle/input/carvana-image-masking-challenge/train_masks.zip', 'r') as zip_ref:\n    zip_ref.extractall('/kaggle/working/')","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:50:53.697009Z","iopub.execute_input":"2023-09-22T13:50:53.697612Z","iopub.status.idle":"2023-09-22T13:51:00.094219Z","shell.execute_reply.started":"2023-09-22T13:50:53.697577Z","shell.execute_reply":"2023-09-22T13:51:00.093123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_masks_file_names = sorted(os.listdir('/kaggle/working/train_masks'))\ntrain_file_names = sorted(os.listdir('/kaggle/working/train'))","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:00.096726Z","iopub.execute_input":"2023-09-22T13:51:00.097070Z","iopub.status.idle":"2023-09-22T13:51:00.112112Z","shell.execute_reply.started":"2023-09-22T13:51:00.097043Z","shell.execute_reply":"2023-09-22T13:51:00.111244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"No. of imgs in train_file_names: \", len(train_file_names))\nprint(\"No. of masks in train_masks_file_names:\", len(train_masks_file_names))","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:00.113533Z","iopub.execute_input":"2023-09-22T13:51:00.114334Z","iopub.status.idle":"2023-09-22T13:51:00.912849Z","shell.execute_reply.started":"2023-09-22T13:51:00.114298Z","shell.execute_reply":"2023-09-22T13:51:00.910488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_file_names[:2])\nprint(train_masks_file_names[:2])","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:00.914274Z","iopub.execute_input":"2023-09-22T13:51:00.916132Z","iopub.status.idle":"2023-09-22T13:51:00.923659Z","shell.execute_reply.started":"2023-09-22T13:51:00.916092Z","shell.execute_reply":"2023-09-22T13:51:00.922664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Function to see a particular image and its mask\ndef show_img_and_mask(fname):\n    img = Image.open(\"/kaggle/working/train/\" + fname)\n    img_mask = Image.open(\"/kaggle/working/train_masks/\"+ fname.split('.')[0] + \"_mask.gif\")\n    \n    plt.figure(figsize=(5,5))\n    \n    plt.subplot(1, 2, 1)\n    plt.imshow(img)\n    plt.axis('off')\n    plt.title('Image')\n    \n    plt.subplot(1,2, 2)\n    plt.imshow(img_mask)\n    plt.axis('off')\n    plt.title('Mask')\n\n# Function to see set of images and corresponding masks\ndef show_some_imgs_and_masks(images_number, img_list):\n    sample_img_names = random.sample(img_list, images_number)\n    \n    plt.figure(figsize=(5,5))\n    \n    for ind, img_name in enumerate(sample_img_names):\n        path1 = \"/kaggle/working/train/\" + img_name\n        path2 = \"/kaggle/working/train_masks/\" + img_name.split('.')[0] + \"_mask.gif\"\n        \n        img1 = Image.open(path1)\n        img2 = Image.open(path2)\n        \n        plt.subplot(images_number, 2, 2*ind +1)\n        plt.imshow(img1)\n        plt.title(\"Image\")\n        plt.axis('off')\n        \n        plt.subplot(images_number, 2, 2*ind+2)\n        plt.imshow(img2)\n        plt.title(\"Mask\")\n        plt.axis('off')\n        \n    plt.tight_layout()\n    plt.show()\n    ","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:00.925007Z","iopub.execute_input":"2023-09-22T13:51:00.925842Z","iopub.status.idle":"2023-09-22T13:51:00.938339Z","shell.execute_reply.started":"2023-09-22T13:51:00.925806Z","shell.execute_reply":"2023-09-22T13:51:00.937460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_some_imgs_and_masks(3, train_file_names)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:00.940501Z","iopub.execute_input":"2023-09-22T13:51:00.941550Z","iopub.status.idle":"2023-09-22T13:51:03.206009Z","shell.execute_reply.started":"2023-09-22T13:51:00.941517Z","shell.execute_reply":"2023-09-22T13:51:03.201164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Prepare Data**","metadata":{}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, img_names, mask_names = None, transforms = None, train=True):\n        super(CustomDataset, self).__init__()\n        self.img_names = img_names\n        self.mask_names = mask_names\n        self.train = train\n        self.transforms = transforms\n        \n    def __getitem__(self, idx):\n        img_path = \"/kaggle/working/train/\" + self.img_names[idx]\n        img = np.array(Image.open(img_path).convert(\"RGB\"), dtype = np.float32) / 255.0\n        \n        if self.train:\n            mask_path = \"/kaggle/working/train_masks/\" + self.mask_names[idx]\n            mask = np.array(Image.open(mask_path).convert(\"L\"), dtype = np.float32)\n            mask[mask == 255.0] = 1.0\n            \n            if self.transforms:\n                augmentations = self.transforms(image=img, mask=mask)\n                img = augmentations[\"image\"]\n                mask = augmentations[\"mask\"]\n            mask = mask.unsqueeze(0)\n            return img, mask\n        \n        if self.transforms:\n            img = self.transforms(image=img)\n        \n        mask = mask.unsqueeze(0)\n        return img, mask\n        \n    def __len__(self):\n        return len(self.img_names)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:03.207538Z","iopub.execute_input":"2023-09-22T13:51:03.208027Z","iopub.status.idle":"2023-09-22T13:51:03.218664Z","shell.execute_reply.started":"2023-09-22T13:51:03.207998Z","shell.execute_reply":"2023-09-22T13:51:03.217796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Data Augmentations**","metadata":{}},{"cell_type":"code","source":"IMAGE_HEIGHT = 256\nIMAGE_WIDTH = 256\n\ntrain_transforms = A.Compose([\n    A.Resize(height=IMAGE_HEIGHT, width=IMAGE_WIDTH),\n    A.Rotate(limit=35, p=1.0),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.1),\n    A.RandomBrightnessContrast(brightness_limit=0.3, contrast_limit=0.3, p=0.5),\n    ToTensorV2()\n])\n\nval_transforms = A.Compose([\n    A.Resize(height=IMAGE_HEIGHT, width=IMAGE_WIDTH),\n    ToTensorV2()\n])","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:03.224420Z","iopub.execute_input":"2023-09-22T13:51:03.224695Z","iopub.status.idle":"2023-09-22T13:51:03.242990Z","shell.execute_reply.started":"2023-09-22T13:51:03.224671Z","shell.execute_reply":"2023-09-22T13:51:03.242050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Training and Validation Data**","metadata":{}},{"cell_type":"code","source":"split = int(len(train_file_names) * 0.8)\n\ntrain_imgs = train_file_names[:split]\ntrain_masks = train_masks_file_names[:split]\n\nval_imgs = train_file_names[split:]\nval_masks = train_masks_file_names[split:]","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:03.244272Z","iopub.execute_input":"2023-09-22T13:51:03.245317Z","iopub.status.idle":"2023-09-22T13:51:03.252659Z","shell.execute_reply.started":"2023-09-22T13:51:03.245280Z","shell.execute_reply":"2023-09-22T13:51:03.252051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_imgs), len(val_imgs)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:03.254074Z","iopub.execute_input":"2023-09-22T13:51:03.254587Z","iopub.status.idle":"2023-09-22T13:51:03.263391Z","shell.execute_reply.started":"2023-09-22T13:51:03.254554Z","shell.execute_reply":"2023-09-22T13:51:03.262547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 8\n\ntrain_ds = CustomDataset(train_imgs, train_masks, transforms=train_transforms)\nval_ds = CustomDataset(val_imgs, val_masks, transforms=val_transforms)\n\ntrain_dl = DataLoader(train_ds, batch_size, shuffle=True)\nval_dl = DataLoader(val_ds, batch_size)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:03.264618Z","iopub.execute_input":"2023-09-22T13:51:03.265716Z","iopub.status.idle":"2023-09-22T13:51:03.272062Z","shell.execute_reply.started":"2023-09-22T13:51:03.265683Z","shell.execute_reply":"2023-09-22T13:51:03.271316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_batch(dl):\n    for batch in dl:\n        image, mask = batch\n        \n        fig, (ax1, ax2) = plt.subplots(nrows=2, ncols=1, figsize=(12,4))\n        ax1.imshow(make_grid(image).permute(1,2,0))\n        ax1.set_title('Input Images')\n        ax1.axis('off')\n        \n        ax2.set_title('Input Masks')\n        ax2.imshow(make_grid(mask).permute(1,2,0))  \n        ax2.axis('off')\n        \n        fig.tight_layout()\n        break","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:03.273320Z","iopub.execute_input":"2023-09-22T13:51:03.274504Z","iopub.status.idle":"2023-09-22T13:51:03.284209Z","shell.execute_reply.started":"2023-09-22T13:51:03.274391Z","shell.execute_reply":"2023-09-22T13:51:03.283258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_batch(train_dl)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:03.287300Z","iopub.execute_input":"2023-09-22T13:51:03.287610Z","iopub.status.idle":"2023-09-22T13:51:04.168643Z","shell.execute_reply.started":"2023-09-22T13:51:03.287574Z","shell.execute_reply":"2023-09-22T13:51:04.167602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"show_batch(val_dl)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:04.170125Z","iopub.execute_input":"2023-09-22T13:51:04.170582Z","iopub.status.idle":"2023-09-22T13:51:04.980596Z","shell.execute_reply.started":"2023-09-22T13:51:04.170545Z","shell.execute_reply":"2023-09-22T13:51:04.979680Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Using GPUs**","metadata":{}},{"cell_type":"code","source":"def get_default_device():\n    \"\"\"Pick GPU if available, else CPU\"\"\"\n    if torch.cuda.is_available():\n        return torch.device('cuda')\n    else:\n        return torch.device('cpu')\n    \ndef to_device(data, device):\n    \"\"\"Move tensor(s) to chosen device\"\"\"\n    if isinstance(data, (list,tuple)):\n        return [to_device(x, device) for x in data]\n    return data.to(device, non_blocking=True)\n\nclass DeviceDataLoader():\n    \"\"\"Wrap a dataloader to move data to a device\"\"\"\n    def __init__(self, dl, device):\n        self.dl = dl\n        self.device = device\n        \n    def __iter__(self):\n        \"\"\"Yield a batch of data after moving it to device\"\"\"\n        for b in self.dl: \n            yield to_device(b, self.device)\n\n    def __len__(self):\n        \"\"\"Number of batches\"\"\"\n        return len(self.dl)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:04.982189Z","iopub.execute_input":"2023-09-22T13:51:04.982837Z","iopub.status.idle":"2023-09-22T13:51:04.992528Z","shell.execute_reply.started":"2023-09-22T13:51:04.982802Z","shell.execute_reply":"2023-09-22T13:51:04.991412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = get_default_device()\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:04.994085Z","iopub.execute_input":"2023-09-22T13:51:04.994736Z","iopub.status.idle":"2023-09-22T13:51:05.030147Z","shell.execute_reply.started":"2023-09-22T13:51:04.994701Z","shell.execute_reply":"2023-09-22T13:51:05.029083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dl = DeviceDataLoader(train_dl, device)\nval_dl = DeviceDataLoader(val_dl, device)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:05.031677Z","iopub.execute_input":"2023-09-22T13:51:05.032031Z","iopub.status.idle":"2023-09-22T13:51:05.040374Z","shell.execute_reply.started":"2023-09-22T13:51:05.031990Z","shell.execute_reply":"2023-09-22T13:51:05.039384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Wraper Functions to use during Training**","metadata":{}},{"cell_type":"code","source":"class BaseClass(nn.Module):\n    def training_step(self, batch):\n        inputs, targets = batch        \n        preds = self(inputs)\n        loss_fn = nn.BCEWithLogitsLoss()\n        loss = loss_fn(preds, targets)\n        return loss\n    \n    def validation_step(self, batch, score_fn):\n        inputs, targets = batch\n        preds = self(inputs)\n        loss_fn = nn.BCEWithLogitsLoss()\n        loss = loss_fn(preds, targets)\n        score = score_fn(preds, targets)\n        return {'val_loss': loss.detach(), 'val_score': score}\n    \n    def validation_epoch_end(self, outputs):\n        batch_losses = [x['val_loss'] for x in outputs]\n        epoch_loss = sum(batch_losses)/len(batch_losses)\n        \n        batch_scores = [x['val_score'] for x in outputs]\n        epoch_score = sum(batch_scores)/len(batch_scores)\n        \n        return {'val_loss': epoch_loss.item(), 'val_score': epoch_score}\n    \n    def epoch_end(self, epoch, nEpochs, results):\n        print(\"Epoch: [{}/{}], train_loss: {:.4f}, val_loss: {:.4f}, val_score:{:.4f}\".format(\n                        epoch+1, nEpochs, results['train_loss'], results['val_loss'], results['val_score']))\n","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:05.041929Z","iopub.execute_input":"2023-09-22T13:51:05.042326Z","iopub.status.idle":"2023-09-22T13:51:05.052623Z","shell.execute_reply.started":"2023-09-22T13:51:05.042293Z","shell.execute_reply":"2023-09-22T13:51:05.051041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Model Architecture**","metadata":{}},{"cell_type":"code","source":"class conv_block(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(conv_block, self).__init__()\n        self.block = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size = 3, stride = 1, padding = 1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n\n            nn.Conv2d(out_channels, out_channels, kernel_size = 3,stride = 1, padding = 1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n        \n    def forward(self, x):\n        return self.block(x)\n\ndef copy_and_crop(down_layer, up_layer):\n    b, ch, h, w = up_layer.shape\n    crop = T.CenterCrop((h, w))(down_layer)\n    return crop\n    \nclass UNet(BaseClass):\n    def __init__(self, in_channels, out_channels):\n        super(UNet, self).__init__()\n        \n        self.encoder = nn.ModuleList([\n            conv_block(in_channels, 64),\n            conv_block(64, 128),\n            conv_block(128, 256),\n            conv_block(256, 512)\n        ])\n        \n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)\n        self.bottle_neck = conv_block(512, 1024)\n        \n        self.up_samples = nn.ModuleList([\n            nn.ConvTranspose2d(1024, 512, kernel_size=2, stride=2),\n            nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2),\n            nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2),\n            nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)\n        ])\n        \n        self.decoder = nn.ModuleList([\n            conv_block(1024, 512),\n            conv_block(512, 256),\n            conv_block(256, 128),\n            conv_block(128, 64)\n        ])\n        \n        self.final_layer = nn.Conv2d(64, out_channels, 1, 1)\n        \n    def forward(self, x):\n        skip_connections = []\n        \n        for layer in self.encoder:\n            x = layer(x)\n            skip_connections.append(x)\n            x = self.pool(x)\n        \n        x = self.bottle_neck(x)\n        \n        for ind, layer in enumerate(self.decoder):\n            x = self.up_samples[ind](x)\n            y = copy_and_crop(skip_connections.pop(), x)\n            x = layer(torch.cat([y, x], dim=1))\n        \n        x = self.final_layer(x)\n        \n        return x","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:05.054186Z","iopub.execute_input":"2023-09-22T13:51:05.054566Z","iopub.status.idle":"2023-09-22T13:51:05.071319Z","shell.execute_reply.started":"2023-09-22T13:51:05.054533Z","shell.execute_reply":"2023-09-22T13:51:05.070683Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Training**","metadata":{}},{"cell_type":"code","source":"def dice_score(preds, targets):\n    preds = F.sigmoid(preds)\n    preds = (preds > 0.5).float()\n    score = (2. * (preds * targets).sum()) / (preds + targets).sum()\n    return torch.mean(score).item()\n\ndef evaluation(model, val_dl):\n    model.eval()\n    outputs = [model.validation_step(batch, dice_score) for batch in val_dl]\n    return model.validation_epoch_end(outputs)\n\ndef show_predicted_images(model, dl, n_images = batch_size):\n    imgs, masks = next(iter(dl))\n    logits = model(imgs)\n    preds = F.sigmoid(logits)\n    preds = (preds>0.5).float()\n    \n    fig,axs = plt.subplots(nrows=3, ncols=1)\n    ax1, ax2, ax3 = axs\n    fig.tight_layout()\n\n    ax1.imshow(make_grid(imgs[:n_images].detach().cpu(), nrow=n_images).permute(1,2,0))\n    ax1.set_title('Input Image')\n    ax1.axis('off')\n    \n    ax2.imshow(make_grid(masks[:n_images].detach().cpu(), nrow=n_images).permute(1,2,0))\n    ax2.set_title('Mask')\n    ax2.axis('off')\n    \n    ax3.imshow(make_grid(preds[:n_images].detach().cpu(), nrow=n_images).permute(1,2,0))\n    ax3.set_title('Pred Mask')\n    ax3.axis('off')\n    \n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:05.072492Z","iopub.execute_input":"2023-09-22T13:51:05.073654Z","iopub.status.idle":"2023-09-22T13:51:05.085366Z","shell.execute_reply.started":"2023-09-22T13:51:05.073627Z","shell.execute_reply":"2023-09-22T13:51:05.084677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fit(epochs, lr, model, train_dl, val_dl, opt_func, print_after):\n    history = []\n    optimizer = opt_func(model.parameters(), lr=lr)\n    \n    for epoch in range(epochs):\n        model.train()\n        train_losses = []\n        \n        for batch in tqdm(train_dl):\n            loss = model.training_step(batch)\n            \n            optimizer.zero_grad()\n            loss.backward()\n            optimizer.step()\n            train_losses.append(loss.detach())\n        \n        # Validation\n        result = evaluation(model, val_dl)\n        result['train_loss'] = torch.stack(train_losses).mean().item()\n        history.append(result)\n        \n        model.epoch_end(epoch, epochs, result)\n        if epoch%print_after == 0:\n            show_predicted_images(model, val_dl, n_images=4)\n        \n    return history","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:05.086529Z","iopub.execute_input":"2023-09-22T13:51:05.087630Z","iopub.status.idle":"2023-09-22T13:51:05.097980Z","shell.execute_reply.started":"2023-09-22T13:51:05.087598Z","shell.execute_reply":"2023-09-22T13:51:05.097109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = to_device(UNet(3, 1), device)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:05.099233Z","iopub.execute_input":"2023-09-22T13:51:05.099714Z","iopub.status.idle":"2023-09-22T13:51:06.792480Z","shell.execute_reply.started":"2023-09-22T13:51:05.099680Z","shell.execute_reply":"2023-09-22T13:51:06.791306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr = 1e-4\nepochs = 5\nopt_func = torch.optim.Adam","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:06.794147Z","iopub.execute_input":"2023-09-22T13:51:06.794519Z","iopub.status.idle":"2023-09-22T13:51:06.800152Z","shell.execute_reply.started":"2023-09-22T13:51:06.794485Z","shell.execute_reply":"2023-09-22T13:51:06.798841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = [evaluation(model, val_dl)]","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:51:06.803343Z","iopub.execute_input":"2023-09-22T13:51:06.803628Z","iopub.status.idle":"2023-09-22T13:52:07.846207Z","shell.execute_reply.started":"2023-09-22T13:51:06.803599Z","shell.execute_reply":"2023-09-22T13:52:07.845075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:52:07.848039Z","iopub.execute_input":"2023-09-22T13:52:07.848437Z","iopub.status.idle":"2023-09-22T13:52:07.856455Z","shell.execute_reply.started":"2023-09-22T13:52:07.848402Z","shell.execute_reply":"2023-09-22T13:52:07.855426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history += fit(epochs, lr, model, train_dl, val_dl, opt_func, print_after=2)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T13:52:07.863056Z","iopub.execute_input":"2023-09-22T13:52:07.863937Z","iopub.status.idle":"2023-09-22T14:10:15.893845Z","shell.execute_reply.started":"2023-09-22T13:52:07.863902Z","shell.execute_reply":"2023-09-22T14:10:15.892792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Plots** ","metadata":{}},{"cell_type":"code","source":"# Plot Figures\ntrain_losses = [ele.get('train_loss') for ele in history]\nval_losses = [ele.get('val_loss') for ele in history]\n\nplt.plot(train_losses, '-bx')\nplt.plot(val_losses, '-rx')\nplt.xlabel('epoch')\nplt.ylabel('loss')\nplt.legend(['Train', 'Validation'])\nplt.title('Loss vs. No. of epochs');","metadata":{"execution":{"iopub.status.busy":"2023-09-22T14:10:15.898707Z","iopub.execute_input":"2023-09-22T14:10:15.901123Z","iopub.status.idle":"2023-09-22T14:10:16.255606Z","shell.execute_reply.started":"2023-09-22T14:10:15.901084Z","shell.execute_reply":"2023-09-22T14:10:16.254664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Prediction**","metadata":{}},{"cell_type":"code","source":"def predict_mask(img_path):\n    img = np.array(Image.open(img_path).convert(\"RGB\"), dtype=np.float32) / 255.0\n    img_t = val_transforms(image=img)['image'].to(device)\n    \n    logits = model(img_t.unsqueeze(0)).detach().cpu()\n    preds = F.sigmoid(logits)\n    preds = (preds>0.5).float().detach().cpu()\n    \n    plt.imshow(preds[0].permute(1,2,0), cmap='gray')\n    plt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2023-09-22T14:10:16.257042Z","iopub.execute_input":"2023-09-22T14:10:16.258101Z","iopub.status.idle":"2023-09-22T14:10:16.265427Z","shell.execute_reply.started":"2023-09-22T14:10:16.258065Z","shell.execute_reply":"2023-09-22T14:10:16.264411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_path = \"/kaggle/working/train/\"+val_imgs[8]\npredict_mask(img_path)","metadata":{"execution":{"iopub.status.busy":"2023-09-22T14:11:54.321039Z","iopub.execute_input":"2023-09-22T14:11:54.321463Z","iopub.status.idle":"2023-09-22T14:11:54.472295Z","shell.execute_reply.started":"2023-09-22T14:11:54.321424Z","shell.execute_reply":"2023-09-22T14:11:54.471068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), 'unet_segmentation.pth')","metadata":{"execution":{"iopub.status.busy":"2023-09-22T14:10:19.941065Z","iopub.execute_input":"2023-09-22T14:10:19.941399Z","iopub.status.idle":"2023-09-22T14:10:20.147924Z","shell.execute_reply.started":"2023-09-22T14:10:19.941365Z","shell.execute_reply":"2023-09-22T14:10:20.146878Z"},"trusted":true},"execution_count":null,"outputs":[]}]}