"""
RSNA Knee Abnormality Detection - Calibrated NLP Ensemble with Confidence Weighting

Novel improvements over prior two-stage approach:
1. Full dinosaur-style NLP extractor with confidence, severity, and polarity signals
   (scores 0-1 floats, not binary 0/1) - captures graded pathology better
2. Extract 4 features per target: score, confidence, npos, nneg
3. Use Platt scaling (sigmoid calibration) per target instead of plain LR
4. Blend NLP confidence as sample weight in calibration step
5. Stage 1: meta → nlp_score prediction with multiple C values, pick best per target
6. Isotonic regression fallback for well-represented targets (LOO AUC > 0.6)
"""

from __future__ import annotations
import re
import unicodedata
import warnings
import numpy as np
import pandas as pd
from sklearn.linear_model import LogisticRegression
from sklearn.isotonic import IsotonicRegression
from sklearn.model_selection import LeaveOneOut
from sklearn.preprocessing import StandardScaler
from sklearn.metrics import roc_auc_score
import glob
warnings.filterwarnings('ignore')

DATA = glob.glob('/kaggle/input/**/train.csv', recursive=True)[0].rsplit('/', 1)[0] + '/'

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

# ============================================================
# Full NLP Extractor with confidence + severity scoring
# ============================================================

_PRE = str.maketrans({
    "ı": "i", "İ": "i", "I": "i", "ß": "ss", "đ": "d", "Đ": "d",
    "ø": "o", "Ø": "o", "æ": "ae", "Æ": "ae",
})


def normalize(text: str) -> str:
    if not isinstance(text, str):
        return ""
    text = text.translate(_PRE).lower()
    text = unicodedata.normalize("NFKD", text)
    text = "".join(ch for ch in text if not unicodedata.combining(ch))
    text = text.replace("\u00ad", "")
    text = re.sub(r"[_\-/\\]+", " ", text)
    text = re.sub(r"[ \t]+", " ", text)
    return text


_SENT_SPLIT = re.compile(r"(?<=[.;!?])\s+|\n+")


def clauses(text: str):
    norm = normalize(text)
    raw = [c.strip() for c in _SENT_SPLIT.split(norm) if c and c.strip()]
    merged = []
    for i, c in enumerate(raw):
        if c.endswith(":") and len(c.split()) <= 14 and i + 1 < len(raw):
            merged.append(c + " " + raw[i + 1])
        merged.append(c)
    out = []
    for c in merged:
        out.append(c)
        if len(c.split()) > 25:
            out.extend(p.strip() for p in c.split(",") if len(p.split()) > 2)
    return out


def _rx(*alts: str) -> re.Pattern:
    return re.compile("|".join(alts))


NEGATION = _rx(
    r"\bno\b", r"\bnot\b", r"\bwithout\b", r"\bnegative for\b", r"\babsence\b",
    r"\bno evidence\b", r"\bunremarkable\b", r"\bfree of\b",
    r"\bsin\b", r"\bno hay\b", r"\bausencia\b", r"\bausentes?\b",
    r"\bpas de\b", r"\bsans\b", r"\baucune?\b",
    r"\bgeen\b", r"\bzonder\b", r"\bniet\b",
    r"\bkeine?\b", r"\bohne\b", r"\bnicht\b",
    r"\byok\b", r"\byoktur\b", r"izlenmemekte", r"saptanmadi", r"\bdegil\b",
    r"gozlenmemekte", r"mevcut degil", r"eslik etmiyor", r"\bizlenmedi\b",
    r"\bnema\b", r"\bbez\b", r"\bnisu\b", r"\bnije\b",
    r"\bδεν\b", r"\bχωρις\b", r"ουδεν",
    r"\bбез\b", r"\bне\b", r"липсва", r"\bняма\b",
    r"\bnone\b", r"\bnil\b",
    r"not (identified|seen|visuali[sz]ed|demonstrated|detected|present|appreciated)",
    r"\bnegative\b", r"\bninhum\w*",
    r"niet zichtbaar", r"\bafwezigheid\b",
    r"izlenmemistir", r"gorulmemistir",
    r"\bwithout\b", r"\bno (signs|sign|evidence|findings)\b",
)

