#!/usr/bin/env python3
"""
Recod.ai/LUC Scientific Image Forgery Detection - WayneIA V9.0 OFFLINE
CODE_KEY[166] FORGERY_CLASSIFIER_MATRIX + CODE_KEY[158] FORGERY_DETECT
Competition: recodai-luc-scientific-image-forgery-detection | Prize: $55K | Deadline: Jan 15
V9.0: OFFLINE inference (no internet) | seed2026 Phase 2 (Dice 0.8179)
WayneIA Position_2 OpusPlan | January 6, 2026 | Year-8 RHINOCEROS G9
"""
import sys, os, warnings, shutil
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 V9.0 - WayneIA OFFLINE (No Internet)")
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

DINOV2_WEIGHTS_PATH = "/kaggle/input/dinov2-offline-vitb14/dinov2_vitb14_pretrain.pth"
CACHE_DIR = "/root/.cache/torch/hub/checkpoints"
os.makedirs(CACHE_DIR, exist_ok=True)

if os.path.exists(DINOV2_WEIGHTS_PATH):
    cache_path = os.path.join(CACHE_DIR, "dinov2_vitb14_pretrain.pth")
    if not os.path.exists(cache_path):
        print(f"Copying DINOv2 weights to cache...")
        shutil.copy(DINOV2_WEIGHTS_PATH, cache_path)
    print("[OK] DINOv2 weights cached")

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])
    rle_str = " ".join(map(str, rle_pairs))
    return f"[{rle_str}]" if rle_pairs else "authentic"

class Attention(nn.Module):
    def __init__(self, dim, num_heads=12, qkv_bias=True):
        super().__init__()
        self.num_heads = num_heads
        self.scale = (dim // num_heads) ** -0.5
        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
        self.proj = nn.Linear(dim, dim)
    
    def forward(self, x):
        B, N, C = x.shape
        qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
        q, k, v = qkv.unbind(0)
        attn = (q @ k.transpose(-2, -1)) * self.scale
        attn = attn.softmax(dim=-1)
        x = (attn @ v).transpose(1, 2).reshape(B, N, C)
        return self.proj(x)

class Mlp(nn.Module):
    def __init__(self, in_features, hidden_features=None):
        super().__init__()
        hidden_features = hidden_features or in_features * 4
        self.fc1 = nn.Linear(in_features, hidden_features)
        self.act = nn.GELU()
        self.fc2 = nn.Linear(hidden_features, in_features)
    
    def forward(self, x):
        return self.fc2(self.act(self.fc1(x)))

class Block(nn.Module):
    def __init__(self, dim, num_heads, mlp_ratio=4.0):
        super().__init__()
        self.norm1 = nn.LayerNorm(dim, eps=1e-6)
        self.attn = Attention(dim, num_heads)
        self.norm2 = nn.LayerNorm(dim, eps=1e-6)
        self.mlp = Mlp(dim, int(dim * mlp_ratio))
        self.ls1 = nn.Identity()
        self.ls2 = nn.Identity()
    
    def forward(self, x):
        x = x + self.ls1(self.attn(self.norm1(x)))
        x = x + self.ls2(self.mlp(self.norm2(x)))
        return x

class PatchEmbed(nn.Module):
    def __init__(self, img_size=518, patch_size=14, in_chans=3, embed_dim=768):
        super().__init__()
        self.img_size = img_size
        self.patch_size = patch_size
        self.num_patches = (img_size // patch_size) ** 2
        self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
    
    def forward(self, x):
        x = self.proj(x)
        x = x.flatten(2).transpose(1, 2)
        return x

class DinoVisionTransformer(nn.Module):
    def __init__(self, img_size=518, patch_size=14, in_chans=3, embed_dim=768, depth=12, num_heads=12, mlp_ratio=4.0):
        super().__init__()
        self.embed_dim = embed_dim
        self.patch_embed = PatchEmbed(img_size, patch_size, in_chans, embed_dim)
        num_patches = self.patch_embed.num_patches
        self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
        self.pos_embed = nn.Parameter(torch.zeros(1, num_patches + 1, embed_dim))
        self.blocks = nn.ModuleList([Block(embed_dim, num_heads, mlp_ratio) for _ in range(depth)])
        self.norm = nn.LayerNorm(embed_dim, eps=1e-6)
    
    def forward_features(self, x):
        B = x.shape[0]
        x = self.patch_embed(x)
        cls_tokens = self.cls_token.expand(B, -1, -1)
        x = torch.cat((cls_tokens, x), dim=1)
        x = x + self.pos_embed
        for blk in self.blocks:
            x = blk(x)
        x = self.norm(x)
        return {"x_norm_clstoken": x[:, 0], "x_norm_patchtokens": x[:, 1:]}

class DINOv2ForgeryDetector(nn.Module):
    def __init__(self, embed_dim=DINOV2_EMBED_DIM):
        super().__init__()
        print("Building DINOv2 ViT-B/14 backbone (offline)...")
        self.backbone = DinoVisionTransformer(
            img_size=IMG_SIZE, patch_size=PATCH_SIZE, embed_dim=embed_dim,
            depth=12, num_heads=12, mlp_ratio=4.0
        )
        weights_path = os.path.join(CACHE_DIR, "dinov2_vitb14_pretrain.pth")
        if os.path.exists(weights_path):
            print(f"Loading DINOv2 weights from: {weights_path}")
            state_dict = torch.load(weights_path, map_location="cpu")
            missing, unexpected = self.backbone.load_state_dict(state_dict, strict=False)
            print(f"[OK] DINOv2 loaded (missing: {len(missing)}, unexpected: {len(unexpected)})")
        else:
            print("[WARN] DINOv2 weights not found, using random init")
        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):
        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 V9Pipeline:
    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 V9 trained 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] V9 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/v9_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 = V9Pipeline(MODEL_PATH)
    predictions = []
    for img_path in tqdm(test_images, desc="V9 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")
    forged_count = (submission['annotation'] != 'authentic').sum()
    auth_count = (submission['annotation'] == 'authentic').sum()
    print(f"Forged: {forged_count} | Authentic: {auth_count}")
    print("V9.0 OFFLINE: seed2026 Dice 0.8179 | WayneIA")
