{"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)\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\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":"2021-12-08T17:50:10.696189Z","iopub.execute_input":"2021-12-08T17:50:10.696502Z","iopub.status.idle":"2021-12-08T17:50:10.722231Z","shell.execute_reply.started":"2021-12-08T17:50:10.696422Z","shell.execute_reply":"2021-12-08T17:50:10.72158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Importing the required libraries**","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch import nn\nfrom torch import optim\nfrom torch.utils.tensorboard import SummaryWriter\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\n# import torchvision.transforms.functional as TF\n\nimport random\nimport os, shutil\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm.auto import tqdm\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\n\nimport os\nfrom os.path import join\nimport matplotlib.pyplot as plt\nplt.rcParams.update({'font.size': 18})\nimport cv2\n\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset, sampler\nfrom albumentations import (HorizontalFlip, VerticalFlip, ShiftScaleRotate, Normalize, Resize, Compose, GaussNoise)\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, f1_score","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:10.724055Z","iopub.execute_input":"2021-12-08T17:50:10.724337Z","iopub.status.idle":"2021-12-08T17:50:14.925293Z","shell.execute_reply.started":"2021-12-08T17:50:10.724302Z","shell.execute_reply":"2021-12-08T17:50:14.924538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_train = pd.read_csv('../input/sartorius-cell-instance-segmentation/train.csv')","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:14.926684Z","iopub.execute_input":"2021-12-08T17:50:14.926921Z","iopub.status.idle":"2021-12-08T17:50:15.471518Z","shell.execute_reply.started":"2021-12-08T17:50:14.926887Z","shell.execute_reply":"2021-12-08T17:50:15.470794Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_train.head()","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.472756Z","iopub.execute_input":"2021-12-08T17:50:15.473015Z","iopub.status.idle":"2021-12-08T17:50:15.494132Z","shell.execute_reply.started":"2021-12-08T17:50:15.472983Z","shell.execute_reply":"2021-12-08T17:50:15.493362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Visualization of the images**","metadata":{}},{"cell_type":"code","source":"# def imshow(num_to_show=9):\n    \n#     plt.figure(figsize=(20,20))\n    \n#     for i in range(num_to_show):\n#         plt.subplot(3, 3, i+1)\n#         plt.grid(False)\n#         plt.xticks([])\n#         plt.yticks([])\n        \n#         img = mpimg.imread(f'../input/sartorius-cell-instance-segmentation/train/{data_train.iloc[i,0]}.png')\n#         plt.imshow(img, cmap='plasma')\n\n# imshow()","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.496622Z","iopub.execute_input":"2021-12-08T17:50:15.496878Z","iopub.status.idle":"2021-12-08T17:50:15.500293Z","shell.execute_reply.started":"2021-12-08T17:50:15.496844Z","shell.execute_reply":"2021-12-08T17:50:15.499464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_PATH = '../input/sartorius-cell-instance-segmentation'\nSAMPLE_SUBMISSION = join(DATA_PATH,'train')\nTRAIN_CSV = join(DATA_PATH,'train.csv')\nTRAIN_PATH = join(DATA_PATH,'train')\nTEST_PATH = join(DATA_PATH,'test')\n\ndf_train = pd.read_csv(TRAIN_CSV)\nprint(f'Training Set Shape: {df_train.shape} - {df_train[\"id\"].nunique()} \\\nImages - Memory Usage: {df_train.memory_usage().sum() / 1024 ** 2:.2f} MB')","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.501809Z","iopub.execute_input":"2021-12-08T17:50:15.502147Z","iopub.status.idle":"2021-12-08T17:50:15.799868Z","shell.execute_reply.started":"2021-12-08T17:50:15.50211Z","shell.execute_reply":"2021-12-08T17:50:15.799149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_decode(mask_rle, shape, color=1):\n    '''\n    mask_rle: run-length as string formated (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.float32)\n    for lo, hi in zip(starts, ends):\n        img[lo : hi] = color\n    return img.reshape(shape)\n\ndef build_masks(df_train, image_id, input_shape):\n    height, width = input_shape\n    labels = df_train[df_train[\"id\"] == image_id][\"annotation\"].tolist()\n    mask = np.zeros((height, width))\n    for label in labels:\n        mask += rle_decode(label, shape=(height, width))\n    mask = mask.clip(0, 1)\n    return np.array(mask)","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.801447Z","iopub.execute_input":"2021-12-08T17:50:15.801863Z","iopub.status.idle":"2021-12-08T17:50:15.81083Z","shell.execute_reply.started":"2021-12-08T17:50:15.801824Z","shell.execute_reply":"2021-12-08T17:50:15.810055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Dataset class**","metadata":{}},{"cell_type":"code","source":"class CellDataset(Dataset):\n    def __init__(self, df: pd.core.frame.DataFrame, train:bool):\n        self.IMAGE_RESIZE = (224, 224)\n        self.RESNET_MEAN = (0.485, 0.456, 0.406)\n        self.RESNET_STD = (0.229, 0.224, 0.225)\n        self.df = df\n        self.base_path = TRAIN_PATH\n        self.gb = self.df.groupby('id')\n        self.transforms = Compose([Resize( self.IMAGE_RESIZE[0],  self.IMAGE_RESIZE[1]), \n                                   Normalize(mean=self.RESNET_MEAN, std= self.RESNET_STD, p=1), \n                                   HorizontalFlip(p=0.5),\n                                   VerticalFlip(p=0.5)])\n        \n        # Split train and val set\n        all_image_ids = np.array(df_train.id.unique())\n        np.random.seed(42)\n        iperm = np.random.permutation(len(all_image_ids))\n        num_train_samples = int(len(all_image_ids) * 0.9)\n\n        if train:\n            self.image_ids = all_image_ids[iperm[:num_train_samples]]\n        else:\n             self.image_ids = all_image_ids[iperm[num_train_samples:]]\n\n    def __getitem__(self, idx: int) -> dict:\n\n        image_id = self.image_ids[idx]\n        df = self.gb.get_group(image_id)\n\n        # Read image\n        image_path = os.path.join(self.base_path, image_id + \".png\")\n        image = cv2.imread(image_path)\n\n        # Create the mask\n        mask = build_masks(df_train, image_id, input_shape=(520, 704))\n        mask = (mask >= 1).astype('float32')\n        augmented = self.transforms(image=image, mask=mask)\n        image = augmented['image']\n        mask = augmented['mask']\n        # print(np.moveaxis(image,0,2).shape)\n        return np.moveaxis(np.array(image),2,0), mask.reshape((1, self.IMAGE_RESIZE[0], self.IMAGE_RESIZE[1]))\n\n\n    def __len__(self):\n        return len(self.image_ids)","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.812011Z","iopub.execute_input":"2021-12-08T17:50:15.812391Z","iopub.status.idle":"2021-12-08T17:50:15.826709Z","shell.execute_reply.started":"2021-12-08T17:50:15.812354Z","shell.execute_reply":"2021-12-08T17:50:15.826007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Data Loaders**","metadata":{}},{"cell_type":"code","source":"ds_train = CellDataset(df_train, train=True)\ndl_train = DataLoader(ds_train, batch_size=16, num_workers=2, pin_memory=True, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.827939Z","iopub.execute_input":"2021-12-08T17:50:15.828593Z","iopub.status.idle":"2021-12-08T17:50:15.845428Z","shell.execute_reply.started":"2021-12-08T17:50:15.828557Z","shell.execute_reply":"2021-12-08T17:50:15.84481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_test = CellDataset(df_train, train=False)\ndl_test = DataLoader(ds_test, batch_size=4, num_workers=2, pin_memory=True, shuffle=False)\n","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.846656Z","iopub.execute_input":"2021-12-08T17:50:15.847465Z","iopub.status.idle":"2021-12-08T17:50:15.858818Z","shell.execute_reply.started":"2021-12-08T17:50:15.847428Z","shell.execute_reply":"2021-12-08T17:50:15.858132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Visualization of the images and masks**","metadata":{}},{"cell_type":"code","source":"# # plot simages and mask from dataloader\n# batch = next(iter(dl_train))\n# images, masks = batch\n# print(f\"image shape: {images.shape},\\nmask shape:{masks.shape},\\nbatch len: {len(batch)}\")\n\n# plt.figure(figsize=(20, 20))\n        \n# plt.subplot(1, 3, 1)\n# plt.xticks([])\n# plt.yticks([])\n# plt.imshow(images[1][1])\n# plt.title('Original image')\n\n# plt.subplot( 1, 3, 2)\n# plt.xticks([])\n# plt.yticks([])\n# print(masks[1][0])\n# plt.imshow(masks[1][0])\n# plt.title('Mask')\n\n# plt.subplot( 1, 3, 3)\n# plt.xticks([])\n# plt.yticks([])\n# plt.imshow(images[1][1])\n# plt.imshow(masks[1][0],alpha=0.2)\n# plt.title('Both')\n# plt.tight_layout()\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.860354Z","iopub.execute_input":"2021-12-08T17:50:15.860576Z","iopub.status.idle":"2021-12-08T17:50:15.868686Z","shell.execute_reply.started":"2021-12-08T17:50:15.860554Z","shell.execute_reply":"2021-12-08T17:50:15.867946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Attention Unet","metadata":{}},{"cell_type":"code","source":"class conv_block(nn.Module):\n    \"\"\"\n    Convolution Block \n    \"\"\"\n    def __init__(self, in_ch, out_ch):\n        super(conv_block, self).__init__()\n        \n        self.conv = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=1, padding=1, bias=True),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, kernel_size=3, stride=1, padding=1, bias=True),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True))\n\n    def forward(self, x):\n\n        x = self.conv(x)\n        return x\n\n\nclass up_conv(nn.Module):\n    \"\"\"\n    Up Convolution Block\n    \"\"\"\n    def __init__(self, in_ch, out_ch):\n        super(up_conv, self).__init__()\n        self.up = nn.Sequential(\n            nn.Upsample(scale_factor=2),\n            nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=1, padding=1, bias=True),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        x = self.up(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.87026Z","iopub.execute_input":"2021-12-08T17:50:15.870452Z","iopub.status.idle":"2021-12-08T17:50:15.883335Z","shell.execute_reply.started":"2021-12-08T17:50:15.870411Z","shell.execute_reply":"2021-12-08T17:50:15.882447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Attention_block(nn.Module):\n    \"\"\"\n    Attention Block\n    \"\"\"\n\n    def __init__(self, F_g, F_l, F_int):\n        super(Attention_block, self).__init__()\n\n        self.W_g = nn.Sequential(\n            nn.Conv2d(F_l, F_int, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(F_int)\n        )\n\n        self.W_x = nn.Sequential(\n            nn.Conv2d(F_g, F_int, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(F_int)\n        )\n\n        self.psi = nn.Sequential(\n            nn.Conv2d(F_int, 1, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(1),\n            nn.Sigmoid()\n        )\n\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, g, x):\n        g1 = self.W_g(g)\n        x1 = self.W_x(x)\n        psi = self.relu(g1 + x1)\n        psi = self.psi(psi)\n        out = x * psi\n        return out\n\n\nclass AttU_Net(nn.Module):\n    \"\"\"\n    Attention Unet implementation\n    Paper: https://arxiv.org/abs/1804.03999\n    \"\"\"\n    def __init__(self, img_ch=3, output_ch=1):\n        super(AttU_Net, self).__init__()\n\n        n1 = 64\n        filters = [n1, n1 * 2, n1 * 4, n1 * 8, n1 * 16]\n\n        self.Maxpool1 = nn.MaxPool2d(kernel_size=2, stride=2)\n        self.Maxpool2 = nn.MaxPool2d(kernel_size=2, stride=2)\n        self.Maxpool3 = nn.MaxPool2d(kernel_size=2, stride=2)\n        self.Maxpool4 = nn.MaxPool2d(kernel_size=2, stride=2)\n\n        self.Conv1 = conv_block(img_ch, filters[0])\n        self.Conv2 = conv_block(filters[0], filters[1])\n        self.Conv3 = conv_block(filters[1], filters[2])\n        self.Conv4 = conv_block(filters[2], filters[3])\n        self.Conv5 = conv_block(filters[3], filters[4])\n\n        self.Up5 = up_conv(filters[4], filters[3])\n        self.Att5 = Attention_block(F_g=filters[3], F_l=filters[3], F_int=filters[2])\n        self.Up_conv5 = conv_block(filters[4], filters[3])\n\n        self.Up4 = up_conv(filters[3], filters[2])\n        self.Att4 = Attention_block(F_g=filters[2], F_l=filters[2], F_int=filters[1])\n        self.Up_conv4 = conv_block(filters[3], filters[2])\n\n        self.Up3 = up_conv(filters[2], filters[1])\n        self.Att3 = Attention_block(F_g=filters[1], F_l=filters[1], F_int=filters[0])\n        self.Up_conv3 = conv_block(filters[2], filters[1])\n\n        self.Up2 = up_conv(filters[1], filters[0])\n        self.Att2 = Attention_block(F_g=filters[0], F_l=filters[0], F_int=32)\n        self.Up_conv2 = conv_block(filters[1], filters[0])\n\n        self.Conv = nn.Conv2d(filters[0], output_ch, kernel_size=1, stride=1, padding=0)\n\n        #self.active = torch.nn.Sigmoid()\n\n\n    def forward(self, x):\n\n        e1 = self.Conv1(x)\n\n        e2 = self.Maxpool1(e1)\n        e2 = self.Conv2(e2)\n\n        e3 = self.Maxpool2(e2)\n        e3 = self.Conv3(e3)\n\n        e4 = self.Maxpool3(e3)\n        e4 = self.Conv4(e4)\n\n        e5 = self.Maxpool4(e4)\n        e5 = self.Conv5(e5)\n\n        #print(x5.shape)\n        d5 = self.Up5(e5)\n        #print(d5.shape)\n        x4 = self.Att5(g=d5, x=e4)\n        d5 = torch.cat((x4, d5), dim=1)\n        d5 = self.Up_conv5(d5)\n\n        d4 = self.Up4(d5)\n        x3 = self.Att4(g=d4, x=e3)\n        d4 = torch.cat((x3, d4), dim=1)\n        d4 = self.Up_conv4(d4)\n\n        d3 = self.Up3(d4)\n        x2 = self.Att3(g=d3, x=e2)\n        d3 = torch.cat((x2, d3), dim=1)\n        d3 = self.Up_conv3(d3)\n\n        d2 = self.Up2(d3)\n        x1 = self.Att2(g=d2, x=e1)\n        d2 = torch.cat((x1, d2), dim=1)\n        d2 = self.Up_conv2(d2)\n\n        out = self.Conv(d2)\n        print(out.shape)\n\n      #  out = self.active(out)\n\n        return out","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.884811Z","iopub.execute_input":"2021-12-08T17:50:15.885306Z","iopub.status.idle":"2021-12-08T17:50:15.912251Z","shell.execute_reply.started":"2021-12-08T17:50:15.885264Z","shell.execute_reply":"2021-12-08T17:50:15.911595Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Train the model**","metadata":{}},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(DoubleConv, self).__init__()\n        self.conv = nn.Sequential( \n            nn.Conv2d(in_ch, out_ch, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_ch, out_ch, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(out_ch),\n            nn.ReLU(inplace=True)\n         )\n    def forward(self, x):\n        x = self.conv(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.915928Z","iopub.execute_input":"2021-12-08T17:50:15.916272Z","iopub.status.idle":"2021-12-08T17:50:15.92585Z","shell.execute_reply.started":"2021-12-08T17:50:15.916237Z","shell.execute_reply":"2021-12-08T17:50:15.925127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class InConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(InConv, self).__init__()\n        self.conv = DoubleConv(in_ch, out_ch)\n    def forward(self, x):\n        x = self.conv(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.92732Z","iopub.execute_input":"2021-12-08T17:50:15.92763Z","iopub.status.idle":"2021-12-08T17:50:15.938195Z","shell.execute_reply.started":"2021-12-08T17:50:15.927593Z","shell.execute_reply":"2021-12-08T17:50:15.937467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Down(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(Down, self).__init__()\n        self.mpconv = nn.Sequential( \n            nn.MaxPool2d(2,2),\n            DoubleConv(in_ch, out_ch)\n         )\n    def forward(self, x):\n        x = self.mpconv(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.939513Z","iopub.execute_input":"2021-12-08T17:50:15.939861Z","iopub.status.idle":"2021-12-08T17:50:15.947651Z","shell.execute_reply.started":"2021-12-08T17:50:15.939827Z","shell.execute_reply":"2021-12-08T17:50:15.946909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Up(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(Up, self).__init__()\n        self.up = nn.ConvTranspose2d(in_ch // 2, in_ch // 2, kernel_size=2, stride=2)\n        self.conv = DoubleConv(in_ch, out_ch)\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        x = torch.cat([x2, x1], dim=1)\n        x = self.conv(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.949115Z","iopub.execute_input":"2021-12-08T17:50:15.949493Z","iopub.status.idle":"2021-12-08T17:50:15.958338Z","shell.execute_reply.started":"2021-12-08T17:50:15.949411Z","shell.execute_reply":"2021-12-08T17:50:15.957588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class OutConv(nn.Module):\n    def __init__(self, in_ch, out_ch):\n        super(OutConv, self).__init__()\n        self.conv = nn.Conv2d(in_ch, out_ch, 1)\n        self.sigmoid = nn.Sigmoid()\n    def forward(self, x):\n        x = self.conv(x)\n        x = self.sigmoid(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.959775Z","iopub.execute_input":"2021-12-08T17:50:15.959968Z","iopub.status.idle":"2021-12-08T17:50:15.967721Z","shell.execute_reply.started":"2021-12-08T17:50:15.959946Z","shell.execute_reply":"2021-12-08T17:50:15.96687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, in_channels, num_classes):\n        super(UNet, self).__init__()\n        self.inc = InConv(in_channels, 64)\n        self.down1 = Down(64, 128)\n        self.down2 = Down(128, 256)\n        self.down3 = Down(256, 512)\n        self.down4 = Down(512, 512)\n        self.up1 = Up(1024, 256)\n        self.up2 = Up(512, 128)\n        self.up3 = Up(256, 64)\n        self.up4 = Up(128, 64)\n        self.outc = OutConv(64, num_classes)\n    def forward(self, x):\n        # print(x.shape)\n        x1 = self.inc(x)\n        # print(x1.shape)\n        x2 = self.down1(x1)\n        # print(x2.shape)\n        x3 = self.down2(x2)\n        # print(x3.shape)\n        x4 = self.down3(x3)\n        # print(x4.shape)\n        x5 = self.down4(x4)\n        # print(x5.shape)\n        # print('up')\n        x = self.up1(x5, x4)\n        # print(x.shape)\n        x = self.up2(x, x3)\n        # print(x.shape)\n        x = self.up3(x, x2)\n        # print(x.shape)\n        x = self.up4(x, x1)\n        # print(x.shape)\n        x = self.outc(x)\n        print(x.shape)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.969039Z","iopub.execute_input":"2021-12-08T17:50:15.969495Z","iopub.status.idle":"2021-12-08T17:50:15.98058Z","shell.execute_reply.started":"2021-12-08T17:50:15.969458Z","shell.execute_reply":"2021-12-08T17:50:15.979788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:15.981585Z","iopub.execute_input":"2021-12-08T17:50:15.982215Z","iopub.status.idle":"2021-12-08T17:50:16.04576Z","shell.execute_reply.started":"2021-12-08T17:50:15.982183Z","shell.execute_reply":"2021-12-08T17:50:16.044973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_loop(model, optimizer, criterion, train_loader, device=device):\n    running_loss = 0\n    model.train()\n    pbar = tqdm(train_loader, desc='Iterating over train data')\n    for imgs, masks in pbar:\n        # pass to device\n        imgs = imgs.to(device)\n        masks = masks.to(device)\n        # forward\n        out = model(imgs)\n        loss = criterion(out, masks)\n        running_loss += loss.item()*imgs.shape[0]  # += loss * current batch size\n        # optimize\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n    running_loss /= len(train_loader.sampler)\n    return running_loss","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:16.046962Z","iopub.execute_input":"2021-12-08T17:50:16.047231Z","iopub.status.idle":"2021-12-08T17:50:16.055945Z","shell.execute_reply.started":"2021-12-08T17:50:16.047196Z","shell.execute_reply":"2021-12-08T17:50:16.055135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def eval_loop(model, criterion, eval_loader, device=device):\n    running_loss = 0\n    model.eval()\n    with torch.no_grad():\n        accuracy, f1_scores = [], []\n        pbar = tqdm(eval_loader, desc='Iterating over evaluation data')\n        \n        for imgs, masks in pbar:\n            print(imgs.shape)\n            print(masks.shape)\n            # pass to device\n            li=imgs\n            lm=masks\n            imgs = imgs.to(device)\n            masks = masks.to(device)\n            # forward\n            out = model(imgs)\n            print(out.shape)\n            loss = criterion(out, masks)\n            running_loss += loss.item()*imgs.shape[0]\n            # calculate predictions using output\n            predicted = (out > 0.5).float()\n            print(predicted.shape)\n            predicted = predicted.view(-1).cpu().numpy()\n            labels = masks.view(-1).cpu().numpy()\n            print(predicted.shape)\n            print(labels.shape)\n            accuracy.append(accuracy_score(labels, predicted))\n            f1_scores.append(f1_score(labels, predicted))\n    acc = sum(accuracy)/len(accuracy)\n    f1 = sum(f1_scores)/len(f1_scores)\n    running_loss /= len(eval_loader.sampler)\n    return {\n        'accuracy':acc,\n        'f1_macro':f1, \n        'loss':running_loss,\n        'img': li,\n        'masks': lm,\n        'out':out\n        \n    }","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:16.05781Z","iopub.execute_input":"2021-12-08T17:50:16.057988Z","iopub.status.idle":"2021-12-08T17:50:16.070569Z","shell.execute_reply.started":"2021-12-08T17:50:16.057967Z","shell.execute_reply":"2021-12-08T17:50:16.06915Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(model, optimizer, criterion, train_loader, valid_loader,\n          device=device, \n          num_epochs=5, \n          valid_loss_min=np.inf,\n          logdir='logdir'):\n    \n    tb_writer = SummaryWriter(log_dir=logdir)\n    for e in range(num_epochs):\n        # train for epoch\n        train_loss = train_loop(\n            model, optimizer, criterion, train_loader, device=device)\n        # evaluate on validation set\n        metrics = eval_loop(\n            model, criterion, valid_loader, device=device\n        )\n        # show progress\n        print_string = f'Epoch: {e+1} '\n        print_string+= f'TrainLoss: {train_loss:.5f} '\n        print_string+= f'ValidLoss: {metrics[\"loss\"]:.5f} '\n        print_string+= f'ACC: {metrics[\"accuracy\"]:.5f} '\n        print_string+= f'F1: {metrics[\"f1_macro\"]:.3f}'\n        print(print_string)\n\n        # Tensorboards Logging\n        tb_writer.add_scalar('UNet/Train Loss', train_loss, e)\n        tb_writer.add_scalar('UNet/Valid Loss', metrics[\"loss\"], e)\n        tb_writer.add_scalar('UNet/Accuracy', metrics[\"accuracy\"], e)\n        tb_writer.add_scalar('UNet/F1 Macro', metrics[\"f1_macro\"], e)\n\n        # save the model \n        if metrics[\"loss\"] <= valid_loss_min:\n            torch.save(model.state_dict(), 'UNet.pt')\n            valid_loss_min = metrics[\"loss\"]","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:16.071821Z","iopub.execute_input":"2021-12-08T17:50:16.072328Z","iopub.status.idle":"2021-12-08T17:50:16.081919Z","shell.execute_reply.started":"2021-12-08T17:50:16.07229Z","shell.execute_reply":"2021-12-08T17:50:16.081206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# set_seed(21)\nmodel = AttU_Net(3, 1).to(device)\noptimizer = optim.Adam(model.parameters(), lr=0.01)\ncriterion = nn.BCELoss()\ntrain(model, optimizer, criterion, dl_train, dl_test)","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:51:05.316367Z","iopub.execute_input":"2021-12-08T17:51:05.316989Z","iopub.status.idle":"2021-12-08T17:51:14.367349Z","shell.execute_reply.started":"2021-12-08T17:51:05.31695Z","shell.execute_reply":"2021-12-08T17:51:14.366022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Evaluation**","metadata":{}},{"cell_type":"code","source":"# Load the latest model\nmodel.load_state_dict(torch.load('UNet.pt'))\nmetrics = eval_loop(model, criterion, dl_test)\nprint('accuracy:', metrics['accuracy'])\nprint('f1 macro:', metrics['f1_macro'])\nprint('test loss:', metrics['loss'])","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:38.190943Z","iopub.status.idle":"2021-12-08T17:50:38.192049Z","shell.execute_reply.started":"2021-12-08T17:50:38.191801Z","shell.execute_reply":"2021-12-08T17:50:38.191826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20, 20))\nout=metrics['out'].cpu()     \nplt.subplot(1, 3, 1)\nplt.xticks([])\nplt.yticks([])\nplt.imshow(metrics['img'][0][1])\nplt.title('Original image')\n\nplt.subplot( 1, 3, 2)\nplt.xticks([])\nplt.yticks([])\n\nplt.imshow(out[0][0])\nplt.title('Mask')\n\nplt.subplot( 1, 3, 3)\nplt.xticks([])\nplt.yticks([])\nprint(metrics['img'].shape)\nplt.imshow(metrics['img'][0][1])\nplt.imshow(out[0][0],alpha=0.2)\nplt.title('Both')\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-12-08T17:50:38.193418Z","iopub.status.idle":"2021-12-08T17:50:38.193816Z","shell.execute_reply.started":"2021-12-08T17:50:38.193602Z","shell.execute_reply":"2021-12-08T17:50:38.193623Z"},"trusted":true},"execution_count":null,"outputs":[]}]}