#!/usr/bin/env python

## Imports

import os
import warnings
import numpy as np
import pandas as pd
import pickle
import optuna
import xgboost as xgb
import shap
import matplotlib.pyplot as plt
from collections import Counter, defaultdict
from scipy.stats import mannwhitneyu
from statsmodels.stats.multitest import multipletests
from IPython.display import display
from joblib import Parallel, delayed
from matplotlib import gridspec

# sklearn
from sklearn.experimental import enable_iterative_imputer
from sklearn.impute import IterativeImputer
from sklearn.linear_model import BayesianRidge, LogisticRegression
from sklearn.preprocessing import StandardScaler
from sklearn.svm import SVC
from sklearn.ensemble import RandomForestClassifier
from sklearn.neural_network import MLPClassifier
from sklearn.metrics import (
    roc_auc_score,
    f1_score,
    matthews_corrcoef,
    confusion_matrix,
    ConfusionMatrixDisplay,
    classification_report,
)
from sklearn.model_selection import (
    StratifiedKFold,
    cross_val_predict,
    cross_val_score,
    check_cv,
)
from sklearn.pipeline import Pipeline
from sklearn.base import clone

warnings.filterwarnings("ignore")

print("imports ready")


## Data load and prep

INPUT_FILE = "data/input.xlsx"
SHEET_NAME = 0
ID_COL = "src_subject_id"
TARGETS = ["Internalizing problems", "Externalizing problems"]
OUTPUT_XLSX = "results/ML_resilient_vs_auffaellig_processed.xlsx"
UNIQUE_THRESHOLD = 10
RANDOM_STATE = 37
N_SPLITS = 5
print("\nStarting data pipeline...")

if not os.path.exists(INPUT_FILE):
    raise FileNotFoundError(f"Input file not found: {INPUT_FILE}")
df_raw = pd.read_excel(INPUT_FILE, sheet_name=SHEET_NAME, dtype=object)
print(f"Loaded {df_raw.shape} from {INPUT_FILE}")

null_strings = ["#NULL!", "#N/A", "NA", "NaN", "nan", ""] # replace excel null placeholders
df_raw = df_raw.replace(null_strings, np.nan)

if ID_COL not in df_raw.columns: # validate ID column and targets
    raise KeyError(f"ID column '{ID_COL}' missing. Available: {list(df_raw.columns)}")
for t in TARGETS:
    if t not in df_raw.columns:
        raise KeyError(f"Target '{t}' missing")

feature_cols = [c for c in df_raw.columns if c != ID_COL]
df = df_raw.copy()
for c in feature_cols:
    df[c] = pd.to_numeric(df[c], errors="coerce")

print("Missing values (top 10 columns, %):")
print((df.isna().mean() * 100).sort_values(ascending=False).head(10))

# feature classification
cat_manual = [
    "Child´s gender",
    "Child´s race",
    "Child´s ethnicity",
    "Bullying",
    "Domestic violence",
    "Traumatic loss",
    "Polyvictimization",
    "Food expenses",
    "Evicted from home",
    "Physical abuse",
    "Sexual abuse",
    "Caregiver´s employment",
    "Parents: alcohol/drug problems",
    "Accidents/natural disasters/fire",
    "Terrorism/war/community violence",
    "Medical expenses",
    "Other expenses",
]
cont_manual = [
    "Child´s age",
    "Caregiver´s education ",
    "Family income ",
    "Caregiver: internalizing problems",
    "Caregiver: externalizing problems",
    "BIS",
    "Reward Responsiveness ",
    "Drive",
    "Fun Seeking ",
    "Family conflicts",
    "School: environment",
    "School: involvement",
    "School: disengagement",
    "Neighbourhood safety",
    "PTEs (number)",
    "Parenting/acceptance",
    "Parental monitoring",
    "Prosocial behaviour",
    "Physical activity",
    "Screentime",
]

feature_types = {} # auto classify remaining features (if any)
cat_cols, cont_cols = [], []
for col in feature_cols:
    if col in TARGETS:
        feature_types[col] = "target"
        continue
    if col in cat_manual:
        feature_types[col] = "categorical"
        cat_cols.append(col)
        continue
    if col in cont_manual:
        feature_types[col] = "continuous"
        cont_cols.append(col)
        continue

    s = df[col]
    n_unique = s.dropna().nunique()
    vals = s.dropna().values
    has_decimals = False
    if len(vals) > 0:
        try:
            has_decimals = np.any(np.mod(vals.astype(float), 1) != 0)
        except Exception:
            has_decimals = False

    if n_unique > UNIQUE_THRESHOLD or has_decimals:
        feature_types[col] = "continuous"
        cont_cols.append(col)
    else:
        feature_types[col] = "categorical"
        cat_cols.append(col)

ft_df = pd.DataFrame.from_dict(feature_types, orient="index", columns=["type"])
ft_df["n_unique"] = [df[c].nunique(dropna=True) for c in ft_df.index]
display(ft_df)
print(f"\nCategorical: {len(cat_cols)} | Continuous: {len(cont_cols)}")

CATEGORICAL_PREDICTORS = cat_cols
CONTINUOUS_PREDICTORS = cont_cols
impute_cols = CATEGORICAL_PREDICTORS + CONTINUOUS_PREDICTORS + TARGETS
print(
    f"\nProcessing {len(impute_cols)} columns: {len(CATEGORICAL_PREDICTORS)} categorical, {len(CONTINUOUS_PREDICTORS)} continuous, {len(TARGETS)} targets."
)


def _sanitize_filename(s: str) -> str:
    """Convert target names to filesystem-safe strings."""
    return "".join(
        ch if ch.isalnum() or ch in (" ", "-", "_") else "_" for ch in s
    ).replace(" ", "_")


IMPUTED_DATASETS = {}
SCALED_DATASETS = {}

# ensure targets are binary (0/1)
for t in TARGETS:
    observed = df_raw[t].dropna().unique()
    if not set(observed).issubset({0, 1}):
        print(f"Target '{t}' has non-binary values: {observed} --- rounding to 0/1.")
        non_na_mask = df_raw[t].notna()
        df_raw.loc[non_na_mask, t] = (
            pd.to_numeric(df_raw.loc[non_na_mask, t], errors="coerce")
            .round()
            .astype(int)
        )

