{"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 datetime\nimport os\nfrom tqdm import tqdm\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\nfrom torch.nn import BCELoss, Sigmoid\nfrom torch.optim import Adam\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau \nfrom torchvision.transforms import Normalize\nfrom torch.multiprocessing import Pool","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-13T15:42:41.768429Z","iopub.execute_input":"2023-05-13T15:42:41.768969Z","iopub.status.idle":"2023-05-13T15:42:41.779048Z","shell.execute_reply.started":"2023-05-13T15:42:41.768914Z","shell.execute_reply":"2023-05-13T15:42:41.777805Z"},"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-13T15:42:41.781774Z","iopub.execute_input":"2023-05-13T15:42:41.782689Z","iopub.status.idle":"2023-05-13T15:42:42.205061Z","shell.execute_reply.started":"2023-05-13T15:42:41.782645Z","shell.execute_reply":"2023-05-13T15:42:42.203818Z"},"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, time_frame:int=4):\n    _T11_BOUNDS = (243, 303)\n    _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n    _TDIFF_BOUNDS = (-4, 2)\n    if time_frame is not None:\n        x = x[:,time_frame,:,:]\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) or (3,H,W)","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:42.206898Z","iopub.execute_input":"2023-05-13T15:42:42.207336Z","iopub.status.idle":"2023-05-13T15:42:42.218682Z","shell.execute_reply.started":"2023-05-13T15:42:42.207295Z","shell.execute_reply":"2023-05-13T15:42:42.217336Z"},"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            x_path = os.path.join(self.base_dir, record_id, band)\n            x.append(np.load(x_path).transpose(self.permute))\n        x = np.stack(x,axis=1) ## X.shape = (Time_frame,channel,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            y_path = os.path.join(self.base_dir, record_id,'human_pixel_masks.npy')\n            y = torch.from_numpy(np.load(y_path).transpose(self.permute).astype(np.float32))\n\n            return x, y\n\n    def __len__(self):\n        return len(self.ids)\n        \n    ","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:42.222323Z","iopub.execute_input":"2023-05-13T15:42:42.222782Z","iopub.status.idle":"2023-05-13T15:42:42.239557Z","shell.execute_reply.started":"2023-05-13T15:42:42.222721Z","shell.execute_reply":"2023-05-13T15:42:42.238306Z"},"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\" : 16,\n    \"shuffle\" : True,\n    \"num_workers\": 6,\n    \"drop_last\": True\n#     \"pin_memory\": True\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-13T15:42:42.24105Z","iopub.execute_input":"2023-05-13T15:42:42.241439Z","iopub.status.idle":"2023-05-13T15:42:42.259658Z","shell.execute_reply.started":"2023-05-13T15:42:42.241398Z","shell.execute_reply":"2023-05-13T15:42:42.258355Z"},"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 (1, H, W)\n    time_frame = int, default = 4\n    '''\n    if x.ndim == 4:\n        x = x[time_frame]\n    \n    plt.figure(figsize=(18, 6))\n    ax = plt.subplot(1, 3, 1)\n    ax.imshow(x.permute(1,2,0))\n    ax.set_title('False color image')\n\n    ax = plt.subplot(1, 3, 2)\n    ax.imshow(y.permute(1,2,0), interpolation='none')\n    ax.set_title('Ground truth contrail mask')\n\n    ax = plt.subplot(1, 3, 3)\n    ax.imshow(x.permute(1,2,0))\n    ax.imshow(y.permute(1,2,0), cmap='Reds', alpha=.4, interpolation='none')\n    ax.set_title('Contrail mask on false color image');\n\n    plt.show()\n\ndef plot_contrail_comparision(x, y_true, y_pred, time_frame = 4):\n    '''\n    x = false color img of shape (3, H, W) or (8, 3, H, W)\n    y_true = target contrail mask of shape (1, H, W)\n    y_pred = predicted contrail mask of shape (1, H, W)\n    time_frame = int, default = 4\n    '''\n    if x.ndim == 4:\n        x = x[time_frame]\n    \n    plt.figure(figsize=(18, 6))\n    ax = plt.subplot(1, 5, 1)\n    ax.imshow(x.permute(1,2,0))\n    ax.set_title('False color image(x)')\n    ax.axis('off')\n    \n    ax = plt.subplot(1, 5, 2)\n    ax.imshow(y_true.permute(1,2,0), interpolation='none')\n    ax.set_title('True contrail mask(y_true)')\n    ax.axis('off')\n    \n    ax = plt.subplot(1, 5, 3)\n    ax.imshow(x.permute(1,2,0))\n    ax.imshow(y_true.permute(1,2,0), cmap='Reds', alpha=.4, interpolation='none')\n    ax.set_title('y_true mask on x')\n    ax.axis('off')\n    \n    ax = plt.subplot(1, 5, 4)\n    ax.imshow(y_pred.permute(1,2,0), interpolation='none')\n    ax.set_title('Pred contrail mask(y_pred)')\n    ax.axis('off')\n\n    ax = plt.subplot(1, 5, 5)\n    ax.imshow(x.permute(1,2,0))\n    ax.imshow(y_pred.permute(1,2,0), cmap='Reds', alpha=.4, interpolation='none')\n    ax.set_title('y_pred mask on x')\n    ax.axis('off')\n    \n    plt.show()\n\ndef animate_contrail(x):\n    '''\n    x = false color img of shape (8, 3, H, W)\n    '''\n    if x.ndim !=4:\n        print(f\"Incorrect input dimensions, Expected 4 recievied {x.ndim}.\")\n        return\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-13T15:42:42.261535Z","iopub.execute_input":"2023-05-13T15:42:42.262067Z","iopub.status.idle":"2023-05-13T15:42:42.283285Z","shell.execute_reply.started":"2023-05-13T15:42:42.262032Z","shell.execute_reply":"2023-05-13T15:42:42.282359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# x, y = train_dataset[train_ids.index('1704010292581573769')]\n# plot_contrail(x, y)\n# animate_contrail(x)","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:42.284514Z","iopub.execute_input":"2023-05-13T15:42:42.285244Z","iopub.status.idle":"2023-05-13T15:42:42.301474Z","shell.execute_reply.started":"2023-05-13T15:42:42.28521Z","shell.execute_reply":"2023-05-13T15:42:42.300187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Essential functions\n- Dice coefficient\n- Callbacks to track loss and other params\n- Callbacks to save best state","metadata":{}},{"cell_type":"code","source":"# Dice Coefficient\ndef dice_coeff(mask1, mask2):\n    intersect = torch.sum(mask1 * mask2)\n    m1sum = torch.sum(mask1)\n    m2sum = torch.sum(mask2)\n    dice = (2 * intersect ) / (m1sum + m2sum)\n    return dice.item()\n\n# # Example\n# mask1 = torch.randint(0,2,(10, 256,256))\n# mask2 = torch.randint(0,2,(10, 256,256))\n# print(dice_coeff(mask1, mask2))","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2023-05-13T15:42:42.302427Z","iopub.execute_input":"2023-05-13T15:42:42.302744Z","iopub.status.idle":"2023-05-13T15:42:42.314332Z","shell.execute_reply.started":"2023-05-13T15:42:42.302704Z","shell.execute_reply":"2023-05-13T15:42:42.312963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Callbacks\nclass History:\n    def __init__(self, print_prefix, save_to_disk=True):\n        self.train_batch_history = []\n        self.val_batch_histroy = []\n        self.train_epoch_history = []\n        self.val_epoch_history = []\n        self.running_train_batch_history = []\n        self.running_val_batch_history = []\n        self.print_prefix = print_prefix\n        self.save_to_disk = save_to_disk\n        if save_to_disk:\n            self.save_path = os.path.join(os.getcwd(),\"saved_states\",self.print_prefix)\n            if not os.path.exists(self.save_path):\n                os.makedirs(self.save_path)\n    def on_train_batch_end(self, data):\n        self.running_train_batch_history.append(data)\n        \n    def on_val_batch_end(self, data):\n        self.running_val_batch_history.append(data)\n        \n    def on_epoch_end(self):\n        self.train_epoch_history.append(np.mean(self.running_train_batch_history))\n        self.train_batch_history.extend(self.running_train_batch_history)\n        self.running_train_batch_history=[]\n        self.val_epoch_history.append(np.mean(self.running_val_batch_history))\n        self.val_batch_histroy.extend(self.running_val_batch_history)\n        self.running_val_batch_history=[]\n        print(f\"{self.print_prefix}: Train = {self.train_epoch_history[-1]:.6f} \\\n        | Val = {self.val_epoch_history[-1]:.6f}\")\n    \n    def on_end(self):\n        if self.save_to_disk:\n            dt = datetime.datetime.now().strftime(\"%Y-%m-%d_%H:%M:%S\")\n            np.save(os.path.join(self.save_path,\"train_batch.npy\"),self.train_batch_history)\n            np.save(os.path.join(self.save_path,\"val_batch.npy\"),self.val_batch_histroy)\n            np.save(os.path.join(self.save_path,\"train_epoch.npy\"),self.train_epoch_history)\n            np.save(os.path.join(self.save_path,\"val_epoch.npy\"),self.val_epoch_history)","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2023-05-13T15:42:42.318254Z","iopub.execute_input":"2023-05-13T15:42:42.318874Z","iopub.status.idle":"2023-05-13T15:42:42.335027Z","shell.execute_reply.started":"2023-05-13T15:42:42.318692Z","shell.execute_reply":"2023-05-13T15:42:42.333685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BestStateTracker:\n    def __init__(self, model, optim, trigger:History, save_to_disk:bool = False):\n        self.trigger = trigger\n        self.model = model\n        self.optim = optim\n        self.optim_state = None\n        self.model_state = None\n        self.best_loss = np.inf\n        self.save_to_disk = save_to_disk\n        if save_to_disk:\n            self.save_path = os.path.join(os.getcwd(),\"saved_states\")\n            if not os.path.exists(self.save_path):\n                os.mkdir(self.save_path)\n                \n    def on_epoch_end(self):\n        if self.trigger.val_epoch_history[-1] < self.best_loss:\n            self.model_state = self.model.state_dict()\n            self.optim_state = self.optim.state_dict()\n            \n    def on_end(self):\n        if self.save_to_disk:\n            dt = datetime.datetime.now().strftime(\"%Y-%m-%d_%H:%M:%S\")\n            torch.save(self.model_state,os.path.join(self.save_path,f\"model_state_{dt}.pt\"))\n            torch.save(self.optim_state,os.path.join(self.save_path,f\"optim_state_{dt}.pt\"))","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2023-05-13T15:42:42.336312Z","iopub.execute_input":"2023-05-13T15:42:42.336641Z","iopub.status.idle":"2023-05-13T15:42:42.350685Z","shell.execute_reply.started":"2023-05-13T15:42:42.336604Z","shell.execute_reply":"2023-05-13T15:42:42.349819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model Architecture","metadata":{}},{"cell_type":"code","source":"# Config\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nLR = 1e-2\nEPOCHS = 100","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:42.351817Z","iopub.execute_input":"2023-05-13T15:42:42.352813Z","iopub.status.idle":"2023-05-13T15:42:42.367794Z","shell.execute_reply.started":"2023-05-13T15:42:42.352771Z","shell.execute_reply":"2023-05-13T15:42:42.366789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = torch.hub.load('mateuszbuda/brain-segmentation-pytorch', 'unet',\n                       in_channels=3, out_channels=1, init_features=32, \n                       pretrained=False).to(DEVICE)\noptimizer = Adam(params=model.parameters(), lr=LR)\nloss_fn = BCELoss()\nscheduler = ReduceLROnPlateau(optimizer)\nloss_tracker = History(print_prefix=\"Loss\")\ndice_tracker = History(print_prefix=\"Dice\")\nsave_state = BestStateTracker(model,optimizer,loss_tracker,save_to_disk=True)","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:42.369001Z","iopub.execute_input":"2023-05-13T15:42:42.370012Z","iopub.status.idle":"2023-05-13T15:42:46.451231Z","shell.execute_reply.started":"2023-05-13T15:42:42.369977Z","shell.execute_reply":"2023-05-13T15:42:46.450386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model Train","metadata":{}},{"cell_type":"code","source":"# pbar = tqdm(range(EPOCHS))\n\n# for epoch in pbar:\n#     print(f\"\\nEPOCH: {epoch+1}/{EPOCHS}\")\n#     model.train()\n#     for idx, (data, target) in enumerate(train_loader):\n#         data, target = data.to(DEVICE), target.to(DEVICE) #Copy to GPU if available\n#         y_pred = model(data)\n#         loss = loss_fn(y_pred, target)\n#         optimizer.zero_grad()\n#         loss.backward()\n#         optimizer.step()\n        \n#         loss_tracker.on_train_batch_end(loss.item())\n#         dice_tracker.on_train_batch_end(dice_coeff(y_pred>0.5, target))\n#         pbar.set_description(f\"Train Batch: {idx+1}/{len(train_loader)}\\\n#         | Loss: {loss.item():.6f}\")\n    \n#     plot_contrail_comparision(data[0].cpu().detach(),\n#                               target[0].cpu().detach(),\n#                               y_pred[0].cpu().detach()>0.5)\n#     model.eval()    \n#     with torch.no_grad():\n#         for idx, (data, target) in enumerate(val_loader):\n#             data, target = data.to(DEVICE), target.to(DEVICE) #Copy to GPU if available\n#             y_pred = model(data)\n#             loss = loss_fn(y_pred, target)\n#             loss_tracker.on_val_batch_end(loss.item())\n#             dice_tracker.on_val_batch_end(dice_coeff(y_pred>0.5, target))\n#             pbar.set_description(f\"Val Batch:   {idx+1}/{len(val_loader)}\\\n#             | Loss: {loss.item():.6f}\")\n#         plot_contrail_comparision(data[0].cpu().detach(),\n#                                   target[0].cpu().detach(),\n#                                   y_pred[0].cpu().detach()>0.5)\n#     print(\"Train sample(above), Val sample(below)\")\n#     scheduler.step(loss.item())\n#     loss_tracker.on_epoch_end()\n#     dice_tracker.on_epoch_end()\n#     save_state.on_epoch_end()\n\n# save_state.on_end()\n# loss_tracker.on_end()\n# dice_tracker.on_end()","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:46.452513Z","iopub.execute_input":"2023-05-13T15:42:46.453422Z","iopub.status.idle":"2023-05-13T15:42:46.459706Z","shell.execute_reply.started":"2023-05-13T15:42:46.453384Z","shell.execute_reply":"2023-05-13T15:42:46.458503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load pretrained states:\nmodel.load_state_dict(torch.load(\"/kaggle/input/unet-trained-weights-contrails/contrails_save_states/saved_states/model_state_2023-05-12_06_31_20.pt\",map_location=torch.device(DEVICE)))\noptimizer.load_state_dict(torch.load(\"/kaggle/input/unet-trained-weights-contrails/contrails_save_states/saved_states/optim_state_2023-05-12_06_31_20.pt\",map_location=torch.device(DEVICE)))\nloss_tracker.train_batch_history=np.load(\"/kaggle/input/unet-trained-weights-contrails/contrails_save_states/saved_states/Loss/train_batch.npy\")\nloss_tracker.train_epoch_history=np.load(\"/kaggle/input/unet-trained-weights-contrails/contrails_save_states/saved_states/Loss/train_epoch.npy\")\nloss_tracker.val_batch_histroy=np.load(\"/kaggle/input/unet-trained-weights-contrails/contrails_save_states/saved_states/Loss/val_batch.npy\")\nloss_tracker.val_epoch_history=np.load(\"/kaggle/input/unet-trained-weights-contrails/contrails_save_states/saved_states/Loss/val_epoch.npy\")\ndice_tracker.train_batch_history=np.load(\"/kaggle/input/unet-trained-weights-contrails/contrails_save_states/saved_states/Dice/train_batch.npy\")\ndice_tracker.train_epoch_history=np.load(\"/kaggle/input/unet-trained-weights-contrails/contrails_save_states/saved_states/Dice/train_epoch.npy\")\ndice_tracker.val_batch_histroy=np.load(\"/kaggle/input/unet-trained-weights-contrails/contrails_save_states/saved_states/Dice/val_batch.npy\")\ndice_tracker.val_epoch_history=np.load(\"/kaggle/input/unet-trained-weights-contrails/contrails_save_states/saved_states/Dice/val_epoch.npy\")\n","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:46.461585Z","iopub.execute_input":"2023-05-13T15:42:46.462042Z","iopub.status.idle":"2023-05-13T15:42:47.636392Z","shell.execute_reply.started":"2023-05-13T15:42:46.461988Z","shell.execute_reply":"2023-05-13T15:42:47.635272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,10))\nax = plt.subplot(2,2,1)\nax.plot(loss_tracker.train_batch_history[len(train_loader):], label=\"train batch loss\")\nax.legend()\nax = plt.subplot(2,2,2)\nax.plot(loss_tracker.val_batch_histroy[len(val_loader):], label=\"val batch loss\")\nax.legend()\nax = plt.subplot(2,2,3)\nax.plot(dice_tracker.train_batch_history[len(train_loader):], label=\"train batch dice coeff\")\nax.legend()\nax = plt.subplot(2,2,4)\nax.plot(dice_tracker.val_batch_histroy[len(val_loader):], label=\"val batch dice coeff\")\nax.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:47.638756Z","iopub.execute_input":"2023-05-13T15:42:47.639096Z","iopub.status.idle":"2023-05-13T15:42:48.773363Z","shell.execute_reply.started":"2023-05-13T15:42:47.639067Z","shell.execute_reply":"2023-05-13T15:42:48.772201Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(loss_tracker.train_epoch_history, label=\"train epoch loss\")\nplt.plot(loss_tracker.val_epoch_history, label=\"val epoch loss\")\nplt.legend()\nplt.title(\"Epoch Loss\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:48.77494Z","iopub.execute_input":"2023-05-13T15:42:48.775409Z","iopub.status.idle":"2023-05-13T15:42:49.066829Z","shell.execute_reply.started":"2023-05-13T15:42:48.775364Z","shell.execute_reply":"2023-05-13T15:42:49.065527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(dice_tracker.train_epoch_history, label=\"train epoch dice coeff\")\nplt.plot(dice_tracker.val_epoch_history, label=\"val epoch dice coeff\")\nplt.legend()\nplt.title(\"Epoch Dice Coeff\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:49.068373Z","iopub.execute_input":"2023-05-13T15:42:49.068824Z","iopub.status.idle":"2023-05-13T15:42:49.334983Z","shell.execute_reply.started":"2023-05-13T15:42:49.068775Z","shell.execute_reply":"2023-05-13T15:42:49.333793Z"},"trusted":true},"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-13T15:42:49.33637Z","iopub.execute_input":"2023-05-13T15:42:49.337318Z","iopub.status.idle":"2023-05-13T15:42:49.349104Z","shell.execute_reply.started":"2023-05-13T15:42:49.337285Z","shell.execute_reply":"2023-05-13T15:42:49.347789Z"},"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(base_dir, \"sample_submission.csv\"), \n                                index_col='record_id')\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)\n    submission_df.to_csv(\"submission.csv\")\n    print(\"Submitted\")","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:49.350461Z","iopub.execute_input":"2023-05-13T15:42:49.350836Z","iopub.status.idle":"2023-05-13T15:42:49.366492Z","shell.execute_reply.started":"2023-05-13T15:42:49.350805Z","shell.execute_reply":"2023-05-13T15:42:49.365198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Example submission\n# y_pred = torch.randint(0,2,(2, 256, 256))\n# submit_to_csv(test_ids, y_pred)","metadata":{"execution":{"iopub.status.busy":"2023-05-13T15:42:49.367958Z","iopub.execute_input":"2023-05-13T15:42:49.368939Z","iopub.status.idle":"2023-05-13T15:42:49.379205Z","shell.execute_reply.started":"2023-05-13T15:42:49.368905Z","shell.execute_reply":"2023-05-13T15:42:49.378348Z"},"trusted":true},"execution_count":null,"outputs":[]}]}