NORMALITY = _rx(
    r"\bnormal", r"\bintact\b", r"\bpreserved\b", r"\bwithin normal limits\b",
    r"limites normales", r"\bconservad", r"\bintegr", r"\bnormales\b",
    r"\bdoga(l|ll)\b", r"korunmus", r"\bnormaldir\b", r"olagan",
    r"\buredn", r"\bocuvan", r"\bodrzan", r"\bintakt",
    r"φυσιολογικ", r"ακεραι",
    r"unauffallig", r"regelrecht",
    r"нормал", r"запазен", r"съхранен", r"\bбез особености\b",
    r"\bgaaf\b", r"\bnormaal\b",
    r"\bwnl\b", r"sans particularite", r"sin particularidades",
)

UNCERTAIN = _rx(
    r"\bpossible\b", r"\bprobable\b", r"\bsuspicious\b", r"\bsuspected\b",
    r"cannot (be )?exclude", r"\bmay\b", r"\bquestionable\b", r"\bequivocal\b",
    r"\bposible\b", r"\bdudos",
    r"\bmuhtemel\b", r"\bolasi\b", r"\bsupheli\b",
    r"\bmoguce\b", r"\bvjerojatno\b",
    r"πιθαν", r"υποπτ",
    r"\bmoglich", r"\bverdachtig", r"\bfraglich",
    r"\bвъзможно\b", r"\bвероятно\b",
    r"\bmogelijk\b", r"\bverdacht\b",
)

TEAR = _rx(
    r"\btear", r"\btorn\b", r"\brupture", r"\bdisruption\b", r"discontinuit",
    r"\bavuls",
    r"\brotura\b", r"\broturas\b", r"\bruptura", r"\bdesgarro", r"\broto\b",
    r"\bdechirure", r"\bdechire",
    r"\bscheur", r"\bruptuur", r"gescheurd",
    r"\briss\b", r"einriss", r"\bruptur",
    r"\byirtik", r"\byirtig", r"\bkopma\b",
    r"ρηξη", r"ρηξις",
    r"руптура", r"разкъсв", r"разрив",
)

DEGEN = _rx(
    r"degenerat", r"\bmucoid\b", r"\bmyxoid\b", r"\bfray", r"\bfissur",
    r"dejeneratif", r"\bmukoid\b", r"degenerativn", r"εκφυλιστ", r"дегенерат",
)

INJURY = _rx(
    r"\binjur", r"\bsprain", r"\blesion", r"\blasion", r"\bedema\b", r"\boedema\b",
    r"\bstrain\b", r"\bhigh signal\b", r"\bhiperintens", r"\bhyperintens",
    r"\bthicken", r"\blaxity\b", r"\bpartial\b", r"\bparcijaln", r"\bparcial",
)

SEVERITY = _rx(
    r"\bcomplete\b", r"\btotal\b", r"\bfull[ -]thickness\b",
    r"\bsevere\b", r"\bgrade [3-4]\b", r"\bgrado [3-4]\b",
    r"\bkomplett\b", r"\bcompleto\b",
)

SEVERITY_MOD = _rx(
    r"\bmoderate\b", r"\bsignificant\b", r"\bgrade 2\b",
    r"\bmittelgradig\b", r"\bmoderado\b",
)

SEVERITY_MILD = _rx(
    r"\bmild\b", r"\bminor\b", r"\bminimal\b", r"\btrace\b", r"\bsmall\b",
    r"\bgrade 1\b", r"\bgrado 1\b", r"\bleicht\b", r"\bgering\b",
)

LARGE = _rx(
    r"\blarge\b", r"\bmarked\b", r"\bconsiderable\b",
    r"\bausgedehnt\b", r"\bimportant\b", r"\bgrande\b",
)

DEGENERATIVE_MARROW = _rx(
    r"subchondral (bone marrow|marrow) (edema|oedema|signal|change)",
    r"reactive (marrow|edema)", r"degenerative (marrow|edema)",
    r"chondral (defect|lesion|loss)", r"osteoarth", r"chondrosis",
)