# main
for target in TARGETS:
    print(f"\n--- Processing: {target} ---")
    df_t = df_raw.copy()
    before_nrows = df_t.shape[0]
    df_t = df_t[df_t[target].notna()].copy()
    after_nrows = df_t.shape[0]
    print(f"Rows: {before_nrows} transformed to {after_nrows} after filtering missing '{target}'")

    audit = pd.DataFrame( # missingness audit
        {
            "column": impute_cols,
            "missing_before": df_raw[impute_cols].isna().sum(),
            "missing_after": df_t[impute_cols].isna().sum(),
        }
    )
    audit["still_missing"] = audit["missing_after"] > 0
    print("\nMissing values (before vs. after filtering):")
    display(audit.sort_values("missing_before", ascending=False).head(10))

    other_targets = [t for t in TARGETS if t != target]
    if other_targets:
        print(f"Ignoring other targets: {other_targets}")

    predictors = CATEGORICAL_PREDICTORS + CONTINUOUS_PREDICTORS

    df_imputed = pd.DataFrame(index=df_t.index)
    df_imputed[ID_COL] = df_t[ID_COL]
    df_imputed[target] = df_t[target].astype(int).values

    impute_predictors = [c for c in predictors if df_t[c].notna().any()]
    skip_predictors = [c for c in predictors if c not in impute_predictors]
    if skip_predictors:
        print(f"Skipping entirely missing predictors: {skip_predictors}")
        for c in skip_predictors:
            df_imputed[c] = np.nan

    outer_cv = StratifiedKFold(
        n_splits=N_SPLITS, shuffle=True, random_state=RANDOM_STATE
    )
    if len(impute_predictors) > 0:
        for train_idx, test_idx in outer_cv.split(
            df_t, df_t[target].astype(int).values
        ):
            train_idx_labels = df_t.index[train_idx]
            test_idx_labels = df_t.index[test_idx]

            X_train_raw = df_t.loc[train_idx_labels, impute_predictors].astype(float)
            imp = IterativeImputer(
                estimator=BayesianRidge(),
                random_state=RANDOM_STATE,
                max_iter=30,
                sample_posterior=False,
            )
            X_train_imp = imp.fit_transform(X_train_raw.values)
            df_imputed.loc[train_idx_labels, impute_predictors] = X_train_imp

            X_test_raw = df_t.loc[test_idx_labels, impute_predictors].astype(float)
            X_test_imp = imp.transform(X_test_raw.values)
            df_imputed.loc[test_idx_labels, impute_predictors] = X_test_imp
        print(f"Imputed {len(impute_predictors)} predictors for '{target}'.")
    else:
        print("No predictors to impute (all missing).")

    # rounding binary categorical variables
    for col in CATEGORICAL_PREDICTORS:
        original_unique = df[col].dropna().unique()
        try:
            original_set = set(np.array(original_unique, dtype=int))
        except Exception:
            original_set = set(original_unique)
        if original_set.issubset({0, 1}) and col in df_imputed.columns:
            df_imputed[col] = df_imputed[col].round().astype("Int64")

    # scaling continuous features (except those with zero-variance)
    cont_present = [c for c in CONTINUOUS_PREDICTORS if c in df_imputed.columns]
    if cont_present:
        cont_std = df_imputed[cont_present].std(ddof=0, skipna=True)
        zero_var_mask = cont_std == 0
        if zero_var_mask.any():
            print(
                f"Zero-variance columns (not scaled): {zero_var_mask[zero_var_mask].index.tolist()}"
            )
            cont_to_scale = [
                c for c in cont_present if c not in zero_var_mask[zero_var_mask].index
            ]
        else:
            cont_to_scale = cont_present
    else:
        cont_to_scale = []

    if cont_to_scale:
        scaler = StandardScaler()
        scaled_array = scaler.fit_transform(
            df_imputed[cont_to_scale].astype(float).values
        )
        df_scaled_cont = pd.DataFrame(
            scaled_array, columns=cont_to_scale, index=df_imputed.index
        )
    else:
        df_scaled_cont = pd.DataFrame(index=df_imputed.index)

    unscaled_cont = [c for c in cont_present if c not in cont_to_scale]

    df_scaled = pd.concat( # combining all data
        [
            df_imputed[[ID_COL]],
            df_scaled_cont.reindex(index=df_imputed.index),
            (
                df_imputed[unscaled_cont]
                if unscaled_cont
                else pd.DataFrame(index=df_imputed.index)
            ),
            df_imputed[[c for c in CATEGORICAL_PREDICTORS if c in df_imputed.columns]],
            df_imputed[[target]],
        ],
        axis=1,
    )

    # scaling sanity check
    if cont_to_scale:
        max_abs_mean = df_scaled[cont_to_scale].mean().abs().max()
        min_std = df_scaled[cont_to_scale].std().min()
        print(
            f"Scaling: max |mean| = {max_abs_mean:.2e} (≈0), min std = {min_std:.3f} (≈1)"
        )
    else:
        print("No continuous predictors scaled (all zero-variance or none).")
    print(f"Final shape: {df_scaled.shape}")

    # store results
    IMPUTED_DATASETS[target] = df_imputed.copy()
    SCALED_DATASETS[target] = df_scaled.copy()

    PREDICTOR_COLS = [c for c in predictors if c in df_imputed.columns]
    print(f"Predictors to inspect: {len(PREDICTOR_COLS)}")

    # compute correlations
    if cont_present:
        corr_continuous = pd.concat(
            [df_scaled_cont, df_imputed[unscaled_cont]], axis=1
        ).corr()
    else:
        corr_continuous = pd.DataFrame()
    corr_full = df_imputed[PREDICTOR_COLS].corr()

    auc_rows = []
    y = df_imputed[target].astype(int).values
    two_classes_present = len(np.unique(y)) == 2
    for feat in PREDICTOR_COLS:
        x = df_imputed[feat].values
        mask = ~pd.isna(x)
        x_masked = x[mask]
        y_masked = y[mask]
        if (not two_classes_present) or len(x_masked) == 0 or np.nanstd(x_masked) == 0:
            auc = np.nan
        else:
            try:
                if len(np.unique(y_masked)) < 2:
                    auc = np.nan
                else:
                    auc = roc_auc_score(y_masked, x_masked)
            except Exception:
                auc = np.nan
        auc_rows.append({"feature": feat, "target": target, "auc": auc})
    auc_df = pd.DataFrame(auc_rows)
    auc_pivot = auc_df.pivot(index="feature", columns="target", values="auc")
    print("\nUnivariate AUC (top 20 features):")
    display(auc_pivot.head(20))

    # save results
    corr_matrices = {"continuous_only": corr_continuous, "full_predictors": corr_full}
    os.makedirs(os.path.dirname(OUTPUT_XLSX), exist_ok=True)
    out_name = f"{os.path.splitext(OUTPUT_XLSX)[0]}___{_sanitize_filename(target)}.xlsx"
    with pd.ExcelWriter(out_name, engine="openpyxl") as writer:
        df_imputed.to_excel(writer, sheet_name="imputed_data", index=False)
        df_scaled.to_excel(writer, sheet_name="scaled_continuous", index=False)
        ft_df.to_excel(writer, sheet_name="feature_types")
        corr_matrices["continuous_only"].to_excel(writer, sheet_name="corr_continuous")
        corr_matrices["full_predictors"].to_excel(writer, sheet_name="corr_full")
        auc_pivot.to_excel(writer, sheet_name="auc_per_feature")

        pct_missing_before = df_raw[impute_cols].isna().sum().sum() / (
            df_raw[impute_cols].shape[0] * df_raw[impute_cols].shape[1]
        )
        present_after_cols = [c for c in impute_cols if c in df_imputed.columns]
        pct_missing_after = (
            (
                df_imputed[present_after_cols].isna().sum().sum()
                / (
                    df_imputed[present_after_cols].shape[0]
                    * df_imputed[present_after_cols].shape[1]
                )
            )
            if len(present_after_cols) > 0
            else 0.0
        )

        summary = pd.DataFrame(
            {
                "n_rows": [df_imputed.shape[0]],
                "n_input_features": [len(PREDICTOR_COLS)],
                "pct_missing_before": [round(pct_missing_before * 100, 2)],
                "pct_missing_after": [round(pct_missing_after * 100, 2)],
            }
        )
        summary.to_excel(writer, sheet_name="summary", index=False)

    print(f"Results for '{target}' saved to '{out_name}'.")

print("\nAll targets processed.")


## Coarse Exploratory model CV


# fallback config if preprocessing cell wasn't run
try:
    TARGETS, OUTPUT_XLSX, OPTUNA_TRIALS, N_SPLITS, RANDOM_STATE
except NameError:
    TARGETS = ["Internalizing problems", "Externalizing problems"]
    OUTPUT_XLSX = "results/exploratory_cv_results.xlsx"
    OPTUNA_TRIALS = 30
    N_SPLITS = 5
    RANDOM_STATE = 37

# ensuring preprocessed datasets are available
if "IMPUTED_DATASETS" not in globals() or "SCALED_DATASETS" not in globals():
    raise RuntimeError(
        "IMPUTED_DATASETS and SCALED_DATASETS must be in memory (run preprocessing)."
    )


def _sanitize_filename(s: str) -> str:
    # converting target names to filesystem-safe strings
    return "".join(
        ch if ch.isalnum() or ch in (" ", "-", "_") else "_" for ch in s
    ).replace(" ", "_")


def calc_scale_pos_weight(y: np.ndarray) -> float:
    # calculating scale_pos_weight = neg / pos (safe when pos==0)
    pos = int(np.sum(y == 1))
    neg = int(np.sum(y == 0))
    return float(neg) / float(pos) if pos > 0 else 1.0


def suggest_params(trial, model_name: str) -> dict:
    # suggesting hyper-parameters for the given model using optuna
    if model_name == "xgboost":
        return {
            "n_estimators": trial.suggest_int("n_estimators", 50, 500),
            "max_depth": trial.suggest_int("max_depth", 3, 8),
            "learning_rate": trial.suggest_float("learning_rate", 1e-3, 0.3, log=True),
            "subsample": trial.suggest_float("subsample", 0.6, 1.0),
            "colsample_bytree": trial.suggest_float("colsample_bytree", 0.6, 1.0),
            "reg_alpha": trial.suggest_float("reg_alpha", 1e-8, 10.0, log=True),
            "reg_lambda": trial.suggest_float("reg_lambda", 1e-8, 10.0, log=True),
        }
    if model_name == "random_forest":
        return {
            "n_estimators": trial.suggest_int("n_estimators", 50, 500),
            "max_depth": trial.suggest_int("max_depth", 3, 30),
        }
    if model_name == "logistic":
        return {"C": trial.suggest_float("C", 1e-3, 1e3, log=True)}
    if model_name == "svm":
        return {
            "C": trial.suggest_float("C", 1e-3, 1e3, log=True),
            "kernel": trial.suggest_categorical("kernel", ["rbf", "linear"]),
        }
    if model_name == "mlp":
        return {
            "n_units": trial.suggest_int("n_units", 20, 200),
            "alpha": trial.suggest_float("alpha", 1e-6, 1e-1, log=True),
        }
    return {}


def build_clf_from_params(
    model_name: str, params: dict, random_state: int = RANDOM_STATE, scale_pos_w=None
):
    # building a classifier from suggested parameters
    params_local = dict(params) if params is not None else {}
    if model_name == "xgboost":
        return xgb.XGBClassifier(
            eval_metric="logloss",
            scale_pos_weight=scale_pos_w,
            random_state=random_state,
            **params_local,
        )
    elif model_name == "random_forest":
        return RandomForestClassifier(
            class_weight="balanced", random_state=random_state, **params_local
        )
    elif model_name == "logistic":
        return LogisticRegression(
            class_weight="balanced",
            max_iter=2000,
            random_state=random_state,
            **params_local,
        )
    elif model_name == "svm":
        return SVC(probability=True, class_weight="balanced", **params_local)
    elif model_name == "mlp":
        n_units = int(params_local.pop("n_units", 100))
        alpha = float(params_local.pop("alpha", 1e-4))
        return MLPClassifier(
            hidden_layer_sizes=(n_units,),
            alpha=alpha,
            max_iter=1000,
            random_state=random_state,
            **params_local,
        )
    else:
        raise ValueError("Unsupported model name")


def optuna_objective(trial, X: np.ndarray, y: np.ndarray, model_name: str) -> float:
    # evaluating model performance using inner cv macro-f1
    params = suggest_params(trial, model_name)
    skf = StratifiedKFold(n_splits=N_SPLITS, shuffle=True, random_state=RANDOM_STATE)
    f1s = []
    for tr, val in skf.split(X, y):
        clf = build_clf_from_params(
            model_name,
            params,
            random_state=RANDOM_STATE,
            scale_pos_w=calc_scale_pos_weight(y) if model_name == "xgboost" else None,
        )
        pipe = Pipeline([("clf", clf)])  # features are prepped
        pipe.fit(X[tr], y[tr])
        pred = pipe.predict(X[val])
        f1s.append(float(f1_score(y[val], pred, average="macro")))
    return float(np.mean(f1s))


MODEL_NAMES = ["xgboost", "svm", "random_forest", "logistic", "mlp"]
exploratory_results = {}

