{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# pydicom de doc anh RSNA (dinh dang DICOM), cap nhat transformers/huggingface_hub\n# Cai rieng huggingface_hub truoc (force-reinstall) de tranh loi lech phien ban voi transformers moi\n!pip install -q -U --force-reinstall --no-deps huggingface_hub\n!pip install -q -U transformers accelerate pydicom","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-07-11T12:49:00.809293Z","iopub.execute_input":"2026-07-11T12:49:00.809578Z","iopub.status.idle":"2026-07-11T12:49:22.635031Z","shell.execute_reply.started":"2026-07-11T12:49:00.809532Z","shell.execute_reply":"2026-07-11T12:49:22.634215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from kaggle_secrets import UserSecretsClient\nfrom huggingface_hub import login\n\n# Lấy token từ Secrets của Kaggle\nuser_secrets = UserSecretsClient()\ntry:\n    HF_TOKEN = user_secrets.get_secret('Jasper')\nexcept:\n    HF_TOKEN = user_secrets.get_secret('HF_TOKEN')\n\n# Tiến hành đăng nhập vào hệ thống HuggingFace\nlogin(token=HF_TOKEN)\nprint('Da dang nhap HuggingFace thanh cong.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T12:53:41.219792Z","iopub.execute_input":"2026-07-11T12:53:41.220349Z","iopub.status.idle":"2026-07-11T12:53:41.763995Z","shell.execute_reply.started":"2026-07-11T12:53:41.220318Z","shell.execute_reply":"2026-07-11T12:53:41.763222Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\n\ndef find_root(base, marker_file):\n    for root, dirs, files_ in os.walk(base):\n        if marker_file in files_:\n            return root\n    return None\n\nRSNA_ROOT = find_root('/kaggle/input', 'stage_2_train_labels.csv')\nassert RSNA_ROOT is not None, 'Khong tim thay stage_2_train_labels.csv.'\nprint('RSNA_ROOT =', RSNA_ROOT)\n\nNIH_ROOT = find_root('/kaggle/input', 'Data_Entry_2017.csv')\nassert NIH_ROOT is not None, 'Khong tim thay Data_Entry_2017.csv.'\nprint('NIH_ROOT =', NIH_ROOT)\n\nrsna_img_dir = os.path.join(RSNA_ROOT, 'stage_2_train_images')\nassert os.path.isdir(rsna_img_dir), f'Khong tim thay thu muc anh RSNA tai {rsna_img_dir}'\nprint('RSNA images dir:', rsna_img_dir, '-', len(os.listdir(rsna_img_dir)), 'file .dcm')\n\nprint('Dang lap index anh NIH, co the mat 1-2 phut...')\nnih_image_paths = glob.glob(os.path.join(NIH_ROOT, '**', '*.png'), recursive=True)\nnih_filename_to_path = {os.path.basename(p): p for p in nih_image_paths}\nprint(f'NIH: {len(nih_filename_to_path):,} anh.')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T12:55:16.91588Z","iopub.execute_input":"2026-07-11T12:55:16.916384Z","iopub.status.idle":"2026-07-11T12:58:13.330502Z","shell.execute_reply.started":"2026-07-11T12:55:16.916339Z","shell.execute_reply":"2026-07-11T12:58:13.329584Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"MODEL_NAME = 'facebook/dinov3-vitb16-pretrain-lvd1689m'\nNUM_CLASSES = 2\nCLASS_NAMES = ['NORMAL', 'PNEUMONIA']\n\nNUM_EPOCHS = 16        # Giam xuong 16 de tranh overfitting phan sau va tiet kiem thoi gian\nBATCH_SIZE = 64        # Dung batch lon hon nho AMP, train nhanh hon\nBACKBONE_LR = 2e-6\nHEAD_LR = 1e-4\nWARMUP_EPOCHS = 2\nGRAD_CLIP_NORM = 1.0\nWEIGHT_DECAY = 0.01\nNUM_WORKERS = 4\nUSE_AMP = True         # Kich hoat Mixed Precision de tang toc do tren T4 GPU\nUSE_MULTI_GPU = True   # Dung nn.DataParallel tren ca 2 GPU neu co san\n\nTARGET_TOTAL_IMAGES = 50000   \nTRAIN_RATIO = 0.8\nVAL_RATIO = 0.1\nTEST_RATIO = 0.1              \n\nSEED = 42\nSPIKE_THRESHOLD = 1.5\nSPIKE_PATIENCE = 3            # Early stopping neu AUC khong tang sau 3 epoch lien tiep\n\nEXCLUDE_LIST_PATH = None  \n\nWORKING_DIR = '/kaggle/working'\nCHECKPOINT_DIR = os.path.join(WORKING_DIR, 'checkpoints')\nLOG_FILE = os.path.join(WORKING_DIR, 'log.txt')\nHISTORY_FILE = os.path.join(WORKING_DIR, 'history.json')\nos.makedirs(CHECKPOINT_DIR, exist_ok=True)\n\nimport torch\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint('Device dang su dung:', DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T12:59:00.58575Z","iopub.execute_input":"2026-07-11T12:59:00.586465Z","iopub.status.idle":"2026-07-11T12:59:06.256424Z","shell.execute_reply.started":"2026-07-11T12:59:00.586431Z","shell.execute_reply":"2026-07-11T12:59:06.255474Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"RESUME_INPUT_DIR = None  \n\nif RESUME_INPUT_DIR is not None and os.path.isdir(RESUME_INPUT_DIR):\n    import shutil\n    for fname in os.listdir(RESUME_INPUT_DIR):\n        shutil.copy(os.path.join(RESUME_INPUT_DIR, fname), CHECKPOINT_DIR)\n    print(f'Da copy checkpoint tu {RESUME_INPUT_DIR} vao {CHECKPOINT_DIR}, san sang resume.')\nelse:\n    print('Khong resume - train tu dau (backbone goc tu HuggingFace Meta).')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\nlabels_df = pd.read_csv(os.path.join(RSNA_ROOT, 'stage_2_train_labels.csv'))\nclass_info_df = pd.read_csv(os.path.join(RSNA_ROOT, 'stage_2_detailed_class_info.csv'))\n\n# Gop, moi patientId co the co nhieu dong (nhieu bbox) - lay duy nhat 1 dong/patient\nmerged = labels_df[['patientId', 'Target']].drop_duplicates(subset='patientId').merge(\n    class_info_df.drop_duplicates(subset='patientId'), on='patientId', how='left'\n)\n\n# Chi giu Normal va Lung Opacity, bo No Lung Opacity / Not Normal\nmerged = merged[merged['class'].isin(['Normal', 'Lung Opacity'])].copy()\nmerged['binary_label'] = (merged['class'] == 'Lung Opacity').astype(int)\nmerged['filepath'] = merged['patientId'].apply(lambda pid: os.path.join(rsna_img_dir, f'{pid}.dcm'))\nmerged = merged[merged['filepath'].apply(os.path.exists)]\nmerged['unique_patient_id'] = 'rsna_' + merged['patientId'].astype(str)\nmerged['source'] = 'rsna'\n\nrsna_samples_df = merged[['filepath', 'binary_label', 'unique_patient_id', 'source']]\nprint('RSNA sau khi loc:')\nprint(rsna_samples_df['binary_label'].value_counts().rename({0: 'NORMAL', 1: 'PNEUMONIA'}))\nprint('Tong RSNA su dung:', len(rsna_samples_df))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"nih_df = pd.read_csv(os.path.join(NIH_ROOT, 'Data_Entry_2017.csv'))\nnih_df = nih_df.rename(columns={'Image Index': 'filename', 'Patient ID': 'patient_id', 'Finding Labels': 'labels'})\n\nnih_normal_df = nih_df[nih_df['labels'].str.strip() == 'No Finding'].copy()\nnih_normal_df['filepath'] = nih_normal_df['filename'].map(nih_filename_to_path)\nnih_normal_df = nih_normal_df.dropna(subset=['filepath'])\n\nn_before = len(nih_normal_df)\nif EXCLUDE_LIST_PATH is not None and os.path.exists(EXCLUDE_LIST_PATH):\n    with open(EXCLUDE_LIST_PATH, 'r') as f:\n        excluded_filenames = set(line.strip() for line in f if line.strip())\n    nih_normal_df = nih_normal_df[~nih_normal_df['filename'].isin(excluded_filenames)]\n    print(f'Da loai {n_before - len(nih_normal_df)} anh NIH trung voi danh sach Stage 1 pretrain.')\nelse:\n    print('CANH BAO: EXCLUDE_LIST_PATH chua duoc dat - chua loai tru anh Stage 1 pretrain khoi phan NIH bo sung nay.')\n\nn_needed = max(TARGET_TOTAL_IMAGES - len(rsna_samples_df), 0)\nn_available = len(nih_normal_df)\nn_sample = min(n_needed, n_available)\nprint(f'Can bo sung {n_needed:,} anh NORMAL de du {TARGET_TOTAL_IMAGES:,} tong, NIH con {n_available:,} anh No Finding kha dung -> lay {n_sample:,} anh.')\n\nnih_supplement_df = nih_normal_df.sample(n=n_sample, random_state=SEED).copy()\nnih_supplement_df['binary_label'] = 0\nnih_supplement_df['unique_patient_id'] = 'nih_' + nih_supplement_df['patient_id'].astype(str)\nnih_supplement_df['source'] = 'nih'\n\nnih_samples_df = nih_supplement_df[['filepath', 'binary_label', 'unique_patient_id', 'source']]\n\nall_samples_df = pd.concat([rsna_samples_df, nih_samples_df], ignore_index=True)\nprint('\\nTong hop cuoi cung:')\nprint(all_samples_df['binary_label'].value_counts().rename({0: 'NORMAL', 1: 'PNEUMONIA'}))\nprint('Tong so anh:', len(all_samples_df))","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import GroupShuffleSplit\n\ngss1 = GroupShuffleSplit(n_splits=1, test_size=TEST_RATIO, random_state=SEED)\ntrainval_idx, test_idx = next(gss1.split(all_samples_df, all_samples_df['binary_label'], groups=all_samples_df['unique_patient_id']))\ntrainval_df = all_samples_df.iloc[trainval_idx].reset_index(drop=True)\ntest_df = all_samples_df.iloc[test_idx].reset_index(drop=True)\n\nval_fraction_of_trainval = VAL_RATIO / (TRAIN_RATIO + VAL_RATIO)\ngss2 = GroupShuffleSplit(n_splits=1, test_size=val_fraction_of_trainval, random_state=SEED)\ntrain_idx, val_idx = next(gss2.split(trainval_df, trainval_df['binary_label'], groups=trainval_df['unique_patient_id']))\ntrain_df = trainval_df.iloc[train_idx].reset_index(drop=True)\nval_df = trainval_df.iloc[val_idx].reset_index(drop=True)\n\n# ĐOẠN ĐÃ ĐƯỢC SỬA LẠI THẲNG HÀNG, XUỐNG DÒNG CHUẨN PYTHON:\nfor name, d in [('train', train_df), ('val', val_df), ('test', test_df)]:\n    counts = d['binary_label'].value_counts().rename({0: 'NORMAL', 1: 'PNEUMONIA'})\n    print(f'{name}: {len(d)} anh | {dict(counts)} | {d.unique_patient_id.nunique()} benh nhan | nguon: {dict(d.source.value_counts())}')\n\ntrain_samples = list(zip(train_df['filepath'], train_df['binary_label'], train_df['source']))\nval_samples = list(zip(val_df['filepath'], val_df['binary_label'], val_df['source']))\ntest_samples = list(zip(test_df['filepath'], test_df['binary_label'], test_df['source']))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T13:23:02.997692Z","iopub.execute_input":"2026-07-11T13:23:02.998523Z","iopub.status.idle":"2026-07-11T13:23:04.301428Z","shell.execute_reply.started":"2026-07-11T13:23:02.998488Z","shell.execute_reply":"2026-07-11T13:23:04.300405Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from transformers import AutoImageProcessor\nfrom torchvision import transforms\n\ntry:\n    processor = AutoImageProcessor.from_pretrained(MODEL_NAME, token=HF_TOKEN)\n    img_size = processor.crop_size['height'] if hasattr(processor, 'crop_size') and processor.crop_size else processor.size.get('height', 224)\n    mean = processor.image_mean\n    std = processor.image_std\n    print(f'Da tai AutoImageProcessor: size={img_size}, mean={mean}, std={std}')\nexcept Exception as e:\n    print(f'Fallback ImageNet normalization 224x224 ({e})')\n    img_size = 224\n    mean = [0.485, 0.456, 0.406]\n    std = [0.229, 0.224, 0.225]\n\ntrain_transform = transforms.Compose([\n    transforms.Resize((int(img_size * 1.15), int(img_size * 1.15))),\n    transforms.Grayscale(num_output_channels=3),\n    transforms.RandomResizedCrop(img_size, scale=(0.85, 1.0), ratio=(0.95, 1.05)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomRotation(degrees=10),\n    transforms.ColorJitter(brightness=0.15, contrast=0.15),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=mean, std=std),\n])\n\neval_transform = transforms.Compose([\n    transforms.Resize((img_size, img_size)),\n    transforms.Grayscale(num_output_channels=3),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=mean, std=std),\n])","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pydicom\nfrom PIL import Image\nfrom torch.utils.data import Dataset, DataLoader\n\nclass CombinedXrayDataset(Dataset):\n    def __init__(self, samples, transform=None):\n        self.samples = samples  # list of (filepath, label, source)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        path, label, source = self.samples[idx]\n        if source == 'rsna':\n            ds = pydicom.dcmread(path)\n            arr = ds.pixel_array.astype('float32')\n            arr = (arr - arr.min()) / (arr.max() - arr.min() + 1e-8) * 255.0\n            img = Image.fromarray(arr.astype('uint8')).convert('RGB')\n        else:\n            img = Image.open(path).convert('RGB')\n        if self.transform:\n            img = self.transform(img)\n        return img, label\n\ntrain_ds = CombinedXrayDataset(train_samples, transform=train_transform)\nval_ds = CombinedXrayDataset(val_samples, transform=eval_transform)\ntest_ds = CombinedXrayDataset(test_samples, transform=eval_transform)\n\n# SỬA ĐỔI: Chuyển về shuffle=True truyền thống, loại bỏ Sampler để tăng tốc và tăng Precision\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True, num_workers=NUM_WORKERS, pin_memory=True)\nval_loader = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)\ntest_loader = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS, pin_memory=True)\n\n# Sử dụng duy nhất trọng số này trong Loss để xử lý mất cân bằng mẫu dữ liệu\ntrain_labels = [lbl for _, lbl, _ in train_samples]\nclass_sample_count = np.array([train_labels.count(0), train_labels.count(1)])\ntotal_count = class_sample_count.sum()\nclass_weights = torch.tensor(\n    [total_count / (NUM_CLASSES * c) if c > 0 else 0.0 for c in class_sample_count],\n    dtype=torch.float32,\n).to(DEVICE)\nprint(f'Class weights dung trong Loss: NORMAL={class_weights[0]:.3f}, PNEUMONIA={class_weights[1]:.3f}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nimport time\nimport matplotlib.pyplot as plt\n\ndef log_msg(msg):\n    ts = time.strftime('%Y-%m-%d %H:%M:%S')\n    line = f'[{ts}] {msg}'\n    print(line)\n    with open(LOG_FILE, 'a') as f:\n        f.write(line + '\\n')\n\ndef plot_curves(history, save_path=None, title_suffix=''):\n    fig, axes = plt.subplots(1, 3, figsize=(18, 5))\n    axes[0].plot(history['train_loss'], label='Train Loss')\n    axes[0].plot(history['val_loss'], label='Val Loss')\n    axes[0].set_title(f'Loss theo epoch{title_suffix}')\n    axes[0].set_xlabel('Epoch'); axes[0].set_ylabel('Loss'); axes[0].legend()\n\n    axes[1].plot(history['train_acc'], label='Train Acc')\n    axes[1].plot(history['val_acc'], label='Val Acc')\n    axes[1].set_title(f'Accuracy theo epoch{title_suffix}')\n    axes[1].set_xlabel('Epoch'); axes[1].set_ylabel('Accuracy'); axes[1].legend()\n\n    # Thêm trục AUC để tiện theo dõi trực quan xu hướng chuẩn xác nhất\n    axes[2].plot(history.get('val_auc', []), label='Val AUC', color='green')\n    axes[2].set_title(f'Validation ROC AUC{title_suffix}')\n    axes[2].set_xlabel('Epoch'); axes[2].set_ylabel('AUC'); axes[2].legend()\n\n    plt.tight_layout()\n    if save_path:\n        plt.savefig(save_path, dpi=150)\n    plt.show()\n\ndef save_checkpoint(model, optimizer, scheduler, epoch, best_val_auc, history, is_best=False):\n    os.makedirs(CHECKPOINT_DIR, exist_ok=True)\n    model_to_save = get_base_model() if 'get_base_model' in globals() else model\n    state = {\n        'epoch': epoch,\n        'model_state_dict': model_to_save.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'scheduler_state_dict': scheduler.state_dict() if scheduler is not None else None,\n        'best_val_auc': best_val_auc,\n    }\n    torch.save(state, os.path.join(CHECKPOINT_DIR, 'last.pt'))\n    if is_best:\n        torch.save(state, os.path.join(CHECKPOINT_DIR, 'best.pt'))\n    with open(HISTORY_FILE, 'w') as f:\n        json.dump(history, f, indent=2)\n\ndef load_checkpoint_if_exists(model, optimizer, scheduler):\n    last_path = os.path.join(CHECKPOINT_DIR, 'last.pt')\n    # Tích hợp thêm val_auc và f1-macro vào lịch sử nhật ký lưu trữ\n    history = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': [], 'val_auc': [], 'val_f1_macro': []}\n    start_epoch = 0\n    best_val_auc = 0.0\n    if os.path.exists(last_path):\n        log_msg(f'Tim thay checkpoint tai {last_path}, dang resume...')\n        state = torch.load(last_path, map_location=DEVICE)\n        model_to_load = get_base_model() if 'get_base_model' in globals() else model\n        model_to_load.load_state_dict(state['model_state_dict'])\n        optimizer.load_state_dict(state['optimizer_state_dict'])\n        if scheduler is not None and state.get('scheduler_state_dict') is not None:\n            scheduler.load_state_dict(state['scheduler_state_dict'])\n        start_epoch = state['epoch'] + 1\n        best_val_auc = state.get('best_val_auc', 0.0)\n        if os.path.exists(HISTORY_FILE):\n            with open(HISTORY_FILE, 'r') as f:\n                history = json.load(f)\n        log_msg(f'Resume tu epoch {start_epoch}, best_val_auc hien tai = {best_val_auc:.4f}')\n    else:\n        log_msg('Khong tim thay checkpoint cu. Tien hanh fine-tune tu dau.')\n        \n    # SỬA ĐỔI: Khôi phục lệnh return lề chuẩn tránh trả về None gây sập unpack\n    return start_epoch, best_val_auc, history","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\nfrom transformers import AutoModel\n\nclass DinoV3Classifier(nn.Module):\n    def __init__(self, model_name, num_classes, hf_token=None):\n        super().__init__()\n        self.backbone = AutoModel.from_pretrained(model_name, token=hf_token)\n        hidden_size = self.backbone.config.hidden_size\n        self.classifier = nn.Linear(hidden_size, num_classes)\n        for p in self.backbone.parameters():\n            p.requires_grad = True\n\n    def forward(self, pixel_values):\n        outputs = self.backbone(pixel_values=pixel_values)\n        cls_token = outputs.last_hidden_state[:, 0, :]\n        return self.classifier(cls_token)\n\nbase_model = DinoV3Classifier(MODEL_NAME, NUM_CLASSES, hf_token=HF_TOKEN).to(DEVICE)\n\nn_total = sum(p.numel() for p in base_model.parameters())\nn_trainable = sum(p.numel() for p in base_model.parameters() if p.requires_grad)\nlog_msg(f'Xac nhan: Su dung backbone goc cua Meta, khong nạp tiep tu trong so pretrained.')\nlog_msg(f'Tong tham so: {n_total:,} | Co the train (full fine-tune): {n_trainable:,}')\n\nn_gpus = torch.cuda.device_count()\nif USE_MULTI_GPU and n_gpus > 1:\n    model = nn.DataParallel(base_model)\n    log_msg(f'Dang dung {n_gpus} GPU qua nn.DataParallel.')\nelse:\n    model = base_model\n    log_msg(f'Dung 1 GPU ({n_gpus} GPU kha dung, USE_MULTI_GPU={USE_MULTI_GPU}).')\n\ndef get_base_model():\n    return model.module if isinstance(model, nn.DataParallel) else model","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.optim as optim\n\noptimizer = optim.AdamW([\n    {'params': get_base_model().backbone.parameters(), 'lr': BACKBONE_LR},\n    {'params': get_base_model().classifier.parameters(), 'lr': HEAD_LR},\n], weight_decay=WEIGHT_DECAY)\n\nwarmup_scheduler = optim.lr_scheduler.LinearLR(optimizer, start_factor=0.1, end_factor=1.0, total_iters=WARMUP_EPOCHS)\ncosine_scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=max(NUM_EPOCHS - WARMUP_EPOCHS, 1))\nscheduler = optim.lr_scheduler.SequentialLR(optimizer, schedulers=[warmup_scheduler, cosine_scheduler], milestones=[WARMUP_EPOCHS])\n\ncriterion = nn.CrossEntropyLoss(weight=class_weights)\nscaler = torch.amp.GradScaler('cuda', enabled=USE_AMP)\n\nlog_msg(f'Optimizer: backbone_lr={BACKBONE_LR}, head_lr={HEAD_LR}, warmup_epochs={WARMUP_EPOCHS}, grad_clip_norm={GRAD_CLIP_NORM}, use_amp={USE_AMP}')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm.auto import tqdm\nfrom sklearn.metrics import roc_auc_score, f1_score\n\ndef run_epoch(loader, train_mode):\n    model.train() if train_mode else model.eval()\n    total_loss, total_correct, total_samples = 0.0, 0, 0\n    all_probs, all_labels = [], []\n    \n    context = torch.enable_grad() if train_mode else torch.no_grad()\n    with context:\n        for images, labels in tqdm(loader, leave=False):\n            images, labels = images.to(DEVICE), labels.to(DEVICE)\n            if train_mode:\n                optimizer.zero_grad()\n            with torch.amp.autocast('cuda', enabled=USE_AMP):\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n            if train_mode:\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=GRAD_CLIP_NORM)\n                scaler.step(optimizer)\n                scaler.update()\n                \n            total_loss += loss.item() * images.size(0)\n            total_correct += (outputs.argmax(dim=1) == labels).sum().item()\n            total_samples += images.size(0)\n            \n            if not train_mode:\n                probs = torch.softmax(outputs, dim=1)[:, 1]\n                all_probs.extend(probs.detach().cpu().numpy())\n                all_labels.extend(labels.cpu().numpy())\n\n    avg_loss = total_loss / total_samples\n    avg_acc = total_correct / total_samples\n    \n    if train_mode:\n        return avg_loss, avg_acc\n\n    # SỬA ĐỔI: Do lường bổ sung AUC không phụ thuộc ngưỡng và F1-macro cho Validation Set\n    auc = roc_auc_score(all_labels, all_probs) if len(set(all_labels)) > 1 else float('nan')\n    preds_bin = [1 if p >= 0.5 else 0 for p in all_probs]\n    f1_macro = f1_score(all_labels, preds_bin, average='macro')\n    return avg_loss, avg_acc, auc, f1_macro","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"start_epoch = 0\nbest_val_auc = 0.0\nhistory = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': [], 'val_auc': [], 'val_f1_macro': []}\n\ntry:\n    start_epoch, best_val_auc, history = load_checkpoint_if_exists(model, optimizer, scheduler)\nexcept Exception as e:\n    log_msg(f'Luu y an toan nap checkpoint: {e}. Train tu dau.')\n\nbest_val_loss = min(history['val_loss']) if history['val_loss'] else float('inf')\nno_improve_epochs = 0\nSPIKE_PATIENCE = 3  # Tự động dừng sớm nếu qua 3 epoch mà AUC không tăng thêm\n\nfor epoch in range(start_epoch, NUM_EPOCHS):\n    epoch_start = time.time()\n    train_loss, train_acc = run_epoch(train_loader, train_mode=True)\n    val_loss, val_acc, val_auc, val_f1 = run_epoch(val_loader, train_mode=False)\n    scheduler.step()\n    elapsed = time.time() - epoch_start\n\n    history['train_loss'].append(train_loss)\n    history['val_loss'].append(val_loss)\n    history['train_acc'].append(train_acc)\n    history['val_acc'].append(val_acc)\n    history['val_auc'].append(val_auc)\n    history['val_f1_macro'].append(val_f1)\n\n    # SỬA ĐỔI CHÍNH: Lựa chọn mô hình lưu trữ tốt nhất dựa trên Validation AUC bền vững\n    is_best = val_auc > best_val_auc\n    if is_best:\n        best_val_auc = val_auc\n        no_improve_epochs = 0\n    else:\n        no_improve_epochs += 1\n\n    log_msg(\n        f'Epoch {epoch+1}/{NUM_EPOCHS} | '\n        f'train_loss={train_loss:.4f} train_acc={train_acc:.4f} | '\n        f'val_loss={val_loss:.4f} val_acc={val_acc:.4f} val_auc={val_auc:.4f} val_f1_macro={val_f1:.4f} | '\n        f'best_val_auc={best_val_auc:.4f} | thoi_gian={elapsed:.1f}s' + (' *BEST*' if is_best else '')\n    )\n\n    save_checkpoint(model, optimizer, scheduler, epoch, best_val_auc, history, is_best=is_best)\n\n    # Thêm cơ chế Early Stopping để ngăn mô hình học thuộc lòng tập train (Overfitting)\n    if no_improve_epochs >= SPIKE_PATIENCE:\n        log_msg(f'EARLY STOPPING: Dung huan luyen som o epoch {epoch+1} do chi so AUC dung im suot {SPIKE_PATIENCE} epochs.')\n        break\n\n    is_spike = val_loss > best_val_loss * SPIKE_THRESHOLD\n    if is_spike:\n        log_msg(f'CANH BAO: val_loss={val_loss:.4f} bien dong.')\n        plot_curves(history, title_suffix=f' - CANH BAO tai epoch {epoch+1}')\n    else:\n        best_val_loss = min(best_val_loss, val_loss)\n\nlog_msg('HOAN TAT TRAINING BASELINE GỐC.')\nprint('\\nNHO BAM \"Save Version\" (goc tren phai) de luu checkpoint/log/history vao output cua notebook!')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import (accuracy_score, precision_score, recall_score,\n                             f1_score, confusion_matrix, classification_report, roc_auc_score)\nimport seaborn as sns\n\nbest_path = os.path.join(CHECKPOINT_DIR, 'best.pt')\nif os.path.exists(best_path):\n    state = torch.load(best_path, map_location=DEVICE)\n    get_base_model().load_state_dict(state['model_state_dict'])\n    log_msg(f\"Da load best checkpoint (Chon loc theo val_auc) de danh gia test set.\")\n\nmodel.eval()\nall_preds, all_labels, all_probs = [], [], []\nwith torch.no_grad():\n    for images, labels in tqdm(test_loader, desc='Danh gia test set'):\n        images = images.to(DEVICE)\n        outputs = model(images)\n        probs = torch.softmax(outputs, dim=1)[:, 1]\n        preds = outputs.argmax(dim=1).cpu().numpy()\n        all_preds.extend(preds)\n        all_labels.extend(labels.numpy())\n        all_probs.extend(probs.cpu().numpy())\n\ntest_acc = accuracy_score(all_labels, all_preds)\ntest_precision = precision_score(all_labels, all_preds)\ntest_recall = recall_score(all_labels, all_preds)\ntest_f1 = f1_score(all_labels, all_preds)\ntest_auc = roc_auc_score(all_labels, all_probs)\ncm = confusion_matrix(all_labels, all_preds)\n\nlog_msg(f'TEST - Accuracy={test_acc:.4f} Precision={test_precision:.4f} Recall={test_recall:.4f} F1={test_f1:.4f} AUC={test_auc:.4f}')\nlog_msg('Confusion matrix:\\n' + str(cm))\nlog_msg('Classification report:\\n' + classification_report(all_labels, all_preds, target_names=CLASS_NAMES))\n\nresults_summary = {\n    'test_accuracy': test_acc, 'test_precision': test_precision, 'test_recall': test_recall,\n    'test_f1': test_f1, 'test_auc': test_auc, 'confusion_matrix': cm.tolist(),\n    'backbone_source': 'HuggingFace Meta nguyen ban (Baseline Goc)'\n}\nwith open(os.path.join(WORKING_DIR, 'test_results.json'), 'w') as f:\n    json.dump(results_summary, f, indent=2)\n\n# Xuất biểu đồ ma trận nhầm lẫn\nplt.figure(figsize=(5, 4))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=CLASS_NAMES, yticklabels=CLASS_NAMES)\nplt.xlabel('Predicted'); plt.ylabel('Actual'); plt.title('Confusion Matrix - Baseline Goc')\nplt.tight_layout()\nplt.savefig(os.path.join(WORKING_DIR, 'confusion_matrix.png'), dpi=150)\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_curves(history, save_path=os.path.join(WORKING_DIR, 'loss_acc_curves.png'))","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}