import os
import glob
import numpy as np
import pandas as pd
from pathlib import Path
from sklearn.ensemble import RandomForestRegressor
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.metrics import roc_auc_score
from sklearn.model_selection import cross_val_predict
import pydicom
from tqdm import tqdm
import warnings
from scipy.sparse import hstack

warnings.filterwarnings("ignore")

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

train = pd.read_csv("/kaggle/input/rsna-knee-abnormality-detection/train.csv")
train_series = pd.read_csv("/kaggle/input/rsna-knee-abnormality-detection/train_series.csv")
test = pd.read_csv("/kaggle/input/rsna-knee-abnormality-detection/test.csv")
test_series = pd.read_csv("/kaggle/input/rsna-knee-abnormality-detection/test_series.csv")

DATA_ROOT = Path("/kaggle/input/rsna-knee-abnormality-detection")

def aggregate_series(series_df):
    agg = series_df.groupby("StudyInstanceUID").agg(
        num_series=("SeriesInstanceUID", "count"),
        fluid_sensitive_mean=("Fluid_Sensitive", "mean"),
        fat_suppression_mean=("Fat_Suppression", "mean"),
    ).reset_index()
    planes = series_df.pivot_table(
        index="StudyInstanceUID",
        columns="Anatomical_Plane",
        values="SeriesInstanceUID",
        aggfunc="count",
        fill_value=0
    ).reset_index()
    planes.columns.name = None
    planes = planes.rename(columns={
        "Axial": "axial_count",
        "Coronal": "coronal_count",
        "Sagittal": "sagittal_count"
    })
    return agg.merge(planes, on="StudyInstanceUID", how="left")

def process_study_images(study_uid, series_df, max_slices_per_series=2):
    series_uids = series_df["SeriesInstanceUID"].tolist()
    features = []
    for series_uid in series_uids:
        series_dir = DATA_ROOT / "train_series" / study_uid / series_uid
        if not series_dir.exists():
            continue
        dicom_files = sorted(series_dir.glob("*.dcm"))
        if not dicom_files:
            continue
        selected = dicom_files[::max(len(dicom_files)//max_slices_per_series, 1)][:max_slices_per_series]
        for dcm_path in selected:
            try:
                dcm = pydicom.dcmread(dcm_path, stop_before_pixels=True)
                rows = int(dcm.Rows) if hasattr(dcm, 'Rows') else 0
                cols = int(dcm.Columns) if hasattr(dcm, 'Columns') else 0
                features.append({"rows": rows, "cols": cols})
            except Exception:
                features.append({"rows": 0, "cols": 0})
    if not features:
        return {"img_rows_mean": 0, "img_cols_mean": 0, "img_slices": 0}
    df = pd.DataFrame(features)
    return {
        "img_rows_mean": df["rows"].mean(),
        "img_cols_mean": df["cols"].mean(),
        "img_slices": len(features)
    }

print("Extracting series features...")
train_feat = aggregate_series(train_series)
test_feat = aggregate_series(test_series)

print("Extracting image metadata features for train...")
train_img_feats = []
for study_uid, group in tqdm(train_series.groupby("StudyInstanceUID"), total=train_series["StudyInstanceUID"].nunique()):
    feats = process_study_images(study_uid, group)
    feats["StudyInstanceUID"] = study_uid
    train_img_feats.append(feats)
train_img_df = pd.DataFrame(train_img_feats)

print("Extracting image metadata features for test...")
test_img_feats = []
for study_uid, group in tqdm(test_series.groupby("StudyInstanceUID"), total=test_series["StudyInstanceUID"].nunique()):
    feats = process_study_images(study_uid, group)
    feats["StudyInstanceUID"] = study_uid
    test_img_feats.append(feats)
test_img_df = pd.DataFrame(test_img_feats)

print("Merging features...")
train_all = train.merge(train_feat, on="StudyInstanceUID", how="left").merge(train_img_df, on="StudyInstanceUID", how="left")
test_all = test.merge(test_feat, on="StudyInstanceUID", how="left").merge(test_img_df, on="StudyInstanceUID", how="left")

feature_cols = [
    "num_series", "fluid_sensitive_mean", "fat_suppression_mean",
    "axial_count", "coronal_count", "sagittal_count",
    "img_rows_mean", "img_cols_mean", "img_slices"
]

for col in feature_cols:
    median_val = train_all[col].median()
    train_all[col] = train_all[col].fillna(median_val)
    test_all[col] = test_all[col].fillna(median_val)

labeled_mask = train_all[label_cols].notna().all(axis=1)
train_labeled = train_all[labeled_mask].copy()
train_unlabeled = train_all[~labeled_mask].copy()

print("Training text model for pseudo-labeling...")
vectorizer = TfidfVectorizer(max_features=5000, ngram_range=(1,2), stop_words="english")
train_all["Report"] = train_all["Report"].fillna("")
test_all["Report"] = test_all["Report"].fillna("")

X_text_labeled = vectorizer.fit_transform(train_labeled["Report"])
X_text_unlabeled = vectorizer.transform(train_unlabeled["Report"])

X_labeled = hstack([X_text_labeled, train_labeled[feature_cols].values])
X_unlabeled = hstack([X_text_unlabeled, train_unlabeled[feature_cols].values])

clf = RandomForestRegressor(n_estimators=200, max_depth=5, random_state=42)
clf.fit(X_labeled, train_labeled[label_cols].astype(float))

unlabeled_proba = clf.predict(X_unlabeled)
if isinstance(unlabeled_proba, list):
    unlabeled_proba = np.column_stack([p for p in unlabeled_proba])

max_conf = unlabeled_proba.max(axis=1)
high_conf_mask = (max_conf > 0.3) & (max_conf < 0.7)
pseudo_labeled = train_unlabeled[high_conf_mask].copy()
for i, col in enumerate(label_cols):
    pseudo_labeled[col] = unlabeled_proba[high_conf_mask, i]

print(f"Pseudo-labeled studies: {high_conf_mask.sum()}")

combined = pd.concat([train_labeled, pseudo_labeled], ignore_index=True)
for col in feature_cols:
    median_val = combined[col].median()
    combined[col] = combined[col].fillna(median_val)

X_combined_text = vectorizer.transform(combined["Report"].fillna(""))
X_combined = hstack([X_combined_text, combined[feature_cols].values])
y_combined = combined[label_cols].astype(float)

print("Training final model...")
clf_final = RandomForestRegressor(n_estimators=200, max_depth=5, random_state=42)
clf_final.fit(X_combined, y_combined)

X_test_text = vectorizer.transform(test_all["Report"].fillna(""))
X_test = hstack([X_test_text, test_all[feature_cols].values])
test_proba = clf_final.predict(X_test)
if isinstance(test_proba, list):
    test_proba = np.column_stack([p for p in test_proba])

submission = test_all[["StudyInstanceUID"]].copy()
for i, col in enumerate(label_cols):
    submission[col] = test_proba[:, i]

submission[label_cols] = submission[label_cols].clip(0.01, 0.99)
submission.to_csv("submission.csv", index=False)
print(submission.head())
