# %% [code]
from pathlib import Path
import io
import numpy as np
import pandas as pd
import pydicom
import torch
import torch.nn.functional as F
import timm
from safetensors.torch import load_file

weight_candidates = list(Path("/kaggle/input").rglob("coatnet_rmlp_2_rw_384_model.safetensors"))
if len(weight_candidates) != 1:
    raise RuntimeError(f"expected one pinned CoAtNet weight, found {len(weight_candidates)}")
MODEL = weight_candidates[0].parent
competition_candidates = [
    path.parent
    for path in Path("/kaggle/input").rglob("test.csv")
    if (path.parent / "test_series.csv").is_file()
    and (path.parent / "sample_submission.csv").is_file()
]
if len(competition_candidates) != 1:
    raise RuntimeError(f"expected one competition input root, found {len(competition_candidates)}")
COMP = competition_candidates[0]
TARGETS = ["ACL", "MCL", "Medial Meniscus", "Lateral Meniscus", "Medial OA", "Lateral OA", "PF OA", "Effusion", "Synovitis", "Baker's", "Contusion", "Fracture"]
SEEDS = [20260816, 20260817, 20260818, 20260819, 20260820]
PLANES = ["Axial", "Coronal", "Sagittal"]


def percentile_value(values, percentile):
    flat = np.asarray(values, dtype=np.float32).reshape(-1)
    flat = flat[np.isfinite(flat)]
    if not flat.size:
        raise ValueError("empty finite values")
    position = (flat.size - 1) * np.clip(percentile / 100.0, 0.0, 1.0)
    lower, upper = int(np.floor(position)), int(np.ceil(position))
    selected = np.partition(flat, (lower, upper))
    if lower == upper:
        return np.float32(selected[lower])
    fraction = np.float32(position - lower)
    return np.float32(selected[lower]) + fraction * (np.float32(selected[upper]) - np.float32(selected[lower]))


def prepare(dataset):
    image = np.asarray(dataset.pixel_array).T.astype(np.float32)
    image = image * np.float32(getattr(dataset, "RescaleSlope", 1.0)) + np.float32(getattr(dataset, "RescaleIntercept", 0.0))
    finite = image[np.isfinite(image)]
    if not finite.size:
        image = np.zeros_like(image)
    else:
        low, high = percentile_value(finite, 1.0), percentile_value(finite, 99.0)
        image = np.zeros_like(image) if high <= low else np.clip((image - low) / np.float32(high - low), 0, 1).astype(np.float32)
    height, width = image.shape
    scale = min(384 / height, 384 / width)
    new_height, new_width = int(np.clip(round(height * scale), 1, 384)), int(np.clip(round(width * scale), 1, 384))
    tensor = torch.from_numpy(np.ascontiguousarray(image))[None, None]
    resized = F.interpolate(tensor, size=(new_height, new_width), mode="bicubic", align_corners=False, antialias=True)[0, 0].numpy()
    output = np.zeros((384, 384), np.float32)
    top, left = (384 - new_height) // 2, (384 - new_width) // 2
    output[top:top + new_height, left:left + new_width] = resized
    return output


def coordinate(dataset):
    position, orientation = getattr(dataset, "ImagePositionPatient", None), getattr(dataset, "ImageOrientationPatient", None)
    if position is None or orientation is None or len(position) != 3 or len(orientation) != 6:
        return None
    normal = np.cross(np.asarray(orientation[:3], float), np.asarray(orientation[3:], float))
    norm = np.linalg.norm(normal)
    return None if norm <= np.finfo(float).eps else float(np.dot(np.asarray(position, float), normal / norm))


def order_key(path):
    dataset = pydicom.dcmread(path, stop_before_pixels=True)
    value, sop = coordinate(dataset), str(dataset.SOPInstanceUID)
    if value is not None:
        return 1, value, sop
    if hasattr(dataset, "SliceLocation"):
        return 2, float(dataset.SliceLocation), sop
    if hasattr(dataset, "InstanceNumber"):
        return 3, float(dataset.InstanceNumber), sop
    raise RuntimeError(f"unorderable DICOM: {path}")


def centers(count):
    if count <= 8:
        return list(range(count))
    return list(dict.fromkeys((np.rint(np.linspace(1.0, float(count), 8)).astype(int) - 1).tolist()))


def valid_decode(paths, index, cache):
    if index not in cache:
        try:
            dataset = pydicom.dcmread(paths[index])
            pixels = np.asarray(dataset.pixel_array)
            cache[index] = dataset if pixels.ndim == 2 and np.isfinite(pixels.astype(np.float32)).all() else None
        except Exception:
            cache[index] = None
    return cache[index] is not None


