{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"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"},"accelerator":"GPU","colab":{"gpuType":"T4","provenance":[]},"widgets":{"application/vnd.jupyter.widget-state+json":{"3143cd7b879b4959a9a360150a8a2e1a":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"420ea6820e354cf4a46d8cae8587c727":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_6fb918b09ea94dfeb4909d761d23ac90","IPY_MODEL_dbb17dec4d784354870409bab8c013e8","IPY_MODEL_930151ed68c74cc7b3b70e4f3bcd53c3"],"layout":"IPY_MODEL_3143cd7b879b4959a9a360150a8a2e1a"}},"6fb918b09ea94dfeb4909d761d23ac90":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_e4880b158a7a4d61a8351a6e4968fe18","placeholder":"​","style":"IPY_MODEL_af7a09d20cd24692b20f8e2792ef8777","value":"model.safetensors: 100%"}},"706a72ba6cb94e99a705a3cae914d360":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"8f21ca60ad3b482494ab9398c69d34b8":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"930151ed68c74cc7b3b70e4f3bcd53c3":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_a71652ba98754d418443ad03d6fff600","placeholder":"​","style":"IPY_MODEL_8f21ca60ad3b482494ab9398c69d34b8","value":" 60.6M/60.6M [00:14&lt;00:00, 4.62MB/s]"}},"a71652ba98754d418443ad03d6fff600":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"af7a09d20cd24692b20f8e2792ef8777":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"b0bdce41376a4e26ab425ea6e8a17cc2":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"dbb17dec4d784354870409bab8c013e8":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"success","description":"","description_tooltip":null,"layout":"IPY_MODEL_b0bdce41376a4e26ab425ea6e8a17cc2","max":60641304,"min":0,"orientation":"horizontal","style":"IPY_MODEL_706a72ba6cb94e99a705a3cae914d360","value":60641304}},"e4880b158a7a4d61a8351a6e4968fe18":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}}}},"kaggle":{"accelerator":"none","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431,"isSourceIdPinned":false}],"dockerImageVersionId":31328,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# =============================================\n# Kaggle Setup - APTOS 2019 Blindness Detection\n# =============================================\n# Dataset: Add \"aptos2019-blindness-detection\" as a Kaggle dataset\n# Go to: Add Data -> Competition Data -> aptos2019-blindness-detection\n# The data will be at: /kaggle/input/aptos2019-blindness-detection/\n\nimport os\n\nDATA_DIR = '/kaggle/input/competitions/aptos2019-blindness-detection'\nSAVE_DIR = '/kaggle/working/'\n\n# Ensure output directory exists\nos.makedirs(SAVE_DIR, exist_ok=True)\n\nprint(f\"Data directory: {DATA_DIR}\")\nprint(f\"Output directory: {SAVE_DIR}\")\nprint(f\"Data files: {os.listdir(DATA_DIR)}\")","metadata":{"id":"hgeGLH-KpfWH","outputId":"7fba20f8-9163-481e-c0a9-0fe27174315b","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:15.987652Z","iopub.execute_input":"2026-04-07T21:48:15.988051Z","iopub.status.idle":"2026-04-07T21:48:15.999438Z","shell.execute_reply.started":"2026-04-07T21:48:15.988027Z","shell.execute_reply":"2026-04-07T21:48:15.998349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install timm\n\nimport os\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\nimport random\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix, classification_report, cohen_kappa_score\nfrom sklearn.utils.class_weight import compute_class_weight\n\nimport timm\nfrom tqdm import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# Set seeds for reproducibility\ndef set_seed(seed=42):\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\nset_seed(42)","metadata":{"id":"LcSee08-pfWK","outputId":"75b2c9a1-52f4-44e1-fb95-56758e0db666","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:16.001698Z","iopub.execute_input":"2026-04-07T21:48:16.002004Z","iopub.status.idle":"2026-04-07T21:48:36.242942Z","shell.execute_reply.started":"2026-04-07T21:48:16.001982Z","shell.execute_reply":"2026-04-07T21:48:36.241889Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    print(dirname)","metadata":{"id":"0WGiIKAQpfWK","outputId":"4cbdfd6d-b4a0-497b-b8c3-74ecaf6eb5a3","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:36.244000Z","iopub.execute_input":"2026-04-07T21:48:36.244416Z","iopub.status.idle":"2026-04-07T21:48:43.363453Z","shell.execute_reply.started":"2026-04-07T21:48:36.244393Z","shell.execute_reply":"2026-04-07T21:48:43.362511Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = pd.read_csv(os.path.join(DATA_DIR, \"train.csv\"))\nprint(f\"Total samples: {len(df)}\")\nprint(f\"Class distribution:\\n{df['diagnosis'].value_counts().sort_index()}\")\n\ntrain_df, temp_df = train_test_split(\n    df, test_size=0.2, stratify=df[\"diagnosis\"], random_state=42\n)\n\nval_df, test_df = train_test_split(\n    temp_df, test_size=0.5, stratify=temp_df[\"diagnosis\"], random_state=42\n)\n\nprint(f\"\\nTrain: {len(train_df)}, Val: {len(val_df)}, Test: {len(test_df)}\")","metadata":{"id":"KOhviK_ppfWK","outputId":"4a66e2cd-68f2-436a-e843-371d91019dea","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:43.364952Z","iopub.execute_input":"2026-04-07T21:48:43.365241Z","iopub.status.idle":"2026-04-07T21:48:43.430074Z","shell.execute_reply.started":"2026-04-07T21:48:43.365211Z","shell.execute_reply":"2026-04-07T21:48:43.429219Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_clahe(pil_img):\n    img = np.array(pil_img)\n    img = cv2.cvtColor(img, cv2.COLOR_RGB2LAB)\n\n    l, a, b = cv2.split(img)\n    clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8,8))\n    cl = clahe.apply(l)\n\n    merged = cv2.merge((cl,a,b))\n    img = cv2.cvtColor(merged, cv2.COLOR_LAB2RGB)\n\n    return Image.fromarray(img)","metadata":{"id":"K3rDn50hpfWL","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:43.432230Z","iopub.execute_input":"2026-04-07T21:48:43.432475Z","iopub.status.idle":"2026-04-07T21:48:43.438086Z","shell.execute_reply.started":"2026-04-07T21:48:43.432453Z","shell.execute_reply":"2026-04-07T21:48:43.437113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ImageNet normalization - CRITICAL for pretrained models\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD  = [0.229, 0.224, 0.225]\nIMG_SIZE = 224\n\ntrain_transform = transforms.Compose([\n    transforms.Lambda(apply_clahe),\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(10),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2, saturation=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\nval_transform = transforms.Compose([\n    transforms.Lambda(apply_clahe),\n    transforms.Resize((IMG_SIZE, IMG_SIZE)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\nprint(\"Transforms configured: 224x224 with conservative augmentation for CoAtNet\")","metadata":{"id":"OLW23fqGpfWL","outputId":"89c40d74-f2b7-4d2c-f4c6-851b67875080","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:43.439259Z","iopub.execute_input":"2026-04-07T21:48:43.439605Z","iopub.status.idle":"2026-04-07T21:48:43.460257Z","shell.execute_reply.started":"2026-04-07T21:48:43.439554Z","shell.execute_reply":"2026-04-07T21:48:43.459356Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class APTOSDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_name = self.df.iloc[idx][\"id_code\"]\n        label = self.df.iloc[idx][\"diagnosis\"]\n\n        # FIXED: Use full Kaggle data path instead of relative path\n        path = os.path.join(DATA_DIR, \"train_images\", f\"{img_name}.png\")\n        image = Image.open(path).convert(\"RGB\")\n\n        if self.transform:\n            image = self.transform(image)\n\n        return image, label","metadata":{"id":"WE-UzcmxpfWL","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:43.461249Z","iopub.execute_input":"2026-04-07T21:48:43.461464Z","iopub.status.idle":"2026-04-07T21:48:43.482688Z","shell.execute_reply.started":"2026-04-07T21:48:43.461443Z","shell.execute_reply":"2026-04-07T21:48:43.481514Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 16\nNUM_WORKERS = 4  # ⚡ OPTIMIZED from 2 → faster data loading (30-50% speedup)\n\ntrain_loader = DataLoader(\n    APTOSDataset(train_df, train_transform),\n    batch_size=BATCH_SIZE, shuffle=True,\n    num_workers=NUM_WORKERS, pin_memory=True, drop_last=True\n)\nval_loader = DataLoader(\n    APTOSDataset(val_df, val_transform),\n    batch_size=BATCH_SIZE,\n    num_workers=NUM_WORKERS, pin_memory=True\n)\ntest_loader = DataLoader(\n    APTOSDataset(test_df, val_transform),\n    batch_size=BATCH_SIZE,\n    num_workers=NUM_WORKERS, pin_memory=True\n)\n\nprint(f\"Train batches: {len(train_loader)}, Val batches: {len(val_loader)}, Test batches: {len(test_loader)}\")","metadata":{"id":"6r0wZ7D7pfWM","outputId":"0d788bdf-d99e-464a-ae3b-e61e6831d6a1","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:43.483887Z","iopub.execute_input":"2026-04-07T21:48:43.484478Z","iopub.status.idle":"2026-04-07T21:48:43.510902Z","shell.execute_reply.started":"2026-04-07T21:48:43.484445Z","shell.execute_reply":"2026-04-07T21:48:43.509834Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# =============================================\n# CoAtNet Model - smallest 224x224 variant\n# =============================================\n\n# Check available CoAtNet variants in current timm version\navailable_coatnet = [m for m in timm.list_models('coatnet*') if '224' in m]\nprint(f\"Available CoAtNet 224 variants: {available_coatnet}\")\n\n# Prefer coatnet_nano_rw_224; fallback to smallest available 224 variant\npreferred_model = \"coatnet_nano_rw_224\"\nif preferred_model in timm.list_models('coatnet*'):\n    model_name = preferred_model\nelse:\n    # Fallback: pick the smallest available CoAtNet 224 variant\n    all_coatnet = timm.list_models('coatnet*')\n    coatnet_224 = [m for m in all_coatnet if '224' in m]\n    if coatnet_224:\n        # Prefer nano/pico/tiny variants\n        for keyword in ['nano', 'pico', 'tiny', '0']:\n            matches = [m for m in coatnet_224 if keyword in m]\n            if matches:\n                model_name = matches[0]\n                break\n        else:\n            model_name = coatnet_224[0]\n    else:\n        # Last resort: any CoAtNet variant\n        model_name = all_coatnet[0] if all_coatnet else 'coatnet_0_rw_224'\n\nprint(f\"\\nUsing model: {model_name}\")\n\nmodel = timm.create_model(model_name, pretrained=True, num_classes=5).to(device)\n\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"Total params: {total_params:,}\")\nprint(f\"Trainable params: {trainable_params:,}\")","metadata":{"id":"sRECXy92pfWM","outputId":"fc5d0a64-bd1b-4194-eaad-306392a76a71","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:43.512021Z","iopub.execute_input":"2026-04-07T21:48:43.512265Z","iopub.status.idle":"2026-04-07T21:48:53.141290Z","shell.execute_reply.started":"2026-04-07T21:48:43.512243Z","shell.execute_reply":"2026-04-07T21:48:53.139764Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class_weights = compute_class_weight(\n    \"balanced\",\n    classes=np.unique(train_df[\"diagnosis\"]),\n    y=train_df[\"diagnosis\"]\n)\n\nclass_weights = torch.tensor(class_weights, dtype=torch.float).to(device)\nprint(f\"Class weights: {class_weights}\")","metadata":{"id":"NRFv_L-EpfWM","outputId":"c9f3dcbf-5208-4b72-e934-4d1e0680143d","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:53.142669Z","iopub.execute_input":"2026-04-07T21:48:53.142952Z","iopub.status.idle":"2026-04-07T21:48:53.186640Z","shell.execute_reply.started":"2026-04-07T21:48:53.142924Z","shell.execute_reply":"2026-04-07T21:48:53.185658Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def freeze_backbone(model):\n    \"\"\"Freeze all layers except the classification head.\"\"\"\n    for name, param in model.named_parameters():\n        if 'head' not in name:\n            param.requires_grad = False\n        else:\n            param.requires_grad = True\n    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Backbone frozen. Trainable params (head only): {trainable:,}\")\n\ndef unfreeze_all(model):\n    \"\"\"Unfreeze all layers for fine-tuning.\"\"\"\n    for param in model.parameters():\n        param.requires_grad = True\n    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"All layers unfrozen. Trainable params: {trainable:,}\")","metadata":{"id":"_8fW3DnApfWN","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:48:53.187689Z","iopub.execute_input":"2026-04-07T21:48:53.187933Z","iopub.status.idle":"2026-04-07T21:48:53.193979Z","shell.execute_reply.started":"2026-04-07T21:48:53.187912Z","shell.execute_reply":"2026-04-07T21:48:53.193230Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def mixup_data(x, y, alpha=0.4):\n    \"\"\"Mixup augmentation to prevent overfitting.\"\"\"\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1\n    \n    batch_size = x.size(0)\n    index = torch.randperm(batch_size).to(device)\n    \n    mixed_x = lam * x + (1 - lam) * x[index, :]\n    y_a, y_b = y, y[index]\n    return mixed_x, y_a, y_b, lam\n\ndef mixup_criterion(criterion, pred, y_a, y_b, lam):\n    \"\"\"Mixup loss calculation.\"\"\"\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\ndef train_model(model, train_loader, val_loader, total_epochs=20):\n\n    # Loss with class weights and label smoothing\n    criterion = nn.CrossEntropyLoss(weight=class_weights, label_smoothing=0.1)\n    \n    # ⚡ MIXED PRECISION for 20-40% speedup\n    scaler = torch.cuda.amp.GradScaler()\n\n    # ==============================\n    # STAGE 1: Warmup (head only, 1 epoch - FURTHER REDUCED)\n    # ==============================\n    freeze_backbone(model)\n    warmup_params = [p for p in model.parameters() if p.requires_grad]\n\n    warmup_optimizer = torch.optim.Adam(warmup_params, lr=5e-4, weight_decay=5e-4)\n\n    print(\"\\n\" + \"=\"*60)\n    print(\"STAGE 1: Warmup - Training classification head (1 epoch - OPTIMIZED)\")\n    print(\"=\"*60)\n\n    for epoch in range(1):\n        model.train()\n        running_loss = 0\n        loop = tqdm(train_loader, desc=f\"Warmup Epoch {epoch+1}/1\")\n        for images, labels in loop:\n            images, labels = images.to(device), labels.to(device)\n            warmup_optimizer.zero_grad()\n            # ⚡ Mixed precision wrapper\n            with torch.cuda.amp.autocast():\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n            scaler.scale(loss).backward()\n            torch.nn.utils.clip_grad_norm_(warmup_params, max_norm=1.0)\n            scaler.step(warmup_optimizer)\n            scaler.update()\n            running_loss += loss.item()\n            loop.set_postfix(loss=loss.item())\n        print(f\"  Warmup Epoch {epoch+1} avg loss: {running_loss/len(train_loader):.4f}\")\n\n    # ==============================\n    # STAGE 2: Full Fine-tuning\n    # ==============================\n    unfreeze_all(model)\n\n    # Conservative LR for CoAtNet on small medical datasets\n    # INCREASED weight decay to 5e-4 to combat overfitting\n    optimizer = torch.optim.Adam(model.parameters(), lr=3e-5, weight_decay=5e-4)\n\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        optimizer, T_max=total_epochs, eta_min=1e-7\n    )\n\n    train_losses, val_losses = [], []\n    train_accs, val_accs = [], []\n\n    best_val_qwk = 0\n    patience = 5\n    patience_counter = 0\n\n    save_path = os.path.join(SAVE_DIR, 'coatnet_best.pth')\n\n    print(\"\\n\" + \"=\"*60)\n    print(f\"STAGE 2: Full Fine-tuning ({total_epochs} epochs, patience={patience})\")\n    print(f\"OVERFITTING FIXES: weight_decay=5e-4, Mixup augmentation, reduced epochs\")\n    print(f\"⚡ SPEED OPTIMIZATIONS: Mixed Precision (AMP), NUM_WORKERS=4, 1-epoch warmup\")\n    print(\"=\"*60 + \"\\n\")\n\n    for epoch in range(total_epochs):\n\n        # ===== TRAIN =====\n        model.train()\n        running_loss, correct, total = 0, 0, 0\n\n        loop = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{total_epochs}\")\n\n        for images, labels in loop:\n            images, labels = images.to(device), labels.to(device)\n\n            optimizer.zero_grad()\n            \n            # ⚡ Mixed precision training\n            with torch.cuda.amp.autocast():\n                # Apply Mixup 50% of batches to prevent overfitting\n                if np.random.rand() < 0.5:\n                    images, labels_a, labels_b, lam = mixup_data(images, labels, alpha=0.4)\n                    outputs = model(images)\n                    loss = mixup_criterion(criterion, outputs, labels_a, labels_b, lam)\n                else:\n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n\n            scaler.scale(loss).backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            running_loss += loss.item()\n            _, preds_batch = torch.max(outputs, 1)\n            correct += (preds_batch == labels).sum().item()\n            total += labels.size(0)\n\n            loop.set_postfix(loss=loss.item(), acc=correct/total)\n\n        train_loss = running_loss / len(train_loader)\n        train_acc = correct / total\n\n        # ===== VALIDATION =====\n        model.eval()\n        val_loss_sum, correct, total = 0, 0, 0\n        all_preds = []\n        all_labels = []\n\n        with torch.no_grad():\n            for images, labels in val_loader:\n                images, labels = images.to(device), labels.to(device)\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n\n                val_loss_sum += loss.item()\n                _, preds_batch = torch.max(outputs, 1)\n                correct += (preds_batch == labels).sum().item()\n                total += labels.size(0)\n\n                all_preds.extend(preds_batch.cpu().numpy())\n                all_labels.extend(labels.cpu().numpy())\n\n        val_loss = val_loss_sum / len(val_loader)\n        val_acc = correct / total\n        val_qwk = cohen_kappa_score(all_labels, all_preds, weights='quadratic')\n\n        scheduler.step()\n\n        train_losses.append(train_loss)\n        val_losses.append(val_loss)\n        train_accs.append(train_acc)\n        val_accs.append(val_acc)\n\n        # Print epoch results\n        current_lr = optimizer.param_groups[0]['lr']\n        print(f\"\\nEpoch {epoch+1}/{total_epochs} | LR: {current_lr:.2e}\")\n        print(f\"  Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}\")\n        print(f\"  Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val QWK: {val_qwk:.4f}\")\n        print(f\"  Gap (Train-Val Acc): {(train_acc - val_acc):.4f}\")\n        print(\"-\"*60)\n\n        # Save best model based on validation QWK\n        if val_qwk > best_val_qwk:\n            best_val_qwk = val_qwk\n            patience_counter = 0\n            torch.save({\n                'epoch': epoch + 1,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_acc': val_acc,\n                'val_qwk': val_qwk,\n                'val_loss': val_loss,\n                'train_acc': train_acc,\n                'model_name': model_name,\n            }, save_path)\n            file_size = os.path.getsize(save_path) / (1024 * 1024)  # Convert to MB\n            print(f\"  >>> Best model saved! Val QWK: {val_qwk:.4f}, Val Acc: {val_acc:.4f}\")\n            print(f\"      Checkpoint: {save_path} ({file_size:.1f} MB)\")\n        else:\n            patience_counter += 1\n            print(f\"  No improvement. Patience: {patience_counter}/{patience}\")\n\n        # Early stopping\n        if patience_counter >= patience:\n            print(f\"\\nEarly stopping triggered at epoch {epoch+1}!\")\n            break\n\n    print(f\"\\nTraining complete. Best Val QWK: {best_val_qwk:.4f}\")\n    print(f\"Best model saved to: {save_path}\")\n    \n    # Verify file was saved\n    if os.path.exists(save_path):\n        file_size = os.path.getsize(save_path) / (1024 * 1024)\n        print(f\"✓ Confirmed: Model file exists ({file_size:.1f} MB)\")\n        print(f\"✓ Location: /kaggle/working/coatnet_best.pth\")\n        print(f\"✓ This will be available in Kaggle output files!\")\n    else:\n        print(f\"✗ WARNING: Model file not found at {save_path}\")\n\n    return train_losses, val_losses, train_accs, val_accs","metadata":{"id":"50IGg2qypfWN","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:51:24.196637Z","iopub.execute_input":"2026-04-07T21:51:24.196994Z","iopub.status.idle":"2026-04-07T21:51:24.217423Z","shell.execute_reply.started":"2026-04-07T21:51:24.196970Z","shell.execute_reply":"2026-04-07T21:51:24.216534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history = train_model(model, train_loader, val_loader, total_epochs=20)","metadata":{"id":"AwaVf0--pfWN","outputId":"3ca2c19e-eeef-4431-8006-d91b76586449","trusted":true,"execution":{"iopub.status.busy":"2026-04-07T21:51:24.909862Z","iopub.execute_input":"2026-04-07T21:51:24.910184Z","iopub.status.idle":"2026-04-08T01:55:22.868838Z","shell.execute_reply.started":"2026-04-07T21:51:24.910161Z","shell.execute_reply":"2026-04-08T01:55:22.857193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_history(history):\n    train_losses, val_losses, train_accs, val_accs = history\n\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n\n    # Loss plot\n    axes[0].plot(train_losses, label=\"Train Loss\", linewidth=2)\n    axes[0].plot(val_losses, label=\"Validation Loss\", linewidth=2)\n    axes[0].set_xlabel(\"Epochs\")\n    axes[0].set_ylabel(\"Loss\")\n    axes[0].set_title(\"Training vs Validation Loss\")\n    axes[0].legend()\n    axes[0].grid(True, alpha=0.3)\n\n    # Accuracy plot\n    axes[1].plot(train_accs, label=\"Train Accuracy\", linewidth=2)\n    axes[1].plot(val_accs, label=\"Validation Accuracy\", linewidth=2)\n    axes[1].set_xlabel(\"Epochs\")\n    axes[1].set_ylabel(\"Accuracy\")\n    axes[1].set_title(\"Training vs Validation Accuracy\")\n    axes[1].legend()\n    axes[1].grid(True, alpha=0.3)\n\n    # Gap plot (overfitting indicator)\n    gaps = [t - v for t, v in zip(train_accs, val_accs)]\n    axes[2].plot(gaps, label=\"Train-Val Acc Gap\", linewidth=2, color='red')\n    axes[2].axhline(y=0.05, color='green', linestyle='--', label='Acceptable gap (5%)')\n    axes[2].set_xlabel(\"Epochs\")\n    axes[2].set_ylabel(\"Accuracy Gap\")\n    axes[2].set_title(\"Overfitting Gap (Lower is Better)\")\n    axes[2].legend()\n    axes[2].grid(True, alpha=0.3)\n\n    plt.tight_layout()\n    plt.show()\n\nplot_history(history)","metadata":{"id":"0GhC7iVApfWO","outputId":"fcb2d4ef-4812-49c4-accf-c8cf6eb8b436","trusted":true,"execution":{"iopub.status.busy":"2026-04-08T01:55:22.896896Z","iopub.execute_input":"2026-04-08T01:55:22.897543Z","iopub.status.idle":"2026-04-08T01:55:23.801773Z","shell.execute_reply.started":"2026-04-08T01:55:22.897480Z","shell.execute_reply":"2026-04-08T01:55:23.800723Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load best model for evaluation\ncheckpoint = torch.load(os.path.join(SAVE_DIR, 'coatnet_best.pth'), map_location=device, weights_only=False)\nmodel.load_state_dict(checkpoint['model_state_dict'])\nprint(f\"Loaded best model from epoch {checkpoint['epoch']}\")\nprint(f\"Best Val Acc: {checkpoint['val_acc']:.4f}, Best Val QWK: {checkpoint['val_qwk']:.4f}\")\n\nmodel.eval()\n\npreds, labels_all = [], []\n\nwith torch.no_grad():\n    for images, labels in tqdm(test_loader, desc=\"Testing\"):\n        images = images.to(device)\n        outputs = model(images)\n\n        preds.extend(outputs.argmax(1).cpu().numpy())\n        labels_all.extend(labels.numpy())\n\nprint(\"\\n\" + \"=\"*60)\nprint(\"TEST SET RESULTS - CoAtNet\")\nprint(\"=\"*60)\nprint(classification_report(labels_all, preds,\n      target_names=['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative']))\nprint(f\"Quadratic Weighted Kappa: {cohen_kappa_score(labels_all, preds, weights='quadratic'):.4f}\")\nprint(f\"Test Accuracy: {np.mean(np.array(preds) == np.array(labels_all)):.4f}\")","metadata":{"id":"BhThaJFnpfWO","outputId":"bce4ebb0-3e66-492a-b8a8-71c9559b7023","trusted":true,"execution":{"iopub.status.busy":"2026-04-08T01:55:23.803141Z","iopub.execute_input":"2026-04-08T01:55:23.803478Z","iopub.status.idle":"2026-04-08T01:56:09.803967Z","shell.execute_reply.started":"2026-04-08T01:55:23.803443Z","shell.execute_reply":"2026-04-08T01:56:09.802935Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cm = confusion_matrix(labels_all, preds)\n\nplt.figure(figsize=(8, 6))\nsns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\",\n            xticklabels=['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative'],\n            yticklabels=['No DR', 'Mild', 'Moderate', 'Severe', 'Proliferative'])\nplt.xlabel(\"Predicted\")\nplt.ylabel(\"Actual\")\nplt.title(\"Confusion Matrix - CoAtNet\")\nplt.tight_layout()\nplt.show()","metadata":{"id":"wRliUBEApfWO","trusted":true,"execution":{"iopub.status.busy":"2026-04-08T01:56:09.806399Z","iopub.execute_input":"2026-04-08T01:56:09.806643Z","iopub.status.idle":"2026-04-08T01:56:10.041469Z","shell.execute_reply.started":"2026-04-08T01:56:09.806618Z","shell.execute_reply":"2026-04-08T01:56:10.040392Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from IPython.display import FileLink\n\nFileLink('/kaggle/working/coatnet_best.pth')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-08T02:45:01.411092Z","iopub.execute_input":"2026-04-08T02:45:01.411621Z","iopub.status.idle":"2026-04-08T02:45:01.418161Z","shell.execute_reply.started":"2026-04-08T02:45:01.411564Z","shell.execute_reply":"2026-04-08T02:45:01.417276Z"}},"outputs":[],"execution_count":null}]}