{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<p style=\"font-size: 32px; font-family: consolas;\"> <b>DeepLabv3+ on Ash ConTrails</b> </p>     \n\n***\n\n<p style=\"font-size: 18px; font-family: consolas;\"> This NB is an end-to-end pytorch pipeline. It contains the following: </p>    \n\n* <p style=\"font-size: 13px; font-family: consolas;\"> A Dataset class, \n* <p style=\"font-size: 13px; font-family: consolas;\"> An efficient Dataloader \n* <p style=\"font-size: 13px; font-family: consolas;\"> Training-validation-testing loops \n* <p style=\"font-size: 13px; font-family: consolas;\"> Metrics: F1 score and IOU \n* <p style=\"font-size: 13px; font-family: consolas;\"> Loss for training: Dice loss \n* <p style=\"font-size: 13px; font-family: consolas;\"> Visualization code for losses and model output \n* <p style=\"font-size: 13px; font-family: consolas;\"> Run-Length Encoding  & Submission\n\n***\n    \n<p style=\"font-size: 18px; font-family: consolas;\">1. This should be a good starting point or framework for beginners, like myself :).</p>\n\n<p style=\"font-size: 18px; font-family: consolas;\">2. I'm using the power of Atrous convolutions <i>(model: DeepLabV3+)</i> to detect contrails using the <i>false ash color images(as input)</i> as specified in the competition. Note that the dataset class passes the ash image as an additional item in the dictionary from the `__getitem__` function.</p>\n\n***   ","metadata":{"id":"jJnqDELD8t1x"}},{"cell_type":"markdown","source":"<p style=\"font-family: consolas; font-size: 20px\"><b>NOTE 1:</b> To see all visualizations, metrics, plots and enable model checkpointing </p>\n\n* <p style=\"font-family: consolas; font-size: 16px\"> set `SUBMIT` to `False` (cell below)    \n* <p style=\"font-family: consolas; font-size: 16px\"> else `SUBMIT` = `True` keeps most IO and logging (for future: competition time-complexity purposes)\n\n<p style=\"font-family: consolas; font-size: 20px\"><b>NOTE 2:</b> This NB can not be submitted to the competition because segmentation-models-pytorch is not a library that kaggle has in it's environment by default. :(. </p>\n\n","metadata":{}},{"cell_type":"code","source":"SUBMIT = False","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:58:02.512689Z","iopub.execute_input":"2023-05-15T10:58:02.513226Z","iopub.status.idle":"2023-05-15T10:58:02.526117Z","shell.execute_reply.started":"2023-05-15T10:58:02.513193Z","shell.execute_reply":"2023-05-15T10:58:02.524885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Imports","metadata":{"id":"BKl81qZA8yvl"}},{"cell_type":"markdown","source":"### standard imports","metadata":{"id":"t1EQePfV83O2"}},{"cell_type":"code","source":"import os\nimport shutil\nimport pathlib\n\nfrom PIL import Image\nimport pandas as pd\nimport numpy as np\nimport cv2 as cv\nimport random\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch import optim\n\nfrom torch.utils.data import DataLoader, random_split\nfrom torch.utils.data import Dataset\n\nimport torchvision\nfrom torchvision import datasets","metadata":{"id":"6AevVvdl85Um","execution":{"iopub.status.busy":"2023-05-15T10:58:02.53017Z","iopub.execute_input":"2023-05-15T10:58:02.530831Z","iopub.status.idle":"2023-05-15T10:58:06.760838Z","shell.execute_reply.started":"2023-05-15T10:58:02.530798Z","shell.execute_reply":"2023-05-15T10:58:06.75945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torchvision.transforms as T\nfrom torchvision.transforms import Compose, ToTensor, Resize\nfrom torchvision.utils import make_grid","metadata":{"id":"Dzq3qiLS9-u8","execution":{"iopub.status.busy":"2023-05-15T10:58:06.766742Z","iopub.execute_input":"2023-05-15T10:58:06.772555Z","iopub.status.idle":"2023-05-15T10:58:06.783529Z","shell.execute_reply.started":"2023-05-15T10:58:06.772493Z","shell.execute_reply":"2023-05-15T10:58:06.782175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Albumentations","metadata":{"id":"Y7NgVvbE-AHU"}},{"cell_type":"code","source":"try:\n    import albumentations as A\n    from albumentations.pytorch import ToTensorV2\n    import segmentation_models_pytorch as smp\nexcept:\n    !pip install -q -U segmentation-models-pytorch albumentations > /dev/null\n    import albumentations as A\n    import segmentation_models_pytorch as smp\n    from albumentations.pytorch import ToTensorV2","metadata":{"id":"NOCIrRda-BYM","execution":{"iopub.status.busy":"2023-05-15T10:58:06.785807Z","iopub.execute_input":"2023-05-15T10:58:06.78622Z","iopub.status.idle":"2023-05-15T10:58:27.992439Z","shell.execute_reply.started":"2023-05-15T10:58:06.786181Z","shell.execute_reply":"2023-05-15T10:58:27.991453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config File, Seeds & Devices","metadata":{"id":"qigozplb-TDs"}},{"cell_type":"code","source":"# Log this config file to wandb\nCONFIG = dict(\n    seed=42,\n    DATA_ROOT = '/kaggle/input/google-research-identify-contrails-reduce-global-warming/',\n    BATCH_SIZE = 16,\n    IMG_SIZE = (256,256),\n    NUM_EPOCHS = 5,\n    lr = 0.0003,\n    n_channels = 9)","metadata":{"id":"uWjUuUh8-PQF","execution":{"iopub.status.busy":"2023-05-15T10:58:27.995291Z","iopub.execute_input":"2023-05-15T10:58:27.99559Z","iopub.status.idle":"2023-05-15T10:58:28.000398Z","shell.execute_reply.started":"2023-05-15T10:58:27.99556Z","shell.execute_reply":"2023-05-15T10:58:27.999561Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# device = torch.device('cpu')\nif torch.cuda.is_available():\n    device = torch.device('cuda')\nelse:\n    device = torch.device('cpu')\n\ndevice","metadata":{"id":"sfi7e5KrSrsk","outputId":"4fffb542-588e-4cfd-bab3-1d92c84d913d","execution":{"iopub.status.busy":"2023-05-15T10:58:28.002461Z","iopub.execute_input":"2023-05-15T10:58:28.003093Z","iopub.status.idle":"2023-05-15T10:58:28.039322Z","shell.execute_reply.started":"2023-05-15T10:58:28.003045Z","shell.execute_reply":"2023-05-15T10:58:28.038322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Transforms","metadata":{"id":"Z0bspYGB851c"}},{"cell_type":"code","source":"# train_transform = A.Compose([    \n#         A.RandomCrop(height=256, width=256, always_apply=True),\n#         A.OneOf(\n#             [\n#                 A.HorizontalFlip(p=1),\n#                 A.VerticalFlip(p=1),\n#                 A.RandomRotate90(p=1),\n#             ],\n#             p=0.75,\n#         ),\n#         ToTensorV2(),\n#     ])\n\n# test_transform = A.Compose([\n#      A.PadIfNeeded(min_height=1536, min_width=1536, always_apply=True, border_mode=0),\n#      A.ToTensorV2()\n# ])\n\ntrain_transform = Compose([ToTensor()])\nval_transform = Compose([ToTensor()])\ntest_transform = Compose([ToTensor()])","metadata":{"id":"zvbPlyvZ89rF","execution":{"iopub.status.busy":"2023-05-15T10:58:28.04099Z","iopub.execute_input":"2023-05-15T10:58:28.041708Z","iopub.status.idle":"2023-05-15T10:58:28.04905Z","shell.execute_reply.started":"2023-05-15T10:58:28.041674Z","shell.execute_reply":"2023-05-15T10:58:28.0482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Count of train records\n! ls -l /kaggle/input/google-research-identify-contrails-reduce-global-warming/train | wc -l\n# Count of val records\n! ls -l /kaggle/input/google-research-identify-contrails-reduce-global-warming/validation | wc -l\n# Count of test records: 2 (verification)\n! ls -l /kaggle/input/google-research-identify-contrails-reduce-global-warming/test | wc -l","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:58:28.051157Z","iopub.execute_input":"2023-05-15T10:58:28.051424Z","iopub.status.idle":"2023-05-15T10:58:35.740368Z","shell.execute_reply.started":"2023-05-15T10:58:28.051402Z","shell.execute_reply":"2023-05-15T10:58:35.739095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset & Dataloader\n\n<p style='font-family: consolas'>  Reference: <a href=\"https://www.kaggle.com/code/thomasrochefort/pytorch-dataloader-example\"> Thomas' Dataloader & Dataset Code <a/>   \n\n<p style='font-family: consolas'> Thanks Thomas! <3","metadata":{"id":"QXCjhCis8-Ft"}},{"cell_type":"code","source":"class contrailsDataset(Dataset):\n    def __init__(self, base_dir=CONFIG['DATA_ROOT'], mode='test', transform=None):\n        super().__init__()\n        \n        assert mode in ['train', 'test', 'validation'], \"Please pass in train, test or validation\"\n        \n        self.base_dir = base_dir\n        self.mode = mode\n        self.transform = transform\n        \n        self.records = os.listdir(self.base_dir+self.mode)\n        \n        if self.mode == 'train': \n            select_cnt = 100 # Change for # samples you want \n            self.records = np.random.choice(self.records, select_cnt, replace=False)\n        \n        if self.mode == 'validation':\n            select_cnt = 50 # Change for # samples you want \n            self.records = np.random.choice(self.records, select_cnt, replace=False)\n    \n    def get_ash_img(self, bands):\n        band14 = bands[:,:,0,14-8]\n        band15 = bands[:,:,0,15-8]\n        band11 = bands[:,:,0,11-8]\n        r = normalize_range(band15 - band14, _TDIFF_BOUNDS)\n        g = normalize_range(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\n        b = normalize_range(band14, _T11_BOUNDS)\n        false_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n        return false_color\n        \n        \n    def __getitem__(self, idx):\n        record_id = self.records[idx]\n        record_dir = os.path.join(self.base_dir, self.mode, record_id)\n        \n        bands_data = []\n        for i in range(8, 17):\n            band_file = os.path.join(record_dir, f'band_{str(i).zfill(2)}.npy')\n            band_data = np.load(band_file)\n            bands_data.append(band_data)\n\n        # Stack band data along the channel axis\n        bands_data = np.stack(bands_data, axis=-1)\n        ash = self.get_ash_img(bands_data)\n\n        # If the data type is 'train' or 'validation', load the masks\n        if self.mode in ['train', 'validation']:\n            pixel_masks_file = os.path.join(record_dir, 'human_pixel_masks.npy')\n            pixel_masks = np.load(pixel_masks_file)\n        else:\n            pixel_masks = None  \n\n        if self.transform:\n            ash = self.transform(ash)\n\n            if self.mode != 'test':\n                pixel_masks = self.transform(pixel_masks)\n                sample = {'bands': bands_data, 'mask': pixel_masks, 'ash': ash}\n            else:\n                sample = {'bands': bands_data, 'ash': ash}\n        \n        return sample\n    \n    \n    def __len__(self):\n        return len(self.records)","metadata":{"id":"fyutMQiH-U8U","execution":{"iopub.status.busy":"2023-05-15T10:58:35.74404Z","iopub.execute_input":"2023-05-15T10:58:35.744409Z","iopub.status.idle":"2023-05-15T10:58:35.768855Z","shell.execute_reply.started":"2023-05-15T10:58:35.744371Z","shell.execute_reply":"2023-05-15T10:58:35.767839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainData = contrailsDataset(mode='train', transform=train_transform)\nvalData = contrailsDataset(mode='validation', transform=val_transform)\ntestData = contrailsDataset(mode='test', transform=test_transform)","metadata":{"id":"A-rOZjAmEZJf","outputId":"50a08da0-f5d5-4edc-cda4-38f7043d4147","execution":{"iopub.status.busy":"2023-05-15T10:58:35.773744Z","iopub.execute_input":"2023-05-15T10:58:35.776473Z","iopub.status.idle":"2023-05-15T10:58:35.81031Z","shell.execute_reply.started":"2023-05-15T10:58:35.776439Z","shell.execute_reply":"2023-05-15T10:58:35.809311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(trainData), len(valData), len(testData)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:58:35.818216Z","iopub.execute_input":"2023-05-15T10:58:35.820399Z","iopub.status.idle":"2023-05-15T10:58:35.830796Z","shell.execute_reply.started":"2023-05-15T10:58:35.820366Z","shell.execute_reply":"2023-05-15T10:58:35.829557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## DataLoader","metadata":{}},{"cell_type":"code","source":"train_dataloader = DataLoader(trainData, \n                              batch_size=CONFIG['BATCH_SIZE'], \n                              shuffle=True)\n\nval_dataloader = DataLoader(valData, \n                            batch_size=CONFIG['BATCH_SIZE'],\n                            shuffle=True)\n\ntest_dataloader = DataLoader(testData, \n                             batch_size=CONFIG['BATCH_SIZE'], \n                             shuffle=False)","metadata":{"id":"w4taEkO_EoEg","execution":{"iopub.status.busy":"2023-05-15T10:58:35.835305Z","iopub.execute_input":"2023-05-15T10:58:35.837391Z","iopub.status.idle":"2023-05-15T10:58:35.845279Z","shell.execute_reply.started":"2023-05-15T10:58:35.83736Z","shell.execute_reply":"2023-05-15T10:58:35.844333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Sanity Visualizations   \n\n\nCode Reference: https://www.kaggle.com/code/inversion/visualizing-contrails/notebook     \nAsh color scheme: https://rammb.cira.colostate.edu/training/visit/quick_guides/GOES_Ash_RGB.pdf ","metadata":{"id":"8IdfCGJn9EYU"}},{"cell_type":"code","source":"_T11_BOUNDS = (243, 303)\n_CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n_TDIFF_BOUNDS = (-4, 2)\n\ndef normalize_range(data, bounds):\n    return (data - bounds[0]) / (bounds[1] - bounds[0])\n\ndef get_ash_img(band11, band14, band15):\n    r = normalize_range(band15 - band14, _TDIFF_BOUNDS)\n    g = normalize_range(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\n    b = normalize_range(band14, _T11_BOUNDS)\n    false_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n    return false_color","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:58:35.850005Z","iopub.execute_input":"2023-05-15T10:58:35.852131Z","iopub.status.idle":"2023-05-15T10:58:35.860593Z","shell.execute_reply.started":"2023-05-15T10:58:35.852099Z","shell.execute_reply":"2023-05-15T10:58:35.859728Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"s = next(iter(train_dataloader))\nbands = s['bands']\nmask = s['mask']\nash = s['ash']\nprint(bands.shape, mask.shape, ash.shape) # bands.shape ==  (N x H x W x T x Bands)","metadata":{"id":"qP4uDops9Grt","outputId":"828d68ac-6fb7-4dc1-c717-22bbcd96fc64","execution":{"iopub.status.busy":"2023-05-15T10:58:35.86514Z","iopub.execute_input":"2023-05-15T10:58:35.867574Z","iopub.status.idle":"2023-05-15T10:58:42.722647Z","shell.execute_reply.started":"2023-05-15T10:58:35.867542Z","shell.execute_reply":"2023-05-15T10:58:42.721671Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sampler_val = np.random.randint(0,bands.size(0))\n\nimg = ash[sampler_val, :, :, :].permute(1,2,0)\nmask_show = mask[sampler_val, :,:,:].permute(1,2,0)\n\nplt.figure(figsize=(18, 6))\nax = plt.subplot(1, 3, 1)\nax.imshow(img)\nax.set_title('False color image')\n\nax = plt.subplot(1, 3, 2)\nax.imshow(mask_show, interpolation='none')\nax.set_title('Ground truth contrail mask')\n\nax = plt.subplot(1, 3, 3)\nax.imshow(img)\nax.imshow(mask_show, cmap='Reds', alpha=0.3, interpolation='none')\nax.set_title('Contrail mask on false color image');","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:58:42.724064Z","iopub.execute_input":"2023-05-15T10:58:42.725765Z","iopub.status.idle":"2023-05-15T10:58:43.64829Z","shell.execute_reply.started":"2023-05-15T10:58:42.72573Z","shell.execute_reply":"2023-05-15T10:58:43.643823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Animation: Traversing the temporal dimension","metadata":{}},{"cell_type":"code","source":"false_imgs = []\n\nfor i in range(8):\n    band11 = bands[0,:,:,i,3]\n    band14 = bands[0,:,:,i,6]\n    band15 = bands[0,:,:,i,7]\n    false_imgs.append(get_ash_img(band11, band14, band15))\n    \ndef draw(i):\n    im.set_array(false_imgs[i])\n    return [im]\n\nfrom matplotlib import animation\nfrom IPython import display\n\nfig = plt.figure(figsize=(6, 6))\nim = plt.imshow(false_imgs[0])\n\nanim = animation.FuncAnimation(\n    fig, draw, frames=len(false_imgs), interval=500, blit=True\n)\nplt.close()\ndisplay.HTML(anim.to_jshtml())","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:58:43.649605Z","iopub.execute_input":"2023-05-15T10:58:43.649944Z","iopub.status.idle":"2023-05-15T10:58:45.687735Z","shell.execute_reply.started":"2023-05-15T10:58:43.649915Z","shell.execute_reply":"2023-05-15T10:58:45.686484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utilities for Segmentation","metadata":{"id":"8xEbBzpU9Hil"}},{"cell_type":"code","source":"from torchvision import models\ntry:\n    from torchsummary import summary\nexcept:\n    !pip install torchsummary > /dev/null\n    from torchsummary import summary\n    \n\nfrom torchvision import models","metadata":{"id":"eagh9GG-SbwA","execution":{"iopub.status.busy":"2023-05-15T10:58:45.689271Z","iopub.execute_input":"2023-05-15T10:58:45.689935Z","iopub.status.idle":"2023-05-15T10:58:56.994069Z","shell.execute_reply.started":"2023-05-15T10:58:45.689895Z","shell.execute_reply":"2023-05-15T10:58:56.992754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class_names = ['background', 'contrails']\nclass_rgb_values = [[0.,0.,0.], [255.,255.,255.]]","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:58:56.996254Z","iopub.execute_input":"2023-05-15T10:58:56.996685Z","iopub.status.idle":"2023-05-15T10:58:57.003694Z","shell.execute_reply.started":"2023-05-15T10:58:56.996645Z","shell.execute_reply":"2023-05-15T10:58:57.002536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def one_hot_encode(label, label_values):\n    semantic_map = []\n    for colour in label_values:\n        equality = np.equal(label, colour)\n        class_map = np.all(equality, axis = -1)\n        semantic_map.append(class_map)\n    semantic_map = np.stack(semantic_map, axis=-1)\n    return semantic_map\n\ndef reverse_one_hot(image):\n    x = np.argmax(image, axis = -1)\n    return x\n\ndef colourize_seg(image, label_values):\n    colour_codes = np.array(label_values)\n    x = colour_codes[image.astype(int)]\n    return x","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:58:57.006616Z","iopub.execute_input":"2023-05-15T10:58:57.006999Z","iopub.status.idle":"2023-05-15T10:58:57.014374Z","shell.execute_reply.started":"2023-05-15T10:58:57.006975Z","shell.execute_reply":"2023-05-15T10:58:57.013164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model    \n\n<p style=\"font-family: consolas; color: red; font-size: 20px\"> <b>TODO:</b> Replace with a model without using any external library to submit NB to competition.","metadata":{}},{"cell_type":"code","source":"ENCODER = 'resnet101'\nENCODER_WEIGHTS = 'imagenet'\nCLASSES = class_names\nACTIVATION = 'sigmoid' # could be None for logits or 'softmax2d' for multiclass segmentation\n\n# create segmentation model with pretrained encoder\nmodel = smp.DeepLabV3Plus(\n    encoder_name=ENCODER, \n    encoder_weights=ENCODER_WEIGHTS, \n    classes=len(CLASSES), \n    activation=ACTIVATION,\n).to(device)\n\npreprocessing_fn = smp.encoders.get_preprocessing_fn(ENCODER, ENCODER_WEIGHTS)\n\n# summary(model, (ash.size(3),ash.size(1),ash.size(2)))","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:58:57.016409Z","iopub.execute_input":"2023-05-15T10:58:57.017339Z","iopub.status.idle":"2023-05-15T10:59:22.021348Z","shell.execute_reply.started":"2023-05-15T10:58:57.017306Z","shell.execute_reply":"2023-05-15T10:59:22.020358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf /kaggle/working/checkpoints\n!mkdir /kaggle/working/checkpoints","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:59:22.022693Z","iopub.execute_input":"2023-05-15T10:59:22.023042Z","iopub.status.idle":"2023-05-15T10:59:24.075917Z","shell.execute_reply.started":"2023-05-15T10:59:22.023009Z","shell.execute_reply":"2023-05-15T10:59:24.074609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Some more settings: (epochs, visualization frequency optimizer etc.)","metadata":{}},{"cell_type":"code","source":"# Set flag to train the model or not. If set to 'False', only prediction is performed (using an older model checkpoint)\nTRAINING = True\n\n# Set num of epochs\nEPOCHS = 10\n\n# Checkpoint saving path\nPATH = '/kaggle/working/checkpoints/'\n\n# Visualize masks, input, output after every 'vis_masks_after' \"epochs\"\nvis_after = 2\n\n# define loss function\ndice_loss = smp.losses.DiceLoss(mode='binary', from_logits=False)\n\n# define optimizer\noptimizer = torch.optim.Adam([ \n    dict(params=model.parameters(), lr=0.0001),\n])\n\n# define learning rate scheduler (not used in this NB)\n# lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n#     optimizer, T_0=1, T_mult=2, eta_min=5e-5,\n# )\n\n# load best saved model checkpoint from previous commit (if present)\nif os.path.exists('../input/deeplabv3-efficientnetb4-frontend-using-pytorch/best_model.pth'):\n    model = torch.load('../input/deeplabv3-efficientnetb4-frontend-using-pytorch/best_model.pth', map_location=device)","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:59:24.077886Z","iopub.execute_input":"2023-05-15T10:59:24.078277Z","iopub.status.idle":"2023-05-15T10:59:24.091072Z","shell.execute_reply.started":"2023-05-15T10:59:24.078238Z","shell.execute_reply":"2023-05-15T10:59:24.090204Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def visualize(out: torch.Tensor, ash: torch.Tensor, mask: torch.Tensor, title_str=\"\"):\n    out_ = np.array(out.cpu())\n    ash_ = np.array(ash.cpu())\n    mask_ = np.array(mask.cpu())\n\n    out_imgs = []\n    ash_imgs = []\n    mask_imgs = []\n    for i in range(out_.shape[0]):\n        out_imgs.append(reverse_one_hot(out_[i].transpose(1,2,0)))\n        ash_imgs.append(ash_[i].transpose(1,2,0))\n        mask_imgs.append(mask_[i].transpose(1,2,0))\n  \n    fig = plt.figure(figsize=(16, 4))\n\n    for i in range(len(out_imgs)):\n        ax = plt.subplot(3, len(out_imgs), i+1)\n        image1 = out_imgs[i]\n        ax.imshow(image1)\n        ax.axis('off')\n        ax.set_title(f\"{i+1}\")\n        \n    for i in range(len(ash_imgs)):\n        ax = plt.subplot(3, len(ash_imgs), len(out_imgs)+len(ash_imgs)+i+1)\n        image1 = ash_imgs[i]\n        ax.imshow(image1)\n        ax.axis('off')\n        ax.set_title(f\"{i+1}\")\n        \n    for i in range(len(mask_imgs)):\n        ax = plt.subplot(3, len(mask_imgs), len(out_imgs)+i+1)\n        image1 = mask_imgs[i]\n        ax.imshow(image1)\n        ax.axis('off')\n        ax.set_title(f\"{i+1}\")\n        \n    plt.savefig(title_str)\n    plt.show()\n    ","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:59:24.093618Z","iopub.execute_input":"2023-05-15T10:59:24.094318Z","iopub.status.idle":"2023-05-15T10:59:24.106413Z","shell.execute_reply.started":"2023-05-15T10:59:24.094287Z","shell.execute_reply":"2023-05-15T10:59:24.10551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TRAINING:\n    best_iou_score = 0.0\n    \n    train_logs, valid_logs = {}, {}\n    for e in ['loss', 'f1', 'iou']:\n        train_logs[e] = []\n        valid_logs[e] = []\n\n    for epoch in range(EPOCHS):\n        # Perform training & validation steps\n        print('\\nEpoch: {}'.format(epoch+1))\n        \n        model.train()\n        # Training loop\n        train_iou, train_f1, train_loss = 0., 0., 0.\n        val_iou, val_f1, val_loss = 0., 0., 0.\n        \n        for idx, batch  in enumerate(tqdm(train_dataloader)):  \n            bands, mask, ash = batch['bands'], batch['mask'], batch['ash']\n            bands, ash = bands.to(device), ash.to(device)\n            # One-hot encode mask on the fly: for loss computations\n            # TODO: Move this transformation to contrailsDataset class\n            mask_oh = np.array([one_hot_encode(np.array(m).transpose(1,2,0), [0,1]).transpose(2,0,1) for m in mask])\n            mask_oh = torch.from_numpy(mask_oh).to(device)\n            \n            optimizer.zero_grad()\n        \n            # forward\n            out = model(ash)\n            \n            # loss\n            loss = dice_loss(out, mask_oh)\n            \n            # backpropagate gradients\n            loss.backward()\n            \n            # optimizer step\n            optimizer.step()\n            \n            # store eval metrics\n            tp, fp, fn, tn = smp.metrics.get_stats(out, mask_oh, mode='binary', threshold=0.5)\n            \n            iou_score = smp.metrics.iou_score(tp, fp, fn, tn, reduction=\"micro\")\n            f1_score = smp.metrics.f1_score(tp, fp, fn, tn, reduction=\"micro\")\n            \n            train_iou += iou_score.cpu().numpy()\n            train_f1 += f1_score.cpu().numpy()\n            train_loss += loss.detach().cpu().numpy()\n            \n            print(f\"step dice loss: {loss:.3f}\")\n            print(f\"\\nTrain Batch {idx+1}: Metrics\")\n            print(f\"-----------------------\\nStep f1 score: {f1_score:.3f}\")\n            print(f\"Step IoU score: {iou_score:.3f}\\n-----------------------\")\n        \n            # Visualize masks after every 'vis_after' epochs\n            if idx == 0 and not SUBMIT and (epoch % vis_after == 0):\n                visualize(out.detach(), ash, mask, f\"ep-{epoch}-batch-{idx}\")\n                \n        n = len(train_dataloader)\n        train_logs['iou'].append(1.*train_iou/n)\n        train_logs['f1'].append(1.*train_f1 / n)\n        train_logs['loss'].append(1.*train_loss/n) \n        \n        print(\"\\nValidating...\")\n        # Validation loop\n        model.eval()\n        with torch.inference_mode():\n            valid_iou, valid_f1, valid_loss = 0., 0., 0.\n            for idx, batch in enumerate(tqdm(val_dataloader)):\n                bands2, mask2, ash2 = batch['bands'], batch['mask'], batch['ash']\n                \n                bands2, ash2 = bands2.to(device), ash2.to(device)\n                # One-hot encode mask on the fly: for loss computations\n                # TODO: Move this transformation to contrailsDataset class\n                mask_oh2 = np.array([one_hot_encode(np.array(m).transpose(1,2,0), [0,1]).transpose(2,0,1) for m in mask2])\n                mask_oh2 = torch.from_numpy(mask_oh2).to(device)\n\n                out2 = model(ash2)\n                loss_val = dice_loss(out2, mask_oh2)\n                \n                # store eval metric\n                tp2, fp2, fn2, tn2 = smp.metrics.get_stats(out2, mask_oh2, mode='binary', threshold=0.5)\n            \n                iou_score2 = smp.metrics.iou_score(tp2, fp2, fn2, tn2, reduction=\"micro\")\n                f1_score2 = smp.metrics.f1_score(tp2, fp2, fn2, tn2, reduction=\"micro\")\n                \n                valid_loss += loss_val\n                valid_f1 += f1_score2.cpu().numpy()\n                valid_iou += iou_score2.cpu().numpy()\n                \n                print(f\"Val Batch {idx+1}: Metrics\")\n                print(f\"-----------------------\\nStep f1 score: {f1_score2:.3f}\")\n                print(f\"Step IoU score: {iou_score2:.3f}\\n-----------------------\")\n            \n            n_val = len(val_dataloader)\n            valid_logs['iou'].append(valid_iou / n_val)\n            valid_logs['f1'].append(valid_f1 / n_val)\n            valid_logs['loss'].append(valid_f1 / n_val)\n            \n            # Save model if a better val IoU score is obtained\n            if valid_logs['iou'][-1] > best_iou_score:\n                best_iou_score = valid_logs['iou'][-1]\n                print(f\"\\n*** Best Avg. Validation IoU Score: {valid_logs['iou'][-1]} ***\\nSaving model...\")\n                if not SUBMIT:\n                    torch.save({'epoch': epoch, \n                        'model_state_dict': model.state_dict(), \n                        'optimizer_state_dict': optimizer.state_dict(),\n                        'train_loss': train_logs['loss'][-1], \n                        'valid_iou': valid_logs['iou'][-1],\n                        'valid_f1': valid_logs['f1'][-1]\n                       }, PATH+f\"epoch-{epoch}_iou-{valid_logs['iou'][-1]:.2f}\")","metadata":{"execution":{"iopub.status.busy":"2023-05-15T10:59:24.107983Z","iopub.execute_input":"2023-05-15T10:59:24.108417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Metric Plots (train vs validation)","metadata":{}},{"cell_type":"code","source":"if not SUBMIT:\n    train_logs_df = pd.DataFrame(train_logs)\n    valid_logs_df = pd.DataFrame(valid_logs)\ntrain_logs_df.T","metadata":{"id":"Ifa_lObzThab","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMIT:\n    plt.figure(figsize=(20,8))\n    plt.plot(train_logs_df.index.tolist(), train_logs_df.iou.tolist(), lw=3, label = 'Train')\n    plt.plot(valid_logs_df.index.tolist(), valid_logs_df.iou.tolist(), lw=3, label = 'Valid')\n    plt.xlabel('Epochs', fontsize=20)\n    plt.ylabel('IoU Score', fontsize=20)\n    plt.title('IoU Score Plot', fontsize=20)\n    plt.legend(loc='best', fontsize=16)\n    plt.grid()\n    plt.savefig('iou_score_plot.png')\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if not SUBMIT:\n    plt.figure(figsize=(20,8))\n    plt.plot(train_logs_df.index.tolist(), train_logs_df.loss.tolist(), lw=3, label = 'Train')\n    plt.plot(valid_logs_df.index.tolist(), valid_logs_df.loss.tolist(), lw=3, label = 'Valid')\n    plt.xlabel('Epochs', fontsize=20)\n    plt.ylabel('Dice Loss', fontsize=20)\n    plt.title('Dice Loss Plot', fontsize=20)\n    plt.legend(loc='best', fontsize=16)\n    plt.grid()\n    plt.savefig('dice_loss_plot.png')\n    plt.show()","metadata":{"id":"mRnj6r29gHr-","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Run-Length Encoding and Decoding Code Reference: [Contrails RLE Submission by inversion](https://www.kaggle.com/code/inversion/contrails-rle-submission)","metadata":{}},{"cell_type":"code","source":"def rle_encode(y_pred, fg_val=1):\n    \"\"\"\n    Args:\n        x:  numpy array of shape (height, width), 1 - mask, 0 - background\n    Returns: run length encoding as list\n    \"\"\"\n    dots = np.where(\n        y_pred.T.flatten() == fg_val)[0]  # .T sets Fortran order down-then-right\n    run_lengths = []\n    prev = -2\n    for b in dots:\n        if b > prev + 1:\n            run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n    \n    def list_to_string(x):\n        if x: # non-empty list\n            s = str(x).replace(\"[\", \"\").replace(\"]\", \"\").replace(\",\", \"\")\n        else:\n            s = '-'\n        return s\n    \n    return list_to_string(run_lengths)\n\ndef rle_decode(mask_rle, shape=(256, 256)):\n    '''\n    mask_rle: run-length as string formatted (start length)\n              empty predictions need to be encoded with '-'\n    shape: (height, width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n    '''\n\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    if mask_rle != '-': \n        s = mask_rle.split()\n        starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n        starts -= 1\n        ends = starts + lengths\n        for lo, hi in zip(starts, ends):\n            img[lo:hi] = 1\n    return img.reshape(shape, order='F')  # Needed to align to RLE direction","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def submit_to_csv(test_ids, y_pred):\n    submission_df = pd.read_csv(os.path.join(CONFIG['DATA_ROOT'], \"sample_submission.csv\"), \n                                index_col='record_id')\n#     print(y_pred.shape)\n#     print(y_pred[0].shape)\n    y_encoded = [rle_encode(y) for y in y_pred]\n    submission_df[\"encoded_pixels\"] = y_encoded\n    submission_df.index = test_ids\n    print(submission_df.head(2))\n    submission_df.to_csv(\"submission.csv\")\n    print(\"Submitted\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Test:","metadata":{}},{"cell_type":"code","source":"out_rle = []\ntest_ids = []\nmodel.eval()\nwith torch.inference_mode():\n    for idx, batch in enumerate(tqdm(test_dataloader)):  \n        print(\"Test idx:\", idx)\n        bands,ash = batch['bands'], batch['ash']\n        ash = ash.to(device)\n        \n        out = model(ash)\n#         print(out.shape)\n#         print(reverse_one_hot(out[0].permute(1,2,0)).shape)\n        out_reversed = np.array([reverse_one_hot(o.cpu().permute(1,2,0).numpy()) for o in out])\n        for e in out_reversed:\n            out_rle.append(e)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\ntest_ids = list(os.listdir(\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/test\"))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit_to_csv(test_ids, out_reversed)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}