for target in TARGETS:
    print("\n" + "=" * 70)
    print("exploratory cv -- target:", target)
    if target not in IMPUTED_DATASETS or target not in SCALED_DATASETS:
        raise RuntimeError(f"missing preprocessed datasets for '{target}'.")

    df_imputed = IMPUTED_DATASETS[target]
    df_scaled = SCALED_DATASETS[target]
    other_targets = [t for t in TARGETS if t != target]
    predictors_cont = [
        c
        for c in df_scaled.columns
        if c not in (ID_COL, target) and c not in other_targets
    ]
    predictors_cat = [
        c
        for c in df_imputed.columns
        if c not in (ID_COL, target) and c not in other_targets
    ]
    predictors_cat = [c for c in predictors_cat if c not in set(predictors_cont)]
    print(
        f" - continuous predictors: {len(predictors_cont)}  categorical predictors: {len(predictors_cat)}"
    )

    # dropping predictors with remaining nans
    usable_cont = [c for c in predictors_cont if not df_scaled[c].isna().any()]
    usable_cat = [c for c in predictors_cat if not df_imputed[c].isna().any()]
    dropped_all = [
        c
        for c in (predictors_cont + predictors_cat)
        if c not in (usable_cont + usable_cat)
    ]
    if dropped_all:
        print(
            f"   dropping {len(dropped_all)} predictors with remaining nans: {dropped_all}"
        )

    if len(usable_cont) == 0 and len(usable_cat) == 0:
        print("   no usable predictors left -- skipping.")
        exploratory_results[target] = {
            "per_model_eval": {},
            "top3": [],
            "dropped_predictors": dropped_all,
        }
        out_name = (
            f"{os.path.splitext(OUTPUT_XLSX)[0]}__{_sanitize_filename(target)}.xlsx"
        )
        os.makedirs(os.path.dirname(out_name) or ".", exist_ok=True)
        with pd.ExcelWriter(
            out_name, engine="openpyxl", mode="a" if os.path.exists(out_name) else "w"
        ) as writer:
            pd.DataFrame().to_excel(
                writer,
                sheet_name=f"exploratory_model_eval_{_sanitize_filename(target)}",
            )
        continue

    # stacking continuous and categorical predictors
    X_parts = []
    if usable_cont:
        X_parts.append(df_scaled[usable_cont].to_numpy(dtype=float))
    if usable_cat:
        X_parts.append(df_imputed[usable_cat].astype(float).to_numpy(dtype=float))
    X_all = np.hstack(X_parts)
    y = df_imputed[target].astype(int).to_numpy()
    print(
        "   final predictor count:",
        X_all.shape[1],
        "  class distribution:",
        dict(zip(*np.unique(y, return_counts=True))),
    )

    # tuning each model
    studies = {}
    per_model_eval = {}
    for m in MODEL_NAMES:
        print(f"   tuning : {m}")
        study = optuna.create_study(direction="maximize")
        study.optimize(
            lambda trial: optuna_objective(trial, X_all, y, m),
            n_trials=OPTUNA_TRIALS,
            show_progress_bar=False,
        )
        studies[m] = study
        print(f"   best macro-f1 ({m}): {round(study.best_value, 4)}")

    # evaluating models with cross-validation
    skf = StratifiedKFold(n_splits=N_SPLITS, shuffle=True, random_state=RANDOM_STATE)
    for m, study in studies.items():
        best_params = dict(study.best_params)
        clf = build_clf_from_params(
            m,
            best_params,
            scale_pos_w=calc_scale_pos_weight(y) if m == "xgboost" else None,
        )
        pipe = Pipeline([("clf", clf)])
        y_pred = cross_val_predict(pipe, X_all, y, cv=skf, method="predict")
        try:
            y_proba = cross_val_predict(pipe, X_all, y, cv=skf, method="predict_proba")[
                :, 1
            ]
        except Exception:
            y_proba = None
        mac_f1 = float(f1_score(y, y_pred, average="macro"))
        mcc = float(matthews_corrcoef(y, y_pred))
        auc = (
            float(roc_auc_score(y, y_proba))
            if (y_proba is not None and len(np.unique(y_proba)) > 1)
            else np.nan
        )
        per_model_eval[m] = {
            "macro_f1": mac_f1,
            "mcc": mcc,
            "auc": auc,
            "best_params": best_params,
        }
        print(
            f"   cv – {m}: f1={mac_f1:.3f}, mcc={mcc:.3f}, auc={auc if not np.isnan(auc) else 'NA'}"
        )

    # selecting best model and saving results
    best_model = max(per_model_eval, key=lambda k: per_model_eval[k]["macro_f1"])
    best_params = dict(per_model_eval[best_model]["best_params"])
    print("   --- selected model:", best_model)
    final_clf = build_clf_from_params(
        best_model,
        best_params,
        scale_pos_w=calc_scale_pos_weight(y) if best_model == "xgboost" else None,
    )
    final_pipe = Pipeline([("clf", final_clf)])
    final_pipe.fit(X_all, y)

    out_name = f"{os.path.splitext(OUTPUT_XLSX)[0]}__{_sanitize_filename(target)}.xlsx"
    os.makedirs(os.path.dirname(out_name) or ".", exist_ok=True)
    with pd.ExcelWriter(
        out_name,
        engine="openpyxl",
        mode="a" if os.path.exists(out_name) else "w",
        if_sheet_exists="replace",
    ) as writer:
        pd.DataFrame(per_model_eval).T.to_excel(
            writer, sheet_name=f"model_eval_{_sanitize_filename(target)}"
        )
        class_report = classification_report(
            y, final_pipe.predict(X_all), output_dict=True
        )
        pd.DataFrame(class_report).T.to_excel(
            writer, sheet_name=f"class_report_{_sanitize_filename(target)}"
        )

    exploratory_results[target] = {
        "per_model_eval": per_model_eval,
        "best_model": best_model,
        "best_params": best_params,
        "final_pipe": final_pipe,
        "dropped_predictors": dropped_all,
        "workbook": out_name,
    }

print(
    "\nExploratory modelling finished. Results in per-target workbooks and `exploratory_results`."
)


## Nested CV + SFS


# fallback config for nested cv with sfs and checkpointing
OUTPUT_XLSX_NESTED = globals().get(
    "OUTPUT_XLSX_NESTED", "results/full_nested_cv_sfs_results.xlsx"
)

# settings
RANDOM_STATE = globals().get("RANDOM_STATE", 37)
OUTER_SPLITS = 10
INNER_SPLITS = 10
OPTUNA_TRIALS = int(globals().get("OPTUNA_TRIALS", 30))
OPTUNA_N_JOBS = int(globals().get("OPTUNA_N_JOBS", 4))
N_JOBS = globals().get("N_JOBS", 1)
MAX_FEATURES_TO_SELECT = int(globals().get("MAX_FEATURES_TO_SELECT", 30))
MIN_FEATURES_TO_SELECT = int(globals().get("MIN_FEATURES_TO_SELECT", 20))
FEATURE_FREQUENCY_THRESHOLD = float(globals().get("FEATURE_FREQUENCY_THRESHOLD", 0.2))
MODELS_TO_RUN = globals().get("MODELS_TO_RUN", ["xgboost", "random_forest", "svm"])
SCORING = globals().get("SCORING", "f1_macro")
BOOT_N = int(globals().get("BOOTSTRAP_N", 10000))

# ensuring preprocessed datasets are available
if "IMPUTED_DATASETS" not in globals() or "SCALED_DATASETS" not in globals():
    raise RuntimeError(
        "IMPUTED_DATASETS and SCALED_DATASETS must exist (run preprocessing first)."
    )


def _sanitize_filename(s: str) -> str:
    # converting target names to filesystem-safe strings
    return "".join(
        ch if ch.isalnum() or ch in (" ", "-", "_") else "_" for ch in str(s)
    ).replace(" ", "_")


def calc_scale_pos_weight(y: np.ndarray) -> float:
    # calculating scale_pos_weight = neg / pos (safe when pos==0)
    pos = int(np.sum(y == 1))
    neg = int(np.sum(y == 0))
    return float(neg) / float(pos) if pos > 0 else 1.0


def suggest_params(trial, model_name, n_features, target=None):
    # suggesting hyper-parameters including k_features for sfs
    k_max = min(MAX_FEATURES_TO_SELECT, max(1, int(n_features)))
    k = trial.suggest_int("k_features", max(MIN_FEATURES_TO_SELECT, 1), k_max)

    if model_name == "xgboost":
        if target and "internal" in str(target).lower():
            return {
                "k_features": k,
                "n_estimators": trial.suggest_int("n_estimators", 250, 500),
                "max_depth": trial.suggest_int("max_depth", 5, 7),
                "learning_rate": trial.suggest_float("learning_rate", 0.001, 0.004),
                "subsample": trial.suggest_float("subsample", 0.65, 0.85),
                "colsample_bytree": trial.suggest_float("colsample_bytree", 0.70, 0.90),
                "reg_alpha": trial.suggest_float("reg_alpha", 1e-4, 5e-2, log=True),
                "reg_lambda": trial.suggest_float("reg_lambda", 1e-6, 1e-3, log=True),
                "n_jobs": 4,
            }
        if target and "external" in str(target).lower():
            return {
                "k_features": k,
                "n_estimators": trial.suggest_int("n_estimators", 400, 500),
                "max_depth": trial.suggest_int("max_depth", 6, 7),
                "learning_rate": trial.suggest_float("learning_rate", 0.008, 0.02),
                "subsample": trial.suggest_float("subsample", 0.72, 0.80),
                "colsample_bytree": trial.suggest_float("colsample_bytree", 0.65, 0.75),
                "reg_alpha": trial.suggest_float("reg_alpha", 1e-3, 2e-2, log=True),
                "reg_lambda": trial.suggest_float("reg_lambda", 1.0, 10.0),
                "n_jobs": 4,
            }
        return {
            "k_features": k,
            "n_estimators": trial.suggest_int("n_estimators", 50, 300),
            "max_depth": trial.suggest_int("max_depth", 6, 9),
            "learning_rate": trial.suggest_float("learning_rate", 1e-4, 0.01, log=True),
            "subsample": trial.suggest_float("subsample", 0.6, 0.95),
            "colsample_bytree": trial.suggest_float("colsample_bytree", 0.7, 0.8),
            "reg_alpha": trial.suggest_float("reg_alpha", 1e-8, 1e-4, log=True),
            "reg_lambda": trial.suggest_float("reg_lambda", 1e-8, 0.1, log=True),
            "n_jobs": 4,
        }

    if model_name == "random_forest":
        if target and "internal" in str(target).lower():
            return {
                "k_features": k,
                "n_estimators": trial.suggest_int("n_estimators", 150, 300),
                "max_depth": trial.suggest_int("max_depth", 4, 7),
                "n_jobs": 4,
            }
        return {
            "k_features": k,
            "n_estimators": trial.suggest_int("n_estimators", 50, 200),
            "max_depth": trial.suggest_int("max_depth", 5, 10),
            "n_jobs": 4,
        }

    if model_name == "logistic":
        if target and "external" in str(target).lower():
            return {"k_features": k, "C": trial.suggest_float("C", 0.05, 0.2)}
        return {"k_features": k, "C": trial.suggest_float("C", 1e-4, 1e0, log=True)}

    if model_name == "svm":
        if target and "internal" in str(target).lower():
            return {
                "k_features": k,
                "kernel": "rbf",
                "C": trial.suggest_float("C", 0.3, 2.0),
            }
        if target and "external" in str(target).lower():
            return {
                "k_features": k,
                "kernel": "linear",
                "C": trial.suggest_float("C", 0.001, 0.02),
            }
        return {
            "k_features": k,
            "kernel": trial.suggest_categorical("kernel", ["rbf", "linear"]),
            "C": trial.suggest_float("C", 1e-3, 1e2, log=True),
        }

    return {"k_features": k}


