import os
import glob
import pydicom
import cv2
import numpy as np
import pandas as pd
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader
from pathlib import Path
import yaml
import timm

# ----------------- DATA PREPROCESSING & DATASET -----------------

def load_dicom_slice(path):
    """Loads a single DICOM file and returns its pixel array."""
    try:
        dicom = pydicom.dcmread(path)
        pixel_array = dicom.pixel_array.astype(float)
        
        z_pos = 0.0
        if "ImagePositionPatient" in dicom:
            z_pos = float(dicom.ImagePositionPatient[2])
        elif "SliceLocation" in dicom:
            z_pos = float(dicom.SliceLocation)
        elif "InstanceNumber" in dicom:
            z_pos = float(dicom.InstanceNumber)
            
        return pixel_array, z_pos
    except Exception as e:
        return None, 0.0

def process_series(series_dir, target_slices=32, target_size=224):
    """Loads, sorts, resamples, and normalizes a DICOM series."""
    dicom_files = glob.glob(os.path.join(series_dir, "*.dcm"))
    if len(dicom_files) == 0:
        return None
        
    slices = []
    for f in dicom_files:
        res = load_dicom_slice(f)
        if res[0] is not None:
            slices.append(res)
            
    if len(slices) == 0:
        return None
        
    # Sort slices physically by position
    slices = sorted(slices, key=lambda x: x[1])
    pixel_arrays = [s[0] for s in slices]
    
    # Resize slices
    resized_slices = []
    for arr in pixel_arrays:
        if arr.shape != (target_size, target_size):
            resized = cv2.resize(arr, (target_size, target_size), interpolation=cv2.INTER_AREA)
        else:
            resized = arr
        resized_slices.append(resized)
        
    volume = np.stack(resized_slices, axis=0) # (S, H, W)
    
    # Normalize volume intensities
    p1, p99 = np.percentile(volume, [1, 99])
    if p99 > p1:
        volume = np.clip(volume, p1, p99)
        volume = (volume - p1) / (p99 - p1)
    else:
        volume = volume - np.min(volume)
        denom = np.max(volume)
        if denom > 0:
            volume = volume / denom
            
    # Resample slice count to target_slices
    s_count = volume.shape[0]
    if s_count >= target_slices:
        indices = np.round(np.linspace(0, s_count - 1, target_slices)).astype(int)
        volume = volume[indices]
    else:
        pad_width = target_slices - s_count
        pad_before = pad_width // 2
        pad_after = pad_width - pad_before
        
        pad_arr_before = np.repeat(volume[0:1], pad_before, axis=0) if s_count > 0 else np.zeros((pad_before, target_size, target_size))
        pad_arr_after = np.repeat(volume[-1:], pad_after, axis=0) if s_count > 0 else np.zeros((pad_after, target_size, target_size))
        
        volume = np.concatenate([pad_arr_before, volume, pad_arr_after], axis=0)
        
    return volume

def select_best_series(study_df):
    """Selects the best series for Sagittal, Coronal, and Axial planes."""
    planes = ["Sagittal", "Coronal", "Axial"]
    selected = {}
    
    for plane in planes:
        plane_df = study_df[study_df["Anatomical_Plane"] == plane]
        if plane_df.empty:
            selected[plane] = None
            continue
            
        scores = []
        for idx, row in plane_df.iterrows():
            score = 0
            fs = str(row.get("Fluid_Sensitive", "")).lower()
            if fs in ["1", "true", "yes", "fluid-sensitive"]:
                score += 2
            fatsat = str(row.get("Fat_Suppression", "")).lower()
            if fatsat in ["1", "true", "yes", "fat-suppressed", "fat-suppression"]:
                score += 2
                
            slices = row.get("SliceCount", 0)
            if pd.isna(slices):
                slices = 0
            if slices >= 15:
                score += 1
                
            scores.append((score, slices, row["SeriesInstanceUID"]))
            
        scores = sorted(scores, key=lambda x: (x[0], x[1]), reverse=True)
        selected[plane] = scores[0][2]
        
    return selected

