{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import pandas as pd\nfrom torch.utils.data import Dataset, DataLoader\nimport cv2\nimport torch\nfrom torch import nn\nimport os\nimport sys\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensor","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"package_path = '../input/pytorch-image-models/pytorch-image-models-master'\nsys.path.append(package_path) ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import timm\nSEED = 2484\ntorch.manual_seed(SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = True\ntorch.cuda.manual_seed_all(SEED)\ntorch.is_deterministic=True","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_tta(image_size=380, p=0.5):\n    imagenet_stats = {\"mean\": [0.485, 0.456, 0.406],\n                      \"std\": [0.229, 0.224, 0.225]}\n    augs = A.Compose(\n        [\n            A.Resize(428, 428, cv2.INTER_CUBIC),\n            A.RandomCrop(image_size, image_size),\n            A.RandomRotate90(p=p),\n            A.RandomBrightnessContrast(\n                brightness_limit=0.2, contrast_limit=0.2, p=p),\n            A.OneOf(\n                [\n                    A.MotionBlur(p=p),\n                    A.MedianBlur(blur_limit=3, p=p),\n                    A.Blur(blur_limit=3, p=p),\n                    A.GaussianBlur(blur_limit=(3, 5), p=p)\n                ],\n                p=p,\n            ),\n            A.OneOf(\n                [\n                    A.OpticalDistortion(p=p),\n                    A.GridDistortion(p=p)\n                ],\n                p=p,\n            ),\n            ToTensor(normalize=imagenet_stats),\n        ]\n    )\n    return augs","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaLeafDataset(Dataset):\n    def __init__(self, root_dir, transforms):\n        self.root_dir = root_dir\n        self.transform = transforms\n        self.dataframe = pd.read_csv('../input/cassava-leaf-disease-classification/sample_submission.csv')\n\n    def __len__(self):\n        return self.dataframe.shape[0]\n\n    def __getitem__(self, idx):\n        data = self.dataframe.iloc[idx]\n        img_name = os.path.join(self.root_dir, data[0])\n        image = cv2.imread(img_name)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = self.transform(image=image)['image']\n        return image, data[1]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_transforms = get_tta(380)\ntest_data = CassavaLeafDataset('../input/cassava-leaf-disease-classification/test_images', test_transforms)\ntest_loader = DataLoader(test_data, batch_size=16, pin_memory=True, num_workers=4)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df = pd.read_csv('../input/cassava-leaf-disease-classification/sample_submission.csv')\ndf.shape[0]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def get_probabilities(model):\n    model.eval()\n    model.to(device)\n    res = torch.zeros((df.shape[0], 5), device=device)\n    epochs = 10\n    with torch.no_grad():\n        for _ in range(epochs):\n            predictions = torch.tensor([], device=device)\n            for images, labels in test_loader:\n                images, labels = images.to(device), labels.to(device)\n                outputs = model(images)\n                predictions = torch.cat((predictions, outputs.data), dim=0)\n            res = res + predictions\n        res = res / epochs\n    return res","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"resnest = timm.create_model('resnest50d_1s4x24d', pretrained=False)\nresnest.fc = nn.Linear(in_features=2048, out_features=5, bias=True)\nresnest.load_state_dict(torch.load('../input/mycassavamodels/resnest50.pth'))\nresnest_probs = get_probabilities(resnest)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"xception = timm.create_model('xception65', pretrained=False)\nxception.head.fc = nn.Linear(in_features=2048, out_features=5, bias=True)\nxception.load_state_dict(torch.load('../input/mycassavamodels/xception65.pth'))\nxception_probs = get_probabilities(xception)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"efficientb4 = timm.create_model('tf_efficientnet_b4_ns', pretrained=False)\nefficientb4.classifier = nn.Linear(in_features=1792, out_features=5, bias=True)\nefficientb4.load_state_dict(torch.load('../input/mycassavamodels/efficientnetB4.pth'))\nefficientb4_probs = get_probabilities(efficientb4)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"res = resnest_probs + efficientb4_probs + xception_probs\nres = res / 3.\npred = torch.argmax(res, dim=1)\nprint(pred)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"labels = [elem.item() for elem in pred]\ndf['label'] = labels\ndf.to_csv('./submission.csv', index=False)\ndf.head()","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}