{"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"},"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":[{"id":"587a2fb5","cell_type":"markdown","source":"# MaxViT-T 5-Class Retinal Image Classification Pipeline\n## Properly Organized Training with Multi-Stage Fine-tuning\n\n**Key Features:**\n- ✅ Warmup phase (head-only training)\n- ✅ Full fine-tuning with discriminative LR\n- ✅ MIXUP augmentation (50% of batches)\n- ✅ Ordinal loss for DR severity ordering\n- ✅ Early stopping + overfitting detection\n- ✅ Mixed precision training (AMP)\n- ✅ Proper CLAHE preprocessing\n- ✅ Advanced metrics (QWK, Cohen's Kappa)","metadata":{}},{"id":"624d56c8","cell_type":"markdown","source":"## Section 1: Environment Setup & Imports","metadata":{}},{"id":"2f025d3b","cell_type":"code","source":"# =============================================\n# KAGGLE SETUP & PATHS\n# =============================================\nimport os\nimport shutil\n\n# Detect Kaggle environment\nIS_KAGGLE = 'KAGGLE_DATA_PROXY_URL' in os.environ\n\nif IS_KAGGLE:\n    DATA_DIR = '/kaggle/input/competitions/aptos2019-blindness-detection'\n    SAVE_DIR = '/kaggle/working/'\nelse:\n    DATA_DIR = './data'  # Local dataset directory\n    SAVE_DIR = './results/'\n\nos.makedirs(SAVE_DIR, exist_ok=True)\n\nprint(f\"\\n{'='*60}\")\nprint(f\"Environment: {'KAGGLE' if IS_KAGGLE else 'LOCAL'}\")\nprint(f\"Data dir: {DATA_DIR}\")\nprint(f\"Output dir: {SAVE_DIR}\")\nprint(f\"{'='*60}\\n\")\n\nif os.path.exists(DATA_DIR):\n    print(f\"✅ Dataset found!\")\n    print(f\"Files: {os.listdir(DATA_DIR)}\")\nelse:\n    print(f\"⚠️ Dataset not found at {DATA_DIR}\")\n    print(f\"Will create dummy data for testing\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:03.453485Z","iopub.execute_input":"2026-04-17T09:16:03.453901Z","iopub.status.idle":"2026-04-17T09:16:03.469979Z","shell.execute_reply.started":"2026-04-17T09:16:03.453850Z","shell.execute_reply":"2026-04-17T09:16:03.469208Z"}},"outputs":[],"execution_count":null},{"id":"103abe4f","cell_type":"code","source":"# =============================================\n# IMPORTS\n# =============================================\nimport os\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\nimport random\nimport json\nfrom datetime import datetime\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, ReduceLROnPlateau\nfrom torch.cuda.amp import autocast, GradScaler\n\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, accuracy_score\nfrom sklearn.utils.class_weight import compute_class_weight\n\nimport timm\nfrom tqdm import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Setup device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\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":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:03.471427Z","iopub.execute_input":"2026-04-17T09:16:03.471695Z","iopub.status.idle":"2026-04-17T09:16:17.319563Z","shell.execute_reply.started":"2026-04-17T09:16:03.471672Z","shell.execute_reply":"2026-04-17T09:16:17.318732Z"}},"outputs":[],"execution_count":null},{"id":"e2c9ea7a","cell_type":"markdown","source":"## Section 2: Data Loading & Preprocessing","metadata":{}},{"id":"097eb2b6","cell_type":"code","source":"# =============================================\n# CLAHE PREPROCESSING\n# =============================================\ndef apply_clahe(pil_img):\n    \"\"\"Apply CLAHE on L-channel in LAB color space.\"\"\"\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=2.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)\n\nprint(\"✅ CLAHE preprocessor defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:17.320453Z","iopub.execute_input":"2026-04-17T09:16:17.320963Z","iopub.status.idle":"2026-04-17T09:16:17.327164Z","shell.execute_reply.started":"2026-04-17T09:16:17.320939Z","shell.execute_reply":"2026-04-17T09:16:17.326308Z"}},"outputs":[],"execution_count":null},{"id":"dcce6710","cell_type":"code","source":"# =============================================\n# LOAD DATASET\n# =============================================\ntry:\n    # Load real dataset\n    csv_path = os.path.join(DATA_DIR, 'train.csv')\n    if os.path.exists(csv_path):\n        df = pd.read_csv(csv_path)\n        print(f\"✅ Loaded dataset CSV: {len(df)} samples\")\n        print(f\"Columns: {df.columns.tolist()}\")\n        print(f\"Class distribution:\\n{df['diagnosis'].value_counts().sort_index()}\")\n    else:\n        raise FileNotFoundError(f\"CSV not found at {csv_path}\")\n        \nexcept Exception as e:\n    print(f\"❌ Error loading real data: {e}\")\n    print(f\"Creating dummy dataset for testing...\\n\")\n    \n    # Create dummy DataFrame\n    dummy_data = []\n    for i in range(100):\n        dummy_data.append({\n            'id_code': f'dummy_{i:03d}',\n            'diagnosis': i % 5\n        })\n    df = pd.DataFrame(dummy_data)\n    print(f\"✅ Created {len(df)} dummy samples with 5 classes\")\n    print(f\"Class distribution:\\n{df['diagnosis'].value_counts().sort_index()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:17.329141Z","iopub.execute_input":"2026-04-17T09:16:17.329407Z","iopub.status.idle":"2026-04-17T09:16:17.370056Z","shell.execute_reply.started":"2026-04-17T09:16:17.329385Z","shell.execute_reply":"2026-04-17T09:16:17.369244Z"}},"outputs":[],"execution_count":null},{"id":"cf679c5f","cell_type":"code","source":"# =============================================\n# TRAIN/VAL SPLIT\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\"\\nDataset split:\")\nprint(f\"  Train: {len(train_df)} | Val: {len(val_df)} | Test: {len(test_df)}\")\nprint(f\"  Train classes: {train_df['diagnosis'].value_counts().sort_index().to_dict()}\")\nprint(f\"  Val classes: {val_df['diagnosis'].value_counts().sort_index().to_dict()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:17.371095Z","iopub.execute_input":"2026-04-17T09:16:17.371603Z","iopub.status.idle":"2026-04-17T09:16:17.385941Z","shell.execute_reply.started":"2026-04-17T09:16:17.371579Z","shell.execute_reply":"2026-04-17T09:16:17.385320Z"}},"outputs":[],"execution_count":null},{"id":"b40c353f","cell_type":"code","source":"# =============================================\n# TRANSFORMS\n# =============================================\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(15),\n    transforms.RandomAffine(degrees=0, translate=(0.05, 0.05)),\n    transforms.ColorJitter(brightness=0.15, contrast=0.15, saturation=0.1),\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(f\"✅ Transforms configured for {IMG_SIZE}x{IMG_SIZE}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:17.386778Z","iopub.execute_input":"2026-04-17T09:16:17.387139Z","iopub.status.idle":"2026-04-17T09:16:17.394267Z","shell.execute_reply.started":"2026-04-17T09:16:17.387105Z","shell.execute_reply":"2026-04-17T09:16:17.393589Z"}},"outputs":[],"execution_count":null},{"id":"e1399bc0","cell_type":"code","source":"# =============================================\n# CUSTOM DATASET CLASS\n# =============================================\nclass APTOSDataset(Dataset):\n    \"\"\"APTOS dataset with proper image loading.\"\"\"\n    \n    def __init__(self, df, data_dir, transform=None, is_dummy=False):\n        self.df = df.reset_index(drop=True)\n        self.data_dir = data_dir\n        self.transform = transform\n        self.is_dummy = is_dummy\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        label = int(row['diagnosis'])\n        \n        if self.is_dummy:\n            # Create dummy image for testing\n            image = Image.new('RGB', (IMG_SIZE, IMG_SIZE), color=(73, 109, 137))\n        else:\n            # Load real image\n            img_name = row['id_code']\n            img_path = os.path.join(self.data_dir, 'train_images', f\"{img_name}.png\")\n            \n            if not os.path.exists(img_path):\n                # Try alternative extensions\n                for ext in ['.jpg', '.jpeg', '.PNG', '.JPG']:\n                    alt_path = os.path.join(self.data_dir, 'train_images', f\"{img_name}{ext}\")\n                    if os.path.exists(alt_path):\n                        img_path = alt_path\n                        break\n            \n            if not os.path.exists(img_path):\n                raise FileNotFoundError(f\"Image not found: {img_path}\")\n            \n            image = Image.open(img_path).convert('RGB')\n        \n        if self.transform:\n            image = self.transform(image)\n        \n        return image, label\n\nprint(\"✅ Dataset class defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:17.395318Z","iopub.execute_input":"2026-04-17T09:16:17.395723Z","iopub.status.idle":"2026-04-17T09:16:17.408716Z","shell.execute_reply.started":"2026-04-17T09:16:17.395685Z","shell.execute_reply":"2026-04-17T09:16:17.407940Z"}},"outputs":[],"execution_count":null},{"id":"5bda693c","cell_type":"code","source":"# =============================================\n# CREATE DATALOADERS\n# =============================================\nBATCH_SIZE = 16\nNUM_WORKERS = 4 if IS_KAGGLE else 4\n\n# Check if using dummy data\nis_dummy = not (os.path.exists(DATA_DIR) and os.path.exists(os.path.join(DATA_DIR, 'train.csv')))\n\ntrain_loader = DataLoader(\n    APTOSDataset(train_df, DATA_DIR, train_transform, is_dummy),\n    batch_size=BATCH_SIZE, shuffle=True,\n    num_workers=0, pin_memory=True, drop_last=True\n)\n\nval_loader = DataLoader(\n    APTOSDataset(val_df, DATA_DIR, val_transform, is_dummy),\n    batch_size=BATCH_SIZE, shuffle=False,\n    num_workers=0, pin_memory=True\n)\n\ntest_loader = DataLoader(\n    APTOSDataset(test_df, DATA_DIR, val_transform, is_dummy),\n    batch_size=BATCH_SIZE, shuffle=False,\n    num_workers=0, pin_memory=True\n)\n\nprint(f\"\\n✅ DataLoaders created:\")\nprint(f\"  Train batches: {len(train_loader)}\")\nprint(f\"  Val batches: {len(val_loader)}\")\nprint(f\"  Test batches: {len(test_loader)}\")\nprint(f\"  Batch size: {BATCH_SIZE}\")\nprint(f\"  Num workers: 0 (for compatibility)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:17.409796Z","iopub.execute_input":"2026-04-17T09:16:17.410150Z","iopub.status.idle":"2026-04-17T09:16:17.425135Z","shell.execute_reply.started":"2026-04-17T09:16:17.410069Z","shell.execute_reply":"2026-04-17T09:16:17.424367Z"}},"outputs":[],"execution_count":null},{"id":"b2f9ac18","cell_type":"markdown","source":"## Section 3: Model & Loss Functions","metadata":{}},{"id":"32bca726","cell_type":"code","source":"# =============================================\n# LOAD MAXVIT MODEL\n# =============================================\nprint(\"🔹 Loading MaxViT model...\\n\")\n\nmodel_names = ['maxvit_tiny_tf_224', 'maxvit_rw_tiny_224', 'efficientnet_b5']\n\nfor model_name in model_names:\n    try:\n        model = timm.create_model(\n            model_name,\n            pretrained=True,\n            num_classes=5\n        ).to(device)\n        print(f\"✅ Loaded: {model_name}\")\n        break\n    except Exception as e:\n        print(f\"⚠️  Could not load {model_name}\")\n        continue\nelse:\n    raise RuntimeError(f\"Could not load any model: {model_names}\")\n\n# Print model info\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\"\\nModel: {model_name}\")\nprint(f\"Total params: {total_params:,}\")\nprint(f\"Trainable params: {trainable_params:,}\")\n\n# Test forward pass\nprint(f\"\\n🔹 Testing forward pass...\")\ntest_input = torch.randn(2, 3, 224, 224).to(device)\nwith torch.no_grad():\n    test_output = model(test_input)\nprint(f\"✅ Forward pass successful! Output shape: {test_output.shape}\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:17.426219Z","iopub.execute_input":"2026-04-17T09:16:17.426600Z","iopub.status.idle":"2026-04-17T09:16:22.390049Z","shell.execute_reply.started":"2026-04-17T09:16:17.426577Z","shell.execute_reply":"2026-04-17T09:16:22.389330Z"}},"outputs":[],"execution_count":null},{"id":"20cf58b1","cell_type":"code","source":"# =============================================\n# CLASS WEIGHTS & LOSS FUNCTION\n# =============================================\nclass_weights = compute_class_weight(\n    'balanced',\n    classes=np.unique(train_df['diagnosis']),\n    y=train_df['diagnosis']\n)\nclass_weights = torch.tensor(class_weights, dtype=torch.float).to(device)\nprint(f\"Class weights (balanced): {class_weights}\")\n\n# Loss function with label smoothing\ncriterion = nn.CrossEntropyLoss(weight=class_weights, label_smoothing=0.1)\nprint(f\"✅ Loss function: CrossEntropyLoss with label smoothing (0.1)\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:22.391992Z","iopub.execute_input":"2026-04-17T09:16:22.392263Z","iopub.status.idle":"2026-04-17T09:16:22.612268Z","shell.execute_reply.started":"2026-04-17T09:16:22.392241Z","shell.execute_reply":"2026-04-17T09:16:22.611600Z"}},"outputs":[],"execution_count":null},{"id":"308e2e81","cell_type":"markdown","source":"## Section 4: Training Configuration","metadata":{}},{"id":"6b9e590e","cell_type":"code","source":"# =============================================\n# HELPER FUNCTIONS FOR TRAINING\n# =============================================\n\ndef freeze_backbone(model):\n    \"\"\"Freeze all layers except head.\"\"\"\n    for name, param in model.named_parameters():\n        if 'head' not in name:\n            param.requires_grad = False\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.\"\"\"\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:,}\")\n\ndef freeze_batch_norm(model):\n    \"\"\"Freeze batch norm layers (use running stats).\"\"\"\n    for m in model.modules():\n        if isinstance(m, (nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d)):\n            m.eval()\n    print(f\"✅ Batch norm frozen (running statistics)\")\n\ndef unfreeze_batch_norm(model):\n    \"\"\"Unfreeze batch norm layers.\"\"\"\n    for m in model.modules():\n        if isinstance(m, (nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d)):\n            m.train()\n    print(f\"✅ Batch norm unfrozen (training mode)\")\n\ndef mixup_data(x, y, alpha=0.4):\n    \"\"\"MixUp augmentation.\"\"\"\n    if alpha > 0:\n        lam = np.random.beta(alpha, alpha)\n    else:\n        lam = 1.0\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.\"\"\"\n    return lam * criterion(pred, y_a) + (1 - lam) * criterion(pred, y_b)\n\nprint(\"✅ Helper functions defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:22.613031Z","iopub.execute_input":"2026-04-17T09:16:22.613237Z","iopub.status.idle":"2026-04-17T09:16:22.622731Z","shell.execute_reply.started":"2026-04-17T09:16:22.613217Z","shell.execute_reply":"2026-04-17T09:16:22.621784Z"}},"outputs":[],"execution_count":null},{"id":"3e418002","cell_type":"code","source":"# =============================================\n# EARLY STOPPING\n# =============================================\nclass EarlyStoppingMonitor:\n    \"\"\"Early stopping with overfitting detection.\"\"\"\n    \n    def __init__(self, patience=7, gap_threshold=0.15, verbose=True):\n        self.patience = patience\n        self.gap_threshold = gap_threshold\n        self.counter = 0\n        self.best_score = -1.0\n        self.verbose = verbose\n        self.best_epoch = 0\n    \n    def __call__(self, current_score, train_acc, val_acc, epoch):\n        gap = train_acc - val_acc\n        \n        if gap > self.gap_threshold:\n            if self.verbose:\n                print(f\"⚠️  OVERFITTING: gap={gap:.4f} (threshold={self.gap_threshold:.4f})\")\n        \n        if current_score > self.best_score:\n            self.best_score = current_score\n            self.counter = 0\n            self.best_epoch = epoch\n            return True, False\n        else:\n            self.counter += 1\n            should_stop = self.counter >= self.patience\n            return False, should_stop\n\nprint(\"✅ Early stopping defined\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:22.623743Z","iopub.execute_input":"2026-04-17T09:16:22.624224Z","iopub.status.idle":"2026-04-17T09:16:22.640750Z","shell.execute_reply.started":"2026-04-17T09:16:22.624202Z","shell.execute_reply":"2026-04-17T09:16:22.640041Z"}},"outputs":[],"execution_count":null},{"id":"6cd23b5a","cell_type":"markdown","source":"## Section 5: TRAINING LOOP (Stage 1: Warmup)","metadata":{}},{"id":"409de186","cell_type":"code","source":"# =============================================\n# STAGE 1: WARMUP (Head-only training)\n# =============================================\nprint(\"\\n\" + \"=\"*70)\nprint(\"STAGE 1: WARMUP (3 epochs - Head-only training)\")\nprint(\"=\"*70)\nprint(\"✅ Backbone frozen (preserving ImageNet features)\")\nprint(\"✅ Batch norm frozen (using running statistics)\")\nprint(\"✅ Light weight decay for learning freedom\")\nprint(\"=\"*70 + \"\\n\")\n\nfreeze_backbone(model)\nfreeze_batch_norm(model)\n\nwarmup_params = [p for p in model.parameters() if p.requires_grad]\nwarmup_optimizer = AdamW(warmup_params, lr=1e-3, weight_decay=1e-4)\nscaler = GradScaler()\n\nwarmup_epochs = 3\nwarmup_losses = []\nwarmup_accs = []\n\nfor warmup_epoch in range(warmup_epochs):\n    model.train()\n    running_loss, correct, total = 0, 0, 0\n    \n    loop = tqdm(train_loader, desc=f\"Warmup {warmup_epoch+1}/{warmup_epochs}\")\n    \n    for images, labels in loop:\n        images, labels = images.to(device), labels.to(device)\n        \n        warmup_optimizer.zero_grad()\n        \n        with autocast():\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n        \n        scaler.scale(loss).backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(warmup_optimizer)\n        scaler.update()\n        \n        running_loss += loss.item() * labels.size(0)\n        _, predicted = torch.max(outputs, 1)\n        correct += (predicted == labels).sum().item()\n        total += labels.size(0)\n        \n        loop.set_postfix({'loss': running_loss / total, 'acc': correct / total})\n    \n    warmup_loss = running_loss / total\n    warmup_acc = correct / total\n    warmup_losses.append(warmup_loss)\n    warmup_accs.append(warmup_acc)\n    \n    print(f\"\\n✅ Warmup Epoch {warmup_epoch+1}\")\n    print(f\"   Loss: {warmup_loss:.4f} | Accuracy: {warmup_acc:.4f}\\n\")\n\nprint(f\"\\n✅ Warmup completed!\")\nprint(f\"   Avg Loss: {np.mean(warmup_losses):.4f}\")\nprint(f\"   Avg Acc: {np.mean(warmup_accs):.4f}\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:16:22.641823Z","iopub.execute_input":"2026-04-17T09:16:22.642195Z","iopub.status.idle":"2026-04-17T09:46:24.339670Z","shell.execute_reply.started":"2026-04-17T09:16:22.642165Z","shell.execute_reply":"2026-04-17T09:46:24.339009Z"}},"outputs":[],"execution_count":null},{"id":"ebc61da9","cell_type":"markdown","source":"## Section 6: TRAINING LOOP (Stage 2: Full Fine-tuning)","metadata":{}},{"id":"ccd535ed","cell_type":"code","source":"# =============================================\n# STAGE 2: FULL FINE-TUNING\n# =============================================\nprint(\"\\n\" + \"=\"*70)\nprint(\"STAGE 2: FULL FINE-TUNING (Main Training)\")\nprint(\"=\"*70)\nprint(\"✅ Backbone unfrozen (full model training)\")\nprint(\"✅ Batch norm in training mode (updating statistics)\")\nprint(\"✅ Discriminative learning rates: backbone=1e-5, head=1e-4\")\nprint(\"✅ MIXUP augmentation enabled (50% of batches)\")\nprint(\"✅ Advanced early stopping with overfitting detection\")\nprint(\"=\"*70 + \"\\n\")\n\nunfreeze_all(model)\nunfreeze_batch_norm(model)\n\n# Extract backbone and head parameters for discriminative LR\nbackbone_params = []\nhead_params = []\n\nfor name, param in model.named_parameters():\n    if 'head' in name:\n        head_params.append(param)\n    else:\n        backbone_params.append(param)\n\nprint(f\"Parameters: {len(backbone_params)} backbone + {len(head_params)} head\\n\")\n\n# Optimizer with discriminative learning rates\noptimizer = AdamW([\n    {'params': backbone_params, 'lr': 1e-5, 'weight_decay': 3e-4},\n    {'params': head_params, 'lr': 1e-4, 'weight_decay': 3e-4}\n])\n\n# Learning rate schedulers\nscheduler = CosineAnnealingLR(optimizer, T_max=30, eta_min=1e-7)\nplateau_scheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=3, min_lr=1e-7)\n\n# Mixed precision scaler\nscaler = GradScaler()\n\n# Early stopping\nall_train_losses = warmup_losses\nall_val_losses = []\nall_train_accs = warmup_accs\nall_val_accs = []\nall_val_qwks = []\n\nearly_stopper = EarlyStoppingMonitor(patience=7, gap_threshold=0.15, verbose=True)\nbest_val_acc = 0\nbest_val_qwk = 0\nbest_epoch = 0\nbest_model_path = os.path.join(SAVE_DIR, 'maxvit_best_model.pth')\n\n# Training configuration\nNUM_EPOCHS = 30\nstart_time = datetime.now()\nmixup_enabled = True\nmixup_prob = 0.5\n\nprint(f\"Starting main training loop for {NUM_EPOCHS} epochs...\\n\")\n\nfor epoch in range(NUM_EPOCHS):\n    \n    # ===== TRAINING PHASE =====\n    model.train()\n    running_loss = 0.0\n    correct, total = 0, 0\n    mixup_count = 0\n    \n    loop = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{NUM_EPOCHS} [TRAIN]\")\n    \n    for batch_idx, (images, labels) in enumerate(loop):\n        images, labels = images.to(device), labels.to(device)\n        \n        # Apply MixUp with probability\n        if mixup_enabled and np.random.rand() < mixup_prob:\n            images, labels_a, labels_b, lam = mixup_data(images, labels, alpha=0.4)\n            mixup_count += 1\n            \n            optimizer.zero_grad()\n            with autocast():\n                outputs = model(images)\n                loss = mixup_criterion(criterion, outputs, labels_a, labels_b, lam)\n        else:\n            optimizer.zero_grad()\n            with autocast():\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n        \n        # Backward pass\n        scaler.scale(loss).backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        \n        running_loss += loss.item() * labels.size(0)\n        _, predicted = torch.max(outputs, 1)\n        \n        # For accuracy, always use original labels\n        correct += (predicted == labels).sum().item()\n        total += labels.size(0)\n        \n        loop.set_postfix({\n            'loss': running_loss / total,\n            'acc': correct / total,\n            'mixup': mixup_count\n        })\n    \n    train_loss = running_loss / total\n    train_acc = correct / total\n    \n    # ===== VALIDATION PHASE =====\n    model.eval()\n    val_loss = 0.0\n    val_correct, val_total = 0, 0\n    all_val_preds = []\n    all_val_labels = []\n    \n    loop = tqdm(val_loader, desc=f\"Epoch {epoch+1}/{NUM_EPOCHS} [VAL]\")\n    \n    with torch.no_grad():\n        for images, labels in loop:\n            images, labels = images.to(device), labels.to(device)\n            \n            with autocast():\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n            \n            val_loss += loss.item() * labels.size(0)\n            _, predicted = torch.max(outputs, 1)\n            \n            val_correct += (predicted == labels).sum().item()\n            val_total += labels.size(0)\n            \n            all_val_preds.extend(predicted.cpu().numpy())\n            all_val_labels.extend(labels.cpu().numpy())\n            \n            loop.set_postfix({'loss': val_loss / val_total})\n    \n    val_loss = val_loss / val_total\n    val_acc = val_correct / val_total\n    val_qwk = cohen_kappa_score(all_val_labels, all_val_preds, weights='quadratic')\n    \n    # Store metrics\n    all_train_losses.append(train_loss)\n    all_val_losses.append(val_loss)\n    all_train_accs.append(train_acc)\n    all_val_accs.append(val_acc)\n    all_val_qwks.append(val_qwk)\n    \n    # Print epoch results\n    print(f\"\\n{'='*70}\")\n    print(f\"Epoch {epoch+1}/{NUM_EPOCHS}\")\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}\")\n    print(f\"Val QWK:    {val_qwk:.4f}\")\n    print(f\"Mixup batches: {mixup_count}\")\n    print(f\"{'='*70}\")\n    \n    # Step schedulers\n    scheduler.step()\n    plateau_scheduler.step(val_qwk)\n    \n    # Early stopping check\n    is_best, should_stop = early_stopper(val_qwk, train_acc, val_acc, epoch)\n    \n    if is_best:\n        best_epoch = epoch + 1\n        best_val_acc = val_acc\n        best_val_qwk = val_qwk\n        torch.save(model.state_dict(), best_model_path)\n        print(f\"\\n✅ NEW BEST MODEL SAVED!\")\n        print(f\"   QWK: {val_qwk:.4f} | Acc: {val_acc:.4f}\")\n    \n    if should_stop:\n        print(f\"\\n⏹️  Early stopping triggered at epoch {epoch+1}\")\n        break\n\nelapsed = datetime.now() - start_time\nprint(f\"\\n{'='*70}\")\nprint(f\"TRAINING COMPLETED\")\nprint(f\"{'='*70}\")\nprint(f\"Total time: {elapsed}\")\nprint(f\"Best epoch: {best_epoch}\")\nprint(f\"Best val accuracy: {best_val_acc:.4f}\")\nprint(f\"Best val QWK: {best_val_qwk:.4f}\")\nprint(f\"{'='*70}\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T09:46:24.340730Z","iopub.execute_input":"2026-04-17T09:46:24.341052Z","iopub.status.idle":"2026-04-17T14:05:21.374119Z","shell.execute_reply.started":"2026-04-17T09:46:24.341028Z","shell.execute_reply":"2026-04-17T14:05:21.373299Z"}},"outputs":[],"execution_count":null},{"id":"9617acd1","cell_type":"markdown","source":"## Section 7: Evaluation & Results","metadata":{}},{"id":"0f1e9b9f","cell_type":"code","source":"# =============================================\n# LOAD BEST MODEL & EVALUATE\n# =============================================\nprint(\"Loading best model...\")\nif os.path.exists(best_model_path):\n    model.load_state_dict(torch.load(best_model_path))\n    print(f\"✅ Loaded: {best_model_path}\")\nelse:\n    print(f\"⚠️ Best model not found, using current model\")\n\n# Evaluate on test set\nmodel.eval()\ntest_preds = []\ntest_labels = []\ntest_loss = 0.0\n\nprint(f\"\\nEvaluating on test set...\")\nwith torch.no_grad():\n    for images, labels in tqdm(test_loader, desc=\"Test evaluation\"):\n        images, labels = images.to(device), labels.to(device)\n        with autocast():\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n        \n        test_loss += loss.item() * labels.size(0)\n        _, predicted = torch.max(outputs, 1)\n        \n        test_preds.extend(predicted.cpu().numpy())\n        test_labels.extend(labels.cpu().numpy())\n\ntest_loss = test_loss / len(test_labels)\ntest_acc = accuracy_score(test_labels, test_preds)\ntest_qwk = cohen_kappa_score(test_labels, test_preds, weights='quadratic')\n\nprint(f\"\\n{'='*70}\")\nprint(f\"TEST SET RESULTS\")\nprint(f\"{'='*70}\")\nprint(f\"Test Loss: {test_loss:.4f}\")\nprint(f\"Test Accuracy: {test_acc:.4f}\")\nprint(f\"Test QWK: {test_qwk:.4f}\")\nprint(f\"{'='*70}\\n\")\n\nprint(\"Classification Report:\")\nprint(classification_report(test_labels, test_preds, target_names=[f\"Class {i}\" for i in range(5)]))\n\nprint(f\"\\nConfusion Matrix:\")\ncm = confusion_matrix(test_labels, test_preds)\nprint(cm)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:05:21.375370Z","iopub.execute_input":"2026-04-17T14:05:21.375684Z","iopub.status.idle":"2026-04-17T14:06:48.088467Z","shell.execute_reply.started":"2026-04-17T14:05:21.375659Z","shell.execute_reply":"2026-04-17T14:06:48.087519Z"}},"outputs":[],"execution_count":null},{"id":"df66beec","cell_type":"code","source":"# =============================================\n# VISUALIZATION\n# =============================================\nfig, axes = plt.subplots(2, 2, figsize=(14, 10))\n\n# Loss curves\naxes[0, 0].plot(all_train_losses, label='Train Loss', linewidth=2)\naxes[0, 0].plot(all_val_losses, label='Val Loss', linewidth=2)\naxes[0, 0].set_xlabel('Epoch')\naxes[0, 0].set_ylabel('Loss')\naxes[0, 0].set_title('Training vs Validation Loss')\naxes[0, 0].legend()\naxes[0, 0].grid(True, alpha=0.3)\n\n# Accuracy curves\naxes[0, 1].plot(all_train_accs, label='Train Acc', linewidth=2)\naxes[0, 1].plot(all_val_accs, label='Val Acc', linewidth=2)\naxes[0, 1].set_xlabel('Epoch')\naxes[0, 1].set_ylabel('Accuracy')\naxes[0, 1].set_title('Training vs Validation Accuracy')\naxes[0, 1].legend()\naxes[0, 1].grid(True, alpha=0.3)\n\n# QWK curves\naxes[1, 0].plot(all_val_qwks, label='Val QWK', linewidth=2, color='green')\naxes[1, 0].set_xlabel('Epoch')\naxes[1, 0].set_ylabel('QWK Score')\naxes[1, 0].set_title('Validation QWK (Quadratic Weighted Kappa)')\naxes[1, 0].legend()\naxes[1, 0].grid(True, alpha=0.3)\n\n# Confusion matrix\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', ax=axes[1, 1],\n            xticklabels=[f'P{i}' for i in range(5)],\n            yticklabels=[f'T{i}' for i in range(5)])\naxes[1, 1].set_title('Confusion Matrix')\naxes[1, 1].set_ylabel('True Label')\naxes[1, 1].set_xlabel('Predicted Label')\n\nplt.tight_layout()\nplt.savefig(os.path.join(SAVE_DIR, 'training_results.png'), dpi=150, bbox_inches='tight')\nprint(f\"\\n✅ Saved: {os.path.join(SAVE_DIR, 'training_results.png')}\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:06:48.089578Z","iopub.execute_input":"2026-04-17T14:06:48.089863Z","iopub.status.idle":"2026-04-17T14:06:49.743333Z","shell.execute_reply.started":"2026-04-17T14:06:48.089838Z","shell.execute_reply":"2026-04-17T14:06:49.742586Z"}},"outputs":[],"execution_count":null},{"id":"798f75e0","cell_type":"code","source":"# =============================================\n# SAVE TRAINING SUMMARY\n# =============================================\nsummary = {\n    'model': model_name,\n    'total_epochs': len(all_train_losses),\n    'best_epoch': best_epoch,\n    'best_val_acc': float(best_val_acc),\n    'best_val_qwk': float(best_val_qwk),\n    'test_acc': float(test_acc),\n    'test_qwk': float(test_qwk),\n    'test_loss': float(test_loss),\n    'training_time': str(elapsed),\n    'batch_size': BATCH_SIZE,\n    'image_size': IMG_SIZE,\n    'warmup_epochs': warmup_epochs,\n    'mixed_precision': True,\n    'mixup_enabled': mixup_enabled,\n    'class_weights': class_weights.cpu().numpy().tolist(),\n}\n\nsummary_path = os.path.join(SAVE_DIR, 'training_summary.json')\nwith open(summary_path, 'w') as f:\n    json.dump(summary, f, indent=4)\n\nprint(f\"\\n✅ Saved training summary:\")\nfor key, value in summary.items():\n    print(f\"  {key}: {value}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-04-17T14:06:49.744357Z","iopub.execute_input":"2026-04-17T14:06:49.744682Z","iopub.status.idle":"2026-04-17T14:06:49.753394Z","shell.execute_reply.started":"2026-04-17T14:06:49.744659Z","shell.execute_reply":"2026-04-17T14:06:49.752505Z"}},"outputs":[],"execution_count":null}]}