{"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"}],"dockerImageVersionId":31154,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ==================== CELL 1: IMPORTS & SETUP ====================\nimport os, random, warnings\nfrom pathlib import Path\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\nfrom tqdm import tqdm\n\n# sklearn\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score, classification_report, confusion_matrix, cohen_kappa_score\n\n# PyTorch\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\n# misc\nwarnings.filterwarnings(\"ignore\")\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)\n\n# Define circular_crop early so worker processes always have it\ndef circular_crop(img):\n    \"\"\"Crop to the largest bright contour (retina). Input: BGR image from cv2.\"\"\"\n    if img is None:\n        return None\n    try:\n        gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    except Exception:\n        return None\n    _, thresh = cv2.threshold(gray, 10, 255, cv2.THRESH_BINARY)\n    contours, _ = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    if not contours:\n        return img\n    cnt = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(cnt)\n    cropped = img[y:y+h, x:x+w]\n    return cropped\n\n# small helper for display/testing\ndef preprocess_for_display(path, img_size=512):\n    img = cv2.imread(path)\n    if img is None:\n        return np.zeros((img_size, img_size, 3), dtype=np.float32)\n    c = circular_crop(img)\n    if c is None:\n        return np.zeros((img_size, img_size, 3), dtype=np.float32)\n    img = cv2.cvtColor(c, cv2.COLOR_BGR2RGB)\n    img = cv2.resize(img, (img_size, img_size))\n    img = img.astype(np.float32) / 255.0\n    return img\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-10-16T19:19:09.852759Z","iopub.execute_input":"2025-10-16T19:19:09.852998Z","iopub.status.idle":"2025-10-16T19:19:20.808878Z","shell.execute_reply.started":"2025-10-16T19:19:09.852974Z","shell.execute_reply":"2025-10-16T19:19:20.808116Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 2: PATHS & DATAFRAME ====================\nDATA_ROOT = Path('/kaggle/input/aptos2019-blindness-detection')  # change if needed\nassert DATA_ROOT.exists(), f\"Dataset not found at {DATA_ROOT}\"\n\nTRAIN_CSV = DATA_ROOT / 'train.csv'\nTRAIN_DIR = DATA_ROOT / 'train_images'\n\ntrain_df = pd.read_csv(TRAIN_CSV)\ntrain_df['path'] = train_df['id_code'].map(lambda x: str(TRAIN_DIR / f\"{x}.png\"))\ntrain_df['diagnosis'] = train_df['diagnosis'].astype(int)\ntrain_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T19:19:27.393208Z","iopub.execute_input":"2025-10-16T19:19:27.393771Z","iopub.status.idle":"2025-10-16T19:19:27.465281Z","shell.execute_reply.started":"2025-10-16T19:19:27.393741Z","shell.execute_reply":"2025-10-16T19:19:27.464684Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 3: QUICK EDA ====================\nprint(\"Class counts:\\n\", train_df['diagnosis'].value_counts())\nplt.figure(figsize=(6,4))\ntrain_df['diagnosis'].value_counts().sort_index().plot(kind='bar')\nplt.title('Class distribution')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T19:19:35.040108Z","iopub.execute_input":"2025-10-16T19:19:35.040807Z","iopub.status.idle":"2025-10-16T19:19:35.390559Z","shell.execute_reply.started":"2025-10-16T19:19:35.040772Z","shell.execute_reply":"2025-10-16T19:19:35.389767Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 4: TRANSFORMS (PRE-RESIZE) ====================\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nIMG_SIZE = 512  # change to 384 if OOM\n\n# Because we pre-resize in Dataset, use RandomCrop not RandomResizedCrop\ntrain_transforms = A.Compose([\n    A.RandomCrop(height=IMG_SIZE, width=IMG_SIZE, p=1.0),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.2),\n    A.ShiftScaleRotate(shift_limit=0.06, scale_limit=0.06, rotate_limit=15, p=0.5),\n    A.RandomBrightnessContrast(p=0.5),\n    A.OneOf([A.GaussNoise(), A.MultiplicativeNoise()], p=0.2),\n    A.Normalize(mean=(0.485,0.456,0.406), std=(0.229,0.224,0.225)),\n    ToTensorV2()\n])\n\nvalid_transforms = A.Compose([\n    A.Resize(height=IMG_SIZE, width=IMG_SIZE),\n    A.Normalize(mean=(0.485,0.456,0.406), std=(0.229,0.224,0.225)),\n    ToTensorV2()\n])\n\nprint(\"Transforms ready. IMG_SIZE =\", IMG_SIZE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T19:19:44.912146Z","iopub.execute_input":"2025-10-16T19:19:44.912439Z","iopub.status.idle":"2025-10-16T19:19:46.683322Z","shell.execute_reply.started":"2025-10-16T19:19:44.912417Z","shell.execute_reply":"2025-10-16T19:19:46.682512Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 5: ROBUST DATASET (PRE-RESIZE) ====================\nclass APTOSDataset(Dataset):\n    def __init__(self, df, transforms=None, img_size=IMG_SIZE, debug_failures=0):\n        self.df = df.reset_index(drop=True)\n        self.transforms = transforms\n        self.img_size = int(img_size)\n        self._fail_count = 0\n        self._debug_limit = debug_failures\n\n    def __len__(self):\n        return len(self.df)\n\n    def _safe_resize_and_to_tensor(self, img):\n        if img is None:\n            img = np.zeros((self.img_size, self.img_size, 3), dtype=np.uint8)\n        img_resized = cv2.resize(img, (self.img_size, self.img_size), interpolation=cv2.INTER_AREA)\n        img_resized = img_resized.astype(np.float32) / 255.0\n        mean = np.array((0.485, 0.456, 0.406), dtype=np.float32)\n        std  = np.array((0.229, 0.224, 0.225), dtype=np.float32)\n        img_resized = (img_resized - mean) / std\n        img_tensor = torch.from_numpy(img_resized).permute(2, 0, 1).float()\n        return img_tensor\n\n    def __getitem__(self, idx):\n        row = self.df.loc[idx]\n        path = row['path']\n        label = int(row['diagnosis'])\n\n        img = cv2.imread(path)\n        if img is None:\n            img = np.zeros((self.img_size, self.img_size, 3), dtype=np.uint8)\n        else:\n            img = circular_crop(img)\n            if img is None:\n                img = np.zeros((self.img_size, self.img_size, 3), dtype=np.uint8)\n            else:\n                img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        # Pre-resize to fixed size BEFORE passing to Albumentations\n        try:\n            img = cv2.resize(img, (self.img_size, self.img_size), interpolation=cv2.INTER_AREA)\n        except Exception:\n            return self._safe_resize_and_to_tensor(img), label\n\n        if self.transforms:\n            try:\n                sample = self.transforms(image=img)\n                img_out = sample['image']\n                if isinstance(img_out, np.ndarray):\n                    img_tensor = torch.from_numpy(img_out.astype(np.float32)).permute(2,0,1)\n                elif torch.is_tensor(img_out):\n                    img_tensor = img_out\n                else:\n                    img_tensor = self._safe_resize_and_to_tensor(img)\n                if img_tensor.dim() != 3 or img_tensor.shape[1] != self.img_size or img_tensor.shape[2] != self.img_size:\n                    img_tensor = self._safe_resize_and_to_tensor(img)\n            except Exception:\n                img_tensor = self._safe_resize_and_to_tensor(img)\n        else:\n            img_tensor = self._safe_resize_and_to_tensor(img)\n\n        img_tensor = img_tensor.float()\n        if img_tensor.dim() == 2:\n            img_tensor = img_tensor.unsqueeze(0).repeat(3,1,1)\n        if img_tensor.shape[0] != 3:\n            if img_tensor.shape[-1] == 3 and img_tensor.dim() == 3:\n                img_tensor = img_tensor.permute(2,0,1)\n            else:\n                img_tensor = self._safe_resize_and_to_tensor(img)\n\n        return img_tensor, label\n\nprint(\"Dataset class ready.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T19:19:56.488532Z","iopub.execute_input":"2025-10-16T19:19:56.489061Z","iopub.status.idle":"2025-10-16T19:19:56.500194Z","shell.execute_reply.started":"2025-10-16T19:19:56.489034Z","shell.execute_reply":"2025-10-16T19:19:56.499498Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 6: SPLIT & DATALOADERS ====================\nSEED = 42\nBATCH_SIZE = 8  # reduce to 4 if OOM\n\ntrain_df_, test_df = train_test_split(train_df, test_size=0.10, stratify=train_df['diagnosis'], random_state=SEED)\ntrain_df, val_df = train_test_split(train_df_, test_size=0.10, stratify=train_df_['diagnosis'], random_state=SEED)\n\nprint(f\"Sizes -> train: {len(train_df)}, val: {len(val_df)}, test: {len(test_df)}\")\n\n# debug_failures=0 to silence fallback prints\ntrain_ds = APTOSDataset(train_df, transforms=train_transforms, img_size=IMG_SIZE, debug_failures=0)\nval_ds   = APTOSDataset(val_df,   transforms=valid_transforms, img_size=IMG_SIZE, debug_failures=0)\ntest_ds  = APTOSDataset(test_df,  transforms=valid_transforms, img_size=IMG_SIZE, debug_failures=0)\n\n# num_workers=0 avoids worker-process import issues during debugging\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=0, pin_memory=True)\nval_loader   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False, num_workers=0, pin_memory=True)\ntest_loader  = DataLoader(test_ds,  batch_size=BATCH_SIZE, shuffle=False, num_workers=0, pin_memory=True)\n\n# quick sanity\nimgs, labels = next(iter(train_loader))\nprint(\"Sample batch shapes:\", imgs.shape, labels.shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T19:20:06.441551Z","iopub.execute_input":"2025-10-16T19:20:06.442271Z","iopub.status.idle":"2025-10-16T19:20:07.935745Z","shell.execute_reply.started":"2025-10-16T19:20:06.442247Z","shell.execute_reply":"2025-10-16T19:20:07.934928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 7: FOCAL LOSS ====================\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=2.0, alpha=None, reduction='mean', logits=True):\n        super().__init__()\n        self.gamma = gamma\n        self.alpha = alpha\n        self.reduction = reduction\n        self.logits = logits\n\n    def forward(self, inputs, targets):\n        if self.logits:\n            ce_loss = F.cross_entropy(inputs, targets, reduction='none')\n            pt = torch.exp(-ce_loss)\n            loss = ((1 - pt) ** self.gamma) * ce_loss\n            if self.alpha is not None:\n                at = self.alpha.gather(0, targets)\n                loss = at * loss\n        else:\n            pt = inputs.gather(1, targets.unsqueeze(1)).squeeze(1)\n            loss = -((1 - pt) ** self.gamma) * torch.log(pt + 1e-12)\n            if self.alpha is not None:\n                at = self.alpha.gather(0, targets)\n                loss = at * loss\n\n        if self.reduction == 'mean':\n            return loss.mean()\n        elif self.reduction == 'sum':\n            return loss.sum()\n        return loss\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T19:20:26.136167Z","iopub.execute_input":"2025-10-16T19:20:26.136474Z","iopub.status.idle":"2025-10-16T19:20:26.144033Z","shell.execute_reply.started":"2025-10-16T19:20:26.136453Z","shell.execute_reply":"2025-10-16T19:20:26.143438Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 8: MODEL (timm EfficientNet) ====================\n# If timm not installed on Kaggle/Colab, run: !pip install timm\nimport timm\n\ndef get_efficientnet(model_name='tf_efficientnet_b3_ns', num_classes=5, pretrained=True):\n    model = timm.create_model(model_name, pretrained=pretrained, num_classes=num_classes)\n    return model\n\nmodel = get_efficientnet('tf_efficientnet_b3_ns', num_classes=5, pretrained=True)\nmodel = model.to(DEVICE)\nprint(\"Model loaded.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T19:20:38.430962Z","iopub.execute_input":"2025-10-16T19:20:38.431250Z","iopub.status.idle":"2025-10-16T19:20:48.100778Z","shell.execute_reply.started":"2025-10-16T19:20:38.431226Z","shell.execute_reply":"2025-10-16T19:20:48.099957Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 9: METRICS & EVAL HELPERS ====================\nscaler = torch.amp.GradScaler(device=\"cuda\") if torch.cuda.is_available() else torch.amp.GradScaler()\n\ndef qwk(y_true, y_pred):\n    return cohen_kappa_score(y_true, y_pred, weights='quadratic')\n\n@torch.no_grad()\ndef evaluate(model, loader, criterion=None):\n    model.eval()\n    all_preds = []\n    all_targets = []\n    running_loss = 0.0\n    if criterion is None:\n        criterion = FocalLoss(gamma=2.0, logits=True)\n    device_type = \"cuda\" if DEVICE.type == \"cuda\" else \"cpu\"\n    for imgs, targets in loader:\n        imgs = imgs.to(DEVICE, non_blocking=True)\n        targets = targets.to(DEVICE, non_blocking=True)\n        with torch.amp.autocast(device_type=device_type):\n            logits = model(imgs)\n            loss = criterion(logits, targets)\n        preds = torch.softmax(logits, dim=1).argmax(dim=1).cpu().numpy()\n        all_preds.extend(preds.tolist())\n        all_targets.extend(targets.cpu().numpy().tolist())\n        running_loss += loss.item() * imgs.size(0)\n    avg_loss = running_loss / len(loader.dataset)\n    acc = accuracy_score(all_targets, all_preds)\n    q = qwk(all_targets, all_preds)\n    return avg_loss, acc, q, np.array(all_targets), np.array(all_preds)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T19:20:51.111290Z","iopub.execute_input":"2025-10-16T19:20:51.111984Z","iopub.status.idle":"2025-10-16T19:20:51.118406Z","shell.execute_reply.started":"2025-10-16T19:20:51.111959Z","shell.execute_reply":"2025-10-16T19:20:51.117826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 10: TRAIN LOOP (MIXED PRECISION + BEST CHECKPOINT) ====================\nEPOCHS = 12\ncriterion = FocalLoss(gamma=2.0, logits=True)\noptimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-5)\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=EPOCHS)\n\nhistory = {'train_loss':[], 'train_acc':[], 'val_loss':[], 'val_acc':[], 'val_qwk':[]}\nbest_qwk = -1.0\n\nfor epoch in range(EPOCHS):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    pbar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{EPOCHS}\", leave=False)\n    device_type = \"cuda\" if DEVICE.type == \"cuda\" else \"cpu\"\n    for imgs, targets in pbar:\n        imgs = imgs.to(DEVICE, non_blocking=True)\n        targets = targets.to(DEVICE, non_blocking=True)\n        optimizer.zero_grad()\n        with torch.amp.autocast(device_type=device_type):\n            logits = model(imgs)\n            loss = criterion(logits, targets)\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item() * imgs.size(0)\n        preds = logits.softmax(dim=1).argmax(dim=1)\n        correct += (preds == targets).sum().item()\n        total += imgs.size(0)\n        pbar.set_postfix({'loss': f'{running_loss/total:.4f}', 'acc': f'{correct/total:.4f}'})\n\n    train_loss = running_loss / len(train_loader.dataset)\n    train_acc = correct / total\n\n    val_loss, val_acc, val_qwk, _, _ = evaluate(model, val_loader, criterion=criterion)\n\n    history['train_loss'].append(train_loss)\n    history['train_acc'].append(train_acc)\n    history['val_loss'].append(val_loss)\n    history['val_acc'].append(val_acc)\n    history['val_qwk'].append(val_qwk)\n\n    print(f\"Epoch {epoch+1}: train_loss={train_loss:.4f} train_acc={train_acc:.4f} val_loss={val_loss:.4f} val_acc={val_acc:.4f} val_qwk={val_qwk:.4f}\")\n\n    if val_qwk > best_qwk:\n        best_qwk = val_qwk\n        torch.save(model.state_dict(), \"efficientnet_b3_best.pth\")\n        print(f\"Saved best model with val_qwk={best_qwk:.4f}\")\n\n    scheduler.step()\n\ntorch.save(model.state_dict(), \"efficientnet_b3_last.pth\")\nprint(\"Training finished.\")\n   ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T19:21:05.817005Z","iopub.execute_input":"2025-10-16T19:21:05.817787Z","iopub.status.idle":"2025-10-16T20:50:04.736879Z","shell.execute_reply.started":"2025-10-16T19:21:05.817762Z","shell.execute_reply":"2025-10-16T20:50:04.736155Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 11: PLOT TRAINING METRICS ====================\nimport matplotlib.pyplot as plt\n\nplt.figure(figsize=(14,5))\nplt.subplot(1,2,1)\nplt.plot(history['train_loss'], label='train_loss')\nplt.plot(history['val_loss'], label='val_loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.title('Training & Validation Loss')\n\nplt.subplot(1,2,2)\nplt.plot(history['train_acc'], label='train_acc')\nplt.plot(history['val_acc'], label='val_acc')\nplt.plot(history['val_qwk'], label='val_qwk')\nplt.xlabel('Epoch')\nplt.ylabel('Score')\nplt.legend()\nplt.title('Training Accuracy / Val Accuracy / Val QWK')\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T20:51:54.555986Z","iopub.execute_input":"2025-10-16T20:51:54.556579Z","iopub.status.idle":"2025-10-16T20:51:54.875330Z","shell.execute_reply.started":"2025-10-16T20:51:54.556559Z","shell.execute_reply":"2025-10-16T20:51:54.874582Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 12: TEST EVALUATION ====================\nfrom sklearn.metrics import classification_report, confusion_matrix\n\nckpt_path = \"efficientnet_b3_best.pth\"\nmodel.load_state_dict(torch.load(ckpt_path, map_location=DEVICE))\nmodel.to(DEVICE)\nmodel.eval()\n\n# Evaluate on test set\ntest_loss, test_acc, test_qwk, y_true, y_pred = evaluate(\n    model, test_loader, criterion=FocalLoss(gamma=2.0, logits=True)\n)\n\nprint(f\"✅ Loaded checkpoint: {ckpt_path}\")\nprint(f\"Test Loss: {test_loss:.4f}\")\nprint(f\"Test Accuracy: {test_acc:.4f}\")\nprint(f\"Test QWK: {test_qwk:.4f}\\n\")\n\nprint(\"Classification Report:\")\nprint(classification_report(y_true, y_pred, digits=4))\n\n# Confusion Matrix\ncm = confusion_matrix(y_true, y_pred)\nplt.figure(figsize=(6,5))\nplt.imshow(cm, cmap='Blues')\nfor i in range(cm.shape[0]):\n    for j in range(cm.shape[1]):\n        plt.text(j, i, str(cm[i,j]), ha='center', va='center')\nplt.xlabel('Predicted')\nplt.ylabel('Actual')\nplt.title('Confusion Matrix')\nplt.colorbar()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T20:52:16.004165Z","iopub.execute_input":"2025-10-16T20:52:16.004446Z","iopub.status.idle":"2025-10-16T20:53:07.536598Z","shell.execute_reply.started":"2025-10-16T20:52:16.004420Z","shell.execute_reply":"2025-10-16T20:53:07.535944Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== A: Robust Grad-CAM class (batch-safe, smoothing) =====\nimport torch\nimport torch.nn.functional as F\nimport numpy as np\nimport cv2\nfrom scipy.ndimage import gaussian_filter\nimport matplotlib.pyplot as plt\n\nclass GradCAM:\n    def __init__(self, model, target_layer=None, use_cuda=True, smooth=True, sigma=3):\n        self.model = model\n        self.device = next(model.parameters()).device if use_cuda else torch.device(\"cpu\")\n        self.model.to(self.device)\n        self.model.eval()\n        self.gradients = None\n        self.activations = None\n        self.smooth = smooth\n        self.sigma = sigma\n\n        if target_layer is None:\n            for module in reversed(list(self.model.modules())):\n                if isinstance(module, torch.nn.Conv2d):\n                    target_layer = module\n                    break\n            if target_layer is None:\n                raise ValueError(\"No Conv2d layer found in the model.\")\n        self.target_layer = target_layer\n        self._register_hooks()\n\n    def _register_hooks(self):\n        def forward_hook(module, inp, out):\n            self.activations = out.detach()\n        def backward_hook(module, grad_in, grad_out):\n            self.gradients = grad_out[0].detach()\n        self.target_layer.register_forward_hook(forward_hook)\n        try:\n            self.target_layer.register_full_backward_hook(backward_hook)\n        except Exception:\n            self.target_layer.register_backward_hook(backward_hook)\n\n    def generate_cam(self, input_tensor, target_class=None):\n        \"\"\"\n        input_tensor: (B,C,H,W) torch tensor on same device as model\n        target_class: None -> predicted class; or list/array of class indices length B\n        returns: np.array (B, H, W) values in [0,1]\n        \"\"\"\n        assert input_tensor.dim() == 4, \"input_tensor must be 4D (B,C,H,W)\"\n        input_tensor = input_tensor.to(self.device)\n        self.model.zero_grad()\n        logits = self.model(input_tensor)          # (B, num_classes)\n        B = logits.shape[0]\n        if target_class is None:\n            preds = logits.softmax(dim=1).argmax(dim=1)\n            target_class = preds.cpu().tolist()\n        elif isinstance(target_class, (int, np.integer)):\n            target_class = [int(target_class)] * B\n        else:\n            target_class = [int(x) for x in target_class]\n\n        cams = []\n        for i in range(B):\n            score = logits[i, target_class[i]]\n            score.backward(retain_graph=True)\n            grads = self.gradients[i]   # (C, h, w)\n            acts  = self.activations[i] # (C, h, w)\n            weights = grads.mean(dim=(1,2))   # (C,)\n            cam = (weights.view(-1,1,1) * acts).sum(dim=0)  # (h, w)\n            cam = F.relu(cam)\n            cam_np = cam.cpu().numpy()\n            if cam_np.max() > 0:\n                cam_np = (cam_np - cam_np.min()) / (cam_np.max() - cam_np.min() + 1e-8)\n            else:\n                cam_np = np.zeros_like(cam_np)\n            if self.smooth:\n                cam_np = gaussian_filter(cam_np, sigma=self.sigma)\n            cams.append(cam_np)\n            self.model.zero_grad()\n        cams = np.stack(cams, axis=0)  # (B, h, w)\n        _, _, H, W = input_tensor.shape\n        cams_t = torch.from_numpy(cams).unsqueeze(1).to(self.device)  # (B,1,h,w)\n        cams_resized = F.interpolate(cams_t, size=(H,W), mode='bilinear', align_corners=False)\n        cams_resized = cams_resized.squeeze(1).cpu().numpy()  # (B, H, W)\n        out = []\n        for c in cams_resized:\n            if c.max() > 0:\n                c = (c - c.min()) / (c.max() - c.min() + 1e-8)\n            else:\n                c = np.zeros_like(c)\n            out.append(c)\n        return np.stack(out, axis=0)  # (B, H, W)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T21:11:03.263086Z","iopub.execute_input":"2025-10-16T21:11:03.263646Z","iopub.status.idle":"2025-10-16T21:11:03.277742Z","shell.execute_reply.started":"2025-10-16T21:11:03.263622Z","shell.execute_reply":"2025-10-16T21:11:03.277118Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===== B: Use GradCAM to generate and save overlays =====\nimport os\nos.makedirs(\"gradcam_outputs\", exist_ok=True)\n\n# init tool (uses model already on DEVICE)\ncam_tool = GradCAM(model, smooth=True, sigma=3)\n\n# choose mode:\n#  - indices: list of test indices to visualize (0-based from test_ds)\n#  - by_pred_class: if True, chooses one sample per predicted class automatically\nindices = [0, 10, 25, 50]   # <-- change these to visualize particular test images\nby_pred_class = False       # set True to pick one image per predicted class\n\ndef unnormalize(img_tensor):\n    img = img_tensor.permute(1,2,0).cpu().numpy()\n    mean = np.array([0.485,0.456,0.406])\n    std  = np.array([0.229,0.224,0.225])\n    img = np.clip(img * std + mean, 0, 1)\n    return (img * 255).astype(np.uint8)\n\nto_process = []\nif by_pred_class:\n    # pick one example per predicted class\n    preds_all = []\n    with torch.no_grad():\n        for i in range(len(test_ds)):\n            img_t, label = test_ds[i]\n            logits = model(img_t.unsqueeze(0).to(DEVICE))\n            preds_all.append(logits.softmax(dim=1).argmax(dim=1).item())\n    preds_all = np.array(preds_all)\n    for cls in range(5):\n        idxs = np.where(preds_all==cls)[0]\n        if len(idxs)>0:\n            to_process.append(int(idxs[0]))\nelse:\n    to_process = indices\n\nprint(\"Will produce overlays for indices:\", to_process)\n\nfor idx in to_process:\n    img_t, label = test_ds[idx]           # (C,H,W)\n    inp = img_t.unsqueeze(0).to(DEVICE)   # (1,C,H,W)\n    cam_maps = cam_tool.generate_cam(inp) # (1,H,W)\n    cam = cam_maps[0]                     # HxW float [0,1]\n    # prepare heatmap\n    cam_uint8 = np.uint8(255 * cam)\n    heatmap = cv2.applyColorMap(cam_uint8, cv2.COLORMAP_JET)          # BGR\n    heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB) / 255.0        # RGB float\n    orig = unnormalize(img_t)                                        # HxW x3 uint8\n    orig_f = orig.astype(np.float32)/255.0\n    overlay = 0.4 * heatmap + 0.6 * orig_f\n    overlay = np.clip(overlay, 0, 1)\n    # save images: original, heatmap, overlay\n    base = f\"idx{idx}_lbl{label}\"\n    cv2.imwrite(f\"gradcam_outputs/{base}_orig.png\", cv2.cvtColor(orig, cv2.COLOR_RGB2BGR))\n    cv2.imwrite(f\"gradcam_outputs/{base}_heatmap.png\", cv2.cvtColor((heatmap*255).astype(np.uint8), cv2.COLOR_RGB2BGR))\n    # overlay converted to BGR uint8 for cv2\n    cv2.imwrite(f\"gradcam_outputs/{base}_overlay.png\", cv2.cvtColor((overlay*255).astype(np.uint8), cv2.COLOR_RGB2BGR))\n    # also display inline\n    plt.figure(figsize=(10,4))\n    plt.subplot(1,3,1); plt.imshow(orig); plt.title(f\"Orig idx={idx} label={label}\"); plt.axis('off')\n    plt.subplot(1,3,2); plt.imshow(heatmap); plt.title(\"Heatmap\"); plt.axis('off')\n    plt.subplot(1,3,3); plt.imshow(overlay); plt.title(\"Overlay\"); plt.axis('off')\n    plt.show()\n\nprint(\"Saved Grad-CAM images to ./gradcam_outputs/\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T21:11:14.222433Z","iopub.execute_input":"2025-10-16T21:11:14.222764Z","iopub.status.idle":"2025-10-16T21:11:16.516812Z","shell.execute_reply.started":"2025-10-16T21:11:14.222743Z","shell.execute_reply":"2025-10-16T21:11:16.516125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 14: SAVE TEST PREDICTIONS ====================\nout_df = test_df.copy().reset_index(drop=True)\nout_df['true'] = y_true\nout_df['pred'] = y_pred\nout_df.to_csv('efficientnet_b3_test_predictions.csv', index=False)\nprint(\"Saved test predictions → efficientnet_b3_test_predictions.csv\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-10-16T21:11:46.611706Z","iopub.execute_input":"2025-10-16T21:11:46.612279Z","iopub.status.idle":"2025-10-16T21:11:46.631364Z","shell.execute_reply.started":"2025-10-16T21:11:46.612256Z","shell.execute_reply":"2025-10-16T21:11:46.630742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}