{"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":"# Imports","metadata":{}},{"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\n!pip install /kaggle/input/torchsummary/torchsummary-1.5.1-py3-none-any.whl\nfrom torchsummary import summary","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-05-31T10:09:21.357308Z","iopub.execute_input":"2023-05-31T10:09:21.358016Z","iopub.status.idle":"2023-05-31T10:09:56.645067Z","shell.execute_reply.started":"2023-05-31T10:09:21.35798Z","shell.execute_reply":"2023-05-31T10:09:56.643864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Preprocessing Utilites","metadata":{}},{"cell_type":"code","source":"data_dir: str = '/kaggle/input/google-research-identify-contrails-reduce-global-warming'","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-31T10:09:56.647603Z","iopub.execute_input":"2023-05-31T10:09:56.648869Z","iopub.status.idle":"2023-05-31T10:09:56.655865Z","shell.execute_reply.started":"2023-05-31T10:09:56.648811Z","shell.execute_reply":"2023-05-31T10:09:56.654929Z"},"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":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-31T10:09:56.657418Z","iopub.execute_input":"2023-05-31T10:09:56.657711Z","iopub.status.idle":"2023-05-31T10:09:57.08657Z","shell.execute_reply.started":"2023-05-31T10:09:56.657688Z","shell.execute_reply":"2023-05-31T10:09:57.085504Z"},"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-05-31T10:09:57.089306Z","iopub.execute_input":"2023-05-31T10:09:57.0897Z","iopub.status.idle":"2023-05-31T10:09:57.094643Z","shell.execute_reply.started":"2023-05-31T10:09:57.089616Z","shell.execute_reply":"2023-05-31T10:09:57.093386Z"},"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\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-05-31T10:09:57.096052Z","iopub.execute_input":"2023-05-31T10:09:57.096653Z","iopub.status.idle":"2023-05-31T10:09:57.107119Z","shell.execute_reply.started":"2023-05-31T10:09:57.096623Z","shell.execute_reply":"2023-05-31T10:09:57.106179Z"},"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-05-31T10:09:57.108563Z","iopub.execute_input":"2023-05-31T10:09:57.109141Z","iopub.status.idle":"2023-05-31T10:09:57.117722Z","shell.execute_reply.started":"2023-05-31T10:09:57.109112Z","shell.execute_reply":"2023-05-31T10:09:57.116865Z"},"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\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    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":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-31T10:09:57.119081Z","iopub.execute_input":"2023-05-31T10:09:57.119648Z","iopub.status.idle":"2023-05-31T10:11:58.135247Z","shell.execute_reply.started":"2023-05-31T10:09:57.119618Z","shell.execute_reply":"2023-05-31T10:11:58.134314Z"},"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\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    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":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-31T10:11:58.136725Z","iopub.execute_input":"2023-05-31T10:11:58.137171Z","iopub.status.idle":"2023-05-31T10:12:08.162434Z","shell.execute_reply.started":"2023-05-31T10:11:58.137136Z","shell.execute_reply":"2023-05-31T10:12:08.161544Z"},"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('Presentage of Dataset')\nax.set_title('With Contrails vs No Contrails')\n\nplt.legend()\nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-31T10:12:08.167296Z","iopub.execute_input":"2023-05-31T10:12:08.16953Z","iopub.status.idle":"2023-05-31T10:12:08.57116Z","shell.execute_reply.started":"2023-05-31T10:12:08.169499Z","shell.execute_reply":"2023-05-31T10:12:08.570253Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"There is a significantly higher percentage of images with contrails in the train dataset than in the validation.","metadata":{}},{"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']})\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')\naxes[0].set_xlabel('')\naxes[0].set_ylabel('Presentage of Contrails Pixels')\naxes[0].set_title('Presentage of Contrails pixels in Images')\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')\naxes[1].set_xlabel('')\naxes[1].set_ylabel('Presentage of Contrails Pixels')\naxes[1].set_title('Presentage of Contrails pixels in Images wiht Contrails Present')\n\nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-31T10:12:08.575771Z","iopub.execute_input":"2023-05-31T10:12:08.576057Z","iopub.status.idle":"2023-05-31T10:12:08.975785Z","shell.execute_reply.started":"2023-05-31T10:12:08.576035Z","shell.execute_reply":"2023-05-31T10:12:08.974885Z"},"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:\nAfter 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:\nAnother 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.","metadata":{}},{"cell_type":"markdown","source":"Additionally, 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, 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_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":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-31T10:12:08.977048Z","iopub.execute_input":"2023-05-31T10:12:08.97806Z","iopub.status.idle":"2023-05-31T10:12:08.988025Z","shell.execute_reply.started":"2023-05-31T10:12:08.978028Z","shell.execute_reply":"2023-05-31T10:12:08.986931Z"},"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[730 + i]['idx']), 'train')\n    axes[i].imshow(images[:,:,:,4])\n    axes[i].axis('off')","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-31T10:12:08.989573Z","iopub.execute_input":"2023-05-31T10:12:08.990018Z","iopub.status.idle":"2023-05-31T10:12:17.672355Z","shell.execute_reply.started":"2023-05-31T10:12:08.989988Z","shell.execute_reply":"2023-05-31T10:12:17.671185Z"},"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":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-31T10:12:17.673676Z","iopub.execute_input":"2023-05-31T10:12:17.674106Z","iopub.status.idle":"2023-05-31T10:12:18.118092Z","shell.execute_reply.started":"2023-05-31T10:12:17.674065Z","shell.execute_reply":"2023-05-31T10:12:18.117218Z"},"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# Animation\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]\nanim = animation.FuncAnimation(\n    fig, draw, frames=false_color.shape[-1], interval=500, blit=True\n)\nplt.close()\ndisplay.HTML(anim.to_jshtml())","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-05-31T10:12:18.119395Z","iopub.execute_input":"2023-05-31T10:12:18.119893Z","iopub.status.idle":"2023-05-31T10:12:19.704096Z","shell.execute_reply.started":"2023-05-31T10:12:18.119863Z","shell.execute_reply":"2023-05-31T10:12:19.700913Z"},"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-05-31T10:12:19.706223Z","iopub.execute_input":"2023-05-31T10:12:19.706666Z","iopub.status.idle":"2023-05-31T10:12:19.742924Z","shell.execute_reply.started":"2023-05-31T10:12:19.706626Z","shell.execute_reply":"2023-05-31T10:12:19.742089Z"},"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    def forward(self, x):\n        return self.maxpool_conv(x)\n\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_channels // 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\n\nclass UNet(nn.Module):\n    def __init__(self):\n        super(UNet, self).__init__()\n        # Define your 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-05-31T10:12:19.744334Z","iopub.execute_input":"2023-05-31T10:12:19.744837Z","iopub.status.idle":"2023-05-31T10:12:19.774186Z","shell.execute_reply.started":"2023-05-31T10:12:19.744791Z","shell.execute_reply":"2023-05-31T10:12:19.77305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"summary(UNet().to(device), input_size=(24, 256, 256))","metadata":{"execution":{"iopub.status.busy":"2023-05-31T10:12:19.776059Z","iopub.execute_input":"2023-05-31T10:12:19.776875Z","iopub.status.idle":"2023-05-31T10:12:27.473751Z","shell.execute_reply.started":"2023-05-31T10:12:19.776835Z","shell.execute_reply":"2023-05-31T10:12:27.472753Z"},"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-05-31T10:12:27.47571Z","iopub.execute_input":"2023-05-31T10:12:27.47645Z","iopub.status.idle":"2023-05-31T10:12:27.485585Z","shell.execute_reply.started":"2023-05-31T10:12:27.476412Z","shell.execute_reply":"2023-05-31T10:12:27.484395Z"},"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'epotch: {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))","metadata":{"execution":{"iopub.status.busy":"2023-05-31T10:12:27.487338Z","iopub.execute_input":"2023-05-31T10:12:27.487996Z","iopub.status.idle":"2023-05-31T10:12:27.506789Z","shell.execute_reply.started":"2023-05-31T10:12:27.487964Z","shell.execute_reply":"2023-05-31T10:12:27.505852Z"},"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-05-31T10:12:27.50831Z","iopub.execute_input":"2023-05-31T10:12:27.50868Z","iopub.status.idle":"2023-05-31T10:12:27.519617Z","shell.execute_reply.started":"2023-05-31T10:12:27.508628Z","shell.execute_reply":"2023-05-31T10:12:27.518741Z"},"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-05-31T10:12:27.521038Z","iopub.execute_input":"2023-05-31T10:12:27.521413Z","iopub.status.idle":"2023-05-31T10:12:27.546848Z","shell.execute_reply.started":"2023-05-31T10:12:27.521383Z","shell.execute_reply":"2023-05-31T10:12:27.546044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"train = True","metadata":{"execution":{"iopub.status.busy":"2023-05-31T10:12:27.549354Z","iopub.execute_input":"2023-05-31T10:12:27.549655Z","iopub.status.idle":"2023-05-31T10:12:27.555488Z","shell.execute_reply.started":"2023-05-31T10:12:27.549626Z","shell.execute_reply":"2023-05-31T10:12:27.55453Z"},"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 = 11\n\n    trainer = MyTrainer(model, optimizer, criterion, lr_scheduler)\n    trainer.fit(data_loader_train, data_loader_validation, epochs=num_epochs)\nelse:\n    model = UNet()\n    model.load_state_dict(torch.load('/kaggle/input/contrails-unet-pretrained/unet.pt'))\n    model.eval()\n    model.to(device)","metadata":{"execution":{"iopub.status.busy":"2023-05-31T10:12:27.55688Z","iopub.execute_input":"2023-05-31T10:12:27.557261Z"},"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":{"_kg_hide-input":true,"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":{"_kg_hide-input":true,"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":{"_kg_hide-input":true,"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":{"_kg_hide-input":true,"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":{"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dice_threshold_tester = DiceThresholdTester(model, data_loader_validation)\ndice_threshold_tester.precalculate_prediction()","metadata":{"_kg_hide-input":true,"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":{"_kg_hide-input":true,"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":{"_kg_hide-input":true,"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":{"_kg_hide-input":true,"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":{"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":{"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":{"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":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.to_csv('submission.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}