{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":1335230,"sourceType":"datasetVersion","datasetId":775382},{"sourceId":11420633,"sourceType":"datasetVersion","datasetId":7152513}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import cv2\nimport torch\nimport numpy as np\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nfrom torch.optim.lr_scheduler import StepLR\nimg_path = '/kaggle/input/isic2018-challenge-task1-data-segmentation'","metadata":{"execution":{"iopub.status.busy":"2025-05-15T09:12:54.211839Z","iopub.execute_input":"2025-05-15T09:12:54.212138Z","iopub.status.idle":"2025-05-15T09:13:01.729646Z","shell.execute_reply.started":"2025-05-15T09:12:54.212103Z","shell.execute_reply":"2025-05-15T09:13:01.728719Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = 'cuda'\n\nepochs = 25\nlearning_rate = 1e-3\nimage_size = 256\nbatch_size = 16","metadata":{"execution":{"iopub.status.busy":"2025-05-15T09:13:03.843094Z","iopub.execute_input":"2025-05-15T09:13:03.843369Z","iopub.status.idle":"2025-05-15T09:13:03.847295Z","shell.execute_reply.started":"2025-05-15T09:13:03.843347Z","shell.execute_reply":"2025-05-15T09:13:03.846315Z"},"trusted":true},"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/isic2018-testdata/ISIC2018_Task1_Test_GroundTruth/ISIC2018_Task1_Test_GroundTruth/\"\nval_truth_path = \"/kaggle/input/isic2018-testdata/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":{"execution":{"iopub.status.busy":"2025-05-15T09:13:05.631981Z","iopub.execute_input":"2025-05-15T09:13:05.632261Z","iopub.status.idle":"2025-05-15T09:13:05.766203Z","shell.execute_reply.started":"2025-05-15T09:13:05.632239Z","shell.execute_reply":"2025-05-15T09:13:05.765394Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-05-15T09:14:25.227418Z","iopub.execute_input":"2025-05-15T09:14:25.227799Z","iopub.status.idle":"2025-05-15T09:14:25.726914Z","shell.execute_reply.started":"2025-05-15T09:14:25.227731Z","shell.execute_reply":"2025-05-15T09:14:25.726069Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-05-15T09:13:15.423867Z","iopub.execute_input":"2025-05-15T09:13:15.424153Z","iopub.status.idle":"2025-05-15T09:13:16.899713Z","shell.execute_reply.started":"2025-05-15T09:13:15.424132Z","shell.execute_reply":"2025-05-15T09:13:16.898685Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-05-15T09:13:21.404422Z","iopub.execute_input":"2025-05-15T09:13:21.404903Z","iopub.status.idle":"2025-05-15T09:13:21.411229Z","shell.execute_reply.started":"2025-05-15T09:13:21.404874Z","shell.execute_reply":"2025-05-15T09:13:21.410334Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-05-15T09:13:32.120557Z","iopub.execute_input":"2025-05-15T09:13:32.120913Z","iopub.status.idle":"2025-05-15T09:13:32.126943Z","shell.execute_reply.started":"2025-05-15T09:13:32.120882Z","shell.execute_reply":"2025-05-15T09:13:32.126263Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"index = 35\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":{"execution":{"iopub.status.busy":"2025-05-15T09:13:36.355902Z","iopub.execute_input":"2025-05-15T09:13:36.356177Z","iopub.status.idle":"2025-05-15T09:13:37.188792Z","shell.execute_reply.started":"2025-05-15T09:13:36.356156Z","shell.execute_reply":"2025-05-15T09:13:37.187822Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2025-05-15T09:14:33.093172Z","iopub.execute_input":"2025-05-15T09:14:33.093454Z","iopub.status.idle":"2025-05-15T09:14:33.100355Z","shell.execute_reply.started":"2025-05-15T09:14:33.093432Z","shell.execute_reply":"2025-05-15T09:14:33.099491Z"},"trusted":true},"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 * g\n    \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        # Store pre-attention features\n        pre_attention = x\n        \n        # Apply attention\n        post_attention = self.attention_gate(x, skip_connection)\n        \n        # Process and return all outputs\n        output = self.layers(torch.cat([skip_connection, post_attention], dim=1))\n        \n        return output, pre_attention, post_attention\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 = False\n\n        # backbone.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        self.encoder0 = nn.Sequential(\n                        backbone.conv1, \n                        backbone.bn1, \n                        backbone.relu\n                        )\n        \n        self.maxpool = backbone.maxpool\n\n        self.encoder1 = backbone.layer1\n        self.encoder2 = backbone.layer2\n        self.encoder3 = backbone.layer3\n        self.encoder4 = backbone.layer4\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, num_classes, kernel_size=3, stride=1, padding=1),\n            nn.BatchNorm2d(num_classes),\n            nn.ReLU(inplace=True),\n            )\n\n        # For storing intermediate results\n        self.intermediate_outputs = {}\n\n    def forward(self, x, return_intermediates=False):\n        a1 = self.encoder0(x)\n        a1_pooled = self.maxpool(a1)\n        a2 = self.encoder1(a1_pooled)\n        a3 = self.encoder2(a2)\n        a4 = self.encoder3(a3)\n        a5 = self.encoder4(a4)\n\n        mid = self.middle(a5)\n\n        # Capture both outputs and intermediate attention results\n        d1, pre_att0, post_att0 = self.decoder0(a4, mid)\n        d2, pre_att1, post_att1 = self.decoder1(a3, d1)\n        d3, pre_att2, post_att2 = self.decoder2(a2, d2)\n        d4, pre_att3, post_att3 = 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":{"execution":{"iopub.status.busy":"2025-05-15T09:14:52.443581Z","iopub.execute_input":"2025-05-15T09:14:52.443920Z","iopub.status.idle":"2025-05-15T09:14:55.471833Z","shell.execute_reply.started":"2025-05-15T09:14:52.443893Z","shell.execute_reply":"2025-05-15T09:14:55.470870Z"},"trusted":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":{"execution":{"iopub.status.busy":"2025-05-15T09:14:58.257251Z","iopub.execute_input":"2025-05-15T09:14:58.257721Z","iopub.status.idle":"2025-05-15T09:14:58.266311Z","shell.execute_reply.started":"2025-05-15T09:14:58.257691Z","shell.execute_reply":"2025-05-15T09:14:58.265578Z"},"trusted":true},"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-15T09:15:01.604262Z","iopub.execute_input":"2025-05-15T09:15:01.604548Z","iopub.status.idle":"2025-05-15T09:15:01.616121Z","shell.execute_reply.started":"2025-05-15T09:15:01.604525Z","shell.execute_reply":"2025-05-15T09:15:01.615130Z"}},"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":{"execution":{"iopub.status.busy":"2025-05-15T09:15:05.684258Z","iopub.execute_input":"2025-05-15T09:15:05.684647Z","iopub.status.idle":"2025-05-15T09:16:01.883672Z","shell.execute_reply.started":"2025-05-15T09:15:05.684612Z","shell.execute_reply":"2025-05-15T09:16:01.882396Z"},"trusted":true},"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":{},"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":{"execution":{"iopub.execute_input":"2025-04-16T02:31:10.328446Z","iopub.status.busy":"2025-04-16T02:31:10.328075Z","iopub.status.idle":"2025-04-16T02:33:20.695888Z","shell.execute_reply":"2025-04-16T02:33:20.694965Z","shell.execute_reply.started":"2025-04-16T02:31:10.328417Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Plot all metrics after training\nmetrics_tracker.plot_metrics()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Code to Visualize Outputs of Various Encoder-Decoder Layers After & Before Attention","metadata":{}},{"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":{},"outputs":[],"execution_count":null}]}