{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.13","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":2822650,"datasetId":1715304,"databundleVersionId":2869088}],"dockerImageVersionId":31401,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"b8300798-d430-4df2-aa70-50a19675e5a4","cell_type":"markdown","source":"# 🔬 Diabetic Retinopathy Classification\n## Hybrid EfficientNetB4 + DeiT-Small (CNN-Transformer Fusion)\n**Dataset:** APTOS 2019 Blindness Detection  \n**Task:** 5-Class Severity Classification (0–4)  \n**Hardware:** Kaggle GPU T4 x2  \n**Strategy:** Feature-level fusion of CNN (local texture) + ViT (global context)","metadata":{}},{"id":"874bd123-d60b-4abf-8757-180382032bd7","cell_type":"markdown","source":"## 1. Install & Import Dependencies","metadata":{}},{"id":"622c4831-85ea-4e34-8aa7-04c0adea8fba","cell_type":"code","source":"# Install timm for DeiT and other transformer models\n!pip install timm -q  \n!pip install albumentations -q","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T10:40:55.216339Z","iopub.execute_input":"2026-05-30T10:40:55.216627Z","iopub.status.idle":"2026-05-30T10:41:05.491810Z","shell.execute_reply.started":"2026-05-30T10:40:55.216597Z","shell.execute_reply":"2026-05-30T10:41:05.490753Z"}},"outputs":[],"execution_count":null},{"id":"5845662d-c454-4a6d-ac6f-2b4a2d27659d","cell_type":"code","source":"import os\nimport random\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom PIL import Image\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import cohen_kappa_score, classification_report, confusion_matrix\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.amp import autocast, GradScaler\nimport torchvision.transforms as T\n\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# ── Reproducibility ──────────────────────────────────────────────────────────\nSEED = 42\ndef seed_everything(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\nseed_everything(SEED)\nprint('Libraries loaded successfully.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T10:41:50.306571Z","iopub.execute_input":"2026-05-30T10:41:50.307225Z","iopub.status.idle":"2026-05-30T10:41:50.316053Z","shell.execute_reply.started":"2026-05-30T10:41:50.307190Z","shell.execute_reply":"2026-05-30T10:41:50.315152Z"}},"outputs":[],"execution_count":null},{"id":"8289cf90-31c0-41ec-8233-057f10ae4536","cell_type":"markdown","source":"## 2. Configuration","metadata":{}},{"id":"7850ec41-7477-4787-b7e7-e35ad5b7f492","cell_type":"code","source":"class CFG:\n    # ── Paths ─────────────────────────────────────────────────────────────────\n    DATA_DIR        = Path('/kaggle/input/datasets/mariaherrerot/aptos2019')\n    TRAIN_CSV       = DATA_DIR / 'train_1.csv'\n    TRAIN_IMG_DIR   = DATA_DIR / 'train_images' / 'train_images'\n    TEST_CSV        = DATA_DIR / 'valid.csv'\n    TEST_IMG_DIR    = DATA_DIR / 'test_images'  / 'test_images'\n    VAL_IMG_DIR     = DATA_DIR / 'val_images'   / 'val_images'\n    OUTPUT_DIR      = Path('/kaggle/working')\n\n    # ── Model ─────────────────────────────────────────────────────────────────\n    CNN_MODEL       = 'efficientnet_b4'       # timm model name\n    VIT_MODEL       = 'deit_small_patch16_224' # timm model name\n    NUM_CLASSES     = 5\n    FUSION_DIM      = 512                      # projection dim before fusion head\n    DROPOUT         = 0.3\n\n    # ── Training ──────────────────────────────────────────────────────────────\n    IMG_SIZE        = 224   # DeiT requires 224; EfficientNetB4 works well at 224\n    BATCH_SIZE      = 16    # per GPU — effective 32 with T4x2\n    EPOCHS          = 20\n    N_FOLDS         = 5\n    TRAIN_FOLD      = 0     # which fold to train (0-indexed)\n\n    # ── Optimizer / Scheduler ─────────────────────────────────────────────────\n    LR_CNN          = 1e-4\n    LR_VIT          = 5e-5  # lower LR for transformer (more sensitive)\n    LR_HEAD         = 3e-4\n    WEIGHT_DECAY    = 1e-2\n    WARMUP_EPOCHS   = 2\n\n    # ── Loss ──────────────────────────────────────────────────────────────────\n    LABEL_SMOOTHING = 0.1\n\n    # ── Hardware ──────────────────────────────────────────────────────────────\n    DEVICE          = 'cuda' if torch.cuda.is_available() else 'cpu'\n    NUM_WORKERS     = 4\n    USE_AMP         = True   # mixed precision\n\nprint(f'Device: {CFG.DEVICE}')\nprint(f'GPU count: {torch.cuda.device_count()}')\nif torch.cuda.is_available():\n    for i in range(torch.cuda.device_count()):\n        print(f'  GPU {i}: {torch.cuda.get_device_name(i)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T10:41:52.758767Z","iopub.execute_input":"2026-05-30T10:41:52.759611Z","iopub.status.idle":"2026-05-30T10:41:53.076110Z","shell.execute_reply.started":"2026-05-30T10:41:52.759556Z","shell.execute_reply":"2026-05-30T10:41:53.075236Z"}},"outputs":[],"execution_count":null},{"id":"f18afe93-47cd-45dd-9cfa-132930d6cf24","cell_type":"markdown","source":"## 3. Data Loading & Preprocessing","metadata":{}},{"id":"d0482260-8ae4-4375-ad92-884f01b7c248","cell_type":"code","source":"df = pd.read_csv(CFG.TRAIN_CSV)\nprint(f'Total samples: {len(df)}')\nprint(df['diagnosis'].value_counts().sort_index())\n\n# Stratified K-Fold split\nskf = StratifiedKFold(n_splits=CFG.N_FOLDS, shuffle=True, random_state=SEED)\nfor fold, (train_idx, val_idx) in enumerate(skf.split(df, df['diagnosis'])):\n    df.loc[val_idx, 'fold'] = fold\ndf['fold'] = df['fold'].astype(int)\n\n# Class distribution plot\nfig, ax = plt.subplots(1, 2, figsize=(12, 4))\ndf['diagnosis'].value_counts().sort_index().plot(\n    kind='bar', ax=ax[0], color='steelblue', edgecolor='black'\n)\nax[0].set_title('Class Distribution (APTOS 2019)')\nax[0].set_xlabel('DR Severity (0=No DR, 4=Proliferative DR)')\nax[0].set_ylabel('Count')\nax[0].tick_params(axis='x', rotation=0)\n\ndf['fold'].value_counts().sort_index().plot(\n    kind='bar', ax=ax[1], color='coral', edgecolor='black'\n)\nax[1].set_title('Fold Distribution')\nax[1].set_xlabel('Fold')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T10:41:55.239103Z","iopub.execute_input":"2026-05-30T10:41:55.239767Z","iopub.status.idle":"2026-05-30T10:41:55.958887Z","shell.execute_reply.started":"2026-05-30T10:41:55.239733Z","shell.execute_reply":"2026-05-30T10:41:55.957959Z"}},"outputs":[],"execution_count":null},{"id":"e8b1b7b1-d3ab-4bbc-9373-f9b58a884784","cell_type":"code","source":"# ── Ben Graham-style preprocessing (green channel enhancement) ───────────────\ndef ben_graham_preprocess(img: np.ndarray, sigmaX: int = 10) -> np.ndarray:\n    \"\"\"\n    Removes low-frequency lighting variations and enhances local\n    structure — standard preprocessing for fundus images.\n    \"\"\"\n    import cv2\n    img = cv2.addWeighted(\n        img, 4,\n        cv2.GaussianBlur(img, (0, 0), sigmaX), -4,\n        128\n    )\n    return img\n\n\n# ── Augmentation pipelines ────────────────────────────────────────────────────\ndef get_train_transforms():\n    return A.Compose([\n        A.Resize(CFG.IMG_SIZE, CFG.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,\n                           rotate_limit=30, p=0.6),\n        A.OneOf([\n            A.GridDistortion(p=0.5),\n            A.OpticalDistortion(p=0.5),\n            A.ElasticTransform(p=0.5),\n        ], p=0.3),\n        A.OneOf([\n            A.GaussNoise(p=0.5),\n            A.ISONoise(p=0.5),\n        ], p=0.3),\n        A.ColorJitter(brightness=0.2, contrast=0.2,\n                      saturation=0.2, hue=0.1, p=0.5),\n        A.CLAHE(clip_limit=4.0, tile_grid_size=(8, 8), p=0.4),\n        A.CoarseDropout(num_holes_range=(1, 8), hole_height_range=(1, 32), \n                        hole_width_range=(1, 32), fill=0, p=0.3),\n        A.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225]\n        ),\n        ToTensorV2(),\n    ])\n\ndef get_val_transforms():\n    return A.Compose([\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n        A.Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225]\n        ),\n        ToTensorV2(),\n    ])\n\nprint('Transforms defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T10:41:58.905027Z","iopub.execute_input":"2026-05-30T10:41:58.905503Z","iopub.status.idle":"2026-05-30T10:41:58.915348Z","shell.execute_reply.started":"2026-05-30T10:41:58.905470Z","shell.execute_reply":"2026-05-30T10:41:58.914425Z"}},"outputs":[],"execution_count":null},{"id":"f6837b28-2cc6-448c-ae79-c5e9f2071e65","cell_type":"code","source":"class APTOSDataset(Dataset):\n    def __init__(self, df, img_dir, transforms=None, apply_ben_graham=True):\n        self.df             = df.reset_index(drop=True)\n        self.img_dir        = Path(img_dir)\n        self.transforms     = transforms\n        self.apply_ben_graham = apply_ben_graham\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        import cv2\n        row      = self.df.iloc[idx]\n        img_path = self.img_dir / f\"{row['id_code']}.png\"\n        img      = cv2.imread(str(img_path))\n        if img is None:\n            raise FileNotFoundError(f'Image not found: {img_path}')\n        img      = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n\n        if self.apply_ben_graham:\n            img = ben_graham_preprocess(img)\n\n        if self.transforms:\n            img = self.transforms(image=img)['image']\n\n        label = torch.tensor(row['diagnosis'], dtype=torch.long)\n        return img, label\n\n\ndef get_dataloaders(df, fold):\n    train_df = df[df['fold'] != fold]\n    val_df   = df[df['fold'] == fold]\n\n    train_ds = APTOSDataset(train_df, CFG.TRAIN_IMG_DIR, get_train_transforms())\n    val_ds   = APTOSDataset(val_df,   CFG.TRAIN_IMG_DIR, get_val_transforms())\n\n    train_loader = DataLoader(\n        train_ds, batch_size=CFG.BATCH_SIZE, shuffle=True,\n        num_workers=CFG.NUM_WORKERS, pin_memory=True, drop_last=True\n    )\n    val_loader = DataLoader(\n        val_ds, batch_size=CFG.BATCH_SIZE * 2, shuffle=False,\n        num_workers=CFG.NUM_WORKERS, pin_memory=True\n    )\n    print(f'Fold {fold} | Train: {len(train_ds)} | Val: {len(val_ds)}')\n    return train_loader, val_loader\n\n\n# Quick sanity check\ntrain_loader, val_loader = get_dataloaders(df, CFG.TRAIN_FOLD)\nimgs, labels = next(iter(train_loader))\nprint(f'Batch shape: {imgs.shape} | Labels: {labels[:8].tolist()}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-30T10:42:01.559968Z","iopub.execute_input":"2026-05-30T10:42:01.560577Z","iopub.status.idle":"2026-05-30T10:42:16.745119Z","shell.execute_reply.started":"2026-05-30T10:42:01.560540Z","shell.execute_reply":"2026-05-30T10:42:16.743971Z"}},"outputs":[],"execution_count":null},{"id":"1e28d51a-df75-44d7-91a0-b76225c5fa86","cell_type":"markdown","source":"## 4. Hybrid Model: EfficientNetB4 + DeiT-Small","metadata":{}},{"id":"a91a40fc-e72e-4512-8e76-82b7e0b1f55a","cell_type":"code","source":"class HybridDRModel(nn.Module):\n    \"\"\"\n    Hybrid CNN-Transformer model for diabetic retinopathy classification.\n\n    Architecture:\n      ┌─────────────────────┐    ┌──────────────────────┐\n      │  EfficientNetB4      │    │  DeiT-Small           │\n      │  (local features)    │    │  (global context)     │\n      │  1792-d features     │    │  384-d [CLS] token    │\n      └────────┬────────────┘    └──────────┬───────────┘\n               │  Linear → 512-d            │  Linear → 512-d\n               └──────────────┬─────────────┘\n                          Concatenate\n                          1024-d fused\n                          MLP Head → 5 classes\n    \"\"\"\n\n    def __init__(self):\n        super().__init__()\n\n        # ── CNN Branch: EfficientNetB4 ────────────────────────────────────────\n        self.cnn = timm.create_model(\n            CFG.CNN_MODEL,\n            pretrained=True,\n            num_classes=0,        # remove classifier head\n            global_pool='avg'     # global average pooling\n        )\n        cnn_feat_dim = self.cnn.num_features  # 1792 for B4\n\n        # ── Transformer Branch: DeiT-Small ────────────────────────────────────\n        self.vit = timm.create_model(\n            CFG.VIT_MODEL,\n            pretrained=True,\n            num_classes=0,        # remove classifier head\n        )\n        vit_feat_dim = self.vit.num_features  # 384 for DeiT-Small\n\n        # ── Projection layers (align both branches to FUSION_DIM) ─────────────\n        self.cnn_proj = nn.Sequential(\n            nn.Linear(cnn_feat_dim, CFG.FUSION_DIM),\n            nn.BatchNorm1d(CFG.FUSION_DIM),\n            nn.GELU(),\n            nn.Dropout(CFG.DROPOUT)\n        )\n        self.vit_proj = nn.Sequential(\n            nn.Linear(vit_feat_dim, CFG.FUSION_DIM),\n            nn.BatchNorm1d(CFG.FUSION_DIM),\n            nn.GELU(),\n            nn.Dropout(CFG.DROPOUT)\n        )\n\n        # ── Fusion attention gate (learned weighting between branches) ─────────\n        self.attention_gate = nn.Sequential(\n            nn.Linear(CFG.FUSION_DIM * 2, 2),\n            nn.Softmax(dim=-1)\n        )\n\n        # ── Classifier head ───────────────────────────────────────────────────\n        self.head = nn.Sequential(\n            nn.Linear(CFG.FUSION_DIM * 2, 512),\n            nn.BatchNorm1d(512),\n            nn.GELU(),\n            nn.Dropout(CFG.DROPOUT),\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.GELU(),\n            nn.Dropout(CFG.DROPOUT / 2),\n            nn.Linear(256, CFG.NUM_CLASSES)\n        )\n\n    def forward(self, x):\n        # CNN features\n        cnn_feat = self.cnn(x)          # (B, 1792)\n        cnn_feat = self.cnn_proj(cnn_feat) # (B, 512)\n\n        # ViT features (DeiT returns CLS token)\n        vit_feat = self.vit(x)           # (B, 384)\n        vit_feat = self.vit_proj(vit_feat) # (B, 512)\n\n        # Concatenate\n        fused = torch.cat([cnn_feat, vit_feat], dim=-1)  # (B, 1024)\n\n        # Classify\n        logits = self.head(fused)  # (B, 5)\n        return logits\n\n\n# ── Test forward pass ─────────────────────────────────────────────────────────\nmodel = HybridDRModel()\ndummy = torch.randn(2, 3, CFG.IMG_SIZE, CFG.IMG_SIZE)\nout   = model(dummy)\nprint(f'Model output shape: {out.shape}')  # Should be (2, 5)\n\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable    = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'Total params:     {total_params:,}')\nprint(f'Trainable params: {trainable:,}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"45982c8d-4786-4c57-8669-e59480992f63","cell_type":"markdown","source":"## 5. Loss, Optimizer & Scheduler","metadata":{}},{"id":"044db2cd-b829-4e28-86e8-c5ea5a04d35e","cell_type":"code","source":"# ── Class-weighted CrossEntropy + Label Smoothing ─────────────────────────────\ndef get_class_weights(df, fold):\n    \"\"\"Inverse-frequency weights to handle class imbalance.\"\"\"\n    train_df = df[df['fold'] != fold]\n    counts   = train_df['diagnosis'].value_counts().sort_index().values\n    weights  = 1.0 / counts\n    weights  = weights / weights.sum() * CFG.NUM_CLASSES\n    return torch.tensor(weights, dtype=torch.float)\n\n\nclass SmoothCrossEntropy(nn.Module):\n    def __init__(self, smoothing=0.1, weight=None):\n        super().__init__()\n        self.smoothing = smoothing\n        self.weight    = weight\n\n    def forward(self, logits, targets):\n        n_classes = logits.size(-1)\n        log_probs = F.log_softmax(logits, dim=-1)\n\n        # Smooth targets\n        with torch.no_grad():\n            smooth_targets = torch.full_like(log_probs, self.smoothing / (n_classes - 1))\n            smooth_targets.scatter_(1, targets.unsqueeze(1), 1.0 - self.smoothing)\n\n        loss = -(smooth_targets * log_probs)\n\n        if self.weight is not None:\n            loss = loss * self.weight.to(logits.device).unsqueeze(0)\n\n        return loss.sum(dim=-1).mean()\n\n\ndef build_optimizer(model):\n    \"\"\"\n    Layer-wise learning rates:\n      - CNN backbone  → LR_CNN  (1e-4)\n      - ViT backbone  → LR_VIT  (5e-5)\n      - Projection + head → LR_HEAD (3e-4)\n    \"\"\"\n    param_groups = [\n        {'params': model.cnn.parameters(),      'lr': CFG.LR_CNN},\n        {'params': model.vit.parameters(),      'lr': CFG.LR_VIT},\n        {'params': list(model.cnn_proj.parameters()) +\n                   list(model.vit_proj.parameters()) +\n                   list(model.attention_gate.parameters()) +\n                   list(model.head.parameters()),\n         'lr': CFG.LR_HEAD}\n    ]\n    return torch.optim.AdamW(param_groups, weight_decay=CFG.WEIGHT_DECAY)\n\n\ndef build_scheduler(optimizer, total_steps):\n    \"\"\"Cosine annealing with linear warmup.\"\"\"\n    from torch.optim.lr_scheduler import OneCycleLR\n    return OneCycleLR(\n        optimizer,\n        max_lr=[CFG.LR_CNN, CFG.LR_VIT, CFG.LR_HEAD],\n        total_steps=total_steps,\n        pct_start=CFG.WARMUP_EPOCHS / CFG.EPOCHS,\n        anneal_strategy='cos',\n        div_factor=10,\n        final_div_factor=100\n    )\n\n\nprint('Loss, optimizer and scheduler functions ready.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"f085fa84-5fed-4b97-b056-53b6c7286587","cell_type":"markdown","source":"## 6. Training & Validation Loop","metadata":{}},{"id":"d73b59e0-f9d4-40cf-995e-353bf5756f15","cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, criterion, scaler, scheduler):\n    model.train()\n    total_loss, correct, total = 0.0, 0, 0\n\n    for batch_idx, (imgs, labels) in enumerate(loader):\n        imgs, labels = imgs.to(CFG.DEVICE), labels.to(CFG.DEVICE)\n\n        optimizer.zero_grad()\n\n        with autocast('cuda', enabled=CFG.USE_AMP):\n            logits = model(imgs)\n            loss   = criterion(logits, labels)\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        old_scale = scaler.get_scale()\n        scaler.step(optimizer)\n        scaler.update()\n        if scaler.get_scale() >= old_scale:\n            scheduler.step()\n\n        preds       = logits.argmax(dim=1)\n        correct    += (preds == labels).sum().item()\n        total      += labels.size(0)\n        total_loss += loss.item() * labels.size(0)\n\n    return total_loss / total, correct / total\n\n\n@torch.no_grad()\ndef validate(model, loader, criterion):\n    model.eval()\n    total_loss, correct, total = 0.0, 0, 0\n    all_preds, all_labels = [], []\n\n    for imgs, labels in loader:\n        imgs, labels = imgs.to(CFG.DEVICE), labels.to(CFG.DEVICE)\n\n        with autocast('cuda', enabled=CFG.USE_AMP):\n            logits = model(imgs)\n            loss   = criterion(logits, labels)\n\n        preds       = logits.argmax(dim=1)\n        correct    += (preds == labels).sum().item()\n        total      += labels.size(0)\n        total_loss += loss.item() * labels.size(0)\n\n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\n    kappa = cohen_kappa_score(all_labels, all_preds, weights='quadratic')\n    return total_loss / total, correct / total, kappa, all_preds, all_labels\n\n\nprint('Training loop functions ready.')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"eae3af9a-38a7-4d7b-a200-50551ab61b7d","cell_type":"code","source":"def train_fold(df, fold):\n    print(f'\\n{\"=\"*60}')\n    print(f'  TRAINING FOLD {fold}')\n    print(f'{\"=\"*60}')\n\n    # ── Data ──────────────────────────────────────────────────────────────────\n    train_loader, val_loader = get_dataloaders(df, fold)\n\n    # ── Model ─────────────────────────────────────────────────────────────────\n    model = HybridDRModel().to(CFG.DEVICE)\n\n    # Multi-GPU support for T4x2\n    if torch.cuda.device_count() > 1:\n        print(f'Using {torch.cuda.device_count()} GPUs with DataParallel')\n        model = nn.DataParallel(model)\n\n    # ── Loss / Optimizer / Scheduler ──────────────────────────────────────────\n    class_weights = get_class_weights(df, fold)\n    criterion     = SmoothCrossEntropy(\n        smoothing=CFG.LABEL_SMOOTHING,\n        weight=class_weights\n    )\n    optimizer     = build_optimizer(model.module if hasattr(model, 'module') else model)\n    total_steps   = len(train_loader) * CFG.EPOCHS\n    scheduler     = build_scheduler(optimizer, total_steps)\n    scaler        = GradScaler('cuda', enabled=CFG.USE_AMP)\n\n    # ── Training loop ─────────────────────────────────────────────────────────\n    best_kappa    = -1.0\n    history       = {'train_loss': [], 'val_loss': [], 'val_kappa': [], 'val_acc': []}\n    save_path     = CFG.OUTPUT_DIR / f'hybrid_dr_fold{fold}_best.pth'\n\n    for epoch in range(CFG.EPOCHS):\n        train_loss, train_acc = train_one_epoch(\n            model, train_loader, optimizer, criterion, scaler, scheduler\n        )\n        val_loss, val_acc, kappa, val_preds, val_labels = validate(\n            model, val_loader, criterion\n        )\n\n        history['train_loss'].append(train_loss)\n        history['val_loss'].append(val_loss)\n        history['val_kappa'].append(kappa)\n        history['val_acc'].append(val_acc)\n\n        print(f'Epoch [{epoch+1:02d}/{CFG.EPOCHS}] '\n              f'Train Loss: {train_loss:.4f} | Acc: {train_acc:.4f} || '\n              f'Val Loss: {val_loss:.4f} | Acc: {val_acc:.4f} | '\n              f'Kappa: {kappa:.4f}', end='')\n\n        if kappa > best_kappa:\n            best_kappa = kappa\n            torch.save(model.state_dict(), save_path)\n            print(f'  ✅ Saved (best kappa: {best_kappa:.4f})')\n        else:\n            print()\n\n    print(f'\\nFold {fold} — Best Quadratic Kappa: {best_kappa:.4f}')\n    return history, best_kappa, val_preds, val_labels, save_path\n\n\n# ── Run training ──────────────────────────────────────────────────────────────\nhistory, best_kappa, val_preds, val_labels, save_path = train_fold(df, CFG.TRAIN_FOLD)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"5bfe3e0b-7907-4450-b111-a383ae4c3507","cell_type":"markdown","source":"## 7. Training Visualizations","metadata":{}},{"id":"fc502815-f14b-4d2c-bc08-233d822f8cf9","cell_type":"code","source":"fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n# Loss curve\naxes[0].plot(history['train_loss'], label='Train Loss', color='steelblue')\naxes[0].plot(history['val_loss'],   label='Val Loss',   color='coral')\naxes[0].set_title('Loss Curve')\naxes[0].set_xlabel('Epoch')\naxes[0].set_ylabel('Loss')\naxes[0].legend()\naxes[0].grid(alpha=0.3)\n\n# Kappa curve\naxes[1].plot(history['val_kappa'], label='Val Quadratic Kappa',\n             color='green', marker='o', markersize=4)\naxes[1].axhline(y=best_kappa, color='red', linestyle='--',\n                label=f'Best: {best_kappa:.4f}')\naxes[1].set_title('Quadratic Weighted Kappa')\naxes[1].set_xlabel('Epoch')\naxes[1].set_ylabel('Kappa')\naxes[1].legend()\naxes[1].grid(alpha=0.3)\n\n# Confusion matrix\ncm = confusion_matrix(val_labels, val_preds)\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', ax=axes[2],\n            xticklabels=[f'Grade {i}' for i in range(5)],\n            yticklabels=[f'Grade {i}' for i in range(5)])\naxes[2].set_title('Confusion Matrix (Best Val Epoch)')\naxes[2].set_ylabel('True Label')\naxes[2].set_xlabel('Predicted Label')\n\nplt.suptitle(f'Fold {CFG.TRAIN_FOLD} — Hybrid EfficientNetB4 + DeiT-Small',\n             fontsize=14, y=1.02)\nplt.tight_layout()\nplt.savefig(CFG.OUTPUT_DIR / 'training_results.png', dpi=150, bbox_inches='tight')\nplt.show()\n\nprint('\\nClassification Report:')\nprint(classification_report(\n    val_labels, val_preds,\n    target_names=[f'Grade {i}' for i in range(5)]\n))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"2d4c609b-07be-406c-9165-323d97403ca2","cell_type":"markdown","source":"## 8. Inference & Submission","metadata":{}},{"id":"e12c8d70-d64f-4d6e-b81e-e630f91ab7fe","cell_type":"code","source":"class APTOSTestDataset(Dataset):\n    def __init__(self, df, img_dir, transforms=None):\n        self.df         = df.reset_index(drop=True)\n        self.img_dir    = Path(img_dir)\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        import cv2\n        row      = self.df.iloc[idx]\n        img_path = self.img_dir / f\"{row['id_code']}.png\"\n        img      = cv2.imread(str(img_path))\n        img      = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        img      = ben_graham_preprocess(img)\n        if self.transforms:\n            img  = self.transforms(image=img)['image']\n        return img\n\n\ndef get_tta_transforms():\n    return A.Compose([\n        A.Resize(CFG.IMG_SIZE, CFG.IMG_SIZE),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.RandomRotate90(p=0.5),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ToTensorV2(),\n    ])\n\n\ndef tta_predict(model, test_df, n_tta=5):\n    \"\"\"\n    Test-Time Augmentation: rebuild a fresh loader with random augmentations\n    each pass so every iteration actually sees different views of the images.\n    \"\"\"\n    model.eval()\n    all_probs = []\n\n    for _ in range(n_tta):\n        ds = APTOSTestDataset(test_df, CFG.TEST_IMG_DIR, get_tta_transforms())\n        loader = DataLoader(\n            ds, batch_size=CFG.BATCH_SIZE * 2,\n            shuffle=False, num_workers=CFG.NUM_WORKERS, pin_memory=True\n        )\n        probs_tta = []\n        with torch.no_grad():\n            for imgs in loader:\n                imgs = imgs.to(CFG.DEVICE)\n                with autocast('cuda', enabled=CFG.USE_AMP):\n                    logits = model(imgs)\n                probs_tta.append(F.softmax(logits, dim=-1).cpu().numpy())\n        all_probs.append(np.concatenate(probs_tta, axis=0))\n\n    return np.mean(all_probs, axis=0)  # average over TTA runs\n\n\ndef run_inference():\n    test_df = pd.read_csv(CFG.TEST_CSV)\n\n    # Load best model\n    model = HybridDRModel().to(CFG.DEVICE)\n    if torch.cuda.device_count() > 1:\n        model = nn.DataParallel(model)\n    model.load_state_dict(torch.load(save_path))\n\n    # TTA inference — pass test_df so each pass gets fresh augmented views\n    print('Running TTA inference...')\n    probs = tta_predict(model, test_df, n_tta=5)\n    preds = probs.argmax(axis=1)\n\n    # Save submission\n    submission = pd.DataFrame({\n        'id_code':   test_df['id_code'],\n        'diagnosis': preds\n    })\n    sub_path = CFG.OUTPUT_DIR / 'submission.csv'\n    submission.to_csv(sub_path, index=False)\n    print(f'Submission saved to {sub_path}')\n    print(submission['diagnosis'].value_counts().sort_index())\n    return submission\n\n\nsubmission = run_inference()\nsubmission.head()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"id":"873dde85-25ac-404a-b36b-8884c38431dd","cell_type":"markdown","source":"## 9. (Optional) Full Cross-Validation","metadata":{}},{"id":"979eb9be-5206-4fed-8509-eb89e5e885dd","cell_type":"code","source":"# ── Uncomment to train all 5 folds and ensemble ───────────────────────────────\n\n# fold_results = []\n# all_save_paths = []\n# \n# for fold in range(CFG.N_FOLDS):\n#     hist, kappa, preds, labels, spath = train_fold(df, fold)\n#     fold_results.append(kappa)\n#     all_save_paths.append(spath)\n#\n# print('\\n' + '='*40)\n# print('CROSS-VALIDATION RESULTS')\n# print('='*40)\n# for i, k in enumerate(fold_results):\n#     print(f'  Fold {i}: Kappa = {k:.4f}')\n# print(f'  Mean:   {np.mean(fold_results):.4f} ± {np.std(fold_results):.4f}')\n#\n# # Ensemble inference (average across all fold models)\n# # ... add ensemble inference loop using all_save_paths\n\nprint('Uncomment the block above to run full CV + ensemble.')","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}