{"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":"## Getting Started | EDA -> Model -> Train -> Submit\n# JAKUB NOWACKI\n# https://www.kaggle.com/code/mnokno/getting-started-eda-model-train-submit/commentshttps://www.kaggle.com/code/mnokno/getting-started-eda-model-train-submit/comments","metadata":{}},{"cell_type":"markdown","source":"Acceleratator = GPU T4X2","metadata":{}},{"cell_type":"code","source":"import sys\nprint(sys.version)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:31.869059Z","iopub.execute_input":"2023-07-30T03:21:31.869416Z","iopub.status.idle":"2023-07-30T03:21:31.876226Z","shell.execute_reply.started":"2023-07-30T03:21:31.869388Z","shell.execute_reply":"2023-07-30T03:21:31.87493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install torchsummary","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:32.210365Z","iopub.execute_input":"2023-07-30T03:21:32.210693Z","iopub.status.idle":"2023-07-30T03:21:32.214807Z","shell.execute_reply.started":"2023-07-30T03:21:32.210667Z","shell.execute_reply":"2023-07-30T03:21:32.213875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install numpy == 1.16.5","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:32.345045Z","iopub.execute_input":"2023-07-30T03:21:32.345663Z","iopub.status.idle":"2023-07-30T03:21:32.350162Z","shell.execute_reply.started":"2023-07-30T03:21:32.345629Z","shell.execute_reply":"2023-07-30T03:21:32.349163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport os\n\nfrom matplotlib import animation\nfrom IPython import display\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch import Tensor\nfrom torch.utils.data import TensorDataset\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import DataLoader\n# from torchsummary import summary","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:32.481515Z","iopub.execute_input":"2023-07-30T03:21:32.481783Z","iopub.status.idle":"2023-07-30T03:21:32.487507Z","shell.execute_reply.started":"2023-07-30T03:21:32.481761Z","shell.execute_reply":"2023-07-30T03:21:32.486582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_dir: str = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:32.664293Z","iopub.execute_input":"2023-07-30T03:21:32.665257Z","iopub.status.idle":"2023-07-30T03:21:32.669656Z","shell.execute_reply.started":"2023-07-30T03:21:32.665221Z","shell.execute_reply":"2023-07-30T03:21:32.668715Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_idx = pd.DataFrame({'idx': os.listdir('/kaggle/input/google-research-identify-contrails-reduce-global-warming/train')})\ndf_validation_idx = pd.DataFrame({'idx': os.listdir('/kaggle/input/google-research-identify-contrails-reduce-global-warming/validation')})\ndf_test_idx = pd.DataFrame({'idx': os.listdir('/kaggle/input/google-research-identify-contrails-reduce-global-warming/test')})","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:32.817995Z","iopub.execute_input":"2023-07-30T03:21:32.818384Z","iopub.status.idle":"2023-07-30T03:21:33.106497Z","shell.execute_reply.started":"2023-07-30T03:21:32.818357Z","shell.execute_reply":"2023-07-30T03:21:33.105342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_idx.tail()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:33.108387Z","iopub.execute_input":"2023-07-30T03:21:33.10937Z","iopub.status.idle":"2023-07-30T03:21:33.130376Z","shell.execute_reply.started":"2023-07-30T03:21:33.109335Z","shell.execute_reply":"2023-07-30T03:21:33.129501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_band_images(idx: str, parrent_folder: str, band: str) -> np.array:\n    return np.load(os.path.join(data_dir, parrent_folder, idx, f'band_{band}.npy'))","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:33.183884Z","iopub.execute_input":"2023-07-30T03:21:33.184165Z","iopub.status.idle":"2023-07-30T03:21:33.188968Z","shell.execute_reply.started":"2023-07-30T03:21:33.184141Z","shell.execute_reply":"2023-07-30T03:21:33.187868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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])\n\ndef get_ash_color_images(idx: str, parrent_folder: str, get_mask_frame_only=False) -> np.array:\n    band11 = get_band_images(idx, parrent_folder, '11')\n    band14 = get_band_images(idx, parrent_folder, '14')\n    band15 = get_band_images(idx, parrent_folder, '15')\n    \n    if get_mask_frame_only:\n        band11 = band11[:, :, 4]\n        band14 = band14[:, : ,4]\n        band15 = band15[:, : ,4]\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 = np.clip(np.stack([r, g, b], axis=2), 0, 1)\n    return false_color","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:33.365673Z","iopub.execute_input":"2023-07-30T03:21:33.36634Z","iopub.status.idle":"2023-07-30T03:21:33.375306Z","shell.execute_reply.started":"2023-07-30T03:21:33.366309Z","shell.execute_reply":"2023-07-30T03:21:33.374299Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_mask_image(idx: str, parrent_folder: str) -> np.array:\n    return np.load(os.path.join(data_dir, parrent_folder, idx, 'human_pixel_masks.npy'))","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:33.543518Z","iopub.execute_input":"2023-07-30T03:21:33.544362Z","iopub.status.idle":"2023-07-30T03:21:33.549553Z","shell.execute_reply.started":"2023-07-30T03:21:33.544324Z","shell.execute_reply":"2023-07-30T03:21:33.548545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## EDA","metadata":{}},{"cell_type":"code","source":"train_images_with_contrails = 0\ntrain_images_without_contrails = 0\ntrain_contrail_pixel_count = 0\ntrain_non_contrail_pixel_count = 0\ntrain_contrail_pixel_count_conly = 0\ntrain_non_contrail_pixel_count_conly = 0\nimg_pixel_count = 256 * 256\nreal_data_train_idx = []\n\nfor idx in df_train_idx['idx']:\n    mask = get_mask_image(idx, 'train')\n    contrail_pixel_count = np.sum(mask > 0)\n    \n    if contrail_pixel_count > 0:\n        train_images_with_contrails += 1\n        train_contrail_pixel_count_conly += contrail_pixel_count\n        train_non_contrail_pixel_count_conly += (img_pixel_count - contrail_pixel_count)\n        real_data_train_idx.append(idx)\n        \n    else:\n        train_images_without_contrails += 1\n        \n    train_contrail_pixel_count += contrail_pixel_count\n    train_non_contrail_pixel_count += (img_pixel_count - contrail_pixel_count)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:21:33.992477Z","iopub.execute_input":"2023-07-30T03:21:33.99359Z","iopub.status.idle":"2023-07-30T03:24:05.914714Z","shell.execute_reply.started":"2023-07-30T03:21:33.993546Z","shell.execute_reply":"2023-07-30T03:24:05.913697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"validation_images_with_contrails = 0\nvalidation_images_without_contrails = 0\nvalidation_contrail_pixel_count = 0\nvalidation_non_contrail_pixel_count = 0\nvalidation_contrail_pixel_count_conly = 0\nvalidation_non_contrail_pixel_count_conly = 0\nimg_pixel_count = 256 * 256\nreal_data_valid_idx = []\n\nfor idx in df_validation_idx['idx']:\n    mask = get_mask_image(idx, 'validation')\n    contrail_pixel_count = np.sum(mask > 0)\n    \n    if contrail_pixel_count > 0:\n        validation_images_with_contrails += 1\n        validation_contrail_pixel_count_conly += contrail_pixel_count\n        validation_non_contrail_pixel_count_conly += (img_pixel_count - contrail_pixel_count)\n        real_data_valid_idx.append(idx)\n        \n    else:\n        validation_images_without_contrails += 1\n        \n    validation_contrail_pixel_count += contrail_pixel_count\n    validation_non_contrail_pixel_count += (img_pixel_count - contrail_pixel_count)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:05.916799Z","iopub.execute_input":"2023-07-30T03:24:05.917397Z","iopub.status.idle":"2023-07-30T03:24:19.883289Z","shell.execute_reply.started":"2023-07-30T03:24:05.917362Z","shell.execute_reply":"2023-07-30T03:24:19.88233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask.shape","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:19.884525Z","iopub.execute_input":"2023-07-30T03:24:19.885737Z","iopub.status.idle":"2023-07-30T03:24:19.892785Z","shell.execute_reply.started":"2023-07-30T03:24:19.885705Z","shell.execute_reply":"2023-07-30T03:24:19.891674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(real_data_train_idx)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:19.895853Z","iopub.execute_input":"2023-07-30T03:24:19.896336Z","iopub.status.idle":"2023-07-30T03:24:19.903739Z","shell.execute_reply.started":"2023-07-30T03:24:19.896213Z","shell.execute_reply":"2023-07-30T03:24:19.90276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(real_data_valid_idx)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:19.905359Z","iopub.execute_input":"2023-07-30T03:24:19.905816Z","iopub.status.idle":"2023-07-30T03:24:19.914391Z","shell.execute_reply.started":"2023-07-30T03:24:19.905785Z","shell.execute_reply":"2023-07-30T03:24:19.913239Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"validation_contrail_pixel_count:\",validation_contrail_pixel_count)\nprint(\"validation_non_contrail_pixel_count:\",validation_non_contrail_pixel_count)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:19.91603Z","iopub.execute_input":"2023-07-30T03:24:19.91648Z","iopub.status.idle":"2023-07-30T03:24:19.9248Z","shell.execute_reply.started":"2023-07-30T03:24:19.916451Z","shell.execute_reply":"2023-07-30T03:24:19.923842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_with_contrails = train_images_with_contrails / (train_images_with_contrails + train_images_without_contrails)\ntrain_without_contrails = train_images_without_contrails / (train_images_with_contrails + train_images_without_contrails)\nvalidation_with_contrails = validation_images_with_contrails / (validation_images_with_contrails + validation_images_without_contrails)\nvalidation_without_contrails = validation_images_without_contrails / (validation_images_with_contrails + validation_images_without_contrails)\ndata = pd.DataFrame({'Type': ['With Contrails', 'No Contrails', 'With Contrails', 'No Contrails'],\n        'Data': [train_with_contrails, train_without_contrails, validation_with_contrails, validation_without_contrails],\n        'Data Set': ['train', 'train', 'validation', 'validation']})\n\nax = sns.barplot(data=data, y='Data', x=\"Type\", hue=\"Data Set\", orient='v')\n\nfor p in ax.patches:\n    ax.annotate(format(p.get_height() * 100, '.0f') + '%',\n                (p.get_x() + p.get_width() / 2., p.get_height()),\n                ha = 'center', va = 'center',\n                xytext = (0, 5),\n                textcoords = 'offset points')\n    \nax.set_xlabel('')\nax.set_ylabel('Percentage of Dataset')\nax.set_title('With Contrails vs No Contrails')\n\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:19.927281Z","iopub.execute_input":"2023-07-30T03:24:19.927547Z","iopub.status.idle":"2023-07-30T03:24:20.248671Z","shell.execute_reply.started":"2023-07-30T03:24:19.927526Z","shell.execute_reply":"2023-07-30T03:24:20.247795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(train_with_contrails)\nprint(train_without_contrails)\nprint(validation_with_contrails)\nprint(validation_without_contrails)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:20.250234Z","iopub.execute_input":"2023-07-30T03:24:20.250569Z","iopub.status.idle":"2023-07-30T03:24:20.256015Z","shell.execute_reply.started":"2023-07-30T03:24:20.250537Z","shell.execute_reply":"2023-07-30T03:24:20.255123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:20.257467Z","iopub.execute_input":"2023-07-30T03:24:20.258109Z","iopub.status.idle":"2023-07-30T03:24:20.271422Z","shell.execute_reply.started":"2023-07-30T03:24:20.258079Z","shell.execute_reply":"2023-07-30T03:24:20.270391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(nrows=1, ncols=2, figsize=(12,6))\nfig.subplots_adjust(wspace=0.3)\naxes = axes.flatten()\n\ntrain_with_contrails_pix = train_contrail_pixel_count / (train_contrail_pixel_count + train_non_contrail_pixel_count)\nvalidation_with_contrails_pix = validation_contrail_pixel_count / (validation_contrail_pixel_count + validation_non_contrail_pixel_count)\ndata = pd.DataFrame({'Data': [train_with_contrails_pix, validation_with_contrails_pix], \n                    'Data Set': ['train', 'validation']})\n\nsns.barplot(data=data, y='Data', x='Data Set', orient='v', ax=axes[0])\nfor p in axes[0].patches:\n    axes[0].annotate(format(p.get_height(), '.4f'),\n                    (p.get_x() + p.get_width() / 2., p.get_height()),\n                    ha = 'center', va = 'center',\n                    xytext = (0, 5),\n                    textcoords = 'offset points')\n    \naxes[0].set_xlabel('')\naxes[0].set_ylabel('Percentage of Contrails Pixels')\naxes[0].set_title('Percentage of Contrails pixels in Images')\n\n####################################################\n\ntrain_with_contrails_pix_conly = train_contrail_pixel_count_conly / (train_contrail_pixel_count_conly + train_non_contrail_pixel_count_conly)\nvalidation_with_contrails_pix_conly = validation_contrail_pixel_count_conly / (validation_contrail_pixel_count_conly + validation_non_contrail_pixel_count_conly)\ndata = pd.DataFrame({'Data': [train_with_contrails_pix_conly, validation_with_contrails_pix_conly],\n                    'Data Set': ['train', 'validation']})\nsns.barplot(data=data, y='Data', x='Data Set', orient='v', ax=axes[1])\nfor p in axes[1].patches:\n    axes[1].annotate(format(p.get_height(), '.4f'),\n                    (p.get_x() + p.get_width() / 2., p.get_height()),\n                    ha = 'center', va = 'center',\n                    xytext = (0, 5),\n                    textcoords = 'offset points')\n    \naxes[1].set_xlabel('')\naxes[1].set_ylabel('Percentage of Contrails Pixels')\naxes[1].set_title('Percentage of Contrails pixels in Images with Contrails Percent')\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:20.276452Z","iopub.execute_input":"2023-07-30T03:24:20.276714Z","iopub.status.idle":"2023-07-30T03:24:20.689069Z","shell.execute_reply.started":"2023-07-30T03:24:20.276692Z","shell.execute_reply":"2023-07-30T03:24:20.688124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"There is a significant imbalance in the number of pixels between the negative class (non-contrail) and the positive class (contrail). In the training and validation datasets, the ratio of negative to positive class pixels is 85:1 and 164:1 respectively. This severe class imbalance can lead to biased models that prioritize the majority class, ultimately reducing the overall prediction quality.\n\nTo address this issue, we can employ two strategies: optimizing the confidence threshold during post-processing or incorporating class weights into the loss function during training.\n\nOptimizing confidence threshold during post-processing: After training the model, we can adjust the confidence threshold used to determine class predictions. By carefully selecting the threshold, we can increase the sensitivity to the positive class, thereby improving the detection of contrails. This approach allows us to fine-tune the model's predictions without retraining it.\n\nAdding class weights to the loss function during training: Another way to handle the class imbalance is by assigning appropriate weights to the different classes during model training. By assigning higher weights to the minority class (contrail), we can increase its influence on the loss function. This adjustment ensures that the model pays more attention to the positive class and helps mitigate the bias towards the majority class (non-contrail).\n\nThis notebook will utilize both strategies.\n\nAdditionally, I thought it was worth pointing out that there are significantly lass contrails in the validation images than in the train images.","metadata":{}},{"cell_type":"code","source":"def show_band_images(idx: str, parrent_folder: str, band: str):\n    data = get_band_images(idx, parrent_folder, band)\n    fig, axes = plt.subplots(nrows=2, nocles=4, figsize=(20,10))\n    axes = axes.flatten()\n    for i in range(8):\n        axes[i].imshow(data[:,:,i])\n        axes[i].axis('off')\n    plt.show()\n    \ndef show_ash_images(idx: str, parrent_folder: str):\n    data = get_ash_color_images(idx, parrent_folder)\n    fig, axes = plt.subplots(nrows=2, ncols=4, figsize=(20, 10))\n    axes = axes.flatten()\n    for i in range(8):\n        axes[i].imshow(data[:, :, :, i])\n        axes[i].axis('off')\n    plt.show()\n    \ndef show_ash_frame(idx: str, parrent_folder: str, frame: int):\n    data = get_ash_color_images(idx, parrent_folder)\n    plt.imshow(data[:,:,:,frame])\n    plt.show()\n    \ndef show_mask_image(idx: str, parrent_folder: str):\n    plt.imshow(get_mask_image(idx, parrent_folder))\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:20.692063Z","iopub.execute_input":"2023-07-30T03:24:20.69235Z","iopub.status.idle":"2023-07-30T03:24:20.70264Z","shell.execute_reply.started":"2023-07-30T03:24:20.692326Z","shell.execute_reply":"2023-07-30T03:24:20.701612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(nrows=4, ncols=4, figsize=(40, 40))\naxes = axes.flatten()\n\nfor i in range(len(axes)):\n    images = get_ash_color_images(str(df_train_idx.iloc[683 + i]['idx']), 'train')\n    axes[i].imshow(images[:,:,:,4])\n    axes[i].axis('off')","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:20.70389Z","iopub.execute_input":"2023-07-30T03:24:20.704423Z","iopub.status.idle":"2023-07-30T03:24:29.492744Z","shell.execute_reply.started":"2023-07-30T03:24:20.704391Z","shell.execute_reply":"2023-07-30T03:24:29.491451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images = get_ash_color_images(str(df_train_idx.iloc[683]['idx']), 'train')\nfig, axes = plt.subplots(nrows=1, ncols=2, figsize=(10, 5))\naxes = axes.flatten()\n\naxes[0].imshow(images[:,:,:,4])\naxes[0].axis('off')\naxes[1].imshow(get_mask_image(str(df_train_idx.iloc[683]['idx']), 'train'))\naxes[1].axis('off')\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:29.494102Z","iopub.execute_input":"2023-07-30T03:24:29.494581Z","iopub.status.idle":"2023-07-30T03:24:29.83165Z","shell.execute_reply.started":"2023-07-30T03:24:29.494536Z","shell.execute_reply":"2023-07-30T03:24:29.830775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"false_color = get_ash_color_images(str(df_train_idx.iloc[683]['idx']), 'train')\n\nfig, ax = plt.subplots(figsize=(6, 6))\nax.set_axis_off()\nim = plt.imshow(false_color[..., 0])\ndef draw(i):\n    im.set_array(false_color[..., i])\n    return [im]\n\nanim = animation.FuncAnimation(\n    fig, draw, frames=false_color.shape[-1], interval = 500, blit=True)\nplt.close()\ndisplay.HTML(anim.to_jshtml())","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:29.83309Z","iopub.execute_input":"2023-07-30T03:24:29.8337Z","iopub.status.idle":"2023-07-30T03:24:31.476083Z","shell.execute_reply.started":"2023-07-30T03:24:29.833667Z","shell.execute_reply":"2023-07-30T03:24:31.475168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ndevice","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:31.477369Z","iopub.execute_input":"2023-07-30T03:24:31.482403Z","iopub.status.idle":"2023-07-30T03:24:31.516857Z","shell.execute_reply.started":"2023-07-30T03:24:31.482372Z","shell.execute_reply":"2023-07-30T03:24:31.51591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DoubleConv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(DoubleConv, self).__init__()\n        self.double_conv = nn.Sequential(\n        nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n        nn.BatchNorm2d(out_channels),\n        nn.ReLU(inplace=True),\n        nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n        nn.BatchNorm2d(out_channels),\n        nn.ReLU(inplace=True)\n        )\n        \n    def forward(self, x):\n        return self.double_conv(x)\n    \n    \nclass Down(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(Down, self).__init__()\n        self.maxpool_conv = nn.Sequential(\n            nn.MaxPool2d(2),\n            DoubleConv(in_channels, out_channels)\n        )\n        \n        \n    def forward(self, x):\n        return self.maxpool_conv(x)\n    \nclass Up(nn.Module):\n    def __init__(self, in_channels, out_channels, bilinear=True):\n        super(Up, self).__init__()\n        \n        if bilinear:\n            self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n        else:\n            self.up = nn.ConvTranspose2d(in_channels // 2, in_channles // 2, kernel_size=2, stride = 2)\n        \n        self.conv = DoubleConv(in_channels, out_channels)\n        \n    def forward(self, x1, x2):\n        x1 = self.up(x1)\n        \n        diffY = x2.size()[2] - x1.size()[2]\n        diffX = x2.size()[3] - x1.size()[3]\n        \n        x1 = nn.functional.pad(x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2])\n        \n        x = torch.cat([x2, x1], dim=1)\n        return self.conv(x)\n    \nclass UNet(nn.Module):\n    def __init__(self):\n        super(UNet, self).__init__()\n        # Define layers\n        self.inc = DoubleConv(24, 64)\n        self.down1 = Down(64, 128)\n        self.down2 = Down(128, 256)\n        self.down3 = Down(256, 512)\n        self.down4 = Down(512, 512)\n        self.up1 = Up(1024, 256)\n        self.up2 = Up(512, 128)\n        self.up3 = Up(256, 64)\n        self.up4 = Up(128, 64)\n        self.outc = nn.Conv2d(64, 1, kernel_size=1)\n        \n    def forward(self, x):\n        # Forward pass through the layers\n        x1 = self.inc(x)\n        x2 = self.down1(x1)\n        x3 = self.down2(x2)\n        x4 = self.down3(x3)\n        x5 = self.down4(x4)\n        x = self.up1(x5, x4)\n        x = self.up2(x, x3)\n        x = self.up3(x, x2)\n        x = self.up4(x, x1)\n        x = self.outc(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:31.518691Z","iopub.execute_input":"2023-07-30T03:24:31.519448Z","iopub.status.idle":"2023-07-30T03:24:31.541612Z","shell.execute_reply.started":"2023-07-30T03:24:31.519297Z","shell.execute_reply":"2023-07-30T03:24:31.540694Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# summary(UNet().to(device), input_size=(24, 256, 256))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Trainer","metadata":{}},{"cell_type":"code","source":"class Dice(nn.Module):\n    def __init__(self, use_sigmoid=True):\n        super(Dice, self).__init__()\n        self.sigmoid = nn.Sigmoid()\n        self.use_sigmoid = use_sigmoid\n        \n    def forward(self, inputs, targets, smooth=1):\n        if self.use_sigmoid:\n            inputs = self.sigmoid(inputs)\n            \n        inputs = inputs.view(-1)\n        targets = targets.view(-1)\n        \n        intersection = (inputs * targets).sum()\n        dice = (2.0 * intersection + smooth)/(inputs.sum() + targets.sum() + smooth)\n        \n        return dice\n    \ndice = Dice()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:31.543068Z","iopub.execute_input":"2023-07-30T03:24:31.543727Z","iopub.status.idle":"2023-07-30T03:24:31.558774Z","shell.execute_reply.started":"2023-07-30T03:24:31.543697Z","shell.execute_reply":"2023-07-30T03:24:31.557932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MyTrainer:\n    def __init__(self, model, optimizer, loss_fn, lr_scheduler):\n        self.validation_losses = []\n        self.batch_losses = []\n        self.epoch_losses = []\n        self.learning_rates = []\n        self.model = model\n        self.optimizer = optimizer\n        self.loss_fn = loss_fn\n        self.lr_scheduler = lr_scheduler\n        self._check_optim_net_aligned()\n        \n    # Ensures that the given optimizer points to the given model\n    def _check_optim_net_aligned(self):\n        assert self.optimizer.param_groups[0]['params'] == list(self.model.parameters())\n        \n    # Trains the model\n    def fit(self,\n           train_dataloader: DataLoader,\n           test_dataloader: DataLoader,\n           epochs: int = 10,\n           eval_every: int = 1,\n           ):\n        \n        for e in range(epochs):\n            print(\"New learning rate: {}\".format(self.lr_scheduler.get_last_lr()))\n            self.learning_rates.append(self.lr_scheduler.get_last_lr()[0])\n            \n            # Stores data about the batch\n            batch_losses = []\n            sub_batch_losses = []\n            \n            for i, data in enumerate(train_dataloader):\n                self.model.train()\n                if i % 100 == 0:\n                    print(f'epoch: {e} batch: {i}/{len(train_dataloader)} loss: {torch.Tensor(sub_batch_losses).mean()}')\n                    sub_batch_losses.clear()\n                # Every data instance is an input + label pair\n                images, mask = data\n                \n                if torch.cuda.is_available():\n                    images = images.cuda()\n                    mask = mask.cuda()\n                    \n                # Zero your gradients for every batch\n                self.optimizer.zero_grad()\n                # Make predictions for this batch\n                outputs = self.model(images)\n                # Compute the loss and its gradients\n                loss = self.loss_fn(outputs, mask)\n                loss.backward()\n                # Adjust learning weights\n                self.optimizer.step()\n                \n                # Saves data\n                self.batch_losses.append(loss.item())\n                batch_losses.append(loss)\n                sub_batch_losses.append(loss)\n                \n            # Adjusts learning rate\n            if self.lr_scheduler is not None:\n                self.lr_scheduler.step()\n                \n            # Reports on the path\n            mean_epoch_loss = torch.Tensor(batch_losses).mean()\n            self.epoch_losses.append(mean_epoch_loss.item())\n            print('Train Epoch: {} Average Loss: {:.6f}'.format(e, mean_epoch_loss))\n            \n            # Reports on the training progress\n            if (e + 1) % eval_every == 0:\n                torch.save(self.model.state_dict(), \"model_checkpoint_e\" + str(e) + \".pt\")\n                with torch.no_grad():\n                    self.model.eval()\n                    losses = []\n                    for i, data in enumerate(test_dataloader):\n                        # Every data instance is an input + label pair\n                        images, mask = data\n                        \n                        if torch.cuda.is_available():\n                            images = images.cuda()\n                            mask = mask.cuda()\n                            \n                        output = self.model(images)\n                        loss = self.loss_fn(output, mask)\n                        losses.append(loss.item())\n                        \n                    avg_loss = torch.Tensor(losses).mean().item()\n                    self.validation_losses.append(avg_loss)\n                    print(\"Validation loss after\",  (e + 1), \"epochs was\", round(avg_loss, 4))\n                        ","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:31.560512Z","iopub.execute_input":"2023-07-30T03:24:31.561223Z","iopub.status.idle":"2023-07-30T03:24:31.583567Z","shell.execute_reply.started":"2023-07-30T03:24:31.561169Z","shell.execute_reply":"2023-07-30T03:24:31.582631Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class ContrailsAshDataset(torch.utils.data.Dataset):\n    def __init__(self, parrent_folder: str):\n        self.df_idx: pd.DataFrame = pd.DataFrame({'idx': os.listdir(f'/kaggle/input/google-research-identify-contrails-reduce-global-warming/{parrent_folder}')})\n        self.parrent_folder: str = parrent_folder\n            \n    def __len__(self):\n        return len(self.df_idx)\n    \n    def __getitem__(self, idx):\n        image_id: str = str(self.df_idx.iloc[idx]['idx'])\n        images = torch.tensor(np.reshape(get_ash_color_images(image_id, self.parrent_folder, get_mask_frame_only=False), (256, 256, 24))).to(torch.float32).permute(2, 0, 1)\n        mask = torch.tensor(get_mask_image(image_id, self.parrent_folder)).to(torch.float32).permute(2, 0, 1)\n        return images, mask","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:31.585154Z","iopub.execute_input":"2023-07-30T03:24:31.585888Z","iopub.status.idle":"2023-07-30T03:24:31.596778Z","shell.execute_reply.started":"2023-07-30T03:24:31.585855Z","shell.execute_reply":"2023-07-30T03:24:31.595858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_train = ContrailsAshDataset('train')\ndataset_validation = ContrailsAshDataset('validation')\n\ndata_loader_train = DataLoader(dataset_train, batch_size=16, shuffle=True, num_workers=2)\ndata_loader_validation = DataLoader(dataset_validation, batch_size=16, shuffle=True, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:31.598444Z","iopub.execute_input":"2023-07-30T03:24:31.599196Z","iopub.status.idle":"2023-07-30T03:24:31.620101Z","shell.execute_reply.started":"2023-07-30T03:24:31.599144Z","shell.execute_reply":"2023-07-30T03:24:31.619479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(dataset_train))\nprint(len(dataset_validation))","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:31.62138Z","iopub.execute_input":"2023-07-30T03:24:31.621957Z","iopub.status.idle":"2023-07-30T03:24:31.62847Z","shell.execute_reply.started":"2023-07-30T03:24:31.621927Z","shell.execute_reply":"2023-07-30T03:24:31.627316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(type(dataset_train))\nprint(type(data_loader_train))","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:24:31.630085Z","iopub.execute_input":"2023-07-30T03:24:31.630859Z","iopub.status.idle":"2023-07-30T03:24:31.636653Z","shell.execute_reply.started":"2023-07-30T03:24:31.630827Z","shell.execute_reply":"2023-07-30T03:24:31.635572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## mask 있는 데이터만 train시켜볼까?\n## 이렇게하니까 정확도 내려감 0.422","metadata":{}},{"cell_type":"code","source":"# class ContrailsAshDataset_train(torch.utils.data.Dataset):\n#     def __init__(self, parrent_folder: str):\n#         self.df_idx: pd.DataFrame = pd.DataFrame({'idx': real_data_train_idx})\n#         self.parrent_folder: str = parrent_folder\n            \n#     def __len__(self):\n#         return len(self.df_idx)\n    \n#     def __getitem__(self, idx):\n#         image_id: str = str(self.df_idx.iloc[idx]['idx'])\n#         images = torch.tensor(np.reshape(get_ash_color_images(image_id, self.parrent_folder, get_mask_frame_only=False), (256, 256, 24))).to(torch.float32).permute(2, 0, 1)\n#         mask = torch.tensor(get_mask_image(image_id, self.parrent_folder)).to(torch.float32).permute(2, 0, 1)\n#         return images, mask","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:27:17.257895Z","iopub.execute_input":"2023-07-30T03:27:17.259412Z","iopub.status.idle":"2023-07-30T03:27:17.270977Z","shell.execute_reply.started":"2023-07-30T03:27:17.25937Z","shell.execute_reply":"2023-07-30T03:27:17.269991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dataset_train = ContrailsAshDataset_train('train')","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:27:32.191078Z","iopub.execute_input":"2023-07-30T03:27:32.19168Z","iopub.status.idle":"2023-07-30T03:27:32.197798Z","shell.execute_reply.started":"2023-07-30T03:27:32.191649Z","shell.execute_reply":"2023-07-30T03:27:32.19674Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class ContrailsAshDataset_valid(torch.utils.data.Dataset):\n#     def __init__(self, parrent_folder: str):\n#         self.df_idx: pd.DataFrame = pd.DataFrame({'idx': real_data_valid_idx})\n#         self.parrent_folder: str = parrent_folder\n            \n#     def __len__(self):\n#         return len(self.df_idx)\n    \n#     def __getitem__(self, idx):\n#         image_id: str = str(self.df_idx.iloc[idx]['idx'])\n#         images = torch.tensor(np.reshape(get_ash_color_images(image_id, self.parrent_folder, get_mask_frame_only=False), (256, 256, 24))).to(torch.float32).permute(2, 0, 1)\n#         mask = torch.tensor(get_mask_image(image_id, self.parrent_folder)).to(torch.float32).permute(2, 0, 1)\n#         return images, mask","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:27:52.668809Z","iopub.execute_input":"2023-07-30T03:27:52.669172Z","iopub.status.idle":"2023-07-30T03:27:52.676753Z","shell.execute_reply.started":"2023-07-30T03:27:52.669141Z","shell.execute_reply":"2023-07-30T03:27:52.675812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dataset_validation = ContrailsAshDataset_valid('validation')","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:28:17.915212Z","iopub.execute_input":"2023-07-30T03:28:17.915569Z","iopub.status.idle":"2023-07-30T03:28:17.920382Z","shell.execute_reply.started":"2023-07-30T03:28:17.915541Z","shell.execute_reply":"2023-07-30T03:28:17.919412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(dataset_train))\nprint(len(dataset_validation))","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:28:52.199212Z","iopub.execute_input":"2023-07-30T03:28:52.199563Z","iopub.status.idle":"2023-07-30T03:28:52.204693Z","shell.execute_reply.started":"2023-07-30T03:28:52.199536Z","shell.execute_reply":"2023-07-30T03:28:52.203781Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data_loader_train = DataLoader(dataset_train, batch_size=16, shuffle=True, num_workers=2)\n# data_loader_validation = DataLoader(dataset_validation, batch_size=16, shuffle=True, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:29:04.621856Z","iopub.execute_input":"2023-07-30T03:29:04.622238Z","iopub.status.idle":"2023-07-30T03:29:04.629831Z","shell.execute_reply.started":"2023-07-30T03:29:04.622205Z","shell.execute_reply":"2023-07-30T03:29:04.628857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"train = True","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:29:05.858445Z","iopub.execute_input":"2023-07-30T03:29:05.858807Z","iopub.status.idle":"2023-07-30T03:29:05.863514Z","shell.execute_reply.started":"2023-07-30T03:29:05.858778Z","shell.execute_reply":"2023-07-30T03:29:05.862319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    model = UNet()\n    model.to(device)\n    \n    criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor(100))\n    optimizer = optim.Adam(model.parameters(), lr=0.01)\n    lr_scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, 0.70)\n    \n    num_epochs = 15\n    \n    trainer = MyTrainer(model, optimizer, criterion, lr_scheduler)\n    trainer.fit(data_loader_train, data_loader_validation, epochs=num_epochs)\n    \nelse:\n    model = UNet()\n    model.load_state_dict(torch.load('/kaggle/input/contrails-unet-pretraind/unet.pt'))\n    model.eval()\n    model.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:29:06.891859Z","iopub.execute_input":"2023-07-30T03:29:06.892227Z","iopub.status.idle":"2023-07-30T03:37:34.951793Z","shell.execute_reply.started":"2023-07-30T03:29:06.892192Z","shell.execute_reply":"2023-07-30T03:37:34.944161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training Overview","metadata":{}},{"cell_type":"code","source":"if train:\n    df_data = pd.DataFrame({'Batch Losses': trainer.batch_losses})\n\n    sns.lineplot(data=df_data)\n    plt.xlabel('Batch')\n    plt.ylabel('Loss')\n    plt.title('Batch Loss')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.952814Z","iopub.status.idle":"2023-07-30T03:37:34.953174Z","shell.execute_reply.started":"2023-07-30T03:37:34.952999Z","shell.execute_reply":"2023-07-30T03:37:34.953016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    df_data = pd.DataFrame({'Loss': trainer.epoch_losses})\n\n    sns.lineplot(data=df_data)\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Model Argavgre Training Loss over Epochs')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.958093Z","iopub.status.idle":"2023-07-30T03:37:34.958471Z","shell.execute_reply.started":"2023-07-30T03:37:34.958297Z","shell.execute_reply":"2023-07-30T03:37:34.958314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    df_data = pd.DataFrame({'Loss': trainer.validation_losses})\n\n    sns.lineplot(data=df_data)\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.title('Model Validation Loss over Epochs')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.959877Z","iopub.status.idle":"2023-07-30T03:37:34.961676Z","shell.execute_reply.started":"2023-07-30T03:37:34.961409Z","shell.execute_reply":"2023-07-30T03:37:34.961433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if train:\n    df_data = pd.DataFrame({'Learning rates': trainer.learning_rates})\n\n    sns.lineplot(data=df_data)\n    plt.xlabel('Epoch')\n    plt.ylabel('Learinig Rate')\n    plt.title('Learinig Rate over Epochs')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.963169Z","iopub.status.idle":"2023-07-30T03:37:34.96403Z","shell.execute_reply.started":"2023-07-30T03:37:34.963785Z","shell.execute_reply":"2023-07-30T03:37:34.963808Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Find Optimal Threshold","metadata":{}},{"cell_type":"code","source":"class DiceThresholdTester:\n    \n    def __init__(self, model: nn.Module, data_loader: torch.utils.data.DataLoader):\n        self.model = model\n        self.data_loader = data_loader\n        self.cumulative_mask_pred = []\n        self.cumulative_mask_true = []\n        \n    def precalculate_prediction(self) -> None:\n        sigmoid = nn.Sigmoid()\n        \n        for images, mask_true in self.data_loader:\n            if torch.cuda.is_available():\n                images = images.cuda()\n\n            mask_pred = sigmoid(model.forward(images))\n\n            self.cumulative_mask_pred.append(mask_pred.cpu().detach().numpy())\n            self.cumulative_mask_true.append(mask_true.cpu().detach().numpy())\n            \n        self.cumulative_mask_pred = np.concatenate(self.cumulative_mask_pred, axis=0)\n        self.cumulative_mask_true = np.concatenate(self.cumulative_mask_true, axis=0)\n\n        self.cumulative_mask_pred = torch.flatten(torch.from_numpy(self.cumulative_mask_pred))\n        self.cumulative_mask_true = torch.flatten(torch.from_numpy(self.cumulative_mask_true))\n    \n    def test_threshold(self, threshold: float) -> float:\n        _dice = Dice(use_sigmoid=False)\n        after_threshold = np.zeros(self.cumulative_mask_pred.shape)\n        after_threshold[self.cumulative_mask_pred[:] > threshold] = 1\n        after_threshold[self.cumulative_mask_pred[:] < threshold] = 0\n        after_threshold = torch.flatten(torch.from_numpy(after_threshold))\n        return _dice(self.cumulative_mask_true, after_threshold).item()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.965559Z","iopub.status.idle":"2023-07-30T03:37:34.966103Z","shell.execute_reply.started":"2023-07-30T03:37:34.965854Z","shell.execute_reply":"2023-07-30T03:37:34.965877Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dice_threshold_tester = DiceThresholdTester(model, data_loader_validation)\ndice_threshold_tester.precalculate_prediction()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.967793Z","iopub.status.idle":"2023-07-30T03:37:34.969556Z","shell.execute_reply.started":"2023-07-30T03:37:34.969297Z","shell.execute_reply":"2023-07-30T03:37:34.96932Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"thresholds_to_test = [round(x * 0.01, 2) for x in range(101)]\n\noptim_threshold = 0.975\nbest_dice_score = -1\n\nthresholds = []\ndice_scores = []\n\nfor t in thresholds_to_test:\n    dice_score = dice_threshold_tester.test_threshold(t)\n    if dice_score > best_dice_score:\n        best_dice_score = dice_score\n        optim_threshold = t\n    \n    thresholds.append(t)\n    dice_scores.append(dice_score)\n    \nprint(f'Best Threshold: {optim_threshold} with dice: {best_dice_score}')\ndf_threshold_data = pd.DataFrame({'Threshold': thresholds, 'Dice Score': dice_scores})","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.970977Z","iopub.status.idle":"2023-07-30T03:37:34.971793Z","shell.execute_reply.started":"2023-07-30T03:37:34.971543Z","shell.execute_reply":"2023-07-30T03:37:34.971567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.lineplot(data=df_threshold_data, x='Threshold', y='Dice Score')\nplt.axhline(y=best_dice_score, color='green')\nplt.axvline(x=optim_threshold, color='green')\nplt.text(-0.02, best_dice_score * 0.96, f'{best_dice_score:.3f}', va='center', ha='left', color='green')\nplt.text(optim_threshold - 0.01, 0.02, f'{optim_threshold}', va='center', ha='right', color='green')\nplt.ylim(bottom=0)\nplt.title('Threshold vs Dice Score')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.973239Z","iopub.status.idle":"2023-07-30T03:37:34.974021Z","shell.execute_reply.started":"2023-07-30T03:37:34.973761Z","shell.execute_reply":"2023-07-30T03:37:34.973784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preview Models Predictions on Validation","metadata":{}},{"cell_type":"code","source":"def sigmoid(x):\n    return 1 / (1 + np.exp(-x))\n\nbatches_to_show = 4\nmodel.eval()\n\nfor i, data in enumerate(data_loader_validation):\n    images, mask = data\n    \n    # Predict mask for this instance\n    if torch.cuda.is_available():\n        images = images.cuda()\n    predicated_mask = sigmoid(model.forward(images[:, :, :, :]).cpu().detach().numpy())\n    \n    # Apply threshold\n    predicated_mask_with_threshold = np.zeros((images.shape[0], 256, 256))\n    predicated_mask_with_threshold[predicated_mask[:, 0, :, :] < optim_threshold] = 0\n    predicated_mask_with_threshold[predicated_mask[:, 0, :, :] > optim_threshold] = 1\n    \n    images = images.cpu()\n        \n    for img_num in range(0, images.shape[0]):\n        fig, axes = plt.subplots(nrows=1, ncols=4, figsize=(20,10))\n        axes = axes.flatten()\n        \n        # Show groud trought \n        axes[0].imshow(mask[img_num, 0, :, :])\n        axes[0].axis('off')\n        axes[0].set_title('Ground Truth')\n        \n        # Show ash color scheme input image\n        axes[1].imshow( np.concatenate(\n            (\n            np.expand_dims(images[img_num, 4, :, :], axis=2),\n            np.expand_dims(images[img_num, 12, :, :], axis=2),\n            np.expand_dims(images[img_num, 20, :, :], axis=2)\n        ), axis=2))\n        axes[1].axis('off')\n        axes[1].set_title('Ash color scheeme input - Frame 4')\n\n        # Show predicted mask\n        axes[2].imshow(predicated_mask[img_num, 0, :, :], vmin=0, vmax=1)\n        axes[2].axis('off')\n        axes[2].set_title('Predicted probability mask')\n\n        # Show predicted mask after threshold\n        axes[3].imshow(predicated_mask_with_threshold[img_num, :, :])\n        axes[3].axis('off')\n        axes[3].set_title('Predicted mask with threshold')\n        plt.show()\n    \n    if i + 1 >= batches_to_show:\n        break","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.975455Z","iopub.status.idle":"2023-07-30T03:37:34.97624Z","shell.execute_reply.started":"2023-07-30T03:37:34.975964Z","shell.execute_reply":"2023-07-30T03:37:34.975987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"# Same a the ContrailsAshDataset but does not load the mask (since its not avalble for test) and returns image_id instaed of the mask to assamble submission\nclass ContrailsAshTestDataset(torch.utils.data.Dataset):\n    def __init__(self):\n        self.df_idx: pd.DataFrame = pd.DataFrame({'idx': os.listdir(f'/kaggle/input/google-research-identify-contrails-reduce-global-warming/test')})\n        self.parrent_folder: str = 'test'\n\n    def __len__(self):\n        return len(self.df_idx)\n\n    def __getitem__(self, idx):\n        image_id: int = int(self.df_idx.iloc[idx]['idx'])\n        images = torch.tensor(np.reshape(get_ash_color_images(str(image_id), self.parrent_folder, get_mask_frame_only=False), (256, 256, 24))).to(torch.float32).permute(2, 0, 1)\n        return images,  torch.tensor(image_id)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.977635Z","iopub.status.idle":"2023-07-30T03:37:34.978432Z","shell.execute_reply.started":"2023-07-30T03:37:34.978172Z","shell.execute_reply":"2023-07-30T03:37:34.97821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataset_test = ContrailsAshTestDataset()\ndata_loader_test = DataLoader(dataset_test, batch_size=16, shuffle=True, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.979815Z","iopub.status.idle":"2023-07-30T03:37:34.980625Z","shell.execute_reply.started":"2023-07-30T03:37:34.980379Z","shell.execute_reply":"2023-07-30T03:37:34.980401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#source https://www.kaggle.com/code/inversion/contrails-rle-submission?scriptVersionId=128527711&cellId=4\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\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","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.982014Z","iopub.status.idle":"2023-07-30T03:37:34.982783Z","shell.execute_reply.started":"2023-07-30T03:37:34.982541Z","shell.execute_reply":"2023-07-30T03:37:34.982563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = 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-30T03:37:34.98422Z","iopub.status.idle":"2023-07-30T03:37:34.984986Z","shell.execute_reply.started":"2023-07-30T03:37:34.984741Z","shell.execute_reply":"2023-07-30T03:37:34.984764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, data in enumerate(data_loader_test):\n    images, image_id = data\n    \n    # Predict mask for this instance\n    if torch.cuda.is_available():\n        images = images.cuda()\n    predicated_mask = sigmoid(model.forward(images[:, :, :, :]).cpu().detach().numpy())\n    \n    # Apply threshold\n    predicated_mask_with_threshold = np.zeros((images.shape[0], 256, 256))\n    predicated_mask_with_threshold[predicated_mask[:, 0, :, :] < optim_threshold] = 0\n    predicated_mask_with_threshold[predicated_mask[:, 0, :, :] > optim_threshold] = 1\n    \n    for img_num in range(0, images.shape[0]):\n        current_mask = predicated_mask_with_threshold[img_num, :, :]\n        current_image_id = image_id[img_num].item()\n        \n        submission.loc[int(current_image_id), 'encoded_pixels'] = list_to_string(rle_encode(current_mask))","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.986401Z","iopub.status.idle":"2023-07-30T03:37:34.98717Z","shell.execute_reply.started":"2023-07-30T03:37:34.986906Z","shell.execute_reply":"2023-07-30T03:37:34.98693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.988568Z","iopub.status.idle":"2023-07-30T03:37:34.989387Z","shell.execute_reply.started":"2023-07-30T03:37:34.989092Z","shell.execute_reply":"2023-07-30T03:37:34.989115Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"execution":{"iopub.status.busy":"2023-07-30T03:37:34.990787Z","iopub.status.idle":"2023-07-30T03:37:34.991584Z","shell.execute_reply.started":"2023-07-30T03:37:34.991336Z","shell.execute_reply":"2023-07-30T03:37:34.991359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}