{"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":"import os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        if 'best.pt' in filename:\n            print(os.path.join(dirname, filename))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-07-11T07:45:12.123275Z","iopub.execute_input":"2026-07-11T07:45:12.123823Z","iopub.status.idle":"2026-07-11T07:46:18.509493Z","shell.execute_reply.started":"2026-07-11T07:45:12.123794Z","shell.execute_reply":"2026-07-11T07:46:18.508599Z"}},"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()\nHF_TOKEN = user_secrets.get_secret('Jasper')\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-11T07:46:52.841196Z","iopub.execute_input":"2026-07-11T07:46:52.841644Z","iopub.status.idle":"2026-07-11T07:46:53.075768Z","shell.execute_reply.started":"2026-07-11T07:46:52.841603Z","shell.execute_reply":"2026-07-11T07:46:53.074996Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 3. NẠP TRỌNG SỐ TỪ CHECKPOINT VÒNG 1 (Ép theo thứ tự danh sách tham số)\nimport torch\nimport torch.nn as nn\nfrom transformers import AutoModel\nfrom kaggle_secrets import UserSecretsClient\n\n# Khai báo đường dẫn và thiết bị trực tiếp tại đây để tránh lỗi tuần tự\ncheckpoint_path = \"/kaggle/input/notebooks/vonguyenkhang/dinov3-finetune-xray-nih-rsna/checkpoints/best.pt\"\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nMODEL_NAME = 'facebook/dinov3-vitb16-pretrain-lvd1689m'\nNUM_CLASSES = 2\n\n# Khởi tạo cấu trúc mạng chuẩn của Khang\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\n# Đăng nhập token\ntry:\n    HF_TOKEN = UserSecretsClient().get_secret('Jasper')\nexcept:\n    HF_TOKEN = None\n\n# Tạo mô hình và nạp trọng số\nmodel = DinoV3Classifier(MODEL_NAME, NUM_CLASSES, hf_token=HF_TOKEN).to(DEVICE)\nmodel_dict = model.state_dict()\n\ncheckpoint = torch.load(checkpoint_path, map_location=DEVICE)\nkhang_state_dict = checkpoint['model_state_dict']\n\nkhang_keys = list(khang_state_dict.keys())\nmy_keys = list(model_dict.keys())\n\ncount = 0\nfor i in range(min(len(khang_keys), len(my_keys))):\n    khang_k = khang_keys[i]\n    my_k = my_keys[i]\n    if khang_state_dict[khang_k].shape == model_dict[my_k].shape:\n        model_dict[my_k] = khang_state_dict[khang_k]\n        count += 1\n\nmodel.load_state_dict(model_dict, strict=False)\nprint(f\"=> XUẤT SẮC! Đã ép nạp thành công trọn vẹn {count}/{len(my_keys)} layers theo đúng thứ tự cấu trúc mạng ViT!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T07:46:58.980026Z","iopub.execute_input":"2026-07-11T07:46:58.980773Z","iopub.status.idle":"2026-07-11T07:46:59.628195Z","shell.execute_reply.started":"2026-07-11T07:46:58.980732Z","shell.execute_reply":"2026-07-11T07:46:59.627457Z"}},"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 - kiem tra da Add Data rsna-pneumonia-detection-challenge chua.'\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 - kiem tra da Add Data nih-chest-xrays/data chua.'\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-11T07:47:07.844110Z","iopub.execute_input":"2026-07-11T07:47:07.844768Z","iopub.status.idle":"2026-07-11T07:47:56.632490Z","shell.execute_reply.started":"2026-07-11T07:47:07.844737Z","shell.execute_reply":"2026-07-11T07:47:56.631560Z"}},"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 = 20        # giam tu 50 -> 20: dataset ~50k anh da co nhieu luot cap nhat gradient hon Kermany roi\nBATCH_SIZE = 64        # tang tu 32 -> 64, AMP giai phong bo nho de dung batch lon hon, 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         # mixed precision - tan dung tensor core cua T4, nhanh hon ~1.5-2x, khong doi chat luong\nUSE_MULTI_GPU = True   # Kaggle T4 x2 co 2 GPU - tu dong dung ca 2 qua nn.DataParallel neu co san\n\nTARGET_TOTAL_IMAGES = 50000   # tong so anh muc tieu (RSNA + NIH-Normal bo sung)\nTRAIN_RATIO = 0.8\nVAL_RATIO = 0.1\nTEST_RATIO = 0.1              # con lai sau train+val\n\nSEED = 42\nSPIKE_THRESHOLD = 1.5\nSPIKE_PATIENCE = 3\n\nEXCLUDE_LIST_PATH = None  # vd: '/kaggle/input/stage1-pretrain-filelist/pretrain_filenames.txt' (ap dung cho phan NIH)\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:', DEVICE)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T07:26:26.208032Z","iopub.execute_input":"2026-07-11T07:26:26.208623Z","iopub.status.idle":"2026-07-11T07:26:26.216063Z","shell.execute_reply.started":"2026-07-11T07:26:26.208592Z","shell.execute_reply":"2026-07-11T07:26:26.215169Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"RESUME_INPUT_DIR = None  # vd: '/kaggle/input/ten-notebook-cu/checkpoints'\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 (hoac tiep tuc checkpoint co san trong /kaggle/working neu chua Restart session).')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T07:26:29.440034Z","iopub.execute_input":"2026-07-11T07:26:29.440561Z","iopub.status.idle":"2026-07-11T07:26:29.445535Z","shell.execute_reply.started":"2026-07-11T07:26:29.440525Z","shell.execute_reply":"2026-07-11T07:26:29.444683Z"}},"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,"execution":{"iopub.status.busy":"2026-07-11T07:26:33.594784Z","iopub.execute_input":"2026-07-11T07:26:33.595175Z","iopub.status.idle":"2026-07-11T07:26:42.666883Z","shell.execute_reply.started":"2026-07-11T07:26:33.595146Z","shell.execute_reply":"2026-07-11T07:26:42.666045Z"}},"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,"execution":{"iopub.status.busy":"2026-07-11T07:27:17.519236Z","iopub.execute_input":"2026-07-11T07:27:17.519666Z","iopub.status.idle":"2026-07-11T07:27:17.811215Z","shell.execute_reply.started":"2026-07-11T07:27:17.519636Z","shell.execute_reply":"2026-07-11T07:27:17.810431Z"}},"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\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-11T07:27:40.350670Z","iopub.execute_input":"2026-07-11T07:27:40.351055Z","iopub.status.idle":"2026-07-11T07:27:40.489266Z","shell.execute_reply.started":"2026-07-11T07:27:40.351027Z","shell.execute_reply":"2026-07-11T07:27:40.488446Z"}},"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,"execution":{"iopub.status.busy":"2026-07-11T07:27:53.615839Z","iopub.execute_input":"2026-07-11T07:27:53.616587Z","iopub.status.idle":"2026-07-11T07:27:54.121249Z","shell.execute_reply.started":"2026-07-11T07:27:53.616556Z","shell.execute_reply":"2026-07-11T07:27:54.120653Z"}},"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# WeightedRandomSampler cho train de moi batch co du 2 lop du van con mat can bang sau khi gop du lieu\ntrain_labels = [lbl for _, lbl, _ in train_samples]\nclass_sample_count = np.array([train_labels.count(0), train_labels.count(1)])\nclass_weights_for_sampler = 1.0 / class_sample_count\nsample_weights = np.array([class_weights_for_sampler[lbl] for lbl in train_labels])\nsampler = torch.utils.data.WeightedRandomSampler(\n    weights=torch.from_numpy(sample_weights).double(), num_samples=len(sample_weights), replacement=True\n)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler, 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# Class weight cho loss (bo sung, dung kem WeightedRandomSampler cho chac)\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 (loss): NORMAL={class_weights[0]:.3f}, PNEUMONIA={class_weights[1]:.3f}')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T07:28:18.652090Z","iopub.execute_input":"2026-07-11T07:28:18.652480Z","iopub.status.idle":"2026-07-11T07:28:18.675558Z","shell.execute_reply.started":"2026-07-11T07:28:18.652451Z","shell.execute_reply":"2026-07-11T07:28:18.674983Z"}},"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, 2, figsize=(14, 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    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_acc, 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_acc': best_val_acc,\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    history = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': []}\n    start_epoch = 0\n    best_val_acc = 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_acc = state.get('best_val_acc', 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_acc hien tai = {best_val_acc:.4f}')\n    else:\n        log_msg('Khong tim thay checkpoint, bat dau train tu dau (epoch 0).')\n        \n    # Lệnh return chuẩn vị trí lề để tránh trả về NoneType gây lỗi treo máy ảo\n    return start_epoch, best_val_acc, history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T07:28:38.638861Z","iopub.execute_input":"2026-07-11T07:28:38.639700Z","iopub.status.idle":"2026-07-11T07:28:38.651134Z","shell.execute_reply.started":"2026-07-11T07:28:38.639661Z","shell.execute_reply":"2026-07-11T07:28:38.650526Z"}},"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'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    \"\"\"Tra ve model goc (khong bi DataParallel boc ngoai) - dung khi luu/doc checkpoint hoac truy cap .backbone/.classifier.\"\"\"\n    return model.module if isinstance(model, nn.DataParallel) else model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T07:29:03.287908Z","iopub.execute_input":"2026-07-11T07:29:03.288408Z","iopub.status.idle":"2026-07-11T07:29:04.025355Z","shell.execute_reply.started":"2026-07-11T07:29:03.288367Z","shell.execute_reply":"2026-07-11T07:29:04.024450Z"}},"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)\n\nscaler = torch.cuda.amp.GradScaler(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,"execution":{"iopub.status.busy":"2026-07-11T07:29:44.027978Z","iopub.execute_input":"2026-07-11T07:29:44.028511Z","iopub.status.idle":"2026-07-11T07:29:44.035438Z","shell.execute_reply.started":"2026-07-11T07:29:44.028482Z","shell.execute_reply":"2026-07-11T07:29:44.034754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm.auto import tqdm\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    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.cuda.amp.autocast(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            total_loss += loss.item() * images.size(0)\n            total_correct += (outputs.argmax(dim=1) == labels).sum().item()\n            total_samples += images.size(0)\n    return total_loss / total_samples, total_correct / total_samples","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T07:30:54.770966Z","iopub.execute_input":"2026-07-11T07:30:54.771832Z","iopub.status.idle":"2026-07-11T07:30:54.779053Z","shell.execute_reply.started":"2026-07-11T07:30:54.771794Z","shell.execute_reply":"2026-07-11T07:30:54.778042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Thiết lập mặc định an toàn cho biến\nstart_epoch = 0\nbest_val_acc = 0.0\nhistory = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': []}\n\n# Gọi hàm load_checkpoint có bọc bộ lọc tránh treo máy ảo\ntry:\n    start_epoch, best_val_acc, history = load_checkpoint_if_exists(model, optimizer, scheduler)\nexcept Exception as e:\n    log_msg(f'Luu y khi goi checkpoint: {e}. Tien hanh huan luyen voi gia tri mac dinh.')\n\nbest_val_loss = min(history['val_loss']) if history['val_loss'] else float('inf')\nspike_streak = 0\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 = 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\n    is_best = val_acc > best_val_acc\n    if is_best:\n        best_val_acc = val_acc\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} | '\n        f'best_val_acc={best_val_acc:.4f} | thoi_gian={elapsed:.1f}s' + (' *BEST*' if is_best else '')\n    )\n\n    save_checkpoint(model, optimizer, scheduler, epoch, best_val_acc, history, is_best=is_best)\n\n    is_spike = val_loss > best_val_loss * SPIKE_THRESHOLD\n    if is_spike:\n        spike_streak += 1\n        log_msg(f'CANH BAO: val_loss={val_loss:.4f} vuot nguong on dinh (best={best_val_loss:.4f} x {SPIKE_THRESHOLD}) - lan {spike_streak}/{SPIKE_PATIENCE}.')\n        plot_curves(history, title_suffix=f' - CANH BAO tai epoch {epoch+1}')\n        if spike_streak >= SPIKE_PATIENCE:\n            log_msg(f'DUNG TRAINING SOM tai epoch {epoch+1} do {spike_streak} lan dao dong bat thuong lien tiep.')\n            break\n    else:\n        spike_streak = 0\n        best_val_loss = min(best_val_loss, val_loss)\n\nlog_msg('HOAN TAT TRAINING (hoac da dung som do canh bao instability - xem log ben tren).')\nprint('\\nNHO BAM \"Save Version\" (goc tren phai) de luu checkpoint/log/history vao output cua notebook!')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-07-11T07:33:07.546721Z","iopub.execute_input":"2026-07-11T07:33:07.547109Z"}},"outputs":[],"execution_count":null}]}