{"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":31193,"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","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-11-16T06:27:13.432792Z","iopub.execute_input":"2025-11-16T06:27:13.433469Z","iopub.status.idle":"2025-11-16T06:27:21.230477Z","shell.execute_reply.started":"2025-11-16T06:27:13.433446Z","shell.execute_reply":"2025-11-16T06:27:21.229796Z"}},"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()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T06:27:29.921517Z","iopub.execute_input":"2025-11-16T06:27:29.922217Z","iopub.status.idle":"2025-11-16T06:27:29.986165Z","shell.execute_reply.started":"2025-11-16T06:27:29.922194Z","shell.execute_reply":"2025-11-16T06:27:29.985394Z"}},"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-16T06:27:51.087176Z","iopub.execute_input":"2025-11-16T06:27:51.088051Z","iopub.status.idle":"2025-11-16T06:27:51.382620Z","shell.execute_reply.started":"2025-11-16T06:27:51.088007Z","shell.execute_reply":"2025-11-16T06:27:51.382029Z"}},"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)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T06:27:56.679689Z","iopub.execute_input":"2025-11-16T06:27:56.680648Z","iopub.status.idle":"2025-11-16T06:27:58.312624Z","shell.execute_reply.started":"2025-11-16T06:27:56.680610Z","shell.execute_reply":"2025-11-16T06:27:58.311852Z"}},"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.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T06:28:02.480716Z","iopub.execute_input":"2025-11-16T06:28:02.481217Z","iopub.status.idle":"2025-11-16T06:28:02.492381Z","shell.execute_reply.started":"2025-11-16T06:28:02.481191Z","shell.execute_reply":"2025-11-16T06:28:02.491804Z"}},"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-11-16T06:28:07.584609Z","iopub.execute_input":"2025-11-16T06:28:07.585338Z","iopub.status.idle":"2025-11-16T06:28:09.050447Z","shell.execute_reply.started":"2025-11-16T06:28:07.585314Z","shell.execute_reply":"2025-11-16T06:28:09.049698Z"}},"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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T06:28:31.340126Z","iopub.execute_input":"2025-11-16T06:28:31.340402Z","iopub.status.idle":"2025-11-16T06:28:31.347550Z","shell.execute_reply.started":"2025-11-16T06:28:31.340382Z","shell.execute_reply":"2025-11-16T06:28:31.346623Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 8: MODEL (timm DenseNet) ====================\n# If timm not installed: !pip install timm\nimport timm\n\ndef get_densenet(num_classes=5, pretrained=True):\n    model = timm.create_model('densenet121', pretrained=pretrained, num_classes=num_classes)\n    return model\n\nmodel = get_densenet(num_classes=5, pretrained=True)\nmodel = model.to(DEVICE)\nprint(\"DenseNet121 model loaded.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T06:28:57.985883Z","iopub.execute_input":"2025-11-16T06:28:57.986156Z","iopub.status.idle":"2025-11-16T06:29:07.382335Z","shell.execute_reply.started":"2025-11-16T06:28:57.986135Z","shell.execute_reply":"2025-11-16T06:29:07.381716Z"}},"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-16T06:29:23.750166Z","iopub.execute_input":"2025-11-16T06:29:23.750693Z","iopub.status.idle":"2025-11-16T06:29:23.757313Z","shell.execute_reply.started":"2025-11-16T06:29:23.750644Z","shell.execute_reply":"2025-11-16T06:29:23.756757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 10: TRAIN LOOP (DenseNet) ====================\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\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\n        pbar.set_postfix({'loss': f'{running_loss/total:.4f}',\n                          '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} \"\n          f\"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(), \"densenet_best.pth\")\n        print(f\"Saved best model with val_qwk={best_qwk:.4f}\")\n\n    scheduler.step()\n\ntorch.save(model.state_dict(), \"densenet_last.pth\")\nprint(\"Training finished.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T06:29:37.809473Z","iopub.execute_input":"2025-11-16T06:29:37.810225Z","iopub.status.idle":"2025-11-16T07:57:24.152194Z","shell.execute_reply.started":"2025-11-16T06:29:37.810202Z","shell.execute_reply":"2025-11-16T07:57:24.151362Z"}},"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()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:00:10.956285Z","iopub.execute_input":"2025-11-16T08:00:10.956977Z","iopub.status.idle":"2025-11-16T08:00:11.292315Z","shell.execute_reply.started":"2025-11-16T08:00:10.956953Z","shell.execute_reply":"2025-11-16T08:00:11.291754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==================== CELL 12: TEST EVALUATION (DenseNet) ====================\nckpt_path = \"densenet_best.pth\"\nmodel.load_state_dict(torch.load(ckpt_path, map_location=DEVICE))\nmodel.to(DEVICE)\nmodel.eval()\n\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\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()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-11-16T08:00:18.926510Z","iopub.execute_input":"2025-11-16T08:00:18.926773Z","iopub.status.idle":"2025-11-16T08:01:08.899341Z","shell.execute_reply.started":"2025-11-16T08:00:18.926753Z","shell.execute_reply":"2025-11-16T08:01:08.898667Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}