{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"}],"dockerImageVersionId":30627,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Importing the Libraries","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\nfrom torchvision import transforms\nimport tifffile as tiff\nimport cv2\nimport torch.nn as nn\nimport albumentations as A\nimport numpy as np\nimport os\nimport time\nimport torch.nn.functional as F\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-12-19T09:52:48.414019Z","iopub.execute_input":"2023-12-19T09:52:48.414704Z","iopub.status.idle":"2023-12-19T09:52:53.583441Z","shell.execute_reply.started":"2023-12-19T09:52:48.414666Z","shell.execute_reply":"2023-12-19T09:52:53.582410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"! nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:52:53.585175Z","iopub.execute_input":"2023-12-19T09:52:53.586243Z","iopub.status.idle":"2023-12-19T09:52:54.620060Z","shell.execute_reply.started":"2023-12-19T09:52:53.586214Z","shell.execute_reply":"2023-12-19T09:52:54.619083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"base_path = \"/kaggle/input/blood-vessel-segmentation/train\"\ndataset = \"kidney_1_dense\"\n\nimages_path = os.path.join(base_path,dataset,\"images\")\nlabel_path = os.path.join(base_path,dataset,\"labels\")\n\nimage_files= sorted([os.path.join(images_path,f) for f in os.listdir(images_path) if f.endswith('.tif')])\nlabels_files= sorted([os.path.join(label_path,f) for f in os.listdir(label_path) if f.endswith('.tif')])\n","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:52:54.621622Z","iopub.execute_input":"2023-12-19T09:52:54.622000Z","iopub.status.idle":"2023-12-19T09:52:54.784708Z","shell.execute_reply.started":"2023-12-19T09:52:54.621964Z","shell.execute_reply":"2023-12-19T09:52:54.783956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualizing the images","metadata":{}},{"cell_type":"code","source":"def show_images(images,titles=None,cmap='gray'):\n    n =len(images)\n    fig,axes = plt.subplots(1,n,figsize=(20,10))\n    for idx,ax in enumerate(axes):\n        ax.imshow(images[idx],cmap=cmap)\n        if titles:\n            ax.set_title(titles[idx])\n        ax.axis('off')\n    plt.tight_layout()\n    plt.show()\n    \nfirst_image = tiff.imread(image_files[100])\nfirst_label = tiff.imread(labels_files[100])\n\nshow_images([first_image,first_label])\n    ","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:52:54.787269Z","iopub.execute_input":"2023-12-19T09:52:54.787573Z","iopub.status.idle":"2023-12-19T09:52:55.681444Z","shell.execute_reply.started":"2023-12-19T09:52:54.787548Z","shell.execute_reply":"2023-12-19T09:52:55.680468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data preprocessing and augmentation","metadata":{}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self,image_files, mask_files, input_size=(256, 256), augmentation_transforms=None):\n        self.image_files=image_files\n        self.mask_files=mask_files\n        self.input_size=input_size\n        self.augmentation_transforms=augmentation_transforms\n    def __len__(self):\n        return len(self.image_files)\n\n    def __getitem__(self,idx):\n        image_path=self.image_files[idx]\n        mask_path=self.mask_files[idx]\n        \n        image = preprocess_image(image_path)\n        mask = preprocess_mask(mask_path)\n        if self.augmentation_transforms:\n            image,mask=self.augmentation_transforms(image,mask)\n        return image,mask\n    \n        ","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:52:55.682660Z","iopub.execute_input":"2023-12-19T09:52:55.682952Z","iopub.status.idle":"2023-12-19T09:52:55.689644Z","shell.execute_reply.started":"2023-12-19T09:52:55.682927Z","shell.execute_reply":"2023-12-19T09:52:55.688771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_image(path):\n    img = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n    img = np.tile(img[...,None],[1,1,3])\n    img = img.astype('float32')\n    mx = np.max(img)\n    if mx:\n        img/=mx\n    img = np.transpose(img,(2,0,1))\n    img_ten = torch.tensor(img)\n    return img_ten","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:52:55.690922Z","iopub.execute_input":"2023-12-19T09:52:55.691501Z","iopub.status.idle":"2023-12-19T09:52:55.705634Z","shell.execute_reply.started":"2023-12-19T09:52:55.691451Z","shell.execute_reply":"2023-12-19T09:52:55.704679Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_mask(path):\n    \n    msk = cv2.imread(path, cv2.IMREAD_UNCHANGED)\n    msk = msk.astype('float32')\n    msk/=255.0\n    msk_ten = torch.tensor(msk)\n    \n    return msk_ten","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:52:55.706900Z","iopub.execute_input":"2023-12-19T09:52:55.707184Z","iopub.status.idle":"2023-12-19T09:52:55.719159Z","shell.execute_reply.started":"2023-12-19T09:52:55.707152Z","shell.execute_reply":"2023-12-19T09:52:55.718411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def augment_image(image, mask):\n    \n    image_np = image.permute(1, 2, 0).numpy()\n    mask_np = mask.numpy()\n\n    transform = A.Compose([\n        A.Resize(256,256, interpolation=cv2.INTER_NEAREST),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.ShiftScaleRotate(scale_limit=0.5, rotate_limit=0, shift_limit=0.1, p=1, border_mode=0),\n        A.RandomCrop(height=256, width=256, always_apply=True),\n        A.RandomBrightness(p=1),\n        A.OneOf(\n            [\n                A.Blur(blur_limit=3, p=1),\n                A.MotionBlur(blur_limit=3, p=1),\n            ],\n            p=0.9,\n        ),\n    \n    ])\n    augmented = transform(image = image_np,mask = mask_np)\n    augmented_image , augmented_mask = augmented['image'],augmented['mask']\n    \n    augmented_image = torch.tensor(augmented_image, dtype=torch.float32).permute(2, 0, 1)\n    augmented_mask  = torch.tensor(augmented_mask,dtype=torch.float32)\n    \n    return augmented_image,augmented_mask\n","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:52:55.720253Z","iopub.execute_input":"2023-12-19T09:52:55.720679Z","iopub.status.idle":"2023-12-19T09:52:55.729771Z","shell.execute_reply.started":"2023-12-19T09:52:55.720633Z","shell.execute_reply":"2023-12-19T09:52:55.728842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_image_files, val_image_files, train_mask_files, val_mask_files = train_test_split(\n    image_files, labels_files, test_size=0.2, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:52:55.730997Z","iopub.execute_input":"2023-12-19T09:52:55.731252Z","iopub.status.idle":"2023-12-19T09:52:55.741138Z","shell.execute_reply.started":"2023-12-19T09:52:55.731230Z","shell.execute_reply":"2023-12-19T09:52:55.740264Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CustomDataset(train_image_files, train_mask_files, augmentation_transforms=augment_image)\nval_dataset = CustomDataset(val_image_files, val_mask_files, augmentation_transforms=augment_image)","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:52:55.744780Z","iopub.execute_input":"2023-12-19T09:52:55.745040Z","iopub.status.idle":"2023-12-19T09:52:55.751253Z","shell.execute_reply.started":"2023-12-19T09:52:55.745015Z","shell.execute_reply":"2023-12-19T09:52:55.750387Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataloader= DataLoader(train_dataset,batch_size=8,shuffle=True)\nval_dataloader = DataLoader(val_dataset,batch_size=8,shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:52:55.752603Z","iopub.execute_input":"2023-12-19T09:52:55.752908Z","iopub.status.idle":"2023-12-19T09:52:55.761551Z","shell.execute_reply.started":"2023-12-19T09:52:55.752880Z","shell.execute_reply":"2023-12-19T09:52:55.760611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for batch_idx, (batch_images, batch_masks) in enumerate(train_dataloader):\n    print(\"Batch\", batch_idx + 1)\n    print(\"Image batch shape:\", batch_images.shape)\n    print(\"Mask batch shape:\", batch_masks.shape)\n    \n    for image, mask, image_path, mask_path in zip(batch_images, batch_masks, train_image_files, train_mask_files):\n       \n        image = image.permute((1, 2, 0)).numpy()*255.0\n        image = image.astype('uint8')\n        mask = (mask*255).numpy().astype('uint8')\n        \n        image_filename = os.path.basename(image_path)\n        mask_filename = os.path.basename(mask_path)\n        \n        plt.figure(figsize=(15, 10))\n        \n        plt.subplot(2, 4, 1)\n        plt.imshow(image, cmap='gray')\n        plt.title(f\"Original Image - {image_filename}\")\n        \n        plt.subplot(2, 4, 2)\n        plt.imshow(mask, cmap='gray')\n        plt.title(f\"Mask Image - {mask_filename}\")\n        \n        plt.tight_layout()\n        plt.show()\n    break","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-12-19T09:52:55.762574Z","iopub.execute_input":"2023-12-19T09:52:55.762862Z","iopub.status.idle":"2023-12-19T09:53:00.902609Z","shell.execute_reply.started":"2023-12-19T09:52:55.762836Z","shell.execute_reply":"2023-12-19T09:53:00.901777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for batch_idx, (batch_images, batch_masks) in enumerate(train_dataloader):\n    print(\"Batch\", batch_idx + 1)\n    print(\"Image batch shape:\", batch_images.shape)\n    print(\"Mask batch shape:\", batch_masks.shape)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-12-19T09:53:00.903711Z","iopub.execute_input":"2023-12-19T09:53:00.903961Z","iopub.status.idle":"2023-12-19T09:55:24.872168Z","shell.execute_reply.started":"2023-12-19T09:53:00.903940Z","shell.execute_reply":"2023-12-19T09:55:24.871151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_default_device():\n    if torch.cuda.is_available():\n        return torch.device('cuda')\n    else:\n        return torch.device('cpu')\ndevice = get_default_device()\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:55:24.875125Z","iopub.execute_input":"2023-12-19T09:55:24.875410Z","iopub.status.idle":"2023-12-19T09:55:24.905685Z","shell.execute_reply.started":"2023-12-19T09:55:24.875386Z","shell.execute_reply":"2023-12-19T09:55:24.904801Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ConvBlock(nn.Module):\n    def __init__(self,in_channels,out_channels):\n        super().__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_channels,out_channels, kernel_size=3, stride=1, padding=1, bias=True),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=True),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n    def forward(self,x):\n        x = self.conv(x)\n        return x\n    \n\nclass UpConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n\n        self.up = nn.Sequential(\n            nn.Upsample(scale_factor=2),\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=True),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        x = self.up(x)\n        return x\n","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:55:24.906914Z","iopub.execute_input":"2023-12-19T09:55:24.907189Z","iopub.status.idle":"2023-12-19T09:55:24.918193Z","shell.execute_reply.started":"2023-12-19T09:55:24.907166Z","shell.execute_reply":"2023-12-19T09:55:24.917267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model\n<div style=\"text-align:center\"><img src=\"https://www.mdpi.com/machines/machines-10-00327/article_deploy/html/images/machines-10-00327-g001.png\" /></div>\nThe U-Net has demonstrated remarkable performance in medical image segmentation applications, especially in scenarios where accurate localization and delineation of structures within images are crucial. Its ability to capture both local details and global context, facilitated by the skip connections, has contributed to its widespread adoption in various semantic segmentation tasks.\n\n### Attention Mechanisms in Deep Learning\n\n<div style=\"text-align:center\"><img src=\"https://www.sciltp.com/journals/public/site/images/ijndi/pic/173-3.jpg\" /></div>\nAttention mechanisms are of significant importance in the field of deep learning, providing models with the ability to focus on specific parts of the input data while filtering out irrelevant information. This capability is particularly useful in scenarios involving complex and high-dimensional data, where the model needs to selectively attend to the most relevant features for making predictions or classifications. In image recognition tasks, attention mechanisms (self attention,CBAM etc) enable convolutional neural networks (CNNs) to adaptively concentrate on salient image regions, enhancing the model's capability to capture fine-grained details and improving its overall performance. This selective focus is especially valuable when dealing with large images or when precise localization of objects within the images is necessary.\n\nIn natural language processing (NLP), attention mechanisms have revolutionized the field, providing models with the ability to selectively attend to different words or phrases within input sequences, allowing for more effective language understanding and generation. The introduction of attention mechanisms in deep learning architectures has led to significant advancements in various domains, enabling models to process information more effectively and produce more accurate and detailed results.\n\nNow, transitioning to the Attention U-Net - a deep learning architecture designed for semantic image segmentation tasks, particularly in medical imaging. The Attention U-Net incorporates attention mechanisms, allowing the model to focus on relevant parts of the input image while suppressing irrelevant noisy information. This integration enhances the model's segmentation performance, particularly in cases where precise localization and delineation of structures within medical images are crucial.  \n\n### Attention U-Net\n<div style=\"text-align:center\"><img src=\"https://miro.medium.com/v2/resize:fit:1400/1*PdYEf-OuUWkRsm2Lfrmy6A.png\" /></div>\n\nThe Attention U-Net is a deep learning architecture designed for semantic image segmentation tasks, particularly in the field of medical imaging. It is an extension of the widely used U-Net architecture, which is known for its effectiveness in biomedical image segmentation.\n\nThe key feature of the Attention U-Net is the incorporation of attention mechanisms, which allow the model to focus on relevant parts of the input image while suppressing irrelevant or noisy information. This is achieved through the integration of attention gates within the U-Net architecture, enabling the model to adaptively weigh the importance of different spatial locations within the feature maps. By incorporating attention mechanisms, the Attention U-Net aims to enhance the segmentation performance, particularly in cases where precise localization and delineation of structures within medical images are crucial. The attention modules help the model to selectively attend to informative image regions, improving the accuracy of segmentation results.","metadata":{}},{"cell_type":"code","source":"class AttentionBlock(nn.Module):\n    \"\"\"Attention block with learnable parameters\"\"\"\n\n    def __init__(self, F_g, F_l, n_coefficients):\n        super(AttentionBlock, self).__init__()\n\n        self.W_gate = nn.Sequential(\n            nn.Conv2d(F_g, n_coefficients, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(n_coefficients)\n        )\n\n        self.W_x = nn.Sequential(\n            nn.Conv2d(F_l, n_coefficients, kernel_size=1, stride=1, padding=0, bias=True),\n            nn.BatchNorm2d(n_coefficients)\n        )\n\n        self.psi = nn.Sequential(\n            nn.Conv2d(n_coefficients, 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, gate, skip_connection):\n\n        g1 = self.W_gate(gate)\n        x1 = self.W_x(skip_connection)\n        psi = self.relu(g1 + x1)\n        psi = self.psi(psi)\n        out = skip_connection * psi\n        return out\n\n\nclass AttentionUNet(nn.Module):\n\n    def __init__(self, img_ch=3, output_ch=1):\n        super(AttentionUNet, self).__init__()\n\n        self.MaxPool = nn.MaxPool2d(kernel_size=2, stride=2)\n\n        self.Conv1 = ConvBlock(img_ch, 64)\n        self.Conv2 = ConvBlock(64, 128)\n        self.Conv3 = ConvBlock(128, 256)\n        self.Conv4 = ConvBlock(256, 512)\n        self.Conv5 = ConvBlock(512, 1024)\n\n        self.Up5 = UpConv(1024, 512)\n        self.Att5 = AttentionBlock(F_g=512, F_l=512, n_coefficients=256)\n        self.UpConv5 = ConvBlock(1024, 512)\n\n        self.Up4 = UpConv(512, 256)\n        self.Att4 = AttentionBlock(F_g=256, F_l=256, n_coefficients=128)\n        self.UpConv4 = ConvBlock(512, 256)\n\n        self.Up3 = UpConv(256, 128)\n        self.Att3 = AttentionBlock(F_g=128, F_l=128, n_coefficients=64)\n        self.UpConv3 = ConvBlock(256, 128)\n\n        self.Up2 = UpConv(128, 64)\n        self.Att2 = AttentionBlock(F_g=64, F_l=64, n_coefficients=32)\n        self.UpConv2 = ConvBlock(128, 64)\n\n        self.Conv = nn.Conv2d(64, output_ch, kernel_size=1, stride=1, padding=0)\n\n    def forward(self, x):\n\n        e1 = self.Conv1(x)\n\n        e2 = self.MaxPool(e1)\n        e2 = self.Conv2(e2)\n\n        e3 = self.MaxPool(e2)\n        e3 = self.Conv3(e3)\n\n        e4 = self.MaxPool(e3)\n        e4 = self.Conv4(e4)\n\n        e5 = self.MaxPool(e4)\n        e5 = self.Conv5(e5)\n\n        d5 = self.Up5(e5)\n\n        s4 = self.Att5(gate=d5, skip_connection=e4)\n        d5 = torch.cat((s4, d5), dim=1) \n        d5 = self.UpConv5(d5)\n\n        d4 = self.Up4(d5)\n        s3 = self.Att4(gate=d4, skip_connection=e3)\n        d4 = torch.cat((s3, d4), dim=1)\n        d4 = self.UpConv4(d4)\n\n        d3 = self.Up3(d4)\n        s2 = self.Att3(gate=d3, skip_connection=e2)\n        d3 = torch.cat((s2, d3), dim=1)\n        d3 = self.UpConv3(d3)\n\n        d2 = self.Up2(d3)\n        s1 = self.Att2(gate=d2, skip_connection=e1)\n        d2 = torch.cat((s1, d2), dim=1)\n        d2 = self.UpConv2(d2)\n\n        out = self.Conv(d2)\n\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:55:24.919549Z","iopub.execute_input":"2023-12-19T09:55:24.919991Z","iopub.status.idle":"2023-12-19T09:55:24.941271Z","shell.execute_reply.started":"2023-12-19T09:55:24.919958Z","shell.execute_reply":"2023-12-19T09:55:24.940530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss Function\n\n### Dice Loss\n<div style=\"text-align:center\"><img src=\"https://miro.medium.com/v2/resize:fit:514/1*EF3VCtk-VbTIKhriaQF0YQ.png\" /></div>\nDice Loss is a widely used loss function in the field of image segmentation, particularly in medical image analysis tasks. The Dice coefficient, also known as the Sørensen-Dice coefficient, is a statistic used to gauge the similarity between two samples. The Dice Loss is particularly well-suited for segmentation tasks due to its ability to effectively handle class imbalance in the data. It addresses the issue of class skew by focusing on the relative overlap between predicted and ground truth segmentation masks, rather than absolute pixel-wise differences.\n\n### Focal Loss\n\n<div style=\"text-align:center\"><img src=\"https://miro.medium.com/v2/resize:fit:606/0*fDwafFNWavy5TrPF.png\" /></div>\nThe choice between Focal Loss and Dice Loss depends on the specific characteristics of the data and the requirements of the segmentation task.\n\n1. Class Imbalance: Focal Loss addresses the issue of class imbalance by assigning higher weights to difficult, misclassified examples. This can be particularly beneficial when dealing with imbalanced datasets where certain classes or regions of interest are underrepresented.\n\n2. Handling Misclassifications: Focal Loss focuses on correcting the misclassified examples by down-weighting easy examples. In scenarios where there are significant variations in the difficulty of examples, Focal Loss can help the model better focus on correcting these challenging cases.\n\n3. Enhanced Training Stability: Focal Loss has shown effectiveness in improving training stability, especially in scenarios where the dataset is challenging or noisy. The dynamic focusing of the loss function can assist in more stable convergence during training.\n ","metadata":{}},{"cell_type":"code","source":"class FocalLoss(nn.modules.loss._WeightedLoss):\n\n    def __init__(self, gamma=0, size_average=None, ignore_index=-100,\n                 reduce=None, balance_param=1.0):\n        super(FocalLoss, self).__init__(size_average)\n        self.gamma = gamma\n        self.size_average = size_average\n        self.ignore_index = ignore_index\n        self.balance_param = balance_param\n\n    def forward(self, input, target):\n        \n        assert len(input.shape) == len(target.shape)\n        assert input.size(0) == target.size(0)\n        assert input.size(1) == target.size(1)\n\n        logpt = - F.binary_cross_entropy_with_logits(input, target)\n        pt = torch.exp(logpt)\n\n        focal_loss = -((1 - pt) ** self.gamma) * logpt\n        balanced_focal_loss = self.balance_param * focal_loss\n        return balanced_focal_loss","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:55:24.942282Z","iopub.execute_input":"2023-12-19T09:55:24.942559Z","iopub.status.idle":"2023-12-19T09:55:24.954769Z","shell.execute_reply.started":"2023-12-19T09:55:24.942536Z","shell.execute_reply":"2023-12-19T09:55:24.954018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#PyTorch\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 = F.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":"2023-12-19T09:55:24.955703Z","iopub.execute_input":"2023-12-19T09:55:24.955995Z","iopub.status.idle":"2023-12-19T09:55:24.967347Z","shell.execute_reply.started":"2023-12-19T09:55:24.955971Z","shell.execute_reply":"2023-12-19T09:55:24.966512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice_coeff(prediction, target):\n\n    mask = np.zeros_like(prediction)\n    mask[prediction >= 0.5] = 1\n\n    inter = np.sum(mask * target)\n    union = np.sum(mask) + np.sum(target)\n    epsilon = 1e-6\n    result = np.mean(2 * inter / (union + epsilon))\n    return result","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:55:24.968468Z","iopub.execute_input":"2023-12-19T09:55:24.968794Z","iopub.status.idle":"2023-12-19T09:55:24.975881Z","shell.execute_reply.started":"2023-12-19T09:55:24.968748Z","shell.execute_reply":"2023-12-19T09:55:24.975180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloaders = {\n    'training': train_dataloader,\n    'test': val_dataloader\n}","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:55:24.976896Z","iopub.execute_input":"2023-12-19T09:55:24.977147Z","iopub.status.idle":"2023-12-19T09:55:24.985373Z","shell.execute_reply.started":"2023-12-19T09:55:24.977125Z","shell.execute_reply":"2023-12-19T09:55:24.984631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def train_and_test(model,dataloaders,optimizer,criterion,num_epochs=100, show_images=False):\n    since = time.time()\n    best_loss=1e10\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    model.to(device)\n\n    fieldnames = ['epoch', 'training_loss', 'test_loss', 'training_dice_coeff', 'test_dice_coeff']\n    train_epoch_losses = []\n    test_epoch_losses = []\n    for epoch in range(1,num_epochs+1):\n        print(f'Epoch {epoch}/{num_epochs}')\n        print('-' * 10)\n        batchsummary = {a: [0] for a in fieldnames}\n        batch_train_loss= 0.0\n        batch_test_loss = 0.0\n        \n        for phase in ['training','test']:\n            if phase =='training':\n                model.train()\n            else:\n                model.eval()\n            for sample in iter(dataloaders[phase]):\n                if show_images:\n                    grid_img = make_grid(sample[0])\n                    grid_img = grid_img.permute(1, 2, 0)\n                    plt.imshow(grid_img)\n                    plt.show()\n\n                inputs = sample[0].to(device)\n                masks = sample[1].to(device)\n                \n                masks = masks.unsqueeze(1)\n                \n                optimizer.zero_grad()\n\n                with torch.set_grad_enabled(phase == 'training'):\n                    outputs = model(inputs)\n\n                    loss = criterion(outputs, masks)\n\n                    y_pred = outputs.data.cpu().numpy().ravel()\n                    y_true = masks.data.cpu().numpy().ravel()\n\n                    batchsummary[f'{phase}_dice_coeff'].append(dice_coeff(y_pred, y_true))\n\n                    if phase == 'training':\n                        loss.backward()\n                        optimizer.step()\n\n                        batch_train_loss += loss.item() * sample[0].size(0)\n\n                    else:\n                        batch_test_loss += loss.item() * sample[0].size(0)\n\n            if phase == 'training':\n                epoch_train_loss = batch_train_loss / len(dataloaders['training'])\n                train_epoch_losses.append(epoch_train_loss)\n            else:\n                epoch_test_loss = batch_test_loss / len(dataloaders['test'])\n                test_epoch_losses.append(epoch_test_loss)\n\n            batchsummary['epoch'] = epoch\n            \n            print('{} Loss: {:.4f}'.format(phase, loss))\n\n        best_loss = np.max(batchsummary['test_dice_coeff'])\n        for field in fieldnames[3:]:\n            batchsummary[field] = np.mean(batchsummary[field])\n        print(\n            f'\\t\\t\\t train_dice_coeff: {batchsummary[\"training_dice_coeff\"]}, test_dice_coeff: {batchsummary[\"test_dice_coeff\"]}')\n\n    print('Best dice coefficient: {:4f}'.format(best_loss))\n\n    return model, train_epoch_losses, test_epoch_losses","metadata":{"execution":{"iopub.status.busy":"2023-12-19T09:55:24.986577Z","iopub.execute_input":"2023-12-19T09:55:24.986837Z","iopub.status.idle":"2023-12-19T09:55:25.001470Z","shell.execute_reply.started":"2023-12-19T09:55:24.986815Z","shell.execute_reply":"2023-12-19T09:55:25.000644Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epochs = 25\ndef train():\n    model = AttentionUNet()\n    optimizer = torch.optim.Adam(model.parameters(),lr = 1e5)\n    criterion = DiceLoss()\n    trained_model, train_epoch_losses, test_epoch_losses = train_and_test(model, dataloaders,optimizer, criterion, num_epochs= epochs)\n    return trained_model, train_epoch_losses, test_epoch_losses\n\n\ntrained_model, train_epoch_losses, test_epoch_losses = train()","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-12-19T09:55:25.002710Z","iopub.execute_input":"2023-12-19T09:55:25.003030Z","iopub.status.idle":"2023-12-19T10:54:53.270684Z","shell.execute_reply.started":"2023-12-19T09:55:25.002999Z","shell.execute_reply":"2023-12-19T10:54:53.269727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(trained_model.state_dict(), 'trained_model.pth')","metadata":{"execution":{"iopub.status.busy":"2023-12-19T10:54:53.272252Z","iopub.execute_input":"2023-12-19T10:54:53.272624Z","iopub.status.idle":"2023-12-19T10:54:53.512244Z","shell.execute_reply.started":"2023-12-19T10:54:53.272591Z","shell.execute_reply":"2023-12-19T10:54:53.511277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Results","metadata":{}},{"cell_type":"code","source":"train_plot, = plt.plot(range(1, len(train_epoch_losses) + 1), train_epoch_losses, label='train loss')\ntest_plot, = plt.plot(range(1, len(test_epoch_losses) + 1), test_epoch_losses, label='test loss')\nplt.legend(handles=[train_plot, test_plot])\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.title('Training and Test Loss Over Epochs')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-19T10:54:53.513624Z","iopub.execute_input":"2023-12-19T10:54:53.513926Z","iopub.status.idle":"2023-12-19T10:54:53.760325Z","shell.execute_reply.started":"2023-12-19T10:54:53.513902Z","shell.execute_reply":"2023-12-19T10:54:53.759495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_plot, = plt.plot(range(len(train_epoch_losses)-15), train_epoch_losses[15:], label='train loss')\ntest_plot, = plt.plot(range(len(test_epoch_losses)-15), test_epoch_losses[15:], label='test loss')\nplt.legend(handles=[train_plot, test_plot])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-19T10:54:53.761550Z","iopub.execute_input":"2023-12-19T10:54:53.761802Z","iopub.status.idle":"2023-12-19T10:54:53.932659Z","shell.execute_reply.started":"2023-12-19T10:54:53.761780Z","shell.execute_reply":"2023-12-19T10:54:53.931850Z"},"trusted":true},"execution_count":null,"outputs":[]}]}