{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Imports and gpu","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport json\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    classification_report,\n    confusion_matrix,\n    accuracy_score,\n    precision_score,\n    recall_score,\n    f1_score\n)\n\ndevice = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nprint(\"Device:\", device)\n\nif torch.cuda.is_available():\n    print(\"GPU:\", torch.cuda.get_device_name(0))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-09-07T10:58:30.508540Z","iopub.execute_input":"2026-09-07T10:58:30.509428Z","iopub.status.idle":"2026-09-07T10:58:30.516274Z","shell.execute_reply.started":"2026-09-07T10:58:30.509395Z","shell.execute_reply":"2026-09-07T10:58:30.515258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 42\n\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\n\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\nprint(\"Seed:\", SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T10:59:13.418651Z","iopub.execute_input":"2026-09-07T10:59:13.419323Z","iopub.status.idle":"2026-09-07T10:59:13.425280Z","shell.execute_reply.started":"2026-09-07T10:59:13.419294Z","shell.execute_reply":"2026-09-07T10:59:13.424553Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for dirname, _, filenames in os.walk(\"/kaggle/input\"):\n    print(dirname)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T10:59:22.784311Z","iopub.execute_input":"2026-09-07T10:59:22.784949Z","iopub.status.idle":"2026-09-07T10:59:32.950727Z","shell.execute_reply.started":"2026-09-07T10:59:22.784919Z","shell.execute_reply":"2026-09-07T10:59:32.950008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/competitions/aptos2019-blindness-detection\"\n\nCSV_PATH = os.path.join(\n    DATA_DIR,\n    \"train.csv\"\n)\n\nIMAGE_DIR = os.path.join(\n    DATA_DIR,\n    \"train_images\"\n)\n\nprint(\"CSV:\", CSV_PATH)\nprint(\"Images:\", IMAGE_DIR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T10:59:36.229014Z","iopub.execute_input":"2026-09-07T10:59:36.229818Z","iopub.status.idle":"2026-09-07T10:59:36.234243Z","shell.execute_reply.started":"2026-09-07T10:59:36.229789Z","shell.execute_reply":"2026-09-07T10:59:36.233489Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(CSV_PATH)\n\nprint(df.head())\n\nprint(\"\\nNumber of labeled images:\", len(df))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T10:59:49.349956Z","iopub.execute_input":"2026-09-07T10:59:49.350207Z","iopub.status.idle":"2026-09-07T10:59:49.385603Z","shell.execute_reply.started":"2026-09-07T10:59:49.350185Z","shell.execute_reply":"2026-09-07T10:59:49.384997Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_counts = (\n    df[\"diagnosis\"]\n    .value_counts()\n    .sort_index()\n)\n\nprint(class_counts)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T10:59:58.834021Z","iopub.execute_input":"2026-09-07T10:59:58.834756Z","iopub.status.idle":"2026-09-07T10:59:58.857345Z","shell.execute_reply.started":"2026-09-07T10:59:58.834727Z","shell.execute_reply":"2026-09-07T10:59:58.856744Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_names = [\n    \"No DR\",\n    \"Mild\",\n    \"Moderate\",\n    \"Severe\",\n    \"Proliferative DR\"\n]\n\nplt.figure(figsize=(8, 5))\n\nplt.bar(\n    class_names,\n    class_counts.values\n)\n\nplt.xticks(rotation=20)\nplt.ylabel(\"Number of Images\")\nplt.title(\"APTOS Class Distribution\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:00:09.817038Z","iopub.execute_input":"2026-09-07T11:00:09.817943Z","iopub.status.idle":"2026-09-07T11:00:10.009203Z","shell.execute_reply.started":"2026-09-07T11:00:09.817911Z","shell.execute_reply":"2026-09-07T11:00:10.008613Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df[\"image_path\"] = df[\"id_code\"].apply(\n    lambda x: os.path.join(\n        IMAGE_DIR,\n        x + \".png\"\n    )\n)\n\nprint(df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:00:28.094100Z","iopub.execute_input":"2026-09-07T11:00:28.094364Z","iopub.status.idle":"2026-09-07T11:00:28.105120Z","shell.execute_reply.started":"2026-09-07T11:00:28.094341Z","shell.execute_reply":"2026-09-07T11:00:28.104436Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"missing_images = df[\n    ~df[\"image_path\"].apply(os.path.exists)\n]\n\nprint(\n    \"Missing images:\",\n    len(missing_images)\n)\n\nif len(missing_images) > 0:\n    print(missing_images.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:00:42.071377Z","iopub.execute_input":"2026-09-07T11:00:42.072161Z","iopub.status.idle":"2026-09-07T11:00:45.563797Z","shell.execute_reply.started":"2026-09-07T11:00:42.072133Z","shell.execute_reply":"2026-09-07T11:00:45.563017Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df, temp_df = train_test_split(\n    df,\n    test_size=0.30,\n    stratify=df[\"diagnosis\"],\n    random_state=SEED\n)\n\nval_df, test_df = train_test_split(\n    temp_df,\n    test_size=0.50,\n    stratify=temp_df[\"diagnosis\"],\n    random_state=SEED\n)\n\nprint(\"Train:\", len(train_df))\nprint(\"Validation:\", len(val_df))\nprint(\"Test:\", len(test_df))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:00:56.093879Z","iopub.execute_input":"2026-09-07T11:00:56.094296Z","iopub.status.idle":"2026-09-07T11:00:56.107591Z","shell.execute_reply.started":"2026-09-07T11:00:56.094267Z","shell.execute_reply":"2026-09-07T11:00:56.106847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"TRAIN\")\nprint(\n    train_df[\"diagnosis\"]\n    .value_counts()\n    .sort_index()\n)\n\nprint(\"\\nVALIDATION\")\nprint(\n    val_df[\"diagnosis\"]\n    .value_counts()\n    .sort_index()\n)\n\nprint(\"\\nTEST\")\nprint(\n    test_df[\"diagnosis\"]\n    .value_counts()\n    .sort_index()\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:01:06.165055Z","iopub.execute_input":"2026-09-07T11:01:06.165624Z","iopub.status.idle":"2026-09-07T11:01:06.174301Z","shell.execute_reply.started":"2026-09-07T11:01:06.165592Z","shell.execute_reply":"2026-09-07T11:01:06.173634Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n\n    transforms.RandomHorizontalFlip(\n        p=0.5\n    ),\n\n    transforms.RandomRotation(10),\n\n    transforms.ColorJitter(\n        brightness=0.10,\n        contrast=0.10,\n        saturation=0.05\n    ),\n\n    transforms.ToTensor(),\n\n    transforms.Normalize(\n        mean=[\n            0.485,\n            0.456,\n            0.406\n        ],\n        std=[\n            0.229,\n            0.224,\n            0.225\n        ]\n    )\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:01:27.993397Z","iopub.execute_input":"2026-09-07T11:01:27.994002Z","iopub.status.idle":"2026-09-07T11:01:27.998966Z","shell.execute_reply.started":"2026-09-07T11:01:27.993973Z","shell.execute_reply":"2026-09-07T11:01:27.998089Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_test_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n\n    transforms.ToTensor(),\n\n    transforms.Normalize(\n        mean=[\n            0.485,\n            0.456,\n            0.406\n        ],\n        std=[\n            0.229,\n            0.224,\n            0.225\n        ]\n    )\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:01:37.740117Z","iopub.execute_input":"2026-09-07T11:01:37.740948Z","iopub.status.idle":"2026-09-07T11:01:37.744843Z","shell.execute_reply.started":"2026-09-07T11:01:37.740916Z","shell.execute_reply":"2026-09-07T11:01:37.744147Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RetinopathyDataset(Dataset):\n\n    def __init__(\n        self,\n        dataframe,\n        transform=None\n    ):\n\n        self.dataframe = (\n            dataframe\n            .reset_index(drop=True)\n        )\n\n        self.transform = transform\n\n    def __len__(self):\n\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n\n        row = self.dataframe.iloc[index]\n\n        image = Image.open(\n            row[\"image_path\"]\n        ).convert(\"RGB\")\n\n        label = int(\n            row[\"diagnosis\"]\n        )\n\n        if self.transform:\n\n            image = self.transform(\n                image\n            )\n\n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:01:48.281529Z","iopub.execute_input":"2026-09-07T11:01:48.281789Z","iopub.status.idle":"2026-09-07T11:01:48.288198Z","shell.execute_reply.started":"2026-09-07T11:01:48.281759Z","shell.execute_reply":"2026-09-07T11:01:48.287262Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = RetinopathyDataset(\n    train_df,\n    transform=train_transform\n)\n\nval_dataset = RetinopathyDataset(\n    val_df,\n    transform=val_test_transform\n)\n\ntest_dataset = RetinopathyDataset(\n    test_df,\n    transform=val_test_transform\n)\n\nprint(\"Train dataset:\", len(train_dataset))\nprint(\"Validation dataset:\", len(val_dataset))\nprint(\"Test dataset:\", len(test_dataset))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:01:58.298186Z","iopub.execute_input":"2026-09-07T11:01:58.298979Z","iopub.status.idle":"2026-09-07T11:01:58.306465Z","shell.execute_reply.started":"2026-09-07T11:01:58.298946Z","shell.execute_reply":"2026-09-07T11:01:58.305332Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 32\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=2,\n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=2,\n    pin_memory=True\n)\n\nprint(\"DataLoaders ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:02:09.550009Z","iopub.execute_input":"2026-09-07T11:02:09.550268Z","iopub.status.idle":"2026-09-07T11:02:09.555872Z","shell.execute_reply.started":"2026-09-07T11:02:09.550248Z","shell.execute_reply":"2026-09-07T11:02:09.555112Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images, labels = next(\n    iter(train_loader)\n)\n\nprint(\n    \"Images shape:\",\n    images.shape\n)\n\nprint(\n    \"Labels shape:\",\n    labels.shape\n)\n\nprint(\n    \"Labels:\",\n    labels[:10]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:02:21.693607Z","iopub.execute_input":"2026-09-07T11:02:21.693971Z","iopub.status.idle":"2026-09-07T11:02:30.828499Z","shell.execute_reply.started":"2026-09-07T11:02:21.693943Z","shell.execute_reply":"2026-09-07T11:02:30.827477Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"weights = models.DenseNet121_Weights.DEFAULT\n\nmodel = models.densenet121(\n    weights=weights\n)\n\nnum_features = (\n    model.classifier.in_features\n)\n\nmodel.classifier = nn.Sequential(\n    nn.Dropout(0.4),\n    nn.Linear(\n        num_features,\n        5\n    )\n)\n\nmodel = model.to(device)\n\nprint(model.classifier)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:02:46.954028Z","iopub.execute_input":"2026-09-07T11:02:46.954646Z","iopub.status.idle":"2026-09-07T11:02:47.540165Z","shell.execute_reply.started":"2026-09-07T11:02:46.954610Z","shell.execute_reply":"2026-09-07T11:02:47.539574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_counts_train = (\n    train_df[\"diagnosis\"]\n    .value_counts()\n    .sort_index()\n)\n\nclass_weights = (\n    len(train_df)\n    /\n    (\n        len(class_counts_train)\n        *\n        class_counts_train.values\n    )\n)\n\nclass_weights = torch.tensor(\n    class_weights,\n    dtype=torch.float32\n).to(device)\n\nprint(\"Class counts:\")\nprint(class_counts_train)\n\nprint(\"\\nClass weights:\")\nprint(class_weights)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:03:03.444797Z","iopub.execute_input":"2026-09-07T11:03:03.445535Z","iopub.status.idle":"2026-09-07T11:03:03.770562Z","shell.execute_reply.started":"2026-09-07T11:03:03.445498Z","shell.execute_reply":"2026-09-07T11:03:03.769848Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss(\n    weight=class_weights\n)\n\nprint(criterion)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:03:13.947009Z","iopub.execute_input":"2026-09-07T11:03:13.947861Z","iopub.status.idle":"2026-09-07T11:03:13.952511Z","shell.execute_reply.started":"2026-09-07T11:03:13.947828Z","shell.execute_reply":"2026-09-07T11:03:13.951500Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=1e-4,\n    weight_decay=1e-4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:03:32.224865Z","iopub.execute_input":"2026-09-07T11:03:32.225405Z","iopub.status.idle":"2026-09-07T11:03:32.230861Z","shell.execute_reply.started":"2026-09-07T11:03:32.225376Z","shell.execute_reply":"2026-09-07T11:03:32.229950Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer,\n    mode=\"min\",\n    factor=0.5,\n    patience=2\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:03:44.179653Z","iopub.execute_input":"2026-09-07T11:03:44.180319Z","iopub.status.idle":"2026-09-07T11:03:44.184044Z","shell.execute_reply.started":"2026-09-07T11:03:44.180289Z","shell.execute_reply":"2026-09-07T11:03:44.183413Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(\n    model,\n    loader,\n    criterion,\n    optimizer,\n    device\n):\n\n    model.train()\n\n    running_loss = 0.0\n\n    correct = 0\n    total = 0\n\n    for images, labels in loader:\n\n        images = images.to(\n            device,\n            non_blocking=True\n        )\n\n        labels = labels.to(\n            device,\n            non_blocking=True\n        )\n\n        optimizer.zero_grad()\n\n        outputs = model(images)\n\n        loss = criterion(\n            outputs,\n            labels\n        )\n\n        loss.backward()\n\n        optimizer.step()\n\n        running_loss += (\n            loss.item()\n            *\n            images.size(0)\n        )\n\n        predictions = torch.argmax(\n            outputs,\n            dim=1\n        )\n\n        total += labels.size(0)\n\n        correct += (\n            predictions == labels\n        ).sum().item()\n\n    epoch_loss = (\n        running_loss / total\n    )\n\n    epoch_accuracy = (\n        correct / total\n    )\n\n    return (\n        epoch_loss,\n        epoch_accuracy\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:03:54.387832Z","iopub.execute_input":"2026-09-07T11:03:54.388573Z","iopub.status.idle":"2026-09-07T11:03:54.395217Z","shell.execute_reply.started":"2026-09-07T11:03:54.388540Z","shell.execute_reply":"2026-09-07T11:03:54.394500Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(\n    model,\n    loader,\n    criterion,\n    device\n):\n\n    model.eval()\n\n    running_loss = 0.0\n\n    total = 0\n\n    all_predictions = []\n    all_labels = []\n\n    with torch.no_grad():\n\n        for images, labels in loader:\n\n            images = images.to(\n                device,\n                non_blocking=True\n            )\n\n            labels = labels.to(\n                device,\n                non_blocking=True\n            )\n\n            outputs = model(images)\n\n            loss = criterion(\n                outputs,\n                labels\n            )\n\n            running_loss += (\n                loss.item()\n                *\n                images.size(0)\n            )\n\n            total += labels.size(0)\n\n            predictions = torch.argmax(\n                outputs,\n                dim=1\n            )\n\n            all_predictions.extend(\n                predictions\n                .cpu()\n                .numpy()\n            )\n\n            all_labels.extend(\n                labels\n                .cpu()\n                .numpy()\n            )\n\n    val_loss = (\n        running_loss / total\n    )\n\n    all_predictions = np.array(\n        all_predictions\n    )\n\n    all_labels = np.array(\n        all_labels\n    )\n\n    val_accuracy = np.mean(\n        all_predictions == all_labels\n    )\n\n    # ==========================================\n    # DR VS NO DR\n    # ==========================================\n\n    actual_dr = (\n        all_labels != 0\n    )\n\n    predicted_dr = (\n        all_predictions != 0\n    )\n\n    tp = np.sum(\n        actual_dr & predicted_dr\n    )\n\n    fn = np.sum(\n        actual_dr & ~predicted_dr\n    )\n\n    tn = np.sum(\n        ~actual_dr & ~predicted_dr\n    )\n\n    fp = np.sum(\n        ~actual_dr & predicted_dr\n    )\n\n    dr_sensitivity = (\n        tp / (tp + fn)\n        if (tp + fn) > 0\n        else 0\n    )\n\n    no_dr_specificity = (\n        tn / (tn + fp)\n        if (tn + fp) > 0\n        else 0\n    )\n\n    # ==========================================\n    # PER-CLASS RECALL\n    # ==========================================\n\n    class_recalls = []\n\n    for class_id in range(5):\n\n        actual_class = (\n            all_labels == class_id\n        )\n\n        predicted_class = (\n            all_predictions == class_id\n        )\n\n        class_tp = np.sum(\n            actual_class &\n            predicted_class\n        )\n\n        class_total = np.sum(\n            actual_class\n        )\n\n        recall = (\n            class_tp / class_total\n            if class_total > 0\n            else 0\n        )\n\n        class_recalls.append(\n            recall\n        )\n\n    # ==========================================\n    # MACRO F1\n    # ==========================================\n\n    macro_f1 = f1_score(\n        all_labels,\n        all_predictions,\n        average=\"macro\",\n        zero_division=0\n    )\n\n    return (\n        val_loss,\n        val_accuracy,\n        dr_sensitivity,\n        no_dr_specificity,\n        class_recalls,\n        macro_f1,\n        all_labels,\n        all_predictions\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:04:11.770908Z","iopub.execute_input":"2026-09-07T11:04:11.771820Z","iopub.status.idle":"2026-09-07T11:04:11.782403Z","shell.execute_reply.started":"2026-09-07T11:04:11.771787Z","shell.execute_reply":"2026-09-07T11:04:11.781513Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NUM_EPOCHS = 15\nPATIENCE = 3\n\nMODEL_PATH = (\n    \"/kaggle/working/\"\n    \"densenet121_dr_best.pth\"\n)\n\nHISTORY_PATH = (\n    \"/kaggle/working/\"\n    \"densenet121_training_history.json\"\n)\n\nMETRICS_PATH = (\n    \"/kaggle/working/\"\n    \"densenet121_best_metrics.json\"\n)\n\nVAL_RESULTS_PATH = (\n    \"/kaggle/working/\"\n    \"densenet121_best_validation_results.csv\"\n)\n\n\n# ============================================================\n# BEST MODEL TRACKING\n# ============================================================\n\nbest_dr_sensitivity = -1.0\nbest_macro_f1 = -1.0\nbest_val_loss = float(\"inf\")\n\nbest_epoch = 0\n\nepochs_without_improvement = 0\n\n\n# ============================================================\n# TRAINING HISTORY\n# ============================================================\n\ntrain_losses = []\nval_losses = []\n\ntrain_accuracies = []\nval_accuracies = []\n\ndr_sensitivities = []\nno_dr_specificities = []\n\nmacro_f1_scores = []\n\nclass_recalls_history = []\n\n\n# ============================================================\n# TRAINING LOOP\n# ============================================================\n\nfor epoch in range(NUM_EPOCHS):\n\n    print(\"\\n\" + \"=\" * 60)\n    print(\n        f\"Epoch {epoch + 1}/{NUM_EPOCHS}\"\n    )\n    print(\"=\" * 60)\n\n    # --------------------------------------------------------\n    # TRAIN\n    # --------------------------------------------------------\n\n    train_loss, train_acc = train_one_epoch(\n        model,\n        train_loader,\n        criterion,\n        optimizer,\n        device\n    )\n\n    # --------------------------------------------------------\n    # VALIDATION\n    # --------------------------------------------------------\n\n    (\n        val_loss,\n        val_acc,\n        dr_sensitivity,\n        no_dr_specificity,\n        class_recalls,\n        macro_f1,\n        val_labels,\n        val_predictions\n    ) = validate(\n        model,\n        val_loader,\n        criterion,\n        device\n    )\n\n    # --------------------------------------------------------\n    # SCHEDULER\n    # --------------------------------------------------------\n\n    scheduler.step(val_loss)\n\n    # --------------------------------------------------------\n    # STORE HISTORY\n    # --------------------------------------------------------\n\n    train_losses.append(\n        float(train_loss)\n    )\n\n    val_losses.append(\n        float(val_loss)\n    )\n\n    train_accuracies.append(\n        float(train_acc)\n    )\n\n    val_accuracies.append(\n        float(val_acc)\n    )\n\n    dr_sensitivities.append(\n        float(dr_sensitivity)\n    )\n\n    no_dr_specificities.append(\n        float(no_dr_specificity)\n    )\n\n    macro_f1_scores.append(\n        float(macro_f1)\n    )\n\n    class_recalls_history.append(\n        [\n            float(x)\n            for x in class_recalls\n        ]\n    )\n\n    # --------------------------------------------------------\n    # PRINT\n    # --------------------------------------------------------\n\n    print(\n        f\"Train Loss:           {train_loss:.4f}\"\n    )\n\n    print(\n        f\"Train Accuracy:       {train_acc:.4f}\"\n    )\n\n    print(\n        f\"Validation Loss:      {val_loss:.4f}\"\n    )\n\n    print(\n        f\"Validation Accuracy:  {val_acc:.4f}\"\n    )\n\n    print(\n        f\"DR Sensitivity:       {dr_sensitivity:.4f}\"\n    )\n\n    print(\n        f\"No-DR Specificity:    {no_dr_specificity:.4f}\"\n    )\n\n    print(\n        f\"Macro F1:             {macro_f1:.4f}\"\n    )\n\n    print(\"\\nClass Recall:\")\n\n    print(\n        f\"  No DR:             {class_recalls[0]:.4f}\"\n    )\n\n    print(\n        f\"  Mild:              {class_recalls[1]:.4f}\"\n    )\n\n    print(\n        f\"  Moderate:          {class_recalls[2]:.4f}\"\n    )\n\n    print(\n        f\"  Severe:            {class_recalls[3]:.4f}\"\n    )\n\n    print(\n        f\"  Proliferative DR:  {class_recalls[4]:.4f}\"\n    )\n\n    # --------------------------------------------------------\n    # BEST MODEL DECISION\n    #\n    # 1. Highest DR sensitivity\n    # 2. If tied -> highest Macro F1\n    # 3. If tied -> lowest validation loss\n    # --------------------------------------------------------\n\n    is_better = False\n\n    if (\n        dr_sensitivity\n        >\n        best_dr_sensitivity\n    ):\n\n        is_better = True\n\n    elif (\n        dr_sensitivity\n        ==\n        best_dr_sensitivity\n        and\n        macro_f1\n        >\n        best_macro_f1\n    ):\n\n        is_better = True\n\n    elif (\n        dr_sensitivity\n        ==\n        best_dr_sensitivity\n        and\n        macro_f1\n        ==\n        best_macro_f1\n        and\n        val_loss\n        <\n        best_val_loss\n    ):\n\n        is_better = True\n\n    # --------------------------------------------------------\n    # SAVE BEST MODEL\n    # --------------------------------------------------------\n\n    if is_better:\n\n        best_dr_sensitivity = float(\n            dr_sensitivity\n        )\n\n        best_macro_f1 = float(\n            macro_f1\n        )\n\n        best_val_loss = float(\n            val_loss\n        )\n\n        best_epoch = epoch + 1\n\n        epochs_without_improvement = 0\n\n        torch.save(\n            {\n                \"epoch\": best_epoch,\n\n                \"model_state_dict\":\n                    model.state_dict(),\n\n                \"optimizer_state_dict\":\n                    optimizer.state_dict(),\n\n                \"best_dr_sensitivity\":\n                    best_dr_sensitivity,\n\n                \"best_macro_f1\":\n                    best_macro_f1,\n\n                \"best_val_loss\":\n                    best_val_loss,\n\n                \"class_recalls\":\n                    [\n                        float(x)\n                        for x in class_recalls\n                    ]\n            },\n            MODEL_PATH\n        )\n\n        # Save validation predictions\n        best_val_results = pd.DataFrame({\n            \"actual\": val_labels,\n            \"predicted\": val_predictions\n        })\n\n        best_val_results.to_csv(\n            VAL_RESULTS_PATH,\n            index=False\n        )\n\n        print(\"\\n✓ NEW BEST MODEL SAVED\")\n\n        print(\n            f\"  Epoch: {best_epoch}\"\n        )\n\n        print(\n            f\"  DR Sensitivity: \"\n            f\"{best_dr_sensitivity:.4f}\"\n        )\n\n        print(\n            f\"  Macro F1: \"\n            f\"{best_macro_f1:.4f}\"\n        )\n\n        print(\n            f\"  Validation Loss: \"\n            f\"{best_val_loss:.4f}\"\n        )\n\n    else:\n\n        epochs_without_improvement += 1\n\n        print(\n            f\"\\nNo improvement \"\n            f\"({epochs_without_improvement}/\"\n            f\"{PATIENCE})\"\n        )\n\n    # --------------------------------------------------------\n    # SAVE HISTORY AFTER EVERY EPOCH\n    # --------------------------------------------------------\n\n    history = {\n\n        \"train_losses\":\n            train_losses,\n\n        \"val_losses\":\n            val_losses,\n\n        \"train_accuracies\":\n            train_accuracies,\n\n        \"val_accuracies\":\n            val_accuracies,\n\n        \"dr_sensitivities\":\n            dr_sensitivities,\n\n        \"no_dr_specificities\":\n            no_dr_specificities,\n\n        \"macro_f1_scores\":\n            macro_f1_scores,\n\n        \"class_recalls\":\n            class_recalls_history,\n\n        \"best_epoch\":\n            best_epoch,\n\n        \"best_dr_sensitivity\":\n            best_dr_sensitivity,\n\n        \"best_macro_f1\":\n            best_macro_f1,\n\n        \"best_val_loss\":\n            best_val_loss\n    }\n\n    with open(\n        HISTORY_PATH,\n        \"w\"\n    ) as f:\n\n        json.dump(\n            history,\n            f,\n            indent=4\n        )\n\n    # --------------------------------------------------------\n    # EARLY STOPPING\n    # --------------------------------------------------------\n\n    if (\n        epochs_without_improvement\n        >= PATIENCE\n    ):\n\n        print(\n            \"\\nEarly stopping triggered.\"\n        )\n\n        break\n\n\n# ============================================================\n# SAVE BEST METRICS\n# ============================================================\n\nbest_metrics = {\n\n    \"best_epoch\":\n        int(best_epoch),\n\n    \"best_dr_sensitivity\":\n        float(best_dr_sensitivity),\n\n    \"best_macro_f1\":\n        float(best_macro_f1),\n\n    \"best_validation_loss\":\n        float(best_val_loss),\n\n    \"best_validation_accuracy\":\n        float(\n            val_accuracies[\n                best_epoch - 1\n            ]\n        )\n        if best_epoch > 0\n        else None,\n\n    \"best_no_dr_specificity\":\n        float(\n            no_dr_specificities[\n                best_epoch - 1\n            ]\n        )\n        if best_epoch > 0\n        else None,\n\n    \"best_class_recalls\":\n        class_recalls_history[\n            best_epoch - 1\n        ]\n        if best_epoch > 0\n        else None\n}\n\n\nwith open(\n    METRICS_PATH,\n    \"w\"\n) as f:\n\n    json.dump(\n        best_metrics,\n        f,\n        indent=4\n    )\n\n\n# ============================================================\n# VERIFY FILES\n# ============================================================\n\nprint(\"\\n\")\nprint(\"=\" * 60)\nprint(\"TRAINING COMPLETE\")\nprint(\"=\" * 60)\n\nprint(\n    f\"\\nBest Epoch: \"\n    f\"{best_epoch}\"\n)\n\nprint(\n    f\"Best DR Sensitivity: \"\n    f\"{best_dr_sensitivity:.4f}\"\n)\n\nprint(\n    f\"Best Macro F1: \"\n    f\"{best_macro_f1:.4f}\"\n)\n\nprint(\n    f\"Best Validation Loss: \"\n    f\"{best_val_loss:.4f}\"\n)\n\n\nprint(\"\\nSaved files:\")\n\nfor path in [\n    MODEL_PATH,\n    HISTORY_PATH,\n    METRICS_PATH,\n    VAL_RESULTS_PATH\n]:\n\n    exists = os.path.exists(path)\n\n    print(\"\\n\", path)\n    print(\"Exists:\", exists)\n\n    if exists:\n\n        size = (\n            os.path.getsize(path)\n            / (1024 * 1024)\n        )\n\n        print(\n            f\"Size: {size:.2f} MB\"\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:04:41.624322Z","iopub.execute_input":"2026-09-07T11:04:41.625051Z","iopub.status.idle":"2026-09-07T11:48:54.905987Z","shell.execute_reply.started":"2026-09-07T11:04:41.625020Z","shell.execute_reply":"2026-09-07T11:48:54.905188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"checkpoint = torch.load(\n    MODEL_PATH,\n    map_location=device\n)\n\nmodel.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nmodel = model.to(device)\n\nmodel.eval()\n\nprint(\"Best model loaded successfully.\")\n\nprint(\n    \"Best epoch:\",\n    checkpoint[\"epoch\"]\n)\n\nprint(\n    \"Best DR sensitivity:\",\n    checkpoint[\"best_dr_sensitivity\"]\n)\n\nprint(\n    \"Best Macro F1:\",\n    checkpoint[\"best_macro_f1\"]\n)\n\nprint(\n    \"Best validation loss:\",\n    checkpoint[\"best_val_loss\"]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:50:03.186554Z","iopub.execute_input":"2026-09-07T11:50:03.186860Z","iopub.status.idle":"2026-09-07T11:50:03.633784Z","shell.execute_reply.started":"2026-09-07T11:50:03.186831Z","shell.execute_reply":"2026-09-07T11:50:03.632980Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.eval()\n\nall_test_labels = []\nall_test_predictions = []\n\nwith torch.no_grad():\n\n    for images, labels in test_loader:\n\n        images = images.to(\n            device,\n            non_blocking=True\n        )\n\n        labels = labels.to(\n            device,\n            non_blocking=True\n        )\n\n        outputs = model(images)\n\n        predictions = torch.argmax(\n            outputs,\n            dim=1\n        )\n\n        all_test_labels.extend(\n            labels.cpu().numpy()\n        )\n\n        all_test_predictions.extend(\n            predictions.cpu().numpy()\n        )\n\n\nall_test_labels = np.array(\n    all_test_labels\n)\n\nall_test_predictions = np.array(\n    all_test_predictions\n)\n\nprint(\n    \"Test predictions:\",\n    len(all_test_predictions)\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:50:24.923644Z","iopub.execute_input":"2026-09-07T11:50:24.923911Z","iopub.status.idle":"2026-09-07T11:51:03.584124Z","shell.execute_reply.started":"2026-09-07T11:50:24.923889Z","shell.execute_reply":"2026-09-07T11:51:03.583355Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_accuracy = accuracy_score(\n    all_test_labels,\n    all_test_predictions\n)\n\ntest_precision = precision_score(\n    all_test_labels,\n    all_test_predictions,\n    average=\"macro\",\n    zero_division=0\n)\n\ntest_recall = recall_score(\n    all_test_labels,\n    all_test_predictions,\n    average=\"macro\",\n    zero_division=0\n)\n\ntest_f1 = f1_score(\n    all_test_labels,\n    all_test_predictions,\n    average=\"macro\",\n    zero_division=0\n)\n\n\n# ============================================================\n# DR VS NO DR\n# ============================================================\n\nactual_dr = (\n    all_test_labels != 0\n)\n\npredicted_dr = (\n    all_test_predictions != 0\n)\n\ntp = np.sum(\n    actual_dr & predicted_dr\n)\n\nfn = np.sum(\n    actual_dr & ~predicted_dr\n)\n\ntn = np.sum(\n    ~actual_dr & ~predicted_dr\n)\n\nfp = np.sum(\n    ~actual_dr & predicted_dr\n)\n\n\ntest_dr_sensitivity = (\n    tp / (tp + fn)\n    if (tp + fn) > 0\n    else 0\n)\n\ntest_no_dr_specificity = (\n    tn / (tn + fp)\n    if (tn + fp) > 0\n    else 0\n)\n\n\nprint(\"=\" * 60)\nprint(\"FINAL DENSENET121 TEST RESULTS\")\nprint(\"=\" * 60)\n\nprint(\n    f\"\\nAccuracy:             \"\n    f\"{test_accuracy:.4f}\"\n)\n\nprint(\n    f\"Macro Precision:      \"\n    f\"{test_precision:.4f}\"\n)\n\nprint(\n    f\"Macro Recall:         \"\n    f\"{test_recall:.4f}\"\n)\n\nprint(\n    f\"Macro F1:             \"\n    f\"{test_f1:.4f}\"\n)\n\nprint(\n    f\"\\nDR Sensitivity:       \"\n    f\"{test_dr_sensitivity:.4f}\"\n)\n\nprint(\n    f\"No-DR Specificity:    \"\n    f\"{test_no_dr_specificity:.4f}\"\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:51:36.599731Z","iopub.execute_input":"2026-09-07T11:51:36.600605Z","iopub.status.idle":"2026-09-07T11:51:36.617808Z","shell.execute_reply.started":"2026-09-07T11:51:36.600570Z","shell.execute_reply":"2026-09-07T11:51:36.616980Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"report = classification_report(\n    all_test_labels,\n    all_test_predictions,\n    target_names=class_names,\n    zero_division=0\n)\n\nprint(report)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:51:58.466604Z","iopub.execute_input":"2026-09-07T11:51:58.467562Z","iopub.status.idle":"2026-09-07T11:51:58.486460Z","shell.execute_reply.started":"2026-09-07T11:51:58.467528Z","shell.execute_reply":"2026-09-07T11:51:58.485703Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cm = confusion_matrix(\n    all_test_labels,\n    all_test_predictions\n)\n\nprint(\"Confusion Matrix:\")\nprint(cm)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:52:22.214822Z","iopub.execute_input":"2026-09-07T11:52:22.215735Z","iopub.status.idle":"2026-09-07T11:52:22.227269Z","shell.execute_reply.started":"2026-09-07T11:52:22.215702Z","shell.execute_reply":"2026-09-07T11:52:22.226610Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(8, 6))\n\nplt.imshow(cm)\n\nplt.title(\n    \"DenseNet121 Confusion Matrix\"\n)\n\nplt.xlabel(\"Predicted Label\")\nplt.ylabel(\"True Label\")\n\nplt.xticks(\n    range(5),\n    class_names,\n    rotation=30\n)\n\nplt.yticks(\n    range(5),\n    class_names\n)\n\nplt.colorbar()\n\nfor i in range(5):\n\n    for j in range(5):\n\n        plt.text(\n            j,\n            i,\n            cm[i, j],\n            ha=\"center\",\n            va=\"center\"\n        )\n\nplt.tight_layout()\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:52:42.889932Z","iopub.execute_input":"2026-09-07T11:52:42.890789Z","iopub.status.idle":"2026-09-07T11:52:43.100915Z","shell.execute_reply.started":"2026-09-07T11:52:42.890749Z","shell.execute_reply":"2026-09-07T11:52:43.100078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TEST_METRICS_PATH = (\n    \"/kaggle/working/\"\n    \"densenet121_test_metrics.json\"\n)\n\ntest_metrics = {\n\n    \"accuracy\":\n        float(test_accuracy),\n\n    \"macro_precision\":\n        float(test_precision),\n\n    \"macro_recall\":\n        float(test_recall),\n\n    \"macro_f1\":\n        float(test_f1),\n\n    \"dr_sensitivity\":\n        float(test_dr_sensitivity),\n\n    \"no_dr_specificity\":\n        float(test_no_dr_specificity),\n\n    \"confusion_matrix\":\n        cm.tolist(),\n\n    \"best_epoch\":\n        int(checkpoint[\"epoch\"]),\n\n    \"best_validation_dr_sensitivity\":\n        float(\n            checkpoint[\n                \"best_dr_sensitivity\"\n            ]\n        ),\n\n    \"best_validation_macro_f1\":\n        float(\n            checkpoint[\n                \"best_macro_f1\"\n            ]\n        ),\n\n    \"best_validation_loss\":\n        float(\n            checkpoint[\n                \"best_val_loss\"\n            ]\n        )\n}\n\n\nwith open(\n    TEST_METRICS_PATH,\n    \"w\"\n) as f:\n\n    json.dump(\n        test_metrics,\n        f,\n        indent=4\n    )\n\nprint(\n    \"Saved:\",\n    TEST_METRICS_PATH\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:53:01.857841Z","iopub.execute_input":"2026-09-07T11:53:01.858541Z","iopub.status.idle":"2026-09-07T11:53:01.865477Z","shell.execute_reply.started":"2026-09-07T11:53:01.858510Z","shell.execute_reply":"2026-09-07T11:53:01.864638Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"TEST_PREDICTIONS_PATH = (\n    \"/kaggle/working/\"\n    \"densenet121_test_predictions.csv\"\n)\n\ntest_predictions_df = pd.DataFrame({\n\n    \"actual\":\n        all_test_labels,\n\n    \"predicted\":\n        all_test_predictions\n})\n\ntest_predictions_df.to_csv(\n    TEST_PREDICTIONS_PATH,\n    index=False\n)\n\nprint(\n    \"Saved:\",\n    TEST_PREDICTIONS_PATH\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:53:16.431737Z","iopub.execute_input":"2026-09-07T11:53:16.432453Z","iopub.status.idle":"2026-09-07T11:53:16.440167Z","shell.execute_reply.started":"2026-09-07T11:53:16.432408Z","shell.execute_reply":"2026-09-07T11:53:16.439249Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"\\n\")\nprint(\"=\" * 60)\nprint(\"FINAL FILE CHECK\")\nprint(\"=\" * 60)\n\nfiles_to_check = [\n\n    MODEL_PATH,\n\n    HISTORY_PATH,\n\n    METRICS_PATH,\n\n    VAL_RESULTS_PATH,\n\n    TEST_METRICS_PATH,\n\n    TEST_PREDICTIONS_PATH\n]\n\n\nfor path in files_to_check:\n\n    exists = os.path.exists(path)\n\n    print(\"\\n\", path)\n\n    print(\n        \"Exists:\",\n        exists\n    )\n\n    if exists:\n\n        size_mb = (\n            os.path.getsize(path)\n            / (1024 * 1024)\n        )\n\n        print(\n            f\"Size: {size_mb:.2f} MB\"\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-07T11:53:31.155627Z","iopub.execute_input":"2026-09-07T11:53:31.156339Z","iopub.status.idle":"2026-09-07T11:53:31.162585Z","shell.execute_reply.started":"2026-09-07T11:53:31.156308Z","shell.execute_reply":"2026-09-07T11:53:31.161686Z"}},"outputs":[],"execution_count":null}]}