{"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, StratifiedKFold\nfrom sklearn.metrics import accuracy_score, classification_report, confusion_matrix, cohen_kappa_score, f1_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# circular crop helper\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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-06T13:28:12.263048Z","iopub.execute_input":"2025-11-06T13:28:12.263199Z","iopub.status.idle":"2025-11-06T13:28:18.054462Z","shell.execute_reply.started":"2025-11-06T13:28:12.263185Z","shell.execute_reply":"2025-11-06T13:28:18.053676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 2: PATHS & DATAFRAME ====================\nDATA_ROOT = Path('/kaggle/input/aptos2019-blindness-detection')\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)\nprint(\"Total samples:\", len(train_df))\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T13:28:26.427484Z","iopub.execute_input":"2025-11-06T13:28:26.428106Z","iopub.status.idle":"2025-11-06T13:28:26.479949Z","shell.execute_reply.started":"2025-11-06T13:28:26.428074Z","shell.execute_reply":"2025-11-06T13:28:26.479130Z"}},"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()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T13:28:30.569761Z","iopub.execute_input":"2025-11-06T13:28:30.570025Z","iopub.status.idle":"2025-11-06T13:28:30.814272Z","shell.execute_reply.started":"2025-11-06T13:28:30.570004Z","shell.execute_reply":"2025-11-06T13:28:30.813527Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 4: TRANSFORMS (PRE-RESIZE) ====================\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# InceptionV3 expects 299×299 input images\nIMG_SIZE = 299\nBATCH_SIZE = 8   # reduce to 4 if GPU runs out of memory\n\n# --- TRAINING AUGMENTATIONS ---\n# Using RandomCrop (stable across albumentations versions)\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([\n        A.GaussNoise(var_limit=(10.0, 50.0), p=0.5),\n        A.MultiplicativeNoise(multiplier=(0.9, 1.1), p=0.5)\n    ], p=0.2),\n    A.Normalize(mean=(0.485, 0.456, 0.406),\n                std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\n# --- VALIDATION AUGMENTATIONS ---\nvalid_transforms = A.Compose([\n    A.Resize(height=IMG_SIZE, width=IMG_SIZE),\n    A.Normalize(mean=(0.485, 0.456, 0.406),\n                std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\nprint(\"✅ Transforms ready.\")\nprint(f\"IMG_SIZE = {IMG_SIZE},  BATCH_SIZE = {BATCH_SIZE}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T13:32:56.046697Z","iopub.execute_input":"2025-11-06T13:32:56.047250Z","iopub.status.idle":"2025-11-06T13:32:56.062021Z","shell.execute_reply.started":"2025-11-06T13:32:56.047226Z","shell.execute_reply":"2025-11-06T13:32:56.061284Z"}},"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):\n        self.df = df.reset_index(drop=True)\n        self.transforms = transforms\n        self.img_size = int(img_size)\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 (helps stability)\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.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T13:33:19.279775Z","iopub.execute_input":"2025-11-06T13:33:19.280353Z","iopub.status.idle":"2025-11-06T13:33:19.291215Z","shell.execute_reply.started":"2025-11-06T13:33:19.280328Z","shell.execute_reply":"2025-11-06T13:33:19.290388Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 6: SPLIT & DATALOADERS (single split) ====================\nSEED = 42\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\ntrain_ds = APTOSDataset(train_df, transforms=train_transforms, img_size=IMG_SIZE)\nval_ds   = APTOSDataset(val_df,   transforms=valid_transforms, img_size=IMG_SIZE)\ntest_ds  = APTOSDataset(test_df,  transforms=valid_transforms, img_size=IMG_SIZE)\n\nnum_workers = 2\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=num_workers, pin_memory=True)\nval_loader   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False, num_workers=num_workers, pin_memory=True)\ntest_loader  = DataLoader(test_ds,  batch_size=BATCH_SIZE, shuffle=False, num_workers=num_workers, pin_memory=True)\n\nimgs, labels = next(iter(train_loader))\nprint(\"Sample batch shapes:\", imgs.shape, labels.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T13:33:31.645416Z","iopub.execute_input":"2025-11-06T13:33:31.646157Z","iopub.status.idle":"2025-11-06T13:33:34.240054Z","shell.execute_reply.started":"2025-11-06T13:33:31.646130Z","shell.execute_reply":"2025-11-06T13:33:34.239121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 7: FOCAL LOSS & CLASS WEIGHTS ====================\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\n# Compute class weights (inverse frequency) for optional use in loss\ncounts = train_df['diagnosis'].value_counts().sort_index().values\nclass_weights = torch.tensor(1.0 / (counts + 1e-8), dtype=torch.float32)\nclass_weights = class_weights / class_weights.sum() * len(counts)  # normalize scale\nprint(\"Class counts:\", counts)\nprint(\"Class weights (normalized):\", class_weights.numpy())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T13:33:43.500989Z","iopub.execute_input":"2025-11-06T13:33:43.501293Z","iopub.status.idle":"2025-11-06T13:33:43.524858Z","shell.execute_reply.started":"2025-11-06T13:33:43.501270Z","shell.execute_reply":"2025-11-06T13:33:43.524228Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 8: MODEL (InceptionV3) - robust for different torchvision versions ====================\nimport torchvision.models as models\nimport torch.nn as nn\nimport torchvision\nimport warnings\n\ndef get_inception_v3(num_classes=5, pretrained=True, aux_logits=False, freeze_backbone=False):\n    \"\"\"\n    Robust loader for torchvision's inception_v3 across versions.\n    - If torchvision expects aux_logits=True for the pretrained weights, we still load the weights\n      but then wrap the model so forward() returns a single logits tensor.\n    - Returns a nn.Module that outputs (B, num_classes) logits only.\n    \"\"\"\n    # Helper wrapper to make forward always return single logits tensor\n    class InceptionWrapper(nn.Module):\n        def __init__(self, model):\n            super().__init__()\n            self.model = model\n\n        def forward(self, x):\n            out = self.model(x)\n            # torchvision Inception may return (logits, aux_logits) depending on aux_logits flag\n            if isinstance(out, tuple) or isinstance(out, list):\n                return out[0]\n            return out\n\n    model = None\n    # Newer torchvision uses 'weights=' API\n    try:\n        # Try to use the weights enum if available (torchvision >= 0.13+)\n        if pretrained:\n            # find available attribute names for weights enum\n            # This will fail if torchvision older; wrapped in try/except\n            try:\n                Weights = models.Inception_V3_Weights  # type: ignore\n                weights = Weights.IMAGENET1K_V1\n                # Some torchvision versions will force aux_logits=True for pretrained weights.\n                # We pass aux_logits=True when required and later wrap to return a single output.\n                model = models.inception_v3(weights=weights, aux_logits=aux_logits)\n            except Exception:\n                # fallback: older signature may be models.inception_v3(pretrained=True, aux_logits=...)\n                model = models.inception_v3(pretrained=True, aux_logits=aux_logits)\n        else:\n            # no pretrained weights requested\n            model = models.inception_v3(pretrained=False, aux_logits=aux_logits)\n    except ValueError as e:\n        # Common case: pretrained weights require aux_logits=True in this torchvision version.\n        # Re-call with aux_logits=True, then wrap.\n        warnings.warn(f\"InceptionV3 loader raised ValueError: {e}. Retrying with aux_logits=True and wrapping the output.\")\n        try:\n            # try with weights=... if available\n            try:\n                Weights = models.Inception_V3_Weights  # type: ignore\n                weights = Weights.IMAGENET1K_V1 if pretrained else None\n                model = models.inception_v3(weights=weights, aux_logits=True)\n            except Exception:\n                model = models.inception_v3(pretrained=pretrained, aux_logits=True)\n        except Exception as e2:\n            # As a last fallback, create without pretrained weights\n            warnings.warn(f\"Failed to load pretrained inception_v3 due to: {e2}. Creating model without pretrained weights.\")\n            model = models.inception_v3(pretrained=False, aux_logits=False)\n\n    # Replace final fc to match num_classes\n    # InceptionV3 has attribute `fc`\n    in_features = model.fc.in_features\n    model.fc = nn.Linear(in_features, num_classes)\n\n    # Optionally freeze backbone parameters (keep final fc trainable)\n    if freeze_backbone:\n        for p in model.parameters():\n            p.requires_grad = False\n        for p in model.fc.parameters():\n            p.requires_grad = True\n\n    # Wrap to ensure forward returns single logits tensor (works for both aux_logits True/False)\n    model = InceptionWrapper(model)\n    return model\n\n# Create model (set freeze_backbone=True to train only fc first)\nmodel = get_inception_v3(num_classes=5, pretrained=True, aux_logits=False, freeze_backbone=False)\nmodel = model.to(DEVICE)\nprint(\"InceptionV3 model loaded. Trainable params:\", sum(p.numel() for p in model.parameters() if p.requires_grad))\n\nbest_ckpt = \"inception_v3_best.pth\"\nlast_ckpt = \"inception_v3_last.pth\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T13:36:08.292428Z","iopub.execute_input":"2025-11-06T13:36:08.293036Z","iopub.status.idle":"2025-11-06T13:36:09.318419Z","shell.execute_reply.started":"2025-11-06T13:36:08.293013Z","shell.execute_reply":"2025-11-06T13:36:09.317685Z"}},"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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T13:36:36.312733Z","iopub.execute_input":"2025-11-06T13:36:36.313425Z","iopub.status.idle":"2025-11-06T13:36:36.319895Z","shell.execute_reply.started":"2025-11-06T13:36:36.313401Z","shell.execute_reply":"2025-11-06T13:36:36.319114Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 10: TRAIN LOOP (MIXED PRECISION + BEST CHECKPOINT) ====================\nEPOCHS = 12\n# use focal loss, optionally pass class_weights as alpha\nalpha = class_weights.to(DEVICE)  # try using this or set to None\ncriterion = FocalLoss(gamma=2.0, alpha=alpha, logits=True)\n\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\n# Optional: MixUp helper (commented out by default)\n# def mixup_data(x, y, alpha=0.4):\n#     if alpha <= 0:\n#         return x, y, None, 1.0\n#     lam = np.random.beta(alpha, alpha)\n#     batch_size = x.size()[0]\n#     index = torch.randperm(batch_size).to(x.device)\n#     mixed_x = lam * x + (1 - lam) * x[index, :]\n#     y_a, y_b = y, y[index]\n#     return mixed_x, y_a, y_b, lam\n\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\n        optimizer.zero_grad()\n        with torch.amp.autocast(device_type=device_type):\n            logits = model(imgs)\n            loss = criterion(logits, targets)\n\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(), best_ckpt)\n        print(f\"Saved best model with val_qwk={best_qwk:.4f}\")\n\n    scheduler.step()\n\ntorch.save(model.state_dict(), last_ckpt)\nprint(\"Training finished.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T13:37:04.819579Z","iopub.execute_input":"2025-11-06T13:37:04.820139Z","iopub.status.idle":"2025-11-06T14:22:06.488077Z","shell.execute_reply.started":"2025-11-06T13:37:04.820116Z","shell.execute_reply":"2025-11-06T14:22:06.487075Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 11: PLOT TRAINING METRICS ====================\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-11-06T14:24:02.055537Z","iopub.execute_input":"2025-11-06T14:24:02.055951Z","iopub.status.idle":"2025-11-06T14:24:02.413216Z","shell.execute_reply.started":"2025-11-06T14:24:02.055923Z","shell.execute_reply":"2025-11-06T14:24:02.412502Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 12: TEST EVALUATION (with optional TTA) ====================\nfrom sklearn.metrics import classification_report, confusion_matrix\n\n# load best\nckpt_path = best_ckpt\nmodel.load_state_dict(torch.load(ckpt_path, map_location=DEVICE))\nmodel.to(DEVICE)\nmodel.eval()\n\n# Simple inference function\n@torch.no_grad()\ndef inference(model, loader):\n    all_preds = []\n    all_targets = []\n    for imgs, targets in loader:\n        imgs = imgs.to(DEVICE)\n        logits = model(imgs)\n        preds = torch.softmax(logits, dim=1).argmax(dim=1).cpu().numpy()\n        all_preds.extend(preds.tolist())\n        all_targets.extend(targets.numpy().tolist())\n    return np.array(all_targets), np.array(all_preds)\n\ny_true, y_pred = inference(model, test_loader)\ntest_acc = accuracy_score(y_true, y_pred)\ntest_qwk = qwk(y_true, y_pred)\ntest_f1_macro = f1_score(y_true, y_pred, average='macro')\nprint(f\"✅ Loaded checkpoint: {ckpt_path}\")\nprint(f\"Test Accuracy: {test_acc:.4f}\")\nprint(f\"Test QWK: {test_qwk:.4f}\")\nprint(f\"Test F1 (macro): {test_f1_macro:.4f}\\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'); plt.ylabel('Actual'); plt.title('Confusion Matrix'); plt.colorbar(); plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T14:24:26.828721Z","iopub.execute_input":"2025-11-06T14:24:26.829424Z","iopub.status.idle":"2025-11-06T14:24:55.815362Z","shell.execute_reply.started":"2025-11-06T14:24:26.829402Z","shell.execute_reply":"2025-11-06T14:24:55.814663Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 13: GRAD-CAM VISUALIZATION ====================\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\nimport os\n\n# --- Grad-CAM class ---\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        # find last Conv2d layer if not specified\n        if target_layer is None:\n            for m in reversed(list(self.model.modules())):\n                if isinstance(m, torch.nn.Conv2d):\n                    target_layer = m\n                    break\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        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)\n        B = logits.shape[0]\n\n        if target_class is None:\n            target_class = logits.softmax(dim=1).argmax(dim=1).cpu().tolist()\n        elif isinstance(target_class, int):\n            target_class = [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]\n            acts = self.activations[i]\n            weights = grads.mean(dim=(1, 2))\n            cam = (weights.view(-1, 1, 1) * acts).sum(dim=0)\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            if self.smooth:\n                cam_np = gaussian_filter(cam_np, sigma=self.sigma)\n            cams.append(cam_np)\n            self.model.zero_grad()\n\n        cams = np.stack(cams, axis=0)\n        _, _, H, W = input_tensor.shape\n        cams_t = torch.from_numpy(cams).unsqueeze(1).to(self.device)\n        cams_resized = F.interpolate(cams_t, size=(H, W), mode=\"bilinear\", align_corners=False)\n        cams_resized = cams_resized.squeeze(1).cpu().numpy()\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            out.append(c)\n        return np.stack(out, axis=0)\n\n# --- helper to unnormalize ---\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\n# --- Generate and save overlays ---\nos.makedirs(\"gradcam_inception_outputs\", exist_ok=True)\ncam_tool = GradCAM(model, smooth=True, sigma=3)\n\n# choose indices to visualize\nindices = [0, 10, 25, 50]\nprint(\"Producing Grad-CAM overlays for indices:\", indices)\n\nfor idx in indices:\n    img_t, label = test_ds[idx]\n    inp = img_t.unsqueeze(0).to(DEVICE)\n    cam_map = cam_tool.generate_cam(inp)[0]\n    cam_uint8 = np.uint8(255 * cam_map)\n    heatmap = cv2.applyColorMap(cam_uint8, cv2.COLORMAP_JET)\n    heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB) / 255.0\n    orig = unnormalize(img_t)\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\n    base = f\"idx{idx}_lbl{label}\"\n    cv2.imwrite(f\"gradcam_inception_outputs/{base}_orig.png\", cv2.cvtColor(orig, cv2.COLOR_RGB2BGR))\n    cv2.imwrite(f\"gradcam_inception_outputs/{base}_heatmap.png\", cv2.cvtColor((heatmap * 255).astype(np.uint8), cv2.COLOR_RGB2BGR))\n    cv2.imwrite(f\"gradcam_inception_outputs/{base}_overlay.png\", cv2.cvtColor((overlay * 255).astype(np.uint8), cv2.COLOR_RGB2BGR))\n\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_inception_outputs/\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T14:33:24.155051Z","iopub.execute_input":"2025-11-06T14:33:24.155798Z","iopub.status.idle":"2025-11-06T14:33:26.406322Z","shell.execute_reply.started":"2025-11-06T14:33:24.155773Z","shell.execute_reply":"2025-11-06T14:33:26.405360Z"}},"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(\"inception_v3_test_predictions.csv\", index=False)\nprint(\"✅ Saved test predictions → inception_v3_test_predictions.csv\")\n\n# quick peek\nprint(out_df.head())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-06T14:33:42.584698Z","iopub.execute_input":"2025-11-06T14:33:42.584989Z","iopub.status.idle":"2025-11-06T14:33:42.601655Z","shell.execute_reply.started":"2025-11-06T14:33:42.584969Z","shell.execute_reply":"2025-11-06T14:33:42.600773Z"}},"outputs":[],"execution_count":null}]}