{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":14774,"databundleVersionId":875431,"sourceType":"competition"},{"sourceId":418031,"sourceType":"datasetVersion","datasetId":131128},{"sourceId":111254110,"sourceType":"kernelVersion"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<a name='1-1'></a>\n## 1.1 Import Packages","metadata":{}},{"cell_type":"code","source":"#importing libraries \nimport numpy as np\nimport pandas as pd\n\nfrom PIL import Image, ImageFile\nImageFile.LOAD_TRUNCATED_IMAGES = True\nimport cv2\n\nfrom tqdm import tqdm_notebook as tqdm\nfrom functools import partial\nimport scipy as sp\n\nimport random\nimport time\nimport sys\nimport os\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n%matplotlib inline\n\nfrom sklearn import metrics\nfrom sklearn.metrics import confusion_matrix\nimport torch\nimport torchvision\n\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torchvision import transforms, models, datasets\nfrom torch.utils.data import Dataset\nfrom torch.autograd import Variable\n\n!pip install efficientnet_pytorch\nfrom efficientnet_pytorch import EfficientNet\n\nimport warnings\nwarnings.filterwarnings('ignore')\n!mkdir models","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:23:14.515003Z","iopub.execute_input":"2025-06-09T14:23:14.515686Z","iopub.status.idle":"2025-06-09T14:23:21.320795Z","shell.execute_reply.started":"2025-06-09T14:23:14.515653Z","shell.execute_reply":"2025-06-09T14:23:21.320026Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name='1-2'></a>\n\n## 1.2 Create Classes and Function","metadata":{}},{"cell_type":"code","source":"# seed function\ndef seed_everything(seed = 23):\n    # tests\n    assert isinstance(seed, int), 'seed has to be an integer'\n    \n    # randomness\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:23:27.749677Z","iopub.execute_input":"2025-06-09T14:23:27.751086Z","iopub.status.idle":"2025-06-09T14:23:27.757055Z","shell.execute_reply.started":"2025-06-09T14:23:27.751046Z","shell.execute_reply":"2025-06-09T14:23:27.756266Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_size = 256\n#IMAGE PREPROCESSING\n\ndef prepare_image(path, \n                  sigmaX         = 10, \n                  do_random_crop = False):\n    \n    '''\n    Preprocess image\n    '''\n    \n    # import image\n    image = cv2.imread(path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    \n    # perform smart crops\n    image = crop_black(image, tol = 7)\n    if do_random_crop == True:\n        image = random_crop(image, size = (0.9, 1))\n    \n    # resize and color\n    image = cv2.resize(image, (int(image_size), int(image_size)))\n    image = cv2.addWeighted(image, 4, cv2.GaussianBlur(image, (0, 0), sigmaX), -4, 128)\n    \n    # circular crop\n    image = circle_crop(image, sigmaX = sigmaX)\n\n    # convert to tensor    \n    image = torch.tensor(image)\n    image = image.permute(2, 1, 0)\n    return image","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:23:31.375315Z","iopub.execute_input":"2025-06-09T14:23:31.3756Z","iopub.status.idle":"2025-06-09T14:23:31.381636Z","shell.execute_reply.started":"2025-06-09T14:23:31.375579Z","shell.execute_reply":"2025-06-09T14:23:31.380783Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#CROP FUNCTIONS\n\ndef crop_black(img, \n               tol = 7):\n    \n    '''\n    Perform automatic crop of black areas\n    '''\n    \n    if img.ndim == 2:\n        mask = img > tol\n        return img[np.ix_(mask.any(1),mask.any(0))]\n    \n    elif img.ndim == 3:\n        gray_img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n        mask = gray_img > tol\n        check_shape = img[:,:,0][np.ix_(mask.any(1),mask.any(0))].shape[0]\n        \n        if (check_shape == 0): \n            return img \n        else:\n            img1 = img[:,:,0][np.ix_(mask.any(1),mask.any(0))]\n            img2 = img[:,:,1][np.ix_(mask.any(1),mask.any(0))]\n            img3 = img[:,:,2][np.ix_(mask.any(1),mask.any(0))]\n            img  = np.stack([img1, img2, img3], axis = -1)\n            return img\n","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:23:35.229473Z","iopub.execute_input":"2025-06-09T14:23:35.230287Z","iopub.status.idle":"2025-06-09T14:23:35.236103Z","shell.execute_reply.started":"2025-06-09T14:23:35.23026Z","shell.execute_reply":"2025-06-09T14:23:35.235307Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def circle_crop(img, \n                sigmaX = 10):   \n    \n    '''\n    Perform circular crop around image center\n    '''\n        \n    height, width, depth = img.shape\n    \n    largest_side = np.max((height, width))\n    img = cv2.resize(img, (largest_side, largest_side))\n\n    height, width, depth = img.shape\n    \n    x = int(width / 2)\n    y = int(height / 2)\n    r = np.amin((x,y))\n    \n    circle_img = np.zeros((height, width), np.uint8)\n    cv2.circle(circle_img, (x,y), int(r), 1, thickness = -1)\n    \n    img = cv2.bitwise_and(img, img, mask = circle_img)\n    return img \n","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:23:38.449054Z","iopub.execute_input":"2025-06-09T14:23:38.449343Z","iopub.status.idle":"2025-06-09T14:23:38.454638Z","shell.execute_reply.started":"2025-06-09T14:23:38.449322Z","shell.execute_reply":"2025-06-09T14:23:38.45387Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def random_crop(img, \n                size = (0.9, 1)):\n    \n    '''\n    Random crop\n    '''\n\n    height, width, depth = img.shape\n    \n    cut = 1 - random.uniform(size[0], size[1])\n    \n    i = random.randint(0, int(cut * height))\n    j = random.randint(0, int(cut * width))\n    h = i + int((1 - cut) * height)\n    w = j + int((1 - cut) * width)\n\n    img = img[i:h, j:w, :]    \n    \n    return img","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:23:42.039544Z","iopub.execute_input":"2025-06-09T14:23:42.039835Z","iopub.status.idle":"2025-06-09T14:23:42.044723Z","shell.execute_reply.started":"2025-06-09T14:23:42.039812Z","shell.execute_reply":"2025-06-09T14:23:42.043948Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class EyeData(Dataset):\n    \n    # initialize\n    def __init__(self, data, directory, transform = None, do_random_crop = True, itype = '.png'):\n        self.data      = data\n        self.directory = directory\n        self.transform = transform\n        self.do_random_crop = do_random_crop\n        self.itype = itype\n    # length\n    def __len__(self):\n        return len(self.data)\n    \n    # get items    \n    def __getitem__(self, idx):\n        img_name = os.path.join(self.directory, self.data.loc[idx, 'id_code'] + self.itype)\n        image    = prepare_image(img_name, do_random_crop = self.do_random_crop)\n        image    = self.transform(image)\n        label    = torch.tensor(self.data.loc[idx, 'diagnosis'])\n        return {'image': image, 'label': label}","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:23:46.464352Z","iopub.execute_input":"2025-06-09T14:23:46.465091Z","iopub.status.idle":"2025-06-09T14:23:46.470291Z","shell.execute_reply.started":"2025-06-09T14:23:46.465065Z","shell.execute_reply":"2025-06-09T14:23:46.469582Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Data(Dataset):\n    \n    # initialize\n    def __init__(self, data, directory, transform = None, do_random_crop = True, itype = '.png'):\n        self.data      = data\n        self.directory = directory\n        self.transform = transform\n        self.do_random_crop = do_random_crop\n        self.itype = itype\n    # length\n    def __len__(self):\n        return len(self.data)\n    \n    # get items    \n    def __getitem__(self, idx):\n        img_name = os.path.join(self.directory, self.data.loc[idx, 'id_code'] + self.itype)\n        image = cv2.imread(img_name)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = crop_black(image, tol = 7)\n        image = cv2.resize(image, (int(image_size), int(image_size)))\n        image = circle_crop(image, sigmaX = 10)\n        image = torch.tensor(image)\n        image = image.permute(2, 1, 0)\n        image    = self.transform(image)\n        label    = torch.tensor(self.data.loc[idx, 'diagnosis'])\n        return {'image': image, 'label': label}","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:23:50.544417Z","iopub.execute_input":"2025-06-09T14:23:50.545149Z","iopub.status.idle":"2025-06-09T14:23:50.551351Z","shell.execute_reply.started":"2025-06-09T14:23:50.545123Z","shell.execute_reply":"2025-06-09T14:23:50.550513Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def init_model(train= True, \n               trn_layers = 2,\n               model_name = 'enet_b7'):\n    \n    '''\n    Initialize the model\n    '''\n    \n    ### training mode\n    if train == True:\n        \n        # load pre-trained model\n        model = EfficientNet.from_pretrained('efficientnet-b7', num_classes = 5)\n        model.load_state_dict(torch.load('../input/diabetic-retinopathy-pre-training/models/model_{}.bin'.format(model_name, 1)))   \n        \n        # freeze first layers\n        for child in list(model.children())[:-trn_layers]:\n            for param in child.parameters():\n                param.requires_grad = False\n        \n        \n    #inference mode\n    if train == False:\n        \n        # load pre-trained model\n        model = EfficientNet.from_pretrained('efficientnet-b7', num_classes = 5)\n        model.load_state_dict(torch.load('../input/diabetic-retinopathy-pre-training/models/model_{}.bin'.format(model_name, 1)))   \n\n        # freeze all layers\n        for param in model.parameters():\n            param.requires_grad = False\n            \n            \n    ### return model\n    return model","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:23:54.95901Z","iopub.execute_input":"2025-06-09T14:23:54.959278Z","iopub.status.idle":"2025-06-09T14:23:54.964459Z","shell.execute_reply.started":"2025-06-09T14:23:54.959258Z","shell.execute_reply":"2025-06-09T14:23:54.963715Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#GPU CHECK\ntrain_on_gpu = torch.cuda.is_available()\nif not train_on_gpu:\n    print('CUDA is not available. Training on CPU...')\n    device = torch.device('cpu')\nelse:\n    print('CUDA is available. Training on GPU...')\n    device = torch.device('cuda:0')","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:24:03.739806Z","iopub.execute_input":"2025-06-09T14:24:03.740385Z","iopub.status.idle":"2025-06-09T14:24:03.744596Z","shell.execute_reply.started":"2025-06-09T14:24:03.740361Z","shell.execute_reply":"2025-06-09T14:24:03.743866Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#RANDOMNESS\n\nseed = 23\nseed_everything(seed)","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:24:14.170056Z","iopub.execute_input":"2025-06-09T14:24:14.170325Z","iopub.status.idle":"2025-06-09T14:24:14.175375Z","shell.execute_reply.started":"2025-06-09T14:24:14.170305Z","shell.execute_reply":"2025-06-09T14:24:14.174844Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name='2'></a>\n# 2 - Pre-Training on Diabetic Retinopathy 2015 data\n<center><img src=\"https://www.mdpi.com/applsci/applsci-10-07274/article_deploy/html/images/applsci-10-07274-g001.png\" width=\"1800px\" height=\"200px\"></center>","metadata":{}},{"cell_type":"markdown","source":"<a name='2-1'></a>\n## 2.1 Import Data","metadata":{}},{"cell_type":"code","source":"# import data\ntrain = pd.read_csv('../input/diabetic-retinopathy-resized/trainLabels.csv')\ntrain.columns = ['id_code', 'diagnosis']\ntest = pd.read_csv('../input/aptos2019-blindness-detection/train.csv')\n\n# check shape\nprint(train.shape, test.shape)\nprint('-' * 15)\nprint(train['diagnosis'].value_counts())\nprint('-' * 15)\nprint(test['diagnosis'].value_counts())","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:24:19.072181Z","iopub.execute_input":"2025-06-09T14:24:19.072473Z","iopub.status.idle":"2025-06-09T14:24:19.110657Z","shell.execute_reply.started":"2025-06-09T14:24:19.072452Z","shell.execute_reply":"2025-06-09T14:24:19.110124Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name='2-2'></a>\n## 2.2 Explore Data","metadata":{}},{"cell_type":"code","source":"# CLASS DISTRIBUTION\n\n# plot\nfig = plt.figure(figsize = (15, 5))\nplt.hist(train['diagnosis'])\nplt.title('Class Distribution')\nplt.ylabel('Number of examples')\nplt.xlabel('Diagnosis')","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:24:23.88981Z","iopub.execute_input":"2025-06-09T14:24:23.890108Z","iopub.status.idle":"2025-06-09T14:24:24.116947Z","shell.execute_reply.started":"2025-06-09T14:24:23.890088Z","shell.execute_reply":"2025-06-09T14:24:24.116349Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# transformations\nsample_trans = transforms.Compose([transforms.ToPILImage(),\n                                   transforms.ToTensor(),\n                                  ])\nsample = Data(data       = train.iloc[0:10], \n                      directory  = '../input/diabetic-retinopathy-resized/resized_train/resized_train',\n                      transform  = sample_trans,\n                      itype ='.jpeg')\n\n# data loader\nsample_loader = torch.utils.data.DataLoader(dataset     = sample, \n                                            batch_size  = 10, \n                                            shuffle     = False, \n                                            num_workers = 4)\n\n# display images\nfor batch_i, data in enumerate(sample_loader):\n\n    # extract data\n    inputs = data['image']\n    labels = data['label'].view(-1, 1)\n    \n    # create plot\n    fig = plt.figure(figsize = (15, 7))\n    for i in range(len(labels)):\n        ax = fig.add_subplot(2, int(len(labels)/2), i + 1, xticks = [], yticks = [])     \n        plt.imshow(inputs[i].numpy().transpose(1, 2, 0))\n        ax.set_title(labels.numpy()[i])\n\n    break","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:24:28.200393Z","iopub.execute_input":"2025-06-09T14:24:28.201015Z","iopub.status.idle":"2025-06-09T14:24:29.189308Z","shell.execute_reply.started":"2025-06-09T14:24:28.200991Z","shell.execute_reply":"2025-06-09T14:24:29.188539Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# IMAGE SIZES\n\n# placeholder\nimage_stats = []\n\n# import loop\nfor index, observation in tqdm(train.iterrows(), total = len(train)):\n    \n    # import image\n    img = cv2.imread('../input/diabetic-retinopathy-resized/resized_train/resized_train/{}.jpeg'.format(observation['id_code']))\n\n    # compute stats\n    height, width, channels = img.shape\n    ratio = width / height\n    \n    # save\n    image_stats.append(np.array((observation['diagnosis'], height, width, channels, ratio)))\n\n# construct DF\nimage_stats = pd.DataFrame(image_stats)\nimage_stats.columns = ['diagnosis', 'height', 'width', 'channels', 'ratio']","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:24:37.706329Z","iopub.execute_input":"2025-06-09T14:24:37.707058Z","iopub.status.idle":"2025-06-09T14:27:20.019396Z","shell.execute_reply.started":"2025-06-09T14:24:37.707022Z","shell.execute_reply":"2025-06-09T14:27:20.018819Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# IMAGE SIZE DISTRIBUTION\n\nfig = plt.figure(figsize = (15, 5))\n\n# width\nplt.subplot(1, 3, 1)\nplt.hist(image_stats['width'])\nplt.title('(a) Image Width')\nplt.ylabel('Number of examples')\nplt.xlabel('Width')\n\n# height\nplt.subplot(1, 3, 2)\nplt.hist(image_stats['height'])\nplt.title('(b) Image Height')\nplt.ylabel('Number of examples')\nplt.xlabel('Height')\n\n# ratio\nplt.subplot(1, 3, 3)\nplt.hist(image_stats['ratio'])\nplt.title('(c) Aspect Ratio')\nplt.ylabel('Number of examples')\nplt.xlabel('Ratio')","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:27:32.333586Z","iopub.execute_input":"2025-06-09T14:27:32.334214Z","iopub.status.idle":"2025-06-09T14:27:32.781449Z","shell.execute_reply.started":"2025-06-09T14:27:32.334187Z","shell.execute_reply":"2025-06-09T14:27:32.780594Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name='2-3'></a>\n## 2.3  Preprocess and Augment Training Data","metadata":{}},{"cell_type":"code","source":"#TRANSFORMATIONS\n\n# parameters\nbatch_size = 16\nimage_size = 256\n\n# train transformations\ntrain_trans = transforms.Compose([transforms.ToPILImage(),\n                                  transforms.RandomRotation((-360, 360)),\n                                  transforms.RandomHorizontalFlip(),\n                                  transforms.RandomVerticalFlip(),\n                                  transforms.ToTensor()\n                                 ])\n\n# validation transformations\nvalid_trans = transforms.Compose([transforms.ToPILImage(),\n                                  transforms.ToTensor(),\n                                 ])\n\n# test transformations\ntest_trans = valid_trans","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:27:41.68355Z","iopub.execute_input":"2025-06-09T14:27:41.684165Z","iopub.status.idle":"2025-06-09T14:27:41.689272Z","shell.execute_reply.started":"2025-06-09T14:27:41.684139Z","shell.execute_reply":"2025-06-09T14:27:41.688441Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#EXAMINE FIRST BATCH (TRAIN)\n\n# get dataset\nsample = EyeData(data       = train.iloc[0:10], \n                      directory  = '../input/diabetic-retinopathy-resized/resized_train/resized_train',\n                      transform  = train_trans,\n                      itype ='.jpeg')\n\n# data loader\nsample_loader = torch.utils.data.DataLoader(dataset     = sample, \n                                            batch_size  = batch_size, \n                                            shuffle     = True, \n                                            num_workers = 4)\n\n# display images\nfor batch_i, data in enumerate(sample_loader):\n\n    # extract data\n    inputs = data['image']\n    labels = data['label'].view(-1, 1)\n    \n    # create plot\n    fig = plt.figure(figsize = (20,10))\n    for i in range(len(labels)):\n        ax = fig.add_subplot(2, int(len(labels)/2), i + 1, xticks = [], yticks = [])     \n        plt.imshow(inputs[i].numpy().transpose(1, 2, 0))\n        ax.set_title(labels.numpy()[i])\n\n    break","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:27:47.763944Z","iopub.execute_input":"2025-06-09T14:27:47.764231Z","iopub.status.idle":"2025-06-09T14:27:49.028543Z","shell.execute_reply.started":"2025-06-09T14:27:47.76421Z","shell.execute_reply":"2025-06-09T14:27:49.02764Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#EXAMINE FIRST BATCH (TEST)\n\n# get dataset\nsample = EyeData(data       = test.iloc[0:10], \n                      directory  = '../input/aptos2019-blindness-detection/train_images',\n                      transform  = test_trans,\n                      itype ='.png',\n                      do_random_crop = False)\n\n# data loader\nsample_loader = torch.utils.data.DataLoader(dataset     = sample, \n                                            batch_size  = batch_size, \n                                            shuffle     = False, \n                                            num_workers = 4)\n\n# display images\nfor batch_i, data in enumerate(sample_loader):\n\n    # extract data\n    inputs = data['image']\n    \n    # create plot\n    fig = plt.figure(figsize = (20,10))\n    for i in range(10):\n        ax = fig.add_subplot(2, int(10/2), i + 1, xticks = [], yticks = [])     \n        plt.imshow(inputs[i].numpy().transpose(1, 2, 0))\n\n    break","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:28:00.039401Z","iopub.execute_input":"2025-06-09T14:28:00.040312Z","iopub.status.idle":"2025-06-09T14:28:03.498988Z","shell.execute_reply.started":"2025-06-09T14:28:00.040279Z","shell.execute_reply":"2025-06-09T14:28:03.49772Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name='2-4'></a>\n## 2.4 - Using EFfienctNet for Transfer Learning\n<center><img src=\"https://1.bp.blogspot.com/-DjZT_TLYZok/XO3BYqpxCJI/AAAAAAAAEKM/BvV53klXaTUuQHCkOXZZGywRMdU9v9T_wCLcBGAs/s1600/image2.png\" width=\"1400\" height=\"500x\"></center>","metadata":{}},{"cell_type":"markdown","source":"<a name='2-4-1'></a>\n### 2.4.1 - Setup Model and Choose Hyperparameters","metadata":{}},{"cell_type":"code","source":"#MODEL ARCHITECTURE\n\n# model name\nmodel_name = 'enet_b7'\n\n# initialization function\ndef init_pre_model(train = True):\n    \n    '''\n    Initialize the model\n    '''\n    \n    ### training mode\n    if train == True:\n        \n        # load pre-trained model\n        model = EfficientNet.from_pretrained('efficientnet-b7', num_classes = 5)\n        \n    ### inference mode\n    if train == False:\n        \n        # load pre-trained model\n        model = EfficientNet.from_name('efficientnet-b7')\n        model._fc = nn.Linear(model._fc.in_features, 5)\n\n        # freeze  layers\n        for param in model.parameters():\n            param.requires_grad = False\n            \n    ### return model\n    return model\n\n\n# check architecture\nmodel = init_pre_model()\nprint(model)","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2025-06-09T14:28:34.391442Z","iopub.execute_input":"2025-06-09T14:28:34.392188Z","iopub.status.idle":"2025-06-09T14:28:35.531349Z","shell.execute_reply.started":"2025-06-09T14:28:34.392158Z","shell.execute_reply":"2025-06-09T14:28:35.530675Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#VALIDATION SETTINGS\n\n# placeholders\noof_preds = np.zeros((len(test), 5))\n\n# timer\ncv_start = time.time()","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:28:44.702894Z","iopub.execute_input":"2025-06-09T14:28:44.703658Z","iopub.status.idle":"2025-06-09T14:28:44.707719Z","shell.execute_reply.started":"2025-06-09T14:28:44.703624Z","shell.execute_reply":"2025-06-09T14:28:44.706993Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#PARAMETERS\n\n# loss function\ncriterion = nn.CrossEntropyLoss()\n\n# epochs\nmax_epochs = 15\nearly_stop = 5\n\n# learning rates\neta = 1e-3\n\n# scheduler\nstep  = 5\ngamma = 0.5","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:28:49.746229Z","iopub.execute_input":"2025-06-09T14:28:49.746768Z","iopub.status.idle":"2025-06-09T14:28:49.750588Z","shell.execute_reply.started":"2025-06-09T14:28:49.746746Z","shell.execute_reply":"2025-06-09T14:28:49.74969Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name='2-4-2'></a>\n## 2.4.2 - Train Model on Diabetic Retinopathy 2015 data","metadata":{}},{"cell_type":"code","source":"#DATA PREPARATION\n\n# load splits\ndata_train = train\ndata_valid = test\n\n# create datasets\ntrain_dataset = EyeData(data      = data_train, \n                             directory = '../input/diabetic-retinopathy-resized/resized_train/resized_train',\n                             transform = train_trans,\n                             itype ='.jpeg')\nvalid_dataset = EyeData(data       = data_valid, \n                            directory  = '../input/aptos2019-blindness-detection/train_images',\n                            transform  = valid_trans,\n                            itype ='.png')\n\n# create data loaders\ntrain_loader = torch.utils.data.DataLoader(train_dataset, \n                                           batch_size  = batch_size, \n                                           shuffle     = True, \n                                           num_workers = 4)\nvalid_loader = torch.utils.data.DataLoader(valid_dataset, \n                                           batch_size  = batch_size, \n                                           shuffle     = False, \n                                           num_workers = 4)","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:29:02.051469Z","iopub.execute_input":"2025-06-09T14:29:02.052174Z","iopub.status.idle":"2025-06-09T14:29:02.056616Z","shell.execute_reply.started":"2025-06-09T14:29:02.052148Z","shell.execute_reply":"2025-06-09T14:29:02.05594Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#MODELING EPOCHS\n\n# placeholders\nval_kappas = []\nval_losses = []\ntrn_losses = []\nbad_epochs = 0\n\n# initialize and send to GPU\nmodel = init_pre_model()\nmodel = model.to(device)\n\n# optimizer\noptimizer = optim.Adam(model.parameters(), lr = eta)\nscheduler = lr_scheduler.StepLR(optimizer, step_size = step, gamma = gamma)\n\n# training and validation loop\nfor epoch in range(max_epochs):\n    ### PREPARATION\n\n    # timer\n    epoch_start = time.time()\n\n    # reset losses\n    trn_loss = 0.0\n    val_loss = 0.0\n\n    # placeholders\n    fold_preds = np.zeros((len(data_valid), 5))\n\n\n    #TRAINING\n\n    # switch regime\n    model.train()\n\n    # loop through batches\n    for batch_i, data in enumerate(train_loader):\n\n        # extract inputs and labels\n        inputs = data['image']\n        labels = data['label'].view(-1)\n        inputs = inputs.to(device, dtype = torch.float)\n        labels = labels.to(device, dtype = torch.long)\n        optimizer.zero_grad()\n\n        # forward and backward pass\n        with torch.set_grad_enabled(True):\n            preds = model(inputs)\n            loss  = criterion(preds, labels)\n            loss.backward()\n            optimizer.step()\n\n        # compute loss\n        trn_loss += loss.item() * inputs.size(0)\n        \n        \n    #INFERENCE\n\n    # switch regime\n    model.eval()\n    \n    # loop through batches\n    for batch_i, data in enumerate(valid_loader):\n        \n        # extract inputs and labels\n        inputs = data['image']\n        labels = data['label'].view(-1)\n        inputs = inputs.to(device, dtype = torch.float)\n        labels = labels.to(device, dtype = torch.long)\n\n        # compute predictions\n        with torch.set_grad_enabled(False):\n            preds = model(inputs).detach()\n            fold_preds[batch_i * batch_size:(batch_i + 1) * batch_size, :] = preds.cpu().numpy()\n\n        # compute loss\n        loss      = criterion(preds, labels)\n        val_loss += loss.item() * inputs.size(0)\n        \n    # save predictions\n    oof_preds = fold_preds\n\n    # scheduler step\n    scheduler.step()\n\n\n    #EVALUATION\n\n    # evaluate performance\n    fold_preds_round = fold_preds.argmax(axis = 1)\n    val_kappa = metrics.cohen_kappa_score(data_valid['diagnosis'], fold_preds_round.astype('int'), weights = 'quadratic')\n\n    # save perfoirmance values\n    val_kappas.append(val_kappa)\n    val_losses.append(val_loss / len(data_valid))\n    trn_losses.append(trn_loss / len(data_train))\n\n\n    #EARLY STOPPING\n\n    # display info\n    print('- epoch {}/{} | lr = {} | trn_loss = {:.4f} | val_loss = {:.4f} | val_kappa = {:.4f} | {:.2f} min'.format(\n        epoch + 1, max_epochs, scheduler.get_lr()[len(scheduler.get_lr()) - 1],\n        trn_loss / len(data_train), val_loss / len(data_valid), val_kappa,\n        (time.time() - epoch_start) / 60))\n\n    # check if there is any improvement\n    if epoch > 0:       \n        if val_kappas[epoch] < val_kappas[epoch - bad_epochs - 1]:\n            bad_epochs += 1\n        else:\n            bad_epochs = 0\n\n    # save model weights if improvement\n    if bad_epochs == 0:\n        oof_preds_best = oof_preds.copy()\n        torch.save(model.state_dict(), 'models/model_{}.bin'.format(model_name))\n\n    # break if early stop\n    if bad_epochs == early_stop:\n        print('Early stopping. Best results: loss = {:.4f}, kappa = {:.4f} (epoch {})'.format(\n            np.min(val_losses), val_kappas[np.argmin(val_losses)], np.argmin(val_losses) + 1))\n        print('')\n        break\n\n    # break if max epochs\n    if epoch == (max_epochs - 1):\n        print('Did not met early stopping. Best results: loss = {:.4f}, kappa = {:.4f} (epoch {})'.format(\n            np.min(val_losses), val_kappas[np.argmin(val_losses)], np.argmin(val_losses) + 1))\n        print('')\n        break\n\n\n# load best predictions\noof_preds = oof_preds_best\n\n# print performance\nprint('')\nprint('Finished in {:.2f} minutes'.format((time.time() - cv_start) / 60)) ","metadata":{"execution":{"iopub.status.busy":"2025-06-09T14:29:16.132485Z","iopub.execute_input":"2025-06-09T14:29:16.132753Z","iopub.status.idle":"2025-06-10T00:45:51.32189Z","shell.execute_reply.started":"2025-06-09T14:29:16.132734Z","shell.execute_reply":"2025-06-10T00:45:51.320967Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name='2-4-3'></a>\n\n## 2.4.3 - Evaluation","metadata":{}},{"cell_type":"code","source":"#PLOT LOSS AND KAPPA DYNAMICS\nsns.set()\n# plot size\nfig = plt.figure(figsize = (15, 5))\n\n# plot loss dynamics\nplt.subplot(1, 2, 1)\nplt.plot(trn_losses, 'red',   label = 'Training')\nplt.plot(val_losses, 'green', label = 'Validation')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\n\n# plot kappa dynamics\nplt.subplot(1, 2, 2)\nplt.plot(val_kappas, 'blue', label = 'Kappa')\nplt.xlabel('Epoch') \nplt.ylabel('Kappa')\nplt.legend()","metadata":{"execution":{"iopub.status.busy":"2025-06-10T00:57:03.696915Z","iopub.execute_input":"2025-06-10T00:57:03.697235Z","iopub.status.idle":"2025-06-10T00:57:04.175587Z","shell.execute_reply.started":"2025-06-10T00:57:03.697207Z","shell.execute_reply":"2025-06-10T00:57:04.174934Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#RECHECK PERFORMANCE\n\n# rounding\noof_preds_round = oof_preds.argmax(axis = 1)\ncoef = [0.5, 1.5, 2.5, 3.5]\nfor i, pred in enumerate(oof_preds_round):\n    if pred < coef[0]:\n        oof_preds_round[i] = 0\n    elif pred >= coef[0] and pred < coef[1]:\n        oof_preds_round[i] = 1\n    elif pred >= coef[1] and pred < coef[2]:\n        oof_preds_round[i] = 2\n    elif pred >= coef[2] and pred < coef[3]:\n        oof_preds_round[i] = 3\n    else:\n        oof_preds_round[i] = 4\n\n# compute kappa\noof_loss  = criterion(torch.tensor(oof_preds), torch.tensor(test['diagnosis']).view(-1).type(torch.long))\noof_kappa = metrics.cohen_kappa_score(test['diagnosis'], oof_preds_round.astype('int'), weights = 'quadratic')\nprint('OOF loss  = {:.4f}'.format(oof_loss))\nprint('OOF kappa = {:.4f}'.format(oof_kappa))","metadata":{"execution":{"iopub.status.busy":"2025-06-10T00:57:15.04683Z","iopub.execute_input":"2025-06-10T00:57:15.047639Z","iopub.status.idle":"2025-06-10T00:57:15.162158Z","shell.execute_reply.started":"2025-06-10T00:57:15.047608Z","shell.execute_reply":"2025-06-10T00:57:15.161384Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#CONFUSION MATRIX\n\n# construct confusion matrix\ncm = confusion_matrix(test['diagnosis'], oof_preds_round)\ncm = cm.astype('float') / cm.sum(axis = 1)[:, np.newaxis]\nannot = np.around(cm, 2)\n\n# plot matrix\nfig, ax = plt.subplots(figsize = (10, 10))\nsns.heatmap(cm, cmap = 'Blues', annot = annot, lw = 0.5)\nax.set_xlabel('Prediction')\nax.set_ylabel('Ground Truth')\nax.set_aspect('equal')","metadata":{"execution":{"iopub.status.busy":"2025-06-10T00:57:21.866568Z","iopub.execute_input":"2025-06-10T00:57:21.866838Z","iopub.status.idle":"2025-06-10T00:57:22.173082Z","shell.execute_reply.started":"2025-06-10T00:57:21.866818Z","shell.execute_reply":"2025-06-10T00:57:22.172359Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report\n\n#Classification Report Test\nprint('\\n Classification Report in Test: \\n',classification_report(test['diagnosis'], oof_preds_round))","metadata":{"execution":{"iopub.status.busy":"2025-06-10T00:57:29.216925Z","iopub.execute_input":"2025-06-10T00:57:29.217209Z","iopub.status.idle":"2025-06-10T00:57:29.239134Z","shell.execute_reply.started":"2025-06-10T00:57:29.217187Z","shell.execute_reply":"2025-06-10T00:57:29.23832Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name =\"3\"> </a>\n# 3 - Training on Diabetic Retinopathy 2019 data\n<center><img src = \"https://i.imgur.com/pzwzeQj.png\"/></center>","metadata":{}},{"cell_type":"markdown","source":"<a name =\"3-1\"> </a>\n## 3.1 Import Data","metadata":{}},{"cell_type":"code","source":"# import data\ntrain = pd.read_csv('../input/aptos2019-blindness-detection/train.csv')\ntest  = pd.read_csv('../input/aptos2019-blindness-detection/sample_submission.csv')\n\n# check shape\nprint(train.shape, test.shape)\nprint('-' * 15)\nprint(train['diagnosis'].value_counts(normalize = True))","metadata":{"execution":{"iopub.status.busy":"2025-06-10T00:57:47.207178Z","iopub.execute_input":"2025-06-10T00:57:47.207483Z","iopub.status.idle":"2025-06-10T00:57:47.231503Z","shell.execute_reply.started":"2025-06-10T00:57:47.207458Z","shell.execute_reply":"2025-06-10T00:57:47.230668Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name='3-2'></a>\n## 3.2 Explore Data","metadata":{}},{"cell_type":"code","source":"# CLASS DISTRIBUTION\n\n# plot\nfig = plt.figure(figsize = (15, 5))\nplt.hist(train['diagnosis'])\nplt.title('Class Distribution')\nplt.ylabel('Number of examples')\nplt.xlabel('Diagnosis')","metadata":{"execution":{"iopub.status.busy":"2025-06-10T00:58:02.02731Z","iopub.execute_input":"2025-06-10T00:58:02.027993Z","iopub.status.idle":"2025-06-10T00:58:02.322038Z","shell.execute_reply.started":"2025-06-10T00:58:02.027868Z","shell.execute_reply":"2025-06-10T00:58:02.32116Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# transformations\nsample_trans = transforms.Compose([transforms.ToPILImage(),\n                                   transforms.ToTensor(),\n                                  ])\nsample = Data(data       = train.iloc[0:10], \n                      directory  = '../input/aptos2019-blindness-detection/train_images',\n                      transform  = sample_trans,\n                      itype ='.png')\n\n# data loader\nsample_loader = torch.utils.data.DataLoader(dataset     = sample, \n                                            batch_size  = 10, \n                                            shuffle     = False, \n                                            num_workers = 4)\n\n# display images\nfor batch_i, data in enumerate(sample_loader):\n\n    # extract data\n    inputs = data['image']\n    labels = data['label'].view(-1, 1)\n    \n    # create plot\n    fig = plt.figure(figsize = (15, 7))\n    for i in range(len(labels)):\n        ax = fig.add_subplot(2, int(len(labels)/2), i + 1, xticks = [], yticks = [])     \n        plt.imshow(inputs[i].numpy().transpose(1, 2, 0))\n        ax.set_title(labels.numpy()[i])\n\n    break","metadata":{"execution":{"iopub.status.busy":"2025-06-10T00:58:10.462567Z","iopub.execute_input":"2025-06-10T00:58:10.463282Z","iopub.status.idle":"2025-06-10T00:58:13.868829Z","shell.execute_reply.started":"2025-06-10T00:58:10.463257Z","shell.execute_reply":"2025-06-10T00:58:13.867937Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# IMAGE SIZES\n\n# placeholder\nimage_stats = []\n\n# import loop\nfor index, observation in tqdm(train.iterrows(), total = len(train)):\n    \n    # import image\n    img = cv2.imread('../input/aptos2019-blindness-detection/train_images/{}.png'.format(observation['id_code']))\n\n    # compute stats\n    height, width, channels = img.shape\n    ratio = width / height\n    \n    # save\n    image_stats.append(np.array((observation['diagnosis'], height, width, channels, ratio)))\n\n# construct DF\nimage_stats = pd.DataFrame(image_stats)\nimage_stats.columns = ['diagnosis', 'height', 'width', 'channels', 'ratio']","metadata":{"execution":{"iopub.status.busy":"2025-06-10T00:58:21.867804Z","iopub.execute_input":"2025-06-10T00:58:21.868469Z","iopub.status.idle":"2025-06-10T01:03:57.261829Z","shell.execute_reply.started":"2025-06-10T00:58:21.86844Z","shell.execute_reply":"2025-06-10T01:03:57.260946Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# IMAGE SIZE DISTRIBUTION\n\nfig = plt.figure(figsize = (15, 5))\n\n# width\nplt.subplot(1, 3, 1)\nplt.hist(image_stats['width'])\nplt.title('(a) Image Width')\nplt.ylabel('Number of examples')\nplt.xlabel('Width')\n\n# height\nplt.subplot(1, 3, 2)\nplt.hist(image_stats['height'])\nplt.title('(b) Image Height')\nplt.ylabel('Number of examples')\nplt.xlabel('Height')\n\n# ratio\nplt.subplot(1, 3, 3)\nplt.hist(image_stats['ratio'])\nplt.title('(c) Aspect Ratio')\nplt.ylabel('Number of examples')\nplt.xlabel('Ratio')","metadata":{"execution":{"iopub.status.busy":"2025-06-10T01:04:03.398007Z","iopub.execute_input":"2025-06-10T01:04:03.398662Z","iopub.status.idle":"2025-06-10T01:04:04.039555Z","shell.execute_reply.started":"2025-06-10T01:04:03.398637Z","shell.execute_reply":"2025-06-10T01:04:04.038688Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name =\"3-3\"> </a>\n## 3.3 Preprocess and Augment Data","metadata":{}},{"cell_type":"code","source":"#TRANSFORMATIONS\n\n# parameters\nbatch_size = 25\nimage_size = 256\n\n# train transformations\ntrain_trans = transforms.Compose([transforms.ToPILImage(),\n                                  transforms.RandomRotation((-360, 360)),\n                                  transforms.RandomHorizontalFlip(),\n                                  transforms.RandomVerticalFlip(),\n                                  transforms.ToTensor(),\n                                 ])\n\n# valid transformations\nvalid_trans = transforms.Compose([transforms.ToPILImage(),\n                                  transforms.ToTensor(),\n                                 ])\n                                 \n# test transformations\ntest_trans = valid_trans","metadata":{"execution":{"iopub.status.busy":"2025-06-10T01:04:15.317244Z","iopub.execute_input":"2025-06-10T01:04:15.317534Z","iopub.status.idle":"2025-06-10T01:04:15.322413Z","shell.execute_reply.started":"2025-06-10T01:04:15.317513Z","shell.execute_reply":"2025-06-10T01:04:15.321672Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#EXAMINE FIRST BATCH (TRAIN)\n\n# get dataset\nsample = EyeData(data = train.iloc[0:10], \n                      directory  = '../input/aptos2019-blindness-detection/train_images',\n                      transform  = train_trans,\n                      itype ='.png' )\n\n# data loader\nsample_loader = torch.utils.data.DataLoader(dataset     = sample, \n                                            batch_size  = batch_size, \n                                            shuffle     = True, \n                                            num_workers = 4)\n\n# display images\nfor batch_i, data in enumerate(sample_loader):\n\n    # extract data\n    inputs = data['image']\n    labels = data['label'].view(-1, 1)\n    \n    # create plot\n    fig = plt.figure(figsize = (20,10))\n    for i in range(len(labels)):\n        ax = fig.add_subplot(2, int(len(labels)/2), i + 1, xticks = [], yticks = [])     \n        plt.imshow(inputs[i].numpy().transpose(1, 2, 0))\n        ax.set_title(labels.numpy()[i])\n\n    break","metadata":{"execution":{"iopub.status.busy":"2025-06-10T01:04:21.483099Z","iopub.execute_input":"2025-06-10T01:04:21.483351Z","iopub.status.idle":"2025-06-10T01:04:25.421478Z","shell.execute_reply.started":"2025-06-10T01:04:21.483334Z","shell.execute_reply":"2025-06-10T01:04:25.420674Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#EXAMINE FIRST BATCH (TEST)\n\n# get dataset\nsample = EyeData(data       = test.iloc[0:10], \n                      directory  = '../input/aptos2019-blindness-detection/test_images',\n                      transform  = test_trans,\n                      itype = '.png',\n                      do_random_crop = False)\n\n# data loader\nsample_loader = torch.utils.data.DataLoader(dataset     = sample, \n                                            batch_size  = batch_size, \n                                            shuffle     = False, \n                                            num_workers = 4)\n\n# display images\nfor batch_i, data in enumerate(sample_loader):\n\n    # extract data\n    inputs = data['image']\n    \n    # create plot\n    fig = plt.figure(figsize = (20,10))\n    for i in range(10):\n        ax = fig.add_subplot(2, int(10/2), i + 1, xticks = [], yticks = [])     \n        plt.imshow(inputs[i].numpy().transpose(1, 2, 0))\n\n    break","metadata":{"execution":{"iopub.status.busy":"2025-06-10T01:04:35.39782Z","iopub.execute_input":"2025-06-10T01:04:35.398565Z","iopub.status.idle":"2025-06-10T01:04:37.180158Z","shell.execute_reply.started":"2025-06-10T01:04:35.39854Z","shell.execute_reply":"2025-06-10T01:04:37.179008Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name=\"3-4\"></a> \n## 3.4 - Load Pre-Trained EFfienctNetB7 trained on 2015 data\n<center><img src = \"https://miro.medium.com/max/1400/1*8oE4jOMfOXeEzgsHjSB5ww.png\"/></center>","metadata":{}},{"cell_type":"markdown","source":"<a name =\"3-4-1\"></a>\n### 3.4.1 Setup Model and Choose Hyperparameters","metadata":{}},{"cell_type":"code","source":"#MODEL ARCHITECTURE\n\n# model name\nmodel_name = 'enet_b7'\n\n# check architecture\nmodel = init_model(model_name = model_name)\nprint(model)","metadata":{"execution":{"iopub.status.busy":"2025-06-10T01:06:10.883912Z","iopub.execute_input":"2025-06-10T01:06:10.884204Z","iopub.status.idle":"2025-06-10T01:06:11.751925Z","shell.execute_reply.started":"2025-06-10T01:06:10.884183Z","shell.execute_reply":"2025-06-10T01:06:11.750989Z"},"scrolled":true,"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#VALIDATION SETTINGS\nfrom sklearn.model_selection import KFold, StratifiedKFold\n# no. folds\nnum_folds = 4\n\n# creating splits\nskf    = StratifiedKFold(n_splits = num_folds, shuffle = True, random_state = seed)\nsplits = list(skf.split(train['id_code'], train['diagnosis']))\n\n# placeholders\noof_preds = np.zeros((len(train), 1))\n\n# timer\ncv_start = time.time()","metadata":{"execution":{"iopub.status.busy":"2025-06-10T01:05:49.264086Z","iopub.execute_input":"2025-06-10T01:05:49.264554Z","iopub.status.idle":"2025-06-10T01:05:49.303136Z","shell.execute_reply.started":"2025-06-10T01:05:49.264522Z","shell.execute_reply":"2025-06-10T01:05:49.302614Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#PARAMETERS\n\n# loss function\ncriterion = nn.CrossEntropyLoss()\n\n# epochs\nmax_epochs = 15\nearly_stop = 5\n\n# learning rates\neta = 1e-3\n\n# scheduler\nstep  = 5\ngamma = 0.5","metadata":{"execution":{"iopub.status.busy":"2025-06-10T01:05:58.753118Z","iopub.execute_input":"2025-06-10T01:05:58.753804Z","iopub.status.idle":"2025-06-10T01:05:58.757587Z","shell.execute_reply.started":"2025-06-10T01:05:58.753782Z","shell.execute_reply":"2025-06-10T01:05:58.756831Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name = \"3-4-2\"></a>\n## 3.4.2 Train Model on Diabetic Retinopathy 2019 dat","metadata":{}},{"cell_type":"code","source":"#CROSS-VALIDATION LOOP\nfor fold in tqdm(range(num_folds)):\n    \n    \n    #DATA PREPARATION\n\n    # display information\n    print('-' * 30)\n    print('FOLD {}/{}'.format(fold + 1, num_folds))\n    print('-' * 30)\n\n    # load splits\n    data_train = train.iloc[splits[fold][0]].reset_index(drop = True)\n    data_valid = train.iloc[splits[fold][1]].reset_index(drop = True)\n\n    # create datasets\n    train_dataset = EyeData(data      = data_train, \n                                 directory = '../input/aptos2019-blindness-detection/train_images',\n                                 transform = train_trans,\n                                 itype = '.png')\n    valid_dataset = EyeData(data      = data_valid, \n                                 directory = '../input/aptos2019-blindness-detection/train_images',\n                                 transform = valid_trans,\n                                 itype = '.png')\n\n    # create data loaders\n    train_loader = torch.utils.data.DataLoader(train_dataset, \n                                               batch_size  = batch_size, \n                                               shuffle     = True, \n                                               num_workers = 4)\n    valid_loader = torch.utils.data.DataLoader(valid_dataset, \n                                               batch_size  = batch_size, \n                                               shuffle     = False, \n                                               num_workers = 4)\n    \n    \n    #MODEL PREPARATION\n    \n    # placeholders\n    val_kappas = []\n    val_losses = []\n    trn_losses = []\n    bad_epochs = 0\n    \n    # load best OOF predictions\n    if fold > 0:\n        oof_preds = oof_preds_best.copy()\n    \n    # initialize and send to GPU\n    model = init_model(train = True)\n    model = model.to(device)\n\n    # optimizer\n    optimizer = optim.Adam(model._fc.parameters(), lr = eta)\n    scheduler = lr_scheduler.StepLR(optimizer, step_size = step, gamma = gamma)\n    \n    \n    #TRAINING AND VALIDATION LOOP\n    for epoch in range(max_epochs):\n\n        ## PREPARATION\n\n        # timer\n        epoch_start = time.time()\n\n        # reset losses\n        trn_loss = 0.0\n        val_loss = 0.0\n\n        # placeholders\n        fold_preds = np.zeros((len(data_valid), 1))\n\n\n        # TRAINING\n\n        # switch regime\n        model.train()\n        \n        # loop through batches\n        for batch_i, data in enumerate(train_loader):\n\n            # extract inputs and labels\n            inputs = data['image']\n            labels = data['label'].view(-1)\n            inputs = inputs.to(device, dtype = torch.float)\n            labels = labels.to(device, dtype = torch.long)\n            optimizer.zero_grad()\n\n            # forward and backward pass\n            with torch.set_grad_enabled(True):\n                preds = model(inputs)\n                loss  = criterion(preds, labels)\n                loss.backward()\n                optimizer.step()\n\n            # compute loss\n            trn_loss += loss.item() * inputs.size(0)\n\n\n        # INFERENCE\n        \n        # initialize\n        model.eval()\n\n        # loop through batches\n        for batch_i, data in enumerate(valid_loader):\n\n            # extract inputs and labels\n            inputs = data['image']\n            labels = data['label'].view(-1)\n            inputs = inputs.to(device, dtype = torch.float)\n            labels = labels.to(device, dtype = torch.long)\n\n            # compute predictions\n            with torch.set_grad_enabled(False):\n                preds = model(inputs).detach()\n                _, class_preds = preds.topk(1)\n                fold_preds[batch_i * batch_size:(batch_i + 1) * batch_size, :] = class_preds.cpu().numpy()\n\n            # compute loss\n            loss      = criterion(preds, labels)\n            val_loss += loss.item() * inputs.size(0)\n\n        # save predictions\n        oof_preds[splits[fold][1]] = fold_preds\n        \n        # scheduler step\n        scheduler.step()\n\n\n        # EVALUATION\n\n        # evaluate performance\n        fold_preds_round = fold_preds\n        val_kappa = metrics.cohen_kappa_score(data_valid['diagnosis'], fold_preds_round.astype('int'), weights = 'quadratic')\n        \n        # save perfoirmance values\n        val_kappas.append(val_kappa)\n        val_losses.append(val_loss / len(data_valid))\n        trn_losses.append(trn_loss / len(data_train))\n\n        \n        # EARLY STOPPING\n        \n        # display info\n        print('- epoch {}/{} | lr = {} | trn_loss = {:.4f} | val_loss = {:.4f} | val_kappa = {:.4f} | {:.2f} min'.format(\n            epoch + 1, max_epochs, scheduler.get_lr()[len(scheduler.get_lr()) - 1],\n            trn_loss / len(data_train), val_loss / len(data_valid), val_kappa,\n            (time.time() - epoch_start) / 60))\n        \n        # check if there is any improvement\n        if epoch > 0:       \n            if val_kappas[epoch] < val_kappas[epoch - bad_epochs - 1]:\n                bad_epochs += 1\n            else:\n                bad_epochs = 0\n\n        # save model weights if improvement\n        if bad_epochs == 0:\n            oof_preds_best = oof_preds.copy()\n            torch.save(model.state_dict(), 'models/model_{}_fold{}.bin'.format(model_name, fold + 1))\n\n        # break if early stop\n        if bad_epochs == early_stop:\n            print('Early stopping. Best results: loss = {:.4f}, kappa = {:.4f} (epoch {})'.format(\n                np.min(val_losses), val_kappas[np.argmin(val_losses)], np.argmin(val_losses) + 1))\n            print('')\n            break\n\n        # break if max epochs\n        if epoch == (max_epochs - 1):\n            print('Did not meet early stopping. Best results: loss = {:.4f}, kappa = {:.4f} (epoch {})'.format(\n                np.min(val_losses), val_kappas[np.argmin(val_losses)], np.argmin(val_losses) + 1))\n            print('')\n            break\n        \n\n# load best predictions\noof_preds = oof_preds_best\n\n# print performance\nprint('')\nprint('Finished in {:.2f} minutes'.format((time.time() - cv_start) / 60))","metadata":{"execution":{"iopub.status.busy":"2022-11-17T08:16:04.228251Z","iopub.execute_input":"2022-11-17T08:16:04.228646Z","iopub.status.idle":"2022-11-17T13:18:18.088969Z","shell.execute_reply.started":"2022-11-17T08:16:04.228612Z","shell.execute_reply":"2022-11-17T13:18:18.087833Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a name =\"3-4-3\"></a>\n## 3.4.3 - Evaluation\n","metadata":{}},{"cell_type":"code","source":"# PLOT LOSS AND KAPPA DYNAMICS\nsns.set()\n# plot size\nfig = plt.figure(figsize = (15, 5))\n\n# plot loss dynamics\nplt.subplot(1, 2, 1)\nplt.plot(trn_losses, 'red',   label = 'Training')\nplt.plot(val_losses, 'green', label = 'Validation')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\n\n# plot kappa dynamics\nplt.subplot(1, 2, 2)\nplt.plot(val_kappas, 'blue', label = 'Kappa')\nplt.xlabel('Epoch')\nplt.ylabel('Kappa')\nplt.legend()","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:26:56.839058Z","iopub.execute_input":"2022-11-17T13:26:56.83943Z","iopub.status.idle":"2022-11-17T13:26:57.322471Z","shell.execute_reply.started":"2022-11-17T13:26:56.839397Z","shell.execute_reply":"2022-11-17T13:26:57.321535Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#RECHECK PERFORMANCE\n\n# evaluate performance\noof_preds_round = oof_preds.copy()\noof_kappa = metrics.cohen_kappa_score(train['diagnosis'], oof_preds_round.astype('int'), weights = 'quadratic')\nprint('OOF kappa = {:.4f}'.format(oof_kappa))","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:29:03.343367Z","iopub.execute_input":"2022-11-17T13:29:03.344379Z","iopub.status.idle":"2022-11-17T13:29:03.35379Z","shell.execute_reply.started":"2022-11-17T13:29:03.344339Z","shell.execute_reply":"2022-11-17T13:29:03.352626Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#CONFUSION MATRIX\n\n# construct confusion matrx\ncm = confusion_matrix(train['diagnosis'], oof_preds_round)\ncm = cm.astype('float') / cm.sum(axis = 1)[:, np.newaxis]\nannot = np.around(cm, 2)\n\n# plot matrix\nfig, ax = plt.subplots(figsize = (8, 6))\nsns.heatmap(cm, cmap = 'Blues', annot = annot, lw = 0.5)\nax.set_xlabel('Prediction')\nax.set_ylabel('Ground Truth')\nax.set_aspect('equal')","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:29:33.93938Z","iopub.execute_input":"2022-11-17T13:29:33.939781Z","iopub.status.idle":"2022-11-17T13:29:34.265431Z","shell.execute_reply.started":"2022-11-17T13:29:33.939749Z","shell.execute_reply":"2022-11-17T13:29:34.264411Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report\n\n#Classification Report Test\nprint('\\n Classification Report in Test: \\n',classification_report(train['diagnosis'], oof_preds_round))","metadata":{"execution":{"iopub.status.busy":"2022-11-17T13:32:16.488389Z","iopub.execute_input":"2022-11-17T13:32:16.488805Z","iopub.status.idle":"2022-11-17T13:32:16.507564Z","shell.execute_reply.started":"2022-11-17T13:32:16.488771Z","shell.execute_reply":"2022-11-17T13:32:16.506468Z"},"trusted":true},"outputs":[],"execution_count":null}]}