{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"import pandas as pd\n\nimport time\nimport torchvision\nimport torch.nn as nn\nfrom tqdm import tqdm_notebook as tqdm\n\nfrom PIL import Image, ImageFile\nfrom torch.utils.data import Dataset\nimport torch\nimport torch.optim as optim\nfrom torchvision import transforms\nfrom torch.optim import lr_scheduler\nimport os\nfrom torch.autograd import Variable\nfrom PIL import ImageOps \ndevice = torch.device(\"cuda:0\")\n# ImageFile.LOAD_TRUNCATED_IMAGES = True","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"torch.cuda.get_device_name()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**defining dataset**"},{"metadata":{"trusted":true},"cell_type":"code","source":"class RetinopathyDataset(Dataset):\n\n    def __init__(self, csv_file, transform, test=False):\n        self.test_ = test\n        self.transform = transform\n        self.data = pd.read_csv(csv_file)\n        self.data_test = pd.read_csv('../input/aptos2019-blindness-detection/test.csv')\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        if not self.test_:\n            img_name = os.path.join('../input/aptos2019-blindness-detection/train_images', self.data.loc[idx, 'id_code'] + '.png')\n            image = Image.open(img_name)\n            image = image.resize((350, 350), resample=Image.BILINEAR)   \n            image = ImageOps.equalize(image)\n            label = torch.tensor(self.data.loc[idx, 'diagnosis'])\n            return {'image': self.transform(image),\n                    'labels': label\n                    }\n        else:\n            img_name = os.path.join('../input/aptos2019-blindness-detection/test_images', self.data_test.loc[idx, 'id_code'] + '.png')\n            image = Image.open(img_name)\n            image = image.resize((350, 350), resample=Image.BILINEAR)\n            image = ImageOps.equalize(image)\n            return  self.transform(image)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Replacing classifier layer**"},{"metadata":{"trusted":true},"cell_type":"code","source":"model = torchvision.models.resnet101(pretrained=False)\nmodel.load_state_dict(torch.load(\"../input/pytorch-pretrained-models/resnet101-5d3b4d8f.pth\"))\nnum_features = model.fc.in_features\nmodel.fc = nn.Sequential(\n                          nn.BatchNorm1d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True),\n                          nn.Dropout(p=0.25),\n                          nn.Linear(in_features=2048, out_features=2048, bias=True),\n                          nn.ReLU(),\n                          nn.BatchNorm1d(2048, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True),\n                          nn.Dropout(p=0.5),\n                          nn.Linear(in_features=2048, out_features=5, bias=True),\n                         )\n\nmodel = model.to(device)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model.eval()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"for name, param in model.named_parameters():\n    print(name, param.requires_grad)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Defining some transformations**"},{"metadata":{"trusted":true},"cell_type":"code","source":"train_transform = transforms.Compose([\n    transforms.RandomChoice(\n        [\n            transforms.RandomHorizontalFlip(),\n            transforms.RandomVerticalFlip(),\n            transforms.RandomAffine(20),\n        ]\n    ),\n    transforms.ToTensor(),\n    transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n])\ntest_dataset = RetinopathyDataset(csv_file='../input/aptos2019-blindness-detection/sample_submission.csv',\n                                      transform=train_transform)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Hyperparameters and optimizer**"},{"metadata":{"trusted":true},"cell_type":"code","source":"train_dataset = RetinopathyDataset(csv_file='../input/aptos2019-blindness-detection/train.csv', transform=train_transform)\ndata_loader = torch.utils.data.DataLoader(train_dataset, batch_size=24, shuffle=True)\n\noptimizer = optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-5)\nscheduler = lr_scheduler.StepLR(optimizer, step_size=10)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"int(len(data_loader))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def accuracy_score(output, labels):\n    score = 0\n    for c, cc in zip(output, labels):\n        if c == cc:\n            score += 1\n    return score/len(output)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**trainining**"},{"metadata":{"trusted":true},"cell_type":"code","source":"def train():\n    since = time.time()\n    criterion = nn.criterion = nn.CrossEntropyLoss()\n    num_epochs = 50\n    for epoch in range(num_epochs):\n        print('Epoch {}/{}'.format(epoch, num_epochs - 1))\n        print('-' * 10)\n        scheduler.step()\n        model.train()\n        running_loss = 0.0\n        tk0 = tqdm(data_loader, total=int(len(data_loader)))\n        counter = 0\n        for bi, d in enumerate(tk0):\n            inputs = d[\"image\"]\n            labels = d[\"labels\"].view(-1, 1)\n            inputs = inputs.to(device)\n            labels = labels.to(device, dtype=torch.long)\n            optimizer.zero_grad()\n            with torch.set_grad_enabled(True):\n                outputs = model(inputs)\n                loss = criterion(outputs, labels.squeeze_())\n                loss.backward()\n                optimizer.step()\n            running_loss += loss.item() * inputs.size(0)\n            counter += 1\n            tk0.set_postfix(loss=(running_loss / (counter * data_loader.batch_size)))\n        epoch_loss = running_loss / len(data_loader)\n        \n        print('Training Loss: {:.4f}'.format(epoch_loss))\n        print('Epoch Accuracy: ', accuracy_score(outputs.argmax(1), labels))\n    time_elapsed = time.time() - since\n    print('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))\n    torch.save(model.state_dict(), \"model.bin\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_dataset = RetinopathyDataset(csv_file='../input/aptos2019-blindness-detection/test.csv', transform=train_transform, test=True)\ntest_data_loader = torch.utils.data.DataLoader(test_dataset, batch_size=16, shuffle=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"result = []\n\ntk0 = tqdm(test_data_loader, total=int(len(test_data_loader)))\nfor d in enumerate(tk0):\n    inputs = Variable(d[1])\n    inputs = inputs.to(device, dtype=torch.float)\n    with torch.set_grad_enabled(False):\n        outputs = model(inputs)\n        for pred in outputs.argmax(1).tolist():\n            result.append(pred)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"result[:5]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample = pd.read_csv(\"../input/aptos2019-blindness-detection/sample_submission.csv\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample.diagnosis = result","execution_count":null,"outputs":[]},{"metadata":{"trusted":false},"cell_type":"code","source":"sample.to_csv(\"submission.csv\", index=False)","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.7.3"}},"nbformat":4,"nbformat_minor":1}