{"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":13451,"datasetId":654585,"databundleVersionId":1188070},{"sourceType":"kernelVersion","sourceId":298700524},{"sourceType":"kernelVersion","sourceId":298967632}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"914521b7","cell_type":"markdown","source":"# Notebook 3 — Model Training (Session Chaining)\n**RSNA Intracranial Hemorrhage Detection — EfficientNet-B0**\n\nThis notebook handles training across multiple Kaggle sessions.\nEach session trains a fixed number of epochs, saves a checkpoint,\nand the next session resumes from where the previous one left off.\n\n### Session chaining workflow\n| Session | `PREV_CHECKPOINT_DIR` | Trains epochs |\n|---------|----------------------|---------------|\n| 1 (first run) | `None` (starts fresh) | 0 → N-1 |\n| 2 | `/kaggle/input/<nb3-session1-name>` | N → 2N-1 |\n| 3 | `/kaggle/input/<nb3-session2-name>` | 2N → 3N-1 |\n\n> **Before each session**: update `PREV_CHECKPOINT_DIR` to point to the\n> previous session's notebook output name.\n\n### Key improvements\n- **Patient-level split**: Train/val split is done at the patient level (NB02).\n  No slices from the same patient appear in both sets.\n- **Dataset normalization**: Uses dataset-specific mean/std from NB02 (with ImageNet fallback).\n- **NaN divergence guard**: Training aborts early if loss becomes NaN/Inf.\n- **Backbone freezing**: Classifier head trains alone for the first N epochs,\n  then backbone unfreezes with a lower learning rate (discriminative LR).\n- **Early stopping**: Halts training when val_loss hasn't improved for `PATIENCE` epochs.\n- **Discriminative LR**: Backbone uses `BASE_LR × BACKBONE_LR_FACTOR` to prevent\n\n  catastrophic forgetting of ImageNet features.- Previous session checkpoint (all sessions except the first)\n\n- Preprocessing cache (Notebook 02 output) — contains `cache/` NPY arrays, `manifest.csv`, `normalization_stats.json`\n### Required input datasets","metadata":{}},{"id":"096be66d","cell_type":"code","source":"from pathlib import Path\nimport os\nimport torch\nimport json\n\n# ═══════════════════════════════════════════════════════════════════════════\n# ██  CONFIG — edit these values at the start of each session  ██\n# ═══════════════════════════════════════════════════════════════════════════\n\nPREV_CHECKPOINT_DIR = \"/kaggle/input/notebooks/harshitghosh/03nbeda/\"\nN_EPOCHS_THIS_SESSION = 5\nTOTAL_EPOCHS = 20\n\nBATCH_SIZE   = 16\nBASE_LR      = 1e-4\nWEIGHT_DECAY = 1e-5\nNUM_WORKERS  = 4\nIMG_SIZE     = 256\nSEED         = 42\n\nARCH = 'efficientnet_b0'  # or 'resnet50'\n\n# ─── Anti-overfitting controls ───────────────────────────────────────────\nFREEZE_BACKBONE_EPOCHS = 2    # Train only classifier head for first N epochs\nBACKBONE_LR_FACTOR     = 0.1  # After unfreezing: backbone LR = BASE_LR × this\nPATIENCE               = 5    # Early stopping: epochs without val_loss improvement\nDROPOUT                = 0.4  # Classifier head dropout (default was 0.3)\n\n# Path to NB02 output dataset\nCACHE_INPUT_DIR = Path('/kaggle/input/notebooks/harshitghosh/nb02eda')  # <-- verify this name\nMANIFEST_PATH   = CACHE_INPUT_DIR / 'manifest.csv'\nNPY_CACHE_DIR   = CACHE_INPUT_DIR / 'cache'\nSTATS_PATH      = CACHE_INPUT_DIR / 'normalization_stats.json'\n\nCHECKPOINT_PATH  = Path('/kaggle/working/checkpoint.pth')\nMETRICS_LOG_PATH = Path('/kaggle/working/training_metrics.csv')\n\n# ═══════════════════════════════════════════════════════════════════════════\n# ██  ENVIRONMENT INFO  ██\n# ═══════════════════════════════════════════════════════════════════════════\n\nDEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(f'Device : {DEVICE}')\nif DEVICE == 'cuda':\n    print(f'GPU    : {torch.cuda.get_device_name(0)}')\n    print(f'VRAM   : {torch.cuda.get_device_properties(0).total_memory/1e9:.1f} GB')\n\n# ═══════════════════════════════════════════════════════════════════════════\n# ██  VERIFY NB02 OUTPUT MOUNT  ██\n# ═══════════════════════════════════════════════════════════════════════════\n\nprint(\"\\nChecking mounted input datasets...\")\nprint(\"Available input folders:\")\nprint(os.listdir('/kaggle/input'))\n\nassert CACHE_INPUT_DIR.exists(), f\"❌ CACHE_INPUT_DIR not found: {CACHE_INPUT_DIR}\"\nassert MANIFEST_PATH.exists(),   f\"❌ manifest.csv not found at {MANIFEST_PATH}\"\nassert NPY_CACHE_DIR.exists(),   f\"❌ cache directory not found at {NPY_CACHE_DIR}\"\nassert STATS_PATH.exists(),      f\"❌ normalization_stats.json not found\"\n\nprint(\"✅ NB02 dataset found successfully.\")\n\n# Check number of cached files\nnpy_files = list(NPY_CACHE_DIR.glob(\"*.npy\"))\nprint(f\"NPY files found: {len(npy_files)}\")\n\nassert len(npy_files) > 0, \"❌ No NPY files detected in cache directory.\"\n\nprint(\"✅ Cache directory verified.\")\n\n# Load normalization stats\nwith open(STATS_PATH) as f:\n\n    norm_stats = json.load(f)\n\nprint(\"Std :\", norm_stats[\"std\"])\n\nprint(\"Normalization stats loaded:\")\nprint(\"Mean:\", norm_stats[\"mean\"])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:20:19.463818Z","iopub.execute_input":"2026-02-20T17:20:19.46409Z","iopub.status.idle":"2026-02-20T17:20:24.335414Z","shell.execute_reply.started":"2026-02-20T17:20:19.464067Z","shell.execute_reply":"2026-02-20T17:20:24.334653Z"}},"outputs":[],"execution_count":null},{"id":"4d1f9612","cell_type":"code","source":"# ── 1. Imports ────────────────────────────────────────────────────────────\nimport os, gc, time, random\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\n# AMP: use torch.amp.autocast / torch.amp.GradScaler (non-deprecated)\n\nimport torchvision.transforms as T\nimport torchvision.models as models\n\nfrom sklearn.metrics import (\n    roc_auc_score, confusion_matrix,\n    roc_curve, precision_recall_curve\n)\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\n\n\ndef seed_everything(s: int):\n    random.seed(s)\n    np.random.seed(s)\n    torch.manual_seed(s)\n    torch.cuda.manual_seed_all(s)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nseed_everything(SEED)\nprint('Imports OK.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:20:24.33671Z","iopub.execute_input":"2026-02-20T17:20:24.337047Z","iopub.status.idle":"2026-02-20T17:20:28.589457Z","shell.execute_reply.started":"2026-02-20T17:20:24.337023Z","shell.execute_reply":"2026-02-20T17:20:28.588858Z"}},"outputs":[],"execution_count":null},{"id":"b0593e4c","cell_type":"code","source":"# ── 2. Dataset ────────────────────────────────────────────────────────────\nimport json as _json\n\nSUBTYPES = ['any', 'epidural', 'intraparenchymal',\n            'intraventricular', 'subarachnoid', 'subdural']\n\n# ─── Load dataset-specific normalization stats (with ImageNet fallback) ──\n_norm_path = os.path.join(CACHE_INPUT_DIR, 'normalization_stats.json')\nif os.path.exists(_norm_path):\n    with open(_norm_path) as f:\n        _norm = _json.load(f)\n    MEAN = _norm['mean']\n    STD  = _norm['std']\n    print(f'Loaded dataset-specific normalization: mean={MEAN}, std={STD}')\nelse:\n    MEAN = [0.485, 0.456, 0.406]    # ImageNet fallback\n    STD  = [0.229, 0.224, 0.225]\n    print(f'WARNING: normalization_stats.json not found. Using ImageNet defaults.')\n    print(f'  mean={MEAN}, std={STD}')\n\n\nclass ICHDataset(Dataset):\n    \"\"\"Loads cached NPY arrays (uint8 [0,255], H×W×3) and returns (image, label) pairs.\n\n    For binary detection of *any* hemorrhage.  If multi-label targets\n    are needed later, change `label_col` or return the full 6-dim vector.\n    \"\"\"\n\n    def __init__(self, df: pd.DataFrame, npy_root: str, transform, label_col: str = 'any'):\n        self.df        = df.reset_index(drop=True)\n        self.npy_root  = npy_root\n        self.transform = transform\n        self.label_col = label_col\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row  = self.df.iloc[idx]\n        path = os.path.join(self.npy_root, f'{row[\"image_id\"]}.npy')\n        try:\n            img = np.load(path)                        # uint8 H×W×3 [0,255]\n        except Exception:\n            img = np.zeros((IMG_SIZE, IMG_SIZE, 3), dtype=np.uint8)\n        img  = self.transform(img)                     # → Tensor C×H×W normalised\n        label = torch.tensor(float(row[self.label_col]), dtype=torch.float32)\n        return img, label\n\n\n# ─── Transforms ──────────────────────────────────────────────────────────\ntrain_transform = T.Compose([\n    T.ToPILImage(),\n    T.RandomHorizontalFlip(),\n    T.RandomRotation(degrees=10),\n    T.ColorJitter(brightness=0.1, contrast=0.1),\n    T.ToTensor(),\n    T.Normalize(mean=MEAN, std=STD),\n])\n\nval_transform = T.Compose([\n    T.ToPILImage(),\n    T.ToTensor(),\n    T.Normalize(mean=MEAN, std=STD),\n])\n\nprint('Dataset class + transforms defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:20:28.590405Z","iopub.execute_input":"2026-02-20T17:20:28.590796Z","iopub.status.idle":"2026-02-20T17:20:28.601369Z","shell.execute_reply.started":"2026-02-20T17:20:28.590772Z","shell.execute_reply":"2026-02-20T17:20:28.600659Z"}},"outputs":[],"execution_count":null},{"id":"1c33f1ae","cell_type":"code","source":"# ── 3. Load manifest and build DataLoaders ────────────────────────────────\nmanifest = pd.read_csv(MANIFEST_PATH)\nprint(f'Manifest loaded: {len(manifest):,} rows')\nprint(f'Cache files: {len(os.listdir(NPY_CACHE_DIR)):,} .npy files found')\n\ntrain_df = manifest[manifest['split'] == 'train'].reset_index(drop=True)\nval_df   = manifest[manifest['split'] == 'val'].reset_index(drop=True)\nprint(f'Train: {len(train_df):,}  |  Val: {len(val_df):,}')\n\n# ─── Weighted sampler to handle class imbalance ──────────────────────────\nfrom torch.utils.data import WeightedRandomSampler\n\npos_count = int(train_df['any'].sum())\nneg_count = len(train_df) - pos_count\npos_weight_val = neg_count / pos_count     # for BCEWithLogitsLoss\nprint(f'Pos: {pos_count:,}  Neg: {neg_count:,}  pos_weight={pos_weight_val:.2f}')\n\nsample_weights = np.where(train_df['any'].values == 1,\n                          neg_count / pos_count,\n                          1.0)\nsampler = WeightedRandomSampler(\n    weights=sample_weights.tolist(),\n    num_samples=len(train_df),\n    replacement=True\n)\n\ntrain_ds = ICHDataset(train_df, NPY_CACHE_DIR, train_transform)\nval_ds   = ICHDataset(val_df,   NPY_CACHE_DIR, val_transform)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                          num_workers=NUM_WORKERS, pin_memory=True)\nval_loader   = DataLoader(val_ds,   batch_size=BATCH_SIZE * 2, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=True)\n\nprint(f'Batches per epoch — train: {len(train_loader)}, val: {len(val_loader)}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:20:28.603488Z","iopub.execute_input":"2026-02-20T17:20:28.603789Z","iopub.status.idle":"2026-02-20T17:20:28.869741Z","shell.execute_reply.started":"2026-02-20T17:20:28.603767Z","shell.execute_reply":"2026-02-20T17:20:28.868986Z"}},"outputs":[],"execution_count":null},{"id":"442bacf9","cell_type":"code","source":"# ── 4. Model definition ───────────────────────────────────────────────────\ndef build_model(arch: str, pretrained: bool = True) -> nn.Module:\n    \"\"\"\n    Build a binary classifier on top of EfficientNet-B0 or ResNet-50.\n    Output: single logit (no sigmoid — use BCEWithLogitsLoss).\n    \"\"\"\n    if arch == 'efficientnet_b0':\n        weights = models.EfficientNet_B0_Weights.DEFAULT if pretrained else None\n        m = models.efficientnet_b0(weights=weights)\n        in_features = m.classifier[1].in_features\n        m.classifier = nn.Sequential(\n            nn.Dropout(p=DROPOUT),\n            nn.Linear(in_features, 1)\n        )\n\n    elif arch == 'resnet50':\n        weights = models.ResNet50_Weights.DEFAULT if pretrained else None\n        m = models.resnet50(weights=weights)\n        in_features = m.fc.in_features\n        m.fc = nn.Sequential(\n            nn.Dropout(p=DROPOUT),\n            nn.Linear(in_features, 1)\n        )\n\n    else:\n        raise ValueError(f'Unknown architecture: {arch}')\n\n    return m\n\n\n# ─── Backbone freeze / unfreeze helper ───────────────────────────────────\ndef set_backbone_frozen(model, arch, freeze: bool):\n    \"\"\"Freeze or unfreeze all backbone layers (keep classifier head trainable).\"\"\"\n    head_names = ('classifier',) if arch == 'efficientnet_b0' else ('fc',)\n    for name, p in model.named_parameters():\n        if not any(h in name for h in head_names):\n            p.requires_grad = not freeze\n    tag = 'FROZEN' if freeze else 'UNFROZEN'\n    n_locked = sum(1 for p in model.parameters() if not p.requires_grad)\n    print(f'  Backbone {tag} — {n_locked} parameter tensors locked')\n\n\nmodel = build_model(ARCH, pretrained=True).to(DEVICE)\ntotal_params = sum(p.numel() for p in model.parameters())\ntrain_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f'Model    : {ARCH}')\nprint(f'Params   : {total_params:,}  (trainable: {train_params:,})')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:20:28.870588Z","iopub.execute_input":"2026-02-20T17:20:28.870851Z","iopub.status.idle":"2026-02-20T17:20:29.513197Z","shell.execute_reply.started":"2026-02-20T17:20:28.870828Z","shell.execute_reply":"2026-02-20T17:20:29.512568Z"}},"outputs":[],"execution_count":null},{"id":"313df8ba","cell_type":"code","source":"# ── 5. Loss, optimiser, scheduler ─────────────────────────────────────────\npos_weight_tensor = torch.tensor([pos_weight_val], device=DEVICE)\ncriterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight_tensor)\n\n# ─── Discriminative LR: backbone trains slower than classifier head ──────\nhead_names = ('classifier',) if ARCH == 'efficientnet_b0' else ('fc',)\nbackbone_params = [p for n, p in model.named_parameters()\n                   if not any(h in n for h in head_names)]\nhead_params     = [p for n, p in model.named_parameters()\n                   if any(h in n for h in head_names)]\nprint(f'Param groups — backbone: {len(backbone_params)} tensors '\n      f'(lr={BASE_LR * BACKBONE_LR_FACTOR:.1e}), '\n      f'head: {len(head_params)} tensors (lr={BASE_LR:.1e})')\n\noptimizer = optim.AdamW([\n    {'params': backbone_params, 'lr': BASE_LR * BACKBONE_LR_FACTOR},\n    {'params': head_params,     'lr': BASE_LR},\n], weight_decay=WEIGHT_DECAY)\n\n# CosineAnnealingLR over the TOTAL planned epochs\nscheduler = optim.lr_scheduler.CosineAnnealingLR(\n    optimizer, T_max=TOTAL_EPOCHS, eta_min=1e-6\n)\n\nscaler = torch.amp.GradScaler('cuda')   # mixed precision (non-deprecated)\n\nprint('Loss / optimizer / scheduler ready.')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:20:29.513992Z","iopub.execute_input":"2026-02-20T17:20:29.514278Z","iopub.status.idle":"2026-02-20T17:20:29.522799Z","shell.execute_reply.started":"2026-02-20T17:20:29.514249Z","shell.execute_reply":"2026-02-20T17:20:29.522059Z"}},"outputs":[],"execution_count":null},{"id":"60da953b","cell_type":"code","source":"# ── 6. Load previous checkpoint (session chaining) ────────────────────────\nSTART_EPOCH = 0\nmetrics_history = []\nbest_val_loss = float('inf')\npatience_counter = 0\nearly_stopped = False\n\nif PREV_CHECKPOINT_DIR is not None:\n    ckpt_path = Path(PREV_CHECKPOINT_DIR) / 'checkpoint.pth'\n    log_path  = Path(PREV_CHECKPOINT_DIR) / 'training_metrics.csv'\n\n    if ckpt_path.exists():\n        print(f'Loading checkpoint: {ckpt_path}')\n        ckpt = torch.load(str(ckpt_path), map_location=DEVICE)\n\n        model.load_state_dict(ckpt['model_state_dict'])\n        START_EPOCH = ckpt['epoch'] + 1\n\n        # Check optimizer compatibility (param group count must match)\n        saved_groups = len(ckpt['optimizer_state_dict']['param_groups'])\n        if saved_groups == len(optimizer.param_groups):\n            optimizer.load_state_dict(ckpt['optimizer_state_dict'])\n            scheduler.load_state_dict(ckpt['scheduler_state_dict'])\n            scaler.load_state_dict(ckpt['scaler_state_dict'])\n        else:\n            print(f'⚠ Optimizer layout changed ({saved_groups} → '\n                  f'{len(optimizer.param_groups)} param groups) — '\n                  f'model weights loaded, optimizer/scheduler reset.')\n\n        # Restore early-stopping state (backward-compatible)\n        best_val_loss    = ckpt.get('best_val_loss', float('inf'))\n        patience_counter = ckpt.get('patience_counter', 0)\n\n        # Check if previous session already triggered early stopping\n        if ckpt.get('early_stopped', False):\n            print('⚠ Previous session already triggered early stopping.')\n            print('  Set PATIENCE higher or adjust hyperparams to continue.')\n\n        print(f'Resuming from epoch {START_EPOCH}  '\n              f'(best_val_loss={best_val_loss:.5f}, '\n              f'patience={patience_counter}/{PATIENCE})')\n    else:\n        print(f'WARNING: checkpoint not found at {ckpt_path}. Starting from scratch.')\n\n    if log_path.exists():\n        prev_log = pd.read_csv(log_path)\n        metrics_history = prev_log.to_dict('records')\n        print(f'Loaded {len(metrics_history)} previous epoch records')\n        # Fallback: compute best_val_loss from history if not in checkpoint\n        if best_val_loss == float('inf') and 'val_loss' in prev_log.columns:\n            best_val_loss = float(prev_log['val_loss'].min())\n            print(f'  (best_val_loss inferred from log: {best_val_loss:.5f})')\nelse:\n    print('No previous checkpoint. Starting from scratch (epoch 0).')\n\n# ─── Freeze backbone if still within the freeze period ────────────────────\nif START_EPOCH < FREEZE_BACKBONE_EPOCHS:\n    set_backbone_frozen(model, ARCH, freeze=True)\n    print(f'🧊 Backbone frozen for epochs 0–{FREEZE_BACKBONE_EPOCHS - 1} (head-only training)')\nelse:\n    print(f'Backbone trainable (freeze period was epochs 0–{FREEZE_BACKBONE_EPOCHS - 1})')\n\nEND_EPOCH = START_EPOCH + N_EPOCHS_THIS_SESSION\nprint(f'This session: epoch {START_EPOCH} → {END_EPOCH - 1}  (of 0-{TOTAL_EPOCHS-1} total)')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:20:29.523594Z","iopub.execute_input":"2026-02-20T17:20:29.523861Z","iopub.status.idle":"2026-02-20T17:20:30.420573Z","shell.execute_reply.started":"2026-02-20T17:20:29.523833Z","shell.execute_reply":"2026-02-20T17:20:30.419824Z"}},"outputs":[],"execution_count":null},{"id":"d130f5d4","cell_type":"code","source":"# ── 7. Evaluation helpers ─────────────────────────────────────────────────\n@torch.no_grad()\ndef evaluate(model: nn.Module, loader: DataLoader) -> dict:\n    \"\"\"Run inference on loader and return metrics dict.\"\"\"\n    model.eval()\n    all_logits, all_labels = [], []\n\n    for imgs, labels in tqdm(loader, desc='Val', leave=False):\n        imgs   = imgs.to(DEVICE)\n        with torch.amp.autocast('cuda'):\n            logits = model(imgs).squeeze(1)   # (B,)\n        all_logits.append(logits.cpu().float())\n        all_labels.append(labels)\n\n    all_logits = torch.cat(all_logits).numpy()\n    all_labels = torch.cat(all_labels).numpy()\n    all_probs  = torch.sigmoid(torch.tensor(all_logits)).numpy()\n\n    # Loss\n    val_loss = nn.BCEWithLogitsLoss()(\n        torch.tensor(all_logits),\n        torch.tensor(all_labels)\n    ).item()\n\n    # AUC\n    auc = roc_auc_score(all_labels, all_probs)\n\n    # Sensitivity & Specificity at Youden-optimal threshold\n    fpr, tpr, thresholds = roc_curve(all_labels, all_probs)\n    youden_idx = np.argmax(tpr - fpr)\n    best_thresh = float(thresholds[youden_idx])\n    preds = (all_probs >= best_thresh).astype(int)\n    tn, fp, fn, tp = confusion_matrix(all_labels, preds).ravel()\n    sensitivity  = tp / (tp + fn + 1e-9)\n    specificity  = tn / (tn + fp + 1e-9)\n    precision    = tp / (tp + fp + 1e-9)\n    f1           = 2 * precision * sensitivity / (precision + sensitivity + 1e-9)\n\n    return {\n        'val_loss'    : round(val_loss, 5),\n        'val_auc'     : round(float(auc), 5),\n        'sensitivity' : round(float(sensitivity), 5),\n        'specificity' : round(float(specificity), 5),\n        'precision'   : round(float(precision), 5),\n        'f1'          : round(float(f1), 5),\n        'best_thresh' : round(float(best_thresh), 4),\n        'tp': int(tp), 'tn': int(tn), 'fp': int(fp), 'fn': int(fn),\n        'all_probs'   : all_probs,\n        'all_labels'  : all_labels,\n        'fpr': fpr, 'tpr': tpr,\n    }\n\n\nprint('Evaluation function defined.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:20:30.421675Z","iopub.execute_input":"2026-02-20T17:20:30.422003Z","iopub.status.idle":"2026-02-20T17:20:30.432256Z","shell.execute_reply.started":"2026-02-20T17:20:30.421977Z","shell.execute_reply":"2026-02-20T17:20:30.431546Z"}},"outputs":[],"execution_count":null},{"id":"7e6b762d","cell_type":"code","source":"# ── 8. Training loop ──────────────────────────────────────────────────────\nbest_auc = max((r['val_auc'] for r in metrics_history), default=0.0)\n\nfor epoch in range(START_EPOCH, END_EPOCH):\n    # ── Unfreeze backbone at the scheduled epoch ──────────────────────\n    if epoch == FREEZE_BACKBONE_EPOCHS:\n        set_backbone_frozen(model, ARCH, freeze=False)\n        trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n        print(f'🔓 Backbone unfrozen at epoch {epoch} — {trainable:,} trainable params')\n\n    epoch_start = time.time()\n    model.train()\n\n    train_loss  = 0.0\n    n_batches   = 0\n    nan_count   = 0   # track NaN/Inf losses\n\n    pbar = tqdm(train_loader, desc=f'Epoch {epoch}/{TOTAL_EPOCHS-1} [train]', leave=True)\n    for imgs, labels in pbar:\n        imgs   = imgs.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.amp.autocast('cuda'):\n            logits = model(imgs).squeeze(1)        # (B,)\n            loss   = criterion(logits, labels)\n\n        # ── NaN / Inf divergence guard ────────────────────────────────\n        if not torch.isfinite(loss):\n            nan_count += 1\n            if nan_count >= 5:\n                raise RuntimeError(\n                    f'Training diverged: {nan_count} NaN/Inf losses in epoch {epoch}. '\n                    f'Try reducing lr ({BASE_LR}) or batch size ({BATCH_SIZE}).'\n                )\n            print(f'  ⚠ NaN/Inf loss at batch {n_batches} — skipping update')\n            optimizer.zero_grad(set_to_none=True)\n            continue\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)\n        scaler.update()\n\n        train_loss += loss.item()\n        n_batches  += 1\n        pbar.set_postfix(loss=f'{train_loss/n_batches:.4f}')\n\n    scheduler.step()\n\n    if n_batches == 0:\n        raise RuntimeError(f'No valid batches in epoch {epoch} — all losses were NaN/Inf!')\n\n    train_loss /= n_batches\n    metrics = evaluate(model, val_loader)\n\n    elapsed = time.time() - epoch_start\n    lrs = scheduler.get_last_lr()\n    row = {\n        'epoch'       : epoch,\n        'train_loss'  : round(train_loss, 5),\n        'lr_backbone' : lrs[0],\n        'lr_head'     : lrs[1] if len(lrs) > 1 else lrs[0],\n        'elapsed_s'   : round(elapsed, 1),\n        'nan_batches' : nan_count,\n        **{k: v for k, v in metrics.items()\n           if k not in ('all_probs', 'all_labels', 'fpr', 'tpr')}\n    }\n    metrics_history.append(row)\n\n    # ── Save metrics log ──────────────────────────────────────────────────\n    pd.DataFrame(metrics_history).to_csv(METRICS_LOG_PATH, index=False)\n\n    # ── Early stopping check ──────────────────────────────────────────────\n    if metrics['val_loss'] < best_val_loss:\n        best_val_loss = metrics['val_loss']\n        patience_counter = 0\n    else:\n        patience_counter += 1\n\n    # ── Save checkpoint (always, after every epoch) ───────────────────────\n    torch.save({\n        'epoch'                : epoch,\n        'arch'                 : ARCH,\n        'model_state_dict'     : model.state_dict(),\n        'optimizer_state_dict' : optimizer.state_dict(),\n        'scheduler_state_dict' : scheduler.state_dict(),\n        'scaler_state_dict'    : scaler.state_dict(),\n        'best_thresh'          : metrics['best_thresh'],\n        'best_val_loss'        : best_val_loss,\n        'patience_counter'     : patience_counter,\n        'early_stopped'        : False,\n    }, CHECKPOINT_PATH)\n\n    # ── Save best model separately ────────────────────────────────────────\n    if metrics['val_auc'] > best_auc:\n        best_auc = metrics['val_auc']\n        torch.save(model.state_dict(), '/kaggle/working/best_model.pth')\n        print(f'  ★ New best AUC: {best_auc:.5f} — saved best_model.pth')\n\n    frozen_tag = ' [head-only]' if epoch < FREEZE_BACKBONE_EPOCHS else ''\n    print(f'  Epoch {epoch:02d}{frozen_tag} | '\n          f'train_loss={train_loss:.4f} | '\n          f'val_loss={metrics[\"val_loss\"]:.4f} | '\n          f'AUC={metrics[\"val_auc\"]:.4f} | '\n          f'Sens={metrics[\"sensitivity\"]:.4f} | '\n          f'Spec={metrics[\"specificity\"]:.4f} | '\n          f'patience={patience_counter}/{PATIENCE} | '\n          f'NaN={nan_count} | '\n          f'{elapsed:.0f}s')\n\n    # ── Break if patience exhausted ───────────────────────────────────────\n    if patience_counter >= PATIENCE:\n        print(f'\\n⛔ Early stopping at epoch {epoch} '\n              f'(no val_loss improvement for {PATIENCE} epochs)')\n        early_stopped = True\n        # Update checkpoint to record early stopping\n        ckpt_es = torch.load(CHECKPOINT_PATH, map_location='cpu')\n        ckpt_es['early_stopped'] = True\n        torch.save(ckpt_es, CHECKPOINT_PATH)\n        break\n\nprint('\\nSession training complete.')\nif early_stopped:\n    print(f'  ⛔ Stopped early at epoch {epoch} (best val_loss={best_val_loss:.5f})')\nelse:\n    print(f'  Completed epochs {START_EPOCH}–{END_EPOCH - 1}')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:20:30.433243Z","iopub.execute_input":"2026-02-20T17:20:30.433515Z","iopub.status.idle":"2026-02-20T17:46:35.245729Z","shell.execute_reply.started":"2026-02-20T17:20:30.433485Z","shell.execute_reply":"2026-02-20T17:46:35.244683Z"}},"outputs":[],"execution_count":null},{"id":"2948fa1f","cell_type":"code","source":"# ── 9. Plot learning curves ───────────────────────────────────────────────\nlog = pd.read_csv(METRICS_LOG_PATH)\n\nfig, axes = plt.subplots(1, 3, figsize=(15, 4))\n\naxes[0].plot(log['epoch'], log['train_loss'], label='Train loss')\naxes[0].plot(log['epoch'], log['val_loss'],   label='Val loss')\naxes[0].set(title='Loss', xlabel='Epoch', ylabel='BCE Loss')\naxes[0].legend()\n\naxes[1].plot(log['epoch'], log['val_auc'], color='tab:orange')\naxes[1].set(title='Validation AUC-ROC', xlabel='Epoch', ylabel='AUC')\naxes[1].set_ylim([0.5, 1.0])\n\naxes[2].plot(log['epoch'], log['sensitivity'], label='Sensitivity')\naxes[2].plot(log['epoch'], log['specificity'], label='Specificity')\naxes[2].plot(log['epoch'], log['f1'],          label='F1')\naxes[2].set(title='Sensitivity / Specificity / F1', xlabel='Epoch')\naxes[2].legend()\n\nplt.tight_layout()\nplt.savefig('/kaggle/working/learning_curves.png', bbox_inches='tight')\nplt.show()\n\nprint(log[['epoch', 'train_loss', 'val_loss', 'val_auc',\n           'sensitivity', 'specificity', 'f1']].to_string(index=False))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:46:35.249241Z","iopub.execute_input":"2026-02-20T17:46:35.24953Z","iopub.status.idle":"2026-02-20T17:46:36.035209Z","shell.execute_reply.started":"2026-02-20T17:46:35.249497Z","shell.execute_reply":"2026-02-20T17:46:36.034663Z"}},"outputs":[],"execution_count":null},{"id":"24cc9bf6","cell_type":"code","source":"# ── 10. ROC curve for last validation run ────────────────────────────────\n# Rerun evaluate on best model to get ROC arrays\nbest_model_path = '/kaggle/working/best_model.pth'\nmodel.load_state_dict(torch.load(best_model_path, map_location=DEVICE))\nval_metrics = evaluate(model, val_loader)\n\nplt.figure(figsize=(6, 6))\nplt.plot(val_metrics['fpr'], val_metrics['tpr'],\n         label=f'AUC = {val_metrics[\"val_auc\"]:.4f}', color='tab:blue')\nplt.plot([0, 1], [0, 1], 'k--', linewidth=0.8)\nplt.scatter([1 - val_metrics['specificity']],\n            [val_metrics['sensitivity']],\n            color='red', zorder=5,\n            label=f'Youden pt (Thr={val_metrics[\"best_thresh\"]:.3f})')\nplt.xlabel('False Positive Rate'); plt.ylabel('True Positive Rate')\nplt.title('ROC Curve (Best Model)')\nplt.legend()\nplt.tight_layout()\nplt.savefig('/kaggle/working/roc_curve.png', bbox_inches='tight')\nplt.show()\n\nprint(f'Best model AUC   : {val_metrics[\"val_auc\"]:.5f}')\nprint(f'Sensitivity      : {val_metrics[\"sensitivity\"]:.5f}')\nprint(f'Specificity      : {val_metrics[\"specificity\"]:.5f}')\nprint(f'F1               : {val_metrics[\"f1\"]:.5f}')\nprint(f'Optimal threshold: {val_metrics[\"best_thresh\"]:.4f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:46:36.035972Z","iopub.execute_input":"2026-02-20T17:46:36.036168Z","iopub.status.idle":"2026-02-20T17:46:53.632138Z","shell.execute_reply.started":"2026-02-20T17:46:36.036149Z","shell.execute_reply":"2026-02-20T17:46:53.631304Z"}},"outputs":[],"execution_count":null},{"id":"c2085054","cell_type":"code","source":"# ── 11. Confusion matrix ──────────────────────────────────────────────────\nimport seaborn as sns\n\ncm = np.array([[val_metrics['tn'], val_metrics['fp']],\n               [val_metrics['fn'], val_metrics['tp']]])\n\nplt.figure(figsize=(5, 4))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n            xticklabels=['Pred Neg', 'Pred Pos'],\n            yticklabels=['True Neg', 'True Pos'])\nplt.title(f'Confusion Matrix (threshold={val_metrics[\"best_thresh\"]:.3f})')\nplt.tight_layout()\nplt.savefig('/kaggle/working/confusion_matrix.png', bbox_inches='tight')\nplt.show()\n\nprint('\\nOutputs saved to /kaggle/working/')\nprint('Files: checkpoint.pth, best_model.pth, training_metrics.csv,'\n      ' learning_curves.png, roc_curve.png, confusion_matrix.png')\nprint()\nprint('NEXT STEP:')\nprint(' Save Version → Save & Run All → commit this output')\nprint(' In the next session: set PREV_CHECKPOINT_DIR to this notebook output path')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:46:53.633901Z","iopub.execute_input":"2026-02-20T17:46:53.634641Z","iopub.status.idle":"2026-02-20T17:46:54.093298Z","shell.execute_reply.started":"2026-02-20T17:46:53.634582Z","shell.execute_reply":"2026-02-20T17:46:54.092662Z"}},"outputs":[],"execution_count":null},{"id":"9cbf91cd","cell_type":"code","source":"# ── HEALTH CHECK — automated output validation ────────────────────────────\nimport json as _json_hc\n\nerrors = []\n\n# Check checkpoint\nif not os.path.exists(CHECKPOINT_PATH):\n    errors.append('checkpoint.pth is MISSING')\nelse:\n    ckpt_hc = torch.load(CHECKPOINT_PATH, map_location='cpu')\n    if 'model_state_dict' not in ckpt_hc:\n        errors.append('checkpoint.pth missing model_state_dict')\n    if 'best_thresh' not in ckpt_hc:\n        errors.append('checkpoint.pth missing best_thresh')\n    ckpt_epoch = ckpt_hc.get('epoch', -1)\n    expected_epoch = END_EPOCH - 1\n    if ckpt_hc.get('early_stopped', False):\n        print(f'  ℹ Training stopped early at epoch {ckpt_epoch}')\n    elif ckpt_epoch != expected_epoch:\n        errors.append(f'checkpoint epoch={ckpt_epoch} != expected {expected_epoch}')\n\n# Check best model\nif not os.path.exists('/kaggle/working/best_model.pth'):\n    errors.append('best_model.pth is MISSING')\n\n# Check metrics log\nif not os.path.exists(METRICS_LOG_PATH):\n    errors.append('training_metrics.csv is MISSING')\nelse:\n    log_hc = pd.read_csv(METRICS_LOG_PATH)\n    if not early_stopped and len(log_hc) < N_EPOCHS_THIS_SESSION:\n        errors.append(f'metrics log has {len(log_hc)} rows, expected >= {N_EPOCHS_THIS_SESSION}')\n    last_auc = log_hc['val_auc'].iloc[-1]\n    if last_auc < 0.55:\n        errors.append(f'Final AUC={last_auc:.4f} is suspiciously low — possible training issue')\n    nan_total = log_hc.get('nan_batches', pd.Series([0])).sum()\n    if nan_total > 10:\n        errors.append(f'{nan_total} NaN batches detected across all epochs')\n\n# Check plots\nfor plot in ['learning_curves.png', 'roc_curve.png', 'confusion_matrix.png']:\n    if not os.path.exists(f'/kaggle/working/{plot}'):\n        errors.append(f'Missing plot: {plot}')\n\nactual_last = metrics_history[-1]['epoch'] if metrics_history else START_EPOCH\nhealth = {\n    'notebook'          : '03_train_session',\n    'status'            : 'PASS' if not errors else 'FAIL',\n    'errors'            : errors,\n    'session_epochs'    : f'{START_EPOCH}-{actual_last}',\n    'early_stopped'     : early_stopped,\n    'best_auc'          : round(best_auc, 5),\n    'best_val_loss'     : round(best_val_loss, 5) if best_val_loss < float('inf') else None,\n    'final_train_loss'  : round(metrics_history[-1]['train_loss'], 5) if metrics_history else None,\n}\n\nwith open('/kaggle/working/health_check_nb03.json', 'w') as f:\n    _json_hc.dump(health, f, indent=2)\n\nif errors:\n    print('❌ HEALTH CHECK FAILED:')\n    for e in errors:\n        print(f'   • {e}')\nelse:\n    print('✅ HEALTH CHECK PASSED')\n    print(f'   Epochs trained : {START_EPOCH} → {actual_last}'\n          + (' (early stopped)' if early_stopped else ''))\n    print(f'   Best AUC       : {best_auc:.5f}')\n    print(f'   Best val_loss  : {best_val_loss:.5f}')\n    print(f'   Checkpoint     : saved')\n    print(f'   Metrics log    : {len(log_hc)} rows')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-20T17:46:54.094273Z","iopub.execute_input":"2026-02-20T17:46:54.094572Z","iopub.status.idle":"2026-02-20T17:46:54.230562Z","shell.execute_reply.started":"2026-02-20T17:46:54.094549Z","shell.execute_reply":"2026-02-20T17:46:54.229868Z"}},"outputs":[],"execution_count":null},{"id":"b42552c1-9c4b-49a7-a9d3-663543f1785c","cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}