{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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"},"accelerator":"GPU","kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =============================================\n# Kaggle Setup - APTOS 2019 Blindness Detection\n# =============================================\n# Dataset: Add \"aptos2019-blindness-detection\" as a Kaggle dataset\n# Go to: Add Data -> Competition Data -> aptos2019-blindness-detection\n# The data will be at: /kaggle/input/aptos2019-blindness-detection/\n\nimport os\n\nDATA_DIR = '/kaggle/input/competitions/aptos2019-blindness-detection'\nSAVE_DIR = '/kaggle/working/'\n\nprint(f\"Data directory: {DATA_DIR}\")\nprint(f\"Output directory: {SAVE_DIR}\")\nprint(f\"Data files: {os.listdir(DATA_DIR)}\")","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:19.550087Z","iopub.execute_input":"2026-04-09T14:05:19.550791Z","iopub.status.idle":"2026-04-09T14:05:19.559887Z","shell.execute_reply.started":"2026-04-09T14:05:19.550758Z","shell.execute_reply":"2026-04-09T14:05:19.559020Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    print(dirname)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:19.561106Z","iopub.execute_input":"2026-04-09T14:05:19.561457Z","iopub.status.idle":"2026-04-09T14:05:25.701774Z","shell.execute_reply.started":"2026-04-09T14:05:19.561432Z","shell.execute_reply":"2026-04-09T14:05:25.701020Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q timm\n\nimport os\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\nimport random\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix, classification_report, cohen_kappa_score, f1_score, precision_score, recall_score\nfrom sklearn.utils.class_weight import compute_class_weight\n\nimport timm\nfrom tqdm import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# Set seeds for reproducibility\ndef 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    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(42)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:25.702637Z","iopub.execute_input":"2026-04-09T14:05:25.702840Z","iopub.status.idle":"2026-04-09T14:05:44.174076Z","shell.execute_reply.started":"2026-04-09T14:05:25.702819Z","shell.execute_reply":"2026-04-09T14:05:44.173107Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(os.path.join(DATA_DIR, \"train.csv\"))\nprint(f\"Total samples: {len(df)}\")\nprint(f\"Class distribution:\\n{df['diagnosis'].value_counts().sort_index()}\")\n\ntrain_df, temp_df = train_test_split(\n    df, test_size=0.2, stratify=df[\"diagnosis\"], random_state=42\n)\n\nval_df, test_df = train_test_split(\n    temp_df, test_size=0.5, stratify=temp_df[\"diagnosis\"], random_state=42\n)\n\nprint(f\"\\nTrain: {len(train_df)}, Val: {len(val_df)}, Test: {len(test_df)}\")","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:44.176139Z","iopub.execute_input":"2026-04-09T14:05:44.176700Z","iopub.status.idle":"2026-04-09T14:05:44.220926Z","shell.execute_reply.started":"2026-04-09T14:05:44.176669Z","shell.execute_reply":"2026-04-09T14:05:44.220199Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_clahe(pil_img):\n    img = np.array(pil_img)\n    img = cv2.cvtColor(img, cv2.COLOR_RGB2LAB)\n\n    l, a, b = cv2.split(img)\n    clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8,8))\n    cl = clahe.apply(l)\n\n    merged = cv2.merge((cl,a,b))\n    img = cv2.cvtColor(merged, cv2.COLOR_LAB2RGB)\n\n    return Image.fromarray(img)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:44.221831Z","iopub.execute_input":"2026-04-09T14:05:44.222222Z","iopub.status.idle":"2026-04-09T14:05:44.227396Z","shell.execute_reply.started":"2026-04-09T14:05:44.222187Z","shell.execute_reply":"2026-04-09T14:05:44.226711Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ImageNet normalization - CRITICAL for pretrained models\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD  = [0.229, 0.224, 0.225]\nIMG_SIZE = 456\n\ntrain_transform = transforms.Compose([\n    transforms.Lambda(apply_clahe),\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(10),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\nval_transform = transforms.Compose([\n    transforms.Lambda(apply_clahe),\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    # CHANGE: Removed validation augmentation (RandomHorizontalFlip)\n    # REASON: Validation must reflect real, unseen data without augmentation\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\nprint(\"Transforms configured with ImageNet normalization (no val augmentation)\")","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:44.228355Z","iopub.execute_input":"2026-04-09T14:05:44.228548Z","iopub.status.idle":"2026-04-09T14:05:44.242971Z","shell.execute_reply.started":"2026-04-09T14:05:44.228528Z","shell.execute_reply":"2026-04-09T14:05:44.242177Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class APTOSDataset(Dataset):\n    def __init__(self, df, data_dir, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_name = self.df.iloc[idx][\"id_code\"]\n        label = self.df.iloc[idx][\"diagnosis\"]\n\n        path = os.path.join(self.data_dir, \"train_images\", f\"{img_name}.png\")\n        image = Image.open(path).convert(\"RGB\")\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n\n\ndef mixup_data(x, y, alpha=0.4):\n    \"\"\"Mixup augmentation: blends pairs of images and their labels.\"\"\"\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1.0\n\n    batch_size = x.size(0)\n    index = torch.randperm(batch_size).to(x.device)\n\n    mixed_x = lam * x + (1 - lam) * x[index]\n    y_a, y_b = y, y[index]\n    return mixed_x, y_a, y_b, lam\n\n\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    \"\"\"Mixup loss: weighted combination of losses for both label sets.\"\"\"\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:44.244564Z","iopub.execute_input":"2026-04-09T14:05:44.244863Z","iopub.status.idle":"2026-04-09T14:05:44.257374Z","shell.execute_reply.started":"2026-04-09T14:05:44.244840Z","shell.execute_reply":"2026-04-09T14:05:44.256712Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 8\nNUM_WORKERS = 2\n\n# SIMPLIFIED: Use simple shuffle instead of complex stratified sampling\n# Class weights in loss function handle imbalance\ntrain_loader = DataLoader(\n    APTOSDataset(train_df, DATA_DIR, train_transform),\n    batch_size=BATCH_SIZE, shuffle=True,  # Simple random shuffle\n    num_workers=NUM_WORKERS, pin_memory=True, drop_last=True\n)\nval_loader = DataLoader(\n    APTOSDataset(val_df, DATA_DIR, val_transform),\n    batch_size=BATCH_SIZE,\n    num_workers=NUM_WORKERS, pin_memory=True\n)\ntest_loader = DataLoader(\n    APTOSDataset(test_df, DATA_DIR, val_transform),\n    batch_size=BATCH_SIZE,\n    num_workers=NUM_WORKERS, pin_memory=True\n)\n\nprint(f\"Train batches: {len(train_loader)}, Val batches: {len(val_loader)}, Test batches: {len(test_loader)}\")","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:44.258243Z","iopub.execute_input":"2026-04-09T14:05:44.258568Z","iopub.status.idle":"2026-04-09T14:05:44.274218Z","shell.execute_reply.started":"2026-04-09T14:05:44.258539Z","shell.execute_reply":"2026-04-09T14:05:44.273518Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# EfficientNet-B5 with pretrained ImageNet weights\n# =============================================\n\nmodel = timm.create_model(\"efficientnet_b5\", pretrained=True, num_classes=5,\n                          drop_rate=0.4, drop_path_rate=0.2).to(device)\n\n# Replace classifier with balanced regularization\n# CHANGE: Reduced dropout from (0.5 + 0.3) to just 0.3\n# WHY: Double dropout was over-regularizing, preventing the classifier from learning\n# EXPECTED: Faster convergence, higher accuracy\nin_features = model.classifier.in_features\n# CHANGE: Use single dropout layer (0.3) instead of dual (0.4 + 0.3)\n# REASON: EfficientNet has internal regularization; dual dropout causes underfitting\n# EXPECTED: Higher accuracy + better generalization (no over-regularization)\nmodel.classifier = nn.Sequential(\n    nn.Dropout(p=0.3),  # Single, balanced dropout\n    nn.Linear(in_features, 256),\n    nn.ReLU(inplace=True),\n    nn.BatchNorm1d(256),\n    nn.Linear(256, 5)\n).to(device)\n\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(f\"Total params: {total_params:,}\")\nprint(f\"Trainable params: {trainable_params:,}\")","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:44.275047Z","iopub.execute_input":"2026-04-09T14:05:44.275362Z","iopub.status.idle":"2026-04-09T14:05:47.357904Z","shell.execute_reply.started":"2026-04-09T14:05:44.275318Z","shell.execute_reply":"2026-04-09T14:05:47.357214Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_weights = compute_class_weight(\n    \"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)\nprint(f\"Class weights: {class_weights}\")","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:47.360020Z","iopub.execute_input":"2026-04-09T14:05:47.360323Z","iopub.status.idle":"2026-04-09T14:05:47.627697Z","shell.execute_reply.started":"2026-04-09T14:05:47.360287Z","shell.execute_reply":"2026-04-09T14:05:47.626821Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def freeze_backbone(model):\n    \"\"\"Freeze all layers except the classifier head.\"\"\"\n    for name, param in model.named_parameters():\n        if 'classifier' not in name:\n            param.requires_grad = False\n        else:\n            param.requires_grad = True\n\ndef unfreeze_all(model):\n    \"\"\"Unfreeze all layers for fine-tuning.\"\"\"\n    for param in model.parameters():\n        param.requires_grad = True\n\nprint(\"Helper functions defined: freeze_backbone, unfreeze_all\")","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:47.628781Z","iopub.execute_input":"2026-04-09T14:05:47.629503Z","iopub.status.idle":"2026-04-09T14:05:47.634647Z","shell.execute_reply.started":"2026-04-09T14:05:47.629473Z","shell.execute_reply":"2026-04-09T14:05:47.633996Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(model, train_loader, val_loader, epochs=30, use_mixup=True, mixup_alpha=0.2):\n    \"\"\"\n    BALANCED TRAINING APPROACH:\n    - Simple CrossEntropyLoss with class weights (no Focal Loss)\n    - Weak Mixup for early epochs only (alpha=0.2)\n    - Single dropout (0.3) - no over-regularization\n    - Backbone LR = 1e-4 (moderate fine-tuning)\n    - Simple CosineAnnealingLR scheduler\n    - Early stopping with patience=7\n    \"\"\"\n    \n    # CHANGE: Use only CrossEntropyLoss (no Focal Loss)\n    # REASON: Dataset not extremely imbalanced; class weights handle imbalance\n    criterion = nn.CrossEntropyLoss(weight=class_weights)\n\n    # ==============================\n    # STAGE 1: Warmup (classifier only)\n    # ==============================\n    freeze_backbone(model)\n    warmup_params = [p for p in model.parameters() if p.requires_grad]\n    print(f\"Stage 1 - Trainable params (classifier only): {sum(p.numel() for p in warmup_params):,}\")\n\n    warmup_optimizer = torch.optim.Adam(warmup_params, lr=1e-3, weight_decay=1e-4)\n\n    print(\"\\n\" + \"=\"*60)\n    print(\"STAGE 1: Warmup - Training classifier head (3 epochs)\")\n    print(\"=\"*60)\n\n    for epoch in range(3):\n        model.train()\n        loop = tqdm(train_loader, desc=f\"Warmup Epoch {epoch+1}/3\")\n        for images, labels in loop:\n            images, labels = images.to(device), labels.to(device)\n            warmup_optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(warmup_params, max_norm=1.0)\n            warmup_optimizer.step()\n            loop.set_postfix(loss=loss.item())\n\n    # ==============================\n    # STAGE 2: Full Fine-tuning with Balanced Approach\n    # ==============================\n    unfreeze_all(model)\n\n    all_params = list(model.parameters())\n    print(f\"\\nStage 2 - All params unfrozen: {sum(p.numel() for p in all_params):,}\")\n\n    # CHANGE: Backbone LR = 1e-4 (moderate, not conservative 1e-5)\n    # REASON: Allows better adaptation of pretrained features to medical images\n    backbone_params = [p for n, p in model.named_parameters() if 'classifier' not in n]\n    head_params = [p for n, p in model.named_parameters() if 'classifier' in n]\n\n    optimizer = torch.optim.Adam([\n        {'params': backbone_params, 'lr': 1e-4},  # CHANGED: 1e-4 (not 1e-5) for faster adaptation\n        {'params': head_params, 'lr': 1e-4}\n    ], weight_decay=1e-4)\n\n    # CHANGE: Simple CosineAnnealingLR (not custom multi-stage scheduler)\n    # REASON: Simpler, more stable, easier to tune\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, T_max=epochs, eta_min=1e-7\n    )\n\n    train_losses, val_losses = [], []\n    train_accs, val_accs = [], []\n    val_f1s = []  # Track F1 scores for better minority class monitoring\n\n    best_val_acc = 0\n    patience = 7  # REVERTED: Back to 7 (balanced, not aggressive 5)\n    patience_counter = 0\n\n    save_path = os.path.join(SAVE_DIR, 'eff_best.pth')\n\n    print(\"\\n\" + \"=\"*60)\n    print(f\"STAGE 2: Full Fine-tuning ({epochs} epochs, patience={patience})\")\n    print(f\"Loss: CrossEntropyLoss with class weights (no Focal Loss)\")\n    print(f\"Scheduler: CosineAnnealingLR\")\n    print(f\"Backbone LR: 1e-4 | Head LR: 1e-4\")\n    print(\"=\"*60 + \"\\n\")\n\n    for epoch in range(epochs):\n\n        # ===== TRAIN =====\n        model.train()\n        running_loss, correct, total = 0, 0, 0\n\n        loop = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs}\")\n\n        for images, labels in loop:\n            images, labels = images.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            \n            # KEEP: Conditional Mixup for early epochs only (alpha=0.2, weak)\n            # This helps generalization without blurring class boundaries\n            if use_mixup and epoch < 5 and random.random() < 0.5:\n                mixed_images, y_a, y_b, lam = mixup_data(images, labels, mixup_alpha)\n                outputs = model(mixed_images)\n                loss = mixup_criterion(criterion, outputs, y_a, y_b, lam)\n                _, preds_batch = torch.max(outputs, 1)\n                correct += (lam * (preds_batch == y_a).sum().item() +\n                           (1 - lam) * (preds_batch == y_b).sum().item())\n            else:\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                _, preds_batch = torch.max(outputs, 1)\n                correct += (preds_batch == labels).sum().item()\n\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n\n            running_loss += loss.item()\n            total += labels.size(0)\n            loop.set_postfix(loss=loss.item(), acc=correct/total)\n\n        train_loss = running_loss / len(train_loader)\n        train_acc = correct / total\n\n        # ===== VALIDATION =====\n        model.eval()\n        val_loss_sum, correct, total = 0, 0, 0\n        all_preds, all_labels = [], []\n\n        with torch.no_grad():\n            for images, labels in val_loader:\n                images, labels = images.to(device), labels.to(device)\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n\n                val_loss_sum += loss.item()\n                _, preds_batch = torch.max(outputs, 1)\n                correct += (preds_batch == labels).sum().item()\n                total += labels.size(0)\n\n                all_preds.extend(preds_batch.cpu().numpy())\n                all_labels.extend(labels.cpu().numpy())\n\n        val_loss = val_loss_sum / len(val_loader)\n        val_acc = correct / total\n        val_f1 = f1_score(all_labels, all_preds, average='weighted', zero_division=0)\n\n        scheduler.step()\n\n        train_losses.append(train_loss)\n        val_losses.append(val_loss)\n        train_accs.append(train_acc)\n        val_accs.append(val_acc)\n        val_f1s.append(val_f1)\n\n        # Print epoch results\n        current_lr = optimizer.param_groups[0]['lr']\n        print(f\"\\nEpoch {epoch+1}/{epochs} | LR: {current_lr:.2e}\")\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} | Val F1: {val_f1:.4f}\")\n        \n        # OPTIONAL: Monitor train-val accuracy gap (warning if >0.1)\n        acc_gap = train_acc - val_acc\n        if acc_gap > 0.10:\n            print(f\"  ⚠️  Overfitting indicator: Acc gap = {acc_gap:.4f} (>0.10)\")\n        print(\"-\"*60)\n\n        # Save best model based on val accuracy\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            patience_counter = 0\n            torch.save({\n                'epoch': epoch + 1,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_acc': val_acc,\n                'val_f1': val_f1,\n                'train_acc': train_acc,\n            }, save_path)\n            print(f\"  ✓ Best model saved! Val Acc: {val_acc:.4f}, Val F1: {val_f1:.4f}\")\n        else:\n            patience_counter += 1\n            print(f\"  No improvement. Patience: {patience_counter}/{patience}\")\n\n        # Early stopping\n        if patience_counter >= patience:\n            print(f\"\\nEarly stopping triggered at epoch {epoch+1}!\")\n            break\n\n    print(f\"\\nTraining complete. Best Val Acc: {best_val_acc:.4f}\")\n    print(f\"Best model saved to: {save_path}\")\n\n    return train_losses, val_losses, train_accs, val_accs, val_f1s","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:47.635587Z","iopub.execute_input":"2026-04-09T14:05:47.635916Z","iopub.status.idle":"2026-04-09T14:05:47.658002Z","shell.execute_reply.started":"2026-04-09T14:05:47.635868Z","shell.execute_reply":"2026-04-09T14:05:47.657477Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trainable = [name for name, p in model.named_parameters() if p.requires_grad]\nprint(f\"Total trainable layers: {len(trainable)}\")\nprint(\"Last 10 trainable:\", trainable[-10:])","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:47.659659Z","iopub.execute_input":"2026-04-09T14:05:47.660406Z","iopub.status.idle":"2026-04-09T14:05:47.676312Z","shell.execute_reply.started":"2026-04-09T14:05:47.660381Z","shell.execute_reply":"2026-04-09T14:05:47.675721Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history = train_model(model, train_loader, val_loader, epochs=30, use_mixup=True, mixup_alpha=0.4)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T14:05:47.677337Z","iopub.execute_input":"2026-04-09T14:05:47.677929Z","iopub.status.idle":"2026-04-09T16:30:35.371062Z","shell.execute_reply.started":"2026-04-09T14:05:47.677905Z","shell.execute_reply":"2026-04-09T16:30:35.369931Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_history(history):\n    train_losses, val_losses, train_accs, val_accs, val_f1s = history\n\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n    # Loss plot\n    axes[0].plot(train_losses, label=\"Train Loss\", linewidth=2)\n    axes[0].plot(val_losses, label=\"Validation Loss\", linewidth=2)\n    axes[0].set_xlabel(\"Epochs\")\n    axes[0].set_ylabel(\"Loss\")\n    axes[0].set_title(\"Training vs Validation Loss\")\n    axes[0].legend()\n    axes[0].grid(True, alpha=0.3)\n\n    # Accuracy plot\n    axes[1].plot(train_accs, label=\"Train Accuracy\", linewidth=2)\n    axes[1].plot(val_accs, label=\"Validation Accuracy\", linewidth=2)\n    axes[1].set_xlabel(\"Epochs\")\n    axes[1].set_ylabel(\"Accuracy\")\n    axes[1].set_title(\"Training vs Validation Accuracy\")\n    axes[1].legend()\n    axes[1].grid(True, alpha=0.3)\n\n    # F1 plot (weighted)\n    axes[2].plot(val_f1s, label=\"Validation F1 (weighted)\", linewidth=2, color='green')\n    axes[2].set_xlabel(\"Epochs\")\n    axes[2].set_ylabel(\"F1 Score\")\n    axes[2].set_title(\"Weighted F1 Score (Better for Imbalance)\")\n    axes[2].legend()\n    axes[2].grid(True, alpha=0.3)\n\n    plt.tight_layout()\n    plt.savefig(os.path.join(SAVE_DIR, 'training_curves.png'), dpi=150, bbox_inches='tight')\n    plt.show()\n\nplot_history(history)","metadata":{"execution":{"iopub.status.busy":"2026-04-09T16:30:35.372718Z","iopub.execute_input":"2026-04-09T16:30:35.372997Z","iopub.status.idle":"2026-04-09T16:30:36.421008Z","shell.execute_reply.started":"2026-04-09T16:30:35.372967Z","shell.execute_reply":"2026-04-09T16:30:36.420418Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load best model for evaluation\ncheckpoint = torch.load(os.path.join(SAVE_DIR, 'eff_best.pth'), map_location=device, weights_only=False)\nmodel.load_state_dict(checkpoint['model_state_dict'])\nprint(f\"Loaded best model from epoch {checkpoint['epoch']}\")\nprint(f\"Best Val Acc: {checkpoint['val_acc']:.4f}, Best Val F1: {checkpoint['val_f1']:.4f}\")\n\nmodel.eval()\n\npreds, labels_all = [], []\n\nwith torch.no_grad():\n    for images, labels in tqdm(test_loader, desc=\"Testing\"):\n        images = images.to(device)\n        outputs = model(images)\n\n        preds.extend(outputs.argmax(1).cpu().numpy())\n        labels_all.extend(labels.numpy())\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"TEST SET RESULTS - EfficientNet-B5\")\nprint(\"=\"*60)\nprint(classification_report(labels_all, preds,\n      target_names=['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative']))\n\n# Comprehensive metrics\ntest_accuracy = np.mean(np.array(preds) == np.array(labels_all))\ntest_qwk = cohen_kappa_score(labels_all, preds, weights='quadratic')\ntest_f1 = f1_score(labels_all, preds, average='weighted')\ntest_precision = precision_score(labels_all, preds, average='weighted')\ntest_recall = recall_score(labels_all, preds, average='weighted')\n\nprint(f\"\\n{'='*60}\")\nprint(f\"Aggregated Metrics:\")\nprint(f\"{'='*60}\")\nprint(f\"Test Accuracy:              {test_accuracy:.4f}\")\nprint(f\"Test F1 Score (weighted):   {test_f1:.4f}\")\nprint(f\"Test Precision (weighted):  {test_precision:.4f}\")\nprint(f\"Test Recall (weighted):     {test_recall:.4f}\")\nprint(f\"Quadratic Weighted Kappa:   {test_qwk:.4f}\")\nprint(f\"{'='*60}\")","metadata":{"execution":{"iopub.status.busy":"2026-04-09T16:32:53.931697Z","iopub.execute_input":"2026-04-09T16:32:53.932349Z","iopub.status.idle":"2026-04-09T16:33:40.365278Z","shell.execute_reply.started":"2026-04-09T16:32:53.932315Z","shell.execute_reply":"2026-04-09T16:33:40.364503Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cm = confusion_matrix(labels_all, preds)\n\nplt.figure(figsize=(8, 6))\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\",\n            xticklabels=['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative'],\n            yticklabels=['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative'])\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"Actual\")\nplt.title(\"Confusion Matrix - EfficientNet-B5\")\nplt.tight_layout()\nplt.savefig(os.path.join(SAVE_DIR, 'confusion_matrix.png'), dpi=150, bbox_inches='tight')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2026-04-09T16:33:40.366799Z","iopub.execute_input":"2026-04-09T16:33:40.367033Z","iopub.status.idle":"2026-04-09T16:33:40.802189Z","shell.execute_reply.started":"2026-04-09T16:33:40.367007Z","shell.execute_reply":"2026-04-09T16:33:40.801491Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import HTML\n\nHTML('<a href=\"/kaggle/working/eff_best.pth\" download>Click here to download</a>')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-09T16:30:37.060073Z","iopub.status.idle":"2026-04-09T16:30:37.060409Z","shell.execute_reply.started":"2026-04-09T16:30:37.060273Z","shell.execute_reply":"2026-04-09T16:30:37.060289Z"}},"outputs":[],"execution_count":null}]}