{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[],"dockerImageVersionId":28755,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"#!/usr/bin/env python3\n# -*- coding: utf-8 -*-\n\"\"\"\n============================================================================\n EKSPERIMEN B -- METODOLOGI + PERBAIKAN (IMPROVED)\n Klasifikasi Retinopati Diabetik, APTOS 2019\n============================================================================\nPasangan ablasi untuk eksperimen_A_baseline.py. SEED, split, dan fungsi\nmetrik dibuat IDENTIK dengan Eksperimen A agar angkanya sebanding langsung.\n\nPERBEDAAN TERHADAP EKSPERIMEN A (inilah variabel yang diuji):\n\n  No | Komponen        | A (Baseline)            | B (Improved)\n  ---|-----------------|-------------------------|---------------------------\n   1 | Imbalance       | class weight saja       | + hybrid under/oversampling\n     |                 |                         |   (target = mean, 512/kelas)\n   2 | Loss            | Weighted Cross-Entropy  | OrdinalFocalLoss\n     |                 |                         |   (focal + target ordinal)\n   3 | Attention       | tidak ada               | CBAM (Woo et al., 2018)\n   4 | Skema training  | Fase1 frozen -> Fase2   | Stage1 supervised penuh ->\n     |                 | unfreeze -> FSL penuh   | Stage2 FSL head (backbone\n     |                 |                         | SELALU frozen)\n   5 | Optimizer       | Adam + ReduceLROnPlateau| AdamW + warmup-cosine\n   6 | Stabilisasi     | tidak ada               | EMA (+ perbaikan penangkapan\n     |                 |                         |   checkpoint saat EMA aktif)\n   7 | Kriteria best   | AUC-ROC val             | sensitivity macro val\n   8 | Augmentasi      | 4 transformasi dasar    | + RandomResizedCrop,\n     |                 |                         |   ColorJitter, GaussianBlur,\n     |                 |                         |   GridDistortion, CoarseDropout\n   9 | Prototype       | rata-rata + Euclidean   | rectification BD-CSPN +\n     |                 |                         |   cosine + temperature\n  10 | Inferensi       | single pass             | TTA x8\n\nCATATAN PENTING TENTANG UKURAN INPUT (agar ablasi tetap jujur):\n  Dokumen metodologi menyebut EfficientNet-B4 pada 380x380, dan Eksperimen A\n  memakai itu. Eksperimen B secara default memakai 300x300 -- SAMA dengan\n  konfigurasi yang benar-benar kamu jalankan selama ini (hasil 0.836/0.688),\n  sehingga angkanya langsung nyambung dengan riwayat eksperimenmu.\n  Konsekuensinya: ukuran input menjadi variabel perancu (confound) dalam\n  ablasi ini.\n  -> Kalau ingin ablasi yang BENAR-BENAR bersih (hanya 10 komponen di atas\n     yang berbeda), set EFF_IMG_SIZE = 380 di bawah, lalu latih ulang.\n     Gunakan cache 380px yang sama dengan Eksperimen A.\n  -> Kalau memilih tetap 300, sebutkan hal ini apa adanya di bab pembahasan\n     sebagai keterbatasan -- jangan disembunyikan.\n\nCARA PAKAI DI KAGGLE:\n  Sama seperti Eksperimen A. Hasil -> hasil_improved.json\n============================================================================\n\"\"\"\n\nimport os\nimport gc\nimport glob\nimport json\nimport time\nimport random\nimport copy\nimport pickle\nimport warnings\nimport concurrent.futures\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR\nfrom torch.amp import autocast, GradScaler\nfrom torchvision import models\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (accuracy_score, confusion_matrix, roc_auc_score,\n                             f1_score, recall_score, precision_score,\n                             cohen_kappa_score)\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom tqdm import tqdm\n\nwarnings.filterwarnings('ignore')\n\n# ============================================================================\n# SEED & DEVICE -- IDENTIK dengan Eksperimen A\n# ============================================================================\n\nSEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\ntorch.cuda.manual_seed_all(SEED)\ntorch.backends.cudnn.deterministic = True\ntorch.backends.cudnn.benchmark = False\n\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nos.environ[\"PYTORCH_ALLOC_CONF\"] = \"expandable_segments:True\"\n\nTARGET_ACC, TARGET_SENS, TARGET_SPEC = 0.80, 0.85, 0.75\n\nCKPT_DIR = '/kaggle/working/improved_checkpoints'\nos.makedirs(CKPT_DIR, exist_ok=True)\n\n# Lihat \"CATATAN PENTING TENTANG UKURAN INPUT\" di header file.\nEFF_IMG_SIZE = 300      # ganti ke 380 untuk ablasi yang benar-benar bersih\n\n\n# ============================================================================\n# KONFIGURASI IMPROVED\n# ============================================================================\n\nCONFIG_RESNET = {\n    'BACKBONE': 'resnet50', 'IMG_SIZE': 224, 'PROJ_DIM': 256, 'N_WAY': 5,\n    'USE_AMP': True, 'WEIGHT_DECAY': 2e-4, 'GRAD_CLIP': 1.0,\n\n    # Stage 1 -- supervised penuh\n    'S1_EPOCHS': 150, 'S1_LR': 3e-4, 'S1_BATCH_SIZE': 32, 'S1_PATIENCE': 20,\n    'S1_LR_WARMUP': 3, 'S1_FOCAL_GAMMA': 2.0, 'S1_ORDINAL_SIGMA': 1.0,\n    'S1_SELECT_BY': 'sensitivity',\n\n    # Stage 2 -- FSL head (backbone frozen)\n    'S2_EPOCHS': 30, 'S2_LR': 5e-4, 'S2_K_SHOT': 5, 'S2_N_QUERY': 10,\n    'S2_EPISODES': 150, 'S2_EPISODES_VAL': 80, 'S2_PATIENCE': 10,\n    'S2_LAMBDA_ORD': 0.3, 'S2_TEMP_INIT': 8.0, 'S2_TEMP_MIN': 5.0, 'S2_TEMP_MAX': 15.0,\n\n    'USE_OVERSAMPLE': True, 'RESAMPLE_STRATEGY': 'mean',\n    'VAL_LOSS_EMA_ALPHA': 0.3, 'FORCE_RESTART': True,\n}\n\nCONFIG_EFF = dict(CONFIG_RESNET)\nCONFIG_EFF.update({\n    'BACKBONE': 'efficientnet_b4',\n    'IMG_SIZE': EFF_IMG_SIZE,\n    'S1_BATCH_SIZE': 16 if EFF_IMG_SIZE <= 300 else 8,\n    'S2_N_QUERY': 8, 'S2_EPISODES': 80, 'S2_EPISODES_VAL': 40,\n})\n\n\n# ============================================================================\n# PRAPROSES -- IDENTIK dengan Eksperimen A (bukan variabel yang diuji)\n# ============================================================================\n\nCLAHE_CLIP = 2.0\nCLAHE_TILE = (8, 8)\nIMAGENET_MEAN = np.array([0.485, 0.456, 0.406], dtype=np.float32)\nIMAGENET_STD = np.array([0.229, 0.224, 0.225], dtype=np.float32)\n\n\ndef circular_crop(img: np.ndarray, thr: int = 10) -> np.ndarray:\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    mask = gray > thr\n    if mask.sum() == 0:\n        return img\n    coords = np.argwhere(mask)\n    y0, x0 = coords.min(axis=0)\n    y1, x1 = coords.max(axis=0) + 1\n    cropped = img[y0:y1, x0:x1]\n    h, w = cropped.shape[:2]\n    if h < 10 or w < 10:\n        return img\n    size = min(h, w)\n    cy, cx = h // 2, w // 2\n    half = size // 2\n    sq = cropped[max(cy - half, 0):min(cy + half, h), max(cx - half, 0):min(cx + half, w)]\n    if sq.shape[0] < 10 or sq.shape[1] < 10:\n        return img\n    mask_c = np.zeros(sq.shape[:2], dtype=np.uint8)\n    cv2.circle(mask_c, (sq.shape[1] // 2, sq.shape[0] // 2), min(sq.shape[:2]) // 2, 255, -1)\n    return cv2.bitwise_and(sq, sq, mask=mask_c)\n\n\ndef apply_green_clahe(img: np.ndarray) -> np.ndarray:\n    green = img[:, :, 1]\n    clahe = cv2.createCLAHE(clipLimit=CLAHE_CLIP, tileGridSize=CLAHE_TILE)\n    return np.stack([clahe.apply(green)] * 3, axis=-1)\n\n\ndef preprocess_fundus(img_path: str, img_size: int) -> np.ndarray:\n    img = cv2.imread(img_path)\n    if img is None:\n        raise FileNotFoundError(f\"Tidak ditemukan: {img_path}\")\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    img = circular_crop(img)\n    img = apply_green_clahe(img)\n    return cv2.resize(img, (img_size, img_size), interpolation=cv2.INTER_AREA)\n\n\ndef _preprocess_one(args):\n    id_code, img_path, img_size = args\n    try:\n        return id_code, preprocess_fundus(img_path, img_size), None\n    except Exception as e:\n        return id_code, None, str(e)\n\n\ndef build_cache(df, img_dir, img_size, n_workers=8, cache_dir=None):\n    cache_file = os.path.join(cache_dir or CKPT_DIR, f'cache_imp_{img_size}px.pkl')\n    if os.path.exists(cache_file):\n        print(f\"  Cache {img_size}px ditemukan, memuat dari disk...\")\n        with open(cache_file, 'rb') as f:\n            return pickle.load(f)\n    print(f\"  Membangun cache {img_size}px ({n_workers} workers)...\")\n    args_list = [(r['id_code'], os.path.join(img_dir, f\"{r['id_code']}.png\"), img_size)\n                 for _, r in df.iterrows()]\n    cache, errors, t0 = {}, [], time.time()\n    with concurrent.futures.ThreadPoolExecutor(max_workers=n_workers) as ex:\n        futures = {ex.submit(_preprocess_one, a): a[0] for a in args_list}\n        for fut in tqdm(concurrent.futures.as_completed(futures),\n                        total=len(futures), desc=f\"Cache {img_size}px\"):\n            id_code, result, err = fut.result()\n            if result is not None:\n                cache[id_code] = result\n            else:\n                errors.append((id_code, err))\n    print(f\"  Selesai: {len(cache)} gambar ({time.time()-t0:.0f}s)\")\n    if errors:\n        print(f\"  {len(errors)} gambar gagal\")\n    with open(cache_file, 'wb') as f:\n        pickle.dump(cache, f, protocol=4)\n    return cache\n\n\n# ============================================================================\n# PERBAIKAN #8 -- AUGMENTASI AGRESIF\n# Ditulis manual dgn cv2/numpy (bukan albumentations) agar bebas dari\n# perubahan nama parameter antar-versi albumentations, misalnya\n# CoarseDropout(max_holes=...) yang sudah di-rename di versi >= 1.4.\n# ============================================================================\n\ndef _rand_resized_crop(img, rng, scale=(0.75, 1.0), ratio=(0.9, 1.1)):\n    h, w = img.shape[:2]\n    area = h * w\n    for _ in range(10):\n        ta = area * float(rng.uniform(*scale))\n        ar = float(rng.uniform(*ratio))\n        nw, nh = int(round(np.sqrt(ta * ar))), int(round(np.sqrt(ta / ar)))\n        if nw <= w and nh <= h:\n            x0 = int(rng.integers(0, w - nw + 1))\n            y0 = int(rng.integers(0, h - nh + 1))\n            crop = img[y0:y0 + nh, x0:x0 + nw]\n            return cv2.resize(crop, (w, h), interpolation=cv2.INTER_LINEAR)\n    return img\n\n\ndef _color_jitter(img, rng, brightness=0.2, contrast=0.2, saturation=0.2, hue=0.1):\n    out = img.astype(np.float32)\n    out = out * (1.0 + float(rng.uniform(-contrast, contrast)))\n    out = out + 255.0 * float(rng.uniform(-brightness, brightness)) * 0.5\n    out = np.clip(out, 0, 255).astype(np.uint8)\n    hsv = cv2.cvtColor(out, cv2.COLOR_RGB2HSV).astype(np.float32)\n    hsv[:, :, 1] *= (1.0 + float(rng.uniform(-saturation, saturation)))\n    hsv[:, :, 0] = (hsv[:, :, 0] + 180.0 * float(rng.uniform(-hue, hue))) % 180.0\n    hsv[:, :, 1] = np.clip(hsv[:, :, 1], 0, 255)\n    return cv2.cvtColor(hsv.astype(np.uint8), cv2.COLOR_HSV2RGB)\n\n\ndef _grid_distortion(img, rng, num_steps=5, distort_limit=0.2):\n    h, w = img.shape[:2]\n    xs = np.linspace(0, w, num_steps + 1)\n    ys = np.linspace(0, h, num_steps + 1)\n    xs_d = xs + rng.uniform(-distort_limit, distort_limit, xs.shape) * (w / num_steps)\n    ys_d = ys + rng.uniform(-distort_limit, distort_limit, ys.shape) * (h / num_steps)\n    xs_d[0], xs_d[-1] = 0, w\n    ys_d[0], ys_d[-1] = 0, h\n    map_x = np.interp(np.arange(w), xs_d, xs).astype(np.float32)\n    map_y = np.interp(np.arange(h), ys_d, ys).astype(np.float32)\n    map_x = np.tile(map_x, (h, 1))\n    map_y = np.tile(map_y.reshape(-1, 1), (1, w))\n    return cv2.remap(img, map_x, map_y, interpolation=cv2.INTER_LINEAR,\n                     borderMode=cv2.BORDER_CONSTANT, borderValue=(0, 0, 0))\n\n\ndef _coarse_dropout(img, rng, max_holes=8, max_h=16, max_w=16):\n    out = img.copy()\n    h, w = out.shape[:2]\n    for _ in range(int(rng.integers(1, max_holes + 1))):\n        hh = int(rng.integers(4, max_h + 1))\n        ww = int(rng.integers(4, max_w + 1))\n        y0 = int(rng.integers(0, max(h - hh, 1)))\n        x0 = int(rng.integers(0, max(w - ww, 1)))\n        out[y0:y0 + hh, x0:x0 + ww] = 0\n    return out\n\n\ndef _rotate(img, rng, limit=360):\n    h, w = img.shape[:2]\n    angle = float(rng.uniform(-limit, limit))\n    M = cv2.getRotationMatrix2D((w / 2, h / 2), angle, 1.0)\n    return cv2.warpAffine(img, M, (w, h), flags=cv2.INTER_LINEAR,\n                          borderMode=cv2.BORDER_CONSTANT, borderValue=(0, 0, 0))\n\n\ndef _brightness_contrast(img, rng, b_limit=0.25, c_limit=0.25):\n    alpha = 1.0 + float(rng.uniform(-c_limit, c_limit))\n    beta = 255.0 * float(rng.uniform(-b_limit, b_limit)) * 0.5\n    return np.clip(img.astype(np.float32) * alpha + beta, 0, 255).astype(np.uint8)\n\n\ndef augment_s1(img: np.ndarray, rng: np.random.Generator) -> np.ndarray:\n    \"\"\"Augmentasi agresif untuk Stage 1 (backbone masih belajar).\"\"\"\n    out = img\n    if rng.random() < 0.8:\n        out = _rand_resized_crop(out, rng, scale=(0.75, 1.0), ratio=(0.9, 1.1))\n    if rng.random() < 0.5:\n        out = cv2.flip(out, 1)\n    if rng.random() < 0.5:\n        out = cv2.flip(out, 0)\n    if rng.random() < 0.8:\n        out = _rotate(out, rng, limit=360)\n    if rng.random() < 0.6:\n        out = _brightness_contrast(out, rng, 0.25, 0.25)\n    if rng.random() < 0.5:\n        out = _color_jitter(out, rng)\n    if rng.random() < 0.3:\n        k = int(rng.choice([3, 5, 7]))\n        out = cv2.GaussianBlur(out, (k, k), 0)\n    if rng.random() < 0.2:\n        out = _grid_distortion(out, rng)\n    if rng.random() < 0.2:\n        out = _coarse_dropout(out, rng)\n    return out\n\n\ndef augment_s2(img: np.ndarray, rng: np.random.Generator) -> np.ndarray:\n    \"\"\"Augmentasi sedang untuk Stage 2 (backbone sudah frozen).\"\"\"\n    out = img\n    if rng.random() < 0.7:\n        out = _rand_resized_crop(out, rng, scale=(0.8, 1.0), ratio=(0.9, 1.1))\n    if rng.random() < 0.5:\n        out = cv2.flip(out, 1)\n    if rng.random() < 0.5:\n        out = cv2.flip(out, 0)\n    if rng.random() < 0.5:\n        out = _rotate(out, rng, limit=30)\n    if rng.random() < 0.4:\n        out = _brightness_contrast(out, rng, 0.15, 0.15)\n    return out\n\n\ndef normalize_to_tensor(img: np.ndarray) -> torch.Tensor:\n    f = img.astype(np.float32) / 255.0\n    f = (f - IMAGENET_MEAN) / IMAGENET_STD\n    return torch.from_numpy(np.transpose(f, (2, 0, 1))).float()\n\n\n# ============================================================================\n# PERBAIKAN #1 -- HYBRID UNDER/OVERSAMPLING\n# ============================================================================\n\ndef compute_resample_target(train_df, label_col='diagnosis', strategy='mean'):\n    \"\"\"'max' = oversampling murni (perilaku lama); 'mean'/'median' = gabungan\n    undersampling kelas mayoritas + oversampling kelas minoritas.\"\"\"\n    counts = train_df[label_col].value_counts()\n    if isinstance(strategy, int):\n        target = strategy\n    elif strategy == 'max':\n        target = int(counts.max())\n    elif strategy == 'mean':\n        target = int(counts.mean())\n    elif strategy == 'median':\n        target = int(counts.median())\n    else:\n        raise ValueError(f\"strategy tidak dikenal: {strategy}\")\n    print(f\"  Resample target ({strategy}) = {target}\")\n    for c in sorted(counts.index):\n        n = counts[c]\n        aksi = \"OVERSAMPLE\" if n < target else (\"UNDERSAMPLE\" if n > target else \"tetap\")\n        print(f\"    Grade {c}: {n:5d} -> {target:5d}  ({aksi}, x{target/n:.2f})\")\n    return target\n\n\nclass APTOSDataset(Dataset):\n    def __init__(self, df, cache, aug_fn=None):\n        self.df = df.reset_index(drop=True)\n        self.cache = cache\n        self.aug_fn = aug_fn\n        self.rng = np.random.default_rng(SEED)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = self.cache[row['id_code']]\n        if self.aug_fn is not None:\n            img = self.aug_fn(img, self.rng)\n        return normalize_to_tensor(img), int(row['diagnosis'])\n\n\nclass ResampledDataset(Dataset):\n    \"\"\"Setiap repetisi mendapat augmentasi berbeda -> duplikat tidak identik.\"\"\"\n\n    def __init__(self, df, cache, aug_fn, target_per_class):\n        self.cache = cache\n        self.aug_fn = aug_fn\n        self.rng = np.random.default_rng(SEED)\n        df = df.reset_index(drop=True)\n        labels = df['diagnosis'].values\n        self.samples = []\n        for c in np.unique(labels):\n            idx = np.where(labels == c)[0]\n            n = len(idx)\n            chosen = np.random.choice(idx, target_per_class, replace=(n < target_per_class))\n            for i in chosen:\n                self.samples.append((df.iloc[i]['id_code'], int(c)))\n        random.shuffle(self.samples)\n        print(f\"  Total sampel per epoch: {len(self.samples)}\")\n\n    def __len__(self):\n        return len(self.samples)\n\n    def __getitem__(self, idx):\n        id_code, label = self.samples[idx]\n        img = self.cache[id_code]\n        if self.aug_fn is not None:\n            img = self.aug_fn(img, self.rng)\n        return normalize_to_tensor(img), label\n\n\nclass EpisodicSampler:\n    def __init__(self, labels, n_way, k_shot, n_query, n_episodes):\n        self.n_way, self.k_shot = n_way, k_shot\n        self.n_query, self.n_episodes = n_query, n_episodes\n        labels = np.array(labels)\n        self.classes = np.unique(labels)\n        self.class_idx = {c: np.where(labels == c)[0] for c in self.classes}\n\n    def __len__(self):\n        return self.n_episodes\n\n    def __iter__(self):\n        for _ in range(self.n_episodes):\n            batch, n = [], self.k_shot + self.n_query\n            for c in self.classes:\n                idx = self.class_idx[c]\n                batch.extend(np.random.choice(idx, n, replace=len(idx) < n).tolist())\n            yield batch\n\n\n# ============================================================================\n# PERBAIKAN #3 -- CBAM (Woo et al., 2018, ECCV)\n# ============================================================================\n\nclass ChannelAttention(nn.Module):\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n        self.avgpool = nn.AdaptiveAvgPool2d(1)\n        self.maxpool = nn.AdaptiveMaxPool2d(1)\n        self.fc = nn.Sequential(\n            nn.Linear(channels, channels // reduction, bias=False), nn.ReLU(inplace=True),\n            nn.Linear(channels // reduction, channels, bias=False))\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        b, c, _, _ = x.shape\n        avg = self.fc(self.avgpool(x).view(b, c))\n        mx = self.fc(self.maxpool(x).view(b, c))\n        return x * self.sigmoid(avg + mx).view(b, c, 1, 1)\n\n\nclass SpatialAttention(nn.Module):\n    def __init__(self, kernel_size=7):\n        super().__init__()\n        self.conv = nn.Conv2d(2, 1, kernel_size=kernel_size,\n                              padding=kernel_size // 2, bias=False)\n        self.sigmoid = nn.Sigmoid()\n\n    def forward(self, x):\n        avg = torch.mean(x, dim=1, keepdim=True)\n        mx, _ = torch.max(x, dim=1, keepdim=True)\n        return x * self.sigmoid(self.conv(torch.cat([avg, mx], dim=1)))\n\n\nclass CBAM(nn.Module):\n    def __init__(self, channels, reduction=16, spatial_kernel=7):\n        super().__init__()\n        self.ch = ChannelAttention(channels, reduction)\n        self.sp = SpatialAttention(spatial_kernel)\n\n    def forward(self, x):\n        return self.sp(self.ch(x))\n\n\n# ============================================================================\n# PERBAIKAN #4 -- MODEL TWO-STAGE\n# ============================================================================\n\nclass TwoStageModel(nn.Module):\n    def __init__(self, backbone_name: str, proj_dim: int = 256, n_classes: int = 5):\n        super().__init__()\n        self.backbone_name = backbone_name\n        if backbone_name == 'resnet50':\n            base = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2)\n            self.feature_dim = base.fc.in_features\n            base.fc = nn.Identity()\n            self.encoder = base\n        elif backbone_name == 'efficientnet_b4':\n            base = models.efficientnet_b4(weights=models.EfficientNet_B4_Weights.IMAGENET1K_V1)\n            self.feature_dim = base.classifier[1].in_features\n            base.classifier = nn.Identity()\n            self.encoder = base\n        else:\n            raise ValueError(f\"Backbone tidak dikenal: {backbone_name}\")\n\n        self.cbam = CBAM(self.feature_dim, reduction=16, spatial_kernel=7)\n        self.classifier = nn.Sequential(\n            nn.Linear(self.feature_dim, 512), nn.BatchNorm1d(512), nn.ReLU(inplace=True),\n            nn.Dropout(0.4), nn.Linear(512, n_classes))\n        self.projection = nn.Sequential(\n            nn.Linear(self.feature_dim, 512), nn.BatchNorm1d(512), nn.ReLU(inplace=True),\n            nn.Dropout(0.3), nn.Linear(512, proj_dim))\n        self.temperature = nn.Parameter(torch.tensor(8.0))\n\n    def get_last_conv_features(self, x):\n        if self.backbone_name == 'resnet50':\n            x = self.encoder.conv1(x); x = self.encoder.bn1(x)\n            x = self.encoder.relu(x); x = self.encoder.maxpool(x)\n            x = self.encoder.layer1(x); x = self.encoder.layer2(x)\n            x = self.encoder.layer3(x); x = self.encoder.layer4(x)\n        else:\n            x = self.encoder.features(x)\n        return x\n\n    def _extract_features(self, x):\n        x = self.get_last_conv_features(x)\n        x = self.cbam(x)\n        x = self.encoder.avgpool(x)\n        return torch.flatten(x, 1)\n\n    def forward_supervised(self, x):\n        return self.classifier(self._extract_features(x))\n\n    def forward_embedding(self, x):\n        return F.normalize(self.projection(self._extract_features(x)), dim=1)\n\n    def freeze_backbone(self):\n        for p in self.encoder.parameters():\n            p.requires_grad = False\n        for p in self.cbam.parameters():\n            p.requires_grad = False\n\n    def unfreeze_all_for_s1(self):\n        for p in self.parameters():\n            p.requires_grad = True\n        for p in self.projection.parameters():\n            p.requires_grad = False\n\n    def unfreeze_projection_only(self):\n        self.freeze_backbone()\n        for p in self.classifier.parameters():\n            p.requires_grad = False\n        for p in self.projection.parameters():\n            p.requires_grad = True\n        self.temperature.requires_grad = True\n\n    def count_trainable(self):\n        n = sum(p.numel() for p in self.parameters() if p.requires_grad)\n        print(f\"  Trainable params: {n:,}\")\n        return n\n\n\n# ============================================================================\n# PERBAIKAN #2 -- ORDINAL FOCAL LOSS\n# ============================================================================\n\nclass OrdinalFocalLoss(nn.Module):\n    \"\"\"Focal Loss + target lunak berbasis jarak ordinal (pengganti label\n    smoothing seragam). Dihitung MURNI per-sampel -- tidak bergantung\n    statistik antar-batch, jadi stabil di batch kecil sekalipun.\"\"\"\n\n    def __init__(self, alpha=None, gamma=2.0, n_classes=5, sigma=1.0, soft_weight=1.0):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.n_classes = n_classes\n        self.sigma = sigma\n        self.soft_weight = soft_weight\n        self.register_buffer('classes', torch.arange(n_classes, dtype=torch.float32))\n\n    def forward(self, logits, targets):\n        y = targets.float().unsqueeze(1)\n        dist2 = (self.classes.unsqueeze(0) - y) ** 2\n        soft = torch.softmax(-dist2 / (2 * self.sigma ** 2), dim=1)\n        onehot = F.one_hot(targets, self.n_classes).float()\n        target_dist = self.soft_weight * soft + (1 - self.soft_weight) * onehot\n        logp = F.log_softmax(logits, dim=1)\n        ce = -(target_dist * logp).sum(dim=1)\n        pt = logp.exp().gather(1, targets.unsqueeze(1)).squeeze(1)\n        w = (1 - pt) ** self.gamma\n        if self.alpha is not None:\n            w = w * self.alpha[targets]\n        return (w * ce).mean()\n\n\n# ============================================================================\n# PERBAIKAN #9 -- PROTOTYPICAL LOSS + RECTIFICATION (BD-CSPN)\n# ============================================================================\n\ndef prototypical_loss_improved(support_emb, query_emb, n_way, k_shot, n_query,\n                                class_weights, lambda_ordinal=0.3,\n                                temperature=10.0, top_z=3):\n    query_labels = torch.arange(n_way, device=query_emb.device).repeat_interleave(n_query)\n    support_v = support_emb.view(n_way, k_shot, -1)\n    proto_init = support_v.mean(dim=1)\n\n    q_norm = F.normalize(query_emb, dim=1)\n    p_norm = F.normalize(proto_init, dim=1)\n    probs = F.softmax(torch.matmul(q_norm, p_norm.t()) * temperature, dim=1)\n    pseudo = probs.argmax(dim=1)\n    conf = probs.max(dim=1).values\n\n    prototypes = []\n    for c in range(n_way):\n        sc = support_v[c]\n        mc = (pseudo == c)\n        if mc.sum() > 0:\n            qc = query_emb[mc]\n            cc = conf[mc]\n            top = torch.topk(cc, min(top_z, qc.size(0))).indices\n            combined = torch.cat([sc, qc[top]], dim=0)\n        else:\n            combined = sc\n        cn = F.normalize(combined, dim=1)\n        w = F.softmax(torch.matmul(cn, p_norm[c].unsqueeze(1)).squeeze(1) * temperature, dim=0)\n        prototypes.append((w.unsqueeze(1) * combined).sum(dim=0))\n    prototypes = torch.stack(prototypes, dim=0)\n\n    logits = torch.matmul(q_norm, F.normalize(prototypes, dim=1).t()) * temperature\n    log_p_y = F.log_softmax(logits, dim=1)\n    w = class_weights[query_labels]\n    loss = (F.nll_loss(log_p_y, query_labels, reduction='none') * w).sum() / w.sum()\n\n    if lambda_ordinal > 0:\n        probs_q = log_p_y.exp()\n        cv = torch.arange(n_way, device=query_emb.device).float()\n        loss = loss + lambda_ordinal * F.mse_loss((probs_q * cv).sum(1), query_labels.float())\n\n    preds = log_p_y.argmax(dim=1)\n    return loss, (preds == query_labels).float().mean().item(), preds, query_labels\n\n\n# ============================================================================\n# PERBAIKAN #5 & #6 -- WARMUP-COSINE + EMA\n# ============================================================================\n\ndef build_warmup_cosine(optimizer, warmup_ep, total_ep, eta_min=1e-6):\n    warmup = LinearLR(optimizer, start_factor=0.1, end_factor=1.0, total_iters=warmup_ep)\n    cosine = CosineAnnealingLR(optimizer, T_max=max(total_ep - warmup_ep, 1), eta_min=eta_min)\n    return SequentialLR(optimizer, schedulers=[warmup, cosine], milestones=[warmup_ep])\n\n\nclass EMA:\n    def __init__(self, model: nn.Module, decay: float = 0.999):\n        self.model = model\n        self.decay = decay\n        self.shadow = {}\n        self._backup = {}\n        for name, p in model.named_parameters():\n            if p.requires_grad:\n                self.shadow[name] = p.data.clone()\n\n    def update(self):\n        for name, p in self.model.named_parameters():\n            if p.requires_grad and name in self.shadow:\n                self.shadow[name] = self.decay * self.shadow[name] + (1 - self.decay) * p.data\n\n    def apply_shadow(self):\n        for name, p in self.model.named_parameters():\n            if name in self.shadow:\n                self._backup[name] = p.data.clone()\n                p.data = self.shadow[name]\n\n    def restore(self):\n        for name, p in self.model.named_parameters():\n            if name in self._backup:\n                p.data = self._backup[name]\n        self._backup.clear()\n\n    def state_dict(self):\n        return {'shadow': {k: v.clone() for k, v in self.shadow.items()}, 'decay': self.decay}\n\n    def load_state_dict(self, sd):\n        self.shadow = {k: v.clone() for k, v in sd['shadow'].items()}\n        self.decay = sd.get('decay', self.decay)\n\n\ndef maybe_backup_checkpoint(ckpt_path, force_restart):\n    if force_restart and os.path.exists(ckpt_path):\n        backup = ckpt_path + f'.bak_{int(time.time())}'\n        os.rename(ckpt_path, backup)\n        print(f\"  FORCE_RESTART aktif -- checkpoint lama dipindah ke: {backup}\")\n\n\n# ============================================================================\n# METRIK -- IDENTIK PERSIS dengan Eksperimen A\n# ============================================================================\n\ndef compute_metrics(y_true, y_pred, n_way=5):\n    acc = accuracy_score(y_true, y_pred)\n    cm = confusion_matrix(y_true, y_pred, labels=list(range(n_way)))\n    try:\n        auc = roc_auc_score(y_true, np.eye(n_way)[y_pred],\n                            multi_class='ovr', average='weighted')\n    except Exception:\n        auc = 0.0\n    sens_l, spec_l = [], []\n    for c in range(n_way):\n        TP = cm[c, c]\n        FN = cm[c, :].sum() - TP\n        FP = cm[:, c].sum() - TP\n        TN = cm.sum() - TP - FN - FP\n        sens_l.append(TP / (TP + FN) if (TP + FN) > 0 else 0.0)\n        spec_l.append(TN / (TN + FP) if (TN + FP) > 0 else 0.0)\n    return {\n        'accuracy': acc,\n        'weighted_f1': f1_score(y_true, y_pred, average='weighted', zero_division=0),\n        'weighted_precision': precision_score(y_true, y_pred, average='weighted', zero_division=0),\n        # CATATAN: recall 'weighted' secara matematis IDENTIK dengan accuracy.\n        # Dilaporkan karena dokumen memintanya, tapi sensitivity_macro-lah\n        # yang benar-benar mencerminkan performa kelas minoritas.\n        'sensitivity_weighted': recall_score(y_true, y_pred, average='weighted', zero_division=0),\n        'sensitivity_macro': float(np.mean(sens_l)),\n        'specificity_macro': float(np.mean(spec_l)),\n        'sensitivity_per_class': [float(s) for s in sens_l],\n        'specificity_per_class': [float(s) for s in spec_l],\n        'auc': auc,\n        'qwk': cohen_kappa_score(y_true, y_pred, weights='quadratic'),\n        'confusion_matrix': cm.tolist(),\n    }\n\n\n# ============================================================================\n# STAGE 1 -- Supervised penuh\n# ============================================================================\n\ndef train_stage1(model, train_df, val_df, cache, cfg):\n    backbone = cfg['BACKBONE']\n    ckpt_path = os.path.join(CKPT_DIR, f's1_{backbone}.pt')\n    maybe_backup_checkpoint(ckpt_path, cfg.get('FORCE_RESTART', False))\n\n    print(f\"\\n{'='*65}\\n  STAGE 1 -- Supervised: {backbone}\")\n    print(f\"  Epochs: {cfg['S1_EPOCHS']} | Batch: {cfg['S1_BATCH_SIZE']} | LR: {cfg['S1_LR']}\")\n    print(f\"{'='*65}\")\n\n    model.unfreeze_all_for_s1()\n    model.count_trainable()\n\n    target_cls = compute_resample_target(train_df, strategy=cfg.get('RESAMPLE_STRATEGY', 'mean'))\n    train_ds = ResampledDataset(train_df, cache, augment_s1, target_cls)\n    val_ds = APTOSDataset(val_df, cache, aug_fn=None)\n    train_loader = DataLoader(train_ds, batch_size=cfg['S1_BATCH_SIZE'], shuffle=True,\n                              num_workers=2, pin_memory=True, drop_last=True)\n    val_loader = DataLoader(val_ds, batch_size=cfg['S1_BATCH_SIZE'] * 2, shuffle=False,\n                            num_workers=2, pin_memory=True)\n\n    cw = compute_class_weight('balanced', classes=np.unique(train_df['diagnosis']),\n                              y=train_df['diagnosis'])\n    cw_t = torch.tensor(cw, dtype=torch.float32).to(DEVICE)\n    criterion = OrdinalFocalLoss(alpha=cw_t, gamma=cfg['S1_FOCAL_GAMMA'],\n                                 sigma=cfg.get('S1_ORDINAL_SIGMA', 1.0)).to(DEVICE)\n\n    optimizer = torch.optim.AdamW(filter(lambda p: p.requires_grad, model.parameters()),\n                                  lr=cfg['S1_LR'], weight_decay=cfg['WEIGHT_DECAY'])\n    scheduler = build_warmup_cosine(optimizer, cfg['S1_LR_WARMUP'], cfg['S1_EPOCHS'])\n    scaler = GradScaler('cuda') if cfg['USE_AMP'] else None\n    ema = EMA(model, decay=0.999)\n\n    history = {'train_loss': [], 'val_loss': [], 'val_acc': [], 'val_sens': []}\n    best_val_acc, best_state, patience = 0.0, None, 0\n    best_val_sens, best_state_sens = 0.0, None\n    best_state_ema, best_state_sens_ema = None, None\n\n    for ep in range(cfg['S1_EPOCHS']):\n        if ep > 0 and ep % 10 == 0:\n            train_ds = ResampledDataset(train_df, cache, augment_s1, target_cls)\n            train_loader = DataLoader(train_ds, batch_size=cfg['S1_BATCH_SIZE'], shuffle=True,\n                                      num_workers=2, pin_memory=True, drop_last=True)\n\n        model.train()\n        tot_loss, correct, total = 0.0, 0, 0\n        for imgs, labels in tqdm(train_loader, desc=f\"  S1 Ep{ep+1}/{cfg['S1_EPOCHS']}\",\n                                 leave=False):\n            imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n            optimizer.zero_grad()\n            with autocast('cuda', enabled=scaler is not None):\n                logits = model.forward_supervised(imgs)\n                loss = criterion(logits, labels)\n            if scaler:\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), cfg['GRAD_CLIP'])\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), cfg['GRAD_CLIP'])\n                optimizer.step()\n            ema.update()\n            tot_loss += loss.item() * labels.size(0)\n            correct += (logits.argmax(1) == labels).sum().item()\n            total += labels.size(0)\n        train_loss, train_acc = tot_loss / total, correct / total\n        scheduler.step()\n\n        # ---- Validasi dgn bobot EMA ----\n        ema.apply_shadow()\n        model.eval()\n        vl, vc, vt = 0.0, 0, 0\n        y_true, y_pred = [], []\n        with torch.no_grad():\n            for imgs, labels in val_loader:\n                imgs, labels = imgs.to(DEVICE), labels.to(DEVICE)\n                logits = model.forward_supervised(imgs)\n                vl += criterion(logits, labels).item() * labels.size(0)\n                preds = logits.argmax(1)\n                vc += (preds == labels).sum().item()\n                vt += labels.size(0)\n                y_true.extend(labels.cpu().numpy())\n                y_pred.extend(preds.cpu().numpy())\n        # model MASIH memegang bobot EMA di titik ini -- jangan restore dulu\n\n        val_loss, val_acc = vl / vt, vc / vt\n        val_sens = recall_score(y_true, y_pred, average='macro', zero_division=0)\n\n        history['train_loss'].append(train_loss)\n        history['val_loss'].append(val_loss)\n        history['val_acc'].append(val_acc)\n        history['val_sens'].append(val_sens)\n\n        print(f\"  S1 Ep {ep+1:3d}/{cfg['S1_EPOCHS']} | train: loss={train_loss:.4f} \"\n              f\"acc={train_acc:.4f} | val: loss={val_loss:.4f} acc={val_acc:.4f} \"\n              f\"sens={val_sens:.4f} | lr={optimizer.param_groups[0]['lr']:.2e}\")\n\n        # PERBAIKAN #6: tangkap checkpoint SELAGI bobot EMA masih aktif\n        new_best_acc = val_acc > best_val_acc\n        new_best_sens = val_sens > best_val_sens\n        if new_best_acc:\n            best_state_ema = copy.deepcopy(model.state_dict())\n        if new_best_sens:\n            best_state_sens_ema = copy.deepcopy(model.state_dict())\n\n        ema.restore()\n\n        if new_best_acc:\n            best_val_acc = val_acc\n            best_state = copy.deepcopy(model.state_dict())\n            patience = 0\n            print(f\"    -> Best by ACCURACY (val_acc={val_acc:.4f})\")\n        else:\n            patience += 1\n        if new_best_sens:\n            best_val_sens = val_sens\n            best_state_sens = copy.deepcopy(model.state_dict())\n            print(f\"    -> Best by SENSITIVITY (val_sens={val_sens:.4f})\")\n\n        torch.save({\n            'model_state': model.state_dict(), 'best_state': best_state,\n            'best_state_sens': best_state_sens, 'best_state_ema': best_state_ema,\n            'best_state_sens_ema': best_state_sens_ema,\n            'best_val_acc': best_val_acc, 'best_val_sens': best_val_sens,\n            'history': history, 'epoch': ep + 1, 'patience': patience,\n            'optimizer_state': optimizer.state_dict(),\n            'scheduler_state': scheduler.state_dict(),\n            'scaler_state': scaler.state_dict() if scaler else None,\n            'ema_state': ema.state_dict(),\n        }, ckpt_path)\n\n        if patience >= cfg['S1_PATIENCE']:\n            print(f\"  Early stop Stage 1 pada epoch {ep+1}\")\n            break\n        torch.cuda.empty_cache()\n\n    # PERBAIKAN #7: muat checkpoint berbasis SENSITIVITY, versi EMA\n    if cfg.get('S1_SELECT_BY') == 'sensitivity' and best_state_sens_ema is not None:\n        model.load_state_dict(best_state_sens_ema)\n        print(f\"\\n  Stage 1 selesai. Checkpoint BY SENSITIVITY versi EMA \"\n              f\"(val_sens={best_val_sens:.4f})\")\n    elif best_state_ema is not None:\n        model.load_state_dict(best_state_ema)\n        print(f\"\\n  Stage 1 selesai. Checkpoint BY ACCURACY versi EMA \"\n              f\"(val_acc={best_val_acc:.4f})\")\n    else:\n        model.load_state_dict(best_state)\n        print(f\"\\n  Stage 1 selesai (fallback versi RAW).\")\n    return model, history\n\n\n# ============================================================================\n# STAGE 2 -- FSL projection head, backbone SELALU frozen\n# ============================================================================\n\ndef train_stage2(model, train_df, val_df, cache, cfg, class_weights_tensor):\n    backbone = cfg['BACKBONE']\n    ckpt_path = os.path.join(CKPT_DIR, f's2_{backbone}.pt')\n    maybe_backup_checkpoint(ckpt_path, cfg.get('FORCE_RESTART', False))\n\n    print(f\"\\n{'='*65}\\n  STAGE 2 -- FSL Head: {backbone}\")\n    print(f\"  Backbone FROZEN | K_SHOT={cfg['S2_K_SHOT']} N_QUERY={cfg['S2_N_QUERY']}\")\n    print(f\"{'='*65}\")\n\n    model.unfreeze_projection_only()\n    model.count_trainable()\n    with torch.no_grad():\n        model.temperature.fill_(cfg['S2_TEMP_INIT'])\n\n    target_cls = compute_resample_target(train_df, strategy=cfg.get('RESAMPLE_STRATEGY', 'mean'))\n    train_ds = ResampledDataset(train_df, cache, augment_s2, target_cls)\n    t_labels = [s[1] for s in train_ds.samples]\n    t_sampler = EpisodicSampler(t_labels, cfg['N_WAY'], cfg['S2_K_SHOT'],\n                                cfg['S2_N_QUERY'], cfg['S2_EPISODES'])\n    val_ds = APTOSDataset(val_df, cache, aug_fn=None)\n    v_sampler = EpisodicSampler(val_df['diagnosis'].values, cfg['N_WAY'], cfg['S2_K_SHOT'],\n                                cfg['S2_N_QUERY'], cfg['S2_EPISODES_VAL'])\n\n    optimizer = torch.optim.AdamW(filter(lambda p: p.requires_grad, model.parameters()),\n                                  lr=cfg['S2_LR'], weight_decay=cfg['WEIGHT_DECAY'])\n    scheduler = build_warmup_cosine(optimizer, 3, cfg['S2_EPOCHS'], eta_min=1e-6)\n    scaler = GradScaler('cuda') if cfg['USE_AMP'] else None\n    ema = EMA(model, decay=0.999)\n\n    history = {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': [],\n               'val_auc': [], 'val_sens': []}\n    best_sens, best_state, best_state_ema, patience = 0.0, None, None, 0\n    n_way, k, q = cfg['N_WAY'], cfg['S2_K_SHOT'], cfg['S2_N_QUERY']\n\n    for ep in range(cfg['S2_EPOCHS']):\n        if ep > 0 and ep % 8 == 0:\n            train_ds = ResampledDataset(train_df, cache, augment_s2, target_cls)\n            t_labels = [s[1] for s in train_ds.samples]\n            t_sampler = EpisodicSampler(t_labels, cfg['N_WAY'], cfg['S2_K_SHOT'],\n                                        cfg['S2_N_QUERY'], cfg['S2_EPISODES'])\n\n        model.train()\n        tl, ta, nb = 0.0, 0.0, 0\n        t_loader = DataLoader(train_ds, batch_sampler=t_sampler, num_workers=2, pin_memory=True)\n        for imgs, _ in tqdm(t_loader, desc=f\"  S2 Ep{ep+1}/{cfg['S2_EPOCHS']}\", leave=False):\n            imgs = imgs.to(DEVICE).view(n_way, k + q, *imgs.shape[1:])\n            support = imgs[:, :k].reshape(-1, *imgs.shape[2:])\n            query = imgs[:, k:].reshape(-1, *imgs.shape[2:])\n            with autocast('cuda', enabled=scaler is not None):\n                s_emb = model.forward_embedding(support)\n                q_emb = model.forward_embedding(query)\n                temp = model.temperature.clamp(cfg['S2_TEMP_MIN'], cfg['S2_TEMP_MAX'])\n                loss, acc, _, _ = prototypical_loss_improved(\n                    s_emb, q_emb, n_way, k, q, class_weights_tensor,\n                    lambda_ordinal=cfg['S2_LAMBDA_ORD'], temperature=temp)\n            optimizer.zero_grad()\n            if scaler:\n                scaler.scale(loss).backward()\n                scaler.unscale_(optimizer)\n                torch.nn.utils.clip_grad_norm_(model.parameters(), cfg['GRAD_CLIP'])\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                loss.backward()\n                torch.nn.utils.clip_grad_norm_(model.parameters(), cfg['GRAD_CLIP'])\n                optimizer.step()\n            ema.update()\n            tl += loss.item(); ta += acc; nb += 1\n        tr_loss, tr_acc = tl / nb, ta / nb\n        scheduler.step()\n\n        ema.apply_shadow()\n        model.eval()\n        vl, va, nvb = 0.0, 0.0, 0\n        y_true, y_pred = [], []\n        v_loader = DataLoader(val_ds, batch_sampler=v_sampler, num_workers=2, pin_memory=True)\n        with torch.no_grad():\n            for imgs, _ in v_loader:\n                imgs = imgs.to(DEVICE).view(n_way, k + q, *imgs.shape[1:])\n                support = imgs[:, :k].reshape(-1, *imgs.shape[2:])\n                query = imgs[:, k:].reshape(-1, *imgs.shape[2:])\n                s_emb = model.forward_embedding(support)\n                q_emb = model.forward_embedding(query)\n                temp = model.temperature.clamp(cfg['S2_TEMP_MIN'], cfg['S2_TEMP_MAX'])\n                loss, acc, preds, lbls = prototypical_loss_improved(\n                    s_emb, q_emb, n_way, k, q, class_weights_tensor,\n                    lambda_ordinal=cfg['S2_LAMBDA_ORD'], temperature=temp)\n                vl += loss.item(); va += acc; nvb += 1\n                y_true.extend(lbls.cpu().numpy())\n                y_pred.extend(preds.detach().cpu().numpy())\n\n        va_loss, va_acc = vl / nvb, va / nvb\n        try:\n            va_auc = roc_auc_score(y_true, np.eye(n_way)[y_pred],\n                                   multi_class='ovr', average='weighted')\n        except Exception:\n            va_auc = 0.0\n        sens = recall_score(y_true, y_pred, average='macro', zero_division=0)\n\n        history['train_loss'].append(tr_loss); history['train_acc'].append(tr_acc)\n        history['val_loss'].append(va_loss); history['val_acc'].append(va_acc)\n        history['val_auc'].append(va_auc); history['val_sens'].append(sens)\n\n        print(f\"  S2 Ep {ep+1:2d}/{cfg['S2_EPOCHS']} | T={tr_loss:.4f}/{tr_acc:.4f} | \"\n              f\"V={va_loss:.4f} AUC={va_auc:.4f} Sens={sens:.4f} | \"\n              f\"lr={optimizer.param_groups[0]['lr']:.2e}\")\n\n        new_best = sens > best_sens\n        if new_best:\n            best_state_ema = copy.deepcopy(model.state_dict())\n        ema.restore()\n        if new_best:\n            best_sens = sens\n            best_state = copy.deepcopy(model.state_dict())\n            patience = 0\n            print(f\"    -> Best Stage 2 (val_sens={sens:.4f})\")\n        else:\n            patience += 1\n\n        torch.save({\n            'model_state': model.state_dict(), 'best_state': best_state,\n            'best_state_ema': best_state_ema, 'best_sens': best_sens,\n            'history': history, 'epoch': ep + 1, 'patience': patience,\n            'optimizer_state': optimizer.state_dict(),\n            'scheduler_state': scheduler.state_dict(),\n            'scaler_state': scaler.state_dict() if scaler else None,\n            'ema_state': ema.state_dict(),\n        }, ckpt_path)\n\n        if patience >= cfg['S2_PATIENCE']:\n            print(f\"  Early stop Stage 2 pada epoch {ep+1}\")\n            break\n        torch.cuda.empty_cache()\n\n    if best_state_ema is not None:\n        model.load_state_dict(best_state_ema)\n        print(f\"\\n  Stage 2 selesai. Checkpoint versi EMA (best val_sens={best_sens:.4f})\")\n    elif best_state is not None:\n        model.load_state_dict(best_state)\n        print(f\"\\n  Stage 2 selesai (fallback RAW). Best val_sens={best_sens:.4f}\")\n    return model, history\n\n\n# ============================================================================\n# PERBAIKAN #10 -- EVALUASI DENGAN TTA\n# ============================================================================\n\n@torch.no_grad()\ndef evaluate_supervised(model, test_df, cache, cfg):\n    model.eval()\n    ds = APTOSDataset(test_df, cache, aug_fn=None)\n    loader = DataLoader(ds, batch_size=32, shuffle=False, num_workers=2, pin_memory=True)\n    y_true, y_pred, y_prob = [], [], []\n    for imgs, labels in tqdm(loader, desc=\"  Eval supervised\", leave=False):\n        imgs = imgs.to(DEVICE)\n        logits = model.forward_supervised(imgs)\n        y_prob.extend(F.softmax(logits, dim=1).cpu().numpy())\n        y_pred.extend(logits.argmax(1).cpu().numpy())\n        y_true.extend(labels.numpy())\n    return compute_metrics(np.array(y_true), np.array(y_pred)), np.array(y_prob)\n\n\n@torch.no_grad()\ndef evaluate_fsl_prototype(model, train_df, test_df, cache, cfg, tta_n=8):\n    \"\"\"Prototype dari projection head + TTA x8.\"\"\"\n    model.eval()\n    rng = np.random.default_rng(SEED)\n\n    ds = APTOSDataset(train_df, cache, aug_fn=None)\n    loader = DataLoader(ds, batch_size=32, shuffle=False, num_workers=2, pin_memory=True)\n    embs = {c: [] for c in range(5)}\n    for imgs, labels in tqdm(loader, desc=\"  Hitung prototype\", leave=False):\n        imgs = imgs.to(DEVICE)\n        with autocast('cuda', enabled=cfg['USE_AMP']):\n            e = model.forward_embedding(imgs).float()\n        for emb, lbl in zip(e.cpu(), labels):\n            embs[lbl.item()].append(emb)\n    sorted_cls = sorted([c for c, v in embs.items() if v])\n    protos = F.normalize(torch.stack([torch.stack(embs[c]).mean(0)\n                                      for c in sorted_cls]), dim=1).to(DEVICE)\n\n    df_r = test_df.reset_index(drop=True)\n    y_true, y_pred = [], []\n    for idx in tqdm(range(len(df_r)), desc=f\"  Eval FSL (TTA x{tta_n})\", leave=False):\n        row = df_r.iloc[idx]\n        raw = cache[row['id_code']]\n        views = [normalize_to_tensor(raw)]\n        for _ in range(tta_n - 1):\n            views.append(normalize_to_tensor(augment_s2(raw, rng)))\n        batch = torch.stack(views).to(DEVICE)\n        with autocast('cuda', enabled=cfg['USE_AMP']):\n            e = model.forward_embedding(batch).float()\n        emb = F.normalize(e.mean(0, keepdim=True), dim=1)\n        dist = 1.0 - torch.matmul(emb, protos.T)\n        y_pred.append(sorted_cls[dist.argmin(dim=1).item()])\n        y_true.append(int(row['diagnosis']))\n    return compute_metrics(np.array(y_true), np.array(y_pred))\n\n\ndef print_results(all_results: dict):\n    CLASS_NAMES = ['Grade 0', 'Grade 1', 'Grade 2', 'Grade 3', 'Grade 4']\n    print(\"\\n\" + \"=\" * 70)\n    print(\"  HASIL AKHIR -- EKSPERIMEN B (METODOLOGI + PERBAIKAN)\")\n    print(\"=\" * 70)\n    for name, m in all_results.items():\n        ok = (m['accuracy'] >= TARGET_ACC and m['sensitivity_macro'] >= TARGET_SENS\n              and m['specificity_macro'] >= TARGET_SPEC)\n        print(f\"\\n{'[LULUS]' if ok else '[BELUM]'} {name}:\")\n        print(f\"  accuracy             {m['accuracy']:.4f}  (>=0.80: \"\n              f\"{'LULUS' if m['accuracy'] >= TARGET_ACC else 'BELUM'})\")\n        print(f\"  sensitivity_macro    {m['sensitivity_macro']:.4f}  (>=0.85: \"\n              f\"{'LULUS' if m['sensitivity_macro'] >= TARGET_SENS else 'BELUM'})\")\n        print(f\"  specificity_macro    {m['specificity_macro']:.4f}  (>=0.75: \"\n              f\"{'LULUS' if m['specificity_macro'] >= TARGET_SPEC else 'BELUM'})\")\n        print(f\"  QWK: {m['qwk']:.4f} | AUC: {m['auc']:.4f} | F1: {m['weighted_f1']:.4f}\")\n        print(f\"  Sensitivity per kelas:\")\n        for c, s in enumerate(m['sensitivity_per_class']):\n            print(f\"    {CLASS_NAMES[c]}: {s:.4f}  |{'#' * int(s * 20):<20}|\")\n\n\n# ============================================================================\n# RUNNER UTAMA\n# ============================================================================\n\ndef find_dataset(base_path='/kaggle/input'):\n    csvs = glob.glob(os.path.join(base_path, '**', 'train.csv'), recursive=True)\n    if not csvs:\n        raise FileNotFoundError(f\"train.csv tidak ditemukan di {base_path}\")\n    csv_path = csvs[0]\n    img_dir = os.path.join(os.path.dirname(csv_path), 'train_images')\n    if not os.path.exists(img_dir):\n        cands = glob.glob(os.path.join(base_path, '**', 'train_images'), recursive=True)\n        if not cands:\n            raise FileNotFoundError(\"train_images tidak ditemukan\")\n        img_dir = cands[0]\n    return csv_path, img_dir\n\n# ============================================================================\n# VISUALISASI HASIL -- tempel SEBELUM `def main():` (identik untuk A dan B)\n# ============================================================================\nimport math\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import roc_curve, auc as sk_auc\n\nPLOT_CLASSES = ['Grade 0', 'Grade 1', 'Grade 2', 'Grade 3', 'Grade 4']\n\n\ndef _finish(fig, out_dir, fname):\n    fig.tight_layout()\n    path = os.path.join(out_dir, fname)\n    fig.savefig(path, dpi=140, bbox_inches='tight')\n    plt.show()\n    plt.close(fig)\n    print(f\"  Grafik disimpan: {path}\")\n\n\ndef plot_confusion_matrices(results, out_dir, tag):\n    n = len(results); cols = min(n, 3); rows = math.ceil(n / cols)\n    fig, axes = plt.subplots(rows, cols, figsize=(6 * cols, 5.4 * rows), squeeze=False)\n    for ax, (name, m) in zip(axes.ravel(), results.items()):\n        cm = np.array(m['confusion_matrix'])\n        norm = cm / np.maximum(cm.sum(axis=1, keepdims=True), 1)\n        ax.imshow(norm, cmap='Blues', vmin=0, vmax=1)\n        for i in range(cm.shape[0]):\n            for j in range(cm.shape[1]):\n                ax.text(j, i, f\"{cm[i, j]}\\n{norm[i, j]*100:.0f}%\", ha='center', va='center',\n                        fontsize=9, color='white' if norm[i, j] > 0.5 else 'black')\n        ax.set_xticks(range(5)); ax.set_yticks(range(5))\n        ax.set_xticklabels(PLOT_CLASSES, rotation=45, ha='right'); ax.set_yticklabels(PLOT_CLASSES)\n        ax.set_xlabel('Prediksi'); ax.set_ylabel('Sebenarnya')\n        ax.set_title(f\"{name}\\nAcc={m['accuracy']:.3f} | F1={m['weighted_f1']:.3f} | QWK={m['qwk']:.3f}\",\n                     fontsize=10)\n    for ax in axes.ravel()[n:]:\n        ax.axis('off')\n    fig.suptitle(f'Confusion Matrix (uji) -- {tag}', fontsize=14, fontweight='bold')\n    _finish(fig, out_dir, f'{tag}_1_confusion_matrix.png')\n\n\ndef plot_summary_metrics(results, out_dir, tag):\n    keys = ['accuracy', 'weighted_f1', 'sensitivity_macro', 'specificity_macro', 'qwk', 'auc']\n    labels = ['Accuracy', 'F1 (weighted)', 'Sens. macro', 'Spec. macro', 'QWK', 'AUC*']\n    names = list(results.keys())\n    x = np.arange(len(keys)); w = 0.8 / len(names)\n    fig, ax = plt.subplots(figsize=(13, 5.5))\n    for i, nm in enumerate(names):\n        vals = [results[nm][k] for k in keys]\n        bars = ax.bar(x + i * w - 0.4 + w / 2, vals, w, label=nm)\n        for b, v in zip(bars, vals):\n            ax.text(b.get_x() + b.get_width() / 2, v + 0.01, f\"{v:.2f}\", ha='center', fontsize=7)\n    for val, col, txt in [(TARGET_ACC, 'red', 'target acc'), (TARGET_SENS, 'orange', 'target sens'),\n                          (TARGET_SPEC, 'green', 'target spec')]:\n        ax.axhline(val, color=col, ls='--', lw=1, alpha=0.7)\n        ax.text(len(keys) - 0.45, val + 0.005, f'{txt} {val}', color=col, fontsize=8, ha='right')\n    ax.set_xticks(x); ax.set_xticklabels(labels); ax.set_ylim(0, 1.1)\n    ax.set_title(f'Ringkasan metrik uji -- {tag}   (*AUC dari label prediksi, bukan probabilitas)')\n    ax.legend(fontsize=8, loc='lower right'); ax.grid(axis='y', alpha=0.3)\n    _finish(fig, out_dir, f'{tag}_2_ringkasan_metrik.png')\n\n\ndef plot_per_class(results, out_dir, tag):\n    names = list(results.keys()); w = 0.8 / len(names)\n    fig, axes = plt.subplots(1, 2, figsize=(15, 5))\n    for ax, key, ttl in zip(axes, ['sensitivity_per_class', 'specificity_per_class'],\n                            ['Sensitivitas per kelas', 'Spesifisitas per kelas']):\n        for i, nm in enumerate(names):\n            ax.bar(np.arange(5) + i * w - 0.4 + w / 2, results[nm][key], w, label=nm)\n        ax.set_xticks(range(5)); ax.set_xticklabels(PLOT_CLASSES)\n        ax.set_ylim(0, 1.05); ax.set_title(f'{ttl} -- {tag}'); ax.grid(axis='y', alpha=0.3)\n    axes[0].axhline(TARGET_SENS, color='orange', ls='--', lw=1)\n    axes[1].axhline(TARGET_SPEC, color='green', ls='--', lw=1)\n    axes[0].legend(fontsize=7, loc='upper right')\n    _finish(fig, out_dir, f'{tag}_3_per_kelas.png')\n\n\ndef plot_training_curves(histories, out_dir, tag):\n    n = len(histories)\n    fig, axes = plt.subplots(n, 2, figsize=(13, 3.8 * n), squeeze=False)\n    for r, (name, h) in enumerate(histories.items()):\n        for k, v in h.items():\n            ax = axes[r, 0] if 'loss' in k else axes[r, 1]\n            ax.plot(range(1, len(v) + 1), v, label=k)\n        axes[r, 0].set_title(f'{name} -- loss'); axes[r, 1].set_title(f'{name} -- metrik')\n        for ax in axes[r]:\n            ax.set_xlabel('epoch'); ax.legend(fontsize=8); ax.grid(alpha=0.3)\n    _finish(fig, out_dir, f'{tag}_4_kurva_training.png')\n\n\ndef plot_roc(roc_data, out_dir, tag):\n    \"\"\"roc_data: {nama: (y_true, y_prob[N,5])} -- ROC sungguhan dari probabilitas softmax.\"\"\"\n    if not roc_data:\n        return\n    n = len(roc_data); cols = min(n, 3); rows = math.ceil(n / cols)\n    fig, axes = plt.subplots(rows, cols, figsize=(6 * cols, 5.2 * rows), squeeze=False)\n    for ax, (name, (y_true, y_prob)) in zip(axes.ravel(), roc_data.items()):\n        y_true = np.asarray(y_true); y_prob = np.asarray(y_prob)\n        aucs = []\n        for c in range(5):\n            yb = (y_true == c).astype(int)\n            if yb.sum() == 0:\n                continue\n            fpr, tpr, _ = roc_curve(yb, y_prob[:, c]); a = sk_auc(fpr, tpr); aucs.append(a)\n            ax.plot(fpr, tpr, label=f'{PLOT_CLASSES[c]} (AUC={a:.3f})')\n        ax.plot([0, 1], [0, 1], 'k--', lw=0.8)\n        ax.set_title(f'{name}\\nROC OvR, macro AUC={np.mean(aucs):.3f}', fontsize=10)\n        ax.set_xlabel('FPR'); ax.set_ylabel('TPR'); ax.legend(fontsize=8); ax.grid(alpha=0.3)\n    for ax in axes.ravel()[n:]:\n        ax.axis('off')\n    _finish(fig, out_dir, f'{tag}_5_roc.png')\n\n\ndef run_all_plots(results, histories, roc_data, out_dir, tag):\n    print(\"\\n  Membuat grafik evaluasi...\")\n    plot_confusion_matrices(results, out_dir, tag)\n    plot_summary_metrics(results, out_dir, tag)\n    plot_per_class(results, out_dir, tag)\n    if histories:\n        plot_training_curves(histories, out_dir, tag)\n    plot_roc(roc_data, out_dir, tag)\n\ndef main():\n    print(\"\\n\" + \"=\" * 70)\n    print(\"  EKSPERIMEN B -- METODOLOGI + PERBAIKAN (IMPROVED)\")\n    print(\"=\" * 70)\n    print(f\"Device: {DEVICE} | EfficientNet input: {EFF_IMG_SIZE}px\")\n    if EFF_IMG_SIZE != 380:\n        print(\"  CATATAN: Eksperimen A memakai 380px untuk EfficientNet-B4.\")\n        print(\"  Ukuran input menjadi variabel perancu -- sebutkan di pembahasan,\")\n        print(\"  atau set EFF_IMG_SIZE=380 untuk ablasi yang benar-benar bersih.\")\n\n    CSV_PATH, IMG_DIR = find_dataset()\n    df = pd.read_csv(CSV_PATH)\n    print(f\"\\nDistribusi kelas:\\n{df['diagnosis'].value_counts().sort_index().to_string()}\")\n\n    # Split IDENTIK dengan Eksperimen A\n    train_df, temp_df = train_test_split(df, test_size=0.30,\n                                         stratify=df['diagnosis'], random_state=SEED)\n    val_df, test_df = train_test_split(temp_df, test_size=0.50,\n                                       stratify=temp_df['diagnosis'], random_state=SEED)\n    print(f\"\\nSplit: Train={len(train_df)} | Val={len(val_df)} | Test={len(test_df)}\")\n\n    all_df = pd.concat([train_df, val_df, test_df]).drop_duplicates(subset='id_code')\n    print(\"\\n[1/2] Cache 224px (ResNet-50)...\")\n    cache_rn = build_cache(all_df, IMG_DIR, 224)\n    print(f\"\\n[2/2] Cache {EFF_IMG_SIZE}px (EfficientNet-B4)...\")\n    cache_eff = build_cache(all_df, IMG_DIR, EFF_IMG_SIZE)\n\n    cw = compute_class_weight('balanced', classes=np.unique(train_df['diagnosis']),\n                            y=train_df['diagnosis'])\n    cw_tensor = torch.tensor(cw, dtype=torch.float32).to(DEVICE)\n\n    results = {}\n    probs_store = {}\n\n    # ---------------- ResNet-50 ----------------\n    print(\"\\n\" + \"#\" * 70 + \"\\n  RESNET-50 (IMPROVED)\\n\" + \"#\" * 70)\n    m_rn = TwoStageModel('resnet50', proj_dim=CONFIG_RESNET['PROJ_DIM']).to(DEVICE)\n    m_rn,h_rn_s1 = train_stage1(m_rn, train_df, val_df, cache_rn, CONFIG_RESNET)\n    torch.save({'model_state': m_rn.state_dict(), 'backbone_name': 'resnet50',\n                'proj_dim': 256, 'img_size': 224, 'stage': 1},\n               os.path.join(CKPT_DIR, 'model_rn_after_s1.pt'))\n    results['ResNet-50 (Supervised)'], probs_store['rn'] = evaluate_supervised(\n        m_rn, test_df, cache_rn, CONFIG_RESNET)\n\n    m_rn,h_rn_s2 = train_stage2(m_rn, train_df, val_df, cache_rn, CONFIG_RESNET, cw_tensor)\n    torch.save({'model_state': m_rn.state_dict(), 'backbone_name': 'resnet50',\n                'proj_dim': 256, 'img_size': 224, 'stage': 2},\n               os.path.join(CKPT_DIR, 'model_rn_final.pt'))\n    results['ResNet-50 (FSL Prototype + TTA)'] = evaluate_fsl_prototype(\n        m_rn, train_df, test_df, cache_rn, CONFIG_RESNET, tta_n=8)\n\n    del m_rn\n    torch.cuda.empty_cache(); gc.collect()\n\n    # ---------------- EfficientNet-B4 ----------------\n    print(\"\\n\" + \"#\" * 70 + \"\\n  EFFICIENTNET-B4 (IMPROVED)\\n\" + \"#\" * 70)\n    m_eff = TwoStageModel('efficientnet_b4', proj_dim=CONFIG_EFF['PROJ_DIM']).to(DEVICE)\n    m_eff,h_eff_s1 = train_stage1(m_eff, train_df, val_df, cache_eff, CONFIG_EFF)\n    torch.save({'model_state': m_eff.state_dict(), 'backbone_name': 'efficientnet_b4',\n                'proj_dim': 256, 'img_size': EFF_IMG_SIZE, 'stage': 1},\n               os.path.join(CKPT_DIR, 'model_eff_after_s1.pt'))\n    results['EfficientNet-B4 (Supervised)'], probs_store['eff'] = evaluate_supervised(\n        m_eff, test_df, cache_eff, CONFIG_EFF)\n\n    m_eff,h_eff_s2 = train_stage2(m_eff, train_df, val_df, cache_eff, CONFIG_EFF, cw_tensor)\n    torch.save({'model_state': m_eff.state_dict(), 'backbone_name': 'efficientnet_b4',\n                'proj_dim': 256, 'img_size': EFF_IMG_SIZE, 'stage': 2},\n               os.path.join(CKPT_DIR, 'model_eff_final.pt'))\n    results['EfficientNet-B4 (FSL Prototype + TTA)'] = evaluate_fsl_prototype(\n        m_eff, train_df, test_df, cache_eff, CONFIG_EFF, tta_n=8)\n\n    # ---------------- Ensemble supervised (varian terbaik) ----------------\n    print(\"\\n\" + \"#\" * 70 + \"\\n  ENSEMBLE SUPERVISED\\n\" + \"#\" * 70)\n    prob_ens = (probs_store['rn'] + probs_store['eff']) / 2\n    results['Ensemble (Supervised)'] = compute_metrics(\n        test_df['diagnosis'].values, prob_ens.argmax(axis=1))\n\n    print_results(results)\n\n    y_test = test_df['diagnosis'].values\n    histories = {\n        'ResNet-50 Stage 1':       h_rn_s1,\n        'ResNet-50 Stage 2 (FSL)': h_rn_s2,\n        'EffNet-B4 Stage 1':       h_eff_s1,\n        'EffNet-B4 Stage 2 (FSL)': h_eff_s2,\n    }\n    run_all_plots(results, histories,\n                  roc_data={'ResNet-50 (Supervised)': (y_test, probs_store['rn']),\n                            'EfficientNet-B4 (Supervised)': (y_test, probs_store['eff']),\n                            'Ensemble (Supervised)': (y_test, prob_ens)},\n                  out_dir=CKPT_DIR, tag='B_improved')\n\n    out_path = os.path.join(CKPT_DIR, 'hasil_improved.json')\n    with open(out_path, 'w') as f:\n        json.dump({'eksperimen': 'B -- Metodologi + Perbaikan (Improved)',\n                   'split': '70/15/15', 'seed': SEED,\n                   'eff_img_size': EFF_IMG_SIZE, 'hasil': results,\n                   'histories': histories}, f, indent=2)\n    print(f\"\\nHasil tersimpan: {out_path}\")\n    print(\"Jalankan bandingkan_hasil.py untuk membuat tabel ablasi A vs B.\")\n    return results\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}