{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":19991,"databundleVersionId":1117522}],"dockerImageVersionId":31089,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# ==============================================================================\n# Baseline_Sanity.ipynb\n# Chuẩn hóa pipeline & baseline cho ALASKA2 steganalysis\n# ==============================================================================\n\n# ==============================================================================\n# BƯỚC 1: CÀI ĐẶT VÀ CHUẨN BỊ MÔI TRƯỜNG\n# ==============================================================================\n\nprint(\"--- Cài đặt các thư viện cần thiết ---\")\n!pip install -q clip faiss-cpu scikit-learn torch torchvision torchaudio\n\n# ==============================================================================\n# BƯỚC 2: IMPORT THƯ VIỆN CƠ BẢN\n# ==============================================================================\n\nimport os\nimport sys\nimport random\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport cv2\nimport torch\nfrom torch.nn import Identity\nfrom torchvision.models import mobilenet_v2\nfrom torch.utils.data import DataLoader\nfrom tqdm import tqdm\n\n# ==============================================================================\n# BƯỚC 3: CLONE REPO\n# ==============================================================================\n\nrepo_path = \"/kaggle/working/alaska2-steganalysis\"\n\nif os.path.exists(repo_path):\n    print(\"\\nThư mục 'alaska2-steganalysis' đã tồn tại. Đang xóa...\")\n    !rm -r /kaggle/working/alaska2-steganalysis\n\nprint(\"\\n--- Clone repository 'alaska2-steganalysis' ---\")\n!git clone https://github.com/Rinovative/alaska2-steganalysis.git {repo_path}\n\nprint(\"\\n--- Cài đặt thư viện phụ thuộc 'conseal' ---\")\n!pip install -q git+https://github.com/Rinovative/conseal.git\n\nsys.path.append(repo_path)\n\n# ==============================================================================\n# BƯỚC 4: GHIM SEED CỐ ĐỊNH\n# ==============================================================================\n\ndef set_seed(seed=42):\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(42)\nprint(\"✓ Seed cố định thành công (42)\")\n\n# ==============================================================================\n# BƯỚC 5: ĐỊNH NGHĨA LẠI HÀM create_model\n# ==============================================================================\n\nprint(\"\\n--- Định nghĩa lại hàm create_model ---\")\ndef create_model(name: str):\n    \"\"\"Factory for model creation\"\"\"\n    if name == 'conseal_mobilenet':\n        model = mobilenet_v2(weights=\"IMAGENET1K_V1\")\n        model.classifier = Identity()\n        model.out_features = 1280\n    else:\n        raise ValueError(f\"Unknown model name: {name}\")\n    return model\n\n# ==============================================================================\n# BƯỚC 6: IMPORT MODULE TRONG DỰ ÁN\n# ==============================================================================\n\nprint(\"\\n--- Import các modules còn lại ---\")\ntry:\n    from src.util import util_data, util_nb as util_pipeline\n    print(\"✓ Các modules đã được import thành công.\")\nexcept Exception as e:\n    print(f\"✗ Lỗi khi import modules: {e}\")\n    sys.exit()\n\n# ==============================================================================\n# BƯỚC 7: KIỂM TRA DỮ LIỆU\n# ==============================================================================\n\nDATA_DIR = \"/kaggle/input/alaska2-image-steganalysis\"\nutil_data.DATA_DIR = DATA_DIR\nprint(f\"Dataset path đã được thiết lập: {util_data.DATA_DIR}\")\n\ntry:\n    img_path = f\"{DATA_DIR}/Cover/00001.jpg\"\n    img = cv2.imread(img_path)\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    plt.imshow(img)\n    plt.title(\"Sample Cover Image\")\n    plt.axis(\"off\")\n    plt.show()\n    print(\"✓ Hiển thị ảnh mẫu thành công.\")\nexcept Exception as e:\n    print(f\"✗ Lỗi khi hiển thị ảnh: {e}\")\n    sys.exit()\n\n# ==============================================================================\n# BƯỚC 8: TẠO INDEX, CHIA DATASET (SANITY-CHECK)\n# ==============================================================================\n\nprint(\"\\n--- Tạo index và chia dataset ---\")\ntry:\n    dataset_root = DATA_DIR\n    class_labels = {\"Cover\": 0, \"JMiPOD\": 1, \"JUNIWARD\": 2, \"UERD\": 3}\n    \n    print(\"Đang tạo file index với 1% dữ liệu (sanity-check)...\")\n    df = util_data.build_file_index(dataset_root, class_labels, subsample_percent=0.01)\n    df_meta = util_data.add_jpeg_metadata(df, quiet=True)\n    train_df, val_df, test_df = util_data.split_dataset_by_filename(df_meta)\n\n    print(\"✓ Index và chia dataset thành công.\")\n    print(f\"Kích thước train set: {len(train_df)}\")\n    print(f\"Kích thước validation set: {len(val_df)}\")\n    print(f\"Kích thước test set: {len(test_df)}\")\n    \n    print(\"\\nDataFrame đầu tiên:\")\n    print(train_df.head())\n\nexcept Exception as e:\n    print(f\"✗ Đã xảy ra lỗi khi tạo index: {e}\")\n    sys.exit()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T06:21:56.729899Z","iopub.execute_input":"2026-03-31T06:21:56.730058Z","iopub.status.idle":"2026-03-31T06:24:26.092681Z","shell.execute_reply.started":"2026-03-31T06:21:56.730042Z","shell.execute_reply":"2026-03-31T06:24:26.091954Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# STEP 5: FEATURE-BASED BASELINE (SRM + GLCM + SVM/RF)\n# ==============================================================================\n\nfrom skimage.feature import graycomatrix, graycoprops\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.svm import SVC\nfrom sklearn.ensemble import RandomForestClassifier\nfrom sklearn.metrics import accuracy_score, f1_score, confusion_matrix\n\n# Map label_name sang binary: Cover=0, Stego=1\nlabel_map = {\"Cover\":0, \"JMiPOD\":1, \"JUNIWARD\":1, \"UERD\":1}\nfor df_ in [train_df, val_df, test_df]:\n    df_['label'] = df_['label_name'].map(label_map)\n\ndef srm_features(img):\n    gray = cv2.cvtColor(np.array(img), cv2.COLOR_RGB2GRAY)\n    kernels = [np.array([[0,1,0],[1,-4,1],[0,1,0]]),\n               np.array([[1,-2,1],[-2,4,-2],[1,-2,1]])]\n    feats = []\n    for k in kernels:\n        conv = cv2.filter2D(gray, -1, k)\n        feats.extend([conv.mean(), conv.std(), np.median(conv)])\n    return np.array(feats)\n\ndef glcm_features(img, distances=[1], angles=[0,np.pi/4,np.pi/2,3*np.pi/4]):\n    gray = cv2.cvtColor(np.array(img), cv2.COLOR_RGB2GRAY)\n    glcm = graycomatrix(gray, distances=distances, angles=angles, symmetric=True, normed=True)\n    feats = []\n    for prop in ['contrast','dissimilarity','homogeneity','energy','correlation','ASM']:\n        feats.extend(graycoprops(glcm, prop).flatten())\n    return np.array(feats)\n\ndef extract_features(img):\n    return np.concatenate([srm_features(img), glcm_features(img)])\n\ndef build_feature_matrix(df, img_col='path', label_col='label'):\n    X, y = [], []\n    for _, row in tqdm(df.iterrows(), total=len(df)):\n        img_path = row[img_col]\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        X.append(extract_features(img))\n        y.append(row[label_col])\n    return np.array(X), np.array(y)\n\nX_train, y_train = build_feature_matrix(train_df)\nX_val, y_val     = build_feature_matrix(val_df)\nX_test, y_test   = build_feature_matrix(test_df)\n\nscaler = StandardScaler()\nX_train = scaler.fit_transform(X_train)\nX_val   = scaler.transform(X_val)\nX_test  = scaler.transform(X_test)\n\n# --- SVM ---\nsvm_clf = SVC(kernel='linear', probability=True)\nsvm_clf.fit(X_train, y_train)\ny_pred_svm = svm_clf.predict(X_test)\nprint(\"SVM Accuracy:\", accuracy_score(y_test, y_pred_svm))\nprint(\"SVM F1:\", f1_score(y_test, y_pred_svm))\nprint(\"SVM Confusion Matrix:\\n\", confusion_matrix(y_test, y_pred_svm))\n\n# --- Random Forest ---\nrf_clf = RandomForestClassifier(n_estimators=200, random_state=42)\nrf_clf.fit(X_train, y_train)\ny_pred_rf = rf_clf.predict(X_test)\nprint(\"RF Accuracy:\", accuracy_score(y_test, y_pred_rf))\nprint(\"RF F1:\", f1_score(y_test, y_pred_rf))\nprint(\"RF Confusion Matrix:\\n\", confusion_matrix(y_test, y_pred_rf))\n\n# ==============================================================================\n# STEP 6: TINY CNN BASELINE\n# ==============================================================================\n\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nimport torch.nn as nn\nimport torch.optim as optim\nfrom PIL import Image\n\nclass StegoDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        row = self.df.loc[idx]\n        img = Image.open(row['path']).convert('RGB')\n        if self.transform: img = self.transform(img)\n        label = row['label']\n        return img, int(label)\n\ntransform = transforms.Compose([\n    transforms.Resize((128,128)),\n    transforms.ToTensor()\n])\n\ntrain_ds = StegoDataset(train_df, transform)\nval_ds   = StegoDataset(val_df, transform)\ntest_ds  = StegoDataset(test_df, transform)\n\ntrain_loader = DataLoader(train_ds, batch_size=16, shuffle=True)\nval_loader   = DataLoader(val_ds, batch_size=16, shuffle=False)\ntest_loader  = DataLoader(test_ds, batch_size=16, shuffle=False)\n\nclass TinyCNN(nn.Module):\n    def __init__(self, num_classes=2):\n        super().__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(3,16,3,padding=1), nn.ReLU(), nn.MaxPool2d(2),\n            nn.Conv2d(16,32,3,padding=1), nn.ReLU(), nn.MaxPool2d(2),\n            nn.Conv2d(32,64,3,padding=1), nn.ReLU(), nn.MaxPool2d(2)\n        )\n        self.fc = nn.Linear(64*16*16, num_classes)\n    def forward(self,x):\n        x = self.conv(x)\n        x = x.view(x.size(0),-1)\n        return self.fc(x)\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = TinyCNN().to(device)\ncriterion = nn.CrossEntropyLoss()\noptimizer = optim.Adam(model.parameters(), lr=1e-3)\n\n# --- Training loop ---\nfor epoch in range(2):\n    model.train()\n    total, correct = 0,0\n    for imgs, labels in train_loader:\n        imgs, labels = imgs.to(device), labels.to(device)\n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        _, preds = torch.max(outputs,1)\n        total += labels.size(0)\n        correct += (preds==labels).sum().item()\n    print(f\"Epoch {epoch+1} - Train Acc: {correct/total:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T06:24:26.094580Z","iopub.execute_input":"2026-03-31T06:24:26.095288Z","iopub.status.idle":"2026-03-31T06:27:57.752130Z","shell.execute_reply.started":"2026-03-31T06:24:26.095263Z","shell.execute_reply":"2026-03-31T06:27:57.751294Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# FINAL (complete) - SRNet + Balanced sampler + Binary Focal Loss + TTA + Youden+F1 threshold tuning\nimport os, math, copy, random, io\nimport numpy as np\nimport cv2\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torchvision import transforms\n\nfrom sklearn.metrics import (accuracy_score, precision_score, recall_score,\n                             f1_score, confusion_matrix, roc_curve, auc, classification_report,\n                             balanced_accuracy_score)\n\n# NOTE: this script assumes util_data, train_df, val_df, test_df already defined (as in your notebook)\n\n# ----------------------------\n# 0) SEED + DEVICE\n# ----------------------------\ndef set_seed(seed=42):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = False\n    torch.backends.cudnn.benchmark = True\n\nSEED = 42\nset_seed(SEED)\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nUSE_PIN_MEMORY = True if DEVICE.type == \"cuda\" else False\nprint(\"Device:\", DEVICE)\n\n# ----------------------------\n# 1) SRM kernel + transforms (PIL-level -> tensor-level)\n# ----------------------------\nclass RandomJPEG:\n    def __init__(self, p=0.35, qualities=(95,85,75)):\n        self.p = p\n        self.qualities = qualities\n    def __call__(self, pil_img):\n        if random.random() < self.p:\n            q = random.choice(self.qualities)\n            buf = io.BytesIO()\n            pil_img.save(buf, format='JPEG', quality=int(q))\n            buf.seek(0)\n            return Image.open(buf).convert(\"RGB\")\n        return pil_img\n\nclass AddGaussianNoise:\n    def __init__(self, p=0.25, std=(0.002, 0.01)):\n        self.p = p\n        self.std_low, self.std_high = std\n    def __call__(self, pil_img):\n        if random.random() >= self.p:\n            return pil_img\n        arr = np.array(pil_img).astype(np.float32) / 255.0\n        std = random.uniform(self.std_low, self.std_high)\n        noise = np.random.normal(0.0, std, arr.shape).astype(np.float32)\n        arr = np.clip(arr + noise, 0.0, 1.0)\n        arr = (arr * 255.0).astype(np.uint8)\n        return Image.fromarray(arr)\n\nclass ToTensorGray:\n    def __call__(self, pil_img):\n        return transforms.ToTensor()(pil_img.convert('L'))\n\npre_tensor_transforms = [\n    transforms.RandomResizedCrop(size=256, scale=(0.9, 1.0)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomRotation(10),\n    RandomJPEG(p=0.35, qualities=(95,85,75)),\n    AddGaussianNoise(p=0.25, std=(0.002, 0.01)),\n    transforms.ColorJitter(brightness=0.10, contrast=0.10),\n]\n\ntrain_tf = transforms.Compose([\n    *pre_tensor_transforms,\n    ToTensorGray(),\n    transforms.RandomErasing(p=0.15, scale=(0.02, 0.06), ratio=(0.3, 3.3), value=0)\n])\n\nval_tf = transforms.Compose([\n    transforms.Resize((256,256)),\n    ToTensorGray()\n])\n\n# ----------------------------\n# 2) DATASET (grayscale)\n# ----------------------------\nclass StegoDataset(Dataset):\n    def __init__(self, df, transform=None, path_col='path', label_col='label'):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n        self.path_col = path_col\n        self.label_col = label_col\n        self.labels = self.df[self.label_col].astype(int).values\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self, idx):\n        r = self.df.loc[idx]\n        p = r[self.path_col]\n        if not os.path.isabs(p):\n            p = os.path.join(util_data.DATA_DIR, p)\n        pil = Image.open(p).convert(\"RGB\")\n        x = self.transform(pil) if self.transform else ToTensorGray()(pil)\n        y = int(r[self.label_col])\n        return x, y\n\ntrain_ds = StegoDataset(train_df, transform=train_tf)\nval_ds   = StegoDataset(val_df,   transform=val_tf)\ntest_ds  = StegoDataset(test_df,  transform=val_tf)\n\n# ----------------------------\n# 3) Balanced sampler (from train_df)\n# ----------------------------\nlabels_np = train_df['label'].values.astype(int)\nclasses, counts = np.unique(labels_np, return_counts=True)\nprint(\"Train class counts:\", dict(zip(classes, counts)))\n\ninv_freq = {int(c): (1.0 / float(cnt)) for c, cnt in zip(classes, counts)}\nsample_weights = np.array([inv_freq[int(l)] for l in labels_np], dtype=np.float32)\nsampler = WeightedRandomSampler(sample_weights, num_samples=len(sample_weights), replacement=True)\n\nBATCH_SIZE = 32\nNUM_WORKERS = 4\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                          num_workers=NUM_WORKERS, pin_memory=USE_PIN_MEMORY)\nval_loader   = DataLoader(val_ds, batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=USE_PIN_MEMORY)\ntest_loader  = DataLoader(test_ds, batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=USE_PIN_MEMORY)\n\n# ----------------------------\n# 4) MODEL: SRNet (Steganalysis Residual Network) - Full Version\n# ----------------------------\nclass ResidualBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, stride=1):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False)\n        self.bn1 = nn.BatchNorm2d(out_channels)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n        \n        # Shortcut connection to handle different dimensions\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n        else:\n            self.shortcut = nn.Identity()\n\n    def forward(self, x):\n        residual = x\n        \n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n        \n        out = self.conv2(out)\n        out = self.bn2(out)\n        \n        out += self.shortcut(residual)\n        out = self.relu(out)\n        return out\n\n\nclass SRNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Layer 1: SRM convolution\n        srm_kernel = self._get_srm_kernel()\n        self.srm_conv = nn.Conv2d(1, 3, kernel_size=5, stride=1, padding=2, bias=False)\n        self.srm_conv.weight = nn.Parameter(srm_kernel, requires_grad=False)\n        self.srm_bn = nn.BatchNorm2d(3)\n        self.srm_relu = nn.ReLU(inplace=True)\n\n        # Layers with pooling\n        self.conv1 = nn.Conv2d(3, 30, kernel_size=3, padding=1)\n        self.bn1 = nn.BatchNorm2d(30)\n        self.conv2 = nn.Conv2d(30, 30, kernel_size=3, padding=1)\n        self.bn2 = nn.BatchNorm2d(30)\n        self.pool1 = nn.AvgPool2d(kernel_size=3, stride=2, padding=1)\n\n        # Residual blocks\n        self.res1 = ResidualBlock(30, 32, stride=1)\n        self.res2 = ResidualBlock(32, 32, stride=2)\n        self.res3 = ResidualBlock(32, 64, stride=1)\n        self.res4 = ResidualBlock(64, 64, stride=2)\n        self.res5 = ResidualBlock(64, 128, stride=1)\n        self.res6 = ResidualBlock(128, 128, stride=2)\n\n        # Final layers\n        self.conv_final = nn.Conv2d(128, 512, kernel_size=3, padding='same')\n        self.bn_final = nn.BatchNorm2d(512)\n        self.avgpool = nn.AdaptiveAvgPool2d((1,1))\n        self.fc = nn.Linear(512, 1)\n\n    def _get_srm_kernel(self):\n        srm_kernel = np.array([\n            [0, 0, 0, 0, 0],\n            [0, -1, 2, -1, 0],\n            [0, 2, -4, 2, 0],\n            [0, -1, 2, -1, 0],\n            [0, 0, 0, 0, 0]\n        ], dtype=np.float32)\n        srm_kernel = srm_kernel / 4.0\n        srm_kernel = np.stack([srm_kernel, srm_kernel, srm_kernel]) # 3x5x5\n        srm_kernel = np.expand_dims(srm_kernel, 1) # 3x1x5x5\n        return torch.from_numpy(srm_kernel)\n\n    def forward(self, x):\n        x = self.srm_conv(x)\n        x = self.srm_bn(x)\n        x = self.srm_relu(x)\n\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.srm_relu(x)\n        x = self.conv2(x)\n        x = self.bn2(x)\n        x = self.srm_relu(x)\n        x = self.pool1(x)\n        \n        x = self.res1(x)\n        x = self.res2(x)\n        x = self.res3(x)\n        x = self.res4(x)\n        x = self.res5(x)\n        x = self.res6(x)\n        \n        x = self.conv_final(x)\n        x = self.bn_final(x)\n        x = self.srm_relu(x)\n        x = self.avgpool(x)\n        x = x.flatten(1)\n        x = self.fc(x).squeeze(1)\n        return x\n\nmodel = SRNet().to(DEVICE)\nprint(\"Using SRNet (Full) architecture.\")\n# ----------------------------\n# 5) Binary focal loss (class-aware alpha computed from inv_freq)\n# ----------------------------\nclass BinaryFocalLoss(nn.Module):\n    def __init__(self, alpha=(0.5,0.5), gamma=2.0, reduction='mean'):\n        super().__init__()\n        self.register_buffer('alpha', torch.tensor(alpha, dtype=torch.float32))\n        self.gamma = gamma\n        self.reduction = reduction\n    def forward(self, logits, targets):\n        probs = torch.sigmoid(logits)\n        targets = targets.float()\n        pt = probs * targets + (1 - probs) * (1 - targets)\n        alpha_t = self.alpha[1] * targets + self.alpha[0] * (1 - targets)\n        bce = F.binary_cross_entropy_with_logits(logits, targets, reduction='none')\n        loss = alpha_t * ((1 - pt) ** self.gamma) * bce\n        if self.reduction == 'mean':\n            return loss.mean()\n        elif self.reduction == 'sum':\n            return loss.sum()\n        return loss\n\nneg_if = inv_freq.get(0, 1.0)\npos_if = inv_freq.get(1, 1.0)\nalpha_neg = neg_if / (neg_if + pos_if)\nalpha_pos = pos_if / (neg_if + pos_if)\nprint(f\"Binary focal alpha (neg,pos) = ({alpha_neg:.3f}, {alpha_pos:.3f})\")\ncriterion = BinaryFocalLoss(alpha=(alpha_neg, alpha_pos), gamma=2.0).to(DEVICE)\n\n# ----------------------------\n# 6) OPTIM + LR schedule (fixed epochs)\n# ----------------------------\n# For SRNet, we unfreeze all layers from the beginning\nEPOCHS_FULL = 25\nGRAD_CLIP = 1.0\nLR_FULL = 1e-5 # Tốc độ học ban đầu đã được điều chỉnh\n\noptim_full = torch.optim.Adam(model.parameters(), lr=LR_FULL, weight_decay=1e-5)\n\n# Sửa lại hàm cosine_lr để thêm warmup\nWARMUP_EPOCHS = 2 # Đặt 2 epoch đầu tiên để warmup\ndef cosine_lr_with_warmup(base_lr, step, total_steps, warmup_steps, min_lr=1e-7): # min_lr giảm xuống\n    if step < warmup_steps:\n        return base_lr * (step / warmup_steps)\n    cos = 0.5 * (1 + math.cos(math.pi * (step - warmup_steps) / (total_steps - warmup_steps)))\n    return min_lr + (base_lr - min_lr) * cos\n\n# ----------------------------\n# 7) helpers: predict probs (sigmoid), threshold tuning (Youden+F1), eval, TTA\n# ----------------------------\ndef predict_probs(model, loader):\n    model.eval()\n    ys, probs = [], []\n    with torch.no_grad():\n        for x,y in loader:\n            x = x.to(DEVICE)\n            logits = model(x)\n            p = torch.sigmoid(logits).cpu().numpy()\n            ys.extend(y.numpy().tolist())\n            probs.extend(p.tolist())\n    return np.array(ys), np.array(probs)\n\ndef find_best_threshold(model, loader):\n    y_true, y_prob = predict_probs(model, loader)\n    if len(y_true) == 0:\n        return 0.5\n    if len(np.unique(y_true)) < 2:\n        best_t, best_f1 = 0.5, -1.0\n        for t in np.linspace(0.01,0.99,99):\n            y_pred = (y_prob >= t).astype(int)\n            f1m = f1_score(y_true, y_pred, average='macro', zero_division=0)\n            if f1m > best_f1:\n                best_f1, best_t = f1m, float(t)\n        return float(best_t)\n    fpr, tpr, thr = roc_curve(y_true, y_prob)\n    youden = tpr - fpr\n    best_thr_youden = thr[np.argmax(youden)]\n    best_thr_f1, best_f1 = 0.5, -1\n    for t in np.linspace(0.1,0.9,81):\n        preds = (y_prob>=t).astype(int)\n        f1m = f1_score(y_true, preds, average='macro', zero_division=0)\n        if f1m > best_f1:\n            best_f1, best_thr_f1 = f1m, float(t)\n    if best_f1 > youden.max():\n        return float(best_thr_f1)\n    return float(best_thr_youden)\n\ndef evaluate_with_thresh(model, loader, threshold):\n    model.eval()\n    y_true, y_pred, y_prob = [], [], []\n    with torch.no_grad():\n        for x,y in loader:\n            x = x.to(DEVICE)\n            logits = model(x)\n            p = torch.sigmoid(logits).cpu().numpy()\n            preds = (p >= threshold).astype(int)\n            y_true.extend(y.numpy().tolist()); y_pred.extend(preds.tolist()); y_prob.extend(p.tolist())\n    y_true = np.array(y_true); y_pred = np.array(y_pred); y_prob = np.array(y_prob)\n    metrics = {}\n    if len(y_true) == 0:\n        return metrics\n    metrics['acc'] = accuracy_score(y_true, y_pred)\n    metrics['bal_acc'] = balanced_accuracy_score(y_true, y_pred)\n    metrics['f1_macro'] = f1_score(y_true, y_pred, average='macro', zero_division=0)\n    metrics['f1_weighted'] = f1_score(y_true, y_pred, average='weighted', zero_division=0)\n    metrics['prec_w'] = precision_score(y_true, y_pred, average='weighted', zero_division=0)\n    metrics['recall_w'] = recall_score(y_true, y_pred, average='weighted', zero_division=0)\n    metrics['cm'] = confusion_matrix(y_true, y_pred)\n    metrics['y_true'] = y_true; metrics['y_pred'] = y_pred; metrics['y_prob'] = y_prob\n    try:\n        if len(np.unique(y_true)) == 2:\n            fpr, tpr, _ = roc_curve(y_true, y_prob)\n            metrics['auc'] = auc(fpr, tpr)\n        else:\n            metrics['auc'] = np.nan\n    except Exception:\n        metrics['auc'] = np.nan\n    return metrics\n\ndef tta_predict_probs(model, x_tensor):\n    model.eval()\n    with torch.no_grad():\n        logits_list = []\n        logits_list.append(model(x_tensor))\n        logits_list.append(model(torch.flip(x_tensor, dims=[3]))) # hflip\n        logits_list.append(model(torch.flip(x_tensor, dims=[2]))) # vflip\n        try:\n            x_rot = x_tensor.transpose(2,3).contiguous()\n            logits_list.append(model(x_rot))\n        except Exception:\n            pass\n        probs = torch.stack([torch.sigmoid(l) for l in logits_list], dim=0).mean(0)\n    return probs.cpu().numpy()\n\n# ----------------------------\n# 😎 TRAIN: full finetune for ~25 epochs\n# ----------------------------\nSAVE_DIR = \"./checkpoints\"; os.makedirs(SAVE_DIR, exist_ok=True)\nRUN_NAME = \"srnet_binary_focal\"\nbest_state, best_epoch, best_val_f1 = None, -1, -1.0\nbest_threshold = 0.5\n\nprint(\"\\n[Phase 1] full training with SRNet\")\ntotal_steps = EPOCHS_FULL * len(train_loader)\nstep = 0\nfor ep in range(1, EPOCHS_FULL + 1):\n    model.train()\n    total_loss = 0.0; total_samples = 0; corrects = 0\n    for x, y in train_loader:\n        x = x.to(DEVICE); y = y.to(DEVICE)\n        \n        # Sửa lại hàm gọi cosine_lr\n        for gi, pg in enumerate(optim_full.param_groups):\n            pg['lr'] = cosine_lr_with_warmup(LR_FULL, step, total_steps, WARMUP_EPOCHS * len(train_loader))\n        \n        step += 1\n        optim_full.zero_grad()\n        logits = model(x)\n        loss = criterion(logits, y)\n        loss.backward()\n        nn.utils.clip_grad_norm_(model.parameters(), GRAD_CLIP)\n        optim_full.step()\n        total_loss += loss.item() * x.size(0)\n        total_samples += x.size(0)\n        with torch.no_grad():\n            preds = (torch.sigmoid(logits) >= 0.5).long()\n            corrects += (preds == y).sum().item()\n    train_acc = corrects / float(total_samples) if total_samples > 0 else 0.0\n    t_best = find_best_threshold(model, val_loader)\n    val_m = evaluate_with_thresh(model, val_loader, t_best)\n    val_f1m = val_m.get('f1_macro', 0.0)\n    print(f\"Ep{ep} | train_acc={train_acc:.4f} | train_loss={total_loss/total_samples if total_samples>0 else 0:.4f} | val_f1m={val_f1m:.4f} | bal_acc={val_m.get('bal_acc',np.nan):.4f} | AUC={val_m.get('auc',np.nan):.3f} | t={t_best:.3f}\")\n    if val_f1m > best_val_f1:\n        best_val_f1 = val_f1m; best_epoch = ep\n        best_state = copy.deepcopy(model.state_dict()); best_threshold = float(t_best)\n        torch.save({'state':best_state,'threshold':best_threshold}, os.path.join(SAVE_DIR, f\"{RUN_NAME}_best.pth\"))\n\nprint(f\"\\nDone. Best val macro-F1: {best_val_f1:.4f} @ epoch {best_epoch} (threshold={best_threshold:.3f})\")\n\n# ----------------------------\n# 9) Test (TTA) + metrics + ROC/AUC using best checkpoint + threshold\n# ----------------------------\nck = torch.load(os.path.join(SAVE_DIR, f\"{RUN_NAME}_best.pth\"))\nmodel.load_state_dict(ck['state'])\nbest_threshold = float(ck['threshold'])\nmodel.to(DEVICE)\nmodel.eval()\n\ny_true_all, y_prob_all = [], []\nwith torch.no_grad():\n    for x,y in test_loader:\n        x = x.to(DEVICE)\n        probs = tta_predict_probs(model, x)\n        y_prob_all.extend(probs.tolist())\n        y_true_all.extend(y.numpy().tolist())\n\ny_true_all = np.array(y_true_all); y_prob_all = np.array(y_prob_all)\ny_pred_all = (y_prob_all >= best_threshold).astype(int)\n\nacc    = accuracy_score(y_true_all, y_pred_all)\nbal    = balanced_accuracy_score(y_true_all, y_pred_all)\nf1m    = f1_score(y_true_all, y_pred_all, average='macro', zero_division=0)\nf1w    = f1_score(y_true_all, y_pred_all, average='weighted', zero_division=0)\nprecw = precision_score(y_true_all, y_pred_all, average='weighted', zero_division=0)\nrecw  = recall_score(y_true_all, y_pred_all, average='weighted', zero_division=0)\ncm    = confusion_matrix(y_true_all, y_pred_all)\n\nprint(\"\\n--- TEST METRICS (best checkpoint + tuned threshold) ---\")\nprint(f\"Threshold used: {best_threshold:.3f}\")\nprint(f\"Accuracy: {acc:.4f} | Balanced Acc: {bal:.4f}\")\nprint(f\"Macro F1: {f1m:.4f} | Weighted F1: {f1w:.4f}\")\nprint(f\"Precision (weighted): {precw:.4f} | Recall (weighted): {recw:.4f}\")\nprint(\"Confusion matrix:\\n\", cm)\nprint(\"\\nClassification Report:\\n\", classification_report(y_true_all, y_pred_all, target_names=['Cover','Stego'], zero_division=0))\n\ntry:\n    if len(np.unique(y_true_all)) == 2:\n        fpr, tpr, _ = roc_curve(y_true_all, y_prob_all)\n        roc_auc = auc(fpr, tpr)\n        print(f\"AUC: {roc_auc:.4f}\")\n        plt.figure(figsize=(5,4))\n        plt.plot(fpr, tpr, label=f\"AUC={roc_auc:.3f}\")\n        plt.plot([0,1],[0,1],'--', color='gray')\n        plt.xlabel(\"FPR\"); plt.ylabel(\"TPR\"); plt.title(\"ROC (test)\")\n        plt.legend()\n        plt.show()\n    else:\n        print(\"ROC/AUC: only one class present in test labels, skipping ROC plot.\")\nexcept Exception as e:\n    print(\"ROC/AUC failed:\", e)\n\nos.makedirs(\"./checkpoints\", exist_ok=True)\ntorch.save({'state_dict':model.state_dict(), 'threshold':best_threshold}, os.path.join(\"./checkpoints\", \"srnet_final.pth\"))\nprint(\"Saved final model to:\", os.path.join(\"./checkpoints\", \"srnet_final.pth\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T06:27:57.753226Z","iopub.execute_input":"2026-03-31T06:27:57.753817Z","iopub.status.idle":"2026-03-31T06:35:27.370278Z","shell.execute_reply.started":"2026-03-31T06:27:57.753794Z","shell.execute_reply":"2026-03-31T06:35:27.369325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ==============================================================================\n# FINAL CLEAN SRNET (WORKING VERSION - STEGANALYSIS CORRECT)\n# ==============================================================================\n\nimport os, copy, random\nimport numpy as np\nfrom PIL import Image\nimport matplotlib.pyplot as plt\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\n\nfrom sklearn.metrics import accuracy_score, f1_score, confusion_matrix, roc_curve, auc\n\n# ----------------------------\n# 0) SEED\n# ----------------------------\ndef set_seed(seed=42):\n    random.seed(seed); np.random.seed(seed)\n    torch.manual_seed(seed); torch.cuda.manual_seed_all(seed)\n\nset_seed(42)\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device:\", DEVICE)\n\n# ----------------------------\n# 1) TRANSFORM (CRITICAL FIX)\n# ----------------------------\nclass ToTensorGray:\n    def __call__(self, img):\n        return transforms.ToTensor()(img.convert(\"L\"))\n\ntrain_tf = transforms.Compose([\n    transforms.RandomHorizontalFlip(p=0.5),  # ONLY safe aug\n    ToTensorGray()\n])\n\nval_tf = transforms.Compose([\n    ToTensorGray()\n])\n\n# ----------------------------\n# 2) DATASET (BINARY FIX)\n# ----------------------------\nclass StegoDataset(Dataset):\n    def __init__(self, df, tf):\n        self.df = df.reset_index(drop=True)\n        self.tf = tf\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, i):\n        r = self.df.loc[i]\n        p = r['path']\n\n        if not os.path.isabs(p):\n            p = os.path.join(util_data.DATA_DIR, p)\n\n        img = Image.open(p).convert(\"RGB\")\n        x = self.tf(img)\n\n        # 🔥 FIX: binary classification (IMPORTANT)\n        y = 0 if r['label'] == 0 else 1\n        return x, y\n\ntrain_ds = StegoDataset(train_df, train_tf)\nval_ds   = StegoDataset(val_df, val_tf)\ntest_ds  = StegoDataset(test_df, val_tf)\n\ntrain_loader = DataLoader(train_ds, batch_size=32, shuffle=True, num_workers=2)\nval_loader   = DataLoader(val_ds, batch_size=32, shuffle=False, num_workers=2)\ntest_loader  = DataLoader(test_ds, batch_size=32, shuffle=False, num_workers=2)\n\n# ----------------------------\n# 3) SRM 30 FILTERS (CRITICAL)\n# ----------------------------\ndef get_srm_30():\n    filters = []\n    base = [\n        [[0,0,0,0,0],[0,-1,2,-1,0],[0,2,-4,2,0],[0,-1,2,-1,0],[0,0,0,0,0]],\n        [[0,0,0,0,0],[0,0,-1,0,0],[0,-1,4,-1,0],[0,0,-1,0,0],[0,0,0,0,0]]\n    ]\n\n    for i in range(15):\n        k = np.array(base[i % 2], dtype=np.float32)\n        filters.append(k)\n        filters.append(np.rot90(k))\n\n    filters = np.array(filters[:30]) / 4.0\n    return torch.tensor(filters[:, None, :, :], dtype=torch.float32)\n\n# ----------------------------\n# 4) TLU (VERY IMPORTANT)\n# ----------------------------\nclass TLU(nn.Module):\n    def __init__(self, t=3.0):\n        super().__init__()\n        self.t = t\n\n    def forward(self, x):\n        return torch.clamp(x, -self.t, self.t)\n\n# ----------------------------\n# 5) MODEL (CORRECT SRNET CORE)\n# ----------------------------\nclass ResidualBlock(nn.Module):\n    def __init__(self, in_c, out_c, stride=1):\n        super().__init__()\n\n        self.conv1 = nn.Conv2d(in_c, out_c, 3, stride, 1, bias=False)\n        self.bn1 = nn.BatchNorm2d(out_c)\n\n        self.conv2 = nn.Conv2d(out_c, out_c, 3, 1, 1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_c)\n\n        self.shortcut = nn.Identity()\n        if stride != 1 or in_c != out_c:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_c, out_c, 1, stride, bias=False),\n                nn.BatchNorm2d(out_c)\n            )\n\n        self.relu = nn.ReLU(inplace=True)\n\n    def forward(self, x):\n        out = self.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        out += self.shortcut(x)\n        return self.relu(out)\n\nclass SRNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        self.srm = nn.Conv2d(1, 30, 5, padding=2, bias=False)\n        self.srm.weight = nn.Parameter(get_srm_30(), requires_grad=False)\n\n        self.tlu = TLU(3.0)\n\n        self.layer1 = nn.Sequential(\n            nn.Conv2d(30, 32, 3, padding=1),\n            nn.BatchNorm2d(32),\n            nn.ReLU()\n        )\n\n        self.res = nn.Sequential(\n            ResidualBlock(32, 32),\n            ResidualBlock(32, 64, stride=2),\n            ResidualBlock(64, 64),\n            ResidualBlock(64, 128, stride=2),\n            ResidualBlock(128, 128),\n        )\n\n        self.head = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.Linear(128, 1)\n        )\n\n    def forward(self, x):\n        x = self.srm(x)\n        x = self.tlu(x)\n        x = self.layer1(x)\n        x = self.res(x)\n        return self.head(x).squeeze(1)\n\nmodel = SRNet().to(DEVICE)\n\n# ----------------------------\n# 6) LOSS + OPTIM (FIX LR)\n# ----------------------------\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)  # 🔥 FIX lớn\n\n# ----------------------------\n# 7) TRAIN\n# ----------------------------\nEPOCHS = 20\nbest_auc = 0\nbest_state = None\n\nfor ep in range(1, EPOCHS+1):\n    model.train()\n    total_loss = 0\n\n    for x,y in train_loader:\n        x = x.to(DEVICE)\n        y = y.float().to(DEVICE)\n\n        optimizer.zero_grad()\n        logits = model(x)\n        loss = criterion(logits, y)\n        loss.backward()\n        optimizer.step()\n\n        total_loss += loss.item()\n\n    # VALIDATION\n    model.eval()\n    y_true, y_prob = [], []\n\n    with torch.no_grad():\n        for x,y in val_loader:\n            x = x.to(DEVICE)\n            prob = torch.sigmoid(model(x)).cpu().numpy()\n            y_true.extend(y.numpy())\n            y_prob.extend(prob)\n\n    y_true = np.array(y_true)\n    y_prob = np.array(y_prob)\n\n    fpr, tpr, _ = roc_curve(y_true, y_prob)\n    roc_auc = auc(fpr, tpr)\n\n    print(f\"Ep{ep} | loss={total_loss:.4f} | val_AUC={roc_auc:.4f}\")\n\n    if roc_auc > best_auc:\n        best_auc = roc_auc\n        best_state = copy.deepcopy(model.state_dict())\n\nprint(\"Best AUC:\", best_auc)\n\n# ----------------------------\n# 8) TEST\n# ----------------------------\nmodel.load_state_dict(best_state)\nmodel.eval()\n\ny_true, y_prob = [], []\n\nwith torch.no_grad():\n    for x,y in test_loader:\n        x = x.to(DEVICE)\n        prob = torch.sigmoid(model(x)).cpu().numpy()\n        y_true.extend(y.numpy())\n        y_prob.extend(prob)\n\ny_true = np.array(y_true)\ny_prob = np.array(y_prob)\ny_pred = (y_prob >= 0.5).astype(int)\n\nprint(\"\\n--- TEST ---\")\nprint(\"Accuracy:\", accuracy_score(y_true, y_pred))\nprint(\"F1:\", f1_score(y_true, y_pred))\nprint(\"Confusion:\\n\", confusion_matrix(y_true, y_pred))\n\nfpr, tpr, _ = roc_curve(y_true, y_prob)\nroc_auc = auc(fpr, tpr)\nprint(\"AUC:\", roc_auc)\n\nplt.plot(fpr, tpr)\nplt.title(f\"ROC AUC={roc_auc:.3f}\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-31T07:45:57.338916Z","iopub.execute_input":"2026-03-31T07:45:57.339222Z","iopub.status.idle":"2026-03-31T07:46:04.926362Z","shell.execute_reply.started":"2026-03-31T07:45:57.339202Z","shell.execute_reply":"2026-03-31T07:46:04.924992Z"}},"outputs":[],"execution_count":null}]}