{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.14"},"accelerator":"GPU","colab":{"gpuType":"T4","name":"Pretarained_2d_resnetUnet","provenance":[]},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":8677798,"sourceType":"datasetVersion","datasetId":5201742},{"sourceId":12440538,"sourceType":"datasetVersion","datasetId":7847542},{"sourceId":12440676,"sourceType":"datasetVersion","datasetId":7847647},{"sourceId":12440809,"sourceType":"datasetVersion","datasetId":7847744},{"sourceId":12445375,"sourceType":"datasetVersion","datasetId":7850532},{"sourceId":12446744,"sourceType":"datasetVersion","datasetId":7851411},{"sourceId":12448779,"sourceType":"datasetVersion","datasetId":7852797}],"dockerImageVersionId":30762,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":16847.177412,"end_time":"2024-06-17T17:22:18.105705","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-06-17T12:41:30.928293","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"\n\nimport cv2\nimport torch\nfrom typing import Dict\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\nimport numpy as np\nimport torch.nn as nn3\nfrom sklearn.model_selection import train_test_split\nfrom torchvision.models import resnet50\nimport pandas as pd\nimport os\nfrom torch.cuda.amp import GradScaler, autocast\nimport glob\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport torch.nn as nn\nimport torchvision.models as models\nfrom sklearn.metrics import jaccard_score, f1_score, precision_score, recall_score\nimport torch.optim as optim\n","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","id":"Ft2GG7i26QTN","papermill":{"duration":7.295519,"end_time":"2024-06-17T12:41:40.996282","exception":false,"start_time":"2024-06-17T12:41:33.700763","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-07-10T17:59:26.271194Z","iopub.execute_input":"2025-07-10T17:59:26.271682Z","iopub.status.idle":"2025-07-10T17:59:26.539468Z","shell.execute_reply.started":"2025-07-10T17:59:26.271644Z","shell.execute_reply":"2025-07-10T17:59:26.537655Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MyDataset(Dataset):\n    def __init__(self, meta, image_size=(512, 512), augment=False):\n        self.meta = meta\n        self.image_size = image_size\n        self.augment = augment\n        self.transforms = T.Compose([\n        T.ToPILImage(),\n        T.RandomHorizontalFlip(),\n        T.RandomVerticalFlip(),\n        T.RandomRotation(90),\n        #.RandomResizedCrop(size=(512, 512), scale=(0.8, 1.0)),\n        T.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1),\n        T.ToTensor()\n    ]) if augment else None\n\n\n    def __len__(self):\n        return len(self.meta)\n\n    def __getitem__(self, index):\n        f = self.meta[index]\n        if f is None:\n            raise ValueError(\"File path not found in metadata\")\n\n        # Load image\n        image = cv2.imread(f, cv2.IMREAD_GRAYSCALE)\n        image = cv2.resize(image, self.image_size)\n        image = torch.from_numpy(image).float().unsqueeze(0)  # Convert to tensor and add channel dimension\n        image = norm_by_percentile(image)\n\n        # Load mask\n        maskfile = f.replace('images', 'labels')  # Replace 'images' with 'labels' in the file path\n        if os.path.exists(maskfile):\n            mask = cv2.imread(maskfile, cv2.IMREAD_GRAYSCALE)\n            mask = cv2.resize(mask, self.image_size)\n            mask = torch.from_numpy(mask).float().unsqueeze(0)  # Convert to tensor and add channel dimension\n            mask = norm_by_percentile(mask)\n        else:\n            mask = torch.zeros_like(image)  # Create a zero mask with same shape as image\n\n        if self.augment:\n            # Apply the same transformation to both image and mask\n            seed = torch.random.seed()  # Set the seed for reproducibility\n            torch.random.manual_seed(seed)\n            image = self.transforms(image)\n            torch.random.manual_seed(seed)\n            mask = self.transforms(mask)\n\n        return image, mask\n\n\n\n\nclass DotDict(dict):\n    \"\"\"dot.notation access to dictionary attributes\"\"\"\n    def __getattr__(self, attr):\n        return self.get(attr)\n\n    __setattr__= dict.__setitem__\n    __delattr__= dict.__delitem__","metadata":{"id":"BrngayU16QTP","papermill":{"duration":0.019353,"end_time":"2024-06-17T12:41:41.020354","exception":false,"start_time":"2024-06-17T12:41:41.001001","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-07-10T17:59:26.540179Z","iopub.status.idle":"2025-07-10T17:59:26.540486Z","shell.execute_reply.started":"2025-07-10T17:59:26.540336Z","shell.execute_reply":"2025-07-10T17:59:26.540351Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_images(images, masks, n=5):\n    \"\"\"\n    Display a batch of images and their corresponding masks.\n\n    Args:\n    - images (torch.Tensor): Batch of images.\n    - masks (torch.Tensor): Batch of masks.\n    - n (int): Number of images to display. Default is 5.\n    \"\"\"\n    images = images.cpu().numpy()\n    masks = masks.cpu().numpy()\n\n    batch_size = images.shape[0]\n    n = min(n, batch_size)  # Ensure n does not exceed batch size\n\n    plt.figure(figsize=(15, 10))\n    for i in range(n):\n        plt.subplot(n, 2, 2 * i + 1)\n        plt.imshow(images[i, 0], cmap='gray')\n        plt.title('Image')\n        plt.axis('off')\n\n        plt.subplot(n, 2, 2 * i + 2)\n        plt.imshow(masks[i, 0], cmap='gray')\n        plt.title('Mask')\n        plt.axis('off')\n    plt.show()","metadata":{"id":"0ZOs7BeP6QT3","papermill":{"duration":0.01377,"end_time":"2024-06-17T12:41:41.038327","exception":false,"start_time":"2024-06-17T12:41:41.024557","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-07-10T17:59:26.541842Z","iopub.status.idle":"2025-07-10T17:59:26.542161Z","shell.execute_reply.started":"2025-07-10T17:59:26.542012Z","shell.execute_reply":"2025-07-10T17:59:26.542028Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_checkpoint(epoch, model, optimizer, loss, file_path):\n    torch.save({\n        'epoch': epoch,\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'loss': loss,\n    }, file_path)\n\ndef load_checkpoint(file_path, model, optimizer):\n    checkpoint = torch.load(file_path)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n    epoch = checkpoint['epoch']\n    loss = checkpoint['loss']\n    return epoch, model, optimizer, loss","metadata":{"id":"VDTtFpPA6QT8","papermill":{"duration":0.012729,"end_time":"2024-06-17T12:41:41.055128","exception":false,"start_time":"2024-06-17T12:41:41.042399","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-07-10T17:59:26.543957Z","iopub.status.idle":"2025-07-10T17:59:26.544249Z","shell.execute_reply.started":"2025-07-10T17:59:26.544106Z","shell.execute_reply":"2025-07-10T17:59:26.54412Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def norm_by_percentile(x, low=10, high=99.8, alpha=0.01):\n    xmin = np.percentile(x, low)\n    xmax = np.percentile(x, high)\n    x = (x - xmin) / (xmax - xmin + 1e-5)  # Avoid division by zero\n    x = np.clip(x, 0, 1)  # Ensure values are in the range [0, 1]\n    if 1:\n        x[x > 1] = (x[x > 1] - 1) * alpha + 1\n        x[x < 0] = x[x < 0] * alpha\n    return x\n    #return x / 65536.0","metadata":{"id":"C_8fMV5m6QUy","papermill":{"duration":0.012447,"end_time":"2024-06-17T12:41:41.071645","exception":false,"start_time":"2024-06-17T12:41:41.059198","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-07-10T17:59:26.545112Z","iopub.status.idle":"2025-07-10T17:59:26.545389Z","shell.execute_reply.started":"2025-07-10T17:59:26.545258Z","shell.execute_reply":"2025-07-10T17:59:26.545272Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Net(nn.Module):\n    def __init__(self, weights= None):\n        super(Net, self).__init__()\n\n        self.backbone = models.resnet50(weights=weights)\n        self.backbone.conv1 = nn.Conv2d(1, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n\n        self.encoder = nn.Sequential(*list(self.backbone.children())[:-2])\n\n        self.upconv1 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv1 = nn.Conv2d(2048 + 1024, 1024, kernel_size=3, padding=1)\n\n        self.upconv2 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv2 = nn.Conv2d(1024 + 512, 512, kernel_size=3, padding=1)\n\n        self.upconv3 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv3 = nn.Conv2d(512 + 256, 256, kernel_size=3, padding=1)\n\n        self.upconv4 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv4 = nn.Conv2d(256 + 64, 64, kernel_size=3, padding=1)\n\n        self.upconv5 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        self.conv5 = nn.Conv2d(64, 32, kernel_size=3, padding=1)\n\n        self.vessel_out = nn.Conv2d(32, 1, kernel_size=1)\n\n        # Apply dropout in the model definition\n        self.dropout = nn.Dropout(p=0.4)  # Add dropout layer\n\n    def forward(self, x):\n        # Encoder\n        enc1 = self.backbone.conv1(x)\n        enc1 = self.backbone.bn1(enc1)\n        enc1 = self.backbone.relu(enc1)\n        enc1 = self.backbone.maxpool(enc1)\n\n        enc2 = self.backbone.layer1(enc1)\n        enc3 = self.backbone.layer2(enc2)\n        enc4 = self.backbone.layer3(enc3)\n        enc5 = self.backbone.layer4(enc4)\n\n        # Decoder with debugging prints\n        dec1 = self.upconv1(enc5)\n        if dec1.size()[2:] != enc4.size()[2:]:\n            dec1 = F.interpolate(dec1, size=enc4.shape[2:], mode='bilinear', align_corners=True)\n        dec1 = torch.cat((dec1, enc4), dim=1)\n        dec1 = self.conv1(dec1)\n        dec1 = self.dropout(dec1)  # Apply dropout before passing through a layer\n\n        dec2 = self.upconv2(dec1)\n        if dec2.size()[2:] != enc3.size()[2:]:\n            dec2 = F.interpolate(dec2, size=enc3.shape[2:], mode='bilinear', align_corners=True)\n        dec2 = torch.cat((dec2, enc3), dim=1)\n        dec2 = self.conv2(dec2)\n        dec2 = self.dropout(dec2)  # Apply dropout before passing through a layer\n\n        dec3 = self.upconv3(dec2)\n        if dec3.size()[2:] != enc2.size()[2:]:\n            dec3 = F.interpolate(dec3, size=enc2.shape[2:], mode='bilinear', align_corners=True)\n        dec3 = torch.cat((dec3, enc2), dim=1)\n        dec3 = self.conv3(dec3)\n        dec3 = self.dropout(dec3)  # Apply dropout before passing through a layer\n\n        dec4 = self.upconv4(dec3)\n        if dec4.size()[2:] != enc1.size()[2:]:\n            dec4 = F.interpolate(dec4, size=enc1.shape[2:], mode='bilinear', align_corners=True)\n        dec4 = torch.cat((dec4, enc1), dim=1)\n        dec4 = self.conv4(dec4)\n        dec4 = self.dropout(dec4)  # Apply dropout before passing through a layer\n\n        dec5 = self.upconv5(dec4)\n        if dec5.size()[2:] != x.size()[2:]:\n            dec5 = F.interpolate(dec5, size=x.shape[2:], mode='bilinear', align_corners=True)\n        dec5 = self.conv5(dec5)\n        dec5 = self.dropout(dec5)  # Apply dropout before passing through a layer\n\n        vessel = self.vessel_out(dec5)\n\n        return vessel\n\n","metadata":{"id":"Bq8qwlpN6QU7","outputId":"279b615e-8506-4efd-c6b5-8a2195b3c91a","papermill":{"duration":0.027133,"end_time":"2024-06-17T12:41:41.10288","exception":false,"start_time":"2024-06-17T12:41:41.075747","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-07-10T17:59:26.547212Z","iopub.status.idle":"2025-07-10T17:59:26.547496Z","shell.execute_reply.started":"2025-07-10T17:59:26.547363Z","shell.execute_reply":"2025-07-10T17:59:26.547377Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_loss(train_losses, val_losses):\n    plt.figure(figsize=(10, 5))\n    plt.plot(train_losses, label='Training Loss')\n    plt.plot(val_losses, label='Validation Loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.title('Training and Validation Loss')\n    plt.legend()\n    plt.show()\n\ndef train(net, train_loader, val_loader, criterion, optimizer, start_epoch=0, total_epochs=10,\n          accumulation_steps=4, gradient_clip=1.0, checkpoint_dir='/kaggle/working/',\n          scheduler=None, min_delta=0.001, patience=4):\n    \n    net.train().cuda()\n    scaler = GradScaler()\n    best_val_loss = float('inf')\n    early_stop_counter = 0\n    train_losses = []\n    val_losses = []\n\n    for epoch in range(start_epoch, total_epochs):\n        net.train()\n        epoch_loss = 0\n        optimizer.zero_grad()\n\n        for i, (images, masks) in enumerate(train_loader):\n            images, masks = images.cuda(), masks.cuda()\n\n            with autocast():\n                vessel_pred = net(images)\n                loss = criterion(vessel_pred, masks)\n            \n            scaler.scale(loss).backward()\n\n            if (i + 1) % accumulation_steps == 0:\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(net.parameters(), gradient_clip)\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n\n            epoch_loss += loss.item()\n\n        train_loss = epoch_loss / len(train_loader)\n        val_loss = validate(net, val_loader, criterion)\n\n        train_losses.append(train_loss)\n        val_losses.append(val_loss)\n\n        print(f'Epoch [{epoch + 1}/{total_epochs}], Training Loss: {train_loss:.6f}, Validation Loss: {val_loss:.6f}')\n\n        # Save best model\n        if val_loss + min_delta < best_val_loss:\n            best_val_loss = val_loss\n            early_stop_counter = 0\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': net.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict()\n            }, os.path.join(checkpoint_dir, 'best_model.pth'))\n        else:\n            early_stop_counter += 1\n            print(f\"EarlyStopping counter: {early_stop_counter} / {patience}\")\n\n        # Save checkpoint every epoch\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': net.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict()\n        }, os.path.join(checkpoint_dir, f'checkpoint_epoch_{epoch+1}.pth'))\n\n        if scheduler:\n            scheduler.step(val_loss)\n\n        if early_stop_counter >= patience:\n            print(f\"⏹️ Early stopping triggered at epoch {epoch + 1}\")\n            break\n\n    plot_loss(train_losses, val_losses)\n    print(\"✅ Training Complete\")\n\ndef validate(net, val_loader, criterion):\n    net.eval()\n    val_loss = 0\n    with torch.no_grad():\n        for images, masks in val_loader:\n            images, masks = images.cuda(), masks.cuda()\n            vessel_pred = net(images)\n            loss = criterion(vessel_pred, masks)\n            val_loss += loss.item()\n    return val_loss / len(val_loader)","metadata":{"id":"IMHsf1Wk6QVQ","papermill":{"duration":0.023018,"end_time":"2024-06-17T12:41:41.130118","exception":false,"start_time":"2024-06-17T12:41:41.1071","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-07-10T17:59:26.548605Z","iopub.status.idle":"2025-07-10T17:59:26.548915Z","shell.execute_reply.started":"2025-07-10T17:59:26.548781Z","shell.execute_reply":"2025-07-10T17:59:26.548796Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pip install torchmetrics\n","metadata":{"id":"O4S3O1lA6QVS","outputId":"0791b727-c38f-4082-c16e-b47ebfa81f70","papermill":{"duration":13.22388,"end_time":"2024-06-17T12:41:54.358069","exception":false,"start_time":"2024-06-17T12:41:41.134189","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-07-10T17:59:26.549722Z","iopub.status.idle":"2025-07-10T17:59:26.549988Z","shell.execute_reply.started":"2025-07-10T17:59:26.549859Z","shell.execute_reply":"2025-07-10T17:59:26.549872Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load the model and optimizer\nimport time\nnet = Net(weights= None).cuda()\noptimizer = torch.optim.Adam(net.parameters(), lr=1e-5, weight_decay=1e-5)\n\ntotal_epochs = 100\ncheckpoint_path = \"/kaggle/input/exp8-75checkpoints/checkpoint_epoch_75.pth\"\n\n# Load checkpoint if available\nif os.path.exists(checkpoint_path):\n    checkpoint = torch.load(checkpoint_path)\n    net.load_state_dict(checkpoint['model_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n    start_epoch = checkpoint['epoch']  # Resume from saved epoch\n    print(f\"Resuming training from epoch {start_epoch}\")\nelse:\n    start_epoch = 0  # Start from scratch if no checkpoint is found\n\n\n\n# Collect all image files from the training directory\ntrain_files = []\npatterns = [\n    \n    '/kaggle/input/kidney-3-upgraded/kidney_3_labels_combined/images/*.tif',\n    '/kaggle/input/blood-vessel-segmentation/train/kidney_1_dense/images/*.tif',\n    '/kaggle/input/blood-vessel-segmentation/train/kidney_1_voi/images/*.tif'\n]\n\nfor pattern in patterns:\n    train_files.extend(glob.glob(pattern, recursive=True))\n\ntrain_meta = [\n    DotDict(\n        file=train_files,\n        shape=(len(train_files), 512, 512),\n        id=[os.path.splitext(os.path.basename(f))[0] for f in train_files]\n    )\n]\n\n# Split the dataset into training and validation sets\ntrain_files, val_files = train_test_split(train_files, test_size=0.2, random_state=42)\n\n# Prepare training DataLoader with augmentation\ntrain_dataset = MyDataset(train_files, augment=True)\ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)\n\n# Prepare validation DataLoader without augmentation\nval_dataset = MyDataset(val_files)\nval_loader = DataLoader(val_dataset, batch_size=16, shuffle=False)\n\n\n\n\n#Learning Rate Scheduling: Utilize advanced learning rate schedulers like OneCycleLR or CosineAnnealingLR.\nscheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer, max_lr=0.001, steps_per_epoch=len(train_loader), epochs=total_epochs)\n\n\n# Define ComboLoss using dice_score from torchmetrics\nclass ComboLoss(nn.Module):\n    def __init__(self, weight_dice=0.5, weight_bce=0.5):\n        super(ComboLoss, self).__init__()\n        self.weight_dice = weight_dice\n        self.weight_bce = weight_bce\n        self.bce_loss = nn.BCEWithLogitsLoss()\n\n    def dice_loss(self, inputs, targets):\n        inputs = inputs.sigmoid().view(-1)\n        targets = targets.view(-1)\n        intersection = (inputs * targets).sum()\n        dice = (2. * intersection + 1e-5) / (inputs.sum() + targets.sum() + 1e-5)\n        return 1 - dice\n\n    def forward(self, inputs, targets):\n        dice = self.dice_loss(inputs, targets)\n        bce = self.bce_loss(inputs, targets)\n        return self.weight_dice * dice + self.weight_bce * bce\n\n\n\ncriterion = ComboLoss()\n\n# Continue training for additional epochs\n#train(net, train_loader, val_loader, criterion, optimizer, total_epochs=total_epochs, scheduler=scheduler)\n# train(net, train_loader, val_loader, criterion, optimizer, start_epoch=start_epoch, total_epochs=total_epochs, scheduler=scheduler)\ntrain(\n    net, \n    train_loader, \n    val_loader, \n    criterion, \n    optimizer, \n    start_epoch=start_epoch, \n    total_epochs=total_epochs, \n    scheduler=scheduler,\n    min_delta=0.001, \n    patience=4\n)","metadata":{"id":"ManVwMiY6QVT","outputId":"3452bb92-85c7-4d7e-f179-943aaf4c0d74","papermill":{"duration":15988.865731,"end_time":"2024-06-17T17:08:23.228801","exception":false,"start_time":"2024-06-17T12:41:54.36307","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2025-07-10T17:59:26.551489Z","iopub.status.idle":"2025-07-10T17:59:26.551805Z","shell.execute_reply.started":"2025-07-10T17:59:26.551661Z","shell.execute_reply":"2025-07-10T17:59:26.551677Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_metrics(y_true, y_pred):\n    y_true = np.nan_to_num(y_true)\n    y_true = (y_true > 0.5).astype(int)\n    y_pred = np.nan_to_num(y_pred)\n    y_pred = (y_pred > 0.5).astype(int)\n    y_true_flat = y_true.flatten()\n    y_pred_flat = y_pred.flatten()\n    \n    iou = jaccard_score(y_true_flat, y_pred_flat, average='binary')\n    dice = f1_score(y_true_flat, y_pred_flat, average='binary')\n    precision = precision_score(y_true_flat, y_pred_flat, average='binary')\n    recall = recall_score(y_true_flat, y_pred_flat, average='binary')\n\n    # Calculate accuracy\n    accuracy = np.mean(y_true_flat == y_pred_flat)\n\n    return iou, dice, precision, recall, accuracy\n\n\n\n\n\ndef evaluate_model(net, test_loader, checkpoint_dir=\"/kaggle/working/\"):\n    if checkpoint_dir:\n        checkpoint_path = os.path.join(checkpoint_dir, 'best_model.pth')\n        checkpoint = torch.load(checkpoint_path)\n        net.load_state_dict(checkpoint['model_state_dict'])\n        net.eval().cuda()\n\n    all_vessel_preds = []\n    all_vessel_masks = []\n\n    with torch.no_grad():\n        for images, masks in test_loader:\n            images, masks = images.cuda(), masks.cuda()\n            vessel_pred = net(images)  # Only get vessel predictions\n            vessel_pred = torch.sigmoid(vessel_pred).cpu().numpy()\n            vessel_masks = masks.cpu().numpy()\n\n            all_vessel_preds.append(vessel_pred)\n            all_vessel_masks.append(vessel_masks)\n\n    all_vessel_preds = np.concatenate(all_vessel_preds)\n    all_vessel_masks = np.concatenate(all_vessel_masks)\n\n    iou_vessel, dice_vessel, precision_vessel, recall_vessel, accuracy_vessel = evaluate_metrics(all_vessel_masks, all_vessel_preds)\n\n    return {\n        'iou_vessel': iou_vessel,\n        'dice_vessel': dice_vessel,\n        'precision_vessel': precision_vessel,\n        'recall_vessel': recall_vessel,\n        'accuracy_vessel': accuracy_vessel  # Include accuracy in the returned metrics\n    }\n\n\n\n\n# Collect all image files from the test directory\ntest_files = glob.glob('/kaggle/input/blood-vessel-segmentation/train/kidney_2/images/*.tif', recursive=True)\n# Prepare test DataLoader without augmentation\ntest_files = test_files[:500]\ntest_dataset = MyDataset(test_files)\ntest_loader = DataLoader(test_dataset, batch_size=8, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-10T17:59:26.552869Z","iopub.status.idle":"2025-07-10T17:59:26.553161Z","shell.execute_reply.started":"2025-07-10T17:59:26.553021Z","shell.execute_reply":"2025-07-10T17:59:26.553036Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"metrics = evaluate_model(net, test_loader)\nprint(metrics)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-10T17:59:26.554545Z","iopub.status.idle":"2025-07-10T17:59:26.554981Z","shell.execute_reply.started":"2025-07-10T17:59:26.554762Z","shell.execute_reply":"2025-07-10T17:59:26.554784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch\n# from torch.utils.data import DataLoader\n# import numpy as np\n# import matplotlib.pyplot as plt\n# import os\n\n# def get_sample_from_dataset(dataset, idx):\n#     image_tensor, mask = dataset[idx]\n#     return image_tensor.unsqueeze(0), mask.numpy()\n\n# def integrated_gradients(inputs, model, target_class=0, baseline=None, steps=50, device='cuda'):\n#     if baseline is None:\n#         baseline = torch.zeros_like(inputs).to(device)\n#     inputs = inputs.to(device)\n#     baseline = baseline.to(device)\n\n#     scaled_inputs = [baseline + (float(i) / steps) * (inputs - baseline) for i in range(steps + 1)]\n\n#     grads = []\n#     for scaled_input in scaled_inputs:\n#         scaled_input.requires_grad = True\n#         output = model(scaled_input)\n#         score = output[0, target_class].mean()\n#         model.zero_grad()\n#         score.backward(retain_graph=True)\n#         grads.append(scaled_input.grad.detach().cpu().numpy())\n\n#     grads = np.array(grads)\n#     avg_grads = (grads[:-1] + grads[1:]) / 2.0\n#     avg_grads = avg_grads.mean(axis=0).squeeze()\n\n#     integrated_grads = (inputs.cpu().numpy().squeeze() - baseline.cpu().numpy().squeeze()) * avg_grads\n\n#     if integrated_grads.ndim == 3:\n#         integrated_grads = integrated_grads[0]\n\n#     integrated_grads = np.abs(integrated_grads)\n#     integrated_grads -= integrated_grads.min()\n#     integrated_grads /= integrated_grads.max() + 1e-8\n\n#     # ✅ Apply gamma correction\n#     gamma = 0.5\n#     integrated_grads = integrated_grads ** gamma\n\n#     return integrated_grads\n\n# def save_visualizations(image_tensor, mask, attribution_map, idx, output_dir='output'):\n#     image_np = image_tensor.squeeze().cpu().numpy()\n#     mask_2d = mask.squeeze()\n\n#     folder = os.path.join(output_dir, f\"image_{idx}\")\n#     os.makedirs(folder, exist_ok=True)\n\n#     # Save input image\n#     plt.imshow(image_np, cmap='gray')\n#     plt.axis('off')\n#     plt.savefig(os.path.join(folder, f\"image_{idx}_input.png\"), bbox_inches='tight', pad_inches=0)\n#     plt.close()\n\n#     # Save mask image\n#     plt.imshow(mask_2d, cmap='gray')\n#     plt.axis('off')\n#     plt.savefig(os.path.join(folder, f\"image_{idx}_mask.png\"), bbox_inches='tight', pad_inches=0)\n#     plt.close()\n\n#     # Save attribution overlay with plasma colormap\n#     fig, ax = plt.subplots()\n#     heatmap = ax.imshow(image_np, cmap='gray')\n#     overlay = ax.imshow(attribution_map, cmap='plasma', alpha=0.5)\n#     plt.axis('off')\n#     cbar = plt.colorbar(overlay, ax=ax, fraction=0.046, pad=0.04)\n#     cbar.set_label('Attribution Intensity', rotation=270, labelpad=15)\n#     plt.savefig(os.path.join(folder, f\"image_{idx}_integrated_gradients.png\"), bbox_inches='tight', pad_inches=0)\n#     plt.close()\n\n\n# # Main loop\n# net.eval()\n# device = 'cuda' if torch.cuda.is_available() else 'cpu'\n# net.to(device)\n# num_images = 100\n\n# import random\n# indices = random.sample(range(len(test_dataset)), num_images)\n# for idx in indices:\n#     image_tensor, mask = get_sample_from_dataset(test_dataset, idx)\n#     image_tensor = image_tensor.to(device)\n\n#     attribution_map = integrated_gradients(image_tensor, net, target_class=0, steps=100, device=device)\n\n#     print(f\"Image {idx} - Attribution min:\", attribution_map.min())\n#     print(f\"Image {idx} - Attribution max:\", attribution_map.max())\n\n#     plt.hist(attribution_map.ravel(), bins=50)\n#     plt.title(f\"Attribution Histogram for Image {idx}\")\n#     plt.xlabel(\"Attribution Value\")  # ✅ Axis label\n#     plt.ylabel(\"Pixel Count\")        # ✅ Axis label\n#     plt.show()\n\n#     save_visualizations(image_tensor.cpu(), mask, attribution_map, idx)\n\n\n\n\n\n\n# print(f\"Saved visualizations for {num_images} images under './output/' folder.\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-07-10T17:59:26.556935Z","iopub.status.idle":"2025-07-10T17:59:26.557364Z","shell.execute_reply.started":"2025-07-10T17:59:26.557143Z","shell.execute_reply":"2025-07-10T17:59:26.557166Z"}},"outputs":[],"execution_count":null}]}