{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"!mkdir -p /tmp/pip/cache/\n!mkdir -p /root/.cache/torch/hub/checkpoints\n!cp ../input/effficientnet/efficientnet_pytorch-0.7.0.xyz /tmp/pip/cache/efficientnet_pytorch-0.7.0.tar.gz\n!cp ../input/effficientnet/efficientnet-b4-6ed6700e.pth /root/.cache/torch/hub/checkpoints/efficientnet-b4-6ed6700e.pth\n!pip install --no-index --find-links /tmp/pip/cache/ efficientnet_pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport cv2\nfrom torch.utils.data import Dataset\nimport torch\nimport albumentations\nfrom PIL import Image\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import efficientnet_pytorch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# checking if cuda is available\nfrom torch import device as device_\n\ndevice = device_(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"sample_submission = pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")\nsample_submission","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class leaf_test_classification(Dataset):\n    def __init__(self, ids, df):\n        self.ids = ids\n        self.image_ids = df.image_id.values\n        \n        self.aug = albumentations.Compose([\n                    albumentations.RandomResizedCrop(256, 256),\n                    albumentations.Transpose(p=0.5),\n                    albumentations.HorizontalFlip(p=0.5),\n                    albumentations.VerticalFlip(p=0.5),\n                    albumentations.HueSaturationValue(\n                        hue_shift_limit=0.2, \n                        sat_shift_limit=0.2,\n                        val_shift_limit=0.2, \n                        p=0.5\n                    ),\n                    albumentations.RandomBrightnessContrast(\n                        brightness_limit=(-0.1,0.1), \n                        contrast_limit=(-0.1, 0.1), \n                        p=0.5\n                    ),\n                    albumentations.Normalize(\n                        mean=[0.485, 0.456, 0.406], \n                        std=[0.229, 0.224, 0.225], \n                        max_pixel_value=255.0, \n                        p=1.0\n                    )\n                ], p=1.)\n    def __len__(self):\n        return len(self.ids)\n    \n    def __getitem__(self, index):\n        # converting jpg format of images to numpy array\n        img = np.array(Image.open('../input/cassava-leaf-disease-classification/test_images/' + self.image_ids[index]))\n        img = self.aug(image = img)['image']\n        img = np.transpose(img , (2,0,1)).astype(np.float32) # 2,0,1 because pytorch excepts image channel first then dimension of image\n        \n       \n        return torch.tensor(img, dtype = torch.float) , self.image_ids[index]\n    \n\n    \ntest_data = leaf_test_classification(ids = [i for i in range(len(sample_submission))], df = sample_submission)\n\ntest_dataloader = DataLoader(test_data,\n                        num_workers=4,\n                        batch_size=32,\n                        drop_last=False\n                       )\nidx = 0 \nimg = test_data[idx][0]\n\nprint(test_data[idx][1])\nnpimg = img.numpy()\nplt.imshow(np.transpose(npimg, (1,2,0)))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class EfficientNet_b4(nn.Module):\n    def __init__(self):\n        super(EfficientNet_b3, self).__init__()\n        self.model = efficientnet_pytorch.EfficientNet.from_pretrained('efficientnet-b4')\n        self.dropout = nn.Dropout(0.1)\n        self.final_layer = nn.Linear(1792 , 5)\n        \n    def forward(self, inputs):\n        batch_size, _, _, _ = inputs.shape\n        \n        x = self.model.extract_features(inputs)\n\n        # Pooling and final linear layer\n        x = self.model._avg_pooling(x)\n        \n        x = F.adaptive_avg_pool2d(x, 1).reshape(batch_size, -1)\n        outputs = self.final_layer(self.dropout(x))\n\n        return outputs\n    \nmodel = efficientnet_pytorch.EfficientNet.from_pretrained('efficientnet-b4')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"checkpoint_path = \"../input/cassava-efficientnetb3-trained-on-fold-1/leaf_Classsification_GPU_CutMix_EfficientNet-B4.pt\"\n\ncheckpoint = torch.load(checkpoint_path)\n\nmodel.load_state_dict(checkpoint['state_dict'], strict=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# run inference 5 times for test time augmentation implementation\nmodel.eval()\n\nfinal_preds = None\nfor j in range(5):\n    for image,image_id in test_dataloader:\n        image = image.to(device, dtype=torch.float)\n\n        with torch.no_grad():\n            preds = model(image)\n    temp_preds = None\n    for p in preds:\n        if temp_preds is None:\n            temp_preds = p\n        else:\n            temp_preds = np.vstack((temp_preds, p))\n    if final_preds is None:\n        final_preds = temp_preds\n    else:\n        final_preds += temp_preds\nfinal_preds /= 5\nfinal_preds = final_preds.detach().cpu().numpy()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"final_preds = np.argmax(final_preds)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_submission.label = final_preds \nsample_submission","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_submission.to_csv(\"submission.csv\", index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}