#!/usr/bin/env python3
"""Self-contained inference script for the RSNA Knee Abnormality Detection
Kaggle code competition.

Meant to run as the single cell of a Kaggle Notebook attached to:
  - the competition dataset (provides /kaggle/input/rsna-knee-abnormality-detection/
    with test.csv, test_series.csv, test_series/<study>/<series>/*.dcm -- the full
    hidden test set only exists when Kaggle actually executes this kernel)
  - a Kaggle Dataset containing best_model.pt (the trained MRNetStyleClassifier
    checkpoint from experiments/P1_alt_backbones/mrnet_style_finetune.py)

It reimplements, self-contained (no dependency on this project's other files,
since the Kaggle kernel sandbox cannot see our own storage):
  - per-plane series selection, copied from the exact sort key in
    rsna_knee_solution/prepare_rsna_knee.py's scan_study_manifest: prefer
    Fluid_Sensitive/Fat_Suppression series, then more slices (capped at 96)
  - per-series percentile windowing (p1-p99) + MONOCHROME1 inversion + 8-bit
    quantization, copied from export_dicom_png.py, so pixel statistics match
    what the PNGs the model was trained on actually looked like
  - deterministic 8-slices-per-plane sampling and the MRNetStyleClassifier
    architecture, copied from experiments/P1_alt_backbones/mrnet_style_finetune.py

A single corrupted/undecodable DICOM does not fail the whole study: it is
skipped and substituted with the previously decoded slice (same fallback used
in mae_knee_full_pipeline.py after we hit real corrupted files in this
dataset), and a study where every slice fails falls back to a mid-gray image
rather than crashing the whole submission.
"""

from __future__ import annotations

import math
import os
from pathlib import Path
from typing import Any

import numpy as np
import pandas as pd
import pydicom
import torch
import torch.nn as nn
import torchvision
from PIL import Image

# ---------------------------------------------------------------------------
# Paths -- adjust MODEL_PATH's dataset slug to whatever you name the Kaggle
# Dataset holding best_model.pt when attaching it to this kernel.
# ---------------------------------------------------------------------------
COMPETITION_DIR = Path(os.environ.get(
    "COMPETITION_DIR", "/kaggle/input/competitions/rsna-knee-abnormality-detection"
))
MODEL_PATH = Path(os.environ.get(
    "MODEL_PATH", "/kaggle/input/datasets/juejiezeng/rsna-knee-mrnet-resnet50/best_model.pt"
))
OUTPUT_PATH = Path(os.environ.get("OUTPUT_PATH", "submission.csv"))

LABELS = [
    "ACL", "MCL", "Medial Meniscus", "Lateral Meniscus", "Medial OA",
    "Lateral OA", "PF OA", "Effusion", "Synovitis", "Baker's", "Contusion", "Fracture",
]
PLANES = ("sagittal", "coronal", "axial")
IMAGENET_MEAN = (0.485, 0.456, 0.406)
IMAGENET_STD = (0.229, 0.224, 0.225)
IMAGE_SIZE = 224
SLICES_PER_PLANE = 8  # must match --slices-per-plane used to train best_model.pt
BATCH_STUDIES = 4

BACKBONES = {
    "resnet18": (torchvision.models.resnet18, 512, 4),
    "resnet34": (torchvision.models.resnet34, 512, 4),
    "resnet50": (torchvision.models.resnet50, 2048, 4),
}


# ---------------------------------------------------------------------------
# Series selection -- copied from rsna_knee_solution/prepare_rsna_knee.py
# scan_study_manifest's candidate sort (fluid_sensitive/fat_suppression first,
# then slice count capped at 96), series_per_plane=1.
# ---------------------------------------------------------------------------

def normalize_plane(value: Any) -> str:
    text = str(value).strip().lower()
    if text in PLANES:
        return text
    return "unknown"


def select_series_per_plane(series_df: pd.DataFrame) -> dict[str, dict[str, str]]:
    result: dict[str, dict[str, str]] = {}
    for study_id, group in series_df.groupby("StudyInstanceUID"):
        records = []
        for _, row in group.iterrows():
            records.append({
                "series_uid": str(row["SeriesInstanceUID"]),
                "plane": normalize_plane(row.get("Anatomical_Plane", "unknown")),
                "fluid_or_fat": bool(row.get("Fluid_Sensitive", 0)) or bool(row.get("Fat_Suppression", 0)),
            })
        chosen: dict[str, str] = {}
        for plane in PLANES:
            candidates = [r for r in records if r["plane"] == plane]
            if not candidates:
                continue
            candidates.sort(key=lambda r: int(r["fluid_or_fat"]), reverse=True)
            chosen[plane] = candidates[0]["series_uid"]
        result[str(study_id)] = chosen
    return result


