{"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 = [\n    \"../input/pytorch-image-library/pytorch-image-models-master/pytorch-image-models-master\",\n]\nimport sys;\n\nfor pth in package_paths:\n    sys.path.append(pth)\n\nimport timm","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:47.751414Z","iopub.execute_input":"2021-05-24T18:01:47.751765Z","iopub.status.idle":"2021-05-24T18:01:47.757181Z","shell.execute_reply.started":"2021-05-24T18:01:47.751734Z","shell.execute_reply":"2021-05-24T18:01:47.756298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport cv2\nimport torch\nimport torch.nn as nn\nimport albumentations as A\nimport pytorch_lightning as pl\nimport matplotlib.pyplot as plt\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom albumentations.core.composition import Compose, OneOf\nfrom albumentations.pytorch import ToTensorV2\n\nfrom pytorch_lightning import Trainer, seed_everything\nfrom pytorch_lightning import Callback\nfrom pytorch_lightning.loggers import CSVLogger\nfrom pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\nfrom sklearn.model_selection import StratifiedKFold","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:49.426916Z","iopub.execute_input":"2021-05-24T18:01:49.427251Z","iopub.status.idle":"2021-05-24T18:01:49.433453Z","shell.execute_reply.started":"2021-05-24T18:01:49.427223Z","shell.execute_reply":"2021-05-24T18:01:49.432332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"PyTorch Lightning version: {pl.__version__}\")","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:49.951135Z","iopub.execute_input":"2021-05-24T18:01:49.951473Z","iopub.status.idle":"2021-05-24T18:01:49.956725Z","shell.execute_reply.started":"2021-05-24T18:01:49.951444Z","shell.execute_reply":"2021-05-24T18:01:49.955902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"DEBUG = False\n\nclass CFG:\n    seed = 42\n    model_name = 'tf_efficientnet_b5_ns'\n    pretrained = True\n    img_size = 512\n    num_classes = 6\n    lr = 1e-4\n    max_lr = 1e-3\n    pct_start = 0.3\n    div_factor = 1.0e+3\n    final_div_factor = 1.0e+3\n    num_epochs = 5\n    batch_size = 16\n    accum = 1\n    precision = 16\n    n_fold = 5\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:50.354132Z","iopub.execute_input":"2021-05-24T18:01:50.354447Z","iopub.status.idle":"2021-05-24T18:01:50.360454Z","shell.execute_reply.started":"2021-05-24T18:01:50.354419Z","shell.execute_reply":"2021-05-24T18:01:50.359339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"PATH = \"../input/plant-pathology-2021-fgvc8/\"\nTEST_DIR = PATH + 'test_images/'","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:50.553290Z","iopub.execute_input":"2021-05-24T18:01:50.553589Z","iopub.status.idle":"2021-05-24T18:01:50.557491Z","shell.execute_reply.started":"2021-05-24T18:01:50.553561Z","shell.execute_reply":"2021-05-24T18:01:50.556641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:50.779035Z","iopub.execute_input":"2021-05-24T18:01:50.779357Z","iopub.status.idle":"2021-05-24T18:01:50.785824Z","shell.execute_reply.started":"2021-05-24T18:01:50.779329Z","shell.execute_reply":"2021-05-24T18:01:50.784963Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Dataset","metadata":{}},{"cell_type":"code","source":"class PlantDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.image_id = df['image'].values\n        self.labels = df.iloc[:, 2:].values\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.labels)\n\n    def __getitem__(self, idx):\n        image_id = self.image_id[idx]\n        label = torch.tensor(self.labels[idx].astype('int8'), dtype=torch.float32)\n        \n        image_path = TEST_DIR + image_id\n        image = cv2.imread(image_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        augmented = self.transform(image=image)\n        image = augmented['image']\n        return {'image':image, 'target': label}","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:51.159749Z","iopub.execute_input":"2021-05-24T18:01:51.160093Z","iopub.status.idle":"2021-05-24T18:01:51.167670Z","shell.execute_reply.started":"2021-05-24T18:01:51.160064Z","shell.execute_reply":"2021-05-24T18:01:51.166828Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transform(phase: str):\n    if phase == 'train':\n        return Compose([\n            A.RandomResizedCrop(height=CFG.img_size, width=CFG.img_size),\n            A.Flip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.ShiftScaleRotate(p=0.5),\n            A.HueSaturationValue(p=0.5),\n            A.OneOf([\n                A.RandomBrightnessContrast(p=0.5),\n                A.RandomGamma(p=0.5),\n            ], p=0.5),\n            A.OneOf([\n                A.Blur(p=0.1),\n                A.GaussianBlur(p=0.1),\n                A.MotionBlur(p=0.1),\n            ], p=0.1),\n            A.OneOf([\n                A.GaussNoise(p=0.1),\n                A.ISONoise(p=0.1),\n                A.GridDropout(ratio=0.5, p=0.2),\n                A.CoarseDropout(max_holes=16, min_holes=8, max_height=16, max_width=16, min_height=8, min_width=8, p=0.2)\n            ], p=0.2),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n    else:\n        return Compose([\n            A.Resize(height=CFG.img_size, width=CFG.img_size),\n            A.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:51.641076Z","iopub.execute_input":"2021-05-24T18:01:51.641389Z","iopub.status.idle":"2021-05-24T18:01:51.651530Z","shell.execute_reply.started":"2021-05-24T18:01:51.641358Z","shell.execute_reply":"2021-05-24T18:01:51.650547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Model","metadata":{}},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    def __init__(self, model_name='tf_efficientnet_b0_ns', pretrained=True):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        in_features = self.model.get_classifier().in_features\n        fc = nn.Sequential(\n            nn.Linear(in_features, in_features),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.5),\n            nn.Linear(in_features, CFG.num_classes)\n        )\n        \n        if 'efficientnet' in model_name:\n            self.model.classifier = fc\n        elif 'res' in model_name:\n            self.model.fc = fc\n\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:52.025647Z","iopub.execute_input":"2021-05-24T18:01:52.026007Z","iopub.status.idle":"2021-05-24T18:01:52.033003Z","shell.execute_reply.started":"2021-05-24T18:01:52.025974Z","shell.execute_reply":"2021-05-24T18:01:52.032054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"code","source":"checkpoint = \"../input/pp-2021-efficientnet-model/last.ckpt\"\nmodel = CustomModel(model_name=CFG.model_name, pretrained=False)\nmodel.load_state_dict(torch.load(checkpoint)['state_dict'])","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:52.407347Z","iopub.execute_input":"2021-05-24T18:01:52.407673Z","iopub.status.idle":"2021-05-24T18:01:53.324494Z","shell.execute_reply.started":"2021-05-24T18:01:52.407635Z","shell.execute_reply":"2021-05-24T18:01:53.323807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.read_csv(PATH + \"sample_submission.csv\")\nsub","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:53.326205Z","iopub.execute_input":"2021-05-24T18:01:53.326472Z","iopub.status.idle":"2021-05-24T18:01:53.341675Z","shell.execute_reply.started":"2021-05-24T18:01:53.326446Z","shell.execute_reply":"2021-05-24T18:01:53.340862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_labels = ['healthy', 'scab', 'frog_eye_leaf_spot', 'complex', 'rust', 'powdery_mildew']","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:53.342906Z","iopub.execute_input":"2021-05-24T18:01:53.343399Z","iopub.status.idle":"2021-05-24T18:01:53.349733Z","shell.execute_reply.started":"2021-05-24T18:01:53.343361Z","shell.execute_reply":"2021-05-24T18:01:53.349039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tmp = pd.DataFrame(np.zeros([len(sub), len(new_labels)]), columns=new_labels)\nsub = pd.concat([sub, tmp], axis=1)\nsub","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:53.350886Z","iopub.execute_input":"2021-05-24T18:01:53.351256Z","iopub.status.idle":"2021-05-24T18:01:53.368391Z","shell.execute_reply.started":"2021-05-24T18:01:53.351219Z","shell.execute_reply":"2021-05-24T18:01:53.367611Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = PlantDataset(sub, get_transform('valid'))\ntest_loader = DataLoader(test_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:53.370794Z","iopub.execute_input":"2021-05-24T18:01:53.371157Z","iopub.status.idle":"2021-05-24T18:01:53.376487Z","shell.execute_reply.started":"2021-05-24T18:01:53.371123Z","shell.execute_reply":"2021-05-24T18:01:53.375018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.cuda()\nmodel.eval()\n\nsigmoid = nn.Sigmoid()\n\npredictions = []\nfor batch in test_loader:\n    image = batch['image'].cuda()\n    with torch.no_grad():\n        outputs = model(image)\n        preds = outputs.detach().cpu()\n        # The probability of 0.5 or more is considered positive.\n        predictions.append(sigmoid(preds).numpy() > 0.55)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:53.956289Z","iopub.execute_input":"2021-05-24T18:01:53.956627Z","iopub.status.idle":"2021-05-24T18:01:54.688217Z","shell.execute_reply.started":"2021-05-24T18:01:53.956588Z","shell.execute_reply":"2021-05-24T18:01:54.687272Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = pd.DataFrame(np.concatenate(predictions).astype(np.int), columns=new_labels)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:54.690052Z","iopub.execute_input":"2021-05-24T18:01:54.690431Z","iopub.status.idle":"2021-05-24T18:01:54.695531Z","shell.execute_reply.started":"2021-05-24T18:01:54.690402Z","shell.execute_reply":"2021-05-24T18:01:54.694390Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.iloc[:, 2:] = predictions\nsub","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:54.773826Z","iopub.execute_input":"2021-05-24T18:01:54.774102Z","iopub.status.idle":"2021-05-24T18:01:54.791623Z","shell.execute_reply.started":"2021-05-24T18:01:54.774077Z","shell.execute_reply":"2021-05-24T18:01:54.790629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = []\nfor i, row in sub.iloc[:, 2:].iterrows():\n    if ((row['healthy'] == 1) or row.sum() == 0):\n        tmp = 'healthy'\n    else:\n        tmp = ' '.join(np.array(new_labels)[row==row.max()])\n    labels.append(tmp)","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:55.127113Z","iopub.execute_input":"2021-05-24T18:01:55.127426Z","iopub.status.idle":"2021-05-24T18:01:55.138028Z","shell.execute_reply.started":"2021-05-24T18:01:55.127396Z","shell.execute_reply":"2021-05-24T18:01:55.137241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub['labels'] = labels\nsub[['image', 'labels']].to_csv('submission.csv', index=False)\nsub","metadata":{"execution":{"iopub.status.busy":"2021-05-24T18:01:56.716915Z","iopub.execute_input":"2021-05-24T18:01:56.717295Z","iopub.status.idle":"2021-05-24T18:01:56.740292Z","shell.execute_reply.started":"2021-05-24T18:01:56.717263Z","shell.execute_reply":"2021-05-24T18:01:56.738439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}