{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":14774,"databundleVersionId":875431,"sourceType":"competition"},{"sourceId":418031,"sourceType":"datasetVersion","datasetId":131128},{"sourceId":111254110,"sourceType":"kernelVersion"}],"dockerImageVersionId":30301,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<center><h1>Diabetic Retinopathy Detection</h1></center>","metadata":{}},{"cell_type":"markdown","source":"## Table of Content\n<ul>\n    <li><a href= \"#1\">1 - Introduction and Create Workspace</a></li>\n    <ul>\n        <li><a href= \"#1-1\">    1.1 Import Packages</a></li>\n        <li><a href= \"#1-2\">    1.2 Create Classes and Function</a></li>\n    </ul>\n    <li><a href= \"#2\">2 - Pre-Training on Diabetic Retinopathy 2015 data</a></li>\n    <ul>\n    <li><a href= \"#2-1\">    2.1 Import Data</a></li>\n    <li><a href= \"#2-2\">    2.2 Explore Data</a></li>\n    <li><a href= \"#2-3\">    2.3 Preprocess and Augment Data</a></li>\n    <li><a href= \"#2-4\">    2.4 - Using EFfienctNetB7 on Diabetic Retinopathy 2015 data</a></li>\n        <ul>\n    <li><a href= \"#2-4-1\">        2.4.1 Setup Model and Choose Hyperparameters</a></li>\n    <li><a href= \"#2-4-2\">        2.4.2 Train Model on Diabetic Retinopathy 2015 data</a></li>\n    <li><a href= \"#2-4-3\">        2.4.3 Evaluation</a></li>\n         </ul>   \n    </ul>\n    <li><a href= \"#3\">3 - Training on Diabetic Retinopathy 2019 data</a></li>\n    <ul>\n    <li><a href= \"#3-1\">    3.1 Import Data</a></li>\n    <li><a href= \"#3-2\">    3.2 Explore Data</a></li>\n    <li><a href= \"#3-3\">    3.3 Preprocess and Augment Data</a></li>\n    <li><a href= \"#3-4\">    3.4 - Load Pre-Trained EFfienctNetB7 trained on 2015 data</a></li>\n        <ul>\n    <li><a href= \"#3-4-1\">        3.4.1 Setup Model and Choose Hyperparameters</a></li>\n    <li><a href= \"#3-4-2\">        3.4.2 Train Model on Diabetic Retinopathy 2019 data</a></li>\n    <li><a href= \"#3-4-3\">        3.4.3 Evaluation</a></li>\n        </ul>\n        </ul>\n</ul>","metadata":{}},{"cell_type":"markdown","source":"# 1 - Introducation and Create Workspace\n","metadata":{}},{"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\n\n!pip install ipywidgets==7.7.2\n","metadata":{"execution":{"iopub.status.busy":"2025-12-02T15:13:52.250928Z","iopub.execute_input":"2025-12-02T15:13:52.251480Z","iopub.status.idle":"2025-12-02T15:14:20.746454Z","shell.execute_reply.started":"2025-12-02T15:13:52.251338Z","shell.execute_reply":"2025-12-02T15:14:20.745137Z"},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.749514Z","iopub.execute_input":"2025-12-02T15:14:20.750434Z","iopub.status.idle":"2025-12-02T15:14:20.757403Z","shell.execute_reply.started":"2025-12-02T15:14:20.750394Z","shell.execute_reply":"2025-12-02T15:14:20.756406Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_size = 244\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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.758679Z","iopub.execute_input":"2025-12-02T15:14:20.759011Z","iopub.status.idle":"2025-12-02T15:14:20.771934Z","shell.execute_reply.started":"2025-12-02T15:14:20.758943Z","shell.execute_reply":"2025-12-02T15:14:20.770798Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.773481Z","iopub.execute_input":"2025-12-02T15:14:20.774136Z","iopub.status.idle":"2025-12-02T15:14:20.789624Z","shell.execute_reply.started":"2025-12-02T15:14:20.774094Z","shell.execute_reply":"2025-12-02T15:14:20.788521Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.791031Z","iopub.execute_input":"2025-12-02T15:14:20.791382Z","iopub.status.idle":"2025-12-02T15:14:20.804476Z","shell.execute_reply.started":"2025-12-02T15:14:20.791354Z","shell.execute_reply":"2025-12-02T15:14:20.803479Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.807070Z","iopub.execute_input":"2025-12-02T15:14:20.807765Z","iopub.status.idle":"2025-12-02T15:14:20.852021Z","shell.execute_reply.started":"2025-12-02T15:14:20.807717Z","shell.execute_reply":"2025-12-02T15:14:20.851157Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.855437Z","iopub.execute_input":"2025-12-02T15:14:20.855792Z","iopub.status.idle":"2025-12-02T15:14:20.869394Z","shell.execute_reply.started":"2025-12-02T15:14:20.855762Z","shell.execute_reply":"2025-12-02T15:14:20.868317Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.870814Z","iopub.execute_input":"2025-12-02T15:14:20.871265Z","iopub.status.idle":"2025-12-02T15:14:20.885415Z","shell.execute_reply.started":"2025-12-02T15:14:20.871226Z","shell.execute_reply":"2025-12-02T15:14:20.884231Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def init_model(train= True, \n               trn_layers = 2,\n               model_name = 'enet_b0'):\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-b0', 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-b0', 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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.887003Z","iopub.execute_input":"2025-12-02T15:14:20.887424Z","iopub.status.idle":"2025-12-02T15:14:20.908568Z","shell.execute_reply.started":"2025-12-02T15:14:20.887384Z","shell.execute_reply":"2025-12-02T15:14:20.907321Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.909905Z","iopub.execute_input":"2025-12-02T15:14:20.910316Z","iopub.status.idle":"2025-12-02T15:14:20.926886Z","shell.execute_reply.started":"2025-12-02T15:14:20.910284Z","shell.execute_reply":"2025-12-02T15:14:20.925663Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#RANDOMNESS\n\nseed = 23\nseed_everything(seed)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.928451Z","iopub.execute_input":"2025-12-02T15:14:20.928881Z","iopub.status.idle":"2025-12-02T15:14:20.940027Z","shell.execute_reply.started":"2025-12-02T15:14:20.928841Z","shell.execute_reply":"2025-12-02T15:14:20.939043Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:20.941626Z","iopub.execute_input":"2025-12-02T15:14:20.942514Z","iopub.status.idle":"2025-12-02T15:14:21.012170Z","shell.execute_reply.started":"2025-12-02T15:14:20.942481Z","shell.execute_reply":"2025-12-02T15:14:21.010973Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:21.013361Z","iopub.execute_input":"2025-12-02T15:14:21.013675Z","iopub.status.idle":"2025-12-02T15:14:21.472172Z","shell.execute_reply.started":"2025-12-02T15:14:21.013647Z","shell.execute_reply":"2025-12-02T15:14:21.470802Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:21.473790Z","iopub.execute_input":"2025-12-02T15:14:21.474317Z","iopub.status.idle":"2025-12-02T15:14:22.813424Z","shell.execute_reply.started":"2025-12-02T15:14:21.474274Z","shell.execute_reply":"2025-12-02T15:14:22.812316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport cv2\nfrom tqdm import tqdm\n\nimage_stats = []\n\nfor index, observation in tqdm(train.iterrows(), total=len(train)):\n    path = f\"../input/diabetic-retinopathy-resized/resized_train/resized_train/{observation['id_code']}.jpeg\"\n    img = cv2.imread(path)\n\n    if img is None:\n        print (\"Image not found:\", path)\n        # Skip if image not found or unreadable\n        continue\n\n    height, width, channels = img.shape\n    ratio = width / height\n\n    image_stats.append([observation['diagnosis'], height, width, channels, ratio])\n\n# construct DF\nimage_stats = pd.DataFrame(image_stats, columns=['diagnosis', 'height', 'width', 'channels', 'ratio'])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:14:22.815418Z","iopub.execute_input":"2025-12-02T15:14:22.816246Z","iopub.status.idle":"2025-12-02T15:24:17.845228Z","shell.execute_reply.started":"2025-12-02T15:14:22.816190Z","shell.execute_reply":"2025-12-02T15:24:17.843995Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:24:17.846985Z","iopub.execute_input":"2025-12-02T15:24:17.847441Z","iopub.status.idle":"2025-12-02T15:24:18.341496Z","shell.execute_reply.started":"2025-12-02T15:24:17.847398Z","shell.execute_reply":"2025-12-02T15:24:18.340293Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:24:18.342891Z","iopub.execute_input":"2025-12-02T15:24:18.343332Z","iopub.status.idle":"2025-12-02T15:24:18.350106Z","shell.execute_reply.started":"2025-12-02T15:24:18.343300Z","shell.execute_reply":"2025-12-02T15:24:18.349003Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:24:18.351446Z","iopub.execute_input":"2025-12-02T15:24:18.351788Z","iopub.status.idle":"2025-12-02T15:24:19.696717Z","shell.execute_reply.started":"2025-12-02T15:24:18.351760Z","shell.execute_reply":"2025-12-02T15:24:19.695540Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:24:19.698446Z","iopub.execute_input":"2025-12-02T15:24:19.698803Z","iopub.status.idle":"2025-12-02T15:24:23.955571Z","shell.execute_reply.started":"2025-12-02T15:24:19.698770Z","shell.execute_reply":"2025-12-02T15:24:23.954417Z"}},"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_b0'\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-b0', num_classes = 5)\n        \n    ### inference mode\n    if train == False:\n        \n        # load pre-trained model\n        model = EfficientNet.from_name('efficientnet-b0')\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,"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:24:23.957346Z","iopub.execute_input":"2025-12-02T15:24:23.958182Z","iopub.status.idle":"2025-12-02T15:24:24.458252Z","shell.execute_reply.started":"2025-12-02T15:24:23.958142Z","shell.execute_reply":"2025-12-02T15:24:24.456886Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:24:24.459512Z","iopub.execute_input":"2025-12-02T15:24:24.459924Z","iopub.status.idle":"2025-12-02T15:24:24.465161Z","shell.execute_reply.started":"2025-12-02T15:24:24.459893Z","shell.execute_reply":"2025-12-02T15:24:24.463869Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:24:24.469674Z","iopub.execute_input":"2025-12-02T15:24:24.470083Z","iopub.status.idle":"2025-12-02T15:24:24.480103Z","shell.execute_reply.started":"2025-12-02T15:24:24.470051Z","shell.execute_reply":"2025-12-02T15:24:24.479012Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:24:24.481475Z","iopub.execute_input":"2025-12-02T15:24:24.481834Z","iopub.status.idle":"2025-12-02T15:24:24.495546Z","shell.execute_reply.started":"2025-12-02T15:24:24.481805Z","shell.execute_reply":"2025-12-02T15:24:24.494404Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-02T15:24:52.933942Z","iopub.execute_input":"2025-12-02T15:24:52.935151Z"}},"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":{"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":{"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":{"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":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3 - Training on Diabetic Retinopathy 2019 data","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":{"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":{"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":{"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":{"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":{"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":{"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":{"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":{"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_b0'\n\n# check architecture\nmodel = init_model(model_name = model_name)\nprint(model)","metadata":{"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":{"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":{"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":{"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":{"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":{"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":{"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":{"trusted":true},"outputs":[],"execution_count":null}]}