{"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":"import os\nimport gc\nimport numpy as np\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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-05-25T04:29:43.24174Z","iopub.execute_input":"2023-05-25T04:29:43.242249Z","iopub.status.idle":"2023-05-25T04:29:46.892539Z","shell.execute_reply.started":"2023-05-25T04:29:43.242211Z","shell.execute_reply":"2023-05-25T04:29:46.891572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch_size = 8\nlr = 1e-5\nepochs = 10","metadata":{"execution":{"iopub.status.busy":"2023-05-25T04:29:46.894468Z","iopub.execute_input":"2023-05-25T04:29:46.895105Z","iopub.status.idle":"2023-05-25T04:29:46.900037Z","shell.execute_reply.started":"2023-05-25T04:29:46.89505Z","shell.execute_reply":"2023-05-25T04:29:46.899145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_ids(tar_path):\n    ids = []\n    for img_id in os.listdir(tar_path):\n        ids.append(img_id)\n    print(f\"{len(ids)} samples in {tar_path}\")\n    return ids\n\ntar_path = \"/kaggle/input/google-research-identify-contrails-reduce-global-warming/train\"\nids = get_ids(tar_path)","metadata":{"execution":{"iopub.status.busy":"2023-05-25T04:29:46.901462Z","iopub.execute_input":"2023-05-25T04:29:46.902786Z","iopub.status.idle":"2023-05-25T04:29:47.146494Z","shell.execute_reply.started":"2023-05-25T04:29:46.902746Z","shell.execute_reply":"2023-05-25T04:29:47.145514Z"},"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    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        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).to(torch.int64)\n        label = label.permute(2, 0, 1)\n        return image, label\n    \nids_train, ids_valid = train_test_split(ids, test_size=0.1, random_state=42)\nids_train, ids_valid = ids_train[:100], ids_valid[:100]  # for saving a version faster\nprint(f\"TrainSize: {len(ids_train)}, ValidSize: {len(ids_valid)}\")\ntrain_dataset = ICRGWDataset(tar_path, ids_train, 100)\nvalid_dataset = ICRGWDataset(tar_path, ids_valid, 100)\ntrain_dataloader = DataLoader(train_dataset, batch_size, shuffle=True, num_workers=1)\nvalid_dataloader = DataLoader(valid_dataset, 1, shuffle=None, num_workers=1)","metadata":{"execution":{"iopub.status.busy":"2023-05-25T04:29:47.149575Z","iopub.execute_input":"2023-05-25T04:29:47.150254Z","iopub.status.idle":"2023-05-25T04:29:47.171022Z","shell.execute_reply.started":"2023-05-25T04:29:47.15022Z","shell.execute_reply":"2023-05-25T04:29:47.170128Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"del ids_train, ids_valid, train_dataset, valid_dataset\ngc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-05-25T04:29:47.172405Z","iopub.execute_input":"2023-05-25T04:29:47.173461Z","iopub.status.idle":"2023-05-25T04:29:47.339688Z","shell.execute_reply.started":"2023-05-25T04:29:47.173426Z","shell.execute_reply":"2023-05-25T04:29:47.338775Z"},"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.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, 3, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n    def forward(self, x):\n        return self.layers(x)\n\n\nclass U_Net(nn.Module):\n    def __init__(self, n_channels, n_classes):\n        super(U_Net, 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(1024, 512)\n        self.conv6 = Conv(512, 256)\n        self.conv7 = Conv(256, 128)\n        self.conv8 = Conv(128, 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        \n        \n    def forward(self, x):\n        # contracting path\n        x0 = self.conv0(x)\n        x1 = self.conv1(self.maxpool(x0))\n        x2 = self.conv2(self.maxpool(x1))\n        x3 = self.conv3(self.maxpool(x2))\n        x = self.conv4(self.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        return self.outconv(x)\n\n#         # contracting path\n#         x0 = C.checkpoint(self.conv0, x)\n# #         print(x0.size())\n#         x1 = C.checkpoint(self.conv1, C.checkpoint(self.maxpool, x0))\n# #         print(x1.size())\n#         x2 = C.checkpoint(self.conv2, C.checkpoint(self.maxpool, x1))\n# #         print(x2.size())\n#         x3 = C.checkpoint(self.conv3, C.checkpoint(self.maxpool, x2))\n# #         print(x3.size())\n#         x = C.checkpoint(self.conv4, C.checkpoint(self.maxpool, x3))\n# #         print(x.size())\n#         # expanding path\n#         x = C.checkpoint(self.conv5, self.concat(C.checkpoint(self.convT0, x), x3))\n# #         print(x.size())\n#         x = C.checkpoint(self.conv6, self.concat(C.checkpoint(self.convT1, x), x2))\n# #         print(x.size())\n#         x = C.checkpoint(self.conv7, self.concat(C.checkpoint(self.convT2, x), x1))\n# #         print(x.size())\n#         x = C.checkpoint(self.conv8, self.concat(C.checkpoint(self.convT3, x), x0))\n# #         print(x.size())\n#         return C.checkpoint(self.outconv, x)\n    \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)","metadata":{"execution":{"iopub.status.busy":"2023-05-25T04:29:47.343006Z","iopub.execute_input":"2023-05-25T04:29:47.343388Z","iopub.status.idle":"2023-05-25T04:29:47.362035Z","shell.execute_reply.started":"2023-05-25T04:29:47.343349Z","shell.execute_reply":"2023-05-25T04:29:47.361129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = U_Net(3, 2).to('cuda')\noptimizer = optim.Adam(model.parameters(), lr=lr)","metadata":{"execution":{"iopub.status.busy":"2023-05-25T04:29:47.363059Z","iopub.execute_input":"2023-05-25T04:29:47.363365Z","iopub.status.idle":"2023-05-25T04:29:51.020324Z","shell.execute_reply.started":"2023-05-25T04:29:47.363336Z","shell.execute_reply":"2023-05-25T04:29:51.019374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def ce_loss(y_p, y_t):\n    y_p = y_p[:, :, 2:-2, 2:-2]\n    y_t = y_t.squeeze(dim=1)\n    weight = torch.Tensor([0.57, 4.17]).to('cuda')\n    criterion = nn.CrossEntropyLoss(weight)\n    loss = criterion(y_p, y_t)\n    return loss\n\n# def dice_loss(y_p, y_t, smooth=1e-6):\n#     y_p = y_p[:, :, 2:-2, 2:-2]\n#     y_p = y_p.reshape(y_p.size(0), -1)\n#     y_t = y_t.reshape(y_t.size(0), -1)\n#     i = (y_p * y_t).sum()\n#     return 1 - (2. * i + smooth) / (y_p.sum() + y_t.sum() + smooth)\n\n\ndef dice_score(y_p, y_t, smooth=1e-6):\n    y_p = y_p[:, :, 2:-2, 2:-2]\n    y_p = F.softmax(y_p, dim=1)\n    y_p = torch.argmax(y_p, dim=1, keepdim=True)\n    i = torch.sum(y_p * y_t, dim=(2, 3))\n    u = torch.sum(y_p, dim=(2, 3)) + torch.sum(y_t, dim=(2, 3))\n    score = (2 * i + smooth)/(u + smooth)\n    return torch.mean(score)","metadata":{"execution":{"iopub.status.busy":"2023-05-25T04:29:51.021632Z","iopub.execute_input":"2023-05-25T04:29:51.021996Z","iopub.status.idle":"2023-05-25T04:29:51.030922Z","shell.execute_reply.started":"2023-05-25T04:29:51.021959Z","shell.execute_reply":"2023-05-25T04:29:51.030061Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot(y_p, y_t):\n    if y_p.size()[0] != 1:\n        y_p = y_p[0:1, :, :, :]\n        y_t = y_t[0:1, :, :, :]\n    y_p = y_p[:, :, 2:-2, 2:-2]\n    y_p = F.softmax(y_p, dim=1)\n    y_p = torch.argmax(y_p, dim=1, keepdim=True)\n    y_p, y_t = y_p.squeeze(), y_t.squeeze()\n    y_p, y_t = y_p.cpu().numpy(), y_t.cpu().numpy()\n    print(y_p.shape, y_t.shape)\n    plt.figure(figsize=(5, 3))\n    ax = plt.subplot(1, 2, 1)\n    ax.imshow(y_p, interpolation='none')\n    ax.set_title('Pred')\n    ax = plt.subplot(1, 2, 2)\n    ax.imshow(y_t, interpolation='none')\n    ax.set_title('GT')\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-05-25T04:29:51.032362Z","iopub.execute_input":"2023-05-25T04:29:51.03345Z","iopub.status.idle":"2023-05-25T04:29:51.045897Z","shell.execute_reply.started":"2023-05-25T04:29:51.033417Z","shell.execute_reply":"2023-05-25T04:29:51.044875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"bst_dice = 0\nfor epoch in range(epochs):\n    model.train()\n    bar = tqdm(train_dataloader)\n    tot_loss = 0\n    tot_score = 0\n    count = 0\n    for X, y in bar:\n        X, y = X.to('cuda'), y.to('cuda')\n        pred = model(X)\n        loss = ce_loss(pred, y)\n        loss.backward()\n        optimizer.step()\n        optimizer.zero_grad()\n        tot_loss += loss.item()\n        tot_score += dice_score(pred, y)\n        count += 1\n        bar.set_postfix(TrainLoss=f'{tot_loss/count:.4f}', TrainDice=f'{tot_score/count:.4f}')\n        if count % 200 == 0:\n            plot(pred, y)\n    model.eval()\n    bar = tqdm(valid_dataloader)\n    tot_score = 0\n    count = 0\n    for X, y in bar:\n        X, y = X.to('cuda'), y.to('cuda')\n        pred = model(X)\n        tot_score += dice_score(pred, y)\n        count += 1\n        bar.set_postfix(ValidDice=f'{tot_score/count:.4f}')\n        if count % 200 == 0:\n            plot(pred, y)\n    if tot_score/count > bst_dice:\n        bst_dice = tot_score/count\n        torch.save(model.state_dict(), 'u-net.pth')\n        print(\"current model saved!\")","metadata":{"execution":{"iopub.status.busy":"2023-05-25T04:29:51.048752Z","iopub.execute_input":"2023-05-25T04:29:51.049144Z","iopub.status.idle":"2023-05-25T04:30:08.145301Z","shell.execute_reply.started":"2023-05-25T04:29:51.049113Z","shell.execute_reply":"2023-05-25T04:30:08.143624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(bst_dice)","metadata":{"execution":{"iopub.status.busy":"2023-05-25T04:30:08.146634Z","iopub.status.idle":"2023-05-25T04:30:08.148489Z","shell.execute_reply.started":"2023-05-25T04:30:08.148229Z","shell.execute_reply":"2023-05-25T04:30:08.148258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}