{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":19991,"databundleVersionId":1117522,"sourceType":"competition"},{"sourceId":11422468,"sourceType":"datasetVersion","datasetId":7137354},{"sourceId":234080702,"sourceType":"kernelVersion"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 1. Libaries 💻","metadata":{}},{"cell_type":"code","source":"!pip install -q efficientnet_pytorch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:08.336217Z","iopub.execute_input":"2025-04-15T15:15:08.336623Z","iopub.status.idle":"2025-04-15T15:15:14.574115Z","shell.execute_reply.started":"2025-04-15T15:15:08.336589Z","shell.execute_reply":"2025-04-15T15:15:14.573280Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install -q torchsampler","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:14.575197Z","iopub.execute_input":"2025-04-15T15:15:14.575487Z","iopub.status.idle":"2025-04-15T15:15:18.239647Z","shell.execute_reply.started":"2025-04-15T15:15:14.575464Z","shell.execute_reply":"2025-04-15T15:15:18.238708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import warnings\nwarnings.filterwarnings(\"ignore\", category=ResourceWarning)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:18.241571Z","iopub.execute_input":"2025-04-15T15:15:18.241878Z","iopub.status.idle":"2025-04-15T15:15:18.246583Z","shell.execute_reply.started":"2025-04-15T15:15:18.241840Z","shell.execute_reply":"2025-04-15T15:15:18.245530Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np \nimport pandas as pd\nfrom glob import glob\nfrom tqdm import tqdm\nimport os\nimport cv2\nimport random\nimport time\nimport re\nfrom datetime import datetime","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:18.249882Z","iopub.execute_input":"2025-04-15T15:15:18.250232Z","iopub.status.idle":"2025-04-15T15:15:18.852137Z","shell.execute_reply.started":"2025-04-15T15:15:18.250198Z","shell.execute_reply":"2025-04-15T15:15:18.851215Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torchvision.transforms as transforms\nfrom torchvision import datasets\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.data.sampler import SequentialSampler, RandomSampler\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchsampler import ImbalancedDatasetSampler\nimport timm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:18.853021Z","iopub.execute_input":"2025-04-15T15:15:18.853479Z","iopub.status.idle":"2025-04-15T15:15:27.478594Z","shell.execute_reply.started":"2025-04-15T15:15:18.853446Z","shell.execute_reply":"2025-04-15T15:15:27.477941Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import seaborn as sns\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nfrom skimage.feature import hog\nfrom sklearn import metrics\nfrom sklearn.model_selection import GroupKFold\nimport albumentations as A\nfrom albumentations.pytorch.transforms import ToTensorV2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:27.479414Z","iopub.execute_input":"2025-04-15T15:15:27.479704Z","iopub.status.idle":"2025-04-15T15:15:28.553644Z","shell.execute_reply.started":"2025-04-15T15:15:27.479668Z","shell.execute_reply":"2025-04-15T15:15:28.552754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === Đường dẫn dữ liệu ===\nPATH = \"/kaggle/input/alaska2-image-steganalysis\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:28.554539Z","iopub.execute_input":"2025-04-15T15:15:28.555014Z","iopub.status.idle":"2025-04-15T15:15:28.558400Z","shell.execute_reply.started":"2025-04-15T15:15:28.554990Z","shell.execute_reply":"2025-04-15T15:15:28.557539Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 42\n\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:28.560669Z","iopub.execute_input":"2025-04-15T15:15:28.560886Z","iopub.status.idle":"2025-04-15T15:15:28.576195Z","shell.execute_reply.started":"2025-04-15T15:15:28.560858Z","shell.execute_reply":"2025-04-15T15:15:28.575511Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Preprocessing 📂","metadata":{}},{"cell_type":"markdown","source":"### Dataset Preparation & K-Fold Splitting\n\nTrong bài toán phân loại ảnh giấu tin ALASKA2, tập dữ liệu gồm 4 lớp:\n\n- `Cover` (ảnh gốc)\n- `JMiPOD`, `JUNIWARD`, `UERD` (3 phương pháp giấu tin khác nhau)\n\nChúng được xử lý và gán nhãn lần lượt từ `0 → 3` thông qua danh sách `CLASSES`.","metadata":{}},{"cell_type":"markdown","source":"### Các bước tiền xử lí ảnh\n1. Đọc ảnh và chuyển đổi ảnh sang không gian màu RGB\n2. Thực hiện K fold cross validation\n3. Tăng cường tập dữ liệu\n4. Gắn nhãn vector cho tập dữ liệu","metadata":{}},{"cell_type":"markdown","source":"## 2.1 Dataset\n","metadata":{}},{"cell_type":"markdown","source":"### Đọc ảnh từ tập dữ liệu\n1. Đọc file JPEG với CV\n>  image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n\n2. Với việc mở file JPEG bằng phương pháp này biến file JPEG trở về miền tần số DCT rồi mới biến đổi về ảnh RGB để trích xuất đặc trưng, tức JPEG  -->  YCbCr  --> RGB  -->  --> YCbCr, mỗi bước chuyển đổi vậy khả năng dẫn đến việc thay đổi nhẹ do quá trình làm tròn DCT ","metadata":{}},{"cell_type":"code","source":"class DatasetRetriever(Dataset):\n    \n    def __init__(self, kinds, image_names, labels, transforms=None):\n        super().__init__()\n        self.kinds = kinds\n        self.image_names = image_names\n        self.labels = labels\n        self.transforms = transforms\n\n    def __getitem__(self, index: int):\n        kind, image_name, label = self.kinds[index], self.image_names[index], self.labels[index]\n        image = cv2.imread(f'{PATH}/{kind}/{image_name}', cv2.IMREAD_COLOR)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n        image /= 255.0\n        if self.transforms:\n            sample = {'image': image}\n            sample = self.transforms(**sample)\n            image = sample['image']\n\n        target = torch.zeros(4, dtype=torch.float32)\n        \n        return image, target\n\n    def __len__(self) -> int:\n        return self.image_names.shape[0]\n\n    def get_labels(self):\n        return list(self.labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:28.577832Z","iopub.execute_input":"2025-04-15T15:15:28.578102Z","iopub.status.idle":"2025-04-15T15:15:28.584423Z","shell.execute_reply.started":"2025-04-15T15:15:28.578082Z","shell.execute_reply":"2025-04-15T15:15:28.583612Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2.2 K fold cross validation ","metadata":{}},{"cell_type":"markdown","source":"### Cross-validation: Group K-Fold\n\nĐể đánh giá mô hình một cách ổn định và tránh overfitting, ta sử dụng **Group K-Fold (5 folds)** – với điểm đặc biệt:\n\n> **Các ảnh cùng tên (`image_name`) sẽ không xuất hiện ở cả tập huấn luyện và tập validation cùng lúc**\n\n🔹 Tại sao cần `GroupKFold`?\n- Một ảnh có thể có nhiều phiên bản giấu tin → nếu cùng tên ảnh xuất hiện ở cả train/val thì mô hình dễ **học tủ**\n- Giải pháp: nhóm theo `image_name` để đảm bảo phân tách đúng logic\n\n---\n\n### Quy trình chia fold:\n1. Gán mặc định tất cả ảnh vào `fold = 0`\n2. Dùng `GroupKFold(n_splits=5)` từ `sklearn`:\n   - Input:\n     - `X = dataset.index`\n     - `y = dataset['label']`\n     - `groups = dataset['image_name']`\n   - Output: 5 fold\n3. Gán số `fold` tương ứng cho ảnh thuộc mỗi `val_index`\n\n---\n\n### Kết quả:\n- Cột `fold` trong `DataFrame` chứa giá trị từ 0 → 4\n- Có thể dùng `dataset[dataset.fold != i]` để lấy tập huấn luyện và `dataset[dataset.fold == i]` cho validation tương ứng từng fold\n\n> Điều này giúp mô hình đánh giá được **tính tổng quát** và kiểm soát được **rò rỉ thông tin** trong quá trình huấn luyện.\n","metadata":{}},{"cell_type":"code","source":"CLASSES = ['Cover', 'JMiPOD', 'JUNIWARD', 'UERD']\nN_SPLITS = 5\n\ndataset = []\nfor label, kind in enumerate(CLASSES):\n    image_paths = glob(os.path.join(PATH, kind, '*.jpg'))\n    for path in image_paths:\n        dataset.append({\n            'kind': kind,\n            'image_name': os.path.basename(path),\n            'label': label\n        })\n\n# Shuffle và tạo DataFrame\nrandom.shuffle(dataset)\ndataset = pd.DataFrame(dataset)\n\n# Gán fold mặc định\ndataset.loc[:, 'fold'] = 0\n\n# Chia K-Fold theo image_name (đảm bảo nhóm ảnh không trùng)\ngkf = GroupKFold(n_splits=N_SPLITS)\n\nfor fold_number, (train_index, val_index) in enumerate(gkf.split\n                                                       (X=dataset.index, y=dataset['label'], groups=dataset['image_name'])):\n    dataset.loc[dataset.iloc[val_index].index, 'fold'] = fold_number","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:28.585283Z","iopub.execute_input":"2025-04-15T15:15:28.585580Z","iopub.status.idle":"2025-04-15T15:15:41.143532Z","shell.execute_reply.started":"2025-04-15T15:15:28.585550Z","shell.execute_reply":"2025-04-15T15:15:41.142844Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2.3 Data Agumentation","metadata":{}},{"cell_type":"markdown","source":"### Data Augmentation & Transformations\n\nTrong bài toán phân loại ảnh ALASKA2, việc chuẩn hóa và tăng cường dữ liệu (augmentation) đóng vai trò rất quan trọng, đặc biệt do ảnh stego thường **khác biệt rất nhỏ** so với ảnh cover.\n\nCác phương pháp thực hiện bao gồm việc lật ảnh ngang dọc và chuyển ảnh sang dạng Pytorch Tensor","metadata":{}},{"cell_type":"code","source":"def get_train_transforms():\n    return A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.Resize(height=512, width=512, p=1.0),\n            ToTensorV2(p=1.0),\n        ], p=1.0)\n\ndef get_valid_transforms():\n    return A.Compose([\n            A.Resize(height=512, width=512, p=1.0),\n            ToTensorV2(p=1.0),\n        ], p=1.0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.144336Z","iopub.execute_input":"2025-04-15T15:15:41.144641Z","iopub.status.idle":"2025-04-15T15:15:41.149136Z","shell.execute_reply.started":"2025-04-15T15:15:41.144612Z","shell.execute_reply":"2025-04-15T15:15:41.148182Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2.4 Labeling","metadata":{}},{"cell_type":"markdown","source":"### Dataset Class & One-Hot Encoding\n\nTrong mô hình học sâu với PyTorch, dữ liệu đầu vào thường được đóng gói bằng `torch.utils.data.Dataset`.  \nTại đây, class `DatasetRetriever` được xây dựng để **load ảnh, áp dụng biến đổi, và gắn nhãn đầu ra** một cách có tổ chức và linh hoạt.","metadata":{}},{"cell_type":"code","source":"def onehot(size, target):\n    vec = torch.zeros(size, dtype=torch.float32)\n    vec[target] = 1.\n    return vec\n\nclass DatasetRetriever(Dataset):\n\n    def __init__(self, kinds, image_names, labels, transforms=None):\n        super().__init__()\n        self.kinds = kinds\n        self.image_names = image_names\n        self.labels = labels\n        self.transforms = transforms\n\n    def __getitem__(self, index: int):\n        kind, image_name, label = self.kinds[index], self.image_names[index], self.labels[index]\n        image = cv2.imread(f'{PATH}/{kind}/{image_name}', cv2.IMREAD_COLOR)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n        image /= 255.0\n        if self.transforms:\n            sample = {'image': image}\n            sample = self.transforms(**sample)\n            image = sample['image']\n            \n        target = onehot(4, label)\n        return image, target\n\n    def __len__(self) -> int:\n        return self.image_names.shape[0]\n\n    def get_labels(self):\n        return list(self.labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.150177Z","iopub.execute_input":"2025-04-15T15:15:41.150480Z","iopub.status.idle":"2025-04-15T15:15:41.167213Z","shell.execute_reply.started":"2025-04-15T15:15:41.150437Z","shell.execute_reply":"2025-04-15T15:15:41.166367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Chuyển tập dữ liệu thành các vector one-hot cho từng lớp:\")\nfor i in range(4):\n    print(f\"Class {i}: {onehot(4, i)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:21:32.818381Z","iopub.execute_input":"2025-04-15T15:21:32.818584Z","iopub.status.idle":"2025-04-15T15:21:32.826203Z","shell.execute_reply.started":"2025-04-15T15:21:32.818567Z","shell.execute_reply":"2025-04-15T15:21:32.825413Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Metrics 📐","metadata":{}},{"cell_type":"markdown","source":"### Evaluation Metric: ALASKA Weighted AUC\n\nTrong cuộc thi ALASKA2 (Image Steganalysis), độ chính xác mô hình **không chỉ được đo bằng AUC thông thường**, mà sử dụng một **biến thể có trọng số (Weighted AUC)**, được thiết kế riêng để đánh giá mô hình trong bối cảnh bảo mật ảnh.\n\n#### Cách hoạt động của Alaska Weighted AUC:\n- Thay vì tính toàn bộ diện tích dưới đường ROC, ta chia **trục TPR (True Positive Rate)** thành 2 đoạn:\n  - `[0.0 – 0.4]` → **trọng số 2** (quan trọng hơn)\n  - `[0.4 – 1.0]` → **trọng số 1**\n\n> Điều này phản ánh thực tế là:  \n> **Khả năng phát hiện ảnh stego khi TPR còn thấp (tức là mô hình chưa quá chắc chắn)** quan trọng hơn khả năng phát hiện khi TPR cao (đã quá rõ ràng).\n\n#### Công thức:\n1. **Tính đường ROC** (tpr, fpr) với `sklearn.metrics.roc_curve`\n2. **Chia đoạn TPR theo các mốc** `[0.0, 0.4, 1.0]`\n3. **Tính AUC trên từng đoạn con**, nhân với trọng số tương ứng\n4. **Chuẩn hóa tổng lại**, để kết quả nằm trong [0, 1]\n\n#### Lớp `RocAucMeter` hoạt động như sau:\n- Lưu trữ `y_true` và `y_pred` liên tục qua từng batch\n- Dùng `softmax` để chuyển logits thành xác suất\n- Dùng hàm `alaska_weighted_auc()` để tính điểm cuối cùng\n\n> **Kết luận:**  \n> Metric này yêu cầu mô hình **ổn định và chính xác ngay từ TPR thấp**, rất phù hợp với bài toán phát hiện các mẫu giấu tin (steganography), vốn rất tinh vi và khó phân biệt.\n","metadata":{}},{"cell_type":"code","source":"class AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.254464Z","iopub.execute_input":"2025-04-15T15:15:41.254768Z","iopub.status.idle":"2025-04-15T15:15:41.259223Z","shell.execute_reply.started":"2025-04-15T15:15:41.254737Z","shell.execute_reply":"2025-04-15T15:15:41.258507Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RocAucMeter(object):\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.y_true = np.array([0,1])\n        self.y_pred = np.array([0.5,0.5])\n        self.score = 0\n\n    def update(self, y_true, y_pred):\n        y_true = y_true.cpu().numpy().argmax(axis=1).clip(min=0, max=1).astype(int)\n        y_pred = 1 - nn.functional.softmax(y_pred, dim=1).data.cpu().numpy()[:,0]\n        self.y_true = np.hstack((self.y_true, y_true))\n        self.y_pred = np.hstack((self.y_pred, y_pred))\n        self.score = alaska_weighted_auc(self.y_true, self.y_pred)\n    \n    @property\n    def avg(self):\n        return self.score","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.259954Z","iopub.execute_input":"2025-04-15T15:15:41.260195Z","iopub.status.idle":"2025-04-15T15:15:41.273570Z","shell.execute_reply.started":"2025-04-15T15:15:41.260176Z","shell.execute_reply":"2025-04-15T15:15:41.272795Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def alaska_weighted_auc(y_true, y_valid):\n    tpr_thresholds = [0.0, 0.4, 1.0]\n    weights = [2, 1]\n    \n    fpr, tpr, thresholds = metrics.roc_curve(y_true, y_valid, pos_label=1)\n    \n    # size of subsets\n    areas = np.array(tpr_thresholds[1:]) - np.array(tpr_thresholds[:-1])\n\n    # The total area is normalized by the sum of weights such that the final weighted AUC is between 0 and 1.\n    normalization = np.dot(areas, weights)\n\n    competition_metric = 0\n    for idx, weight in enumerate(weights):\n        y_min = tpr_thresholds[idx]\n        y_max = tpr_thresholds[idx + 1]\n        mask = (y_min < tpr) & (tpr < y_max)\n\n        if np.sum(mask) == 0:\n            continue\n\n        x_padding = np.linspace(fpr[mask][-1], 1, 100)\n            \n        x = np.concatenate([fpr[mask], x_padding])\n        y = np.concatenate([tpr[mask], [y_max] * len(x_padding)])\n        y = y - y_min \n            \n        score = metrics.auc(x, y)\n        submetric = score * weight\n        competition_metric += submetric\n\n    return competition_metric / normalization","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.274454Z","iopub.execute_input":"2025-04-15T15:15:41.274753Z","iopub.status.idle":"2025-04-15T15:15:41.290824Z","shell.execute_reply.started":"2025-04-15T15:15:41.274725Z","shell.execute_reply":"2025-04-15T15:15:41.290150Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Label Smoothing 🧂","metadata":{}},{"cell_type":"markdown","source":"### Custom Loss Function: Label Smoothing\n\nTrong bài toán phân loại ảnh steganalysis, đôi khi mô hình quá tự tin vào một nhãn duy nhất → dễ dẫn đến **overfitting**.\n\n#### Vấn đề:\n- Dùng `CrossEntropy` thông thường → label dạng one-hot `[0, 0, 1, 0]`\n- Mô hình bị ép học “tuyệt đối” vào 1 class\n- Điều này gây hại khi dữ liệu nhiễu, hoặc lớp khó phân biệt như ảnh Stego\n\n---\n\n### Giải pháp: **Label Smoothing**\n\n```python\nLabelSmoothing(smoothing=0.1)\n```\n- Nhằm chia bớt độ tự tin của một nhãn cho đều các nhãn khác\n","metadata":{}},{"cell_type":"code","source":"class LabelSmoothing(nn.Module):\n    def __init__(self, smoothing = 0.1):\n        super(LabelSmoothing, self).__init__()\n        self.confidence = 1.0 - smoothing\n        self.smoothing = smoothing\n\n    def forward(self, x, target):\n        if self.training:\n            x = x.float()\n            target = target.float()\n            logprobs = torch.nn.functional.log_softmax(x, dim = -1)\n\n            nll_loss = -logprobs * target\n            nll_loss = nll_loss.sum(-1)\n    \n            smooth_loss = -logprobs.mean(dim=-1)\n\n            loss = self.confidence * nll_loss + self.smoothing * smooth_loss\n\n            return loss.mean()\n        else:\n            return torch.nn.functional.cross_entropy(x, target)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.291501Z","iopub.execute_input":"2025-04-15T15:15:41.291686Z","iopub.status.idle":"2025-04-15T15:15:41.304231Z","shell.execute_reply.started":"2025-04-15T15:15:41.291665Z","shell.execute_reply":"2025-04-15T15:15:41.303557Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Fitter 🧠","metadata":{}},{"cell_type":"code","source":"### Lớp Fitter – Pipeline Huấn Luyện Tuỳ Chỉnh Cho Bài Toán Phát Hiện Giấu Tin\n\nLớp `Fitter` được thiết kế như một vòng lặp huấn luyện tuỳ chỉnh trong PyTorch, bao trùm toàn bộ quy trình train/val và được trang bị để:\n- Quản lý optimizer, scheduler\n- Theo dõi loss/AUC\n- Lưu/khôi phục checkpoint\n- Tiếp tục training sau khi ngắt (resume)\n\n---\n\n#### 1. Khởi Tạo (`__init__`)\nHàm khởi tạo bao gồm:\n- Nhận `model`, `device`, và `config`\n- Tạo optimizer: `AdamW`\n- Thiết lập scheduler (vd: `ReduceLROnPlateau`)\n- Loss function: `LabelSmoothing`\n- Cấu trúc lưu log\n\n```python\nself.optimizer = torch.optim.AdamW(self.model.parameters(), lr=config.lr)\nself.scheduler = config.SchedulerClass(self.optimizer, **config.scheduler_params)\nself.criterion = LabelSmoothing().to(self.device)\n```\n\n---\n\n#### 2. Vòng Huấn Luyện Chính (`fit`)\nHàm quản lý toàn bộ training qua nhiều epoch:\n\n1. Trong mỗi epoch:\n   - In learning rate\n   - Gọ `train_model()` huấn luyện\n   - Ghi log loss + AUC training\n   - Gọ `validation()` để đánh giá\n   - Lưu best checkpoint nếu loss giảm\n   - Giữ lại 3 checkpoint tốt nhất (xóa cái cũ)\n   - Cập nhật learning rate (nếu scheduler được bật)\n\n---\n\n#### 3. Giai Đoạn Huấn Luyện (`train_model`)\n\n- Mô hình chuyển sang `train()`\n- Vòng lặp qua batches:\n  - Load data lên GPU\n  - Tính loss (label smoothing)\n  - Cập nhật optimizer, metrics (loss + AUC)\n\nNếu `step_scheduler=True` thì gọ `scheduler.step()` sau mỗi batch.\n\n---\n\n## 4. Giai Đoạn Đánh Giá (`validation`)\n\n- Mô hình chuyển sang `eval()`\n- Dữ liệu được dòng qua `torch.no_grad()`\n- Tính loss + ROC AUC\n- Trả về trung bình metrics\n\n---\n\n## 5. Lưu & Load Checkpoint (`save` / `load`)\n\n### `save(path)`\n- Lưu: state_dict của model, optimizer, scheduler, epoch, best loss\n\n### `load(path)`\n- Load từ checkpoint trước và resume training từ epoch +1\n\n```python\nself.epoch = checkpoint['epoch'] + 1\n```\n\n---\n\n## 6. Ghi Log (`log`)\n\nGhi các thông tin training ra cả terminal và file `log.txt`, hỮfu ích khi training trên Kaggle hoặc server.\n\n---\n\n## Tổng Kết\nLớp `Fitter` giúc tổ chức pipeline training một cách linh hoạt, bền vững:\n\n- Hỗ trợ resume đúng epoch\n- Ghi log chi tiết\n- Cập nhật learning rate động\n- Theo dõi metric (loss, AUC)\n\nRất phù hợp cho bài toán giới hạn thời gian như ALASKA2 (giới hạn 12h/session).","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Fitter:\n    def __init__(self, model, device, config):\n        self.config = config\n        self.epoch = 0\n        \n        self.base_dir = './'\n        self.log_path = f'{self.base_dir}/log.txt'\n        self.best_summary_loss = 10**5 \n\n        self.model = model\n        self.device = device\n\n        param_optimizer = list(self.model.named_parameters())\n        no_decay = ['bias', 'LayerNorm.bias', 'LayerNorm.weight']\n        optimizer_grouped_parameters = [\n            {'params': [p for n, p in param_optimizer if not any(nd in n for nd in no_decay)], 'weight_decay': 0.001},\n            {'params': [p for n, p in param_optimizer if any(nd in n for nd in no_decay)], 'weight_decay': 0.0}\n        ] \n\n        self.optimizer = torch.optim.AdamW(self.model.parameters(), lr=config.lr)\n        self.scheduler = config.SchedulerClass(self.optimizer, **config.scheduler_params)\n        self.criterion = LabelSmoothing().to(self.device)\n        self.log(f'Fitter prepared. Device is {self.device}')\n\n    def fit(self, train_loader, validation_loader):\n        for e in range(self.epoch, self.config.n_epochs):\n            if self.config.verbose:\n                lr = self.optimizer.param_groups[0]['lr']\n                timestamp = datetime.utcnow().isoformat()\n                self.log(f'\\n{timestamp}\\nLR: {lr}')\n\n            # Training\n            t = time.time()\n            summary_loss, final_scores = self.train_model(train_loader)\n\n            self.log(f'[RESULT]: Train. Epoch: {self.epoch},summary_loss: {summary_loss.avg:.5f},final_score: {final_scores.avg:.5f},time: {(time.time() - t):.5f}')\n            self.save(f'{self.base_dir}/last-checkpoint.bin')\n\n            # Validation\n            t = time.time()\n            summary_loss, final_scores = self.validation(validation_loader)\n\n            self.log(f'[RESULT]: Val. Epoch: {self.epoch},summary_loss: {summary_loss.avg:.5f},final_score: {final_scores.avg:.5f},time: {(time.time() - t):.5f}')\n            if summary_loss.avg < self.best_summary_loss:\n                self.best_summary_loss = summary_loss.avg\n                self.model.eval()\n                self.save(f'{self.base_dir}/best-checkpoint-{str(self.epoch).zfill(3)}epoch.bin')\n                for path in sorted(glob(f'{self.base_dir}/best-checkpoint-*epoch.bin'))[:-3]:\n                    os.remove(path)\n\n            # Next epoch\n            if self.config.validation_scheduler:\n                self.scheduler.step(metrics=summary_loss.avg)\n            self.epoch += 1\n\n    def validation(self, val_loader):\n        self.model.eval()\n        summary_loss = AverageMeter()\n        final_scores = RocAucMeter()\n        t = time.time()\n        for step, (images, targets) in enumerate(val_loader):\n            if self.config.verbose:\n                if step % self.config.verbose_step == 0:\n                    print(\n                        f'Val Step {step}/{len(val_loader)}, ' + \\\n                        f'summary_loss: {summary_loss.avg:.5f}, final_score: {final_scores.avg:.5f}, ' + \\\n                        f'time: {(time.time() - t):.5f}', end='\\r'\n                    )\n            with torch.no_grad():\n                targets = targets.to(self.device).float()\n                batch_size = images.shape[0]\n                images = images.to(self.device).float()\n                outputs = self.model(images)\n                loss = self.criterion(outputs, targets)\n                final_scores.update(targets, outputs)\n                summary_loss.update(loss.detach().item(), batch_size)\n\n        return summary_loss, final_scores\n\n    def train_model(self, train_loader):\n        self.model.train()\n        summary_loss = AverageMeter()\n        final_scores = RocAucMeter()\n        t = time.time()\n        for step, (images, targets) in enumerate(train_loader):\n            if self.config.verbose:\n                if step % self.config.verbose_step == 0:\n                    print(\n                        f'Train Step {step}/{len(train_loader)}, ' + \\\n                        f'summary_loss: {summary_loss.avg:.5f}, final_score: {final_scores.avg:.5f}, ' + \\\n                        f'time: {(time.time() - t):.5f}', end='\\r'\n                    )\n            \n            targets = targets.to(self.device).float()\n            images = images.to(self.device).float()\n            batch_size = images.shape[0]\n\n            self.optimizer.zero_grad()\n            outputs = self.model(images)\n            loss = self.criterion(outputs, targets)\n            loss.backward()\n            \n            final_scores.update(targets, outputs)\n            summary_loss.update(loss.detach().item(), batch_size)\n\n            self.optimizer.step()\n\n            if self.config.step_scheduler:\n                self.scheduler.step()\n\n        return summary_loss, final_scores\n    \n    def save(self, path):\n        self.model.eval()\n        torch.save({\n            'model_state_dict': self.model.state_dict(),\n            'optimizer_state_dict': self.optimizer.state_dict(),\n            'scheduler_state_dict': self.scheduler.state_dict(),\n            'best_summary_loss': self.best_summary_loss,\n            'epoch': self.epoch,\n        }, path)\n        \n    def load(self, path):\n        checkpoint = torch.load(path, map_location=self.device)\n        self.model.load_state_dict(checkpoint['model_state_dict'], strict=False)\n        self.optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n        self.scheduler.load_state_dict(checkpoint['scheduler_state_dict'])\n        self.best_summary_loss = checkpoint['best_summary_loss']\n        self.epoch = checkpoint['epoch'] + 1\n        self.log(f'🔁 Đã load checkpoint từ {path}, resume từ epoch {self.epoch}')\n    \n    def log(self, message):\n        if self.config.verbose:\n            print(message)\n        with open(self.log_path, 'a+') as logger:\n            logger.write(f'{message}\\n')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.305016Z","iopub.execute_input":"2025-04-15T15:15:41.305234Z","iopub.status.idle":"2025-04-15T15:15:41.321933Z","shell.execute_reply.started":"2025-04-15T15:15:41.305216Z","shell.execute_reply":"2025-04-15T15:15:41.321140Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Custom Training Loop: `Fitter` Class\n\nLớp `Fitter` được thiết kế để bao quát toàn bộ quy trình huấn luyện mô hình một cách có tổ chức và linh hoạt – hỗ trợ đầy đủ các tính năng như:\n\n---\n\n#### **Chức năng chính:**\n- Quản lý mô hình, thiết bị, optimizer, scheduler và loss function\n- Theo dõi epoch hiện tại (`self.epoch`) để **hỗ trợ resume training**\n- Ghi log toàn bộ quá trình huấn luyện vào file `log.txt`\n\n---\n\n#### **Huấn luyện (`fit` method):**\n- Với mỗi epoch:\n  - Gọi `train_model()` để huấn luyện\n  - Gọi `validation()` để đánh giá\n  - Lưu checkpoint tốt nhất (`best-checkpoint`) và mới nhất (`last-checkpoint`)\n  - Giảm learning rate bằng `ReduceLROnPlateau` nếu cần\n\n---\n\n#### **Resume Training – Checkpoint Logic:**\n- `save(path)` lưu:\n  - Trạng thái mô hình, optimizer, scheduler\n  - `best_summary_loss`\n  - Epoch hiện tại\n- `load(path)` đọc checkpoint và **resume training đúng từ epoch +1**\n\n---\n\n#### **Metric & Loss Tracking:**\n- Sử dụng:\n  - `AverageMeter()` để theo dõi loss trung bình theo batch\n  - `RocAucMeter()` để tính AUC score từng epoch\n- Log kết quả theo từng bước và từng epoch\n\n---\n\n>  Lớp `Fitter` giúp chuẩn hoá pipeline huấn luyện – đặc biệt hiệu quả khi làm việc trên nền tảng giới hạn thời gian như Kaggle (12 giờ/session).\n","metadata":{}},{"cell_type":"markdown","source":"# 6. Plot data 🎨 ","metadata":{}},{"cell_type":"markdown","source":"### Trực quan hoá dữ liệu huấn luyện","metadata":{}},{"cell_type":"code","source":"def parse_log_file(log_path):\n    with open(log_path, 'r') as f:\n        lines = f.readlines()\n\n    epochs = []\n    train_loss = []\n    train_score = []\n    train_time = []\n\n    val_loss = []\n    val_score = []\n    val_time = []\n\n    lr_list = []\n\n    current_lr = None\n    for line in lines:\n        if line.startswith('LR:'):\n            current_lr = float(line.strip().split(':')[1])\n        elif '[RESULT]: Train.' in line:\n            epoch = int(re.search(r'Epoch: (\\d+)', line).group(1))\n            summary_loss = float(re.search(r'summary_loss: ([\\d.]+)', line).group(1))\n            final_score = float(re.search(r'final_score: ([\\d.]+)', line).group(1))\n            t_time = float(re.search(r'time: ([\\d.]+)', line).group(1))\n\n            epochs.append(epoch)\n            train_loss.append(summary_loss)\n            train_score.append(final_score)\n            train_time.append(t_time)\n            lr_list.append(current_lr)  # log lr theo mỗi epoch\n\n        elif '[RESULT]: Val.' in line:\n            val_summary_loss = float(re.search(r'summary_loss: ([\\d.]+)', line).group(1))\n            val_final_score = float(re.search(r'final_score: ([\\d.]+)', line).group(1))\n            val_t_time = float(re.search(r'time: ([\\d.]+)', line).group(1))\n\n            val_loss.append(val_summary_loss)\n            val_score.append(val_final_score)\n            val_time.append(val_t_time)\n\n    return {\n        'epochs': epochs,\n        'train_loss': train_loss,\n        'val_loss': val_loss,\n        'train_score': train_score,\n        'val_score': val_score,\n        'train_time': train_time,\n        'val_time': val_time,\n        'lr': lr_list\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.322723Z","iopub.execute_input":"2025-04-15T15:15:41.322931Z","iopub.status.idle":"2025-04-15T15:15:41.339572Z","shell.execute_reply.started":"2025-04-15T15:15:41.322888Z","shell.execute_reply":"2025-04-15T15:15:41.338750Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_log_results(metrics, save_path='log_plots.png'):\n    epochs = metrics['epochs']\n\n    fig, axs = plt.subplots(2, 2, figsize=(14, 10))\n    fig.suptitle('Training Metrics from Log File', fontsize=16)\n\n    # 1. Loss\n    axs[0, 0].plot(epochs, metrics['train_loss'], label='Train Loss', marker='o')\n    axs[0, 0].plot(epochs, metrics['val_loss'], label='Val Loss', marker='x')\n    axs[0, 0].set_title('Loss per Epoch')\n    axs[0, 0].set_xlabel('Epoch')\n    axs[0, 0].set_ylabel('Loss')\n    axs[0, 0].legend()\n\n    # 2. Score\n    axs[0, 1].plot(epochs, metrics['train_score'], label='Train Score', marker='o')\n    axs[0, 1].plot(epochs, metrics['val_score'], label='Val Score', marker='x')\n    axs[0, 1].set_title('AUC score per Epoch')\n    axs[0, 1].set_xlabel('Epoch')\n    axs[0, 1].set_ylabel('Score')\n    axs[0, 1].legend()\n\n    # 3. LR\n    axs[1, 0].plot(epochs, metrics['lr'], label='Learning Rate', marker='o')\n    axs[1, 0].set_title('Learning Rate')\n    axs[1, 0].set_xlabel('Epoch')\n    axs[1, 0].set_ylabel('LR')\n    axs[1, 0].legend()\n\n    # 4. Time\n    axs[1, 1].plot(epochs, metrics['train_time'], label='Train Time (s)', marker='o')\n    axs[1, 1].plot(epochs, metrics['val_time'], label='Val Time (s)', marker='x')\n    axs[1, 1].set_title('Time per Epoch')\n    axs[1, 1].set_xlabel('Epoch')\n    axs[1, 1].set_ylabel('Seconds')\n    axs[1, 1].legend()\n\n    plt.tight_layout(rect=[0, 0.03, 1, 0.95])\n    plt.savefig(save_path)\n    plt.close()\n    print(f\"✅ Đã lưu biểu đồ tại: {save_path}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.340424Z","iopub.execute_input":"2025-04-15T15:15:41.340680Z","iopub.status.idle":"2025-04-15T15:15:41.355224Z","shell.execute_reply.started":"2025-04-15T15:15:41.340648Z","shell.execute_reply":"2025-04-15T15:15:41.354468Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 7. EfficientNet","metadata":{}},{"cell_type":"markdown","source":"1. Xây dựng custom model với backbone là sử dụng mô hình pre-trained EfficientNet-B0.\n2. **Ý tưởng** sử dụng modify model **EfficientNet-B0** với **3 lớp Conv2d** ở đầu để h giữ nguyên kích thước ảnh 512x512 và tăng dần các kênh màu lên 36\n   * Mục đích nhằm trích xuất các đặc trưng cấp thấp tốt hơn việc resize ảnh xuống kích thước 224x224 so với mô hình gốc\n   * Tuy nhiên mô hình gốc chỉ nhận 3 kênh input do đó thực hiện repeat weight 12 lần để thay đổi lớp input đầu của mô hình EfficientNet-B0 giờ sẽ nhận 36 kênh màu\n   * Giữ nguyên các lớp khác để học sâu hơn các đặc trưng, tuy nhiên loại bỏ hoàn toàn block 5 & 6 (biến đổi toàn bộ block 6 thành Conv2d 1x1 để nối thông block 4 đến conv_head) do mô hình chỉ cần nhận diện các thay đổi vừa đủ ko cần quá trừu tượng đồng thời tiết kiệm thời gian train (thực tế chứng minh giảm từ 8000s -> 5600s mỗi vòng lặp)\n   * Classifier cuối với fully-connect đẩy mô hình đầu ra với out-dim = 4\n4.  Mô hình được tham khảo từ tác giả Qisehn Ha [How to Modify Effnet Architecture](http://www.kaggle.com/competitions/alaska2-image-steganalysis/discussion/168542)","metadata":{}},{"cell_type":"code","source":"class EffNet(nn.Module):\n    \n    def __init__(self, out_dim):\n        super(EffNet, self).__init__()\n        self.conv1 = nn.Conv2d(3, 6, 3, stride=1, padding=1, bias=False)\n        self.conv2 = nn.Conv2d(6, 12, 3, stride=1, padding=1, bias=False)\n        self.conv3 = nn.Conv2d(12, 36, 3, stride=1, padding=1, bias=False)\n        self.mybn1 = nn.BatchNorm2d(6)\n        self.mybn2 = nn.BatchNorm2d(12)\n        self.mybn3 = nn.BatchNorm2d(36)\n\n        self.net = timm.create_model('efficientnet_b0', pretrained=True)\n        self.net.conv_stem.weight = nn.Parameter(self.net.conv_stem.weight.repeat(1, 12, 1, 1))\n\n        self.dropout = nn.Dropout(0.5)\n        self.net.blocks[5] = nn.Identity()\n        self.net.blocks[6] = nn.Sequential(\n            nn.Conv2d(self.net.blocks[4][2].conv_pwl.out_channels, self.net.conv_head.in_channels, 1),\n            nn.BatchNorm2d(self.net.conv_head.in_channels),\n            nn.ReLU6(),\n        )\n        self.myfc = nn.Linear(self.net.classifier.in_features, out_dim)\n        self.net.classifier = nn.Identity()\n\n    def extract(self, x):\n        x = F.relu6(self.mybn1(self.conv1(x)))\n        x = F.relu6(self.mybn2(self.conv2(x)))\n        x = F.relu6(self.mybn3(self.conv3(x)))\n        x = self.net(x)\n        return x\n\n    def forward(self, x):\n        x = self.extract(x)\n        x = self.myfc(self.dropout(x))\n        return x","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.355990Z","iopub.execute_input":"2025-04-15T15:15:41.356196Z","iopub.status.idle":"2025-04-15T15:15:41.372107Z","shell.execute_reply.started":"2025-04-15T15:15:41.356168Z","shell.execute_reply":"2025-04-15T15:15:41.371344Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = EffNet(4).cuda()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:41.372854Z","iopub.execute_input":"2025-04-15T15:15:41.373206Z","iopub.status.idle":"2025-04-15T15:15:42.103156Z","shell.execute_reply.started":"2025-04-15T15:15:41.373178Z","shell.execute_reply":"2025-04-15T15:15:42.102471Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7.1 Config","metadata":{}},{"cell_type":"code","source":"# ---------- CONFIG ---------------\nclass Config:\n    batch_size = 16\n    n_epochs = 11 # Thực tế với 12 tiếng session của kaggle chỉ train được tối đa 5 epoch thôi \n    num_workers = 4\n    lr = 0.001\n    # -----------------------------\n    verbose = True\n    verbose_step = 1\n    # -----------------------------\n    step_scheduler = False  # Ko chỉnh lr sau mỗi batch\n    validation_scheduler = True  # Chỉnh sau mỗi epoch\n    #------------------------------\n    SchedulerClass = torch.optim.lr_scheduler.ReduceLROnPlateau\n    scheduler_params = dict(\n        mode='min',\n        factor=0.5,\n        patience=1,\n        verbose=False, \n        threshold=0.0001,\n        threshold_mode='abs',\n        cooldown=0, \n        min_lr=1e-8,\n        eps=1e-08\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:42.106279Z","iopub.execute_input":"2025-04-15T15:15:42.106529Z","iopub.status.idle":"2025-04-15T15:15:42.111174Z","shell.execute_reply.started":"2025-04-15T15:15:42.106507Z","shell.execute_reply":"2025-04-15T15:15:42.110159Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Data loader","metadata":{}},{"cell_type":"code","source":"fold_number = 0\n\ntrain_dataset = DatasetRetriever(\n    kinds=dataset[dataset['fold'] != fold_number].kind.values,\n    image_names=dataset[dataset['fold'] != fold_number].image_name.values,\n    labels=dataset[dataset['fold'] != fold_number].label.values,\n    transforms=get_train_transforms(),\n)\n\nvalidation_dataset = DatasetRetriever(\n    kinds=dataset[dataset['fold'] == fold_number].kind.values,\n    image_names=dataset[dataset['fold'] == fold_number].image_name.values,\n    labels=dataset[dataset['fold'] == fold_number].label.values,\n    transforms=get_valid_transforms(),\n)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    sampler = ImbalancedDatasetSampler(train_dataset, labels=train_dataset.get_labels()),\n    batch_size=Config.batch_size,\n    num_workers=Config.num_workers,\n    pin_memory=False,\n    drop_last=True,  \n)\n\nval_loader = DataLoader(\n    validation_dataset, \n    sampler=SequentialSampler(validation_dataset),\n    batch_size=Config.batch_size,\n    num_workers=Config.num_workers,\n    shuffle=False,\n    pin_memory=False,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:42.112331Z","iopub.execute_input":"2025-04-15T15:15:42.112603Z","iopub.status.idle":"2025-04-15T15:15:43.079846Z","shell.execute_reply.started":"2025-04-15T15:15:42.112583Z","shell.execute_reply":"2025-04-15T15:15:43.079179Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7.2 Training Session","metadata":{}},{"cell_type":"code","source":"class TrainingSession:\n    def __init__(self, model, config, train_loader, val_loader,\n                 ckpt_folder ='/kaggle/input/alaska-checkpoint', # thư mục chứa checkpoint\n                 output_log_path ='/kaggle/working/log.txt'):\n        \n        self.model = model\n        self.device = torch.device('cuda:0')\n        self.config = config\n        self.train_loader = train_loader\n        self.val_loader = val_loader\n        self.ckpt_folder = ckpt_folder\n        self.output_log_path = output_log_path\n\n    def append_previous_log(self):\n        prev_log = os.path.join(self.ckpt_folder, 'log.txt')\n        if os.path.exists(prev_log):\n            with open(prev_log, 'r') as f:\n                old_content = f.read()\n            with open(self.output_log_path, 'a+') as f:\n                f.write('\\n\\n# ==== Log từ phiên trước ====\\n')\n                f.write(old_content)\n                f.write('\\n\\n# ==== Bắt đầu phiên mới ====\\n')\n            print(f'📜 Ghi lại log cũ từ {prev_log} vào {self.output_log_path}')\n        else:\n            print(f'⚠️ Không tìm thấy log.txt trong {self.ckpt_folder}')\n            \n    def get_latest_best_checkpoint(self):\n        pattern = re.compile(r'best-checkpoint-(\\d+)epoch\\.bin')\n        max_epoch = -1\n        best_path = None\n\n        if not os.path.exists(self.ckpt_folder):\n            return None\n\n        for fname in os.listdir(self.ckpt_folder):\n            match = pattern.match(fname)\n            if match:\n                epoch = int(match.group(1))\n                if epoch > max_epoch:\n                    max_epoch = epoch\n                    best_path = os.path.join(self.ckpt_folder, fname)\n\n        return best_path\n    \n    def run(self):\n        fitter = Fitter(self.model, self.device, self.config)\n        self.append_previous_log()\n\n        # Luôn luôn tìm và load best checkpoint nếu có\n        best_ckpt = self.get_latest_best_checkpoint()\n        if best_ckpt is not None:\n            print(f'🔁 Resume từ checkpoint: {best_ckpt}')\n            fitter.load(best_ckpt)\n        else:\n            print(f'🚨 Không có checkpoint nào, bắt đầu từ đầu (epoch 0)')\n\n        print(f'▶️ Bắt đầu training từ epoch {fitter.epoch}')\n        fitter.fit(self.train_loader, self.val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:43.086478Z","iopub.execute_input":"2025-04-15T15:15:43.086717Z","iopub.status.idle":"2025-04-15T15:15:43.105657Z","shell.execute_reply.started":"2025-04-15T15:15:43.086697Z","shell.execute_reply":"2025-04-15T15:15:43.104921Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 8. Training...","metadata":{}},{"cell_type":"code","source":"session = TrainingSession(model, Config, train_loader, val_loader)\n#session.run()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:43.106452Z","iopub.execute_input":"2025-04-15T15:15:43.106729Z","iopub.status.idle":"2025-04-15T15:15:43.121778Z","shell.execute_reply.started":"2025-04-15T15:15:43.106709Z","shell.execute_reply":"2025-04-15T15:15:43.120946Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 9. Inference","metadata":{}},{"cell_type":"markdown","source":"### Load model\nSau quá trình huấn luyện xong lưu các model ở thư mục output lại và tạo thư mục input riêng để upload lại để thực hiện quá trình đánh giá trên tập test","metadata":{}},{"cell_type":"code","source":"#checkpoint = torch.load('../input/alaska-checkpoint/best-checkpoint-004epoch.bin')\n#checkpoint = torch.load('../input/alaska-checkpoint/best-checkpoint-008epoch.bin')\ncheckpoint = torch.load('../input/alaska-checkpoint/best-checkpoint-010epoch.bin')\nmodel.load_state_dict(checkpoint['model_state_dict']);\nmodel.eval();","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:43.122611Z","iopub.execute_input":"2025-04-15T15:15:43.122828Z","iopub.status.idle":"2025-04-15T15:15:44.079322Z","shell.execute_reply.started":"2025-04-15T15:15:43.122799Z","shell.execute_reply":"2025-04-15T15:15:44.078656Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"checkpoint.keys()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:44.080074Z","iopub.execute_input":"2025-04-15T15:15:44.080296Z","iopub.status.idle":"2025-04-15T15:15:44.085007Z","shell.execute_reply.started":"2025-04-15T15:15:44.080269Z","shell.execute_reply":"2025-04-15T15:15:44.084217Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Lưu lại biểu đồ","metadata":{}},{"cell_type":"code","source":"metrics = parse_log_file('/kaggle/input/alaska-checkpoint/log.txt')\nplot_log_results(metrics, save_path='/kaggle/working/log_plots.png')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:44.085696Z","iopub.execute_input":"2025-04-15T15:15:44.085928Z","iopub.status.idle":"2025-04-15T15:15:44.935010Z","shell.execute_reply.started":"2025-04-15T15:15:44.085882Z","shell.execute_reply":"2025-04-15T15:15:44.934126Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 10. Test-Time Augmentation","metadata":{}},{"cell_type":"code","source":"def get_test_transforms(mode):\n    if mode == 0:\n        return A.Compose([\n                A.Resize(height=512, width=512, p=1.0),\n                ToTensorV2(p=1.0),\n            ], p=1.0)\n    elif mode == 1:\n        return A.Compose([\n                A.HorizontalFlip(p=1),\n                A.Resize(height=512, width=512, p=1.0),\n                ToTensorV2(p=1.0),\n            ], p=1.0)    \n    elif mode == 2:\n        return A.Compose([\n                A.VerticalFlip(p=1),\n                A.Resize(height=512, width=512, p=1.0),\n                ToTensorV2(p=1.0),\n            ], p=1.0)\n    else:\n        return A.Compose([\n                A.HorizontalFlip(p=1),\n                A.VerticalFlip(p=1),\n                A.Resize(height=512, width=512, p=1.0),\n                ToTensorV2(p=1.0),\n            ], p=1.0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:44.935757Z","iopub.execute_input":"2025-04-15T15:15:44.936028Z","iopub.status.idle":"2025-04-15T15:15:44.941713Z","shell.execute_reply.started":"2025-04-15T15:15:44.935996Z","shell.execute_reply":"2025-04-15T15:15:44.940819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DatasetSubmissionRetriever(Dataset):\n\n    def __init__(self, image_names, transforms=None):\n        super().__init__()\n        self.image_names = image_names\n        self.transforms = transforms\n\n    def __getitem__(self, index: int):\n        image_name = self.image_names[index]\n        image = cv2.imread(f'{PATH}/Test/{image_name}', cv2.IMREAD_COLOR)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB).astype(np.float32)\n        image /= 255.0\n        if self.transforms:\n            sample = {'image': image}\n            sample = self.transforms(**sample)\n            image = sample['image']\n\n        return image_name, image\n\n    def __len__(self) -> int:\n        return self.image_names.shape[0]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:44.942623Z","iopub.execute_input":"2025-04-15T15:15:44.942861Z","iopub.status.idle":"2025-04-15T15:15:44.958723Z","shell.execute_reply.started":"2025-04-15T15:15:44.942841Z","shell.execute_reply":"2025-04-15T15:15:44.957979Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 11. Save & Prediction","metadata":{}},{"cell_type":"markdown","source":"### Đọc dữ liệu file test và test-time augmentation\n\nĐoạn mã dưới đây thực hiện **dự đoán trên tập kiểm tra (Test Set)** bằng cách áp dụng **4 phiên bản biến đổi (TTA modes)** cho từng ảnh.\n\n#### Chi tiết quy trình:\n1. **Loop qua 4 mode TTA (0 đến 3)**:\n   - Mode 0: Ảnh gốc\n   - Mode 1: Flip ngang\n   - Mode 2: Flip dọc\n   - Mode 3: Flip cả ngang + dọc\n2. Với mỗi mode:\n   - Tạo `DatasetSubmissionRetriever` để load ảnh test và áp dụng transform tương ứng\n   - Dùng `DataLoader` để load ảnh theo batches\n3. Với mỗi batch:\n   - Dự đoán với mô hình (`model(images.cuda())`)\n   - Áp dụng softmax để lấy xác suất (đảo ngược lớp 0 vì bài toán là binary cover/stego)\n   - Lưu tên ảnh và xác suất dự đoán vào `result`\n\n#### Kết quả:\n- Danh sách `results` chứa 4 dictionary kết quả (mỗi mode 1 dict)\n- Có thể dùng tiếp để thực hiện **weighted average** hoặc **ensemble logic**\n\n> Cách này giúp tận dụng hiệu quả TTA để tăng độ ổn định và chính xác khi dự đoán trên tập test.\n","metadata":{}},{"cell_type":"code","source":"results = []\nfor mode in range(0, 4):\n    dataset = DatasetSubmissionRetriever(\n        image_names=np.array([path.split('/')[-1] for path in glob('../input/alaska2-image-steganalysis/Test/*.jpg')]),\n        transforms=get_test_transforms(mode),\n    )\n\n    data_loader = DataLoader(\n        dataset,\n        batch_size=8,\n        shuffle=False,\n        num_workers=2,\n        drop_last=False,\n    )\n\n    result = {'Id': [], 'Label': []}\n    for step, (image_names, images) in enumerate(data_loader):\n        print(step, end='\\r')\n        \n        y_pred = model(images.cuda())\n        y_pred = 1 - nn.functional.softmax(y_pred, dim=1).data.cpu().numpy()[:,0]\n        \n        result['Id'].extend(image_names)\n        result['Label'].extend(y_pred)\n\n    results.append(result)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:15:44.959505Z","iopub.execute_input":"2025-04-15T15:15:44.959727Z","iopub.status.idle":"2025-04-15T15:21:32.648027Z","shell.execute_reply.started":"2025-04-15T15:15:44.959707Z","shell.execute_reply":"2025-04-15T15:21:32.646952Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submissions = []\nfor mode in range(0,4):\n    submission = pd.DataFrame(results[mode])\n    submissions.append(submission)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:21:32.649222Z","iopub.execute_input":"2025-04-15T15:21:32.649534Z","iopub.status.idle":"2025-04-15T15:21:32.663545Z","shell.execute_reply.started":"2025-04-15T15:21:32.649505Z","shell.execute_reply":"2025-04-15T15:21:32.662650Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_pred = model(images.cuda())\ny_pred = 1 - nn.functional.softmax(y_pred, dim=1).data.cpu().numpy()[:,0]\n\nresult['Id'].extend(image_names)\nresult['Label'].extend(y_pred)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:21:32.664583Z","iopub.execute_input":"2025-04-15T15:21:32.664844Z","iopub.status.idle":"2025-04-15T15:21:32.725172Z","shell.execute_reply.started":"2025-04-15T15:21:32.664820Z","shell.execute_reply":"2025-04-15T15:21:32.724345Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submissions = []\nfor mode in range(0,4):\n    submission = pd.DataFrame(results[mode])\n    submissions.append(submission)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:21:32.726007Z","iopub.execute_input":"2025-04-15T15:21:32.726257Z","iopub.status.idle":"2025-04-15T15:21:32.736832Z","shell.execute_reply.started":"2025-04-15T15:21:32.726227Z","shell.execute_reply":"2025-04-15T15:21:32.735962Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for mode in range(0,4):\n    submissions[mode].to_csv(f'submission_{mode}.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:21:32.737786Z","iopub.execute_input":"2025-04-15T15:21:32.738027Z","iopub.status.idle":"2025-04-15T15:21:32.784294Z","shell.execute_reply.started":"2025-04-15T15:21:32.738007Z","shell.execute_reply":"2025-04-15T15:21:32.783475Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Phân bổ weight cho từng tập ảnh test\nHãy thử thay đổi weight để có kết quả tốt hơn","metadata":{}},{"cell_type":"code","source":"weight0=5  \nweight1=1  \nweight2=1  \nweight3=3  \nweight=weight0+weight1+weight2+weight3","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:21:32.785317Z","iopub.execute_input":"2025-04-15T15:21:32.785614Z","iopub.status.idle":"2025-04-15T15:21:32.789166Z","shell.execute_reply.started":"2025-04-15T15:21:32.785584Z","shell.execute_reply":"2025-04-15T15:21:32.788432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submissions[0]['Label'] = (submissions[0]['Label']*weight0 + submissions[1]['Label']*weight1 \n                           + submissions[2]['Label']*weight2 + submissions[3]['Label']*weight3) / weight\nsubmissions[0].to_csv(f'submission.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-15T15:21:32.790141Z","iopub.execute_input":"2025-04-15T15:21:32.790428Z","iopub.status.idle":"2025-04-15T15:21:32.815972Z","shell.execute_reply.started":"2025-04-15T15:21:32.790399Z","shell.execute_reply":"2025-04-15T15:21:32.815380Z"}},"outputs":[],"execution_count":null}]}