{"metadata":{"kernelspec":{"display_name":"Python3 (ipykernel)","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.13"}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"9ff679ac","cell_type":"markdown","source":"## Cell 0 — Install packages","metadata":{}},{"id":"a8093fb2","cell_type":"code","source":"%pip install -q opencv-python-headless scikit-learn pandas matplotlib pillow tqdm openpyxl","metadata":{},"outputs":[],"execution_count":null},{"id":"d1b7eb54","cell_type":"markdown","source":"## Cell 1 — Config","metadata":{}},{"id":"ed0a33d6","cell_type":"code","source":"import os, gc, random, subprocess, shutil\nfrom pathlib import Path\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom tqdm.auto import tqdm\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    confusion_matrix,\n    cohen_kappa_score,\n    accuracy_score,\n    precision_recall_fscore_support,\n    balanced_accuracy_score,\n    roc_auc_score,\n    matthews_corrcoef,\n)\n\nSEED = 42\nWORK_DIR = Path(\"/workspace\")\nSAVE_DIR = WORK_DIR / \"jfsp_vit_no_scheduler_first_outputs\"\nSAVE_DIR.mkdir(parents=True, exist_ok=True)\n\nDATASET_NAME = \"APTOS2019\"\n\n# Thư mục đích chỉ được dùng khi thật sự cần tải dữ liệu.\n# Cell kế tiếp sẽ ưu tiên tìm dataset đã được gắn sẵn trước.\nAPTOS_DOWNLOAD_ROOT = (\n    WORK_DIR\n    / \"data\"\n    / \"aptos2019-blindness-detection\"\n)\n\nAPTOS_ROOT = None\nAPTOS_CSV = None\nAPTOS_IMG_DIR = None\n\nIMG_SIZE = 224\nNUM_CLASSES = 5\n\nPAPER_LR = 0.00019\nBATCH_SIZE = 16\n\nSWEEP_EPOCHS = 18\nSWEEP_PATIENCE = 5\nFULL_EPOCHS = 60\nFULL_PATIENCE = 10\n\nUSE_AMP = True\nNUM_WORKERS = 0\nPIN_MEMORY = False\nEFFECTIVE_BATCH = 16\nACCUM_STEPS = max(1, EFFECTIVE_BATCH // BATCH_SIZE)\n\n# FAST: chỉ quét nhóm gần paper nhất, không scheduler.\n# FULL: thêm vài optimizer/loss và thêm scheduler ở nhóm phụ để chẩn đoán.\nSWEEP_MODE = \"FAST\"\n\nCLASS_NAMES = {\n    0: \"0_No_DR\",\n    1: \"1_Mild\",\n    2: \"2_Moderate\",\n    3: \"3_Severe\",\n    4: \"4_Proliferative_DR\",\n}\n\ndef seed_everything(seed=42):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(SEED)\n\nDEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\nprint(\"DEVICE:\", DEVICE)\nif DEVICE == \"cuda\":\n    print(\"GPU:\", torch.cuda.get_device_name(0))\nprint(\"SAVE_DIR:\", SAVE_DIR)\nprint(\n    \"APTOS dataset sẽ được tự nhận diện \"\n    \"ở Cell 2.\"\n)","metadata":{},"outputs":[],"execution_count":null},{"id":"23b08f63","cell_type":"markdown","source":"## Cell 2 — Download APTOS nếu cần","metadata":{}},{"id":"ec874de1","cell_type":"code","source":"# =========================================================\n# APTOS 2019 DATASET SETUP\n#\n# Chế độ \"auto\":\n# 1. Ưu tiên dùng dataset đã được gắn sẵn.\n# 2. Chỉ tải từ Kaggle nếu chưa tìm thấy dataset.\n# 3. Không bắt buộc kaggle.json khi dataset đã có sẵn.\n# =========================================================\n\nDOWNLOAD_IF_MISSING = True\n\n\ndef is_valid_aptos_root(root):\n    \"\"\"\n    Một thư mục APTOS hợp lệ cần có:\n    - train.csv\n    - thư mục train_images\n    \"\"\"\n    root = Path(root)\n\n    return (\n        (root / \"train.csv\").is_file()\n        and (root / \"train_images\").is_dir()\n    )\n\n\ndef find_existing_aptos_dataset():\n    \"\"\"\n    Tìm APTOS trong các vị trí thường gặp trên:\n    - Kaggle Notebook\n    - /workspace\n    - Google Colab\n    - thư mục hiện hành\n    \"\"\"\n    direct_candidates = [\n        Path(\n            \"/kaggle/input/\"\n            \"aptos2019-blindness-detection\"\n        ),\n        Path(\n            \"/workspace/data/\"\n            \"aptos2019-blindness-detection\"\n        ),\n        Path(\n            \"/workspace/\"\n            \"aptos2019-blindness-detection\"\n        ),\n        Path(\n            \"/content/\"\n            \"aptos2019-blindness-detection\"\n        ),\n        Path.cwd()\n        / \"aptos2019-blindness-detection\",\n    ]\n\n    for candidate in direct_candidates:\n        if is_valid_aptos_root(candidate):\n            return candidate\n\n    # Tìm theo train.csv nhưng giới hạn ở các thư mục dữ liệu phổ biến.\n    search_roots = [\n        Path(\"/kaggle/input\"),\n        Path(\"/workspace/data\"),\n        Path(\"/content\"),\n        Path.cwd(),\n    ]\n\n    for search_root in search_roots:\n        if not search_root.exists():\n            continue\n\n        try:\n            for csv_path in search_root.rglob(\n                \"train.csv\"\n            ):\n                candidate = csv_path.parent\n\n                if is_valid_aptos_root(\n                    candidate\n                ):\n                    return candidate\n\n        except PermissionError:\n            continue\n\n    return None\n\n\ndef prepare_kaggle_credentials():\n    \"\"\"\n    Trả về True khi có thể dùng Kaggle API.\n\n    Hỗ trợ:\n    - /root/.kaggle/kaggle.json\n    - /workspace/kaggle.json\n    - /content/kaggle.json\n    - biến môi trường KAGGLE_USERNAME và KAGGLE_KEY\n    \"\"\"\n    kaggle_dir = Path(\n        \"/root/.kaggle\"\n    )\n\n    kaggle_dir.mkdir(\n        parents=True,\n        exist_ok=True\n    )\n\n    default_token = (\n        kaggle_dir\n        / \"kaggle.json\"\n    )\n\n    token_candidates = [\n        default_token,\n        WORK_DIR / \"kaggle.json\",\n        Path(\"/content/kaggle.json\"),\n        Path.cwd() / \"kaggle.json\",\n    ]\n\n    if not default_token.is_file():\n        for token_path in token_candidates:\n            if not token_path.is_file():\n                continue\n\n            shutil.copy(\n                token_path,\n                default_token\n            )\n\n            break\n\n    if default_token.is_file():\n        os.chmod(\n            default_token,\n            0o600\n        )\n\n        return True\n\n    has_environment_key = bool(\n        os.environ.get(\n            \"KAGGLE_USERNAME\"\n        )\n        and os.environ.get(\n            \"KAGGLE_KEY\"\n        )\n    )\n\n    return has_environment_key\n\n\ndef download_aptos_dataset(\n    destination\n):\n    \"\"\"\n    Tải và giải nén APTOS 2019 bằng Kaggle API.\n    \"\"\"\n    destination = Path(\n        destination\n    )\n\n    destination.mkdir(\n        parents=True,\n        exist_ok=True\n    )\n\n    subprocess.run(\n        [\n            \"pip\",\n            \"install\",\n            \"-q\",\n            \"kaggle\",\n        ],\n        check=True\n    )\n\n    subprocess.run(\n        [\n            \"kaggle\",\n            \"competitions\",\n            \"download\",\n            \"-c\",\n            \"aptos2019-blindness-detection\",\n            \"-p\",\n            str(destination),\n        ],\n        check=True\n    )\n\n    zip_candidates = list(\n        destination.glob(\n            \"*.zip\"\n        )\n    )\n\n    if not zip_candidates:\n        raise FileNotFoundError(\n            \"Kaggle API đã chạy nhưng không tìm thấy \"\n            f\"file ZIP trong: {destination}\"\n        )\n\n    zip_path = zip_candidates[0]\n\n    subprocess.run(\n        [\n            \"unzip\",\n            \"-q\",\n            \"-o\",\n            str(zip_path),\n            \"-d\",\n            str(destination),\n        ],\n        check=True\n    )\n\n    if not is_valid_aptos_root(\n        destination\n    ):\n        raise RuntimeError(\n            \"Đã tải và giải nén nhưng cấu trúc dataset \"\n            \"không có train.csv và train_images.\"\n        )\n\n    return destination\n\n\n# =========================================================\n# 1. TÌM DATASET ĐÃ CÓ SẴN\n# =========================================================\n\ndetected_root = (\n    find_existing_aptos_dataset()\n)\n\n\n# =========================================================\n# 2. NẾU CHƯA CÓ, THỬ TẢI BẰNG KAGGLE API\n# =========================================================\n\nif detected_root is None:\n    if not DOWNLOAD_IF_MISSING:\n        raise FileNotFoundError(\n            \"Không tìm thấy APTOS 2019 trong hệ thống. \"\n            \"Hãy gắn dataset vào notebook hoặc đặt \"\n            \"DOWNLOAD_IF_MISSING=True.\"\n        )\n\n    if not prepare_kaggle_credentials():\n        raise FileNotFoundError(\n            \"\\nKhông tìm thấy APTOS 2019 và cũng chưa có \"\n            \"thông tin đăng nhập Kaggle.\\n\\n\"\n            \"Cách 1 — khi chạy trên Kaggle Notebook:\\n\"\n            \"  Add Input → tìm \"\n            \"'APTOS 2019 Blindness Detection'.\\n\"\n            \"  Sau đó chạy lại cell này; không cần kaggle.json.\\n\\n\"\n            \"Cách 2 — khi chạy trong /workspace hoặc Colab:\\n\"\n            \"  đặt kaggle.json tại một trong các vị trí:\\n\"\n            \"  /workspace/kaggle.json\\n\"\n            \"  /content/kaggle.json\\n\"\n            \"  /root/.kaggle/kaggle.json\\n\\n\"\n            \"Cách 3 — khai báo biến môi trường:\\n\"\n            \"  KAGGLE_USERNAME và KAGGLE_KEY.\"\n        )\n\n    detected_root = download_aptos_dataset(\n        APTOS_DOWNLOAD_ROOT\n    )\n\n\n# =========================================================\n# 3. CẬP NHẬT ĐƯỜNG DẪN CHUNG CHO CÁC CELL SAU\n# =========================================================\n\nAPTOS_ROOT = Path(\n    detected_root\n)\n\nAPTOS_CSV = str(\n    APTOS_ROOT\n    / \"train.csv\"\n)\n\nAPTOS_IMG_DIR = str(\n    APTOS_ROOT\n    / \"train_images\"\n)\n\nassert Path(APTOS_CSV).is_file(), (\n    f\"Không tìm thấy CSV: {APTOS_CSV}\"\n)\n\nassert Path(APTOS_IMG_DIR).is_dir(), (\n    f\"Không tìm thấy thư mục ảnh: \"\n    f\"{APTOS_IMG_DIR}\"\n)\n\nprint(\n    \"APTOS dataset đã sẵn sàng.\"\n)\n\nprint(\n    \"APTOS_ROOT    :\",\n    APTOS_ROOT\n)\n\nprint(\n    \"APTOS_CSV     :\",\n    APTOS_CSV\n)\n\nprint(\n    \"APTOS_IMG_DIR :\",\n    APTOS_IMG_DIR\n)","metadata":{},"outputs":[],"execution_count":null},{"id":"803e1e06","cell_type":"markdown","source":"## Cell 3 — Data split","metadata":{}},{"id":"74f0d82b","cell_type":"code","source":"def find_image_file(img_dir, image_id, suffix=\".png\"):\n    img_dir = Path(img_dir)\n    image_id = str(image_id)\n    stem = Path(image_id).stem\n\n    candidates = []\n    if suffix:\n        candidates.append(img_dir / f\"{image_id}{suffix}\")\n    candidates.append(img_dir / image_id)\n\n    for ext in [\".png\", \".jpg\", \".jpeg\", \".PNG\", \".JPG\", \".JPEG\"]:\n        candidates.append(img_dir / f\"{stem}{ext}\")\n\n    for p in candidates:\n        if p.exists():\n            return str(p)\n\n    matches = list(img_dir.rglob(f\"{stem}.*\"))\n    if matches:\n        return str(matches[0])\n\n    return str(candidates[0])\n\ndef load_aptos_dataframe():\n    df = pd.read_csv(APTOS_CSV)\n    df = df.rename(columns={\"id_code\": \"image_id\", \"diagnosis\": \"label\"})\n    df[\"path\"] = df[\"image_id\"].apply(lambda x: find_image_file(APTOS_IMG_DIR, x, suffix=\".png\"))\n    df = df[[\"image_id\", \"path\", \"label\"]].copy()\n    df[\"label\"] = df[\"label\"].astype(int)\n    return df\n\nfull_df = load_aptos_dataframe()\n\nprint(\"Total:\", len(full_df))\nprint(full_df[\"label\"].value_counts().sort_index())\ndisplay(full_df.head())\n\ntrain_df, temp_df = train_test_split(\n    full_df,\n    test_size=0.20,\n    stratify=full_df[\"label\"],\n    random_state=SEED,\n)\n\nval_df, test_df = train_test_split(\n    temp_df,\n    test_size=0.50,\n    stratify=temp_df[\"label\"],\n    random_state=SEED,\n)\n\ntrain_df = train_df.reset_index(drop=True)\nval_df = val_df.reset_index(drop=True)\ntest_df = test_df.reset_index(drop=True)\n\nprint(\"train:\", len(train_df), \"val:\", len(val_df), \"test:\", len(test_df))\nprint(\"train labels:\\n\", train_df[\"label\"].value_counts().sort_index())\nprint(\"val labels:\\n\", val_df[\"label\"].value_counts().sort_index())\nprint(\"test labels:\\n\", test_df[\"label\"].value_counts().sort_index())\n\ntrain_df.to_csv(SAVE_DIR / \"APTOS2019_train_split.csv\", index=False)\nval_df.to_csv(SAVE_DIR / \"APTOS2019_val_split.csv\", index=False)\ntest_df.to_csv(SAVE_DIR / \"APTOS2019_test_split.csv\", index=False)","metadata":{},"outputs":[],"execution_count":null},{"id":"427df5c5","cell_type":"markdown","source":"## Cell 4 — Preprocess + cache","metadata":{}},{"id":"46d92f00","cell_type":"code","source":"def crop_black_border_rgb(img_rgb, tol=7):\n    gray = cv2.cvtColor(img_rgb, cv2.COLOR_RGB2GRAY)\n    mask = gray > tol\n    if mask.sum() == 0:\n        return img_rgb\n    ys, xs = np.where(mask)\n    return img_rgb[ys.min():ys.max() + 1, xs.min():xs.max() + 1]\n\ndef jfsp_preprocess_from_path(path, img_size=224, clahe_clip=2.0, clahe_grid=(8, 8)):\n    img_bgr = cv2.imread(str(path))\n    if img_bgr is None:\n        raise FileNotFoundError(f\"Cannot read image: {path}\")\n\n    img_rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n    img_rgb = crop_black_border_rgb(img_rgb)\n    img_rgb = cv2.resize(img_rgb, (img_size, img_size), interpolation=cv2.INTER_AREA)\n\n    blur = cv2.GaussianBlur(img_rgb, (0, 0), sigmaX=10)\n    fused = cv2.addWeighted(img_rgb, 4.0, blur, -4.0, 128.0)\n\n    lab = cv2.cvtColor(fused, cv2.COLOR_RGB2LAB)\n    l, a, b = cv2.split(lab)\n    clahe = cv2.createCLAHE(clipLimit=clahe_clip, tileGridSize=clahe_grid)\n    l2 = clahe.apply(l)\n    enhanced = cv2.cvtColor(cv2.merge([l2, a, b]), cv2.COLOR_LAB2RGB)\n    return enhanced\n\nUSE_PREPROCESS_CACHE = True\n\ndef build_preprocess_cache(df, split_name, img_size=224):\n    if not USE_PREPROCESS_CACHE:\n        return df\n\n    cache_dir = WORK_DIR / f\"cache_{DATASET_NAME}_{img_size}\" / split_name\n    cache_dir.mkdir(parents=True, exist_ok=True)\n\n    cached_df = df.copy()\n    new_paths = []\n\n    for _, row in tqdm(cached_df.iterrows(), total=len(cached_df), desc=f\"Cache {split_name}\"):\n        label = int(row[\"label\"])\n        image_id = str(row[\"image_id\"])\n        out_dir = cache_dir / str(label)\n        out_dir.mkdir(parents=True, exist_ok=True)\n        out_path = out_dir / f\"{Path(image_id).stem}.png\"\n\n        if not out_path.exists():\n            img = jfsp_preprocess_from_path(row[\"path\"], img_size)\n            cv2.imwrite(str(out_path), cv2.cvtColor(img, cv2.COLOR_RGB2BGR))\n\n        new_paths.append(str(out_path))\n\n    cached_df[\"original_path\"] = cached_df[\"path\"]\n    cached_df[\"path\"] = new_paths\n    return cached_df\n\nif USE_PREPROCESS_CACHE:\n    train_df = build_preprocess_cache(train_df, \"train\", IMG_SIZE)\n    val_df = build_preprocess_cache(val_df, \"validation\", IMG_SIZE)\n    test_df = build_preprocess_cache(test_df, \"test\", IMG_SIZE)\n\n    train_df.to_csv(SAVE_DIR / \"APTOS2019_train_split_cached.csv\", index=False)\n    val_df.to_csv(SAVE_DIR / \"APTOS2019_val_split_cached.csv\", index=False)\n    test_df.to_csv(SAVE_DIR / \"APTOS2019_test_split_cached.csv\", index=False)\n\nprint(\"Cache ready\")","metadata":{},"outputs":[],"execution_count":null},{"id":"852d91d2","cell_type":"markdown","source":"## Cell 5 — DataLoader","metadata":{}},{"id":"e11c1689","cell_type":"code","source":"class DRDataset(Dataset):\n    def __init__(self, df, img_size=224):\n        self.df = df.reset_index(drop=True)\n        self.img_size = img_size\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        path = str(row[\"path\"])\n\n        if \"cache_\" in path:\n            img_bgr = cv2.imread(path)\n            if img_bgr is None:\n                raise FileNotFoundError(path)\n            img = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB)\n        else:\n            img = jfsp_preprocess_from_path(path, self.img_size)\n\n        x = torch.from_numpy(img).permute(2, 0, 1).float() / 255.0\n        y = torch.tensor(int(row[\"label\"]), dtype=torch.long)\n        return x, y\n\ndef make_loader(df, train=False):\n    return DataLoader(\n        DRDataset(df, IMG_SIZE),\n        batch_size=BATCH_SIZE,\n        shuffle=train,\n        num_workers=NUM_WORKERS,\n        pin_memory=PIN_MEMORY,\n        drop_last=False,\n    )\n\ntrain_loader = make_loader(train_df, train=True)\nval_loader = make_loader(val_df, train=False)\ntest_loader = make_loader(test_df, train=False)\n\nx, y = next(iter(train_loader))\nprint(x.shape, y.shape, x.min().item(), x.max().item())","metadata":{},"outputs":[],"execution_count":null},{"id":"430df548","cell_type":"markdown","source":"## Cell 6 — Model modules","metadata":{}},{"id":"6c9841a9","cell_type":"code","source":"class LayerNorm2d(nn.Module):\n    def __init__(self, channels, eps=1e-6):\n        super().__init__()\n        self.norm = nn.GroupNorm(1, channels, eps=eps)\n\n    def forward(self, x):\n        return self.norm(x)\n\nclass ConvBNAct(nn.Module):\n    def __init__(self, in_ch, out_ch, kernel=3, stride=1, padding=None, groups=1, act=True):\n        super().__init__()\n        if padding is None:\n            padding = kernel // 2\n        self.conv = nn.Conv2d(in_ch, out_ch, kernel, stride, padding, groups=groups, bias=False)\n        self.bn = nn.BatchNorm2d(out_ch)\n        self.act = nn.SiLU(inplace=True) if act else nn.Identity()\n\n    def forward(self, x):\n        return self.act(self.bn(self.conv(x)))\n\nclass InvertedResidualBlock(nn.Module):\n    def __init__(self, in_ch, out_ch, stride=1, expand_ratio=4):\n        super().__init__()\n        hidden = int(round(in_ch * expand_ratio))\n        self.use_residual = stride == 1 and in_ch == out_ch\n        self.expand = ConvBNAct(in_ch, hidden, kernel=1, stride=1, padding=0)\n        self.depthwise = ConvBNAct(hidden, hidden, kernel=3, stride=stride, padding=1, groups=hidden)\n        self.project = ConvBNAct(hidden, out_ch, kernel=1, stride=1, padding=0, act=False)\n\n    def forward(self, x):\n        y = self.project(self.depthwise(self.expand(x)))\n        return x + y if self.use_residual else y\n\nclass EMA(nn.Module):\n    def __init__(self, channels, groups=8):\n        super().__init__()\n        assert channels % groups == 0\n        self.groups = groups\n        group_ch = channels // groups\n        self.softmax = nn.Softmax(dim=-1)\n        self.agp = nn.AdaptiveAvgPool2d((1, 1))\n        self.gn = nn.GroupNorm(group_ch, group_ch)\n        self.conv1x1 = nn.Conv2d(group_ch, group_ch, 1, bias=False)\n        self.conv3x3 = nn.Conv2d(group_ch, group_ch, 3, padding=1, bias=False)\n\n    def forward(self, x):\n        b, c, h, w = x.shape\n        g = self.groups\n        gx = x.reshape(b * g, c // g, h, w)\n\n        x_h = gx.mean(dim=3, keepdim=True)\n        x_w = gx.mean(dim=2, keepdim=True).permute(0, 1, 3, 2)\n\n        hw = self.conv1x1(torch.cat([x_h, x_w], dim=2))\n        x_h, x_w = torch.split(hw, [h, w], dim=2)\n\n        x1 = self.gn(gx * x_h.sigmoid() * x_w.permute(0, 1, 3, 2).sigmoid())\n        x2 = self.conv3x3(gx)\n\n        x11 = self.softmax(self.agp(x1).reshape(b * g, 1, c // g))\n        x12 = x2.reshape(b * g, c // g, -1)\n        x21 = self.softmax(self.agp(x2).reshape(b * g, 1, c // g))\n        x22 = x1.reshape(b * g, c // g, -1)\n\n        weights = (torch.matmul(x11, x12) + torch.matmul(x21, x22)).reshape(b * g, 1, h, w)\n        return (gx * weights.sigmoid()).reshape(b, c, h, w)\n\nclass InvertedResidualAttentionBlock(nn.Module):\n    def __init__(self, in_ch, out_ch, stride=1, expand_ratio=4, ema_groups=8):\n        super().__init__()\n        self.irb = InvertedResidualBlock(in_ch, out_ch, stride=stride, expand_ratio=expand_ratio)\n        self.ema = EMA(out_ch, groups=ema_groups)\n\n    def forward(self, x):\n        return self.ema(self.irb(x))\n\nclass DepthwiseSeparableConv(nn.Module):\n    def __init__(self, channels, kernel_size):\n        super().__init__()\n        padding = kernel_size // 2\n        self.dw = nn.Conv2d(channels, channels, kernel_size, padding=padding, groups=channels, bias=False)\n        self.bn1 = nn.BatchNorm2d(channels)\n        self.pw = nn.Conv2d(channels, channels, 1, bias=False)\n        self.bn2 = nn.BatchNorm2d(channels)\n        self.act = nn.SiLU(inplace=True)\n\n    def forward(self, x):\n        x = self.act(self.bn1(self.dw(x)))\n        return self.act(self.bn2(self.pw(x)))\n\nclass DepthwiseOnlyConv(nn.Module):\n    def __init__(self, channels, kernel_size):\n        super().__init__()\n        padding = kernel_size // 2\n        self.dw = nn.Conv2d(channels, channels, kernel_size, padding=padding, groups=channels, bias=False)\n        self.bn = nn.BatchNorm2d(channels)\n        self.act = nn.SiLU(inplace=True)\n\n    def forward(self, x):\n        return self.act(self.bn(self.dw(x)))\n\ndef amp_disabled_context():\n    if DEVICE == \"cuda\":\n        return torch.amp.autocast(device_type=\"cuda\", enabled=False)\n    return torch.amp.autocast(device_type=\"cpu\", enabled=False)\n\nclass FrequencySelfAttention(nn.Module):\n    def __init__(self, channels, heads=8):\n        super().__init__()\n        assert channels % heads == 0\n        self.channels = channels\n        self.heads = heads\n        self.dim_head = channels // heads\n\n        self.norm = LayerNorm2d(channels)\n        self.q = nn.Conv2d(channels, channels, 1, bias=False)\n        self.k = nn.Conv2d(channels, channels, 1, bias=False)\n        self.v = nn.Conv2d(channels, channels, 1, bias=False)\n\n        self.temperature = nn.Parameter(torch.ones(heads, 1, 1))\n        self.freq_weight = nn.Parameter(torch.zeros(1, channels, 1, 1))\n\n        self.proj = nn.Sequential(\n            nn.Conv2d(channels * 2, channels, 1, bias=False),\n            nn.BatchNorm2d(channels),\n        )\n\n    def _reshape_heads(self, z):\n        b, c, h, w = z.shape\n        return z.reshape(b, self.heads, self.dim_head, h * w)\n\n    def forward(self, x):\n        original_dtype = x.dtype\n\n        # Quan trọng: toàn bộ FFT/FSA chạy FP32, không để AMP biến thành half.\n        with amp_disabled_context():\n            xn = self.norm(x.float())\n\n            q = torch.fft.fft2(self.q(xn), norm=\"ortho\")\n            k = torch.fft.fft2(self.k(xn), norm=\"ortho\")\n            v = torch.fft.fft2(self.v(xn), norm=\"ortho\")\n\n            fw = torch.sigmoid(self.freq_weight.float())\n            q = q * fw.to(q.dtype)\n            k = k * fw.to(k.dtype)\n            v = v * fw.to(v.dtype)\n\n            qh = self._reshape_heads(q)\n            kh = self._reshape_heads(k)\n            vh = self._reshape_heads(v)\n\n            attn = torch.matmul(qh, kh.conj().transpose(-2, -1))\n            attn = attn * self.temperature.view(1, self.heads, 1, 1).float()\n\n            attn_real = torch.softmax(attn.real, dim=-1)\n            attn_imag = torch.softmax(attn.imag, dim=-1)\n            attn_complex = torch.complex(attn_real, attn_imag)\n\n            out = torch.matmul(attn_complex, vh).reshape_as(v)\n            out = torch.fft.ifft2(out, norm=\"ortho\").abs()\n\n            fx = torch.fft.fft2(xn, norm=\"ortho\")\n            res = torch.fft.ifft2(torch.sigmoid(torch.abs(fx)) * fx, norm=\"ortho\").abs()\n\n            y = self.proj(torch.cat([out, res], dim=1))\n\n        return y.to(original_dtype)\n\nclass SpatialSelfAttention(nn.Module):\n    def __init__(self, channels, heads=8):\n        super().__init__()\n        assert channels % heads == 0\n        self.channels = channels\n        self.heads = heads\n\n        self.norm = LayerNorm2d(channels)\n        self.q = nn.Conv2d(channels, channels, 1, bias=False)\n        self.k = nn.Conv2d(channels, channels, 1, bias=False)\n        self.v = nn.Conv2d(channels, channels, 1, bias=False)\n\n        self.q3 = DepthwiseSeparableConv(channels, 3)\n        self.q5 = DepthwiseSeparableConv(channels, 5)\n        self.k3 = DepthwiseSeparableConv(channels, 3)\n        self.k5 = DepthwiseSeparableConv(channels, 5)\n        self.v3 = DepthwiseSeparableConv(channels, 3)\n        self.v5 = DepthwiseSeparableConv(channels, 5)\n\n        self.r3 = DepthwiseOnlyConv(channels, 3)\n        self.r5 = DepthwiseOnlyConv(channels, 5)\n\n        self.temperature = nn.Parameter(torch.ones(heads, 1, 1))\n        self.proj = nn.Sequential(nn.Conv2d(channels * 4, channels, 1, bias=False), nn.BatchNorm2d(channels))\n\n    def _channel_attention(self, q, k, v):\n        b, c2, h, w = q.shape\n        n = h * w\n        dim = c2 // self.heads\n\n        q = q.reshape(b, self.heads, dim, n)\n        k = k.reshape(b, self.heads, dim, n)\n        v = v.reshape(b, self.heads, dim, n)\n\n        q = F.normalize(q, dim=-1)\n        k = F.normalize(k, dim=-1)\n\n        attn = torch.matmul(q, k.transpose(-2, -1))\n        attn = attn * self.temperature.view(1, self.heads, 1, 1)\n        attn = torch.softmax(attn, dim=-1)\n\n        out = torch.matmul(attn, v)\n        return out.reshape(b, c2, h, w)\n\n    def forward(self, x):\n        xn = self.norm(x)\n        q0 = self.q(xn)\n        k0 = self.k(xn)\n        v0 = self.v(xn)\n\n        qs = torch.cat([self.q3(q0), self.q5(q0)], dim=1)\n        ks = torch.cat([self.k3(k0), self.k5(k0)], dim=1)\n        vs = torch.cat([self.v3(v0), self.v5(v0)], dim=1)\n\n        attn_out = self._channel_attention(qs, ks, vs)\n        spatial_res = torch.cat([self.r3(xn), self.r5(xn)], dim=1)\n\n        return self.proj(torch.cat([attn_out, spatial_res], dim=1))\n\nclass CrossDomainAttentionFusion(nn.Module):\n    def __init__(self, channels, heads=8):\n        super().__init__()\n        assert channels % heads == 0\n        self.heads = heads\n        self.dim_head = channels // heads\n\n        self.q = nn.Conv2d(channels, channels, 1, bias=False)\n        self.k = nn.Conv2d(channels, channels, 1, bias=False)\n        self.v = nn.Conv2d(channels, channels, 1, bias=False)\n\n        self.temperature = nn.Parameter(torch.ones(heads, 1, 1))\n        self.proj = nn.Sequential(nn.Conv2d(channels, channels, 1, bias=False), nn.BatchNorm2d(channels))\n\n    def forward(self, xf, xs):\n        b, c, h, w = xf.shape\n        n = h * w\n\n        q = self.q(xf).reshape(b, self.heads, self.dim_head, n)\n        k = self.k(xs).reshape(b, self.heads, self.dim_head, n)\n        v = self.v(xs).reshape(b, self.heads, self.dim_head, n)\n\n        q = F.normalize(q, dim=-1)\n        k = F.normalize(k, dim=-1)\n\n        attn = torch.matmul(q, k.transpose(-2, -1))\n        attn = attn * self.temperature.view(1, self.heads, 1, 1)\n        attn = torch.softmax(attn, dim=-1)\n\n        out = torch.matmul(attn, v).reshape(b, c, h, w)\n        return self.proj(out)\n\nclass DualDomainPerceptionTransformer(nn.Module):\n    def __init__(self, channels, heads=8):\n        super().__init__()\n        self.fsa = FrequencySelfAttention(channels, heads=heads)\n        self.ssa = SpatialSelfAttention(channels, heads=heads)\n        self.cross = CrossDomainAttentionFusion(channels, heads=heads)\n        self.act = nn.GELU()\n\n    def forward(self, x):\n        xf = self.fsa(x)\n        xs = self.ssa(x)\n        fused = self.cross(xf, xs)\n        return self.act(x + fused)","metadata":{},"outputs":[],"execution_count":null},{"id":"8eec249c","cell_type":"markdown","source":"## Cell 7 — Full model","metadata":{}},{"id":"1f2f72fb","cell_type":"code","source":"class JFSPViT(nn.Module):\n    def __init__(self, num_classes=5, ema_groups=8, ddpt_heads=8):\n        super().__init__()\n\n        self.conv_layer_1 = ConvBNAct(3, 16, kernel=3, stride=2, padding=1)\n        self.irb_1 = InvertedResidualBlock(16, 32, stride=1, expand_ratio=4)\n        self.irab_1 = InvertedResidualAttentionBlock(32, 64, stride=2, expand_ratio=4, ema_groups=ema_groups)\n\n        self.irb_2 = nn.Sequential(\n            InvertedResidualBlock(64, 64, stride=1, expand_ratio=4),\n            InvertedResidualBlock(64, 64, stride=1, expand_ratio=4),\n        )\n\n        self.irab_2 = InvertedResidualAttentionBlock(64, 96, stride=2, expand_ratio=4, ema_groups=ema_groups)\n        self.ddpt_96 = nn.Sequential(*[DualDomainPerceptionTransformer(96, heads=ddpt_heads) for _ in range(3)])\n\n        self.irab_3 = InvertedResidualAttentionBlock(96, 128, stride=2, expand_ratio=4, ema_groups=ema_groups)\n        self.ddpt_128 = nn.Sequential(*[DualDomainPerceptionTransformer(128, heads=ddpt_heads) for _ in range(3)])\n\n        self.irab_4 = InvertedResidualAttentionBlock(128, 160, stride=2, expand_ratio=4, ema_groups=ema_groups)\n        self.ddpt_160 = nn.Sequential(*[DualDomainPerceptionTransformer(160, heads=ddpt_heads) for _ in range(3)])\n\n        self.final_conv = ConvBNAct(160, 640, kernel=1, stride=1, padding=0)\n        self.pool = nn.AdaptiveAvgPool2d(1)\n        self.classifier = nn.Linear(640, num_classes)\n\n    def forward_features(self, x):\n        x = self.conv_layer_1(x)\n        x = self.irb_1(x)\n        x = self.irab_1(x)\n        x = self.irb_2(x)\n        x = self.irab_2(x)\n        x = self.ddpt_96(x)\n        x = self.irab_3(x)\n        x = self.ddpt_128(x)\n        x = self.irab_4(x)\n        x = self.ddpt_160(x)\n        x = self.final_conv(x)\n        return x\n\n    def forward(self, x):\n        x = self.forward_features(x)\n        x = self.pool(x).flatten(1)\n        return self.classifier(x)\n\ndef count_parameters(model):\n    return sum(p.numel() for p in model.parameters() if p.requires_grad)\n\ndef count_params_by_top_module(model):\n    rows = []\n    for name, module in model.named_children():\n        params = sum(p.numel() for p in module.parameters() if p.requires_grad)\n        rows.append({\"module\": name, \"params\": params, \"params_M\": params / 1e6})\n    return pd.DataFrame(rows)\n\nmodel = JFSPViT(num_classes=NUM_CLASSES, ema_groups=8, ddpt_heads=8).to(DEVICE)\n\nprint(\"Model:\", model.__class__.__name__)\nprint(f\"Trainable parameters: {count_parameters(model) / 1e6:.3f}M\")\ndisplay(count_params_by_top_module(model))\n\nwith torch.no_grad():\n    dummy = torch.randn(1, 3, IMG_SIZE, IMG_SIZE).to(DEVICE)\n    out = model(dummy)\nprint(\"Output shape:\", out.shape)\n\ndel dummy, out\ngc.collect()\ntorch.cuda.empty_cache()","metadata":{},"outputs":[],"execution_count":null},{"id":"da5ef545","cell_type":"markdown","source":"## Cell 8 — Loss/optimizer/scheduler builders","metadata":{}},{"id":"8584d418","cell_type":"code","source":"def get_sqrt_class_weights(train_df):\n    counts = train_df[\"label\"].value_counts().sort_index()\n    weights = np.zeros(NUM_CLASSES, dtype=np.float32)\n\n    for cls in range(NUM_CLASSES):\n        weights[cls] = len(train_df) / (NUM_CLASSES * max(1, counts.get(cls, 0)))\n\n    weights = np.sqrt(weights)\n    weights = weights / weights.mean()\n    return torch.tensor(weights, dtype=torch.float32).to(DEVICE)\n\n\nCLASS_WEIGHTS = get_sqrt_class_weights(train_df)\nprint(\"sqrt class weights:\", CLASS_WEIGHTS)\n\n\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=1.5, weight=None, label_smoothing=0.0):\n        super().__init__()\n        self.gamma = gamma\n        self.weight = weight\n        self.label_smoothing = label_smoothing\n\n    def forward(self, logits, targets):\n        ce = F.cross_entropy(\n            logits,\n            targets,\n            weight=self.weight,\n            reduction=\"none\",\n            label_smoothing=self.label_smoothing,\n        )\n        pt = torch.exp(-ce)\n        loss = ((1 - pt) ** self.gamma) * ce\n        return loss.mean()\n\n\ndef build_criterion(loss_name):\n    if loss_name == \"CE\":\n        return nn.CrossEntropyLoss()\n    if loss_name == \"WCE\":\n        return nn.CrossEntropyLoss(weight=CLASS_WEIGHTS)\n    if loss_name == \"CE_LS005\":\n        return nn.CrossEntropyLoss(label_smoothing=0.05)\n    if loss_name == \"WCE_LS005\":\n        return nn.CrossEntropyLoss(weight=CLASS_WEIGHTS, label_smoothing=0.05)\n    if loss_name == \"WCE_LS010\":\n        return nn.CrossEntropyLoss(weight=CLASS_WEIGHTS, label_smoothing=0.10)\n    if loss_name == \"FOCAL\":\n        return FocalLoss(gamma=1.5, weight=None, label_smoothing=0.0)\n    if loss_name == \"WFOCAL\":\n        return FocalLoss(gamma=1.5, weight=CLASS_WEIGHTS, label_smoothing=0.05)\n\n    raise ValueError(f\"Unknown loss: {loss_name}\")\n\n\ndef build_optimizer(model, opt_name, lr, weight_decay):\n    if opt_name == \"AdamW\":\n        return torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)\n    if opt_name == \"Adam\":\n        return torch.optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay)\n    if opt_name == \"RAdam\":\n        return torch.optim.RAdam(model.parameters(), lr=lr, weight_decay=weight_decay)\n    if opt_name == \"Adamax\":\n        return torch.optim.Adamax(model.parameters(), lr=lr, weight_decay=weight_decay)\n    if opt_name == \"SGD\":\n        return torch.optim.SGD(model.parameters(), lr=lr, momentum=0.9, weight_decay=weight_decay, nesterov=True)\n    if opt_name == \"RMSprop\":\n        return torch.optim.RMSprop(model.parameters(), lr=lr, momentum=0.9, weight_decay=weight_decay)\n\n    raise ValueError(f\"Unknown optimizer: {opt_name}\")\n\n\ndef build_scheduler(optimizer, sched_name, epochs):\n    if sched_name == \"None\":\n        return None\n    if sched_name == \"Cosine\":\n        return torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)\n    if sched_name == \"Plateau\":\n        return torch.optim.lr_scheduler.ReduceLROnPlateau(\n            optimizer,\n            mode=\"max\",\n            factor=0.5,\n            patience=3,\n            min_lr=1e-6,\n        )\n    if sched_name == \"Step\":\n        return torch.optim.lr_scheduler.StepLR(optimizer, step_size=6, gamma=0.5)\n\n    raise ValueError(f\"Unknown scheduler: {sched_name}\")\n\n\n# Nhóm chính: không scheduler, LR cố định 0.00019.\nFAST_SWEEP_CONFIGS = [\n    {\"name\": \"AdamW_CE_NoSched_paperLR\",        \"opt\": \"AdamW\", \"loss\": \"CE\",        \"sched\": \"None\", \"lr\": PAPER_LR, \"wd\": 0.0},\n    {\"name\": \"AdamW_CE_NoSched_wd1e4_paperLR\",  \"opt\": \"AdamW\", \"loss\": \"CE\",        \"sched\": \"None\", \"lr\": PAPER_LR, \"wd\": 1e-4},\n    {\"name\": \"Adam_CE_NoSched_paperLR\",         \"opt\": \"Adam\",  \"loss\": \"CE\",        \"sched\": \"None\", \"lr\": PAPER_LR, \"wd\": 0.0},\n    {\"name\": \"Adam_CE_NoSched_wd1e4_paperLR\",   \"opt\": \"Adam\",  \"loss\": \"CE\",        \"sched\": \"None\", \"lr\": PAPER_LR, \"wd\": 1e-4},\n    {\"name\": \"RAdam_CE_NoSched_paperLR\",        \"opt\": \"RAdam\", \"loss\": \"CE\",        \"sched\": \"None\", \"lr\": PAPER_LR, \"wd\": 0.0},\n    {\"name\": \"AdamW_WCE_NoSched_paperLR\",       \"opt\": \"AdamW\", \"loss\": \"WCE\",       \"sched\": \"None\", \"lr\": PAPER_LR, \"wd\": 1e-4},\n]\n\n# Nhóm phụ: vẫn giữ LR ban đầu 0.00019 nhưng có scheduler, chỉ dùng để chẩn đoán.\nFULL_SWEEP_CONFIGS = FAST_SWEEP_CONFIGS + [\n    {\"name\": \"AdamW_CE_Cosine_paperLR\",          \"opt\": \"AdamW\", \"loss\": \"CE\",        \"sched\": \"Cosine\",  \"lr\": PAPER_LR, \"wd\": 1e-4},\n    {\"name\": \"AdamW_WCE_Cosine_paperLR\",         \"opt\": \"AdamW\", \"loss\": \"WCE\",       \"sched\": \"Cosine\",  \"lr\": PAPER_LR, \"wd\": 1e-4},\n    {\"name\": \"AdamW_CE_Plateau_paperLR\",         \"opt\": \"AdamW\", \"loss\": \"CE\",        \"sched\": \"Plateau\", \"lr\": PAPER_LR, \"wd\": 1e-4},\n    {\"name\": \"AdamW_WCE_Plateau_paperLR\",        \"opt\": \"AdamW\", \"loss\": \"WCE\",       \"sched\": \"Plateau\", \"lr\": PAPER_LR, \"wd\": 1e-4},\n    {\"name\": \"AdamW_WCE_LS005_NoSched_paperLR\",  \"opt\": \"AdamW\", \"loss\": \"WCE_LS005\", \"sched\": \"None\",    \"lr\": PAPER_LR, \"wd\": 1e-4},\n    {\"name\": \"Adamax_CE_NoSched_paperLR\",        \"opt\": \"Adamax\",\"loss\": \"CE\",        \"sched\": \"None\",    \"lr\": PAPER_LR, \"wd\": 0.0},\n]\n\nSWEEP_CONFIGS = FAST_SWEEP_CONFIGS if SWEEP_MODE == \"FAST\" else FULL_SWEEP_CONFIGS\nprint(\"Number of sweep configs:\", len(SWEEP_CONFIGS))\npd.DataFrame(SWEEP_CONFIGS)","metadata":{},"outputs":[],"execution_count":null},{"id":"44b2b56a","cell_type":"markdown","source":"## Cell 9 — Metrics","metadata":{}},{"id":"4960182b","cell_type":"code","source":"def safe_div(a, b):\n    return a / b if b != 0 else 0.0\n\n\ndef compute_all_metrics(\n    y_true,\n    y_pred,\n    y_prob=None,\n    num_classes=5,\n    class_names=None\n):\n    \"\"\"\n    Tính bộ metric thống nhất cho bài toán phân loại 5 mức DR.\n\n    Specificity được tính theo one-vs-rest:\n        Specificity = TN / (TN + FP)\n    \"\"\"\n    y_true = np.asarray(\n        y_true\n    ).reshape(-1).astype(int)\n\n    y_pred = np.asarray(\n        y_pred\n    ).reshape(-1).astype(int)\n\n    labels = list(\n        range(num_classes)\n    )\n\n    display_names = {\n        0: \"No DR\",\n        1: \"Mild\",\n        2: \"Moderate\",\n        3: \"Severe\",\n        4: \"Proliferative\",\n    }\n\n    if class_names is not None:\n        for label in labels:\n            if label not in display_names:\n                display_names[label] = str(\n                    class_names.get(label, label)\n                )\n\n    # =====================================================\n    # CONFUSION MATRIX\n    # =====================================================\n\n    cm = confusion_matrix(\n        y_true,\n        y_pred,\n        labels=labels\n    )\n\n    total = int(\n        cm.sum()\n    )\n\n    # =====================================================\n    # CLASS-WISE PRECISION / RECALL / F1 / SUPPORT\n    # =====================================================\n\n    (\n        precision_arr,\n        recall_arr,\n        f1_arr,\n        support_arr\n    ) = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        labels=labels,\n        average=None,\n        zero_division=0\n    )\n\n    (\n        macro_precision,\n        macro_recall,\n        macro_f1,\n        _\n    ) = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        labels=labels,\n        average=\"macro\",\n        zero_division=0\n    )\n\n    (\n        weighted_precision,\n        weighted_recall,\n        weighted_f1,\n        _\n    ) = precision_recall_fscore_support(\n        y_true,\n        y_pred,\n        labels=labels,\n        average=\"weighted\",\n        zero_division=0\n    )\n\n    # =====================================================\n    # SPECIFICITY THEO ONE-VS-REST\n    # =====================================================\n\n    specificity_arr = []\n    per_class_rows = []\n\n    for class_index, label in enumerate(labels):\n        tp = int(\n            cm[class_index, class_index]\n        )\n\n        fp = int(\n            cm[:, class_index].sum()\n            - tp\n        )\n\n        fn = int(\n            cm[class_index, :].sum()\n            - tp\n        )\n\n        tn = int(\n            total\n            - tp\n            - fp\n            - fn\n        )\n\n        specificity = safe_div(\n            tn,\n            tn + fp\n        )\n\n        specificity_arr.append(\n            specificity\n        )\n\n        per_class_rows.append(\n            {\n                \"Class\": display_names[label],\n                \"Precision\": precision_arr[class_index],\n                \"Recall\": recall_arr[class_index],\n                \"Specificity\": specificity,\n                \"F1-Score\": f1_arr[class_index],\n                \"Support\": int(support_arr[class_index]),\n                \"TP\": tp,\n                \"FP\": fp,\n                \"FN\": fn,\n                \"TN\": tn,\n            }\n        )\n\n    specificity_arr = np.asarray(\n        specificity_arr,\n        dtype=float\n    )\n\n    specificity_macro = float(\n        np.mean(specificity_arr)\n    )\n\n    specificity_weighted = float(\n        np.average(\n            specificity_arr,\n            weights=support_arr\n        )\n    )\n\n    # =====================================================\n    # OVERALL METRICS\n    # =====================================================\n\n    accuracy = accuracy_score(\n        y_true,\n        y_pred\n    )\n\n    balanced_acc = balanced_accuracy_score(\n        y_true,\n        y_pred\n    )\n\n    mcc = matthews_corrcoef(\n        y_true,\n        y_pred\n    )\n\n    qwk = cohen_kappa_score(\n        y_true,\n        y_pred,\n        labels=labels,\n        weights=\"quadratic\"\n    )\n\n    within_1_grade_acc = float(\n        np.mean(\n            np.abs(y_true - y_pred) <= 1\n        )\n    )\n\n    overall = {\n        \"Accuracy\": accuracy,\n        \"BalancedAcc\": balanced_acc,\n        \"Precision Macro\": macro_precision,\n        \"Precision Weighted\": weighted_precision,\n        \"Recall Macro\": macro_recall,\n        \"Recall Weighted\": weighted_recall,\n        \"Specificity Macro\": specificity_macro,\n        \"Specificity Weighted\": specificity_weighted,\n        \"F1-Score Macro\": macro_f1,\n        \"F1-Score Weighted\": weighted_f1,\n        \"MCC\": mcc,\n        \"QWK\": qwk,\n        \"Within-1-Grade Acc\": within_1_grade_acc,\n    }\n\n    # Giữ AUC nếu xác suất dự đoán có sẵn\n    if y_prob is not None:\n        try:\n            y_prob = np.asarray(\n                y_prob\n            )\n\n            y_onehot = np.eye(\n                num_classes\n            )[y_true]\n\n            overall[\"ROC-AUC Macro OVR\"] = roc_auc_score(\n                y_onehot,\n                y_prob,\n                average=\"macro\",\n                multi_class=\"ovr\"\n            )\n\n            overall[\"ROC-AUC Micro OVR\"] = roc_auc_score(\n                y_onehot,\n                y_prob,\n                average=\"micro\",\n                multi_class=\"ovr\"\n            )\n\n        except Exception as error:\n            overall[\"ROC-AUC Macro OVR\"] = np.nan\n            overall[\"ROC-AUC Micro OVR\"] = np.nan\n\n            print(\n                \"Không thể tính ROC-AUC:\",\n                error\n            )\n\n    # =====================================================\n    # DATAFRAMES\n    # =====================================================\n\n    overall_df = pd.DataFrame(\n        [overall]\n    )\n\n    overall_long_df = pd.DataFrame(\n        {\n            \"Metric\": list(overall.keys()),\n            \"Value\": list(overall.values()),\n        }\n    )\n\n    per_class_df = pd.DataFrame(\n        per_class_rows\n    )\n\n    class_labels = [\n        display_names[label]\n        for label in labels\n    ]\n\n    cm_df = pd.DataFrame(\n        cm,\n        index=class_labels,\n        columns=class_labels\n    )\n\n    cm_df.index.name = \"Actual\"\n    cm_df.columns.name = \"Predicted\"\n\n    cm_normalized = confusion_matrix(\n        y_true,\n        y_pred,\n        labels=labels,\n        normalize=\"true\"\n    )\n\n    cm_normalized_df = pd.DataFrame(\n        cm_normalized,\n        index=class_labels,\n        columns=class_labels\n    )\n\n    cm_normalized_df.index.name = \"Actual\"\n    cm_normalized_df.columns.name = \"Predicted\"\n\n    # Giữ tên biến paper_style_df để không phá cấu trúc cũ\n    paper_style_df = overall_df.copy()\n\n    return (\n        paper_style_df,\n        overall_long_df,\n        per_class_df,\n        cm_df,\n        cm_normalized_df\n    )\n\n\ndef print_standard_metrics(\n    overall_long_df,\n    per_class_df\n):\n    \"\"\"\n    In kết quả theo cùng một mẫu đã dùng cho các notebook DR trước.\n    \"\"\"\n    metric_map = dict(\n        zip(\n            overall_long_df[\"Metric\"],\n            overall_long_df[\"Value\"]\n        )\n    )\n\n    print(\"\\n\" + \"=\" * 31)\n    print(\"OVERALL METRICS\")\n    print(\"=\" * 31)\n\n    print(\n        f\"{'Accuracy':<22}: \"\n        f\"{metric_map['Accuracy']:.4f}\"\n    )\n\n    print(\n        f\"{'BalancedAcc':<22}: \"\n        f\"{metric_map['BalancedAcc']:.4f}\"\n    )\n\n    print(\"-\" * 31)\n\n    print(\n        f\"{'Precision Macro':<22}: \"\n        f\"{metric_map['Precision Macro']:.4f}\"\n    )\n\n    print(\n        f\"{'Precision Weighted':<22}: \"\n        f\"{metric_map['Precision Weighted']:.4f}\"\n    )\n\n    print(\"-\" * 31)\n\n    print(\n        f\"{'Recall Macro':<22}: \"\n        f\"{metric_map['Recall Macro']:.4f}\"\n    )\n\n    print(\n        f\"{'Recall Weighted':<22}: \"\n        f\"{metric_map['Recall Weighted']:.4f}\"\n    )\n\n    print(\"-\" * 31)\n\n    print(\n        f\"{'Specificity Macro':<22}: \"\n        f\"{metric_map['Specificity Macro']:.4f}\"\n    )\n\n    print(\n        f\"{'Specificity Weighted':<22}: \"\n        f\"{metric_map['Specificity Weighted']:.4f}\"\n    )\n\n    print(\"-\" * 31)\n\n    print(\n        f\"{'F1-Score Macro':<22}: \"\n        f\"{metric_map['F1-Score Macro']:.4f}\"\n    )\n\n    print(\n        f\"{'F1-Score Weighted':<22}: \"\n        f\"{metric_map['F1-Score Weighted']:.4f}\"\n    )\n\n    print(\"-\" * 31)\n\n    print(\n        f\"{'MCC':<22}: \"\n        f\"{metric_map['MCC']:.4f}\"\n    )\n\n    print(\n        f\"{'QWK':<22}: \"\n        f\"{metric_map['QWK']:.4f}\"\n    )\n\n    print(\n        f\"{'Within-1-Grade Acc':<22}: \"\n        f\"{metric_map['Within-1-Grade Acc']:.4f}\"\n    )\n\n    if \"ROC-AUC Macro OVR\" in metric_map:\n        print(\"-\" * 31)\n\n        print(\n            f\"{'ROC-AUC Macro OVR':<22}: \"\n            f\"{metric_map['ROC-AUC Macro OVR']:.4f}\"\n        )\n\n        print(\n            f\"{'ROC-AUC Micro OVR':<22}: \"\n            f\"{metric_map['ROC-AUC Micro OVR']:.4f}\"\n        )\n\n    print(\"=\" * 31)\n\n    print(\"\\n--- CLASS-WISE METRICS ---\")\n\n    header = (\n        f\"{'Class':<15} | \"\n        f\"{'Precision':>9} | \"\n        f\"{'Recall':>6} | \"\n        f\"{'Specificity':>11} | \"\n        f\"{'F1-Score':>8} | \"\n        f\"{'Support':>7}\"\n    )\n\n    print(header)\n    print(\"-\" * len(header))\n\n    for _, row in per_class_df.iterrows():\n        print(\n            f\"{row['Class']:<15} | \"\n            f\"{row['Precision']:>9.4f} | \"\n            f\"{row['Recall']:>6.4f} | \"\n            f\"{row['Specificity']:>11.4f} | \"\n            f\"{row['F1-Score']:>8.4f} | \"\n            f\"{int(row['Support']):>7d}\"\n        )\n\n    print(\"=\" * len(header))","metadata":{},"outputs":[],"execution_count":null},{"id":"63dc40ef","cell_type":"markdown","source":"## Cell 10 — Local train/eval function for recipe sweep","metadata":{}},{"id":"345cc09e","cell_type":"code","source":"def autocast_context():\n    return torch.amp.autocast(device_type=\"cuda\", enabled=(USE_AMP and DEVICE == \"cuda\"))\n\ndef get_gpu_memory_gb():\n    if DEVICE != \"cuda\":\n        return None\n    return torch.cuda.max_memory_allocated() / (1024 ** 3)\n\ndef run_one_epoch_local(\n    model,\n    loader,\n    criterion,\n    optimizer=None,\n    scaler=None,\n    train=True,\n    epoch=None,\n    total_epochs=None,\n    stage=None,\n    show_progress=True,\n):\n    model.train(train)\n\n    if stage is None:\n        stage = \"Train\" if train else \"Eval\"\n\n    desc = f\"{stage} | Epoch {epoch}/{total_epochs}\" if epoch is not None and total_epochs is not None else stage\n\n    total_loss = 0.0\n    total_correct = 0\n    total_seen = 0\n\n    all_targets, all_preds, all_probs = [], [], []\n\n    if train:\n        optimizer.zero_grad(set_to_none=True)\n\n    if DEVICE == \"cuda\":\n        torch.cuda.reset_peak_memory_stats()\n\n    pbar = tqdm(loader, desc=desc, leave=False, dynamic_ncols=True, unit=\"batch\", disable=not show_progress)\n\n    for step, (images, targets) in enumerate(pbar):\n        images = images.to(DEVICE, non_blocking=True)\n        targets = targets.to(DEVICE, non_blocking=True)\n\n        with torch.set_grad_enabled(train):\n            with autocast_context():\n                logits = model(images)\n                loss = criterion(logits, targets)\n\n            if train:\n                scaler.scale(loss / ACCUM_STEPS).backward()\n\n                if (step + 1) % ACCUM_STEPS == 0 or (step + 1) == len(loader):\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad(set_to_none=True)\n\n        probs = torch.softmax(logits.detach(), dim=1)\n        preds = probs.argmax(dim=1)\n\n        batch_size = images.size(0)\n        total_loss += loss.item() * batch_size\n        total_correct += (preds == targets).sum().item()\n        total_seen += batch_size\n\n        all_targets.extend(targets.detach().cpu().numpy().tolist())\n        all_preds.extend(preds.detach().cpu().numpy().tolist())\n        all_probs.extend(probs.detach().cpu().numpy().tolist())\n\n        postfix = {\n            \"avg_loss\": f\"{total_loss / max(1, total_seen):.4f}\",\n            \"acc\": f\"{total_correct / max(1, total_seen):.4f}\",\n        }\n        if train and optimizer is not None:\n            postfix[\"lr\"] = f\"{optimizer.param_groups[0]['lr']:.2e}\"\n        if DEVICE == \"cuda\":\n            postfix[\"vram_GB\"] = f\"{get_gpu_memory_gb():.2f}\"\n\n        pbar.set_postfix(postfix)\n\n    avg_loss = total_loss / len(loader.dataset)\n    y_true = np.array(all_targets)\n    y_pred = np.array(all_preds)\n    y_prob = np.array(all_probs)\n\n    acc = accuracy_score(y_true, y_pred)\n    qwk = cohen_kappa_score(y_true, y_pred, weights=\"quadratic\")\n\n    return avg_loss, acc, qwk, y_true, y_pred, y_prob","metadata":{},"outputs":[],"execution_count":null},{"id":"215c46bf","cell_type":"markdown","source":"## Cell 11 — Sweep optimizer/loss/scheduler","metadata":{}},{"id":"675abc30","cell_type":"code","source":"SWEEP_DIR = SAVE_DIR / \"sweep_checkpoints\"\nSWEEP_DIR.mkdir(parents=True, exist_ok=True)\n\ndef make_new_model():\n    seed_everything(SEED)\n    m = JFSPViT(num_classes=NUM_CLASSES, ema_groups=8, ddpt_heads=8).to(DEVICE)\n    return m\n\ndef run_recipe_config(cfg, epochs=SWEEP_EPOCHS, patience=SWEEP_PATIENCE):\n    print(\"=\" * 100)\n    print(\"Running config:\", cfg[\"name\"])\n    print(cfg)\n\n    model_i = make_new_model()\n    criterion_i = build_criterion(cfg[\"loss\"])\n    optimizer_i = build_optimizer(model_i, cfg[\"opt\"], cfg[\"lr\"], cfg[\"wd\"])\n    scheduler_i = build_scheduler(optimizer_i, cfg[\"sched\"], epochs)\n\n    try:\n        scaler_i = torch.amp.GradScaler(\"cuda\", enabled=(USE_AMP and DEVICE == \"cuda\"))\n    except Exception:\n        from torch.cuda.amp import GradScaler\n        scaler_i = GradScaler(enabled=(USE_AMP and DEVICE == \"cuda\"))\n\n    best_qwk = -1\n    best_epoch = -1\n    no_improve = 0\n    history = []\n\n    best_path_i = SWEEP_DIR / f\"{cfg['name']}.pth\"\n    hist_path_i = SWEEP_DIR / f\"{cfg['name']}_history.csv\"\n\n    for epoch in range(1, epochs + 1):\n        train_loss, train_acc, train_qwk, _, _, _ = run_one_epoch_local(\n            model_i,\n            train_loader,\n            criterion_i,\n            optimizer=optimizer_i,\n            scaler=scaler_i,\n            train=True,\n            epoch=epoch,\n            total_epochs=epochs,\n            stage=f\"Train {cfg['name']}\",\n            show_progress=True,\n        )\n\n        val_loss, val_acc, val_qwk, _, _, _ = run_one_epoch_local(\n            model_i,\n            val_loader,\n            criterion_i,\n            train=False,\n            epoch=epoch,\n            total_epochs=epochs,\n            stage=f\"Val {cfg['name']}\",\n            show_progress=True,\n        )\n\n        if scheduler_i is not None:\n            if cfg[\"sched\"] == \"Plateau\":\n                scheduler_i.step(val_qwk)\n            else:\n                scheduler_i.step()\n\n        gap = train_qwk - val_qwk\n\n        row = {\n            \"config_name\": cfg[\"name\"],\n            \"epoch\": epoch,\n            \"lr\": optimizer_i.param_groups[0][\"lr\"],\n            \"train_loss\": train_loss,\n            \"train_acc\": train_acc,\n            \"train_qwk\": train_qwk,\n            \"val_loss\": val_loss,\n            \"val_acc\": val_acc,\n            \"val_qwk\": val_qwk,\n            \"qwk_gap\": gap,\n        }\n        history.append(row)\n\n        if val_qwk > best_qwk + 1e-4:\n            best_qwk = val_qwk\n            best_epoch = epoch\n            no_improve = 0\n            torch.save({\n                \"model_state\": model_i.state_dict(),\n                \"config\": cfg,\n                \"best_qwk\": best_qwk,\n                \"best_epoch\": best_epoch,\n                \"history\": history,\n            }, best_path_i)\n            print(f\"Saved best config model: {best_path_i} | best_qwk={best_qwk:.4f} | epoch={best_epoch}\")\n        else:\n            no_improve += 1\n\n        print(\n            f\"{cfg['name']} | Epoch {epoch:03d}/{epochs} | \"\n            f\"train_qwk={train_qwk:.4f} | val_qwk={val_qwk:.4f} | \"\n            f\"gap={gap:.4f} | val_loss={val_loss:.4f} | \"\n            f\"best={best_qwk:.4f}@{best_epoch} | no_improve={no_improve}\"\n        )\n\n        if epoch >= 8 and no_improve >= patience:\n            print(f\"Early stop config {cfg['name']} at epoch {epoch}\")\n            break\n\n    hist_df = pd.DataFrame(history)\n    hist_df.to_csv(hist_path_i, index=False)\n\n    best_row = hist_df.loc[hist_df[\"val_qwk\"].idxmax()].to_dict()\n    best_row.update({\n        \"config_name\": cfg[\"name\"],\n        \"optimizer\": cfg[\"opt\"],\n        \"loss\": cfg[\"loss\"],\n        \"scheduler\": cfg[\"sched\"],\n        \"init_lr\": cfg[\"lr\"],\n        \"weight_decay\": cfg[\"wd\"],\n        \"checkpoint\": str(best_path_i),\n    })\n\n    del model_i, criterion_i, optimizer_i, scheduler_i, scaler_i\n    gc.collect()\n    torch.cuda.empty_cache()\n\n    return best_row\n\nsweep_results = []\nfor cfg in SWEEP_CONFIGS:\n    result = run_recipe_config(cfg)\n    sweep_results.append(result)\n\nsweep_df = pd.DataFrame(sweep_results)\n\nsweep_df[\"overfit_flag\"] = (sweep_df[\"qwk_gap\"] > 0.15) | (sweep_df[\"val_loss\"] > 1.5)\nsweep_df[\"score_for_selection\"] = sweep_df[\"val_qwk\"] - 0.25 * np.maximum(sweep_df[\"qwk_gap\"], 0)\n\nsweep_df = sweep_df.sort_values(\n    by=[\"score_for_selection\", \"val_qwk\"],\n    ascending=False,\n).reset_index(drop=True)\n\nsweep_df.to_csv(SAVE_DIR / \"recipe_sweep_summary.csv\", index=False)\ndisplay(sweep_df)\n\nBEST_CONFIG = SWEEP_CONFIGS[[c[\"name\"] for c in SWEEP_CONFIGS].index(sweep_df.iloc[0][\"config_name\"])]\nBEST_SWEEP_CHECKPOINT = sweep_df.iloc[0][\"checkpoint\"]\n\nprint(\"BEST_CONFIG:\", BEST_CONFIG)\nprint(\"BEST_SWEEP_CHECKPOINT:\", BEST_SWEEP_CHECKPOINT)","metadata":{},"outputs":[],"execution_count":null},{"id":"85e0eec7","cell_type":"markdown","source":"## Cell 12 — Plot best sweep history","metadata":{}},{"id":"e3174786","cell_type":"code","source":"summary_path = SAVE_DIR / \"recipe_sweep_summary.csv\"\nsweep_df = pd.read_csv(summary_path)\ndisplay(sweep_df)\n\nbest_name = sweep_df.iloc[0][\"config_name\"]\nbest_hist_path = SAVE_DIR / \"sweep_checkpoints\" / f\"{best_name}_history.csv\"\nhistory_df = pd.read_csv(best_hist_path)\n\nplt.figure(figsize=(8, 5))\nplt.plot(history_df[\"epoch\"], history_df[\"train_acc\"], label=\"train_acc\")\nplt.plot(history_df[\"epoch\"], history_df[\"val_acc\"], label=\"val_acc\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Accuracy\")\nplt.title(f\"Best recipe accuracy: {best_name}\")\nplt.legend()\nplt.grid(True)\nplt.savefig(SAVE_DIR / f\"{best_name}_accuracy_curve.png\", dpi=200, bbox_inches=\"tight\")\nplt.show()\n\nplt.figure(figsize=(8, 5))\nplt.plot(history_df[\"epoch\"], history_df[\"train_loss\"], label=\"train_loss\")\nplt.plot(history_df[\"epoch\"], history_df[\"val_loss\"], label=\"val_loss\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(f\"Best recipe loss: {best_name}\")\nplt.legend()\nplt.grid(True)\nplt.savefig(SAVE_DIR / f\"{best_name}_loss_curve.png\", dpi=200, bbox_inches=\"tight\")\nplt.show()\n\nplt.figure(figsize=(8, 5))\nplt.plot(history_df[\"epoch\"], history_df[\"train_qwk\"], label=\"train_qwk\")\nplt.plot(history_df[\"epoch\"], history_df[\"val_qwk\"], label=\"val_qwk\")\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"QWK\")\nplt.title(f\"Best recipe QWK: {best_name}\")\nplt.legend()\nplt.grid(True)\nplt.savefig(SAVE_DIR / f\"{best_name}_qwk_curve.png\", dpi=200, bbox_inches=\"tight\")\nplt.show()","metadata":{},"outputs":[],"execution_count":null},{"id":"c9a8de9b","cell_type":"markdown","source":"## Cell 13 — Test best sweep checkpoint","metadata":{}},{"id":"84f3a499","cell_type":"code","source":"sweep_df = pd.read_csv(\n    SAVE_DIR / \"recipe_sweep_summary.csv\"\n)\n\nbest_row = sweep_df.iloc[0]\nbest_ckpt_path = best_row[\"checkpoint\"]\n\n# Checkpoint này do chính notebook tạo và có chứa cả config,\n# epoch và các giá trị NumPy, không chỉ riêng model weights.\n# PyTorch 2.6+ mặc định weights_only=True nên cần đặt False.\ntry:\n    ckpt = torch.load(\n        best_ckpt_path,\n        map_location=DEVICE,\n        weights_only=False\n    )\nexcept TypeError:\n    # Tương thích với các phiên bản PyTorch cũ\n    ckpt = torch.load(\n        best_ckpt_path,\n        map_location=DEVICE\n    )\n\nbest_cfg = ckpt[\"config\"]\n\nmodel = JFSPViT(\n    num_classes=NUM_CLASSES,\n    ema_groups=8,\n    ddpt_heads=8\n).to(DEVICE)\n\nmodel.load_state_dict(\n    ckpt[\"model_state\"]\n)\n\nmodel.eval()\n\ncriterion_test = build_criterion(\n    best_cfg[\"loss\"]\n)\n\nprint(\"Best config :\", best_cfg)\nprint(\"Best epoch  :\", ckpt.get(\"best_epoch\"))\nprint(\"Best val QWK:\", ckpt.get(\"best_qwk\"))\n\n(\n    test_loss,\n    test_acc,\n    test_qwk,\n    test_true,\n    test_pred,\n    test_prob\n) = run_one_epoch_local(\n    model,\n    test_loader,\n    criterion_test,\n    train=False,\n    stage=\"Test\",\n    show_progress=True\n)\n\n(\n    paper_style_df,\n    overall_df,\n    per_class_df,\n    cm_df,\n    cm_normalized_df\n) = compute_all_metrics(\n    test_true,\n    test_pred,\n    y_prob=test_prob,\n    num_classes=NUM_CLASSES,\n    class_names=CLASS_NAMES\n)\n\nprint(\n    f\"\\nTest Loss: {test_loss:.6f}\"\n)\n\nprint_standard_metrics(\n    overall_df,\n    per_class_df\n)\n\nprint(\"\\n--- CONFUSION MATRIX: COUNTS ---\")\ndisplay(cm_df)\n\nprint(\n    \"\\n--- CONFUSION MATRIX: \"\n    \"NORMALIZED BY TRUE CLASS ---\"\n)\ndisplay(\n    cm_normalized_df.round(4)\n)\n\n# Lưu kết quả\npaper_style_df.to_csv(\n    SAVE_DIR\n    / f\"{best_cfg['name']}_standard_overall_metrics_wide.csv\",\n    index=False\n)\n\noverall_df.to_csv(\n    SAVE_DIR\n    / f\"{best_cfg['name']}_overall_metrics.csv\",\n    index=False\n)\n\nper_class_df.to_csv(\n    SAVE_DIR\n    / f\"{best_cfg['name']}_per_class_detailed_metrics.csv\",\n    index=False\n)\n\ncm_df.to_csv(\n    SAVE_DIR\n    / f\"{best_cfg['name']}_confusion_matrix_counts.csv\"\n)\n\ncm_normalized_df.to_csv(\n    SAVE_DIR\n    / f\"{best_cfg['name']}_confusion_matrix_normalized.csv\"\n)","metadata":{},"outputs":[],"execution_count":null},{"id":"b8802de0","cell_type":"markdown","source":"## Cell 14 — Export confusion matrix and Excel","metadata":{}},{"id":"fad8e5e8","cell_type":"code","source":"def plot_confusion_matrix(\n    matrix_df,\n    title,\n    output_name,\n    value_format=\"d\"\n):\n    matrix = matrix_df.values\n\n    fig, ax = plt.subplots(\n        figsize=(8, 7)\n    )\n\n    image = ax.imshow(\n        matrix\n    )\n\n    ax.set_title(\n        title\n    )\n\n    ax.set_xlabel(\n        \"Predicted label\"\n    )\n\n    ax.set_ylabel(\n        \"True label\"\n    )\n\n    ax.set_xticks(\n        np.arange(\n            matrix.shape[1]\n        )\n    )\n\n    ax.set_yticks(\n        np.arange(\n            matrix.shape[0]\n        )\n    )\n\n    ax.set_xticklabels(\n        matrix_df.columns,\n        rotation=45,\n        ha=\"right\"\n    )\n\n    ax.set_yticklabels(\n        matrix_df.index\n    )\n\n    for row_index in range(\n        matrix.shape[0]\n    ):\n        for column_index in range(\n            matrix.shape[1]\n        ):\n            value = matrix[\n                row_index,\n                column_index\n            ]\n\n            if value_format == \"d\":\n                text = f\"{int(value)}\"\n            else:\n                text = f\"{value:.2f}\"\n\n            ax.text(\n                column_index,\n                row_index,\n                text,\n                ha=\"center\",\n                va=\"center\"\n            )\n\n    fig.colorbar(\n        image,\n        ax=ax\n    )\n\n    plt.tight_layout()\n\n    output_path = (\n        SAVE_DIR\n        / output_name\n    )\n\n    plt.savefig(\n        output_path,\n        dpi=200,\n        bbox_inches=\"tight\"\n    )\n\n    plt.show()\n\n    print(\n        \"Saved:\",\n        output_path\n    )\n\n\nplot_confusion_matrix(\n    cm_df,\n    (\n        f\"{DATASET_NAME} Confusion Matrix - Counts\\n\"\n        f\"{best_cfg['name']}\"\n    ),\n    f\"{best_cfg['name']}_confusion_matrix_counts.png\",\n    value_format=\"d\"\n)\n\nplot_confusion_matrix(\n    cm_normalized_df,\n    (\n        f\"{DATASET_NAME} Confusion Matrix - Normalized\\n\"\n        f\"{best_cfg['name']}\"\n    ),\n    f\"{best_cfg['name']}_confusion_matrix_normalized.png\",\n    value_format=\".2f\"\n)\n\n\nexcel_path = (\n    SAVE_DIR\n    / f\"{DATASET_NAME}_recipe_sweep_all_results.xlsx\"\n)\n\nwith pd.ExcelWriter(\n    excel_path,\n    engine=\"openpyxl\"\n) as writer:\n    sweep_df.to_excel(\n        writer,\n        sheet_name=\"SweepSummary\",\n        index=False\n    )\n\n    paper_style_df.to_excel(\n        writer,\n        sheet_name=\"OverallMetricsWide\",\n        index=False\n    )\n\n    overall_df.to_excel(\n        writer,\n        sheet_name=\"OverallMetrics\",\n        index=False\n    )\n\n    per_class_df.to_excel(\n        writer,\n        sheet_name=\"PerClassDetailed\",\n        index=False\n    )\n\n    cm_df.to_excel(\n        writer,\n        sheet_name=\"CM_Counts\"\n    )\n\n    cm_normalized_df.to_excel(\n        writer,\n        sheet_name=\"CM_Normalized\"\n    )\n\n    history_df.to_excel(\n        writer,\n        sheet_name=\"BestHistory\",\n        index=False\n    )\n\n    count_params_by_top_module(\n        model\n    ).to_excel(\n        writer,\n        sheet_name=\"ParamsByModule\",\n        index=False\n    )\n\nprint(\n    \"Saved Excel:\",\n    excel_path\n)","metadata":{},"outputs":[],"execution_count":null}]}