# ---------------------------------------------------------------------------
# DICOM decode + per-series percentile windowing -- copied from
# export_dicom_png.py (spatial_position, percentile_sample, decode_slice,
# slice_sort_key, the p1-p99 window + MONOCHROME1 inversion + 8-bit quantize).
# ---------------------------------------------------------------------------

def as_float_list(value: Any) -> list[float] | None:
    if value is None:
        return None
    try:
        return [float(item) for item in value]
    except (TypeError, ValueError):
        return None


def spatial_position(dataset: Any) -> float | None:
    orientation = as_float_list(getattr(dataset, "ImageOrientationPatient", None))
    position = as_float_list(getattr(dataset, "ImagePositionPatient", None))
    if orientation and len(orientation) >= 6 and position and len(position) >= 3:
        row = np.asarray(orientation[:3], dtype=np.float64)
        column = np.asarray(orientation[3:6], dtype=np.float64)
        normal = np.cross(row, column)
        norm = float(np.linalg.norm(normal))
        if norm > 1e-8:
            normal /= norm
            return float(np.dot(np.asarray(position[:3], dtype=np.float64), normal))
    try:
        return float(dataset.SliceLocation)
    except (AttributeError, TypeError, ValueError):
        return None


def percentile_sample(pixels: np.ndarray, maximum: int = 65536) -> np.ndarray:
    values = pixels[np.isfinite(pixels)].reshape(-1)
    if values.size <= maximum:
        return values
    stride = int(math.ceil(values.size / maximum))
    return values[::stride]


def list_series_files(study_dir: Path, series_uid: str) -> list[Path]:
    series_dir = study_dir / series_uid
    if series_dir.is_dir():
        return sorted(path for path in series_dir.rglob("*.dcm") if path.is_file())
    return sorted(path for path in study_dir.rglob("*.dcm") if series_uid in str(path))


def decode_series_slices(study_dir: Path, series_uid: str) -> list[np.ndarray]:
    """Returns geometrically-sorted, [0,1]-normalized float32 slices for one series."""
    paths = list_series_files(study_dir, series_uid)
    if not paths:
        return []

    decoded = []
    for order, path in enumerate(paths):
        try:
            dataset = pydicom.dcmread(path, force=True)
            pixels = np.asarray(dataset.pixel_array).astype(np.float32, copy=False)
            slope = float(getattr(dataset, "RescaleSlope", 1.0) or 1.0)
            intercept = float(getattr(dataset, "RescaleIntercept", 0.0) or 0.0)
            if slope != 1.0 or intercept != 0.0:
                pixels = pixels * slope + intercept
            photometric = str(getattr(dataset, "PhotometricInterpretation", "MONOCHROME2")).upper()
            decoded.append((order, spatial_position(dataset), getattr(dataset, "InstanceNumber", None), pixels, photometric))
        except Exception as error:
            print(f"  [decode] skipping {path.name}: {error}")

    if not decoded:
        return []

    def sort_key(item):
        order, pos, instance, _, _ = item
        if pos is not None:
            return (0, pos, int(instance) if instance is not None else 0, order)
        if instance is not None:
            try:
                return (1, float(instance), int(instance), order)
            except (TypeError, ValueError):
                pass
        return (2, float(order), 0, order)

    decoded.sort(key=sort_key)
    samples = [percentile_sample(item[3]) for item in decoded]
    samples = [s for s in samples if s.size]
    if not samples:
        return []
    intensity = np.concatenate(samples)
    low, high = np.percentile(intensity, [1.0, 99.0]).astype(float)
    if not (np.isfinite(low) and np.isfinite(high)) or high <= low:
        low, high = float(intensity.min()), float(intensity.max())
        if high <= low:
            high = low + 1.0

    normalized_slices = []
    for _, _, _, pixels, photometric in decoded:
        normalized = np.clip((pixels - low) / (high - low), 0.0, 1.0)
        if photometric == "MONOCHROME1":
            normalized = 1.0 - normalized
        # match the 8-bit PNG quantization the model was trained on
        quantized = np.rint(normalized * 255.0).astype(np.uint8)
        normalized_slices.append(quantized.astype(np.float32) / 255.0)
    return normalized_slices


