{"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},{"sourceType":"kernelVersion","sourceId":29830115}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"34191ad2","cell_type":"markdown","source":"# Diabetic Retinopathy Classification\n## ConvNeXt-Tiny + CBAM | APTOS 2019 — **ConvNeXt Variant**\n### Backbone: `convnext_tiny` (28M params, ImageNet-1k pretrained)  ·  All memory fixes applied\n\n| # | Change vs EfficientNet-B3 notebook | Detail |\n|---|-------------------------------------|--------|\n| 1 | Backbone | `convnext_tiny` replaces `efficientnet_b3` |\n| 2 | IMG_SIZE | **224** (ConvNeXt standard) vs 300 for B3 |\n| 3 | BATCH_SIZE | **32** (lighter activations at 224²) vs 16 |\n| 4 | CBAM placement | After **each of the 4 ConvNeXt stages** (96→192→384→768 ch) |\n| 5 | Head normalisation | `LayerNorm(768)` copied from pretrained head (ConvNeXt standard) |\n| 6 | Cache dir | Separate `/preprocessed_cache_224` (different resolution) |\n\n**Memory fixes inherited from fixed EfficientNet notebook (all present):**\npre-cache preprocessing · `prefetch_factor=1` · `torch.cuda.empty_cache()` per epoch · `torch.stack` CORN loss\n","metadata":{}},{"id":"cbc268a9","cell_type":"markdown","source":"## Cell 1 — Install Dependencies","metadata":{}},{"id":"f2080651","cell_type":"code","source":"!pip install -q coral-pytorch opencv-python-headless albumentations timm tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T11:48:02.126843Z","iopub.execute_input":"2026-05-22T11:48:02.127771Z","iopub.status.idle":"2026-05-22T11:48:07.010378Z","shell.execute_reply.started":"2026-05-22T11:48:02.127732Z","shell.execute_reply":"2026-05-22T11:48:07.009647Z"}},"outputs":[],"execution_count":null},{"id":"bd7a2500","cell_type":"markdown","source":"## Cell 2 — Imports","metadata":{}},{"id":"4aa05c68","cell_type":"code","source":"import os, cv2, gc, math, random, time\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import cohen_kappa_score, classification_report\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport timm\nimport matplotlib.pyplot as plt\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\nSEED = 42\nrandom.seed(SEED); np.random.seed(SEED)\ntorch.manual_seed(SEED); torch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.benchmark = True\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nN_GPUS = torch.cuda.device_count()\n\nprint(f\"Device    : {DEVICE}\")\nprint(f\"GPU count : {N_GPUS}\")\nif N_GPUS > 0:\n    for i in range(N_GPUS):\n        p = torch.cuda.get_device_properties(i)\n        print(f\"  GPU {i}  : {p.name}  ({p.total_memory/1024**3:.1f} GB VRAM)\")\nprint(f\"PyTorch   : {torch.__version__}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T11:48:07.012067Z","iopub.execute_input":"2026-05-22T11:48:07.012308Z","iopub.status.idle":"2026-05-22T11:48:20.865677Z","shell.execute_reply.started":"2026-05-22T11:48:07.012281Z","shell.execute_reply":"2026-05-22T11:48:20.864997Z"}},"outputs":[],"execution_count":null},{"id":"f5ffeaad","cell_type":"markdown","source":"## Cell 3 — Configuration","metadata":{}},{"id":"ebd5f39a","cell_type":"code","source":"class CFG:\n    TRAIN_CSV = \"/kaggle/input/competitions/aptos2019-blindness-detection/train.csv\"\n    TEST_CSV  = \"/kaggle/input/competitions/aptos2019-blindness-detection/test.csv\"\n    TRAIN_IMG = \"/kaggle/input/competitions/aptos2019-blindness-detection/train_images\"\n    TEST_IMG  = \"/kaggle/input/competitions/aptos2019-blindness-detection/test_images\"\n\n    # Separate cache dir -- ConvNeXt uses 224x224, not 300x300\n    CACHE_DIR = \"/kaggle/working/preprocessed_cache_224\"\n\n    IMG_SIZE      = 224   # ConvNeXt standard (vs 300 for B3)\n    EPOCHS        = 20\n    BATCH_SIZE    = 32    # safe at 224x224 with AMP on T4; was 16 for B3@300\n    NUM_WORKERS   = 4\n    VAL_SPLIT     = 0.15\n    FREEZE_EPOCHS = 5\n    CORN_WEIGHT   = 0.5\n    WCE_WEIGHT    = 0.5\n    NUM_CLASSES   = 5\n    LR            = 1e-4\n    WEIGHT_DECAY  = 1e-5\n    CLAHE_CLIP    = 2.0\n    CLAHE_GRID    = (8, 8)\n    BEN_SIGMA     = 10\n    BEN_ALPHA     = 4\n    BEN_BETA      = -4\n    BEN_GAMMA     = 128\n\nos.makedirs(CFG.CACHE_DIR, exist_ok=True)\n\nn_images   = 3662 + 1928\ncache_mb   = n_images * CFG.IMG_SIZE * CFG.IMG_SIZE * 3 / 1024**2\n\nprint(\"CFG loaded  (ConvNeXt-Tiny variant)\")\nprint(f\"  IMG_SIZE    : {CFG.IMG_SIZE}x{CFG.IMG_SIZE}\")\nprint(f\"  BATCH_SIZE  : {CFG.BATCH_SIZE}  ({CFG.BATCH_SIZE // max(N_GPUS,1)} per GPU)\")\nprint(f\"  num_workers : {CFG.NUM_WORKERS}  (safe after caching)\")\nprint(f\"  Cache dir   : {CFG.CACHE_DIR}\")\nprint(f\"  Cache size  : ~{cache_mb:.0f} MB for {n_images} images\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T11:48:20.866574Z","iopub.execute_input":"2026-05-22T11:48:20.867019Z","iopub.status.idle":"2026-05-22T11:48:20.87386Z","shell.execute_reply.started":"2026-05-22T11:48:20.866997Z","shell.execute_reply":"2026-05-22T11:48:20.873146Z"}},"outputs":[],"execution_count":null},{"id":"13c19b1f","cell_type":"markdown","source":"## Cell 4 — Preprocessing Utilities","metadata":{}},{"id":"29cf20e6","cell_type":"code","source":"def crop_black_borders(image, tol=7):\n    mask = image.max(axis=2) > tol if image.ndim == 3 else image > tol\n    rows, cols = np.any(mask, axis=1), np.any(mask, axis=0)\n    if not rows.any() or not cols.any():\n        return image\n    rmin, rmax = np.where(rows)[0][[0,-1]]\n    cmin, cmax = np.where(cols)[0][[0,-1]]\n    return image[rmin:rmax+1, cmin:cmax+1]\n\ndef apply_circular_mask(image):\n    h, w = image.shape[:2]\n    cx, cy = w//2, h//2\n    Y, X = np.ogrid[:h, :w]\n    mask = np.sqrt((X-cx)**2 + (Y-cy)**2) <= min(cx, cy)\n    out = image.copy(); out[~mask] = 0\n    return out\n\ndef extract_green_channel(image):\n    # [PIPELINE CHECKPOINT]\n    # Green channel has highest contrast for haemorrhages and microaneurysms.\n    # Replicated to 3 channels so ImageNet pretrained weights still apply.\n    g = image[:,:,1]\n    return cv2.merge([g, g, g])\n\ndef apply_clahe(image):\n    lab = cv2.cvtColor(image, cv2.COLOR_RGB2LAB)\n    clahe = cv2.createCLAHE(clipLimit=CFG.CLAHE_CLIP, tileGridSize=CFG.CLAHE_GRID)\n    lab[:,:,0] = clahe.apply(lab[:,:,0])\n    return cv2.cvtColor(lab, cv2.COLOR_LAB2RGB)\n\ndef ben_graham_preprocessing(image):\n    blur = cv2.GaussianBlur(image, (0,0), sigmaX=CFG.BEN_SIGMA, sigmaY=CFG.BEN_SIGMA)\n    out  = cv2.addWeighted(image, CFG.BEN_ALPHA, blur, CFG.BEN_BETA, CFG.BEN_GAMMA)\n    return np.clip(out, 0, 255).astype(np.uint8)\n\ndef preprocess_fundus_image(image_path, img_size=CFG.IMG_SIZE):\n    # Full pipeline (runs once per image, saved to cache):\n    # Read -> Crop -> Resize -> CircleMask -> GreenChannel -> CLAHE -> BenGraham\n    img = cv2.imread(str(image_path))\n    if img is None:\n        raise FileNotFoundError(f\"Cannot read: {image_path}\")\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    img = crop_black_borders(img)\n    img = cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_LANCZOS4)\n    img = apply_circular_mask(img)\n    img = extract_green_channel(img)   # <-- PIPELINE CHECKPOINT\n    img = apply_clahe(img)\n    img = ben_graham_preprocessing(img)\n    return img   # uint8, shape (img_size, img_size, 3)\n\nprint(\"Preprocessing functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T11:48:20.874927Z","iopub.execute_input":"2026-05-22T11:48:20.875239Z","iopub.status.idle":"2026-05-22T11:48:20.8977Z","shell.execute_reply.started":"2026-05-22T11:48:20.875189Z","shell.execute_reply":"2026-05-22T11:48:20.897116Z"}},"outputs":[],"execution_count":null},{"id":"bd91d96d","cell_type":"markdown","source":"## Cell 5 — Preprocessing Visualisation  (Raw through Green Channel Checkpoint)","metadata":{}},{"id":"5bba26ea","cell_type":"code","source":"def visualise_preprocessing_pipeline(image_path):\n    SEP = \"-\" * 56\n\n    raw_bgr = cv2.imread(str(image_path))\n    raw     = cv2.cvtColor(raw_bgr, cv2.COLOR_BGR2RGB)\n    print(f\"STEP 1 | Raw Image Loaded\")\n    print(f\"       | Shape : {raw.shape}   dtype: {raw.dtype}\")\n    print(f\"       | Size  : {raw.nbytes/1024**2:.1f} MB in RAM  <- this is what workers loaded before fix\")\n    print(SEP)\n\n    cropped  = crop_black_borders(raw)\n    retained = cropped.shape[0]*cropped.shape[1] / (raw.shape[0]*raw.shape[1])\n    print(f\"STEP 2 | Black Border Crop\")\n    print(f\"       | Shape : {cropped.shape}   retained: {retained*100:.1f}%\")\n    print(SEP)\n\n    resized = cv2.resize(cropped, (CFG.IMG_SIZE, CFG.IMG_SIZE), interpolation=cv2.INTER_LANCZOS4)\n    print(f\"STEP 3 | Lanczos Resize -> {CFG.IMG_SIZE}x{CFG.IMG_SIZE}\")\n    print(f\"       | Shape : {resized.shape}   size: {resized.nbytes/1024:.0f} KB\")\n    print(SEP)\n\n    masked = apply_circular_mask(resized)\n    h, w   = resized.shape[:2]\n    Y, X   = np.ogrid[:h, :w]\n    zeroed = int(np.sum(np.sqrt((X-w//2)**2+(Y-h//2)**2) > min(w//2,h//2)))\n    print(f\"STEP 4 | Circular Fundus Mask\")\n    print(f\"       | Shape : {masked.shape}   zeroed: {zeroed:,} px ({zeroed/(h*w)*100:.1f}%)\")\n    print(SEP)\n\n    green = extract_green_channel(masked)\n    print(f\"STEP 5 | GREEN CHANNEL EXTRACTION  [PIPELINE CHECKPOINT]\")\n    print(f\"       | 2D slice   : {masked[:,:,1].shape}\")\n    print(f\"       | Merged RGB : {green.shape}   range: [{green.min()}, {green.max()}]\")\n    print(f\"       | Cached as  : {green.nbytes/1024:.0f} KB   <- workers now load THIS, not the {raw.nbytes/1024**2:.1f} MB raw PNG\")\n    print(f\"       | Speedup    : {raw.nbytes // green.nbytes}x smaller per worker\")\n    print(SEP)\n\n    stages = [\n        (raw,     f\"1 . Raw Input\\n{raw.shape}\\n{raw.nbytes/1024**2:.1f} MB\"),\n        (cropped, f\"2 . Border Cropped\\n{cropped.shape}\"),\n        (resized, f\"3 . Resized {CFG.IMG_SIZE}x{CFG.IMG_SIZE}\\n{resized.shape}\"),\n        (masked,  f\"4 . Circle Masked\\n{masked.shape}\\n{zeroed:,}px zeroed\"),\n        (green,   f\"5 . Green Ch  [CHECKPOINT]\\n{green.shape}\\n{green.nbytes/1024:.0f} KB cached\"),\n    ]\n    fig, axes = plt.subplots(1, 5, figsize=(22, 5))\n    for ax, (img, title) in zip(axes, stages):\n        ax.imshow(img)\n        ax.set_title(title, fontsize=9, fontweight=\"bold\", pad=6)\n        ax.axis(\"off\")\n    for spine in axes[4].spines.values():\n        spine.set_edgecolor(\"green\"); spine.set_linewidth(3)\n    plt.suptitle(\n        f\"Preprocessing Pipeline . {os.path.basename(str(image_path))}\\n\"\n        f\"Cache: {raw.nbytes/1024**2:.1f} MB raw PNG -> {green.nbytes/1024:.0f} KB .npy  \"\n        f\"({raw.nbytes//green.nbytes}x smaller per DataLoader worker)\",\n        fontsize=11, fontweight=\"bold\", y=1.03)\n    plt.tight_layout(); plt.show()\n    return green\n\nif os.path.exists(CFG.TRAIN_IMG):\n    sample = os.path.join(CFG.TRAIN_IMG, os.listdir(CFG.TRAIN_IMG)[0])\n    green_checkpoint = visualise_preprocessing_pipeline(sample)\nelse:\n    print(f\"Path not found: {CFG.TRAIN_IMG}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T11:48:20.89976Z","iopub.execute_input":"2026-05-22T11:48:20.90018Z","iopub.status.idle":"2026-05-22T11:48:21.763131Z","shell.execute_reply.started":"2026-05-22T11:48:20.900157Z","shell.execute_reply":"2026-05-22T11:48:21.762368Z"}},"outputs":[],"execution_count":null},{"id":"d7c8f441","cell_type":"markdown","source":"## Cell 6 — Pre-Cache All Images  *(runs once, ~4-6 min, permanently fixes the 1-hour bottleneck)*\n\n**Why this is the most important cell:**\n\n| | Before (original) | After (this fix) |\n|--|--|--|\n| When preprocessing happens | Every `__getitem__` call | **Once**, here, before any training |\n| Preprocessing calls total | `3662 x 20 epochs = 73,240` | **3,662** |\n| Data read per worker per image | 20 MB raw PNG | 264 KB `.npy` file |\n| Intermediate numpy arrays per image in worker | ~6 arrays x ~6 MB = ~36 MB | **0** |\n| Expected epoch time | ~60 min | **~2-3 min** |\n\nSafe to re-run: skips images already cached.\n","metadata":{}},{"id":"6e7e4950","cell_type":"code","source":"def build_cache(df_combined, img_dir, size=CFG.IMG_SIZE):\n    ids        = df_combined[\"id_code\"].unique().tolist()\n    to_process = [i for i in ids\n                  if not os.path.exists(os.path.join(CFG.CACHE_DIR, f\"{i}.npy\"))]\n\n    if not to_process:\n        print(f\"Cache already complete: {len(ids)} files exist in {CFG.CACHE_DIR}\")\n        return\n\n    print(f\"Caching {len(to_process)}/{len(ids)} images -> {CFG.CACHE_DIR}\")\n    print(f\"Estimated disk: {len(to_process) * size * size * 3 / 1024**2:.0f} MB\")\n\n    t0 = time.time(); errors = []\n    for img_id in tqdm(to_process, desc=\"Preprocessing\", unit=\"img\"):\n        cache_path = os.path.join(CFG.CACHE_DIR, f\"{img_id}.npy\")\n        raw_path   = os.path.join(img_dir, f\"{img_id}.png\")\n        try:\n            arr = preprocess_fundus_image(raw_path, size)\n            np.save(cache_path, arr)\n        except Exception as e:\n            errors.append((img_id, str(e)))\n\n    elapsed = time.time() - t0\n    n_cached = len([f for f in os.listdir(CFG.CACHE_DIR) if f.endswith(\".npy\")])\n    disk_mb  = sum(os.path.getsize(os.path.join(CFG.CACHE_DIR, f))\n                   for f in os.listdir(CFG.CACHE_DIR) if f.endswith(\".npy\")) / 1024**2\n    print(f\"\\nDone in {elapsed/60:.1f} min  |  {n_cached} files  |  {disk_mb:.0f} MB on disk\")\n    print(f\"Avg speed: {elapsed/max(len(to_process),1)*1000:.0f} ms/image\")\n    if errors:\n        print(f\"Errors: {len(errors)}\")\n        for img_id, msg in errors[:5]:\n            print(f\"  {img_id}: {msg}\")\n\ndf_train = pd.read_csv(CFG.TRAIN_CSV)\ndf_test  = pd.read_csv(CFG.TEST_CSV)\ndf_test[\"diagnosis\"] = -1\n\nprint(\"=== TRAIN ===\")\nbuild_cache(df_train, CFG.TRAIN_IMG)\nprint(\"\\n=== TEST ===\")\nbuild_cache(df_test,  CFG.TEST_IMG)\n\nprint(f\"\\n{'='*56}\")\nprint(\"Cache complete. DataLoader workers now load .npy files.\")\nprint(f\"Worker RAM/image : ~264 KB  (was ~20,000 KB = 75x reduction)\")\nprint(f\"{'='*56}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T11:48:21.764109Z","iopub.execute_input":"2026-05-22T11:48:21.764371Z","iopub.status.idle":"2026-05-22T12:05:00.735488Z","shell.execute_reply.started":"2026-05-22T11:48:21.764349Z","shell.execute_reply":"2026-05-22T12:05:00.734704Z"}},"outputs":[],"execution_count":null},{"id":"21f242a2","cell_type":"markdown","source":"## Cell 7 — Dataset & Augmentation","metadata":{}},{"id":"fb79a36b","cell_type":"code","source":"def get_augmentation(phase):\n    if phase == \"train\":\n        return A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.Rotate(limit=30, p=0.6),\n            A.RandomBrightnessContrast(p=0.3),\n            A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.1, rotate_limit=15, p=0.4),\n            A.CoarseDropout(max_holes=8, max_height=16, max_width=16, p=0.2),\n            A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n            ToTensorV2(),\n        ])\n    return A.Compose([\n        A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n        ToTensorV2(),\n    ])\n\n\nclass APTOSDataset(Dataset):\n    def __init__(self, df, phase=\"train\"):\n        self.df        = df.reset_index(drop=True)\n        self.phase     = phase\n        self.transform = get_augmentation(phase)\n\n    def __len__(self):\n        return len(self.df)\n\n    def _load(self, idx):\n        # FIX: load from .npy cache -- np.load is ~0.5 ms vs ~70 ms for raw preprocessing\n        img_id     = self.df.iloc[idx][\"id_code\"]\n        cache_path = os.path.join(CFG.CACHE_DIR, f\"{img_id}.npy\")\n        if os.path.exists(cache_path):\n            return np.load(cache_path)          # fast path (always, after Cell 6)\n        # Fallback: live preprocessing if cache missing (should not happen)\n        img_dir = CFG.TEST_IMG if self.df.iloc[idx].get(\"diagnosis\",-1) == -1 else CFG.TRAIN_IMG\n        return preprocess_fundus_image(os.path.join(img_dir, f\"{img_id}.png\"))\n\n    def __getitem__(self, idx):\n        img   = self._load(idx)\n        label = int(self.df.iloc[idx][\"diagnosis\"]) if \"diagnosis\" in self.df.columns else -1\n        aug   = self.transform(image=img)\n        return aug[\"image\"], torch.tensor(label, dtype=torch.long)\n\n\ndef build_weighted_sampler(labels):\n    counts = np.bincount(labels)\n    w      = 1.0 / counts.astype(np.float32)\n    sw     = torch.from_numpy(w[labels])\n    return WeightedRandomSampler(sw, len(sw), replacement=True)\n\n\n# Benchmark: how fast is __getitem__ now?\n_df_tmp = pd.read_csv(CFG.TRAIN_CSV).iloc[:30]\n_ds_tmp = APTOSDataset(_df_tmp, \"val\")\nt0 = time.time()\nfor i in range(30): _ = _ds_tmp._load(i)\nms_per = (time.time()-t0)/30*1000\ndel _ds_tmp, _df_tmp\n\nprint(f\"APTOSDataset defined\")\nprint(f\"Load speed from cache : {ms_per:.1f} ms/image\")\nprint(f\"  (raw preprocessing was ~70 ms/image = {70/ms_per:.0f}x faster now)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T12:05:00.736703Z","iopub.execute_input":"2026-05-22T12:05:00.737101Z","iopub.status.idle":"2026-05-22T12:05:00.768111Z","shell.execute_reply.started":"2026-05-22T12:05:00.737073Z","shell.execute_reply":"2026-05-22T12:05:00.767208Z"}},"outputs":[],"execution_count":null},{"id":"9ce11448","cell_type":"markdown","source":"## Cell 8 — CBAM Attention Module","metadata":{}},{"id":"76291dc2","cell_type":"code","source":"class ChannelAttention(nn.Module):\n    def __init__(self, in_channels, reduction_ratio=16):\n        super().__init__()\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.max_pool = nn.AdaptiveMaxPool2d(1)\n        mid = max(in_channels // reduction_ratio, 1)\n        self.mlp = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(in_channels, mid, bias=False),\n            nn.ReLU(inplace=True),\n            nn.Linear(mid, in_channels, bias=False),\n        )\n    def forward(self, x):\n        s = torch.sigmoid(self.mlp(self.avg_pool(x)) + self.mlp(self.max_pool(x)))\n        return x * s.unsqueeze(-1).unsqueeze(-1)\n\nclass SpatialAttention(nn.Module):\n    def __init__(self, kernel_size=7):\n        super().__init__()\n        self.conv = nn.Conv2d(2, 1, kernel_size, padding=kernel_size//2, bias=False)\n        self.bn   = nn.BatchNorm2d(1)\n    def forward(self, x):\n        avg = x.mean(dim=1, keepdim=True)\n        mx  = x.max(dim=1, keepdim=True).values\n        return x * torch.sigmoid(self.bn(self.conv(torch.cat([avg, mx], dim=1))))\n\nclass CBAM(nn.Module):\n    def __init__(self, in_channels, reduction_ratio=16, kernel_size=7):\n        super().__init__()\n        self.channel_att = ChannelAttention(in_channels, reduction_ratio)\n        self.spatial_att = SpatialAttention(kernel_size)\n    def forward(self, x):\n        return self.spatial_att(self.channel_att(x))\n\n_x = torch.randn(2, 96, 19, 19)\nprint(f\"CBAM input  : {tuple(_x.shape)}\")\nprint(f\"CBAM output : {tuple(CBAM(96)(_x).shape)}  (shape preserved - attention only scales)\")\ndel _x\nprint(\"CBAM defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T12:05:00.769092Z","iopub.execute_input":"2026-05-22T12:05:00.769609Z","iopub.status.idle":"2026-05-22T12:05:00.972462Z","shell.execute_reply.started":"2026-05-22T12:05:00.769586Z","shell.execute_reply":"2026-05-22T12:05:00.971549Z"}},"outputs":[],"execution_count":null},{"id":"75faee89","cell_type":"markdown","source":"## Cell 9 — ConvNeXt-Tiny + CBAM Model","metadata":{}},{"id":"be199b0e","cell_type":"code","source":"class ConvNeXtCBAM(nn.Module):\n    \"\"\"\n    ConvNeXt-Tiny backbone with CBAM attention injected after every stage.\n\n    ConvNeXt-Tiny stage channels : [96, 192, 384, 768]\n    Feature dim (after pool)     : 768\n    Dual head                    : CE classifier + CORN ordinal regression\n\n    Key differences from EfficientNet-B3 version:\n      - 4 CBAM modules (one per ConvNeXt stage) instead of 4 on specific blocks\n      - LayerNorm(768) after global-average-pool (ConvNeXt standard)\n      - ConvNeXt blocks are channel-last internally; CBAM sees BCHW at stage output\n    \"\"\"\n    _STAGE_CH = [96, 192, 384, 768]\n    _HEAD_CH  = 768\n\n    def __init__(self, num_classes=CFG.NUM_CLASSES, pretrained=True, dropout=0.4):\n        super().__init__()\n        bb = timm.create_model(\"convnext_tiny\", pretrained=pretrained)\n\n        # Backbone components\n        self.stem   = bb.stem    # Patchify stem: 4x4 conv, output 96 ch\n        self.stages = bb.stages  # 4 ConvNeXtStage objects\n\n        # Head normalisation: ConvNeXt applies LayerNorm AFTER global pool.\n        # Copy pretrained weights so the norm starts from a good initialisation.\n        self.head_norm = nn.LayerNorm(self._HEAD_CH)\n        if hasattr(bb, \"head\") and hasattr(bb.head, \"norm\") and bb.head.norm is not None:\n            self.head_norm.load_state_dict(bb.head.norm.state_dict())\n\n        # CBAM after each stage (96, 192, 384, 768 channels)\n        self.cbam0 = CBAM(self._STAGE_CH[0])\n        self.cbam1 = CBAM(self._STAGE_CH[1])\n        self.cbam2 = CBAM(self._STAGE_CH[2])\n        self.cbam3 = CBAM(self._STAGE_CH[3])\n\n        self.global_pool = nn.AdaptiveAvgPool2d(1)\n        self.dropout     = nn.Dropout(p=dropout)\n        self.classifier  = nn.Linear(self._HEAD_CH, num_classes)\n        self.corn_head   = nn.Linear(self._HEAD_CH, num_classes - 1)\n\n        for layer in (self.classifier, self.corn_head):\n            nn.init.kaiming_normal_(layer.weight)\n            nn.init.zeros_(layer.bias)\n\n    def forward_features(self, x):\n        x = self.stem(x)   # (B, 96, H/4, W/4)\n        for stage, cbam in zip(\n            self.stages,\n            [self.cbam0, self.cbam1, self.cbam2, self.cbam3]\n        ):\n            x = stage(x)   # ConvNeXtStage handles downsampling internally\n            x = cbam(x)    # CBAM on BCHW output\n        x = self.global_pool(x).flatten(1)  # (B, 768)\n        x = self.head_norm(x)               # LayerNorm — ConvNeXt standard\n        return self.dropout(x)\n\n    def forward(self, x):\n        f = self.forward_features(x)\n        return self.classifier(f), self.corn_head(f)\n\n\n# Sanity check\n_m = ConvNeXtCBAM(pretrained=False).to(DEVICE)\n_d = torch.randn(2, 3, CFG.IMG_SIZE, CFG.IMG_SIZE).to(DEVICE)\nwith torch.no_grad():\n    _lo, _co = _m(_d)\ntotal = sum(p.numel() for p in _m.parameters())\nprint(f\"Backbone        : ConvNeXt-Tiny\")\nprint(f\"Input shape     : {tuple(_d.shape)}\")\nprint(f\"CE  output shape: {tuple(_lo.shape)}\")\nprint(f\"CORN output shape: {tuple(_co.shape)}\")\nprint(f\"Total parameters: {total:,}\")\ndel _m, _d, _lo, _co\ntorch.cuda.empty_cache(); gc.collect()\nprint(\"Sanity check passed  (test model freed)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T12:05:00.973703Z","iopub.execute_input":"2026-05-22T12:05:00.974124Z","iopub.status.idle":"2026-05-22T12:05:05.74676Z","shell.execute_reply.started":"2026-05-22T12:05:00.974088Z","shell.execute_reply":"2026-05-22T12:05:05.745742Z"}},"outputs":[],"execution_count":null},{"id":"096d6b5e","cell_type":"markdown","source":"## Cell 10 — Loss Functions","metadata":{}},{"id":"c3c46c0c","cell_type":"code","source":"def corn_loss(logits, targets, num_classes):\n    # FIX 5: use torch.stack instead of += to avoid accumulating graph nodes\n    # across the loop. Each F.bce call is independent; stack+mean is cleaner.\n    K     = num_classes - 1\n    terms = []\n    for i in range(K):\n        sub = (targets > (i-1)) if i > 0 else torch.ones(\n            len(targets), dtype=torch.bool, device=logits.device)\n        if sub.sum() == 0:\n            continue\n        terms.append(F.binary_cross_entropy_with_logits(\n            logits[sub, i], (targets[sub] > i).float()))\n    if not terms:\n        return torch.tensor(0.0, device=logits.device)\n    return torch.stack(terms).mean()\n\nclass CombinedLoss(nn.Module):\n    def __init__(self, class_weights, num_classes=CFG.NUM_CLASSES, alpha=CFG.CORN_WEIGHT):\n        super().__init__()\n        self.num_classes   = num_classes\n        self.alpha         = alpha\n        self.class_weights = class_weights\n    def forward(self, logits, corn_logits, targets):\n        l_corn = corn_loss(corn_logits, targets, self.num_classes)\n        l_wce  = nn.CrossEntropyLoss(weight=self.class_weights.to(logits.device))(logits, targets)\n        return self.alpha * l_corn + (1 - self.alpha) * l_wce\n\nprint(\"Loss functions defined  (alpha*CORN + (1-alpha)*WCE)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T12:05:05.748649Z","iopub.execute_input":"2026-05-22T12:05:05.749157Z","iopub.status.idle":"2026-05-22T12:05:05.757041Z","shell.execute_reply.started":"2026-05-22T12:05:05.749132Z","shell.execute_reply":"2026-05-22T12:05:05.756271Z"}},"outputs":[],"execution_count":null},{"id":"b6932dc1","cell_type":"markdown","source":"## Cell 11 — Evaluation Metrics","metadata":{}},{"id":"090d8442","cell_type":"code","source":"def compute_qwk(y_true, y_pred):\n    return cohen_kappa_score(y_true, y_pred, weights=\"quadratic\")\n\ndef corn_to_preds(corn_logits):\n    return (torch.sigmoid(corn_logits) > 0.5).long().sum(dim=1)\n\ndef print_per_class_metrics(y_true, y_pred):\n    names = [f\"Grade {i}\" for i in range(CFG.NUM_CLASSES)]\n    print(\"\\n\" + \"=\"*60)\n    print(\"PER-CLASS PERFORMANCE\")\n    print(\"=\"*60)\n    print(classification_report(y_true, y_pred,\n                                 labels=list(range(CFG.NUM_CLASSES)),\n                                 target_names=names, digits=4))\n\nprint(\"Metrics defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T12:05:05.758033Z","iopub.execute_input":"2026-05-22T12:05:05.758977Z","iopub.status.idle":"2026-05-22T12:05:05.774537Z","shell.execute_reply.started":"2026-05-22T12:05:05.758943Z","shell.execute_reply":"2026-05-22T12:05:05.773835Z"}},"outputs":[],"execution_count":null},{"id":"f519eb41","cell_type":"markdown","source":"## Cell 12 — Train / Validate Loops","metadata":{}},{"id":"6120b812","cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, criterion, scaler, scheduler):\n    model.train()\n    run_loss, preds_all, labels_all = 0.0, [], []\n\n    pbar = tqdm(loader, desc=\"  train\", leave=False, unit=\"batch\")\n    for images, labels in pbar:\n        images = images.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n        with torch.cuda.amp.autocast():\n            logits, corn_logits = model(images)\n            loss = criterion(logits, corn_logits, labels)\n\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        scaler.step(optimizer); scaler.update()\n\n        run_loss += loss.item() * images.size(0)\n        preds_all.extend(logits.detach().argmax(1).cpu().tolist())\n        labels_all.extend(labels.cpu().tolist())\n        pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n\n    if scheduler: scheduler.step()\n    return {\"loss\": run_loss / len(loader.dataset),\n            \"qwk\":  compute_qwk(labels_all, preds_all)}\n\n@torch.no_grad()\ndef validate(model, loader, criterion):\n    model.eval()\n    run_loss, preds_all, labels_all = 0.0, [], []\n\n    for images, labels in tqdm(loader, desc=\"  val  \", leave=False, unit=\"batch\"):\n        images = images.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n        with torch.cuda.amp.autocast():\n            logits, corn_logits = model(images)\n            loss = criterion(logits, corn_logits, labels)\n\n        run_loss += loss.item() * images.size(0)\n        preds = ((logits.argmax(1).float() + corn_to_preds(corn_logits).float()) / 2).round().long()\n        preds_all.extend(preds.cpu().tolist())\n        labels_all.extend(labels.cpu().tolist())\n\n    return {\"loss\":   run_loss / len(loader.dataset),\n            \"qwk\":    compute_qwk(labels_all, preds_all),\n            \"preds\":  preds_all,\n            \"labels\": labels_all}\n\nprint(\"Train / validate loops defined  (tqdm progress bars added)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T12:05:05.775438Z","iopub.execute_input":"2026-05-22T12:05:05.775698Z","iopub.status.idle":"2026-05-22T12:05:05.792551Z","shell.execute_reply.started":"2026-05-22T12:05:05.775676Z","shell.execute_reply":"2026-05-22T12:05:05.79187Z"}},"outputs":[],"execution_count":null},{"id":"6bdb75e9","cell_type":"markdown","source":"## Cell 13 — Data Loading","metadata":{}},{"id":"6e0d652a","cell_type":"code","source":"df = pd.read_csv(CFG.TRAIN_CSV)\nprint(f\"Total samples : {len(df)}\")\nprint(f\"\\nClass distribution:\")\nprint(df[\"diagnosis\"].value_counts().sort_index().to_string())\n\ntrain_df, val_df = train_test_split(\n    df, test_size=CFG.VAL_SPLIT, stratify=df[\"diagnosis\"], random_state=SEED)\nprint(f\"\\nTrain : {len(train_df)}  |  Val : {len(val_df)}\")\n\ntrain_dataset = APTOSDataset(train_df, phase=\"train\")\nval_dataset   = APTOSDataset(val_df,   phase=\"val\")\nsampler       = build_weighted_sampler(train_df[\"diagnosis\"].tolist())\n\n# FIX 4/6: persistent_workers avoids forking overhead; prefetch_factor=1 limits pinned-memory accumulation\n_loader_kw = dict(\n    num_workers        = CFG.NUM_WORKERS,\n    pin_memory         = True,\n    persistent_workers = (CFG.NUM_WORKERS > 0),\n    prefetch_factor    = 1 if CFG.NUM_WORKERS > 0 else None,  # FIX 6: was 2; 4 workers x 2 batches x ~164 MB pinned = ~1.3 GB that never freed between epochs\n)\ntrain_loader = DataLoader(train_dataset, batch_size=CFG.BATCH_SIZE, sampler=sampler, **_loader_kw)\nval_loader   = DataLoader(val_dataset,   batch_size=CFG.BATCH_SIZE, shuffle=False,   **_loader_kw)\n\nclass_counts  = np.bincount(train_df[\"diagnosis\"].tolist(), minlength=CFG.NUM_CLASSES)\nclass_weights = torch.tensor(1.0/(class_counts/class_counts.sum()), dtype=torch.float32)\nclass_weights = class_weights / class_weights.sum() * CFG.NUM_CLASSES\nprint(f\"\\nClass weights : {class_weights.numpy().round(4)}\")\nprint(f\"Train batches : {len(train_loader)}  |  Val batches : {len(val_loader)}\")\n\n# How fast is one batch now?\nt0 = time.time()\n_img, _lbl = next(iter(train_loader))\nms = (time.time()-t0)*1000\nprint(f\"\\nFirst batch : {tuple(_img.shape)}  loaded in {ms:.0f} ms  (was several minutes)\")\ndel _img, _lbl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T12:05:05.793565Z","iopub.execute_input":"2026-05-22T12:05:05.793872Z","iopub.status.idle":"2026-05-22T12:05:06.1768Z","shell.execute_reply.started":"2026-05-22T12:05:05.79385Z","shell.execute_reply":"2026-05-22T12:05:06.175345Z"}},"outputs":[],"execution_count":null},{"id":"6f34eb82","cell_type":"markdown","source":"## Cell 14 — Model, Optimizer & Scheduler","metadata":{}},{"id":"20c81f2a","cell_type":"code","source":"model = ConvNeXtCBAM(num_classes=CFG.NUM_CLASSES, pretrained=True).to(DEVICE)\n\nif N_GPUS > 1:\n    model = nn.DataParallel(model, device_ids=list(range(N_GPUS)))\n    print(f\"DataParallel across {N_GPUS} GPUs\")\nelse:\n    print(\"Single GPU mode\")\n\n# ConvNeXt head keys: CBAMs + head norm + classifier heads\n_HEAD_KEYS = (\"cbam0\", \"cbam1\", \"cbam2\", \"cbam3\",\n               \"head_norm\", \"classifier\", \"corn_head\")\n\n_m = model.module if hasattr(model, \"module\") else model\nbackbone_p = [p for n,p in _m.named_parameters() if not any(k in n for k in _HEAD_KEYS)]\nhead_p     = [p for n,p in _m.named_parameters() if     any(k in n for k in _HEAD_KEYS)]\n\noptimizer = torch.optim.AdamW([\n    {\"params\": backbone_p, \"lr\": CFG.LR * 0.1},\n    {\"params\": head_p,     \"lr\": CFG.LR},\n], weight_decay=CFG.WEIGHT_DECAY)\n\nscheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, T_max=CFG.EPOCHS, eta_min=1e-6)\n\ncriterion = CombinedLoss(class_weights=class_weights, num_classes=CFG.NUM_CLASSES).to(DEVICE)\nscaler    = torch.cuda.amp.GradScaler()\n\ntotal     = sum(p.numel() for p in model.parameters())\ntrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"\\nBackbone params  : {sum(p.numel() for p in backbone_p):,}\")\nprint(f\"Head params      : {sum(p.numel() for p in head_p):,}\")\nprint(f\"Total params     : {total:,}\")\nprint(f\"Trainable params : {trainable:,}\")\nprint(f\"Freeze epochs    : {CFG.FREEZE_EPOCHS}\")\n\ndef freeze_backbone(model):\n    m = model.module if hasattr(model,\"module\") else model\n    for n,p in m.named_parameters():\n        p.requires_grad = any(k in n for k in _HEAD_KEYS)\n    frozen = sum(p.numel() for p in m.parameters() if not p.requires_grad)\n    live   = sum(p.numel() for p in m.parameters() if     p.requires_grad)\n    print(f\"  [FREEZE]   backbone {frozen:,} frozen  |  head {live:,} trainable\")\n\ndef unfreeze_all(model, optimizer):\n    m = model.module if hasattr(model,\"module\") else model\n    for p in m.parameters(): p.requires_grad = True\n    bp = [p for n,p in m.named_parameters() if not any(k in n for k in _HEAD_KEYS)]\n    hp = [p for n,p in m.named_parameters() if     any(k in n for k in _HEAD_KEYS)]\n    optimizer.param_groups.clear()\n    optimizer.add_param_group({\"params\": bp, \"lr\": CFG.LR * 0.1})\n    optimizer.add_param_group({\"params\": hp, \"lr\": CFG.LR})\n    print(f\"  [UNFREEZE] all {sum(p.numel() for p in m.parameters()):,} params trainable\")\n\nfreeze_backbone(model)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T12:05:06.18513Z","iopub.execute_input":"2026-05-22T12:05:06.185489Z","iopub.status.idle":"2026-05-22T12:05:08.979768Z","shell.execute_reply.started":"2026-05-22T12:05:06.185451Z","shell.execute_reply":"2026-05-22T12:05:08.978813Z"}},"outputs":[],"execution_count":null},{"id":"817198d1","cell_type":"markdown","source":"## Cell 15 — Training Loop  (Two-Stage + Timed)","metadata":{}},{"id":"7ad1056d","cell_type":"code","source":"best_qwk, best_weights, history = -float(\"inf\"), None, []\n\nprint(f\"Two-stage: freeze {CFG.FREEZE_EPOCHS} ep -> full fine-tune\")\nhdr = f\"{'Ep':>4}  {'Phase':>10}  {'TrLoss':>8}  {'TrQWK':>7}  {'VlLoss':>8}  {'VlQWK':>7}  {'Time':>6}\"\nprint(f\"\\n{hdr}\\n\" + \"-\"*68)\n\nfor epoch in range(1, CFG.EPOCHS + 1):\n    if epoch == CFG.FREEZE_EPOCHS + 1:\n        unfreeze_all(model, optimizer)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, T_max=CFG.EPOCHS - CFG.FREEZE_EPOCHS, eta_min=1e-6)\n\n    phase = \"head-only\" if epoch <= CFG.FREEZE_EPOCHS else \"full-ft\"\n    t0    = time.time()\n\n    tr = train_one_epoch(model, train_loader, optimizer, criterion, scaler, scheduler)\n    vl = validate(model, val_loader, criterion)\n\n    elapsed = time.time() - t0\n    is_best = vl[\"qwk\"] > best_qwk\n\n    if is_best:\n        best_qwk     = vl[\"qwk\"]\n        m_sd         = model.module if hasattr(model,\"module\") else model\n        best_weights = {k: v.cpu().clone() for k,v in m_sd.state_dict().items()}\n        torch.save(best_weights, \"/kaggle/working/best_model.pt\")\n\n    history.append({\"epoch\":epoch, \"phase\":phase,\n                    \"train_loss\":tr[\"loss\"], \"train_qwk\":tr[\"qwk\"],\n                    \"val_loss\":vl[\"loss\"],   \"val_qwk\":vl[\"qwk\"],\n                    \"time_s\":elapsed})\n\n    flag = \" BEST\" if is_best else \"\"\n    print(f\"{epoch:>4}  {phase:>10}  {tr['loss']:>8.4f}  {tr['qwk']:>7.4f}\"\n          f\"  {vl['loss']:>8.4f}  {vl['qwk']:>7.4f}  {elapsed:>5.0f}s{flag}\")\n\n    # VRAM check every 5 epochs\n    if epoch % 5 == 0 and N_GPUS > 0:\n        for gi in range(N_GPUS):\n            used = torch.cuda.memory_allocated(gi)/1024**3\n            res  = torch.cuda.memory_reserved(gi)/1024**3\n            print(f\"       GPU{gi} VRAM: {used:.2f} GB used / {res:.2f} GB reserved\")\n\n    # FIX 5: flush CUDA allocator cache between epochs\n    # Without this, reserved-but-freed tensors from epoch N pile up\n    # and the allocator runs out of contiguous space at epoch N+1\n    torch.cuda.empty_cache()\n    gc.collect()\n\navg_s = sum(h[\"time_s\"] for h in history) / len(history)\nprint(f\"\\nBest Validation QWK : {best_qwk:.4f}\")\nprint(f\"Avg epoch time      : {avg_s:.0f}s ({avg_s/60:.1f} min)  <- was ~60 min before caching\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-05-22T12:05:08.980964Z","iopub.execute_input":"2026-05-22T12:05:08.981597Z","execution_failed":"2026-05-22T13:01:05.775Z"}},"outputs":[],"execution_count":null},{"id":"4dfb2e83","cell_type":"markdown","source":"## Cell 16 — Final Evaluation","metadata":{}},{"id":"c46bf820","cell_type":"code","source":"_m = model.module if hasattr(model,\"module\") else model\n_m.load_state_dict(best_weights)\n\nval_stats = validate(model, val_loader, criterion)\nqwk       = compute_qwk(val_stats[\"labels\"], val_stats[\"preds\"])\nprint(f\"Final Validation QWK : {qwk:.4f}\")\nprint_per_class_metrics(val_stats[\"labels\"], val_stats[\"preds\"])","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-22T13:01:05.777Z"}},"outputs":[],"execution_count":null},{"id":"d875b561","cell_type":"markdown","source":"## Cell 17 — Test Inference & Submission","metadata":{}},{"id":"b9dc728c","cell_type":"code","source":"test_df              = pd.read_csv(CFG.TEST_CSV)\ntest_df[\"diagnosis\"] = -1\ntest_dataset = APTOSDataset(test_df, phase=\"val\")\ntest_loader  = DataLoader(test_dataset, batch_size=CFG.BATCH_SIZE, shuffle=False, **_loader_kw)\n\nmodel.eval(); all_preds = []\nwith torch.no_grad():\n    for images, _ in tqdm(test_loader, desc=\"inference\", unit=\"batch\"):\n        images = images.to(DEVICE, non_blocking=True)\n        with torch.cuda.amp.autocast():\n            logits, corn_logits = model(images)\n        preds = ((logits.argmax(1).float() + corn_to_preds(corn_logits).float())/2).round().long()\n        all_preds.extend(preds.cpu().tolist())\n\nsubmission = pd.DataFrame({\"id_code\": test_df[\"id_code\"], \"diagnosis\": all_preds})\nsubmission.to_csv(\"/kaggle/working/submission.csv\", index=False)\nprint(f\"Saved -- shape: {submission.shape}\")\nprint(submission[\"diagnosis\"].value_counts().sort_index().to_string())","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-22T13:01:05.778Z"}},"outputs":[],"execution_count":null},{"id":"c59ffd30","cell_type":"markdown","source":"## Cell 18 — Training History & Time Analysis","metadata":{}},{"id":"136cb6b6","cell_type":"code","source":"%matplotlib inline\nhist_df = pd.DataFrame(history)\nfig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\naxes[0].plot(hist_df[\"epoch\"], hist_df[\"train_loss\"], \"o-\", label=\"Train\")\naxes[0].plot(hist_df[\"epoch\"], hist_df[\"val_loss\"],   \"s-\", label=\"Val\")\naxes[0].axvline(CFG.FREEZE_EPOCHS+0.5, color=\"gray\", linestyle=\"--\", alpha=0.6, label=\"Unfreeze\")\naxes[0].set_title(\"Loss Curve\"); axes[0].set_xlabel(\"Epoch\"); axes[0].legend(); axes[0].grid(alpha=0.3)\n\naxes[1].plot(hist_df[\"epoch\"], hist_df[\"train_qwk\"], \"o-\", label=\"Train\")\naxes[1].plot(hist_df[\"epoch\"], hist_df[\"val_qwk\"],   \"s-\", label=\"Val\")\naxes[1].axhline(best_qwk, color=\"red\", linestyle=\"--\", label=f\"Best={best_qwk:.4f}\")\naxes[1].axvline(CFG.FREEZE_EPOCHS+0.5, color=\"gray\", linestyle=\"--\", alpha=0.6)\naxes[1].set_title(\"QWK\"); axes[1].set_xlabel(\"Epoch\"); axes[1].legend(); axes[1].grid(alpha=0.3)\n\naxes[2].bar(hist_df[\"epoch\"], hist_df[\"time_s\"]/60, color=\"steelblue\", alpha=0.7)\naxes[2].axhline(hist_df[\"time_s\"].mean()/60, color=\"red\", linestyle=\"--\",\n                label=f\"Avg {hist_df['time_s'].mean()/60:.1f} min/epoch\")\naxes[2].set_title(\"Epoch Time (min)\"); axes[2].set_xlabel(\"Epoch\")\naxes[2].legend(); axes[2].grid(alpha=0.3)\n\nplt.suptitle(\"ConvNeXt-Tiny + CBAM  .  APTOS 2019  .  ConvNeXt Variant\",\n             fontsize=13, fontweight=\"bold\")\nplt.tight_layout()\nplt.savefig(\"/kaggle/working/training_history.png\", dpi=150, bbox_inches=\"tight\")\nplt.show()\n\ntotal_h   = hist_df[\"time_s\"].sum() / 3600\nbefore_h  = CFG.EPOCHS * 60 / 60   # 1 hour per epoch\nprint(f\"Total training time : {total_h:.1f} h  (was ~{before_h:.0f} h before fix)\")\nprint(f\"Time saved          : {before_h - total_h:.1f} h  ({(before_h-total_h)/before_h*100:.0f}%)\")","metadata":{"trusted":true,"execution":{"execution_failed":"2026-05-22T13:01:05.778Z"}},"outputs":[],"execution_count":null}]}