{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"83905618","cell_type":"markdown","source":"## Import, cấu hình và siêu tham số","metadata":{}},{"id":"02177b55","cell_type":"code","source":"import os\nimport random\nfrom collections import defaultdict\nfrom pathlib import Path\n\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import models\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom skimage.morphology import reconstruction as skimage_reconstruct\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    accuracy_score,\n    balanced_accuracy_score,\n    classification_report,\n    confusion_matrix,\n    precision_recall_fscore_support,\n    matthews_corrcoef,\n    cohen_kappa_score,\n    roc_auc_score,\n)\n\nAPTOS_CSV = Path('/kaggle/input/competitions/aptos2019-blindness-detection/train.csv')\nAPTOS_IMAGE_DIR = Path('/kaggle/input/competitions/aptos2019-blindness-detection/train_images')\nDDR_CSV = Path('/kaggle/input/datasets/mariaherrerot/ddrdataset/DR_grading.csv')\nDDR_IMAGE_ROOT = Path('/kaggle/input/datasets/mariaherrerot/ddrdataset/DR_grading')\n\nWORK_DIR = Path('/kaggle/working/dr_set1_paper_esrgan_two_stage')\nPRE_ESRGAN_DIR = WORK_DIR / 'pre_esrgan_hr224'\nPREPROCESSED_DIR = WORK_DIR / 'preprocessed_paper_esrgan'\nMANIFEST_DIR = WORK_DIR / 'manifests'\nESRGAN_CHECKPOINT_DIR = WORK_DIR / 'esrgan_checkpoints'\nCHECKPOINT_PATH = WORK_DIR / 'best_set1_resnet50.pth'\nESRGAN_PRETRAIN_CHECKPOINT = ESRGAN_CHECKPOINT_DIR / 'generator_pretrain_l1.pth'\nESRGAN_GAN_CHECKPOINT = ESRGAN_CHECKPOINT_DIR / 'generator_ragan_final.pth'\n\n# Checkpoint ESRGAN đã huấn luyện và lưu trên Kaggle Models.\nEXTERNAL_ESRGAN_PRETRAIN_CHECKPOINT = Path(\n    '/kaggle/input/models/laimochuy/pretrainl1-ragan/pytorch/default/1/'\n    'generator_pretrain_l1.pth'\n)\nEXTERNAL_ESRGAN_GAN_CHECKPOINT = Path(\n    '/kaggle/input/models/laimochuy/pretrainl1-ragan/pytorch/default/1/'\n    'generator_ragan_final.pth'\n)\n\n# True: tải checkpoint Kaggle Models và bỏ qua huấn luyện ESRGAN.\n# False: quay lại cơ chế huấn luyện/checkpoint trong /kaggle/working.\nUSE_EXTERNAL_ESRGAN_CHECKPOINTS = True\n\nfor directory in (\n    WORK_DIR,\n    PRE_ESRGAN_DIR,\n    PREPROCESSED_DIR,\n    MANIFEST_DIR,\n    ESRGAN_CHECKPOINT_DIR,\n):\n    directory.mkdir(parents=True, exist_ok=True)\n\nSEED = 42\nNUM_CLASSES = 5\nCLASS_NAMES = ['No DR', 'Mild NPDR', 'Moderate NPDR', 'Severe NPDR', 'PDR']\n\nIMAGE_SIZE = 224\nEPOCHS = 25\nLEARNING_RATE = 1e-3\nBATCH_SIZE = 5\nLR_STEP_SIZE = 3\nLR_GAMMA = 0.5  # Bài báo không công bố; giả định phổ biến để tái lập.\nNUM_WORKERS = 2\n\n\nCLAHE_CLIP_LIMIT = 2.0\nCLAHE_TILE_GRID = (8, 8)\nGAUSSIAN_KERNEL = (5, 5)  # Bài báo chỉ nói Gaussian blur, không báo kernel.\nGAUSSIAN_SIGMA = 0\n\n# ESRGAN — TABLE 5\nESRGAN_SCALE = 4\nESRGAN_NUM_RRDB = 23\nESRGAN_FEATURES = 64\nESRGAN_GROWTH_CHANNELS = 32\nESRGAN_RESIDUAL_SCALE = 0.2\nESRGAN_HR_PATCH = 128\nESRGAN_LR_PATCH = ESRGAN_HR_PATCH // ESRGAN_SCALE\nESRGAN_BATCH_SIZE = 16\nESRGAN_LR = 1e-4\nESRGAN_BETAS = (0.9, 0.999)\nESRGAN_PIXEL_WEIGHT = 1.0\nESRGAN_PERCEPTUAL_WEIGHT = 1.0\nESRGAN_ADVERSARIAL_WEIGHT = 0.005\n\n# Bài báo không công bố số epoch ESRGAN. 10 + 10 là giả định thực thi.\nESRGAN_PRETRAIN_EPOCHS = 10\nESRGAN_GAN_EPOCHS = 10\nESRGAN_MAX_STEPS_PER_EPOCH = None  \nESRGAN_VAL_IMAGES = 512            \nFORCE_RETRAIN_ESRGAN = False\n\n# Batch inference cho toàn bộ ảnh sau khi tải checkpoint.\n# T4 x2: 32 thường phù hợp; giảm còn 16 nếu CUDA báo hết bộ nhớ.\nESRGAN_INFERENCE_BATCH_SIZE = 32\n\nMAX_IMAGES_PER_SOURCE = None\n\nDEVICE = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')\nUSE_AMP = DEVICE.type == 'cuda'\nGPU_COUNT = torch.cuda.device_count() if torch.cuda.is_available() else 0\nUSE_MULTI_GPU_ESRGAN = GPU_COUNT > 1\n\n\ndef seed_everything(seed: int = 42) -> None:\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\n\nseed_everything(SEED)\nprint('Device:', DEVICE)\nprint('GPU count:', GPU_COUNT)\nif USE_MULTI_GPU_ESRGAN:\n    print('ESRGAN sẽ dùng DataParallel trên', GPU_COUNT, 'GPU.')\nelse:\n    print('ESRGAN chạy trên một GPU/CPU.')\nACTIVE_ESRGAN_PRETRAIN_CHECKPOINT = (\n    EXTERNAL_ESRGAN_PRETRAIN_CHECKPOINT\n    if USE_EXTERNAL_ESRGAN_CHECKPOINTS\n    else ESRGAN_PRETRAIN_CHECKPOINT\n)\nACTIVE_ESRGAN_GAN_CHECKPOINT = (\n    EXTERNAL_ESRGAN_GAN_CHECKPOINT\n    if USE_EXTERNAL_ESRGAN_CHECKPOINTS\n    else ESRGAN_GAN_CHECKPOINT\n)\n\nprint('Work dir:', WORK_DIR)\nprint('Dùng checkpoint ESRGAN ngoài:', USE_EXTERNAL_ESRGAN_CHECKPOINTS)\nprint('Pretrain checkpoint:', ACTIVE_ESRGAN_PRETRAIN_CHECKPOINT)\nprint('RaGAN checkpoint:', ACTIVE_ESRGAN_GAN_CHECKPOINT)\n","metadata":{"execution":{"iopub.execute_input":"2026-06-24T05:55:59.407858Z","iopub.status.busy":"2026-06-24T05:55:59.406716Z","iopub.status.idle":"2026-06-24T05:55:59.426526Z","shell.execute_reply":"2026-06-24T05:55:59.425714Z","shell.execute_reply.started":"2026-06-24T05:55:59.407823Z"}},"outputs":[],"execution_count":null},{"id":"ad637bd1","cell_type":"markdown","source":"## Đọc APTOS và DDR, kiểm tra file ảnh","metadata":{}},{"id":"8dd0f32e","cell_type":"code","source":"IMAGE_EXTENSIONS = {'.jpg', '.jpeg', '.png', '.tif', '.tiff', '.bmp'}\n\n\ndef require_path(path: Path, description: str) -> None:\n    if not path.exists():\n        raise FileNotFoundError(\n            f'{description} không tồn tại: {path}\\n'\n            'Hãy sửa đường dẫn trong Cell cấu hình trước khi chạy tiếp.'\n        )\n\n\ndef build_image_index(root: Path) -> dict:\n    require_path(root, 'Thư mục ảnh')\n    index = {}\n    collisions = defaultdict(list)\n\n    for path in tqdm(root.rglob('*'), desc=f'Lập chỉ mục {root.name}'):\n        if path.is_file() and path.suffix.lower() in IMAGE_EXTENSIONS:\n            keys = {path.name.lower(), path.stem.lower()}\n            for key in keys:\n                if key in index and index[key] != path:\n                    collisions[key].append(path)\n                else:\n                    index[key] = path\n\n    if collisions:\n        print(f'Cảnh báo: có {len(collisions)} khóa tên ảnh bị trùng; ưu tiên file gặp đầu tiên.')\n    return index\n\n\ndef resolve_image_path(raw_value, root: Path, image_index: dict) -> Path | None:\n    raw = str(raw_value).strip().replace('\\\\', '/')\n    direct = root / raw\n    if direct.exists():\n        return direct\n\n    name = Path(raw).name.lower()\n    stem = Path(raw).stem.lower()\n    for key in (name, stem):\n        if key in image_index:\n            return image_index[key]\n\n    if Path(raw).suffix == '':\n        for ext in IMAGE_EXTENSIONS:\n            candidate = root / f'{raw}{ext}'\n            if candidate.exists():\n                return candidate\n    return None\n\n\ndef load_aptos() -> pd.DataFrame:\n    require_path(APTOS_CSV, 'APTOS CSV')\n    require_path(APTOS_IMAGE_DIR, 'APTOS image directory')\n\n    df = pd.read_csv(APTOS_CSV)\n    required = {'id_code', 'diagnosis'}\n    if not required.issubset(df.columns):\n        raise ValueError(f'APTOS CSV thiếu cột {required}. Cột hiện có: {df.columns.tolist()}')\n\n    df = df[['id_code', 'diagnosis']].copy()\n    df['source'] = 'APTOS'\n    df['image_id'] = df['id_code'].astype(str)\n    df['label'] = df['diagnosis'].astype(int)\n    df['original_path'] = df['image_id'].map(lambda x: str(APTOS_IMAGE_DIR / f'{x}.png'))\n    return df[['source', 'image_id', 'original_path', 'label']]\n\n\ndef infer_ddr_columns(df: pd.DataFrame) -> tuple[str, str]:\n    image_candidates = ['image', 'image_id', 'id_code', 'filename', 'file_name', 'path', 'id']\n    label_candidates = ['label', 'diagnosis', 'grade', 'dr_grade', 'level', 'class']\n\n    lower_map = {str(c).lower(): c for c in df.columns}\n    image_col = next((lower_map[c] for c in image_candidates if c in lower_map), None)\n    label_col = next((lower_map[c] for c in label_candidates if c in lower_map), None)\n\n    if image_col is None:\n        image_col = df.columns[0]\n    if label_col is None:\n        if len(df.columns) < 2:\n            raise ValueError('DDR CSV cần ít nhất hai cột: tên ảnh và nhãn.')\n        label_col = df.columns[1]\n\n    return image_col, label_col\n\n\ndef load_ddr() -> pd.DataFrame:\n    require_path(DDR_CSV, 'DDR CSV')\n    require_path(DDR_IMAGE_ROOT, 'DDR image root')\n\n    raw = pd.read_csv(DDR_CSV)\n    image_col, label_col = infer_ddr_columns(raw)\n    print(f'DDR image column: {image_col!r} | label column: {label_col!r}')\n\n    image_index = build_image_index(DDR_IMAGE_ROOT)\n    rows = []\n    missing = []\n\n    for _, row in tqdm(raw.iterrows(), total=len(raw), desc='Ghép DDR CSV với ảnh'):\n        resolved = resolve_image_path(row[image_col], DDR_IMAGE_ROOT, image_index)\n        if resolved is None:\n            missing.append(str(row[image_col]))\n            continue\n        rows.append({\n            'source': 'DDR',\n            'image_id': resolved.stem,\n            'original_path': str(resolved),\n            'label': int(row[label_col]),\n        })\n\n    if missing:\n        sample = '\\n'.join(missing[:10])\n        raise FileNotFoundError(\n            f'Không tìm thấy {len(missing)} ảnh DDR. Ví dụ:\\n{sample}\\n'\n            'Kiểm tra DDR_CSV và DDR_IMAGE_ROOT.'\n        )\n    return pd.DataFrame(rows)\n\n\naptos_df = load_aptos()\nddr_df = load_ddr()\n\nfor name, frame in [('APTOS', aptos_df), ('DDR', ddr_df)]:\n    if not frame['label'].isin(range(NUM_CLASSES)).all():\n        bad = sorted(frame.loc[~frame['label'].isin(range(NUM_CLASSES)), 'label'].unique())\n        raise ValueError(f'{name} có nhãn ngoài 0..4: {bad}')\n\n    missing_paths = frame.loc[~frame['original_path'].map(lambda p: Path(p).exists())]\n    if len(missing_paths):\n        raise FileNotFoundError(f'{name} có {len(missing_paths)} đường dẫn ảnh không tồn tại.')\n\n    frame.drop_duplicates(subset=['source', 'image_id'], inplace=True)\n\nif MAX_IMAGES_PER_SOURCE is not None:\n    # Chỉ dùng để thử pipeline; vẫn lấy mẫu phân tầng theo nhãn.\n    def stratified_limit(df, n):\n        if len(df) <= n:\n            return df.copy()\n        parts = []\n        for _, g in df.groupby('label'):\n            k = max(2, round(n * len(g) / len(df)))\n            parts.append(g.sample(min(k, len(g)), random_state=SEED))\n        return pd.concat(parts).sample(frac=1, random_state=SEED).head(n).reset_index(drop=True)\n\n    aptos_df = stratified_limit(aptos_df, MAX_IMAGES_PER_SOURCE)\n    ddr_df = stratified_limit(ddr_df, MAX_IMAGES_PER_SOURCE)\n\nprint('APTOS:', len(aptos_df))\nprint(aptos_df['label'].value_counts().sort_index())\nprint('\\nDDR:', len(ddr_df))\nprint(ddr_df['label'].value_counts().sort_index())","metadata":{"execution":{"iopub.execute_input":"2026-06-24T05:55:59.428549Z","iopub.status.busy":"2026-06-24T05:55:59.428227Z","iopub.status.idle":"2026-06-24T05:56:09.623286Z","shell.execute_reply":"2026-06-24T05:56:09.621838Z","shell.execute_reply.started":"2026-06-24T05:55:59.428523Z"}},"outputs":[],"execution_count":null},{"id":"952b22dc","cell_type":"code","source":"# KIỂM TRA PHÂN BỐ DATASET SO VỚI TABLE 3 VÀ TABLE 6 CỦA BÀI BÁO\n\nPAPER_TABLE3 = {\n    'APTOS': {\n        'reported_total_images': 3662,\n        'reported_total_dr': 1483,  # Bài báo ghi vậy, nhưng tổng lớp 1..4 thực ra là 1857.\n        'class_counts': {0: 1805, 1: 370, 2: 999, 3: 193, 4: 295},\n    },\n    'DDR': {\n        'reported_total_images': 13673,\n        'reported_total_dr': 6256,\n        'class_counts': {0: 6266, 1: 630, 2: 4477, 3: 236, 4: 913},\n    },\n}\n\nPAPER_TABLE6_ORIGINAL = {\n    'APTOS': {0: 1805, 1: 370, 2: 999, 3: 193, 4: 295},\n    'DDR':   {0: 6372, 1: 708, 2: 725, 3: 352, 4: 453},\n}\n\nPAPER_TABLE6_REPORTED_TOTALS = {\n    'APTOS': {'original': 3662, 'augmented': 5363, 'balanced': 9025},\n    'DDR':   {'original': 8608, 'augmented': 23250, 'balanced': 31858},\n}\n\nDATASETS_TO_CHECK = {\n    'APTOS': aptos_df,\n    'DDR': ddr_df,\n}\n\nCLASS_LABELS = {\n    0: 'No DR',\n    1: 'Mild',\n    2: 'Moderate',\n    3: 'Severe',\n    4: 'PDR',\n}\n\n\ndef class_count_dict(frame: pd.DataFrame) -> dict[int, int]:\n    counts = frame['label'].value_counts().reindex(range(NUM_CLASSES), fill_value=0)\n    return {int(label): int(count) for label, count in counts.items()}\n\n\ndef distribution_distance(actual: dict[int, int], expected: dict[int, int]) -> int:\n    return int(sum(abs(actual.get(label, 0) - expected.get(label, 0))\n                   for label in range(NUM_CLASSES)))\n\n\ncomparison_rows = []\nsummary_rows = []\nintegrity_rows = []\n\nfor source, frame in DATASETS_TO_CHECK.items():\n    actual = class_count_dict(frame)\n    actual_total = int(len(frame))\n    actual_dr = int(sum(actual[label] for label in range(1, NUM_CLASSES)))\n\n    invalid_labels = sorted(\n        frame.loc[~frame['label'].isin(range(NUM_CLASSES)), 'label'].unique().tolist()\n    )\n    duplicate_ids = int(frame.duplicated(subset=['source', 'image_id']).sum())\n    missing_files = int((~frame['original_path'].map(lambda p: Path(p).exists())).sum())\n\n    integrity_rows.append({\n        'Dataset': source,\n        'Rows loaded': actual_total,\n        'Invalid labels': invalid_labels if invalid_labels else 'None',\n        'Duplicate image_id': duplicate_ids,\n        'Missing image files': missing_files,\n        'Integrity': 'PASS' if not invalid_labels and duplicate_ids == 0 and missing_files == 0 else 'CHECK',\n    })\n\n    references = {\n        'Table 3': PAPER_TABLE3[source]['class_counts'],\n        'Table 6 - Original': PAPER_TABLE6_ORIGINAL[source],\n    }\n\n    for reference_name, expected in references.items():\n        for label in range(NUM_CLASSES):\n            comparison_rows.append({\n                'Dataset': source,\n                'Reference': reference_name,\n                'Class': label,\n                'Class name': CLASS_LABELS[label],\n                'Paper count': int(expected[label]),\n                'Loaded count': int(actual[label]),\n                'Delta': int(actual[label] - expected[label]),\n                'Match': 'YES' if actual[label] == expected[label] else 'NO',\n            })\n\n    table3_class_sum = int(sum(PAPER_TABLE3[source]['class_counts'].values()))\n    table6_original_sum = int(sum(PAPER_TABLE6_ORIGINAL[source].values()))\n\n    summary_rows.append({\n        'Dataset': source,\n        'Loaded total': actual_total,\n        'Loaded DR (class 1..4)': actual_dr,\n        'Table 3 total images (reported)': PAPER_TABLE3[source]['reported_total_images'],\n        'Table 3 sum of class 0..4': table3_class_sum,\n        'Table 3 Total DR (reported)': PAPER_TABLE3[source]['reported_total_dr'],\n        'Table 6 Original total (reported)': PAPER_TABLE6_REPORTED_TOTALS[source]['original'],\n        'Table 6 Original sum from rows': table6_original_sum,\n        'Distance to Table 3': distribution_distance(\n            actual, PAPER_TABLE3[source]['class_counts']\n        ),\n        'Distance to Table 6': distribution_distance(\n            actual, PAPER_TABLE6_ORIGINAL[source]\n        ),\n    })\n\n\nintegrity_df = pd.DataFrame(integrity_rows)\ncomparison_df = pd.DataFrame(comparison_rows)\ndataset_summary_df = pd.DataFrame(summary_rows)\n\nprint('1) KIỂM TRA TÍNH TOÀN VẸN FILE/NHÃN')\ndisplay(integrity_df)\n\nprint('\\n2) SO SÁNH TỪNG LỚP VỚI TABLE 3 VÀ TABLE 6')\ndisplay(comparison_df)\n\nprint('\\n3) SO SÁNH TỔNG VÀ ĐỘ LỆCH PHÂN BỐ')\ndisplay(dataset_summary_df)\n\nprint('\\n4) KIỂM TRA SỐ HỌC NỘI BỘ CỦA TABLE 6')\ntable6_math_rows = []\nfor source, original_counts in PAPER_TABLE6_ORIGINAL.items():\n    original_sum = int(sum(original_counts.values()))\n    majority = int(max(original_counts.values()))\n    balanced_sum_from_rows = int(majority * NUM_CLASSES)\n\n    reported = PAPER_TABLE6_REPORTED_TOTALS[source]\n    generated_needed = int(sum(majority - original_counts[label]\n                               for label in range(NUM_CLASSES)))\n\n    table6_math_rows.append({\n        'Dataset': source,\n        'Original sum from class rows': original_sum,\n        'Original total reported': reported['original'],\n        'Augmented sum implied by rows': generated_needed,\n        'Augmented total reported': reported['augmented'],\n        'Balanced sum from class rows': balanced_sum_from_rows,\n        'Balanced total reported': reported['balanced'],\n        'Internal arithmetic': (\n            'PASS'\n            if original_sum == reported['original']\n            and generated_needed == reported['augmented']\n            and balanced_sum_from_rows == reported['balanced']\n            else 'INCONSISTENT'\n        ),\n    })\n\ntable6_math_df = pd.DataFrame(table6_math_rows)\ndisplay(table6_math_df)\n\nprint('\\n5) KẾT LUẬN TỰ ĐỘNG')\nfor source, frame in DATASETS_TO_CHECK.items():\n    actual = class_count_dict(frame)\n    d3 = distribution_distance(actual, PAPER_TABLE3[source]['class_counts'])\n    d6 = distribution_distance(actual, PAPER_TABLE6_ORIGINAL[source])\n\n    if d3 == 0 and d6 == 0:\n        verdict = 'khớp hoàn toàn cả Table 3 và Table 6.'\n    elif d3 == 0:\n        verdict = 'khớp Table 3 nhưng không khớp Table 6.'\n    elif d6 == 0:\n        verdict = 'khớp Table 6 nhưng không khớp Table 3.'\n    else:\n        closer = 'Table 3' if d3 < d6 else 'Table 6'\n        verdict = f'không khớp hoàn toàn; gần {closer} hơn (distance: T3={d3}, T6={d6}).'\n    print(f'- {source}: {verdict}')\n\nprint(\n    '\\nLưu ý: notebook chia dữ liệu gốc 70/15/15 trước rồi chỉ cân bằng train. '\n    'Vì vậy số train sau cân bằng sẽ không bằng tổng Balanced toàn dataset trong Table 6; '\n    'đây là chủ ý để tránh data leakage.'\n)\n","metadata":{"execution":{"iopub.execute_input":"2026-06-24T05:56:09.626194Z","iopub.status.busy":"2026-06-24T05:56:09.625714Z","iopub.status.idle":"2026-06-24T05:56:09.940602Z","shell.execute_reply":"2026-06-24T05:56:09.939414Z","shell.execute_reply.started":"2026-06-24T05:56:09.626137Z"}},"outputs":[],"execution_count":null},{"id":"fb519843","cell_type":"markdown","source":"## Chia 70/15/15 riêng cho từng dataset\n","metadata":{}},{"id":"834aab23","cell_type":"code","source":"def split_70_15_15(df: pd.DataFrame, seed: int) -> tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:\n    train_df, temp_df = train_test_split(\n        df,\n        test_size=0.30,\n        random_state=seed,\n        stratify=df['label'],\n    )\n    val_df, test_df = train_test_split(\n        temp_df,\n        test_size=0.50,\n        random_state=seed,\n        stratify=temp_df['label'],\n    )\n    return train_df.copy(), val_df.copy(), test_df.copy()\n\n\nsplits = []\nfor source_df in (aptos_df, ddr_df):\n    source = source_df['source'].iloc[0]\n    train_part, val_part, test_part = split_70_15_15(source_df, SEED)\n    train_part['split'] = 'train'\n    val_part['split'] = 'val'\n    test_part['split'] = 'test'\n    splits.extend([train_part, val_part, test_part])\n    print(source, {'train': len(train_part), 'val': len(val_part), 'test': len(test_part)})\n\nmanifest = pd.concat(splits, ignore_index=True)\n\n# Kiểm tra không rò rỉ ảnh giữa các split.\nfor source, group in manifest.groupby('source'):\n    split_sets = {\n        split: set(part['original_path'])\n        for split, part in group.groupby('split')\n    }\n    assert split_sets['train'].isdisjoint(split_sets['val'])\n    assert split_sets['train'].isdisjoint(split_sets['test'])\n    assert split_sets['val'].isdisjoint(split_sets['test'])\n\nmanifest.to_csv(MANIFEST_DIR / 'original_splits.csv', index=False)\n\nsummary = (\n    manifest.groupby(['source', 'split', 'label'])\n    .size()\n    .rename('count')\n    .reset_index()\n)\ndisplay(summary.pivot_table(index=['source', 'split'], columns='label', values='count', fill_value=0))","metadata":{"execution":{"iopub.execute_input":"2026-06-24T05:56:09.943621Z","iopub.status.busy":"2026-06-24T05:56:09.942856Z","iopub.status.idle":"2026-06-24T05:56:10.131232Z","shell.execute_reply":"2026-06-24T05:56:10.130291Z","shell.execute_reply.started":"2026-06-24T05:56:09.943573Z"}},"outputs":[],"execution_count":null},{"id":"da716f3d","cell_type":"markdown","source":"## ESRGAN theo Table 5: RRDB + PixelShuffle + RaGAN\n","metadata":{}},{"id":"58616da1","cell_type":"code","source":"def initialize_conv(module: nn.Module, scale: float = 1.0) -> None:\n    if isinstance(module, nn.Conv2d):\n        nn.init.kaiming_normal_(module.weight, a=0.2, mode='fan_in', nonlinearity='leaky_relu')\n        module.weight.data *= scale\n        if module.bias is not None:\n            nn.init.zeros_(module.bias)\n    elif isinstance(module, nn.Linear):\n        nn.init.kaiming_normal_(module.weight, a=0.2, mode='fan_in', nonlinearity='leaky_relu')\n        if module.bias is not None:\n            nn.init.zeros_(module.bias)\n\n\nclass ResidualDenseBlock5C(nn.Module):\n    def __init__(self, channels=64, growth_channels=32, residual_scale=0.2):\n        super().__init__()\n        self.residual_scale = residual_scale\n        self.conv1 = nn.Conv2d(channels, growth_channels, 3, 1, 1)\n        self.conv2 = nn.Conv2d(channels + growth_channels, growth_channels, 3, 1, 1)\n        self.conv3 = nn.Conv2d(channels + 2 * growth_channels, growth_channels, 3, 1, 1)\n        self.conv4 = nn.Conv2d(channels + 3 * growth_channels, growth_channels, 3, 1, 1)\n        self.conv5 = nn.Conv2d(channels + 4 * growth_channels, channels, 3, 1, 1)\n        self.activation = nn.LeakyReLU(0.2, inplace=True)\n\n        for layer in (self.conv1, self.conv2, self.conv3, self.conv4, self.conv5):\n            initialize_conv(layer, scale=0.1)\n\n    def forward(self, x):\n        x1 = self.activation(self.conv1(x))\n        x2 = self.activation(self.conv2(torch.cat([x, x1], dim=1)))\n        x3 = self.activation(self.conv3(torch.cat([x, x1, x2], dim=1)))\n        x4 = self.activation(self.conv4(torch.cat([x, x1, x2, x3], dim=1)))\n        x5 = self.conv5(torch.cat([x, x1, x2, x3, x4], dim=1))\n        return x + x5 * self.residual_scale\n\n\nclass RRDB(nn.Module):\n    def __init__(self, channels=64, growth_channels=32, residual_scale=0.2):\n        super().__init__()\n        self.residual_scale = residual_scale\n        self.rdb1 = ResidualDenseBlock5C(channels, growth_channels, residual_scale)\n        self.rdb2 = ResidualDenseBlock5C(channels, growth_channels, residual_scale)\n        self.rdb3 = ResidualDenseBlock5C(channels, growth_channels, residual_scale)\n\n    def forward(self, x):\n        out = self.rdb1(x)\n        out = self.rdb2(out)\n        out = self.rdb3(out)\n        return x + out * self.residual_scale\n\n\nclass PixelShuffleUpsample(nn.Module):\n    def __init__(self, channels=64):\n        super().__init__()\n        self.conv = nn.Conv2d(channels, channels * 4, 3, 1, 1)\n        self.shuffle = nn.PixelShuffle(2)\n        self.activation = nn.LeakyReLU(0.2, inplace=True)\n        initialize_conv(self.conv)\n\n    def forward(self, x):\n        return self.activation(self.shuffle(self.conv(x)))\n\n\nclass PaperESRGANGenerator(nn.Module):\n    def __init__(\n        self,\n        in_channels=3,\n        out_channels=3,\n        features=64,\n        num_rrdb=23,\n        growth_channels=32,\n        residual_scale=0.2,\n    ):\n        super().__init__()\n        self.conv_first = nn.Conv2d(in_channels, features, 3, 1, 1)\n        self.trunk = nn.Sequential(*[\n            RRDB(features, growth_channels, residual_scale)\n            for _ in range(num_rrdb)\n        ])\n        self.trunk_conv = nn.Conv2d(features, features, 3, 1, 1)\n        self.up1 = PixelShuffleUpsample(features)\n        self.up2 = PixelShuffleUpsample(features)\n        self.hr_conv = nn.Conv2d(features, features, 3, 1, 1)\n        self.conv_last = nn.Conv2d(features, out_channels, 3, 1, 1)\n        self.activation = nn.LeakyReLU(0.2, inplace=True)\n\n        for layer in (self.conv_first, self.trunk_conv, self.hr_conv, self.conv_last):\n            initialize_conv(layer)\n\n    def forward(self, x):\n        first = self.conv_first(x)\n        trunk = self.trunk_conv(self.trunk(first))\n        features = first + trunk\n        features = self.up1(features)\n        features = self.up2(features)\n        features = self.activation(self.hr_conv(features))\n        return self.conv_last(features)\n\n\nclass DiscriminatorBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, stride, use_bn=True):\n        super().__init__()\n        layers = [nn.Conv2d(in_channels, out_channels, 3, stride, 1)]\n        if use_bn:\n            layers.append(nn.BatchNorm2d(out_channels))\n        layers.append(nn.LeakyReLU(0.2, inplace=True))\n        self.block = nn.Sequential(*layers)\n        self.block.apply(initialize_conv)\n\n    def forward(self, x):\n        return self.block(x)\n\n\nclass PaperRaGANDiscriminator(nn.Module):\n    def __init__(self, in_channels=3):\n        super().__init__()\n        self.features = nn.Sequential(\n            DiscriminatorBlock(in_channels, 64, 1, use_bn=False),\n            DiscriminatorBlock(64, 64, 2),\n            DiscriminatorBlock(64, 128, 1),\n            DiscriminatorBlock(128, 128, 2),\n            DiscriminatorBlock(128, 256, 1),\n            DiscriminatorBlock(256, 256, 2),\n            DiscriminatorBlock(256, 512, 1),\n            DiscriminatorBlock(512, 512, 2),\n        )\n        self.pool = nn.AdaptiveAvgPool2d((4, 4))\n        self.classifier = nn.Sequential(\n            nn.Flatten(),\n            nn.Linear(512 * 4 * 4, 100),\n            nn.LeakyReLU(0.2, inplace=True),\n            nn.Linear(100, 1),\n        )\n        self.classifier.apply(initialize_conv)\n\n    def forward(self, x):\n        features = self.pool(self.features(x))\n        return self.classifier(features)  # logits; sigmoid nằm trong BCEWithLogitsLoss\n\n\nclass VGG19PerceptualExtractor(nn.Module):\n    def __init__(self):\n        super().__init__()\n        weights = models.VGG19_Weights.IMAGENET1K_V1\n        # [:35] kết thúc ở conv5_4 trước ReLU5_4 trong torchvision VGG-19.\n        self.features = models.vgg19(weights=weights).features[:35].eval()\n        for parameter in self.features.parameters():\n            parameter.requires_grad = False\n        self.register_buffer('mean', torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1))\n        self.register_buffer('std', torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1))\n\n    def forward(self, x_minus1_1):\n        x_0_1 = torch.clamp((x_minus1_1 + 1.0) / 2.0, 0.0, 1.0)\n        normalized = (x_0_1 - self.mean) / self.std\n        return self.features(normalized)\n\n\ndef maybe_data_parallel(model: nn.Module) -> nn.Module:\n    model = model.to(DEVICE)\n    if USE_MULTI_GPU_ESRGAN:\n        model = nn.DataParallel(model)\n    return model\n\n\ndef unwrap_model(model: nn.Module) -> nn.Module:\n    return model.module if isinstance(model, nn.DataParallel) else model\n\n\ndef set_requires_grad(model: nn.Module, enabled: bool) -> None:\n    for parameter in model.parameters():\n        parameter.requires_grad = enabled\n\n\ndef count_module_type(model: nn.Module, module_type) -> int:\n    return sum(isinstance(module, module_type) for module in unwrap_model(model).modules())\n\n\ngenerator = maybe_data_parallel(PaperESRGANGenerator(\n    features=ESRGAN_FEATURES,\n    num_rrdb=ESRGAN_NUM_RRDB,\n    growth_channels=ESRGAN_GROWTH_CHANNELS,\n    residual_scale=ESRGAN_RESIDUAL_SCALE,\n))\ndiscriminator = maybe_data_parallel(PaperRaGANDiscriminator())\n\nprint('Generator parameters:', f'{sum(p.numel() for p in generator.parameters()):,}')\nprint('Discriminator parameters:', f'{sum(p.numel() for p in discriminator.parameters()):,}')\nprint('RRDB blocks:', count_module_type(generator, RRDB))\nprint('PixelShuffle layers:', count_module_type(generator, nn.PixelShuffle))\nprint('BatchNorm layers trong generator:', count_module_type(generator, nn.BatchNorm2d))\n\nassert count_module_type(generator, RRDB) == 23\nassert count_module_type(generator, nn.PixelShuffle) == 2\nassert count_module_type(generator, nn.BatchNorm2d) == 0\n","metadata":{"execution":{"iopub.execute_input":"2026-06-24T05:56:10.134040Z","iopub.status.busy":"2026-06-24T05:56:10.133721Z","iopub.status.idle":"2026-06-24T05:56:11.237284Z","shell.execute_reply":"2026-06-24T05:56:11.236303Z","shell.execute_reply.started":"2026-06-24T05:56:10.134012Z"}},"outputs":[],"execution_count":null},{"id":"fce04f9d","cell_type":"markdown","source":"## Tiền xử lý trước ESRGAN và cache ảnh HR 224×224\n1. Crop nền đen và resize về `224×224`.\n2. Tách green channel.\n3. Morphological operators đa tỉ lệ theo Table 4.\n4. CLAHE.\n5. Gaussian blur.\n6. Nhân thành ba kênh và lưu làm ảnh HR tham chiếu cho ESRGAN.","metadata":{}},{"id":"113228b3","cell_type":"code","source":"def crop_fundus(image_rgb: np.ndarray, threshold: int = 7) -> np.ndarray:\n    gray = cv2.cvtColor(image_rgb, cv2.COLOR_RGB2GRAY)\n    mask = (gray > threshold).astype(np.uint8) * 255\n    contours, _ = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    if not contours:\n        return image_rgb\n\n    contour = max(contours, key=cv2.contourArea)\n    x, y, w, h = cv2.boundingRect(contour)\n    if w == 0 or h == 0:\n        return image_rgb\n    return image_rgb[y:y + h, x:x + w]\n\n\ndef opening_by_reconstruction(gray: np.ndarray, kernel_size: int) -> np.ndarray:\n    kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (kernel_size, kernel_size))\n    marker = cv2.erode(gray, kernel)\n    reconstructed = skimage_reconstruct(\n        marker.astype(np.float32) / 255.0,\n        gray.astype(np.float32) / 255.0,\n        method='dilation',\n    )\n    return np.rint(reconstructed * 255.0).astype(np.uint8)\n\n\ndef multiscale_max(images: list[np.ndarray]) -> np.ndarray:\n    return np.maximum.reduce(images)\n\n\ndef paper_morphology(green: np.ndarray) -> np.ndarray:\n    obr = multiscale_max([\n        opening_by_reconstruction(green, 3),\n        opening_by_reconstruction(green, 5),\n    ])\n\n    wth = multiscale_max([\n        cv2.morphologyEx(\n            green,\n            cv2.MORPH_TOPHAT,\n            cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3)),\n        ),\n        cv2.morphologyEx(\n            green,\n            cv2.MORPH_TOPHAT,\n            cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (7, 7)),\n        ),\n    ])\n\n    bth = multiscale_max([\n        cv2.morphologyEx(\n            green,\n            cv2.MORPH_BLACKHAT,\n            cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (5, 5)),\n        ),\n        cv2.morphologyEx(\n            green,\n            cv2.MORPH_BLACKHAT,\n            cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (9, 9)),\n        ),\n    ])\n\n    combined = obr.astype(np.int16) + wth.astype(np.int16) - bth.astype(np.int16)\n    return np.clip(combined, 0, 255).astype(np.uint8)\n\n\ndef preprocess_set1_before_esrgan(image_rgb: np.ndarray) -> np.ndarray:\n    image = crop_fundus(image_rgb)\n    image = cv2.resize(image, (IMAGE_SIZE, IMAGE_SIZE), interpolation=cv2.INTER_AREA)\n\n    green = image[:, :, 1]\n    morph = paper_morphology(green)\n\n    clahe = cv2.createCLAHE(\n        clipLimit=CLAHE_CLIP_LIMIT,\n        tileGridSize=CLAHE_TILE_GRID,\n    )\n    contrast = clahe.apply(morph)\n    denoised = cv2.GaussianBlur(contrast, GAUSSIAN_KERNEL, GAUSSIAN_SIGMA)\n    return cv2.merge([denoised, denoised, denoised])\n\n\ndef safe_cache_name(image_id: str) -> str:\n    keep = ''.join(ch if ch.isalnum() or ch in {'-', '_'} else '_' for ch in str(image_id))\n    return keep[:180] + '.png'\n\n\ndef cache_pre_esrgan_images(df: pd.DataFrame) -> pd.DataFrame:\n    output = df.copy()\n    cached_paths = []\n    failures = []\n\n    for row in tqdm(output.itertuples(index=False), total=len(output), desc='Cache trước ESRGAN'):\n        source_dir = PRE_ESRGAN_DIR / row.source\n        source_dir.mkdir(parents=True, exist_ok=True)\n        destination = source_dir / safe_cache_name(row.image_id)\n        cached_paths.append(str(destination))\n\n        if destination.exists() and destination.stat().st_size > 0:\n            continue\n\n        try:\n            bgr = cv2.imread(row.original_path, cv2.IMREAD_COLOR)\n            if bgr is None:\n                raise FileNotFoundError(f'OpenCV không đọc được ảnh: {row.original_path}')\n            rgb = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)\n            processed = preprocess_set1_before_esrgan(rgb)\n            ok = cv2.imwrite(str(destination), cv2.cvtColor(processed, cv2.COLOR_RGB2BGR))\n            if not ok:\n                raise IOError(f'cv2.imwrite trả về False: {destination}')\n        except Exception as exc:\n            failures.append((row.original_path, repr(exc)))\n\n    if failures:\n        report_path = WORK_DIR / 'pre_esrgan_failures.csv'\n        pd.DataFrame(failures, columns=['path', 'error']).to_csv(report_path, index=False)\n        raise RuntimeError(f'Có {len(failures)} ảnh lỗi. Xem: {report_path}')\n\n    output['pre_esrgan_path'] = cached_paths\n    return output\n\n\npre_esrgan_manifest_path = MANIFEST_DIR / 'pre_esrgan_splits.csv'\nif pre_esrgan_manifest_path.exists():\n    pre_esrgan_manifest = pd.read_csv(pre_esrgan_manifest_path)\n    cache_ok = pre_esrgan_manifest['pre_esrgan_path'].map(lambda p: Path(p).exists()).all()\n    if not cache_ok:\n        pre_esrgan_manifest = cache_pre_esrgan_images(manifest)\n        pre_esrgan_manifest.to_csv(pre_esrgan_manifest_path, index=False)\nelse:\n    pre_esrgan_manifest = cache_pre_esrgan_images(manifest)\n    pre_esrgan_manifest.to_csv(pre_esrgan_manifest_path, index=False)\n\nprint('Đã cache', len(pre_esrgan_manifest), 'ảnh trước ESRGAN.')\n","metadata":{"execution":{"iopub.execute_input":"2026-06-24T05:56:11.239152Z","iopub.status.busy":"2026-06-24T05:56:11.238391Z"}},"outputs":[],"execution_count":null},{"id":"1c58a456","cell_type":"markdown","source":"## Tạo cặp LR–HR cho ESRGAN\n\n- HR patch: `128×128`.\n- LR patch: `32×32` vì scale ×4.\n- LR được tạo bằng bicubic downsampling.\n- LR và HR đều được chuẩn hóa về `[-1,1]`.\n- ESRGAN chỉ học từ train split; validation split chỉ dùng theo dõi L1/PSNR.\n","metadata":{}},{"id":"2740223d","cell_type":"code","source":"def uint8_rgb_to_minus1_1_tensor(image: np.ndarray) -> torch.Tensor:\n    tensor = torch.from_numpy(np.ascontiguousarray(image)).permute(2, 0, 1).float()\n    return tensor / 127.5 - 1.0\n\n\nclass FundusSRPatchDataset(Dataset):\n    def __init__(self, dataframe: pd.DataFrame, random_crop: bool):\n        self.df = dataframe.reset_index(drop=True)\n        self.random_crop = random_crop\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        path = self.df.iloc[index]['pre_esrgan_path']\n        bgr = cv2.imread(path, cv2.IMREAD_COLOR)\n        if bgr is None:\n            raise FileNotFoundError(f'Không đọc được ảnh cache trước ESRGAN: {path}')\n        hr_full = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB)\n\n        height, width = hr_full.shape[:2]\n        if height < ESRGAN_HR_PATCH or width < ESRGAN_HR_PATCH:\n            hr_full = cv2.resize(\n                hr_full,\n                (max(width, ESRGAN_HR_PATCH), max(height, ESRGAN_HR_PATCH)),\n                interpolation=cv2.INTER_CUBIC,\n            )\n            height, width = hr_full.shape[:2]\n\n        if self.random_crop:\n            top = np.random.randint(0, height - ESRGAN_HR_PATCH + 1)\n            left = np.random.randint(0, width - ESRGAN_HR_PATCH + 1)\n        else:\n            top = (height - ESRGAN_HR_PATCH) // 2\n            left = (width - ESRGAN_HR_PATCH) // 2\n\n        hr_patch = hr_full[\n            top:top + ESRGAN_HR_PATCH,\n            left:left + ESRGAN_HR_PATCH,\n        ]\n        lr_patch = cv2.resize(\n            hr_patch,\n            (ESRGAN_LR_PATCH, ESRGAN_LR_PATCH),\n            interpolation=cv2.INTER_CUBIC,\n        )\n\n        return (\n            uint8_rgb_to_minus1_1_tensor(lr_patch),\n            uint8_rgb_to_minus1_1_tensor(hr_patch),\n        )\n\n\ndef seed_worker(worker_id: int) -> None:\n    worker_seed = (torch.initial_seed() + worker_id) % (2 ** 32)\n    np.random.seed(worker_seed)\n    random.seed(worker_seed)\n\n\nsr_train_df = pre_esrgan_manifest[pre_esrgan_manifest['split'] == 'train'].copy()\nsr_val_df = pre_esrgan_manifest[pre_esrgan_manifest['split'] == 'val'].copy()\nif len(sr_val_df) > ESRGAN_VAL_IMAGES:\n    sr_val_df = sr_val_df.sample(ESRGAN_VAL_IMAGES, random_state=SEED)\n\nsr_train_dataset = FundusSRPatchDataset(sr_train_df, random_crop=True)\nsr_val_dataset = FundusSRPatchDataset(sr_val_df, random_crop=False)\n\ngenerator_seed = torch.Generator()\ngenerator_seed.manual_seed(SEED)\n\nsr_loader_kwargs = dict(\n    batch_size=ESRGAN_BATCH_SIZE,\n    num_workers=NUM_WORKERS,\n    pin_memory=(DEVICE.type == 'cuda'),\n    persistent_workers=(NUM_WORKERS > 0),\n    worker_init_fn=seed_worker,\n)\n\nsr_train_loader = DataLoader(\n    sr_train_dataset,\n    shuffle=True,\n    drop_last=True,\n    generator=generator_seed,\n    **sr_loader_kwargs,\n)\nsr_val_loader = DataLoader(\n    sr_val_dataset,\n    shuffle=False,\n    drop_last=False,\n    **sr_loader_kwargs,\n)\n\nprint('SR train images:', len(sr_train_dataset))\nprint('SR validation images:', len(sr_val_dataset))\nprint('SR batches:', len(sr_train_loader), len(sr_val_loader))\nprint('LR patch:', ESRGAN_LR_PATCH, '→ HR patch:', ESRGAN_HR_PATCH)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"8bff453a","cell_type":"code","source":"# Hàm loss, checkpoint và đánh giá ESRGAN\npixel_criterion = nn.L1Loss()\nadversarial_criterion = nn.BCEWithLogitsLoss()\n\n\ndef relativistic_discriminator_loss(real_logits, fake_logits):\n    real_relative = real_logits - fake_logits.mean()\n    fake_relative = fake_logits - real_logits.mean()\n    real_loss = adversarial_criterion(real_relative, torch.ones_like(real_relative))\n    fake_loss = adversarial_criterion(fake_relative, torch.zeros_like(fake_relative))\n    return 0.5 * (real_loss + fake_loss)\n\n\ndef relativistic_generator_loss(real_logits, fake_logits):\n    real_relative = real_logits - fake_logits.mean()\n    fake_relative = fake_logits - real_logits.mean()\n    real_loss = adversarial_criterion(real_relative, torch.zeros_like(real_relative))\n    fake_loss = adversarial_criterion(fake_relative, torch.ones_like(fake_relative))\n    return 0.5 * (real_loss + fake_loss)\n\n\ndef save_esrgan_checkpoint(path: Path, stage: str, epoch: int, history: list[dict]):\n    payload = {\n        'stage': stage,\n        'epoch': epoch,\n        'generator_state_dict': unwrap_model(generator).state_dict(),\n        'discriminator_state_dict': unwrap_model(discriminator).state_dict(),\n        'history': history,\n        'config': {\n            'scale': ESRGAN_SCALE,\n            'num_rrdb': ESRGAN_NUM_RRDB,\n            'features': ESRGAN_FEATURES,\n            'growth_channels': ESRGAN_GROWTH_CHANNELS,\n            'hr_patch': ESRGAN_HR_PATCH,\n            'batch_size': ESRGAN_BATCH_SIZE,\n            'learning_rate': ESRGAN_LR,\n            'betas': ESRGAN_BETAS,\n            'pixel_weight': ESRGAN_PIXEL_WEIGHT,\n            'perceptual_weight': ESRGAN_PERCEPTUAL_WEIGHT,\n            'adversarial_weight': ESRGAN_ADVERSARIAL_WEIGHT,\n        },\n    }\n    torch.save(payload, path)\n\n\ndef load_esrgan_checkpoint(path: Path, load_discriminator: bool = True):\n    checkpoint = torch.load(path, map_location='cpu')\n    unwrap_model(generator).load_state_dict(checkpoint['generator_state_dict'], strict=True)\n    if load_discriminator and 'discriminator_state_dict' in checkpoint:\n        unwrap_model(discriminator).load_state_dict(\n            checkpoint['discriminator_state_dict'], strict=True\n        )\n    print(f\"Loaded ESRGAN checkpoint: {path} | stage={checkpoint.get('stage')} | epoch={checkpoint.get('epoch')}\")\n    return checkpoint\n\n\n@torch.inference_mode()\ndef evaluate_sr_generator(loader):\n    generator.eval()\n    total_l1 = 0.0\n    total_squared_error = 0.0\n    total_elements = 0\n    total_samples = 0\n\n    for lr_images, hr_images in tqdm(loader, desc='SR validation', leave=False):\n        lr_images = lr_images.to(DEVICE, non_blocking=True)\n        hr_images = hr_images.to(DEVICE, non_blocking=True)\n        with torch.autocast(device_type='cuda', enabled=USE_AMP):\n            sr_images = generator(lr_images)\n            l1 = F.l1_loss(sr_images, hr_images, reduction='sum')\n        error = (sr_images.float() - hr_images.float()) ** 2\n        total_l1 += l1.item()\n        total_squared_error += error.sum().item()\n        total_elements += error.numel()\n        total_samples += hr_images.size(0)\n\n    mean_l1 = total_l1 / max(total_elements, 1)\n    mse = total_squared_error / max(total_elements, 1)\n    psnr = 20.0 * np.log10(2.0 / np.sqrt(max(mse, 1e-12)))  # peak-to-peak của [-1,1] là 2\n    return {'val_l1': mean_l1, 'val_psnr': float(psnr), 'samples': total_samples}\n\n\ndef should_stop_step(step_index: int) -> bool:\n    return (\n        ESRGAN_MAX_STEPS_PER_EPOCH is not None\n        and step_index >= ESRGAN_MAX_STEPS_PER_EPOCH\n    )\n","metadata":{},"outputs":[],"execution_count":null},{"id":"61052ae7","cell_type":"markdown","source":"## Giai đoạn 1 — tải checkpoint pretrain L1 đã lưu\n\nVới `USE_EXTERNAL_ESRGAN_CHECKPOINTS=True`, cell dưới đây chỉ kiểm tra và tải checkpoint pretrain L1 từ Kaggle Models; không huấn luyện lại ESRGAN.\n","metadata":{}},{"id":"5be920bf","cell_type":"code","source":"pretrain_history = []\n\nif USE_EXTERNAL_ESRGAN_CHECKPOINTS:\n    if not EXTERNAL_ESRGAN_PRETRAIN_CHECKPOINT.exists():\n        raise FileNotFoundError(\n            'Không tìm thấy checkpoint pretrain L1: '\n            f'{EXTERNAL_ESRGAN_PRETRAIN_CHECKPOINT}'\n        )\n\n    checkpoint = load_esrgan_checkpoint(\n        EXTERNAL_ESRGAN_PRETRAIN_CHECKPOINT,\n        load_discriminator=False,\n    )\n    pretrain_history = checkpoint.get('history', [])\n\n    print(\n        'Bỏ qua huấn luyện pretrain. '\n        f\"Checkpoint stage={checkpoint.get('stage')} | \"\n        f\"epoch={checkpoint.get('epoch')}\"\n    )\n\nelif ESRGAN_GAN_CHECKPOINT.exists() and not FORCE_RETRAIN_ESRGAN:\n    print('Đã có checkpoint GAN cục bộ; bỏ qua giai đoạn pretrain.')\n    load_esrgan_checkpoint(\n        ESRGAN_GAN_CHECKPOINT,\n        load_discriminator=True,\n    )\n\nelif ESRGAN_PRETRAIN_CHECKPOINT.exists() and not FORCE_RETRAIN_ESRGAN:\n    print('Đã có checkpoint pretrain cục bộ; bỏ qua huấn luyện.')\n    checkpoint = load_esrgan_checkpoint(\n        ESRGAN_PRETRAIN_CHECKPOINT,\n        load_discriminator=False,\n    )\n    pretrain_history = checkpoint.get('history', [])\n\nelse:\n    optimizer_pretrain = torch.optim.Adam(\n        generator.parameters(),\n        lr=ESRGAN_LR,\n        betas=ESRGAN_BETAS,\n    )\n    scaler_pretrain = torch.amp.GradScaler('cuda', enabled=USE_AMP)\n\n    for epoch in range(1, ESRGAN_PRETRAIN_EPOCHS + 1):\n        generator.train()\n        running_loss = 0.0\n        seen = 0\n\n        progress = tqdm(\n            sr_train_loader,\n            desc=f'ESRGAN pretrain {epoch}/{ESRGAN_PRETRAIN_EPOCHS}',\n            mininterval=15,\n            miniters=50,\n            leave=False,\n        )\n\n        for step, (lr_images, hr_images) in enumerate(progress, start=1):\n            lr_images = lr_images.to(DEVICE, non_blocking=True)\n            hr_images = hr_images.to(DEVICE, non_blocking=True)\n            optimizer_pretrain.zero_grad(set_to_none=True)\n\n            with torch.autocast(device_type='cuda', enabled=USE_AMP):\n                sr_images = generator(lr_images)\n                loss = pixel_criterion(sr_images, hr_images)\n\n            scaler_pretrain.scale(loss).backward()\n            scaler_pretrain.step(optimizer_pretrain)\n            scaler_pretrain.update()\n\n            batch_size = lr_images.size(0)\n            running_loss += loss.item() * batch_size\n            seen += batch_size\n\n            if step % 50 == 0:\n                progress.set_postfix(l1=f'{loss.item():.4f}')\n\n            if should_stop_step(step):\n                break\n\n        validation = evaluate_sr_generator(sr_val_loader)\n        row = {\n            'stage': 'pretrain_l1',\n            'epoch': epoch,\n            'train_l1': running_loss / max(seen, 1),\n            **validation,\n        }\n        pretrain_history.append(row)\n        print(row)\n\n        save_esrgan_checkpoint(\n            ESRGAN_PRETRAIN_CHECKPOINT,\n            stage='pretrain_l1',\n            epoch=epoch,\n            history=pretrain_history,\n        )\n\nprint('Checkpoint pretrain đang dùng:', ACTIVE_ESRGAN_PRETRAIN_CHECKPOINT)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"3a026359","cell_type":"markdown","source":"## Giai đoạn 2 — tải checkpoint RaGAN đã lưu\n\nCheckpoint RaGAN từ Kaggle Models được tải trực tiếp. Phiên trước đã lưu đến epoch 8; bài báo không công bố số epoch ESRGAN nên notebook dùng checkpoint này làm mô hình SR cuối.\n","metadata":{}},{"id":"4510d572","cell_type":"code","source":"gan_history = []\n\nif USE_EXTERNAL_ESRGAN_CHECKPOINTS:\n    if not EXTERNAL_ESRGAN_GAN_CHECKPOINT.exists():\n        raise FileNotFoundError(\n            'Không tìm thấy checkpoint RaGAN: '\n            f'{EXTERNAL_ESRGAN_GAN_CHECKPOINT}'\n        )\n\n    checkpoint = load_esrgan_checkpoint(\n        EXTERNAL_ESRGAN_GAN_CHECKPOINT,\n        load_discriminator=True,\n    )\n    gan_history = checkpoint.get('history', [])\n\n    print(\n        'Bỏ qua huấn luyện RaGAN. '\n        f\"Checkpoint stage={checkpoint.get('stage')} | \"\n        f\"epoch={checkpoint.get('epoch')}\"\n    )\n\nelif ESRGAN_GAN_CHECKPOINT.exists() and not FORCE_RETRAIN_ESRGAN:\n    checkpoint = load_esrgan_checkpoint(\n        ESRGAN_GAN_CHECKPOINT,\n        load_discriminator=True,\n    )\n    gan_history = checkpoint.get('history', [])\n\nelse:\n    if ESRGAN_PRETRAIN_CHECKPOINT.exists():\n        load_esrgan_checkpoint(\n            ESRGAN_PRETRAIN_CHECKPOINT,\n            load_discriminator=False,\n        )\n    else:\n        raise FileNotFoundError(\n            'Chưa có pretrain checkpoint. Hãy chạy cell pretrain trước.'\n        )\n\n    perceptual_extractor = maybe_data_parallel(\n        VGG19PerceptualExtractor()\n    )\n    perceptual_extractor.eval()\n\n    optimizer_g = torch.optim.Adam(\n        generator.parameters(),\n        lr=ESRGAN_LR,\n        betas=ESRGAN_BETAS,\n    )\n    optimizer_d = torch.optim.Adam(\n        discriminator.parameters(),\n        lr=ESRGAN_LR,\n        betas=ESRGAN_BETAS,\n    )\n    scaler_g = torch.amp.GradScaler('cuda', enabled=USE_AMP)\n    scaler_d = torch.amp.GradScaler('cuda', enabled=USE_AMP)\n\n    for epoch in range(1, ESRGAN_GAN_EPOCHS + 1):\n        generator.train()\n        discriminator.train()\n        totals = defaultdict(float)\n        seen = 0\n\n        progress = tqdm(\n            sr_train_loader,\n            desc=f'ESRGAN RaGAN {epoch}/{ESRGAN_GAN_EPOCHS}',\n            mininterval=15,\n            miniters=50,\n            leave=False,\n        )\n\n        for step, (lr_images, hr_images) in enumerate(progress, start=1):\n            lr_images = lr_images.to(DEVICE, non_blocking=True)\n            hr_images = hr_images.to(DEVICE, non_blocking=True)\n\n            with torch.autocast(device_type='cuda', enabled=USE_AMP):\n                sr_images = generator(lr_images)\n\n            set_requires_grad(discriminator, True)\n            optimizer_d.zero_grad(set_to_none=True)\n\n            with torch.autocast(device_type='cuda', enabled=USE_AMP):\n                real_logits_d = discriminator(hr_images)\n                fake_logits_d = discriminator(sr_images.detach())\n                d_loss = relativistic_discriminator_loss(\n                    real_logits_d,\n                    fake_logits_d,\n                )\n\n            scaler_d.scale(d_loss).backward()\n            scaler_d.step(optimizer_d)\n            scaler_d.update()\n\n            set_requires_grad(discriminator, False)\n            optimizer_g.zero_grad(set_to_none=True)\n\n            with torch.autocast(device_type='cuda', enabled=USE_AMP):\n                pixel_loss = pixel_criterion(sr_images, hr_images)\n                sr_features = perceptual_extractor(sr_images)\n\n                with torch.no_grad():\n                    hr_features = perceptual_extractor(hr_images)\n                    real_logits_g = discriminator(hr_images)\n\n                fake_logits_g = discriminator(sr_images)\n                perceptual_loss = F.l1_loss(\n                    sr_features,\n                    hr_features,\n                )\n                adversarial_loss = relativistic_generator_loss(\n                    real_logits_g,\n                    fake_logits_g,\n                )\n                g_loss = (\n                    ESRGAN_PIXEL_WEIGHT * pixel_loss\n                    + ESRGAN_PERCEPTUAL_WEIGHT * perceptual_loss\n                    + ESRGAN_ADVERSARIAL_WEIGHT * adversarial_loss\n                )\n\n            scaler_g.scale(g_loss).backward()\n            scaler_g.step(optimizer_g)\n            scaler_g.update()\n            set_requires_grad(discriminator, True)\n\n            batch_size = lr_images.size(0)\n            seen += batch_size\n            totals['d_loss'] += d_loss.item() * batch_size\n            totals['g_loss'] += g_loss.item() * batch_size\n            totals['pixel'] += pixel_loss.item() * batch_size\n            totals['perceptual'] += perceptual_loss.item() * batch_size\n            totals['adversarial'] += adversarial_loss.item() * batch_size\n\n            if step % 50 == 0:\n                progress.set_postfix(\n                    D=f'{d_loss.item():.3f}',\n                    G=f'{g_loss.item():.3f}',\n                    L1=f'{pixel_loss.item():.3f}',\n                )\n\n            if should_stop_step(step):\n                break\n\n        validation = evaluate_sr_generator(sr_val_loader)\n        row = {\n            'stage': 'ragan',\n            'epoch': epoch,\n            **{\n                name: value / max(seen, 1)\n                for name, value in totals.items()\n            },\n            **validation,\n        }\n        gan_history.append(row)\n        print(row)\n\n        save_esrgan_checkpoint(\n            ESRGAN_GAN_CHECKPOINT,\n            stage='ragan',\n            epoch=epoch,\n            history=gan_history,\n        )\n\nprint('Checkpoint RaGAN đang dùng:', ACTIVE_ESRGAN_GAN_CHECKPOINT)\n","metadata":{},"outputs":[],"execution_count":null},{"id":"14dfbb03","cell_type":"markdown","source":"## Dùng ESRGAN đã lưu để tạo ảnh Set 1 theo batch\n\nMỗi ảnh tiền xử lý `224×224` được thu nhỏ bicubic xuống `56×56`, sau đó ESRGAN ×4 tạo lại ảnh `224×224`. Notebook xử lý nhiều ảnh cùng lúc để tận dụng T4 ×2.\n","metadata":{}},{"id":"520e8b8a","cell_type":"code","source":"if not ACTIVE_ESRGAN_GAN_CHECKPOINT.exists():\n    raise FileNotFoundError(\n        'Không tìm thấy checkpoint ESRGAN đang được chọn: '\n        f'{ACTIVE_ESRGAN_GAN_CHECKPOINT}'\n    )\n\nloaded_checkpoint = load_esrgan_checkpoint(\n    ACTIVE_ESRGAN_GAN_CHECKPOINT,\n    load_discriminator=False,\n)\ngenerator.eval()\n\nprint(\n    'ESRGAN inference checkpoint:',\n    ACTIVE_ESRGAN_GAN_CHECKPOINT,\n)\nprint(\n    'Stage/epoch:',\n    loaded_checkpoint.get('stage'),\n    loaded_checkpoint.get('epoch'),\n)\n\n\ndef processed_destination(source: str, image_id: str) -> Path:\n    source_dir = PREPROCESSED_DIR / source\n    source_dir.mkdir(parents=True, exist_ok=True)\n    return source_dir / safe_cache_name(image_id)\n\n\ndef preprocess_manifest_with_esrgan(\n    df: pd.DataFrame,\n) -> pd.DataFrame:\n    output = df.copy()\n\n    output['processed_path'] = [\n        str(processed_destination(row.source, row.image_id))\n        for row in output.itertuples(index=False)\n    ]\n\n    cached = output['processed_path'].map(\n        lambda p: Path(p).exists() and Path(p).stat().st_size > 0\n    )\n    pending = output.loc[~cached].reset_index(drop=True)\n\n    print(\n        f'Ảnh SR đã có: {int(cached.sum())}/{len(output)} | '\n        f'cần tạo: {len(pending)}'\n    )\n\n    if pending.empty:\n        return output\n\n    failures = []\n    batch_size = ESRGAN_INFERENCE_BATCH_SIZE\n\n    progress = tqdm(\n        total=len(pending),\n        desc=f'ESRGAN inference ×4, batch={batch_size}',\n        unit='img',\n        mininterval=10,\n    )\n\n    with torch.inference_mode():\n        for start in range(0, len(pending), batch_size):\n            batch_rows = pending.iloc[start:start + batch_size]\n\n            lr_tensors = []\n            destinations = []\n            source_paths = []\n\n            for row in batch_rows.itertuples(index=False):\n                try:\n                    bgr = cv2.imread(\n                        row.pre_esrgan_path,\n                        cv2.IMREAD_COLOR,\n                    )\n                    if bgr is None:\n                        raise FileNotFoundError(\n                            'Không đọc được ảnh trước ESRGAN: '\n                            f'{row.pre_esrgan_path}'\n                        )\n\n                    hr_reference = cv2.cvtColor(\n                        bgr,\n                        cv2.COLOR_BGR2RGB,\n                    )\n                    lr_size = IMAGE_SIZE // ESRGAN_SCALE\n                    lr_image = cv2.resize(\n                        hr_reference,\n                        (lr_size, lr_size),\n                        interpolation=cv2.INTER_CUBIC,\n                    )\n\n                    lr_tensors.append(\n                        uint8_rgb_to_minus1_1_tensor(lr_image)\n                    )\n                    destinations.append(Path(row.processed_path))\n                    source_paths.append(row.pre_esrgan_path)\n\n                except Exception as exc:\n                    failures.append(\n                        (row.pre_esrgan_path, repr(exc))\n                    )\n\n            if lr_tensors:\n                lr_batch = torch.stack(lr_tensors).to(\n                    DEVICE,\n                    non_blocking=True,\n                )\n\n                with torch.autocast(\n                    device_type='cuda',\n                    enabled=USE_AMP,\n                ):\n                    sr_batch = generator(lr_batch)\n\n                sr_batch = (\n                    ((sr_batch.float().clamp(-1.0, 1.0) + 1.0) * 127.5)\n                    .round()\n                    .to(torch.uint8)\n                    .permute(0, 2, 3, 1)\n                    .cpu()\n                    .numpy()\n                )\n\n                for sr_image, destination, source_path in zip(\n                    sr_batch,\n                    destinations,\n                    source_paths,\n                ):\n                    try:\n                        ok = cv2.imwrite(\n                            str(destination),\n                            cv2.cvtColor(\n                                sr_image,\n                                cv2.COLOR_RGB2BGR,\n                            ),\n                        )\n                        if not ok:\n                            raise IOError(\n                                f'cv2.imwrite trả về False: {destination}'\n                            )\n                    except Exception as exc:\n                        failures.append(\n                            (source_path, repr(exc))\n                        )\n\n                del lr_batch, sr_batch\n\n            progress.update(len(batch_rows))\n\n    progress.close()\n\n    if failures:\n        report_path = (\n            WORK_DIR / 'esrgan_inference_failures.csv'\n        )\n        pd.DataFrame(\n            failures,\n            columns=['path', 'error'],\n        ).to_csv(report_path, index=False)\n\n        raise RuntimeError(\n            f'Có {len(failures)} ảnh ESRGAN lỗi. '\n            f'Xem: {report_path}'\n        )\n\n    return output\n\n\nprocessed_manifest_path = (\n    MANIFEST_DIR / 'processed_splits_paper_esrgan.csv'\n)\n\nif processed_manifest_path.exists():\n    processed_manifest = pd.read_csv(processed_manifest_path)\n\n    all_cached = processed_manifest['processed_path'].map(\n        lambda p: Path(p).exists() and Path(p).stat().st_size > 0\n    ).all()\n\n    if not all_cached:\n        processed_manifest = preprocess_manifest_with_esrgan(\n            pre_esrgan_manifest\n        )\n        processed_manifest.to_csv(\n            processed_manifest_path,\n            index=False,\n        )\nelse:\n    processed_manifest = preprocess_manifest_with_esrgan(\n        pre_esrgan_manifest\n    )\n    processed_manifest.to_csv(\n        processed_manifest_path,\n        index=False,\n    )\n\nprint(\n    'Đã có',\n    len(processed_manifest),\n    'ảnh Set 1 bằng ESRGAN.',\n)\n\nsample_row = processed_manifest.sample(\n    1,\n    random_state=SEED,\n).iloc[0]\n\noriginal_bgr = cv2.imread(sample_row['original_path'])\npre_bgr = cv2.imread(sample_row['pre_esrgan_path'])\nsr_bgr = cv2.imread(sample_row['processed_path'])\n\nfig, axes = plt.subplots(1, 3, figsize=(14, 4))\n\naxes[0].imshow(\n    cv2.cvtColor(original_bgr, cv2.COLOR_BGR2RGB)\n)\naxes[0].set_title('Original')\n\naxes[1].imshow(\n    cv2.cvtColor(pre_bgr, cv2.COLOR_BGR2RGB)\n)\naxes[1].set_title('Morphology + CLAHE + Gaussian')\n\naxes[2].imshow(\n    cv2.cvtColor(sr_bgr, cv2.COLOR_BGR2RGB)\n)\naxes[2].set_title('Bicubic 56 → ESRGAN ×4 → 224')\n\nfor axis in axes:\n    axis.axis('off')\n\nplt.tight_layout()\nplt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"id":"f36988ea","cell_type":"markdown","source":"## Cân bằng train riêng cho APTOS và DDR","metadata":{}},{"id":"bef4f01b","cell_type":"code","source":"def balance_train_per_source(train_df: pd.DataFrame, seed: int = 42) -> pd.DataFrame:\n    rng = np.random.RandomState(seed)\n    balanced_parts = []\n\n    for source, source_df in train_df.groupby('source', sort=False):\n        counts = source_df['label'].value_counts().sort_index()\n        target = int(counts.max())\n        print(f'\\n{source}: target mỗi lớp = {target}')\n        print('Trước cân bằng:', counts.to_dict())\n\n        original = source_df.copy()\n        original['is_augmented'] = False\n        original['replica_id'] = 0\n        balanced_parts.append(original)\n\n        for label in range(NUM_CLASSES):\n            class_df = source_df[source_df['label'] == label]\n            if class_df.empty:\n                raise ValueError(f'{source} train không có lớp {label}.')\n            needed = target - len(class_df)\n            if needed <= 0:\n                continue\n\n            chosen_indices = rng.choice(class_df.index.to_numpy(), size=needed, replace=True)\n            synthetic = source_df.loc[chosen_indices].copy()\n            synthetic['is_augmented'] = True\n            synthetic['replica_id'] = np.arange(1, needed + 1)\n            balanced_parts.append(synthetic)\n\n    balanced = pd.concat(balanced_parts, ignore_index=True)\n    balanced = balanced.sample(frac=1.0, random_state=seed).reset_index(drop=True)\n    return balanced\n\n\ntrain_df = processed_manifest[processed_manifest['split'] == 'train'].copy()\nval_df = processed_manifest[processed_manifest['split'] == 'val'].copy()\ntest_df = processed_manifest[processed_manifest['split'] == 'test'].copy()\n\ntrain_balanced_df = balance_train_per_source(train_df, SEED)\ntrain_balanced_df.to_csv(MANIFEST_DIR / 'train_balanced.csv', index=False)\n\nprint('\\nSau cân bằng:')\ndisplay(\n    train_balanced_df.groupby(['source', 'label']).size()\n    .rename('count')\n    .unstack(fill_value=0)\n)\nprint(f'Train balanced: {len(train_balanced_df)} | Val: {len(val_df)} | Test: {len(test_df)}')","metadata":{},"outputs":[],"execution_count":null},{"id":"fad33caa","cell_type":"markdown","source":"## Augmentation và DataLoader\n\n- `rotation=±15°`\n- `width/height shift=±0.1`\n- `zoom=0.7..1.3` tương ứng zoom range `0.3`\n- `brightness factor=0.9..1.4`\n- `shear=±0.2 rad ≈ ±11.46°`\n- horizontal/vertical flip\n- contrast `±30%`\n- Gaussian noise\n- chuẩn hóa ảnh về `[-1,1]`\n","metadata":{}},{"id":"6e358955","cell_type":"code","source":"class RandomBrightnessFactor(A.ImageOnlyTransform):\n    def __init__(self, factor_range=(0.9, 1.4), p=0.5):\n        super().__init__(p=p)\n        self.factor_range = factor_range\n\n    def apply(self, img, **params):\n        factor = np.random.uniform(*self.factor_range)\n        return np.clip(img.astype(np.float32) * factor, 0, 255).astype(np.uint8)\n\n    def get_transform_init_args_names(self):\n        return ('factor_range',)\n\n\nclass RandomContrastFactor(A.ImageOnlyTransform):\n    def __init__(self, factor_range=(0.7, 1.3), p=0.5):\n        super().__init__(p=p)\n        self.factor_range = factor_range\n\n    def apply(self, img, **params):\n        factor = np.random.uniform(*self.factor_range)\n        mean = img.astype(np.float32).mean(axis=(0, 1), keepdims=True)\n        output = (img.astype(np.float32) - mean) * factor + mean\n        return np.clip(output, 0, 255).astype(np.uint8)\n\n    def get_transform_init_args_names(self):\n        return ('factor_range',)\n\n\nclass RandomGaussianNoise(A.ImageOnlyTransform):\n    def __init__(self, std_range=(3.0, 15.0), p=0.3):\n        super().__init__(p=p)\n        self.std_range = std_range\n\n    def apply(self, img, **params):\n        std = np.random.uniform(*self.std_range)\n        noise = np.random.normal(0.0, std, size=img.shape).astype(np.float32)\n        return np.clip(img.astype(np.float32) + noise, 0, 255).astype(np.uint8)\n\n    def get_transform_init_args_names(self):\n        return ('std_range',)\n\n\nnormalize_to_minus1_1 = [\n    A.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5), max_pixel_value=255.0),\n    ToTensorV2(),\n]\n\npaper_augmentation = A.Compose([\n    A.Affine(\n        scale=(0.7, 1.3),\n        translate_percent={'x': (-0.1, 0.1), 'y': (-0.1, 0.1)},\n        rotate=(-15, 15),\n        shear=(-11.46, 11.46),\n        interpolation=cv2.INTER_LINEAR,\n        border_mode=cv2.BORDER_REFLECT_101,\n        p=0.9,\n    ),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    RandomBrightnessFactor((0.9, 1.4), p=0.5),\n    RandomContrastFactor((0.7, 1.3), p=0.5),\n    RandomGaussianNoise((3.0, 15.0), p=0.3),\n    *normalize_to_minus1_1,\n])\n\nbase_transform = A.Compose(normalize_to_minus1_1)\n\n\nclass FundusDataset(Dataset):\n    def __init__(self, dataframe: pd.DataFrame, augmented_rows: bool):\n        self.df = dataframe.reset_index(drop=True)\n        self.augmented_rows = augmented_rows\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, index):\n        row = self.df.iloc[index]\n        path = row['processed_path']\n        image = cv2.imread(path, cv2.IMREAD_COLOR)\n        if image is None:\n            raise FileNotFoundError(f'Không đọc được ảnh cache: {path}')\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n\n        use_aug = self.augmented_rows and bool(row.get('is_augmented', False))\n        transform = paper_augmentation if use_aug else base_transform\n        tensor = transform(image=image)['image']\n        label = torch.tensor(int(row['label']), dtype=torch.long)\n        return tensor, label\n\n\ntrain_dataset = FundusDataset(train_balanced_df, augmented_rows=True)\nval_dataset = FundusDataset(val_df, augmented_rows=False)\ntest_dataset = FundusDataset(test_df, augmented_rows=False)\n\nloader_kwargs = dict(\n    batch_size=BATCH_SIZE,\n    num_workers=NUM_WORKERS,\n    pin_memory=(DEVICE.type == 'cuda'),\n    persistent_workers=(NUM_WORKERS > 0),\n)\n\ntrain_loader = DataLoader(train_dataset, shuffle=True, drop_last=True, **loader_kwargs)\nval_loader = DataLoader(val_dataset, shuffle=False, drop_last=False, **loader_kwargs)\ntest_loader = DataLoader(test_dataset, shuffle=False, drop_last=False, **loader_kwargs)\n\nprint('Batches:', len(train_loader), len(val_loader), len(test_loader))","metadata":{},"outputs":[],"execution_count":null},{"id":"e5a61aba","cell_type":"markdown","source":"## Modified ResNet-50 theo Figure 17\n","metadata":{}},{"id":"2e09fea8","cell_type":"code","source":"class ModifiedResNet50(nn.Module):\n    def __init__(self, num_classes=5, pretrained=True):\n        super().__init__()\n        weights = models.ResNet50_Weights.DEFAULT if pretrained else None\n        backbone = models.resnet50(weights=weights)\n        backbone.fc = nn.Identity()\n        self.backbone = backbone\n        self.activation = nn.ReLU(inplace=True)\n        self.batch_norm = nn.BatchNorm1d(2048)\n        self.classifier = nn.Linear(2048, num_classes)\n\n        nn.init.xavier_uniform_(self.classifier.weight)\n        nn.init.zeros_(self.classifier.bias)\n\n    def forward(self, x):\n        features = self.backbone(x)       # avg pool + flatten nằm trong torchvision ResNet\n        features = self.activation(features)\n        features = self.batch_norm(features)\n        logits = self.classifier(features)\n        return logits\n\n\nmodel = ModifiedResNet50(NUM_CLASSES, pretrained=True).to(DEVICE)\ntrainable = sum(p.numel() for p in model.parameters() if p.requires_grad)\ntotal = sum(p.numel() for p in model.parameters())\nprint(f'Parameters: {trainable:,} trainable / {total:,} total')","metadata":{},"outputs":[],"execution_count":null},{"id":"6b02555c","cell_type":"markdown","source":"## Training loop classifier\n\n- Adam\n- CrossEntropyLoss\n- StepLR(step size = 3, gamma = 0.1)\n- 25 epochs\n- lưu checkpoint có validation loss tốt nhất\n- mixed precision khi có GPU\n","metadata":{}},{"id":"cf3fcb88","cell_type":"code","source":"def run_epoch(model, loader, criterion, optimizer=None, scaler=None):\n    training = optimizer is not None\n    model.train(training)\n\n    total_loss = 0.0\n    total_correct = 0\n    total_samples = 0\n\n    progress = tqdm(loader, leave=False)\n    for images, labels in progress:\n        images = images.to(DEVICE, non_blocking=True)\n        labels = labels.to(DEVICE, non_blocking=True)\n\n        if training:\n            optimizer.zero_grad(set_to_none=True)\n\n        with torch.cuda.amp.autocast(enabled=USE_AMP):\n            logits = model(images)\n            loss = criterion(logits, labels)\n\n        if training:\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n\n        batch_size = images.size(0)\n        total_loss += loss.item() * batch_size\n        total_correct += (logits.argmax(dim=1) == labels).sum().item()\n        total_samples += batch_size\n        progress.set_postfix(loss=f'{loss.item():.4f}')\n\n    return total_loss / total_samples, total_correct / total_samples\n\n\ndef train_model(model, train_loader, val_loader):\n    criterion = nn.CrossEntropyLoss()\n    optimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)\n    scheduler = torch.optim.lr_scheduler.StepLR(\n        optimizer,\n        step_size=LR_STEP_SIZE,\n        gamma=LR_GAMMA,\n    )\n    scaler = torch.cuda.amp.GradScaler(enabled=USE_AMP)\n\n    history = []\n    best_val_loss = float('inf')\n\n    for epoch in range(1, EPOCHS + 1):\n        train_loss, train_acc = run_epoch(\n            model, train_loader, criterion, optimizer=optimizer, scaler=scaler\n        )\n        with torch.inference_mode():\n            val_loss, val_acc = run_epoch(model, val_loader, criterion)\n\n        lr_now = optimizer.param_groups[0]['lr']\n        row = {\n            'epoch': epoch,\n            'lr': lr_now,\n            'train_loss': train_loss,\n            'train_acc': train_acc,\n            'val_loss': val_loss,\n            'val_acc': val_acc,\n        }\n        history.append(row)\n        print(\n            f\"Epoch {epoch:02d}/{EPOCHS} | lr={lr_now:.2e} | \"\n            f\"train loss={train_loss:.4f}, acc={train_acc:.4f} | \"\n            f\"val loss={val_loss:.4f}, acc={val_acc:.4f}\"\n        )\n\n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'best_val_loss': best_val_loss,\n                'class_names': CLASS_NAMES,\n                'config': {\n                    'image_size': IMAGE_SIZE,\n                    'epochs': EPOCHS,\n                    'learning_rate': LEARNING_RATE,\n                    'batch_size': BATCH_SIZE,\n                    'lr_step_size': LR_STEP_SIZE,\n                    'lr_gamma': LR_GAMMA,\n                    'seed': SEED,\n                },\n            }, CHECKPOINT_PATH)\n            print('  → Saved best checkpoint')\n\n        scheduler.step()\n\n    history_df = pd.DataFrame(history)\n    history_df.to_csv(WORK_DIR / 'training_history.csv', index=False)\n    return history_df\n\n\nhistory_df = train_model(model, train_loader, val_loader)\nhistory_df.tail()","metadata":{},"outputs":[],"execution_count":null},{"id":"d785f0e6","cell_type":"markdown","source":"## Learning curves classifier\n","metadata":{}},{"id":"b0baee33","cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(12, 4))\naxes[0].plot(history_df['epoch'], history_df['train_loss'], label='Train')\naxes[0].plot(history_df['epoch'], history_df['val_loss'], label='Validation')\naxes[0].set_title('Loss')\naxes[0].set_xlabel('Epoch')\naxes[0].legend()\n\naxes[1].plot(history_df['epoch'], history_df['train_acc'], label='Train')\naxes[1].plot(history_df['epoch'], history_df['val_acc'], label='Validation')\naxes[1].set_title('Accuracy')\naxes[1].set_xlabel('Epoch')\naxes[1].legend()\n\nplt.tight_layout()\nplt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"e40d1cfe","cell_type":"markdown","source":"## Đánh giá test: tổng hợp và theo từng nguồn","metadata":{}},{"id":"5908294d","cell_type":"code","source":"checkpoint = torch.load(\n    CHECKPOINT_PATH,\n    map_location=DEVICE\n)\n\nmodel.load_state_dict(\n    checkpoint[\"model_state_dict\"]\n)\n\nmodel.eval()\n\nprint(\n    \"Loaded checkpoint from epoch:\",\n    checkpoint[\"epoch\"]\n)\n\n\n# =========================================================\n# 1. INFERENCE\n# =========================================================\n\ndef predict(\n    model,\n    loader\n):\n    y_true = []\n    y_pred = []\n    y_prob = []\n\n    with torch.inference_mode():\n        for images, labels in tqdm(\n            loader,\n            desc=\"Inference\",\n            leave=False\n        ):\n            images = images.to(\n                DEVICE,\n                non_blocking=True\n            )\n\n            logits = model(images)\n\n            probabilities = torch.softmax(\n                logits,\n                dim=1\n            )\n\n            predictions = probabilities.argmax(\n                dim=1\n            )\n\n            y_true.extend(\n                labels.cpu().numpy().tolist()\n            )\n\n            y_pred.extend(\n                predictions.cpu().numpy().tolist()\n            )\n\n            y_prob.extend(\n                probabilities.cpu().numpy().tolist()\n            )\n\n    return (\n        np.asarray(y_true, dtype=int),\n        np.asarray(y_pred, dtype=int),\n        np.asarray(y_prob, dtype=float)\n    )\n\n\n# =========================================================\n# 2. HÀM TÍNH BỘ METRIC THỐNG NHẤT\n# =========================================================\n\nEVAL_LABELS = [0, 1, 2, 3, 4]\n\nEVAL_CLASS_NAMES = [\n    \"No DR\",\n    \"Mild\",\n    \"Moderate\",\n    \"Severe\",\n    \"Proliferative\"\n]\n\n\ndef compute_standard_metrics(\n    y_true,\n    y_pred,\n    y_prob=None,\n    subset_name=\"APTOS + DDR\"\n):\n    y_true = np.asarray(\n        y_true\n    ).reshape(-1).astype(int)\n\n    y_pred = np.asarray(\n        y_pred\n    ).reshape(-1).astype(int)\n\n    if len(y_true) == 0:\n        raise ValueError(\n            f\"Subset {subset_name} không có mẫu để đánh giá.\"\n        )\n\n    # -----------------------------------------------------\n    # Confusion matrix\n    # -----------------------------------------------------\n\n    cm = confusion_matrix(\n        y_true,\n        y_pred,\n        labels=EVAL_LABELS\n    )\n\n    cm_normalized = confusion_matrix(\n        y_true,\n        y_pred,\n        labels=EVAL_LABELS,\n        normalize=\"true\"\n    )\n\n    # -----------------------------------------------------\n    # Precision, Recall, F1 và Support theo từng lớp\n    # -----------------------------------------------------\n\n    (\n        precision_class,\n        recall_class,\n        f1_class,\n        support_class\n    ) = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        labels=EVAL_LABELS,\n        average=None,\n        zero_division=0\n    )\n\n    (\n        precision_macro,\n        recall_macro,\n        f1_macro,\n        _\n    ) = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        labels=EVAL_LABELS,\n        average=\"macro\",\n        zero_division=0\n    )\n\n    (\n        precision_weighted,\n        recall_weighted,\n        f1_weighted,\n        _\n    ) = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        labels=EVAL_LABELS,\n        average=\"weighted\",\n        zero_division=0\n    )\n\n    # -----------------------------------------------------\n    # Specificity one-vs-rest\n    # Specificity = TN / (TN + FP)\n    # -----------------------------------------------------\n\n    specificity_class = []\n\n    for class_index in range(\n        len(EVAL_LABELS)\n    ):\n        true_positive = int(\n            cm[class_index, class_index]\n        )\n\n        false_negative = int(\n            cm[class_index, :].sum()\n            - true_positive\n        )\n\n        false_positive = int(\n            cm[:, class_index].sum()\n            - true_positive\n        )\n\n        true_negative = int(\n            cm.sum()\n            - true_positive\n            - false_negative\n            - false_positive\n        )\n\n        denominator = (\n            true_negative\n            + false_positive\n        )\n\n        specificity = (\n            true_negative / denominator\n            if denominator > 0\n            else 0.0\n        )\n\n        specificity_class.append(\n            specificity\n        )\n\n    specificity_class = np.asarray(\n        specificity_class,\n        dtype=float\n    )\n\n    specificity_macro = float(\n        np.mean(specificity_class)\n    )\n\n    specificity_weighted = float(\n        np.average(\n            specificity_class,\n            weights=support_class\n        )\n    )\n\n    # -----------------------------------------------------\n    # Overall metrics\n    # -----------------------------------------------------\n\n    accuracy = accuracy_score(\n        y_true,\n        y_pred\n    )\n\n    balanced_acc = balanced_accuracy_score(\n        y_true,\n        y_pred\n    )\n\n    mcc = matthews_corrcoef(\n        y_true,\n        y_pred\n    )\n\n    qwk = cohen_kappa_score(\n        y_true,\n        y_pred,\n        labels=EVAL_LABELS,\n        weights=\"quadratic\"\n    )\n\n    within_1_grade_acc = float(\n        np.mean(\n            np.abs(y_true - y_pred) <= 1\n        )\n    )\n\n    overall_metrics = {\n        \"Subset\": subset_name,\n        \"Support\": int(len(y_true)),\n        \"Accuracy\": accuracy,\n        \"BalancedAcc\": balanced_acc,\n        \"Precision Macro\": precision_macro,\n        \"Precision Weighted\": precision_weighted,\n        \"Recall Macro\": recall_macro,\n        \"Recall Weighted\": recall_weighted,\n        \"Specificity Macro\": specificity_macro,\n        \"Specificity Weighted\": specificity_weighted,\n        \"F1-Score Macro\": f1_macro,\n        \"F1-Score Weighted\": f1_weighted,\n        \"MCC\": mcc,\n        \"QWK\": qwk,\n        \"Within-1-Grade Acc\": within_1_grade_acc\n    }\n\n    # Giữ thêm AUC vì notebook có xác suất dự đoán\n    if y_prob is not None:\n        try:\n            y_prob = np.asarray(\n                y_prob,\n                dtype=float\n            )\n\n            y_true_one_hot = np.eye(\n                len(EVAL_LABELS)\n            )[y_true]\n\n            overall_metrics[\n                \"ROC-AUC Macro OVR\"\n            ] = roc_auc_score(\n                y_true_one_hot,\n                y_prob,\n                average=\"macro\",\n                multi_class=\"ovr\"\n            )\n\n            overall_metrics[\n                \"ROC-AUC Weighted OVR\"\n            ] = roc_auc_score(\n                y_true_one_hot,\n                y_prob,\n                average=\"weighted\",\n                multi_class=\"ovr\"\n            )\n\n        except Exception as error:\n            print(\n                f\"Không thể tính ROC-AUC cho \"\n                f\"{subset_name}: {error}\"\n            )\n\n            overall_metrics[\n                \"ROC-AUC Macro OVR\"\n            ] = np.nan\n\n            overall_metrics[\n                \"ROC-AUC Weighted OVR\"\n            ] = np.nan\n\n    class_metrics_df = pd.DataFrame(\n        {\n            \"Class\": EVAL_CLASS_NAMES,\n            \"Precision\": precision_class,\n            \"Recall\": recall_class,\n            \"Specificity\": specificity_class,\n            \"F1-Score\": f1_class,\n            \"Support\": support_class.astype(int)\n        }\n    )\n\n    cm_df = pd.DataFrame(\n        cm,\n        index=EVAL_CLASS_NAMES,\n        columns=EVAL_CLASS_NAMES\n    )\n\n    cm_df.index.name = \"Actual\"\n    cm_df.columns.name = \"Predicted\"\n\n    cm_normalized_df = pd.DataFrame(\n        cm_normalized,\n        index=EVAL_CLASS_NAMES,\n        columns=EVAL_CLASS_NAMES\n    )\n\n    cm_normalized_df.index.name = \"Actual\"\n    cm_normalized_df.columns.name = \"Predicted\"\n\n    return {\n        \"overall\": overall_metrics,\n        \"class_metrics\": class_metrics_df,\n        \"cm\": cm_df,\n        \"cm_normalized\": cm_normalized_df\n    }\n\n\n# =========================================================\n# 3. HÀM IN KẾT QUẢ ĐÚNG MẪU\n# =========================================================\n\ndef print_standard_report(\n    evaluation_result\n):\n    metrics = evaluation_result[\n        \"overall\"\n    ]\n\n    class_metrics_df = evaluation_result[\n        \"class_metrics\"\n    ]\n\n    print(\"\\n\" + \"=\" * 60)\n    print(\n        f\"FINAL EVALUATION: \"\n        f\"{metrics['Subset']}\"\n    )\n    print(\"=\" * 60)\n\n    print(\n        f\"{'Accuracy':<22}: \"\n        f\"{metrics['Accuracy']:.4f}\"\n    )\n\n    print(\n        f\"{'BalancedAcc':<22}: \"\n        f\"{metrics['BalancedAcc']:.4f}\"\n    )\n\n    print(\"-\" * 31)\n\n    print(\n        f\"{'Precision Macro':<22}: \"\n        f\"{metrics['Precision Macro']:.4f}\"\n    )\n\n    print(\n        f\"{'Precision Weighted':<22}: \"\n        f\"{metrics['Precision Weighted']:.4f}\"\n    )\n\n    print(\"-\" * 31)\n\n    print(\n        f\"{'Recall Macro':<22}: \"\n        f\"{metrics['Recall Macro']:.4f}\"\n    )\n\n    print(\n        f\"{'Recall Weighted':<22}: \"\n        f\"{metrics['Recall Weighted']:.4f}\"\n    )\n\n    print(\"-\" * 31)\n\n    print(\n        f\"{'Specificity Macro':<22}: \"\n        f\"{metrics['Specificity Macro']:.4f}\"\n    )\n\n    print(\n        f\"{'Specificity Weighted':<22}: \"\n        f\"{metrics['Specificity Weighted']:.4f}\"\n    )\n\n    print(\"-\" * 31)\n\n    print(\n        f\"{'F1-Score Macro':<22}: \"\n        f\"{metrics['F1-Score Macro']:.4f}\"\n    )\n\n    print(\n        f\"{'F1-Score Weighted':<22}: \"\n        f\"{metrics['F1-Score Weighted']:.4f}\"\n    )\n\n    print(\"-\" * 31)\n\n    print(\n        f\"{'MCC':<22}: \"\n        f\"{metrics['MCC']:.4f}\"\n    )\n\n    print(\n        f\"{'QWK':<22}: \"\n        f\"{metrics['QWK']:.4f}\"\n    )\n\n    print(\n        f\"{'Within-1-Grade Acc':<22}: \"\n        f\"{metrics['Within-1-Grade Acc']:.4f}\"\n    )\n\n    if \"ROC-AUC Macro OVR\" in metrics:\n        print(\"-\" * 31)\n\n        print(\n            f\"{'ROC-AUC Macro OVR':<22}: \"\n            f\"{metrics['ROC-AUC Macro OVR']:.4f}\"\n        )\n\n        print(\n            f\"{'ROC-AUC Weighted OVR':<22}: \"\n            f\"{metrics['ROC-AUC Weighted OVR']:.4f}\"\n        )\n\n    print(\"=\" * 31)\n\n    print(\"\\n--- CLASS-WISE METRICS ---\")\n\n    header = (\n        f\"{'Class':<15} | \"\n        f\"{'Precision':>9} | \"\n        f\"{'Recall':>6} | \"\n        f\"{'Specificity':>11} | \"\n        f\"{'F1-Score':>8} | \"\n        f\"{'Support':>7}\"\n    )\n\n    print(header)\n    print(\"-\" * len(header))\n\n    for _, row in class_metrics_df.iterrows():\n        print(\n            f\"{row['Class']:<15} | \"\n            f\"{row['Precision']:>9.4f} | \"\n            f\"{row['Recall']:>6.4f} | \"\n            f\"{row['Specificity']:>11.4f} | \"\n            f\"{row['F1-Score']:>8.4f} | \"\n            f\"{int(row['Support']):>7d}\"\n        )\n\n    print(\"=\" * len(header))\n\n\n# =========================================================\n# 4. HÀM VẼ CONFUSION MATRIX\n# =========================================================\n\ndef plot_confusion_matrix(\n    matrix_df,\n    title,\n    normalized=False\n):\n    matrix = matrix_df.to_numpy()\n\n    fig, ax = plt.subplots(\n        figsize=(8, 7)\n    )\n\n    image = ax.imshow(\n        matrix,\n        vmin=0,\n        vmax=1 if normalized else None\n    )\n\n    ax.set_title(\n        title\n    )\n\n    ax.set_xlabel(\n        \"Predicted\"\n    )\n\n    ax.set_ylabel(\n        \"Actual\"\n    )\n\n    ax.set_xticks(\n        range(len(EVAL_CLASS_NAMES))\n    )\n\n    ax.set_yticks(\n        range(len(EVAL_CLASS_NAMES))\n    )\n\n    ax.set_xticklabels(\n        EVAL_CLASS_NAMES,\n        rotation=45,\n        ha=\"right\"\n    )\n\n    ax.set_yticklabels(\n        EVAL_CLASS_NAMES\n    )\n\n    for row_index in range(\n        len(EVAL_CLASS_NAMES)\n    ):\n        for column_index in range(\n            len(EVAL_CLASS_NAMES)\n        ):\n            value = matrix[\n                row_index,\n                column_index\n            ]\n\n            text = (\n                f\"{value:.2f}\"\n                if normalized\n                else f\"{int(value)}\"\n            )\n\n            ax.text(\n                column_index,\n                row_index,\n                text,\n                ha=\"center\",\n                va=\"center\"\n            )\n\n    fig.colorbar(\n        image,\n        ax=ax,\n        fraction=0.046,\n        pad=0.04\n    )\n\n    plt.tight_layout()\n    plt.show()\n\n\n# =========================================================\n# 5. ĐÁNH GIÁ TỔNG HỢP APTOS + DDR\n# =========================================================\n\ny_true, y_pred, y_prob = predict(\n    model,\n    test_loader\n)\n\noverall_result = compute_standard_metrics(\n    y_true,\n    y_pred,\n    y_prob=y_prob,\n    subset_name=\"APTOS + DDR\"\n)\n\nprint_standard_report(\n    overall_result\n)\n\nprint(\n    \"\\n--- CONFUSION MATRIX: COUNTS ---\"\n)\n\ndisplay(\n    overall_result[\"cm\"]\n)\n\nprint(\n    \"\\n--- CONFUSION MATRIX: \"\n    \"NORMALIZED BY TRUE CLASS ---\"\n)\n\ndisplay(\n    overall_result[\n        \"cm_normalized\"\n    ].round(4)\n)\n\nplot_confusion_matrix(\n    overall_result[\"cm\"],\n    \"Set 1 — Confusion Matrix Counts (APTOS + DDR)\",\n    normalized=False\n)\n\nplot_confusion_matrix(\n    overall_result[\"cm_normalized\"],\n    \"Set 1 — Confusion Matrix Normalized (APTOS + DDR)\",\n    normalized=True\n)\n\n\n# =========================================================\n# 6. ĐÁNH GIÁ RIÊNG THEO NGUỒN APTOS VÀ DDR\n# =========================================================\n\ntest_metadata = test_df.reset_index(\n    drop=True\n)\n\nsource_results = []\nsource_class_metrics = {}\n\nfor source_name in [\n    \"APTOS\",\n    \"DDR\"\n]:\n    source_mask = (\n        test_metadata[\"source\"].to_numpy()\n        == source_name\n    )\n\n    source_result = compute_standard_metrics(\n        y_true[source_mask],\n        y_pred[source_mask],\n        y_prob=y_prob[source_mask],\n        subset_name=source_name\n    )\n\n    source_results.append(\n        source_result[\"overall\"]\n    )\n\n    source_class_metrics[\n        source_name\n    ] = source_result[\"class_metrics\"]\n\nsource_summary_df = pd.DataFrame(\n    [\n        overall_result[\"overall\"],\n        *source_results\n    ]\n)\n\nprint(\n    \"\\n--- SUMMARY BY DATA SOURCE ---\"\n)\n\ndisplay(\n    source_summary_df\n)\n\n\n# =========================================================\n# 7. LƯU KẾT QUẢ\n# =========================================================\n\noverall_metrics_df = pd.DataFrame(\n    {\n        \"Metric\": [\n            key\n            for key in overall_result[\n                \"overall\"\n            ].keys()\n            if key not in {\n                \"Subset\",\n                \"Support\"\n            }\n        ],\n        \"Value\": [\n            value\n            for key, value in overall_result[\n                \"overall\"\n            ].items()\n            if key not in {\n                \"Subset\",\n                \"Support\"\n            }\n        ]\n    }\n)\n\noverall_metrics_df.to_csv(\n    WORK_DIR\n    / \"test_overall_metrics.csv\",\n    index=False\n)\n\noverall_result[\n    \"class_metrics\"\n].to_csv(\n    WORK_DIR\n    / \"test_class_metrics.csv\",\n    index=False\n)\n\noverall_result[\n    \"cm\"\n].to_csv(\n    WORK_DIR\n    / \"test_confusion_matrix_counts.csv\"\n)\n\noverall_result[\n    \"cm_normalized\"\n].to_csv(\n    WORK_DIR\n    / \"test_confusion_matrix_normalized.csv\"\n)\n\nsource_summary_df.to_csv(\n    WORK_DIR\n    / \"test_metrics_by_source.csv\",\n    index=False\n)\n\nfor source_name, class_df in (\n    source_class_metrics.items()\n):\n    class_df.to_csv(\n        WORK_DIR\n        / (\n            f\"test_class_metrics_\"\n            f\"{source_name.lower()}.csv\"\n        ),\n        index=False\n    )\n\n\n# Giữ lại file dự đoán chi tiết như notebook gốc\npredictions_df = test_metadata[\n    [\n        \"source\",\n        \"image_id\",\n        \"original_path\",\n        \"processed_path\",\n        \"label\"\n    ]\n].copy()\n\npredictions_df[\n    \"prediction\"\n] = y_pred\n\nfor class_index, class_name in enumerate(\n    EVAL_CLASS_NAMES\n):\n    safe_name = (\n        class_name.lower()\n        .replace(\" \", \"_\")\n    )\n\n    predictions_df[\n        f\"prob_{class_index}_{safe_name}\"\n    ] = y_prob[:, class_index]\n\npredictions_df.to_csv(\n    WORK_DIR\n    / \"test_predictions.csv\",\n    index=False\n)\n\nprint(\n    \"\\nĐã lưu kết quả tại:\",\n    WORK_DIR\n)","metadata":{},"outputs":[],"execution_count":null}]}