{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"! pip install ../input/tta-torch","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"! pip install ../input/timm-package/timm-0.1.26-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 torch\nimport torchvision\nimport skimage.io as io\nimport os\nimport numpy as np\nfrom PIL import Image\nimport torch.nn as nn\nimport pandas as pd\nimport pickle\nimport pytorch_lightning as pl\nimport timm\nimport ttach as tta","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class CassavaNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        backbone = timm.create_model(TIMM_MODEL, pretrained=True)\n        n_features = backbone.fc.in_features\n        self.backbone = nn.Sequential(*backbone.children())[:-2]\n        self.classifier = nn.Linear(n_features, 5)\n        self.pool = nn.AdaptiveAvgPool2d((1, 1))\n\n    def forward_features(self, x):\n        x = self.backbone(x)\n        return x\n\n    def forward(self, x):\n        feats = self.forward_features(x)\n        x = self.pool(feats).view(x.size(0), -1)\n        x = self.classifier(x)\n        return x\n    \ndef load_checkpoint(filepath):\n    checkpoint = torch.load(filepath)\n    model = checkpoint['model']\n    model.load_state_dict(checkpoint['state_dict'])\n    for parameter in model.parameters():\n        parameter.requires_grad = False\n\n    model.eval()\n    return model\n\ndevice = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model_0 = load_checkpoint('../input/cassavamodels2/ckpt_resnet50-512-0.pth')\nmodel_0 = model_0.cuda()\n\nmodel_1 = load_checkpoint('../input/cassavamodels2/ckpt_resnet50-512-0-snap.pth')\nmodel_1 = model_1.cuda()\n\nmodel_2 = load_checkpoint('../input/cassavamodels2/ckpt_resnet50-512-0.pth')\nmodel_2 = model_2.cuda()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#We say the model to use tta to make better predictions\n\ntta_transforms = tta.Compose(\n    [\n        tta.HorizontalFlip(),\n        tta.Rotate90(angles=[0, 180]),\n        #tta.Scale(scales=[1, 2]),\n        tta.Multiply(factors=[0.95, 1, 1.05])       \n    ]\n)\n\nmerge_mode = 'mean'\n#gmean (geometric mean)\n#sum\n#max\n#min\n#tsharpen (temperature sharpen with t=0.5)\n\ntta_model_0 = tta.ClassificationTTAWrapper(model_0, tta_transforms, merge_mode = merge_mode)\ntta_model_1 = tta.ClassificationTTAWrapper(model_1, tta_transforms, merge_mode = merge_mode)\ntta_model_2 = tta.ClassificationTTAWrapper(model_2, tta_transforms, merge_mode = merge_mode)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"database_base_path = '../input/cassava-leaf-disease-classification/'\nn_classes = 5\nfiles_path = f'{database_base_path}test_images/'\ntest_size = len(os.listdir(files_path))\ntest_preds = np.zeros((test_size, n_classes))\nloader = torchvision.transforms.Compose([\n        torchvision.transforms.Resize(512),\n        torchvision.transforms.CenterCrop(448),\n        torchvision.transforms.ToTensor(),\n        torchvision.transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n    ])\n\nimage_names =  os.listdir(files_path)\n\ndef image_loader(image_name):\n    image = Image.open(image_name)\n    image = loader(image).float()\n    image = image.unsqueeze(0)  #this is for VGG, may not be needed for ResNet\n    return image.cuda()  #assumes that you're using GPU\n\ndef ensembled_prediction(model_0, model_1, model_2, files_path):\n    labels = []\n    for img in os.listdir(files_path):\n        image = image_loader(files_path + img)\n        model_0.eval()\n        model_1.eval()\n        model_2.eval()\n        with torch.no_grad():\n            \n            pred_0 = model_0(image.cuda())\n            pred_1 = model_1(image.cuda())\n            pred_2 = model_2(image.cuda())\n            pred = pred_0 + pred_1 + pred_2\n            label = torch.argmax(pred)\n            labels.append(label.item())\n    return np.array(labels) ","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_preds = ensembled_prediction(tta_model_0, tta_model_1, tta_model_2, files_path)\ntest_preds","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"submission = pd.DataFrame({'image_id': image_names, 'label': test_preds})\nsubmission.to_csv('submission.csv', index=False)\ndisplay(submission.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}