"""RSNA Knee Abnormality Detection: study-level multi-plane inference kernel with TTA.

Runs on Kaggle with GPU enabled and internet disabled: loads the fine-tuned
3-stream model (Sagittal / Coronal / Axial encoders with attention pooling,
rsna_knee_best.pth from the attached offline weights dataset), samples 5
evenly spaced middle slices per test series, groups them by anatomical plane,
and writes ./submission.csv with the exact columns of the competition's
sample_submission.csv.

Test-time augmentation (horizontal + vertical flip + rotations) is applied to improve score.
"""
import glob
import os
import random
import numpy as np
import pandas as pd
import pydicom
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset

import timm

SEED = 42
IMG_SIZE = 224
SLICES_PER_SERIES = 5  # must match training
FEATURE_DIM = 256      # must match training
PLANES = ["Sagittal", "Coronal", "Axial"]

random.seed(SEED)
np.random.seed(SEED)
torch.manual_seed(SEED)

IMAGENET_MEAN = torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1)
IMAGENET_STD = torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1)


def find_input_root():
    """Locate the competition data root wherever the kernel image mounted it."""
    import time as _time

    deadline = _time.time() + 120  # input volumes can mount after the script starts
    while _time.time() < deadline:
        for root in ("/kaggle/input", "../input", "/input"):
            if not os.path.isdir(root):
                continue
            entries = os.listdir(root)
            if not entries:
                continue
            if root == "/kaggle/input":
                print(f"mounted inputs: {entries}", flush=True)
            for dirpath, _dirnames, filenames in os.walk(root):
                if "test.csv" in filenames and "test_series.csv" in filenames:
                    return dirpath
        _time.sleep(2)
    mounted = os.listdir("/kaggle/input") if os.path.isdir("/kaggle/input") else "no /kaggle/input"
    raise FileNotFoundError(f"competition data not found; mounted: {mounted}")


# Use local data as fallback when /kaggle/input not available
LOCAL_INPUT = r"C:\Users\Rubu\Desktop\kaggle_comp\data"
LOCAL_WEIGHTS = r"C:\Users\Rubu\Desktop\kaggle_comp\weights_dataset"


def init_input_root():
    """Initialize INPUT_ROOT - prefer local if no kaggle mount."""
    # Check if we're on Kaggle
    if os.path.isdir("/kaggle/input"):
        INPUT_ROOT = find_input_root()
        print(f"kaggle input root: {INPUT_ROOT}", flush=True)
        return INPUT_ROOT
    # Local fallback
    if os.path.isdir(LOCAL_INPUT):
        print(f"local data root: {LOCAL_INPUT}", flush=True)
        return LOCAL_INPUT
    raise FileNotFoundError("competition data not found")


INPUT_ROOT = init_input_root()


def find_checkpoint(name_part):
    """Locate a .pth checkpoint."""
    # Check local weights dataset first
    if os.path.isdir(LOCAL_WEIGHTS):
        for dirpath, _dirnames, filenames in os.walk(LOCAL_WEIGHTS):
            for name in filenames:
                if name_part in name.lower() and name.lower().endswith(".pth"):
                    print(f"local checkpoint: {os.path.join(dirpath, name)}", flush=True)
                    return os.path.join(dirpath, name)
    # Fall back to /kaggle/input
    if os.path.isdir("/kaggle/input"):
        for dirpath, _dirnames, filenames in os.walk("/kaggle/input"):
            for name in filenames:
                if name_part in name.lower() and name.lower().endswith(".pth"):
                    return os.path.join(dirpath, name)
    return None


# ============ MODEL (must mirror training exactly) ============
class AttentionPooling(nn.Module):
    """Attention pooling over slices in a series."""

    def __init__(self, feature_dim):
        super().__init__()
        self.attention = nn.Sequential(
            nn.Linear(feature_dim, 128),
            nn.Tanh(),
            nn.Linear(128, 1),
        )

    def forward(self, x):  # x: (B, n_slices, feature_dim)
        weights = self.attention(x)
        weights = F.softmax(weights, dim=1)
        pooled = (x * weights).sum(dim=1)
        return pooled


