{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\nimport numpy as np\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.nn.functional as F\nimport torchvision.transforms as transforms\nfrom torch.utils.data import Dataset, DataLoader","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-09-06T14:36:52.604125Z","iopub.execute_input":"2023-09-06T14:36:52.604502Z","iopub.status.idle":"2023-09-06T14:36:52.61241Z","shell.execute_reply.started":"2023-09-06T14:36:52.604469Z","shell.execute_reply":"2023-09-06T14:36:52.611307Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    thr = 0.46\n    encoder_path = \"/kaggle/input/gn-u-net-ae/ae (1).pth\"\n    u_net_path = \"/kaggle/input/gn-unets/u-net_22.pth\"","metadata":{"execution":{"iopub.status.busy":"2023-09-06T14:36:52.614274Z","iopub.execute_input":"2023-09-06T14:36:52.615278Z","iopub.status.idle":"2023-09-06T14:36:52.626725Z","shell.execute_reply.started":"2023-09-06T14:36:52.615246Z","shell.execute_reply":"2023-09-06T14:36:52.625751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Conv(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(Conv, self).__init__()\n        self.layers = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, 3, bias=False),\n            nn.GroupNorm(8, out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, 3, bias=False),\n            nn.GroupNorm(8, out_channels),\n            nn.ReLU(inplace=True)\n        )\n        \n    def forward(self, x):\n        return self.layers(x)\n\n\nclass AE(nn.Module):\n    def __init__(self, n_channels, n_classes):\n        super(AE, self).__init__()\n        self.conv0 = Conv(n_channels, 64)\n        self.conv1 = Conv(64, 128)\n        self.conv2 = Conv(128, 256)\n        self.conv3 = Conv(256, 512)\n        self.conv4 = Conv(512, 1024)\n        self.conv5 = Conv(512, 512)\n        self.conv6 = Conv(256, 256)\n        self.conv7 = Conv(128, 128)\n        self.conv8 = Conv(64, 64)\n        self.maxpool = nn.MaxPool2d(2)\n        self.convT0 = nn.ConvTranspose2d(1024, 512, kernel_size=2, stride=2)\n        self.convT1 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2)\n        self.convT2 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2)\n        self.convT3 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)\n        self.outconv = nn.Conv2d(64, n_classes, 1)\n        self.sigmoid =  nn.Sigmoid()\n        \n        \nclass U_Net(nn.Module):\n    def __init__(self, n_channels, n_classes):\n        super(U_Net, self).__init__()\n        self.ae = AE(3, 3)\n        self.conv5 = Conv(1024, 512)\n        self.conv6 = Conv(512, 256)\n        self.conv7 = Conv(256, 128)\n        self.conv8 = Conv(128, 64)\n        self.convT0 = nn.ConvTranspose2d(1024, 512, kernel_size=2, stride=2)\n        self.convT1 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2)\n        self.convT2 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2)\n        self.convT3 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2)\n        self.ae.load_state_dict(torch.load(CFG.encoder_path))\n        self.outconv = nn.Conv2d(64, n_classes, 1)\n        self.sigmoid = nn.Sigmoid()\n    \n    def forward(self, x):\n        # contracting path\n        x0 = self.ae.conv0(x)\n        x1 = self.ae.conv1(self.ae.maxpool(x0))\n        x2 = self.ae.conv2(self.ae.maxpool(x1))\n        x3 = self.ae.conv3(self.ae.maxpool(x2))\n        x = self.ae.conv4(self.ae.maxpool(x3))\n        # expanding path\n        x = self.conv5(self.concat(self.convT0(x), x3))\n        x = self.conv6(self.concat(self.convT1(x), x2))\n        x = self.conv7(self.concat(self.convT2(x), x1))\n        x = self.conv8(self.concat(self.convT3(x), x0))\n        x = self.outconv(x)\n        x = self.sigmoid(x)\n        x = x[:, :, 2:-2, 2:-2]\n        return x\n    \n    @staticmethod\n    def concat(x_e, x_c):\n        diff_h = x_c.size()[2] - x_e.size()[2]\n        diff_w = x_c.size()[3] - x_e.size()[3]\n        x_c = x_c[:, :, diff_h//2:-(diff_h - diff_h//2), diff_w//2:-(diff_w - diff_w//2)]\n        return torch.cat([x_c, x_e], dim=1)\n        ","metadata":{"execution":{"iopub.status.busy":"2023-09-06T14:36:52.654892Z","iopub.execute_input":"2023-09-06T14:36:52.655305Z","iopub.status.idle":"2023-09-06T14:36:52.677274Z","shell.execute_reply.started":"2023-09-06T14:36:52.655281Z","shell.execute_reply":"2023-09-06T14:36:52.676385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = U_Net(3, 1)\nmodel.load_state_dict(torch.load(CFG.u_net_path))\nmodel.to(\"cuda\")\nmodel.eval()","metadata":{"execution":{"iopub.status.busy":"2023-09-06T14:36:52.679329Z","iopub.execute_input":"2023-09-06T14:36:52.67965Z","iopub.status.idle":"2023-09-06T14:36:53.387761Z","shell.execute_reply.started":"2023-09-06T14:36:52.679621Z","shell.execute_reply":"2023-09-06T14:36:53.386612Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def false_color(band11, band14, band15):\n    def normalize(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 = normalize(band15 - band14, _TDIFF_BOUNDS)\n    g = normalize(band14 - band11, _CLOUD_TOP_TDIFF_BOUNDS)\n    b = normalize(band14, _T11_BOUNDS)\n    return np.clip(np.stack([r, g, b], axis=2), 0, 1)\n\n\nclass ICRGWDataset(Dataset):\n    def __init__(self, tar_path, ids, padding_size):\n        self.tar_path = tar_path\n        self.ids = ids\n        self.padding_size = padding_size\n        self.normalize = transforms.Normalize(\n            mean=[0.2189, 0.5505, 0.5626],\n            std=[0.1803, 0.2279, 0.2927]\n        )\n    def __len__(self):\n        return len(self.ids)\n    def __getitem__(self, idx):\n        N_TIMES_BEFORE = 4\n        sample_path = f\"{tar_path}/{self.ids[idx]}\"\n        band11 = np.load(f\"{sample_path}/band_11.npy\")[..., N_TIMES_BEFORE]\n        band14 = np.load(f\"{sample_path}/band_14.npy\")[..., N_TIMES_BEFORE]\n        band15 = np.load(f\"{sample_path}/band_15.npy\")[..., N_TIMES_BEFORE]\n        image = false_color(band11, band14, band15)\n        image = torch.Tensor(image)\n        image = image.permute(2, 0, 1)\n        o_image = image.clone()\n        image = self.normalize(image)\n        padding_size = self.padding_size\n        image = F.pad(image, (padding_size, padding_size, padding_size, padding_size), mode='reflect')\n        label = np.load(f\"{sample_path}/human_pixel_masks.npy\")\n        label = torch.Tensor(label)\n        label = label.permute(2, 0, 1)\n        return o_image, label, image\n    \n    \ntar_path = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/validation\"\nids = os.listdir(tar_path)\ndataloader = DataLoader(ICRGWDataset(tar_path, ids, 100), 1, shuffle=False, num_workers=1)","metadata":{"execution":{"iopub.status.busy":"2023-09-06T14:36:53.389573Z","iopub.execute_input":"2023-09-06T14:36:53.38999Z","iopub.status.idle":"2023-09-06T14:36:53.407151Z","shell.execute_reply.started":"2023-09-06T14:36:53.389953Z","shell.execute_reply":"2023-09-06T14:36:53.405929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot(y_p, y_t, label, thr):\n    y_p, y_t, label = y_p[0], y_t[0], label[0]\n    y_t = y_t.permute(1, 2, 0)\n    y_p = torch.where(y_p > thr, torch.tensor(1), torch.tensor(0))\n    y_p, y_t = y_p.squeeze().cpu().numpy(), y_t.squeeze().cpu().numpy()\n    label = label.squeeze().cpu().numpy()\n    plt.figure(figsize=(5, 3))\n    ax = plt.subplot(1, 3, 1)\n    ax.imshow(label, interpolation='none')\n    ax.set_title('Label')\n    ax = plt.subplot(1, 3, 2)\n    ax.imshow(y_p, interpolation='none')\n    ax.set_title('Pred')\n    ax = plt.subplot(1, 3, 3)\n    ax.imshow(y_t)\n    ax.set_title('Orgin')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-09-06T14:43:09.30855Z","iopub.execute_input":"2023-09-06T14:43:09.308929Z","iopub.status.idle":"2023-09-06T14:43:09.320275Z","shell.execute_reply.started":"2023-09-06T14:43:09.308877Z","shell.execute_reply":"2023-09-06T14:43:09.31927Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for o_image, label, image in tqdm(dataloader):\n    image = image.to(\"cuda\")\n    pred = model(image)\n    pred = torch.where(pred > CFG.thr, torch.tensor(1), torch.tensor(0))\n    plot(pred, o_image, label, CFG.thr)","metadata":{"execution":{"iopub.status.busy":"2023-09-06T14:43:09.436349Z","iopub.execute_input":"2023-09-06T14:43:09.436631Z","iopub.status.idle":"2023-09-06T14:43:22.018343Z","shell.execute_reply.started":"2023-09-06T14:43:09.436608Z","shell.execute_reply":"2023-09-06T14:43:22.016827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}],"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"}}