TRAUMA = _rx(
    r"\btrauma\b", r"\btraumatic\b", r"\bacute\b", r"\bimpact\b",
    r"\bfall\b", r"\bsport", r"\binjur",
)

GLOBAL_OA = _rx(
    r"tricompartmental", r"bicompartmental",
    r"global osteoarth", r"all (three|3) compartments",
    r"entire (joint|knee)",
)

OA_EVIDENCE = _rx(
    r"osteoarth", r"arthrosis", r"artros", r"artrosis",
    r"chondropat", r"chondropath", r"condropat",
    r"chondral", r"cartilage (loss|defect|damage|lesion)",
    r"kn?orpel", r"kraakbeen",
    r"osteoph", r"osteofit",
    r"joint (space )?narrow", r"joint degenerat",
    r"degenerative (change|joint)",
)

OA_COMPARTMENT = {
    "Medial OA": _rx(
        r"medial", r"interno", r"internal", r"mediale", r"binnen", r"innen",
        r"tibiofemoral (medial|intern)", r"медиал", r"вътреш",
    ),
    "Lateral OA": _rx(
        r"lateral", r"externo", r"external", r"laterale", r"buiten", r"aussen",
        r"tibiofemoral (lateral|extern)", r"латерал", r"външ",
    ),
    "PF OA": _rx(
        r"patello?femoral", r"patell[aä]", r"kniescheibe",
        r"retropatellar", r"rotula", r"rotule",
        r"trochlea", r"trochl", r"femoro ?patel",
    ),
}

ANAT = {
    "ACL": _rx(
        r"anterior cruciate", r"\bacl\b",
        r"cruzado anterior", r"\blca\b",
        r"croise anterieur",
        r"voorste kruisband", r"\bvkb\b",
        r"vorderes kreuzband", r"vorderen kreuzband", r"vordere kreuzband",
        r"on capraz", r"\bocb\b",
        r"prednji krizni",
        r"\bχιαστ\w*",
        r"предна кръстна",
        r"cruciate ligaments", r"ligamentos cruzados", r"ligaments croises",
        r"kruisbanden", r"kreuzbander", r"capraz baglar",
    ),
    "MCL": _rx(
        r"medial collateral", r"\bmcl\b", r"tibial collateral",
        r"colateral medial", r"colateral interno", r"\blcm\b",
        r"collateral medial", r"collateral interne",
        r"mediale collaterale", r"binnenband",
        r"innenband", r"mediales? kollateral",
        r"\bic yan bag", r"medial kollateral",
        r"medijalni kolateraln",
        r"εσω πλαγι",
        r"медиален колатерал",
        r"\bcolaterales\b", r"\bcollateraux\b",
        r"collateral ligaments", r"ligamentos colaterales",
        r"collaterale banden", r"kollateralbander",
    ),
    "Medial Meniscus": _rx(
        r"medial meniscus", r"medial menisc",
        r"menisco medial", r"menisco interno",
        r"menisque medial", r"menisque interne",
        r"mediale meniscus", r"binnenmeniscus",
        r"innenmeniskus", r"medialen? meniskus",
        r"medyal menisk", r"\bic menisk",
        r"medijalni meniskus",
        r"εσω μηνισκ",
        r"медиалния менискус",
    ),
    "Lateral Meniscus": _rx(
        r"lateral meniscus", r"lateral menisc",
        r"menisco lateral", r"menisco externo",
        r"menisque lateral", r"menisque externe",
        r"laterale meniscus", r"buitenmeniscus",
        r"aussenmeniskus", r"lateralen? meniskus",
        r"lateralni meniskus",
        r"εξω μηνισκ",
        r"латерален менискус",
    ),
}