class KneeMRIDataset(Dataset):
    def __init__(self, df, df_series, raw_images_dir, target_slices=32, target_size=224):
        self.df = df.copy()
        self.df_series = df_series
        self.raw_images_dir = Path(raw_images_dir)
        self.target_slices = target_slices
        self.target_size = target_size
        self.study_uids = self.df["StudyInstanceUID"].values
        
    def __len__(self):
        return len(self.study_uids)
        
    def __getitem__(self, idx):
        study_uid = self.study_uids[idx]
        
        study_df = self.df_series[self.df_series["StudyInstanceUID"] == study_uid]
        best_series = select_best_series(study_df)
        study_folder = self.raw_images_dir / str(study_uid)
        
        volumes = {}
        for plane in ["Sagittal", "Coronal", "Axial"]:
            series_uid = best_series.get(plane)
            if series_uid is not None:
                series_dir = study_folder / str(series_uid)
                if series_dir.exists():
                    vol = process_series(series_dir, self.target_slices, self.target_size)
                    if vol is not None:
                        volumes[plane.lower()] = vol
                        
            if plane.lower() not in volumes:
                volumes[plane.lower()] = np.zeros((self.target_slices, self.target_size, self.target_size), dtype=np.float32)
                
        sagittal = volumes["sagittal"]
        coronal = volumes["coronal"]
        axial = volumes["axial"]
        
        return {
            "sagittal": torch.tensor(sagittal, dtype=torch.float32),
            "coronal": torch.tensor(coronal, dtype=torch.float32),
            "axial": torch.tensor(axial, dtype=torch.float32),
            "study_uid": str(study_uid)
        }

# ----------------- MODEL ARCHITECTURE -----------------

class ImageBranch(nn.Module):
    def __init__(self, backbone_name="efficientnet_b0", pretrained=False, feature_dim=128):
        super().__init__()
        self.backbone = timm.create_model(
            backbone_name,
            pretrained=pretrained,
            num_classes=0,
            in_chans=3
        )
        
        # Determine feature dim dynamically
        with torch.no_grad():
            dummy = torch.zeros(1, 3, 224, 224)
            dummy_out = self.backbone(dummy)
            backbone_features = dummy_out.shape[1]
            
        self.proj = nn.Sequential(
            nn.Linear(backbone_features, feature_dim),
            nn.LayerNorm(feature_dim),
            nn.ReLU(),
            nn.Dropout(0.2)
        )
        
    def forward(self, x):
        B, S, H, W = x.shape
        x = x.unsqueeze(2).repeat(1, 1, 3, 1, 1) # (B, S, 3, H, W)
        x = x.view(B * S, 3, H, W)
        
        # Batch slice processing to prevent CUDA OOM
        chunk_size = 16
        feats_list = []
        for i in range(0, B * S, chunk_size):
            chunk = x[i : i + chunk_size]
            feats_list.append(self.backbone(chunk))
        feats = torch.cat(feats_list, dim=0) # (B * S, C)
        
        feats = feats.view(B, S, -1)
        max_feats, _ = torch.max(feats, dim=1)
        mean_feats = torch.mean(feats, dim=1)
        
        combined_feats = max_feats + mean_feats
        return self.proj(combined_feats)

class KneeMRIModel(nn.Module):
    def __init__(self, backbone_name="efficientnet_b0", pretrained=False, image_feature_dim=128, num_classes=12):
        super().__init__()
        self.sagittal_branch = ImageBranch(backbone_name, pretrained, image_feature_dim)
        self.coronal_branch = ImageBranch(backbone_name, pretrained, image_feature_dim)
        self.axial_branch = ImageBranch(backbone_name, pretrained, image_feature_dim)
        
        total_feats = image_feature_dim * 3
        self.classifier = nn.Sequential(
            nn.Linear(total_feats, 128),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(128, num_classes)
        )
        
    def forward(self, sagittal, coronal, axial):
        sag_feats = self.sagittal_branch(sagittal)
        cor_feats = self.coronal_branch(coronal)
        axi_feats = self.axial_branch(axial)
        
        img_feats = torch.cat([sag_feats, cor_feats, axi_feats], dim=1)
        return self.classifier(img_feats)