def resolve(paths, center, cache):
    total = len(paths)
    valid = lambda index: valid_decode(paths, index, cache)
    if valid(center):
        actual_center = center
    else:
        candidates = [index for index in range(total) if valid(index)]
        if not candidates:
            raise RuntimeError(f"unrepresentable series: {paths[0].parent}")
        actual_center = min(candidates, key=lambda index: (abs(index - center), index))
    resolved = []
    for direction in (-1, 0, 1):
        if direction == 0:
            resolved.append(actual_center)
            continue
        requested = center + direction
        if requested < 0 or requested >= total:
            resolved.append(actual_center)
        elif valid(requested):
            resolved.append(requested)
        else:
            index = requested + direction
            while 0 <= index < total and not valid(index):
                index += direction
            resolved.append(index if 0 <= index < total else actual_center)
    return resolved


if not torch.cuda.is_available():
    raise RuntimeError("CUDA is required for Stage11R inference; refusing CPU fallback")
if torch.cuda.get_device_capability(0)[0] < 7:
    raise RuntimeError("GPU compute capability is insufficient for Stage11R inference")
device = torch.device("cuda")
print(f"INFERENCE_DEVICE={device}")
encoder = timm.create_model("coatnet_rmlp_2_rw_384", pretrained=False)
state = load_file(MODEL / "coatnet_rmlp_2_rw_384_model.safetensors")
encoder.load_state_dict(state, strict=True)
encoder.reset_classifier(0)
encoder.eval().to(device)
for parameter in encoder.parameters():
    parameter.requires_grad_(False)

test = pd.read_csv(COMP / "test.csv", dtype=str)
series_metadata = pd.read_csv(COMP / "test_series.csv", dtype=str)
series_vectors = {}
with torch.inference_mode():
    for row in series_metadata.itertuples(index=False):
        directory = COMP / "test_series" / row.StudyInstanceUID / row.SeriesInstanceUID
        paths = sorted(directory.glob("*.dcm"), key=order_key)
        selected = centers(len(paths))
        if len(selected) != 8:
            raise RuntimeError(f"uniform8 failure: {directory}")
        cache, images = {}, []
        for center in selected:
            triplet = resolve(paths, center, cache)
            images.append(np.stack([prepare(cache[index]) for index in triplet]).astype(np.float32))
        batch = torch.from_numpy(np.stack(images)).to(device)
        features = encoder((batch - 0.5) / 0.5).float().cpu().numpy()
        if features.shape != (8, 1024) or not np.isfinite(features).all():
            raise RuntimeError("encoder feature gate")
        series_vectors[(row.StudyInstanceUID, row.SeriesInstanceUID)] = (row.Anatomical_Plane, features.mean(axis=0, dtype=np.float32))

study_features = []
for study in test.StudyInstanceUID:
    values = [(plane, vector) for (uid, _), (plane, vector) in series_vectors.items() if uid == study]
    if not values:
        raise RuntimeError(f"missing study features: {study}")
    parts = []
    for plane in PLANES:
        selected = [vector for current, vector in values if current == plane]
        parts.append(np.stack(selected).mean(axis=0, dtype=np.float32) if selected else np.zeros(1024, np.float32))
    parts.append(np.stack([vector for _, vector in values]).mean(axis=0, dtype=np.float32))
    study_features.append(np.concatenate(parts))
study_features = torch.from_numpy(np.stack(study_features).astype(np.float32))

predictions = []
for seed in SEEDS:
    head = torch.load(MODEL / f"stage11q_baseline_seed{seed}.pt", map_location="cpu", weights_only=True)
    predictions.append(torch.sigmoid(study_features @ head["weight"].T + head["bias"]).numpy())
ensemble = np.mean(predictions, axis=0, dtype=np.float32)
assert ensemble.shape == (len(test), 12) and np.isfinite(ensemble).all()
submission = pd.DataFrame(ensemble, columns=TARGETS)
submission.insert(0, "StudyInstanceUID", test.StudyInstanceUID)
sample = pd.read_csv(COMP / "sample_submission.csv", dtype={"StudyInstanceUID": str})
assert list(submission.columns) == list(sample.columns)
assert submission.StudyInstanceUID.tolist() == sample.StudyInstanceUID.tolist()
submission.to_csv("/kaggle/working/submission.csv", index=False)
print(f"STAGE11R_INFERENCE_COMPLETE studies={len(submission)}")