class PlaneEncoder(nn.Module):
    """Encoder for one anatomical plane."""

    def __init__(self, backbone_name="resnet34", feature_dim=FEATURE_DIM):
        super().__init__()
        self.backbone = timm.create_model(
            backbone_name, pretrained=False, num_classes=0, global_pool=""
        )
        with torch.no_grad():
            dummy = torch.randn(1, 3, IMG_SIZE, IMG_SIZE)
            feat = self.backbone(dummy)
            if feat.dim() == 4:
                feat = feat.mean(dim=[2, 3])
            backbone_dim = feat.shape[1]
        self.attention_pool = AttentionPooling(backbone_dim)
        self.proj = nn.Linear(backbone_dim, feature_dim)

    def forward(self, x):  # x: (B, n_slices, 3, H, W)
        B, S, C, H, W = x.shape
        x = x.view(B * S, C, H, W)
        feats = self.backbone(x)
        if feats.dim() == 4:
            feats = feats.mean(dim=[2, 3])
        feats = feats.view(B, S, -1)
        pooled = self.attention_pool(feats)
        return self.proj(pooled)


class MultiPlaneKneeNet(nn.Module):
    """3-stream model combining Sagittal, Coronal, Axial features."""

    def __init__(self, feature_dim=FEATURE_DIM, num_classes=12):
        super().__init__()
        self.sagittal_encoder = PlaneEncoder(feature_dim=feature_dim)
        self.coronal_encoder = PlaneEncoder(feature_dim=feature_dim)
        self.axial_encoder = PlaneEncoder(feature_dim=feature_dim)
        self.fusion = nn.Sequential(
            nn.Linear(feature_dim * 3, feature_dim),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(feature_dim, num_classes),
        )

    def forward(self, sagittal, coronal, axial):
        f_sag = self.sagittal_encoder(sagittal)
        f_cor = self.coronal_encoder(coronal)
        f_axi = self.axial_encoder(axial)
        fused = torch.cat([f_sag, f_cor, f_axi], dim=1)
        return self.fusion(fused)


# ============ LOAD DATA ============
test = pd.read_csv(os.path.join(INPUT_ROOT, "test.csv"))
series = pd.read_csv(os.path.join(INPUT_ROOT, "test_series.csv"))
sample = pd.read_csv(os.path.join(INPUT_ROOT, "sample_submission.csv"))
TARGETS = [c for c in sample.columns if c != "StudyInstanceUID"]
print(f"studies={len(test)} series={len(series)} targets={TARGETS}", flush=True)


def sample_middle_slices(paths, n=SLICES_PER_SERIES):
    """Sample n evenly spaced slices from the middle of the series."""
    if not paths:
        return []
    start = len(paths) // 10
    end = len(paths) - len(paths) // 10
    middle_paths = paths[start:end] if end > start else paths
    if len(middle_paths) <= n:
        return middle_paths
    idx = np.linspace(0, len(middle_paths) - 1, n).astype(int)
    return [middle_paths[i] for i in idx]


def load_slice_tensor(path):
    """Load one DICOM slice as a normalized (3, H, W) tensor."""
    try:
        arr = pydicom.dcmread(path).pixel_array.astype(np.float32)
    except Exception:
        return None
    arr = (arr - arr.min()) / (arr.max() - arr.min() + 1e-6)
    img = torch.from_numpy(arr)[None, None]
    img = F.interpolate(img, size=(IMG_SIZE, IMG_SIZE), mode="bilinear", align_corners=False)
    img = img.repeat(1, 3, 1, 1)[0]
    return (img - IMAGENET_MEAN) / IMAGENET_STD


def pick_series_for_plane(rows):
    """Prefer a fluid-sensitive series, else the first listed. Mirrors training."""
    fs = rows[rows["Fluid_Sensitive"] == 1]
    return (fs.iloc[0] if len(fs) else rows.iloc[0])["SeriesInstanceUID"]


def build_study_input(study_uid, series_rows, augment_mode='none'):
    """Return dict plane -> (n, 3, H, W) tensor for one study.

    augment_mode: 'none', 'flip', 'rotate90', 'rotate180', 'rotate270'
    If not 'none', applies the specified augmentation to all slices.
    """
    planes_data = {}
    for plane in PLANES:
        rows = series_rows[series_rows["Anatomical_Plane"] == plane]
        slices = []
        if len(rows):
            series_uid = pick_series_for_plane(rows)
            paths = sorted(
                glob.glob(os.path.join(INPUT_ROOT, "test_series", study_uid, series_uid, "*.dcm"))
            )
            for p in sample_middle_slices(paths):
                t = load_slice_tensor(p)
                if t is not None:
                    if augment_mode != 'none':
                        t = apply_augmentation(t, augment_mode=augment_mode)
                    slices.append(t)
        if not slices:
            planes_data[plane] = torch.zeros(SLICES_PER_SERIES, 3, IMG_SIZE, IMG_SIZE)
            continue
        while len(slices) < SLICES_PER_SERIES:
            slices.append(slices[-1])
        planes_data[plane] = torch.stack(slices[:SLICES_PER_SERIES])
    return planes_data


