{"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 numpy as np\n\nfrom PIL import Image\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms as tsfm\nfrom tqdm.notebook import tqdm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-05-24T18:56:14.210876Z","iopub.execute_input":"2021-05-24T18:56:14.211395Z","iopub.status.idle":"2021-05-24T18:56:14.62675Z","shell.execute_reply.started":"2021-05-24T18:56:14.21128Z","shell.execute_reply":"2021-05-24T18:56:14.625888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    # data path\n    train_csv_path = '../input/plant-pathology-2021-fgvc8/train.csv'\n    train_imgs_dir = '../input/pp2021-train-images-resized/224_square_not_crop'\n    save_path = \"/kaggle/working/images\"\n    seed = 77\n    batch_size = 32\n    num_workers = 2\n    device = torch.device(f'cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2021-05-24T19:06:08.233363Z","iopub.execute_input":"2021-05-24T19:06:08.233726Z","iopub.status.idle":"2021-05-24T19:06:08.240115Z","shell.execute_reply.started":"2021-05-24T19:06:08.233694Z","shell.execute_reply":"2021-05-24T19:06:08.237614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import models\n\ndef convrelu(in_channels, out_channels, kernel, padding):\n    return nn.Sequential(\n        nn.Conv2d(in_channels, out_channels, kernel, padding=padding),\n        nn.ReLU(inplace=True),\n    )\n\nclass ResNetUNet(nn.Module):\n    def __init__(self, n_class):\n        super().__init__()\n\n        self.base_model = models.resnet18(pretrained=False)\n        self.base_layers = list(self.base_model.children())\n\n        self.layer0 = nn.Sequential(*self.base_layers[:3]) # size=(N, 64, x.H/2, x.W/2)\n        self.layer0_1x1 = convrelu(64, 64, 1, 0)\n        self.layer1 = nn.Sequential(*self.base_layers[3:5]) # size=(N, 64, x.H/4, x.W/4)\n        self.layer1_1x1 = convrelu(64, 64, 1, 0)\n        self.layer2 = self.base_layers[5]  # size=(N, 128, x.H/8, x.W/8)\n        self.layer2_1x1 = convrelu(128, 128, 1, 0)\n        self.layer3 = self.base_layers[6]  # size=(N, 256, x.H/16, x.W/16)\n        self.layer3_1x1 = convrelu(256, 256, 1, 0)\n        self.layer4 = self.base_layers[7]  # size=(N, 512, x.H/32, x.W/32)\n        self.layer4_1x1 = convrelu(512, 512, 1, 0)\n\n        self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n\n        self.conv_up3 = convrelu(256 + 512, 512, 3, 1)\n        self.conv_up2 = convrelu(128 + 512, 256, 3, 1)\n        self.conv_up1 = convrelu(64 + 256, 256, 3, 1)\n        self.conv_up0 = convrelu(64 + 256, 128, 3, 1)\n\n        self.conv_original_size0 = convrelu(3, 64, 3, 1)\n        self.conv_original_size1 = convrelu(64, 64, 3, 1)\n        self.conv_original_size2 = convrelu(64 + 128, 64, 3, 1)\n\n        self.conv_last = nn.Conv2d(64, n_class, 1)\n        self.dropout = nn.Dropout(p=0.25)\n        \n    def forward(self, input):\n        x_original = self.conv_original_size0(input)\n        x_original = self.conv_original_size1(x_original)\n\n        layer0 = self.layer0(input)\n        layer1 = self.layer1(layer0)\n        layer2 = self.layer2(layer1)\n        layer3 = self.layer3(layer2)\n        layer4 = self.layer4(layer3)\n\n        layer4 = self.layer4_1x1(layer4)\n        x = self.upsample(layer4)\n        layer3 = self.layer3_1x1(layer3)\n        x = torch.cat([x, layer3], dim=1)\n        x = self.conv_up3(x)\n\n        x = self.upsample(x)\n        layer2 = self.layer2_1x1(layer2)\n        x = torch.cat([x, layer2], dim=1)\n        x = self.conv_up2(x)\n\n        x = self.upsample(x)\n        layer1 = self.layer1_1x1(layer1)\n        x = torch.cat([x, layer1], dim=1)\n        x = self.conv_up1(x)\n\n        x = self.upsample(x)\n        layer0 = self.layer0_1x1(layer0)\n        x = torch.cat([x, layer0], dim=1)\n        x = self.conv_up0(x)\n\n        x = self.upsample(x)\n        x = torch.cat([x, x_original], dim=1)\n        x = self.conv_original_size2(x)\n\n        out = self.conv_last(x)\n\n        return out","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:56:14.662848Z","iopub.execute_input":"2021-05-24T18:56:14.663305Z","iopub.status.idle":"2021-05-24T18:56:14.690746Z","shell.execute_reply.started":"2021-05-24T18:56:14.663259Z","shell.execute_reply":"2021-05-24T18:56:14.689647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seg_model = ResNetUNet(n_class=1)\nseg_model = seg_model.to(CFG.device)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:56:14.692004Z","iopub.execute_input":"2021-05-24T18:56:14.692365Z","iopub.status.idle":"2021-05-24T18:56:16.763624Z","shell.execute_reply.started":"2021-05-24T18:56:14.692328Z","shell.execute_reply":"2021-05-24T18:56:16.76254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seg_model.load_state_dict(torch.load('../input/leafsegweights/best_val_weights.pth'))","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:56:16.766141Z","iopub.execute_input":"2021-05-24T18:56:16.766468Z","iopub.status.idle":"2021-05-24T18:56:16.87183Z","shell.execute_reply.started":"2021-05-24T18:56:16.766432Z","shell.execute_reply":"2021-05-24T18:56:16.870735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_csv = pd.read_csv(CFG.train_csv_path)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:56:16.873437Z","iopub.execute_input":"2021-05-24T18:56:16.8738Z","iopub.status.idle":"2021-05-24T18:56:16.895925Z","shell.execute_reply.started":"2021-05-24T18:56:16.873756Z","shell.execute_reply":"2021-05-24T18:56:16.895267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nDefine dataset class\n\"\"\"\nclass PlantDataset(Dataset):\n    def __init__(self, csv_file, image_loc):\n        self.csv_file = csv_file\n        self.image_loc = image_loc\n\n    def __len__(self):\n        return len(self.csv_file)\n\n    def __getitem__(self, idx):\n        img_name = self.csv_file.iloc[idx, 0]\n        img_path = os.path.join(self.image_loc,\n                                img_name)\n        \n        img = Image.open(img_path).convert('RGB')\n        img = tsfm.ToTensor()(img)\n        return img, img_name","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:56:16.898763Z","iopub.execute_input":"2021-05-24T18:56:16.899018Z","iopub.status.idle":"2021-05-24T18:56:16.906741Z","shell.execute_reply.started":"2021-05-24T18:56:16.898988Z","shell.execute_reply":"2021-05-24T18:56:16.904199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = PlantDataset(data_csv, CFG.train_imgs_dir)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:56:16.90815Z","iopub.execute_input":"2021-05-24T18:56:16.908525Z","iopub.status.idle":"2021-05-24T18:56:16.914672Z","shell.execute_reply.started":"2021-05-24T18:56:16.908474Z","shell.execute_reply":"2021-05-24T18:56:16.913598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds_dataloader = DataLoader(ds, batch_size = CFG.batch_size, shuffle=False, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:56:16.916047Z","iopub.execute_input":"2021-05-24T18:56:16.916478Z","iopub.status.idle":"2021-05-24T18:56:16.92372Z","shell.execute_reply.started":"2021-05-24T18:56:16.916441Z","shell.execute_reply":"2021-05-24T18:56:16.92272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def reverse_transform(inp):\n    inp = inp.numpy().transpose((1, 2, 0))\n    inp = np.clip(inp, 0, 1)\n    inp = (inp * 255).astype(np.uint8)\n    return inp","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:56:16.925017Z","iopub.execute_input":"2021-05-24T18:56:16.925431Z","iopub.status.idle":"2021-05-24T18:56:16.932462Z","shell.execute_reply.started":"2021-05-24T18:56:16.925393Z","shell.execute_reply":"2021-05-24T18:56:16.931825Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_dir(dir_path):\n    if not os.path.exists(dir_path):\n        os.makedirs(dir_path)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T19:05:18.401974Z","iopub.execute_input":"2021-05-24T19:05:18.402294Z","iopub.status.idle":"2021-05-24T19:05:18.407424Z","shell.execute_reply.started":"2021-05-24T19:05:18.402264Z","shell.execute_reply":"2021-05-24T19:05:18.406308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport imageio\n\nseg_model.eval()\nfor i, batch_pair in enumerate(tqdm(ds_dataloader)):\n    img_batch = batch_pair[0].to(CFG.device)\n    img_names = batch_pair[1]\n    \n    seg_batch = seg_model(img_batch)\n    seg_batch = torch.sigmoid(seg_batch)\n    for img, seg, filename in zip(img_batch, seg_batch, img_names):\n        seg_np = seg.cpu().detach()\n        seg_np = reverse_transform(seg_np)\n        seg_np = np.where(seg_np > 220, 1, 0)\n        \n        img_np = img.cpu()\n        img_np = reverse_transform(img_np)\n        prod_img = np.multiply(seg_np, img_np)\n#         plt.figure()\n#         plt.imshow(prod_img)\n        make_dir(CFG.save_path)\n        savename = os.path.join(CFG.save_path, filename)\n        imageio.imwrite(savename, prod_img)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import shutil\nshutil.make_archive(\"leaf-segmented-224\", 'zip', CFG.save_path)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T19:28:20.558126Z","iopub.execute_input":"2021-05-24T19:28:20.558455Z","iopub.status.idle":"2021-05-24T19:28:22.16707Z","shell.execute_reply.started":"2021-05-24T19:28:20.558423Z","shell.execute_reply":"2021-05-24T19:28:22.166306Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf ./images/*.jpg","metadata":{"execution":{"iopub.status.busy":"2021-05-24T19:49:28.361223Z","iopub.execute_input":"2021-05-24T19:49:28.361559Z","iopub.status.idle":"2021-05-24T19:49:29.163200Z","shell.execute_reply.started":"2021-05-24T19:49:28.361525Z","shell.execute_reply":"2021-05-24T19:49:29.162021Z"},"trusted":true},"execution_count":null,"outputs":[]}]}