{"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":[{"sourceId":113558,"databundleVersionId":14456136,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-08T22:34:54.128991Z","iopub.execute_input":"2025-12-08T22:34:54.129642Z","iopub.status.idle":"2025-12-08T22:35:02.575937Z","shell.execute_reply.started":"2025-12-08T22:34:54.129614Z","shell.execute_reply":"2025-12-08T22:35:02.575129Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#检查数据平衡\nimport os\n\ntrain_dir = '/kaggle/input/recodai-luc-scientific-image-forgery-detection/train_images'\n\nnum_forged = len(os.listdir(os.path.join(train_dir, 'forged')))\nnum_authentic = len(os.listdir(os.path.join(train_dir, 'authentic')))\n\nprint(\"Forged images:\", num_forged)\nprint(\"Authentic images:\", num_authentic)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:18:46.780378Z","iopub.execute_input":"2025-12-09T16:18:46.780565Z","iopub.status.idle":"2025-12-09T16:18:46.838208Z","shell.execute_reply.started":"2025-12-09T16:18:46.780549Z","shell.execute_reply":"2025-12-09T16:18:46.837450Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport random\n\n#查看部分数据，规范数据为01分型\ntrain_dir = '/kaggle/input/recodai-luc-scientific-image-forgery-detection/train_images'\n\nimage_paths = []\nlabels = []\n\nfor label_name, label_id in [('forged', 1), ('authentic', 0)]:\n    folder = os.path.join(train_dir, label_name)\n    for fname in os.listdir(folder):\n        if fname.lower().endswith(('.png', '.jpg', '.jpeg', '.tif')):\n            image_paths.append(os.path.join(folder, fname))\n            labels.append(label_id)\n\ndf = pd.DataFrame({\n    'image_path': image_paths,\n    'label': labels\n})\n\nprint(df.shape)\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:22:51.158199Z","iopub.execute_input":"2025-12-09T16:22:51.158467Z","iopub.status.idle":"2025-12-09T16:22:51.440507Z","shell.execute_reply.started":"2025-12-09T16:22:51.158446Z","shell.execute_reply":"2025-12-09T16:22:51.439754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df['label'].value_counts(), df['label'].value_counts(normalize=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:22:53.550585Z","iopub.execute_input":"2025-12-09T16:22:53.551182Z","iopub.status.idle":"2025-12-09T16:22:53.564225Z","shell.execute_reply.started":"2025-12-09T16:22:53.551153Z","shell.execute_reply":"2025-12-09T16:22:53.563435Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n#数据分布直方图\ndf['label'].replace({0: 'authentic', 1: 'forged'}).value_counts().plot(kind='bar')\nplt.ylabel('count')\nplt.title('Class distribution')\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:22:55.344591Z","iopub.execute_input":"2025-12-09T16:22:55.344862Z","iopub.status.idle":"2025-12-09T16:22:55.575437Z","shell.execute_reply.started":"2025-12-09T16:22:55.344840Z","shell.execute_reply":"2025-12-09T16:22:55.574815Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from PIL import Image\n#预览数据集中的图片\ndef show_examples(df, label, n=4):\n    subset = df[df['label'] == label].sample(n, random_state=42)\n    plt.figure(figsize=(10, 3))\n    for i, (_, row) in enumerate(subset.iterrows(), 1):\n        img = Image.open(row['image_path'])\n        plt.subplot(1, n, i)\n        plt.imshow(img, cmap='gray')\n        plt.axis('off')\n        plt.title('forged' if label == 1 else 'authentic')\n    plt.show()\n\nshow_examples(df, 1)  # forged\nshow_examples(df, 0)  # authentic\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:22:57.587154Z","iopub.execute_input":"2025-12-09T16:22:57.587701Z","iopub.status.idle":"2025-12-09T16:22:59.969152Z","shell.execute_reply.started":"2025-12-09T16:22:57.587676Z","shell.execute_reply":"2025-12-09T16:22:59.968449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\nfrom sklearn.model_selection import train_test_split\n\n# 如果之前还没保存 df，这里假设 df 已经在内存中\n# df: columns = ['image_path', 'label']\n\ntrain_df, val_df = train_test_split(\n    df,\n    test_size=0.2,\n    stratify=df['label'],\n    random_state=42\n)\n\nprint(len(train_df), len(val_df))\n\n# 图像增强 & 预处理（ImageNet 风格）\ntrain_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(10),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n\nval_transform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                         std=[0.229, 0.224, 0.225])\n])\n\nclass ImageForgeryDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df.reset_index(drop=True)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img = Image.open(row['image_path']).convert('RGB')\n        if self.transform:\n            img = self.transform(img)\n        label = torch.tensor(row['label'], dtype=torch.float32)\n        return img, label\n\ntrain_dataset = ImageForgeryDataset(train_df, transform=train_transform)\nval_dataset = ImageForgeryDataset(val_df, transform=val_transform)\n\ntrain_loader = DataLoader(train_dataset, batch_size=32, shuffle=True, num_workers=2)\nval_loader = DataLoader(val_dataset, batch_size=32, shuffle=False, num_workers=2)\n\nlen(train_loader), len(val_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:23:01.214199Z","iopub.execute_input":"2025-12-09T16:23:01.214678Z","iopub.status.idle":"2025-12-09T16:23:08.407951Z","shell.execute_reply.started":"2025-12-09T16:23:01.214653Z","shell.execute_reply":"2025-12-09T16:23:08.407197Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Baseline Model：ResNet18","metadata":{}},{"cell_type":"code","source":"import torch.nn as nn\nfrom torchvision import models\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\ndevice","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:23:11.290910Z","iopub.execute_input":"2025-12-09T16:23:11.291773Z","iopub.status.idle":"2025-12-09T16:23:11.367020Z","shell.execute_reply.started":"2025-12-09T16:23:11.291741Z","shell.execute_reply":"2025-12-09T16:23:11.366353Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 加载预训练 ResNet18\nresnet = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)\n\n# 替换最后一层为二分类（输出一个 logit）\nin_features = resnet.fc.in_features\nresnet.fc = nn.Linear(in_features, 1)\n\nresnet = resnet.to(device)\n\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(resnet.parameters(), lr=1e-4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:23:13.551330Z","iopub.execute_input":"2025-12-09T16:23:13.551596Z","iopub.status.idle":"2025-12-09T16:23:14.303559Z","shell.execute_reply.started":"2025-12-09T16:23:13.551575Z","shell.execute_reply":"2025-12-09T16:23:14.302771Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm\n\ndef train_one_epoch(model, loader, optimizer, criterion, device):\n    model.train()\n    running_loss = 0.0\n    for imgs, labels in tqdm(loader, leave=False):\n        imgs = imgs.to(device)\n        labels = labels.to(device).unsqueeze(1)  # (B,) -> (B,1)\n\n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * imgs.size(0)\n    return running_loss / len(loader.dataset)\n\ndef eval_one_epoch(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs = imgs.to(device)\n            labels = labels.to(device).unsqueeze(1)\n\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            running_loss += loss.item() * imgs.size(0)\n\n            preds = (torch.sigmoid(outputs) > 0.5).float()\n            correct += (preds == labels).sum().item()\n            total += labels.size(0)\n    return running_loss / len(loader.dataset), correct / total\n\nnum_epochs = 5\nfor epoch in range(1, num_epochs + 1):\n    train_loss = train_one_epoch(resnet, train_loader, optimizer, criterion, device)\n    val_loss, val_acc = eval_one_epoch(resnet, val_loader, criterion, device)\n    print(f\"Epoch {epoch}: train_loss={train_loss:.4f}, val_loss={val_loss:.4f}, val_acc={val_acc:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:23:15.857021Z","iopub.execute_input":"2025-12-09T16:23:15.857513Z","iopub.status.idle":"2025-12-09T16:32:59.143999Z","shell.execute_reply.started":"2025-12-09T16:23:15.857489Z","shell.execute_reply":"2025-12-09T16:32:59.143097Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### 使用train_mask","metadata":{}},{"cell_type":"code","source":"import os\n#检查trai_mask的数据\nmask_dir = '/kaggle/input/recodai-luc-scientific-image-forgery-detection/train_masks'\n\nfor root, dirs, files in os.walk(mask_dir):\n    print('ROOT:', root)\n    print('  subdirs:', dirs[:5])\n    print('  num_files_here:', len(files))\n    print('  sample_files:', files[:5])\n    # 只看第一层就停\n    break\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:39:49.699718Z","iopub.execute_input":"2025-12-09T16:39:49.700471Z","iopub.status.idle":"2025-12-09T16:39:52.202676Z","shell.execute_reply.started":"2025-12-09T16:39:49.700437Z","shell.execute_reply":"2025-12-09T16:39:52.201922Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nimport numpy as np\nfrom PIL import Image\nimport matplotlib.pyplot as plt\nimport os\n\ntrain_img_dir = '/kaggle/input/recodai-luc-scientific-image-forgery-detection/train_images'\ntrain_mask_dir = '/kaggle/input/recodai-luc-scientific-image-forgery-detection/train_masks'\n\n# 随机抽一张 forged 图像\nsample_row = df[df['label'] == 1].sample(1, random_state=random.randint(0, 10_000)).iloc[0]\nimg_path = sample_row['image_path']\nimg_name = os.path.basename(img_path)          # e.g. \"59069.png\"\nstem = os.path.splitext(img_name)[0]           # \"59069\"\n\nmask_path = os.path.join(train_mask_dir, stem + '.npy')\n\nprint(\"Image path:\", img_path)\nprint(\"Mask path :\", mask_path)\nprint(\"Mask exists?\", os.path.exists(mask_path))\n\n# 读图 & mask\nimg = Image.open(img_path).convert('RGB')\nimg_np = np.array(img)\n\nmask_np = np.load(mask_path)\nprint(\"Raw mask shape:\", mask_np.shape, \"dtype:\", mask_np.dtype)\n\n# ---- 把 mask 转成 2D 灰度 ----\nif mask_np.ndim == 2:\n    mask_2d = mask_np\nelif mask_np.ndim == 3:\n    # 如果是 (C, H, W)\n    if mask_np.shape[0] in [1, 3]:\n        mask_2d = mask_np.mean(axis=0)\n    # 如果是 (H, W, C)\n    elif mask_np.shape[-1] in [1, 3]:\n        mask_2d = mask_np.mean(axis=-1)\n    else:\n        # 兜底：在第 0 维求均值\n        mask_2d = mask_np.mean(axis=0)\nelse:\n    # 形状更奇怪时，先压扁再 reshape 成和图像大小一样（不太可能用到）\n    mask_2d = mask_np.reshape(img_np.shape[0], img_np.shape[1])\n\nprint(\"Mask 2D shape:\", mask_2d.shape)\n\n# 归一化到 [0,1]\nmask_norm = mask_2d.astype(np.float32)\nif mask_norm.max() > 0:\n    mask_norm /= mask_norm.max()\n\n# 构造叠加：红色区域为伪造\noverlay = img_np.astype(np.float32) / 255.0\nred_layer = np.zeros_like(overlay)\nred_layer[..., 0] = 1.0  # R 通道为 1\n\nalpha = 0.6\noverlay = overlay * (1 - mask_norm[..., None]) + red_layer * (mask_norm[..., None] * alpha)\n\n# 画三张图\nplt.figure(figsize=(12, 4))\nplt.subplot(1, 3, 1)\nplt.imshow(img_np)\nplt.axis('off')\nplt.title('Forged image')\n\nplt.subplot(1, 3, 2)\nplt.imshow(mask_norm, cmap='gray')\nplt.axis('off')\nplt.title('Mask (2D)')\n\nplt.subplot(1, 3, 3)\nplt.imshow(overlay)\nplt.axis('off')\nplt.title('Overlay')\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:39:53.790789Z","iopub.execute_input":"2025-12-09T16:39:53.791649Z","iopub.status.idle":"2025-12-09T16:39:54.399757Z","shell.execute_reply.started":"2025-12-09T16:39:53.791614Z","shell.execute_reply":"2025-12-09T16:39:54.399019Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### patch level数据集","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nfrom PIL import Image\nfrom skimage.measure import label, regionprops\n\ndef extract_forged_patches(img_path, mask_path, patch_sizes=[128, 160], max_patches=5):\n    \"\"\"\n    输入: 原图路径 + mask 路径\n    输出: patch 列表 [(patch_array, label=1), ...]\n    \"\"\"\n    img = np.array(Image.open(img_path).convert('RGB'))\n    H, W = img.shape[:2]\n\n    mask = np.load(mask_path)\n\n    # ---- 把 mask 稳定地转成 2D ----\n    if mask.ndim == 2:\n        mask_2d = mask\n    elif mask.ndim == 3:\n        # 如果有某一维是 1（比如 (1,256,320) 或 (256,320,1)），就 squeeze 掉\n        if 1 in mask.shape:\n            mask_2d = np.squeeze(mask)\n        else:\n            # 否则对通道求平均，得到 2D\n            # 假设是 (C,H,W) 或 (H,W,C)，都兼容\n            if mask.shape[0] in [3, 4]:\n                mask_2d = mask.mean(axis=0)\n            elif mask.shape[-1] in [3, 4]:\n                mask_2d = mask.mean(axis=-1)\n            else:\n                # 兜底：直接在第 0 维求平均\n                mask_2d = mask.mean(axis=0)\n    else:\n        # 形状更怪时的兜底逻辑：reshape 成与图像同大小\n        mask_2d = mask.reshape(H, W)\n\n    # 二值化\n    m = (mask_2d > 0).astype(np.uint8)\n\n    # 如果没有任何伪造区域，直接返回空\n    if m.sum() == 0:\n        return []\n\n    # 寻找连通区域（伪造区域）\n    lbl = label(m)\n    props = regionprops(lbl)\n\n    patches = []\n\n    for region in props:\n        # 2D 情况下 bbox 一定是 4 个值\n        minr, minc, maxr, maxc = region.bbox\n        cy = (minr + maxr) // 2\n        cx = (minc + maxc) // 2\n\n        for size in patch_sizes:\n            half = size // 2\n            r1 = max(0, cy - half)\n            r2 = min(H, cy + half)\n            c1 = max(0, cx - half)\n            c2 = min(W, cx + half)\n\n            patch = img[r1:r2, c1:c2]\n\n            # patch 太小就跳过\n            if patch.shape[0] < size // 2 or patch.shape[1] < size // 2:\n                continue\n\n            patches.append((patch, 1))\n\n            if len(patches) >= max_patches:\n                return patches\n\n    return patches\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:39:58.918526Z","iopub.execute_input":"2025-12-09T16:39:58.918836Z","iopub.status.idle":"2025-12-09T16:39:58.976501Z","shell.execute_reply.started":"2025-12-09T16:39:58.918812Z","shell.execute_reply":"2025-12-09T16:39:58.975523Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_authentic_patches(img_path, patch_sizes=[128, 160], num_patches=3):\n    img = np.array(Image.open(img_path).convert('RGB'))\n    H, W = img.shape[:2]\n\n    patches = []\n\n    # 只保留比图像尺寸小的 patch size\n    valid_sizes = [s for s in patch_sizes if H > s and W > s]\n\n    if len(valid_sizes) == 0:\n        # 图像太小，直接跳过，不生成负样本\n        return patches\n\n    for _ in range(num_patches):\n        size = np.random.choice(valid_sizes)\n        half = size // 2\n\n        # 再次保护，防止极端尺寸\n        if H <= size or W <= size:\n            continue\n\n        # cy/cx 的取值范围要保证 patch 完全在图内\n        low_y, high_y = half, H - half\n        low_x, high_x = half, W - half\n\n        if low_y >= high_y or low_x >= high_x:\n            continue  # 安全兜底\n\n        cy = np.random.randint(low_y, high_y)\n        cx = np.random.randint(low_x, high_x)\n\n        patch = img[cy-half:cy+half, cx-half:cx+half]\n\n        # 再检查一次尺寸是否正确\n        if patch.shape[0] == size and patch.shape[1] == size:\n            patches.append((patch, 0))\n\n    return patches\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:40:03.098628Z","iopub.execute_input":"2025-12-09T16:40:03.099305Z","iopub.status.idle":"2025-12-09T16:40:03.105843Z","shell.execute_reply.started":"2025-12-09T16:40:03.099280Z","shell.execute_reply":"2025-12-09T16:40:03.105003Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom tqdm import tqdm\n\npatch_data = []\n\nfor idx, row in tqdm(df.iterrows(), total=len(df)):\n    img_path = row['image_path']\n    img_name = os.path.basename(img_path)\n    stem = os.path.splitext(img_name)[0]\n\n    if row['label'] == 1:\n        # forged → 找 mask\n        mask_path = f\"/kaggle/input/recodai-luc-scientific-image-forgery-detection/train_masks/{stem}.npy\"\n        if os.path.exists(mask_path):\n            patches = extract_forged_patches(img_path, mask_path)\n            for p, l in patches:\n                patch_data.append((p, l))\n    else:\n        # authentic → 随机采样\n        patches = extract_authentic_patches(img_path)\n        for p, l in patches:\n            patch_data.append((p, l))\n\nlen(patch_data)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T16:40:06.648047Z","iopub.execute_input":"2025-12-09T16:40:06.648418Z","iopub.status.idle":"2025-12-09T16:44:16.770729Z","shell.execute_reply.started":"2025-12-09T16:40:06.648394Z","shell.execute_reply":"2025-12-09T16:44:16.769906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.model_selection import train_test_split\n\n# patch_data: [(np.array(H,W,3), label), ...]\nlabels = [l for _, l in patch_data]\n\ntrain_patches, val_patches = train_test_split(\n    patch_data,\n    test_size=0.2,\n    stratify=labels,\n    random_state=42\n)\n\nlen(train_patches), len(val_patches)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T17:49:50.680304Z","iopub.execute_input":"2025-12-09T17:49:50.680557Z","iopub.status.idle":"2025-12-09T17:49:51.444842Z","shell.execute_reply.started":"2025-12-09T17:49:50.680537Z","shell.execute_reply":"2025-12-09T17:49:51.443791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom PIL import Image\n\nclass PatchDataset(Dataset):\n    def __init__(self, patch_list, train=True):\n        self.patch_list = patch_list\n        self.train = train\n\n        if train:\n            self.transform = transforms.Compose([\n                transforms.ToPILImage(),\n                transforms.RandomHorizontalFlip(),\n                transforms.RandomVerticalFlip(),\n                transforms.RandomRotation(15),\n                transforms.Resize((128, 128)),\n                transforms.ToTensor(),\n                transforms.Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225]\n                )\n            ])\n        else:\n            self.transform = transforms.Compose([\n                transforms.ToPILImage(),\n                transforms.Resize((128, 128)),\n                transforms.ToTensor(),\n                transforms.Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225]\n                )\n            ])\n\n    def __len__(self):\n        return len(self.patch_list)\n\n    def __getitem__(self, idx):\n        patch, label = self.patch_list[idx]   # patch: np.array(H,W,3), label: 0/1\n        patch = self.transform(patch)         # (3,128,128)\n        label = torch.tensor(label, dtype=torch.float32)\n        return patch, label\n\ntrain_patch_dataset = PatchDataset(train_patches, train=True)\nval_patch_dataset   = PatchDataset(val_patches, train=False)\n\ntrain_patch_loader = DataLoader(train_patch_dataset, batch_size=64, shuffle=True, num_workers=2)\nval_patch_loader   = DataLoader(val_patch_dataset,   batch_size=64, shuffle=False, num_workers=2)\n\nlen(train_patch_loader), len(val_patch_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T22:42:24.393866Z","iopub.execute_input":"2025-12-08T22:42:24.394456Z","iopub.status.idle":"2025-12-08T22:42:24.405848Z","shell.execute_reply.started":"2025-12-08T22:42:24.394431Z","shell.execute_reply":"2025-12-08T22:42:24.405225Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ResNet18 on patch","metadata":{}},{"cell_type":"code","source":"# 预训练 ResNet18，输入 3 通道，输出 1 维（伪造概率 logit）\npatch_model = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)\nin_features = patch_model.fc.in_features\npatch_model.fc = nn.Linear(in_features, 1)\npatch_model = patch_model.to(device)\n\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(patch_model.parameters(), lr=1e-4)\n\n\ndef train_one_epoch(model, loader, optimizer, criterion, device):\n    model.train()\n    running_loss = 0.0\n    for imgs, labels in tqdm(loader, leave=False):\n        imgs = imgs.to(device)\n        labels = labels.to(device).unsqueeze(1)\n\n        optimizer.zero_grad()\n        outputs = model(imgs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item() * imgs.size(0)\n    return running_loss / len(loader.dataset)\n\n\ndef eval_one_epoch(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs = imgs.to(device)\n            labels = labels.to(device).unsqueeze(1)\n\n            outputs = model(imgs)\n            loss = criterion(outputs, labels)\n            running_loss += loss.item() * imgs.size(0)\n\n            probs = torch.sigmoid(outputs)\n            preds = (probs > 0.5).float()\n            correct += (preds == labels).sum().item()\n            total += labels.size(0)\n\n    return running_loss / len(loader.dataset), correct / total\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T22:42:30.774310Z","iopub.execute_input":"2025-12-08T22:42:30.775060Z","iopub.status.idle":"2025-12-08T22:42:30.983545Z","shell.execute_reply.started":"2025-12-08T22:42:30.775028Z","shell.execute_reply":"2025-12-08T22:42:30.982764Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"history = {\n    \"train_loss\": [],\n    \"val_loss\": [],\n    \"val_acc\": []\n}\n\nnum_epochs = 10  # 顺便可以跑到 10 轮，看过拟合趋势更明显\nfor epoch in range(1, num_epochs + 1):\n    train_loss = train_one_epoch(patch_model, train_patch_loader, optimizer, criterion, device)\n    val_loss, val_acc = eval_one_epoch(patch_model, val_patch_loader, criterion, device)\n\n    history[\"train_loss\"].append(train_loss)\n    history[\"val_loss\"].append(val_loss)\n    history[\"val_acc\"].append(val_acc)\n\n    print(f\"Epoch {epoch}: train_loss={train_loss:.4f}, val_loss={val_loss:.4f}, val_acc={val_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T05:24:40.329325Z","iopub.execute_input":"2025-12-08T05:24:40.329995Z","iopub.status.idle":"2025-12-08T05:27:58.935791Z","shell.execute_reply.started":"2025-12-08T05:24:40.329968Z","shell.execute_reply":"2025-12-08T05:27:58.934813Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\n\nepochs = np.arange(1, len(history[\"train_loss\"])+1)\n\nplt.figure(figsize=(10,4))\nplt.subplot(1,2,1)\nplt.plot(epochs, history[\"train_loss\"], label='train_loss')\nplt.plot(epochs, history[\"val_loss\"], label='val_loss')\nplt.xlabel('epoch'); plt.ylabel('loss')\nplt.legend(); plt.title('Loss')\n\nplt.subplot(1,2,2)\nplt.plot(epochs, history[\"val_acc\"], label='val_acc')\nplt.xlabel('epoch'); plt.ylabel('accuracy')\nplt.legend(); plt.title('Val Accuracy')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix, classification_report\n\ndef eval_with_details(model, loader, device):\n    model.eval()\n    all_labels = []\n    all_preds = []\n    with torch.no_grad():\n        for imgs, labels in loader:\n            imgs = imgs.to(device)\n            labels = labels.to(device).unsqueeze(1)\n\n            outputs = model(imgs)\n            probs = torch.sigmoid(outputs)\n            preds = (probs > 0.5).float()\n\n            all_labels.append(labels.cpu())\n            all_preds.append(preds.cpu())\n\n    all_labels = torch.cat(all_labels).numpy().astype(int).ravel()\n    all_preds  = torch.cat(all_preds).numpy().astype(int).ravel()\n\n    print(\"Confusion matrix:\")\n    print(confusion_matrix(all_labels, all_preds))\n    print(\"\\nClassification report:\")\n    print(classification_report(all_labels, all_preds, digits=4))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T05:28:50.900546Z","iopub.execute_input":"2025-12-08T05:28:50.900828Z","iopub.status.idle":"2025-12-08T05:28:50.907007Z","shell.execute_reply.started":"2025-12-08T05:28:50.900807Z","shell.execute_reply":"2025-12-08T05:28:50.906214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"eval_with_details(patch_model, val_patch_loader, device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T05:29:01.205649Z","iopub.execute_input":"2025-12-08T05:29:01.206162Z","iopub.status.idle":"2025-12-08T05:29:04.338160Z","shell.execute_reply.started":"2025-12-08T05:29:01.206138Z","shell.execute_reply":"2025-12-08T05:29:04.337130Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model：EfficientNet-B0","metadata":{}},{"cell_type":"code","source":"import timm\nimport torch.nn as nn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T22:43:44.182350Z","iopub.execute_input":"2025-12-08T22:43:44.183000Z","iopub.status.idle":"2025-12-08T22:43:47.901149Z","shell.execute_reply.started":"2025-12-08T22:43:44.182975Z","shell.execute_reply":"2025-12-08T22:43:47.900556Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_name = \"efficientnet_b0\"\neffnet = timm.create_model(model_name, pretrained=True, num_classes=1)\neffnet = effnet.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T22:43:54.296792Z","iopub.execute_input":"2025-12-08T22:43:54.297450Z","iopub.status.idle":"2025-12-08T22:43:55.630188Z","shell.execute_reply.started":"2025-12-08T22:43:54.297422Z","shell.execute_reply":"2025-12-08T22:43:55.629390Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model_name = \"efficientnet_b0\"\neffnet = timm.create_model(model_name, pretrained=True, num_classes=1)\neffnet = effnet.to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T22:44:03.681051Z","iopub.execute_input":"2025-12-08T22:44:03.681350Z","iopub.status.idle":"2025-12-08T22:44:03.862990Z","shell.execute_reply.started":"2025-12-08T22:44:03.681329Z","shell.execute_reply":"2025-12-08T22:44:03.862385Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"criterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(effnet.parameters(), lr=1e-4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T22:44:10.701101Z","iopub.execute_input":"2025-12-08T22:44:10.701692Z","iopub.status.idle":"2025-12-08T22:44:10.706229Z","shell.execute_reply.started":"2025-12-08T22:44:10.701667Z","shell.execute_reply":"2025-12-08T22:44:10.705369Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 10\nfor epoch in range(1, num_epochs + 1):\n    train_loss = train_one_epoch(effnet, train_patch_loader, optimizer, criterion, device)\n    val_loss, val_acc = eval_one_epoch(effnet, val_patch_loader, criterion, device)\n\n    print(f\"Epoch {epoch}: train_loss={train_loss:.4f}, val_loss={val_loss:.4f}, val_acc={val_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T22:44:18.274053Z","iopub.execute_input":"2025-12-08T22:44:18.274325Z","iopub.status.idle":"2025-12-08T22:49:28.209295Z","shell.execute_reply.started":"2025-12-08T22:44:18.274305Z","shell.execute_reply":"2025-12-08T22:49:28.208214Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## ResNet18调优","metadata":{}},{"cell_type":"code","source":"train_transforms = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(degrees=15),\n    transforms.ColorJitter(brightness=0.2, contrast=0.2),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])\n\nval_transforms = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor(),\n    transforms.Normalize(\n        mean=[0.485, 0.456, 0.406],\n        std=[0.229, 0.224, 0.225]\n    )\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T23:53:01.776075Z","iopub.execute_input":"2025-12-08T23:53:01.776372Z","iopub.status.idle":"2025-12-08T23:53:01.782934Z","shell.execute_reply.started":"2025-12-08T23:53:01.776348Z","shell.execute_reply":"2025-12-08T23:53:01.782280Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torchvision.models as models\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\n\n# 创建 ResNet18（patch 模型用这个）\nresnet = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)\nnum_features = resnet.fc.in_features\nresnet.fc = nn.Linear(num_features, 1)   # 二分类，用 BCEWithLogitsLoss\nresnet = resnet.to(device)\n\ncriterion = nn.BCEWithLogitsLoss()\n\n# 关键：加入 weight_decay\noptimizer = optim.Adam(resnet.parameters(), lr=1e-4, weight_decay=1e-4)\n\n# 学习率调度器：验证集 loss 连续几轮不下降就减小 lr\nscheduler = ReduceLROnPlateau(optimizer, mode='min',\n                              factor=0.5, patience=2,\n                              verbose=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T23:53:09.438712Z","iopub.execute_input":"2025-12-08T23:53:09.439243Z","iopub.status.idle":"2025-12-08T23:53:09.677142Z","shell.execute_reply.started":"2025-12-08T23:53:09.439202Z","shell.execute_reply":"2025-12-08T23:53:09.676232Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"num_epochs = 15\n\nfor epoch in range(1, num_epochs + 1):\n    train_loss = train_one_epoch(resnet, train_patch_loader,\n                                 optimizer, criterion, device)\n    val_loss, val_acc = eval_one_epoch(resnet, val_patch_loader,\n                                       criterion, device)\n\n    # 调度器用验证集 loss 决定是否降学习率\n    scheduler.step(val_loss)\n\n    print(f\"Epoch {epoch}: train_loss={train_loss:.4f}, \"\n          f\"val_loss={val_loss:.4f}, val_acc={val_acc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T23:53:26.913091Z","iopub.execute_input":"2025-12-08T23:53:26.913354Z","iopub.status.idle":"2025-12-08T23:58:27.413229Z","shell.execute_reply.started":"2025-12-08T23:53:26.913335Z","shell.execute_reply":"2025-12-08T23:58:27.412400Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"search_lr = [1e-4, 3e-4]\nsearch_wd = [1e-4, 5e-4]\n\nresults = []\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T00:15:27.263662Z","iopub.execute_input":"2025-12-09T00:15:27.264507Z","iopub.status.idle":"2025-12-09T00:15:27.268535Z","shell.execute_reply.started":"2025-12-09T00:15:27.264473Z","shell.execute_reply":"2025-12-09T00:15:27.267786Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for lr in search_lr:\n    for wd in search_wd:\n        \n        print(f\"\\n===== Training with lr={lr}, weight_decay={wd} =====\")\n\n        # 重建模型\n        resnet = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)\n        num_features = resnet.fc.in_features\n        resnet.fc = nn.Linear(num_features, 1)\n        resnet = resnet.to(device)\n\n        criterion = nn.BCEWithLogitsLoss()\n        optimizer = optim.Adam(resnet.parameters(), lr=lr, weight_decay=wd)\n\n        # 跑5个epoch的快速验证\n        for epoch in range(5):\n            train_loss = train_one_epoch(resnet, train_patch_loader, optimizer, criterion, device)\n            val_loss, val_acc = eval_one_epoch(resnet, val_patch_loader, criterion, device)\n\n        print(f\"Result: lr={lr}, wd={wd}, val_acc={val_acc:.4f}\")\n\n        results.append((lr, wd, val_acc))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-09T00:15:33.085848Z","iopub.execute_input":"2025-12-09T00:15:33.086357Z","iopub.status.idle":"2025-12-09T00:22:18.953583Z","shell.execute_reply.started":"2025-12-09T00:15:33.086333Z","shell.execute_reply":"2025-12-09T00:22:18.952579Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Final upload for Kaggle","metadata":{}},{"cell_type":"code","source":"def extract_test_patches(img_path, patch_sizes=[128,160], num_patches=8):\n    img = np.array(Image.open(img_path).convert(\"RGB\"))\n    H, W = img.shape[:2]\n\n    patches = []\n\n    valid_sizes = [s for s in patch_sizes if H > s and W > s]\n    if len(valid_sizes) == 0:\n        return patches\n\n    for _ in range(num_patches):\n        size = np.random.choice(valid_sizes)\n        half = size // 2\n        low_y, high_y = half, H - half\n        low_x, high_x = half, W - half\n        if low_y >= high_y or low_x >= high_x:\n            continue\n\n        cy = np.random.randint(low_y, high_y)\n        cx = np.random.randint(low_x, high_x)\n\n        patch = img[cy-half:cy+half, cx-half:cx+half]\n        if patch.shape[0] == size and patch.shape[1] == size:\n            patches.append(patch)\n\n    return patches\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T05:09:23.124203Z","iopub.execute_input":"2025-12-08T05:09:23.124802Z","iopub.status.idle":"2025-12-08T05:09:23.131938Z","shell.execute_reply.started":"2025-12-08T05:09:23.124774Z","shell.execute_reply":"2025-12-08T05:09:23.131131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_dir = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection/test_images\"\n\npatch_model.eval()\nresults = []\n\nfor fname in os.listdir(test_dir):\n    img_path = os.path.join(test_dir, fname)\n\n    patches = extract_test_patches(img_path, num_patches=10)\n\n    scores = []\n    for p in patches:\n        x = transforms.ToTensor()(Image.fromarray(p).resize((128,128)))\n        x = transforms.Normalize(\n            [0.485,0.456,0.406],\n            [0.229,0.224,0.225]\n        )(x).unsqueeze(0).to(device)\n\n        with torch.no_grad():\n            logit = patch_model(x)\n            prob = torch.sigmoid(logit).item()\n            scores.append(prob)\n\n    if len(scores) == 0:\n        final_score = 0.01\n    else:\n        final_score = max(scores)\n\n    results.append((fname, final_score))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T05:09:32.689239Z","iopub.execute_input":"2025-12-08T05:09:32.689903Z","iopub.status.idle":"2025-12-08T05:09:32.819796Z","shell.execute_reply.started":"2025-12-08T05:09:32.689884Z","shell.execute_reply":"2025-12-08T05:09:32.819298Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\n\ndf_submit = pd.DataFrame(results, columns=[\"filename\", \"label\"])\ndf_submit.to_csv(\"submission.csv\", index=False)\n\ndf_submit.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T05:09:43.424218Z","iopub.execute_input":"2025-12-08T05:09:43.424809Z","iopub.status.idle":"2025-12-08T05:09:43.443708Z","shell.execute_reply.started":"2025-12-08T05:09:43.424783Z","shell.execute_reply":"2025-12-08T05:09:43.443140Z"}},"outputs":[],"execution_count":null}]}