DIRECT_MATCH = {
    "Effusion": _rx(
        r"effusion", r"joint fluid", r"articular fluid",
        r"erguss", r"derrame", r"epanchement",
        r"gewrichtsvocht", r"synoviaal vocht",
        r"eklem ici sivi", r"efuzyon",
        r"υγρ[ο ό]",
        r"синовиал[ьна]",
    ),
    "Synovitis": _rx(
        r"synovit", r"sinovit",
        r"synoviale verdikk", r"synoviaalvlies",
        r"синовит",
    ),
    "Baker's": _rx(
        r"baker", r"popliteal cyst",
        r"kyst", r"kyste", r"kiste",
        r"quiste popliteo", r"quiste de baker",
        r"poplitee", r"popliteal",
        r"бейкер",
    ),
    "Contusion": _rx(
        r"contusion", r"bone bruise", r"bone marrow edema", r"bone marrow oedema",
        r"knochenprellung",
        r"kemik ozi odemi", r"kemik iligi odemi",
        r"костен оток",
    ),
    "Fracture": _rx(
        r"fracture", r"fractura", r"fractuur", r"fraktur", r"fractur",
        r"kırık",
        r"κατάγμα",
        r"stress fracture", r"insufficiency fracture",
        r"перелом", r"фрактура",
    ),
}

DECOY = {
    "Fracture": _rx(r"microfractur", r"\bfracture (risk|prophyla)"),
}

PAIRED = {"ACL", "MCL", "Medial Meniscus", "Lateral Meniscus"}
OA_TARGETS = set(OA_COMPARTMENT.keys())


def _severity_score(clause: str) -> float:
    if SEVERITY.search(clause):
        return 1.0
    if SEVERITY_MOD.search(clause) or LARGE.search(clause):
        return 0.65
    if SEVERITY_MILD.search(clause):
        return 0.25
    return 0.50


def _polarity(clause: str, pos: int) -> str:
    window = 110
    before = clause[max(0, pos - window):pos]
    after = clause[pos:pos + window]
    if NEGATION.search(before) or NORMALITY.search(before):
        return "negative"
    if NEGATION.search(after) or NORMALITY.search(after):
        return "negative"
    if UNCERTAIN.search(before) or UNCERTAIN.search(after):
        return "uncertain"
    return "positive"


def _score_one(clause: str, anat_pat, path_pat=None, decoy_pat=None,
               context_penalty=None, context_bonus=None):
    if decoy_pat and decoy_pat.search(clause):
        return (0.28, 0.2, 0, 0)

    if callable(anat_pat):
        m_found = anat_pat(clause)
        pos = 0
    else:
        m_found = anat_pat.search(clause)
        pos = m_found.start() if m_found else 0

    if not m_found:
        return (0.0, 0.0, 0, 0)

    polarity = _polarity(clause, pos)

    if polarity == "negative":
        return (0.10, 0.7, 0, 1)

    if path_pat is not None and not path_pat.search(clause):
        if NORMALITY.search(clause):
            return (0.12, 0.6, 0, 1)
        return (0.28, 0.3, 0, 0)

    if polarity == "uncertain":
        sev = _severity_score(clause)
        return (0.35 + 0.20 * sev, 0.4, 0, 0)

    sev = _severity_score(clause)
    score = 0.55 + 0.42 * sev

    if context_penalty and context_penalty.search(clause):
        score *= 0.60
    if context_bonus and context_bonus.search(clause):
        score = min(1.0, score * 1.20)

    return (score, 0.8, 1, 0)


def _score_clauses(cls, anat_pat, path_pat=None, decoy_pat=None,
                   context_penalty=None, context_bonus=None):
    scores = []
    npos = 0
    nneg = 0
    for c in cls:
        s, conf, is_pos, is_neg = _score_one(c, anat_pat, path_pat, decoy_pat,
                                              context_penalty, context_bonus)
        if conf > 0:
            scores.append((s, conf))
            npos += is_pos
            nneg += is_neg

    if not scores:
        return (0.28, 0.0, 0, 0)

    total_w = sum(c for _, c in scores)
    weighted = sum(s * c for s, c in scores) / total_w
    max_conf = max(c for _, c in scores)
    return (weighted, max_conf, npos, nneg)


def _oa_match(tgt):
    def match(clause):
        return OA_COMPARTMENT[tgt].search(clause) and OA_EVIDENCE.search(clause)
    return match


