{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":113558,"databundleVersionId":14174843,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### Revealing Forgeries through a CNN–DINOv2 Hybrid Model\n\nA segmentation-based framework for detecting scientific image forgeries — integrating deep feature embeddings with adaptive CNN reconstruction.\n\n### A. Introduction\n\nThis approach introduces a segmentation-driven hybrid architecture that integrates the self-supervised visual representation power of **DINOv2** with a compact yet effective **CNN decoder**. The DINOv2 encoder extracts rich, semantically meaningful embeddings, while the CNN component refines these features to accurately reconstruct spatial details. Together, they enable precise **pixel-level detection and localization** of image forgeries. This design balances **deep contextual understanding** with **computational efficiency**, making it suitable for scientific image authenticity verification.\n\n### Model Blueprint\n\n1. **Visual Encoder —** responsible for capturing rich, high-level semantic representations from input images using **DINOv2 embeddings**. By leveraging self-supervised learning, it learns to understand intricate visual patterns such as texture, structure, and contextual relationships without the need for labeled data. This enables the encoder to generate robust feature maps that serve as the foundation for precise forgery detection and segmentation in subsequent stages.\n\n2. **CNN Decoder —** transforms the deep feature embeddings produced by the encoder into a **binary segmentation mask** through a progressive decoding process (768 → 256 → 64 → 1). Each stage gradually upsamples and refines the spatial information, reconstructing fine-grained details lost during encoding. The final output highlights manipulated or forged regions at the **pixel level**, enabling precise localization of tampered areas. This lightweight decoder design ensures efficient computation while maintaining high segmentation accuracy.\n\n3. **Resizing —** all input images and their corresponding masks are standardized to a resolution of **256×256 pixels** to ensure consistency across the dataset. This uniform scaling allows the model to process images efficiently and maintain a fixed input size for the encoder–decoder pipeline. By normalizing dimensions, the resizing step also helps stabilize training, reduce computational overhead, and prevent distortion-related inconsistencies that could affect segmentation accuracy.\n\n### Optimization Phase\n\n1. **Forged Images —** each manipulated image is accompanied by a corresponding **binary `.npy` mask** that precisely marks the tampered regions. These masks serve as ground truth during training, guiding the model to distinguish between authentic and altered pixels. Represented in NumPy array format, they enable efficient data loading and seamless integration with deep learning frameworks, ensuring accurate supervision for segmentation-based forgery detection.\n\n2. **Authentic Images —** genuine, untampered images are paired with **empty zero-valued masks**, indicating the complete absence of manipulation. These masks act as the ground truth for unaltered samples, training the model to recognize and classify regions with no signs of forgery. By including such clean examples, the system learns to balance detection sensitivity and specificity, effectively distinguishing authentic visuals from those containing subtle or localized edits.\n\n3. **Loss Function —** the model employs **Binary Cross-Entropy with Logits Loss (BCEWithLogitsLoss)**, which combines a sigmoid activation with binary cross-entropy in a numerically stable form. This loss effectively measures the pixel-wise discrepancy between predicted and ground truth masks, making it well-suited for binary segmentation tasks such as forgery detection.\n\n4. **Optimizer —** training is performed using the **AdamW optimizer**, an enhanced variant of Adam that decouples weight decay from gradient updates. This helps prevent overfitting and ensures more stable, efficient convergence during optimization, leading to improved generalization across diverse image datasets.\n\n5. **Training Strategy —** only the **CNN decoder head** is actively trained, while the **DINOv2 encoder** remains **frozen** throughout the process. This approach leverages the pretrained self-supervised representations from DINOv2, which already capture rich semantic and structural information. By keeping the encoder fixed, the model focuses computational resources on refining the decoder’s ability to map these high-level embeddings into precise segmentation masks. This not only accelerates training but also prevents overfitting, ensuring the system maintains strong generalization across unseen images.\n\n### Prediction & Refinement\n\n1. The CNN head produces a **probability map** that highlights regions within the image most likely to contain **suspicious or manipulated areas**. Each pixel in this map represents the model’s confidence level — higher values correspond to regions the network deems more likely to exhibit tampering or anomalies. This spatial probability distribution effectively localizes potential areas of concern, serving as the foundation for subsequent **post-processing and mask refinement** steps.\n\n2. An **adaptive refinement stage** is applied to enhance the clarity and precision of detected boundaries. This stage combines **Sobel gradient filtering** to emphasize edge details and **Gaussian blurring** to smooth out noise and irregularities. The interplay between these two operations sharpens the transition zones between manipulated and authentic regions, producing cleaner, more coherent mask contours. As a result, the final probability map attains improved visual sharpness and structural consistency, making it better suited for accurate segmentation or visualization.\n\n3. A **dynamic thresholding mechanism** is employed, defined as **μ + 0.3σ**, where *μ* represents the mean and *σ* the standard deviation of the probability map values. This adaptive criterion adjusts automatically to the distribution of predictions in each image, ensuring that the threshold remains context-sensitive rather than fixed. By incorporating a fraction of the standard deviation, the method effectively enhances the separation between true positives and false positives—highlighting genuinely suspicious regions while suppressing background noise or uncertain detections.\n\n4. A **rule-based classification criterion** is applied to filter out insignificant detections. Specifically, if a detected region’s **area** is smaller than **400 pixels** or its **mean probability value** (mean_inside) falls below **0.35**, the image is categorized as **“authentic.”** This dual-condition rule helps eliminate minor or low-confidence artifacts that might arise from noise or weak model responses. By enforcing these thresholds, the system prioritizes only substantial, high-confidence regions as potential manipulations, thereby improving the overall precision and reliability of authenticity assessment.\n\n### Performance Assessment\n\nThe model’s performance is assessed using a **validation subset composed of forged images**, ensuring that evaluation focuses on manipulated content. For each sample, the **predicted mask** generated by the network is directly compared to its corresponding **ground-truth mask** using the **F1-score**, a harmonic mean of precision and recall that balances detection accuracy and completeness. This metric effectively captures how well the model identifies manipulated regions while minimizing both false positives and false negatives. The **mean F1-score** computed across all validation samples serves as a robust indicator of the model’s overall segmentation capability and consistency in distinguishing authentic from tampered areas.\n\n### Highlights\n\n1. The architecture employs a **hybrid design** that integrates the strengths of **DINOv2** and a **Convolutional Neural Network (CNN)**. The **DINOv2 encoder** contributes rich **semantic understanding**, capturing high-level contextual and representational features from the image, such as global patterns and object relationships. In contrast, the **CNN head** focuses on **spatial precision**, refining these semantic embeddings to localize subtle and fine-grained manipulations within specific regions. This combination enables the model to leverage both **global comprehension** and **localized accuracy**, resulting in a balanced and effective framework for detecting and segmenting forged or suspicious areas.\n\n2. The model is designed to be **lightweight and memory-efficient**, making it particularly well-suited for environments with **limited GPU resources**. Its compact architecture minimizes computational overhead and reduces memory consumption without compromising detection accuracy. This efficiency allows for **faster inference**, **lower power usage**, and **deployment on edge devices or mid-range hardware**, enabling practical use in real-world scenarios where high-end GPUs or cloud infrastructure may not be readily available.\n\n3. An **edge-aware post-processing stage** is applied to ensure that the resulting masks are **clean, sharp, and structurally consistent**. This stage emphasizes boundary preservation by integrating gradient-based cues and localized smoothing, allowing the model to refine transitions between manipulated and authentic regions. By selectively enhancing edge continuity while suppressing noise or fragmented artifacts, the process produces **well-defined segmentation masks** that more accurately represent the true contours of suspicious areas.\n","metadata":{}},{"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},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#!/usr/bin/env python3\n\"\"\"\nScientific Image Forgery Detection - Detection and segmentation for copy-move forgeries in biomedical images\n\"\"\"\n\nimport numpy as np\nimport pandas as pd\nimport os\nimport json\nfrom pathlib import Path\nfrom typing import List, Dict, Tuple, Optional\nimport random\nfrom collections import defaultdict\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, random_split\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import OneCycleLR\n\nimport cv2\nfrom PIL import Image\nfrom tqdm import tqdm\n\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Random seeds setting for reproducibility\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n\nset_seed(42)\n\n# Configuration\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"Memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.2f} GB\")\n\n# ============================================================================\n# CONFIGURATION\n# ============================================================================\n\nclass Config:\n    \"\"\"Configuration with optimized parameters\"\"\"\n    \n    # Base_Paths\n    \n    BASE_PATH = Path('/kaggle/input/recodai-luc-scientific-image-forgery-detection')\n    TRAIN_IMAGES_DIR = BASE_PATH / 'train_images'\n    TRAIN_MASKS_DIR = BASE_PATH / 'train_masks'\n    TEST_IMAGES_DIR = BASE_PATH / 'test_images'\n    SAMPLE_SUB_PATH = BASE_PATH / 'sample_submission.csv'\n    \n    # Models\n    \n    IMAGE_SIZE = 384\n    BATCH_SIZE = 8 if torch.cuda.is_available() else 2\n    VAL_BATCH_SIZE = 12 if torch.cuda.is_available() else 2\n    NUM_WORKERS = 0\n    \n    # Parameters\n    \n    EPOCHS = 12\n    LEARNING_RATE = 2e-3\n    WEIGHT_DECAY = 1e-5\n    VALIDATION_SPLIT = 0.15\n    \n    # Loss weights\n    \n    DICE_WEIGHT = 0.6\n    BCE_WEIGHT = 0.4\n    \n    # Detection threshold\n    \n    CLASSIFICATION_THRESHOLD = 0.35\n    SEGMENTATION_THRESHOLD = 0.45\n    MIN_AREA = 100\n    \n    # Test Time Augmentation\n    \n    TTA_ENABLED = True\n\n# ============================================================================\n# DATA DISCOVERY'S\n# ============================================================================\n\ndef discover_data():\n    \"\"\"Discover all training and test data\"\"\"\n    \n    config = Config()\n    print(\"\\n\" + \"=\"*70)\n    print(\"DATA DISCOVERY\")\n    print(\"=\"*70)\n    \n    # Discover training images\n    \n    authentic_images = []\n    forged_images = []\n    \n    authentic_dir = config.TRAIN_IMAGES_DIR / 'authentic'\n    forged_dir = config.TRAIN_IMAGES_DIR / 'forged'\n    \n    # Check authentic images\n    if authentic_dir.exists():\n        authentic_images = sorted(list(authentic_dir.glob('*.[jpJP][npNP][gG]*')))\n        print(f\"\\nAuthentic images found: {len(authentic_images)}\")\n    \n    # Check forged images\n    if forged_dir.exists():\n        forged_images = sorted(list(forged_dir.glob('*.[jpJP][npNP][gG]*')))\n        print(f\"Forged images found: {len(forged_images)}\")\n    \n    # Discover mask files\n    \n    mask_mapping = {}\n    \n    if config.TRAIN_MASKS_DIR.exists():\n        mask_files = sorted(list(config.TRAIN_MASKS_DIR.glob('*.npy')))\n        print(f\"Mask files (.npy) found: {len(mask_files)}\")\n        \n    # Create mask mapping\n    \n    for mask_file in mask_files:\n        mask_stem = mask_file.stem\n            \n    # Match with forged images\n    \n    for forged_img in forged_images:\n        if forged_img.stem == mask_stem:\n        if forged_img.stem not in mask_mapping:\n        mask_mapping[forged_img.stem] = []\n        mask_mapping[forged_img.stem].append(mask_file)\n                    \n        break\n    \n    print(f\"Images with masks: {len(mask_mapping)}\")\n    \n    # Discover test images\n    test_images = sorted(list(config.TEST_IMAGES_DIR.glob('*.[jpJP][npNP][gG]*')))\n    print(f\"Test images found: {len(test_images)}\")\n    \n    print(\"\\n\" + \"=\"*70)\n    print(f\"Total training images: {len(authentic_images) + len(forged_images)}\")\n    print(f\"  - Authentic: {len(authentic_images)}\")\n    print(f\"  - Forged: {len(forged_images)}\")\n    print(\"=\"*70)\n    \n    return authentic_images, forged_images, mask_mapping, test_images\n\n# ============================================================================\n# DATASET - FIXED VERSION\n# ============================================================================\n\nclass ForgeryDataset(Dataset):\n    \"\"\"Dataset for copy-move forgery detection with FIXED authentic/forged handling\"\"\"\n    \n    def __init__(self, authentic_paths: List[Path], forged_paths: List[Path], \n                 mask_mapping: Dict, image_size: int = 384, augment: bool = False):\n        # Store authentic and forged separately\n        self.authentic_paths = authentic_paths\n        self.forged_paths = forged_paths\n        self.all_paths = authentic_paths + forged_paths\n        self.mask_mapping = mask_mapping\n        self.image_size = image_size\n        self.augment = augment\n        \n        print(f\"\\nDataset initialized:\")\n        print(f\"  Total images: {len(self.all_paths)}\")\n        print(f\"  Authentic images: {len(self.authentic_paths)}\")\n        print(f\"  Forged images: {len(self.forged_paths)}\")\n    \n    def __len__(self):\n        return len(self.all_paths)\n    \n    def augment_image(self, image, mask):\n        \"\"\"Apply random augmentations\"\"\"\n        # Random horizontal flip\n        if random.random() > 0.5:\n            image = cv2.flip(image, 1)\n            mask = cv2.flip(mask, 1)\n        \n        # Random vertical flip\n        if random.random() > 0.5:\n            image = cv2.flip(image, 0)\n            mask = cv2.flip(mask, 0)\n        \n        # Random rotation (90, 180, 270)\n        if random.random() > 0.5:\n            k = random.choice([1, 2, 3])\n            image = np.rot90(image, k)\n            mask = np.rot90(mask, k)\n        \n        # Random brightness/contrast\n        if random.random() > 0.5:\n            alpha = random.uniform(0.8, 1.2)  # Contrast\n            beta = random.uniform(-20, 20)     # Brightness\n            image = np.clip(alpha * image + beta, 0, 255).astype(np.uint8)\n        \n        return image, mask\n    \n    def __getitem__(self, idx):\n        img_path = self.all_paths[idx]\n        \n        # Load image\n        img = cv2.imread(str(img_path))\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        orig_h, orig_w = img.shape[:2]\n        \n        # Determine if this is an authentic or forged image\n        is_forged = img_path in self.forged_paths\n        has_mask = img_path.stem in self.mask_mapping\n        \n        # Load mask\n        if is_forged and has_mask:\n            mask_paths = self.mask_mapping[img_path.stem]\n            mask = np.zeros((orig_h, orig_w), dtype=np.uint8)\n            \n            for mask_path in mask_paths:\n                # Load .npy mask\n                single_mask = np.load(str(mask_path))\n                \n                # Ensure mask is 2D\n                if single_mask.ndim > 2:\n                    single_mask = single_mask[:, :, 0] if single_mask.shape[2] == 1 else single_mask.max(axis=2)\n                \n                # Resize if needed\n                if single_mask.shape[:2] != (orig_h, orig_w):\n                    single_mask = cv2.resize(single_mask, (orig_w, orig_h), interpolation=cv2.INTER_NEAREST)\n                \n                # Binary threshold\n                single_mask = (single_mask > 0).astype(np.uint8)\n                mask = np.maximum(mask, single_mask)\n            \n            label = 1  # Forged\n        else:\n            # Authentic image - no mask\n            mask = np.zeros((orig_h, orig_w), dtype=np.uint8)\n            label = 0  # Authentic\n        \n        # Apply augmentation\n        if self.augment:\n            img, mask = self.augment_image(img, mask)\n        \n        # Resize\n        img = cv2.resize(img, (self.image_size, self.image_size))\n        mask = cv2.resize(mask, (self.image_size, self.image_size), interpolation=cv2.INTER_NEAREST)\n        \n        # Normalize image\n        img = img.astype(np.float32) / 255.0\n        \n        # Convert to tensors\n        img_tensor = torch.from_numpy(img).permute(2, 0, 1).float()\n        mask_tensor = torch.from_numpy(mask).unsqueeze(0).float()\n        label_tensor = torch.tensor([label], dtype=torch.float32)\n        \n        return img_tensor, mask_tensor, label_tensor\n\n# ============================================================================\n# MODEL ARCHITECTURE\n# ============================================================================\n\nclass ConvBlock(nn.Module):\n    \"\"\"Convolutional block with BatchNorm and ReLU\"\"\"\n    def __init__(self, in_channels, out_channels):\n        super().__init__()\n        self.conv = nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, 3, padding=1, bias=False),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        return self.conv(x)\n\nclass ForgeryDetectionModel(nn.Module):\n    \"\"\"U-Net style model for segmentation with classification head\"\"\"\n    \n    def __init__(self, in_channels=3, num_classes=1):\n        super().__init__()\n        \n        # Encoder\n        self.enc1 = ConvBlock(in_channels, 32)\n        self.pool1 = nn.MaxPool2d(2)\n        \n        self.enc2 = ConvBlock(32, 64)\n        self.pool2 = nn.MaxPool2d(2)\n        \n        self.enc3 = ConvBlock(64, 128)\n        self.pool3 = nn.MaxPool2d(2)\n        \n        self.enc4 = ConvBlock(128, 256)\n        self.pool4 = nn.MaxPool2d(2)\n        \n        # Bottleneck\n        self.bottleneck = ConvBlock(256, 512)\n        \n        # Decoder\n        self.upconv4 = nn.ConvTranspose2d(512, 256, 2, stride=2)\n        self.dec4 = ConvBlock(512, 256)\n        \n        self.upconv3 = nn.ConvTranspose2d(256, 128, 2, stride=2)\n        self.dec3 = ConvBlock(256, 128)\n        \n        self.upconv2 = nn.ConvTranspose2d(128, 64, 2, stride=2)\n        self.dec2 = ConvBlock(128, 64)\n        \n        self.upconv1 = nn.ConvTranspose2d(64, 32, 2, stride=2)\n        self.dec1 = ConvBlock(64, 32)\n        \n        # Segmentation head\n        self.seg_head = nn.Conv2d(32, num_classes, 1)\n        \n        # Classification head\n        self.global_pool = nn.AdaptiveAvgPool2d(1)\n        self.classifier = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.ReLU(inplace=True),\n            nn.Dropout(0.3),\n            nn.Linear(256, 1)\n        )\n    \n    def forward(self, x):\n        # Encoder\n        enc1 = self.enc1(x)\n        enc2 = self.enc2(self.pool1(enc1))\n        enc3 = self.enc3(self.pool2(enc2))\n        enc4 = self.enc4(self.pool3(enc3))\n        \n        # Bottleneck\n        bottleneck = self.bottleneck(self.pool4(enc4))\n        \n        # Decoder with skip connections\n        dec4 = self.upconv4(bottleneck)\n        dec4 = torch.cat([dec4, enc4], dim=1)\n        dec4 = self.dec4(dec4)\n        \n        dec3 = self.upconv3(dec4)\n        dec3 = torch.cat([dec3, enc3], dim=1)\n        dec3 = self.dec3(dec3)\n        \n        dec2 = self.upconv2(dec3)\n        dec2 = torch.cat([dec2, enc2], dim=1)\n        dec2 = self.dec2(dec2)\n        \n        dec1 = self.upconv1(dec2)\n        dec1 = torch.cat([dec1, enc1], dim=1)\n        dec1 = self.dec1(dec1)\n        \n        # Segmentation output\n        seg_out = self.seg_head(dec1)\n        \n        # Classification output\n        pooled = self.global_pool(bottleneck)\n        pooled = pooled.view(pooled.size(0), -1)\n        cls_out = self.classifier(pooled)\n        \n        return seg_out, cls_out\n\n# ============================================================================\n# LOSS FUNCTIONS\n# ============================================================================\n\nclass DiceLoss(nn.Module):\n    \"\"\"Dice loss for segmentation\"\"\"\n    def __init__(self, smooth=1.0):\n        super().__init__()\n        self.smooth = smooth\n    \n    def forward(self, pred, target):\n        pred = torch.sigmoid(pred)\n        pred = pred.view(-1)\n        target = target.view(-1)\n        \n        intersection = (pred * target).sum()\n        dice = (2. * intersection + self.smooth) / (pred.sum() + target.sum() + self.smooth)\n        \n        return 1 - dice\n\nclass CombinedLoss(nn.Module):\n    \"\"\"Combined loss for segmentation and classification\"\"\"\n    def __init__(self, dice_weight=0.6, bce_weight=0.4):\n        super().__init__()\n        self.dice_loss = DiceLoss()\n        self.bce_loss = nn.BCEWithLogitsLoss()\n        self.dice_weight = dice_weight\n        self.bce_weight = bce_weight\n    \n    def forward(self, seg_pred, cls_pred, seg_target, cls_target):\n        # Segmentation loss (only for forged images)\n        seg_loss = self.dice_loss(seg_pred, seg_target)\n        seg_bce = self.bce_loss(seg_pred, seg_target)\n        seg_combined = self.dice_weight * seg_loss + self.bce_weight * seg_bce\n        \n        # Classification loss\n        cls_loss = self.bce_loss(cls_pred, cls_target)\n        \n        # Combine losses\n        total_loss = seg_combined + cls_loss\n        \n        return total_loss, seg_loss, cls_loss\n\n# ============================================================================\n# TRAINING\n# ============================================================================\n\ndef train_epoch(model, dataloader, criterion, optimizer, scheduler, device):\n    \"\"\"Train for one epoch\"\"\"\n    model.train()\n    total_loss = 0\n    total_seg_loss = 0\n    total_cls_loss = 0\n    \n    pbar = tqdm(dataloader, desc=\"Training\")\n    for images, masks, labels in pbar:\n        images = images.to(device)\n        masks = masks.to(device)\n        labels = labels.to(device)\n        \n        optimizer.zero_grad()\n        \n        seg_out, cls_out = model(images)\n        loss, seg_loss, cls_loss = criterion(seg_out, cls_out, masks, labels)\n        \n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n        \n        total_loss += loss.item()\n        total_seg_loss += seg_loss.item()\n        total_cls_loss += cls_loss.item()\n        \n        pbar.set_postfix({\n            'loss': f'{loss.item():.4f}',\n            'seg': f'{seg_loss.item():.4f}',\n            'cls': f'{cls_loss.item():.4f}',\n            'lr': f'{scheduler.get_last_lr()[0]:.6f}'\n        })\n    \n    return total_loss / len(dataloader)\n\ndef validate(model, dataloader, criterion, device):\n    \"\"\"Validate the model\"\"\"\n    model.eval()\n    total_loss = 0\n    \n    with torch.no_grad():\n        for images, masks, labels in tqdm(dataloader, desc=\"Validation\"):\n            images = images.to(device)\n            masks = masks.to(device)\n            labels = labels.to(device)\n            \n            seg_out, cls_out = model(images)\n            loss, _, _ = criterion(seg_out, cls_out, masks, labels)\n            \n            total_loss += loss.item()\n    \n    return total_loss / len(dataloader)\n\n# ============================================================================\n# INFERENCE\n# ============================================================================\n\ndef predict_with_tta(model, image, config):\n    \"\"\"Prediction with Test Time Augmentation\"\"\"\n    model.eval()\n    \n    if not config.TTA_ENABLED:\n        with torch.no_grad():\n            seg_out, cls_out = model(image.unsqueeze(0))\n        return seg_out[0], cls_out[0]\n    \n    predictions_seg = []\n    predictions_cls = []\n    \n    # Original\n    with torch.no_grad():\n        seg, cls = model(image.unsqueeze(0))\n        predictions_seg.append(seg[0])\n        predictions_cls.append(cls[0])\n    \n    # Horizontal flip\n    img_flip = torch.flip(image, [2])\n    with torch.no_grad():\n        seg, cls = model(img_flip.unsqueeze(0))\n        seg = torch.flip(seg[0], [2])\n        predictions_seg.append(seg)\n        predictions_cls.append(cls[0])\n    \n    # Vertical flip\n    img_flip = torch.flip(image, [1])\n    with torch.no_grad():\n        seg, cls = model(img_flip.unsqueeze(0))\n        seg = torch.flip(seg[0], [1])\n        predictions_seg.append(seg)\n        predictions_cls.append(cls[0])\n    \n    # Average predictions\n    seg_final = torch.stack(predictions_seg).mean(0)\n    cls_final = torch.stack(predictions_cls).mean(0)\n    \n    return seg_final, cls_final\n\ndef rle_encode(mask):\n    \"\"\"Run-length encode a binary mask\"\"\"\n    dots = np.where(mask.T.flatten() == 1)[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\n# ============================================================================\n# MAIN PIPELINE\n# ============================================================================\n\ndef main():\n    config = Config()\n    \n    print(\"\\n\" + \"=\"*70)\n    print(\"SCIENTIFIC IMAGE FORGERY DETECTION\")\n    print(\"=\"*70)\n    \n    # Discover data\n    authentic_images, forged_images, mask_mapping, test_images = discover_data()\n    \n    # Create dataset\n    all_authentic = authentic_images\n    all_forged = forged_images\n    \n    # Split into train and validation\n    num_val_authentic = int(len(all_authentic) * config.VALIDATION_SPLIT)\n    num_val_forged = int(len(all_forged) * config.VALIDATION_SPLIT)\n    \n    # Shuffle\n    random.shuffle(all_authentic)\n    random.shuffle(all_forged)\n    \n    train_authentic = all_authentic[num_val_authentic:]\n    val_authentic = all_authentic[:num_val_authentic]\n    \n    train_forged = all_forged[num_val_forged:]\n    val_forged = all_forged[:num_val_forged]\n    \n    print(f\"\\nDataset split:\")\n    print(f\"  Training: {len(train_authentic)} authentic + {len(train_forged)} forged = {len(train_authentic) + len(train_forged)}\")\n    print(f\"  Validation: {len(val_authentic)} authentic + {len(val_forged)} forged = {len(val_authentic) + len(val_forged)}\")\n    \n    # Create datasets\n    train_dataset = ForgeryDataset(\n        train_authentic, train_forged, mask_mapping,\n        image_size=config.IMAGE_SIZE, augment=True\n    )\n    \n    val_dataset = ForgeryDataset(\n        val_authentic, val_forged, mask_mapping,\n        image_size=config.IMAGE_SIZE, augment=False\n    )\n    \n    # Create dataloaders\n    train_loader = DataLoader(\n        train_dataset, batch_size=config.BATCH_SIZE,\n        shuffle=True, num_workers=config.NUM_WORKERS, pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset, batch_size=config.VAL_BATCH_SIZE,\n        shuffle=False, num_workers=config.NUM_WORKERS, pin_memory=True\n    )\n    \n    # Initialize model\n    print(\"\\nInitializing model...\")\n    model = ForgeryDetectionModel().to(device)\n    total_params = sum(p.numel() for p in model.parameters())\n    print(f\"Model parameters: {total_params:,}\")\n    \n    # Loss and optimizer\n    criterion = CombinedLoss(\n        dice_weight=config.DICE_WEIGHT,\n        bce_weight=config.BCE_WEIGHT\n    )\n    \n    optimizer = AdamW(\n        model.parameters(),\n        lr=config.LEARNING_RATE,\n        weight_decay=config.WEIGHT_DECAY\n    )\n    \n    scheduler = OneCycleLR(\n        optimizer,\n        max_lr=config.LEARNING_RATE,\n        epochs=config.EPOCHS,\n        steps_per_epoch=len(train_loader),\n        pct_start=0.3\n    )\n    \n    # Training loop\n    print(\"\\nStarting training...\")\n    best_val_loss = float('inf')\n    \n    for epoch in range(config.EPOCHS):\n        print(f\"\\nEpoch {epoch + 1}/{config.EPOCHS}\")\n        \n        train_loss = train_epoch(model, train_loader, criterion, optimizer, scheduler, device)\n        val_loss = validate(model, val_loader, criterion, device)\n        \n        print(f\"Epoch {epoch + 1}: Train Loss = {train_loss:.4f}, Val Loss = {val_loss:.4f}\")\n        \n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            torch.save(model.state_dict(), 'best_model.pth')\n            print(f\"  -> Best model saved (Val Loss: {val_loss:.4f})\")\n    \n    # Load best model for inference\n    print(\"\\nLoading best model for inference...\")\n    model.load_state_dict(torch.load('best_model.pth'))\n    model.eval()\n    \n    # Generate predictions\n    print(\"\\nGenerating predictions...\")\n    \n    sample_sub = pd.read_csv(config.SAMPLE_SUB_PATH)\n    results = []\n    \n    for idx, row in tqdm(sample_sub.iterrows(), total=len(sample_sub), desc=\"Predicting\"):\n        case_id = str(row['case_id'])\n        \n        # Find test image\n        test_img_path = None\n        for img_path in test_images:\n            if img_path.stem == case_id:\n                test_img_path = img_path\n                break\n        \n        if test_img_path and test_img_path.exists():\n            # Load image\n            img = cv2.imread(str(test_img_path))\n            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n            orig_h, orig_w = img.shape[:2]\n            \n            # Preprocess\n            img_resized = cv2.resize(img, (config.IMAGE_SIZE, config.IMAGE_SIZE))\n            img_normalized = img_resized.astype(np.float32) / 255.0\n            img_tensor = torch.from_numpy(img_normalized).permute(2, 0, 1).float().to(device)\n            \n            # Predict with TTA\n            seg_output, cls_output = predict_with_tta(model, img_tensor, config)\n            \n            cls_prob = torch.sigmoid(cls_output).cpu().item()\n            \n            if cls_prob < config.CLASSIFICATION_THRESHOLD:\n                results.append({'case_id': int(case_id), 'annotation': 'authentic'})\n            else:\n                seg_prob = torch.sigmoid(seg_output).cpu().numpy()[0]\n                seg_mask = (seg_prob > config.SEGMENTATION_THRESHOLD).astype(np.uint8)\n                seg_mask = cv2.resize(seg_mask, (orig_w, orig_h), interpolation=cv2.INTER_NEAREST)\n                \n                # Morphological operations\n                kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3))\n                seg_mask = cv2.morphologyEx(seg_mask, cv2.MORPH_CLOSE, kernel, iterations=2)\n                seg_mask = cv2.morphologyEx(seg_mask, cv2.MORPH_OPEN, kernel)\n                \n                if seg_mask.sum() > config.MIN_AREA:\n                    # RLE encoding\n                    run_lengths = rle_encode(seg_mask)\n                    \n                    if len(run_lengths) > 0:\n                        results.append({\n                            'case_id': int(case_id),\n                            'annotation': json.dumps([int(x) for x in run_lengths])\n                        })\n                    else:\n                        results.append({'case_id': int(case_id), 'annotation': 'authentic'})\n                else:\n                    results.append({'case_id': int(case_id), 'annotation': 'authentic'})\n        else:\n            results.append({'case_id': int(case_id), 'annotation': 'authentic'})\n    \n    # Create submission\n    submission_df = pd.DataFrame(results)\n    submission_df.to_csv('submission.csv', index=False)\n    \n    print(\"\\n\" + \"=\"*70)\n    print(\"SUBMISSION COMPLETE\")\n    print(\"=\"*70)\n    print(f\"Total predictions: {len(submission_df)}\")\n    print(f\"Authentic: {(submission_df['annotation'] == 'authentic').sum()}\")\n    print(f\"Forgeries: {(submission_df['annotation'] != 'authentic').sum()}\")\n    \n    return submission_df\n\nif __name__ == \"__main__\":\n    submission = main()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}