{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.12"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"1b24cce5-1600-40f6-897d-3d7d8531b8db","cell_type":"markdown","source":"# APTOS 2019 — Diabetic Retinopathy Classification\n### ResNet50 + Heavy Augmentation + Class-Weighted Focal Loss\n","metadata":{}},{"id":"7a42ec2e-384d-4712-9b9d-673fb4565b0e","cell_type":"markdown","source":"## Cell 1 — Imports","metadata":{}},{"id":"9e1ec034-571c-403b-9266-20d0e1b3afc7","cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\nimport torchvision.transforms as transforms\n\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    classification_report,\n    confusion_matrix,\n    roc_auc_score,\n    roc_curve\n)\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint('Device:', device)","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:23:32.543063Z","iopub.execute_input":"2026-04-11T11:23:32.543916Z","iopub.status.idle":"2026-04-11T11:23:32.550696Z","shell.execute_reply.started":"2026-04-11T11:23:32.543805Z","shell.execute_reply":"2026-04-11T11:23:32.549716Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"e4ba0694-a5b2-4166-9842-14ee324ec785","cell_type":"markdown","source":"## Cell 2 — Load CSV & Create Binary Label","metadata":{}},{"id":"b09eba25-4c3b-413e-820a-bf4c7e9983cc","cell_type":"code","source":"DATA_PATH = '/kaggle/input/competitions/aptos2019-blindness-detection'\nIMAGE_DIR = os.path.join(DATA_PATH, 'train_images')\n\ndf = pd.read_csv(os.path.join(DATA_PATH, 'train.csv'))\n\n\ndf['label'] = df['diagnosis'].apply(lambda x: 1 if x == 4 else 0)\n\nprint('Total samples:', len(df))\nprint(df['label'].value_counts())\nprint(f'\\nClass imbalance ratio: {df[\"label\"].value_counts()[0] / df[\"label\"].value_counts()[1]:.1f}:1')","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:23:32.552740Z","iopub.execute_input":"2026-04-11T11:23:32.553213Z","iopub.status.idle":"2026-04-11T11:23:32.577470Z","shell.execute_reply.started":"2026-04-11T11:23:32.553150Z","shell.execute_reply":"2026-04-11T11:23:32.576915Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"3b3d0fa8-69da-4a31-b1c5-996f9c96ac36","cell_type":"markdown","source":"## Cell 3 — Visualise Class Distribution","metadata":{}},{"id":"92f2eba2-f64b-417d-beee-87c3020a2cfe","cell_type":"code","source":"counts = df['label'].value_counts()\nplt.figure(figsize=(6, 4))\nplt.bar(['Non-PDR (0)', 'PDR (1)'], [counts[0], counts[1]],\n        color=['steelblue', 'tomato'])\nplt.title('Binary Class Distribution (PDR vs Non-PDR)')\nplt.ylabel('Count')\nfor i, v in enumerate([counts[0], counts[1]]):\n    plt.text(i, v + 10, str(v), ha='center', fontweight='bold')\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:23:32.578749Z","iopub.execute_input":"2026-04-11T11:23:32.579060Z","iopub.status.idle":"2026-04-11T11:23:32.696794Z","shell.execute_reply.started":"2026-04-11T11:23:32.579036Z","shell.execute_reply":"2026-04-11T11:23:32.696238Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"d88c2661-69a5-4f88-b57f-608eaf5690b0","cell_type":"markdown","source":"## Cell 4 — Train / Val / Test Split (stratified 70/15/15)","metadata":{}},{"id":"7ad943c0-eb43-4abd-a10a-4062d87dd3ac","cell_type":"code","source":"train_df, temp_df = train_test_split(\n    df, test_size=0.30, stratify=df['label'], random_state=42\n)\nval_df, test_df = train_test_split(\n    temp_df, test_size=0.50, stratify=temp_df['label'], random_state=42\n)\n\ntrain_df = train_df.reset_index(drop=True)\nval_df   = val_df.reset_index(drop=True)\ntest_df  = test_df.reset_index(drop=True)\n\nprint(f'Train: {len(train_df)} | Val: {len(val_df)} | Test: {len(test_df)}')\nprint('\\nTrain distribution:')\nprint(train_df['label'].value_counts())\nprint('\\nVal distribution:')\nprint(val_df['label'].value_counts())\nprint('\\nTest distribution:')\nprint(test_df['label'].value_counts())","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:23:32.697590Z","iopub.execute_input":"2026-04-11T11:23:32.697965Z","iopub.status.idle":"2026-04-11T11:23:32.713499Z","shell.execute_reply.started":"2026-04-11T11:23:32.697941Z","shell.execute_reply":"2026-04-11T11:23:32.712678Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"7ab79e84-39a6-40e5-8654-1cc4396d57ed","cell_type":"markdown","source":"## Cell 5 — Transforms (heavy augmentation for training)","metadata":{}},{"id":"d750cca1-9742-4d80-8e12-8ccb6a428e8e","cell_type":"code","source":"\ntrain_transform = transforms.Compose([\n    transforms.Resize((256, 256)),\n    transforms.RandomCrop(224),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(20),\n    transforms.ColorJitter(\n        brightness=0.3,\n        contrast=0.3,\n        saturation=0.2,\n        hue=0.05\n    ),\n    transforms.RandomGrayscale(p=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n\n\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n\nprint('Transforms ready.')","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:23:32.715265Z","iopub.execute_input":"2026-04-11T11:23:32.715515Z","iopub.status.idle":"2026-04-11T11:23:32.722020Z","shell.execute_reply.started":"2026-04-11T11:23:32.715494Z","shell.execute_reply":"2026-04-11T11:23:32.721283Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"e77c8cbe-479c-4ffd-ac50-ec42c652b196","cell_type":"markdown","source":"## Cell 6 — Dataset Class","metadata":{}},{"id":"7581120d-1691-44fd-acf1-fe3f40d18cad","cell_type":"code","source":"class APTOSDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.df = dataframe.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        img_path = os.path.join(IMAGE_DIR,\n                                self.df.loc[idx, 'id_code'] + '.png')\n        image = Image.open(img_path).convert('RGB')\n        label = int(self.df.loc[idx, 'label'])\n        if self.transform:\n            image = self.transform(image)\n        return image, label\n\nprint('APTOSDataset ready.')","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:23:32.722803Z","iopub.execute_input":"2026-04-11T11:23:32.723100Z","iopub.status.idle":"2026-04-11T11:23:32.736482Z","shell.execute_reply.started":"2026-04-11T11:23:32.723069Z","shell.execute_reply":"2026-04-11T11:23:32.735671Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"435064a0-e162-4f38-aa20-824e28753f9e","cell_type":"markdown","source":"## Cell 7 — Weighted Sampler (oversamples PDR during training)","metadata":{}},{"id":"c3f63ba6-3806-408e-a6a0-7e49752f92a6","cell_type":"code","source":"\nclass_counts = train_df['label'].value_counts().to_dict()\nsample_weights = train_df['label'].map(\n    lambda x: 1.0 / class_counts[x]\n).values\n\nsampler = WeightedRandomSampler(\n    weights=sample_weights,\n    num_samples=len(sample_weights),\n    replacement=True\n)\n\ntrain_dataset = APTOSDataset(train_df, transform=train_transform)\nval_dataset   = APTOSDataset(val_df,   transform=val_transform)\ntest_dataset  = APTOSDataset(test_df,  transform=val_transform)\n\n\ntrain_loader = DataLoader(train_dataset, batch_size=32,\n                          sampler=sampler, num_workers=0)\nval_loader   = DataLoader(val_dataset,   batch_size=32,\n                          shuffle=False,  num_workers=0)\ntest_loader  = DataLoader(test_dataset,  batch_size=32,\n                          shuffle=False,  num_workers=0)\n\n\nimgs, lbls = next(iter(train_loader))\nprint('Batch shape:', imgs.shape)\nprint('Labels in first batch:', lbls.tolist())","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:23:32.737318Z","iopub.execute_input":"2026-04-11T11:23:32.737577Z","iopub.status.idle":"2026-04-11T11:23:37.562397Z","shell.execute_reply.started":"2026-04-11T11:23:32.737541Z","shell.execute_reply":"2026-04-11T11:23:37.561658Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"50b2fa13-d8dc-4f4e-a1da-d3a08ac61958","cell_type":"markdown","source":"## Cell 8 — Model (ResNet50, pretrained)","metadata":{}},{"id":"2165854f-1f34-4b62-ae75-12dabb6dd990","cell_type":"code","source":"model = models.resnet50(weights=models.ResNet50_Weights.DEFAULT)\n\n\nfor param in model.parameters():\n    param.requires_grad = False\n\n\nmodel.fc = nn.Sequential(\n    nn.Linear(model.fc.in_features, 256),\n    nn.ReLU(),\n    nn.Dropout(0.4),\n    nn.Linear(256, 2)\n)\nmodel = model.to(device)\nprint('Model ready.')","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:23:37.563431Z","iopub.execute_input":"2026-04-11T11:23:37.563767Z","iopub.status.idle":"2026-04-11T11:23:37.958161Z","shell.execute_reply.started":"2026-04-11T11:23:37.563741Z","shell.execute_reply":"2026-04-11T11:23:37.957309Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"eec2a804-675a-4939-ab4d-da8ff7997265","cell_type":"markdown","source":"## Cell 9 — Class-Weighted Focal Loss & Helpers","metadata":{}},{"id":"9c06f55b-f88f-4f42-a167-0e8639532797","cell_type":"code","source":"# Compute class weights from training set\nn_non_pdr = class_counts[0]\nn_pdr     = class_counts[1]\ntotal     = n_non_pdr + n_pdr\nweight_non_pdr = total / (2 * n_non_pdr)\nweight_pdr     = total / (2 * n_pdr)\nclass_weights  = torch.tensor([weight_non_pdr, weight_pdr],\n                               dtype=torch.float).to(device)\nprint(f'Class weights — Non-PDR: {weight_non_pdr:.3f} | PDR: {weight_pdr:.3f}')\n\n\nclass FocalLoss(nn.Module):\n    \"\"\"Weighted Focal Loss — handles class imbalance.\"\"\"\n    def __init__(self, alpha=1, gamma=2, weight=None):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.ce = nn.CrossEntropyLoss(weight=weight)\n\n    def forward(self, inputs, targets):\n        ce_loss = self.ce(inputs, targets)\n        pt = torch.exp(-ce_loss)\n        return self.alpha * (1 - pt) ** self.gamma * ce_loss\n\n\ncriterion = FocalLoss(alpha=1, gamma=2, weight=class_weights)\noptimizer = torch.optim.Adam(model.fc.parameters(), lr=1e-4)\n\n\ndef train_one_epoch(model, loader, optimizer, criterion):\n    model.train()\n    total_loss = 0\n    for images, labels in loader:\n        images, labels = images.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        total_loss += loss.item()\n    return total_loss / len(loader)\n\n\ndef evaluate(model, loader):\n    model.eval()\n    all_preds, all_labels, all_probs = [], [], []\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device)\n            outputs = model(images)\n            probs = torch.softmax(outputs, dim=1)[:, 1].cpu().numpy()\n            _, preds = torch.max(outputs, 1)\n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.numpy())\n            all_probs.extend(probs)\n    return all_labels, all_preds, all_probs\n\nprint('Loss, optimizer, helpers ready.')","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:23:37.959225Z","iopub.execute_input":"2026-04-11T11:23:37.959768Z","iopub.status.idle":"2026-04-11T11:23:37.970585Z","shell.execute_reply.started":"2026-04-11T11:23:37.959742Z","shell.execute_reply":"2026-04-11T11:23:37.969925Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"4a649bed-cafd-4bd0-94c6-c06dc850c6b2","cell_type":"markdown","source":"## Cell 10 — Phase 1: Warmup (head only, 3 epochs)","metadata":{}},{"id":"38eab308-cdd5-4cda-a106-598b7caacd2d","cell_type":"code","source":"print('=== Phase 1: Warmup — training head only ===')\nfor epoch in range(3):\n    loss = train_one_epoch(model, train_loader, optimizer, criterion)\n    val_true, val_pred, val_probs = evaluate(model, val_loader)\n    auc = roc_auc_score(val_true, val_probs)\n    print(f'Epoch {epoch+1}/3 | Loss: {loss:.4f} | Val AUC: {auc:.4f}')","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:23:37.972964Z","iopub.execute_input":"2026-04-11T11:23:37.973224Z","iopub.status.idle":"2026-04-11T11:45:29.991676Z","shell.execute_reply.started":"2026-04-11T11:23:37.973204Z","shell.execute_reply":"2026-04-11T11:45:29.990900Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"013e10f0-ee1e-4d27-a438-6ad52d3a581b","cell_type":"markdown","source":"## Cell 11 — Phase 2: Fine-tune all layers (10 epochs)","metadata":{}},{"id":"f43dcd4a-464e-4a94-91c9-637627df803e","cell_type":"code","source":"# Unfreeze everything\nfor param in model.parameters():\n    param.requires_grad = True\n\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-5)\nscheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=4, gamma=0.5)\n\nbest_auc  = 0.0\ntrain_losses, val_aucs = [], []\n\nprint('=== Phase 2: Fine-tuning all layers ===')\nfor epoch in range(10):\n    loss = train_one_epoch(model, train_loader, optimizer, criterion)\n    val_true, val_pred, val_probs = evaluate(model, val_loader)\n    auc = roc_auc_score(val_true, val_probs)\n    scheduler.step()\n\n    train_losses.append(loss)\n    val_aucs.append(auc)\n\n    \n    if auc > best_auc:\n        best_auc = auc\n        torch.save(model.state_dict(), '/kaggle/working/best_model.pth')\n        print(f'Epoch {epoch+1}/10 | Loss: {loss:.4f} | Val AUC: {auc:.4f}  ← best saved')\n    else:\n        print(f'Epoch {epoch+1}/10 | Loss: {loss:.4f} | Val AUC: {auc:.4f}')\n\nprint(f'\\nBest Val AUC: {best_auc:.4f}')","metadata":{"execution":{"iopub.status.busy":"2026-04-11T11:45:29.992985Z","iopub.execute_input":"2026-04-11T11:45:29.993274Z","iopub.status.idle":"2026-04-11T13:01:26.118437Z","shell.execute_reply.started":"2026-04-11T11:45:29.993248Z","shell.execute_reply":"2026-04-11T13:01:26.117593Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"79c5ee3c-11d9-4ae2-9209-fce2a4640729","cell_type":"markdown","source":"## Cell 12 — Plot Training Curves","metadata":{}},{"id":"7d7d24a1-a038-4864-b193-4b70d12925b7","cell_type":"code","source":"fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))\n\nax1.plot(range(1, 11), train_losses, marker='o', color='steelblue')\nax1.set_title('Training Loss')\nax1.set_xlabel('Epoch')\nax1.set_ylabel('Loss')\nax1.grid(True)\n\nax2.plot(range(1, 11), val_aucs, marker='o', color='tomato')\nax2.set_title('Validation AUC-ROC')\nax2.set_xlabel('Epoch')\nax2.set_ylabel('AUC')\nax2.set_ylim([0, 1])\nax2.grid(True)\n\nplt.suptitle('Training Curves')\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2026-04-11T13:01:26.119689Z","iopub.execute_input":"2026-04-11T13:01:26.119973Z","iopub.status.idle":"2026-04-11T13:01:26.374234Z","shell.execute_reply.started":"2026-04-11T13:01:26.119948Z","shell.execute_reply":"2026-04-11T13:01:26.373631Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"6f8e437f-5846-4580-af41-ef3f2de4c54e","cell_type":"markdown","source":"## Cell 13 — Load Best Model & Tune Decision Threshold","metadata":{}},{"id":"2a33babf-cc33-49b3-94d4-7b6cba2cc551","cell_type":"code","source":"\nmodel.load_state_dict(torch.load('/kaggle/working/best_model.pth'))\n\n\nval_true, _, val_probs = evaluate(model, val_loader)\n\n\nfrom sklearn.metrics import f1_score\n\nbest_thresh, best_f1 = 0.5, 0.0\nfor thresh in np.arange(0.1, 0.9, 0.01):\n    preds = (np.array(val_probs) >= thresh).astype(int)\n    f1 = f1_score(val_true, preds, pos_label=1, zero_division=0)\n    if f1 > best_f1:\n        best_f1 = f1\n        best_thresh = thresh\n\nprint(f'Best threshold: {best_thresh:.2f} (Val PDR F1: {best_f1:.4f})')","metadata":{"execution":{"iopub.status.busy":"2026-04-11T13:01:26.375228Z","iopub.execute_input":"2026-04-11T13:01:26.375483Z","iopub.status.idle":"2026-04-11T13:02:29.080898Z","shell.execute_reply.started":"2026-04-11T13:01:26.375460Z","shell.execute_reply":"2026-04-11T13:02:29.080029Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"e44732cc-adb7-4dc9-9c7e-815b1c64d539","cell_type":"markdown","source":"## Cell 14 — Final Test Evaluation","metadata":{}},{"id":"76e30f3a-0853-431e-83b9-ed6604e77375","cell_type":"code","source":"test_true, _, test_probs = evaluate(model, test_loader)\n\n\ntest_pred = (np.array(test_probs) >= best_thresh).astype(int)\n\n\ncm = confusion_matrix(test_true, test_pred)\nplt.figure(figsize=(5, 4))\nsns.heatmap(cm, annot=True, fmt='d',\n            xticklabels=['Non-PDR', 'PDR'],\n            yticklabels=['Non-PDR', 'PDR'],\n            cmap='Blues')\nplt.title('Test Confusion Matrix')\nplt.ylabel('Actual')\nplt.xlabel('Predicted')\nplt.tight_layout()\nplt.show()\n\n\nfpr, tpr, _ = roc_curve(test_true, test_probs)\nauc_score   = roc_auc_score(test_true, test_probs)\n\nplt.figure(figsize=(5, 4))\nplt.plot(fpr, tpr, color='tomato', label=f'AUC = {auc_score:.4f}')\nplt.plot([0,1], [0,1], 'k--')\nplt.xlabel('False Positive Rate')\nplt.ylabel('True Positive Rate')\nplt.title('ROC Curve (Test Set)')\nplt.legend()\nplt.tight_layout()\nplt.show()\n\n\nprint(f'Test AUC-ROC: {auc_score:.4f}')\nprint()\nprint(classification_report(test_true, test_pred,\n      target_names=['Non-PDR', 'PDR'], digits=4))","metadata":{"execution":{"iopub.status.busy":"2026-04-11T13:02:29.082069Z","iopub.execute_input":"2026-04-11T13:02:29.082394Z","iopub.status.idle":"2026-04-11T13:03:42.376667Z","shell.execute_reply.started":"2026-04-11T13:02:29.082367Z","shell.execute_reply":"2026-04-11T13:03:42.375828Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"9ced41c1-56a1-465d-878f-f47a448893a9","cell_type":"markdown","source":"## Cell 15 — Visualise Sample Predictions","metadata":{}},{"id":"3f26380a-d8b4-4b3f-8efd-0ff5f5af7fd7","cell_type":"code","source":"\nmodel.eval()\nimages_shown = 0\nfig, axes = plt.subplots(2, 4, figsize=(14, 7))\naxes = axes.flatten()\n\ninv_normalize = transforms.Normalize(\n    mean=[-0.485/0.229, -0.456/0.224, -0.406/0.225],\n    std=[1/0.229, 1/0.224, 1/0.225]\n)\n\nwith torch.no_grad():\n    for images, labels in test_loader:\n        images_gpu = images.to(device)\n        outputs = model(images_gpu)\n        probs = torch.softmax(outputs, dim=1)[:, 1].cpu().numpy()\n        preds = (probs >= best_thresh).astype(int)\n\n        for i in range(len(images)):\n            if images_shown >= 8:\n                break\n            img = inv_normalize(images[i]).permute(1, 2, 0).numpy()\n            img = np.clip(img, 0, 1)\n            axes[images_shown].imshow(img)\n            color = 'green' if preds[i] == labels[i].item() else 'red'\n            axes[images_shown].set_title(\n                f'True: {\"PDR\" if labels[i]==1 else \"Non-PDR\"}\\n'\n                f'Pred: {\"PDR\" if preds[i]==1 else \"Non-PDR\"} ({probs[i]:.2f})',\n                color=color, fontsize=9\n            )\n            axes[images_shown].axis('off')\n            images_shown += 1\n        if images_shown >= 8:\n            break\n\nplt.suptitle('Sample Predictions (Green=Correct, Red=Wrong)', fontsize=12)\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2026-04-11T13:03:42.377747Z","iopub.execute_input":"2026-04-11T13:03:42.378377Z","iopub.status.idle":"2026-04-11T13:03:46.652743Z","shell.execute_reply.started":"2026-04-11T13:03:42.378351Z","shell.execute_reply":"2026-04-11T13:03:46.651824Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"cd31eda7-47a8-4fa3-8beb-5224eb629f0f","cell_type":"markdown","source":"## Cell 16 — Save Final Model","metadata":{}},{"id":"6b80025a-b74b-4f75-bd0b-2a4b1f2fbaf4","cell_type":"code","source":"torch.save(model.state_dict(), '/kaggle/working/resnet50_final.pth')\nprint('Model saved to /kaggle/working/resnet50_final.pth')\nprint(f'Final Test AUC: {auc_score:.4f}')\nprint(f'Decision threshold used: {best_thresh:.2f}')","metadata":{"execution":{"iopub.status.busy":"2026-04-11T13:03:46.654084Z","iopub.execute_input":"2026-04-11T13:03:46.654895Z","iopub.status.idle":"2026-04-11T13:03:46.784490Z","shell.execute_reply.started":"2026-04-11T13:03:46.654834Z","shell.execute_reply":"2026-04-11T13:03:46.783577Z"},"trusted":true},"outputs":[],"execution_count":null},{"id":"e0175c8f-93b8-4466-9f14-51178cec26f5","cell_type":"markdown","source":"---\n# Part 2 — EfficientNet-B4 + CBAM\n### Upgrading from ResNet50\n\n**Changes:**\n- `timm` EfficientNet-B4 backbone at native 380×380 resolution\n- CBAM (Channel + Spatial Attention) on last feature map (1792 ch)\n- AdamW + CosineAnnealingLR with differential learning rates\n- AMP mixed precision (CUDA only)\n- Batch size 16 (B4 at 380px is more VRAM-heavy)\n\nEverything else (focal loss, weighted sampler, threshold tuning) is unchanged.\n\n---","metadata":{}},{"id":"e85b6a6c-655c-4723-afce-7bd131cf5f84","cell_type":"markdown","source":"## Cell 17 — Install & Import timm","metadata":{}},{"id":"a6acd0e1-ecc6-4f8b-99a2-ccd724203265","cell_type":"code","source":"\nimport timm\nprint('timm version:', timm.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T13:03:46.785501Z","iopub.execute_input":"2026-04-11T13:03:46.785789Z","iopub.status.idle":"2026-04-11T13:03:46.793151Z","shell.execute_reply.started":"2026-04-11T13:03:46.785755Z","shell.execute_reply":"2026-04-11T13:03:46.791677Z"}},"outputs":[],"execution_count":null},{"id":"6c92a0e0-e018-46c6-b537-1620678d5019","cell_type":"markdown","source":"## Cell 18 — Updated Transforms (380×380)\n> EfficientNet-B4 was trained at 380×380 — matching this gives the best feature quality.","metadata":{}},{"id":"18a302bd-f6a9-4d7c-8a98-c821dc158c44","cell_type":"code","source":"IMG_SIZE = 380  \n\ntrain_transform_b4 = transforms.Compose([\n    transforms.Resize((IMG_SIZE + 32, IMG_SIZE + 32)),\n    transforms.RandomCrop(IMG_SIZE),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(20),\n    transforms.ColorJitter(brightness=0.3, contrast=0.3, saturation=0.2, hue=0.05),\n    transforms.RandomGrayscale(p=0.05),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\nval_transform_b4 = transforms.Compose([\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])\n\nprint('Transforms ready — input size:', IMG_SIZE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T13:03:46.794709Z","iopub.execute_input":"2026-04-11T13:03:46.795094Z","iopub.status.idle":"2026-04-11T13:03:46.805647Z","shell.execute_reply.started":"2026-04-11T13:03:46.795054Z","shell.execute_reply":"2026-04-11T13:03:46.804692Z"}},"outputs":[],"execution_count":null},{"id":"00621e3c-82a8-461e-a52a-e17278cb22d9","cell_type":"markdown","source":"## Cell 19 — Rebuild DataLoaders (batch 16)\n> Batch size 32 → 16 because B4 at 380px needs more VRAM. Sampler and splits are reused as-is.","metadata":{}},{"id":"c7b4d489-11dd-40ba-ad7b-5b30a46cbb49","cell_type":"code","source":"BATCH_SIZE = 16\n\ntrain_dataset_b4 = APTOSDataset(train_df, transform=train_transform_b4)\nval_dataset_b4   = APTOSDataset(val_df,   transform=val_transform_b4)\ntest_dataset_b4  = APTOSDataset(test_df,  transform=val_transform_b4)\n\n\ntrain_loader_b4 = DataLoader(train_dataset_b4, batch_size=BATCH_SIZE, sampler=sampler,\n                             num_workers=2, pin_memory=True)\nval_loader_b4   = DataLoader(val_dataset_b4,   batch_size=BATCH_SIZE, shuffle=False,\n                             num_workers=2, pin_memory=True)\ntest_loader_b4  = DataLoader(test_dataset_b4,  batch_size=BATCH_SIZE, shuffle=False,\n                             num_workers=2, pin_memory=True)\n\nimgs, lbls = next(iter(train_loader_b4))\nprint('Batch shape:', imgs.shape)   # expect (16, 3, 380, 380)\nprint('Labels sample:', lbls.tolist())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T13:03:46.806653Z","iopub.execute_input":"2026-04-11T13:03:46.807017Z","iopub.status.idle":"2026-04-11T13:03:52.351027Z","shell.execute_reply.started":"2026-04-11T13:03:46.806970Z","shell.execute_reply":"2026-04-11T13:03:52.350004Z"}},"outputs":[],"execution_count":null},{"id":"9ae5ad5e-2221-46be-8b73-3195b0e9e12c","cell_type":"markdown","source":"## Cell 20 — CBAM + EfficientNet-B4 Model\n\n```\nInput (B, 3, 380, 380)\n        |\n  EfficientNet-B4 backbone  (timm, pretrained)\n        |\n  Last feature map  (B, 1792, ~12, ~12)\n        |\n      CBAM\n   +-------+-------+\n Channel         Spatial\n Attention       Attention\n (what)          (where)\n        |\n  AdaptiveAvgPool2d  ->  (B, 1792)\n        |\n  Linear(1792->256) -> ReLU -> Dropout(0.4) -> Linear(256->2)\n```","metadata":{}},{"id":"e1f2ff19-8b37-40a6-ba94-110690c1df0b","cell_type":"code","source":"\nclass ChannelAttention(nn.Module):\n    \"\"\"WHAT features matter — uses both avg and max pool signals.\"\"\"\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n        # Shared MLP\n        self.mlp = nn.Sequential(\n            nn.Linear(channels, channels // reduction, bias=False),\n            nn.ReLU(),\n            nn.Linear(channels // reduction, channels, bias=False)\n        )\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        B, C, _, _ = x.shape\n        avg = self.mlp(self.avg_pool(x).view(B, C))\n        mx  = self.mlp(self.max_pool(x).view(B, C))\n        scale = self.sigmoid(avg + mx).view(B, C, 1, 1)\n        return x * scale\n\n\n\nclass SpatialAttention(nn.Module):\n    \"\"\"WHERE to look — highlights diagnostically important regions.\"\"\"\n    def __init__(self, kernel_size=7):\n        super().__init__()\n        self.conv = nn.Conv2d(2, 1, kernel_size,\n                              padding=kernel_size // 2, bias=False)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg    = x.mean(dim=1, keepdim=True)      # (B, 1, H, W)\n        mx, _  = x.max(dim=1,  keepdim=True)      # (B, 1, H, W)\n        attn   = self.sigmoid(self.conv(torch.cat([avg, mx], dim=1)))\n        return x * attn\n\n\n\nclass CBAM(nn.Module):\n    \"\"\"Convolutional Block Attention Module (Woo et al., 2018).\"\"\"\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n        self.channel = ChannelAttention(channels, reduction)\n        self.spatial = SpatialAttention()\n\n    def forward(self, x):\n        x = self.channel(x)   # channel first\n        x = self.spatial(x)   # then spatial\n        return x\n\n\n# ─── EfficientNet-B4 + CBAM ───────────────────────────────────────────────────\nclass EfficientNetB4_CBAM(nn.Module):\n    def __init__(self, num_classes=2, dropout=0.4):\n        super().__init__()\n\n       \n        self.backbone = timm.create_model(\n            'efficientnet_b4', pretrained=True, features_only=True\n        )\n        last_ch = self.backbone.feature_info[-1]['num_chs']  # 1792\n        print(f'Backbone last-stage channels: {last_ch}')\n\n       \n        self.cbam = CBAM(last_ch)\n\n    \n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.head = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(last_ch, 256),\n            nn.ReLU(),\n            nn.Dropout(dropout),\n            nn.Linear(256, num_classes)\n        )\n\n        self._freeze_backbone(True)   \n\n    def _freeze_backbone(self, freeze: bool):\n        for p in self.backbone.parameters():\n            p.requires_grad = not freeze\n\n    def forward(self, x):\n        feat = self.backbone(x)[-1]   \n        feat = self.cbam(feat)         \n        feat = self.pool(feat)       \n        return self.head(feat)\n\n\n\nmodel_b4 = EfficientNetB4_CBAM(num_classes=2, dropout=0.4).to(device)\n\ntotal     = sum(p.numel() for p in model_b4.parameters())\ntrainable = sum(p.numel() for p in model_b4.parameters() if p.requires_grad)\nprint(f'Total params    : {total:,}')\nprint(f'Trainable params: {trainable:,}  (backbone frozen)')\n\n\nwith torch.no_grad():\n    dummy = torch.randn(2, 3, 380, 380).to(device)\n    out = model_b4(dummy)\n    print(f'Forward pass OK — output shape: {out.shape}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-11T13:03:52.352451Z","iopub.execute_input":"2026-04-11T13:03:52.352716Z","execution_failed":"2026-04-11T14:27:02.670Z"}},"outputs":[],"execution_count":null},{"id":"967130a9-ad68-4026-ab9a-dfa2e62dd6d1","cell_type":"markdown","source":"## Cell 21 — Loss, Optimizer Helpers & AMP\n> Reusing FocalLoss and class_weights from Cell 9. Adding AMP scaler for CUDA.","metadata":{}},{"id":"3bd37fb7-530d-49c0-b298-8d598d0d63e7","cell_type":"code","source":"\ncriterion_b4 = FocalLoss(alpha=1, gamma=2, weight=class_weights)\n\n\ndef train_one_epoch_b4(model, loader, optimizer, criterion, scaler=None):\n    model.train()\n    total_loss = 0\n    for images, labels in loader:\n        images, labels = images.to(device), labels.to(device)\n        optimizer.zero_grad()\n        if scaler is not None:\n            with torch.cuda.amp.autocast():\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            optimizer.step()\n        total_loss += loss.item()\n    return total_loss / len(loader)\n\n\n\nscaler_b4 = torch.cuda.amp.GradScaler() if device.type == 'cuda' else None\nprint(f'Mixed-precision (AMP): {scaler_b4 is not None}')","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.676Z"}},"outputs":[],"execution_count":null},{"id":"3598df65-3134-42b2-b211-f0a2b6f4a26e","cell_type":"markdown","source":"## Cell 22 — Phase 1: Warmup (CBAM + head only, 5 epochs)","metadata":{}},{"id":"3444af3d-9420-4509-98bc-73de24e51feb","cell_type":"code","source":"optimizer_warmup_b4 = torch.optim.AdamW(\n    filter(lambda p: p.requires_grad, model_b4.parameters()),\n    lr=3e-4, weight_decay=1e-4\n)\n\nprint('=== Phase 1: Warmup — backbone frozen ===')\nfor epoch in range(5):\n    loss = train_one_epoch_b4(model_b4, train_loader_b4,\n                              optimizer_warmup_b4, criterion_b4, scaler_b4)\n    val_true, val_pred, val_probs = evaluate(model_b4, val_loader_b4)\n    auc = roc_auc_score(val_true, val_probs)\n    print(f'Epoch {epoch+1}/5 | Loss: {loss:.4f} | Val AUC: {auc:.4f}')","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.676Z"}},"outputs":[],"execution_count":null},{"id":"b15b386c-6505-41a2-aa4e-f52ebd943b1e","cell_type":"markdown","source":"## Cell 23 — Phase 2: Fine-tune all layers (15 epochs)\n> **Differential LRs** — backbone `1e-5`, CBAM + head `1e-4`. Cosine annealing decays both smoothly to near-zero.","metadata":{}},{"id":"dfae2092-8d7f-42fc-94ee-12c6881471fb","cell_type":"code","source":"NUM_EPOCHS_B4 = 15\n\nmodel_b4._freeze_backbone(False)  \n\nbackbone_params     = list(model_b4.backbone.parameters())\nnon_backbone_params = [p for n, p in model_b4.named_parameters()\n                       if 'backbone' not in n]\n\noptimizer_b4 = torch.optim.AdamW([\n    {'params': backbone_params,     'lr': 1e-5},   \n    {'params': non_backbone_params, 'lr': 1e-4},   \n], weight_decay=1e-4)\n\nscheduler_b4 = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer_b4, T_max=NUM_EPOCHS_B4, eta_min=1e-7\n)\n\nbest_auc_b4 = 0.0\ntrain_losses_b4, val_aucs_b4, lrs_b4 = [], [], []\n\nprint('=== Phase 2: Fine-tuning all layers ===')\nfor epoch in range(NUM_EPOCHS_B4):\n    loss = train_one_epoch_b4(model_b4, train_loader_b4,\n                              optimizer_b4, criterion_b4, scaler_b4)\n    val_true, val_pred, val_probs = evaluate(model_b4, val_loader_b4)\n    auc    = roc_auc_score(val_true, val_probs)\n    lr_now = optimizer_b4.param_groups[0]['lr']\n    scheduler_b4.step()\n\n    train_losses_b4.append(loss)\n    val_aucs_b4.append(auc)\n    lrs_b4.append(lr_now)\n\n    if auc > best_auc_b4:\n        best_auc_b4 = auc\n        torch.save(model_b4.state_dict(), '/kaggle/working/best_efnb4_cbam.pth')\n        tag = '  <- best saved'\n    else:\n        tag = ''\n    print(f'Epoch {epoch+1:02d}/{NUM_EPOCHS_B4} | Loss: {loss:.4f} | '\n          f'Val AUC: {auc:.4f} | LR: {lr_now:.2e}{tag}')\n\nprint(f'\\nBest Val AUC: {best_auc_b4:.4f}')","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.676Z"}},"outputs":[],"execution_count":null},{"id":"837edcb7-66e0-49f3-baef-a793bbb353ff","cell_type":"markdown","source":"## Cell 24 — Training Curves","metadata":{}},{"id":"f4e38d28-f831-4434-8c92-097bb8028dfa","cell_type":"code","source":"fig, axes = plt.subplots(1, 3, figsize=(15, 4))\n\naxes[0].plot(range(1, NUM_EPOCHS_B4+1), train_losses_b4, marker='o', color='steelblue')\naxes[0].set_title('Training Loss')\naxes[0].set_xlabel('Epoch'); axes[0].set_ylabel('Focal Loss'); axes[0].grid(True)\n\naxes[1].plot(range(1, NUM_EPOCHS_B4+1), val_aucs_b4, marker='o', color='tomato')\naxes[1].set_title('Validation AUC-ROC')\naxes[1].set_xlabel('Epoch'); axes[1].set_ylabel('AUC')\naxes[1].set_ylim([0, 1]); axes[1].grid(True)\n\naxes[2].plot(range(1, NUM_EPOCHS_B4+1), lrs_b4, marker='.', color='purple')\naxes[2].set_title('LR Schedule (backbone)')\naxes[2].set_xlabel('Epoch'); axes[2].set_ylabel('LR'); axes[2].grid(True)\n\nplt.suptitle('EfficientNet-B4 + CBAM — Training Curves')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.677Z"}},"outputs":[],"execution_count":null},{"id":"8c63be86-7fd7-4828-8ce2-f10b7d952a94","cell_type":"markdown","source":"## Cell 25 — Load Best Model & Tune Decision Threshold","metadata":{}},{"id":"7fbf9469-7b64-43ab-85fb-072d4148e23b","cell_type":"code","source":"model_b4.load_state_dict(torch.load('/kaggle/working/best_efnb4_cbam.pth'))\n\nval_true, _, val_probs = evaluate(model_b4, val_loader_b4)\n\nbest_thresh_b4, best_f1_b4 = 0.5, 0.0\nfor thresh in np.arange(0.1, 0.9, 0.01):\n    preds = (np.array(val_probs) >= thresh).astype(int)\n    f1 = f1_score(val_true, preds, pos_label=1, zero_division=0)\n    if f1 > best_f1_b4:\n        best_f1_b4     = f1\n        best_thresh_b4 = thresh\n\nprint(f'Best threshold: {best_thresh_b4:.2f}  (Val PDR F1: {best_f1_b4:.4f})')","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.677Z"}},"outputs":[],"execution_count":null},{"id":"7414f4e3-ba49-4798-85a0-af314fc12d25","cell_type":"markdown","source":"## Cell 26 — Final Test Evaluation","metadata":{}},{"id":"e383a1c3-b6f8-404c-a170-a6dd1d8ce0ae","cell_type":"code","source":"test_true_b4, _, test_probs_b4 = evaluate(model_b4, test_loader_b4)\ntest_pred_b4 = (np.array(test_probs_b4) >= best_thresh_b4).astype(int)\n\ncm_b4 = confusion_matrix(test_true_b4, test_pred_b4)\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(11, 4))\n\nsns.heatmap(cm_b4, annot=True, fmt='d',\n            xticklabels=['Non-PDR', 'PDR'],\n            yticklabels=['Non-PDR', 'PDR'],\n            cmap='Blues', ax=ax1)\nax1.set_title('Test Confusion Matrix — EfficientNet-B4 + CBAM')\nax1.set_ylabel('Actual'); ax1.set_xlabel('Predicted')\n\nfpr_b4, tpr_b4, _ = roc_curve(test_true_b4, test_probs_b4)\nauc_b4 = roc_auc_score(test_true_b4, test_probs_b4)\nax2.plot(fpr_b4, tpr_b4, color='tomato', label=f'B4+CBAM  AUC = {auc_b4:.4f}')\n# Overlay ResNet50 curve for direct comparison\nax2.plot(fpr, tpr, color='steelblue', linestyle='--',\n         label=f'ResNet50  AUC = {auc_score:.4f}')\nax2.plot([0,1],[0,1],'k--', linewidth=0.8)\nax2.set_xlabel('False Positive Rate'); ax2.set_ylabel('True Positive Rate')\nax2.set_title('ROC Curve Comparison (Test Set)'); ax2.legend()\n\nplt.tight_layout()\nplt.show()\n\nprint(f'EfficientNet-B4 + CBAM  Test AUC : {auc_b4:.4f}')\nprint(f'ResNet50 (baseline)     Test AUC : {auc_score:.4f}')\nprint(f'Improvement                       : {auc_b4 - auc_score:+.4f}')\nprint(f'Decision threshold: {best_thresh_b4:.2f}\\n')\nprint(classification_report(test_true_b4, test_pred_b4,\n      target_names=['Non-PDR', 'PDR'], digits=4))","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.677Z"}},"outputs":[],"execution_count":null},{"id":"5aed7502-bd3e-42a6-898e-9e95e8c290e9","cell_type":"markdown","source":"## Cell 27 — Visualise Sample Predictions","metadata":{}},{"id":"dc2dd13a-b5ee-4126-a474-e4f8dac73935","cell_type":"code","source":"model_b4.eval()\nfig, axes = plt.subplots(2, 4, figsize=(14, 7))\naxes = axes.flatten()\n\ninv_normalize = transforms.Normalize(\n    mean=[-0.485/0.229, -0.456/0.224, -0.406/0.225],\n    std=[1/0.229, 1/0.224, 1/0.225]\n)\n\nshown = 0\nwith torch.no_grad():\n    for images, labels in test_loader_b4:\n        probs = torch.softmax(model_b4(images.to(device)), dim=1)[:, 1].cpu().numpy()\n        preds = (probs >= best_thresh_b4).astype(int)\n        for i in range(len(images)):\n            if shown >= 8:\n                break\n            img = np.clip(inv_normalize(images[i]).permute(1,2,0).numpy(), 0, 1)\n            axes[shown].imshow(img)\n            color = 'green' if preds[i] == labels[i].item() else 'red'\n            true_lbl = 'PDR' if labels[i] == 1 else 'Non-PDR'\n            pred_lbl = 'PDR' if preds[i] == 1 else 'Non-PDR'\n            axes[shown].set_title(\n                f'True: {true_lbl}\\nPred: {pred_lbl} ({probs[i]:.2f})',\n                color=color, fontsize=9\n            )\n            axes[shown].axis('off')\n            shown += 1\n        if shown >= 8:\n            break\n\nplt.suptitle('EfficientNet-B4 + CBAM — Sample Predictions (Green=Correct, Red=Wrong)')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.677Z"}},"outputs":[],"execution_count":null},{"id":"ed67d7c6-daaf-4575-a95d-3799fff44aad","cell_type":"markdown","source":"## Cell 28 — Save Final Model & Config","metadata":{}},{"id":"e2417039-bb4a-4c18-9202-abb4e33be22c","cell_type":"code","source":"import json as _json\n\ntorch.save(model_b4.state_dict(), '/kaggle/working/efnb4_cbam_final.pth')\n\nconfig_b4 = {\n    'architecture': 'EfficientNet-B4 + CBAM',\n    'backbone':     'efficientnet_b4',\n    'img_size':      380,\n    'batch_size':    16,\n    'cbam_reduction': 16,\n    'best_val_auc':  round(best_auc_b4, 4),\n    'test_auc':      round(auc_b4, 4),\n    'threshold':     round(float(best_thresh_b4), 2)\n}\n\nwith open('/kaggle/working/efnb4_cbam_config.json', 'w') as f:\n    _json.dump(config_b4, f, indent=2)\n\nprint('Saved: efnb4_cbam_final.pth')\nprint('Saved: efnb4_cbam_config.json')\nprint()\nfor k, v in config_b4.items():\n    print(f'  {k}: {v}')","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.677Z"}},"outputs":[],"execution_count":null},{"id":"023e8e9a-080d-4f80-b5ad-89aefd8faf75","cell_type":"markdown","source":"---\n# Part 3 — Explainability & Evaluation\n### GradCAM Heatmaps + Model Comparison Table\n\n---","metadata":{}},{"id":"b225d45d-18a9-4b5d-b385-25ba30fa8ab2","cell_type":"markdown","source":"## Cell 29 — GradCAM Implementation\n\n**What is GradCAM?**\n\nGradCAM (Gradient-weighted Class Activation Mapping) shows *where* the model\nis looking when it makes a prediction. It produces a heatmap overlaid on the\noriginal image — red/warm regions = high attention, blue/cool = ignored.\n\nFor a retinal disease project this is critical — it lets you verify the model\nis focusing on medically relevant regions (optic disc, macula, blood vessels)\nrather than noise or image artefacts.\n\n**How it works:**\n1. Run a forward pass and get the predicted class\n2. Backpropagate the class score back to the last conv layer\n3. Weight each feature map channel by its gradient magnitude\n4. Average and ReLU → heatmap","metadata":{}},{"id":"595a80e0-dc7d-43fb-852a-7ece16f7fe77","cell_type":"code","source":"import cv2\nimport torch.nn.functional as F\n\n\nclass GradCAM:\n    \"\"\"\n    GradCAM for EfficientNetB4_CBAM.\n    Hooks onto the last conv layer of the backbone.\n    \"\"\"\n    def __init__(self, model):\n        self.model = model\n        self.gradients  = None\n        self.activations = None\n        self._register_hooks()\n\n    def _register_hooks(self):\n        # Last stage of EfficientNet-B4 backbone\n        target_layer = self.model.backbone.blocks[-1]\n\n        def forward_hook(module, input, output):\n            self.activations = output.detach()\n\n        def backward_hook(module, grad_input, grad_output):\n            self.gradients = grad_output[0].detach()\n\n        target_layer.register_forward_hook(forward_hook)\n        target_layer.register_full_backward_hook(backward_hook)\n\n    def generate(self, image_tensor, class_idx=None):\n        \"\"\"\n        Args:\n            image_tensor : (1, 3, H, W) on device\n            class_idx    : 0=Non-PDR, 1=PDR. If None, uses predicted class.\n        Returns:\n            heatmap (H, W) numpy array in [0, 1]\n        \"\"\"\n        self.model.eval()\n        self.model.zero_grad()\n\n        output = self.model(image_tensor)          \n        if class_idx is None:\n            class_idx = output.argmax(dim=1).item()\n\n        # Backprop the target class score\n        output[0, class_idx].backward()\n\n        # Global average pool the gradients\n        weights = self.gradients.mean(dim=(2, 3), keepdim=True)  # (1, C, 1, 1)\n        cam     = (weights * self.activations).sum(dim=1, keepdim=True)  # (1, 1, H, W)\n        cam     = F.relu(cam)\n\n       \n        cam = cam.squeeze().cpu().numpy()\n        cam = (cam - cam.min()) / (cam.max() - cam.min() + 1e-8)\n        return cam, class_idx\n\n\ngradcam = GradCAM(model_b4)\nprint('GradCAM ready.')","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.677Z"}},"outputs":[],"execution_count":null},{"id":"4122b631-5492-4e91-8625-5da8073999f5","cell_type":"markdown","source":"## Cell 30 — GradCAM Heatmaps on Test Images\n\nShows 8 test images with the attention heatmap overlaid.\nWarm colours (red/yellow) = regions the model weighted most heavily.\nFor PDR cases you should see focus around the optic disc and vascular regions.","metadata":{}},{"id":"5e4a8e17-5bf7-4348-8804-2787f60746ea","cell_type":"code","source":"inv_normalize = transforms.Normalize(\n    mean=[-0.485/0.229, -0.456/0.224, -0.406/0.225],\n    std=[1/0.229, 1/0.224, 1/0.225]\n)\n\nfig, axes = plt.subplots(4, 4, figsize=(16, 16))\n\n# Row layout: (orig, overlay, orig, overlay) x 4 rows = 8 images\n\nshown = 0\nfor images, labels in test_loader_b4:\n    for i in range(len(images)):\n        if shown >= 8:\n            break\n\n        img_tensor = images[i].unsqueeze(0).to(device)\n        img_tensor.requires_grad_(False)\n\n        \n        cam, pred_idx = gradcam.generate(img_tensor, class_idx=None)\n\n        \n        orig = np.clip(inv_normalize(images[i]).permute(1,2,0).numpy(), 0, 1)\n        orig_uint8 = (orig * 255).astype(np.uint8)\n\n        \n        cam_resized = cv2.resize(cam, (orig.shape[1], orig.shape[0]))\n        heatmap = cv2.applyColorMap(\n            (cam_resized * 255).astype(np.uint8), cv2.COLORMAP_JET\n        )\n        heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n\n       \n        overlay = (0.5 * orig_uint8 + 0.5 * heatmap).astype(np.uint8)\n\n        true_lbl = 'PDR' if labels[i] == 1 else 'Non-PDR'\n        pred_lbl = 'PDR' if pred_idx == 1 else 'Non-PDR'\n        color    = 'green' if pred_idx == labels[i].item() else 'red'\n\n        row = shown // 2\n        col_orig    = (shown % 2) * 2\n        col_overlay = col_orig + 1\n\n        axes[row, col_orig].imshow(orig)\n        axes[row, col_orig].set_title(f'Original\\nTrue: {true_lbl}', fontsize=8)\n        axes[row, col_orig].axis('off')\n\n        axes[row, col_overlay].imshow(overlay)\n        axes[row, col_overlay].set_title(\n            f'GradCAM\\nPred: {pred_lbl}',\n            color=color, fontsize=8\n        )\n        axes[row, col_overlay].axis('off')\n\n        shown += 1\n    if shown >= 8:\n        break\n\nplt.suptitle(\n    'GradCAM Heatmaps — Warm regions = model attention\\n'\n    '(Green title = correct, Red title = wrong)',\n    fontsize=12\n)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.677Z"}},"outputs":[],"execution_count":null},{"id":"b718bdac-952a-4a32-8bbf-92f107470d05","cell_type":"markdown","source":"## Cell 31 — Model Comparison Table\n\nSide-by-side comparison of ResNet50 baseline vs EfficientNet-B4 + CBAM\nacross all key metrics.","metadata":{}},{"id":"c5caa2ca-f773-4cc2-8e7d-43acbfd5b60d","cell_type":"code","source":"from sklearn.metrics import precision_score, recall_score\n\n# ── ResNet50 metrics \nres50_precision = precision_score(test_true, test_pred,  pos_label=1, zero_division=0)\nres50_recall    = recall_score(   test_true, test_pred,  pos_label=1, zero_division=0)\nres50_f1        = f1_score(       test_true, test_pred,  pos_label=1, zero_division=0)\nres50_auc       = auc_score\nres50_specificity = recall_score( test_true, test_pred,  pos_label=0, zero_division=0)\n\n# ── EfficientNet-B4 + CBAM metrics \nb4_precision    = precision_score(test_true_b4, test_pred_b4, pos_label=1, zero_division=0)\nb4_recall       = recall_score(   test_true_b4, test_pred_b4, pos_label=1, zero_division=0)\nb4_f1           = f1_score(       test_true_b4, test_pred_b4, pos_label=1, zero_division=0)\nb4_auc          = auc_b4\nb4_specificity  = recall_score(   test_true_b4, test_pred_b4, pos_label=0, zero_division=0)\n\n# ── Build comparison DataFrame ────────────────────────────────────────────\ncomparison = pd.DataFrame({\n    'Metric': [\n        'Backbone',\n        'Input Resolution',\n        'Attention Module',\n        'Trainable Params',\n        'AUC-ROC',\n        'PDR Precision',\n        'PDR Recall (Sensitivity)',\n        'PDR F1-Score',\n        'Specificity (Non-PDR Recall)',\n        'Decision Threshold'\n    ],\n    'ResNet50 (Baseline)': [\n        'ResNet50',\n        '224 x 224',\n        'None',\n        '~2M (head only, Phase 1)',\n        f'{res50_auc:.4f}',\n        f'{res50_precision:.4f}',\n        f'{res50_recall:.4f}',\n        f'{res50_f1:.4f}',\n        f'{res50_specificity:.4f}',\n        f'{best_thresh:.2f}'\n    ],\n    'EfficientNet-B4 + CBAM': [\n        'EfficientNet-B4',\n        '380 x 380',\n        'CBAM (Channel + Spatial)',\n        '~19M total',\n        f'{b4_auc:.4f}',\n        f'{b4_precision:.4f}',\n        f'{b4_recall:.4f}',\n        f'{b4_f1:.4f}',\n        f'{b4_specificity:.4f}',\n        f'{best_thresh_b4:.2f}'\n    ]\n})\n\n\nprint('=' * 70)\nprint(comparison.to_string(index=False))\nprint('=' * 70)\n\n\nmetrics  = ['AUC-ROC', 'PDR F1-Score', 'PDR Recall\\n(Sensitivity)', 'Specificity']\nres50_vals = [res50_auc, res50_f1, res50_recall, res50_specificity]\nb4_vals    = [b4_auc,    b4_f1,    b4_recall,    b4_specificity]\n\nx = np.arange(len(metrics))\nwidth = 0.35\n\nfig, ax = plt.subplots(figsize=(10, 5))\nbars1 = ax.bar(x - width/2, res50_vals, width, label='ResNet50 (Baseline)',\n               color='steelblue', alpha=0.85)\nbars2 = ax.bar(x + width/2, b4_vals,    width, label='EfficientNet-B4 + CBAM',\n               color='tomato',    alpha=0.85)\n\nax.set_ylim([0, 1.12])\nax.set_xticks(x)\nax.set_xticklabels(metrics)\nax.set_ylabel('Score')\nax.set_title('Model Comparison — ResNet50 vs EfficientNet-B4 + CBAM')\nax.legend()\nax.grid(axis='y', alpha=0.3)\n\n\nfor bar in bars1:\n    ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.01,\n            f'{bar.get_height():.3f}', ha='center', va='bottom', fontsize=9)\nfor bar in bars2:\n    ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.01,\n            f'{bar.get_height():.3f}', ha='center', va='bottom', fontsize=9)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"execution_failed":"2026-04-11T14:27:02.677Z"}},"outputs":[],"execution_count":null}]}