{"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,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":316980196,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":317003019,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":318826856,"isSourceIdPinned":false}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"d1890c8a-f0d5-443a-b7f0-0e41428d671a","cell_type":"markdown","source":"# Notebook 04 — I-JEPA Fine-tuning trên RSNA\n\n**Fixed version** — pos_weight, LR scheduler, tăng epochs, unfreeze 4 blocks, checkpoint cleanup.","metadata":{}},{"id":"eb3eacf7-ebaf-4e72-ba83-61abb548abb7","cell_type":"code","source":"# ============================================================\n# CELL 1: IMPORT THƯ VIỆN VÀ CẤU HÌNH CHUNG\n# ============================================================\n\nimport os, gc, json, math, random, shutil, zipfile, time\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\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as T\n\nfrom sklearn.metrics import (\n    roc_auc_score, accuracy_score, precision_score,\n    recall_score, f1_score, confusion_matrix, classification_report\n)\nimport matplotlib.pyplot as plt\n\n# ── Seed ──────────────────────────────────────────────────\nSEED = 42\n\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.benchmark = True\n\nset_seed(SEED)\n\n# ── Device ────────────────────────────────────────────────\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    print(\"CUDA:\", torch.version.cuda)\n\n# ── Output dirs ───────────────────────────────────────────\nOUTPUT_DIR = Path(\"/kaggle/working/notebook04_ijepa_finetune\")\nCKPT_DIR   = OUTPUT_DIR / \"checkpoints\"\nPRED_DIR   = OUTPUT_DIR / \"predictions\"\nLOG_DIR    = OUTPUT_DIR / \"logs\"\nFIG_DIR    = OUTPUT_DIR / \"figures\"\n\nfor d in [OUTPUT_DIR, CKPT_DIR, PRED_DIR, LOG_DIR, FIG_DIR]:\n    d.mkdir(parents=True, exist_ok=True)\n\nprint(\"Output dir:\", OUTPUT_DIR)","metadata":{},"outputs":[],"execution_count":null},{"id":"891159bc-0b31-4857-b4df-ed75da807f34","cell_type":"code","source":"# ============================================================\n# CELL 2: CÀI / IMPORT TIMM VÀ PYDICOM\n# ============================================================\n\ntry:\n    import timm\nexcept ImportError:\n    !pip install -q timm\n    import timm\n\ntry:\n    import pydicom\nexcept ImportError:\n    !pip install -q pydicom\n    import pydicom\n\nprint(\"timm:\", timm.__version__)\nprint(\"pydicom: OK\")","metadata":{},"outputs":[],"execution_count":null},{"id":"7d5325e9-5674-4201-b533-f68ba9645a0e","cell_type":"code","source":"# ============================================================\n# CELL 3: TIME GUARD\n# ============================================================\n\nSESSION_SAFE_HOURS   = 11.5\nSESSION_SAFE_SECONDS = SESSION_SAFE_HOURS * 3600\nNOTEBOOK_START_TIME  = time.time()\n\nprint(f\"Time guard: {SESSION_SAFE_HOURS}h\")","metadata":{},"outputs":[],"execution_count":null},{"id":"565c25fb-926c-4823-8c0c-5e5388910153","cell_type":"code","source":"# ============================================================\n# CELL 4: TÌM METADATA RSNA TỪ NOTEBOOK 01\n# ============================================================\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\nRSNA_TRAIN_CSV = find_file(\"rsna_train.csv\")\nRSNA_VAL_CSV   = find_file(\"rsna_val.csv\")\nRSNA_TEST_CSV  = find_file(\"rsna_test.csv\")\n\nprint(\"RSNA_TRAIN_CSV:\", RSNA_TRAIN_CSV)\nprint(\"RSNA_VAL_CSV  :\", RSNA_VAL_CSV)\nprint(\"RSNA_TEST_CSV :\", RSNA_TEST_CSV)\n\nassert RSNA_TRAIN_CSV, \"Không tìm thấy rsna_train.csv\"\nassert RSNA_VAL_CSV,   \"Không tìm thấy rsna_val.csv\"\nassert RSNA_TEST_CSV,  \"Không tìm thấy rsna_test.csv\"\n\nprint(\"Metadata RSNA OK.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"bd7a1f2a-493d-4e1b-9241-bed8d00f0311","cell_type":"code","source":"# ============================================================\n# CELL 5: ĐỌC METADATA RSNA\n# ============================================================\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(\"Train:\", train_df.shape)\nprint(\"Val  :\", val_df.shape)\nprint(\"Test :\", test_df.shape)\n\nprint(\"Train label distribution:\")\nprint(train_df[\"label\"].value_counts())\n\ndisplay(train_df.head(3))","metadata":{},"outputs":[],"execution_count":null},{"id":"043c471e-5eff-4cff-a180-52242f7b6116","cell_type":"code","source":"# ============================================================\n# CELL 6: KIỂM TRA / SỬA ĐƯỜNG DẪN ẢNH RSNA\n# ============================================================\n\ndef count_missing(df):\n    return (~df[\"image_path\"].apply(lambda x: Path(x).exists())).sum()\n\ndef find_dir(root, dirname):\n    matches = [p for p in Path(root).rglob(dirname) if p.is_dir()]\n    return matches[0] if matches else None\n\nprint(\"Missing trước fix - train:\", count_missing(train_df),\n      \"val:\", count_missing(val_df), \"test:\", count_missing(test_df))\n\nif count_missing(train_df) + count_missing(val_df) + count_missing(test_df) > 0:\n    RSNA_IMG_DIR = find_dir(INPUT_ROOT, \"stage_2_train_images\")\n    assert RSNA_IMG_DIR, \"Không tìm thấy stage_2_train_images — hãy add RSNA dataset.\"\n    print(\"RSNA_IMG_DIR:\", RSNA_IMG_DIR)\n\n    for df in [train_df, val_df, test_df]:\n        df[\"image_path\"] = df[\"patientId\"].apply(\n            lambda x: str(RSNA_IMG_DIR / f\"{x}.dcm\")\n        )\n\nprint(\"Missing sau fix - train:\", count_missing(train_df),\n      \"val:\", count_missing(val_df), \"test:\", count_missing(test_df))\n\nassert count_missing(train_df) == 0\nassert count_missing(val_df)   == 0\nassert count_missing(test_df)  == 0\n\nprint(\"Đường dẫn ảnh RSNA OK.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"9fa8b47b-6baf-4531-ab2a-74bfed47e288","cell_type":"code","source":"# ============================================================\n# CELL 7: TÌM I-JEPA ENCODER CHECKPOINT TỪ NOTEBOOK 03\n# ============================================================\n# Tự động tìm — không cần nhập tay đường dẫn.\n# Chỉ cần đã add output Notebook 03 làm Input Dataset.\n# ============================================================\n\ndef find_ijepa_encoder():\n    patterns = [\n        \"ijepa_vit_small_nih_50k_best_encoder.pth\",\n        \"ijepa_vit_small_nih_30k_best_encoder.pth\",\n        \"ijepa_vit_small_nih_*_best_encoder.pth\",\n        \"ijepa_vit_small_nih_*_encoder_epoch_*.pth\",\n    ]\n    candidates = []\n    for pat in patterns:\n        candidates += list(WORKING_ROOT.rglob(pat))\n        candidates += list(INPUT_ROOT.rglob(pat))\n\n    # deduplicate\n    seen, unique = set(), []\n    for p in candidates:\n        if str(p) not in seen:\n            unique.append(p); seen.add(str(p))\n    return unique\n\ncandidates = find_ijepa_encoder()\nprint(\"Encoder candidates found:\")\nfor i, p in enumerate(candidates):\n    print(f\"  [{i}] {p}\")\n\nassert candidates, (\n    \"Không tìm thấy I-JEPA encoder checkpoint!\"\n    \"Hãy add output Notebook 03 làm Input Dataset trước khi chạy.\"\n)\n\n# Ưu tiên: 50k best > 30k best > epoch mới nhất\nIJEPA_ENCODER_CKPT = None\nfor p in candidates:\n    if \"50k_best_encoder\" in p.name:\n        IJEPA_ENCODER_CKPT = p; break\nif IJEPA_ENCODER_CKPT is None:\n    for p in candidates:\n        if \"best_encoder\" in p.name:\n            IJEPA_ENCODER_CKPT = p; break\nif IJEPA_ENCODER_CKPT is None:\n    IJEPA_ENCODER_CKPT = candidates[0]\n\nprint(\"Selected:\", IJEPA_ENCODER_CKPT)","metadata":{},"outputs":[],"execution_count":null},{"id":"10a590a0-7423-4be0-8bd1-8b9f3de91a66","cell_type":"code","source":"# ============================================================\n# CELL 8: HÀM ĐỌC ẢNH DICOM\n# ============================================================\n\ndef read_dicom_as_pil(path):\n    dicom = pydicom.dcmread(path)\n    img   = dicom.pixel_array.astype(np.float32)\n    img   = img - np.min(img)\n    mx    = np.max(img)\n    if mx > 0:\n        img = img / mx\n    img = (img * 255).astype(np.uint8)\n    return Image.fromarray(img).convert(\"RGB\")","metadata":{},"outputs":[],"execution_count":null},{"id":"d3a46bcd-1565-452f-9a49-64753681b1b9","cell_type":"code","source":"# ============================================================\n# CELL 9: TRANSFORMS\n# ============================================================\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.RandomRotation(degrees=7),\n    T.RandomHorizontalFlip(p=0.5),\n    T.ColorJitter(brightness=0.10, contrast=0.10),\n    T.ToTensor(),\n    T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\neval_transform = T.Compose([\n    T.Resize((IMG_SIZE, IMG_SIZE)),\n    T.ToTensor(),\n    T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n])\n\nprint(\"Transforms ready.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"be45e24f-40c3-4faa-a774-75eb9b00e0b5","cell_type":"code","source":"# ============================================================\n# CELL 10: DATASET CLASS\n# ============================================================\n\nclass RSNADataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.df        = dataframe.reset_index(drop=True).copy()\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row   = self.df.iloc[idx]\n        image = read_dicom_as_pil(row[\"image_path\"])\n        label = float(row[\"label\"])\n        if self.transform:\n            image = self.transform(image)\n        return {\n            \"image\":     image,\n            \"label\":     torch.tensor(label, dtype=torch.float32),\n            \"patientId\": row[\"patientId\"]\n        }","metadata":{},"outputs":[],"execution_count":null},{"id":"b6778238-dd79-413d-9516-97f225862ea2","cell_type":"code","source":"# ============================================================\n# CELL 11: DATALOADER + POS_WEIGHT (class imbalance fix)\n# ============================================================\n\nBATCH_SIZE  = 16\nNUM_WORKERS = 2\n\ntrain_dataset = RSNADataset(train_df, transform=train_transform)\nval_dataset   = RSNADataset(val_df,   transform=eval_transform)\ntest_dataset  = RSNADataset(test_df,  transform=eval_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE,\n                          shuffle=True,  num_workers=NUM_WORKERS, pin_memory=True)\nval_loader   = DataLoader(val_dataset,   batch_size=BATCH_SIZE,\n                          shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)\ntest_loader  = DataLoader(test_dataset,  batch_size=BATCH_SIZE,\n                          shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)\n\n# ── Tính pos_weight để xử lý class imbalance 3.44:1 ──────\n# pos_weight = n_negative / n_positive\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).to(DEVICE)\n\nprint(f\"Train neg={n_neg:,}, pos={n_pos:,}, ratio={n_neg/n_pos:.3f}\")\nprint(f\"POS_WEIGHT = {POS_WEIGHT.item():.4f}\")\nprint(f\"Train batches: {len(train_loader)} | Val: {len(val_loader)} | Test: {len(test_loader)}\")","metadata":{},"outputs":[],"execution_count":null},{"id":"50d0c8e0-7d58-418e-9f0a-72a666373b0d","cell_type":"code","source":"# ============================================================\n# CELL 12: TEST MỘT BATCH\n# ============================================================\n\nbatch = next(iter(train_loader))\nprint(\"Image shape :\", batch[\"image\"].shape)\nprint(\"Label shape :\", batch[\"label\"].shape)\nprint(\"PatientIds  :\", batch[\"patientId\"][:3])","metadata":{},"outputs":[],"execution_count":null},{"id":"12986ced-5155-4621-8433-85f5a972aee4","cell_type":"code","source":"# ============================================================\n# CELL 13: XÂY DỰNG VIT-SMALL ENCODER\n# ============================================================\n\ndef build_vit_small_encoder(pretrained=False):\n    model = timm.create_model(\n        \"vit_small_patch16_224\",\n        pretrained=pretrained,\n        num_classes=0          # bỏ classification head, trả về [B, 384]\n    )\n    return model\n\n_test_enc = build_vit_small_encoder()\nprint(\"ViT-Small embed_dim:\", _test_enc.num_features)\ndel _test_enc","metadata":{},"outputs":[],"execution_count":null},{"id":"5cf72b3c-4e0d-4508-b740-b5b60b391274","cell_type":"code","source":"# ============================================================\n# CELL 14: LOAD I-JEPA ENCODER CHECKPOINT\n# ============================================================\n\nijepa_encoder = build_vit_small_encoder(pretrained=False)\n\ntry:\n    enc_ckpt = torch.load(IJEPA_ENCODER_CKPT, map_location=\"cpu\", weights_only=False)\nexcept TypeError:\n    enc_ckpt = torch.load(IJEPA_ENCODER_CKPT, map_location=\"cpu\")\n\nprint(\"Checkpoint keys:\", list(enc_ckpt.keys()))\n\nif \"encoder_state_dict\" in enc_ckpt:\n    state_dict = enc_ckpt[\"encoder_state_dict\"]\nelif \"student_encoder_state_dict\" in enc_ckpt:\n    state_dict = enc_ckpt[\"student_encoder_state_dict\"]\nelse:\n    state_dict = enc_ckpt\n\nmissing, unexpected = ijepa_encoder.load_state_dict(state_dict, strict=False)\nprint(\"Missing keys    :\", missing)\nprint(\"Unexpected keys :\", unexpected)\nprint(\"Pretrain epoch  :\", enc_ckpt.get(\"epoch\", \"?\"))\nprint(\"Pretrain loss   :\", enc_ckpt.get(\"avg_loss\", \"?\"))\n\ncfg = enc_ckpt.get(\"config\", {})\nprint(\"Pretrain config :\", {k: v for k, v in cfg.items() if k in [\"mask_ratio\",\"lr_max\",\"epochs_total\"]})","metadata":{},"outputs":[],"execution_count":null},{"id":"af31976f-02da-49cc-8d7f-4ee05342b4fe","cell_type":"code","source":"# ============================================================\n# CELL 15: CLASSIFIER MODEL\n# ============================================================\n\nclass IJEPAClassifier(nn.Module):\n    \"\"\"I-JEPA encoder + binary classification head.\"\"\"\n\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]\n        logits   = self.classifier(features).squeeze(1)\n        return logits\n\nEMBED_DIM    = ijepa_encoder.num_features   # 384 cho ViT-Small\nijepa_model  = IJEPAClassifier(ijepa_encoder, embed_dim=EMBED_DIM, dropout=0.2)\nprint(\"Classifier ready. embed_dim:\", EMBED_DIM)","metadata":{},"outputs":[],"execution_count":null},{"id":"e910d1e9-6c97-4c7b-87df-004a02a925de","cell_type":"code","source":"# ============================================================\n# CELL 16: FREEZE / UNFREEZE HELPERS\n# ============================================================\n\ndef freeze_encoder(model):\n    for p in model.encoder.parameters():\n        p.requires_grad = False\n    for p in model.classifier.parameters():\n        p.requires_grad = True\n    print(\"Encoder FROZEN — chỉ train classifier head.\")\n\n\ndef unfreeze_last_n_blocks(model, n=4):\n    for p in model.encoder.parameters():\n        p.requires_grad = False\n    if hasattr(model.encoder, \"blocks\"):\n        for block in model.encoder.blocks[-n:]:\n            for p in block.parameters():\n                p.requires_grad = True\n    if hasattr(model.encoder, \"norm\"):\n        for p in model.encoder.norm.parameters():\n            p.requires_grad = True\n    for p in model.classifier.parameters():\n        p.requires_grad = True\n    print(f\"Partial FT — unfreeze {n} ViT blocks + norm + classifier.\")\n\n\ndef unfreeze_all(model):\n    for p in model.parameters():\n        p.requires_grad = True\n    print(\"Full FT — toàn bộ params trainable.\")\n\n\ndef count_trainable_params(model):\n    trainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    total     = sum(p.numel() for p in model.parameters())\n    print(f\"Trainable: {trainable:,} / {total:,} ({trainable/total*100:.2f}%)\")\n    return trainable, total","metadata":{},"outputs":[],"execution_count":null},{"id":"d9d4d33b-47db-4747-8929-04bb009270e3","cell_type":"code","source":"# ============================================================\n# CELL 17: HÀM FACTORY — TẠO MODEL TỪNG LẦN\n# ============================================================\n\ndef create_ijepa_classifier():\n    \"\"\"Load fresh encoder từ I-JEPA checkpoint, trả về IJEPAClassifier.\"\"\"\n    encoder = build_vit_small_encoder(pretrained=False)\n\n    try:\n        ckpt = torch.load(IJEPA_ENCODER_CKPT, map_location=\"cpu\", weights_only=False)\n    except TypeError:\n        ckpt = torch.load(IJEPA_ENCODER_CKPT, map_location=\"cpu\")\n\n    sd = (ckpt.get(\"encoder_state_dict\") or\n          ckpt.get(\"student_encoder_state_dict\") or ckpt)\n    encoder.load_state_dict(sd, strict=False)\n\n    return IJEPAClassifier(encoder, embed_dim=encoder.num_features, dropout=0.2)\n\nprint(\"Factory function create_ijepa_classifier() ready.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"51db99cc-06bf-401c-a73c-c641a8c880ac","cell_type":"code","source":"# ============================================================\n# CELL 18: HÀM TÍNH METRICS\n# ============================================================\n\ndef compute_binary_metrics(y_true, y_prob, threshold=0.5):\n    y_true = np.array(y_true).astype(int)\n    y_prob = np.array(y_prob)\n    y_pred = (y_prob >= threshold).astype(int)\n\n    try:\n        auc = roc_auc_score(y_true, y_prob)\n    except ValueError:\n        auc = np.nan\n\n    acc       = accuracy_score(y_true, y_pred)\n    precision = precision_score(y_true, y_pred, zero_division=0)\n    recall    = recall_score(y_true, y_pred, zero_division=0)\n    f1        = f1_score(y_true, y_pred, zero_division=0)\n    cm        = confusion_matrix(y_true, y_pred, labels=[0, 1])\n    tn, fp, fn, tp = cm.ravel()\n    specificity = tn / (tn + fp + 1e-8)\n\n    return {\n        \"auc\": auc, \"accuracy\": acc,\n        \"precision\": precision, \"recall_pneumonia\": recall,\n        \"specificity\": specificity, \"f1\": f1,\n        \"tn\": int(tn), \"fp\": int(fp), \"fn\": int(fn), \"tp\": int(tp),\n        \"threshold\": threshold\n    }","metadata":{},"outputs":[],"execution_count":null},{"id":"7137f4a5-3913-42d0-9a7f-6c8a4a1bba4b","cell_type":"code","source":"# ============================================================\n# CELL 19: HÀM EVALUATE\n# ============================================================\n\n@torch.no_grad()\ndef evaluate_model(model, loader, criterion=None):\n    model.eval()\n    all_labels, all_probs, all_pids = [], [], []\n    total_loss, total_samples = 0.0, 0\n\n    for batch in tqdm(loader, desc=\"Eval\", leave=False):\n        images = batch[\"image\"].to(DEVICE, non_blocking=True)\n        labels = batch[\"label\"].to(DEVICE, non_blocking=True)\n\n        logits = model(images)\n\n        if criterion is not None:\n            total_loss += criterion(logits, labels).item() * images.size(0)\n\n        probs = torch.sigmoid(logits)\n        all_labels.extend(labels.cpu().numpy().tolist())\n        all_probs.extend(probs.cpu().numpy().tolist())\n        all_pids.extend(batch[\"patientId\"])\n        total_samples += images.size(0)\n\n    avg_loss = total_loss / total_samples if criterion is not None else None\n    metrics  = compute_binary_metrics(all_labels, all_probs)\n\n    pred_df = pd.DataFrame({\n        \"patientId\":     all_pids,\n        \"label\":         all_labels,\n        \"prob_pneumonia\": all_probs\n    })\n    pred_df[\"pred\"] = (pred_df[\"prob_pneumonia\"] >= 0.5).astype(int)\n\n    return avg_loss, metrics, pred_df","metadata":{},"outputs":[],"execution_count":null},{"id":"afe3b698-0118-4dd2-851e-4f22558edd2d","cell_type":"code","source":"# ============================================================\n# CELL 20: HÀM TRAIN ONE EPOCH\n# ============================================================\n\ndef train_one_epoch(model, loader, optimizer, criterion, scaler,\n                    accumulation_steps=2, max_grad_norm=1.0):\n    model.train()\n    running_loss, total_samples = 0.0, 0\n    optimizer.zero_grad(set_to_none=True)\n\n    for step, batch in enumerate(tqdm(loader, desc=\"Train\", leave=False)):\n        images = batch[\"image\"].to(DEVICE, non_blocking=True)\n        labels = batch[\"label\"].to(DEVICE, non_blocking=True)\n\n        with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):\n            logits = model(images)\n            loss   = criterion(logits, labels) / accumulation_steps\n\n        scaler.scale(loss).backward()\n\n        if (step + 1) % accumulation_steps == 0:\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(\n                [p for p in model.parameters() if p.requires_grad],\n                max_grad_norm\n            )\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad(set_to_none=True)\n\n        running_loss  += loss.item() * accumulation_steps * images.size(0)\n        total_samples += images.size(0)\n\n    return running_loss / total_samples","metadata":{},"outputs":[],"execution_count":null},{"id":"a86c8f99-71aa-4ac7-aec2-ff97671bc4fc","cell_type":"code","source":"# ============================================================\n# CELL 21: LR SCHEDULER (COSINE + WARMUP)\n# ============================================================\n\ndef make_lr_scheduler(optimizer, num_epochs, warmup_ratio=0.1):\n    \"\"\"\n    Linear warmup trong warmup_ratio*num_epochs epoch đầu,\n    sau đó cosine decay về 0.\n    \"\"\"\n    warmup_epochs = max(1, int(num_epochs * warmup_ratio))\n\n    def lr_lambda(current_epoch):\n        if current_epoch < warmup_epochs:\n            return float(current_epoch + 1) / float(warmup_epochs)\n        progress = (current_epoch - warmup_epochs) / max(1, num_epochs - warmup_epochs)\n        return 0.5 * (1.0 + math.cos(math.pi * progress))\n\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)\n\nprint(\"make_lr_scheduler() ready.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"3b0fa54e-241c-4027-89ed-6b8d97ba3fa0","cell_type":"code","source":"# ============================================================\n# CELL 22: CHECKPOINT CLEANUP HELPER\n# ============================================================\n\ndef cleanup_last_checkpoints(model_name, keep_last=3):\n    \"\"\"Giữ chỉ keep_last epoch cuối, xóa cũ hơn để tiết kiệm disk.\"\"\"\n    def epoch_num(p):\n        try:\n            return int(p.stem.split(\"_epoch_\")[-1])\n        except ValueError:\n            return -1\n\n    all_ckpts = sorted(\n        CKPT_DIR.glob(f\"{model_name}_last_epoch_*.pth\"),\n        key=epoch_num\n    )\n    for old in all_ckpts[:-keep_last]:\n        old.unlink(missing_ok=True)\n\ndef disk_used_gb():\n    total, used, _ = shutil.disk_usage(\"/kaggle/working\")\n    return used / 1e9\n\nprint(\"Checkpoint cleanup helper ready.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"69125fa0-e47d-4627-aece-0e3f0c3ead51","cell_type":"code","source":"# ============================================================\n# CELL 23: HÀM TRAIN MODEL HOÀN CHỈNH\n# ============================================================\n\ndef train_model(\n    model,\n    model_name,\n    train_loader,\n    val_loader,\n    num_epochs    = 10,\n    lr            = 1e-4,\n    encoder_lr    = None,\n    weight_decay  = 1e-4,\n    accum_steps   = 2,\n    patience      = 5,\n    warmup_ratio  = 0.10,\n    pos_weight    = None,\n):\n    model = model.to(DEVICE)\n\n    # ── Loss với pos_weight (xử lý class imbalance 3.44:1) ──\n    criterion = nn.BCEWithLogitsLoss(\n        pos_weight=pos_weight.to(DEVICE) if pos_weight is not None else None\n    )\n\n    # ── Optimizer (differential LR nếu có encoder_lr) ──────\n    if encoder_lr is not None:\n        optimizer = torch.optim.AdamW([\n            {\"params\": [p for p in model.encoder.parameters()    if p.requires_grad], \"lr\": encoder_lr},\n            {\"params\": [p for p in model.classifier.parameters() if p.requires_grad], \"lr\": lr},\n        ], weight_decay=weight_decay)\n    else:\n        optimizer = torch.optim.AdamW(\n            [p for p in model.parameters() if p.requires_grad],\n            lr=lr, weight_decay=weight_decay\n        )\n\n    # ── LR Scheduler (cosine + warmup) ─────────────────────\n    scheduler = make_lr_scheduler(optimizer, num_epochs, warmup_ratio)\n\n    scaler = torch.cuda.amp.GradScaler(enabled=torch.cuda.is_available())\n\n    best_auc, best_epoch = -1, -1\n    patience_counter     = 0\n    history              = []\n    best_ckpt_path       = CKPT_DIR / f\"{model_name}_best.pth\"\n\n    pw_str = f\"{pos_weight.item():.4f}\" if pos_weight is not None else \"None\"\n    print(f\"{'='*58}\")\n    print(f\"  {model_name}\")\n    print(f\"  epochs={num_epochs} | lr_head={lr} | lr_enc={encoder_lr}\")\n    print(f\"  warmup={warmup_ratio*100:.0f}% epochs | patience={patience}\")\n    print(f\"  pos_weight={pw_str} | weight_decay={weight_decay}\")\n    print(f\"{'='*58}\")\n\n    for epoch in range(1, num_epochs + 1):\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 {epoch}.\")\n            break\n\n        current_lr = optimizer.param_groups[0][\"lr\"]\n        print(f\"\\n[{model_name}] Epoch {epoch}/{num_epochs} | LR={current_lr:.2e}\")\n\n        # ── Train ──────────────────────────────────────────\n        train_loss = train_one_epoch(\n            model, train_loader, optimizer, criterion,\n            scaler, accum_steps\n        )\n\n        # ── Scheduler step (sau mỗi epoch) ─────────────────\n        scheduler.step()\n\n        # ── Validate ───────────────────────────────────────\n        val_loss, val_metrics, _ = evaluate_model(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} | F1={val_metrics['f1']:.4f} | \"              f\"Recall={val_metrics['recall_pneumonia']:.4f} | Spec={val_metrics['specificity']:.4f}\")\n\n        history.append({\n            \"model_name\": model_name, \"epoch\": epoch,\n            \"train_loss\": train_loss, \"val_loss\": val_loss,\n            \"lr\": current_lr,\n            **{f\"val_{k}\": v for k, v in val_metrics.items()}\n        })\n\n        # ── Lưu last checkpoint ────────────────────────────\n        last_ckpt = CKPT_DIR / f\"{model_name}_last_epoch_{epoch}.pth\"\n        torch.save({\n            \"model_state_dict\": model.state_dict(),\n            \"model_name\": model_name, \"epoch\": epoch,\n            \"val_auc\": val_metrics[\"auc\"],\n            \"encoder_checkpoint\": str(IJEPA_ENCODER_CKPT),\n            \"config\": {\n                \"lr\": lr, \"encoder_lr\": encoder_lr,\n                \"weight_decay\": weight_decay, \"num_epochs\": num_epochs,\n                \"accum_steps\": accum_steps, \"pos_weight\": pw_str,\n                \"seed\": SEED, \"img_size\": IMG_SIZE\n            }\n        }, last_ckpt)\n\n        cleanup_last_checkpoints(model_name, keep_last=3)\n        print(f\"  Disk used: {disk_used_gb():.2f} GB\")\n\n        # ── Best checkpoint ────────────────────────────────\n        current_auc = val_metrics[\"auc\"]\n        if current_auc > best_auc:\n            best_auc, best_epoch = current_auc, epoch\n            patience_counter     = 0\n            torch.save({\n                \"model_state_dict\": model.state_dict(),\n                \"model_name\": model_name, \"epoch\": epoch,\n                \"best_auc\": best_auc,\n                \"encoder_checkpoint\": str(IJEPA_ENCODER_CKPT),\n                \"config\": {\n                    \"lr\": lr, \"encoder_lr\": encoder_lr,\n                    \"weight_decay\": weight_decay, \"num_epochs\": num_epochs,\n                    \"accum_steps\": accum_steps, \"pos_weight\": pw_str,\n                    \"seed\": SEED, \"img_size\": IMG_SIZE\n                }\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(f\"  Early stopping.\")\n                break\n\n        # ── Save history ───────────────────────────────────\n        pd.DataFrame(history).to_csv(\n            LOG_DIR / f\"{model_name}_train_history.csv\", index=False\n        )\n\n    print(f\"\\nBest {model_name} AUC: {best_auc:.4f} at epoch {best_epoch}\")\n    return best_ckpt_path, pd.DataFrame(history)","metadata":{},"outputs":[],"execution_count":null},{"id":"7f0c34ae-5ff0-48a4-942b-830bd611bcd8","cell_type":"markdown","source":"## Phase 1 — Linear Probing\nĐóng băng toàn bộ encoder, chỉ train classifier head.","metadata":{}},{"id":"2e44d7af-670d-437e-bae0-42be12de2521","cell_type":"code","source":"# ============================================================\n# CELL 24: LINEAR PROBING\n# ============================================================\n\nlinear_model = create_ijepa_classifier()\nfreeze_encoder(linear_model)\ncount_trainable_params(linear_model)\n\nlinear_ckpt_path, linear_history = train_model(\n    model        = linear_model,\n    model_name   = \"ijepa_linear_probe\",\n    train_loader = train_loader,\n    val_loader   = val_loader,\n    num_epochs   = 15,      # ← tăng từ 5 → 15\n    lr           = 1e-3,\n    encoder_lr   = None,\n    weight_decay = 1e-4,\n    accum_steps  = 2,\n    patience     = 5,       # ← tăng từ 3 → 5\n    warmup_ratio = 0.10,\n    pos_weight   = POS_WEIGHT,  # ← class imbalance fix\n)\n\ndisplay(linear_history)","metadata":{},"outputs":[],"execution_count":null},{"id":"45880fe5-cade-421d-937f-42079a291c5f","cell_type":"code","source":"# ============================================================\n# CELL 25: EVALUATE LINEAR PROBE TRÊN TEST SET\n# ============================================================\n\nlinear_best = create_ijepa_classifier()\nfreeze_encoder(linear_best)\n\ntry:\n    ckpt = torch.load(linear_ckpt_path, map_location=DEVICE, weights_only=False)\nexcept TypeError:\n    ckpt = torch.load(linear_ckpt_path, map_location=DEVICE)\n\nlinear_best.load_state_dict(ckpt[\"model_state_dict\"])\nlinear_best = linear_best.to(DEVICE)\n\ncriterion_eval = nn.BCEWithLogitsLoss(pos_weight=POS_WEIGHT.to(DEVICE))\n\nlinear_test_loss, linear_test_metrics, linear_pred_df = evaluate_model(\n    linear_best, test_loader, criterion_eval\n)\n\nprint(\"── I-JEPA Linear Probe — Test Results ──\")\nfor k, v in linear_test_metrics.items():\n    print(f\"  {k}: {v}\")\n\nlinear_pred_df.to_csv(PRED_DIR / \"ijepa_linear_probe_predictions.csv\", index=False)\nprint(\"\", classification_report(\n    linear_pred_df[\"label\"].astype(int),\n    linear_pred_df[\"pred\"].astype(int),\n    target_names=[\"Non-pneumonia\", \"Pneumonia\"], zero_division=0\n))","metadata":{},"outputs":[],"execution_count":null},{"id":"0f75115a-36ca-4f2a-ac9e-545f0dd6b770","cell_type":"code","source":"# ============================================================\n# CELL 26: DỌN GPU SAU LINEAR PROBE\n# ============================================================\n\ndel linear_model, linear_best\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\nprint(\"GPU cleared.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"2eb3c0aa-1d34-4138-87db-5a7f004f99d9","cell_type":"markdown","source":"## Phase 2 — Partial Fine-tuning\nUnfreeze 4 blocks cuối + norm + classifier head.","metadata":{}},{"id":"933d8548-e699-4bcc-9001-d58bfaa8258e","cell_type":"code","source":"# ============================================================\n# CELL 27: PARTIAL FINE-TUNING\n# ============================================================\n\npartial_model     = create_ijepa_classifier()\nN_UNFREEZE_BLOCKS = 4    # ← tăng từ 2 → 4\n\nunfreeze_last_n_blocks(partial_model, n=N_UNFREEZE_BLOCKS)\ncount_trainable_params(partial_model)\n\npartial_ckpt_path, partial_history = train_model(\n    model        = partial_model,\n    model_name   = \"ijepa_partial_finetune\",\n    train_loader = train_loader,\n    val_loader   = val_loader,\n    num_epochs   = 20,      # ← tăng từ 10 → 20\n    lr           = 1e-4,\n    encoder_lr   = 1e-5,\n    weight_decay = 1e-4,\n    accum_steps  = 2,\n    patience     = 6,       # ← tăng từ 4 → 6\n    warmup_ratio = 0.10,\n    pos_weight   = POS_WEIGHT,  # ← class imbalance fix\n)\n\ndisplay(partial_history)","metadata":{},"outputs":[],"execution_count":null},{"id":"79b9a8fc-88a7-4c7d-8bc7-5d6e29e968dd","cell_type":"code","source":"# ============================================================\n# CELL 28: EVALUATE PARTIAL FT TRÊN TEST SET\n# ============================================================\n\npartial_best = create_ijepa_classifier()\nunfreeze_last_n_blocks(partial_best, n=N_UNFREEZE_BLOCKS)\n\ntry:\n    ckpt = torch.load(partial_ckpt_path, map_location=DEVICE, weights_only=False)\nexcept TypeError:\n    ckpt = torch.load(partial_ckpt_path, map_location=DEVICE)\n\npartial_best.load_state_dict(ckpt[\"model_state_dict\"])\npartial_best = partial_best.to(DEVICE)\n\npartial_test_loss, partial_test_metrics, partial_pred_df = evaluate_model(\n    partial_best, test_loader, criterion_eval\n)\n\nprint(\"── I-JEPA Partial Fine-tune — Test Results ──\")\nfor k, v in partial_test_metrics.items():\n    print(f\"  {k}: {v}\")\n\npartial_pred_df.to_csv(PRED_DIR / \"ijepa_partial_finetune_predictions.csv\", index=False)\nprint(\"\", classification_report(\n    partial_pred_df[\"label\"].astype(int),\n    partial_pred_df[\"pred\"].astype(int),\n    target_names=[\"Non-pneumonia\", \"Pneumonia\"], zero_division=0\n))","metadata":{},"outputs":[],"execution_count":null},{"id":"9feded1d-2d8c-49c1-8835-38af82e04278","cell_type":"code","source":"# ============================================================\n# CELL 29: DỌN GPU SAU PARTIAL FT\n# ============================================================\n\ndel partial_model, partial_best\ngc.collect()\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\nprint(\"GPU cleared.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"841e0e9d-1dac-4982-bfcd-eada7ff28906","cell_type":"markdown","source":"## Phase 3 — Full Fine-tuning\nUnfreeze toàn bộ encoder.","metadata":{}},{"id":"b9198133-a2a3-4ae6-93a2-d87df8fc7a50","cell_type":"code","source":"# ============================================================\n# CELL 30: CONFIG FULL FINE-TUNING\n# ============================================================\n\nRUN_FULL_FINETUNE = True   # Đổi False nếu muốn bỏ qua\nprint(\"RUN_FULL_FINETUNE:\", RUN_FULL_FINETUNE)","metadata":{},"outputs":[],"execution_count":null},{"id":"94108c08-7dcc-44b8-abd3-a2040980bd06","cell_type":"code","source":"# ============================================================\n# CELL 31: FULL FINE-TUNING\n# ============================================================\n\nif RUN_FULL_FINETUNE:\n    full_model = create_ijepa_classifier()\n    unfreeze_all(full_model)\n    count_trainable_params(full_model)\n\n    full_ckpt_path, full_history = train_model(\n        model        = full_model,\n        model_name   = \"ijepa_full_finetune\",\n        train_loader = train_loader,\n        val_loader   = val_loader,\n        num_epochs   = 20,      # ← tăng từ 5 → 20\n        lr           = 5e-5,\n        encoder_lr   = 5e-6,\n        weight_decay = 0.05,    # ← tăng từ 1e-4 → 0.05 cho full FT\n        accum_steps  = 2,\n        patience     = 6,       # ← tăng từ 3 → 6\n        warmup_ratio = 0.15,    # ← tăng 0.10 → 0.15 cho full FT (an toàn hơn)\n        pos_weight   = POS_WEIGHT,  # ← class imbalance fix\n    )\n    display(full_history)\nelse:\n    full_ckpt_path   = None\n    full_history     = None\n    print(\"Skipped.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"0ab9398d-eefc-412a-8e95-f92cbc228c15","cell_type":"code","source":"# ============================================================\n# CELL 32: EVALUATE FULL FT TRÊN TEST SET\n# ============================================================\n\nif RUN_FULL_FINETUNE:\n    full_best = create_ijepa_classifier()\n    unfreeze_all(full_best)\n\n    try:\n        ckpt = torch.load(full_ckpt_path, map_location=DEVICE, weights_only=False)\n    except TypeError:\n        ckpt = torch.load(full_ckpt_path, map_location=DEVICE)\n\n    full_best.load_state_dict(ckpt[\"model_state_dict\"])\n    full_best = full_best.to(DEVICE)\n\n    full_test_loss, full_test_metrics, full_pred_df = evaluate_model(\n        full_best, test_loader, criterion_eval\n    )\n\n    print(\"── I-JEPA Full Fine-tune — Test Results ──\")\n    for k, v in full_test_metrics.items():\n        print(f\"  {k}: {v}\")\n\n    full_pred_df.to_csv(PRED_DIR / \"ijepa_full_finetune_predictions.csv\", index=False)\n    print(\"\", classification_report(\n        full_pred_df[\"label\"].astype(int),\n        full_pred_df[\"pred\"].astype(int),\n        target_names=[\"Non-pneumonia\", \"Pneumonia\"], zero_division=0\n    ))\nelse:\n    full_test_loss, full_test_metrics, full_pred_df = None, None, None\n    print(\"Skipped.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"449f29d6-64bf-424a-b54f-c3b5f7a8fce3","cell_type":"markdown","source":"## Tổng hợp kết quả","metadata":{}},{"id":"add72b22-2694-4636-8698-d3fe5f8d9211","cell_type":"code","source":"# ============================================================\n# CELL 33: BẢNG TỔNG HỢP METRICS I-JEPA\n# ============================================================\n\nrows = [\n    {\"model\": \"I-JEPA Linear Probe\",    \"test_loss\": linear_test_loss,  **linear_test_metrics},\n    {\"model\": \"I-JEPA Partial FT (4blk)\",\"test_loss\": partial_test_loss, **partial_test_metrics},\n]\nif RUN_FULL_FINETUNE and full_test_metrics:\n    rows.append({\"model\": \"I-JEPA Full FT\", \"test_loss\": full_test_loss, **full_test_metrics})\n\nijepa_metrics_df = pd.DataFrame(rows)\nijepa_metrics_df.to_csv(OUTPUT_DIR / \"ijepa_finetune_metrics.csv\", index=False)\n\ndisplay(ijepa_metrics_df[[\"model\",\"auc\",\"f1\",\"recall_pneumonia\",\"specificity\",\"precision\"]])\nprint(\"Saved: ijepa_finetune_metrics.csv\")","metadata":{},"outputs":[],"execution_count":null},{"id":"958dee0a-37d0-4a88-a2bd-7bd49c8c5363","cell_type":"code","source":"# ============================================================\n# CELL 34: CONFUSION MATRICES\n# ============================================================\n\ndef plot_cm(pred_df, title, save_path):\n    y_true = pred_df[\"label\"].astype(int).values\n    y_pred = pred_df[\"pred\"].astype(int).values\n    cm     = confusion_matrix(y_true, y_pred, labels=[0, 1])\n\n    fig, ax = plt.subplots(figsize=(5, 4))\n    im = ax.imshow(cm, cmap=\"Blues\")\n    ax.set_xticks([0,1]); ax.set_yticks([0,1])\n    ax.set_xticklabels([\"Non-pneu\",\"Pneumonia\"], rotation=15)\n    ax.set_yticklabels([\"Non-pneu\",\"Pneumonia\"])\n    ax.set_xlabel(\"Predicted\"); ax.set_ylabel(\"True\")\n    ax.set_title(title)\n    for i in range(2):\n        for j in range(2):\n            ax.text(j, i, cm[i,j], ha=\"center\", va=\"center\",\n                    color=\"white\" if cm[i,j] > cm.max()/2 else \"black\")\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150)\n    plt.show()\n    print(\"Saved:\", save_path)\n\nplot_cm(linear_pred_df,  \"Linear Probe\",  FIG_DIR/\"cm_linear_probe.png\")\nplot_cm(partial_pred_df, \"Partial FT\",    FIG_DIR/\"cm_partial_ft.png\")\nif RUN_FULL_FINETUNE and full_pred_df is not None:\n    plot_cm(full_pred_df, \"Full FT\", FIG_DIR/\"cm_full_ft.png\")","metadata":{},"outputs":[],"execution_count":null},{"id":"30460d77-8d04-4656-a67c-5cf6946b2131","cell_type":"code","source":"# ============================================================\n# CELL 35: ROC CURVES\n# ============================================================\n\nfrom sklearn.metrics import roc_curve\n\ndef roc_data(pred_df):\n    yt = pred_df[\"label\"].astype(int).values\n    yp = pred_df[\"prob_pneumonia\"].values\n    return roc_curve(yt, yp) + (roc_auc_score(yt, yp),)\n\nplt.figure(figsize=(6, 5))\n\nfpr, tpr, _, auc = roc_data(linear_pred_df)\nplt.plot(fpr, tpr, label=f\"Linear Probe   AUC={auc:.4f}\")\n\nfpr, tpr, _, auc = roc_data(partial_pred_df)\nplt.plot(fpr, tpr, label=f\"Partial FT     AUC={auc:.4f}\")\n\nif RUN_FULL_FINETUNE and full_pred_df is not None:\n    fpr, tpr, _, auc = roc_data(full_pred_df)\n    plt.plot(fpr, tpr, label=f\"Full FT        AUC={auc:.4f}\")\n\nplt.plot([0,1],[0,1], \"k--\", label=\"Random\")\nplt.xlabel(\"FPR\"); plt.ylabel(\"TPR\")\nplt.title(\"I-JEPA ROC Curves — RSNA Test Set\")\nplt.legend(); plt.tight_layout()\nroc_path = FIG_DIR / \"ijepa_roc_curves.png\"\nplt.savefig(roc_path, dpi=150); plt.show()\nprint(\"Saved:\", roc_path)","metadata":{},"outputs":[],"execution_count":null},{"id":"6773c38c-a804-47d4-bda6-4f826d8ed3eb","cell_type":"code","source":"# ============================================================\n# CELL 36: SO SÁNH NHANH VỚI BASELINE (NẾU CÓ)\n# ============================================================\n\nbaseline_csv = find_file(\"baseline_metrics.csv\")\n\nif baseline_csv:\n    baseline_df = pd.read_csv(baseline_csv)\n    compare_df  = pd.concat([baseline_df, ijepa_metrics_df], ignore_index=True)\n    compare_df.to_csv(OUTPUT_DIR / \"all_models_compare.csv\", index=False)\n    display(compare_df[[\"model\",\"auc\",\"f1\",\"recall_pneumonia\",\"specificity\"]])\n    print(\"Saved: all_models_compare.csv\")\nelse:\n    print(\"baseline_metrics.csv không tìm thấy — bỏ qua so sánh.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"9215f42e-cc9a-4611-8dcc-425e70762359","cell_type":"code","source":"# ============================================================\n# CELL 37: LƯU CONFIG\n# ============================================================\n\nconfig = {\n    \"seed\": SEED, \"img_size\": IMG_SIZE, \"batch_size\": BATCH_SIZE,\n    \"pos_weight\": float(POS_WEIGHT.item()),\n    \"encoder_checkpoint\": str(IJEPA_ENCODER_CKPT),\n    \"backbone\": \"vit_small_patch16_224\",\n    \"linear\":  {\"epochs\": 15, \"lr\": 1e-3, \"patience\": 5, \"warmup\": 0.10},\n    \"partial\":  {\"epochs\": 20, \"lr_head\": 1e-4, \"lr_enc\": 1e-5,\n                  \"n_blocks\": 4, \"patience\": 6, \"warmup\": 0.10},\n    \"full\":     {\"epochs\": 20, \"lr_head\": 5e-5, \"lr_enc\": 5e-6,\n                  \"weight_decay\": 0.05, \"patience\": 6, \"warmup\": 0.15},\n    \"loss\": \"BCEWithLogitsLoss+pos_weight\",\n    \"optimizer\": \"AdamW\", \"scheduler\": \"cosine+warmup\",\n}\n\nwith open(OUTPUT_DIR / \"ijepa_finetune_config.json\", \"w\") as f:\n    json.dump(config, f, indent=2)\nprint(\"Config saved.\")","metadata":{},"outputs":[],"execution_count":null},{"id":"eeb2be21-e8fe-4ece-b384-de0167944610","cell_type":"code","source":"# ============================================================\n# CELL 38: NÉN OUTPUT\n# ============================================================\n\nprint(\"Files trong output:\")\nfor p in sorted(OUTPUT_DIR.rglob(\"*\")):\n    if p.is_file():\n        size_mb = p.stat().st_size / 1e6\n        print(f\"  {p.relative_to(OUTPUT_DIR)}  ({size_mb:.1f} MB)\")\n\nzip_base = \"/kaggle/working/notebook04_ijepa_finetune\"\nshutil.make_archive(zip_base, \"zip\", OUTPUT_DIR)\nprint(\"Zip created:\", zip_base + \".zip\")","metadata":{},"outputs":[],"execution_count":null}]}