{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":8677798,"sourceType":"datasetVersion","datasetId":5201742},{"sourceId":10529254,"sourceType":"datasetVersion","datasetId":6516179},{"sourceId":12402910,"sourceType":"datasetVersion","datasetId":7821624}],"dockerImageVersionId":30805,"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":"\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":{},"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":[],"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":[],"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":[],"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":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Net(nn.Module):\n    def __init__(self, weights='IMAGENET1K_V1'):\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":[],"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, accumulation_steps=4, gradient_clip=1.0, checkpoint_dir='/kaggle/working/', scheduler=None):\n    net.train().cuda()\n    scaler = GradScaler()\n    best_val_loss = float('inf')\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            if torch.isnan(images).any() or torch.isnan(masks).any():\n                print(\"NaNs found in input data\")\n                continue\n\n            with autocast():\n                vessel_pred = net(images)\n                if torch.isnan(vessel_pred).any():\n                    print(\"NaNs found in model outputs\")\n                    continue\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                has_nan = False\n                for param in net.parameters():\n                    if param.grad is not None and torch.isnan(param.grad).any():\n                        has_nan = True\n                        break\n                if has_nan:\n                    print(f\"NaNs found in gradients at epoch {epoch + 1}, step {i + 1}\")\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)  # Calculate average training loss for the epoch\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}, Validation Loss: {val_loss}')\n\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save({'epoch': epoch, 'model_state_dict': net.state_dict(), 'optimizer_state_dict': optimizer.state_dict()}, os.path.join(checkpoint_dir, 'best_model.pth'))\n\n        torch.save({'epoch': epoch, 'model_state_dict': net.state_dict(), 'optimizer_state_dict': optimizer.state_dict()}, os.path.join(checkpoint_dir, f'checkpoint_epoch_{epoch+1}.pth'))\n\n        if scheduler:\n            scheduler.step(val_loss)\n\n    plot_loss(train_losses, val_losses)  # Plot both training and validation losses\n\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":[],"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":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load the model and optimizer\nimport time\nnet = Net(weights='IMAGENET1K_V1').cuda()\noptimizer = torch.optim.Adam(net.parameters(), lr=1e-5, weight_decay=1e-5)\n\ntotal_epochs = 51  # Continue training until epoch 30\ncheckpoint_path = \"/kaggle/input/exp-7-checkpoints30/checkpoint_epoch_30 (1).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)\ntrain(net, train_loader, val_loader, criterion, optimizer, start_epoch=start_epoch, total_epochs=total_epochs, scheduler=scheduler)\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":[],"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\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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"metrics = evaluate_model(net, test_loader)\nprint(metrics)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}