{"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_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Cool Imports","metadata":{}},{"cell_type":"code","source":"!pip install pytorch_pretrained_vit","metadata":{"execution":{"iopub.status.busy":"2022-04-12T02:17:09.008021Z","iopub.execute_input":"2022-04-12T02:17:09.008407Z","iopub.status.idle":"2022-04-12T02:17:20.017629Z","shell.execute_reply.started":"2022-04-12T02:17:09.008322Z","shell.execute_reply":"2022-04-12T02:17:20.016811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport random\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\n\ndevice = torch.device(\"cuda:0\")\nImageFile.LOAD_TRUNCATED_IMAGES = True\n\n\nos.environ['CUDA_VISIBLE_DEVICE'] = '0'\n\ntorch.cuda.set_device(0)\n# 设置全局参数\nmodellr = 1e-6\nBATCH_SIZE = 64\nEPOCHS = 100\nDEVICE = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-04-12T02:17:20.019843Z","iopub.execute_input":"2022-04-12T02:17:20.0201Z","iopub.status.idle":"2022-04-12T02:17:21.771972Z","shell.execute_reply.started":"2022-04-12T02:17:20.020068Z","shell.execute_reply":"2022-04-12T02:17:21.771231Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def setup_seed(seed):\n     torch.manual_seed(seed)\n     torch.cuda.manual_seed_all(seed)\n     np.random.seed(seed)\n     random.seed(seed)\n     torch.backends.cudnn.deterministic = True","metadata":{"execution":{"iopub.status.busy":"2022-04-12T02:17:21.779343Z","iopub.execute_input":"2022-04-12T02:17:21.780335Z","iopub.status.idle":"2022-04-12T02:17:21.785699Z","shell.execute_reply.started":"2022-04-12T02:17:21.780296Z","shell.execute_reply":"2022-04-12T02:17:21.78505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"setup_seed(2333)","metadata":{"execution":{"iopub.status.busy":"2022-04-12T02:17:21.787316Z","iopub.execute_input":"2022-04-12T02:17:21.787932Z","iopub.status.idle":"2022-04-12T02:17:21.795549Z","shell.execute_reply.started":"2022-04-12T02:17:21.787893Z","shell.execute_reply":"2022-04-12T02:17:21.79487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset Class","metadata":{}},{"cell_type":"code","source":"class RetinopathyDatasetTrain(Dataset):\n\n    def __init__(self, csv_file):\n\n        self.data = pd.read_csv(csv_file)\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/train_images', self.data.loc[idx, 'id_code'] + '.png')\n        image = Image.open(img_name)\n        image = image.resize((224, 224), resample=Image.BILINEAR)\n        label = torch.tensor(self.data.loc[idx, 'diagnosis'])\n        return {'image': transforms.ToTensor()(image),\n                'labels': label\n                }\n    \nclass RetinopathyDatasetTest(Dataset):\n\n    def __init__(self, csv_file):\n        self.data = pd.read_csv(csv_file)\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 = image.resize((224, 224), resample=Image.BILINEAR)\n        return transforms.ToTensor()(image)","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","execution":{"iopub.status.busy":"2022-04-12T02:17:21.797006Z","iopub.execute_input":"2022-04-12T02:17:21.797501Z","iopub.status.idle":"2022-04-12T02:17:21.808807Z","shell.execute_reply.started":"2022-04-12T02:17:21.797465Z","shell.execute_reply":"2022-04-12T02:17:21.808108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Get the model","metadata":{}},{"cell_type":"code","source":"from pytorch_pretrained_vit import ViT\nmodel_name = 'B_16'\nmodel = ViT(model_name, pretrained=True, num_classes=5)\nmodel.load_state_dict(torch.load('../input/aptos2019-pretrained-best-2/best_model.pt'))\nmodel = model.to(device)","metadata":{"execution":{"iopub.status.busy":"2022-04-12T02:17:21.809708Z","iopub.execute_input":"2022-04-12T02:17:21.811874Z","iopub.status.idle":"2022-04-12T02:18:25.255913Z","shell.execute_reply.started":"2022-04-12T02:17:21.811847Z","shell.execute_reply":"2022-04-12T02:18:25.255203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\n\noptimizer = optim.Adam(model.parameters(), lr=modellr)\n\n\ndef adjust_learning_rate(optimizer, epoch):\n    \"\"\"Sets the learning rate to the initial LR decayed by 10 every 30 epochs\"\"\"\n    modellrnew = modellr * (0.1 ** (epoch // 50))\n    print(\"lr:\", modellrnew)\n    for param_group in optimizer.param_groups:\n        param_group['lr'] = modellrnew","metadata":{"execution":{"iopub.status.busy":"2022-04-12T02:18:25.257301Z","iopub.execute_input":"2022-04-12T02:18:25.257544Z","iopub.status.idle":"2022-04-12T02:18:25.264983Z","shell.execute_reply.started":"2022-04-12T02:18:25.257513Z","shell.execute_reply":"2022-04-12T02:18:25.263749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create dataset + optimizer","metadata":{}},{"cell_type":"code","source":"dataset = RetinopathyDatasetTrain(csv_file='../input/aptos2019-blindness-detection/train.csv')\ntrain_dataset, valid_dataset = torch.utils.data.random_split(dataset, [int(0.8 * len(dataset)), len(dataset) - int(0.8*len(dataset))])\ndata_loader = torch.utils.data.DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2, drop_last=False)\nvalid_data_loader = torch.utils.data.DataLoader(valid_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=2, drop_last=False)","metadata":{"execution":{"iopub.status.busy":"2022-04-12T02:18:25.266538Z","iopub.execute_input":"2022-04-12T02:18:25.266794Z","iopub.status.idle":"2022-04-12T02:18:25.295509Z","shell.execute_reply.started":"2022-04-12T02:18:25.266758Z","shell.execute_reply":"2022-04-12T02:18:25.294893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Loop","metadata":{}},{"cell_type":"code","source":"best_acc = 0\ntrain_loss, valid_loss, valid_acc = [], [], []","metadata":{"execution":{"iopub.status.busy":"2022-04-12T02:18:25.298093Z","iopub.execute_input":"2022-04-12T02:18:25.298477Z","iopub.status.idle":"2022-04-12T02:18:25.302221Z","shell.execute_reply.started":"2022-04-12T02:18:25.298441Z","shell.execute_reply":"2022-04-12T02:18:25.301446Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef Valid():\n    print('Valid')\n    tk1 = tqdm(valid_data_loader, total=int(len(valid_data_loader)))\n    acc_num, sum_num, loss_sum = 0, 0, 0\n    for bi, d in enumerate(tk1):\n        inputs = d[\"image\"]\n        labels = d[\"labels\"]\n        inputs = inputs.to(device, dtype=torch.float)\n        labels = labels.to(device, dtype=torch.long)\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        acc_num += (labels==outputs.argmax(-1)).sum()\n        sum_num += labels.size(0)\n        loss_sum += loss.item() * inputs.size(0)\n    acc = acc_num / sum_num\n    print('Valid Acc: {:.4f}'.format(acc))\n    global best_acc\n    if acc > best_acc:\n        best_acc = acc\n        torch.save(model.state_dict(), 'best_model.pt')\n        print('save best model with acc:{:.4f}'.format(acc))\n    valid_loss.append(loss_sum / len(valid_data_loader))\n    valid_acc.append(acc)","metadata":{"execution":{"iopub.status.busy":"2022-04-12T02:18:25.303594Z","iopub.execute_input":"2022-04-12T02:18:25.303954Z","iopub.status.idle":"2022-04-12T02:18:25.314303Z","shell.execute_reply.started":"2022-04-12T02:18:25.30392Z","shell.execute_reply":"2022-04-12T02:18:25.313458Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nsince = time.time()\nfor epoch in range(EPOCHS):\n    print('Epoch {}/{}'.format(epoch, EPOCHS - 1))\n    print('-' * 30)\n    adjust_learning_rate(optimizer, epoch)\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\"]\n        inputs = inputs.to(device, dtype=torch.float)\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)\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    print('Training Loss: {:.4f}'.format(epoch_loss))\n    Valid()\n    train_loss.append(epoch_loss)\n    \n\ntime_elapsed = time.time() - since\nprint('Training complete in {:.0f}m {:.0f}s'.format(time_elapsed // 60, time_elapsed % 60))\ntorch.save(model.state_dict(), \"model.bin\")","metadata":{"execution":{"iopub.status.busy":"2022-04-12T02:18:25.315726Z","iopub.execute_input":"2022-04-12T02:18:25.316135Z","iopub.status.idle":"2022-04-12T02:36:14.297384Z","shell.execute_reply.started":"2022-04-12T02:18:25.316097Z","shell.execute_reply":"2022-04-12T02:36:14.296041Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), './ckp.pt')","metadata":{"execution":{"iopub.status.busy":"2022-04-12T02:36:14.298861Z","iopub.status.idle":"2022-04-12T02:36:14.299477Z","shell.execute_reply.started":"2022-04-12T02:36:14.299247Z","shell.execute_reply":"2022-04-12T02:36:14.299271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"valid_acc = [i.cpu() for i in valid_acc]","metadata":{"execution":{"iopub.status.busy":"2022-03-27T14:11:30.866264Z","iopub.execute_input":"2022-03-27T14:11:30.867351Z","iopub.status.idle":"2022-03-27T14:11:30.878125Z","shell.execute_reply.started":"2022-03-27T14:11:30.867311Z","shell.execute_reply":"2022-03-27T14:11:30.876946Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nplt.figure(figsize=(8,6))\nplt.plot([i for i in range(len(train_loss))],train_loss,'',label=\"train_loss\")\nplt.plot([i for i in range(len(valid_loss))],valid_loss,'',label=\"valid_loss\")\nplt.title('loss')\nplt.legend(loc='upper right')\nplt.xlabel('epoch')\nplt.ylabel('')\nplt.grid(len(train_loss))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-27T14:11:33.221128Z","iopub.execute_input":"2022-03-27T14:11:33.22185Z","iopub.status.idle":"2022-03-27T14:11:33.635239Z","shell.execute_reply.started":"2022-03-27T14:11:33.221798Z","shell.execute_reply":"2022-03-27T14:11:33.63411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(8,6))\nplt.plot([i for i in range(len(valid_acc))],valid_acc,'',label=\"acc\")\n\nplt.title('acc')\nplt.legend(loc='upper right')\nplt.xlabel('epoch')\nplt.ylabel('')\nplt.grid(len(train_loss))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-03-27T14:11:42.682844Z","iopub.execute_input":"2022-03-27T14:11:42.683191Z","iopub.status.idle":"2022-03-27T14:11:42.9383Z","shell.execute_reply.started":"2022-03-27T14:11:42.683155Z","shell.execute_reply":"2022-03-27T14:11:42.937339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.chdir('/kaggle/working')\nprint(os.getcwd())\nprint(os.listdir(\"/kaggle/working\"))\nfrom IPython.display import FileLink\nFileLink('ckp.pt')","metadata":{"execution":{"iopub.status.busy":"2022-03-27T14:13:52.604409Z","iopub.execute_input":"2022-03-27T14:13:52.60533Z","iopub.status.idle":"2022-03-27T14:13:52.616534Z","shell.execute_reply.started":"2022-03-27T14:13:52.60529Z","shell.execute_reply":"2022-03-27T14:13:52.615303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Testing","metadata":{}},{"cell_type":"code","source":"test_dataset = RetinopathyDatasetTest(csv_file='../input/aptos2019-blindness-detection/test.csv')\n\ntest_dataloader = torch.utils.data.DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=2)","metadata":{"execution":{"iopub.status.busy":"2022-03-27T06:47:51.887333Z","iopub.execute_input":"2022-03-27T06:47:51.887617Z","iopub.status.idle":"2022-03-27T06:47:51.896887Z","shell.execute_reply.started":"2022-03-27T06:47:51.887587Z","shell.execute_reply":"2022-03-27T06:47:51.896224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def testInTestDataset(dataloader, model):\n    \n    model.eval()\n    \n    labels = []\n    \n    with torch.no_grad():\n        \n        for batch,x in enumerate(dataloader):\n            output = model(x.to(device))\n            predictions = output.argmax(dim=1).cpu().detach().tolist() #Predicted labels for an image batch.\n            labels.extend(predictions)\n                \n    print('Testing has completed')\n            \n    return labels     ","metadata":{"execution":{"iopub.status.busy":"2022-03-27T06:47:54.299386Z","iopub.execute_input":"2022-03-27T06:47:54.299867Z","iopub.status.idle":"2022-03-27T06:47:54.305098Z","shell.execute_reply.started":"2022-03-27T06:47:54.299829Z","shell.execute_reply":"2022-03-27T06:47:54.30444Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = testInTestDataset(test_dataloader, model) ","metadata":{"execution":{"iopub.status.busy":"2022-03-27T06:47:57.810138Z","iopub.execute_input":"2022-03-27T06:47:57.811009Z","iopub.status.idle":"2022-03-27T06:49:17.358692Z","shell.execute_reply.started":"2022-03-27T06:47:57.81096Z","shell.execute_reply":"2022-03-27T06:49:17.357758Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nlabels = np.array(labels)\ntest_df = pd.read_csv('../input/aptos2019-blindness-detection/test.csv')\nsubmission_df = pd.DataFrame({'id_code':test_df['id_code'],'diagnosis':labels}) #DataFrame with id_code's and predicted labels.\n\nsubmission_df.to_csv('submission.csv',index=False) #csv file for submission.\nprint(\"Your submission was successfully saved!\")","metadata":{"execution":{"iopub.status.busy":"2022-03-27T06:51:12.58318Z","iopub.execute_input":"2022-03-27T06:51:12.583446Z","iopub.status.idle":"2022-03-27T06:51:12.604991Z","shell.execute_reply.started":"2022-03-27T06:51:12.583416Z","shell.execute_reply":"2022-03-27T06:51:12.604172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}