{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":113558,"databundleVersionId":14878066,"sourceType":"competition"}],"dockerImageVersionId":31234,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q timm segmentation-models-pytorch albumentations opencv-python faiss-cpu\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-23T11:22:50.387470Z","iopub.execute_input":"2025-12-23T11:22:50.387755Z","iopub.status.idle":"2025-12-23T11:22:53.713292Z","shell.execute_reply.started":"2025-12-23T11:22:50.387733Z","shell.execute_reply":"2025-12-23T11:22:53.712257Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Image Visualization","metadata":{}},{"cell_type":"code","source":"import os\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom glob import glob\nimport json\nfrom scipy import ndimage\nfrom scipy.fftpack import fft2, fftshift\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Set paths\nROOT = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection\"\nTRAIN_AUTH = f\"{ROOT}/train_images/authentic\"\nTRAIN_FORG = f\"{ROOT}/train_images/forged\"\nTRAIN_MASK = f\"{ROOT}/train_masks\"\n\nprint(\"=\"*80)\nprint(\"SCIENTIFIC IMAGE FORGERY DETECTION - DATA VISUALIZATION\")\nprint(\"=\"*80)\n\n# Get file lists\nauth_images = sorted(glob(os.path.join(TRAIN_AUTH, \"*.png\")))\nforg_images = sorted(glob(os.path.join(TRAIN_FORG, \"*.png\")))\n\nprint(f\"Found {len(auth_images)} authentic images\")\nprint(f\"Found {len(forg_images)} forged images\")\n\n# Function to get corresponding mask path\ndef get_mask_path(image_path):\n    \"\"\"Get corresponding mask path for an image\"\"\"\n    basename = os.path.basename(image_path)\n    mask_name = basename.replace(\".png\", \".npy\")\n    mask_path = os.path.join(TRAIN_MASK, mask_name)\n    return mask_path if os.path.exists(mask_path) else None\n\n# Select 2 random authentic and 2 random forged images\nnp.random.seed(42)\nselected_auth = np.random.choice(auth_images, 2, replace=False)\nselected_forg = selected_auth\n\nprint(f\"\\nSelected Authentic Images:\")\nfor i, img_path in enumerate(selected_auth):\n    print(f\"  {i+1}. {os.path.basename(img_path)}\")\n    \nprint(f\"\\nSelected Forged Images:\")\nfor i, img_path in enumerate(selected_forg):\n    print(f\"  {i+1}. {os.path.basename(img_path)}\")\n\n# ============================================================================\n# RLE DECODING FUNCTIONS (from your evaluation code)\n# ============================================================================\ndef rle_decode(mask_rle, shape):\n    \"\"\"\n    Decode RLE encoded mask\n    \"\"\"\n    if mask_rle == 'authentic' or mask_rle == '':\n        return np.zeros(shape, dtype=np.uint8)\n    \n    try:\n        mask_rle = json.loads(mask_rle)\n        mask_rle = np.asarray(mask_rle, dtype=np.int32)\n        \n        # Decode RLE\n        img = np.zeros(shape[0] * shape[1], dtype=np.uint8)\n        for i in range(0, len(mask_rle), 2):\n            start = mask_rle[i] - 1\n            length = mask_rle[i + 1]\n            img[start:start+length] = 1\n        \n        return img.reshape(shape, order='F')\n    except:\n        return np.zeros(shape, dtype=np.uint8)\n\n# ============================================================================\n# IMAGE ANALYSIS FUNCTIONS\n# ============================================================================\ndef analyze_image_characteristics(img):\n    \"\"\"Analyze various image characteristics\"\"\"\n    if len(img.shape) == 3:\n        gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    else:\n        gray = img\n    \n    results = {}\n    \n    # Basic statistics\n    results['mean_intensity'] = np.mean(gray)\n    results['std_intensity'] = np.std(gray)\n    results['min_intensity'] = np.min(gray)\n    results['max_intensity'] = np.max(gray)\n    \n    # Histogram analysis\n    hist = cv2.calcHist([gray], [0], None, [256], [0, 256])\n    hist = hist / hist.sum()\n    \n    # Entropy\n    entropy = -np.sum(hist * np.log2(hist + 1e-10))\n    results['entropy'] = entropy\n    \n    # Edge analysis\n    sobelx = cv2.Sobel(gray, cv2.CV_64F, 1, 0, ksize=3)\n    sobely = cv2.Sobel(gray, cv2.CV_64F, 0, 1, ksize=3)\n    edge_magnitude = np.sqrt(sobelx**2 + sobely**2)\n    results['edge_density'] = np.mean(edge_magnitude > 10)\n    \n    # Texture analysis (Laplacian variance)\n    laplacian = cv2.Laplacian(gray, cv2.CV_64F)\n    results['texture_variance'] = np.var(laplacian)\n    \n    # Frequency analysis\n    fft = fft2(gray)\n    fft_shift = fftshift(fft)\n    magnitude = np.log1p(np.abs(fft_shift))\n    \n    # High frequency content\n    h, w = gray.shape\n    center_h, center_w = h//2, w//2\n    radius = min(h, w) // 4\n    \n    # Create circular mask for high frequencies\n    y, x = np.ogrid[:h, :w]\n    dist_from_center = np.sqrt((x - center_w)**2 + (y - center_h)**2)\n    high_freq_mask = dist_from_center > radius\n    \n    results['high_freq_energy'] = np.mean(magnitude[high_freq_mask]) if np.any(high_freq_mask) else 0\n    results['low_freq_energy'] = np.mean(magnitude[~high_freq_mask]) if np.any(~high_freq_mask) else 0\n    \n    return results\n\ndef load_and_analyze_mask(mask_path, img_shape):\n    \"\"\"Load and analyze mask file\"\"\"\n    if not mask_path or not os.path.exists(mask_path):\n        return None, None, None\n    \n    try:\n        mask_data = np.load(mask_path)\n        \n        analysis = {\n            'file_exists': True,\n            'raw_shape': mask_data.shape,\n            'raw_dtype': str(mask_data.dtype),\n            'num_masks': mask_data.shape[0] if mask_data.ndim == 3 else 1\n        }\n        \n        # Handle different mask formats\n        if mask_data.ndim == 3:\n            # Multiple masks stacked\n            masks = []\n            for i in range(mask_data.shape[0]):\n                mask = mask_data[i]\n                # Resize to match original image if needed\n                if mask.shape != img_shape[:2]:\n                    mask = cv2.resize(mask, (img_shape[1], img_shape[0]), interpolation=cv2.INTER_NEAREST)\n                masks.append(mask)\n            \n            # Create combined mask\n            combined_mask = np.zeros(img_shape[:2], dtype=np.float32)\n            for mask in masks:\n                combined_mask = np.maximum(combined_mask, mask)\n            \n            mask_display = combined_mask\n            \n            # Calculate statistics for each mask\n            mask_stats = []\n            for i, mask in enumerate(masks):\n                mask_binary = (mask > 0).astype(np.float32)\n                stats = {\n                    'mask_index': i,\n                    'coverage': mask_binary.mean(),\n                    'total_pixels': mask_binary.sum(),\n                    'unique_values': np.unique(mask),\n                    'shape': mask.shape\n                }\n                mask_stats.append(stats)\n            \n            analysis['mask_stats'] = mask_stats\n            \n        elif mask_data.ndim == 2:\n            # Single mask\n            mask = mask_data\n            if mask.shape != img_shape[:2]:\n                mask = cv2.resize(mask, (img_shape[1], img_shape[0]), interpolation=cv2.INTER_NEAREST)\n            \n            mask_display = mask\n            mask_binary = (mask > 0).astype(np.float32)\n            \n            mask_stats = [{\n                'mask_index': 0,\n                'coverage': mask_binary.mean(),\n                'total_pixels': mask_binary.sum(),\n                'unique_values': np.unique(mask),\n                'shape': mask.shape\n            }]\n            analysis['mask_stats'] = mask_stats\n            \n        else:\n            # Unexpected format\n            return None, None, None\n        \n        analysis['combined_coverage'] = mask_display.mean()\n        analysis['combined_nonzero'] = np.sum(mask_display > 0)\n        \n        return analysis, mask_display, masks if mask_data.ndim == 3 else [mask_display]\n        \n    except Exception as e:\n        print(f\"Error loading mask {mask_path}: {e}\")\n        return None, None, None\n\n# ============================================================================\n# VISUALIZATION FUNCTION FOR SINGLE IMAGE\n# ============================================================================\ndef visualize_single_image_comprehensive(img_path, mask_path, image_type=\"Unknown\"):\n    \"\"\"Create comprehensive visualization for a single image\"\"\"\n    \n    # Load image\n    img = cv2.imread(img_path)\n    if img is None:\n        print(f\"ERROR: Could not load image {img_path}\")\n        return None, None\n    \n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    original_h, original_w = img.shape[:2]\n    gray = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)\n    \n    # Analyze image characteristics\n    img_analysis = analyze_image_characteristics(img)\n    \n    # Load and analyze mask\n    mask_analysis, mask_display, individual_masks = load_and_analyze_mask(mask_path, img.shape)\n    \n    # Create figure\n    fig = plt.figure(figsize=(20, 15))\n    \n    # 1. Original Image\n    ax1 = plt.subplot(3, 5, 1)\n    ax1.imshow(img)\n    ax1.set_title(f'{image_type}\\nOriginal\\n{original_w}x{original_h}', \n                  fontsize=10, fontweight='bold',\n                  color='green' if image_type=='Authentic' else 'red')\n    ax1.axis('off')\n    \n    # Add text overlay with basic info\n    ax1.text(0.02, 0.02, f'File: {os.path.basename(img_path)}', \n             transform=ax1.transAxes, fontsize=8, color='white',\n             bbox=dict(boxstyle=\"round,pad=0.3\", facecolor=\"black\", alpha=0.7))\n    \n    # 2. Grayscale with intensity analysis\n    ax2 = plt.subplot(3, 5, 2)\n    im2 = ax2.imshow(gray, cmap='gray')\n    ax2.set_title(f'Grayscale\\nMean={gray.mean():.1f}, Std={gray.std():.1f}', fontsize=10)\n    ax2.axis('off')\n    plt.colorbar(im2, ax=ax2, fraction=0.046, pad=0.04)\n    \n    # 3. Color Histograms\n    ax3 = plt.subplot(3, 5, 3)\n    colors = ('red', 'green', 'blue')\n    for i, color in enumerate(colors):\n        hist = cv2.calcHist([img], [i], None, [256], [0, 256])\n        ax3.plot(hist, color=color, alpha=0.7, linewidth=1)\n    ax3.set_title('Color Histograms', fontsize=10)\n    ax3.set_xlabel('Intensity')\n    ax3.set_ylabel('Frequency')\n    ax3.grid(True, alpha=0.3)\n    ax3.set_xlim([0, 256])\n    \n    # 4. Grayscale Histogram\n    ax4 = plt.subplot(3, 5, 4)\n    ax4.hist(gray.flatten(), bins=256, range=[0, 256], density=True, \n             alpha=0.7, color='gray', edgecolor='black', linewidth=0.5)\n    ax4.set_title(f'Grayscale Histogram\\nEntropy={img_analysis[\"entropy\"]:.2f}', fontsize=10)\n    ax4.set_xlabel('Intensity')\n    ax4.set_ylabel('Density')\n    ax4.grid(True, alpha=0.3)\n    \n    # 5. FFT Magnitude Spectrum\n    ax5 = plt.subplot(3, 5, 5)\n    fft = fft2(gray)\n    fft_shift = fftshift(fft)\n    magnitude = np.log1p(np.abs(fft_shift))\n    im5 = ax5.imshow(magnitude, cmap='hot')\n    ax5.set_title(f'FFT Magnitude\\nHF={img_analysis[\"high_freq_energy\"]:.2f}', fontsize=10)\n    ax5.axis('off')\n    plt.colorbar(im5, ax=ax5, fraction=0.046, pad=0.04)\n    \n    # 6. Sobel Edge Detection\n    ax6 = plt.subplot(3, 5, 6)\n    sobelx = cv2.Sobel(gray, cv2.CV_64F, 1, 0, ksize=3)\n    sobely = cv2.Sobel(gray, cv2.CV_64F, 0, 1, ksize=3)\n    sobel_mag = np.sqrt(sobelx**2 + sobely**2)\n    im6 = ax6.imshow(sobel_mag, cmap='hot')\n    ax6.set_title(f'Sobel Edges\\nDensity={img_analysis[\"edge_density\"]:.3f}', fontsize=10)\n    ax6.axis('off')\n    plt.colorbar(im6, ax=ax6, fraction=0.046, pad=0.04)\n    \n    # 7. Laplacian (Texture)\n    ax7 = plt.subplot(3, 5, 7)\n    laplacian = cv2.Laplacian(gray, cv2.CV_64F)\n    laplacian_abs = np.abs(laplacian)\n    im7 = ax7.imshow(laplacian_abs, cmap='hot')\n    ax7.set_title(f'Laplacian\\nVariance={img_analysis[\"texture_variance\"]:.1f}', fontsize=10)\n    ax7.axis('off')\n    plt.colorbar(im7, ax=ax7, fraction=0.046, pad=0.04)\n    \n    # 8. Noise Residual\n    ax8 = plt.subplot(3, 5, 8)\n    blur = cv2.GaussianBlur(gray, (5, 5), 0)\n    residual = cv2.absdiff(gray, blur)\n    im8 = ax8.imshow(residual, cmap='hot')\n    ax8.set_title('Noise Residual', fontsize=10)\n    ax8.axis('off')\n    plt.colorbar(im8, ax=ax8, fraction=0.046, pad=0.04)\n    \n    # 9. CLAHE Enhanced\n    ax9 = plt.subplot(3, 5, 9)\n    clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n    clahe_img = clahe.apply(gray)\n    im9 = ax9.imshow(clahe_img, cmap='gray')\n    ax9.set_title('CLAHE Enhanced', fontsize=10)\n    ax9.axis('off')\n    plt.colorbar(im9, ax=ax9, fraction=0.046, pad=0.04)\n    \n    # 10. Combined Mask Visualization\n    ax10 = plt.subplot(3, 5, 10)\n    if mask_display is not None and np.any(mask_display > 0):\n        im10 = ax10.imshow(mask_display, cmap='hot')\n        coverage = mask_display.mean()\n        ax10.set_title(f'Combined Mask\\nCoverage={coverage:.6f}', fontsize=10)\n        plt.colorbar(im10, ax=ax10, fraction=0.046, pad=0.04)\n    else:\n        ax10.text(0.5, 0.5, 'NO MASK\\nOR ALL ZEROS', \n                 horizontalalignment='center', verticalalignment='center',\n                 transform=ax10.transAxes, fontsize=12, color='red',\n                 fontweight='bold')\n        ax10.set_title('Mask Status', fontsize=10)\n    ax10.axis('off')\n    \n    # 11. Image Overlay with Mask\n    ax11 = plt.subplot(3, 5, 11)\n    ax11.imshow(img)\n    if mask_display is not None and np.any(mask_display > 0):\n        # Create colored overlay\n        overlay = np.zeros((original_h, original_w, 3), dtype=np.float32)\n        overlay[..., 0] = mask_display  # Red channel for mask\n        ax11.imshow(overlay, alpha=0.5)\n        ax11.set_title('Image + Mask Overlay', fontsize=10)\n        \n        # Highlight mask areas with boxes\n        y_coords, x_coords = np.where(mask_display > 0)\n        if len(y_coords) > 0:\n            y_min, y_max = y_coords.min(), y_coords.max()\n            x_min, x_max = x_coords.min(), x_coords.max()\n            \n            # Draw bounding box\n            rect = plt.Rectangle((x_min, y_min), x_max-x_min, y_max-y_min,\n                                linewidth=2, edgecolor='yellow', facecolor='none')\n            ax11.add_patch(rect)\n    else:\n        ax11.set_title('No Mask Overlay', fontsize=10)\n    ax11.axis('off')\n    \n    # 12. Individual Masks (if multiple)\n    ax12 = plt.subplot(3, 5, 12)\n    if individual_masks and len(individual_masks) > 1:\n        # Show first individual mask\n        mask1 = individual_masks[0]\n        im12 = ax12.imshow(mask1, cmap='hot')\n        coverage1 = mask1.mean() if np.any(mask1 > 0) else 0\n        ax12.set_title(f'Mask 1/2\\nCoverage={coverage1:.6f}', fontsize=10)\n        plt.colorbar(im12, ax=ax12, fraction=0.046, pad=0.04)\n    elif individual_masks and len(individual_masks) == 1:\n        # Show the single mask\n        mask1 = individual_masks[0]\n        im12 = ax12.imshow(mask1, cmap='hot')\n        coverage1 = mask1.mean() if np.any(mask1 > 0) else 0\n        ax12.set_title(f'Single Mask\\nCoverage={coverage1:.6f}', fontsize=10)\n        plt.colorbar(im12, ax=ax12, fraction=0.046, pad=0.04)\n    else:\n        ax12.text(0.5, 0.5, 'NO\\nINDIVIDUAL\\nMASKS', \n                 horizontalalignment='center', verticalalignment='center',\n                 transform=ax12.transAxes, fontsize=10, color='gray')\n        ax12.set_title('Individual Masks', fontsize=10)\n    ax12.axis('off')\n    \n    # 13. Second Individual Mask (if exists)\n    ax13 = plt.subplot(3, 5, 13)\n    if individual_masks and len(individual_masks) > 1:\n        # Show second individual mask\n        mask2 = individual_masks[1]\n        im13 = ax13.imshow(mask2, cmap='hot')\n        coverage2 = mask2.mean() if np.any(mask2 > 0) else 0\n        ax13.set_title(f'Mask 2/2\\nCoverage={coverage2:.6f}', fontsize=10)\n        plt.colorbar(im13, ax=ax13, fraction=0.046, pad=0.04)\n    else:\n        ax13.text(0.5, 0.5, 'NO\\nSECOND\\nMASK', \n                 horizontalalignment='center', verticalalignment='center',\n                 transform=ax13.transAxes, fontsize=10, color='gray')\n        ax13.set_title('Second Mask', fontsize=10)\n    ax13.axis('off')\n    \n    # 14. Mask Histogram\n    ax14 = plt.subplot(3, 5, 14)\n    if mask_display is not None:\n        mask_flat = mask_display.flatten()\n        ax14.hist(mask_flat, bins=50, alpha=0.7, color='red')\n        ax14.set_title('Mask Value Distribution', fontsize=10)\n        ax14.set_xlabel('Mask Value')\n        ax14.set_ylabel('Frequency')\n        ax14.grid(True, alpha=0.3)\n        \n        # Add statistics text\n        stats_text = f\"Non-zero: {np.sum(mask_flat > 0):,}\\n\"\n        stats_text += f\"Max: {mask_flat.max():.3f}\\n\"\n        stats_text += f\"Min: {mask_flat.min():.3f}\"\n        ax14.text(0.95, 0.95, stats_text, transform=ax14.transAxes,\n                 fontsize=8, verticalalignment='top',\n                 horizontalalignment='right',\n                 bbox=dict(boxstyle=\"round,pad=0.3\", facecolor=\"white\", alpha=0.8))\n    else:\n        ax14.axis('off')\n        ax14.set_title('No Mask Data', fontsize=10)\n    \n    # 15. Summary Statistics\n    ax15 = plt.subplot(3, 5, 15)\n    ax15.axis('off')\n    \n    # Create comprehensive summary text\n    summary_text = f\"\"\"\n    {image_type.upper()} IMAGE ANALYSIS\n    {'='*40}\n    File: {os.path.basename(img_path)[:20]}...\n    Size: {original_w}x{original_h}\n    \n    IMAGE STATISTICS:\n    Mean Intensity: {img_analysis['mean_intensity']:.1f}\n    Contrast (Std): {img_analysis['std_intensity']:.1f}\n    Entropy: {img_analysis['entropy']:.2f}\n    Edge Density: {img_analysis['edge_density']:.3f}\n    Texture Variance: {img_analysis['texture_variance']:.1f}\n    HF Energy: {img_analysis['high_freq_energy']:.2f}\n    \n    MASK INFORMATION:\n    \"\"\"\n    \n    if mask_analysis:\n        summary_text += f\"Mask File: {'Exists' if mask_analysis['file_exists'] else 'Missing'}\\n\"\n        summary_text += f\"Raw Shape: {mask_analysis['raw_shape']}\\n\"\n        summary_text += f\"Num Masks: {mask_analysis['num_masks']}\\n\"\n        summary_text += f\"Combined Coverage: {mask_analysis.get('combined_coverage', 0):.6f}\\n\"\n        summary_text += f\"Non-zero Pixels: {mask_analysis.get('combined_nonzero', 0):,}\\n\"\n        \n        if 'mask_stats' in mask_analysis:\n            for i, stats in enumerate(mask_analysis['mask_stats']):\n                summary_text += f\"\\nMask {i+1}:\\n\"\n                summary_text += f\"  Coverage: {stats['coverage']:.6f}\\n\"\n                summary_text += f\"  Pixels: {int(stats['total_pixels']):,}\\n\"\n    else:\n        summary_text += \"No mask analysis available\\n\"\n    \n    if mask_display is not None and np.any(mask_display > 0):\n        # Calculate mask centroid\n        y_coords, x_coords = np.where(mask_display > 0)\n        if len(y_coords) > 0:\n            centroid_y = int(np.mean(y_coords))\n            centroid_x = int(np.mean(x_coords))\n            summary_text += f\"\\nForged Region:\\n\"\n            summary_text += f\"  Centroid: ({centroid_x}, {centroid_y})\\n\"\n            summary_text += f\"  Bounding Box: \"\n            summary_text += f\"{x_coords.min()}-{x_coords.max()} x \"\n            summary_text += f\"{y_coords.min()}-{y_coords.max()}\\n\"\n    \n    ax15.text(0, 1, summary_text, fontsize=9, verticalalignment='top',\n             transform=ax15.transAxes, fontfamily='monospace',\n             bbox=dict(boxstyle=\"round,pad=0.5\", facecolor=\"lightgray\", alpha=0.8))\n    \n    plt.suptitle(f'{image_type} - {os.path.basename(img_path)}', \n                fontsize=16, fontweight='bold', y=0.98)\n    plt.tight_layout()\n    plt.show()\n    \n    return img_analysis, mask_analysis\n\n# ============================================================================\n# VISUALIZE SELECTED IMAGES\n# ============================================================================\nprint(\"\\n\" + \"=\"*80)\nprint(\"VISUALIZING AUTHENTIC IMAGES\")\nprint(\"=\"*80)\n\nauthentic_analyses = []\nfor i, img_path in enumerate(selected_auth):\n    print(f\"\\n{'='*70}\")\n    print(f\"AUTHENTIC IMAGE {i+1}/{len(selected_auth)}\")\n    print(f\"{'='*70}\")\n    \n    mask_path = get_mask_path(img_path)\n    print(f\"Image: {os.path.basename(img_path)}\")\n    print(f\"Mask: {os.path.basename(mask_path) if mask_path else 'None'}\")\n    \n    img_analysis, mask_analysis = visualize_single_image_comprehensive(\n        img_path, mask_path, image_type=\"Authentic\"\n    )\n    \n    if img_analysis:\n        authentic_analyses.append(img_analysis)\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"VISUALIZING FORGED IMAGES\")\nprint(\"=\"*80)\n\nforged_analyses = []\nfor i, img_path in enumerate(selected_forg):\n    print(f\"\\n{'='*70}\")\n    print(f\"FORGED IMAGE {i+1}/{len(selected_forg)}\")\n    print(f\"{'='*70}\")\n    \n    mask_path = get_mask_path(img_path)\n    print(f\"Image: {os.path.basename(img_path)}\")\n    print(f\"Mask: {os.path.basename(mask_path) if mask_path else 'None'}\")\n    \n    img_analysis, mask_analysis = visualize_single_image_comprehensive(\n        img_path, mask_path, image_type=\"Forged\"\n    )\n    \n    if img_analysis:\n        forged_analyses.append(img_analysis)\n\n# ============================================================================\n# COMPARATIVE ANALYSIS\n# ============================================================================\nif authentic_analyses and forged_analyses:\n    print(\"\\n\" + \"=\"*80)\n    print(\"COMPARATIVE ANALYSIS: AUTHENTIC vs FORGED\")\n    print(\"=\"*80)\n    \n    # Create summary statistics\n    import pandas as pd\n    \n    auth_df = pd.DataFrame(authentic_analyses)\n    forg_df = pd.DataFrame(forged_analyses)\n    \n    print(\"\\nAVERAGE STATISTICS:\")\n    print(\"\\nAuthentic Images:\")\n    for col in auth_df.columns:\n        print(f\"  {col:20s}: {auth_df[col].mean():.3f} ± {auth_df[col].std():.3f}\")\n    \n    print(\"\\nForged Images:\")\n    for col in forg_df.columns:\n        print(f\"  {col:20s}: {forg_df[col].mean():.3f} ± {forg_df[col].std():.3f}\")\n    \n    # Create comparison visualization\n    fig, axes = plt.subplots(2, 3, figsize=(15, 10))\n    \n    metrics_to_compare = ['mean_intensity', 'std_intensity', 'entropy', \n                         'edge_density', 'texture_variance', 'high_freq_energy']\n    \n    for idx, metric in enumerate(metrics_to_compare):\n        ax = axes[idx//3, idx%3]\n        \n        auth_values = auth_df[metric].values if metric in auth_df.columns else []\n        forg_values = forg_df[metric].values if metric in forg_df.columns else []\n        \n        # Box plot\n        bp = ax.boxplot([auth_values, forg_values], \n                       positions=[1, 2], \n                       patch_artist=True,\n                       widths=0.6,\n                       showmeans=True,\n                       meanline=True,\n                       meanprops=dict(color='black', linewidth=2),\n                       medianprops=dict(color='yellow', linewidth=2))\n        \n        # Color boxes\n        bp['boxes'][0].set_facecolor('green')\n        bp['boxes'][0].set_alpha(0.7)\n        bp['boxes'][1].set_facecolor('red')\n        bp['boxes'][1].set_alpha(0.7)\n        \n        # Add individual points\n        for i, (values, color) in enumerate(zip([auth_values, forg_values], ['green', 'red'])):\n            x = np.random.normal(i+1, 0.1, len(values))\n            ax.scatter(x, values, alpha=0.6, color=color, s=50, edgecolor='black')\n        \n        ax.set_xticks([1, 2])\n        ax.set_xticklabels(['Authentic', 'Forged'], fontweight='bold')\n        ax.set_title(metric.replace('_', ' ').title(), fontsize=12, fontweight='bold')\n        ax.grid(True, alpha=0.3)\n        ax.set_ylabel('Value')\n    \n    plt.suptitle('Comparison: Authentic vs Forged Image Characteristics', \n                fontsize=14, fontweight='bold', y=1.02)\n    plt.tight_layout()\n    plt.show()\n    \n    # Statistical significance\n    print(\"\\nDIFFERENCES (Forged - Authentic):\")\n    print(\"-\" * 40)\n    for col in auth_df.columns:\n        if col in forg_df.columns:\n            auth_mean = auth_df[col].mean()\n            forg_mean = forg_df[col].mean()\n            diff = forg_mean - auth_mean\n            pct_diff = (diff / auth_mean) * 100 if auth_mean != 0 else 0\n            \n            significance = \"\"\n            if abs(pct_diff) > 10:\n                significance = \"** SIGNIFICANT **\"\n            \n            print(f\"{col:20s}: {diff:+.3f} ({pct_diff:+.1f}%) {significance}\")\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"VISUALIZATION COMPLETE\")\nprint(\"=\"*80)\nprint(\"\\nKey observations from this analysis:\")\nprint(\"1. Mask coverage is VERY low (typically < 1%)\")\nprint(\"2. Masks mark tiny duplicated regions (copy-move forgery)\")\nprint(\"3. Some forged images have multiple masks (multiple regions)\")\nprint(\"4. Authentic images typically don't have masks (or have all-zero masks)\")\nprint(\"5. This is a LOCAL ANOMALY DETECTION problem, not semantic segmentation\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T13:44:13.248748Z","iopub.execute_input":"2025-12-23T13:44:13.249288Z","iopub.status.idle":"2025-12-23T13:44:24.177961Z","shell.execute_reply.started":"2025-12-23T13:44:13.249258Z","shell.execute_reply":"2025-12-23T13:44:24.177412Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training - Inference - Evalusation","metadata":{}},{"cell_type":"code","source":"# ==================================================================================\n# ADVANCED DUAL-STREAM FORGERY DETECTION PIPELINE\n# Features: SimCLR Pre-training, ASPP, Cross-Attention, Forensic Augmentations\n# ==================================================================================\nimport os\nimport cv2\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport timm\nfrom PIL import Image, ImageChops, ImageEnhance\nimport io\nfrom tqdm import tqdm\nfrom glob import glob\nimport random\nimport warnings\n\n# Suppress warnings for cleaner output\nwarnings.filterwarnings(\"ignore\")\n\n# ==================================================================================\n# 1. CENTRAL CONFIGURATION\n# ==================================================================================\nclass Config:\n    # --- System ---\n    SEED = 42\n    DEVICE = \"cuda\" if torch.cuda.is_available() else \"cpu\"\n    NUM_WORKERS = 4\n    \n    # --- Paths (Adjust as needed) ---\n    ROOT_DIR = \"/kaggle/input/recodai-luc-scientific-image-forgery-detection\" \n    TRAIN_AUTH = os.path.join(ROOT_DIR, \"train_images/authentic\")\n    TRAIN_FORG = os.path.join(ROOT_DIR, \"train_images/forged\")\n    TRAIN_MASK = os.path.join(ROOT_DIR, \"train_masks\")\n    \n    # --- Model Hyperparameters ---\n    IMG_SIZE = 384\n    BATCH_SIZE = 8       # Lower this if OOM occurs\n    LR_PRETRAIN = 1e-3   # Learning rate for SimCLR\n    LR_FINETUNE = 1e-4   # Learning rate for Dual Stream\n    \n    # --- Training Control ---\n    EPOCHS_PRETRAIN =20\n    PATIENCE_PRETRAIN = 5\n    \n    EPOCHS_FINETUNE = 60\n    PATIENCE_FINETUNE = 7\n    \n    # --- Feature Engineering ---\n    ELA_QUALITY = 90     # Quality level for Error Level Analysis difference\n\nprint(f\"Using Device: {Config.DEVICE}\")\n\n# Seeding for Reproducibility\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n\nseed_everything(Config.SEED)\n\n# ==================================================================================\n# 2. UTILITY FUNCTIONS\n# ==================================================================================\nclass EarlyStopping:\n    \"\"\"Stops training if validation loss doesn't improve after a given patience.\"\"\"\n    def __init__(self, patience=5, min_delta=0, path='checkpoint.pth'):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.path = path\n        self.counter = 0\n        self.best_loss = None\n        self.early_stop = False\n\n    def __call__(self, val_loss, model):\n        if self.best_loss is None:\n            self.best_loss = val_loss\n            self.save_checkpoint(val_loss, model)\n        elif val_loss > self.best_loss - self.min_delta:\n            self.counter += 1\n            print(f'EarlyStopping counter: {self.counter} out of {self.patience}')\n            if self.counter >= self.patience:\n                self.early_stop = True\n        else:\n            self.best_loss = val_loss\n            self.save_checkpoint(val_loss, model)\n            self.counter = 0\n\n    def save_checkpoint(self, val_loss, model):\n        torch.save(model.state_dict(), self.path)\n        print(f'Validation loss decreased ({self.best_loss:.6f} --> {val_loss:.6f}).  Saving model...')\n\ndef compute_ela(img_np, quality=90):\n    \"\"\"Computes Error Level Analysis to highlight compression artifacts.\"\"\"\n    pil_img = Image.fromarray(img_np)\n    buffer = io.BytesIO()\n    pil_img.save(buffer, 'JPEG', quality=quality)\n    buffer.seek(0)\n    compressed_img = Image.open(buffer)\n    \n    ela_img = ImageChops.difference(pil_img, compressed_img)\n    \n    extrema = ela_img.getextrema()\n    max_diff = max([ex[1] for ex in extrema])\n    if max_diff == 0: max_diff = 1\n    scale = 255.0 / max_diff\n    \n    ela_img = ImageEnhance.Brightness(ela_img).enhance(scale)\n    return np.array(ela_img)\n\ndef compute_of1_score(pred_masks, gt_masks, threshold=0.5):\n    \"\"\"Metric: Optical F1 Score.\"\"\"\n    scores = []\n    pred_masks = pred_masks.detach().cpu().numpy()\n    gt_masks = gt_masks.detach().cpu().numpy()\n    \n    for pred, gt in zip(pred_masks, gt_masks):\n        pred_bin = (pred > threshold).astype(np.float32).flatten()\n        gt_bin = gt.flatten()\n        \n        if gt_bin.sum() == 0:\n            scores.append(1.0 if pred_bin.sum() == 0 else 0.0)\n            continue\n            \n        tp = np.sum((pred_bin == 1) & (gt_bin == 1))\n        fp = np.sum((pred_bin == 1) & (gt_bin == 0))\n        fn = np.sum((pred_bin == 0) & (gt_bin == 1))\n        \n        precision = tp / (tp + fp + 1e-6)\n        recall = tp / (tp + fn + 1e-6)\n        f1 = 2 * (precision * recall) / (precision + recall + 1e-6)\n        scores.append(f1)\n        \n    return np.mean(scores)\n\n# ==================================================================================\n# 3. ADVANCED DATASET WITH FORENSIC AUGMENTATIONS\n# ==================================================================================\nclass AdvancedForgeryDataset(Dataset):\n    def __init__(self, img_paths, mask_paths=None, mode='finetune'):\n        self.img_paths = img_paths\n        self.mask_paths = mask_paths\n        self.mode = mode\n        \n        # Base Transform\n        self.base_transform = A.Compose([\n            A.Resize(height=Config.IMG_SIZE, width=Config.IMG_SIZE),\n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2()\n        ], additional_targets={'ela': 'image'})\n        \n        # Forensic-Aware Augmentations (Replacing standard SimCLR)\n        # These mimic copy-move artifacts without destroying high-freq details\n        self.forensic_transform = A.Compose([\n            A.Resize(height=Config.IMG_SIZE, width=Config.IMG_SIZE),\n            A.HorizontalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            \n            # Forensic specifics:\n            A.ImageCompression(quality_lower=60, quality_upper=100, p=0.5),\n            A.GaussNoise(var_limit=(10.0, 50.0), p=0.3),\n            A.Downscale(scale_min=0.5, scale_max=0.9, p=0.3),\n            \n            A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n            ToTensorV2()\n        ])\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        img = cv2.imread(img_path)\n        if img is None:\n            img = np.zeros((Config.IMG_SIZE, Config.IMG_SIZE, 3), dtype=np.uint8)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        \n        # --- PHASE 1: PRE-TRAINING ---\n        if self.mode == 'pretrain':\n            # Create two views with forensic distortions\n            view1 = self.forensic_transform(image=img)['image']\n            view2 = self.forensic_transform(image=img)['image']\n            return view1, view2\n            \n        # --- PHASE 2: FINE-TUNING ---\n        else:\n            ela_img = compute_ela(img, quality=Config.ELA_QUALITY)\n            mask = np.zeros(img.shape[:2], dtype=np.float32)\n            \n            if self.mask_paths and self.mask_paths[idx] and os.path.exists(self.mask_paths[idx]):\n                try:\n                    m = np.load(self.mask_paths[idx])\n                    if m.ndim == 3: m = np.max(m, axis=0)\n                    mask = m.astype(np.float32)\n                except: pass\n\n            transformed = self.base_transform(image=img, mask=mask, ela=ela_img)\n            \n            img_tensor = transformed['image']       # [3, H, W]\n            mask_tensor = transformed['mask'].unsqueeze(0) # [1, H, W]\n            \n            # Create Artifact Stream Input (Green + ELA)\n            ela_tensor = transformed['ela']\n            green_channel = img_tensor[1:2, :, :]\n            artifact_input = torch.cat([green_channel, ela_tensor], dim=0) # [4, H, W]\n\n            label = 0.0 if \"authentic\" in img_path else 1.0\n            \n            return {\n                'rgb': img_tensor,\n                'artifact': artifact_input,\n                'mask': mask_tensor,\n                'label': torch.tensor(label).float()\n            }\n\n# ==================================================================================\n# 4. ADVANCED MODEL COMPONENTS\n# ==================================================================================\n\n# --- 1. Spectral Gating (Frequency Domain Analysis) ---\n# --- 1. Spectral Gating (Frequency Domain Analysis) ---\nclass SpectralFilter(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.complex_weight = nn.Parameter(torch.randn(dim, 2, dtype=torch.float32) * 0.02)\n\n    def forward(self, x):\n        # FIX: Force Float32 for FFT to avoid cuFFT power-of-2 error in FP16\n        x_in = x.float() \n        \n        x_fft = torch.fft.rfft2(x_in, norm='ortho')\n        weight = torch.view_as_complex(self.complex_weight)\n        \n        # Gating\n        x_fft = x_fft * weight.view(1, -1, 1, 1)\n        \n        # Inverse FFT\n        x = torch.fft.irfft2(x_fft, norm='ortho')\n        \n        # Cast back to original type (FP16) if needed\n        return x.type_as(x)\n\n# --- 2. Cross-Attention Fusion ---\nclass CrossAttentionFusion(nn.Module):\n    def __init__(self, dim):\n        super().__init__()\n        self.query = nn.Conv2d(dim, dim // 8, 1)\n        self.key   = nn.Conv2d(dim, dim // 8, 1)\n        self.value = nn.Conv2d(dim, dim, 1)\n        self.gamma = nn.Parameter(torch.zeros(1)) \n\n    def forward(self, rgb, artifact):\n        B, C, H, W = rgb.size()\n        proj_query = self.query(rgb).view(B, -1, W*H).permute(0, 2, 1)\n        proj_key = self.key(artifact).view(B, -1, W*H)\n        energy = torch.bmm(proj_query, proj_key)\n        attention = F.softmax(energy, dim=-1)\n        proj_value = self.value(artifact).view(B, -1, W*H)\n        out = torch.bmm(proj_value, attention.permute(0, 2, 1))\n        out = out.view(B, C, H, W)\n        out = self.gamma * out + rgb\n        return out\n\n# --- 3. ASPP (Atrous Spatial Pyramid Pooling) ---\nclass ASPP(nn.Module):\n    def __init__(self, in_channels, out_channels):\n        super(ASPP, self).__init__()\n        self.conv1 = nn.Conv2d(in_channels, out_channels, 1, bias=False)\n        self.bn1 = nn.BatchNorm2d(out_channels)\n        self.conv2 = nn.Conv2d(in_channels, out_channels, 3, padding=6, dilation=6, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n        self.conv3 = nn.Conv2d(in_channels, out_channels, 3, padding=12, dilation=12, bias=False)\n        self.bn3 = nn.BatchNorm2d(out_channels)\n        self.conv4 = nn.Conv2d(in_channels, out_channels, 3, padding=18, dilation=18, bias=False)\n        self.bn4 = nn.BatchNorm2d(out_channels)\n        self.avg_pool = nn.AdaptiveAvgPool2d(1)\n        self.conv5 = nn.Conv2d(in_channels, out_channels, 1, bias=False)\n        self.bn5 = nn.BatchNorm2d(out_channels)\n        self.conv_f = nn.Conv2d(out_channels * 5, out_channels, 1, bias=False)\n        self.bn_f = nn.BatchNorm2d(out_channels)\n        self.relu = nn.ReLU()\n\n    def forward(self, x):\n        x1 = self.relu(self.bn1(self.conv1(x)))\n        x2 = self.relu(self.bn2(self.conv2(x)))\n        x3 = self.relu(self.bn3(self.conv3(x)))\n        x4 = self.relu(self.bn4(self.conv4(x)))\n        x5 = self.avg_pool(x)\n        x5 = self.relu(self.bn5(self.conv5(x5)))\n        x5 = F.interpolate(x5, size=x4.size()[2:], mode='bilinear', align_corners=True)\n        x = torch.cat((x1, x2, x3, x4, x5), dim=1)\n        x = self.relu(self.bn_f(self.conv_f(x)))\n        return x\n\n# --- 4. Main Dual-Stream Network ---\nclass DualStreamForgeryNet(nn.Module):\n    def __init__(self, rgb_backbone='efficientnet_b3', artifact_backbone='efficientnet_b0'):\n        super().__init__()\n        # Streams\n        self.rgb_stream = timm.create_model(rgb_backbone, pretrained=True, features_only=True)\n        self.art_stream = timm.create_model(artifact_backbone, pretrained=True, features_only=True, in_chans=4)\n        \n        rgb_ch = self.rgb_stream.feature_info.channels()[-1]\n        art_ch = self.art_stream.feature_info.channels()[-1]\n        \n        # Modules\n        self.freq_gate = SpectralFilter(art_ch)\n        self.adapter = nn.Conv2d(rgb_ch + art_ch, 256, 1) # Align dims for ASPP\n        self.fusion = CrossAttentionFusion(256)\n        self.aspp = ASPP(256, 256)\n        \n        # Decoder\n        self.decoder = nn.Sequential(\n            nn.ConvTranspose2d(256, 128, 4, 2, 1), nn.ReLU(),\n            nn.ConvTranspose2d(128, 64, 4, 2, 1), nn.ReLU(),\n            nn.ConvTranspose2d(64, 32, 4, 2, 1), nn.ReLU(),\n            nn.ConvTranspose2d(32, 16, 4, 2, 1), nn.ReLU(),\n            nn.ConvTranspose2d(16, 1, 4, 2, 1)\n        )\n        \n        # Classification Head\n        self.classifier = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1), nn.Flatten(),\n            nn.Linear(256, 128), nn.ReLU(), nn.Dropout(0.3),\n            nn.Linear(128, 1)\n        )\n\n    def forward(self, rgb, artifact):\n        # Extract Features\n        f_rgb = self.rgb_stream(rgb)[-1]\n        f_art = self.art_stream(artifact)[-1]\n        \n        # Enhance Artifacts\n        f_art = self.freq_gate(f_art)\n        \n        # Align Dimensions\n        if f_rgb.shape[2:] != f_art.shape[2:]:\n            f_art = F.interpolate(f_art, size=f_rgb.shape[2:], mode='bilinear')\n            \n        # Initial Concatenation & Adapt\n        combined = torch.cat([f_rgb, f_art], dim=1)\n        combined = self.adapter(combined)\n        \n        # Advanced Fusion & Context\n        # Note: CrossAttention expects separate inputs, we split or approximate\n        # Here we apply self-attention on the combined feature map for stability\n        fused = self.fusion(combined, combined) \n        fused = self.aspp(fused)\n        \n        # Heads\n        cls_logits = self.classifier(fused)\n        mask_logits = self.decoder(fused)\n        mask_logits = F.interpolate(mask_logits, size=rgb.shape[2:], mode='bilinear')\n        \n        return mask_logits, cls_logits\n\n# --- 5. SimCLR Wrapper ---\nclass SimCLR(nn.Module):\n    def __init__(self, base_model_name='efficientnet_b3', hidden_dim=128):\n        super().__init__()\n        self.encoder = timm.create_model(base_model_name, pretrained=True, num_classes=0)\n        self.projector = nn.Sequential(\n            nn.Linear(self.encoder.num_features, self.encoder.num_features),\n            nn.ReLU(),\n            nn.Linear(self.encoder.num_features, hidden_dim)\n        )\n    def forward(self, x):\n        return self.projector(self.encoder(x))\n\n# ==================================================================================\n# 5. LOSS AND TRAINING LOGIC\n# ==================================================================================\ndef nt_xent_loss(z1, z2, temperature=0.5):\n    \"\"\"Contrastive Loss (FP16 Safe)\"\"\"\n    B = z1.size(0)\n    z = torch.cat([z1, z2], dim=0)\n    z = F.normalize(z, dim=1)\n    sim_matrix = torch.mm(z, z.t()) / temperature\n    mask = torch.eye(2 * B, dtype=torch.bool, device=z.device)\n    labels = torch.cat([torch.arange(B, 2*B), torch.arange(0, B)], dim=0).to(z.device)\n    sim_matrix.masked_fill_(mask, -1e4) # Safe value\n    return F.cross_entropy(sim_matrix, labels)\n\nclass Trainer:\n    def __init__(self, train_loader, val_loader):\n        self.train_loader = train_loader\n        self.val_loader = val_loader\n        self.scaler = GradScaler()\n\n    def run_simclr(self):\n        print(f\"\\n>>> Starting Phase 1: SimCLR Pre-training ({Config.EPOCHS_PRETRAIN} Epochs)\")\n        model = SimCLR().to(Config.DEVICE)\n        optimizer = optim.Adam(model.parameters(), lr=Config.LR_PRETRAIN)\n        early_stopping = EarlyStopping(patience=Config.PATIENCE_PRETRAIN, path='simclr_encoder.pth')\n        \n        for epoch in range(Config.EPOCHS_PRETRAIN):\n            model.train()\n            total_loss = 0\n            pbar = tqdm(self.train_loader, desc=f\"SimCLR Epoch {epoch+1}\")\n            \n            for v1, v2 in pbar:\n                v1, v2 = v1.to(Config.DEVICE), v2.to(Config.DEVICE)\n                with autocast():\n                    z1, z2 = model(v1), model(v2)\n                    loss = nt_xent_loss(z1, z2)\n                \n                optimizer.zero_grad()\n                self.scaler.scale(loss).backward()\n                self.scaler.step(optimizer)\n                self.scaler.update()\n                total_loss += loss.item()\n                pbar.set_postfix({'loss': loss.item()})\n            \n            avg_loss = total_loss / len(self.train_loader)\n            print(f\"Epoch {epoch+1} Loss: {avg_loss:.4f}\")\n            \n            # Use training loss for patience in SSL (since no val set for pretext task)\n            early_stopping(avg_loss, model.encoder)\n            if early_stopping.early_stop:\n                print(\"Early stopping triggered.\")\n                break\n        return 'simclr_encoder.pth'\n\n    def run_finetune(self, encoder_path):\n        print(f\"\\n>>> Starting Phase 2: Dual-Stream Fine-tuning ({Config.EPOCHS_FINETUNE} Epochs)\")\n        model = DualStreamForgeryNet().to(Config.DEVICE)\n        \n        # Load Pre-trained Weights\n        if os.path.exists(encoder_path):\n            print(f\"Loading SimCLR weights from {encoder_path}\")\n            model.rgb_stream.load_state_dict(torch.load(encoder_path), strict=False)\n        \n        optimizer = optim.AdamW(model.parameters(), lr=Config.LR_FINETUNE)\n        early_stopping = EarlyStopping(patience=Config.PATIENCE_FINETUNE, path='best_forgery_model.pth')\n        bce = nn.BCEWithLogitsLoss()\n        \n        for epoch in range(Config.EPOCHS_FINETUNE):\n            model.train()\n            train_loss = 0\n            pbar = tqdm(self.train_loader, desc=f\"FineTune Epoch {epoch+1}\")\n            \n            for batch in pbar:\n                rgb = batch['rgb'].to(Config.DEVICE)\n                art = batch['artifact'].to(Config.DEVICE)\n                mask = batch['mask'].to(Config.DEVICE)\n                label = batch['label'].to(Config.DEVICE).unsqueeze(1)\n                \n                with autocast():\n                    m_pred, c_pred = model(rgb, art)\n                    loss = 0.6 * bce(m_pred, mask) + 0.4 * bce(c_pred, label)\n                \n                optimizer.zero_grad()\n                self.scaler.scale(loss).backward()\n                self.scaler.step(optimizer)\n                self.scaler.update()\n                train_loss += loss.item()\n                pbar.set_postfix({'loss': loss.item()})\n            \n            # Validation\n            model.eval()\n            val_loss = 0\n            val_of1 = 0\n            with torch.no_grad():\n                for batch in self.val_loader:\n                    rgb = batch['rgb'].to(Config.DEVICE)\n                    art = batch['artifact'].to(Config.DEVICE)\n                    mask = batch['mask'].to(Config.DEVICE)\n                    label = batch['label'].to(Config.DEVICE).unsqueeze(1)\n                    \n                    m_pred, c_pred = model(rgb, art)\n                    loss = 0.6 * bce(m_pred, mask) + 0.4 * bce(c_pred, label)\n                    val_loss += loss.item()\n                    val_of1 += compute_of1_score(torch.sigmoid(m_pred), mask)\n            \n            avg_val = val_loss / len(self.val_loader)\n            print(f\"Epoch {epoch+1} | Val Loss: {avg_val:.4f} | Val oF1: {val_of1/len(self.val_loader):.4f}\")\n            \n            early_stopping(avg_val, model)\n            if early_stopping.early_stop:\n                print(\"Early stopping triggered.\")\n                break\n\n# ==================================================================================\n# 6. MAIN EXECUTION\n# ==================================================================================\n\nif __name__ == \"__main__\":\n    # 1. Prepare Data\n    print(\"Initializing Data...\")\n    all_imgs = glob(os.path.join(Config.TRAIN_AUTH, \"*.png\")) + glob(os.path.join(Config.TRAIN_FORG, \"*.png\"))\n    all_masks = []\n    for p in all_imgs:\n        m_path = os.path.join(Config.TRAIN_MASK, os.path.basename(p).replace(\".png\", \".npy\"))\n        all_masks.append(m_path if os.path.exists(m_path) else None)\n        \n    if not all_imgs:\n        raise ValueError(\"No images found! Check Config paths.\")\n        \n    # Split\n    split = int(0.8 * len(all_imgs))\n    indices = list(range(len(all_imgs)))\n    random.shuffle(indices)\n    train_idx, val_idx = indices[:split], indices[split:]\n    \n    train_imgs = [all_imgs[i] for i in train_idx]\n    train_masks = [all_masks[i] for i in train_idx]\n    val_imgs = [all_imgs[i] for i in val_idx]\n    val_masks = [all_masks[i] for i in val_idx]\n    \n    # 2. Dataloaders\n    ds_pre = AdvancedForgeryDataset(train_imgs, mode='pretrain')\n    dl_pre = DataLoader(ds_pre, batch_size=Config.BATCH_SIZE, shuffle=True, \n                        num_workers=Config.NUM_WORKERS, pin_memory=True, drop_last=True)\n    \n    ds_train = AdvancedForgeryDataset(train_imgs, train_masks, mode='finetune')\n    ds_val = AdvancedForgeryDataset(val_imgs, val_masks, mode='finetune')\n    \n    dl_train = DataLoader(ds_train, batch_size=Config.BATCH_SIZE, shuffle=True, \n                          num_workers=Config.NUM_WORKERS, pin_memory=True)\n    dl_val = DataLoader(ds_val, batch_size=Config.BATCH_SIZE, shuffle=False, \n                        num_workers=Config.NUM_WORKERS, pin_memory=True)\n    \n    # 3. Run Pipeline\n    trainer = Trainer(dl_train, dl_val)\n    \n    # Phase 1\n    trainer.train_loader = dl_pre\n    encoder_path = trainer.run_simclr()\n    \n    # Phase 2\n    trainer.train_loader = dl_train\n    trainer.run_finetune(encoder_path)\n    \n    print(\"\\n>>> Pipeline Complete. Best model saved as 'best_forgery_model.pth'\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-29T18:08:40.901761Z","iopub.execute_input":"2025-12-29T18:08:40.902726Z","execution_failed":"2025-12-29T13:45:14.226Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install faiss-cpu\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-23T15:21:54.031503Z","iopub.execute_input":"2025-12-23T15:21:54.032401Z","iopub.status.idle":"2025-12-23T15:21:58.561507Z","shell.execute_reply.started":"2025-12-23T15:21:54.032366Z","shell.execute_reply":"2025-12-23T15:21:58.560787Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}