{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":6927,"databundleVersionId":45059,"sourceType":"competition"}],"dockerImageVersionId":30512,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"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","_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-18T09:55:23.214895Z","iopub.execute_input":"2023-06-18T09:55:23.215564Z","iopub.status.idle":"2023-06-18T09:55:23.231353Z","shell.execute_reply.started":"2023-06-18T09:55:23.215527Z","shell.execute_reply":"2023-06-18T09:55:23.230459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import Libaries","metadata":{}},{"cell_type":"code","source":"!pip install segmentation-models-pytorch\nimport segmentation_models_pytorch as smp\n","metadata":{"scrolled":true,"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-06-18T09:55:23.233964Z","iopub.execute_input":"2023-06-18T09:55:23.235008Z","iopub.status.idle":"2023-06-18T09:55:52.029105Z","shell.execute_reply.started":"2023-06-18T09:55:23.234983Z","shell.execute_reply":"2023-06-18T09:55:52.028007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport random\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision\nimport torchvision.utils as vutils\n\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau, CosineAnnealingLR, StepLR, MultiStepLR, CyclicLR\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms as T, datasets as dset\nfrom sklearn.model_selection import train_test_split\nfrom matplotlib import pyplot as plt\n\nfrom zipfile import ZipFile\nfrom tqdm import tqdm\nfrom glob import glob\nfrom PIL import Image","metadata":{"execution":{"iopub.status.busy":"2023-06-18T09:55:52.030627Z","iopub.execute_input":"2023-06-18T09:55:52.031081Z","iopub.status.idle":"2023-06-18T09:55:52.859594Z","shell.execute_reply.started":"2023-06-18T09:55:52.031036Z","shell.execute_reply":"2023-06-18T09:55:52.858545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"batch_size = 4\nn_iters = 10000\nepochs = 10\nlearning_rate = 0.0002\nn_workers = 2\n\n\nwidth = 256\nheight = 256\nchannels = 3\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nseed = 44\nrandom.seed(seed)\ntorch.manual_seed(seed)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T09:55:52.864585Z","iopub.execute_input":"2023-06-18T09:55:52.868658Z","iopub.status.idle":"2023-06-18T09:55:52.973280Z","shell.execute_reply.started":"2023-06-18T09:55:52.868616Z","shell.execute_reply":"2023-06-18T09:55:52.972037Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load Data","metadata":{}},{"cell_type":"code","source":"# extract train data\nwith ZipFile('../input/carvana-image-masking-challenge/train.zip', 'r') as zf:\n    zf.extractall('../working')\n     \nwith ZipFile('../input/carvana-image-masking-challenge/train_masks.zip', 'r') as zf:\n    zf.extractall('../working')","metadata":{"execution":{"iopub.status.busy":"2023-06-18T09:55:52.979197Z","iopub.execute_input":"2023-06-18T09:55:52.981499Z","iopub.status.idle":"2023-06-18T09:56:01.922109Z","shell.execute_reply.started":"2023-06-18T09:55:52.981463Z","shell.execute_reply":"2023-06-18T09:56:01.921144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" class MyDataset(Dataset):\n    def __init__(self, root_dir: str, train=True, transforms=None):\n        super(MyDataset, self).__init__()\n        self.train = train\n        self.transforms = transforms\n        \n        file_path = root_dir + 'train/*.*'\n        file_mask_path = root_dir + 'train_masks/*.*'\n        \n        self.images = sorted(glob(file_path))\n        self.image_mask = sorted(glob(file_mask_path))\n        \n        # manually split the train/valid data\n        split_ratio = int(len(self.images) * 0.7)\n        if train:\n            self.images = self.images[:split_ratio]\n            self.image_mask = self.image_mask[:split_ratio]\n        else:\n            self.images = self.images[split_ratio:]\n            self.image_mask = self.image_mask[split_ratio:]\n        \n        \n    def __getitem__(self, index: int):\n        image = Image.open(self.images[index]).convert('RGB')\n        image_mask = Image.open(self.image_mask[index]).convert('L')\n\n        \n        if self.transforms:\n            image = self.transforms(image)\n            image_mask = self.transforms(image_mask)\n                        \n        return {'img': image, 'mask': image_mask}\n    \n    def __len__(self):\n        return len(self.images)\n    \n\ntransforms = T.Compose([\n    T.Resize((width, height)),\n    T.ToTensor(),\n#     T.Normalize(mean=[0.485, 0.456, 0.406],\n#                 std=[0.229, 0.224, 0.225]),\n#     T.RandomHorizontalFlip()\n])\ntrain_dataset = MyDataset(root_dir='../working/',\n                                 train=True,\n                                 transforms=transforms)\nval_dataset = MyDataset(root_dir='../working/',\n                                train=False,\n                                transforms=transforms)\n\ntrain_dataset_loader = DataLoader(dataset=train_dataset,\n                                 batch_size=batch_size,\n                                 shuffle=True,\n                                 num_workers=n_workers)\nval_dataset_loader = DataLoader(dataset=val_dataset,\n                                batch_size=batch_size,\n                                shuffle=True,\n                                num_workers=n_workers)\n\n        ","metadata":{"execution":{"iopub.status.busy":"2023-06-18T09:56:01.923447Z","iopub.execute_input":"2023-06-18T09:56:01.923797Z","iopub.status.idle":"2023-06-18T09:56:02.026943Z","shell.execute_reply.started":"2023-06-18T09:56:01.923765Z","shell.execute_reply":"2023-06-18T09:56:02.026093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Sample Data","metadata":{}},{"cell_type":"code","source":"samples = next(iter(train_dataset_loader))\n\nfig, (ax1, ax2) = plt.subplots(nrows=2, ncols=1, figsize=(12, 4))\nfig.tight_layout()\n\n\nax1.axis('off')\nax1.set_title('input image')\nax1.imshow(np.transpose(vutils.make_grid(samples['img'], padding=2).numpy(),\n                       (1, 2, 0)))\n\nax2.axis('off')\nax2.set_title('input mask')\nax2.imshow(np.transpose(vutils.make_grid(samples['mask'], padding=2).numpy(),\n                       (1, 2, 0)), cmap='gray')\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-06-18T09:56:02.028173Z","iopub.execute_input":"2023-06-18T09:56:02.028524Z","iopub.status.idle":"2023-06-18T09:56:03.357441Z","shell.execute_reply.started":"2023-06-18T09:56:02.028492Z","shell.execute_reply":"2023-06-18T09:56:03.356417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#unet part\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\nclass DoubleConv(nn.Module):\n    \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),\n            nn.BatchNorm2d(mid_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1),\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\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        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":"2023-06-18T09:56:03.358847Z","iopub.execute_input":"2023-06-18T09:56:03.359204Z","iopub.status.idle":"2023-06-18T09:56:03.380347Z","shell.execute_reply.started":"2023-06-18T09:56:03.359162Z","shell.execute_reply":"2023-06-18T09:56:03.379242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#unet_att\nimport torch.nn as nn\n\nclass MultiConv(nn.Module):\n    def __init__(self, in_ch, out_ch, attn=True):\n        super(MultiConv, self).__init__()\n\n        self.fuse_attn = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.PReLU(),\n            nn.Conv2d(out_ch, out_ch, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.PReLU(),\n            nn.Conv2d(out_ch, out_ch, kernel_size=1),\n            nn.BatchNorm2d(out_ch),\n            nn.Softmax2d() if attn else nn.PReLU()\n        )\n\n    def forward(self, x):\n        return self.fuse_attn(x)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T09:56:03.382068Z","iopub.execute_input":"2023-06-18T09:56:03.382482Z","iopub.status.idle":"2023-06-18T09:56:03.398247Z","shell.execute_reply.started":"2023-06-18T09:56:03.382450Z","shell.execute_reply":"2023-06-18T09:56:03.397437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#unet_transformer\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn\n\n\nclass PositionEmbeddingLearned(nn.Module):\n    \n    def __init__(self, num_pos_feats=256, len_embedding=32):\n        super().__init__()\n        self.row_embed = nn.Embedding(len_embedding, num_pos_feats)\n        self.col_embed = nn.Embedding(len_embedding, num_pos_feats)\n        self.reset_parameters()\n\n    def reset_parameters(self):\n        nn.init.uniform_(self.row_embed.weight)\n        nn.init.uniform_(self.col_embed.weight)\n\n    def forward(self, tensor_list):\n        x = tensor_list\n        h, w = x.shape[-2:]\n        i = torch.arange(w, device=x.device)\n        j = torch.arange(h, device=x.device)\n\n        x_emb = self.col_embed(i)\n        y_emb = self.row_embed(j)\n\n        pos = torch.cat([\n            x_emb.unsqueeze(0).repeat(h, 1, 1),\n            y_emb.unsqueeze(1).repeat(1, w, 1),\n        ], dim=-1).permute(2, 0, 1).unsqueeze(0).repeat(x.shape[0], 1, 1, 1)\n\n        return pos\n\nclass DoubleConv(nn.Module):\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),\n            nn.BatchNorm2d(mid_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.double_conv(x)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T09:56:03.399774Z","iopub.execute_input":"2023-06-18T09:56:03.400112Z","iopub.status.idle":"2023-06-18T09:56:03.425358Z","shell.execute_reply.started":"2023-06-18T09:56:03.400073Z","shell.execute_reply":"2023-06-18T09:56:03.424336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UNet_Attention_Transformer_Multiscale(nn.Module):\n    def __init__(self, n_channels, n_classes, bilinear=True):\n        super(UNet_Attention_Transformer_Multiscale, 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(1024, 256 // factor, bilinear)\n        self.up3 = Up(512, 128 // factor, bilinear)\n        self.up4 = Up(256, 64, bilinear)\n        self.outc = OutConv(128, n_classes)\n\n        self.pos = PositionEmbeddingLearned(512 // factor)\n\n        self.fuse1 = MultiConv(768, 256)\n        self.fuse2 = MultiConv(384, 128)\n        self.fuse3 = MultiConv(192, 64)\n        self.fuse4 = MultiConv(128, 64)\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\n        x6 = self.up1(x5, x4)\n        x5_scale = F.interpolate(x5, size=x6.shape[2:], mode='bilinear', align_corners=True)\n        x6_cat = torch.cat((x5_scale, x6), 1)\n\n        x7 = self.up2(x6_cat, x3)\n        x6_scale = F.interpolate(x6, size=x7.shape[2:], mode='bilinear', align_corners=True)\n        x7_cat = torch.cat((x6_scale, x7), 1)\n\n        x8 = self.up3(x7_cat, x2)\n        x7_scale = F.interpolate(x7, size=x8.shape[2:], mode='bilinear', align_corners=True)\n        x8_cat = torch.cat((x7_scale, x8), 1)\n\n        x9 = self.up4(x8_cat, x1)\n        x8_scale = F.interpolate(x8, size=x9.shape[2:], mode='bilinear', align_corners=True)\n        x9 = torch.cat((x8_scale, x9), 1)\n\n        logits = self.outc(x9)\n        return logits\n","metadata":{"execution":{"iopub.status.busy":"2023-06-18T10:15:45.061747Z","iopub.execute_input":"2023-06-18T10:15:45.062152Z","iopub.status.idle":"2023-06-18T10:15:45.077064Z","shell.execute_reply.started":"2023-06-18T10:15:45.062117Z","shell.execute_reply":"2023-06-18T10:15:45.076017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"It is a binary segmentation problem. We use *BCEWithLogitsLoss* loss function. You can change the loss function to *BCELoss* by adding sigmoid layer to the output.","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        self.double_conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\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    def __init__(self, in_channels, out_channels):\n        super(Down, self).__init__()\n        self.max_pool_conv = nn.Sequential(\n            nn.MaxPool2d(2),\n            DoubleConv(in_channels, out_channels)\n        )\n\n    def forward(self, x):\n        return self.max_pool_conv(x)\n\n\nclass Up(nn.Module):\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super(Up, self).__init__()\n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n            self.conv = DoubleConv(in_channels + in_channels // 2, out_channels)\n        else:\n            self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, kernel_size=2, stride=2)\n            self.conv = DoubleConv(in_channels + in_channels // 2, out_channels)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n        x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,\n                        diffY // 2, diffY - diffY // 2])\n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n\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\nclass TransUNet(nn.Module):\n    def __init__(self, n_channels, n_classes, bilinear=True):\n        super(TransUNet, 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(1024, 256 // factor, bilinear)\n        self.up3 = Up(512, 128 // factor, bilinear)\n        self.up4 = Up(256, 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\n\n\n# Example usage:\nn_channels = 3\nn_classes = 2\nmodel = TransUNet(n_channels, n_classes)\ninput_tensor = torch.randn(1, n_channels, 256, 256)\noutput = model(input_tensor)\nprint(\"Output shape:\", output.shape)\n","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:05:01.037681Z","iopub.execute_input":"2023-06-18T11:05:01.038238Z","iopub.status.idle":"2023-06-18T11:05:01.849859Z","shell.execute_reply.started":"2023-06-18T11:05:01.038203Z","shell.execute_reply":"2023-06-18T11:05:01.848367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Metrics","metadata":{}},{"cell_type":"code","source":"def dice_score(pred: torch.Tensor, mask: torch.Tensor):\n    dice = (2 * (pred * mask).sum()) / (pred + mask).sum()\n    return np.mean(dice.cpu().numpy())\n\ndef iou_score(pred: torch.Tensor, mask: torch.Tensor):\n    intersection = torch.logical_and(pred, mask).sum()\n    union = torch.logical_or(pred, mask).sum()\n    iou = intersection.float() / union.float()\n    return iou\n\ndef pixel_accuracy(pred: torch.Tensor, mask: torch.Tensor):\n    correct = torch.eq(pred, mask).int()\n    return float(correct.sum()) / float(correct.numel())","metadata":{"execution":{"iopub.status.busy":"2023-06-18T10:15:47.369340Z","iopub.execute_input":"2023-06-18T10:15:47.370354Z","iopub.status.idle":"2023-06-18T10:15:47.378454Z","shell.execute_reply.started":"2023-06-18T10:15:47.370308Z","shell.execute_reply":"2023-06-18T10:15:47.377111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"def plot_pred_img(samples, pred):\n    fig, (ax1, ax2, ax3) = plt.subplots(nrows=3, ncols=1, figsize=(12, 6))\n    fig.tight_layout()\n\n\n    ax1.axis('off')\n    ax1.set_title('input image')\n    ax1.imshow(np.transpose(vutils.make_grid(samples['img'], padding=2).numpy(),\n                           (1, 2, 0)))\n\n    ax2.axis('off')\n    ax2.set_title('input mask')\n    ax2.imshow(np.transpose(vutils.make_grid(samples['mask'], padding=2).numpy(),\n                           (1, 2, 0)), cmap='gray')\n    \n    ax3.axis('off')\n    ax3.set_title('predicted mask')\n    ax3.imshow(np.transpose(vutils.make_grid(pred, padding=2).cpu().numpy(),\n                           (1, 2, 0)), cmap='gray')\n\n    plt.show()\n    \n    \ndef plot_train_progress(model):\n#     model.eval()\n\n#     with torch.no_grad():\n    samples = next(iter(val_dataset_loader))\n    val_img = samples['img'].to(device)\n    val_mask = samples['mask'].to(device)\n\n    pred = model(val_img)\n\n\n    plot_pred_img(samples, pred.detach())","metadata":{"execution":{"iopub.status.busy":"2023-06-18T10:15:48.026843Z","iopub.execute_input":"2023-06-18T10:15:48.027208Z","iopub.status.idle":"2023-06-18T10:15:48.036808Z","shell.execute_reply.started":"2023-06-18T10:15:48.027177Z","shell.execute_reply":"2023-06-18T10:15:48.035781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, optimizer, criteration, scheduler=None):\n    train_losses = []\n    val_lossess = []\n    lr_rates = []\n    \n    # calculate train epochs\n    epochs = int(n_iters / (len(train_dataset) / batch_size))\n\n\n    for epoch in range(epochs):\n        model.train()        \n        train_total_loss = 0\n        train_iterations = 0\n        \n        for idx, data in enumerate(tqdm(train_dataset_loader)):\n            train_iterations += 1\n            train_img = data['img'].to(device)\n            train_mask = data['mask'].to(device)\n            \n            optimizer.zero_grad()\n            # speed up the training\n            with torch.autocast(device_type='cuda'):\n                train_output_mask = model(train_img)\n                train_loss = criterion(train_output_mask, train_mask)\n                train_total_loss += train_loss.item()\n\n            train_loss.backward()\n            optimizer.step()\n\n\n\n        train_epoch_loss = train_total_loss / train_iterations\n        train_losses.append(train_epoch_loss)\n        \n        # evaluate mode\n        model.eval()\n        with torch.no_grad():\n            val_total_loss = 0\n            val_iterations = 0\n            scores = 0\n            score_iou = 0\n            score_acc = 0\n\n            for vidx, val_data in enumerate(tqdm(val_dataset_loader)):\n                val_iterations += 1\n                val_img = val_data['img'].to(device)\n                val_mask = val_data['mask'].to(device)\n\n                with torch.autocast(device_type='cuda'):\n                    pred = model(val_img)\n                    val_loss = criterion(pred, val_mask)\n                    val_total_loss += val_loss.item()\n                    scores += dice_score(pred, val_mask)\n                    score_iou += iou_score(pred, val_mask)\n                    \n\n\n            val_epoch_loss = val_total_loss / val_iterations\n            dice_coef_scroe = scores / val_iterations\n            score_iou_coef = score_iou/val_iterations\n\n            val_lossess.append(val_epoch_loss)           \n\n            plot_train_progress(model)\n            print('epochs - {}/{} [{}/{}], dice score: {}, train loss: {}, val loss: {}, iou score:{}'.format(\n                epoch+1, epochs,\n                idx+1, len(train_dataset_loader),\n                dice_coef_scroe, train_epoch_loss, val_epoch_loss, score_iou_coef\n            )) \n            \n        lr_rates.append(optimizer.param_groups[0]['lr'])\n        if scheduler:\n            scheduler.step() # decay learning rate\n            print('LR rate:', scheduler.get_last_lr())\n            \n    return {\n        'lr': lr_rates,\n        'train_loss': train_losses,\n        'valid_loss': val_lossess\n    }","metadata":{"execution":{"iopub.status.busy":"2023-06-18T10:15:48.527276Z","iopub.execute_input":"2023-06-18T10:15:48.528691Z","iopub.status.idle":"2023-06-18T10:15:48.542512Z","shell.execute_reply.started":"2023-06-18T10:15:48.528644Z","shell.execute_reply":"2023-06-18T10:15:48.541639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = UNet_Attention_Transformer_Multiscale(n_channels=3, n_classes=1).to(device)\ncriterion = nn.BCEWithLogitsLoss()\n# criterion = smp.losses.DiceLoss(mode='binary')\noptimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)\n\nhistory = train(model, optimizer, criterion)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T10:15:50.078877Z","iopub.execute_input":"2023-06-18T10:15:50.079276Z","iopub.status.idle":"2023-06-18T10:59:27.448237Z","shell.execute_reply.started":"2023-06-18T10:15:50.079245Z","shell.execute_reply":"2023-06-18T10:59:27.447260Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#accuracy\ndef pixel_accuracy(pred: torch.Tensor, mask: torch.Tensor):\n    correct = torch.eq(pred, val_mask).int()\n    return float(correct.sum()) / float(correct.numel())\n\ndef train_A(model, optimizer, criteration, scheduler=None):\n    train_losses = []\n    val_lossess = []\n    lr_rates = []\n    \n    # calculate train epochs\n    epochs = int(n_iters / (len(train_dataset) / batch_size))\n\n\n    for epoch in range(epochs):\n        model.train()        \n        train_total_loss = 0\n        train_iterations = 0\n        \n        for idx, img_mask in enumerate(tqdm(train_dataset_loader)):\n            train_iterations += 1\n            train_img = samples['img'].float().to(device)\n            train_mask = samples['mask'].float().to(device)\n            \n            optimizer.zero_grad()\n            # speed up the training\n            with torch.autocast(device_type='cuda'):\n                train_output_mask = model(train_img)\n                train_loss = criterion(train_output_mask, train_mask)\n                train_total_loss += train_loss.item()\n\n            train_loss.backward()\n            optimizer.step()\n\n\n\n        train_epoch_loss = train_total_loss / train_iterations\n        train_losses.append(train_epoch_loss)\n        \n        # evaluate mode\n        model.eval()\n        with torch.no_grad():\n            val_total_loss = 0\n            val_iterations = 0\n            scores_acc = 0\n\n            for vidx, img_mask in enumerate(tqdm(val_dataset_loader)):\n                val_iterations += 1\n                val_img = samples['img'].float().to(device)\n                val_mask = samples['mask'].float().to(device)\n                def pixel_accuracy(pred: torch.Tensor, mask: torch.Tensor):\n                    correct = torch.eq(pred, val_mask).int()\n                    return float(correct.sum()) / float(correct.numel())\n\n                with torch.autocast(device_type='cuda'):\n                    pred = model(val_img)\n                    scores_acc += pixel_accuracy(pred, val_mask)\n\n\n            scores_acc_coeff = scores_acc/ val_iterations\n                      \n\n            #plot_train_progress(model)\n            print('epochs - {}/{} [{}/{}], Accuracy : {}'.format(\n                epoch+1, epochs,\n                idx+1, len(train_dataset_loader),\n                scores_acc_coeff\n            )) \n            \n        lr_rates.append(optimizer.param_groups[0]['lr'])\n        if scheduler:\n            scheduler.step() # decay learning rate\n            print('LR rate:', scheduler.get_last_lr())\n            \n    return {\n        'lr': lr_rates,\n        'train_loss': train_losses,\n        'valid_loss': val_lossess\n    }","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:41.737573Z","iopub.execute_input":"2023-06-18T11:00:41.738637Z","iopub.status.idle":"2023-06-18T11:00:41.755921Z","shell.execute_reply.started":"2023-06-18T11:00:41.738595Z","shell.execute_reply":"2023-06-18T11:00:41.754984Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Mymodel().to(device)\ncriterion = nn.BCEWithLogitsLoss()\n# criterion = smp.losses.DiceLoss(mode='binary')\noptimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)\n\nhistory = train_A(model, optimizer, criterion)","metadata":{"execution":{"iopub.status.busy":"2023-06-18T11:00:44.030678Z","iopub.execute_input":"2023-06-18T11:00:44.031039Z","iopub.status.idle":"2023-06-18T11:00:44.087931Z","shell.execute_reply.started":"2023-06-18T11:00:44.031009Z","shell.execute_reply":"2023-06-18T11:00:44.086375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Result","metadata":{}},{"cell_type":"code","source":"plt.plot(history['train_loss'], label='constant LR train loss')\nplt.plot(history['valid_loss'], label='constant LR val loss')\nplt.xlabel('epoch')\nplt.ylabel('loss')\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-06-18T10:15:07.247202Z","iopub.status.idle":"2023-06-18T10:15:07.247949Z","shell.execute_reply.started":"2023-06-18T10:15:07.247692Z","shell.execute_reply":"2023-06-18T10:15:07.247718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}