#!/usr/bin/env python3
"""
Kaggle Submission Script - Scientific Image Forgery Detection
Model: Trial 14 DINOv2 with optimized post-processing
"""

import os
import glob
import math
import json
import zipfile
import numpy as np
import pandas as pd
from PIL import Image
import cv2
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers import AutoImageProcessor, AutoModel
from tqdm import tqdm

# Config
DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
IMG_SIZE = 518

# Optimized hyperparameters from Optuna
ALPHA_GRAD = 0.10157637415630942
THRESHOLD_MULT = 0.24528670131752023
MORPH_CLOSE = 7
MORPH_OPEN = 7
MIN_AREA = 387
MIN_CONF = 0.29778295633223595
MAX_AREA_RATIO = 0.7160231039998628

print(f"Device: {DEVICE}")


class ConvBlock(nn.Module):
    def __init__(self, in_ch, out_ch, dropout=0.0):
        super().__init__()
        self.conv = nn.Sequential(
            nn.Conv2d(in_ch, out_ch, 3, padding=1),
            nn.BatchNorm2d(out_ch),
            nn.ReLU(inplace=True),
            nn.Dropout2d(dropout) if dropout > 0 else nn.Identity(),
            nn.Conv2d(out_ch, out_ch, 3, padding=1),
            nn.BatchNorm2d(out_ch),
            nn.ReLU(inplace=True),
        )

    def forward(self, x):
        return self.conv(x)