def build_clf_from_params(
    model_name,
    params,
    random_state=RANDOM_STATE,
    scale_pos_w=None,
    n_jobs_override=None,
):
    # building a classifier from suggested parameters
    p = dict(params) if params is not None else {}
    p.pop("k_features", None)

    if model_name == "xgboost":
        p_local = dict(p)
        if scale_pos_w is not None:
            p_local["scale_pos_weight"] = scale_pos_w
        if n_jobs_override is not None:
            p_local["n_jobs"] = n_jobs_override
        return xgb.XGBClassifier(
            eval_metric="logloss", random_state=random_state, **p_local
        )

    if model_name == "random_forest":
        p_local = dict(p)
        if n_jobs_override is not None:
            p_local["n_jobs"] = n_jobs_override
        return RandomForestClassifier(
            class_weight="balanced", random_state=random_state, **p_local
        )

    if model_name == "logistic":
        p_local = dict(p)
        if n_jobs_override is not None:
            p_local["n_jobs"] = n_jobs_override
        return LogisticRegression(
            class_weight="balanced", max_iter=2000, random_state=random_state, **p_local
        )

    if model_name == "svm":
        p_local = dict(p)
        p_local.pop("n_jobs", None)
        return SVC(
            probability=True,
            class_weight="balanced",
            random_state=random_state,
            **p_local,
        )

    if model_name == "mlp":
        p_local = dict(p)
        p_local.pop("n_jobs", None)
        n_units = int(p_local.pop("n_units", 100)) if "n_units" in p_local else 100
        alpha = float(p_local.pop("alpha", 1e-4)) if "alpha" in p_local else 1e-4
        return MLPClassifier(
            hidden_layer_sizes=(n_units,),
            alpha=alpha,
            max_iter=1000,
            random_state=random_state,
            **p_local,
        )

    raise ValueError(f"Unknown model {model_name}")


def sfs_order_parallel(X, y, estimator, k, cv, scoring=SCORING, n_jobs=1):
    # greedy forward sfs using joblib for candidate evaluation
    n_samples, n_features = X.shape
    remaining = list(range(n_features))
    selected = []
    scores_step = []
    cv_obj = (
        check_cv(cv, y, classifier=True)
        if not isinstance(cv, (int, type(None)))
        else cv
    )
    for step in range(min(k, n_features)):

        def eval_cand(cand):
            cols = selected + [cand]
            try:
                v = float(
                    np.mean(
                        cross_val_score(
                            clone(estimator),
                            X[:, cols],
                            y,
                            cv=cv_obj,
                            scoring=scoring,
                            n_jobs=1,
                        )
                    )
                )
            except Exception:
                v = -np.inf
            return cand, v

        results = Parallel(n_jobs=n_jobs, prefer="threads")(
            delayed(eval_cand)(cand) for cand in remaining
        )
        best_cand, best_score = max(results, key=lambda t: t[1])
        if best_cand is None or best_score == -np.inf:
            break
        selected.append(best_cand)
        remaining.remove(best_cand)
        scores_step.append(best_score)
    return selected, scores_step


def bootstrap_ci(y_true, y_val, metric_fn, n_boot=1000, seed=RANDOM_STATE):
    # non-parametric bootstrap ci for a metric function
    rng = np.random.RandomState(seed)
    n = len(y_true)
    boots = []
    for _ in range(int(n_boot)):
        idx = rng.randint(0, n, n)
        try:
            boots.append(metric_fn(np.asarray(y_true)[idx], np.asarray(y_val)[idx]))
        except Exception:
            boots.append(np.nan)
    boots = np.array(boots)
    boots = boots[~np.isnan(boots)]
    if len(boots) == 0:
        return (np.nan, np.nan, np.nan)
    return (
        float(np.mean(boots)),
        float(np.percentile(boots, 2.5)),
        float(np.percentile(boots, 97.5)),
    )


# containers
master_sheets = {}
results_summary = {}
boot_all_results = {}
shap_top_features = {}

outer_cv = StratifiedKFold(
    n_splits=OUTER_SPLITS, shuffle=True, random_state=RANDOM_STATE
)

CHECKPOINT_PKL = OUTPUT_XLSX_NESTED.replace(".xlsx", "_checkpoint.pkl")


def save_progress(
    checkpoint_path=CHECKPOINT_PKL,
    excel_path=OUTPUT_XLSX_NESTED,
    master_sheets_local=None,
):
    # saving current state and writing master_sheets to excel
    try:
        chk_dir = os.path.dirname(checkpoint_path) or "."
        excel_dir = os.path.dirname(excel_path) or "."
        os.makedirs(chk_dir, exist_ok=True)
        os.makedirs(excel_dir, exist_ok=True)

        with open(checkpoint_path, "wb") as f:
            pickle.dump(
                {
                    "results_summary": results_summary,
                    "master_sheets": (
                        master_sheets
                        if master_sheets_local is None
                        else master_sheets_local
                    ),
                    "boot_all_results": boot_all_results,
                    "shap_top_features": shap_top_features,
                },
                f,
            )

        ms = master_sheets if master_sheets_local is None else master_sheets_local
        with pd.ExcelWriter(excel_path, engine="openpyxl", mode="w") as writer:
            for name, df in ms.items():
                try:
                    df.to_excel(writer, sheet_name=name[:31], index=False)
                except Exception:
                    try:
                        pd.DataFrame(df).to_excel(
                            writer, sheet_name=name[:31], index=False
                        )
                    except Exception:
                        pd.DataFrame({"data": [str(df)]}).to_excel(
                            writer, sheet_name=name[:31], index=False
                        )
        print(
            f"Progress saved to checkpoint: {checkpoint_path} and Excel: {excel_path}"
        )
    except Exception as e:
        print("Failed to save progress:", e)


