{"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":"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)\nimport matplotlib.pyplot as plt\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":"2022-09-10T05:39:55.109964Z","iopub.execute_input":"2022-09-10T05:39:55.110475Z","iopub.status.idle":"2022-09-10T05:39:55.140683Z","shell.execute_reply.started":"2022-09-10T05:39:55.110389Z","shell.execute_reply":"2022-09-10T05:39:55.139717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:39:55.142461Z","iopub.execute_input":"2022-09-10T05:39:55.142916Z","iopub.status.idle":"2022-09-10T05:39:56.239914Z","shell.execute_reply.started":"2022-09-10T05:39:55.142878Z","shell.execute_reply":"2022-09-10T05:39:56.238340Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%matplotlib inline","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:39:56.241654Z","iopub.execute_input":"2022-09-10T05:39:56.243857Z","iopub.status.idle":"2022-09-10T05:39:56.261222Z","shell.execute_reply.started":"2022-09-10T05:39:56.243798Z","shell.execute_reply":"2022-09-10T05:39:56.259563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from zipfile import ZipFile\nfrom fastai.vision.all import Path, get_image_files\nfrom skimage import io, transform","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:39:56.268162Z","iopub.execute_input":"2022-09-10T05:39:56.269325Z","iopub.status.idle":"2022-09-10T05:40:00.538947Z","shell.execute_reply.started":"2022-09-10T05:39:56.269274Z","shell.execute_reply":"2022-09-10T05:40:00.537702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#Copy and extract the train file\nwith ZipFile('../input/carvana-image-masking-challenge/train.zip', 'r') as zip_ref:\n  zip_ref.extractall('')","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:40:00.541227Z","iopub.execute_input":"2022-09-10T05:40:00.543722Z","iopub.status.idle":"2022-09-10T05:40:11.144675Z","shell.execute_reply.started":"2022-09-10T05:40:00.543675Z","shell.execute_reply":"2022-09-10T05:40:11.143480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with ZipFile('../input/carvana-image-masking-challenge/train_masks.zip', 'r') as zip_ref:\n  zip_ref.extractall('')","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:40:11.146102Z","iopub.execute_input":"2022-09-10T05:40:11.151731Z","iopub.status.idle":"2022-09-10T05:40:12.567542Z","shell.execute_reply.started":"2022-09-10T05:40:11.151690Z","shell.execute_reply":"2022-09-10T05:40:12.566534Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with ZipFile('../input/carvana-image-masking-challenge/sample_submission.csv.zip', 'r') as zip_ref:\n  zip_ref.extractall('')","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:40:12.569184Z","iopub.execute_input":"2022-09-10T05:40:12.569555Z","iopub.status.idle":"2022-09-10T05:40:12.592583Z","shell.execute_reply.started":"2022-09-10T05:40:12.569519Z","shell.execute_reply":"2022-09-10T05:40:12.591578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with ZipFile('../input/carvana-image-masking-challenge/test.zip', 'r') as zip_ref:\n  zip_ref.extractall('')","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:40:12.594120Z","iopub.execute_input":"2022-09-10T05:40:12.594498Z","iopub.status.idle":"2022-09-10T05:43:05.590844Z","shell.execute_reply.started":"2022-09-10T05:40:12.594461Z","shell.execute_reply":"2022-09-10T05:43:05.589834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#View the file\npath = Path('')\nfnames = get_image_files(path/'train')\nlbl_names = get_image_files(path/'train_masks')","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:05.592351Z","iopub.execute_input":"2022-09-10T05:43:05.592708Z","iopub.status.idle":"2022-09-10T05:43:05.667974Z","shell.execute_reply.started":"2022-09-10T05:43:05.592670Z","shell.execute_reply":"2022-09-10T05:43:05.666992Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fnames.sort()\nlbl_names.sort()","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:05.671989Z","iopub.execute_input":"2022-09-10T05:43:05.672371Z","iopub.status.idle":"2022-09-10T05:43:05.741839Z","shell.execute_reply.started":"2022-09-10T05:43:05.672338Z","shell.execute_reply":"2022-09-10T05:43:05.740968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f, axarr = plt.subplots(1,2)\nidx = 8\nim = io.imread(lbl_names[idx])\naxarr[0].imshow(im)\nim = io.imread(fnames[idx])\naxarr[1].imshow(im)","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:05.743289Z","iopub.execute_input":"2022-09-10T05:43:05.743651Z","iopub.status.idle":"2022-09-10T05:43:07.502827Z","shell.execute_reply.started":"2022-09-10T05:43:05.743616Z","shell.execute_reply":"2022-09-10T05:43:07.501801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"type(lbl_names[0])","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:07.507104Z","iopub.execute_input":"2022-09-10T05:43:07.509715Z","iopub.status.idle":"2022-09-10T05:43:07.519770Z","shell.execute_reply.started":"2022-09-10T05:43:07.509659Z","shell.execute_reply":"2022-09-10T05:43:07.518880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:07.524209Z","iopub.execute_input":"2022-09-10T05:43:07.526689Z","iopub.status.idle":"2022-09-10T05:43:07.533817Z","shell.execute_reply.started":"2022-09-10T05:43:07.526649Z","shell.execute_reply":"2022-09-10T05:43:07.532822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset\nfrom torchvision import transforms\nfrom PIL import Image\nclass ImageMasksDataset(Dataset):\n    \"\"\"Face Landmarks dataset.\"\"\"\n\n    def __init__(self, fnames, lbl_names, transform=None):\n        \"\"\"\n        Args:\n            csv_file (string): Path to the csv file with annotations.\n            root_dir (string): Directory with all the images.\n            transform (callable, optional): Optional transform to be applied\n                on a sample.\n        \"\"\"\n        self.fnames = fnames\n        self.lbl_names = lbl_names\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.fnames)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n        \n        sample = (Image.open(self.fnames[idx]),Image.open(self.lbl_names[idx]))\n\n        if self.transform:\n#             convert = transforms.ToTensor()\n            if len(self.transform)>1:\n                img = Image.fromarray(np.array(sample[1])>0)\n                sample = (self.transform[0](sample[0]), torch.squeeze(self.transform[1](img)).long())\n            else:\n                sample = (self.transform[0](sample[0]), torch.squeeze(self.transform[0](sample[1])))\n\n#             sample = self.transform(image=sample[0], mask=sample[1])\n#             sample = (sample['image'], sample['mask'])\n#             sample = (convert(sample[0]), convert(sample[1]))\n        return sample","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:07.536216Z","iopub.execute_input":"2022-09-10T05:43:07.538146Z","iopub.status.idle":"2022-09-10T05:43:07.556015Z","shell.execute_reply.started":"2022-09-10T05:43:07.538071Z","shell.execute_reply":"2022-09-10T05:43:07.554972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = ImageMasksDataset(fnames, lbl_names)\nf, axarr = plt.subplots(1,2)\naxarr[0].imshow(dataset[0][0])\naxarr[1].imshow(dataset[0][1])","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:07.560573Z","iopub.execute_input":"2022-09-10T05:43:07.563320Z","iopub.status.idle":"2022-09-10T05:43:08.396090Z","shell.execute_reply.started":"2022-09-10T05:43:07.563275Z","shell.execute_reply":"2022-09-10T05:43:08.394993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import albumentations as A\nimport albumentations.pytorch as alb_pytorch\nimport cv2\nfrom torchvision import transforms, utils\nfrom torchvision.transforms import InterpolationMode\n\nresize_num = 256\n\n# transform_seg = A.Compose([\n#     alb_pytorch.ToTensorV2(),\n#     A.Resize(width=256, height=256),\n# #     A.HorizontalFlip(p=0.5),\n# #     A.RandomBrightnessContrast(p=0.2),\n# ])\ntransform_seg_img=transforms.Compose([transforms.Resize((resize_num,resize_num)),transforms.ToTensor()])\ntransform_seg_mask=transforms.Compose([transforms.ToTensor(),transforms.Resize((resize_num,resize_num), interpolation=InterpolationMode.NEAREST)])","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:08.397390Z","iopub.execute_input":"2022-09-10T05:43:08.399279Z","iopub.status.idle":"2022-09-10T05:43:09.516103Z","shell.execute_reply.started":"2022-09-10T05:43:08.399239Z","shell.execute_reply":"2022-09-10T05:43:09.515039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\ndataset = ImageMasksDataset(fnames, lbl_names, transform=[transform_seg_img, transform_seg_mask])\ntrainloader = DataLoader(dataset, batch_size=4, shuffle=True, num_workers=0)\n# dataset = ImageMasksDataset(fnames_train[:5], lbl_names_train[:5], transform_seg)\n# trainloader = DataLoader(dataset, batch_size=4, shuffle=True, num_workers=0)\nfor i, (image_in, image_mask) in enumerate(trainloader):\n#     print(np.unique(image_mask[0]))\n    f, axarr = plt.subplots(1,2)\n    axarr[0].imshow(image_in[0].permute(1, 2, 0))\n    axarr[1].imshow(image_mask[0])\n    if i == 4:\n        break","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:09.517412Z","iopub.execute_input":"2022-09-10T05:43:09.519608Z","iopub.status.idle":"2022-09-10T05:43:12.325815Z","shell.execute_reply.started":"2022-09-10T05:43:09.519579Z","shell.execute_reply":"2022-09-10T05:43:12.324699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.nn import Flatten, Linear, ReLU, Conv2d, MaxPool2d, Sigmoid, Upsample","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:12.327488Z","iopub.execute_input":"2022-09-10T05:43:12.328130Z","iopub.status.idle":"2022-09-10T05:43:12.333355Z","shell.execute_reply.started":"2022-09-10T05:43:12.328089Z","shell.execute_reply":"2022-09-10T05:43:12.332043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\" Parts of the U-Net model \"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\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)","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:12.334778Z","iopub.execute_input":"2022-09-10T05:43:12.335416Z","iopub.status.idle":"2022-09-10T05:43:13.491555Z","shell.execute_reply.started":"2022-09-10T05:43:12.335375Z","shell.execute_reply":"2022-09-10T05:43:13.490446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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 = self.outc(x)\n        return logits","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:13.492858Z","iopub.execute_input":"2022-09-10T05:43:13.493414Z","iopub.status.idle":"2022-09-10T05:43:13.905805Z","shell.execute_reply.started":"2022-09-10T05:43:13.493377Z","shell.execute_reply":"2022-09-10T05:43:13.904577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet(3,2)\n# model = torch.nn.modules.Sequential(Conv2d(in_channels=3, out_channels=16, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=16, out_channels=32, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     MaxPool2d(2),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=32, out_channels=64, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=64, out_channels=128, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     MaxPool2d(2),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=128, out_channels=64, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     Upsample(scale_factor=2, mode='bilinear'),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=64, out_channels=32, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=32, out_channels=16, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     Upsample(scale_factor=2, mode='bilinear'),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=16, out_channels=2, kernel_size=3, padding=1, stride=1),\n#                                     )\n\n# model = torch.nn.modules.Sequential(Conv2d(in_channels=3, out_channels=64, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=64, out_channels=128, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=128, out_channels=256, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=256, out_channels=128, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=128, out_channels=64, kernel_size=3, padding=1, stride=1),\n#                                     ReLU(),\n#                                     Conv2d(in_channels=64, out_channels=2, kernel_size=3, padding=1, stride=1),\n# )","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:13.907840Z","iopub.execute_input":"2022-09-10T05:43:13.908694Z","iopub.status.idle":"2022-09-10T05:43:14.226890Z","shell.execute_reply.started":"2022-09-10T05:43:13.908653Z","shell.execute_reply":"2022-09-10T05:43:14.225721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nidx = list(range(len(fnames)))\nX_train, X_val, _, _ = train_test_split(idx, idx, test_size=0.2, random_state=123)\nfnames_train = [fnames[i] for i in X_train]\nlbl_names_train = [lbl_names[i] for i in X_train]\nfnames_val = [fnames[i] for i in X_val]\nlbl_names_val = [lbl_names[i] for i in X_val]\n\ntrain_dataset = ImageMasksDataset(fnames_train, lbl_names_train, transform=[transform_seg_img, transform_seg_mask])\ntrainloader = DataLoader(train_dataset, batch_size=30, shuffle=True, num_workers=0)\n\ntest_dataset = ImageMasksDataset(fnames_val, lbl_names_val, transform=[transform_seg_img, transform_seg_mask])\ntestloader = DataLoader(test_dataset, batch_size=30, shuffle=True, num_workers=0)\n\ndataloaders = {'train':trainloader, 'val':testloader}\ndataset_sizes = {'train':len(train_dataset), 'val':len(test_dataset)}","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:14.228668Z","iopub.execute_input":"2022-09-10T05:43:14.229101Z","iopub.status.idle":"2022-09-10T05:43:14.254668Z","shell.execute_reply.started":"2022-09-10T05:43:14.229059Z","shell.execute_reply":"2022-09-10T05:43:14.253577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(len(trainloader))\n# for i, (image_in, image_mask) in enumerate(trainloader):\n# #     print((image_in[0]).shape)\n#     f, axarr = plt.subplots(1,2)\n#     axarr[0].imshow(image_in[0].permute(1,2,0))\n#     axarr[1].imshow(image_mask[0].reshape(256,256))\n#     if i == 4:\n#         break","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:14.256211Z","iopub.execute_input":"2022-09-10T05:43:14.256669Z","iopub.status.idle":"2022-09-10T05:43:14.270711Z","shell.execute_reply.started":"2022-09-10T05:43:14.256602Z","shell.execute_reply":"2022-09-10T05:43:14.269679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.tensorboard import SummaryWriter\n%load_ext tensorboard","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:14.273202Z","iopub.execute_input":"2022-09-10T05:43:14.273884Z","iopub.status.idle":"2022-09-10T05:43:14.553837Z","shell.execute_reply.started":"2022-09-10T05:43:14.273849Z","shell.execute_reply":"2022-09-10T05:43:14.552932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\nclass DiceLoss(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(DiceLoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        \n        #comment out if your model contains a sigmoid or equivalent activation layer\n        inputs = torch.sigmoid(inputs)       \n        \n        #flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()                            \n        dice = (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n        \n        return 1 - dice","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:14.555098Z","iopub.execute_input":"2022-09-10T05:43:14.555481Z","iopub.status.idle":"2022-09-10T05:43:14.563040Z","shell.execute_reply.started":"2022-09-10T05:43:14.555440Z","shell.execute_reply":"2022-09-10T05:43:14.561527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:14.564259Z","iopub.execute_input":"2022-09-10T05:43:14.564613Z","iopub.status.idle":"2022-09-10T05:43:14.576660Z","shell.execute_reply.started":"2022-09-10T05:43:14.564578Z","shell.execute_reply":"2022-09-10T05:43:14.575487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x,y = next(iter(testloader))\ny.shape","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:14.579646Z","iopub.execute_input":"2022-09-10T05:43:14.579951Z","iopub.status.idle":"2022-09-10T05:43:16.467160Z","shell.execute_reply.started":"2022-09-10T05:43:14.579926Z","shell.execute_reply":"2022-09-10T05:43:16.465902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir models","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:16.472850Z","iopub.execute_input":"2022-09-10T05:43:16.473151Z","iopub.status.idle":"2022-09-10T05:43:17.508576Z","shell.execute_reply.started":"2022-09-10T05:43:16.473124Z","shell.execute_reply":"2022-09-10T05:43:17.507222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls models","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:17.511261Z","iopub.execute_input":"2022-09-10T05:43:17.512500Z","iopub.status.idle":"2022-09-10T05:43:18.626690Z","shell.execute_reply.started":"2022-09-10T05:43:17.512459Z","shell.execute_reply":"2022-09-10T05:43:18.625137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\nimport os\nimport copy\n\ndef train_model(model, criterion, optimizer, scheduler, dice_loss, num_epochs=25):\n    since = time.time()\n    train_losses = []\n    val_losses = []\n    train_accuracies = []\n    val_accuracies = []\n    train_dice = []\n    val_dice = []\n\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_acc = 0.0\n\n    for epoch in range(num_epochs):\n        print(f'Epoch {epoch}/{num_epochs - 1}')\n        print('-' * 10)\n\n        # Each epoch has a training and validation phase\n        for phase in ['train', 'val']:\n            if phase == 'train':\n                model.train()  # Set model to training mode\n            else:\n                model.eval()   # Set model to evaluate mode\n\n            running_loss = 0.0\n            running_corrects = 0\n            running_dice = 0.0\n\n            # Iterate over data.\n            for inputs, labels in dataloaders[phase]:\n                inputs = inputs.cuda()\n                labels = labels.cuda()\n\n                # zero the parameter gradients\n                optimizer.zero_grad()\n\n                # forward\n                # track history if only in train\n                with torch.set_grad_enabled(phase == 'train'):                    \n                    out = model(inputs)\n            \n                    preds = torch.argmax(out, axis=1)\n\n                    dice_l = dice_loss(preds,labels)\n                    loss = criterion(out, labels)\n                        \n\n                    # backward + optimize only if in training phase\n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()\n\n                # statistics\n                running_loss += loss.item() * inputs.size(0)\n                running_corrects += torch.sum(preds == labels.data)\n                running_dice += dice_l\n            \n\n            epoch_loss = running_loss / dataset_sizes[phase]\n            epoch_acc = running_corrects.double().item() / (dataset_sizes[phase]*(resize_num*resize_num))\n            epoch_dice = running_dice.item() / dataset_sizes[phase]\n            \n            if phase == 'train':\n                scheduler.step(epoch_loss)\n                \n                train_losses.append(epoch_loss)\n                train_accuracies.append(epoch_acc)\n                train_dice.append(epoch_dice)\n                \n                torch.save({\n                    'epoch': epoch,\n                    'model_state_dict': model.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict(),\n                    }, 'models/latest_checkpoint.pt')\n            else:\n                val_losses.append(epoch_loss)\n                val_accuracies.append(epoch_acc)\n                val_dice.append(epoch_dice)\n            \n            print(f'{phase} Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}')\n\n            # deep copy the model\n            if phase == 'val' and epoch_acc > best_acc:\n                best_acc = epoch_acc\n                best_model_wts = copy.deepcopy(model.state_dict())\n                print('Best checkpoint saved!')\n                torch.save({\n                    'epoch': epoch,\n                    'model_state_dict': model.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict(),\n                    'best_acc': best_acc\n                    }, 'models/best_checkpoint.pt')\n\n        print()\n\n    time_elapsed = time.time() - since\n    print(f'Training complete in {time_elapsed // 60:.0f}m {time_elapsed % 60:.0f}s')\n    print(f'Best val Acc: {best_acc:4f}')\n\n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    return model, [train_losses, train_accuracies, train_dice], [val_losses, val_accuracies, val_dice]","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:18.632096Z","iopub.execute_input":"2022-09-10T05:43:18.634317Z","iopub.status.idle":"2022-09-10T05:43:18.661117Z","shell.execute_reply.started":"2022-09-10T05:43:18.634271Z","shell.execute_reply":"2022-09-10T05:43:18.660084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = model.cuda()\n\nnum_epochs = 80\noptimizer = torch.optim.Adam(model.parameters(), lr=3e-4)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=10, factor=0.01)    \n# print_every = 20\nloss_func = torch.nn.functional.cross_entropy\ndice = DiceLoss()\n\nmodel, train_stats, val_stats = train_model(model, loss_func, optimizer, scheduler, dice, num_epochs)","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:43:18.662863Z","iopub.execute_input":"2022-09-10T05:43:18.663291Z","iopub.status.idle":"2022-09-10T05:51:29.865624Z","shell.execute_reply.started":"2022-09-10T05:43:18.663253Z","shell.execute_reply":"2022-09-10T05:51:29.863968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import csv\n\nxs = list(range(num_epochs))\n\nwith open('out.csv', 'w', encoding='utf-8', newline='') as out:\n    writer = csv.writer(out)\n    writer.writerow(['Epoch', 'Loss', 'Accuracy', 'Dice Loss', 'Dice Score'])\n    for x in xs:\n        writer.writerow([str(x)+'_train', train_stats[0][x], train_stats[1][x], train_stats[2][x], 1-float(train_stats[2][x])])\n    for x in xs:\n        writer.writerow([str(x)+'_val', val_stats[0][x], val_stats[1][x], val_stats[2][x], 1-float(val_stats[2][x])])","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:51:40.295752Z","iopub.execute_input":"2022-09-10T05:51:40.296141Z","iopub.status.idle":"2022-09-10T05:51:40.323384Z","shell.execute_reply.started":"2022-09-10T05:51:40.296106Z","shell.execute_reply":"2022-09-10T05:51:40.322077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nf, axarr = plt.subplots(1,3, figsize=(20, 10))\naxarr[0].plot(xs, train_stats[0], c='blue')\naxarr[0].plot(xs, val_stats[0], c='orange')\naxarr[0].set_xlabel(\"Epochs\")\naxarr[0].set_ylabel(\"Training Loss\")\naxarr[1].plot(xs, train_stats[1], c='blue')\naxarr[1].plot(xs, val_stats[1], c='orange')\naxarr[1].set_xlabel(\"Epochs\")\naxarr[1].set_ylabel(\"Training Accuracy\")\naxarr[2].plot(xs, train_stats[2], c='blue')\naxarr[2].plot(xs, val_stats[2], c='orange')\naxarr[2].set_xlabel(\"Epochs\")\naxarr[2].set_ylabel(\"Training Dice Loss\")\nf.savefig('visuals.png')","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:51:40.581612Z","iopub.execute_input":"2022-09-10T05:51:40.581984Z","iopub.status.idle":"2022-09-10T05:51:40.942132Z","shell.execute_reply.started":"2022-09-10T05:51:40.581951Z","shell.execute_reply":"2022-09-10T05:51:40.940953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"f, axarr = plt.subplots(1,3, figsize=(20, 10))\naxarr[0].plot(xs, val_stats[0])\naxarr[0].set_xlabel(\"Epochs\")\naxarr[0].set_ylabel(\"Validation Loss\")\naxarr[1].plot(xs, val_stats[1])\naxarr[1].set_xlabel(\"Epochs\")\naxarr[1].set_ylabel(\"Validation Accuracy\")\naxarr[2].plot(xs, val_stats[2])\naxarr[2].set_xlabel(\"Epochs\")\naxarr[2].set_ylabel(\"Validation Dice Loss\")\nf.savefig('validation.png')","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:51:41.850970Z","iopub.execute_input":"2022-09-10T05:51:41.852131Z","iopub.status.idle":"2022-09-10T05:51:42.207582Z","shell.execute_reply.started":"2022-09-10T05:51:41.852083Z","shell.execute_reply":"2022-09-10T05:51:42.206284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load checkpoint\ncheckpoint = torch.load('models/best_checkpoint.pt') # or latest_checkpoint.pt\n\nmodel.load_state_dict(checkpoint['model_state_dict'])","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:51:42.503344Z","iopub.execute_input":"2022-09-10T05:51:42.504002Z","iopub.status.idle":"2022-09-10T05:51:42.788047Z","shell.execute_reply.started":"2022-09-10T05:51:42.503958Z","shell.execute_reply":"2022-09-10T05:51:42.786934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = Path('')\ntest_names = get_image_files(path/'test')\ntest_names.sort()\nplt.imshow(io.imread(test_names[16]))\nio.imread(test_names[16]).shape","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:51:46.438263Z","iopub.execute_input":"2022-09-10T05:51:46.439195Z","iopub.status.idle":"2022-09-10T05:51:49.161862Z","shell.execute_reply.started":"2022-09-10T05:51:46.439154Z","shell.execute_reply":"2022-09-10T05:51:49.160982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image\n\nmodel = model.cuda()\nconvert = transforms.ToTensor()\nf, axarr = plt.subplots(5,5, figsize=(10, 10))\n# plt.figure(figsize=(10,10))\nfor x in range(5):\n    for i in range(5):\n        img = transform_seg_img(Image.open(test_names[i+x]))\n        img = img.cuda()\n#         img = convert(img['image']).cuda().reshape(1,3,1000,1000)\n#         print(img.shape)\n        l = model(img.reshape(1,img.shape[0],img.shape[1],img.shape[2]))\n        preds = torch.argmax(l, axis=1)\n#         print(preds)\n    #     print(l[0].shape)\n        l = preds.cpu().detach().numpy()\n        axarr[x,i].imshow(l.reshape(resize_num,resize_num), cmap='gray')","metadata":{"execution":{"iopub.status.busy":"2022-09-10T06:00:46.762021Z","iopub.execute_input":"2022-09-10T06:00:46.762998Z","iopub.status.idle":"2022-09-10T06:00:49.947716Z","shell.execute_reply.started":"2022-09-10T06:00:46.762961Z","shell.execute_reply":"2022-09-10T06:00:49.946793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_encode(mask_image):\n    pixels = mask_image.flatten()\n    # We avoid issues with '1' at the start or end (at the corners of \n    # the original image) by setting those pixels to '0' explicitly.\n    # We do not expect these to be non-zero for an accurate mask, \n    # so this should not harm the score.\n    pixels[0] = 0\n    pixels[-1] = 0\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    runs[1::2] = runs[1::2] - runs[:-1:2]\n    return runs\n\ndef rle2mask(mask_rle: str, label=1, shape=[1280, 1918]):\n    \"\"\"\n    mask_rle: run-length as string formatted (start length)\n    shape: (height,width) of array to return\n    Returns numpy array, 1 - mask, 0 - background\n\n    \"\"\"\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0] * shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = label\n    return img.reshape(shape)  # Needed to align to RLE direction","metadata":{"execution":{"iopub.status.busy":"2022-09-10T06:00:55.359621Z","iopub.execute_input":"2022-09-10T06:00:55.359993Z","iopub.status.idle":"2022-09-10T06:00:55.370995Z","shell.execute_reply.started":"2022-09-10T06:00:55.359960Z","shell.execute_reply":"2022-09-10T06:00:55.369897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import csv\nimport pandas as pd\n\nout = pd.DataFrame(columns=['img', 'rle_mask'])\n\nmissing = []\nmodel.cuda()\n\nwith open('sample_submission.csv', 'r', encoding='utf-8') as sample:\n    reader = csv.reader(sample)\n    next(reader)\n    for i, row in enumerate(reader):\n        if (i%1000 == 0):\n            print(i)\n        x = []\n        x.append(row[0])\n        img = transform_seg_img(Image.open('test/'+row[0]))\n        img = img.cuda()\n        with torch.no_grad():\n            l = model(img.reshape(1,img.shape[0],img.shape[1],img.shape[2]))\n        preds = torch.argmax(l, axis=1)\n        l = preds.reshape(1,resize_num,resize_num)\n        trans = transforms.Resize((1280,1918))\n        l = trans(l)\n        l = l.cpu().detach().squeeze()\n        rle = rle_encode(np.array(l))\n        rle = ' '.join([str(x) for x in rle])\n        x.append(rle)\n        out = out.append({'img':x[0], 'rle_mask':x[1]}, ignore_index=True)\n        \n        if i == 10:\n            break\n\nout.to_csv('end.csv', index = False, sep=',')","metadata":{"execution":{"iopub.status.busy":"2022-09-10T06:03:47.024802Z","iopub.execute_input":"2022-09-10T06:03:47.025204Z","iopub.status.idle":"2022-09-10T06:03:47.869349Z","shell.execute_reply.started":"2022-09-10T06:03:47.025166Z","shell.execute_reply":"2022-09-10T06:03:47.868248Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv('end.csv')\nplt.imshow(rle2mask(df['rle_mask'][5]), cmap='gray')","metadata":{"execution":{"iopub.status.busy":"2022-09-10T06:03:52.434973Z","iopub.execute_input":"2022-09-10T06:03:52.435576Z","iopub.status.idle":"2022-09-10T06:03:52.839546Z","shell.execute_reply.started":"2022-09-10T06:03:52.435540Z","shell.execute_reply":"2022-09-10T06:03:52.837503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:51:51.164160Z","iopub.status.idle":"2022-09-10T05:51:51.164933Z","shell.execute_reply.started":"2022-09-10T05:51:51.164674Z","shell.execute_reply":"2022-09-10T05:51:51.164698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf test train train_masks","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:51:29.890194Z","iopub.status.idle":"2022-09-10T05:51:29.891273Z","shell.execute_reply.started":"2022-09-10T05:51:29.890998Z","shell.execute_reply":"2022-09-10T05:51:29.891039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:51:29.892523Z","iopub.status.idle":"2022-09-10T05:51:29.894485Z","shell.execute_reply.started":"2022-09-10T05:51:29.894217Z","shell.execute_reply":"2022-09-10T05:51:29.894245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('-'*20)\nprint('end')\nprint('-'*20)","metadata":{"execution":{"iopub.status.busy":"2022-09-10T05:51:29.895904Z","iopub.status.idle":"2022-09-10T05:51:29.896394Z","shell.execute_reply.started":"2022-09-10T05:51:29.896152Z","shell.execute_reply":"2022-09-10T05:51:29.896174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}