# RSNA Knee — TEST INFERENCE v2 (trained EfficientNet-B3 → submission.csv)
# Loads rsna_model_b3.pt (state_dict from the image-v2 train kernel, val AUC
# 0.6824), runs the 3 public test studies through the SAME preprocessing
# pipeline, and writes submission.csv. CPU-only (3 studies × 12 slices = trivial).
import os, subprocess, sys, glob
import numpy as np
import pandas as pd
import torch
import torch.nn as nn
import pydicom
import cv2

for _pkg in ("timm", "pydicom", "opencv-python-headless"):
    try:
        __import__(_pkg.replace("-", "_"))
    except ImportError:
        subprocess.check_call([sys.executable, "-m", "pip", "install", "-q", _pkg])

LABELS = ['ACL','MCL','Medial Meniscus','Lateral Meniscus','Medial OA','Lateral OA',
          'PF OA','Effusion','Synovitis',"Baker's",'Contusion','Fracture']
N_SLICES = 12
IMG_SIZE = 224

# ── auto-detect competition data dir ──
print("Mounted inputs:", os.listdir("/kaggle/input"))
_candidates = glob.glob("/kaggle/input/**/test_series.csv", recursive=True)
assert _candidates, "test_series.csv not found — add competition input"
DATA_DIR = os.path.dirname(_candidates[0])
SERIES_DIR = os.path.join(DATA_DIR, "test_series")
print("Data dir:", DATA_DIR)

# ── series selection: sagittal + fluid-sensitive, largest per study ──
series = pd.read_csv(os.path.join(DATA_DIR, "test_series.csv"))
series = series[(series["Anatomical_Plane"] == "Sagittal") & (series["Fluid_Sensitive"] == 1)]
from collections import Counter
_counts = Counter()
for _root, _dirs, _files in os.walk(SERIES_DIR):
    if _files:
        _counts[os.path.basename(_root)] += len(_files)
series["n_files"] = series["SeriesInstanceUID"].map(lambda s: _counts.get(s, 0))
series = series.sort_values("n_files", ascending=False).drop_duplicates("StudyInstanceUID")
print(f"sagittal-fluid series selected: {len(series)}")

test = pd.read_csv(os.path.join(DATA_DIR, "test.csv"))
print(f"test.csv studies: {len(test)}")
# NOTE: keep ALL test UIDs — studies without a usable series get 0.5 priors
# below. Kaggle requires every sample_submission row in the file.

# ── DICOM → middle slices (identical to train kernel) ──
def load_slices(study_uid):
    sel = series[series["StudyInstanceUID"] == study_uid]
    if len(sel) == 0:
        return None
    series_uid = sel.iloc[0]["SeriesInstanceUID"]
    files = sorted(glob.glob(os.path.join(SERIES_DIR, study_uid, series_uid, "*.dcm")))
    if len(files) == 0:
        return None
    arrays = []
    for f in files:
        try:
            dcm = pydicom.dcmread(f, force=True)
            arr = dcm.pixel_array.astype(np.float32)
            lo, hi = np.percentile(arr, 2), np.percentile(arr, 98)
            if hi - lo < 1e-3:
                continue
            arr = np.clip((arr - lo) / (hi - lo), 0, 1)
            arr = cv2.resize(arr, (IMG_SIZE, IMG_SIZE), interpolation=cv2.INTER_AREA)
            arrays.append(arr)
        except Exception:
            continue
    if len(arrays) == 0:
        return None
    arrays = np.stack(arrays)
    n = len(arrays)
    mid = n // 2
    start = max(0, mid - N_SLICES // 2)
    end = min(n, mid + N_SLICES // 2)
    if end - start < N_SLICES:
        idxs = np.linspace(0, n - 1, N_SLICES).astype(int)
        out = arrays[idxs]
    else:
        out = arrays[start:end]
    return out

# ── model (private dataset rsna-assets — no external egress) ──
# 🔴 v11 (Sep 6): DISTILL2 B3 weights — rsna_model_b3_distill2.pt
# (8-epoch distill vs 0.9487 XLM-R-large teacher, gold-val AUC 0.6990 —
# vs 0.6824 plain B3 / distill1 teacher-val 0.7699). Falls back to distill1.
import timm
model = timm.create_model("tf_efficientnet_b3", pretrained=False, num_classes=12)
model.eval()
_ckpts = glob.glob("/kaggle/input/**/rsna_model_b3_distill2.pt", recursive=True)
if not _ckpts:
    _ckpts = glob.glob("/kaggle/input/**/rsna_model_b3_distill.pt", recursive=True)
assert _ckpts, f"distill2/distill checkpoint not found under /kaggle/input ({os.listdir('/kaggle/input')})"
ckpt = _ckpts[0]
sd = torch.load(ckpt, map_location="cpu")
model.load_state_dict(sd)
print("Model loaded from", ckpt, flush=True)

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

def predict_study(model, slices):
    """slices: (S,224,224) float32 0-1 → (12,) sigmoid probs"""
    imgs = np.repeat(slices[:, None, :, :], 3, axis=1).astype(np.float32)
    imgs = torch.from_numpy(imgs)
    with torch.no_grad():
        imgs = (imgs - IMAGENET_MEAN) / IMAGENET_STD
        B = 32
        chunks = [model(imgs[i:i+B]) for i in range(0, len(imgs), B)]
        logits = torch.cat(chunks, 0) if len(chunks) > 1 else chunks[0]
        return torch.sigmoid(logits.mean(0)).numpy()

# ── predict all test studies ──
rows = []
for uid in test["StudyInstanceUID"]:
    sl = load_slices(uid)
    if sl is None:
        print(f"⚠️ no usable slices for {uid} — using 0.5 priors")
        rows.append([uid] + [0.5]*12)
        continue
    probs = predict_study(model, sl)
    rows.append([uid] + [float(x) for x in probs])
    print(f"{uid}: pred mean {probs.mean():.3f} std {probs.std():.3f}", flush=True)

sub = pd.DataFrame(rows, columns=["StudyInstanceUID"] + LABELS)
sub.to_csv("submission.csv", index=False)
print("✅ submission.csv saved:", sub.shape)
print(sub.to_string(index=False))
