{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":19991,"databundleVersionId":1117522,"sourceType":"competition"}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np \nimport pandas as pd\nimport torch\nimport os\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom PIL import Image\nfrom torchvision.models.efficientnet import efficientnet_b2\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import accuracy_score, roc_curve\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport os\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.559447Z","iopub.execute_input":"2025-05-11T16:24:14.560059Z","iopub.status.idle":"2025-05-11T16:24:14.564759Z","shell.execute_reply.started":"2025-05-11T16:24:14.560033Z","shell.execute_reply":"2025-05-11T16:24:14.564052Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"base_path = '../input/alaska2-image-steganalysis'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.566074Z","iopub.execute_input":"2025-05-11T16:24:14.566352Z","iopub.status.idle":"2025-05-11T16:24:14.594663Z","shell.execute_reply.started":"2025-05-11T16:24:14.566328Z","shell.execute_reply":"2025-05-11T16:24:14.593855Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def read_images_path(dir_name, label):\n    folder_path = os.path.join(base_path, dir_name)\n    return [[os.path.join(folder_path, filename), label] for filename in os.listdir(folder_path)]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.595946Z","iopub.execute_input":"2025-05-11T16:24:14.596468Z","iopub.status.idle":"2025-05-11T16:24:14.610304Z","shell.execute_reply.started":"2025-05-11T16:24:14.596441Z","shell.execute_reply":"2025-05-11T16:24:14.609649Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_image(image_path, size=(512, 512)):\n    # Load image\n    img = Image.open(image_path)\n    \n    # Convert to YCbCr color space (better for detecting steganography)\n    img = img.convert(\"YCbCr\")\n    \n    # Convert to numpy array with high precision (float64 to preserve subtle details)\n    img_array = np.array(img, dtype=np.float64)\n    \n    # Extract DCT residuals to better detect steganography artifacts\n    # This helps highlight the noise patterns that steganography introduces\n    y_channel = img_array[:, :, 0]\n    high_freq = y_channel - cv2.GaussianBlur(y_channel, (3, 3), 0)\n    img_array[:, :, 0] = high_freq\n    \n    # Transpose channels from (H, W, C) to (C, H, W) for PyTorch\n    img_array = np.transpose(img_array, (2, 0, 1))\n    \n    # Normalize to [0, 1] range but avoid standard ImageNet normalization\n    # which could destroy subtle steganographic patterns\n    img_array = img_array / 255.0\n    \n    # Convert to PyTorch tensor\n    img_tensor = torch.tensor(img_array, dtype=torch.float32)\n    \n    return img_tensor","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.611079Z","iopub.execute_input":"2025-05-11T16:24:14.611308Z","iopub.status.idle":"2025-05-11T16:24:14.637672Z","shell.execute_reply.started":"2025-05-11T16:24:14.611284Z","shell.execute_reply":"2025-05-11T16:24:14.637042Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SteganalysisDataset(Dataset):\n    def __init__(self, image_paths, labels, augment=False):\n        self.image_paths = image_paths\n        self.labels = labels\n        self.augment = augment\n        \n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self, index):\n        # Get image path and label\n        img_path = self.image_paths[index]\n        label = self.labels[index]\n        \n        # Process image\n        image = preprocess_image(img_path)\n        \n        # Apply data augmentation (only for training)\n        if self.augment:\n            # Flip horizontally with 50% probability (safe for steganalysis)\n            if np.random.random() > 0.5:\n                image = torch.flip(image, [2])\n                \n            # Random crop and resize back (safe for steganalysis as it preserves patterns)\n            if np.random.random() > 0.5:\n                i, j = torch.randint(0, 32, (2,))\n                image = image[:, i:512-32+i, j:512-32+j]\n                image = F.interpolate(image.unsqueeze(0), size=(512, 512), \n                                     mode='bilinear', align_corners=False).squeeze(0)\n                \n        return image, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.639091Z","iopub.execute_input":"2025-05-11T16:24:14.63932Z","iopub.status.idle":"2025-05-11T16:24:14.65458Z","shell.execute_reply.started":"2025-05-11T16:24:14.639304Z","shell.execute_reply":"2025-05-11T16:24:14.653858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === 1. SRMConv2D with 30 Filters ===\nclass SRMConv2D_30(nn.Module):\n    def __init__(self):\n        super(SRMConv2D_30, self).__init__()\n        kernels = self._get_30_srm_kernels()  # (30, 1, 5, 5)\n        self.conv = nn.Conv2d(1, 30, kernel_size=5, padding=2, bias=False)\n        self.conv.weight = nn.Parameter(kernels, requires_grad=False)\n\n    def forward(self, x):\n        y = self.conv(x[:, 0:1, :, :])\n        cb = self.conv(x[:, 1:2, :, :])\n        cr = self.conv(x[:, 2:3, :, :])\n        return torch.cat([y, cb, cr], dim=1)  # Output: (B, 90, H, W)\n\n    def _get_30_srm_kernels(self):\n        base_kernels = [\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    [[-1, 2, -2, 2, -1], [2, -6, 8, -6, 2], [-2, 8, -12, 8, -2], [2, -6, 8, -6, 2], [-1, 2, -2, 2, -1]],\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    [[1, -2, 0, 2, -1], [-2, 4, 0, -4, 2], [0, 0, 0, 0, 0], [2, -4, 0, 4, -2], [-1, 2, 0, -2, 1]],\n    [[-1, 2, -2, 2, -1], [2, -6, 6, -6, 2], [-2, 6, -6, 6, -2], [2, -6, 6, -6, 2], [-1, 2, -2, 2, -1]],\n    [[0, 0, -1, 0, 0], [0, -2, 0, -2, 0], [-1, 0, 6, 0, -1], [0, -2, 0, -2, 0], [0, 0, -1, 0, 0]],\n    [[0, 0, -1, 0, 0], [0, 2, 0, 2, 0], [-1, 0, -4, 0, -1], [0, 2, 0, 2, 0], [0, 0, -1, 0, 0]],\n    [[-1, 0, 2, 0, -1], [0, 2, 0, 2, 0], [2, 0, -4, 0, 2], [0, 2, 0, 2, 0], [-1, 0, 2, 0, -1]],\n    [[1, -4, 6, -4, 1], [-4, 16, -24, 16, -4], [6, -24, 36, -24, 6], [-4, 16, -24, 16, -4], [1, -4, 6, -4, 1]],\n    [[-1, 2, -1, 0, 0], [2, -6, 2, 0, 0], [-1, 2, -1, 0, 0], [0, 0, 0, 0, 0], [0, 0, 0, 0, 0]],\n    [[0, 0, -1, 2, -1], [0, 0, 2, -6, 2], [0, 0, -1, 2, -1], [0, 0, 0, 0, 0], [0, 0, 0, 0, 0]],\n    [[-1, 2, -1, 2, -1], [2, -6, 2, -6, 2], [-1, 2, -1, 2, -1], [2, -6, 2, -6, 2], [-1, 2, -1, 2, -1]],\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    [[-1, 2, 0, -2, 1], [2, -6, 0, 6, -2], [0, 0, 0, 0, 0], [-2, 6, 0, -6, 2], [1, -2, 0, 2, -1]],\n    [[0, -1, 0, -1, 0], [-1, 4, -4, 4, -1], [0, -4, 0, -4, 0], [-1, 4, -4, 4, -1], [0, -1, 0, -1, 0]],\n    [[-1, 0, 2, 0, -1], [0, 2, 0, 2, 0], [2, 0, -4, 0, 2], [0, 2, 0, 2, 0], [-1, 0, 2, 0, -1]],\n    [[0, 0, 1, 0, 0], [0, 1, -4, 1, 0], [1, -4, 12, -4, 1], [0, 1, -4, 1, 0], [0, 0, 1, 0, 0]],\n    [[1, -2, 1, -2, 1], [-2, 4, -2, 4, -2], [1, -2, 1, -2, 1], [-2, 4, -2, 4, -2], [1, -2, 1, -2, 1]],\n    [[0, 0, 1, 0, 0], [0, -2, 0, -2, 0], [1, 0, 4, 0, 1], [0, -2, 0, -2, 0], [0, 0, 1, 0, 0]],\n    [[0, 0, -1, 0, 0], [0, 2, 0, 2, 0], [-1, 0, -4, 0, -1], [0, 2, 0, 2, 0], [0, 0, -1, 0, 0]],\n    [[-1, 0, 2, 0, -1], [0, -2, 0, -2, 0], [2, 0, 4, 0, 2], [0, -2, 0, -2, 0], [-1, 0, 2, 0, -1]],\n    [[1, -2, 0, 2, -1], [-2, 4, 0, -4, 2], [0, 0, 0, 0, 0], [2, -4, 0, 4, -2], [-1, 2, 0, -2, 1]],\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, -1, 2, -1, 0], [0, 2, -4, 2, 0], [0, -1, 2, -1, 0], [0, 0, 0, 0, 0]],\n    [[-1, 2, -1, 2, -1], [2, -6, 2, -6, 2], [-1, 2, -1, 2, -1], [2, -6, 2, -6, 2], [-1, 2, -1, 2, -1]],\n    [[-1, 2, -2, 2, -1], [2, -6, 8, -6, 2], [-2, 8, -12, 8, -2], [2, -6, 8, -6, 2], [-1, 2, -2, 2, -1]],\n    [[1, 0, -2, 0, 1], [0, -2, 0, -2, 0], [-2, 0, 4, 0, -2], [0, -2, 0, -2, 0], [1, 0, -2, 0, 1]],\n    [[0, 0, 0, 0, 0], [0, 2, -4, 2, 0], [0, -4, 8, -4, 0], [0, 2, -4, 2, 0], [0, 0, 0, 0, 0]],\n    [[0, 0, -1, 0, 0], [0, -2, 0, -2, 0], [-1, 0, 6, 0, -1], [0, -2, 0, -2, 0], [0, 0, -1, 0, 0]],\n    [[-1, 0, 2, 0, -1], [0, 2, 0, 2, 0], [2, 0, -4, 0, 2], [0, 2, 0, 2, 0], [-1, 0, 2, 0, -1]]\n]\n  # Place the full 30 5×5 SRM kernel list here\n        return torch.tensor(base_kernels, dtype=torch.float32).unsqueeze(1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.678397Z","iopub.execute_input":"2025-05-11T16:24:14.678785Z","iopub.status.idle":"2025-05-11T16:24:14.705854Z","shell.execute_reply.started":"2025-05-11T16:24:14.678769Z","shell.execute_reply":"2025-05-11T16:24:14.705131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === 2. Transformer block for arbitrary H × W ===\nclass SpatialTransformerBlock(nn.Module):\n    def __init__(self, in_channels, embed_dim=512, num_heads=8, num_layers=2):\n        super().__init__()\n        self.project = nn.Conv2d(in_channels, embed_dim, 1)\n        encoder_layer = nn.TransformerEncoderLayer(d_model=embed_dim, nhead=num_heads)\n        self.encoder = nn.TransformerEncoder(encoder_layer, num_layers)\n        self.restore = nn.Conv2d(embed_dim, in_channels, 1)\n\n    def forward(self, x):\n        B, C, H, W = x.size()\n        x = self.project(x)              # (B, embed_dim, H, W)\n        x = x.flatten(2).transpose(1, 2) # (B, H*W, embed_dim)\n        x = self.encoder(x)              # (B, H*W, embed_dim)\n        x = x.transpose(1, 2).reshape(B, -1, H, W)\n        return self.restore(x)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.707257Z","iopub.execute_input":"2025-05-11T16:24:14.707797Z","iopub.status.idle":"2025-05-11T16:24:14.724584Z","shell.execute_reply.started":"2025-05-11T16:24:14.707777Z","shell.execute_reply":"2025-05-11T16:24:14.72404Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === 3. Final Model ===\nclass SteganalysisModelHybrid(nn.Module):\n    def __init__(self, num_classes=2):\n        super().__init__()\n        self.srm = SRMConv2D_30()\n\n        self.backbone = timm.create_model(\"efficientnet_b2\", pretrained=True, features_only=True)\n        self.backbone.conv_stem = nn.Conv2d(90, 32, kernel_size=3, stride=2, padding=1, bias=False)\n\n        self.reducer = nn.Sequential(\n            nn.Conv2d(352, 512, kernel_size=1),\n            nn.BatchNorm2d(512),\n            nn.ReLU()\n        )\n\n        self.transformer = SpatialTransformerBlock(in_channels=512, embed_dim=512)\n\n        self.attention = nn.Sequential(\n            nn.Conv2d(512, 128, kernel_size=1),\n            nn.ReLU(),\n            nn.Conv2d(128, 1, kernel_size=1),\n            nn.Sigmoid()\n        )\n\n        self.classifier = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(256, num_classes)\n        )\n\n    def forward(self, x):\n        x = self.srm(x)               # (B, 90, 512, 512)\n        x = self.backbone(x)[-1]      # (B, 352, H, W)\n        x = self.reducer(x)           # (B, 512, H, W)\n        x = self.transformer(x)       # (B, 512, H, W)\n        attn = self.attention(x)      # (B, 1, H, W)\n        x = x * attn                  # (B, 512, H, W)\n        x = F.adaptive_avg_pool2d(x, 1).squeeze(-1).squeeze(-1)  # (B, 512)\n        return self.classifier(x)     # (B, 2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.725208Z","iopub.execute_input":"2025-05-11T16:24:14.725449Z","iopub.status.idle":"2025-05-11T16:24:14.748231Z","shell.execute_reply.started":"2025-05-11T16:24:14.725428Z","shell.execute_reply":"2025-05-11T16:24:14.747534Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_balanced_dataset(sample_size=15000):\n    # Load cover images\n    cover_img = read_images_path('Cover', 0)[:sample_size]\n    \n    # Load stego images (balanced among different techniques)\n    stego_per_type = sample_size\n    jmipod_img = read_images_path('JMiPOD', 1)[:stego_per_type]\n    juniward_img = read_images_path('JUNIWARD', 1)[:stego_per_type]\n    uerd_img = read_images_path('UERD', 1)[:stego_per_type]\n    \n    # Combine all images\n    data = cover_img + jmipod_img + juniward_img + uerd_img\n    df = pd.DataFrame(data=data, columns=['path', 'label'])\n    \n    # Shuffle data\n    df = df.sample(frac=1, random_state=42).reset_index(drop=True)\n    \n    # Split into train and validation\n    train_size = int(len(df) * 0.8)\n    train = df.iloc[:train_size]\n    val = df.iloc[train_size:]\n    \n    return train, val","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.77094Z","iopub.execute_input":"2025-05-11T16:24:14.771151Z","iopub.status.idle":"2025-05-11T16:24:14.788616Z","shell.execute_reply.started":"2025-05-11T16:24:14.771138Z","shell.execute_reply":"2025-05-11T16:24:14.788009Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def weighted_auc(y_true, y_score, tpr_thresholds=[0.0, 0.4, 1.0], weights=[2, 1]):\n    fpr, tpr, _ = roc_curve(y_true, y_score)\n    \n    # Calculate AUC for each region\n    auc_scores = []\n    for i in range(len(tpr_thresholds) - 1):\n        start = tpr_thresholds[i]\n        end = tpr_thresholds[i+1]\n        \n        # Filter points in current region\n        mask = (tpr >= start) & (tpr < end)\n        auc_scores.append(np.trapz(tpr[mask], fpr[mask]))\n    \n    # Calculate weighted AUC\n    weighted_auc = np.sum(np.multiply(auc_scores, weights)) / np.sum(weights)\n    \n    return weighted_auc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.789349Z","iopub.execute_input":"2025-05-11T16:24:14.78983Z","iopub.status.idle":"2025-05-11T16:24:14.807067Z","shell.execute_reply.started":"2025-05-11T16:24:14.789812Z","shell.execute_reply":"2025-05-11T16:24:14.806503Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_metrics(model, dataloader, device='cpu'):\n    model.to(device).eval()\n    labels, probs, preds = [], [], []\n\n    with torch.no_grad():\n        for images, batch_labels in dataloader:\n            outputs = model(images.to(device))\n            batch_probs = F.softmax(outputs, dim=1)[:, 1]\n            batch_preds = torch.argmax(outputs, dim=1)\n\n            labels.extend(batch_labels.cpu().numpy())\n            probs.extend(batch_probs.cpu().numpy())\n            preds.extend(batch_preds.cpu().numpy())\n\n    labels = np.array(labels)\n    probs = np.array(probs)\n    preds = np.array(preds)\n\n    return {\n        'Weighted AUC': weighted_auc(labels, probs),\n        'Accuracy': accuracy_score(labels, preds)\n    }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.807736Z","iopub.execute_input":"2025-05-11T16:24:14.807919Z","iopub.status.idle":"2025-05-11T16:24:14.822165Z","shell.execute_reply.started":"2025-05-11T16:24:14.807905Z","shell.execute_reply":"2025-05-11T16:24:14.821608Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\n\n# Import OpenCV for noise extraction\nimport cv2\n\n# Set hyperparameters\nBATCH_SIZE = 16  # Smaller batch size for better gradient updates\nEPOCHS = 20  # More epochs for better training\nLR = 5e-5  # Learning rate\nWEIGHT_DECAY = 1e-5  # Add regularization","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.822924Z","iopub.execute_input":"2025-05-11T16:24:14.823119Z","iopub.status.idle":"2025-05-11T16:24:14.902126Z","shell.execute_reply.started":"2025-05-11T16:24:14.823098Z","shell.execute_reply":"2025-05-11T16:24:14.901328Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Prepare datasets\nprint(\"Preparing datasets...\")\ntrain, val = prepare_balanced_dataset(sample_size=35000)\n\n# Create datasets\ntrain_dataset = SteganalysisDataset(train['path'].tolist(), train['label'].tolist(), augment=True)\nval_dataset = SteganalysisDataset(val['path'].tolist(), val['label'].tolist(), augment=False)\n\n# Create dataloaders\ntrain_dataloader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, \n                              num_workers=4, pin_memory=True, prefetch_factor=2, persistent_workers=True)\nval_dataloader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, \n                           num_workers=4, pin_memory=True)\n\n# Create model\nmodel = SteganalysisModelHybrid(num_classes=2).to(device)\n\n# Set up class weights to handle imbalance\n# Higher weight for stego images (class 1)\nWEIGHTS = torch.tensor([3, 1], dtype=torch.float32).to(device)\n\n# Create optimizer with weight decay\noptimizer = torch.optim.AdamW(model.parameters(), lr=LR, weight_decay=WEIGHT_DECAY)\n\n# Learning rate scheduler\nscheduler = torch.optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=1e-4,\n    epochs=EPOCHS,\n    steps_per_epoch=len(train_dataloader),\n    pct_start=0.3\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:14.902917Z","iopub.execute_input":"2025-05-11T16:24:14.903213Z","iopub.status.idle":"2025-05-11T16:24:19.098934Z","shell.execute_reply.started":"2025-05-11T16:24:14.90319Z","shell.execute_reply":"2025-05-11T16:24:19.098166Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport torch\nimport torch.nn.functional as F\nfrom tqdm import tqdm\n\n# === Set Kaggle-safe checkpoint path ===\ncheckpoint_path = \"/kaggle/working/checkpoint.pth\"\n\n# === Initialize training state ===\nstart_epoch = 0\ntrain_losses = []\naccuracies = []\nweighted_aucs = []\nbest_auc = 0\n\n# === Resume training if checkpoint exists ===\nif os.path.exists(checkpoint_path):\n    checkpoint = torch.load(checkpoint_path, map_location=device)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n    scheduler.load_state_dict(checkpoint['scheduler_state_dict'])\n    start_epoch = checkpoint['epoch'] + 1\n    train_losses = checkpoint['train_losses']\n    accuracies = checkpoint['accuracies']\n    weighted_aucs = checkpoint['weighted_aucs']\n    best_auc = checkpoint['best_auc']\n    print(f\"Resumed training from epoch {start_epoch}\")\n\n# === Begin training ===\nfor epoch in range(start_epoch, EPOCHS):\n    print(f'\\nEPOCH: {epoch + 1}/{EPOCHS}')\n    if torch.cuda.is_available():\n        print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n        print(f\"Memory allocated: {torch.cuda.memory_allocated(0)/1e9:.2f} GB\")\n        print(f\"Memory reserved: {torch.cuda.memory_reserved(0)/1e9:.2f} GB\")\n    \n    # === Training phase ===\n    model.train()\n    train_loss = 0\n    \n    for images, labels in tqdm(train_dataloader, desc=\"Training\"):\n        images, labels = images.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = F.cross_entropy(outputs, labels, weight=WEIGHTS)\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optimizer.step()\n        scheduler.step()\n        \n        train_loss += loss.item()\n        \n    avg_train_loss = train_loss / len(train_dataloader)\n    train_losses.append(avg_train_loss)\n    print(f'Average Training Loss: {avg_train_loss:.4f}')\n    \n    # === Validation phase ===\n    metrics = calculate_metrics(model, val_dataloader, device)\n    acc = metrics['Accuracy']\n    current_auc = metrics['Weighted AUC']\n    \n    accuracies.append(acc)\n    weighted_aucs.append(current_auc)\n    \n    print(f'Validation Accuracy: {acc:.4f}')\n    print(f'Validation Weighted AUC: {current_auc:.4f}')\n    \n    scheduler.step(current_auc)\n    \n    # === Save best model separately ===\n    if current_auc > best_auc:\n        best_auc = current_auc\n        torch.save(model.state_dict(), '/kaggle/working/best_stego_model.pth')\n        print(f\" New best model saved with Weighted AUC: {best_auc:.4f}\")\n    \n    # === Save full checkpoint every epoch ===\n    checkpoint = {\n        'epoch': epoch,\n        'model_state_dict': model.state_dict(),\n        'optimizer_state_dict': optimizer.state_dict(),\n        'scheduler_state_dict': scheduler.state_dict(),\n        'train_losses': train_losses,\n        'accuracies': accuracies,\n        'weighted_aucs': weighted_aucs,\n        'best_auc': best_auc\n    }\n    torch.save(checkpoint, checkpoint_path)\n    print(f\"Checkpoint saved at epoch {epoch}\")\n\n# === Load and save final model ===\nmodel.load_state_dict(torch.load('/kaggle/working/best_stego_model.pth'))\nprint(\" Best model loaded!\")\n\ntorch.save(model, '/kaggle/working/complete_steganalysis_model.pth')\nprint(\" Complete model saved to 'complete_steganalysis_model.pth'\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T16:24:19.099704Z","iopub.execute_input":"2025-05-11T16:24:19.099942Z","iopub.status.idle":"2025-05-11T21:49:35.967843Z","shell.execute_reply.started":"2025-05-11T16:24:19.099925Z","shell.execute_reply":"2025-05-11T21:49:35.966733Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":" # Plot training metrics\nplt.figure(figsize=(15, 5))\n\n# Plot Loss\nplt.subplot(1, 3, 1)\nplt.plot(range(1, EPOCHS+1), train_losses, marker='o', label='Training Loss')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.title('Loss per Epoch')\nplt.legend()\n\n# Plot Accuracy\nplt.subplot(1, 3, 2)\nplt.plot(range(1, EPOCHS+1), accuracies, marker='o', label='Accuracy', color='green')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.title('Accuracy per Epoch')\nplt.legend()\n\n# Plot Weighted AUC\nplt.subplot(1, 3, 3)\nplt.plot(range(1, EPOCHS+1), weighted_aucs, marker='o', label='Weighted AUC', color='red')\nplt.xlabel('Epoch')\nplt.ylabel('AUC')\nplt.title('Weighted AUC per Epoch')\nplt.legend()\n\nplt.tight_layout()\nplt.savefig('training_metrics.png')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-11T21:55:38.870906Z","iopub.execute_input":"2025-05-11T21:55:38.871524Z","iopub.status.idle":"2025-05-11T21:55:39.643663Z","shell.execute_reply.started":"2025-05-11T21:55:38.871493Z","shell.execute_reply":"2025-05-11T21:55:39.642901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"| Metric       | Trend         | Interpretation                                                     |\n| ------------ | ------------- | ------------------------------------------------------------------ |\n| Loss         | ↓ good        | Model is converging properly                                       |\n| Accuracy     | ↓ fluctuating | May reflect threshold instability or overfitting on majority class |\n| Weighted AUC | ↑ very good   | Strong signal that model is learning subtle differences            |\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}