{"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":"!pip install segmentation_models_pytorch","metadata":{"execution":{"iopub.status.busy":"2023-08-08T09:30:53.524847Z","iopub.execute_input":"2023-08-08T09:30:53.52556Z","iopub.status.idle":"2023-08-08T09:31:12.926167Z","shell.execute_reply.started":"2023-08-08T09:30:53.525523Z","shell.execute_reply":"2023-08-08T09:31:12.925022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nfrom tqdm.notebook import tqdm\nimport torch\nimport torch.nn as nn\nimport numpy as np\nimport os \nfrom functools import partial\nimport warnings\nimport torch.nn.functional as F\nimport torchvision.transforms as T\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport segmentation_models_pytorch as smp\n\n\nclass ContrailsDataset(torch.utils.data.Dataset):\n    def __init__(self, df, train=True):\n\n        self.df = df\n        self.trn = train\n        self.normalize_image = T.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225))\n        self.upsample = T.Resize((384,384), interpolation=T.InterpolationMode.BILINEAR)\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        con_path = row.path\n        con = np.load(str(con_path))\n\n        img_raw = con[..., :-1]\n        label = con[..., -1]\n\n        label = torch.tensor(label)\n\n        img = torch.tensor(np.reshape(img_raw, (256, 256, 3))).to(torch.float32).permute(2, 0, 1)\n        #img = self.upsample(img)\n        img = self.normalize_image(img)\n        \n        return img.float(), label.float(), img_raw\n\n    def __len__(self):\n        return len(self.df)\n\n\n#data_path = r\"data\\ashcolor\\ashcolor\" \ndata_path = \"/kaggle/input/contrails-images-ash-color\"\n\ncontrails = os.path.join(data_path, \"contrails/\")\ntrain_path = os.path.join(data_path, \"train_df.csv\")\nvalid_path = os.path.join(data_path, \"valid_df.csv\")\n\ntrain_df = pd.read_csv(train_path)\nvalid_df = pd.read_csv(valid_path)\n\ntrain_df[\"path\"] = contrails + train_df[\"record_id\"].astype(str) + \".npy\"\nvalid_df[\"path\"] = contrails + valid_df[\"record_id\"].astype(str) + \".npy\"\n\n#dset_train = ContrailsDataset(train_df, train=True)\ndset_val = ContrailsDataset(valid_df, train=False)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-08-08T09:31:12.929365Z","iopub.execute_input":"2023-08-08T09:31:12.93057Z","iopub.status.idle":"2023-08-08T09:31:19.567545Z","shell.execute_reply.started":"2023-08-08T09:31:12.930529Z","shell.execute_reply":"2023-08-08T09:31:19.56658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Gathering contrails from the validation set which are not empty, to see the model predictions on this imgs.","metadata":{}},{"cell_type":"code","source":"dloader_val = DataLoader(dset_val, batch_size=32, shuffle=True)\n\n#(b,H,H)\nxs_full = []\nys_full = []\ncount_pixels = []\nraws = []\n\nfor idx, (x,y,raw) in enumerate(dloader_val):\n    mask = y.sum(dim=(1,2))>0    \n    xs_full.append(x[mask])\n    contains_c = y[mask] #(n,256,256)\n    pixels_per_n = contains_c.sum(dim=(1,2)) #(n,)\n    count_pixels.append(pixels_per_n)\n    ys_full.append(y[mask])\n    raws.append(raw[mask])\n\n    if len(xs_full) > 10:\n        break\n\n# order x,y by number of pixels in descending order \ncount_pixels = torch.cat(count_pixels, dim=0)\norder = torch.argsort(count_pixels, descending=True)  \nx_ = torch.cat(xs_full, dim=0)\ny_ = torch.cat(ys_full, dim=0)","metadata":{"execution":{"iopub.status.busy":"2023-08-08T09:31:19.569027Z","iopub.execute_input":"2023-08-08T09:31:19.569374Z","iopub.status.idle":"2023-08-08T09:31:24.856487Z","shell.execute_reply.started":"2023-08-08T09:31:19.569343Z","shell.execute_reply":"2023-08-08T09:31:24.85531Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"raws = np.concatenate(raws)\nx_ = x_[order]\ny_ = y_[order]\nraws = raws[order.numpy()]","metadata":{"execution":{"iopub.status.busy":"2023-08-08T09:31:24.859221Z","iopub.execute_input":"2023-08-08T09:31:24.859895Z","iopub.status.idle":"2023-08-08T09:31:24.967892Z","shell.execute_reply.started":"2023-08-08T09:31:24.859861Z","shell.execute_reply":"2023-08-08T09:31:24.966815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model definition","metadata":{}},{"cell_type":"markdown","source":"My best score on the test set so far was achieved by calculating a weighted sum of the probabilities from two models. Both utilize Unet as the upsampling decoder architecture; however, the first model employs a convolution-based encoder, whereas the second one utilizes a transformer-based encoder. I think that the combination of a transformer and a convolution-based encoder yields a strong model: the transformer pays global attention to the input features, while the convolution-based encoder might be better at detecting small-scale patterns and possesses desirable properties like translation invariance of input objects.\n\nHence, combining these two methods might result in a more robust model than combining two models that both employ convolution-based encoders.\n\nSingle models performance:  \n\nmit-b1 - Unet: 0.652  \nresnet26d - Unet: 0.64  \nresnet26d + mit-b1: 0.663  ","metadata":{}},{"cell_type":"code","source":"model1 = smp.Unet(\n    encoder_name=\"timm-resnest26d\",\n    encoder_weights=None,\n    in_channels=3,\n    classes=1)\n\nmodel2 = smp.Unet(\n    encoder_name=\"mit_b1\",\n    encoder_weights=None,\n    in_channels=3,\n    classes=1)\n\nwei_path1 = \"/kaggle/input/model-weights-contrails/0714_resnet26dUnet_epoch_36.pt\"\nwei_path2 = \"/kaggle/input/model-weights-contrails/0713_segformer_epoch_58.pt\"\n\ndevice = \"cuda\"\n\nmodel1.to(device)\nmodel2.to(device)\nweis1 = torch.load(wei_path1, map_location=device) \nweis2 = torch.load(wei_path2, map_location=device)\nmodel1.load_state_dict(weis1)\nmodel2.load_state_dict(weis2)\nmodel1.eval()\nmodel2.eval()\nprint(\"\")","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-08-08T09:31:24.969326Z","iopub.execute_input":"2023-08-08T09:31:24.969672Z","iopub.status.idle":"2023-08-08T09:31:31.42056Z","shell.execute_reply.started":"2023-08-08T09:31:24.96964Z","shell.execute_reply":"2023-08-08T09:31:31.41944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# import gc\n\n# gc.collect()\n# torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-08-08T09:31:31.421924Z","iopub.execute_input":"2023-08-08T09:31:31.422289Z","iopub.status.idle":"2023-08-08T09:31:31.429192Z","shell.execute_reply.started":"2023-08-08T09:31:31.422253Z","shell.execute_reply":"2023-08-08T09:31:31.42685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x_ = x_[:40,...].to(device)\ny_ = y_[:40,...].float()\n\nout1 = model1(x_)\nout2 = model2(x_)\nprobs1 = torch.sigmoid(out1.squeeze(1))\nprobs2 = torch.sigmoid(out2.squeeze(1))\n\np_m1 = 0.6\nprobs = probs1*(1-p_m1) + probs2*p_m1\nthreshold = 0.35\npredict = (probs > threshold).float()\n\nnp.save(\"gt.npy\", y_.cpu().numpy())\nnp.save(\"predict.npy\", predict.cpu().numpy())","metadata":{"execution":{"iopub.status.busy":"2023-08-08T09:31:31.430823Z","iopub.execute_input":"2023-08-08T09:31:31.4312Z","iopub.status.idle":"2023-08-08T09:31:37.912074Z","shell.execute_reply.started":"2023-08-08T09:31:31.431166Z","shell.execute_reply":"2023-08-08T09:31:37.911077Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predictions visualization","metadata":{}},{"cell_type":"markdown","source":"You can download the clip if you use shift+right click on the video and select open in a new tab.","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.animation as animation\nfrom IPython.display import HTML\n\n# Prepare the figure\nfig, axs = plt.subplots(2, 2, figsize=(14,14))\n\n# Reducing space between plots horizontally and vertically\nplt.subplots_adjust(wspace=0.08, hspace=0.08)\n\n# Super title for the whole figure\nfig.suptitle(\"Contrails validation set predictions\", size=25, y=0.95)\n\n# Function to update the figure\ndef update(i):\n    # Clear the axes\n    for ax in axs.ravel():\n        ax.cla()\n\n    # Display input in the top left\n    axs[0, 0].imshow(raws.astype(np.float32)[i,...])\n    axs[0, 0].set_title(f'Fake RGB, human friendly image {i+1}', size=20)\n\n    # Display ground truth in the top right\n    axs[0, 1].imshow(y_[i,...].cpu().numpy(), alpha=1, cmap=\"gray\")\n    axs[0, 1].set_title(f'Desired output {i+1}', size=20)\n\n    # Display prediction in the bottom left\n    axs[1, 1].imshow(predict[i,...].cpu().numpy(), alpha=1, cmap=\"gray\")\n    axs[1, 1].set_title(f'Mask prediction', size=20)\n\n    # Display probabilities in the bottom right\n    axs[1, 0].imshow(probs.detach().cpu().numpy()[i,...],alpha=1,cmap=\"gray\")\n    axs[1, 0].set_title(f'Probabilities', size=20)\n\n    # Optional: Remove axes for a cleaner look\n    for ax in axs.ravel():\n        ax.axis('off')\n\nani = animation.FuncAnimation(fig, update, frames=range(40), interval=2000)\nplt.close(ani._fig)\n\nHTML(ani.to_html5_video())\n\n","metadata":{"execution":{"iopub.status.busy":"2023-08-08T09:47:23.405172Z","iopub.execute_input":"2023-08-08T09:47:23.405536Z","iopub.status.idle":"2023-08-08T09:47:47.263251Z","shell.execute_reply.started":"2023-08-08T09:47:23.405506Z","shell.execute_reply":"2023-08-08T09:47:47.256811Z"},"trusted":true},"execution_count":null,"outputs":[]}]}