{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":86142,"databundleVersionId":9786425,"sourceType":"competition"},{"sourceId":135443,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":114567,"modelId":137832}],"dockerImageVersionId":30786,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 🧠 Image Segmentation on Indian Roads Dataset\n\nThis notebook implements an image segmentation pipeline tailored to Indian road scenes for autonomous driving applications. The pipeline includes:\n\n- Dataset loading from image masks and JSON annotations.\n- Custom Dataset class for PyTorch.\n- Transformations using Albumentations.\n- U-Net (or other architectures) from `segmentation-models-pytorch`.\n- Training, evaluation, and visualization of predictions.\n","metadata":{}},{"cell_type":"code","source":"!pip install -U segmentation-models-pytorch","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:08.605618Z","iopub.execute_input":"2025-06-27T07:19:08.606451Z","iopub.status.idle":"2025-06-27T07:19:16.913435Z","shell.execute_reply.started":"2025-06-27T07:19:08.606417Z","shell.execute_reply":"2025-06-27T07:19:16.912338Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Standard libraries\nimport os\nimport json\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom PIL import Image\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\n\n# PyTorch and torchvision\nimport torch\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision\n\n# Albumentations for augmentations\nimport albumentations as A\n\n# Image processing\nimport cv2\n\n# Segmentation models\nimport segmentation_models_pytorch as smp\n\n# Sklearn for splitting data\nfrom sklearn.model_selection import train_test_split\n","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:16.915302Z","iopub.execute_input":"2025-06-27T07:19:16.915599Z","iopub.status.idle":"2025-06-27T07:19:16.921255Z","shell.execute_reply.started":"2025-06-27T07:19:16.91557Z","shell.execute_reply":"2025-06-27T07:19:16.920364Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Class to Colour mapping","metadata":{}},{"cell_type":"code","source":"class Label:\n    def __init__(self, name, id, csId, csTrainId, level4id, level3Id, category, level2Id, level1Id, hasInstances, ignoreInEval, color):\n        self.name = name\n        self.id = id\n        self.csId = csId\n        self.csTrainId = csTrainId\n        self.level4id = level4id\n        self.level3Id = level3Id\n        self.category = category\n        self.level2Id = level2Id\n        self.level1Id = level1Id\n        self.hasInstances = hasInstances\n        self.ignoreInEval = ignoreInEval\n        self.color = color\n\nlabels = [\n    #       name                     id    csId     csTrainId level4id        level3Id  category           level2Id      level1Id  hasInstances   ignoreInEval   color\n    Label(  'road'                 ,  0   ,  7 ,     0 ,       0   ,     0  ,   'drivable'            , 0           , 0      , False        , False        , (128, 64,128)  ),\n    Label(  'parking'              ,  1   ,  9 ,   255 ,       1   ,     1  ,   'drivable'            , 1           , 0      , False        , False         , (250,170,160)  ),\n    Label(  'drivable fallback'    ,  2   ,  255 ,   255 ,     2   ,       1  ,   'drivable'            , 1           , 0      , False        , False         , ( 81,  0, 81)  ),\n    Label(  'sidewalk'             ,  3   ,  8 ,     1 ,       3   ,     2  ,   'non-drivable'        , 2           , 1      , False        , False        , (244, 35,232)  ),\n    Label(  'rail track'           ,  4   , 10 ,   255 ,       3   ,     3  ,   'non-drivable'        , 3           , 1      , False        , False         , (230,150,140)  ),\n    Label(  'non-drivable fallback',  5   , 255 ,     9 ,      4   ,      3  ,   'non-drivable'        , 3           , 1      , False        , False        , (152,251,152)  ),\n    Label(  'person'               ,  6   , 24 ,    11 ,       5   ,     4  ,   'living-thing'        , 4           , 2      , True         , False        , (220, 20, 60)  ),\n    Label(  'animal'               ,  7   , 255 ,   255 ,      6   ,      4  ,   'living-thing'        , 4           , 2      , True         , True        , (246, 198, 145)),\n    Label(  'rider'                ,  8   , 25 ,    12 ,       7   ,     5  ,   'living-thing'        , 5           , 2      , True         , False        , (255,  0,  0)  ),\n    Label(  'motorcycle'           ,  9   , 32 ,    17 ,       8   ,     6  ,   '2-wheeler'           , 6           , 3      , True         , False        , (  0,  0,230)  ),\n    Label(  'bicycle'              , 10   , 33 ,    18 ,       9   ,     7  ,   '2-wheeler'           , 6           , 3      , True         , False        , (119, 11, 32)  ),\n    Label(  'autorickshaw'         , 11   , 255 ,   255 ,     10   ,      8  ,   'autorickshaw'        , 7           , 3      , True         , False        , (255, 204, 54) ),\n    Label(  'car'                  , 12   , 26 ,    13 ,      11   ,     9  ,   'car'                 , 7           , 3      , True         , False        , (  0,  0,142)  ),\n    Label(  'truck'                , 13   , 27 ,    14 ,      12   ,     10 ,   'large-vehicle'       , 8           , 3      , True         , False        , (  0,  0, 70)  ),\n    Label(  'bus'                  , 14   , 28 ,    15 ,      13   ,     11 ,   'large-vehicle'       , 8           , 3      , True         , False        , (  0, 60,100)  ),\n    Label(  'caravan'              , 15   , 29 ,   255 ,      14   ,     12 ,   'large-vehicle'       , 8           , 3      , True         , True         , (  0,  0, 90)  ),\n    Label(  'trailer'              , 16   , 30 ,   255 ,      15   ,     12 ,   'large-vehicle'       , 8           , 3      , True         , True         , (  0,  0,110)  ),\n    Label(  'train'                , 17   , 31 ,    16 ,      15   ,     12 ,   'large-vehicle'       , 8           , 3      , True         , True        , (  0, 80,100)  ),\n    Label(  'vehicle fallback'     , 18   , 355 ,   255 ,     15   ,      12 ,   'large-vehicle'       , 8           , 3      , True         , False        , (136, 143, 153)),  \n    Label(  'curb'                 , 19   ,255 ,   255 ,      16   ,     13 ,   'barrier'             , 9           , 4      , False        , False        , (220, 190, 40)),\n    Label(  'wall'                 , 20   , 12 ,     3 ,      17   ,     14 ,   'barrier'             , 9           , 4      , False        , False        , (102,102,156)  ),\n    Label(  'fence'                , 21   , 13 ,     4 ,      18   ,     15 ,   'barrier'             , 10           , 4      , False        , False        , (190,153,153)  ),\n    Label(  'guard rail'           , 22   , 14 ,   255 ,      19   ,     16 ,   'barrier'             , 10          , 4      , False        , False         , (180,165,180)  ),\n    Label(  'billboard'            , 23   , 255 ,   255 ,     20   ,      17 ,   'structures'          , 11           , 4      , False        , False        , (174, 64, 67) ),\n    Label(  'traffic sign'         , 24   , 20 ,     7 ,      21   ,     18 ,   'structures'          , 11          , 4      , False        , False        , (220,220,  0)  ),\n    Label(  'traffic light'        , 25   , 19 ,     6 ,      22   ,     19 ,   'structures'          , 11          , 4      , False        , False        , (250,170, 30)  ),\n    Label(  'pole'                 , 26   , 17 ,     5 ,      23   ,     20 ,   'structures'          , 12          , 4      , False        , False        , (153,153,153)  ),\n    Label(  'polegroup'            , 27   , 18 ,   255 ,      23   ,     20 ,   'structures'          , 12          , 4      , False        , False         , (153,153,153)  ),\n    Label(  'obs-str-bar-fallback' , 28   , 255 ,   255 ,     24   ,      21 ,   'structures'          , 12          , 4      , False        , False        , (169, 187, 214) ),  \n    Label(  'building'             , 29   , 11 ,     2 ,      25   ,     22 ,   'construction'        , 13          , 5      , False        , False        , ( 70, 70, 70)  ),\n    Label(  'bridge'               , 30   , 15 ,   255 ,      26   ,     23 ,   'construction'        , 13          , 5      , False        , False         , (150,100,100)  ),\n    Label(  'tunnel'               , 31   , 16 ,   255 ,      26   ,     23 ,   'construction'        , 13          , 5      , False        , False         , (150,120, 90)  ),\n    Label(  'vegetation'           , 32   , 21 ,     8 ,      27   ,     24 ,   'vegetation'          , 14          , 5      , False        , False        , (107,142, 35)  ),\n    Label(  'sky'                  , 33   , 23 ,    10 ,      28   ,     25 ,   'sky'                 , 15          , 6      , False        , False        , ( 70,130,180)  ),\n    Label(  'fallback background'  , 34   , 255 ,   255 ,     29   ,      25 ,   'object fallback'     , 15          , 6      , False        , False        , (169, 187, 214)),\n    Label(  'unlabeled'            , 35   ,  0  ,     255 ,   255   ,      255 ,   'void'                , 255         , 255    , False        , True         , (  0,  0,  0)  ),\n    Label(  'ego vehicle'          , 36   ,  1  ,     255 ,   255   ,      255 ,   'void'                , 255         , 255    , False        , True         , (  0,  0,  0)  ),\n    Label(  'rectification border' , 37   ,  2  ,     255 ,   255   ,      255 ,   'void'                , 255         , 255    , False        , True         , (  0,  0,  0)  ),\n    Label(  'out of roi'           , 38   ,  3  ,     255 ,   255   ,      255 ,   'void'                , 255         , 255    , False        , True         , (  0,  0,  0)  ),\n    Label(  'license plate'        , 39   , 255 ,     255 ,   255   ,      255 ,   'vehicle'             , 255         , 255    , False        , True         , (  0,  0,142)  ),\n    \n]           \n\n\nclass_to_color = {label.id: label.color for label in labels}\nclass_to_color = {k: (np.array(v) / 255.0).astype(np.float32) for k, v in class_to_color.items()}\n\n","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:16.922475Z","iopub.execute_input":"2025-06-27T07:19:16.922802Z","iopub.status.idle":"2025-06-27T07:19:16.946269Z","shell.execute_reply.started":"2025-06-27T07:19:16.922767Z","shell.execute_reply":"2025-06-27T07:19:16.94563Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Helper Functions","metadata":{}},{"cell_type":"code","source":"def one_hot_to_color(one_hot, class_to_color):\n    \"\"\"\n    Convert a one-hot encoded segmentation map back to an RGB image using the class_to_color mapping.\n    Args:\n        one_hot (torch.Tensor): The one-hot encoded tensor of shape (40, H, W)\n        class_to_color (dict): The mapping from class index to color.\n\n    Returns:\n        color_image (np.ndarray): RGB image of shape (H, W, 3)\n    \"\"\"\n    one_hot_np = one_hot.cpu().numpy()  # Convert to numpy (40, H, W)\n    h, w = one_hot_np.shape[1:]  # Get height and width\n    color_image = np.zeros((h, w, 3), dtype=np.uint8)  # Initialize blank color image\n\n    for class_index, color in class_to_color.items():\n        mask = one_hot_np[class_index] == 1  # Find all pixels of the current class\n        color_image[mask] = np.array(color) * 255  # Set corresponding pixels to class color\n\n    return color_image\n\n# Function to display the original image, true label, and prediction\ndef visualize_predictions(model, dataloader, class_to_color, device):\n    model.eval()  # Set the model to evaluation mode\n\n    with torch.inference_mode():  # No need to calculate gradients\n        for images, true_labels in dataloader:  # Get a batch of images and true labels\n            # Move images and labels to the device\n            images = images.to(device)\n            true_labels = true_labels.to(device)\n\n            # Get model output (predictions)\n            preds = model(images)  # (2, 40, 512, 512)\n\n            # Take the first image in the batch\n            image = images[0].cpu().permute(1, 2, 0).numpy()  # Convert to (H, W, 3) format\n            true_label = true_labels[0]  # True label (40, 512, 512)\n            pred = torch.argmax(preds[0], dim=0)  # Predicted label (512, 512), argmax over class dim\n\n            # Convert true label and predicted label back to color images\n            true_label_img = one_hot_to_color(true_label, class_to_color)  # Ground truth image (H, W, 3)\n            pred_img = one_hot_to_color(torch.nn.functional.one_hot(pred, num_classes=40).permute(2, 0, 1), class_to_color)  # Predicted image (H, W, 3)\n\n            # Plot original image, true label, and predicted label\n            fig, axs = plt.subplots(1, 3, figsize=(15, 5))\n\n            axs[0].imshow(image)\n            axs[0].set_title('Original Image')\n            axs[0].axis('off')\n\n            axs[1].imshow(true_label_img)\n            axs[1].set_title('True Label')\n            axs[1].axis('off')\n\n            axs[2].imshow(pred_img)\n            axs[2].set_title('Predicted Output')\n            axs[2].axis('off')\n\n            plt.show()\n\n            break  # Display the first batch and stop\n            \ndef convert_to_one_hot(segmented_img, class_to_color):\n    # Ensure the tensor is on CPU and convert to numpy (H, W, 3)\n    segmented_img_np = segmented_img.permute(1, 2, 0).cpu().numpy()  # Convert (3, H, W) -> (H, W, 3)\n\n    # Initialize an empty array for the one-hot encoded result (40, H, W)\n    H, W, _ = segmented_img_np.shape\n    one_hot_encoded = np.zeros((len(class_to_color), H, W), dtype=np.uint8)\n\n    # Iterate over each class and its corresponding color\n    for class_index, color in class_to_color.items():\n        # Create a mask where all channels (R, G, B) match the current class color\n        mask = np.all(segmented_img_np == color, axis=-1)  # Compare pixel-wise across channels\n\n        # Check if the color is (0, 0, 0) and set the corresponding channel to class 35\n        if np.array_equal(color, (0, 0, 0)):\n            one_hot_encoded[35, mask] = 1  # Set class 35 for (0, 0, 0) pixels\n        else:\n            # Set the corresponding channel to 1 where the mask is True\n            one_hot_encoded[class_index, mask] = 1\n\n    # Convert the result to a torch tensor (shape: (40, H, W))\n    one_hot_encoded_tensor = torch.tensor(one_hot_encoded, dtype=torch.float32)\n\n    return one_hot_encoded_tensor\n\ndef visualize_predictions_TEST(model, dataloader, class_to_color, device):\n    model.eval()  # Set the model to evaluation mode\n\n    with torch.inference_mode():  # No need to calculate gradients\n        for images in dataloader:  # Get a batch of images and true labels\n            # Move images and labels to the device\n            images = images.to(device)\n            #true_labels = true_labels.to(device)\n\n            # Get model output (predictions)\n            preds = model(images)  # (2, 40, 512, 512)\n\n            # Take the first image in the batch\n            image = images[0].cpu().permute(1, 2, 0).numpy()  # Convert to (H, W, 3) format\n            #true_label = true_labels[0]  # True label (40, 512, 512)\n            pred = torch.argmax(preds[0], dim=0)  # Predicted label (512, 512), argmax over class dim\n\n            # Convert true label and predicted label back to color images\n            #true_label_img = one_hot_to_color(true_label, class_to_color)  # Ground truth image (H, W, 3)\n            pred_img = one_hot_to_color(torch.nn.functional.one_hot(pred, num_classes=40).permute(2, 0, 1), class_to_color)  # Predicted image (H, W, 3)\n\n            # Plot original image, true label, and predicted label\n            fig, axs = plt.subplots(1, 3, figsize=(15, 5))\n\n            axs[0].imshow(image)\n            axs[0].set_title('Original Image')\n            axs[0].axis('off')\n\n            #axs[1].imshow(true_label_img)\n            #axs[1].set_title('True Label')\n            #axs[1].axis('off')\n\n            axs[2].imshow(pred_img)\n            axs[2].set_title('Predicted Output')\n            axs[2].axis('off')\n\n            plt.show()\n\n            break \n            \ndef save_predictions(model, dataloader, output_dir, device):\n    \"\"\"\n    Predict on test images, convert predictions to images, and save them in the specified directory.\n    \n    Args:\n        model: The trained model.\n        dataloader: DataLoader for the test set (returns both image tensors and their file paths).\n        output_dir: Directory to save the images.\n        device: Device to run the model on (CPU/GPU).\n    \"\"\"\n    # Ensure the output directory exists\n    os.makedirs(output_dir, exist_ok=True)\n\n    model.eval()  # Set the model to evaluation mode\n    with torch.no_grad():  # Disable gradient calculation for inference\n        for batch_idx, (images, image_paths) in tqdm(enumerate(dataloader), total=len(dataloader), desc=\"Predicting\"):\n            images = images.to(device)\n            \n            # Make predictions (output is one-hot encoded of shape (bs, 40, 512, 512))\n            outputs = model(images)  # Shape: (bs, 40, 512, 512)\n            \n            # Get predicted classes (argmax along the class dimension)\n            predictions = torch.argmax(outputs, dim=1)  # Shape: (bs, 512, 512)\n            \n            # Process each image in the batch\n            for i in range(predictions.shape[0]):  \n                one_hot_encoded = predictions[i] \n                \n                # Convert one-hot encoded output to a color image using one_hot_to_color\n                color_image_numpy = one_hot_to_color(torch.nn.functional.one_hot(one_hot_encoded, num_classes=40).permute(2, 0, 1), class_to_color)  # Should return a numpy array of shape (512, 512, 3)\n                \n                # Convert numpy array to a PIL image\n                color_image_pil = Image.fromarray(np.uint8(color_image_numpy))  # Assuming the numpy array is in [0, 255] range\n                \n                # Resize the image to (1080, 1920)\n                resized_image = color_image_pil.resize((1920, 1080), Image.Resampling.LANCZOS)\n                \n                # Extract the file name from the path (without extension)\n                image_path = image_paths[i]\n                image_name = os.path.splitext(os.path.basename(image_path))[0]  # Extract name like 'frame11003'\n                \n                # Save the resized image with the original name as its output\n                output_image_path = os.path.join(output_dir, f\"{image_name}.png\")\n                resized_image.save(output_image_path)\n                print(f\"Saved: {output_image_path}\")\n\n                \n# Function to convert the segmented image into polygon-based format and prepare for CSV\ndef image_to_polygon_format(image_path):\n    if not os.path.exists(image_path):\n        raise FileNotFoundError(f\"Image path '{image_path}' not found.\")\n    \n    segmented_image = cv2.imread(image_path)  # Load the segmented image\n    if segmented_image is None:\n        raise ValueError(f\"Failed to load image at '{image_path}'. Make sure the file is a valid image.\")\n\n    height, width, _ = segmented_image.shape\n\n    objects = []  # To hold the polygon data for each object\n\n    # Loop through all the labels\n    for label in labels:\n        if label.ignoreInEval:\n            continue  # Skip labels that should be ignored\n\n        # Convert the label color to a binary mask\n        mask = cv2.inRange(segmented_image, np.array(label.color), np.array(label.color))\n\n        # Find contours (polygons) in the binary mask\n        contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n\n        for contour in contours:\n            if len(contour) < 3:  # Skip small or invalid polygons\n                continue\n\n            # Simplify the contour to polygon format\n            polygon = contour.reshape(-1, 2).tolist()  # Flatten the contour to a list of points\n\n            # Create the object data for the current label and polygon\n            object_data = {\n                \"label\": label.name,\n                \"polygon\": polygon\n            }\n\n            # Append the object to the result\n            objects.append(object_data)\n\n    return objects\n\n# Function to save to the CSV format as requested\ndef save_to_csv(objects_dict, output_csv_path):\n    # Create a DataFrame with id and objects\n    data = [{\"id\": filename, \"objects\": json.dumps(objects)} for filename, objects in objects_dict.items()]\n    df = pd.DataFrame(data)\n    \n    # Save to CSV\n    df.to_csv(output_csv_path, index=False)\n\n# Main function to process multiple images and save the output CSV\ndef process_images(image_folder, output_csv_path):\n    if not os.path.exists(image_folder):\n        raise FileNotFoundError(f\"Image folder '{image_folder}' not found.\")\n    \n    objects_dict = {}\n\n    # Process each image and store the objects for each row ID\n    segmented_image_paths = [os.path.join(image_folder, f) for f in os.listdir(image_folder) if os.path.isfile(os.path.join(image_folder, f))]\n\n    for image_path in segmented_image_paths:\n        # Get the filename without the extension\n        filename = os.path.splitext(os.path.basename(image_path))[0].replace('_leftImg8bit', '')\n        polygon_data = image_to_polygon_format(image_path)\n        objects_dict[filename] = polygon_data\n\n    # Save the results to CSV\n    save_to_csv(objects_dict, output_csv_path)","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:16.947889Z","iopub.execute_input":"2025-06-27T07:19:16.94814Z","iopub.status.idle":"2025-06-27T07:19:16.971709Z","shell.execute_reply.started":"2025-06-27T07:19:16.948117Z","shell.execute_reply":"2025-06-27T07:19:16.970943Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset Class","metadata":{}},{"cell_type":"code","source":"class SegmentationDataset(Dataset):\n    def __init__ (self, image_dir, label_dir, class_to_color, train = True):\n        super().__init__()\n        \n        self.image_paths = list(sorted(glob(image_dir +\"/*/*\")))\n        self.label_paths =  sorted([\n                                os.path.join(label_dir, scene_dir, file_name)\n                                for scene_dir in os.listdir(label_dir)\n                                if os.path.isdir(os.path.join(label_dir, scene_dir))  # Ensure it's a directory\n                                for file_name in os.listdir(os.path.join(label_dir, scene_dir))  # List files in the scene directory\n                                if file_name.endswith('_gtFine_labelColors.png')  # Check for the desired file ending\n                            ])\n         \n            \n        image_train, image_val, label_train, label_val = train_test_split(self.image_paths, self.label_paths, shuffle=True, test_size = 0.2, random_state=2000)\n\n        if train:\n            self.image_paths = image_train\n            self.label_paths = label_train\n        else:\n            self.image_paths = image_val\n            self.label_paths = label_val\n            \n            \n        self.transform = A.Compose([\n            A.Resize(512, 512)\n        ])\n    \n    @staticmethod\n    def convert_to_one_hot(segmented_img, class_to_color):\n        # Ensure the tensor is on CPU and convert to numpy (H, W, 3)\n        segmented_img_np = segmented_img.permute(1, 2, 0).cpu().numpy()  # Convert (3, H, W) -> (H, W, 3)\n\n        # Initialize an empty array for the one-hot encoded result (40, H, W)\n        H, W, _ = segmented_img_np.shape\n        one_hot_encoded = np.zeros((len(class_to_color), H, W), dtype=np.uint8)\n\n        # Iterate over each class and its corresponding color\n        for class_index, color in class_to_color.items():\n            # Create a mask where all channels (R, G, B) match the current class color\n            mask = np.all(segmented_img_np == color, axis=-1)  # Compare pixel-wise across channels\n\n            # Check if the color is (0, 0, 0) and set the corresponding channel to class 35\n            if np.array_equal(color, (0, 0, 0)):\n                one_hot_encoded[35, mask] = 1  # Set class 35 for (0, 0, 0) pixels\n            else:\n                # Set the corresponding channel to 1 where the mask is True\n                one_hot_encoded[class_index, mask] = 1\n\n        # Convert the result to a torch tensor (shape: (40, H, W))\n        one_hot_encoded_tensor = torch.tensor(one_hot_encoded, dtype=torch.float32)\n\n        return one_hot_encoded_tensor\n\n\n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self, idx):\n        \n        img_path = self.image_paths[idx]\n        img = Image.open(img_path)\n        label_path = self.label_paths[idx]\n        label = Image.open(label_path).convert(\"RGB\")\n        \n        transformed = self.transform(image = np.array(img) , mask = np.array(label))\n        \n        img = transformed['image']\n        label = transformed['mask']\n        \n        to_tensor = torchvision.transforms.ToTensor()\n        img = to_tensor(img)\n        label = to_tensor(label)\n        label = self.convert_to_one_hot(label, class_to_color)\n        \n        return img,label","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:16.972702Z","iopub.execute_input":"2025-06-27T07:19:16.972952Z","iopub.status.idle":"2025-06-27T07:19:16.989011Z","shell.execute_reply.started":"2025-06-27T07:19:16.972928Z","shell.execute_reply":"2025-06-27T07:19:16.988208Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img_dir = \"/kaggle/input/iitg-ai-overnight-hackathon-2024/dataset/dataset/train\"\nlabel_dir = \"/kaggle/input/iitg-ai-overnight-hackathon-2024/dataset/dataset/labels\"\n\ntrain_dataset = SegmentationDataset(img_dir, label_dir,class_to_color, train=True)\ntest_dataset = SegmentationDataset(img_dir, label_dir,class_to_color, train=False)\n\nbatch_size = 2\n\ntrain_dataloader = DataLoader(dataset = train_dataset, batch_size = batch_size, shuffle = True )\ntest_dataloader = DataLoader(dataset = test_dataset, batch_size = batch_size, shuffle = True )","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:16.990052Z","iopub.execute_input":"2025-06-27T07:19:16.990285Z","iopub.status.idle":"2025-06-27T07:19:23.135715Z","shell.execute_reply.started":"2025-06-27T07:19:16.990263Z","shell.execute_reply":"2025-06-27T07:19:23.13485Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the UNet model with pretrained encoder and 40 output classes (for one-hot style)\nmodel = smp.Unet(\n    encoder_name=\"resnet34\",     # You can choose other encoders too\n    encoder_weights=\"imagenet\",  # Use pretrained weights\n    in_channels=3,               # RGB input\n    classes=40                   # Number of classes\n)\n\n# Move model to device (GPU if available)\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = model.to(device)\n\n# Loss function — note: CrossEntropyLoss expects raw logits and class indices (not one-hot)\nloss = nn.CrossEntropyLoss()\n\n# Optimizer\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-06-27T07:24:54.206942Z","iopub.execute_input":"2025-06-27T07:24:54.207266Z","iopub.status.idle":"2025-06-27T07:24:54.687374Z","shell.execute_reply.started":"2025-06-27T07:24:54.207238Z","shell.execute_reply":"2025-06-27T07:24:54.686353Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"state_dict = torch.load(\"/kaggle/input/final_unet/pytorch/default/1/model_epoch_2 (1).pth\")\nmodel.load_state_dict(state_dict)","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:24:56.434931Z","iopub.execute_input":"2025-06-27T07:24:56.435527Z","iopub.status.idle":"2025-06-27T07:24:56.547712Z","shell.execute_reply.started":"2025-06-27T07:24:56.435497Z","shell.execute_reply":"2025-06-27T07:24:56.546793Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training Loop\n","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\nfrom tqdm import tqdm\n\n# Assuming you've already defined your DataLoader, model, loss function, and optimizer.\n# You should have `train_dataloader`, `test_dataloader`, `model`, `loss`, `optimizer` defined.\n\n# Number of epochs\nnum_epochs = 20\n\n# Move model to GPU if available\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(device)\n\n# Training loop\nfor epoch in range(num_epochs):\n    model.train()  # Set the model to training mode\n    running_train_loss = 0.0\n    \n    # Training phase\n    with tqdm(total=len(train_dataloader), desc=f'Train Epoch {epoch + 1}/{num_epochs}', unit='batch') as pbar_train:\n        for images, labels in train_dataloader:\n            images = images.to(device)\n            labels = labels.to(device)\n\n            optimizer.zero_grad()  # Zero the gradients\n\n            # Forward pass\n            outputs = model(images)\n            \n            # Compute loss\n            loss_value = loss(outputs, labels)\n            \n            # Backward pass and optimization\n            loss_value.backward()\n            optimizer.step()\n\n            # Accumulate training loss\n            running_train_loss += loss_value.item()\n\n            # Update progress bar\n            pbar_train.set_postfix(train_loss=loss_value.item())\n            pbar_train.update(1)\n\n    # Calculate average training loss for the epoch\n    avg_train_loss = running_train_loss / len(train_dataloader)\n    print(f\"Epoch [{epoch + 1}/{num_epochs}], Average Training Loss: {avg_train_loss:.4f}\")\n    \n    with tqdm(total=len(test_dataloader), desc=f'Test Epoch {epoch + 1}/{num_epochs}', unit='batch') as pbar_test:# Validation phase (evaluate on test set)\n        model.eval()  # Set the model to evaluation mode\n        running_test_loss = 0.0\n\n        with torch.no_grad():\n            for images, labels in test_dataloader:\n                images = images.to(device)\n                labels = labels.to(device)\n\n                # Forward pass\n                outputs = model(images)\n\n                # Compute loss\n                loss_value = loss(outputs, labels)\n\n                # Accumulate test loss\n                running_test_loss += loss_value.item()\n\n                # Update progress bar\n                pbar_test.set_postfix(test_loss=loss_value.item())\n                pbar_test.update(1)\n\n    # Calculate average test loss for the epoch\n    avg_test_loss = running_test_loss / len(test_dataloader)\n    print(f\"Epoch [{epoch + 1}/{num_epochs}], Average Test Loss: {avg_test_loss:.4f}\")\n\n    # Optional: Save the model after every epoch\n    torch.save(model.state_dict(), f\"model_epoch_{epoch + 1}.pth\")\n\nprint(\"Training complete.\")\n","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:24:58.488825Z","iopub.execute_input":"2025-06-27T07:24:58.489676Z","iopub.status.idle":"2025-06-27T07:25:43.046267Z","shell.execute_reply.started":"2025-06-27T07:24:58.489639Z","shell.execute_reply":"2025-06-27T07:25:43.044896Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(device)","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.490744Z","iopub.status.idle":"2025-06-27T07:19:24.491061Z","shell.execute_reply.started":"2025-06-27T07:19:24.490917Z","shell.execute_reply":"2025-06-27T07:19:24.490932Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_predictions(model, test_dataloader, class_to_color, device)","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.492625Z","iopub.status.idle":"2025-06-27T07:19:24.49308Z","shell.execute_reply.started":"2025-06-27T07:19:24.492854Z","shell.execute_reply":"2025-06-27T07:19:24.492877Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, test_dir):\n        super().__init__()\n        \n        self.image_paths = list(sorted(glob(test_dir +\"/*\")))\n        \n        \n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self,idx):\n        \n        img_path = self.image_paths[idx]\n        img = Image.open(img_path)\n        img = img.resize((512,512))\n        \n        to_tensor = torchvision.transforms.ToTensor()\n        img = to_tensor(img)\n        \n        return img,img_path\n        ","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.494099Z","iopub.status.idle":"2025-06-27T07:19:24.494508Z","shell.execute_reply.started":"2025-06-27T07:19:24.494294Z","shell.execute_reply":"2025-06-27T07:19:24.494316Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TEST_DATASET = TestDataset(test_dir = '/kaggle/input/iitg-ai-overnight-hackathon-2024/dataset/dataset/test')","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.495673Z","iopub.status.idle":"2025-06-27T07:19:24.496104Z","shell.execute_reply.started":"2025-06-27T07:19:24.495889Z","shell.execute_reply":"2025-06-27T07:19:24.495911Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(TEST_DATASET)","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.497523Z","iopub.status.idle":"2025-06-27T07:19:24.497785Z","shell.execute_reply.started":"2025-06-27T07:19:24.497656Z","shell.execute_reply":"2025-06-27T07:19:24.49767Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TEST_DATALOADER = DataLoader(dataset = TEST_DATASET, batch_size = 1, shuffle = False)","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.498586Z","iopub.status.idle":"2025-06-27T07:19:24.49891Z","shell.execute_reply.started":"2025-06-27T07:19:24.498729Z","shell.execute_reply":"2025-06-27T07:19:24.498743Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pred_dir = \"/kaggle/working/predictions\"\nsave_predictions(model = model, dataloader = TEST_DATALOADER, output_dir = pred_dir, device = device)","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.499972Z","iopub.status.idle":"2025-06-27T07:19:24.500236Z","shell.execute_reply.started":"2025-06-27T07:19:24.500107Z","shell.execute_reply":"2025-06-27T07:19:24.50012Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_folder = \"/kaggle/working/predictions\"\noutput_csv_path = \"/kaggle/working/Submission_2.csv\"\n\nprocess_images(image_folder, output_csv_path)","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.501019Z","iopub.status.idle":"2025-06-27T07:19:24.501291Z","shell.execute_reply.started":"2025-06-27T07:19:24.501158Z","shell.execute_reply":"2025-06-27T07:19:24.501171Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img = Image.open(\"/kaggle/input/iitg-ai-overnight-hackathon-2024/dataset/dataset/test/frame11003_leftImg8bit.jpg\")","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.502225Z","iopub.status.idle":"2025-06-27T07:19:24.502513Z","shell.execute_reply.started":"2025-06-27T07:19:24.502374Z","shell.execute_reply":"2025-06-27T07:19:24.502389Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img = np.array(img)","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.503874Z","iopub.status.idle":"2025-06-27T07:19:24.504164Z","shell.execute_reply.started":"2025-06-27T07:19:24.504026Z","shell.execute_reply":"2025-06-27T07:19:24.50404Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"img.shape","metadata":{"execution":{"iopub.status.busy":"2025-06-27T07:19:24.504994Z","iopub.status.idle":"2025-06-27T07:19:24.505257Z","shell.execute_reply.started":"2025-06-27T07:19:24.505126Z","shell.execute_reply":"2025-06-27T07:19:24.50514Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}