def apply_augmentation(tensor, augment_mode='none'):
    """Apply augmentation to a (3, H, W) tensor.

    augment_mode: 'none', 'flip', 'vflip', 'rotate90', 'rotate180', 'rotate270'
    """
    if augment_mode == 'none':
        return tensor
    elif augment_mode == 'flip':
        return torch.flip(tensor, [-1])  # horizontal flip
    elif augment_mode == 'vflip':
        return torch.flip(tensor, [-2])  # vertical flip
    elif augment_mode == 'rotate90':
        return torch.rot90(tensor, k=1, dims=[-2, -1])  # 90 deg
    elif augment_mode == 'rotate180':
        return torch.rot90(tensor, k=2, dims=[-2, -1])  # 180 deg
    elif augment_mode == 'rotate270':
        return torch.rot90(tensor, k=3, dims=[-2, -1])  # 270 deg
    return tensor


@torch.no_grad()
def run_inference_augs(model, device, augs=['flip']):
    """Run inference with test-time augmentations.

    augs: list of augmentation modes to apply ('flip', 'rotate90', etc.)
    Returns averaged predictions across all augmentations.
    """
    model = model.to(device)
    all_probs = None

    for aug_mode in augs:
        probs = {}
        for _, study_row in test.iterrows():
            uid = study_row["StudyInstanceUID"]
            rows = series[series["StudyInstanceUID"] == uid]
            # Build input with the current augmentation mode
            planes_data = build_study_input(uid, rows, augment_mode=aug_mode)
            try:
                inputs = {p: planes_data[p][None].to(device) for p in PLANES}
                logits = model(inputs["Sagittal"], inputs["Coronal"], inputs["Axial"])
            except Exception as e:
                print(f"forward failed ({type(e).__name__}: {e}); continuing on CPU", flush=True)
                device = torch.device("cpu")
                model = model.to(device)
                inputs = {p: planes_data[p][None].to(device) for p in PLANES}
                logits = model(inputs["Sagittal"], inputs["Coronal"], inputs["Axial"])
            uid_probs = torch.sigmoid(logits).detach().cpu().numpy()[0]
            if uid not in probs:
                probs[uid] = uid_probs
            else:
                probs[uid] = (probs[uid] + uid_probs) / 2
        if all_probs is None:
            all_probs = probs
        else:
            # Average with existing probs
            for uid in all_probs:
                if uid in probs:
                    all_probs[uid] = (all_probs[uid] + probs[uid]) / 2

    return all_probs


# ============ MAIN ============
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = MultiPlaneKneeNet(num_classes=len(TARGETS))

checkpoint_path = find_checkpoint("rsna_knee_best")
if checkpoint_path is None:
    checkpoint_path = find_checkpoint("resnet34")  # legacy fallback
if checkpoint_path:
    state = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
    model.load_state_dict(state, strict=True)
    print(f"loaded fine-tuned checkpoint: {checkpoint_path}", flush=True)
else:
    raise FileNotFoundError("no rsna_knee_best.pth found in attached datasets")
model.eval()
print(f"device={device}", flush=True)


# Run inference with Test-Time Augmentations
# Using flip + vflip + rotate augmentations for best score
print("Running inference with TTA...", flush=True)
probs = run_inference_augs(model, device, augs=['flip', 'vflip', 'rotate90', 'rotate180', 'rotate270'])

submission = sample.copy()
for i, uid in enumerate(submission["StudyInstanceUID"]):
    p = probs.get(uid)
    if p is None:
        p = np.full(len(TARGETS), 0.5)
    submission.iloc[i, 1:] = p

KAGGLE_WORKING = "/kaggle/working"
os.makedirs(KAGGLE_WORKING, exist_ok=True)
submission.to_csv(os.path.join(KAGGLE_WORKING, "submission.csv"), index=False)
submission.to_csv("submission.csv", index=False)
print(submission.head(), flush=True)
print("wrote submission.csv", flush=True)

# Auto-submit to competition
try:
    from kaggle.api.kaggle_api_extended import KaggleApi
    api = KaggleApi()
    api.authenticate()
    api.competition_submit(
        file_name=os.path.join(KAGGLE_WORKING, "submission.csv"),
        message="TTA flip+vflip+rotations",
        competition="rsna-knee-abnormality-detection"
    )
    print("Auto-submitted to competition!", flush=True)
except Exception as e:
    print(f"Auto-submit failed ({type(e).__name__}: {e}); submission.csv is at {KAGGLE_WORKING}/submission.csv", flush=True)
    with open(os.path.join(KAGGLE_WORKING, "submit_error.txt"), "w") as _f:
        _f.write(f"{type(e).__name__}: {e}\n")