{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-06-16T17:40:09.071475Z","iopub.execute_input":"2023-06-16T17:40:09.071851Z","iopub.status.idle":"2023-06-16T17:40:09.077116Z","shell.execute_reply.started":"2023-06-16T17:40:09.071821Z","shell.execute_reply":"2023-06-16T17:40:09.075974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install segmentation-models-pytorch\n#!pip install torch-summary\n#from torchsummary import summary\nimport segmentation_models_pytorch as smp","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport os\nimport random\nimport pandas as pd\nfrom tqdm.notebook import tqdm\nimport matplotlib.pyplot as plt\nfrom IPython import display\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nimport torch.utils.checkpoint as C\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.utils.class_weight import compute_class_weight","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.214058Z","iopub.execute_input":"2023-06-16T17:40:31.214446Z","iopub.status.idle":"2023-06-16T17:40:31.222777Z","shell.execute_reply.started":"2023-06-16T17:40:31.214409Z","shell.execute_reply":"2023-06-16T17:40:31.221711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_path=\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train\"\nvalidation_path=\"/kaggle/input/google-research-identify-contrails-reduce-global-warming/validation\"\ntrain_ids=[]\nvalidation_ids=[]\nfor id in os.listdir(train_path):\n    train_ids.append(id)\nfor id in os.listdir(validation_path):\n    validation_ids.append(id)\nprint(len(train_ids))\nprint(len(validation_ids))","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.223945Z","iopub.execute_input":"2023-06-16T17:40:31.224428Z","iopub.status.idle":"2023-06-16T17:40:31.517314Z","shell.execute_reply.started":"2023-06-16T17:40:31.224401Z","shell.execute_reply":"2023-06-16T17:40:31.5163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if torch.cuda.is_available():\n    device=\"cuda\"\nelse:\n    device=\"cpu\"","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.52009Z","iopub.execute_input":"2023-06-16T17:40:31.520441Z","iopub.status.idle":"2023-06-16T17:40:31.524761Z","shell.execute_reply.started":"2023-06-16T17:40:31.520408Z","shell.execute_reply":"2023-06-16T17:40:31.52369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define ash_color conversion\nConversion to ash color provided by https://www.kaggle.com/code/inversion/visualizing-contrails <br />\nFollowing guidelines set in https://eumetrain.org/sites/default/files/2020-05/RGB_recipes.pdf","metadata":{}},{"cell_type":"code","source":"def ash_color(b11,b14,b15):\n    def correction(band,bounds):\n        return (band - bounds[0]) / (bounds[1] - bounds[0])\n    _T11_BOUNDS = (243, 303)\n    _CLOUD_TOP_TDIFF_BOUNDS = (-4, 5)\n    _TDIFF_BOUNDS = (-4, 2)\n    r = correction(b15 - b14, _TDIFF_BOUNDS)\n    g = correction(b14 - b11, _CLOUD_TOP_TDIFF_BOUNDS)\n    b = correction(b14, _T11_BOUNDS)\n    return np.clip(np.stack([r, g, b], axis=2), 0, 1)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.526269Z","iopub.execute_input":"2023-06-16T17:40:31.526622Z","iopub.status.idle":"2023-06-16T17:40:31.535668Z","shell.execute_reply.started":"2023-06-16T17:40:31.526591Z","shell.execute_reply":"2023-06-16T17:40:31.534758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Define Dataset and Dataloader\nDefine a dataloader to provide 1 or 3 frames based on input","metadata":{}},{"cell_type":"code","source":"class ContrailDataset(Dataset):\n    def __init__(self,ids,target_path,frames=False):\n        self.ids=ids\n        self.target_path=target_path\n        self.frames=frames\n    \n    def __len__(self):\n        return len(self.ids)\n    def __getitem__(self,index):\n        N_before=4\n        N_after=3\n        source=f\"{self.target_path}/{self.ids[index]}\"\n        b11_central = np.load(f\"{source}/band_11.npy\")[..., N_before]\n\n        b14_central = np.load(f\"{source}/band_14.npy\")[..., N_before]\n        \n        b15_central = np.load(f\"{source}/band_15.npy\")[..., N_before]\n        \n        image_central=ash_color(b11_central,b14_central,b15_central)\n        image_central=torch.Tensor(image_central)\n        image_central=image_central.permute(2,0,1)\n        label=np.load(f\"{source}/human_pixel_masks.npy\")\n        label=torch.Tensor(label)\n        label=label.permute(2,0,1)\n        if not self.frames:\n            return image_central,label\n        else:\n            b11_prev=np.load(f\"{source}/band_11.npy\")[..., N_before-2]\n            b11_future=np.load(f\"{source}/band_11.npy\")[..., N_before+2]\n            b14_prev = np.load(f\"{source}/band_14.npy\")[..., N_before-2]\n            b14_future = np.load(f\"{source}/band_14.npy\")[..., N_before+2]\n            b15_prev = np.load(f\"{source}/band_15.npy\")[..., N_before-2]\n            b15_future = np.load(f\"{source}/band_15.npy\")[..., N_before+2]\n            image_prev=ash_color(b11_prev,b14_prev,b15_prev)\n            image_prev=torch.Tensor(image_prev)\n            image_prev=image_prev.permute(2,0,1)\n            image_future=ash_color(b11_future,b14_future,b15_future)\n            image_future=torch.Tensor(image_future)\n            image_future=image_future.permute(2,0,1)\n            label=np.load(f\"{source}/human_pixel_masks.npy\")\n            label=torch.Tensor(label)\n            label=label.permute(2,0,1)\n            return image_prev,image_central,image_future,label","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.537024Z","iopub.execute_input":"2023-06-16T17:40:31.538089Z","iopub.status.idle":"2023-06-16T17:40:31.551703Z","shell.execute_reply.started":"2023-06-16T17:40:31.538059Z","shell.execute_reply":"2023-06-16T17:40:31.550705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size=16\nlr=0.0001\nepochs=10\n# We lack the compute to efficiently utilize the entire dataset\ntrain_size=10000","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.553138Z","iopub.execute_input":"2023-06-16T17:40:31.553907Z","iopub.status.idle":"2023-06-16T17:40:31.563547Z","shell.execute_reply.started":"2023-06-16T17:40:31.553868Z","shell.execute_reply":"2023-06-16T17:40:31.562672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize Contrails","metadata":{}},{"cell_type":"code","source":"random.shuffle(train_ids)\ntrain_data=ContrailDataset(train_ids[:100],train_path)\ntrain_loader=DataLoader(train_data,batch_size,shuffle=True,num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:44:52.798481Z","iopub.execute_input":"2023-06-16T17:44:52.798891Z","iopub.status.idle":"2023-06-16T17:44:52.824666Z","shell.execute_reply.started":"2023-06-16T17:44:52.798859Z","shell.execute_reply":"2023-06-16T17:44:52.823773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for batch in train_loader:\n    images,labels=batch\n    break","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:47:35.767409Z","iopub.execute_input":"2023-06-16T17:47:35.767775Z","iopub.status.idle":"2023-06-16T17:47:36.258864Z","shell.execute_reply.started":"2023-06-16T17:47:35.767726Z","shell.execute_reply":"2023-06-16T17:47:36.257379Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Visualization from https://www.kaggle.com/code/inversion/visualizing-contrails\nplt.figure(figsize=(18, 6))\nax = plt.subplot(1, 3, 1)\nax.imshow(images[0].permute(1,2,0))\nax.set_title('False color image')\n\nax = plt.subplot(1, 3, 2)\nax.imshow(labels[0].permute(1,2,0), interpolation='none')\nax.set_title('Ground truth contrail mask')\n\nax = plt.subplot(1, 3, 3)\nax.imshow(images[0].permute(1,2,0))\nax.imshow(labels[0].permute(1,2,0), cmap='Reds', alpha=.4, interpolation='none')\nax.set_title('Contrail mask on false color image');","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:47:46.689393Z","iopub.execute_input":"2023-06-16T17:47:46.689807Z","iopub.status.idle":"2023-06-16T17:47:47.974841Z","shell.execute_reply.started":"2023-06-16T17:47:46.689766Z","shell.execute_reply":"2023-06-16T17:47:47.973964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Define loss function and validation metric\nDice Coefficient and Dice Loss are much more effective than cross-entropy for segmentation tasks as they measure the overlap between predicted and ground truth masks.</br>\nCross-entropy on the other hand struggles with heavily imbalanced data, as is common in segmentation","metadata":{}},{"cell_type":"code","source":"loss_fn=smp.losses.DiceLoss(mode='binary')","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.564807Z","iopub.execute_input":"2023-06-16T17:40:31.565482Z","iopub.status.idle":"2023-06-16T17:40:31.576092Z","shell.execute_reply.started":"2023-06-16T17:40:31.565451Z","shell.execute_reply":"2023-06-16T17:40:31.575176Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice_coeff(pred_mask,truth):\n    truth=truth.flatten(1)\n    intersection=torch.logical_and(pred_mask,truth.squeeze(1))\n    union=pred_mask.sum()+truth.sum()\n    dice=2*torch.sum(intersection)/(union+1e-8)\n    return dice","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.577476Z","iopub.execute_input":"2023-06-16T17:40:31.578293Z","iopub.status.idle":"2023-06-16T17:40:31.585667Z","shell.execute_reply.started":"2023-06-16T17:40:31.57826Z","shell.execute_reply":"2023-06-16T17:40:31.584779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define building blocks of U-Net and TCN","metadata":{}},{"cell_type":"code","source":"class convolutional_block(nn.Module):\n    def __init__(self,in_channels,out_channels,device='cpu'):\n        super().__init__()\n        self.device=device\n        self.conv1=nn.Conv2d(in_channels,out_channels,kernel_size=3,padding=1).to(self.device)\n        self.norm1=nn.BatchNorm2d(out_channels).to(self.device)\n        self.conv2=nn.Conv2d(out_channels,out_channels,kernel_size=3,padding=1).to(self.device)\n        self.norm2=nn.BatchNorm2d(out_channels).to(self.device)\n        self.relu=nn.ReLU().to(self.device)\n        \n    def forward(self,x):\n        x=self.conv1(x)\n        x=self.norm1(x)\n        x=self.relu(x)\n        \n        x=self.conv2(x)\n        x=self.norm2(x)\n        x=self.relu(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.589444Z","iopub.execute_input":"2023-06-16T17:40:31.590349Z","iopub.status.idle":"2023-06-16T17:40:31.598598Z","shell.execute_reply.started":"2023-06-16T17:40:31.590319Z","shell.execute_reply":"2023-06-16T17:40:31.59769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class encoder_block(nn.Module):\n    def __init__(self,in_channels,out_channels,device='cpu'):\n        super().__init__()\n        self.device=device\n        self.pooling=nn.MaxPool2d(2,2).to(self.device)\n        self.conv=convolutional_block(in_channels,out_channels,device=self.device)\n    \n    def forward(self,x):\n        skip=self.conv(x)\n        pool=self.pooling(skip)\n        return skip,pool","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.599836Z","iopub.execute_input":"2023-06-16T17:40:31.60041Z","iopub.status.idle":"2023-06-16T17:40:31.608215Z","shell.execute_reply.started":"2023-06-16T17:40:31.600338Z","shell.execute_reply":"2023-06-16T17:40:31.607228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"ConvTranspose2d allows for trainable up-convolutions as opposed to nn.upsample </br>\nnn.Upsample uses standard interpolation methods such as linear, bilinear, bicubic, nearest neighbor etc.</br>\nhttps://discuss.pytorch.org/t/torch-nn-convtranspose2d-vs-torch-nn-upsample/30574","metadata":{}},{"cell_type":"code","source":"class decoder_block(nn.Module):\n    def __init__(self,in_channels,out_channels,device='cpu'):\n        super().__init__()\n        self.device=device\n        self.conv=convolutional_block(out_channels*2,out_channels,self.device)\n        self.upconv=nn.ConvTranspose2d(in_channels,out_channels,kernel_size=2,stride=2,padding=0).to(self.device)\n        \n    def forward(self,x,pool):\n        x=self.upconv(x)\n        x=torch.concat([x,pool],axis=1)\n        x=self.conv(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.609838Z","iopub.execute_input":"2023-06-16T17:40:31.610613Z","iopub.status.idle":"2023-06-16T17:40:31.617854Z","shell.execute_reply.started":"2023-06-16T17:40:31.610581Z","shell.execute_reply":"2023-06-16T17:40:31.616972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define U-net (https://arxiv.org/abs/1505.04597)","metadata":{}},{"cell_type":"code","source":"class uNet(nn.Module):\n    def __init__(self,in_channels,num_classes=1,device=\"cpu\"):\n        super().__init__()\n        self.device=device\n        # Encoders\n        self.encoder1=encoder_block(in_channels,64,self.device)\n        self.encoder2=encoder_block(64,128,self.device)\n        self.encoder3=encoder_block(128,256,self.device)\n        self.encoder4=encoder_block(256,512,self.device)\n        \n        # Bottleneck\n        self.bottleneck=convolutional_block(512,1024).to(self.device)\n        \n        # Decoders\n        self.decoder1=decoder_block(1024,512,self.device)\n        self.decoder2=decoder_block(512,256,self.device)\n        self.decoder3=decoder_block(256,128,self.device)\n        self.decoder4=decoder_block(128,64,self.device)\n        \n        # Classifier\n        self.output=nn.Conv2d(64,1,1).to(self.device)\n        \n        \n    def forward(self,x):\n        # Encoders\n        skip1,pool1=self.encoder1(x)\n        skip2,pool2=self.encoder2(pool1)\n        skip3,pool3=self.encoder3(pool2)\n        skip4,pool4=self.encoder4(pool3)\n        \n        # Bottleneck\n        bn=self.bottleneck(pool4)\n        \n        # Decoders\n        decode1=self.decoder1(bn,skip4)\n        decode2=self.decoder2(decode1,skip3)\n        decode3=self.decoder3(decode2,skip2)\n        decode4=self.decoder4(decode3,skip1)\n        \n        # Output\n        op=self.output(decode4)\n        \n        return op\n        \n        ","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.620076Z","iopub.execute_input":"2023-06-16T17:40:31.621505Z","iopub.status.idle":"2023-06-16T17:40:31.632134Z","shell.execute_reply.started":"2023-06-16T17:40:31.621481Z","shell.execute_reply":"2023-06-16T17:40:31.631172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define building block of TCN\nDilation TCN is utilized with 3d convolutions to handle 2D images with the temporal axis being the third dimension</br>\nhttps://www.kaggle.com/code/ceshine/pytorch-temporal-convolutional-networks","metadata":{}},{"cell_type":"code","source":"class TC_Block(nn.Module):\n    def __init__(self,in_channels,out_channels,kernel_size,dilation,device='cpu'):\n        super().__init__()\n        self.device=device\n        # No dilation over time dimension since there's only 3 samples\n        self.conv=nn.Conv3d(in_channels,out_channels,kernel_size=(3,kernel_size,kernel_size),\n                            padding=(1,1,1),dilation=(1,dilation,dilation)).to(self.device)\n        self.norm=nn.BatchNorm3d(out_channels).to(self.device)\n        self.relu=nn.ReLU().to(self.device)\n        self.dropout=nn.Dropout3d().to(self.device)\n        \n    def forward(self,x):\n        out=self.conv(x)\n        out=self.norm(out)\n        out=self.relu(out)\n        out=self.dropout(out)\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.634004Z","iopub.execute_input":"2023-06-16T17:40:31.634713Z","iopub.status.idle":"2023-06-16T17:40:31.644303Z","shell.execute_reply.started":"2023-06-16T17:40:31.63464Z","shell.execute_reply":"2023-06-16T17:40:31.643634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TCN(nn.Module):\n    def __init__(self,num_inputs,num_channels,kernel_size=2,dropout=0.2,device='cpu'):\n        super().__init__()\n        self.num_inputs=num_inputs\n        self.num_channels=num_channels\n        self.kernel_size=kernel_size\n        self.dropout=dropout\n        self.layers=nn.ModuleList()\n        self.device=device\n        for i in range(len(num_channels)):\n            dilation_size=2**i\n            in_channels=num_inputs if i==0 else num_channels[i-1]\n            out_channels=num_channels[i]\n            layer=TC_Block(in_channels,out_channels,kernel_size,dilation_size,device=self.device)\n            self.layers.append(layer)\n        self.final_conv=nn.Conv3d(in_channels=num_channels[-1],out_channels=1,kernel_size=1).to(self.device)\n    def forward(self,x):\n        x=x.permute(0,2,1,3,4)\n        for layer in self.layers:\n            x=layer(x)\n        out=self.final_conv(x)\n        # Squeeze over extra dimension\n        out=out.squeeze(dim=2)\n        out=out.permute(0,2,1,3,4)\n        return out","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.645765Z","iopub.execute_input":"2023-06-16T17:40:31.646407Z","iopub.status.idle":"2023-06-16T17:40:31.658194Z","shell.execute_reply.started":"2023-06-16T17:40:31.646376Z","shell.execute_reply":"2023-06-16T17:40:31.657541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define TCN augmented Unet","metadata":{}},{"cell_type":"code","source":"class TCN_uNet_Central(nn.Module):\n    def __init__(self,in_channels,num_classes=1,rgb_channels=3,num_inputs=3,num_channels=[16,32,64],kernel_size=3,device=\"cpu\"):\n        super().__init__()\n        self.num_inputs=num_inputs\n        self.num_channels = num_channels\n        self.kernel_size = kernel_size\n        self.in_channels=in_channels\n        self.num_inputs=num_inputs,\n        self.num_classes=num_classes\n        self.device=device\n        dropout=0.2\n        self.TCN_block=TCN(num_inputs,num_channels,kernel_size,dropout,device=self.device)\n        self.unet=uNet(6,num_classes,device=self.device)\n        \n    def forward(self,prev,central,future):\n        temporal_map=self.TCN_block(torch.cat([prev.unsqueeze(1),central.unsqueeze(1),future.unsqueeze(1)],axis=1))\n        temporal_map=F.pad(temporal_map,(0,prev.shape[-2]-temporal_map.shape[-2],0,prev.shape[-1]-temporal_map.shape[-2]))\n        op=torch.cat([central,temporal_map.squeeze(2)],dim=1)\n        op=self.unet(op)\n        return op","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.659699Z","iopub.execute_input":"2023-06-16T17:40:31.660427Z","iopub.status.idle":"2023-06-16T17:40:31.670543Z","shell.execute_reply.started":"2023-06-16T17:40:31.660394Z","shell.execute_reply":"2023-06-16T17:40:31.669789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"op_path=\"/kaggle/working\"","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:40:31.672072Z","iopub.execute_input":"2023-06-16T17:40:31.672794Z","iopub.status.idle":"2023-06-16T17:40:31.68276Z","shell.execute_reply.started":"2023-06-16T17:40:31.67276Z","shell.execute_reply":"2023-06-16T17:40:31.681789Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Base Unet taking the central image as input","metadata":{}},{"cell_type":"code","source":"random.shuffle(train_ids)\nrandom.shuffle(validation_ids)\ntrain_data=ContrailDataset(train_ids[:train_size],train_path)\nvalidation_data=ContrailDataset(validation_ids,validation_path)\ntrain_loader=DataLoader(train_data,batch_size,shuffle=True,num_workers=2)\nvalidation_loader=DataLoader(validation_data,batch_size,shuffle=None,num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:49:40.164492Z","iopub.execute_input":"2023-06-16T08:49:40.164768Z","iopub.status.idle":"2023-06-16T08:49:40.194037Z","shell.execute_reply.started":"2023-06-16T08:49:40.164734Z","shell.execute_reply":"2023-06-16T08:49:40.19302Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model=uNet(3,1,device=device)\noptimizer=torch.optim.Adam(model.parameters(),lr=lr)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:49:40.195301Z","iopub.execute_input":"2023-06-16T08:49:40.195685Z","iopub.status.idle":"2023-06-16T08:49:43.448522Z","shell.execute_reply.started":"2023-06-16T08:49:40.195653Z","shell.execute_reply":"2023-06-16T08:49:43.447605Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_val=0\nPATH=f'{op_path}/best_uNet_{train_size}.pt'\nfor epoch in range(epochs):\n    total_loss=0\n    model.train()\n    for batch in train_loader:\n        optimizer.zero_grad()\n        img,labels=batch\n        labels=labels.to(device)\n        preds=model(img.to(device))\n        loss=loss_fn(preds,labels)\n        total_loss+=loss.item()\n        loss.backward()\n        optimizer.step()\n    model.eval()\n    total_dice=0\n    total_samples=0\n    with torch.no_grad():\n        for batch in validation_loader:\n            img,labels=batch\n            batch_len=len(img)\n            labels=labels\n            preds=model(img.to(device))\n            prediction_mask=(preds.detach().to(\"cpu\")>0.5).flatten(1)\n            dice_score=dice_coeff(prediction_mask,labels.squeeze(1))\n            total_dice+=dice_score*batch_len\n            total_samples+=batch_len\n        epoch_dice=total_dice/total_samples\n        if epoch_dice>best_val or epoch==0:\n            torch.save({\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'loss': total_loss,\n            }, PATH)\n    print(f\"Epoch: {epoch+1}, Train loss:{total_loss}, Validation score:{epoch_dice}, Total Dice:{total_dice}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Best Validation Score over 10 epochs\n\n| No. Training Samples | Best Validation Score |\n| --- | --- | \n| 1000 | 0.17877590656280518 |\n| 5000 | 0.45024675130844116 |\n| 10000 | 0.5069330930709839 |","metadata":{}},{"cell_type":"code","source":"del model\ndel optimizer","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:51:36.515559Z","iopub.status.idle":"2023-06-16T08:51:36.517797Z","shell.execute_reply.started":"2023-06-16T08:51:36.517546Z","shell.execute_reply":"2023-06-16T08:51:36.517576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Unet taking 3 frames as input","metadata":{}},{"cell_type":"code","source":"model_3frame=uNet(9,1,device=device)\noptimizer = optim.Adam(model_3frame.parameters(), lr=lr)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:51:36.519258Z","iopub.status.idle":"2023-06-16T08:51:36.520033Z","shell.execute_reply.started":"2023-06-16T08:51:36.519789Z","shell.execute_reply":"2023-06-16T08:51:36.519812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"random.shuffle(train_ids)\nrandom.shuffle(validation_ids)\ntrain_data=ContrailDataset(train_ids[:train_size],train_path,frames=True)\nvalidation_data=ContrailDataset(validation_ids,validation_path,frames=True)\ntrain_loader=DataLoader(train_data,batch_size,shuffle=True,num_workers=2)\nvalidation_loader=DataLoader(validation_data,batch_size,shuffle=None,num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:51:36.52149Z","iopub.status.idle":"2023-06-16T08:51:36.522271Z","shell.execute_reply.started":"2023-06-16T08:51:36.522031Z","shell.execute_reply":"2023-06-16T08:51:36.522054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_val=0\nPATH=f'{op_path}/best_uNet_3frames_{train_size}.pt'\nfor epoch in range(epochs):\n    total_loss=0\n    model_3frame.train()\n    for batch in train_loader:\n        optimizer.zero_grad()\n        prev,central,future,labels=batch\n        img=torch.cat([prev,central,future],dim=1)\n        labels=labels.to(device)\n        preds=model_3frame(img.to(device))\n        loss=loss_fn(preds,labels)\n        total_loss+=loss.item()\n        loss.backward()\n        optimizer.step()\n    model_3frame.eval()\n    total_dice=0\n    total_samples=0\n    with torch.no_grad():\n        for batch in validation_loader:\n            prev,central,future,labels=batch\n            img=torch.cat([prev,central,future],dim=1)\n            batch_len=len(img)\n            labels=labels\n            preds=model_3frame(img.to(device))\n            prediction_mask=(preds.detach().to(\"cpu\")>0.5).flatten(1)\n            dice_score=dice_coeff(prediction_mask,labels.squeeze(1))\n            total_dice+=dice_score*batch_len\n            total_samples+=batch_len\n        epoch_dice=total_dice/total_samples\n        if epoch_dice>best_val or epoch==0:\n            torch.save({\n            'epoch': epoch,\n            'model_state_dict': model_3frame.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'loss': total_loss,\n            }, PATH)\n    print(f\"Epoch: {epoch+1}, Train loss:{total_loss}, Validation score:{epoch_dice}, Total Dice:{total_dice}\")","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:51:36.523788Z","iopub.status.idle":"2023-06-16T08:51:36.524579Z","shell.execute_reply.started":"2023-06-16T08:51:36.524308Z","shell.execute_reply":"2023-06-16T08:51:36.524331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Best Validation Score over 10 epochs\n\n| No. Training Samples | Best Validation Score |\n| --- | --- | \n| 1000 | 0.2078145146369934 |\n| 5000 | 0.41773277521133423 |\n| 10000 | 0.4545910358428955 |","metadata":{}},{"cell_type":"code","source":"del model_3frame\ndel optimizer","metadata":{"execution":{"iopub.status.busy":"2023-06-16T08:51:36.525974Z","iopub.status.idle":"2023-06-16T08:51:36.526747Z","shell.execute_reply.started":"2023-06-16T08:51:36.526495Z","shell.execute_reply":"2023-06-16T08:51:36.526517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## TCN augmented Unet","metadata":{}},{"cell_type":"code","source":"tcn_unet_central=TCN_uNet_Central(3,1,3,3,device=\"cuda\")","metadata":{"execution":{"iopub.status.busy":"2023-06-16T15:10:49.799858Z","iopub.execute_input":"2023-06-16T15:10:49.800224Z","iopub.status.idle":"2023-06-16T15:10:52.779133Z","shell.execute_reply.started":"2023-06-16T15:10:49.800196Z","shell.execute_reply":"2023-06-16T15:10:52.778144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"optimizer = optim.Adam(tcn_unet_central.parameters(), lr=lr)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"random.shuffle(train_ids)\nrandom.shuffle(validation_ids)\ntrain_data=ContrailDataset(train_ids[:10000],train_path,frames=True)\nvalidation_data=ContrailDataset(validation_ids,validation_path,frames=True)\ntrain_loader=DataLoader(train_data,8,shuffle=True,num_workers=2)\nvalidation_loader=DataLoader(validation_data,8,shuffle=None,num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T15:10:00.998636Z","iopub.execute_input":"2023-06-16T15:10:00.9993Z","iopub.status.idle":"2023-06-16T15:10:01.026448Z","shell.execute_reply.started":"2023-06-16T15:10:00.999268Z","shell.execute_reply":"2023-06-16T15:10:01.025493Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_val=0\nPATH=f'{op_path}/best_tcn_uNet_central_10000.pt'\nfor epoch in range(epochs):\n    total_loss=0\n    tcn_unet_central.train()\n    for batch in train_loader:\n        prev,central,future,labels=batch\n        preds=tcn_unet_central(prev.to(device),central.to(device),future.to(device))\n        labels=labels.to(device)\n        loss=loss_fn(preds,labels)\n        total_loss+=loss.item()\n        loss.backward()\n        optimizer.step()\n    total_dice=0\n    total_samples=0\n    tcn_unet_central.eval()\n    with torch.no_grad():\n        for batch in validation_loader:\n            prev,central,future,labels=batch\n            preds=tcn_unet_central(prev.to(device),central.to(device),future.to(device))\n            labels=labels\n            batch_len=len(prev)\n            prediction_mask=(preds.detach().to(\"cpu\")>0.5).flatten(1)\n            dice_score=dice_coeff(prediction_mask,labels.squeeze(1))\n            total_dice+=dice_score*batch_len\n            total_samples+=batch_len\n        epoch_dice=total_dice/total_samples\n        if epoch_dice>best_val or epoch==0:\n            torch.save({\n            'epoch': epoch,\n            'model_state_dict': tcn_unet_central.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'loss': total_loss,\n            }, PATH)\n    print(f\"Epoch: {epoch+1}, Train loss:{total_loss}, Validation score:{epoch_dice}, Total Dice:{total_dice}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Best Validation Score over 10 epochs\n\n| No. Training Samples | Best Validation Score |\n| --- | --- | \n| 1000 | 0.029899990186095238 |\n| 5000 | 0.05545173957943916 |\n| 10000 | 0.0943593829870224 |","metadata":{}},{"cell_type":"code","source":"del tcn_unet_central\ndel optimizer","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Comparing Improvement in Performance with increase in Training Data","metadata":{}},{"cell_type":"code","source":"statistics = {\n    'Training Samples': [1000, 5000, 10000],\n    'unet_base': [0.17877590656280518, 0.45024675130844116, 0.5069330930709839],\n    'unet_3frame': [0.2078145146369934, 0.41773277521133423, 0.4545910358428955],\n    'tcn_unet': [0.029899990186095238,0.05545173957943916, 0.0943593829870224]\n}","metadata":{"execution":{"iopub.status.busy":"2023-06-16T16:47:19.34769Z","iopub.execute_input":"2023-06-16T16:47:19.348129Z","iopub.status.idle":"2023-06-16T16:47:19.354378Z","shell.execute_reply.started":"2023-06-16T16:47:19.348099Z","shell.execute_reply":"2023-06-16T16:47:19.353251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x=statistics['Training Samples']","metadata":{"execution":{"iopub.status.busy":"2023-06-16T16:54:24.081683Z","iopub.execute_input":"2023-06-16T16:54:24.082584Z","iopub.status.idle":"2023-06-16T16:54:24.087537Z","shell.execute_reply.started":"2023-06-16T16:54:24.082551Z","shell.execute_reply":"2023-06-16T16:54:24.086256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot the scores for each model\nplt.plot(x,statistics['unet_base'], label='unet_base')\nplt.plot(x, statistics['unet_3frame'], label='unet_3frame')\nplt.plot(x, statistics['tcn_unet'], label='tcn_unet')\n\n# Set labels and title\nplt.xlabel('Training Samples')\nplt.ylabel('Validation Scores')\nplt.title('Model Scores')\n# Set the x-axis tick labels\nplt.xticks(x, ['1000', '5000','10000'])\nfor num_samples in x:\n    plt.axvline(x=num_samples, linestyle='dotted', color='gray')\n# Add a legend\nplt.legend()\n\n# Show the plot\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:02:22.451263Z","iopub.execute_input":"2023-06-16T17:02:22.451627Z","iopub.status.idle":"2023-06-16T17:02:22.739643Z","shell.execute_reply.started":"2023-06-16T17:02:22.451597Z","shell.execute_reply":"2023-06-16T17:02:22.738655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We can clearly see that our proposed model struggles to perform at the same level as a basic Unet. \nThis is possibly because the additional information provided by the extra frames passing through the TCN doesn't provide significant information to outweigh the additional noise added by the increase in number of parameters.\nThe higher superior performance of the U-net with 3 frames as input can possibly be due to better random initialization.","metadata":{}},{"cell_type":"markdown","source":"## Model Predictions","metadata":{}},{"cell_type":"code","source":"random.shuffle(train_ids)\ntrain_data=ContrailDataset(train_ids[:100],train_path,frames=True)\ntrain_loader=DataLoader(train_data,batch_size,shuffle=True,num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:50:56.218272Z","iopub.execute_input":"2023-06-16T17:50:56.218641Z","iopub.status.idle":"2023-06-16T17:50:56.244018Z","shell.execute_reply.started":"2023-06-16T17:50:56.218612Z","shell.execute_reply":"2023-06-16T17:50:56.242966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Ground Truth and Image sample","metadata":{}},{"cell_type":"code","source":"for batch in train_loader:\n    prev,central,future,labels=batch\n    break","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Visualization from https://www.kaggle.com/code/inversion/visualizing-contrails\nplt.figure(figsize=(18, 6))\nax = plt.subplot(1, 3, 1)\nax.imshow(central[0].permute(1,2,0))\nax.set_title('False color image')\n\nax = plt.subplot(1, 3, 2)\nax.imshow(labels[0].permute(1,2,0), interpolation='none')\nax.set_title('Ground truth contrail mask')\n\nax = plt.subplot(1, 3, 3)\nax.imshow(central[0].permute(1,2,0))\nax.imshow(labels[0].permute(1,2,0), cmap='Reds', alpha=.4, interpolation='none')\nax.set_title('Contrail mask on false color image');","metadata":{"execution":{"iopub.status.busy":"2023-06-16T18:15:27.839267Z","iopub.execute_input":"2023-06-16T18:15:27.839665Z","iopub.status.idle":"2023-06-16T18:15:30.177353Z","shell.execute_reply.started":"2023-06-16T18:15:27.839632Z","shell.execute_reply":"2023-06-16T18:15:30.176046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load saved models","metadata":{}},{"cell_type":"code","source":"model_paths='/kaggle/input/trained-contrail-models'\nbase_unet=uNet(3,1,device=device)\nunet_3frame=uNet(9,1,device=device)\ntcn_unet=TCN_uNet_Central(3,1,3,3,device=\"cuda\")","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:55:25.340654Z","iopub.execute_input":"2023-06-16T17:55:25.341145Z","iopub.status.idle":"2023-06-16T17:55:26.201484Z","shell.execute_reply.started":"2023-06-16T17:55:25.341107Z","shell.execute_reply":"2023-06-16T17:55:26.200496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"unet_state=torch.load(f'{model_paths}/best_uNet_10000.pt')\nunet_3frame_state=torch.load(f'{model_paths}/best_uNet_3frames_10000.pt')\ntcn_unet_state=torch.load(f'{model_paths}/best_tcn_uNet_central_10000 (1).pt')\n\nbase_unet.load_state_dict(unet_state['model_state_dict'])\nbase_unet.eval()\nunet_3frame.load_state_dict(unet_3frame_state['model_state_dict'])\nunet_3frame.eval()\ntcn_unet.load_state_dict(tcn_unet_state['model_state_dict'])\ntcn_unet.eval()\nprint(\"Loaded models\")","metadata":{"execution":{"iopub.status.busy":"2023-06-16T17:58:44.980261Z","iopub.execute_input":"2023-06-16T17:58:44.980696Z","iopub.status.idle":"2023-06-16T17:58:45.03542Z","shell.execute_reply.started":"2023-06-16T17:58:44.980663Z","shell.execute_reply":"2023-06-16T17:58:45.034523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Generate prediction masks","metadata":{}},{"cell_type":"code","source":"base_unet_preds=base_unet(central[0].unsqueeze(dim=0).to(device)).detach()\nunet_3frame_preds=unet_3frame(torch.cat((prev,central,future),dim=1)[0].unsqueeze(dim=0).to(device)).detach()\ntcn_unet_preds=tcn_unet(prev[0].unsqueeze(dim=0).to(device),\n                        central[0].unsqueeze(dim=0).to(device),\n                        future[0].unsqueeze(dim=0).to(device)).detach()\nbase_unet_prediction_mask=(base_unet_preds.squeeze(1).to(\"cpu\")>0.5)\nunet_3frame_prediction_mask=(unet_3frame_preds.squeeze(1).to(\"cpu\")>0.5)\ntcn_unet_prediction_mask=(tcn_unet_preds.squeeze(1).to(\"cpu\")>0.5)","metadata":{"execution":{"iopub.status.busy":"2023-06-16T18:15:39.111627Z","iopub.execute_input":"2023-06-16T18:15:39.112595Z","iopub.status.idle":"2023-06-16T18:15:39.175062Z","shell.execute_reply.started":"2023-06-16T18:15:39.112559Z","shell.execute_reply":"2023-06-16T18:15:39.174067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## Visualization from https://www.kaggle.com/code/inversion/visualizing-contrails\nplt.figure(figsize=(12, 10))\n\nax = plt.subplot(2, 2, 1)\nax.imshow(central[0].permute(1,2,0))\nax.set_title('False Color image')\n\nax = plt.subplot(2, 2, 2)\nax.imshow(labels[0].permute(1,2,0))\nax.set_title('Ground truth contrail mask')\n\nax = plt.subplot(2, 3, 4)\nax.imshow(base_unet_prediction_mask.permute(1,2,0), interpolation='none')\nax.set_title('Standard U-net prediction')\n\nax = plt.subplot(2, 3, 5)\nax.imshow(unet_3frame_prediction_mask.permute(1,2,0), interpolation='none')\nax.set_title('3 Frame U-net prediction')\n\nax = plt.subplot(2, 3, 6)\nax.imshow(tcn_unet_prediction_mask.permute(1,2,0), interpolation='none')\nax.set_title('TCN+U-net prediction')\n\n# Adjust spacing between subplots\nplt.subplots_adjust(wspace=0.2, hspace=0.01)\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-06-16T18:20:02.260867Z","iopub.execute_input":"2023-06-16T18:20:02.261236Z","iopub.status.idle":"2023-06-16T18:20:03.907515Z","shell.execute_reply.started":"2023-06-16T18:20:02.261207Z","shell.execute_reply":"2023-06-16T18:20:03.906676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"As we can see, both versions of the standard U-net perform \"visibly\" decently. The proposed TCN augmented U-net on the other hand manages to find the general location of the contrail but fails to generate precise outputs along with generating false positive.","metadata":{}}]}