{"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":"gpu","dataSources":[{"sourceId":19991,"databundleVersionId":1117522,"sourceType":"competition"}],"dockerImageVersionId":31041,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nos.environ[\"PYTORCH_CUDA_ALLOC_CONF\"] = \"expandable_segments:True\"\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:30:28.87739Z","iopub.execute_input":"2025-05-15T12:30:28.877649Z","iopub.status.idle":"2025-05-15T12:30:28.88114Z","shell.execute_reply.started":"2025-05-15T12:30:28.877627Z","shell.execute_reply":"2025-05-15T12:30:28.880464Z"}},"outputs":[],"execution_count":5},{"cell_type":"code","source":"BASE_PATH = \"/kaggle/input/alaska2-image-steganalysis\"\nCLASSES = [\"Cover\", \"JMiPOD\", \"JUNIWARD\", \"UERD\"]\nLABEL_MAP = {cls: idx for idx, cls in enumerate(CLASSES)}\nSAMPLES_PER_CLASS = 75000\nQF_PER_CLASS = 25000  # 25k images per QF (75, 90, 95)\n\ndef load_image_paths():\n    \"\"\"\n    Assigns 12-class labels based on ordered file positions:\n    - Index 0–24999 → QF75\n    - Index 25000–49999 → QF90\n    - Index 50000–74999 → QF95\n    Label = (class_id * 3) + qf_id ∈ [0, 11]\n    \"\"\"\n    data = []\n    for cls in CLASSES:\n        folder = os.path.join(BASE_PATH, cls)\n        assert os.path.exists(folder), f\"Missing folder: {folder}\"\n        files = sorted([f for f in os.listdir(folder) if f.lower().endswith(\".jpg\")])\n        files = files[:SAMPLES_PER_CLASS]  # Ensure max 75k/class\n\n        for i, f in enumerate(files):\n            qf_id = i // QF_PER_CLASS  # 0: QF75, 1: QF90, 2: QF95\n            label = LABEL_MAP[cls] * 3 + qf_id\n            path = os.path.join(folder, f)\n            data.append((path, label))\n    return data\n\n\n# def load_image_paths():\n#     data = []\n#     for cls in CLASSES:\n#         folder = os.path.join(BASE_PATH, cls)\n#         files = sorted([f for f in os.listdir(folder) if f.lower().endswith(\".jpg\")])\n#         files = files[:SAMPLES_PER_CLASS]  # limit to first 75,000\n#         for f in files:\n#             path = os.path.join(folder, f)\n#             label = LABEL_MAP[cls]\n#             data.append((path, label))\n#     return data\n\n# Load and check total samples\nall_data = load_image_paths()\nprint(f\"Total images loaded: {len(all_data)}\")\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:30:30.753893Z","iopub.execute_input":"2025-05-15T12:30:30.754186Z","iopub.status.idle":"2025-05-15T12:30:34.446538Z","shell.execute_reply.started":"2025-05-15T12:30:30.754162Z","shell.execute_reply":"2025-05-15T12:30:34.445588Z"}},"outputs":[{"name":"stdout","text":"Total images loaded: 300000\n","output_type":"stream"}],"execution_count":6},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedShuffleSplit\nimport numpy as np\nimport pandas as pd\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:30:40.129578Z","iopub.execute_input":"2025-05-15T12:30:40.130217Z","iopub.status.idle":"2025-05-15T12:30:41.01259Z","shell.execute_reply.started":"2025-05-15T12:30:40.130192Z","shell.execute_reply":"2025-05-15T12:30:41.011855Z"}},"outputs":[],"execution_count":7},{"cell_type":"code","source":"all_df = pd.DataFrame(all_data, columns=[\"path\", \"label\"])\nprint(all_df['label'].value_counts())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:30:42.909738Z","iopub.execute_input":"2025-05-15T12:30:42.910553Z","iopub.status.idle":"2025-05-15T12:30:43.029242Z","shell.execute_reply.started":"2025-05-15T12:30:42.91052Z","shell.execute_reply":"2025-05-15T12:30:43.028588Z"}},"outputs":[{"name":"stdout","text":"label\n0     25000\n1     25000\n2     25000\n3     25000\n4     25000\n5     25000\n6     25000\n7     25000\n8     25000\n9     25000\n10    25000\n11    25000\nName: count, dtype: int64\n","output_type":"stream"}],"execution_count":8},{"cell_type":"code","source":"def make_splits(df, n_splits=4, train_size=65000, val_size=10000, seed_base=3):\n    splits = {}\n    y = df[\"label\"].values\n\n    for fold_id in [3, 4, 6, 7]:\n        sss = StratifiedShuffleSplit(n_splits=1, train_size=train_size, test_size=val_size, random_state=seed_base + fold_id)\n        train_idx, val_idx = next(sss.split(df[\"path\"], y))\n\n        splits[f\"split{fold_id}\"] = {\n            \"train\": df.iloc[train_idx].reset_index(drop=True),\n            \"val\": df.iloc[val_idx].reset_index(drop=True)\n        }\n\n    return splits\n\nsplits = make_splits(all_df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:30:45.220965Z","iopub.execute_input":"2025-05-15T12:30:45.221655Z","iopub.status.idle":"2025-05-15T12:30:45.464714Z","shell.execute_reply.started":"2025-05-15T12:30:45.221617Z","shell.execute_reply":"2025-05-15T12:30:45.463906Z"}},"outputs":[],"execution_count":9},{"cell_type":"code","source":"print(\"Split 3 — train label distribution:\")\nprint(splits[\"split3\"][\"train\"][\"label\"].value_counts())\nprint(\"Split 3 — val label distribution:\")\nprint(splits[\"split3\"][\"val\"][\"label\"].value_counts())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:30:47.741417Z","iopub.execute_input":"2025-05-15T12:30:47.742069Z","iopub.status.idle":"2025-05-15T12:30:47.748709Z","shell.execute_reply.started":"2025-05-15T12:30:47.742044Z","shell.execute_reply":"2025-05-15T12:30:47.747903Z"}},"outputs":[{"name":"stdout","text":"Split 3 — train label distribution:\nlabel\n0     5417\n6     5417\n8     5417\n2     5417\n1     5417\n7     5417\n5     5417\n11    5417\n10    5416\n4     5416\n3     5416\n9     5416\nName: count, dtype: int64\nSplit 3 — val label distribution:\nlabel\n9     834\n4     834\n3     834\n10    834\n0     833\n8     833\n7     833\n11    833\n6     833\n2     833\n5     833\n1     833\nName: count, dtype: int64\n","output_type":"stream"}],"execution_count":10},{"cell_type":"code","source":"!pip install -q jpegio\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:30:50.581903Z","iopub.execute_input":"2025-05-15T12:30:50.582151Z","iopub.status.idle":"2025-05-15T12:31:27.299208Z","shell.execute_reply.started":"2025-05-15T12:30:50.582134Z","shell.execute_reply":"2025-05-15T12:31:27.298289Z"}},"outputs":[{"name":"stdout","text":"\u001b[2K     \u001b[90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[32m73.5/73.5 MB\u001b[0m \u001b[31m23.9 MB/s\u001b[0m eta \u001b[36m0:00:00\u001b[0m:00:01\u001b[0m00:01\u001b[0m\n\u001b[?25h  Installing build dependencies ... \u001b[?25l\u001b[?25hdone\n  Getting requirements to build wheel ... \u001b[?25l\u001b[?25hdone\n  Preparing metadata (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n  Building wheel for jpegio (pyproject.toml) ... \u001b[?25l\u001b[?25hdone\n","output_type":"stream"}],"execution_count":11},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\nfrom PIL import Image\nimport jpegio as jio\nimport numpy as np\n\ndef block_dct_map(coeffs, value_range=(-50, 50)):\n    \"\"\"\n    Extract the top-left (DC) value from each 8x8 block (i.e., 64x64 grid from 512x512 coefficients)\n    \"\"\"\n    h, w = coeffs.shape\n    blocks = coeffs.reshape(h//8, 8, w//8, 8).transpose(0, 2, 1, 3)\n    dc_map = blocks[:, :, 0, 0]  # Take the DC coefficient from each block\n    return dc_map  # shape: [64, 64]\n\ndef one_hot_encode_dct(coeffs, value_range=(-50, 50)):\n    coeffs = block_dct_map(coeffs, value_range)\n    min_val, max_val = value_range\n    coeffs = np.clip(coeffs, min_val, max_val)\n    indices = coeffs - min_val\n    n_values = max_val - min_val + 1\n    one_hot = np.eye(n_values, dtype=np.float32)[indices]\n    return np.transpose(one_hot, (2, 0, 1))  # [C, 64, 64]\n\n\nclass Alaska2Dataset(Dataset):\n    def __init__(self, df, mode=\"spatial\", augment=None):\n        self.df = df.reset_index(drop=True)\n        self.mode = mode  # 'spatial' or 'dct'\n        self.augment = augment\n        self.dct_range = (-50, 50)\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        path = self.df.loc[idx, \"path\"]\n        label = self.df.loc[idx, \"label\"]\n\n        if self.mode == \"spatial\":\n            image = Image.open(path).convert(\"YCbCr\")\n            image = np.array(image, dtype=np.float32) / 255.0\n            image = np.transpose(image, (2, 0, 1))  # [C, H, W]\n            return torch.tensor(image, dtype=torch.float32), label\n\n        elif self.mode == \"dct\":\n            jpeg = jio.read(path)\n            Y = one_hot_encode_dct(jpeg.coef_arrays[0], self.dct_range)\n            Cb = one_hot_encode_dct(jpeg.coef_arrays[1], self.dct_range)\n            Cr = one_hot_encode_dct(jpeg.coef_arrays[2], self.dct_range)\n\n            # Align sizes: pad Cb/Cr to match Y (64x64 blocks)\n            def pad64(x):  # [C, H, W] → pad to [C, 64, 64]\n                ch, h, w = x.shape\n                padded = np.zeros((ch, 64, 64), dtype=np.float32)\n                padded[:, :h, :w] = x\n                return padded\n\n            Cb = pad64(Cb)\n            Cr = pad64(Cr)\n\n            dct = np.concatenate([Y, Cb, Cr], axis=0)  # shape [channels, 64, 64]\n            return torch.tensor(dct, dtype=torch.float32), label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:31:30.859733Z","iopub.execute_input":"2025-05-15T12:31:30.860572Z","iopub.status.idle":"2025-05-15T12:31:32.496951Z","shell.execute_reply.started":"2025-05-15T12:31:30.860535Z","shell.execute_reply":"2025-05-15T12:31:32.496191Z"}},"outputs":[],"execution_count":12},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\ntrain_df = splits[\"split3\"][\"train\"]\nds_spatial = Alaska2Dataset(train_df, mode=\"spatial\")\nds_dct = Alaska2Dataset(train_df, mode=\"dct\")\n\nimg, lbl = ds_spatial[0]\nprint(\"Spatial image shape:\", img.shape)  # should be [3, H, W]\n\ndct, lbl = ds_dct[0]\nprint(\"DCT shape:\", dct.shape)  # should be [~303, 64, 64]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:31:36.764647Z","iopub.execute_input":"2025-05-15T12:31:36.765307Z","iopub.status.idle":"2025-05-15T12:31:36.906505Z","shell.execute_reply.started":"2025-05-15T12:31:36.765279Z","shell.execute_reply":"2025-05-15T12:31:36.905642Z"}},"outputs":[{"name":"stdout","text":"Spatial image shape: torch.Size([3, 512, 512])\nDCT shape: torch.Size([303, 64, 64])\n","output_type":"stream"}],"execution_count":13},{"cell_type":"code","source":"import torch.nn as nn\nimport torch.nn.functional as F\n\nclass SEBlock(nn.Module):\n    def __init__(self, channels, reduction=16):\n        super().__init__()\n        self.fc1 = nn.Linear(channels, channels // reduction)\n        self.fc2 = nn.Linear(channels // reduction, channels)\n\n    def forward(self, x):\n        b, c, _, _ = x.shape\n        y = F.adaptive_avg_pool2d(x, 1).view(b, c)\n        y = F.relu(self.fc1(y))\n        y = torch.sigmoid(self.fc2(y)).view(b, c, 1, 1)\n        return x * y\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:31:39.37601Z","iopub.execute_input":"2025-05-15T12:31:39.376707Z","iopub.status.idle":"2025-05-15T12:31:39.381645Z","shell.execute_reply.started":"2025-05-15T12:31:39.376681Z","shell.execute_reply":"2025-05-15T12:31:39.380992Z"}},"outputs":[],"execution_count":14},{"cell_type":"code","source":"class SEBasicBlock(nn.Module):\n    expansion = 1\n\n    def __init__(self, in_channels, out_channels, stride=1, downsample=None):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_channels, out_channels, 3, stride, 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, 3, 1, 1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n        self.se = SEBlock(out_channels)\n        self.downsample = downsample\n\n    def forward(self, x):\n        identity = x\n        out = self.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        out = self.se(out)\n\n        if self.downsample is not None:\n            identity = self.downsample(x)\n\n        return self.relu(out + identity)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:31:41.923678Z","iopub.execute_input":"2025-05-15T12:31:41.924161Z","iopub.status.idle":"2025-05-15T12:31:41.930191Z","shell.execute_reply.started":"2025-05-15T12:31:41.924136Z","shell.execute_reply":"2025-05-15T12:31:41.929379Z"}},"outputs":[],"execution_count":15},{"cell_type":"code","source":"from torch.hub import load_state_dict_from_url\n\n\nclass SEResNet18(nn.Module):\n    def __init__(self, num_classes=12):\n        super().__init__()\n        self.in_channels = 64\n\n        # First layer: STRIDE=1, no maxpool\n        self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=1, padding=3, bias=False)\n        self.bn1 = nn.BatchNorm2d(64)\n        self.relu = nn.ReLU(inplace=True)\n\n        self.layer1 = self._make_layer(64, 2, stride=1)\n        self.layer2 = self._make_layer(128, 2, stride=2)\n        self.layer3 = self._make_layer(256, 2, stride=2)\n        self.layer4 = self._make_layer(512, 2, stride=2)\n\n        self.avgpool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(512, num_classes)\n\n    def _make_layer(self, out_channels, blocks, stride):\n        downsample = None\n        if stride != 1 or self.in_channels != out_channels:\n            downsample = nn.Sequential(\n                nn.Conv2d(self.in_channels, out_channels, 1, stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n        layers = [\n            SEBasicBlock(self.in_channels, out_channels, stride, downsample)\n        ]\n        self.in_channels = out_channels\n        for _ in range(1, blocks):\n            layers.append(SEBasicBlock(out_channels, out_channels))\n        return nn.Sequential(*layers)\n\n    def forward(self, x):\n        x = self.relu(self.bn1(self.conv1(x)))\n        # no maxpool\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n\n        feat = self.avgpool(x).flatten(1)\n        out = self.fc(feat)\n        return out\n\n    def extract_features(self, x):\n        x = self.relu(self.bn1(self.conv1(x)))\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n        return self.avgpool(x).flatten(1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:31:45.401801Z","iopub.execute_input":"2025-05-15T12:31:45.402373Z","iopub.status.idle":"2025-05-15T12:31:45.412182Z","shell.execute_reply.started":"2025-05-15T12:31:45.402345Z","shell.execute_reply":"2025-05-15T12:31:45.411416Z"}},"outputs":[],"execution_count":16},{"cell_type":"code","source":"model = SEResNet18(num_classes=12)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:31:49.01141Z","iopub.execute_input":"2025-05-15T12:31:49.012076Z","iopub.status.idle":"2025-05-15T12:31:49.137008Z","shell.execute_reply.started":"2025-05-15T12:31:49.012051Z","shell.execute_reply":"2025-05-15T12:31:49.136225Z"}},"outputs":[],"execution_count":17},{"cell_type":"code","source":"class DCTResNet(nn.Module):\n    def __init__(self, num_classes=12, in_channels=303):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_channels, 64, 3, padding=1)\n        self.bn1 = nn.BatchNorm2d(64)\n        self.relu = nn.ReLU(inplace=True)\n\n        self.conv2 = nn.Conv2d(64, 64, 3, padding=1)\n        self.bn2 = nn.BatchNorm2d(64)\n        self.conv3 = nn.Conv2d(64, 64, 3, padding=1)\n        self.bn3 = nn.BatchNorm2d(64)\n        self.se1 = SEBlock(64)\n\n        self.conv4 = nn.Conv2d(64, 64, 3, padding=1)\n        self.bn4 = nn.BatchNorm2d(64)\n        self.conv5 = nn.Conv2d(64, 64, 3, padding=1)\n        self.bn5 = nn.BatchNorm2d(64)\n        self.se2 = SEBlock(64)\n\n        self.conv6 = nn.Conv2d(64, 64, 3, padding=1)\n        self.bn6 = nn.BatchNorm2d(64)\n        self.conv7 = nn.Conv2d(64, 64, 3, padding=1)\n        self.bn7 = nn.BatchNorm2d(64)\n        self.se3 = SEBlock(64)\n\n        self.avgpool = nn.AdaptiveAvgPool2d(1)\n        self.fc = nn.Linear(64, num_classes)\n\n    def forward(self, x):\n        x = self.relu(self.bn1(self.conv1(x)))\n\n        # Block 1\n        identity = x\n        x = self.relu(self.bn2(self.conv2(x)))\n        x = self.bn3(self.conv3(x))\n        x = self.se1(x) + identity\n        x = self.relu(x)\n\n        # Block 2\n        identity = x\n        x = self.relu(self.bn4(self.conv4(x)))\n        x = self.bn5(self.conv5(x))\n        x = self.se2(x) + identity\n        x = self.relu(x)\n\n        # Block 3\n        identity = x\n        x = self.relu(self.bn6(self.conv6(x)))\n        x = self.bn7(self.conv7(x))\n        x = self.se3(x) + identity\n        x = self.relu(x)\n\n        x = self.avgpool(x).flatten(1)\n        return self.fc(x)\n\n    def extract_features(self, x):\n        x = self.relu(self.bn1(self.conv1(x)))\n\n        identity = x\n        x = self.relu(self.bn2(self.conv2(x)))\n        x = self.bn3(self.conv3(x))\n        x = self.se1(x) + identity\n        x = self.relu(x)\n\n        identity = x\n        x = self.relu(self.bn4(self.conv4(x)))\n        x = self.bn5(self.conv5(x))\n        x = self.se2(x) + identity\n        x = self.relu(x)\n\n        identity = x\n        x = self.relu(self.bn6(self.conv6(x)))\n        x = self.bn7(self.conv7(x))\n        x = self.se3(x) + identity\n        x = self.relu(x)\n\n        return self.avgpool(x).flatten(1)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:31:52.705948Z","iopub.execute_input":"2025-05-15T12:31:52.706607Z","iopub.status.idle":"2025-05-15T12:31:52.841695Z","shell.execute_reply.started":"2025-05-15T12:31:52.706581Z","shell.execute_reply":"2025-05-15T12:31:52.84099Z"}},"outputs":[],"execution_count":19},{"cell_type":"code","source":"model = DCTResNet(num_classes=12)\nprint(model(torch.randn(1, 303, 64, 64)).shape)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:31:56.933545Z","iopub.execute_input":"2025-05-15T12:31:56.934008Z","iopub.status.idle":"2025-05-15T12:31:57.109101Z","shell.execute_reply.started":"2025-05-15T12:31:56.933984Z","shell.execute_reply":"2025-05-15T12:31:57.108302Z"}},"outputs":[{"name":"stdout","text":"torch.Size([1, 12])\n","output_type":"stream"}],"execution_count":20},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score\nimport numpy as np\n\ndef compute_auc(y_true, y_score):\n    \"\"\"\n    Computes weighted multi-class AUC.\n    Handles the case when not all classes are present in y_true.\n\n    y_true: List or array of shape [N], with class indices (e.g. 0-11)\n    y_score: Array of shape [N, C], softmax probabilities over all classes\n    \"\"\"\n    try:\n        y_true = np.array(y_true)\n        y_score = np.array(y_score)\n        num_classes = y_score.shape[1]\n\n        # One-hot encode y_true into shape [N, C]\n        y_true_onehot = np.zeros((len(y_true), num_classes))\n        y_true_onehot[np.arange(len(y_true)), y_true] = 1\n\n        return roc_auc_score(y_true_onehot, y_score, average=\"weighted\")\n    except Exception as e:\n        print(\"⚠️ AUC computation failed:\", e)\n        return 0.0\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:31:59.183889Z","iopub.execute_input":"2025-05-15T12:31:59.184166Z","iopub.status.idle":"2025-05-15T12:31:59.189677Z","shell.execute_reply.started":"2025-05-15T12:31:59.184144Z","shell.execute_reply":"2025-05-15T12:31:59.188894Z"}},"outputs":[],"execution_count":21},{"cell_type":"code","source":"#===========================Training with checkpoints=========================\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport numpy as np\nfrom tqdm import tqdm\nimport os\nimport gc\n\ndef train_model(model, train_loader, val_loader, config, start_epoch=0):\n    device = config.get(\"device\", \"cuda\")\n    epochs = config.get(\"epochs\", 10)\n    lr = config.get(\"lr\", 1e-4)\n    alpha = config.get(\"alpha\", 1.0)\n    mix_type = config.get(\"mix_type\", None)  # None, \"cutmix\", \"mixup\"\n    save_path = config.get(\"save_path\", \"checkpoint.pth\")\n    grad_acc_steps = config.get(\"grad_acc_steps\", 1)\n\n    model = model.to(device)\n    criterion = nn.CrossEntropyLoss()\n    optimizer = optim.Adam(model.parameters(), lr=lr)\n\n    best_auc = 0.0\n    for epoch in range(start_epoch, epochs):\n        model.train()\n        total_loss = 0.0\n        optimizer.zero_grad()\n\n        for step, (images, labels) in enumerate(tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs}\")):\n            images, labels = images.to(device), labels.to(device)\n\n            if mix_type in [\"cutmix\", \"mixup\"]:\n                lam = np.random.beta(alpha, alpha)\n                rand_index = torch.randperm(images.size(0)).to(device)\n                images2 = images[rand_index]\n                labels2 = labels[rand_index]\n\n                if mix_type == \"cutmix\":\n                    _, _, H, W = images.shape\n                    rx, ry = np.random.randint(W), np.random.randint(H)\n                    rw = int(W * np.sqrt(1 - lam))\n                    rh = int(H * np.sqrt(1 - lam))\n                    x1 = np.clip(rx - rw // 2, 0, W)\n                    y1 = np.clip(ry - rh // 2, 0, H)\n                    x2 = np.clip(rx + rw // 2, 0, W)\n                    y2 = np.clip(ry + rh // 2, 0, H)\n                    images[:, :, y1:y2, x1:x2] = images2[:, :, y1:y2, x1:x2]\n                    lam = 1 - ((x2 - x1) * (y2 - y1) / (W * H))\n                elif mix_type == \"mixup\":\n                    images = lam * images + (1 - lam) * images2\n\n                outputs = model(images)\n                loss = lam * criterion(outputs, labels) + (1 - lam) * criterion(outputs, labels2)\n            else:\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n\n            loss = loss / grad_acc_steps\n            loss.backward()\n\n            if (step + 1) % grad_acc_steps == 0 or (step + 1) == len(train_loader):\n                optimizer.step()\n                optimizer.zero_grad()\n\n            total_loss += loss.item() * grad_acc_steps  # reverse scaled for logging\n\n        print(f\"[Epoch {epoch+1}] Loss: {total_loss/len(train_loader):.4f}\")\n\n        # Validation\n        model.eval()\n        all_preds, all_labels = [], []\n        with torch.no_grad():\n            for images, labels in val_loader:\n                images = images.to(device)\n                labels = labels.to(device)\n                outputs = model(images)\n                probs = torch.softmax(outputs, dim=1)\n                all_preds.extend(probs.cpu().numpy())\n                all_labels.extend(labels.cpu().numpy())\n\n        val_auc = compute_auc(all_labels, all_preds)\n        print(f\"Validation AUC: {val_auc:.4f}\")\n\n        # Save current epoch checkpoint\n        epoch_path = save_path.replace(\".pth\", f\"_epoch{epoch+1}.pth\")\n        torch.save(model.state_dict(), epoch_path)\n        print(f\"✅ Epoch {epoch+1} checkpoint saved: {epoch_path}\")\n\n        # Save best model\n        if val_auc > best_auc:\n            best_auc = val_auc\n            torch.save(model.state_dict(), save_path)\n            print(f\"✅ ✅ New best model saved: {save_path}\")\n\n        print(\"✅ Gradient flow OK\\n\")\n\n        # free memory\n        gc.collect()\n        torch.cuda.empty_cache()\n\n    return model\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:32:02.031533Z","iopub.execute_input":"2025-05-15T12:32:02.032401Z","iopub.status.idle":"2025-05-15T12:32:02.044918Z","shell.execute_reply.started":"2025-05-15T12:32:02.032375Z","shell.execute_reply":"2025-05-15T12:32:02.044139Z"}},"outputs":[],"execution_count":22},{"cell_type":"code","source":"# import torch\n# import torch.nn as nn\n# import torch.optim as optim\n# import numpy as np\n# from tqdm import tqdm\n# import os\n\n# def train_model(model, train_loader, val_loader, config):\n#     device = config.get(\"device\", \"cuda\")\n#     epochs = config.get(\"epochs\", 10)\n#     lr = config.get(\"lr\", 1e-4)\n#     alpha = config.get(\"alpha\", 1.0)\n#     mix_type = config.get(\"mix_type\", None)  # None, \"cutmix\", \"mixup\"\n#     save_path = config.get(\"save_path\", \"checkpoint.pth\")\n#     grad_acc_steps = config.get(\"grad_acc_steps\", 1)  # NEW\n\n#     model = model.to(device)\n#     criterion = nn.CrossEntropyLoss()\n#     optimizer = optim.Adam(model.parameters(), lr=lr)\n\n#     best_auc = 0.0\n#     for epoch in range(epochs):\n#         model.train()\n#         total_loss = 0.0\n#         optimizer.zero_grad()\n\n#         for step, (images, labels) in enumerate(tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs}\")):\n#             images, labels = images.to(device), labels.to(device)\n\n#             if mix_type in [\"cutmix\", \"mixup\"]:\n#                 lam = np.random.beta(alpha, alpha)\n#                 rand_index = torch.randperm(images.size(0)).to(device)\n#                 images2 = images[rand_index]\n#                 labels2 = labels[rand_index]\n\n#                 if mix_type == \"cutmix\":\n#                     _, _, H, W = images.shape\n#                     rx, ry = np.random.randint(W), np.random.randint(H)\n#                     rw = int(W * np.sqrt(1 - lam))\n#                     rh = int(H * np.sqrt(1 - lam))\n#                     x1 = np.clip(rx - rw // 2, 0, W)\n#                     y1 = np.clip(ry - rh // 2, 0, H)\n#                     x2 = np.clip(rx + rw // 2, 0, W)\n#                     y2 = np.clip(ry + rh // 2, 0, H)\n#                     images[:, :, y1:y2, x1:x2] = images2[:, :, y1:y2, x1:x2]\n#                     lam = 1 - ((x2 - x1) * (y2 - y1) / (W * H))\n#                 elif mix_type == \"mixup\":\n#                     images = lam * images + (1 - lam) * images2\n\n#                 outputs = model(images)\n#                 loss = lam * criterion(outputs, labels) + (1 - lam) * criterion(outputs, labels2)\n#             else:\n#                 outputs = model(images)\n#                 loss = criterion(outputs, labels)\n\n#             # Gradient accumulation logic\n#             loss = loss / grad_acc_steps\n#             loss.backward()\n\n#             if (step + 1) % grad_acc_steps == 0 or (step + 1) == len(train_loader):\n#                 optimizer.step()\n#                 optimizer.zero_grad()\n\n#             total_loss += loss.item() * grad_acc_steps  # reverse scaling for logging\n\n#         print(f\"[Epoch {epoch+1}] Loss: {total_loss/len(train_loader):.4f}\")\n\n#         # Validation\n#         model.eval()\n#         all_preds, all_labels = [], []\n#         with torch.no_grad():\n#             for images, labels in val_loader:\n#                 images = images.to(device)\n#                 labels = labels.to(device)\n#                 outputs = model(images)\n#                 probs = torch.softmax(outputs, dim=1)\n#                 all_preds.extend(probs.cpu().numpy())\n#                 all_labels.extend(labels.cpu().numpy())\n\n#         val_auc = compute_auc(all_labels, all_preds)\n#         print(f\"Validation AUC: {val_auc:.4f}\")\n\n#         if val_auc > best_auc:\n#             best_auc = val_auc\n#             torch.save(model.state_dict(), save_path)\n#             print(f\"✅ Checkpoint saved: {save_path}\")\n\n#         print(\"✅ Gradient flow OK\\n\")\n\n#     return model\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import DataLoader\n\n# Use split3\ntrain_df = splits[\"split3\"][\"train\"]\nval_df = splits[\"split3\"][\"val\"]\n\n# Create datasets (SPATIAL mode)\ntrain_ds = Alaska2Dataset(train_df, mode=\"spatial\")\nval_ds = Alaska2Dataset(val_df, mode=\"spatial\")\n\n# DataLoaders\ntrain_loader = DataLoader(train_ds, batch_size=8, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_ds, batch_size=8, shuffle=False, num_workers=2)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:32:08.643276Z","iopub.execute_input":"2025-05-15T12:32:08.64354Z","iopub.status.idle":"2025-05-15T12:32:08.652612Z","shell.execute_reply.started":"2025-05-15T12:32:08.64352Z","shell.execute_reply":"2025-05-15T12:32:08.651853Z"}},"outputs":[],"execution_count":23},{"cell_type":"code","source":"model = SEResNet18(num_classes=12)\n\nconfig = {\n    \"device\": \"cuda\",\n    \"epochs\": 12,\n    \"lr\": 1e-4,\n    \"alpha\": 1.0,\n    \"mix_type\": \"cutmix\",  # or \"mixup\" or None\n    \"save_path\": \"seresnet_split3.pth\",\n    \"grad_acc_steps\": 2\n}\n\n# --------- EĞİTİME KALDIĞI YERDEN DEVAM ETMEK İÇİN ---------\n# model.load_state_dict(torch.load(\"seresnet_split3_epoch1.pth\"))\n# train_model(model, train_loader, val_loader, config, start_epoch=1)\n\n# --------- EĞİTİME SIFIRDAN BAŞLAMAK İÇİN ---------\n#train_model(model, train_loader, val_loader, config)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:32:12.6587Z","iopub.execute_input":"2025-05-15T12:32:12.658951Z","iopub.status.idle":"2025-05-15T12:32:12.751084Z","shell.execute_reply.started":"2025-05-15T12:32:12.658932Z","shell.execute_reply":"2025-05-15T12:32:12.75035Z"}},"outputs":[],"execution_count":24},{"cell_type":"code","source":"import os\n\n# Kayıt yapılması beklenen dizin genellikle çalışma dizinidir\nfor file in os.listdir():\n    if file.endswith(\".pth\") or \"seresnet\" in file:\n        print(file)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T14:33:41.049024Z","iopub.execute_input":"2025-05-15T14:33:41.049748Z","iopub.status.idle":"2025-05-15T14:33:41.055014Z","shell.execute_reply.started":"2025-05-15T14:33:41.049718Z","shell.execute_reply":"2025-05-15T14:33:41.054184Z"}},"outputs":[{"name":"stdout","text":"seresnet_split3_epoch1.pth\nseresnet_split3_epoch4.pth\nseresnet_split3_epoch10.pth\nseresnet_split3_epoch5.pth\nseresnet_split3_epoch9.pth\nseresnet_split3_epoch3.pth\nseresnet_split3_epoch11.pth\nseresnet_split3_epoch7.pth\nseresnet_split3_epoch12.pth\nseresnet_split3.pth\nseresnet_split3_epoch6.pth\nseresnet_split3_epoch8.pth\n","output_type":"stream"}],"execution_count":27},{"cell_type":"code","source":"model.load_state_dict(torch.load(\"seresnet_split3_epoch11.pth\"))\ntrain_model(model, train_loader, val_loader, config, start_epoch=11)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T12:32:21.376367Z","iopub.execute_input":"2025-05-15T12:32:21.376973Z","iopub.status.idle":"2025-05-15T14:23:16.564409Z","shell.execute_reply.started":"2025-05-15T12:32:21.376947Z","shell.execute_reply":"2025-05-15T14:23:16.563621Z"}},"outputs":[{"name":"stderr","text":"Epoch 12/12: 100%|██████████| 8125/8125 [1:45:42<00:00,  1.28it/s]  ","output_type":"stream"},{"name":"stdout","text":"[Epoch 12] Loss: 2.3851\n","output_type":"stream"},{"name":"stderr","text":"\n","output_type":"stream"},{"name":"stdout","text":"Validation AUC: 0.6894\n✅ Epoch 12 checkpoint saved: seresnet_split3_epoch12.pth\n✅ ✅ New best model saved: seresnet_split3.pth\n✅ Gradient flow OK\n\n","output_type":"stream"},{"execution_count":26,"output_type":"execute_result","data":{"text/plain":"SEResNet18(\n  (conv1): Conv2d(3, 64, kernel_size=(7, 7), stride=(1, 1), padding=(3, 3), bias=False)\n  (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n  (relu): ReLU(inplace=True)\n  (layer1): Sequential(\n    (0): SEBasicBlock(\n      (conv1): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (relu): ReLU(inplace=True)\n      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (se): SEBlock(\n        (fc1): Linear(in_features=64, out_features=4, bias=True)\n        (fc2): Linear(in_features=4, out_features=64, bias=True)\n      )\n    )\n    (1): SEBasicBlock(\n      (conv1): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn1): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (relu): ReLU(inplace=True)\n      (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn2): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (se): SEBlock(\n        (fc1): Linear(in_features=64, out_features=4, bias=True)\n        (fc2): Linear(in_features=4, out_features=64, bias=True)\n      )\n    )\n  )\n  (layer2): Sequential(\n    (0): SEBasicBlock(\n      (conv1): Conv2d(64, 128, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n      (bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (relu): ReLU(inplace=True)\n      (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (se): SEBlock(\n        (fc1): Linear(in_features=128, out_features=8, bias=True)\n        (fc2): Linear(in_features=8, out_features=128, bias=True)\n      )\n      (downsample): Sequential(\n        (0): Conv2d(64, 128, kernel_size=(1, 1), stride=(2, 2), bias=False)\n        (1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      )\n    )\n    (1): SEBasicBlock(\n      (conv1): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn1): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (relu): ReLU(inplace=True)\n      (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn2): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (se): SEBlock(\n        (fc1): Linear(in_features=128, out_features=8, bias=True)\n        (fc2): Linear(in_features=8, out_features=128, bias=True)\n      )\n    )\n  )\n  (layer3): Sequential(\n    (0): SEBasicBlock(\n      (conv1): Conv2d(128, 256, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n      (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (relu): ReLU(inplace=True)\n      (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (se): SEBlock(\n        (fc1): Linear(in_features=256, out_features=16, bias=True)\n        (fc2): Linear(in_features=16, out_features=256, bias=True)\n      )\n      (downsample): Sequential(\n        (0): Conv2d(128, 256, kernel_size=(1, 1), stride=(2, 2), bias=False)\n        (1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      )\n    )\n    (1): SEBasicBlock(\n      (conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn1): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (relu): ReLU(inplace=True)\n      (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn2): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (se): SEBlock(\n        (fc1): Linear(in_features=256, out_features=16, bias=True)\n        (fc2): Linear(in_features=16, out_features=256, bias=True)\n      )\n    )\n  )\n  (layer4): Sequential(\n    (0): SEBasicBlock(\n      (conv1): Conv2d(256, 512, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1), bias=False)\n      (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (relu): ReLU(inplace=True)\n      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (se): SEBlock(\n        (fc1): Linear(in_features=512, out_features=32, bias=True)\n        (fc2): Linear(in_features=32, out_features=512, bias=True)\n      )\n      (downsample): Sequential(\n        (0): Conv2d(256, 512, kernel_size=(1, 1), stride=(2, 2), bias=False)\n        (1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      )\n    )\n    (1): SEBasicBlock(\n      (conv1): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn1): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (relu): ReLU(inplace=True)\n      (conv2): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1), bias=False)\n      (bn2): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)\n      (se): SEBlock(\n        (fc1): Linear(in_features=512, out_features=32, bias=True)\n        (fc2): Linear(in_features=32, out_features=512, bias=True)\n      )\n    )\n  )\n  (avgpool): AdaptiveAvgPool2d(output_size=1)\n  (fc): Linear(in_features=512, out_features=12, bias=True)\n)"},"metadata":{}}],"execution_count":26},{"cell_type":"code","source":"torch.save(model.state_dict(), \"/kaggle/working/seresnet_split3_epoch3.pth\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T20:46:23.574473Z","iopub.execute_input":"2025-05-14T20:46:23.574996Z","iopub.status.idle":"2025-05-14T20:46:23.715088Z","shell.execute_reply.started":"2025-05-14T20:46:23.574973Z","shell.execute_reply":"2025-05-14T20:46:23.714532Z"}},"outputs":[],"execution_count":33},{"cell_type":"code","source":"import os\n\nfor file in os.listdir(\"/kaggle/working\"):\n    print(file)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-15T10:02:54.066901Z","iopub.execute_input":"2025-05-15T10:02:54.067159Z","iopub.status.idle":"2025-05-15T10:02:54.07143Z","shell.execute_reply.started":"2025-05-15T10:02:54.067141Z","shell.execute_reply":"2025-05-15T10:02:54.070876Z"}},"outputs":[{"name":"stdout","text":"seresnet_split3_epoch9.pth\nseresnet_split3_epoch8.pth\nseresnet_split3_epoch5.pth\nseresnet_split3_epoch3.pth\nseresnet_split3_epoch7.pth\nseresnet_split3_epoch10.pth\nseresnet_split3.pth\n.virtual_documents\nseresnet_split3_epoch1.pth\nseresnet_split3_epoch6.pth\nseresnet_split3_epoch4.pth\n","output_type":"stream"}],"execution_count":36},{"cell_type":"code","source":"model.load_state_dict(torch.load(\"seresnet_split3_epoch5.pth\"))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T20:49:19.573857Z","iopub.execute_input":"2025-05-14T20:49:19.574455Z","iopub.status.idle":"2025-05-14T20:49:19.648244Z","shell.execute_reply.started":"2025-05-14T20:49:19.574431Z","shell.execute_reply":"2025-05-14T20:49:19.647586Z"}},"outputs":[{"execution_count":35,"output_type":"execute_result","data":{"text/plain":"<All keys matched successfully>"},"metadata":{}}],"execution_count":35},{"cell_type":"code","source":"!ls -lh | grep epoch5\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-14T20:49:30.69153Z","iopub.execute_input":"2025-05-14T20:49:30.692073Z","iopub.status.idle":"2025-05-14T20:49:30.858637Z","shell.execute_reply.started":"2025-05-14T20:49:30.692047Z","shell.execute_reply":"2025-05-14T20:49:30.857912Z"}},"outputs":[{"name":"stdout","text":"-rw-r--r-- 1 root root 44M May 14 20:30 seresnet_split3_epoch5.pth\n","output_type":"stream"}],"execution_count":36},{"cell_type":"code","source":"# model = SEResNet18(num_classes=12)\n# config = {\n#     \"device\": \"cuda\",\n#     \"epochs\": 5,\n#     \"lr\": 1e-4,\n#     \"alpha\": 1.0,\n#     \"mix_type\": \"cutmix\",  # or \"mixup\" or None\n#     \"save_path\": \"seresnet_split3.pth\",\n#     \"grad_acc_steps\": 8\n# }\n# train_model(model, train_loader, val_loader, config)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}