# ---------------------------------------------------------------------------
# Slice sampling -- copied from experiments/P1_alt_backbones/mrnet_style_finetune.py
# ---------------------------------------------------------------------------

def deterministic_bin_centers(count: int, samples: int) -> np.ndarray:
    if count <= 1:
        return np.zeros(samples, dtype=np.int64)
    edges = np.linspace(0, count, samples + 1)
    positions = (edges[:-1] + edges[1:]) / 2.0
    return np.clip(np.floor(positions).astype(np.int64), 0, count - 1)


def slices_to_tensor(slices: list[np.ndarray]) -> torch.Tensor:
    """One [0,1] float32 slice -> normalized 3xHxW tensor, resized to IMAGE_SIZE."""
    tensors = []
    normalize = torchvision.transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD)
    for array in slices:
        image = Image.fromarray(np.rint(array * 255.0).astype(np.uint8), mode="L")
        if image.size != (IMAGE_SIZE, IMAGE_SIZE):
            image = image.resize((IMAGE_SIZE, IMAGE_SIZE), Image.Resampling.BILINEAR)
        tensor = torch.from_numpy(np.asarray(image, dtype=np.float32) / 255.0).unsqueeze(0).repeat(3, 1, 1)
        tensors.append(normalize(tensor))
    return torch.stack(tensors, dim=0)


def build_study_tensor(study_dir: Path, chosen_series: dict[str, str]) -> torch.Tensor | None:
    per_plane_tensors = []
    for plane in PLANES:
        series_uid = chosen_series.get(plane)
        normalized_slices = decode_series_slices(study_dir, series_uid) if series_uid else []
        if not normalized_slices:
            # No usable series for this plane: fall back to mid-gray placeholders
            # so the study still produces a prediction instead of being dropped.
            placeholder = np.full((IMAGE_SIZE, IMAGE_SIZE), 0.5, dtype=np.float32)
            selected = [placeholder] * SLICES_PER_PLANE
        else:
            indices = deterministic_bin_centers(len(normalized_slices), SLICES_PER_PLANE)
            selected = [normalized_slices[int(i)] for i in indices]
        per_plane_tensors.append(slices_to_tensor(selected))
    return torch.cat(per_plane_tensors, dim=0)  # (3*SLICES_PER_PLANE, 3, H, W)


# ---------------------------------------------------------------------------
# Model -- copied from experiments/P1_alt_backbones/mrnet_style_finetune.py
# ---------------------------------------------------------------------------

