{"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":"package_paths = ['../input/timm-pytorch-image-models/pytorch-image-models-master']\nimport sys\nfor pth in package_paths:\n    sys.path.append(pth)\nimport timm","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:10.586213Z","iopub.execute_input":"2021-08-26T06:16:10.586549Z","iopub.status.idle":"2021-08-26T06:16:13.6422Z","shell.execute_reply.started":"2021-08-26T06:16:10.58648Z","shell.execute_reply":"2021-08-26T06:16:13.641333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pytorch_lightning as pl\nimport torch\nimport pandas as pd\nimport timm\nimport torch.nn as nn\n\nfrom PIL import Image\nfrom sklearn.model_selection import KFold\nfrom torchvision import transforms as tsfm\nfrom torch.utils.data import Dataset, DataLoader\nfrom pytorch_lightning import Trainer, seed_everything\nfrom pytorch_lightning.loggers import CSVLogger\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\nfrom pytorch_lightning.metrics import Metric","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:13.643626Z","iopub.execute_input":"2021-08-26T06:16:13.643947Z","iopub.status.idle":"2021-08-26T06:16:15.198675Z","shell.execute_reply.started":"2021-08-26T06:16:13.643913Z","shell.execute_reply":"2021-08-26T06:16:15.197788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    test_imgs_dir = '../input/plant-pathology-2021-fgvc8/test_images/'\n    submit_csv_path = '../input/plant-pathology-2021-fgvc8/sample_submission.csv'\n    # data info\n    label_num2str = {0: 'powdery_mildew',\n                     1: 'scab',\n                     2: 'complex',\n                     3: 'frog_eye_leaf_spot',\n                     4: 'rust'}\n    \n    label_str2num = {'powdery_mildew': 0,\n                     'scab': 1,\n                     'complex': 2,\n                     'frog_eye_leaf_spot': 3,\n                     'rust': 4}\n    # model info\n    model_name = 'tf_efficientnetv2_s_in21k'\n    # training hyper-parameters\n    seed = 77\n    num_classes = 5\n    n_fold = 6\n    img_size = [512, 512]\n    tta_step = 5","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:15.20047Z","iopub.execute_input":"2021-08-26T06:16:15.200839Z","iopub.status.idle":"2021-08-26T06:16:15.207151Z","shell.execute_reply.started":"2021-08-26T06:16:15.200803Z","shell.execute_reply":"2021-08-26T06:16:15.206102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nDefine dataset class\n\"\"\"\n\nclass PlantDataset(Dataset):\n    def __init__(self, img_dir, img_names: list, labels: list, transform=None):\n        self.img_dir = img_dir\n        self.img_names = img_names\n        self.labels = labels\n        self.transform = transform\n    \n    def __len__(self):\n        return len(self.img_names)\n    \n    def __getitem__(self, idx):\n        img_path = os.path.join(self.img_dir, self.img_names[idx])\n        img = Image.open(img_path).convert('RGB')\n        img_ts = self.transform(img)\n        label_ts = self.labels[idx]\n        return img_ts, label_ts","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:15.208891Z","iopub.execute_input":"2021-08-26T06:16:15.209329Z","iopub.status.idle":"2021-08-26T06:16:15.217958Z","shell.execute_reply.started":"2021-08-26T06:16:15.209266Z","shell.execute_reply":"2021-08-26T06:16:15.216813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nDefine Focal-Loss\n\"\"\"\n\nclass FocalLoss(nn.Module):\n    \"\"\"\n    The focal loss for fighting against class-imbalance\n    \"\"\"\n    def __init__(self, alpha=1, gamma=2):\n        super(FocalLoss, self).__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.epsilon = 1e-12  # prevent training from Nan-loss error \n    \n    def forward(self, logits, target):\n        \"\"\"\n        logits & target should be tensors with shape [batch_size, num_classes]\n        \"\"\"\n        probs = torch.sigmoid(logits)\n        one_subtract_probs = 1.0 - probs\n        # add epsilon\n        probs_new = probs + self.epsilon\n        one_subtract_probs_new = one_subtract_probs + self.epsilon\n        # calculate focal loss\n        log_pt =  target * torch.log(probs_new) + (1.0 - target) * torch.log(one_subtract_probs_new)\n        pt = torch.exp(log_pt)\n        focal_loss = -1.0 * (self.alpha * (1 - pt) ** self.gamma) * log_pt\n        return torch.mean(focal_loss)\n        ","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:15.219405Z","iopub.execute_input":"2021-08-26T06:16:15.219965Z","iopub.status.idle":"2021-08-26T06:16:15.229906Z","shell.execute_reply.started":"2021-08-26T06:16:15.21993Z","shell.execute_reply":"2021-08-26T06:16:15.229125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nDefine F1 score metric\n\"\"\"\nclass MyF1Score(Metric):\n    def __init__(self, cfg, threshold: float=0.5, dist_sync_on_step=False):\n        super().__init__(dist_sync_on_step=dist_sync_on_step)\n        self.cfg = cfg\n        self.threshold = threshold\n        self.add_state(\"tp\", default=torch.tensor(0), dist_reduce_fx=\"sum\")\n        self.add_state(\"fp\", default=torch.tensor(0), dist_reduce_fx=\"sum\")\n        self.add_state(\"fn\", default=torch.tensor(0), dist_reduce_fx=\"sum\")\n\n    def update(self, preds: torch.Tensor, target: torch.Tensor):\n        assert preds.shape == target.shape\n        preds_str_batch = self.num_to_str(preds)\n        target_str_batch = self.num_to_str(target)\n        tp, fp, fn = 0, 0, 0\n        for pred_str_list, target_str_list in zip(preds_str_batch, target_str_batch):\n            for pred_str in pred_str_list:\n                if pred_str in target_str_list:\n                    tp += 1\n                if pred_str not in target_str_list:\n                    fp += 1\n            \n            for target_str in target_str_list:\n                if target_str not in pred_str_list:\n                    fn += 1\n        self.tp += tp\n        self.fp += fp\n        self.fn += fn\n\n    def compute(self):\n        f1 = 2.0 * self.tp / (2.0 * self.tp + self.fn + self.fp)\n        return f1\n    \n    def num_to_str(self, ts: torch.Tensor) -> list:\n        batch_bool_list = (ts > self.threshold).detach().cpu().numpy().tolist()\n        batch_str_list = []\n        for one_sample_bool in batch_bool_list:\n            lb_str_list = [self.cfg.label_num2str[lb_idx] for lb_idx, bool_val in enumerate(one_sample_bool) if bool_val]\n            if len(lb_str_list) == 0:\n                lb_str_list = ['healthy']\n            batch_str_list.append(lb_str_list)\n        return batch_str_list","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:15.231245Z","iopub.execute_input":"2021-08-26T06:16:15.231836Z","iopub.status.idle":"2021-08-26T06:16:15.245264Z","shell.execute_reply.started":"2021-08-26T06:16:15.231794Z","shell.execute_reply":"2021-08-26T06:16:15.244461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nDefine test image transformation\n\"\"\"\n#center_crop_size = [int(CFG.img_size[0] * 0.8), int(CFG.img_size[1] * 0.8)]\n\ntest_transform_normal = tsfm.Compose([\n                                tsfm.Resize(CFG.img_size),\n                                tsfm.ToTensor(),\n                                tsfm.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),])\n\ntest_transform_tta = tsfm.Compose([\n                                tsfm.Resize(CFG.img_size),\n                               # tsfm.RandomApply([tsfm.ColorJitter(0.4, 0.4, 0.4),\n                                 #                tsfm.RandomAffine(degrees=10),], p=0.3),\n                                tsfm.RandomHorizontalFlip(p=0.3),\n                                tsfm.RandomVerticalFlip(p=0.3),\n                                tsfm.ToTensor(),\n                                tsfm.Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),])","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:15.246513Z","iopub.execute_input":"2021-08-26T06:16:15.246905Z","iopub.status.idle":"2021-08-26T06:16:15.256656Z","shell.execute_reply.started":"2021-08-26T06:16:15.246871Z","shell.execute_reply":"2021-08-26T06:16:15.255776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_img_names = os.listdir(CFG.test_imgs_dir)\nif CFG.tta_step > 0:\n    test_dataset = PlantDataset(CFG.test_imgs_dir, test_img_names, range(len(test_img_names)), test_transform_tta)\nelse:\n    test_dataset = PlantDataset(CFG.test_imgs_dir, test_img_names, range(len(test_img_names)), test_transform_normal)\n    \ntest_loader = DataLoader(test_dataset, batch_size=1, num_workers=0, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:15.259973Z","iopub.execute_input":"2021-08-26T06:16:15.260915Z","iopub.status.idle":"2021-08-26T06:16:15.274323Z","shell.execute_reply.started":"2021-08-26T06:16:15.26089Z","shell.execute_reply":"2021-08-26T06:16:15.27336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\"\"\"\nDefine neural network model\n\"\"\"\n\nclass MyNetwork(pl.LightningModule):\n    def __init__(self, cfg):\n        super(MyNetwork, self).__init__()\n        self.cfg = cfg\n        self.model = timm.create_model(cfg.model_name, pretrained=False, num_classes=cfg.num_classes)\n        self.criterion = FocalLoss()\n        self.metric = MyF1Score(cfg)\n       \n    def forward(self, x):\n        return self.model(x)\n    \n    def configure_optimizers(self):\n        self.optimizer = torch.optim.Adam(self.model.parameters(), lr=self.cfg.lr)\n        self.scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(self.optimizer,\n                                                                    T_max=self.cfg.t_max,\n                                                                    eta_min=self.cfg.min_lr,\n                                                                    verbose=True)\n        return {'optimizer': self.optimizer, 'lr_scheduler': self.scheduler}\n    \n    def training_step(self, batch, batch_idx):\n        img_ts, lb_ts = batch\n        pred_ts = self.model(img_ts)\n        loss = self.criterion(pred_ts, lb_ts)\n        score = self.metric(pred_ts, lb_ts)\n        logs = {'train_loss': loss, 'train_f1': score, 'lr': self.optimizer.param_groups[0]['lr']}\n        self.log_dict(logs, on_step=False, on_epoch=True, prog_bar=True, logger=True)\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        img_ts, lb_ts = batch\n        pred_ts = self.model(img_ts)\n        loss = self.criterion(pred_ts, lb_ts)\n        score = self.metric(pred_ts, lb_ts)\n        logs = {'valid_loss': loss, 'valid_f1': score}\n        self.log_dict(logs, on_step=False, on_epoch=True, prog_bar=True, logger=True)\n        return loss","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:15.275796Z","iopub.execute_input":"2021-08-26T06:16:15.276158Z","iopub.status.idle":"2021-08-26T06:16:15.287695Z","shell.execute_reply.started":"2021-08-26T06:16:15.276106Z","shell.execute_reply":"2021-08-26T06:16:15.286807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoints = os.listdir('../input/plant-pathology-2021-weights-321')\ncheckpoints","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:15.289233Z","iopub.execute_input":"2021-08-26T06:16:15.290026Z","iopub.status.idle":"2021-08-26T06:16:15.308922Z","shell.execute_reply.started":"2021-08-26T06:16:15.289988Z","shell.execute_reply":"2021-08-26T06:16:15.308182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models_list = []\nfor checkpoint in checkpoints:\n    ckpt_path =f'../input/plant-pathology-2021-weights-321/'+checkpoint\n    model = MyNetwork.load_from_checkpoint(ckpt_path, cfg=CFG)\n    model.cuda()\n    model.eval()\n    models_list.append(model)","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:15.31Z","iopub.execute_input":"2021-08-26T06:16:15.310358Z","iopub.status.idle":"2021-08-26T06:16:30.796578Z","shell.execute_reply.started":"2021-08-26T06:16:15.310321Z","shell.execute_reply":"2021-08-26T06:16:30.795711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"treshold = np.array([0.5, 0.5, 0.5, 0.5, 0.5])\nsubmit_df = pd.read_csv(CFG.submit_csv_path)\n\nwith torch.no_grad():\n    test_img_idx = 0\n    for img_ts, lb_ts in test_loader:\n        img_ts = img_ts.cuda()\n        n_fold_pred_list = []\n        for model in models_list:\n            pred_ts = torch.sigmoid(model(img_ts)).detach().cpu()\n            n_fold_pred_list.append(pred_ts)\n        pred_np = torch.cat(n_fold_pred_list).mean(dim=0).numpy()*1.3\n        print(np.round(pred_np, decimals=2))\n        pred = (pred_np > treshold).tolist()\n        img_name = test_img_names[test_img_idx]\n        lb_str_list = []\n        for lb_idx, bool_val in enumerate(pred):\n            if bool_val:\n                lb_str = CFG.label_num2str[lb_idx]\n                lb_str_list.append(lb_str)\n        if len(lb_str_list) == 0:\n            final_label = 'healthy'\n        else:\n            final_label = ' '.join(lb_str_list)\n        row_idx = submit_df[submit_df.image == img_name].index.tolist()[0]\n        submit_df.iloc[row_idx, 1] = final_label\n        test_img_idx += 1    ","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:30.797923Z","iopub.execute_input":"2021-08-26T06:16:30.798259Z","iopub.status.idle":"2021-08-26T06:16:33.328318Z","shell.execute_reply.started":"2021-08-26T06:16:30.798223Z","shell.execute_reply":"2021-08-26T06:16:33.326857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit_df","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:33.329607Z","iopub.execute_input":"2021-08-26T06:16:33.329998Z","iopub.status.idle":"2021-08-26T06:16:33.342186Z","shell.execute_reply.started":"2021-08-26T06:16:33.32996Z","shell.execute_reply":"2021-08-26T06:16:33.340942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submit_df.to_csv('./submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2021-08-26T06:16:33.34376Z","iopub.execute_input":"2021-08-26T06:16:33.344175Z","iopub.status.idle":"2021-08-26T06:16:33.35713Z","shell.execute_reply.started":"2021-08-26T06:16:33.344137Z","shell.execute_reply":"2021-08-26T06:16:33.356318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}