{"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,"isSourceIdPinned":false},{"sourceType":"datasetVersion","sourceId":532013,"datasetId":253160,"databundleVersionId":548363}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport cv2\nimport time\nimport random\nimport numpy as np\nimport pandas as pd","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:44:39.769881Z","iopub.execute_input":"2026-03-08T07:44:39.770460Z","iopub.status.idle":"2026-03-08T07:44:40.732574Z","shell.execute_reply.started":"2026-03-08T07:44:39.770429Z","shell.execute_reply":"2026-03-08T07:44:40.731945Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\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 f1_score, accuracy_score, cohen_kappa_score\nfrom tqdm.auto import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:44:46.779863Z","iopub.execute_input":"2026-03-08T07:44:46.780298Z","iopub.status.idle":"2026-03-08T07:44:58.493854Z","shell.execute_reply.started":"2026-03-08T07:44:46.780270Z","shell.execute_reply":"2026-03-08T07:44:58.493056Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import timm\nimport gc\nfrom PIL import Image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:44:58.495003Z","iopub.execute_input":"2026-03-08T07:44:58.495561Z","iopub.status.idle":"2026-03-08T07:45:03.108422Z","shell.execute_reply.started":"2026-03-08T07:44:58.495537Z","shell.execute_reply":"2026-03-08T07:45:03.107590Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"CONFIG = {\n    'seed': 42,\n    'img_size': 256,       # Must be 256 for Swin V2 Tiny\n    'batch_size': 16,      # Fits in 16GB VRAM\n    'accum_steps': 2,      # Effective batch size = 32\n    'num_classes': 5,\n    'device': torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:45:06.559818Z","iopub.execute_input":"2026-03-08T07:45:06.560594Z","iopub.status.idle":"2026-03-08T07:45:06.825703Z","shell.execute_reply.started":"2026-03-08T07:45:06.560564Z","shell.execute_reply":"2026-03-08T07:45:06.824886Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"2015 --- > https://www.kaggle.com/datasets/benjaminwarner/resized-2015-2019-blindness-detection-images\n2019 --- > https://www.kaggle.com/competitions/aptos2019-blindness-detection/data","metadata":{}},{"cell_type":"code","source":"DATA_2015_PATH = '/kaggle/input/datasets/benjaminwarner/resized-2015-2019-blindness-detection-images/resized train 15'\nLABEL_2015_PATH = '/kaggle/input/datasets/benjaminwarner/resized-2015-2019-blindness-detection-images/labels/trainLabels15.csv'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:50:09.547552Z","iopub.execute_input":"2026-03-08T07:50:09.548070Z","iopub.status.idle":"2026-03-08T07:50:09.551180Z","shell.execute_reply.started":"2026-03-08T07:50:09.548042Z","shell.execute_reply":"2026-03-08T07:50:09.550506Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"DATA_2019_PATH = '/kaggle/input/competitions/aptos2019-blindness-detection/train_images'\nLABEL_2019_PATH = '/kaggle/input/competitions/aptos2019-blindness-detection/train.csv'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:50:09.811968Z","iopub.execute_input":"2026-03-08T07:50:09.812642Z","iopub.status.idle":"2026-03-08T07:50:09.815681Z","shell.execute_reply.started":"2026-03-08T07:50:09.812614Z","shell.execute_reply":"2026-03-08T07:50:09.815148Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:50:30.694943Z","iopub.execute_input":"2026-03-08T07:50:30.695261Z","iopub.status.idle":"2026-03-08T07:50:30.699629Z","shell.execute_reply.started":"2026-03-08T07:50:30.695234Z","shell.execute_reply":"2026-03-08T07:50:30.699077Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seed_everything(CONFIG['seed'])\nprint(f\"Device: {CONFIG['device']}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:50:38.847856Z","iopub.execute_input":"2026-03-08T07:50:38.848532Z","iopub.status.idle":"2026-03-08T07:50:38.862950Z","shell.execute_reply.started":"2026-03-08T07:50:38.848503Z","shell.execute_reply":"2026-03-08T07:50:38.862421Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class h_sigmoid(nn.Module):\n    def __init__(self, inplace=True):\n        super(h_sigmoid, self).__init__()\n        self.relu = nn.ReLU6(inplace=inplace)\n    def forward(self, x):\n        return self.relu(x + 3) / 6","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:51:27.892917Z","iopub.execute_input":"2026-03-08T07:51:27.893619Z","iopub.status.idle":"2026-03-08T07:51:27.897929Z","shell.execute_reply.started":"2026-03-08T07:51:27.893588Z","shell.execute_reply":"2026-03-08T07:51:27.897238Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class h_swish(nn.Module):\n    def __init__(self, inplace=True):\n        super(h_swish, self).__init__()\n        self.sigmoid = h_sigmoid(inplace=inplace)\n    def forward(self, x):\n        return x * self.sigmoid(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:51:29.100477Z","iopub.execute_input":"2026-03-08T07:51:29.100764Z","iopub.status.idle":"2026-03-08T07:51:29.104978Z","shell.execute_reply.started":"2026-03-08T07:51:29.100739Z","shell.execute_reply":"2026-03-08T07:51:29.104389Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CoordinateAttention(nn.Module):\n    \"\"\"Paper Eq 10: Coordinate Attention mechanism\"\"\"\n    def __init__(self, inp, reduction=32):\n        super(CoordinateAttention, self).__init__()\n        self.pool_h = nn.AdaptiveAvgPool2d((None, 1))\n        self.pool_w = nn.AdaptiveAvgPool2d((1, None))\n\n        mip = max(8, inp // reduction)\n\n        self.conv1 = nn.Conv2d(inp, mip, kernel_size=1, stride=1, padding=0)\n        self.bn1 = nn.BatchNorm2d(mip)\n        self.act = h_swish()\n        \n        self.conv_h = nn.Conv2d(mip, inp, kernel_size=1, stride=1, padding=0)\n        self.conv_w = nn.Conv2d(mip, inp, kernel_size=1, stride=1, padding=0)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        identity = x\n        n, c, h, w = x.size()\n        x_h = self.pool_h(x)\n        x_w = self.pool_w(x).permute(0, 1, 3, 2)\n\n        y = torch.cat([x_h, x_w], dim=2)\n        y = self.conv1(y)\n        y = self.bn1(y)\n        y = self.act(y) \n        \n        x_h, x_w = torch.split(y, [h, w], dim=2)\n        x_w = x_w.permute(0, 1, 3, 2)\n\n        a_h = self.sigmoid(self.conv_h(x_h))\n        a_w = self.sigmoid(self.conv_w(x_w))\n\n        return identity * a_h * a_w","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:51:36.412919Z","iopub.execute_input":"2026-03-08T07:51:36.413252Z","iopub.status.idle":"2026-03-08T07:51:36.420526Z","shell.execute_reply.started":"2026-03-08T07:51:36.413224Z","shell.execute_reply":"2026-03-08T07:51:36.419981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class HybridModel(nn.Module):\n    def __init__(self, num_classes=5):\n        super(HybridModel, self).__init__()\n        \n        # --- Branch 1: EfficientNet-B3 ---\n        self.effnet = timm.create_model('efficientnet_b3', pretrained=True)\n        self.eff_n_features = self.effnet.classifier.in_features\n        self.effnet.classifier = nn.Identity() \n        self.effnet.global_pool = nn.Identity() \n        \n        # MHSA for EfficientNet\n        self.eff_mhsa = nn.MultiheadAttention(embed_dim=self.eff_n_features, num_heads=8, batch_first=True)\n\n        # --- Branch 2: Swin Transformer V2 ---\n        self.swin = timm.create_model('swinv2_tiny_window8_256', pretrained=True)\n        self.swin_n_features = self.swin.head.in_features\n        self.swin.head = nn.Identity() \n        \n        # Coordinate Attention for Swin\n        self.swin_coord_att = CoordinateAttention(inp=self.swin_n_features)\n        \n        # --- Fusion & Classifier ---\n        fusion_dim = self.eff_n_features + self.swin_n_features\n        self.classifier = nn.Sequential(\n            nn.Dropout(0.2), \n            nn.Linear(fusion_dim, 512),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x):\n        # --- EfficientNet Forward ---\n        eff_feat = self.effnet.forward_features(x) # (B, C, H, W)\n        \n        # Reshape for MHSA: (B, C, H, W) -> (B, H*W, C)\n        b, c, h, w = eff_feat.shape\n        eff_tokens = eff_feat.flatten(2).transpose(1, 2)\n        \n        # Apply MHSA\n        eff_att, _ = self.eff_mhsa(eff_tokens, eff_tokens, eff_tokens)\n        eff_out = torch.mean(eff_att, dim=1) \n        \n        # --- Swin Forward ---\n        swin_feat = self.swin.forward_features(x) # (B, H, W, C)\n        \n        # Fix Dimensions for Coordinate Attention: (B, H, W, C) -> (B, C, H, W)\n        if swin_feat.dim() == 4:\n            swin_feat = swin_feat.permute(0, 3, 1, 2)\n        \n        # Apply Coordinate Attention\n        swin_att = self.swin_coord_att(swin_feat)\n        \n        # Global Pooling\n        swin_out = torch.mean(swin_att.flatten(2), dim=2)\n\n        # --- Fusion ---\n        combined = torch.cat((eff_out, swin_out), dim=1)\n        output = self.classifier(combined)\n        return output","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:51:44.057148Z","iopub.execute_input":"2026-03-08T07:51:44.057864Z","iopub.status.idle":"2026-03-08T07:51:44.065870Z","shell.execute_reply.started":"2026-03-08T07:51:44.057831Z","shell.execute_reply":"2026-03-08T07:51:44.065158Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DRDataset(Dataset):\n    def __init__(self, df, base_path, extension=\".png\", transform=None):\n        self.df = df\n        self.base_path = base_path\n        self.transform = transform\n        self.extension = extension\n        \n        # Sanity Check\n        if len(df) > 0:\n            sample_id = df.iloc[0, 0] if 'image' in df.columns else df.iloc[0, 0]\n            # Verify if at least one file exists\n            found = False\n            for ext in [extension, '.jpg', '.jpeg', '.png']:\n                if os.path.exists(os.path.join(base_path, f\"{sample_id}{ext}\")):\n                    found = True\n                    break\n            if not found:\n                print(f\"⚠️ WARNING: Could not find sample image: {sample_id} in {base_path}\")\n                print(f\"⚠️ Expected path: {os.path.join(base_path, sample_id + extension)}\")\n\n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        # Handle column name differences\n        if 'image' in self.df.columns:\n            id_code = self.df.loc[idx, 'image']\n            label = self.df.loc[idx, 'level']\n        else:\n            id_code = self.df.loc[idx, 'id_code']\n            label = self.df.loc[idx, 'diagnosis']\n            \n        # Try loading\n        image = None\n        # Try preferred extension first\n        try:\n            p = os.path.join(self.base_path, f\"{id_code}{self.extension}\")\n            if os.path.exists(p):\n                image = cv2.imread(p)\n        except: pass\n        \n        # Fallback to other extensions\n        if image is None:\n            for ext in ['.jpg', '.jpeg', '.png']:\n                try:\n                    p = os.path.join(self.base_path, f\"{id_code}{ext}\")\n                    if os.path.exists(p):\n                        image = cv2.imread(p)\n                        if image is not None: break\n                except: pass\n\n        # Safety Net: Black Image\n        if image is None:\n            image = np.zeros((CONFIG['img_size'], CONFIG['img_size'], 3), dtype=np.uint8)\n        else:\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        image = cv2.resize(image, (CONFIG['img_size'], CONFIG['img_size']))\n        img = Image.fromarray(image)\n        \n        if self.transform:\n            img = self.transform(img)\n            \n        return img, torch.tensor(label, dtype=torch.long)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:51:56.425283Z","iopub.execute_input":"2026-03-08T07:51:56.425801Z","iopub.status.idle":"2026-03-08T07:51:56.436720Z","shell.execute_reply.started":"2026-03-08T07:51:56.425775Z","shell.execute_reply":"2026-03-08T07:51:56.436155Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_full_training():\n    \n    # --- Transforms ---\n    train_transforms = transforms.Compose([\n        transforms.RandomRotation(180),\n        transforms.RandomHorizontalFlip(),\n        transforms.RandomVerticalFlip(),\n        transforms.ColorJitter(brightness=0.15, contrast=0.15),\n        transforms.RandomResizedCrop(CONFIG['img_size'], scale=(0.85, 1.0)),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n    ])\n    \n    val_transforms = transforms.Compose([\n        transforms.Resize((CONFIG['img_size'], CONFIG['img_size'])),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])\n    ])\n\n    # Initialize Model & Tools\n    model = HybridModel(num_classes=5).to(CONFIG['device'])\n    scaler = torch.amp.GradScaler('cuda')\n    criterion = nn.CrossEntropyLoss()\n\n    # ====================================================\n    # STAGE 1: BALANCED PRE-TRAINING (2015 DATASET)\n    # ====================================================\n    print(\"\\n\" + \"=\"*60)\n    print(\"STAGE 1: Pre-training on 2015 Dataset (15 Epochs)\")\n    print(\"=\"*60)\n    \n    df_2015 = pd.read_csv(LABEL_2015_PATH)\n    # Filter valid labels\n    df_2015 = df_2015[df_2015['level'].isin([0,1,2,3,4])].reset_index(drop=True)\n    \n    # Create Weighted Sampler for 2015\n    class_counts = df_2015['level'].value_counts().sort_index().values\n    weights = 1. / class_counts\n    samples_weights = torch.DoubleTensor([weights[t] for t in df_2015['level']])\n    # Sample 15,000 images per epoch\n    sampler_2015 = WeightedRandomSampler(samples_weights, num_samples=15000, replacement=True)\n    \n    ds_2015 = DRDataset(df_2015, DATA_2015_PATH, extension=\".jpg\", transform=train_transforms)\n    dl_2015 = DataLoader(ds_2015, batch_size=CONFIG['batch_size'], sampler=sampler_2015, \n                         num_workers=2, pin_memory=True, persistent_workers=True)\n    \n    optimizer = optim.AdamW(model.parameters(), lr=3e-4, weight_decay=1e-4)\n    \n    for epoch in range(15):\n        model.train()\n        running_loss = 0.0\n        train_preds, train_targets = [], []\n        \n        pbar = tqdm(dl_2015, desc=f\"Pre-train {epoch+1}/15\")\n        \n        for batch_idx, (images, labels) in enumerate(pbar):\n            images, labels = images.to(CONFIG['device']), labels.to(CONFIG['device'])\n            \n            with torch.amp.autocast('cuda'):\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                loss = loss / CONFIG['accum_steps']\n            \n            scaler.scale(loss).backward()\n            \n            if ((batch_idx + 1) % CONFIG['accum_steps'] == 0):\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n            \n            # Metrics tracking\n            current_loss = loss.item() * CONFIG['accum_steps']\n            running_loss += current_loss\n            _, predicted = torch.max(outputs.data, 1)\n            train_preds.extend(predicted.detach().cpu().numpy())\n            train_targets.extend(labels.detach().cpu().numpy())\n            \n            pbar.set_postfix({'loss': f\"{current_loss:.4f}\"})\n        \n        # Calculate Training Metrics\n        train_loss = running_loss / len(dl_2015)\n        train_acc = accuracy_score(train_targets, train_preds)\n        train_f1 = f1_score(train_targets, train_preds, average='weighted')\n        train_kappa = cohen_kappa_score(train_targets, train_preds, weights='quadratic')\n        \n        print(f\"Epoch {epoch+1} Results:\")\n        print(f\"  [Train] Loss: {train_loss:.4f} | Acc: {train_acc:.4f} | F1: {train_f1:.4f} | Kappa: {train_kappa:.4f}\")\n            \n    print(\"Saving Pre-trained weights...\")\n    torch.save(model.state_dict(), 'pretrained_2015.pth')\n\n    # ====================================================\n    # STAGE 2: FINE-TUNING (APTOS 2019)\n    # ====================================================\n    print(\"\\n\" + \"=\"*60)\n    print(\"STAGE 2: Fine-tuning on APTOS 2019 (30 Epochs)\")\n    print(\"=\"*60)\n    \n    df_2019 = pd.read_csv(LABEL_2019_PATH)\n    train_sub, val_sub = train_test_split(df_2019, test_size=0.2, stratify=df_2019['diagnosis'], random_state=CONFIG['seed'])\n    train_sub = train_sub.reset_index(drop=True)\n    val_sub = val_sub.reset_index(drop=True)\n    \n    # Weighted Sampler for 2019\n    class_counts_19 = train_sub['diagnosis'].value_counts().sort_index().values\n    weights_19 = 1. / class_counts_19\n    samples_weights_19 = torch.DoubleTensor([weights_19[t] for t in train_sub['diagnosis']])\n    sampler_19 = WeightedRandomSampler(samples_weights_19, len(samples_weights_19))\n    \n    ds_train = DRDataset(train_sub, DATA_2019_PATH, extension=\".png\", transform=train_transforms)\n    ds_val = DRDataset(val_sub, DATA_2019_PATH, extension=\".png\", transform=val_transforms)\n    \n    dl_train = DataLoader(ds_train, batch_size=CONFIG['batch_size'], sampler=sampler_19, \n                          num_workers=2, pin_memory=True, persistent_workers=True)\n    dl_val = DataLoader(ds_val, batch_size=CONFIG['batch_size'], shuffle=False, \n                        num_workers=2, pin_memory=True, persistent_workers=True)\n    \n    # Lower LR for fine-tuning\n    optimizer = optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)\n    # Scheduler: Reduces LR if F1 score stagnates\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=3)\n    \n    best_f1 = 0.0\n    train_losses = []\n    val_losses = []\n    train_accs = []\n    val_accs = []\n    for epoch in range(30):\n        model.train()\n        running_loss = 0.0\n        train_preds, train_targets = [], []\n        \n        pbar = tqdm(dl_train, desc=f\"Fine-tune {epoch+1}/30\")\n        optimizer.zero_grad()\n        \n        for batch_idx, (images, labels) in enumerate(pbar):\n            images, labels = images.to(CONFIG['device']), labels.to(CONFIG['device'])\n            \n            with torch.amp.autocast('cuda'):\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                loss = loss / CONFIG['accum_steps']\n                \n            scaler.scale(loss).backward()\n            \n            if ((batch_idx + 1) % CONFIG['accum_steps'] == 0):\n                scaler.step(optimizer)\n                scaler.update()\n                optimizer.zero_grad()\n            \n            # Track training metrics\n            current_loss = loss.item() * CONFIG['accum_steps']\n            running_loss += current_loss\n            _, predicted = torch.max(outputs.data, 1)\n            train_preds.extend(predicted.detach().cpu().numpy())\n            train_targets.extend(labels.detach().cpu().numpy())\n            \n            pbar.set_postfix({'loss': f\"{current_loss:.4f}\"})\n            \n        # Calculate Training Stats\n        train_loss = running_loss / len(dl_train)\n        train_acc = accuracy_score(train_targets, train_preds)\n        train_f1 = f1_score(train_targets, train_preds, average='weighted')\n        train_kappa = cohen_kappa_score(train_targets, train_preds, weights='quadratic')\n        train_losses.append(train_loss)\n        train_accs.append(train_acc)\n\n        # Validation with Test Time Augmentation (TTA)\n        model.eval()\n        val_preds, val_targets = [], []\n        with torch.no_grad():\n            for images, labels in dl_val:\n                images = images.to(CONFIG['device'])\n                with torch.amp.autocast('cuda'):\n                    # TTA: Predict twice (normal + flip) and average\n                    out1 = model(images)\n                    out2 = model(torch.flip(images, dims=[3]))\n                    outputs = (out1 + out2) / 2.0\n                    \n                _, predicted = torch.max(outputs.data, 1)\n                val_preds.extend(predicted.cpu().numpy())\n                val_targets.extend(labels.cpu().numpy())\n                \n        val_acc = accuracy_score(val_targets, val_preds)\n        val_f1 = f1_score(val_targets, val_preds, average='weighted')\n        val_kappa = cohen_kappa_score(val_targets, val_preds, weights='quadratic')\n        val_accs.append(val_acc)\n        \n        print(f\"Epoch {epoch+1} Summary:\")\n        print(f\"  [Train] Loss: {train_loss:.4f} | Acc: {train_acc:.4f} | F1: {train_f1:.4f} | Kappa: {train_kappa:.4f}\")\n        print(f\"  [Valid] ---     | Acc: {val_acc:.4f} | F1: {val_f1:.4f} | Kappa: {val_kappa:.4f}\")\n        \n        scheduler.step(val_f1)\n        \n        if val_f1 > best_f1:\n            best_f1 = val_f1\n            torch.save(model.state_dict(), 'final_hybrid_model.pth')\n            print(\"  [+] Best Model Saved\")\n\n# Run the complete pipeline\nrun_full_training()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_full_training():\n\n    import matplotlib.pyplot as plt\n    from sklearn.metrics import confusion_matrix, classification_report\n    from sklearn.metrics import precision_score, recall_score, f1_score\n    from sklearn.metrics import roc_curve, auc\n    from sklearn.preprocessing import label_binarize\n    import seaborn as sns\n\n    # --- Transforms ---\n    train_transforms = transforms.Compose([\n        transforms.RandomRotation(180),\n        transforms.RandomHorizontalFlip(),\n        transforms.RandomVerticalFlip(),\n        transforms.ColorJitter(brightness=0.15, contrast=0.15),\n        transforms.RandomResizedCrop(CONFIG['img_size'], scale=(0.85, 1.0)),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n    ])\n\n    val_transforms = transforms.Compose([\n        transforms.Resize((CONFIG['img_size'],CONFIG['img_size'])),\n        transforms.ToTensor(),\n        transforms.Normalize([0.485,0.456,0.406],[0.229,0.224,0.225])\n    ])\n\n    model = HybridModel(num_classes=5).to(CONFIG['device'])\n    criterion = nn.CrossEntropyLoss()\n    scaler = torch.amp.GradScaler('cuda')\n\n    df_2019 = pd.read_csv(LABEL_2019_PATH)\n\n    train_sub, val_sub = train_test_split(\n        df_2019,\n        test_size=0.2,\n        stratify=df_2019['diagnosis'],\n        random_state=CONFIG['seed']\n    )\n\n    train_sub = train_sub.reset_index(drop=True)\n    val_sub = val_sub.reset_index(drop=True)\n\n    ds_train = DRDataset(train_sub, DATA_2019_PATH, extension=\".png\", transform=train_transforms)\n    ds_val   = DRDataset(val_sub, DATA_2019_PATH, extension=\".png\", transform=val_transforms)\n\n    dl_train = DataLoader(ds_train,batch_size=CONFIG['batch_size'],shuffle=True)\n    dl_val   = DataLoader(ds_val,batch_size=CONFIG['batch_size'],shuffle=False)\n\n    optimizer = optim.AdamW(model.parameters(), lr=1e-4)\n\n    train_losses=[]\n    val_losses=[]\n    train_accs=[]\n    val_accs=[]\n\n    best_f1=0\n\n    for epoch in range(30):\n\n        model.train()\n        running_loss=0\n        train_preds=[]\n        train_targets=[]\n\n        for images,labels in tqdm(dl_train):\n\n            images=images.to(CONFIG['device'])\n            labels=labels.to(CONFIG['device'])\n\n            optimizer.zero_grad()\n\n            with torch.amp.autocast('cuda'):\n                outputs=model(images)\n                loss=criterion(outputs,labels)\n\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n            running_loss+=loss.item()\n\n            _,pred=torch.max(outputs,1)\n            train_preds.extend(pred.cpu().numpy())\n            train_targets.extend(labels.cpu().numpy())\n\n        train_loss=running_loss/len(dl_train)\n        train_acc=accuracy_score(train_targets,train_preds)\n\n        train_losses.append(train_loss)\n        train_accs.append(train_acc)\n\n        # ----- Validation -----\n\n        model.eval()\n\n        val_preds=[]\n        val_targets=[]\n        val_running_loss=0\n\n        with torch.no_grad():\n\n            for images,labels in dl_val:\n\n                images=images.to(CONFIG['device'])\n                labels=labels.to(CONFIG['device'])\n\n                outputs=model(images)\n                loss=criterion(outputs,labels)\n\n                val_running_loss+=loss.item()\n\n                _,pred=torch.max(outputs,1)\n\n                val_preds.extend(pred.cpu().numpy())\n                val_targets.extend(labels.cpu().numpy())\n\n        val_loss=val_running_loss/len(dl_val)\n        val_acc=accuracy_score(val_targets,val_preds)\n\n        val_losses.append(val_loss)\n        val_accs.append(val_acc)\n\n        val_f1=f1_score(val_targets,val_preds,average=\"weighted\")\n\n        print(f\"\\nEpoch {epoch+1}\")\n        print(\"Train Loss:\",train_loss,\" Train Acc:\",train_acc)\n        print(\"Val Loss:\",val_loss,\" Val Acc:\",val_acc)\n\n        if val_f1>best_f1:\n            best_f1=val_f1\n            torch.save(model.state_dict(),\"best_model.pth\")\n            print(\"Best model saved\")\n\n    # =============================\n    # Plot Accuracy / Loss Graphs\n    # =============================\n\n    plt.figure(figsize=(12,5))\n\n    plt.subplot(1,2,1)\n    plt.plot(train_accs,label=\"Train Acc\")\n    plt.plot(val_accs,label=\"Val Acc\")\n    plt.title(\"Accuracy Curve\")\n    plt.legend()\n\n    plt.subplot(1,2,2)\n    plt.plot(train_losses,label=\"Train Loss\")\n    plt.plot(val_losses,label=\"Val Loss\")\n    plt.title(\"Loss Curve\")\n    plt.legend()\n\n    plt.show()\n\n    # =============================\n    # Confusion Matrix\n    # =============================\n\n    cm=confusion_matrix(val_targets,val_preds)\n\n    plt.figure(figsize=(6,6))\n    sns.heatmap(cm,annot=True,cmap=\"Blues\")\n    plt.title(\"Confusion Matrix\")\n    plt.show()\n\n    # =============================\n    # Classification Report\n    # =============================\n\n    print(\"\\nClassification Report\\n\")\n    print(classification_report(val_targets,val_preds))\n\n    # =============================\n    # Precision Recall F1\n    # =============================\n\n    precision=precision_score(val_targets,val_preds,average=\"weighted\")\n    recall=recall_score(val_targets,val_preds,average=\"weighted\")\n    f1=f1_score(val_targets,val_preds,average=\"weighted\")\n\n    print(\"Precision:\",precision)\n    print(\"Recall:\",recall)\n    print(\"F1 Score:\",f1)\n\n    # =============================\n    # Sensitivity Specificity\n    # =============================\n\n    cm=confusion_matrix(val_targets,val_preds)\n\n    sensitivity=cm[1,1]/(cm[1,1]+cm[1,0])\n    specificity=cm[0,0]/(cm[0,0]+cm[0,1])\n\n    print(\"Sensitivity:\",sensitivity)\n    print(\"Specificity:\",specificity)\n\n    # =============================\n    # ROC Curve\n    # =============================\n\n    y_true=label_binarize(val_targets,classes=[0,1,2,3,4])\n    y_pred=label_binarize(val_preds,classes=[0,1,2,3,4])\n\n    fpr,tpr,_=roc_curve(y_true.ravel(),y_pred.ravel())\n    roc_auc=auc(fpr,tpr)\n\n    plt.figure()\n    plt.plot(fpr,tpr,label=\"ROC curve (AUC=%0.2f)\"%roc_auc)\n    plt.plot([0,1],[0,1],'k--')\n    plt.xlabel(\"False Positive Rate\")\n    plt.ylabel(\"True Positive Rate\")\n    plt.title(\"ROC Curve\")\n    plt.legend()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:59:02.686200Z","iopub.execute_input":"2026-03-08T07:59:02.686782Z","iopub.status.idle":"2026-03-08T07:59:02.704763Z","shell.execute_reply.started":"2026-03-08T07:59:02.686754Z","shell.execute_reply":"2026-03-08T07:59:02.704172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"run_full_training()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-08T07:59:09.910682Z","iopub.execute_input":"2026-03-08T07:59:09.911271Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}