{"metadata":{"kernelspec":{"display_name":"Python 3 (ipykernel)","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.3"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":113558,"databundleVersionId":14878066,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"5635fbde-ccc6-42c2-b90d-372247288541","cell_type":"code","source":"import matplotlib.pyplot as plt\nimport numpy as np\nimport cv2\nfrom pathlib import Path\nfrom PIL import Image\nimport math\n\n# ================= 配置路径 =================\nCOMP_ROOT = Path(\"/kaggle/input/recodai-luc-scientific-image-forgery-detection\")  #  \"/kaggle/input/...\"\nTRAIN_IMG_FORGED = COMP_ROOT / \"train_images\" / \"forged\"\nTRAIN_MASK_DIR = COMP_ROOT / \"train_masks\"\n\nSUP_IMG_DIR = COMP_ROOT / \"supplemental_images\"\nSUP_MASK_DIR = COMP_ROOT / \"supplemental_masks\"\n\nclass DatasetScanner:\n    def __init__(self, img_dir, mask_dir, dataset_name=\"Dataset\"):\n        self.img_dir = Path(img_dir)\n        self.mask_dir = Path(mask_dir)\n        self.dataset_name = dataset_name\n        self.pairs = self._load_pairs()\n        self.total = len(self.pairs)\n        self.current_idx = 0 # 记录当前阅读进度\n        print(f\"[{dataset_name}] 共加载 {self.total} 张样本。\")\n\n    def _load_pairs(self):\n        valid_ext = {\".jpg\", \".jpeg\", \".png\", \".tif\", \".tiff\"}\n        if not self.img_dir.exists():\n            return []\n        # 按文件名排序，确保顺序遍历\n        all_imgs = sorted([p for p in self.img_dir.glob(\"*\") if p.suffix.lower() in valid_ext])\n        pairs = []\n        for p in all_imgs:\n            m = self.mask_dir / f\"{p.stem}.npy\"\n            if m.exists():\n                pairs.append((p, m))\n        return pairs\n\n    def reset(self):\n        self.current_idx = 0\n        print(\"进度已重置。\")\n\n    def show_next_batch(self, batch_size=100, cols=5):\n        \"\"\"\n        核心显示函数：显示 batch_size 张叠加图\n        cols: 每行显示几张图（建议5或10）\n        \"\"\"\n        if self.current_idx >= self.total:\n            print(\">>> 已经检查完所有图片！(All done)\")\n            return\n\n        end_idx = min(self.current_idx + batch_size, self.total)\n        batch_pairs = self.pairs[self.current_idx : end_idx]\n        \n        # 计算行数\n        rows = math.ceil(len(batch_pairs) / cols)\n        \n        # 设置画布大小：每张小图宽3寸，高3寸\n        fig, axes = plt.subplots(rows, cols, figsize=(cols * 3, rows * 3))\n        axes = axes.flatten()\n        \n        print(f\"正在显示: {self.current_idx} - {end_idx} (共 {self.total} 张)\")\n\n        for i, ax in enumerate(axes):\n            if i < len(batch_pairs):\n                img_path, mask_path = batch_pairs[i]\n                try:\n                    # 读取\n                    img = Image.open(str(img_path)).convert(\"RGB\")\n                    img_np = np.array(img)\n                    mask_np = np.load(mask_path)\n                    if mask_np.ndim == 3: mask_np = mask_np[0]\n                    \n                    # 制作叠加图 (Overlay)\n                    heatmap = np.zeros_like(img_np)\n                    heatmap[:, :, 0] = 255 # 纯红\n                    \n                    overlay = img_np.copy()\n                    mask_bool = mask_np > 0\n                    \n                    # 在掩码区域混合红色\n                    overlay[mask_bool] = cv2.addWeighted(img_np[mask_bool], 0.5, heatmap[mask_bool], 0.5, 0)\n                    \n                    # 画绿色轮廓增强显示\n                    mask_u8 = (mask_bool).astype(np.uint8) * 255\n                    contours, _ = cv2.findContours(mask_u8, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n                    cv2.drawContours(overlay, contours, -1, (0, 255, 0), 3)\n\n                    ax.imshow(overlay)\n                    # 标题显示索引和ID，方便定位\n                    ax.set_title(f\"Idx:{self.current_idx + i}\\n{img_path.stem}\", fontsize=9)\n                except Exception as e:\n                    ax.text(0.5, 0.5, \"Error\", ha='center')\n                    print(f\"Error: {e}\")\n            \n            ax.axis(\"off\") # 隐藏坐标轴\n\n        plt.tight_layout()\n        plt.show()\n        \n        # 更新进度\n        self.current_idx = end_idx\n\n# ================= 初始化扫描器 =================\n# 你可以创建两个扫描器，分别对应 Phase 1 和 Phase 2 数据\nscanner_p1 = DatasetScanner(TRAIN_IMG_FORGED, TRAIN_MASK_DIR, \"Phase1_Base\")\nscanner_p2 = DatasetScanner(SUP_IMG_DIR, SUP_MASK_DIR, \"Phase2_Sup\")","metadata":{},"outputs":[],"execution_count":null},{"id":"bbd4c323-7842-48f7-bdbb-6424f9b11fbf","cell_type":"code","source":"scanner_p1.show_next_batch(batch_size=20, cols=4)","metadata":{},"outputs":[],"execution_count":null},{"id":"ecd09b4d-3fa3-41b1-af79-a29b09af8961","cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null}]}