{"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":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","scrolled":true,"execution":{"iopub.status.busy":"2023-06-22T12:54:09.784667Z","iopub.execute_input":"2023-06-22T12:54:09.785191Z"},"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"markdown","source":"## Identifying Contrails to Reduce Global Warming\n\nCondensation trails, are line-shaped clouds of ice crystals which come from airplane's exhaust when they're flying through super humid regions in the atmosphere. It has been discovered that contrails contribute approximately 1% of all human caused global warming. ","metadata":{}},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\nfrom matplotlib import animation\nimport matplotlib.pyplot as plt\nfrom IPython import display\nfrom PIL import Image, ImageDraw\nfrom pathlib import Path\nimport torch\n\n\nBASE_DIR = Path('/kaggle/input/google-research-identify-contrails-reduce-global-warming')","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:07:30.789476Z","iopub.execute_input":"2023-07-04T14:07:30.789991Z","iopub.status.idle":"2023-07-04T14:07:34.616504Z","shell.execute_reply.started":"2023-07-04T14:07:30.789944Z","shell.execute_reply":"2023-07-04T14:07:34.615277Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if torch.cuda.is_available():\n    device = torch.device('cuda')\nelse:\n    device = torch.device('cpu')\n\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:07:37.556032Z","iopub.execute_input":"2023-07-04T14:07:37.55736Z","iopub.status.idle":"2023-07-04T14:07:37.593861Z","shell.execute_reply.started":"2023-07-04T14:07:37.557277Z","shell.execute_reply":"2023-07-04T14:07:37.592777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Information about Contrails and Visualizations:\nFor this competition we will be using geostationary satellite images to identify aviation contrails.\n\n- Contrails must contain at least 10 pixels\n- At some time in their life, Contrails must be at least 3x longer than they are wide\n- Contrails must either appear suddenly or enter from the sides of the image\n- Contrails should be visible in at least two image\n\nGround truth was determined by (generally) 4+ different labelers annotating each image. Pixels were considered a contrail when >50% of the labelers annotated it as such. Individual annotations (human_individual_masks.npy) as well as the aggregated ground truth annotations (human_pixel_masks.npy)\n\nI've found this pinned notebook (https://www.kaggle.com/code/inversion/visualizing-contrails) very helpful in order to understand and visualize the samples in the training set...","metadata":{}},{"cell_type":"code","source":"N_TIMES_BEFORE = 4\nN_TIMES_AFTER = 3\n\nrecord_id = '1000660467359258186'\n\nwith open(os.path.join(BASE_DIR, 'train', record_id, 'band_11.npy'), 'rb') as f:\n    band11 = np.load(f)\nwith open(os.path.join(BASE_DIR,  'train', record_id, 'band_14.npy'), 'rb') as f:\n    band14 = np.load(f)\nwith open(os.path.join(BASE_DIR, 'train', record_id, 'band_15.npy'), 'rb') as f:\n    band15 = np.load(f)\nwith open(os.path.join(BASE_DIR, 'train', record_id, 'human_pixel_masks.npy'), 'rb') as f:\n    human_pixel_mask = np.load(f)\nwith open(os.path.join(BASE_DIR, 'train', record_id, 'human_individual_masks.npy'), 'rb') as f:\n    human_individual_mask = np.load(f)","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:48:38.022229Z","iopub.execute_input":"2023-07-04T13:48:38.022582Z","iopub.status.idle":"2023-07-04T13:48:38.172889Z","shell.execute_reply.started":"2023-07-04T13:48:38.022555Z","shell.execute_reply":"2023-07-04T13:48:38.171686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"- band_{08-16}.npy: array with size of H x W x T, where T = n_times_before + n_times_after + 1, representing the number of images in the sequence. There are n_times_before and n_times_after images before and after the labeled frame respectively. In this dataset all examples have  n_times_before=4 and n_times_after=3. Each band represents an infrared channel at different wavelengths and is converted to brightness temperatures based on the calibration parameters. The number in the filename corresponds to the GOES-16 ABI band number.","metadata":{}},{"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    \"\"\"Maps data to the range [0, 1].\"\"\"\n    return (data - bounds[0]) / (bounds[1] - bounds[0])","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:48:43.289613Z","iopub.execute_input":"2023-07-04T13:48:43.290623Z","iopub.status.idle":"2023-07-04T13:48:43.296616Z","shell.execute_reply.started":"2023-07-04T13:48:43.290577Z","shell.execute_reply":"2023-07-04T13:48:43.295643Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"r = normalize_range(band15 - band14, _TDIFF_BOUNDS)\ng = normalize_range(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\nb = normalize_range(band14, _T11_BOUNDS)\nfalse_color = np.clip(np.stack([r, g, b], axis=2), 0, 1)","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:48:46.21931Z","iopub.execute_input":"2023-07-04T13:48:46.219678Z","iopub.status.idle":"2023-07-04T13:48:46.243854Z","shell.execute_reply.started":"2023-07-04T13:48:46.219645Z","shell.execute_reply":"2023-07-04T13:48:46.242928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = false_color[..., N_TIMES_BEFORE]\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(human_pixel_mask, interpolation='none')\nax.set_title('Ground truth contrail mask')\n\nax = plt.subplot(1, 3, 3)\nax.imshow(img)\nax.imshow(human_pixel_mask, cmap='Reds', alpha=.4, interpolation='none')\nax.set_title('Contrail mask on false color image');","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:48:48.314407Z","iopub.execute_input":"2023-07-04T13:48:48.314784Z","iopub.status.idle":"2023-07-04T13:48:49.349278Z","shell.execute_reply.started":"2023-07-04T13:48:48.314755Z","shell.execute_reply":"2023-07-04T13:48:49.347719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Individual human masks\nn = human_individual_mask.shape[-1]\nplt.figure(figsize=(16, 4))\nfor i in range(n):\n    plt.subplot(1, n, i+1)\n    plt.imshow(human_individual_mask[..., i], interpolation='none')","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:48:53.787947Z","iopub.execute_input":"2023-07-04T13:48:53.788535Z","iopub.status.idle":"2023-07-04T13:48:54.520393Z","shell.execute_reply.started":"2023-07-04T13:48:53.788503Z","shell.execute_reply":"2023-07-04T13:48:54.519526Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Animation\nfig = plt.figure(figsize=(6, 6))\nim = plt.imshow(false_color[..., 0])\ndef draw(i):\n    im.set_array(false_color[..., i])\n    return [im]\nanim = animation.FuncAnimation(\n    fig, draw, frames=false_color.shape[-1], interval=500, blit=True\n)\nplt.close()\ndisplay.HTML(anim.to_jshtml())","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:48:57.65858Z","iopub.execute_input":"2023-07-04T13:48:57.658958Z","iopub.status.idle":"2023-07-04T13:49:00.033635Z","shell.execute_reply.started":"2023-07-04T13:48:57.658927Z","shell.execute_reply":"2023-07-04T13:49:00.032784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We can start checking the number of samples on each dataset.","metadata":{}},{"cell_type":"code","source":"train_list = os.listdir(BASE_DIR / 'train')\nval_list = os.listdir(BASE_DIR / 'validation/')\ntest_list = os.listdir(BASE_DIR / 'test/')\n\n\nprint (f\"Number of training examples: {len(train_list)}\")\nprint(f\"Number of validation examples: {len(val_list)}\")\nprint(f\"Number of test examples: {len(test_list)}\")","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:49:05.412239Z","iopub.execute_input":"2023-07-04T13:49:05.412605Z","iopub.status.idle":"2023-07-04T13:49:05.813687Z","shell.execute_reply.started":"2023-07-04T13:49:05.412576Z","shell.execute_reply":"2023-07-04T13:49:05.812587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_metadata = pd.read_json(\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train_metadata.json\")\nval_metadata = pd.read_json(\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/validation_metadata.json\")","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:49:09.470881Z","iopub.execute_input":"2023-07-04T13:49:09.471227Z","iopub.status.idle":"2023-07-04T13:49:10.008242Z","shell.execute_reply.started":"2023-07-04T13:49:09.471201Z","shell.execute_reply":"2023-07-04T13:49:10.007196Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm import tqdm\nfrom pathlib import Path\n\nrows = []\nfor example in tqdm(train_list):\n    im = np.load(f\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train/{example}/human_pixel_masks.npy\")\n    if len(np.unique(im)) > 1:\n        rows.append({'record_id': int(example), 'contrails': True})\n    else:\n        rows.append({'record_id': int(example), 'contrails': False})\n    \ntrain_class_df = pd.DataFrame(rows)\ndel rows\n\nfilepath = '/kaggle/working/train_record_class.csv'\ntrain_class_df.to_csv(filepath)\n","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-07-04T13:49:12.621574Z","iopub.execute_input":"2023-07-04T13:49:12.622068Z","iopub.status.idle":"2023-07-04T13:52:25.767091Z","shell.execute_reply.started":"2023-07-04T13:49:12.622028Z","shell.execute_reply":"2023-07-04T13:52:25.766132Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rows = []\nfor example in tqdm(val_list):\n    im = np.load(f\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/validation/{example}/human_pixel_masks.npy\")\n    if len(np.unique(im)) > 1:\n        rows.append({'record_id': int(example), 'contrails': True})\n    else:\n        rows.append({'record_id': int(example), 'contrails': False})\n    \nval_class_df = pd.DataFrame(rows)\nval_class_df.astype({'record_id': 'int64'}).dtypes\ndel rows\n\nfilepath = '/kaggle/working/val_record_class.csv'\nval_class_df.to_csv(filepath)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-07-04T13:53:12.921227Z","iopub.execute_input":"2023-07-04T13:53:12.921608Z","iopub.status.idle":"2023-07-04T13:53:28.326484Z","shell.execute_reply.started":"2023-07-04T13:53:12.921578Z","shell.execute_reply":"2023-07-04T13:53:28.325533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv('/kaggle/working/train_record_class.csv')\nval_df = pd.read_csv('/kaggle/working/val_record_class.csv')","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:07:49.374487Z","iopub.execute_input":"2023-07-04T14:07:49.37494Z","iopub.status.idle":"2023-07-04T14:07:49.417624Z","shell.execute_reply.started":"2023-07-04T14:07:49.374906Z","shell.execute_reply":"2023-07-04T14:07:49.416498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_class_df['contrails'].value_counts().plot(kind='pie', y='contrails', title='Training Data Class Distr.', autopct='%1.1f%%', \\\n                   shadow=True, startangle=0, labels=['Contrails Not Detected', 'Contrails Detected'])","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:53:35.882329Z","iopub.execute_input":"2023-07-04T13:53:35.88268Z","iopub.status.idle":"2023-07-04T13:53:36.075786Z","shell.execute_reply.started":"2023-07-04T13:53:35.882652Z","shell.execute_reply":"2023-07-04T13:53:36.07458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_class_df['contrails'].value_counts().plot(kind='pie', y='contrails', title='Val Data Class Distr.', autopct='%1.1f%%',\\\n                   shadow=True, startangle=0, labels=['Contrails Not Detected', 'Contrails Detected'])","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:53:38.73541Z","iopub.execute_input":"2023-07-04T13:53:38.735775Z","iopub.status.idle":"2023-07-04T13:53:38.910816Z","shell.execute_reply.started":"2023-07-04T13:53:38.735745Z","shell.execute_reply":"2023-07-04T13:53:38.909658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As we can see that, this dataset is very unbalanced and we have way more samples which does not include contrails than samples which contains contrails. For this reason we should hold this information in a dataframe for determining further sampling strategies in the training to achieve better results.","metadata":{}},{"cell_type":"code","source":"train_metadata = train_metadata.merge(train_class_df, how='inner', on='record_id')","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:53:45.335083Z","iopub.execute_input":"2023-07-04T13:53:45.335437Z","iopub.status.idle":"2023-07-04T13:53:45.357392Z","shell.execute_reply.started":"2023-07-04T13:53:45.335409Z","shell.execute_reply":"2023-07-04T13:53:45.356452Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"samples_with_contrails_df = train_metadata[train_metadata['contrails'] == True]\nsamples_with_contrails_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:53:48.017986Z","iopub.execute_input":"2023-07-04T13:53:48.018414Z","iopub.status.idle":"2023-07-04T13:53:48.039727Z","shell.execute_reply.started":"2023-07-04T13:53:48.018381Z","shell.execute_reply":"2023-07-04T13:53:48.038674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Pytorch Dataset Creation","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset\nimport torchvision.transforms.functional as TF\nfrom torchvision import transforms\nimport random\nimport albumentations as A\n\n\n_T11_BOUNDS = (243, 303)\n_CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n_TDIFF_BOUNDS = (-4, 2)\n\nclass SegmentationDataset(Dataset):\n    def __init__(self, root_dir=BASE_DIR, mode='train', should_transform=False):\n        self.root_dir = root_dir\n        self.mode = mode\n        self.should_transform = should_transform\n        self.id2label =  {'0': 'background', '1': 'contrails'}\n        self.records = os.listdir(self.root_dir / self.mode)\n        \n#         if self.mode == 'train': \n#             select_cnt = 700\n#             self.records = np.random.choice(self.records, select_cnt, replace=False)\n        \n#         if self.mode == 'validation':\n#             select_cnt = 180\n#             self.records = np.random.choice(self.records, select_cnt, replace=False)\n\n    def __len__(self):\n        return len(self.records)\n    \n    def transform(self, image, mask):\n        aug = A.Compose([\n                A.VerticalFlip(p=0.3),\n                A.HorizontalFlip(p=0.3),      \n                A.RandomRotate90(p=0.3),\n                ]\n            )\n\n        random.seed(7)\n        augmented = aug(image=image, mask=mask)\n\n        image_transformed = augmented['image']\n        mask_transformed = augmented['mask']\n        \n        return image_transformed, mask_transformed\n    \n    def normalize_range(self, data, bounds):\n        return (data - bounds[0]) / (bounds[1] - bounds[0])\n\n    def get_ash_img(self, bands):\n        band11 = bands[:,:,0,11-8]\n        band14 = bands[:,:,0,14-8]\n        band15 = bands[:,:,0,15-8]\n        r = self.normalize_range(band15 - band14, _TDIFF_BOUNDS)\n        g = self.normalize_range(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\n        b = self.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    def __getitem__(self, idx):\n        record_id = self.records[idx]\n        record_dir = os.path.join(self.root_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        # 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            if self.should_transform:\n                ash, pixel_masks = self.transform(ash, pixel_masks)\n                \n            sample = {'pixel_values': ash, 'labels': pixel_masks}\n        else:\n            sample = {'record_ids': str(record_dir.split('/')[-1]), 'pixel_values': ash}\n      \n\n        return sample\n","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:07:52.950599Z","iopub.execute_input":"2023-07-04T14:07:52.951032Z","iopub.status.idle":"2023-07-04T14:07:55.601731Z","shell.execute_reply.started":"2023-07-04T14:07:52.951Z","shell.execute_reply":"2023-07-04T14:07:55.60039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Sampling","metadata":{}},{"cell_type":"markdown","source":"I will use Weighted Random Sampler for handling inbalanced classes.","metadata":{}},{"cell_type":"code","source":"def get_class_weights(dataframe):\n    # calculate class and sample weights\n    class_counts = dataframe.contrails.value_counts()\n    class_weights = 1 / class_counts\n    sample_weights = [class_weights[i] for i in dataframe.contrails.values]\n    return sample_weights","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:08:05.726136Z","iopub.execute_input":"2023-07-04T14:08:05.72654Z","iopub.status.idle":"2023-07-04T14:08:05.73526Z","shell.execute_reply.started":"2023-07-04T14:08:05.726509Z","shell.execute_reply":"2023-07-04T14:08:05.731272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import WeightedRandomSampler\n\ndef create_sampler(mode):\n    if mode == 'train':\n        sample_w = get_class_weights(train_df)\n        return WeightedRandomSampler(weights=sample_w,num_samples=len(train_df))\n    elif mode=='val':\n        sample_w = get_class_weights(val_df)\n        return WeightedRandomSampler(weights=sample_w,num_samples=len(val_df))\n    else:\n        return None","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:08:03.147368Z","iopub.execute_input":"2023-07-04T14:08:03.147789Z","iopub.status.idle":"2023-07-04T14:08:03.155785Z","shell.execute_reply.started":"2023-07-04T14:08:03.147758Z","shell.execute_reply":"2023-07-04T14:08:03.153687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Creating Dataloaders with Weighted Random Sampling","metadata":{}},{"cell_type":"code","source":"from transformers import SegformerImageProcessor\nfrom torch.utils.data import DataLoader\n                                            \nbatch_size = 8\n\ntrain_dataset = SegmentationDataset(mode='train', should_transform=True)\nval_dataset = SegmentationDataset(mode='validation', should_transform=True)\ntest_dataset = SegmentationDataset(mode='test', should_transform=False)\n\ntrain_sampler = create_sampler('train')\nval_sampler = create_sampler('val')\n\ntrain_dataloader = DataLoader(train_dataset, batch_size=batch_size, num_workers=2, sampler=train_sampler, prefetch_factor=8)\nval_dataloader = DataLoader(val_dataset, batch_size=batch_size, num_workers=2, sampler=val_sampler, prefetch_factor=8)\ntest_dataloader = DataLoader(test_dataset, batch_size=batch_size, num_workers=2, prefetch_factor=8)","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:08:10.264374Z","iopub.execute_input":"2023-07-04T14:08:10.26497Z","iopub.status.idle":"2023-07-04T14:08:19.891898Z","shell.execute_reply.started":"2023-07-04T14:08:10.264932Z","shell.execute_reply":"2023-07-04T14:08:19.890625Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"s = next(iter(val_dataloader))\nash = s['pixel_values']\nmask = s['labels']\nprint(ash.shape) \nprint(mask.shape) ","metadata":{"execution":{"iopub.status.busy":"2023-07-04T13:56:17.564913Z","iopub.execute_input":"2023-07-04T13:56:17.565372Z","iopub.status.idle":"2023-07-04T13:56:22.117787Z","shell.execute_reply.started":"2023-07-04T13:56:17.565337Z","shell.execute_reply":"2023-07-04T13:56:22.11646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sampler_val = np.random.randint(0, 8)\n\nimg = ash[sampler_val, :, :, :]\nmask_show = mask[sampler_val, :,:,:]\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-07-04T13:56:38.71089Z","iopub.execute_input":"2023-07-04T13:56:38.711929Z","iopub.status.idle":"2023-07-04T13:56:39.988918Z","shell.execute_reply.started":"2023-07-04T13:56:38.711889Z","shell.execute_reply":"2023-07-04T13:56:39.987882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Custom Loss","metadata":{}},{"cell_type":"code","source":"ALPHA = 0.5\nBETA = 0.5\nGAMMA = 1\n\nclass FocalTverskyLoss(nn.Module):\n    def __init__(self, weight=None, size_average=True):\n        super(FocalTverskyLoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=0.5, alpha=ALPHA, beta=BETA, gamma=GAMMA):\n        #flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n    \n        TP = (inputs * targets).sum()    \n        FP = ((1-targets) * inputs).sum()\n        FN = (targets * (1-inputs)).sum()\n        \n        Tversky = (TP + smooth) / (TP + alpha*FP + beta*FN + smooth)  \n        FocalTversky = (1 - Tversky)**gamma\n        return FocalTversky","metadata":{"execution":{"iopub.status.busy":"2023-06-17T12:22:44.686445Z","iopub.execute_input":"2023-06-17T12:22:44.686984Z","iopub.status.idle":"2023-06-17T12:22:44.701338Z","shell.execute_reply.started":"2023-06-17T12:22:44.686946Z","shell.execute_reply":"2023-06-17T12:22:44.700116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model","metadata":{}},{"cell_type":"code","source":"import wandb\nfrom kaggle_secrets import UserSecretsClient\n\nuser_secrets = UserSecretsClient()\nwandb_api = user_secrets.get_secret(\"wandb.api.key\") \n\nwandb.login(key=wandb_api)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-07-04T14:08:19.894154Z","iopub.execute_input":"2023-07-04T14:08:19.895381Z","iopub.status.idle":"2023-07-04T14:08:25.482519Z","shell.execute_reply.started":"2023-07-04T14:08:19.895344Z","shell.execute_reply":"2023-07-04T14:08:25.481392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize(out: torch.Tensor, ash: torch.Tensor, mask=None, title='', test=False):\n    out_ = np.array(out.cpu())\n    ash_ = np.array(ash.cpu())\n    if not test:\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(out_[i])\n        ash_imgs.append(ash_[i])\n        if not test:\n            mask_imgs.append(mask_[i])\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\"Predicted_{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\"False Color_{i+1}\")\n    \n    if not test:\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\"Label_{i+1}\")\n        \n    plt.savefig(title)\n    plt.show() ","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:08:26.660897Z","iopub.execute_input":"2023-07-04T14:08:26.661788Z","iopub.status.idle":"2023-07-04T14:08:26.67515Z","shell.execute_reply.started":"2023-07-04T14:08:26.66175Z","shell.execute_reply":"2023-07-04T14:08:26.673923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -U git+https://github.com/qubvel/segmentation_models.pytorch\nimport segmentation_models_pytorch as smp\n","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-07-04T14:08:29.547487Z","iopub.execute_input":"2023-07-04T14:08:29.547876Z","iopub.status.idle":"2023-07-04T14:09:05.497021Z","shell.execute_reply.started":"2023-07-04T14:08:29.547846Z","shell.execute_reply":"2023-07-04T14:09:05.495745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn\nfrom transformers import SegformerForSemanticSegmentation\nfrom datasets import load_metric\nfrom torchmetrics.functional import dice\nimport pytorch_lightning as pl\n\n\nclass SegFormerModel(pl.LightningModule):\n    def __init__(self, id2label, train_dataloader=None, val_dataloader=None, test_dataloader=None, metrics_interval=10, model=None):\n        super(SegFormerModel, self).__init__()\n        self.id2label = id2label\n        self.metrics_interval = metrics_interval\n        self.train_dl = train_dataloader\n        self.val_dl = val_dataloader\n        self.test_dl = test_dataloader\n        self.label2id = {v:k for k,v in self.id2label.items()}\n        self.num_classes = len(id2label.keys())\n        self.model = model or get_initial_model(self.num_classes)\n        self.loss_module = smp.losses.DiceLoss(mode=\"binary\", smooth=1.0, from_logits=True)\n        self.train_step_ious= []\n        self.validation_step_ious = []\n        self.validation_step_outputs = []\n        self.test_step_outputs = []\n        self.save_hyperparameters()\n        \n        \n    def forward(self, images, masks=None):\n        outputs = self.model(pixel_values=images)\n        return outputs\n    \n    def training_step(self, batch, batch_idx):\n        masks =  torch.squeeze(batch['labels']).long().to(device)\n        masks = nn.functional.one_hot(masks, num_classes=self.num_classes).permute(0, 3, 1, 2).contiguous().to(device)\n        images = batch['pixel_values'].permute(0, 3, 1, 2).to(device)\n        \n        outputs = self.model(pixel_values=images, return_dict=True)\n        \n        upsampled_logits = nn.functional.interpolate(\n            outputs.logits, \n            size=masks.shape[-2:], \n            mode=\"bilinear\", \n            align_corners=False\n        ).contiguous().to(device)\n\n        # predicted = upsampled_logits.argmax(dim=1)\n        loss = self.loss_module(upsampled_logits, masks)\n        tp, fp, fn, tn = smp.metrics.get_stats((upsampled_logits.sigmoid()>0.5).long(), masks.long(), mode='binary')\n        iou = smp.metrics.iou_score(tp, fp, fn, tn, reduction=\"micro-imagewise\")\n        self.train_step_ious.append(iou)\n        \n\n        if batch_idx % self.metrics_interval == 0:\n            mean_iou = torch.stack(self.train_step_ious).mean()\n            # Log loss and metric\n            self.log('train_loss', loss)\n            self.log('train_mean_iou',  mean_iou)\n            \n            print(f\"Training loss: {loss:.5f}\")\n            print(\"\\n-----------------------\")\n\n        return {'loss': loss}\n\n    def validation_step(self, batch, batch_idx):\n        masks =  torch.squeeze(batch['labels']).long().to(device)\n        masks = nn.functional.one_hot(masks, num_classes=self.num_classes).permute(0, 3, 1, 2).contiguous().to(device)\n        images = batch['pixel_values'].permute(0, 3, 1, 2).to(device)\n        \n        outputs = self.model(pixel_values=images, return_dict=True)\n        \n        upsampled_logits = nn.functional.interpolate(\n            outputs.logits, \n            size=masks.shape[-2:], \n            mode=\"bilinear\", \n            align_corners=False\n        ).contiguous()\n\n        predicted = upsampled_logits.argmax(dim=1).to(device)\n        loss = self.loss_module(upsampled_logits, masks)\n        \n        tp, fp, fn, tn = smp.metrics.get_stats((upsampled_logits.sigmoid()>0.5).long(), masks.long(), mode='binary')\n        iou = smp.metrics.iou_score(tp, fp, fn, tn, reduction=\"micro-imagewise\")\n        \n        self.validation_step_ious.append(iou)\n        self.validation_step_outputs.append(loss)\n        \n        # Log loss and metric\n        self.log('val_loss', loss)\n        self.log(f\"IoU\", iou)\n        \n        print(f\"Val Batch {batch_idx+1}: Metrics\")\n        print(f\"-----------------------\\nStep Validation Loss: {loss:.5f}\")\n        print(\"\\n-----------------------\")\n        \n        return {'val_loss': loss, 'predicted': predicted}\n    \n    def on_validation_epoch_end(self):\n        epoch_average_loss = torch.stack(self.validation_step_outputs).mean()\n        val_step_mean_iou = torch.stack(self.validation_step_ious).mean()\n \n        metrics = {\"val_loss\": epoch_average_loss, \"val_mean_iou\":val_step_mean_iou, }\n        \n        print(f\"Val Epoch Metrics\")\n        print(f\"Epoch IoU score: {val_step_mean_iou:.3f}\\n-----------------------\")    \n        self.validation_step_outputs.clear()  # free memory\n        return metrics\n    \n       \n    def predict_step(self, batch, batch_idx, dataloader_idx=0):\n        images = batch['pixel_values'].permute(0, 3, 1, 2).to(device)\n        return self.model(images, return_dict=True)\n    \n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam([p for p in self.parameters() if p.requires_grad], lr=2e-05, eps=1e-08)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max= 500, eta_min= 1e-06, last_epoch= -1)\n        return {\"optimizer\": optimizer, \"lr_scheduler\": {\"scheduler\": scheduler, \"interval\": \"step\"}, \"monitor\": \"val_loss\"}\n    \n    def train_dataloader(self):\n        return self.train_dl\n    \n    def val_dataloader(self):\n        return self.val_dl\n    \n    def test_dataloader(self):\n        return self.test_dl","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:09:14.199558Z","iopub.execute_input":"2023-07-04T14:09:14.200051Z","iopub.status.idle":"2023-07-04T14:09:16.959015Z","shell.execute_reply.started":"2023-07-04T14:09:14.200001Z","shell.execute_reply":"2023-07-04T14:09:16.957991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_initial_model(num_classes):\n    return SegformerForSemanticSegmentation.from_pretrained(\n            \"nvidia/mit-b3\", \n            return_dict=True, \n            num_labels=num_classes,\n            ignore_mismatched_sizes=True,\n        )","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:09:26.399899Z","iopub.execute_input":"2023-07-04T14:09:26.401191Z","iopub.status.idle":"2023-07-04T14:09:26.408575Z","shell.execute_reply.started":"2023-07-04T14:09:26.401137Z","shell.execute_reply":"2023-07-04T14:09:26.40749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"segformer = SegFormerModel(\n    train_dataset.id2label, \n    train_dataloader=train_dataloader, \n    val_dataloader=val_dataloader,\n    test_dataloader=test_dataloader,\n    metrics_interval=5,\n)","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:09:29.009799Z","iopub.execute_input":"2023-07-04T14:09:29.010208Z","iopub.status.idle":"2023-07-04T14:09:33.394794Z","shell.execute_reply.started":"2023-07-04T14:09:29.010175Z","shell.execute_reply":"2023-07-04T14:09:33.393668Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"from pytorch_lightning.callbacks import Callback\n\nclass LogPredictionsCallback(Callback):\n    def __init__(self, num_samples=8):\n        super().__init__()\n        self.num_samples = num_samples\n\n    \n    def on_validation_batch_end(self, trainer, pl_module, outputs, batch, batch_idx, dataloader_idx=0):\n        \"\"\"Called when the validation batch ends.\"\"\"\n        ash=batch['pixel_values']\n        mask=batch['labels']\n        out=outputs['predicted'][:, :, :, None]\n\n        visualize(outputs['predicted'], batch['pixel_values'], batch['labels'], f\"batch-{batch_idx}\", False)\nlog_predictions_callback = LogPredictionsCallback()\n","metadata":{"execution":{"iopub.status.busy":"2023-07-04T14:09:36.559095Z","iopub.execute_input":"2023-07-04T14:09:36.55953Z","iopub.status.idle":"2023-07-04T14:09:36.568943Z","shell.execute_reply.started":"2023-07-04T14:09:36.559496Z","shell.execute_reply":"2023-07-04T14:09:36.567259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pytorch_lightning.callbacks.early_stopping import EarlyStopping\nfrom pytorch_lightning.callbacks.model_checkpoint import ModelCheckpoint\nfrom pytorch_lightning.loggers import TensorBoardLogger\nimport pytorch_lightning as pl\n\nearly_stop_callback = EarlyStopping(\n    monitor=\"val_loss\", \n    min_delta=0.00, \n    patience=3, \n    verbose=False, \n    mode=\"min\",\n)\n\n\n# default logger used by trainer (if tensorboard is installed)\n# logger = TensorBoardLogger(save_dir=os.getcwd(), version=1, name=\"lightning_logs\")\ncheckpoint_callback = ModelCheckpoint(save_top_k=1, monitor=\"val_loss\")\nlog_predictions_callback = LogPredictionsCallback()\nwandb_logger = pl.loggers.WandbLogger(project='Contrails', log_model='all') \n\ntrainer = pl.Trainer(\n    callbacks=[early_stop_callback, log_predictions_callback, checkpoint_callback],\n    max_epochs=6,\n    val_check_interval=80,\n    log_every_n_steps=10,\n    logger=wandb_logger,\n)\n\ntrainer.fit(segformer)\nwandb.finish()","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-07-04T14:11:13.462903Z","iopub.execute_input":"2023-07-04T14:11:13.463329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load a checkpoint","metadata":{}},{"cell_type":"code","source":"# artifact = run.use_artifact('xxx', type='model')\n# artifact_dir = artifact.download()","metadata":{"execution":{"iopub.status.busy":"2023-07-01T16:37:17.235448Z","iopub.execute_input":"2023-07-01T16:37:17.236023Z","iopub.status.idle":"2023-07-01T16:37:19.234218Z","shell.execute_reply.started":"2023-07-01T16:37:17.235968Z","shell.execute_reply":"2023-07-01T16:37:19.233391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# load checkpoint\n\n# checkpoint_model = segformer.load_from_checkpoint(Path(artifact_dir) / \"model.ckpt\", map_location=torch.device('cpu'))\n# checkpoint_model.eval()","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2023-07-01T16:37:26.78248Z","iopub.execute_input":"2023-07-01T16:37:26.782987Z","iopub.status.idle":"2023-07-01T16:37:28.713969Z","shell.execute_reply.started":"2023-07-01T16:37:26.782945Z","shell.execute_reply":"2023-07-01T16:37:28.712785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"segformer.eval()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Submission","metadata":{}},{"cell_type":"code","source":"def get_test_image(batch_files):\n    _T11_BOUNDS = (243, 303)\n    _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n    _TDIFF_BOUNDS = (-4, 2)\n\n    band11 = tf.py_function(lambda path: np.load(path.numpy()), [batch_files[0]], tf.float32)\n    band14 = tf.py_function(lambda path: np.load(path.numpy()), [batch_files[1]], tf.float32)\n    band15 = tf.py_function(lambda path: np.load(path.numpy()), [batch_files[2]], tf.float32)\n    \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 = tf.py_function(lambda r, g, b: np.clip(np.stack([r.numpy(), g.numpy(), b.numpy()], axis=2), 0, 1), [r, g, b], tf.float32)\n\n    false_color = false_color[...,4]\n    \n    false_color.set_shape((256, 256, 3)) \n    \n    false_color = tf.expand_dims(false_color, axis=0)\n    return false_color","metadata":{"execution":{"iopub.status.busy":"2023-07-01T11:15:39.459477Z","iopub.execute_input":"2023-07-01T11:15:39.460261Z","iopub.status.idle":"2023-07-01T11:15:39.473078Z","shell.execute_reply.started":"2023-07-01T11:15:39.46022Z","shell.execute_reply":"2023-07-01T11:15:39.471583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\n    Helper functions for submission in Run-Length-Encoding.\n    Link: https://www.kaggle.com/code/inversion/contrails-rle-submission\n\"\"\"\ndef rle_encode(x, 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        x.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    return run_lengths\n\n\ndef list_to_string(x):\n    \"\"\"\n        Converts list to a string representation\n        Empty list returns '-'\n    \"\"\"\n    if x: # non-empty list\n        s = str(x).replace(\"[\", \"\").replace(\"]\", \"\").replace(\",\", \"\")\n    else:\n        s = '-'\n    return s\n\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    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":{"execution":{"iopub.status.busy":"2023-07-01T16:37:35.585089Z","iopub.execute_input":"2023-07-01T16:37:35.585556Z","iopub.status.idle":"2023-07-01T16:37:35.601387Z","shell.execute_reply.started":"2023-07-01T16:37:35.585519Z","shell.execute_reply":"2023-07-01T16:37:35.600447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_first = pd.read_csv('/kaggle/input/google-research-identify-contrails-reduce-global-warming/sample_submission.csv', index_col='record_id')","metadata":{"execution":{"iopub.status.busy":"2023-07-01T13:07:18.787965Z","iopub.execute_input":"2023-07-01T13:07:18.788445Z","iopub.status.idle":"2023-07-01T13:07:18.804731Z","shell.execute_reply.started":"2023-07-01T13:07:18.788411Z","shell.execute_reply":"2023-07-01T13:07:18.80336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tr = pl.Trainer()\nsegformer.eval()\noutputs = tr.predict(checkpoint_model, test_dataloader)\n\nfor i, data in enumerate(test_dataloader):\n    image = data['pixel_values'].permute(0, 3, 1, 2).to('cpu')\n    upsampled_logits = nn.functional.interpolate(\n        outputs[i].logits, \n        size=image.shape[-2:], \n        mode=\"bilinear\", \n        align_corners=False\n    ).to('cpu')\n\n    predicted = (upsampled_logits.sigmoid()>0.5).long()\n    pr = upsampled_logits.argmax(dim=1).to('cpu')\n    visualize(pr, data['pixel_values'], None, f\"test-batch\", True)\n#     for i, record_id in enumerate(test_dataloader['record_ids']):\n#         submission.loc[record_id, 'encoded_pixels'] = list_to_string(rle_encode(predicted[i, :, :, :]))","metadata":{"execution":{"iopub.status.busy":"2023-07-01T16:37:41.826977Z","iopub.execute_input":"2023-07-01T16:37:41.827662Z","iopub.status.idle":"2023-07-01T16:37:44.747806Z","shell.execute_reply.started":"2023-07-01T16:37:41.827618Z","shell.execute_reply":"2023-07-01T16:37:44.746048Z"},"trusted":true},"execution_count":null,"outputs":[]}]}