{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"datasetVersion","sourceId":3009448,"datasetId":1843391,"databundleVersionId":3057323}],"dockerImageVersionId":31401,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\n\nfrom tqdm import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.030397Z","iopub.execute_input":"2026-05-30T12:06:33.030891Z","iopub.status.idle":"2026-05-30T12:06:33.036401Z","shell.execute_reply.started":"2026-05-30T12:06:33.030859Z","shell.execute_reply":"2026-05-30T12:06:33.035400Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nprint(DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.148633Z","iopub.execute_input":"2026-05-30T12:06:33.148967Z","iopub.status.idle":"2026-05-30T12:06:33.154475Z","shell.execute_reply.started":"2026-05-30T12:06:33.148930Z","shell.execute_reply":"2026-05-30T12:06:33.153762Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMAGE_DIR = \"/kaggle/input/datasets/ipythonx/carvana-image-masking-png/train_images\"\nMASK_DIR = \"/kaggle/input/datasets/ipythonx/carvana-image-masking-png/train_masks\"\n\nIMAGE_SIZE = 256\nBATCH_SIZE = 8\nLR = 1e-4\nEPOCHS = 10","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.155934Z","iopub.execute_input":"2026-05-30T12:06:33.157005Z","iopub.status.idle":"2026-05-30T12:06:33.167057Z","shell.execute_reply.started":"2026-05-30T12:06:33.156951Z","shell.execute_reply":"2026-05-30T12:06:33.166345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_transform = transforms.Compose([\n    transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),\n    transforms.ToTensor()\n])\n\nmask_transform = transforms.Compose([\n    transforms.Resize(\n        (IMAGE_SIZE, IMAGE_SIZE),\n        interpolation=Image.NEAREST\n    ),\n    transforms.ToTensor()\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.168102Z","iopub.execute_input":"2026-05-30T12:06:33.168375Z","iopub.status.idle":"2026-05-30T12:06:33.179828Z","shell.execute_reply.started":"2026-05-30T12:06:33.168355Z","shell.execute_reply":"2026-05-30T12:06:33.179139Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CarvanaDataset(Dataset):\n\n    def __init__(self, image_dir, mask_dir):\n        self.image_dir = image_dir\n        self.mask_dir = mask_dir\n\n        self.images = sorted(os.listdir(image_dir))\n\n    def __len__(self):\n        return len(self.images)\n\n    def __getitem__(self, idx):\n\n        img_name = self.images[idx]\n\n        image_path = os.path.join(\n            self.image_dir,\n            img_name\n        )\n\n        # IMPORTANT CHANGE\n        mask_name = img_name.replace(\n            \".jpg\",\n            \".png\"\n        )\n\n        mask_path = os.path.join(\n            self.mask_dir,\n            mask_name\n        )\n\n        image = Image.open(image_path).convert(\"RGB\")\n        mask = Image.open(mask_path).convert(\"L\")\n\n        image = image_transform(image)\n        mask = mask_transform(mask)\n\n        mask = (mask > 0).float()\n\n        return image, mask","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.181623Z","iopub.execute_input":"2026-05-30T12:06:33.181970Z","iopub.status.idle":"2026-05-30T12:06:33.191721Z","shell.execute_reply.started":"2026-05-30T12:06:33.181939Z","shell.execute_reply":"2026-05-30T12:06:33.191148Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset = CarvanaDataset(\n    IMAGE_DIR,\n    MASK_DIR\n)\n\nprint(\"Dataset Size:\", len(dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.192442Z","iopub.execute_input":"2026-05-30T12:06:33.192837Z","iopub.status.idle":"2026-05-30T12:06:33.212233Z","shell.execute_reply.started":"2026-05-30T12:06:33.192813Z","shell.execute_reply":"2026-05-30T12:06:33.211345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image, mask = dataset[0]\n\nprint(image.shape)\nprint(mask.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.213450Z","iopub.execute_input":"2026-05-30T12:06:33.213766Z","iopub.status.idle":"2026-05-30T12:06:33.256252Z","shell.execute_reply.started":"2026-05-30T12:06:33.213744Z","shell.execute_reply":"2026-05-30T12:06:33.255452Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image, mask = dataset[0]\n\nplt.figure(figsize=(12,5))\n\nplt.subplot(1,2,1)\nplt.imshow(image.permute(1,2,0))\nplt.title(\"Image\")\n\nplt.subplot(1,2,2)\nplt.imshow(mask.squeeze(), cmap=\"gray\")\nplt.title(\"Mask\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.257193Z","iopub.execute_input":"2026-05-30T12:06:33.257526Z","iopub.status.idle":"2026-05-30T12:06:33.611389Z","shell.execute_reply.started":"2026-05-30T12:06:33.257502Z","shell.execute_reply":"2026-05-30T12:06:33.610458Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import random_split\n\ntrain_size = int(0.8 * len(dataset))\nval_size = len(dataset) - train_size\n\ntrain_dataset, val_dataset = random_split(\n    dataset,\n    [train_size, val_size]\n)\n\nprint(len(train_dataset))\nprint(len(val_dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.612507Z","iopub.execute_input":"2026-05-30T12:06:33.613142Z","iopub.status.idle":"2026-05-30T12:06:33.623771Z","shell.execute_reply.started":"2026-05-30T12:06:33.613119Z","shell.execute_reply":"2026-05-30T12:06:33.623110Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.624686Z","iopub.execute_input":"2026-05-30T12:06:33.625028Z","iopub.status.idle":"2026-05-30T12:06:33.632410Z","shell.execute_reply.started":"2026-05-30T12:06:33.625007Z","shell.execute_reply":"2026-05-30T12:06:33.631740Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images, masks = next(iter(train_loader))\n\nprint(images.shape)\nprint(masks.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:33.634720Z","iopub.execute_input":"2026-05-30T12:06:33.635083Z","iopub.status.idle":"2026-05-30T12:06:34.755265Z","shell.execute_reply.started":"2026-05-30T12:06:33.635050Z","shell.execute_reply":"2026-05-30T12:06:34.754281Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, 3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n\n            nn.Conv2d(out_channels, out_channels, 3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n\n    def forward(self, x):\n        return self.conv(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:34.756506Z","iopub.execute_input":"2026-05-30T12:06:34.757054Z","iopub.status.idle":"2026-05-30T12:06:34.762346Z","shell.execute_reply.started":"2026-05-30T12:06:34.757022Z","shell.execute_reply":"2026-05-30T12:06:34.761695Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DownSample(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n\n        self.conv = DoubleConv(in_channels, out_channels)\n        self.pool = nn.MaxPool2d(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":"2026-05-30T12:06:34.763455Z","iopub.execute_input":"2026-05-30T12:06:34.763790Z","iopub.status.idle":"2026-05-30T12:06:34.777815Z","shell.execute_reply.started":"2026-05-30T12:06:34.763767Z","shell.execute_reply":"2026-05-30T12:06:34.776955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UpSample(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n\n        self.up = nn.ConvTranspose2d(\n            in_channels,\n            out_channels,\n            kernel_size=2,\n            stride=2\n        )\n\n        self.conv = DoubleConv(\n            in_channels,\n            out_channels\n        )\n\n    def forward(self, x1, x2):\n\n        x1 = self.up(x1)\n\n        x = torch.cat(\n            [x1, x2],\n            dim=1\n        )\n\n        return self.conv(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:34.778855Z","iopub.execute_input":"2026-05-30T12:06:34.779169Z","iopub.status.idle":"2026-05-30T12:06:34.791631Z","shell.execute_reply.started":"2026-05-30T12:06:34.779136Z","shell.execute_reply":"2026-05-30T12:06:34.790929Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNet(nn.Module):\n    def __init__(self, n_channels=3, n_classes=1):\n        super().__init__()\n\n        self.down_conv1 = DownSample(n_channels, 64)\n        self.down_conv2 = DownSample(64, 128)\n        self.down_conv3 = DownSample(128, 256)\n        self.down_conv4 = DownSample(256, 512)\n\n        self.bottleneck = DoubleConv(512, 1024)\n\n        self.up_conv1 = UpSample(1024, 512)\n        self.up_conv2 = UpSample(512, 256)\n        self.up_conv3 = UpSample(256, 128)\n        self.up_conv4 = UpSample(128, 64)\n\n        self.final_conv = nn.Conv2d(\n            64,\n            n_classes,\n            kernel_size=1\n        )\n\n    def forward(self, x):\n\n        down1, p1 = self.down_conv1(x)\n        down2, p2 = self.down_conv2(p1)\n        down3, p3 = self.down_conv3(p2)\n        down4, p4 = self.down_conv4(p3)\n\n        b = self.bottleneck(p4)\n\n        up1 = self.up_conv1(b, down4)\n        up2 = self.up_conv2(up1, down3)\n        up3 = self.up_conv3(up2, down2)\n        up4 = self.up_conv4(up3, down1)\n\n        return self.final_conv(up4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:34.792799Z","iopub.execute_input":"2026-05-30T12:06:34.793099Z","iopub.status.idle":"2026-05-30T12:06:34.804700Z","shell.execute_reply.started":"2026-05-30T12:06:34.793077Z","shell.execute_reply":"2026-05-30T12:06:34.804079Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = UNet(\n    n_channels=3,\n    n_classes=1\n).to(DEVICE)\n\nprint(\n    sum(p.numel() for p in model.parameters())\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:34.805692Z","iopub.execute_input":"2026-05-30T12:06:34.806011Z","iopub.status.idle":"2026-05-30T12:06:35.098188Z","shell.execute_reply.started":"2026-05-30T12:06:34.805979Z","shell.execute_reply":"2026-05-30T12:06:35.097461Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.BCEWithLogitsLoss()\n\noptimizer = optim.Adam(\n    model.parameters(),\n    lr=LR\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:35.099088Z","iopub.execute_input":"2026-05-30T12:06:35.099419Z","iopub.status.idle":"2026-05-30T12:06:35.104863Z","shell.execute_reply.started":"2026-05-30T12:06:35.099376Z","shell.execute_reply":"2026-05-30T12:06:35.104174Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def dice_score(preds, masks):\n\n    preds = torch.sigmoid(preds)\n    preds = (preds > 0.5).float()\n\n    intersection = (preds * masks).sum()\n\n    return (\n        2 * intersection + 1e-8\n    ) / (\n        preds.sum() + masks.sum() + 1e-8\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:35.105853Z","iopub.execute_input":"2026-05-30T12:06:35.106226Z","iopub.status.idle":"2026-05-30T12:06:35.116429Z","shell.execute_reply.started":"2026-05-30T12:06:35.106190Z","shell.execute_reply":"2026-05-30T12:06:35.115649Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_fn(loader):\n\n    model.train()\n\n    total_loss = 0\n\n    loop = tqdm(loader)\n\n    for images, masks in loop:\n\n        images = images.to(DEVICE)\n        masks = masks.to(DEVICE)\n\n        preds = model(images)\n\n        loss = criterion(\n            preds,\n            masks\n        )\n\n        optimizer.zero_grad()\n\n        loss.backward()\n\n        optimizer.step()\n\n        total_loss += loss.item()\n\n        loop.set_postfix(\n            loss=loss.item()\n        )\n\n    return total_loss / len(loader)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:35.117251Z","iopub.execute_input":"2026-05-30T12:06:35.118069Z","iopub.status.idle":"2026-05-30T12:06:35.129638Z","shell.execute_reply.started":"2026-05-30T12:06:35.118038Z","shell.execute_reply":"2026-05-30T12:06:35.128947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"@torch.no_grad()\ndef eval_fn(loader):\n\n    model.eval()\n\n    dice_avg = 0\n\n    for images, masks in loader:\n\n        images = images.to(DEVICE)\n        masks = masks.to(DEVICE)\n\n        preds = model(images)\n\n        dice_avg += dice_score(\n            preds,\n            masks\n        )\n\n    return dice_avg / len(loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:35.131211Z","iopub.execute_input":"2026-05-30T12:06:35.131718Z","iopub.status.idle":"2026-05-30T12:06:35.144214Z","shell.execute_reply.started":"2026-05-30T12:06:35.131667Z","shell.execute_reply":"2026-05-30T12:06:35.143314Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_dice = 0\n\nfor epoch in range(EPOCHS):\n\n    train_loss = train_fn(\n        train_loader\n    )\n\n    dice = eval_fn(\n        val_loader\n    )\n\n    print(\n        f\"Epoch {epoch+1}/{EPOCHS}\"\n    )\n\n    print(\n        f\"Train Loss: {train_loss:.4f}\"\n    )\n\n    print(\n        f\"Dice Score: {dice:.4f}\"\n    )\n\n    if dice > best_dice:\n\n        best_dice = dice\n\n        torch.save(\n            model.state_dict(),\n            \"best_unet.pth\"\n        )\n\n        print(\"Best Model Saved\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T12:06:35.145167Z","iopub.execute_input":"2026-05-30T12:06:35.145498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.load_state_dict(\n    torch.load(\"best_unet.pth\")\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\n\nimage, mask = dataset[0]\n\nwith torch.no_grad():\n\n    pred = model(\n        image.unsqueeze(0).to(DEVICE)\n    )\n\npred = torch.sigmoid(pred)\n\npred = (pred > 0.5).float()\n\npred = pred.squeeze().cpu()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(15,5))\n\nplt.subplot(1,3,1)\nplt.imshow(image.permute(1,2,0))\nplt.title(\"Image\")\n\nplt.subplot(1,3,2)\nplt.imshow(mask.squeeze(), cmap=\"gray\")\nplt.title(\"Ground Truth\")\n\nplt.subplot(1,3,3)\nplt.imshow(pred, cmap=\"gray\")\nplt.title(\"Prediction\")\n\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}