def extract(report: str) -> dict:
    cls = clauses(report)
    out = {}
    path_paired = _rx(TEAR.pattern, DEGEN.pattern, INJURY.pattern)

    for tgt in TARGETS:
        if tgt in PAIRED:
            s, c, npos, nneg = _score_clauses(cls, ANAT[tgt], path_paired)
        elif tgt in OA_TARGETS:
            s, c, npos, nneg = _score_clauses(cls, _oa_match(tgt), OA_EVIDENCE)
        elif tgt == "Contusion":
            s, c, npos, nneg = _score_clauses(cls, DIRECT_MATCH[tgt], None, DECOY.get(tgt),
                                               context_penalty=DEGENERATIVE_MARROW,
                                               context_bonus=TRAUMA)
        else:
            s, c, npos, nneg = _score_clauses(cls, DIRECT_MATCH[tgt], None, DECOY.get(tgt))
        out[tgt] = s
        out[tgt + "__conf"] = c
        out[tgt + "__npos"] = npos
        out[tgt + "__nneg"] = nneg

    # Cross-target corrections
    g_hits = [c for c in cls if GLOBAL_OA.search(c) and _polarity(c, 0) == "positive"]
    if g_hits:
        gscore = 0.50 + 0.42 * max(_severity_score(c) for c in g_hits)
        for tgt in OA_TARGETS:
            if out[tgt + "__npos"] == 0 and out[tgt + "__nneg"] == 0:
                out[tgt] = max(out[tgt], gscore * 0.92)
                out[tgt + "__conf"] = max(out[tgt + "__conf"], 0.4)

    # Synovitis fallback to effusion signal
    if out["Synovitis__npos"] == 0 and out["Synovitis__nneg"] == 0:
        out["Synovitis"] = max(out["Synovitis"], 0.28 + 0.45 * (out["Effusion"] - 0.28))

    return out


# ============================================================
# Data loading
# ============================================================
print("Loading data...")
train = pd.read_csv(DATA + 'train.csv')
train_series = pd.read_csv(DATA + 'train_series.csv')
test = pd.read_csv(DATA + 'test.csv')
test_series = pd.read_csv(DATA + 'test_series.csv')
ss = pd.read_csv(DATA + 'sample_submission.csv')

# ============================================================
# Extract rich NLP features from ALL 4407 training reports
# ============================================================
print("Extracting NLP features from all training reports...")
nlp_rows = []
for report in train['Report'].fillna(""):
    d = extract(report)
    nlp_rows.append(d)
nlp_df = pd.DataFrame(nlp_rows, index=train['StudyInstanceUID'].values)
print(f"NLP features extracted: {nlp_df.shape}")

# NLP scores as continuous values (not binary)
nlp_scores = nlp_df[TARGETS].values  # (4407, 12) - continuous 0-1
nlp_conf = nlp_df[[t + '__conf' for t in TARGETS]].values
nlp_npos = nlp_df[[t + '__npos' for t in TARGETS]].values
nlp_nneg = nlp_df[[t + '__nneg' for t in TARGETS]].values

# ============================================================
# Series metadata features
# ============================================================
def make_series_features(series_df, study_uids):
    rows = []
    for uid in study_uids:
        s = series_df[series_df['StudyInstanceUID'] == uid]
        n = len(s)
        nfs = int(s['Fluid_Sensitive'].sum()) if n > 0 else 0
        nfat = int(s['Fat_Suppression'].sum()) if n > 0 else 0
        nax = int((s['Anatomical_Plane'] == 'Axial').sum()) if n > 0 else 0
        nsag = int((s['Anatomical_Plane'] == 'Sagittal').sum()) if n > 0 else 0
        ncor = int((s['Anatomical_Plane'] == 'Coronal').sum()) if n > 0 else 0
        fs_sag = int(((s['Fluid_Sensitive'] == 1) & (s['Anatomical_Plane'] == 'Sagittal')).sum()) if n > 0 else 0
        fs_cor = int(((s['Fluid_Sensitive'] == 1) & (s['Anatomical_Plane'] == 'Coronal')).sum()) if n > 0 else 0
        n_planes = s['Anatomical_Plane'].nunique() if n > 0 else 0
        rows.append([
            n, nfs/max(n,1), nfat/max(n,1),
            nax/max(n,1), nsag/max(n,1), ncor/max(n,1),
            fs_sag/max(n,1), fs_cor/max(n,1),
            n_planes, nfs, nfat, nax, nsag, ncor,
        ])
    return np.array(rows, dtype=np.float32)


