{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":13333,"databundleVersionId":862146,"sourceType":"competition"},{"sourceId":6536639,"sourceType":"datasetVersion","datasetId":3778994},{"sourceId":7214967,"sourceType":"datasetVersion","datasetId":4175149},{"sourceId":164240128,"sourceType":"kernelVersion"},{"sourceId":164804886,"sourceType":"kernelVersion"}],"dockerImageVersionId":30498,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Table of Contents <a class=\"anchor\"  id=\"contents\"></a>\n* [Installations and Imports](#imports)\n* [Setting up Kaggle API](#kaggle)\n* [Setting up wandb for Experiment Tracking](#wandb)\n* [Read the Data](#read_data)\n    * [Preaparing Train Dataframe](#prep_train)\n    * [Preaparing Test Dataframe](#prep_test)\n* [Visualizing the Images](#viz_img)\n* [Visualizing the Segmentation Masks](#viz_mask)\n* [Visualizing the Images with Segmentation Masks](#viz_img_mask)\n* [Dataset and Data Loader](#data_gen)\n* [Train Test Split](#data_split)\n* [Model Definitions](#model_def)\n* [Custom Losses and Metrics](#custom_objects)\n* [Training the Model](#model_training)\n* [Managing CUDA memory](#manage_cuda)\n* [Post Processing for output masks](#post_process)\n* [Model Evaluation and Exploring predicticted mask on Validation data](#model_evaluation)\n* [Model Prediction on Test data and make submission](#model_prediction)\n* [Acknowledgements](#ack)","metadata":{}},{"cell_type":"markdown","source":" # Installations and Imports <a class=\"anchor\" id=\"imports\"></a>\n [Go back to the Table of Contents](#contents)","metadata":{}},{"cell_type":"code","source":"import os\nimport io\nimport sys\nimport cv2\nimport time\nimport math\nimport tqdm\nimport random\nimport logging\nimport warnings\n\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:24:49.347501Z","iopub.execute_input":"2024-02-29T05:24:49.347856Z","iopub.status.idle":"2024-02-29T05:24:49.498971Z","shell.execute_reply.started":"2024-02-29T05:24:49.347830Z","shell.execute_reply":"2024-02-29T05:24:49.498117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport pandas as pd\nimport albumentations as A\nimport matplotlib.pyplot as plt\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:24:49.500620Z","iopub.execute_input":"2024-02-29T05:24:49.500906Z","iopub.status.idle":"2024-02-29T05:24:54.811736Z","shell.execute_reply.started":"2024-02-29T05:24:49.500881Z","shell.execute_reply":"2024-02-29T05:24:54.810150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"GPU available!\" if torch.cuda.is_available() else \"GPU is not available!\")","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:24:54.813447Z","iopub.execute_input":"2024-02-29T05:24:54.813876Z","iopub.status.idle":"2024-02-29T05:24:54.852697Z","shell.execute_reply.started":"2024-02-29T05:24:54.813837Z","shell.execute_reply":"2024-02-29T05:24:54.851781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install segmentation-models-pytorch\n!pip install efficientnet_pytorch\n#!pip install pytorch-toolbelt\n#!pip install torchvision\n!pip install torchinfo\n!pip install torchview\n#!pip install catalyst","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-02-29T05:24:54.855977Z","iopub.execute_input":"2024-02-29T05:24:54.856352Z","iopub.status.idle":"2024-02-29T05:25:46.641791Z","shell.execute_reply.started":"2024-02-29T05:24:54.856317Z","shell.execute_reply":"2024-02-29T05:25:46.640491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\n#import pytorch_toolbelt\n#import torchvision\nimport torchinfo\nimport torchview\n#import catalyst","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:25:46.643611Z","iopub.execute_input":"2024-02-29T05:25:46.644006Z","iopub.status.idle":"2024-02-29T05:25:49.342328Z","shell.execute_reply.started":"2024-02-29T05:25:46.643964Z","shell.execute_reply":"2024-02-29T05:25:49.341460Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(sys.version)\nprint(torch.__version__)\n#print(torchvision.__version__)\n#print(catalyst.__version__)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:25:49.343520Z","iopub.execute_input":"2024-02-29T05:25:49.343823Z","iopub.status.idle":"2024-02-29T05:25:49.349193Z","shell.execute_reply.started":"2024-02-29T05:25:49.343798Z","shell.execute_reply":"2024-02-29T05:25:49.348229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Gloabal Configuration\n\nSEED = 42\n\nIMG_WIDTH = 2100\nIMG_HEIGHT = 1400\n\n# training params\nR_WIDTH = 1152\nR_HEIGHT = 768\nNUM_CHANNELS = 3\nNUM_CLASSES = 4\nNUM_EPOCHS = 16\nBATCH_SIZE = 4\nTEST_BATCH_SIZE = 32\n\n\n# focal loss params\nALPHA = 0.8\nGAMMA = 2\n\nLABELS = ['Fish', 'Flower', 'Gravel', 'Sugar']\nMINSIZES = [25000 ,21000, 2000, 10000]\nTHRESHOLDS = [0.5, 0.6, 0.3, 0.5]","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:25:49.350232Z","iopub.execute_input":"2024-02-29T05:25:49.350523Z","iopub.status.idle":"2024-02-29T05:25:49.360355Z","shell.execute_reply.started":"2024-02-29T05:25:49.350495Z","shell.execute_reply":"2024-02-29T05:25:49.359352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:25:49.361397Z","iopub.execute_input":"2024-02-29T05:25:49.361655Z","iopub.status.idle":"2024-02-29T05:25:49.370964Z","shell.execute_reply.started":"2024-02-29T05:25:49.361633Z","shell.execute_reply":"2024-02-29T05:25:49.370095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed = SEED\nrandom.seed(seed)\nnp.random.seed(seed)\ntorch.manual_seed(seed)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:25:49.372019Z","iopub.execute_input":"2024-02-29T05:25:49.372311Z","iopub.status.idle":"2024-02-29T05:25:49.386858Z","shell.execute_reply.started":"2024-02-29T05:25:49.372287Z","shell.execute_reply":"2024-02-29T05:25:49.386011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setting up Kaggle API  <a class=\"anchor\" id=\"kaggle\"></a>\n[Go back to the Table of Contents](#contents)","metadata":{}},{"cell_type":"code","source":"import os\nimport json\nfrom kaggle_secrets import UserSecretsClient\n\ndef set_kaggle_api():\n    user_secrets = UserSecretsClient()\n    kaggle_key = user_secrets.get_secret(\"KAGGLE_KEY\")\n    kaggle_username = user_secrets.get_secret(\"KAGGLE_USERNAME\")\n    \n    kaggle_dict = dict(username=kaggle_username, key=kaggle_key)\n    kaggle_json = json.dumps(kaggle_dict)\n\n    os.makedirs('/root/.kaggle', exist_ok=True)\n    with open('/root/.kaggle/kaggle.json', 'w') as f:\n        f.write(kaggle_json)\n    os.chmod('/root/.kaggle/kaggle.json', 0o600)\n    \n    return None\n\nset_kaggle_api()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:25:49.390894Z","iopub.execute_input":"2024-02-29T05:25:49.391186Z","iopub.status.idle":"2024-02-29T05:25:50.046341Z","shell.execute_reply.started":"2024-02-29T05:25:49.391162Z","shell.execute_reply":"2024-02-29T05:25:50.045520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!kaggle competitions submissions understanding_cloud_organization ","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:25:50.047471Z","iopub.execute_input":"2024-02-29T05:25:50.047767Z","iopub.status.idle":"2024-02-29T05:25:51.840299Z","shell.execute_reply.started":"2024-02-29T05:25:50.047742Z","shell.execute_reply":"2024-02-29T05:25:51.839114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Read the Data <a class=\"anchor\" id=\"read_data\"></a>\n[Go back to the Table of Contents](#contents)","metadata":{}},{"cell_type":"code","source":"work_dir = \"/kaggle/working/\"\ndata_dir = '../input/understanding_cloud_organization'\ntrain_csv_path = os.path.join(data_dir,'train.csv')\ntest_csv_path = os.path.join(data_dir,\"sample_submission.csv\")\ntrain_image_path = os.path.join(data_dir,'train_images')\ntest_image_path = os.path.join(data_dir,'test_images')","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:25:51.842208Z","iopub.execute_input":"2024-02-29T05:25:51.842634Z","iopub.status.idle":"2024-02-29T05:25:51.849076Z","shell.execute_reply.started":"2024-02-29T05:25:51.842595Z","shell.execute_reply":"2024-02-29T05:25:51.847886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":" ### Preaparing Train Dataframe <a class=\"anchor\" id=\"prep_train\"></a>\n [Go back to the Table of Contents](#contents)","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(train_csv_path).fillna(-1)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:25:51.850390Z","iopub.execute_input":"2024-02-29T05:25:51.851065Z","iopub.status.idle":"2024-02-29T05:25:55.997683Z","shell.execute_reply.started":"2024-02-29T05:25:51.851015Z","shell.execute_reply":"2024-02-29T05:25:55.996785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['Image_Id'] = train_df['Image_Label'].apply(lambda x: x.split('_')[0])\ntrain_df['Label'] = train_df['Image_Label'].apply(lambda x: x.split('_')[1])\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:12.489872Z","iopub.execute_input":"2024-02-29T05:53:12.490684Z","iopub.status.idle":"2024-02-29T05:53:12.530119Z","shell.execute_reply.started":"2024-02-29T05:53:12.490647Z","shell.execute_reply":"2024-02-29T05:53:12.529122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df['Label_EncodedPixels'] = train_df.apply(lambda row: (row['Label'], row['EncodedPixels']), axis = 1)\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:12.729945Z","iopub.execute_input":"2024-02-29T05:53:12.730712Z","iopub.status.idle":"2024-02-29T05:53:13.065601Z","shell.execute_reply.started":"2024-02-29T05:53:12.730680Z","shell.execute_reply":"2024-02-29T05:53:13.064657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"grouped_EncodedPixels = train_df.groupby('Image_Id')['Label_EncodedPixels'].apply(list)\ngrouped_EncodedPixels.head()\ngrouped_EncodedPixels.info()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:13.067248Z","iopub.execute_input":"2024-02-29T05:53:13.067551Z","iopub.status.idle":"2024-02-29T05:53:13.248570Z","shell.execute_reply.started":"2024-02-29T05:53:13.067527Z","shell.execute_reply":"2024-02-29T05:53:13.247674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = grouped_EncodedPixels.to_frame().reset_index()\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:13.250314Z","iopub.execute_input":"2024-02-29T05:53:13.250964Z","iopub.status.idle":"2024-02-29T05:53:13.270274Z","shell.execute_reply.started":"2024-02-29T05:53:13.250930Z","shell.execute_reply":"2024-02-29T05:53:13.269312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = ['Fish', 'Flower', 'Gravel', 'Sugar']\n\nfor label in labels:\n    train_df = train_df.assign(**{label: 0})\nfor index, row in train_df.iterrows():\n    for item in row['Label_EncodedPixels']:\n        label, value = item\n        if value == -1:\n            bool_value = 0\n        else:\n            bool_value = 1\n        train_df.loc[index, label] = bool_value\n\ntrain_df['classes'] = train_df.apply(lambda row: [col for col in labels if row[col] == 1], axis=1)\n\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:13.775371Z","iopub.execute_input":"2024-02-29T05:53:13.775994Z","iopub.status.idle":"2024-02-29T05:53:16.506015Z","shell.execute_reply.started":"2024-02-29T05:53:13.775964Z","shell.execute_reply":"2024-02-29T05:53:16.505080Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df[labels].sum().plot(kind='bar')\nplt.xlabel('Columns')\nplt.ylabel('Frequency')","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:16.507681Z","iopub.execute_input":"2024-02-29T05:53:16.507985Z","iopub.status.idle":"2024-02-29T05:53:16.797325Z","shell.execute_reply.started":"2024-02-29T05:53:16.507959Z","shell.execute_reply":"2024-02-29T05:53:16.796310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.info()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:16.798487Z","iopub.execute_input":"2024-02-29T05:53:16.798775Z","iopub.status.idle":"2024-02-29T05:53:16.815192Z","shell.execute_reply.started":"2024-02-29T05:53:16.798737Z","shell.execute_reply":"2024-02-29T05:53:16.814289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Finding the index of images having all the masks\nfor ix,item in enumerate(train_df['Label_EncodedPixels'][:100]):\n    c1=item[0][-1]!=-1\n    c2=item[1][-1]!=-1\n    c3=item[2][-1]!=-1\n    c4=item[3][-1]!=-1\n    if c1 and c2 and c3 and c4:\n        print(ix)\n        ","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:16.817696Z","iopub.execute_input":"2024-02-29T05:53:16.818282Z","iopub.status.idle":"2024-02-29T05:53:16.826226Z","shell.execute_reply.started":"2024-02-29T05:53:16.818255Z","shell.execute_reply":"2024-02-29T05:53:16.825378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.loc[28]","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:16.827255Z","iopub.execute_input":"2024-02-29T05:53:16.827530Z","iopub.status.idle":"2024-02-29T05:53:16.838959Z","shell.execute_reply.started":"2024-02-29T05:53:16.827506Z","shell.execute_reply":"2024-02-29T05:53:16.837943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for ix,item in enumerate(os.listdir(train_image_path)):\n    if item == \"015aa06.jpg\":\n        print(ix)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:16.839997Z","iopub.execute_input":"2024-02-29T05:53:16.840299Z","iopub.status.idle":"2024-02-29T05:53:17.402878Z","shell.execute_reply.started":"2024-02-29T05:53:16.840276Z","shell.execute_reply":"2024-02-29T05:53:17.401958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for j,item in enumerate(train_df['Label_EncodedPixels'][:100]):\n    c1=item[0][-1]!=-1\n    c2=item[1][-1]!=-1\n    c3=item[2][-1]!=-1\n    c4=item[3][-1]!=-1\n    if c1 and c2 and c3 and c4:\n        for ix,item in enumerate(os.listdir(train_image_path)):\n            if item == train_df.loc[j][\"Image_Id\"]:\n                print(ix)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:17.403964Z","iopub.execute_input":"2024-02-29T05:53:17.404292Z","iopub.status.idle":"2024-02-29T05:53:20.758966Z","shell.execute_reply.started":"2024-02-29T05:53:17.404265Z","shell.execute_reply":"2024-02-29T05:53:20.758093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Preaparing Test Dataframe <a class=\"anchor\" id=\"prep_test\"></a>\n[Go back to the Table of Contents](#contents)","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv(test_csv_path).fillna(-1)\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:20.761868Z","iopub.execute_input":"2024-02-29T05:53:20.762191Z","iopub.status.idle":"2024-02-29T05:53:20.798551Z","shell.execute_reply.started":"2024-02-29T05:53:20.762166Z","shell.execute_reply":"2024-02-29T05:53:20.797703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df['Image_Id'] = test_df['Image_Label'].apply(lambda x: x.split('_')[0])\ntest_df['Label'] = test_df['Image_Label'].apply(lambda x: x.split('_')[1])\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:20.799732Z","iopub.execute_input":"2024-02-29T05:53:20.800020Z","iopub.status.idle":"2024-02-29T05:53:20.828102Z","shell.execute_reply.started":"2024-02-29T05:53:20.799996Z","shell.execute_reply":"2024-02-29T05:53:20.827095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df['Label_EncodedPixels'] = test_df.apply(lambda row: (row['Label'], row['EncodedPixels']), axis = 1)\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:20.829515Z","iopub.execute_input":"2024-02-29T05:53:20.830433Z","iopub.status.idle":"2024-02-29T05:53:21.090089Z","shell.execute_reply.started":"2024-02-29T05:53:20.830375Z","shell.execute_reply":"2024-02-29T05:53:21.089088Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"grouped_EncodedPixels = test_df.groupby('Image_Id')['Label_EncodedPixels'].apply(list)\ngrouped_EncodedPixels.head()\ngrouped_EncodedPixels.info()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:21.091236Z","iopub.execute_input":"2024-02-29T05:53:21.091568Z","iopub.status.idle":"2024-02-29T05:53:21.209189Z","shell.execute_reply.started":"2024-02-29T05:53:21.091528Z","shell.execute_reply":"2024-02-29T05:53:21.208243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = grouped_EncodedPixels.to_frame().reset_index()\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:21.210206Z","iopub.execute_input":"2024-02-29T05:53:21.210501Z","iopub.status.idle":"2024-02-29T05:53:21.229739Z","shell.execute_reply.started":"2024-02-29T05:53:21.210471Z","shell.execute_reply":"2024-02-29T05:53:21.228716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"indexes =  [0, 1, 2, 3]\nlabels = ['Fish', 'Flower', 'Gravel', 'Sugar']\ncolors = ['maroon', 'darkblue', 'purple','teal']\ncolormaps = ['PuRd_r', 'Blues_r', 'Purples_r','winter_r']\nrgb_colors = [(56, 255, 255),(255, 70, 90),(48, 255, 99),(255, 255, 102)]\n\nlabel_to_idx = dict(zip(labels,indexes))\nidx_to_label = dict(zip(indexes,labels))\n\nlabel_to_color = dict(zip(labels,colors))\nidx_to_color = dict(zip(indexes,colors))\n\nlabel_to_rgb_color =  dict(zip(labels,rgb_colors))\nidx_to_rgb_color = dict(zip(indexes,rgb_colors))\n\nlabel_to_colormap = dict(zip(labels,colormaps))\nidx_to_colormap = dict(zip(indexes,colormaps))","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:21.231091Z","iopub.execute_input":"2024-02-29T05:53:21.231404Z","iopub.status.idle":"2024-02-29T05:53:21.240290Z","shell.execute_reply.started":"2024-02-29T05:53:21.231375Z","shell.execute_reply":"2024-02-29T05:53:21.239381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualizing the Images  <a class=\"anchor\" id=\"viz_img\"></a>\n[Go back to the Table of Contents](#contents) <br>\n`PIL or CV2 Image (image_width, image_height) ==  Numpy Array (image_height, image_width)`","metadata":{}},{"cell_type":"code","source":"def batchDataLoader(image_dir,img_w= 512, img_h=512, num_channel =4, Batch_Size=32):\n    \n    while True:\n        k=0\n        image_ids = os.listdir(image_dir)\n        num_batches = math.ceil(len(image_ids)/Batch_Size)\n        \n        for batch_no in range(1,num_batches+1): \n            if batch_no < num_batches:\n                batch_size = Batch_Size\n                batch_image_ids = image_ids[k:k+batch_size]\n                image_batch = np.zeros((batch_size, img_h, img_w, num_channel),dtype=np.uint8)\n                for i in range(batch_size):\n                    path = os.path.join(image_dir, image_ids[i])\n                    img = cv2.imread(path)\n                    img =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n                    image_batch[i] = img\n            # for the last batch which could be fractional\n            if batch_no == num_batches:\n                batch_image_ids = image_ids[k:]\n                batch_size = len(batch_image_ids)\n                image_batch = np.zeros((batch_size, img_h, img_w, num_channel),dtype=np.uint8)\n                for i in range(batch_size):\n                    path = os.path.join(image_dir, image_ids[i])\n                    img = cv2.imread(path)\n                    img =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n                    image_batch[i] = img\n            \n            k = k+batch_size\n            print(f\"batch_no = {batch_no}\")\n            yield image_batch","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:21.242727Z","iopub.execute_input":"2024-02-29T05:53:21.242986Z","iopub.status.idle":"2024-02-29T05:53:21.254598Z","shell.execute_reply.started":"2024-02-29T05:53:21.242963Z","shell.execute_reply":"2024-02-29T05:53:21.253811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_width = 2100\nimg_height = 1400\nnum_channel = 3\nbatch_size = 32\ncurrent_batch = batchDataLoader(train_image_path,img_width,img_height, num_channel, batch_size)\n","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:21.373810Z","iopub.execute_input":"2024-02-29T05:53:21.374116Z","iopub.status.idle":"2024-02-29T05:53:21.378534Z","shell.execute_reply.started":"2024-02-29T05:53:21.374082Z","shell.execute_reply":"2024-02-29T05:53:21.377608Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images = next(current_batch)\nprint(images.shape)\nplt.figure(figsize=(24,8))\nfor i in range(8):\n    ax = plt.subplot(2,4, i+1)\n    plt.imshow(images[i])\n    #plt.title(labels[i])\n    plt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:10:40.966407Z","iopub.execute_input":"2024-02-28T18:10:40.967298Z","iopub.status.idle":"2024-02-28T18:10:46.876196Z","shell.execute_reply.started":"2024-02-28T18:10:40.967260Z","shell.execute_reply":"2024-02-28T18:10:46.874853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = images[7]\nplt.imshow(img)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:20:02.953958Z","iopub.execute_input":"2024-02-28T18:20:02.954854Z","iopub.status.idle":"2024-02-28T18:20:03.736675Z","shell.execute_reply.started":"2024-02-28T18:20:02.954819Z","shell.execute_reply":"2024-02-28T18:20:03.735718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cloud_mask = (img>115).astype(np.float32)\nplt.imshow(cloud_mask)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:32:19.624856Z","iopub.execute_input":"2024-02-28T18:32:19.625243Z","iopub.status.idle":"2024-02-28T18:32:20.640777Z","shell.execute_reply.started":"2024-02-28T18:32:19.625213Z","shell.execute_reply":"2024-02-28T18:32:20.639865Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualizing the Segmentation Masks <a class=\"anchor\" id=\"viz_mask\"></a>\n[Go back to the Table of Contents](#contents) <br>","metadata":{}},{"cell_type":"code","source":"def rle_to_mask(rle_string, height, width):\n    '''\n    convert RLE(run length encoding) string to numpy array\n\n    Parameters: \n    rle_string (str): string of rle encoded mask\n    height (int): height of the mask\n    width (int): width of the mask \n\n    Returns: \n    numpy.array: numpy array of the mask\n    '''\n    \n    rows, cols = height, width\n    \n    if rle_string == -1:\n        return np.zeros((height,width))\n    else:\n        rle_numbers = [int(num_string) for num_string in rle_string.split(' ')]\n        rle_pairs = np.array(rle_numbers).reshape(-1,2)\n        img = np.zeros(rows*cols, dtype=np.uint8)\n        for index, length in rle_pairs:\n            index -= 1\n            img[index:index+length] = 255\n        img = img.reshape(cols,rows)\n        img = img.T\n        img = img/255.0\n        return img\n","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:18.718499Z","iopub.execute_input":"2024-02-28T18:44:18.719235Z","iopub.status.idle":"2024-02-28T18:44:18.726400Z","shell.execute_reply.started":"2024-02-28T18:44:18.719198Z","shell.execute_reply":"2024-02-28T18:44:18.725543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_masks_by_img_id(dataframe, image_id):\n    masks = np.zeros((img_height,img_width,4))\n    rle_masks = list(dataframe[dataframe['Image_Id'] == image_id]['Label_EncodedPixels'])[0]\n    fish_mask = rle_to_mask(rle_masks[0][1], img_height, img_width)\n    flower_mask = rle_to_mask(rle_masks[1][1], img_height, img_width)\n    gravel_mask = rle_to_mask(rle_masks[2][1], img_height, img_width)\n    sugar_mask = rle_to_mask(rle_masks[3][1], img_height, img_width)\n    mask_list = [fish_mask,flower_mask,gravel_mask,sugar_mask]\n    for ix, mask in enumerate(mask_list):\n        masks[:,:,ix] = mask\n    return masks","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:19.170660Z","iopub.execute_input":"2024-02-28T18:44:19.171302Z","iopub.status.idle":"2024-02-28T18:44:19.177725Z","shell.execute_reply.started":"2024-02-28T18:44:19.171274Z","shell.execute_reply":"2024-02-28T18:44:19.176774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = '0011165.jpg'\npath = os.path.join(train_image_path,image_id)\nimg = cv2.imread(path)\nimg =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\nrle = list(train_df[train_df['Image_Id'] == image_id]['Label_EncodedPixels'])[0][0][1]\n\nm = rle_to_mask(rle,img_height,img_width)\nm = cv2.resize(m, (384,256),interpolation=cv2.INTER_LINEAR)\nm = (m>0).astype(int)\nplt.imshow(m)\nprint(m.shape)\nprint(np.unique(m))\nprint(np.argwhere(m==1)[0])","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:19.770665Z","iopub.execute_input":"2024-02-28T18:44:19.771519Z","iopub.status.idle":"2024-02-28T18:44:20.052113Z","shell.execute_reply.started":"2024-02-28T18:44:19.771482Z","shell.execute_reply":"2024-02-28T18:44:20.051206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = os.listdir(train_image_path)[1]\n#image_id = 'f516a20.jpg'\nmasks = get_masks_by_img_id(train_df, image_id)\nmasks.shape","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:20.069013Z","iopub.execute_input":"2024-02-28T18:44:20.069700Z","iopub.status.idle":"2024-02-28T18:44:20.211836Z","shell.execute_reply.started":"2024-02-28T18:44:20.069672Z","shell.execute_reply":"2024-02-28T18:44:20.210917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = os.listdir(train_image_path)[5390]\n#image_id = 'f516a20.jpg'\nmasks = get_masks_by_img_id(train_df, image_id)\nprint(image_id)\nplt.figure(figsize=(24,4))\nfor ix in range(masks.shape[-1]):\n    ax = plt.subplot(1,4, ix+1)\n    plt.imshow(masks[:,:,ix],cmap=None)\n    plt.axis(\"off\");","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:20.530207Z","iopub.execute_input":"2024-02-28T18:44:20.531076Z","iopub.status.idle":"2024-02-28T18:44:21.913495Z","shell.execute_reply.started":"2024-02-28T18:44:20.531035Z","shell.execute_reply":"2024-02-28T18:44:21.912526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualizing the Images with Segmentation Masks <a class=\"anchor\" id=\"viz_img_mask\"></a>\n[Go back to the Table of Contents](#contents) <br>","metadata":{}},{"cell_type":"code","source":"def draw_label_on_mask(mask, label, obj=plt):\n    '''\n    Function to add labels to the image.\n    '''\n    if np.sum(mask) > 0:\n        y,x = 0,0\n        y,x = np.argwhere(mask==1)[0]\n        y,x = y+50,x+20      \n        obj.text(x,y,label,color='white',)\n    return None","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:24.927860Z","iopub.execute_input":"2024-02-28T18:44:24.928241Z","iopub.status.idle":"2024-02-28T18:44:24.934103Z","shell.execute_reply.started":"2024-02-28T18:44:24.928210Z","shell.execute_reply":"2024-02-28T18:44:24.933162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = os.listdir(train_image_path)[5390]\npath = os.path.join(train_image_path,image_id)\nimg = cv2.imread(path)\nimg = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\nimg = cv2.resize(img,(384,256))\nimg = img.astype(np.float32)\nimg = img/255.0\n#img -= img.mean()\n#img /= img.std()\n#standarization changes the color\nprint(img.shape)\nplt.imshow(img);","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:25.263629Z","iopub.execute_input":"2024-02-28T18:44:25.263896Z","iopub.status.idle":"2024-02-28T18:44:25.674563Z","shell.execute_reply.started":"2024-02-28T18:44:25.263873Z","shell.execute_reply":"2024-02-28T18:44:25.673625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = os.listdir(train_image_path)[5390]\npath = os.path.join(train_image_path,image_id)\nimg = cv2.imread(path)\nimg =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\nmasks = get_masks_by_img_id(train_df, image_id)\nmask = masks[:,:,1]\nmask = np.clip(mask,0,1)\nmask = np.ma.masked_where(mask == 0, mask)\nplt.imshow(img)\nplt.imshow(mask,alpha=0.7,cmap='PuRd_r')\ndraw_label_on_mask(mask,\"Flower\")\nplt.axis('off');","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:25.885334Z","iopub.execute_input":"2024-02-28T18:44:25.886094Z","iopub.status.idle":"2024-02-28T18:44:27.206746Z","shell.execute_reply.started":"2024-02-28T18:44:25.886058Z","shell.execute_reply":"2024-02-28T18:44:27.205590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = os.listdir(train_image_path)[5390]\npath = os.path.join(train_image_path,image_id)\nimg = cv2.imread(path)\nimg =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\nmasks = get_masks_by_img_id(train_df, image_id)\nmask = masks[:,:,1]\nmask = np.clip(mask,0,1)\nmask = np.ma.masked_where(mask == 0, mask)\nbbox = cv2.boundingRect(mask.astype(np.uint8))\ncv2.rectangle(img, bbox, (0, 255, 0), 5)\nplt.imshow(img)\n#plt.imshow(mask,alpha=0.7,cmap='PuRd_r')\ndraw_label_on_mask(mask,\"Flower\")\nplt.axis('off');","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:27.208451Z","iopub.execute_input":"2024-02-28T18:44:27.208824Z","iopub.status.idle":"2024-02-28T18:44:28.486151Z","shell.execute_reply.started":"2024-02-28T18:44:27.208789Z","shell.execute_reply":"2024-02-28T18:44:28.485207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = os.listdir(train_image_path)[5390]\nmasks = get_masks_by_img_id(train_df, image_id)\nmasks = (masks[:,:,0], masks[:,:,1],masks[:,:,2],masks[:,:,3])\n\npath = os.path.join(train_image_path,image_id)\nimg = cv2.imread(path)\nimg =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\ncolormaps = ['PuRd_r', 'Blues_r', 'Purples_r','winter_r'] # colormap_r = inverse colormap\nmask_labels = ['Fish', 'Flower', 'Gravel', 'Sugar']\nplt.figure(figsize=(15,10))\nfor i,(mask,cmap,label)in enumerate(zip(masks,colormaps,mask_labels)):\n    mask = np.clip(mask,0,1)\n    mask = np.ma.masked_where(mask == 0, mask)\n    ax = plt.subplot(2,2, i+1)\n    plt.imshow(img)\n    plt.imshow(mask,alpha=0.7,cmap=cmap) \n    draw_label_on_mask(mask,label)\n    plt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:28.487556Z","iopub.execute_input":"2024-02-28T18:44:28.487909Z","iopub.status.idle":"2024-02-28T18:44:32.911818Z","shell.execute_reply.started":"2024-02-28T18:44:28.487880Z","shell.execute_reply":"2024-02-28T18:44:32.910892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = os.listdir(train_image_path)[5390]\nmasks = get_masks_by_img_id(train_df, image_id)\nmasks = (masks[:,:,0], masks[:,:,1],masks[:,:,2],masks[:,:,3])\n\npath = os.path.join(train_image_path,image_id)\nimg = cv2.imread(path)\nimg =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\ncolormaps = ['PuRd_r', 'Blues_r', 'Purples_r','winter_r']\nmask_labels = ['Fish', 'Flower', 'Gravel', 'Sugar']\nplt.figure(figsize=(32,8))\nplt.imshow(img)\nfor i,(mask,cmap,label) in enumerate(zip(masks,colormaps,mask_labels)):\n    mask = np.clip(mask,0,1)\n    mask = np.ma.masked_where(mask == 0, mask)\n    plt.imshow(mask,alpha=0.7,cmap=cmap) # colormap_r = inverse colormap\n    draw_label_on_mask(mask,label)\n    plt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:32.913767Z","iopub.execute_input":"2024-02-28T18:44:32.914118Z","iopub.status.idle":"2024-02-28T18:44:35.977038Z","shell.execute_reply.started":"2024-02-28T18:44:32.914088Z","shell.execute_reply":"2024-02-28T18:44:35.976114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport numpy as np\nimport matplotlib\n\ndef show_bounding_boxes(image, mask, labels, colors):\n    \"\"\"Shows the bounding boxes surrounding the polygon in the image, and\n    adds labels to the bounding boxes.\n\n    Args:\n    image: The image.\n    mask: The binary polygon mask.\n    labels: The labels of the objects in the mask.\n    colors: A list of colors to use for the bounding boxes.\n\n    Returns:\n    The image with the bounding boxes and labels drawn on it.\n    \"\"\"\n\n    # Find the bounding boxes of the polygon.\n    bounding_boxes = []\n    for i in range(mask.shape[-1]):\n        bbox = cv2.boundingRect(mask[:, :, i])\n        bounding_boxes.append(bbox)\n\n    # Draw the bounding boxes on the image.\n    for bbox, label, color_name in zip(bounding_boxes, labels, colors):\n        rgb_color = matplotlib.colors.to_rgb(color_name)\n        rgb_color = tuple(value * 255 for value in rgb_color)\n        cv2.rectangle(image, bbox, rgb_color, 10)\n        cv2.putText(image, label, (bbox[0], bbox[1] + 50),\n                    cv2.FONT_HERSHEY_SIMPLEX, 2, (255,255,255), 3)\n\n    return image\n","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:40.371853Z","iopub.execute_input":"2024-02-28T18:44:40.373214Z","iopub.status.idle":"2024-02-28T18:44:40.381183Z","shell.execute_reply.started":"2024-02-28T18:44:40.373175Z","shell.execute_reply":"2024-02-28T18:44:40.380237Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib\nrgb_color = matplotlib.colors.to_rgb('darkblue')\nrgb_color = tuple(value * 255 for value in rgb_color)\nrgb_color","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:40.944347Z","iopub.execute_input":"2024-02-28T18:44:40.945075Z","iopub.status.idle":"2024-02-28T18:44:40.952119Z","shell.execute_reply.started":"2024-02-28T18:44:40.945039Z","shell.execute_reply":"2024-02-28T18:44:40.951167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = os.listdir(train_image_path)[5390]\nmasks = get_masks_by_img_id(train_df, image_id)\nmasks = masks.astype(np.uint8)\n\n\npath = os.path.join(train_image_path,image_id)\nimg = cv2.imread(path)\nimg =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\ncolors = ['maroon', 'darkblue', 'purple','teal']\nlabels = ['Fish', 'Flower', 'Gravel', 'Sugar']\n\nimg = show_bounding_boxes(img,masks,labels,colors)\n\nplt.figure(figsize=(32,8))\nplt.imshow(img)\nplt.axis(\"off\");","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:44:51.064430Z","iopub.execute_input":"2024-02-28T18:44:51.065433Z","iopub.status.idle":"2024-02-28T18:44:52.208798Z","shell.execute_reply.started":"2024-02-28T18:44:51.065385Z","shell.execute_reply":"2024-02-28T18:44:52.207735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_img_with_masks(img,masks,channel_first=False, comment=\"\"):\n    \n    colormaps = ['PuRd_r', 'Blues_r', 'Purples_r','winter_r']\n    mask_labels = ['Fish', 'Flower', 'Gravel', 'Sugar']\n    \n    if channel_first:\n        img = np.transpose(img, (1,2,0))\n        masks = np.transpose(masks, (1,2,0))\n\n    masks = (masks[:,:,0], masks[:,:,1],masks[:,:,2],masks[:,:,3])\n    \n    fig, axes = plt.subplots(1,6,figsize=(36,4))\n    axes = axes.ravel()\n    \n    if img.shape[-1]!=3:\n        img_cmap = 'gray'\n    else: \n        img_cmap=None\n        \n        \n    for ix,axis in enumerate(axes):\n        ix = ix%6\n        axis.imshow(img,cmap=img_cmap)\n        axis.axis('off')\n        if ix==0:\n            axis.set_title(\"Main Image\")\n        elif ix==1:\n            for i,(mask,cmap,label) in enumerate(zip(masks,colormaps,mask_labels)):\n                mask = np.clip(mask,0,1)\n                mask = np.ma.masked_where(mask == 0, mask)\n                axis.imshow(mask,alpha=0.7,cmap=cmap)\n                axis.set_title(f\"All the mask {comment}\")\n                draw_label_on_mask(mask,label,axis)\n        elif ix>=2:\n            for i,(mask,cmap,label) in enumerate(zip(masks,colormaps,mask_labels)):\n                mask = np.clip(mask,0,1)\n                mask = np.ma.masked_where(mask == 0, mask)\n                axis = axes[2+i]\n                axis.imshow(mask,alpha=0.4,cmap=cmap)\n                axis.set_title(f\"{label} {comment}\")\n                draw_label_on_mask(mask,label,axis)\n    plt.show()\n    \n    return None","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:36.778085Z","iopub.execute_input":"2024-02-29T05:53:36.779235Z","iopub.status.idle":"2024-02-29T05:53:36.792622Z","shell.execute_reply.started":"2024-02-29T05:53:36.779192Z","shell.execute_reply":"2024-02-29T05:53:36.791562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_id = os.listdir(train_image_path)[5390]\npath = os.path.join(train_image_path,image_id)\nimg = cv2.imread(path)\nimg =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\nmasks = get_masks_by_img_id(train_df, image_id)\nshow_img_with_masks(img,masks,channel_first=False)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:45:14.613146Z","iopub.execute_input":"2024-02-28T18:45:14.613999Z","iopub.status.idle":"2024-02-28T18:45:23.497801Z","shell.execute_reply.started":"2024-02-28T18:45:14.613965Z","shell.execute_reply":"2024-02-28T18:45:23.496912Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_ids = os.listdir(train_image_path)[13:16]\nfor image_id in image_ids:\n    path = os.path.join(train_image_path,image_id)\n    img = cv2.imread(path)\n    img =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    masks = get_masks_by_img_id(train_df, image_id)\n    show_img_with_masks(img,masks,channel_first=False)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:45:23.499339Z","iopub.execute_input":"2024-02-28T18:45:23.499640Z","iopub.status.idle":"2024-02-28T18:45:49.655655Z","shell.execute_reply.started":"2024-02-28T18:45:23.499614Z","shell.execute_reply":"2024-02-28T18:45:49.654737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset and Data Loader  <a class=\"anchor\" id=\"data_gen\"></a>\n[Go back to the Table of Contents](#contents) <br>","metadata":{"execution":{"iopub.status.busy":"2023-06-03T00:01:18.601333Z","iopub.execute_input":"2023-06-03T00:01:18.602029Z","iopub.status.idle":"2023-06-03T00:01:18.606731Z","shell.execute_reply.started":"2023-06-03T00:01:18.602004Z","shell.execute_reply":"2023-06-03T00:01:18.605097Z"}}},{"cell_type":"code","source":"def rle_to_mask(rle_string, height, width):\n    '''\n    convert RLE(run length encoding) string to numpy array\n\n    Parameters: \n    rle_string (str): string of rle encoded mask\n    height (int): height of the mask\n    width (int): width of the mask \n\n    Returns: \n    numpy.array: numpy array of the mask\n    '''\n    \n    rows, cols = height, width\n    \n    if rle_string == -1:\n        return np.zeros((height,width))\n    else:\n        rle_numbers = [int(num_string) for num_string in rle_string.split(' ')]\n        rle_pairs = np.array(rle_numbers).reshape(-1,2)\n        img = np.zeros(rows*cols, dtype=np.uint8)\n        for index, length in rle_pairs:\n            index -= 1\n            img[index:index+length] = 255\n        img = img.reshape(cols,rows)\n        img = img.T\n        img = img/255.0\n        return img\n    \ndef get_masks_by_img_id(dataframe, image_id):\n    masks = np.zeros((img_height,img_width,4))\n    rle_masks = list(dataframe[dataframe['Image_Id'] == image_id]['Label_EncodedPixels'])[0]\n    fish_mask = rle_to_mask(rle_masks[0][1], img_height, img_width)\n    flower_mask = rle_to_mask(rle_masks[1][1], img_height, img_width)\n    gravel_mask = rle_to_mask(rle_masks[2][1], img_height, img_width)\n    sugar_mask = rle_to_mask(rle_masks[3][1], img_height, img_width)\n    mask_list = [fish_mask,flower_mask,gravel_mask,sugar_mask]\n    for ix, mask in enumerate(mask_list):\n        masks[:,:,ix] = mask\n    return masks\n\ndef draw_label_on_mask(mask, label, obj=plt):\n    '''\n    Function to add labels to the image.\n    '''\n    if np.sum(mask) > 0:\n        y,x = 0,0\n        y,x = np.argwhere(mask==1)[0]\n        y,x = y+10,x+5      \n        obj.text(x,y,label,color='white',)\n    return None\n","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:39.813542Z","iopub.execute_input":"2024-02-29T05:53:39.814348Z","iopub.status.idle":"2024-02-29T05:53:39.830175Z","shell.execute_reply.started":"2024-02-29T05:53:39.814306Z","shell.execute_reply":"2024-02-29T05:53:39.829201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.tensor(train_df[labels].iloc[0])","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:40.877945Z","iopub.execute_input":"2024-02-29T05:53:40.878783Z","iopub.status.idle":"2024-02-29T05:53:40.896000Z","shell.execute_reply.started":"2024-02-29T05:53:40.878747Z","shell.execute_reply":"2024-02-29T05:53:40.895067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport math\nimport torch\n\nclass CloudDatasetPre(torch.utils.data.Dataset):\n    \n    def __init__(self,\n                 dataframe=None,\n                 root_dir=\".\",\n                 img_width=2100,\n                 img_height=1400,\n                 resize = False,\n                 resize_width=384,\n                 resize_height=256,\n                 augmentations=None,\n                 channel_first=True,\n                 mode = 'fit',\n                 threshold = 115,\n                 num_channels = 3,\n                 num_classes = 4,\n                 class_names = None): \n         \n        if class_names is not None and len(class_names)!=num_classes:\n            raise ValueError(\"num class and length of class_names must be same\")\n                \n        self.dataframe = dataframe\n        self.filenames = list(dataframe['Image_Id']) if not dataframe is None else os.listdir(root_dir)\n        self.root_dir = root_dir\n        self.img_width = img_width\n        self.img_height = img_height\n        self.resize = resize\n        self.resize_width = resize_width\n        self.resize_height = resize_height\n        self.augmentations = augmentations\n        self.mode = mode\n        self.threshold = threshold\n        self.channel_first = channel_first\n        self.num_channels = num_channels\n        self.num_classes = num_classes\n        self.class_names = class_names\n        self.total_samples = len(self.filenames)\n        self.indexes = np.arange(len(self.filenames))\n        \n        \n        \n    @property\n    def image_shape(self):\n        if not self.resize:\n            img_shape = (self.num_channels, self.img_height, self.img_width)\n        else:\n            img_shape = (self.num_channels, self.resize_height, self.resize_width)\n        return img_shape\n    \n    \n    def __len__(self):\n        no_of_total_samples = self.total_samples\n        return no_of_total_samples\n    \n    def __getitem__(self,index):\n        \n        filename = self.filenames[index]\n           \n        if self.mode == 'fit':\n            img = self.__generate_X(filename)\n            masks = self.__generate_y(filename)\n            if self.augmentations is not None:\n                img, masks = self.augment_sample(img, masks)\n            if self.channel_first:\n                img = np.transpose(img, (2,0,1))\n                masks = np.transpose(masks, (2,0,1))\n            if self.class_names is not None:\n                label = torch.tensor(self.dataframe[labels].iloc[0])\n                return img, masks, label.float()\n            return img,masks\n        elif self.mode == 'predict':\n            img = self.__generate_X(filename)\n            img = np.transpose(img, (2,0,1))        \n            return img\n        else:\n            raise AttributeError('The mode parameter should be set to \"fit\" or \"predict\".')\n    \n    def __generate_X(self,filename):\n        if self.resize:\n            img_size = (self.resize_height, self.resize_width)\n        else:\n            img_size = (self.img_height, self.img_width)\n            \n        img_path = os.path.join(self.root_dir, filename)\n        if self.num_channels==3:\n            img = cv2.imread(img_path)\n            img =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        else:\n            img = cv2.imread(img_path,0)\n            #img = np.expand_dims(img, axis=-1)\n        if self.resize:\n            img = cv2.resize(img,tuple(reversed(img_size)))\n            if self.num_channels==1:\n                img = np.expand_dims(img, axis=-1)\n                        \n        img = img.astype(np.float32)\n        img = img/255\n                \n        return img\n    \n    def __generate_y(self,filename):\n        if self.resize:\n            img_size = (self.resize_height, self.resize_width)\n        else:\n            img_size = (self.img_height, self.img_width)\n            \n        img_path = os.path.join(self.root_dir, filename)\n        if self.num_channels==3:\n            img = cv2.imread(img_path)\n            img =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        else:\n            img = cv2.imread(img_path,0)\n            #img = np.expand_dims(img, axis=-1)\n        if self.resize:\n            img = cv2.resize(img,tuple(reversed(img_size)))\n            if self.num_channels==1:\n                img = np.expand_dims(img, axis=-1)\n                        \n        img = img.astype(np.float32)\n        img = (img>self.threshold).astype(np.float32)\n        return img\n            \n        \n    def get_masks_by_img_id(self,dataframe,image_id):\n        rle_masks = list(dataframe[dataframe['Image_Id'] == image_id]['Label_EncodedPixels'])[0]\n        fish_mask = rle_to_mask(rle_masks[0][1], self.img_height, self.img_width)\n        flower_mask = rle_to_mask(rle_masks[1][1], self.img_height, self.img_width)\n        gravel_mask = rle_to_mask(rle_masks[2][1], self.img_height, self.img_width)\n        sugar_mask = rle_to_mask(rle_masks[3][1], self.img_height, self.img_width)\n        mask_list = [fish_mask,flower_mask,gravel_mask,sugar_mask]\n        if self.resize:\n            resized_mask_list = []\n            for mask in mask_list:\n                mask = cv2.resize(mask, (self.resize_width,self.resize_height))\n                resized_mask_list.append(mask)\n            masks = np.zeros((self.resize_height,self.resize_width,self.num_classes))\n            \n            for ix, mask in enumerate(resized_mask_list):\n                masks[:,:,ix] = mask\n        else:\n            masks = np.zeros((self.img_height,self.img_width,self.num_classes))   \n            for ix, mask in enumerate(mask_list):\n                masks[:,:,ix] = mask\n                \n        return masks\n    \n\n    def augment_sample(self, img, masks):\n        augmented = self.augmentations(image=img, mask=masks)\n        img = augmented['image']\n        masks = augmented['mask']\n        return img, masks\n        \n","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:41.445685Z","iopub.execute_input":"2024-02-29T05:53:41.446445Z","iopub.status.idle":"2024-02-29T05:53:41.485668Z","shell.execute_reply.started":"2024-02-29T05:53:41.446394Z","shell.execute_reply":"2024-02-29T05:53:41.484632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport math\nimport torch\n\nclass CloudDataset(torch.utils.data.Dataset):\n    \n    def __init__(self,\n                 dataframe=None,\n                 root_dir=\".\",\n                 img_width=2100,\n                 img_height=1400,\n                 resize = False,\n                 resize_width=384,\n                 resize_height=256,\n                 augmentations=None,\n                 channel_first=True,\n                 mode = 'fit',\n                 num_channels = 3,\n                 num_classes = 4,\n                 class_names = None): \n         \n        if class_names is not None and len(class_names)!=num_classes:\n            raise ValueError(\"num class and length of class_names must be same\")\n                \n        self.dataframe = dataframe\n        self.filenames = list(dataframe['Image_Id']) if not dataframe is None else os.listdir(root_dir)\n        self.root_dir = root_dir\n        self.img_width = img_width\n        self.img_height = img_height\n        self.resize = resize\n        self.resize_width = resize_width\n        self.resize_height = resize_height\n        self.augmentations = augmentations\n        self.mode = mode\n        self.channel_first = channel_first\n        self.num_channels = num_channels\n        self.num_classes = num_classes\n        self.class_names = class_names\n        self.total_samples = len(self.filenames)\n        self.indexes = np.arange(len(self.filenames))\n        \n        \n        \n    @property\n    def image_shape(self):\n        if not self.resize:\n            img_shape = (self.num_channels, self.img_height, self.img_width)\n        else:\n            img_shape = (self.num_channels, self.resize_height, self.resize_width)\n        return img_shape\n    \n    \n    def __len__(self):\n        no_of_total_samples = self.total_samples\n        return no_of_total_samples\n    \n    def __getitem__(self,index):\n        \n        filename = self.filenames[index]\n           \n        if self.mode == 'fit':\n            img = self.__generate_X(filename)\n            masks = self.__generate_y(filename)\n            if self.augmentations is not None:\n                img, masks = self.augment_sample(img, masks)\n            if self.channel_first:\n                img = np.transpose(img, (2,0,1))\n                masks = np.transpose(masks, (2,0,1))\n            if self.class_names is not None:\n                label = torch.tensor(self.dataframe[labels].iloc[0])\n                return img, masks, label.float()\n            return img,masks\n        elif self.mode == 'predict':\n            img = self.__generate_X(filename)\n            img = np.transpose(img, (2,0,1))        \n            return img\n        else:\n            raise AttributeError('The mode parameter should be set to \"fit\" or \"predict\".')\n    \n    def __generate_X(self,filename):\n        if self.resize:\n            img_size = (self.resize_height, self.resize_width)\n        else:\n            img_size = (self.img_height, self.img_width)\n            \n        img_path = os.path.join(self.root_dir, filename)\n        if self.num_channels==3:\n            img = cv2.imread(img_path)\n            img =  cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        else:\n            img = cv2.imread(img_path,0)\n            #img = np.expand_dims(img, axis=-1)\n        if self.resize:\n            img = cv2.resize(img,tuple(reversed(img_size)))\n            if self.num_channels==1:\n                img = np.expand_dims(img, axis=-1)\n                        \n        img = img.astype(np.float32)\n        img = img/255\n                \n        return img\n    \n    def __generate_y(self,filename):\n        masks = self.get_masks_by_img_id(self.dataframe, filename)\n        masks = (masks > 0).astype(int)\n        masks = masks.astype(np.float32)\n        return masks\n            \n        \n    def get_masks_by_img_id(self,dataframe,image_id):\n        rle_masks = list(dataframe[dataframe['Image_Id'] == image_id]['Label_EncodedPixels'])[0]\n        fish_mask = rle_to_mask(rle_masks[0][1], self.img_height, self.img_width)\n        flower_mask = rle_to_mask(rle_masks[1][1], self.img_height, self.img_width)\n        gravel_mask = rle_to_mask(rle_masks[2][1], self.img_height, self.img_width)\n        sugar_mask = rle_to_mask(rle_masks[3][1], self.img_height, self.img_width)\n        mask_list = [fish_mask,flower_mask,gravel_mask,sugar_mask]\n        if self.resize:\n            resized_mask_list = []\n            for mask in mask_list:\n                mask = cv2.resize(mask, (self.resize_width,self.resize_height))\n                resized_mask_list.append(mask)\n            masks = np.zeros((self.resize_height,self.resize_width,self.num_classes))\n            \n            for ix, mask in enumerate(resized_mask_list):\n                masks[:,:,ix] = mask\n        else:\n            masks = np.zeros((self.img_height,self.img_width,self.num_classes))   \n            for ix, mask in enumerate(mask_list):\n                masks[:,:,ix] = mask\n                \n        return masks\n    \n\n    def augment_sample(self, img, masks):\n        augmented = self.augmentations(image=img, mask=masks)\n        img = augmented['image']\n        masks = augmented['mask']\n        return img, masks\n        \n","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:45.765607Z","iopub.execute_input":"2024-02-29T05:53:45.766329Z","iopub.status.idle":"2024-02-29T05:53:45.791943Z","shell.execute_reply.started":"2024-02-29T05:53:45.766290Z","shell.execute_reply":"2024-02-29T05:53:45.791015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"augmentations = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.ShiftScaleRotate(scale_limit=0.5, rotate_limit=0, shift_limit=0.1, p=0.5, border_mode=0),\n    A.GridDistortion(p=0.5),\n    A.OpticalDistortion(p=0.5, distort_limit=2, shift_limit=0.5),\n    \n])","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:46.051352Z","iopub.execute_input":"2024-02-29T05:53:46.052111Z","iopub.status.idle":"2024-02-29T05:53:46.057819Z","shell.execute_reply.started":"2024-02-29T05:53:46.052039Z","shell.execute_reply":"2024-02-29T05:53:46.056919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset = CloudDataset(train_df,\n                        train_image_path,\n                        resize=True,\n                        resize_width=576,\n                        resize_height=384,\n                        num_channels=3,\n                        augmentations=augmentations,\n                        class_names=LABELS\n                        )\n\nprint(dataset.total_samples)\nprint(len(dataset.indexes))\nprint(dataset.__len__())\nprint(dataset.image_shape)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:46.390438Z","iopub.execute_input":"2024-02-29T05:53:46.391130Z","iopub.status.idle":"2024-02-29T05:53:46.398241Z","shell.execute_reply.started":"2024-02-29T05:53:46.391100Z","shell.execute_reply":"2024-02-29T05:53:46.397272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image, masks,label = dataset.__getitem__(1)\nprint(image.shape)\nprint(masks.shape)\nprint(image.dtype)\nprint(masks.dtype)\nprint(label.shape)\nprint(label.dtype)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:46.630589Z","iopub.execute_input":"2024-02-29T05:53:46.630879Z","iopub.status.idle":"2024-02-29T05:53:46.847186Z","shell.execute_reply.started":"2024-02-29T05:53:46.630855Z","shell.execute_reply":"2024-02-29T05:53:46.846216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.initial_seed() % 2**32","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:46.876821Z","iopub.execute_input":"2024-02-29T05:53:46.877576Z","iopub.status.idle":"2024-02-29T05:53:46.883113Z","shell.execute_reply.started":"2024-02-29T05:53:46.877547Z","shell.execute_reply":"2024-02-29T05:53:46.882269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_worker(worker_id):\n    worker_seed = torch.initial_seed() % 2**32\n    np.random.seed(worker_seed)\n    random.seed(worker_seed)\n\ngenerator = torch.Generator()\ngenerator.manual_seed(SEED)\n\nimport multiprocessing \nnum_workers = multiprocessing.cpu_count()\nprint(num_workers)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:47.076536Z","iopub.execute_input":"2024-02-29T05:53:47.077300Z","iopub.status.idle":"2024-02-29T05:53:47.083468Z","shell.execute_reply.started":"2024-02-29T05:53:47.077270Z","shell.execute_reply":"2024-02-29T05:53:47.082307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloader = torch.utils.data.DataLoader(dataset,\n                                         shuffle=True,\n                                         batch_size=BATCH_SIZE*2,\n                                         generator=generator,\n                                         num_workers=num_workers,\n                                         worker_init_fn=seed_worker,\n                                         )\n\nbatch_X, batch_y, batch_label = next(iter(dataloader))\n\nprint(batch_X.shape)\nprint(batch_y.shape)\nprint(batch_label.shape)\nprint(batch_X.dtype)\nprint(batch_y.dtype)\nprint(batch_label.dtype)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:47.311466Z","iopub.execute_input":"2024-02-29T05:53:47.312145Z","iopub.status.idle":"2024-02-29T05:53:51.159342Z","shell.execute_reply.started":"2024-02-29T05:53:47.312114Z","shell.execute_reply":"2024-02-29T05:53:51.158213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(np.transpose(batch_X[1],(1,2,0)));","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:51.161490Z","iopub.execute_input":"2024-02-29T05:53:51.161800Z","iopub.status.idle":"2024-02-29T05:53:51.614817Z","shell.execute_reply.started":"2024-02-29T05:53:51.161772Z","shell.execute_reply":"2024-02-29T05:53:51.613842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#plt.imshow(np.transpose(batch_y[1],(1,2,0)))","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:51.616041Z","iopub.execute_input":"2024-02-29T05:53:51.616396Z","iopub.status.idle":"2024-02-29T05:53:51.620716Z","shell.execute_reply.started":"2024-02-29T05:53:51.616368Z","shell.execute_reply":"2024-02-29T05:53:51.619702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.imshow(np.array(batch_y[1][1,:,:]));","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:51.623242Z","iopub.execute_input":"2024-02-29T05:53:51.623603Z","iopub.status.idle":"2024-02-29T05:53:51.884086Z","shell.execute_reply.started":"2024-02-29T05:53:51.623570Z","shell.execute_reply":"2024-02-29T05:53:51.883168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for img,masks in zip(batch_X[3:6],batch_y[3:6]):\n    show_img_with_masks(img,masks,channel_first=True)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:51.885411Z","iopub.execute_input":"2024-02-29T05:53:51.885699Z","iopub.status.idle":"2024-02-29T05:53:56.939370Z","shell.execute_reply.started":"2024-02-29T05:53:51.885674Z","shell.execute_reply":"2024-02-29T05:53:56.938497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train Test Split <a class=\"anchor\" id=\"data_split\"></a>\n[Go back to the Table of Contents](#contents) <br>","metadata":{}},{"cell_type":"code","source":"df_train, df_val = train_test_split(train_df,test_size=0.1,random_state=42, stratify=train_df['classes'])\ndf_train = df_train.reset_index(drop=True)\ndf_val = df_val.reset_index(drop=True)\nprint(df_train.shape)\nprint(df_val.shape)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:56.940573Z","iopub.execute_input":"2024-02-29T05:53:56.940858Z","iopub.status.idle":"2024-02-29T05:53:56.962933Z","shell.execute_reply.started":"2024-02-29T05:53:56.940833Z","shell.execute_reply":"2024-02-29T05:53:56.962035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train[labels].sum().plot(kind='bar')\nplt.xlabel('Columns')\nplt.ylabel('Frequency')","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:56.964038Z","iopub.execute_input":"2024-02-29T05:53:56.964402Z","iopub.status.idle":"2024-02-29T05:53:57.222712Z","shell.execute_reply.started":"2024-02-29T05:53:56.964375Z","shell.execute_reply":"2024-02-29T05:53:57.221797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for j,item in enumerate(df_train['Label_EncodedPixels'][:100]):\n    c1=item[0][-1]!=-1\n    c2=item[1][-1]!=-1\n    c3=item[2][-1]!=-1\n    c4=item[3][-1]!=-1\n    if c1 and c2 and c3 and c4:\n        for ix,item in enumerate(os.listdir(train_image_path)):\n            if item == df_train.loc[j][\"Image_Id\"]:\n                print(ix)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:57.223825Z","iopub.execute_input":"2024-02-29T05:53:57.224129Z","iopub.status.idle":"2024-02-29T05:53:58.354448Z","shell.execute_reply.started":"2024-02-29T05:53:57.224104Z","shell.execute_reply":"2024-02-29T05:53:58.353569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_val[labels].sum().plot(kind='bar')\nplt.xlabel('Columns')\nplt.ylabel('Frequency')","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:58.355714Z","iopub.execute_input":"2024-02-29T05:53:58.356099Z","iopub.status.idle":"2024-02-29T05:53:58.608911Z","shell.execute_reply.started":"2024-02-29T05:53:58.356062Z","shell.execute_reply":"2024-02-29T05:53:58.607999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for j,item in enumerate(df_val['Label_EncodedPixels'][:100]):\n    c1=item[0][-1]!=-1\n    c2=item[1][-1]!=-1\n    c3=item[2][-1]!=-1\n    c4=item[3][-1]!=-1\n    if c1 and c2 and c3 and c4:\n        for ix,item in enumerate(os.listdir(train_image_path)):\n            if item == df_val.loc[j][\"Image_Id\"]:\n                print(ix)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:53:58.611651Z","iopub.execute_input":"2024-02-29T05:53:58.611928Z","iopub.status.idle":"2024-02-29T05:54:02.044597Z","shell.execute_reply.started":"2024-02-29T05:53:58.611903Z","shell.execute_reply":"2024-02-29T05:54:02.043575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = CloudDataset(dataframe=df_train,\n                            root_dir=train_image_path,\n                            mode=\"fit\",\n                            resize=True,\n                            resize_width=R_WIDTH,\n                            resize_height=R_HEIGHT,\n                            num_channels=NUM_CHANNELS,\n                            class_names = LABELS,\n                            augmentations=augmentations)\n\nprint(train_dataset.total_samples)\nprint(len(train_dataset.indexes))\nprint(train_dataset.__len__())\n\nimage, masks, label = train_dataset.__getitem__(1)\nprint(image.shape)\nprint(masks.shape)\nprint(label.shape)\nprint(image.dtype)\nprint(masks.dtype)\nprint(label.dtype)\n\n\ntrain_dataloader = torch.utils.data.DataLoader(train_dataset,\n                                               shuffle=True,\n                                               batch_size=BATCH_SIZE,\n                                               generator=generator,\n                                               worker_init_fn=seed_worker)\n\nbatch_X, batch_y,batch_l = next(iter(train_dataloader))\n\nprint(batch_X.shape)\nprint(batch_y.shape)\nprint(batch_l.shape)\nprint(batch_X.dtype)\nprint(batch_y.dtype)\nprint(batch_l.dtype)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:54:02.045879Z","iopub.execute_input":"2024-02-29T05:54:02.046188Z","iopub.status.idle":"2024-02-29T05:54:03.194761Z","shell.execute_reply.started":"2024-02-29T05:54:02.046164Z","shell.execute_reply":"2024-02-29T05:54:03.193743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_dataset = CloudDataset(dataframe=df_val,\n                            root_dir=train_image_path,\n                            mode=\"fit\",\n                            resize=True,\n                            resize_width=R_WIDTH,\n                            resize_height=R_HEIGHT,\n                            num_channels=NUM_CHANNELS,\n                            class_names = LABELS,\n                            augmentations=None)\n\nprint(val_dataset.total_samples)\nprint(len(val_dataset.indexes))\nprint(val_dataset.__len__())\n\nimage, masks, label = val_dataset.__getitem__(1)\nprint(image.shape)\nprint(masks.shape)\nprint(label.shape)\nprint(image.dtype)\nprint(masks.dtype)\nprint(label.dtype)\n\nval_dataloader = torch.utils.data.DataLoader(val_dataset,\n                                             shuffle=True,\n                                             batch_size=BATCH_SIZE,\n                                             generator=generator,\n                                             worker_init_fn=seed_worker)\n\nbatch_X, batch_y,batch_l = next(iter(val_dataloader))\n\nprint(batch_X.shape)\nprint(batch_y.shape)\nprint(batch_l.shape)\nprint(batch_X.dtype)\nprint(batch_y.dtype)\nprint(batch_l.dtype)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:54:03.196110Z","iopub.execute_input":"2024-02-29T05:54:03.196747Z","iopub.status.idle":"2024-02-29T05:54:04.066458Z","shell.execute_reply.started":"2024-02-29T05:54:03.196710Z","shell.execute_reply":"2024-02-29T05:54:04.065304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = CloudDataset(dataframe=test_df,\n                            root_dir=test_image_path,\n                            mode=\"predict\",\n                            resize=True,\n                            resize_width=R_WIDTH,\n                            resize_height=R_HEIGHT,\n                            num_channels=NUM_CHANNELS,\n                            augmentations=None)\n\nprint(test_dataset.total_samples)\nprint(len(test_dataset.indexes))\nprint(test_dataset.__len__())\n\nimage = test_dataset.__getitem__(1)\nprint(image.shape)\nprint(image.dtype)\n\n\ntest_dataloader = torch.utils.data.DataLoader(test_dataset,\n                                              shuffle=False,\n                                              batch_size=TEST_BATCH_SIZE,\n                                              generator=generator,\n                                              worker_init_fn=seed_worker)\nbatch_X = next(iter(test_dataloader))\n\nprint(batch_X.shape)\nprint(batch_X.dtype)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:54:04.067819Z","iopub.execute_input":"2024-02-29T05:54:04.068488Z","iopub.status.idle":"2024-02-29T05:54:06.112160Z","shell.execute_reply.started":"2024-02-29T05:54:04.068447Z","shell.execute_reply":"2024-02-29T05:54:06.110989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Definitions <a class=\"anchor\" id=\"model_def\"></a>\n[Go back to the Table of Contents](#contents) <br>\n\n- Model Summary: https://medium.com/the-owl/how-to-get-model-summary-in-pytorch-57db7824d1e3","metadata":{}},{"cell_type":"code","source":"from PIL import Image\nimport torchview\n\ndef viz_model(model,input_size,flip=True,download=False, name=\"model_plot\",expand_nested=True):\n    model_plot = torchview.draw_graph(model, input_size=input_size, expand_nested=expand_nested)\n    digraph = model_plot.visual_graph\n    png_bytes = digraph.pipe(format='png')\n    model_plot = io.BytesIO(png_bytes)\n    model_plot = Image.open(model_plot)\n    model_plot = np.array(model_plot)\n    if flip:\n        model_plot = np.transpose(model_plot, (1, 0, 2))\n        model_plot = Image.fromarray(model_plot)\n        model_plot = model_plot.transpose(Image.FLIP_TOP_BOTTOM)\n        model_plot = np.array(model_plot)\n    figsize = (40,20) if flip else (20,40)\n    fig = plt.figure(figsize=figsize)\n    ax = fig.add_subplot(111)\n    ax.imshow(model_plot)\n    ax.axis('off')\n    plt.show()\n    if download:\n        model_plot = Image.fromarray(model_plot)\n        model_plot.save(f\"{name}.png\")\n    \n    \ndef model_info(model):\n    param_size = 0\n    for param in model.parameters():\n        param_size += param.nelement() * param.element_size()\n        \n    buffer_size = 0\n    for buffer in model.buffers():\n        buffer_size += buffer.nelement() * buffer.element_size()\n    \n    param_size_mb = param_size / 1024**2\n    buffer_size_mb = buffer_size / 1024**2\n    total_size_mb = (param_size + buffer_size) / 1024**2\n    \n    total_params = sum(dict((p.data_ptr(), p.numel()) for p in model.parameters()).values())\n    trainable_params = sum(dict((p.data_ptr(), p.numel()) for p in model.parameters() if p.requires_grad).values())\n    non_trainable_params = total_params - trainable_params\n\n    print(f\"Total Parameters = {format(total_params,',')}\")\n    print(f\"Trainable Parameters = {format(trainable_params,',')}\")\n    print(f\"Non-trainable Parameters = {format(non_trainable_params,',')}\")\n    print(\"=\"*50)\n    print('Model size: {:.3f} MB'.format(total_size_mb))\n    ","metadata":{"execution":{"iopub.status.busy":"2024-02-29T05:54:06.113552Z","iopub.execute_input":"2024-02-29T05:54:06.113931Z","iopub.status.idle":"2024-02-29T05:54:06.129726Z","shell.execute_reply.started":"2024-02-29T05:54:06.113895Z","shell.execute_reply":"2024-02-29T05:54:06.128687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = UNet()\n# input_size = (1,3,256, 384)\n\n# viz_model(model,input_size,flip=True,download=True)\n\n# torchinfo.summary(model,\n#                   input_size=input_size,\n#                   col_names=(\"input_size\", \"output_size\", \"num_params\"),\n#                   depth=1,\n#                   verbose=0)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T04:42:49.319250Z","iopub.execute_input":"2024-02-28T04:42:49.319606Z","iopub.status.idle":"2024-02-28T04:42:49.330329Z","shell.execute_reply.started":"2024-02-28T04:42:49.319572Z","shell.execute_reply":"2024-02-28T04:42:49.329445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Custom Losses and Metrics <a class=\"anchor\" id=\"custom_objects\"></a>\n[Go back to the Table of Contents](#contents) <br>\n\nhttps://www.kaggle.com/code/dhananjay3/image-segmentation-from-scratch-in-pytorch <br>\nhttps://www.kaggle.com/code/bigironsphere/loss-function-library-keras-pytorch  <br>\n\nhttps://github.com/JunMa11/SegLoss <br>\nhttps://arxiv.org/ftp/arxiv/papers/2209/2209.00729.pdf <br>\n\n\n`Also, a big gotcha: while all NumPy/TensorFlow/JAX/Keras APIs as well as Python unittest APIs use the argument order convention fn(y_true, y_pred) (reference values first, predicted values second), PyTorch actually uses fn(y_pred, y_true) for its losses. So make sure to invert the order of logits and targets.` https://keras.io/keras_core/guides/writing_a_custom_training_loop_in_torch/","metadata":{}},{"cell_type":"code","source":"import torch \n\n\nBINARY_MODE = \"binary\"\nMULTICLASS_MODE = \"multiclass\"\nMULTILABEL_MODE = \"multilabel\"\n        \nclass Loss(torch.nn.Module):\n    \n    def __init__(self):\n        super(Loss, self).__init__()\n        self.mode = \"binary\"\n        self.from_logits = False\n        self.reduction = \"mean\"\n        self.smooth = 0.0\n        self.epsilon  = 1e-7\n        self.ignore_index = None\n        self.log_loss = False\n        self.classes = None\n        self.log_norm = True\n        self.one_hot = True\n        \n    def compute_loss(self, y_pred, y_true):\n        raise NotImplementedError(\"Subclasses must implement the compute_loss method.\")\n    \n    def format_tensors(self, y_pred, y_true):\n        batch_size = y_pred.size(0)\n        num_classes = y_pred.size(1)\n        dims = (0,2)\n        \n        if self.from_logits:\n            if self.mode == MULTICLASS_MODE:\n                if self.log_norm:\n                    y_pred = torch.nn.functional.logsigmoid(y_pred).exp()\n                else:\n                    y_pred = torch.nn.functional.sigmoid(y_pred)\n            else:\n                if self.log_norm:\n                    y_pred = torch.nn.functional.log_softmax(y_pred, dim=1).exp()\n                else:\n                    y_pred = torch.nn.functional.softmax(y_pred, dim=1)\n        \n        if self.mode == BINARY_MODE:\n            y_true = y_true.view(batch_size, 1, -1)\n            y_pred = y_pred.view(batch_size, 1, -1)\n            if self.ignore_index is not None:\n                mask = y_true != self.ignore_index\n                y_pred = y_pred * mask\n                y_true = y_true * mask\n        if self.mode == MULTILABEL_MODE:\n            y_true = y_true.view(batch_size, num_classes, -1)\n            y_pred = y_pred.view(batch_size, num_classes, -1)\n            if self.ignore_index is not None:\n                mask = y_true != self.ignore_index\n                y_pred = y_pred * mask\n                y_true = y_true * mask\n        if self.mode == MULTICLASS_MODE:\n            y_true = y_true.view(batch_size, -1)\n            y_pred = y_pred.view(batch_size, num_classes, -1)\n            if self.ignore_index is not None:\n                mask = y_true != self.ignore_index\n                y_pred = y_pred * mask.unsqueeze(1)\n                if self.one_hot:\n                    y_true = torch.nn.functional.one_hot((y_true * mask).long(), num_classes)  # N,H*W -> N,H*W, C\n                    y_true = y_true.permute(0, 2, 1) * mask.unsqueeze(1)  # N, C, H*W  \n            elif self.one_hot:\n                y_true = torch.nn.functional.one_hot(y_true.long(), num_classes)  # N,H*W -> N,H*W, C\n                y_true = y_true.permute(0, 2, 1).float()  # N, C, H*W\n            \n        return y_pred, y_true\n\n    \n    def aggregate_loss(self, losses, batch_size):\n        if self.reduction == \"mean\":\n            loss = torch.mean(losses)\n        elif self.reduction == \"sum\":\n            loss = torch.sum(losses)\n        elif self.reduction == \"sum_over_batch_size\":\n            loss = torch.sum(losses)/batch_size\n        elif callable(self.reduction):\n            loss = self.reduction(losses)\n        else:\n            loss = losses\n        return loss\n\n    def forward(self, y_pred, y_true):\n        batch_size = y_pred.size(0)\n        num_classes = y_pred.size(1)\n        y_pred, y_true = self.format_tensors(y_pred, y_true)\n        loss = self.compute_loss(y_pred,y_true)\n        loss = self.aggregate_loss(loss, batch_size) \n        return loss\n    \n    \n    \n\nclass JointLoss(torch.nn.Module):\n    \n    def __init__(self,losses = [], weights=None):\n        super(JointLoss, self).__init__()\n        self.losses = losses\n        self.weights = weights\n        self.__name__ = self.set_name()\n        if not weights==None:\n            if len(losses)!=len(weights):\n                raise ValueError(\"Length of parameter 'losses' and 'weights' should be same\")\n    \n    def set_name(self):\n        name =  \"_\".join([loss.__name__.split(\"_\")[0] for loss in self.losses])+\"_loss\"\n        return name\n        \n    def forward(self,inputs, targets):\n\n        joint_loss = 0\n        \n        if self.weights:\n            for weight, loss in zip(self.weights, self.losses):\n                joint_loss += weight*loss(inputs ,targets)\n        else:\n            for loss in self.losses:\n                joint_loss += loss(inputs,targets)\n                \n        return joint_loss","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:10:45.314630Z","iopub.execute_input":"2024-02-29T06:10:45.314999Z","iopub.status.idle":"2024-02-29T06:10:45.338775Z","shell.execute_reply.started":"2024-02-29T06:10:45.314969Z","shell.execute_reply":"2024-02-29T06:10:45.337753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class IoULoss(Loss):\n    \n    __name__ = 'iou_loss'\n        \n    def __init__(self,\n                 mode= 'binary',\n                 from_logits=False,\n                 reduction=\"mean\",\n                 smooth=0.0,\n                 epsilon=1e-7,\n                 ignore_index=None,\n                 log_loss=False,\n                 classes=None):\n        super(IoULoss, self).__init__()\n        self.mode = mode\n        self.from_logits = from_logits\n        self.reduction = reduction\n        self.smooth = smooth\n        self.epsilon  = epsilon\n        self.ignore_index = ignore_index\n        self.log_loss = log_loss\n        self.classes = classes\n        \n        assert mode in {BINARY_MODE, MULTILABEL_MODE, MULTICLASS_MODE},f'''Invalid mode '{mode}'. \n                Supported modes are: 'binary', 'multilabel', and 'multiclass'.'''\n        \n    def compute_loss(self, y_pred, y_true):\n        dims = (0,2)\n        eps = self.epsilon\n        intersection = torch.sum((y_pred * y_true), dim=dims)\n        total = torch.sum((y_pred + y_true), dim=dims)\n        union = total - intersection \n        score = (intersection + self.smooth)/(union + self.smooth).clamp_min(eps)\n        if self.log_loss:\n            loss = -torch.log(score.clamp_min(eps))\n        else:\n            loss = 1.0 - score\n        return loss\n        \n    \nclass DiceLoss(Loss):\n    \n    __name__ = 'dice_loss'\n    \n    def __init__(self,\n                 mode= 'binary',\n                 from_logits=False,\n                 reduction=\"mean\",\n                 smooth=0.0,\n                 epsilon=1e-7,\n                 ignore_index=None,\n                 log_loss=False,\n                 classes=None):\n        super(DiceLoss, self).__init__()\n        self.mode = mode\n        self.from_logits = from_logits\n        self.reduction = reduction\n        self.smooth = smooth\n        self.epsilon  = epsilon\n        self.ignore_index = ignore_index\n        self.log_loss = log_loss\n        self.classes = classes\n        \n        assert mode in {BINARY_MODE, MULTILABEL_MODE, MULTICLASS_MODE},f'''Invalid mode '{mode}'. \n                Supported modes are: 'binary', 'multilabel', and 'multiclass'.'''\n        \n    def compute_loss(self, y_pred, y_true):\n        dims = (0,2)\n        eps = self.epsilon\n        intersection = torch.sum((y_pred * y_true),dim=dims)\n        cardinality = torch.sum((y_pred + y_true), dim=dims)\n        score = (2.*intersection + self.smooth)/(cardinality + self.smooth).clamp_min(eps) \n        if self.log_loss:\n            loss = -torch.log(score.clamp_min(eps))\n        else:\n            loss = 1.0 - score\n        # Dice loss is undefined for non-empty classes\n        # So we zero contribution of channel that does not have true pixels\n        # NOTE: A better workaround would be to use loss term `mean(y_pred)`\n        # for this case, however it will be a modified jaccard loss\n        mask = y_true.sum(dims) > 0\n        loss *= mask.to(loss.dtype)\n\n        if self.classes is not None:\n            loss = loss[self.classes]\n\n        return loss\n    \n    \nclass CELoss(Loss):\n    __name__ = 'ce_loss'\n    \n    def __init__(self,\n                 mode= 'binary',\n                 from_logits=False,\n                 reduction=\"mean\",\n                 smooth=0.0,\n                 epsilon=1e-7,\n                 ignore_index=None):\n        super(CELoss, self).__init__()\n        self.mode = mode\n        self.from_logits = from_logits\n        self.reduction = reduction\n        self.smooth = smooth\n        self.epsilon  = epsilon\n        self.ignore_index = ignore_index\n        self.log_loss = False\n        self.one_hot = False\n        \n        assert mode in {BINARY_MODE, MULTILABEL_MODE, MULTICLASS_MODE},f'''Invalid mode '{mode}'. \n                Supported modes are: 'binary', 'multilabel', and 'multiclass'.'''\n        \n    def compute_loss(self, y_pred, y_true):\n        if self.mode == MULTICLASS_MODE:\n            loss = torch.nn.functional.cross_entropy(y_pred, y_true, reduction=\"none\")\n        else:\n            loss = torch.nn.functional.binary_cross_entropy(y_pred, y_true, reduction=\"none\")\n        return loss\n    \n    \nclass FocalLoss(Loss):\n    \n    __name__ = 'focal_loss'\n    \n    def __init__(self,\n                 mode= 'binary',\n                 from_logits=False,\n                 reduction=\"mean\",\n                 smooth=0.0,\n                 epsilon=1e-7,\n                 ignore_index=None,\n                 alpha=0.25,\n                 gamma=2.0,\n                 log_loss=False,\n                 normalized=False):\n        \n        super(FocalLoss, self).__init__()\n        self.mode = mode\n        self.from_logits = from_logits\n        self.reduction = reduction\n        self.smooth = smooth\n        self.epsilon  = epsilon\n        self.ignore_index = ignore_index\n        self.alpha = alpha\n        self.gamma = gamma\n        self.log_loss = log_loss\n        self.normalized = normalized\n        self.log_norm = False\n        self.one_hot = False\n        \n        \n        assert mode in {BINARY_MODE, MULTILABEL_MODE, MULTICLASS_MODE},f'''Invalid mode '{mode}'. \n                Supported modes are: 'binary', 'multilabel', and 'multiclass'.'''\n        \n\n    def compute_loss(self, y_pred, y_true):\n        eps = self.epsilon\n        if self.mode == MULTICLASS_MODE:\n            log_loss = torch.nn.functional.cross_entropy(y_pred, y_true, reduction=\"none\")\n        else:\n            log_loss = torch.nn.functional.binary_cross_entropy(y_pred, y_true, reduction=\"none\")\n        pt = torch.exp(-log_loss)\n        focal_term = (1.0 - pt).pow(self.gamma)\n        if self.normalized:\n            norm_factor = focal_term.sum().clamp_min(eps)\n            log_loss = log_loss / norm_factor\n        loss = self.alpha * focal_term * log_loss\n        return loss\n    \n    ","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:10:47.291226Z","iopub.execute_input":"2024-02-29T06:10:47.291565Z","iopub.status.idle":"2024-02-29T06:10:47.318903Z","shell.execute_reply.started":"2024-02-29T06:10:47.291539Z","shell.execute_reply":"2024-02-29T06:10:47.318013Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_pred = torch.rand((8,4,256,384)).float()\ny_true = torch.randint(0,2,(8,4,256,384)).float()\n\n\ncriterion = DiceLoss(from_logits=True, mode=\"multilabel\")\nloss = criterion.forward(y_pred, y_true)\nprint(f\"{criterion.__name__} = {loss}\")\n\ncriterion = FocalLoss(from_logits=True, mode=\"multilabel\")\nloss = criterion.forward(y_pred, y_true)\nprint(f\"{criterion.__name__} = {loss}\")\n\ncriterion = CELoss(from_logits=True, mode=\"multilabel\")\nloss = criterion.forward(y_pred, y_true)\nprint(f\"{criterion.__name__} = {loss}\")\n\n\n\nlosses = [DiceLoss(from_logits=True, mode=\"multilabel\"),\n          FocalLoss(from_logits=True, mode=\"multilabel\", alpha=ALPHA)]\nweights = [1,2]\ncriterion = JointLoss(losses,weights)\nloss = criterion.forward(y_pred, y_true)\nprint(f\"{criterion.__name__} = {loss}\")","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:10:48.368963Z","iopub.execute_input":"2024-02-29T06:10:48.369394Z","iopub.status.idle":"2024-02-29T06:10:48.645090Z","shell.execute_reply.started":"2024-02-29T06:10:48.369364Z","shell.execute_reply":"2024-02-29T06:10:48.644162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class IoU(Loss):\n    \n    __name__ = 'iou_score'\n        \n    def __init__(self,\n                 mode= 'binary',\n                 from_logits=False,\n                 reduction=\"mean\",\n                 smooth=0.0,\n                 epsilon=1e-7,\n                 ignore_index=None,\n                 log_loss=False,\n                 classes=None):\n        super(IoU, self).__init__()\n        self.mode = mode\n        self.from_logits = from_logits\n        self.reduction = reduction\n        self.smooth = smooth\n        self.epsilon  = epsilon\n        self.ignore_index = ignore_index\n        self.log_loss = log_loss\n        self.classes = classes\n        \n        assert mode in {BINARY_MODE, MULTILABEL_MODE, MULTICLASS_MODE},f'''Invalid mode '{mode}'. \n                Supported modes are: 'binary', 'multilabel', and 'multiclass'.'''\n        \n    def compute_loss(self, y_pred, y_true):\n        dims = (0,2)\n        eps = self.epsilon\n        intersection = torch.sum((y_pred * y_true), dim=dims)\n        total = torch.sum((y_pred + y_true), dim=dims)\n        union = total - intersection \n        score = (intersection + self.smooth)/(union + self.smooth).clamp_min(eps)\n        return score\n        \n    \nclass Dice(Loss):\n    \n    __name__ = 'dice_score'\n    \n    def __init__(self,\n                 mode= 'binary',\n                 from_logits=False,\n                 reduction=\"mean\",\n                 smooth=0.0,\n                 epsilon=1e-7,\n                 ignore_index=None,\n                 log_loss=False,\n                 classes=None):\n        super(Dice, self).__init__()\n        self.mode = mode\n        self.from_logits = from_logits\n        self.reduction = reduction\n        self.smooth = smooth\n        self.epsilon  = epsilon\n        self.ignore_index = ignore_index\n        self.log_loss = log_loss\n        self.classes = classes\n        \n        assert mode in {BINARY_MODE, MULTILABEL_MODE, MULTICLASS_MODE},f'''Invalid mode '{mode}'. \n                Supported modes are: 'binary', 'multilabel', and 'multiclass'.'''\n        \n    def compute_loss(self, y_pred, y_true):\n        dims = (0,2)\n        eps = self.epsilon\n        #intersection = torch.sum((y_pred * y_true),dim=dims)\n        #cardinality = torch.sum((y_pred + y_true), dim=dims)\n        #score = (2.*intersection + self.smooth)/(cardinality + self.smooth).clamp_min(eps)\n        intersection = torch.sum((y_pred * y_true))\n        cardinality = torch.sum((y_pred + y_true))\n        score = (2.*intersection + self.smooth)/(cardinality + self.smooth)\n        return score","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:10:51.124384Z","iopub.execute_input":"2024-02-29T06:10:51.125035Z","iopub.status.idle":"2024-02-29T06:10:51.139361Z","shell.execute_reply.started":"2024-02-29T06:10:51.125002Z","shell.execute_reply":"2024-02-29T06:10:51.138344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = Dice(from_logits=True, mode=\"multilabel\")\nloss = criterion.forward(y_pred, y_true)\nprint(f\"{criterion.__name__} = {loss}\")\n\ncriterion = IoU(from_logits=True, mode=\"multilabel\")\nloss = criterion.forward(y_pred, y_true)\nprint(f\"{criterion.__name__} = {loss}\")","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:10:52.127871Z","iopub.execute_input":"2024-02-29T06:10:52.128641Z","iopub.status.idle":"2024-02-29T06:10:52.180334Z","shell.execute_reply.started":"2024-02-29T06:10:52.128610Z","shell.execute_reply":"2024-02-29T06:10:52.179245Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setting up wandb for Experiment Tracking<a class=\"anchor\" id=\"wandb\"></a>\n[Go back to the Table of Contents](#contents) <br>","metadata":{}},{"cell_type":"code","source":"import wandb\nfrom kaggle_secrets import UserSecretsClient\n\nuser_secrets = UserSecretsClient()\nwandb_api_key = user_secrets.get_secret(\"wandb_api_key\") \n\nwandb.login(key=wandb_api_key)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:10:59.825658Z","iopub.execute_input":"2024-02-29T06:10:59.826288Z","iopub.status.idle":"2024-02-29T06:11:00.527946Z","shell.execute_reply.started":"2024-02-29T06:10:59.826256Z","shell.execute_reply":"2024-02-29T06:11:00.526987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = dict(\n    epochs = NUM_EPOCHS,\n    img_height = R_HEIGHT,\n    img_width = R_WIDTH,\n    num_channels =NUM_CHANNELS,\n    num_classes = NUM_CLASSES,\n    batch_size = BATCH_SIZE,\n    labels = LABELS, \n    minsizes = MINSIZES,\n    thresholds = THRESHOLDS,\n    model_name = \"gated_deeplabv3+_efficientnetb0\",\n    encoder = 'efficientnet-b0',\n    decoder = 'deeplabv3+',\n    optimizer = 'RAdam',\n    loss = 'dice_focal_ce_loss',\n    scheduler = 'ReduceLROnPlateu',\n    framework = 'pytorch',\n    comment = 'Gated network with classification head - First Try!',\n    run_id = None\n    \n)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:11:00.529424Z","iopub.execute_input":"2024-02-29T06:11:00.529715Z","iopub.status.idle":"2024-02-29T06:11:00.535381Z","shell.execute_reply.started":"2024-02-29T06:11:00.529690Z","shell.execute_reply":"2024-02-29T06:11:00.534422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run_id = wandb.util.generate_id()\nprint(run_id)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#### don't use wandb.init() for catalyst. Instead use catalyst.dl.WandbLogger\nrun = wandb.init(entity='taki',\n                 project = 'cloud_formation_segmentation',\n                 name=config['model_name'],\n                 config=config,\n                 id=run_id,\n                 save_code = True\n                )\n\n#run.finish()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training the Model <a class=\"anchor\" id=\"model_training\"></a>\n[Go back to the Table of Contents](#contents) <br>","metadata":{}},{"cell_type":"code","source":"# ENCODER = config['encoder']\n# ENCODER_WEIGHTS = 'imagenet'\n# DEVICE = 'cuda'\n# ACTIVATION = 'sigmoid'\n# INPUT_SHAPE = (BATCH_SIZE,NUM_CHANNELS,R_HEIGHT,R_WIDTH)\n\n# model = smp.DeepLabV3Plus(\n#     encoder_name=ENCODER, \n#     encoder_weights=ENCODER_WEIGHTS,\n#     classes=4, \n#     activation=ACTIVATION,\n# )","metadata":{"execution":{"iopub.status.busy":"2024-02-28T04:43:07.087551Z","iopub.execute_input":"2024-02-28T04:43:07.087927Z","iopub.status.idle":"2024-02-28T04:43:07.092760Z","shell.execute_reply.started":"2024-02-28T04:43:07.087898Z","shell.execute_reply":"2024-02-28T04:43:07.091764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ENCODER = config['encoder']\nENCODER_WEIGHTS = 'imagenet'\nACTIVATION = 'sigmoid'\nINPUT_SHAPE = (BATCH_SIZE,NUM_CHANNELS,R_HEIGHT,R_WIDTH)\n\nAUX_PARAMS = dict(classes=NUM_CLASSES,\n                  activation=ACTIVATION,\n                  dropout=0.5,\n                  pooling='avg')\n\n\n\nclass GatedDeepLabV3PLus(smp.DeepLabV3Plus):\n    def forward(self, x):\n        mask, label = super().forward(x)\n        mask=mask*label.reshape(*label.size(), 1, 1)\n        return mask.float(), label.float()\n\n\n\n\nmodel = GatedDeepLabV3PLus(\n    encoder_name=ENCODER, \n    encoder_weights=ENCODER_WEIGHTS,\n    classes=NUM_CLASSES, \n    activation=ACTIVATION,\n    aux_params=AUX_PARAMS \n)\n\n\nmodel_info(model)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:22:17.748868Z","iopub.execute_input":"2024-02-29T06:22:17.749504Z","iopub.status.idle":"2024-02-29T06:22:17.973899Z","shell.execute_reply.started":"2024-02-29T06:22:17.749474Z","shell.execute_reply":"2024-02-29T06:22:17.972543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# run_id = \"fctcn055\"\n# run_id = \"ilbskcig\"\n# run_id = \"kbmhavbk\"\n# run_id = \"5k6wr2fi\"\n# entity = \"taki\"\n# project = \"cloud_formation_segmentation\"\n# run_path = os.path.join(entity,project,run_id)\n# filepath_in_wandb = \"ckpts/model_best.pt\"\n# restored_file = wandb.restore(filepath_in_wandb, run_path)\n# state_dict = torch.load(restored_file.name)[\"model_state_dict\"]\n# new_state_dict = dict(state_dict)\n# for key in state_dict.keys():\n#     new_state_dict[key.replace(\"model.\",\"\")] = state_dict[key]\n#     del new_state_dict[key]\n# state_dict_path = \"model_state_dict.pt\"\n# torch.save(new_state_dict, state_dict_path)\n\n# import sys\n# try:\n#     del sys.modules[\"torchmate\"]\n#     del sys.modules[\"modules\"]\n#     print(\"Torchmate and modules are removed for fresh import\")\n# except:\n#     pass\n\n#from modules import EfficientAttentionDeepLabV3PlusV2\n\n\n#from modules import *\n# model = EfficientAttentionDeepLabV3PlusV2(model_name=\"efficientnet-b0\",\n#                                             in_channels=3,\n#                                             out_channels=32,\n#                                             num_classes=4,\n#                                             activation=\"sigmoid\",\n#                                             encoder_weights=state_dict_path)\n\n\n# # input_size = (1,3, R_HEIGHT, R_WIDTH)\n# # viz_model(model,input_size)\n# # model = model.cuda()\n# model_info(model)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torchinfo.summary(model,\n                  input_size=(1,3, R_HEIGHT, R_WIDTH),\n                  col_names=(\"input_size\", \"output_size\", \"num_params\"),\n                  depth=2,\n                  verbose=0)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:18:43.793948Z","iopub.execute_input":"2024-02-29T06:18:43.794658Z","iopub.status.idle":"2024-02-29T06:18:44.847530Z","shell.execute_reply.started":"2024-02-29T06:18:43.794625Z","shell.execute_reply":"2024-02-29T06:18:44.846619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training using torchmate","metadata":{}},{"cell_type":"code","source":"import sys\ntry:\n    del sys.modules[\"torchmate\"]\n    print(\"Torchmate is removed for fresh import\")\nexcept:\n    pass","metadata":{"execution":{"iopub.status.busy":"2024-02-28T19:09:46.649736Z","iopub.execute_input":"2024-02-28T19:09:46.650653Z","iopub.status.idle":"2024-02-28T19:09:46.655748Z","shell.execute_reply.started":"2024-02-28T19:09:46.650613Z","shell.execute_reply":"2024-02-28T19:09:46.654762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchmate import Trainer,EarlyStopper, GradientAccumulator, WandbModelCheckpoint, WandbLogger","metadata":{"execution":{"iopub.status.busy":"2024-02-28T19:09:47.918614Z","iopub.execute_input":"2024-02-28T19:09:47.919326Z","iopub.status.idle":"2024-02-28T19:09:47.938291Z","shell.execute_reply.started":"2024-02-28T19:09:47.919291Z","shell.execute_reply":"2024-02-28T19:09:47.937214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_name = config['model_name']\nwork_dir = \"/kaggle/working/\"\nmonitor = \"val_dice_score\"\n\noptimizer = torch.optim.RAdam([\n    {'params': model.parameters(), 'lr': 1e-2}, \n])\n\nwandb_logger = WandbLogger()\nwandb_model_checkpoint = WandbModelCheckpoint(checkpoint_dir=work_dir,\n                                              save_best_only=True,\n                                              monitor=monitor,\n                                              mode=\"max\",\n                                              save_state_dict_only=True)\n\nearly_stopper = EarlyStopper(monitor=monitor, mode=\"max\", patience=5, min_delta=0.0001)\ngradient_accumulator = GradientAccumulator(num_accum_steps=8)\n\nnum_epochs = config['epochs']\nmetrics = [IoU(mode=\"multilabel\"),Dice(mode=\"multilabel\")]\ncallbacks = [early_stopper,gradient_accumulator,wandb_logger,wandb_model_checkpoint]\nlosses = [DiceLoss(mode=\"multilabel\", epsilon=1.0),\n          CELoss(mode=\"multilabel\",epsilon=1.0),\n          FocalLoss(mode=\"multilabel\", alpha=ALPHA,epsilon=1.0)]\n\nweights = [1,1,1]\ncriterion = [JointLoss(losses,weights),CELoss(mode=\"multilabel\",epsilon=1.0)]\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.5, patience=2, mode=\"max\")\n#scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=5, verbose=True)\ndevice = \"cuda\" if torch.cuda.is_available() else \"cpu\"","metadata":{"execution":{"iopub.status.busy":"2024-02-28T19:10:28.229642Z","iopub.execute_input":"2024-02-28T19:10:28.230040Z","iopub.status.idle":"2024-02-28T19:10:28.242101Z","shell.execute_reply.started":"2024-02-28T19:10:28.230005Z","shell.execute_reply":"2024-02-28T19:10:28.241253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"\nscheduler = torch.optim.lr_scheduler.OneCycleLR(optimizer,\n                                                max_lr=0.01,\n                                                steps_per_epoch=len(loaders[\"train\"]),\n                                                epochs=num_epochs)","metadata":{}},{"cell_type":"code","source":"trainer = Trainer(model=model,\n                  train_dataloader=train_dataloader,\n                  val_dataloader=val_dataloader,\n                  loss_fn=criterion,\n                  optimizer=optimizer,\n                  scheduler=scheduler,\n                  schedule_monitor=monitor,\n                  metrics=metrics,\n                  num_epochs=num_epochs,\n                  callbacks=callbacks,\n                  device=device\n                 )","metadata":{"execution":{"iopub.status.busy":"2024-02-28T17:14:52.105781Z","iopub.execute_input":"2024-02-28T17:14:52.106634Z","iopub.status.idle":"2024-02-28T17:14:52.111884Z","shell.execute_reply.started":"2024-02-28T17:14:52.106597Z","shell.execute_reply":"2024-02-28T17:14:52.110805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"history = trainer.fit()","metadata":{"execution":{"iopub.status.busy":"2024-02-28T17:14:55.683309Z","iopub.execute_input":"2024-02-28T17:14:55.683687Z","iopub.status.idle":"2024-02-28T17:42:46.063630Z","shell.execute_reply.started":"2024-02-28T17:14:55.683656Z","shell.execute_reply":"2024-02-28T17:42:46.062832Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config.update(run_id=run_id)\n#wandb.config.update(run_id=run_id)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_metrics_loss(history):\n    f,ax = plt.subplots(1,3,figsize=(16,4))\n    ax = ax.ravel()\n    \n    ax[0].plot([None]+history['loss'],'o-')\n    ax[0].plot([None]+history['val_loss'],'o-')\n    ax[0].legend(['Train Loss','Validation Loss'],loc = 0)\n    ax[0].set_title('Training & Validation Loss')\n    ax[0].set_xlabel('Epoch')\n    ax[0].set_ylabel('Loss')\n    ax[0].grid(True)\n\n    ax[1].plot([None]+history['dice_score'],'o-')\n    ax[1].plot([None]+history['val_dice_score'],'o-')\n    ax[1].legend(['Training Dice','Validation Dice'],loc = 0)\n    ax[1].set_title('Training & Validation Dice Score')\n    ax[1].set_xlabel('Epoch')\n    ax[1].set_ylabel('Dice Score')\n    ax[1].grid(True)\n    \n    ax[2].plot([None]+history['iou_score'],'o-')\n    ax[2].plot([None]+history['val_iou_score'],'o-')\n    ax[2].legend(['Train IOU','Validation IOU'],loc = 0)\n    ax[2].set_title('Training & Validation IOU Score')\n    ax[2].set_xlabel('Epoch')\n    ax[2].set_ylabel('IOU Score')\n    ax[2].grid(True)\n    \n    #plt.style.use('ggplot')\n    plt.tight_layout()\n    plt.show()\n    \n    return None","metadata":{"execution":{"iopub.status.busy":"2024-02-28T17:47:02.904877Z","iopub.execute_input":"2024-02-28T17:47:02.905611Z","iopub.status.idle":"2024-02-28T17:47:02.916627Z","shell.execute_reply.started":"2024-02-28T17:47:02.905573Z","shell.execute_reply":"2024-02-28T17:47:02.915719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_metrics_loss(history)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T17:47:08.783561Z","iopub.execute_input":"2024-02-28T17:47:08.783945Z","iopub.status.idle":"2024-02-28T17:47:09.692971Z","shell.execute_reply.started":"2024-02-28T17:47:08.783897Z","shell.execute_reply":"2024-02-28T17:47:09.692082Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import time\ntime.sleep(180)\n\nrun.finish()","metadata":{"execution":{"iopub.status.busy":"2024-02-27T19:54:37.476690Z","iopub.status.idle":"2024-02-27T19:54:37.477141Z","shell.execute_reply.started":"2024-02-27T19:54:37.476912Z","shell.execute_reply":"2024-02-27T19:54:37.476934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Managing CUDA memory <a class=\"anchor\" id=\"manage_cuda\"></a>\n\n[Go back to the Table of Contents](#contents) <br>","metadata":{}},{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2024-02-28T17:47:23.820754Z","iopub.execute_input":"2024-02-28T17:47:23.821140Z","iopub.status.idle":"2024-02-28T17:47:24.846016Z","shell.execute_reply.started":"2024-02-28T17:47:23.821109Z","shell.execute_reply":"2024-02-28T17:47:24.844986Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# free gpu memory\nimport gc\n#del history\nfor i in range(5):\n    torch.cuda.empty_cache()\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2024-02-28T17:47:32.000966Z","iopub.execute_input":"2024-02-28T17:47:32.001916Z","iopub.status.idle":"2024-02-28T17:47:33.281476Z","shell.execute_reply.started":"2024-02-28T17:47:32.001878Z","shell.execute_reply":"2024-02-28T17:47:33.280396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint('Using device:', device)\nprint()\n\n#Additional Info when using cuda\nif device.type == 'cuda':\n    current_device = torch.cuda.current_device()\n    print(torch.cuda.get_device_name(current_device))\n    print('Memory Usage:')\n    print('Allocated:', round(torch.cuda.memory_allocated(current_device)/1024**3,2), 'GB')\n    print('Cached:   ', round(torch.cuda.memory_reserved(current_device)/1024**3,2), 'GB')","metadata":{"execution":{"iopub.status.busy":"2024-02-28T17:47:26.599080Z","iopub.execute_input":"2024-02-28T17:47:26.599351Z","iopub.status.idle":"2024-02-28T17:47:26.610036Z","shell.execute_reply.started":"2024-02-28T17:47:26.599327Z","shell.execute_reply":"2024-02-28T17:47:26.609165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(torch.cuda.memory_summary(device=None, abbreviated=False))","metadata":{"execution":{"iopub.status.busy":"2024-02-27T21:10:22.889302Z","iopub.execute_input":"2024-02-27T21:10:22.890133Z","iopub.status.idle":"2024-02-27T21:10:22.896900Z","shell.execute_reply.started":"2024-02-27T21:10:22.890093Z","shell.execute_reply":"2024-02-27T21:10:22.895776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Post Processing for output masks <a class=\"anchor\" id=\"post_process\"></a>\n\n[Go back to the Table of Contents](#contents) <br>","metadata":{}},{"cell_type":"markdown","source":"`Classwise best threshold and minsize params: {0: (0.5, 25000), 1: (0.6, 21000), 2: (0.3, 20000), 3: (0.5, 10000)}` `Collected from` https://www.kaggle.com/code/artgor/classification-in-catalyst-with-utility-scripts#Post-processing","metadata":{}},{"cell_type":"code","source":"loaded_model = model.eval()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#minsizes = [20000 ,20000, 22500, 10000]\nminsizes = [25000 ,21000, 2000, 10000]\nthresholds = [0.5, 0.6, 0.3, 0.5]\nsigmoid = lambda x: 1 / (1 + np.exp(-x))\n\n\ndef mask2rle(img):\n    '''\n    Convert mask to rle.\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels= img.T.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\ndef post_process(probability, threshold, min_size):\n    \"\"\"\n    Post processing of each predicted mask, components with lesser number of pixels\n    than `min_size` are ignored\n    \"\"\"\n    \n    mask = cv2.threshold(probability, threshold, 1, cv2.THRESH_BINARY)[1]\n    \n    num_component, component = cv2.connectedComponents(mask.astype(np.uint8))\n    predictions = np.zeros((350, 525), np.float32)\n    num = 0\n    for c in range(1, num_component):\n        p = (component == c)\n        if p.sum() > min_size:\n            predictions[p] = 1\n            num += 1\n    return predictions, num\n            \n    \ndef draw_convex_hull(mask, mode='approx'):\n   \n    img = np.zeros(mask.shape)\n    contours, hier = cv2.findContours(mask, cv2.RETR_TREE, cv2.CHAIN_APPROX_SIMPLE)\n    \n    for c in contours:\n        if mode=='rect': # simple rectangle\n            x, y, w, h = cv2.boundingRect(c)\n            cv2.rectangle(img, (x, y), (x+w, y+h), (255, 255, 255), -1)\n        elif mode=='convex': # minimum convex hull\n            hull = cv2.convexHull(c)\n            cv2.drawContours(img, [hull], 0, (255, 255, 255),-1)\n        elif mode=='approx':\n            epsilon = 0.02*cv2.arcLength(c,True)\n            approx = cv2.approxPolyDP(c,epsilon,True)\n            cv2.drawContours(img, [approx], 0, (255, 255, 255),-1)\n        elif mode=='min': # minimum area rectangle\n            rect = cv2.minAreaRect(c)\n            box = cv2.boxPoints(rect)\n            box = np.int0(box)\n            cv2.drawContours(img, [box], 0, (255, 255, 255),-1)\n        else:\n            raise AttributeError('The mode parameter should be set to \"approx\", \"convex\", \"min\" or \"rect\".')\n    return img/255.0","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:11:19.589479Z","iopub.execute_input":"2024-02-29T06:11:19.590382Z","iopub.status.idle":"2024-02-29T06:11:19.605647Z","shell.execute_reply.started":"2024-02-29T06:11:19.590337Z","shell.execute_reply":"2024-02-29T06:11:19.604526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import cv2\nimport numpy as np\ndef batch_post_process(batch_pred_masks,\n                       thresholds=[0.5, 0.6, 0.3, 0.5],\n                       minsizes=[20000, 20000, 22500, 10000],\n                       mode=\"convex\",\n                       channel_first=True):\n    if channel_first:\n        batch_size, num_channel, height, width = batch_pred_masks.shape\n    else:\n        batch_size, height, width, num_channel = batch_pred_masks.shape\n    \n    batch_processed_masks = np.zeros(batch_pred_masks.shape)\n    \n    for k in range(batch_size):\n        for i in range(num_channel):\n            if channel_first:\n                probability = batch_pred_masks[k, i, :, :]\n            else:\n                probability = batch_pred_masks[k, :, :, i]\n            \n            min_size = minsizes[i]\n            threshold = thresholds[i]\n\n            mask = cv2.threshold(probability, threshold, 1, cv2.THRESH_BINARY)[1]\n            num_component, component = cv2.connectedComponents(mask.astype(np.uint8))\n            predictions = np.zeros((height, width), np.float32)\n            num = 0\n            for c in range(1, num_component):\n                p = (component == c)\n                if p.sum() > min_size:\n                    predictions[p] = 1\n                    num += 1\n            mask = predictions\n            if mode is not None:\n                mask = draw_convex_hull(mask.astype(np.uint8), mode=mode)\n            \n            if channel_first:\n                batch_processed_masks[k, i, :, :] = mask.astype(np.uint8)\n            else:\n                batch_processed_masks[k, :, :, i] = mask.astype(np.uint8)\n    \n    return batch_processed_masks\n\n# Example usage\nbatch_size = 2\nnum_channels = 4\nheight = 384\nwidth = 576\n\n# Simulating batch_pred_masks with random values\n\n# Using channel_first=True\nbatch_pred_masks = np.random.rand(batch_size, num_channels, height, width)\nbatch_processed_masks = batch_post_process(batch_pred_masks, channel_first=True)\nprint(batch_processed_masks.shape)  # Should print (batch_size, num_channels, height, width)\n\n# Using channel_first=False\nbatch_pred_masks = np.random.rand(batch_size, height, width,num_channels)\nbatch_processed_masks = batch_post_process(batch_pred_masks, channel_first=False)\nprint(batch_processed_masks.shape)  # Should print (batch_size, height, width, num_channels)","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:11:22.816623Z","iopub.execute_input":"2024-02-29T06:11:22.817681Z","iopub.status.idle":"2024-02-29T06:11:28.157300Z","shell.execute_reply.started":"2024-02-29T06:11:22.817632Z","shell.execute_reply":"2024-02-29T06:11:28.156326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Prediction on Test data and make submission <a class=\"anchor\" id=\"model_prediction\"></a>\n\n[Go back to the Table of Contents](#contents) <br>","metadata":{}},{"cell_type":"code","source":"#don't use sigmoid for probability --> from logits=False\n\nminsizes = MINSIZES\nthresholds = THRESHOLDS\n\nsigmoid = lambda x: 1 / (1 + np.exp(-x))\n\nencoded_pixels = []\nimage_id = 0\n\nfor i, batch_images in enumerate(tqdm.tqdm(test_dataloader)):\n    with torch.inference_mode():\n        batch_predicted_masks,_= loaded_model(batch_images.cuda())\n    for i, batch in enumerate(batch_predicted_masks):\n        for probability in batch:\n            probability = probability.cpu().detach().numpy()\n            if probability.shape != (350, 525):\n                probability = cv2.resize(probability, dsize=(525, 350), interpolation=cv2.INTER_LINEAR)\n            predict, num_predict = post_process(probability, thresholds[image_id%4], minsizes[image_id%4])\n            if num_predict == 0:\n                encoded_pixels.append('')\n            else:\n                r = mask2rle(predict)\n                encoded_pixels.append(r)\n            image_id += 1","metadata":{"execution":{"iopub.status.busy":"2024-02-28T19:06:47.408867Z","iopub.execute_input":"2024-02-28T19:06:47.409260Z","iopub.status.idle":"2024-02-28T19:07:15.604592Z","shell.execute_reply.started":"2024-02-28T19:06:47.409227Z","shell.execute_reply":"2024-02-28T19:07:15.603202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df = pd.read_csv(\"/kaggle/input/understanding_cloud_organization/sample_submission.csv\")\nsub_df['EncodedPixels'] = encoded_pixels\nsub_df.to_csv('submission.csv', columns=['Image_Label', 'EncodedPixels'], index=False)","metadata":{"execution":{"iopub.status.busy":"2024-02-27T19:54:37.529749Z","iopub.status.idle":"2024-02-27T19:54:37.530304Z","shell.execute_reply.started":"2024-02-27T19:54:37.529997Z","shell.execute_reply":"2024-02-27T19:54:37.530022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.head(10)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import subprocess\n\ndef submit_to_kaggle(filename,competition,message=\"\"):\n    command=f'kaggle competitions submit -q {competition} -f {filename} -m \"{message}\"'\n    output = subprocess.run(command, shell=True, capture_output=True, text=True)\n    print(output.stdout)\n    print(output.stderr)\n    return None\n\ncompetition_url_suffix = \"understanding_cloud_organization\"\nsubmission_filename = \"submission.csv\"\nsubmission_message = config.__str__()\n\nsubmit_to_kaggle(submission_filename,\n                competition_url_suffix,\n                submission_message)","metadata":{"execution":{"iopub.status.busy":"2024-02-27T19:54:37.533819Z","iopub.status.idle":"2024-02-27T19:54:37.534279Z","shell.execute_reply.started":"2024-02-27T19:54:37.534037Z","shell.execute_reply":"2024-02-27T19:54:37.534059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Evaluation and Visualizing predicticted mask on Validation data <a class=\"anchor\" id=\"model_evaluation\"></a>\n[Go back to the Table of Contents](#contents) <br>","metadata":{}},{"cell_type":"code","source":"# run_id = \"vchiflrb\"\n# entity = \"taki\"\n# project = \"cloud_formation_segmentation\"\n# run_path = os.path.join(entity,project,run_id)\n# filepath_in_wandb = \"ckpts/model_best.pt\"\n# restored_file = wandb.restore(filepath_in_wandb, run_path)\n# loaded_model = torch.load(restored_file.name)\n# loaded_model.eval()\n# model_info(loaded_model)\n# loaded_model = model.eval()\n#loaded_model = torch.load(\"/kaggle/input/model-file-wandb/model_best.pt\")\n#loaded_model.eval()","metadata":{"execution":{"iopub.status.busy":"2024-02-29T06:12:54.101264Z","iopub.execute_input":"2024-02-29T06:12:54.102158Z","iopub.status.idle":"2024-02-29T06:12:54.678307Z","shell.execute_reply.started":"2024-02-29T06:12:54.102125Z","shell.execute_reply":"2024-02-29T06:12:54.676974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_images,batch_masks,batch_label = next(iter(val_dataloader))\nwith torch.inference_mode():\n    batch_predicted_masks,pred_label = loaded_model(batch_images.cuda())\n\nbatch_images = batch_images.numpy()\nbatch_masks = batch_masks.numpy()\nbatch_predicted_masks = batch_predicted_masks.to('cpu').numpy()\nbatch_predicted_masks = batch_post_process(batch_predicted_masks,\n                                           mode=None,\n                                           channel_first=True,\n                                           minsizes = [20000 ,20000, 22500, 10000],\n                                           thresholds = [0.5, 0.6, 0.3, 0.5])\nprint(batch_predicted_masks.shape)\nprint(batch_masks.shape)\nprint(batch_images.shape)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T17:49:09.318793Z","iopub.execute_input":"2024-02-28T17:49:09.319429Z","iopub.status.idle":"2024-02-28T17:49:10.295321Z","shell.execute_reply.started":"2024-02-28T17:49:09.319391Z","shell.execute_reply":"2024-02-28T17:49:10.294196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"t_slice = slice(0,3)\nfor img,masks,pred_masks in zip(batch_images[t_slice],batch_masks[t_slice],batch_predicted_masks[t_slice]):\n    print(\"Image, Masks and Predicted Masks\")\n    show_img_with_masks(img,masks,comment=\"(ground truth)\",channel_first=True)\n    show_img_with_masks(img,pred_masks,comment=\"(predicted)\",channel_first=True)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T17:49:14.619623Z","iopub.execute_input":"2024-02-28T17:49:14.619996Z","iopub.status.idle":"2024-02-28T17:49:32.833786Z","shell.execute_reply.started":"2024-02-28T17:49:14.619965Z","shell.execute_reply":"2024-02-28T17:49:32.832898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Explore raw and processed predicted masks","metadata":{}},{"cell_type":"code","source":"# run_id = \"nwxn6bqc\"\n# entity = \"taki\"\n# project = \"cloud_formation_segmentation\"\n# run_path = os.path.join(entity,project,run_id)\n# filepath_in_wandb = \"ckpts/model_best.pt\"\n# restored_file = wandb.restore(filepath_in_wandb, run_path)\n# loaded_model = torch.load(restored_file.name)\n# loaded_model.eval()\n# model_info(loaded_model)\n# loaded_model = model.eval()","metadata":{"execution":{"iopub.status.busy":"2024-02-28T01:39:04.533122Z","iopub.execute_input":"2024-02-28T01:39:04.533840Z","iopub.status.idle":"2024-02-28T01:39:04.537957Z","shell.execute_reply.started":"2024-02-28T01:39:04.533806Z","shell.execute_reply":"2024-02-28T01:39:04.537044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"eval_dataset = CloudDataset(dataframe=df_val,\n                            root_dir=train_image_path,\n                            mode=\"fit\",\n                            resize=True,\n                            resize_width=R_WIDTH,\n                            resize_height=R_HEIGHT,\n                            num_channels=NUM_CHANNELS,\n                            augmentations=None)\n\nprint(eval_dataset.total_samples)\nprint(len(eval_dataset.indexes))\nprint(eval_dataset.__len__())\n\nimage, masks = eval_dataset.__getitem__(1)\nprint(image.shape)\nprint(masks.shape)\nprint(image.dtype)\nprint(masks.dtype)\n\n\neval_dataloader = torch.utils.data.DataLoader(eval_dataset,\n                                             shuffle=False,\n                                             batch_size=BATCH_SIZE*4,\n                                             generator=generator,\n                                             worker_init_fn=seed_worker)\n\nbatch_X, batch_y = next(iter(eval_dataloader))\n\nprint(batch_X.shape)\nprint(batch_y.shape)\nprint(batch_X.dtype)\nprint(batch_y.dtype)\nprint(eval_dataset.total_samples)\nprint(len(eval_dataset.indexes))\nprint(eval_dataset.__len__())\n\nimage, masks = eval_dataset.__getitem__(1)\nprint(image.shape)\nprint(masks.shape)\nprint(image.dtype)\nprint(masks.dtype)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:56:56.698373Z","iopub.execute_input":"2024-02-28T18:56:56.699133Z","iopub.status.idle":"2024-02-28T18:56:59.748647Z","shell.execute_reply.started":"2024-02-28T18:56:56.699098Z","shell.execute_reply":"2024-02-28T18:56:59.747691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for j,item in enumerate(df_val['Label_EncodedPixels'][:100]):\n    c1=item[0][-1]!=-1\n    c2=item[1][-1]!=-1\n    c3=item[2][-1]!=-1\n    c4=item[3][-1]!=-1\n    if c1 and c2 and c3 and c4:\n        for ix,item in enumerate(os.listdir(train_image_path)):\n            if item == df_val.loc[j][\"Image_Id\"]:\n                print(f\"File Index = {ix} Dataframe Index = {j}\")","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:57:17.959373Z","iopub.execute_input":"2024-02-28T18:57:17.960175Z","iopub.status.idle":"2024-02-28T18:57:21.336108Z","shell.execute_reply.started":"2024-02-28T18:57:17.960135Z","shell.execute_reply":"2024-02-28T18:57:21.335276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"m_image, m_masks = eval_dataset.__getitem__(100)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:57:21.337631Z","iopub.execute_input":"2024-02-28T18:57:21.337941Z","iopub.status.idle":"2024-02-28T18:57:21.461369Z","shell.execute_reply.started":"2024-02-28T18:57:21.337903Z","shell.execute_reply":"2024-02-28T18:57:21.460335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"minsizes = [20000 ,20000, 22500, 10000]\nthresholds = [0.5, 0.6, 0.3, 0.5]\nthresholds = [0.5, 0.6, 0.3, 0.5]     \nsigmoid = lambda x: 1 / (1 + np.exp(-x))\n\n\nbatch_images = torch.from_numpy(np.expand_dims(m_image,axis=0))\nbatch_masks = torch.from_numpy(np.expand_dims(m_masks,axis=0))\n\nprint(batch_images.shape)\nprint(batch_masks.shape)\nprint(batch_images.dtype)\nprint(batch_masks.dtype)\n\nwith torch.inference_mode():\n    batch_predicted_masks_raw,_ = loaded_model(batch_images.cuda())\n\nbatch_images = batch_images.numpy()\nbatch_masks = batch_masks.numpy()\nbatch_predicted_masks_raw = batch_predicted_masks_raw.to('cpu').numpy()\nbatch_predicted_masks_pp = batch_post_process(batch_predicted_masks_raw,\n                                            thresholds=thresholds,\n                                            minsizes=minsizes,\n                                            mode='rect',\n                                            channel_first=True)\n\nprint(batch_predicted_masks_raw.shape)\nprint(batch_predicted_masks_pp.shape)\nprint(batch_masks.shape)\nprint(batch_images.shape)\n\nimage = np.squeeze(batch_images,axis=0)\nmasks = np.squeeze(batch_masks,axis=0)\npredicted_masks_raw = np.squeeze(batch_predicted_masks_raw,axis=0)\npredicted_masks_pp = np.squeeze(batch_predicted_masks_pp,axis=0)\n\n\nprint(predicted_masks_raw.shape)\nprint(predicted_masks_pp.shape)\nprint(masks.shape)\nprint(image.shape)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:57:52.096697Z","iopub.execute_input":"2024-02-28T18:57:52.097592Z","iopub.status.idle":"2024-02-28T18:57:52.205666Z","shell.execute_reply.started":"2024-02-28T18:57:52.097556Z","shell.execute_reply":"2024-02-28T18:57:52.204750Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_with_raw(image, ground_truth_mask, raw_pred_mask, processed_pred_mask, channel_first=True):\n    \"\"\"\n    Plot image and masks.\n    If two pairs of images and masks are passes, show both.\n    \"\"\"\n    fontsize = 14\n    class_dict = {0: 'Fish', 1: 'Flower', 2: 'Gravel', 3: 'Sugar'}\n    \n    if channel_first:\n        image = np.transpose(image, (1,2,0))\n        ground_truth_mask = np.transpose(ground_truth_mask, (1,2,0))\n        raw_pred_mask = np.transpose(raw_pred_mask, (1,2,0))\n        processed_pred_mask = np.transpose(processed_pred_mask, (1,2,0))\n\n    f, ax = plt.subplots(3, 5, figsize=(24, 10))\n\n    ax[0, 0].imshow(image)\n    ax[0, 0].axis(\"off\")\n    ax[0, 0].set_title('Original image', fontsize=fontsize)\n    for i in range(4):\n        ax[0, i + 1].imshow(ground_truth_mask[:, :, i],cmap=\"gray\")\n        ax[0, i + 1].axis(\"off\")\n        ax[0, i + 1].set_title(f'{class_dict[i]} - ground truth mask', fontsize=fontsize)\n\n\n    ax[1, 0].imshow(image)\n    ax[1, 0].axis(\"off\")\n    ax[1, 0].set_title('Original image', fontsize=fontsize)\n    for i in range(4):\n        ax[1, i + 1].imshow(raw_pred_mask[:, :, i],cmap=\"gray\")\n        ax[1, i + 1].axis(\"off\")\n        ax[1, i + 1].set_title(f'{class_dict[i]} - raw prediction', fontsize=fontsize)\n        \n        \n    ax[2, 0].imshow(image)\n    ax[2, 0].axis(\"off\")\n    ax[2, 0].set_title('Original image', fontsize=fontsize)\n    for i in range(4):\n        ax[2, i + 1].imshow(processed_pred_mask[:, :, i],cmap=\"gray\")\n        ax[2, i + 1].axis(\"off\")\n        ax[2, i + 1].set_title(f'{class_dict[i]} - processed prediction', fontsize=fontsize)\n        \n    \n    return None","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:58:01.093306Z","iopub.execute_input":"2024-02-28T18:58:01.093977Z","iopub.status.idle":"2024-02-28T18:58:01.107850Z","shell.execute_reply.started":"2024-02-28T18:58:01.093944Z","shell.execute_reply":"2024-02-28T18:58:01.106735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualize_with_raw(image=image,\n                   ground_truth_mask=masks,\n                   raw_pred_mask=predicted_masks_raw,\n                   processed_pred_mask=predicted_masks_pp)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:58:02.022997Z","iopub.execute_input":"2024-02-28T18:58:02.023894Z","iopub.status.idle":"2024-02-28T18:58:04.824318Z","shell.execute_reply.started":"2024-02-28T18:58:02.023860Z","shell.execute_reply":"2024-02-28T18:58:04.823352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dpi = 100\n# plt.figure(figsize=(R_WIDTH/dpi, R_HEIGHT/dpi), dpi=dpi)\n# ix = 0\n# plt.imshow(predicted_masks_pp[ix,:, :],cmap=\"gray\")\n# plt.axis(\"off\");\n#plt.savefig(f\"{work_dir}{idx_to_label[ix]}_gt_mask.png\",transparent=True,bbox_inches='tight', pad_inches=0)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:58:14.525380Z","iopub.execute_input":"2024-02-28T18:58:14.526124Z","iopub.status.idle":"2024-02-28T18:58:14.819735Z","shell.execute_reply.started":"2024-02-28T18:58:14.526094Z","shell.execute_reply":"2024-02-28T18:58:14.818811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ix=1\nimg = image.transpose(1,2,0)\nmask = predicted_masks_raw[ix,:,:]*255.0\nmask = cv2.threshold(mask, 50, 255, cv2.THRESH_TOZERO)[1]\nmask = np.ma.masked_where(mask ==0 , mask)\nplt.imshow(img)\nplt.imshow(mask,alpha=0.7,cmap='winter')\nplt.axis('off');\n#plt.savefig(f\"{work_dir}{idx_to_label[ix]}_rp_mask_ol.png\",transparent=True,bbox_inches='tight', pad_inches=0)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:58:59.709650Z","iopub.execute_input":"2024-02-28T18:58:59.710619Z","iopub.status.idle":"2024-02-28T18:59:00.209639Z","shell.execute_reply.started":"2024-02-28T18:58:59.710575Z","shell.execute_reply":"2024-02-28T18:59:00.208555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ix=1\nmask = masks[ix,:,:]\nmask = np.clip(mask,0,1)\nmask = np.ma.masked_where(mask == 0, mask)\nplt.imshow(img)\nplt.imshow(mask,alpha=0.6,cmap='winter')\nplt.axis('off');\n#plt.savefig(f\"{work_dir}{idx_to_label[ix]}_gt_mask_ol.png\",transparent=True,bbox_inches='tight', pad_inches=0)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:58:43.574169Z","iopub.execute_input":"2024-02-28T18:58:43.575068Z","iopub.status.idle":"2024-02-28T18:58:44.056692Z","shell.execute_reply.started":"2024-02-28T18:58:43.575024Z","shell.execute_reply":"2024-02-28T18:58:44.055633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Evaluation ","metadata":{}},{"cell_type":"code","source":"# from torchmate import ProgressBar, RunningAverage\n# loss_fn = criterion\n# model.to(\"cuda\")\n# model.eval()\n# progress_bar = ProgressBar(total=len(eval_dataloader), prefix=\"evaluation\")\n# val_loss_avg = RunningAverage()\n# running_avg_dict = dict()\n# history = dict()\n# prefix = \"val_\"\n\n# if metrics is not None:\n#     for metric in metrics:\n#         running_avg_dict[f\"{prefix}{metric.__name__}_avg\"] = RunningAverage()\n        \n# for batch_ix, (X_val, y_val) in enumerate(eval_dataloader):\n#     X_val = X_val.to(device)\n#     y_val = y_val.to(device)\n#     with torch.inference_mode():\n#         y_pred_val = model(X_val)\n#         batch_val_loss = loss_fn(y_pred_val, y_val)\n#     # update value + message\n#     val_loss_avg.update(batch_val_loss.item())\n#     message = f\"{prefix}loss: {round(val_loss_avg(), 5)}\"\n#     y_pred_val = y_pred_val.to('cpu').numpy()\n#     y_pred_val = batch_post_process(y_pred_val,\n#                    thresholds=[0.5, 0.6, 0.3, 0.5],\n#                    minsizes=[20000, 20000, 22500, 10000],\n#                    mode=\"convex\",\n#                    channel_first=True\n#                 )\n#     y_pred_val = torch.tensor(y_pred_val)\n#     y_val = y_val.cpu()\n#     for metric in metrics:\n#         running_avg_dict[f\"{prefix}{metric.__name__}_avg\"].update(metric(y_pred_val, y_val).item())\n#         metric_value = round(running_avg_dict[f\"{prefix}{metric.__name__}_avg\"](), 5)\n#         message += f\" | {prefix}{metric.__name__}: {metric_value}\"\n#     progress_bar.update(batch_ix + 1, message)\n\n# # update history\n# history[f\"{prefix}loss\"] = val_loss_avg()\n# if metrics is not None:\n#     for metric in metrics:\n#         history[f\"{prefix}{metric.__name__}\"] = running_avg_dict[f\"{prefix}{metric.__name__}_avg\"]()\n# ####################################","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:59:29.056996Z","iopub.execute_input":"2024-02-28T18:59:29.057674Z","iopub.status.idle":"2024-02-28T18:59:29.062972Z","shell.execute_reply.started":"2024-02-28T18:59:29.057640Z","shell.execute_reply":"2024-02-28T18:59:29.062112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model.to(\"cpu\")","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:59:29.714390Z","iopub.execute_input":"2024-02-28T18:59:29.715262Z","iopub.status.idle":"2024-02-28T18:59:29.719060Z","shell.execute_reply.started":"2024-02-28T18:59:29.715229Z","shell.execute_reply":"2024-02-28T18:59:29.718127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# #don't use sigmoid for probability --> from logits=False\n# minsizes = [10000]*4\n# thresholds = [0.6,0.4,0.6,0.4]\n# minsizes = [20000 ,20000, 22500, 10000]\n# thresholds = [0.5, 0.6, 0.3, 0.5]\n# sigmoid = lambda x: 1 / (1 + np.exp(-x))\n\n# encoded_pixels = []\n# image_id = 0\n# progress_bar = ProgressBar(total=len(test_dataloader), prefix=\"testing\")\n# for batch_ix, batch_images in enumerate(test_dataloader):\n#     with torch.inference_mode():\n#         batch_predicted_masks = loaded_model(batch_images.cuda())\n#     for i, batch in enumerate(batch_predicted_masks):\n#         for probability in batch:\n#             probability = probability.cpu().detach().numpy()\n#             if probability.shape != (350, 525):\n#                 probability = cv2.resize(probability, dsize=(525, 350), interpolation=cv2.INTER_LINEAR)\n#             predict, num_predict = post_process(probability, thresholds[image_id%4], minsizes[image_id%4])\n#             if num_predict == 0:\n#                 encoded_pixels.append('')\n#             else:\n#                 r = mask2rle(predict)\n#                 encoded_pixels.append(r)\n#             image_id += 1\n#     progress_bar.update(batch_ix + 1)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T18:59:30.113508Z","iopub.execute_input":"2024-02-28T18:59:30.114303Z","iopub.status.idle":"2024-02-28T18:59:30.119440Z","shell.execute_reply.started":"2024-02-28T18:59:30.114267Z","shell.execute_reply":"2024-02-28T18:59:30.118411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"https://www.kaggle.com/code/ratthachat/cloud-convexhull-polygon-postprocessing-no-gpu?scriptVersionId=20977692","metadata":{}},{"cell_type":"code","source":"# batch_images = next(iter(test_dataloader))\n# batch_images = batch_images.cuda()\n\n# with torch.inference_mode():\n#     batch_predicted_masks = loaded_model(batch_images)\n    \n# batch_predicted_masks = batch_predicted_masks.round()\n# batch_predicted_masks = batch_predicted_masks.to('cpu').numpy()\n# batch_predicted_masks = np.transpose(batch_predicted_masks, (0,2,3,1))\n\n# resized_batch_predicted_masks = np.zeros((32, 350, 525, 4))\n# for i in range(batch_predicted_masks.shape[0]):\n#     resized_batch_predicted_masks[i] = cv2.resize(batch_predicted_masks[i], (525,350))\n\n# resized_batch_predicted_masks = resized_batch_predicted_masks.astype(np.int64)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T19:00:28.331825Z","iopub.execute_input":"2024-02-28T19:00:28.332273Z","iopub.status.idle":"2024-02-28T19:00:28.337203Z","shell.execute_reply.started":"2024-02-28T19:00:28.332239Z","shell.execute_reply":"2024-02-28T19:00:28.336195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for ix,item in enumerate(test_dataloader):\n#     batch_images = next(iter(test_dataloader))\n#     batch_images = batch_images.cuda()\n\n#     with torch.inference_mode():\n#         batch_predicted_masks = loaded_model(batch_images)\n\n#     batch_predicted_masks = batch_predicted_masks.round()\n#     batch_predicted_masks = batch_predicted_masks.to('cpu').numpy()\n#     batch_predicted_masks = np.transpose(batch_predicted_masks, (0,2,3,1))\n\n#     resized_batch_predicted_masks = np.zeros((32, 350, 525, 4))\n#     for i in range(batch_predicted_masks.shape[0]):\n#         resized_batch_predicted_masks[i] = cv2.resize(batch_predicted_masks[i], (525,350))\n#     resized_batch_predicted_masks = resized_batch_predicted_masks.astype(np.int64)","metadata":{"execution":{"iopub.status.busy":"2024-02-28T19:00:28.863674Z","iopub.execute_input":"2024-02-28T19:00:28.864066Z","iopub.status.idle":"2024-02-28T19:00:28.868781Z","shell.execute_reply.started":"2024-02-28T19:00:28.864029Z","shell.execute_reply":"2024-02-28T19:00:28.867768Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"execution":{"iopub.status.busy":"2024-02-27T19:54:37.531750Z","iopub.status.idle":"2024-02-27T19:54:37.532224Z","shell.execute_reply.started":"2024-02-27T19:54:37.531968Z","shell.execute_reply":"2024-02-27T19:54:37.531989Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Acknowledgements <a class=\"anchor\" id=\"ack\"></a>\n[Go back to the Table of Contents](#contents) <br>","metadata":{}},{"cell_type":"markdown","source":"* [Segmentation in PyTorch using convenient tools](https://www.kaggle.com/code/artgor/segmentation-in-pytorch-using-convenient-tools)\n* [Satellite Clouds: U-Net with ResNet Encoder](https://www.kaggle.com/code/xhlulu/satellite-clouds-u-net-with-resnet-encoder)\n* [Cloud: ConvexHull& Polygon PostProcessing (No GPU)](https://www.kaggle.com/code/ratthachat/cloud-convexhull-polygon-postprocessing-no-gpu?scriptVersionId=20977692)\n* [TTA tutorial](https://www.kaggle.com/code/joshi98kishan/let-s-understand-tta-in-segmentation/notebook)\n* [Jupyter Notebook Tricks](https://www.kaggle.com/code/tientd95/jupyter-notebook-tricks)\n* [torch tensor - numpy array conversion (cuda and cpu)](https://stackoverflow.com/questions/49768306/pytorch-tensor-to-numpy-array)","metadata":{}}]}