{"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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":13992641,"sourceType":"datasetVersion","datasetId":8917690},{"sourceId":13992840,"sourceType":"datasetVersion","datasetId":8917793},{"sourceId":672552,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":509578,"modelId":524244},{"sourceId":672941,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":509914,"modelId":524579}],"dockerImageVersionId":31193,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install --no-deps /kaggle/input/segmentation-models-pytorch/segmentation_models_pytorch-0.5.0-py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T09:51:39.825776Z","iopub.execute_input":"2025-12-08T09:51:39.826016Z","iopub.status.idle":"2025-12-08T09:51:42.311169Z","shell.execute_reply.started":"2025-12-08T09:51:39.825997Z","shell.execute_reply":"2025-12-08T09:51:42.310265Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install --no-deps /kaggle/input/albumentations/albumentations-2.0.8-py3-none-any.whl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T09:51:42.313167Z","iopub.execute_input":"2025-12-08T09:51:42.313427Z","iopub.status.idle":"2025-12-08T09:51:43.813127Z","shell.execute_reply.started":"2025-12-08T09:51:42.313402Z","shell.execute_reply":"2025-12-08T09:51:43.812290Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport glob\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nfrom tqdm import tqdm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import train_test_split\nimport json","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T09:51:43.814225Z","iopub.execute_input":"2025-12-08T09:51:43.814537Z","iopub.status.idle":"2025-12-08T09:52:27.151238Z","shell.execute_reply.started":"2025-12-08T09:51:43.814502Z","shell.execute_reply":"2025-12-08T09:52:27.150618Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n    SEED = 42\n    IMAGE_SIZE = (512, 512)\n    BATCH_SIZE = 8\n    EPOCHS = 10\n    LEARNING_RATE = 1e-4\n    ENCODER = 'efficientnet-b0'\n    WEIGHTS = None\n    DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    LOCAL_WEIGHTS_PATH = \"/kaggle/input/efficientnet/pytorch/default/1/efficientnet-b0-355c32eb.pth\"\n    \n    # === 原有训练集路径 ===\n    TRAIN_IMG_DIR = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection/train_images\"\n    TRAIN_MASK_DIR = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection/train_masks\"\n    \n    # === ✅ 新增：补充数据集路径 ===\n    # 注意：请确保路径名称与你的 Kaggle Input 目录完全一致\n    SUPP_IMG_DIR = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection/supplemental_images\"\n    SUPP_MASK_DIR = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection/supplemental_masks\"","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T09:52:27.151985Z","iopub.execute_input":"2025-12-08T09:52:27.152466Z","iopub.status.idle":"2025-12-08T09:52:27.234223Z","shell.execute_reply.started":"2025-12-08T09:52:27.152446Z","shell.execute_reply":"2025-12-08T09:52:27.233304Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ForensicsDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\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        image_path = row['image_path']\n        mask_path = row['mask_path']\n        \n        # 1. 读取图像 (Image)\n        image = cv2.imread(image_path)\n        if image is None:\n             raise FileNotFoundError(f\"无法读取图片: {image_path}\")\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        target_h, target_w = image.shape[:2]\n        \n        # 2. 读取掩码 (Mask)\n        mask = None\n        \n        if os.path.exists(mask_path):\n            try:\n                # 加载 npy\n                mask_data = np.load(mask_path)\n                \n                # --- 修复核心：强制降维 ---\n                # 无论读进来是什么形状，我们都要把它弄成 2D 的 (H, W)\n                if mask_data.ndim == 3:\n                    # 如果是 (H, W, C) 或者 (C, H, W)\n                    # 我们取所有通道的最大值 (只要任意一个通道是1，就是篡改)\n                    # 也可以用 sum 或 mean，但 max 对于二分类最保险\n                    mask_data = np.max(mask_data, axis=-1) if mask_data.shape[-1] < mask_data.shape[0] else np.max(mask_data, axis=0)\n                \n                # 现在 mask_data 应该是 2D 的了，检查并 resize\n                if mask_data.ndim == 2:\n                    if mask_data.shape[0] != target_h or mask_data.shape[1] != target_w:\n                        mask_data = cv2.resize(mask_data.astype(np.float32), (target_w, target_h), interpolation=cv2.INTER_NEAREST)\n                else:\n                    # 如果维度还是不对 (比如 1D 或 4D)，丢弃\n                    mask_data = None\n                \n                mask = mask_data\n                \n            except Exception as e:\n                # 打印错误但不中断训练，返回全黑\n                # print(f\"Error processing mask {mask_path}: {e}\")\n                mask = None\n\n        # 3. 处理空掩码 (Authentic 或 读取失败)\n        if mask is None:\n            mask = np.zeros((target_h, target_w), dtype=np.float32)\n        else:\n            # 确保是 float32 且二值化\n            mask = mask.astype(np.float32)\n            mask = np.where(mask > 0, 1.0, 0.0)\n            \n        # 4. 数据增强\n        if self.transform:\n            try:\n                augmented = self.transform(image=image, mask=mask)\n                image = augmented['image']\n                mask = augmented['mask']\n            except Exception as e:\n                # 如果增强失败，回退到不增强\n                print(f\"Augmentation failed for {image_path}, using raw image.\")\n                image = torch.from_numpy(image.transpose(2, 0, 1)).float()\n                mask = torch.from_numpy(mask).unsqueeze(0).float()\n            \n        # 5. 最终形状检查与转换\n        # 确保 mask 是 Tensor\n        if not isinstance(mask, torch.Tensor):\n            mask = torch.from_numpy(mask)\n            \n        # 确保 Image 是 Tensor\n        if not isinstance(image, torch.Tensor):\n            image = torch.from_numpy(image.transpose(2, 0, 1)).float()\n\n        # 统一 Mask 维度为 (1, H, W)\n        if mask.ndim == 2:\n            mask = mask.unsqueeze(0)\n            \n        return image, mask, row['id']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T09:52:27.236075Z","iopub.execute_input":"2025-12-08T09:52:27.236308Z","iopub.status.idle":"2025-12-08T09:52:27.253970Z","shell.execute_reply.started":"2025-12-08T09:52:27.236289Z","shell.execute_reply":"2025-12-08T09:52:27.253380Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_transforms(data):\n    if data == 'train':\n        return A.Compose([\n            A.Resize(Config.IMAGE_SIZE[0], Config.IMAGE_SIZE[1]),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.1, rotate_limit=45, p=0.5),\n            A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n            ToTensorV2(),\n        ])\n    elif data == 'valid':\n        return A.Compose([\n            A.Resize(Config.IMAGE_SIZE[0], Config.IMAGE_SIZE[1]),\n            A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n            ToTensorV2(),\n        ])\n\n# ==========================================\n# 4. 准备数据 DataFrame\n# ==========================================\ndef prepare_data():\n    # 1. 定义要加载的数据源列表： [(图片文件夹, 掩码文件夹), ...]\n    data_sources = [\n        (Config.TRAIN_IMG_DIR, Config.TRAIN_MASK_DIR),      # 原始训练集\n        (Config.SUPP_IMG_DIR, Config.SUPP_MASK_DIR)         # 补充数据集\n    ]\n    \n    data = []\n    \n    print(f\"{'='*10} 开始准备数据 {'='*10}\")\n    \n    # 2. 遍历每个数据源\n    for img_dir, mask_dir in data_sources:\n        if not os.path.exists(img_dir):\n            print(f\"⚠️ 警告: 路径不存在，跳过: {img_dir}\")\n            continue\n            \n        # 搜索该目录下的所有 png 图片\n        search_path = os.path.join(img_dir, \"**/*.png\")\n        all_files = glob.glob(search_path, recursive=True)\n        \n        print(f\"📂 正在扫描目录: {os.path.basename(img_dir)} | 找到图片: {len(all_files)} 张\")\n        \n        for img_path in all_files:\n            if not os.path.isfile(img_path):\n                continue\n                \n            file_name = os.path.basename(img_path)\n            img_id = os.path.splitext(file_name)[0]\n            \n            # 假设补充数据的 Mask 命名规则也是 id.npy\n            # 如果补充数据的 Mask 是 .png 格式，请在这里修改后缀为 .png\n            mask_filename = img_id + \".npy\" \n            mask_path = os.path.join(mask_dir, mask_filename)\n            \n            data.append({\n                'id': img_id,\n                'image_path': img_path,\n                'mask_path': mask_path\n            })\n    \n    # 3. 构建 DataFrame\n    df = pd.DataFrame(data)\n    \n    if len(df) == 0:\n        raise RuntimeError(\"❌ 未找到任何图片，请检查路径配置！\")\n        \n    print(f\"✅ 总计加载数据: {len(df)} 条\")\n\n    # 4. 划分训练集和验证集\n    # 这里要小心：如果补充数据量很大且分布不同，混在一起切分可能导致验证集“过易”或“过难”\n    # 但作为 Baseline，直接随机切分通常是可以接受的\n    train_df, valid_df = train_test_split(df, test_size=0.2, random_state=Config.SEED)\n    \n    print(f\"📊 训练集数量: {len(train_df)} | 验证集数量: {len(valid_df)}\")\n    return train_df.reset_index(drop=True), valid_df.reset_index(drop=True)\n\n# ==========================================\n# 5. 训练辅助函数\n# ==========================================\ndef train_fn(loader, model, optimizer, loss_fn, scaler):\n    model.train()\n    train_loss = 0\n    loop = tqdm(loader, desc=\"Training\")\n    \n    for images, masks, _ in loop:\n        images = images.to(Config.DEVICE)\n        masks = masks.to(Config.DEVICE)\n        \n        # 混合精度训练 (AMP) - 节省显存，加快速度\n        with autocast():\n            outputs = model(images)\n            loss = loss_fn(outputs, masks)\n            \n        optimizer.zero_grad()\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        train_loss += loss.item()\n        loop.set_postfix(loss=loss.item())\n        \n    return train_loss / len(loader)\n\ndef valid_fn(loader, model, loss_fn):\n    model.eval()\n    valid_loss = 0\n    loop = tqdm(loader, desc=\"Validating\")\n    \n    pred_rles = []\n    gt_rles = []\n    ids = []\n    shapes = []\n    \n    with torch.no_grad():\n        for images, masks, img_ids in loop:\n            images = images.to(Config.DEVICE)\n            masks = masks.to(Config.DEVICE)\n            \n            outputs = model(images)\n            loss = loss_fn(outputs, masks)\n            valid_loss += loss.item()\n            \n            # 1. Sigmoid & 阈值化\n            preds = torch.sigmoid(outputs)\n            preds = (preds > 0.5).float()\n            \n            preds = preds.cpu().numpy()\n            masks = masks.cpu().numpy()\n            \n            # 2. 逐张处理，适配提交格式\n            for i in range(images.shape[0]):\n                pred_img = preds[i][0]\n                gt_img = masks[i][0]\n                \n                # --- 核心修改：处理 \"authentic\" 标签 ---\n                \n                # 预测处理：如果预测全是背景(0)，或者是极小噪点(比如小于10个像素)，标记为 authentic\n                if np.sum(pred_img) < 10: \n                    p_rle = \"authentic\"\n                else:\n                    p_rle = rle_encode([pred_img])\n                    # 双重保险：如果编码结果是空列表字符串，也改成 authentic\n                    if p_rle == \"[]\" or p_rle == \"\":\n                        p_rle = \"authentic\"\n\n                # 真值处理：如果 Ground Truth 是全黑，标记为 authentic\n                if np.sum(gt_img) == 0:\n                    g_rle = \"authentic\"\n                else:\n                    g_rle = rle_encode([gt_img])\n                # -------------------------------------\n                \n                pred_rles.append(p_rle)\n                gt_rles.append(g_rle)\n                ids.append(img_ids[i])\n                shapes.append(json.dumps([pred_img.shape[0], pred_img.shape[1]]))\n                \n    # --- 构造符合要求的 DataFrame ---\n    # 截图要求列名为 case_id, annotation\n    solution_df = pd.DataFrame({\n        'case_id': ids,\n        'annotation': gt_rles,\n        'shape': shapes\n    })\n    \n    submission_df = pd.DataFrame({\n        'case_id': ids,\n        'annotation': pred_rles\n    })\n\n    # 计算分数\n    try:\n        # 注意：这里我们传递 'case_id' 作为 ID 列名\n        f1_score = score(solution_df, submission_df, 'case_id')\n    except Exception as e:\n        print(f\"Metrics calculation failed: {e}\")\n        f1_score = 0.0\n\n    return valid_loss / len(loader), f1_score","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T09:52:27.254771Z","iopub.execute_input":"2025-12-08T09:52:27.255023Z","iopub.status.idle":"2025-12-08T09:52:27.277032Z","shell.execute_reply.started":"2025-12-08T09:52:27.255001Z","shell.execute_reply":"2025-12-08T09:52:27.276475Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import json\nimport numba\nimport numpy as np\nimport pandas as pd\nimport scipy.optimize\nfrom numba import types\nimport numpy.typing as npt\n\n# --- 1. RLE 编码相关 (必须有 numba) ---\n\n@numba.jit(nopython=True)\ndef _rle_encode_jit(x: npt.NDArray, fg_val: int = 1) -> list[int]:\n    \"\"\"Numba-jitted RLE encoder.\"\"\"\n    dots = np.where(x.T.flatten() == fg_val)[0]\n    run_lengths = []\n    prev = -2\n    for b in dots:\n        if b > prev + 1:\n            run_lengths.extend((b + 1, 0))\n        run_lengths[-1] += 1\n        prev = b\n    return run_lengths\n\ndef rle_encode(masks, fg_val: int = 1) -> str:\n    \"\"\"\n    Args:\n        masks: list of numpy array of shape (height, width), 1 - mask, 0 - background\n    Returns: run length encodings as a string\n    \"\"\"\n    return ';'.join([json.dumps(_rle_encode_jit(x, fg_val)) for x in masks])\n\n@numba.njit\ndef _rle_decode_jit(mask_rle: npt.NDArray, height: int, width: int) -> npt.NDArray:\n    if len(mask_rle) % 2 != 0:\n        raise ValueError('One or more rows has an odd number of values.')\n    starts, lengths = mask_rle[0::2], mask_rle[1::2]\n    starts -= 1\n    ends = starts + lengths\n    for i in range(len(starts) - 1):\n        if ends[i] > starts[i + 1]:\n            raise ValueError('Pixels must not be overlapping.')\n    img = np.zeros(height * width, dtype=np.bool_)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img\n\nclass ParticipantVisibleError(Exception):\n    pass\n\ndef rle_decode(mask_rle: str, shape: tuple[int, int]) -> npt.NDArray:\n    mask_rle = json.loads(mask_rle)\n    mask_rle = np.asarray(mask_rle, dtype=np.int32)\n    try:\n        return _rle_decode_jit(mask_rle, shape[0], shape[1]).reshape(shape, order='F')\n    except ValueError as e:\n        raise ParticipantVisibleError(str(e))\n\n# --- 2. 分数计算相关 (Score Metrics) ---\n\ndef calculate_f1_score(pred_mask: npt.NDArray, gt_mask: npt.NDArray):\n    pred_flat = pred_mask.flatten()\n    gt_flat = gt_mask.flatten()\n    tp = np.sum((pred_flat == 1) & (gt_flat == 1))\n    fp = np.sum((pred_flat == 1) & (gt_flat == 0))\n    fn = np.sum((pred_flat == 0) & (gt_flat == 1))\n    precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n    recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n    if (precision + recall) > 0:\n        return 2 * (precision * recall) / (precision + recall)\n    else:\n        return 0\n\ndef calculate_f1_matrix(pred_masks, gt_masks):\n    num_instances_pred = len(pred_masks)\n    num_instances_gt = len(gt_masks)\n    f1_matrix = np.zeros((num_instances_pred, num_instances_gt))\n    for i in range(num_instances_pred):\n        for j in range(num_instances_gt):\n            pred_flat = pred_masks[i].flatten()\n            gt_flat = gt_masks[j].flatten()\n            f1_matrix[i, j] = calculate_f1_score(pred_mask=pred_flat, gt_mask=gt_flat)\n    if f1_matrix.shape[0] < len(gt_masks):\n        f1_matrix = np.vstack((f1_matrix, np.zeros((len(gt_masks) - len(f1_matrix), num_instances_gt))))\n    return f1_matrix\n\ndef oF1_score(pred_masks, gt_masks):\n    f1_matrix = calculate_f1_matrix(pred_masks, gt_masks)\n    row_ind, col_ind = scipy.optimize.linear_sum_assignment(-f1_matrix)\n    excess_predictions_penalty = len(gt_masks) / max(len(pred_masks), len(gt_masks))\n    return np.mean(f1_matrix[row_ind, col_ind]) * excess_predictions_penalty\n\ndef evaluate_single_image(label_rles: str, prediction_rles: str, shape_str: str) -> float:\n    shape = json.loads(shape_str)\n    label_rles = [rle_decode(x, shape=shape) for x in label_rles.split(';')]\n    prediction_rles = [rle_decode(x, shape=shape) for x in prediction_rles.split(';')]\n    return oF1_score(prediction_rles, label_rles)\n\ndef score(solution: pd.DataFrame, submission: pd.DataFrame, row_id_column_name: str) -> float:\n    df = solution.copy()\n    df = df.rename(columns={'annotation': 'label'})\n    df['prediction'] = submission['annotation']\n    # Check for correct 'authentic' label\n    authentic_indices = (df['label'] == 'authentic') | (df['prediction'] == 'authentic')\n    df['image_score'] = ((df['label'] == df['prediction']) & authentic_indices).astype(float)\n    \n    # 只对非 authentic 的行进行复杂计算\n    mask_indices = ~authentic_indices\n    if mask_indices.any():\n        df.loc[mask_indices, 'image_score'] = df.loc[mask_indices].apply(\n            lambda row: evaluate_single_image(row['label'], row['prediction'], row['shape']), axis=1\n        )\n    return float(np.mean(df['image_score']))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T09:52:27.277755Z","iopub.execute_input":"2025-12-08T09:52:27.278009Z","iopub.status.idle":"2025-12-08T09:52:28.247018Z","shell.execute_reply.started":"2025-12-08T09:52:27.277988Z","shell.execute_reply":"2025-12-08T09:52:28.246252Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, img_paths, transform=None):\n        self.img_paths = img_paths\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.img_paths)\n    \n    def __getitem__(self, idx):\n        img_path = self.img_paths[idx]\n        file_name = os.path.basename(img_path)\n        img_id = os.path.splitext(file_name)[0] # 获取 case_id，比如 \"image_123\" 或 \"123\"\n        \n        # 读取图片\n        image = cv2.imread(img_path)\n        if image is None:\n            raise FileNotFoundError(f\"无法读取测试图: {img_path}\")\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        # 记录原始尺寸 (H, W)，用于最后把 Mask 还原回去\n        orig_h, orig_w = image.shape[:2]\n        \n        if self.transform:\n            # Albumentations\n            augmented = self.transform(image=image)\n            image = augmented['image']\n            \n        # 转 Tensor\n        if not isinstance(image, torch.Tensor):\n            image = torch.from_numpy(image.transpose(2, 0, 1)).float()\n            \n        return image, img_id, orig_h, orig_w\n\ndef generate_submission():\n    print(\"\\nStarting Inference on Test Set...\")\n    \n    # 1. 准备测试集路径\n    # 根据你的截图，test_images 和 train_images 在同一级\n    # 假设路径结构类似，你需要根据实际情况调整 TEST_DIR\n    # 通常 Kaggle 路径是: /kaggle/input/[competition-name]/test_images/\n    # 这里我们尝试自动推断或者手动指定\n    TEST_DIR = Config.TRAIN_IMG_DIR.replace(\"train_images\", \"test_images\")\n    \n    test_img_paths = glob.glob(os.path.join(TEST_DIR, \"*\"))\n    \n    # 过滤非图片文件\n    valid_exts = {'.jpg', '.jpeg', '.png', '.tif', '.tiff'}\n    test_img_paths = [p for p in test_img_paths if os.path.splitext(p)[1].lower() in valid_exts]\n    \n    if len(test_img_paths) == 0:\n        print(\"Warning: No test images found! Generating empty submission for debugging.\")\n        # 创建一个 dummy submission 以防报错\n        pd.DataFrame({'case_id': [], 'annotation': []}).to_csv('submission.csv', index=False)\n        return\n\n    # 2. Dataset & DataLoader\n    test_dataset = TestDataset(test_img_paths, transform=get_transforms('valid'))\n    test_loader = DataLoader(test_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=0) # 保持0以防崩\n    \n    # 3. 加载最优模型\n    model = smp.Unet(\n        encoder_name=Config.ENCODER, \n        encoder_weights=None, # 推理时不需要下载权重\n        in_channels=3, \n        classes=1\n    )\n    # 加载刚训练好的权重\n    if os.path.exists('best_model.pth'):\n        model.load_state_dict(torch.load('best_model.pth', map_location=Config.DEVICE))\n        print(\"Loaded best_model.pth\")\n    else:\n        print(\"Warning: best_model.pth not found, using current model weights.\")\n        \n    model.to(Config.DEVICE)\n    model.eval()\n    \n    submission_data = []\n    \n    # 4. 推理循环\n    with torch.no_grad():\n        for images, ids, orig_hs, orig_ws in tqdm(test_loader, desc=\"Inference\"):\n            images = images.to(Config.DEVICE)\n            outputs = model(images)\n            \n            preds = torch.sigmoid(outputs)\n            preds = (preds > 0.5).float()\n            preds = preds.cpu().numpy()\n            \n            # 逐张处理\n            for i in range(len(ids)):\n                pred_mask = preds[i][0] # (512, 512)\n                case_id = ids[i]\n                h, w = orig_hs[i].item(), orig_ws[i].item()\n                \n                # --- 关键步骤：还原到原始尺寸 ---\n                # 因为 RLE 编码必须对应原图大小\n                if pred_mask.shape[0] != h or pred_mask.shape[1] != w:\n                    pred_mask = cv2.resize(pred_mask, (w, h), interpolation=cv2.INTER_NEAREST)\n                \n                # --- 格式处理 ---\n                # 如果像素点太少，或者全黑，视为 authentic\n                if np.sum(pred_mask) < 10: \n                    rle_str = \"authentic\"\n                else:\n                    rle_str = rle_encode([pred_mask]) # 注意这里要传 list\n                    if rle_str == \"[]\" or rle_str == \"\":\n                        rle_str = \"authentic\"\n                \n                submission_data.append({\n                    'case_id': case_id, # 根据截图要求列名是 case_id\n                    'annotation': rle_str\n                })\n                \n    # 5. 保存 CSV\n    sub_df = pd.DataFrame(submission_data)\n    \n    # 再次根据截图确认 ID 格式\n    # 如果截图里 ID 是数字 (如 1, 2)，可能需要转 int，如果是文件名 (image_123) 则保持 str\n    # 截图显示 \"1, authentic\"，说明 ID 可能是纯数字\n    # 尝试把 ID 转为数字排序，这样比较整齐（可选）\n    try:\n        sub_df['case_id'] = pd.to_numeric(sub_df['case_id'])\n        sub_df = sub_df.sort_values('case_id')\n    except:\n        pass # 如果转不了数字（含字符），就保持原样\n        \n    sub_df.to_csv('submission.csv', index=False)\n    print(f\"Submission saved to submission.csv with {len(sub_df)} rows.\")\n    print(sub_df.head())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T09:52:28.247874Z","iopub.execute_input":"2025-12-08T09:52:28.248186Z","iopub.status.idle":"2025-12-08T09:52:28.262251Z","shell.execute_reply.started":"2025-12-08T09:52:28.248156Z","shell.execute_reply":"2025-12-08T09:52:28.261625Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 6. 主程序\n# ==========================================\nif __name__ == '__main__':\n    # 1. 准备数据\n    train_df, valid_df = prepare_data()\n    print(f\"Training on {len(train_df)} images, Validating on {len(valid_df)} images.\")\n    \n    train_dataset = ForensicsDataset(train_df, transform=get_transforms('train'))\n    valid_dataset = ForensicsDataset(valid_df, transform=get_transforms('valid'))\n    \n    train_loader = DataLoader(train_dataset, batch_size=Config.BATCH_SIZE, shuffle=True, num_workers=0, pin_memory=True)\n    valid_loader = DataLoader(valid_dataset, batch_size=Config.BATCH_SIZE, shuffle=False, num_workers=0, pin_memory=True)\n    \n    # 2. 定义模型\n    # 使用 Unet + EfficientNet-b0\n    print(\"Initializing model...\")\n    model = smp.Unet(\n        encoder_name=Config.ENCODER, \n        encoder_weights=None,  # 🚨 核心修改 1: 强制设为 None，禁止联网下载\n        in_channels=3, \n        classes=1, \n        activation=None \n    )\n\n    # 🚨 核心修改 2: 手动加载你上传的本地权重\n    # 只有当 Config.WEIGHTS 为 None 时才执行这一步 (假设你在 Config 里把 WEIGHTS 改成了 None)\n    if Config.LOCAL_WEIGHTS_PATH and os.path.exists(Config.LOCAL_WEIGHTS_PATH):\n        try:\n            print(f\"Loading local weights from: {Config.LOCAL_WEIGHTS_PATH}\")\n            # 读取权重文件\n            state_dict = torch.load(Config.LOCAL_WEIGHTS_PATH, map_location='cpu')\n            \n            # 将权重载入模型的 encoder 部分\n            # strict=False 非常重要：因为 HF 的权重包含很多分类层参数，\n            # 而 SMP 的 encoder 只需要特征提取部分，strict=False 会自动忽略多余或不匹配的键\n            model.encoder.load_state_dict(state_dict, strict=False)\n            print(\"✅ Successfully loaded local pre-trained weights to encoder!\")\n        except Exception as e:\n            print(f\"⚠️ Warning: Failed to load local weights. Training from scratch. Error: {e}\")\n    else:\n        print(\"ℹ️ No local weights found or path not set. Training from scratch.\")\n\n    model.to(Config.DEVICE)\n    \n    # 3. 损失函数和优化器\n    # 组合 DiceLoss (关注重叠) 和 BCE (关注分类)\n    # Jaccard/Dice Loss 对不平衡数据（篡改区域很小）非常有效\n    loss_fn = smp.losses.DiceLoss(mode='binary', from_logits=True)\n    \n    optimizer = optim.AdamW(model.parameters(), lr=Config.LEARNING_RATE)\n    scaler = GradScaler() # 混合精度\n    scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=Config.EPOCHS, eta_min=1e-6)\n    \n    # 4. 训练循环\n    best_score = 0\n    \n    for epoch in range(Config.EPOCHS):\n        print(f\"\\nEpoch: {epoch + 1}/{Config.EPOCHS}\")\n        \n        train_loss = train_fn(train_loader, model, optimizer, loss_fn, scaler)\n        valid_loss, val_score = valid_fn(valid_loader, model, loss_fn)\n        \n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}\")\n        print(f\"Valid Loss: {valid_loss:.4f} | oF1 Score: {val_score:.4f}\")\n        \n        # 保存最优模型\n        if val_score > best_score:\n            print(f\"Score Improved ({best_score:.4f} ---> {val_score:.4f}). Saving Model...\")\n            best_score = val_score\n            torch.save(model.state_dict(), 'best_model.pth')\n    generate_submission()        \n    print(\"Training Complete!\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-08T09:52:28.263150Z","iopub.execute_input":"2025-12-08T09:52:28.263447Z","iopub.status.idle":"2025-12-08T10:53:59.297314Z","shell.execute_reply.started":"2025-12-08T09:52:28.263420Z","shell.execute_reply":"2025-12-08T10:53:59.296652Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}