{"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":"datasetVersion","sourceId":2819730,"datasetId":1723812,"databundleVersionId":2866107}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport torch\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport timm\nfrom torchvision import transforms\nfrom PIL import Image\n\n# ==========================================\n# 1. CONFIGURATION & KAGGLE PATHS\n# ==========================================\nclass Config:\n    APTOS_IMG_DIR = '/kaggle/input/competitions/aptos2019-blindness-detection/train_images/'\n    MESSIDOR_IMG_DIR = '/kaggle/input/datasets/mariaherrerot/messidor2preprocess/messidor-2/messidor-2/preprocess/'\n    \n    IMG_SIZE = 896         # High resolution required by TMIL\n    PATCH_SIZE = 224       # ViT-Small standard input\n    NUM_PATCHES = 16       # (896/224) * (896/224)\n    NUM_CLASSES = 5        # Adjust to 2 if mapping Messidor to Binary\n    BATCH_SIZE = 4         # Keep small due to 16 instances per image\n    EPOCHS = 30\n    LR = 1e-4\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# ==========================================\n# 2. PREPROCESSING & DATASET\n# ==========================================\ndef apply_clahe(image):\n    \"\"\"Enhances foreground lesions and blood vessels.\"\"\"\n    lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n    l_channel, a, b = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    cl = clahe.apply(l_channel)\n    limg = cv2.merge((cl,a,b))\n    return cv2.cvtColor(limg, cv2.COLOR_LAB2RGB)\n\nclass RetinalPatchDataset(Dataset):\n    def __init__(self, image_paths, labels, transform=None):\n        self.image_paths = image_paths\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        img_path = self.image_paths[idx]\n        \n        # Load and preprocess\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        img = apply_clahe(img)\n        \n        # Convert to tensor: (C, H, W)\n        img_tensor = transforms.ToTensor()(img)\n        \n        if self.transform:\n            img_tensor = self.transform(img_tensor)\n\n        # Slice the 896x896 image into 16 non-overlapping 224x224 patches\n        # Output shape will be: (16, 3, 224, 224)\n        patches = img_tensor.unfold(1, Config.PATCH_SIZE, Config.PATCH_SIZE)\\\n                            .unfold(2, Config.PATCH_SIZE, Config.PATCH_SIZE)\n        patches = patches.contiguous().view(3, Config.NUM_PATCHES, Config.PATCH_SIZE, Config.PATCH_SIZE)\n        patches = patches.permute(1, 0, 2, 3) \n        \n        label = torch.tensor(self.labels[idx], dtype=torch.long)\n        return patches, label\n\n# ==========================================\n# 3. TMIL ARCHITECTURE (ViT + GICB)\n# ==========================================\nclass GICB(nn.Module):\n    \"\"\"Global Instance Computing Block to aggregate features across patches.\"\"\"\n    def __init__(self, in_dim, out_dim, heads):\n        super().__init__()\n        self.norm = nn.LayerNorm(in_dim)\n        self.attn = nn.MultiheadAttention(embed_dim=in_dim, num_heads=heads, batch_first=True)\n        self.mlp = nn.Sequential(\n            nn.Linear(in_dim, out_dim),\n            nn.GELU()\n        )\n        self.res_proj = nn.Linear(in_dim, out_dim) if in_dim != out_dim else nn.Identity()\n\n    def forward(self, x):\n        # Multi-Head Self Attention with Residual Connection\n        nx = self.norm(x)\n        attn_out, _ = self.attn(nx, nx, nx)\n        a = attn_out + x \n        \n        # Dimension Reduction MLP with Residual Connection\n        m = self.mlp(a) + self.res_proj(a)\n        return m\n\nclass TMIL_Net(nn.Module):\n    def __init__(self, num_classes=Config.NUM_CLASSES):\n        super().__init__()\n        # Weight-shared feature extractor \n        self.vit = timm.create_model('vit_small_patch16_224', pretrained=True, num_classes=0)\n        \n        # Learnable Position Embedding for the 16 grid locations\n        self.pos_embed = nn.Parameter(torch.zeros(1, Config.NUM_PATCHES, 384))\n        \n        # GICB Blocks (Dimensions match the TMIL paper)\n        self.gicb1 = GICB(in_dim=384, out_dim=128, heads=6)\n        self.gicb2 = GICB(in_dim=128, out_dim=128, heads=4)\n        \n        # Final Classification Head\n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(Config.NUM_PATCHES * 128, 512),\n            nn.ReLU(),\n            nn.Dropout(0.4), # High dropout to prevent overfitting on tiny dataset\n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x):\n        B, N, C, H, W = x.shape\n        # Flatten batch and instances to process all patches simultaneously\n        x = x.view(B * N, C, H, W) \n        \n        features = self.vit(x) # Shape: (B*16, 384)\n        features = features.view(B, N, 384) \n        \n        # Add spatial awareness\n        features = features + self.pos_embed\n        \n        # Calculate global context across all patches\n        out = self.gicb1(features)\n        out = self.gicb2(out)\n        \n        logits = self.classifier(out)\n        return logits\n\n# ==========================================\n# 4. FN-CRUSHING LOSS & MIXUP\n# ==========================================\nclass WeightedFocalLoss(nn.Module):\n    \"\"\"Heavily penalizes False Negatives by scaling disease classes.\"\"\"\n    def __init__(self, alpha, gamma=2.0):\n        super().__init__()\n        self.alpha = alpha.to(Config.DEVICE)\n        self.gamma = gamma\n\n    def forward(self, inputs, targets):\n        ce_loss = F.cross_entropy(inputs, targets, reduction='none', weight=self.alpha)\n        pt = torch.exp(-ce_loss)\n        focal_loss = ((1 - pt) ** self.gamma) * ce_loss\n        return focal_loss.mean()\n\ndef mixup_data(x, y, alpha=0.2):\n    \"\"\"Blends two high-res multi-instance inputs.\"\"\"\n    lam = np.random.beta(alpha, alpha) if alpha > 0 else 1\n    index = torch.randperm(x.size()[0]).to(Config.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\n# ==========================================\n# 5. INFERENCE WITH TEST-TIME AUGMENTATION\n# ==========================================\ndef tta_inference(model, image_tensor, num_augmentations=4):\n    \"\"\"\n    Squeezes out extra accuracy by predicting on multiple variations \n    of the same image and averaging the probabilities.\n    \"\"\"\n    model.eval()\n    tta_transforms = [\n        transforms.RandomHorizontalFlip(p=1.0),\n        transforms.RandomRotation(15),\n        transforms.ColorJitter(brightness=0.1, contrast=0.1)\n    ]\n    \n    probabilities = []\n    with torch.no_grad():\n        # 1. Base prediction\n        base_logits = model(image_tensor.unsqueeze(0).to(Config.DEVICE))\n        probabilities.append(F.softmax(base_logits, dim=1))\n        \n        # 2. Augmented predictions\n        for t in tta_transforms:\n            # We must apply the transform to the patches or the base image\n            # For simplicity, applying transform directly to the (16,3,224,224) tensor\n            aug_tensor = t(image_tensor)\n            logits = model(aug_tensor.unsqueeze(0).to(Config.DEVICE))\n            probabilities.append(F.softmax(logits, dim=1))\n            \n    # Average the probabilities across all TTA passes\n    avg_probs = torch.mean(torch.stack(probabilities), dim=0)\n    final_pred = torch.argmax(avg_probs, dim=1).item()\n    return final_pred\n\n# ==========================================\n# MAIN EXECUTION MOCKUP\n# ==========================================\nif __name__ == '__main__':\n    # Define harsh weights to aggressively penalize missing a disease (Classes 1-4)\n    weights = torch.tensor([1.0, 3.5, 3.5, 3.5, 3.5]) \n    criterion = WeightedFocalLoss(alpha=weights)\n    \n    model = TMIL_Net().to(Config.DEVICE)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    \n    # Note: Wrap your dataloader and execute standard MixUp training loop here:\n    # inputs, targets_a, targets_b, lam = mixup_data(inputs, targets)\n    # outputs = model(inputs)\n    # loss = lam * criterion(outputs, targets_a) + (1 - lam) * criterion(outputs, targets_b)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T08:41:07.374766Z","iopub.execute_input":"2026-02-19T08:41:07.375256Z","iopub.status.idle":"2026-02-19T08:41:17.908245Z","shell.execute_reply.started":"2026-02-19T08:41:07.375227Z","shell.execute_reply":"2026-02-19T08:41:17.907583Z"}},"outputs":[],"execution_count":null},{"cell_type":"raw","source":"","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\nfrom tqdm import tqdm\nimport numpy as np\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader\n\n# --- Re-using the classes from the previous step (Config, TMIL_Net, etc.) ---\n# Ensure you have run the class definitions cell first!\n\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    all_preds = []\n    all_targets = []\n    \n    pbar = tqdm(loader, desc=\"Training\")\n    for images, labels in pbar:\n        images, labels = images.to(device), labels.to(device)\n        \n        # --- Apply MixUp Augmentation ---\n        images, targets_a, targets_b, lam = mixup_data(images, labels, alpha=0.2)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        \n        # MixUp Loss\n        loss = mixup_criterion(criterion, outputs, targets_a, targets_b, lam)\n        \n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        \n        # Store predictions for monitoring (argmax is approximation during MixUp)\n        _, preds = torch.max(outputs, 1)\n        all_preds.extend(preds.cpu().numpy())\n        all_targets.extend(labels.cpu().numpy())\n        \n        pbar.set_postfix({'loss': running_loss / (len(all_preds) / Config.BATCH_SIZE)})\n        \n    return running_loss / len(loader)\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    all_preds = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for images, labels in tqdm(loader, desc=\"Validating\"):\n            images, labels = images.to(device), labels.to(device)\n            \n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, preds = torch.max(outputs, 1)\n            \n            all_preds.extend(preds.cpu().numpy())\n            all_targets.extend(labels.cpu().numpy())\n            \n    # Calculate Metrics\n    acc = accuracy_score(all_targets, all_preds)\n    # Quadratic Weighted Kappa (The Gold Standard for DR)\n    kappa = cohen_kappa_score(all_targets, all_preds, weights='quadratic')\n    \n    return running_loss / len(loader), acc, kappa\n\n# ==========================================\n# MAIN EXECUTION LOOP\n# ==========================================\nif __name__ == '__main__':\n    print(f\"Using device: {Config.DEVICE}\")\n    \n    # 1. Load Data (Adjusting for your specific Kaggle paths)\n    # Using APTOS for the main training as it is the standard 5-class benchmark\n    # If the file isn't found, try: '/kaggle/input/aptos2019-blindness-detection/train.csv'\n    base_path = '/kaggle/input/competitions/aptos2019-blindness-detection/'\n    if not os.path.exists(base_path):\n        base_path = '/kaggle/input/aptos2019-blindness-detection/'\n        \n    df = pd.read_csv(os.path.join(base_path, 'train.csv'))\n    \n    # Update Config with dynamic paths based on where the code is running\n    Config.APTOS_IMG_DIR = os.path.join(base_path, 'train_images/')\n    \n    # Add full path to filenames\n    df['path'] = df['id_code'].apply(lambda x: os.path.join(Config.APTOS_IMG_DIR, f\"{x}.png\"))\n    \n    # Split Data (80% Train, 20% Val)\n    train_df, val_df = train_test_split(df, test_size=0.2, random_state=42, stratify=df['diagnosis'])\n    \n    print(f\"Training on {len(train_df)} images, Validating on {len(val_df)} images.\")\n\n    # 2. Create Datasets & Loaders\n    train_dataset = RetinalPatchDataset(\n        train_df['path'].values, \n        train_df['diagnosis'].values\n    )\n    val_dataset = RetinalPatchDataset(\n        val_df['path'].values, \n        val_df['diagnosis'].values\n    )\n\n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n\n    # 3. Initialize Model & Loss\n    # Class weights: Heavily penalize missing Proliferative DR (Class 4)\n    weights = torch.tensor([0.5, 2.0, 2.0, 3.0, 4.0]).to(Config.DEVICE)\n    criterion = WeightedFocalLoss(alpha=weights)\n    \n    model = TMIL_Net(num_classes=5).to(Config.DEVICE)\n    optimizer = optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-4)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n\n    # 4. Training Loop\n    best_kappa = -1.0\n    \n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        \n        train_loss = train_one_epoch(model, train_loader, criterion, optimizer, Config.DEVICE)\n        val_loss, val_acc, val_kappa = validate(model, val_loader, criterion, Config.DEVICE)\n        \n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n        print(f\"Val Accuracy: {val_acc:.4f} | Val Kappa: {val_kappa:.4f}\")\n        \n        # Save Best Model\n        if val_kappa > best_kappa:\n            best_kappa = val_kappa\n            torch.save(model.state_dict(), 'best_tmil_model.pth')\n            print(\">>> Best Model Saved!\")\n            \n    print(\"\\nTraining Complete. Best Kappa:\", best_kappa)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T08:41:26.07839Z","iopub.execute_input":"2026-02-19T08:41:26.07879Z","iopub.status.idle":"2026-02-19T12:56:41.83182Z","shell.execute_reply.started":"2026-02-19T08:41:26.078762Z","shell.execute_reply":"2026-02-19T12:56:41.830935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def mixup_criterion(criterion, pred, y_a, y_b, lam):\n    \"\"\"Calculates the loss for the blended MixUp images.\"\"\"\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T08:40:49.90787Z","iopub.execute_input":"2026-02-19T08:40:49.908207Z","iopub.status.idle":"2026-02-19T08:40:49.912484Z","shell.execute_reply.started":"2026-02-19T08:40:49.90818Z","shell.execute_reply":"2026-02-19T08:40:49.9119Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix\n\ndef evaluate_best_model(model_path, loader, device):\n    model.load_state_dict(torch.load(model_path))\n    model.eval()\n    \n    all_preds = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device)\n            outputs = model(images)\n            _, preds = torch.max(outputs, 1)\n            all_preds.extend(preds.cpu().numpy())\n            all_targets.extend(labels.numpy())\n            \n    cm = confusion_matrix(all_targets, all_preds)\n    plt.figure(figsize=(8,6))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues')\n    plt.xlabel('Predicted')\n    plt.ylabel('Actual')\n    plt.title('Confusion Matrix - Best TMIL Model')\n    plt.show()\n\n# Run it on your validation loader\nevaluate_best_model('best_tmil_model.pth', val_loader, Config.DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-19T13:39:52.914549Z","iopub.execute_input":"2026-02-19T13:39:52.914916Z","iopub.status.idle":"2026-02-19T13:40:50.458402Z","shell.execute_reply.started":"2026-02-19T13:39:52.914882Z","shell.execute_reply":"2026-02-19T13:40:50.457634Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# The exact data extracted from your 30-epoch training logs\nepochs = list(range(1, 31))\n\ntrain_loss = [1.8701, 1.6938, 1.5884, 1.5945, 1.5379, 1.5277, 1.4973, 1.4663, 1.4546, 1.4282, \n              1.3875, 1.3320, 1.2776, 1.1998, 1.1520, 1.0771, 0.9660, 0.9153, 0.8081, 0.7216, \n              0.6612, 0.5808, 0.4989, 0.4708, 0.4286, 0.3943, 0.4046, 0.3953, 0.3594, 0.4043]\n\nval_loss = [1.7345, 1.5882, 1.5963, 1.5012, 1.4726, 1.5554, 1.4778, 1.4937, 1.4032, 1.3942, \n            1.3696, 1.3277, 1.3725, 1.3238, 1.3220, 1.3775, 1.3881, 1.4482, 1.5701, 1.6850, \n            1.7374, 1.7543, 1.8793, 1.9761, 1.8684, 1.9095, 1.9854, 1.9363, 1.9453, 1.9412]\n\nval_kappa = [0.0000, 0.3004, 0.2284, 0.5703, 0.6706, 0.5551, 0.6530, 0.4553, 0.6401, 0.7063, \n             0.6664, 0.6157, 0.6804, 0.6652, 0.6353, 0.6351, 0.6831, 0.7030, 0.6914, 0.7067, \n             0.6899, 0.6842, 0.7155, 0.7092, 0.7222, 0.7142, 0.7151, 0.7192, 0.7273, 0.7234]\n\n# Set up the figure with 2 subplots side-by-side\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 5))\n\n# Plot 1: Training vs Validation Loss\nax1.plot(epochs, train_loss, 'b-', label='Training Loss', linewidth=2)\nax1.plot(epochs, val_loss, 'r-', label='Validation Loss', linewidth=2)\nax1.axvline(x=14, color='gray', linestyle='--', alpha=0.7, label='Overfitting Begins')\nax1.set_title('Training vs Validation Loss', fontsize=14)\nax1.set_xlabel('Epochs', fontsize=12)\nax1.set_ylabel('Loss', fontsize=12)\nax1.grid(True, linestyle='--', alpha=0.7)\nax1.legend(fontsize=11)\n\n# Plot 2: Validation Kappa Score\nax2.plot(epochs, val_kappa, 'g-', label='Validation Kappa', linewidth=2)\n# Highlight the epoch with the best Kappa\nbest_epoch = val_kappa.index(max(val_kappa)) + 1\nax2.scatter(best_epoch, max(val_kappa), color='red', s=100, zorder=5, \n            label=f'Best Kappa ({max(val_kappa):.4f})')\n\nax2.set_title('Validation Quadratic Weighted Kappa', fontsize=14)\nax2.set_xlabel('Epochs', fontsize=12)\nax2.set_ylabel('Kappa Score', fontsize=12)\nax2.grid(True, linestyle='--', alpha=0.7)\nax2.legend(fontsize=11)\n\n# Adjust layout and display\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T06:20:59.517834Z","iopub.execute_input":"2026-02-20T06:20:59.518121Z","iopub.status.idle":"2026-02-20T06:21:00.133724Z","shell.execute_reply.started":"2026-02-20T06:20:59.518088Z","shell.execute_reply":"2026-02-20T06:21:00.132833Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport torch\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import transforms\nfrom tqdm import tqdm\nimport timm\n\n# ==========================================\n# 1. CONFIGURATION & KAGGLE PATHS\n# ==========================================\nclass Config:\n    # Handle Kaggle paths dynamically\n    BASE_PATH = '/kaggle/input/competitions/aptos2019-blindness-detection/'\n    if not os.path.exists(BASE_PATH):\n        BASE_PATH = '/kaggle/input/aptos2019-blindness-detection/'\n        \n    APTOS_IMG_DIR = os.path.join(BASE_PATH, 'train_images/')\n    TRAIN_CSV = os.path.join(BASE_PATH, 'train.csv')\n    \n    IMG_SIZE = 896\n    PATCH_SIZE = 224\n    NUM_PATCHES = 16\n    NUM_CLASSES = 5\n    BATCH_SIZE = 4\n    EPOCHS = 30\n    LR = 1e-4\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# ==========================================\n# 2. PREPROCESSING & DATASET\n# ==========================================\ndef apply_clahe(image):\n    \"\"\"Enhances foreground lesions like microaneurysms.\"\"\"\n    lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n    l_channel, a, b = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    cl = clahe.apply(l_channel)\n    limg = cv2.merge((cl,a,b))\n    return cv2.cvtColor(limg, cv2.COLOR_LAB2RGB)\n\nclass RetinalPatchDataset(Dataset):\n    def __init__(self, image_paths, labels, transform=None):\n        self.image_paths = image_paths\n        self.labels = labels\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        img = cv2.imread(self.image_paths[idx])\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        img = apply_clahe(img)\n        \n        img_tensor = transforms.ToTensor()(img)\n        if self.transform:\n            img_tensor = self.transform(img_tensor)\n\n        # Slice into 16 non-overlapping 224x224 patches\n        patches = img_tensor.unfold(1, Config.PATCH_SIZE, Config.PATCH_SIZE)\\\n                            .unfold(2, Config.PATCH_SIZE, Config.PATCH_SIZE)\n        patches = patches.contiguous().view(3, Config.NUM_PATCHES, Config.PATCH_SIZE, Config.PATCH_SIZE)\n        patches = patches.permute(1, 0, 2, 3) \n        \n        label = torch.tensor(self.labels[idx], dtype=torch.long)\n        return patches, label\n\n# ==========================================\n# 3. TMIL ARCHITECTURE\n# ==========================================\nclass GICB(nn.Module):\n    \"\"\"Global Instance Computing Block.\"\"\"\n    def __init__(self, in_dim, out_dim, heads):\n        super().__init__()\n        self.norm = nn.LayerNorm(in_dim)\n        self.attn = nn.MultiheadAttention(embed_dim=in_dim, num_heads=heads, batch_first=True)\n        self.mlp = nn.Sequential(\n            nn.Linear(in_dim, out_dim),\n            nn.GELU()\n        )\n        self.res_proj = nn.Linear(in_dim, out_dim) if in_dim != out_dim else nn.Identity()\n\n    def forward(self, x):\n        nx = self.norm(x)\n        attn_out, _ = self.attn(nx, nx, nx)\n        a = attn_out + x \n        m = self.mlp(a) + self.res_proj(a)\n        return m\n\nclass TMIL_Net(nn.Module):\n    def __init__(self, num_classes=Config.NUM_CLASSES):\n        super().__init__()\n        self.vit = timm.create_model('vit_small_patch16_224', pretrained=True, num_classes=0)\n        self.pos_embed = nn.Parameter(torch.zeros(1, Config.NUM_PATCHES, 384))\n        \n        self.gicb1 = GICB(in_dim=384, out_dim=128, heads=6)\n        self.gicb2 = GICB(in_dim=128, out_dim=128, heads=4)\n        \n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(Config.NUM_PATCHES * 128, 512),\n            nn.ReLU(),\n            ### NEW ADJUSTMENT 1: Dropout increased to 0.5 to stop memorization ###\n            nn.Dropout(0.5), \n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x):\n        B, N, C, H, W = x.shape\n        x = x.view(B * N, C, H, W) \n        \n        features = self.vit(x) \n        features = features.view(B, N, 384) \n        features = features + self.pos_embed\n        \n        out = self.gicb1(features)\n        out = self.gicb2(out)\n        logits = self.classifier(out)\n        return logits\n\n# ==========================================\n# 4. LOSS & MIXUP\n# ==========================================\nclass WeightedFocalLoss(nn.Module):\n    def __init__(self, alpha, gamma=2.0):\n        super().__init__()\n        self.alpha = alpha.to(Config.DEVICE)\n        self.gamma = gamma\n\n    def forward(self, inputs, targets):\n        ce_loss = F.cross_entropy(inputs, targets, reduction='none', weight=self.alpha)\n        pt = torch.exp(-ce_loss)\n        focal_loss = ((1 - pt) ** self.gamma) * ce_loss\n        return focal_loss.mean()\n\ndef mixup_data(x, y, alpha=0.2):\n    lam = np.random.beta(alpha, alpha) if alpha > 0 else 1\n    index = torch.randperm(x.size()[0]).to(Config.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\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\n# ==========================================\n# 5. TRAINING LOOPS\n# ==========================================\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    all_preds, all_targets = [], []\n    \n    pbar = tqdm(loader, desc=\"Training\")\n    for images, labels in pbar:\n        images, labels = images.to(device), labels.to(device)\n        \n        ### NEW ADJUSTMENT 2: MixUp Alpha locked to 0.2 to preserve tiny lesions ###\n        images, targets_a, targets_b, lam = mixup_data(images, labels, alpha=0.2)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = mixup_criterion(criterion, outputs, targets_a, targets_b, lam)\n        \n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, preds = torch.max(outputs, 1)\n        all_preds.extend(preds.cpu().numpy())\n        \n        pbar.set_postfix({'loss': running_loss / (len(all_preds) / Config.BATCH_SIZE)})\n        \n    return running_loss / len(loader)\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    all_preds, all_targets = [], []\n    \n    with torch.no_grad():\n        for images, labels in tqdm(loader, desc=\"Validating\"):\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, preds = torch.max(outputs, 1)\n            \n            all_preds.extend(preds.cpu().numpy())\n            all_targets.extend(labels.cpu().numpy())\n            \n    acc = accuracy_score(all_targets, all_preds)\n    kappa = cohen_kappa_score(all_targets, all_preds, weights='quadratic')\n    return running_loss / len(loader), acc, kappa\n\n# ==========================================\n# 6. MAIN EXECUTION\n# ==========================================\nif __name__ == '__main__':\n    print(f\"Using device: {Config.DEVICE}\")\n    \n    # Load Data\n    df = pd.read_csv(Config.TRAIN_CSV)\n    df['path'] = df['id_code'].apply(lambda x: os.path.join(Config.APTOS_IMG_DIR, f\"{x}.png\"))\n    train_df, val_df = train_test_split(df, test_size=0.2, random_state=42, stratify=df['diagnosis'])\n    print(f\"Training on {len(train_df)} images, Validating on {len(val_df)} images.\")\n\n    # Datasets & Loaders\n    train_dataset = RetinalPatchDataset(train_df['path'].values, train_df['diagnosis'].values)\n    val_dataset = RetinalPatchDataset(val_df['path'].values, val_df['diagnosis'].values)\n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n\n    ### NEW ADJUSTMENT 3: Penalize minority classes (1,3,4) heavily, reduce penalty for majority class (2) ###\n    weights = torch.tensor([0.5, 3.5, 1.0, 4.0, 4.0]).to(Config.DEVICE)\n    criterion = WeightedFocalLoss(alpha=weights)\n    \n    model = TMIL_Net(num_classes=5).to(Config.DEVICE)\n    \n    ### NEW ADJUSTMENT 4: Weight decay increased from 1e-4 to 1e-2 ###\n    optimizer = optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=1e-2)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n\n    best_kappa = -1.0\n    \n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        \n        train_loss = train_one_epoch(model, train_loader, criterion, optimizer, Config.DEVICE)\n        val_loss, val_acc, val_kappa = validate(model, val_loader, criterion, Config.DEVICE)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n        print(f\"Val Accuracy: {val_acc:.4f} | Val Kappa: {val_kappa:.4f}\")\n        \n        if val_kappa > best_kappa:\n            best_kappa = val_kappa\n            torch.save(model.state_dict(), 'best_tmil_model_v2.pth')\n            print(\">>> Best Model Saved (v2)!\")\n            \n    print(\"\\nTraining Complete. Best Kappa:\", best_kappa)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T06:37:36.331463Z","iopub.execute_input":"2026-02-20T06:37:36.332107Z","iopub.status.idle":"2026-02-20T11:30:15.3864Z","shell.execute_reply.started":"2026-02-20T06:37:36.332074Z","shell.execute_reply":"2026-02-20T11:30:15.385308Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport torch\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import transforms\nfrom tqdm import tqdm\nimport timm\n\n# ==========================================\n# 1. CONFIGURATION & KAGGLE PATHS\n# ==========================================\nclass Config:\n    BASE_PATH = '/kaggle/input/competitions/aptos2019-blindness-detection/'\n    if not os.path.exists(BASE_PATH):\n        BASE_PATH = '/kaggle/input/aptos2019-blindness-detection/'\n        \n    APTOS_IMG_DIR = os.path.join(BASE_PATH, 'train_images/')\n    TRAIN_CSV = os.path.join(BASE_PATH, 'train.csv')\n    \n    IMG_SIZE = 896\n    PATCH_SIZE = 224\n    NUM_PATCHES = 16\n    NUM_CLASSES = 5\n    BATCH_SIZE = 4\n    EPOCHS = 40  # Increased to allow the new LR scheduler to work\n    LR = 1e-4\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# ==========================================\n# 2. PREPROCESSING & DATASET\n# ==========================================\ndef apply_clahe(image):\n    lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n    l_channel, a, b = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    cl = clahe.apply(l_channel)\n    limg = cv2.merge((cl,a,b))\n    return cv2.cvtColor(limg, cv2.COLOR_LAB2RGB)\n\nclass RetinalPatchDataset(Dataset):\n    def __init__(self, image_paths, labels, is_train=True):\n        self.image_paths = image_paths\n        self.labels = labels\n        self.is_train = is_train\n        \n        # EXACT TMIL PAPER AUGMENTATIONS (-180 to 180 degree rotation)\n        self.train_transforms = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomRotation(180), \n            transforms.ToTensor(),\n        ])\n        self.val_transforms = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.ToTensor(),\n        ])\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        img = cv2.imread(self.image_paths[idx])\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        img = apply_clahe(img)\n        \n        if self.is_train:\n            img_tensor = self.train_transforms(img)\n        else:\n            img_tensor = self.val_transforms(img)\n\n        # Slice into 16 non-overlapping 224x224 patches\n        patches = img_tensor.unfold(1, Config.PATCH_SIZE, Config.PATCH_SIZE)\\\n                            .unfold(2, Config.PATCH_SIZE, Config.PATCH_SIZE)\n        patches = patches.contiguous().view(3, Config.NUM_PATCHES, Config.PATCH_SIZE, Config.PATCH_SIZE)\n        patches = patches.permute(1, 0, 2, 3) \n        \n        label = torch.tensor(self.labels[idx], dtype=torch.long)\n        return patches, label\n\n# ==========================================\n# 3. TMIL ARCHITECTURE (No changes)\n# ==========================================\nclass GICB(nn.Module):\n    def __init__(self, in_dim, out_dim, heads):\n        super().__init__()\n        self.norm = nn.LayerNorm(in_dim)\n        self.attn = nn.MultiheadAttention(embed_dim=in_dim, num_heads=heads, batch_first=True)\n        self.mlp = nn.Sequential(\n            nn.Linear(in_dim, out_dim),\n            nn.GELU()\n        )\n        self.res_proj = nn.Linear(in_dim, out_dim) if in_dim != out_dim else nn.Identity()\n\n    def forward(self, x):\n        nx = self.norm(x)\n        attn_out, _ = self.attn(nx, nx, nx)\n        a = attn_out + x \n        m = self.mlp(a) + self.res_proj(a)\n        return m\n\nclass TMIL_Net(nn.Module):\n    def __init__(self, num_classes=Config.NUM_CLASSES):\n        super().__init__()\n        self.vit = timm.create_model('vit_small_patch16_224', pretrained=True, num_classes=0)\n        self.pos_embed = nn.Parameter(torch.zeros(1, Config.NUM_PATCHES, 384))\n        \n        self.gicb1 = GICB(in_dim=384, out_dim=128, heads=6)\n        self.gicb2 = GICB(in_dim=128, out_dim=128, heads=4)\n        \n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(Config.NUM_PATCHES * 128, 512),\n            nn.ReLU(),\n            nn.Dropout(0.4), # Reverted to paper's baseline\n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x):\n        B, N, C, H, W = x.shape\n        x = x.view(B * N, C, H, W) \n        features = self.vit(x) \n        features = features.view(B, N, 384) \n        features = features + self.pos_embed\n        out = self.gicb1(features)\n        out = self.gicb2(out)\n        logits = self.classifier(out)\n        return logits\n\n# ==========================================\n# 4. TRAINING LOOPS\n# ==========================================\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    all_preds, all_targets = [], []\n    \n    pbar = tqdm(loader, desc=\"Training\")\n    for images, labels in pbar:\n        images, labels = images.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        \n        # Standard loss computation without MixUp to hit raw accuracy targets\n        loss = criterion(outputs, labels)\n        \n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, preds = torch.max(outputs, 1)\n        all_preds.extend(preds.cpu().numpy())\n        \n        pbar.set_postfix({'loss': running_loss / (len(all_preds) / Config.BATCH_SIZE)})\n        \n    return running_loss / len(loader)\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    all_preds, all_targets = [], []\n    \n    with torch.no_grad():\n        for images, labels in tqdm(loader, desc=\"Validating\"):\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, preds = torch.max(outputs, 1)\n            \n            all_preds.extend(preds.cpu().numpy())\n            all_targets.extend(labels.cpu().numpy())\n            \n    acc = accuracy_score(all_targets, all_preds)\n    kappa = cohen_kappa_score(all_targets, all_preds, weights='quadratic')\n    return running_loss / len(loader), acc, kappa\n\n# ==========================================\n# 5. MAIN EXECUTION\n# ==========================================\nif __name__ == '__main__':\n    print(f\"Using device: {Config.DEVICE}\")\n    \n    df = pd.read_csv(Config.TRAIN_CSV)\n    df['path'] = df['id_code'].apply(lambda x: os.path.join(Config.APTOS_IMG_DIR, f\"{x}.png\"))\n    train_df, val_df = train_test_split(df, test_size=0.2, random_state=42, stratify=df['diagnosis'])\n\n    train_dataset = RetinalPatchDataset(train_df['path'].values, train_df['diagnosis'].values, is_train=True)\n    val_dataset = RetinalPatchDataset(val_df['path'].values, val_df['diagnosis'].values, is_train=False)\n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n\n    # ---------------------------------------------------------\n    # EXACT PAPER HYPERPARAMETERS IMPLEMENTED HERE\n    # ---------------------------------------------------------\n    # 1. Label Smoothing Cross Entropy (0.05) instead of Focal Loss\n    criterion = nn.CrossEntropyLoss(label_smoothing=0.05)\n    \n    model = TMIL_Net(num_classes=5).to(Config.DEVICE)\n    \n    # 2. Massive weight decay (0.3)\n    optimizer = optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=0.3)\n    \n    # 3. Step LR Scheduler (Reduces LR by half at epochs 20 and 35)\n    scheduler = optim.lr_scheduler.MultiStepLR(optimizer, milestones=[20, 35], gamma=0.5)\n\n    best_acc = 0.0\n    best_kappa = -1.0\n    \n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        \n        train_loss = train_one_epoch(model, train_loader, criterion, optimizer, Config.DEVICE)\n        val_loss, val_acc, val_kappa = validate(model, val_loader, criterion, Config.DEVICE)\n        \n        # Step the learning rate\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n        print(f\"Val Accuracy: {val_acc:.4f} | Val Kappa: {val_kappa:.4f}\")\n        print(f\"Current Learning Rate: {scheduler.get_last_lr()[0]}\")\n        \n        # Save model based on Accuracy now instead of Kappa\n        if val_acc > best_acc:\n            best_acc = val_acc\n            best_kappa = val_kappa\n            torch.save(model.state_dict(), 'best_tmil_model_v3.pth')\n            print(\">>> Best Model Saved (v3)!\")\n            \n    print(f\"\\nTraining Complete. Best Accuracy: {best_acc:.4f} | Associated Kappa: {best_kappa:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-15T06:14:46.67431Z","iopub.execute_input":"2026-03-15T06:14:46.674545Z","iopub.status.idle":"2026-03-15T12:05:53.592718Z","shell.execute_reply.started":"2026-03-15T06:14:46.674523Z","shell.execute_reply":"2026-03-15T12:05:53.591755Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport torch\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import transforms\nfrom tqdm import tqdm\nimport timm\n\n# ==========================================\n# 1. CONFIGURATION\n# ==========================================\nclass Config:\n    BASE_PATH = '/kaggle/input/competitions/aptos2019-blindness-detection/'\n    if not os.path.exists(BASE_PATH):\n        BASE_PATH = '/kaggle/input/aptos2019-blindness-detection/'\n        \n    APTOS_IMG_DIR = os.path.join(BASE_PATH, 'train_images/')\n    TRAIN_CSV = os.path.join(BASE_PATH, 'train.csv')\n    \n    # Swin Transformer optimal high-resolution input\n    IMG_SIZE = 384 \n    NUM_CLASSES = 5\n    BATCH_SIZE = 8 # Swin is memory intensive; keep batch size manageable\n    EPOCHS = 45\n    LR = 5e-5 # Lower learning rate for fine-tuning a massive Swin Base model\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# ==========================================\n# 2. PREPROCESSING & DATASET\n# ==========================================\ndef apply_clahe(image):\n    \"\"\"Enhances foreground lesions like microaneurysms.\"\"\"\n    lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n    l_channel, a, b = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    cl = clahe.apply(l_channel)\n    limg = cv2.merge((cl,a,b))\n    return cv2.cvtColor(limg, cv2.COLOR_LAB2RGB)\n\nclass SwinRetinalDataset(Dataset):\n    def __init__(self, image_paths, labels, is_train=True):\n        self.image_paths = image_paths\n        self.labels = labels\n        self.is_train = is_train\n        \n        # Heavy augmentations to prevent Swin memorization\n        self.train_transforms = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomVerticalFlip(p=0.5),\n            transforms.RandomRotation(180), \n            transforms.ColorJitter(brightness=0.1, contrast=0.2),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n        ])\n        self.val_transforms = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n        ])\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        img = cv2.imread(self.image_paths[idx])\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n        img = apply_clahe(img)\n        \n        if self.is_train:\n            img_tensor = self.train_transforms(img)\n        else:\n            img_tensor = self.val_transforms(img)\n        \n        label = torch.tensor(self.labels[idx], dtype=torch.long)\n        return img_tensor, label\n\n# ==========================================\n# 3. SWIN TRANSFORMER ARCHITECTURE\n# ==========================================\nclass Swin_DR_Net(nn.Module):\n    def __init__(self, num_classes=Config.NUM_CLASSES):\n        super().__init__()\n        # Load Swin-Base optimized for 384x384 resolution\n        self.swin = timm.create_model('swin_base_patch4_window12_384', pretrained=True, num_classes=0)\n        \n        # Swin outputs a 1D pooled feature vector (dim=1024 for base model)\n        self.classifier = nn.Sequential(\n            nn.Linear(1024, 512),\n            nn.GELU(),\n            nn.Dropout(0.5), # Aggressive dropout\n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x):\n        features = self.swin(x)\n        logits = self.classifier(features)\n        return logits\n\n# ==========================================\n# 4. LOSS & MIXUP\n# ==========================================\nclass WeightedFocalLoss(nn.Module):\n    \"\"\"Crushes False Negatives by heavily penalizing disease misclassification.\"\"\"\n    def __init__(self, alpha, gamma=2.0):\n        super().__init__()\n        self.alpha = alpha.to(Config.DEVICE)\n        self.gamma = gamma\n\n    def forward(self, inputs, targets):\n        ce_loss = F.cross_entropy(inputs, targets, reduction='none', weight=self.alpha)\n        pt = torch.exp(-ce_loss)\n        focal_loss = ((1 - pt) ** self.gamma) * ce_loss\n        return focal_loss.mean()\n\ndef mixup_data(x, y, alpha=0.2):\n    lam = np.random.beta(alpha, alpha) if alpha > 0 else 1\n    index = torch.randperm(x.size()[0]).to(Config.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\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\n# ==========================================\n# 5. TRAINING LOOPS\n# ==========================================\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    all_preds, all_targets = [], []\n    \n    pbar = tqdm(loader, desc=\"Training\")\n    for images, labels in pbar:\n        images, labels = images.to(device), labels.to(device)\n        \n        # Apply MixUp augmentation\n        images, targets_a, targets_b, lam = mixup_data(images, labels, alpha=0.2)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = mixup_criterion(criterion, outputs, targets_a, targets_b, lam)\n        \n        loss.backward()\n        \n        # Gradient clipping to stabilize Swin training\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optimizer.step()\n        \n        running_loss += loss.item()\n        _, preds = torch.max(outputs, 1)\n        all_preds.extend(preds.cpu().numpy())\n        \n        pbar.set_postfix({'loss': running_loss / (len(all_preds) / Config.BATCH_SIZE)})\n        \n    return running_loss / len(loader)\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    all_preds, all_targets = [], []\n    \n    # Switch out of focal loss for clean validation loss tracking\n    val_criterion = nn.CrossEntropyLoss() \n    \n    with torch.no_grad():\n        for images, labels in tqdm(loader, desc=\"Validating\"):\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = val_criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            _, preds = torch.max(outputs, 1)\n            \n            all_preds.extend(preds.cpu().numpy())\n            all_targets.extend(labels.cpu().numpy())\n            \n    acc = accuracy_score(all_targets, all_preds)\n    kappa = cohen_kappa_score(all_targets, all_preds, weights='quadratic')\n    return running_loss / len(loader), acc, kappa\n\n# ==========================================\n# 6. MAIN EXECUTION\n# ==========================================\nif __name__ == '__main__':\n    print(f\"Using device: {Config.DEVICE}\")\n    \n    df = pd.read_csv(Config.TRAIN_CSV)\n    df['path'] = df['id_code'].apply(lambda x: os.path.join(Config.APTOS_IMG_DIR, f\"{x}.png\"))\n    train_df, val_df = train_test_split(df, test_size=0.2, random_state=42, stratify=df['diagnosis'])\n\n    train_dataset = SwinRetinalDataset(train_df['path'].values, train_df['diagnosis'].values, is_train=True)\n    val_dataset = SwinRetinalDataset(val_df['path'].values, val_df['diagnosis'].values, is_train=False)\n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n\n    # Reintroduce the class weights to force the model to identify diseases\n    weights = torch.tensor([0.5, 3.5, 1.0, 4.0, 4.0]).to(Config.DEVICE)\n    criterion = WeightedFocalLoss(alpha=weights)\n    \n    model = Swin_DR_Net(num_classes=5).to(Config.DEVICE)\n    \n    # High weight decay (0.05) to combat Swin's massive parameter count\n    optimizer = optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=0.05)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n\n    best_acc = 0.0\n    best_kappa = -1.0\n    \n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        \n        train_loss = train_one_epoch(model, train_loader, criterion, optimizer, Config.DEVICE)\n        val_loss, val_acc, val_kappa = validate(model, val_loader, criterion, Config.DEVICE)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f}\")\n        print(f\"Val Accuracy: {val_acc:.4f} | Val Kappa: {val_kappa:.4f}\")\n        \n        # Save based on Kappa as it is the critical metric for imbalanced DR tasks\n        if val_kappa > best_kappa:\n            best_kappa = val_kappa\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_swin_model.pth')\n            print(\">>> Best Swin Model Saved!\")\n            \n    print(f\"\\nTraining Complete. Best Kappa: {best_kappa:.4f} | Associated Accuracy: {best_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-25T19:49:22.556588Z","iopub.execute_input":"2026-03-25T19:49:22.556923Z","iopub.status.idle":"2026-03-26T00:49:38.079962Z","shell.execute_reply.started":"2026-03-25T19:49:22.556878Z","shell.execute_reply":"2026-03-26T00:49:38.079037Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport time\nimport glob\nimport torch\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score, roc_curve, auc\nfrom sklearn.model_selection import train_test_split\nfrom torchvision import transforms\nfrom tqdm import tqdm\nimport timm\nimport matplotlib.pyplot as plt\n\n# Start the total runtime timer\nTOTAL_START_TIME = time.time()\n\n# ==========================================\n# 1. CONFIGURATION (AUTO-DETECT PATHS)\n# ==========================================\nclass Config:\n    # 1. Automatically find the CSV file\n    csv_search = glob.glob('/kaggle/input/**/messidor_data.csv', recursive=True)\n    if not csv_search:\n        raise FileNotFoundError(\"Could not find messidor_data.csv. Please ensure the Messidor dataset is added.\")\n    CSV_PATH = csv_search[0]\n    \n    # 2. Automatically find the image directory\n    img_search = glob.glob('/kaggle/input/**/preprocess/', recursive=True)\n    if not img_search:\n        img_search = glob.glob('/kaggle/input/**/messidor-2/preprocess/', recursive=True)\n        \n    IMG_DIR = img_search[0] if img_search else '/kaggle/input/datasets/mariaherrerot/messidor2preprocess/messidor-2/messidor-2/preprocess/'\n    \n    IMG_SIZE = 384 \n    NUM_CLASSES = 2 # Binary: Normal vs. Referable DR\n    BATCH_SIZE = 8 \n    EPOCHS = 35\n    LR = 5e-5 \n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# ==========================================\n# 2. PREPROCESSING & DATASET\n# ==========================================\ndef apply_clahe(image):\n    lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n    l_channel, a, b = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8,8))\n    cl = clahe.apply(l_channel)\n    limg = cv2.merge((cl,a,b))\n    return cv2.cvtColor(limg, cv2.COLOR_LAB2RGB)\n\nclass MessidorSwinDataset(Dataset):\n    def __init__(self, image_paths, labels, is_train=True):\n        self.image_paths = image_paths\n        self.labels = labels\n        self.is_train = is_train\n        \n        self.train_transforms = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomVerticalFlip(p=0.5),\n            transforms.RandomRotation(180), \n            transforms.ColorJitter(brightness=0.1, contrast=0.2),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n        ])\n        self.val_transforms = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n        ])\n\n    def __len__(self):\n        return len(self.image_paths)\n\n    def __getitem__(self, idx):\n        img_path = self.image_paths[idx]\n        img = cv2.imread(img_path)\n        \n        if img is None:\n            img = np.zeros((Config.IMG_SIZE, Config.IMG_SIZE, 3), dtype=np.uint8)\n        else:\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n            img = cv2.resize(img, (Config.IMG_SIZE, Config.IMG_SIZE))\n            img = apply_clahe(img)\n        \n        if self.is_train:\n            img_tensor = self.train_transforms(img)\n        else:\n            img_tensor = self.val_transforms(img)\n        \n        label = torch.tensor(self.labels[idx], dtype=torch.long)\n        return img_tensor, label\n\n# ==========================================\n# 3. SWIN ARCHITECTURE\n# ==========================================\nclass Swin_DR_Net(nn.Module):\n    def __init__(self, num_classes=Config.NUM_CLASSES):\n        super().__init__()\n        self.swin = timm.create_model('swin_base_patch4_window12_384', pretrained=True, num_classes=0)\n        self.classifier = nn.Sequential(\n            nn.Linear(1024, 512),\n            nn.GELU(),\n            nn.Dropout(0.5), \n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x):\n        features = self.swin(x)\n        logits = self.classifier(features)\n        return logits\n\n# ==========================================\n# 4. BINARY FOCAL LOSS & MIXUP\n# ==========================================\nclass BinaryWeightedFocalLoss(nn.Module):\n    def __init__(self, alpha, gamma=2.0):\n        super().__init__()\n        self.alpha = alpha.to(Config.DEVICE)\n        self.gamma = gamma\n\n    def forward(self, inputs, targets):\n        ce_loss = F.cross_entropy(inputs, targets, reduction='none', weight=self.alpha)\n        pt = torch.exp(-ce_loss)\n        focal_loss = ((1 - pt) ** self.gamma) * ce_loss\n        return focal_loss.mean()\n\ndef mixup_data(x, y, alpha=0.2):\n    lam = np.random.beta(alpha, alpha) if alpha > 0 else 1\n    index = torch.randperm(x.size()[0]).to(Config.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\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\n# ==========================================\n# 5. TRAINING & VALIDATION LOOPS\n# ==========================================\ndef train_one_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    all_preds, all_targets = [], []\n    \n    pbar = tqdm(loader, desc=\"Training\")\n    for images, labels in pbar:\n        images, labels = images.to(device), labels.to(device)\n        images, targets_a, targets_b, lam = mixup_data(images, labels, alpha=0.2)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = mixup_criterion(criterion, outputs, targets_a, targets_b, lam)\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optimizer.step()\n        \n        running_loss += loss.item()\n        \n        _, preds = torch.max(outputs, 1)\n        all_preds.extend(preds.cpu().numpy())\n        dominant_targets = targets_a if lam > 0.5 else targets_b\n        all_targets.extend(dominant_targets.cpu().numpy())\n        \n        pbar.set_postfix({'loss': running_loss / (len(all_preds) / Config.BATCH_SIZE)})\n        \n    epoch_loss = running_loss / len(loader)\n    epoch_acc = accuracy_score(all_targets, all_preds)\n    return epoch_loss, epoch_acc\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    all_preds, all_targets, all_probs = [], [], []\n    val_criterion = nn.CrossEntropyLoss() \n    \n    with torch.no_grad():\n        for images, labels in tqdm(loader, desc=\"Validating\"):\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = val_criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            probs = F.softmax(outputs, dim=1)\n            _, preds = torch.max(probs, 1)\n            \n            all_preds.extend(preds.cpu().numpy())\n            all_targets.extend(labels.cpu().numpy())\n            all_probs.extend(probs[:, 1].cpu().numpy()) \n            \n    acc = accuracy_score(all_targets, all_preds)\n    kappa = cohen_kappa_score(all_targets, all_preds, weights='quadratic')\n    return running_loss / len(loader), acc, kappa, all_targets, all_probs\n\n# ==========================================\n# 6. RESEARCH PLOTTING FUNCTION\n# ==========================================\ndef plot_research_graphs(history, fpr, tpr, roc_auc):\n    epochs = range(1, len(history['train_loss']) + 1)\n    fig, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=(20, 5))\n    \n    # 1. Loss Curve\n    ax1.plot(epochs, history['train_loss'], 'b-', label='Train Loss', linewidth=2)\n    ax1.plot(epochs, history['val_loss'], 'r-', label='Val Loss', linewidth=2)\n    ax1.set_title('Training vs Validation Loss', fontsize=14)\n    ax1.set_xlabel('Epochs', fontsize=12)\n    ax1.set_ylabel('Loss', fontsize=12)\n    ax1.legend()\n    ax1.grid(True, linestyle='--', alpha=0.7)\n    \n    # 2. Accuracy Curve\n    ax2.plot(epochs, history['train_acc'], 'b-', label='Train Accuracy', linewidth=2)\n    ax2.plot(epochs, history['val_acc'], 'g-', label='Val Accuracy', linewidth=2)\n    ax2.set_title('Training vs Validation Accuracy', fontsize=14)\n    ax2.set_xlabel('Epochs', fontsize=12)\n    ax2.set_ylabel('Accuracy', fontsize=12)\n    ax2.legend()\n    ax2.grid(True, linestyle='--', alpha=0.7)\n    \n    # 3. ROC AUC Curve\n    ax3.plot(fpr, tpr, color='darkorange', lw=2, label=f'ROC curve (area = {roc_auc:.3f})')\n    ax3.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--')\n    ax3.set_xlim([0.0, 1.0])\n    ax3.set_ylim([0.0, 1.05])\n    ax3.set_xlabel('False Positive Rate', fontsize=12)\n    ax3.set_ylabel('True Positive Rate', fontsize=12)\n    ax3.set_title('Receiver Operating Characteristic (ROC)', fontsize=14)\n    ax3.legend(loc=\"lower right\")\n    ax3.grid(True, linestyle='--', alpha=0.7)\n    \n    plt.tight_layout()\n    plt.show()\n\n# ==========================================\n# 7. MAIN EXECUTION (WITH COLUMN AUTO-DETECT)\n# ==========================================\nif __name__ == '__main__':\n    print(f\"Using device: {Config.DEVICE}\")\n    print(f\"Loading CSV from: {Config.CSV_PATH}\")\n    print(f\"Loading Images from: {Config.IMG_DIR}\")\n    \n    df = pd.read_csv(Config.CSV_PATH)\n    \n    # Auto-detect columns to prevent KeyErrors\n    grade_cols = ['adjudicated_dr_grade', 'Retinopathy grade', 'dr_grade', 'diagnosis', 'grade']\n    id_cols = ['image_id', 'Image ID', 'id_code', 'image_name', 'image']\n    \n    grade_col = next((col for col in grade_cols if col in df.columns), None)\n    id_col = next((col for col in id_cols if col in df.columns), None)\n    \n    if not grade_col or not id_col:\n        raise ValueError(f\"Could not find required columns. CSV has: {df.columns.tolist()}\")\n        \n    print(f\"Using Column '{id_col}' for Image IDs and '{grade_col}' for Labels.\")\n\n    def map_to_binary(grade):\n        return 0 if grade in [0, 1] else 1 \n        \n    df['binary_label'] = df[grade_col].apply(map_to_binary)\n    df['path'] = df[id_col].apply(lambda x: os.path.join(Config.IMG_DIR, f\"{x}.png\" if not str(x).endswith('.png') else str(x)))\n    \n    df = df[df['path'].apply(os.path.exists)]\n    print(f\"Found {len(df)} valid images for training.\")\n\n    train_df, val_df = train_test_split(df, test_size=0.2, random_state=42, stratify=df['binary_label'])\n\n    train_dataset = MessidorSwinDataset(train_df['path'].values, train_df['binary_label'].values, is_train=True)\n    val_dataset = MessidorSwinDataset(val_df['path'].values, val_df['binary_label'].values, is_train=False)\n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n\n    weights = torch.tensor([1.0, 3.5]).to(Config.DEVICE)\n    criterion = BinaryWeightedFocalLoss(alpha=weights)\n    \n    model = Swin_DR_Net(num_classes=Config.NUM_CLASSES).to(Config.DEVICE)\n    optimizer = optim.AdamW(model.parameters(), lr=Config.LR, weight_decay=0.05)\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS)\n\n    history = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': []}\n    best_acc = 0.0\n    best_roc_data = None \n    \n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        \n        train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, Config.DEVICE)\n        val_loss, val_acc, val_kappa, val_targets, val_probs = validate(model, val_loader, criterion, Config.DEVICE)\n        scheduler.step()\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        \n        print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}\")\n        print(f\"Val Loss:   {val_loss:.4f} | Val Acc:   {val_acc:.4f} | Val Kappa: {val_kappa:.4f}\")\n        \n        if val_acc > best_acc:\n            best_acc = val_acc\n            torch.save(model.state_dict(), 'best_messidor_swin.pth')\n            print(\">>> Best Model Saved!\")\n            fpr, tpr, _ = roc_curve(val_targets, val_probs)\n            roc_auc = auc(fpr, tpr)\n            best_roc_data = (fpr, tpr, roc_auc)\n\n    # Calculate Total Runtime\n    total_time = time.time() - TOTAL_START_TIME\n    hours, rem = divmod(total_time, 3600)\n    minutes, seconds = divmod(rem, 60)\n    \n    print(\"\\n\" + \"=\"*50)\n    print(f\"TRAINING COMPLETE\")\n    print(f\"Best Validation Accuracy: {best_acc:.4f}\")\n    if best_roc_data:\n        print(f\"Best ROC AUC Score:       {best_roc_data[2]:.4f}\")\n    print(f\"Total Runtime:            {int(hours)}h {int(minutes)}m {int(seconds)}s\")\n    print(\"=\"*50)\n    \n    if best_roc_data:\n        plot_research_graphs(history, best_roc_data[0], best_roc_data[1], best_roc_data[2])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-26T12:18:50.752045Z","iopub.execute_input":"2026-03-26T12:18:50.752878Z","iopub.status.idle":"2026-03-26T13:19:54.23318Z","shell.execute_reply.started":"2026-03-26T12:18:50.752845Z","shell.execute_reply":"2026-03-26T13:19:54.232456Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nimport seaborn as sns\nfrom scipy.stats import chi2_contingency\nfrom collections import Counter\n\n# ==========================================\n# STEP 1: LOAD THE DATASET\n# ==========================================\ncsv_search = glob.glob('/kaggle/input/**/messidor_data.csv', recursive=True)\nCSV_PATH   = csv_search[0]\ndf         = pd.read_csv(CSV_PATH)\n\n# Auto-detect columns (same logic as your training code)\ngrade_cols = ['adjudicated_dr_grade', 'Retinopathy grade', \n              'dr_grade', 'diagnosis', 'grade']\ngrade_col  = next((col for col in grade_cols if col in df.columns), None)\n\nprint(\"=\"*55)\nprint(\"        MESSIDOR DATASET IMBALANCE ANALYSIS\")\nprint(\"=\"*55)\nprint(f\"\\nDataset Shape : {df.shape}\")\nprint(f\"Grade Column  : '{grade_col}'\")\nprint(f\"Total Samples : {len(df)}\\n\")\n\n# ==========================================\n# STEP 2: RAW GRADE DISTRIBUTION\n# ==========================================\nprint(\"─\"*55)\nprint(\" ORIGINAL DR GRADES (0-4 Scale)\")\nprint(\"─\"*55)\n\ngrade_counts = df[grade_col].value_counts().sort_index()\ngrade_labels = {\n    0: 'No DR',\n    1: 'Mild DR',\n    2: 'Moderate DR',\n    3: 'Severe DR',\n    4: 'Proliferative DR'\n}\n\nfor grade, count in grade_counts.items():\n    label      = grade_labels.get(grade, f'Grade {grade}')\n    pct        = (count / len(df)) * 100\n    bar        = '█' * int(pct / 2)\n    print(f\"  Grade {grade} | {label:<20} | {count:>4} samples \"\n          f\"| {pct:5.1f}% | {bar}\")\n\n# ==========================================\n# STEP 3: BINARY LABEL DISTRIBUTION\n# ==========================================\nprint(\"\\n\" + \"─\"*55)\nprint(\" BINARY LABELS (Normal vs Referable DR)\")\nprint(\"─\"*55)\n\ndf['binary_label'] = df[grade_col].apply(lambda x: 0 if x in [0,1] else 1)\nbinary_counts      = df['binary_label'].value_counts().sort_index()\nbinary_labels      = {0: 'Normal (Grade 0-1)', 1: 'Referable DR (Grade 2-4)'}\n\nfor label, count in binary_counts.items():\n    pct = (count / len(df)) * 100\n    bar = '█' * int(pct / 2)\n    print(f\"  Class {label} | {binary_labels[label]:<25} | \"\n          f\"{count:>4} samples | {pct:5.1f}% | {bar}\")\n\n# ==========================================\n# STEP 4: IMBALANCE METRICS\n# ==========================================\nprint(\"\\n\" + \"─\"*55)\nprint(\" IMBALANCE METRICS\")\nprint(\"─\"*55)\n\nmajority_count = binary_counts.max()\nminority_count = binary_counts.min()\nmajority_class = binary_counts.idxmax()\nminority_class = binary_counts.idxmin()\n\n# Imbalance Ratio\nimbalance_ratio = majority_count / minority_count\n\n# Minority Class Percentage\nminority_pct = (minority_count / len(df)) * 100\n\n# Imbalance Degree\ndef get_imbalance_degree(ratio):\n    if ratio < 1.5:\n        return \"✅ BALANCED\"\n    elif ratio < 2.5:\n        return \"🟡 SLIGHTLY IMBALANCED\"\n    elif ratio < 4.0:\n        return \"🟠 MODERATELY IMBALANCED\"\n    elif ratio < 10.0:\n        return \"🔴 IMBALANCED\"\n    else:\n        return \"🚨 SEVERELY IMBALANCED\"\n\ndegree = get_imbalance_degree(imbalance_ratio)\n\nprint(f\"\\n  Majority Class  : Class {majority_class} \"\n      f\"({binary_labels[majority_class]}) → {majority_count} samples\")\nprint(f\"  Minority Class  : Class {minority_class} \"\n      f\"({binary_labels[minority_class]}) → {minority_count} samples\")\nprint(f\"\\n  Imbalance Ratio : {imbalance_ratio:.2f}:1\")\nprint(f\"  Minority %      : {minority_pct:.1f}%\")\nprint(f\"  Verdict         : {degree}\")\n\n# ==========================================\n# STEP 5: STATISTICAL TESTS FOR IMBALANCE\n# ==========================================\nprint(\"\\n\" + \"─\"*55)\nprint(\" STATISTICAL TESTS\")\nprint(\"─\"*55)\n\n# Chi-Square Test\n# H0: Distribution is uniform (balanced)\n# H1: Distribution is NOT uniform (imbalanced)\nobserved   = np.array(list(binary_counts.values))\nexpected   = np.array([len(df)/2, len(df)/2])  # perfectly balanced expectation\n\nchi2_stat  = np.sum((observed - expected)**2 / expected)\n# degrees of freedom = n_classes - 1 = 1\nfrom scipy.stats import chi2\np_value    = 1 - chi2.cdf(chi2_stat, df=1)\n\nprint(f\"\\n  Chi-Square Test (vs. Uniform Distribution)\")\nprint(f\"  χ² Statistic : {chi2_stat:.4f}\")\nprint(f\"  p-value      : {p_value:.6f}\")\nif p_value < 0.05:\n    print(f\"  Result       : ✅ Significant imbalance detected (p < 0.05)\")\nelse:\n    print(f\"  Result       : ❌ No significant imbalance (p ≥ 0.05)\")\n\n# Entropy-based imbalance measure\n# Max entropy = perfectly balanced, lower = more imbalanced\nprobs          = binary_counts.values / binary_counts.values.sum()\nentropy        = -np.sum(probs * np.log2(probs + 1e-10))\nmax_entropy    = np.log2(len(binary_counts))\nentropy_ratio  = entropy / max_entropy  # 1.0 = perfectly balanced, 0.0 = all one class\n\nprint(f\"\\n  Entropy-Based Imbalance Score\")\nprint(f\"  Entropy       : {entropy:.4f} bits\")\nprint(f\"  Max Entropy   : {max_entropy:.4f} bits (perfect balance)\")\nprint(f\"  Balance Score : {entropy_ratio:.4f}  \"\n      f\"(1.0 = perfect, 0.0 = all one class)\")\n\n# ==========================================\n# STEP 6: TRAIN/VAL SPLIT IMBALANCE CHECK\n# ==========================================\nprint(\"\\n\" + \"─\"*55)\nprint(\" SPLIT-WISE IMBALANCE CHECK\")\nprint(\"─\"*55)\n\nfrom sklearn.model_selection import train_test_split\n\ntrain_df, val_df = train_test_split(\n    df, test_size=0.2, random_state=42, stratify=df['binary_label']\n)\n\nfor split_name, split_df in [(\"Train (80%)\", train_df), (\"Val (20%)\", val_df)]:\n    split_counts = split_df['binary_label'].value_counts().sort_index()\n    print(f\"\\n  {split_name}\")\n    for label, count in split_counts.items():\n        pct = (count / len(split_df)) * 100\n        print(f\"    Class {label}: {count:>4} samples ({pct:.1f}%)\")\n    split_ratio = split_counts.max() / split_counts.min()\n    print(f\"    Imbalance Ratio: {split_ratio:.2f}:1  \"\n          f\"→ {get_imbalance_degree(split_ratio)}\")\n\n# ==========================================\n# STEP 7: VISUALIZATIONS\n# ==========================================\nfig, axes = plt.subplots(2, 3, figsize=(18, 11))\nfig.suptitle('Messidor-2 Dataset Imbalance Analysis', \n             fontsize=16, fontweight='bold', y=0.98)\n\ncolors_grade  = ['#2ecc71','#f1c40f','#e67e22','#e74c3c','#8e44ad']\ncolors_binary = ['#3498db','#e74c3c']\n\n# Plot 1: Raw grade bar chart\nax1 = axes[0, 0]\nbars = ax1.bar(\n    [grade_labels.get(g, str(g)) for g in grade_counts.index],\n    grade_counts.values,\n    color=colors_grade[:len(grade_counts)],\n    edgecolor='black', linewidth=0.8\n)\nfor bar, val in zip(bars, grade_counts.values):\n    ax1.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 5,\n             str(val), ha='center', va='bottom', fontweight='bold', fontsize=10)\nax1.set_title('DR Grade Distribution (Original)', fontsize=12, fontweight='bold')\nax1.set_xlabel('DR Grade')\nax1.set_ylabel('Sample Count')\nax1.tick_params(axis='x', rotation=15)\nax1.grid(axis='y', linestyle='--', alpha=0.5)\n\n# Plot 2: Binary distribution bar chart\nax2 = axes[0, 1]\nbinary_names = [binary_labels[k] for k in binary_counts.index]\nbars2 = ax2.bar(binary_names, binary_counts.values,\n                color=colors_binary, edgecolor='black', linewidth=0.8)\nfor bar, val in zip(bars2, binary_counts.values):\n    pct = val / len(df) * 100\n    ax2.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 5,\n             f'{val}\\n({pct:.1f}%)', ha='center', va='bottom',\n             fontweight='bold', fontsize=11)\nax2.set_title('Binary Class Distribution', fontsize=12, fontweight='bold')\nax2.set_ylabel('Sample Count')\nax2.set_xlabel('Class')\nax2.grid(axis='y', linestyle='--', alpha=0.5)\nax2.axhline(y=len(df)/2, color='green', linestyle='--',\n            linewidth=2, label='Ideal Balance')\nax2.legend()\n\n# Plot 3: Pie chart\nax3 = axes[0, 2]\nwedge_props = {'edgecolor': 'white', 'linewidth': 2}\nax3.pie(binary_counts.values,\n        labels=[f\"{binary_labels[k]}\\n({v} samples)\" \n                for k, v in binary_counts.items()],\n        colors=colors_binary,\n        autopct='%1.1f%%',\n        startangle=90,\n        wedgeprops=wedge_props,\n        textprops={'fontsize': 10})\nax3.set_title('Binary Class Proportion', fontsize=12, fontweight='bold')\n\n# Plot 4: Grade stacked bar (Normal vs DR breakdown)\nax4 = axes[1, 0]\nnormal_grades = grade_counts[grade_counts.index.isin([0,1])]\ndr_grades     = grade_counts[grade_counts.index.isin([2,3,4])]\n\nax4.bar(['Normal\\n(Grade 0-1)'], [normal_grades.sum()],\n        color='#3498db', label='Normal', edgecolor='black')\nbottom = 0\nfor grade, count in dr_grades.items():\n    ax4.bar(['Referable DR\\n(Grade 2-4)'], [count], bottom=bottom,\n            color=colors_grade[grade], label=grade_labels[grade],\n            edgecolor='black')\n    ax4.text(1, bottom + count/2, f'G{grade}: {count}',\n             ha='center', va='center', fontweight='bold',\n             color='white', fontsize=9)\n    bottom += count\n\nax4.set_title('Grade Breakdown Within Binary Classes',\n              fontsize=12, fontweight='bold')\nax4.set_ylabel('Sample Count')\nax4.legend(loc='upper right', fontsize=8)\nax4.grid(axis='y', linestyle='--', alpha=0.5)\n\n# Plot 5: Train/Val split comparison\nax5 = axes[1, 1]\nx      = np.arange(2)\nwidth  = 0.3\ntrain_vals = [train_df['binary_label'].value_counts().get(i, 0) for i in [0,1]]\nval_vals   = [val_df['binary_label'].value_counts().get(i, 0) for i in [0,1]]\n\nax5.bar(x - width/2, train_vals, width, label='Train', \n        color='#3498db', edgecolor='black', alpha=0.85)\nax5.bar(x + width/2, val_vals,   width, label='Val',\n        color='#e74c3c', edgecolor='black', alpha=0.85)\nax5.set_xticks(x)\nax5.set_xticklabels(['Normal', 'Referable DR'])\nax5.set_title('Train vs Val Split Distribution',\n              fontsize=12, fontweight='bold')\nax5.set_ylabel('Sample Count')\nax5.legend()\nax5.grid(axis='y', linestyle='--', alpha=0.5)\n\n# Plot 6: Imbalance summary metrics\nax6 = axes[1, 2]\nax6.axis('off')\nmetrics_text = [\n    (\"Total Samples\",    f\"{len(df)}\"),\n    (\"Majority Class\",   f\"Class {majority_class} → {majority_count}\"),\n    (\"Minority Class\",   f\"Class {minority_class} → {minority_count}\"),\n    (\"Imbalance Ratio\",  f\"{imbalance_ratio:.2f} : 1\"),\n    (\"Minority %\",       f\"{minority_pct:.1f}%\"),\n    (\"Balance Score\",    f\"{entropy_ratio:.4f} / 1.0\"),\n    (\"χ² p-value\",       f\"{p_value:.6f}\"),\n    (\"Verdict\",          degree),\n]\ny_pos = 0.95\nax6.text(0.5, 1.02, 'Summary Report', transform=ax6.transAxes,\n         fontsize=13, fontweight='bold', ha='center', va='top')\nfor key, val in metrics_text:\n    color = '#e74c3c' if key == 'Verdict' else 'black'\n    weight = 'bold' if key == 'Verdict' else 'normal'\n    ax6.text(0.05, y_pos, f\"{key}:\", transform=ax6.transAxes,\n             fontsize=10, fontweight='bold', va='top')\n    ax6.text(0.55, y_pos, val, transform=ax6.transAxes,\n             fontsize=10, color=color, fontweight=weight, va='top')\n    y_pos -= 0.11\n\nplt.tight_layout()\nplt.savefig('messidor_imbalance_analysis.png', dpi=150, bbox_inches='tight')\nplt.show()\n\n# ==========================================\n# STEP 8: RECOMMENDED HANDLING STRATEGIES\n# ==========================================\nprint(\"\\n\" + \"=\"*55)\nprint(\" RECOMMENDED HANDLING STRATEGIES\")\nprint(\"=\"*55)\n\nif imbalance_ratio >= 4.0:\n    print(\"\"\"\n  Your dataset IS imbalanced. Here's what to apply:\n\n  1. CLASS WEIGHTS IN LOSS (already done ✅)\n     weights = [1.0, imbalance_ratio]\n     \n  2. OVERSAMPLE MINORITY CLASS\n     from imblearn.over_sampling import SMOTE\n     # or simply duplicate minority samples\n     \n  3. INCREASE CLASS WEIGHT FURTHER\n     Current: [1.0, 3.5] → Try [1.0, imbalance_ratio]\n     \n  4. USE STRATIFIED SAMPLING (already done ✅)\n     train_test_split(..., stratify=df['binary_label'])\n     \n  5. MONITOR RECALL > ACCURACY\n     A model predicting all-Normal gets high accuracy\n     but zero clinical value — always check recall/F1\n    \"\"\")\nelif imbalance_ratio >= 2.5:\n    print(\"\"\"\n  Your dataset is MODERATELY imbalanced.\n  Current handling (focal loss + class weights) \n  should be sufficient. Monitor validation recall.\n    \"\"\")\nelse:\n    print(\"\"\"\n  Your dataset is relatively BALANCED.\n  No special handling needed beyond current setup.\n    \"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T19:57:56.122113Z","iopub.execute_input":"2026-03-29T19:57:56.122509Z","iopub.status.idle":"2026-03-29T19:58:09.517874Z","shell.execute_reply.started":"2026-03-29T19:57:56.122462Z","shell.execute_reply":"2026-03-29T19:58:09.517162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nimport seaborn as sns\nfrom collections import Counter\nfrom scipy.stats import chi2\nfrom sklearn.model_selection import train_test_split\n\n# ==========================================\n# STEP 1: LOAD THE APTOS DATASET\n# ==========================================\n# Search for APTOS CSV\ncsv_search = glob.glob('/kaggle/input/**/train.csv', recursive=True)\nif not csv_search:\n    raise FileNotFoundError(\"Could not find train.csv. \"\n                            \"Make sure APTOS 2019 dataset is added.\")\nCSV_PATH = csv_search[0]\ndf       = pd.read_csv(CSV_PATH)\n\n# APTOS fixed column names\ngrade_col = 'diagnosis'\nid_col    = 'id_code'\n\n# Validate columns exist\nif grade_col not in df.columns:\n    raise ValueError(f\"'diagnosis' column not found. \"\n                     f\"Available columns: {df.columns.tolist()}\")\n\nprint(\"=\"*55)\nprint(\"        APTOS 2019 DATASET IMBALANCE ANALYSIS\")\nprint(\"=\"*55)\nprint(f\"\\nDataset Shape  : {df.shape}\")\nprint(f\"Grade Column   : '{grade_col}'\")\nprint(f\"ID Column      : '{id_col}'\")\nprint(f\"Total Samples  : {len(df)}\")\nprint(f\"Unique Grades  : {sorted(df[grade_col].unique())}\\n\")\n\n# ==========================================\n# STEP 2: RAW GRADE DISTRIBUTION\n# ==========================================\nprint(\"─\"*55)\nprint(\" ORIGINAL DR GRADES (0-4 Scale)\")\nprint(\"─\"*55)\n\ngrade_counts = df[grade_col].value_counts().sort_index()\ngrade_labels = {\n    0: 'No DR',\n    1: 'Mild DR',\n    2: 'Moderate DR',\n    3: 'Severe DR',\n    4: 'Proliferative DR'\n}\n\nfor grade, count in grade_counts.items():\n    label = grade_labels.get(grade, f'Grade {grade}')\n    pct   = (count / len(df)) * 100\n    bar   = '█' * int(pct / 2)\n    print(f\"  Grade {grade} | {label:<20} | {count:>4} samples \"\n          f\"| {pct:5.1f}% | {bar}\")\n\n# ==========================================\n# STEP 3: BINARY LABEL DISTRIBUTION\n# ==========================================\nprint(\"\\n\" + \"─\"*55)\nprint(\" BINARY LABELS (Normal vs Referable DR)\")\nprint(\"─\"*55)\n\ndf['binary_label'] = df[grade_col].apply(lambda x: 0 if x in [0, 1] else 1)\nbinary_counts      = df['binary_label'].value_counts().sort_index()\nbinary_labels      = {0: 'Normal (Grade 0-1)', 1: 'Referable DR (Grade 2-4)'}\n\nfor label, count in binary_counts.items():\n    pct = (count / len(df)) * 100\n    bar = '█' * int(pct / 2)\n    print(f\"  Class {label} | {binary_labels[label]:<25} | \"\n          f\"{count:>4} samples | {pct:5.1f}% | {bar}\")\n\n# ==========================================\n# STEP 4: MULTI-CLASS IMBALANCE (All 5 grades)\n# APTOS is a 5-class problem unlike Messidor\n# ==========================================\nprint(\"\\n\" + \"─\"*55)\nprint(\" MULTI-CLASS IMBALANCE (All 5 Grades)\")\nprint(\"─\"*55)\n\nmax_grade_count = grade_counts.max()\nmin_grade_count = grade_counts.min()\nmax_grade       = grade_counts.idxmax()\nmin_grade       = grade_counts.idxmin()\nmulticlass_ratio = max_grade_count / min_grade_count\n\nprint(f\"\\n  Most Common  : Grade {max_grade} \"\n      f\"({grade_labels[max_grade]}) → {max_grade_count} samples\")\nprint(f\"  Least Common : Grade {min_grade} \"\n      f\"({grade_labels[min_grade]}) → {min_grade_count} samples\")\nprint(f\"  Multi-Class Imbalance Ratio : {multiclass_ratio:.2f}:1\")\n\n# ==========================================\n# STEP 5: BINARY IMBALANCE METRICS\n# ==========================================\nprint(\"\\n\" + \"─\"*55)\nprint(\" BINARY IMBALANCE METRICS\")\nprint(\"─\"*55)\n\nmajority_count  = binary_counts.max()\nminority_count  = binary_counts.min()\nmajority_class  = binary_counts.idxmax()\nminority_class  = binary_counts.idxmin()\nimbalance_ratio = majority_count / minority_count\nminority_pct    = (minority_count / len(df)) * 100\n\ndef get_imbalance_degree(ratio):\n    if ratio < 1.5:\n        return \"✅ BALANCED\"\n    elif ratio < 2.5:\n        return \"🟡 SLIGHTLY IMBALANCED\"\n    elif ratio < 4.0:\n        return \"🟠 MODERATELY IMBALANCED\"\n    elif ratio < 10.0:\n        return \"🔴 IMBALANCED\"\n    else:\n        return \"🚨 SEVERELY IMBALANCED\"\n\ndegree = get_imbalance_degree(imbalance_ratio)\n\nprint(f\"\\n  Majority Class  : Class {majority_class} \"\n      f\"({binary_labels[majority_class]}) → {majority_count} samples\")\nprint(f\"  Minority Class  : Class {minority_class} \"\n      f\"({binary_labels[minority_class]}) → {minority_count} samples\")\nprint(f\"\\n  Imbalance Ratio : {imbalance_ratio:.2f}:1\")\nprint(f\"  Minority %      : {minority_pct:.1f}%\")\nprint(f\"  Verdict         : {degree}\")\n\n# ==========================================\n# STEP 6: STATISTICAL TESTS\n# ==========================================\nprint(\"\\n\" + \"─\"*55)\nprint(\" STATISTICAL TESTS\")\nprint(\"─\"*55)\n\n# --- Binary Chi-Square Test ---\nobserved_bin = np.array(list(binary_counts.values))\nexpected_bin = np.array([len(df) / 2, len(df) / 2])\nchi2_bin     = np.sum((observed_bin - expected_bin)**2 / expected_bin)\np_binary     = 1 - chi2.cdf(chi2_bin, df=1)\n\nprint(f\"\\n  [Binary] Chi-Square Test (Normal vs Referable DR)\")\nprint(f\"  χ² Statistic : {chi2_bin:.4f}\")\nprint(f\"  p-value      : {p_binary:.6f}\")\nprint(f\"  Result       : {'✅ Significant imbalance (p<0.05)' if p_binary < 0.05 else '❌ No significant imbalance'}\")\n\n# --- Multi-Class Chi-Square Test ---\nobserved_mc  = np.array(list(grade_counts.values))\nexpected_mc  = np.full(len(grade_counts), len(df) / len(grade_counts))\nchi2_mc      = np.sum((observed_mc - expected_mc)**2 / expected_mc)\np_multiclass = 1 - chi2.cdf(chi2_mc, df=len(grade_counts) - 1)\n\nprint(f\"\\n  [Multi-Class] Chi-Square Test (All 5 Grades)\")\nprint(f\"  χ² Statistic : {chi2_mc:.4f}\")\nprint(f\"  p-value      : {p_multiclass:.6f}\")\nprint(f\"  Result       : {'✅ Significant imbalance (p<0.05)' if p_multiclass < 0.05 else '❌ No significant imbalance'}\")\n\n# --- Entropy Score ---\nprobs_bin      = binary_counts.values / binary_counts.values.sum()\nentropy_bin    = -np.sum(probs_bin * np.log2(probs_bin + 1e-10))\nmax_entropy    = np.log2(2)\nentropy_ratio  = entropy_bin / max_entropy\n\n# Multi-class entropy\nprobs_mc       = grade_counts.values / grade_counts.values.sum()\nentropy_mc     = -np.sum(probs_mc * np.log2(probs_mc + 1e-10))\nmax_entropy_mc = np.log2(5)\nentropy_mc_ratio = entropy_mc / max_entropy_mc\n\nprint(f\"\\n  Entropy Score (Binary)      : {entropy_ratio:.4f} / 1.0\")\nprint(f\"  Entropy Score (Multi-Class) : {entropy_mc_ratio:.4f} / 1.0\")\nprint(f\"  (1.0 = perfect balance, 0.0 = all one class)\")\n\n# ==========================================\n# STEP 7: PER-GRADE IMBALANCE RATIO TABLE\n# Unique to APTOS — shows each grade vs \n# the ideal uniform count\n# ==========================================\nprint(\"\\n\" + \"─\"*55)\nprint(\" PER-GRADE DEVIATION FROM BALANCE\")\nprint(\"─\"*55)\n\nideal_count = len(df) / 5\nprint(f\"\\n  Ideal count per grade (if balanced): {ideal_count:.0f}\\n\")\nprint(f\"  {'Grade':<8} {'Label':<20} {'Count':>6} \"\n      f\"{'Ideal':>6} {'Ratio':>7} {'Over/Under'}\")\nprint(f\"  {'─'*65}\")\n\nfor grade, count in grade_counts.items():\n    label     = grade_labels[grade]\n    ratio     = count / ideal_count\n    status    = f\"OVER  by {ratio-1:+.1%}\" if ratio > 1 else f\"UNDER by {ratio-1:+.1%}\"\n    print(f\"  Grade {grade}  | {label:<20} | {count:>5} \"\n          f\"| {ideal_count:>5.0f} | {ratio:>6.2f}x | {status}\")\n\n# ==========================================\n# STEP 8: TRAIN/VAL SPLIT CHECK\n# ==========================================\nprint(\"\\n\" + \"─\"*55)\nprint(\" SPLIT-WISE IMBALANCE CHECK\")\nprint(\"─\"*55)\n\ntrain_df, val_df = train_test_split(\n    df, test_size=0.2, random_state=42, stratify=df['binary_label']\n)\n\nfor split_name, split_df in [(\"Train (80%)\", train_df), (\"Val (20%)\", val_df)]:\n    split_counts = split_df['binary_label'].value_counts().sort_index()\n    print(f\"\\n  {split_name}  (Total: {len(split_df)})\")\n    for label, count in split_counts.items():\n        pct = (count / len(split_df)) * 100\n        print(f\"    Class {label} ({binary_labels[label]:<25}): \"\n              f\"{count:>4} samples ({pct:.1f}%)\")\n    split_ratio = split_counts.max() / split_counts.min()\n    print(f\"    Imbalance Ratio : {split_ratio:.2f}:1 \"\n          f\"→ {get_imbalance_degree(split_ratio)}\")\n\n# Multi-class split check\nprint(f\"\\n  Grade-level split check:\")\nfor split_name, split_df in [(\"Train\", train_df), (\"Val\", val_df)]:\n    gc = split_df[grade_col].value_counts().sort_index()\n    print(f\"\\n  {split_name}:\")\n    for g, c in gc.items():\n        print(f\"    Grade {g} ({grade_labels[g]:<20}): {c:>4} \"\n              f\"({c/len(split_df)*100:.1f}%)\")\n\n# ==========================================\n# STEP 9: VISUALIZATIONS (7 panels)\n# Extra panel for APTOS multi-class analysis\n# ==========================================\nfig, axes = plt.subplots(2, 4, figsize=(22, 11))\nfig.suptitle('APTOS 2019 Dataset Imbalance Analysis',\n             fontsize=16, fontweight='bold', y=0.98)\n\ncolors_grade  = ['#2ecc71', '#f1c40f', '#e67e22', '#e74c3c', '#8e44ad']\ncolors_binary = ['#3498db', '#e74c3c']\n\n# Plot 1: Raw Grade Bar Chart\nax1 = axes[0, 0]\nbars = ax1.bar(\n    [grade_labels.get(g, str(g)) for g in grade_counts.index],\n    grade_counts.values,\n    color=colors_grade,\n    edgecolor='black', linewidth=0.8\n)\nfor bar, val in zip(bars, grade_counts.values):\n    ax1.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 5,\n             str(val), ha='center', va='bottom', fontweight='bold', fontsize=9)\nax1.axhline(y=ideal_count, color='blue', linestyle='--',\n            linewidth=1.5, label=f'Ideal ({ideal_count:.0f})')\nax1.set_title('Grade Distribution (All 5 Classes)', fontsize=11, fontweight='bold')\nax1.set_xlabel('DR Grade')\nax1.set_ylabel('Sample Count')\nax1.tick_params(axis='x', rotation=20)\nax1.legend(fontsize=9)\nax1.grid(axis='y', linestyle='--', alpha=0.5)\n\n# Plot 2: Binary Distribution Bar Chart\nax2 = axes[0, 1]\nbars2 = ax2.bar(\n    [binary_labels[k] for k in binary_counts.index],\n    binary_counts.values,\n    color=colors_binary, edgecolor='black', linewidth=0.8\n)\nfor bar, val in zip(bars2, binary_counts.values):\n    pct = val / len(df) * 100\n    ax2.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 10,\n             f'{val}\\n({pct:.1f}%)', ha='center', va='bottom',\n             fontweight='bold', fontsize=10)\nax2.axhline(y=len(df)/2, color='green', linestyle='--',\n            linewidth=2, label='Ideal Balance')\nax2.set_title('Binary Class Distribution', fontsize=11, fontweight='bold')\nax2.set_ylabel('Sample Count')\nax2.legend()\nax2.grid(axis='y', linestyle='--', alpha=0.5)\n\n# Plot 3: Pie Chart (Binary)\nax3 = axes[0, 2]\nax3.pie(\n    binary_counts.values,\n    labels=[f\"{binary_labels[k]}\\n({v})\" for k, v in binary_counts.items()],\n    colors=colors_binary,\n    autopct='%1.1f%%',\n    startangle=90,\n    wedgeprops={'edgecolor': 'white', 'linewidth': 2},\n    textprops={'fontsize': 9}\n)\nax3.set_title('Binary Class Proportion', fontsize=11, fontweight='bold')\n\n# Plot 4: Pie Chart (Multi-Class)\nax4 = axes[0, 3]\nax4.pie(\n    grade_counts.values,\n    labels=[f\"G{g}: {grade_labels[g]}\\n({v})\" \n            for g, v in grade_counts.items()],\n    colors=colors_grade,\n    autopct='%1.1f%%',\n    startangle=90,\n    wedgeprops={'edgecolor': 'white', 'linewidth': 2},\n    textprops={'fontsize': 8}\n)\nax4.set_title('Grade-Level Proportion (5 Classes)', fontsize=11, fontweight='bold')\n\n# Plot 5: Deviation from Ideal Balance\nax5 = axes[1, 0]\ndeviations = [(grade_counts[g] / ideal_count - 1) * 100 \n              for g in grade_counts.index]\nbar_colors = ['#e74c3c' if d > 0 else '#3498db' for d in deviations]\nbars5 = ax5.bar(\n    [f\"G{g}\\n{grade_labels[g][:8]}\" for g in grade_counts.index],\n    deviations,\n    color=bar_colors, edgecolor='black', linewidth=0.8\n)\nfor bar, val in zip(bars5, deviations):\n    ax5.text(bar.get_x() + bar.get_width()/2,\n             bar.get_height() + (2 if val >= 0 else -5),\n             f'{val:+.1f}%', ha='center', va='bottom',\n             fontweight='bold', fontsize=9)\nax5.axhline(y=0, color='black', linewidth=1.5)\nax5.set_title('% Deviation from Perfect Balance', fontsize=11, fontweight='bold')\nax5.set_ylabel('% Over/Under Ideal Count')\nax5.grid(axis='y', linestyle='--', alpha=0.5)\nred_patch  = mpatches.Patch(color='#e74c3c', label='Overrepresented')\nblue_patch = mpatches.Patch(color='#3498db', label='Underrepresented')\nax5.legend(handles=[red_patch, blue_patch], fontsize=8)\n\n# Plot 6: Train vs Val Split\nax6 = axes[1, 1]\nx      = np.arange(2)\nwidth  = 0.3\nt_vals = [train_df['binary_label'].value_counts().get(i, 0) for i in [0, 1]]\nv_vals = [val_df['binary_label'].value_counts().get(i, 0) for i in [0, 1]]\nax6.bar(x - width/2, t_vals, width, label='Train',\n        color='#3498db', edgecolor='black', alpha=0.85)\nax6.bar(x + width/2, v_vals, width, label='Val',\n        color='#e74c3c', edgecolor='black', alpha=0.85)\nax6.set_xticks(x)\nax6.set_xticklabels(['Normal', 'Referable DR'])\nax6.set_title('Train vs Val Split Distribution', fontsize=11, fontweight='bold')\nax6.set_ylabel('Sample Count')\nax6.legend()\nax6.grid(axis='y', linestyle='--', alpha=0.5)\n\n# Plot 7: Multi-Class Train/Val\nax7 = axes[1, 2]\nx2     = np.arange(5)\nt_mc   = [train_df[grade_col].value_counts().get(g, 0) for g in range(5)]\nv_mc   = [val_df[grade_col].value_counts().get(g, 0) for g in range(5)]\nax7.bar(x2 - width/2, t_mc, width, label='Train',\n        color='#3498db', edgecolor='black', alpha=0.85)\nax7.bar(x2 + width/2, v_mc, width, label='Val',\n        color='#e74c3c', edgecolor='black', alpha=0.85)\nax7.set_xticks(x2)\nax7.set_xticklabels([f'G{g}' for g in range(5)])\nax7.set_title('Grade-Level Train/Val Split', fontsize=11, fontweight='bold')\nax7.set_ylabel('Sample Count')\nax7.legend()\nax7.grid(axis='y', linestyle='--', alpha=0.5)\n\n# Plot 8: Summary Report Card\nax8 = axes[1, 3]\nax8.axis('off')\nmetrics = [\n    (\"Total Samples\",          f\"{len(df)}\"),\n    (\"Majority Class\",         f\"Class {majority_class} → {majority_count}\"),\n    (\"Minority Class\",         f\"Class {minority_class} → {minority_count}\"),\n    (\"Binary Imbalance Ratio\", f\"{imbalance_ratio:.2f} : 1\"),\n    (\"Multi-Class Ratio\",      f\"{multiclass_ratio:.2f} : 1\"),\n    (\"Minority %\",             f\"{minority_pct:.1f}%\"),\n    (\"Binary Balance Score\",   f\"{entropy_ratio:.4f} / 1.0\"),\n    (\"Multi-Class Bal. Score\", f\"{entropy_mc_ratio:.4f} / 1.0\"),\n    (\"χ² p-value (Binary)\",    f\"{p_binary:.6f}\"),\n    (\"χ² p-value (Multi)\",     f\"{p_multiclass:.6f}\"),\n    (\"Verdict\",                degree),\n]\nax8.text(0.5, 1.02, 'Summary Report', transform=ax8.transAxes,\n         fontsize=12, fontweight='bold', ha='center', va='top')\ny_pos = 0.93\nfor key, val in metrics:\n    color  = '#e74c3c' if key == 'Verdict' else 'black'\n    weight = 'bold'    if key == 'Verdict' else 'normal'\n    ax8.text(0.02, y_pos, f\"{key}:\", transform=ax8.transAxes,\n             fontsize=8.5, fontweight='bold', va='top')\n    ax8.text(0.60, y_pos, val, transform=ax8.transAxes,\n             fontsize=8.5, color=color, fontweight=weight, va='top')\n    y_pos -= 0.085\n\nplt.tight_layout()\nplt.savefig('aptos_imbalance_analysis.png', dpi=150, bbox_inches='tight')\nplt.show()\n\n# ==========================================\n# STEP 10: RECOMMENDED CLASS WEIGHTS\n# Compute exact weights based on actual data\n# ==========================================\nprint(\"\\n\" + \"=\"*55)\nprint(\" COMPUTED CLASS WEIGHTS FOR TRAINING\")\nprint(\"=\"*55)\n\n# Binary weights\ntotal      = len(df)\nn_classes  = 2\nw0 = total / (n_classes * binary_counts[0])\nw1 = total / (n_classes * binary_counts[1])\n\nprint(f\"\\n  Balanced Binary Class Weights:\")\nprint(f\"    Normal weight      : {w0:.4f}\")\nprint(f\"    Referable DR weight: {w1:.4f}\")\nprint(f\"\\n  → In your training code use:\")\nprint(f\"    weights = torch.tensor([{w0:.2f}, {w1:.2f}])\")\n\n# Multi-class weights (for 5-class training)\nprint(f\"\\n  Balanced Multi-Class Weights (if training 5-class):\")\nfor grade, count in grade_counts.items():\n    w = total / (5 * count)\n    print(f\"    Grade {grade} ({grade_labels[grade]:<20}): {w:.4f}\")\n\n# ==========================================\n# STEP 11: RECOMMENDATION\n# ==========================================\nprint(\"\\n\" + \"=\"*55)\nprint(\" RECOMMENDATION\")\nprint(\"=\"*55)\n\nif imbalance_ratio >= 4.0:\n    print(f\"\"\"\n  ❗ Dataset is IMBALANCED (ratio = {imbalance_ratio:.2f}:1)\n  \n  Apply ALL of the following:\n  1. Use computed weights above in your loss function\n  2. Use stratified train/val split (already coded above)\n  3. Consider oversampling minority class with augmentation\n  4. Monitor Recall + F1 score, not just Accuracy\n  5. Use Focal Loss (gamma=2) — already in your pipeline ✅\n    \"\"\")\nelif imbalance_ratio >= 2.5:\n    print(f\"\"\"\n  ⚠️  Dataset is MODERATELY IMBALANCED (ratio = {imbalance_ratio:.2f}:1)\n  \n  Apply the following:\n  1. Use computed weights above in your loss function\n  2. Focal Loss is sufficient — already in pipeline ✅\n  3. Monitor Quadratic Kappa + Recall alongside Accuracy\n    \"\"\")\nelse:\n    print(f\"\"\"\n  ✅ Dataset is relatively BALANCED (ratio = {imbalance_ratio:.2f}:1)\n  \n  Current setup is fine. Just monitor:\n  1. Per-class accuracy (not just overall)\n  2. Quadratic Kappa score\n    \"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T20:08:32.138572Z","iopub.execute_input":"2026-03-29T20:08:32.13887Z","iopub.status.idle":"2026-03-29T20:08:46.118269Z","shell.execute_reply.started":"2026-03-29T20:08:32.138844Z","shell.execute_reply":"2026-03-29T20:08:46.11748Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport time\nimport torch\nimport numpy as np\nimport pandas as pd\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom sklearn.metrics import cohen_kappa_score, accuracy_score, roc_curve, auc\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import label_binarize\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm import tqdm\nimport timm\nimport matplotlib.pyplot as plt\n\n# Start runtime tracker\nTOTAL_START_TIME = time.time()\n\n# ==========================================\n# 1. BEAT-THE-PAPER CONFIGURATION\n# ==========================================\nclass Config:\n    # APTOS 2019 Paths\n    BASE_PATH = '/kaggle/input/aptos2019-blindness-detection/'\n    if not os.path.exists(BASE_PATH):\n        BASE_PATH = '/kaggle/input/competitions/aptos2019-blindness-detection/'\n        \n    IMG_DIR = os.path.join(BASE_PATH, 'train_images/')\n    CSV_PATH = os.path.join(BASE_PATH, 'train.csv')\n    \n    # Advanced Architecture & Hyperparameters\n    MODEL_NAME = 'swin_large_patch4_window12_384'\n    IMG_SIZE = 512 \n    NUM_CLASSES = 5\n    BATCH_SIZE = 4  # Keep very small for Swin-Large at 512x512 to prevent OOM\n    EPOCHS = 50\n    MAX_LR = 1e-4\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Class Weights for Focal Loss [0, 1, 2, 3, 4]\n    CLASS_WEIGHTS = [0.41, 1.98, 0.73, 3.79, 2.48]\n\n# ==========================================\n# 2. ALBUMENTATIONS & DATASET\n# ==========================================\n# Heavy augmentations to prevent Swin-Large memorization\ntrain_transforms = A.Compose([\n    A.Resize(Config.IMG_SIZE, Config.IMG_SIZE),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomRotate90(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.15, rotate_limit=60, p=0.5),\n    A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1, p=0.5),\n    A.CLAHE(clip_limit=2.0, tile_grid_size=(8, 8), p=0.5),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\nval_transforms = A.Compose([\n    A.Resize(Config.IMG_SIZE, Config.IMG_SIZE),\n    A.CLAHE(clip_limit=2.0, tile_grid_size=(8, 8), p=1.0),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\nclass APTOSDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_path = self.df.iloc[idx]['path']\n        label = self.df.iloc[idx]['diagnosis']\n        \n        image = cv2.imread(img_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n            \n        return image, torch.tensor(label, dtype=torch.long)\n\n# ==========================================\n# 3. SWIN-LARGE ARCHITECTURE\n# ==========================================\nclass SwinLarge_DR_Net(nn.Module):\n    def __init__(self, num_classes=Config.NUM_CLASSES):\n        super().__init__()\n        # Initialize Swin-Large with dynamic resolution capability (384 -> 512)\n        self.swin = timm.create_model(\n            Config.MODEL_NAME, \n            pretrained=True, \n            num_classes=0,\n            img_size=Config.IMG_SIZE \n        )\n        self.classifier = nn.Sequential(\n            nn.Linear(1536, 512), # Swin-Large outputs 1536-dim features\n            nn.GELU(),\n            nn.Dropout(0.5),\n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x):\n        features = self.swin(x)\n        return self.classifier(features)\n\n# ==========================================\n# 4. FOCAL LOSS WITH LABEL SMOOTHING\n# ==========================================\nclass FocalLossWithSmoothing(nn.Module):\n    def __init__(self, weights, alpha=1.0, gamma=2.0, label_smoothing=0.05):\n        super().__init__()\n        self.weights = torch.tensor(weights).to(Config.DEVICE)\n        self.gamma = gamma\n        self.smoothing = label_smoothing\n\n    def forward(self, inputs, targets):\n        ce_loss = F.cross_entropy(\n            inputs, targets, \n            weight=self.weights, \n            label_smoothing=self.smoothing, \n            reduction='none'\n        )\n        pt = torch.exp(-ce_loss)\n        focal_loss = ((1 - pt) ** self.gamma) * ce_loss\n        return focal_loss.mean()\n\n# ==========================================\n# 5. TRAINING & VALIDATION LOOPS\n# ==========================================\ndef train_one_epoch(model, loader, criterion, optimizer, scheduler, device):\n    model.train()\n    running_loss = 0.0\n    all_preds, all_targets = [], []\n    \n    pbar = tqdm(loader, desc=\"Training\")\n    for images, labels in pbar:\n        images, labels = images.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        \n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optimizer.step()\n        scheduler.step() # OneCycleLR steps per batch\n        \n        running_loss += loss.item()\n        _, preds = torch.max(outputs, 1)\n        all_preds.extend(preds.cpu().numpy())\n        all_targets.extend(labels.cpu().numpy())\n        \n        pbar.set_postfix({'loss': running_loss / len(all_preds) * Config.BATCH_SIZE})\n        \n    return running_loss / len(loader), accuracy_score(all_targets, all_preds)\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    all_preds, all_targets, all_probs = [], [], []\n    \n    with torch.no_grad():\n        for images, labels in tqdm(loader, desc=\"Validating\"):\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            probs = F.softmax(outputs, dim=1)\n            _, preds = torch.max(probs, 1)\n            \n            all_preds.extend(preds.cpu().numpy())\n            all_targets.extend(labels.cpu().numpy())\n            all_probs.extend(probs.cpu().numpy())\n            \n    acc = accuracy_score(all_targets, all_preds)\n    kappa = cohen_kappa_score(all_targets, all_preds, weights='quadratic')\n    return running_loss / len(loader), acc, kappa, np.array(all_targets), np.array(all_probs)\n\n# ==========================================\n# 6. PLOTTING FUNCTIONS\n# ==========================================\ndef plot_research_graphs(history, val_targets, val_probs):\n    epochs = range(1, len(history['train_loss']) + 1)\n    fig, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=(20, 5))\n    \n    # 1. Loss Curve\n    ax1.plot(epochs, history['train_loss'], 'b-', label='Train Loss', linewidth=2)\n    ax1.plot(epochs, history['val_loss'], 'r-', label='Val Loss', linewidth=2)\n    ax1.set_title('Training vs Validation Loss')\n    ax1.set_xlabel('Epochs')\n    ax1.set_ylabel('Loss')\n    ax1.legend()\n    ax1.grid(True, linestyle='--', alpha=0.7)\n    \n    # 2. Accuracy & Kappa Curve\n    ax2.plot(epochs, history['train_acc'], 'b--', label='Train Acc', alpha=0.6)\n    ax2.plot(epochs, history['val_acc'], 'g-', label='Val Acc', linewidth=2)\n    ax2.plot(epochs, history['val_kappa'], 'm-', label='Val Kappa', linewidth=2)\n    ax2.set_title('Accuracy and Quadratic Weighted Kappa')\n    ax2.set_xlabel('Epochs')\n    ax2.set_ylabel('Score')\n    ax2.legend()\n    ax2.grid(True, linestyle='--', alpha=0.7)\n    \n    # 3. Multi-Class ROC AUC Curve (Macro Average)\n    binarized_targets = label_binarize(val_targets, classes=[0, 1, 2, 3, 4])\n    fpr, tpr, roc_auc = dict(), dict(), dict()\n    \n    for i in range(Config.NUM_CLASSES):\n        fpr[i], tpr[i], _ = roc_curve(binarized_targets[:, i], val_probs[:, i])\n        roc_auc[i] = auc(fpr[i], tpr[i])\n        \n    # Compute Macro-Average ROC\n    all_fpr = np.unique(np.concatenate([fpr[i] for i in range(Config.NUM_CLASSES)]))\n    mean_tpr = np.zeros_like(all_fpr)\n    for i in range(Config.NUM_CLASSES):\n        mean_tpr += np.interp(all_fpr, fpr[i], tpr[i])\n    mean_tpr /= Config.NUM_CLASSES\n    \n    macro_auc = auc(all_fpr, mean_tpr)\n    \n    ax3.plot(all_fpr, mean_tpr, color='darkorange', lw=2, label=f'Macro-Average ROC (AUC = {macro_auc:.3f})')\n    ax3.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--')\n    ax3.set_xlim([0.0, 1.0])\n    ax3.set_ylim([0.0, 1.05])\n    ax3.set_xlabel('False Positive Rate')\n    ax3.set_ylabel('True Positive Rate')\n    ax3.set_title('Multi-Class ROC Curve (OvR)')\n    ax3.legend(loc=\"lower right\")\n    ax3.grid(True, linestyle='--', alpha=0.7)\n    \n    plt.tight_layout()\n    plt.show()\n\n# ==========================================\n# 7. MAIN EXECUTION\n# ==========================================\nif __name__ == '__main__':\n    print(f\"Using device: {Config.DEVICE}\")\n    \n    # Load APTOS Data\n    df = pd.read_csv(Config.CSV_PATH)\n    df['path'] = df['id_code'].apply(lambda x: os.path.join(Config.IMG_DIR, f\"{x}.png\"))\n    df = df[df['path'].apply(os.path.exists)] # Filter valid\n    \n    train_df, val_df = train_test_split(df, test_size=0.2, random_state=42, stratify=df['diagnosis'])\n    print(f\"Training on {len(train_df)} images, Validating on {len(val_df)} images.\")\n\n    # Apply WeightedRandomSampler to fix Grade 3 imbalance\n    class_counts = train_df['diagnosis'].value_counts().sort_index().values\n    class_weights = 1.0 / class_counts\n    sample_weights = [class_weights[label] for label in train_df['diagnosis'].values]\n    \n    sampler = WeightedRandomSampler(\n        weights=sample_weights, \n        num_samples=len(sample_weights), \n        replacement=True\n    )\n\n    # Datasets & Loaders\n    train_dataset = APTOSDataset(train_df, transform=train_transforms)\n    val_dataset = APTOSDataset(val_df, transform=val_transforms)\n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, sampler=sampler, num_workers=2)\n    val_loader = DataLoader(val_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=2)\n\n    # Initialize Model, Loss, Optimizer\n    model = SwinLarge_DR_Net().to(Config.DEVICE)\n    criterion = FocalLossWithSmoothing(weights=Config.CLASS_WEIGHTS)\n    \n    # Differential LR: Lower LR for backbone, higher for classifier\n    optimizer = optim.AdamW([\n        {'params': model.swin.parameters(), 'lr': Config.MAX_LR * 0.1},\n        {'params': model.classifier.parameters(), 'lr': Config.MAX_LR}\n    ], weight_decay=0.05)\n    \n    scheduler = optim.lr_scheduler.OneCycleLR(\n        optimizer, max_lr=[Config.MAX_LR * 0.1, Config.MAX_LR], \n        epochs=Config.EPOCHS, steps_per_epoch=len(train_loader), \n        pct_start=0.1\n    )\n\n    # Tracking\n    history = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': [], 'val_kappa': []}\n    best_kappa = -1.0\n    best_roc_data = None \n    \n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{Config.EPOCHS}\")\n        \n        train_loss, train_acc = train_one_epoch(model, train_loader, criterion, optimizer, scheduler, Config.DEVICE)\n        val_loss, val_acc, val_kappa, val_targets, val_probs = validate(model, val_loader, criterion, Config.DEVICE)\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_kappa'].append(val_kappa)\n        \n        print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}\")\n        print(f\"Val Loss:   {val_loss:.4f} | Val Acc:   {val_acc:.4f} | Val Kappa: {val_kappa:.4f}\")\n        \n        if val_kappa > best_kappa:\n            best_kappa = val_kappa\n            best_acc = val_acc\n            best_roc_data = (val_targets, val_probs)\n            torch.save(model.state_dict(), 'best_aptos_swin_large.pth')\n            print(\">>> Best Model Saved!\")\n\n    # Calculate Total Runtime\n    total_time = time.time() - TOTAL_START_TIME\n    hours, rem = divmod(total_time, 3600)\n    minutes, seconds = divmod(rem, 60)\n    \n    print(\"\\n\" + \"=\"*50)\n    print(f\"TRAINING COMPLETE (50 Epochs)\")\n    print(f\"Best Validation Accuracy: {best_acc:.4f}\")\n    print(f\"Best Validation Kappa:    {best_kappa:.4f}\")\n    print(f\"Total Runtime:            {int(hours)}h {int(minutes)}m {int(seconds)}s\")\n    print(\"=\"*50)\n    \n    # Plot final research graphs\n    if best_roc_data:\n        plot_research_graphs(history, best_roc_data[0], best_roc_data[1])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-29T20:36:30.816326Z","iopub.execute_input":"2026-03-29T20:36:30.816678Z","execution_failed":"2026-03-30T07:56:27.919Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# APTOS 2019 - COMPLETE TRAINING PIPELINE\n# Swin-Large | Focal Loss | TTA | Full Metrics Per Epoch\n# ============================================================\n\nimport os\nimport glob\nimport time\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision import transforms\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (cohen_kappa_score, accuracy_score,\n                             roc_curve, auc, roc_auc_score,\n                             confusion_matrix, classification_report)\nimport cv2\nimport timm\nimport matplotlib.pyplot as plt\nimport matplotlib.gridspec as gridspec\nimport seaborn as sns\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\nTOTAL_START = time.time()\n\n# ============================================================\n# 1. CONFIGURATION\n# ============================================================\nclass Config:\n    CSV_PATH    = glob.glob('/kaggle/input/**/train.csv',\n                            recursive=True)[0]\n    IMG_DIR     = glob.glob('/kaggle/input/**/train_images/',\n                            recursive=True)\n    IMG_DIR     = IMG_DIR[0] if IMG_DIR else \\\n                  '/kaggle/input/aptos2019-blindness-detection/train_images/'\n\n    IMG_SIZE    = 384          # Phase 1: 384. Change to 512 in Phase 3\n    NUM_CLASSES = 5\n    BATCH_SIZE  = 8\n    EPOCHS      = 50\n    LR          = 3e-5\n    DEVICE      = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    SEED        = 42\n\n    # APTOS class weights from imbalance analysis\n    CLASS_WEIGHTS = torch.tensor([0.4058, 1.9795, 0.7331, 3.7948, 2.4827])\n    GRADE_NAMES   = ['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative']\n\ntorch.manual_seed(Config.SEED)\nnp.random.seed(Config.SEED)\nprint(f\"Device      : {Config.DEVICE}\")\nprint(f\"Image Size  : {Config.IMG_SIZE}x{Config.IMG_SIZE}\")\nprint(f\"CSV Path    : {Config.CSV_PATH}\")\nprint(f\"Image Dir   : {Config.IMG_DIR}\")\n\n# ============================================================\n# 2. PREPROCESSING\n# ============================================================\ndef apply_clahe(img):\n    lab = cv2.cvtColor(img, cv2.COLOR_RGB2LAB)\n    l, a, b = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8, 8))\n    cl = clahe.apply(l)\n    return cv2.cvtColor(cv2.merge((cl, a, b)), cv2.COLOR_LAB2RGB)\n\ndef crop_black_border(img, tol=7):\n    \"\"\"Remove black borders common in fundus images.\"\"\"\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    mask = gray > tol\n    coords = np.argwhere(mask)\n    if coords.size == 0:\n        return img\n    r0, c0 = coords.min(axis=0)\n    r1, c1 = coords.max(axis=0) + 1\n    return img[r0:r1, c0:c1]\n\n# ============================================================\n# 3. DATASET\n# ============================================================\nclass APTOSDataset(Dataset):\n    def __init__(self, image_ids, labels, img_dir,\n                 img_size=Config.IMG_SIZE, is_train=True):\n        self.image_ids = image_ids\n        self.labels    = labels\n        self.img_dir   = img_dir\n        self.img_size  = img_size\n        self.is_train  = is_train\n\n        self.train_tf = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.RandomHorizontalFlip(p=0.5),\n            transforms.RandomVerticalFlip(p=0.5),\n            transforms.RandomRotation(180),\n            transforms.ColorJitter(brightness=0.2, contrast=0.2,\n                                   saturation=0.1, hue=0.05),\n            transforms.ToTensor(),\n            transforms.Normalize([0.485, 0.456, 0.406],\n                                 [0.229, 0.224, 0.225]),\n            transforms.RandomErasing(p=0.2, scale=(0.02, 0.1))\n        ])\n        self.val_tf = transforms.Compose([\n            transforms.ToPILImage(),\n            transforms.ToTensor(),\n            transforms.Normalize([0.485, 0.456, 0.406],\n                                 [0.229, 0.224, 0.225])\n        ])\n\n    def __len__(self):\n        return len(self.image_ids)\n\n    def __getitem__(self, idx):\n        img_id   = self.image_ids[idx]\n        img_path = os.path.join(self.img_dir, f\"{img_id}.png\")\n\n        img = cv2.imread(img_path)\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            img = crop_black_border(img)\n            img = cv2.resize(img, (self.img_size, self.img_size))\n            img = apply_clahe(img)\n\n        img    = self.train_tf(img) if self.is_train else self.val_tf(img)\n        label  = torch.tensor(self.labels[idx], dtype=torch.long)\n        return img, label\n\n# ============================================================\n# 4. MODEL: Swin-Large\n# ============================================================\nclass APTOS_SwinLarge(nn.Module):\n    def __init__(self, num_classes=Config.NUM_CLASSES):\n        super().__init__()\n        self.backbone = timm.create_model(\n            'swin_large_patch4_window12_384',\n            pretrained=True,\n            num_classes=0       # remove head → raw 1536-dim features\n        )\n        feat_dim = 1536         # Swin-Large output dim\n\n        # Squeeze-Excitation recalibration\n        self.se = nn.Sequential(\n            nn.Linear(feat_dim, feat_dim // 16),\n            nn.ReLU(),\n            nn.Linear(feat_dim // 16, feat_dim),\n            nn.Sigmoid()\n        )\n\n        # Deep classifier head\n        self.head = nn.Sequential(\n            nn.LayerNorm(feat_dim),\n            nn.Linear(feat_dim, 512),\n            nn.GELU(),\n            nn.Dropout(0.4),\n            nn.Linear(512, 256),\n            nn.GELU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, num_classes)\n        )\n\n    def forward(self, x):\n        f     = self.backbone(x)       # (B, 1536)\n        scale = self.se(f)\n        f     = f * scale              # SE recalibration\n        return self.head(f)            # (B, 5)\n\n# ============================================================\n# 5. LOSS: Focal + Label Smoothing + Class Weights\n# ============================================================\nclass APTOSLoss(nn.Module):\n    def __init__(self, weights, gamma=2.0, smoothing=0.05):\n        super().__init__()\n        self.weights   = weights\n        self.gamma     = gamma\n        self.smoothing = smoothing\n        self.n         = Config.NUM_CLASSES\n\n    def forward(self, inputs, targets):\n        w   = self.weights.to(inputs.device)\n        eps = self.smoothing\n\n        # Smooth labels\n        with torch.no_grad():\n            smooth = torch.full_like(inputs, eps / (self.n - 1))\n            smooth.scatter_(1, targets.unsqueeze(1), 1.0 - eps)\n\n        log_p    = F.log_softmax(inputs, dim=1)\n        ce       = -(smooth * log_p).sum(dim=1)\n        pt       = torch.exp(-ce)\n        focal    = ((1 - pt) ** self.gamma) * ce\n        sw       = w[targets]\n        return (focal * sw).mean()\n\n# ============================================================\n# 6. MIXUP\n# ============================================================\ndef mixup(x, y, alpha=0.4):\n    lam   = np.random.beta(alpha, alpha)\n    idx   = torch.randperm(x.size(0)).to(x.device)\n    mixed = lam * x + (1 - lam) * x[idx]\n    return mixed, y, y[idx], lam\n\n# ============================================================\n# 7. TRAIN ONE EPOCH\n# ============================================================\ndef train_epoch(model, loader, criterion, optimizer, scheduler, device):\n    model.train()\n    total_loss = 0.0\n    all_preds, all_targets = [], []\n\n    pbar = tqdm(loader, desc=\"  Train\", leave=False)\n    for imgs, labels in pbar:\n        imgs, labels = imgs.to(device), labels.to(device)\n        mixed, ya, yb, lam = mixup(imgs, labels)\n\n        optimizer.zero_grad()\n        out  = model(mixed)\n        loss = lam * criterion(out, ya) + (1 - lam) * criterion(out, yb)\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        optimizer.step()\n        scheduler.step()\n\n        total_loss += loss.item()\n        preds = out.argmax(dim=1)\n        dominant = ya if lam >= 0.5 else yb\n        all_preds.extend(preds.cpu().numpy())\n        all_targets.extend(dominant.cpu().numpy())\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n\n    epoch_loss  = total_loss / len(loader)\n    epoch_acc   = accuracy_score(all_targets, all_preds)\n    epoch_kappa = cohen_kappa_score(all_targets, all_preds,\n                                    weights='quadratic')\n    return epoch_loss, epoch_acc, epoch_kappa\n\n# ============================================================\n# 8. VALIDATE / TEST\n# ============================================================\ndef evaluate(model, loader, criterion, device, split_name=\"Val\"):\n    model.eval()\n    total_loss = 0.0\n    all_preds, all_targets, all_probs = [], [], []\n    val_crit = nn.CrossEntropyLoss()\n\n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc=f\"  {split_name}\", leave=False):\n            imgs, labels = imgs.to(device), labels.to(device)\n            out   = model(imgs)\n            loss  = val_crit(out, labels)\n            total_loss += loss.item()\n            probs = F.softmax(out, dim=1)\n            preds = probs.argmax(dim=1)\n            all_preds.extend(preds.cpu().numpy())\n            all_targets.extend(labels.cpu().numpy())\n            all_probs.extend(probs.cpu().numpy())\n\n    all_probs   = np.array(all_probs)\n    acc         = accuracy_score(all_targets, all_preds)\n    kappa       = cohen_kappa_score(all_targets, all_preds, weights='quadratic')\n    epoch_loss  = total_loss / len(loader)\n\n    # One-vs-Rest AUC for 5-class\n    from sklearn.preprocessing import label_binarize\n    y_bin   = label_binarize(all_targets, classes=list(range(Config.NUM_CLASSES)))\n    try:\n        macro_auc = roc_auc_score(y_bin, all_probs,\n                                  multi_class='ovr', average='macro')\n    except Exception:\n        macro_auc = 0.0\n\n    return epoch_loss, acc, kappa, macro_auc, \\\n           np.array(all_targets), np.array(all_preds), all_probs\n\n# ============================================================\n# 9. TTA INFERENCE\n# ============================================================\ndef tta_predict(model, loader, device, n_tta=8):\n    model.eval()\n    all_targets = []\n    all_probs   = []\n\n    tta_fns = [\n        lambda x: x,\n        lambda x: torch.flip(x, [-1]),\n        lambda x: torch.flip(x, [-2]),\n        lambda x: torch.rot90(x, 1, [-2, -1]),\n        lambda x: torch.rot90(x, 2, [-2, -1]),\n        lambda x: torch.rot90(x, 3, [-2, -1]),\n        lambda x: torch.flip(torch.rot90(x, 1, [-2,-1]), [-1]),\n        lambda x: torch.flip(torch.rot90(x, 1, [-2,-1]), [-2]),\n    ]\n\n    with torch.no_grad():\n        for imgs, labels in tqdm(loader, desc=\"  TTA Inference\", leave=False):\n            imgs = imgs.to(device)\n            batch_probs = []\n            for fn in tta_fns[:n_tta]:\n                aug  = fn(imgs)\n                out  = model(aug)\n                prob = F.softmax(out, dim=1)\n                batch_probs.append(prob.cpu().numpy())\n            avg = np.mean(batch_probs, axis=0)\n            all_probs.extend(avg)\n            all_targets.extend(labels.numpy())\n\n    return np.array(all_targets), np.array(all_probs)\n\n# ============================================================\n# 10. PLOTTING FUNCTIONS\n# ============================================================\ndef plot_training_curves(history, save_path='training_curves.png'):\n    epochs = range(1, len(history['train_loss']) + 1)\n    fig    = plt.figure(figsize=(20, 12))\n    fig.suptitle('APTOS 2019 — Training Dashboard (Swin-Large)',\n                 fontsize=16, fontweight='bold', y=0.98)\n\n    gs = gridspec.GridSpec(2, 3, figure=fig, hspace=0.35, wspace=0.3)\n\n    # --- Plot 1: Loss ---\n    ax1 = fig.add_subplot(gs[0, 0])\n    ax1.plot(epochs, history['train_loss'], 'b-o', ms=3,\n             label='Train Loss', lw=2)\n    ax1.plot(epochs, history['val_loss'],   'r-o', ms=3,\n             label='Val Loss',   lw=2)\n    ax1.set_title('Loss Curve', fontweight='bold')\n    ax1.set_xlabel('Epoch'); ax1.set_ylabel('Loss')\n    ax1.legend(); ax1.grid(True, alpha=0.4)\n\n    # --- Plot 2: Accuracy ---\n    ax2 = fig.add_subplot(gs[0, 1])\n    ax2.plot(epochs, [a*100 for a in history['train_acc']], 'b-o',\n             ms=3, label='Train Acc', lw=2)\n    ax2.plot(epochs, [a*100 for a in history['val_acc']],   'g-o',\n             ms=3, label='Val Acc',   lw=2)\n    ax2.plot(epochs, [a*100 for a in history['test_acc']],  'm-o',\n             ms=3, label='Test Acc',  lw=2)\n    ax2.set_title('Accuracy Curve', fontweight='bold')\n    ax2.set_xlabel('Epoch'); ax2.set_ylabel('Accuracy (%)')\n    ax2.legend(); ax2.grid(True, alpha=0.4)\n\n    # --- Plot 3: Kappa ---\n    ax3 = fig.add_subplot(gs[0, 2])\n    ax3.plot(epochs, history['train_kappa'], 'b-o', ms=3,\n             label='Train κ', lw=2)\n    ax3.plot(epochs, history['val_kappa'],   'g-o', ms=3,\n             label='Val κ',   lw=2)\n    ax3.plot(epochs, history['test_kappa'],  'm-o', ms=3,\n             label='Test κ',  lw=2)\n    ax3.set_title('Quadratic Weighted Kappa', fontweight='bold')\n    ax3.set_xlabel('Epoch'); ax3.set_ylabel('Kappa Score')\n    ax3.legend(); ax3.grid(True, alpha=0.4)\n\n    # --- Plot 4: AUC ---\n    ax4 = fig.add_subplot(gs[1, 0])\n    ax4.plot(epochs, history['val_auc'],  'g-o', ms=3,\n             label='Val AUC',  lw=2)\n    ax4.plot(epochs, history['test_auc'], 'm-o', ms=3,\n             label='Test AUC', lw=2)\n    ax4.set_title('Macro AUC (OvR)', fontweight='bold')\n    ax4.set_xlabel('Epoch'); ax4.set_ylabel('AUC')\n    ax4.legend(); ax4.grid(True, alpha=0.4)\n\n    # --- Plot 5: Overfitting Gap ---\n    ax5 = fig.add_subplot(gs[1, 1])\n    gap = [tr - vl for tr, vl in\n           zip(history['train_acc'], history['val_acc'])]\n    colors = ['red' if g > 0.08 else 'orange' if g > 0.04 else 'green'\n              for g in gap]\n    ax5.bar(epochs, gap, color=colors, alpha=0.75)\n    ax5.axhline(0.08, color='red',    ls='--', lw=1.5,\n                label='Overfit threshold (8%)')\n    ax5.axhline(0.04, color='orange', ls='--', lw=1.5,\n                label='Warning threshold (4%)')\n    ax5.set_title('Overfitting Gap (Train - Val Acc)',\n                  fontweight='bold')\n    ax5.set_xlabel('Epoch'); ax5.set_ylabel('Gap')\n    ax5.legend(fontsize=8); ax5.grid(True, alpha=0.4)\n\n    # --- Plot 6: LR Schedule ---\n    ax6 = fig.add_subplot(gs[1, 2])\n    ax6.plot(history['lr'], 'darkorange', lw=2)\n    ax6.set_title('Learning Rate Schedule', fontweight='bold')\n    ax6.set_xlabel('Step'); ax6.set_ylabel('LR')\n    ax6.grid(True, alpha=0.4)\n    ax6.set_yscale('log')\n\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.show()\n    print(f\"Saved → {save_path}\")\n\n\ndef plot_roc_curves(targets, probs, save_path='roc_curves.png'):\n    from sklearn.preprocessing import label_binarize\n    y_bin = label_binarize(targets,\n                           classes=list(range(Config.NUM_CLASSES)))\n    colors = ['#2ecc71','#f1c40f','#e67e22','#e74c3c','#8e44ad']\n\n    fig, axes = plt.subplots(1, 2, figsize=(16, 6))\n    fig.suptitle('APTOS 2019 — ROC Curves (TTA Predictions)',\n                 fontsize=14, fontweight='bold')\n\n    # Per-class ROC\n    ax = axes[0]\n    all_fpr, all_tpr, roc_aucs = {}, {}, {}\n    for i in range(Config.NUM_CLASSES):\n        fpr, tpr, _ = roc_curve(y_bin[:, i], probs[:, i])\n        roc_aucs[i] = auc(fpr, tpr)\n        all_fpr[i]  = fpr\n        all_tpr[i]  = tpr\n        ax.plot(fpr, tpr, color=colors[i], lw=2,\n                label=f\"{Config.GRADE_NAMES[i]} (AUC={roc_aucs[i]:.3f})\")\n\n    ax.plot([0,1],[0,1],'k--', lw=1.5)\n    ax.set_xlabel('False Positive Rate', fontsize=11)\n    ax.set_ylabel('True Positive Rate', fontsize=11)\n    ax.set_title('Per-Class ROC', fontweight='bold')\n    ax.legend(fontsize=9); ax.grid(True, alpha=0.4)\n\n    # Macro-average ROC\n    ax2   = axes[1]\n    macro = roc_auc_score(y_bin, probs, multi_class='ovr', average='macro')\n    # Compute macro curve via interpolation\n    all_fpr_vals = np.unique(np.concatenate([all_fpr[i]\n                             for i in range(Config.NUM_CLASSES)]))\n    mean_tpr     = np.zeros_like(all_fpr_vals)\n    for i in range(Config.NUM_CLASSES):\n        mean_tpr += np.interp(all_fpr_vals, all_fpr[i], all_tpr[i])\n    mean_tpr /= Config.NUM_CLASSES\n\n    ax2.plot(all_fpr_vals, mean_tpr, 'b-', lw=3,\n             label=f'Macro-avg (AUC={macro:.3f})')\n    for i in range(Config.NUM_CLASSES):\n        ax2.plot(all_fpr[i], all_tpr[i], color=colors[i],\n                 lw=1, alpha=0.4,\n                 label=f\"{Config.GRADE_NAMES[i]}={roc_aucs[i]:.3f}\")\n    ax2.plot([0,1],[0,1],'k--', lw=1.5)\n    ax2.set_xlabel('False Positive Rate', fontsize=11)\n    ax2.set_ylabel('True Positive Rate', fontsize=11)\n    ax2.set_title('Macro-Average ROC', fontweight='bold')\n    ax2.legend(fontsize=9); ax2.grid(True, alpha=0.4)\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.show()\n    print(f\"Saved → {save_path}\")\n    return roc_aucs, macro\n\n\ndef plot_confusion_matrix(targets, preds, save_path='confusion_matrix.png'):\n    cm   = confusion_matrix(targets, preds)\n    cm_n = cm.astype(float) / cm.sum(axis=1, keepdims=True)\n\n    fig, axes = plt.subplots(1, 2, figsize=(14, 5))\n    fig.suptitle('Confusion Matrix — TTA Predictions',\n                 fontsize=13, fontweight='bold')\n\n    for ax, data, title, fmt in zip(\n        axes,\n        [cm, cm_n],\n        ['Raw Counts', 'Normalized (Row %)'],\n        ['d', '.2%']\n    ):\n        sns.heatmap(data, annot=True, fmt=fmt, cmap='Blues',\n                    xticklabels=Config.GRADE_NAMES,\n                    yticklabels=Config.GRADE_NAMES,\n                    ax=ax, linewidths=0.5)\n        ax.set_xlabel('Predicted', fontsize=11)\n        ax.set_ylabel('True', fontsize=11)\n        ax.set_title(title, fontweight='bold')\n\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.show()\n    print(f\"Saved → {save_path}\")\n\n\ndef print_epoch_table(epoch, total_epochs, metrics):\n    \"\"\"Print a formatted per-epoch results table.\"\"\"\n    if epoch == 1:\n        print(\"\\n\" + \"=\"*100)\n        print(f\"{'Ep':>4} | {'Tr Loss':>8} | {'Tr Acc':>7} | \"\n              f\"{'Tr κ':>7} | {'Va Loss':>8} | {'Va Acc':>7} | \"\n              f\"{'Va κ':>7} | {'Te Acc':>7} | {'Te κ':>7} | \"\n              f\"{'Va AUC':>7} | {'Te AUC':>7}\")\n        print(\"-\"*100)\n\n    m = metrics\n    print(f\"{epoch:>4}/{total_epochs} | \"\n          f\"{m['train_loss']:>8.4f} | \"\n          f\"{m['train_acc']*100:>6.2f}% | \"\n          f\"{m['train_kappa']:>7.4f} | \"\n          f\"{m['val_loss']:>8.4f} | \"\n          f\"{m['val_acc']*100:>6.2f}% | \"\n          f\"{m['val_kappa']:>7.4f} | \"\n          f\"{m['test_acc']*100:>6.2f}% | \"\n          f\"{m['test_kappa']:>7.4f} | \"\n          f\"{m['val_auc']:>7.4f} | \"\n          f\"{m['test_auc']:>7.4f}\"\n          f\"{'  ← BEST' if m.get('is_best') else ''}\")\n\n# ============================================================\n# 11. MAIN\n# ============================================================\nif __name__ == '__main__':\n\n    # ---- Load CSV ----\n    df = pd.read_csv(Config.CSV_PATH)\n    print(f\"\\nLoaded CSV: {df.shape} | \"\n          f\"Columns: {df.columns.tolist()}\")\n\n    # APTOS columns are fixed\n    df['path'] = df['id_code'].apply(\n        lambda x: os.path.join(Config.IMG_DIR, f\"{x}.png\"))\n    df = df[df['path'].apply(os.path.exists)].reset_index(drop=True)\n    print(f\"Valid images found: {len(df)}\")\n\n    # Grade distribution\n    print(\"\\nGrade Distribution:\")\n    for g, cnt in df['diagnosis'].value_counts().sort_index().items():\n        print(f\"  Grade {g} ({Config.GRADE_NAMES[g]}): {cnt} \"\n              f\"({cnt/len(df)*100:.1f}%)\")\n\n    # ---- Splits: 70 / 15 / 15 ----\n    train_df, temp_df = train_test_split(\n        df, test_size=0.30, random_state=Config.SEED,\n        stratify=df['diagnosis'])\n    val_df, test_df   = train_test_split(\n        temp_df, test_size=0.50, random_state=Config.SEED,\n        stratify=temp_df['diagnosis'])\n\n    print(f\"\\nTrain: {len(train_df)} | \"\n          f\"Val: {len(val_df)} | Test: {len(test_df)}\")\n\n    # ---- Weighted Sampler for Grade 3 ----\n    labels_train  = train_df['diagnosis'].values\n    class_counts  = np.bincount(labels_train)\n    class_w       = 1.0 / class_counts\n    sample_w      = class_w[labels_train]\n    sampler       = WeightedRandomSampler(\n        weights=torch.FloatTensor(sample_w),\n        num_samples=len(labels_train),\n        replacement=True\n    )\n\n    # ---- Datasets ----\n    train_ds = APTOSDataset(train_df['id_code'].values,\n                            labels_train, Config.IMG_DIR, is_train=True)\n    val_ds   = APTOSDataset(val_df['id_code'].values,\n                            val_df['diagnosis'].values,\n                            Config.IMG_DIR, is_train=False)\n    test_ds  = APTOSDataset(test_df['id_code'].values,\n                            test_df['diagnosis'].values,\n                            Config.IMG_DIR, is_train=False)\n\n    train_loader = DataLoader(train_ds, batch_size=Config.BATCH_SIZE,\n                              sampler=sampler, num_workers=2,\n                              pin_memory=True)\n    val_loader   = DataLoader(val_ds,   batch_size=Config.BATCH_SIZE,\n                              shuffle=False, num_workers=2,\n                              pin_memory=True)\n    test_loader  = DataLoader(test_ds,  batch_size=Config.BATCH_SIZE,\n                              shuffle=False, num_workers=2,\n                              pin_memory=True)\n\n    # ---- Model ----\n    model     = APTOS_SwinLarge(num_classes=Config.NUM_CLASSES)\n    model     = model.to(Config.DEVICE)\n    criterion = APTOSLoss(Config.CLASS_WEIGHTS, gamma=2.0, smoothing=0.05)\n\n    # ---- Differential LR optimizer ----\n    optimizer = optim.AdamW([\n        {'params': model.backbone.parameters(),\n         'lr': 5e-6,  'weight_decay': 0.05},\n        {'params': list(model.se.parameters()) +\n                   list(model.head.parameters()),\n         'lr': 1e-4,  'weight_decay': 0.01},\n    ])\n\n    total_steps = Config.EPOCHS * len(train_loader)\n    scheduler   = optim.lr_scheduler.OneCycleLR(\n        optimizer,\n        max_lr=[5e-6, 1e-4],\n        total_steps=total_steps,\n        pct_start=0.1,\n        anneal_strategy='cos'\n    )\n\n    # ---- History ----\n    history = {k: [] for k in [\n        'train_loss', 'train_acc', 'train_kappa',\n        'val_loss',   'val_acc',   'val_kappa',   'val_auc',\n        'test_loss',  'test_acc',  'test_kappa',  'test_auc',\n        'lr'\n    ]}\n\n    best_kappa      = 0.0\n    best_val_acc    = 0.0\n    best_roc_data   = None\n    best_epoch      = 0\n\n    print(f\"\\n{'='*60}\")\n    print(f\"  Starting Training — {Config.EPOCHS} Epochs\")\n    print(f\"{'='*60}\")\n\n    # ---- Training Loop ----\n    for epoch in range(1, Config.EPOCHS + 1):\n        ep_start = time.time()\n\n        # Train\n        tr_loss, tr_acc, tr_kappa = train_epoch(\n            model, train_loader, criterion, optimizer, scheduler,\n            Config.DEVICE)\n\n        # LR snapshot (last step LR)\n        cur_lr = optimizer.param_groups[-1]['lr']\n\n        # Validate\n        va_loss, va_acc, va_kappa, va_auc, \\\n            _, _, _ = evaluate(\n                model, val_loader, criterion, Config.DEVICE, \"Val\")\n\n        # Test (every epoch so you can see test curve)\n        te_loss, te_acc, te_kappa, te_auc, \\\n            te_targets, te_preds, te_probs = evaluate(\n                model, test_loader, criterion, Config.DEVICE, \"Test\")\n\n        # Record history\n        history['train_loss'].append(tr_loss)\n        history['train_acc'].append(tr_acc)\n        history['train_kappa'].append(tr_kappa)\n        history['val_loss'].append(va_loss)\n        history['val_acc'].append(va_acc)\n        history['val_kappa'].append(va_kappa)\n        history['val_auc'].append(va_auc)\n        history['test_loss'].append(te_loss)\n        history['test_acc'].append(te_acc)\n        history['test_kappa'].append(te_kappa)\n        history['test_auc'].append(te_auc)\n        history['lr'].append(cur_lr)\n\n        is_best = va_kappa > best_kappa\n        if is_best:\n            best_kappa    = va_kappa\n            best_val_acc  = va_acc\n            best_epoch    = epoch\n            torch.save(model.state_dict(), 'best_aptos_swinlarge.pth')\n            best_roc_data = (te_targets, np.array(te_probs))\n\n        ep_time = time.time() - ep_start\n\n        metrics = {\n            'train_loss': tr_loss, 'train_acc': tr_acc,\n            'train_kappa': tr_kappa,\n            'val_loss':   va_loss, 'val_acc':   va_acc,\n            'val_kappa':  va_kappa, 'val_auc':  va_auc,\n            'test_acc':   te_acc,  'test_kappa': te_kappa,\n            'test_auc':   te_auc,\n            'is_best':    is_best\n        }\n        print_epoch_table(epoch, Config.EPOCHS, metrics)\n        print(f\"         ↳ Time: {ep_time:.1f}s | \"\n              f\"LR: {cur_lr:.2e}\")\n\n    # ============================================================\n    # 12. FINAL EVALUATION WITH TTA\n    # ============================================================\n    print(f\"\\n{'='*60}\")\n    print(f\"  Loading Best Model (Epoch {best_epoch}) for TTA...\")\n    print(f\"{'='*60}\")\n\n    model.load_state_dict(torch.load('best_aptos_swinlarge.pth',\n                                     map_location=Config.DEVICE))\n\n    print(\"\\n--- TTA on Validation Set ---\")\n    val_targets, val_probs = tta_predict(model, val_loader,\n                                         Config.DEVICE, n_tta=8)\n    val_preds_tta  = val_probs.argmax(axis=1)\n    val_acc_tta    = accuracy_score(val_targets, val_preds_tta)\n    val_kappa_tta  = cohen_kappa_score(val_targets, val_preds_tta,\n                                       weights='quadratic')\n\n    print(\"\\n--- TTA on Test Set ---\")\n    tta_targets, tta_probs = tta_predict(model, test_loader,\n                                          Config.DEVICE, n_tta=8)\n    tta_preds   = tta_probs.argmax(axis=1)\n    tta_acc     = accuracy_score(tta_targets, tta_preds)\n    tta_kappa   = cohen_kappa_score(tta_targets, tta_preds,\n                                    weights='quadratic')\n\n    # ============================================================\n    # 13. TOTAL RUNTIME\n    # ============================================================\n    total_time      = time.time() - TOTAL_START\n    hrs, rem        = divmod(total_time, 3600)\n    mins, secs      = divmod(rem, 60)\n\n    # ============================================================\n    # 14. FINAL REPORT\n    # ============================================================\n    print(\"\\n\" + \"=\"*65)\n    print(\"                 FINAL RESULTS SUMMARY\")\n    print(\"=\"*65)\n    print(f\"  Best Epoch               : {best_epoch}/{Config.EPOCHS}\")\n    print(f\"  Best Val Kappa (no TTA)  : {best_kappa:.4f}\")\n    print(f\"  Best Val Acc   (no TTA)  : {best_val_acc*100:.2f}%\")\n    print(f\"\\n  --- With 8-Fold TTA ---\")\n    print(f\"  Val Accuracy             : {val_acc_tta*100:.2f}%\")\n    print(f\"  Val Kappa                : {val_kappa_tta:.4f}\")\n    print(f\"  Test Accuracy            : {tta_acc*100:.2f}%\")\n    print(f\"  Test Kappa               : {tta_kappa:.4f}\")\n    print(f\"\\n  Total Runtime            : \"\n          f\"{int(hrs)}h {int(mins)}m {int(secs)}s\")\n    print(\"=\"*65)\n\n    print(\"\\n--- Classification Report (Test, TTA) ---\")\n    print(classification_report(tta_targets, tta_preds,\n                                target_names=Config.GRADE_NAMES))\n\n    # ============================================================\n    # 15. PLOTS\n    # ============================================================\n    print(\"\\nGenerating plots...\")\n\n    # (A) Training Dashboard\n    plot_training_curves(history, 'training_curves.png')\n\n    # (B) ROC Curves\n    roc_per_class, macro_auc = plot_roc_curves(\n        tta_targets, tta_probs, 'roc_curves.png')\n\n    # (C) Confusion Matrix\n    plot_confusion_matrix(tta_targets, tta_preds, 'confusion_matrix.png')\n\n    # (D) Final AUC summary bar chart\n    fig, ax = plt.subplots(figsize=(9, 5))\n    names   = Config.GRADE_NAMES + ['Macro Avg']\n    aucs    = [roc_per_class[i] for i in range(Config.NUM_CLASSES)] \\\n              + [macro_auc]\n    colors  = ['#2ecc71','#f1c40f','#e67e22','#e74c3c','#8e44ad','#2c3e50']\n    bars    = ax.bar(names, aucs, color=colors, edgecolor='black',\n                     linewidth=0.8)\n    for bar, val in zip(bars, aucs):\n        ax.text(bar.get_x() + bar.get_width()/2,\n                bar.get_height() + 0.005,\n                f'{val:.3f}', ha='center', va='bottom',\n                fontweight='bold', fontsize=10)\n    ax.set_ylim(0.7, 1.02)\n    ax.axhline(0.95, color='red', ls='--', lw=1.5,\n               label='Target (0.95)')\n    ax.set_title('Per-Class & Macro AUC — TTA Predictions',\n                 fontweight='bold', fontsize=13)\n    ax.set_ylabel('AUC Score')\n    ax.legend(); ax.grid(axis='y', alpha=0.4)\n    plt.tight_layout()\n    plt.savefig('auc_summary.png', dpi=150, bbox_inches='tight')\n    plt.show()\n    print(\"Saved → auc_summary.png\")\n\n    print(\"\\n All done.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-30T17:36:39.531949Z","iopub.execute_input":"2026-03-30T17:36:39.53274Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# MedMamba for Diabetic Retinopathy — APTOS 2019 Dataset\n# Kaggle Competition: aptos2019-blindness-detection\n# Dataset path: /kaggle/input/aptos2019-blindness-detection/\n#\n# APTOS layout:\n#   train.csv        → columns: id_code, diagnosis (0-4)\n#   train_images/    → <id_code>.png\n#   test.csv         → columns: id_code  (no labels — competition test set)\n#   test_images/     → <id_code>.png\n#\n# NOTE: APTOS has only 3,662 labelled images (vs 88k in EyePACS).\n# Key adaptations for small dataset:\n#   • Stronger augmentation\n#   • MixUp regularisation\n#   • Test-Time Augmentation (TTA) at inference\n#   • 5-fold cross-validation option\n#   • Smaller batch size (16) to give more gradient update steps\n# =============================================================================\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 1 — Install dependencies\n# ─────────────────────────────────────────────────────────────────────────────\nimport subprocess, sys, os\n\ndef run(cmd):\n    r = subprocess.run(cmd, shell=True, capture_output=True, text=True)\n    tail = (r.stdout + r.stderr).strip()[-300:]\n    if r.returncode != 0:\n        print(f\"[WARN] {tail}\")\n    else:\n        print(\"OK\")\n    return r.returncode == 0\n\nif not os.path.exists(\"/kaggle/working/MedMamba\"):\n    run(\"git clone --depth 1 https://github.com/YubiaoYue/MedMamba.git \"\n        \"/kaggle/working/MedMamba\")\nelse:\n    print(\"MedMamba already cloned — skipping.\")\n\nimport torch\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}  → installing mamba-ssm …\")\n    ok1 = run(\"pip install -q 'causal-conv1d>=1.2.0.post2' --no-build-isolation\")\n    ok2 = run(\"pip install -q 'mamba-ssm>=1.2.0'           --no-build-isolation\")\n    MAMBA_AVAILABLE = ok1 and ok2\nelse:\n    print(\"No GPU — mamba-ssm skipped, Swin-Tiny fallback will be used.\")\n    MAMBA_AVAILABLE = False\n\nrun(\"pip install -q timm einops scikit-learn matplotlib seaborn \"\n    \"albumentations opencv-python-headless\")\n\nsys.path.insert(0, \"/kaggle/working/MedMamba\")\nprint(f\"\\nSetup complete.  MAMBA_AVAILABLE={MAMBA_AVAILABLE}\")\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 2 — Imports & config (With Deep-Search Path Detection)\n# ─────────────────────────────────────────────────────────────────────────────\nimport random, warnings, os\nimport numpy as np\nimport pandas as pd\nimport cv2\nfrom pathlib import Path\nfrom collections import Counter\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.cuda.amp import GradScaler, autocast\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nfrom sklearn.metrics import (\n    accuracy_score, f1_score, roc_auc_score,\n    recall_score, precision_score,\n    confusion_matrix, classification_report,\n    roc_curve, auc,\n    cohen_kappa_score,          # competition metric for APTOS\n)\nfrom sklearn.preprocessing import label_binarize\nfrom sklearn.model_selection import train_test_split, StratifiedKFold\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport timm\n\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 42\ndef seed_everything(s):\n    random.seed(s); np.random.seed(s)\n    torch.manual_seed(s); torch.cuda.manual_seed_all(s)\n    torch.backends.cudnn.deterministic = True\nseed_everything(SEED)\n\n# ── Deep Search for APTOS Path ────────────────────────────────────────────────\n_BASE = None\n\nprint(\"Scanning /kaggle/input/ for the APTOS dataset...\")\n# Deep search through all directories for train.csv\nfor path in Path('/kaggle/input').rglob('train.csv'):\n    parent_dir = path.parent\n    # Verify it has the train_images folder next to it\n    if (parent_dir / 'train_images').exists():\n        _BASE = str(parent_dir)\n        break\n\nif _BASE is None:\n    print(\"\\n[DIAGNOSTIC] Here is exactly what Python sees in your /kaggle/input/ folder:\")\n    for dirname, _, filenames in os.walk('/kaggle/input'):\n        print(f\"Directory: {dirname}\")\n        for filename in filenames[:3]: # print first 3 files per dir\n            print(f\"  - {filename}\")\n        if len(filenames) > 3:\n            print(f\"  ... and {len(filenames)-3} more files\")\n            \n    raise FileNotFoundError(\n        \"\\n✗ Could not find the APTOS dataset! \\n\"\n        \"FIX 1: Go to https://www.kaggle.com/c/aptos2019-blindness-detection/rules and accept the rules.\\n\"\n        \"FIX 2: Restart your Kaggle session to force the data to mount.\"\n    )\n\nprint(f\"✓ Dynamically resolved dataset base path: {_BASE}\")\n\n_TRAIN_CSV = f\"{_BASE}/train.csv\"\n_TRAIN_IMG = f\"{_BASE}/train_images\"\n_TEST_CSV  = f\"{_BASE}/test.csv\"\n_TEST_IMG  = f\"{_BASE}/test_images\"\n\n# ── Configuration Dictionary ──────────────────────────────────────────────────\nCFG = dict(\n    TRAIN_CSV   = _TRAIN_CSV,\n    TRAIN_IMG   = _TRAIN_IMG,\n    TEST_CSV    = _TEST_CSV,\n    TEST_IMG    = _TEST_IMG,\n    IMG_SIZE    = 224,\n    NUM_CLASSES = 5,\n    BATCH_SIZE  = 16,\n    EPOCHS      = 60,\n    LR          = 5e-5,          \n    WEIGHT_DECAY= 1e-3,          \n    MODEL       = \"MedMamba-T\",\n    AMP         = torch.cuda.is_available(),\n    NUM_WORKERS = 2,\n    PATIENCE    = 15,            \n    GRAD_CLIP   = 1.0,\n    MIXUP_ALPHA = 0.4,           \n    TTA_STEPS   = 5,             \n    SAVE_PATH   = \"/kaggle/working/best_aptos_model.pth\",\n    CLASS_NAMES = [\"No DR\",\"Mild\",\"Moderate\",\"Severe\",\"Proliferative\"],\n)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Device : {DEVICE}\")\n\n# Verify paths explicitly\nfor p in [CFG[\"TRAIN_CSV\"], CFG[\"TRAIN_IMG\"]]:\n    assert os.path.exists(p), f\"Still not found: {p}\"\nprint(\"Paths verified successfully ✓\")\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 3 — Load data & class-balance analysis\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef load_aptos(csv_path: str, img_dir: str) -> pd.DataFrame:\n    img_dir = Path(img_dir)\n    df = pd.read_csv(csv_path)\n    df.columns = df.columns.str.strip().str.lower()\n    print(f\"CSV shape : {df.shape}\")\n    print(f\"Columns   : {df.columns.tolist()}\")\n\n    # APTOS standard columns: id_code, diagnosis\n    img_col   = next(c for c in df.columns if any(k in c for k in\n                     (\"id_code\",\"image\",\"name\",\"file\",\"id\")))\n    label_col = next(c for c in df.columns if any(k in c for k in\n                     (\"diagnosis\",\"level\",\"label\",\"grade\")))\n    print(f\"  image col → '{img_col}'  |  label col → '{label_col}'\")\n\n    def resolve(name):\n        name = str(name).strip()\n        for ext in (\".png\", \".jpeg\", \".jpg\"):\n            p = img_dir / (name + ext)\n            if p.exists():\n                return str(p)\n        return None\n\n    df[\"path\"]  = df[img_col].apply(resolve)\n    df[\"label\"] = df[label_col].astype(int)\n\n    missing = df[\"path\"].isna().sum()\n    if missing:\n        print(f\"[WARN] {missing} images missing from disk — dropped.\")\n    df = df.dropna(subset=[\"path\"])[[\"path\",\"label\"]].reset_index(drop=True)\n    print(f\"Usable samples : {len(df):,}\")\n    return df\n\n\nfull_df = load_aptos(CFG[\"TRAIN_CSV\"], CFG[\"TRAIN_IMG\"])\n\n# Stratified 70 / 15 / 15 (APTOS is small — keep more for training)\ntrain_df, tmp = train_test_split(full_df, test_size=0.30,\n                                 random_state=SEED, stratify=full_df[\"label\"])\nval_df, test_df = train_test_split(tmp, test_size=0.50,\n                                   random_state=SEED, stratify=tmp[\"label\"])\ntrain_df = train_df.reset_index(drop=True)\nval_df   = val_df.reset_index(drop=True)\ntest_df  = test_df.reset_index(drop=True)\n\nprint(f\"\\nSplit → Train: {len(train_df)}  |  Val: {len(val_df)}  |  Test: {len(test_df)}\")\n\n# ── Class distribution ────────────────────────────────────────────────────────\nlabel_counts = Counter(train_df[\"label\"].values)\ncounts = [label_counts.get(k, 0) for k in range(CFG[\"NUM_CLASSES\"])]\n\nprint(\"\\nClass distribution (train):\")\nfor k in range(CFG[\"NUM_CLASSES\"]):\n    pct = 100 * counts[k] / len(train_df)\n    bar = \"█\" * int(pct / 2)\n    print(f\"  Grade {k} | {bar:<40} {counts[k]:>5}  ({pct:4.1f}%)\")\n\nratio = max(counts) / (min(c for c in counts if c > 0))\nprint(f\"\\nImbalance ratio : {ratio:.1f}x\")\nprint(f\"Total labelled  : {len(full_df)} images  (much smaller than EyePACS)\")\nprint(\"→ Using: WeightedSampler + class weights + MixUp + strong augmentation\")\n\nfig, ax = plt.subplots(figsize=(8, 4))\ncolors = [\"#4C72B0\",\"#DD8452\",\"#55A868\",\"#C44E52\",\"#8172B2\"]\nbars = ax.bar(CFG[\"CLASS_NAMES\"], counts, color=colors)\nax.set_title(\"APTOS 2019 — Class Distribution (train split)\", fontsize=12)\nax.set_ylabel(\"Number of images\")\nax.tick_params(axis=\"x\", rotation=15)\nfor bar, c in zip(bars, counts):\n    ax.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 2,\n            str(c), ha=\"center\", va=\"bottom\", fontsize=10)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/aptos_class_distribution.png\", dpi=120)\nplt.show()\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 4 — Augmentation pipeline + Dataset + DataLoaders\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef get_transforms(split: str, img_size: int):\n    mean, std = [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]\n    size = (img_size, img_size)\n\n    if split == \"train\":\n        return A.Compose([\n            A.RandomResizedCrop(size=size, scale=(0.7, 1.0)),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.2),\n            A.RandomRotate90(p=0.5),\n            A.Rotate(limit=20, p=0.6),\n            # Fundus-specific: CLAHE boosts microaneurysm visibility\n            A.CLAHE(clip_limit=3.0, tile_grid_size=(8, 8), p=0.6),\n            A.OneOf([\n                A.ColorJitter(brightness=0.3, contrast=0.3,\n                              saturation=0.3, hue=0.07, p=1.0),\n                A.HueSaturationValue(hue_shift_limit=15,\n                                     sat_shift_limit=30,\n                                     val_shift_limit=25, p=1.0),\n            ], p=0.8),\n            A.OneOf([\n                A.GaussianBlur(blur_limit=(3, 7), p=1.0),\n                A.GaussNoise(var_limit=(10, 60), p=1.0),\n                A.MotionBlur(blur_limit=5, p=1.0),\n            ], p=0.4),\n            A.RandomBrightnessContrast(brightness_limit=0.2,\n                                       contrast_limit=0.2, p=0.5),\n            A.CoarseDropout(\n                num_holes_range=(1, 8),\n                hole_height_range=(1, img_size // 16),\n                hole_width_range=(1, img_size // 16),\n                fill_value=0, p=0.3\n            ),\n            A.Normalize(mean=mean, std=std),\n            ToTensorV2(),\n        ])\n\n    # Val / Test / TTA-base\n    return A.Compose([\n        A.Resize(height=img_size, width=img_size),\n        A.Normalize(mean=mean, std=std),\n        ToTensorV2(),\n    ])\n\n\ndef get_tta_transform(img_size: int):\n    \"\"\"Light random augmentation for test-time augmentation passes.\"\"\"\n    mean, std = [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]\n    return A.Compose([\n        A.Resize(height=img_size, width=img_size),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.2),\n        A.RandomRotate90(p=0.5),\n        A.CLAHE(clip_limit=2.0, tile_grid_size=(8, 8), p=0.4),\n        A.Normalize(mean=mean, std=std),\n        ToTensorV2(),\n    ])\n\n\nclass APTOSDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df        = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = cv2.imread(row[\"path\"])\n        if img is None:\n            img = np.zeros((CFG[\"IMG_SIZE\"], CFG[\"IMG_SIZE\"], 3), dtype=np.uint8)\n        else:\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            img = self.transform(image=img)[\"image\"]\n        return img, int(row[\"label\"])\n\n\ndef make_weighted_sampler(df):\n    labels    = df[\"label\"].values\n    class_cnt = np.bincount(labels, minlength=CFG[\"NUM_CLASSES\"]).astype(float)\n    class_wt  = 1.0 / (class_cnt + 1e-6)\n    return WeightedRandomSampler(class_wt[labels],\n                                 num_samples=len(labels), replacement=True)\n\n\ntrain_ds = APTOSDataset(train_df, get_transforms(\"train\", CFG[\"IMG_SIZE\"]))\nval_ds   = APTOSDataset(val_df,   get_transforms(\"val\",   CFG[\"IMG_SIZE\"]))\ntest_ds  = APTOSDataset(test_df,  get_transforms(\"test\",  CFG[\"IMG_SIZE\"]))\n\ntrain_loader = DataLoader(train_ds, batch_size=CFG[\"BATCH_SIZE\"],\n                          sampler=make_weighted_sampler(train_df),\n                          num_workers=CFG[\"NUM_WORKERS\"], pin_memory=True)\nval_loader   = DataLoader(val_ds,   batch_size=CFG[\"BATCH_SIZE\"], shuffle=False,\n                          num_workers=CFG[\"NUM_WORKERS\"], pin_memory=True)\ntest_loader  = DataLoader(test_ds,  batch_size=CFG[\"BATCH_SIZE\"], shuffle=False,\n                          num_workers=CFG[\"NUM_WORKERS\"], pin_memory=True)\n\nprint(f\"DataLoaders ready — \"\n      f\"train: {len(train_loader)} batches | \"\n      f\"val: {len(val_loader)} | test: {len(test_loader)}\")\n\n# ── Sanity-check: show 8 augmented samples ────────────────────────────────────\nimgs_s, labels_s = next(iter(train_loader))\nfig, axes = plt.subplots(1, 8, figsize=(18, 3))\nmean_t = torch.tensor([0.485, 0.456, 0.406]).view(3,1,1)\nstd_t  = torch.tensor([0.229, 0.224, 0.225]).view(3,1,1)\nfor ax, img, lbl in zip(axes, imgs_s[:8], labels_s[:8]):\n    show = (img * std_t + mean_t).permute(1,2,0).clamp(0,1).numpy()\n    ax.imshow(show); ax.axis(\"off\")\n    ax.set_title(CFG[\"CLASS_NAMES\"][lbl.item()], fontsize=7)\nplt.suptitle(\"APTOS — Augmented training samples\", fontsize=11)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/aptos_samples.png\", dpi=100)\nplt.show()\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 5 — Build model (Strictly MedMamba - Fixed Import)\n# ─────────────────────────────────────────────────────────────────────────────\n\n# 1. Hard Check: Stop the notebook immediately if Mamba is missing\nif not MAMBA_AVAILABLE:\n    raise RuntimeError(\n        \"\\n❌ FATAL ERROR: mamba-ssm did not install correctly in Cell 1! \\n\"\n        \"The script has been stopped because the Swin Transformer fallback has been disabled.\\n\"\n        \"Please factory-reset your Kaggle session and re-run Cell 1.\"\n    )\n\n# 2. Correct Import from the MedMamba repository\nfrom MedMamba import VSSM\n\n# 3. Model Factory based on the original paper's architecture (Table 1)\ndef create_medmamba(variant, num_classes):\n    if variant == \"MedMamba-T\":\n        # Tiny: 96 base channels, depths = [2, 2, 4, 2]\n        return VSSM(patch_size=4, in_chans=3, num_classes=num_classes, \n                    depths=[2, 2, 4, 2], embed_dim=96)\n    elif variant == \"MedMamba-S\":\n        # Small: 96 base channels, depths = [2, 2, 8, 2]\n        return VSSM(patch_size=4, in_chans=3, num_classes=num_classes, \n                    depths=[2, 2, 8, 2], embed_dim=96)\n    elif variant == \"MedMamba-B\":\n        # Base: 128 base channels, depths = [2, 2, 12, 2]\n        return VSSM(patch_size=4, in_chans=3, num_classes=num_classes, \n                    depths=[2, 2, 12, 2], embed_dim=128)\n    else:\n        raise ValueError(f\"Unknown model variant: {variant}\")\n\nMODEL_NAME = CFG[\"MODEL\"]\nprint(f\"Initializing {MODEL_NAME} architecture...\")\n\n# Instantiate the Model\nmodel = create_medmamba(MODEL_NAME, CFG[\"NUM_CLASSES\"])\n\n# ─────────────────────────────────────────────────────────────────────────────\n# OPTIONAL BUT HIGHLY RECOMMENDED: Load Pre-Trained Weights here\n# If you upload the MedMamba weights to Kaggle, uncomment the lines below:\n# ─────────────────────────────────────────────────────────────────────────────\n# PRETRAINED_PATH = \"/kaggle/input/your-medmamba-weights-dataset/vssm_tiny.pth\"\n# if os.path.exists(PRETRAINED_PATH):\n#     print(\"Loading pre-trained ImageNet weights...\")\n#     checkpoint = torch.load(PRETRAINED_PATH, map_location=\"cpu\")\n#     # Note: strict=False ignores the classifier head size difference (1000 vs 5 classes)\n#     model.load_state_dict(checkpoint['model'], strict=False) \n# else:\n#     print(\"⚠️ WARNING: Training from scratch! No pre-trained weights found.\")\n# ─────────────────────────────────────────────────────────────────────────────\n\n# 4. Move to GPU\nmodel = model.to(DEVICE)\nn_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(f\"✓ Successfully loaded {MODEL_NAME} strictly. No fallbacks.\")\nprint(f\"Trainable parameters: {n_params/1e6:.2f}M\")\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 6 — Loss, optimiser, scheduler\n# ─────────────────────────────────────────────────────────────────────────────\n\n# Inverse-frequency class weights\nlabel_arr  = train_df[\"label\"].values\nclass_cnt  = np.bincount(label_arr, minlength=CFG[\"NUM_CLASSES\"]).astype(float)\nclass_wts  = torch.tensor(1.0 / (class_cnt + 1e-6), dtype=torch.float32)\nclass_wts  = (class_wts / class_wts.sum() * CFG[\"NUM_CLASSES\"]).to(DEVICE)\nprint(\"Class weights:\", np.round(class_wts.cpu().numpy(), 3))\n\ncriterion = nn.CrossEntropyLoss(weight=class_wts, label_smoothing=0.10)\n# Slightly higher label smoothing (0.10 vs 0.05) because APTOS is smaller\n# and the model will otherwise overfit the hard labels quickly.\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=CFG[\"LR\"],\n    weight_decay=CFG[\"WEIGHT_DECAY\"],\n    betas=(0.9, 0.999)\n)\n\n# OneCycleLR: ramps up then decays — very effective on small datasets\nsteps_per_epoch = len(train_loader)\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=CFG[\"LR\"] * 10,          # peak LR = 10× base\n    steps_per_epoch=steps_per_epoch,\n    epochs=CFG[\"EPOCHS\"],\n    pct_start=0.2,                   # 20% of training = warm-up\n    anneal_strategy=\"cos\",\n    div_factor=10.0,\n    final_div_factor=100.0,\n)\n\nscaler = GradScaler(enabled=CFG[\"AMP\"])\nprint(\"Optimiser / scheduler / scaler ready ✓\")\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 7 — MixUp helper\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef mixup_data(x, y, alpha=0.4):\n    \"\"\"\n    MixUp: blends two random training images and their labels.\n    Forces the model to learn smooth decision boundaries rather than\n    memorising individual samples — very useful for small datasets.\n    Returns mixed inputs, and the two label sets + mixing lambda.\n    \"\"\"\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1.0\n    batch_size = x.size(0)\n    idx = torch.randperm(batch_size, device=x.device)\n    mixed_x = lam * x + (1 - lam) * x[idx]\n    y_a, y_b = y, y[idx]\n    return mixed_x, y_a, y_b, lam\n\n\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 8 — Training & validation loop\n# ─────────────────────────────────────────────────────────────────────────────\n\ndef run_epoch(loader, model, criterion,\n              optimizer=None, scaler=None, scheduler=None,\n              train=True, use_mixup=False):\n\n    model.train() if train else model.eval()\n    total_loss = 0.0\n    all_preds, all_labels, all_probs = [], [], []\n\n    ctx = torch.enable_grad() if train else torch.no_grad()\n    with ctx:\n        for imgs, labels in loader:\n            imgs   = imgs.to(DEVICE, non_blocking=True)\n            labels = labels.to(DEVICE)\n\n            with autocast(enabled=CFG[\"AMP\"]):\n                if train and use_mixup:\n                    imgs, y_a, y_b, lam = mixup_data(\n                        imgs, labels, CFG[\"MIXUP_ALPHA\"])\n                    logits = model(imgs)\n                    loss   = mixup_criterion(criterion, logits, y_a, y_b, lam)\n                else:\n                    logits = model(imgs)\n                    loss   = criterion(logits, labels)\n\n            if train:\n                optimizer.zero_grad(set_to_none=True)\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(\n                    model.parameters(), CFG[\"GRAD_CLIP\"])\n                scaler.step(optimizer)\n                scaler.update()\n                if scheduler is not None:\n                    scheduler.step()   # OneCycleLR steps per batch\n\n            total_loss += loss.item() * len(labels)\n            probs = torch.softmax(logits.detach(), dim=-1).cpu().numpy()\n            preds = probs.argmax(axis=1)\n            all_preds.extend(preds.tolist())\n            all_labels.extend(labels.cpu().numpy().tolist())\n            all_probs.extend(probs.tolist())\n\n    avg_loss = total_loss / len(all_labels)\n    acc      = accuracy_score(all_labels, all_preds)\n    f1       = f1_score(all_labels, all_preds,\n                        average=\"weighted\", zero_division=0)\n    recall   = recall_score(all_labels, all_preds,\n                             average=\"weighted\", zero_division=0)\n    kappa    = cohen_kappa_score(all_labels, all_preds,\n                                  weights=\"quadratic\")   # APTOS metric\n    try:\n        auc_val = roc_auc_score(all_labels, np.array(all_probs),\n                                multi_class=\"ovr\", average=\"weighted\")\n    except Exception:\n        auc_val = 0.0\n\n    return (avg_loss, acc, f1, recall, kappa, auc_val,\n            all_preds, all_labels, all_probs)\n\n\nhistory = {k: [] for k in [\"train_loss\",\"val_loss\",\"train_acc\",\"val_acc\",\n                            \"train_kappa\",\"val_kappa\",\"val_auc\",\"lr\"]}\nbest_val_kappa = -1.0\npatience_cnt   = 0\n\nprint(f\"\\nModel : {MODEL_NAME}  |  Dataset : APTOS 2019  ({len(train_df)} train images)\")\nprint(f\"{'Ep':>4} | {'T-Loss':>7} | {'V-Loss':>7} | {'T-Acc':>6} | \"\n      f\"{'V-Acc':>6} | {'T-Kap':>6} | {'V-Kap':>6} | {'V-AUC':>6}\")\nprint(\"─\" * 74)\n\nfor epoch in range(1, CFG[\"EPOCHS\"] + 1):\n\n    (tr_loss, tr_acc, tr_f1, tr_rec, tr_kap, tr_auc,\n     _, _, _) = run_epoch(\n        train_loader, model, criterion,\n        optimizer=optimizer, scaler=scaler, scheduler=scheduler,\n        train=True, use_mixup=True\n    )\n\n    (vl_loss, vl_acc, vl_f1, vl_rec, vl_kap, vl_auc,\n     vl_preds, vl_labels, vl_probs) = run_epoch(\n        val_loader, model, criterion, train=False\n    )\n\n    lr_now = optimizer.param_groups[0][\"lr\"]\n    for k, v in zip(\n        [\"train_loss\",\"val_loss\",\"train_acc\",\"val_acc\",\n         \"train_kappa\",\"val_kappa\",\"val_auc\",\"lr\"],\n        [tr_loss, vl_loss, tr_acc, vl_acc,\n         tr_kap, vl_kap, vl_auc, lr_now]\n    ):\n        history[k].append(v)\n\n    print(f\"{epoch:>4} | {tr_loss:>7.4f} | {vl_loss:>7.4f} | {tr_acc:>6.4f} | \"\n          f\"{vl_acc:>6.4f} | {tr_kap:>6.4f} | {vl_kap:>6.4f} | {vl_auc:>6.4f}\")\n\n    # Save on best quadratic kappa (the official APTOS competition metric)\n    if vl_kap > best_val_kappa:\n        best_val_kappa = vl_kap\n        patience_cnt   = 0\n        torch.save({\n            \"epoch\"        : epoch,\n            \"model_state\"  : model.state_dict(),\n            \"optimizer_state\": optimizer.state_dict(),\n            \"val_kappa\"    : best_val_kappa,\n            \"val_auc\"      : vl_auc,\n            \"model_name\"   : MODEL_NAME,\n        }, CFG[\"SAVE_PATH\"])\n        print(f\"        ✓ Saved  (Val Kappa = {best_val_kappa:.4f}  \"\n              f\"Val AUC = {vl_auc:.4f})\")\n    else:\n        patience_cnt += 1\n        if patience_cnt >= CFG[\"PATIENCE\"]:\n            print(f\"\\nEarly stopping at epoch {epoch}.\")\n            break\n\nprint(f\"\\nTraining complete.  Best Val Quadratic Kappa = {best_val_kappa:.4f}\")\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 9 — Training curves + overfitting diagnosis\n# ─────────────────────────────────────────────────────────────────────────────\n\nep = list(range(1, len(history[\"train_loss\"]) + 1))\n\nfig, axes = plt.subplots(2, 2, figsize=(14, 10))\nfig.suptitle(f\"{MODEL_NAME} — Training Curves (APTOS 2019)\", fontsize=14)\n\nax = axes[0, 0]\nax.plot(ep, history[\"train_loss\"], label=\"Train\", color=\"#4C72B0\", lw=2)\nax.plot(ep, history[\"val_loss\"],   label=\"Val\",   color=\"#DD8452\", lw=2)\nax.set_title(\"Cross-Entropy Loss\"); ax.set_xlabel(\"Epoch\")\nax.set_ylabel(\"Loss\"); ax.legend(); ax.grid(alpha=0.3)\n\nax = axes[0, 1]\nax.plot(ep, history[\"train_acc\"], label=\"Train\", color=\"#4C72B0\", lw=2)\nax.plot(ep, history[\"val_acc\"],   label=\"Val\",   color=\"#DD8452\", lw=2)\nax.set_title(\"Accuracy\"); ax.set_xlabel(\"Epoch\")\nax.set_ylabel(\"Accuracy\"); ax.legend(); ax.grid(alpha=0.3)\n\nax = axes[1, 0]\nax.plot(ep, history[\"train_kappa\"], label=\"Train Kappa\", color=\"#4C72B0\", lw=2)\nax.plot(ep, history[\"val_kappa\"],   label=\"Val Kappa\",   color=\"#DD8452\", lw=2)\nax.set_title(\"Quadratic Weighted Kappa (APTOS metric)\")\nax.set_xlabel(\"Epoch\"); ax.set_ylabel(\"Kappa\")\nax.legend(); ax.grid(alpha=0.3)\n\nax  = axes[1, 1]\nax2 = ax.twinx()\nax.plot(ep, history[\"val_auc\"], color=\"#55A868\", lw=2, label=\"Val AUC\")\ngap = [t - v for t, v in zip(history[\"train_loss\"], history[\"val_loss\"])]\nax2.plot(ep, gap, color=\"#C44E52\", lw=1.5, ls=\"--\",\n         alpha=0.7, label=\"Loss gap (overfit indicator)\")\nax.set_title(\"Val AUC + Overfitting Indicator\")\nax.set_xlabel(\"Epoch\"); ax.set_ylabel(\"AUC\")\nax2.set_ylabel(\"Train−Val loss gap\", color=\"#C44E52\")\nax.legend(loc=\"lower right\"); ax2.legend(loc=\"upper right\")\nax.grid(alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/aptos_training_curves.png\", dpi=150)\nplt.show()\n\nfinal_gap    = history[\"train_loss\"][-1] - history[\"val_loss\"][-1]\nfinal_tr_acc = history[\"train_acc\"][-1]\nfinal_vl_acc = history[\"val_acc\"][-1]\nprint(\"\\n── Overfitting / Underfitting Diagnosis ──────────────────────\")\nif final_gap > 0.30:\n    print(f\"⚠ OVERFITTING  (gap = {final_gap:.3f})\")\n    print(\"  Try: higher WEIGHT_DECAY, stronger augmentation, \"\n          \"higher MIXUP_ALPHA, fewer epochs.\")\nelif final_tr_acc < 0.55 and final_vl_acc < 0.55:\n    print(f\"⚠ UNDERFITTING  (train acc = {final_tr_acc:.3f})\")\n    print(\"  Try: more epochs, larger model, lower LR, \"\n          \"disable label smoothing.\")\nelse:\n    print(f\"✓ Healthy  (train acc={final_tr_acc:.3f}, \"\n          f\"val acc={final_vl_acc:.3f}, gap={final_gap:.3f})\")\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 10 — Load best checkpoint & Test-Time Augmentation (TTA) inference\n# ─────────────────────────────────────────────────────────────────────────────\n\nckpt = torch.load(CFG[\"SAVE_PATH\"], map_location=DEVICE, weights_only=False)\nmodel.load_state_dict(ckpt[\"model_state\"])\nprint(f\"Loaded checkpoint from epoch {ckpt['epoch']}  \"\n      f\"(Val Kappa = {ckpt['val_kappa']:.4f})\")\n\n\ndef predict_with_tta(loader, model, n_tta: int):\n    \"\"\"\n    Run n_tta forward passes with random augmentation, average the\n    softmax probabilities. Reduces variance on small test sets.\n    \"\"\"\n    model.eval()\n    tta_transform = get_tta_transform(CFG[\"IMG_SIZE\"])\n    all_probs  = []\n    all_labels = []\n\n    # Collect raw numpy images first (before any transform)\n    raw_imgs, raw_lbls = [], []\n    for imgs_t, lbls in loader:\n        # Reverse normalisation to get back to [0,1] pixel images\n        mean_t = torch.tensor([0.485,0.456,0.406]).view(1,3,1,1)\n        std_t  = torch.tensor([0.229,0.224,0.225]).view(1,3,1,1)\n        imgs_np = ((imgs_t * std_t + mean_t) * 255).clamp(0,255) \\\n                    .permute(0,2,3,1).byte().numpy()\n        raw_imgs.append(imgs_np)\n        raw_lbls.extend(lbls.numpy().tolist())\n    raw_imgs = np.concatenate(raw_imgs, axis=0)\n\n    # n_tta passes\n    tta_probs = np.zeros((len(raw_imgs), CFG[\"NUM_CLASSES\"]), dtype=np.float32)\n    for t in range(n_tta):\n        batch_probs = []\n        for i in range(0, len(raw_imgs), CFG[\"BATCH_SIZE\"]):\n            batch = raw_imgs[i : i + CFG[\"BATCH_SIZE\"]]\n            tensors = torch.stack([\n                tta_transform(image=img)[\"image\"] for img in batch\n            ]).to(DEVICE)\n            with torch.no_grad(), autocast(enabled=CFG[\"AMP\"]):\n                logits = model(tensors)\n            batch_probs.append(\n                torch.softmax(logits, dim=-1).cpu().numpy()\n            )\n        tta_probs += np.concatenate(batch_probs, axis=0)\n\n    tta_probs /= n_tta\n    preds = tta_probs.argmax(axis=1).tolist()\n    return preds, raw_lbls, tta_probs.tolist()\n\n\nprint(f\"\\nRunning TTA inference ({CFG['TTA_STEPS']} passes) on test set …\")\nts_preds, ts_labels, ts_probs = predict_with_tta(\n    test_loader, model, CFG[\"TTA_STEPS\"])\n\nts_acc   = accuracy_score(ts_labels, ts_preds)\nts_f1    = f1_score(ts_labels, ts_preds, average=\"weighted\", zero_division=0)\nts_rec   = recall_score(ts_labels, ts_preds, average=\"weighted\", zero_division=0)\nts_prec  = precision_score(ts_labels, ts_preds, average=\"weighted\", zero_division=0)\nts_kappa = cohen_kappa_score(ts_labels, ts_preds, weights=\"quadratic\")\nts_probs_arr = np.array(ts_probs)\ntry:\n    ts_auc = roc_auc_score(ts_labels, ts_probs_arr,\n                            multi_class=\"ovr\", average=\"weighted\")\nexcept Exception:\n    ts_auc = 0.0\n\nprint(\"\\n\" + \"═\"*54)\nprint(f\"  TEST SET RESULTS — {MODEL_NAME} + TTA×{CFG['TTA_STEPS']}\")\nprint(\"═\"*54)\nprint(f\"  Quadratic Kappa : {ts_kappa:.4f}  ← APTOS competition metric\")\nprint(f\"  Accuracy        : {ts_acc:.4f}  ({ts_acc*100:.2f}%)\")\nprint(f\"  Precision       : {ts_prec:.4f}  (weighted)\")\nprint(f\"  Recall          : {ts_rec:.4f}  (weighted)\")\nprint(f\"  F1 Score        : {ts_f1:.4f}  (weighted)\")\nprint(f\"  AUC (OvR)       : {ts_auc:.4f}  (weighted)\")\nprint(\"═\"*54)\nprint(\"\\nPer-class report:\")\nprint(classification_report(ts_labels, ts_preds,\n                             target_names=CFG[\"CLASS_NAMES\"], digits=4))\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 11 — Confusion matrix\n# ─────────────────────────────────────────────────────────────────────────────\n\ncm      = confusion_matrix(ts_labels, ts_preds)\ncm_norm = cm.astype(\"float\") / cm.sum(axis=1, keepdims=True)\n\nfig, axes = plt.subplots(1, 2, figsize=(16, 6))\nfor ax, data, fmt, title in zip(\n    axes,\n    [cm,     cm_norm],\n    [\"d\",    \".3f\"],\n    [\"Confusion Matrix (counts)\",\n     \"Confusion Matrix (normalised recall)\"]\n):\n    sns.heatmap(data, annot=True, fmt=fmt, cmap=\"Blues\",\n                xticklabels=CFG[\"CLASS_NAMES\"],\n                yticklabels=CFG[\"CLASS_NAMES\"],\n                ax=ax, linewidths=0.5, linecolor=\"white\")\n    ax.set_title(title, fontsize=12)\n    ax.set_xlabel(\"Predicted\"); ax.set_ylabel(\"True\")\n    ax.tick_params(axis=\"x\", rotation=20)\n\nplt.suptitle(f\"{MODEL_NAME} — APTOS Test Set Confusion Matrix\", fontsize=13)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/aptos_confusion_matrix.png\", dpi=150)\nplt.show()\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 12 — ROC / AUC curves\n# ─────────────────────────────────────────────────────────────────────────────\n\nts_labels_arr = np.array(ts_labels)\nts_labels_bin = label_binarize(ts_labels_arr,\n                               classes=list(range(CFG[\"NUM_CLASSES\"])))\ncolors_roc = [\"#4C72B0\",\"#DD8452\",\"#55A868\",\"#C44E52\",\"#8172B2\"]\nroc_aucs   = []\n\nfig, ax = plt.subplots(figsize=(9, 7))\nfor i, (cls_name, color) in enumerate(zip(CFG[\"CLASS_NAMES\"], colors_roc)):\n    fpr, tpr, _ = roc_curve(ts_labels_bin[:, i], ts_probs_arr[:, i])\n    cls_auc     = auc(fpr, tpr)\n    roc_aucs.append(cls_auc)\n    ax.plot(fpr, tpr, color=color, lw=2,\n            label=f\"{cls_name}  (AUC = {cls_auc:.4f})\")\n\nmicro_auc = roc_auc_score(ts_labels_bin, ts_probs_arr, average=\"micro\")\nax.plot([0,1],[0,1], \"k--\", lw=1, label=\"Random (AUC = 0.50)\")\nax.set_xlim([0,1]); ax.set_ylim([0,1.02])\nax.set_xlabel(\"False Positive Rate\", fontsize=12)\nax.set_ylabel(\"True Positive Rate\",  fontsize=12)\nax.set_title(f\"ROC Curves — {MODEL_NAME} on APTOS 2019\\n\"\n             f\"Micro-avg AUC = {micro_auc:.4f}\", fontsize=13)\nax.legend(loc=\"lower right\", fontsize=10)\nax.grid(alpha=0.3)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/aptos_roc_curves.png\", dpi=150)\nplt.show()\n\nprint(\"\\nPer-class AUC:\")\nfor cls_name, a_val in zip(CFG[\"CLASS_NAMES\"], roc_aucs):\n    print(f\"  {cls_name:<16} {a_val:.4f}\")\nprint(f\"  {'Micro-average':<16} {micro_auc:.4f}\")\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 13 — Per-class bar chart\n# ─────────────────────────────────────────────────────────────────────────────\n\nprec_pc = precision_score(ts_labels, ts_preds, average=None, zero_division=0)\nrec_pc  = recall_score(ts_labels,   ts_preds, average=None, zero_division=0)\nf1_pc   = f1_score(ts_labels,       ts_preds, average=None, zero_division=0)\n\nx     = np.arange(CFG[\"NUM_CLASSES\"])\nwidth = 0.25\nfig, ax = plt.subplots(figsize=(11, 5))\nax.bar(x - width, prec_pc, width, label=\"Precision\", color=\"#4C72B0\", alpha=0.85)\nax.bar(x,          rec_pc, width, label=\"Recall\",    color=\"#DD8452\", alpha=0.85)\nax.bar(x + width,  f1_pc,  width, label=\"F1\",        color=\"#55A868\", alpha=0.85)\nax.set_xticks(x)\nax.set_xticklabels(CFG[\"CLASS_NAMES\"], rotation=15, ha=\"right\")\nax.set_ylim(0, 1.05)\nax.set_ylabel(\"Score\")\nax.set_title(f\"Per-class Precision / Recall / F1 — {MODEL_NAME} on APTOS\")\nax.legend(); ax.grid(axis=\"y\", alpha=0.3)\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/aptos_per_class_metrics.png\", dpi=150)\nplt.show()\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n# CELL 14 — Summary CSV\n# ─────────────────────────────────────────────────────────────────────────────\n\nsupport = [int(np.sum(ts_labels_arr == i)) for i in range(CFG[\"NUM_CLASSES\"])]\nsummary_df = pd.DataFrame({\n    \"Class\"    : CFG[\"CLASS_NAMES\"],\n    \"Precision\": np.round(prec_pc,  4),\n    \"Recall\"   : np.round(rec_pc,   4),\n    \"F1\"       : np.round(f1_pc,    4),\n    \"AUC\"      : np.round(roc_aucs, 4),\n    \"Support\"  : support,\n})\noverall_row = pd.DataFrame([{\n    \"Class\"    : \"OVERALL (weighted)\",\n    \"Precision\": round(ts_prec,  4),\n    \"Recall\"   : round(ts_rec,   4),\n    \"F1\"       : round(ts_f1,    4),\n    \"AUC\"      : round(ts_auc,   4),\n    \"Support\"  : len(ts_labels),\n}])\nkappa_row = pd.DataFrame([{\n    \"Class\"    : \"Quadratic Kappa\",\n    \"Precision\": \"—\",\n    \"Recall\"   : \"—\",\n    \"F1\"       : \"—\",\n    \"AUC\"      : \"—\",\n    \"Support\"  : round(ts_kappa, 4),\n}])\nresults_df = pd.concat([summary_df, overall_row, kappa_row], ignore_index=True)\nresults_df.to_csv(\"/kaggle/working/aptos_test_results.csv\", index=False)\n\nprint(\"Final Results:\")\nprint(results_df.to_string(index=False))\n\nprint(\"\\n✅ Outputs saved to /kaggle/working/\")\nfor fname in [\"aptos_class_distribution.png\", \"aptos_samples.png\",\n              \"aptos_training_curves.png\", \"aptos_confusion_matrix.png\",\n              \"aptos_roc_curves.png\", \"aptos_per_class_metrics.png\",\n              \"aptos_test_results.csv\", \"best_aptos_model.pth\"]:\n    fpath = f\"/kaggle/working/{fname}\"\n    size  = f\"{os.path.getsize(fpath)//1024} KB\" if os.path.exists(fpath) else \"missing\"\n    print(f\"  {fname:<42} {size}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-24T11:07:10.33472Z","iopub.execute_input":"2026-05-24T11:07:10.335551Z","iopub.status.idle":"2026-05-24T13:08:21.549244Z","shell.execute_reply.started":"2026-05-24T11:07:10.335516Z","shell.execute_reply":"2026-05-24T13:08:21.548033Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 1 — Environment Setup & Dependencies\n# =============================================================================\nimport subprocess, sys, os\n\ndef run(cmd):\n    r = subprocess.run(cmd, shell=True, capture_output=True, text=True)\n    if r.returncode != 0: print(f\"[WARN] {r.stderr.strip()[-300:]}\")\n    return r.returncode == 0\n\nprint(\"1. Cloning MedMamba Repository...\")\nif not os.path.exists(\"/kaggle/working/MedMamba\"):\n    run(\"git clone --depth 1 https://github.com/YubiaoYue/MedMamba.git /kaggle/working/MedMamba\")\n\nimport torch\nprint(f\"2. GPU Check: {torch.cuda.get_device_name(0) if torch.cuda.is_available() else 'CPU'}\")\nif torch.cuda.is_available():\n    print(\"3. Compiling State-Space Models (mamba-ssm)...\")\n    run(\"pip install -q ninja packaging\")\n    run(\"pip install -q 'causal-conv1d>=1.2.0.post2' 'mamba-ssm>=1.2.0' --no-build-isolation\")\n\nprint(\"4. Installing ML/CV Libraries...\")\nrun(\"pip install -q timm einops scikit-learn matplotlib seaborn albumentations opencv-python-headless\")\n\nsys.path.insert(0, \"/kaggle/working/MedMamba\")\nprint(\"✅ Environment Setup Complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-24T20:38:31.966393Z","iopub.execute_input":"2026-05-24T20:38:31.966752Z","iopub.status.idle":"2026-05-24T20:43:46.011274Z","shell.execute_reply.started":"2026-05-24T20:38:31.966717Z","shell.execute_reply":"2026-05-24T20:43:46.010402Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 2 — Configuration & Path Resolution\n# =============================================================================\nimport os, random, warnings\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom pathlib import Path\n\nwarnings.filterwarnings(\"ignore\")\n\n# ── Reproducibility ──\nSEED = 42\ndef seed_everything(s):\n    random.seed(s); os.environ['PYTHONHASHSEED'] = str(s); np.random.seed(s)\n    torch.manual_seed(s); torch.cuda.manual_seed(s); torch.cuda.manual_seed_all(s)\n    torch.backends.cudnn.deterministic = True; torch.backends.cudnn.benchmark = False\nseed_everything(SEED)\n\n# ── Dynamic Path Resolution ──\n_BASE = None\nfor path in Path('/kaggle/input').rglob('train.csv'):\n    if (path.parent / 'train_images').exists():\n        _BASE = str(path.parent)\n        break\nif not _BASE: raise FileNotFoundError(\"APTOS Dataset not found. Accept rules or remount.\")\n\n# ── Research Configuration ──\nCFG = dict(\n    TRAIN_CSV   = f\"{_BASE}/train.csv\",\n    TRAIN_IMG   = f\"{_BASE}/train_images\",\n    IMG_SIZE    = 224,\n    NUM_CLASSES = 5,\n    BATCH_SIZE  = 32,\n    EPOCHS      = 35,          # 35 is optimal for pretrained Swin; Mamba might need more if from scratch\n    LR_SWIN     = 5e-5,        # Lower LR for pretrained Swin\n    LR_MAMBA    = 1e-4,        # Higher LR for MedMamba (especially if no pretrained weights)\n    WEIGHT_DECAY= 1e-3,\n    MIXUP_ALPHA = 0.4,\n    TTA_STEPS   = 5,           # Test-Time Augmentation passes\n    DEVICE      = \"cuda\" if torch.cuda.is_available() else \"cpu\",\n    CLASS_NAMES = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferative\"],\n    # Model Save Paths\n    SWIN_PATH   = \"/kaggle/working/best_swin.pth\",\n    MAMBA_PATH  = \"/kaggle/working/best_mamba.pth\",\n)\nprint(f\"Dataset mapped to: {_BASE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-25T06:14:28.03495Z","iopub.execute_input":"2026-05-25T06:14:28.03556Z","iopub.status.idle":"2026-05-25T06:14:28.049055Z","shell.execute_reply.started":"2026-05-25T06:14:28.035527Z","shell.execute_reply":"2026-05-25T06:14:28.048356Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 3 — Data Augmentation & Loaders (Fully Fixed API)\n# =============================================================================\nimport cv2\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom sklearn.model_selection import train_test_split\nimport numpy as np\nimport pandas as pd\n\n# 1. Dataset Class\nclass APTOSDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self): return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = cv2.imread(row[\"path\"])\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) if img is not None else np.zeros((CFG[\"IMG_SIZE\"], CFG[\"IMG_SIZE\"], 3), dtype=np.uint8)\n        if self.transform: img = self.transform(image=img)[\"image\"]\n        return img, int(row[\"label\"])\n\n# 2. Augmentations (Research Grade)\nmean, std = [0.485, 0.456, 0.406], [0.229, 0.224, 0.225]\n\ntransform_train = A.Compose([\n    # RandomResizedCrop REQUIRES the 'size' tuple in the latest version\n    A.RandomResizedCrop(size=(CFG[\"IMG_SIZE\"], CFG[\"IMG_SIZE\"]), scale=(0.8, 1.0)),\n    A.HorizontalFlip(p=0.5), A.VerticalFlip(p=0.5), A.RandomRotate90(p=0.5),\n    A.CLAHE(clip_limit=3.0, tile_grid_size=(8, 8), p=0.6), # Enhances blood vessels\n    A.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2, hue=0.05, p=0.5),\n    A.CoarseDropout(max_holes=8, max_height=16, max_width=16, fill_value=0, p=0.3),\n    A.Normalize(mean=mean, std=std), ToTensorV2(),\n])\n\ntransform_val = A.Compose([\n    # Resize STILL REQUIRES separate 'height' and 'width' \n    A.Resize(height=CFG[\"IMG_SIZE\"], width=CFG[\"IMG_SIZE\"]),\n    A.Normalize(mean=mean, std=std), ToTensorV2(),\n])\n\n# 3. Load & Split Data\ndf = pd.read_csv(CFG[\"TRAIN_CSV\"])\ndf['path'] = df['id_code'].apply(lambda x: f\"{CFG['TRAIN_IMG']}/{x}.png\")\ndf['label'] = df['diagnosis']\n\ntrain_df, val_df = train_test_split(df, test_size=0.20, random_state=SEED, stratify=df[\"label\"])\n\n# 4. Weighted Sampler (To fix the imbalance)\nclass_counts = np.bincount(train_df[\"label\"].values)\nweights = 1.0 / class_counts\nsample_weights = weights[train_df[\"label\"].values]\nsampler = WeightedRandomSampler(weights=sample_weights, num_samples=len(sample_weights), replacement=True)\n\ntrain_loader = DataLoader(APTOSDataset(train_df, transform_train), batch_size=CFG[\"BATCH_SIZE\"], sampler=sampler, num_workers=2, pin_memory=True)\nval_loader   = DataLoader(APTOSDataset(val_df, transform_val), batch_size=CFG[\"BATCH_SIZE\"], shuffle=False, num_workers=2)\n\nprint(f\"Train: {len(train_df)} | Val: {len(val_df)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-25T06:14:31.868085Z","iopub.execute_input":"2026-05-25T06:14:31.86887Z","iopub.status.idle":"2026-05-25T06:14:32.008792Z","shell.execute_reply.started":"2026-05-25T06:14:31.868839Z","shell.execute_reply":"2026-05-25T06:14:32.008182Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 4 — Training Engine\n# =============================================================================\nimport torch.nn as nn\nfrom torch.cuda.amp import GradScaler, autocast\nfrom sklearn.metrics import accuracy_score, cohen_kappa_score, roc_auc_score\nfrom copy import deepcopy\n\n# MixUp Regularization\ndef mixup_data(x, y, alpha=0.4):\n    if alpha > 0: lam = np.random.beta(alpha, alpha)\n    else: lam = 1\n    idx = torch.randperm(x.size(0)).to(x.device)\n    mixed_x = lam * x + (1 - lam) * x[idx]\n    return mixed_x, y, y[idx], lam\n\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\n# Training Function\ndef train_model(model, name, lr, save_path):\n    print(f\"\\n{'='*50}\\nTraining {name}\\n{'='*50}\")\n    model = model.to(CFG[\"DEVICE\"])\n    \n    # We DO NOT use class weights here because the Sampler handles it.\n    criterion = nn.CrossEntropyLoss(label_smoothing=0.1) \n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=CFG[\"WEIGHT_DECAY\"])\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CFG[\"EPOCHS\"], eta_min=1e-6)\n    scaler = GradScaler()\n    \n    # Exponential Moving Average (EMA) for stability\n    ema_model = deepcopy(model).eval()\n    ema_decay = 0.999\n    \n    best_kappa = -1.0\n    \n    for epoch in range(1, CFG[\"EPOCHS\"] + 1):\n        model.train()\n        tr_loss = 0.0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(CFG[\"DEVICE\"]), labels.to(CFG[\"DEVICE\"])\n            \n            with autocast():\n                imgs, y_a, y_b, lam = mixup_data(imgs, labels, CFG[\"MIXUP_ALPHA\"])\n                logits = model(imgs)\n                loss = mixup_criterion(criterion, logits, y_a, y_b, lam)\n                \n            optimizer.zero_grad()\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n            tr_loss += loss.item()\n            \n            # Update EMA\n            with torch.no_grad():\n                for param_train, param_ema in zip(model.parameters(), ema_model.parameters()):\n                    param_ema.data.mul_(ema_decay).add_(param_train.data, alpha=1 - ema_decay)\n        \n        scheduler.step()\n        \n        # Validation using EMA model\n        ema_model.eval()\n        all_preds, all_labels = [], []\n        with torch.no_grad():\n            for imgs, labels in val_loader:\n                imgs = imgs.to(CFG[\"DEVICE\"])\n                logits = ema_model(imgs)\n                preds = logits.argmax(dim=1).cpu().numpy()\n                all_preds.extend(preds)\n                all_labels.extend(labels.numpy())\n                \n        val_acc = accuracy_score(all_labels, all_preds)\n        val_kappa = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n        \n        print(f\"Epoch {epoch:02d} | Loss: {tr_loss/len(train_loader):.4f} | Val Acc: {val_acc:.4f} | Val Kappa: {val_kappa:.4f}\")\n        \n        if val_kappa > best_kappa:\n            best_kappa = val_kappa\n            torch.save(ema_model.state_dict(), save_path)\n            print(f\"  --> Saved {name} Checkpoint (Kappa: {best_kappa:.4f})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-24T20:48:05.252371Z","iopub.execute_input":"2026-05-24T20:48:05.252919Z","iopub.status.idle":"2026-05-24T20:48:05.264607Z","shell.execute_reply.started":"2026-05-24T20:48:05.25289Z","shell.execute_reply":"2026-05-24T20:48:05.263951Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 5 — Execute Training\n# =============================================================================\nimport timm\nfrom MedMamba import VSSM\n\n# 1. Initialize Swin Transformer\nprint(\"Loading Swin-Tiny...\")\nmodel_swin = timm.create_model(\"swin_tiny_patch4_window7_224\", pretrained=True, num_classes=CFG[\"NUM_CLASSES\"])\ntrain_model(model_swin, \"Swin-Transformer\", CFG[\"LR_SWIN\"], CFG[\"SWIN_PATH\"])\n\n# 2. Initialize MedMamba\nprint(\"\\nLoading MedMamba-T...\")\nmodel_mamba = VSSM(patch_size=4, in_chans=3, num_classes=CFG[\"NUM_CLASSES\"], depths=[2, 2, 4, 2], embed_dim=96)\n\n# [OPTIONAL BUT HIGHLY RECOMMENDED] Load MedMamba Pretrained Weights Here if you have them uploaded\n# ckpt = torch.load(\"/kaggle/input/your-weights/vssm_tiny.pth\", map_location=\"cpu\", weights_only=False)\n# model_mamba.load_state_dict(ckpt['model'], strict=False)\n\ntrain_model(model_mamba, \"MedMamba\", CFG[\"LR_MAMBA\"], CFG[\"MAMBA_PATH\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-24T20:48:10.891292Z","iopub.execute_input":"2026-05-24T20:48:10.892046Z","iopub.status.idle":"2026-05-25T01:34:41.166605Z","shell.execute_reply.started":"2026-05-24T20:48:10.892017Z","shell.execute_reply":"2026-05-25T01:34:41.165398Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 6 — Ensemble Evaluation & Research Metrics\n# =============================================================================\nfrom sklearn.metrics import classification_report, confusion_matrix\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Load Best Models\nprint(\"Loading Best Checkpoints for Ensemble...\")\nswin = timm.create_model(\"swin_tiny_patch4_window7_224\", num_classes=CFG[\"NUM_CLASSES\"]).to(CFG[\"DEVICE\"])\nswin.load_state_dict(torch.load(CFG[\"SWIN_PATH\"], map_location=CFG[\"DEVICE\"], weights_only=False))\nswin.eval()\n\nmamba = VSSM(patch_size=4, in_chans=3, num_classes=CFG[\"NUM_CLASSES\"], depths=[2, 2, 4, 2], embed_dim=96).to(CFG[\"DEVICE\"])\nmamba.load_state_dict(torch.load(CFG[\"MAMBA_PATH\"], map_location=CFG[\"DEVICE\"], weights_only=False))\nmamba.eval()\n\nall_labels, all_ensemble_preds, all_ensemble_probs = [], [], []\n\nprint(\"Running Dual-Architecture Inference...\")\nwith torch.no_grad():\n    for imgs, labels in val_loader:\n        imgs = imgs.to(CFG[\"DEVICE\"])\n        \n        # Get raw logits\n        logits_swin = swin(imgs)\n        logits_mamba = mamba(imgs)\n        \n        # Convert to probabilities\n        probs_swin = torch.softmax(logits_swin, dim=1)\n        probs_mamba = torch.softmax(logits_mamba, dim=1)\n        \n        # ENSEMBLE: Average the probabilities (Soft Voting)\n        # You can adjust weights (e.g., 0.6 * Swin + 0.4 * Mamba) based on validation\n        probs_ensemble = (probs_swin + probs_mamba) / 2.0 \n        \n        preds = probs_ensemble.argmax(dim=1).cpu().numpy()\n        \n        all_labels.extend(labels.numpy())\n        all_ensemble_preds.extend(preds)\n        all_ensemble_probs.extend(probs_ensemble.cpu().numpy())\n\n# Calculate Metrics\nacc = accuracy_score(all_labels, all_ensemble_preds)\nkappa = cohen_kappa_score(all_labels, all_ensemble_preds, weights=\"quadratic\")\n\nprint(\"\\n\" + \"=\"*50)\nprint(f\"🏆 ENSEMBLE RESULTS (Swin + MedMamba)\")\nprint(\"=\"*50)\nprint(f\"Accuracy        : {acc:.4f} ({(acc*100):.2f}%)\")\nprint(f\"Quadratic Kappa : {kappa:.4f}\")\nprint(\"=\"*50)\n\nprint(\"\\nClassification Report:\")\nprint(classification_report(all_labels, all_ensemble_preds, target_names=CFG[\"CLASS_NAMES\"]))\n\n# Generate Confusion Matrix for the Research Paper\ncm = confusion_matrix(all_labels, all_ensemble_preds)\nplt.figure(figsize=(8, 6))\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", xticklabels=CFG[\"CLASS_NAMES\"], yticklabels=CFG[\"CLASS_NAMES\"])\nplt.title(f\"Ensemble Confusion Matrix (Kappa: {kappa:.4f})\")\nplt.ylabel(\"Actual Diagnosis\")\nplt.xlabel(\"Predicted Diagnosis\")\nplt.savefig(\"/kaggle/working/ensemble_confusion_matrix.png\", dpi=300)\nplt.show()\nprint(\"Saved publication-ready confusion matrix to /kaggle/working/\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-25T05:23:14.045732Z","iopub.execute_input":"2026-05-25T05:23:14.046146Z","iopub.status.idle":"2026-05-25T05:23:15.676269Z","shell.execute_reply.started":"2026-05-25T05:23:14.046116Z","shell.execute_reply":"2026-05-25T05:23:15.67528Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 1 — Environment Setup (Official Vision Mamba)\n# =============================================================================\nimport subprocess, sys, os\n\ndef run(cmd):\n    r = subprocess.run(cmd, shell=True, capture_output=True, text=True)\n    if r.returncode != 0: print(f\"[WARN] {r.stderr.strip()[-300:]}\")\n    return r.returncode == 0\n\nprint(\"1. Cloning Official Vision Mamba (Vim) Repository...\")\nif not os.path.exists(\"/kaggle/working/Vim\"):\n    run(\"git clone --depth 1 https://github.com/hustvl/Vim.git /kaggle/working/Vim\")\n\nimport torch\nprint(f\"2. GPU Check: {torch.cuda.get_device_name(0) if torch.cuda.is_available() else 'CPU'}\")\nif torch.cuda.is_available():\n    print(\"3. Compiling State-Space Models (mamba-ssm)...\")\n    run(\"pip install -q ninja packaging\")\n    run(\"pip install -q 'causal-conv1d>=1.2.0.post2' 'mamba-ssm>=1.2.0' --no-build-isolation\")\n\nprint(\"4. Installing ML/CV Libraries...\")\nrun(\"pip install -q timm einops scikit-learn matplotlib seaborn albumentations opencv-python-headless\")\n\n# Add Vim to path\nsys.path.insert(0, \"/kaggle/working/Vim\")\nprint(\"✅ Environment Setup Complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-25T06:06:51.118326Z","iopub.execute_input":"2026-05-25T06:06:51.118975Z","iopub.status.idle":"2026-05-25T06:11:58.374979Z","shell.execute_reply.started":"2026-05-25T06:06:51.118935Z","shell.execute_reply":"2026-05-25T06:11:58.374213Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 4 — Training Engine (With Anti-Collapse Warmup)\n# =============================================================================\nimport torch.nn as nn\nfrom torch.cuda.amp import GradScaler, autocast\nfrom sklearn.metrics import accuracy_score, cohen_kappa_score\nfrom copy import deepcopy\n\ndef mixup_data(x, y, alpha=0.4):\n    if alpha > 0: lam = np.random.beta(alpha, alpha)\n    else: lam = 1\n    idx = torch.randperm(x.size(0)).to(x.device)\n    mixed_x = lam * x + (1 - lam) * x[idx]\n    return mixed_x, y, y[idx], lam\n\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\ndef train_vision_mamba(model, name, lr, save_path):\n    print(f\"\\n{'='*50}\\nTraining {name}\\n{'='*50}\")\n    model = model.to(CFG[\"DEVICE\"])\n    \n    criterion = nn.CrossEntropyLoss(label_smoothing=0.1) \n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=CFG[\"WEIGHT_DECAY\"])\n    \n    # THE FIX: 30% Warmup Scheduler to prevent Mode Collapse\n    steps_per_epoch = len(train_loader)\n    scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        optimizer,\n        max_lr=lr,\n        steps_per_epoch=steps_per_epoch,\n        epochs=CFG[\"EPOCHS\"],\n        pct_start=0.3,  # Spends the first 30% of training slowly ramping up LR\n        anneal_strategy=\"cos\"\n    )\n    \n    scaler = GradScaler()\n    ema_model = deepcopy(model).eval()\n    ema_decay = 0.999\n    best_kappa = -1.0\n    \n    for epoch in range(1, CFG[\"EPOCHS\"] + 1):\n        model.train()\n        tr_loss = 0.0\n        for imgs, labels in train_loader:\n            imgs, labels = imgs.to(CFG[\"DEVICE\"]), labels.to(CFG[\"DEVICE\"])\n            \n            with autocast():\n                imgs, y_a, y_b, lam = mixup_data(imgs, labels, CFG[\"MIXUP_ALPHA\"])\n                logits = model(imgs)\n                loss = mixup_criterion(criterion, logits, y_a, y_b, lam)\n                \n            optimizer.zero_grad()\n            scaler.scale(loss).backward()\n            \n            # Gradient Clipping provides an extra layer of protection against exploding gradients\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            \n            scaler.step(optimizer)\n            scaler.update()\n            \n            # Step the warm-up scheduler EVERY BATCH\n            scheduler.step()\n            \n            tr_loss += loss.item()\n            \n            with torch.no_grad():\n                for param_train, param_ema in zip(model.parameters(), ema_model.parameters()):\n                    param_ema.data.mul_(ema_decay).add_(param_train.data, alpha=1 - ema_decay)\n        \n        ema_model.eval()\n        all_preds, all_labels = [], []\n        with torch.no_grad():\n            for imgs, labels in val_loader:\n                imgs = imgs.to(CFG[\"DEVICE\"])\n                logits = ema_model(imgs)\n                preds = logits.argmax(dim=1).cpu().numpy()\n                all_preds.extend(preds)\n                all_labels.extend(labels.numpy())\n                \n        val_acc = accuracy_score(all_labels, all_preds)\n        val_kappa = cohen_kappa_score(all_labels, all_preds, weights=\"quadratic\")\n        \n        print(f\"Epoch {epoch:02d} | Loss: {tr_loss/len(train_loader):.4f} | Val Acc: {val_acc:.4f} | Val Kappa: {val_kappa:.4f}\")\n        \n        if val_kappa > best_kappa:\n            best_kappa = val_kappa\n            torch.save(ema_model.state_dict(), save_path)\n            print(f\"  --> Saved {name} Checkpoint (Kappa: {best_kappa:.4f})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-25T06:19:50.83084Z","iopub.execute_input":"2026-05-25T06:19:50.831517Z","iopub.status.idle":"2026-05-25T06:19:50.844154Z","shell.execute_reply.started":"2026-05-25T06:19:50.831484Z","shell.execute_reply":"2026-05-25T06:19:50.843384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================================================\n# CELL 5 — Execute Pure Vision Mamba (VMamba-Tiny) - Path Fixed\n# =============================================================================\nimport sys\nimport os\nimport subprocess\n\n# 1. Safety Check: If the folder got deleted, clone it again instantly\nif not os.path.exists(\"/kaggle/working/MedMamba\"):\n    print(\"Re-cloning MedMamba repository...\")\n    subprocess.run(\"git clone --depth 1 https://github.com/YubiaoYue/MedMamba.git /kaggle/working/MedMamba\", shell=True)\n\n# 2. Force Python to recognize the MedMamba folder\nif \"/kaggle/working/MedMamba\" not in sys.path:\n    sys.path.insert(0, \"/kaggle/working/MedMamba\")\n\n# 3. Import the Pure Vision Mamba architecture\nfrom MedMamba import VSSM\n\nprint(\"\\nLoading Pure Vision Mamba (VMamba-Tiny / VSSM)...\")\n\n# This explicitly initializes the 4-way scanning Vision State Space Model\n# Parameters: 96 base channels, depths=[2, 2, 4, 2] equates to the ~14.4M Tiny variant\nmodel_vmamba = VSSM(\n    patch_size=4, \n    in_chans=3, \n    num_classes=CFG[\"NUM_CLASSES\"], \n    depths=[2, 2, 4, 2], \n    embed_dim=96\n)\n\n# Start training! We give it a slightly higher max_lr (5e-4) because the \n# OneCycleLR scheduler will throttle it down heavily at the start to protect it.\ntrain_vision_mamba(model_vmamba, \"Vision-Mamba-Pure\", lr=5e-4, save_path=CFG[\"MAMBA_PATH\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-25T06:20:44.738674Z","iopub.execute_input":"2026-05-25T06:20:44.739255Z","iopub.status.idle":"2026-05-25T06:21:19.118294Z","shell.execute_reply.started":"2026-05-25T06:20:44.739225Z","shell.execute_reply":"2026-05-25T06:21:19.116419Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}