{"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":[{"sourceId":1144795,"sourceType":"datasetVersion","datasetId":645942}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **Semantic Segmentation-UNet(Resnet34)-CityScapes**","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":"2025-12-07T12:04:31.668218Z","iopub.execute_input":"2025-12-07T12:04:31.668858Z","iopub.status.idle":"2025-12-07T12:04:31.674830Z","shell.execute_reply.started":"2025-12-07T12:04:31.668818Z","shell.execute_reply":"2025-12-07T12:04:31.674211Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# testing\nmask_path = \"/kaggle/input/d/xiaose/cityscapes/Cityspaces/gtFine/train/aachen/aachen_000000_000019_gtFine_labelTrainIds.png\"\nmask = np.array(Image.open(mask_path), dtype=np.uint8)\nprint(mask)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:31.689567Z","iopub.execute_input":"2025-12-07T12:04:31.689779Z","iopub.status.idle":"2025-12-07T12:04:31.711578Z","shell.execute_reply.started":"2025-12-07T12:04:31.689765Z","shell.execute_reply":"2025-12-07T12:04:31.710968Z"}},"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":"class EncoderResnet34(nn.Module):\n    def __init__(self, backbone='resnet34', use_pretrained=True):\n        super(EncoderResnet34, self).__init__()\n        self.resnet = models.resnet34(weights=models.ResNet34_Weights.IMAGENET1K_V1 if use_pretrained else None)\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, # 64 - H/4\n            self.resnet.layer2, # 128 - H/8\n            self.resnet.layer3, # 256 - H/16\n        ])\n        # take the last layer of Resnet as the bottleneck\n        self.bottleneck = self.resnet.layer4 # 512 - 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:\n                skip_connections.append(x)\n\n        x = self.bottleneck(x)\n            \n        return x, skip_connections   ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:31.712483Z","iopub.execute_input":"2025-12-07T12:04:31.712672Z","iopub.status.idle":"2025-12-07T12:04:31.718689Z","shell.execute_reply.started":"2025-12-07T12:04:31.712657Z","shell.execute_reply":"2025-12-07T12:04:31.717722Z"}},"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":"2025-12-07T12:04:31.719747Z","iopub.execute_input":"2025-12-07T12:04:31.719959Z","iopub.status.idle":"2025-12-07T12:04:31.734068Z","shell.execute_reply.started":"2025-12-07T12:04:31.719946Z","shell.execute_reply":"2025-12-07T12:04:31.733363Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Decoder(nn.Module):\n    def __init__(self, features=[512, 256, 128, 64, 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        # reverse the skip connection list\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":"2025-12-07T12:04:31.735390Z","iopub.execute_input":"2025-12-07T12:04:31.735615Z","iopub.status.idle":"2025-12-07T12:04:31.748550Z","shell.execute_reply.started":"2025-12-07T12:04:31.735599Z","shell.execute_reply":"2025-12-07T12:04:31.747856Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## UNet","metadata":{}},{"cell_type":"code","source":"class UNetResnet34(nn.Module):\n    def __init__(self, in_channels=3, out_channels=1, features=[512, 256, 128, 64, 64]):\n        super(UNetResnet34, self).__init__()\n        self.in_channels = in_channels\n        self.out_channels = out_channels\n        self.features = features \n\n        self.encoder = EncoderResnet34()\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)\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:31.755522Z","iopub.execute_input":"2025-12-07T12:04:31.755703Z","iopub.status.idle":"2025-12-07T12:04:31.767567Z","shell.execute_reply.started":"2025-12-07T12:04:31.755689Z","shell.execute_reply":"2025-12-07T12:04:31.766869Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def test():\n    x = torch.randn((3, 3, 512, 512))\n    model = UNetResnet34(in_channels=1, 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, 161, 161))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:31.824592Z","iopub.execute_input":"2025-12-07T12:04:31.825231Z","iopub.status.idle":"2025-12-07T12:04:35.867951Z","shell.execute_reply.started":"2025-12-07T12:04:31.825207Z","shell.execute_reply":"2025-12-07T12:04:35.867121Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **3 - Dataset**","metadata":{}},{"cell_type":"code","source":"class CarvanaDataset():\n    def __init__(self, image_dir, mask_dir, transforms=None):\n        self.image_dir = image_dir\n        self.mask_dir = mask_dir\n        self.image_list = []\n        self.transforms = transforms\n\n        #/kaggle/input/d/xiaose/cityscapes/Cityspaces/gtFine/train/aachen/aachen_000000_000019_gtFine_labelTrainIds.png\n        #/kaggle/input/d/xiaose/cityscapes/Cityspaces/images/train/aachen/aachen_000000_000019_leftImg8bit.png\n        for city in os.listdir(self.image_dir):\n            city_dir = os.path.join(self.image_dir, city)\n            for f in os.listdir(city_dir):\n                if f.endswith(\"_leftImg8bit.png\"):\n                    self.image_list.append(os.path.join(city, f))\n\n    def __len__(self):\n        return len(self.image_list)\n\n    def __getitem__(self, idx):\n        rel_path = self.image_list[idx]\n        city = rel_path.split(\"/\")[0]\n        image_name = rel_path.split(\"/\")[1]\n\n        image_path = os.path.join(self.image_dir, city, image_name)\n        mask_path = os.path.join(self.mask_dir, city, image_name.replace(\"_leftImg8bit.png\", \"_gtFine_labelTrainIds.png\"))\n\n        image = np.array(Image.open(image_path).convert(\"RGB\"))\n        mask = np.array(Image.open(mask_path), dtype=np.uint8) # so Albumentations doesn’t interpolate or modify class IDs.\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, np.ndarray):\n            mask = torch.from_numpy(mask).long()\n\n        return image, mask\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:35.869427Z","iopub.execute_input":"2025-12-07T12:04:35.869648Z","iopub.status.idle":"2025-12-07T12:04:35.876521Z","shell.execute_reply.started":"2025-12-07T12:04:35.869631Z","shell.execute_reply":"2025-12-07T12:04:35.875879Z"}},"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.tar\"):\n    print(\"=> Saving checkpoint\")\n    torch.save(state, filename)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:35.877281Z","iopub.execute_input":"2025-12-07T12:04:35.877514Z","iopub.status.idle":"2025-12-07T12:04:35.890165Z","shell.execute_reply.started":"2025-12-07T12:04:35.877491Z","shell.execute_reply":"2025-12-07T12:04:35.889398Z"}},"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":"2025-12-07T12:04:35.891729Z","iopub.execute_input":"2025-12-07T12:04:35.891972Z","iopub.status.idle":"2025-12-07T12:04:35.902207Z","shell.execute_reply.started":"2025-12-07T12:04:35.891957Z","shell.execute_reply":"2025-12-07T12:04:35.901601Z"}},"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.long().to(device) \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":"2025-12-07T12:04:35.902825Z","iopub.execute_input":"2025-12-07T12:04:35.902984Z","iopub.status.idle":"2025-12-07T12:04:35.917069Z","shell.execute_reply.started":"2025-12-07T12:04:35.902971Z","shell.execute_reply":"2025-12-07T12:04:35.916463Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def val_fn(val_loader, model, device, num_cls=19): \n    # evaluate by pixel accuracy and dice \n    model.eval()\n    \n    num_correct = 0\n    num_pixels = 0\n    dice_sum = torch.zeros(num_cls, device=device)\n    iou_sum = torch.zeros(num_cls, device=device)\n    cls_count = torch.zeros(num_cls, device=device)\n\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.long().to(device) \n            \n            preds = model(imgs).argmax(dim=1)\n\n            # ignore 225 pixels \n            valid_filter = (targets != 255) # position of pixels != 255\n            preds = preds[valid_filter]\n            targets = targets[valid_filter]\n\n            num_correct += (preds == targets).sum()\n            num_pixels += torch.numel(preds) # sum of all pixels in the mask \n            \n            for cls in range(num_cls):\n                pred_cls = (preds == cls)\n                target_cls = (targets == cls)\n\n                intersection = (pred_cls & target_cls).sum()\n                union = pred_cls.sum() + target_cls.sum()\n\n                TP = ( pred_cls &  target_cls).sum()\n                FP = ( pred_cls & ~target_cls).sum()\n                FN = (~pred_cls &  target_cls).sum()\n                \n                denom_dice = (2 * TP + FP + FN).float()\n                denom_iou  = (TP + FP + FN).float()\n\n                if denom_iou > 0:   # only count classes that appear\n                    dice_sum[cls] += (2 * TP.float()) / (denom_dice + 1e-8)\n                    iou_sum[cls]  += TP.float() / (denom_iou + 1e-8)\n                    cls_count[cls] += 1\n\n    pixel_acc = (num_correct / num_pixels).item() * 100\n    mean_dice = (dice_sum / (cls_count + 1e-8)).mean().item()\n    mean_iou  = (iou_sum  / (cls_count + 1e-8)).mean().item()\n\n    print(f\"Pixel Accuracy: {pixel_acc:.2f}%\")\n    print(f\"Mean Dice: {mean_dice:.4f}\")\n    print(f\"Mean IoU :  {mean_iou:.4f}\")\n\n    return mean_dice, mean_iou     ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:35.917807Z","iopub.execute_input":"2025-12-07T12:04:35.918054Z","iopub.status.idle":"2025-12-07T12:04:35.932842Z","shell.execute_reply.started":"2025-12-07T12:04:35.918034Z","shell.execute_reply":"2025-12-07T12:04:35.932245Z"}},"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":"2025-12-07T12:04:35.933584Z","iopub.execute_input":"2025-12-07T12:04:35.933743Z","iopub.status.idle":"2025-12-07T12:04:35.947015Z","shell.execute_reply.started":"2025-12-07T12:04:35.933724Z","shell.execute_reply":"2025-12-07T12:04:35.946437Z"}},"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":"2025-12-07T12:04:35.947502Z","iopub.execute_input":"2025-12-07T12:04:35.947697Z","iopub.status.idle":"2025-12-07T12:04:35.954316Z","shell.execute_reply.started":"2025-12-07T12:04:35.947682Z","shell.execute_reply":"2025-12-07T12:04:35.953708Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Get Loaders ","metadata":{}},{"cell_type":"code","source":"def get_loaders(train_img_dir, train_mask_dir, val_img_dir, val_mask_dir,\n                batch_size, train_transforms, val_transforms,\n                num_workers=4, pin_memory=True):\n    \n    train_dataset = CarvanaDataset(train_img_dir, train_mask_dir, train_transforms)\n    val_dataset = CarvanaDataset(val_img_dir, val_mask_dir, 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":"2025-12-07T12:04:35.954920Z","iopub.execute_input":"2025-12-07T12:04:35.955130Z","iopub.status.idle":"2025-12-07T12:04:35.965779Z","shell.execute_reply.started":"2025-12-07T12:04:35.955108Z","shell.execute_reply":"2025-12-07T12:04:35.965019Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# **5 - Training**","metadata":{}},{"cell_type":"markdown","source":"## Hyperparameters","metadata":{}},{"cell_type":"code","source":"seed = 123\ntorch.manual_seed(seed)\n\n# Hyperparameters\nLEARNING_RATE_A = 1e-4\nLEARNING_RATE_B = 1e-5\nWEIGHT_DECAY = 1e-4\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nNUM_CLASSES = 19\nBATCH_SIZE = 16\nNUM_EPOCHS = 15\nNUM_WORKERS = 4\nIMAGE_HEIGHT = 512 # 1280 originally\nIMAGE_WIDTH = 1024  # 1918 originally\nPIN_MEMORY = True\nLOAD_MODEL = False\nLOAD_MODEL_FILE = \"/kaggle/working/Cityscapes.pth\"\nTRAIN_IMG_DIR  = \"/kaggle/input/d/xiaose/cityscapes/Cityspaces/images/train\"\nTRAIN_MASK_DIR = \"/kaggle/input/d/xiaose/cityscapes/Cityspaces/gtFine/train\"\nVAL_IMG_DIR  = \"/kaggle/input/d/xiaose/cityscapes/Cityspaces/images/val\"\nVAL_MASK_DIR = \"/kaggle/input/d/xiaose/cityscapes/Cityspaces/gtFine/val\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:35.967822Z","iopub.execute_input":"2025-12-07T12:04:35.968194Z","iopub.status.idle":"2025-12-07T12:04:35.981092Z","shell.execute_reply.started":"2025-12-07T12:04:35.968177Z","shell.execute_reply":"2025-12-07T12:04:35.980419Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Processing","metadata":{}},{"cell_type":"markdown","source":"### Augmentations ","metadata":{}},{"cell_type":"code","source":"train_transforms = A.Compose([    \n            A.Resize(height=IMAGE_HEIGHT, width=IMAGE_WIDTH, interpolation=cv2.INTER_NEAREST),\n            # --- Augmentations ---\n            # geometry\n            A.RandomScale(scale_limit=(0.5, 2.0), p=1.0),\n            A.HorizontalFlip(p=0.5),\n            A.RandomCrop(height=IMAGE_HEIGHT, width=IMAGE_WIDTH, p=1.0),\n            # color transforms\n            A.ColorJitter(p=0.3),\n            A.RandomBrightnessContrast(p=0.5),\n            A.HueSaturationValue(p=0.5),\n            A.RandomGamma(p=0.4),\n            # blur / distortion\n            A.GaussianBlur(blur_limit=(3, 7), p=0.3),\n            A.ElasticTransform(alpha=20, sigma=5, alpha_affine=10, p=0.2),\n            # occlusion\n            A.CoarseDropout(max_holes=6, max_height=64, max_width=64, p=0.4),\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, interpolation=cv2.INTER_NEAREST),\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":"2025-12-07T12:04:35.981767Z","iopub.execute_input":"2025-12-07T12:04:35.981931Z","iopub.status.idle":"2025-12-07T12:04:36.000262Z","shell.execute_reply.started":"2025-12-07T12:04:35.981917Z","shell.execute_reply":"2025-12-07T12:04:35.999425Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Loaders","metadata":{}},{"cell_type":"code","source":"train_dataset, train_loader, val_dataset, val_loader = get_loaders(TRAIN_IMG_DIR, TRAIN_MASK_DIR, VAL_IMG_DIR, VAL_MASK_DIR, \n                                                                    BATCH_SIZE, train_transforms, val_transforms,\n                                                                    NUM_WORKERS, PIN_MEMORY)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:36.000999Z","iopub.execute_input":"2025-12-07T12:04:36.001249Z","iopub.status.idle":"2025-12-07T12:04:36.025687Z","shell.execute_reply.started":"2025-12-07T12:04:36.001228Z","shell.execute_reply":"2025-12-07T12:04:36.025005Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Loop","metadata":{}},{"cell_type":"code","source":"model = UNetResnet34(in_channels=3, out_channels=NUM_CLASSES).to(DEVICE)\nmodel = nn.DataParallel(model)\nloss_fn = torch.nn.CrossEntropyLoss(ignore_index=255) # Binary Cross Entropy + Sigmoid\noptimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE_A, weight_decay=WEIGHT_DECAY)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:36.026391Z","iopub.execute_input":"2025-12-07T12:04:36.026622Z","iopub.status.idle":"2025-12-07T12:04:36.409248Z","shell.execute_reply.started":"2025-12-07T12:04:36.026603Z","shell.execute_reply":"2025-12-07T12:04:36.408607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scaler = torch.amp.GradScaler(\"cuda\")\nbest_mean_iou = 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    mean_dice, mean_iou = val_fn(val_loader, model, device=DEVICE)\n    if mean_iou >= best_mean_iou:\n        best_mean_iou = mean_iou\n        save_checkpoint(checkpoint)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T12:04:36.409974Z","iopub.execute_input":"2025-12-07T12:04:36.410280Z","iopub.status.idle":"2025-12-07T13:00:16.645445Z","shell.execute_reply.started":"2025-12-07T12:04:36.410261Z","shell.execute_reply":"2025-12-07T13:00:16.644607Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE_B, weight_decay=WEIGHT_DECAY)\n\nunfreeze_backbone(model)\nfor epoch in range(25):\n    print(f\"[Epoch: {epoch+1+NUM_EPOCHS}/{NUM_EPOCHS+25}]\")\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    mean_dice, mean_iou = val_fn(val_loader, model, device=DEVICE)\n    if mean_iou >= best_mean_iou:\n        best_mean_iou = mean_iou\n        save_checkpoint(checkpoint)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T13:00:16.646376Z","iopub.execute_input":"2025-12-07T13:00:16.646610Z","iopub.status.idle":"2025-12-07T14:35:33.673959Z","shell.execute_reply.started":"2025-12-07T13:00:16.646592Z","shell.execute_reply":"2025-12-07T14:35:33.673011Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"new_lr = (1e-5) / 5\nfor g in optimizer.param_groups:\n    g['lr'] = new_lr\n\nfor epoch in range(10):\n    print(f\"[Epoch: {epoch+1+NUM_EPOCHS+25}/{NUM_EPOCHS+25+10}]\")\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    mean_dice, mean_iou = val_fn(val_loader, model, device=DEVICE)\n    if mean_iou >= best_mean_iou:\n        best_mean_iou = mean_iou\n        save_checkpoint(checkpoint)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T14:42:16.013701Z","iopub.execute_input":"2025-12-07T14:42:16.014504Z","iopub.status.idle":"2025-12-07T15:18:36.833849Z","shell.execute_reply.started":"2025-12-07T14:42:16.014477Z","shell.execute_reply":"2025-12-07T15:18:36.832909Z"}},"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,"execution":{"iopub.status.busy":"2025-12-07T15:22:23.152349Z","iopub.execute_input":"2025-12-07T15:22:23.152636Z","iopub.status.idle":"2025-12-07T15:22:23.540070Z","shell.execute_reply.started":"2025-12-07T15:22:23.152608Z","shell.execute_reply":"2025-12-07T15:22:23.539443Z"}},"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,"execution":{"iopub.status.busy":"2025-12-07T15:22:23.541206Z","iopub.execute_input":"2025-12-07T15:22:23.541407Z","iopub.status.idle":"2025-12-07T15:22:23.544803Z","shell.execute_reply.started":"2025-12-07T15:22:23.541392Z","shell.execute_reply":"2025-12-07T15:22:23.544284Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T15:22:23.926982Z","iopub.execute_input":"2025-12-07T15:22:23.927258Z","iopub.status.idle":"2025-12-07T15:22:23.931183Z","shell.execute_reply.started":"2025-12-07T15:22:23.927238Z","shell.execute_reply":"2025-12-07T15:22:23.930444Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport matplotlib.pyplot as plt\nimport random\nimport numpy as np\n\n# -------------------------\n# Cityscapes COLOR MAP (19 classes)\n# -------------------------\nCITYSCAPES_COLORS = [\n    (128, 64,128), (244, 35,232), ( 70, 70, 70), (102,102,156),\n    (190,153,153), (153,153,153), (250,170, 30), (220,220,  0),\n    (107,142, 35), (152,251,152), ( 70,130,180), (220, 20, 60),\n    (255,  0,  0), (  0,  0,142), (  0,  0, 70), (  0, 60,100),\n    (  0, 80,100), (  0,  0,230), (119, 11, 32)\n]\n\ndef colorize_mask(mask):\n    \"\"\"Convert trainId mask (HxW) to RGB while ignoring 255.\"\"\"\n    h, w = mask.shape\n    color_mask = np.zeros((h, w, 3), dtype=np.uint8)\n\n    for cid, col in enumerate(CITYSCAPES_COLORS):\n        color_mask[mask == cid] = col\n\n    # ignored area: paint black\n    color_mask[mask == 255] = (0, 0, 0)\n\n    return color_mask\n\n\n# -------------------------\n# ResNet Normalization Undo\n# -------------------------\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\ndef denormalize(img_tensor):\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\n# -------------------------\n# Display N Samples\n# -------------------------\nN = 20\n\nfig, axes = plt.subplots(N, 4, figsize=(4 * 6, N * 5))\nrandom_indices = random.sample(range(len(val_dataset)), N)\n\nfor row, idx in enumerate(random_indices):\n\n    img, gt_mask = val_dataset[idx]\n    img_np = denormalize(img)\n\n    gt_np = gt_mask.cpu().numpy().astype(np.uint8)\n\n    # VALID PIXELS = NOT IGNORED\n    valid = (gt_np != 255)\n\n    # color GT normally (ignored=black)\n    gt_color = colorize_mask(gt_np)\n\n    # ----- Model -----\n    img_input = img.unsqueeze(0).to(DEVICE)\n\n    with torch.no_grad():\n        logits = model(img_input)[0]  \n        pred_mask = torch.argmax(logits, dim=0).cpu().numpy().astype(np.uint8)\n\n    # mask out ignored pixels in prediction\n    pred_mask_vis = pred_mask.copy()\n    pred_mask_vis[~valid] = 255  # mark ignored\n    pred_color = colorize_mask(pred_mask_vis)\n\n    # ----- Error Map -----\n    diff = np.zeros((*gt_np.shape, 3), dtype=np.uint8)\n\n    correct = (pred_mask == gt_np) & valid\n    incorrect = (pred_mask != gt_np) & valid\n\n    diff[correct] = (180, 180, 180)  # correct = gray\n    diff[incorrect] = (255, 0, 0)    # incorrect = red\n    diff[~valid] = (0, 0, 0)         # ignored = black\n\n    # ----- Show panels -----\n    axes[row][0].imshow(img_np)\n    axes[row][0].set_title(\"Image\")\n    axes[row][0].set_axis_off()\n\n    axes[row][1].imshow(gt_color)\n    axes[row][1].set_title(\"GT Mask\")\n    axes[row][1].set_axis_off()\n\n    axes[row][2].imshow(pred_color)\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 (Red=Wrong)\")\n    axes[row][3].set_axis_off()\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-07T15:23:44.353894Z","iopub.execute_input":"2025-12-07T15:23:44.354670Z","iopub.status.idle":"2025-12-07T15:24:01.670799Z","shell.execute_reply.started":"2025-12-07T15:23:44.354645Z","shell.execute_reply":"2025-12-07T15:24:01.669532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}