"""
================================================================================
   VESUVIUS V13 - ENHANCED

   Improvements over V12 (0.523):
   - Lower threshold (0.35 instead of hysteresis)
   - Connected component filtering (min_size=5000)
   - 8x TTA: 4 rotations × 2 flips
   Target: 0.55+
================================================================================
"""

import os
os.environ["KERAS_BACKEND"] = "tensorflow"
os.environ["TF_CPP_MIN_LOG_LEVEL"] = "2"

# ============================================================================
# SETUP PACKAGES
# ============================================================================
var = "/kaggle/input/vesuvius25-packages-offline-installer-v20251226/whls"
if os.path.exists(var):
    print(f"Installing packages from: {var}")
    import subprocess
    subprocess.run([
        "pip", "install", "--quiet",
        f"{var}/keras_nightly-3.12.0.dev2025100703-py3-none-any.whl",
        f"{var}/tifffile-2025.10.16-py3-none-any.whl",
        f"{var}/imagecodecs-2025.11.11-cp311-abi3-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl",
        f"{var}/medicai-0.0.3-py3-none-any.whl",
        "--no-index",
        "--find-links", var
    ], check=False, capture_output=True)
else:
    print(f"Package directory not found: {var}")

# Protobuf patch
try:
    from google.protobuf import message_factory as _message_factory
    if not hasattr(_message_factory.MessageFactory, "GetPrototype"):
        from google.protobuf.message_factory import GetMessageClass
        def _GetPrototype(self, descriptor):
            return GetMessageClass(descriptor)
        _message_factory.MessageFactory.GetPrototype = _GetPrototype
except:
    pass

import keras
from medicai.transforms import Compose, ScaleIntensityRange
from medicai.models import TransUNet
from medicai.utils.inference import SlidingWindowInference

import numpy as np
import pandas as pd
import zipfile
import tifffile
import scipy.ndimage as ndi
from scipy.ndimage import label, generate_binary_structure

print("="*60)
print("VESUVIUS V13 - ENHANCED")
print("="*60)
print(f"Keras backend: {keras.config.backend()}, version: {keras.version()}")

# ============================================================================
# CONFIG
# ============================================================================
root_dir = "/kaggle/input/vesuvius-challenge-surface-detection"
test_dir = f"{root_dir}/test_images"
output_dir = "/kaggle/working/submission_masks"
zip_path = "/kaggle/working/submission.zip"
os.makedirs(output_dir, exist_ok=True)

# Model config
NUM_CLASSES = 3
PATCH_SIZE = (160, 160, 160)
OVERLAP = 0.50

# Post-processing config - IMPROVED
THRESHOLD = 0.35  # Lower threshold (was hysteresis 0.45/0.85)
MIN_COMPONENT_SIZE = 5000  # Remove small noise (was 100)

# ============================================================================
# MODEL
# ============================================================================
def get_model():
    """Load TransUNet with fine-tuned weights."""
    model = TransUNet(
        input_shape=(160, 160, 160, 1),
        encoder_name='seresnext50',
        classifier_activation='softmax',
        num_classes=NUM_CLASSES,
    )

    # Try multiple model paths
    model_paths = [
        "/kaggle/input/train-transunet-baseline-lb-0-537/fine_tuning_epoch_20.weights.h5",
        "/kaggle/input/train-vesuvius-surface-3d-detection-on-tpu/model.weights.h5",
        "/kaggle/input/colab-a-162v4-gpu-transunet-seresnext101-x160/model.weights.h5"
    ]

    for path in model_paths:
        if os.path.exists(path):
            print(f"Loading weights from: {path}")
            model.load_weights(path)
            break

    print(f"Model params: {model.count_params() / 1e6:.1f}M")
    return model

# ============================================================================
# TRANSFORMS
# ============================================================================
def val_transformation(image):
    """Normalize image to [0, 1]."""
    data = {"image": image}
    pipeline = Compose([
        ScaleIntensityRange(
            keys=["image"],
            a_min=0,
            a_max=255,
            b_min=0,
            b_max=1,
            clip=True,
        ),
    ])
    result = pipeline(data)
    return result["image"]

# ============================================================================
# HELPERS
# ============================================================================
def load_volume(path):
    """Load 3D TIFF."""
    vol = tifffile.imread(path)
    vol = vol.astype(np.float32)
    vol = vol[None, ..., None]  # (1, D, H, W, 1)
    return vol

