{"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":"# Google Research - Identify Contrails to Reduce Global Warming\n\n### Import modules","metadata":{}},{"cell_type":"code","source":"# basic modules\nimport pandas as pd\nimport numpy as np\nimport os\n\n# Visualization\nimport matplotlib.pyplot as plt\nfrom matplotlib import animation\nfrom IPython import display\n\n# pytorch modules\nimport torch\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-11T10:05:10.126313Z","iopub.execute_input":"2023-05-11T10:05:10.127803Z","iopub.status.idle":"2023-05-11T10:05:11.992683Z","shell.execute_reply.started":"2023-05-11T10:05:10.127741Z","shell.execute_reply":"2023-05-11T10:05:11.991485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"base_dir = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/\"\ntrain_path = os.path.join(base_dir,\"train\")\ntest_path = os.path.join(base_dir,\"test\")\nval_path = os.path.join(base_dir,\"validation\")\n\ntrain_ids = os.listdir(train_path)\ntest_ids = os.listdir(test_path)\nval_ids = os.listdir(val_path)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:05:11.994958Z","iopub.execute_input":"2023-05-11T10:05:11.995749Z","iopub.status.idle":"2023-05-11T10:05:12.39799Z","shell.execute_reply.started":"2023-05-11T10:05:11.995702Z","shell.execute_reply":"2023-05-11T10:05:12.397026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Combine bands into a false color image (ASH Transform)\n\n[In order to view contrails in GOES, we use the \"ash\" color scheme. This color scheme was originally developed for viewing volcanic ash in the atmosphere but is also useful for viewing thin cirrus, including contrails. In this color scheme, contrails appear in the image as dark blue.](https://www.kaggle.com/code/inversion/visualizing-contrails)\n\n- `Input` = np.ndarray of shape `(Bands=3, Time_frame=8, H=256, W=256)` where Bands corresponds to bands `[11, 14, 15]`.\n- `Output` = np.ndarray of shape `(Time_frame=8, Channel=3, H=256, W=256)` where Channel corresponds to rgb colorscheme.","metadata":{}},{"cell_type":"code","source":"def ash_transform(x):\n    _T11_BOUNDS = (243, 303)\n    _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n    _TDIFF_BOUNDS = (-4, 2)\n\n    def normalize_range(data, bounds):\n        \"\"\"Maps data to the range [0, 1].\"\"\"\n        return (data - bounds[0]) / (bounds[1] - bounds[0])\n\n    r = normalize_range(x[2] - x[1], _TDIFF_BOUNDS)\n    g = normalize_range(x[1] - x[0], _CLOUD_TOP_TDIFF_BOUNDS)\n    b = normalize_range(x[1], _T11_BOUNDS)\n    return np.clip(np.stack([r, g, b], axis=-3), 0, 1) # (T,3,H,W)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:05:14.180743Z","iopub.execute_input":"2023-05-11T10:05:14.18116Z","iopub.status.idle":"2023-05-11T10:05:14.190816Z","shell.execute_reply.started":"2023-05-11T10:05:14.18113Z","shell.execute_reply":"2023-05-11T10:05:14.189338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### CustomDataset and DataLoader\n\n- `ids` = record ids\n- `base_dir` = path to parent(train/val/test) dir\n- `bands` = list of integers corresponding bands(band_xx.npy). Default:loads all bands \n- `transforms` = list of function applied to np.ndarray of shape `(Band, Time_frame, H, W)`\n- `Output` = torch.tensor, dtype = `torch.float32`","metadata":{}},{"cell_type":"code","source":"class ContrailDataset(Dataset):\n    def __init__(self, ids, base_dir, bands=None, transforms:list=[], test_mode:bool=False):\n        self.ids = ids\n        self.base_dir = base_dir\n        self.transforms = transforms\n        self.bands = bands\n        self.permute = (2,0,1)\n        self.test_mode = test_mode\n        \n    def __getitem__(self, index):\n        record_id = self.ids[index]\n        \n        if self.bands is None:\n            band_list = [f'band_{band:02d}.npy' for band in range(8,17)]\n        else :\n            band_list = [f'band_{int(band):02d}.npy' for band in self.bands]\n        \n        x = list()\n        for band in band_list:\n            with open(os.path.join(self.base_dir, record_id, band), 'rb') as f:\n                x.append(np.load(f).transpose(self.permute))\n        x = np.stack(x,axis=0) ## X.shape = (Band,Time_frame,H,W)\n        \n        for transformation in self.transforms:\n            x = transformation(x)\n        x = torch.from_numpy(x.astype(np.float32))\n        \n        if self.test_mode:\n            return x\n        else:\n            with open(os.path.join(self.base_dir, record_id,'human_pixel_masks.npy'), 'rb') as f:\n                y = torch.from_numpy(np.load(f).squeeze().astype(np.float32))\n\n            return x, y\n\n    def __len__(self):\n        return len(self.ids)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:28:13.031275Z","iopub.execute_input":"2023-05-11T10:28:13.031706Z","iopub.status.idle":"2023-05-11T10:28:13.048546Z","shell.execute_reply.started":"2023-05-11T10:28:13.03167Z","shell.execute_reply":"2023-05-11T10:28:13.045174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Initiate Dataset and DataLoader","metadata":{}},{"cell_type":"code","source":"# Datasets \ndataset_params = {\n    \"bands\" : [11,14,15], \n    \"transforms\" : [ash_transform]\n}\ntrain_dataset = ContrailDataset(train_ids, train_path, **dataset_params)\ntest_dataset = ContrailDataset(test_ids, test_path, test_mode=True, **dataset_params)\nval_dataset = ContrailDataset(val_ids, val_path, **dataset_params)\n\n# DalaLoaders\ndataloader_params = {\n    \"batch_size\" : 32, \n    \"shuffle\" : True,\n    \"num_workers\": 0\n}\ntrain_loader = DataLoader(train_dataset, **dataloader_params)\ntest_loader = DataLoader(test_dataset, shuffle=False, batch_size=2)\nval_loader = DataLoader(val_dataset, **dataloader_params)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:42:00.806817Z","iopub.execute_input":"2023-05-11T10:42:00.807288Z","iopub.status.idle":"2023-05-11T10:42:00.816057Z","shell.execute_reply.started":"2023-05-11T10:42:00.807255Z","shell.execute_reply":"2023-05-11T10:42:00.814952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualization","metadata":{}},{"cell_type":"code","source":"def plot_contrail(x, y, time_frame = 4):\n    '''\n    x = false color img of shape (8, 3, H, W)\n    y = contrail mask of shape (H, W)\n    time_frame = int, default = 4\n    '''\n    plt.figure(figsize=(18, 6))\n    ax = plt.subplot(1, 3, 1)\n    ax.imshow(x[time_frame].permute(1,2,0))\n    ax.set_title('False color image')\n\n    ax = plt.subplot(1, 3, 2)\n    ax.imshow(y, interpolation='none')\n    ax.set_title('Ground truth contrail mask')\n\n    ax = plt.subplot(1, 3, 3)\n    ax.imshow(x[time_frame].permute(1,2,0))\n    ax.imshow(y, cmap='Reds', alpha=.4, interpolation='none')\n    ax.set_title('Contrail mask on false color image');\n\n    plt.show()\n\ndef animate_contrail(x):\n    '''\n    x = false color img of shape (8, 3, H, W)\n    '''\n    # Animation\n    fig = plt.figure(figsize=(4, 4))\n    im = plt.imshow(x[0].permute(1,2,0))\n    def draw(i):\n        im.set_array(x[i].permute(1,2,0))\n        return [im]\n    anim = animation.FuncAnimation(\n        fig, draw, frames=x.shape[0], interval=100, blit=True\n    )\n    plt.close()\n    return display.HTML(anim.to_jshtml())","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:42:25.68512Z","iopub.execute_input":"2023-05-11T10:42:25.685909Z","iopub.status.idle":"2023-05-11T10:42:25.698903Z","shell.execute_reply.started":"2023-05-11T10:42:25.685871Z","shell.execute_reply":"2023-05-11T10:42:25.697289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x, y = train_dataset[train_ids.index('1704010292581573769')]\nplot_contrail(x, y)\nanimate_contrail(x)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:42:25.700978Z","iopub.execute_input":"2023-05-11T10:42:25.701948Z","iopub.status.idle":"2023-05-11T10:42:28.044423Z","shell.execute_reply.started":"2023-05-11T10:42:25.701901Z","shell.execute_reply":"2023-05-11T10:42:28.038701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Essential Functions","metadata":{}},{"cell_type":"code","source":"def dice_coeff(mask1, mask2):\n    intersect = torch.sum(mask1 * mask2)\n    fsum = torch.sum(mask1)\n    ssum = torch.sum(mask2)\n    dice = (2 * intersect ) / (fsum + ssum)\n    dice = torch.mean(dice)\n    return dice.item()\n\nmask1 = torch.randint(0,2,(256,256))\nmask2 = torch.randint(0,2,(256,256))\nprint(dice_coeff(mask1, mask2))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model Architecture","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model Train","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model prediction","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Submission\n\n`rle_encode` and `rle_decode` -> [Reference](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":{"execution":{"iopub.status.busy":"2023-05-11T10:42:30.730435Z","iopub.execute_input":"2023-05-11T10:42:30.731409Z","iopub.status.idle":"2023-05-11T10:42:30.743395Z","shell.execute_reply.started":"2023-05-11T10:42:30.731361Z","shell.execute_reply":"2023-05-11T10:42:30.741443Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def submit_to_csv(test_ids, y_pred):\n    submission_df = pd.read_csv(\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/sample_submission.csv\",index_col='record_id')\n    y_encoded = [rle_encode(mask) for mask in y_pred]\n    for record_id, mask in zip(test_ids, y_encoded):\n        submission_df.loc[int(record_id), 'encoded_pixels'] = mask\n    submission_df.to_csv(\"submission.csv\")\n    print(submission_df)\n    print(\"Submitted\")","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:42:30.985273Z","iopub.execute_input":"2023-05-11T10:42:30.985663Z","iopub.status.idle":"2023-05-11T10:42:30.992896Z","shell.execute_reply.started":"2023-05-11T10:42:30.985633Z","shell.execute_reply":"2023-05-11T10:42:30.991646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Example submission\ny_pred = torch.randint(0,2,(2, 256, 256))\nsubmit_to_csv(test_ids, y_pred)","metadata":{"execution":{"iopub.status.busy":"2023-05-11T10:42:31.965666Z","iopub.execute_input":"2023-05-11T10:42:31.966101Z","iopub.status.idle":"2023-05-11T10:42:32.036289Z","shell.execute_reply.started":"2023-05-11T10:42:31.966054Z","shell.execute_reply":"2023-05-11T10:42:32.034853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}