{"cells":[{"cell_type":"markdown","metadata":{},"source":"# 🕵️‍♂️ Scientific Image Forgery Detection - Inference\n====================================================\n\nThis notebook runs inference using the **Optimized Model** trained in [this kernel](https://www.kaggle.com/code/hossam82/scientific-forgery-progressive-decoder-tta).\n\n## 🚀 Approach\n1. **Load Trained Model**: Recover the `best_model.pt` saved during training.\n2. **TTA (Test-Time Augmentation)**: Predict 4 variations (Original, H-Flip, V-Flip, Rot90) and average them.\n3. **Post-Processing**: Filter masks based on area percentage and confidence thresholds (Top-14 technique).\n4. **Submission**: Generate RLE-encoded CSV for the leaderboard.\n\n### Model Architecture\n- **Encoder**: Multi-layer DINOv2 (Layers 9-12)\n- **Decoder**: Progressive Bilinear Decoder (4-stage upsampling)\n"},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"import os\nimport gc\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom transformers import AutoModel\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\nimport warnings\nwarnings.filterwarnings('ignore')"},{"cell_type":"markdown","metadata":{},"source":"## 1. Configuration & Model Architecture 🏗️\nWe must define the exact same model class used in training to load the weights."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"class CFG:\n    # Input Paths\n    data_dir = Path(\"/kaggle/input/recodai-luc-scientific-image-forgery-detection\")\n    test_images_dir = data_dir / \"test_images\"\n    \n    # Model Path (from the training kernel output)\n    model_path = Path(\"/kaggle/input/scientific-forgery-progressive-decoder-tta/best_model.pt\")\n    \n    model_name = \"facebook/dinov2-base\"\n    dinov2_layers = [9, 10, 11, 12]\n    image_size = 518\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    \n    # Post-processing\n    min_area_percent = 0.00031\n    min_confidence = 0.3348\n    threshold = 0.5\n\nprint(f\"Device: {CFG.device}\")\n\n# ============================================================================\n# Model Classes\n# ============================================================================\nclass ProgressiveBilinearDecoder(nn.Module):\n    def __init__(self, in_features=768, target_size=518):\n        super().__init__()\n        self.target_size = target_size\n        self.block1 = nn.Sequential(nn.Conv2d(in_features, 384, 3, 1), nn.BatchNorm2d(384), nn.ReLU(True), nn.Dropout2d(0.1))\n        self.block2 = nn.Sequential(nn.Conv2d(384, 192, 3, 1), nn.BatchNorm2d(192), nn.ReLU(True), nn.Dropout2d(0.1))\n        self.block3 = nn.Sequential(nn.Conv2d(192, 96, 3, 1), nn.BatchNorm2d(96), nn.ReLU(True))\n        self.block4 = nn.Sequential(nn.Conv2d(96, 48, 3, 1), nn.BatchNorm2d(48), nn.ReLU(True))\n        self.head = nn.Conv2d(48, 1, 1)\n    \n    def forward(self, x):\n        x = self.block1(x)\n        x = F.interpolate(x, size=(74, 74), mode='bilinear', align_corners=False)\n        x = self.block2(x)\n        x = F.interpolate(x, size=(148, 148), mode='bilinear', align_corners=False)\n        x = self.block3(x)\n        x = F.interpolate(x, size=(296, 296), mode='bilinear', align_corners=False)\n        x = self.block4(x)\n        x = F.interpolate(x, size=(self.target_size, self.target_size), mode='bilinear', align_corners=False)\n        return self.head(x)\n\nclass ForgeryDetector(nn.Module):\n    def __init__(self, model_name=\"facebook/dinov2-base\", layers=[9, 10, 11, 12], target_size=518):\n        super().__init__()\n        self.encoder = AutoModel.from_pretrained(model_name)\n        self.layers = layers\n        self.decoder = ProgressiveBilinearDecoder(768, target_size)\n    \n    def forward(self, x):\n        outputs = self.encoder(x, output_hidden_states=True)\n        features = torch.stack([outputs.hidden_states[l][:, 1:, :] for l in self.layers], 0).mean(0)\n        B, N, C = features.shape\n        H = W = int(N ** 0.5)\n        x = features.permute(0, 2, 1).reshape(B, C, H, W)\n        return self.decoder(x)"},{"cell_type":"markdown","metadata":{},"source":"## 2. Utils: Transforms & RLE Encoding\nWe use the same validation transforms as training."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"def get_transforms():\n    return A.Compose([\n        A.Resize(CFG.image_size, CFG.image_size),\n        A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n        ToTensorV2(),\n    ])\n\ndef rle_encode(mask):\n    pixels = mask.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)"},{"cell_type":"markdown","metadata":{},"source":"## 3. Load Model Weights\nLoading the weights from the training kernel output."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"print(\"Loading model...\")\nmodel = ForgeryDetector(target_size=CFG.image_size)\ntry:\n    # Load weights from the previous kernel output\n    state_dict = torch.load(CFG.model_path, map_location=CFG.device)\n    model.load_state_dict(state_dict)\n    print(\"✅ Model weights loaded successfully!\")\nexcept FileNotFoundError:\n    print(f\"❌ Model file not found at {CFG.model_path}. Please check kernel sources.\")\n    # Fallback to init weights for testing if file missing (should not happen in prod)\nexcept Exception as e:\n    print(f\"❌ Error loading weights: {e}\")\n\nmodel = model.to(CFG.device)\nmodel.eval()"},{"cell_type":"markdown","metadata":{},"source":"## 4. Run Inference with TTA 🔮\nWe perform 4-way Test-Time Augmentation (Original, Flip H, Flip V, Rotate 90) for maximum accuracy."},{"cell_type":"code","execution_count":null,"metadata":{},"outputs":[],"source":"test_images = sorted(list(CFG.test_images_dir.glob(\"*.png\") if CFG.test_images_dir.exists() else []))\nprint(f\"Found {len(test_images)} test images.\")\n\nresults = []\ntransforms = get_transforms()\n\nfor path in tqdm(test_images, desc=\"Inference\"):\n    # Read image\n    img = cv2.imread(str(path))\n    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n    h, w = img.shape[:2]\n    \n    # Prepare tensor\n    tensor = transforms(image=img)['image'].unsqueeze(0).to(CFG.device)\n    \n    with torch.no_grad():\n        # 4-Way TTA\n        # 1. Original\n        p1 = torch.sigmoid(model(tensor))\n        \n        # 2. Horizontal Flip\n        p2 = torch.flip(torch.sigmoid(model(torch.flip(tensor, [3]))), [3])\n        \n        # 3. Vertical Flip\n        p3 = torch.flip(torch.sigmoid(model(torch.flip(tensor, [2]))), [2])\n        \n        # 4. Rotate 90\n        p4 = torch.rot90(torch.sigmoid(model(torch.rot90(tensor, 1, [2,3]))), -1, [2,3])\n        \n        # Average predictions\n        pred = (p1 + p2 + p3 + p4) / 4.0\n    \n    # Resize to original resolution\n    pred = F.interpolate(pred, size=(h, w), mode='bilinear', align_corners=False).squeeze().cpu().numpy()\n    \n    # Post-processing: Percentage-based filtering\n    mask_area = (pred > 0.5).sum()\n    mean_conf = pred[pred > 0.5].mean() if mask_area > 0 else 0\n    \n    rle = \"\"\n    if mask_area >= (h * w * CFG.min_area_percent) and mean_conf >= CFG.min_confidence:\n        rle = rle_encode((pred > 0.5).astype(np.uint8))\n    \n    results.append({'case_id': path.stem, 'annotation': rle if rle else 'authentic'})\n\n# Determine output filename\noutput_file = \"submission.csv\"\npd.DataFrame(results).to_csv(output_file, index=False)\nprint(f\"\\nSaved submission to {output_file}\")\n\n# Check stats\ndf = pd.DataFrame(results)\nforged_count = len(df[df['annotation'] != 'authentic'])\nprint(f\"Total predicted forged: {forged_count} / {len(df)} ({forged_count/len(df)*100:.2f}%)\")"}],"metadata":{"kernelspec":{"display_name":"Python 3","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.10.12"}},"nbformat":4,"nbformat_minor":5}