{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431},{"sourceType":"modelInstanceVersion","sourceId":785339,"databundleVersionId":16065457,"modelInstanceId":599249,"modelId":611506}],"dockerImageVersionId":31286,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 📊 Models Development & Performance Analysis\n\nIn this notebook we experimented with two architectures: **ResNet50** and **EfficientNet-B4** to solve the **APTOS Diabetic Retinopathy Classification** problem.\n\nThe evaluation metric used is **Quadratic Weighted Kappa (QWK)**, which is the official competition metric and is well suited for **ordinal classification problems** where prediction distance matters.\n\n---\n\n### 1️⃣ ResNet50 Baseline\n\nThe first experiment used **ResNet50 pretrained on ImageNet** with minimal modifications.\n\n**Configuration**\n\n- Resolution: `224 × 224`\n- Transfer learning\n- Standard classifier layer\n- No strong augmentation\n\n**Results**\n\nEpoch 10/10  \nTrain Loss: 0.1205 | Train Acc: 0.9713 | Train QWK: 0.9823  \nVal   Loss: 1.2603 | Val   Acc: 0.7926 | Val   QWK: 0.8563  \n\n**Observation**\n\nThe model clearly **overfits the training data**.  \nThe large gap between training and validation metrics indicates poor generalization.\n\n---\n\n### 2️⃣ ResNet50 + Strong Augmentation + Dropout\n\nTo reduce overfitting, stronger regularization was introduced.\n\n**Changes**\n\n- Strong **Albumentations data augmentation**\n- Added **Dropout layer** to the classifier\n- Improved regularization\n\n**Results**\n\nEpoch 10/10  \nTrain Loss: 0.7125 | Train Acc: 0.7945 | Train QWK: 0.8567  \nVal   Loss: 0.8896 | Val   Acc: 0.7872 | Val   QWK: 0.8693  \n\n**Observation**\n\n- Overfitting reduced significantly\n- Training and validation metrics became closer\n- Validation performance improved\n\nThis configuration became the **best ResNet setup**.\n\n---\n\n### 3️⃣ Increasing Image Resolution\n\nTo capture more fine-grained retinal features, image resolution was increased.\n\n**Changes**\n\n- Resolution increased to `320 × 320`\n- Same augmentation pipeline\n\n**Results**\n\nEpoch 6/15  \nTrain Loss: 0.7713 | Train Acc: 0.7685 | Train QWK: 0.8662  \nVal   Loss: 0.8569 | Val   Acc: 0.7708 | Val   QWK: 0.8541  \n\nEarly stopping triggered.\n\n**Observation**\n\nHigher resolution **did not improve ResNet performance**.\n\nThe best ResNet configuration remained:\n\n`224 resolution + strong augmentation + dropout`\n\n---\n\n### 4️⃣ EfficientNet-B4 Baseline\n\nNext we experimented with a stronger architecture: **EfficientNet-B4**.\n\n**Configuration**\n\n- Resolution: `380 × 380`\n- Transfer learning\n- Only classifier trained initially\n\n**Results**\n\nEpoch 12/12  \nTrain Loss: 0.9135 | Train Acc: 0.7231 | Train QWK: 0.8014  \nVal   Loss: 0.9587 | Val   Acc: 0.7544 | Val   QWK: 0.8432  \n\n**Observation**\n\nInitial performance was lower than ResNet because the backbone had not yet been properly fine-tuned.\n\n---\n\n### 5️⃣ EfficientNet Fine-Tuning\n\nTo improve performance, deeper layers of EfficientNet were unfrozen.\n\n**Changes**\n\n- Reduced learning rate to `5e-5`\n- Unfroze **3 EfficientNet blocks**\n- Continued training with **CosineAnnealingLR scheduler**\n\n**Results**\n\nEpoch 20/20  \nTrain Loss: 0.7962 | Train Acc: 0.7545 | Train QWK: 0.8460  \nVal   Loss: 0.9002 | Val   Acc: 0.7735 | Val   QWK: 0.8724  \n\n**Observation**\n\nAfter proper fine-tuning, **EfficientNet slightly outperformed ResNet**.\n\n---\n\n### 6️⃣ Threshold Optimization\n\nInstead of using `argmax` directly, predictions were converted into **continuous severity scores**, and **threshold optimization** was applied to maximize QWK.\n\n**Result**\n\nQWK after threshold optimization ≈ **0.8905**\n\nThis step significantly improved performance **without retraining the model**.\n\n---\n\n### 7️⃣ Test Time Augmentation (TTA)\n\nTo further stabilize predictions, **Test Time Augmentation** was applied.\n\n**Augmentations used**\n\n- Original image\n- Horizontal flip\n- Vertical flip\n\n**Result**\n\nQWK with TTA ≈ **0.8918**\n\nThis produced a small but consistent improvement.\n\n---\n\n### 8️⃣ Model Ensemble\n\nFinally, predictions from **ResNet50** and **EfficientNet-B4** were combined.\n\nInstead of averaging class labels, we averaged **continuous severity scores** from both models.\n\n**Result**\n\nFinal Ensemble QWK ≈ **0.8989**\n\n---\n\n### 🏆 Final Pipeline\n\nThe final prediction pipeline is:\n\nImage  \n↓  \nTest Time Augmentation (TTA)  \n↓  \nResNet50 Prediction  \nEfficientNet-B4 Prediction  \n↓  \nAverage Severity Scores  \n↓  \nOptimized Thresholds  \n↓  \nFinal DR Class Prediction  \n\n---\n\n### 📈 Final Performance Summary\n\n| Model Stage | QWK |\n|-------------|------|\nResNet Baseline | 0.856 |\nResNet + Augmentation + Dropout | 0.869 |\nEfficientNet Fine-Tuned | 0.872 |\nThreshold Optimization | 0.890 |\nTest Time Augmentation | 0.892 |\n**ResNet + EfficientNet Ensemble** | **0.8989** |\n\n---\n\nThis workflow demonstrates how **progressive experimentation, strong regularization, fine-tuning, threshold optimization, test-time augmentation, and ensembling** can significantly improve model performance for medical image classification tasks.","metadata":{}},{"cell_type":"code","source":"import os\nimport random\nimport time\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport cv2\nfrom PIL import Image\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, f1_score, confusion_matrix, classification_report\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:01.618977Z","iopub.execute_input":"2026-03-13T21:00:01.619804Z","iopub.status.idle":"2026-03-13T21:00:06.625408Z","shell.execute_reply.started":"2026-03-13T21:00:01.619765Z","shell.execute_reply":"2026-03-13T21:00:06.624532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nset_seed(42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T20:58:53.195655Z","iopub.status.idle":"2026-03-13T20:58:53.196084Z","shell.execute_reply.started":"2026-03-13T20:58:53.195877Z","shell.execute_reply":"2026-03-13T20:58:53.195901Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🔍 Exploratory Data Analysis (EDA)","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/competitions/aptos2019-blindness-detection/train.csv\")\n\nprint(train_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:21.344872Z","iopub.execute_input":"2026-03-13T21:00:21.345512Z","iopub.status.idle":"2026-03-13T21:00:21.374543Z","shell.execute_reply.started":"2026-03-13T21:00:21.345480Z","shell.execute_reply":"2026-03-13T21:00:21.373718Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_df['diagnosis'].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:23.042002Z","iopub.execute_input":"2026-03-13T21:00:23.042332Z","iopub.status.idle":"2026-03-13T21:00:23.053795Z","shell.execute_reply.started":"2026-03-13T21:00:23.042305Z","shell.execute_reply":"2026-03-13T21:00:23.053032Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The dataset is highly imbalanced, with class 0 dominating and severe cases underrepresented.","metadata":{}},{"cell_type":"code","source":"# path to images\nIMAGE_DIR = \"/kaggle/input/competitions/aptos2019-blindness-detection/train_images\"\n\nclasses = sorted(train_df['diagnosis'].unique())\n\nfig, axes = plt.subplots(len(classes), 3, figsize=(12, 15))\n\nfor row, cls in enumerate(classes):\n    \n    # get 3 samples from this class\n    samples = train_df[train_df['diagnosis'] == cls].sample(3)\n    \n    for col, (_, sample) in enumerate(samples.iterrows()):\n        \n        img_path = os.path.join(IMAGE_DIR, sample['id_code'] + \".png\")\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        axes[row, col].imshow(img)\n        axes[row, col].set_title(f\"Class {cls}\")\n        axes[row, col].axis(\"off\")\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:24.760438Z","iopub.execute_input":"2026-03-13T21:00:24.761035Z","iopub.status.idle":"2026-03-13T21:00:32.596284Z","shell.execute_reply.started":"2026-03-13T21:00:24.761003Z","shell.execute_reply":"2026-03-13T21:00:32.595378Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sizes = []\n\nfor img_id in train_df['id_code'].sample(100):\n    path = os.path.join(IMAGE_DIR, img_id + \".png\")\n    img = Image.open(path)\n    sizes.append(img.size)\n\nprint(set(sizes))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:32.598311Z","iopub.execute_input":"2026-03-13T21:00:32.598718Z","iopub.status.idle":"2026-03-13T21:00:34.390319Z","shell.execute_reply.started":"2026-03-13T21:00:32.598678Z","shell.execute_reply":"2026-03-13T21:00:34.389496Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Sampled images reveal highly variable resolutions, indicating the need for resizing to a consistent input size before training.","metadata":{}},{"cell_type":"code","source":"train_df[\"image_path\"] = train_df[\"id_code\"].apply(\n    lambda x: os.path.join(IMAGE_DIR, x + \".png\")\n)\n\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:34.417351Z","iopub.execute_input":"2026-03-13T21:00:34.418013Z","iopub.status.idle":"2026-03-13T21:00:34.431992Z","shell.execute_reply.started":"2026-03-13T21:00:34.417984Z","shell.execute_reply":"2026-03-13T21:00:34.431018Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df, val_df = train_test_split(\n    train_df,\n    test_size=0.2,\n    stratify=train_df[\"diagnosis\"],\n    random_state=42\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:34.433046Z","iopub.execute_input":"2026-03-13T21:00:34.434044Z","iopub.status.idle":"2026-03-13T21:00:34.451189Z","shell.execute_reply.started":"2026-03-13T21:00:34.434013Z","shell.execute_reply":"2026-03-13T21:00:34.450254Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The dataset is split into training and validation sets using stratified sampling to preserve the original class distribution.","metadata":{}},{"cell_type":"code","source":"print(train_df[\"diagnosis\"].value_counts())\nprint(val_df[\"diagnosis\"].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:34.452274Z","iopub.execute_input":"2026-03-13T21:00:34.452590Z","iopub.status.idle":"2026-03-13T21:00:34.459349Z","shell.execute_reply.started":"2026-03-13T21:00:34.452564Z","shell.execute_reply":"2026-03-13T21:00:34.458555Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install albumentations","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T20:58:54.140658Z","iopub.execute_input":"2026-03-13T20:58:54.141077Z","iopub.status.idle":"2026-03-13T20:58:58.691952Z","shell.execute_reply.started":"2026-03-13T20:58:54.141031Z","shell.execute_reply":"2026-03-13T20:58:58.690924Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T20:59:04.524180Z","iopub.execute_input":"2026-03-13T20:59:04.524519Z","iopub.status.idle":"2026-03-13T20:59:10.192785Z","shell.execute_reply.started":"2026-03-13T20:59:04.524485Z","shell.execute_reply":"2026-03-13T20:59:10.192147Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🧪 Data Augmentation Strategy\n\nTo improve generalization and reduce overfitting observed in earlier experiments, we applied a **strong augmentation pipeline** using Albumentations.\n\n### Why Strong Augmentation?\n\nIn initial experiments with weak or minimal augmentation, the model achieved very high training performance but failed to generalize well to validation data (clear overfitting).\n\nTo address this, we increased augmentation strength to:\n- Expose the model to more diverse variations\n- Improve robustness to real-world retinal image differences\n- Reduce reliance on memorizing training samples\n\n---\n\n### Applied Transformations\n\n- **Resize (380 × 380)**  \n  Ensures a consistent input size while preserving sufficient detail for EfficientNet.\n\n- **Horizontal & Vertical Flip**  \n  Retinal images are orientation-invariant → safe and effective augmentation.\n\n- **Shift / Scale / Rotate**  \n  Simulates variations in camera positioning and zoom.\n\n- **Brightness & Contrast Adjustment**  \n  Mimics different lighting conditions during image acquisition.\n\n- **Hue / Saturation / Value **  \n  Accounts for color variability across devices and patients.\n\n- **Blur / Sharpen (OneOf)**  \n  Improves robustness to image quality differences.\n\n- **Normalization (ImageNet stats)**  \n  Aligns input distribution with pretrained backbone expectations.\n\n- **Tensor Conversion**  \n  Prepares data for PyTorch models.\n\n---\n\n### Impact\n\n- Reduced overfitting significantly  \n- Improved validation stability  \n- Became a key factor in boosting QWK performance  \n\n---\n\n**Key Insight:**  \nStrong augmentation is essential in medical imaging, where datasets are limited and variability in real-world conditions is high.","metadata":{}},{"cell_type":"code","source":"train_transforms = A.Compose([\n    \n    A.Resize(380,380), #(224,224) ,(320,320)\n\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n\n    A.ShiftScaleRotate(\n        shift_limit=0.2,\n        scale_limit=0.2,\n        rotate_limit=180,\n        border_mode=0,\n        p=0.7\n    ),\n\n    A.RandomBrightnessContrast(\n        brightness_limit=0.2,\n        contrast_limit=0.2,\n        p=0.7\n    ),\n\n    A.HueSaturationValue(\n        hue_shift_limit=10,\n        sat_shift_limit=20,\n        val_shift_limit=20,\n        p=0.5\n    ),\n\n    A.OneOf([\n        A.GaussianBlur(blur_limit=3),\n        A.Sharpen(),\n    ], p=0.3),\n\n    #A.Normalize(),\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406),\n        std=(0.229, 0.224, 0.225)\n    ),\n    ToTensorV2()\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:41.239184Z","iopub.execute_input":"2026-03-13T21:00:41.239503Z","iopub.status.idle":"2026-03-13T21:00:41.251921Z","shell.execute_reply.started":"2026-03-13T21:00:41.239476Z","shell.execute_reply":"2026-03-13T21:00:41.250957Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_transforms = A.Compose([\n    A.Resize(380,380), #(224,224) ,(320,320)\n    #A.Normalize(),\n    A.Normalize(\n        mean=(0.485, 0.456, 0.406),\n        std=(0.229, 0.224, 0.225)\n    ),\n    ToTensorV2()\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:43.771557Z","iopub.execute_input":"2026-03-13T21:00:43.771977Z","iopub.status.idle":"2026-03-13T21:00:43.779521Z","shell.execute_reply.started":"2026-03-13T21:00:43.771942Z","shell.execute_reply":"2026-03-13T21:00:43.778748Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class APTOSDataset(Dataset):\n    \n    def __init__(self, dataframe, transform=None):\n        self.df = dataframe\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        \n        img_path = self.df.iloc[idx][\"image_path\"]\n        label = self.df.iloc[idx][\"diagnosis\"]\n        \n        #image = Image.open(img_path).convert(\"RGB\")\n        image = cv2.imread(img_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        if self.transform:\n            image = self.transform(image=image)[\"image\"]\n        \n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:45.407256Z","iopub.execute_input":"2026-03-13T21:00:45.407921Z","iopub.status.idle":"2026-03-13T21:00:45.414157Z","shell.execute_reply.started":"2026-03-13T21:00:45.407889Z","shell.execute_reply":"2026-03-13T21:00:45.413105Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = APTOSDataset(train_df, transform=train_transforms)\nval_dataset = APTOSDataset(val_df, transform=val_transforms)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:00:47.575637Z","iopub.execute_input":"2026-03-13T21:00:47.575995Z","iopub.status.idle":"2026-03-13T21:00:47.580779Z","shell.execute_reply.started":"2026-03-13T21:00:47.575967Z","shell.execute_reply":"2026-03-13T21:00:47.579832Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_loader = DataLoader(\n    train_dataset,\n    batch_size=32, #16\n    shuffle=True,\n    num_workers=2\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=32, #16\n    shuffle=False,\n    num_workers=2\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:01:00.344349Z","iopub.execute_input":"2026-03-13T21:01:00.344693Z","iopub.status.idle":"2026-03-13T21:01:00.350559Z","shell.execute_reply.started":"2026-03-13T21:01:00.344668Z","shell.execute_reply":"2026-03-13T21:01:00.349665Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### DataLoaders are used to efficiently load data in batches, with shuffling applied during training to improve generalization and disabled during validation for consistent evaluation.","metadata":{}},{"cell_type":"code","source":"images, labels = next(iter(train_loader))\n\nprint(images.shape)\nprint(labels.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:01:02.817206Z","iopub.execute_input":"2026-03-13T21:01:02.817494Z","iopub.status.idle":"2026-03-13T21:01:11.833461Z","shell.execute_reply.started":"2026-03-13T21:01:02.817471Z","shell.execute_reply":"2026-03-13T21:01:11.832452Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### ResNet","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:02:11.920941Z","iopub.execute_input":"2026-03-13T21:02:11.921912Z","iopub.status.idle":"2026-03-13T21:02:12.181139Z","shell.execute_reply.started":"2026-03-13T21:02:11.921867Z","shell.execute_reply":"2026-03-13T21:02:12.180177Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🧬 ResNet50 ","metadata":{}},{"cell_type":"code","source":"import torchvision.models as models\n\nmodel = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1)\n\n# replace final layer\nnum_features = model.fc.in_features\n#model.fc = nn.Linear(num_features, 5)\nmodel.fc = nn.Sequential(\n    nn.Dropout(0.4),\n    nn.Linear(num_features, 5)\n)\n\nmodel = model.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T19:05:32.982739Z","iopub.execute_input":"2026-03-07T19:05:32.983255Z","iopub.status.idle":"2026-03-07T19:05:34.163520Z","shell.execute_reply.started":"2026-03-07T19:05:32.983226Z","shell.execute_reply":"2026-03-07T19:05:34.162830Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### A pretrained ResNet50 model is used (transfer learning), with the final layer replaced by a custom classifier that includes **Dropout** for regularization and a **Linear layer** for 5-class prediction.","metadata":{}},{"cell_type":"markdown","source":"### Freeze backbone (Stage 1)","metadata":{}},{"cell_type":"code","source":"for param in model.parameters():\n    param.requires_grad = False\n\nfor param in model.fc.parameters():\n    param.requires_grad = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T19:05:42.495579Z","iopub.execute_input":"2026-03-07T19:05:42.496651Z","iopub.status.idle":"2026-03-07T19:05:42.500752Z","shell.execute_reply.started":"2026-03-07T19:05:42.496611Z","shell.execute_reply":"2026-03-07T19:05:42.500094Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### All backbone layers are frozen while training only the final classifier (**feature extraction**) to leverage pretrained knowledge and reduce overfitting.","metadata":{}},{"cell_type":"markdown","source":"### Compute class weights","metadata":{}},{"cell_type":"code","source":"from sklearn.utils.class_weight import compute_class_weight\n\nclass_weights = compute_class_weight(\n    class_weight=\"balanced\",\n    classes=np.unique(train_df[\"diagnosis\"]),\n    y=train_df[\"diagnosis\"]\n)\n\nclass_weights = torch.tensor(class_weights, dtype=torch.float).to(device)\n\nprint(class_weights)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:02:50.111383Z","iopub.execute_input":"2026-03-13T21:02:50.111711Z","iopub.status.idle":"2026-03-13T21:02:50.334450Z","shell.execute_reply.started":"2026-03-13T21:02:50.111685Z","shell.execute_reply":"2026-03-13T21:02:50.333524Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Class weights are computed using a balanced strategy to handle class imbalance, giving higher importance to underrepresented classes during training.","metadata":{}},{"cell_type":"code","source":"train_df['diagnosis'].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T19:05:51.457513Z","iopub.execute_input":"2026-03-07T19:05:51.457937Z","iopub.status.idle":"2026-03-07T19:05:51.466949Z","shell.execute_reply.started":"2026-03-07T19:05:51.457900Z","shell.execute_reply":"2026-03-07T19:05:51.465787Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss(weight=class_weights)\n\noptimizer = optim.AdamW(      \n    model.fc.parameters(),\n    lr=1e-3,   # lr for head\n    weight_decay=1e-4\n)   ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T19:05:55.449964Z","iopub.execute_input":"2026-03-07T19:05:55.450641Z","iopub.status.idle":"2026-03-07T19:05:55.456268Z","shell.execute_reply.started":"2026-03-07T19:05:55.450604Z","shell.execute_reply":"2026-03-07T19:05:55.455532Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### CrossEntropyLoss with **class weights** is used to address imbalance, while the **AdamW optimizer** updates only the classifier with weight decay for better generalization.","metadata":{}},{"cell_type":"markdown","source":"### train head","metadata":{}},{"cell_type":"code","source":"EPOCHS = 5\n\ntrain_losses = []\nval_losses = []\ntrain_accs = []\nval_accs = []\n\nfor epoch in range(EPOCHS):\n\n    # ======================\n    # TRAIN\n    # ======================\n    model.train()\n\n    train_loss = 0\n    train_preds = []\n    train_targets = []\n\n    for images, labels in train_loader:\n\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(images)\n\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item()\n\n        preds = torch.argmax(outputs, dim=1)\n\n        train_preds.extend(preds.cpu().numpy())\n        train_targets.extend(labels.cpu().numpy())\n\n    train_loss /= len(train_loader)\n    train_acc = accuracy_score(train_targets, train_preds)\n\n    # ======================\n    # VALIDATION\n    # ======================\n    model.eval()\n\n    val_loss = 0\n    val_preds = []\n    val_targets = []\n\n    with torch.no_grad():\n\n        for images, labels in val_loader:\n\n            images = images.to(device)\n            labels = labels.to(device)\n\n            outputs = model(images)\n\n            loss = criterion(outputs, labels)\n\n            val_loss += loss.item()\n\n            preds = torch.argmax(outputs, dim=1)\n\n            val_preds.extend(preds.cpu().numpy())\n            val_targets.extend(labels.cpu().numpy())\n\n    val_loss /= len(val_loader)\n    val_acc = accuracy_score(val_targets, val_preds)\n\n    # STORE METRICS\n\n    train_losses.append(train_loss)\n    val_losses.append(val_loss)\n\n    train_accs.append(train_acc)\n    val_accs.append(val_acc)\n\n    print(f\"Epoch {epoch+1}/{EPOCHS}\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}\")\n    print(f\"Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f}\")\n    print(\"-\"*40)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T19:06:06.915677Z","iopub.execute_input":"2026-03-07T19:06:06.916005Z","iopub.status.idle":"2026-03-07T19:25:10.334380Z","shell.execute_reply.started":"2026-03-07T19:06:06.915977Z","shell.execute_reply":"2026-03-07T19:25:10.333341Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The model is trained for multiple epochs using a standard training loop, where only the **classifier head** is updated (frozen backbone), and validation is performed without gradients to evaluate generalization.","metadata":{}},{"cell_type":"markdown","source":"### Unfreeze the backbone","metadata":{}},{"cell_type":"code","source":"print(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T14:22:16.953730Z","iopub.execute_input":"2026-03-07T14:22:16.954100Z","iopub.status.idle":"2026-03-07T14:22:16.959856Z","shell.execute_reply.started":"2026-03-07T14:22:16.954060Z","shell.execute_reply":"2026-03-07T14:22:16.959196Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### train layer 4 + head","metadata":{}},{"cell_type":"code","source":"for param in model.parameters():\n    param.requires_grad = False\n\nfor param in model.layer4.parameters():\n    param.requires_grad = True\n\nfor param in model.fc.parameters():\n    param.requires_grad = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T19:27:27.590623Z","iopub.execute_input":"2026-03-07T19:27:27.591582Z","iopub.status.idle":"2026-03-07T19:27:27.597218Z","shell.execute_reply.started":"2026-03-07T19:27:27.591542Z","shell.execute_reply":"2026-03-07T19:27:27.596491Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The backbone is partially unfrozen by training the final block (**layer4**) along with the **classifier head**, enabling deeper fine-tuning while keeping earlier layers frozen.","metadata":{}},{"cell_type":"markdown","source":"### New optimizer","metadata":{}},{"cell_type":"code","source":"optimizer = optim.AdamW(\n    filter(lambda p: p.requires_grad, model.parameters()),\n    lr=1e-4,\n    weight_decay=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=10\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T19:27:32.724203Z","iopub.execute_input":"2026-03-07T19:27:32.725055Z","iopub.status.idle":"2026-03-07T19:27:32.730446Z","shell.execute_reply.started":"2026-03-07T19:27:32.725026Z","shell.execute_reply":"2026-03-07T19:27:32.729746Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The optimizer updates only trainable layers using **AdamW**, while a **CosineAnnealingLR scheduler** gradually reduces the learning rate to improve fine-tuning stability.","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import cohen_kappa_score\n\ndef compute_qwk(y_true, y_pred):\n    return cohen_kappa_score(y_true, y_pred, weights=\"quadratic\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:41:22.705229Z","iopub.execute_input":"2026-03-13T21:41:22.705568Z","iopub.status.idle":"2026-03-13T21:41:22.710419Z","shell.execute_reply.started":"2026-03-13T21:41:22.705537Z","shell.execute_reply":"2026-03-13T21:41:22.709738Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 15\n\nbest_qwk = 0\npatience = 4\ncounter = 0\n\ntrain_losses = []\nval_losses = []\n\ntrain_accs = []\nval_accs = []\n\ntrain_qwks = []\nval_qwks = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T19:36:21.012487Z","iopub.execute_input":"2026-03-07T19:36:21.013578Z","iopub.status.idle":"2026-03-07T19:36:21.018579Z","shell.execute_reply.started":"2026-03-07T19:36:21.013533Z","shell.execute_reply":"2026-03-07T19:36:21.017545Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for epoch in range(EPOCHS):\n\n    # ======================\n    # TRAIN\n    # ======================\n    model.train()\n\n    train_loss = 0\n    train_preds = []\n    train_targets = []\n\n    for images, labels in train_loader:\n\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(images)\n\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item()\n\n        preds = torch.argmax(outputs, dim=1)\n\n        train_preds.extend(preds.cpu().numpy())\n        train_targets.extend(labels.cpu().numpy())\n\n    train_loss /= len(train_loader)\n\n    train_acc = accuracy_score(train_targets, train_preds)\n    train_qwk = compute_qwk(train_targets, train_preds)\n\n    # ======================\n    # VALIDATION\n    # ======================\n    \n    model.eval()\n\n    val_loss = 0\n    val_preds = []\n    val_targets = []\n\n    with torch.no_grad():\n\n        for images, labels in val_loader:\n\n            images = images.to(device)\n            labels = labels.to(device)\n\n            outputs = model(images)\n\n            loss = criterion(outputs, labels)\n\n            val_loss += loss.item()\n\n            preds = torch.argmax(outputs, dim=1)\n\n            val_preds.extend(preds.cpu().numpy())\n            val_targets.extend(labels.cpu().numpy())\n\n    val_loss /= len(val_loader)\n\n    val_acc = accuracy_score(val_targets, val_preds)\n    val_qwk = compute_qwk(val_targets, val_preds)\n\n    scheduler.step()\n\n    # ======================\n    # STORE METRICS\n    # ======================\n\n    train_losses.append(train_loss)\n    val_losses.append(val_loss)\n\n    train_accs.append(train_acc)\n    val_accs.append(val_acc)\n\n    train_qwks.append(train_qwk)\n    val_qwks.append(val_qwk)\n\n    # ======================\n    # SAVE BEST MODEL\n    # ======================\n\n    if val_qwk > best_qwk:\n        best_qwk = val_qwk\n        torch.save(model.state_dict(), \"best_model.pth\")\n        counter = 0\n    else:\n        counter += 1\n\n    # ======================\n    # PRINT\n    # ======================\n\n    print(f\"Epoch {epoch+1}/{EPOCHS}\")\n\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train QWK: {train_qwk:.4f}\")\n\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val   QWK: {val_qwk:.4f}\")\n\n    print(\"-\"*50)\n\n    # ======================\n    # EARLY STOPPING\n    # ======================\n\n    if counter >= patience:\n        print(\"Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T19:36:26.555473Z","iopub.execute_input":"2026-03-07T19:36:26.555815Z","iopub.status.idle":"2026-03-07T19:59:12.494108Z","shell.execute_reply.started":"2026-03-07T19:36:26.555786Z","shell.execute_reply":"2026-03-07T19:59:12.493106Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The model is trained with **fine-tuning**, tracking loss, accuracy, and **QWK** for both training and validation, while using a learning rate scheduler, saving the best model based on validation QWK, and applying **early stopping** to prevent overfitting.","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\n\ncm = confusion_matrix(val_targets, val_preds)\n\nclass_names = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative\"]\n\nplt.figure(figsize=(7,6))\nsns.heatmap(cm,\n            annot=True,\n            fmt=\"d\",\n            cmap=\"Blues\",\n            xticklabels=class_names,\n            yticklabels=class_names)\n\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"True\")\nplt.title(\"Confusion Matrix\")\n\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T20:11:16.162726Z","iopub.execute_input":"2026-03-07T20:11:16.163071Z","iopub.status.idle":"2026-03-07T20:11:16.366710Z","shell.execute_reply.started":"2026-03-07T20:11:16.163044Z","shell.execute_reply":"2026-03-07T20:11:16.366113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report\n\nprint(classification_report(val_targets, val_preds))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T20:11:51.236870Z","iopub.execute_input":"2026-03-07T20:11:51.237627Z","iopub.status.idle":"2026-03-07T20:11:51.250988Z","shell.execute_reply.started":"2026-03-07T20:11:51.237598Z","shell.execute_reply":"2026-03-07T20:11:51.250344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Prediction distribution:\")\nprint(np.bincount(val_preds))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T20:12:16.709055Z","iopub.execute_input":"2026-03-07T20:12:16.709711Z","iopub.status.idle":"2026-03-07T20:12:16.714540Z","shell.execute_reply.started":"2026-03-07T20:12:16.709683Z","shell.execute_reply":"2026-03-07T20:12:16.713621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"True distribution:\")\nprint(np.bincount(val_targets))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-07T20:12:25.131854Z","iopub.execute_input":"2026-03-07T20:12:25.132454Z","iopub.status.idle":"2026-03-07T20:12:25.137076Z","shell.execute_reply.started":"2026-03-07T20:12:25.132427Z","shell.execute_reply":"2026-03-07T20:12:25.136175Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The model achieves strong performance on the majority class (0), while showing weaker precision and recall on minority classes, highlighting the impact of class imbalance despite overall good accuracy (~77%).","metadata":{}},{"cell_type":"markdown","source":"## ⚡ EfficientNet-B4","metadata":{}},{"cell_type":"code","source":"import torchvision.models as models\n\nmodel = models.efficientnet_b4(weights=models.EfficientNet_B4_Weights.IMAGENET1K_V1)\n\nnum_features = model.classifier[1].in_features\n\nmodel.classifier = nn.Sequential(\n    nn.Dropout(0.4),\n    nn.Linear(num_features, 5)\n)\n\nmodel = model.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:02:20.993152Z","iopub.execute_input":"2026-03-13T21:02:20.993485Z","iopub.status.idle":"2026-03-13T21:02:21.765184Z","shell.execute_reply.started":"2026-03-13T21:02:20.993459Z","shell.execute_reply":"2026-03-13T21:02:21.764502Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Freeze backbone","metadata":{}},{"cell_type":"code","source":"for param in model.parameters():\n    param.requires_grad = False\n\nfor param in model.classifier.parameters():\n    param.requires_grad = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:02:26.868619Z","iopub.execute_input":"2026-03-13T21:02:26.868984Z","iopub.status.idle":"2026-03-13T21:02:26.875817Z","shell.execute_reply.started":"2026-03-13T21:02:26.868955Z","shell.execute_reply":"2026-03-13T21:02:26.874885Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss(weight=class_weights)\n\noptimizer = optim.AdamW(      \n    model.classifier.parameters(),\n    lr=1e-3,   # lr for head\n    weight_decay=1e-4\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:03:21.263345Z","iopub.execute_input":"2026-03-13T21:03:21.263649Z","iopub.status.idle":"2026-03-13T21:03:21.268718Z","shell.execute_reply.started":"2026-03-13T21:03:21.263626Z","shell.execute_reply":"2026-03-13T21:03:21.267812Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### train head","metadata":{}},{"cell_type":"code","source":"EPOCHS = 5\n\ntrain_losses = []\nval_losses = []\ntrain_accs = []\nval_accs = []\n\nfor epoch in range(EPOCHS):\n\n    # ======================\n    # TRAIN\n    # ======================\n    model.train()\n\n    train_loss = 0\n    train_preds = []\n    train_targets = []\n\n    for images, labels in train_loader:\n\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad()\n\n        outputs = model(images)\n\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item()\n\n        preds = torch.argmax(outputs, dim=1)\n\n        train_preds.extend(preds.cpu().numpy())\n        train_targets.extend(labels.cpu().numpy())\n\n    train_loss /= len(train_loader)\n    train_acc = accuracy_score(train_targets, train_preds)\n\n    # ======================\n    # VALIDATION\n    # ======================\n    model.eval()\n\n    val_loss = 0\n    val_preds = []\n    val_targets = []\n\n    with torch.no_grad():\n\n        for images, labels in val_loader:\n\n            images = images.to(device)\n            labels = labels.to(device)\n\n            outputs = model(images)\n\n            loss = criterion(outputs, labels)\n\n            val_loss += loss.item()\n\n            preds = torch.argmax(outputs, dim=1)\n\n            val_preds.extend(preds.cpu().numpy())\n            val_targets.extend(labels.cpu().numpy())\n\n    val_loss /= len(val_loader)\n    val_acc = accuracy_score(val_targets, val_preds)\n\n    # STORE METRICS\n\n    train_losses.append(train_loss)\n    val_losses.append(val_loss)\n\n    train_accs.append(train_acc)\n    val_accs.append(val_acc)\n\n    print(f\"Epoch {epoch+1}/{EPOCHS}\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}\")\n    print(f\"Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.4f}\")\n    print(\"-\"*40)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:03:25.294216Z","iopub.execute_input":"2026-03-13T21:03:25.294540Z","iopub.status.idle":"2026-03-13T21:22:38.657452Z","shell.execute_reply.started":"2026-03-13T21:03:25.294514Z","shell.execute_reply":"2026-03-13T21:22:38.656504Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Unfreeze last 4 blocks + head.","metadata":{}},{"cell_type":"code","source":"for param in model.parameters():\n    param.requires_grad = False\n\nfor param in model.features[-4:].parameters():\n    param.requires_grad = True\n\nfor param in model.classifier.parameters():\n    param.requires_grad = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:23:32.845724Z","iopub.execute_input":"2026-03-13T21:23:32.846463Z","iopub.status.idle":"2026-03-13T21:23:32.853071Z","shell.execute_reply.started":"2026-03-13T21:23:32.846431Z","shell.execute_reply":"2026-03-13T21:23:32.852363Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The model is partially fine-tuned by unfreezing the last **4 blocks** of the EfficientNet feature extractor along with the **classifier head**, while keeping earlier layers frozen.","metadata":{}},{"cell_type":"code","source":"optimizer = optim.AdamW(\n    filter(lambda p: p.requires_grad, model.parameters()),\n    lr=5e-5,\n    weight_decay=1e-4\n)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=30 #10\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:23:45.676191Z","iopub.execute_input":"2026-03-13T21:23:45.676467Z","iopub.status.idle":"2026-03-13T21:23:45.682885Z","shell.execute_reply.started":"2026-03-13T21:23:45.676447Z","shell.execute_reply":"2026-03-13T21:23:45.682232Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The optimizer fine-tunes only unfrozen layers using **AdamW** with a lower learning rate, while a **CosineAnnealingLR scheduler** smoothly adjusts the learning rate for better convergence.","metadata":{}},{"cell_type":"code","source":"best_qwk = 0\npatience = 6\ncounter = 0","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:24:14.060396Z","iopub.execute_input":"2026-03-13T21:24:14.061003Z","iopub.status.idle":"2026-03-13T21:24:14.064800Z","shell.execute_reply.started":"2026-03-13T21:24:14.060972Z","shell.execute_reply":"2026-03-13T21:24:14.063937Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"EPOCHS = 15\n\ntrain_losses = []\nval_losses = []\n\ntrain_accs = []\nval_accs = []\n\ntrain_qwks = []\nval_qwks = []\n\nfor epoch in range(EPOCHS):\n\n    # ======================\n    # TRAIN\n    # ======================\n    model.train()\n\n    train_loss = 0\n    train_preds = []\n    train_targets = []\n\n    for images, labels in train_loader:\n\n        images = images.to(device)\n        labels = labels.to(device)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        outputs = model(images)\n\n        loss = criterion(outputs, labels)\n\n        loss.backward()\n        optimizer.step()\n\n        train_loss += loss.item()\n\n        preds = torch.argmax(outputs, dim=1)\n\n        train_preds.extend(preds.cpu().numpy())\n        train_targets.extend(labels.cpu().numpy())\n\n    train_loss /= len(train_loader)\n\n    train_acc = accuracy_score(train_targets, train_preds)\n    train_qwk = compute_qwk(train_targets, train_preds)\n\n    # ======================\n    # VALIDATION\n    # ======================\n\n    model.eval()\n\n    val_loss = 0\n    val_preds = []\n    val_targets = []\n\n    with torch.no_grad():\n\n        for images, labels in val_loader:\n\n            images = images.to(device)\n            labels = labels.to(device)\n\n            outputs = model(images)\n\n            loss = criterion(outputs, labels)\n\n            val_loss += loss.item()\n\n            preds = torch.argmax(outputs, dim=1)\n\n            val_preds.extend(preds.cpu().numpy())\n            val_targets.extend(labels.cpu().numpy())\n\n    val_loss /= len(val_loader)\n\n    val_acc = accuracy_score(val_targets, val_preds)\n    val_qwk = compute_qwk(val_targets, val_preds)\n\n    # scheduler step\n    scheduler.step()\n\n    # ======================\n    # STORE METRICS\n    # ======================\n\n    train_losses.append(train_loss)\n    val_losses.append(val_loss)\n\n    train_accs.append(train_acc)\n    val_accs.append(val_acc)\n\n    train_qwks.append(train_qwk)\n    val_qwks.append(val_qwk)\n\n    # ======================\n    # SAVE BEST MODEL\n    # ======================\n\n    if val_qwk > best_qwk:\n        best_qwk = val_qwk\n        torch.save(model.state_dict(), \"/kaggle/working/best_effb4.pth\")\n        counter = 0\n    else:\n        counter += 1\n\n    # ======================\n    # PRINT\n    # ======================\n\n    print(f\"Epoch {epoch+1}/{EPOCHS}\")\n\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f} | Train QWK: {train_qwk:.4f}\")\n\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val   QWK: {val_qwk:.4f}\")\n\n    print(\"-\"*50)\n\n    # ======================\n    # EARLY STOPPING\n    # ======================\n\n    if counter >= patience:\n        print(\"Early stopping triggered\")\n        break","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T21:41:51.247954Z","iopub.execute_input":"2026-03-13T21:41:51.248707Z","iopub.status.idle":"2026-03-13T22:26:21.999665Z","shell.execute_reply.started":"2026-03-13T21:41:51.248677Z","shell.execute_reply":"2026-03-13T22:26:21.998675Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ⚖️ Threshold Optimization","metadata":{}},{"cell_type":"code","source":"val_preds = []\nval_targets = []\n\nmodel.eval()\n\nwith torch.no_grad():\n    \n    for images, labels in val_loader:\n\n        images = images.to(device)\n        labels = labels.to(device)\n\n        outputs = model(images)\n\n        probs = torch.softmax(outputs, dim=1)\n\n        score = (probs * torch.arange(5, device=probs.device)).sum(dim=1)\n\n        val_preds.extend(score.cpu().numpy())\n        val_targets.extend(labels.cpu().numpy())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T22:41:49.134012Z","iopub.execute_input":"2026-03-13T22:41:49.134349Z","iopub.status.idle":"2026-03-13T22:42:30.619836Z","shell.execute_reply.started":"2026-03-13T22:41:49.134323Z","shell.execute_reply":"2026-03-13T22:42:30.619011Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom scipy.optimize import minimize\nfrom sklearn.metrics import cohen_kappa_score\n\n\ndef apply_thresholds(preds, thresholds):\n\n    thresholds = sorted(thresholds)\n\n    return np.digitize(preds, thresholds)\n\n\ndef qwk_loss(thresholds, preds, targets):\n\n    preds_class = apply_thresholds(preds, thresholds)\n\n    return -cohen_kappa_score(targets, preds_class, weights=\"quadratic\")\n\n\ndef optimize_thresholds(preds, targets):\n\n    preds = np.array(preds)\n    targets = np.array(targets)\n\n    initial_thresholds = [0.5, 1.5, 2.5, 3.5]\n\n    result = minimize(\n        qwk_loss,\n        initial_thresholds,\n        args=(preds, targets),\n        method=\"nelder-mead\"\n    )\n\n    return result.x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T22:42:35.463824Z","iopub.execute_input":"2026-03-13T22:42:35.464196Z","iopub.status.idle":"2026-03-13T22:42:35.471019Z","shell.execute_reply.started":"2026-03-13T22:42:35.464163Z","shell.execute_reply":"2026-03-13T22:42:35.470265Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_preds_continuous = np.array(val_preds)\nval_targets_array = np.array(val_targets)\n\nbest_thresholds = optimize_thresholds(\n    val_preds_continuous,\n    val_targets_array\n)\n\nprint(\"Optimized thresholds:\", best_thresholds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T22:42:49.276214Z","iopub.execute_input":"2026-03-13T22:42:49.276533Z","iopub.status.idle":"2026-03-13T22:42:49.356649Z","shell.execute_reply.started":"2026-03-13T22:42:49.276507Z","shell.execute_reply":"2026-03-13T22:42:49.355811Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_preds = apply_thresholds(\n    val_preds_continuous,\n    best_thresholds\n)\n\nqwk = cohen_kappa_score(\n    val_targets_array,\n    final_preds,\n    weights=\"quadratic\"\n)\n\nprint(\"QWK after optimization:\", qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T22:43:03.748147Z","iopub.execute_input":"2026-03-13T22:43:03.748469Z","iopub.status.idle":"2026-03-13T22:43:03.755579Z","shell.execute_reply.started":"2026-03-13T22:43:03.748444Z","shell.execute_reply":"2026-03-13T22:43:03.754790Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Instead of direct class predictions, model outputs are treated as continuous severity scores and optimal thresholds are learned using optimization to maximize **QWK**, significantly improving performance without retraining.","metadata":{}},{"cell_type":"code","source":"torch.save(model.state_dict(), \"/kaggle/working/efficientnet_b4_best.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T22:50:03.839437Z","iopub.execute_input":"2026-03-13T22:50:03.839768Z","iopub.status.idle":"2026-03-13T22:50:03.979146Z","shell.execute_reply.started":"2026-03-13T22:50:03.839742Z","shell.execute_reply":"2026-03-13T22:50:03.978480Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 🔄 Test Time Augmentation","metadata":{}},{"cell_type":"code","source":"def predict_with_tta(model, loader, device):\n\n    model.eval()\n\n    all_preds = []\n\n    with torch.no_grad():\n\n        for images, _ in loader:\n\n            images = images.to(device)\n\n            # original\n            logits1 = model(images)\n\n            # horizontal flip\n            logits2 = model(torch.flip(images, dims=[3]))\n\n            # vertical flip\n            logits3 = model(torch.flip(images, dims=[2]))\n\n            # average logits\n            logits = (logits1 + logits2 + logits3) / 3\n\n            probs = torch.softmax(logits, dim=1)\n\n            score = (probs * torch.arange(5, device=device)).sum(dim=1)\n\n            all_preds.extend(score.cpu().numpy())\n\n    return np.array(all_preds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T22:52:31.676679Z","iopub.execute_input":"2026-03-13T22:52:31.677340Z","iopub.status.idle":"2026-03-13T22:52:31.683195Z","shell.execute_reply.started":"2026-03-13T22:52:31.677310Z","shell.execute_reply":"2026-03-13T22:52:31.682579Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_preds_tta = predict_with_tta(model, val_loader, device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T22:52:37.525389Z","iopub.execute_input":"2026-03-13T22:52:37.525984Z","iopub.status.idle":"2026-03-13T22:53:24.451903Z","shell.execute_reply.started":"2026-03-13T22:52:37.525957Z","shell.execute_reply":"2026-03-13T22:53:24.451115Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_preds = apply_thresholds(val_preds_tta, best_thresholds)\n\nqwk = cohen_kappa_score(val_targets_array, final_preds, weights=\"quadratic\")\n\nprint(\"QWK with TTA:\", qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T22:53:38.022598Z","iopub.execute_input":"2026-03-13T22:53:38.023454Z","iopub.status.idle":"2026-03-13T22:53:38.031642Z","shell.execute_reply.started":"2026-03-13T22:53:38.023408Z","shell.execute_reply":"2026-03-13T22:53:38.030921Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Test Time Augmentation (TTA) is applied by averaging predictions from original, horizontally flipped, and vertically flipped images, producing more stable outputs and improving QWK.","metadata":{}},{"cell_type":"markdown","source":"## 🔬 ensemble ResNet + EfficientNet","metadata":{}},{"cell_type":"code","source":"import torchvision.models as models\nimport torch.nn as nn\nimport torch\n\nresnet = models.resnet50(weights=None) # to load trained weights, not ImageNet weights.\n\nnum_features = resnet.fc.in_features\n\nresnet.fc = nn.Sequential(\n    nn.Dropout(0.4),\n    nn.Linear(num_features, 5)\n)\n\nresnet = resnet.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T23:04:06.030171Z","iopub.execute_input":"2026-03-13T23:04:06.030507Z","iopub.status.idle":"2026-03-13T23:04:06.404069Z","shell.execute_reply.started":"2026-03-13T23:04:06.030482Z","shell.execute_reply":"2026-03-13T23:04:06.403276Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"resnet.load_state_dict(\n    torch.load(\"/kaggle/input/models/vabdoamr/resnet/pytorch/default/1/best_model(ResNet).pth\")\n)\n\nresnet.eval()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T23:05:18.962177Z","iopub.execute_input":"2026-03-13T23:05:18.962726Z","iopub.status.idle":"2026-03-13T23:05:20.942767Z","shell.execute_reply.started":"2026-03-13T23:05:18.962695Z","shell.execute_reply":"2026-03-13T23:05:20.942042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_model_scores(model, images, device):\n\n    logits = model(images)\n\n    probs = torch.softmax(logits, dim=1)\n\n    scores = (probs * torch.arange(5, device=device)).sum(dim=1)\n\n    return scores","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T23:08:41.809261Z","iopub.execute_input":"2026-03-13T23:08:41.809547Z","iopub.status.idle":"2026-03-13T23:08:41.813758Z","shell.execute_reply.started":"2026-03-13T23:08:41.809527Z","shell.execute_reply":"2026-03-13T23:08:41.813148Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def ensemble_predict_tta(resnet, effnet, loader, device):\n\n    resnet.eval()\n    effnet.eval()\n\n    all_scores = []\n    all_targets = []\n\n    with torch.no_grad():\n\n        for images, labels in loader:\n\n            images = images.to(device)\n\n            # ORIGINAL\n            score1_r = get_model_scores(resnet, images, device)\n            score1_e = get_model_scores(effnet, images, device)\n\n            # HORIZONTAL FLIP\n            images_h = torch.flip(images, dims=[3])\n\n            score2_r = get_model_scores(resnet, images_h, device)\n            score2_e = get_model_scores(effnet, images_h, device)\n\n            # VERTICAL FLIP\n            images_v = torch.flip(images, dims=[2])\n\n            score3_r = get_model_scores(resnet, images_v, device)\n            score3_e = get_model_scores(effnet, images_v, device)\n\n            # AVERAGE TTA FOR EACH MODEL\n            resnet_score = (score1_r + score2_r + score3_r) / 3\n            effnet_score = (score1_e + score2_e + score3_e) / 3\n\n            # ENSEMBLE MODELS\n            final_score = (resnet_score + effnet_score) / 2\n\n            all_scores.extend(final_score.cpu().numpy())\n            all_targets.extend(labels.numpy())\n\n    return np.array(all_scores), np.array(all_targets)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T23:09:06.459631Z","iopub.execute_input":"2026-03-13T23:09:06.460349Z","iopub.status.idle":"2026-03-13T23:09:06.466390Z","shell.execute_reply.started":"2026-03-13T23:09:06.460320Z","shell.execute_reply":"2026-03-13T23:09:06.465760Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Predictions from ResNet50 and EfficientNet-B4 are combined using both **TTA** and **model ensembling**, where each model’s augmented predictions are averaged first, then merged via score averaging to produce more robust final predictions.","metadata":{}},{"cell_type":"code","source":"scores, targets = ensemble_predict_tta(\n    resnet,\n    model,      # EfficientNet model\n    val_loader,\n    device\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T23:09:53.792005Z","iopub.execute_input":"2026-03-13T23:09:53.792587Z","iopub.status.idle":"2026-03-13T23:10:42.079677Z","shell.execute_reply.started":"2026-03-13T23:09:53.792559Z","shell.execute_reply":"2026-03-13T23:10:42.078838Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_thresholds = optimize_thresholds(scores, targets)\n\nprint(\"Optimized thresholds:\", best_thresholds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T23:10:51.389875Z","iopub.execute_input":"2026-03-13T23:10:51.390232Z","iopub.status.idle":"2026-03-13T23:10:51.477546Z","shell.execute_reply.started":"2026-03-13T23:10:51.390198Z","shell.execute_reply":"2026-03-13T23:10:51.476912Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_preds = apply_thresholds(scores, best_thresholds)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T23:10:56.484373Z","iopub.execute_input":"2026-03-13T23:10:56.484667Z","iopub.status.idle":"2026-03-13T23:10:56.488920Z","shell.execute_reply.started":"2026-03-13T23:10:56.484643Z","shell.execute_reply":"2026-03-13T23:10:56.488152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import cohen_kappa_score\n\nqwk = cohen_kappa_score(\n    targets,\n    final_preds,\n    weights=\"quadratic\"\n)\n\nprint(\"Ensemble QWK:\", qwk)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-13T23:11:00.411117Z","iopub.execute_input":"2026-03-13T23:11:00.411847Z","iopub.status.idle":"2026-03-13T23:11:00.417817Z","shell.execute_reply.started":"2026-03-13T23:11:00.411820Z","shell.execute_reply":"2026-03-13T23:11:00.417257Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### The final ensemble achieves a strong **QWK of 0.8989**, demonstrating the effectiveness of combining models and prediction strategies.","metadata":{}}]}