{"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":"code","source":"!pip3 install torchvision==0.11.1+cu113  -f https://download.pytorch.org/whl/cu113/torch_stable.html\n\n\n","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:25:19.671958Z","iopub.execute_input":"2022-07-03T08:25:19.672312Z","iopub.status.idle":"2022-07-03T08:28:13.325283Z","shell.execute_reply.started":"2022-07-03T08:25:19.672231Z","shell.execute_reply":"2022-07-03T08:28:13.324381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv\nimport os\nimport torchvision\nimport matplotlib.pyplot as plt #For plotting.\nimport PIL.Image as Image #For working with image files.\n#Importing torch\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\nfrom torch.utils.data import Dataset,DataLoader #For working with data.\nimport statistics\n\nfrom torchvision import models,transforms #For pretrained models,image transformations.\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu') #Use GPU if it's available or else use CPU.\nprint(device) #Prints the device we're using.","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:30:43.850623Z","iopub.execute_input":"2022-07-03T08:30:43.851341Z","iopub.status.idle":"2022-07-03T08:30:43.864168Z","shell.execute_reply.started":"2022-07-03T08:30:43.8513Z","shell.execute_reply":"2022-07-03T08:30:43.861786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torchvision.__version__","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:28:32.876231Z","iopub.execute_input":"2022-07-03T08:28:32.876437Z","iopub.status.idle":"2022-07-03T08:28:32.883188Z","shell.execute_reply.started":"2022-07-03T08:28:32.876403Z","shell.execute_reply":"2022-07-03T08:28:32.882497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%mkdir /kaggle/working/outputs\n\npath = \"/kaggle/input/aptos2019-blindness-detection/\"\n\ntrain_df = pd.read_csv(f\"{path}train.csv\")\nprint(f'No.of.training_samples: {len(train_df)}')\n\ntest_df = pd.read_csv(f'{path}test.csv')\nprint(f'No.of.testing_samples: {len(test_df)}')","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:28:39.994798Z","iopub.execute_input":"2022-07-03T08:28:39.995228Z","iopub.status.idle":"2022-07-03T08:28:40.783616Z","shell.execute_reply.started":"2022-07-03T08:28:39.995182Z","shell.execute_reply":"2022-07-03T08:28:40.782798Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n%matplotlib inline\nrgb_list = []\nimage = cv2.imread('../input/aptos2019-blindness-detection/train_images/0124dffecf29.png')\nRGB_img = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\nrgb_list.append(RGB_img)\nplt.imshow(RGB_img)\n\nGaussian = cv2.GaussianBlur(image, (7, 7), 0)\nRGB_img = cv2.cvtColor(Gaussian, cv2.COLOR_BGR2RGB)\nrgb_list.append(RGB_img)\nplt.imshow(RGB_img)\n\nmedian = cv2.medianBlur(image, 5)\nRGB_img = cv2.cvtColor(median, cv2.COLOR_BGR2RGB)\nrgb_list.append(RGB_img)\nplt.imshow(RGB_img)\n  \nbilateral = cv2.bilateralFilter(image, 9, 75, 75)\nRGB_img = cv2.cvtColor(bilateral, cv2.COLOR_BGR2RGB)\nrgb_list.append(RGB_img)\nplt.imshow(RGB_img)\n\n  \n# generation of a dictionary of (title, images)\n\ndef plot_figures(figures, nrows = 1, ncols=1):\n    \"\"\"Plot a dictionary of figures.\n\n    Parameters\n    ----------\n    figures : <title, figure> dictionary\n    ncols : number of columns of subplots wanted in the display\n    nrows : number of rows of subplots wanted in the figure\n    \"\"\"\n\n    fig, axeslist = plt.subplots(ncols=ncols, nrows=nrows)\n    for ind,title in enumerate(figures):\n        axeslist.ravel()[ind].imshow(figures[title], cmap=plt.gray())\n        axeslist.ravel()[ind].set_title(title)\n        axeslist.ravel()[ind].set_axis_off()\n        f = plt.gcf()\n        f.set_figwidth(15) # Sets overall figure width to 10 inches\n        f.set_figheight(15) # Sets overall figure height to 10 inches    \n    plt.tight_layout() # optional\n    #lt.show()\n\nnumber_of_im = 4\ntitle = ['original','gaussian blur', 'median blur', 'bilateral blur']\n\nfigures = { title[i]:rgb_list[i]  for i in range(number_of_im)}\n\n# plot of the images in a figure, with 2 rows and 3 columns\nplot_figures(figures, 1,4)\n","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:28:40.785871Z","iopub.execute_input":"2022-07-03T08:28:40.786434Z","iopub.status.idle":"2022-07-03T08:28:50.020136Z","shell.execute_reply.started":"2022-07-03T08:28:40.786394Z","shell.execute_reply":"2022-07-03T08:28:50.019017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df.diagnosis.hist()\nplt.xticks([0,1,2,3,4])\nplt.grid(False)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:28:50.021998Z","iopub.execute_input":"2022-07-03T08:28:50.022485Z","iopub.status.idle":"2022-07-03T08:28:50.329622Z","shell.execute_reply.started":"2022-07-03T08:28:50.02245Z","shell.execute_reply":"2022-07-03T08:28:50.328919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#As you can see,the data is imbalanced.\n#So we've to calculate weights for each class,which can be used in calculating loss.\n\nfrom sklearn.utils import class_weight #For calculating weights for each class.\nclass_weights = class_weight.compute_class_weight(class_weight='balanced',classes=np.array([0,1,2,3,4]),y=train_df['diagnosis'].values)\nclass_weights = torch.tensor(class_weights,dtype=torch.float).to(device)\n \nprint(class_weights) #Prints the calculated weights for the classes.","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:28:50.330916Z","iopub.execute_input":"2022-07-03T08:28:50.33189Z","iopub.status.idle":"2022-07-03T08:28:54.177192Z","shell.execute_reply.started":"2022-07-03T08:28:50.331844Z","shell.execute_reply":"2022-07-03T08:28:54.176399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"dataset preprocessing","metadata":{}},{"cell_type":"code","source":"class dataset(Dataset): # Inherits from the Dataset class.\n    '''\n    dataset class overloads the __init__, __len__, __getitem__ methods of the Dataset class. \n    \n    Attributes :\n        df:  DataFrame object for the csv file.\n        data_path: Location of the dataset.\n        image_transform: Transformations to apply to the image.\n        train: A boolean indicating whether it is a training_set or not.\n    '''\n    \n    def __init__(self,df,data_path,image_transform=None,train=True): # Constructor.\n        super(Dataset,self).__init__() #Calls the constructor of the Dataset class.\n        self.df = df\n        self.data_path = data_path\n        self.image_transform = image_transform\n        self.train = train\n        \n    def __len__(self):\n        return len(self.df) #Returns the number of samples in the dataset.\n    \n    def __getitem__(self,index):\n        image_id = self.df['id_code'][index]\n        image = Image.open(f'{self.data_path}/{image_id}.png') #Image.\n        if self.image_transform :\n            image = self.image_transform(image) #Applies transformation to the image.\n        \n        if self.train :\n            label = self.df['diagnosis'][index] #Label.\n            return image,label #If train == True, return image & label.\n        \n        else:\n            return image #If train != True, return image.","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:28:54.178961Z","iopub.execute_input":"2022-07-03T08:28:54.179671Z","iopub.status.idle":"2022-07-03T08:28:54.18939Z","shell.execute_reply.started":"2022-07-03T08:28:54.17963Z","shell.execute_reply":"2022-07-03T08:28:54.188527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMAGE_SIZE = 224 # Image size of resize when applying transforms.\nBATCH_SIZE = 32 \nNUM_WORKERS = 2 # Number of parallel processes for data preparation.\ntrain_image_transform = transforms.Compose([\n        transforms.Resize((IMAGE_SIZE, IMAGE_SIZE)),\n        transforms.RandomHorizontalFlip(p=0.5),\n        transforms.RandomVerticalFlip(p=0.5),\n        transforms.GaussianBlur(kernel_size=(5, 9), sigma=(0.1, 5)),\n        transforms.RandomAdjustSharpness(sharpness_factor=2, p=0.5),\n        transforms.ToTensor(),\n        transforms.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225]\n            )\n    ])\n\n#forming the dataset \n\ndata_set = dataset(train_df,f'{path}train_images',image_transform=train_image_transform)\n\n#splitting into train and validation\n\ntrain_set,valid_set = torch.utils.data.random_split(data_set,[3302,360])\ntrain_set, test_set = torch.utils.data.random_split(train_set,[2972,330])\n\n#dataloaders \n\ntrain_loader = DataLoader(\n        train_set, batch_size=BATCH_SIZE, \n        shuffle=True, num_workers=NUM_WORKERS\n    )\nvalid_loader = DataLoader(\n        valid_set, batch_size=BATCH_SIZE, \n        shuffle=False, num_workers=NUM_WORKERS\n    )\n\n","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-07-03T08:28:54.191056Z","iopub.execute_input":"2022-07-03T08:28:54.191357Z","iopub.status.idle":"2022-07-03T08:28:54.204616Z","shell.execute_reply.started":"2022-07-03T08:28:54.191322Z","shell.execute_reply":"2022-07-03T08:28:54.203907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model = models.resnet34(pretrained=True)\n#model.fc.out_features = 5\n#print(model)","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:29:04.123372Z","iopub.execute_input":"2022-07-03T08:29:04.12364Z","iopub.status.idle":"2022-07-03T08:29:04.126657Z","shell.execute_reply.started":"2022-07-03T08:29:04.12361Z","shell.execute_reply":"2022-07-03T08:29:04.12593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model = models.densenet121(pretrained=True)\n#model.classifier = nn.Sequential(\n           #nn.Linear(1024, 128),\n           #nn.ReLU(inplace=True),\n           #nn.Linear(128, 5)).to(device)\n        \ncount = 0\n\"\"\"\nfor child in model.children():\n    count+=1\n    print(\"CHILD \", child)\n    #for param in child.parameters():\n        #print(\"PARAM: \", param)\nprint(\"count: \", count)\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:29:04.464662Z","iopub.execute_input":"2022-07-03T08:29:04.465202Z","iopub.status.idle":"2022-07-03T08:29:04.470436Z","shell.execute_reply.started":"2022-07-03T08:29:04.465164Z","shell.execute_reply":"2022-07-03T08:29:04.46977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#model.classifier.out_features = 5","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:29:04.982008Z","iopub.execute_input":"2022-07-03T08:29:04.98246Z","iopub.status.idle":"2022-07-03T08:29:04.985969Z","shell.execute_reply.started":"2022-07-03T08:29:04.982425Z","shell.execute_reply":"2022-07-03T08:29:04.98528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#building model \n\ndef build_model(pretrained=True, fine_tune=True, num_classes=5):\n    if pretrained: \n        print('loading pretrained weights')\n    else: \n        print('not loading pretrained weights')\n    model = models.densenet121(pretrained=True).to(device)\n    if fine_tune:\n        print('[INFO]: Fine-tuning all layers...')\n        for params in model.parameters():\n            params.requires_grad = True\n    elif not fine_tune:\n        print('[INFO]: Freezing hidden layers...')\n        #ct = 0\n        #for child in model.children():\n            #ct += 1\n            #if ct < 10:\n                #for param in child.parameters():\n                    #param.requires_grad = False\n            #else: \n                #print(\"count = 10, requires grad is on for the last layer \")\n        for params in model.parameters():\n            params.requires_grad = False\n        model.classifier = nn.Sequential(\n               nn.Linear(1024, 128),\n               nn.ReLU(inplace=True),\n               nn.Linear(128, 5)).to(device)\n        \n    #\n\n    \n    #model.fc.out_features = 5\n    return model\n","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:29:05.850144Z","iopub.execute_input":"2022-07-03T08:29:05.850813Z","iopub.status.idle":"2022-07-03T08:29:05.858787Z","shell.execute_reply.started":"2022-07-03T08:29:05.850772Z","shell.execute_reply":"2022-07-03T08:29:05.857964Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#train\nimport torch\nimport argparse\nimport torch.nn as nn\nimport torch.optim as optim\nimport time\nfrom tqdm.auto import tqdm\n\n\n#model = models.resnet34(pretrained=True)\n\n    \n#model.fc.out_features = 5\n#model.to(device)\n\ndef train(trainloader, model, criterion,  optimizer):\n    model.train()\n    print('Training')\n    train_running_loss = 0.0\n    train_running_correct = 0\n    counter = 0\n    for i, data in tqdm(enumerate(trainloader), total=len(trainloader)):\n        counter += 1\n        image, labels = data\n        image = image.to(device)\n        labels = labels.to(device)\n        optimizer.zero_grad()\n        # Forward pass.\n        outputs = model(image)\n        # Calculate the loss.\n        loss = criterion(outputs, labels)\n        #print(\"losss: \", loss)\n        train_running_loss += loss.item()\n        # Calculate the accuracy.\n        _, preds = torch.max(outputs.data, 1)\n        train_running_correct += (preds == labels).sum().item()\n        # Backpropagation\n        loss.backward()\n        # Update the weights.\n        optimizer.step()\n    \n    # Loss and accuracy for the complete epoch.\n    epoch_loss = train_running_loss / counter\n    epoch_acc = 100. * (train_running_correct / len(trainloader.dataset))\n    return epoch_loss, epoch_acc\n\ndef validate( testloader, model, criterion):\n    model.eval()\n    print('Validation')\n    valid_running_loss = 0.0\n    valid_running_correct = 0\n    counter = 0\n    with torch.no_grad():\n        for i, data in tqdm(enumerate(testloader), total=len(testloader)):\n            counter += 1\n            \n            image, labels = data\n            image = image.to(device)\n            labels = labels.to(device)\n            # Forward pass.\n            outputs = model(image)\n            # Calculate the loss.\n            loss = criterion(outputs, labels)\n            valid_running_loss += loss.item()\n            # Calculate the accuracy.\n            _, preds = torch.max(outputs.data, 1)\n            valid_running_correct += (preds == labels).sum().item()\n        \n    # Loss and accuracy for the complete epoch.\n    epoch_loss = valid_running_loss / counter\n    epoch_acc = 100. * (valid_running_correct / len(testloader.dataset))\n    return epoch_loss, epoch_acc\n    \n\n\n\n","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:29:06.589869Z","iopub.execute_input":"2022-07-03T08:29:06.590112Z","iopub.status.idle":"2022-07-03T08:29:06.604382Z","shell.execute_reply.started":"2022-07-03T08:29:06.590081Z","shell.execute_reply":"2022-07-03T08:29:06.60244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport matplotlib\nimport matplotlib.pyplot as plt\nmatplotlib.style.use('ggplot')\ndef save_model(epochs, model, optimizer, criterion):\n    \"\"\"\n    Function to save the trained model to disk.\n    \"\"\"\n    torch.save({\n                'epoch': epochs,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'loss': criterion,\n                }, f\"/kaggle/working/outputs/model.pth\")\ndef save_plots(train_acc, valid_acc, train_loss, valid_loss):\n    \"\"\"\n    Function to save the loss and accuracy plots to disk.\n    \"\"\"\n    # accuracy plots\n    plt.figure(figsize=(10, 7))\n    plt.plot(\n        train_acc, color='green', linestyle='-', \n        label='train accuracy'\n    )\n    plt.plot(\n        valid_acc, color='blue', linestyle='-', \n        label='validataion accuracy'\n    )\n    plt.xlabel('Epochs')\n    plt.ylabel('Accuracy')\n    plt.legend()\n    plt.savefig(f\"/kaggle/working/outputs/accuracy.png\")\n    \n    # loss plots\n    plt.figure(figsize=(10, 7))\n    plt.plot(\n        train_loss, color='orange', linestyle='-', \n        label='train loss'\n    )\n    plt.plot(\n        valid_loss, color='red', linestyle='-', \n        label='validataion loss'\n    )\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.savefig(f\"/kaggle/working/outputs/loss.png\")","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:29:07.141941Z","iopub.execute_input":"2022-07-03T08:29:07.142195Z","iopub.status.idle":"2022-07-03T08:29:07.151768Z","shell.execute_reply.started":"2022-07-03T08:29:07.142165Z","shell.execute_reply":"2022-07-03T08:29:07.151032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == '__main__':\n    # Load the training and validation datasets.\n    lr = 0.0001\n    epochs = 40\n    device = ('cuda' if torch.cuda.is_available() else 'gpu')\n    print(f\"Computation device: {device}\")\n    print(f\"Learning rate: {lr}\")\n    print(f\"Epochs to train for: {epochs}\\n\")\n    \n    model = build_model(\n        pretrained=True,\n        fine_tune=False, \n        num_classes=5\n    ).to(device)\n    \n    # Total parameters and trainable parameters.\n    total_params = sum(p.numel() for p in model.parameters())\n    print(f\"{total_params:,} total parameters.\")\n    total_trainable_params = sum(\n        p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"{total_trainable_params:,} training parameters.\")\n    # Optimizer.\n    optimizer = optim.Adam(model.parameters(), lr=lr)\n    # Loss function.\n    criterion = nn.CrossEntropyLoss()\n    # Lists to keep track of losses and accuracies.\n    train_loss, valid_loss = [], []\n    train_acc, valid_acc = [], []\n    # Start the training.\n    for epoch in range(epochs):\n        print(f\"[INFO]: Epoch {epoch+1} of {epochs}\")\n        train_epoch_loss, train_epoch_acc = train( train_loader, model,  \n                                                criterion, optimizer)\n        valid_epoch_loss, valid_epoch_acc = validate(valid_loader, model,  \n                                                    criterion)\n        train_loss.append(train_epoch_loss)\n        valid_loss.append(valid_epoch_loss)\n        train_acc.append(train_epoch_acc)\n        valid_acc.append(valid_epoch_acc)\n        print(f\"Training loss: {train_epoch_loss:.3f}, training acc: {train_epoch_acc:.3f}\")\n        print(f\"Validation loss: {valid_epoch_loss:.3f}, validation acc: {valid_epoch_acc:.3f}\")\n        print('-'*50)\n        time.sleep(5)\n        \n    # Save the trained model weights.\n    save_model(epochs, model, optimizer, criterion)\n    # Save the loss and accuracy plots.\n    save_plots(train_acc, valid_acc, train_loss, valid_loss)\n    print('TRAINING COMPLETE')\n    print('avergae training accuracy of all folds: ', statistics.mean(train_acc))\n    print('avergae validation accuracy of all folds: ', statistics.mean(valid_acc))","metadata":{"execution":{"iopub.status.busy":"2022-07-03T08:32:02.832532Z","iopub.execute_input":"2022-07-03T08:32:02.832821Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":" #save_plots(train_acc, valid_acc, train_loss, valid_loss)","metadata":{"execution":{"iopub.status.busy":"2022-06-28T12:50:41.555357Z","iopub.execute_input":"2022-06-28T12:50:41.555642Z","iopub.status.idle":"2022-06-28T12:50:42.335361Z","shell.execute_reply.started":"2022-06-28T12:50:41.555607Z","shell.execute_reply":"2022-06-28T12:50:42.334527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"epochs = range(40)\nplt.plot(epochs, train_loss, 'g', label='Training loss')\nplt.plot(epochs, valid_loss, 'b', label='validation loss')\nplt.title('Training and Validation loss')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.legend()\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nfrom sklearn.metrics import classification_report\nimport seaborn as sns\n\ndef test(dataloader,model):\n    '''\n    test function predicts the labels given an image batches.\n    \n    Args :\n         dataloader: DataLoader for the test_set.\n         model: Given an input produces an output by multiplying the input with the model weights.\n         \n    Returns :\n         List of predicted labels.\n    '''\n    model = build_model(pretrained=False, fine_tune=False, num_classes=4)\n    checkpoint = torch.load('./outputs/model.pth', map_location=device)\n    #model.eval() #Sets the model for evaluation.\n    model.load_state_dict(checkpoint['model_state_dict'])\n    test_running_correct = 0\n    #labels = [] #List to store the predicted labels.\n    counter = 0\n    y_true = []\n    y_pred = []\n    #with torch.no_grad():\n        \n        #for batch,x in enumerate(dataloader):\n            \n            #output = model(x.to(device))\n            \n            #predictions = output.argmax(dim=1).cpu().detach().tolist() #Predicted labels for an image batch.\n            #labels.extend(predictions)\n    for i, data in tqdm(enumerate(dataloader), total=len(dataloader)):\n        counter += 1\n        image, labels = data\n        image = image.to(device)\n        labels = labels.to(device)\n        y_true.extend(labels.cpu().numpy())\n        outputs = model(image)\n        #outputs = outputs.detach().numpy()\n        _, preds = torch.max(outputs.data, 1)\n        y_pred.extend(preds.cpu().numpy())\n        test_running_correct += (preds == labels).sum().item()\n    epoch_acc = 100. * (test_running_correct / len(dataloader.dataset))\n    \n    print(classification_report(y_true, y_pred))\n    cf_matrix = confusion_matrix(y_true, y_pred)\n    class_names = ('no_dr', 'mild', 'moderate', 'severe', 'proliferate_dr')\n    dataframe = pd.DataFrame(cf_matrix, index=class_names, columns=class_names)\n    \n    plt.figure(figsize=(8, 6))\n    sns.heatmap(dataframe, annot=True, cbar=None,cmap=\"YlGnBu\",fmt=\"d\")\n\n    plt.title(\"Confusion Matrix\"), plt.tight_layout()\n\n    plt.ylabel(\"True Class\"), \n    plt.xlabel(\"Predicted Class\")\n    plt.show()\n    plt.savefig(f\"/kaggle/working/outputs/confusion_matrix.png\")\n                \n    print('Testing has completed : ', epoch_acc)\n            \n    #return labels               ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataloader = DataLoader(test_set, batch_size=32, shuffle=False)\ntest(test_dataloader,model)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}