{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.10.0"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":10338,"databundleVersionId":862042,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":317003019,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":318826856,"isSourceIdPinned":false},{"sourceType":"kernelVersion","sourceId":319949819,"isSourceIdPinned":false}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Notebook 04 v2 — I-JEPA Fine-tuning trên RSNA\n\n**v2** — Layer-wise LR Decay (LLRD), config mới Phase 2 & 3, không resume Phase 3.","metadata":{}},{"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, 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')\n_AMP_DEV   = '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_v2')\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},{"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},{"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},{"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},{"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},{"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},{"cell_type":"code","source":"# ============================================================\n# CELL 7: TÌM I-JEPA ENCODER CHECKPOINT TỪ NOTEBOOK 03\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    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\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},{"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},{"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},{"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},{"cell_type":"code","source":"# ============================================================\n# CELL 11: DATALOADER + POS_WEIGHT\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\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},{"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},{"cell_type":"code","source":"# ============================================================\n# CELL 13: CLASSIFIER MODEL\n# v2: dropout giảm từ 0.2 → 0.1 (phối hợp với weight_decay)\n# ============================================================\n\ndef build_vit_small_encoder(pretrained=False):\n    return timm.create_model(\n        'vit_small_patch16_224',\n        pretrained=pretrained,\n        num_classes=0\n    )\n\nclass IJEPAClassifier(nn.Module):\n    def __init__(self, encoder, embed_dim=384, dropout=0.1):  # v2: dropout=0.1\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        return self.classifier(self.encoder(x)).squeeze(-1)\n\n_test = build_vit_small_encoder()\nprint('ViT-Small embed_dim:', _test.num_features)\nprint('ViT-Small num_blocks:', len(_test.blocks))  # phải là 12\ndel _test","metadata":{},"outputs":[],"execution_count":null},{"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', '?'))","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 15: FACTORY + FREEZE/UNFREEZE HELPERS\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    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    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    return IJEPAClassifier(encoder, embed_dim=encoder.num_features, dropout=0.1)\n\n\ndef freeze_encoder(model):\n    for p in model.encoder.parameters(): p.requires_grad = False\n    for p in model.classifier.parameters(): p.requires_grad = True\n    print('Encoder FROZEN — chỉ train classifier head.')\n\n\ndef unfreeze_last_n_blocks(model, n=4):\n    \"\"\"Unfreeze n blocks cuối + norm + classifier. Các block còn lại frozen.\"\"\"\n    for p in model.encoder.parameters(): p.requires_grad = False\n    if hasattr(model.encoder, 'blocks'):\n        for block in model.encoder.blocks[-n:]:\n            for p in block.parameters(): p.requires_grad = True\n    if hasattr(model.encoder, 'norm'):\n        for p in model.encoder.norm.parameters(): p.requires_grad = True\n    for p in model.classifier.parameters(): 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(): 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\n\nprint('Factory + freeze helpers ready.')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 16: LAYER-WISE LR DECAY (LLRD)\n# ============================================================\n# weight_decay nằm trong từng group (không truyền global vào AdamW)\n# seen_ids dedup tránh stem_params trùng với block params\n# unfreeze_from_block: truyền vào để Cell 24 không crash,\n#   nhưng logic thực không cần vì filter requires_grad đã đủ\n# ============================================================\n\nNUM_LAYERS  = 12\nLAYER_DECAY = 0.65\n\ndef get_llrd_param_groups(model, lr_head, lr_enc_base,\n                           layer_decay=0.65, num_layers=12,\n                           weight_decay=0.05,\n                           unfreeze_from_block=None):  # giữ param để tương thích Cell 24/27\n    \"\"\"\n    Tạo param groups với Layer-wise LR Decay.\n    unfreeze_from_block: không dùng trong logic (requires_grad filter đủ),\n                         chỉ giữ signature để Cell 24 (partial FT) không crash.\n    \"\"\"\n    param_groups = []\n    seen_ids = set()  # dedup: tránh cùng 1 param xuất hiện trong 2 groups\n\n    # 1. Classifier head — LR cao nhất\n    head_params = [p for p in model.classifier.parameters() if p.requires_grad]\n    if head_params:\n        param_groups.append({\n            'params': head_params, 'lr': lr_head,\n            'weight_decay': weight_decay, 'name': 'classifier_head'\n        })\n        seen_ids.update(id(p) for p in head_params)\n\n    # 2. Encoder norm\n    if hasattr(model.encoder, 'norm'):\n        norm_params = [p for p in model.encoder.norm.parameters()\n                       if p.requires_grad and id(p) not in seen_ids]\n        if norm_params:\n            param_groups.append({\n                'params': norm_params,\n                'lr': lr_enc_base * (layer_decay ** 0),\n                'weight_decay': weight_decay, 'name': 'encoder_norm'\n            })\n            seen_ids.update(id(p) for p in norm_params)\n\n    # 3. Transformer blocks — LLRD\n    #    Block frozen (requires_grad=False) tự động bị bỏ qua\n    #    → Phase 2 partial FT: chỉ block 8-11 được add (do unfreeze_last_n_blocks)\n    #    → Phase 3 full FT: tất cả 12 blocks được add\n    if hasattr(model.encoder, 'blocks'):\n        for block_idx, block in enumerate(model.encoder.blocks):\n            block_params = [p for p in block.parameters()\n                            if p.requires_grad and id(p) not in seen_ids]\n            if not block_params:\n                continue\n            exponent = (num_layers - 1) - block_idx\n            lr_block = lr_enc_base * (layer_decay ** exponent)\n            param_groups.append({\n                'params': block_params, 'lr': lr_block,\n                'weight_decay': weight_decay,\n                'name': f'encoder_block_{block_idx:02d}'\n            })\n            seen_ids.update(id(p) for p in block_params)\n\n    # 4. Stem (patch_embed, cls_token, pos_embed) — LR thấp nhất\n    stem_params = [\n        p for name, p in model.encoder.named_parameters()\n        if p.requires_grad\n        and id(p) not in seen_ids\n        and any(k in name for k in ['patch_embed', 'cls_token', 'pos_embed'])\n    ]\n    if stem_params:\n        param_groups.append({\n            'params': stem_params,\n            'lr': lr_enc_base * (layer_decay ** num_layers),\n            'weight_decay': weight_decay, 'name': 'encoder_stem'\n        })\n\n    return param_groups\n\n\ndef print_llrd_summary(param_groups):\n    print(f'{\"Group\":<25} {\"Params\":>10} {\"LR\":>12}')\n    print('-' * 50)\n    for g in param_groups:\n        n = sum(p.numel() for p in g['params'])\n        print(f\"{g.get('name','?'):<25} {n:>10,} {g['lr']:>12.2e}\")\n    print('-' * 50)\n\n\nprint('LLRD helper ready.')\nprint(f'LAYER_DECAY={LAYER_DECAY}, NUM_LAYERS={NUM_LAYERS}')\nprint(f'Phase 3 block 0 LR (base=1e-6) ≈ {1e-6 * LAYER_DECAY**11:.2e}')\n\n# Demo kiểm tra Phase 3 (full FT — tất cả blocks)\n_demo = create_ijepa_classifier()\nunfreeze_all(_demo)\n_demo_groups = get_llrd_param_groups(\n    _demo, lr_head=3e-5, lr_enc_base=1e-6,\n    layer_decay=LAYER_DECAY, num_layers=NUM_LAYERS\n)\nprint_llrd_summary(_demo_groups)\ndel _demo\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 17: 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    cm = confusion_matrix(y_true, y_pred, labels=[0, 1])\n    tn, fp, fn, tp = cm.ravel()\n\n    return {\n        'auc':              auc,\n        'accuracy':         accuracy_score(y_true, y_pred),\n        'precision':        precision_score(y_true, y_pred, zero_division=0),\n        'recall_pneumonia': recall_score(y_true, y_pred, zero_division=0),\n        'specificity':      tn / (tn + fp + 1e-8),\n        'f1':               f1_score(y_true, y_pred, zero_division=0),\n        'tn': int(tn), 'fp': int(fp), 'fn': int(fn), 'tp': int(tp),\n        'threshold': threshold\n    }","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 18: 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        with torch.amp.autocast(_AMP_DEV, enabled=torch.cuda.is_available()):\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},{"cell_type":"code","source":"# ============================================================\n# CELL 19: HÀM TRAIN ONE EPOCH\n# v2: grad_clip=0.5 (giảm từ 1.0), torch.amp thay cuda.amp\n# ============================================================\n\ndef train_one_epoch(model, loader, optimizer, criterion, scaler,\n                    accumulation_steps=2, max_grad_norm=0.5):  # v2: 0.5\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.amp.autocast(_AMP_DEV, 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    # Flush gradient batch cuối nếu tổng steps không chia hết cho accum\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    return running_loss / total_samples","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 20: LR SCHEDULER (COSINE + WARMUP) + CHECKPOINT HELPERS\n# ============================================================\n\ndef make_lr_scheduler(optimizer, num_epochs, warmup_ratio=0.1):\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\n\ndef cleanup_last_checkpoints(model_name, keep_last=3):\n    def epoch_num(p):\n        try:    return int(p.stem.split('_epoch_')[-1])\n        except: return -1\n    all_ckpts = sorted(CKPT_DIR.glob(f'{model_name}_last_epoch_*.pth'), key=epoch_num)\n    for old in all_ckpts[:-keep_last]:\n        old.unlink(missing_ok=True)\n\n\ndef disk_used_gb():\n    import shutil\n    _, used, _ = shutil.disk_usage('/kaggle/working')\n    return used / 1e9\n\nprint('Scheduler + checkpoint helpers ready.')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 21: HÀM TRAIN MODEL HOÀN CHỈNH\n# v2: nhận param_groups thay vì lr/encoder_lr riêng lẻ\n# ============================================================\n\ndef train_model(\n    model,\n    model_name,\n    train_loader,\n    val_loader,\n    param_groups,       # v2: truyền param_groups từ LLRD thay vì lr đơn\n    num_epochs   = 10,\n    weight_decay = 0.05,   # đã embed trong param_groups — tham số này chỉ để document\n    accum_steps  = 2,\n    patience     = 8,   # v2: tăng từ 5/6 → 8\n    warmup_ratio = 0.15,\n    grad_clip    = 0.5, # v2: giảm từ 1.0 → 0.5\n    pos_weight   = None,\n    config_extra = None,\n):\n    model = model.to(DEVICE)\n    criterion = nn.BCEWithLogitsLoss(\n        pos_weight=pos_weight.to(DEVICE) if pos_weight is not None else None\n    )\n\n    # v2: optimizer nhận param_groups với LR đã tính sẵn\n    # LỖI 1 fix: weight_decay đã có trong group — không truyền global\n    optimizer = torch.optim.AdamW(param_groups)\n    scheduler = make_lr_scheduler(optimizer, num_epochs, warmup_ratio)\n    scaler    = torch.amp.GradScaler(_AMP_DEV, 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    # Print config\n    head_lr = next((g['lr'] for g in param_groups if 'head' in g.get('name','')), '?')\n    enc_lrs = [g['lr'] for g in param_groups if 'block' in g.get('name','')]\n    pw_str  = f\"{pos_weight.item():.4f}\" if pos_weight is not None else 'None'\n    print(f\"{'='*60}\")\n    print(f\"  {model_name}\")\n    print(f\"  epochs={num_epochs} | warmup={warmup_ratio*100:.0f}% | patience={patience}\")\n    print(f\"  lr_head={head_lr:.2e} | lr_enc range=[{min(enc_lrs):.2e}, {max(enc_lrs):.2e}]\" if enc_lrs else f\"  lr_head={head_lr}\")\n    print(f\"  grad_clip={grad_clip} | pos_weight={pw_str} | wd={weight_decay}\")\n    print(f\"{'='*60}\")\n\n    for epoch in range(1, num_epochs + 1):\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']  # head LR làm đại diện\n        print(f'\\n[{model_name}] Epoch {epoch}/{num_epochs} | head_LR={current_lr:.2e}')\n\n        train_loss = train_one_epoch(\n            model, train_loader, optimizer, criterion,\n            scaler, accum_steps, max_grad_norm=grad_clip\n        )\n        scheduler.step()\n\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} | '\n              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            'head_lr': current_lr,\n            **{f'val_{k}': v for k, v in val_metrics.items()}\n        })\n\n        # Last checkpoint\n        last_ckpt = CKPT_DIR / f'{model_name}_last_epoch_{epoch}.pth'\n        save_dict = {\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        }\n        if config_extra: save_dict['config'] = config_extra\n        torch.save(save_dict, last_ckpt)\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            save_dict['best_auc'] = best_auc\n            torch.save(save_dict, 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        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},{"cell_type":"markdown","source":"## Phase 1 — Linear Probing\nGiữ nguyên config cũ. Nếu đã có `ijepa_linear_probe_best.pth`, bỏ qua để tiết kiệm ~1h GPU.","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# CELL 22: LINEAR PROBING\n# Nếu đã có checkpoint từ NB04 cũ → load lại, không train lại\n# ============================================================\n\nLINEAR_CKPT_EXISTING = find_file('ijepa_linear_probe_best.pth')\nprint('Existing linear probe ckpt:', LINEAR_CKPT_EXISTING)\n\nif LINEAR_CKPT_EXISTING:\n    print('Found existing checkpoint — bỏ qua training, dùng lại.')\n    linear_ckpt_path = LINEAR_CKPT_EXISTING\n    linear_history   = pd.read_csv(find_file('ijepa_linear_probe_train_history.csv')) \\\n                       if find_file('ijepa_linear_probe_train_history.csv') else pd.DataFrame()\nelse:\n    print('Không tìm thấy — chạy Linear Probe mới...')\n    linear_model = create_ijepa_classifier()\n    freeze_encoder(linear_model)\n    count_trainable_params(linear_model)\n\n    # Linear probe: chỉ có head params, dùng flat LR\n    _lp_groups = [{\n        'params': [p for p in linear_model.classifier.parameters() if p.requires_grad],\n        'lr': 1e-3,\n        'weight_decay': 1e-4,\n        'name': 'classifier_head'\n    }]\n\n    linear_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        param_groups = _lp_groups,\n        num_epochs   = 15,\n        weight_decay = 1e-4,\n        accum_steps  = 2,\n        patience     = 5,\n        warmup_ratio = 0.10,\n        grad_clip    = 0.5,\n        pos_weight   = POS_WEIGHT,\n        config_extra = {'phase': 'linear_probe', 'lr': 1e-3}\n    )\n\n    del linear_model; gc.collect()\n    if torch.cuda.is_available(): torch.cuda.empty_cache()\n\ndisplay(linear_history.tail(5) if len(linear_history) else pd.DataFrame())","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 23: 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))\n\ndel linear_best; gc.collect()\nif torch.cuda.is_available(): torch.cuda.empty_cache()\nprint('GPU cleared.')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 2 — Partial Fine-tuning (v2)\nUnfreeze 4 blocks cuối + LLRD. LR_HEAD=3e-5, LR_ENC_BASE=5e-6, EPOCHS=25, WARMUP=15%","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# CELL 24: PARTIAL FINE-TUNING — v2\n# Thay đổi: LR_HEAD 1e-4→3e-5, LR_ENC 1e-5→5e-6+LLRD\n#           EPOCHS 20→25, WARMUP 10%→15%, PATIENCE 6→8\n# ============================================================\n\nN_UNFREEZE_BLOCKS = 4   # unfreeze block 8, 9, 10, 11\n\n# Phase 2 config\nP2_LR_HEAD     = 3e-5\nP2_LR_ENC_BASE = 5e-6\nP2_EPOCHS      = 25\nP2_WARMUP      = 0.15\n\npartial_model = create_ijepa_classifier()\nunfreeze_last_n_blocks(partial_model, n=N_UNFREEZE_BLOCKS)\ncount_trainable_params(partial_model)\n\n# LLRD chỉ trên 4 blocks được unfreeze (block 8-11)\n# unfreeze_from_block=8 → bỏ qua block 0-7 trong LLRD\np2_param_groups = get_llrd_param_groups(\n    partial_model,\n    lr_head          = P2_LR_HEAD,\n    lr_enc_base      = P2_LR_ENC_BASE,\n    layer_decay      = LAYER_DECAY,\n    num_layers       = NUM_LAYERS,\n    weight_decay     = 0.05,\n    unfreeze_from_block = NUM_LAYERS - N_UNFREEZE_BLOCKS  # =8\n)\n\nprint('\\nPhase 2 LLRD param groups:')\nprint_llrd_summary(p2_param_groups)\n\npartial_ckpt_path, partial_history = train_model(\n    model        = partial_model,\n    model_name   = 'ijepa_partial_finetune_v2',\n    train_loader = train_loader,\n    val_loader   = val_loader,\n    param_groups = p2_param_groups,\n    num_epochs   = P2_EPOCHS,\n    weight_decay = 0.05,\n    accum_steps  = 2,\n    patience     = 8,\n    warmup_ratio = P2_WARMUP,\n    grad_clip    = 0.5,\n    pos_weight   = POS_WEIGHT,\n    config_extra = {\n        'phase': 'partial_ft_v2',\n        'lr_head': P2_LR_HEAD, 'lr_enc_base': P2_LR_ENC_BASE,\n        'n_unfreeze_blocks': N_UNFREEZE_BLOCKS,\n        'layer_decay': LAYER_DECAY, 'warmup': P2_WARMUP\n    }\n)\n\ndisplay(partial_history[['epoch','train_loss','val_loss','val_auc','val_recall_pneumonia']])","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 25: 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 FT v2 — 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_v2_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))\n\ndel partial_model, partial_best; gc.collect()\nif torch.cuda.is_available(): torch.cuda.empty_cache()\nprint('GPU cleared.')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Phase 3 — Full Fine-tuning (v2)\nKHÔNG resume từ Phase 2. Train thẳng từ encoder pretrain.\nLR_HEAD=3e-5, LR_ENC_BASE=1e-6, LAYER_DECAY=0.65, EPOCHS=35, WARMUP=20%","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# CELL 26: CONFIG FULL FINE-TUNING v2\n# ============================================================\n\nRUN_FULL_FINETUNE = True  # Đổi False nếu muốn bỏ qua\n\n# Phase 3 config\nP3_LR_HEAD     = 3e-5\nP3_LR_ENC_BASE = 1e-6   # base LR cho encoder (block 11)\n                         # block 0 sẽ nhận: 1e-6 × 0.65^11 ≈ 3.2e-9\nP3_EPOCHS      = 35\nP3_WARMUP      = 0.20   # 7 epochs warmup — encoder cần ổn định trước\n\nprint(f'Full FT config v2:')\nprint(f'  LR_HEAD={P3_LR_HEAD:.2e}, LR_ENC_BASE={P3_LR_ENC_BASE:.2e}')\nprint(f'  LAYER_DECAY={LAYER_DECAY} → block 0 LR ≈ {P3_LR_ENC_BASE * LAYER_DECAY**11:.2e}')\nprint(f'  EPOCHS={P3_EPOCHS}, WARMUP={P3_WARMUP*100:.0f}% ({int(P3_EPOCHS*P3_WARMUP)} epochs)')\nprint(f'  GRAD_CLIP=0.5, PATIENCE=8, WD=0.05')\nprint(f'  NOTE: KHÔNG resume từ Phase 2 — load thẳng từ encoder pretrain')\nprint(f'RUN_FULL_FINETUNE: {RUN_FULL_FINETUNE}')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 27: FULL FINE-TUNING — v2\n# Thay đổi lớn nhất: LLRD toàn bộ 12 blocks, LR_ENC_BASE=1e-6\n#   → block 0-2: LR ≈ 2e-8 (gần frozen, bảo tồn CXR features)\n#   → block 10-11: LR ≈ 8e-7 (adapt sang RSNA)\n# ============================================================\n\nif RUN_FULL_FINETUNE:\n    # Load fresh từ encoder pretrain — KHÔNG dùng partial_ckpt_path\n    full_model = create_ijepa_classifier()\n    unfreeze_all(full_model)\n    count_trainable_params(full_model)\n\n    # LLRD toàn bộ 12 blocks\n    p3_param_groups = get_llrd_param_groups(\n        full_model,\n        lr_head              = P3_LR_HEAD,\n        lr_enc_base          = P3_LR_ENC_BASE,\n        layer_decay          = LAYER_DECAY,\n        num_layers           = NUM_LAYERS,\n        weight_decay         = 0.05,\n        unfreeze_from_block  = None  # toàn bộ blocks\n    )\n\n    print('\\nPhase 3 LLRD param groups:')\n    print_llrd_summary(p3_param_groups)\n\n    full_ckpt_path, full_history = train_model(\n        model        = full_model,\n        model_name   = 'ijepa_full_finetune_v2',\n        train_loader = train_loader,\n        val_loader   = val_loader,\n        param_groups = p3_param_groups,\n        num_epochs   = P3_EPOCHS,\n        weight_decay = 0.05,\n        accum_steps  = 2,\n        patience     = 8,\n        warmup_ratio = P3_WARMUP,\n        grad_clip    = 0.5,\n        pos_weight   = POS_WEIGHT,\n        config_extra = {\n            'phase': 'full_ft_v2',\n            'lr_head': P3_LR_HEAD, 'lr_enc_base': P3_LR_ENC_BASE,\n            'layer_decay': LAYER_DECAY, 'warmup': P3_WARMUP,\n            'source': 'encoder_pretrain_direct'  # không resume từ phase 2\n        }\n    )\n\n    display(full_history[['epoch','train_loss','val_loss','val_auc','val_recall_pneumonia']])\nelse:\n    full_ckpt_path = None\n    full_history   = None\n    print('Skipped.')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 28: EVALUATE FULL FT TRÊN TEST SET\n# ============================================================\n\nif RUN_FULL_FINETUNE and full_ckpt_path:\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 FT v2 — 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_v2_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    ))\n\n    del full_model, full_best; gc.collect()\n    if torch.cuda.is_available(): torch.cuda.empty_cache()\n    print('GPU cleared.')\nelse:\n    full_test_loss, full_test_metrics, full_pred_df = None, None, None\n    print('Skipped.')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Tổng hợp kết quả","metadata":{}},{"cell_type":"code","source":"# ============================================================\n# CELL 29: BẢNG TỔNG HỢP METRICS I-JEPA\n# ============================================================\n\n# Guard: nếu evaluate cells bị skip do timeout → dict đầy đủ keys (không rỗng)\n# Dùng dict đầy đủ keys thay vì {} để tránh cột auc=NaN khi **unpack\n_EMPTY_METRICS = {\n    'auc': None, 'accuracy': None, 'precision': None,\n    'recall_pneumonia': None, 'specificity': None, 'f1': None,\n    'tn': None, 'fp': None, 'fn': None, 'tp': None, 'threshold': None\n}\n\nlinear_test_loss     = locals().get('linear_test_loss',     None)\nlinear_test_metrics  = locals().get('linear_test_metrics',  None) or _EMPTY_METRICS\npartial_test_loss    = locals().get('partial_test_loss',    None)\npartial_test_metrics = locals().get('partial_test_metrics', None) or _EMPTY_METRICS\nfull_test_loss       = locals().get('full_test_loss',       None)\nfull_test_metrics    = locals().get('full_test_metrics',    None)  # None nếu skip\n\nrows = [\n    {'model': 'I-JEPA Linear Probe',  'test_loss': linear_test_loss,  **linear_test_metrics},\n    {'model': 'I-JEPA Partial FT v2', '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 v2', '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_v2.csv', index=False)\n\ndisplay(ijepa_metrics_df[['model','auc','f1','recall_pneumonia','specificity','precision']])\nprint('Saved: ijepa_finetune_metrics_v2.csv')\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 30: CONFUSION MATRICES + ROC CURVES\n# VĐ 2 fix: guard pred_df — skip plot nếu chưa có (timeout)\n# ============================================================\n\nfrom sklearn.metrics import roc_curve\n\n# Guard pred_df\nlinear_pred_df  = locals().get('linear_pred_df',  None)\npartial_pred_df = locals().get('partial_pred_df',  None)\nfull_pred_df    = locals().get('full_pred_df',     None)\n\ndef plot_cm(pred_df, title, save_path):\n    if pred_df is None:\n        print(f'  [{title}] pred_df chưa có — skip.')\n        return\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    fig, ax = plt.subplots(figsize=(5, 4))\n    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(); plt.savefig(save_path, dpi=150); plt.show()\n\nplot_cm(linear_pred_df,  'Linear Probe',  FIG_DIR/'cm_linear_probe.png')\nplot_cm(partial_pred_df, 'Partial FT v2', FIG_DIR/'cm_partial_ft_v2.png')\nif RUN_FULL_FINETUNE and full_pred_df is not None:\n    plot_cm(full_pred_df, 'Full FT v2',   FIG_DIR/'cm_full_ft_v2.png')\n\n# ROC — chỉ vẽ các model đã có pred_df\nplt.figure(figsize=(6, 5))\nroc_candidates = [\n    (linear_pred_df,  'Linear Probe'),\n    (partial_pred_df, 'Partial FT v2'),\n] + ([(full_pred_df, 'Full FT v2')] if RUN_FULL_FINETUNE and full_pred_df is not None else [])\n\nhas_roc = False\nfor pred_df, label in roc_candidates:\n    if pred_df is None:\n        print(f'  [{label}] pred_df chưa có — skip ROC.')\n        continue\n    yt = pred_df['label'].astype(int).values\n    yp = pred_df['prob_pneumonia'].values\n    fpr, tpr, _ = roc_curve(yt, yp)\n    auc = roc_auc_score(yt, yp)\n    plt.plot(fpr, tpr, label=f'{label}  AUC={auc:.4f}')\n    has_roc = True\n\nif has_roc:\n    plt.plot([0,1],[0,1],'k--',label='Random')\n    plt.xlabel('FPR'); plt.ylabel('TPR')\n    plt.title('I-JEPA ROC Curves v2 — RSNA Test Set')\n    plt.legend(); plt.tight_layout()\n    plt.savefig(FIG_DIR/'ijepa_roc_curves_v2.png', dpi=150); plt.show()\n    print('Saved: ijepa_roc_curves_v2.png')\nelse:\n    print('Không có pred_df nào — bỏ qua ROC plot.')\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 31: PLOT TRAINING CURVES\n# LỖI 4 fix: axes[0]=loss only, axes[1]=AUC only (khác đơn vị)\n# ============================================================\n\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\n\nphase_histories = [('Partial FT v2', partial_history, 'blue')]\nif RUN_FULL_FINETUNE and full_history is not None:\n    phase_histories.append(('Full FT v2', full_history, 'red'))\n\nfor name, hist, color in phase_histories:\n    # axes[0]: train loss + val loss — cùng đơn vị\n    axes[0].plot(hist['epoch'], hist['train_loss'], color=color,\n                 linestyle='--', alpha=0.7, label=f'{name} Train Loss')\n    axes[0].plot(hist['epoch'], hist['val_loss'],   color=color,\n                 label=f'{name} Val Loss')\n    # axes[1]: val AUC — đơn vị riêng\n    axes[1].plot(hist['epoch'], hist['val_auc'], color=color,\n                 marker='o', markersize=3, label=f'{name}')\n\naxes[0].set_title('Train / Val Loss')\naxes[0].set_xlabel('Epoch'); axes[0].set_ylabel('Loss')\naxes[0].legend(); axes[0].grid(True, alpha=0.3)\n\naxes[1].set_title('Val AUC per Epoch')\naxes[1].set_xlabel('Epoch'); axes[1].set_ylabel('AUC')\naxes[1].legend(); axes[1].grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.savefig(FIG_DIR / 'training_curves_v2.png', dpi=150); plt.show()\nprint('Saved: training_curves_v2.png')\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 32: 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_v2.csv', index=False)\n    display(compare_df[['model','auc','f1','recall_pneumonia','specificity']])\n    print('Saved: all_models_compare_v2.csv')\nelse:\n    print('baseline_metrics.csv không tìm thấy — bỏ qua.')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 33: LƯU CONFIG v2\n# ============================================================\n\nconfig_v2 = {\n    'version': 'v2',\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    'layer_decay': LAYER_DECAY,\n    'num_layers': NUM_LAYERS,\n    'dropout': 0.1,\n    'grad_clip': 0.5,\n    'linear': {\n        'epochs': 15, 'lr': 1e-3, 'patience': 5, 'warmup': 0.10\n    },\n    'partial': {\n        'epochs': P2_EPOCHS, 'lr_head': P2_LR_HEAD, 'lr_enc_base': P2_LR_ENC_BASE,\n        'n_blocks': N_UNFREEZE_BLOCKS, 'patience': 8,\n        'warmup': P2_WARMUP, 'layer_decay': LAYER_DECAY,\n        'llrd': True\n    },\n    'full': {\n        'epochs': P3_EPOCHS, 'lr_head': P3_LR_HEAD, 'lr_enc_base': P3_LR_ENC_BASE,\n        'patience': 8, 'warmup': P3_WARMUP, 'layer_decay': LAYER_DECAY,\n        'weight_decay': 0.05, 'llrd': True,\n        'source': 'encoder_pretrain_direct'\n    },\n    'loss': 'BCEWithLogitsLoss+pos_weight',\n    'optimizer': 'AdamW', 'scheduler': 'cosine+warmup',\n}\n\nwith open(OUTPUT_DIR / 'ijepa_finetune_config_v2.json', 'w') as f:\n    json.dump(config_v2, f, indent=2)\nprint('Config v2 saved.')","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================\n# CELL 34: LIỆT KÊ VÀ 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_v2'\nimport shutil\nshutil.make_archive(zip_base, 'zip', OUTPUT_DIR)\nprint('Zip created:', zip_base + '.zip')","metadata":{},"outputs":[],"execution_count":null}]}