{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"import sys\nimport numpy as np\nimport pandas as pd\nimport torchvision\nimport torch.nn as nn\nfrom tqdm import tqdm\nfrom PIL import Image, ImageFile\nfrom torch.utils.data import Dataset\nimport torch\nfrom torchvision import transforms\nimport os\n\n#package_dir = \"../input/pretrained-models/pretrained-models/pretrained-models.pytorch-master/\"\n#sys.path.insert(0, package_dir)\n\n\ndevice = torch.device('cuda:0')\nImageFile.LOAD_TRUNCATED_IMAGES = True","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"!pip install pretrainedmodels","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"import pretrainedmodels","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"class RetinopathyDatasetTest(Dataset):\n    def __init__(self, csv_file, transform):\n        self.data = pd.read_csv(csv_file)\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        img_name = os.path.join('../input/aptos2019-blindness-detection/test_images', self.data.loc[idx,'id_code']+'.png')\n        image = Image.open(img_name)\n        image = self.transform(image)\n        return {'image':image}","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model = pretrainedmodels.__dict__['resnet101'](pretrained = None)\nmodel.avg_pool = nn.AdaptiveAvgPool2d(1)\nmodel.last_linear = nn.Sequential(\n                        nn.BatchNorm1d(2048, eps =1e-5, momentum = 0.1, affine= True, track_running_stats = True),\n                        nn.Dropout(0.25),\n                        nn.Linear(in_features = 2048, out_features = 2048, bias = True),\n                        nn.ReLU(),\n                        nn.BatchNorm1d(2048, eps = 1e-5, momentum =0.1, affine = True, track_running_stats = True ),\n                        nn.Dropout(0.5),\n                        nn.Linear(in_features = 2048, out_features = 1, bias = True )\n                        )\n\nmodel.load_state_dict(torch.load(\"../input/mmmodel/model.bin\"))\nmodel = model.to(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for param in model.parameters():\n    param.requires_grad = False\n    \nmodel.eval()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_transform = transforms.Compose([\n    transforms.Resize((224,224)),\n    transforms.RandomHorizontalFlip(),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n    ])\n\ntest_dataset = RetinopathyDatasetTest(\"../input/aptos2019-blindness-detection/sample_submission.csv\", transform = test_transform)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# TTA"},{"metadata":{"trusted":true},"cell_type":"code","source":"test_preds_all = np.zeros((len(test_dataset),10))\nfor j in range(10):\n    test_data_loader = torch.utils.data.DataLoader(test_dataset, batch_size = 32, shuffle = False, num_workers = 4)\n    test_preds = np.zeros((len(test_dataset),1))\n    tk0 = tqdm(test_data_loader)\n    for i, x_batch in enumerate(tk0):\n        x_batch = x_batch['image']\n        pred = model(x_batch.to(device))\n        test_preds[i*32:(i+1)*32] = pred.detach().cpu().squeeze().numpy().ravel().reshape(-1,1)\n    test_preds = test_preds.flatten()\n    test_preds_all[:,j] = test_preds","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_preds_all","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Average the results\ntest_preds_agg = np.sum(test_preds_all,axis = 1)/10","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_preds_agg","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"coef = [0.5,1.5,2.5,3.5]\n\nfor i, pred in enumerate(test_preds_agg):\n    if pred<coef[0]:\n        test_preds_agg[i] = 0\n    elif pred>=coef[0] and pred <coef[1]:\n        test_preds_agg[i] = 1\n    elif pred>=coef[1] and pred<coef[2]:\n        test_preds_agg[i] = 2\n    elif pred>=coef[2] and pred<coef[3]:\n        test_preds_agg[i] = 3\n    else:\n        test_preds_agg[i] = 4\n        \n\nsample = pd.read_csv('../input/aptos2019-blindness-detection/sample_submission.csv')\nsample.diagnosis = test_preds_agg.astype(int)\nsample.to_csv('submission.csv',index = False)\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"> Reference :\n https://www.kaggle.com/abhishek/pytorch-inference-kernel-lazy-tta/data"},{"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":1}