#!/usr/bin/env python
# -*- coding: utf-8 -*-
"""
RSNA Knee Abnormality Detection - vision baseline (Kaggle notebook)
...
"""
from __future__ import annotations
import os
for _v in ("OMP_NUM_THREADS", "OPENBLAS_NUM_THREADS", "MKL_NUM_THREADS"):
    os.environ.setdefault(_v, "4")

# --- GPU compatibility shim -------------------------------------------------
# Kaggle's free tier sometimes assigns a Tesla P100 (sm_60), which the bundled
# torch build (sm_70+) cannot run. Detect it and install a cu118 build that
# supports sm_60. Runs before `import torch` so the main code gets the new build.
import subprocess, sys
def _ensure_gpu_compat():
    try:
        out = subprocess.check_output(
            ["nvidia-smi", "--query-gpu=name", "--format=csv,noheader"]).decode().strip()
        if "P100" in out:
            print("P100 (sm_60) detected; installing torch cu118 for compatibility", flush=True)
            subprocess.check_call([sys.executable, "-m", "pip", "install", "-q",
                "numpy==1.26.4"])
            subprocess.check_call([sys.executable, "-m", "pip", "install", "-q",
                "torch==2.2.2", "torchvision==0.17.2",
                "--index-url", "https://download.pytorch.org/whl/cu118"])
            print("torch cu118 installed", flush=True)
    except Exception as e:
        print("gpu compat check skipped:", e, flush=True)
_ensure_gpu_compat()

import gc, re, time, hashlib, math, traceback
from copy import deepcopy
from pathlib import Path
from concurrent.futures import ThreadPoolExecutor

import numpy as np
import pandas as pd
import pydicom
import torch
import torch.nn as nn
import torch.nn.functional as F

T0 = time.time()
SEED = 2026
np.random.seed(SEED); torch.manual_seed(SEED)

TARGETS = ["ACL", "MCL", "Medial Meniscus", "Lateral Meniscus", "Medial OA",
           "Lateral OA", "PF OA", "Effusion", "Synovitis", "Baker's",
           "Contusion", "Fracture"]

# ---------------- config -------------------------------------------------
IMG = 224
CROP_MM = 160.0
GROUP = 3
N_GROUP_MAX = 3
CACHE_BUDGET_GB = 12.0
HDR_THREADS = 16
PIX_THREADS = 12

EPOCHS = 12
BATCH_STUDIES = 8
LR_BACKBONE = 8e-6
LR_HEAD = 1e-3
WEIGHT_DECAY = 0.02
UNFREEZE_LAST = 6
EVAL_BATCH = 12
TIME_BUDGET = 4.0 * 3600

N_FOLDS = 4
MAX_FOLDS = "auto"
FOLD_TIME_PAD = 1.15
GOLD_WEIGHT = 3.0

USE_VFLIP = False
AUG_AFFINE = True
AUG_ROT_DEG = 8.0
AUG_SCALE = 0.08
AUG_SHIFT = 0.05
AUG_INTENSITY = 0.10
EMA_DECAY = 0.997
RANK_LOSS_W = 0.05
RANK_POS, RANK_NEG = 0.60, 0.40

LAT_FALLBACK = "auto"
LAT_MIN_AGREEMENT = 0.85
LAT_MIN_OFFSET_MM = 5.0

# DINOv2 weights: local mount path if attached, else HF hub id (needs internet).
DINOV2_LOCAL_PREFIX = "/kaggle/input"
DINOV2_HF_ID = "facebook/dinov2-small"

SLOTS = [
    ("SAG_FLUID_FS", "Sagittal", True, True),
    ("COR_FLUID_FS", "Coronal", True, True),
    ("AX_FLUID_FS", "Axial", True, True),
    ("SAG_FLUID_NOFS", "Sagittal", True, False),
    ("COR_T1", "Coronal", False, False),
    ("SAG_T1", "Sagittal", False, False),
]
N_SLOT = len(SLOTS)