# main nested cv per-target
for t_idx, target in enumerate(TARGETS):
    print("\nProcessing target:", target)
    tnorm = str(target).lower().strip()
    if "internal" in tnorm:
        models_to_run_local = ["svm", "xgboost", "random_forest"]
    elif "external" in tnorm:
        models_to_run_local = ["svm", "xgboost", "logistic"]
    else:
        models_to_run_local = MODELS_TO_RUN
    print(f"  --- normalized target '{tnorm}' mapped to models: {models_to_run_local}")

    if target not in IMPUTED_DATASETS or target not in SCALED_DATASETS:
        raise RuntimeError(f"Missing per-target datasets for '{target}'")
    df_t_imputed = IMPUTED_DATASETS[target].copy()
    df_t_scaled = SCALED_DATASETS[target].copy()

    other_targets = [t for t in TARGETS if t != target]
    predictors_cont = [
        c
        for c in df_t_scaled.columns
        if c not in (globals().get("ID_COL", "src_subject_id"), target)
        and c not in other_targets
    ]
    predictors_cat = [
        c
        for c in df_t_imputed.columns
        if c not in (globals().get("ID_COL", "src_subject_id"), target)
        and c not in other_targets
    ]
    predictors_cat = [c for c in predictors_cat if c not in set(predictors_cont)]
    feature_names = predictors_cont + predictors_cat

    parts = []
    if predictors_cont:
        parts.append(df_t_scaled[predictors_cont].astype(float).values)
    if predictors_cat:
        parts.append(df_t_imputed[predictors_cat].astype(float).values)
    if len(parts) == 0:
        print("No predictors found; skipping target.")
        results_summary[target] = {"skipped": True}
        continue
    X_all = np.hstack(parts)
    if target not in df_t_imputed.columns:
        raise RuntimeError(f"Target '{target}' not present in per-target data.")
    y = df_t_imputed[target].astype(int).values
    if len(np.unique(y)) < 2:
        print("Less than 2 classes; skipping.")
        results_summary[target] = {"skipped": True}
        continue

    per_model_eval = {}
    for model_name in models_to_run_local:
        print("Model:", model_name)
        fold_info = []
        outer_scores = []
        outer_idx = 0

        for train_idx_fold, test_idx_fold in outer_cv.split(X_all, y):
            outer_idx += 1
            print(f" Outer fold {outer_idx}/{outer_cv.get_n_splits()}")

            X_tr_fold, y_tr_fold = X_all[train_idx_fold], y[train_idx_fold]
            X_te_fold, y_te_fold = X_all[test_idx_fold], y[test_idx_fold]
            inner_cv = StratifiedKFold(
                n_splits=INNER_SPLITS, shuffle=True, random_state=RANDOM_STATE
            )
            sampler = optuna.samplers.TPESampler(seed=RANDOM_STATE)
            study = optuna.create_study(direction="maximize", sampler=sampler)

            def inner_obj(trial):
                params = suggest_params(
                    trial, model_name, X_tr_fold.shape[1], target=target
                )
                params_no_k = dict(params)
                params_no_k.pop("k_features", None)
                clf = build_clf_from_params(
                    model_name,
                    params_no_k,
                    random_state=RANDOM_STATE,
                    scale_pos_w=(
                        calc_scale_pos_weight(y_tr_fold)
                        if model_name == "xgboost"
                        else None
                    ),
                    n_jobs_override=1,
                )
                try:
                    score = float(
                        np.mean(
                            cross_val_score(
                                clone(clf),
                                X_tr_fold,
                                y_tr_fold,
                                cv=inner_cv,
                                scoring=SCORING,
                                n_jobs=1,
                            )
                        )
                    )
                except Exception:
                    score = 0.0
                return score

            study.optimize(
                inner_obj,
                n_trials=OPTUNA_TRIALS,
                n_jobs=OPTUNA_N_JOBS,
                show_progress_bar=False,
            )

            best_params_full = dict(study.best_params)
            k_best = int(
                best_params_full.pop(
                    "k_features", min(MAX_FEATURES_TO_SELECT, X_tr_fold.shape[1])
                )
            )
            k_best = min(k_best, max(1, X_tr_fold.shape[1]))

            clf_for_sfs_full = build_clf_from_params(
                model_name,
                dict(best_params_full, **{"k_features": k_best}),
                random_state=RANDOM_STATE,
                scale_pos_w=(
                    calc_scale_pos_weight(y_tr_fold)
                    if model_name == "xgboost"
                    else None
                ),
                n_jobs_override=1,
            )
            try:
                n_jobs_sfs = None if (N_JOBS is None) else int(N_JOBS)
                sel_rel_full_train, scores_full_train = sfs_order_parallel(
                    X_tr_fold,
                    y_tr_fold,
                    clone(clf_for_sfs_full),
                    k_best,
                    cv=max(2, min(3, INNER_SPLITS)),
                    scoring=SCORING,
                    n_jobs=n_jobs_sfs,
                )
            except Exception as e:
                print("SFS failed on outer-train fold:", e)
                sel_rel_full_train, scores_full_train = [], []

            try:
                if not sel_rel_full_train:
                    clf_fold = build_clf_from_params(
                        model_name,
                        best_params_full,
                        random_state=RANDOM_STATE,
                        scale_pos_w=(
                            calc_scale_pos_weight(y_tr_fold)
                            if model_name == "xgboost"
                            else None
                        ),
                        n_jobs_override=(None if N_JOBS is None else int(N_JOBS)),
                    )
                    clf_fold.fit(X_tr_fold, y_tr_fold)
                    y_pred = clf_fold.predict(X_te_fold)
                    try:
                        y_proba = clf_fold.predict_proba(X_te_fold)[:, 1]
                    except Exception:
                        y_proba = None
                else:
                    clf_fold = build_clf_from_params(
                        model_name,
                        best_params_full,
                        random_state=RANDOM_STATE,
                        scale_pos_w=(
                            calc_scale_pos_weight(y_tr_fold)
                            if model_name == "xgboost"
                            else None
                        ),
                        n_jobs_override=(None if N_JOBS is None else int(N_JOBS)),
                    )
                    clf_fold.fit(X_tr_fold[:, sel_rel_full_train], y_tr_fold)
                    y_pred = clf_fold.predict(X_te_fold[:, sel_rel_full_train])
                    try:
                        y_proba = clf_fold.predict_proba(
                            X_te_fold[:, sel_rel_full_train]
                        )[:, 1]
                    except Exception:
                        y_proba = None
            except Exception:
                y_pred = np.zeros_like(y_te_fold)
                y_proba = None

            fold_macro = float(f1_score(y_te_fold, y_pred, average="macro"))
            fold_mcc = float(matthews_corrcoef(y_te_fold, y_pred))
            fold_auc = (
                float(roc_auc_score(y_te_fold, y_proba))
                if (y_proba is not None and len(np.unique(y_proba)) > 1)
                else np.nan
            )

            global_sel_idx = (
                np.atleast_1d(sel_rel_full_train).astype(int)
                if sel_rel_full_train
                else np.array([], dtype=int)
            )
            selected_names_train = (
                [feature_names[int(i)] for i in np.atleast_1d(global_sel_idx)]
                if global_sel_idx.size > 0
                else []
            )

            fold_info.append(
                {
                    "fold_idx": outer_idx,
                    "selected_feature_names_ordered": list(selected_names_train),
                    "selected_feature_relative_indices": list(
                        map(int, np.atleast_1d(global_sel_idx))
                    ),
                    "selection_scores_per_step": (
                        list(scores_full_train)
                        if isinstance(scores_full_train, (list, tuple, np.ndarray))
                        else []
                    ),
                    "best_params": best_params_full.copy(),
                    "y_test": y_te_fold,
                    "y_pred": y_pred,
                    "y_proba": y_proba,
                    "fold_macro_f1": fold_macro,
                    "fold_mcc": fold_mcc,
                    "fold_auc": fold_auc,
                }
            )
            outer_scores.append(fold_macro)

        per_model_eval[model_name] = {
            "outer_fold_scores": outer_scores,
            "mean_macro_f1": (float(np.mean(outer_scores)) if outer_scores else np.nan),
            "std_macro_f1": (float(np.std(outer_scores)) if outer_scores else np.nan),
            "fold_info": fold_info,
        }
        print("Model summary", model_name, per_model_eval[model_name]["mean_macro_f1"])

    # pooled confusion matrices and save pngs
    try:
        for model_name, model_info in per_model_eval.items():
            fold_info = model_info.get("fold_info", [])
            if not fold_info:
                continue
            try:
                y_true_all = np.concatenate([f["y_test"] for f in fold_info])
                y_pred_all = np.concatenate([f["y_pred"] for f in fold_info])
            except Exception:
                continue
            if y_true_all.size == 0 or y_pred_all.size == 0:
                continue
            labels = np.unique(np.concatenate([y_true_all, y_pred_all]))
            cm = confusion_matrix(y_true_all, y_pred_all, labels=labels)
            fig, ax = plt.subplots(figsize=(5, 5))
            disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=labels)
            disp.plot(ax=ax, cmap="Blues", values_format="d")
            ax.set_title(
                f"Confusion matrix -- {_sanitize_filename(target)} -- {model_name}"
            )
            plt.tight_layout()
            fname = f"confmat_{_sanitize_filename(target)}_{model_name}.png"
            try:
                fig.savefig(fname, dpi=150)
                print(f"Saved confusion matrix --- {fname}")
            except Exception as e:
                print(f"Failed to save confusion matrix {fname}: {e}")
            plt.close(fig)
            try:
                cm_df = pd.DataFrame(
                    cm, index=[str(l) for l in labels], columns=[str(l) for l in labels]
                )
                master_sheets[f"confmat_{target}_{model_name}"] = (
                    cm_df.reset_index().rename(columns={"index": "true_label"})
                )
            except Exception:
                pass
    except Exception as e:
        print("Error while creating/saving confusion matrices:", e)

    # aggregate feature selection across outer folds for best model
    best_model = max(
        per_model_eval.keys(), key=lambda k: per_model_eval[k]["mean_macro_f1"]
    )
    print("Best model:", best_model)
    fold_info_best = per_model_eval[best_model]["fold_info"]
    n_folds_done = len(fold_info_best)
    freq = Counter()
    rank_sum = defaultdict(float)
    rank_counts = defaultdict(int)
    detailed_fold_table = []
    for finfo in fold_info_best:
        ordered = finfo["selected_feature_names_ordered"]
        detailed_fold_table.append(
            {
                "fold": finfo["fold_idx"],
                "ordered_features": ordered,
                "scores_per_step": finfo["selection_scores_per_step"],
                "best_params": finfo["best_params"],
            }
        )
        for r, f in enumerate(ordered, start=1):
            freq[f] += 1
            rank_sum[f] += r
            rank_counts[f] += 1

    agg_rows = []
    for fname in feature_names:
        cnt = freq.get(fname, 0)
        fr = cnt / n_folds_done if n_folds_done > 0 else 0.0
        avg_r = (
            (rank_sum[fname] / rank_counts[fname]) if rank_counts[fname] > 0 else np.nan
        )
        agg_rows.append(
            {
                "feature": fname,
                "selection_count": int(cnt),
                "selection_frequency": float(fr),
                "avg_rank_when_selected": float(avg_r),
            }
        )
    df_feature_agg = pd.DataFrame(agg_rows).sort_values(
        ["selection_frequency", "avg_rank_when_selected"], ascending=[False, True]
    )

    chosen_features = [
        r["feature"]
        for r in agg_rows
        if r["selection_frequency"] >= FEATURE_FREQUENCY_THRESHOLD
    ]
    if not chosen_features:
        if fold_info_best:
            best_fold_idx_local = int(
                np.argmax([f["fold_macro_f1"] for f in fold_info_best])
            )
            chosen_features = fold_info_best[best_fold_idx_local][
                "selected_feature_names_ordered"
            ]
        else:
            chosen_features = [
                feature_names[i]
                for i in range(min(MAX_FEATURES_TO_SELECT, len(feature_names)))
            ]

    chosen_indices = (
        [feature_names.index(f) for f in chosen_features]
        if chosen_features
        else list(range(len(feature_names)))
    )
    selected_indices = chosen_indices
    selected_names = (
        chosen_features
        if chosen_features
        else [feature_names[i] for i in selected_indices]
    )

    best_fold_idx = (
        int(np.argmax([f["fold_macro_f1"] for f in fold_info_best]))
        if fold_info_best
        else 0
    )
    final_params = (
        fold_info_best[best_fold_idx]["best_params"].copy() if fold_info_best else {}
    )

    clf_final = build_clf_from_params(
        best_model,
        final_params,
        random_state=RANDOM_STATE,
        scale_pos_w=(calc_scale_pos_weight(y) if best_model == "xgboost" else None),
        n_jobs_override=(None if N_JOBS is None else int(N_JOBS)),
    )
    if not selected_indices:
        selected_indices = list(range(len(feature_names)))
    X_full_selected = X_all[:, selected_indices]
    try:
        clf_final.fit(X_full_selected, y)
    except Exception:
        pass
    final_pipeline = {"clf": clf_final}

    # pooled oof metrics
    y_oof_true = (
        np.concatenate([f["y_test"] for f in fold_info_best])
        if fold_info_best
        else np.array([])
    )
    y_oof_pred = (
        np.concatenate([f["y_pred"] for f in fold_info_best])
        if fold_info_best
        else np.array([])
    )
    y_oof_proba = (
        np.concatenate([f["y_proba"] for f in fold_info_best])
        if fold_info_best and all(f["y_proba"] is not None for f in fold_info_best)
        else None
    )

    pooled_macro_f1 = (
        float(f1_score(y_oof_true, y_oof_pred, average="macro"))
        if y_oof_pred.size > 0
        else np.nan
    )
    pooled_mcc = (
        float(matthews_corrcoef(y_oof_true, y_oof_pred))
        if y_oof_pred.size > 0
        else np.nan
    )
    pooled_auc = (
        float(roc_auc_score(y_oof_true, y_oof_proba))
        if (y_oof_proba is not None and len(np.unique(y_oof_proba)) > 1)
        else np.nan
    )

    if y_oof_pred.size > 0:
        f1_mean, f1_lb, f1_ub = bootstrap_ci(
            y_oof_true,
            y_oof_pred,
            lambda a, b: f1_score(a, b, average="macro"),
            n_boot=BOOT_N,
        )
        mcc_mean, mcc_lb, mcc_ub = bootstrap_ci(
            y_oof_true, y_oof_pred, matthews_corrcoef, n_boot=BOOT_N
        )
    else:
        f1_mean = mcc_mean = f1_lb = f1_ub = np.nan
    if y_oof_proba is not None:
        auc_mean, auc_lb, auc_ub = bootstrap_ci(
            y_oof_true, y_oof_proba, roc_auc_score, n_boot=BOOT_N
        )
    else:
        auc_mean = auc_lb = auc_ub = np.nan

    results_summary[target] = {
        "per_model_eval": per_model_eval,
        "best_model": best_model,
        "final_pipeline": final_pipeline,
        "selected_feature_indices": selected_indices,
        "selected_feature_names": selected_names,
        "final_params": final_params,
        "pooled_oof_metrics": {
            "macro_f1": pooled_macro_f1,
            "mcc": pooled_mcc,
            "auc": pooled_auc,
        },
        "pooled_oof_bootstrap": {
            "f1": (f1_mean, f1_lb, f1_ub),
            "mcc": (mcc_mean, mcc_lb, mcc_ub),
            "auc": (auc_mean, auc_lb, auc_ub),
        },
        "feature_selection_aggregation": df_feature_agg,
        "detailed_fold_selection": detailed_fold_table,
        "final_test_metrics": {
            "macro_f1": pooled_macro_f1,
            "mcc": pooled_mcc,
            "auc": pooled_auc,
        },
        "feature_names_all": list(feature_names),
    }

    # assemble sheets for writing later
    df_eval = pd.DataFrame(
        [
            {
                "model": m,
                "mean_macro_f1": v["mean_macro_f1"],
                "std_macro_f1": v["std_macro_f1"],
                "outer_scores": v["outer_fold_scores"],
            }
            for m, v in per_model_eval.items()
        ]
    ).set_index("model")
    df_boot = pd.DataFrame(
        {
            "metric": ["macro_f1", "mcc", "auc"],
            "pooled": [pooled_macro_f1, pooled_mcc, pooled_auc],
            "boot_mean": [f1_mean, mcc_mean, auc_mean],
            "boot_lb": [f1_lb, mcc_lb, auc_lb],
            "boot_ub": [f1_ub, mcc_ub, auc_ub],
        }
    )
    rows = [
        {
            "fold": d["fold"],
            "ordered_features": ", ".join(d["ordered_features"]),
            "scores_per_step": ", ".join([f"{s:.4f}" for s in d["scores_per_step"]]),
        }
        for d in detailed_fold_table
    ]
    df_detailed = pd.DataFrame(rows)
    df_feature_agg_out = df_feature_agg.copy()
    df_final_metrics = pd.DataFrame(
        [
            {"metric": "macro_f1", "value": pooled_macro_f1},
            {"metric": "mcc", "value": pooled_mcc},
            {"metric": "auc", "value": pooled_auc},
        ]
    )

    master_sheets[f"model_eval_{target}"] = df_eval.reset_index()
    master_sheets[f"pooled_boot_{target}"] = df_boot
    master_sheets[f"feature_agg_{target}"] = df_feature_agg_out
    master_sheets[f"fold_sel_{target}"] = df_detailed
    master_sheets[f"final_metrics_{target}"] = df_final_metrics

    # checkpoint after each target
    try:
        save_progress()
    except Exception as e:
        print("Warning: failed to save progress for target", target, "Error:", e)

    print(f"Completed target {target} (sheets stored and checkpoint saved).")