class MRNetStyleClassifier(nn.Module):
    def __init__(self, backbone_name, pretrained, backbone_train_mode, train_last_n_blocks,
                 slices_per_plane, num_labels, dropout):
        super().__init__()
        constructor, feature_dim, num_stages = BACKBONES[backbone_name]
        backbone = constructor(weights=None)
        self.feature_dim = feature_dim
        self.slices_per_plane = slices_per_plane
        self.stem = nn.Sequential(backbone.conv1, backbone.bn1, backbone.relu, backbone.maxpool)
        self.stages = nn.ModuleList([backbone.layer1, backbone.layer2, backbone.layer3, backbone.layer4])
        self.pool = nn.AdaptiveAvgPool2d((1, 1))
        attention_dim = max(feature_dim // 4, 128)
        self.plane_attention = nn.Sequential(nn.Linear(feature_dim, attention_dim), nn.Tanh(), nn.Linear(attention_dim, 1))
        self.norm = nn.LayerNorm(feature_dim * len(PLANES))
        self.dropout = nn.Dropout(dropout)
        self.head = nn.Linear(feature_dim * len(PLANES), num_labels)

    def encode_slices(self, pixel_values: torch.Tensor) -> torch.Tensor:
        batch_size, total_slices, channels, height, width = pixel_values.shape
        flat = pixel_values.reshape(batch_size * total_slices, channels, height, width)
        hidden = self.stem(flat)
        for stage in self.stages:
            hidden = stage(hidden)
        pooled = self.pool(hidden).flatten(1)
        return pooled.reshape(batch_size, total_slices, self.feature_dim)

    def forward(self, pixel_values: torch.Tensor) -> torch.Tensor:
        features = self.encode_slices(pixel_values)
        plane_vectors = []
        for plane_index in range(len(PLANES)):
            start = plane_index * self.slices_per_plane
            end = start + self.slices_per_plane
            plane_features = features[:, start:end, :]
            attention = torch.softmax(self.plane_attention(plane_features), dim=1)
            plane_vectors.append((plane_features * attention).sum(dim=1))
        combined = torch.cat(plane_vectors, dim=-1)
        return self.head(self.dropout(self.norm(combined)))


def resolve_model_path(preferred: Path) -> Path:
    if preferred.exists():
        return preferred
    print(f"[warn] {preferred} not found; listing /kaggle/input top level (non-recursive)", flush=True)
    root = Path("/kaggle/input")
    top_level = sorted(root.iterdir()) if root.is_dir() else []
    for entry in top_level:
        print(f"  {entry}", flush=True)
    # Only check a few fixed depths -- never rglob the whole tree, since the
    # competition's own mounted data can be a large, deeply-nested DICOM directory
    # and a full recursive walk over it can take a very long time. Kaggle nests
    # dataset mounts as /kaggle/input/datasets/<user>/<slug>/..., hence depth 4.
    candidates = (
        sorted(root.glob("*/best_model.pt"))
        + sorted(root.glob("*/*/best_model.pt"))
        + sorted(root.glob("*/*/*/best_model.pt"))
        + sorted(root.glob("*/*/*/*/best_model.pt"))
    )
    if not candidates:
        raise FileNotFoundError(f"No best_model.pt found under {root} (1-2 levels deep); top level was: {top_level}")
    print(f"[warn] using {candidates[0]} instead", flush=True)
    return candidates[0]


def main() -> None:
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    print(f"device={device}")
    global MODEL_PATH
    MODEL_PATH = resolve_model_path(MODEL_PATH)

    checkpoint = torch.load(MODEL_PATH, map_location="cpu", weights_only=False)
    if checkpoint.get("labels") != LABELS:
        raise ValueError(f"Checkpoint label order mismatch: {checkpoint.get('labels')}")
    model = MRNetStyleClassifier(**checkpoint["model_config"]).to(device)
    model.load_state_dict(checkpoint["model_state_dict"], strict=True)
    model.eval()
    print("model_config:", checkpoint["model_config"])

    competition_dir = COMPETITION_DIR
    if not (competition_dir / "test.csv").exists():
        root = Path("/kaggle/input")
        hits = (
            sorted(root.glob("*/test.csv"))
            + sorted(root.glob("*/*/test.csv"))
            + sorted(root.glob("*/*/*/test.csv"))
        ) if root.is_dir() else []
        if not hits:
            raise FileNotFoundError(f"No test.csv found under {root} (up to 3 levels deep)")
        competition_dir = hits[0].parent
        print(f"[warn] {COMPETITION_DIR}/test.csv not found; using {competition_dir} instead", flush=True)

    test_df = pd.read_csv(competition_dir / "test.csv", dtype=str)
    series_df = pd.read_csv(competition_dir / "test_series.csv", dtype=str)
    series_root = competition_dir / "test_series"
    chosen_by_study = select_series_per_plane(series_df)

    study_ids = test_df["StudyInstanceUID"].astype(str).tolist()
    rows = []
    with torch.inference_mode():
        for start in range(0, len(study_ids), BATCH_STUDIES):
            batch_ids = study_ids[start:start + BATCH_STUDIES]
            batch_tensors = []
            for study_id in batch_ids:
                chosen = chosen_by_study.get(study_id, {})
                tensor = build_study_tensor(series_root / study_id, chosen)
                batch_tensors.append(tensor)
            pixel_values = torch.stack(batch_tensors, dim=0).to(device)
            logits = model(pixel_values)
            probabilities = torch.sigmoid(logits).float().cpu().numpy()
            for study_id, row_probs in zip(batch_ids, probabilities):
                rows.append({"StudyInstanceUID": study_id, **{label: float(p) for label, p in zip(LABELS, row_probs)}})
            print(f"scored {min(start + BATCH_STUDIES, len(study_ids))}/{len(study_ids)}")

    submission = pd.DataFrame(rows, columns=["StudyInstanceUID"] + LABELS)
    submission.to_csv(OUTPUT_PATH, index=False)
    print(f"Wrote {OUTPUT_PATH.resolve()} with {len(submission)} rows")
    print(submission.head())


if __name__ == "__main__":
    main()