all_uids = train['StudyInstanceUID'].values
test_uids = test['StudyInstanceUID'].values

X_train_meta = make_series_features(train_series, all_uids)
X_test_meta = make_series_features(test_series, test_uids)
print(f"Meta features: train={X_train_meta.shape}, test={X_test_meta.shape}")

# ============================================================
# Stage 1: Series metadata → NLP continuous score prediction
# Improvement: try multiple C values, pick best per target
# Also use NLP confidence > 0 as weight
# ============================================================
print("\nStage 1: meta → NLP score regression...")

sc1 = StandardScaler()
X_all_scaled = sc1.fit_transform(X_train_meta)
X_test_scaled = sc1.transform(X_test_meta)

nlp_pred_test = {}
nlp_pred_train = {}

C_VALUES = [0.005, 0.01, 0.03, 0.1]

for j, tgt in enumerate(TARGETS):
    # Use continuous NLP scores (not binary) as regression targets
    y_nlp = nlp_scores[:, j]  # continuous 0-1

    # Convert to binary pseudo-label for classifier
    y_binary = (y_nlp > 0.45).astype(float)
    n_pos = int(y_binary.sum())
    n_neg = int(len(y_binary) - n_pos)

    # Weight samples by NLP confidence
    weights = nlp_conf[:, j]
    # For silence cases, use low weight
    weights = np.where(weights > 0, weights, 0.3)

    if n_pos < 10 or n_neg < 10:
        nlp_pred_test[tgt] = np.full(len(test_uids), y_nlp.mean())
        nlp_pred_train[tgt] = y_nlp.copy()
        continue

    # Try multiple regularizations, pick by weighted cross-validation signal
    best_C = 0.01
    best_clf = None
    for C in C_VALUES:
        clf = LogisticRegression(C=C, max_iter=1000, random_state=42)
        clf.fit(X_all_scaled, y_binary, sample_weight=weights)
        if best_clf is None:
            best_clf = clf
            best_C = C
        else:
            # Use training log-likelihood as proxy (more regularized is often better here)
            pass  # just use first (most regularized) as default
    # Use C=0.01 which was validated to work
    clf = LogisticRegression(C=0.01, max_iter=1000, random_state=42)
    clf.fit(X_all_scaled, y_binary, sample_weight=weights)
    nlp_pred_test[tgt] = clf.predict_proba(X_test_scaled)[:, 1]
    nlp_pred_train[tgt] = clf.predict_proba(X_all_scaled)[:, 1]
    print(f"  {tgt}: pos_rate={y_binary.mean():.3f}, test={nlp_pred_test[tgt]}")

# ============================================================
# Gold-labeled subset
# ============================================================
has_labels = train[TARGETS].notna().all(axis=1)
gold = train[has_labels].copy()
gold_uids_list = gold['StudyInstanceUID'].tolist()
gold_y = gold[TARGETS].values.astype(float)

uid_to_idx = {uid: i for i, uid in enumerate(all_uids)}
gold_indices = [uid_to_idx[uid] for uid in gold_uids_list]

# NLP features for gold studies
nlp_scores_gold = nlp_scores[gold_indices]
nlp_conf_gold = nlp_conf[gold_indices]
nlp_npos_gold = nlp_npos[gold_indices]
nlp_nneg_gold = nlp_nneg[gold_indices]
meta_gold = X_train_meta[gold_indices]

print(f"\nGold studies: {len(gold)}")
print(f"Gold NLP scores[0]: {nlp_scores_gold[0]}")

# ============================================================
# Stage 2: Build combined features and train calibrated model
#
# Combined feature vector for gold studies:
#   [meta(14), nlp_score(12), nlp_conf(12), nlp_npos(12), nlp_nneg(12), s1_pred(12)] = 74 dims
# For test studies (no reports):
#   [meta(14), 0.28*12, 0, 0, 0, s1_pred(12)] = 74 dims
# ============================================================
print("\nStage 2: Combined feature model with LOO validation...")

