#!/usr/bin/env python3
"""
Recod.ai/LUC Scientific Image Forgery Detection - WayneIA V8.0 seed2026
CODE_KEY[166] FORGERY_CLASSIFIER_MATRIX + CODE_KEY[158] FORGERY_DETECT
Competition: recodai-luc-scientific-image-forgery-detection | Prize: $55K | Deadline: Jan 15
V8.0: seed2026 Phase 2 (Dice 0.8179) | Target: 0.40+ F1 (Top 3)
WayneIA Position_2 OpusPlan | January 5, 2026 | Year-8 RHINOCEROS G9
"""
import sys, os, warnings
warnings.filterwarnings("ignore")
import numpy as np
import pandas as pd
from pathlib import Path
import torch
import torch.nn as nn
import torch.nn.functional as F
from PIL import Image
from torchvision import transforms
from tqdm.auto import tqdm

print("="*70)
print("Recod.ai/LUC V8.0 - WayneIA seed2026 (Dice 0.8179)")
print("="*70)

IN_KAGGLE = os.path.exists("/kaggle/input")
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Kaggle: {IN_KAGGLE} | CUDA: {torch.cuda.is_available()} | Device: {DEVICE}")

IMG_SIZE, PATCH_SIZE, DINOV2_EMBED_DIM, NUM_CLASSES = 518, 14, 768, 1
CLASSIFICATION_THRESHOLD, SEGMENTATION_THRESHOLD = 0.5, 0.25

def rle_encode(mask):
    pixels = np.concatenate([[0], mask.flatten(), [0]])
    runs = np.where(pixels[1:] != pixels[:-1])[0]
    if len(runs) == 0: return "authentic"
    runs = runs.reshape(-1, 2)
    rle_pairs = []
    for start, end in runs:
        rle_pairs.extend([start, end - start])
    return f"[{" ".join(map(str, rle_pairs))}]" if rle_pairs else "authentic"

class DINOv2ForgeryDetector(nn.Module):
    def __init__(self, embed_dim=DINOV2_EMBED_DIM):
        super().__init__()
        try:
            self.backbone = torch.hub.load("facebookresearch/dinov2", "dinov2_vitb14", pretrained=True)
        except:
            self.backbone = None
        if self.backbone:
            for param in self.backbone.parameters():
                param.requires_grad = False
        self.cls_head = nn.Sequential(nn.Linear(embed_dim, 256), nn.ReLU(True), nn.Dropout(0.3), nn.Linear(256, NUM_CLASSES))
        self.seg_conv1 = nn.Conv2d(embed_dim, 256, 3, padding=1)
        self.seg_conv2 = nn.Conv2d(256, 128, 3, padding=1)
        self.seg_conv3 = nn.Conv2d(128, 64, 3, padding=1)
        self.seg_out = nn.Conv2d(64, NUM_CLASSES, 1)
        self.bn1, self.bn2, self.bn3 = nn.BatchNorm2d(256), nn.BatchNorm2d(128), nn.BatchNorm2d(64)

    def forward(self, x):
        if not self.backbone:
            return torch.zeros(x.shape[0], NUM_CLASSES, device=x.device), torch.zeros(x.shape[0], NUM_CLASSES, x.shape[2], x.shape[3], device=x.device)
        with torch.no_grad():
            features = self.backbone.forward_features(x)
        cls_out = self.cls_head(features["x_norm_clstoken"])
        pt = features["x_norm_patchtokens"]
        h = w = int(np.sqrt(pt.shape[1]))
        spatial = pt.transpose(1, 2).reshape(pt.shape[0], -1, h, w)
        x = F.relu(self.bn1(self.seg_conv1(spatial)))
        x = F.interpolate(x, scale_factor=2, mode="bilinear", align_corners=False)
        x = F.relu(self.bn2(self.seg_conv2(x)))
        x = F.interpolate(x, scale_factor=2, mode="bilinear", align_corners=False)
        x = F.relu(self.bn3(self.seg_conv3(x)))
        x = F.interpolate(x, scale_factor=4, mode="bilinear", align_corners=False)
        return cls_out, self.seg_out(x)