# aggregate bootstrapped summaries across models
for target, info in results_summary.items():
    if not isinstance(info, dict) or info.get("skipped", False):
        continue
    per_model = info.get("per_model_eval", {}) or {}
    rows = []
    for model_name in list(per_model.keys()) if per_model else MODELS_TO_RUN:
        model_info = per_model.get(model_name, {})
        fold_info = model_info.get("fold_info", [])
        try:
            y_oof_true = (
                np.concatenate([f["y_test"] for f in fold_info])
                if fold_info
                else np.array([])
            )
            y_oof_pred = (
                np.concatenate([f["y_pred"] for f in fold_info])
                if fold_info
                else np.array([])
            )
        except Exception:
            y_oof_true, y_oof_pred = np.array([]), np.array([])
        y_oof_proba = (
            np.concatenate([f["y_proba"] for f in fold_info])
            if fold_info and all(f["y_proba"] is not None for f in fold_info)
            else None
        )
        n_samples = len(y_oof_true) if y_oof_true is not None else 0

        if n_samples == 0 or len(y_oof_pred) == 0:
            f1_stats = (np.nan, np.nan, np.nan)
            mcc_stats = (np.nan, np.nan, np.nan)
        else:
            f1_stats = bootstrap_ci(
                y_oof_true,
                y_oof_pred,
                lambda a, b: f1_score(a, b, average="macro"),
                n_boot=BOOT_N,
            )
            mcc_stats = bootstrap_ci(
                y_oof_true, y_oof_pred, matthews_corrcoef, n_boot=BOOT_N
            )

        if y_oof_proba is None or (
            isinstance(y_oof_proba, np.ndarray) and len(np.unique(y_oof_proba)) <= 1
        ):
            auc_stats = (np.nan, np.nan, np.nan)
        else:
            auc_stats = bootstrap_ci(
                y_oof_true, y_oof_proba, roc_auc_score, n_boot=BOOT_N
            )

        try:
            pooled_f1 = (
                float(f1_score(y_oof_true, y_oof_pred, average="macro"))
                if n_samples > 0
                else np.nan
            )
        except Exception:
            pooled_f1 = np.nan
        try:
            pooled_mcc = (
                float(matthews_corrcoef(y_oof_true, y_oof_pred))
                if n_samples > 0
                else np.nan
            )
        except Exception:
            pooled_mcc = np.nan
        try:
            pooled_auc = (
                float(roc_auc_score(y_oof_true, y_oof_proba))
                if (y_oof_proba is not None and len(np.unique(y_oof_proba)) > 1)
                else np.nan
            )
        except Exception:
            pooled_auc = np.nan

        rows.append(
            {
                "model": model_name,
                "n_oof_samples": int(n_samples),
                "macro_f1_point": pooled_f1,
                "macro_f1_boot_mean": f1_stats[0],
                "macro_f1_boot_lb": f1_stats[1],
                "macro_f1_boot_ub": f1_stats[2],
                "mcc_point": pooled_mcc,
                "mcc_boot_mean": mcc_stats[0],
                "mcc_boot_lb": mcc_stats[1],
                "mcc_boot_ub": mcc_stats[2],
                "auc_point": pooled_auc,
                "auc_boot_mean": auc_stats[0],
                "auc_boot_lb": auc_stats[1],
                "auc_boot_ub": auc_stats[2],
            }
        )
    df_out = pd.DataFrame(rows).set_index("model")
    boot_all_results[target] = df_out
    master_sheets[f"boot_summary_{target}"] = df_out.reset_index()

print("\nBootstrapped summaries prepared and stored in master_sheets.")


## SHAP & Confusion Matrices


# sanity checks
if "results_summary" not in globals() or "master_sheets" not in globals():
    raise RuntimeError(
        "results_summary and master_sheets must exist (run prior cells)."
    )

FIG_DIR = "figures"
os.makedirs(FIG_DIR, exist_ok=True)
RND = globals().get("RANDOM_STATE", 37)

# checking shap availability
try:
    import shap

    shap_available = True
except Exception:
    shap = None
    shap_available = False

# initializing shap_top_features if not present
shap_top_features = {} if "shap_top_features" not in globals() else shap_top_features

