{"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":"datasetVersion","sourceId":2812287,"datasetId":1719146,"databundleVersionId":2858573}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install albumentations timm -q\nprint(\"Done\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-03-14T14:05:28.261037Z","iopub.execute_input":"2026-03-14T14:05:28.261277Z","iopub.status.idle":"2026-03-14T14:05:32.326751Z","shell.execute_reply.started":"2026-03-14T14:05:28.261255Z","shell.execute_reply":"2026-03-14T14:05:32.326052Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import classification_report, roc_auc_score, confusion_matrix, roc_curve\nimport warnings\nwarnings.filterwarnings('ignore')\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Device  : {DEVICE}\")\nprint(f\"PyTorch : {torch.__version__}\")\nprint(f\"Timm    : {timm.__version__}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T14:05:40.939726Z","iopub.execute_input":"2026-03-14T14:05:40.940459Z","iopub.status.idle":"2026-03-14T14:05:56.415737Z","shell.execute_reply.started":"2026-03-14T14:05:40.940421Z","shell.execute_reply":"2026-03-14T14:05:56.415059Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Paths ──\nAPTOS_TRAIN_CSV  = '/kaggle/input/competitions/aptos2019-blindness-detection/train.csv'\nAPTOS_TRAIN_IMGS = '/kaggle/input/competitions/aptos2019-blindness-detection/train_images'\nAPTOS_TEST_CSV   = '/kaggle/input/competitions/aptos2019-blindness-detection/test.csv'\nAPTOS_TEST_IMGS  = '/kaggle/input/competitions/aptos2019-blindness-detection/test_images'\nIDRID_IMGS       = '/kaggle/input/datasets/mariaherrerot/idrid-dataset/Imagenes/Imagenes'\nIDRID_CSV        = '/kaggle/input/datasets/mariaherrerot/idrid-dataset/idrid_labels.csv'\n\n# ── Config ──\nIMG_SIZE    = 300\nBATCH_SIZE  = 16\nLR          = 3e-4\nEPOCHS      = 35\nNUM_WORKERS = 2\nCKPT        = 'best_dr_model.pth'\n\n# ── Verify paths ──\npaths = {\n    'APTOS Train CSV' : APTOS_TRAIN_CSV,\n    'APTOS Train Imgs': APTOS_TRAIN_IMGS,\n    'APTOS Test CSV'  : APTOS_TEST_CSV,\n    'APTOS Test Imgs' : APTOS_TEST_IMGS,\n    'IDRiD Imgs'      : IDRID_IMGS,\n    'IDRiD CSV'       : IDRID_CSV,\n}\nfor name, path in paths.items():\n    print(f\"  {'✓' if os.path.exists(path) else '✗'} {name}: {path}\")\n\nprint(f\"\\nIMG_SIZE   : {IMG_SIZE}\")\nprint(f\"BATCH_SIZE : {BATCH_SIZE}\")\nprint(f\"EPOCHS     : {EPOCHS}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T14:06:09.766291Z","iopub.execute_input":"2026-03-14T14:06:09.767470Z","iopub.status.idle":"2026-03-14T14:06:09.799182Z","shell.execute_reply.started":"2026-03-14T14:06:09.767440Z","shell.execute_reply":"2026-03-14T14:06:09.798597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def crop_black_borders(img, threshold=10):\n    gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    _, binary = cv2.threshold(gray, threshold, 255, cv2.THRESH_BINARY)\n    contours, _ = cv2.findContours(binary, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    if not contours:\n        return img\n    x, y, w, h = cv2.boundingRect(max(contours, key=cv2.contourArea))\n    return img[y:y+h, x:x+w]\n\ndef ben_graham_normalization(img, sigmaX=10):\n    return cv2.addWeighted(img, 4, cv2.GaussianBlur(img, (0, 0), sigmaX), -4, 128)\n\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_BGR2LAB)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    lab[:, :, 0] = clahe.apply(lab[:, :, 0])\n    return cv2.cvtColor(lab, cv2.COLOR_LAB2BGR)\n\ndef preprocess_fundus(img_path, size=IMG_SIZE):\n    img = cv2.imread(img_path)\n    if img is None:\n        raise ValueError(f\"Cannot load: {img_path}\")\n    img = crop_black_borders(img)\n    img = cv2.resize(img, (size, size))\n    img = apply_clahe(img)\n    img = ben_graham_normalization(img)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    return img\n\n# ── Visual check ──\nsamples = glob(os.path.join(APTOS_TRAIN_IMGS, '*.png'))[:2]\nfig, axes = plt.subplots(2, 2, figsize=(12, 10))\nfor i, p in enumerate(samples):\n    raw  = cv2.resize(cv2.cvtColor(cv2.imread(p), cv2.COLOR_BGR2RGB), (IMG_SIZE, IMG_SIZE))\n    proc = preprocess_fundus(p)\n    axes[i][0].imshow(raw);  axes[i][0].set_title('Original');     axes[i][0].axis('off')\n    axes[i][1].imshow(proc); axes[i][1].set_title('Preprocessed'); axes[i][1].axis('off')\nplt.suptitle('CLAHE + Ben Graham Normalization', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.show()\nprint(\"Preprocessing ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T14:06:16.669884Z","iopub.execute_input":"2026-03-14T14:06:16.670716Z","iopub.status.idle":"2026-03-14T14:06:17.952377Z","shell.execute_reply.started":"2026-03-14T14:06:16.670689Z","shell.execute_reply":"2026-03-14T14:06:17.951427Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Load CSVs ──\naptos_df = pd.read_csv(APTOS_TRAIN_CSV)\nidrid_df = pd.read_csv(IDRID_CSV, usecols=['id_code', 'diagnosis']).dropna(subset=['diagnosis'])\n\n# ── Add image paths ──\naptos_df['img_path'] = aptos_df['id_code'].apply(\n    lambda x: os.path.join(APTOS_TRAIN_IMGS, x + '.png'))\nidrid_df['img_path'] = idrid_df['id_code'].apply(\n    lambda x: os.path.join(IDRID_IMGS, x + '.jpg'))\n\n# ── Rename to common column ──\naptos_df = aptos_df.rename(columns={'diagnosis': 'grade'})[['img_path', 'grade']]\nidrid_df = idrid_df.rename(columns={'diagnosis': 'grade'})[['img_path', 'grade']]\naptos_df['source'] = 'aptos'\nidrid_df['source'] = 'idrid'\n\n# ── Filter only existing files ──\naptos_df = aptos_df[aptos_df['img_path'].apply(os.path.exists)].reset_index(drop=True)\nidrid_df = idrid_df[idrid_df['img_path'].apply(os.path.exists)].reset_index(drop=True)\n\n# ── Combine ──\ncombined_df = pd.concat([aptos_df, idrid_df], ignore_index=True)\ncombined_df['grade']  = combined_df['grade'].astype(int)\ncombined_df['binary'] = (combined_df['grade'] > 0).astype(int)\n\n# ── Summary ──\nprint(f\"{'='*40}\")\nprint(f\"COMBINED DATASET\")\nprint(f\"{'='*40}\")\nprint(f\"  APTOS  : {len(aptos_df)}\")\nprint(f\"  IDRiD  : {len(idrid_df)}\")\nprint(f\"  Total  : {len(combined_df)}\")\nprint(f\"\\n  Grade distribution:\")\nprint(combined_df['grade'].value_counts().sort_index().to_string())\nprint(f\"\\n  No DR : {(combined_df['binary']==0).sum()}\")\nprint(f\"  DR    : {(combined_df['binary']==1).sum()}\")\nprint(f\"{'='*40}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T14:06:30.349716Z","iopub.execute_input":"2026-03-14T14:06:30.350283Z","iopub.status.idle":"2026-03-14T14:06:43.641683Z","shell.execute_reply.started":"2026-03-14T14:06:30.350255Z","shell.execute_reply":"2026-03-14T14:06:43.640918Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── Train/Val split stratified by grade ──\ntrain_df, val_df = train_test_split(\n    combined_df, test_size=0.2, stratify=combined_df['grade'], random_state=42\n)\ntrain_df = train_df.reset_index(drop=True)\nval_df   = val_df.reset_index(drop=True)\n\n# ── Dataset class ──\nclass DRDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df        = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = preprocess_fundus(row['img_path'])\n        if self.transform:\n            img = self.transform(image=img)['image']\n        label = torch.tensor(row['binary'], dtype=torch.long)\n        return img, label\n\n# ── Augmentations ──\ntrain_transform = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=20, p=0.5),\n    A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.1, hue=0.05, p=0.5),\n    A.GaussianBlur(blur_limit=3, p=0.2),\n    A.GaussNoise(p=0.2),\n    A.CoarseDropout(max_holes=8, max_height=32, max_width=32, fill_value=0, p=0.3),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\nval_transform = A.Compose([\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\n# ── Datasets ──\ntrain_ds = DRDataset(train_df, transform=train_transform)\nval_ds   = DRDataset(val_df,   transform=val_transform)\n\n# ── WeightedRandomSampler ──\nlabels      = train_df['binary'].values\nclass_count = np.bincount(labels)\nweights     = 1.0 / class_count[labels]\nsampler     = WeightedRandomSampler(weights, num_samples=len(weights), replacement=True)\n\n# ── DataLoaders ──\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                          num_workers=NUM_WORKERS, pin_memory=True)\nval_loader   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=True)\n\nprint(f\"Train : {len(train_ds)} samples | {len(train_loader)} batches\")\nprint(f\"Val   : {len(val_ds)} samples | {len(val_loader)} batches\")\nprint(f\"\\nTrain — No DR: {(train_df['binary']==0).sum()} | DR: {(train_df['binary']==1).sum()}\")\nprint(f\"Val   — No DR: {(val_df['binary']==0).sum()} | DR: {(val_df['binary']==1).sum()}\")\nprint(\"Dataloaders ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T14:06:48.950170Z","iopub.execute_input":"2026-03-14T14:06:48.950569Z","iopub.status.idle":"2026-03-14T14:06:48.982743Z","shell.execute_reply.started":"2026-03-14T14:06:48.950542Z","shell.execute_reply":"2026-03-14T14:06:48.982200Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DRClassifier(nn.Module):\n    def __init__(self, num_classes=2):\n        super().__init__()\n        self.backbone = timm.create_model(\n            'efficientnet_b4', pretrained=True, num_classes=0, global_pool=''\n        )\n        self.backbone.set_grad_checkpointing(enable=True)\n        feat_dim = self.backbone.num_features\n\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.head = nn.Sequential(\n            nn.Flatten(),\n            nn.Dropout(0.4),\n            nn.Linear(feat_dim, 512),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n        self.gradients   = None\n        self.activations = None\n\n    def save_gradient(self, grad):\n        self.gradients = grad\n\n    def forward(self, x):\n        feats = self.backbone(x)\n        if feats.requires_grad:\n            feats.register_hook(self.save_gradient)\n        self.activations = feats\n        return self.head(self.pool(feats))\n\n# ── Initialize everything ──\ntorch.cuda.empty_cache()\nmodel     = DRClassifier(num_classes=2).to(DEVICE)\ncriterion = nn.CrossEntropyLoss(label_smoothing=0.1)\noptimizer = AdamW(model.parameters(), lr=LR, weight_decay=1e-4)\nscheduler = CosineAnnealingLR(optimizer, T_max=EPOCHS, eta_min=1e-6)\nscaler    = torch.cuda.amp.GradScaler()\n\n# ── Forward pass test ──\ndummy = torch.randn(2, 3, IMG_SIZE, IMG_SIZE).to(DEVICE)\nwith torch.no_grad():\n    out = model(dummy)\ndel dummy\ntorch.cuda.empty_cache()\n\nprint(f\"Output shape : {list(out.shape)}\")\nprint(f\"GPU memory   : {torch.cuda.memory_allocated()/1024**3:.2f} GB\")\nprint(f\"Parameters   : {sum(p.numel() for p in model.parameters()):,}\")\nprint(\"Model ready.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T14:06:55.826217Z","iopub.execute_input":"2026-03-14T14:06:55.826934Z","iopub.status.idle":"2026-03-14T14:06:59.234193Z","shell.execute_reply.started":"2026-03-14T14:06:55.826901Z","shell.execute_reply":"2026-03-14T14:06:59.233508Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, criterion, scaler):\n    model.train()\n    total_loss, correct, total = 0, 0, 0\n    for imgs, labels in loader:\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n        optimizer.zero_grad()\n        with torch.cuda.amp.autocast():\n            outputs = model(imgs)\n            loss    = criterion(outputs, labels)\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        total_loss += loss.item()\n        correct    += (outputs.argmax(1) == labels).sum().item()\n        total      += labels.size(0)\n    return total_loss / len(loader), correct / total\n\n\ndef validate(model, loader):\n    model.eval()\n    all_preds, all_probs, all_labels = [], [], []\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs = imgs.to(DEVICE)\n            with torch.cuda.amp.autocast():\n                outputs = model(imgs)\n            probs = F.softmax(outputs, dim=1)[:, 1]\n            all_preds.extend(outputs.argmax(1).cpu().numpy())\n            all_probs.extend(probs.cpu().numpy())\n            all_labels.extend(labels.numpy())\n    auc = roc_auc_score(all_labels, all_probs)\n    acc = np.mean(np.array(all_preds) == np.array(all_labels))\n    return auc, acc, np.array(all_preds), np.array(all_probs), np.array(all_labels)\n\n\n# ── Training loop ──\nbest_auc = 0.0\nhistory  = {'train_loss': [], 'train_acc': [], 'val_auc': [], 'val_acc': []}\n\nfor epoch in range(EPOCHS):\n    train_loss, train_acc     = train_one_epoch(model, train_loader, optimizer, criterion, scaler)\n    val_auc, val_acc, _, _, _ = validate(model, val_loader)\n    scheduler.step()\n\n    history['train_loss'].append(train_loss)\n    history['train_acc'].append(train_acc)\n    history['val_auc'].append(val_auc)\n    history['val_acc'].append(val_acc)\n\n    print(f\"Epoch [{epoch+1:02d}/{EPOCHS}] | Loss: {train_loss:.4f} | \"\n          f\"Train Acc: {train_acc:.4f} | Val AUC: {val_auc:.4f} | Val Acc: {val_acc:.4f}\")\n\n    if val_auc > best_auc:\n        best_auc = val_auc\n        torch.save(model.state_dict(), CKPT)\n        print(f\"  ✓ Saved (AUC: {best_auc:.4f})\")\n\nprint(f\"\\nTraining complete. Best AUC: {best_auc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T14:07:03.261901Z","iopub.execute_input":"2026-03-14T14:07:03.262638Z","iopub.status.idle":"2026-03-14T16:39:34.190610Z","shell.execute_reply.started":"2026-03-14T14:07:03.262611Z","shell.execute_reply":"2026-03-14T16:39:34.189879Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ── File is already in /kaggle/working, skip copy ──\nprint(\"Model location:\", os.path.abspath('best_dr_model.pth'))\nprint(\"File size:\", round(os.path.getsize('best_dr_model.pth') / 1024**2, 2), \"MB\")\n\n# ── Load best weights ──\nmodel.load_state_dict(torch.load('best_dr_model.pth'))\nmodel.eval()\n\n# ── Run validation ──\nall_preds, all_probs, all_labels = [], [], []\nwith torch.no_grad():\n    for imgs, labels in val_loader:\n        imgs = imgs.to(DEVICE)\n        with torch.cuda.amp.autocast():\n            outputs = model(imgs)\n        probs = F.softmax(outputs, dim=1)[:, 1]\n        all_preds.extend(outputs.argmax(1).cpu().numpy())\n        all_probs.extend(probs.cpu().numpy())\n        all_labels.extend(labels.numpy())\n\nall_preds  = np.array(all_preds)\nall_probs  = np.array(all_probs)\nall_labels = np.array(all_labels)\n\n# ── Metrics ──\nprint(\"\\n\" + \"=\"*55)\nprint(\"FINAL EVALUATION\")\nprint(\"=\"*55)\nprint(classification_report(all_labels, all_preds,\n      target_names=['No DR', 'DR'], digits=4))\nprint(f\"AUC-ROC : {roc_auc_score(all_labels, all_probs):.4f}\")\n\n# ── Plots ──\nfig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\ncm = confusion_matrix(all_labels, all_preds)\naxes[0].imshow(cm, cmap='Blues')\naxes[0].set_xticks([0,1]); axes[0].set_yticks([0,1])\naxes[0].set_xticklabels(['No DR','DR'])\naxes[0].set_yticklabels(['No DR','DR'])\naxes[0].set_xlabel('Predicted'); axes[0].set_ylabel('Actual')\naxes[0].set_title('Confusion Matrix')\nfor i in range(2):\n    for j in range(2):\n        axes[0].text(j, i, str(cm[i,j]), ha='center', va='center',\n                     color='white' if cm[i,j] > cm.max()/2 else 'black',\n                     fontsize=16, fontweight='bold')\n\nfpr, tpr, _ = roc_curve(all_labels, all_probs)\naxes[1].plot(fpr, tpr, 'b-', lw=2, label=f'AUC = {roc_auc_score(all_labels, all_probs):.4f}')\naxes[1].plot([0,1],[0,1], 'r--', alpha=0.5, label='Random')\naxes[1].set_xlabel('False Positive Rate')\naxes[1].set_ylabel('True Positive Rate')\naxes[1].set_title('ROC Curve')\naxes[1].legend(); axes[1].grid(True)\n\naxes[2].hist(all_probs[all_labels==0], bins=30, alpha=0.6, color='green', label='No DR')\naxes[2].hist(all_probs[all_labels==1], bins=30, alpha=0.6, color='red',   label='DR')\naxes[2].axvline(0.5, color='black', ls='--', label='Threshold=0.5')\naxes[2].set_xlabel('DR Probability')\naxes[2].set_ylabel('Count')\naxes[2].set_title('Confidence Distribution')\naxes[2].legend(); axes[2].grid(True)\n\nplt.suptitle('Model Evaluation', fontsize=14, fontweight='bold')\nplt.tight_layout()\nplt.show()\n\n# ── Threshold table ──\nprint(\"\\nThreshold Analysis (for hospital calibration):\")\nprint(f\"{'Thresh':>8} | {'Sensitivity':>12} | {'Specificity':>12} | {'PPV':>8} | {'NPV':>8}\")\nprint(\"-\" * 58)\nfor t in [0.3, 0.35, 0.4, 0.5, 0.6, 0.7]:\n    p  = (all_probs >= t).astype(int)\n    tp = ((p==1) & (all_labels==1)).sum()\n    tn = ((p==0) & (all_labels==0)).sum()\n    fp = ((p==1) & (all_labels==0)).sum()\n    fn = ((p==0) & (all_labels==1)).sum()\n    print(f\"{t:>8.2f} | {tp/(tp+fn+1e-8):>12.4f} | {tn/(tn+fp+1e-8):>12.4f} | \"\n          f\"{tp/(tp+fp+1e-8):>8.4f} | {tn/(tn+fn+1e-8):>8.4f}\")\n\nprint(\"\\n  For hospital screening → use threshold 0.35-0.40\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T16:40:56.031937Z","iopub.execute_input":"2026-03-14T16:40:56.032669Z","iopub.status.idle":"2026-03-14T16:41:45.090224Z","shell.execute_reply.started":"2026-03-14T16:40:56.032643Z","shell.execute_reply":"2026-03-14T16:41:45.089372Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class GradCAM:\n    def __init__(self, model):\n        self.model = model\n\n    def generate(self, img_tensor):\n        self.model.eval()\n\n        # Temporarily disable grad checkpointing for Grad-CAM\n        self.model.backbone.set_grad_checkpointing(enable=False)\n\n        img_tensor = img_tensor.unsqueeze(0).to(DEVICE)\n\n        # Forward pass — no autocast, need full precision for gradients\n        features = self.model.backbone(img_tensor)   # (1, C, H, W)\n        features.retain_grad()\n\n        pooled  = self.model.pool(features)\n        output  = self.model.head(pooled)\n\n        pred_class = output.argmax(dim=1).item()\n        prob       = F.softmax(output, dim=1)[0, 1].item()\n\n        # Backward\n        self.model.zero_grad()\n        output[0, pred_class].backward()\n\n        # Grad-CAM\n        grads   = features.grad                      # (1, C, H, W)\n        weights = grads.mean(dim=(2, 3), keepdim=True)\n        cam     = F.relu((weights * features).sum(dim=1, keepdim=True))\n        cam     = F.interpolate(cam, (IMG_SIZE, IMG_SIZE), mode='bilinear', align_corners=False)\n        cam     = cam.squeeze().cpu().detach().numpy()\n        cam     = (cam - cam.min()) / (cam.max() - cam.min() + 1e-8)\n\n        # Re-enable grad checkpointing\n        self.model.backbone.set_grad_checkpointing(enable=True)\n\n        return cam, pred_class, prob\n\n\ndef predict_and_explain(img_path, threshold=0.35):\n    gradcam = GradCAM(model)\n    img_rgb = preprocess_fundus(img_path)\n    tensor  = val_transform(image=img_rgb)['image']\n\n    cam, pred, dr_prob = gradcam.generate(tensor)\n\n    heatmap = cv2.applyColorMap(np.uint8(255 * cam), cv2.COLORMAP_JET)\n    heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n    overlay = cv2.addWeighted(img_rgb, 0.5, heatmap, 0.5, 0)\n\n    fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n    fig.patch.set_facecolor('#111111')\n    titles = ['Preprocessed Image', 'Grad-CAM Heatmap\\n(Red = High Attention)', 'Overlay']\n    for ax, im, title in zip(axes, [img_rgb, heatmap, overlay], titles):\n        ax.imshow(im)\n        ax.set_title(title, color='white', fontsize=12, pad=10)\n        ax.axis('off')\n\n    if dr_prob >= threshold:\n        label, color = f'DR DETECTED  |  Probability: {dr_prob:.1%}', '#FF4444'\n    elif dr_prob >= 0.2:\n        label, color = f'BORDERLINE — REVIEW ADVISED  |  Probability: {dr_prob:.1%}', '#FFA500'\n    else:\n        label, color = f'No DR  |  Probability: {dr_prob:.1%}', '#44FF88'\n\n    conf = 'High Confidence' if abs(dr_prob - 0.5) > 0.25 else 'Low Confidence — Manual Review Advised'\n\n    fig.text(0.5, 0.01, f'{label}   |   {conf}',\n             ha='center', fontsize=13, fontweight='bold', color=color,\n             bbox=dict(boxstyle='round,pad=0.5', facecolor='#222222',\n                       edgecolor=color, linewidth=2))\n    plt.suptitle('DR Detection — AI Diagnostic Report',\n                 color='white', fontsize=15, fontweight='bold', y=1.02)\n    plt.tight_layout()\n    plt.show()\n\n    print(f\"  Prediction  : {'DR' if pred==1 else 'No DR'}\")\n    print(f\"  Probability : {dr_prob:.4f}\")\n    print(f\"  Confidence  : {conf}\")\n    return pred, dr_prob\n\n# ── Test on DR and No DR samples ──\nprint(\"── DR samples ──\")\nfor p in val_df[val_df['binary']==1]['img_path'].values[:2]:\n    predict_and_explain(p)\n\nprint(\"\\n── No DR samples ──\")\nfor p in val_df[val_df['binary']==0]['img_path'].values[:2]:\n    predict_and_explain(p)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T16:44:51.978225Z","iopub.execute_input":"2026-03-14T16:44:51.979016Z","iopub.status.idle":"2026-03-14T16:44:54.711046Z","shell.execute_reply.started":"2026-03-14T16:44:51.978987Z","shell.execute_reply":"2026-03-14T16:44:54.710034Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def batch_predict(img_folder, output_csv='dr_screening_results.csv', threshold=0.35):\n    \"\"\"\n    Screen a folder of hospital fundus images.\n    Outputs a CSV ready for clinical review.\n    \"\"\"\n    model.load_state_dict(torch.load(CKPT))\n    model.eval()\n\n    paths = (glob(os.path.join(img_folder, '*.png')) +\n             glob(os.path.join(img_folder, '*.jpg')) +\n             glob(os.path.join(img_folder, '*.jpeg')))\n\n    if not paths:\n        print(f\"No images found in {img_folder}\")\n        return None\n\n    print(f\"Found {len(paths)} images. Running inference...\")\n    results = []\n\n    for img_path in paths:\n        try:\n            img    = preprocess_fundus(img_path)\n            tensor = val_transform(image=img)['image'].unsqueeze(0).to(DEVICE)\n            with torch.no_grad():\n                with torch.cuda.amp.autocast():\n                    output = model(tensor)\n            prob = F.softmax(output, dim=1)[0, 1].item()\n            pred = int(prob >= threshold)\n\n            if prob >= threshold:\n                decision = 'DR DETECTED'\n            elif prob >= 0.2:\n                decision = 'BORDERLINE'\n            else:\n                decision = 'No DR'\n\n            results.append({\n                'image_file'    : os.path.basename(img_path),\n                'prediction'    : decision,\n                'dr_probability': round(prob, 4),\n                'confidence'    : 'High' if abs(prob - 0.5) > 0.25 else 'Low',\n                'manual_review' : 'YES' if (0.2 <= prob <= 0.6) else 'No',\n                'threshold_used': threshold\n            })\n        except Exception as e:\n            results.append({\n                'image_file'    : os.path.basename(img_path),\n                'prediction'    : 'ERROR',\n                'dr_probability': -1,\n                'confidence'    : 'N/A',\n                'manual_review' : 'YES',\n                'error'         : str(e)\n            })\n\n    df = pd.DataFrame(results)\n    df.to_csv(output_csv, index=False)\n\n    valid = df[df['prediction'] != 'ERROR']\n    print(f\"\\n{'='*45}\")\n    print(f\"SCREENING SUMMARY\")\n    print(f\"{'='*45}\")\n    print(f\"  Total processed : {len(valid)}\")\n    print(f\"  DR Detected     : {(valid['prediction']=='DR DETECTED').sum()}\")\n    print(f\"  Borderline      : {(valid['prediction']=='BORDERLINE').sum()}\")\n    print(f\"  No DR           : {(valid['prediction']=='No DR').sum()}\")\n    print(f\"  Manual review   : {(valid['manual_review']=='YES').sum()}\")\n    print(f\"  Errors          : {len(df) - len(valid)}\")\n    print(f\"  Threshold used  : {threshold}\")\n    print(f\"  Results saved   : {output_csv}\")\n    print(f\"{'='*45}\")\n    return df\n\n# ── Test on APTOS test images ──\nresults_df = batch_predict(APTOS_TEST_IMGS)\nprint(\"\\nSample results:\")\nprint(results_df.head(10).to_string(index=False))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-14T16:46:53.707302Z","iopub.execute_input":"2026-03-14T16:46:53.707754Z","iopub.status.idle":"2026-03-14T16:49:53.319703Z","shell.execute_reply.started":"2026-03-14T16:46:53.707727Z","shell.execute_reply":"2026-03-14T16:49:53.318786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}