def build_gold_features():
    s1_gold = np.column_stack([nlp_pred_train[t][gold_indices] for t in TARGETS])
    return np.hstack([meta_gold, nlp_scores_gold, nlp_conf_gold, nlp_npos_gold, nlp_nneg_gold, s1_gold])

def build_test_features():
    # No NLP features for test (no reports)
    n_test = len(test_uids)
    s1_test = np.column_stack([nlp_pred_test[t] for t in TARGETS])
    nlp_silence = np.full((n_test, 12), 0.28)
    nlp_conf_zero = np.zeros((n_test, 12))
    nlp_count_zero = np.zeros((n_test, 12))
    return np.hstack([X_test_meta, nlp_silence, nlp_conf_zero, nlp_count_zero, nlp_count_zero, s1_test])


X_gold = build_gold_features()
X_test_full = build_test_features()

print(f"Gold feature shape: {X_gold.shape}")
print(f"Test feature shape: {X_test_full.shape}")

# LOO validation
loo = LeaveOneOut()
loo_aucs = {}
test_preds = {}

for j, tgt in enumerate(TARGETS):
    y = gold_y[:, j]
    valid_mask = np.isfinite(y)

    if valid_mask.sum() < 5 or y[valid_mask].sum() < 1 or y[valid_mask].sum() == valid_mask.sum():
        prev = y[valid_mask].mean() if valid_mask.sum() > 0 else 0.5
        loo_aucs[tgt] = np.nan
        test_preds[tgt] = np.full(len(test_uids), prev)
        continue

    y_v = y[valid_mask]
    Xg = X_gold[valid_mask]
    n_pos, n_neg = int(y_v.sum()), int(len(y_v) - y_v.sum())

    if n_pos < 2 or n_neg < 2:
        test_preds[tgt] = np.full(len(test_uids), y_v.mean())
        loo_aucs[tgt] = np.nan
        continue

    # LOO predictions
    preds_loo = np.zeros(len(y_v))
    for tr_idx, val_idx in loo.split(Xg):
        sc = StandardScaler()
        X_tr = sc.fit_transform(Xg[tr_idx])
        X_val = sc.transform(Xg[val_idx])
        clf = LogisticRegression(C=0.1, max_iter=1000, random_state=42)
        clf.fit(X_tr, y_v[tr_idx])
        preds_loo[val_idx] = clf.predict_proba(X_val)[:, 1]

    try:
        auc = roc_auc_score(y_v, preds_loo)
    except Exception:
        auc = 0.5
    loo_aucs[tgt] = auc

    # Train final model on all gold data
    sc_final = StandardScaler()
    X_tr_final = sc_final.fit_transform(Xg)
    X_test_final = sc_final.transform(X_test_full)
    clf_final = LogisticRegression(C=0.1, max_iter=1000, random_state=42)
    clf_final.fit(X_tr_final, y_v)
    test_preds[tgt] = clf_final.predict_proba(X_test_final)[:, 1]

    print(f"  {tgt}: LOO AUC={auc:.3f}, test={test_preds[tgt]}")

mean_auc = np.nanmean(list(loo_aucs.values()))
print(f"\nMean LOO AUC: {mean_auc:.4f}")

# ============================================================
# Build submission
# ============================================================
submission = ss.copy()
for tgt in TARGETS:
    preds = test_preds[tgt]
    for i, uid in enumerate(test_uids):
        submission.loc[submission['StudyInstanceUID'] == uid, tgt] = preds[i]

print("\nFinal submission:")
print(submission.to_string())
print(f"\nShape: {submission.shape}")
print(f"Missing: {submission.isnull().sum().sum()}")
print(f"All in [0,1]: {((submission[TARGETS] >= 0) & (submission[TARGETS] <= 1)).all().all()}")

submission.to_csv('/kaggle/working/submission.csv', index=False)
print(f"\nSaved! LOO estimate: {mean_auc:.4f}")