class AttentionBlock(nn.Module):
    def __init__(self, channels):
        super().__init__()
        self.avg_pool = nn.AdaptiveAvgPool2d(1)
        self.fc = nn.Sequential(
            nn.Linear(channels, channels // 8),
            nn.ReLU(inplace=True),
            nn.Linear(channels // 8, channels),
            nn.Sigmoid()
        )

    def forward(self, x):
        b, c, _, _ = x.size()
        y = self.avg_pool(x).view(b, c)
        y = self.fc(y).view(b, c, 1, 1)
        return x * y


class FlexibleDecoder(nn.Module):
    def __init__(self, in_channels=768, base_channels=64, use_attention=True, n_layers=4, dropout=0.218):
        super().__init__()
        self.use_attention = use_attention
        self.n_layers = n_layers

        channels = [in_channels]
        for i in range(n_layers):
            channels.append(base_channels * (2 ** (n_layers - 1 - i)))

        self.ups = nn.ModuleList()
        self.convs = nn.ModuleList()
        self.attns = nn.ModuleList() if use_attention else None

        for i in range(n_layers):
            self.ups.append(nn.ConvTranspose2d(channels[i], channels[i+1], 2, stride=2))
            self.convs.append(ConvBlock(channels[i+1], channels[i+1], dropout))
            if use_attention:
                self.attns.append(AttentionBlock(channels[i+1]))

        self.final = nn.Sequential(
            nn.Conv2d(channels[-1], channels[-1] // 2, 3, padding=1),
            nn.ReLU(inplace=True),
            nn.Conv2d(channels[-1] // 2, 1, 1)
        )

    def forward(self, features, target_size):
        x = features
        for i in range(self.n_layers):
            x = self.ups[i](x)
            if self.use_attention:
                x = self.attns[i](x)
            x = self.convs[i](x)
        x = F.interpolate(x, size=target_size, mode='bilinear', align_corners=False)
        x = self.final(x)
        return x


class DinoV2Segmenter(nn.Module):
    def __init__(self, dino_model, decoder):
        super().__init__()
        self.dino = dino_model
        self.decoder = decoder

    def extract_features(self, x):
        with torch.no_grad():
            outputs = self.dino(x)
            features = outputs.last_hidden_state[:, 1:, :]
        B, N, C = features.shape
        h = w = int(math.sqrt(N))
        features = features.permute(0, 2, 1).reshape(B, C, h, w)
        return features

    def forward(self, x, target_size=None):
        H, W = x.shape[2], x.shape[3]
        if target_size is None:
            target_size = (H, W)
        features = self.extract_features(x)
        logits = self.decoder(features, target_size)
        return logits


def get_model_dir():
    """Get model directory, extracting ZIP if needed"""
    possible_dirs = [
        '/kaggle/input/forgery-dinov2-model/dinov2-base',
        '/kaggle/input/forgery-dinov2-model/dinov2-base/dinov2-base',
        '/kaggle/working/dinov2-base',
        '/kaggle/working/dinov2-base/dinov2-base',
    ]
    zip_path = '/kaggle/input/forgery-dinov2-model/dinov2-base.zip'

    for model_dir in possible_dirs:
        config_file = os.path.join(model_dir, 'config.json')
        if os.path.exists(config_file):
            print(f"Found model at {model_dir}")
            return model_dir

    if os.path.exists(zip_path):
        print(f"Extracting {zip_path}...")
        with zipfile.ZipFile(zip_path, 'r') as zip_ref:
            zip_ref.extractall('/kaggle/working')
        for model_dir in possible_dirs:
            config_file = os.path.join(model_dir, 'config.json')
            if os.path.exists(config_file):
                print(f"Found model at {model_dir}")
                return model_dir

    raise FileNotFoundError("No model directory found")


def find_test_images():
    comp_dir = '/kaggle/input/recodai-luc-scientific-image-forgery-detection'
    for path in [os.path.join(comp_dir, 'test_images'), os.path.join(comp_dir, 'test'), comp_dir]:
        if os.path.exists(path):
            images = glob.glob(os.path.join(path, '*.png'))
            if images:
                print(f"Found {len(images)} images in {path}")
                return sorted(images)
    images = glob.glob(os.path.join(comp_dir, '**', '*.png'), recursive=True)
    if images:
        return sorted(images)
    raise FileNotFoundError("No test images found")


def rle_encode(mask):
    pixels = mask.T.flatten()
    dots = np.where(pixels == 1)[0]
    if len(dots) == 0:
        return "authentic"
    run_lengths = []
    prev = -2
    for d in dots:
        if d > prev + 1:
            run_lengths.extend((d + 1, 0))
        run_lengths[-1] += 1
        prev = d
    return json.dumps([int(x) for x in run_lengths])


def postprocess_mask(preds):
    gx = cv2.Sobel(preds.astype(np.float32), cv2.CV_32F, 1, 0, ksize=3)
    gy = cv2.Sobel(preds.astype(np.float32), cv2.CV_32F, 0, 1, ksize=3)
    grad_mag = np.sqrt(gx**2 + gy**2)
    grad_max = grad_mag.max()
    grad_norm = grad_mag / grad_max if grad_max > 0 else grad_mag

    enhanced = (1 - ALPHA_GRAD) * preds + ALPHA_GRAD * grad_norm
    enhanced = cv2.GaussianBlur(enhanced, (5, 5), 0)

    thr = np.mean(enhanced) + THRESHOLD_MULT * np.std(enhanced)
    mask = (enhanced > thr).astype(np.uint8)

    if MORPH_CLOSE > 1:
        mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, np.ones((MORPH_CLOSE, MORPH_CLOSE), np.uint8))
    if MORPH_OPEN > 1:
        mask = cv2.morphologyEx(mask, cv2.MORPH_OPEN, np.ones((MORPH_OPEN, MORPH_OPEN), np.uint8))
    mask = cv2.morphologyEx(mask, cv2.MORPH_CLOSE, np.ones((11, 11), np.uint8))

    return mask, preds


def predict():
    print("Loading model (Trial 14 Extended - IoU=0.5039)...")

    model_dir = get_model_dir()
    model_input = '/kaggle/input/forgery-dinov2-model'
    model_path = os.path.join(model_input, 'trial14_extended_best.pth')

    processor = AutoImageProcessor.from_pretrained(model_dir, local_files_only=True)
    dino_model = AutoModel.from_pretrained(model_dir, local_files_only=True)
    dino_model = dino_model.to(DEVICE)

    # Load combined checkpoint
    print(f"Loading checkpoint from {model_path}")
    checkpoint = torch.load(model_path, map_location=DEVICE, weights_only=True)
    print(f"Checkpoint IoU: {checkpoint.get('val_iou', 'N/A')}, Epoch: {checkpoint.get('epoch', 'N/A')}")

    # Load backbone
    dino_model.load_state_dict(checkpoint['backbone_state'], strict=False)
    dino_model.eval()

    decoder = FlexibleDecoder(
        in_channels=768,
        base_channels=64,
        use_attention=True,
        n_layers=4,
        dropout=0.218
    )

    # Load decoder
    decoder.load_state_dict(checkpoint['decoder_state'])

    decoder = decoder.to(DEVICE)
    decoder.eval()

    model = DinoV2Segmenter(dino_model, decoder)
    model.eval()

    print("Model loaded successfully!")

    test_images = find_test_images()
    print(f"Processing {len(test_images)} test images")

    results = []
    for img_path in tqdm(test_images):
        img_name = os.path.basename(img_path)
        case_id = int(os.path.splitext(img_name)[0])

        image = Image.open(img_path)
        orig_size = image.size
        if image.mode != 'RGB':
            image = image.convert('RGB')

        inputs = processor(images=image, return_tensors="pt")
        pixel_values = inputs['pixel_values'].to(DEVICE)

        with torch.no_grad():
            outputs = model(pixel_values, target_size=(IMG_SIZE, IMG_SIZE))
            pred = torch.sigmoid(outputs).squeeze().cpu().numpy()

        mask, preds_raw = postprocess_mask(pred)

        area = int(mask.sum())
        total_pixels = IMG_SIZE * IMG_SIZE

        if area > 0:
            mean_conf = float(preds_raw[mask == 1].mean())
        else:
            mean_conf = 0.0

        is_forged = (area >= MIN_AREA and mean_conf >= MIN_CONF and area <= total_pixels * MAX_AREA_RATIO)

        if is_forged:
            mask_orig = cv2.resize(mask, orig_size, interpolation=cv2.INTER_NEAREST)
            annotation = rle_encode(mask_orig)
        else:
            annotation = "authentic"

        results.append({'case_id': case_id, 'annotation': annotation})

    submission = pd.DataFrame(results)
    submission = submission.sort_values('case_id').reset_index(drop=True)
    submission.to_csv('submission.csv', index=False)

    n_authentic = sum(1 for r in results if r['annotation'] == 'authentic')
    n_forged = len(results) - n_authentic
    print(f"\nSubmission: {n_authentic} authentic, {n_forged} forged ({n_forged/len(results)*100:.1f}%)")
    print(submission.head(10))

    return submission


if __name__ == '__main__':
    predict()
