{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"}],"dockerImageVersionId":30715,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"###-----\n'''\nTODO:\nAugmentation #1, #2\nMetrices\n'''","metadata":{"execution":{"iopub.status.busy":"2024-06-01T19:37:41.954432Z","iopub.execute_input":"2024-06-01T19:37:41.955070Z","iopub.status.idle":"2024-06-01T19:37:41.968931Z","shell.execute_reply.started":"2024-06-01T19:37:41.955034Z","shell.execute_reply":"2024-06-01T19:37:41.967989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import 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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-06-01T22:28:04.528316Z","iopub.execute_input":"2024-06-01T22:28:04.528707Z","iopub.status.idle":"2024-06-01T22:28:04.536569Z","shell.execute_reply.started":"2024-06-01T22:28:04.528677Z","shell.execute_reply":"2024-06-01T22:28:04.535564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyDataset(Dataset):\n    def __init__(self, meta, detailed_meta: Dict[str, str], image_size=(512, 512), augment=False):\n        self.meta = meta\n        self.detailed_meta = detailed_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            T.ToTensor()\n        ]) if augment else None\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        detailed_maskfile = self.detailed_meta.get(os.path.basename(f))\n        if detailed_maskfile and os.path.exists(detailed_maskfile):\n            maskfile = detailed_maskfile\n        else:\n            maskfile = f.replace('images', 'labels')\n\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    \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":{"execution":{"iopub.status.busy":"2024-06-01T22:27:06.161221Z","iopub.execute_input":"2024-06-01T22:27:06.161898Z","iopub.status.idle":"2024-06-01T22:27:06.174723Z","shell.execute_reply.started":"2024-06-01T22:27:06.161865Z","shell.execute_reply":"2024-06-01T22:27:06.173773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-06-01T22:27:10.305922Z","iopub.execute_input":"2024-06-01T22:27:10.306758Z","iopub.status.idle":"2024-06-01T22:27:10.314235Z","shell.execute_reply.started":"2024-06-01T22:27:10.306722Z","shell.execute_reply":"2024-06-01T22:27:10.313193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-06-01T22:27:12.614887Z","iopub.execute_input":"2024-06-01T22:27:12.615521Z","iopub.status.idle":"2024-06-01T22:27:12.621490Z","shell.execute_reply.started":"2024-06-01T22:27:12.615482Z","shell.execute_reply":"2024-06-01T22:27:12.620565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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":{"execution":{"iopub.status.busy":"2024-06-01T22:27:14.913684Z","iopub.execute_input":"2024-06-01T22:27:14.914032Z","iopub.status.idle":"2024-06-01T22:27:14.920631Z","shell.execute_reply.started":"2024-06-01T22:27:14.914007Z","shell.execute_reply":"2024-06-01T22:27:14.919635Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Net(nn.Module):\n    def __init__(self, weights='IMAGENET1K_V1'):\n    #def __init__(self, weights=None):\n\n        super(Net, self).__init__()\n        \n        self.backbone = models.resnet50(weights=weights)\n       \n        #self.initialize_weights(self.backbone)\n        \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        self.kidney_out = nn.Conv2d(32, 1, kernel_size=1)\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        \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        \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        \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        \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        \n        vessel = self.vessel_out(dec5)\n        kidney = self.kidney_out(dec5)\n        \n        return vessel, kidney\n    \n    '''def initialize_weights(self, model):\n        for m in model.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n'''\nprint(\"NET OK!!!\")","metadata":{"execution":{"iopub.status.busy":"2024-06-01T22:27:20.367551Z","iopub.execute_input":"2024-06-01T22:27:20.367880Z","iopub.status.idle":"2024-06-01T22:27:20.388896Z","shell.execute_reply.started":"2024-06-01T22:27:20.367855Z","shell.execute_reply":"2024-06-01T22:27:20.387912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Added schedualer and loss plotting****","metadata":{}},{"cell_type":"code","source":"\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            if i == 0:\n                show_images(images, masks, n=1)  # Assuming show_images is defined elsewhere\n            with autocast():\n                vessel_pred, kidney_pred = net(images)\n                if torch.isnan(vessel_pred).any() or torch.isnan(kidney_pred).any():\n                    print(\"NaNs found in model outputs\")\n                    continue\n                loss_vessel = criterion(vessel_pred, masks)\n                loss_kidney = criterion(kidney_pred, masks)\n                loss = loss_vessel + loss_kidney\n            scaler.scale(loss).backward()\n            if (i + 1) % accumulation_steps == 0:\n                scaler.unscale_(optimizer)\n                # Check for NaNs in gradients\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            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}, 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)\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, kidney_pred = net(images)\n            loss_vessel = criterion(vessel_pred, masks)\n            loss_kidney = criterion(kidney_pred, masks)\n            loss = loss_vessel + loss_kidney\n            val_loss += loss.item()\n    return val_loss / len(val_loader)\n\ndef plot_loss(train_loss, val_loss):\n    plt.figure(figsize=(10, 5))\n    plt.plot(train_loss, label='Training Loss')\n    plt.plot(val_loss, label='Validation Loss')\n    plt.title('Training and Validation Loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.show()\n\n\n","metadata":{"execution":{"iopub.status.busy":"2024-06-01T19:37:48.882814Z","iopub.execute_input":"2024-06-01T19:37:48.883174Z","iopub.status.idle":"2024-06-01T19:37:48.908204Z","shell.execute_reply.started":"2024-06-01T19:37:48.883141Z","shell.execute_reply":"2024-06-01T19:37:48.907275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load the model and optimizer\nnet = Net(weights='IMAGENET1K_V1').cuda()\noptimizer = torch.optim.Adam(net.parameters(), lr=1e-5, weight_decay=1e-5)\nscheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=5)\n\n# Collect all image files from the training directory\ntrain_files = []\npatterns = [\n    '/kaggle/input/blood-vessel-segmentation/train/kidney_3_sparse/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# Create a dictionary for detailed masks\ndetailed_masks = {}\ndetailed_mask_patterns = [\n    '/kaggle/input/blood-vessel-segmentation/train/kidney_3_dense/labels/*.tif'\n]\n\nfor pattern in detailed_mask_patterns:\n    for maskfile in glob.glob(pattern, recursive=True):\n        basename = os.path.basename(maskfile)\n        detailed_masks[basename] = maskfile\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, detailed_masks, augment=True)\ntrain_loader = DataLoader(train_dataset, batch_size=16, shuffle=True)\n\n# Prepare validation DataLoader without augmentation\nval_dataset = MyDataset(val_files, detailed_masks)\nval_loader = DataLoader(val_dataset, batch_size=16, shuffle=False)\n\ntest_files = glob.glob('/kaggle/input/blood-vessel-segmentation/train/kidney_2/images/*.tif', recursive=True)\ntest_dataset = MyDataset(test_files, detailed_masks)\ntest_loader = DataLoader(test_dataset, batch_size=8, shuffle=False)\n\ncriterion = nn.BCEWithLogitsLoss()\n\n# Define the number of epochs you want to train for\ntotal_epochs = 10\n\n# Continue training for additional epochs\n#train(net, train_loader, val_loader, criterion, optimizer, total_epochs=total_epochs, scheduler=scheduler)\n","metadata":{"execution":{"iopub.status.busy":"2024-06-01T22:28:31.822902Z","iopub.execute_input":"2024-06-01T22:28:31.823254Z","iopub.status.idle":"2024-06-01T22:28:32.868701Z","shell.execute_reply.started":"2024-06-01T22:28:31.823225Z","shell.execute_reply":"2024-06-01T22:28:32.867834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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    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    return iou, dice, precision, recall\n\nimport os\n\ndef evaluate_model(net, test_loader, checkpoint_dir=None):\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    all_kidney_preds = []\n    all_kidney_masks = []\n\n    with torch.no_grad():\n        for images, masks in test_loader:\n            images, masks = images.cuda(), masks.cuda()\n            vessel_pred, kidney_pred = net(images)\n            vessel_pred = torch.sigmoid(vessel_pred).cpu().numpy()\n            kidney_pred = torch.sigmoid(kidney_pred).cpu().numpy()\n            vessel_masks = masks.cpu().numpy()\n            kidney_masks = masks.cpu().numpy()\n\n            all_vessel_preds.append(vessel_pred)\n            all_vessel_masks.append(vessel_masks)\n            all_kidney_preds.append(kidney_pred)\n            all_kidney_masks.append(kidney_masks)\n\n    all_vessel_preds = np.concatenate(all_vessel_preds)\n    all_vessel_masks = np.concatenate(all_vessel_masks)\n    all_kidney_preds = np.concatenate(all_kidney_preds)\n    all_kidney_masks = np.concatenate(all_kidney_masks)\n\n    iou_vessel, dice_vessel, precision_vessel, recall_vessel = evaluate_metrics(all_vessel_masks, all_vessel_preds)\n    iou_kidney, dice_kidney, precision_kidney, recall_kidney = evaluate_metrics(all_kidney_masks, all_kidney_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        'iou_kidney': iou_kidney,\n        'dice_kidney': dice_kidney,\n        'precision_kidney': precision_kidney,\n        'recall_kidney': recall_kidney\n    }\n","metadata":{"execution":{"iopub.status.busy":"2024-06-01T22:28:41.974934Z","iopub.execute_input":"2024-06-01T22:28:41.975285Z","iopub.status.idle":"2024-06-01T22:28:41.990115Z","shell.execute_reply.started":"2024-06-01T22:28:41.975257Z","shell.execute_reply":"2024-06-01T22:28:41.989042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"evaluate_model(net,test_loader)","metadata":{"execution":{"iopub.status.busy":"2024-06-01T22:28:46.243532Z","iopub.execute_input":"2024-06-01T22:28:46.243942Z"},"trusted":true},"execution_count":null,"outputs":[]}]}