if shap_available:
    for target, info in results_summary.items():
        if not isinstance(info, dict) or info.get("skipped", False):
            continue
        if "final_pipeline" not in info:
            continue

        final_pipe = info["final_pipeline"]
        clf = (
            final_pipe["clf"]
            if isinstance(final_pipe, dict) and "clf" in final_pipe
            else final_pipe
        )
        if clf is None:
            continue

        # determining candidate feature names
        sel_names_from_summary = info.get("selected_feature_names", None)
        all_feature_names = info.get("feature_names_all", None)
        if all_feature_names is None:
            df_agg = info.get("feature_selection_aggregation", pd.DataFrame())
            if (
                isinstance(df_agg, pd.DataFrame)
                and "feature" in df_agg.columns
                and len(df_agg) > 0
            ):
                all_feature_names = list(df_agg["feature"].values)
            else:
                df_tmp = IMPUTED_DATASETS[target]
                all_feature_names = [
                    c
                    for c in df_tmp.columns
                    if c
                    not in (globals().get("ID_COL", "src_subject_id"),) + tuple(TARGETS)
                ]

        sel_indices = info.get("selected_feature_indices", None)
        sel_names = []
        if sel_indices:
            try:
                arr_idx = np.atleast_1d(sel_indices).astype(int)
            except Exception:
                arr_idx = np.array(sel_indices)
            for ii in arr_idx:
                try:
                    ii_int = int(ii)
                    if ii_int < 0:
                        sel_names.append(all_feature_names[ii_int])
                    elif ii_int < len(all_feature_names):
                        sel_names.append(all_feature_names[ii_int])
                except Exception:
                    continue
            if len(sel_names) == 0 and sel_names_from_summary:
                sel_names = list(sel_names_from_summary)
        else:
            sel_names = (
                list(sel_names_from_summary)
                if sel_names_from_summary
                else list(all_feature_names)
            )

        # building X_selected (preferring scaled then imputed)
        df_t_imputed = IMPUTED_DATASETS[target]
        df_t_scaled = SCALED_DATASETS[target]
        X_cols = []
        final_sel_names = []
        for fname in sel_names:
            if fname in df_t_scaled.columns:
                try:
                    X_cols.append(
                        df_t_scaled[fname].astype(float).values.reshape(-1, 1)
                    )
                    final_sel_names.append(fname)
                except Exception:
                    continue
            elif fname in df_t_imputed.columns:
                try:
                    X_cols.append(
                        df_t_imputed[fname].astype(float).values.reshape(-1, 1)
                    )
                    final_sel_names.append(fname)
                except Exception:
                    continue
        if len(X_cols) == 0:
            continue
        X_selected_raw = np.hstack(X_cols)

        # computing shap values
        shap_vals = None
        explanation_obj = None
        try:
            if hasattr(clf, "predict_proba"):
                explainer = shap.Explainer(
                    lambda X: clf.predict_proba(np.asarray(X))[:, 1],
                    X_selected_raw,
                    feature_names=final_sel_names,
                )
            else:
                explainer = shap.Explainer(
                    clf, X_selected_raw, feature_names=final_sel_names
                )
            explanation_obj = explainer(X_selected_raw)
            shap_vals = (
                explanation_obj.values
                if hasattr(explanation_obj, "values")
                else np.array(explanation_obj)
            )
        except Exception:
            # fallback: LinearExplainer -> TreeExplainer -> KernelExplainer (sampled)
            tried = False
            if (
                "LogisticRegression" in clf.__class__.__name__
                or "Linear" in clf.__class__.__name__
            ):
                try:
                    bg = _sample_background(
                        X_selected_raw, n=min(200, max(10, X_selected_raw.shape[0]))
                    )
                    explainer = shap.LinearExplainer(
                        clf, bg, feature_perturbation="interventional"
                    )
                    explanation_obj = explainer(X_selected_raw)
                    shap_vals = (
                        explanation_obj.values
                        if hasattr(explanation_obj, "values")
                        else np.array(explanation_obj)
                    )
                    tried = True
                except Exception:
                    tried = False
            if not tried:
                try:
                    explainer = shap.TreeExplainer(clf)
                    shap_vals = explainer.shap_values(X_selected_raw)
                except Exception:
                    try:
                        bg = _sample_background(
                            X_selected_raw, n=min(100, max(10, X_selected_raw.shape[0]))
                        )
                        f = (
                            (lambda Xr: clf.predict_proba(np.asarray(Xr))[:, 1])
                            if hasattr(clf, "predict_proba")
                            else (lambda Xr: clf.predict(np.asarray(Xr)).astype(float))
                        )
                        explainer = shap.KernelExplainer(f, bg)
                        explanation_obj = explainer(X_selected_raw, nsamples=200)
                        shap_vals = (
                            explanation_obj.values
                            if hasattr(explanation_obj, "values")
                            else np.array(explanation_obj)
                        )
                    except Exception:
                        shap_vals = None

        if shap_vals is None:
            continue

        # normalizing shap array to 2D and getting mean absolute
        if isinstance(shap_vals, list):
            shap_arr_full = (
                np.asarray(shap_vals[1])
                if len(shap_vals) >= 2
                else np.asarray(shap_vals[0])
            )
        else:
            shap_arr_full = np.asarray(shap_vals)
            if shap_arr_full.ndim == 3:
                shap_arr_full = (
                    shap_arr_full[1]
                    if shap_arr_full.shape[0] >= 2
                    else shap_arr_full.reshape(
                        shap_arr_full.shape[1], shap_arr_full.shape[2]
                    )
                )
        try:
            if shap_arr_full.ndim != 2:
                shap_arr_full = shap_arr_full.reshape(shap_arr_full.shape[0], -1)
        except Exception:
            continue
        mean_abs_shap = np.mean(np.abs(shap_arr_full), axis=0)
        mn = min(len(mean_abs_shap), len(final_sel_names))
        mean_abs_shap = mean_abs_shap[:mn]
        final_sel_names = final_sel_names[:mn]
        top_k = min(30, mn)
        top_idx = np.argsort(-mean_abs_shap)[:top_k]
        top_features = [final_sel_names[int(i)] for i in top_idx]
        shap_top_features[target] = top_features

        # saving shap summary figure
        try:
            plt.figure(figsize=(8, 6))
            if explanation_obj is not None:
                shap.summary_plot(explanation_obj, show=False)
            else:
                feat_ord = np.array(final_sel_names)[top_idx]
                vals = mean_abs_shap[top_idx]
                y_pos = np.arange(len(feat_ord))
                plt.barh(y_pos[::-1], vals[::-1])
                plt.yticks(y_pos[::-1], feat_ord[::-1])
                plt.xlabel("mean |SHAP value|")
            plt.tight_layout()
            fname = os.path.join(
                FIG_DIR, f"shap_summary_{target.replace(' ', '_')}.png"
            )
            plt.savefig(fname, dpi=150)
            plt.close()
        except Exception:
            pass

    # persisting shap top features into master_sheets
    for t, feats in shap_top_features.items():
        master_sheets[f"shap_top_{t}"] = pd.DataFrame({"shap_top_features": feats})
    try:
        save_progress()
    except Exception:
        pass
else:
    print("shap not available; skipped shap block.")

# confusion matrices computed from results_summary
CMAP = "Blues"


def _sample_background(
    X: np.ndarray, n: int = 100, random_state: int = RND
) -> np.ndarray:
    # sampling up to n rows (no replacement) for shap background
    rng = np.random.RandomState(random_state)
    nrows = int(X.shape[0])
    if nrows <= n:
        return X
    return X[rng.choice(nrows, size=n, replace=False)]


def plot_single_confmat(cm: np.ndarray, title: str, vmax: int):
    # plotting single confusion matrix heatmap (counts)
    fig, ax = plt.subplots(figsize=(5, 4.5))
    im = ax.imshow(cm, cmap=CMAP, vmin=0, vmax=vmax)
    ax.set_xlabel("Predicted label")
    ax.set_ylabel("True label")
    ax.set_xticks([0, 1])
    ax.set_yticks([0, 1])
    ax.set_xticklabels(["0", "1"])
    ax.set_yticklabels(["0", "1"])
    ax.set_title(title, fontsize=12, pad=8)
    for (r, c), val in np.ndenumerate(cm):
        ax.text(
            c,
            r,
            f"{int(val)}",
            ha="center",
            va="center",
            color="white" if val > 0.6 * vmax else "black",
        )
    fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
    plt.tight_layout()
    return fig


def plot_grouped_confmats(target: str, mats_dict: dict):
    # plotting a row of confusion matrices for a single target
    models = list(mats_dict.keys())
    n_models = len(models)
    fig = plt.figure(figsize=(5 * n_models, 5))
    gs = gridspec.GridSpec(1, n_models, wspace=0.3)
    vmax = max(cm.max() for cm in mats_dict.values())
    for i, model in enumerate(models):
        ax = fig.add_subplot(gs[i])
        cm = mats_dict[model]
        im = ax.imshow(cm, cmap=CMAP, vmin=0, vmax=vmax)
        ax.set_xlabel("Predicted label")
        ax.set_ylabel("True label")
        ax.set_xticks([0, 1])
        ax.set_yticks([0, 1])
        ax.set_xticklabels(["0", "1"])
        ax.set_yticklabels(["0", "1"])
        ax.set_title(model, fontsize=12, pad=8)
        for (r, c), val in np.ndenumerate(cm):
            ax.text(
                c,
                r,
                f"{int(val)}",
                ha="center",
                va="center",
                color="white" if val > 0.6 * vmax else "black",
            )
    fig.suptitle(f"Confusion Matrices – {target}", fontsize=14, y=1.02)
    plt.tight_layout()
    return fig


# computing confusion matrices from results_summary fold_info
computed_confmats = {}
for target, info in results_summary.items():
    if not isinstance(info, dict) or info.get("skipped", False):
        continue
    per_model = info.get("per_model_eval", {}) or {}
    target_mats = {}
    for model_name, model_info in per_model.items():
        fold_info = model_info.get("fold_info", []) or []
        if not fold_info:
            continue
        try:
            y_true_all = np.concatenate([f["y_test"] for f in fold_info])
            y_pred_all = np.concatenate([f["y_pred"] for f in fold_info])
        except Exception:
            continue
        if y_true_all.size == 0 or y_pred_all.size == 0:
            continue
        labels = np.unique(np.concatenate([y_true_all, y_pred_all]))
        cm = confusion_matrix(y_true_all, y_pred_all, labels=labels)
        target_mats[model_name] = {
            "cm": cm,
            "labels": labels,
            "y_true_all": y_true_all,
            "y_pred_all": y_pred_all,
        }
        # storing confusion matrix for excel output
        try:
            cm_df = pd.DataFrame(
                cm, index=[str(l) for l in labels], columns=[str(l) for l in labels]
            )
            master_sheets[f"confmat_{target}_{model_name}"] = (
                cm_df.reset_index().rename(columns={"index": "true_label"})
            )
        except Exception:
            pass
    if target_mats:
        computed_confmats[target] = target_mats

