{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":12015593,"sourceType":"datasetVersion","datasetId":7559418},{"sourceId":242890618,"sourceType":"kernelVersion"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as transforms\nfrom PIL import Image\nfrom numpy import linalg as LA\nfrom torch import optim, nn\nfrom torch.utils.data import DataLoader, random_split\nfrom torch.utils.data.dataset import Dataset\nfrom torchvision import transforms\nfrom tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:30.693091Z","iopub.execute_input":"2025-05-31T16:42:30.693856Z","iopub.status.idle":"2025-05-31T16:42:30.698543Z","shell.execute_reply.started":"2025-05-31T16:42:30.693833Z","shell.execute_reply":"2025-05-31T16:42:30.697786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SegmentationDataset(Dataset):\n    def __init__(self, root_dir, image_folder='data', mask_folder='mask', transform=None, mask_transform=None):\n        self.root_dir = root_dir\n        self.image_folder = os.path.join(root_dir, image_folder)\n        self.mask_folder = os.path.join(root_dir, mask_folder)\n        self.image_filenames = sorted([\n            fname for fname in os.listdir(self.image_folder) if fname.endswith('.jpg')\n        ])\n        self.transform = transform\n        self.mask_transform = mask_transform\n\n    def __len__(self):\n        return len(self.image_filenames)\n\n    def __getitem__(self, idx):\n        image_name = self.image_filenames[idx]\n        image_path = os.path.join(self.image_folder, image_name)\n        \n        # Mask file assumed to be named with \"_mask\" suffix\n        mask_name = image_name.replace('.jpg', '_mask.gif')\n        mask_path = os.path.join(self.mask_folder, mask_name)\n\n        image = Image.open(image_path).convert(\"RGB\")\n        mask = Image.open(mask_path).convert(\"L\")  # Grayscale mask (class per pixel)\n\n        if self.transform:\n            image = self.transform(image)\n        if self.mask_transform:\n            mask = self.mask_transform(mask)\n\n        return image, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:30.699779Z","iopub.execute_input":"2025-05-31T16:42:30.699999Z","iopub.status.idle":"2025-05-31T16:42:30.717160Z","shell.execute_reply.started":"2025-05-31T16:42:30.699971Z","shell.execute_reply":"2025-05-31T16:42:30.716488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import random_split\n\n# Dataset root and transforms\nroot_dir = '/kaggle/input/carvana-dataset'\n\nimg_transform = transforms.Compose([\n    transforms.Resize((512, 512)),\n    transforms.ToTensor()\n])\nmask_transform = transforms.Compose([\n    transforms.Resize((512, 512)),\n    transforms.ToTensor()\n])\n\n# Load full dataset\nfull_dataset = SegmentationDataset(root_dir, transform=img_transform, mask_transform=mask_transform)\n\n# Compute split lengths\ntrain_len = int(0.8 * len(full_dataset))\nval_len = len(full_dataset) - train_len\n\n# Random split\ntrain_dataset, val_dataset = random_split(full_dataset, [train_len, val_len])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:30.717802Z","iopub.execute_input":"2025-05-31T16:42:30.717997Z","iopub.status.idle":"2025-05-31T16:42:30.754649Z","shell.execute_reply.started":"2025-05-31T16:42:30.717983Z","shell.execute_reply":"2025-05-31T16:42:30.753980Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef show_batch_grid(dataset, num_pairs=6):\n    assert num_pairs % 3 == 0, \"Use a multiple of 3 for clean rows\"\n\n    rows = (num_pairs // 3) * 2\n    cols = 6\n\n    fig, axes = plt.subplots(rows, cols, figsize=(18, rows * 2.5))\n\n    for ax in axes.flat:\n        ax.axis('off')  # Turn off everything first\n\n    for i in range(num_pairs):\n        img, mask = dataset[i]\n        img_np = img.permute(1, 2, 0).numpy()\n        mask_np = mask.squeeze().numpy()\n\n        row = (i // 3) * 2\n        col = (i % 3) * 2\n\n        axes[row, col].imshow(img_np)\n        axes[row, col].set_title(f\"Image {i+1}\")\n        axes[row, col].axis('off')\n\n        axes[row, col+1].imshow(mask_np, cmap='gray')\n        axes[row, col+1].set_title(f\"Mask {i+1}\")\n        axes[row, col+1].axis('off')\n\n    plt.tight_layout()\n    plt.show()\n\nshow_batch_grid(train_dataset)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:30.756106Z","iopub.execute_input":"2025-05-31T16:42:30.756440Z","iopub.status.idle":"2025-05-31T16:42:32.562855Z","shell.execute_reply.started":"2025-05-31T16:42:30.756425Z","shell.execute_reply":"2025-05-31T16:42:32.562093Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\nbatch_size = 8\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=batch_size,\n    shuffle=True,       # Shuffle for training\n    num_workers=4,      # Adjust based on your system\n    pin_memory=True     # Recommended if using GPU\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=batch_size,\n    shuffle=False,      # Don't shuffle for validation\n    num_workers=4,\n    pin_memory=True\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:32.563703Z","iopub.execute_input":"2025-05-31T16:42:32.563949Z","iopub.status.idle":"2025-05-31T16:42:32.568894Z","shell.execute_reply.started":"2025-05-31T16:42:32.563929Z","shell.execute_reply":"2025-05-31T16:42:32.568142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv_op = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.conv_op(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:32.569546Z","iopub.execute_input":"2025-05-31T16:42:32.569788Z","iopub.status.idle":"2025-05-31T16:42:32.587910Z","shell.execute_reply.started":"2025-05-31T16:42:32.569773Z","shell.execute_reply":"2025-05-31T16:42:32.587185Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DownSample(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv = DoubleConv(in_channels, out_channels)\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)\n\n    def forward(self, x):\n        down = self.conv(x)\n        p = self.pool(down)\n\n        return down, p","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:32.588786Z","iopub.execute_input":"2025-05-31T16:42:32.589011Z","iopub.status.idle":"2025-05-31T16:42:32.602573Z","shell.execute_reply.started":"2025-05-31T16:42:32.588988Z","shell.execute_reply":"2025-05-31T16:42:32.601944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UpSample(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.up = nn.ConvTranspose2d(in_channels, in_channels//2, kernel_size=2, stride=2)\n        self.conv = DoubleConv(in_channels, out_channels)\n\n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        x = torch.cat([x1, x2], dim = 1)\n        return self.conv(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:32.604433Z","iopub.execute_input":"2025-05-31T16:42:32.604599Z","iopub.status.idle":"2025-05-31T16:42:32.625204Z","shell.execute_reply.started":"2025-05-31T16:42:32.604587Z","shell.execute_reply":"2025-05-31T16:42:32.624627Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, in_channels, num_classes):\n        super().__init__()\n        self.down_convolution_1 = DownSample(in_channels, 64)\n        self.down_convolution_2 = DownSample(64, 128)\n        self.down_convolution_3 = DownSample(128, 256)\n        self.down_convolution_4 = DownSample(256, 512)\n\n        self.bottle_neck = DoubleConv(512, 1024)\n\n        self.up_convolution_1 = UpSample(1024, 512)\n        self.up_convolution_2 = UpSample(512, 256)\n        self.up_convolution_3 = UpSample(256, 128)\n        self.up_convolution_4 = UpSample(128, 64)\n\n        self.out = nn.Conv2d(in_channels=64, out_channels=num_classes, kernel_size=1)\n\n    def forward(self, x):\n        down_1, p1 = self.down_convolution_1(x)\n        down_2, p2 = self.down_convolution_2(p1)\n        down_3, p3 = self.down_convolution_3(p2)\n        down_4, p4 = self.down_convolution_4(p3)\n\n        b = self.bottle_neck(p4)\n\n        up_1 = self.up_convolution_1(b, down_4)\n        up_2 = self.up_convolution_2(up_1, down_3)\n        up_3 = self.up_convolution_3(up_2, down_2)\n        up_4 = self.up_convolution_4(up_3, down_1)\n\n        out = self.out(up_4)\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:32.625837Z","iopub.execute_input":"2025-05-31T16:42:32.626054Z","iopub.status.idle":"2025-05-31T16:42:32.641013Z","shell.execute_reply.started":"2025-05-31T16:42:32.626040Z","shell.execute_reply":"2025-05-31T16:42:32.640415Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:32.641758Z","iopub.execute_input":"2025-05-31T16:42:32.641990Z","iopub.status.idle":"2025-05-31T16:42:32.656892Z","shell.execute_reply.started":"2025-05-31T16:42:32.641971Z","shell.execute_reply":"2025-05-31T16:42:32.656232Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = UNet(in_channels=3, num_classes=1)\nmodel = model.to(device)  \nif torch.cuda.device_count() > 1:\n    print(f\"Using {torch.cuda.device_count()} GPUs with DataParallel!\")\n    model = nn.DataParallel(model)\n\nfor data, label in train_loader:\n    data = data.to(device)  # move data to device\n    output = model(data)\n    print(output.shape)  # Expected: (batch_size, 1, 256, 256)\n    break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:32.657549Z","iopub.execute_input":"2025-05-31T16:42:32.657739Z","iopub.status.idle":"2025-05-31T16:42:34.659266Z","shell.execute_reply.started":"2025-05-31T16:42:32.657724Z","shell.execute_reply":"2025-05-31T16:42:34.658552Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 25\nLEARNING_RATE = 1e-4","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:34.660465Z","iopub.execute_input":"2025-05-31T16:42:34.660704Z","iopub.status.idle":"2025-05-31T16:42:34.664678Z","shell.execute_reply.started":"2025-05-31T16:42:34.660681Z","shell.execute_reply":"2025-05-31T16:42:34.664047Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE)\ncriterion = nn.BCEWithLogitsLoss()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:34.665545Z","iopub.execute_input":"2025-05-31T16:42:34.665837Z","iopub.status.idle":"2025-05-31T16:42:34.681384Z","shell.execute_reply.started":"2025-05-31T16:42:34.665815Z","shell.execute_reply":"2025-05-31T16:42:34.680665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dice_coefficient(prediction, target, epsilon=1e-07):\n    prediction_copy = prediction.clone()\n\n    prediction_copy[prediction_copy < 0] = 0\n    prediction_copy[prediction_copy > 0] = 1\n\n    intersection = abs(torch.sum(prediction_copy * target))\n    union = abs(torch.sum(prediction_copy) + torch.sum(target))\n    dice = (2. * intersection + epsilon) / (union + epsilon)\n    \n    return dice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:34.682399Z","iopub.execute_input":"2025-05-31T16:42:34.682637Z","iopub.status.idle":"2025-05-31T16:42:34.694715Z","shell.execute_reply.started":"2025-05-31T16:42:34.682617Z","shell.execute_reply":"2025-05-31T16:42:34.694082Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_losses = []\ntrain_dcs = []\nval_losses = []\nval_dcs = []\nbest_val_dice = 0.0\nbest_model_path = \"best_model.pth\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:34.695373Z","iopub.execute_input":"2025-05-31T16:42:34.695532Z","iopub.status.idle":"2025-05-31T16:42:34.709268Z","shell.execute_reply.started":"2025-05-31T16:42:34.695519Z","shell.execute_reply":"2025-05-31T16:42:34.708465Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in tqdm(range(EPOCHS)):\n    model.train()\n    train_running_loss = 0\n    train_running_dc = 0\n    \n    for idx, img_mask in enumerate(tqdm(train_loader, position=0, leave=True)):\n        img = img_mask[0].float().to(device)\n        mask = img_mask[1].float().to(device)\n        \n        y_pred = model(img)\n        optimizer.zero_grad()\n        \n        dc = dice_coefficient(y_pred, mask)\n        loss = criterion(y_pred, mask)\n        \n        train_running_loss += loss.item()\n        train_running_dc += dc.item()\n\n        loss.backward()\n        optimizer.step()\n\n    train_loss = train_running_loss / (idx + 1)\n    train_dc = train_running_dc / (idx + 1)\n    \n    train_losses.append(train_loss)\n    train_dcs.append(train_dc)\n\n    model.eval()\n    val_running_loss = 0\n    val_running_dc = 0\n    \n    with torch.no_grad():\n        for idx, img_mask in enumerate(tqdm(val_loader, position=0, leave=True)):\n            img = img_mask[0].float().to(device)\n            mask = img_mask[1].float().to(device)\n\n            y_pred = model(img)\n            loss = criterion(y_pred, mask)\n            dc = dice_coefficient(y_pred, mask)\n            \n            val_running_loss += loss.item()\n            val_running_dc += dc.item()\n\n        val_loss = val_running_loss / (idx + 1)\n        val_dc = val_running_dc / (idx + 1)\n    \n    val_losses.append(val_loss)\n    val_dcs.append(val_dc)\n\n    print(\"-\" * 30)\n    print(f\"Training Loss EPOCH {epoch + 1}: {train_loss:.4f}\")\n    print(f\"Training DICE EPOCH {epoch + 1}: {train_dc:.4f}\")\n    print(\"\\n\")\n    print(f\"Validation Loss EPOCH {epoch + 1}: {val_loss:.4f}\")\n    print(f\"Validation DICE EPOCH {epoch + 1}: {val_dc:.4f}\")\n    print(\"-\" * 30)\n\n    if val_dc > best_val_dice:\n        best_val_dice = val_dc\n        torch.save(model.state_dict(), best_model_path)\n        print(f\"✅ Saved new best model at epoch {epoch + 1} with Dice: {val_dc:.4f}\")\n\n# Saving the model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:42:34.710170Z","iopub.execute_input":"2025-05-31T16:42:34.710419Z","iopub.status.idle":"2025-05-31T16:51:12.444089Z","shell.execute_reply.started":"2025-05-31T16:42:34.710398Z","shell.execute_reply":"2025-05-31T16:51:12.443328Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"epochs = range(1, len(train_losses) + 1)\n\nplt.figure(figsize=(14, 6))\n\n# Plot Losses\nplt.subplot(1, 2, 1)\nplt.plot(epochs, train_losses, label='Train Loss', marker='o')\nplt.plot(epochs, val_losses, label='Val Loss', marker='o')\nplt.title('Loss over Epochs')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.grid(True)\n\n# Plot Dice Coefficients\nplt.subplot(1, 2, 2)\nplt.plot(epochs, train_dcs, label='Train Dice', marker='o')\nplt.plot(epochs, val_dcs, label='Val Dice', marker='o')\nplt.title('Dice Coefficient over Epochs')\nplt.xlabel('Epoch')\nplt.ylabel('Dice Coefficient')\nplt.legend()\nplt.grid(True)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:51:12.445098Z","iopub.execute_input":"2025-05-31T16:51:12.445368Z","iopub.status.idle":"2025-05-31T16:51:12.813413Z","shell.execute_reply.started":"2025-05-31T16:51:12.445343Z","shell.execute_reply":"2025-05-31T16:51:12.812722Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load the saved model\nmodel.load_state_dict(torch.load(\"/kaggle/working/best_model.pth\"))\nmodel.eval()\n\nprint(\"✅ Loaded model.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:51:12.814155Z","iopub.execute_input":"2025-05-31T16:51:12.814503Z","iopub.status.idle":"2025-05-31T16:51:12.922753Z","shell.execute_reply.started":"2025-05-31T16:51:12.814485Z","shell.execute_reply":"2025-05-31T16:51:12.922084Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Evaluate on validation/test set\ndice_scores = []\nlosses = []\n\nwith torch.no_grad():\n    for img_mask in tqdm(val_loader):  # or test_loader\n        img = img_mask[0].float().to(device)\n        mask = img_mask[1].float().to(device)\n\n        y_pred = model(img)\n        loss = criterion(y_pred, mask)\n        dc = dice_coefficient(y_pred, mask)\n\n        dice_scores.append(dc.item())\n        losses.append(loss.item())\n\n# Final metrics\nmean_dice = sum(dice_scores) / len(dice_scores)\nmean_loss = sum(losses) / len(losses)\n\nprint(\"🧪 Evaluation Results on Best Model:\")\nprint(f\"Average Dice Coefficient: {mean_dice:.4f}\")\nprint(f\"Average Loss: {mean_loss:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:51:12.925183Z","iopub.execute_input":"2025-05-31T16:51:12.925413Z","iopub.status.idle":"2025-05-31T16:52:01.090735Z","shell.execute_reply.started":"2025-05-31T16:51:12.925396Z","shell.execute_reply":"2025-05-31T16:52:01.089826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport torch\nimport random\n\ndef show_batch_grid_with_preds(dataset, model, device, num_pairs=6):\n    assert num_pairs % 3 == 0, \"Use a multiple of 3 for clean rows\"\n\n    model.eval()\n    rows = (num_pairs // 3) * 2\n    cols = 9  # image, prediction, GT mask\n\n    fig, axes = plt.subplots(rows, cols, figsize=(22, rows * 2.5))\n    for ax in axes.flat:\n        ax.axis('off')\n\n    # Randomly sample indices\n    indices = random.sample(range(len(dataset)), num_pairs)\n\n    for i, idx in enumerate(indices):\n        img, mask = dataset[idx]\n        img_input = img.unsqueeze(0).to(device).float()\n\n        with torch.no_grad():\n            pred = model(img_input)\n            pred_np = pred.squeeze().cpu().numpy()\n            pred_np = (pred_np > 0.5).astype(float)\n\n        img_np = img.permute(1, 2, 0).numpy()\n        mask_np = mask.squeeze().numpy()\n\n        row = (i // 3) * 2\n        col = (i % 3) * 3\n\n        axes[row, col].imshow(img_np)\n        axes[row, col].set_title(f\"Image {idx}\")\n\n        axes[row, col+1].imshow(pred_np, cmap='gray')\n        axes[row, col+1].set_title(f\"Pred Mask {idx}\")\n\n        axes[row, col+2].imshow(mask_np, cmap='gray')\n        axes[row, col+2].set_title(f\"Mask {idx}\")\n\n    plt.tight_layout()\n    plt.show()\n\n\nshow_batch_grid_with_preds(val_dataset, model, device, num_pairs=6)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-31T16:52:01.091753Z","iopub.execute_input":"2025-05-31T16:52:01.091994Z","iopub.status.idle":"2025-05-31T16:52:03.738262Z","shell.execute_reply.started":"2025-05-31T16:52:01.091970Z","shell.execute_reply":"2025-05-31T16:52:03.737533Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"reference: https://medium.com/@fernandopalominocobo/mastering-u-net-a-step-by-step-guide-to-segmentation-from-scratch-with-pytorch-6a17c5916114","metadata":{}}]}