{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10338,"databundleVersionId":862042},{"sourceType":"datasetVersion","sourceId":16260904,"datasetId":10424935,"databundleVersionId":17244818},{"sourceType":"kernelVersion","sourceId":316980196}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"3c0ce5d8-5ebc-4dea-b227-78e8a22db564","cell_type":"markdown","source":"# Notebook 04b — I-JEPA Full Fine-tuning RESUME\n**Mục đích:** Resume Full FT từ checkpoint epoch 20 với LR cao hơn (1e-4/1e-5).\n- LR cũ: HEAD=5e-5, ENC=5e-6 → quá thấp, AUC vẫn tăng đến epoch 20\n- LR mới: HEAD=1e-4, ENC=1e-5 → tăng 2×, đúng với Partial FT\n- Train thêm 15 epochs (epoch 21→35), resume từ best checkpoint epoch 20\n","metadata":{}},{"id":"0d35f0ab-fa21-49dd-bce4-7934f31dd808","cell_type":"code","source":"# ============================================================\n# CELL 1: IMPORTS\n# ============================================================\nimport os, gc, json, math, time, shutil\nfrom pathlib import Path\n\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\n\nfrom sklearn.metrics import (\n    roc_auc_score, f1_score, recall_score,\n    precision_score, confusion_matrix, accuracy_score\n)\n\ntry:\n    import timm\nexcept ImportError:\n    import subprocess; subprocess.run([\"pip\",\"install\",\"-q\",\"timm\"])\n    import timm\n\ntry:\n    import pydicom\nexcept ImportError:\n    import subprocess; subprocess.run([\"pip\",\"install\",\"-q\",\"pydicom\"])\n    import pydicom\n\nSEED = 42\ndef set_seed(s=42):\n    import random\n    random.seed(s); np.random.seed(s)\n    torch.manual_seed(s); torch.cuda.manual_seed_all(s)\nset_seed(SEED)\n\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)\nif torch.cuda.is_available():\n    print(\"GPU:\", torch.cuda.get_device_name(0))\n","metadata":{},"outputs":[],"execution_count":null},{"id":"02fd7a52-aba5-4b86-9ffb-3a214c988fa1","cell_type":"code","source":"# ============================================================\n# CELL 2: TIME GUARD + OUTPUT DIRS\n# ============================================================\nSESSION_SAFE_HOURS   = 11.5\nSESSION_SAFE_SECONDS = SESSION_SAFE_HOURS * 3600\nNOTEBOOK_START_TIME  = time.time()\n\nOUTPUT_DIR = Path(\"/kaggle/working/nb04b_resume\")\nCKPT_DIR   = OUTPUT_DIR / \"checkpoints\"\nLOG_DIR    = OUTPUT_DIR / \"logs\"\nFIG_DIR    = OUTPUT_DIR / \"figures\"\nfor d in [OUTPUT_DIR, CKPT_DIR, LOG_DIR, FIG_DIR]:\n    d.mkdir(parents=True, exist_ok=True)\n\nprint(f\"Time guard: {SESSION_SAFE_HOURS}h\")\nprint(\"Output:\", OUTPUT_DIR)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"e5c8d568-92f9-4398-b56f-ab566aa7d62d","cell_type":"code","source":"# ============================================================\n# CELL 3: TÌM INPUT FILES\n# ============================================================\nINPUT_ROOT   = Path(\"/kaggle/input\")\nWORKING_ROOT = Path(\"/kaggle/working\")\n\ndef find_file(filename):\n    for root in [WORKING_ROOT, INPUT_ROOT]:\n        matches = list(root.rglob(filename))\n        if matches:\n            return matches[0]\n    return None\n\n# ── Encoder checkpoint từ Notebook 03 ────────────────────────\nIJEPA_ENCODER_CKPT = find_file(\"ijepa_vit_small_nih_50k_best_encoder.pth\")\nprint(\"Encoder ckpt  :\", IJEPA_ENCODER_CKPT)\n\n# ── Full FT checkpoint epoch 20 (để resume) ──────────────────\n# File này từ output Notebook 04 version cũ\n# Tên file: ijepa_full_finetune_best.pth\nFULL_FT_RESUME_CKPT = find_file(\"ijepa_full_finetune_best.pth\")\nprint(\"Full FT resume:\", FULL_FT_RESUME_CKPT)\n\n# ── RSNA CSVs ─────────────────────────────────────────────────\nRSNA_TRAIN_CSV = find_file(\"rsna_train.csv\")\nRSNA_VAL_CSV   = find_file(\"rsna_val.csv\")\nRSNA_TEST_CSV  = find_file(\"rsna_test.csv\")\nprint(\"Train CSV:\", RSNA_TRAIN_CSV)\nprint(\"Val CSV  :\", RSNA_VAL_CSV)\nprint(\"Test CSV :\", RSNA_TEST_CSV)\n\nassert IJEPA_ENCODER_CKPT   is not None, \"Không tìm thấy encoder checkpoint NB03!\"\nassert FULL_FT_RESUME_CKPT  is not None, \"Không tìm thấy full_ft_best.pth từ NB04!\"\nassert RSNA_TRAIN_CSV       is not None, \"Không tìm thấy rsna_train.csv!\"\nprint(\"\\nAll inputs found ✓\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"f977b1c9-4c59-4938-a0e6-1af20eeef1c5","cell_type":"code","source":"# ============================================================\n# CELL 4: ĐỌC RSNA METADATA + TÍNH POS_WEIGHT\n# ============================================================\ntrain_df = pd.read_csv(RSNA_TRAIN_CSV)\nval_df   = pd.read_csv(RSNA_VAL_CSV)\ntest_df  = pd.read_csv(RSNA_TEST_CSV)\n\nprint(f\"Train: {train_df.shape} | Val: {val_df.shape} | Test: {test_df.shape}\")\n\nn_neg = (train_df[\"label\"] == 0).sum()\nn_pos = (train_df[\"label\"] == 1).sum()\nPOS_WEIGHT = torch.tensor([n_neg / n_pos], dtype=torch.float32)\nprint(f\"Train neg={n_neg} pos={n_pos} → POS_WEIGHT={POS_WEIGHT.item():.4f}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"db6f3e10-bc7f-48b2-8b46-494af39894fe","cell_type":"code","source":"# ============================================================\n# CELL 5: FIX ĐƯỜNG DẪN ẢNH RSNA\n# ============================================================\ndef count_missing(df):\n    return (~df[\"image_path\"].apply(lambda x: Path(x).exists())).sum()\n\nmissing = count_missing(train_df)\nprint(f\"Missing before fix: {missing}\")\n\nif missing > 0:\n    all_imgs = list(INPUT_ROOT.rglob(\"*.dcm\")) + list(INPUT_ROOT.rglob(\"*.png\"))\n    fname2path = {p.name: str(p) for p in all_imgs}\n\n    def fix_path(p):\n        return fname2path.get(Path(p).name, None)\n\n    for df in [train_df, val_df, test_df]:\n        df[\"image_path\"] = df[\"image_path\"].apply(fix_path)\n\n    train_df.dropna(subset=[\"image_path\"], inplace=True)\n    val_df.dropna(subset=[\"image_path\"], inplace=True)\n    test_df.dropna(subset=[\"image_path\"], inplace=True)\n\nprint(f\"Missing after fix: {count_missing(train_df)}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"564b625d-56f5-4dbd-88f1-6d04e3437330","cell_type":"code","source":"# ============================================================\n# CELL 6: DATASET + TRANSFORMS\n# ============================================================\nIMG_SIZE = 224\nIMAGENET_MEAN = [0.485, 0.456, 0.406]\nIMAGENET_STD  = [0.229, 0.224, 0.225]\n\ntrain_transform = T.Compose([\n    T.Resize((IMG_SIZE, IMG_SIZE)),\n    T.RandomHorizontalFlip(p=0.5),\n    T.RandomRotation(degrees=7),\n    T.ColorJitter(brightness=0.1, contrast=0.1),\n    T.ToTensor(),\n    T.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n])\neval_transform = T.Compose([\n    T.Resize((IMG_SIZE, IMG_SIZE)),\n    T.ToTensor(),\n    T.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n])\n\nclass RSNADataset(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): return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        path = row[\"image_path\"]\n        if path.endswith(\".dcm\"):\n            dcm = pydicom.dcmread(path)\n            arr = dcm.pixel_array.astype(np.float32)\n            arr = (arr - arr.min()) / (arr.max() - arr.min() + 1e-8) * 255\n            img = Image.fromarray(arr.astype(np.uint8)).convert(\"RGB\")\n        else:\n            img = Image.open(path).convert(\"RGB\")\n        if self.transform:\n            img = self.transform(img)\n        return img, torch.tensor(row[\"label\"], dtype=torch.float32)\n\nBATCH_SIZE  = 16\nNUM_WORKERS = 2\n\ntrain_loader = DataLoader(RSNADataset(train_df, train_transform),\n    batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True)\nval_loader   = DataLoader(RSNADataset(val_df,   eval_transform),\n    batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)\ntest_loader  = DataLoader(RSNADataset(test_df,  eval_transform),\n    batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)\n\nprint(f\"Train: {len(train_loader)} batches | Val: {len(val_loader)} | Test: {len(test_loader)}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"701cd81d-8a73-47e0-9e0d-2eeadd32860a","cell_type":"code","source":"# ============================================================\n# CELL 7: MODEL ARCHITECTURE\n# ============================================================\nclass IJEPAClassifier(nn.Module):\n    \"\"\"Giữ đúng kiến trúc gốc Notebook 04: LayerNorm → Dropout → Linear.\"\"\"\n    def __init__(self, encoder, embed_dim=384, dropout=0.2):\n        super().__init__()\n        self.encoder    = encoder\n        self.classifier = nn.Sequential(\n            nn.LayerNorm(embed_dim),\n            nn.Dropout(dropout),\n            nn.Linear(embed_dim, 1)\n        )\n\n    def forward(self, x):\n        features = self.encoder(x)   # [B, embed_dim] — dùng forward() giống NB04 gốc\n        logits   = self.classifier(features).squeeze(1)\n        return logits\n\n\ndef create_ijepa_classifier():\n    encoder = timm.create_model(\"vit_small_patch16_224\", pretrained=False, num_classes=0)\n    try:\n        enc_ckpt = torch.load(IJEPA_ENCODER_CKPT, map_location=\"cpu\", weights_only=False)\n    except TypeError:\n        enc_ckpt = torch.load(IJEPA_ENCODER_CKPT, map_location=\"cpu\")\n    state_key = \"encoder_state_dict\" if \"encoder_state_dict\" in enc_ckpt else \"student_encoder_state_dict\"\n    encoder.load_state_dict(enc_ckpt[state_key])\n    embed_dim = encoder.num_features\n    model = IJEPAClassifier(encoder, embed_dim=embed_dim, dropout=0.2)\n    return model\n\ndef unfreeze_all(model):\n    for p in model.parameters():\n        p.requires_grad = True\n    print(\"All parameters unfrozen.\")\n\ndef count_trainable(model):\n    n = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"Trainable params: {n:,}\")\n\nprint(\"Model architecture defined ✓\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"984f2d39-6d7f-465d-8ca6-7bf27a7490b0","cell_type":"code","source":"# ============================================================\n# CELL 8: TRAIN / EVALUATE HELPERS\n# ============================================================\n\ndef train_one_epoch(model, loader, optimizer, criterion, scaler, accum_steps):\n    model.train()\n    total_loss, total_steps = 0.0, 0\n    optimizer.zero_grad(set_to_none=True)\n    for step, (images, labels) in enumerate(tqdm(loader, leave=False)):\n        images, labels = images.to(DEVICE), labels.to(DEVICE)\n        with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):\n            logits = model(images)\n            loss   = criterion(logits, labels) / accum_steps\n        scaler.scale(loss).backward()\n        if (step + 1) % accum_steps == 0:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n        total_loss  += loss.item() * accum_steps\n        total_steps += 1\n    return total_loss / max(total_steps, 1)\n\n\ndef evaluate_model(model, loader, criterion, threshold=0.5):\n    model.eval()\n    all_probs, all_labels, total_loss, n = [], [], 0.0, 0\n    with torch.no_grad():\n        for images, labels in loader:\n            images, labels = images.to(DEVICE), labels.to(DEVICE)\n            with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):\n                logits = model(images)\n                loss   = criterion(logits, labels)\n            probs = torch.sigmoid(logits).cpu().numpy()\n            all_probs.extend(probs); all_labels.extend(labels.cpu().numpy())\n            total_loss += loss.item(); n += 1\n    all_probs  = np.array(all_probs)\n    all_labels = np.array(all_labels)\n    preds      = (all_probs >= threshold).astype(int)\n    auc  = roc_auc_score(all_labels, all_probs)\n    tn, fp, fn, tp = confusion_matrix(all_labels, preds, labels=[0,1]).ravel()\n    metrics = {\n        \"auc\": auc,\n        \"f1\":  f1_score(all_labels, preds, zero_division=0),\n        \"recall_pneumonia\": recall_score(all_labels, preds, zero_division=0),\n        \"specificity\": tn / (tn + fp + 1e-8),\n        \"precision\":   precision_score(all_labels, preds, zero_division=0),\n        \"accuracy\":    accuracy_score(all_labels, preds),\n        \"tn\": int(tn), \"fp\": int(fp), \"fn\": int(fn), \"tp\": int(tp),\n        \"threshold\": threshold\n    }\n    pred_df = pd.DataFrame({\"label\": all_labels, \"prob_pneumonia\": all_probs,\n                             \"pred\": preds})\n    return total_loss / max(n, 1), metrics, pred_df\n\n\ndef make_lr_scheduler(optimizer, num_epochs, warmup_ratio=0.1):\n    warmup_epochs = max(1, int(num_epochs * warmup_ratio))\n    def lr_lambda(ep):\n        if ep < warmup_epochs:\n            return float(ep + 1) / warmup_epochs\n        progress = (ep - warmup_epochs) / max(1, num_epochs - warmup_epochs)\n        return 0.5 * (1.0 + math.cos(math.pi * progress))\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\n\ndef disk_used_gb():\n    _, used, _ = shutil.disk_usage(\"/kaggle/working\")\n    return used / 1e9\n\n\ndef cleanup_checkpoints(model_name, keep_last=3):\n    def ep_num(p):\n        try: return int(p.stem.split(\"_epoch_\")[-1])\n        except: return -1\n    for pattern in [f\"{model_name}_last_epoch_*.pth\"]:\n        files = sorted(CKPT_DIR.glob(pattern), key=ep_num)\n        for old in files[:-keep_last]:\n            old.unlink(missing_ok=True)\n\nprint(\"Helpers defined ✓\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"feefa273-d952-4ec4-9926-9fb4eabec182","cell_type":"code","source":"# ============================================================\n# CELL 9: CONFIG RESUME FULL FT\n# ============================================================\n# ┌─────────────────────────────────────────────────────────┐\n# │ Thay đổi so với lần chạy cũ (notebook 04):              │\n# │   HEAD_LR  : 5e-5 → 1e-4  (tăng 2×)                   │\n# │   ENCODER_LR: 5e-6 → 1e-5  (tăng 2×)                  │\n# │   RESUME_EPOCH: 20 (bắt đầu đặt tên từ epoch 21)       │\n# │   ADDITIONAL_EPOCHS: 15 (train thêm epoch 21→35)       │\n# └─────────────────────────────────────────────────────────┘\n\nHEAD_LR            = 1e-4       # tăng từ 5e-5\nENCODER_LR         = 1e-5       # tăng từ 5e-6\nWEIGHT_DECAY       = 0.05\nACCUM_STEPS        = 2\nPATIENCE           = 6\nADDITIONAL_EPOCHS  = 15         # train thêm epoch 21→35\nRESUME_EPOCH       = 20         # epoch cuối của lần chạy trước\nWARMUP_EPOCHS_ABS  = 2          # warmup tuyến tính 2 epochs đầu của session này\nMODEL_NAME         = \"ijepa_full_finetune_resume\"\n\nprint(f\"HEAD_LR={HEAD_LR} | ENCODER_LR={ENCODER_LR}\")\nprint(f\"Resume từ epoch {RESUME_EPOCH}, train thêm {ADDITIONAL_EPOCHS} epochs\")\nprint(f\"Tổng epochs sau resume: epoch {RESUME_EPOCH + ADDITIONAL_EPOCHS}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"a9c4396a-9c5b-4d4b-b097-80537bfbd023","cell_type":"code","source":"# ============================================================\n# CELL 10: BUILD MODEL + LOAD CHECKPOINT EPOCH 20\n# ============================================================\nfull_model = create_ijepa_classifier()\nunfreeze_all(full_model)\ncount_trainable(full_model)\nfull_model = full_model.to(DEVICE)\n\n# Load state từ best checkpoint epoch 20\ntry:\n    ckpt = torch.load(FULL_FT_RESUME_CKPT, map_location=DEVICE, weights_only=False)\nexcept TypeError:\n    ckpt = torch.load(FULL_FT_RESUME_CKPT, map_location=DEVICE)\n\nfull_model.load_state_dict(ckpt[\"model_state_dict\"])\nprev_auc = ckpt.get(\"best_auc\", ckpt.get(\"val_auc\", None))\nprint(f\"Loaded checkpoint from: {FULL_FT_RESUME_CKPT.name}\")\nprint(f\"Previous best AUC: {prev_auc}\")\n\n# Verify load\nfull_model.eval()\nwith torch.no_grad():\n    dummy = torch.randn(2, 3, 224, 224).to(DEVICE)\n    out   = full_model(dummy)\n    print(f\"Forward pass OK — output shape: {out.shape}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"9af132ce-34c7-4fe1-ac88-1bdc57190ed0","cell_type":"code","source":"# ============================================================\n# CELL 11: OPTIMIZER + LR SCHEDULER (RESUME)\n# ============================================================\n# LambdaLR với warmup ngắn 2 epochs đầu của session mới\n# Sau warmup: cosine decay từ peak LR → ~0\n\ncriterion = nn.BCEWithLogitsLoss(pos_weight=POS_WEIGHT.to(DEVICE))\n\noptimizer = torch.optim.AdamW([\n    {\"params\": [p for p in full_model.encoder.parameters()    if p.requires_grad], \"lr\": ENCODER_LR},\n    {\"params\": [p for p in full_model.classifier.parameters() if p.requires_grad], \"lr\": HEAD_LR},\n], weight_decay=WEIGHT_DECAY)\n\ndef lr_lambda_resume(current_epoch):\n    if current_epoch < WARMUP_EPOCHS_ABS:\n        return float(current_epoch + 1) / WARMUP_EPOCHS_ABS\n    progress = (current_epoch - WARMUP_EPOCHS_ABS) / max(1, ADDITIONAL_EPOCHS - WARMUP_EPOCHS_ABS)\n    return 0.5 * (1.0 + math.cos(math.pi * progress))\n\nscheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda_resume)\nscaler    = torch.cuda.amp.GradScaler(enabled=torch.cuda.is_available())\n\n# Preview LR schedule\nprint(\"LR schedule preview (epoch = vòng lặp mới, không phải epoch tuyệt đối):\")\nfor ep in [0, 1, 2, 5, 10, 14]:\n    factor = lr_lambda_resume(ep)\n    print(f\"  Local ep {ep:2d}: HEAD_LR={HEAD_LR*factor:.2e} | ENC_LR={ENCODER_LR*factor:.2e}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"94fa9317-9b9f-4008-90de-71bc03b67672","cell_type":"code","source":"# ============================================================\n# CELL 12: VÒNG LẶP RESUME TRAINING\n# ============================================================\nbest_auc         = prev_auc if prev_auc else 0.8193  # AUC tốt nhất từ run cũ\nbest_epoch_abs   = RESUME_EPOCH                       # epoch tuyệt đối\npatience_counter = 0\nhistory          = []\nbest_ckpt_path   = CKPT_DIR / f\"{MODEL_NAME}_best.pth\"\n\nprint(f\"Starting resume from epoch {RESUME_EPOCH+1} → {RESUME_EPOCH+ADDITIONAL_EPOCHS}\")\nprint(f\"Best AUC to beat: {best_auc:.4f}\")\nprint(\"=\" * 58)\n\nfor local_ep in range(ADDITIONAL_EPOCHS):\n    abs_epoch = RESUME_EPOCH + local_ep + 1  # epoch tuyệt đối: 21, 22, ...\n\n    # Time guard\n    if time.time() - NOTEBOOK_START_TIME > SESSION_SAFE_SECONDS:\n        print(f\"\\n⏱  Time guard — dừng an toàn trước epoch {abs_epoch}.\")\n        break\n\n    current_lr_head = optimizer.param_groups[1][\"lr\"]\n    current_lr_enc  = optimizer.param_groups[0][\"lr\"]\n    print(f\"\\nEpoch {abs_epoch} (local {local_ep+1}/{ADDITIONAL_EPOCHS}) | \"\n          f\"HEAD_LR={current_lr_head:.2e} | ENC_LR={current_lr_enc:.2e}\")\n\n    # ── Train ────────────────────────────────────────────────\n    train_loss = train_one_epoch(\n        full_model, train_loader, optimizer, criterion, scaler, ACCUM_STEPS\n    )\n    scheduler.step()\n\n    # ── Validate ─────────────────────────────────────────────\n    val_loss, val_metrics, _ = evaluate_model(full_model, val_loader, criterion)\n\n    print(f\"  Train loss: {train_loss:.4f} | Val loss: {val_loss:.4f}\")\n    print(f\"  AUC={val_metrics['auc']:.4f} | Recall={val_metrics['recall_pneumonia']:.4f} | \"\n          f\"Spec={val_metrics['specificity']:.4f} | F1={val_metrics['f1']:.4f}\")\n\n    history.append({\n        \"abs_epoch\": abs_epoch, \"local_epoch\": local_ep + 1,\n        \"train_loss\": train_loss, \"val_loss\": val_loss,\n        \"lr_head\": current_lr_head, \"lr_enc\": current_lr_enc,\n        **{f\"val_{k}\": v for k, v in val_metrics.items()}\n    })\n\n    # ── Last checkpoint + cleanup ────────────────────────────\n    last_ckpt = CKPT_DIR / f\"{MODEL_NAME}_last_epoch_{abs_epoch}.pth\"\n    torch.save({\"model_state_dict\": full_model.state_dict(),\n                \"abs_epoch\": abs_epoch, \"val_auc\": val_metrics[\"auc\"],\n                \"best_auc\": best_auc}, last_ckpt)\n    cleanup_checkpoints(MODEL_NAME, keep_last=3)\n    print(f\"  Disk: {disk_used_gb():.2f} GB\")\n\n    # ── Best checkpoint ───────────────────────────────────────\n    current_auc = val_metrics[\"auc\"]\n    if current_auc > best_auc:\n        best_auc       = current_auc\n        best_epoch_abs = abs_epoch\n        patience_counter = 0\n        torch.save({\"model_state_dict\": full_model.state_dict(),\n                    \"abs_epoch\": abs_epoch, \"best_auc\": best_auc,\n                    \"encoder_checkpoint\": str(IJEPA_ENCODER_CKPT)},\n                   best_ckpt_path)\n        print(f\"  ✓ New best AUC {best_auc:.4f} saved.\")\n    else:\n        patience_counter += 1\n        print(f\"  No improvement. Patience: {patience_counter}/{PATIENCE}\")\n        if patience_counter >= PATIENCE:\n            print(\"  Early stopping.\")\n            break\n\n    # ── Save history CSV ─────────────────────────────────────\n    pd.DataFrame(history).to_csv(LOG_DIR / f\"{MODEL_NAME}_history.csv\", index=False)\n\nprint(f\"\\nBest AUC: {best_auc:.4f} at abs epoch {best_epoch_abs}\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"c309683b-0300-4b4f-a1c0-cf30ad0e89d0","cell_type":"code","source":"# ============================================================\n# CELL 13: EVALUATE TRÊN TEST SET\n# ============================================================\n# Load best checkpoint của lần resume này\ntry:\n    ckpt = torch.load(best_ckpt_path, map_location=DEVICE, weights_only=False)\nexcept TypeError:\n    ckpt = torch.load(best_ckpt_path, map_location=DEVICE)\n\nfull_model.load_state_dict(ckpt[\"model_state_dict\"])\nprint(f\"Loaded best checkpoint: abs_epoch={ckpt['abs_epoch']} | AUC={ckpt['best_auc']:.4f}\")\n\ntest_loss, test_metrics, test_pred_df = evaluate_model(\n    full_model, test_loader, criterion, threshold=0.5\n)\n\nprint(\"\\n=== TEST SET RESULTS (threshold=0.5) ===\")\nfor k, v in test_metrics.items():\n    print(f\"  {k:25s}: {v}\")\n\n# So sánh với lần chạy cũ\nprint(\"\\n=== SO SÁNH VỚI RUN CŨ ===\")\nold_metrics = {\"auc\": 0.8137, \"recall_pneumonia\": 0.7461, \"specificity\": 0.7440, \"f1\": 0.5682}\nfor k, old_v in old_metrics.items():\n    new_v = test_metrics.get(k, 0)\n    delta = new_v - old_v\n    print(f\"  {k:25s}: {old_v:.4f} → {new_v:.4f}  ({delta:+.4f})\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"d20bd141-c368-4f8b-8680-26f78c1512ce","cell_type":"code","source":"# ============================================================\n# CELL 14: LƯU KẾT QUẢ\n# ============================================================\nimport matplotlib.pyplot as plt\n\n# ── Save metrics CSV ─────────────────────────────────────────\nresults_df = pd.DataFrame([{\n    \"model\": \"I-JEPA Full FT Resume (ep21-35)\",\n    \"test_loss\": test_loss,\n    **test_metrics\n}])\nresults_df.to_csv(OUTPUT_DIR / \"full_ft_resume_test_metrics.csv\", index=False)\nprint(\"Saved: full_ft_resume_test_metrics.csv\")\n\n# ── Plot training curve ───────────────────────────────────────\nhistory_df = pd.DataFrame(history)\n\nfig, axes = plt.subplots(1, 2, figsize=(12, 4))\naxes[0].plot(history_df[\"abs_epoch\"], history_df[\"train_loss\"], label=\"Train\", marker=\"o\", markersize=3)\naxes[0].plot(history_df[\"abs_epoch\"], history_df[\"val_loss\"],   label=\"Val\",   marker=\"o\", markersize=3)\naxes[0].set_xlabel(\"Epoch (absolute)\"); axes[0].set_ylabel(\"Loss\")\naxes[0].set_title(\"Train vs Val Loss (Resume)\"); axes[0].legend(); axes[0].grid(True)\n\naxes[1].plot(history_df[\"abs_epoch\"], history_df[\"val_auc\"], color=\"steelblue\", marker=\"o\", markersize=3)\naxes[1].axhline(y=0.8137, color=\"orange\", linestyle=\"--\", label=\"Cũ: 0.8137\")\naxes[1].axhline(y=0.886,  color=\"red\",    linestyle=\":\",  label=\"ResNet50: 0.886\")\naxes[1].set_xlabel(\"Epoch (absolute)\"); axes[1].set_ylabel(\"Val AUC\")\naxes[1].set_title(\"Val AUC — Resume\"); axes[1].legend(); axes[1].grid(True)\n\nplt.tight_layout()\nplt.savefig(FIG_DIR / \"full_ft_resume_curves.png\", dpi=150)\nplt.show()\nprint(\"Saved figure.\")\n","metadata":{},"outputs":[],"execution_count":null},{"id":"8b93d795-1344-4fe7-a27d-448509a1037e","cell_type":"code","source":"# ============================================================\n# CELL 15: SUMMARY + FILES OUTPUT\n# ============================================================\nprint(\"=== Files output ===\")\nfor p in sorted(OUTPUT_DIR.rglob(\"*\")):\n    if p.is_file():\n        print(f\"  {p.relative_to(OUTPUT_DIR)}  ({p.stat().st_size/1e6:.1f} MB)\")\n\nprint(\"\\n=== Checkpoint để dùng cho Notebook 05 ===\")\nprint(f\"  {best_ckpt_path.name}\")\nprint(\"  → Add output notebook này vào Notebook 05 như 1 Kaggle dataset\")\n","metadata":{},"outputs":[],"execution_count":null}]}