{"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":"markdown","source":"# Plant 2021 with PyTorch Lightning\nThis notebook uses the models learned in the following notebooks for inference.\n[Training notebook](https://www.kaggle.com/pegasos/plant2021-pytorch-lightning-starter-training)","metadata":{"papermill":{"duration":0.01403,"end_time":"2021-03-19T05:17:29.053979","exception":false,"start_time":"2021-03-19T05:17:29.039949","status":"completed"},"tags":[]}},{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Import","metadata":{"papermill":{"duration":0.015954,"end_time":"2021-03-19T05:17:37.702125","exception":false,"start_time":"2021-03-19T05:17:37.686171","status":"completed"},"tags":[]}},{"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.augmentations.transforms import CLAHE, GaussNoise, ISONoise\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\n\nfrom sklearn.model_selection import StratifiedKFold","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2021-03-19T05:17:37.742104Z","iopub.status.busy":"2021-03-19T05:17:37.739742Z","iopub.status.idle":"2021-03-19T05:17:42.768223Z","shell.execute_reply":"2021-03-19T05:17:42.767186Z"},"papermill":{"duration":5.050454,"end_time":"2021-03-19T05:17:42.76836","exception":false,"start_time":"2021-03-19T05:17:37.717906","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{"papermill":{"duration":0.015956,"end_time":"2021-03-19T05:17:42.846847","exception":false,"start_time":"2021-03-19T05:17:42.830891","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class CFG:\n    seed = 42\n    model_name = 'resnet50'\n    pretrained = False\n    img_size = 512\n    num_classes = 12\n    batch_size = 32\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.execute_input":"2021-03-19T05:17:43.229654Z","iopub.status.busy":"2021-03-19T05:17:43.227833Z","iopub.status.idle":"2021-03-19T05:17:43.232573Z","shell.execute_reply":"2021-03-19T05:17:43.232953Z"},"papermill":{"duration":0.370106,"end_time":"2021-03-19T05:17:43.233104","exception":false,"start_time":"2021-03-19T05:17:42.862998","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load images that have been pre-resized by AnkurSingh to speed up the learning process. https://www.kaggle.com/c/plant-pathology-2021-fgvc8/discussion/227032","metadata":{"papermill":{"duration":0.016077,"end_time":"2021-03-19T05:17:43.266137","exception":false,"start_time":"2021-03-19T05:17:43.25006","status":"completed"},"tags":[]}},{"cell_type":"code","source":"PATH = \"../input/plant-pathology-2021-fgvc8/\"\nTEST_DIR = PATH + 'test_images/'","metadata":{"execution":{"iopub.execute_input":"2021-03-19T05:17:43.30231Z","iopub.status.busy":"2021-03-19T05:17:43.301824Z","iopub.status.idle":"2021-03-19T05:17:43.305421Z","shell.execute_reply":"2021-03-19T05:17:43.305036Z"},"papermill":{"duration":0.023219,"end_time":"2021-03-19T05:17:43.305587","exception":false,"start_time":"2021-03-19T05:17:43.282368","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(CFG.seed)","metadata":{"execution":{"iopub.execute_input":"2021-03-19T05:17:43.342331Z","iopub.status.busy":"2021-03-19T05:17:43.341811Z","iopub.status.idle":"2021-03-19T05:17:43.354245Z","shell.execute_reply":"2021-03-19T05:17:43.353819Z"},"papermill":{"duration":0.032063,"end_time":"2021-03-19T05:17:43.354343","exception":false,"start_time":"2021-03-19T05:17:43.32228","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_all = pd.read_csv(PATH + \"train.csv\")\nlabels = list(df_all['labels'].value_counts().keys())\nlabels_dict = dict(zip(labels, range(12)))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = pd.read_csv(PATH + \"sample_submission.csv\")\nsub.head()","metadata":{"execution":{"iopub.execute_input":"2021-03-19T05:17:43.393808Z","iopub.status.busy":"2021-03-19T05:17:43.393267Z","iopub.status.idle":"2021-03-19T05:17:43.437337Z","shell.execute_reply":"2021-03-19T05:17:43.437784Z"},"papermill":{"duration":0.066448,"end_time":"2021-03-19T05:17:43.437937","exception":false,"start_time":"2021-03-19T05:17:43.371489","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Dataset","metadata":{"papermill":{"duration":0.019487,"end_time":"2021-03-19T05:17:43.734692","exception":false,"start_time":"2021-03-19T05:17:43.715205","status":"completed"},"tags":[]}},{"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['labels'].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 = self.labels[idx]\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.execute_input":"2021-03-19T05:17:43.780348Z","iopub.status.busy":"2021-03-19T05:17:43.779846Z","iopub.status.idle":"2021-03-19T05:17:43.783498Z","shell.execute_reply":"2021-03-19T05:17:43.783064Z"},"papermill":{"duration":0.029578,"end_time":"2021-03-19T05:17:43.783653","exception":false,"start_time":"2021-03-19T05:17:43.754075","status":"completed"},"tags":[],"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.HorizontalFlip(p=0.5),\n            A.ShiftScaleRotate(p=0.5),\n            A.RandomBrightnessContrast(p=0.5),\n            A.Normalize(),\n            ToTensorV2(),\n        ])\n    else:\n        return Compose([\n            A.Resize(height=CFG.img_size, width=CFG.img_size),\n            A.Normalize(),\n            ToTensorV2(),\n        ])","metadata":{"execution":{"iopub.execute_input":"2021-03-19T05:17:43.827414Z","iopub.status.busy":"2021-03-19T05:17:43.8261Z","iopub.status.idle":"2021-03-19T05:17:43.828924Z","shell.execute_reply":"2021-03-19T05:17:43.828499Z"},"papermill":{"duration":0.026774,"end_time":"2021-03-19T05:17:43.82902","exception":false,"start_time":"2021-03-19T05:17:43.802246","status":"completed"},"tags":[],"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.execute_input":"2021-03-19T05:17:43.874Z","iopub.status.busy":"2021-03-19T05:17:43.872376Z","iopub.status.idle":"2021-03-19T05:17:43.87467Z","shell.execute_reply":"2021-03-19T05:17:43.875065Z"},"papermill":{"duration":0.026765,"end_time":"2021-03-19T05:17:43.875195","exception":false,"start_time":"2021-03-19T05:17:43.84843","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Define Model","metadata":{"papermill":{"duration":0.018747,"end_time":"2021-03-19T05:17:43.91307","exception":false,"start_time":"2021-03-19T05:17:43.894323","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class CustomResNet(nn.Module):\n    def __init__(self, model_name='resnet18', pretrained=False):\n        super().__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        in_features = self.model.get_classifier().in_features\n        self.model.fc = nn.Linear(in_features, CFG.num_classes)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.execute_input":"2021-03-19T05:17:43.955903Z","iopub.status.busy":"2021-03-19T05:17:43.955258Z","iopub.status.idle":"2021-03-19T05:17:43.95873Z","shell.execute_reply":"2021-03-19T05:17:43.958284Z"},"papermill":{"duration":0.026946,"end_time":"2021-03-19T05:17:43.958842","exception":false,"start_time":"2021-03-19T05:17:43.931896","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import OrderedDict\n\ndef fix_model_state_dict(state_dict):\n    new_state_dict = OrderedDict()\n    for k, v in state_dict.items():\n        name = k\n        if name.startswith('model.'):\n            name = name[6:]  # remove 'model.' of dataparallel\n        new_state_dict[name] = v\n    return new_state_dict","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CustomResNet(model_name=CFG.model_name, pretrained=CFG.pretrained)","metadata":{"execution":{"iopub.execute_input":"2021-03-19T05:17:44.052104Z","iopub.status.busy":"2021-03-19T05:17:44.051606Z","iopub.status.idle":"2021-03-19T05:17:47.545107Z","shell.execute_reply":"2021-03-19T05:17:47.544643Z"},"papermill":{"duration":3.516051,"end_time":"2021-03-19T05:17:47.545235","exception":false,"start_time":"2021-03-19T05:17:44.029184","status":"completed"},"tags":[],"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint = \"../input/plat2021-resnet50/last.ckpt\"\n\nweight = torch.load(checkpoint)['state_dict']\nmodel.load_state_dict(fix_model_state_dict(weight))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference","metadata":{}},{"cell_type":"code","source":"model.cuda()\nmodel.eval()\n\npredictions = []\nfor batch in test_loader:\n    image = batch['image'].cuda()\n    with torch.no_grad():\n        outputs = model(image)\n        preds = outputs.argmax(1).detach().cpu().numpy()\n        predictions.append(preds)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"inv_labels_dict = {v: k for k, v in labels_dict.items()}\ninv_labels_dict","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub['labels'] = np.concatenate(predictions)\nsub = sub.replace({\"labels\": inv_labels_dict})\nsub.to_csv('submission.csv', index=False)\nsub.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}