{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"colab":{"name":"training.ipynb","provenance":[],"include_colab_link":true},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431},{"sourceType":"datasetVersion","sourceId":945851,"datasetId":512846,"databundleVersionId":973598}],"dockerImageVersionId":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<a href=\"https://colab.research.google.com/github/souravs17031999/Retinal_blindness_detection_Pytorch/blob/master/training.ipynb\" target=\"_parent\"><img src=\"https://colab.research.google.com/assets/colab-badge.svg\" alt=\"Open In Colab\"/></a>","metadata":{"id":"view-in-github"}},{"cell_type":"markdown","source":"# Import the essentials","metadata":{"id":"4lVyi6I3wp9-"}},{"cell_type":"code","source":"# Imports here\nfrom __future__ import print_function, division\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom torch.utils import data\nimport torch\nfrom torch import nn\nfrom torch import optim\nimport torchvision\nimport torch.nn.functional as F\nfrom torchvision import datasets, transforms, models\nimport torchvision.models as models\nfrom torch.utils.data.sampler import SubsetRandomSampler\nfrom torch.utils.data import Dataset, DataLoader\nfrom skimage import io, transform\nimport torch.utils.data as data_utils\nfrom PIL import Image, ImageFile\nimport json\nfrom torch.optim import lr_scheduler\nimport time\nimport os\nimport argparse\nimport copy\nimport pandas as pd\nImageFile.LOAD_TRUNCATED_IMAGES = True\nimport cv2\n# Import useful sklearn functions\nimport sklearn\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\n\nimport time\nfrom tqdm import tqdm_notebook\n\n# import os\n# print(os.listdir(\"../input\"))\nbase_dir = \"/kaggle/input/aptos2019-blindness-detection\"","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","id":"-H5XMQt5wp-B","outputId":"8405258a-6edb-48ef-efde-fe87e9d38934","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:05.226387Z","iopub.execute_input":"2025-03-23T17:42:05.226685Z","iopub.status.idle":"2025-03-23T17:42:05.232581Z","shell.execute_reply.started":"2025-03-23T17:42:05.226660Z","shell.execute_reply":"2025-03-23T17:42:05.231567Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(os.listdir(\"../input\"))\n# base_dir = \"../input/aptos2019-blindness-detection/\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:07.901295Z","iopub.execute_input":"2025-03-23T17:42:07.901613Z","iopub.status.idle":"2025-03-23T17:42:07.906699Z","shell.execute_reply.started":"2025-03-23T17:42:07.901585Z","shell.execute_reply":"2025-03-23T17:42:07.905758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(os.listdir(\"../input/kernel4f121f3247\"))","metadata":{"id":"EHKTScVRwp-x","outputId":"aa412cf4-4dcc-4fec-b2b3-1d39e21eb7ed","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:10.851370Z","iopub.execute_input":"2025-03-23T17:42:10.851656Z","iopub.status.idle":"2025-03-23T17:42:10.855504Z","shell.execute_reply.started":"2025-03-23T17:42:10.851634Z","shell.execute_reply":"2025-03-23T17:42:10.854578Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import seaborn as sns","metadata":{"id":"5yxnwe8jwp-2","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:12.196381Z","iopub.execute_input":"2025-03-23T17:42:12.196673Z","iopub.status.idle":"2025-03-23T17:42:12.200280Z","shell.execute_reply.started":"2025-03-23T17:42:12.196651Z","shell.execute_reply":"2025-03-23T17:42:12.199409Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loading Data + EDA","metadata":{"id":"Mkb_neoqwp-6"}},{"cell_type":"code","source":"train_csv = pd.read_csv('/kaggle/input/aptos2019-blindness-detection/train.csv')\ntest_csv = pd.read_csv('/kaggle/input/aptos2019-blindness-detection/test.csv')","metadata":{"_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","id":"hkyWGmomwp-7","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:13.637638Z","iopub.execute_input":"2025-03-23T17:42:13.638065Z","iopub.status.idle":"2025-03-23T17:42:13.649260Z","shell.execute_reply.started":"2025-03-23T17:42:13.638030Z","shell.execute_reply":"2025-03-23T17:42:13.648287Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print('Train Size = {}'.format(len(train_csv)))\nprint('Public Test Size = {}'.format(len(test_csv)))","metadata":{"id":"pkd7u6RRwp-9","outputId":"7fdbb7a4-25c1-43c7-f6f2-d9787b1c7ba2","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:14.141247Z","iopub.execute_input":"2025-03-23T17:42:14.141527Z","iopub.status.idle":"2025-03-23T17:42:14.146043Z","shell.execute_reply.started":"2025-03-23T17:42:14.141505Z","shell.execute_reply":"2025-03-23T17:42:14.145164Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_csv.head()","metadata":{"id":"jgR8VptTwp_C","outputId":"898edfe9-0fd2-4edc-c23a-14a3b5d295be","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:15.196058Z","iopub.execute_input":"2025-03-23T17:42:15.196350Z","iopub.status.idle":"2025-03-23T17:42:15.203919Z","shell.execute_reply.started":"2025-03-23T17:42:15.196327Z","shell.execute_reply":"2025-03-23T17:42:15.203111Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"counts = train_csv['diagnosis'].value_counts()\nclass_list = ['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferate']\n\n# Map numeric indices to class names\ncounts = counts.reindex(range(len(class_list))).fillna(0).astype(int)\ncounts.index = class_list\n\nplt.figure(figsize=(10,5))\nsns.barplot(x=counts.index, y=counts.values, alpha=0.8, palette='bright')\nplt.title('Distribution of Output Classes')\nplt.ylabel('Number of Occurrences', fontsize=12)\nplt.xlabel('Target Classes', fontsize=12)\nplt.show()","metadata":{"id":"vgygoo4Vwp_F","outputId":"dff17c47-0a34-45ce-ff5f-83fcf0463047","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:16.046049Z","iopub.execute_input":"2025-03-23T17:42:16.046339Z","iopub.status.idle":"2025-03-23T17:42:16.225353Z","shell.execute_reply.started":"2025-03-23T17:42:16.046317Z","shell.execute_reply":"2025-03-23T17:42:16.224447Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualizing Training Data","metadata":{"id":"k7jDZc3twp_H"}},{"cell_type":"code","source":"fig = plt.figure(figsize=(30, 6))\n# display 20 images\ntrain_imgs = os.listdir('/kaggle/input/aptos2019-blindness-detection/train_images')\nfor idx, img in enumerate(np.random.choice(train_imgs, 16)):\n    ax = fig.add_subplot(2, 16//2, idx+1, xticks=[], yticks=[])\n    im = Image.open(base_dir+\"/train_images/\" + img)\n    plt.imshow(im)\n    lab = train_csv.loc[train_csv['id_code'] == img.split('.')[0], 'diagnosis'].values[0]\n    ax.set_title('Severity: %s'%lab)","metadata":{"id":"Ax11I6Zswp_I","outputId":"38d1e936-008b-4703-8ea9-110ed0880eb3","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:36.086230Z","iopub.execute_input":"2025-03-23T17:42:36.086552Z","iopub.status.idle":"2025-03-23T17:42:46.575442Z","shell.execute_reply.started":"2025-03-23T17:42:36.086525Z","shell.execute_reply":"2025-03-23T17:42:46.574566Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualizing Test Set","metadata":{"id":"1tw7S3-ywp_L"}},{"cell_type":"code","source":"fig = plt.figure(figsize=(30, 6))\n# display 20 images\ntest_imgs = os.listdir(base_dir+\"/test_images\")\nfor idx, img in enumerate(np.random.choice(test_imgs, 16)):\n    ax = fig.add_subplot(2, 16//2, idx+1, xticks=[], yticks=[])\n    im = Image.open(base_dir+\"/test_images/\" + img)\n    plt.imshow(im)","metadata":{"id":"gmy9BNFwwp_M","outputId":"de5d12e0-c57c-451c-e297-5da3d12a3fbd","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:46.576844Z","iopub.execute_input":"2025-03-23T17:42:46.577254Z","iopub.status.idle":"2025-03-23T17:42:52.769403Z","shell.execute_reply.started":"2025-03-23T17:42:46.577217Z","shell.execute_reply":"2025-03-23T17:42:52.768442Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Processing","metadata":{"id":"1ahpFGSqwp_O"}},{"cell_type":"code","source":"# Our own custom class for datasets\nclass CreateDataset(Dataset):\n    def __init__(self, df_data, data_dir = '../input/', transform=None):\n        super().__init__()\n        self.df = df_data.values\n        self.data_dir = data_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_name,label = self.df[index]\n        img_path = os.path.join(self.data_dir, img_name+'.png')\n        image = cv2.imread(img_path)\n        if self.transform is not None:\n            image = self.transform(image)\n        return image, label","metadata":{"id":"QfvxpUdcwp_P","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:52.771151Z","iopub.execute_input":"2025-03-23T17:42:52.771485Z","iopub.status.idle":"2025-03-23T17:42:52.777820Z","shell.execute_reply.started":"2025-03-23T17:42:52.771454Z","shell.execute_reply":"2025-03-23T17:42:52.777080Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_transforms = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(p=0.4),\n    #transforms.ColorJitter(brightness=2, contrast=2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225))\n])","metadata":{"id":"usFGpYV9wp_R","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:52.778597Z","iopub.execute_input":"2025-03-23T17:42:52.778844Z","iopub.status.idle":"2025-03-23T17:42:52.793446Z","shell.execute_reply.started":"2025-03-23T17:42:52.778824Z","shell.execute_reply":"2025-03-23T17:42:52.792543Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_transforms = transforms.Compose([transforms.Resize(256),\n                                      transforms.CenterCrop(224),\n                                      transforms.ToTensor(),\n                                      transforms.Normalize([0.485, 0.456, 0.406],[0.229, 0.224, 0.225])])","metadata":{"id":"3LvQYjjzwp_U","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:42:59.491414Z","iopub.execute_input":"2025-03-23T17:42:59.491719Z","iopub.status.idle":"2025-03-23T17:42:59.495980Z","shell.execute_reply.started":"2025-03-23T17:42:59.491692Z","shell.execute_reply":"2025-03-23T17:42:59.494992Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_path = \"/kaggle/input/aptos2019-blindness-detection/train_images\"\ntest_path = \"/kaggle/input/aptos2019-blindness-detection/test_images\"","metadata":{"id":"mF_yP6FVwp_W","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:04.991450Z","iopub.execute_input":"2025-03-23T17:43:04.991796Z","iopub.status.idle":"2025-03-23T17:43:04.995500Z","shell.execute_reply.started":"2025-03-23T17:43:04.991768Z","shell.execute_reply":"2025-03-23T17:43:04.994580Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = CreateDataset(df_data=train_csv, data_dir=train_path, transform=train_transforms)\ntest_data = CreateDataset(df_data=test_csv, data_dir=test_path, transform=test_transforms)","metadata":{"id":"rkUVBpU7wp_Z","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:24.021393Z","iopub.execute_input":"2025-03-23T17:43:24.021685Z","iopub.status.idle":"2025-03-23T17:43:24.025963Z","shell.execute_reply.started":"2025-03-23T17:43:24.021665Z","shell.execute_reply":"2025-03-23T17:43:24.025016Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"valid_size = 0.2\nnum_train = len(train_data)\nindices = list(range(num_train))\nnp.random.shuffle(indices)\nsplit = int(np.floor(valid_size * num_train))\ntrain_idx, valid_idx = indices[split:], indices[:split]","metadata":{"id":"BE_W-n38wp_c","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:24.951397Z","iopub.execute_input":"2025-03-23T17:43:24.951687Z","iopub.status.idle":"2025-03-23T17:43:24.956763Z","shell.execute_reply.started":"2025-03-23T17:43:24.951667Z","shell.execute_reply":"2025-03-23T17:43:24.955861Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_sampler = SubsetRandomSampler(train_idx)\nvalid_sampler = SubsetRandomSampler(valid_idx)","metadata":{"id":"i14OtQz0wp_e","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:26.391324Z","iopub.execute_input":"2025-03-23T17:43:26.391625Z","iopub.status.idle":"2025-03-23T17:43:26.395451Z","shell.execute_reply.started":"2025-03-23T17:43:26.391602Z","shell.execute_reply":"2025-03-23T17:43:26.394464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainloader = torch.utils.data.DataLoader(train_data, batch_size=64,sampler=train_sampler)\nvalidloader = torch.utils.data.DataLoader(train_data, batch_size=64, sampler=valid_sampler)\ntestloader = torch.utils.data.DataLoader(test_data, batch_size=64)","metadata":{"id":"oSPMv1iewp_h","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:27.793888Z","iopub.execute_input":"2025-03-23T17:43:27.794224Z","iopub.status.idle":"2025-03-23T17:43:27.798600Z","shell.execute_reply.started":"2025-03-23T17:43:27.794195Z","shell.execute_reply":"2025-03-23T17:43:27.797647Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"training examples contain : {len(train_data)}\")\nprint(f\"testing examples contain : {len(test_data)}\")\n\nprint(len(trainloader))\nprint(len(validloader))\nprint(len(testloader))","metadata":{"id":"NvMplySzwp_j","outputId":"4637822a-caa7-4617-91c4-7820badd6978","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:27.896127Z","iopub.execute_input":"2025-03-23T17:43:27.896452Z","iopub.status.idle":"2025-03-23T17:43:27.901755Z","shell.execute_reply.started":"2025-03-23T17:43:27.896425Z","shell.execute_reply":"2025-03-23T17:43:27.901054Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# LOAD ONE BATCH OF TESTING SET TO CHECK THE IMAGES AND THEIR LABELS\nimages, labels = next(iter(trainloader))\n\n# Checking shape of image\nprint(f\"Image shape : {images.shape}\")\nprint(f\"Label shape : {labels.shape}\")\n\n# denormalizing images\ndef imshow(inp, title=None):\n    \"\"\"Imshow for Tensor.\"\"\"\n    inp = inp.numpy().transpose((1, 2, 0))\n    mean = np.array([0.485, 0.456, 0.406])\n    std = np.array([0.229, 0.224, 0.225])\n    inp = std * inp + mean\n    inp = np.clip(inp, 0, 1)\n    plt.imshow(inp)\n    if title is not None:\n        plt.title(title)\n    plt.pause(0.001)","metadata":{"id":"I63353zewp_l","outputId":"229d7b10-ecaf-49bd-f25b-81816dab9a5d","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:28.187760Z","iopub.execute_input":"2025-03-23T17:43:28.188073Z","iopub.status.idle":"2025-03-23T17:43:35.995648Z","shell.execute_reply.started":"2025-03-23T17:43:28.188050Z","shell.execute_reply":"2025-03-23T17:43:35.994854Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# plotting the images of loaded batch with given fig size and frame data    \nimport torchvision\nimport matplotlib.pyplot as plt\nimport numpy as np\ngrid = torchvision.utils.make_grid(images, nrow = 20, padding = 2)\nplt.figure(figsize = (20, 20))  \nplt.imshow(np.transpose(grid, (1, 2, 0)))   \nprint('labels:', labels)    ","metadata":{"id":"UJC-bbEGwp_o","outputId":"420d0427-7736-460b-82da-37015f8ee195","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:35.996850Z","iopub.execute_input":"2025-03-23T17:43:35.997231Z","iopub.status.idle":"2025-03-23T17:43:37.358287Z","shell.execute_reply.started":"2025-03-23T17:43:35.997198Z","shell.execute_reply":"2025-03-23T17:43:37.357522Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_names = ['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative DR']\n\nimages, labels = next(iter(trainloader))\nout = torchvision.utils.make_grid(images)\nimshow(out, title=[class_names[x] for x in labels])","metadata":{"id":"LJ9C48WYwp_s","outputId":"4ebc2eea-87bd-418f-fb96-cccd50d0e62d","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:37.360057Z","iopub.execute_input":"2025-03-23T17:43:37.360360Z","iopub.status.idle":"2025-03-23T17:43:45.780270Z","shell.execute_reply.started":"2025-03-23T17:43:37.360332Z","shell.execute_reply":"2025-03-23T17:43:45.779393Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_on_gpu = torch.cuda.is_available()\n\nif not train_on_gpu:\n    print('CUDA is not available.  Training on CPU ...')\nelse:\n    print('CUDA is available!  Training on GPU ...')","metadata":{"id":"jqPov202wp_w","outputId":"67396e4a-21e5-4007-d67b-1d913cddcdab","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:45.781657Z","iopub.execute_input":"2025-03-23T17:43:45.782003Z","iopub.status.idle":"2025-03-23T17:43:45.866955Z","shell.execute_reply.started":"2025-03-23T17:43:45.781976Z","shell.execute_reply":"2025-03-23T17:43:45.866234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = models.resnet152(pretrained=True) \n\nnum_ftrs = model.fc.in_features \nout_ftrs = 5 \n  \nmodel.fc = nn.Sequential(nn.Linear(num_ftrs, 512),nn.ReLU(),nn.Linear(512,out_ftrs),nn.LogSoftmax(dim=1))\n\ncriterion = nn.NLLLoss()\noptimizer = torch.optim.Adam(filter(lambda p:p.requires_grad,model.parameters()) , lr = 0.00001) \n\nscheduler = lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\nmodel.to(device);","metadata":{"id":"YWTE2JVxwp_1","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:45.867695Z","iopub.execute_input":"2025-03-23T17:43:45.867974Z","iopub.status.idle":"2025-03-23T17:43:48.632410Z","shell.execute_reply.started":"2025-03-23T17:43:45.867944Z","shell.execute_reply":"2025-03-23T17:43:48.631410Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_save_name = 'classifier.pt'\npath = f\"/kaggle/working/{model_save_name}\"\npath","metadata":{"id":"v4Hn2lk4wp_4","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T18:02:07.150669Z","iopub.execute_input":"2025-03-23T18:02:07.151049Z","iopub.status.idle":"2025-03-23T18:02:07.156110Z","shell.execute_reply.started":"2025-03-23T18:02:07.151020Z","shell.execute_reply":"2025-03-23T18:02:07.155382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# to unfreeze more layers \nfor name,child in model.named_children():\n  if name in ['layer2','layer3','layer4','fc']:\n    print(name + 'is unfrozen')\n    for param in child.parameters():\n      param.requires_grad = True\n  else:\n    print(name + 'is frozen')\n    for param in child.parameters():\n      param.requires_grad = False","metadata":{"id":"9ETX_27awp_7","outputId":"b7d030e9-e9bb-4dde-d592-0c8c50167d6f","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:56.706186Z","iopub.execute_input":"2025-03-23T17:43:56.706466Z","iopub.status.idle":"2025-03-23T17:43:56.714988Z","shell.execute_reply.started":"2025-03-23T17:43:56.706445Z","shell.execute_reply":"2025-03-23T17:43:56.714107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.Adam(filter(lambda p:p.requires_grad,model.parameters()) , lr = 0.000001) \nscheduler = lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)","metadata":{"id":"iOF-Abo-wp_-","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:43:59.181477Z","iopub.execute_input":"2025-03-23T17:43:59.181825Z","iopub.status.idle":"2025-03-23T17:43:59.187615Z","shell.execute_reply.started":"2025-03-23T17:43:59.181798Z","shell.execute_reply":"2025-03-23T17:43:59.186754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_model(path):\n  checkpoint = torch.load(path)\n  model.load_state_dict(checkpoint['model_state_dict'])\n  optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n  return model","metadata":{"id":"G0sprIg3wqAB","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:44:23.066151Z","iopub.execute_input":"2025-03-23T17:44:23.066450Z","iopub.status.idle":"2025-03-23T17:44:23.070418Z","shell.execute_reply.started":"2025-03-23T17:44:23.066427Z","shell.execute_reply":"2025-03-23T17:44:23.069432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import kagglehub\n\n# Download latest version\npath = kagglehub.dataset_download(\"souravs17031999/blindness-detection-pretrained-weights-pytorch\")\n\nprint(\"Path to dataset files:\", path)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:46:56.345948Z","iopub.execute_input":"2025-03-23T17:46:56.346269Z","iopub.status.idle":"2025-03-23T17:47:11.395250Z","shell.execute_reply.started":"2025-03-23T17:46:56.346246Z","shell.execute_reply":"2025-03-23T17:47:11.394534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = load_model(\"/kaggle/input/blindness-detection-pretrained-weights-pytorch/classifier.pt\")","metadata":{"id":"7gndNUYIwqAD","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:47:43.271697Z","iopub.execute_input":"2025-03-23T17:47:43.272074Z","iopub.status.idle":"2025-03-23T17:47:48.540589Z","shell.execute_reply.started":"2025-03-23T17:47:43.272047Z","shell.execute_reply":"2025-03-23T17:47:48.539949Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model","metadata":{"id":"JBtQWIRFwqAF","outputId":"c75dfe78-b90b-435e-8edb-4aaf6d011f07","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:47:51.716552Z","iopub.execute_input":"2025-03-23T17:47:51.716868Z","iopub.status.idle":"2025-03-23T17:47:51.726024Z","shell.execute_reply.started":"2025-03-23T17:47:51.716845Z","shell.execute_reply":"2025-03-23T17:47:51.725178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"pytorch_total_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(\"Number of trainable parameters: \\n{}\".format(pytorch_total_params))","metadata":{"id":"3LIwOcNzwqAH","outputId":"fc685b55-eaa3-4e71-b55b-51a0c43f19eb","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:47:53.971063Z","iopub.execute_input":"2025-03-23T17:47:53.971342Z","iopub.status.idle":"2025-03-23T17:47:53.977378Z","shell.execute_reply.started":"2025-03-23T17:47:53.971323Z","shell.execute_reply":"2025-03-23T17:47:53.976416Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_and_test(e):\n    epochs = e\n    train_losses , test_losses, acc = [] , [], []\n    valid_loss_min = np.Inf \n    model.train()\n    print(\"Model Training started.....\")\n    for epoch in range(epochs):\n      running_loss = 0\n      batch = 0\n      for images , labels in trainloader:\n        images, labels = images.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs,labels)\n        loss.backward()\n        optimizer.step()\n        running_loss += loss.item()\n        batch += 1\n        if batch % 10 == 0:\n            print(f\" epoch {epoch + 1} batch {batch} completed\") \n      test_loss = 0\n      accuracy = 0\n      with torch.no_grad():\n        print(f\"validation started for {epoch + 1}\")\n        model.eval() \n        for images , labels in validloader:\n          images, labels = images.to(device), labels.to(device)\n          logps = model(images) \n          test_loss += criterion(logps,labels) \n          ps = torch.exp(logps)\n          top_p , top_class = ps.topk(1,dim=1)\n          equals = top_class == labels.view(*top_class.shape)\n          accuracy += torch.mean(equals.type(torch.FloatTensor))\n      train_losses.append(running_loss/len(trainloader))\n      test_losses.append(test_loss/len(validloader))\n      acc.append(accuracy)\n      scheduler.step()\n      print(\"Epoch: {}/{}.. \".format(epoch+1, epochs),\"Training Loss: {:.3f}.. \".format(running_loss/len(trainloader)),\"Valid Loss: {:.3f}.. \".format(test_loss/len(validloader)),\n        \"Valid Accuracy: {:.3f}\".format(accuracy/len(validloader)))\n      model.train() \n      if test_loss/len(validloader) <= valid_loss_min:\n        print('Validation loss decreased ({:.6f} --> {:.6f}).  Saving model ...'.format(valid_loss_min,test_loss/len(validloader))) \n        torch.save({\n            'epoch': epoch,\n            'model': model,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'loss': valid_loss_min\n            }, path)\n        valid_loss_min = test_loss/len(validloader)    \n    print('Training Completed Succesfully !')    \n    return train_losses, test_losses, acc ","metadata":{"id":"rHB2XhSzwqAK","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T18:02:38.871118Z","iopub.execute_input":"2025-03-23T18:02:38.871411Z","iopub.status.idle":"2025-03-23T18:02:38.879779Z","shell.execute_reply.started":"2025-03-23T18:02:38.871389Z","shell.execute_reply":"2025-03-23T18:02:38.878700Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_losses, valid_losses, acc = train_and_test(5)","metadata":{"id":"9X7zW47awqAM","outputId":"d514b4d2-f42e-4991-f326-ec1c68056c6d","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T18:02:39.195084Z","iopub.execute_input":"2025-03-23T18:02:39.195350Z","iopub.status.idle":"2025-03-23T18:42:24.172865Z","shell.execute_reply.started":"2025-03-23T18:02:39.195328Z","shell.execute_reply":"2025-03-23T18:42:24.172060Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# %matplotlib inline\n# %config InlineBackend.figure_format = 'retina'\n\n# plt.plot(train_losses, label='train_')\n# plt.plot(valid_losses, label='Validation loss')\n# plt.xlabel(\"Epochs\")\n# plt.ylabel(\"Loss\")\n# plt.legend(frameon=False)\n\n\ntrain_losses = [loss.cpu().item() if isinstance(loss, torch.Tensor) else loss for loss in train_losses]\nvalid_losses = [loss.cpu().item() if isinstance(loss, torch.Tensor) else loss for loss in valid_losses]\n\nplt.plot(train_losses, label='Train Loss')\nplt.plot(valid_losses, label='Validation Loss')\n\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"Loss\")\nplt.legend(frameon=False)\nplt.show()\n\n","metadata":{"id":"zq572wFDwqAO","outputId":"b12538f5-6eef-4284-b7af-7bca45cabb9f","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T18:46:19.891190Z","iopub.execute_input":"2025-03-23T18:46:19.891530Z","iopub.status.idle":"2025-03-23T18:46:20.107592Z","shell.execute_reply.started":"2025-03-23T18:46:19.891501Z","shell.execute_reply":"2025-03-23T18:46:20.106663Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%matplotlib inline\n%config InlineBackend.figure_format = 'retina'\n\nplt.plot(acc, label='accuracy')\nplt.legend(\"\")\nplt.xlabel(\"Epochs\")\nplt.ylabel(\"accuracy\")\nplt.legend(frameon=False)","metadata":{"id":"14JeIrsSwqAR","outputId":"98493497-4844-4615-a872-d9d2e367d412","trusted":true,"execution":{"iopub.status.busy":"2025-03-23T18:46:31.555274Z","iopub.execute_input":"2025-03-23T18:46:31.555562Z","iopub.status.idle":"2025-03-23T18:46:31.802207Z","shell.execute_reply.started":"2025-03-23T18:46:31.555540Z","shell.execute_reply":"2025-03-23T18:46:31.801353Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}