{"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":"none","dataSources":[{"sourceId":86142,"databundleVersionId":9786425,"sourceType":"competition"}],"dockerImageVersionId":30786,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport tensorflow as tf\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm import tqdm \nfrom PIL import Image\nimport os\nimport warnings\nfrom tensorflow.keras.layers import Input, Conv2D, Conv2DTranspose, MaxPooling2D, Dropout, concatenate\nfrom tensorflow.keras.callbacks import ModelCheckpoint\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:32:35.356155Z","iopub.execute_input":"2024-10-12T14:32:35.356621Z","iopub.status.idle":"2024-10-12T14:32:35.362937Z","shell.execute_reply.started":"2024-10-12T14:32:35.356581Z","shell.execute_reply":"2024-10-12T14:32:35.361511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import namedtuple\n\nLabel = namedtuple( 'Label' , [\n\n    'name'        ,\n    'id'          ,\n    'csId',\n    'csTrainId',\n    'level4id',\n    'level3id',\n    'category',\n    'level2id',\n    'level1Id',\n    'hasInstances',\n    'ingnoreInEval',\n    '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]  ","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:32:35.395798Z","iopub.execute_input":"2024-10-12T14:32:35.396208Z","iopub.status.idle":"2024-10-12T14:32:35.422510Z","shell.execute_reply.started":"2024-10-12T14:32:35.396170Z","shell.execute_reply":"2024-10-12T14:32:35.421285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"N_FILTERS = 64\nKERNEL_SIZE = 3\nN_CLASSES = len(labels)\nIMAGE_SIZE = [128,128]\nIMAGE_SHAPE = IMAGE_SIZE + [3,]\nEPOCHS = 2\nBATCH_SIZE = 16\nMODEL_CHECKPOINT_FILEPATH = '/kaggle/working/imagesegmented-unet.weights.h5'\n\nid2color = { label.id : np.asarray(label.color) for label in labels }\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:32:35.424349Z","iopub.execute_input":"2024-10-12T14:32:35.424731Z","iopub.status.idle":"2024-10-12T14:32:35.439077Z","shell.execute_reply.started":"2024-10-12T14:32:35.424682Z","shell.execute_reply":"2024-10-12T14:32:35.437911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def closest_labels(mask, mapping):\n    \n    closest_distance = np.full([mask.shape[0], mask.shape[1]], 10000) \n    closest_category = np.full([mask.shape[0], mask.shape[1]], None)   \n\n    for id, color in mapping.items(): \n        dist = np.sqrt(np.linalg.norm(mask - color.reshape([1,1,-1]), axis=-1))\n        is_closer = closest_distance > dist\n        closest_distance = np.where(is_closer, dist, closest_distance)\n        closest_category = np.where(is_closer, id, closest_category)\n    \n    return closest_category","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:32:35.440574Z","iopub.execute_input":"2024-10-12T14:32:35.441175Z","iopub.status.idle":"2024-10-12T14:32:35.448996Z","shell.execute_reply.started":"2024-10-12T14:32:35.441115Z","shell.execute_reply":"2024-10-12T14:32:35.447778Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport numpy as np\nfrom PIL import Image\nfrom tqdm import tqdm\n\ndef load_images(train_root_filepath, mask_root_filepath, image_size=(128, 128)):\n    train_images = []\n    train_masks = []\n    train_masks_enc = []\n    for subdir in tqdm(os.listdir(train_root_filepath), desc='Loading images from subdirectories'):\n        subdir_path = os.path.join(train_root_filepath, subdir)\n        if os.path.isdir(subdir_path):  \n            for train_file in os.listdir(subdir_path):\n                if train_file.endswith('_leftImg8bit.jpg'):\n                    train_image_path = os.path.join(subdir_path, train_file)\n                    train_image = Image.open(train_image_path).convert('RGB').resize(image_size)\n                    train_images.append(np.array(train_image) / 255.0)\n                    mask_file = train_file.replace('_leftImg8bit.jpg', '_gtFine_labelColors.png')\n                    mask_image_path = os.path.join(mask_root_filepath, subdir, mask_file)\n\n                    if os.path.exists(mask_image_path):\n                        mask_image = Image.open(mask_image_path).convert('RGB').resize(image_size)\n                        mask_array = np.array(mask_image)\n\n                        mask_enc = closest_labels(mask_array, id2color)\n                        \n                        train_masks.append(mask_array)\n                        train_masks_enc.append(mask_enc)\n                    else:\n                        print(f\"Warning: Mask file {mask_file} not found in {mask_root_filepath}\")\n\n    return train_images, train_masks, train_masks_enc\n\ntrain_filepath = '/kaggle/input/iitg-ai-overnight-hackathon-2024/dataset/dataset/train'  \nmask_filepath = '/kaggle/input/iitg-ai-overnight-hackathon-2024/dataset/dataset/labels' \n\ntrain_images, train_masks, train_masks_enc = load_images(train_filepath, mask_filepath)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:32:35.451242Z","iopub.execute_input":"2024-10-12T14:32:35.451621Z","iopub.status.idle":"2024-10-12T14:52:13.998742Z","shell.execute_reply.started":"2024-10-12T14:32:35.451584Z","shell.execute_reply":"2024-10-12T14:52:13.997746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=[20, 14])\n\nfor i in range(2):\n    img = train_images[i]\n    msk = train_masks[i]\n    enc = train_masks_enc[i]\n    tmp = np.zeros([enc.shape[0], enc.shape[1], 3])\n    \n    for row in range(enc.shape[0]):\n        for col in range(enc.shape[1]):\n            tmp[row, col, :] = id2color[enc[row, col]]\n            tmp = tmp.astype('uint8')\n            \n    plt.subplot(2, 3, i*3 + 1)\n    plt.imshow(img)\n    plt.axis('off')\n    plt.gca().set_title('Sample Image {}'.format(str(i+1)))\n    \n    plt.subplot(2, 3, i*3 + 2)\n    plt.imshow(msk)\n    plt.axis('off')\n    plt.gca().set_title('Sample Mask {}'.format(str(i+1)))\n    \n    plt.subplot(2, 3, i*3 + 3)\n    plt.imshow(tmp)\n    plt.axis('off')\n    plt.gca().set_title('Sample Encoded Mask {}'.format(str(i+1)))\n    \nplt.subplots_adjust(wspace=0, hspace=0.1)","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:52:14.000187Z","iopub.execute_input":"2024-10-12T14:52:14.000526Z","iopub.status.idle":"2024-10-12T14:52:14.886192Z","shell.execute_reply.started":"2024-10-12T14:52:14.000491Z","shell.execute_reply":"2024-10-12T14:52:14.885135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_images=np.stack(train_images).astype('float32')\ntrain_masks_enc=np.stack(train_masks_enc).astype('float32')","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:52:14.888876Z","iopub.execute_input":"2024-10-12T14:52:14.889289Z","iopub.status.idle":"2024-10-12T14:52:23.451832Z","shell.execute_reply.started":"2024-10-12T14:52:14.889247Z","shell.execute_reply":"2024-10-12T14:52:23.450760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_masks_enc.shape)","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:52:23.455526Z","iopub.execute_input":"2024-10-12T14:52:23.455972Z","iopub.status.idle":"2024-10-12T14:52:23.461888Z","shell.execute_reply.started":"2024-10-12T14:52:23.455924Z","shell.execute_reply":"2024-10-12T14:52:23.460705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"Train Images Shape: {len(train_images)}, Shape: {(train_images[0].shape)}\")\nprint(f\"Train Masks Shape: {len(train_masks)}, Shape: {(train_masks[0].shape)}\")\nprint(f\"Train Masks Encoded Shape: {len(train_masks_enc)}, Shape: {(train_masks_enc[0].shape)}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:52:23.463611Z","iopub.execute_input":"2024-10-12T14:52:23.464906Z","iopub.status.idle":"2024-10-12T14:52:23.475341Z","shell.execute_reply.started":"2024-10-12T14:52:23.464854Z","shell.execute_reply":"2024-10-12T14:52:23.474184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def conv_block(inputs=None, n_filters=32, kernel_size = 3, dropout_prob = 0, max_pooling=True):\n    \n    conv = Conv2D(n_filters, \n                  kernel_size = 3, \n                  activation = 'relu',\n                  padding = 'same',\n                  kernel_initializer = 'he_normal')(inputs)\n    conv = Conv2D(n_filters, \n                  kernel_size = 3,  \n                  activation = 'relu',\n                  padding = 'same',\n                  kernel_initializer = 'he_normal')(conv)\n    if dropout_prob > 0:\n        conv = Dropout(dropout_prob)(conv)\n    if max_pooling:\n        next_layer = MaxPooling2D(pool_size = (2,2))(conv)\n    else:\n        next_layer = conv\n        \n    skip_connection = conv\n    \n    return next_layer, skip_connection\n\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:52:23.477106Z","iopub.execute_input":"2024-10-12T14:52:23.478382Z","iopub.status.idle":"2024-10-12T14:52:23.488881Z","shell.execute_reply.started":"2024-10-12T14:52:23.478339Z","shell.execute_reply":"2024-10-12T14:52:23.487769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def upsampling_block(expansive_input, contractive_input, n_filters=32, kernel_size = 3):\n    \n    up = Conv2DTranspose(\n                 n_filters,  \n                 kernel_size = kernel_size,   \n                 strides = (2,2),\n                 padding = 'same')(expansive_input)\n    merge = concatenate([up, contractive_input], axis=3)\n    \n    conv = Conv2D(n_filters,  \n                 kernel_size = (3,3),  \n                 activation='relu',\n                 padding='same',\n                 kernel_initializer='he_normal')(merge)\n    conv = Conv2D(n_filters, \n                 kernel_size = (3,3), \n                 activation='relu',\n                 padding='same',\n                 kernel_initializer='he_normal')(conv)\n    \n    return conv","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:52:23.490223Z","iopub.execute_input":"2024-10-12T14:52:23.490628Z","iopub.status.idle":"2024-10-12T14:52:23.504486Z","shell.execute_reply.started":"2024-10-12T14:52:23.490581Z","shell.execute_reply":"2024-10-12T14:52:23.502979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_unet_model(image_shape, n_filters, kernel_size, n_classes):\n\n    inputs = Input(image_shape)\n\n    # Contracting Path (encoding)\n    cblock1 = conv_block(inputs, n_filters, kernel_size)\n    cblock2 = conv_block(cblock1[0], n_filters * 2, kernel_size)\n    cblock3 = conv_block(cblock2[0], n_filters * 4, kernel_size, dropout_prob = 0.3)\n    cblock4 = conv_block(cblock3[0], n_filters * 8, kernel_size, dropout_prob = 0.3) # Include a dropout_prob of 0.3 for this layer\n    cblock5 = conv_block(cblock4[0], n_filters * 16, kernel_size, dropout_prob = 0.3, max_pooling=False) \n\n\n    ublock6 = upsampling_block(cblock5[0], cblock4[1], n_filters * 8, kernel_size)\n    ublock7 = upsampling_block(ublock6, cblock3[1], n_filters * 4, kernel_size)\n    ublock8 = upsampling_block(ublock7, cblock2[1], n_filters * 2, kernel_size)\n    ublock9 = upsampling_block(ublock8, cblock1[1], n_filters, kernel_size)\n\n    conv9 = Conv2D(n_filters,\n                 kernel_size = kernel_size,\n                 activation='relu',\n                 padding='same',\n                 kernel_initializer='he_normal')(ublock9)\n\n    conv10 = Conv2D(n_classes, kernel_size = 1, padding='same')(conv9)\n\n    model = tf.keras.Model(inputs=inputs, outputs=conv10)\n    \n    return model","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:52:23.509202Z","iopub.execute_input":"2024-10-12T14:52:23.509649Z","iopub.status.idle":"2024-10-12T14:52:23.519978Z","shell.execute_reply.started":"2024-10-12T14:52:23.509608Z","shell.execute_reply":"2024-10-12T14:52:23.518471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nmodel = create_unet_model(IMAGE_SHAPE, N_FILTERS, KERNEL_SIZE, N_CLASSES)\n\ntf.keras.utils.plot_model(model, show_shapes = True)","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:52:23.521487Z","iopub.execute_input":"2024-10-12T14:52:23.521830Z","iopub.status.idle":"2024-10-12T14:52:27.394572Z","shell.execute_reply.started":"2024-10-12T14:52:23.521795Z","shell.execute_reply":"2024-10-12T14:52:27.391357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_size = int(0.8 * len(train_images)) \nx_train, x_val = train_images[:train_size], train_images[train_size:]\ny_train, y_val = train_masks_enc[:train_size], train_masks_enc[train_size:]","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:52:27.396670Z","iopub.execute_input":"2024-10-12T14:52:27.397139Z","iopub.status.idle":"2024-10-12T14:52:27.416536Z","shell.execute_reply.started":"2024-10-12T14:52:27.397087Z","shell.execute_reply":"2024-10-12T14:52:27.409904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tensorflow.keras.callbacks import ModelCheckpoint\n\n\nmodel.compile(optimizer='adam',\n              loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),\n              metrics=['accuracy'])\n\n# Define the model checkpoint\nmodel_checkpoint = ModelCheckpoint(MODEL_CHECKPOINT_FILEPATH,\n                                   monitor='val_accuracy',\n                                   save_best_only=True,\n                                   save_weights_only=True,\n                                   verbose=1,\n                                   mode='max',\n                                  save_freq=1)\n\ncallbacks = [model_checkpoint]\n\n\nhistory = model.fit(x=x_train,\n                    y=y_train,\n                    validation_data=(x_val, y_val),\n                    batch_size=BATCH_SIZE,\n                    epochs=EPOCHS,\n                    callbacks=callbacks)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T14:52:27.425020Z","iopub.execute_input":"2024-10-12T14:52:27.426995Z","iopub.status.idle":"2024-10-12T17:30:11.739407Z","shell.execute_reply.started":"2024-10-12T14:52:27.426924Z","shell.execute_reply":"2024-10-12T17:30:11.737077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, (ax1, ax2) = plt.subplots(1,2, figsize=(16,6))\ntitle_fontsize = 16\naxis_fontsize = 12\n\nax1.plot(range(1, EPOCHS + 1), history.history['loss'], marker='o', label='Training loss')\nax1.plot(range(1, EPOCHS + 1), history.history['val_loss'], marker='o', label='Validation Loss')\nax1.legend()\nax1.set_xticks(range(1, EPOCHS + 1))\nax1.set_title('Loss', fontsize=title_fontsize)\nax1.set_xlabel('Epoch', fontsize=axis_fontsize)\n\nax2.plot(range(1, EPOCHS + 1), history.history['accuracy'], marker='o', label='Training Accuracy')\nax2.plot(range(1, EPOCHS + 1), history.history['val_accuracy'], marker='o', label='Validation Accuracy')\nax2.legend()\nax2.set_xticks(range(1, EPOCHS + 1))\nax2.set_title('Accuracy', fontsize=title_fontsize)\nax2.set_xlabel('Epoch', fontsize=axis_fontsize);","metadata":{"execution":{"iopub.status.busy":"2024-10-12T17:32:17.305908Z","iopub.execute_input":"2024-10-12T17:32:17.306749Z","iopub.status.idle":"2024-10-12T17:32:17.794895Z","shell.execute_reply.started":"2024-10-12T17:32:17.306705Z","shell.execute_reply":"2024-10-12T17:32:17.793795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=[15, 20])\n\nfor i in range(4):    \n    img = train_images[i]\n    enc = train_masks_enc[i]\n    \n    pred = model.predict(img.reshape([1] + IMAGE_SHAPE))\n    pred = np.squeeze(np.argmax(pred, axis=-1))\n    \n    tmp1 = np.zeros([enc.shape[0], enc.shape[1], 3])\n    tmp2 = np.zeros([enc.shape[0], enc.shape[1], 3])\n    \n    \n    for row in range(enc.shape[0]):\n        for col in range(enc.shape[1]):\n            tmp1[row, col, :] = id2color[enc[row, col]]\n            tmp1 = tmp1.astype('uint8')\n                     \n            tmp2[row, col, :] = id2color[pred[row, col]]\n            tmp2 = tmp2.astype('uint8')\n            \n    plt.subplot(4, 3, i*3 + 1)\n    plt.imshow(img)\n    plt.axis('off')\n    plt.gca().set_title('Image {}'.format(str(i+1)))\n    \n    plt.subplot(4, 3, i*3 + 2)\n    plt.imshow(tmp1)\n    plt.axis('off')\n    plt.gca().set_title('Encoded Mask {}'.format(str(i+1)))\n    \n    plt.subplot(4, 3, i*3 + 3)\n    plt.imshow(tmp2)\n    plt.axis('off')\n    plt.gca().set_title('Model Prediction {}'.format(str(i+1)))\n    \nplt.subplots_adjust(wspace=0, hspace=0.1)","metadata":{"execution":{"iopub.status.busy":"2024-10-12T17:32:37.971462Z","iopub.execute_input":"2024-10-12T17:32:37.972280Z","iopub.status.idle":"2024-10-12T17:32:41.890640Z","shell.execute_reply.started":"2024-10-12T17:32:37.972236Z","shell.execute_reply":"2024-10-12T17:32:41.889205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef load_test_images(test_filepath, image_size=(128,128)):\n    test_images = []\n\n    for test_file in tqdm(os.listdir(test_filepath), desc='Loading Test Images: '):\n        if test_file.endswith('_leftImg8bit.jpg'): \n            test_image_path = os.path.join(test_filepath, test_file)\n            test_image = Image.open(test_image_path).convert('RGB').resize(image_size)\n            test_images.append(np.array(test_image) / 255.0) \n    return test_images\n\ndef visualize_predictions(test_images, model, id2color, IMAGE_SHAPE):\n    plt.figure(figsize=[15, 20])\n\n    for i in range(min(4, len(test_images))):  \n        img = test_images[i] \n        \n\n        pred = model.predict(img.reshape([1] + list(IMAGE_SHAPE))) \n        pred = np.squeeze(np.argmax(pred, axis=-1))\n        \n        tmp2 = np.zeros([pred.shape[0], pred.shape[1], 3])\n\n        for row in range(pred.shape[0]):\n            for col in range(pred.shape[1]):\n                tmp2[row, col, :] = id2color[pred[row, col]] \n        tmp2 = tmp2.astype('uint8')\n\n        plt.subplot(4, 2, i*2 + 1)\n        plt.imshow(img)\n        plt.axis('off')\n        plt.gca().set_title('Test Image {}'.format(str(i+1)))\n\n        plt.subplot(4, 2, i*2 + 2)\n        plt.imshow(tmp2)\n        plt.axis('off')\n        plt.gca().set_title('Model Prediction {}'.format(str(i+1)))\n\n    plt.subplots_adjust(wspace=0, hspace=0.1)\n    plt.show()\n\ntest_filepath = '/kaggle/input/iitg-ai-overnight-hackathon-2024/dataset/dataset/test' \nIMAGE_SHAPE = (128,128, 3)\n\ntest_images = load_test_images(test_filepath)\n\nvisualize_predictions(test_images, model, id2color, IMAGE_SHAPE)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T17:32:57.755986Z","iopub.execute_input":"2024-10-12T17:32:57.756407Z","iopub.status.idle":"2024-10-12T17:33:06.029585Z","shell.execute_reply.started":"2024-10-12T17:32:57.756367Z","shell.execute_reply":"2024-10-12T17:33:06.027400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndef get_predicted_masks(test_images, model, IMAGE_SHAPE):\n    predicted_masks = [] \n    for img in tqdm(test_images, desc='Predicting Masks: '):\n\n        pred = model.predict(img.reshape([1] + list(IMAGE_SHAPE))) \n        pred = np.squeeze(np.argmax(pred, axis=-1)) \n\n        predicted_masks.append(pred)\n\n    return predicted_masks \n\npredicted_masks = get_predicted_masks(test_images, model, IMAGE_SHAPE)\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T17:58:34.616074Z","iopub.execute_input":"2024-10-12T17:58:34.616813Z","iopub.status.idle":"2024-10-12T17:59:08.671426Z","shell.execute_reply.started":"2024-10-12T17:58:34.616770Z","shell.execute_reply":"2024-10-12T17:59:08.670427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=[20, 35])\n\nfor i in range(10):  \n    img = test_images[i]\n    enc = predicted_masks[i]\n    tmp = np.zeros([enc.shape[0], enc.shape[1], 3])\n    \n    for row in range(enc.shape[0]):\n        for col in range(enc.shape[1]):\n            tmp[row, col, :] = id2color[enc[row, col]]\n    tmp = tmp.astype('uint8')\n    \n    # Show test image\n    plt.subplot(10, 2, i*2 + 1) \n    plt.imshow(img)\n    plt.axis('off')\n    plt.gca().set_title(f'TEST IMAGE {i+1}')\n    \n    # Show corresponding output image\n    plt.subplot(10, 2, i*2 + 2) \n    plt.imshow(tmp)\n    plt.axis('off')\n    plt.gca().set_title(f'OUTPUT IMAGE {i+1}')\n    \nplt.subplots_adjust(wspace=0.1, hspace=0.5) \nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T17:59:53.201608Z","iopub.execute_input":"2024-10-12T17:59:53.202032Z","iopub.status.idle":"2024-10-12T17:59:55.582165Z","shell.execute_reply.started":"2024-10-12T17:59:53.201994Z","shell.execute_reply":"2024-10-12T17:59:55.580932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_masks=np.stack(predicted_masks).astype('float32')\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T18:00:01.361546Z","iopub.execute_input":"2024-10-12T18:00:01.361977Z","iopub.status.idle":"2024-10-12T18:00:01.376390Z","shell.execute_reply.started":"2024-10-12T18:00:01.361937Z","shell.execute_reply":"2024-10-12T18:00:01.375138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"Predicted Maks: {len(predicted_masks)}, Shape: {(predicted_masks[0].shape)}\")","metadata":{"execution":{"iopub.status.busy":"2024-10-12T18:00:03.361543Z","iopub.execute_input":"2024-10-12T18:00:03.362620Z","iopub.status.idle":"2024-10-12T18:00:03.367980Z","shell.execute_reply.started":"2024-10-12T18:00:03.362571Z","shell.execute_reply":"2024-10-12T18:00:03.366844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\n\ntest_image = np.expand_dims(test_images[0], axis=0) \n\nstart_time = time.time()\n\npredicted_mask = model.predict(test_image)\n\ninference_time = time.time() - start_time\nprint(f\"Inference Time: {inference_time:.4f} seconds per image\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T18:00:05.760893Z","iopub.execute_input":"2024-10-12T18:00:05.761324Z","iopub.status.idle":"2024-10-12T18:00:06.058827Z","shell.execute_reply.started":"2024-10-12T18:00:05.761285Z","shell.execute_reply":"2024-10-12T18:00:06.057701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nmodel.save('/kaggle/working/imagesegmented-unet.weights.h5')\nmodel_size = os.path.getsize('/kaggle/working/imagesegmented-unet.weights.h5') / (1024 * 1024)\nprint(f\"Model Size: {model_size:.2f} MB\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T17:39:08.981894Z","iopub.execute_input":"2024-10-12T17:39:08.982927Z","iopub.status.idle":"2024-10-12T17:39:09.673297Z","shell.execute_reply.started":"2024-10-12T17:39:08.982884Z","shell.execute_reply":"2024-10-12T17:39:09.672148Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_loss, val_accuracy = model.evaluate(x = x_val, y = y_val) \nprint('\\n\\033[1m' + 'The model had an accuracy score of {}%!!'.format(round(100*val_accuracy, 2)) + '\\033[0m')\n# Earlier I had achieved an accuracy of more than 80 percent but because I lost my all progress twice , and ran only for 2 epochs , i still got 71 percent accuracy\n","metadata":{"execution":{"iopub.status.busy":"2024-10-12T17:41:56.620636Z","iopub.execute_input":"2024-10-12T17:41:56.621509Z","iopub.status.idle":"2024-10-12T17:47:19.267569Z","shell.execute_reply.started":"2024-10-12T17:41:56.621465Z","shell.execute_reply":"2024-10-12T17:47:19.266489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}