{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Note Training and Visualisation Cell Are way down"},{"metadata":{"trusted":true},"cell_type":"code","source":"# # Set your own project id here\n# PROJECT_ID = 'your-google-cloud-project'\n# from google.cloud import storage\n# storage_client = storage.Client(project=PROJECT_ID)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#import data manipulation libraries\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport matplotlib.pyplot as plt","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load in \n\n# Input data files are available in the \"../input/\" directory.\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# for dirname, _, filenames in os.walk('/kaggle/input/diabetic-retinopathy-detection/'):\n#     for filename in filenames:\n# #         print(os.path.join(dirname, filename))\n#           pass\n# # Any results you write to the current directory are saved as output.","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Data Processing for DR-real"},{"metadata":{"trusted":true},"cell_type":"code","source":"#reading in file and checking the head\ndata_1 = pd.read_csv('/kaggle/input/diabetic-retinopathy-detection/trainLabels.csv.zip')\ndata_1.head()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Unpacking Dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"data_1.info()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#checking the unique labels\ndata_1.level.unique()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#checking for nan values\ndata_1.isna().count()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#Create Unzippd_dr_first\ntry:\n    os.mkdir('Unzippd_dr_first')\nexcept:\n    print('Dir Exist')\n    !ls","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#change dir to kaggle\nos.chdir('/kaggle/')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#checkout out the folders\n!ls","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# !rm -r train/","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#extract from input folder folder to output folder\n# !7z x input/diabetic-retinopathy-detection/train.zip.001","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!ls","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#extract from input folder folder to output folder\n# !7z x input/messedor-complete/IMAGES.zip.001","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dr_real = pd.read_csv(\"input/diabetic-retinopathy-detection/trainLabels.csv.zip\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dr_real.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dr_real.level.unique()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dr_real.level.value_counts()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Data Processing for Messidor-dr-grades"},{"metadata":{"trusted":true},"cell_type":"code","source":"!ls","execution_count":null,"outputs":[]},{"metadata":{"_cell_guid":"","_uuid":"","trusted":true},"cell_type":"code","source":"# MESSIDOR-2 DR Grades\nmessidor_data = pd.read_csv(\"input/messidor2-dr-grades/messidor_data.csv\")\nmessidor_data.tail()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"messidor_data.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"messidor_data.adjudicated_dr_grade.unique()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"messidor_data.adjudicated_dr_grade.value_counts()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dr_real.level.value_counts()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#Sum of Each Class That will be used for the project\ndr_real.level.value_counts() + messidor_data.adjudicated_dr_grade.value_counts()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#total sum of all the data\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Process Data"},{"metadata":{"_cell_guid":"","_uuid":"","trusted":true},"cell_type":"code","source":"import pandas as pd\nmessidor_2 = pd.read_csv(\"../input/messedor-complete/messidor-2.csv\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"messidor_2.tail()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"messidor_data.tail()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"875*2","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!ls","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"not_in = []\n#check for missing dataset in both files\nfor each in messidor_data.image_id:\n    try:\n        data = plt.imread('/kaggle/IMAGES/'+each)\n        print(each)\n    except:\n        data = plt.imread('/kaggle/IMAGES/'+each[:-3]+'JPG')\n        print(each[:-3]+'JPG')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#print out the missing values\nfor each in not_in:\n    d = each[:-3]+'JPG'\n    k = plt.imread('/kaggle/IMAGES/'+d)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"not_in","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Creating New Dataset csv File"},{"metadata":{"trusted":true},"cell_type":"code","source":"#creat a new csv File of the merged data from diabetic-retinopathy-detection dataset and messidor-2-fundus dataset\nnew_dset = pd.DataFrame(columns=['image_id', 'target'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"messidor_data","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#loop through the messidor-fundus dataset and add to the new csv file\nfor image_id,target in zip(messidor_data.image_id, messidor_data.adjudicated_dr_grade):\n    try:\n        data = plt.imread('/kaggle/IMAGES/'+image_id)\n        new_dset = new_dset.append({'image_id':'/kaggle/IMAGES/'+image_id,'target':target},ignore_index=True)\n    except:\n        data = plt.imread('/kaggle/IMAGES/'+image_id[:-3]+'JPG')\n        new_dset = new_dset.append({'image_id':'/kaggle/IMAGES/'+image_id[:-3]+'JPG','target':target},ignore_index=True)\n\n#loop through the diabetic-retinopathy-detection dataset and add to the new csv file\nfor image_id,target in zip(dr_real.image, dr_real.level):\n    data = plt.imread('/kaggle/train/'+image_id+'.jpeg')\n    new_dset = new_dset.append({'image_id':'/kaggle/train/'+image_id+'.jpeg','target':target},ignore_index=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"new_dset.iloc[35126:]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"new_dset.tail()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dr_real.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dr_real.shape[0] + messidor_data.shape[0]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"new_dset.shape[0]","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#save the new csv file\nnew_dset.to_csv('working/new_dataset_csv2.csv', index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"dtest = plt.imread('/kaggle/IMAGES/20051208_42314_0400_PP.png')\n# dtest2 = plt.imread('/kaggle/IMAGES/IM004765.JPG')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"plt.figure(figsize=(20,20))\nplt.imshow(dtest)\nplt.show()\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"plt.imshow(dtest2)\nplt.show()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!ls","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"os.chdir('train/')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!ls","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# #other codes i used locally\n# data_set = pd.read_csv('/home/vincycode7/Downloads/new_dataset_csv2 (2).csv')\n# data_set.head()\n# data_set.tail()\n# data_set[data_set.target.isnull()]\n# chk = data_set.dropna()\n# chk.info()\n\n# #reduce from class 5 to 3\n# _5_to_3 = lambda x: {0:0, 1:1, 2:1, 3:2, 4:2}.get(int(x))\n# chk['new_target'] = chk.target.map(_5_to_3)\n# chk.target.nunique()\n# chk.new_target.nunique()\n# #checking to see if all is well\n# chk[chk.target==5]\n# #shuffle and save new dataset\n# shuff_ds = chk.copy()\n\n# def shuffle_reldset(X=None,y=None,test_s=0.2,split_d=True):\n#     from sklearn.model_selection import StratifiedShuffleSplit\n#     #splitting dataset into train,remaining\n#     split = StratifiedShuffleSplit(test_size=test_s)\n#     for trn_idx, tst_idx in split.split(X,y):\n#         x_trn, x_tst = X[trn_idx], X[tst_idx]\n#         y_trn, y_tst = y[trn_idx], y[tst_idx]\n#     if split_d == True:\n#         split = StratifiedShuffleSplit(test_size=0.5)\n\n#         for trn_idx, tst_idx in split.split(x_tst,y_tst):\n#             x_val_, x_tst_ = X[trn_idx], X[tst_idx]\n#             y_val_, y_tst_ = y[trn_idx], y[tst_idx]\n#         return (x_trn, y_trn, x_val_, y_val_, x_tst_, y_tst_) \n#     else:\n#         return (x_trn, y_trn, x_tst, y_tst)\n    \n# (x_trn, y_trn, x_val_, y_val_, x_tst_, y_tst_)  = shuffle_reldset(X=shuff_ds.image_id, y=shuff_ds.new_target)\n\n# # for training set\n# train_ds = pd.DataFrame({'image_id':x_trn, 'target':y_trn})\n# train_ds.reset_index(inplace=True)\n# train_ds.drop(columns='index', inplace=True)\n# train_ds.head()\n\n# #for test set\n# test_ds = pd.DataFrame({'image_id':x_tst_, 'target':y_tst_})\n# test_ds.reset_index(inplace=True)\n# test_ds.drop(columns='index', inplace=True)\n# test_ds.head()\n\n# #for val set\n# val_ds = pd.DataFrame({'image_id':x_val_, 'target':y_val_})\n# val_ds.reset_index(inplace=True)\n# val_ds.drop(columns='index', inplace=True)\n# val_ds.head()\n\n# (train_ds.target.value_counts()/train_ds.target.value_counts().sum())*100,train_ds.target.value_counts()\n\n# (test_ds.target.value_counts()/test_ds.target.value_counts().sum())*100,test_ds.target.value_counts()\n\n# (val_ds.target.value_counts()/val_ds.target.value_counts().sum())*100,val_ds.target.value_counts()\n\n# #saving to csv file\n# trn_ds.to_csv('trainset.csv', index=False)\n# test_ds.to_csv('testset.csv', index=False)\n# val_ds.to_csv('valset.csv', index=False)\n\n# trn_ds = pd.read_csv('trainset.csv')\n# trn_ds['image_id'] = trn_ds.imagee_id\n# trn_ds = trn_ds[['image_id', 'target']]","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Before you Begin UnPacking The main Image might take a while"},{"metadata":{},"cell_type":"markdown","source":"You need to Unpack the images because they are in a Zip file, Code to unpack is down"},{"metadata":{},"cell_type":"markdown","source":" importing The need libraries for training and testing"},{"metadata":{"trusted":true},"cell_type":"code","source":"#load all your packages\n# import data manipulation libraries\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport matplotlib.pyplot as plt\n\n#import frameworks\nimport torch as tch\nimport tensorflow as tf\n# import pytorch_lightning as pi\n\n#image manipulation\nimport zipfile, os\nimport torch as tch\nfrom torchvision import datasets, models, transforms\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import StratifiedShuffleSplit\nfrom sklearn.preprocessing import OneHotEncoder\nimport numpy as np\nimport PIL\nimport csv\nfrom PIL import  Image\nfrom IPython.display import clear_output\nfrom IPython.core.debugger import set_trace","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Unpacking Dataset"},{"metadata":{},"cell_type":"markdown","source":"!Note don't run the next code if you not ready to unpack the dataset"},{"metadata":{"trusted":true},"cell_type":"code","source":"#change dir to kaggle\nos.chdir('/kaggle/')\n\n#unpack First image Dataset\n#extract from input folder folder to output folder\n!7z x input/diabetic-retinopathy-detection/train.zip.001\n\n#Unpack Second Image Dataset\n#extract from input folder folder to output folder\n!7z x input/messedor-complete/IMAGES.zip.001\n\n#change dir to kaggle\nos.chdir('/kaggle/working/')\n!ls","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# #load reference to dataset (train, val, test)\n# testset = pd.read_csv(\"/kaggle/input/train-val-test/testset.csv\")\n# trainset = pd.read_csv(\"/kaggle/input/train-val-test/trainset.csv\")\n# valset = pd.read_csv(\"/kaggle/input/train-val-test/valset.csv\")\n\n# #check shape\n# print(f'train set : {trainset.shape},\\nval set : {valset.shape},\\ntest set : {testset.shape}')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Creating Dataloader"},{"metadata":{"trusted":true},"cell_type":"code","source":"#data loader\n!pip install pytorch-lightning\n!pip install torchsummary\n# !pip install pytorch-model-summary\nimport os, torch\nimport torch as tch\nfrom torch.utils.data import Dataset, DataLoader\nfrom PIL import  Image\nimport matplotlib.pyplot as plt\nclass torch_dset(Dataset):\n\n  #initialize the dataloader instance\n  def __init__(self, csv_file=None, train=True, X=None, y=None):\n    data = pd.read_csv(csv_file)\n    if csv_file:\n        self.X, self.y = data.image_id,data.target\n    else:\n        self.X, self.y = X, y\n    \n    trans = {'train': transforms.Compose([\n#                                             transforms.RandomResizedCrop(224,scale=(0.7, 1.0)),\n                                            transforms.RandomResizedCrop(224),\n                                            transforms.RandomVerticalFlip(p=0.5),\n                                            transforms.RandomHorizontalFlip(),\n                                            transforms.ColorJitter(),\n                                            transforms.RandomRotation(30),\n                                            transforms.ToTensor(),\n                                            transforms.Normalize([0.485, 0.456, 0.406],\n                                                                 [0.229, 0.224, 0.225]),\n                                           ]),\n                'val' :transforms.Compose([ transforms.Resize((224,224)),\n                                         transforms.ToTensor()])\n              }\n    self.trans = trans['train'] if train else trans['val']\n\n  def __getitem__(self, idx):\n    #check if last sample \n    if idx == (self.__len__()):raise StopIteration\n    if torch.is_tensor(idx):idx = idx.tolist()\n      \n    x_,y_ = self.X[idx], self.y[idx]\n    trans = self.trans\n\n    #open images and resize\n    img = trans(Image.open(x_)).float()\n    y_ = torch.tensor(y_.astype('float32')).long()\n    return img, y_\n\n  def __len__(self):\n    return len(self.X)\n               \n# dataset_sizes = {x: len(dataloaders[x]) for x in ['train', 'val']}\n# device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Creating the Pytorch Lightning Model class"},{"metadata":{"trusted":true},"cell_type":"code","source":"#model creation\n# import libraries\nimport numpy as numpy\nimport pandas as pd\nimport matplotlib.pyplot as pyplot\nfrom sklearn.metrics import accuracy_score\n\n\nimport torch as tch\nimport torch\nimport torchvision as tv\nfrom torchvision.transforms import transforms\nfrom torchvision import datasets\nfrom torchvision.datasets import MNIST\nfrom torchsummary import summary\nimport pytorch_lightning as pl\nfrom pytorch_lightning import Trainer\nfrom torch.nn import functional as F\nfrom torch import nn\nfrom torch.nn import Sequential,Dropout, Linear, Identity\nfrom collections import OrderedDict\nfrom torch.utils.data import DataLoader\nfrom torch.optim import SGD\nimport os\n\n#building the pytorch lighting class Model\nclass Dia_ret(pl.LightningModule):\n    def __init__(self,model_name=None,every_b=3,best_acc=0.0):\n        super(Dia_ret, self).__init__()\n\n        self.prepare_data()\n        self.n_batch,self.n_v_batch,self.trn_acc,self.trn_lss, self.val_acc, self.val_lss=0,0,0,0,0,0\n        self.model_name,self.every_b = model_name, every_b\n        self.best_acc,self.epoch,self.still_epoch = best_acc,0,False\n        self.device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n        self.feature2 = models.resnet34(pretrained=False)\n        self.feature2.fc =  Sequential(OrderedDict([\n                                          ('drop', Dropout(p=0.25)),\n                                          ('cassifier1', Linear(512,3)),\n        ]))\n\n    def forward(self, x):\n        x = self.feature2(x.view(-1, 3, 224, 224))\n        return x\n\n    def configure_optimizers(self):\n        return SGD([{'params' : self.feature2.parameters(), 'lr':0.001}],\n                            lr=0.001, momentum=0.9)\n\n    def my_loss(self, y_hat, y):\n        return F.cross_entropy(y_hat, y)\n\n    def training_step(self, batch, batch_size):\n        if self.n_batch == 0 and not self.still_epoch:\n          clear_output()\n          self.still_epoch = True\n          print(f'<============= Epoch {self.epoch} ==============>')\n        x,y = batch\n        #move to cuda\n        x.to(self.device)\n        self.to(self.device)\n        \n        #forward pass\n        logits = self.forward(x)\n        \n        #get predictions and loss\n        _, preds = torch.max(logits, 1)\n        loss = self.my_loss(logits.cpu(), y.cpu())\n        acc = torch.tensor(accuracy_score(preds.cpu(), y.cpu()))\n        #add up train loss and accuracy\n        self.trn_lss += loss\n        self.trn_acc += acc\n        \n        #save model state\n        model_info = {'state_dict_cnn' : model.state_dict(),'best_acc':self.best_acc}\n        torch.save(model_info, './models'+self.model_name+'_train.pt')\n        \n        #print out something if it is on every third batch\n        if self.n_batch%self.every_b==(self.every_b-1):\n            b_n = int(self.n_batch/self.every_b)+1\n            of_ = self.dataset_sizes['train']//self.every_b\n            print(f'batch {b_n}/{of_} loss {self.trn_lss/(b_n*self.every_b)} acc {self.trn_acc/(b_n*self.every_b)}')\n        self.n_batch+=1\n        return {'loss' : loss}\n        # return loss (also works)\n\n    def validation_step(self, batch, batch_size):\n        x,y = batch\n        x.to(self.device)\n        self.to(self.device)\n        logits = self.forward(x)\n        \n        #get predictions and loss\n        _, preds = torch.max(logits, 1)\n        loss = self.my_loss(logits.cpu(), y.cpu())\n        acc = torch.tensor(accuracy_score(preds.cpu(), y.cpu()))\n        #add up train loss and accuracy\n        self.val_lss += loss\n        self.val_acc += acc\n        \n        #print out something if it is on every third batch\n        if self.n_v_batch%self.every_b==(self.every_b-1):\n            b_n = int(self.n_v_batch/self.every_b)+1\n            of_ = self.dataset_sizes['val']//self.every_b\n            print(f'batch {b_n}/{of_} loss {self.val_lss/(b_n*self.every_b)} acc {self.val_acc/(b_n*self.every_b)}')\n        self.n_v_batch += 1\n        return {'val_loss' :loss, 'val_acc':acc}\n\n    def test_step(self, batch, batch_size):\n        x,y = batch\n        x.to(self.device)\n        self.to(self.device)\n        logits = self.forward(x)\n        _, preds = torch.max(logits, 1)\n        acc = torch.tensor(accuracy_score(preds.cpu(), y.cpu()))\n        return {'test_loss' : self.my_loss(logits, y), 'test_acc':acc}\n\n    def validation_epoch_end(self, outputs):\n        # OPTIONAL\n        self.n_batch,self.n_v_batch= 0,0\n        self.epoch += 1\n        self.still_epoch = False\n        self.trn_lss,self.trn_acc,self.val_lss,self.val_acc = 0,0,0,0\n        avg_loss = torch.stack([x['val_loss'] for x in outputs]).mean()\n        tensorboard_logs = {'val_loss': avg_loss}\n        avg_acc = torch.stack([x['val_acc'] for x in outputs]).mean()\n        print(f'avg_val_loss:  {avg_loss} avg_val_acc : {avg_acc}')\n        if avg_acc > self.best_acc:\n          self.best_acc = avg_acc\n          model_info = {'state_dict_cnn' : model.state_dict(),'best_acc':self.best_acc}\n          torch.save(model_info, './models'+self.model_name+'_val.pt')\n        return {'avg_val_loss': avg_loss, 'log': tensorboard_logs,'avg_val_acc':avg_acc}\n\n    def test_epoch_end(self, outputs):\n        # OPTIONAL\n        avg_loss = torch.stack([x['test_loss'] for x in outputs]).mean()\n        logs = {'test_loss': avg_loss}\n        avg_acc = torch.stack([x['test_acc'] for x in outputs]).mean()\n        return {'avg_test_loss': avg_loss, 'log': logs, 'progress_bar': logs,'avg_test_acc':avg_acc}\n\n    \n\n    def prepare_data(self):\n        # download only        \n        train = torch_dset(csv_file=\"/kaggle/input/train-val-test/trainset.csv\")\n        val = torch_dset(csv_file=\"/kaggle/input/train-val-test/valset.csv\")\n        test = torch_dset(csv_file=\"/kaggle/input/train-val-test/testset.csv\")\n\n        self.dataloaders = {'train':DataLoader(train, batch_size=64, shuffle=True, num_workers=1), \n                            'val':  DataLoader(val, batch_size=32, shuffle=True, num_workers=1),\n                            'test' :DataLoader(test, batch_size=20, shuffle=True, num_workers=1)}\n        self.dataset_sizes = {x: len(self.dataloaders[x]) for x in ['train', 'val','test']}\n    def train_dataloader(self):\n        #convert data to torch.FloatTensor\n        #load the training dataset\n        return self.dataloaders['train']\n\n    def val_dataloader(self):\n        return self.dataloaders['val']\n\n    def test_dataloader(self):\n        return self.dataloaders['test']\n    \n# create Instance of the Object Autoencoder\nstates = torch.load('../input/modelnew/modelsmodel1_train (6).pt',map_location='cpu')\n# states = torch.load('modelsmodel0_train.pt',map_location='cpu')\nmodel = Dia_ret(model_name='model1',every_b=3,best_acc=states['best_acc'])\nmodel.load_state_dict(states['state_dict_cnn'])\nfor param in model.parameters():\n    param.requires_grad_(True)\nmodel.cuda()\nsummary(model, input_size=(3,224,224))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Creating the data Visualization Function"},{"metadata":{"trusted":true},"cell_type":"code","source":"# obtain one batch of test images\nimport matplotlib.pyplot as plt\nimport matplotlib.gridspec as gridspec\nimport numpy as np\ndef viz(kwarg,rows=4,cols=5):\n    images, labels = kwarg\n    # # plot the first ten input images and then reconstructed images\n    fig, axes = plt.subplots(nrows=rows, ncols=cols, sharex=True, sharey=True, figsize=(10,10))\n\n    # # input images on top row, reconstructions on bottom\n    for each_row in range(rows):\n        for each_col in range(cols):\n            idx = cols*each_row\n            img = images[idx+each_col]\n            ax = axes[each_row, each_col]\n            ax.imshow(img.numpy().transpose((1,2,0)))\n            ax.get_xaxis().set_visible(False)\n            ax.get_yaxis().set_visible(False)\n            ax.title.set_text(labels[idx+each_col])\n            plt.tight_layout()\n    print(labels)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"#loading all the datset into a dataloader"},{"metadata":{"trusted":true},"cell_type":"code","source":"ds = iter(model.test_dataloader())","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"#visualizing the dataset(keep rerunning the cell to load diffferent images)"},{"metadata":{"trusted":true},"cell_type":"code","source":"viz(ds.next())","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# Creater a Trainer to train the Model"},{"metadata":{"trusted":true},"cell_type":"code","source":"#trainer\ntrainer = Trainer(max_epochs=1, \n                  check_val_every_n_epoch=1,\n                  gpus=-1,\n                  reload_dataloaders_every_epoch=True\n                 )","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# train the model"},{"metadata":{"trusted":true},"cell_type":"code","source":"trainer.fit(model=model)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"#save the model state in a variable"},{"metadata":{"trusted":true},"cell_type":"code","source":"state_dict = model.state_dict()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"model","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"load the model state back into the model"},{"metadata":{"trusted":true},"cell_type":"code","source":"model.load_state_dict(state_dict)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"#test the model on the train set"},{"metadata":{"trusted":true},"cell_type":"code","source":"trainer.test(model)","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}