FATSAT_OPTS = {"FS", "FATSAT", "FAT_SAT", "FSAT"}
_SEP = re.compile(r"[_\-.]")
_FATSAT_RX = re.compile(r"\bfs\b|fatsat|fat sat|\bstir\b|\bspair\b|\bspir\b|\bwe\b|"
                        r"water excit|\btirm\b|\bsting\b|\bfatsup\b")
_T1_RX = re.compile(r"\bt1\b|\bt1w\b")
_T2_RX = re.compile(r"\bt2\b|\bt2w\b")
_PD_RX = re.compile(r"\bpd\b|\bpdw\b|proton|\bdp\b|dens")


def log(msg):
    print(f"[{time.time() - T0:7.1f}s] {msg}", flush=True)


def find_root():
    for c in [Path("/kaggle/input/competitions/rsna-knee-abnormality-detection"),
              Path("/kaggle/input/rsna-knee-abnormality-detection"),
              Path("data"), Path(".")]:
        if (c / "test.csv").is_file() and (c / "test_series").is_dir():
            return c
    for d1 in sorted(p for p in Path("/kaggle/input").iterdir() if p.is_dir()):
        for cand in [d1] + sorted(p for p in d1.iterdir() if p.is_dir()):
            if (cand / "test.csv").is_file():
                return cand
    raise FileNotFoundError("competition mount not found")


def find_dinov2_local(variant="small"):
    base = Path(DINOV2_LOCAL_PREFIX)
    if not base.is_dir():
        return None
    hits = []
    for root, dirs, files in os.walk(base):
        dirs[:] = [d for d in dirs if d not in ("train_series", "test_series")]
        if "config.json" in files and "dinov2" in root.lower():
            hits.append(Path(root))
    for h in hits:
        if variant in str(h).lower():
            return h
    return hits[0] if hits else None


def load_backbone():
    from transformers import AutoModel
    p = find_dinov2_local("small")
    if p is not None:
        log(f"using local DINOv2 weights at {p}")
        return AutoModel.from_pretrained(str(p))
    log(f"no local DINOv2 found; downloading {DINOV2_HF_ID} from HF hub")
    return AutoModel.from_pretrained(DINOV2_HF_ID)


ROOT = find_root()
log(f"input root: {ROOT}")


