{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# APTOS 2019 — CNN + Swin Transformer Hybrid\n**Architecture:** EfficientNet-B4 (local CNN features) + Swin Transformer (hierarchical global context)  \n**Dataset:** APTOS 2019 Blindness Detection  \n**Key design choices:**\n- EfficientNet-B4 backbone extracts multi-scale local features (microaneurysms, exudates, haemorrhages)\n- Swin Transformer encoder captures hierarchical global context via shifted window attention\n- Feature fusion via a learned gating mechanism before classification head\n- Same preprocessing, augmentation, splits, and training recipe as baseline for fair comparison\n- Expected QWK target: > 0.88 (baseline VGG-19-BN = 0.8813)","metadata":{}},{"cell_type":"markdown","source":"## Cell 1 — Install dependencies","metadata":{}},{"cell_type":"code","source":"import subprocess, sys\n\nsubprocess.run([\n    sys.executable, \"-m\", \"pip\", \"install\", \"-q\",\n    \"timm==1.0.3\",\n    \"albumentations==1.4.3\",\n], check=True)\n\nprint(\"Dependencies installed ✓\")\nprint(\"Note: Swin Transformer is implemented in pure PyTorch — no extra CUDA extensions needed.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:34:58.936109Z","iopub.execute_input":"2026-04-12T04:34:58.936443Z","iopub.status.idle":"2026-04-12T04:35:05.90348Z","shell.execute_reply.started":"2026-04-12T04:34:58.936419Z","shell.execute_reply":"2026-04-12T04:35:05.902753Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 2 — Imports","metadata":{}},{"cell_type":"code","source":"import os, random, time, json, math, warnings\nfrom pathlib import Path\nfrom copy import deepcopy\nfrom functools import partial\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\n\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import (\n    cohen_kappa_score, roc_auc_score,\n    classification_report, confusion_matrix\n)\nfrom sklearn.preprocessing import label_binarize\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nwarnings.filterwarnings('ignore')\nSEED = 42\n\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    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(SEED)\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'Device: {DEVICE}')\nif torch.cuda.is_available():\n    print(f'GPU   : {torch.cuda.get_device_name(0)}')\n    print(f'VRAM  : {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:35:05.90443Z","iopub.execute_input":"2026-04-12T04:35:05.904743Z","iopub.status.idle":"2026-04-12T04:35:18.963968Z","shell.execute_reply.started":"2026-04-12T04:35:05.904719Z","shell.execute_reply":"2026-04-12T04:35:18.963157Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 3 — Configuration","metadata":{}},{"cell_type":"code","source":"CFG = dict(\n    # Paths (same as baseline)\n    data_dir   = '/kaggle/input/competitions/aptos2019-blindness-detection',\n    output_dir = '/kaggle/working/outputs',\n\n    # Image (same as baseline for fair comparison)\n    img_size       = 384,\n    use_ben_graham = True,\n\n    # Split — SAME fold as baseline\n    n_folds    = 5,\n    train_fold = 0,\n\n    # Training\n    epochs        = 25,\n    batch_size    = 8,    # same as baseline\n    accum_steps   = 4,    # effective batch = 32\n    lr            = 2e-4,\n    backbone_lr_scale = 0.1,   # CNN backbone lr\n    swin_lr_scale     = 0.3,   # Swin Transformer lr\n    min_lr        = 1e-6,\n    weight_decay  = 1e-4,\n    warmup_epochs = 3,\n    label_smooth  = 0.05,\n    patience      = 8,\n\n    # Task\n    num_classes = 5,\n    drop_rate   = 0.3,\n    amp         = True,\n\n    # Hybrid architecture\n    cnn_backbone    = 'efficientnet_b4',\n    cnn_stage3_ch   = 160,   # EfficientNet-B4 stage-3 channels @ 24x24 (384 input)\n    cnn_global_ch   = 448,   # EfficientNet-B4 stage-4 channels @ 12x12\n\n    # Swin Transformer encoder config\n    swin_embed_dim  = 96,    # base embed dim (matches Swin-T)\n    swin_depths     = [2, 2],  # 2 stages x 2 blocks each\n    swin_num_heads  = [3, 6],  # heads per stage (must divide embed dims)\n    swin_window_size= 6,     # window size for 24x24 feature map (24 divisible by 6)\n    swin_mlp_ratio  = 4.0,\n    fusion_dim      = 512,\n)\n\nos.makedirs(CFG['output_dir'], exist_ok=True)\nOUT_DIR = Path(CFG['output_dir'])\n\nGRADE_NAMES  = {0:'No DR', 1:'Mild', 2:'Moderate', 3:'Severe', 4:'Proliferative'}\nGRADE_COLORS = ['#4CAF50','#FFC107','#FF9800','#F44336','#9C27B0']\n\nprint('Config loaded ✓')\nprint(f'CNN backbone   : {CFG[\"cnn_backbone\"]}')\nprint(f'Swin stages    : {CFG[\"swin_depths\"]}  heads={CFG[\"swin_num_heads\"]}')\nprint(f'Window size    : {CFG[\"swin_window_size\"]}x{CFG[\"swin_window_size\"]}  spatial=24x24')\nprint(f'Eff. batch     : {CFG[\"batch_size\"] * CFG[\"accum_steps\"]}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:35:18.965727Z","iopub.execute_input":"2026-04-12T04:35:18.966584Z","iopub.status.idle":"2026-04-12T04:35:18.974438Z","shell.execute_reply.started":"2026-04-12T04:35:18.966558Z","shell.execute_reply":"2026-04-12T04:35:18.973613Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/input","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:43:54.335551Z","iopub.execute_input":"2026-04-12T04:43:54.336374Z","iopub.status.idle":"2026-04-12T04:43:54.477256Z","shell.execute_reply.started":"2026-04-12T04:43:54.336343Z","shell.execute_reply":"2026-04-12T04:43:54.47631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls /kaggle/input/competitions/aptos2019-blindness-detection","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:46:20.471468Z","iopub.execute_input":"2026-04-12T04:46:20.47201Z","iopub.status.idle":"2026-04-12T04:46:20.615458Z","shell.execute_reply.started":"2026-04-12T04:46:20.471974Z","shell.execute_reply":"2026-04-12T04:46:20.614828Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 4 — Load Data","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nimport pandas as pd\n\nCFG['data_dir'] = '/kaggle/input/competitions/aptos2019-blindness-detection'\n\nDATA_DIR = Path(CFG['data_dir'])\n\ndf = pd.read_csv(DATA_DIR / 'train.csv')\n\ndf['path'] = df['id_code'].apply(\n    lambda x: str(DATA_DIR / 'train_images' / f'{x}.png')\n)\n\ndf.rename(columns={'diagnosis': 'label'}, inplace=True)\n\nprint(f'Total images : {len(df)}')\n\ncounts = df['label'].value_counts().sort_index()\nprint('\\nLabel distribution:')\n\nfor g, n in counts.items():\n    pct = n / len(df) * 100\n    print(f'  Grade {g}: {n} ({pct:.1f}%)')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:47:40.543953Z","iopub.execute_input":"2026-04-12T04:47:40.544437Z","iopub.status.idle":"2026-04-12T04:47:40.613712Z","shell.execute_reply.started":"2026-04-12T04:47:40.544401Z","shell.execute_reply":"2026-04-12T04:47:40.613053Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 5 — Preprocessing (same as baseline)","metadata":{}},{"cell_type":"code","source":"def circle_crop(img: np.ndarray) -> np.ndarray:\n    h, w = img.shape[:2]\n    mask = np.zeros((h, w), dtype=np.uint8)\n    cv2.circle(mask, (w//2, h//2), min(h, w)//2, 255, -1)\n    return cv2.bitwise_and(img, img, mask=mask)\n\n\ndef ben_graham(img: np.ndarray, size: int) -> np.ndarray:\n    img = cv2.resize(img, (size, size))\n    blurred = cv2.GaussianBlur(img, (0, 0), size / 30)\n    img = cv2.addWeighted(img, 4, blurred, -4, 128)\n    return img\n\n\ndef load_image(path: str, size: int, use_bg: bool = True) -> np.ndarray:\n    img = cv2.imread(str(path))\n    if img is None:\n        raise FileNotFoundError(f'Cannot read: {path}')\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    img = circle_crop(img)\n    if use_bg:\n        img = ben_graham(img, size)\n    else:\n        img = cv2.resize(img, (size, size))\n    return img\n\n\nprint('Preprocessing functions defined ✓')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:48:42.925298Z","iopub.execute_input":"2026-04-12T04:48:42.925989Z","iopub.status.idle":"2026-04-12T04:48:42.933319Z","shell.execute_reply.started":"2026-04-12T04:48:42.925958Z","shell.execute_reply":"2026-04-12T04:48:42.932537Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 6 — Stratified Splits (identical fold to baseline)","metadata":{}},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=CFG['n_folds'], shuffle=True, random_state=SEED)\ndf['fold'] = -1\nfor fi, (_, vi) in enumerate(skf.split(df, df['label'])):\n    df.loc[vi, 'fold'] = fi\n\nTRAIN_DF = df[df['fold'] != CFG['train_fold']].reset_index(drop=True)\nVAL_DF   = df[df['fold'] == CFG['train_fold']].reset_index(drop=True)\n\nprint(f'Train: {len(TRAIN_DF)}  |  Val: {len(VAL_DF)}')\nprint(f'Ratio: {len(TRAIN_DF)/len(df)*100:.0f}% / {len(VAL_DF)/len(df)*100:.0f}%\\n')\nprint(f'{\"Grade\":<15} {\"Train\":>7} {\"Train%\":>8} {\"Val\":>7} {\"Val%\":>7}')\nprint('-'*47)\nfor g in range(5):\n    tn = (TRAIN_DF['label']==g).sum()\n    vn = (VAL_DF['label']==g).sum()\n    print(f'{GRADE_NAMES[g]:<15} {tn:>7} {tn/len(TRAIN_DF)*100:>7.1f}%'\n          f' {vn:>7} {vn/len(VAL_DF)*100:>6.1f}%')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:49:11.946233Z","iopub.execute_input":"2026-04-12T04:49:11.946559Z","iopub.status.idle":"2026-04-12T04:49:11.97142Z","shell.execute_reply.started":"2026-04-12T04:49:11.946531Z","shell.execute_reply":"2026-04-12T04:49:11.970834Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 7 — Dataset & Augmentation","metadata":{}},{"cell_type":"code","source":"MEAN = [0.485, 0.456, 0.406]\nSTD  = [0.229, 0.224, 0.225]\n\n\ndef get_train_transforms(size: int) -> A.Compose:\n    return A.Compose([\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.2),\n        A.RandomRotate90(p=0.5),\n        A.ShiftScaleRotate(\n            shift_limit=0.1, scale_limit=0.15,\n            rotate_limit=30,\n            border_mode=cv2.BORDER_REFLECT, p=0.6\n        ),\n        A.OneOf([\n            A.CLAHE(clip_limit=4.0, tile_grid_size=(8, 8), p=1.0),\n            A.RandomBrightnessContrast(\n                brightness_limit=0.2, contrast_limit=0.2, p=1.0),\n            A.HueSaturationValue(\n                hue_shift_limit=10, sat_shift_limit=20,\n                val_shift_limit=10, p=1.0)\n        ], p=0.7),\n        A.OneOf([\n            A.GaussNoise(p=1.0),\n            A.GaussianBlur(blur_limit=(3, 5), p=1.0),\n        ], p=0.3),\n        A.CoarseDropout(\n            max_holes=8, min_holes=4,\n            max_height=32, min_height=16,\n            max_width=32, min_width=16,\n            p=0.3\n        ),\n        A.Normalize(mean=MEAN, std=STD),\n        ToTensorV2(),\n    ])\n\n\ndef get_val_transforms(size: int) -> A.Compose:\n    return A.Compose([\n        A.Normalize(mean=MEAN, std=STD),\n        ToTensorV2(),\n    ])\n\n\nclass APTOSDataset(Dataset):\n    def __init__(self, df, cfg, transforms=None, is_train=True):\n        self.df         = df.reset_index(drop=True)\n        self.cfg        = cfg\n        self.transforms = transforms\n        self.is_train   = is_train\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row   = self.df.iloc[idx]\n        img   = load_image(row['path'], self.cfg['img_size'], self.cfg['use_ben_graham'])\n        label = int(row['label'])\n        if self.transforms:\n            img = self.transforms(image=img)['image']\n        return img, torch.tensor(label, dtype=torch.long)\n\n\ndef get_loaders(cfg):\n    # Class weights for WeightedRandomSampler\n    label_counts = TRAIN_DF['label'].value_counts().sort_index().values\n    class_weights = 1.0 / label_counts\n    sample_weights = class_weights[TRAIN_DF['label'].values]\n    sampler = WeightedRandomSampler(\n        weights=sample_weights.tolist(),\n        num_samples=len(TRAIN_DF),\n        replacement=True\n    )\n\n    train_ds = APTOSDataset(TRAIN_DF, cfg, get_train_transforms(cfg['img_size']), is_train=True)\n    val_ds   = APTOSDataset(VAL_DF,   cfg, get_val_transforms(cfg['img_size']),   is_train=False)\n\n    train_loader = DataLoader(\n        train_ds, batch_size=cfg['batch_size'], sampler=sampler,\n        num_workers=4, pin_memory=True, persistent_workers=True, drop_last=True\n    )\n    val_loader = DataLoader(\n        val_ds, batch_size=cfg['batch_size']*2, shuffle=False,\n        num_workers=4, pin_memory=True, persistent_workers=True, drop_last=False\n    )\n    return train_loader, val_loader\n\n\n# Class weights tensor for loss\nlabel_counts_all = df['label'].value_counts().sort_index().values\nclass_w = torch.tensor(1.0 / label_counts_all, dtype=torch.float32)\nclass_w = class_w / class_w.sum()\nCLASS_WEIGHTS_T = class_w\n\nprint('Dataset & DataLoader functions defined ✓')\nprint(f'Class weights: {CLASS_WEIGHTS_T.numpy().round(3)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:50:10.827167Z","iopub.execute_input":"2026-04-12T04:50:10.827875Z","iopub.status.idle":"2026-04-12T04:50:10.842957Z","shell.execute_reply.started":"2026-04-12T04:50:10.827844Z","shell.execute_reply":"2026-04-12T04:50:10.842229Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 8 — Swin Transformer Components\n\nWe implement a lightweight Swin Transformer encoder that operates on the 24×24 feature map\nextracted from EfficientNet-B4's stage-3 output.  \nKey properties:\n- **Window-based self-attention**: divides 24×24 map into non-overlapping 6×6 windows → 16 windows, each with 36 tokens.\n- **Shifted windows**: alternating layers use a cyclic shift to capture cross-window context.\n- **Relative position bias**: learned per head for each (Δrow, Δcol) pair within a window.\n- **Patch merging**: between stages, halves spatial size and doubles channels (24→12, 96→192).\n- Pure PyTorch — no external CUDA extensions required.","metadata":{}},{"cell_type":"code","source":"import math as _math\n\n# ── Window partitioning helpers ───────────────────────────────────────────────\ndef window_partition(x, window_size):\n    \"\"\"\n    Partition feature map into non-overlapping windows.\n    x: (B, H, W, C)  →  (num_windows*B, window_size, window_size, C)\n    \"\"\"\n    B, H, W, C = x.shape\n    x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)\n    windows = x.permute(0, 1, 3, 2, 4, 5).contiguous()\n    return windows.view(-1, window_size, window_size, C)\n\n\ndef window_reverse(windows, window_size, H, W):\n    \"\"\"\n    Reconstruct feature map from windows.\n    windows: (num_windows*B, window_size, window_size, C)  →  (B, H, W, C)\n    \"\"\"\n    B_times_nw = windows.shape[0]\n    nw_h = H // window_size\n    nw_w = W // window_size\n    B = B_times_nw // (nw_h * nw_w)\n    x = windows.view(B, nw_h, nw_w, window_size, window_size, -1)\n    return x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)\n\n\n# ── Window Multi-Head Self-Attention ──────────────────────────────────────────\nclass WindowAttention(nn.Module):\n    \"\"\"\n    Window-based multi-head self-attention with relative position bias.\n    Supports both regular windows and shifted windows (via attn_mask).\n    \"\"\"\n    def __init__(self, dim, window_size, num_heads, qkv_bias=True, attn_drop=0., proj_drop=0.):\n        super().__init__()\n        self.dim         = dim\n        self.window_size = window_size   # Wh, Ww\n        self.num_heads   = num_heads\n        head_dim         = dim // num_heads\n        self.scale       = head_dim ** -0.5\n\n        # Relative position bias table: (2*Wh-1) * (2*Ww-1), num_heads\n        self.relative_position_bias_table = nn.Parameter(\n            torch.zeros((2 * window_size - 1) ** 2, num_heads)\n        )\n        nn.init.trunc_normal_(self.relative_position_bias_table, std=0.02)\n\n        # Pre-compute relative position index for each token pair in a window\n        coords_h = torch.arange(window_size)\n        coords_w = torch.arange(window_size)\n        coords   = torch.stack(torch.meshgrid(coords_h, coords_w, indexing='ij'))  # 2, Wh, Ww\n        coords_flat = torch.flatten(coords, 1)   # 2, Wh*Ww\n        rel_coords = coords_flat[:, :, None] - coords_flat[:, None, :]  # 2, N, N\n        rel_coords = rel_coords.permute(1, 2, 0).contiguous()           # N, N, 2\n        rel_coords[:, :, 0] += window_size - 1\n        rel_coords[:, :, 1] += window_size - 1\n        rel_coords[:, :, 0] *= 2 * window_size - 1\n        rel_pos_index = rel_coords.sum(-1)   # N, N\n        self.register_buffer('relative_position_index', rel_pos_index)\n\n        self.qkv      = nn.Linear(dim, dim * 3, bias=qkv_bias)\n        self.attn_drop = nn.Dropout(attn_drop)\n        self.proj      = nn.Linear(dim, dim)\n        self.proj_drop = nn.Dropout(proj_drop)\n        self.softmax   = nn.Softmax(dim=-1)\n\n    def forward(self, x, attn_mask=None):\n        # x: (num_windows*B, N, C)\n        B_, N, C = x.shape\n        qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads)\n        q, k, v = qkv.permute(2, 0, 3, 1, 4).unbind(0)  # each (B_, heads, N, head_dim)\n\n        attn = (q @ k.transpose(-2, -1)) * self.scale\n\n        # Add relative position bias\n        rel_bias = self.relative_position_bias_table[\n            self.relative_position_index.view(-1)\n        ].view(N, N, -1).permute(2, 0, 1).contiguous()  # heads, N, N\n        attn = attn + rel_bias.unsqueeze(0)\n\n        if attn_mask is not None:\n            nW = attn_mask.shape[0]\n            attn = attn.view(B_ // nW, nW, self.num_heads, N, N)\n            attn = attn + attn_mask.unsqueeze(1).unsqueeze(0)\n            attn = attn.view(-1, self.num_heads, N, N)\n\n        attn = self.softmax(attn)\n        attn = self.attn_drop(attn)\n\n        x = (attn @ v).transpose(1, 2).reshape(B_, N, C)\n        return self.proj_drop(self.proj(x))\n\n\n# ── Swin Transformer Block ────────────────────────────────────────────────────\nclass SwinTransformerBlock(nn.Module):\n    \"\"\"\n    One Swin Transformer block.\n    shift_size=0 → regular window attention.\n    shift_size=window_size//2 → shifted window attention.\n    \"\"\"\n    def __init__(self, dim, num_heads, window_size=6, shift_size=0,\n                 mlp_ratio=4., drop=0., attn_drop=0.):\n        super().__init__()\n        self.dim         = dim\n        self.window_size = window_size\n        self.shift_size  = shift_size\n        self.norm1 = nn.LayerNorm(dim)\n        self.norm2 = nn.LayerNorm(dim)\n        self.attn  = WindowAttention(\n            dim, window_size=window_size, num_heads=num_heads,\n            attn_drop=attn_drop, proj_drop=drop\n        )\n        mlp_hidden = int(dim * mlp_ratio)\n        self.mlp = nn.Sequential(\n            nn.Linear(dim, mlp_hidden),\n            nn.GELU(),\n            nn.Dropout(drop),\n            nn.Linear(mlp_hidden, dim),\n            nn.Dropout(drop),\n        )\n        self._attn_mask = None\n        self._H = self._W = None\n\n    def _build_attn_mask(self, H, W, device):\n        \"\"\"Build cyclic-shift attention mask (only needed when shift_size > 0).\"\"\"\n        if self.shift_size == 0:\n            return None\n        img_mask = torch.zeros(1, H, W, 1, device=device)\n        h_slices = (slice(0, -self.window_size),\n                    slice(-self.window_size, -self.shift_size),\n                    slice(-self.shift_size, None))\n        w_slices = (slice(0, -self.window_size),\n                    slice(-self.window_size, -self.shift_size),\n                    slice(-self.shift_size, None))\n        cnt = 0\n        for h in h_slices:\n            for w in w_slices:\n                img_mask[:, h, w, :] = cnt\n                cnt += 1\n        mask_windows = window_partition(img_mask, self.window_size)   # nW, ws, ws, 1\n        mask_windows = mask_windows.view(-1, self.window_size * self.window_size)\n        attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)  # nW, N, N\n        return attn_mask.masked_fill(attn_mask != 0, -100.0).masked_fill(attn_mask == 0, 0.0)\n\n    def forward(self, x):\n        # x: (B, H, W, C)\n        B, H, W, C = x.shape\n\n        # Rebuild mask only when spatial size changes\n        if self._H != H or self._W != W:\n            self._attn_mask = self._build_attn_mask(H, W, x.device)\n            self._H, self._W = H, W\n\n        shortcut = x\n        x = self.norm1(x)\n\n        # Cyclic shift\n        if self.shift_size > 0:\n            x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))\n\n        # Partition → attend → reverse\n        x_windows = window_partition(x, self.window_size)            # nW*B, ws, ws, C\n        x_windows = x_windows.view(-1, self.window_size ** 2, C)     # nW*B, N, C\n        attn_out = self.attn(x_windows, attn_mask=self._attn_mask)   # nW*B, N, C\n        attn_out = attn_out.view(-1, self.window_size, self.window_size, C)\n        x = window_reverse(attn_out, self.window_size, H, W)         # B, H, W, C\n\n        # Reverse shift\n        if self.shift_size > 0:\n            x = torch.roll(x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))\n\n        x = shortcut + x\n        x = x + self.mlp(self.norm2(x))\n        return x\n\n\n# ── Patch Merging (between Swin stages) ───────────────────────────────────────\nclass PatchMerging(nn.Module):\n    \"\"\"\n    Downsample spatial resolution by 2x, double channels.\n    (B, H, W, C) → (B, H/2, W/2, 2C)\n    \"\"\"\n    def __init__(self, dim):\n        super().__init__()\n        self.norm = nn.LayerNorm(4 * dim)\n        self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)\n\n    def forward(self, x):\n        B, H, W, C = x.shape\n        x0 = x[:, 0::2, 0::2, :]  # top-left\n        x1 = x[:, 1::2, 0::2, :]  # bottom-left\n        x2 = x[:, 0::2, 1::2, :]  # top-right\n        x3 = x[:, 1::2, 1::2, :]  # bottom-right\n        return self.reduction(self.norm(torch.cat([x0, x1, x2, x3], dim=-1)))\n\n\n# ── Swin Stage (one resolution level) ────────────────────────────────────────\nclass SwinStage(nn.Module):\n    \"\"\"\n    One Swin stage: `depth` alternating W-MSA / SW-MSA blocks,\n    optionally followed by PatchMerging for downsampling.\n    \"\"\"\n    def __init__(self, dim, depth, num_heads, window_size,\n                 mlp_ratio=4., drop=0., attn_drop=0., downsample=True):\n        super().__init__()\n        self.blocks = nn.ModuleList([\n            SwinTransformerBlock(\n                dim=dim, num_heads=num_heads,\n                window_size=window_size,\n                shift_size=0 if (i % 2 == 0) else window_size // 2,\n                mlp_ratio=mlp_ratio, drop=drop, attn_drop=attn_drop\n            )\n            for i in range(depth)\n        ])\n        self.downsample = PatchMerging(dim) if downsample else None\n\n    def forward(self, x):\n        for blk in self.blocks:\n            x = blk(x)\n        if self.downsample is not None:\n            x = self.downsample(x)\n        return x\n\n\nprint(\"Swin Transformer components defined ✓\")\nprint(f\"  WindowAttention: window_size={CFG['swin_window_size']}, with relative position bias\")\nprint(f\"  SwinTransformerBlock: regular + shifted window attention\")\nprint(f\"  PatchMerging: 2× spatial downsampling, 2× channel expansion\")\nprint(f\"  SwinStage: {CFG['swin_depths']} blocks per stage\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:51:06.616109Z","iopub.execute_input":"2026-04-12T04:51:06.616697Z","iopub.status.idle":"2026-04-12T04:51:06.643851Z","shell.execute_reply.started":"2026-04-12T04:51:06.616668Z","shell.execute_reply":"2026-04-12T04:51:06.643136Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 9 — CNN + Swin Transformer Hybrid Model\n\nArchitecture overview:\n```\nInput (384×384)\n    │\n    ▼\nEfficientNet-B4 backbone\n    │  stage3: (B×160×24×24)  → Swin Transformer encoder\n    │  stage4: (B×448×12×12)  → Global Avg Pool → CNN global branch\n    ▼\nInput projection  160→96  →  (B, 24, 24, 96)\n    │\n    ▼\nSwin Stage 1  [2 blocks, 3 heads, ws=6, W-MSA + SW-MSA]  →  (B, 12, 12, 192)\n    │\n    ▼\nSwin Stage 2  [2 blocks, 6 heads, ws=6, W-MSA + SW-MSA]  →  (B, 12, 12, 192)\n    │\n    ▼\nGlobal Avg Pool  →  (B, 192)\n    │\n    ▼\nGated Fusion with CNN global features (B×448 → B×512)\n    │\n    ▼\nFusion MLP → (B×512) → Classifier → 5 classes\n```","metadata":{}},{"cell_type":"code","source":"# ── GPU memory flush — run this before building the model ──────────────────\nimport gc, torch\n\n# Kill any leftover model/tensor from previous failed runs\nfor _var in ['_test_model', 'model', '_dummy', '_out']:\n    if _var in dir():\n        del globals()[_var]\n\ngc.collect()\ntorch.cuda.empty_cache()\ntorch.cuda.synchronize()\n\nfree, total = torch.cuda.mem_get_info()\nprint(f'GPU memory free : {free/1e9:.2f} GB / {total/1e9:.2f} GB')\nprint('If free < 8 GB, use Kernel → Restart & Run All to clear the session.')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:52:24.606948Z","iopub.execute_input":"2026-04-12T04:52:24.607429Z","iopub.status.idle":"2026-04-12T04:52:25.127457Z","shell.execute_reply.started":"2026-04-12T04:52:24.607399Z","shell.execute_reply":"2026-04-12T04:52:25.126758Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class CNNSwinHybrid(nn.Module):\n    \"\"\"\n    CNN + Swin Transformer Hybrid for APTOS DR grading.\n\n    EfficientNet-B4 stages:\n      Stage 3 → 160ch @ 24x24  → Swin Transformer encoder (spatial tokens)\n      Stage 4 → 448ch @ 12x12  → global avg pool → CNN global branch\n\n    Swin Transformer encoder:\n      Stage 1: 2 blocks, 3 heads, embed_dim=96,  window=6  (24x24 → 12x12 after PatchMerging)\n      Stage 2: 2 blocks, 6 heads, embed_dim=192, window=6  (12x12 → no downsampling)\n\n    Fusion: learned gate merges CNN global + Swin global before classifier.\n    \"\"\"\n    def __init__(self, cfg):\n        super().__init__()\n        stage3_ch   = cfg['cnn_stage3_ch']    # 160\n        global_ch   = cfg['cnn_global_ch']    # 448\n        embed_dim   = cfg['swin_embed_dim']   # 96\n        depths      = cfg['swin_depths']      # [2, 2]\n        num_heads   = cfg['swin_num_heads']   # [3, 6]\n        window_size = cfg['swin_window_size'] # 6\n        mlp_ratio   = cfg['swin_mlp_ratio']   # 4.0\n        fusion_dim  = cfg['fusion_dim']       # 512\n        num_classes = cfg['num_classes']      # 5\n        drop_rate   = cfg['drop_rate']        # 0.3\n\n        # ── CNN Backbone: EfficientNet-B4, feature stages 3 & 4 ─────────────\n        self.cnn = timm.create_model(\n            cfg['cnn_backbone'],\n            pretrained=True,\n            features_only=True,\n            out_indices=(3, 4),\n        )\n\n        # ── Project stage3 channels to Swin embed_dim ────────────────────────\n        # 160 → 96 (stage 1 input)\n        self.input_proj = nn.Sequential(\n            nn.LayerNorm(stage3_ch),\n            nn.Linear(stage3_ch, embed_dim, bias=False),\n        )\n\n        # ── Swin Stage 1: 24x24 → 12x12, 96 → 192 ───────────────────────────\n        self.swin_stage1 = SwinStage(\n            dim=embed_dim,\n            depth=depths[0],\n            num_heads=num_heads[0],\n            window_size=window_size,\n            mlp_ratio=mlp_ratio,\n            drop=drop_rate * 0.5,\n            attn_drop=drop_rate * 0.3,\n            downsample=True,       # PatchMerging: 24x24→12x12, 96→192\n        )\n\n        # ── Swin Stage 2: 12x12, 192 channels ────────────────────────────────\n        self.swin_stage2 = SwinStage(\n            dim=embed_dim * 2,    # 192 (output of stage1 PatchMerging)\n            depth=depths[1],\n            num_heads=num_heads[1],\n            window_size=window_size,\n            mlp_ratio=mlp_ratio,\n            drop=drop_rate * 0.5,\n            attn_drop=drop_rate * 0.3,\n            downsample=False,     # no downsampling at final stage\n        )\n\n        swin_out_dim = embed_dim * 2  # 192\n        self.swin_norm = nn.LayerNorm(swin_out_dim)\n\n        # ── Global pooling ─────────────────────────────────────────────────────\n        self.gap = nn.AdaptiveAvgPool2d(1)\n\n        # ── Project both branches to fusion_dim ──────────────────────────────\n        self.cnn_proj  = nn.Linear(global_ch,   fusion_dim, bias=False)  # 448→512\n        self.swin_proj = nn.Linear(swin_out_dim, fusion_dim, bias=False)  # 192→512\n\n        # ── Gated Fusion ───────────────────────────────────────────────────────\n        self.gate = nn.Sequential(\n            nn.Linear(fusion_dim * 2, fusion_dim),\n            nn.Sigmoid()\n        )\n        self.fusion_norm = nn.LayerNorm(fusion_dim)\n\n        # ── Classifier Head ───────────────────────────────────────────────────\n        self.head = nn.Sequential(\n            nn.Linear(fusion_dim, fusion_dim // 2),\n            nn.GELU(),\n            nn.Dropout(drop_rate),\n            nn.Linear(fusion_dim // 2, num_classes),\n        )\n\n        self._init_weights()\n\n    def _init_weights(self):\n        for m in [self.cnn_proj, self.swin_proj, self.gate, self.head, self.input_proj]:\n            for layer in m.modules():\n                if isinstance(layer, nn.Linear):\n                    nn.init.trunc_normal_(layer.weight, std=0.02)\n                    if layer.bias is not None:\n                        nn.init.zeros_(layer.bias)\n\n    def forward(self, x):\n        # x: (B, 3, H, W)\n\n        # 1. CNN feature extraction\n        stage3, stage4 = self.cnn(x)\n        # stage3: (B, 160, 24, 24)   stage4: (B, 448, 12, 12)\n\n        # 2. CNN global branch (deep semantics from stage4)\n        cnn_embed = self.cnn_proj(self.gap(stage4).flatten(1))   # (B, 512)\n\n        # 3. Swin Transformer branch\n        # Convert (B, C, H, W) → (B, H, W, C), project to embed_dim\n        tokens = self.input_proj(stage3.permute(0, 2, 3, 1))     # (B, 24, 24, 96)\n\n        # Stage 1: W-MSA + SW-MSA × 2, then PatchMerging → (B, 12, 12, 192)\n        tokens = self.swin_stage1(tokens)\n\n        # Stage 2: W-MSA + SW-MSA × 2, no downsampling → (B, 12, 12, 192)\n        tokens = self.swin_stage2(tokens)\n        tokens = self.swin_norm(tokens)                           # (B, 12, 12, 192)\n\n        # Global average pooling over spatial dims\n        swin_embed = self.swin_proj(\n            self.gap(tokens.permute(0, 3, 1, 2)).flatten(1)\n        )                                                         # (B, 512)\n\n        # 4. Gated fusion\n        gate_w = self.gate(torch.cat([cnn_embed, swin_embed], dim=-1))  # (B, 512)\n        fused  = self.fusion_norm(gate_w * cnn_embed + (1 - gate_w) * swin_embed)\n\n        # 5. Classify\n        return self.head(fused)                                   # (B, 5)\n\n\n# ── Sanity check ──────────────────────────────────────────────────────────────\ndef count_params(model):\n    total   = sum(p.numel() for p in model.parameters())\n    trainab = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    return total, trainab\n\n\n# Count params on CPU — avoids wasting GPU memory on a test forward pass\n_cpu_model = CNNSwinHybrid(CFG)   # CPU\n_total, _trainable = count_params(_cpu_model)\nprint(f'Model           : CNN + Swin Transformer Hybrid')\nprint(f'Total params    : {_total/1e6:.1f}M')\nprint(f'Trainable params: {_trainable/1e6:.1f}M')\ndel _cpu_model\nprint('\\nModel architecture verified ✓  (forward-pass test skipped to save GPU memory)')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:53:21.852613Z","iopub.execute_input":"2026-04-12T04:53:21.853462Z","iopub.status.idle":"2026-04-12T04:53:24.819177Z","shell.execute_reply.started":"2026-04-12T04:53:21.853432Z","shell.execute_reply":"2026-04-12T04:53:24.818444Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 10 — Loss, Optimizer, Scheduler","metadata":{}},{"cell_type":"code","source":"# ── Ordinal CE + QWK-proxy loss ────────────────────────────────────────────────\n# Standard CE treats grades as unordered. For DR grading, predicting grade 0\n# when truth is grade 4 is much worse than predicting grade 3.\n# We combine two losses:\n#   1. OrdinalCE: each grade k is split into K-1 binary tasks \"is grade > k?\"\n#      This gives the model an ordinal inductive bias matching QWK.\n#   2. Standard CE with label smoothing (keeps multi-class calibration).\n# Final loss = 0.5 * ordinal_ce + 0.5 * label_smooth_ce\n\nclass OrdinalCE(nn.Module):\n    \"\"\"\n    Ordinal cross-entropy via K-1 binary sub-tasks.\n    For K=5 grades: predicts P(y>0), P(y>1), P(y>2), P(y>3).\n    Forces monotone ordinal structure; gradients align with QWK.\n    \"\"\"\n    def __init__(self, num_classes=5, weight=None):\n        super().__init__()\n        self.K      = num_classes\n        self.weight = weight   # per-class weights tensor\n\n    def forward(self, logits, targets):\n        # logits: (B, K),  targets: (B,) int in [0, K-1]\n        B, K = logits.shape\n        # Cumulative logits: logit[k] approximates log-odds of P(y > k)\n        # Use running cumsum of softmax as ordinal probabilities\n        probs = torch.softmax(logits, dim=1)                 # (B, K)\n        cum   = 1.0 - torch.cumsum(probs, dim=1)[:, :-1]   # (B, K-1): P(y > k)\n        cum   = cum.clamp(1e-6, 1 - 1e-6)\n\n        # Binary target: for each threshold k, target = 1 if y > k\n        t = targets.unsqueeze(1)                             # (B, 1)\n        thresholds = torch.arange(K-1, device=logits.device).unsqueeze(0)  # (1, K-1)\n        binary_tgt = (t > thresholds).float()               # (B, K-1)\n\n        bce = -(binary_tgt * torch.log(cum) + (1 - binary_tgt) * torch.log(1 - cum))\n\n        if self.weight is not None:\n            w   = self.weight.to(logits.device)[targets]    # (B,)\n            bce = bce * w.unsqueeze(1)\n\n        return bce.mean()\n\n\nclass CombinedOrdinalLoss(nn.Module):\n    \"\"\"0.5 * OrdinalCE + 0.5 * LabelSmoothCE — best of both.\"\"\"\n    def __init__(self, num_classes=5, smoothing=0.05, weight=None, ordinal_w=0.5):\n        super().__init__()\n        self.ordinal_w  = ordinal_w\n        self.ordinal    = OrdinalCE(num_classes, weight)\n        self.ce_smooth  = LabelSmoothCE(smoothing, weight)\n\n    def forward(self, logits, targets):\n        return (self.ordinal_w       * self.ordinal(logits, targets) +\n                (1 - self.ordinal_w) * self.ce_smooth(logits, targets))\n\n\nclass LabelSmoothCE(nn.Module):\n    \"\"\"Cross-entropy with label smoothing + per-class weights.\"\"\"\n    def __init__(self, smoothing=0.05, weight=None):\n        super().__init__()\n        self.smoothing = smoothing\n        self.weight    = weight\n\n    def forward(self, logits, targets):\n        n        = logits.size(1)\n        log_prob = F.log_softmax(logits, dim=1)\n        smooth   = torch.full_like(log_prob, self.smoothing / (n - 1))\n        smooth.scatter_(1, targets.unsqueeze(1), 1.0 - self.smoothing)\n        loss = -(smooth * log_prob).sum(dim=1)\n        if self.weight is not None:\n            w    = self.weight.to(logits.device)[targets]\n            loss = loss * w\n        return loss.mean()\n\n\ndef build_optimizer_scheduler(model, cfg, steps_per_epoch):\n    \"\"\"3-group differential LR: CNN backbone | VMamba | head.\"\"\"\n    lr = cfg['lr']\n\n    cnn_params, vmamba_params, head_params = [], [], []  # vmamba_params reused for swin\n\n    for name, param in model.named_parameters():\n        if not param.requires_grad:\n            continue\n        if name.startswith('cnn.'):\n            cnn_params.append(param)\n        elif any(name.startswith(p) for p in ['swin', 'input_proj', 'cnn_proj', 'swin_proj', 'gate', 'fusion', 'gap']):\n            vmamba_params.append(param)\n        else:\n            head_params.append(param)\n\n    param_groups = [\n        {'params': cnn_params,    'lr': lr * cfg['backbone_lr_scale'], 'name': 'cnn_backbone'},\n        {'params': vmamba_params, 'lr': lr * cfg['swin_lr_scale'],     'name': 'swin'},\n        {'params': head_params,   'lr': lr,                              'name': 'head'},\n    ]\n    optimizer = torch.optim.AdamW(param_groups, weight_decay=cfg['weight_decay'])\n\n    total_steps  = steps_per_epoch * cfg['epochs'] // cfg['accum_steps']\n    warmup_steps = steps_per_epoch * cfg['warmup_epochs'] // cfg['accum_steps']\n\n    def lr_lambda(step):\n        if step < warmup_steps:\n            return step / max(warmup_steps, 1)\n        progress = (step - warmup_steps) / max(total_steps - warmup_steps, 1)\n        cosine   = 0.5 * (1 + math.cos(math.pi * progress))\n        return max(cosine, cfg['min_lr'] / lr)\n\n    scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n    return optimizer, scheduler\n\n\nprint('Loss / optimizer / scheduler defined ✓')\nprint('  → CombinedOrdinalLoss: 0.5*OrdinalCE + 0.5*LabelSmoothCE(0.05)')\nprint('  → 3-group LR: CNN=lr*0.1  Swin=lr*0.3  head=lr')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:55:53.953091Z","iopub.execute_input":"2026-04-12T04:55:53.95342Z","iopub.status.idle":"2026-04-12T04:55:53.970947Z","shell.execute_reply.started":"2026-04-12T04:55:53.953398Z","shell.execute_reply":"2026-04-12T04:55:53.97022Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 11 — Train & Validation Loops","metadata":{}},{"cell_type":"code","source":"def compute_metrics(y_true, y_pred, y_prob):\n    kappa = cohen_kappa_score(y_true, y_pred, weights='quadratic')\n    acc   = float((y_true == y_pred).mean())\n\n    unique = np.unique(y_true)\n    if len(unique) == 5:\n        try:\n            auc = roc_auc_score(y_true, y_prob,\n                                multi_class='ovr', average='macro')\n        except Exception:\n            auc = float('nan')\n    else:\n        auc = float('nan')\n\n    report = classification_report(\n        y_true, y_pred,\n        target_names=[GRADE_NAMES[i] for i in range(5)],\n        output_dict=True, zero_division=0\n    )\n    return {\n        'kappa'    : round(kappa, 4),\n        'accuracy' : round(acc,   4),\n        'auc'      : round(auc,   4),\n        'f1_macro' : round(report['macro avg']['f1-score'], 4),\n        'prec_macro': round(report['macro avg']['precision'], 4),\n        'rec_macro' : round(report['macro avg']['recall'], 4),\n        'report'   : report,\n    }\n\n\ndef train_epoch(model, loader, criterion, optimizer, scheduler, scaler, cfg, step_count):\n    model.train()\n    total_loss = 0.0\n    optimizer.zero_grad()\n\n    for batch_idx, (imgs, labels) in enumerate(loader):\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n\n        with torch.amp.autocast('cuda', enabled=cfg['amp']):\n            logits = model(imgs)\n            loss   = criterion(logits, labels) / cfg['accum_steps']\n\n        scaler.scale(loss).backward()\n        total_loss += loss.item() * cfg['accum_steps']\n\n        if (batch_idx + 1) % cfg['accum_steps'] == 0:\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n            scheduler.step()\n            step_count += 1\n\n    return total_loss / len(loader), step_count\n\n\n@torch.no_grad()\ndef val_epoch(model, loader, criterion, cfg):\n    model.eval()\n    total_loss = 0.0\n    all_preds, all_probs, all_labels = [], [], []\n\n    for imgs, labels in loader:\n        imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n        with torch.amp.autocast('cuda', enabled=cfg['amp']):\n            logits = model(imgs)\n            loss   = criterion(logits, labels)\n\n        total_loss += loss.item()\n        probs  = torch.softmax(logits, dim=1).cpu().numpy()\n        preds  = np.argmax(probs, axis=1)\n        all_probs.append(probs)\n        all_preds.append(preds)\n        all_labels.append(labels.cpu().numpy())\n\n    all_preds  = np.concatenate(all_preds)\n    all_probs  = np.concatenate(all_probs)\n    all_labels = np.concatenate(all_labels)\n\n    metrics = compute_metrics(all_labels, all_preds, all_probs)\n    metrics['loss'] = round(total_loss / len(loader), 4)\n    return metrics, all_labels, all_preds, all_probs\n\n\nprint('Train/val loop functions defined ✓')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:56:11.992346Z","iopub.execute_input":"2026-04-12T04:56:11.993089Z","iopub.status.idle":"2026-04-12T04:56:12.005734Z","shell.execute_reply.started":"2026-04-12T04:56:11.993059Z","shell.execute_reply":"2026-04-12T04:56:12.004785Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 12 — Train the Hybrid Model","metadata":{}},{"cell_type":"code","source":"print('=' * 65)\nprint('  CNN + Swin Transformer Hybrid — Training')\nprint('=' * 65)\n\ntrain_loader, val_loader = get_loaders(CFG)\nsteps_per_epoch = len(train_loader)\nprint(f'Steps/epoch : {steps_per_epoch}')\nprint(f'Opt steps   : ~{steps_per_epoch * CFG[\"epochs\"] // CFG[\"accum_steps\"]}\\n')\n\nmodel     = CNNSwinHybrid(CFG).to(DEVICE)\ncriterion = CombinedOrdinalLoss(\n    num_classes=CFG['num_classes'],\n    smoothing=CFG['label_smooth'],\n    weight=CLASS_WEIGHTS_T,\n    ordinal_w=0.5\n)\noptimizer, scheduler = build_optimizer_scheduler(model, CFG, steps_per_epoch)\nscaler    = torch.amp.GradScaler('cuda', enabled=CFG['amp'])\n\nbest_kappa   = -1.0\nbest_weights = None\npatience_cnt = 0\nstep_count   = 0\nhistory      = []\nt_start      = time.time()\n\nfor epoch in range(1, CFG['epochs'] + 1):\n    t0 = time.time()\n\n    tr_loss, step_count = train_epoch(\n        model, train_loader, criterion, optimizer, scheduler, scaler, CFG, step_count\n    )\n    val_metrics, val_labels, val_preds, val_probs = val_epoch(\n        model, val_loader, criterion, CFG\n    )\n\n    kappa = val_metrics['kappa']\n    dt    = time.time() - t0\n    lr_now = optimizer.param_groups[2]['lr']  # head LR\n\n    history.append({\n        'epoch'    : epoch,\n        'tr_loss'  : round(tr_loss, 4),\n        **val_metrics,\n    })\n\n    flag = ''\n    if kappa > best_kappa:\n        best_kappa   = kappa\n        best_weights = deepcopy(model.state_dict())\n        patience_cnt = 0\n        flag = ' ← best'\n        torch.save(best_weights, OUT_DIR / 'best_swin_hybrid_model.pth')\n    else:\n        patience_cnt += 1\n\n    print(\n        f'Ep {epoch:02d}/{CFG[\"epochs\"]}  '\n        f'tr_loss={tr_loss:.4f}  '\n        f'val_loss={val_metrics[\"loss\"]:.4f}  '\n        f'QWK={kappa:.4f}  '\n        f'Acc={val_metrics[\"accuracy\"]:.4f}  '\n        f'F1={val_metrics[\"f1_macro\"]:.4f}  '\n        f'LR={lr_now:.2e}  '\n        f't={dt:.0f}s{flag}'\n    )\n\n    if patience_cnt >= CFG['patience']:\n        print(f'\\nEarly stopping at epoch {epoch} (patience={CFG[\"patience\"]})')\n        break\n\ntotal_mins = (time.time() - t_start) / 60\nprint(f'\\nTraining complete — {total_mins:.1f} min')\nprint(f'Best QWK: {best_kappa:.4f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T04:56:26.03205Z","iopub.execute_input":"2026-04-12T04:56:26.032989Z","iopub.status.idle":"2026-04-12T06:29:39.318104Z","shell.execute_reply.started":"2026-04-12T04:56:26.032958Z","shell.execute_reply":"2026-04-12T06:29:39.317227Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 13 — Evaluate Best Model","metadata":{}},{"cell_type":"code","source":"# ── Test-Time Augmentation (TTA) ───────────────────────────────────────────\n# Average predictions over 5 views: original + 4 flips/rotations.\n# Costs 5x inference time but typically gains +0.01–0.02 QWK for free.\n\ndef tta_predict(model, loader, cfg, n_tta=5):\n    \"\"\"Run TTA: average softmax over n_tta augmented views per image.\"\"\"\n    import albumentations as A\n    from albumentations.pytorch import ToTensorV2\n\n    tta_transforms = [\n        # 0: original (no aug)\n        A.Compose([A.Normalize(mean=MEAN, std=STD), ToTensorV2()]),\n        # 1: horizontal flip\n        A.Compose([A.HorizontalFlip(p=1.0), A.Normalize(mean=MEAN, std=STD), ToTensorV2()]),\n        # 2: vertical flip\n        A.Compose([A.VerticalFlip(p=1.0), A.Normalize(mean=MEAN, std=STD), ToTensorV2()]),\n        # 3: 90-degree rotation\n        A.Compose([A.RandomRotate90(p=1.0), A.Normalize(mean=MEAN, std=STD), ToTensorV2()]),\n        # 4: hflip + vflip\n        A.Compose([A.HorizontalFlip(p=1.0), A.VerticalFlip(p=1.0),\n                   A.Normalize(mean=MEAN, std=STD), ToTensorV2()]),\n    ]\n\n    model.eval()\n    all_probs_list = []   # one entry per TTA view\n    all_labels     = []\n\n    with torch.no_grad():\n        for tta_idx, tfm in enumerate(tta_transforms[:n_tta]):\n            # Rebuild dataset with this TTA transform\n            tta_ds     = APTOSDataset(VAL_DF, cfg, tfm, is_train=False)\n            tta_loader = DataLoader(tta_ds, batch_size=cfg['batch_size']*2,\n                                    shuffle=False, num_workers=4, pin_memory=True)\n            view_probs = []\n            view_labels= []\n            for imgs, labels in tta_loader:\n                imgs = imgs.to(DEVICE)\n                with torch.amp.autocast('cuda', enabled=cfg['amp']):\n                    logits = model(imgs)\n                probs = torch.softmax(logits, dim=1).cpu().numpy()\n                view_probs.append(probs)\n                if tta_idx == 0:\n                    view_labels.append(labels.numpy())\n            all_probs_list.append(np.concatenate(view_probs))\n            if tta_idx == 0:\n                all_labels = np.concatenate(view_labels)\n\n    avg_probs = np.mean(all_probs_list, axis=0)   # average over TTA views\n    avg_preds = np.argmax(avg_probs, axis=1)\n    return avg_preds, avg_probs, all_labels\n\n\nprint('TTA function defined ✓  (5 views: original + hflip + vflip + rot90 + hflip+vflip)')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T06:53:18.014965Z","iopub.execute_input":"2026-04-12T06:53:18.015699Z","iopub.status.idle":"2026-04-12T06:53:18.026693Z","shell.execute_reply.started":"2026-04-12T06:53:18.015658Z","shell.execute_reply":"2026-04-12T06:53:18.026077Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load best weights\nmodel.load_state_dict(torch.load(OUT_DIR / 'best_swin_hybrid_model.pth', map_location=DEVICE))\n\n# Standard eval (no TTA)\nstd_metrics, std_labels, std_preds, std_probs = val_epoch(\n    model, val_loader, criterion, CFG\n)\nprint(f'Without TTA — QWK: {std_metrics[\"kappa\"]:.4f}  Acc: {std_metrics[\"accuracy\"]:.4f}')\n\n# TTA eval (5 views)\ntta_preds, tta_probs, tta_labels = tta_predict(model, val_loader, CFG, n_tta=5)\ntta_metrics = compute_metrics(tta_labels, tta_preds, tta_probs)\nprint(f'With TTA    — QWK: {tta_metrics[\"kappa\"]:.4f}  Acc: {tta_metrics[\"accuracy\"]:.4f}')\nprint(f'TTA gain    : {tta_metrics[\"kappa\"]-std_metrics[\"kappa\"]:+.4f}')\n\n# Use TTA result as final\nfinal_metrics = tta_metrics\nfinal_labels, final_preds, final_probs = tta_labels, tta_preds, tta_probs\n\nprint('=' * 55)\nprint('  CNN + Swin Transformer Hybrid — Final Validation Results')\nprint('=' * 55)\nprint(f'  QWK (primary)      : {final_metrics[\"kappa\"]:.4f}')\nprint(f'  Accuracy           : {final_metrics[\"accuracy\"]:.4f}')\nprint(f'  AUC-ROC (macro)    : {final_metrics[\"auc\"]:.4f}')\nprint(f'  F1  (macro)        : {final_metrics[\"f1_macro\"]:.4f}')\nprint(f'  Precision (macro)  : {final_metrics[\"prec_macro\"]:.4f}')\nprint(f'  Recall (macro)     : {final_metrics[\"rec_macro\"]:.4f}')\nprint('=' * 55)\n\n# Comparison with baselines\nbaselines = {\n    'ResNet-50':       0.7473,\n    'EfficientNet-B4': 0.7766,\n    'VGG-19-BN':       0.8813,\n    'DenseNet-121':    0.8291,\n}\nprint('\\nQWK Comparison:')\nfor name, kappa in baselines.items():\n    diff = final_metrics['kappa'] - kappa\n    sign = '+' if diff >= 0 else ''\n    print(f'  vs {name:<18}: {kappa:.4f}  ({sign}{diff:.4f})')\nprint(f'  Hybrid (ours)       : {final_metrics[\"kappa\"]:.4f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T06:53:30.246917Z","iopub.execute_input":"2026-04-12T06:53:30.247199Z","iopub.status.idle":"2026-04-12T06:58:46.047756Z","shell.execute_reply.started":"2026-04-12T06:53:30.247177Z","shell.execute_reply":"2026-04-12T06:58:46.046846Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 14 — Per-Class Analysis & Confusion Matrix","metadata":{}},{"cell_type":"code","source":"report = final_metrics['report']\nprint('Per-class metrics:')\nprint(f'{\"Grade\":<15} {\"Precision\":>10} {\"Recall\":>9} {\"F1\":>9} {\"Support\":>9}')\nprint('-' * 56)\nfor g in range(5):\n    name = GRADE_NAMES[g]\n    r    = report[name]\n    print(f'{name:<15} {r[\"precision\"]:>10.4f} {r[\"recall\"]:>9.4f} '\n          f'{r[\"f1-score\"]:>9.4f} {int(r[\"support\"]):>9}')\n\n# Confusion matrix\ncm  = confusion_matrix(final_labels, final_preds)\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\n\nim = axes[0].imshow(cm, cmap='Blues')\naxes[0].set_xticks(range(5)); axes[0].set_yticks(range(5))\naxes[0].set_xticklabels([GRADE_NAMES[i] for i in range(5)], rotation=30, ha='right')\naxes[0].set_yticklabels([GRADE_NAMES[i] for i in range(5)])\naxes[0].set_xlabel('Predicted'); axes[0].set_ylabel('True')\naxes[0].set_title('Confusion Matrix — CNN+Swin Hybrid')\nfor i in range(5):\n    for j in range(5):\n        axes[0].text(j, i, str(cm[i, j]),\n                     ha='center', va='center',\n                     color='white' if cm[i, j] > cm.max()/2 else 'black')\nplt.colorbar(im, ax=axes[0])\n\n# Training history\nhist_df = pd.DataFrame(history)\naxes[1].plot(hist_df['epoch'], hist_df['kappa'],    label='Val QWK',  lw=2)\naxes[1].plot(hist_df['epoch'], hist_df['accuracy'], label='Val Acc',  lw=2, ls='--')\naxes[1].plot(hist_df['epoch'], hist_df['f1_macro'], label='Val F1',   lw=2, ls=':')\naxes[1].axhline(0.8813, color='red', ls='--', lw=1, label='VGG baseline (0.8813)')\naxes[1].set_xlabel('Epoch'); axes[1].set_ylabel('Score')\naxes[1].set_title('Training History — CNN+Swin Hybrid')\naxes[1].legend(); axes[1].grid(alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(OUT_DIR / 'swin_hybrid_results.png', dpi=150, bbox_inches='tight')\nplt.show()\nprint('Plot saved ✓')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T06:59:41.903641Z","iopub.execute_input":"2026-04-12T06:59:41.904109Z","iopub.status.idle":"2026-04-12T06:59:42.75716Z","shell.execute_reply.started":"2026-04-12T06:59:41.904074Z","shell.execute_reply":"2026-04-12T06:59:42.756491Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Cell 15 — Save Full Results","metadata":{}},{"cell_type":"code","source":"results_dict = {\n    'model'          : 'CNN + Swin Transformer Hybrid',\n    'backbone'       : CFG['cnn_backbone'],\n    'swin_depths'    : CFG['swin_depths'],\n    'swin_embed_dim' : CFG['swin_embed_dim'],\n    'kappa_qwk'      : final_metrics['kappa'],\n    'accuracy'       : final_metrics['accuracy'],\n    'auc_roc_macro'  : final_metrics['auc'],\n    'f1_macro'       : final_metrics['f1_macro'],\n    'precision_macro': final_metrics['prec_macro'],\n    'recall_macro'   : final_metrics['rec_macro'],\n    'history'        : history,\n}\n\nwith open(OUT_DIR / 'swin_hybrid_results.json', 'w') as f:\n    json.dump(results_dict, f, indent=2)\n\n# Also save as CSV for direct comparison with baseline results_comparison.csv\ncomparison_row = pd.DataFrame([{\n    'Model'            : 'CNN+Swin Hybrid',\n    'Kappa (QWK)'      : final_metrics['kappa'],\n    'Accuracy'         : final_metrics['accuracy'],\n    'AUC-ROC (macro)'  : final_metrics['auc'],\n    'F1 (macro)'       : final_metrics['f1_macro'],\n    'Precision (macro)': final_metrics['prec_macro'],\n    'Recall (macro)'   : final_metrics['rec_macro'],\n    'Time (min)'       : round(total_mins, 1),\n}])\ncomparison_row.to_csv(OUT_DIR / 'swin_hybrid_results_row.csv', index=False)\n\nprint('Results saved:')\nprint(f'  {OUT_DIR}/best_swin_hybrid_model.pth')\nprint(f'  {OUT_DIR}/swin_hybrid_results.json')\nprint(f'  {OUT_DIR}/swin_hybrid_results_row.csv')\nprint(f'  {OUT_DIR}/swin_hybrid_results.png')\nprint()\nprint('Done! Upload hybrid_results_row.csv alongside your baseline results_comparison.csv to compare.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-12T06:59:55.8662Z","iopub.execute_input":"2026-04-12T06:59:55.866962Z","iopub.status.idle":"2026-04-12T06:59:55.878719Z","shell.execute_reply.started":"2026-04-12T06:59:55.866934Z","shell.execute_reply":"2026-04-12T06:59:55.878063Z"}},"outputs":[],"execution_count":null}]}