def rot90_volume(vol, k):
    """Rotate volume k times 90° clockwise in HW plane."""
    if vol.ndim == 5:
        return np.rot90(vol, k=-k, axes=(2, 3))
    else:
        return np.rot90(vol, k=-k, axes=(1, 2))

def unrot90_volume(vol, k):
    """Inverse rotation."""
    return rot90_volume(vol, (4 - k) % 4)

def flip_volume(vol, axis):
    """Flip volume along specified axis."""
    if vol.ndim == 5:
        return np.flip(vol, axis=axis)
    else:
        return np.flip(vol, axis=axis-1)

# ============================================================================
# POST-PROCESSING - IMPROVED
# ============================================================================
def connected_component_filter(mask, min_size=5000):
    """
    Remove small connected components (noise).
    """
    if mask.sum() == 0:
        return mask

    # 3D connectivity (6-connected)
    struct = generate_binary_structure(3, 1)
    labeled_array, num_features = label(mask, structure=struct)

    if num_features == 0:
        return mask

    # Get sizes of each component
    sizes = np.bincount(labeled_array.ravel())

    # Keep only components larger than min_size
    mask_sizes = sizes >= min_size
    mask_sizes[0] = 0  # Background is always 0

    # Rebuild mask
    filtered = mask_sizes[labeled_array].astype(np.uint8)

    return filtered

# ============================================================================
# INFERENCE WITH ENHANCED TTA
# ============================================================================
print("\nLoading model...")
model = get_model()

pred = SlidingWindowInference(
    model,
    roi_size=PATCH_SIZE,
    num_classes=NUM_CLASSES,
    mode="gaussian",
    overlap=OVERLAP,
    sw_batch_size=1
)

def predict_probs_tta_enhanced(sample):
    """
    8x TTA: 4 rotations × 2 (original + horizontal flip)
    Returns averaged probabilities.
    """
    probs_accum = []
    sample_np = np.asarray(sample)  # Convert tensor to numpy

    for flip in [False, True]:
        s = sample_np.copy()
        if flip:
            s = flip_volume(s, axis=3)  # Flip along W axis

        for k in range(4):
            s_rot = rot90_volume(s, k)
            out = pred(s_rot)
            out = np.asarray(out)

            # CLASS INDEX 1 for foreground
            probs = out[0, ..., 1]  # (D, H, W)

            probs = unrot90_volume(probs, k)
            if flip:
                probs = flip_volume(probs, axis=2)  # Unflip

            probs_accum.append(probs)

    return np.mean(probs_accum, axis=0)

def predict(sample):
    """Full prediction pipeline with enhanced TTA and post-processing."""
    # 8x TTA
    probs_fg = predict_probs_tta_enhanced(sample)

    # Simple thresholding
    mask = (probs_fg >= THRESHOLD).astype(np.uint8)
    initial = mask.sum()

    # Connected component filtering
    mask = connected_component_filter(mask, min_size=MIN_COMPONENT_SIZE)
    final = mask.sum()

    print(f"    Before CC filter: {initial:,}, After: {final:,} ({final/initial*100:.1f}% kept)")

    return mask

# ============================================================================
# MAIN INFERENCE
# ============================================================================
print("\n" + "="*60)
print("V13 ENHANCED INFERENCE")
print("="*60)
print(f"Threshold={THRESHOLD}")
print(f"8x TTA (4 rot × 2 flip)")
print(f"Min component size: {MIN_COMPONENT_SIZE}")

test_df = pd.read_csv(f"{root_dir}/test.csv")
print(f"\nTest samples: {len(test_df)}")

with zipfile.ZipFile(zip_path, "w", compression=zipfile.ZIP_DEFLATED) as z:
    for idx, row in test_df.iterrows():
        image_id = row["id"]
        tif_path = f"{test_dir}/{image_id}.tif"

        print(f"\n[{idx+1}/{len(test_df)}] Processing: {image_id}.tif")

        # Load and transform
        volume = load_volume(tif_path)
        print(f"    Shape: {volume.shape[1:-1]}")
        volume = val_transformation(volume)

        # Predict
        output = predict(volume)
        print(f"    Positives: {output.sum():,} voxels ({output.sum()/output.size*100:.2f}%)")

        # Save
        out_path = f"{output_dir}/{image_id}.tif"
        tifffile.imwrite(out_path, output.astype(np.uint8))

        z.write(out_path, arcname=f"{image_id}.tif")
        os.remove(out_path)

print("\n" + "="*60)
print("V13 ENHANCED COMPLETE!")
print(f"Submission: {zip_path}")
print("="*60)
