{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":14774,"databundleVersionId":875431,"sourceType":"competition"},{"sourceId":13220810,"sourceType":"datasetVersion","datasetId":8380137},{"sourceId":13220849,"sourceType":"datasetVersion","datasetId":8380162}],"dockerImageVersionId":31090,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import cv2\nimport os\nfrom tqdm import tqdm\nimport random\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import accuracy_score\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, f1_score, roc_auc_score, classification_report, cohen_kappa_score, confusion_matrix, roc_curve, auc\nimport time\n\n# Pytorch Libraries\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torchvision.models as models\nfrom torch.utils.data import Dataset\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport torch.nn.functional as F","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T08:36:26.757073Z","iopub.execute_input":"2025-12-16T08:36:26.757724Z","iopub.status.idle":"2025-12-16T08:36:36.759406Z","shell.execute_reply.started":"2025-12-16T08:36:26.757699Z","shell.execute_reply":"2025-12-16T08:36:36.758773Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset Overview: APTOS 2019","metadata":{}},{"cell_type":"code","source":"dataset = pd.read_csv(\"/kaggle/input/aptos2019-blindness-detection/train.csv\")\nprint(f\"Total Samples in Dataset: {len(dataset)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T08:36:36.760393Z","iopub.execute_input":"2025-12-16T08:36:36.760809Z","iopub.status.idle":"2025-12-16T08:36:36.779429Z","shell.execute_reply.started":"2025-12-16T08:36:36.760788Z","shell.execute_reply":"2025-12-16T08:36:36.778898Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Path to APTOS-2019 dataset\naptos_root = \"/kaggle/input/aptos2019-blindness-detection\"\nimg_dir = os.path.join(aptos_root, \"train_images\")\ncsv_path = os.path.join(aptos_root, \"train.csv\")\n\n# create full path\ndataset['path'] = dataset['id_code'].apply(lambda x: os.path.join(img_dir, f\"{x}.png\"))\n\ndataset.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T08:36:36.780114Z","iopub.execute_input":"2025-12-16T08:36:36.780305Z","iopub.status.idle":"2025-12-16T08:36:36.807071Z","shell.execute_reply.started":"2025-12-16T08:36:36.780285Z","shell.execute_reply":"2025-12-16T08:36:36.806565Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"\n#### Class 0: No DR\n#### Class 1: Mild\n#### Class 2: Moderate\n#### Class 3: Severe\n#### Class 4: Proliferative\n","metadata":{}},{"cell_type":"code","source":"# Count class distribution\nclass_counts = dataset['diagnosis'].value_counts().sort_index()\n\n# Bar plot\nplt.figure(figsize=(6,4))\nsns.barplot(x=class_counts.index, y=class_counts.values, palette=\"viridis\")\nplt.title(\"Class Distribution (APTOS-2019)\")\nplt.xlabel(\"DR Class\")\nplt.ylabel(\"Number of Images\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T08:36:36.808450Z","iopub.execute_input":"2025-12-16T08:36:36.809029Z","iopub.status.idle":"2025-12-16T08:36:37.043229Z","shell.execute_reply.started":"2025-12-16T08:36:36.809004Z","shell.execute_reply":"2025-12-16T08:36:37.042563Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Pie chart\nplt.figure(figsize=(6,6))\nplt.pie(class_counts.values, labels=class_counts.index, autopct='%1.1f%%', startangle=90, colors=sns.color_palette(\"viridis\", len(class_counts)))\nplt.title(\"Class Distribution (Pie Chart)\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T08:36:37.043879Z","iopub.execute_input":"2025-12-16T08:36:37.044124Z","iopub.status.idle":"2025-12-16T08:36:37.158464Z","shell.execute_reply.started":"2025-12-16T08:36:37.044105Z","shell.execute_reply":"2025-12-16T08:36:37.157799Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Example Images for each Diabetic Retinopathy Class","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(12, 8))\nfor i, c in enumerate(sorted(dataset['diagnosis'].unique())):\n    sample_paths = dataset[dataset['diagnosis']==c]['path'].sample(3, random_state=42)\n    for j, p in enumerate(sample_paths):\n        img = cv2.imread(p); img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        plt.subplot(len(class_counts), 3, i*3 + j + 1)\n        plt.imshow(img)\n        plt.axis(\"off\")\n        if j==1: plt.title(f\"Class: {c}\")\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T08:36:37.159186Z","iopub.execute_input":"2025-12-16T08:36:37.159551Z","iopub.status.idle":"2025-12-16T08:36:44.601715Z","shell.execute_reply.started":"2025-12-16T08:36:37.159507Z","shell.execute_reply":"2025-12-16T08:36:44.600951Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Image Preprocessing and Augmentation","metadata":{}},{"cell_type":"code","source":"# Training transform\ntrain_tf = A.Compose([\n    A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=15,\n                       border_mode=cv2.BORDER_CONSTANT, fill=0, fill_mask=0, p=0.5),\n    A.HorizontalFlip(p=0.5),\n    A.RandomBrightnessContrast(0.15,0.15,p=0.5),\n    A.GaussNoise(std_range=(0.02,0.08), p=0.2),\n    A.Normalize(mean=(0.485,0.456,0.406), std=(0.229,0.224,0.225)),\n    ToTensorV2()\n], additional_targets={\n    'mask_ma':'mask','mask_he':'mask','mask_ex':'mask','mask_se':'mask'\n})\n\n# Validation/test transform\nval_tf = A.Compose([\n    A.Normalize(mean=(0.485,0.456,0.406), std=(0.229,0.224,0.225)),\n    ToTensorV2()\n], additional_targets={\n    'mask_ma':'mask','mask_he':'mask','mask_ex':'mask','mask_se':'mask'\n})","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T08:36:44.602505Z","iopub.execute_input":"2025-12-16T08:36:44.602769Z","iopub.status.idle":"2025-12-16T08:36:44.616905Z","shell.execute_reply.started":"2025-12-16T08:36:44.602747Z","shell.execute_reply":"2025-12-16T08:36:44.616241Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def retina_crop(img):\n    \"\"\"Crop black borders around retina using green channel + thresholding.\"\"\"\n    g = cv2.GaussianBlur(img[:,:,1], (0,0), 5)\n    _, th = cv2.threshold(g, 0, 255, cv2.THRESH_OTSU)\n    cnts, _ = cv2.findContours(th, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    if len(cnts) > 0:\n        c = max(cnts, key=cv2.contourArea)\n        x,y,w,h = cv2.boundingRect(c)\n        pad = int(0.02*max(w,h))\n        x = max(0,x-pad); y = max(0,y-pad)\n        return img[y:y+h+2*pad, x:x+w+2*pad]\n    return img\n\ndef illumination_correction(img):\n    \"\"\"Correct uneven lighting using Gaussian blur background division.\"\"\"\n    imgf = img.astype(np.float32) / 255.0\n    k = int(round(min(img.shape[:2]) * 0.05)) | 1  # kernel size ~5% of min dim\n    bg = cv2.GaussianBlur(imgf, (k,k), 0)\n    corr = np.clip((imgf/(bg+1e-3))*bg.mean(), 0, 1)\n    return (corr*255).astype(np.uint8)\n\ndef clahe_green(img):\n    \"\"\"Apply CLAHE on green channel.\"\"\"\n    g = img[:,:,1]\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    g2 = clahe.apply(g)\n    img[:,:,1] = g2\n    return img\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-16T08:36:44.617632Z","iopub.execute_input":"2025-12-16T08:36:44.617869Z","iopub.status.idle":"2025-12-16T08:36:44.625995Z","shell.execute_reply.started":"2025-12-16T08:36:44.617850Z","shell.execute_reply":"2025-12-16T08:36:44.625309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cache_dir = \"data/aptos2019_preprocessed\"\nos.makedirs(cache_dir, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T08:36:44.626752Z","iopub.execute_input":"2025-12-16T08:36:44.627025Z","iopub.status.idle":"2025-12-16T08:36:44.639107Z","shell.execute_reply.started":"2025-12-16T08:36:44.626999Z","shell.execute_reply":"2025-12-16T08:36:44.638382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_and_save(src_paths, dst_dir):\n    for path in tqdm(src_paths, desc=\"Preprocessing & Saving\"):\n        img = cv2.imread(path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        # Apply your preprocessing pipeline\n        img = retina_crop(img)\n        img = illumination_correction(img)\n        img = clahe_green(img)\n\n        # Resize to 224x224 (or whatever you train with)\n        img = cv2.resize(img, (224, 224), interpolation=cv2.INTER_CUBIC)\n\n        # Save as PNG\n        filename = os.path.basename(path)\n        save_path = os.path.join(dst_dir, filename)\n        cv2.imwrite(save_path, cv2.cvtColor(img, cv2.COLOR_RGB2BGR))  # back to BGR for OpenCV save\n\n# Example usage\npreprocess_and_save(dataset['path'].values, cache_dir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T08:36:44.641540Z","iopub.execute_input":"2025-12-16T08:36:44.641790Z","iopub.status.idle":"2025-12-16T09:07:52.375635Z","shell.execute_reply.started":"2025-12-16T08:36:44.641775Z","shell.execute_reply":"2025-12-16T09:07:52.374812Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"dataset['cached_path'] = dataset['id_code'].apply(lambda x: os.path.join(cache_dir, f\"{x}.png\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:52.376615Z","iopub.execute_input":"2025-12-16T09:07:52.376843Z","iopub.status.idle":"2025-12-16T09:07:52.384564Z","shell.execute_reply.started":"2025-12-16T09:07:52.376825Z","shell.execute_reply":"2025-12-16T09:07:52.383977Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Extract labels & paths\nimg_paths = dataset['cached_path'].values\nlabels = dataset['diagnosis'].values\n\ndataset.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:52.385192Z","iopub.execute_input":"2025-12-16T09:07:52.385408Z","iopub.status.idle":"2025-12-16T09:07:52.409227Z","shell.execute_reply.started":"2025-12-16T09:07:52.385383Z","shell.execute_reply":"2025-12-16T09:07:52.408572Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class FundusDataset(Dataset):\n    def __init__(self, img_paths, labels=None, masks=None, transform=None, is_train=True):\n        \"\"\"\n        img_paths: list of image file paths\n        labels: list of int labels (for classification)\n        masks: dict of mask paths { 'ma':[], 'he':[], 'ex':[], 'se':[] } (for IDRiD)\n        transform: albumentations transform\n        is_train: bool, whether training or not\n        \"\"\"\n        self.img_paths = img_paths\n        self.labels = labels\n        self.masks = masks\n        self.transform = transform\n        self.is_train = is_train\n\n    def __len__(self):\n        return len(self.img_paths)\n\n    def __getitem__(self, idx):\n        # Load image\n        img = cv2.imread(self.img_paths[idx])\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        # Prepare masks if available\n        mask_dict = {}\n        if self.masks is not None:\n            for key in self.masks.keys():\n                mask_path = self.masks[key][idx]\n                mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\n                mask = (mask > 127).astype('uint8')  # binarize {0,1}\n                mask_dict[f\"mask_{key}\"] = mask\n\n        # Apply transforms\n        if self.transform:\n            if mask_dict:\n                aug = self.transform(image=img, **mask_dict)\n                img = aug['image']\n                masks_out = {k: aug[k] for k in mask_dict.keys()}\n                return img, masks_out\n            else:\n                aug = self.transform(image=img)\n                img = aug['image']\n                label = self.labels[idx]\n                return img, torch.tensor(label, dtype=torch.long)\n\n        # Fallback: return raw\n        if mask_dict:\n            return img, mask_dict\n        else:\n            return img, self.labels[idx]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:52.410052Z","iopub.execute_input":"2025-12-16T09:07:52.410320Z","iopub.status.idle":"2025-12-16T09:07:52.422447Z","shell.execute_reply.started":"2025-12-16T09:07:52.410295Z","shell.execute_reply":"2025-12-16T09:07:52.421966Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Stratified split into train/val/test (70/15/15)\ntrain_paths, temp_paths, train_labels, temp_labels = train_test_split(\n    img_paths, labels, test_size=0.20, stratify=labels, random_state=42\n)\nval_paths, test_paths, val_labels, test_labels = train_test_split(\n    temp_paths, temp_labels, test_size=0.50, stratify=temp_labels, random_state=42\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:52.423057Z","iopub.execute_input":"2025-12-16T09:07:52.423241Z","iopub.status.idle":"2025-12-16T09:07:52.440109Z","shell.execute_reply.started":"2025-12-16T09:07:52.423226Z","shell.execute_reply":"2025-12-16T09:07:52.439569Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# For APTOS classification\ntrain_dataset = FundusDataset(\n    img_paths=list(train_paths),\n    labels=list(train_labels),\n    transform=train_tf\n)\n\n# Validation dataset (no augmentation)\nval_dataset = FundusDataset(\n    img_paths=list(val_paths),\n    labels=list(val_labels),\n    transform=val_tf\n)\n\n# Test dataset (no augmentation)\ntest_dataset = FundusDataset(\n    img_paths=list(test_paths),\n    labels=list(test_labels),\n    transform=val_tf\n)\n\n# For IDRiD lesion validation\n# idrid_dataset = FundusDataset(\n#     img_paths=idrid_img_paths,\n#     masks={\n#         'ma': ma_mask_paths,\n#         'he': he_mask_paths,\n#         'ex': ex_mask_paths,\n#         'se': se_mask_paths\n#     },\n#     transform=val_tf,\n#     is_train=False\n# )\n\ntrain_loader = torch.utils.data.DataLoader(train_dataset, batch_size=32, shuffle=True)\nval_loader   = torch.utils.data.DataLoader(val_dataset, batch_size=32, shuffle=False)\ntest_loader  = torch.utils.data.DataLoader(test_dataset, batch_size=32, shuffle=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:52.440906Z","iopub.execute_input":"2025-12-16T09:07:52.441185Z","iopub.status.idle":"2025-12-16T09:07:52.451212Z","shell.execute_reply.started":"2025-12-16T09:07:52.441159Z","shell.execute_reply":"2025-12-16T09:07:52.450462Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_samples(dataset, num_samples=4, with_masks=False):\n    \"\"\"\n    Visualize original vs transformed fundus images.\n    \n    Args:\n        dataset: instance of FundusDataset (defined earlier).\n        num_samples: how many samples to show.\n        with_masks: True if dataset returns masks (IDRiD).\n    \"\"\"\n    plt.figure(figsize=(12, num_samples * 4))\n    \n    for i in range(num_samples):\n        img_path = dataset.img_paths[i]\n        \n        # --- Load original (without transforms)\n        orig = cv2.imread(img_path)\n        orig = cv2.cvtColor(orig, cv2.COLOR_BGR2RGB)\n        \n        # --- Get transformed sample from dataset\n        sample = dataset[i]\n        \n        if with_masks:\n            img, masks = sample\n            img = img.permute(1,2,0).cpu().numpy()  # CHW → HWC\n            img = (img - img.min()) / (img.max() - img.min())  # normalize for viewing\n        else:\n            img, label = sample\n            img = img.permute(1,2,0).cpu().numpy()\n            img = (img - img.min()) / (img.max() - img.min())\n        \n        # --- Plot original and transformed\n        plt.subplot(num_samples, 2, 2*i+1)\n        plt.imshow(orig)\n        plt.title(f\"Original {i}\")\n        plt.axis(\"off\")\n        \n        plt.subplot(num_samples, 2, 2*i+2)\n        plt.imshow(img)\n        if with_masks:\n            # Overlay one example mask if available\n            for k, m in masks.items():\n                m = m.squeeze().cpu().numpy()\n                plt.contour(m, colors='r', linewidths=0.5)\n        plt.title(f\"Transformed {i}\")\n        plt.axis(\"off\")\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:52.452157Z","iopub.execute_input":"2025-12-16T09:07:52.452353Z","iopub.status.idle":"2025-12-16T09:07:52.465323Z","shell.execute_reply.started":"2025-12-16T09:07:52.452337Z","shell.execute_reply":"2025-12-16T09:07:52.464790Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_samples(train_dataset, num_samples=4, with_masks=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:52.466021Z","iopub.execute_input":"2025-12-16T09:07:52.466337Z","iopub.status.idle":"2025-12-16T09:07:54.143127Z","shell.execute_reply.started":"2025-12-16T09:07:52.466310Z","shell.execute_reply":"2025-12-16T09:07:54.142026Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Resnet50 Model for 5 class DR Classification","metadata":{}},{"cell_type":"code","source":"# Device setup\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:54.143984Z","iopub.execute_input":"2025-12-16T09:07:54.144208Z","iopub.status.idle":"2025-12-16T09:07:54.222716Z","shell.execute_reply.started":"2025-12-16T09:07:54.144190Z","shell.execute_reply":"2025-12-16T09:07:54.221935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load pretrained ResNet50\nresnet50 = models.resnet50(weights=\"ResNet50_Weights.DEFAULT\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:54.223375Z","iopub.execute_input":"2025-12-16T09:07:54.223611Z","iopub.status.idle":"2025-12-16T09:07:55.196423Z","shell.execute_reply.started":"2025-12-16T09:07:54.223585Z","shell.execute_reply":"2025-12-16T09:07:55.195643Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Modify final fully connected layer for 5 DR classes\nnum_features = resnet50.fc.in_features\nresnet50.fc = nn.Linear(num_features, 5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:55.197248Z","iopub.execute_input":"2025-12-16T09:07:55.197451Z","iopub.status.idle":"2025-12-16T09:07:55.201211Z","shell.execute_reply.started":"2025-12-16T09:07:55.197434Z","shell.execute_reply":"2025-12-16T09:07:55.200572Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Move model to GPU\nresnet50 = resnet50.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:55.201932Z","iopub.execute_input":"2025-12-16T09:07:55.202372Z","iopub.status.idle":"2025-12-16T09:07:55.430655Z","shell.execute_reply.started":"2025-12-16T09:07:55.202348Z","shell.execute_reply":"2025-12-16T09:07:55.430091Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define loss function and optimizer\ncriterion = nn.CrossEntropyLoss()   # Later you can add class weights if needed\noptimizer = torch.optim.Adam(resnet50.parameters(), lr=1e-4, weight_decay=1e-5)\n\n# Optional: learning rate scheduler\nscheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=5, gamma=0.1)\n\nprint(resnet50)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:55.431380Z","iopub.execute_input":"2025-12-16T09:07:55.431653Z","iopub.status.idle":"2025-12-16T09:07:55.437952Z","shell.execute_reply.started":"2025-12-16T09:07:55.431628Z","shell.execute_reply":"2025-12-16T09:07:55.437262Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(model, train_loader, val_loader, criterion, optimizer, scheduler, num_epochs=20, device=\"cuda\"):\n    since = time.time()\n    \n    best_model_wts = model.state_dict()\n    best_acc = 0.0\n\n    train_losses, val_losses = [], []\n    train_accs, val_accs = [], []\n\n    for epoch in range(num_epochs):\n        print(f\"\\nEpoch {epoch+1}/{num_epochs}\")\n        print(\"-\" * 30)\n\n        # ---- Training phase ----\n        model.train()\n        running_loss, running_corrects, total = 0.0, 0, 0\n\n        for inputs, labels in tqdm(train_loader, desc=\"Training\", leave=False):\n            inputs, labels = inputs.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(inputs)\n            loss = criterion(outputs, labels)\n\n            _, preds = torch.max(outputs, 1)\n            loss.backward()\n            optimizer.step()\n\n            running_loss += loss.item() * inputs.size(0)\n            running_corrects += torch.sum(preds == labels.data)\n            total += labels.size(0)\n\n        epoch_loss = running_loss / total\n        epoch_acc = running_corrects.double() / total\n        train_losses.append(epoch_loss)\n        train_accs.append(epoch_acc.item())\n\n        print(f\"Train Loss: {epoch_loss:.4f} Acc: {epoch_acc:.4f}\")\n\n        # ---- Validation phase ----\n        model.eval()\n        running_loss, running_corrects, total = 0.0, 0, 0\n\n        with torch.no_grad():\n            for inputs, labels in tqdm(val_loader, desc=\"Validating\", leave=False):\n                inputs, labels = inputs.to(device), labels.to(device)\n\n                outputs = model(inputs)\n                loss = criterion(outputs, labels)\n\n                _, preds = torch.max(outputs, 1)\n                running_loss += loss.item() * inputs.size(0)\n                running_corrects += torch.sum(preds == labels.data)\n                total += labels.size(0)\n\n        val_loss = running_loss / total\n        val_acc = running_corrects.double() / total\n        val_losses.append(val_loss)\n        val_accs.append(val_acc.item())\n\n        print(f\"Val Loss: {val_loss:.4f} Acc: {val_acc:.4f}\")\n\n        # ---- Scheduler step ----\n        if scheduler:\n            scheduler.step()\n\n        # ---- Deep copy best model ----\n        if val_acc > best_acc:\n            best_acc = val_acc\n            best_model_wts = model.state_dict()\n\n    time_elapsed = time.time() - since\n    print(f\"\\nTraining complete in {time_elapsed//60:.0f}m {time_elapsed%60:.0f}s\")\n    print(f\"Best val Acc: {best_acc:.4f}\")\n\n    # Load best weights\n    model.load_state_dict(best_model_wts)\n\n    return model, (train_losses, val_losses, train_accs, val_accs)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:55.438794Z","iopub.execute_input":"2025-12-16T09:07:55.439308Z","iopub.status.idle":"2025-12-16T09:07:55.454100Z","shell.execute_reply.started":"2025-12-16T09:07:55.439284Z","shell.execute_reply":"2025-12-16T09:07:55.453357Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train the ResNet50 model\nresnet50, history = train_model(\n    model=resnet50,\n    train_loader=train_loader,\n    val_loader=val_loader,\n    criterion=criterion,\n    optimizer=optimizer,\n    scheduler=scheduler,\n    num_epochs=30,\n    device=device\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:07:55.454813Z","iopub.execute_input":"2025-12-16T09:07:55.455027Z","iopub.status.idle":"2025-12-16T09:29:39.731027Z","shell.execute_reply.started":"2025-12-16T09:07:55.455011Z","shell.execute_reply":"2025-12-16T09:29:39.730385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Unpack history\ntrain_losses, val_losses, train_accs, val_accs = history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:29:39.731715Z","iopub.execute_input":"2025-12-16T09:29:39.731914Z","iopub.status.idle":"2025-12-16T09:29:39.735710Z","shell.execute_reply.started":"2025-12-16T09:29:39.731897Z","shell.execute_reply":"2025-12-16T09:29:39.735035Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Loss curve\nplt.figure(figsize=(8,5))\nplt.plot(train_losses, label=\"Train Loss\")\nplt.plot(val_losses, label=\"Val Loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training vs Validation Loss\")\nplt.legend()\nplt.show()\n\n# Accuracy curve\nplt.figure(figsize=(8,5))\nplt.plot(train_accs, label=\"Train Accuracy\")\nplt.plot(val_accs, label=\"Val Accuracy\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy\")\nplt.title(\"Training vs Validation Accuracy\")\nplt.legend()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:29:39.736390Z","iopub.execute_input":"2025-12-16T09:29:39.736611Z","iopub.status.idle":"2025-12-16T09:29:40.085483Z","shell.execute_reply.started":"2025-12-16T09:29:39.736595Z","shell.execute_reply":"2025-12-16T09:29:40.084821Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Evaluation","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef evaluate_model(model, test_loader, device=\"cuda\", save_cm_path=\"confusion_matrix.png\", class_names=None):\n    model.eval()\n    all_logits, all_probs, all_preds, all_labels = [], [], [], []\n\n    for inputs, labels in test_loader:\n        inputs = inputs.to(device)\n        labels = labels.to(device)\n\n        logits = model(inputs)                            # [B, 5]\n        probs = torch.softmax(logits, dim=1)              # [B, 5]\n        preds = probs.argmax(dim=1)                       # [B]\n\n        all_logits.append(logits.cpu())\n        all_probs.append(probs.cpu())\n        all_preds.append(preds.cpu())\n        all_labels.append(labels.cpu())\n\n    all_logits = torch.cat(all_logits).numpy()\n    all_probs  = torch.cat(all_probs).numpy()\n    all_preds  = torch.cat(all_preds).numpy()\n    all_labels = torch.cat(all_labels).numpy()\n\n    # --- Metrics ---\n    acc = accuracy_score(all_labels, all_preds)\n    f1m = f1_score(all_labels, all_preds, average=\"macro\")\n\n    # One-vs-Rest macro AUC (requires prob estimates)\n    try:\n        auc_macro = roc_auc_score(all_labels, all_probs, multi_class=\"ovr\", average=\"macro\")\n    except ValueError:\n        auc_macro = float(\"nan\")  # if a class missing in test, AUC may fail\n\n    # Quadratic Weighted Kappa\n    qwk = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n\n    # Per-class AUC (OvR)\n    per_class_auc = {}\n    num_classes = all_probs.shape[1]\n    for c in range(num_classes):\n        y_true = (all_labels == c).astype(int)\n        y_score = all_probs[:, c]\n        try:\n            per_class_auc[c] = roc_auc_score(y_true, y_score)\n        except ValueError:\n            per_class_auc[c] = float(\"nan\")\n\n    # Classification report\n    print(\"\\n=== Test Metrics ===\")\n    print(f\"Accuracy: {acc:.4f}\")\n    print(f\"Macro F1: {f1m:.4f}\")\n    print(f\"Macro AUC (OvR): {auc_macro:.4f}\")\n    print(f\"Quadratic Weighted Kappa: {qwk:.4f}\")\n\n    # Optional: pretty per-class names\n    if class_names is None:\n        class_names = [str(i) for i in range(num_classes)]\n\n    print(\"\\n=== Classification Report ===\")\n    print(classification_report(all_labels, all_preds, target_names=class_names, digits=4))\n\n    # Confusion matrix\n    cm = confusion_matrix(all_labels, all_preds, labels=list(range(num_classes)))\n    plt.figure(figsize=(6,5))\n    sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\",\n                xticklabels=class_names, yticklabels=class_names)\n    plt.xlabel(\"Predicted\")\n    plt.ylabel(\"True\")\n    plt.title(\"Confusion Matrix (APTOS Test)\")\n    plt.tight_layout()\n    plt.savefig(save_cm_path, dpi=200)\n    plt.show()\n\n    # Return everything you’ll need for uncertainty/calibration later\n    results = {\n        \"acc\": acc,\n        \"f1_macro\": f1m,\n        \"auc_macro\": auc_macro,\n        \"qwk\": qwk,\n        \"per_class_auc\": per_class_auc,\n        \"labels\": all_labels,\n        \"preds\": all_preds,\n        \"probs\": all_probs,\n        \"logits\": all_logits,\n        \"confusion_matrix\": cm\n    }\n    return results","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:29:40.086405Z","iopub.execute_input":"2025-12-16T09:29:40.086736Z","iopub.status.idle":"2025-12-16T09:29:40.096421Z","shell.execute_reply.started":"2025-12-16T09:29:40.086719Z","shell.execute_reply":"2025-12-16T09:29:40.095741Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_names = [\"No DR (0)\", \"Mild (1)\", \"Moderate (2)\", \"Severe (3)\", \"Proliferative (4)\"]\ntest_results = evaluate_model(resnet50, test_loader, device=device, class_names=class_names)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:29:40.100054Z","iopub.execute_input":"2025-12-16T09:29:40.100243Z","iopub.status.idle":"2025-12-16T09:29:42.445316Z","shell.execute_reply.started":"2025-12-16T09:29:40.100228Z","shell.execute_reply":"2025-12-16T09:29:42.444683Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Uncertainity Estimation with Monte Carlo Dropout","metadata":{}},{"cell_type":"code","source":"def enable_dropout(model):\n    \"\"\"Enable dropout layers during test-time\"\"\"\n    for m in model.modules():\n        if m.__class__.__name__.startswith('Dropout'):\n            m.train()\n\n@torch.no_grad()\ndef mc_dropout_predictions_all(model, dataloader, device=\"cuda\", T=30):\n    model.eval()\n    enable_dropout(model)  # keep dropout active\n\n    all_probs = []\n    all_labels = []\n\n    for inputs, labels in dataloader:\n        inputs, labels = inputs.to(device), labels.to(device)\n        batch_probs = []\n\n        for _ in range(T):\n            outputs = model(inputs)\n            probs = torch.softmax(outputs, dim=1)\n            batch_probs.append(probs.unsqueeze(0))  # [1,B,C]\n\n        batch_probs = torch.cat(batch_probs, dim=0)  # [T,B,C]\n        all_probs.append(batch_probs.cpu().numpy())\n        all_labels.append(labels.cpu().numpy())\n\n    all_probs = np.concatenate(all_probs, axis=1)  # [T,N,C]\n    all_labels = np.concatenate(all_labels, axis=0)\n\n    # Predictive mean\n    mean_probs = all_probs.mean(axis=0)  # [N,C]\n\n    # --- Uncertainty Measures ---\n    # 1. Entropy\n    entropy = -np.sum(mean_probs * np.log(mean_probs + 1e-8), axis=1)\n\n    # 2. Variance\n    variance = all_probs.var(axis=0).mean(axis=1)\n\n    # 3. Mutual Information (MI)\n    # Predictive entropy - expected entropy\n    expected_entropy = -np.mean(np.sum(all_probs * np.log(all_probs + 1e-8), axis=2), axis=0)\n    MI = entropy - expected_entropy\n\n    # 4. Max Probability (Confidence Score)\n    max_prob = mean_probs.max(axis=1)\n\n    preds = mean_probs.argmax(axis=1)\n    return {\n        \"mean_probs\": mean_probs,\n        \"entropy\": entropy,\n        \"variance\": variance,\n        \"mutual_info\": MI,\n        \"max_prob\": max_prob,\n        \"preds\": preds,\n        \"labels\": all_labels\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:29:42.446116Z","iopub.execute_input":"2025-12-16T09:29:42.446393Z","iopub.status.idle":"2025-12-16T09:29:42.455299Z","shell.execute_reply.started":"2025-12-16T09:29:42.446369Z","shell.execute_reply":"2025-12-16T09:29:42.454662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_uncertainty_all(results):\n    labels = results[\"labels\"]\n    preds = results[\"preds\"]\n\n    # Binary: 1 = error, 0 = correct\n    errors = (preds != labels).astype(int)\n\n    metrics = {}\n    for name, scores in {\n        \"Entropy\": results[\"entropy\"],\n        \"Variance\": results[\"variance\"],\n        \"Mutual Information\": results[\"mutual_info\"],\n        \"Max Probability\": -results[\"max_prob\"]  # invert since low prob = high uncertainty\n    }.items():\n        try:\n            auroc = roc_auc_score(errors, scores)\n            metrics[name] = auroc\n            print(f\"{name} AUROC: {auroc:.4f}\")\n        except ValueError:\n            metrics[name] = None\n            print(f\"{name}: could not compute AUROC (maybe missing error cases).\")\n\n    return metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:29:42.456197Z","iopub.execute_input":"2025-12-16T09:29:42.456439Z","iopub.status.idle":"2025-12-16T09:29:42.470560Z","shell.execute_reply.started":"2025-12-16T09:29:42.456411Z","shell.execute_reply":"2025-12-16T09:29:42.470080Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Run MC Dropout with all uncertainty measures\nmc_results_all = mc_dropout_predictions_all(resnet50, test_loader, device=device, T=30)\n\n# Evaluate AUROC for all measures\nunc_metrics = evaluate_uncertainty_all(mc_results_all)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:29:42.471167Z","iopub.execute_input":"2025-12-16T09:29:42.471371Z","iopub.status.idle":"2025-12-16T09:30:20.923376Z","shell.execute_reply.started":"2025-12-16T09:29:42.471348Z","shell.execute_reply":"2025-12-16T09:30:20.922498Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Histogram of Uncertainity (Correct vs Incorrect)","metadata":{}},{"cell_type":"code","source":"def plot_uncertainty_hist(results, metric=\"entropy\", bins=30):\n    labels = results[\"labels\"]\n    preds = results[\"preds\"]\n    errors = (preds != labels).astype(int)\n\n    scores = results[metric]\n\n    correct_scores = scores[errors == 0]\n    wrong_scores = scores[errors == 1]\n\n    plt.figure(figsize=(7,5))\n    sns.histplot(correct_scores, bins=bins, color=\"green\", alpha=0.6, label=\"Correct\")\n    sns.histplot(wrong_scores, bins=bins, color=\"red\", alpha=0.6, label=\"Wrong\")\n    plt.xlabel(f\"{metric.capitalize()} Score\")\n    plt.ylabel(\"Count\")\n    plt.title(f\"Histogram of {metric.capitalize()} for Correct vs Incorrect Predictions\")\n    plt.legend()\n    plt.tight_layout()\n    plt.show()\n\n# Example: plot entropy histogram\nplot_uncertainty_hist(mc_results_all, metric=\"entropy\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:30:20.924323Z","iopub.execute_input":"2025-12-16T09:30:20.924637Z","iopub.status.idle":"2025-12-16T09:30:21.232506Z","shell.execute_reply.started":"2025-12-16T09:30:20.924607Z","shell.execute_reply":"2025-12-16T09:30:21.231853Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### ROC Curves for Error Detection","metadata":{}},{"cell_type":"code","source":"def plot_uncertainty_rocs(results):\n    labels = results[\"labels\"]\n    preds = results[\"preds\"]\n    errors = (preds != labels).astype(int)\n\n    plt.figure(figsize=(7,6))\n\n    metrics_to_plot = {\n        \"Entropy\": results[\"entropy\"],\n        \"Variance\": results[\"variance\"],\n        \"Mutual Information\": results[\"mutual_info\"],\n        \"Max Probability\": -results[\"max_prob\"]  # invert: low prob = uncertain\n    }\n\n    for name, scores in metrics_to_plot.items():\n        fpr, tpr, _ = roc_curve(errors, scores)\n        roc_auc = auc(fpr, tpr)\n        plt.plot(fpr, tpr, label=f\"{name} (AUROC = {roc_auc:.2f})\")\n\n    plt.plot([0,1],[0,1], \"k--\", label=\"Random (0.5)\")\n    plt.xlabel(\"False Positive Rate\")\n    plt.ylabel(\"True Positive Rate\")\n    plt.title(\"ROC Curves for Error Detection via Uncertainty\")\n    plt.legend(loc=\"lower right\")\n    plt.tight_layout()\n    plt.show()\n\n# Example: plot ROC curves\nplot_uncertainty_rocs(mc_results_all)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:30:21.233184Z","iopub.execute_input":"2025-12-16T09:30:21.233426Z","iopub.status.idle":"2025-12-16T09:30:21.451470Z","shell.execute_reply.started":"2025-12-16T09:30:21.233400Z","shell.execute_reply":"2025-12-16T09:30:21.450881Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_true = mc_results_all[\"labels\"]\ny_pred = mc_results_all[\"preds\"]\nentropy = mc_results_all[\"entropy\"]\nvariance = mc_results_all[\"variance\"]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:33:07.852193Z","iopub.execute_input":"2025-12-16T09:33:07.852974Z","iopub.status.idle":"2025-12-16T09:33:07.857115Z","shell.execute_reply.started":"2025-12-16T09:33:07.852949Z","shell.execute_reply":"2025-12-16T09:33:07.856389Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def risk_coverage_curve(y_true, y_pred, uncertainty_scores):\n    y_true = np.array(y_true)\n    y_pred = np.array(y_pred)\n    uncertainty_scores = np.array(uncertainty_scores)\n\n    # Sort by uncertainty (descending = reject first)\n    idx = np.argsort(uncertainty_scores)[::-1]\n\n    y_true_sorted = y_true[idx]\n    y_pred_sorted = y_pred[idx]\n\n    n = len(y_true)\n    coverages = []\n    accuracies = []\n\n    for k in range(n):\n        y_true_kept = y_true_sorted[k:]\n        y_pred_kept = y_pred_sorted[k:]\n\n        if len(y_true_kept) == 0:\n            break\n\n        acc = np.mean(y_true_kept == y_pred_kept)\n        cov = len(y_true_kept) / n\n\n        accuracies.append(acc)\n        coverages.append(cov)\n\n    return np.array(coverages), np.array(accuracies)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:33:27.574743Z","iopub.execute_input":"2025-12-16T09:33:27.574986Z","iopub.status.idle":"2025-12-16T09:33:27.580655Z","shell.execute_reply.started":"2025-12-16T09:33:27.574969Z","shell.execute_reply":"2025-12-16T09:33:27.579822Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cov_e, acc_e = risk_coverage_curve(y_true, y_pred, entropy)\n\nplt.figure(figsize=(7,5))\nplt.plot(cov_e, acc_e, linewidth=2)\nplt.xlabel(\"Coverage (Fraction of Samples Kept)\")\nplt.ylabel(\"Accuracy on Kept Samples\")\nplt.title(\"Risk–Coverage Curve using Predictive Entropy\")\nplt.grid(True)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:35:44.610988Z","iopub.execute_input":"2025-12-16T09:35:44.611563Z","iopub.status.idle":"2025-12-16T09:35:44.805020Z","shell.execute_reply.started":"2025-12-16T09:35:44.611540Z","shell.execute_reply":"2025-12-16T09:35:44.804367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cov_v, acc_v = risk_coverage_curve(y_true, y_pred, variance)\n\nplt.figure(figsize=(7,5))\nplt.plot(cov_e, acc_e, label=\"Entropy\", linewidth=2)\nplt.plot(cov_v, acc_v, label=\"Variance\", linewidth=2)\n\nplt.xlabel(\"Coverage\")\nplt.ylabel(\"Accuracy\")\nplt.title(\"Risk–Coverage Comparison of Uncertainty Measures\")\nplt.legend()\nplt.grid(True)\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:35:39.198314Z","iopub.execute_input":"2025-12-16T09:35:39.199083Z","iopub.status.idle":"2025-12-16T09:35:39.418755Z","shell.execute_reply.started":"2025-12-16T09:35:39.199054Z","shell.execute_reply":"2025-12-16T09:35:39.418063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target_acc = 0.95\n\nvalid_idx = np.where(acc_e >= target_acc)[0]\n\nif len(valid_idx) > 0:\n    i = valid_idx[0]\n    print(f\"Target Accuracy: {acc_e[i]:.3f}\")\n    print(f\"Coverage: {cov_e[i]:.3f}\")\n    print(f\"Rejection Rate: {1 - cov_e[i]:.3f}\")\nelse:\n    print(\"Target accuracy not achievable.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:36:14.206640Z","iopub.execute_input":"2025-12-16T09:36:14.207242Z","iopub.status.idle":"2025-12-16T09:36:14.212087Z","shell.execute_reply.started":"2025-12-16T09:36:14.207219Z","shell.execute_reply":"2025-12-16T09:36:14.211307Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\n\n# Convert coverage to rejection rate\nrejection = 1 - cov_e\n\n# Optional: smooth accuracy for readability\ndef smooth(y, window=5):\n    return np.convolve(y, np.ones(window)/window, mode=\"same\")\n\nacc_smooth = smooth(acc_e, window=7)\n\nplt.figure(figsize=(8, 5))\n\n# Main curve\nplt.plot(\n    rejection,\n    acc_smooth,\n    linewidth=2.5,\n    label=\"Entropy-based Referral\"\n)\n\n# Baseline accuracy point\nplt.scatter(\n    0,\n    acc_e[-1],\n    color=\"red\",\n    zorder=5,\n    label=\"No Referral (Baseline Accuracy)\"\n)\n\n# Clinical safety thresholds\nplt.axhline(0.95, linestyle=\"--\", color=\"green\", alpha=0.7, label=\"95% Accuracy Target\")\nplt.axhline(0.98, linestyle=\"--\", color=\"purple\", alpha=0.7, label=\"98% Accuracy Target\")\n\nplt.xlabel(\"Rejection Rate (Fraction of Samples Referred)\")\nplt.ylabel(\"Accuracy on Retained Samples\")\nplt.title(\"Risk–Coverage Curve (Accuracy–Rejection) using Predictive Entropy\")\n\nplt.legend()\nplt.grid(alpha=0.3)\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:37:10.873973Z","iopub.execute_input":"2025-12-16T09:37:10.874662Z","iopub.status.idle":"2025-12-16T09:37:11.122309Z","shell.execute_reply.started":"2025-12-16T09:37:10.874638Z","shell.execute_reply":"2025-12-16T09:37:11.121566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def risk_coverage_curve_with_mask(y_true, y_pred, uncertainty, mask=None):\n    y_true = np.array(y_true)\n    y_pred = np.array(y_pred)\n    uncertainty = np.array(uncertainty)\n\n    if mask is None:\n        mask = np.ones_like(y_true, dtype=bool)\n\n    y_true = y_true[mask]\n    y_pred = y_pred[mask]\n    uncertainty = uncertainty[mask]\n\n    idx = np.argsort(uncertainty)[::-1]  # high uncertainty rejected first\n    y_true = y_true[idx]\n    y_pred = y_pred[idx]\n\n    n = len(y_true)\n    coverages, accuracies = [], []\n\n    for k in range(n):\n        yt = y_true[k:]\n        yp = y_pred[k:]\n\n        if len(yt) == 0:\n            break\n\n        coverages.append(len(yt) / n)\n        accuracies.append(np.mean(yt == yp))\n\n    return np.array(coverages), np.array(accuracies)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:37:59.829702Z","iopub.execute_input":"2025-12-16T09:37:59.829979Z","iopub.status.idle":"2025-12-16T09:37:59.835907Z","shell.execute_reply.started":"2025-12-16T09:37:59.829957Z","shell.execute_reply":"2025-12-16T09:37:59.835283Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEVERE_CLASS = 3\n\nmask_severe = (mc_results_all[\"labels\"] == SEVERE_CLASS)\n\ncov_s, acc_s = risk_coverage_curve_with_mask(\n    mc_results_all[\"labels\"],\n    mc_results_all[\"preds\"],\n    mc_results_all[\"entropy\"],\n    mask=mask_severe\n)\n\nplt.figure(figsize=(7,5))\nplt.plot(1 - cov_s, acc_s, linewidth=2)\nplt.xlabel(\"Rejection Rate\")\nplt.ylabel(\"Accuracy on Severe DR\")\nplt.title(\"Class-wise Risk–Coverage (Severe DR)\")\nplt.grid(alpha=0.3)\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:40:14.740433Z","iopub.execute_input":"2025-12-16T09:40:14.741085Z","iopub.status.idle":"2025-12-16T09:40:14.910902Z","shell.execute_reply.started":"2025-12-16T09:40:14.741061Z","shell.execute_reply":"2025-12-16T09:40:14.910168Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"entropy_scores = mc_results_all[\"entropy\"]\nmaxprob_uncertainty = -mc_results_all[\"max_prob\"]  # invert\n\ncov_e, acc_e = risk_coverage_curve_with_mask(\n    mc_results_all[\"labels\"],\n    mc_results_all[\"preds\"],\n    entropy_scores\n)\n\ncov_m, acc_m = risk_coverage_curve_with_mask(\n    mc_results_all[\"labels\"],\n    mc_results_all[\"preds\"],\n    maxprob_uncertainty\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:40:55.219585Z","iopub.execute_input":"2025-12-16T09:40:55.220070Z","iopub.status.idle":"2025-12-16T09:40:55.230871Z","shell.execute_reply.started":"2025-12-16T09:40:55.220046Z","shell.execute_reply":"2025-12-16T09:40:55.230063Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(8,5))\nplt.plot(1 - cov_e, acc_e, label=\"Entropy\", linewidth=2)\nplt.plot(1 - cov_m, acc_m, label=\"Max Probability\", linewidth=2)\n\nplt.xlabel(\"Rejection Rate\")\nplt.ylabel(\"Accuracy on Retained Samples\")\nplt.title(\"Risk–Coverage Comparison: Entropy vs Max Probability\")\nplt.legend()\nplt.grid(alpha=0.3)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:41:05.092767Z","iopub.execute_input":"2025-12-16T09:41:05.093682Z","iopub.status.idle":"2025-12-16T09:41:05.310984Z","shell.execute_reply.started":"2025-12-16T09:41:05.093651Z","shell.execute_reply":"2025-12-16T09:41:05.310340Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def coverage_vs_sensitivity(y_true, y_pred, uncertainty, positive_class):\n    y_true = np.array(y_true)\n    y_pred = np.array(y_pred)\n    uncertainty = np.array(uncertainty)\n\n    idx = np.argsort(uncertainty)[::-1]\n    y_true = y_true[idx]\n    y_pred = y_pred[idx]\n\n    n = len(y_true)\n    coverages, recalls = [], []\n\n    for k in range(n):\n        yt = y_true[k:]\n        yp = y_pred[k:]\n\n        if len(yt) == 0:\n            break\n\n        mask = (yt == positive_class) | (yp == positive_class)\n        if mask.sum() == 0:\n            continue\n\n        recall = recall_score(\n            yt == positive_class,\n            yp == positive_class,\n            zero_division=0\n        )\n\n        coverages.append(len(yt) / n)\n        recalls.append(recall)\n\n    return np.array(coverages), np.array(recalls)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:41:28.193960Z","iopub.execute_input":"2025-12-16T09:41:28.194674Z","iopub.status.idle":"2025-12-16T09:41:28.200106Z","shell.execute_reply.started":"2025-12-16T09:41:28.194648Z","shell.execute_reply":"2025-12-16T09:41:28.199464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import recall_score\n\ncov_r, rec_r = coverage_vs_sensitivity(\n    mc_results_all[\"labels\"],\n    mc_results_all[\"preds\"],\n    mc_results_all[\"entropy\"],\n    positive_class=SEVERE_CLASS\n)\n\nplt.figure(figsize=(7,5))\nplt.plot(1 - cov_r, rec_r, linewidth=2)\nplt.xlabel(\"Rejection Rate\")\nplt.ylabel(\"Sensitivity (Recall) for Severe DR\")\nplt.title(\"Coverage vs Sensitivity (Severe DR)\")\nplt.grid(alpha=0.3)\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:42:18.328622Z","iopub.execute_input":"2025-12-16T09:42:18.328908Z","iopub.status.idle":"2025-12-16T09:42:18.658221Z","shell.execute_reply.started":"2025-12-16T09:42:18.328886Z","shell.execute_reply":"2025-12-16T09:42:18.657550Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import matplotlib.pyplot as plt\n# import numpy as np\n# import cv2\n\n# def show_case_studies_with_original(results, dataset, original_paths, class_names, metric=\"entropy\"):\n#     \"\"\"\n#     Show case studies with both original and preprocessed images.\n\n#     Args:\n#         results: dict from mc_dropout_predictions_all\n#         dataset: PyTorch dataset (returns preprocessed image, label)\n#         original_paths: list of file paths to original fundus images (aligned with dataset order)\n#         class_names: list of class names (e.g. [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative\"])\n#         metric: uncertainty metric to display (\"entropy\", \"max_prob\", etc.)\n#     \"\"\"\n\n#     preds = results[\"preds\"]\n#     labels = results[\"labels\"]\n#     scores = results[metric]\n\n#     # Pick indices for 3 case types\n#     idx_correct_confident = np.argmin(scores + (preds != labels)*10)  # correct + confident\n#     idx_wrong_uncertain = np.argmax(scores * (preds != labels))      # wrong + uncertain\n#     idx_wrong_confident = np.argmin(scores + (preds == labels)*10)   # wrong + confident\n#     chosen_idxs = [idx_correct_confident, idx_wrong_uncertain, idx_wrong_confident]\n\n#     plt.figure(figsize=(12, 8))\n\n#     for i, idx in enumerate(chosen_idxs):\n#         # --- Load original image ---\n#         orig_img = cv2.imread(original_paths[idx])\n#         orig_img = cv2.cvtColor(orig_img, cv2.COLOR_BGR2RGB)\n\n#         # --- Load preprocessed image (from dataset) ---\n#         img_tensor, _ = dataset[idx]  # preprocessed tensor\n#         img = img_tensor.permute(1, 2, 0).numpy()\n#         img = (img - img.min()) / (img.max() - img.min())  # normalize for display\n\n#         true_label = class_names[labels[idx]]\n#         pred_label = class_names[preds[idx]]\n#         score = scores[idx]\n\n#         # Plot original (top row)\n#         plt.subplot(2, 3, i+1)\n#         plt.imshow(orig_img)\n#         plt.axis(\"off\")\n#         plt.title(f\"Original\\nTrue: {true_label}\\nPred: {pred_label}\\n{metric}: {score:.3f}\")\n\n#         # Plot preprocessed (bottom row)\n#         plt.subplot(2, 3, i+4)\n#         plt.imshow(img)\n#         plt.axis(\"off\")\n#         plt.title(\"Preprocessed\")\n\n#     plt.tight_layout()\n#     plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:30:21.500474Z","iopub.status.idle":"2025-12-16T09:30:21.500746Z","shell.execute_reply.started":"2025-12-16T09:30:21.500639Z","shell.execute_reply":"2025-12-16T09:30:21.500649Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# class_names = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative\"]\n\n# # Example: using entropy\n# show_case_studies_with_original(\n#     mc_results_all,\n#     test_dataset,\n#     test_img_paths,      # list of paths to original fundus images\n#     class_names,\n#     metric=\"entropy\"\n# )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:30:21.501483Z","iopub.status.idle":"2025-12-16T09:30:21.501758Z","shell.execute_reply.started":"2025-12-16T09:30:21.501642Z","shell.execute_reply":"2025-12-16T09:30:21.501655Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"1st Figure represents -> A Correct and Confident Case\\\n2nd Figure represents -> A Wrong and Uncertain Case\\\n3rd Figure represents -> A Wrong and Confident Case","metadata":{}},{"cell_type":"code","source":"# show_case_studies(mc_results_all, test_dataset, metric=\"max_prob\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:30:21.503256Z","iopub.status.idle":"2025-12-16T09:30:21.503457Z","shell.execute_reply.started":"2025-12-16T09:30:21.503359Z","shell.execute_reply":"2025-12-16T09:30:21.503368Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Calibration","metadata":{}},{"cell_type":"code","source":"class ModelWithTemperature(nn.Module):\n    def __init__(self, model):\n        super(ModelWithTemperature, self).__init__()\n        self.model = model\n        self.temperature = nn.Parameter(torch.ones(1) * 1.5)\n\n    def forward(self, input):\n        logits = self.model(input)\n        return self.temperature_scale(logits)\n\n    def temperature_scale(self, logits):\n        # logits: [N, C]\n        temperature = self.temperature.unsqueeze(1).expand(logits.size(0), logits.size(1))\n        return logits / temperature\n\n    def set_temperature(self, valid_loader, device=\"cuda\"):\n        self.to(device)\n        nll_criterion = nn.CrossEntropyLoss().to(device)\n\n        logits_list, labels_list = [], []\n        with torch.no_grad():\n            for inputs, labels in valid_loader:\n                inputs, labels = inputs.to(device), labels.to(device)\n                logits = self.model(inputs)\n                logits_list.append(logits)\n                labels_list.append(labels)\n        logits = torch.cat(logits_list).to(device)\n        labels = torch.cat(labels_list).to(device)\n\n        optimizer = optim.LBFGS([self.temperature], lr=0.01, max_iter=50)\n\n        def eval():\n            optimizer.zero_grad()\n            loss = nll_criterion(self.temperature_scale(logits), labels)\n            loss.backward()\n            return loss\n\n        optimizer.step(eval)\n        print(f\"Optimal temperature: {self.temperature.item():.3f}\")\n        return self","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:43:08.004243Z","iopub.execute_input":"2025-12-16T09:43:08.004489Z","iopub.status.idle":"2025-12-16T09:43:08.011685Z","shell.execute_reply.started":"2025-12-16T09:43:08.004472Z","shell.execute_reply":"2025-12-16T09:43:08.010926Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Expected Calibration Error (ECE)","metadata":{}},{"cell_type":"code","source":"def expected_calibration_error(probs, labels, n_bins=15):\n    # probs: torch tensor [N, C], labels: torch tensor [N]\n    confidences, predictions = probs.max(dim=1)  # unpack values & indices\n    accuracies = predictions.eq(labels)\n\n    ece = torch.zeros(1, device=probs.device)\n    bin_boundaries = torch.linspace(0, 1, n_bins + 1, device=probs.device)\n\n    for i in range(n_bins):\n        start, end = bin_boundaries[i], bin_boundaries[i+1]\n        mask = (confidences > start) & (confidences <= end)\n        if mask.sum() > 0:\n            bin_acc = accuracies[mask].float().mean()\n            bin_conf = confidences[mask].mean()\n            ece += (mask.float().mean()) * torch.abs(bin_acc - bin_conf)\n\n    return ece.item()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:43:11.738824Z","iopub.execute_input":"2025-12-16T09:43:11.739507Z","iopub.status.idle":"2025-12-16T09:43:11.744650Z","shell.execute_reply.started":"2025-12-16T09:43:11.739483Z","shell.execute_reply":"2025-12-16T09:43:11.743895Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Maximum Calibration Error","metadata":{}},{"cell_type":"code","source":"def maximum_calibration_error(probs, labels, n_bins=15):\n    \"\"\"\n    probs: torch.tensor [N, C]\n    labels: torch.tensor [N]\n    \"\"\"\n    confidences, predictions = probs.max(dim=1)\n    accuracies = predictions.eq(labels)\n\n    mce = torch.zeros(1, device=probs.device)\n    bin_boundaries = torch.linspace(0, 1, n_bins + 1, device=probs.device)\n\n    for i in range(n_bins):\n        start, end = bin_boundaries[i], bin_boundaries[i+1]\n        mask = (confidences > start) & (confidences <= end)\n        if mask.sum() > 0:\n            bin_acc = accuracies[mask].float().mean()\n            bin_conf = confidences[mask].mean()\n            mce = torch.max(mce, torch.abs(bin_acc - bin_conf))\n\n    return mce.item()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:43:15.273348Z","iopub.execute_input":"2025-12-16T09:43:15.273603Z","iopub.status.idle":"2025-12-16T09:43:15.278814Z","shell.execute_reply.started":"2025-12-16T09:43:15.273586Z","shell.execute_reply":"2025-12-16T09:43:15.278078Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Class wise ECE","metadata":{}},{"cell_type":"code","source":"def classwise_ece(probs, labels, n_bins=15, num_classes=5):\n    \"\"\"\n    Computes per-class ECE\n    probs: torch.tensor [N, C]\n    labels: torch.tensor [N]\n    \"\"\"\n    class_ece = {}\n    for c in range(num_classes):\n        # Consider binary correct vs incorrect for class c\n        confidences = probs[:, c]\n        predictions = (probs.argmax(dim=1) == c).long()\n        accuracies = (labels == c)\n\n        ece = torch.zeros(1, device=probs.device)\n        bin_boundaries = torch.linspace(0, 1, n_bins + 1, device=probs.device)\n\n        for i in range(n_bins):\n            start, end = bin_boundaries[i], bin_boundaries[i+1]\n            mask = (confidences > start) & (confidences <= end)\n            if mask.sum() > 0:\n                bin_acc = accuracies[mask].float().mean()\n                bin_conf = confidences[mask].mean()\n                ece += (mask.float().mean()) * torch.abs(bin_acc - bin_conf)\n\n        class_ece[c] = ece.item()\n\n    return class_ece","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:43:19.108423Z","iopub.execute_input":"2025-12-16T09:43:19.108703Z","iopub.status.idle":"2025-12-16T09:43:19.115292Z","shell.execute_reply.started":"2025-12-16T09:43:19.108681Z","shell.execute_reply":"2025-12-16T09:43:19.114674Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Brier Score","metadata":{}},{"cell_type":"code","source":"def brier_score(probs, labels):\n    labels_onehot = np.zeros_like(probs)\n    labels_onehot[np.arange(len(labels)), labels] = 1\n    return np.mean(np.sum((probs - labels_onehot)**2, axis=1))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:43:23.329844Z","iopub.execute_input":"2025-12-16T09:43:23.330501Z","iopub.status.idle":"2025-12-16T09:43:23.334142Z","shell.execute_reply.started":"2025-12-16T09:43:23.330473Z","shell.execute_reply":"2025-12-16T09:43:23.333559Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Reliability Diagram","metadata":{}},{"cell_type":"code","source":"def plot_reliability_diagram(probs, labels, n_bins=10):\n    \"\"\"\n    probs: torch.tensor [N, C] (logits after softmax)\n    labels: torch.tensor [N]\n    \"\"\"\n    # Ensure CPU + NumPy for plotting\n    probs = probs.detach().cpu()\n    labels = labels.detach().cpu()\n\n    confidences, predictions = probs.max(dim=1)\n    accuracies = predictions.eq(labels)\n\n    bin_boundaries = np.linspace(0.0, 1.0, n_bins + 1)\n    bin_accs, bin_confs = [], []\n\n    for i in range(n_bins):\n        start, end = bin_boundaries[i], bin_boundaries[i+1]\n        mask = (confidences > start) & (confidences <= end)\n        if mask.sum() > 0:\n            bin_acc = accuracies[mask].float().mean().item()\n            bin_conf = confidences[mask].mean().item()\n            bin_accs.append(bin_acc)\n            bin_confs.append(bin_conf)\n\n    # Plot\n    plt.figure(figsize=(6,6))\n    plt.plot([0,1],[0,1],\"k--\")\n    plt.plot(bin_confs, bin_accs, marker=\"o\", label=\"Model\")\n    plt.xlabel(\"Predicted Confidence\")\n    plt.ylabel(\"True Accuracy\")\n    plt.title(\"Reliability Diagram\")\n    plt.legend()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:43:29.026825Z","iopub.execute_input":"2025-12-16T09:43:29.027646Z","iopub.status.idle":"2025-12-16T09:43:29.034676Z","shell.execute_reply.started":"2025-12-16T09:43:29.027619Z","shell.execute_reply":"2025-12-16T09:43:29.034075Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = \"cuda\"\n\n# Wrap model with temperature scaling\nscaled_model = ModelWithTemperature(resnet50).set_temperature(val_loader, device=device)\n\n# Collect test logits & probs (before and after scaling)\nresnet50.eval()\nscaled_model.eval()\nlogits_list, labels_list = [], []\nwith torch.no_grad():\n    for inputs, labels in test_loader:\n        inputs, labels = inputs.to(device), labels.to(device)\n        logits = resnet50(inputs)\n        logits_scaled = scaled_model.temperature_scale(logits)\n\n        logits_list.append(logits)\n        labels_list.append(labels)\n\nlogits = torch.cat(logits_list)\nlabels = torch.cat(labels_list)\nprobs = torch.softmax(logits, dim=1)\nprobs_scaled = torch.softmax(scaled_model.temperature_scale(logits), dim=1)\n\n# Brier Score\nbrier_before = brier_score(probs.cpu().numpy(), labels.cpu().numpy())\nbrier_after = brier_score(probs_scaled.detach().cpu().numpy(), labels.detach().cpu().numpy())\nprint(f\"Brier Score before scaling: {brier_before:.4f}, after: {brier_after:.4f}\")\n\n# Compute ECE, MCE, Class-wise ECE\nece_before = expected_calibration_error(probs, labels)\nmce_before = maximum_calibration_error(probs, labels)\nclass_ece_before = classwise_ece(probs, labels, num_classes=5)\n\nece_after = expected_calibration_error(probs_scaled, labels)\nmce_after = maximum_calibration_error(probs_scaled, labels)\nclass_ece_after = classwise_ece(probs_scaled, labels, num_classes=5)\n\nprint(\"ECE before:\", ece_before, \"after:\", ece_after)\nprint(\"MCE before:\", mce_before, \"after:\", mce_after)\nprint(\"Class-wise ECE before:\", class_ece_before)\nprint(\"Class-wise ECE after:\", class_ece_after)\n\n# Reliability Diagram\nplot_reliability_diagram(probs, labels)\nplot_reliability_diagram(probs_scaled, labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:43:33.729466Z","iopub.execute_input":"2025-12-16T09:43:33.729775Z","iopub.status.idle":"2025-12-16T09:43:36.853052Z","shell.execute_reply.started":"2025-12-16T09:43:33.729753Z","shell.execute_reply":"2025-12-16T09:43:36.852434Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_classwise_reliability(probs, labels, class_names, n_bins=10, title=\"\"):\n    \"\"\"\n    probs: torch.tensor [N, C] (softmax probabilities)\n    labels: torch.tensor [N]\n    class_names: list of class names, length C\n    \"\"\"\n\n    probs = probs.detach().cpu()\n    labels = labels.detach().cpu()\n    num_classes = probs.shape[1]\n\n    fig, axes = plt.subplots(1, num_classes, figsize=(4*num_classes, 4), sharey=True)\n\n    for c in range(num_classes):\n        ax = axes[c]\n        confs = probs[:, c]\n        truths = (labels == c).float()\n\n        bin_boundaries = np.linspace(0.0, 1.0, n_bins + 1)\n        bin_accs, bin_confs = [], []\n\n        for i in range(n_bins):\n            start, end = bin_boundaries[i], bin_boundaries[i+1]\n            mask = (confs > start) & (confs <= end)\n            if mask.sum() > 0:\n                bin_acc = truths[mask].mean().item()\n                bin_conf = confs[mask].mean().item()\n                bin_accs.append(bin_acc)\n                bin_confs.append(bin_conf)\n\n        ax.plot([0,1],[0,1],\"k--\")\n        ax.plot(bin_confs, bin_accs, marker=\"o\", label=class_names[c])\n        ax.set_title(class_names[c])\n        ax.set_xlabel(\"Confidence\")\n        if c == 0:\n            ax.set_ylabel(\"Accuracy\")\n        ax.set_xlim([0,1])\n        ax.set_ylim([0,1])\n\n    plt.suptitle(f\"Class-wise Reliability Diagrams: {title}\")\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:44:08.419633Z","iopub.execute_input":"2025-12-16T09:44:08.420184Z","iopub.status.idle":"2025-12-16T09:44:08.427327Z","shell.execute_reply.started":"2025-12-16T09:44:08.420162Z","shell.execute_reply":"2025-12-16T09:44:08.426578Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_names = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative\"]\n\n# Before scaling\nplot_classwise_reliability(probs, labels, class_names, title=\"Before Temperature Scaling\")\n\n# After scaling\nplot_classwise_reliability(probs_scaled, labels, class_names, title=\"After Temperature Scaling\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:44:26.398580Z","iopub.execute_input":"2025-12-16T09:44:26.398841Z","iopub.status.idle":"2025-12-16T09:44:27.596361Z","shell.execute_reply.started":"2025-12-16T09:44:26.398819Z","shell.execute_reply":"2025-12-16T09:44:27.595640Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_combined_calibration_all(probs, probs_scaled, labels, class_names, n_bins=10):\n    \"\"\"\n    Combined calibration plot:\n    - Top: global reliability before & after scaling\n    - Bottom: class-wise reliability before & after scaling\n    \"\"\"\n    probs = probs.detach().cpu()\n    probs_scaled = probs_scaled.detach().cpu()\n    labels = labels.detach().cpu()\n    num_classes = len(class_names)\n\n    # --- Helper for bin stats ---\n    def compute_bin_stats_from_confidences(confs, truths, n_bins):\n        bin_boundaries = np.linspace(0.0, 1.0, n_bins + 1)\n        bin_accs, bin_confs = [], []\n        for i in range(n_bins):\n            start, end = bin_boundaries[i], bin_boundaries[i+1]\n            mask = (confs > start) & (confs <= end)\n            if mask.sum() > 0:\n                bin_accs.append(truths[mask].float().mean().item())\n                bin_confs.append(confs[mask].mean().item())\n        return bin_confs, bin_accs\n\n    def compute_global_bin_stats(probs, labels, n_bins):\n        confidences, predictions = probs.max(dim=1)\n        accuracies = predictions.eq(labels)\n        return compute_bin_stats_from_confidences(confidences, accuracies, n_bins)\n\n    # --- Global stats ---\n    bin_confs_before, bin_accs_before = compute_global_bin_stats(probs, labels, n_bins)\n    bin_confs_after, bin_accs_after = compute_global_bin_stats(probs_scaled, labels, n_bins)\n\n    fig = plt.figure(figsize=(16, 10))\n\n    # --- Top: Global reliability ---\n    ax1 = plt.subplot2grid((2, num_classes), (0, 0), colspan=num_classes)\n    ax1.plot([0,1],[0,1],\"k--\")\n    ax1.plot(bin_confs_before, bin_accs_before, marker=\"o\", label=\"Before Scaling\")\n    ax1.plot(bin_confs_after, bin_accs_after, marker=\"o\", label=\"After Scaling\")\n    ax1.set_title(\"Global Reliability Diagram\")\n    ax1.set_xlabel(\"Predicted Confidence\")\n    ax1.set_ylabel(\"True Accuracy\")\n    ax1.set_xlim([0,1]); ax1.set_ylim([0,1])\n    ax1.legend()\n\n    # --- Bottom: Class-wise reliability before & after ---\n    for c in range(num_classes):\n        ax = plt.subplot2grid((2, num_classes), (1, c))\n\n        # Before scaling\n        confs_before = probs[:, c]\n        truths = (labels == c).float()\n        bin_confs_b, bin_accs_b = compute_bin_stats_from_confidences(confs_before, truths, n_bins)\n\n        # After scaling\n        confs_after = probs_scaled[:, c]\n        bin_confs_a, bin_accs_a = compute_bin_stats_from_confidences(confs_after, truths, n_bins)\n\n        # Plot both\n        ax.plot([0,1],[0,1],\"k--\")\n        ax.plot(bin_confs_b, bin_accs_b, marker=\"o\", label=\"Before\")\n        ax.plot(bin_confs_a, bin_accs_a, marker=\"o\", label=\"After\")\n        ax.set_title(class_names[c])\n        ax.set_xlim([0,1]); ax.set_ylim([0,1])\n        ax.set_xlabel(\"Confidence\")\n        if c == 0:\n            ax.set_ylabel(\"True Accuracy\")\n        if c == num_classes-1:\n            ax.legend()\n\n    plt.suptitle(\"Calibration: Global and Class-wise (Before vs After Scaling)\", fontsize=14)\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:44:34.912129Z","iopub.execute_input":"2025-12-16T09:44:34.912405Z","iopub.status.idle":"2025-12-16T09:44:34.923251Z","shell.execute_reply.started":"2025-12-16T09:44:34.912382Z","shell.execute_reply":"2025-12-16T09:44:34.922657Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_names = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative\"]\n\nplot_combined_calibration_all(probs, probs_scaled, labels, class_names, n_bins=10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:44:40.566350Z","iopub.execute_input":"2025-12-16T09:44:40.566925Z","iopub.status.idle":"2025-12-16T09:44:41.480183Z","shell.execute_reply.started":"2025-12-16T09:44:40.566900Z","shell.execute_reply":"2025-12-16T09:44:41.479373Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Grad-Cam","metadata":{}},{"cell_type":"code","source":"class GradCAM:\n    def __init__(self, model, target_layer):\n        self.model = model\n        self.model.eval()\n        self.target_layer = target_layer\n\n        # placeholders for gradients and activations\n        self.gradients = None\n        self.activations = None\n\n        # hook for forward\n        def forward_hook(module, input, output):\n            self.activations = output.detach()\n        # hook for backward\n        def backward_hook(module, grad_input, grad_output):\n            self.gradients = grad_output[0].detach()\n\n        # register hooks\n        target_layer.register_forward_hook(forward_hook)\n        target_layer.register_full_backward_hook(backward_hook)\n\n    def generate_cam(self, input_tensor, target_class=None):\n        \"\"\"\n        input_tensor: (1, C, H, W) torch tensor\n        target_class: int, class index (if None, uses predicted class)\n        \"\"\"\n        # forward pass\n        output = self.model(input_tensor)\n        if target_class is None:\n            target_class = output.argmax(dim=1).item()\n\n        # backward pass\n        self.model.zero_grad()\n        loss = output[0, target_class]\n        loss.backward()\n\n        # compute weights\n        weights = self.gradients.mean(dim=(2, 3), keepdim=True)  # GAP over H,W\n        cam = (weights * self.activations).sum(dim=1).squeeze()\n\n        # relu + normalize\n        cam = torch.relu(cam)\n        cam = cam - cam.min()\n        cam = cam / (cam.max() + 1e-8)\n        cam = cam.cpu().numpy()\n        return cam","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:44:53.047998Z","iopub.execute_input":"2025-12-16T09:44:53.048604Z","iopub.status.idle":"2025-12-16T09:44:53.054921Z","shell.execute_reply.started":"2025-12-16T09:44:53.048568Z","shell.execute_reply":"2025-12-16T09:44:53.054178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Example: setup Grad-CAM\ntarget_layer = resnet50.layer4[-1]\ngradcam = GradCAM(resnet50, target_layer)\n\n# pick a test image\nimg, label = test_dataset[0]   # returns preprocessed tensor + label\ninput_tensor = img.unsqueeze(0).to(device)\n\n# generate cam\ncam = gradcam.generate_cam(input_tensor, target_class=None)\n\n# convert tensor to displayable image\ndef tensor_to_image(tensor):\n    img = tensor.permute(1, 2, 0).cpu().numpy()\n    img = (img - img.min()) / (img.max() - img.min())\n    return (img * 255).astype(np.uint8)\n\norig_img = tensor_to_image(img)\n\n# resize CAM to match image\nH, W, _ = orig_img.shape\ncam_resized = cv2.resize(cam, (W, H))\n\n# heatmap + overlay\nheatmap = cv2.applyColorMap(np.uint8(255*cam_resized), cv2.COLORMAP_JET)\nheatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\noverlay = cv2.addWeighted(orig_img, 0.5, heatmap, 0.5, 0)\n\nplt.figure(figsize=(12,4))\nplt.subplot(1,3,1); plt.imshow(orig_img); plt.title(\"Original\"); plt.axis(\"off\")\nplt.subplot(1,3,2); plt.imshow(cam_resized, cmap='jet'); plt.title(\"Grad-CAM Map\"); plt.axis(\"off\")\nplt.subplot(1,3,3); plt.imshow(overlay); plt.title(\"Overlay\"); plt.axis(\"off\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-16T09:44:57.133320Z","iopub.execute_input":"2025-12-16T09:44:57.133684Z","iopub.status.idle":"2025-12-16T09:44:57.901136Z","shell.execute_reply.started":"2025-12-16T09:44:57.133661Z","shell.execute_reply":"2025-12-16T09:44:57.900463Z"}},"outputs":[],"execution_count":null}]}