# plotting & saving all confusion matrices (grouped + individual)
for target, mats_info in computed_confmats.items():
    mats_dict = {m: v["cm"] for m, v in mats_info.items()}
    if not mats_dict:
        continue
    # grouped figure
    fig_group = plot_grouped_confmats(target, mats_dict)
    fpath_group = os.path.join(
        FIG_DIR, f"confmats_grouped_{target.replace(' ', '_')}.png"
    )
    try:
        fig_group.savefig(fpath_group, dpi=150, bbox_inches="tight")
        plt.close(fig_group)
    except Exception:
        plt.close(fig_group)
    # individual figures
    vmax_target = max(cm.max() for cm in mats_dict.values())
    for model, info_dict in mats_info.items():
        cm = info_dict["cm"]
        fig_single = plot_single_confmat(
            cm, title=f"{model} – {target}", vmax=vmax_target
        )
        fpath_single = os.path.join(
            FIG_DIR, f"confmat_{target.replace(' ', '_')}_{model.replace(' ', '_')}.png"
        )
        try:
            fig_single.savefig(fpath_single, dpi=150, bbox_inches="tight")
            plt.close(fig_single)
        except Exception:
            plt.close(fig_single)

# ensuring master_sheets updated and checkpoint
try:
    save_progress()
except Exception:
    pass

print("shap and confusion matrices completed; figures saved to:", FIG_DIR)


## Statistical tests & export


# ensuring containers exist
results_summary = globals().get("results_summary", {})
IMPUTED_DATASETS = globals().get("IMPUTED_DATASETS", {})
master_sheets = globals().get("master_sheets", {})

OUTPUT_XLSX_NESTED = globals().get(
    "OUTPUT_XLSX_NESTED", "results/nested_cv_sfs_results_with_final_model.xlsx"
)
os.makedirs(os.path.dirname(OUTPUT_XLSX_NESTED) or ".", exist_ok=True)


# computing common language effect size (cles) from U statistic
def cles(U, n1, n2):
    if n1 <= 0 or n2 <= 0:
        return np.nan
    return float(U) / (n1 * n2)


# computing rank-biserial correlation from U statistic
def rank_biserial(U, n1, n2):
    if n1 <= 0 or n2 <= 0:
        return np.nan
    return 1.0 - (2.0 * U) / (n1 * n2)


# benjamini-hochberg fdr correction
def benjamini_hochberg(pvals):
    if len(pvals) == 0:
        return np.array([])
    pvals = np.array(pvals, dtype=float)
    n = len(pvals)
    order = np.argsort(pvals)
    ranks = np.empty(n, int)
    ranks[order] = np.arange(1, n + 1)
    adj = pvals * n / ranks
    adj_corrected = np.minimum.accumulate(adj[::-1])[::-1]
    adj_corrected[adj_corrected > 1.0] = 1.0
    return adj_corrected


stats_for_excel = {}
shap_top_features = globals().get("shap_top_features", {})

for target in globals().get("TARGETS", []):
    print(f"\n=== stats for target: {target} ===")
    df_imputed = IMPUTED_DATASETS.get(target)
    if df_imputed is None:
        print("no imputed dataset for", target, "- skipping.")
        continue

    # choosing candidate top features (shap -> aggregated selection -> univariate auc)
    top_feats = None
    source = None
    if target in shap_top_features and shap_top_features[target]:
        top_feats = list(shap_top_features[target][:20])
        source = "shap"
    if top_feats is None:
        info = results_summary.get(target, {})
        df_agg = info.get("feature_selection_aggregation")
        if isinstance(df_agg, pd.DataFrame) and "selection_frequency" in df_agg.columns:
            df_sorted = df_agg.sort_values(
                by=["selection_frequency", "avg_rank_when_selected"],
                ascending=[False, True],
            )
            top_feats = list(df_sorted["feature"].head(20).values)
            source = "feature_selection_aggregation"

    if top_feats is None:
        # fallback: univariate auc ranking
        from sklearn.metrics import roc_auc_score

        auc_scores = {}
        y = pd.to_numeric(df_imputed[target], errors="coerce").astype(float).values
        all_features = [
            c
            for c in df_imputed.columns
            if c
            not in (globals().get("ID_COL", "src_subject_id"),)
            + tuple(globals().get("TARGETS", []))
        ]
        for feat in all_features:
            col = pd.to_numeric(df_imputed[feat], errors="coerce").astype(float)
            mask = ~np.isnan(col.values) & ~np.isnan(y)
            if mask.sum() < 10:
                auc_scores[feat] = np.nan
                continue
            vals = col.values[mask]
            yt = y[mask]
            try:
                auc_scores[feat] = (
                    roc_auc_score(yt, vals) if len(np.unique(vals)) > 1 else np.nan
                )
            except Exception:
                auc_scores[feat] = np.nan
        sorted_feats = [
            f
            for f, s in sorted(
                auc_scores.items(),
                key=lambda kv: (np.nan_to_num(kv[1], nan=-np.inf)),
                reverse=True,
            )
        ]
        top_feats = sorted_feats[:20]
        source = "univariate_auc"
    print(f" using top features from: {source}. selected: {len(top_feats)}")

    rows = []
    pvals = []
    feats_done = []
    for feat in top_feats:
        col = pd.to_numeric(df_imputed.get(feat), errors="coerce").astype(float)
        a = col[df_imputed[target] == 1].dropna().values
        b = col[df_imputed[target] == 0].dropna().values
        n1 = int(len(a))
        n2 = int(len(b))
        if n1 == 0 or n2 == 0:
            rows.append(
                {
                    "feature": feat,
                    "n1": n1,
                    "n2": n2,
                    "median_pos": np.nan,
                    "median_neg": np.nan,
                    "U": np.nan,
                    "p": np.nan,
                    "CLES": np.nan,
                    "RBC": np.nan,
                }
            )
            pvals.append(np.nan)
            feats_done.append(feat)
            continue
        try:
            try:
                U_stat, p_val = mannwhitneyu(a, b, alternative="two-sided")
            except TypeError:
                U_stat, p_val = mannwhitneyu(a, b)
                p_val = min(1.0, 2.0 * p_val)
        except Exception:
            U_stat, p_val = np.nan, np.nan
        med_a = float(np.median(a)) if len(a) > 0 else np.nan
        med_b = float(np.median(b)) if len(b) > 0 else np.nan
        cl = cles(U_stat, n1, n2) if np.isfinite(U_stat) else np.nan
        rb = rank_biserial(U_stat, n1, n2) if np.isfinite(U_stat) else np.nan
        rows.append(
            {
                "feature": feat,
                "n1": n1,
                "n2": n2,
                "median_pos": med_a,
                "median_neg": med_b,
                "U": (float(U_stat) if np.isfinite(U_stat) else np.nan),
                "p": (float(p_val) if np.isfinite(p_val) else np.nan),
                "CLES": float(cl) if np.isfinite(cl) else np.nan,
                "RBC": float(rb) if np.isfinite(rb) else np.nan,
            }
        )
        pvals.append(
            float(p_val)
            if (isinstance(p_val, (int, float)) and np.isfinite(p_val))
            else np.nan
        )
        feats_done.append(feat)

    pvals_arr = np.array(
        [
            pv if (isinstance(pv, (int, float)) and np.isfinite(pv)) else np.nan
            for pv in pvals
        ],
        dtype=float,
    )
    nonnan_idx = np.where(~np.isnan(pvals_arr))[0]
    adj_pvals = np.full_like(pvals_arr, np.nan)
    if len(nonnan_idx) > 0:
        try:
            if _has_statsmodels:
                rej, pvals_corr, _, _ = multipletests(
                    pvals_arr[nonnan_idx], method="fdr_bh"
                )
                adj_pvals[nonnan_idx] = pvals_corr
            else:
                adj = benjamini_hochberg(pvals_arr[nonnan_idx])
                adj_pvals[nonnan_idx] = adj
        except Exception:
            adj = benjamini_hochberg(pvals_arr[nonnan_idx])
            adj_pvals[nonnan_idx] = adj

    for i, feat in enumerate(feats_done):
        rows[i]["p_adj_fdr_bh"] = (
            float(adj_pvals[i])
            if (i < len(adj_pvals) and np.isfinite(adj_pvals[i]))
            else np.nan
        )

    stats_df = pd.DataFrame(rows)
    sort_col = "p_adj_fdr_bh" if "p_adj_fdr_bh" in stats_df.columns else "p"
    stats_df = stats_df.sort_values(by=[sort_col], na_position="last")
    stats_for_excel[target] = stats_df
    master_sheets[f"stats_{target}"] = stats_df
    print("top 10 feature statistics for", target)
    print(stats_df.head(10))

# writing master_sheets to excel
from pandas import ExcelWriter

with ExcelWriter(OUTPUT_XLSX_NESTED, engine="openpyxl", mode="w") as writer:
    for name, df in master_sheets.items():
        try:
            df.to_excel(writer, sheet_name=name[:31], index=False)
        except Exception:
            try:
                df.to_excel(writer, sheet_name=name[:31])
            except Exception:
                pd.DataFrame({"data": [str(df)]}).to_excel(
                    writer, sheet_name=name[:31], index=False
                )

print("\nwrote all sheets to:", OUTPUT_XLSX_NESTED)
print(
    "results available in `results_summary`, `boot_all_results`, `shap_top_features`, and `stats_for_excel`."
)

