{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431},{"sourceType":"kernelVersion","sourceId":29830115}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Diabetic Retinopathy Classification\n## ResNet50 + CBAM | APTOS 2019 Blindness Detection\n### Pipeline: CLAHE · Ben Graham · CORN Loss · QWK Evaluation","metadata":{"_uuid":"3395bd66-9ebf-43ca-bdd5-c2a7f77e95eb","_cell_guid":"1b86ac60-7272-45f7-ad79-3159c58be563","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"markdown","source":"## Cell 1 — Install Dependencies","metadata":{"_uuid":"a4551993-682a-482e-afdc-807fa8df164e","_cell_guid":"663eaba0-4ed6-4dd9-9bac-77e0261cd318","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"!pip install -q coral-pytorch opencv-python-headless albumentations","metadata":{"_uuid":"b872f4d5-b760-47d9-a759-bddbeaaff073","_cell_guid":"3e4be1b3-106a-47af-a733-b63971d904e1","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:06:09.940459Z","iopub.execute_input":"2026-05-06T12:06:09.940714Z","iopub.status.idle":"2026-05-06T12:06:13.188132Z","shell.execute_reply.started":"2026-05-06T12:06:09.940692Z","shell.execute_reply":"2026-05-06T12:06:13.187129Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nprint(os.listdir(\"/kaggle/input/competitions/aptos2019-blindness-detection\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-06T12:06:13.190027Z","iopub.execute_input":"2026-05-06T12:06:13.190770Z","iopub.status.idle":"2026-05-06T12:06:13.195193Z","shell.execute_reply.started":"2026-05-06T12:06:13.190731Z","shell.execute_reply":"2026-05-06T12:06:13.194602Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 2 — Imports","metadata":{"_uuid":"f44b9b3d-35e4-4e7b-b927-3dc575b7f13d","_cell_guid":"a834c44d-bfbe-46d6-bb61-3c3c19e2f917","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"import os\nimport cv2\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom PIL import Image\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision import models, transforms\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    cohen_kappa_score,\n    classification_report,\n    precision_recall_fscore_support,\n)\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n# Reproducibility\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {DEVICE}\")","metadata":{"_uuid":"c6454581-b69b-4fb6-8758-77289b5e2d22","_cell_guid":"f8710495-5e6a-4c9c-9bc2-f4bd8989e7ff","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:06:13.196177Z","iopub.execute_input":"2026-05-06T12:06:13.196447Z","iopub.status.idle":"2026-05-06T12:06:13.210297Z","shell.execute_reply.started":"2026-05-06T12:06:13.196420Z","shell.execute_reply":"2026-05-06T12:06:13.209722Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 3 — Configuration","metadata":{"_uuid":"e2b08c75-e522-4ebb-bc29-bfbc526c3bbb","_cell_guid":"64025bfb-ecc6-45b5-b353-2b6b2dc81f09","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class CFG:\n    # Paths (Kaggle environment)\n    TRAIN_CSV  =\"/kaggle/input/competitions/aptos2019-blindness-detection/train.csv\"\n    TEST_CSV   = \"/kaggle/input/competitions/aptos2019-blindness-detection/test.csv\"\n    TRAIN_IMG  = \"/kaggle/input/competitions/aptos2019-blindness-detection/train_images\"\n    TEST_IMG   = \"/kaggle/input/competitions/aptos2019-blindness-detection/test_images\"\n\n    # Image\n    IMG_SIZE   = 224          # ResNet50 default input size\n    CHANNELS   = 3\n\n    # Training\n    EPOCHS        = 30\n    BATCH_SIZE    = 32\n    NUM_WORKERS   = 4\n    LR            = 1e-4\n    WEIGHT_DECAY  = 1e-5\n    VAL_SPLIT     = 0.15\n\n    # Loss blending\n    CORN_WEIGHT   = 0.5       # weight for CORN loss\n    WCE_WEIGHT    = 0.5       # weight for Weighted Cross-Entropy loss\n\n    # Classes\n    NUM_CLASSES = 5           # DR grades 0-4\n\n    # CLAHE\n    CLAHE_CLIP_LIMIT    = 2.0\n    CLAHE_TILE_GRID     = (8, 8)\n\n    # Gaussian blur sigma for Ben Graham\n    BEN_GRAHAM_SIGMA    = 10\n    BEN_GRAHAM_ALPHA    = 4\n    BEN_GRAHAM_BETA     = -4\n    BEN_GRAHAM_GAMMA    = 128","metadata":{"_uuid":"168d1bd9-3b0e-41fe-84ed-baffb49d621c","_cell_guid":"14480b84-9915-4302-844f-4e91ce031a43","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:11:26.103597Z","iopub.execute_input":"2026-05-06T12:11:26.104221Z","iopub.status.idle":"2026-05-06T12:11:26.109189Z","shell.execute_reply.started":"2026-05-06T12:11:26.104192Z","shell.execute_reply":"2026-05-06T12:11:26.108356Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 4 — Preprocessing Utilities","metadata":{"_uuid":"edd22241-e20e-4424-96e6-f27580054967","_cell_guid":"6b5824a1-d64e-4162-aa10-248f9ca957e1","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def crop_black_borders(image: np.ndarray, tol: int = 7) -> np.ndarray:\n    \"\"\"Remove black borders by thresholding pixel brightness.\"\"\"\n    if image.ndim == 2:\n        mask = image > tol\n    else:\n        mask = image.max(axis=2) > tol\n    rows = np.any(mask, axis=1)\n    cols = np.any(mask, axis=0)\n    if not rows.any() or not cols.any():\n        return image\n    rmin, rmax = np.where(rows)[0][[0, -1]]\n    cmin, cmax = np.where(cols)[0][[0, -1]]\n    return image[rmin:rmax+1, cmin:cmax+1]\n\n\ndef apply_circular_mask(image: np.ndarray) -> np.ndarray:\n    \"\"\"Zero-out pixels outside the inscribed circle (fundus mask).\"\"\"\n    h, w = image.shape[:2]\n    cx, cy = w // 2, h // 2\n    radius = min(cx, cy)\n    Y, X = np.ogrid[:h, :w]\n    dist = np.sqrt((X - cx) ** 2 + (Y - cy) ** 2)\n    mask = dist <= radius\n    result = image.copy()\n    if image.ndim == 3:\n        result[~mask] = 0\n    else:\n        result[~mask] = 0\n    return result\n\n\ndef extract_green_channel(image: np.ndarray) -> np.ndarray:\n    \"\"\"Extract green channel (highest contrast for DR features) and return as RGB.\"\"\"\n    green = image[:, :, 1]\n    return cv2.merge([green, green, green])\n\n\ndef apply_clahe(image: np.ndarray) -> np.ndarray:\n    \"\"\"Apply CLAHE per channel in LAB colour space.\"\"\"\n    lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n    clahe = cv2.createCLAHE(\n        clipLimit=CFG.CLAHE_CLIP_LIMIT,\n        tileGridSize=CFG.CLAHE_TILE_GRID,\n    )\n    lab[:, :, 0] = clahe.apply(lab[:, :, 0])\n    return cv2.cvtColor(lab, cv2.COLOR_LAB2RGB)\n\n\ndef ben_graham_preprocessing(image: np.ndarray, sigma: int = None) -> np.ndarray:\n    \"\"\"\n    Ben Graham's technique:\n        output = alpha * image + beta * gaussian_blur(image) + gamma\n    Enhances local contrast while suppressing global illumination variations.\n    \"\"\"\n    sigma = sigma or CFG.BEN_GRAHAM_SIGMA\n    blurred = cv2.GaussianBlur(\n        image,\n        (0, 0),\n        sigmaX=sigma,\n        sigmaY=sigma,\n    )\n    output = cv2.addWeighted(\n        image,\n        CFG.BEN_GRAHAM_ALPHA,\n        blurred,\n        CFG.BEN_GRAHAM_BETA,\n        CFG.BEN_GRAHAM_GAMMA,\n    )\n    return np.clip(output, 0, 255).astype(np.uint8)\n\n\ndef preprocess_fundus_image(image_path: str, img_size: int = CFG.IMG_SIZE) -> np.ndarray:\n    \"\"\"\n    Full preprocessing pipeline:\n      1. Read image (BGR → RGB)\n      2. Crop black borders\n      3. Resize\n      4. Circular mask\n      5. Green channel extraction\n      6. CLAHE\n      7. Ben Graham preprocessing\n    \"\"\"\n    img = cv2.imread(str(image_path))\n    if img is None:\n        raise FileNotFoundError(f\"Cannot read image: {image_path}\")\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n    # Step 1: Crop black borders\n    img = crop_black_borders(img)\n\n    # Step 2: Resize\n    img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_LANCZOS4)\n\n    # Step 3: Circular mask\n    img = apply_circular_mask(img)\n\n    # Step 4: Green channel\n    img = extract_green_channel(img)\n\n    # Step 5: CLAHE\n    img = apply_clahe(img)\n\n    # Step 6: Ben Graham\n    img = ben_graham_preprocessing(img)\n\n    return img","metadata":{"_uuid":"a1a2e7df-3e4a-4d8b-9778-5caf64c0dd15","_cell_guid":"774436a8-7ace-4a14-bec4-79a9ca662587","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:11:28.064956Z","iopub.execute_input":"2026-05-06T12:11:28.065660Z","iopub.status.idle":"2026-05-06T12:11:28.077676Z","shell.execute_reply.started":"2026-05-06T12:11:28.065630Z","shell.execute_reply":"2026-05-06T12:11:28.076800Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 5 — Dataset & Augmentation","metadata":{"_uuid":"e2bf7f77-ad58-4c01-a115-4022c3b0f49d","_cell_guid":"95c6e434-468a-42d2-b6d1-69b713877d2f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def get_augmentation(phase: str) -> A.Compose:\n    \"\"\"Return albumentations pipeline for train/val phases.\"\"\"\n    if phase == \"train\":\n        return A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.Rotate(limit=30, p=0.6),\n            A.RandomBrightnessContrast(p=0.3),\n            A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=15, p=0.4),\n            A.CoarseDropout(max_holes=8, max_height=16, max_width=16, p=0.2),\n            A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n            ToTensorV2(),\n        ])\n    else:  # val / test\n        return A.Compose([\n            A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n            ToTensorV2(),\n        ])\n\n\nclass APTOSDataset(Dataset):\n    def __init__(\n        self,\n        df: pd.DataFrame,\n        img_dir,                   # accepts str or Path\n        phase: str = \"train\",\n        img_size: int = CFG.IMG_SIZE,\n        cache: bool = False,\n    ):\n        self.df        = df.reset_index(drop=True)\n        self.img_dir   = Path(img_dir)          # ← force Path here\n        self.phase     = phase\n        self.img_size  = img_size\n        self.transform = get_augmentation(phase)\n        self.cache     = cache\n        self._cache: dict = {}\n\n    def __len__(self):\n        return len(self.df)\n\n    def _load(self, idx: int) -> np.ndarray:\n        if self.cache and idx in self._cache:\n            return self._cache[idx]\n        row  = self.df.iloc[idx]\n        path = self.img_dir / f\"{row['id_code']}.png\"   # now works\n        img  = preprocess_fundus_image(path, self.img_size)\n        if self.cache:\n            self._cache[idx] = img\n        return img\n\n    def __getitem__(self, idx: int):\n        img   = self._load(idx)\n        label = int(self.df.iloc[idx][\"diagnosis\"]) if \"diagnosis\" in self.df.columns else -1\n        aug   = self.transform(image=img)\n        return aug[\"image\"], torch.tensor(label, dtype=torch.long)\n\n\ndef build_weighted_sampler(labels: list) -> WeightedRandomSampler:\n    \"\"\"\n    Compute per-sample weights inversely proportional to class frequency\n    to tackle class imbalance (WeightedRandomSampler).\n    \"\"\"\n    class_counts = np.bincount(labels)\n    class_weights = 1.0 / class_counts.astype(np.float32)\n    sample_weights = class_weights[labels]\n    sampler = WeightedRandomSampler(\n        weights=torch.from_numpy(sample_weights),\n        num_samples=len(sample_weights),\n        replacement=True,\n    )\n    return sampler","metadata":{"_uuid":"3589d3cd-3871-4fd7-a76f-8a89f6f555f1","_cell_guid":"3c9fd08e-376a-4392-b20a-996441bb1d26","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:14:49.870821Z","iopub.execute_input":"2026-05-06T12:14:49.871211Z","iopub.status.idle":"2026-05-06T12:14:49.882231Z","shell.execute_reply.started":"2026-05-06T12:14:49.871182Z","shell.execute_reply":"2026-05-06T12:14:49.881653Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 6 — CBAM Module","metadata":{"_uuid":"56846284-3c74-4f2d-b125-ace5b563ffb8","_cell_guid":"cc00006d-2ee9-4840-885e-5109eb61b57e","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class ChannelAttention(nn.Module):\n    \"\"\"\n    CBAM — Channel Attention sub-module.\n    Generates channel descriptors via both average-pool and max-pool branches,\n    then blends through a shared MLP and sigmoid gate.\n    \"\"\"\n    def __init__(self, in_channels: int, reduction_ratio: int = 16):\n        super().__init__()\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n        mid = max(in_channels // reduction_ratio, 1)\n        self.mlp = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(in_channels, mid, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Linear(mid, in_channels, bias=False),\n        )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        avg_out = self.mlp(self.avg_pool(x))\n        max_out = self.mlp(self.max_pool(x))\n        scale   = torch.sigmoid(avg_out + max_out).unsqueeze(-1).unsqueeze(-1)\n        return x * scale\n\n\nclass SpatialAttention(nn.Module):\n    \"\"\"\n    CBAM — Spatial Attention sub-module.\n    Aggregates channel-wise statistics (avg + max) and convolves to a\n    single-channel spatial gate.\n    \"\"\"\n    def __init__(self, kernel_size: int = 7):\n        super().__init__()\n        self.conv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size // 2, bias=False)\n        self.bn   = nn.BatchNorm2d(1)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        avg_out = x.mean(dim=1, keepdim=True)\n        max_out = x.max(dim=1, keepdim=True).values\n        concat  = torch.cat([avg_out, max_out], dim=1)\n        scale   = torch.sigmoid(self.bn(self.conv(concat)))\n        return x * scale\n\n\nclass CBAM(nn.Module):\n    \"\"\"Convolutional Block Attention Module (Woo et al., 2018).\"\"\"\n    def __init__(self, in_channels: int, reduction_ratio: int = 16, kernel_size: int = 7):\n        super().__init__()\n        self.channel_att = ChannelAttention(in_channels, reduction_ratio)\n        self.spatial_att = SpatialAttention(kernel_size)\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x = self.channel_att(x)\n        x = self.spatial_att(x)\n        return x","metadata":{"_uuid":"9bdb8f23-7a32-4b36-9fa1-7deb75182afd","_cell_guid":"691fe6aa-b2f3-4cf6-8388-e40d6506aedc","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:11:36.794307Z","iopub.execute_input":"2026-05-06T12:11:36.794598Z","iopub.status.idle":"2026-05-06T12:11:36.804053Z","shell.execute_reply.started":"2026-05-06T12:11:36.794565Z","shell.execute_reply":"2026-05-06T12:11:36.803327Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 7 — ResNet50 + CBAM Model","metadata":{"_uuid":"43180431-b291-4580-81d1-1cbd1e1b26a4","_cell_guid":"21eccb7b-2711-4509-823e-bbeb8a07104a","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"class ResNet50CBAM(nn.Module):\n    \"\"\"\n    ResNet50 backbone with CBAM attention injected after every residual stage.\n    Head: GlobalAvgPool → Dropout → FC(num_classes).\n\n    The model can output either:\n      - raw logits  (for Weighted Cross-Entropy)\n      - cumulative logits via the CORN head (for CORN loss)\n    \"\"\"\n    def __init__(\n        self,\n        num_classes: int = CFG.NUM_CLASSES,\n        pretrained: bool = True,\n        dropout: float = 0.4,\n    ):\n        super().__init__()\n        backbone = models.resnet50(\n            weights=models.ResNet50_Weights.IMAGENET1K_V1 if pretrained else None\n        )\n\n        # ── Stem + Pool (unchanged) ──────────────────────────────────────\n        self.stem = nn.Sequential(\n            backbone.conv1, backbone.bn1, backbone.relu, backbone.maxpool\n        )\n\n        # ── Residual stages with CBAM injected after each ───────────────\n        self.layer1 = backbone.layer1\n        self.cbam1  = CBAM(256)\n\n        self.layer2 = backbone.layer2\n        self.cbam2  = CBAM(512)\n\n        self.layer3 = backbone.layer3\n        self.cbam3  = CBAM(1024)\n\n        self.layer4 = backbone.layer4\n        self.cbam4  = CBAM(2048)\n\n        # ── Classification head ──────────────────────────────────────────\n        self.global_pool = nn.AdaptiveAvgPool2d(1)\n        self.dropout     = nn.Dropout(p=dropout)\n        self.classifier  = nn.Linear(2048, num_classes)\n\n        # ── CORN ordinal head (num_classes − 1 binary tasks) ────────────\n        self.corn_head   = nn.Linear(2048, num_classes - 1)\n\n        self._init_head()\n\n    def _init_head(self):\n        nn.init.kaiming_normal_(self.classifier.weight)\n        nn.init.zeros_(self.classifier.bias)\n        nn.init.kaiming_normal_(self.corn_head.weight)\n        nn.init.zeros_(self.corn_head.bias)\n\n    def forward_features(self, x: torch.Tensor) -> torch.Tensor:\n        x = self.stem(x)\n        x = self.cbam1(self.layer1(x))\n        x = self.cbam2(self.layer2(x))\n        x = self.cbam3(self.layer3(x))\n        x = self.cbam4(self.layer4(x))\n        x = self.global_pool(x).flatten(1)\n        return self.dropout(x)\n\n    def forward(self, x: torch.Tensor):\n        feat      = self.forward_features(x)\n        logits    = self.classifier(feat)      # (B, num_classes) — for WCE\n        corn_logits = self.corn_head(feat)     # (B, num_classes-1) — for CORN\n        return logits, corn_logits","metadata":{"_uuid":"8f18a557-a2db-4b7f-aa7b-b0fe0be1686d","_cell_guid":"84923d14-8842-41ca-86a1-405488aada9a","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:11:39.697728Z","iopub.execute_input":"2026-05-06T12:11:39.697981Z","iopub.status.idle":"2026-05-06T12:11:39.707243Z","shell.execute_reply.started":"2026-05-06T12:11:39.697961Z","shell.execute_reply":"2026-05-06T12:11:39.706439Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 8 — Loss Functions","metadata":{"_uuid":"2f85b6ae-7fb8-4892-b1ce-972c17025248","_cell_guid":"4892976e-6692-4e71-9a64-d08178a39a66","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# ────────────────────────────────────────────────────────────────────────────\n# CORN Loss  (Shi et al., 2021 — https://arxiv.org/abs/2111.08851)\n# Pure-PyTorch implementation without the coral_pytorch dependency.\n# ────────────────────────────────────────────────────────────────────────────\n\ndef corn_loss(logits: torch.Tensor, targets: torch.Tensor, num_classes: int) -> torch.Tensor:\n    \"\"\"\n    Conditional Ordinal Regression for Neural Networks (CORN) loss.\n\n    Args:\n        logits  : (B, K-1)  raw outputs from the CORN head\n        targets : (B,)      integer class labels in [0, K-1]\n        num_classes: K\n    Returns:\n        Scalar loss.\n    \"\"\"\n    K = num_classes - 1  # number of binary tasks\n    sets = []\n    for i in range(K):\n        # Task i: predict P(y > i) — only on samples where y > i-1\n        label_i = (targets > i).float()\n        sets.append((logits[:, i], label_i))\n\n    num_examples = targets.size(0)\n    total_loss   = torch.zeros(1, device=logits.device)\n\n    for i, (task_logits, task_labels) in enumerate(sets):\n        if i == 0:\n            # All samples participate in task 0\n            subset = torch.ones(num_examples, dtype=torch.bool, device=logits.device)\n        else:\n            # Only samples where y > i-1 participate\n            subset = targets > (i - 1)\n\n        if subset.sum() == 0:\n            continue\n\n        subset_logits = task_logits[subset]\n        subset_labels = task_labels[subset]\n        total_loss += F.binary_cross_entropy_with_logits(subset_logits, subset_labels)\n\n    return total_loss / K\n\n\ndef weighted_cross_entropy_loss(\n    logits: torch.Tensor,\n    targets: torch.Tensor,\n    class_weights: torch.Tensor,\n) -> torch.Tensor:\n    \"\"\"Standard cross-entropy with per-class weights.\"\"\"\n    criterion = nn.CrossEntropyLoss(weight=class_weights.to(logits.device))\n    return criterion(logits, targets)\n\n\nclass CombinedLoss(nn.Module):\n    \"\"\"\n    alpha * CORN_loss  +  (1-alpha) * WeightedCE_loss\n    \"\"\"\n    def __init__(\n        self,\n        class_weights: torch.Tensor,\n        num_classes: int = CFG.NUM_CLASSES,\n        alpha: float     = CFG.CORN_WEIGHT,\n    ):\n        super().__init__()\n        self.num_classes   = num_classes\n        self.alpha         = alpha\n        self.class_weights = class_weights\n\n    def forward(\n        self,\n        logits:      torch.Tensor,   # (B, num_classes)\n        corn_logits: torch.Tensor,   # (B, num_classes - 1)\n        targets:     torch.Tensor,   # (B,)\n    ) -> torch.Tensor:\n        loss_corn = corn_loss(corn_logits, targets, self.num_classes)\n        loss_wce  = weighted_cross_entropy_loss(logits, targets, self.class_weights)\n        return self.alpha * loss_corn + (1 - self.alpha) * loss_wce","metadata":{"_uuid":"24b4dca0-ce7c-446d-86e6-417234cca83d","_cell_guid":"00d2c7ee-baa9-4ef4-9f53-ac35d89b2487","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:11:42.687546Z","iopub.execute_input":"2026-05-06T12:11:42.688245Z","iopub.status.idle":"2026-05-06T12:11:42.697104Z","shell.execute_reply.started":"2026-05-06T12:11:42.688216Z","shell.execute_reply":"2026-05-06T12:11:42.696425Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 9 — Evaluation Metrics","metadata":{"_uuid":"9a761477-d67f-4e60-ad72-ab1aa4e82d78","_cell_guid":"508eae86-429d-499e-b2c9-11a493d9e240","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def compute_qwk(y_true: list, y_pred: list) -> float:\n    \"\"\"Quadratic Weighted Kappa — primary competition metric.\"\"\"\n    return cohen_kappa_score(y_true, y_pred, weights=\"quadratic\")\n\n\ndef corn_logits_to_predictions(corn_logits: torch.Tensor) -> torch.Tensor:\n    \"\"\"\n    Convert CORN cumulative logits to ordinal class predictions.\n    P(y > k) = sigmoid(logit_k)\n    Predicted class = number of thresholds exceeded.\n    \"\"\"\n    probs = torch.sigmoid(corn_logits)          # (B, K-1)\n    # Predicted rank = first k where P(y > k) < 0.5\n    exceed = (probs > 0.5).long()               # (B, K-1)\n    return exceed.sum(dim=1)                    # (B,)\n\n\ndef print_per_class_metrics(y_true: list, y_pred: list, num_classes: int = CFG.NUM_CLASSES):\n    labels = list(range(num_classes))\n    names  = [f\"Grade {i}\" for i in labels]\n    print(\"\\n\" + \"=\"*60)\n    print(\"PER-CLASS PERFORMANCE\")\n    print(\"=\"*60)\n    print(classification_report(y_true, y_pred, labels=labels, target_names=names, digits=4))\n    p, r, f, _ = precision_recall_fscore_support(y_true, y_pred, labels=labels, zero_division=0)\n    print(f\"{'Grade':<10} {'Precision':>10} {'Recall':>10} {'F1':>10}\")\n    print(\"-\"*42)\n    for i, name in enumerate(names):\n        print(f\"{name:<10} {p[i]:>10.4f} {r[i]:>10.4f} {f[i]:>10.4f}\")\n    print(\"=\"*60 + \"\\n\")","metadata":{"_uuid":"8c08f2fb-8598-4a3e-9df6-e8f813c6b52a","_cell_guid":"292eb6a9-b9ad-4006-83c2-60723f531e67","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:11:46.840690Z","iopub.execute_input":"2026-05-06T12:11:46.840941Z","iopub.status.idle":"2026-05-06T12:11:46.848606Z","shell.execute_reply.started":"2026-05-06T12:11:46.840920Z","shell.execute_reply":"2026-05-06T12:11:46.847953Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 10 — Training & Validation Loops","metadata":{"_uuid":"9283c6f7-5556-4d59-a14e-fbe2e2aa865d","_cell_guid":"d18df8eb-6903-42ee-9dd5-d69315cfca9b","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"def train_one_epoch(\n    model:     nn.Module,\n    loader:    DataLoader,\n    optimizer: torch.optim.Optimizer,\n    criterion: nn.Module,\n    scaler:    torch.cuda.amp.GradScaler,\n    scheduler,\n) -> dict:\n    model.train()\n    running_loss, all_preds, all_labels = 0.0, [], []\n\n    for images, labels in loader:\n        images = images.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.cuda.amp.autocast():\n            logits, corn_logits = model(images)\n            loss = criterion(logits, corn_logits, labels)\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_loss += loss.item() * images.size(0)\n\n        # Use soft-max prediction from CE head for QWK during training\n        preds = logits.argmax(dim=1).cpu().tolist()\n        all_preds.extend(preds)\n        all_labels.extend(labels.cpu().tolist())\n\n    if scheduler is not None:\n        scheduler.step()\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_qwk  = compute_qwk(all_labels, all_preds)\n    return {\"loss\": epoch_loss, \"qwk\": epoch_qwk}\n\n\n@torch.no_grad()\ndef validate(\n    model:     nn.Module,\n    loader:    DataLoader,\n    criterion: nn.Module,\n) -> dict:\n    model.eval()\n    running_loss, all_preds, all_labels = 0.0, [], []\n\n    for images, labels in loader:\n        images = images.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n\n        with torch.cuda.amp.autocast():\n            logits, corn_logits = model(images)\n            loss = criterion(logits, corn_logits, labels)\n\n        running_loss += loss.item() * images.size(0)\n\n        # Ensemble CE prediction + CORN prediction\n        ce_preds   = logits.argmax(dim=1)\n        corn_preds = corn_logits_to_predictions(corn_logits)\n        preds      = ((ce_preds.float() + corn_preds.float()) / 2).round().long()\n\n        all_preds.extend(preds.cpu().tolist())\n        all_labels.extend(labels.cpu().tolist())\n\n    epoch_loss = running_loss / len(loader.dataset)\n    epoch_qwk  = compute_qwk(all_labels, all_preds)\n    return {\"loss\": epoch_loss, \"qwk\": epoch_qwk, \"preds\": all_preds, \"labels\": all_labels}","metadata":{"_uuid":"0d4f64e1-e5b8-44cc-a958-93e6f317fad2","_cell_guid":"9933f1d6-b7bd-4ccf-a070-de0926708a5c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:42:06.591302Z","iopub.execute_input":"2026-05-06T12:42:06.591902Z","iopub.status.idle":"2026-05-06T12:42:06.603213Z","shell.execute_reply.started":"2026-05-06T12:42:06.591864Z","shell.execute_reply":"2026-05-06T12:42:06.602424Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 11 — Data Loading & Preparation","metadata":{"_uuid":"c377e18f-225a-4f19-8092-2c9920ba0b70","_cell_guid":"59e4b751-4c90-4439-8a67-11b464918288","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"df = pd.read_csv(CFG.TRAIN_CSV)\nprint(f\"Total samples : {len(df)}\")\nprint(f\"Class distribution:\\n{df['diagnosis'].value_counts().sort_index()}\\n\")\n\ntrain_df, val_df = train_test_split(\n    df,\n    test_size=CFG.VAL_SPLIT,\n    stratify=df[\"diagnosis\"],\n    random_state=SEED,\n)\nprint(f\"Train : {len(train_df)} | Val : {len(val_df)}\")\n\n# Datasets\ntrain_dataset = APTOSDataset(train_df, CFG.TRAIN_IMG, phase=\"train\")\nval_dataset   = APTOSDataset(val_df,   CFG.TRAIN_IMG, phase=\"val\")\n\n# WeightedRandomSampler for class imbalance\ntrain_labels  = train_df[\"diagnosis\"].tolist()\nsampler       = build_weighted_sampler(train_labels)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=CFG.BATCH_SIZE,\n    sampler=sampler,\n    num_workers=CFG.NUM_WORKERS,\n    pin_memory=True,\n)\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=CFG.BATCH_SIZE,\n    shuffle=False,\n    num_workers=CFG.NUM_WORKERS,\n    pin_memory=True,\n)\n\n# Compute class weights for Weighted Cross-Entropy\nclass_counts  = np.bincount(train_labels, minlength=CFG.NUM_CLASSES)\nclass_weights = torch.tensor(\n    1.0 / (class_counts / class_counts.sum()),\n    dtype=torch.float32,\n)\nclass_weights = class_weights / class_weights.sum() * CFG.NUM_CLASSES\nprint(f\"Class weights (WCE) : {class_weights.numpy().round(4)}\")","metadata":{"_uuid":"9cae981b-bda4-4b8b-8a02-d9bdcc5f230f","_cell_guid":"da3b6453-47e5-4d6b-9310-1d3d15e1ad8b","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:15:04.241751Z","iopub.execute_input":"2026-05-06T12:15:04.242108Z","iopub.status.idle":"2026-05-06T12:15:24.296040Z","shell.execute_reply.started":"2026-05-06T12:15:04.242082Z","shell.execute_reply":"2026-05-06T12:15:24.295280Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 12 — Model, Optimizer & Scheduler","metadata":{"_uuid":"ba9b91f4-aecd-4c24-944f-759df657bedf","_cell_guid":"7246db64-524d-4e38-8699-285cc374a16f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"model = ResNet50CBAM(num_classes=CFG.NUM_CLASSES, pretrained=True).to(DEVICE)\n\n# Differential learning rates: backbone vs head\nbackbone_params = [p for n, p in model.named_parameters()\n                   if not any(k in n for k in [\"classifier\", \"corn_head\", \"cbam\"])]\nhead_params     = [p for n, p in model.named_parameters()\n                   if any(k in n for k in [\"classifier\", \"corn_head\", \"cbam\"])]\n\noptimizer = torch.optim.AdamW([\n    {\"params\": backbone_params, \"lr\": CFG.LR * 0.1},\n    {\"params\": head_params,     \"lr\": CFG.LR},\n], weight_decay=CFG.WEIGHT_DECAY)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer,\n    T_max=CFG.EPOCHS,\n    eta_min=1e-6,\n)\n\ncriterion = CombinedLoss(\n    class_weights=class_weights,\n    num_classes=CFG.NUM_CLASSES,\n    alpha=CFG.CORN_WEIGHT,\n).to(DEVICE)\n\nscaler = torch.cuda.amp.GradScaler()\n\ntotal_params     = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"Total params     : {total_params:,}\")\nprint(f\"Trainable params : {trainable_params:,}\")","metadata":{"_uuid":"643b14a0-9081-4a50-ba13-a5729d6653c8","_cell_guid":"ab5a95dc-f4b0-4fe3-8d3c-49335d730c33","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:30:29.495136Z","iopub.status.idle":"2026-05-06T12:30:29.495496Z","shell.execute_reply.started":"2026-05-06T12:30:29.495299Z","shell.execute_reply":"2026-05-06T12:30:29.495335Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 13 — Training Loop","metadata":{"_uuid":"2ead42c6-bd34-4f2f-84d5-80568d1e91b4","_cell_guid":"fccb84e0-4f45-43ce-920e-ba70b786513c","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"best_qwk     = -np.inf\nbest_weights = None\nhistory      = []\n\nprint(f\"\\n{'Epoch':>6}  {'Train Loss':>10}  {'Train QWK':>10}  {'Val Loss':>10}  {'Val QWK':>10}  {'Status':>8}\")\nprint(\"-\" * 62)\n\nfor epoch in range(1, CFG.EPOCHS + 1):\n    train_stats = train_one_epoch(model, train_loader, optimizer, criterion, scaler, scheduler)\n    val_stats   = validate(model, val_loader, criterion)\n\n    is_best = val_stats[\"qwk\"] > best_qwk\n    if is_best:\n        best_qwk     = val_stats[\"qwk\"]\n        best_weights = {k: v.cpu().clone() for k, v in model.state_dict().items()}\n        torch.save(best_weights, \"/kaggle/working/best_model.pt\")\n        flag = \"✓ BEST\"\n    else:\n        flag = \"\"\n\n    history.append({\n        \"epoch\":      epoch,\n        \"train_loss\": train_stats[\"loss\"],\n        \"train_qwk\":  train_stats[\"qwk\"],\n        \"val_loss\":   val_stats[\"loss\"],\n        \"val_qwk\":    val_stats[\"qwk\"],\n    })\n\n    print(\n        f\"{epoch:>6}  {train_stats['loss']:>10.4f}  {train_stats['qwk']:>10.4f}  \"\n        f\"{val_stats['loss']:>10.4f}  {val_stats['qwk']:>10.4f}  {flag:>8}\"\n    )\n\nprint(f\"\\nBest Validation QWK : {best_qwk:.4f}\")","metadata":{"_uuid":"026845bc-8255-474d-a1ce-522803d2fdb0","_cell_guid":"454fdb13-43b3-4484-9d83-8e6bc31b3819","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"iopub.status.busy":"2026-05-06T12:42:17.767300Z","iopub.execute_input":"2026-05-06T12:42:17.767639Z","execution_failed":"2026-05-06T17:11:24.483Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 14 — Final Evaluation on Validation Set","metadata":{"_uuid":"eb8b3208-734a-4042-acb8-504755f54bad","_cell_guid":"af6464ee-d26d-4837-95d1-9cd4111060ff","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"# Load best weights\nmodel.load_state_dict({k: v.to(DEVICE) for k, v in best_weights.items()})\n\nval_stats = validate(model, val_loader, criterion)\nqwk       = compute_qwk(val_stats[\"labels\"], val_stats[\"preds\"])\n\nprint(f\"\\nFinal Validation QWK : {qwk:.4f}\")\nprint_per_class_metrics(val_stats[\"labels\"], val_stats[\"preds\"])","metadata":{"_uuid":"c3b57fd9-886f-4dc5-bb90-3a30cbad93f2","_cell_guid":"6fcb2747-1fb1-492a-bf63-055956a91369","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"execution_failed":"2026-05-06T17:11:24.574Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 15 — Test Inference & Submission","metadata":{"_uuid":"53e87a52-dffb-4e89-a4b3-b5f7c0c36646","_cell_guid":"74104e21-01eb-47df-be07-b73a5331410f","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"test_df = pd.read_csv(CFG.TEST_CSV)\n\n# Add a dummy diagnosis column so the Dataset works without labels\ntest_df[\"diagnosis\"] = -1\n\ntest_dataset = APTOSDataset(test_df, CFG.TEST_IMG, phase=\"val\")\ntest_loader  = DataLoader(\n    test_dataset,\n    batch_size=CFG.BATCH_SIZE,\n    shuffle=False,\n    num_workers=CFG.NUM_WORKERS,\n    pin_memory=True,\n)\n\nmodel.eval()\nall_test_preds = []\n\nwith torch.no_grad():\n    for images, _ in test_loader:\n        images = images.to(DEVICE, non_blocking=True)\n        with torch.cuda.amp.autocast():\n            logits, corn_logits = model(images)\n\n        ce_preds   = logits.argmax(dim=1)\n        corn_preds = corn_logits_to_predictions(corn_logits)\n        preds      = ((ce_preds.float() + corn_preds.float()) / 2).round().long()\n        all_test_preds.extend(preds.cpu().tolist())\n\nsubmission = pd.DataFrame({\n    \"id_code\":   test_df[\"id_code\"],\n    \"diagnosis\": all_test_preds,\n})\nsubmission.to_csv(\"/kaggle/working/submission.csv\", index=False)\nprint(f\"Submission saved — shape: {submission.shape}\")\nprint(submission[\"diagnosis\"].value_counts().sort_index())","metadata":{"_uuid":"eaab1935-3ff4-4c4a-beb7-adcac7c92301","_cell_guid":"e3211a7c-c08d-4e5a-914e-d63be73b6e77","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"execution_failed":"2026-05-06T17:11:24.580Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 16 — Training History Plot (Optional)","metadata":{"_uuid":"07272441-8fb4-4d65-9350-82bb0c00a7a4","_cell_guid":"a632f478-2258-4362-89de-27c764e123ab","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false}}},{"cell_type":"code","source":"try:\n    import matplotlib.pyplot as plt\n\n    hist_df = pd.DataFrame(history)\n\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n\n    axes[0].plot(hist_df[\"epoch\"], hist_df[\"train_loss\"], label=\"Train Loss\", marker=\"o\")\n    axes[0].plot(hist_df[\"epoch\"], hist_df[\"val_loss\"],   label=\"Val Loss\",   marker=\"s\")\n    axes[0].set_title(\"Loss Curve\")\n    axes[0].set_xlabel(\"Epoch\")\n    axes[0].set_ylabel(\"Loss\")\n    axes[0].legend()\n    axes[0].grid(True, alpha=0.3)\n\n    axes[1].plot(hist_df[\"epoch\"], hist_df[\"train_qwk\"], label=\"Train QWK\", marker=\"o\")\n    axes[1].plot(hist_df[\"epoch\"], hist_df[\"val_qwk\"],   label=\"Val QWK\",   marker=\"s\")\n    axes[1].axhline(y=best_qwk, color=\"red\", linestyle=\"--\", label=f\"Best QWK = {best_qwk:.4f}\")\n    axes[1].set_title(\"Quadratic Weighted Kappa\")\n    axes[1].set_xlabel(\"Epoch\")\n    axes[1].set_ylabel(\"QWK\")\n    axes[1].legend()\n    axes[1].grid(True, alpha=0.3)\n\n    plt.suptitle(\"ResNet50 + CBAM — APTOS 2019 Diabetic Retinopathy\", fontsize=14, fontweight=\"bold\")\n    plt.tight_layout()\n    plt.savefig(\"/kaggle/working/training_history.png\", dpi=150, bbox_inches=\"tight\")\n    plt.show()\n    print(\"Plot saved.\")\nexcept ImportError:\n    print(\"matplotlib not available.\")","metadata":{"_uuid":"545c91bd-0952-4a57-b526-d01400d1913e","_cell_guid":"acab519e-1815-4990-8556-414fe881e251","trusted":true,"collapsed":false,"jupyter":{"outputs_hidden":false},"execution":{"execution_failed":"2026-05-06T17:11:24.582Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}