class V8Pipeline:
    def __init__(self, model_path=None):
        self.device = DEVICE
        self.transform = transforms.Compose([
            transforms.Resize((IMG_SIZE, IMG_SIZE)),
            transforms.ToTensor(),
            transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
        ])
        self.model = DINOv2ForgeryDetector().to(self.device)
        if model_path and Path(model_path).exists():
            print(f"Loading V8 model: {model_path}")
            try:
                sd = torch.load(model_path, map_location=self.device)
                if "model_state_dict" in sd: sd = sd["model_state_dict"]
                self.model.load_state_dict(sd, strict=False)
                print("[OK] V8 Model loaded (Dice 0.8179)")
            except Exception as e:
                print(f"Warning: {e}")
        self.model.eval()

    def predict(self, image_path):
        img = Image.open(image_path).convert("RGB")
        orig_size = img.size[::-1]
        with torch.no_grad():
            cls_logits, seg_logits = self.model(self.transform(img).unsqueeze(0).to(self.device))
        cls_prob = torch.sigmoid(cls_logits).cpu().numpy()[0, 0]
        seg_prob = torch.sigmoid(seg_logits).cpu().numpy()[0, 0]
        seg_resized = np.array(Image.fromarray((seg_prob * 255).astype(np.uint8)).resize((orig_size[1], orig_size[0]), Image.BILINEAR)) / 255.0
        mask = (seg_resized > SEGMENTATION_THRESHOLD).astype(np.uint8)
        is_forged = cls_prob > CLASSIFICATION_THRESHOLD
        forgery_ratio = mask.sum() / mask.size
        if not is_forged and forgery_ratio > 0.03:
            is_forged = True
        return {"is_forged": is_forged, "confidence": float(cls_prob), "rle": rle_encode(mask) if is_forged else "authentic", "ratio": forgery_ratio}

if __name__ == "__main__":
    if IN_KAGGLE:
        DATA_DIR = Path("/kaggle/input/recodai-luc-scientific-image-forgery-detection")
        OUTPUT_DIR = Path("/kaggle/working")
        MODEL_PATH = "/kaggle/input/recod-luc-v8-seed2026/best_model.pth"
    else:
        DATA_DIR = Path("/mnt/wayne/competitions/recod_ai_luc/extracted")
        OUTPUT_DIR = Path("/mnt/wayne/competitions/recod_ai_luc/v8_output")
        MODEL_PATH = "/mnt/wayne/competitions/recod_ai_luc/experiments/ensemble/seed2026_waynepc/phase2_20260105_144802/best_model.pth"

    print(f"Data: {DATA_DIR} | Output: {OUTPUT_DIR} | Model: {MODEL_PATH}")
    OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
    
    test_dir = DATA_DIR / "test_images" if (DATA_DIR / "test_images").exists() else DATA_DIR / "test"
    test_images = sorted(list(test_dir.glob("*.png")) + list(test_dir.glob("*.jpg")))
    print(f"Test images: {len(test_images)}")
    
    pipeline = V8Pipeline(MODEL_PATH)
    predictions = []
    for img_path in tqdm(test_images, desc="V8 Inference"):
        try:
            result = pipeline.predict(img_path)
            predictions.append({"case_id": img_path.stem, "annotation": result["rle"]})
        except Exception as e:
            print(f"Error {img_path.stem}: {e}")
            predictions.append({"case_id": img_path.stem, "annotation": "authentic"})
    
    submission = pd.DataFrame(predictions)
    submission.to_csv(OUTPUT_DIR / "submission.csv", index=False)
    print(f"[SUCCESS] {len(submission)} predictions saved")
    print(f"Forged: {(submission['annotation'] != 'authentic').sum()} | Authentic: {(submission['annotation'] == 'authentic').sum()}")
    print("V8.0: seed2026 Dice 0.8179 | Target: 0.40+ F1 | WayneIA")