def plan_cache(n_study):
    per_slice = n_study * N_SLOT * IMG * IMG
    afford = int(CACHE_BUDGET_GB * 1024 ** 3 // max(per_slice, 1))
    return max(1, min(N_GROUP_MAX, afford // GROUP))


N_GROUP = plan_cache(len(pd.read_csv(ROOT / "train.csv")))
CACHE_SLICES = GROUP * N_GROUP
log(f"cache layout: {N_GROUP} groups x {GROUP} slices = {CACHE_SLICES}/slot")


# ---------------- header pass --------------------------------------------
HDR_TAGS = ["SeriesDescription", "SequenceName", "ScanOptions", "ScanningSequence",
            "RepetitionTime", "EchoTime", "Laterality", "ImageLaterality",
            "ImagePositionPatient", "PixelSpacing", "Rows", "Columns"]


def probe(item):
    split, study, series, path = item
    row = {"split": split, "StudyInstanceUID": study, "SeriesInstanceUID": series, "dir": path}
    try:
        files = sorted(e.name for e in os.scandir(path) if e.name.endswith(".dcm"))
        row["files"] = files
        row["n_slices"] = len(files)
        if files:
            ds = pydicom.dcmread(os.path.join(path, files[len(files) // 2]),
                                 stop_before_pixels=True, force=True)
            for t in HDR_TAGS:
                v = getattr(ds, t, None)
                row[t] = "|".join(str(x) for x in v) if v is not None else None
    except Exception as exc:
        row["err"] = str(exc)[:120]
    return row


def walk(split):
    base = ROOT / split
    items = []
    if base.is_dir():
        for st in os.scandir(base):
            if st.is_dir():
                for se in os.scandir(st.path):
                    if se.is_dir():
                        items.append((split, st.name, se.name, se.path))
    with ThreadPoolExecutor(max_workers=HDR_THREADS) as pool:
        return pd.DataFrame(list(pool.map(probe, items)))


def annotate(df):
    desc = (df["SeriesDescription"].fillna("") + " " + df["SequenceName"].fillna(""))
    desc = desc.str.lower().str.replace(_SEP, " ", regex=True)
    opts = df["ScanOptions"].fillna("").str.upper().str.split("|")
    opts_fs = opts.apply(lambda ts: any(t.strip() in FATSAT_OPTS for t in ts))
    df["fatsat"] = desc.str.contains(_FATSAT_RX) | opts_fs
    tr = pd.to_numeric(df["RepetitionTime"], errors="coerce")
    te = pd.to_numeric(df["EchoTime"], errors="coerce")
    gre = df["ScanningSequence"].fillna("").str.upper().str.contains("GR")
    t1, t2, pdw = desc.str.contains(_T1_RX), desc.str.contains(_T2_RX), desc.str.contains(_PD_RX)
    df["weight"] = np.where(t1 & ~t2 & ~pdw, "T1",
                     np.where(t2 & ~pdw, "T2",
                       np.where(pdw, "PD",
                         np.where(gre, "GRE",
                           np.where(tr < 800, "T1",
                             np.where(te > 60, "T2", np.where(tr >= 800, "PD", "UNK")))))))
    df["fluid"] = np.isin(df["weight"], ["PD", "T2"])
    df["px"] = pd.to_numeric(
        df["PixelSpacing"].fillna("").str.split("|").str[0].replace("", np.nan), errors="coerce")
    return df


# ---------------- laterality ---------------------------------------------
def _tag_side(g):
    v = [str(x).strip().upper() for x in g["Laterality"].dropna()]
    if "ImageLaterality" in g.columns:
        v += [str(x).strip().upper() for x in g["ImageLaterality"].dropna()]
    v = [x[0] for x in v if x and x[0] in ("L", "R")]
    return v[0] if v else None


def _position_side(g):
    xs = []
    for s in g.get("ImagePositionPatient", pd.Series(dtype=object)).dropna():
        try:
            xs.append(float(str(s).split("|")[0]))
        except Exception:
            pass
    if not xs:
        return None
    x = float(np.median(xs))
    return None if abs(x) < LAT_MIN_OFFSET_MM else ("R" if x < 0 else "L")


def laterality_maps(h):
    tag, pos = {}, {}
    for st, g in h.groupby("StudyInstanceUID"):
        tag[st] = _tag_side(g); pos[st] = _position_side(g)
    both = [st for st in tag if tag[st] and pos[st]]
    agree = float(np.mean([tag[st] == pos[st] for st in both])) if both else np.nan
    if LAT_FALLBACK == "on":
        use = True
    elif LAT_FALLBACK == "off":
        use = False
    else:
        use = bool(both) and np.isfinite(agree) and agree >= LAT_MIN_AGREEMENT
    side = {st: (tag[st] or (pos[st] if use else None)) for st in tag}
    covered = float(np.mean([v is not None for v in side.values()]))
    log(f"laterality: tag on {np.mean([v is not None for v in tag.values()]):.1%}, "
        f"x agrees {agree:.1%} of {len(both)}, fallback {'on' if use else 'off'} "
        f"-> {covered:.1%} normalised")
    return side


# ---------------- slot picking & decoding --------------------------------
def pick_slots(series_df, plane_map):
    series_df = series_df.copy()
    series_df["plane"] = series_df["SeriesInstanceUID"].map(plane_map)
    out = {}
    for study, g in series_df.groupby("StudyInstanceUID"):
        chosen = {}
        for name, plane, fluid, fs in SLOTS:
            sel = (g["plane"] == plane) & (g["fatsat"] == fs)
            if fluid is not None:
                sel &= (g["fluid"] == fluid)
            cand = g[sel]
            if len(cand) == 0 and fluid is False:
                cand = g[(g["plane"] == plane) & (~g["fatsat"])]
            if len(cand):
                chosen[name] = cand.sort_values("n_slices", ascending=False).iloc[0]
        out[study] = chosen
    return out


def read_slot(rec, n_slice=None, out_size=None):
    n_slice = CACHE_SLICES if n_slice is None else n_slice
    out_size = IMG if out_size is None else out_size
    files, d, px = rec["files"], rec["dir"], rec["px"]
    n = len(files)
    if n == 0:
        return None
    lo, hi = int(0.20 * (n - 1)), int(0.80 * (n - 1))
    idx = np.unique(np.linspace(lo, hi, n_slice).astype(int)) if hi > lo else np.array([n // 2])
    while len(idx) < n_slice:
        idx = np.append(idx, idx[-1])
    planes = []
    for i in idx[:n_slice]:
        try:
            ds = pydicom.dcmread(os.path.join(d, files[int(i)]), force=True)
            a = ds.pixel_array.astype(np.float32)
            sl = float(getattr(ds, "RescaleSlope", 1) or 1)
            ic = float(getattr(ds, "RescaleIntercept", 0) or 0)
            a = a * sl + ic
        except Exception:
            a = np.zeros((out_size, out_size), dtype=np.float32)
        planes.append(a)
    shp = planes[0].shape
    planes = [p if p.shape == shp else np.zeros(shp, np.float32) for p in planes]
    vol = np.stack(planes)
    if px and np.isfinite(px) and px > 0:
        want = int(round(CROP_MM / px))
        h, w = shp
        if 16 < want < min(h, w):
            cy, cx = h // 2, w // 2
            half = want // 2
            vol = vol[:, max(0, cy - half):cy + half, max(0, cx - half):cx + half]
    lo_v, hi_v = np.percentile(vol, [1, 99])
    vol = np.clip((vol - lo_v) / max(hi_v - lo_v, 1e-6), 0, 1)
    t = torch.from_numpy(np.ascontiguousarray(vol)).unsqueeze(0)
    t = F.interpolate(t, size=(out_size, out_size), mode="bilinear", align_corners=False)
    return (t.squeeze(0) * 255).round().clamp(0, 255).to(torch.uint8)


def normalise_laterality(img, plane, lat):
    if lat != "R":
        return img
    if plane in ("Coronal", "Axial"):
        return torch.flip(img, dims=[-1])
    return torch.flip(img, dims=[0])


def build_cache(slot_map, plane_map, lat_map, tag):
    studies = sorted(slot_map)
    sidx = {s: i for i, s in enumerate(studies)}
    cache = np.zeros((len(studies), N_SLOT, CACHE_SLICES, IMG, IMG), np.uint8)
    mask = np.zeros((len(studies), N_SLOT), np.float32)
    log(f"{tag}: cache {cache.shape} = {cache.nbytes / 1024 ** 3:.1f} GB")
    jobs = [(st, k, plane, slot_map[st][name])
            for st in studies for k, (name, plane, _, _) in enumerate(SLOTS)
            if name in slot_map[st]]
    log(f"{tag}: decoding {len(jobs)} slot-series")
    done = 0
    with ThreadPoolExecutor(max_workers=PIX_THREADS) as pool:
        for c0 in range(0, len(jobs), 512):
            block = jobs[c0:c0 + 512]
            for (st, k, plane, _), img in zip(
                    block, pool.map(lambda j: read_slot(j[3], CACHE_SLICES, IMG), block)):
                done += 1
                if img is None:
                    continue
                cache[sidx[st], k] = normalise_laterality(img, plane, lat_map.get(st)).numpy()
                mask[sidx[st], k] = 1.0
            if done % 4096 < 512:
                log(f"  {tag} {done}/{len(jobs)}")
            if time.time() - T0 > TIME_BUDGET:
                break
    gc.collect()
    return studies, cache, mask


# ---------------- model ---------------------------------------------------
class SlotHead(nn.Module):
    def __init__(self, dim, n_slot, n_out, hidden=256, p=0.2):
        super().__init__()
        self.proj = nn.Sequential(nn.LayerNorm(dim), nn.Linear(dim, hidden), nn.GELU())
        self.slot_emb = nn.Parameter(torch.randn(n_slot, hidden) * 0.02)
        self.query = nn.Parameter(torch.randn(n_out, hidden) * 0.02)
        self.drop = nn.Dropout(p)
        self.out = nn.Linear(hidden, n_out)
        self.hidden = hidden

    def forward(self, x, mask):
        h = self.proj(x) + self.slot_emb
        att = torch.einsum("bsh,oh->bos", h, self.query) / self.hidden ** 0.5
        att = att.masked_fill(mask.unsqueeze(1) < 0.5, -1e4).softmax(-1)
        ctx = self.drop(torch.einsum("bos,bsh->boh", att, h))
        return (ctx * self.out.weight.unsqueeze(0)).sum(-1) + self.out.bias


class Model(nn.Module):
    def __init__(self, backbone, dim):
        super().__init__()
        self.backbone = backbone
        self.head = SlotHead(dim, N_SLOT, len(TARGETS))
        self.register_buffer("mean", torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1))
        self.register_buffer("std", torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1))

    def forward(self, imgs, mask):
        B, S = imgs.shape[:2]
        x = imgs.reshape(B * S, *imgs.shape[2:]).float().div_(255.0)
        x = (x - self.mean) / self.std
        out = self.backbone(pixel_values=x).last_hidden_state
        feat = torch.cat([out[:, 0], out[:, 1:].mean(1)], dim=1).reshape(B, S, -1)
        return self.head(feat, mask)


def build_model():
    bb = load_backbone()
    n_layer = len(bb.encoder.layer)
    for prm in bb.parameters():
        prm.requires_grad = False
    for blk in bb.encoder.layer[max(0, n_layer - UNFREEZE_LAST):]:
        for prm in blk.parameters():
            prm.requires_grad = True
    for prm in bb.layernorm.parameters():
        prm.requires_grad = True
    dim = bb.config.hidden_size * 2
    log(f"backbone {n_layer} blocks, last {UNFREEZE_LAST} trainable "
        f"({sum(p.numel() for p in bb.parameters() if p.requires_grad) / 1e6:.1f}M), dim {dim}")
    return Model(bb, dim)


def take_group(cache_rows, g):
    return cache_rows[:, :, g * GROUP:(g + 1) * GROUP]


@torch.no_grad()
def predict(model, cache, mask, idx, dev):
    model.eval()
    out = []
    for b in range(0, len(idx), EVAL_BATCH):
        sel = idx[b:b + EVAL_BATCH]
        rows = torch.from_numpy(cache[sel]).to(dev)
        m = torch.from_numpy(mask[sel]).to(dev)
        acc = None
        for g in range(N_GROUP):
            with torch.autocast("cuda", enabled=dev.type == "cuda"):
                z = model(take_group(rows, g), m).float()
            acc = z if acc is None else acc + z
        out.append(torch.sigmoid(acc / N_GROUP).cpu().numpy())
    return np.concatenate(out) if out else np.zeros((0, len(TARGETS)), np.float32)


def macro_auc(y, p):
    from sklearn.metrics import roc_auc_score
    return float(np.nanmean([roc_auc_score(y[:, j], p[:, j])
                             if len(set(y[:, j])) > 1 else np.nan
                             for j in range(y.shape[1])]))


def augment(imgs):
    B, S, C, H, W = imgs.shape
    x = imgs.float()
    if AUG_AFFINE:
        ang = math.radians((torch.rand(1).item() - 0.5) * 2 * AUG_ROT_DEG)
        sc = 1.0 + (torch.rand(1).item() - 0.5) * 2 * AUG_SCALE
        tx = (torch.rand(1).item() - 0.5) * 2 * AUG_SHIFT
        ty = (torch.rand(1).item() - 0.5) * 2 * AUG_SHIFT
        cos, sin = math.cos(ang) / sc, math.sin(ang) / sc
        theta = torch.tensor([[cos, -sin, tx], [sin, cos, ty]],
                             dtype=torch.float32, device=imgs.device).unsqueeze(0).repeat(B * S, 1, 1)
        flat = x.reshape(B * S, C, H, W)
        grid = F.affine_grid(theta, flat.shape, align_corners=False)
        x = F.grid_sample(flat, grid, mode="bilinear", padding_mode="zeros",
                          align_corners=False).reshape(B, S, C, H, W)
    scale = 1.0 + (torch.rand(1, device=imgs.device) - 0.5) * 2 * AUG_INTENSITY
    x = (x * scale).clamp(0, 255)
    return x.round().to(imgs.dtype)


class Ema:
    def __init__(self, model, decay):
        self.decay = decay
        self.step = 0
        self.model = deepcopy(model).eval()
        for p in self.model.parameters():
            p.requires_grad_(False)

    @torch.no_grad()
    def update(self, model):
        if self.decay <= 0:
            return
        self.step += 1
        decay = min(self.decay, (1.0 + self.step) / (10.0 + self.step))
        src = model.state_dict()
        for k, v in self.model.state_dict().items():
            s = src[k]
            if v.dtype.is_floating_point:
                v.mul_(decay).add_(s.detach(), alpha=1.0 - decay)
            else:
                v.copy_(s)

    def target(self, model):
        return self.model if self.decay > 0 else model


def rank_loss(logits, y, w):
    parts = []
    usable = w > 0
    for j in range(logits.shape[1]):
        pos = logits[(y[:, j] > RANK_POS) & usable[:, j], j]
        neg = logits[(y[:, j] < RANK_NEG) & usable[:, j], j]
        if len(pos) and len(neg):
            parts.append(F.softplus(-(pos[:, None] - neg[None, :])).mean())
    return torch.stack(parts).mean() if parts else logits.new_tensor(0.0)


# ---------------- training ------------------------------------------------
def write_benchmark_submission():
    t = pd.read_csv(ROOT / "test.csv")
    for c in TARGETS:
        t[c] = 0.5
    t.to_csv("submission.csv", index=False)


def write_submission(rank_sum, n_models, st_te, test_df):
    P = rank_sum / max(n_models, 1)
    sub = pd.DataFrame(P, columns=TARGETS)
    sub.insert(0, "StudyInstanceUID", st_te)
    sub = test_df[["StudyInstanceUID"]].merge(sub, on="StudyInstanceUID", how="left")
    sub[TARGETS] = sub[TARGETS].fillna(0.5)
    sub.to_csv("submission.csv", index=False)
    return sub


def train_one_fold(fold, Ctr, Mtr, Y, W, tr, va, gi_va, gold_y_va, dev):
    model = build_model().to(dev)
    ema = Ema(model, EMA_DECAY)
    opt = torch.optim.AdamW([
        {"params": [p for p in model.backbone.parameters() if p.requires_grad], "lr": LR_BACKBONE},
        {"params": model.head.parameters(), "lr": LR_HEAD},
    ], weight_decay=WEIGHT_DECAY)
    steps = max(EPOCHS * (len(tr) // BATCH_STUDIES), 1)
    sched = torch.optim.lr_scheduler.OneCycleLR(
        opt, max_lr=[LR_BACKBONE, LR_HEAD], total_steps=steps, pct_start=0.15)
    scaler = torch.cuda.amp.GradScaler(enabled=dev.type == "cuda")
    yv = (Y[va] > 0.5).astype(int)
    best, best_state, history = -1.0, None, []

    for ep in range(EPOCHS):
        model.train()
        perm = np.random.permutation(tr)
        tot, nstep = 0.0, 0
        for b in range(0, len(perm) - BATCH_STUDIES + 1, BATCH_STUDIES):
            sel = perm[b:b + BATCH_STUDIES]
            rows = torch.from_numpy(Ctr[sel]).to(dev)
            g = int(torch.randint(N_GROUP, (1,)).item())
            imgs = augment(take_group(rows, g))
            m = torch.from_numpy(Mtr[sel]).to(dev)
            y = torch.from_numpy(Y[sel]).to(dev)
            w = torch.from_numpy(W[sel]).to(dev)
            with torch.autocast("cuda", enabled=dev.type == "cuda"):
                z = model(imgs, m)
                loss = (F.binary_cross_entropy_with_logits(z, y, reduction="none") * w).mean()
                if RANK_LOSS_W > 0:
                    loss = loss + RANK_LOSS_W * rank_loss(z.float(), y, w)
            opt.zero_grad(set_to_none=True)
            scaler.scale(loss).backward()
            scaler.step(opt); scaler.update(); sched.step(); ema.update(model)
            tot += loss.item(); nstep += 1

        eval_model = ema.target(model)
        pv = predict(eval_model, Ctr, Mtr, va, dev)
        d = macro_auc(yv, pv)
        g_auc = float("nan")
        if gold_y_va is not None and len(gi_va) >= 6:
            g_auc = macro_auc(gold_y_va, predict(eval_model, Ctr, Mtr, gi_va, dev))
        history.append({"fold": fold, "epoch": ep + 1, "loss": tot / max(nstep, 1),
                        "derived": d, "annot": g_auc})
        log(f"fold {fold} ep {ep + 1}/{EPOCHS} loss {tot / max(nstep, 1):.4f} "
            f"derived {d:.4f} annot {g_auc:.4f}")
        score = d if not np.isfinite(g_auc) else min(d, g_auc)
        if score > best:
            best = score
            best_state = {k: v.detach().cpu().clone() for k, v in eval_model.state_dict().items()}
        if time.time() - T0 > TIME_BUDGET:
            break
    if best_state is None:
        best_state = {k: v.detach().cpu().clone() for k, v in ema.target(model).state_dict().items()}
    del model, opt, sched, scaler
    gc.collect()
    if torch.cuda.is_available():
        torch.cuda.empty_cache()
    return best_state, history, best


def main():
    write_benchmark_submission()
    test_df = pd.read_csv(ROOT / "test.csv")
    test_series = pd.read_csv(ROOT / "test_series.csv")
    train_df = pd.read_csv(ROOT / "train.csv")
    train_series = pd.read_csv(ROOT / "train_series.csv")
    log(f"train {train_df.shape} test {test_df.shape}")

    both = pd.concat([train_series, test_series])
    plane_map = dict(zip(both["SeriesInstanceUID"], both["Anatomical_Plane"]))

    log("header pass: test"); hte = annotate(walk("test_series")); log(f"  {len(hte)} series")
    log("header pass: train"); htr = annotate(walk("train_series")); log(f"  {len(htr)} series")
    lat_tr = laterality_maps(htr); lat_te = laterality_maps(hte)
    slots_te, slots_tr = pick_slots(hte, plane_map), pick_slots(htr, plane_map)
    cov = pd.Series([len(v) for v in slots_tr.values()]).describe()
    log(f"train slots/study mean {cov['mean']:.2f} min {cov['min']:.0f} max {cov['max']:.0f}")

    st_tr, Ctr, Mtr = build_cache(slots_tr, plane_map, lat_tr, "train")
    st_te, Cte, Mte = build_cache(slots_te, plane_map, lat_te, "test")

    # ---- labels from the published datasets ------------------------------ #
    # Prefer the LLM-derived labels (pilkwang/rsna-knee-llm-labels): they score
    # ~0.87 macro AUC vs the 58 gold studies, vs ~0.73 for the rule-based ones.
    # Fall back to our own report targets if the LLM file is not attached.
    lab = None
    for cand in Path(DINOV2_LOCAL_PREFIX).rglob("report_labels_v2.csv"):
        lab = pd.read_csv(cand)
        log(f"loaded LLM targets from {cand}")
        break
    if lab is None:
        for cand in Path(DINOV2_LOCAL_PREFIX).rglob("train_targets.csv"):
            lab = pd.read_csv(cand)
            log(f"loaded rule-based targets from {cand}")
            break
    if lab is None:
        raise FileNotFoundError("no label source found under /kaggle/input; "
                                "attach pilkwang/rsna-knee-llm-labels or gabrielep09/rsna-knee-report-targets")
    has_conf = [t + "__conf" for t in TARGETS if (t + "__conf") in lab.columns]
    lab = lab.set_index("StudyInstanceUID")
    log(f"label source for {len(lab)} studies (conf cols: {len(has_conf)})")
    gold = train_df.set_index("StudyInstanceUID")[TARGETS]
    gold = gold[gold.notna().all(axis=1)]

    Y = np.zeros((len(st_tr), len(TARGETS)), np.float32)
    W = np.zeros_like(Y)
    for i, st in enumerate(st_tr):
        if st in gold.index:
            Y[i], W[i] = gold.loc[st].values, GOLD_WEIGHT
        elif st in lab.index:
            r = lab.loc[st]
            Y[i] = r[TARGETS].values
            if len(has_conf) == len(TARGETS):
                # 0.25 + 0.75*conf: uncertain labels pull toward the derived floor
                W[i] = 0.25 + 0.75 * r[has_conf].values
            else:
                W[i] = np.clip(0.5 + 3.0 * np.abs(Y[i] - 0.5), 0.25, 1.0)
    keep = np.where(W.sum(1) > 0)[0]
    log(f"supervised {len(keep)} of {len(st_tr)} (gold {len(gold)})")

    # report-hash folds keep identical reports together
    rep = train_df.set_index("StudyInstanceUID")["Report"].fillna("")
    grp = np.array([int(hashlib.md5(rep.get(s, s).encode()).hexdigest()[:8], 16) % N_FOLDS
                    for s in st_tr])
    gpos = {s: i for i, s in enumerate(st_tr)}
    gi_all = np.array([gpos[s] for s in gold.index if s in gpos])
    gold_y_all = (gold.loc[[st_tr[i] for i in gi_all]].values.astype(int) if len(gi_all) else None)

    dev = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    rank_sum = np.zeros((len(st_te), len(TARGETS)), np.float64)
    n_models, sub = 0, None
    oof_gold = np.full((len(st_tr), len(TARGETS)), np.nan, np.float32)
    histories, fold_scores = [], []

    max_folds = N_FOLDS if MAX_FOLDS == "auto" else int(MAX_FOLDS)
    for fold in range(min(N_FOLDS, max_folds)):
        va = np.array([i for i in keep if grp[i] == fold])
        tr = np.array([i for i in keep if grp[i] != fold])
        if len(va) == 0 or len(tr) < BATCH_STUDIES:
            continue
        gi_va = np.array([i for i in gi_all if grp[i] == fold])
        gold_y_va = (gold.loc[[st_tr[i] for i in gi_va]].values.astype(int) if len(gi_va) else None)
        log(f"=== fold {fold}: train {len(tr)} / holdout {len(va)} (annot held out {len(gi_va)}) ===")
        t_fold = time.time()
        state, history, best = train_one_fold(fold, Ctr, Mtr, Y, W, tr, va, gi_va, gold_y_va, dev)
        fold_time = time.time() - t_fold
        histories.extend(history); fold_scores.append(best)

        model = build_model().to(dev)
        model.load_state_dict(state)
        if len(gi_va):
            oof_gold[gi_va] = predict(model, Ctr, Mtr, gi_va, dev)
        P = predict(model, Cte, Mte, np.arange(len(st_te)), dev)
        rank_sum += pd.DataFrame(P).rank(pct=True).values
        n_models += 1
        del model, state
        gc.collect()
        if torch.cuda.is_available():
            torch.cuda.empty_cache()
        sub = write_submission(rank_sum, n_models, st_te, test_df)
        log(f"fold {fold} done in {fold_time / 60:.1f} min; {n_models} model(s) in submission")
        elapsed = time.time() - T0
        if elapsed + FOLD_TIME_PAD * fold_time > TIME_BUDGET:
            log("stopping: another fold would exceed the time budget")
            break

    if gold_y_all is not None and len(gi_all):
        seen = np.isfinite(oof_gold[gi_all]).all(axis=1)
        if seen.sum() >= 8:
            log(f"out-of-fold macro AUC on {int(seen.sum())} gold studies: "
                f"{macro_auc(gold_y_all[seen], oof_gold[gi_all][seen]):.4f}")
    log(f"fold selection scores: {[round(s, 4) for s in fold_scores]}")
    if sub is None:
        log("no fold completed; benchmark submission stands")
        return
    log(f"submission.csv {sub.shape}; models {n_models}")
    print(sub.head().to_string())


try:
    main()
except Exception:
    traceback.print_exc()
    try:
        t = pd.read_csv(ROOT / "test.csv")
        for c in TARGETS:
            t[c] = 0.5
        t.to_csv("submission.csv", index=False)
    except Exception:
        pass
