# %% [code]
"""
RSNA Knee Abnormality Detection: v30 INFERENCE
Loads best_model_v29.pth (val AUC 0.742) and runs inference on ALL test studies.
Outputs submission.csv to /kaggle/working.
"""
import os
import sys
import glob
import numpy as np
import pandas as pd
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from pydicom.pixel_data_handlers.util import apply_voi_lut
import pydicom
import cv2

print(f"PyTorch: {torch.__version__}, CUDA: {torch.cuda.is_available()}")
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

BASE = '/kaggle/input/competitions/rsna-knee-abnormality-detection'
TARGETS = ['ACL', 'MCL', 'Medial Meniscus', 'Lateral Meniscus', 'Medial OA',
           'Lateral OA', 'PF OA', 'Effusion', 'Synovitis', "Baker's",
           'Contusion', 'Fracture']

class CFG:
    img_size = 128
    num_slices = 3
    num_planes = 2

def find_csv(name, root='/kaggle/input'):
    for dirpath, _, files in os.walk(root):
        if name in files:
            return os.path.join(dirpath, name)
    raise FileNotFoundError(name)

WEIGHTS = find_csv('best_model_v29.pth')
print('[Model] weights:', WEIGHTS)

# ------------- Model (same as training) -------------
class KneeModel(nn.Module):
    def __init__(self):
        super().__init__()
        from torchvision.models import resnet18
        base = resnet18(weights=None)
        fp = find_csv('resnet18-f37072fd.pth')
        base.load_state_dict(torch.load(fp, map_location='cpu', weights_only=True))
        base.fc = nn.Identity()
        self.backbone = base
        self.att = nn.Sequential(nn.Linear(512, 64), nn.ReLU())
        self.gate = nn.Linear(64, 1)
        self.head = nn.Sequential(
            nn.Linear(512, 256), nn.ReLU(), nn.Dropout(0.2),
            nn.Linear(256, 12))

    def forward(self, x):
        x = x.float()
        B = x.size(0)
        x = x.view(B, -1, 3, x.size(-2), x.size(-1))
        K = x.size(1)
        s = x.view(B * K, 3, x.size(-2), x.size(-1))
        feats = self.backbone(s).view(B, K, -1)
        a = self.att(feats)
        g = torch.softmax(self.gate(a), dim=1)
        fused = (feats * g).sum(dim=1)
        return self.head(fused)

# ------------- Image loading -------------
def prep_stack(files, max_slices=CFG.num_slices, img_size=CFG.img_size):
    imgs = []
    for f in files:
        try:
            dcm = pydicom.dcmread(f, stop_before_pixels=False, force=True)
            px = dcm.pixel_array.astype(np.float32)
            px = apply_voi_lut(px, dcm)
            if px.max() > px.min():
                px = (px - px.min()) / (px.max() - px.min())
            else:
                continue
            px = np.clip(px, 0, 1)
            px = cv2.resize(px, (img_size, img_size), interpolation=cv2.INTER_AREA)
            imgs.append(px)
        except Exception:
            continue
    if not imgs:
        return None
    n = len(imgs)
    idx = np.linspace(0, n - 1, max_slices, dtype=int) if n >= max_slices else np.arange(n)
    stack = np.stack([imgs[i] for i in idx], axis=0)
    rgb = np.repeat(stack[:, None, :, :], 3, axis=1)
    return torch.from_numpy(rgb.astype(np.float32))

# ------------- Dataset -------------
class KneeDataset(Dataset):
    def __init__(self, df, series_df, img_dir):
        self.ids = df['StudyInstanceUID'].astype(str).str.strip().values
        self.img_dir = img_dir
        self.series_map = {}
        self.fluid = {}
        for _, r in series_df.iterrows():
            s_id = str(r.iloc[0]).strip()
            ser = str(r.iloc[1]).strip()
            self.series_map.setdefault(s_id, [])
            if ser not in self.series_map[s_id]:
                self.series_map[s_id].append(ser)
            self.fluid[ser] = int(r.iloc[2]) if 'Fluid_Sensitive' in series_df.columns else 0
        print(f'[Dataset] {len(self.ids)} studies, {len(self.series_map)} mapped')

    def __len__(self):
        return len(self.ids)

    def __getitem__(self, i):
        sid = str(self.ids[i]).strip()
        series_list = self.series_map.get(sid, [])
        fluid = [s for s in series_list if self.fluid.get(s, 0) == 1]
        non_fluid = [s for s in series_list if self.fluid.get(s, 0) == 0]
        selected = []
        for pool in [fluid, non_fluid]:
            for s in pool:
                if len(selected) >= CFG.num_planes:
                    break
                selected.append(s)
            if len(selected) >= CFG.num_planes:
                break
        if not selected:
            selected = series_list[:CFG.num_planes]
        stacks = []
        zero = torch.zeros((CFG.num_slices, 3, CFG.img_size, CFG.img_size))
        for s in selected:
            path = os.path.join(self.img_dir, sid, s)
            if not os.path.exists(path):
                stacks.append(zero)
                continue
            files = sorted(glob.glob(os.path.join(path, '*.dcm')))
            if not files:
                stacks.append(zero)
                continue
            t = prep_stack(files)
            stacks.append(t if t is not None else zero)
        while len(stacks) < CFG.num_planes:
            stacks.append(zero)
        x = torch.stack(stacks[:CFG.num_planes], dim=0)
        return x.view(-1, 3, CFG.img_size, CFG.img_size)

# ------------- Main -------------
def main():
    test_df = pd.read_csv(os.path.join(BASE, 'test.csv'))
    print(f'[Data] test studies: {len(test_df)}')
    img_dir = os.path.join(BASE, 'train_series')
    test_series = pd.read_csv(find_csv('test_series.csv'))
    print(f'[Data] test series rows: {len(test_series)}')

    ds = KneeDataset(test_df, test_series, img_dir)
    loader = DataLoader(ds, batch_size=16, shuffle=False, num_workers=0)

    model = KneeModel().to(device)
    model.load_state_dict(torch.load(WEIGHTS, map_location=device))
    model.eval()

    print('[Inference] running...')
    preds = []
    with torch.no_grad():
        for x in loader:
            x = x.to(device)
            preds.append(torch.sigmoid(model(x)).cpu().numpy())
    preds = np.vstack(preds)
    print(f'[Inference] shape={preds.shape}, mean={preds.mean():.4f}')

    sub = pd.DataFrame(preds, columns=TARGETS)
    sub.insert(0, 'StudyInstanceUID', test_df['StudyInstanceUID'].astype(str).str.strip().values[:len(sub)])
    sub.to_csv('/kaggle/working/submission.csv', index=False)
    print('[Saved] /kaggle/working/submission.csv')
    print(sub.to_string())
    sys.stdout.flush()

if __name__ == '__main__':
    main()
