{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"!pip install ../input/efficientnet-pytorch-070/efficientnet_pytorch-0.7.0-py3-none-any.whl","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import os\n\nimport albumentations\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nfrom torchvision import models, transforms\n\n\nfrom efficientnet_pytorch import EfficientNet","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model_path = \"../input/en-b4-tta-calr-clahe-v3/eff_epoch_11.pth\"\nsample_sub_path = \"../input/cassava-leaf-disease-classification/sample_submission.csv\"\ntest_images_path = \"../input/cassava-leaf-disease-classification/test_images\"","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Creating model and loading weights"},{"metadata":{"trusted":true},"cell_type":"code","source":"# Loading Efficient B4 and update it's fc layer output to 5 classes\n\nmodel = EfficientNet.from_name('efficientnet-b4', num_classes=5)\n\n\nmodel.to(torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"))\nmodel.load_state_dict(torch.load(model_path))\nmodel.eval()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Submission Augmentations"},{"metadata":{"trusted":true},"cell_type":"code","source":"# sub_aug = albumentations.Compose([\n#                 albumentations.CenterCrop(512, 512, p=1.),\n#                 albumentations.CLAHE(clip_limit=(1.0, 20.0), tile_grid_size=(32, 32), always_apply=True, p=1),\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\nsub_aug = albumentations.Compose([\n                albumentations.RandomResizedCrop(512, 512, scale=(0.5, 1.0)),\n                albumentations.Transpose(p=0.5),\n                albumentations.HorizontalFlip(p=0.5),\n                albumentations.VerticalFlip(p=0.5),\n                albumentations.ShiftScaleRotate(p=0.8),\n#                 albumentations.HueSaturationValue(\n#                     hue_shift_limit=20, \n#                     sat_shift_limit=30, \n#                     val_shift_limit=20, \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.CLAHE(p=0.5),\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.)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Creating the Submission file"},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_sub = pd.read_csv(sample_sub_path)\ntta_count = 10\n\npredictions = []\nfor _, sample_row in sample_sub.iterrows():\n    \n    image_pred = 0\n    for j in range(tta_count):\n        image = np.array(Image.open(os.path.join(test_images_path, sample_row.image_id)))\n        image = sub_aug(image=image)[\"image\"]\n        image = transforms.ToTensor()(np.array(image))\n        image = image.to(torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\"))\n        outputs = model(image.unsqueeze(0))\n        \n        image_pred += outputs\n    image_pred /= tta_count\n    _, pred_label = torch.max(image_pred, 1)\n        \n    predictions.append([sample_row.image_id, pred_label.item()])\n\nsub_df = pd.DataFrame(predictions,columns=['image_id', 'label'])\nsub_df.to_csv('submission.csv', index=False)\nprint(sub_df.head())\n","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}