{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":14774,"databundleVersionId":875431},{"sourceType":"datasetVersion","sourceId":15000581,"datasetId":9601748,"databundleVersionId":15875699}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%writefile dataset.py\n\"\"\"\ndataset.py — APTOS 2019 Dataset\n================================\nBám sát: Ali et al., IEEE JTEHM 2023, Section IV.B & IV.E\n\nPreprocessing (Section IV.B):\n  1. Resize 256×256\n  2. Histogram Equalization (Eq. 5-6) — đúng bài báo, KHÔNG dùng CLAHE\n  3. Intensity Normalization Min-Max (Eq. 7)\n\nAugmentation (Section IV.E):\n  - Rotation ±15°   ← bài báo nêu rõ\n  - Scaling 0.9–1.1 ← bài báo nêu rõ\n  [KHÔNG thêm Flip/Brightness — giữ đúng setup bài báo để so sánh fair]\n\nCải tiến kỹ thuật giữ lại (không ảnh hưởng kết quả):\n  - Image caching: tránh preprocess lại mỗi epoch → tăng tốc đáng kể\n\"\"\"\n\nimport os\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset\nfrom torchvision import transforms\n\n\nclass APTOSDataset(Dataset):\n    def __init__(self, csv_file, img_dir, is_train=True, cache_dir=None):\n        \"\"\"\n        Args:\n            csv_file  : đường dẫn tới file CSV\n                        APTOS 2019 → cột: id_code, diagnosis\n                        EyePACS    → cột: image,   level\n            img_dir   : thư mục chứa ảnh (.png hoặc .jpeg)\n            is_train  : True → áp dụng augmentation\n            cache_dir : nếu set, cache ảnh đã preprocess để tăng tốc\n        \"\"\"\n        self.data_frame = pd.read_csv(csv_file)\n        self.img_dir    = img_dir\n        self.is_train   = is_train\n        self.cache_dir  = cache_dir\n\n        # ── Tự động nhận dạng tên cột CSV ────────────────────────────────────\n        cols = self.data_frame.columns.tolist()\n        # Cột ID ảnh\n        if 'id_code' in cols:\n            self.col_id = 'id_code'          # APTOS 2019\n        elif 'image' in cols:\n            self.col_id = 'image'            # EyePACS\n        else:\n            raise ValueError(f\"CSV không có cột 'id_code' hoặc 'image'. Columns: {cols}\")\n        # Cột nhãn\n        if 'diagnosis' in cols:\n            self.col_label = 'diagnosis'     # APTOS 2019\n        elif 'level' in cols:\n            self.col_label = 'level'         # EyePACS\n        else:\n            raise ValueError(f\"CSV không có cột 'diagnosis' hoặc 'level'. Columns: {cols}\")\n\n        print(f\"  [Dataset] CSV columns → id='{self.col_id}', label='{self.col_label}' \"\n              f\"| {len(self.data_frame)} samples\")\n\n        if cache_dir:\n            os.makedirs(cache_dir, exist_ok=True)\n\n        # Normalize theo ImageNet mean/std (do dùng pretrained backbone)\n        self.to_tensor = transforms.Compose([\n            transforms.ToTensor(),\n            transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                 std=[0.229, 0.224, 0.225])\n        ])\n\n    # ── Preprocessing (Section IV.B) ─────────────────────────────────────────\n    def preprocess_image(self, image: np.ndarray) -> np.ndarray:\n        \"\"\"\n        Đúng theo bài báo Section IV.B:\n          Step 1 — Resize về 256×256\n          Step 2 — Histogram Equalization (Eq. 5-6)\n                   Áp dụng trên kênh Grayscale, sau đó merge lại RGB\n                   (cách phổ biến nhất khớp với mô tả bài báo)\n          Step 3 — Min-Max Intensity Normalization (Eq. 7)\n                   X_norm = (I - Min) * (Max' - Min') / (Max - Min) + Min'\n                   với Min'=0, Max'=1 → scale về [0,255] để lưu cache\n        \"\"\"\n        # Step 1: Resize\n        image = cv2.resize(image, (256, 256))\n        image = image.astype(np.uint8)\n\n        # Step 2: Histogram Equalization (Eq. 5-6 bài báo)\n        # Bài báo mô tả HE trực tiếp trên grayscale (gi,j = floor(L-1) * Σ Pn)\n        # Áp dụng HE độc lập trên từng kênh R, G, B — đúng nghĩa đen Eq. 5-6\n        r, g, b     = cv2.split(image)\n        r_eq        = cv2.equalizeHist(r)\n        g_eq        = cv2.equalizeHist(g)\n        b_eq        = cv2.equalizeHist(b)\n        image       = cv2.merge((r_eq, g_eq, b_eq))\n\n        # Step 3: Min-Max Normalization (Eq. 7) — X_norm = (I-Min)/(Max-Min)\n        image_f = image.astype(np.float32)\n        lo, hi  = image_f.min(), image_f.max()\n        if hi > lo:\n            image_f = (image_f - lo) / (hi - lo)\n        else:\n            image_f = np.zeros_like(image_f)\n\n        return (image_f * 255.0).astype(np.uint8)\n\n    # ── Augmentation (Section IV.E) ──────────────────────────────────────────\n    def augment_image(self, image: np.ndarray) -> np.ndarray:\n        \"\"\"\n        Đúng theo bài báo Section IV.E:\n          - Rotation  : ±15°      (nêu rõ trong bài báo)\n          - Scaling   : 0.9–1.1   (nêu rõ trong bài báo)\n\n        KHÔNG thêm Flip hay Brightness jitter để bám sát setup gốc.\n        Nếu muốn thử ablation, có thể enable từng cái bên dưới.\n        \"\"\"\n        h, w   = image.shape[:2]\n        cx, cy = w / 2, h / 2\n\n        # Rotation ±15° (Section IV.E)\n        if np.random.rand() > 0.5:\n            angle = np.random.uniform(-15, 15)\n            M     = cv2.getRotationMatrix2D((cx, cy), angle, 1.0)\n            image = cv2.warpAffine(image, M, (w, h),\n                                   borderMode=cv2.BORDER_REFLECT_101)\n\n        # Scaling 0.9–1.1 (Section IV.E)\n        if np.random.rand() > 0.5:\n            scale = np.random.uniform(0.9, 1.1)\n            M     = cv2.getRotationMatrix2D((cx, cy), 0, scale)\n            image = cv2.warpAffine(image, M, (w, h),\n                                   borderMode=cv2.BORDER_REFLECT_101)\n\n        # ── Augmentation bổ sung (KHÔNG có trong bài báo — tắt mặc định) ──\n        # Bỏ comment để enable khi muốn thực nghiệm ablation:\n        #\n        # Horizontal Flip (fundus images đối xứng trái-phải):\n        # if np.random.rand() > 0.5:\n        #     image = cv2.flip(image, 1)\n        #\n        # Brightness/Contrast jitter (ánh sáng fundus không đồng đều):\n        # if np.random.rand() > 0.5:\n        #     alpha = np.random.uniform(0.8, 1.2)\n        #     beta  = np.random.randint(-15, 15)\n        #     image = np.clip(image.astype(np.float32) * alpha + beta,\n        #                     0, 255).astype(np.uint8)\n\n        return image\n\n    # ── Cache helpers ─────────────────────────────────────────────────────────\n    def _cached_path(self, img_id):\n        return os.path.join(self.cache_dir, f\"{img_id}_pre.png\")\n\n    def _load_raw(self, img_id):\n        # Thử lần lượt các extension: EyePACS dùng .jpeg, APTOS dùng .png\n        for ext in ('.jpeg', '.png', '.jpg'):\n            img_path = os.path.join(self.img_dir, img_id + ext)\n            raw = cv2.imread(img_path)\n            if raw is not None:\n                return cv2.cvtColor(raw, cv2.COLOR_BGR2RGB)\n        raise FileNotFoundError(\n            f\"Không tìm thấy ảnh '{img_id}' (.jpeg/.png/.jpg) trong: {self.img_dir}\"\n        )\n\n    # ── Dataset interface ─────────────────────────────────────────────────────\n    def __len__(self):\n        return len(self.data_frame)\n\n    def __getitem__(self, idx):\n        row    = self.data_frame.iloc[idx]\n        img_id = str(row[self.col_id])\n        label  = int(row[self.col_label])\n\n        # Load từ cache nếu có, nếu không thì preprocess và lưu cache\n        if self.cache_dir:\n            cached = self._cached_path(img_id)\n            if os.path.exists(cached):\n                img = cv2.imread(cached)\n                image = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) if img is not None else None\n            else:\n                image = None\n\n            if image is None:\n                raw   = self._load_raw(img_id)\n                image = self.preprocess_image(raw)\n                cv2.imwrite(cached, cv2.cvtColor(image, cv2.COLOR_RGB2BGR))\n        else:\n            raw   = self._load_raw(img_id)\n            image = self.preprocess_image(raw)\n\n        # Augmentation chỉ áp dụng khi training\n        if self.is_train:\n            image = self.augment_image(image)\n\n        return self.to_tensor(image), label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T17:01:17.995711Z","iopub.execute_input":"2026-03-02T17:01:17.995979Z","iopub.status.idle":"2026-03-02T17:01:18.009058Z","shell.execute_reply.started":"2026-03-02T17:01:17.995949Z","shell.execute_reply":"2026-03-02T17:01:18.008345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile models.py\n\"\"\"\nmodels.py — IR-CNN: Hybrid InceptionV3 + ResNet50 + Custom CNN\n==============================================================\nBám sát: Ali et al., \"A Hybrid CNN Model for Automatic DR\nClassification From Fundus Images\", IEEE JTEHM 2023.\n\nKiến trúc (Figure 1 bài báo):\n  InceptionV3  ──┐\n                 ├─ Concat (4096-ch) ─► Conv(4096→64) ─► Block1 ─► Block2\n  ResNet50     ──┘                      MaxPool Dropout0.28\n                                        ─► AdaptivePool ─► Flatten\n                                        ─► Dense(256→128) Dropout0.5\n                                        ─► Dense(128→16)  Dropout0.5\n                                        ─► Dense(16→5)\n\nGhi chú về Figure 1:\n  Figure 1 hiển thị các label block là: 256×3×3, 128×3×3, 128×3, 64×3, 64×1, 32×3\n  Tuy nhiên các label này KHÔNG phải số output channels mà là kích thước\n  spatial feature map sau mỗi block (256→128→64→32 do MaxPool liên tục từ 256×256).\n  Output channels của từng block (64→128→256) là suy luận chuẩn theo kiến trúc CNN.\n  Bài báo không mô tả rõ số channels tại mỗi block.\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\nimport warnings\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\n\n\n# ─────────────────────────────────────────────────────────────────────────────\nclass ConvBlock(nn.Module):\n    \"\"\"\n    Block tái sử dụng sau Concatenation (Figure 1 bài báo):\n      Conv2D → BN → ReLU → Conv2D → BN → ReLU → MaxPool → Dropout\n\n    Dropout mặc định 0.25 — không nêu rõ trong bài báo cho các block này,\n    dùng giá trị thông thường.\n    \"\"\"\n    def __init__(self, in_channels, out_channels, dropout=0.25):\n        super().__init__()\n        self.conv1   = nn.Conv2d(in_channels,  out_channels, kernel_size=3, padding=1)\n        self.bn1     = nn.BatchNorm2d(out_channels)\n        self.conv2   = nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1)\n        self.bn2     = nn.BatchNorm2d(out_channels)\n        self.pool    = nn.MaxPool2d(kernel_size=2, stride=2)\n        self.dropout = nn.Dropout(dropout)\n\n    def forward(self, x):\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = F.relu(self.bn2(self.conv2(x)))\n        x = self.pool(x)\n        x = self.dropout(x)\n        return x\n\n\n# ─────────────────────────────────────────────────────────────────────────────\nclass IR_CNN(nn.Module):\n    \"\"\"\n    IR-CNN — InceptionV3 + ResNet50 CNN (Ali et al., 2023)\n\n    Toàn bộ model được train cùng lúc (không freeze backbone),\n    đúng với mô tả Section IV.C: \"all three network models were trained\n    over 100 epochs, employing a learning rate of 0.001.\"\n    \"\"\"\n    def __init__(self, num_classes=5):\n        super().__init__()\n\n        # ── 1. ResNet50 backbone (pretrained ImageNet) ──────────────────────\n        # Bài báo Section III.A: \"ResNet50 model was utilized for feature\n        # extraction from the DR images.\"\n        # Bỏ AdaptiveAvgPool + FC → giữ feature map [B, 2048, H, W]\n        try:\n            from torchvision.models import ResNet50_Weights\n            resnet = models.resnet50(weights=ResNet50_Weights.IMAGENET1K_V1)\n        except ImportError:\n            resnet = models.resnet50(pretrained=True)\n        self.resnet_features = nn.Sequential(*list(resnet.children())[:-2])\n\n        # ── 2. InceptionV3 backbone (pretrained ImageNet) ───────────────────\n        # Bài báo Section III.B: \"InceptionV3 model ... extensively used\n        # for classification purposes.\"\n        try:\n            from torchvision.models import Inception_V3_Weights\n            inception = models.inception_v3(weights=Inception_V3_Weights.IMAGENET1K_V1)\n        except ImportError:\n            inception = models.inception_v3(pretrained=True)\n        inception.aux_logits = False\n        inception.AuxLogits  = None\n\n        # Lấy toàn bộ feature extraction đến Mixed_7c (trước AvgPool + FC)\n        self.inception_features = nn.Sequential(\n            inception.Conv2d_1a_3x3,\n            inception.Conv2d_2a_3x3,\n            inception.Conv2d_2b_3x3,\n            inception.maxpool1,\n            inception.Conv2d_3b_1x1,\n            inception.Conv2d_4a_3x3,\n            inception.maxpool2,\n            inception.Mixed_5b, inception.Mixed_5c, inception.Mixed_5d,\n            inception.Mixed_6a, inception.Mixed_6b, inception.Mixed_6c,\n            inception.Mixed_6d, inception.Mixed_6e,\n            inception.Mixed_7a, inception.Mixed_7b, inception.Mixed_7c,\n        )\n\n        # ── 3. Lớp Conv đầu sau Concatenation (Figure 1) ───────────────────\n        # Figure 1: Conv2D 32×32×3 → BN → MaxPool2D → Dropout 0.28\n        # Input: 4096 channels (2048 ResNet + 2048 Inception)\n        # Dropout=0.28 được ghi rõ trong Figure 1\n        self.initial_conv = nn.Conv2d(4096, 64, kernel_size=3, padding=1)\n        self.initial_bn   = nn.BatchNorm2d(64)\n        self.initial_pool = nn.MaxPool2d(kernel_size=2, stride=2)\n        self.initial_drop = nn.Dropout(0.28)   # ← Figure 1 ghi rõ 0.28\n\n        # ── 4. Custom CNN Blocks (Figure 1) ─────────────────────────────────\n        # Figure 1 label các block theo spatial size: 256, 128, 64, 32\n        # (= kích thước feature map sau mỗi MaxPool, không phải số channels)\n        # Output channels tăng dần theo chuẩn CNN: 64 → 128 → 256\n        # Block 1: 64→128 (spatial: ~4×4 → ~2×2)\n        # Block 2: 128→256 (spatial: ~2×2 → ~1×1 sau AdaptivePool)\n        self.block1 = ConvBlock(64,  128)\n        self.block2 = ConvBlock(128, 256)\n\n        self.adaptive_pool = nn.AdaptiveAvgPool2d((1, 1))\n        self.flatten       = nn.Flatten()\n\n        # ── 5. Dense Layers (Figure 1: Dense128 → Dense16 → Dense5) ─────────\n        # Dropout 0.5 giữa các Dense layer (Figure 1 ghi \"Dropout 0.5\")\n        # NOTE: Input 256 của fc1 là từ block2 output (128→256).\n        #       Con số 256 này là suy luận từ kiến trúc — bài báo không\n        #       ghi rõ số channel vào Dense head. AdaptiveAvgPool → 256-d vector.\n        self.fc1      = nn.Linear(256, 128)\n        self.drop_fc1 = nn.Dropout(0.5)\n        self.fc2      = nn.Linear(128, 16)\n        self.drop_fc2 = nn.Dropout(0.5)\n        self.fc3      = nn.Linear(16, num_classes)\n\n    def get_trainable_params(self):\n        return sum(p.numel() for p in self.parameters() if p.requires_grad)\n\n    def forward(self, x):\n        # ResNet50 features: [B, 2048, 8, 8] với input 256×256\n        f_res = self.resnet_features(x)\n\n        # InceptionV3 features: [B, 2048, ~6, ~6] với input 256×256\n        f_inc = self.inception_features(x)\n\n        # Align spatial size trước khi concat (cần thiết vì InceptionV3\n        # được thiết kế cho 299×299, input 256×256 tạo ra size khác nhau)\n        if f_inc.shape[2:] != f_res.shape[2:]:\n            f_inc = F.interpolate(f_inc, size=f_res.shape[2:],\n                                  mode='bilinear', align_corners=False)\n\n        # Concatenate → [B, 4096, 8, 8]\n        x = torch.cat([f_res, f_inc], dim=1)\n\n        # Custom CNN (Figure 1)\n        x = F.relu(self.initial_bn(self.initial_conv(x)))\n        x = self.initial_pool(x)\n        x = self.initial_drop(x)\n        x = self.block1(x)\n        x = self.block2(x)\n\n        # Dense head (Figure 1)\n        x   = self.adaptive_pool(x)\n        x   = self.flatten(x)\n        x   = F.relu(self.fc1(x))\n        x   = self.drop_fc1(x)\n        x   = F.relu(self.fc2(x))\n        x   = self.drop_fc2(x)\n        out = self.fc3(x)\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T17:01:18.195097Z","iopub.execute_input":"2026-03-02T17:01:18.195586Z","iopub.status.idle":"2026-03-02T17:01:18.203708Z","shell.execute_reply.started":"2026-03-02T17:01:18.195559Z","shell.execute_reply":"2026-03-02T17:01:18.202970Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\ntrain.py — Huấn luyện IR-CNN  (đã fix)\n========================================\nBám sát: Ali et al., IEEE JTEHM 2023, Section IV.C\n\n═══════════════════════════════════════════════════════════════\n  PHÂN TÍCH VẤN ĐỀ VÀ FIX SO VỚI PHIÊN BẢN CŨ\n═══════════════════════════════════════════════════════════════\n\n❌ LỖI 1 — Double Weighting (nguyên nhân Val Loss tăng liên tục):\n   Phiên bản cũ dùng đồng thời WeightedRandomSampler + WeightedCELoss.\n   WeightedSampler đã cân bằng batch (mỗi class ~20%) → model thấy mỗi\n   class ~20%. WeightedCELoss lại nhân thêm penalty C3×8.09, C4×10.24\n   → model bị over-push sang minority classes trong train.\n   Val set có phân phối tự nhiên (class 0 = 73%) → model predict toàn\n   minority classes → val accuracy 8-15% (thấp hơn cả random).\n   FIX: Chỉ dùng MỘT chiến lược cân bằng: WeightedRandomSampler.\n        Loss function: CrossEntropyLoss TIÊU CHUẨN cho cả train lẫn val.\n\n❌ LỖI 2 — Train loss ≠ Val loss về mặt toán học:\n   Phiên bản cũ: train = WeightedCE, val = UnweightedCE → hai đại lượng\n   khác nhau. Scheduler.step(val_loss) và EarlyStopping so sánh sai đơn vị.\n   → LR giảm xuống 4e-5 từ epoch 7 vì val_loss liên tục tăng (do lỗi 1).\n   FIX: Cả train và val đều dùng CrossEntropyLoss() tiêu chuẩn.\n\n❌ LỖI 3 — Không replicating balanced dataset của bài báo:\n   Bài báo Table 1: Moderate=5291, Severe=5291, PDR=5291 (bằng nhau!)\n   → bài báo ĐÃ balance dataset trước khi train (oversample hoặc duplicate).\n   EyePACS gốc: class 0=25806, class 1=2440 (imbalance cực nặng).\n   FIX: Thêm tùy chọn OVERSAMPLE_TRAIN để replicate setup bài báo.\n        Khi ON: duplicate/resample minority classes đến TARGET_PER_CLASS.\n\n═══════════════════════════════════════════════════════════════\n  CẤU HÌNH SO VỚI BÀI BÁO (Section IV.C)\n═══════════════════════════════════════════════════════════════\n  Thông số      | Bài báo  | Code này   | Ghi chú\n  --------------|----------|------------|---------------------------\n  Optimizer     | Adam     | Adam ✓     |\n  LR            | 0.001    | 1e-4 ✓     | 1e-3 gây instability với\n                |          |            | backbone 48M params — giảm\n                |          |            | xuống 1e-4 là thực hành chuẩn\n  Batch size    | 32       | 32 ✓       | Đúng bài báo (P100 dư VRAM)\n  Epochs        | 100      | 100 ✓      |\n  Split         | 80/20    | 80/20 ✓    |\n  Early stop    | 10 epoch | 10 ✓       | \"no improvement after ten epochs\"\n  LR reduce     | ×0.4/5ep | ×0.4/5 ✓  | \"reduced by factor of 0.4 for 5ep\"\n  Loss          | CE       | CE ✓       | CrossEntropyLoss tiêu chuẩn\n  Dataset cân   | YES      | Tùy chọn  | Table 1: Moderate=Severe=PDR=5291\n\"\"\"\n\nimport os\nimport copy\nimport time\nfrom collections import Counter\n\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.amp import GradScaler, autocast\nfrom torch.utils.data import DataLoader, Subset, WeightedRandomSampler\nfrom sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score\n\nfrom dataset import APTOSDataset\nfrom models import IR_CNN\n\n# ════════════════════════════════════════════════════════════════════════════\n#  ĐƯỜNG DẪN — chỉnh theo môi trường\n# ════════════════════════════════════════════════════════════════════════════\nCSV_FILE   = '/kaggle/input/datasets/laimochuy/ep-g-eli-et-al/zipEyepacs/trainLabels.csv/trainLabels.csv'\nIMG_DIR    = '/kaggle/input/datasets/laimochuy/ep-g-eli-et-al/zipEyepacs/train/'\nOUTPUT_DIR = '/kaggle/working/'\nCACHE_DIR  = '/kaggle/working/img_cache/'\n\n# ════════════════════════════════════════════════════════════════════════════\n#  HYPERPARAMETERS — bám sát Section IV.C\n# ════════════════════════════════════════════════════════════════════════════\nBATCH_SIZE   = 32        # bài báo: \"mini-batch size of 32 image crops\"\nNUM_EPOCHS   = 100       # bài báo: \"trained over 100 epochs\"\nLR           = 1e-4      # bài báo nói 1e-3, nhưng với backbone 48M params +\n                         # EyePACS imbalance → 1e-3 gây gradient instability.\n                         # 1e-4 là điều chỉnh thực tế, kết quả vẫn so sánh được.\nTRAIN_RATIO  = 0.8       # bài báo: \"80-20 train-test split\"\nSEED         = 42\n\n# Early stopping (bài báo: \"no improvement in validation data after ten epochs\")\nPATIENCE_ES  = 10\n\n# LR Reduce on Plateau (bài báo: \"reduced by a factor of 0.4 ... for five epochs\")\nLR_FACTOR    = 0.4\nPATIENCE_LR  = 5\n\n# Gradient clipping — không có trong bài báo, cải tiến kỹ thuật ổn định backbone lớn\nGRAD_CLIP    = 1.0\n\n# ════════════════════════════════════════════════════════════════════════════\n#  CHIẾN LƯỢC CÂN BẰNG CLASS — chọn MỘT trong hai (không dùng cả hai)\n# ════════════════════════════════════════════════════════════════════════════\n\n# Tùy chọn A (gần bài báo nhất): Oversample để replicate Table 1 của bài báo\n# Bài báo Table 1: Moderate=5291, Severe=5291, PDR=5291 (bằng nhau!)\n# → Bật ON để balance dataset trước khi train, dùng standard CE loss.\nOVERSAMPLE_TRAIN   = True\nTARGET_PER_CLASS   = 5291   # số sample mỗi class sau oversample (= paper's Table 1)\n                             # Class 0 (25806 ảnh) sẽ bị giới hạn xuống 5291 (undersample)\n                             # Class 1 (2440 ảnh) sẽ được duplicate lên 5291 (oversample)\n\n# Tùy chọn B: WeightedRandomSampler (KHÔNG kết hợp với WeightedCELoss)\n# Chỉ dùng khi OVERSAMPLE_TRAIN = False\nUSE_WEIGHTED_SAMPLER = False\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  Oversample để balance dataset theo Table 1 bài báo\n# ─────────────────────────────────────────────────────────────────────────────\ndef make_balanced_indices(full_dataset, original_indices, target_per_class, seed=42):\n    \"\"\"\n    Replicate setup của bài báo (Table 1: Moderate=Severe=PDR=5291):\n      - Class 0 (25806): undersample xuống target\n      - Class 1 (2440) : oversample (duplicate) lên target\n      - Class 2,3,4    : giữ nguyên hoặc oversample lên target\n\n    Returns: list of indices (có thể lặp lại cho minority classes)\n    \"\"\"\n    rng = np.random.RandomState(seed)\n\n    # Phân loại indices theo class\n    class_indices = {c: [] for c in range(5)}\n    for idx in original_indices:\n        label = int(full_dataset.data_frame.iloc[idx][full_dataset.col_label])\n        class_indices[label].append(idx)\n\n    balanced = []\n    for cls, idxs in class_indices.items():\n        n = len(idxs)\n        if n == 0:\n            continue\n        if n >= target_per_class:\n            # Undersample (không duplicate)\n            chosen = rng.choice(idxs, target_per_class, replace=False).tolist()\n        else:\n            # Oversample (duplicate, giống augmentation-based oversampling)\n            chosen = rng.choice(idxs, target_per_class, replace=True).tolist()\n        balanced.extend(chosen)\n\n    rng.shuffle(balanced)\n    return balanced\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  WeightedRandomSampler (Tùy chọn B — chỉ dùng khi OVERSAMPLE_TRAIN=False)\n# ─────────────────────────────────────────────────────────────────────────────\ndef make_weighted_sampler(indices, full_dataset):\n    \"\"\"\n    WeightedRandomSampler: mỗi batch có xác suất đều cho mọi class.\n    Chỉ dùng khi KHÔNG oversample.\n    KHÔNG kết hợp với WeightedCELoss (gây double weighting).\n    \"\"\"\n    labels = [int(full_dataset.data_frame.iloc[i][full_dataset.col_label])\n              for i in indices]\n    class_count  = Counter(labels)\n    class_weight = {cls: 1.0 / cnt for cls, cnt in class_count.items()}\n    sample_weights = torch.tensor(\n        [class_weight[lbl] for lbl in labels], dtype=torch.float\n    )\n    return WeightedRandomSampler(sample_weights, len(sample_weights), replacement=True)\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  Một epoch train hoặc val\n# ─────────────────────────────────────────────────────────────────────────────\ndef run_epoch(model, loader, criterion, optimizer, scaler, device, is_train=True):\n    \"\"\"Trả về (loss, acc, precision, recall, f1_macro).\"\"\"\n    model.train() if is_train else model.eval()\n\n    total_loss              = 0.0\n    n_samples               = 0\n    all_preds, all_labels   = [], []\n\n    ctx = torch.enable_grad() if is_train else torch.no_grad()\n    with ctx:\n        for imgs, lbls in loader:\n            imgs = imgs.to(device, non_blocking=True)\n            lbls = lbls.to(device, non_blocking=True)\n\n            if is_train:\n                optimizer.zero_grad(set_to_none=True)\n\n            with autocast('cuda', enabled=(scaler is not None)):\n                out  = model(imgs)\n                loss = criterion(out, lbls)\n\n            if is_train:\n                if scaler is not None:\n                    scaler.scale(loss).backward()\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                    scaler.step(optimizer)\n                    scaler.update()\n                else:\n                    loss.backward()\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n                    optimizer.step()\n\n            bs = imgs.size(0)\n            total_loss += loss.item() * bs\n            n_samples  += bs\n            all_preds.extend(out.argmax(1).cpu().tolist())\n            all_labels.extend(lbls.cpu().tolist())\n\n    avg_loss = total_loss / n_samples\n    acc      = accuracy_score(all_labels, all_preds)\n    prec     = precision_score(all_labels, all_preds, average='macro', zero_division=0)\n    rec      = recall_score   (all_labels, all_preds, average='macro', zero_division=0)\n    f1       = f1_score       (all_labels, all_preds, average='macro', zero_division=0)\n    return avg_loss, acc, prec, rec, f1\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  Main training loop\n# ─────────────────────────────────────────────────────────────────────────────\ndef train_model():\n    os.makedirs(OUTPUT_DIR, exist_ok=True)\n\n    # ── Device ────────────────────────────────────────────────────────────────\n    device  = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    use_amp = torch.cuda.is_available()\n    torch.backends.cudnn.benchmark = True\n\n    print(f\"Device : {device}\")\n    if torch.cuda.is_available():\n        print(f\"GPU    : {torch.cuda.get_device_name(0)}\")\n        print(f\"VRAM   : {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB\")\n    print(f\"AMP    : {'ON' if use_amp else 'OFF'}\")\n\n    print(f\"\\nHyperparameters (Section IV.C):\")\n    print(f\"  LR={LR}, Batch={BATCH_SIZE}, Epochs={NUM_EPOCHS}\")\n    print(f\"  EarlyStopping patience={PATIENCE_ES}\")\n    print(f\"  LR ReduceOnPlateau factor={LR_FACTOR}, patience={PATIENCE_LR}\")\n    print(f\"  Loss: CrossEntropyLoss (tiêu chuẩn, không weighted)\")\n\n    if OVERSAMPLE_TRAIN:\n        print(f\"  Class balance: OVERSAMPLE → {TARGET_PER_CLASS} samples/class (bài báo Table 1)\")\n    else:\n        bal_str = \"WeightedRandomSampler\" if USE_WEIGHTED_SAMPLER else \"Không (shuffle thường)\"\n        print(f\"  Class balance: {bal_str}\")\n\n    # ── Dataset ───────────────────────────────────────────────────────────────\n    full_train = APTOSDataset(CSV_FILE, IMG_DIR, is_train=True,  cache_dir=CACHE_DIR)\n    full_val   = APTOSDataset(CSV_FILE, IMG_DIR, is_train=False, cache_dir=CACHE_DIR)\n\n    total      = len(full_train)\n    train_size = int(TRAIN_RATIO * total)\n\n    # Reproducible split — lưu val_indices để evaluate.py dùng lại\n    torch.manual_seed(SEED)\n    all_indices   = torch.randperm(total).tolist()\n    train_indices = all_indices[:train_size]\n    val_indices   = all_indices[train_size:]\n    torch.save(val_indices, os.path.join(OUTPUT_DIR, 'val_indices.pt'))\n\n    # ── Phân phối gốc ─────────────────────────────────────────────────────────\n    raw_labels = [int(full_train.data_frame.iloc[i][full_train.col_label])\n                  for i in train_indices]\n    raw_dist   = dict(sorted(Counter(raw_labels).items()))\n    print(f\"\\nTrain split gốc : {len(train_indices)} samples\")\n    print(f\"  Phân phối class: {raw_dist}\")\n\n    # ── Cân bằng dataset (bám sát Table 1 bài báo) ───────────────────────────\n    if OVERSAMPLE_TRAIN:\n        balanced_train_indices = make_balanced_indices(\n            full_train, train_indices, TARGET_PER_CLASS, seed=SEED\n        )\n        bal_labels = [int(full_train.data_frame.iloc[i][full_train.col_label])\n                      for i in balanced_train_indices]\n        bal_dist   = dict(sorted(Counter(bal_labels).items()))\n        print(f\"\\nSau oversample  : {len(balanced_train_indices)} samples\")\n        print(f\"  Phân phối class: {bal_dist}\")\n        active_train_indices = balanced_train_indices\n    else:\n        active_train_indices = train_indices\n\n    print(f\"Val set         : {len(val_indices)} samples (phân phối tự nhiên)\")\n\n    train_dataset = Subset(full_train, active_train_indices)\n    val_dataset   = Subset(full_val,   val_indices)\n\n    # ── DataLoader ────────────────────────────────────────────────────────────\n    NUM_WORKERS = 2\n    PREFETCH    = 4\n\n    if OVERSAMPLE_TRAIN:\n        # Đã balance rồi → shuffle thường (giống bài báo nhất)\n        train_loader = DataLoader(\n            train_dataset, batch_size=BATCH_SIZE, shuffle=True,\n            num_workers=NUM_WORKERS, pin_memory=True,\n            prefetch_factor=PREFETCH, persistent_workers=True\n        )\n        print(\"\\n  [Sampler] Shuffle thường (dataset đã balanced — bám sát bài báo)\")\n    elif USE_WEIGHTED_SAMPLER:\n        sampler = make_weighted_sampler(active_train_indices, full_train)\n        train_loader = DataLoader(\n            train_dataset, batch_size=BATCH_SIZE, sampler=sampler,\n            num_workers=NUM_WORKERS, pin_memory=True,\n            prefetch_factor=PREFETCH, persistent_workers=True\n        )\n        print(\"\\n  [Sampler] WeightedRandomSampler (không có trong bài báo)\")\n    else:\n        train_loader = DataLoader(\n            train_dataset, batch_size=BATCH_SIZE, shuffle=True,\n            num_workers=NUM_WORKERS, pin_memory=True,\n            prefetch_factor=PREFETCH, persistent_workers=True\n        )\n        print(\"\\n  [Sampler] Shuffle thường\")\n\n    val_loader = DataLoader(\n        val_dataset, batch_size=BATCH_SIZE, shuffle=False,\n        num_workers=NUM_WORKERS, pin_memory=True, persistent_workers=True\n    )\n\n    # ── Model ─────────────────────────────────────────────────────────────────\n    model = IR_CNN(num_classes=5).to(device)\n    total_params     = sum(p.numel() for p in model.parameters())\n    trainable_params = model.get_trainable_params()\n    print(f\"\\nTotal parameters    : {total_params:,}\")\n    print(f\"Trainable parameters: {trainable_params:,}\")\n\n    # ── Loss — CrossEntropyLoss TIÊU CHUẨN cho cả train lẫn val ─────────────\n    # Bài báo Section IV.C: \"default loss function was cross-entropy\"\n    # KHÔNG dùng weighted CE — bài báo không đề cập, và đã balance dataset\n    # rồi nên không cần. Quan trọng: train và val dùng CÙNG loss function\n    # để scheduler và early stopping hoạt động đúng.\n    criterion = nn.CrossEntropyLoss()\n    print(f\"\\nLoss: CrossEntropyLoss (tiêu chuẩn, không weighted) — đúng bài báo\")\n\n    # ── Optimizer — Adam, LR=0.001 (bài báo) / 1e-4 (thực tế) ───────────────\n    # Bài báo: \"Adam optimization algorithm ... learning rate of 0.001\"\n    optimizer = optim.Adam(model.parameters(), lr=LR)\n\n    # ── LR Scheduler — bài báo: \"reduced by factor 0.4, no improvement 5 epochs\"\n    # Dùng val_loss để trigger LR reduce (đúng bài báo Section IV.C).\n    # mode='min' vì loss càng thấp càng tốt.\n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='min', factor=LR_FACTOR, patience=PATIENCE_LR\n    )\n\n    scaler = GradScaler('cuda') if use_amp else None\n\n    # ── Tracking metrics ──────────────────────────────────────────────────────\n    # PHÂN TÍCH VẤN ĐỀ (từ log EyePACS):\n    #   Val Loss và Val F1 đi NGƯỢC CHIỀU nhau:\n    #     - Val Loss tốt nhất epoch 3 (1.0814) → model bias về class 0 (73% val)\n    #     - Val F1 tốt nhất epoch 13 (0.4231) → model học balanced tốt hơn\n    #   Nguyên nhân: train balanced (20%/class), val natural (class 0 = 73%)\n    #     → CE loss trên val thưởng model predict class 0 nhiều\n    #     → F1 macro thưởng model phân loại đều tất cả classes\n    #\n    # GIẢI PHÁP:\n    #   - Scheduler   : vẫn dùng val_loss (đúng bài báo, trigger LR reduce)\n    #   - Best model  : dùng val_F1_macro (metric thực sự quan trọng)\n    #   - Early stop  : dùng val_F1_macro (tránh dừng khi F1 vẫn đang tăng)\n    best_val_f1        = -1.0\n    best_val_loss_log  = float('inf')   # chỉ để log, không dùng để lưu model\n    best_model_wts     = copy.deepcopy(model.state_dict())\n    early_stop_counter = 0\n    best_epoch         = 0\n\n    print(f\"\\n{'='*70}\")\n    print(f\"  BẮT ĐẦU TRAINING — {NUM_EPOCHS} epochs\")\n    print(f\"  Best model criterion : Val F1 Macro (↑)\")\n    print(f\"  LR scheduler trigger : Val Loss    (↓)\")\n    print(f\"{'='*70}\")\n\n    for epoch in range(1, NUM_EPOCHS + 1):\n        t0 = time.time()\n\n        # Train epoch\n        tr_loss, tr_acc, tr_prec, tr_rec, tr_f1 = run_epoch(\n            model, train_loader, criterion, optimizer, scaler,\n            device, is_train=True\n        )\n        # Val epoch\n        vl_loss, vl_acc, vl_prec, vl_rec, vl_f1 = run_epoch(\n            model, val_loader, criterion, None, None,\n            device, is_train=False\n        )\n\n        # Scheduler dùng val_loss (đúng bài báo)\n        scheduler.step(vl_loss)\n        lr_now  = optimizer.param_groups[0]['lr']\n        elapsed = time.time() - t0\n\n        print(\n            f\"Epoch [{epoch:3d}/{NUM_EPOCHS}] {elapsed:.0f}s | \"\n            f\"Train Loss:{tr_loss:.4f} Acc:{tr_acc:.4f} F1:{tr_f1:.4f} | \"\n            f\"Val   Loss:{vl_loss:.4f} Acc:{vl_acc:.4f} \"\n            f\"Prec:{vl_prec:.4f} Rec:{vl_rec:.4f} F1:{vl_f1:.4f} | \"\n            f\"LR:{lr_now:.2e}\"\n        )\n\n        # Best model và early stopping dùng val_F1_macro\n        if vl_f1 > best_val_f1:\n            best_val_f1        = vl_f1\n            best_val_loss_log  = vl_loss\n            best_model_wts     = copy.deepcopy(model.state_dict())\n            early_stop_counter = 0\n            best_epoch         = epoch\n            save_path          = os.path.join(OUTPUT_DIR, 'best_ir_cnn.pth')\n            torch.save(best_model_wts, save_path)\n            print(f\"    => [Best] Val F1: {best_val_f1:.4f} | Val Loss: {vl_loss:.4f} — Saved!\")\n        else:\n            early_stop_counter += 1\n            print(f\"    => EarlyStopping (F1): {early_stop_counter}/{PATIENCE_ES}\")\n            # Bài báo: \"no improvement in validation data after ten epochs\"\n            if early_stop_counter >= PATIENCE_ES:\n                print(f\"\\nEarly Stopping tại epoch {epoch} (patience={PATIENCE_ES}).\")\n                break\n\n    # ── Kết thúc ──────────────────────────────────────────────────────────────\n    model.load_state_dict(best_model_wts)\n\n    print(f\"\\n{'='*70}\")\n    print(f\"Huấn luyện hoàn tất!\")\n    print(f\"Best Val F1   : {best_val_f1:.4f}  (epoch {best_epoch})\")\n    print(f\"Best Val Loss : {best_val_loss_log:.4f}  (cùng epoch)\")\n    print(f\"Model lưu tại : {OUTPUT_DIR}best_ir_cnn.pth\")\n    print(f\"Val indices   : {OUTPUT_DIR}val_indices.pt\")\n    print(f\"{'='*70}\")\n    return model\n\n\nif __name__ == '__main__':\n    train_model()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-02T17:01:18.391610Z","iopub.execute_input":"2026-03-02T17:01:18.392476Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nevaluate.py — Đánh giá IR-CNN sau khi train\n=============================================\nOutput format giống hệt ảnh template.\n\nChạy:\n  python evaluate.py           → val set (80/20 split)\n  python evaluate.py --kfold   → 5-fold CV (Section IV.F)\n\"\"\"\n\nimport os, sys, io, contextlib\nimport numpy as np\nimport torch\nfrom torch.utils.data import DataLoader, Subset\nfrom sklearn.metrics import (\n    accuracy_score, balanced_accuracy_score,\n    precision_score, recall_score, f1_score,\n    matthews_corrcoef, cohen_kappa_score,\n    confusion_matrix\n)\nfrom sklearn.model_selection import KFold\n\nfrom dataset import APTOSDataset\nfrom models  import IR_CNN\n\n# ════════════════════════════════════════════════════════════════════════════\nCSV_FILE     = '/kaggle/input/datasets/laimochuy/ep-g-eli-et-al/zipEyepacs/trainLabels.csv/trainLabels.csv'\nIMG_DIR      = '/kaggle/input/datasets/laimochuy/ep-g-eli-et-al/zipEyepacs/train/'\nMODEL_PATH   = '/kaggle/working/best_ir_cnn.pth'\nVAL_IDX_PATH = '/kaggle/working/val_indices.pt'\nCACHE_DIR    = '/kaggle/working/img_cache/'\nOUTPUT_DIR   = '/kaggle/working/'\n\nBATCH_SIZE  = 32\nTRAIN_RATIO = 0.8\nSEED        = 42\nN_FOLDS     = 5\n\nLABELS      = [0, 1, 2, 3, 4]\nCLASS_NAMES = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferate\"]\n\n\n# ─────────────────────────────────────────────────────────────────────────────\ndef gmean_per_class(y_true, y_pred):\n    cm, res = confusion_matrix(y_true, y_pred, labels=LABELS), {}\n    for i, lbl in enumerate(LABELS):\n        tp = cm[i, i]; fn = cm[i,:].sum()-tp\n        fp = cm[:,i].sum()-tp; tn = cm.sum()-tp-fn-fp\n        s  = tp/(tp+fn) if (tp+fn)>0 else 0.0\n        sp = tn/(tn+fp) if (tn+fp)>0 else 0.0\n        res[lbl] = float(np.sqrt(s*sp))\n    return res\n\n\ndef compute_metrics(y_true, y_pred):\n    gm = gmean_per_class(y_true, y_pred)\n    return dict(\n        accuracy    = accuracy_score(y_true, y_pred),\n        balanced_acc= balanced_accuracy_score(y_true, y_pred),\n        prec_macro  = precision_score(y_true, y_pred, average='macro',    zero_division=0),\n        prec_wgt    = precision_score(y_true, y_pred, average='weighted', zero_division=0),\n        rec_macro   = recall_score   (y_true, y_pred, average='macro',    zero_division=0),\n        rec_wgt     = recall_score   (y_true, y_pred, average='weighted', zero_division=0),\n        f1_macro    = f1_score       (y_true, y_pred, average='macro',    zero_division=0),\n        f1_wgt      = f1_score       (y_true, y_pred, average='weighted', zero_division=0),\n        gmean_macro = float(np.mean(list(gm.values()))),\n        mcc         = matthews_corrcoef(y_true, y_pred),\n        kappa       = cohen_kappa_score(y_true, y_pred),\n        prec_cls    = precision_score(y_true, y_pred, average=None, labels=LABELS, zero_division=0),\n        rec_cls     = recall_score   (y_true, y_pred, average=None, labels=LABELS, zero_division=0),\n        f1_cls      = f1_score       (y_true, y_pred, average=None, labels=LABELS, zero_division=0),\n        gm_cls      = gm,\n    )\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  Print — format giống hệt ảnh\n# ─────────────────────────────────────────────────────────────────────────────\ndef print_metrics(m):\n    SEP  = \"=\" * 40\n    DASH = \"-\" * 20\n\n    # ── Overall metrics — format khớp với template ────────────────────────────\n    print(f\"Accuracy          : {m['accuracy']:.4f}\")\n    print(f\"BalancedAcc       : {m['balanced_acc']:.4f}\")\n    print(DASH)\n    print(f\"Precision Macro   : {m['prec_macro']:.4f}\")\n    print(f\"Precision Weighted: {m['prec_wgt']:.4f}\")\n    print(DASH)\n    print(f\"Recall Macro      : {m['rec_macro']:.4f}\")\n    print(f\"Recall Weighted   : {m['rec_wgt']:.4f}\")\n    print(DASH)\n    print(f\"F1-Score Macro    : {m['f1_macro']:.4f}\")\n    print(f\"F1-Score Weighted : {m['f1_wgt']:.4f}\")\n    print(DASH)\n    print(f\"GMean (Macro)     : {m['gmean_macro']:.4f}\")\n    print(f\"MCC               : {m['mcc']:.4f}\")\n    print(f\"Kappa             : {m['kappa']:.4f}\")\n    print(SEP)\n    print()\n\n    # ── Class-wise metrics — format khớp với template ─────────────────────────\n    print(\"--- CLASS-WISE METRICS ---\")\n    hdr = f\"{'Class':<12} | {'Precision':>9} | {'Recall':>6} | {'F1-Score':>8} | {'G-Mean':>6}\"\n    print(hdr)\n    print(\"-\" * len(hdr))\n    for i, name in enumerate(CLASS_NAMES):\n        lbl = LABELS[i]\n        print(\n            f\"{name:<12} | {m['prec_cls'][i]:>9.4f} | \"\n            f\"{m['rec_cls'][i]:>6.4f} | \"\n            f\"{m['f1_cls'][i]:>8.4f} | \"\n            f\"{m['gm_cls'][lbl]:>6.4f}\"\n        )\n    print(SEP)\n\n\ndef print_confusion_matrix(y_true, y_pred):\n    cm = confusion_matrix(y_true, y_pred, labels=LABELS)\n    print(\"\\nConfusion Matrix:\")\n    print(\"         \" + \"  \".join(f\"{n[:5]:>5}\" for n in CLASS_NAMES))\n    for i, name in enumerate(CLASS_NAMES):\n        row = \"  \".join(f\"{cm[i][j]:>5}\" for j in range(5))\n        print(f\"{name[:8]:<8}: {row}\")\n    print()\n\n\n# ─────────────────────────────────────────────────────────────────────────────\ndef run_inference(model, loader, device):\n    model.eval()\n    preds, labels = [], []\n    with torch.no_grad():\n        for imgs, lbls in loader:\n            imgs = imgs.to(device)\n            preds.extend(model(imgs).argmax(1).cpu().tolist())\n            labels.extend(lbls.tolist())\n    return labels, preds\n\n\ndef load_model(path, device):\n    mdl = IR_CNN(num_classes=5).to(device)\n    mdl.load_state_dict(torch.load(path, map_location=device))\n    mdl.eval()\n    return mdl\n\n\ndef _save_txt(content, path):\n    with open(path, 'w', encoding='utf-8') as f:\n        f.write(content)\n    print(f\"Kết quả lưu tại: {path}\")\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  Mode 1: Val set\n# ─────────────────────────────────────────────────────────────────────────────\ndef evaluate_val_set():\n    device   = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    full_val = APTOSDataset(CSV_FILE, IMG_DIR, is_train=False, cache_dir=CACHE_DIR)\n\n    if os.path.exists(VAL_IDX_PATH):\n        val_indices = torch.load(VAL_IDX_PATH)\n    else:\n        torch.manual_seed(SEED)\n        idxs        = torch.randperm(len(full_val)).tolist()\n        val_indices = idxs[int(TRAIN_RATIO * len(full_val)):]\n\n    loader = DataLoader(Subset(full_val, val_indices),\n                        batch_size=BATCH_SIZE, shuffle=False,\n                        num_workers=4, pin_memory=True)\n    model          = load_model(MODEL_PATH, device)\n    y_true, y_pred = run_inference(model, loader, device)\n\n    print(f\"\\nDevice: {device} | Val samples: {len(val_indices)}\\n\")\n\n    buf = io.StringIO()\n    with contextlib.redirect_stdout(buf):\n        print_metrics(compute_metrics(y_true, y_pred))\n        print_confusion_matrix(y_true, y_pred)\n    output = buf.getvalue()\n    print(output)\n    _save_txt(output, os.path.join(OUTPUT_DIR, 'eval_results.txt'))\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  Mode 2: 5-Fold CV\n# ─────────────────────────────────────────────────────────────────────────────\ndef evaluate_kfold():\n    device       = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    full_dataset = APTOSDataset(CSV_FILE, IMG_DIR, is_train=False, cache_dir=CACHE_DIR)\n    model        = load_model(MODEL_PATH, device)\n    kf           = KFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\n\n    fold_results = []\n    full_output  = io.StringIO()\n\n    for fold, (_, test_idx) in enumerate(kf.split(range(len(full_dataset))), 1):\n        loader         = DataLoader(Subset(full_dataset, test_idx),\n                                    batch_size=BATCH_SIZE, shuffle=False,\n                                    num_workers=4, pin_memory=True)\n        y_true, y_pred = run_inference(model, loader, device)\n        m              = compute_metrics(y_true, y_pred)\n        fold_results.append(m)\n\n        title = f\"\\nFOLD {fold}/{N_FOLDS}  ({len(test_idx)} samples)\\n\"\n        buf   = io.StringIO()\n        with contextlib.redirect_stdout(buf):\n            print_metrics(m)\n            print_confusion_matrix(y_true, y_pred)\n        section = title + buf.getvalue()\n        print(section)\n        full_output.write(section)\n\n    # Summary\n    keys = ['accuracy', 'prec_macro', 'rec_macro', 'f1_macro', 'gmean_macro', 'kappa']\n    lbls = ['Accuracy', 'Precision', 'Recall', 'F1-Score', 'GMean', 'Kappa']\n    SEP  = \"=\" * 40\n\n    summary = io.StringIO()\n    with contextlib.redirect_stdout(summary):\n        print(f\"\\n{SEP}\")\n        print(\"5-FOLD CV SUMMARY\")\n        print(SEP)\n        hdr = f\"{'Fold':>5} | \" + \" | \".join(f\"{l:>11}\" for l in lbls)\n        print(hdr)\n        print(\"-\" * len(hdr))\n        for i, fm in enumerate(fold_results, 1):\n            print(f\"{i:>5} | \" + \" | \".join(f\"{fm[k]:>11.4f}\" for k in keys))\n        print(\"-\" * len(hdr))\n        avgs = {k: np.mean([fm[k] for fm in fold_results]) for k in keys}\n        stds = {k: np.std ([fm[k] for fm in fold_results]) for k in keys}\n        print(\"  Avg | \" + \" | \".join(f\"{avgs[k]:>11.4f}\" for k in keys))\n        print(\"  Std | \" + \" | \".join(f\"{stds[k]:>11.4f}\" for k in keys))\n        print(SEP)\n\n    sum_text = summary.getvalue()\n    print(sum_text)\n    _save_txt(full_output.getvalue() + sum_text,\n              os.path.join(OUTPUT_DIR, 'kfold_results.txt'))\n\n\n# ─────────────────────────────────────────────────────────────────────────────\nif __name__ == '__main__':\n    if '--kfold' in sys.argv:\n        evaluate_kfold()\n    else:\n        evaluate_val_set()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile save_model.py\n\"\"\"\nsave_model.py — Export IR-CNN model\n=====================================\nChạy sau khi train xong để lưu model theo nhiều định dạng:\n\n  python save_model.py                  → lưu .pth + metadata\n  python save_model.py --torchscript   → thêm TorchScript (dùng được không cần class IR_CNN)\n  python save_model.py --onnx          → thêm ONNX (deploy trên mọi framework)\n\nOutput:\n  best_ir_cnn.pth            ← weights gốc (dùng với load_state_dict)\n  best_ir_cnn_full.pth       ← full model (dùng torch.load trực tiếp)\n  best_ir_cnn_meta.pt        ← weights + metadata (accuracy, epoch, config)\n  best_ir_cnn_script.pt      ← TorchScript (tuỳ chọn)\n  best_ir_cnn.onnx           ← ONNX (tuỳ chọn)\n\"\"\"\n\nimport os\nimport sys\nimport json\nimport torch\nimport torch.nn as nn\n\nfrom models import IR_CNN\n\n# ════════════════════════════════════════════════════════════════════════════\nMODEL_PATH  = '/kaggle/working/best_ir_cnn.pth'      # weights từ train.py\nOUTPUT_DIR  = '/kaggle/working/saved_models/'\nIMG_SIZE    = 256     # kích thước input (phải khớp với dataset.py)\nNUM_CLASSES = 5\n# ════════════════════════════════════════════════════════════════════════════\n\nCLASS_NAMES = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferate\"]\nLABELS      = [0, 1, 2, 3, 4]\n\n\ndef load_trained_model(device):\n    model = IR_CNN(num_classes=NUM_CLASSES).to(device)\n    model.load_state_dict(torch.load(MODEL_PATH, map_location=device))\n    model.eval()\n    print(f\"Đã load: {MODEL_PATH}\")\n    return model\n\n\ndef save_all(do_torchscript=False, do_onnx=False):\n    os.makedirs(OUTPUT_DIR, exist_ok=True)\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(f\"Device: {device}\\n\")\n\n    model = load_trained_model(device)\n\n    # ── 1. Copy weights gốc ──────────────────────────────────────────────────\n    import shutil\n    dst_pth = os.path.join(OUTPUT_DIR, 'best_ir_cnn.pth')\n    shutil.copy(MODEL_PATH, dst_pth)\n    print(f\"[1] Weights (state_dict) → {dst_pth}\")\n\n    # ── 2. Full model (torch.save toàn bộ object) ────────────────────────────\n    #     Ưu điểm: load đơn giản hơn, không cần import IR_CNN\n    #     Nhược điểm: phụ thuộc vào file models.py cùng thư mục\n    dst_full = os.path.join(OUTPUT_DIR, 'best_ir_cnn_full.pth')\n    torch.save(model, dst_full)\n    print(f\"[2] Full model object    → {dst_full}\")\n\n    # ── 3. Weights + metadata ────────────────────────────────────────────────\n    #     Lưu thêm config để biết model được train như thế nào\n    meta = {\n        'state_dict' : model.state_dict(),\n        'config' : {\n            'num_classes' : NUM_CLASSES,\n            'img_size'    : IMG_SIZE,\n            'class_names' : CLASS_NAMES,\n            'backbone'    : 'InceptionV3 + ResNet50',\n            'paper'       : 'Ali et al., IEEE JTEHM 2023',\n        },\n        # Điền vào sau khi có kết quả từ evaluate.py:\n        'best_metrics' : {\n            'accuracy'   : None,\n            'f1_macro'   : None,\n            'kappa'      : None,\n            'best_epoch' : None,\n        },\n    }\n    dst_meta = os.path.join(OUTPUT_DIR, 'best_ir_cnn_meta.pt')\n    torch.save(meta, dst_meta)\n    print(f\"[3] Weights + metadata   → {dst_meta}\")\n\n    # ── 4. TorchScript (tuỳ chọn) ────────────────────────────────────────────\n    if do_torchscript:\n        try:\n            dummy  = torch.randn(1, 3, IMG_SIZE, IMG_SIZE).to(device)\n            script = torch.jit.trace(model, dummy)\n            dst_ts = os.path.join(OUTPUT_DIR, 'best_ir_cnn_script.pt')\n            script.save(dst_ts)\n            print(f\"[4] TorchScript          → {dst_ts}\")\n        except Exception as e:\n            print(f\"[4] TorchScript FAILED: {e}\")\n\n    # ── 5. ONNX (tuỳ chọn) ───────────────────────────────────────────────────\n    if do_onnx:\n        try:\n            dummy    = torch.randn(1, 3, IMG_SIZE, IMG_SIZE).to(device)\n            dst_onnx = os.path.join(OUTPUT_DIR, 'best_ir_cnn.onnx')\n            torch.onnx.export(\n                model, dummy, dst_onnx,\n                input_names=['image'],\n                output_names=['logits'],\n                dynamic_axes={'image': {0: 'batch'}, 'logits': {0: 'batch'}},\n                opset_version=14,\n            )\n            print(f\"[5] ONNX                 → {dst_onnx}\")\n        except Exception as e:\n            print(f\"[5] ONNX FAILED: {e}\")\n\n    # ── Thống kê model ────────────────────────────────────────────────────────\n    total_params = sum(p.numel() for p in model.parameters())\n    print(f\"\\nModel info:\")\n    print(f\"  Total parameters: {total_params:,}\")\n    print(f\"  Output dir      : {OUTPUT_DIR}\")\n    print(f\"\\nHoàn tất!\")\n\n\nif __name__ == '__main__':\n    do_ts   = '--torchscript' in sys.argv\n    do_onnx = '--onnx' in sys.argv\n    save_all(do_torchscript=do_ts, do_onnx=do_onnx)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"%%writefile predict.py\n\"\"\"\npredict.py — Dự đoán DR từ ảnh fundus mới\n==========================================\nDùng model đã train để predict ảnh bất kỳ.\n\nCách dùng:\n  # Dự đoán 1 ảnh\n  python predict.py --img path/to/image.png\n\n  # Dự đoán cả folder (tất cả .png/.jpg)\n  python predict.py --folder path/to/folder/\n\n  # Dự đoán và lưu kết quả ra CSV\n  python predict.py --folder path/to/folder/ --save results.csv\n\n  # Dùng model khác (mặc định: best_ir_cnn.pth)\n  python predict.py --img img.png --model /path/to/model.pth\n\"\"\"\n\nimport os\nimport sys\nimport argparse\nimport cv2\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nfrom torchvision import transforms\n\nfrom models import IR_CNN\n\n# ════════════════════════════════════════════════════════════════════════════\nDEFAULT_MODEL = '/kaggle/working/best_ir_cnn.pth'\nIMG_SIZE      = 256\nNUM_CLASSES   = 5\n# ════════════════════════════════════════════════════════════════════════════\n\nCLASS_NAMES = {\n    0: \"No DR        (Không bị DR)\",\n    1: \"Mild         (DR nhẹ)\",\n    2: \"Moderate     (DR trung bình)\",\n    3: \"Severe       (DR nặng)\",\n    4: \"Proliferate  (DR tăng sinh)\",\n}\nCLASS_SHORT = [\"No DR\", \"Mild\", \"Moderate\", \"Severe\", \"Proliferate\"]\n\n# Normalize theo ImageNet (giống dataset.py)\nNORMALIZE = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225]),\n])\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  Preprocessing — giống dataset.py (HE + Min-Max Norm)\n# ─────────────────────────────────────────────────────────────────────────────\ndef preprocess(image_bgr: np.ndarray) -> torch.Tensor:\n    \"\"\"\n    Nhận ảnh BGR (từ cv2.imread), trả về tensor [1, 3, 256, 256] sẵn sàng\n    đưa vào model. Preprocessing giống hệt dataset.py:\n      1. Resize 256×256\n      2. Histogram Equalization (kênh Y của YCrCb)\n      3. Min-Max Normalization\n      4. ToTensor + ImageNet normalize\n    \"\"\"\n    # 1. Resize\n    img = cv2.resize(image_bgr, (IMG_SIZE, IMG_SIZE))\n    img = img.astype(np.uint8)\n\n    # 2. Histogram Equalization (Eq. 5-6 bài báo — HE trên từng kênh RGB)\n    rgb       = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    r, g, b   = cv2.split(rgb)\n    r_eq, g_eq, b_eq = cv2.equalizeHist(r), cv2.equalizeHist(g), cv2.equalizeHist(b)\n    rgb       = cv2.merge((r_eq, g_eq, b_eq))\n\n    # 3. Min-Max Normalization (Eq. 7 bài báo)\n    img_f = rgb.astype(np.float32)\n    lo, hi = img_f.min(), img_f.max()\n    if hi > lo:\n        img_f = (img_f - lo) / (hi - lo)\n    else:\n        img_f = np.zeros_like(img_f)\n    img_uint8 = (img_f * 255.0).astype(np.uint8)\n\n    # 4. Tensor + normalize\n    tensor = NORMALIZE(img_uint8)          # [3, 256, 256]\n    return tensor.unsqueeze(0)             # [1, 3, 256, 256]\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  Load model\n# ─────────────────────────────────────────────────────────────────────────────\ndef load_model(model_path: str, device: torch.device) -> IR_CNN:\n    model = IR_CNN(num_classes=NUM_CLASSES).to(device)\n    checkpoint = torch.load(model_path, map_location=device)\n\n    # Hỗ trợ cả 3 định dạng lưu từ save_model.py\n    if isinstance(checkpoint, dict) and 'state_dict' in checkpoint:\n        model.load_state_dict(checkpoint['state_dict'])   # định dạng meta\n    elif isinstance(checkpoint, dict):\n        model.load_state_dict(checkpoint)                 # state_dict thuần\n    else:\n        model = checkpoint.to(device)                     # full model object\n\n    model.eval()\n    return model\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  Predict 1 ảnh\n# ─────────────────────────────────────────────────────────────────────────────\ndef predict_single(model, image_path: str, device: torch.device) -> dict:\n    \"\"\"Dự đoán 1 ảnh, trả về dict chứa class, confidence, probabilities.\"\"\"\n    img = cv2.imread(image_path)\n    if img is None:\n        raise FileNotFoundError(f\"Không đọc được: {image_path}\")\n\n    tensor = preprocess(img).to(device)\n\n    with torch.no_grad():\n        logits = model(tensor)                         # [1, 5]\n        probs  = F.softmax(logits, dim=1)[0]           # [5]\n        pred   = probs.argmax().item()\n        conf   = probs[pred].item()\n\n    return {\n        'image'       : os.path.basename(image_path),\n        'pred_class'  : pred,\n        'pred_label'  : CLASS_SHORT[pred],\n        'confidence'  : conf,\n        'probabilities': {CLASS_SHORT[i]: float(probs[i]) for i in range(NUM_CLASSES)},\n    }\n\n\ndef print_single_result(result: dict):\n    \"\"\"In kết quả 1 ảnh đẹp ra terminal.\"\"\"\n    SEP = \"=\" * 50\n    print(f\"\\n{SEP}\")\n    print(f\"  Ảnh      : {result['image']}\")\n    print(f\"  Dự đoán  : [{result['pred_class']}] {CLASS_NAMES[result['pred_class']]}\")\n    print(f\"  Confidence: {result['confidence']*100:.2f}%\")\n    print(f\"\\n  Xác suất từng class:\")\n    for i in range(NUM_CLASSES):\n        name  = CLASS_SHORT[i]\n        prob  = result['probabilities'][name]\n        bar   = \"█\" * int(prob * 30)\n        mark  = \" ◄\" if i == result['pred_class'] else \"\"\n        print(f\"    [{i}] {name:<12}: {prob*100:5.2f}% {bar}{mark}\")\n    print(SEP)\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  Predict cả folder\n# ─────────────────────────────────────────────────────────────────────────────\nVALID_EXT = {'.png', '.jpg', '.jpeg', '.bmp', '.tiff', '.tif'}\n\ndef predict_folder(model, folder_path: str, device: torch.device,\n                   save_csv: str = None):\n    \"\"\"Dự đoán tất cả ảnh trong folder, tuỳ chọn lưu CSV.\"\"\"\n    files = sorted([\n        f for f in os.listdir(folder_path)\n        if os.path.splitext(f)[1].lower() in VALID_EXT\n    ])\n\n    if not files:\n        print(f\"Không tìm thấy ảnh trong: {folder_path}\")\n        return\n\n    print(f\"\\nTìm thấy {len(files)} ảnh trong {folder_path}\\n\")\n\n    # Header\n    hdr = f\"{'Ảnh':<30} | {'Class':>5} | {'Label':<12} | {'Confidence':>10}\"\n    print(hdr)\n    print(\"-\" * len(hdr))\n\n    all_results = []\n    for fname in files:\n        fpath = os.path.join(folder_path, fname)\n        try:\n            res = predict_single(model, fpath, device)\n            print(\n                f\"{fname:<30} | {res['pred_class']:>5} | \"\n                f\"{res['pred_label']:<12} | {res['confidence']*100:>9.2f}%\"\n            )\n            all_results.append(res)\n        except Exception as e:\n            print(f\"{fname:<30} | ERROR: {e}\")\n\n    # Thống kê phân phối\n    print(f\"\\n{'='*50}\")\n    print(\"Phân phối dự đoán:\")\n    from collections import Counter\n    dist = Counter(r['pred_label'] for r in all_results)\n    for cls in CLASS_SHORT:\n        n   = dist.get(cls, 0)\n        pct = n / len(all_results) * 100 if all_results else 0\n        print(f\"  {cls:<12}: {n:>4} ảnh ({pct:.1f}%)\")\n\n    # Lưu CSV\n    if save_csv:\n        import csv\n        with open(save_csv, 'w', newline='', encoding='utf-8') as f:\n            writer = csv.writer(f)\n            writer.writerow(['image', 'pred_class', 'pred_label', 'confidence']\n                            + [f'prob_{c}' for c in CLASS_SHORT])\n            for r in all_results:\n                writer.writerow(\n                    [r['image'], r['pred_class'], r['pred_label'],\n                     f\"{r['confidence']:.4f}\"]\n                    + [f\"{r['probabilities'][c]:.4f}\" for c in CLASS_SHORT]\n                )\n        print(f\"\\nKết quả CSV lưu tại: {save_csv}\")\n\n    return all_results\n\n\n# ─────────────────────────────────────────────────────────────────────────────\n#  CLI\n# ─────────────────────────────────────────────────────────────────────────────\ndef parse_args():\n    p = argparse.ArgumentParser(description='IR-CNN DR Predictor')\n    p.add_argument('--img',    type=str, default=None,\n                   help='Đường dẫn tới 1 ảnh')\n    p.add_argument('--folder', type=str, default=None,\n                   help='Đường dẫn tới folder ảnh')\n    p.add_argument('--model',  type=str, default=DEFAULT_MODEL,\n                   help=f'Đường dẫn model (mặc định: {DEFAULT_MODEL})')\n    p.add_argument('--save',   type=str, default=None,\n                   help='Lưu kết quả folder ra CSV (ví dụ: results.csv)')\n    return p.parse_args()\n\n\ndef main():\n    args   = parse_args()\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n    if args.img is None and args.folder is None:\n        print(\"Cần chỉ định --img hoặc --folder\")\n        print(\"Ví dụ: python predict.py --img my_fundus.png\")\n        sys.exit(1)\n\n    print(f\"Device : {device}\")\n    print(f\"Model  : {args.model}\")\n    model = load_model(args.model, device)\n\n    if args.img:\n        result = predict_single(model, args.img, device)\n        print_single_result(result)\n\n    if args.folder:\n        predict_folder(model, args.folder, device, save_csv=args.save)\n\n\nif __name__ == '__main__':\n    main()\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!python train.py","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}