{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.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":"competition","sourceId":6927,"databundleVersionId":45059}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **Carvana Challenge - UNet with Resnet backbone**","metadata":{}},{"cell_type":"markdown","source":"# **1 - Libraries**","metadata":{}},{"cell_type":"code","source":"import os \nimport cv2 \nimport random \nimport zipfile\nfrom tqdm import tqdm\n\n# Pytorch \nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms.functional as TF\nimport torchvision.models as models\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, random_split\nfrom torchvision import transforms\n\n# Augmentation \nimport albumentations as A \nfrom albumentations.pytorch import ToTensorV2\n\n# Data Processing \nimport numpy as np \nimport pandas as pd\nfrom PIL import Image\n\n# Visualization \nimport matplotlib.pyplot as plt","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:04:13.463387Z","iopub.execute_input":"2026-05-11T10:04:13.463624Z","iopub.status.idle":"2026-05-11T10:04:23.025816Z","shell.execute_reply.started":"2026-05-11T10:04:13.463600Z","shell.execute_reply":"2026-05-11T10:04:23.024960Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **2 - UNet model Resnet backbone**","metadata":{}},{"cell_type":"markdown","source":"\nResnet34 Structure\nResNet(\n  (conv1): Conv2d(3, 64, kernel_size=7, stride=2, padding=3)\n  (bn1): BatchNorm2d(64)\n  (relu): ReLU(inplace=True)\n  (maxpool): MaxPool2d(kernel_size=3, stride=2, padding=1)\n  (layer1): Sequential(\n      (0): BasicBlock(...)\n      (1): BasicBlock(...)\n      (2): BasicBlock(...)\n  )\n  (layer2): Sequential(\n      (0): BasicBlock(...)\n      ...\n  )\n  (layer3): Sequential(...)\n  (layer4): Sequential(...)\n  (avgpool): AdaptiveAvgPool2d(...)\n  (fc): Linear(...)\n)","metadata":{}},{"cell_type":"markdown","source":"## Resnet34 Encoder","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nfrom torchvision import models\n\nclass EncoderResnet50(nn.Module):\n    def __init__(self, use_pretrained=True):\n        super(EncoderResnet50, self).__init__()\n        self.resnet = models.resnet50(\n            weights=models.ResNet50_Weights.IMAGENET1K_V1 if use_pretrained else None\n        )\n        \n        self.encoder = nn.ModuleList([\n            nn.Sequential(self.resnet.conv1, self.resnet.bn1, self.resnet.relu),  # 64, H/2\n            self.resnet.maxpool,                                                   # H/4\n            self.resnet.layer1,                                                    # 256, H/4\n            self.resnet.layer2,                                                    # 512, H/8\n            self.resnet.layer3,                                                    # 1024, H/16\n        ])\n        \n        self.bottleneck = self.resnet.layer4                                       # 2048, H/32\n\n    def forward(self, x):\n        skip_connections = []\n\n        for i, layer in enumerate(self.encoder):\n            x = layer(x)\n            if i != 1:  # 跳过 maxpool，不作为 skip\n                skip_connections.append(x)\n\n        x = self.bottleneck(x)\n            \n        return x, skip_connections","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:04:23.027412Z","iopub.execute_input":"2026-05-11T10:04:23.027810Z","iopub.status.idle":"2026-05-11T10:04:23.034106Z","shell.execute_reply.started":"2026-05-11T10:04:23.027791Z","shell.execute_reply":"2026-05-11T10:04:23.033268Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Decoder ","metadata":{}},{"cell_type":"code","source":"class DecoderModule(nn.Module):\n    # This Decoder Module include conv block + up-sample(conv transpose)\n    def __init__(self, in_channels, out_channels):\n        super(DecoderModule, self).__init__()\n        self.up_sample = nn.ConvTranspose2d(in_channels, out_channels, kernel_size=2 , stride = 2)\n        self.conv = nn.Sequential(\n            nn.Conv2d(out_channels*2, out_channels, kernel_size=3, stride=1, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)    \n        )\n        \n    def forward(self, x, skip_connection):\n        x = self.up_sample(x)\n        \n        if skip_connection.shape[2:] != x.shape[2:]: \n                x = TF.resize(x, size=skip_connection.shape[2:]) \n            \n        x_concat = torch.cat([skip_connection, x], dim=1)\n        return self.conv(x_concat)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:04:23.034885Z","iopub.execute_input":"2026-05-11T10:04:23.035392Z","iopub.status.idle":"2026-05-11T10:04:23.055353Z","shell.execute_reply.started":"2026-05-11T10:04:23.035373Z","shell.execute_reply":"2026-05-11T10:04:23.054757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Decoder(nn.Module):\n    def __init__(self, features=[2048, 1024, 512, 256, 64]):\n        super(Decoder, self).__init__()\n        self.features = features \n        self.decoder = nn.ModuleList()\n\n        for i in range(len(self.features)-1):\n            in_channels = self.features[i]\n            out_channels = self.features[i+1]\n            self.decoder.append(DecoderModule(in_channels, out_channels))\n        \n    def forward(self, x, skip_connections):\n        skip_connections = skip_connections[::-1]\n        for i, layer in enumerate(self.decoder):\n            skip_connection = skip_connections[i]\n            x = layer(x, skip_connection)\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:04:23.056128Z","iopub.execute_input":"2026-05-11T10:04:23.056400Z","iopub.status.idle":"2026-05-11T10:04:23.071639Z","shell.execute_reply.started":"2026-05-11T10:04:23.056365Z","shell.execute_reply":"2026-05-11T10:04:23.071094Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## UNet","metadata":{}},{"cell_type":"code","source":"class UNetResnet50(nn.Module):\n    def __init__(self, in_channels=3, out_channels=1, features=[2048, 1024, 512, 256, 64]):\n        super(UNetResnet50, self).__init__()\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n        self.features = features \n\n        self.encoder = EncoderResnet50()\n        self.decoder = Decoder(self.features)\n\n        last_channel = features[-1]\n        self.final_conv = nn.Sequential(\n            nn.ConvTranspose2d(last_channel, last_channel, kernel_size=2 , stride = 2),\n            \n            nn.Conv2d(last_channel, last_channel, kernel_size=3, stride=1, padding=1, bias=False),\n            nn.BatchNorm2d(last_channel),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(last_channel, last_channel, kernel_size=3, stride=1, padding=1, bias=False),\n            nn.BatchNorm2d(last_channel),\n            nn.ReLU(inplace=True),\n            \n            nn.Conv2d(last_channel, self.out_channels, kernel_size=1, stride=1)\n        )\n\n    def forward(self, x):\n        x, skip_connections = self.encoder(x)\n        x = self.decoder(x, skip_connections)\n        return self.final_conv(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:04:23.072365Z","iopub.execute_input":"2026-05-11T10:04:23.073135Z","iopub.status.idle":"2026-05-11T10:04:23.087874Z","shell.execute_reply.started":"2026-05-11T10:04:23.073107Z","shell.execute_reply":"2026-05-11T10:04:23.086961Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test():\n    x = torch.randn((3, 3, 512, 512))\n    model = UNetResnet50(in_channels=3, out_channels=1) \n    preds = model(x)\n    print(preds.shape)\n    assert preds.shape[2:] == x.shape[2:]\n\ntest()\n\nx = torch.randn((3, 3, 160, 160))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:04:23.088758Z","iopub.execute_input":"2026-05-11T10:04:23.089033Z","iopub.status.idle":"2026-05-11T10:04:32.745153Z","shell.execute_reply.started":"2026-05-11T10:04:23.089009Z","shell.execute_reply":"2026-05-11T10:04:32.744331Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **3 - Dataset**","metadata":{}},{"cell_type":"code","source":"# Unziping data files \ndata_path = \"/kaggle/input/carvana-image-masking-challenge/\"\n\ndef unzip(zip_file_name, extract_dest):\n    with zipfile.ZipFile(os.path.join(data_path, zip_file_name), 'r') as z:\n        z.extractall(extract_dest)\n        print(f\"Extracted {zip_file_name} -> {extract_dest}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:04:32.747258Z","iopub.execute_input":"2026-05-11T10:04:32.747928Z","iopub.status.idle":"2026-05-11T10:04:32.751832Z","shell.execute_reply.started":"2026-05-11T10:04:32.747901Z","shell.execute_reply":"2026-05-11T10:04:32.751194Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"rm -rf /kaggle/working/*","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:04:32.752675Z","iopub.execute_input":"2026-05-11T10:04:32.752934Z","iopub.status.idle":"2026-05-11T10:04:33.358276Z","shell.execute_reply.started":"2026-05-11T10:04:32.752895Z","shell.execute_reply":"2026-05-11T10:04:33.357124Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"unzip(\"test_hq.zip\", \"/kaggle/working/\")\nunzip(\"train_hq.zip\", \"/kaggle/working/\")\nunzip(\"train_masks.zip\", \"/kaggle/working/\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:04:33.359487Z","iopub.execute_input":"2026-05-11T10:04:33.359808Z","iopub.status.idle":"2026-05-11T10:07:26.612604Z","shell.execute_reply.started":"2026-05-11T10:04:33.359771Z","shell.execute_reply":"2026-05-11T10:07:26.611694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(os.listdir(\"/kaggle/working/train_hq\")[:5])\nprint(os.listdir(\"/kaggle/working/train_masks\")[:5])\n\nprint(len(os.listdir(\"/kaggle/working/train_hq\")))\nprint(len(os.listdir(\"/kaggle/working/train_masks\")))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.613470Z","iopub.execute_input":"2026-05-11T10:07:26.613717Z","iopub.status.idle":"2026-05-11T10:07:26.633681Z","shell.execute_reply.started":"2026-05-11T10:07:26.613698Z","shell.execute_reply":"2026-05-11T10:07:26.632789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CarvanaDataset():\n    def __init__(self, image_dir, mask_dir, image_list, transforms=None):\n        self.image_dir = image_dir\n        self.mask_dir = mask_dir\n        self.image_list = image_list\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.image_list)\n\n    def __getitem__(self, idx):\n        image_name = self.image_list[idx]\n\n        image_path = os.path.join(self.image_dir, image_name)\n        mask_path = os.path.join(self.mask_dir, image_name.replace(\".jpg\", \"_mask.gif\"))\n\n        image = np.array(Image.open(image_path).convert(\"RGB\"))\n        mask = np.array(Image.open(mask_path).convert(\"L\"), dtype=np.float32)\n\n        # Ensure mask is binary\n        mask[mask>0] = 1.0\n\n        # Apply Albumentations\n        if self.transforms is not None:\n            augmented = self.transforms(image=image, mask=mask)\n            image = augmented[\"image\"]\n            mask = augmented[\"mask\"]\n\n        if isinstance(mask, torch.Tensor):\n            mask = mask.float()\n\n        return image, mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.634943Z","iopub.execute_input":"2026-05-11T10:07:26.635301Z","iopub.status.idle":"2026-05-11T10:07:26.643751Z","shell.execute_reply.started":"2026-05-11T10:07:26.635275Z","shell.execute_reply":"2026-05-11T10:07:26.642804Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **4 - Ultility Functions**","metadata":{}},{"cell_type":"markdown","source":"## Checkpoint","metadata":{}},{"cell_type":"code","source":"def save_checkpoint(state, filename=\"my_checkpoint.pth\"):\n    print(\"=> Saving checkpoint\")\n    torch.save(state, filename)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.644772Z","iopub.execute_input":"2026-05-11T10:07:26.645108Z","iopub.status.idle":"2026-05-11T10:07:26.655199Z","shell.execute_reply.started":"2026-05-11T10:07:26.645083Z","shell.execute_reply":"2026-05-11T10:07:26.654191Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_checkpoint(checkpoint, model, optimizer):\n    print(\"=> Loading checkpoint\")\n    model.load_state_dict(checkpoint[\"state_dict\"])\n    optimizer.load_state_dict(checkpoint[\"optimizer\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.656111Z","iopub.execute_input":"2026-05-11T10:07:26.656520Z","iopub.status.idle":"2026-05-11T10:07:26.666916Z","shell.execute_reply.started":"2026-05-11T10:07:26.656491Z","shell.execute_reply":"2026-05-11T10:07:26.666017Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train/Validation functions","metadata":{}},{"cell_type":"code","source":"def train_fn(train_loader, model, optimizer, loss_fn, scaler, device):\n    model.train()\n    loop = tqdm(train_loader, desc=\"Training\", leave=True)\n\n    for batch_idx, (imgs, targets) in enumerate(loop):\n        imgs = imgs.to(device)\n        targets = targets.float().unsqueeze(1).to(device) # required by BCE (B,) => (B,1)\n\n        # forward\n        with torch.autocast(\"cuda\"):\n            predictions = model(imgs)\n            loss = loss_fn(predictions, targets)\n\n        # backward\n        optimizer.zero_grad()\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        # update progress bar\n        loop.set_postfix(loss=loss.item())\n    return loss.item()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.667853Z","iopub.execute_input":"2026-05-11T10:07:26.668129Z","iopub.status.idle":"2026-05-11T10:07:26.681781Z","shell.execute_reply.started":"2026-05-11T10:07:26.668110Z","shell.execute_reply":"2026-05-11T10:07:26.680795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def val_fn(val_loader, model, device): \n    # evaluate by pixel accuracy and dice \n    num_correct = 0\n    num_pixels = 0\n    dice_score = 0\n    model.eval()\n    loop = tqdm(val_loader, desc=\"Validating\", leave=True)\n\n    with torch.no_grad():\n        for imgs, targets in loop: \n            imgs = imgs.to(device) # (Batch_size, channels, h, w)\n            targets = targets.to(device).unsqueeze(1) # (Batch_size, h, w) => (Batch_size, 1, h, w)\n            preds = torch.sigmoid(model(imgs))\n            # turn to binary mask\n            preds = (preds > 0.5).float() # (>0.5)=>1, (<0.5)=>0 \n\n            num_correct += (preds == targets).sum() # sum of all correctly predicted pixels\n            num_pixels += torch.numel(preds) # sum of all pixels in the mask \n            \n            intersection = (preds * targets).sum() # since both are binary mask, result=1 <=> both are 1\n            union = (preds + targets).sum()\n            dice_score += (2*intersection) / (union + 1e-8)\n        dice_score = dice_score / len(val_loader)\n\n    print(f\"Correctly predicted {num_correct}/{num_pixels} => Acc: {(num_correct/num_pixels)*100:.2f}\")\n    print(f\"Dice score: {dice_score}\")\n\n    return dice_score     ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.682854Z","iopub.execute_input":"2026-05-11T10:07:26.683159Z","iopub.status.idle":"2026-05-11T10:07:26.696274Z","shell.execute_reply.started":"2026-05-11T10:07:26.683139Z","shell.execute_reply":"2026-05-11T10:07:26.695161Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Freeze/Unfreeze functions","metadata":{}},{"cell_type":"code","source":"def freeze_backbone(model):\n    if isinstance(model, torch.nn.DataParallel):\n        model = model.module\n    # Freeze all backbone params (ResNet) and set BN layers to eval mode\n    for name, param in model.encoder.named_parameters():\n        param.requires_grad = False\n    model.encoder.eval()\n    print(\"Backbone frozen ~ Warm up stage\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.697322Z","iopub.execute_input":"2026-05-11T10:07:26.697570Z","iopub.status.idle":"2026-05-11T10:07:26.709692Z","shell.execute_reply.started":"2026-05-11T10:07:26.697550Z","shell.execute_reply":"2026-05-11T10:07:26.708734Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def unfreeze_backbone(model):\n    if isinstance(model, torch.nn.DataParallel):\n        model = model.module\n    # Unfreeze backbone for fine-tuning\n    for name, param in model.encoder.named_parameters():\n        param.requires_grad = True\n    model.encoder.train()\n    print(\"Backbone unfrozen ~ Fine-tuning stage\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.710550Z","iopub.execute_input":"2026-05-11T10:07:26.711161Z","iopub.status.idle":"2026-05-11T10:07:26.723188Z","shell.execute_reply.started":"2026-05-11T10:07:26.711137Z","shell.execute_reply":"2026-05-11T10:07:26.722290Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Get Loaders ","metadata":{}},{"cell_type":"code","source":"def get_loaders(img_dir, mask_dir, train_img_list, val_img_list,\n                batch_size, train_transforms, val_transforms,\n                num_workers=4, pin_memory=True):\n    \n    train_dataset = CarvanaDataset(img_dir, mask_dir, train_img_list, train_transforms)\n    val_dataset = CarvanaDataset(img_dir, mask_dir, val_img_list, val_transforms)\n    \n    train_loader = DataLoader(train_dataset, \n                             batch_size=batch_size,\n                             num_workers=num_workers,\n                             pin_memory=pin_memory,\n                             shuffle=True,\n                             drop_last=True)\n    val_loader = DataLoader(val_dataset, \n                             batch_size=batch_size,\n                             num_workers=num_workers,\n                             pin_memory=pin_memory,\n                             shuffle=False,\n                             drop_last=True)\n    return train_dataset, train_loader, val_dataset, val_loader\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.723940Z","iopub.execute_input":"2026-05-11T10:07:26.724498Z","iopub.status.idle":"2026-05-11T10:07:26.735997Z","shell.execute_reply.started":"2026-05-11T10:07:26.724477Z","shell.execute_reply":"2026-05-11T10:07:26.735028Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **5 - Training**","metadata":{}},{"cell_type":"code","source":"carvana_path = \"/kaggle/working/\"\ntrainval_images = os.path.join(carvana_path, \"train_hq\")\ntrainval_masks = os.path.join(carvana_path, \"train_masks\")\ntest_images = os.path.join(carvana_path, \"test_hq\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.736926Z","iopub.execute_input":"2026-05-11T10:07:26.737465Z","iopub.status.idle":"2026-05-11T10:07:26.749174Z","shell.execute_reply.started":"2026-05-11T10:07:26.737442Z","shell.execute_reply":"2026-05-11T10:07:26.748251Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Hyperparameters","metadata":{}},{"cell_type":"code","source":"seed = 123\ntorch.manual_seed(seed)\n\n# Hyperparameters\nLEARNING_RATE = 1e-4\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nBATCH_SIZE = 8\nNUM_EPOCHS = 15\nNUM_WORKERS = 4\nIMAGE_HEIGHT = 1024 # 1280 originally\nIMAGE_WIDTH = 1024  # 1918 originally\nPIN_MEMORY = True\nLOAD_MODEL = False\nLOAD_MODEL_FILE = \"/kaggle/working/Semantic_Segmentation_UNet.pth\"\nIMG_DIR = trainval_images\nMASK_DIR = trainval_masks","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:26.749975Z","iopub.execute_input":"2026-05-11T10:07:26.750381Z","iopub.status.idle":"2026-05-11T10:07:27.021943Z","shell.execute_reply.started":"2026-05-11T10:07:26.750359Z","shell.execute_reply":"2026-05-11T10:07:27.021126Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Processing","metadata":{}},{"cell_type":"markdown","source":"### Data Split","metadata":{}},{"cell_type":"code","source":"split_ratio = 0.9\n\nall_images = [f for f in os.listdir(IMG_DIR)]\ntrain_size = int(split_ratio * len(all_images))\n\nrandom.seed(123)\nrandom.shuffle(all_images)\n\ntrain_image_list = all_images[:train_size]\nval_image_list = all_images[train_size:]\n\nprint(f\"({len(train_image_list)} train samples) | ({len(val_image_list)} validation samples)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:27.023133Z","iopub.execute_input":"2026-05-11T10:07:27.023467Z","iopub.status.idle":"2026-05-11T10:07:27.035901Z","shell.execute_reply.started":"2026-05-11T10:07:27.023448Z","shell.execute_reply":"2026-05-11T10:07:27.034969Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Augmentations ","metadata":{}},{"cell_type":"code","source":"train_transforms = A.Compose([    \n            A.Resize(height=IMAGE_HEIGHT, width=IMAGE_WIDTH),\n            # Stretch the longest side to 448px, then stretch the other side propotionally\n            # Then add padding if the image still smaller after stretch\n\n            # --- Augmentations ---\n            A.HorizontalFlip(p=0.5), # randomly flip\n            A.VerticalFlip(p=0.1),\n            A.HueSaturationValue( hue_shift_limit=15, sat_shift_limit=30, val_shift_limit=15, p=0.5),\n            A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1, p=0.4), #randomly adjust color - 40%\n            A.RandomBrightnessContrast(p=0.4), # randomly adjust brightness - 40%\n            A.Blur(blur_limit=3, p=0.2),\n                \n            A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=10, border_mode=cv2.BORDER_CONSTANT, p=0.5),\n            # randomly moves, zooms, rotates\n\n            # --- Normalization ---\n            A.Normalize(\n                mean=(0.485, 0.456, 0.406),\n                std=(0.229, 0.224, 0.225),\n                max_pixel_value=255.0,\n            ),\n            ToTensorV2()\n        ],\n    )\n\nval_transforms = A.Compose(\n        [\n            A.Resize(height=IMAGE_HEIGHT, width=IMAGE_WIDTH),\n            A.Normalize(\n                mean=(0.485, 0.456, 0.406),\n                std=(0.229, 0.224, 0.225),\n                max_pixel_value=255.0,\n            ),\n            ToTensorV2(),\n        ],\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:27.039878Z","iopub.execute_input":"2026-05-11T10:07:27.040353Z","iopub.status.idle":"2026-05-11T10:07:27.063408Z","shell.execute_reply.started":"2026-05-11T10:07:27.040336Z","shell.execute_reply":"2026-05-11T10:07:27.062505Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Loaders","metadata":{}},{"cell_type":"code","source":"train_dataset, train_loader, val_dataset, val_loader = get_loaders(IMG_DIR, MASK_DIR, train_image_list, val_image_list,\n                                    BATCH_SIZE, train_transforms, val_transforms,\n                                    NUM_WORKERS, PIN_MEMORY)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:27.064315Z","iopub.execute_input":"2026-05-11T10:07:27.064584Z","iopub.status.idle":"2026-05-11T10:07:27.075255Z","shell.execute_reply.started":"2026-05-11T10:07:27.064563Z","shell.execute_reply":"2026-05-11T10:07:27.074237Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Loop","metadata":{}},{"cell_type":"code","source":"model = UNetResnet50(in_channels=3, out_channels=1).to(DEVICE)\nmodel = nn.DataParallel(model)\nloss_fn = nn.BCEWithLogitsLoss() # Binary Cross Entropy + Sigmoid\noptimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:27.076020Z","iopub.execute_input":"2026-05-11T10:07:27.076371Z","iopub.status.idle":"2026-05-11T10:07:28.693700Z","shell.execute_reply.started":"2026-05-11T10:07:27.076349Z","shell.execute_reply":"2026-05-11T10:07:28.692742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scaler = torch.amp.GradScaler(\"cuda\")\nbest_dice = 0\n\nfreeze_backbone(model)\nfor epoch in range(NUM_EPOCHS):\n    print(f\"[Epoch: {epoch+1}/{NUM_EPOCHS}]\")\n    train_fn(train_loader, model, optimizer, loss_fn, scaler, device=DEVICE)\n\n    checkpoint = {\n        \"state_dict\": model.state_dict(),\n        \"optimizer\": optimizer.state_dict()\n    }\n\n    dice_score = val_fn(val_loader, model, device=DEVICE)\n    if dice_score >= best_dice:\n        best_dice = dice_score\n        save_checkpoint(checkpoint)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-11T10:07:28.694715Z","iopub.execute_input":"2026-05-11T10:07:28.694993Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = optim.AdamW(model.parameters(), lr=1e-5)\n\nunfreeze_backbone(model)\nfor epoch in range(20):\n    print(f\"[Epoch: {epoch+1+NUM_EPOCHS}/{NUM_EPOCHS+20}]\")\n    train_fn(train_loader, model, optimizer, loss_fn, scaler, device=DEVICE)\n\n    checkpoint = {\n        \"state_dict\": model.state_dict(),\n        \"optimizer\": optimizer.state_dict()\n    }\n\n    dice_score = val_fn(val_loader, model, device=DEVICE)\n    if dice_score >= best_dice:\n        best_dice = dice_score\n        save_checkpoint(checkpoint)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **7 - Visualization**","metadata":{}},{"cell_type":"code","source":"import torch, gc\n\ngc.collect()\ntorch.cuda.empty_cache()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if LOAD_MODEL: \n    checkpoint = torch.load(LOAD_MODEL_FILE, map_location=\"cpu\")\n    load_checkpoint(checkpoint, model, optimizer)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"a, b = 8, 8\n\nfig, axes = plt.subplots(a, b, figsize=(b*8, a*8))\nrandom_indices = random.sample(range(len(val_dataset)), a * b)\n\nidx = 0\nfor i in range(a):\n    for j in range(b):\n\n        img, _ = val_dataset[random_indices[idx]]\n        idx += 1\n\n        # add batch dimension\n        img_input = img.unsqueeze(0).to(device)\n\n        # model prediction\n        with torch.no_grad():\n            pred = torch.sigmoid(model(img_input))[0, 0].cpu()  # remove batch + channel\n\n        ax = axes[i][j]\n        ax.imshow(pred, cmap=\"gray\")\n        ax.set_axis_off()\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport matplotlib.pyplot as plt\nimport random\nimport numpy as np\n\n# Resnet34 Normalization\nmean = torch.tensor([0.485, 0.456, 0.406]).view(3,1,1)\nstd  = torch.tensor([0.229, 0.224, 0.225]).view(3,1,1)\n\n\n\ndef denormalize(img_tensor):\n    \"\"\"Undo normalization: (img * std + mean) and return numpy HWC\"\"\"\n    img = img_tensor.clone().cpu()\n    img = img * std + mean\n    img = img.clamp(0, 1)\n    return img.permute(1, 2, 0).numpy()\n\n\na, b = 10, 1  # number of samples shown = a\n\nfig, axes = plt.subplots(a, 4, figsize=(4 * 6, a * 6))\nrandom_indices = random.sample(range(len(val_dataset)), a)\n\nfor row, idx in enumerate(random_indices):\n\n    img, gt_mask = val_dataset[idx]\n\n    # ---- denormalize BEFORE converting to numpy ----\n    img_np = denormalize(img)\n\n    # ---- GT mask ----\n    gt_mask_np = gt_mask.cpu().numpy()\n\n    # ---- model prediction ----\n    img_input = img.unsqueeze(0).to(DEVICE)\n\n    with torch.no_grad():\n        pred = torch.sigmoid(model(img_input))[0, 0].cpu().numpy()\n\n    # ---- difference map ----\n    diff = np.zeros((*gt_mask_np.shape, 3))\n\n    fp = (pred > 0.5) & (gt_mask_np == 0)      # false positive (red)\n    fn = (pred <= 0.5) & (gt_mask_np == 1)     # false negative (blue)\n    ok = ~(fp | fn)                            # correct pixels (gray)\n\n    diff[ok] = [0.7, 0.7, 0.7]\n    diff[fp] = [1.0, 0.0, 0.0]\n    diff[fn] = [0.0, 0.0, 1.0]\n\n    # ---- plot ----\n    axes[row][0].imshow(img_np)\n    axes[row][0].set_title(\"Image (denormalized)\")\n    axes[row][0].set_axis_off()\n\n    axes[row][1].imshow(gt_mask_np, cmap=\"gray\")\n    axes[row][1].set_title(\"Ground Truth Mask\")\n    axes[row][1].set_axis_off()\n\n    axes[row][2].imshow(pred, cmap=\"gray\")\n    axes[row][2].set_title(\"Predicted Mask\")\n    axes[row][2].set_axis_off()\n\n    axes[row][3].imshow(diff)\n    axes[row][3].set_title(\"Error Map (FP=Red, FN=Blue)\")\n    axes[row][3].set_axis_off()\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_submission_file(model, test_dir, output_csv, device):\n    import cv2\n    import os\n    import numpy as np\n    import pandas as pd\n    from torch.utils.data import Dataset, DataLoader\n    import albumentations as A\n    from albumentations.pytorch import ToTensorV2\n    from tqdm import tqdm\n    import torch\n\n    ORIGINAL_W = 1918\n    ORIGINAL_H = 1280\n\n    model.eval()\n    \n    # Dataset for test images\n    class TestDataset(Dataset):\n        def __init__(self, folder):\n            self.files = sorted(os.listdir(folder))\n            self.folder = folder\n            self.transform = A.Compose([\n                A.Resize(512, 512),\n                A.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),\n                ToTensorV2(),\n            ])\n\n        def __len__(self):\n            return len(self.files)\n\n        def __getitem__(self, idx):\n            fn = self.files[idx]\n            img = cv2.imread(os.path.join(self.folder, fn))\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n            img = self.transform(image=img)[\"image\"]\n            return img, fn\n\n    # Load test dataset\n    test_ds = TestDataset(test_dir)\n    test_loader = DataLoader(\n        test_ds,\n        batch_size=8,\n        num_workers=4,\n        pin_memory=True\n    )\n\n    results = []\n\n    with torch.no_grad():\n        for imgs, fns in tqdm(test_loader):\n            imgs = imgs.to(device)\n\n            # forward pass\n            preds = model(imgs)\n            preds = torch.sigmoid(preds).cpu().numpy()  # [B,1,512,512]\n\n            for mask, fn in zip(preds, fns):\n                mask = mask[0]   # [512,512]\n\n                # threshold\n                mask = (mask > 0.5).astype(np.uint8)\n\n                # ★ ★ ★ RESIZE BACK TO ORIGINAL SIZE ★ ★ ★\n                mask = cv2.resize(mask, (ORIGINAL_W, ORIGINAL_H), interpolation=cv2.INTER_NEAREST)\n\n                # encode RLE\n                rle = mask_to_rle(mask)\n                results.append([fn, rle])\n\n    df = pd.DataFrame(results, columns=[\"img\", \"rle_mask\"])\n    df.to_csv(output_csv, index=False)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def mask_to_rle(mask):\n    pixels = mask.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nmodel.to(device)\n\nTEST_DIR = test_images\nOUTPUT_CSV = \"submission.csv\"\n\ncreate_submission_file(\n    model=model,\n    test_dir=TEST_DIR,\n    output_csv=OUTPUT_CSV,\n    device=device\n)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}