# ----------------- MAIN INFERENCE LOOP -----------------

def main():
    # Setup paths
    raw_path = Path("/kaggle/input/competitions/rsna-knee-abnormality-detection")
    weights_path = Path("/kaggle/input/datasets/ravindrachauhan50/rsna-knee-abnormality-detection-weights")
    config_path = weights_path / "config_snapshot.yaml"
    output_path = Path("submission.csv")
    
    # Load config
    with open(config_path, 'r') as f:
        config = yaml.safe_load(f)
        
    target_cols = config.get("target_columns")
    
    # Checkpoints
    checkpoints = list(weights_path.glob("fold_*_best.pt"))
    print(f"Found {len(checkpoints)} checkpoints for inference.")
    
    # Load datasets
    test_csv_path = raw_path / "test.csv"
    test_series_csv_path = raw_path / "test_series.csv"
    sample_sub_path = raw_path / "sample_submission.csv"
    
    if not test_csv_path.exists():
        print("test.csv not found, using sample_submission.csv...")
        df_test = pd.read_csv(sample_sub_path)
        for col in target_cols:
            df_test[col] = np.nan
    else:
        df_test = pd.read_csv(test_csv_path)
        
    df_test_series = pd.read_csv(test_series_csv_path)
    
    device = torch.device("cpu")
    print(f"Inference device: {device}")
    
    # Load model ensemble
    models = []
    for ckpt_path in checkpoints:
        model = KneeMRIModel(
            backbone_name=config["image_pipeline"]["backbone"],
            pretrained=False,
            num_classes=len(target_cols)
        )
        ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
        model.load_state_dict(ckpt["model_state_dict"])
        model.to(device)
        model.eval()
        models.append(model)
        
    print(f"Loaded {len(models)} models.")
    
    # Initialize Dataset & DataLoader
    img_size = config["image_pipeline"]["image_size"]
    num_slices = config["image_pipeline"]["num_slices"]
    batch_size = config["image_pipeline"]["batch_size"]
    
    test_dataset = KneeMRIDataset(
        df=df_test,
        df_series=df_test_series,
        raw_images_dir=raw_path / "test_images",
        target_slices=num_slices,
        target_size=img_size
    )
    test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=2)
    
    # Run predictions
    study_uids = []
    predictions = []
    
    with torch.no_grad():
        for batch in test_loader:
            sag = batch["sagittal"].to(device)
            cor = batch["coronal"].to(device)
            axi = batch["axial"].to(device)
            uids = batch["study_uid"]
            
            batch_preds = np.zeros((len(uids), len(target_cols)))
            for model in models:
                logits = model(sag, cor, axi)
                probs = torch.sigmoid(logits).cpu().numpy()
                batch_preds += probs / len(models)
                
            predictions.append(batch_preds)
            study_uids.extend(uids)
            
    predictions = np.concatenate(predictions, axis=0)
    
    # Structure submission
    pred_df = pd.DataFrame(predictions, columns=target_cols)
    pred_df["StudyInstanceUID"] = study_uids
    
    ordered_cols = ["StudyInstanceUID"] + target_cols
    pred_df = pred_df[ordered_cols]
    
    # Align with sample submission
    if sample_sub_path.exists():
        sample_df = pd.read_csv(sample_sub_path)
        final_df = sample_df[["StudyInstanceUID"]].merge(pred_df, on="StudyInstanceUID", how="left")
        final_df = final_df.fillna(0.5)
    else:
        final_df = pred_df
        
    # Final check and save
    assert list(final_df.columns) == ordered_cols
    assert not final_df.isna().any().any()
    assert ((final_df[target_cols] >= 0.0) & (final_df[target_cols] <= 1.0)).all().all()
    
    final_df.to_csv(output_path, index=False)
    print(f"Submission saved to {output_path}")

if __name__ == "__main__":
    main()
