{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"accelerator":"GPU","colab":{"gpuType":"T4","provenance":[]},"widgets":{"application/vnd.jupyter.widget-state+json":{"0d7caac0e72e40d58b76ce20cd9ab282":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_3382b64a8829456abeee979c1ade5381","placeholder":"​","style":"IPY_MODEL_cfcb029d87c949f08fde978bb48f8a7a","value":"[teacher_refresh:efficientnet_b0] Train 1/20:   0%"}},"1ec0785447224776aef260f1b68bfcab":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":"2","flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"20f171edfb664a45a3b04f2d84194fc3":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"3382b64a8829456abeee979c1ade5381":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"7267de59b4324f0ea43c3cb7f4caf0da":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"88680a90d4e3483bbe6bd8ab1e610062":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"","description":"","description_tooltip":null,"layout":"IPY_MODEL_1ec0785447224776aef260f1b68bfcab","max":68,"min":0,"orientation":"horizontal","style":"IPY_MODEL_20f171edfb664a45a3b04f2d84194fc3","value":0}},"97b4d95548624db8b45034e69c03d021":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_0d7caac0e72e40d58b76ce20cd9ab282","IPY_MODEL_88680a90d4e3483bbe6bd8ab1e610062","IPY_MODEL_beafa635d28b4403b62d6920991dd342"],"layout":"IPY_MODEL_9bd2f266c22047fe9cbc880774e8c821"}},"9bd2f266c22047fe9cbc880774e8c821":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":"inline-flex","flex":null,"flex_flow":"row wrap","grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":"100%"}},"a30da1a9529547bd89691a12aa11e34c":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"beafa635d28b4403b62d6920991dd342":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_7267de59b4324f0ea43c3cb7f4caf0da","placeholder":"​","style":"IPY_MODEL_a30da1a9529547bd89691a12aa11e34c","value":" 0/68 [00:00&lt;?, ?it/s]"}},"cfcb029d87c949f08fde978bb48f8a7a":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}}}},"phase_d_only_revision":{"description":"Removed dependencies on Phase C screening benchmark_summary.csv; fixed teacher efficientnet_b0 retrained in Phase D."},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":4117,"databundleVersionId":46665},{"sourceType":"datasetVersion","sourceId":16170535,"datasetId":10368284,"databundleVersionId":17146966},{"sourceType":"datasetVersion","sourceId":15822948,"datasetId":10142745,"databundleVersionId":16771899},{"sourceType":"datasetVersion","sourceId":16167391,"datasetId":10366338,"databundleVersionId":17143600}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"386152f3","cell_type":"markdown","source":"# Phase D — Knowledge Distillation\n","metadata":{"id":"386152f3"}},{"id":"SgFyzMyjFEWB","cell_type":"code","source":"from __future__ import annotations\n\nimport copy\nimport json\nimport math\nimport time\nimport warnings\nimport os\nimport sys\nfrom dataclasses import dataclass\nfrom pathlib import Path\nfrom typing import Dict, List, Optional, Sequence, Tuple\n\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport re\n\ntry:\n    import torch\n    import torch.nn as nn\n    import torch.nn.functional as F\n    from torch.utils.data import DataLoader, Dataset, WeightedRandomSampler\n    from torchvision import models, transforms\nexcept Exception as e:  # pragma: no cover\n    torch = None\n    nn = None\n    F = None\n    DataLoader = Dataset = WeightedRandomSampler = None\n    models = None\n    transforms = None\n    _TORCH_IMPORT_ERROR = e\nelse:\n    _TORCH_IMPORT_ERROR = None\n\nfrom sklearn.metrics import (\n    accuracy_score,\n    balanced_accuracy_score,\n    classification_report,\n    cohen_kappa_score,\n    confusion_matrix,\n    f1_score,\n    matthews_corrcoef,\n    precision_recall_fscore_support,\n)\nfrom sklearn.model_selection import train_test_split\n\ntry:\n    from tqdm.auto import tqdm\nexcept Exception:  # pragma: no cover\n    def tqdm(x, **kwargs):\n        return x\n\nIMAGENET_MEAN_RGB = np.array([0.485, 0.456, 0.406], dtype=np.float32)\nIMAGENET_STD_RGB = np.array([0.229, 0.224, 0.225], dtype=np.float32)\nGRAY_MEAN = np.array([0.5], dtype=np.float32)\nGRAY_STD = np.array([0.5], dtype=np.float32)\nIMAGE_PATH_CANDIDATES = [\"copied_image_path\", \"image_path\", \"image_path_used\"]\n\nTorchDatasetBase = Dataset if Dataset is not None else object\nTorchModuleBase = nn.Module if nn is not None else object\n\ndef ensure_dir(path: Path) -> Path:\n    path.mkdir(parents=True, exist_ok=True)\n    return path\n\ndef save_json(data: Dict, path: Path) -> Path:\n    ensure_dir(path.parent)\n    with path.open(\"w\", encoding=\"utf-8\") as f:\n        json.dump(data, f, ensure_ascii=False, indent=2)\n    return path\n\ndef require_torch():\n    if _TORCH_IMPORT_ERROR is not None:\n        raise RuntimeError(\n            \"Không import được torch/torchvision. \"\n            f\"Lỗi gốc: {_TORCH_IMPORT_ERROR}\"\n        )\n\ndef set_seed(seed: int = 42):\n    require_torch()\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\ndef infer_sample_id(df: pd.DataFrame) -> pd.Series:\n    if \"sample_id\" in df.columns:\n        return df[\"sample_id\"].astype(str)\n    if \"Id\" in df.columns:\n        return df[\"Id\"].astype(str)\n    if \"file_name\" in df.columns:\n        return df[\"file_name\"].astype(str).map(lambda x: Path(x).stem)\n    for col in IMAGE_PATH_CANDIDATES:\n        if col in df.columns:\n            return df[col].astype(str).map(lambda x: Path(x).stem)\n    raise ValueError(\n        \"Không tìm được cột để suy ra sample_id. Cần một trong: \"\n        \"sample_id, Id, file_name, image_path, copied_image_path, image_path_used.\"\n    )\n\ndef _normalize_sample_id(x: str) -> str:\n    \"\"\"Chuẩn hóa ID mẫu để khớp nhãn MS Malware với ảnh RGB đã convert.\"\"\"\n    x = str(x).strip()\n    if x == \"\" or x.lower() == \"nan\":\n        return \"\"\n\n    # Lấy stem để loại folder + extension trước.\n    # Ví dụ: /kaggle/input/.../01Iso..._rgb.png -> 01Iso..._rgb\n    x = Path(x).stem\n\n    # 01Iso..._rgb -> 01Iso... để khớp với cột Id trong trainLabels.csv.\n    for suffix in [\"_rgb\", \"_gray\", \"_greyscale\", \"_grayscale\"]:\n        if x.lower().endswith(suffix):\n            x = x[: -len(suffix)]\n            break\n\n    # Fallback cho các file gốc nếu có.\n    for ext in [\".bytes\", \".asm\", \".png\", \".jpg\", \".jpeg\", \".npy\", \".pt\"]:\n        if x.lower().endswith(ext):\n            x = x[: -len(ext)]\n    return x\n\ndef prepare_training_dataframe(\n    manifest_path,\n    labels_path,\n    dropna_class: bool = True,\n) -> Tuple[pd.DataFrame, List[int], Dict[int, int], Dict[int, int]]:\n    manifest_path = Path(manifest_path)\n    labels_path = Path(labels_path)\n\n    if not manifest_path.exists():\n        raise FileNotFoundError(f\"Manifest not found: {manifest_path}\")\n    manifest_df = pd.read_csv(manifest_path)\n    labels_df = pd.read_csv(labels_path) if labels_path.exists() else None\n\n    print(\"manifest_df shape:\", manifest_df.shape)\n    print(\"labels_df shape  :\", labels_df.shape if labels_df is not None else None)\n    print(\"manifest columns :\", manifest_df.columns.tolist())\n    print(\"labels columns   :\", labels_df.columns.tolist() if labels_df is not None else None)\n\n    if manifest_df.empty:\n        raise ValueError(\n            f\"Manifest is empty: {manifest_path}. \"\n            \"Hãy kiểm tra phase filtering hoặc fallback sang MANIFEST_PATH gốc.\"\n        )\n\n    if labels_df is not None:\n        labels_df = labels_df.rename(columns={\"Id\": \"sample_id\", \"Class\": \"class\"})\n\n        if \"sample_id\" not in labels_df.columns:\n            raise KeyError(\"labels_df must contain 'Id' or 'sample_id'\")\n        if \"class\" not in labels_df.columns:\n            raise KeyError(\"labels_df must contain 'Class' or 'class'\")\n\n        labels_df[\"sample_id\"] = labels_df[\"sample_id\"].astype(str).map(_normalize_sample_id)\n\n    manifest_id_candidates = [\"sample_id\", \"stem\", \"file_stem\", \"file_id\", \"id\", \"file_name\", \"image_path\", \"image_path_used\", \"copied_image_path\"]\n    manifest_id_col = None\n    for c in manifest_id_candidates:\n        if c in manifest_df.columns:\n            manifest_id_col = c\n            break\n\n    if manifest_id_col is None:\n        raise KeyError(\n            \"Không tìm thấy cột ID trong manifest. \"\n            \"Cần một trong các cột: sample_id, stem, file_stem, file_id, id, file_name, image_path, image_path_used, copied_image_path\"\n        )\n\n    manifest_df = manifest_df.copy()\n    manifest_df[\"sample_id\"] = manifest_df[manifest_id_col].astype(str).map(_normalize_sample_id)\n\n    # Tránh lỗi duplicate column: manifest từ class-subfolder có thể có cả \"class\" và \"Class\".\n    # Nếu rename trực tiếp Class -> class_manifest rồi lại rename class_manifest -> class,\n    # DataFrame sẽ có 2 cột tên \"class\"; khi đó df[\"class\"] trả về DataFrame, không có .unique().\n    if \"Class\" in manifest_df.columns:\n        manifest_df[\"class_manifest\"] = manifest_df[\"Class\"]\n        manifest_df = manifest_df.drop(columns=[\"Class\"])\n    if \"class\" in manifest_df.columns and \"class_manifest\" in manifest_df.columns:\n        # Ưu tiên nhãn lấy từ tên thư mục cha, lưu ở class_manifest.\n        manifest_df = manifest_df.drop(columns=[\"class\"])\n\n    # Nếu dataset có class-subfolders và CFG.USE_FOLDER_LABELS_IF_AVAILABLE=True,\n    # ưu tiên nhãn trong manifest thay vì merge trainLabels.csv.\n    # Điều này phù hợp với anhTanSuatDataSet: IMAGE_DIR/1 ... IMAGE_DIR/9.\n    use_folder_labels = bool(globals().get(\"CFG\", None) is not None and getattr(CFG, \"USE_FOLDER_LABELS_IF_AVAILABLE\", True))\n\n    if use_folder_labels and \"class_manifest\" in manifest_df.columns and manifest_df[\"class_manifest\"].notna().any():\n        df = manifest_df.rename(columns={\"class_manifest\": \"class\"}).copy()\n        print(\"Using folder labels from manifest/Class column; trainLabels.csv is not required for labels.\")\n    elif labels_df is not None:\n        df = manifest_df.merge(\n            labels_df[[\"sample_id\", \"class\"]],\n            on=\"sample_id\",\n            how=\"left\",\n            validate=\"many_to_one\"\n        )\n\n        if \"class_manifest\" in df.columns:\n            df[\"class\"] = df[\"class\"].fillna(df[\"class_manifest\"])\n            df = df.drop(columns=[\"class_manifest\"])\n    else:\n        if \"class_manifest\" in manifest_df.columns:\n            df = manifest_df.rename(columns={\"class_manifest\": \"class\"}).copy()\n        elif \"class\" in manifest_df.columns:\n            df = manifest_df.copy()\n        else:\n            raise FileNotFoundError(\n                f\"Labels not found: {labels_path}. Manifest cũng không có cột class/Class.\"\n            )\n\n    # Guard bổ sung: nếu vẫn còn duplicate column names, giữ cột đầu tiên không rỗng.\n    if df.columns.duplicated().any():\n        fixed_cols = {}\n        for col in pd.unique(df.columns):\n            same = df.loc[:, df.columns == col]\n            if same.shape[1] == 1:\n                fixed_cols[col] = same.iloc[:, 0]\n            else:\n                fixed_cols[col] = same.bfill(axis=1).iloc[:, 0]\n        df = pd.DataFrame(fixed_cols)\n\n    print(\"after merge shape:\", df.shape)\n    print(\"matched class rows:\", df[\"class\"].notna().sum())\n    print(\"unmatched rows    :\", df[\"class\"].isna().sum())\n\n    if dropna_class:\n        df = df.dropna(subset=[\"class\"]).copy()\n\n    if df.empty:\n        raise ValueError(\n            \"DataFrame is empty after merging manifest and labels. \"\n            \"Khả năng cao là sample_id giữa 2 file không khớp.\"\n        )\n\n    def _resolve_train_path(row):\n        for col in [\"copied_image_path\", \"image_path_used\", \"image_path\"]:\n            if col not in row.index:\n                continue\n            value = row[col]\n            if value is None or pd.isna(value):\n                continue\n            value = str(value).strip()\n            if value == \"\":\n                continue\n            if Path(value).exists():\n                return value\n        return None\n\n    df[\"image_path_used\"] = df.apply(_resolve_train_path, axis=1)\n    print(\"rows with valid image path:\", int(df[\"image_path_used\"].notna().sum()), \"/\", len(df))\n    df = df[df[\"image_path_used\"].notna()].copy()\n\n    if df.empty:\n        raise ValueError(\"DataFrame rỗng sau khi resolve image_path_used hợp lệ.\")\n\n    # Nếu class lấy từ tên thư mục là chuỗi, encode thành số ổn định.\n    try:\n        df[\"class\"] = df[\"class\"].astype(int)\n    except Exception:\n        df[\"class\"] = df[\"class\"].astype(str)\n\n    classes_sorted = sorted(df[\"class\"].unique().tolist())\n    class_to_idx = {cls: i for i, cls in enumerate(classes_sorted)}\n    idx_to_class = {i: cls for cls, i in class_to_idx.items()}\n    df[\"label\"] = df[\"class\"].map(class_to_idx)\n\n    df = df.reset_index(drop=True)\n\n    print(\"final df shape:\", df.shape)\n    print(\"label distribution:\")\n    print(df[\"label\"].value_counts().sort_index())\n\n    return df, classes_sorted, class_to_idx, idx_to_class\n\ndef create_stratified_splits(\n    df: pd.DataFrame,\n    test_size: float = 0.15,\n    val_size_from_train: float = 0.15,\n    seed: int = 42,\n) -> Tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:\n    if df is None or len(df) == 0:\n        raise ValueError(\"Input dataframe is empty before split.\")\n\n    if \"label\" not in df.columns:\n        raise KeyError(\"DataFrame must contain column 'label'.\")\n\n    class_counts = df[\"label\"].value_counts().sort_index()\n    print(\"class counts before split:\")\n    print(class_counts)\n\n    # class có < 2 mẫu thì stratify chắc chắn lỗi\n    too_small = class_counts[class_counts < 2]\n    if len(too_small) > 0:\n        raise ValueError(\n            \"Có class có ít hơn 2 samples, không thể stratify. \"\n            f\"Too small classes: {too_small.to_dict()}\"\n        )\n\n    train_val_df, test_df = train_test_split(\n        df,\n        test_size=test_size,\n        random_state=seed,\n        stratify=df[\"label\"],\n    )\n\n    train_counts = train_val_df[\"label\"].value_counts()\n    too_small_train = train_counts[train_counts < 2]\n    if len(too_small_train) > 0:\n        raise ValueError(\n            \"Sau split test, có class trong train_val còn < 2 mẫu nên không thể tiếp tục tách val. \"\n            f\"Problem classes: {too_small_train.to_dict()}\"\n        )\n\n    train_df, val_df = train_test_split(\n        train_val_df,\n        test_size=val_size_from_train,\n        random_state=seed,\n        stratify=train_val_df[\"label\"],\n    )\n\n    return (\n        train_df.reset_index(drop=True),\n        val_df.reset_index(drop=True),\n        test_df.reset_index(drop=True),\n    )\n\n\ndef create_group_stratified_splits(\n    df: pd.DataFrame,\n    test_size: float = 0.15,\n    val_size_from_train: float = 0.15,\n    seed: int = 42,\n    group_col: str = \"sample_id\",\n    label_col: str = \"label\",\n) -> Tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame]:\n    \"\"\"\n    Split chống data leak theo sample_id.\n\n    Khác với train_test_split theo từng dòng, hàm này tách theo group_col trước,\n    sau đó mới lấy toàn bộ các dòng thuộc cùng sample_id vào cùng một split.\n    Vì vậy nếu manifest có nhiều ảnh/biến thể/tensor cho cùng một malware sample,\n    chúng sẽ không bị rơi vào nhiều split khác nhau.\n    \"\"\"\n    if df is None or len(df) == 0:\n        raise ValueError(\"Input dataframe is empty before group split.\")\n    if label_col not in df.columns:\n        raise KeyError(f\"DataFrame must contain column '{label_col}'.\")\n\n    df = df.copy()\n    if group_col not in df.columns:\n        df[group_col] = infer_sample_id(df)\n    df[group_col] = df[group_col].astype(str).map(_normalize_sample_id)\n    df[label_col] = df[label_col].astype(int)\n\n    group_label_counts = df.groupby(group_col)[label_col].nunique()\n    ambiguous = group_label_counts[group_label_counts > 1]\n    if len(ambiguous) > 0:\n        raise ValueError(\n            \"Một số sample_id xuất hiện với nhiều label khác nhau. Không thể split an toàn. \"\n            f\"Ví dụ: {ambiguous.head(10).to_dict()}\"\n        )\n\n    group_df = (\n        df[[group_col, label_col]]\n        .drop_duplicates(subset=[group_col])\n        .reset_index(drop=True)\n    )\n\n    class_counts = group_df[label_col].value_counts().sort_index()\n    print(\"group-level class counts before split:\")\n    print(class_counts)\n\n    too_small = class_counts[class_counts < 2]\n    if len(too_small) > 0:\n        raise ValueError(\n            \"Có class có ít hơn 2 sample_id, không thể stratify an toàn. \"\n            f\"Too small classes: {too_small.to_dict()}\"\n        )\n\n    train_val_groups, test_groups = train_test_split(\n        group_df,\n        test_size=test_size,\n        random_state=seed,\n        stratify=group_df[label_col],\n    )\n\n    train_val_counts = train_val_groups[label_col].value_counts().sort_index()\n    too_small_train = train_val_counts[train_val_counts < 2]\n    if len(too_small_train) > 0:\n        raise ValueError(\n            \"Sau split test, có class trong train_val còn < 2 sample_id nên không thể tách val. \"\n            f\"Problem classes: {too_small_train.to_dict()}\"\n        )\n\n    train_groups, val_groups = train_test_split(\n        train_val_groups,\n        test_size=val_size_from_train,\n        random_state=seed,\n        stratify=train_val_groups[label_col],\n    )\n\n    train_ids = set(train_groups[group_col].astype(str))\n    val_ids = set(val_groups[group_col].astype(str))\n    test_ids = set(test_groups[group_col].astype(str))\n\n    train_df = df[df[group_col].isin(train_ids)].reset_index(drop=True)\n    val_df = df[df[group_col].isin(val_ids)].reset_index(drop=True)\n    test_df = df[df[group_col].isin(test_ids)].reset_index(drop=True)\n\n    report = leakage_report(train_df, val_df, test_df, group_col=group_col)\n    if any(v > 0 for v in report.values()):\n        raise RuntimeError(f\"Group split vẫn còn leakage: {report}\")\n\n    return train_df, val_df, test_df\n\ndef get_effective_num_class_weights(labels, num_classes, beta=0.9999):\n    labels = np.asarray(labels)\n    counts = np.bincount(labels, minlength=num_classes)\n    weights = []\n    for n_c in counts:\n        if n_c == 0:\n            weights.append(0.0)\n        else:\n            weights.append((1.0 - beta) / (1.0 - beta ** n_c))\n    weights = np.array(weights, dtype=np.float32)\n    weights = weights / max(weights.sum(), 1e-8) * num_classes\n    return weights, counts\n\ndef build_transforms(\n    img_size: int,\n    in_channels: int,\n    use_imagenet_norm: bool = True,\n    weak_aug: bool = False,\n):\n    require_torch()\n    train_ops = [transforms.Resize((img_size, img_size))]\n    if weak_aug:\n        train_ops.extend(\n            [\n                transforms.RandomHorizontalFlip(p=0.5),\n                transforms.RandomRotation(degrees=5),\n            ]\n        )\n    train_ops.append(transforms.ToTensor())\n    eval_ops = [transforms.Resize((img_size, img_size)), transforms.ToTensor()]\n\n    if in_channels == 1:\n        mean = GRAY_MEAN.tolist()\n        std = GRAY_STD.tolist()\n    else:\n        if use_imagenet_norm:\n            mean = IMAGENET_MEAN_RGB.tolist()\n            std = IMAGENET_STD_RGB.tolist()\n        else:\n            mean = [0.5] * in_channels\n            std = [0.5] * in_channels\n\n    norm = transforms.Normalize(mean=mean, std=std)\n    train_ops.append(norm)\n    eval_ops.append(norm)\n    return transforms.Compose(train_ops), transforms.Compose(eval_ops)\n\nclass MalwareImageDataset(TorchDatasetBase):\n    def __init__(self, dataframe: pd.DataFrame, transform=None, in_channels: int = 3):\n        require_torch()\n        self.df = dataframe.reset_index(drop=True).copy()\n        self.transform = transform\n        self.in_channels = in_channels\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        image_path = Path(row[\"image_path_used\"])\n        label = int(row[\"label\"])\n\n        with Image.open(image_path) as img:\n            if self.in_channels == 1:\n                image = img.convert(\"L\")\n            else:\n                image = img.convert(\"RGB\")\n\n        if self.transform is not None:\n            image = self.transform(image)\n        return image, label\n\n\ndef build_datasets(\n    train_df: pd.DataFrame,\n    val_df: pd.DataFrame,\n    test_df: pd.DataFrame,\n    train_transforms=None,\n    eval_transforms=None,\n    in_channels: int = 3,\n):\n    \"\"\"Tạo 3 dataset train/val/test từ dataframe đã chuẩn hóa.\"\"\"\n    train_dataset = MalwareImageDataset(\n        dataframe=train_df,\n        transform=train_transforms,\n        in_channels=in_channels,\n    )\n    val_dataset = MalwareImageDataset(\n        dataframe=val_df,\n        transform=eval_transforms,\n        in_channels=in_channels,\n    )\n    test_dataset = MalwareImageDataset(\n        dataframe=test_df,\n        transform=eval_transforms,\n        in_channels=in_channels,\n    )\n    return train_dataset, val_dataset, test_dataset\n\nclass FocalLoss(TorchModuleBase):\n    def __init__(self, alpha=None, gamma=2.0, reduction=\"mean\"):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n\n    def forward(self, logits, targets):\n        ce = F.cross_entropy(logits, targets, weight=self.alpha, reduction=\"none\")\n        pt = torch.exp(-ce)\n        loss = ((1 - pt) ** self.gamma) * ce\n        if self.reduction == \"mean\":\n            return loss.mean()\n        if self.reduction == \"sum\":\n            return loss.sum()\n        return loss\n\ndef build_criterion(loss_name: str, class_weights=None):\n    loss_name = loss_name.lower()\n    if loss_name == \"ce\":\n        return nn.CrossEntropyLoss(weight=class_weights)\n    if loss_name == \"focal\":\n        return FocalLoss(alpha=class_weights, gamma=2.0)\n    raise ValueError(\"loss_name phải là 'ce' hoặc 'focal'\")\n\ndef compute_metrics(y_true, y_pred, num_classes):\n    y_true = np.asarray(y_true)\n    y_pred = np.asarray(y_pred)\n    labels = list(range(num_classes))\n\n    accuracy = accuracy_score(y_true, y_pred)\n    macro_f1 = f1_score(y_true, y_pred, average=\"macro\", labels=labels, zero_division=0)\n    weighted_f1 = f1_score(y_true, y_pred, average=\"weighted\", labels=labels, zero_division=0)\n    micro_f1 = f1_score(y_true, y_pred, average=\"micro\", labels=labels, zero_division=0)\n    bal_acc = balanced_accuracy_score(y_true, y_pred)\n\n    per_precision, per_recall, per_f1, support = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        average=None,\n        labels=labels,\n        zero_division=0,\n    )\n    macro_precision, macro_recall, _, _ = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        average=\"macro\",\n        labels=labels,\n        zero_division=0,\n    )\n    weighted_precision, weighted_recall, _, _ = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        average=\"weighted\",\n        labels=labels,\n        zero_division=0,\n    )\n\n    try:\n        mcc = matthews_corrcoef(y_true, y_pred)\n    except Exception:\n        mcc = float(\"nan\")\n    try:\n        kappa = cohen_kappa_score(y_true, y_pred, labels=labels)\n    except Exception:\n        kappa = float(\"nan\")\n\n    return {\n        \"accuracy\": float(accuracy),\n        \"macro_precision\": float(macro_precision),\n        \"macro_recall\": float(macro_recall),\n        \"macro_f1\": float(macro_f1),\n        \"weighted_precision\": float(weighted_precision),\n        \"weighted_recall\": float(weighted_recall),\n        \"weighted_f1\": float(weighted_f1),\n        \"micro_f1\": float(micro_f1),\n        \"balanced_acc\": float(bal_acc),\n        \"mcc\": float(mcc),\n        \"cohen_kappa\": float(kappa),\n        \"per_class_precision\": per_precision,\n        \"per_class_recall\": per_recall,\n        \"per_class_f1\": per_f1,\n        \"per_class_support\": support,\n    }\n\ndef _forward_logits(model, images):\n    logits = model(images)\n    if isinstance(logits, tuple):\n        logits = logits[0]\n    if hasattr(logits, \"logits\"):\n        logits = logits.logits\n    return logits\n\ndef count_parameters(model) -> int:\n    return int(sum(p.numel() for p in model.parameters() if p.requires_grad))\n\ndef count_total_parameters(model) -> int:\n    return int(sum(p.numel() for p in model.parameters()))\n\n\ndef estimate_model_footprint_mb(model) -> float:\n    \"\"\"Ước lượng bộ nhớ tĩnh của tham số + buffer theo dtype hiện tại.\"\"\"\n    param_bytes = sum(p.numel() * p.element_size() for p in model.parameters())\n    buffer_bytes = sum(b.numel() * b.element_size() for b in model.buffers())\n    return float((param_bytes + buffer_bytes) / (1024 ** 2))\n\n\ndef _process_rss_mb() -> float:\n    \"\"\"RSS hiện tại của process Python. Dùng được cả CPU-only runtime.\"\"\"\n    try:\n        import psutil  # Colab thường có sẵn; nếu không có sẽ fallback.\n        return float(psutil.Process(os.getpid()).memory_info().rss / (1024 ** 2))\n    except Exception:\n        return float(\"nan\")\n\n\ndef _process_peak_rss_mb() -> float:\n    \"\"\"Peak RSS của process. Trên Linux, ru_maxrss trả về KiB; trên macOS là bytes.\"\"\"\n    try:\n        import resource\n        peak = float(resource.getrusage(resource.RUSAGE_SELF).ru_maxrss)\n        if sys.platform == \"darwin\":\n            return peak / (1024 ** 2)\n        return peak / 1024.0\n    except Exception:\n        return float(\"nan\")\n\n\ndef _move_images_only_to_device(images, runtime: Dict):\n    device = runtime.get(\"device\", getattr(CFG, \"DEVICE\", \"cpu\"))\n    images = images.to(device, non_blocking=bool(runtime.get(\"non_blocking\", False)))\n    if bool(runtime.get(\"channels_last\", False)) and images.ndim == 4:\n        images = images.to(memory_format=torch.channels_last)\n    return images\n\n\n@torch.no_grad()\ndef estimate_forward_flops(\n    model,\n    sample_images,\n    runtime: Optional[Dict] = None,\n    use_amp: bool = False,\n    flop_batch_size: int = 1,\n) -> Dict[str, object]:\n    \"\"\"\n    Ước lượng FLOPs cho một forward pass bằng torch.profiler.\n\n    Ghi chú:\n    - torch.profiler chủ yếu đếm FLOPs cho conv/matmul/mm/addmm nên có thể thấp hơn FLOPs tuyệt đối.\n    - Metric quan trọng để so sánh công bằng là dùng cùng hàm đo, cùng input shape, cùng batch size.\n    \"\"\"\n    if sample_images is None:\n        return {}\n\n    runtime = runtime or {\"device\": getattr(CFG, \"DEVICE\", \"cpu\"), \"channels_last\": False, \"device_type\": \"cpu\", \"use_amp\": False}\n    device = runtime.get(\"device\", getattr(CFG, \"DEVICE\", \"cpu\"))\n    device_type = \"cuda\" if str(device).startswith(\"cuda\") and torch.cuda.is_available() else \"cpu\"\n    amp_enabled = bool(use_amp) and device_type == \"cuda\"\n\n    was_training = model.training\n    model.eval()\n\n    try:\n        bs = max(1, min(int(flop_batch_size), int(sample_images.size(0))))\n        sample_images = sample_images[:bs]\n        sample_images = _move_images_only_to_device(sample_images, runtime)\n\n        activities = [torch.profiler.ProfilerActivity.CPU]\n        if device_type == \"cuda\":\n            activities.append(torch.profiler.ProfilerActivity.CUDA)\n\n        # Warm-up một lượt ngoài profiler để tránh overhead khởi động kernel.\n        with torch.amp.autocast(device_type, enabled=amp_enabled):\n            _ = _forward_logits(model, sample_images)\n        if device_type == \"cuda\":\n            torch.cuda.synchronize()\n\n        with torch.profiler.profile(activities=activities, with_flops=True, record_shapes=False, profile_memory=False) as prof:\n            with torch.amp.autocast(device_type, enabled=amp_enabled):\n                _ = _forward_logits(model, sample_images)\n            if device_type == \"cuda\":\n                torch.cuda.synchronize()\n\n        total_flops = 0\n        for evt in prof.key_averages():\n            total_flops += int(getattr(evt, \"flops\", 0) or 0)\n\n        flops_per_batch = float(total_flops)\n        flops_per_img = float(total_flops / max(1, bs))\n        return {\n            \"flops_method\": \"torch.profiler.with_flops\",\n            \"flops_eval_batch_size\": int(bs),\n            \"forward_flops_per_batch\": flops_per_batch,\n            \"forward_flops_per_img\": flops_per_img,\n            \"forward_mflops_per_img\": flops_per_img / 1e6,\n            \"forward_gflops_per_img\": flops_per_img / 1e9,\n            \"flops_note\": \"Profiler FLOPs are counted mainly for conv/matmul ops; use as a consistent comparative metric.\",\n        }\n    except Exception as e:\n        return {\n            \"flops_method\": \"torch.profiler.with_flops\",\n            \"forward_flops_per_img\": float(\"nan\"),\n            \"forward_mflops_per_img\": float(\"nan\"),\n            \"forward_gflops_per_img\": float(\"nan\"),\n            \"flops_error\": str(e),\n        }\n    finally:\n        if was_training:\n            model.train()\n\n\n@torch.no_grad()\ndef measure_inference_efficiency(\n    model,\n    loader,\n    runtime: Optional[Dict] = None,\n    use_amp: bool = False,\n    max_batches: int = 20,\n    warmup_batches: int = 2,\n    measure_flops: bool = True,\n    flop_batch_size: int = 1,\n) -> Dict[str, object]:\n    \"\"\"\n    Đo hiệu suất suy luận trên một số batch của test_loader.\n\n    Metrics trả về:\n    - latency: ms/image và ms/batch\n    - throughput: images/s\n    - peak memory footprint khi inference: CUDA allocated/reserved hoặc CPU RSS\n    - FLOPs/GFLOPs mỗi ảnh nếu bật CFG.MEASURE_FLOPS\n    \"\"\"\n    if loader is None:\n        return {}\n\n    runtime = runtime or {\"device\": getattr(CFG, \"DEVICE\", \"cpu\"), \"channels_last\": False, \"device_type\": \"cpu\", \"use_amp\": False}\n    device = runtime.get(\"device\", getattr(CFG, \"DEVICE\", \"cpu\"))\n    device_type = \"cuda\" if str(device).startswith(\"cuda\") and torch.cuda.is_available() else \"cpu\"\n    amp_enabled = bool(use_amp) and device_type == \"cuda\"\n\n    was_training = model.training\n    model.eval()\n\n    first_images_for_flops = None\n\n    # Warm-up để phép đo GPU ổn định hơn.\n    for batch_idx, (images, targets) in enumerate(loader):\n        if batch_idx == 0:\n            first_images_for_flops = images.detach().cpu() if hasattr(images, \"detach\") else images\n        if batch_idx >= max(0, int(warmup_batches)):\n            break\n        images, targets = move_batch_to_device(images, targets, runtime)\n        with torch.amp.autocast(device_type, enabled=amp_enabled):\n            _ = _forward_logits(model, images)\n\n    if device_type == \"cuda\":\n        torch.cuda.synchronize()\n        torch.cuda.reset_peak_memory_stats()\n\n    rss_before_mb = _process_rss_mb()\n    peak_rss_before_mb = _process_peak_rss_mb()\n\n    total_images = 0\n    total_batches = 0\n    t0 = time.perf_counter()\n    for batch_idx, (images, targets) in enumerate(loader):\n        if batch_idx >= max(1, int(max_batches)):\n            break\n        if first_images_for_flops is None:\n            first_images_for_flops = images.detach().cpu() if hasattr(images, \"detach\") else images\n        images, targets = move_batch_to_device(images, targets, runtime)\n        with torch.amp.autocast(device_type, enabled=amp_enabled):\n            _ = _forward_logits(model, images)\n        total_images += int(images.size(0))\n        total_batches += 1\n\n    if device_type == \"cuda\":\n        torch.cuda.synchronize()\n    elapsed = max(time.perf_counter() - t0, 1e-9)\n\n    rss_after_mb = _process_rss_mb()\n    peak_rss_after_mb = _process_peak_rss_mb()\n    inference_peak_cpu_rss_mb = peak_rss_after_mb\n    inference_rss_delta_mb = (rss_after_mb - rss_before_mb) if not math.isnan(rss_after_mb) and not math.isnan(rss_before_mb) else float(\"nan\")\n\n    if device_type == \"cuda\":\n        inference_peak_cuda_allocated_mb = float(torch.cuda.max_memory_allocated() / (1024 ** 2))\n        inference_peak_cuda_reserved_mb = float(torch.cuda.max_memory_reserved() / (1024 ** 2))\n        peak_memory_footprint_mb = inference_peak_cuda_allocated_mb\n        peak_memory_footprint_source = \"cuda_max_memory_allocated\"\n    else:\n        inference_peak_cuda_allocated_mb = float(\"nan\")\n        inference_peak_cuda_reserved_mb = float(\"nan\")\n        peak_memory_footprint_mb = inference_peak_cpu_rss_mb\n        peak_memory_footprint_source = \"process_peak_rss\"\n\n    throughput = total_images / elapsed if total_images > 0 else 0.0\n    latency_img_ms = 1000.0 * elapsed / total_images if total_images > 0 else float(\"nan\")\n    latency_batch_ms = 1000.0 * elapsed / total_batches if total_batches > 0 else float(\"nan\")\n\n    flops_metrics = {}\n    if bool(measure_flops):\n        flops_metrics = estimate_forward_flops(\n            model=model,\n            sample_images=first_images_for_flops,\n            runtime=runtime,\n            use_amp=use_amp,\n            flop_batch_size=flop_batch_size,\n        )\n\n    if was_training:\n        model.train()\n\n    return {\n        \"inference_eval_batches\": int(total_batches),\n        \"inference_eval_images\": int(total_images),\n        \"inference_time_sec\": float(elapsed),\n        \"inference_throughput_img_s\": float(throughput),\n        \"inference_latency_ms_per_img\": float(latency_img_ms),\n        \"inference_latency_ms_per_batch\": float(latency_batch_ms),\n        \"inference_peak_memory_footprint_mb\": float(peak_memory_footprint_mb),\n        \"inference_peak_memory_footprint_source\": peak_memory_footprint_source,\n        \"inference_peak_cpu_rss_mb\": float(inference_peak_cpu_rss_mb),\n        \"inference_rss_before_mb\": float(rss_before_mb),\n        \"inference_rss_after_mb\": float(rss_after_mb),\n        \"inference_rss_delta_mb\": float(inference_rss_delta_mb),\n        \"inference_peak_cuda_allocated_mb\": float(inference_peak_cuda_allocated_mb),\n        \"inference_peak_cuda_reserved_mb\": float(inference_peak_cuda_reserved_mb),\n        **flops_metrics,\n    }\n\n\ndef _ensure_sample_id_column(df: pd.DataFrame, split_name: str = \"split\") -> pd.DataFrame:\n    df = df.copy()\n    if \"sample_id\" not in df.columns:\n        df[\"sample_id\"] = infer_sample_id(df)\n    df[\"sample_id\"] = df[\"sample_id\"].astype(str).map(_normalize_sample_id)\n    if df[\"sample_id\"].isna().any() or (df[\"sample_id\"].astype(str).str.len() == 0).any():\n        raise ValueError(f\"{split_name} có sample_id rỗng/không hợp lệ.\")\n    return df\n\n\ndef enforce_no_data_leakage(\n    train_df: pd.DataFrame,\n    val_df: pd.DataFrame,\n    test_df: pd.DataFrame,\n    group_col: str = \"sample_id\",\n    fail_fast: bool = True,\n) -> Dict[str, int]:\n    report = leakage_report(train_df, val_df, test_df, group_col=group_col)\n    leaked = {k: v for k, v in report.items() if int(v) > 0}\n    if leaked and fail_fast:\n        raise RuntimeError(\n            \"Phát hiện data leak giữa train/val/test theo sample_id. \"\n            f\"Chi tiết: {leaked}\"\n        )\n    return report\n\n\ndef leakage_report(\n    train_df: pd.DataFrame,\n    val_df: pd.DataFrame,\n    test_df: pd.DataFrame,\n    group_col: str = \"sample_id\",\n) -> Dict[str, int]:\n    for name, split_df in [(\"train\", train_df), (\"val\", val_df), (\"test\", test_df)]:\n        if group_col not in split_df.columns:\n            raise ValueError(f\"{name}_df thiếu cột {group_col}\")\n\n    train_ids = set(train_df[group_col].astype(str).map(_normalize_sample_id))\n    val_ids = set(val_df[group_col].astype(str).map(_normalize_sample_id))\n    test_ids = set(test_df[group_col].astype(str).map(_normalize_sample_id))\n\n    report = {\n        \"train_val_overlap\": int(len(train_ids & val_ids)),\n        \"train_test_overlap\": int(len(train_ids & test_ids)),\n        \"val_test_overlap\": int(len(val_ids & test_ids)),\n    }\n    return report\n\n","metadata":{"id":"SgFyzMyjFEWB","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T10:56:58.502050Z","iopub.execute_input":"2026-05-09T10:56:58.502482Z","iopub.status.idle":"2026-05-09T10:57:13.042481Z","shell.execute_reply.started":"2026-05-09T10:56:58.502449Z","shell.execute_reply":"2026-05-09T10:57:13.041894Z"}},"outputs":[],"execution_count":null},{"id":"0ba938eb","cell_type":"code","source":"# ============================================================\n# Cell vai trò:\n# - Khai báo cấu hình dự án ở mức toàn cục.\n# - Phiên bản này đã chỉnh để chạy trực tiếp trên Kaggle.\n# - Dataset ảnh đang có cấu trúc class-subfolder:\n#   /kaggle/input/datasets/dungngu/anhtansuatdataset/1, /2, ..., /9\n# - trainLabels.csv được trỏ chính xác tới /kaggle/input/competitions/malware-classification/trainLabels.csv; ảnh lấy từ /kaggle/input/datasets/dungngu/anhtansuatdataset.\n# ============================================================\n\n# ----------------------------\n# Kaggle path helpers\n# ----------------------------\ndef _running_on_kaggle() -> bool:\n    return Path(\"/kaggle/input\").exists()\n\n\ndef _first_existing_path(candidates: Sequence[Path]) -> Optional[Path]:\n    for p in candidates:\n        if Path(p).exists():\n            return Path(p)\n    return None\n\n\ndef _find_kaggle_image_dir() -> Path:\n    \"\"\"Đường dẫn dataset ảnh trên Kaggle. Có thể override bằng MALWARE_IMAGE_DIR.\"\"\"\n    env_path = os.environ.get(\"MALWARE_IMAGE_DIR\", \"\").strip()\n    if env_path:\n        return Path(env_path)\n    return Path(\"/kaggle/input/datasets/dungngu/anhtansuatdataset\")\n\n\ndef _find_kaggle_labels_path() -> Path:\n    \"\"\"Đường dẫn trainLabels.csv trên Kaggle. Có thể override bằng MALWARE_LABELS_PATH.\"\"\"\n    env_path = os.environ.get(\"MALWARE_LABELS_PATH\", \"\").strip()\n    if env_path:\n        return Path(env_path)\n    return Path(\"/kaggle/input/competitions/malware-classification/trainLabels.csv\")\n\n\ndef _default_project_root() -> Path:\n    env_path = os.environ.get(\"MALWARE_PROJECT_ROOT\", \"\").strip()\n    if env_path:\n        return Path(env_path)\n    if _running_on_kaggle():\n        return Path(\"/kaggle/working/malware-prediction\")\n    return Path(\"/content/drive/MyDrive/TDTT/malware-prediction\")\n\n\n@dataclass\nclass ProjectConfig:\n    \"\"\"\n    Cấu hình trung tâm của project.\n\n    Lưu ý Kaggle:\n    - /kaggle/input là read-only, chỉ dùng để đọc ảnh và nhãn.\n    - /kaggle/working là nơi ghi manifest, split, checkpoint, report.\n    \"\"\"\n    # ----------------------------\n    # Project paths\n    # ----------------------------\n    IS_KAGGLE: bool = _running_on_kaggle()\n    PROJECT_ROOT: Path = _default_project_root()\n\n    # Thư mục chứa ảnh. Dataset hiện tại có cấu trúc class-subfolder: IMAGE_DIR/1 ... IMAGE_DIR/9.\n    IMAGE_DIR: Path = _find_kaggle_image_dir()\n\n    # DATA_DIR/INPUT_DIR chỉ để tham chiếu input; không ghi artifact vào đây.\n    DATA_DIR: Path = IMAGE_DIR.parent\n    INPUT_DIR: Path = IMAGE_DIR\n    LABELS_PATH: Path = _find_kaggle_labels_path()\n\n    # Dataset anhTanSuat đang có sẵn nhãn ở tên thư mục cha 1..9.\n    # Đặt True để tránh phụ thuộc vào việc tên ảnh có match Id trong trainLabels.csv hay không.\n    USE_FOLDER_LABELS_IF_AVAILABLE: bool = True\n\n    # ----------------------------\n    # Output folders: phải nằm trong /kaggle/working khi chạy Kaggle\n    # ----------------------------\n    CONVERTED_DIR: Path = PROJECT_ROOT / \"artifacts\" / \"converted\"   # chứa manifest.csv sinh từ folder ảnh\n    FILTERED_DIR: Path = PROJECT_ROOT / \"artifacts\" / \"filtered\"\n    REPORT_DIR: Path = PROJECT_ROOT / \"artifacts\" / \"reports\"\n    BENCHMARK_DIR: Path = PROJECT_ROOT / \"artifacts\" / \"benchmark\"\n\n    # ----------------------------\n    # Phase B conversion config\n    # ----------------------------\n    MODE: str = \"rgb\"\n    FIXED_WIDTH: int | None = None\n    RESIZE_TO: int | None = 256\n    ENTROPY_WINDOW: int = 16\n    SAVE_MASK: bool = False\n    SAVE_TENSOR: bool = False\n\n    # Vì ảnh đã được convert sẵn trên Kaggle nên không chạy lại Phase B.\n    RUN_PHASE_B_BATCH_CONVERSION: bool = False\n    RUN_PHASE_B_REPRESENTATION_ABLATION: bool = False\n\n    # Tự tạo lại manifest từ folder ảnh mỗi lần chạy để tránh dùng nhầm path cũ.\n    FORCE_REBUILD_MANIFEST: bool = True\n\n    # ----------------------------\n    # Phase A / Phase B filtering\n    # ----------------------------\n    MIN_KEEP_PER_CLASS: int = 20\n    HEAD_KEEP_QUANTILE: float = 0.50\n    MEDIUM_KEEP_QUANTILE: float = 0.35\n    TAIL_KEEP_QUANTILE: float = 0.15\n    ABSOLUTE_MISSING_RATIO_MAX: float = 0.98\n\n    # ----------------------------\n    # Phase D training config\n    # ----------------------------\n    IMG_SIZE: int = 224\n    IN_CHANNELS: int = 3\n    BATCH_SIZE: int = 64\n    EPOCHS_SCREENING: int = 12\n    LR: float = 1e-3\n    WEIGHT_DECAY: float = 1e-4\n    LOSS_NAME: str = \"ce\"          # \"ce\" hoặc \"focal\"\n    USE_CLASS_WEIGHTS_IN_LOSS: bool = True\n    USE_WEIGHTED_SAMPLER: bool = True\n    NUM_WORKERS: int = 2           # Kaggle ổn định hơn với 2; tăng lên 4 nếu runtime chịu được\n    PRETRAINED: bool = True\n    WEAK_AUG: bool = False\n    EARLY_STOP_PATIENCE: int = 3\n    TEST_SIZE: float = 0.15\n    VAL_SIZE_FROM_TRAIN: float = 0.1765\n    SEED: int = 42\n\n    # ----------------------------\n    # GPU / runtime config\n    # ----------------------------\n    DEVICE: str = \"cuda\" if (torch is not None and torch.cuda.is_available()) else \"cpu\"\n    USE_AMP: bool = True\n    PIN_MEMORY: bool = True\n    PERSISTENT_WORKERS: bool = True\n    PREFETCH_FACTOR: int = 2\n    CUDNN_BENCHMARK: bool = True\n    ALLOW_TF32: bool = True\n    MATMUL_PRECISION: str = \"high\"\n    CHANNELS_LAST: bool = True\n    COMPILE_MODEL: bool = False\n    COMPILE_MODE: str = \"max-autotune\"\n\n    # ----------------------------\n    # Execution toggles\n    # ----------------------------\n    RUN_PHASE_C_SCREENING: bool = False   # Phase D-only notebook: không chạy screening cũ\n    RUN_PHASE_C_CONFIRMATORY: bool = False # Phase D-only notebook: không chạy confirmatory cũ\n\n    def finalize(self):\n        self.PROJECT_ROOT = Path(self.PROJECT_ROOT)\n        self.IMAGE_DIR = Path(self.IMAGE_DIR)\n        self.DATA_DIR = Path(self.DATA_DIR)\n        self.INPUT_DIR = Path(self.INPUT_DIR)\n        self.LABELS_PATH = Path(self.LABELS_PATH)\n\n        self.CONVERTED_DIR = Path(self.CONVERTED_DIR)\n        self.FILTERED_DIR = Path(self.FILTERED_DIR)\n        self.REPORT_DIR = Path(self.REPORT_DIR)\n        self.BENCHMARK_DIR = Path(self.BENCHMARK_DIR)\n\n        # Chỉ tạo output dirs trong PROJECT_ROOT; không ghi vào /kaggle/input.\n        ensure_dir(self.PROJECT_ROOT)\n        ensure_dir(self.CONVERTED_DIR)\n        ensure_dir(self.FILTERED_DIR)\n        ensure_dir(self.REPORT_DIR)\n        ensure_dir(self.BENCHMARK_DIR)\n\n        print(\"=\" * 100)\n        print(\"Resolved project paths\")\n        print(\"IS_KAGGLE    :\", self.IS_KAGGLE)\n        print(\"PROJECT_ROOT :\", self.PROJECT_ROOT)\n        print(\"IMAGE_DIR    :\", self.IMAGE_DIR, \"| exists =\", self.IMAGE_DIR.exists())\n        print(\"LABELS_PATH  :\", self.LABELS_PATH, \"| exists =\", self.LABELS_PATH.exists())\n        print(\"OUTPUT_ROOT  :\", self.PROJECT_ROOT)\n        print(\"=\" * 100)\n\n        if not self.IMAGE_DIR.exists():\n            raise FileNotFoundError(\n                f\"Không tìm thấy IMAGE_DIR: {self.IMAGE_DIR}. \"\n                \"Hãy kiểm tra lại Kaggle input path hoặc set MALWARE_IMAGE_DIR.\"\n            )\n        if not self.LABELS_PATH.exists():\n            has_class_subfolders = any(p.is_dir() for p in self.IMAGE_DIR.iterdir()) if self.IMAGE_DIR.exists() else False\n            if has_class_subfolders:\n                print(\n                    \"⚠️ Không tìm thấy trainLabels.csv. \"\n                    \"Notebook sẽ thử suy luận nhãn từ tên thư mục con của IMAGE_DIR.\"\n                )\n            else:\n                raise FileNotFoundError(\n                    f\"Không tìm thấy trainLabels.csv: {self.LABELS_PATH}. \"\n                    \"Hãy thêm file trainLabels.csv vào Kaggle dataset, set MALWARE_LABELS_PATH, \"\n                    \"hoặc dùng dataset có cấu trúc class-subfolder.\"\n                )\n        return self\n\n\nCFG = ProjectConfig().finalize()\n","metadata":{"id":"0ba938eb","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T10:57:13.043888Z","iopub.execute_input":"2026-05-09T10:57:13.044522Z","iopub.status.idle":"2026-05-09T10:57:13.316772Z","shell.execute_reply.started":"2026-05-09T10:57:13.044495Z","shell.execute_reply":"2026-05-09T10:57:13.316114Z"}},"outputs":[],"execution_count":null},{"id":"05d1d1d5","cell_type":"code","source":"# ============================================================\n# Cell vai trò:\n# - Chuẩn hóa các đường dẫn artifact thường dùng của pipeline.\n# - Tạo manifest.csv từ dataset Kaggle.\n# - Hỗ trợ cả ảnh nằm phẳng và cấu trúc class-subfolder kiểu ImageFolder.\n# ============================================================\n\nMANIFEST_PATH = CFG.CONVERTED_DIR / \"manifest.csv\"\nTENSOR_DIR = CFG.CONVERTED_DIR / \"preprocessed_tensors\"\nFILTERED_BYTES_DIR = CFG.FILTERED_DIR / \"bytes\"\nFILTERED_ARTIFACT_DIR = CFG.FILTERED_DIR / \"images\"\nFILTERED_TENSOR_DIR = CFG.FILTERED_DIR / \"tensors\"\nFILTERED_KEEP_MANIFEST_PATH = CFG.FILTERED_DIR / \"manifest_filtered_keep.csv\"\nFILTERED_DROP_MANIFEST_PATH = CFG.FILTERED_DIR / \"manifest_filtered_drop.csv\"\nFILTERED_SUMMARY_PATH = CFG.FILTERED_DIR / \"filter_summary.csv\"\n\nDATASET_VERSION_DIR = CFG.REPORT_DIR / \"dataset_versions\"\nSPLIT_DIR = CFG.BENCHMARK_DIR / \"splits\"\n\nfor p in [\n    CFG.CONVERTED_DIR,\n    TENSOR_DIR,\n    FILTERED_BYTES_DIR,\n    FILTERED_ARTIFACT_DIR,\n    FILTERED_TENSOR_DIR,\n    DATASET_VERSION_DIR,\n    SPLIT_DIR,\n]:\n    ensure_dir(p)\n\n\ndef _sample_id_from_converted_image(path: Path) -> str:\n    \"\"\"01Iso..._rgb.png -> 01Iso... để khớp cột Id trong trainLabels.csv.\"\"\"\n    return _normalize_sample_id(path)\n\n\ndef build_manifest_from_single_image_folder(\n    image_dir: Path,\n    manifest_path: Path,\n    labels_path: Optional[Path] = None,\n    use_folder_labels_if_available: bool = True,\n) -> pd.DataFrame:\n    \"\"\"Tạo manifest từ dataset ảnh.\n\n    Dataset anhTanSuat trên Kaggle có cấu trúc:\n        image_dir/1/*.png\n        image_dir/2/*.png\n        ...\n        image_dir/9/*.png\n\n    Vì vậy notebook ưu tiên lấy nhãn từ tên thư mục cha nếu có. trainLabels.csv vẫn được giữ để tham chiếu/đối chiếu, nhưng không bắt buộc để train.\n    \"\"\"\n    image_dir = Path(image_dir)\n    manifest_path = Path(manifest_path)\n\n    exts = (\"*.png\", \"*.jpg\", \"*.jpeg\", \"*.bmp\", \"*.webp\")\n    image_paths = []\n    for ext in exts:\n        image_paths.extend(image_dir.glob(ext))\n\n    # Nếu thư mục gốc không chứa ảnh trực tiếp, tìm đệ quy trong class-subfolders.\n    if len(image_paths) == 0:\n        for ext in exts:\n            image_paths.extend(image_dir.rglob(ext))\n\n    image_paths = sorted(set(image_paths))\n    if len(image_paths) == 0:\n        raise FileNotFoundError(f\"Không tìm thấy ảnh trong {image_dir}\")\n\n    def _class_from_parent(p: Path):\n        rel_parent = p.parent.relative_to(image_dir) if image_dir in p.parents else p.parent\n        if str(rel_parent) in {\".\", \"\"}:\n            return None\n        # Dùng thư mục con cấp đầu tiên làm class. Ví dụ: image_dir/7/abc.png -> class 7.\n        class_name = rel_parent.parts[0]\n        m = re.search(r\"\\d+\", str(class_name))\n        if m:\n            return int(m.group(0))\n        return class_name\n\n    rows = []\n    for p in image_paths:\n        sample_id = _sample_id_from_converted_image(p)\n        inferred_class = _class_from_parent(p)\n        row = {\n            \"sample_id\": sample_id,\n            \"file_name\": p.name,\n            \"image_path\": str(p),\n            \"image_path_used\": str(p),\n            \"copied_image_path\": str(p),\n            \"source_dir\": str(image_dir),\n        }\n        if inferred_class is not None:\n            row[\"class\"] = inferred_class\n            row[\"Class\"] = inferred_class\n            row[\"class_source\"] = \"folder_name\"\n        rows.append(row)\n\n    manifest_df = pd.DataFrame(rows)\n\n    has_folder_labels = \"Class\" in manifest_df.columns and manifest_df[\"Class\"].notna().any()\n    if has_folder_labels and use_folder_labels_if_available:\n        print(\"Dataset có cấu trúc class-subfolder -> dùng nhãn từ tên thư mục cha.\")\n        print(\"Folder-label distribution:\")\n        print(manifest_df[\"Class\"].value_counts().sort_index())\n    elif labels_path is not None and Path(labels_path).exists():\n        labels_df = pd.read_csv(labels_path)\n        if \"Id\" in labels_df.columns:\n            label_ids = set(labels_df[\"Id\"].astype(str).map(_normalize_sample_id))\n            manifest_ids = set(manifest_df[\"sample_id\"].astype(str).map(_normalize_sample_id))\n            matched = len(manifest_ids & label_ids)\n            print(f\"Manifest-label matched IDs: {matched:,}/{len(manifest_ids):,} images\")\n            if matched == 0 and has_folder_labels:\n                print(\"⚠️ Không match trainLabels.csv. Sẽ fallback sang nhãn từ tên thư mục.\")\n    elif has_folder_labels:\n        print(\"Không có trainLabels.csv; dùng nhãn suy luận từ class-subfolders.\")\n\n    ensure_dir(manifest_path.parent)\n    manifest_df.to_csv(manifest_path, index=False)\n    print(f\"Saved manifest: {manifest_path}\")\n    print(f\"Number of images: {len(manifest_df):,}\")\n    print(manifest_df.head())\n    return manifest_df\n\n\nif bool(getattr(CFG, \"FORCE_REBUILD_MANIFEST\", True)) or not MANIFEST_PATH.exists():\n    manifest_df = build_manifest_from_single_image_folder(\n        image_dir=CFG.IMAGE_DIR,\n        manifest_path=MANIFEST_PATH,\n        labels_path=CFG.LABELS_PATH,\n        use_folder_labels_if_available=bool(getattr(CFG, \"USE_FOLDER_LABELS_IF_AVAILABLE\", True)),\n    )\nelse:\n    manifest_df = pd.read_csv(MANIFEST_PATH)\n    print(f\"Loaded existing manifest: {MANIFEST_PATH} | rows = {len(manifest_df):,}\")\n\nprint(f\"Runtime device: {CFG.DEVICE} | AMP: {CFG.USE_AMP}\")\nprint(f\"Image folder: {CFG.IMAGE_DIR}\")\nprint(f\"Labels path : {CFG.LABELS_PATH}\")\nprint(f\"Manifest    : {MANIFEST_PATH}\")\n","metadata":{"id":"05d1d1d5","outputId":"7a634993-1add-418c-bab8-07734a66bc26","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T10:57:13.317916Z","iopub.execute_input":"2026-05-09T10:57:13.318247Z","iopub.status.idle":"2026-05-09T10:57:49.757482Z","shell.execute_reply.started":"2026-05-09T10:57:13.318222Z","shell.execute_reply":"2026-05-09T10:57:49.756617Z"}},"outputs":[],"execution_count":null},{"id":"4d3ee736","cell_type":"code","source":"\n# ============================================================\n# Cell vai trò:\n# - Tối ưu môi trường chạy cho GPU/CPU trước khi bước vào KD.\n# - Thiết lập AMP, pin_memory, persistent_workers, channels_last,\n#   torch.compile (nếu phù hợp) và các lựa chọn runtime liên quan.\n#\n# Ý nghĩa thực nghiệm:\n# - Distillation thường chạy nhiều epoch và cần teacher + student cùng lúc.\n# - Tối ưu runtime tốt sẽ giảm đáng kể thời gian thử nghiệm.\n# ============================================================\n\n# GPU runtime setup + patch training loop (CUDA-first, AMP API mới, channels_last, compile tùy chọn)\n\nfrom pprint import pprint\n\nrequire_torch()\n\ndef get_runtime_config(\n    device: Optional[str] = None,\n    use_amp: Optional[bool] = None,\n    pin_memory: Optional[bool] = None,\n    persistent_workers: Optional[bool] = None,\n    prefetch_factor: Optional[int] = None,\n):\n    device = device or getattr(CFG, \"DEVICE\", None) or (\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    if device == \"auto\":\n        device = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n\n    use_cuda = str(device).startswith(\"cuda\") and torch.cuda.is_available()\n    device_type = \"cuda\" if use_cuda else \"cpu\"\n\n    if use_cuda:\n        if getattr(CFG, \"CUDNN_BENCHMARK\", True) and hasattr(torch.backends, \"cudnn\"):\n            torch.backends.cudnn.benchmark = True\n        if getattr(CFG, \"ALLOW_TF32\", True):\n            try:\n                torch.backends.cuda.matmul.allow_tf32 = True\n                if hasattr(torch.backends, \"cudnn\"):\n                    torch.backends.cudnn.allow_tf32 = True\n            except Exception:\n                pass\n        if hasattr(torch, \"set_float32_matmul_precision\"):\n            try:\n                torch.set_float32_matmul_precision(getattr(CFG, \"MATMUL_PRECISION\", \"high\"))\n            except Exception:\n                pass\n\n    runtime = {\n        \"device\": str(device),\n        \"device_type\": device_type,\n        \"use_cuda\": bool(use_cuda),\n        \"use_amp\": bool(use_amp if use_amp is not None else getattr(CFG, \"USE_AMP\", True)) and use_cuda,\n        \"pin_memory\": bool(pin_memory if pin_memory is not None else getattr(CFG, \"PIN_MEMORY\", True)) and use_cuda,\n        \"persistent_workers\": bool(\n            persistent_workers if persistent_workers is not None else getattr(CFG, \"PERSISTENT_WORKERS\", True)\n        ),\n        \"prefetch_factor\": int(prefetch_factor if prefetch_factor is not None else getattr(CFG, \"PREFETCH_FACTOR\", 2)),\n        \"channels_last\": bool(getattr(CFG, \"CHANNELS_LAST\", True)) and use_cuda,\n        \"compile_model\": bool(getattr(CFG, \"COMPILE_MODEL\", False)),\n        \"compile_mode\": str(getattr(CFG, \"COMPILE_MODE\", \"max-autotune\")),\n    }\n    return runtime\n\n\ndef print_runtime_summary(runtime: Dict) -> None:\n    print(\"Resolved runtime config:\")\n    pprint(runtime)\n\n    if runtime[\"use_cuda\"]:\n        print(f\"GPU detected: {torch.cuda.get_device_name(0)}\")\n        props = torch.cuda.get_device_properties(0)\n        print(f\"Total dedicated VRAM: {props.total_memory / (1024**3):.2f} GB\")\n        print(\"Lưu ý: CUDA/PyTorch chủ yếu dùng dedicated VRAM, không tính shared GPU memory của Windows.\")\n    else:\n        print(\"Không phát hiện GPU CUDA. Notebook sẽ fallback sang CPU.\")\n\n\ndef format_cuda_memory() -> Dict[str, float]:\n    if not torch.cuda.is_available():\n        return {\"allocated_gb\": 0.0, \"reserved_gb\": 0.0, \"max_allocated_gb\": 0.0, \"max_reserved_gb\": 0.0}\n    return {\n        \"allocated_gb\": float(torch.cuda.memory_allocated() / (1024**3)),\n        \"reserved_gb\": float(torch.cuda.memory_reserved() / (1024**3)),\n        \"max_allocated_gb\": float(torch.cuda.max_memory_allocated() / (1024**3)),\n        \"max_reserved_gb\": float(torch.cuda.max_memory_reserved() / (1024**3)),\n    }\n\n\ndef prepare_model_for_runtime(model: nn.Module, runtime: Dict) -> nn.Module:\n    model = model.to(runtime[\"device\"])\n    if runtime[\"channels_last\"]:\n        model = model.to(memory_format=torch.channels_last)\n\n    if runtime[\"compile_model\"] and hasattr(torch, \"compile\"):\n        try:\n            model = torch.compile(model, mode=runtime[\"compile_mode\"])\n            print(f\"torch.compile enabled (mode={runtime['compile_mode']})\")\n        except Exception as exc:\n            print(f\"torch.compile skipped: {exc}\")\n    return model\n\n\ndef move_batch_to_device(images, targets, runtime: Dict):\n    images = images.to(runtime[\"device\"], non_blocking=True)\n    if runtime[\"channels_last\"] and images.ndim == 4:\n        images = images.to(memory_format=torch.channels_last)\n    targets = targets.to(runtime[\"device\"], non_blocking=True)\n    return images, targets\n\n\nruntime_cfg = get_runtime_config()\nprint_runtime_summary(runtime_cfg)\n\n\ndef build_loaders(\n    train_df: pd.DataFrame,\n    val_df: pd.DataFrame,\n    test_df: pd.DataFrame,\n    img_size: int = 224,\n    in_channels: int = 3,\n    batch_size: int = 32,\n    num_workers: int = 2,\n    use_imagenet_norm: bool = True,\n    use_weighted_sampler: bool = True,\n    weak_aug: bool = False,\n    device: Optional[str] = None,\n    pin_memory: Optional[bool] = None,\n    persistent_workers: Optional[bool] = None,\n    prefetch_factor: Optional[int] = None,\n):\n    runtime = get_runtime_config(\n        device=device,\n        pin_memory=pin_memory,\n        persistent_workers=persistent_workers,\n        prefetch_factor=prefetch_factor,\n    )\n\n    train_tfms, eval_tfms = build_transforms(\n        img_size=img_size,\n        in_channels=in_channels,\n        use_imagenet_norm=use_imagenet_norm,\n        weak_aug=weak_aug,\n    )\n\n    train_dataset, val_dataset, test_dataset = build_datasets(\n        train_df=train_df,\n        val_df=val_df,\n        test_df=test_df,\n        train_transforms=train_tfms,\n        eval_transforms=eval_tfms,\n        in_channels=in_channels,\n    )\n\n    loader_kwargs = {\n        \"batch_size\": batch_size,\n        \"num_workers\": num_workers,\n        \"pin_memory\": runtime[\"pin_memory\"],\n    }\n    if num_workers > 0:\n        loader_kwargs[\"persistent_workers\"] = runtime[\"persistent_workers\"]\n        loader_kwargs[\"prefetch_factor\"] = max(2, runtime[\"prefetch_factor\"])\n\n    train_drop_last = len(train_dataset) >= batch_size\n\n    if use_weighted_sampler:\n        num_classes = int(train_df[\"label\"].nunique())\n        class_weights_np, _ = get_effective_num_class_weights(train_df[\"label\"].values, num_classes)\n        sample_weights = np.array([class_weights_np[y] for y in train_df[\"label\"].values], dtype=np.float32)\n        sampler = WeightedRandomSampler(\n            weights=torch.from_numpy(sample_weights),\n            num_samples=len(sample_weights),\n            replacement=True,\n        )\n        train_loader = DataLoader(train_dataset, sampler=sampler, drop_last=train_drop_last, **loader_kwargs)\n    else:\n        train_loader = DataLoader(train_dataset, shuffle=True, drop_last=train_drop_last, **loader_kwargs)\n\n    eval_loader_kwargs = dict(loader_kwargs)\n    val_loader = DataLoader(val_dataset, shuffle=False, drop_last=False, **eval_loader_kwargs)\n    test_loader = DataLoader(test_dataset, shuffle=False, drop_last=False, **eval_loader_kwargs)\n    return train_loader, val_loader, test_loader\n\n\n\ndef train_one_epoch(\n    model,\n    loader,\n    criterion,\n    optimizer,\n    device,\n    num_classes,\n    use_amp: bool = False,\n    scaler=None,\n    model_name: str = \"\",\n    epoch_idx: Optional[int] = None,\n    total_epochs: Optional[int] = None,\n    show_batch_progress: bool = True,\n):\n    model.train()\n    running_loss = 0.0\n    all_preds = []\n    all_targets = []\n    runtime = get_runtime_config(device=device, use_amp=use_amp)\n\n    if epoch_idx is not None and total_epochs is not None:\n        desc = f\"[{model_name}] Train {epoch_idx}/{total_epochs}\"\n    else:\n        desc = f\"[{model_name}] Train\"\n\n    iterable = tqdm(loader, desc=desc, leave=False, dynamic_ncols=True) if show_batch_progress else loader\n\n    for images, targets in iterable:\n        images, targets = move_batch_to_device(images, targets, runtime)\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.amp.autocast(runtime[\"device_type\"], enabled=runtime[\"use_amp\"]):\n            logits = _forward_logits(model, images)\n            loss = criterion(logits, targets)\n\n        if runtime[\"use_amp\"] and scaler is not None:\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            loss.backward()\n            optimizer.step()\n\n        running_loss += loss.item() * images.size(0)\n        preds = torch.argmax(logits, dim=1)\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_targets.extend(targets.detach().cpu().numpy())\n\n        if show_batch_progress:\n            iterable.set_postfix(loss=f\"{loss.item():.4f}\")\n\n    epoch_loss = running_loss / max(len(loader.dataset), 1)\n    metrics = compute_metrics(all_targets, all_preds, num_classes)\n    return epoch_loss, metrics\n\n\n@torch.no_grad()\ndef evaluate(\n    model,\n    loader,\n    criterion,\n    device,\n    num_classes,\n    use_amp: bool = False,\n    model_name: str = \"\",\n    epoch_idx: Optional[int] = None,\n    total_epochs: Optional[int] = None,\n    stage_name: str = \"Val\",\n    show_batch_progress: bool = False,\n):\n    model.eval()\n    running_loss = 0.0\n    all_preds = []\n    all_targets = []\n    runtime = get_runtime_config(device=device, use_amp=use_amp)\n\n    if epoch_idx is not None and total_epochs is not None:\n        desc = f\"[{model_name}] {stage_name} {epoch_idx}/{total_epochs}\"\n    else:\n        desc = f\"[{model_name}] {stage_name}\"\n\n    iterable = tqdm(loader, desc=desc, leave=False, dynamic_ncols=True) if show_batch_progress else loader\n\n    for images, targets in iterable:\n        images, targets = move_batch_to_device(images, targets, runtime)\n\n        with torch.amp.autocast(runtime[\"device_type\"], enabled=runtime[\"use_amp\"]):\n            logits = _forward_logits(model, images)\n            loss = criterion(logits, targets)\n\n        running_loss += loss.item() * images.size(0)\n        preds = torch.argmax(logits, dim=1)\n        all_preds.extend(preds.detach().cpu().numpy())\n        all_targets.extend(targets.detach().cpu().numpy())\n\n        if show_batch_progress:\n            iterable.set_postfix(loss=f\"{loss.item():.4f}\")\n\n    epoch_loss = running_loss / max(len(loader.dataset), 1)\n    metrics = compute_metrics(all_targets, all_preds, num_classes)\n    return epoch_loss, metrics, all_targets, all_preds\n\n\n\ndef benchmark_one_model(\n    model_name: str,\n    train_df: pd.DataFrame,\n    val_df: pd.DataFrame,\n    test_df: pd.DataFrame,\n    num_classes: int,\n    idx_to_class: Dict[int, int],\n    output_dir: Path,\n    img_size: int = 224,\n    in_channels: int = 3,\n    batch_size: int = 32,\n    epochs: int = 5,\n    lr: float = 1e-3,\n    weight_decay: float = 1e-4,\n    loss_name: str = \"ce\",\n    use_class_weights_in_loss: bool = True,\n    use_weighted_sampler: bool = True,\n    num_workers: int = 2,\n    pretrained: bool = True,\n    weak_aug: bool = False,\n    early_stop_patience: int = 3,\n    seed: int = 42,\n    device: Optional[str] = None,\n    use_amp: Optional[bool] = None,\n    pin_memory: Optional[bool] = None,\n    persistent_workers: Optional[bool] = None,\n    prefetch_factor: Optional[int] = None,\n    show_batch_progress: bool = True,\n    show_epoch_summary: bool = True,\n) -> Dict:\n    require_torch()\n    set_seed(seed)\n    output_dir = ensure_dir(output_dir)\n\n    runtime = get_runtime_config(\n        device=device,\n        use_amp=use_amp,\n        pin_memory=pin_memory,\n        persistent_workers=persistent_workers,\n        prefetch_factor=prefetch_factor,\n    )\n    device = runtime[\"device\"]\n\n    print(\"=\" * 100)\n    print(f\"[{model_name}] device={device} | amp={runtime['use_amp']} | batch_size={batch_size} | epochs={epochs}\")\n    print(\"=\" * 100)\n\n    train_loader, val_loader, test_loader = build_loaders(\n        train_df=train_df,\n        val_df=val_df,\n        test_df=test_df,\n        img_size=img_size,\n        in_channels=in_channels,\n        batch_size=batch_size,\n        num_workers=num_workers,\n        use_imagenet_norm=True,\n        use_weighted_sampler=use_weighted_sampler,\n        weak_aug=weak_aug,\n        device=device,\n        pin_memory=runtime[\"pin_memory\"],\n        persistent_workers=runtime[\"persistent_workers\"],\n        prefetch_factor=runtime[\"prefetch_factor\"],\n    )\n\n    model = build_model(\n        model_name=model_name,\n        num_classes=num_classes,\n        in_channels=in_channels,\n        pretrained=pretrained,\n    )\n    model = prepare_model_for_runtime(model, runtime)\n\n    class_weights = None\n    if use_class_weights_in_loss:\n        class_weights_np, _ = get_effective_num_class_weights(train_df[\"label\"].values, num_classes)\n        class_weights = torch.tensor(class_weights_np, dtype=torch.float32).to(device)\n\n    criterion = build_criterion(loss_name, class_weights=class_weights)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer,\n        mode=\"max\",\n        factor=0.5,\n        patience=max(1, early_stop_patience // 2),\n    )\n    scaler = torch.amp.GradScaler(runtime[\"device_type\"], enabled=runtime[\"use_amp\"])\n\n    best_score = -1.0\n    best_epoch = -1\n    best_state = None\n    patience_counter = 0\n    history = []\n\n    if runtime[\"use_cuda\"]:\n        torch.cuda.reset_peak_memory_stats()\n\n    t0 = time.perf_counter()\n    for epoch in range(1, epochs + 1):\n        train_loss, train_metrics = train_one_epoch(\n            model,\n            train_loader,\n            criterion,\n            optimizer,\n            device,\n            num_classes,\n            use_amp=runtime[\"use_amp\"],\n            scaler=scaler,\n            model_name=model_name,\n            epoch_idx=epoch,\n            total_epochs=epochs,\n            show_batch_progress=show_batch_progress,\n        )\n        val_loss, val_metrics, _, _ = evaluate(\n            model,\n            val_loader,\n            criterion,\n            device,\n            num_classes,\n            use_amp=runtime[\"use_amp\"],\n            model_name=model_name,\n            epoch_idx=epoch,\n            total_epochs=epochs,\n            stage_name=\"Val\",\n            show_batch_progress=False,\n        )\n\n        current_score = val_metrics[\"macro_f1\"]\n        scheduler.step(current_score)\n\n        memory_stats = format_cuda_memory() if runtime[\"use_cuda\"] else {}\n        history_row = {\n            \"epoch\": epoch,\n            \"train_loss\": train_loss,\n            \"val_loss\": val_loss,\n            \"train_macro_f1\": train_metrics[\"macro_f1\"],\n            \"val_macro_f1\": val_metrics[\"macro_f1\"],\n            \"train_bal_acc\": train_metrics[\"balanced_acc\"],\n            \"val_bal_acc\": val_metrics[\"balanced_acc\"],\n            \"lr\": optimizer.param_groups[0][\"lr\"],\n            \"cuda_allocated_gb\": memory_stats.get(\"allocated_gb\"),\n            \"cuda_reserved_gb\": memory_stats.get(\"reserved_gb\"),\n            \"cuda_max_allocated_gb\": memory_stats.get(\"max_allocated_gb\"),\n            \"cuda_max_reserved_gb\": memory_stats.get(\"max_reserved_gb\"),\n        }\n        history.append(history_row)\n\n        if show_epoch_summary:\n            peak_text = (\n                f\" | peak_vram={history_row['cuda_max_allocated_gb']:.2f}GB\"\n                if history_row[\"cuda_max_allocated_gb\"] is not None\n                else \"\"\n            )\n            print(\n                f\"[{model_name}] Epoch {epoch:02d}/{epochs} | \"\n                f\"train_loss={train_loss:.4f} | val_loss={val_loss:.4f} | \"\n                f\"train_f1={train_metrics['macro_f1']:.4f} | val_f1={val_metrics['macro_f1']:.4f} | \"\n                f\"val_bacc={val_metrics['balanced_acc']:.4f} | lr={optimizer.param_groups[0]['lr']:.2e}\"\n                f\"{peak_text}\"\n            )\n\n        if current_score > best_score:\n            best_score = current_score\n            best_epoch = epoch\n            best_state = copy.deepcopy(model.state_dict())\n            patience_counter = 0\n            if show_epoch_summary:\n                print(f\"[{model_name}] ↳ New best checkpoint at epoch {epoch} (val_macro_f1={best_score:.4f})\")\n        else:\n            patience_counter += 1\n            if show_epoch_summary:\n                print(f\"[{model_name}] ↳ No improvement. patience={patience_counter}/{early_stop_patience}\")\n            if patience_counter >= early_stop_patience:\n                if show_epoch_summary:\n                    print(f\"[{model_name}] Early stopping triggered at epoch {epoch}.\")\n                break\n\n    total_train_time_sec = time.perf_counter() - t0\n    history_df = pd.DataFrame(history)\n    history_df.to_csv(output_dir / f\"{model_name}_history.csv\", index=False)\n\n    if best_state is None:\n        raise RuntimeError(f\"Model {model_name} không tạo được checkpoint tốt nhất.\")\n    model.load_state_dict(best_state)\n\n    test_loss, test_metrics, test_targets, test_preds = evaluate(\n        model,\n        test_loader,\n        criterion,\n        device,\n        num_classes,\n        use_amp=runtime[\"use_amp\"],\n        model_name=model_name,\n        stage_name=\"Test\",\n        show_batch_progress=False,\n    )\n\n    cm = confusion_matrix(test_targets, test_preds, labels=list(range(num_classes)))\n    cm_df = pd.DataFrame(\n        cm,\n        index=[f\"true_{idx_to_class[i]}\" for i in range(num_classes)],\n        columns=[f\"pred_{idx_to_class[i]}\" for i in range(num_classes)],\n    )\n    cm_df.to_csv(output_dir / f\"{model_name}_confusion_matrix.csv\")\n\n    per_class_df = pd.DataFrame(\n        {\n            \"Class\": [idx_to_class[i] for i in range(num_classes)],\n            \"precision\": test_metrics[\"per_class_precision\"],\n            \"recall\": test_metrics[\"per_class_recall\"],\n            \"f1\": test_metrics[\"per_class_f1\"],\n            \"support\": test_metrics[\"per_class_support\"],\n        }\n    )\n    per_class_df.to_csv(output_dir / f\"{model_name}_per_class_metrics.csv\", index=False)\n\n    ckpt_path = output_dir / f\"{model_name}_best.pth\"\n    torch.save(\n        {\n            \"model_name\": model_name,\n            \"best_epoch\": best_epoch,\n            \"model_state_dict\": model.state_dict(),\n            \"device\": device,\n            \"use_amp\": runtime[\"use_amp\"],\n            \"runtime_config\": runtime,\n        },\n        ckpt_path,\n    )\n\n    final_memory_stats = format_cuda_memory() if runtime[\"use_cuda\"] else {}\n    metrics = {\n        \"model_name\": model_name,\n        \"num_params\": count_parameters(model),\n        \"best_epoch\": int(best_epoch),\n        \"device\": device,\n        \"use_amp\": bool(runtime[\"use_amp\"]),\n        \"channels_last\": bool(runtime[\"channels_last\"]),\n        \"compile_model\": bool(runtime[\"compile_model\"]),\n        \"test_loss\": float(test_loss),\n        \"test_macro_f1\": float(test_metrics[\"macro_f1\"]),\n        \"test_balanced_acc\": float(test_metrics[\"balanced_acc\"]),\n        \"total_train_time_sec\": float(total_train_time_sec),\n        \"cuda_peak_allocated_gb\": float(final_memory_stats.get(\"max_allocated_gb\", 0.0)),\n        \"cuda_peak_reserved_gb\": float(final_memory_stats.get(\"max_reserved_gb\", 0.0)),\n        \"checkpoint_path\": str(ckpt_path),\n        \"history_path\": str(output_dir / f\"{model_name}_history.csv\"),\n        \"confusion_matrix_path\": str(output_dir / f\"{model_name}_confusion_matrix.csv\"),\n        \"per_class_metrics_path\": str(output_dir / f\"{model_name}_per_class_metrics.csv\"),\n        \"classification_report_text\": classification_report(\n            test_targets,\n            test_preds,\n            labels=list(range(num_classes)),\n            target_names=[f\"Class_{idx_to_class[i]}\" for i in range(num_classes)],\n            zero_division=0,\n        ),\n    }\n\n    save_json(metrics, output_dir / f\"{model_name}_test_metrics.json\")\n\n    print(\n        f\"[{model_name}] DONE | best_epoch={best_epoch} | \"\n        f\"test_macro_f1={metrics['test_macro_f1']:.4f} | \"\n        f\"test_balanced_acc={metrics['test_balanced_acc']:.4f} | \"\n        f\"time={metrics['total_train_time_sec'] / 60:.2f} min\"\n    )\n\n    if runtime[\"use_cuda\"]:\n        torch.cuda.empty_cache()\n\n    return metrics\n\n\ndef benchmark_model_list(\n    model_list: Sequence[str],\n    train_df: pd.DataFrame,\n    val_df: pd.DataFrame,\n    test_df: pd.DataFrame,\n    num_classes: int,\n    idx_to_class: Dict[int, int],\n    output_dir: Path,\n    img_size: int = 224,\n    in_channels: int = 3,\n    batch_size: int = 32,\n    epochs: int = 5,\n    lr: float = 1e-3,\n    weight_decay: float = 1e-4,\n    loss_name: str = \"ce\",\n    use_class_weights_in_loss: bool = True,\n    use_weighted_sampler: bool = True,\n    num_workers: int = 2,\n    pretrained: bool = True,\n    weak_aug: bool = False,\n    early_stop_patience: int = 3,\n    seed: int = 42,\n    device: Optional[str] = None,\n    use_amp: Optional[bool] = None,\n    pin_memory: Optional[bool] = None,\n    persistent_workers: Optional[bool] = None,\n    prefetch_factor: Optional[int] = None,\n    show_batch_progress: bool = True,\n    show_epoch_summary: bool = True,\n) -> pd.DataFrame:\n    rows = []\n    ensure_dir(output_dir)\n\n    for model_name in model_list:\n        model_dir = ensure_dir(output_dir / model_name)\n        metrics = benchmark_one_model(\n            model_name=model_name,\n            train_df=train_df,\n            val_df=val_df,\n            test_df=test_df,\n            num_classes=num_classes,\n            idx_to_class=idx_to_class,\n            output_dir=model_dir,\n            img_size=img_size,\n            in_channels=in_channels,\n            batch_size=batch_size,\n            epochs=epochs,\n            lr=lr,\n            weight_decay=weight_decay,\n            loss_name=loss_name,\n            use_class_weights_in_loss=use_class_weights_in_loss,\n            use_weighted_sampler=use_weighted_sampler,\n            num_workers=num_workers,\n            pretrained=pretrained,\n            weak_aug=weak_aug,\n            early_stop_patience=early_stop_patience,\n            seed=seed,\n            device=device,\n            use_amp=use_amp,\n            pin_memory=pin_memory,\n            persistent_workers=persistent_workers,\n            prefetch_factor=prefetch_factor,\n            show_batch_progress=show_batch_progress,\n            show_epoch_summary=show_epoch_summary,\n        )\n        rows.append(metrics)\n\n    summary_df = pd.DataFrame(rows).sort_values(\n        [\"test_macro_f1\", \"test_balanced_acc\", \"num_params\"],\n        ascending=[False, False, True],\n    ).reset_index(drop=True)\n    summary_df.to_csv(output_dir / \"benchmark_summary.csv\", index=False)\n    return summary_df","metadata":{"id":"4d3ee736","outputId":"a59b2516-248b-492e-a2c9-0c09f8cc00df","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T10:57:49.759894Z","iopub.execute_input":"2026-05-09T10:57:49.760396Z","iopub.status.idle":"2026-05-09T10:57:49.868936Z","shell.execute_reply.started":"2026-05-09T10:57:49.760369Z","shell.execute_reply":"2026-05-09T10:57:49.868098Z"}},"outputs":[],"execution_count":null},{"id":"884cae59","cell_type":"markdown","source":"## Bootstrap dữ liệu cho Phase D\n\n### Mục tiêu\nKhôi phục đúng **train / val / test split** đã dùng ở Phase C, hoặc tái dựng split khi artifact cũ không còn.\n\n### Input\n- `benchmark/splits/train_split.csv`\n- `benchmark/splits/val_split.csv`\n- `benchmark/splits/test_split.csv`\n- hoặc `manifest + labels` của pipeline trước đó\n\n### Output trong RAM\n- `train_df`\n- `val_df`\n- `test_df`\n- `classes_sorted`, `class_to_idx`, `idx_to_class`\n\n### Artifact liên quan\n- Không sinh artifact mới nếu split đã tồn tại\n- Có thể ghi lại split nếu notebook phải fallback sang dựng lại\n\n### Lưu ý thực nghiệm\nĐây là cell rất quan trọng để bảo đảm **tính công bằng của phép so sánh**.  \nNếu teacher, baseline student và distilled student dùng khác split, kết quả sẽ không còn so sánh trực tiếp được.\n","metadata":{"id":"884cae59"}},{"id":"f20fdaec","cell_type":"code","source":"# ============================================================\n# Cell vai trò:\n# - Tạo hoặc nạp train/val/test split sạch cho Phase D từ folder ảnh + trainLabels.csv.\n# - Nếu split cũ có overlap sample_id, tự rebuild bằng group-stratified split.\n# - Fallback split luôn tách theo sample_id để tránh data leak từ duplicate rows.\n# ============================================================\n\n# ======================================\n# Bootstrap splits / labels for Phase D\n# ======================================\nrequire_torch()\n\nTRAIN_SPLIT_PATH = SPLIT_DIR / \"train_split.csv\"\nVAL_SPLIT_PATH = SPLIT_DIR / \"val_split.csv\"\nTEST_SPLIT_PATH = SPLIT_DIR / \"test_split.csv\"\n\n# Bật mặc định: nếu split cũ bị leak, notebook sẽ tự dựng lại split sạch.\nCFG.REBUILD_SPLITS_IF_LEAK_FOUND = bool(getattr(CFG, \"REBUILD_SPLITS_IF_LEAK_FOUND\", True))\nCFG.USE_GROUP_STRATIFIED_SPLIT = bool(getattr(CFG, \"USE_GROUP_STRATIFIED_SPLIT\", True))\n\n\ndef _standardize_split_df(split_df: pd.DataFrame, split_name: str) -> pd.DataFrame:\n    split_df = _ensure_sample_id_column(split_df, split_name=split_name)\n    if \"label\" not in split_df.columns:\n        if \"Class\" not in split_df.columns and \"class\" not in split_df.columns:\n            raise KeyError(f\"{split_name} thiếu cột 'label' hoặc 'Class'/'class'.\")\n        class_col = \"Class\" if \"Class\" in split_df.columns else \"class\"\n        split_df[\"Class\"] = split_df[class_col].astype(int)\n    return split_df.reset_index(drop=True)\n\n\ndef _rebuild_phase_d_splits_from_manifest() -> Tuple[pd.DataFrame, pd.DataFrame, pd.DataFrame, str]:\n    print(\"Rebuilding clean Phase D splits from manifest + labels ...\")\n    train_ready_df, classes_sorted_tmp, class_to_idx_tmp, idx_to_class_tmp = prepare_training_dataframe(\n        manifest_path=FILTERED_KEEP_MANIFEST_PATH if FILTERED_KEEP_MANIFEST_PATH.exists() else MANIFEST_PATH,\n        labels_path=CFG.LABELS_PATH,\n        dropna_class=True,\n    )\n    train_ready_df = _ensure_sample_id_column(train_ready_df, split_name=\"full_manifest\")\n\n    if CFG.USE_GROUP_STRATIFIED_SPLIT:\n        train_df, val_df, test_df = create_group_stratified_splits(\n            train_ready_df,\n            test_size=CFG.TEST_SIZE,\n            val_size_from_train=CFG.VAL_SIZE_FROM_TRAIN,\n            seed=CFG.SEED,\n            group_col=\"sample_id\",\n            label_col=\"label\",\n        )\n    else:\n        train_df, val_df, test_df = create_stratified_splits(\n            train_ready_df,\n            test_size=CFG.TEST_SIZE,\n            val_size_from_train=CFG.VAL_SIZE_FROM_TRAIN,\n            seed=CFG.SEED,\n        )\n\n    ensure_dir(SPLIT_DIR)\n    train_df.to_csv(TRAIN_SPLIT_PATH, index=False)\n    val_df.to_csv(VAL_SPLIT_PATH, index=False)\n    test_df.to_csv(TEST_SPLIT_PATH, index=False)\n    print(\"Saved clean splits to SPLIT_DIR.\")\n    return train_df, val_df, test_df, \"rebuilt_group_stratified\" if CFG.USE_GROUP_STRATIFIED_SPLIT else \"rebuilt_row_stratified\"\n\n\ndef _load_phase_d_splits():\n    if TRAIN_SPLIT_PATH.exists() and VAL_SPLIT_PATH.exists() and TEST_SPLIT_PATH.exists():\n        train_df = _standardize_split_df(pd.read_csv(TRAIN_SPLIT_PATH), \"train\")\n        val_df = _standardize_split_df(pd.read_csv(VAL_SPLIT_PATH), \"val\")\n        test_df = _standardize_split_df(pd.read_csv(TEST_SPLIT_PATH), \"test\")\n        print(\"Loaded saved Phase D splits:\")\n        print(\" -\", TRAIN_SPLIT_PATH)\n        print(\" -\", VAL_SPLIT_PATH)\n        print(\" -\", TEST_SPLIT_PATH)\n\n        saved_report = enforce_no_data_leakage(train_df, val_df, test_df, fail_fast=False)\n        print(\"Saved split leakage report:\", saved_report)\n        if any(v > 0 for v in saved_report.values()):\n            if CFG.REBUILD_SPLITS_IF_LEAK_FOUND:\n                print(\"⚠️ Saved splits có overlap sample_id -> rebuild split sạch để tránh data leak.\")\n                return _rebuild_phase_d_splits_from_manifest()\n            raise RuntimeError(f\"Saved splits có data leak: {saved_report}\")\n\n        return train_df, val_df, test_df, \"saved_splits_no_leak\"\n\n    print(\"Saved splits not found.\")\n    return _rebuild_phase_d_splits_from_manifest()\n\n\ntrain_df, val_df, test_df, phase_d_split_source = _load_phase_d_splits()\n\n# build class mapping consistently from TRAIN only, rồi kiểm tra val/test không có class ngoài train.\ntrain_df = _standardize_split_df(train_df, \"train\")\nval_df = _standardize_split_df(val_df, \"val\")\ntest_df = _standardize_split_df(test_df, \"test\")\n\nclass_col = \"Class\" if \"Class\" in train_df.columns else (\"class\" if \"class\" in train_df.columns else None)\nif \"label\" not in train_df.columns:\n    if class_col is None:\n        raise KeyError(\"Không tìm thấy cột 'label' hoặc 'Class'/'class' trong train split.\")\n    classes_sorted = sorted(train_df[class_col].astype(int).unique().tolist())\n    class_to_idx = {cls: i for i, cls in enumerate(classes_sorted)}\n    idx_to_class = {i: cls for cls, i in class_to_idx.items()}\n    for _df in (train_df, val_df, test_df):\n        _src_col = \"Class\" if \"Class\" in _df.columns else \"class\"\n        unseen = set(_df[_src_col].astype(int).unique()) - set(classes_sorted)\n        if unseen:\n            raise ValueError(f\"Split chứa class ngoài train: {unseen}\")\n        _df[\"label\"] = _df[_src_col].astype(int).map(class_to_idx)\nelse:\n    train_df[\"label\"] = train_df[\"label\"].astype(int)\n    val_df[\"label\"] = val_df[\"label\"].astype(int)\n    test_df[\"label\"] = test_df[\"label\"].astype(int)\n    label_values = sorted(train_df[\"label\"].unique().tolist())\n    val_test_unseen = (set(val_df[\"label\"].unique()) | set(test_df[\"label\"].unique())) - set(label_values)\n    if val_test_unseen:\n        raise ValueError(f\"Val/test chứa label ngoài train: {val_test_unseen}\")\n\n    if \"Class\" in train_df.columns:\n        classes_sorted = sorted(train_df[\"Class\"].astype(int).unique().tolist())\n        class_to_idx = {cls: i for i, cls in enumerate(classes_sorted)}\n        idx_to_class = {i: cls for cls, i in class_to_idx.items()}\n    else:\n        classes_sorted = label_values\n        class_to_idx = {cls: i for i, cls in enumerate(classes_sorted)}\n        idx_to_class = {i: cls for cls, i in class_to_idx.items()}\n\nfinal_leakage = enforce_no_data_leakage(train_df, val_df, test_df, fail_fast=True)\n\nprint(f\"Split source: {phase_d_split_source}\")\nprint(\"Train:\", train_df.shape, \"| Val:\", val_df.shape, \"| Test:\", test_df.shape)\nprint(\"Num classes:\", len(classes_sorted))\nprint(\"Leakage:\", final_leakage)\nprint(\"Unique sample_id:\", {\n    \"train\": train_df[\"sample_id\"].nunique(),\n    \"val\": val_df[\"sample_id\"].nunique(),\n    \"test\": test_df[\"sample_id\"].nunique(),\n})\n","metadata":{"id":"f20fdaec","outputId":"7c15cb57-b877-4e4b-d94c-b2327a6f0aae","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T10:57:49.869984Z","iopub.execute_input":"2026-05-09T10:57:49.870290Z","iopub.status.idle":"2026-05-09T10:57:50.999066Z","shell.execute_reply.started":"2026-05-09T10:57:49.870259Z","shell.execute_reply":"2026-05-09T10:57:50.998288Z"}},"outputs":[],"execution_count":null},{"id":"cb6b83be","cell_type":"markdown","source":"## D0. Phase D — Teacher refresh → Student baseline → Knowledge Distillation\n\n### Mục tiêu\nGộp toàn bộ Phase D thành một quy trình khép kín để so sánh công bằng:\n\n1. **Train lại / fine-tune teacher** với số epoch lớn hơn Phase C  \n2. **Train student standalone** (không distillation) trên đúng split đó  \n3. **Train student với KD** theo 3 sources: `logit`, `feature`, `similarity`\n\n### Vì sao cần thiết kế lại như vậy\n- Teacher từ Phase C chỉ train ít epoch có thể **chưa hội tụ**\n- Nếu không có **student standalone baseline**, rất khó kết luận KD có thực sự giúp hay không\n- Cả ba bước dùng chung split sẽ giúp so sánh kết quả **công bằng và có ý nghĩa khoa học**\n\n### Output chính của Phase D mới\n- checkpoint teacher đã refresh\n- checkpoint student supervised baseline\n- checkpoint tốt nhất cho từng nguồn distillation\n- bảng so sánh cuối cùng:\n  - teacher refreshed\n  - student baseline\n  - KD logit\n  - KD feature\n  - KD similarity\n","metadata":{"id":"cb6b83be"}},{"id":"7662a25c","cell_type":"code","source":"# ============================================================\n# Cell vai trò:\n# - Đóng vai trò \"bảng điều khiển\" của Phase D phiên bản đầy đủ.\n# - Phase D mới gồm 3 bước liên tiếp:\n#     1) refresh / retrain teacher\n#     2) train student supervised baseline\n#     3) knowledge distillation theo 3 sources\n#\n# Ghi chú khoa học:\n# - Teacher tốt hơn -> tín hiệu distillation đáng tin cậy hơn.\n# - Baseline student standalone là mốc bắt buộc để kết luận KD có hiệu quả hay không.\n# ============================================================\n\n# ============================================\n# D1. Cấu hình Phase D (full pipeline)\n# ============================================\nCFG.RUN_PHASE_D_TEACHER_RETRAIN = True\nCFG.RUN_PHASE_D_STUDENT_BASELINE = True\nCFG.RUN_PHASE_D_KD = True\n\n# --- output structure ---\nCFG.PHASE_D_DIR = ensure_dir(CFG.BENCHMARK_DIR / \"phase_d_full\")\nCFG.PHASE_D_TEACHER_DIR = ensure_dir(CFG.PHASE_D_DIR / \"teacher_refresh\")\nCFG.PHASE_D_STUDENT_BASELINE_DIR = ensure_dir(CFG.PHASE_D_DIR / \"student_supervised_baseline\")\nCFG.KD_DIR = ensure_dir(CFG.PHASE_D_DIR / \"kd_three_sources\")\n\n# --- distillation sources ---\nCFG.KD_SOURCES = [\"logit\", \"feature\", \"similarity\"]\n\n# --- anti-leak + efficiency metrics ---\nCFG.REBUILD_SPLITS_IF_LEAK_FOUND = True\nCFG.USE_GROUP_STRATIFIED_SPLIT = True\nCFG.EFFICIENCY_MAX_BATCHES = 20\nCFG.EFFICIENCY_WARMUP_BATCHES = 2\nCFG.MEASURE_FLOPS = True\nCFG.FLOP_BATCH_SIZE = 1  # đo FLOPs trên 1 ảnh để so sánh công bằng giữa model\n\n# -------------------------------------------------\n# Teacher retrain trong chính Phase D\n# -------------------------------------------------\n# Chạy độc lập trong Phase D; không cần artifact từ các bước trước.\n# Teacher được chọn cố định và sẽ được train lại trong Phase D từ dữ liệu Kaggle hiện tại.\nCFG.KD_TEACHER_MODEL = \"efficientnet_b0\"\n\n# Optional: nếu đã có checkpoint teacher riêng, set path này để nạp trước khi retrain hoặc dùng trực tiếp khi tắt retrain.\n# Mặc định None để train từ scratch trong Phase D.\nCFG.KD_TEACHER_CKPT_PATH = None\nCFG.PHASE_D_TEACHER_INIT_CKPT_PATH = None\n\nCFG.PHASE_D_TEACHER_INIT_FROM_PHASE_C = False\nCFG.PHASE_D_TEACHER_PRETRAINED = True\nCFG.PHASE_D_TEACHER_EPOCHS = 30\nCFG.PHASE_D_TEACHER_LR = float(getattr(CFG, \"LR\", 1e-3)) * 0.5\nCFG.PHASE_D_TEACHER_WEIGHT_DECAY = float(getattr(CFG, \"WEIGHT_DECAY\", 1e-4))\nCFG.PHASE_D_TEACHER_EARLY_STOP_PATIENCE = 7\nCFG.PHASE_D_TEACHER_LOSS_NAME = getattr(CFG, \"LOSS_NAME\", \"ce\")\nCFG.PHASE_D_TEACHER_USE_CLASS_WEIGHTS_IN_LOSS = getattr(CFG, \"USE_CLASS_WEIGHTS_IN_LOSS\", True)\nCFG.PHASE_D_TEACHER_USE_WEIGHTED_SAMPLER = getattr(CFG, \"USE_WEIGHTED_SAMPLER\", True)\nCFG.PHASE_D_TEACHER_WEAK_AUG = getattr(CFG, \"WEAK_AUG\", False)\n\n# -------------------------------------------------\n# Student undistilled / supervised baseline\n# -------------------------------------------------\nCFG.RUN_PHASE_D_STUDENT_UNDISTILLED = CFG.RUN_PHASE_D_STUDENT_BASELINE\nCFG.PHASE_D_STUDENT_UNDISTILLED_TAG = \"student_undistilled\"\n\n# -------------------------------------------------\n# Student baseline (không distill)\n# -------------------------------------------------\nCFG.KD_STUDENT_MODEL = \"custom_cnn_small\"\nCFG.PHASE_D_STUDENT_BASELINE_INIT_FROM_PHASE_C = False  # legacy flag; Phase D-only không dùng Phase C\nCFG.PHASE_D_STUDENT_BASELINE_INIT_CKPT_PATH = None\nCFG.PHASE_D_STUDENT_BASELINE_PRETRAINED = False\nCFG.PHASE_D_STUDENT_BASELINE_EPOCHS = 30\nCFG.PHASE_D_STUDENT_BASELINE_LR = float(getattr(CFG, \"LR\", 1e-3)) * 0.5\nCFG.PHASE_D_STUDENT_BASELINE_WEIGHT_DECAY = float(getattr(CFG, \"WEIGHT_DECAY\", 1e-4))\nCFG.PHASE_D_STUDENT_BASELINE_EARLY_STOP_PATIENCE = max(4, int(getattr(CFG, \"EARLY_STOP_PATIENCE\", 3)) + 1)\nCFG.PHASE_D_STUDENT_BASELINE_LOSS_NAME = \"ce\"\nCFG.PHASE_D_STUDENT_BASELINE_USE_CLASS_WEIGHTS_IN_LOSS = True\nCFG.PHASE_D_STUDENT_BASELINE_USE_WEIGHTED_SAMPLER = True\nCFG.PHASE_D_STUDENT_BASELINE_WEAK_AUG = getattr(CFG, \"WEAK_AUG\", False)\n\n# -------------------------------------------------\n# Knowledge distillation\n# -------------------------------------------------\nCFG.KD_LOGIT_METHOD = \"vanilla\"   # \"vanilla\" hoặc \"self_mckd\"\nCFG.KD_TEMPERATURE = 4.0\nCFG.KD_ALPHA_HARD = 0.50\nCFG.KD_BETA_TARGET = 1.00\nCFG.KD_BETA_NON_TARGET = 1.00\n\nCFG.KD_FEATURE_POOL = 4\nCFG.KD_FEATURE_LOSS = \"smooth_l1\"\nCFG.KD_SIMILARITY_LOSS = \"mse\"\n\nCFG.EPOCHS_KD = 30\nCFG.KD_LR = float(getattr(CFG, \"LR\", 1e-3)) * 0.5\nCFG.KD_WEIGHT_DECAY = float(getattr(CFG, \"WEIGHT_DECAY\", 1e-4))\nCFG.KD_EARLY_STOP_PATIENCE = max(4, int(getattr(CFG, \"EARLY_STOP_PATIENCE\", 3)) + 1)\nCFG.KD_LOSS_NAME = \"ce\"\n\n# Mặc định KD student vẫn train từ scratch để phép so sánh với baseline công bằng.\n# Nếu sau này muốn thử \"warm-start KD\", chỉ cần bật cờ dưới và truyền checkpoint phù hợp.\nCFG.KD_INIT_STUDENT_FROM_PHASE_C = False  # legacy flag; Phase D-only không dùng Phase C\nCFG.KD_INIT_STUDENT_CKPT_PATH = None\nCFG.KD_INIT_STUDENT_FROM_CHECKPOINT = False\nCFG.KD_USE_CLASS_WEIGHTS_IN_LOSS = True\nCFG.KD_USE_WEIGHTED_SAMPLER = True\nCFG.KD_WEAK_AUG = getattr(CFG, \"WEAK_AUG\", False)\nCFG.KD_PRETRAINED_STUDENT = False\nCFG.KD_USE_RETRAINED_TEACHER = True\n\n# -------------------------------------------------\n# Extra experiment: SqueezeNet student\n# -------------------------------------------------\nCFG.RUN_SQUEEZENET_UNDISTILLED = True\nCFG.RUN_SQUEEZENET_KD_FEATURE = True\nCFG.SQUEEZENET_MODEL_NAME = \"squeezenet1_1\"\nCFG.SQUEEZENET_PRETRAINED = False  # Kaggle-safe: không cần tải ImageNet weights\nCFG.SQUEEZENET_EPOCHS = CFG.PHASE_D_STUDENT_BASELINE_EPOCHS\nCFG.SQUEEZENET_LR = CFG.PHASE_D_STUDENT_BASELINE_LR\nCFG.SQUEEZENET_WEIGHT_DECAY = CFG.PHASE_D_STUDENT_BASELINE_WEIGHT_DECAY\nCFG.SQUEEZENET_EARLY_STOP_PATIENCE = CFG.PHASE_D_STUDENT_BASELINE_EARLY_STOP_PATIENCE\n\nprint(\"Phase D config summary:\")\nprint({\n    \"PHASE_D_DIR\": str(CFG.PHASE_D_DIR),\n    \"KD_DIR\": str(CFG.KD_DIR),\n    \"KD_TEACHER_MODEL\": CFG.KD_TEACHER_MODEL,\n    \"KD_STUDENT_MODEL\": CFG.KD_STUDENT_MODEL,\n    \"KD_SOURCES\": CFG.KD_SOURCES,\n    \"SQUEEZENET_MODEL_NAME\": CFG.SQUEEZENET_MODEL_NAME,\n    \"MEASURE_FLOPS\": CFG.MEASURE_FLOPS,\n    \"FLOP_BATCH_SIZE\": CFG.FLOP_BATCH_SIZE,\n})\n\nprint({\n    \"RUN_PHASE_D_TEACHER_RETRAIN\": CFG.RUN_PHASE_D_TEACHER_RETRAIN,\n    \"PHASE_D_TEACHER_EPOCHS\": CFG.PHASE_D_TEACHER_EPOCHS,\n    \"PHASE_D_TEACHER_LR\": CFG.PHASE_D_TEACHER_LR,\n    \"PHASE_D_TEACHER_INIT_FROM_PHASE_C\": CFG.PHASE_D_TEACHER_INIT_FROM_PHASE_C,\n    \"PHASE_D_TEACHER_INIT_CKPT_PATH\": CFG.PHASE_D_TEACHER_INIT_CKPT_PATH,\n})\n\nprint({\n    \"RUN_PHASE_D_STUDENT_BASELINE\": CFG.RUN_PHASE_D_STUDENT_BASELINE,\n    \"PHASE_D_STUDENT_BASELINE_EPOCHS\": CFG.PHASE_D_STUDENT_BASELINE_EPOCHS,\n    \"PHASE_D_STUDENT_BASELINE_LR\": CFG.PHASE_D_STUDENT_BASELINE_LR,\n})\n\nprint({\n    \"RUN_PHASE_D_KD\": CFG.RUN_PHASE_D_KD,\n    \"EPOCHS_KD\": CFG.EPOCHS_KD,\n    \"KD_LR\": CFG.KD_LR,\n    \"KD_LOGIT_METHOD\": CFG.KD_LOGIT_METHOD,\n    \"KD_USE_RETRAINED_TEACHER\": CFG.KD_USE_RETRAINED_TEACHER,\n})\n\n","metadata":{"id":"7662a25c","outputId":"f5b9c15d-52cb-478a-bc9f-c20cf1d160bf","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T10:57:51.000245Z","iopub.execute_input":"2026-05-09T10:57:51.000535Z","iopub.status.idle":"2026-05-09T10:57:51.018728Z","shell.execute_reply.started":"2026-05-09T10:57:51.000508Z","shell.execute_reply":"2026-05-09T10:57:51.017851Z"}},"outputs":[],"execution_count":null},{"id":"7c3f5739","cell_type":"markdown","source":"## D1.1. Khai báo teacher cho Phase D\n\nNotebook này chạy theo chế độ **Phase D-only**: chỉ dùng folder ảnh hiện tại và `trainLabels.csv`; không yêu cầu artifact nào từ các bước trước.\n\nTeacher được cố định là:\n\n```python\nefficientnet_b0\n```\n\nMặc định teacher sẽ được **train lại trong Phase D** từ folder ảnh hiện tại và `trainLabels.csv`. Checkpoint teacher mới sau khi train sẽ được dùng cho các nhánh KD.\n\n### Output chính\n- `teacher_phase_d_model_name = \"efficientnet_b0\"`\n- `teacher_phase_d_ckpt_path = None` trước khi train\n- `teacher_kd_model_name`, `teacher_kd_ckpt_path` sau bước teacher retrain\n","metadata":{"id":"7c3f5739"}},{"id":"219c04cb","cell_type":"code","source":"# ============================================================\n# Cell vai trò:\n# - Khai báo teacher cố định cho Phase D-only: efficientnet_b0.\n# - Không đọc screening summary, không tìm checkpoint Phase C.\n# - Tạo các biến compatibility để các cell sau chạy được khi notebook cũ còn dùng tên biến phase_c.\n# ============================================================\n\n# ============================================\n# D2. Declare fixed Phase-D teacher\n# ============================================\nrequire_torch()\n\nCFG.KD_TEACHER_MODEL = \"efficientnet_b0\"\nCFG.PHASE_D_TEACHER_INIT_FROM_PHASE_C = False\n\n# Phase D-only: không có screening/phase-c artifact.\nscreening_phase_c_csv_path = None\nscreening_raw_df = pd.DataFrame()\nscreening_phase_c_df = pd.DataFrame()\n\n\ndef _optional_path(value):\n    if value is None:\n        return None\n    value = str(value).strip()\n    if value == \"\" or value.lower() in {\"none\", \"nan\"}:\n        return None\n    p = Path(value)\n    return p if p.exists() else None\n\n\n# Optional manual checkpoint, nếu người dùng tự set trong CFG hoặc env.\nmanual_teacher_ckpt_path = _optional_path(getattr(CFG, \"PHASE_D_TEACHER_INIT_CKPT_PATH\", None))\nif manual_teacher_ckpt_path is None:\n    manual_teacher_ckpt_path = _optional_path(getattr(CFG, \"KD_TEACHER_CKPT_PATH\", None))\nif manual_teacher_ckpt_path is None:\n    manual_teacher_ckpt_path = _optional_path(os.environ.get(\"MALWARE_TEACHER_CKPT_PATH\"))\n\nteacher_phase_d_model_name = str(CFG.KD_TEACHER_MODEL)\nteacher_phase_d_ckpt_path = manual_teacher_ckpt_path\n\n# Compatibility aliases cho các cell cũ: tên phase_c nhưng giá trị là Phase-D-only.\nteacher_phase_c_model_name = teacher_phase_d_model_name\nteacher_phase_c_ckpt_path = teacher_phase_d_ckpt_path\n\n# Trước khi retrain, KD teacher có thể là checkpoint manual nếu có; nếu không sẽ được set sau cell teacher retrain.\nteacher_kd_model_name = teacher_phase_d_model_name\nteacher_kd_ckpt_path = teacher_phase_d_ckpt_path\nteacher_kd_source = \"manual_teacher_checkpoint\" if teacher_kd_ckpt_path is not None else \"phase_d_teacher_retrain_pending\"\n\n\ndef _resolve_phase_c_checkpoint(*args, **kwargs):\n    \"\"\"Phase D-only compatibility: không resolve checkpoint từ Phase C.\"\"\"\n    return None\n\n\ndef _resolve_student_baseline_score(*args, **kwargs):\n    \"\"\"Phase D-only compatibility: baseline chỉ lấy từ student train lại trong Phase D.\"\"\"\n    return None\n\nprint(\"Phase D-only teacher config:\")\nprint({\n    \"teacher_model\": teacher_phase_d_model_name,\n    \"manual_teacher_ckpt_path\": str(teacher_phase_d_ckpt_path) if teacher_phase_d_ckpt_path is not None else None,\n    \"will_retrain_teacher_in_phase_d\": bool(getattr(CFG, \"RUN_PHASE_D_TEACHER_RETRAIN\", True)),\n    \"uses_screening_summary\": False,\n})\n","metadata":{"id":"219c04cb","outputId":"6aaa52be-8e4e-4835-be41-039baafaae0d","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T10:57:51.019741Z","iopub.execute_input":"2026-05-09T10:57:51.020035Z","iopub.status.idle":"2026-05-09T10:57:51.036645Z","shell.execute_reply.started":"2026-05-09T10:57:51.020003Z","shell.execute_reply":"2026-05-09T10:57:51.035681Z"}},"outputs":[],"execution_count":null},{"id":"1416b8db","cell_type":"markdown","source":"## D2.1. Custom student CNN (small only)\n\n### Mục tiêu\nĐịnh nghĩa **student model riêng** cho bài toán malware image classification, thay vì dùng backbone nhỏ có sẵn trong thư viện.\n\n### Input\n- tensor ảnh đã được chuẩn hóa từ pipeline trước\n- số lớp `num_classes`\n\n### Output\n- lớp mô hình `CustomStudentCNN`\n- patch `build_model()` để notebook chỉ hỗ trợ student:\n  - `custom_cnn_small`\n\n### Thiết kế kiến trúc\n- CNN gọn nhẹ nhưng không quá tối giản\n- dùng:\n  - depthwise separable convolution\n  - residual connection\n  - squeeze-excitation nhẹ\n  - global average pooling\n\n### Lý do chỉ giữ `small`\n- đủ capacity để học từ teacher\n- đỡ tốn công so thêm một biến thể rất nhỏ\n- phù hợp mục tiêu hiện tại: **so sánh source of distillation**, không phải so kích thước student\n\n### Artifact\n- Không sinh file ở cell này\n- Nhưng đây là **định nghĩa trung tâm** của student dùng cho cả 3 nhánh KD\n","metadata":{"id":"1416b8db"}},{"id":"7dae76b8","cell_type":"code","source":"\n# ============================================================\n# Cell vai trò:\n# - Định nghĩa custom student CNN cho nhánh 2.\n# - Chỉ giữ lại biến thể \"small\" để notebook tập trung vào việc\n#   so sánh 3 sources of distillation thay vì so sánh thêm kích thước model.\n#\n# Ghi chú:\n# - build_model() được patch để hỗ trợ duy nhất \"custom_cnn_small\".\n# - Nếu vô tình gọi \"custom_cnn_tiny\", cell sẽ raise lỗi có chủ đích.\n# ============================================================\n\n# =====================================================\n# D2.1. Custom student CNN + patch build_model (small only)\n# =====================================================\nrequire_torch()\n\nclass ConvBNAct(nn.Module):\n    \"\"\"Khối cơ sở: Conv -> BN -> activation.\"\"\"\n    def __init__(self, in_ch, out_ch, kernel_size=3, stride=1, groups=1, act_layer=nn.SiLU):\n        super().__init__()\n        padding = kernel_size // 2\n        self.block = nn.Sequential(\n            nn.Conv2d(in_ch, out_ch, kernel_size, stride=stride, padding=padding, groups=groups, bias=False),\n            nn.BatchNorm2d(out_ch),\n            act_layer(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.block(x)\n\n\nclass SqueezeExciteLite(nn.Module):\n    \"\"\"SE gọn nhẹ để tái cân bằng kênh đặc trưng với chi phí thấp.\"\"\"\n    def __init__(self, channels: int, reduction: int = 8):\n        super().__init__()\n        hidden = max(8, channels // reduction)\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Sequential(\n            nn.Conv2d(channels, hidden, kernel_size=1, bias=True),\n            nn.SiLU(inplace=True),\n            nn.Conv2d(hidden, channels, kernel_size=1, bias=True),\n            nn.Sigmoid(),\n        )\n\n    def forward(self, x):\n        w = self.fc(self.pool(x))\n        return x * w\n\n\nclass DepthwiseSeparableResidualBlock(nn.Module):\n    \"\"\"Bottleneck gọn nhẹ cho student: pointwise -> depthwise -> SE -> projection.\"\"\"\n    def __init__(self, in_ch, out_ch, stride=1, expand_ratio=2.0, use_se=True):\n        super().__init__()\n        hidden_ch = int(round(in_ch * expand_ratio))\n        self.use_residual = (stride == 1 and in_ch == out_ch)\n\n        self.expand = ConvBNAct(in_ch, hidden_ch, kernel_size=1, stride=1)\n        self.depthwise = ConvBNAct(hidden_ch, hidden_ch, kernel_size=3, stride=stride, groups=hidden_ch)\n        self.se = SqueezeExciteLite(hidden_ch) if use_se else nn.Identity()\n        self.project = nn.Sequential(\n            nn.Conv2d(hidden_ch, out_ch, kernel_size=1, bias=False),\n            nn.BatchNorm2d(out_ch),\n        )\n        self.act = nn.SiLU(inplace=True)\n\n        if not self.use_residual:\n            self.shortcut = nn.Sequential()\n            if stride != 1 or in_ch != out_ch:\n                self.shortcut = nn.Sequential(\n                    nn.Conv2d(in_ch, out_ch, kernel_size=1, stride=stride, bias=False),\n                    nn.BatchNorm2d(out_ch),\n                )\n        else:\n            self.shortcut = nn.Identity()\n\n    def forward(self, x):\n        identity = x\n        x = self.expand(x)\n        x = self.depthwise(x)\n        x = self.se(x)\n        x = self.project(x)\n        if self.use_residual:\n            x = x + identity\n        else:\n            x = x + self.shortcut(identity)\n        return self.act(x)\n\n\nclass CustomStudentCNN(nn.Module):\n    \"\"\"\n    Student CNN tự thiết kế cho malware image classification.\n\n    Triết lý:\n    - đủ nhỏ để chạy nhanh,\n    - đủ sâu để học texture/pattern của ảnh malware,\n    - đủ ổn định để nhận tri thức từ teacher qua 3 source khác nhau.\n    \"\"\"\n    def __init__(self, num_classes: int, in_channels: int = 3, width_mult: float = 1.0, dropout: float = 0.20):\n        super().__init__()\n\n        def c(ch):\n            return max(16, int(round(ch * width_mult)))\n\n        self.stem = nn.Sequential(\n            ConvBNAct(in_channels, c(24), kernel_size=3, stride=2),\n            ConvBNAct(c(24), c(24), kernel_size=3, stride=1),\n        )\n        self.stage1 = nn.Sequential(\n            DepthwiseSeparableResidualBlock(c(24), c(32), stride=1, expand_ratio=2.0, use_se=True),\n            DepthwiseSeparableResidualBlock(c(32), c(32), stride=1, expand_ratio=2.0, use_se=True),\n        )\n        self.stage2 = nn.Sequential(\n            DepthwiseSeparableResidualBlock(c(32), c(64), stride=2, expand_ratio=2.0, use_se=True),\n            DepthwiseSeparableResidualBlock(c(64), c(64), stride=1, expand_ratio=2.5, use_se=True),\n        )\n        self.stage3 = nn.Sequential(\n            DepthwiseSeparableResidualBlock(c(64), c(96), stride=2, expand_ratio=3.0, use_se=True),\n            DepthwiseSeparableResidualBlock(c(96), c(96), stride=1, expand_ratio=3.0, use_se=True),\n        )\n        self.stage4 = nn.Sequential(\n            DepthwiseSeparableResidualBlock(c(96), c(160), stride=2, expand_ratio=3.0, use_se=True),\n            DepthwiseSeparableResidualBlock(c(160), c(160), stride=1, expand_ratio=3.5, use_se=True),\n        )\n\n        # head conv này sẽ là điểm móc thuận tiện cho feature-based KD.\n        self.head_conv = ConvBNAct(c(160), c(256), kernel_size=1, stride=1)\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.dropout = nn.Dropout(dropout)\n        self.classifier = nn.Linear(c(256), num_classes)\n\n        self._init_weights()\n\n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight, mode=\"fan_out\", nonlinearity=\"relu\")\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.ones_(m.weight)\n                nn.init.zeros_(m.bias)\n            elif isinstance(m, nn.Linear):\n                nn.init.normal_(m.weight, mean=0.0, std=0.01)\n                nn.init.zeros_(m.bias)\n\n    def forward_features(self, x):\n        x = self.stem(x)\n        x = self.stage1(x)\n        x = self.stage2(x)\n        x = self.stage3(x)\n        x = self.stage4(x)\n        x = self.head_conv(x)\n        return x\n\n    def forward(self, x):\n        feat = self.forward_features(x)\n        x = self.pool(feat).flatten(1)\n        x = self.dropout(x)\n        x = self.classifier(x)\n        return x\n\n\n\ndef _make_new_conv_phase_d(conv: nn.Conv2d, in_channels: int) -> nn.Conv2d:\n    new_conv = nn.Conv2d(\n        in_channels=in_channels,\n        out_channels=conv.out_channels,\n        kernel_size=conv.kernel_size,\n        stride=conv.stride,\n        padding=conv.padding,\n        dilation=conv.dilation,\n        groups=conv.groups,\n        bias=(conv.bias is not None),\n        padding_mode=conv.padding_mode,\n    )\n    with torch.no_grad():\n        if conv.weight.shape[1] == in_channels:\n            new_conv.weight.copy_(conv.weight)\n        elif in_channels == 1:\n            new_conv.weight.copy_(conv.weight.mean(dim=1, keepdim=True))\n        elif conv.weight.shape[1] == 1 and in_channels == 3:\n            new_conv.weight.copy_(conv.weight.repeat(1, 3, 1, 1) / 3.0)\n        else:\n            min_c = min(conv.weight.shape[1], in_channels)\n            new_conv.weight.zero_()\n            new_conv.weight[:, :min_c].copy_(conv.weight[:, :min_c])\n        if conv.bias is not None and new_conv.bias is not None:\n            new_conv.bias.copy_(conv.bias)\n    return new_conv\n\ndef _replace_first_conv_phase_d(module: nn.Module, in_channels: int) -> bool:\n    for name, child in module.named_children():\n        if isinstance(child, nn.Conv2d):\n            setattr(module, name, _make_new_conv_phase_d(child, in_channels))\n            return True\n        if _replace_first_conv_phase_d(child, in_channels):\n            return True\n    return False\n\ndef _replace_last_linear_phase_d(module: nn.Module, num_classes: int) -> bool:\n    children = list(module.named_children())\n    for name, child in reversed(children):\n        if isinstance(child, nn.Linear):\n            setattr(module, name, nn.Linear(child.in_features, num_classes))\n            return True\n        if _replace_last_linear_phase_d(child, num_classes):\n            return True\n    return False\n\ndef _set_classifier_head_phase_d(model: nn.Module, model_name: str, num_classes: int) -> nn.Module:\n    name = str(model_name).lower()\n\n    if hasattr(model, \"fc\") and isinstance(model.fc, nn.Linear):\n        model.fc = nn.Linear(model.fc.in_features, num_classes)\n        return model\n\n    if hasattr(model, \"classifier\"):\n        if isinstance(model.classifier, nn.Linear):\n            model.classifier = nn.Linear(model.classifier.in_features, num_classes)\n            return model\n        if isinstance(model.classifier, nn.Sequential):\n            if _replace_last_linear_phase_d(model.classifier, num_classes):\n                return model\n\n    if hasattr(model, \"heads\"):\n        if isinstance(model.heads, nn.Linear):\n            model.heads = nn.Linear(model.heads.in_features, num_classes)\n            return model\n        if hasattr(model.heads, \"head\") and isinstance(model.heads.head, nn.Linear):\n            model.heads.head = nn.Linear(model.heads.head.in_features, num_classes)\n            return model\n        if _replace_last_linear_phase_d(model.heads, num_classes):\n            return model\n\n    if hasattr(model, \"head\") and isinstance(model.head, nn.Linear):\n        model.head = nn.Linear(model.head.in_features, num_classes)\n        return model\n\n    if name.startswith(\"squeezenet\") and hasattr(model, \"classifier\") and isinstance(model.classifier, nn.Sequential):\n        for i, m in enumerate(model.classifier):\n            if isinstance(m, nn.Conv2d):\n                model.classifier[i] = nn.Conv2d(m.in_channels, num_classes, kernel_size=1, stride=1)\n                model.num_classes = num_classes\n                return model\n\n    raise ValueError(f\"Không biết cách thay classifier head cho model: {model_name}\")\n\ndef _build_torchvision_model_phase_d(\n    model_name: str = \"resnet18\",\n    num_classes: int = 9,\n    in_channels: int = 3,\n    pretrained: bool = True,\n):\n    if models is None:\n        raise RuntimeError(\"torchvision.models chưa sẵn sàng.\")\n\n    if not hasattr(models, model_name):\n        raise ValueError(f\"Model '{model_name}' không tồn tại trong torchvision.models\")\n\n    import inspect\n\n    model_fn = getattr(models, model_name)\n    sig = inspect.signature(model_fn)\n    kwargs = {}\n\n    if \"weights\" in sig.parameters:\n        if pretrained:\n            try:\n                weights_enum = models.get_model_weights(model_fn)\n                kwargs[\"weights\"] = weights_enum.DEFAULT\n            except Exception:\n                kwargs[\"weights\"] = None\n        else:\n            kwargs[\"weights\"] = None\n    elif \"pretrained\" in sig.parameters:\n        kwargs[\"pretrained\"] = bool(pretrained)\n\n    model = model_fn(**kwargs)\n\n    if int(in_channels) != 3:\n        ok = _replace_first_conv_phase_d(model, int(in_channels))\n        if not ok:\n            raise ValueError(f\"Không thay được first conv cho model: {model_name}\")\n\n    model = _set_classifier_head_phase_d(model, model_name, int(num_classes))\n    return model\n\nif \"_ORIGINAL_BUILD_MODEL_PHASE_D\" not in globals():\n    if \"build_model\" in globals():\n        _ORIGINAL_BUILD_MODEL_PHASE_D = build_model\n    else:\n        _ORIGINAL_BUILD_MODEL_PHASE_D = _build_torchvision_model_phase_d\n\ndef build_model(model_name=\"resnet18\", num_classes=9, in_channels=3, pretrained=True):\n    \"\"\"Factory thống nhất cho Phase D: hỗ trợ custom student, SqueezeNet và fallback teacher torchvision.\"\"\"\n    name = str(model_name).lower()\n    if name == \"squeezenet\":\n        model_name = \"squeezenet1_1\"\n        name = \"squeezenet1_1\"\n    if name == \"custom_cnn_small\":\n        return CustomStudentCNN(\n            num_classes=num_classes,\n            in_channels=in_channels,\n            width_mult=1.0,\n            dropout=0.20,\n        )\n    if name == \"custom_cnn_tiny\":\n        raise ValueError(\n            \"Notebook này chỉ giữ custom_cnn_small để tập trung so sánh 3 sources. \"\n            \"Hãy dùng CFG.KD_STUDENT_MODEL = 'custom_cnn_small'.\"\n        )\n    return _ORIGINAL_BUILD_MODEL_PHASE_D(\n        model_name=model_name,\n        num_classes=num_classes,\n        in_channels=in_channels,\n        pretrained=pretrained,\n    )\n\nCFG.KD_STUDENT_MODEL = \"custom_cnn_small\"\nCFG.KD_PRETRAINED_STUDENT = False\nCFG.KD_INIT_STUDENT_FROM_PHASE_C = False\nCFG.KD_INIT_STUDENT_CKPT_PATH = None\n\nstudent_model_name = str(CFG.KD_STUDENT_MODEL)\nstudent_phase_d_ckpt_path = None\nstudent_phase_c_ckpt_path = None  # compatibility alias; Phase D-only không dùng Phase C\n\n_tmp_student = build_model(\n    model_name=student_model_name,\n    num_classes=len(classes_sorted),\n    in_channels=CFG.IN_CHANNELS,\n    pretrained=False,\n)\nprint(\"Custom student ready:\", student_model_name)\nprint(\"Number of params   :\", f\"{count_parameters(_tmp_student):,}\")\nprint(\"Student init ckpt  :\", student_phase_d_ckpt_path)\ndel _tmp_student\n","metadata":{"id":"7dae76b8","outputId":"a4149f90-2f3c-47a5-946a-2e4d2d930c85","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T10:57:51.037812Z","iopub.execute_input":"2026-05-09T10:57:51.038469Z","iopub.status.idle":"2026-05-09T10:57:51.181472Z","shell.execute_reply.started":"2026-05-09T10:57:51.038432Z","shell.execute_reply":"2026-05-09T10:57:51.180882Z"}},"outputs":[],"execution_count":null},{"id":"f2ba7aa3","cell_type":"markdown","source":"## D2.2. Huấn luyện có giám sát trong Phase D\n\n### Mục tiêu\nDùng một khung train thống nhất cho:\n1. **teacher refresh**  \n2. **student supervised baseline**\n\n### Ý nghĩa\n- Hai bước này phải dùng **cùng split**, **cùng metric chọn best checkpoint** và **cùng logic early stopping**\n- Khi đó, so sánh giữa baseline và KD mới thực sự chặt chẽ\n","metadata":{"id":"f2ba7aa3"}},{"id":"145830bf","cell_type":"code","source":"# ============================================================\n# Cell vai trò:\n# - Tạo helper train supervised dùng chung cho teacher refresh và student baseline.\n# - Hỗ trợ:\n#   + train từ scratch\n#   + hoặc nạp checkpoint rồi train tiếp\n# - Best checkpoint vẫn chọn theo val_macro_f1 để nhất quán với toàn pipeline.\n# ============================================================\n\ndef load_state_dict_flexible_(model: nn.Module, checkpoint_path: Path, strict: bool = False):\n    state = torch.load(checkpoint_path, map_location=\"cpu\")\n    state_dict = state[\"model_state_dict\"] if isinstance(state, dict) and \"model_state_dict\" in state else state\n\n    cleaned_state_dict = {}\n    for k, v in state_dict.items():\n        nk = str(k)\n        if nk.startswith(\"module.\"):\n            nk = nk[len(\"module.\"):]\n        if nk.startswith(\"_orig_mod.\"):\n            nk = nk[len(\"_orig_mod.\"):]\n        cleaned_state_dict[nk] = v\n\n    load_result = model.load_state_dict(cleaned_state_dict, strict=strict)\n    return state, load_result\n\n\ndef run_phase_d_supervised_stage(\n    stage_name: str,\n    model_name: str,\n    train_df: pd.DataFrame,\n    val_df: pd.DataFrame,\n    test_df: pd.DataFrame,\n    num_classes: int,\n    idx_to_class: Dict[int, int],\n    output_dir: Path,\n    img_size: int = 224,\n    in_channels: int = 3,\n    batch_size: int = 32,\n    epochs: int = 10,\n    lr: float = 1e-3,\n    weight_decay: float = 1e-4,\n    loss_name: str = \"ce\",\n    use_class_weights_in_loss: bool = True,\n    use_weighted_sampler: bool = True,\n    num_workers: int = 2,\n    pretrained: bool = True,\n    weak_aug: bool = False,\n    early_stop_patience: int = 3,\n    seed: int = 42,\n    device: Optional[str] = None,\n    use_amp: Optional[bool] = None,\n    pin_memory: Optional[bool] = None,\n    persistent_workers: Optional[bool] = None,\n    prefetch_factor: Optional[int] = None,\n    init_checkpoint_path: Optional[Path] = None,\n    show_batch_progress: bool = True,\n    show_epoch_summary: bool = True,\n) -> Dict:\n    require_torch()\n    set_seed(seed)\n    output_dir = ensure_dir(output_dir)\n\n    runtime = get_runtime_config(\n        device=device,\n        use_amp=use_amp,\n        pin_memory=pin_memory,\n        persistent_workers=persistent_workers,\n        prefetch_factor=prefetch_factor,\n    )\n    device = runtime[\"device\"]\n\n    print(\"=\" * 110)\n    print(\n        f\"[{stage_name}] model={model_name} | device={device} | \"\n        f\"amp={runtime['use_amp']} | batch_size={batch_size} | epochs={epochs}\"\n    )\n    print(\"=\" * 110)\n\n    train_loader, val_loader, test_loader = build_loaders(\n        train_df=train_df,\n        val_df=val_df,\n        test_df=test_df,\n        img_size=img_size,\n        in_channels=in_channels,\n        batch_size=batch_size,\n        num_workers=num_workers,\n        use_imagenet_norm=True,\n        use_weighted_sampler=use_weighted_sampler,\n        weak_aug=weak_aug,\n        device=device,\n        pin_memory=runtime[\"pin_memory\"],\n        persistent_workers=runtime[\"persistent_workers\"],\n        prefetch_factor=runtime[\"prefetch_factor\"],\n    )\n\n    model = build_model(\n        model_name=model_name,\n        num_classes=num_classes,\n        in_channels=in_channels,\n        pretrained=pretrained,\n    )\n\n    init_mode = \"scratch\"\n    init_missing_keys = []\n    init_unexpected_keys = []\n    if init_checkpoint_path is not None:\n        init_checkpoint_path = Path(init_checkpoint_path)\n        if not init_checkpoint_path.exists():\n            raise FileNotFoundError(f\"Init checkpoint không tồn tại: {init_checkpoint_path}\")\n        _, load_result = load_state_dict_flexible_(model, init_checkpoint_path, strict=False)\n        if hasattr(load_result, \"missing_keys\"):\n            init_missing_keys = list(load_result.missing_keys)\n        if hasattr(load_result, \"unexpected_keys\"):\n            init_unexpected_keys = list(load_result.unexpected_keys)\n        init_mode = \"resume_from_checkpoint\"\n        print(f\"[{stage_name}] Initialized from checkpoint: {init_checkpoint_path}\")\n        if len(init_missing_keys) > 0:\n            print(\n                f\"[{stage_name}] missing_keys:\",\n                init_missing_keys[:10],\n                \"...\" if len(init_missing_keys) > 10 else \"\",\n            )\n        if len(init_unexpected_keys) > 0:\n            print(\n                f\"[{stage_name}] unexpected_keys:\",\n                init_unexpected_keys[:10],\n                \"...\" if len(init_unexpected_keys) > 10 else \"\",\n            )\n\n    model = prepare_model_for_runtime(model, runtime)\n\n    class_weights = None\n    if use_class_weights_in_loss:\n        class_weights_np, _ = get_effective_num_class_weights(train_df[\"label\"].values, num_classes)\n        class_weights = torch.tensor(class_weights_np, dtype=torch.float32).to(device)\n\n    criterion = build_criterion(loss_name, class_weights=class_weights)\n    optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer,\n        mode=\"max\",\n        factor=0.5,\n        patience=max(1, early_stop_patience // 2),\n    )\n    scaler = torch.amp.GradScaler(runtime[\"device_type\"], enabled=runtime[\"use_amp\"])\n\n    best_score = -1.0\n    best_epoch = -1\n    best_state = None\n    patience_counter = 0\n    history = []\n\n    if runtime[\"use_cuda\"]:\n        torch.cuda.reset_peak_memory_stats()\n\n    t0 = time.perf_counter()\n    for epoch in range(1, epochs + 1):\n        train_loss, train_metrics = train_one_epoch(\n            model,\n            train_loader,\n            criterion,\n            optimizer,\n            device,\n            num_classes,\n            use_amp=runtime[\"use_amp\"],\n            scaler=scaler,\n            model_name=f\"{stage_name}:{model_name}\",\n            epoch_idx=epoch,\n            total_epochs=epochs,\n            show_batch_progress=show_batch_progress,\n        )\n\n        val_loss, val_metrics, _, _ = evaluate(\n            model,\n            val_loader,\n            criterion,\n            device,\n            num_classes,\n            use_amp=runtime[\"use_amp\"],\n            model_name=f\"{stage_name}:{model_name}\",\n            epoch_idx=epoch,\n            total_epochs=epochs,\n            stage_name=\"Val\",\n            show_batch_progress=False,\n        )\n\n        current_score = val_metrics[\"macro_f1\"]\n        scheduler.step(current_score)\n\n        memory_stats = format_cuda_memory() if runtime[\"use_cuda\"] else {}\n        history_row = {\n            \"epoch\": epoch,\n            \"train_loss\": train_loss,\n            \"val_loss\": val_loss,\n            \"train_accuracy\": train_metrics[\"accuracy\"],\n            \"val_accuracy\": val_metrics[\"accuracy\"],\n            \"train_macro_precision\": train_metrics[\"macro_precision\"],\n            \"val_macro_precision\": val_metrics[\"macro_precision\"],\n            \"train_macro_recall\": train_metrics[\"macro_recall\"],\n            \"val_macro_recall\": val_metrics[\"macro_recall\"],\n            \"train_macro_f1\": train_metrics[\"macro_f1\"],\n            \"val_macro_f1\": val_metrics[\"macro_f1\"],\n            \"train_weighted_f1\": train_metrics[\"weighted_f1\"],\n            \"val_weighted_f1\": val_metrics[\"weighted_f1\"],\n            \"train_bal_acc\": train_metrics[\"balanced_acc\"],\n            \"val_bal_acc\": val_metrics[\"balanced_acc\"],\n            \"lr\": optimizer.param_groups[0][\"lr\"],\n            \"cuda_allocated_gb\": memory_stats.get(\"allocated_gb\"),\n            \"cuda_reserved_gb\": memory_stats.get(\"reserved_gb\"),\n            \"cuda_max_allocated_gb\": memory_stats.get(\"max_allocated_gb\"),\n            \"cuda_max_reserved_gb\": memory_stats.get(\"max_reserved_gb\"),\n        }\n        history.append(history_row)\n\n        if show_epoch_summary:\n            peak_text = (\n                f\" | peak_vram={history_row['cuda_max_allocated_gb']:.2f}GB\"\n                if history_row[\"cuda_max_allocated_gb\"] is not None\n                else \"\"\n            )\n            print(\n                f\"[{stage_name}] Epoch {epoch:02d}/{epochs} | \"\n                f\"train_loss={train_loss:.4f} | val_loss={val_loss:.4f} | \"\n                f\"train_f1={train_metrics['macro_f1']:.4f} | val_f1={val_metrics['macro_f1']:.4f} | \"\n                f\"val_bacc={val_metrics['balanced_acc']:.4f} | lr={optimizer.param_groups[0]['lr']:.2e}\"\n                f\"{peak_text}\"\n            )\n\n        if current_score > best_score:\n            best_score = current_score\n            best_epoch = epoch\n            best_state = copy.deepcopy(model.state_dict())\n            patience_counter = 0\n            if show_epoch_summary:\n                print(f\"[{stage_name}] ↳ New best checkpoint at epoch {epoch} (val_macro_f1={best_score:.4f})\")\n        else:\n            patience_counter += 1\n            if show_epoch_summary:\n                print(f\"[{stage_name}] ↳ No improvement. patience={patience_counter}/{early_stop_patience}\")\n            if patience_counter >= early_stop_patience:\n                if show_epoch_summary:\n                    print(f\"[{stage_name}] Early stopping triggered at epoch {epoch}.\")\n                break\n\n    total_train_time_sec = time.perf_counter() - t0\n\n    stage_safe = re.sub(r\"[^a-zA-Z0-9_\\-]+\", \"_\", str(stage_name)).strip(\"_\").lower()\n    history_path = output_dir / f\"{stage_safe}_{model_name}_history.csv\"\n    history_df = pd.DataFrame(history)\n    history_df.to_csv(history_path, index=False)\n\n    if best_state is None:\n        raise RuntimeError(f\"{stage_name} không tạo được checkpoint tốt nhất.\")\n\n    model.load_state_dict(best_state)\n\n    test_loss, test_metrics, test_targets, test_preds = evaluate(\n        model,\n        test_loader,\n        criterion,\n        device,\n        num_classes,\n        use_amp=runtime[\"use_amp\"],\n        model_name=f\"{stage_name}:{model_name}\",\n        stage_name=\"Test\",\n        show_batch_progress=False,\n    )\n\n    cm = confusion_matrix(test_targets, test_preds, labels=list(range(num_classes)))\n    cm_df = pd.DataFrame(\n        cm,\n        index=[f\"true_{idx_to_class[i]}\" for i in range(num_classes)],\n        columns=[f\"pred_{idx_to_class[i]}\" for i in range(num_classes)],\n    )\n    cm_path = output_dir / f\"{stage_safe}_{model_name}_confusion_matrix.csv\"\n    cm_df.to_csv(cm_path)\n\n    per_class_df = pd.DataFrame(\n        {\n            \"Class\": [idx_to_class[i] for i in range(num_classes)],\n            \"precision\": test_metrics[\"per_class_precision\"],\n            \"recall\": test_metrics[\"per_class_recall\"],\n            \"f1\": test_metrics[\"per_class_f1\"],\n            \"support\": test_metrics[\"per_class_support\"],\n        }\n    )\n    per_class_path = output_dir / f\"{stage_safe}_{model_name}_per_class_metrics.csv\"\n    per_class_df.to_csv(per_class_path, index=False)\n\n    ckpt_path = output_dir / f\"{stage_safe}_{model_name}_best.pth\"\n    torch.save(\n        {\n            \"stage_name\": stage_name,\n            \"model_name\": model_name,\n            \"best_epoch\": best_epoch,\n            \"model_state_dict\": model.state_dict(),\n            \"device\": device,\n            \"use_amp\": runtime[\"use_amp\"],\n            \"runtime_config\": runtime,\n            \"init_mode\": init_mode,\n            \"init_checkpoint_path\": str(init_checkpoint_path) if init_checkpoint_path is not None else None,\n        },\n        ckpt_path,\n    )\n\n    final_memory_stats = format_cuda_memory() if runtime[\"use_cuda\"] else {}\n    efficiency_metrics = measure_inference_efficiency(\n        model=model,\n        loader=test_loader,\n        runtime=runtime,\n        use_amp=runtime[\"use_amp\"],\n        max_batches=int(getattr(CFG, \"EFFICIENCY_MAX_BATCHES\", 20)),\n        warmup_batches=int(getattr(CFG, \"EFFICIENCY_WARMUP_BATCHES\", 2)),\n        measure_flops=bool(getattr(CFG, \"MEASURE_FLOPS\", True)),\n        flop_batch_size=int(getattr(CFG, \"FLOP_BATCH_SIZE\", 1)),\n    )\n    model_total_params = count_total_parameters(model)\n    model_trainable_params = count_parameters(model)\n    model_footprint_mb = estimate_model_footprint_mb(model)\n    checkpoint_size_mb = float(ckpt_path.stat().st_size / (1024 ** 2)) if ckpt_path.exists() else float(\"nan\")\n\n    metrics = {\n        \"stage_name\": stage_name,\n        \"model_name\": model_name,\n        \"num_params\": int(model_trainable_params),\n        \"total_params\": int(model_total_params),\n        \"trainable_params\": int(model_trainable_params),\n        \"model_footprint_mb\": float(model_footprint_mb),\n        \"static_model_footprint_mb\": float(model_footprint_mb),\n        \"checkpoint_size_mb\": float(checkpoint_size_mb),\n        \"best_epoch\": int(best_epoch),\n        \"device\": device,\n        \"use_amp\": bool(runtime[\"use_amp\"]),\n        \"channels_last\": bool(runtime[\"channels_last\"]),\n        \"compile_model\": bool(runtime[\"compile_model\"]),\n        \"init_mode\": init_mode,\n        \"init_checkpoint_path\": str(init_checkpoint_path) if init_checkpoint_path is not None else None,\n        \"init_missing_keys\": init_missing_keys,\n        \"init_unexpected_keys\": init_unexpected_keys,\n        \"test_loss\": float(test_loss),\n        \"test_accuracy\": float(test_metrics[\"accuracy\"]),\n        \"test_macro_precision\": float(test_metrics[\"macro_precision\"]),\n        \"test_macro_recall\": float(test_metrics[\"macro_recall\"]),\n        \"test_macro_f1\": float(test_metrics[\"macro_f1\"]),\n        \"test_weighted_f1\": float(test_metrics[\"weighted_f1\"]),\n        \"test_micro_f1\": float(test_metrics[\"micro_f1\"]),\n        \"test_balanced_acc\": float(test_metrics[\"balanced_acc\"]),\n        \"test_mcc\": float(test_metrics[\"mcc\"]),\n        \"test_cohen_kappa\": float(test_metrics[\"cohen_kappa\"]),\n        \"total_train_time_sec\": float(total_train_time_sec),\n        \"avg_epoch_time_sec\": float(total_train_time_sec / max(1, len(history))),\n        **efficiency_metrics,\n        \"training_peak_cuda_allocated_mb\": float(final_memory_stats.get(\"max_allocated_gb\", 0.0) * 1024.0),\n        \"training_peak_cuda_reserved_mb\": float(final_memory_stats.get(\"max_reserved_gb\", 0.0) * 1024.0),\n        \"training_peak_memory_footprint_mb\": float(final_memory_stats.get(\"max_allocated_gb\", 0.0) * 1024.0) if runtime[\"use_cuda\"] else float(\"nan\"),\n        \"cuda_peak_allocated_gb\": float(final_memory_stats.get(\"max_allocated_gb\", 0.0)),\n        \"cuda_peak_reserved_gb\": float(final_memory_stats.get(\"max_reserved_gb\", 0.0)),\n        \"checkpoint_path\": str(ckpt_path),\n        \"history_path\": str(history_path),\n        \"confusion_matrix_path\": str(cm_path),\n        \"per_class_metrics_path\": str(per_class_path),\n        \"classification_report_text\": classification_report(\n            test_targets,\n            test_preds,\n            labels=list(range(num_classes)),\n            target_names=[f\"Class_{idx_to_class[i]}\" for i in range(num_classes)],\n            zero_division=0,\n        ),\n    }\n\n    save_json(metrics, output_dir / f\"{stage_safe}_{model_name}_test_metrics.json\")\n\n    print(\n        f\"[{stage_name}] DONE | best_epoch={best_epoch} | \"\n        f\"test_macro_f1={metrics['test_macro_f1']:.4f} | \"\n        f\"test_balanced_acc={metrics['test_balanced_acc']:.4f} | \"\n        f\"time={metrics['total_train_time_sec'] / 60:.2f} min\"\n    )\n\n    if runtime[\"use_cuda\"]:\n        torch.cuda.empty_cache()\n\n    return metrics\n\n","metadata":{"id":"145830bf","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T10:57:51.182746Z","iopub.execute_input":"2026-05-09T10:57:51.182972Z","iopub.status.idle":"2026-05-09T10:57:51.213332Z","shell.execute_reply.started":"2026-05-09T10:57:51.182951Z","shell.execute_reply":"2026-05-09T10:57:51.212653Z"}},"outputs":[],"execution_count":null},{"id":"441c77d5","cell_type":"markdown","source":"## D2.3. Train/retrain teacher trong Phase D trước khi distill\n\n### Mục tiêu\nDùng teacher từ Phase C làm điểm khởi đầu, sau đó train tiếp với số epoch lớn hơn để teacher hội tụ tốt hơn.\n\n### Ý nghĩa khoa học\n- KD phụ thuộc mạnh vào **chất lượng teacher**\n- Nếu teacher Phase C chỉ train ngắn, phân phối soft target có thể chưa đủ tốt\n- Vì vậy bước này giúp tăng độ tin cậy cho toàn bộ nhánh distillation phía sau\n","metadata":{"id":"441c77d5"}},{"id":"83e3cff6","cell_type":"code","source":"# ============================================================\n# D2.3. Train / retrain teacher trong chính Phase D\n# ============================================================\nteacher_refresh_metrics = None\n\nif not getattr(CFG, \"RUN_PHASE_D_TEACHER_RETRAIN\", True):\n    print(\"CFG.RUN_PHASE_D_TEACHER_RETRAIN = False -> bỏ qua bước train/retrain teacher.\")\n    if teacher_kd_ckpt_path is None and getattr(CFG, \"RUN_PHASE_D_KD\", True):\n        raise RuntimeError(\n            \"Không có teacher checkpoint cho KD. Hãy bật CFG.RUN_PHASE_D_TEACHER_RETRAIN=True \"\n            \"hoặc set CFG.KD_TEACHER_CKPT_PATH / CFG.PHASE_D_TEACHER_INIT_CKPT_PATH.\"\n        )\nelse:\n    teacher_init_ckpt = _optional_path(getattr(CFG, \"PHASE_D_TEACHER_INIT_CKPT_PATH\", None))\n    if teacher_init_ckpt is None:\n        teacher_init_ckpt = _optional_path(getattr(CFG, \"KD_TEACHER_CKPT_PATH\", None))\n\n    teacher_pretrained_flag = False if teacher_init_ckpt is not None else bool(getattr(CFG, \"PHASE_D_TEACHER_PRETRAINED\", False))\n\n    teacher_refresh_metrics = run_phase_d_supervised_stage(\n        stage_name=\"teacher_retrain_phase_d\",\n        model_name=teacher_phase_d_model_name,\n        train_df=train_df,\n        val_df=val_df,\n        test_df=test_df,\n        num_classes=len(classes_sorted),\n        idx_to_class=idx_to_class,\n        output_dir=ensure_dir(CFG.PHASE_D_TEACHER_DIR / str(teacher_phase_d_model_name)),\n        img_size=CFG.IMG_SIZE,\n        in_channels=CFG.IN_CHANNELS,\n        batch_size=CFG.BATCH_SIZE,\n        epochs=CFG.PHASE_D_TEACHER_EPOCHS,\n        lr=CFG.PHASE_D_TEACHER_LR,\n        weight_decay=CFG.PHASE_D_TEACHER_WEIGHT_DECAY,\n        loss_name=CFG.PHASE_D_TEACHER_LOSS_NAME,\n        use_class_weights_in_loss=CFG.PHASE_D_TEACHER_USE_CLASS_WEIGHTS_IN_LOSS,\n        use_weighted_sampler=CFG.PHASE_D_TEACHER_USE_WEIGHTED_SAMPLER,\n        num_workers=CFG.NUM_WORKERS,\n        pretrained=teacher_pretrained_flag,\n        weak_aug=CFG.PHASE_D_TEACHER_WEAK_AUG,\n        early_stop_patience=CFG.PHASE_D_TEACHER_EARLY_STOP_PATIENCE,\n        seed=CFG.SEED,\n        device=CFG.DEVICE,\n        use_amp=CFG.USE_AMP,\n        pin_memory=CFG.PIN_MEMORY,\n        persistent_workers=CFG.PERSISTENT_WORKERS,\n        prefetch_factor=CFG.PREFETCH_FACTOR,\n        init_checkpoint_path=teacher_init_ckpt,\n        show_batch_progress=True,\n        show_epoch_summary=True,\n    )\n\nif teacher_refresh_metrics is not None and getattr(CFG, \"KD_USE_RETRAINED_TEACHER\", True):\n    teacher_kd_model_name = teacher_phase_d_model_name\n    teacher_kd_ckpt_path = Path(teacher_refresh_metrics[\"checkpoint_path\"])\n    teacher_kd_source = \"phase_d_teacher_retrain\"\nelse:\n    teacher_kd_model_name = teacher_phase_d_model_name\n    teacher_kd_ckpt_path = teacher_phase_d_ckpt_path\n    teacher_kd_source = \"manual_teacher_checkpoint\" if teacher_kd_ckpt_path is not None else \"missing_teacher_checkpoint\"\n\nif teacher_kd_ckpt_path is None and getattr(CFG, \"RUN_PHASE_D_KD\", True):\n    raise RuntimeError(\n        \"teacher_kd_ckpt_path vẫn None nên chưa thể chạy KD. \"\n        \"Hãy chạy cell train/retrain teacher hoặc set CFG.KD_TEACHER_CKPT_PATH.\"\n    )\n\nprint(\"Teacher used for KD:\")\nprint({\n    \"teacher_kd_model_name\": teacher_kd_model_name,\n    \"teacher_kd_ckpt_path\": str(teacher_kd_ckpt_path) if teacher_kd_ckpt_path is not None else None,\n    \"teacher_kd_source\": teacher_kd_source,\n})\n","metadata":{"id":"83e3cff6","outputId":"e9d2f02b-a35d-428b-b26a-9ac95f97de21","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T10:57:51.215485Z","iopub.execute_input":"2026-05-09T10:57:51.215862Z","iopub.status.idle":"2026-05-09T11:06:18.547544Z","shell.execute_reply.started":"2026-05-09T10:57:51.215836Z","shell.execute_reply":"2026-05-09T11:06:18.546840Z"}},"outputs":[],"execution_count":null},{"id":"d0d17835","cell_type":"markdown","source":"## D2.4. Train student supervised baseline\n\n### Mục tiêu\nHuấn luyện **đúng student architecture** nhưng **không distillation** để lấy mốc so sánh trực tiếp.\n\n### Ý nghĩa khoa học\nĐây là baseline quan trọng nhất của Phase D.  \nNếu KD không vượt được baseline này, thì rất khó nói distillation mang lại lợi ích thực sự.\n","metadata":{"id":"d0d17835"}},{"id":"63959805","cell_type":"code","source":"# ============================================================\n# D2.4. Train student undistilled (supervised baseline) (no distillation)\n# ============================================================\nstudent_baseline_metrics = None\nstudent_baseline_ckpt_path = None\n\nif not getattr(CFG, \"RUN_PHASE_D_STUDENT_BASELINE\", True):\n    print(\"CFG.RUN_PHASE_D_STUDENT_BASELINE = False -> bỏ qua student undistilled.\")\nelse:\n    student_init_ckpt = _optional_path(getattr(CFG, \"PHASE_D_STUDENT_BASELINE_INIT_CKPT_PATH\", None))\n    student_pretrained_flag = False if student_init_ckpt is not None else bool(getattr(CFG, \"PHASE_D_STUDENT_BASELINE_PRETRAINED\", False))\n\n    student_baseline_metrics = run_phase_d_supervised_stage(\n        stage_name=\"student_supervised_baseline\",\n        model_name=student_model_name,\n        train_df=train_df,\n        val_df=val_df,\n        test_df=test_df,\n        num_classes=len(classes_sorted),\n        idx_to_class=idx_to_class,\n        output_dir=ensure_dir(CFG.PHASE_D_STUDENT_BASELINE_DIR / str(student_model_name)),\n        img_size=CFG.IMG_SIZE,\n        in_channels=CFG.IN_CHANNELS,\n        batch_size=CFG.BATCH_SIZE,\n        epochs=CFG.PHASE_D_STUDENT_BASELINE_EPOCHS,\n        lr=CFG.PHASE_D_STUDENT_BASELINE_LR,\n        weight_decay=CFG.PHASE_D_STUDENT_BASELINE_WEIGHT_DECAY,\n        loss_name=CFG.PHASE_D_STUDENT_BASELINE_LOSS_NAME,\n        use_class_weights_in_loss=CFG.PHASE_D_STUDENT_BASELINE_USE_CLASS_WEIGHTS_IN_LOSS,\n        use_weighted_sampler=CFG.PHASE_D_STUDENT_BASELINE_USE_WEIGHTED_SAMPLER,\n        num_workers=CFG.NUM_WORKERS,\n        pretrained=student_pretrained_flag,\n        weak_aug=CFG.PHASE_D_STUDENT_BASELINE_WEAK_AUG,\n        early_stop_patience=CFG.PHASE_D_STUDENT_BASELINE_EARLY_STOP_PATIENCE,\n        seed=CFG.SEED,\n        device=CFG.DEVICE,\n        use_amp=CFG.USE_AMP,\n        pin_memory=CFG.PIN_MEMORY,\n        persistent_workers=CFG.PERSISTENT_WORKERS,\n        prefetch_factor=CFG.PREFETCH_FACTOR,\n        init_checkpoint_path=student_init_ckpt,\n        show_batch_progress=True,\n        show_epoch_summary=True,\n    )\n    student_baseline_ckpt_path = Path(student_baseline_metrics[\"checkpoint_path\"])\n\nprint(\"Student supervised baseline:\")\nprint({\n    \"student_model_name\": student_model_name,\n    \"checkpoint_path\": str(student_baseline_ckpt_path) if student_baseline_ckpt_path is not None else None,\n    \"test_macro_f1\": None if student_baseline_metrics is None else student_baseline_metrics[\"test_macro_f1\"],\n    \"test_balanced_acc\": None if student_baseline_metrics is None else student_baseline_metrics[\"test_balanced_acc\"],\n})\n\nstudent_undistilled_metrics = student_baseline_metrics\nstudent_undistilled_ckpt_path = student_baseline_ckpt_path\n\nprint(\"Student undistilled:\")\nprint({\n    \"student_model_name\": student_model_name,\n    \"checkpoint_path\": str(student_undistilled_ckpt_path) if student_undistilled_ckpt_path is not None else None,\n    \"test_macro_f1\": None if student_undistilled_metrics is None else student_undistilled_metrics[\"test_macro_f1\"],\n    \"test_balanced_acc\": None if student_undistilled_metrics is None else student_undistilled_metrics[\"test_balanced_acc\"],\n})\n","metadata":{"id":"63959805","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T11:06:18.548713Z","iopub.execute_input":"2026-05-09T11:06:18.549077Z","iopub.status.idle":"2026-05-09T11:14:12.273688Z","shell.execute_reply.started":"2026-05-09T11:06:18.549050Z","shell.execute_reply":"2026-05-09T11:14:12.272694Z"}},"outputs":[],"execution_count":null},{"id":"41e20985","cell_type":"markdown","source":"## D3. Distillation theo 3 sources: logit / feature / similarity\n\n### Mục tiêu\nCài toàn bộ loss, hook và vòng train/eval cho ba cơ chế truyền tri thức khác nhau.\n\n### Input\n- `teacher_model`\n- `student_model = custom_cnn_small`\n- dataloader train / val / test\n- cấu hình từ `CFG`\n\n### Output\n- hàm train một epoch\n- hàm evaluate\n- hàm chạy hoàn chỉnh một thí nghiệm Phase D cho từng source\n\n### Ba nguồn tri thức đang được kiểm định\n1. **Logit-based**\n   - Student học từ phân phối đầu ra của teacher\n   - Rẻ nhất và thường ổn định nhất\n2. **Feature-based**\n   - Student học từ feature map trung gian\n   - Cần adapter/projection để khớp chiều\n3. **Similarity-based**\n   - Student học từ quan hệ giữa các mẫu trong batch\n   - Không ép khớp trực tiếp từng giá trị đặc trưng\n\n### Artifact sẽ được sinh ở bước chạy\n- log huấn luyện theo từng epoch\n- best checkpoint của từng source\n- file metrics val/test\n","metadata":{"id":"41e20985"}},{"id":"393e0cdb","cell_type":"code","source":"import math\nfrom typing import Tuple\n\ndef load_checkpoint_to_model(\n    model_name: str,\n    checkpoint_path: Path,\n    num_classes: int,\n    in_channels: int,\n    device: str,\n    pretrained: bool = False,\n):\n    model = build_model(\n        model_name=model_name,\n        num_classes=num_classes,\n        in_channels=in_channels,\n        pretrained=pretrained,\n    )\n    state = torch.load(checkpoint_path, map_location=\"cpu\")\n    state_dict = state[\"model_state_dict\"] if isinstance(state, dict) and \"model_state_dict\" in state else state\n\n    try:\n        model.load_state_dict(state_dict, strict=True)\n    except Exception:\n        if isinstance(state_dict, dict):\n            cleaned_state_dict = {}\n            for k, v in state_dict.items():\n                nk = str(k)\n                if nk.startswith(\"module.\"):\n                    nk = nk[len(\"module.\"):]\n                if nk.startswith(\"_orig_mod.\"):\n                    nk = nk[len(\"_orig_mod.\"):]\n                cleaned_state_dict[nk] = v\n            model.load_state_dict(cleaned_state_dict, strict=False)\n        else:\n            raise\n\n    model = model.to(device)\n    model.eval()\n    return model, state\n\nclass DistillationLoss(nn.Module):\n    \"\"\"\n    Logit-based objective.\n\n    Hỗ trợ 2 biến thể:\n    - vanilla: KL trên toàn bộ phân phối lớp\n    - self_mckd: tách target / non-target để tăng hiệu quả truyền tri thức\n    \"\"\"\n    def __init__(\n        self,\n        method: str = \"vanilla\",\n        temperature: float = 4.0,\n        alpha_hard: float = 0.5,\n        beta_target: float = 1.0,\n        beta_non_target: float = 1.0,\n        hard_criterion: Optional[nn.Module] = None,\n        eps: float = 1e-8,\n    ):\n        super().__init__()\n        self.method = str(method).lower()\n        self.temperature = float(temperature)\n        self.alpha_hard = float(alpha_hard)\n        self.beta_target = float(beta_target)\n        self.beta_non_target = float(beta_non_target)\n        self.hard_criterion = hard_criterion if hard_criterion is not None else nn.CrossEntropyLoss()\n        self.eps = float(eps)\n\n        if self.method not in {\"vanilla\", \"self_mckd\"}:\n            raise ValueError(f\"KD logit method chưa hỗ trợ: {method}\")\n\n    def _vanilla_soft_loss(self, student_logits, teacher_logits):\n        T = self.temperature\n        student_log_probs = F.log_softmax(student_logits / T, dim=1)\n        teacher_probs = F.softmax(teacher_logits / T, dim=1)\n        soft_loss = F.kl_div(student_log_probs, teacher_probs, reduction=\"batchmean\") * (T ** 2)\n        return soft_loss\n\n    def _self_mckd_soft_loss(self, student_logits, teacher_logits, targets):\n        T = self.temperature\n        teacher_probs = F.softmax(teacher_logits / T, dim=1).clamp_min(self.eps)\n        student_probs = F.softmax(student_logits / T, dim=1).clamp_min(self.eps)\n        num_classes = teacher_probs.size(1)\n        target_mask = F.one_hot(targets, num_classes=num_classes).bool()\n\n        p_t = teacher_probs.gather(1, targets.unsqueeze(1)).clamp(self.eps, 1.0 - self.eps)\n        q_t = student_probs.gather(1, targets.unsqueeze(1)).clamp(self.eps, 1.0 - self.eps)\n\n        teacher_binary = torch.cat([p_t, 1.0 - p_t], dim=1)\n        student_binary_log = torch.log(torch.cat([q_t, 1.0 - q_t], dim=1).clamp_min(self.eps))\n        target_loss = F.kl_div(student_binary_log, teacher_binary, reduction=\"batchmean\")\n\n        teacher_nt = teacher_probs.masked_fill(target_mask, 0.0)\n        student_nt = student_probs.masked_fill(target_mask, 0.0)\n        teacher_nt = teacher_nt / teacher_nt.sum(dim=1, keepdim=True).clamp_min(self.eps)\n        student_nt = student_nt / student_nt.sum(dim=1, keepdim=True).clamp_min(self.eps)\n        non_target_loss = F.kl_div(torch.log(student_nt.clamp_min(self.eps)), teacher_nt, reduction=\"batchmean\")\n\n        soft_loss = (self.beta_target * target_loss + self.beta_non_target * non_target_loss) * (T ** 2)\n        return soft_loss, target_loss.detach(), non_target_loss.detach()\n\n    def forward(self, student_logits, teacher_logits, targets):\n        hard_loss = self.hard_criterion(student_logits, targets)\n\n        if self.method == \"vanilla\":\n            soft_loss = self._vanilla_soft_loss(student_logits, teacher_logits)\n            target_loss = torch.tensor(float(\"nan\"), device=student_logits.device)\n            non_target_loss = torch.tensor(float(\"nan\"), device=student_logits.device)\n        else:\n            soft_loss, target_loss, non_target_loss = self._self_mckd_soft_loss(\n                student_logits, teacher_logits, targets\n            )\n\n        total_loss = self.alpha_hard * hard_loss + (1.0 - self.alpha_hard) * soft_loss\n        stats = {\n            \"hard_loss\": float(hard_loss.detach().item()),\n            \"soft_loss\": float(soft_loss.detach().item()),\n            \"target_soft_loss\": float(target_loss.detach().item()) if torch.isfinite(target_loss) else float(\"nan\"),\n            \"non_target_soft_loss\": float(non_target_loss.detach().item()) if torch.isfinite(non_target_loss) else float(\"nan\"),\n        }\n        return total_loss, stats\n\n\nclass FeatureMapHook:\n    \"\"\"Hook đơn giản để lấy output của một layer trung gian trong forward pass.\"\"\"\n    def __init__(self, module: nn.Module):\n        self.feature = None\n        self.handle = module.register_forward_hook(self._hook_fn)\n\n    def _hook_fn(self, module, inputs, output):\n        # Nhiều model trả tensor trực tiếp; nếu có tuple thì lấy phần đầu.\n        self.feature = output[0] if isinstance(output, tuple) else output\n\n    def clear(self):\n        self.feature = None\n\n    def close(self):\n        self.handle.remove()\n\n\ndef find_last_conv_module(model: nn.Module):\n    \"\"\"\n    Tìm Conv2d cuối cùng trong mạng.\n\n    Mục đích:\n    - teacher và student khác kiến trúc nhưng vẫn có thể lấy một feature map\n      mức cao tương đối gần đầu phân loại.\n    \"\"\"\n    last_name, last_module = None, None\n    for name, module in model.named_modules():\n        if isinstance(module, nn.Conv2d):\n            last_name, last_module = name, module\n    if last_module is None:\n        raise ValueError(\"Không tìm thấy nn.Conv2d nào để gắn feature hook.\")\n    return last_name, last_module\n\n\nclass FeatureAdapter(nn.Module):\n    \"\"\"\n    Adapter học được để chiếu feature student sang số kênh của teacher.\n\n    Dùng riêng cho feature-based KD vì teacher/student thường khác chiều kênh.\n    \"\"\"\n    def __init__(self, student_channels: int, teacher_channels: int):\n        super().__init__()\n        self.proj = nn.Sequential(\n            nn.Conv2d(student_channels, teacher_channels, kernel_size=1, bias=False),\n            nn.BatchNorm2d(teacher_channels),\n            nn.SiLU(inplace=True),\n        )\n\n    def forward(self, x):\n        return self.proj(x)\n\n\nclass SourceDistillationObjective(nn.Module):\n    \"\"\"\n    Objective chung cho ba source:\n    - logit: output distribution của teacher\n    - feature: feature map trung gian của teacher\n    - similarity: quan hệ pairwise giữa các mẫu trong batch\n    \"\"\"\n    def __init__(\n        self,\n        source: str,\n        alpha_hard: float,\n        hard_criterion: nn.Module,\n        logit_criterion: Optional[DistillationLoss] = None,\n        feature_loss_name: str = \"smooth_l1\",\n        similarity_loss_name: str = \"mse\",\n        eps: float = 1e-8,\n    ):\n        super().__init__()\n        self.source = str(source).lower()\n        self.alpha_hard = float(alpha_hard)\n        self.hard_criterion = hard_criterion\n        self.logit_criterion = logit_criterion\n        self.eps = float(eps)\n\n        if feature_loss_name == \"mse\":\n            self.feature_criterion = nn.MSELoss()\n        elif feature_loss_name == \"smooth_l1\":\n            self.feature_criterion = nn.SmoothL1Loss(beta=0.5)\n        else:\n            raise ValueError(f\"Feature loss chưa hỗ trợ: {feature_loss_name}\")\n\n        if similarity_loss_name == \"mse\":\n            self.similarity_criterion = nn.MSELoss()\n        elif similarity_loss_name == \"smooth_l1\":\n            self.similarity_criterion = nn.SmoothL1Loss(beta=0.5)\n        else:\n            raise ValueError(f\"Similarity loss chưa hỗ trợ: {similarity_loss_name}\")\n\n        if self.source not in {\"logit\", \"feature\", \"similarity\"}:\n            raise ValueError(f\"Distillation source chưa hỗ trợ: {source}\")\n\n    def _feature_loss(self, student_feature, teacher_feature, feature_adapter):\n        if feature_adapter is None:\n            raise ValueError(\"Feature-based KD cần feature_adapter để khớp chiều kênh.\")\n        pool_size = int(CFG.KD_FEATURE_POOL)\n        teacher_feature = F.adaptive_avg_pool2d(teacher_feature, output_size=(pool_size, pool_size)).detach()\n        student_feature = F.adaptive_avg_pool2d(student_feature, output_size=(pool_size, pool_size))\n        student_feature = feature_adapter(student_feature)\n        return self.feature_criterion(student_feature, teacher_feature)\n\n    def _similarity_loss(self, student_feature, teacher_feature):\n        # Dùng vector đặc trưng sau GAP để tạo sample-wise similarity matrix.\n        t_vec = F.adaptive_avg_pool2d(teacher_feature, output_size=1).flatten(1).detach()\n        s_vec = F.adaptive_avg_pool2d(student_feature, output_size=1).flatten(1)\n\n        t_vec = F.normalize(t_vec, p=2, dim=1, eps=self.eps)\n        s_vec = F.normalize(s_vec, p=2, dim=1, eps=self.eps)\n\n        t_sim = t_vec @ t_vec.t()\n        s_sim = s_vec @ s_vec.t()\n        return self.similarity_criterion(s_sim, t_sim)\n\n    def forward(\n        self,\n        student_logits,\n        teacher_logits,\n        targets,\n        student_feature=None,\n        teacher_feature=None,\n        feature_adapter=None,\n    ):\n        if self.source == \"logit\":\n            total_loss, stats = self.logit_criterion(student_logits, teacher_logits, targets)\n            stats[\"source_loss\"] = stats[\"soft_loss\"]\n            return total_loss, stats\n\n        hard_loss = self.hard_criterion(student_logits, targets)\n\n        if self.source == \"feature\":\n            source_loss = self._feature_loss(student_feature, teacher_feature, feature_adapter)\n        else:\n            source_loss = self._similarity_loss(student_feature, teacher_feature)\n\n        total_loss = self.alpha_hard * hard_loss + (1.0 - self.alpha_hard) * source_loss\n        stats = {\n            \"hard_loss\": float(hard_loss.detach().item()),\n            \"soft_loss\": float(source_loss.detach().item()),\n            \"target_soft_loss\": float(\"nan\"),\n            \"non_target_soft_loss\": float(\"nan\"),\n            \"source_loss\": float(source_loss.detach().item()),\n        }\n        return total_loss, stats\n\n\ndef infer_feature_adapter(\n    teacher_model: nn.Module,\n    student_model: nn.Module,\n    teacher_hook: FeatureMapHook,\n    student_hook: FeatureMapHook,\n    sample_images: torch.Tensor,\n    runtime: Dict,\n):\n    \"\"\"Suy ra adapter cho feature-based KD bằng 1 forward pass khởi tạo.\"\"\"\n    teacher_hook.clear()\n    student_hook.clear()\n\n    with torch.no_grad():\n        with torch.amp.autocast(\"cuda\", enabled=bool(runtime[\"use_amp\"]) and str(runtime[\"device\"]).startswith(\"cuda\")):\n            _ = _forward_logits(teacher_model, sample_images)\n            _ = _forward_logits(student_model, sample_images)\n\n    if teacher_hook.feature is None or student_hook.feature is None:\n        raise RuntimeError(\"Không lấy được feature map để khởi tạo adapter.\")\n\n    teacher_channels = int(teacher_hook.feature.shape[1])\n    student_channels = int(student_hook.feature.shape[1])\n\n    adapter = FeatureAdapter(student_channels=student_channels, teacher_channels=teacher_channels)\n    adapter = adapter.to(runtime[\"device\"])\n    if runtime.get(\"channels_last\", False):\n        adapter = adapter.to(memory_format=torch.channels_last)\n    return adapter\n\n\ndef train_one_epoch_kd_source(\n    student_model: nn.Module,\n    teacher_model: nn.Module,\n    loader: DataLoader,\n    objective: SourceDistillationObjective,\n    optimizer: torch.optim.Optimizer,\n    device: str,\n    num_classes: int,\n    source: str,\n    use_amp: bool = True,\n    scaler: Optional[torch.amp.GradScaler] = None,\n    runtime: Optional[Dict] = None,\n    model_name: str = \"student\",\n    epoch_idx: int = 1,\n    total_epochs: int = 1,\n    show_batch_progress: bool = True,\n    teacher_hook: Optional[FeatureMapHook] = None,\n    student_hook: Optional[FeatureMapHook] = None,\n    feature_adapter: Optional[nn.Module] = None,\n):\n    student_model.train()\n    teacher_model.eval()\n    if feature_adapter is not None:\n        feature_adapter.train()\n\n    total_loss = 0.0\n    running = {\"hard_loss\": 0.0, \"soft_loss\": 0.0, \"target_soft_loss\": 0.0, \"non_target_soft_loss\": 0.0, \"source_loss\": 0.0}\n    all_targets, all_preds = [], []\n\n    iterator = loader\n    if show_batch_progress:\n        iterator = tqdm(loader, desc=f\"[KD:{source}:{model_name}] Epoch {epoch_idx}/{total_epochs}\", leave=False)\n\n    for images, targets in iterator:\n        images, targets = move_batch_to_device(images, targets, runtime or {\"device\": device, \"channels_last\": False})\n\n        if teacher_hook is not None:\n            teacher_hook.clear()\n        if student_hook is not None:\n            student_hook.clear()\n\n        optimizer.zero_grad(set_to_none=True)\n\n        with torch.no_grad():\n            with torch.amp.autocast(\"cuda\", enabled=bool(use_amp) and str(device).startswith(\"cuda\")):\n                teacher_logits = _forward_logits(teacher_model, images)\n                teacher_feature = None if teacher_hook is None else teacher_hook.feature\n\n        with torch.amp.autocast(\"cuda\", enabled=bool(use_amp) and str(device).startswith(\"cuda\")):\n            student_logits = _forward_logits(student_model, images)\n            student_feature = None if student_hook is None else student_hook.feature\n            loss, stats = objective(\n                student_logits=student_logits,\n                teacher_logits=teacher_logits,\n                targets=targets,\n                student_feature=student_feature,\n                teacher_feature=teacher_feature,\n                feature_adapter=feature_adapter,\n            )\n\n        if scaler is not None and bool(use_amp) and str(device).startswith(\"cuda\"):\n            scaler.scale(loss).backward()\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            loss.backward()\n            optimizer.step()\n\n        total_loss += float(loss.detach().item()) * images.size(0)\n        for k in running:\n            v = stats.get(k, float(\"nan\"))\n            if not math.isnan(v):\n                running[k] += float(v) * images.size(0)\n\n        preds = student_logits.argmax(dim=1)\n        all_targets.extend(targets.detach().cpu().tolist())\n        all_preds.extend(preds.detach().cpu().tolist())\n\n    mean_loss = total_loss / max(1, len(loader.dataset))\n    mean_running = {k: v / max(1, len(loader.dataset)) for k, v in running.items()}\n    metrics = compute_metrics(np.array(all_targets), np.array(all_preds), num_classes=num_classes)\n    return mean_loss, metrics, mean_running\n\n\ndef evaluate_kd(\n    student_model: nn.Module,\n    loader: DataLoader,\n    criterion: nn.Module,\n    device: str,\n    num_classes: int,\n    use_amp: bool = True,\n    runtime: Optional[Dict] = None,\n):\n    student_model.eval()\n    total_loss = 0.0\n    all_targets, all_preds = [], []\n\n    for images, targets in loader:\n        images, targets = move_batch_to_device(images, targets, runtime or {\"device\": device, \"channels_last\": False})\n        with torch.amp.autocast(\"cuda\", enabled=bool(use_amp) and str(device).startswith(\"cuda\")):\n            logits = _forward_logits(student_model, images)\n            loss = criterion(logits, targets)\n        total_loss += float(loss.detach().item()) * images.size(0)\n\n        preds = logits.argmax(dim=1)\n        all_targets.extend(targets.detach().cpu().tolist())\n        all_preds.extend(preds.detach().cpu().tolist())\n\n    mean_loss = total_loss / max(1, len(loader.dataset))\n    metrics = compute_metrics(np.array(all_targets), np.array(all_preds), num_classes=num_classes)\n    return mean_loss, metrics, np.array(all_targets), np.array(all_preds)\n\n\ndef _resolve_student_baseline_score(model_name: str, screening_df: pd.DataFrame, *_ignored):\n    if screening_df is None or screening_df.empty or \"model_name\" not in screening_df.columns:\n        return None\n\n    sub = screening_df[screening_df[\"model_name\"].astype(str) == str(model_name)].copy()\n    if len(sub) == 0:\n        return None\n\n    score_col = \"test_macro_f1\" if \"test_macro_f1\" in sub.columns else (\"macro_f1\" if \"macro_f1\" in sub.columns else None)\n    bal_col = \"test_balanced_acc\" if \"test_balanced_acc\" in sub.columns else (\"balanced_acc\" if \"balanced_acc\" in sub.columns else None)\n    params_col = \"num_params\" if \"num_params\" in sub.columns else None\n\n    if score_col is not None:\n        sub[score_col] = pd.to_numeric(sub[score_col], errors=\"coerce\")\n        sub = sub.sort_values(score_col, ascending=False)\n\n    row = sub.iloc[0]\n    return {\n        \"source\": \"screening\",\n        \"macro_f1\": float(row[score_col]) if score_col is not None and pd.notna(row[score_col]) else None,\n        \"bal_acc\": float(row[bal_col]) if bal_col is not None and pd.notna(row[bal_col]) else None,\n        \"num_params\": int(row[params_col]) if params_col is not None and pd.notna(row[params_col]) else None,\n    }\n\n\ndef run_phase_d_kd_source(\n    teacher_model_name: str,\n    student_model_name: str,\n    teacher_checkpoint_path: Path,\n    distill_source: str,\n    student_checkpoint_path: Optional[Path] = None,\n    output_dir: Optional[Path] = None,\n    student_pretrained: Optional[bool] = None,\n    baseline_metrics_override: Optional[Dict] = None,\n    show_batch_progress: bool = True,\n):\n    \"\"\"\n    Điều phối một thí nghiệm KD cho đúng một source.\n\n    Bản cập nhật:\n    - hỗ trợ student torchvision như squeezenet1_1,\n    - có thể truyền pretrained riêng cho student,\n    - nếu có baseline_metrics_override thì delta KD được tính đúng theo baseline cùng kiến trúc,\n    - ghi thêm metrics hiệu suất: footprint, checkpoint size, throughput, latency.\n    \"\"\"\n    require_torch()\n    set_seed(CFG.SEED)\n\n    distill_source = str(distill_source).lower()\n    if output_dir is None:\n        output_dir = ensure_dir(CFG.KD_DIR / f\"{teacher_model_name}_to_{student_model_name}\" / distill_source)\n    else:\n        output_dir = ensure_dir(output_dir)\n\n    runtime = get_runtime_config(\n        device=CFG.DEVICE,\n        use_amp=CFG.USE_AMP,\n        pin_memory=CFG.PIN_MEMORY,\n        persistent_workers=CFG.PERSISTENT_WORKERS,\n        prefetch_factor=CFG.PREFETCH_FACTOR,\n    )\n    device = runtime[\"device\"]\n\n    train_loader, val_loader, test_loader = build_loaders(\n        train_df=train_df,\n        val_df=val_df,\n        test_df=test_df,\n        img_size=CFG.IMG_SIZE,\n        in_channels=CFG.IN_CHANNELS,\n        batch_size=CFG.BATCH_SIZE,\n        num_workers=CFG.NUM_WORKERS,\n        use_imagenet_norm=True,\n        use_weighted_sampler=CFG.KD_USE_WEIGHTED_SAMPLER,\n        weak_aug=CFG.KD_WEAK_AUG,\n        device=device,\n        pin_memory=runtime[\"pin_memory\"],\n        persistent_workers=runtime[\"persistent_workers\"],\n        prefetch_factor=runtime[\"prefetch_factor\"],\n    )\n\n    teacher_model, teacher_state = load_checkpoint_to_model(\n        model_name=teacher_model_name,\n        checkpoint_path=teacher_checkpoint_path,\n        num_classes=len(classes_sorted),\n        in_channels=CFG.IN_CHANNELS,\n        device=device,\n        pretrained=False,\n    )\n    teacher_model = prepare_model_for_runtime(teacher_model, runtime)\n    teacher_model.eval()\n    for p in teacher_model.parameters():\n        p.requires_grad = False\n\n    student_pretrained_flag = bool(getattr(CFG, \"KD_PRETRAINED_STUDENT\", False)) if student_pretrained is None else bool(student_pretrained)\n    student_model = build_model(\n        model_name=student_model_name,\n        num_classes=len(classes_sorted),\n        in_channels=CFG.IN_CHANNELS,\n        pretrained=student_pretrained_flag,\n    )\n\n    student_init_mode = \"scratch_or_torchvision_pretrained\" if student_checkpoint_path is None else \"resume_from_checkpoint\"\n    student_init_missing_keys = []\n    student_init_unexpected_keys = []\n    if student_checkpoint_path is not None:\n        student_checkpoint_path = Path(student_checkpoint_path)\n        if not student_checkpoint_path.exists():\n            raise FileNotFoundError(f\"Student checkpoint không tồn tại: {student_checkpoint_path}\")\n        _, load_result = load_state_dict_flexible_(student_model, student_checkpoint_path, strict=False)\n        if hasattr(load_result, \"missing_keys\"):\n            student_init_missing_keys = list(load_result.missing_keys)\n        if hasattr(load_result, \"unexpected_keys\"):\n            student_init_unexpected_keys = list(load_result.unexpected_keys)\n        print(f\"[KD:{distill_source}] Student initialized from checkpoint: {student_checkpoint_path}\")\n\n    student_model = prepare_model_for_runtime(student_model, runtime)\n\n    class_weights = None\n    if CFG.KD_USE_CLASS_WEIGHTS_IN_LOSS:\n        class_weights_np, _ = get_effective_num_class_weights(train_df[\"label\"].values, len(classes_sorted))\n        class_weights = torch.tensor(class_weights_np, dtype=torch.float32, device=device)\n\n    hard_criterion = build_criterion(CFG.KD_LOSS_NAME, class_weights=class_weights)\n\n    teacher_hook, student_hook, feature_adapter = None, None, None\n    extra_trainable_modules = []\n\n    if distill_source in {\"feature\", \"similarity\"}:\n        teacher_hook_name, teacher_hook_module = find_last_conv_module(teacher_model)\n        student_hook_name, student_hook_module = find_last_conv_module(student_model)\n        teacher_hook = FeatureMapHook(teacher_hook_module)\n        student_hook = FeatureMapHook(student_hook_module)\n        print(f\"[{distill_source}] Teacher hook -> {teacher_hook_name}\")\n        print(f\"[{distill_source}] Student hook -> {student_hook_name}\")\n\n        if distill_source == \"feature\":\n            sample_images, _ = next(iter(train_loader))\n            sample_images, _ = move_batch_to_device(\n                sample_images,\n                torch.zeros(sample_images.size(0), dtype=torch.long),\n                runtime,\n            )\n            feature_adapter = infer_feature_adapter(\n                teacher_model=teacher_model,\n                student_model=student_model,\n                teacher_hook=teacher_hook,\n                student_hook=student_hook,\n                sample_images=sample_images,\n                runtime=runtime,\n            )\n            extra_trainable_modules.append(feature_adapter)\n            print(\"[feature] Adapter initialized.\")\n\n    logit_criterion = DistillationLoss(\n        method=CFG.KD_LOGIT_METHOD,\n        temperature=CFG.KD_TEMPERATURE,\n        alpha_hard=CFG.KD_ALPHA_HARD,\n        beta_target=CFG.KD_BETA_TARGET,\n        beta_non_target=CFG.KD_BETA_NON_TARGET,\n        hard_criterion=hard_criterion,\n    )\n\n    objective = SourceDistillationObjective(\n        source=distill_source,\n        alpha_hard=CFG.KD_ALPHA_HARD,\n        hard_criterion=hard_criterion,\n        logit_criterion=logit_criterion,\n        feature_loss_name=CFG.KD_FEATURE_LOSS,\n        similarity_loss_name=CFG.KD_SIMILARITY_LOSS,\n    )\n\n    params = list(student_model.parameters())\n    for module in extra_trainable_modules:\n        params.extend(list(module.parameters()))\n\n    optimizer = torch.optim.AdamW(params, lr=CFG.KD_LR, weight_decay=CFG.KD_WEIGHT_DECAY)\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer,\n        mode=\"max\",\n        factor=0.5,\n        patience=max(1, CFG.KD_EARLY_STOP_PATIENCE // 2),\n    )\n    scaler = torch.amp.GradScaler(runtime[\"device_type\"], enabled=runtime[\"use_amp\"])\n\n    best_score = -1.0\n    best_epoch = -1\n    best_state = None\n    best_adapter_state = None\n    patience_counter = 0\n    history = []\n\n    if runtime[\"use_cuda\"]:\n        torch.cuda.reset_peak_memory_stats()\n\n    t0 = time.perf_counter()\n    for epoch in range(1, CFG.EPOCHS_KD + 1):\n        train_loss, train_metrics, train_extra = train_one_epoch_kd_source(\n            student_model=student_model,\n            teacher_model=teacher_model,\n            loader=train_loader,\n            objective=objective,\n            optimizer=optimizer,\n            device=device,\n            num_classes=len(classes_sorted),\n            source=distill_source,\n            use_amp=runtime[\"use_amp\"],\n            scaler=scaler,\n            runtime=runtime,\n            model_name=student_model_name,\n            epoch_idx=epoch,\n            total_epochs=CFG.EPOCHS_KD,\n            show_batch_progress=show_batch_progress,\n            teacher_hook=teacher_hook,\n            student_hook=student_hook,\n            feature_adapter=feature_adapter,\n        )\n\n        val_loss, val_metrics, _, _ = evaluate_kd(\n            student_model=student_model,\n            loader=val_loader,\n            criterion=hard_criterion,\n            device=device,\n            num_classes=len(classes_sorted),\n            use_amp=runtime[\"use_amp\"],\n            runtime=runtime,\n        )\n\n        scheduler.step(val_metrics[\"macro_f1\"])\n        memory_stats = format_cuda_memory() if runtime[\"use_cuda\"] else {}\n\n        row = {\n            \"epoch\": epoch,\n            \"distill_source\": distill_source,\n            \"train_loss\": train_loss,\n            \"val_loss\": val_loss,\n            \"train_accuracy\": train_metrics[\"accuracy\"],\n            \"val_accuracy\": val_metrics[\"accuracy\"],\n            \"train_macro_precision\": train_metrics[\"macro_precision\"],\n            \"val_macro_precision\": val_metrics[\"macro_precision\"],\n            \"train_macro_recall\": train_metrics[\"macro_recall\"],\n            \"val_macro_recall\": val_metrics[\"macro_recall\"],\n            \"train_macro_f1\": train_metrics[\"macro_f1\"],\n            \"val_macro_f1\": val_metrics[\"macro_f1\"],\n            \"train_weighted_f1\": train_metrics[\"weighted_f1\"],\n            \"val_weighted_f1\": val_metrics[\"weighted_f1\"],\n            \"train_bal_acc\": train_metrics[\"balanced_acc\"],\n            \"val_bal_acc\": val_metrics[\"balanced_acc\"],\n            \"hard_loss\": train_extra[\"hard_loss\"],\n            \"soft_loss\": train_extra[\"soft_loss\"],\n            \"target_soft_loss\": train_extra[\"target_soft_loss\"],\n            \"non_target_soft_loss\": train_extra[\"non_target_soft_loss\"],\n            \"source_loss\": train_extra[\"source_loss\"],\n            \"lr\": optimizer.param_groups[0][\"lr\"],\n            \"cuda_max_allocated_gb\": memory_stats.get(\"max_allocated_gb\"),\n            \"cuda_max_reserved_gb\": memory_stats.get(\"max_reserved_gb\"),\n        }\n        history.append(row)\n\n        peak_text = (\n            f\" | peak_vram={row['cuda_max_allocated_gb']:.2f}GB\"\n            if row[\"cuda_max_allocated_gb\"] is not None else \"\"\n        )\n        print(\n            f\"[KD:{distill_source}] Epoch {epoch:02d}/{CFG.EPOCHS_KD} | \"\n            f\"train_f1={train_metrics['macro_f1']:.4f} | val_f1={val_metrics['macro_f1']:.4f} | \"\n            f\"train_loss={train_loss:.4f} | val_loss={val_loss:.4f} | \"\n            f\"hard={train_extra['hard_loss']:.4f} | source={train_extra['source_loss']:.4f} | \"\n            f\"lr={optimizer.param_groups[0]['lr']:.2e}{peak_text}\"\n        )\n\n        current_score = val_metrics[\"macro_f1\"]\n        if current_score > best_score:\n            best_score = current_score\n            best_epoch = epoch\n            best_state = copy.deepcopy(student_model.state_dict())\n            best_adapter_state = copy.deepcopy(feature_adapter.state_dict()) if feature_adapter is not None else None\n            patience_counter = 0\n            print(f\"[KD:{distill_source}] ↳ New best checkpoint at epoch {epoch} (val_macro_f1={best_score:.4f})\")\n        else:\n            patience_counter += 1\n            print(f\"[KD:{distill_source}] ↳ No improvement. patience={patience_counter}/{CFG.KD_EARLY_STOP_PATIENCE}\")\n            if patience_counter >= CFG.KD_EARLY_STOP_PATIENCE:\n                print(f\"[KD:{distill_source}] Early stopping at epoch {epoch}.\")\n                break\n\n    total_train_time_sec = time.perf_counter() - t0\n\n    if best_state is None:\n        raise RuntimeError(f\"Phase D ({distill_source}) không lưu được best student state.\")\n\n    student_model.load_state_dict(best_state)\n    if feature_adapter is not None and best_adapter_state is not None:\n        feature_adapter.load_state_dict(best_adapter_state)\n\n    test_loss, test_metrics, test_targets, test_preds = evaluate_kd(\n        student_model=student_model,\n        loader=test_loader,\n        criterion=hard_criterion,\n        device=device,\n        num_classes=len(classes_sorted),\n        use_amp=runtime[\"use_amp\"],\n        runtime=runtime,\n    )\n\n    history_df = pd.DataFrame(history)\n    history_path = output_dir / f\"{student_model_name}_{distill_source}_history.csv\"\n    history_df.to_csv(history_path, index=False)\n\n    cm = confusion_matrix(test_targets, test_preds, labels=list(range(len(classes_sorted))))\n    cm_df = pd.DataFrame(\n        cm,\n        index=[f\"true_{idx_to_class[i]}\" for i in range(len(classes_sorted))],\n        columns=[f\"pred_{idx_to_class[i]}\" for i in range(len(classes_sorted))],\n    )\n    cm_path = output_dir / f\"{student_model_name}_{distill_source}_confusion_matrix.csv\"\n    cm_df.to_csv(cm_path)\n\n    per_class_df = pd.DataFrame(\n        {\n            \"Class\": [idx_to_class[i] for i in range(len(classes_sorted))],\n            \"precision\": test_metrics[\"per_class_precision\"],\n            \"recall\": test_metrics[\"per_class_recall\"],\n            \"f1\": test_metrics[\"per_class_f1\"],\n            \"support\": test_metrics[\"per_class_support\"],\n        }\n    )\n    per_class_path = output_dir / f\"{student_model_name}_{distill_source}_per_class_metrics.csv\"\n    per_class_df.to_csv(per_class_path, index=False)\n\n    ckpt_path = output_dir / f\"{student_model_name}_{distill_source}_best.pth\"\n    payload = {\n        \"teacher_model_name\": teacher_model_name,\n        \"student_model_name\": student_model_name,\n        \"distill_source\": distill_source,\n        \"best_epoch\": best_epoch,\n        \"model_state_dict\": student_model.state_dict(),\n        \"runtime_config\": runtime,\n        \"student_init_mode\": student_init_mode,\n        \"student_pretrained\": bool(student_pretrained_flag),\n        \"student_init_checkpoint_path\": str(student_checkpoint_path) if student_checkpoint_path is not None else None,\n        \"student_init_missing_keys\": student_init_missing_keys,\n        \"student_init_unexpected_keys\": student_init_unexpected_keys,\n        \"kd_config\": {\n            \"source\": distill_source,\n            \"logit_method\": CFG.KD_LOGIT_METHOD,\n            \"temperature\": CFG.KD_TEMPERATURE,\n            \"alpha_hard\": CFG.KD_ALPHA_HARD,\n            \"beta_target\": CFG.KD_BETA_TARGET,\n            \"beta_non_target\": CFG.KD_BETA_NON_TARGET,\n            \"feature_pool\": CFG.KD_FEATURE_POOL,\n        },\n    }\n    if feature_adapter is not None:\n        payload[\"feature_adapter_state_dict\"] = feature_adapter.state_dict()\n    torch.save(payload, ckpt_path)\n\n    final_memory_stats = format_cuda_memory() if runtime[\"use_cuda\"] else {}\n    efficiency_metrics = measure_inference_efficiency(\n        model=student_model,\n        loader=test_loader,\n        runtime=runtime,\n        use_amp=runtime[\"use_amp\"],\n        max_batches=int(getattr(CFG, \"EFFICIENCY_MAX_BATCHES\", 20)),\n        warmup_batches=int(getattr(CFG, \"EFFICIENCY_WARMUP_BATCHES\", 2)),\n        measure_flops=bool(getattr(CFG, \"MEASURE_FLOPS\", True)),\n        flop_batch_size=int(getattr(CFG, \"FLOP_BATCH_SIZE\", 1)),\n    )\n    model_total_params = count_total_parameters(student_model)\n    model_trainable_params = count_parameters(student_model)\n    model_footprint_mb = estimate_model_footprint_mb(student_model)\n    checkpoint_size_mb = float(ckpt_path.stat().st_size / (1024 ** 2)) if ckpt_path.exists() else float(\"nan\")\n\n    baseline_info = None\n    baseline_source = None\n\n    if baseline_metrics_override is not None:\n        baseline_info = {\n            \"macro_f1\": float(baseline_metrics_override[\"test_macro_f1\"]),\n            \"bal_acc\": float(baseline_metrics_override[\"test_balanced_acc\"]),\n            \"source\": str(baseline_metrics_override.get(\"stage_name\", \"baseline_override\")),\n        }\n        baseline_source = baseline_info[\"source\"]\n    elif (\n        \"student_baseline_metrics\" in globals()\n        and student_baseline_metrics is not None\n        and str(student_baseline_metrics.get(\"model_name\")) == str(student_model_name)\n    ):\n        baseline_info = {\n            \"macro_f1\": float(student_baseline_metrics[\"test_macro_f1\"]),\n            \"bal_acc\": float(student_baseline_metrics[\"test_balanced_acc\"]),\n            \"source\": \"phase_d_student_supervised_baseline\",\n        }\n        baseline_source = \"phase_d_student_supervised_baseline\"\n    else:\n        # Phase D-only: không lấy baseline từ screening/Phase C.\n        baseline_info = None\n        baseline_source = None\n\n    delta_macro_f1 = None if baseline_info is None else float(test_metrics[\"macro_f1\"] - baseline_info[\"macro_f1\"])\n    delta_bal_acc = None if baseline_info is None else float(test_metrics[\"balanced_acc\"] - baseline_info[\"bal_acc\"])\n\n    result = {\n        \"teacher_model_name\": teacher_model_name,\n        \"student_model_name\": student_model_name,\n        \"teacher_checkpoint_path\": str(teacher_checkpoint_path),\n        \"student_init_checkpoint_path\": str(student_checkpoint_path) if student_checkpoint_path is not None else None,\n        \"student_init_mode\": student_init_mode,\n        \"student_pretrained\": bool(student_pretrained_flag),\n        \"baseline_source\": baseline_source,\n        \"distill_source\": distill_source,\n        \"kd_method\": CFG.KD_LOGIT_METHOD if distill_source == \"logit\" else distill_source,\n        \"temperature\": float(CFG.KD_TEMPERATURE),\n        \"alpha_hard\": float(CFG.KD_ALPHA_HARD),\n        \"best_epoch\": int(best_epoch),\n        \"student_num_params\": int(model_trainable_params),\n        \"student_total_params\": int(model_total_params),\n        \"student_trainable_params\": int(model_trainable_params),\n        \"model_footprint_mb\": float(model_footprint_mb),\n        \"static_model_footprint_mb\": float(model_footprint_mb),\n        \"checkpoint_size_mb\": float(checkpoint_size_mb),\n        \"test_loss\": float(test_loss),\n        \"test_accuracy\": float(test_metrics[\"accuracy\"]),\n        \"test_macro_precision\": float(test_metrics[\"macro_precision\"]),\n        \"test_macro_recall\": float(test_metrics[\"macro_recall\"]),\n        \"test_macro_f1\": float(test_metrics[\"macro_f1\"]),\n        \"test_weighted_f1\": float(test_metrics[\"weighted_f1\"]),\n        \"test_micro_f1\": float(test_metrics[\"micro_f1\"]),\n        \"test_balanced_acc\": float(test_metrics[\"balanced_acc\"]),\n        \"test_mcc\": float(test_metrics[\"mcc\"]),\n        \"test_cohen_kappa\": float(test_metrics[\"cohen_kappa\"]),\n        \"baseline_macro_f1\": None if baseline_info is None else float(baseline_info[\"macro_f1\"]),\n        \"baseline_bal_acc\": None if baseline_info is None else float(baseline_info[\"bal_acc\"]),\n        \"delta_macro_f1_vs_student_baseline\": delta_macro_f1,\n        \"delta_bal_acc_vs_student_baseline\": delta_bal_acc,\n        # Backward-compatible aliases.\n        \"delta_macro_f1_vs_phase_c\": delta_macro_f1,\n        \"delta_bal_acc_vs_phase_c\": delta_bal_acc,\n        \"total_train_time_sec\": float(total_train_time_sec),\n        \"avg_epoch_time_sec\": float(total_train_time_sec / max(1, len(history))),\n        \"training_peak_cuda_allocated_mb\": float(final_memory_stats.get(\"max_allocated_gb\", 0.0) * 1024.0),\n        \"training_peak_cuda_reserved_mb\": float(final_memory_stats.get(\"max_reserved_gb\", 0.0) * 1024.0),\n        \"training_peak_memory_footprint_mb\": float(final_memory_stats.get(\"max_allocated_gb\", 0.0) * 1024.0) if runtime[\"use_cuda\"] else float(\"nan\"),\n        \"cuda_peak_allocated_gb\": float(final_memory_stats.get(\"max_allocated_gb\", 0.0)),\n        \"cuda_peak_reserved_gb\": float(final_memory_stats.get(\"max_reserved_gb\", 0.0)),\n        **efficiency_metrics,\n        \"checkpoint_path\": str(ckpt_path),\n        \"history_path\": str(history_path),\n        \"confusion_matrix_path\": str(cm_path),\n        \"per_class_metrics_path\": str(per_class_path),\n        \"classification_report_text\": classification_report(\n            test_targets,\n            test_preds,\n            labels=list(range(len(classes_sorted))),\n            target_names=[f\"Class_{idx_to_class[i]}\" for i in range(len(classes_sorted))],\n            zero_division=0,\n        ),\n    }\n\n    save_json(result, output_dir / f\"{student_model_name}_{distill_source}_test_metrics.json\")\n\n    if teacher_hook is not None:\n        teacher_hook.close()\n    if student_hook is not None:\n        student_hook.close()\n\n    return result\n\n","metadata":{"id":"393e0cdb","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T11:14:12.275193Z","iopub.execute_input":"2026-05-09T11:14:12.275785Z","iopub.status.idle":"2026-05-09T11:14:12.356288Z","shell.execute_reply.started":"2026-05-09T11:14:12.275752Z","shell.execute_reply":"2026-05-09T11:14:12.355701Z"}},"outputs":[],"execution_count":null},{"id":"4c379bbb","cell_type":"markdown","source":"## D4. Chạy distillation tách thành 3 cell riêng\n\nỞ phần này, notebook dùng:\n- **teacher đã refresh** nếu `CFG.KD_USE_RETRAINED_TEACHER = True`\n- **student custom_cnn_small**\n- cùng split với teacher refresh và student baseline\n\nBa cell tách riêng giúp bạn:\n- chạy từng phương pháp distillation độc lập,\n- debug từng nhánh dễ hơn,\n- tránh phải chạy lại cả 3 nhánh khi chỉ muốn thử một source.\n","metadata":{"id":"4c379bbb"}},{"id":"d6eae20e","cell_type":"code","source":"student_model_name = globals().get(\"student_model_name\", str(getattr(CFG, \"KD_STUDENT_MODEL\", \"custom_cnn_small\")))\nstudent_phase_d_ckpt_path = globals().get(\"student_phase_d_ckpt_path\", None)\nteacher_kd_model_name = globals().get(\"teacher_kd_model_name\", teacher_phase_d_model_name)\nteacher_kd_ckpt_path = globals().get(\"teacher_kd_ckpt_path\", None)\n\nif teacher_kd_ckpt_path is None:\n    raise RuntimeError(\n        \"Không có teacher_kd_ckpt_path. Hãy chạy cell D2.3 train/retrain teacher trước khi chạy KD.\"\n    )\n\n\n# ============================================================\n# D4.1. Chạy logit-based distillation\n# ============================================================\nif \"phase_d_results\" not in globals():\n    phase_d_results = []\n\nphase_d_results = [r for r in phase_d_results if r.get(\"distill_source\") != \"logit\"]\n\nif not getattr(CFG, \"RUN_PHASE_D_KD\", True):\n    print(\"CFG.RUN_PHASE_D_KD = False -> bỏ qua cell logit.\")\nelif \"logit\" not in CFG.KD_SOURCES:\n    print(\"'logit' không nằm trong CFG.KD_SOURCES -> bỏ qua.\")\nelse:\n    print(\"\\n\" + \"=\" * 110)\n    print(\"Running Phase D with source = logit\")\n    print(\"=\" * 110)\n\n    result = run_phase_d_kd_source(\n        teacher_model_name=teacher_kd_model_name,\n        student_model_name=student_model_name,\n        teacher_checkpoint_path=teacher_kd_ckpt_path,\n        student_checkpoint_path=None,\n        distill_source=\"logit\",\n        output_dir=CFG.KD_DIR / f\"{teacher_kd_model_name}_to_{student_model_name}\" / \"logit\",\n    )\n    phase_d_results.append(result)\n\npd.DataFrame(phase_d_results)\n","metadata":{"id":"d6eae20e","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T11:14:12.357215Z","iopub.execute_input":"2026-05-09T11:14:12.357497Z","iopub.status.idle":"2026-05-09T11:23:55.199748Z","shell.execute_reply.started":"2026-05-09T11:14:12.357473Z","shell.execute_reply":"2026-05-09T11:23:55.198587Z"}},"outputs":[],"execution_count":null},{"id":"28741581","cell_type":"code","source":"student_model_name = globals().get(\"student_model_name\", str(getattr(CFG, \"KD_STUDENT_MODEL\", \"custom_cnn_small\")))\nstudent_phase_d_ckpt_path = globals().get(\"student_phase_d_ckpt_path\", None)\nteacher_kd_model_name = globals().get(\"teacher_kd_model_name\", teacher_phase_d_model_name)\nteacher_kd_ckpt_path = globals().get(\"teacher_kd_ckpt_path\", None)\n\nif teacher_kd_ckpt_path is None:\n    raise RuntimeError(\n        \"Không có teacher_kd_ckpt_path. Hãy chạy cell D2.3 train/retrain teacher trước khi chạy KD.\"\n    )\n\n\n# ============================================================\n# D4.2. Chạy feature-based distillation\n# ============================================================\nif \"phase_d_results\" not in globals():\n    phase_d_results = []\n\nphase_d_results = [r for r in phase_d_results if r.get(\"distill_source\") != \"feature\"]\n\nif not getattr(CFG, \"RUN_PHASE_D_KD\", True):\n    print(\"CFG.RUN_PHASE_D_KD = False -> bỏ qua cell feature.\")\nelif \"feature\" not in CFG.KD_SOURCES:\n    print(\"'feature' không nằm trong CFG.KD_SOURCES -> bỏ qua.\")\nelse:\n    print(\"\\n\" + \"=\" * 110)\n    print(\"Running Phase D with source = feature\")\n    print(\"=\" * 110)\n\n    result = run_phase_d_kd_source(\n        teacher_model_name=teacher_kd_model_name,\n        student_model_name=student_model_name,\n        teacher_checkpoint_path=teacher_kd_ckpt_path,\n        student_checkpoint_path=None,\n        distill_source=\"feature\",\n        output_dir=CFG.KD_DIR / f\"{teacher_kd_model_name}_to_{student_model_name}\" / \"feature\",\n    )\n    phase_d_results.append(result)\n\npd.DataFrame(phase_d_results)\n","metadata":{"id":"28741581","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T11:23:55.201784Z","iopub.execute_input":"2026-05-09T11:23:55.202183Z","iopub.status.idle":"2026-05-09T11:33:52.922599Z","shell.execute_reply.started":"2026-05-09T11:23:55.202147Z","shell.execute_reply":"2026-05-09T11:33:52.921846Z"}},"outputs":[],"execution_count":null},{"id":"e3820e99","cell_type":"code","source":"student_model_name = globals().get(\"student_model_name\", str(getattr(CFG, \"KD_STUDENT_MODEL\", \"custom_cnn_small\")))\nstudent_phase_d_ckpt_path = globals().get(\"student_phase_d_ckpt_path\", None)\nteacher_kd_model_name = globals().get(\"teacher_kd_model_name\", teacher_phase_d_model_name)\nteacher_kd_ckpt_path = globals().get(\"teacher_kd_ckpt_path\", None)\n\nif teacher_kd_ckpt_path is None:\n    raise RuntimeError(\n        \"Không có teacher_kd_ckpt_path. Hãy chạy cell D2.3 train/retrain teacher trước khi chạy KD.\"\n    )\n\n\n# ============================================================\n# D4.3. Chạy similarity-based distillation\n# ============================================================\nif \"phase_d_results\" not in globals():\n    phase_d_results = []\n\nphase_d_results = [r for r in phase_d_results if r.get(\"distill_source\") != \"similarity\"]\n\nif not getattr(CFG, \"RUN_PHASE_D_KD\", True):\n    print(\"CFG.RUN_PHASE_D_KD = False -> bỏ qua cell similarity.\")\nelif \"similarity\" not in CFG.KD_SOURCES:\n    print(\"'similarity' không nằm trong CFG.KD_SOURCES -> bỏ qua.\")\nelse:\n    print(\"\\n\" + \"=\" * 110)\n    print(\"Running Phase D with source = similarity\")\n    print(\"=\" * 110)\n\n    result = run_phase_d_kd_source(\n        teacher_model_name=teacher_kd_model_name,\n        student_model_name=student_model_name,\n        teacher_checkpoint_path=teacher_kd_ckpt_path,\n        student_checkpoint_path=None,\n        distill_source=\"similarity\",\n        output_dir=CFG.KD_DIR / f\"{teacher_kd_model_name}_to_{student_model_name}\" / \"similarity\",\n    )\n    phase_d_results.append(result)\n\npd.DataFrame(phase_d_results)\n","metadata":{"id":"e3820e99","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T11:33:52.923830Z","iopub.execute_input":"2026-05-09T11:33:52.924135Z","iopub.status.idle":"2026-05-09T11:45:57.899125Z","shell.execute_reply.started":"2026-05-09T11:33:52.924100Z","shell.execute_reply":"2026-05-09T11:45:57.898325Z"}},"outputs":[],"execution_count":null},{"id":"cMklLwgctuGL","cell_type":"markdown","source":"## D4.4. Extra student — SqueezeNet baseline và SqueezeNet + KD-feature\n\n### Mục tiêu\nBổ sung đúng hai cấu hình bạn yêu cầu:\n- `squeezenet1_1` **không distill** / supervised baseline\n- `squeezenet1_1` + **feature-based KD**\n\nHai cấu hình này dùng cùng train/val/test split đã qua guard chống data leak. Khi tính delta cho SqueezeNet KD, baseline được lấy từ chính SqueezeNet undistilled, không dùng nhầm baseline của `custom_cnn_small`.\n","metadata":{"id":"cMklLwgctuGL"}},{"id":"wPsDhv6dtuGL","cell_type":"code","source":"# ============================================================\n# D4.4. Extra experiment: SqueezeNet undistilled + SqueezeNet KD-feature\n# ============================================================\nif \"extra_student_results\" not in globals():\n    extra_student_results = []\n\nsqueezenet_model_name = str(getattr(CFG, \"SQUEEZENET_MODEL_NAME\", \"squeezenet1_1\"))\nsqueezenet_baseline_metrics = globals().get(\"squeezenet_baseline_metrics\", None)\nsqueezenet_baseline_ckpt_path = None if squeezenet_baseline_metrics is None else Path(squeezenet_baseline_metrics[\"checkpoint_path\"])\n\nEFFICIENCY_RESULT_KEYS = [\n    \"static_model_footprint_mb\",\n    \"inference_peak_memory_footprint_mb\",\n    \"inference_peak_memory_footprint_source\",\n    \"inference_peak_cpu_rss_mb\",\n    \"inference_peak_cuda_allocated_mb\",\n    \"inference_peak_cuda_reserved_mb\",\n    \"training_peak_memory_footprint_mb\",\n    \"training_peak_cuda_allocated_mb\",\n    \"training_peak_cuda_reserved_mb\",\n    \"inference_time_sec\",\n    \"inference_eval_batches\",\n    \"inference_eval_images\",\n    \"inference_latency_ms_per_batch\",\n    \"forward_flops_per_img\",\n    \"forward_mflops_per_img\",\n    \"forward_gflops_per_img\",\n    \"forward_flops_per_batch\",\n    \"flops_method\",\n    \"flops_eval_batch_size\",\n    \"flops_error\",\n]\n\ndef _copy_efficiency_keys(src_metrics: Dict) -> Dict:\n    return {k: src_metrics.get(k) for k in EFFICIENCY_RESULT_KEYS if isinstance(src_metrics, dict) and k in src_metrics}\n\n# Xóa kết quả cũ của SqueezeNet trong RAM để tránh duplicate nếu chạy lại cell.\nextra_student_results = [\n    r for r in extra_student_results\n    if str(r.get(\"model_name\", r.get(\"student_model_name\", \"\"))) != squeezenet_model_name\n]\nif \"phase_d_results\" not in globals():\n    phase_d_results = []\nphase_d_results = [\n    r for r in phase_d_results\n    if not (str(r.get(\"student_model_name\", \"\")) == squeezenet_model_name and str(r.get(\"distill_source\", \"\")) == \"feature\")\n]\n\nif getattr(CFG, \"RUN_SQUEEZENET_UNDISTILLED\", True):\n    print(\"\\n\" + \"=\" * 110)\n    print(f\"Running SqueezeNet undistilled baseline: {squeezenet_model_name}\")\n    print(\"=\" * 110)\n\n    squeezenet_baseline_metrics = run_phase_d_supervised_stage(\n        stage_name=\"squeezenet_supervised_baseline\",\n        model_name=squeezenet_model_name,\n        train_df=train_df,\n        val_df=val_df,\n        test_df=test_df,\n        num_classes=len(classes_sorted),\n        idx_to_class=idx_to_class,\n        output_dir=ensure_dir(CFG.PHASE_D_STUDENT_BASELINE_DIR / squeezenet_model_name),\n        img_size=CFG.IMG_SIZE,\n        in_channels=CFG.IN_CHANNELS,\n        batch_size=CFG.BATCH_SIZE,\n        epochs=int(getattr(CFG, \"SQUEEZENET_EPOCHS\", CFG.PHASE_D_STUDENT_BASELINE_EPOCHS)),\n        lr=float(getattr(CFG, \"SQUEEZENET_LR\", CFG.PHASE_D_STUDENT_BASELINE_LR)),\n        weight_decay=float(getattr(CFG, \"SQUEEZENET_WEIGHT_DECAY\", CFG.PHASE_D_STUDENT_BASELINE_WEIGHT_DECAY)),\n        loss_name=CFG.PHASE_D_STUDENT_BASELINE_LOSS_NAME,\n        use_class_weights_in_loss=CFG.PHASE_D_STUDENT_BASELINE_USE_CLASS_WEIGHTS_IN_LOSS,\n        use_weighted_sampler=CFG.PHASE_D_STUDENT_BASELINE_USE_WEIGHTED_SAMPLER,\n        num_workers=CFG.NUM_WORKERS,\n        pretrained=bool(getattr(CFG, \"SQUEEZENET_PRETRAINED\", True)),\n        weak_aug=CFG.PHASE_D_STUDENT_BASELINE_WEAK_AUG,\n        early_stop_patience=int(getattr(CFG, \"SQUEEZENET_EARLY_STOP_PATIENCE\", CFG.PHASE_D_STUDENT_BASELINE_EARLY_STOP_PATIENCE)),\n        seed=CFG.SEED,\n        device=CFG.DEVICE,\n        use_amp=CFG.USE_AMP,\n        pin_memory=CFG.PIN_MEMORY,\n        persistent_workers=CFG.PERSISTENT_WORKERS,\n        prefetch_factor=CFG.PREFETCH_FACTOR,\n        init_checkpoint_path=None,\n        show_batch_progress=True,\n        show_epoch_summary=True,\n    )\n    squeezenet_baseline_ckpt_path = Path(squeezenet_baseline_metrics[\"checkpoint_path\"])\n    extra_student_results.append({\n        \"stage\": \"SqueezeNet undistilled\",\n        \"distill_source\": \"baseline\",\n        \"model_name\": squeezenet_model_name,\n        \"macro_f1\": squeezenet_baseline_metrics[\"test_macro_f1\"],\n        \"balanced_acc\": squeezenet_baseline_metrics[\"test_balanced_acc\"],\n        \"accuracy\": squeezenet_baseline_metrics.get(\"test_accuracy\"),\n        \"weighted_f1\": squeezenet_baseline_metrics.get(\"test_weighted_f1\"),\n        \"mcc\": squeezenet_baseline_metrics.get(\"test_mcc\"),\n        \"num_params\": squeezenet_baseline_metrics[\"num_params\"],\n        \"model_footprint_mb\": squeezenet_baseline_metrics.get(\"model_footprint_mb\"),\n        \"checkpoint_size_mb\": squeezenet_baseline_metrics.get(\"checkpoint_size_mb\"),\n        \"inference_throughput_img_s\": squeezenet_baseline_metrics.get(\"inference_throughput_img_s\"),\n        \"inference_latency_ms_per_img\": squeezenet_baseline_metrics.get(\"inference_latency_ms_per_img\"),\n        \"total_train_time_sec\": squeezenet_baseline_metrics.get(\"total_train_time_sec\"),\n        **_copy_efficiency_keys(squeezenet_baseline_metrics),\n        \"checkpoint_path\": squeezenet_baseline_metrics[\"checkpoint_path\"],\n    })\nelse:\n    print(\"CFG.RUN_SQUEEZENET_UNDISTILLED = False -> bỏ qua SqueezeNet baseline.\")\n\nif getattr(CFG, \"RUN_SQUEEZENET_KD_FEATURE\", True):\n    if not getattr(CFG, \"RUN_PHASE_D_KD\", True):\n        print(\"CFG.RUN_PHASE_D_KD = False -> bỏ qua SqueezeNet KD-feature.\")\n    elif \"feature\" not in CFG.KD_SOURCES:\n        print(\"'feature' không nằm trong CFG.KD_SOURCES -> bỏ qua SqueezeNet KD-feature.\")\n    else:\n        teacher_kd_model_name = globals().get(\"teacher_kd_model_name\", teacher_phase_d_model_name)\n        teacher_kd_ckpt_path = globals().get(\"teacher_kd_ckpt_path\", None)\n        if teacher_kd_ckpt_path is None:\n            raise RuntimeError(\"Không có teacher_kd_ckpt_path. Hãy chạy cell D2.3 train/retrain teacher trước khi chạy SqueezeNet KD-feature.\")\n\n        print(\"\\n\" + \"=\" * 110)\n        print(f\"Running SqueezeNet + KD-feature: teacher={teacher_kd_model_name} -> student={squeezenet_model_name}\")\n        print(\"=\" * 110)\n\n        squeezenet_feature_kd_result = run_phase_d_kd_source(\n            teacher_model_name=teacher_kd_model_name,\n            student_model_name=squeezenet_model_name,\n            teacher_checkpoint_path=teacher_kd_ckpt_path,\n            student_checkpoint_path=None,\n            distill_source=\"feature\",\n            output_dir=CFG.KD_DIR / f\"{teacher_kd_model_name}_to_{squeezenet_model_name}\" / \"feature\",\n            student_pretrained=bool(getattr(CFG, \"SQUEEZENET_PRETRAINED\", True)),\n            baseline_metrics_override=squeezenet_baseline_metrics,\n            show_batch_progress=True,\n        )\n        phase_d_results.append(squeezenet_feature_kd_result)\n        extra_student_results.append({\n            \"stage\": \"SqueezeNet KD-feature\",\n            \"distill_source\": \"feature\",\n            \"model_name\": squeezenet_model_name,\n            \"macro_f1\": squeezenet_feature_kd_result[\"test_macro_f1\"],\n            \"balanced_acc\": squeezenet_feature_kd_result[\"test_balanced_acc\"],\n            \"accuracy\": squeezenet_feature_kd_result.get(\"test_accuracy\"),\n            \"weighted_f1\": squeezenet_feature_kd_result.get(\"test_weighted_f1\"),\n            \"mcc\": squeezenet_feature_kd_result.get(\"test_mcc\"),\n            \"num_params\": squeezenet_feature_kd_result[\"student_num_params\"],\n            \"model_footprint_mb\": squeezenet_feature_kd_result.get(\"model_footprint_mb\"),\n            \"checkpoint_size_mb\": squeezenet_feature_kd_result.get(\"checkpoint_size_mb\"),\n            \"inference_throughput_img_s\": squeezenet_feature_kd_result.get(\"inference_throughput_img_s\"),\n            \"inference_latency_ms_per_img\": squeezenet_feature_kd_result.get(\"inference_latency_ms_per_img\"),\n            \"total_train_time_sec\": squeezenet_feature_kd_result.get(\"total_train_time_sec\"),\n            **_copy_efficiency_keys(squeezenet_feature_kd_result),\n            \"delta_macro_f1_vs_baseline\": squeezenet_feature_kd_result.get(\"delta_macro_f1_vs_student_baseline\"),\n            \"delta_bal_acc_vs_baseline\": squeezenet_feature_kd_result.get(\"delta_bal_acc_vs_student_baseline\"),\n            \"checkpoint_path\": squeezenet_feature_kd_result[\"checkpoint_path\"],\n        })\nelse:\n    print(\"CFG.RUN_SQUEEZENET_KD_FEATURE = False -> bỏ qua SqueezeNet KD-feature.\")\n\nextra_student_results_df = pd.DataFrame(extra_student_results)\nif len(extra_student_results_df) > 0:\n    display(extra_student_results_df)\n\n","metadata":{"id":"wPsDhv6dtuGL","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T11:45:57.900921Z","iopub.execute_input":"2026-05-09T11:45:57.901374Z","iopub.status.idle":"2026-05-09T11:59:00.018149Z","shell.execute_reply.started":"2026-05-09T11:45:57.901330Z","shell.execute_reply":"2026-05-09T11:59:00.017330Z"}},"outputs":[],"execution_count":null},{"id":"9ac5b0c1","cell_type":"code","source":"\n# ============================================================\n# D4.4. Tổng hợp nhanh sau khi chạy từng cell benchmark\n# ============================================================\nif \"phase_d_results\" not in globals() or len(phase_d_results) == 0:\n    phase_d_results_df = pd.DataFrame()\n    print(\"Chưa có kết quả nào trong phase_d_results.\")\nelse:\n    phase_d_results_df = pd.DataFrame(phase_d_results).sort_values(\n        by=[\"test_macro_f1\", \"test_balanced_acc\"],\n        ascending=False,\n    ).reset_index(drop=True)\n    display(phase_d_results_df)\n","metadata":{"id":"9ac5b0c1","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T11:59:00.020092Z","iopub.execute_input":"2026-05-09T11:59:00.020397Z","iopub.status.idle":"2026-05-09T11:59:00.057305Z","shell.execute_reply.started":"2026-05-09T11:59:00.020365Z","shell.execute_reply":"2026-05-09T11:59:00.056755Z"}},"outputs":[],"execution_count":null},{"id":"43490376","cell_type":"markdown","source":"## D5. So sánh teacher refresh, student undistilled (supervised baseline) và 3 sources KD\n\n### Mục tiêu\nTạo bảng so sánh cuối cùng cho Phase D mới, gồm:\n- **teacher sau refresh**\n- **student supervised baseline**\n- **student sau KD logit**\n- **student sau KD feature**\n- **student sau KD similarity**\n\n### Lưu ý đọc kết quả\n- Teacher chỉ là **mốc trần tham chiếu**, không phải baseline trực tiếp của student\n- Baseline trực tiếp của KD là **student supervised baseline**\n- Khi kết luận hiệu quả KD, hãy ưu tiên nhìn:\n  - `macro_f1`\n  - `balanced_acc`\n  - `delta` so với student undistilled\n","metadata":{"id":"43490376"}},{"id":"e4c41db2","cell_type":"code","source":"# ============================================================\n# Cell vai trò:\n# - Tổng hợp teacher refresh, student undistilled và các nhánh KD.\n# - Hỗ trợ thêm SqueezeNet baseline + SqueezeNet KD-feature.\n# - Bảng so sánh có cả quality metrics và efficiency metrics.\n# ============================================================\n\n# ==========================\n# D5. Quick comparison\n# ==========================\ndef _pick(row, *keys, default=None):\n    for k in keys:\n        if isinstance(row, dict) and k in row:\n            return row[k]\n    return default\n\n\ndef _row_from_metrics(stage: str, distill_source: str, metrics: Dict) -> Dict:\n    return {\n        \"stage\": stage,\n        \"distill_source\": distill_source,\n        \"model_name\": metrics.get(\"model_name\", metrics.get(\"student_model_name\")),\n        \"accuracy\": _pick(metrics, \"test_accuracy\", \"accuracy\"),\n        \"macro_precision\": _pick(metrics, \"test_macro_precision\", \"macro_precision\"),\n        \"macro_recall\": _pick(metrics, \"test_macro_recall\", \"macro_recall\"),\n        \"macro_f1\": _pick(metrics, \"test_macro_f1\", \"macro_f1\"),\n        \"weighted_f1\": _pick(metrics, \"test_weighted_f1\", \"weighted_f1\"),\n        \"balanced_acc\": _pick(metrics, \"test_balanced_acc\", \"balanced_acc\"),\n        \"mcc\": _pick(metrics, \"test_mcc\", \"mcc\"),\n        \"cohen_kappa\": _pick(metrics, \"test_cohen_kappa\", \"cohen_kappa\"),\n        \"num_params\": _pick(metrics, \"num_params\", \"student_num_params\"),\n        \"model_footprint_mb\": metrics.get(\"model_footprint_mb\"),\n        \"static_model_footprint_mb\": metrics.get(\"static_model_footprint_mb\", metrics.get(\"model_footprint_mb\")),\n        \"checkpoint_size_mb\": metrics.get(\"checkpoint_size_mb\"),\n        \"inference_peak_memory_footprint_mb\": metrics.get(\"inference_peak_memory_footprint_mb\"),\n        \"inference_peak_memory_footprint_source\": metrics.get(\"inference_peak_memory_footprint_source\"),\n        \"inference_peak_cuda_allocated_mb\": metrics.get(\"inference_peak_cuda_allocated_mb\"),\n        \"inference_peak_cuda_reserved_mb\": metrics.get(\"inference_peak_cuda_reserved_mb\"),\n        \"training_peak_memory_footprint_mb\": metrics.get(\"training_peak_memory_footprint_mb\"),\n        \"inference_throughput_img_s\": metrics.get(\"inference_throughput_img_s\"),\n        \"inference_latency_ms_per_img\": metrics.get(\"inference_latency_ms_per_img\"),\n        \"inference_latency_ms_per_batch\": metrics.get(\"inference_latency_ms_per_batch\"),\n        \"forward_gflops_per_img\": metrics.get(\"forward_gflops_per_img\"),\n        \"forward_mflops_per_img\": metrics.get(\"forward_mflops_per_img\"),\n        \"forward_flops_per_img\": metrics.get(\"forward_flops_per_img\"),\n        \"flops_method\": metrics.get(\"flops_method\"),\n        \"flops_eval_batch_size\": metrics.get(\"flops_eval_batch_size\"),\n        \"total_train_time_sec\": metrics.get(\"total_train_time_sec\"),\n        \"cuda_peak_allocated_gb\": metrics.get(\"cuda_peak_allocated_gb\"),\n        \"checkpoint_path\": metrics.get(\"checkpoint_path\"),\n    }\n\n\nteacher_compare_rows = []\nstudent_compare_rows = []\nseen = set()\n\nif \"teacher_refresh_metrics\" in globals() and teacher_refresh_metrics is not None:\n    teacher_compare_rows.append(_row_from_metrics(\"Teacher refresh\", \"teacher\", teacher_refresh_metrics))\nelse:\n    teacher_compare_rows.append({\n        \"stage\": \"Teacher manual / not retrained\",\n        \"distill_source\": \"teacher\",\n        \"model_name\": globals().get(\"teacher_kd_model_name\", teacher_phase_d_model_name),\n        \"macro_f1\": None,\n        \"balanced_acc\": None,\n        \"num_params\": None,\n        \"checkpoint_path\": str(globals().get(\"teacher_kd_ckpt_path\", None)) if globals().get(\"teacher_kd_ckpt_path\", None) is not None else None,\n    })\n\n# Main student baseline.\nif \"student_baseline_metrics\" in globals() and student_baseline_metrics is not None:\n    row = _row_from_metrics(\"Student undistilled\", \"baseline\", student_baseline_metrics)\n    key = (row[\"model_name\"], row[\"distill_source\"], row[\"stage\"])\n    seen.add(key)\n    student_compare_rows.append(row)\nelse:\n    # Phase D-only: không lấy baseline từ screening/Phase C.\n    pass\n\n# Extra baselines, e.g. SqueezeNet undistilled. KD rows từ phase_d_results sẽ được thêm bên dưới.\nif \"extra_student_results\" in globals() and len(extra_student_results) > 0:\n    for r in extra_student_results:\n        if r.get(\"distill_source\") != \"baseline\":\n            continue\n        row = dict(r)\n        key = (row.get(\"model_name\"), row.get(\"distill_source\"), row.get(\"stage\"))\n        if key not in seen:\n            seen.add(key)\n            student_compare_rows.append(row)\n\n# All KD results, including custom_cnn_small and SqueezeNet KD-feature.\nif \"phase_d_results\" in globals() and len(phase_d_results) > 0:\n    for r in phase_d_results:\n        stage = f\"KD - {r['distill_source']} ({r['student_model_name']})\"\n        row = _row_from_metrics(stage, r[\"distill_source\"], r)\n        key = (row.get(\"model_name\"), row.get(\"distill_source\"), row.get(\"stage\"))\n        if key not in seen:\n            seen.add(key)\n            student_compare_rows.append(row)\n\nteacher_compare_df = pd.DataFrame(teacher_compare_rows)\nstudent_compare_df = pd.DataFrame(student_compare_rows)\n\nif not teacher_compare_df.empty:\n    print(\"Teacher reference:\")\n    display(teacher_compare_df)\n\nif student_compare_df.empty:\n    print(\"Chưa có đủ dữ liệu để so sánh student undistilled và KD.\")\nelse:\n    # Delta được tính theo baseline của cùng model_name để tránh so nhầm SqueezeNet với custom_cnn_small.\n    student_compare_df[\"delta_macro_f1_vs_same_model_baseline\"] = np.nan\n    student_compare_df[\"delta_bal_acc_vs_same_model_baseline\"] = np.nan\n\n    for model_name_i, group in student_compare_df.groupby(\"model_name\", dropna=False):\n        baseline_sub = group[group[\"distill_source\"] == \"baseline\"]\n        if len(baseline_sub) == 0:\n            continue\n        base_macro = baseline_sub.iloc[0].get(\"macro_f1\")\n        base_bal = baseline_sub.iloc[0].get(\"balanced_acc\")\n        idxs = student_compare_df[\"model_name\"].astype(str) == str(model_name_i)\n        if pd.notna(base_macro):\n            student_compare_df.loc[idxs, \"delta_macro_f1_vs_same_model_baseline\"] = student_compare_df.loc[idxs, \"macro_f1\"] - float(base_macro)\n        if pd.notna(base_bal):\n            student_compare_df.loc[idxs, \"delta_bal_acc_vs_same_model_baseline\"] = student_compare_df.loc[idxs, \"balanced_acc\"] - float(base_bal)\n\n    preferred_cols = [\n        \"stage\", \"model_name\", \"distill_source\",\n        \"accuracy\", \"macro_precision\", \"macro_recall\", \"macro_f1\", \"weighted_f1\", \"balanced_acc\", \"mcc\", \"cohen_kappa\",\n        \"delta_macro_f1_vs_same_model_baseline\", \"delta_bal_acc_vs_same_model_baseline\",\n        \"num_params\", \"model_footprint_mb\", \"static_model_footprint_mb\", \"checkpoint_size_mb\",\n        \"inference_peak_memory_footprint_mb\", \"inference_peak_memory_footprint_source\",\n        \"inference_peak_cuda_allocated_mb\", \"inference_peak_cuda_reserved_mb\", \"training_peak_memory_footprint_mb\",\n        \"inference_throughput_img_s\", \"inference_latency_ms_per_img\", \"inference_latency_ms_per_batch\",\n        \"forward_gflops_per_img\", \"forward_mflops_per_img\", \"forward_flops_per_img\", \"flops_method\", \"flops_eval_batch_size\",\n        \"total_train_time_sec\", \"cuda_peak_allocated_gb\",\n        \"checkpoint_path\",\n    ]\n    existing_cols = [c for c in preferred_cols if c in student_compare_df.columns]\n    remaining_cols = [c for c in student_compare_df.columns if c not in existing_cols]\n    student_compare_df = student_compare_df[existing_cols + remaining_cols]\n\n    display(student_compare_df)\n\n    compare_csv_path = CFG.PHASE_D_DIR / \"phase_d_student_compare.csv\"\n    student_compare_df.to_csv(compare_csv_path, index=False)\n    print(\"Saved student comparison to:\", compare_csv_path)\n\n    plot_df = student_compare_df.dropna(subset=[\"macro_f1\"]).copy()\n    if len(plot_df) > 0:\n        plt.figure(figsize=(11, 4))\n        sns.barplot(data=plot_df, x=\"stage\", y=\"macro_f1\")\n        plt.title(\"Student comparison: undistilled vs KD\")\n        plt.xticks(rotation=25, ha=\"right\")\n        plt.tight_layout()\n        plt.show()\n\n        plt.figure(figsize=(11, 4))\n        sns.barplot(data=plot_df, x=\"stage\", y=\"balanced_acc\")\n        plt.title(\"Balanced accuracy: undistilled vs KD\")\n        plt.xticks(rotation=25, ha=\"right\")\n        plt.tight_layout()\n        plt.show()\n\n        eff_df = plot_df.dropna(subset=[\"inference_latency_ms_per_img\"])\n        if len(eff_df) > 0:\n            plt.figure(figsize=(11, 4))\n            sns.barplot(data=eff_df, x=\"stage\", y=\"inference_latency_ms_per_img\")\n            plt.title(\"Inference latency per image (lower is better)\")\n            plt.xticks(rotation=25, ha=\"right\")\n            plt.tight_layout()\n            plt.show()\n\n        throughput_df = plot_df.dropna(subset=[\"inference_throughput_img_s\"])\n        if len(throughput_df) > 0:\n            plt.figure(figsize=(11, 4))\n            sns.barplot(data=throughput_df, x=\"stage\", y=\"inference_throughput_img_s\")\n            plt.title(\"Inference throughput, images/sec (higher is better)\")\n            plt.xticks(rotation=25, ha=\"right\")\n            plt.tight_layout()\n            plt.show()\n\n        memory_df = plot_df.dropna(subset=[\"inference_peak_memory_footprint_mb\"])\n        if len(memory_df) > 0:\n            plt.figure(figsize=(11, 4))\n            sns.barplot(data=memory_df, x=\"stage\", y=\"inference_peak_memory_footprint_mb\")\n            plt.title(\"Inference peak memory footprint, MB (lower is better)\")\n            plt.xticks(rotation=25, ha=\"right\")\n            plt.tight_layout()\n            plt.show()\n\n        flops_df = plot_df.dropna(subset=[\"forward_gflops_per_img\"])\n        if len(flops_df) > 0:\n            plt.figure(figsize=(11, 4))\n            sns.barplot(data=flops_df, x=\"stage\", y=\"forward_gflops_per_img\")\n            plt.title(\"Forward GFLOPs per image (lower is better)\")\n            plt.xticks(rotation=25, ha=\"right\")\n            plt.tight_layout()\n            plt.show()\n\n    kd_only_df = student_compare_df[student_compare_df[\"distill_source\"] != \"baseline\"].copy() if \"distill_source\" in student_compare_df.columns else pd.DataFrame()\n    if len(kd_only_df) > 0:\n        kd_only_df = kd_only_df.sort_values([\"macro_f1\", \"balanced_acc\"], ascending=False).reset_index(drop=True)\n        best_row = kd_only_df.iloc[0]\n        print(\"Best KD stage              :\", best_row[\"stage\"])\n        print(\"Best KD macro-F1           :\", best_row[\"macro_f1\"])\n        print(\"Best KD balanced accuracy  :\", best_row[\"balanced_acc\"])\n        print(\"Best KD checkpoint         :\", best_row[\"checkpoint_path\"])\n\n","metadata":{"id":"e4c41db2","trusted":true,"execution":{"iopub.status.busy":"2026-05-09T11:59:00.058459Z","iopub.execute_input":"2026-05-09T11:59:00.058792Z","iopub.status.idle":"2026-05-09T11:59:01.186765Z","shell.execute_reply.started":"2026-05-09T11:59:00.058765Z","shell.execute_reply":"2026-05-09T11:59:01.186100Z"}},"outputs":[],"execution_count":null},{"id":"f2a8b33e","cell_type":"markdown","source":"## D6. Kết luận Phase D\n\n### Sau khi chạy xong, bạn nên chốt 3 ý\n1. **Teacher refresh có cải thiện chất lượng teacher hay không?**\n2. **KD có vượt student baseline hay không?**\n3. **Source nào tốt nhất giữa `logit`, `feature`, `similarity`?**\n\n### Diễn giải khoa học nên dùng\n- Nếu KD vượt baseline rõ rệt: teacher đã truyền được tri thức hữu ích\n- Nếu chỉ một source vượt baseline: source đó phù hợp nhất với student nhỏ hiện tại\n- Nếu teacher mạnh nhưng KD không vượt baseline: vấn đề có thể nằm ở:\n  - design loss\n  - capacity của student\n  - mismatch giữa source tri thức và kiến trúc student\n","metadata":{"id":"f2a8b33e"}}]}