{"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":1335230,"sourceType":"datasetVersion","datasetId":775382},{"sourceId":11421641,"sourceType":"datasetVersion","datasetId":7153093}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import sys\nimport torch\nimport cv2\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport helper\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T10:43:48.493687Z","iopub.execute_input":"2025-05-15T10:43:48.494010Z","iopub.status.idle":"2025-05-15T10:43:53.605293Z","shell.execute_reply.started":"2025-05-15T10:43:48.493984Z","shell.execute_reply":"2025-05-15T10:43:53.604513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = 'cuda'\n\nepochs = 25\nlearning_rate = 1e-3\nimage_size = 256\nbatch_size = 16","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T10:45:12.802641Z","iopub.execute_input":"2025-05-15T10:45:12.802990Z","iopub.status.idle":"2025-05-15T10:45:12.807196Z","shell.execute_reply.started":"2025-05-15T10:45:12.802949Z","shell.execute_reply":"2025-05-15T10:45:12.806165Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\n\ntrain_path = \"/kaggle/input/isic2018-challenge-task1-data-segmentation/ISIC2018_Task1-2_Training_Input/\"\ntest_path = \"/kaggle/input/isic2018-challenge-task1-data-segmentation/ISIC2018_Task1-2_Test_Input/\"\nval_path = \"/kaggle/input/isic2018-challenge-task1-data-segmentation/ISIC2018_Task1-2_Validation_Input/\"\n\ntrain_truth_path=\"/kaggle/input/isic2018-challenge-task1-data-segmentation/ISIC2018_Task1_Training_GroundTruth/\"\ntest_truth_path = \"/kaggle/input/val-and-test-mask/ISIC2018_Task1_Test_GroundTruth/ISIC2018_Task1_Test_GroundTruth/\"\nval_truth_path = \"/kaggle/input/val-and-test-mask/ISIC2018_Task1_Validation_GroundTruth/ISIC2018_Task1_Validation_GroundTruth/\"\n\n# Directory containing images\nimage_dir = Path(train_path)\ntest_dir = Path(test_path)\nval_dir = Path(val_path)\n\n# Image extensions to look for\nimage_extensions = {\".jpg\", \".jpeg\", \".png\", \".gif\", \".bmp\", \".webp\"}\n\n# Get all image paths\ntrain_paths = [str(p) for p in image_dir.iterdir() if p.suffix.lower() in image_extensions]\ntrainm_paths = [(train_truth_path+p[-16:-4]+'_segmentation.png') for p in train_paths]\n\ntest_paths = [str(p) for p in test_dir.iterdir() if p.suffix.lower() in image_extensions]\ntestm_paths = [(test_truth_path+p[-16:-4]+'_segmentation.png') for p in test_paths]\n\nvalid_paths = [str(p) for p in val_dir.iterdir() if p.suffix.lower() in image_extensions]\nvalidm_paths = [(val_truth_path+p[-16:-4]+'_segmentation.png') for p in valid_paths]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T10:47:36.869943Z","iopub.execute_input":"2025-05-15T10:47:36.870348Z","iopub.status.idle":"2025-05-15T10:47:36.897502Z","shell.execute_reply.started":"2025-05-15T10:47:36.870311Z","shell.execute_reply":"2025-05-15T10:47:36.896632Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_path = valid_paths[0]\nmask_path = validm_paths[0]\nimage = cv2.imread(image_path)\nimage = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\nmask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)/255\n\nf, (ax1, ax2) = plt.subplots(1, 2, figsize = (10,5))\nax1.set_title('Image')\nax1.imshow(image)\nax2.set_title('Mask')\nax2.imshow(mask, cmap='gray')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T10:47:46.716054Z","iopub.execute_input":"2025-05-15T10:47:46.716490Z","iopub.status.idle":"2025-05-15T10:47:47.398224Z","shell.execute_reply.started":"2025-05-15T10:47:46.716452Z","shell.execute_reply":"2025-05-15T10:47:47.397272Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations as A\ndef train_augs():\n  return A.Compose([\n      A.Resize(image_size, image_size),\n      A.HorizontalFlip(p = 0.5),\n      A.VerticalFlip(p = 0.5)\n  ])\n\ndef val_augs():\n  return A.Compose([\n      A.Resize(image_size, image_size),\n  ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T11:33:46.523717Z","iopub.execute_input":"2025-05-15T11:33:46.524066Z","iopub.status.idle":"2025-05-15T11:33:46.529225Z","shell.execute_reply.started":"2025-05-15T11:33:46.524039Z","shell.execute_reply":"2025-05-15T11:33:46.528149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Dataset\n\nclass ISIC_2K18_Dataset(Dataset):\n  def __init__(self, df,dft):\n    self.df = df\n    self.dft = dft\n\n  def __len__(self):\n    return len(self.df)\n\n  def __getitem__(self, index):\n    image, mask = self.df[index],self.dft[index]\n    image_path = image\n    mask_path = mask\n\n    image = cv2.imread(image_path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    image = cv2.resize(image,(224,224))\n    mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n    mask = cv2.resize(mask,(224,224))\n    mask = np.expand_dims(mask, axis = -1)\n      \n\n    image = np.transpose(image, (2,0,1)).astype(np.float32)\n    mask = np.transpose(mask, (2,0,1)).astype(np.float32)\n\n    image = torch.Tensor(image) / 255.0\n    mask = torch.round(torch.Tensor(mask) / 255.0)\n\n    return image, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T11:33:46.823495Z","iopub.execute_input":"2025-05-15T11:33:46.823819Z","iopub.status.idle":"2025-05-15T11:33:46.831234Z","shell.execute_reply.started":"2025-05-15T11:33:46.823795Z","shell.execute_reply":"2025-05-15T11:33:46.830203Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainset = ISIC_2K18_Dataset(train_paths, trainm_paths)\nvalidset = ISIC_2K18_Dataset(valid_paths, validm_paths)\ntestset = ISIC_2K18_Dataset(test_paths, testm_paths)\n\nprint(\"The size of the training dataset is: \", len(trainset))\nprint(\"The size of the validation dataset is: \", len(validset))\nprint(\"The size of the test dataset is: \", len(testset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T11:33:48.773759Z","iopub.execute_input":"2025-05-15T11:33:48.774076Z","iopub.status.idle":"2025-05-15T11:33:48.780705Z","shell.execute_reply.started":"2025-05-15T11:33:48.774051Z","shell.execute_reply":"2025-05-15T11:33:48.779678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"index = 0\n\nimage, mask = validset[index]\nf, (ax1, ax2) = plt.subplots(1, 2, figsize=(10,5))\n\nax1.set_title('IMAGE')\nax1.imshow(image.permute(1,2,0),cmap = 'gray')\n\nax2.set_title('GROUND TRUTH')\nax2.imshow(mask.permute(1,2,0),cmap = 'gray')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T11:33:49.068849Z","iopub.execute_input":"2025-05-15T11:33:49.069187Z","iopub.status.idle":"2025-05-15T11:33:49.616205Z","shell.execute_reply.started":"2025-05-15T11:33:49.069158Z","shell.execute_reply":"2025-05-15T11:33:49.615144Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\ntrain_loader = DataLoader(trainset, batch_size = batch_size, shuffle=True)\nvalid_loader = DataLoader(validset, batch_size = batch_size)\ntest_loader = DataLoader(testset, batch_size = batch_size)\n\nprint(\"The total  number of batches in train loader are: \", len(train_loader))\nprint(\"The total  number of batches in valid loader are: \", len(valid_loader))\nprint(\"The total  number of batches in test loader are: \", len(test_loader))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T11:33:51.279453Z","iopub.execute_input":"2025-05-15T11:33:51.279784Z","iopub.status.idle":"2025-05-15T11:33:51.287241Z","shell.execute_reply.started":"2025-05-15T11:33:51.279759Z","shell.execute_reply":"2025-05-15T11:33:51.286351Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch as torch\nimport torch.nn as nn\nfrom torchvision.models import resnet34, ResNet34_Weights\n\nclass AttentionGate(nn.Module):\n    def __init__(self, in_c, out_c):\n        super().__init__()\n\n        self.Wg = nn.Sequential(\n            nn.Conv2d(in_c[0], out_c, kernel_size=1, padding=0),\n            nn.BatchNorm2d(out_c)\n        )\n\n        self.Ws = nn.Sequential(\n            nn.Conv2d(in_c[1], out_c, kernel_size=1, padding=0),\n            nn.BatchNorm2d(out_c)\n        )\n\n        self.relu = nn.ReLU(inplace=True)\n        self.output = nn.Sequential(\n            nn.Conv2d(out_c, 1, kernel_size=1, padding=0),\n            nn.Sigmoid()\n        )\n\n    def forward(self, g, s):\n        Wg = self.Wg(g)\n        Ws = self.Ws(s)\n        out = self.relu(Wg + Ws)\n        out = self.output(out)\n        return out * s\n\nclass AttentionGateM2(nn.Module):\n    def __init__(self, in_c, out_c):\n        super().__init__()\n\n        self.Wg = nn.Sequential(\n            nn.Conv2d(in_c[0], out_c, kernel_size=(1,1), stride=(1,1), padding=0),\n            nn.BatchNorm2d(out_c)\n        )\n\n        self.Ws = nn.Sequential(\n            nn.Conv2d(in_c[1], out_c, kernel_size=(1,1), stride=(2,2), padding=0),\n            nn.BatchNorm2d(out_c)\n        )\n\n        self.relu = nn.ReLU(inplace=True)\n        self.output = nn.Sequential(\n            nn.Conv2d(out_c, 1, kernel_size=1, padding=0),\n            nn.Sigmoid(),\n        )\n\n    def forward(self, s, g):\n        Wg = self.Wg(g)\n        Ws = self.Ws(s)\n        out = self.relu(Wg + Ws)\n        out = self.output(out)\n        return out * g\n\nclass DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels, kernel_size=3, padding=1):\n        super(DoubleConv, self).__init__()\n        self.layers = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=1, padding=padding),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(),\n            nn.Conv2d(out_channels, out_channels, kernel_size=kernel_size, stride=1, padding=padding),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(),\n        )\n\n    def forward(self, x):\n        return self.layers(x)\n\n\nclass Decoder(nn.Module):\n    def __init__(self, up_in_channels, x_in_channels, kernel_size=3, padding=1, dropout=0):\n        super(Decoder, self).__init__()\n\n        self.upsample = nn.ConvTranspose2d(up_in_channels, up_in_channels // 2, kernel_size=3, stride=2, padding=1, output_padding=1)\n\n        up_out_channels = up_in_channels // 2\n        self.attention_gate = AttentionGate([up_out_channels, x_in_channels], up_out_channels)\n        in_channels = up_in_channels // 2 + x_in_channels\n        out_channels = in_channels // 2\n\n        self.layers = DoubleConv(in_channels, out_channels, kernel_size, padding)\n\n    def forward(self, skip_connection, x):\n        x = self.upsample(x)\n        \n        output = self.layers(torch.cat([skip_connection, x], dim=1))\n        \n        return output\n\nclass SelfAttention(nn.Module):\n    def __init__(self, x_channel,g_channel,x_size,y_size):\n        super(SelfAttention, self).__init__()\n        wq = torch.randn((x_size, y_size), requires_grad=True).to(device)\n        wk = torch.randn((x_size, y_size), requires_grad=True).to(device)\n        wv = torch.randn((x_size, y_size), requires_grad=True).to(device)\n        self.wq = nn.Parameter(wq)\n        self.wk = nn.Parameter(wk)\n        self.wv = nn.Parameter(wv)\n        self.softmax = nn.Softmax(dim=1)\n\n        self.attention_gate_m2=AttentionGateM2([g_channel,x_channel],x_channel)\n\n    def forward(self, x, g):\n        xq= torch.matmul(x, self.wq)\n        xk= torch.matmul(x, self.wk)\n        xv= torch.matmul(x, self.wv)\n\n        preAttention = g\n        attention_map = torch.matmul(xq, xk.transpose(2, 3))\n        attention_map = self.softmax(attention_map)\n        attention_map = torch.matmul(attention_map, xv)\n        postAttention = self.attention_gate_m2(attention_map, g)\n        return preattention , postAttention\n    \n    \n\nclass AttenResUnet(nn.Module):\n    def __init__(self, num_classes):\n\n        backbone = resnet34(weights=ResNet34_Weights.IMAGENET1K_V1)\n\n        super(AttenResUnet, self).__init__()\n\n        for param in backbone.parameters():\n            param.requires_grad = True\n\n        self.encoder0 = nn.Sequential(\n                        backbone.conv1, \n                        backbone.bn1, \n                        backbone.relu\n                        )\n        \n        self.maxpool = backbone.maxpool\n        self.encoder1 = backbone.layer1\n        self.atten1=SelfAttention(64,64,112,112) \n        self.encoder2 = backbone.layer2\n        self.atten2=SelfAttention(64,128,56,56)\n        self.encoder3 = backbone.layer3\n        self.atten3=SelfAttention(128,256,28,28)\n        self.encoder4 = backbone.layer4\n        self.atten4=SelfAttention(256,512,14,14)\n\n        self.middle = DoubleConv(512, 512)\n\n        self.decoder0 = Decoder(512, 256)\n        self.decoder1 = Decoder(256, 128)\n        self.decoder2 = Decoder(128, 64)\n        self.decoder3 = Decoder(64, 64)\n\n        self.resize = nn.ConvTranspose2d(48, 16, kernel_size=3, stride=2, padding=1, output_padding=1)\n\n        self.final_layer = nn.Sequential(\n            nn.Conv2d(16, 1, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(1),\n            nn.ReLU(inplace=True),\n            )\n\n\n    def forward(self, x):\n\n        # Capture both outputs and intermediate attention results\n        \n        a1 = self.encoder0(x)\n        a1_pooled = self.maxpool(a1)\n        a2p = self.encoder1(a1_pooled)\n        \n        pre_att0, post_att0=self.atten1(a1,a2p)\n        a2=post_att0\n        a3p = self.encoder2(a2)\n\n        pre_att1, post_att1 = self.atten2(a2,a3p)\n        a3= post_att1\n        a4p = self.encoder3(a3)\n        \n        pre_att2, post_att2 = self.atten2(a3,a4p)\n        a4 = post_att2\n        a5p = self.encoder4(a4)\n        \n        pre_att3, post_att3 = self.atten2(a4,a5p)\n        a5= post_att3\n\n        mid = self.middle(a5)\n\n        d1 = self.decoder0(a4, mid)\n        d2 = self.decoder1(a3, d1)\n        d3 = self.decoder2(a2, d2)\n        d4 = self.decoder3(a1, d3)\n\n        resized = self.resize(d4)\n        final = self.final_layer(resized)\n\n        if return_intermediates:\n            self.intermediate_outputs = {\n                # Encoder outputs\n                'encoder1': a1,\n                'encoder2': a2,\n                'encoder3': a3,\n                'encoder4': a4,\n                'encoder5': a5,\n                'middle': mid,\n                \n                # Decoder outputs\n                'decoder1': d1,\n                'decoder2': d2,\n                'decoder3': d3,\n                'decoder4': d4,\n                \n                # Pre-attention maps\n                'pre_attention1': pre_att0,\n                'pre_attention2': pre_att1,\n                'pre_attention3': pre_att2,\n                'pre_attention4': pre_att3,\n                \n                # Post-attention maps\n                'post_attention1': post_att0,\n                'post_attention2': post_att1,\n                'post_attention3': post_att2,\n                'post_attention4': post_att3,\n            }\n            return final, self.intermediate_outputs\n\n        return final","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T11:35:57.892421Z","iopub.execute_input":"2025-05-15T11:35:57.892786Z","iopub.status.idle":"2025-05-15T11:35:57.916862Z","shell.execute_reply.started":"2025-05-15T11:35:57.892761Z","shell.execute_reply":"2025-05-15T11:35:57.916044Z"},"jupyter":{"source_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport torch.nn.functional as F\n\n# Dice Loss\ndef dice_loss(preds, targets, smooth=1e-6):\n    preds = preds.view(-1)\n    targets = targets.view(-1)\n    intersection = (preds * targets).sum()\n    return 1 - ((2. * intersection + smooth) / (preds.sum() + targets.sum() + smooth))\n\n# IoU Loss\ndef iou_loss(preds, targets, smooth=1e-6):\n    preds = preds.view(-1)\n    targets = targets.view(-1)\n    intersection = (preds * targets).sum()\n    union = preds.sum() + targets.sum() - intersection\n    return 1 - ((intersection + smooth) / (union + smooth))\n\n# Focal Loss (binary)\ndef focal_loss(preds, targets, alpha=0.8, gamma=2):\n    BCE = F.binary_cross_entropy_with_logits(preds, targets, reduction='none')\n    pt = torch.exp(-BCE)\n    FL = alpha * (1 - pt) ** gamma * BCE\n    return FL.mean()\n\n# Tversky Loss\ndef tversky_loss(preds, targets, alpha=0.5, beta=0.5, smooth=1e-6):\n    preds = torch.sigmoid(preds).view(-1)\n    targets = targets.view(-1)\n    TP = (preds * targets).sum()\n    FP = ((1 - targets) * preds).sum()\n    FN = (targets * (1 - preds)).sum()\n    return 1 - ((TP + smooth) / (TP + alpha * FP + beta * FN + smooth))\n\n# Dice Score for Evaluation\ndef dice_score(preds, targets, smooth=1e-6):\n    preds = torch.sigmoid(preds)\n    preds = (preds > 0.5).float()\n\n    intersection = (preds * targets).sum(dim=(1, 2, 3))\n    union = preds.sum(dim=(1, 2, 3)) + targets.sum(dim=(1, 2, 3))\n\n    dice = (2 * intersection + smooth) / (union + smooth)\n    return dice.mean()  # batch average","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T11:36:08.479557Z","iopub.execute_input":"2025-05-15T11:36:08.479884Z","iopub.status.idle":"2025-05-15T11:36:08.489300Z","shell.execute_reply.started":"2025-05-15T11:36:08.479858Z","shell.execute_reply":"2025-05-15T11:36:08.488439Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MetricsTracker:\n    def __init__(self):\n        self.train_losses = []\n        self.val_losses = []\n        self.dice_scores = []\n        self.iou_losses = []\n        self.focal_losses = []\n        self.tversky_losses = []\n        \n    def update(self, train_loss, val_loss, dice_score, iou_loss, focal_loss, tversky_loss):\n        self.train_losses.append(train_loss)\n        self.val_losses.append(val_loss)\n        self.dice_scores.append(dice_score)\n        self.iou_losses.append(iou_loss)\n        self.focal_losses.append(focal_loss)\n        self.tversky_losses.append(tversky_loss)\n    \n    def plot_metrics(self):\n        epochs = range(1, len(self.train_losses) + 1)\n        \n        fig, axs = plt.subplots(2, 2, figsize=(14, 10))\n        \n        # Plot 1: Training and Validation Loss\n        axs[0, 0].plot(epochs, self.train_losses, 'b-', label='Training Loss')\n        axs[0, 0].plot(epochs, self.val_losses, 'r-', label='Validation Loss')\n        axs[0, 0].set_title('Training and Validation Loss')\n        axs[0, 0].set_xlabel('Epochs')\n        axs[0, 0].set_ylabel('Loss')\n        axs[0, 0].legend()\n        axs[0, 0].grid(True)\n        \n        # Plot 2: Dice Scores\n        axs[0, 1].plot(epochs, self.dice_scores, 'g-')\n        axs[0, 1].set_title('Dice Score')\n        axs[0, 1].set_xlabel('Epochs')\n        axs[0, 1].set_ylabel('Score')\n        axs[0, 1].grid(True)\n        \n        # Plot 3: IoU and Tversky Losses\n        axs[1, 0].plot(epochs, self.iou_losses, 'c-', label='IoU Loss')\n        axs[1, 0].plot(epochs, self.tversky_losses, 'm-', label='Tversky Loss')\n        axs[1, 0].set_title('IoU and Tversky Losses')\n        axs[1, 0].set_xlabel('Epochs')\n        axs[1, 0].set_ylabel('Loss')\n        axs[1, 0].legend()\n        axs[1, 0].grid(True)\n        \n        # Plot 4: Focal Loss\n        axs[1, 1].plot(epochs, self.focal_losses, 'y-')\n        axs[1, 1].set_title('Focal Loss')\n        axs[1, 1].set_xlabel('Epochs')\n        axs[1, 1].set_ylabel('Loss')\n        axs[1, 1].grid(True)\n        \n        plt.tight_layout()\n        plt.show()\n        \n        # Additionally plot training vs validation loss in a separate figure\n        plt.figure(figsize=(10, 6))\n        plt.plot(epochs, self.train_losses, 'b-', label='Training Loss')\n        plt.plot(epochs, self.val_losses, 'r-', label='Validation Loss')\n        plt.title('Training and Validation Loss')\n        plt.xlabel('Epochs')\n        plt.ylabel('Loss')\n        plt.legend()\n        plt.grid(True)\n        plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T11:36:10.490344Z","iopub.execute_input":"2025-05-15T11:36:10.490689Z","iopub.status.idle":"2025-05-15T11:36:10.501633Z","shell.execute_reply.started":"2025-05-15T11:36:10.490661Z","shell.execute_reply":"2025-05-15T11:36:10.500761Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize model and training components\nmodel = AttenResUnet(1).to(device)\ncriterian = torch.nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\nscheduler = StepLR(optimizer, step_size=2, gamma=0.1)\n\n# Initialize metrics tracker\nmetrics_tracker = MetricsTracker()\n\nnum_epochs = 10\nfor epoch in range(num_epochs):\n    \n    # --- Training Loop ---\n    model.train()\n    epoch_loss = 0\n    \n    # Use tqdm for progress bar\n    for imgs, masks in tqdm(train_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Train]\"):\n        imgs, masks = imgs.to(device), masks.to(device)\n\n        optimizer.zero_grad()\n        preds = model(imgs)\n        loss = criterian(preds, masks)\n        \n        loss.backward()\n        optimizer.step()\n        \n        epoch_loss += loss.item()\n\n    scheduler.step()\n    avg_train_loss = epoch_loss / len(train_loader)\n\n    # --- Validation Loop ---\n    model.eval()\n    val_loss = 0\n    val_dice = 0\n    last_iou = 0\n    last_focal = 0\n    last_tversky = 0\n    \n    with torch.no_grad():\n        for val_imgs, val_masks in tqdm(valid_loader, desc=f\"Epoch {epoch+1}/{num_epochs} [Valid]\"):\n            val_imgs, val_masks = val_imgs.to(device), val_masks.to(device)\n    \n            val_preds = model(val_imgs)\n            loss = criterian(val_preds, val_masks)\n            val_loss += loss.item()\n    \n            # Compute metrics (we'll save the last batch's metrics for tracking)\n            sigmoid_preds = torch.sigmoid(val_preds)\n            dice = dice_score(val_preds, val_masks)\n            val_dice += dice.item()\n            \n            # Calculate other metrics on last batch\n            last_iou = iou_loss(sigmoid_preds, val_masks).item()\n            last_focal = focal_loss(val_preds, val_masks).item()\n            last_tversky = tversky_loss(val_preds, val_masks).item()\n    \n    avg_val_loss = val_loss / len(valid_loader)\n    avg_val_dice = val_dice / len(valid_loader)\n    \n    # Store metrics\n    metrics_tracker.update(\n        avg_train_loss, \n        avg_val_loss, \n        avg_val_dice, \n        last_iou, \n        last_focal, \n        last_tversky\n    )\n    \n    # Print only the essential information\n    print(f\"Epoch [{epoch+1}/{num_epochs}] - Train Loss: {avg_train_loss:.4f}, Val Loss: {avg_val_loss:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T11:36:58.956659Z","iopub.execute_input":"2025-05-15T11:36:58.957007Z","iopub.status.idle":"2025-05-15T11:36:58.961668Z","shell.execute_reply.started":"2025-05-15T11:36:58.956978Z","shell.execute_reply":"2025-05-15T11:36:58.960659Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_save_path = \"/kaggle/working/attenresunet.pth\"\ntorch.save(model.state_dict(), model_save_path)\nprint(f\"Model weights saved to {model_save_path}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_test_data(model, test_loader, criterion):\n    \"\"\"\n    Evaluates model on test data and returns comprehensive metrics\n    \n    Args:\n        model: The trained model to evaluate\n        test_loader: DataLoader containing test data\n        criterion: Loss function used for evaluation\n    \n    Returns:\n        Dictionary of metrics\n    \"\"\"\n    model.eval()\n    test_loss = 0\n    test_dice = 0\n    test_iou = 0\n    test_focal = 0\n    test_tversky = 0\n    test_samples = 0\n    \n    with torch.no_grad():\n        for test_imgs, test_masks in tqdm(test_loader, desc=\"Evaluating test data\"):\n            test_imgs, test_masks = test_imgs.to(device), test_masks.to(device)\n            \n            # Forward pass\n            test_preds = model(test_imgs)\n            \n            # Calculate primary loss\n            loss = criterion(test_preds, test_masks)\n            test_loss += loss.item() * test_imgs.size(0)\n            \n            # Calculate metrics\n            sigmoid_preds = torch.sigmoid(test_preds)\n            dice = dice_score(test_preds, test_masks)\n            \n            # Add metrics\n            test_dice += dice.item() * test_imgs.size(0)\n            test_iou += (1 - iou_loss(sigmoid_preds, test_masks).item()) * test_imgs.size(0)\n            test_focal += focal_loss(test_preds, test_masks).item() * test_imgs.size(0)\n            test_tversky += tversky_loss(test_preds, test_masks).item() * test_imgs.size(0)\n            \n            test_samples += test_imgs.size(0)\n    \n    # Normalize metrics by total samples\n    metrics = {\n        'loss': test_loss / test_samples,\n        'dice_score': test_dice / test_samples,\n        'iou_score': test_iou / test_samples,\n        'focal_loss': test_focal / test_samples,\n        'tversky_loss': test_tversky / test_samples\n    }\n    \n    return metrics\n\ndef visualize_test_predictions(model, test_loader, num_samples=5):\n    \"\"\"\n    Visualizes test predictions alongside original images and ground truth masks\n    \n    Args:\n        model: The trained model to evaluate\n        test_loader: DataLoader containing test data\n        num_samples: Number of samples to visualize\n    \"\"\"\n    model.eval()\n    with torch.no_grad():\n        # Get a batch from the test loader\n        imgs, masks = next(iter(test_loader))\n        imgs = imgs.to(device)\n        \n        # Generate predictions\n        preds = torch.sigmoid(model(imgs)) > 0.5\n        \n        # Limit to specified number of samples\n        n = min(num_samples, imgs.size(0))\n        \n        # Create visualization figure\n        fig, axs = plt.subplots(3, n, figsize=(n*3, 9))\n        \n        for i in range(n):\n            # Display original image\n            axs[0, i].imshow(imgs[i].cpu().permute(1, 2, 0))\n            axs[0, i].set_title(\"Original Image\")\n            axs[0, i].axis('off')\n            \n            # Display ground truth mask\n            axs[1, i].imshow(masks[i].squeeze(0), cmap='gray')\n            axs[1, i].set_title(\"Ground Truth\")\n            axs[1, i].axis('off')\n            \n            # Display prediction mask\n            axs[2, i].imshow(preds[i].squeeze(0).cpu(), cmap='gray')\n            axs[2, i].set_title(\"Prediction\")\n            axs[2, i].axis('off')\n        \n        plt.tight_layout()\n        plt.show()\n\n# Main evaluation code\ntest_metrics = evaluate_test_data(model, test_loader, criterian)\n\n# Display summary of test metrics\nprint(\"\\n===== Test Dataset Evaluation =====\")\nprint(f\"Loss: {test_metrics['loss']:.4f}\")\nprint(f\"Dice Score: {test_metrics['dice_score']:.4f}\")\nprint(f\"IoU Score: {test_metrics['iou_score']:.4f}\")\nprint(f\"Focal Loss: {test_metrics['focal_loss']:.4f}\")\nprint(f\"Tversky Loss: {test_metrics['tversky_loss']:.4f}\")\n\n# Visualize predictions\nvisualize_test_predictions(model, test_loader, num_samples=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T20:11:31.116258Z","iopub.execute_input":"2025-04-26T20:11:31.116531Z","iopub.status.idle":"2025-04-26T20:12:56.384130Z","shell.execute_reply.started":"2025-04-26T20:11:31.116510Z","shell.execute_reply":"2025-04-26T20:12:56.383215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot all metrics after training\nmetrics_tracker.plot_metrics()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_attention_maps(model, input_batch, sample_idx=0):\n    \"\"\"\n    Visualizes the feature maps before and after attention gates\n    \n    Args:\n        model: Trained AttenResUnet model\n        input_batch: Batch of images [B, C, H, W]\n        sample_idx: Index of the sample to visualize from the batch\n    \"\"\"\n    model.eval()\n    with torch.no_grad():\n        # Forward pass with intermediate outputs\n        _, intermediates = model(input_batch, return_intermediates=True)\n        \n        # Original image for reference\n        original_img = input_batch[sample_idx].cpu().permute(1, 2, 0)\n        \n        # Create figure\n        fig, axes = plt.subplots(4, 3, figsize=(15, 16))\n        \n        # Display original image in first column\n        axes[0, 0].imshow(original_img)\n        axes[0, 0].set_title('Original Image')\n        axes[0, 0].axis('off')\n        \n        # For each decoder level\n        for i in range(4):\n            # Get pre and post attention feature maps for current level\n            pre_att = intermediates[f'pre_attention{i+1}'][sample_idx]\n            post_att = intermediates[f'post_attention{i+1}'][sample_idx]\n            \n            # Reduce channels by taking mean\n            pre_att_mean = pre_att.mean(dim=0).cpu().numpy()\n            post_att_mean = post_att.mean(dim=0).cpu().numpy()\n            \n            # Normalize for better visualization\n            pre_att_mean = (pre_att_mean - pre_att_mean.min()) / (pre_att_mean.max() - pre_att_mean.min() + 1e-8)\n            post_att_mean = (post_att_mean - post_att_mean.min()) / (post_att_mean.max() - post_att_mean.min() + 1e-8)\n            \n            # Show feature maps\n            axes[i, 1].imshow(pre_att_mean, cmap='viridis')\n            axes[i, 1].set_title(f'Level {i+1}: Before Attention')\n            axes[i, 1].axis('off')\n            \n            axes[i, 2].imshow(post_att_mean, cmap='viridis')\n            axes[i, 2].set_title(f'Level {i+1}: After Attention')\n            axes[i, 2].axis('off')\n            \n        plt.tight_layout()\n        plt.show()\n\n# Example usage with trained model\nwith torch.no_grad():\n    sample_batch, _ = next(iter(valid_loader))\n    sample_batch = sample_batch.to(device)\n    visualize_attention_maps(model, sample_batch)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-26T20:12:56.385779Z","iopub.execute_input":"2025-04-26T20:12:56.386085Z","iopub.status.idle":"2025-04-26T20:12:59.010609Z","shell.execute_reply.started":"2025-04-26T20:12:56.386058Z","shell.execute_reply":"2025-04-26T20:12:59.009446Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}