{"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}],"dockerImageVersionId":30805,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":13624.869403,"end_time":"2024-09-17T09:51:12.381189","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2024-09-17T06:04:07.511786","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"97c78838","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\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","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","id":"Ft2GG7i26QTN","papermill":{"duration":7.052563,"end_time":"2024-09-17T06:04:17.329742","exception":false,"start_time":"2024-09-17T06:04:10.277179","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T20:14:21.240134Z","iopub.execute_input":"2024-12-13T20:14:21.240818Z","iopub.status.idle":"2024-12-13T20:14:26.616440Z","shell.execute_reply.started":"2024-12-13T20:14:21.240785Z","shell.execute_reply":"2024-12-13T20:14:26.615721Z"}},"outputs":[],"execution_count":null},{"id":"7a0d3bb3","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.019149,"end_time":"2024-09-17T06:04:17.353147","exception":false,"start_time":"2024-09-17T06:04:17.333998","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T20:14:26.617955Z","iopub.execute_input":"2024-12-13T20:14:26.618331Z","iopub.status.idle":"2024-12-13T20:14:26.628164Z","shell.execute_reply.started":"2024-12-13T20:14:26.618305Z","shell.execute_reply":"2024-12-13T20:14:26.627249Z"}},"outputs":[],"execution_count":null},{"id":"9d9835c4","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.012307,"end_time":"2024-09-17T06:04:17.386092","exception":false,"start_time":"2024-09-17T06:04:17.373785","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T20:14:26.629258Z","iopub.execute_input":"2024-12-13T20:14:26.629510Z","iopub.status.idle":"2024-12-13T20:14:26.651378Z","shell.execute_reply.started":"2024-12-13T20:14:26.629486Z","shell.execute_reply":"2024-12-13T20:14:26.650444Z"}},"outputs":[],"execution_count":null},{"id":"91701145","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.013363,"end_time":"2024-09-17T06:04:17.370212","exception":false,"start_time":"2024-09-17T06:04:17.356849","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T20:14:26.653365Z","iopub.execute_input":"2024-12-13T20:14:26.653632Z","iopub.status.idle":"2024-12-13T20:14:26.661313Z","shell.execute_reply.started":"2024-12-13T20:14:26.653606Z","shell.execute_reply":"2024-12-13T20:14:26.660608Z"}},"outputs":[],"execution_count":null},{"id":"086e1d92","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.011996,"end_time":"2024-09-17T06:04:17.401871","exception":false,"start_time":"2024-09-17T06:04:17.389875","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T20:14:26.662224Z","iopub.execute_input":"2024-12-13T20:14:26.662458Z","iopub.status.idle":"2024-12-13T20:14:26.677224Z","shell.execute_reply.started":"2024-12-13T20:14:26.662434Z","shell.execute_reply":"2024-12-13T20:14:26.676311Z"}},"outputs":[],"execution_count":null},{"id":"b4968efa","cell_type":"code","source":"\nclass UNet(nn.Module):\n    def __init__(self, in_channels=1, out_channels=1):\n        super(UNet, self).__init__()\n        \n        def conv_block(in_channels, out_channels):\n            return nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n                nn.BatchNorm2d(out_channels),\n                nn.ReLU(inplace=True),\n                nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n                nn.BatchNorm2d(out_channels),\n                nn.ReLU(inplace=True)\n            )\n        \n        self.encoder1 = conv_block(in_channels, 64)\n        self.encoder2 = conv_block(64, 128)\n        self.encoder3 = conv_block(128, 256)\n        self.encoder4 = conv_block(256, 512)\n        \n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        self.bottleneck = conv_block(512, 1024)\n        \n        self.upconv4 = nn.ConvTranspose2d(1024, 512, kernel_size=2, stride=2)\n        self.decoder4 = conv_block(1024, 512)\n        \n        self.upconv3 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2)\n        self.decoder3 = conv_block(512, 256)\n        \n        self.upconv2 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2)\n        self.decoder2 = conv_block(256, 128)\n        \n        self.upconv1 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)\n        self.decoder1 = conv_block(128, 64)\n        \n        self.conv_last = nn.Conv2d(64, out_channels, kernel_size=1)\n    \n    def forward(self, x):\n        enc1 = self.encoder1(x)\n        enc2 = self.encoder2(self.pool(enc1))\n        enc3 = self.encoder3(self.pool(enc2))\n        enc4 = self.encoder4(self.pool(enc3))\n        \n        bottleneck = self.bottleneck(self.pool(enc4))\n        \n        dec4 = self.upconv4(bottleneck)\n        dec4 = self.decoder4(torch.cat((dec4, enc4), dim=1))\n        \n        dec3 = self.upconv3(dec4)\n        dec3 = self.decoder3(torch.cat((dec3, enc3), dim=1))\n        \n        dec2 = self.upconv2(dec3)\n        dec2 = self.decoder2(torch.cat((dec2, enc2), dim=1))\n        \n        dec1 = self.upconv1(dec2)\n        dec1 = self.decoder1(torch.cat((dec1, enc1), dim=1))\n        \n        return self.conv_last(dec1)\n","metadata":{"id":"Bq8qwlpN6QU7","outputId":"279b615e-8506-4efd-c6b5-8a2195b3c91a","papermill":{"duration":0.026658,"end_time":"2024-09-17T06:04:17.432179","exception":false,"start_time":"2024-09-17T06:04:17.405521","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T20:14:26.678146Z","iopub.execute_input":"2024-12-13T20:14:26.678405Z","iopub.status.idle":"2024-12-13T20:14:26.689086Z","shell.execute_reply.started":"2024-12-13T20:14:26.678365Z","shell.execute_reply":"2024-12-13T20:14:26.688238Z"}},"outputs":[],"execution_count":null},{"id":"cec13cef","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}], 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.029387,"end_time":"2024-09-17T06:04:17.466577","exception":false,"start_time":"2024-09-17T06:04:17.437190","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T20:14:26.690282Z","iopub.execute_input":"2024-12-13T20:14:26.690567Z","iopub.status.idle":"2024-12-13T20:14:26.704432Z","shell.execute_reply.started":"2024-12-13T20:14:26.690527Z","shell.execute_reply":"2024-12-13T20:14:26.703712Z"}},"outputs":[],"execution_count":null},{"id":"b69ce88d","cell_type":"code","source":"pip install torchmetrics\n","metadata":{"id":"O4S3O1lA6QVS","outputId":"0791b727-c38f-4082-c16e-b47ebfa81f70","papermill":{"duration":13.948444,"end_time":"2024-09-17T06:04:31.418749","exception":false,"start_time":"2024-09-17T06:04:17.470305","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T20:14:26.705214Z","iopub.execute_input":"2024-12-13T20:14:26.705437Z","iopub.status.idle":"2024-12-13T20:14:36.028439Z","shell.execute_reply.started":"2024-12-13T20:14:26.705414Z","shell.execute_reply":"2024-12-13T20:14:36.027370Z"}},"outputs":[],"execution_count":null},{"id":"442362bb","cell_type":"code","source":"# Load the model and optimizer\nnet = UNet(in_channels=1, out_channels=1).cuda()\noptimizer = torch.optim.Adam(net.parameters(), lr=1e-4, weight_decay=1e-4)\ntotal_epochs = 25\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\nimport time\nstart_time = time.time()\n\n# Continue training for additional epochs\ntrain(net, train_loader, val_loader, criterion, optimizer, total_epochs=total_epochs, scheduler=scheduler)\n\nend_time = time.time()\n\ntime_spent = end_time - start_time\nprint(f\"Time spent by model: {time_spent} seconds\")","metadata":{"id":"ManVwMiY6QVT","outputId":"3452bb92-85c7-4d7e-f179-943aaf4c0d74","papermill":{"duration":13597.535247,"end_time":"2024-09-17T09:51:08.958558","exception":false,"start_time":"2024-09-17T06:04:31.423311","status":"completed"},"tags":[],"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T20:14:36.030143Z","iopub.execute_input":"2024-12-13T20:14:36.030897Z"}},"outputs":[],"execution_count":null},{"id":"b6610908-0102-467b-8094-2f22bcb1bd59","cell_type":"markdown","source":"","metadata":{}},{"id":"6d5c03f2-04eb-46ff-a0b9-1a6baf0163fa","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\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 = 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    }\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},{"id":"a0892186-172f-4d61-9d51-74c263da8b81","cell_type":"code","source":"metrics = evaluate_model(net, test_loader)\nprint(metrics)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}