OSCR

Lesion-Level Subtypes of White Matter Hyperintensity Evolution Beyond Spatial Location.

Code ↔ Paper

9 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 9 matches
  1. [1] § Methods › Lesion Clustering ↔ clustering.py, lines 8–64 · score 0.70 · diagnostic categories, optimal, reproducibility, squares, matrices, residualized
  2. [2] § Methods › Lesion-Wise Metrics Extraction ↔ Preprocesing/Whm_mask_only_mean_cluster_fMRI.py, lines 26–41 · score 0.67 · preprocessed BOLD, fMRI, WMH cluster, dilated, masks, space
  3. [3] § Results ↔ suceptibilidad.py, lines 1–29 · score 0.67 · Mann Whitney, lesion volume, WMH volume, Cohen, binomial, log
  4. [4] § Methods › Lesion Identification ↔ Preprocesing/Whm_mask_only_mean_cluster_DWI.py, lines 2–22 · score 0.66 · connected component, WMH cluster mask, nearest, peak, voxel
  5. [5] § Methods › Analysis of Lesion Patterns ↔ suceptibilidad.py, lines 1–29 · score 0.63 · WMH burden, WMH volume, Binomial, covariates, predicted, CSF
  6. [6] § Results ↔ clustering.py, lines 90–99 · score 0.62 · Power slope, fALFF, GFA, ISO, QA, rd1
  7. [7] § Methods › Study Design and Participants ↔ suceptibilidad.py, lines 471–540 · score 0.60 · APOE genotype, HR, PP, MCI, CN, education
  8. [8] § Results ↔ clustering.py, lines 90–99 · score 0.59 · Power slope, fALFF, GFA, ISO, QA, rd1
  9. [9] § Results › Sensitivity Analysis ↔ clustering.py, lines 8–64 · score 0.53 · validation, multiclass, optimal, reproducibility, pipeline, Covariates

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 2,813 lines · 105 KB · no license · 4 matches

  1. # -*- coding: utf-8 -*-
  2. """
  3. Created on Fri Jan 16 22:18:56 2026
  4. @author: rglez
  5. """
  6. # -*- coding: utf-8 -*-
  7. """
  8. WMH PIPELINE COMPLETO (ARREGLADO)
  9. Includes EVERYTHING in a single integrated pipeline:
  10. 1. Excel loading and integrity checks
  11. 2. GLM/OLS residualization applied ONLY to clustering features (excluding Coord-x/y/z)
  12. 3. Z-score normalization
  13. 4. Automatic GMM clustering (optimal k selected by BIC on subset) + FULL dataset labels + maximum posterior probability (probmax)
  14. 5. Main Excel output including raw data, cluster summaries, and statistical tests
  15. 6. Visualization outputs: (BIC curve, PCA 3D, PCA 2D)
  16. 7. Stacked percentage frequency plots by diagnostic group (Gpo)
  17. 8. Interactive 3D HTML visualization using Plotly
  18. 9. Advanced group statistics: Global chi-square test, Cramér’s V effect size, Pairwise post-hoc comparisons, Standardized residuals per cell
  19. 10. Bubble matrix visualization displaying significant pairwise comparisons after FDR correction
  20. 11. Correlation analyses:
  21. • Spearman lower-triangle correlation matrices
  22. • Computed for ALL subjects and stratified by cluster
  23. • Optional partial correlations after residualization by covariates
  24. 12. Machine learning classification framework:
  25. • Repeated cross-validation
  26. • Out-of-fold ROC curves
  27. • Feature importance analysis
  28. • Robust SHAP explainability analysis
  29. 13. Interactive visualization support for exploratory and presentation purposes
  30. 14. Multiple-comparison correction procedures including FDR and Bonferroni adjustments
  31. 15. Cluster confidence estimation using maximum posterior probability (probmax)
  32. 16. Group-balanced statistical evaluation across diagnostic categories
  33. 17. Support for repeated measures and subject-level dependency correction
  34. 18. Robust multiclass LightGBM classification with explainability analyses (feature importance and SHAP)
  35. 19. Fully reproducible end-to-end framework for multimodal WMH clustering, statistical characterization, visualization, correlation analysis, and predictive modeling
  36. """
  37. import os
  38. import re
  39. import warnings
  40. import matplotlib as mpl
  41. import numpy as np
  42. import pandas as pd
  43. import matplotlib.pyplot as plt
  44. from matplotlib import cm, colors as mpl_colors
  45. from sklearn.preprocessing import StandardScaler
  46. from sklearn.decomposition import PCA
  47. from sklearn.mixture import GaussianMixture
  48. from sklearn.metrics import silhouette_score
  49. from sklearn.model_selection import StratifiedKFold, GroupKFold
  50. from sklearn.compose import ColumnTransformer
  51. from sklearn.pipeline import Pipeline
  52. from sklearn.preprocessing import OneHotEncoder
  53. from sklearn.impute import SimpleImputer
  54. from sklearn.metrics import roc_auc_score, roc_curve, auc
  55. from scipy.stats import kruskal, chi2_contingency, spearmanr, norm as normal_dist
  56. from statsmodels.stats.multitest import multipletests
  57. import plotly.graph_objects as go
  58. import plotly.io as pio
  59. from scipy.optimize import linear_sum_assignment
  60. # ---------------- LightGBM + SHAP ----------------
  61. try:
  62. from lightgbm import LGBMClassifier
  63. except Exception as e:
  64. raise ImportError("No puedo importar lightgbm. Instala con: pip install lightgbm") from e
  65. try:
  66. import shap
  67. _HAS_SHAP = True
  68. except Exception:
  69. _HAS_SHAP = False
  70. warnings.warn("No se pudo importar shap. SHAP se omitirá. Instala con: pip install shap")
  71. # ======================================================================================
  72. # ====================================== CONFIG ========================================
  73. # ======================================================================================
  74. excel_path = r"C:\Users\rglez\Documents\Ra\Papers\WMH long caracterization\Ampliacion\WMH_clusters_metrics-all_with_clinic_final.xlsx"
  75. out_dir = r"C:\Users\rglez\Documents\Ra\Papers\WMH long caracterization\Ampliacion\clustering"
  76. os.makedirs(out_dir, exist_ok=True)
  77. base_name = os.path.splitext(os.path.basename(excel_path))[0]
  78. excel_out_main = os.path.join(out_dir, f"{base_name}_CLUSTERING_GLM.xlsx")
  79. # Features para clustering (GLM SOLO a estas, excepto NO_GLM_COLS)
  80. requested_cols = [
  81. "T1", "GFA", "QA", "ISO", "ha", "ad", "fa", "rd", "rd1", "rd2",
  82. "fALFF", "Hurst", "Entropy",
  83. "Power slope", "Autocor",
  84. "WMH",
  85. "Coord-x", "Coord-y", "Coord-z",
  86. ]
  87. glm_covars = ["Age", "Sex", "WMH_t0", "Education", "Size", "TIV", "Site"]
  88. # Comparaciones post-cluster (sin corrección)
  89. compare_vars = ["Sex", "Age", "Coord-x", "Coord-y", "Coord-z", "∆WMH", "Gpo", "Size"]
  90. # NO residualizar coords
  91. NO_GLM_COLS = ["Coord-x", "Coord-y", "Coord-z"]
  92. # Clustering speed
  93. FIT_SUBSET_N = 5000
  94. K_RANGE = list(range(2, 11))
  95. SIL_SAMPLE = 2500
  96. DROP_NA_ROWS = False # si True: drop NaNs en features/covars; si False: error.
  97. # Correlaciones
  98. DO_CORR_BLOCKS = True
  99. vars_corr = [
  100. "WMH", "T1", "GFA", "QA", "ISO", "ha", "ad", "fa", "rd", "rd1", "rd2",
  101. "fALFF", "Hurst", "Entropy", "Power slope", "Autocor",
  102. ]
  103. DO_PARTIAL = True # residualiza vars_corr por covars
  104. CORR_COVARS = glm_covars[:] # usa las mismas covars
  105. CORR_ALPHA = 0.05
  106. CORR_USE_FDR = True
  107. CORR_FDR_METHOD = "fdr_bh"
  108. CORR_VMIN, CORR_VMAX = -0.5, 0.5
  109. CORR_CMAP = "bwr"
  110. CLUSTER_COL = "cluster_gmm_auto"
  111. # ML clasificación
  112. DO_ML = True
  113. TARGET = "cluster_gmm_auto"
  114. GPO_COL = "Gpo" # solo reportes
  115. SUBJ_COL = "subject_id" # si no existe, se crea desde el índice
  116. FEATURES = [
  117. "T1", "GFA", "QA", "ISO", "ha", "ad", "fa", "rd", "rd1", "rd2",
  118. "fALFF", "Hurst", "Entropy", "Power slope", "Autocor",
  119. "WMH",
  120. "Coord-x", "Coord-y", "Coord-z",
  121. "Age", "Sex", "WMH_t0", "Size", "Education", "TIV",
  122. ]
  123. N_SPLITS = 5
  124. N_REPEATS = 20
  125. BASE_SEED = 42
  126. RANDOM_STATE = 42
  127. N_BOOT = 1000
  128. GRID_N = 200
  129. SEED = 42
  130. DO_ROC_BY_GPO = True
  131. DO_FEATURE_IMPORTANCE_AND_SHAP = True
  132. MAX_SHAP_SAMPLES = 1500
  133. SAVE_SHAP_PER_CLASS = True
  134. # ======================================================================================
  135. # ===================================== HELPERS =======================================
  136. # ======================================================================================
  137. def pick_subset(n, subset_n, seed=42):
  138. rng = np.random.default_rng(seed)
  139. if n <= subset_n:
  140. return np.arange(n)
  141. return rng.choice(n, size=subset_n, replace=False)
  142. def safe_silhouette(X, labels, sample_size=2500, seed=42):
  143. labels = np.asarray(labels)
  144. if len(set(labels)) < 2:
  145. return np.nan
  146. ss = min(sample_size, len(labels))
  147. return float(silhouette_score(X, labels, sample_size=ss, random_state=seed))
  148. def plot_curve(x, y, title, xlabel, ylabel, out_png, out_pdf=None):
  149. plt.figure()
  150. plt.plot(x, y, marker="o")
  151. plt.title(title)
  152. plt.xlabel(xlabel)
  153. plt.ylabel(ylabel)
  154. plt.tight_layout()
  155. plt.savefig(out_png, dpi=300, bbox_inches="tight")
  156. # PDF (opcional)
  157. if out_pdf is not None:
  158. plt.savefig(out_pdf, bbox_inches="tight")
  159. plt.close()
  160. def to_subscript(n: int) -> str:
  161. return str(n).translate(str.maketrans("0123456789", "₀₁₂₃₄₅₆₇₈₉"))
  162. def lab_text(lab: int) -> str:
  163. return f"L{to_subscript(int(lab) + 1)}"
  164. def make_cluster_color_map(labels_or_uniq, cmap_name="viridis",
  165. pos_violet=0.10, pos_green=0.55, pos_yellow=0.95):
  166. """
  167. Fuerza:
  168. cluster 0 (L1) -> violeta
  169. cluster 1 (L2) -> verde
  170. cluster 2 (L3) -> amarillo
  171. """
  172. uniq = np.array(sorted(np.unique(np.asarray(labels_or_uniq).astype(int))))
  173. cmap = cm.get_cmap(cmap_name)
  174. pos_map = {0: pos_violet, 1: pos_green, 2: pos_yellow}
  175. # fallback si hubiera más clusters:
  176. default_positions = np.linspace(pos_violet, pos_yellow, len(uniq))
  177. out = {}
  178. for i, lab in enumerate(uniq):
  179. pos = pos_map.get(int(lab), float(default_positions[i]))
  180. out[int(lab)] = cmap(pos)
  181. return out
  182. def rgba_to_plotly_rgba(rgba_tuple):
  183. r, g, b, a = rgba_tuple
  184. return f"rgba({int(r*255)},{int(g*255)},{int(b*255)},{a:.3f})"
  185. def chi2_test(df, a, b):
  186. tab = pd.crosstab(df[a], df[b], dropna=False)
  187. if tab.shape[0] < 2 or tab.shape[1] < 2:
  188. return None
  189. chi2, p, dof, _ = chi2_contingency(tab, correction=False)
  190. return float(chi2), float(p), int(dof)
  191. def cluster_summary_table(df_in, label_col, compare_vars):
  192. df = df_in.copy()
  193. out = []
  194. numeric_cols = [c for c in compare_vars
  195. if c in df.columns and pd.api.types.is_numeric_dtype(df[c])]
  196. for cl in sorted(df[label_col].dropna().unique()):
  197. sub = df[df[label_col] == cl]
  198. n = len(sub)
  199. row = {"cluster": int(cl), "cluster_L": lab_text(int(cl)), "N": int(n)}
  200. if "Sex" in compare_vars and "Sex" in sub.columns:
  201. nF = int((sub["Sex"] == "F").sum())
  202. nM = int((sub["Sex"] == "M").sum())
  203. row.update({
  204. "Sex_F_n": nF,
  205. "Sex_M_n": nM,
  206. "Sex_F_%": (100.0 * nF / n) if n else np.nan,
  207. "Sex_M_%": (100.0 * nM / n) if n else np.nan,
  208. })
  209. for c in numeric_cols:
  210. row[f"{c}_mean"] = float(np.mean(sub[c]))
  211. row[f"{c}_sd"] = float(np.std(sub[c], ddof=1)) if n > 1 else np.nan
  212. out.append(row)
  213. return pd.DataFrame(out)
  214. def run_tests_original(df_in, label_col, compare_vars):
  215. df = df_in.copy()
  216. results = []
  217. for cat in ["Sex", "Gpo"]:
  218. if cat in compare_vars and cat in df.columns:
  219. out = chi2_test(df, label_col, cat)
  220. if out is not None:
  221. chi2, p, dof = out
  222. results.append({
  223. "label_col": label_col,
  224. "variable": cat,
  225. "test": "Chi-square",
  226. "stat": chi2,
  227. "df": dof,
  228. "p_value": p
  229. })
  230. numeric_cols = [c for c in compare_vars
  231. if c in df.columns and c != label_col and pd.api.types.is_numeric_dtype(df[c])]
  232. clusters = sorted(df[label_col].dropna().unique())
  233. if len(clusters) >= 2:
  234. for v in numeric_cols:
  235. groups = [df.loc[df[label_col] == k, v].to_numpy(dtype=float) for k in clusters]
  236. nonempty = [g for g in groups if len(g) > 0]
  237. if len(nonempty) >= 2:
  238. H, p = kruskal(*nonempty)
  239. results.append({
  240. "label_col": label_col,
  241. "variable": v,
  242. "test": "Kruskal-Wallis",
  243. "stat": float(H),
  244. "df": int(len(nonempty) - 1),
  245. "p_value": float(p)
  246. })
  247. res = pd.DataFrame(results)
  248. if not res.empty:
  249. res["p_fdr"] = multipletests(res["p_value"].values, method="fdr_bh")[1]
  250. res["p_bonf"] = multipletests(res["p_value"].values, method="bonferroni")[1]
  251. res = res.sort_values(["p_fdr", "p_value"])
  252. return res
  253. def freq_by_gpo(df_in, label_col, gpo_col="Gpo"):
  254. if (label_col not in df_in.columns) or (gpo_col not in df_in.columns):
  255. return pd.DataFrame()
  256. return pd.crosstab(df_in[label_col], df_in[gpo_col], dropna=False)
  257. def plot_freq_by_gpo_stacked(df_out, label_col, gpo_col, cluster_color,
  258. out_png, out_svg=None, out_pdf=None,
  259. bar_width=0.40, fig_w_per_group=0.90, fig_h=5.0):
  260. tab = pd.crosstab(df_out[label_col], df_out[gpo_col], dropna=False)
  261. if tab.shape[0] == 0 or tab.shape[1] == 0:
  262. raise ValueError("Tabla de contingencia vacía para frecuencias (revisa Gpo / clusters).")
  263. tab = tab.reindex([0, 1, 2], fill_value=0)
  264. desired_order = ["CN", "MCI", "AD", "PD"]
  265. present = [g for g in desired_order if g in tab.columns]
  266. others = [g for g in tab.columns if g not in present]
  267. tab = tab[present + others]
  268. col_sums = tab.sum(axis=0)
  269. tab_pct = tab.div(col_sums.replace(0, np.nan), axis=1) * 100
  270. if np.isfinite(tab_pct.to_numpy()).sum() == 0:
  271. raise ValueError("tab_pct quedó todo NaN (posible grupo con sum=0 o datos inválidos).")
  272. fig_w = max(3.0, fig_w_per_group * len(tab_pct.columns) + 0.6)
  273. fig, ax = plt.subplots(figsize=(fig_w, fig_h))
  274. bottom = np.zeros(len(tab_pct.columns), dtype=float)
  275. x = np.arange(len(tab_pct.columns)) # posiciones
  276. # ✅ apilar en orden ASC para que L1 quede abajo
  277. for lab in [2, 1, 0]:
  278. if lab not in tab_pct.index:
  279. continue
  280. vals = tab_pct.loc[lab].to_numpy(dtype=float)
  281. ax.bar(
  282. x, vals, bottom=bottom, width=bar_width,
  283. color=cluster_color[int(lab)],
  284. label=lab_text(int(lab))
  285. )
  286. bottom += np.nan_to_num(vals, nan=0.0)
  287. ax.set_ylim(0, 100)
  288. ax.set_ylabel("Frequency (% within group)", fontsize=14)
  289. ax.set_xticks(x)
  290. ax.set_xticklabels(tab_pct.columns, rotation=0, fontsize=16)
  291. ax.tick_params(axis="y", labelsize=12)
  292. handles, labels_legend = ax.get_legend_handles_labels()
  293. ax.legend(handles, labels_legend, title="Clusters",
  294. bbox_to_anchor=(1.02, 1), loc="upper left", frameon=False)
  295. fig.tight_layout()
  296. fig.savefig(out_png, dpi=300, bbox_inches="tight")
  297. if out_svg is not None:
  298. fig.savefig(out_svg, format="svg", bbox_inches="tight")
  299. if out_pdf is not None:
  300. fig.savefig(out_pdf, format="pdf", bbox_inches="tight")
  301. plt.show()
  302. plt.close(fig)
  303. print("[DONE] Frequency plot saved:", out_png)
  304. return tab, tab_pct
  305. def encode_sex(series):
  306. if pd.api.types.is_numeric_dtype(series):
  307. return series.astype(float)
  308. sex_map = {"F": 0, "M": 1, "f": 0, "m": 1}
  309. out = series.map(sex_map)
  310. if out.isna().any():
  311. bad = series[out.isna()].unique()
  312. raise ValueError(f"Unexpected Sex values: {bad}. Expected F/M.")
  313. return out.astype(float)
  314. def residualize_matrix(Y, cov_df):
  315. Y = np.asarray(Y, dtype=float)
  316. if Y.ndim == 1:
  317. Y = Y.reshape(-1, 1)
  318. C = cov_df.to_numpy(dtype=float)
  319. Xcov = np.column_stack([np.ones(len(cov_df)), C])
  320. B = np.linalg.lstsq(Xcov, Y, rcond=None)[0]
  321. return Y - (Xcov @ B)
  322. def stars_from_p(p):
  323. if not np.isfinite(p):
  324. return ""
  325. if p < 0.001:
  326. return "***"
  327. if p < 0.01:
  328. return "**"
  329. if p < 0.05:
  330. return "*"
  331. return ""
  332. # ======================================================================================
  333. # ====================================== 1) LOAD ======================================
  334. # ======================================================================================
  335. print(f"[INFO] Reading Excel: {excel_path}")
  336. df = pd.read_excel(excel_path)
  337. df["WMH"] = df["∆WMH"] / df["Follow-up time"]
  338. df["T1"] = df["∆T1"] / df["Follow-up time"]
  339. df["QA"] = df["∆QA"] / df["Follow-up time"]
  340. df["GFA"] = df["∆GFA"] / df["Follow-up time"]
  341. df["ISO"] = df["∆ISO"] / df["Follow-up time"]
  342. df["ha"] = df["∆ha"] / df["Follow-up time"]
  343. df["ad"] = df["∆ad"] / df["Follow-up time"]
  344. df["fa"] = df["∆fa"] / df["Follow-up time"]
  345. df["rd"] = df["∆rd"] / df["Follow-up time"]
  346. df["rd1"] = df["∆rd1"] / df["Follow-up time"]
  347. df["rd2"] = df["∆rd2"] / df["Follow-up time"]
  348. df["fALFF"] = df["∆fALFF"] / df["Follow-up time"]
  349. df["Hurst"] = df["∆Hurst"] / df["Follow-up time"]
  350. df["Entropy"] = df["∆Entropy"] / df["Follow-up time"]
  351. df["Power slope"] = df["∆Power slope"] / df["Follow-up time"]
  352. df["Autocor"] = df["∆Autocor"] / df["Follow-up time"]
  353. needed = list(set(requested_cols + glm_covars + compare_vars + FEATURES + [GPO_COL]))
  354. missing_needed = [c for c in needed if c not in df.columns and c != SUBJ_COL]
  355. if missing_needed:
  356. raise ValueError(f"Missing required columns in Excel: {missing_needed}")
  357. # Si no existe subject_id, créalo
  358. if SUBJ_COL not in df.columns:
  359. df[SUBJ_COL] = np.arange(len(df)).astype(int)
  360. use_cols = [c for c in requested_cols if c in df.columns]
  361. if len(use_cols) < 2:
  362. raise ValueError("Not enough clustering features found in Excel.")
  363. check_cols = use_cols + glm_covars
  364. na_rows = df[check_cols].isna().any(axis=1)
  365. n_na = int(na_rows.sum())
  366. if n_na > 0:
  367. msg = f"[ERROR] Found {n_na} rows with NaNs in clustering features/covars. "
  368. if DROP_NA_ROWS:
  369. print(msg + "Dropping those rows (DROP_NA_ROWS=True).")
  370. df = df.loc[~na_rows].reset_index(drop=True)
  371. else:
  372. raise ValueError(msg + "Set DROP_NA_ROWS=True to drop them, or fix the Excel.")
  373. print(f"[INFO] Rows after NaN handling: {len(df)}")
  374. # ======================================================================================
  375. # ===================== 2) GLM / OLS correction (solo features) ========================
  376. # ======================================================================================
  377. glm_features = [c for c in use_cols if c not in NO_GLM_COLS]
  378. raw_features = [c for c in use_cols if c in NO_GLM_COLS]
  379. if len(glm_features) + len(raw_features) < 2:
  380. raise ValueError("Not enough features after GLM/RAW split.")
  381. print("[INFO] GLM features (corrected):", glm_features)
  382. print("[INFO] RAW features (NOT corrected):", raw_features)
  383. cov_df = df[glm_covars].copy()
  384. if "Sex" in glm_covars:
  385. cov_df["Sex"] = encode_sex(cov_df["Sex"])
  386. C = cov_df[glm_covars].to_numpy(dtype=float)
  387. Xcov = np.column_stack([np.ones(len(df)), C])
  388. # GLM residualize ONLY glm_features
  389. if len(glm_features) > 0:
  390. X_glm_raw = df[glm_features].to_numpy(dtype=float)
  391. B = np.linalg.lstsq(Xcov, X_glm_raw, rcond=None)[0]
  392. X_glm_resid = X_glm_raw - (Xcov @ B)
  393. X_glm_z = StandardScaler().fit_transform(X_glm_resid)
  394. else:
  395. X_glm_z = None
  396. # RAW part: z-score only
  397. if len(raw_features) > 0:
  398. X_rawpart = df[raw_features].to_numpy(dtype=float)
  399. X_raw_z = StandardScaler().fit_transform(X_rawpart)
  400. else:
  401. X_raw_z = None
  402. # Combine for clustering
  403. if X_glm_z is not None and X_raw_z is not None:
  404. X_z = np.column_stack([X_glm_z, X_raw_z])
  405. elif X_glm_z is not None:
  406. X_z = X_glm_z
  407. elif X_raw_z is not None:
  408. X_z = X_raw_z
  409. else:
  410. raise RuntimeError("No features for clustering after processing.")
  411. print("[INFO] X_z shape:", X_z.shape)
  412. # ======================================================================================
  413. # ===================== 3) All lesion for k selection/plots + GMM ======================
  414. # ======================================================================================
  415. n = len(df)
  416. fit_idx = pick_subset(n, FIT_SUBSET_N, seed=42)
  417. X_fit = X_z[fit_idx]
  418. print(f"[INFO] Total rows: {n} | subset for selection/plots: {len(fit_idx)}")
  419. print("\n=== [GMM] Auto-k by min BIC ===")
  420. gmm_bics = []
  421. gmm_models = {}
  422. for k in K_RANGE:
  423. gmm = GaussianMixture(
  424. n_components=k,
  425. covariance_type="full",
  426. random_state=42,
  427. n_init=2,
  428. max_iter=300
  429. )
  430. gmm.fit(X_fit)
  431. bic = gmm.bic(X_fit)
  432. gmm_bics.append(bic)
  433. gmm_models[k] = gmm
  434. print(f" k={k:2d} | BIC={bic:.0f}")
  435. best_k_gmm = K_RANGE[int(np.argmin(gmm_bics))]
  436. gmm_best = gmm_models[best_k_gmm]
  437. print(f"[GMM] Best k = {best_k_gmm}")
  438. bic_png = os.path.join(out_dir, f"{base_name}_GMM_bic.png")
  439. bic_pdf = os.path.join(out_dir, f"{base_name}_GMM_bic.pdf")
  440. plot_curve(
  441. K_RANGE, gmm_bics,
  442. title=f"GMM: BIC vs k (best k={best_k_gmm})",
  443. xlabel="k",
  444. ylabel="BIC",
  445. out_png=bic_png,
  446. out_pdf=bic_pdf
  447. )
  448. print("[DONE] BIC plot:", bic_png)
  449. labels_gmm_full = gmm_best.predict(X_z).astype(int)
  450. labels_gmm_fit = gmm_best.predict(X_fit).astype(int)
  451. gmm_probmax_full = gmm_best.predict_proba(X_z).max(axis=1)
  452. # ===================== RE-LABEL: swap 1 <-> 2 (keep 0) =====================
  453. map_swap = {0: 2, 1: 0, 2: 1}
  454. labels_gmm_full = np.vectorize(map_swap.get)(labels_gmm_full)
  455. labels_gmm_fit = np.vectorize(map_swap.get)(labels_gmm_fit)
  456. # ======================================================================================
  457. # ===================== SAVE TRAINED MODEL =============================================
  458. # ======================================================================================
  459. import joblib
  460. print("\n[INFO] Saving trained clustering model...")
  461. # ----------------------------------------------------------
  462. # SAVE SCALERS
  463. # ----------------------------------------------------------
  464. scaler_glm = None
  465. scaler_raw = None
  466. if len(glm_features) > 0:
  467. scaler_glm = StandardScaler()
  468. scaler_glm.fit(X_glm_resid)
  469. if len(raw_features) > 0:
  470. scaler_raw = StandardScaler()
  471. scaler_raw.fit(X_rawpart)
  472. # ----------------------------------------------------------
  473. # SAVE EVERYTHING NEEDED FOR EXTERNAL APPLICATION
  474. # ----------------------------------------------------------
  475. model_bundle = {
  476. # ================= GMM =================
  477. "gmm_model": gmm_best,
  478. # ================= FEATURES =================
  479. "feature_order": glm_features + raw_features,
  480. "glm_features": glm_features,
  481. "raw_features": raw_features,
  482. # ================= COVARIATES =================
  483. "glm_covars": glm_covars,
  484. "NO_GLM_COLS": NO_GLM_COLS,
  485. # ================= GLM BETAS =================
  486. "glm_betas": B if len(glm_features) > 0 else None,
  487. # ================= SCALERS =================
  488. "scaler_glm": scaler_glm,
  489. "scaler_raw": scaler_raw,
  490. # ================= LABEL MAPPING =================
  491. "map_swap": map_swap,
  492. # ================= CONFIG =================
  493. "requested_cols": requested_cols,
  494. "vars_corr": vars_corr,
  495. # ================= INFO =================
  496. "best_k": best_k_gmm,
  497. "n_features": X_z.shape[1],
  498. }
  499. # ----------------------------------------------------------
  500. # OUTPUT FILE
  501. # ----------------------------------------------------------
  502. model_out = os.path.join(
  503. out_dir,
  504. f"{base_name}_trained_GMM_model.pkl"
  505. )
  506. # ----------------------------------------------------------
  507. # SAVE
  508. # ----------------------------------------------------------
  509. joblib.dump(model_bundle, model_out)
  510. print("[DONE] Trained model saved:")
  511. print(model_out)
  512. # ======================================================================================
  513. # ===================== 3b) GMM PARAMETERS: COMPONENT MEANS + TOP FEATURES + HEATMAP ===
  514. # ======================================================================================
  515. feature_names = glm_features + raw_features # mismo orden que X_z
  516. means = pd.DataFrame(gmm_best.means_, columns=feature_names)
  517. means_swapped = means.copy()
  518. means_swapped.index = [map_swap[i] for i in range(means.shape[0])]
  519. means_swapped = means_swapped.sort_index() # 0,1,2 => L1,L2,L3
  520. # ---- (C) Reordenar columnas: ∆WMH primero y Follow-up time al final ----
  521. cols = list(means_swapped.columns)
  522. # 1) sacar Follow-up time (si está) para agregarlo al final después
  523. follow = ["Follow-up time"] if "Follow-up time" in cols else []
  524. cols = [c for c in cols if c != "Follow-up time"]
  525. # 2) poner ∆WMH primero (si está)
  526. if "WMH" in cols:
  527. cols = ["WMH"] + [c for c in cols if c != "WMH"]
  528. # 3) agregar Follow-up time al final
  529. cols = cols + follow
  530. means_swapped = means_swapped[cols]
  531. # ---- (D) Top features por cluster según |mean| (en z) ----
  532. TOPN = 8
  533. top_list = []
  534. for cl in means_swapped.index:
  535. s = means_swapped.loc[cl] # serie (features) para ese cluster
  536. top = s.abs().sort_values(ascending=False).head(TOPN).index
  537. tmp = pd.DataFrame({
  538. "cluster": int(cl),
  539. "cluster_L": lab_text(int(cl)), # L₁, L₂, L₃...
  540. "feature": top,
  541. "mean_z": s.loc[top].values
  542. })
  543. top_list.append(tmp)
  544. top_df = pd.concat(top_list, ignore_index=True)
  545. # ---- (E) Exportar a Excel ----
  546. out_xlsx = os.path.join(out_dir, f"{base_name}_GMM_component_means.xlsx")
  547. with pd.ExcelWriter(out_xlsx, engine="openpyxl") as writer:
  548. means_swapped.to_excel(writer, sheet_name="component_means_z")
  549. top_df.to_excel(writer, index=False, sheet_name="top_features_by_cluster")
  550. print("[DONE] GMM component means saved:", out_xlsx)
  551. # ---- (F) Heatmap de medias por cluster (z) - VECTOR (PDF) ----
  552. # Para que el texto sea editable en PDF (Illustrator/Inkscape)
  553. mpl.rcParams["pdf.fonttype"] = 42
  554. mpl.rcParams["ps.fonttype"] = 42
  555. M = means_swapped.copy()
  556. M.index = [lab_text(int(i)) for i in M.index] # L₁, L₂, L₃...
  557. # Tamaño base original:
  558. # height_old = 0.8 * M.shape[0] + 2
  559. # Nuevo: 0.9 del tamaño actual
  560. width = 0.55 * M.shape[1] + 5
  561. height = 0.9 * (0.8 * M.shape[0] + 2)
  562. fig, ax = plt.subplots(figsize=(width, height))
  563. # pcolormesh = celdas vectoriales en PDF
  564. x = np.arange(M.shape[1] + 1)
  565. y = np.arange(M.shape[0] + 1)
  566. # ===== Symmetric z-score scale =====
  567. absmax = np.nanmax(np.abs(M.values))
  568. mesh = ax.pcolormesh(
  569. x,
  570. y,
  571. M.values,
  572. cmap="bwr",
  573. vmin=-0.4,
  574. vmax= 0.4,
  575. shading="flat",
  576. edgecolors="none")
  577. # Ticks centrados en cada celda
  578. ax.set_xticks(np.arange(M.shape[1]) + 0.5)
  579. ax.set_yticks(np.arange(M.shape[0]) + 0.5)
  580. ax.set_xticklabels(M.columns, rotation=60, ha="right", fontsize=10)
  581. ax.set_yticklabels(M.index, fontsize=12)
  582. # Para que la primera fila quede arriba (como imshow)
  583. ax.invert_yaxis()
  584. # Limpio (opcional)
  585. for spine in ax.spines.values():
  586. spine.set_visible(False)
  587. cbar = fig.colorbar(mesh, ax=ax, pad=0.01, fraction=0.06)
  588. cbar.set_label("Component mean (z)", fontsize=12)
  589. cbar.ax.tick_params(labelsize=10)
  590. ax.set_title("") # sin título si quieres
  591. fig.tight_layout()
  592. out_png = os.path.join(out_dir, f"{base_name}_GMM_componentMeans_heatmap.png")
  593. out_pdf = os.path.join(out_dir, f"{base_name}_GMM_componentMeans_heatmap.pdf")
  594. # PNG (raster) para vista rápida
  595. fig.savefig(out_png, dpi=300, bbox_inches="tight")
  596. # PDF vector (las celdas quedan vectorizadas)
  597. fig.savefig(out_pdf, bbox_inches="tight")
  598. plt.show()
  599. plt.close(fig)
  600. print("[DONE] Heatmap:", out_png, out_pdf)
  601. # ======================================================================================
  602. # ===================== 4) DF OUTPUT + Summary + Tests =================================
  603. # ======================================================================================
  604. df_out = df.copy()
  605. df_out[CLUSTER_COL] = labels_gmm_full
  606. df_out["cluster_gmm_probmax"] = gmm_probmax_full
  607. gmm_summary = cluster_summary_table(df_out, CLUSTER_COL, compare_vars)
  608. gmm_tests = run_tests_original(df_out, CLUSTER_COL, compare_vars)
  609. gmm_gpo_freq = freq_by_gpo(df_out, CLUSTER_COL, GPO_COL) if GPO_COL in df_out.columns else pd.DataFrame()
  610. with pd.ExcelWriter(excel_out_main, engine="openpyxl") as writer:
  611. df_out.to_excel(writer, index=False, sheet_name="data_with_clusters")
  612. gmm_summary.to_excel(writer, index=False, sheet_name="gmm_cluster_summary")
  613. if not gmm_tests.empty:
  614. gmm_tests.to_excel(writer, index=False, sheet_name="cluster_tests")
  615. if not gmm_gpo_freq.empty:
  616. gmm_gpo_freq.to_excel(writer, sheet_name="gmm_freq_by_Gpo")
  617. print("[DONE] Excel main:", excel_out_main)
  618. # ======================================================================================
  619. # ===================== 5) COLORS CONSISTENTES (matplotlib+plotly) ======================
  620. # ======================================================================================
  621. uniq = np.array(sorted(np.unique(labels_gmm_fit).astype(int)))
  622. cluster_color = make_cluster_color_map(
  623. uniq,
  624. cmap_name="viridis",
  625. pos_violet=0.10,
  626. pos_green=0.55,
  627. pos_yellow=0.95
  628. )
  629. # ======================================================================================
  630. # ===================== 6) PCA plots ===================================================
  631. # ======================================================================================
  632. pc3 = PCA(n_components=3, random_state=42).fit_transform(X_fit)
  633. labels = labels_gmm_fit
  634. views = [(20, 35), (20, 120), (20, 210), (60, 35)]
  635. for elev, azim in views:
  636. fig = plt.figure(figsize=(8, 6))
  637. ax = fig.add_subplot(111, projection="3d")
  638. for lab in uniq:
  639. m = labels == lab
  640. ax.scatter(
  641. pc3[m, 0], pc3[m, 1], pc3[m, 2],
  642. s=10, alpha=0.65,
  643. color=cluster_color[int(lab)],
  644. label=lab_text(int(lab))
  645. )
  646. ax.set_xlabel("PC1")
  647. ax.set_ylabel("PC2")
  648. ax.set_zlabel("PC3")
  649. ax.set_title(f"GMM (subset) k={len(uniq)} | elev={elev} azim={azim}")
  650. ax.view_init(elev=elev, azim=azim)
  651. ax.legend(loc="upper right", title="Clusters", fontsize=9)
  652. out_png = os.path.join(out_dir, f"{base_name}_GMM_PCA3D_e{elev}_a{azim}.png")
  653. out_pdf = os.path.join(out_dir, f"{base_name}_GMM_PCA3D_e{elev}_a{azim}.pdf")
  654. fig.savefig(out_png, dpi=300, bbox_inches="tight")
  655. fig.savefig(out_pdf, format="pdf", bbox_inches="tight")
  656. plt.show()
  657. plt.close(fig)
  658. print("[DONE] PCA3D:", out_png)
  659. pairs = [(0, 1, "PC1", "PC2"), (0, 2, "PC1", "PC3"), (1, 2, "PC2", "PC3")]
  660. # ✅ orden de pintado: L1 -> L2 -> L3 (L3 al final para que quede arriba)
  661. plot_order = [2, 1, 0]
  662. plot_order = [lab for lab in plot_order if lab in set(uniq)] # por si falta alguno
  663. for a, b, xa, xb in pairs:
  664. fig, ax = plt.subplots(figsize=(7, 5))
  665. for lab in plot_order:
  666. m = (labels == lab)
  667. ax.scatter(
  668. pc3[m, a], pc3[m, b],
  669. s=12, alpha=0.65,
  670. color=cluster_color[int(lab)],
  671. label=lab_text(int(lab))
  672. )
  673. ax.set_xlabel(xa)
  674. ax.set_ylabel(xb)
  675. ax.set_title(f"GMM (subset) {xa} vs {xb} | k={len(uniq)}")
  676. ax.legend(title="Clusters", fontsize=9)
  677. out_png = os.path.join(out_dir, f"{base_name}_GMM_{xa}_vs_{xb}.png")
  678. out_pdf = os.path.join(out_dir, f"{base_name}_GMM_{xa}_vs_{xb}.pdf")
  679. fig.savefig(out_png, dpi=300, bbox_inches="tight")
  680. fig.savefig(out_pdf, format="pdf", bbox_inches="tight")
  681. plt.show()
  682. plt.close(fig)
  683. print("[DONE] PCA2D:", out_png)
  684. # ======================================================================================
  685. # ===================== 7) FRECUENCIAS por Gpo (STACKED %) ==============================
  686. # ======================================================================================
  687. if GPO_COL in df_out.columns:
  688. freq_png = os.path.join(out_dir, f"{base_name}_freq.png")
  689. freq_svg = os.path.join(out_dir, f"{base_name}_freq.svg")
  690. freq_pdf = os.path.join(out_dir, f"{base_name}_freq.pdf")
  691. tab_counts, tab_pct = plot_freq_by_gpo_stacked(
  692. df_out, CLUSTER_COL, GPO_COL, cluster_color,
  693. out_png=freq_png, out_svg=freq_svg, out_pdf=freq_pdf,
  694. bar_width=0.6, fig_w_per_group=0.85, fig_h=5
  695. )
  696. # ✅ Asegurar orden L1, L2, L3... en tablas/export (filas = clusters)
  697. tab_counts = tab_counts.reindex([2, 1, 0])
  698. tab_pct = tab_pct.reindex([2, 1, 0])
  699. # ---- Stats avanzadas + Excel ----
  700. obs = tab_counts.values
  701. chi2, p, dof, expected = chi2_contingency(obs, correction=False)
  702. N = obs.sum()
  703. k_eff = min(tab_counts.shape[0] - 1, tab_counts.shape[1] - 1)
  704. cramers_v_global = np.sqrt(chi2 / (N * k_eff)) if (N > 0 and k_eff > 0) else np.nan
  705. global_summary = pd.DataFrame({
  706. "chi2": [float(chi2)],
  707. "df": [int(dof)],
  708. "p": [float(p)],
  709. "N": [int(N)],
  710. "Cramers_V": [float(cramers_v_global)]
  711. })
  712. cols = tab_counts.columns.tolist()
  713. pairs_list, stats_list = [], []
  714. for i in range(len(cols)):
  715. for j in range(i + 1, len(cols)):
  716. g1, g2 = cols[i], cols[j]
  717. sub = tab_counts[[g1, g2]].values
  718. x2, pp, dfp, _ = chi2_contingency(sub, correction=False)
  719. N2 = sub.sum()
  720. k2 = min(sub.shape[0] - 1, sub.shape[1] - 1)
  721. v_pair = np.sqrt(x2 / (N2 * k2)) if (N2 > 0 and k2 > 0) else np.nan
  722. pairs_list.append((g1, g2))
  723. stats_list.append((float(x2), int(dfp), float(pp), float(v_pair), int(N2)))
  724. pair_df = pd.DataFrame(
  725. stats_list,
  726. columns=["chi2", "df", "p", "Cramers_V", "N"],
  727. index=[f"{a} vs {b}" for a, b in pairs_list],
  728. )
  729. if not pair_df.empty:
  730. pair_df["p_fdr"] = multipletests(pair_df["p"].values, method="fdr_bh")[1]
  731. pair_df["p_bonf"] = multipletests(pair_df["p"].values, method="bonferroni")[1]
  732. pair_df = pair_df.sort_values("p_fdr")
  733. # Residuales estandarizados por celda (usa expected de chi2_contingency)
  734. exp = expected
  735. row_sum = obs.sum(axis=1, keepdims=True)
  736. col_sum = obs.sum(axis=0, keepdims=True)
  737. row_prop = row_sum / N if N else np.nan
  738. col_prop = col_sum / N if N else np.nan
  739. with np.errstate(divide="ignore", invalid="ignore"):
  740. std_res = (obs - exp) / np.sqrt(exp * (1 - row_prop) * (1 - col_prop))
  741. std_res_df = pd.DataFrame(std_res, index=tab_counts.index, columns=tab_counts.columns)
  742. p_cell = 2 * normal_dist.sf(np.abs(std_res))
  743. p_cell_adj = multipletests(np.ravel(p_cell), method="fdr_bh")[1].reshape(p_cell.shape)
  744. p_cell_df = pd.DataFrame(p_cell, index=tab_counts.index, columns=tab_counts.columns)
  745. p_cell_fdr_df = pd.DataFrame(p_cell_adj, index=tab_counts.index, columns=tab_counts.columns)
  746. out_xlsx = os.path.join(out_dir, f"{base_name}_GMM_cluster_stats.xlsx")
  747. with pd.ExcelWriter(out_xlsx, engine="openpyxl") as writer:
  748. global_summary.to_excel(writer, sheet_name="global_stats", index=False)
  749. if not pair_df.empty:
  750. pair_df.to_excel(writer, sheet_name="pairwise_stats")
  751. tab_counts.to_excel(writer, sheet_name="counts_numeric")
  752. tab_pct.to_excel(writer, sheet_name="percent_by_group_numeric")
  753. std_res_df.to_excel(writer, sheet_name="std_residuals")
  754. p_cell_df.to_excel(writer, sheet_name="p_cell")
  755. p_cell_fdr_df.to_excel(writer, sheet_name="p_cell_fdr")
  756. print("[DONE] Gpo stats Excel:", out_xlsx)
  757. else:
  758. pair_df = pd.DataFrame() # para bubble matrix si no hay Gpo
  759. # ======================================================================================
  760. # ===================== 8) PROBMAX: mean + SEM por cluster =============================
  761. # ======================================================================================
  762. probmax_xlsx = os.path.join(out_dir, f"{base_name}_GMM_probmax_by_cluster.xlsx")
  763. probmax_png = os.path.join(out_dir, f"{base_name}_GMM_probmax_by_cluster.png")
  764. tmp = df_out[[CLUSTER_COL, "cluster_gmm_probmax"]].copy()
  765. tmp[CLUSTER_COL] = tmp[CLUSTER_COL].astype(int)
  766. summary_prob = (
  767. tmp.groupby(CLUSTER_COL)["cluster_gmm_probmax"]
  768. .agg(["count", "mean", "std"])
  769. .rename(columns={"count": "N", "mean": "probmax_mean", "std": "probmax_sd"})
  770. .reset_index()
  771. .sort_values(CLUSTER_COL)
  772. .reset_index(drop=True)
  773. )
  774. summary_prob["probmax_sem"] = summary_prob["probmax_sd"] / np.sqrt(summary_prob["N"].clip(lower=1))
  775. clusters_prob = summary_prob[CLUSTER_COL].to_numpy()
  776. bar_colors = [cluster_color[int(c)] for c in clusters_prob]
  777. plt.figure(figsize=(3.4 + 0.35 * len(clusters_prob), 3.2))
  778. x = np.arange(len(clusters_prob))
  779. y = summary_prob["probmax_mean"].to_numpy(float)
  780. yerr = summary_prob["probmax_sem"].to_numpy(float)
  781. plt.bar(x, y, yerr=yerr, capsize=3, color=bar_colors)
  782. plt.xticks(x, [lab_text(int(c)) for c in clusters_prob], rotation=0)
  783. plt.xlabel("Cluster")
  784. plt.ylabel("Mean max posterior prob")
  785. plt.ylim(0, 1.0)
  786. plt.tight_layout()
  787. plt.savefig(probmax_png, dpi=300, bbox_inches="tight")
  788. plt.show()
  789. plt.close()
  790. print("[DONE] Probmax plot:", probmax_png)
  791. with pd.ExcelWriter(probmax_xlsx, engine="openpyxl") as writer:
  792. summary_prob.to_excel(writer, index=False, sheet_name="probmax_summary")
  793. tmp.to_excel(writer, index=False, sheet_name="probmax_all_rows")
  794. print("[DONE] Probmax Excel:", probmax_xlsx)
  795. # ======================================================================================
  796. # ===================== 9) INTERACTIVE 3D (HTML) =======================================
  797. # ======================================================================================
  798. fig = go.Figure()
  799. for lab in uniq:
  800. m = labels == lab
  801. fig.add_trace(go.Scatter3d(
  802. x=pc3[m, 0], y=pc3[m, 1], z=pc3[m, 2],
  803. mode="markers",
  804. name=lab_text(int(lab)),
  805. marker=dict(
  806. size=3,
  807. opacity=0.65,
  808. color=rgba_to_plotly_rgba(cluster_color[int(lab)])
  809. )
  810. ))
  811. fig.update_layout(
  812. title=f"GMM clusters (subset) - k={len(uniq)} (interactive)",
  813. scene=dict(xaxis_title="PC1", yaxis_title="PC2", zaxis_title="PC3"),
  814. legend=dict(title="Clusters")
  815. )
  816. html_out = os.path.join(out_dir, f"{base_name}_GMM_PCA3D_interactive.html")
  817. pio.write_html(fig, file=html_out, auto_open=False, include_plotlyjs="cdn")
  818. print("[DONE] Interactive 3D saved:", html_out)
  819. # ======================================================================================
  820. # ===================== 10) BUBBLE MATRIX (SIG ONLY) ===================================
  821. # ======================================================================================
  822. def get_pair_row(a, b, pair_df):
  823. k1 = f"{a} vs {b}"
  824. k2 = f"{b} vs {a}"
  825. if k1 in pair_df.index:
  826. return pair_df.loc[k1]
  827. if k2 in pair_df.index:
  828. return pair_df.loc[k2]
  829. return None
  830. def area_to_diameter_pt(s_area):
  831. return 2.0 * np.sqrt(np.array(s_area, dtype=float) / np.pi)
  832. def map_to_sizes(raw, smin, smax):
  833. raw = np.asarray(raw, float)
  834. m = np.isfinite(raw)
  835. if m.sum() == 0:
  836. return np.array([], float), (np.nan, np.nan, np.nan)
  837. x = raw[m]
  838. lo, hi = (np.percentile(x, [5, 95]) if len(x) > 1 else (x.min(), x.max() + 1e-9))
  839. denom = (hi - lo) if (hi - lo) > 1e-12 else 1.0
  840. z = np.clip((raw - lo) / denom, 0, 1)
  841. return (smin + (smax - smin) * z), (lo, hi, denom)
  842. def ensure_group_order_from_df(df_out, group_col="Gpo"):
  843. return list(pd.unique(df_out[group_col].dropna()))
  844. if not pair_df.empty:
  845. # ---------------- CONFIG ----------------
  846. P_COL = "p_fdr"
  847. ALPHA = 0.05
  848. SIZE_MODE = "cramers_v" # "chi2" o "cramers_v"
  849. # ✅ lo que pediste:
  850. # - eliminar la FILA CN pero NO su columna
  851. # - eliminar la COLUMNA PD pero NO su fila
  852. DROP_ROWS = {"CN"}
  853. DROP_COLS = {"PD"}
  854. S_MIN, S_MAX = 650, 4200
  855. V_MAX = 6.0
  856. LEG_BASE_X = 0.70
  857. LEG_BASE_Y = -0.20
  858. LEG_GAP_PT = 30.0
  859. LEG_SCALE = 0.85
  860. BOTTOM_MARGIN = 0.30
  861. LEG_EDGE = "#666666"
  862. CHI2_LEG = np.array([15, 25, 75], dtype=float)
  863. CHI2_TEXT_SHIFT = 5.0
  864. V_LEG = np.array([0.10, 0.20, 0.40], dtype=float)
  865. V_TEXT_FMT = "{:.2f}"
  866. # ---------------- GROUPS ----------------
  867. groups_all = ensure_group_order_from_df(df_out, group_col=GPO_COL)
  868. rows = [g for g in groups_all if g not in DROP_ROWS] # CN se va solo de filas
  869. cols = [g for g in groups_all if g not in DROP_COLS] # PD se va solo de columnas
  870. R, C = len(rows), len(cols)
  871. if R == 0 or C == 0:
  872. raise ValueError("Después de DROP_ROWS/DROP_COLS, no quedan filas/columnas.")
  873. # ---------------- COLLECT SIGNIFICANT CELLS ----------------
  874. xs, ys, pvals = [], [], []
  875. chi2vals, vvals = [], []
  876. for i in range(R):
  877. for j in range(C):
  878. rr, cc = rows[i], cols[j]
  879. if rr == cc:
  880. continue
  881. # ✅ robusto: mantener solo triángulo inferior respecto al orden global
  882. if groups_all.index(rr) < groups_all.index(cc):
  883. continue
  884. r = get_pair_row(rr, cc, pair_df)
  885. if r is None:
  886. continue
  887. pv = float(r[P_COL]) if P_COL in r.index else np.nan
  888. if np.isfinite(pv) and pv < ALPHA:
  889. xs.append(j)
  890. ys.append(i)
  891. pvals.append(pv)
  892. chi2vals.append(float(r["chi2"]) if "chi2" in r.index else np.nan)
  893. if "Cramers_V" in r.index:
  894. vvals.append(float(r["Cramers_V"]))
  895. else:
  896. vvals.append(np.nan)
  897. pvals = np.array(pvals, float)
  898. chi2vals = np.array(chi2vals, float)
  899. vvals = np.array(vvals, float)
  900. if len(pvals) == 0:
  901. print("[WARN] No hay pares significativos (p < ALPHA). Bubble matrix sin burbujas.")
  902. # ---------------- SIZE DRIVER ----------------
  903. if SIZE_MODE == "cramers_v" and np.any(np.isfinite(vvals)):
  904. size_label = "V"
  905. driver_raw = np.clip(vvals, 0, None)
  906. else:
  907. size_label = "χ²"
  908. driver_raw = np.sqrt(np.clip(chi2vals, 0, None)) if len(chi2vals) else np.array([])
  909. sizes, _ = map_to_sizes(driver_raw, S_MIN, S_MAX)
  910. # ---------------- COLORS (by p) ----------------
  911. val_color = -np.log10(np.clip(pvals, 1e-300, 1.0)) if len(pvals) else np.array([])
  912. cmap_b = mpl_colors.LinearSegmentedColormap.from_list("white_to_red", ["#ffffff", "#ff0000"])
  913. vmin = -np.log10(10) # 1
  914. vmax = V_MAX
  915. norm = mpl_colors.Normalize(vmin=vmin, vmax=vmax)
  916. colors_rgba = cmap_b(norm(np.clip(val_color, vmin, vmax))) if len(val_color) else np.array([])
  917. # ---------------- PLOT ----------------
  918. fig, ax = plt.subplots(figsize=(7.2, 6.0))
  919. ax.set_xlim(-0.5, C - 0.5)
  920. ax.set_ylim(R - 0.5, -0.5)
  921. ax.set_xticks(np.arange(C))
  922. ax.set_yticks(np.arange(R))
  923. ax.set_xticklabels(cols, rotation=35, ha="right")
  924. ax.set_yticklabels(rows)
  925. ax.set_aspect("equal", adjustable="box")
  926. ax.set_xticks(np.arange(-.5, C, 1), minor=True)
  927. ax.set_yticks(np.arange(-.5, R, 1), minor=True)
  928. ax.grid(which="minor", color="#dddddd", linestyle="-", linewidth=1)
  929. ax.tick_params(which="minor", bottom=False, left=False)
  930. for spine in ax.spines.values():
  931. spine.set_visible(False)
  932. if len(xs) > 0:
  933. ax.scatter(xs, ys, s=sizes, c=colors_rgba, edgecolors="none", linewidths=0)
  934. # colorbar
  935. sm = plt.cm.ScalarMappable(cmap=cmap_b, norm=norm)
  936. sm.set_array([])
  937. cbar = fig.colorbar(sm, ax=ax, fraction=0.046, pad=0.04)
  938. p_start = ALPHA
  939. p_end = 10 ** (-vmax)
  940. cbar.set_ticks([vmin, vmax])
  941. cbar.set_ticklabels([f"{p_start:g}", f"{p_end:g}"])
  942. # ---------------- SIZE LEGEND ----------------
  943. def build_size_legend(ax, fig, driver_raw, size_label):
  944. if size_label == "χ²":
  945. leg_driver = np.sqrt(np.clip(CHI2_LEG, 0, None))
  946. leg_sizes, _ = map_to_sizes(leg_driver, S_MIN, S_MAX)
  947. leg_sizes = leg_sizes * LEG_SCALE
  948. leg_text = [f"{int(v - CHI2_TEXT_SHIFT)}" for v in CHI2_LEG]
  949. title_txt = "χ²"
  950. title_index = 1
  951. else:
  952. leg_driver = np.clip(V_LEG, 0, None)
  953. m = np.isfinite(driver_raw)
  954. if m.sum() > 0:
  955. lo = np.percentile(driver_raw[m], 5)
  956. hi = np.percentile(driver_raw[m], 95) if m.sum() > 1 else (driver_raw[m].max() + 1e-9)
  957. denom = (hi - lo) if (hi - lo) > 1e-12 else 1.0
  958. z = np.clip((leg_driver - lo) / denom, 0, 1)
  959. leg_sizes = (S_MIN + (S_MAX - S_MIN) * z) * LEG_SCALE
  960. else:
  961. leg_sizes = (S_MIN + (S_MAX - S_MIN) * np.linspace(0.2, 0.8, len(V_LEG))) * LEG_SCALE
  962. leg_text = [V_TEXT_FMT.format(v) for v in V_LEG]
  963. title_txt = "V"
  964. title_index = 1
  965. diam = area_to_diameter_pt(leg_sizes)
  966. x_pt = np.zeros(len(leg_sizes))
  967. x_pt[0] = 0.0
  968. for k in range(1, len(leg_sizes)):
  969. x_pt[k] = x_pt[k-1] + (diam[k-1]/2 + diam[k]/2) + LEG_GAP_PT
  970. bbox = ax.get_window_extent().transformed(fig.dpi_scale_trans.inverted())
  971. ax_w_pt = bbox.width * 72.0
  972. x_ax = x_pt / ax_w_pt
  973. xs_leg = LEG_BASE_X + x_ax
  974. y_leg = LEG_BASE_Y
  975. for x, s_area, t in zip(xs_leg, leg_sizes, leg_text):
  976. ax.scatter([x], [y_leg], s=[s_area], facecolors="none",
  977. edgecolors=LEG_EDGE, linewidths=1.2,
  978. transform=ax.transAxes, clip_on=False)
  979. ax.text(x, y_leg, t, ha="center", va="center",
  980. fontsize=9, color="#444444",
  981. transform=ax.transAxes, clip_on=False)
  982. ax.text(xs_leg[title_index], y_leg + 0.10, title_txt,
  983. ha="center", va="center",
  984. fontsize=11, color="#444444",
  985. transform=ax.transAxes, clip_on=False)
  986. build_size_legend(ax, fig, driver_raw, size_label)
  987. plt.tight_layout()
  988. plt.subplots_adjust(bottom=BOTTOM_MARGIN)
  989. out_png = os.path.join(out_dir, f"{base_name}_GMM_pairwise_frequency_bubble.png")
  990. plt.savefig(out_png, dpi=300, bbox_inches="tight")
  991. out_pdf = os.path.join(out_dir, f"{base_name}_GMM_pairwise_frequency_bubble.pdf")
  992. fig.savefig(out_pdf, format="pdf", bbox_inches="tight")
  993. plt.close()
  994. print("[DONE] Bubble-matrix saved:", out_png)
  995. # ======================================================================================
  996. # ===================== 11) CORRELACIONES (lower-triangle) ==============================
  997. # ===================== AJUSTADAS POR SUJETO (GEE) ====================================
  998. # ======================================================================================
  999. # ✅ Lo que pediste (aplica a TODOS los plots de correlación):
  1000. # - quitar la FILA de "∆WMH" pero NO su columna
  1001. # - quitar la COLUMNA de "∆Autocor" pero NO su fila
  1002. DROP_ROW_ONLY = "WMH"
  1003. DROP_COL_ONLY = "Autocor"
  1004. cbar_label_size=16
  1005. # ---- imports (solo para esta sección) ----
  1006. import statsmodels.api as sm
  1007. from statsmodels.genmod.generalized_estimating_equations import GEE
  1008. from statsmodels.genmod.cov_struct import Independence # o Exchangeable
  1009. def rank_z(x):
  1010. """Spearman = ranks; luego z-score para que beta ~ r."""
  1011. r = pd.Series(x).rank(method="average").to_numpy(dtype=float)
  1012. sd = np.nanstd(r, ddof=1)
  1013. if sd < 1e-12:
  1014. return (r - np.nanmean(r)) # todo constante -> sd~0
  1015. return (r - np.nanmean(r)) / sd
  1016. def format_lancet(x, nd=2):
  1017. """Formato tipo Lancet: 1.23 -> 1·23 (punto)."""
  1018. return f"{x:.{nd}f}".replace(".", ".")
  1019. def compute_spearman_block_clustered(df_block, label="ALL"):
  1020. """
  1021. Spearman ajustada por dependencia intra-sujeto usando:
  1022. - rank+z de cada variable (Spearman)
  1023. - GEE con groups = subject_id (SUBJ_COL)
  1024. - partial opcional: covars incluidas en el modelo
  1025. Devuelve df_r, df_p, df_padj (lower triangle rellenado).
  1026. """
  1027. # OJO: usa SUBJ_COL ya definido arriba ("subject_id")
  1028. use_cols = vars_corr + ([SUBJ_COL] if SUBJ_COL not in vars_corr else []) + (CORR_COVARS if DO_PARTIAL else [])
  1029. miss = [c for c in use_cols if c not in df_block.columns]
  1030. if miss:
  1031. raise ValueError(f"[{label}] Faltan columnas: {miss}")
  1032. # NaNs: mismo comportamiento que antes
  1033. if df_block[use_cols].isna().any().any():
  1034. bad_n = int(df_block[use_cols].isna().any(axis=1).sum())
  1035. raise ValueError(f"[{label}] Hay {bad_n} filas con NaN en vars/covars/subject. Arregla o filtra.")
  1036. df_use = df_block[use_cols].copy()
  1037. # ranks+z para vars_corr
  1038. Z = {v: rank_z(df_use[v].values) for v in vars_corr}
  1039. # covars para partial (dentro del modelo)
  1040. cov_df = None
  1041. if DO_PARTIAL:
  1042. cov_df = df_use[CORR_COVARS].copy()
  1043. if "Sex" in CORR_COVARS:
  1044. cov_df["Sex"] = encode_sex(cov_df["Sex"])
  1045. # z-score covariables numéricas (recomendado)
  1046. for c in cov_df.columns:
  1047. if pd.api.types.is_numeric_dtype(cov_df[c]):
  1048. sd = cov_df[c].std(ddof=1)
  1049. if sd is None or sd < 1e-12:
  1050. cov_df[c] = cov_df[c] - cov_df[c].mean()
  1051. else:
  1052. cov_df[c] = (cov_df[c] - cov_df[c].mean()) / sd
  1053. groups = df_use[SUBJ_COL].astype("category")
  1054. P = len(vars_corr)
  1055. Rmat = np.full((P, P), np.nan, float)
  1056. PVAL = np.full((P, P), np.nan, float)
  1057. for i in range(P):
  1058. Rmat[i, i] = 1.0
  1059. PVAL[i, i] = 0.0
  1060. y = Z[vars_corr[i]]
  1061. for j in range(i): # lower triangle
  1062. x = Z[vars_corr[j]]
  1063. Xcols = {"x": x}
  1064. if DO_PARTIAL:
  1065. for c in cov_df.columns:
  1066. Xcols[c] = cov_df[c].to_numpy(dtype=float)
  1067. X = pd.DataFrame(Xcols)
  1068. X = sm.add_constant(X, has_constant="add")
  1069. # GEE cluster-robust por sujeto
  1070. model = GEE(
  1071. endog=y,
  1072. exog=X,
  1073. groups=groups,
  1074. cov_struct=Independence()
  1075. )
  1076. res = model.fit()
  1077. r_adj = float(res.params["x"]) # ~ "Spearman r" ajustada
  1078. p_adj = float(res.pvalues["x"]) # p con dependencia corregida
  1079. Rmat[i, j] = r_adj
  1080. PVAL[i, j] = p_adj
  1081. # FDR en triángulo inferior
  1082. PADJ = PVAL.copy()
  1083. if CORR_USE_FDR:
  1084. tri = np.tril_indices(P, k=-1)
  1085. pvec = PVAL[tri]
  1086. padj = multipletests(pvec, method=CORR_FDR_METHOD)[1]
  1087. PADJ[tri] = padj
  1088. df_r = pd.DataFrame(Rmat, index=vars_corr, columns=vars_corr)
  1089. df_p = pd.DataFrame(PVAL, index=vars_corr, columns=vars_corr)
  1090. df_padj = pd.DataFrame(PADJ, index=vars_corr, columns=vars_corr)
  1091. return df_r, df_p, df_padj
  1092. def plot_lower_triangle(df_r, df_p_use, out_png, out_pdf=None,
  1093. drop_row_only=DROP_ROW_ONLY, drop_col_only=DROP_COL_ONLY,
  1094. thr_text_white=0.35, tick_label_size=16, cell_text_size=14,
  1095. cbar_label_size=14, cbar_tick_size=12):
  1096. """
  1097. Heatmap lower-triangle con drops asimétricos:
  1098. - quita SOLO la fila drop_row_only (pero deja su columna)
  1099. - quita SOLO la columna drop_col_only (pero deja su fila)
  1100. """
  1101. rows = [v for v in df_r.index.tolist() if v != drop_row_only]
  1102. cols = [v for v in df_r.columns.tolist() if v != drop_col_only]
  1103. df_rp = df_r.loc[rows, cols]
  1104. df_pp = df_p_use.loc[rows, cols]
  1105. mat = df_rp.to_numpy(dtype=float)
  1106. Pr, Pc = mat.shape
  1107. fig, ax = plt.subplots(figsize=(0.6 * Pc + 6, 0.6 * Pr + 5))
  1108. rr = np.arange(Pr)[:, None]
  1109. cc = np.arange(Pc)[None, :]
  1110. mask = cc > rr
  1111. mat_plot = mat.copy()
  1112. mat_plot[mask] = np.nan
  1113. im = ax.imshow(mat_plot, vmin=CORR_VMIN, vmax=CORR_VMAX, cmap=CORR_CMAP, aspect="equal")
  1114. ax.set_xticks(np.arange(Pc))
  1115. ax.set_yticks(np.arange(Pr))
  1116. ax.set_xticklabels(df_rp.columns.tolist(), rotation=45, ha="right", fontsize=tick_label_size)
  1117. ax.set_yticklabels(df_rp.index.tolist(), fontsize=tick_label_size)
  1118. ax.set_xticks(np.arange(-.5, Pc, 1), minor=True)
  1119. ax.set_yticks(np.arange(-.5, Pr, 1), minor=True)
  1120. ax.grid(which="minor", color="#dddddd", linestyle="-", linewidth=1)
  1121. ax.tick_params(which="minor", bottom=False, left=False)
  1122. # anotaciones (solo triángulo inferior) -> formato "Lancet" con punto medio (·)
  1123. for i in range(Pr):
  1124. for j in range(Pc):
  1125. if j <= i:
  1126. r = df_rp.iat[i, j]
  1127. p = df_pp.iat[i, j]
  1128. if np.isfinite(r):
  1129. txt_color = "white" if abs(r) > thr_text_white else "black"
  1130. r_txt = format_lancet(r, nd=2)
  1131. ax.text(j, i, f"{r_txt}{stars_from_p(p)}",
  1132. ha="center", va="center",
  1133. fontsize=cell_text_size, color=txt_color)
  1134. cbar = fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
  1135. cbar.set_label("Spearman r", fontsize=cbar_label_size)
  1136. cbar.ax.tick_params(labelsize=cbar_tick_size)
  1137. plt.tight_layout()
  1138. fig.savefig(out_png, dpi=300, bbox_inches="tight")
  1139. if out_pdf is not None:
  1140. fig.savefig(out_pdf, format="pdf", bbox_inches="tight")
  1141. plt.show()
  1142. plt.close(fig)
  1143. print("[DONE] Corr heatmap PNG:", out_png)
  1144. if out_pdf is not None:
  1145. print("[DONE] Corr heatmap PDF:", out_pdf)
  1146. if DO_CORR_BLOCKS:
  1147. # ---------------- ALL ----------------
  1148. df_r_all, df_p_all, df_padj_all = compute_spearman_block_clustered(df_out, label="ALL")
  1149. p_used_all = df_padj_all if CORR_USE_FDR else df_p_all
  1150. png_all = os.path.join(out_dir, f"{base_name}_lowertri_spearman_ALL.png")
  1151. pdf_all = os.path.join(out_dir, f"{base_name}_lowertri_spearman_ALL.pdf")
  1152. plot_lower_triangle(df_r_all, p_used_all, png_all, out_pdf=pdf_all)
  1153. # ---------------- BY CLUSTER ----------------
  1154. if CLUSTER_COL not in df_out.columns:
  1155. raise ValueError(f"No existe {CLUSTER_COL} en df_out.")
  1156. clusters = sorted(pd.unique(df_out[CLUSTER_COL].dropna()))
  1157. xlsx_out = os.path.join(out_dir, f"{base_name}_lowertri_spearman_ALL_and_clusters.xlsx")
  1158. sheets = {
  1159. "ALL_r": df_r_all,
  1160. "ALL_p": df_p_all,
  1161. ("ALL_pFDR" if CORR_USE_FDR else "ALL_p_used"): p_used_all
  1162. }
  1163. for cl in clusters:
  1164. block = df_out[df_out[CLUSTER_COL] == cl].copy()
  1165. if len(block) < 5:
  1166. print(f"[WARN] Cluster {cl}: N={len(block)} muy pequeño, salto.")
  1167. continue
  1168. df_r, df_p, df_padj = compute_spearman_block_clustered(block, label=f"cluster_{cl}")
  1169. p_used = df_padj if CORR_USE_FDR else df_p
  1170. png_cl = os.path.join(out_dir, f"{base_name}_lowertri_spearman_{CLUSTER_COL}_{cl}.png")
  1171. pdf_cl = os.path.join(out_dir, f"{base_name}_lowertri_spearman_{CLUSTER_COL}_{cl}.pdf")
  1172. plot_lower_triangle(df_r, p_used, png_cl, out_pdf=pdf_cl)
  1173. sheets[f"cl{cl}_r"[:31]] = df_r
  1174. sheets[f"cl{cl}_p"[:31]] = df_p
  1175. sheets[(f"cl{cl}_pFDR" if CORR_USE_FDR else f"cl{cl}_pused")[:31]] = p_used
  1176. with pd.ExcelWriter(xlsx_out, engine="openpyxl") as writer:
  1177. for name, df_sheet in sheets.items():
  1178. df_sheet.to_excel(writer, sheet_name=str(name)[:31])
  1179. print("[DONE] Corr Excel:", xlsx_out)
  1180. # ======================================================================================
  1181. # ===================== 12) ML CLASIFICACION + ROC (OOF) ===============================
  1182. # ======================================================================================
  1183. # Cambios aplicados:
  1184. # - En ROC: clases mostradas como L₁, L₂, ... (NO 0,1,2)
  1185. # - ROC SIN TÍTULO (plt.title(""))
  1186. # - (Opcional) macro AUC como texto dentro del gráfico (no título)
  1187. def auc_ovr_macro(y_true, proba, labels):
  1188. return roc_auc_score(y_true, proba, multi_class="ovr", average="macro", labels=labels)
  1189. def align_proba_to_global(fold_classes, proba_fold, global_classes):
  1190. out = np.full((proba_fold.shape[0], len(global_classes)), np.nan, float)
  1191. idx_map = {c: i for i, c in enumerate(global_classes)}
  1192. for j, c in enumerate(fold_classes):
  1193. out[:, idx_map[c]] = proba_fold[:, j]
  1194. return out
  1195. def strip_prefixes(feat_names):
  1196. cleaned = []
  1197. for f in feat_names:
  1198. s = str(f)
  1199. s = re.sub(r"^(num|cat)__", "", s)
  1200. s = re.sub(r"^(num|cat)__[^_]+__", "", s)
  1201. s = s.replace("onehot__", "").replace("imputer__", "")
  1202. cleaned.append(s)
  1203. return np.array(cleaned, dtype=object)
  1204. def get_feature_names_from_preprocess(prep: ColumnTransformer):
  1205. try:
  1206. names = prep.get_feature_names_out()
  1207. return strip_prefixes(names)
  1208. except Exception:
  1209. return None
  1210. def build_splits(X, y, groups, n_splits, seed):
  1211. has_repeats = groups.duplicated().any()
  1212. if has_repeats:
  1213. try:
  1214. from sklearn.model_selection import StratifiedGroupKFold
  1215. splitter = StratifiedGroupKFold(n_splits=n_splits, shuffle=True, random_state=seed)
  1216. splits = list(splitter.split(X, y, groups=groups))
  1217. return splits, "StratifiedGroupKFold", True
  1218. except Exception:
  1219. splitter = GroupKFold(n_splits=n_splits)
  1220. splits = list(splitter.split(X, y, groups=groups))
  1221. print("[WARN] No hay StratifiedGroupKFold. Uso GroupKFold (NO estratifica).")
  1222. return splits, "GroupKFold", True
  1223. else:
  1224. splitter = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=seed)
  1225. splits = list(splitter.split(X, y))
  1226. return splits, "StratifiedKFold", False
  1227. def plot_multiclass_roc_with_se(y_true, proba, classes, out_png, cluster_color,
  1228. auc_macro_value=None, n_boot=1000, grid_n=200,
  1229. seed=42, show_auc_text=True):
  1230. rng = np.random.default_rng(seed)
  1231. y_true = np.asarray(y_true)
  1232. proba = np.asarray(proba, float)
  1233. fpr_grid = np.linspace(0, 1, grid_n)
  1234. plt.figure(figsize=(8, 6))
  1235. for i, cl in enumerate(classes):
  1236. y_i = (y_true == cl).astype(int)
  1237. s_i = proba[:, i]
  1238. if len(np.unique(y_i)) < 2:
  1239. continue
  1240. fpr_obs, tpr_obs, _ = roc_curve(y_i, s_i)
  1241. auc_obs = auc(fpr_obs, tpr_obs)
  1242. idx_all = np.arange(len(y_i))
  1243. auc_boot = []
  1244. tpr_boot = []
  1245. for _ in range(n_boot):
  1246. idx = rng.choice(idx_all, size=len(idx_all), replace=True)
  1247. y_b = y_i[idx]
  1248. s_b = s_i[idx]
  1249. if len(np.unique(y_b)) < 2:
  1250. continue
  1251. fpr_b, tpr_b, _ = roc_curve(y_b, s_b)
  1252. auc_boot.append(auc(fpr_b, tpr_b))
  1253. tpr_interp = np.interp(fpr_grid, fpr_b, tpr_b)
  1254. tpr_interp[0] = 0.0
  1255. tpr_boot.append(tpr_interp)
  1256. color = cluster_color.get(int(cl), "black")
  1257. name = lab_text(cl)
  1258. if len(auc_boot) < 10:
  1259. plt.plot(fpr_obs, tpr_obs, color=color, lw=2, label=f"{name} AUC={auc_obs:.3f}")
  1260. continue
  1261. auc_boot = np.asarray(auc_boot, float)
  1262. tpr_boot = np.asarray(tpr_boot, float)
  1263. auc_se = auc_boot.std(ddof=1)
  1264. tpr_mean = tpr_boot.mean(axis=0)
  1265. tpr_se = tpr_boot.std(axis=0, ddof=1) / np.sqrt(tpr_boot.shape[0])
  1266. plt.plot(fpr_grid, tpr_mean, color=color, lw=2,
  1267. label=f"{name} AUC={auc_obs:.3f} ± {auc_se:.3f}")
  1268. plt.fill_between(
  1269. fpr_grid,
  1270. np.clip(tpr_mean - tpr_se, 0, 1),
  1271. np.clip(tpr_mean + tpr_se, 0, 1),
  1272. color=color,
  1273. alpha=0.18,
  1274. linewidth=0
  1275. )
  1276. plt.plot([0, 1], [0, 1], linestyle="--", lw=1)
  1277. plt.xlim(0.0, 1.0)
  1278. plt.ylim(0.0, 1.0)
  1279. plt.xlabel("False Positive Rate (FPR)")
  1280. plt.ylabel("True Positive Rate (TPR)")
  1281. plt.title("")
  1282. if (auc_macro_value is not None) and show_auc_text:
  1283. plt.text(
  1284. 0.98, 0.02,
  1285. f"Macro AUC (OVR) = {auc_macro_value:.3f}",
  1286. ha="right", va="bottom",
  1287. transform=plt.gca().transAxes, fontsize=10
  1288. )
  1289. plt.legend(frameon=False, fontsize=9)
  1290. plt.tight_layout()
  1291. plt.savefig(out_png, dpi=300, bbox_inches="tight")
  1292. plt.show()
  1293. plt.close()
  1294. print("[DONE] ROC guardada:", out_png)
  1295. def compute_shap_global_and_per_class(shap_values, K_expected=None):
  1296. if isinstance(shap_values, list):
  1297. per_class = [np.mean(np.abs(sv), axis=0) for sv in shap_values]
  1298. shap_global = np.mean(np.vstack(per_class), axis=0)
  1299. return shap_global, per_class
  1300. sv = np.asarray(shap_values)
  1301. if sv.ndim == 3:
  1302. shap_global = np.mean(np.abs(sv), axis=(0, 2))
  1303. per_class = [np.mean(np.abs(sv[:, :, k]), axis=0) for k in range(sv.shape[2])]
  1304. if K_expected is not None and sv.shape[2] != K_expected:
  1305. print(f"[WARN] SHAP devolvió K={sv.shape[2]} clases, esperaba K={K_expected}.")
  1306. return shap_global, per_class
  1307. if sv.ndim == 2:
  1308. shap_global = np.mean(np.abs(sv), axis=0)
  1309. return shap_global, None
  1310. raise ValueError(f"Formato shap_values inesperado: shape={sv.shape}")
  1311. def save_shap_summary_plot(shap_vals_2d, X_2d, feat_names, out_png,
  1312. max_display=30, plot_type=None):
  1313. try:
  1314. shap.summary_plot(
  1315. shap_vals_2d, X_2d,
  1316. feature_names=feat_names,
  1317. max_display=max_display,
  1318. plot_type=plot_type,
  1319. show=False
  1320. )
  1321. fig = plt.gcf()
  1322. fig.set_size_inches(10, 7)
  1323. fig.tight_layout()
  1324. fig.savefig(out_png, dpi=300, bbox_inches="tight")
  1325. plt.close(fig)
  1326. print("[DONE] SHAP plot guardado:", out_png)
  1327. except Exception as e:
  1328. print("[WARN] No pude guardar SHAP plot:", repr(e))
  1329. # ======================================================================================
  1330. # ===================================== RUN ============================================
  1331. # ======================================================================================
  1332. if DO_ML:
  1333. needed_ml = [TARGET, GPO_COL, SUBJ_COL] + FEATURES
  1334. miss_ml = [c for c in needed_ml if c not in df_out.columns]
  1335. if miss_ml:
  1336. raise ValueError(f"Faltan columnas para ML: {miss_ml}")
  1337. df_ml = df_out.copy()
  1338. df_ml = df_ml[df_ml[TARGET].notna()].reset_index(drop=True)
  1339. X = df_ml[FEATURES].copy()
  1340. y = df_ml[TARGET].copy()
  1341. groups = df_ml[SUBJ_COL].copy()
  1342. gpo_vec = df_ml[GPO_COL].copy()
  1343. classes_global = np.array(sorted(pd.unique(y)))
  1344. K = len(classes_global)
  1345. if K < 2:
  1346. raise ValueError("El target tiene <2 clases. No se puede clasificar.")
  1347. print(f"[INFO][ML] N={len(df_ml)} | K={K} clases: {classes_global}")
  1348. cat_cols = [c for c in ["Sex"] if c in FEATURES]
  1349. num_cols = [c for c in FEATURES if c not in cat_cols]
  1350. try:
  1351. ohe = OneHotEncoder(handle_unknown="ignore", sparse_output=False)
  1352. except TypeError:
  1353. ohe = OneHotEncoder(handle_unknown="ignore", sparse=False)
  1354. preprocess = ColumnTransformer(
  1355. transformers=[
  1356. ("num", Pipeline(steps=[("imputer", SimpleImputer(strategy="median"))]), num_cols),
  1357. ("cat", Pipeline(steps=[
  1358. ("imputer", SimpleImputer(strategy="most_frequent")),
  1359. ("onehot", ohe)
  1360. ]), cat_cols),
  1361. ],
  1362. remainder="drop"
  1363. )
  1364. clf = LGBMClassifier(
  1365. objective="multiclass",
  1366. n_estimators=600,
  1367. learning_rate=0.03,
  1368. num_leaves=31,
  1369. subsample=0.8,
  1370. colsample_bytree=0.8,
  1371. class_weight="balanced",
  1372. random_state=RANDOM_STATE,
  1373. n_jobs=-1
  1374. )
  1375. pipe = Pipeline(steps=[("prep", preprocess), ("clf", clf)])
  1376. proba_oof_sum = np.zeros((len(df_ml), K), dtype=float)
  1377. proba_oof_count = np.zeros((len(df_ml), K), dtype=int)
  1378. pred_oof_last = np.full((len(df_ml),), None, object)
  1379. auc_per_repeat = []
  1380. cv_name_last = None
  1381. for rep in range(N_REPEATS):
  1382. seed_rep = BASE_SEED + rep
  1383. splits_list, cv_name, has_repeats = build_splits(X, y, groups, N_SPLITS, seed_rep)
  1384. cv_name_last = cv_name
  1385. print(f"[INFO][ML] Repeat {rep+1}/{N_REPEATS} | CV={cv_name} | seed={seed_rep} | repeats={has_repeats}")
  1386. proba_oof_rep = np.full((len(df_ml), K), np.nan, float)
  1387. pred_oof_rep = np.full((len(df_ml),), None, object)
  1388. for fold, (tr, te) in enumerate(splits_list, start=1):
  1389. Xtr, Xte = X.iloc[tr], X.iloc[te]
  1390. ytr, yte = y.iloc[tr], y.iloc[te]
  1391. pipe.set_params(clf__random_state=seed_rep)
  1392. pipe.fit(Xtr, ytr)
  1393. proba = pipe.predict_proba(Xte)
  1394. fold_classes = pipe.named_steps["clf"].classes_
  1395. proba_aligned = align_proba_to_global(fold_classes, proba, classes_global)
  1396. proba_oof_rep[te, :] = proba_aligned
  1397. pred_oof_rep[te] = pipe.predict(Xte)
  1398. if not np.isfinite(proba_oof_rep).all():
  1399. raise RuntimeError("OOF proba (rep) tiene NaNs. Suele pasar si algún fold no vio alguna clase.")
  1400. auc_rep = auc_ovr_macro(y, proba_oof_rep, labels=classes_global)
  1401. auc_per_repeat.append(float(auc_rep))
  1402. print(f"[DONE][ML] Repeat {rep+1}: OOF AUC macro OVR = {auc_rep:.4f}")
  1403. mask_finite = np.isfinite(proba_oof_rep)
  1404. proba_oof_sum[mask_finite] += proba_oof_rep[mask_finite]
  1405. proba_oof_count[mask_finite] += 1
  1406. pred_oof_last = pred_oof_rep
  1407. with np.errstate(divide="ignore", invalid="ignore"):
  1408. proba_oof_mean = proba_oof_sum / np.maximum(proba_oof_count, 1)
  1409. if not np.isfinite(proba_oof_mean).all():
  1410. raise RuntimeError("proba_oof_mean tiene NaNs/inf. Revisa folds con clases faltantes.")
  1411. auc_oof_mean = auc_ovr_macro(y, proba_oof_mean, labels=classes_global)
  1412. print(f"[DONE][ML] OOF AUC global (macro OVR) promedio {N_REPEATS} rep: {auc_oof_mean:.4f}")
  1413. # ---------------- ROC GLOBAL (L₁,L₂,...; sin título) ----------------
  1414. roc_png_all = os.path.join(out_dir, f"{base_name}_ROC_OOF_ALLsubjects_rep{N_REPEATS}.png")
  1415. plot_multiclass_roc_with_se(
  1416. y_true=y,
  1417. proba=proba_oof_mean,
  1418. classes=classes_global,
  1419. out_png=roc_png_all,
  1420. cluster_color=cluster_color,
  1421. auc_macro_value=auc_oof_mean,
  1422. n_boot=N_BOOT,
  1423. grid_n=GRID_N,
  1424. seed=SEED,
  1425. show_auc_text=True
  1426. )
  1427. # ---------------- ROC POR Gpo (L₁,L₂,...; sin título) ----------------
  1428. if DO_ROC_BY_GPO:
  1429. for g in pd.unique(gpo_vec.dropna()):
  1430. m = (np.asarray(gpo_vec) == g)
  1431. n_g = int(m.sum())
  1432. if n_g < 30 or len(np.unique(np.asarray(y)[m])) < 2:
  1433. print(f"[SKIP][ML] Gpo={g}: N={n_g} o 1 clase.")
  1434. continue
  1435. out_png = os.path.join(out_dir, f"{base_name}_ROC_OOFmean_Gpo_{g}_rep{N_REPEATS}.png")
  1436. plot_multiclass_roc_with_se(
  1437. y_true=np.asarray(y)[m],
  1438. proba=proba_oof_mean[m, :],
  1439. classes=classes_global,
  1440. out_png=out_png,
  1441. cluster_color=cluster_color,
  1442. auc_macro_value=None,
  1443. n_boot=N_BOOT,
  1444. grid_n=GRID_N,
  1445. seed=SEED,
  1446. show_auc_text=False
  1447. )
  1448. # ---------------- Feature importance + SHAP (entrena final en todo) ----------------
  1449. if DO_FEATURE_IMPORTANCE_AND_SHAP:
  1450. pipe.set_params(clf__random_state=RANDOM_STATE)
  1451. pipe.fit(X, y)
  1452. prep_fitted = pipe.named_steps["prep"]
  1453. clf_fitted = pipe.named_steps["clf"]
  1454. feat_names = get_feature_names_from_preprocess(prep_fitted)
  1455. if feat_names is None:
  1456. feat_names = np.array([f"f{i}" for i in range(clf_fitted.booster_.num_feature())], dtype=object)
  1457. booster = clf_fitted.booster_
  1458. imp_gain = booster.feature_importance(importance_type="gain")
  1459. fi = (
  1460. pd.DataFrame({"feature": feat_names, "importance_gain": imp_gain})
  1461. .sort_values("importance_gain", ascending=False)
  1462. .reset_index(drop=True)
  1463. )
  1464. featimp_xlsx = os.path.join(out_dir, f"{base_name}_FeatureImportance_LGBM_{TARGET}.xlsx")
  1465. featimp_png = os.path.join(out_dir, f"{base_name}_FeatureImportance_LGBM_{TARGET}.png")
  1466. with pd.ExcelWriter(featimp_xlsx, engine="openpyxl") as writer:
  1467. fi.to_excel(writer, index=False, sheet_name="feature_importance_gain")
  1468. topn = min(30, len(fi))
  1469. plt.figure(figsize=(8, 10))
  1470. plt.barh(fi.loc[:topn-1, "feature"][::-1], fi.loc[:topn-1, "importance_gain"][::-1])
  1471. plt.xlabel("Importance (gain)")
  1472. plt.title("")
  1473. plt.tight_layout()
  1474. plt.savefig(featimp_png, dpi=300, bbox_inches="tight")
  1475. plt.close()
  1476. print("[DONE][ML] Feature importance:", featimp_xlsx)
  1477. print("[DONE][ML] Feature importance plot:", featimp_png)
  1478. shap_dir = os.path.join(out_dir, f"{base_name}_SHAP_{TARGET}")
  1479. os.makedirs(shap_dir, exist_ok=True)
  1480. if _HAS_SHAP:
  1481. X_tr = prep_fitted.transform(X)
  1482. rng = np.random.default_rng(SEED)
  1483. if X_tr.shape[0] > MAX_SHAP_SAMPLES:
  1484. idx = rng.choice(np.arange(X_tr.shape[0]), size=MAX_SHAP_SAMPLES, replace=False)
  1485. X_sh = X_tr[idx]
  1486. else:
  1487. X_sh = X_tr
  1488. explainer = shap.TreeExplainer(clf_fitted)
  1489. shap_values = explainer.shap_values(X_sh)
  1490. shap_global, per_class = compute_shap_global_and_per_class(shap_values, K_expected=K)
  1491. shap_imp = (
  1492. pd.DataFrame({"feature": feat_names, "mean_abs_shap": shap_global})
  1493. .sort_values("mean_abs_shap", ascending=False)
  1494. .reset_index(drop=True)
  1495. )
  1496. shap_xlsx = os.path.join(shap_dir, f"SHAP_importance_{base_name}_{TARGET}.xlsx")
  1497. with pd.ExcelWriter(shap_xlsx, engine="openpyxl") as writer:
  1498. shap_imp.to_excel(writer, index=False, sheet_name="shap_global")
  1499. fi.to_excel(writer, index=False, sheet_name="lgbm_gain_importance")
  1500. if SAVE_SHAP_PER_CLASS and per_class is not None:
  1501. for k_idx, cl in enumerate(classes_global):
  1502. dfk = (
  1503. pd.DataFrame({"feature": feat_names, "mean_abs_shap": per_class[k_idx]})
  1504. .sort_values("mean_abs_shap", ascending=False)
  1505. )
  1506. dfk.to_excel(writer, index=False, sheet_name=f"class_{lab_text(cl)}"[:31])
  1507. print("[DONE][ML] SHAP Excel:", shap_xlsx)
  1508. # Para plots: usa clase 0 por defecto
  1509. if isinstance(shap_values, list):
  1510. shap2d = shap_values[0]
  1511. else:
  1512. sv = np.asarray(shap_values)
  1513. shap2d = sv[:, :, 0] if sv.ndim == 3 else sv
  1514. shap_bee_png = os.path.join(shap_dir, f"SHAP_beeswarm_{base_name}_{TARGET}.png")
  1515. save_shap_summary_plot(shap2d, X_sh, feat_names, shap_bee_png, max_display=30, plot_type=None)
  1516. shap_bar_png = os.path.join(shap_dir, f"SHAP_summaryBar_{base_name}_{TARGET}.png")
  1517. save_shap_summary_plot(shap2d, X_sh, feat_names, shap_bar_png, max_display=30, plot_type="bar")
  1518. if SAVE_SHAP_PER_CLASS and per_class is not None and isinstance(shap_values, list):
  1519. for k_idx, cl in enumerate(classes_global):
  1520. outp = os.path.join(
  1521. shap_dir,
  1522. f"SHAP_beeswarm_class_{lab_text(cl)}_{base_name}_{TARGET}.png"
  1523. )
  1524. save_shap_summary_plot(shap_values[k_idx], X_sh, feat_names, outp, max_display=30, plot_type=None)
  1525. else:
  1526. print("[INFO][ML] SHAP omitido (no está instalado).")
  1527. # ---------------- Guardar OOF predictions ----------------
  1528. oof_xlsx = os.path.join(out_dir, f"{base_name}_OOF_predictions_{TARGET}_rep{N_REPEATS}.xlsx")
  1529. oof_df = df_ml[[TARGET, GPO_COL, SUBJ_COL]].copy()
  1530. oof_df["pred_oof_last_repeat"] = pred_oof_last
  1531. for i, c in enumerate(classes_global):
  1532. oof_df[f"proba_oof_mean_{lab_text(c)}"] = proba_oof_mean[:, i] # ✅ columnas con L₁,L₂,...
  1533. summary_df = pd.DataFrame([{
  1534. "target": TARGET,
  1535. "cv": cv_name_last,
  1536. "n_splits": N_SPLITS,
  1537. "n_repeats": N_REPEATS,
  1538. "AUC_OOF_macro_OVR_meanProba": float(auc_oof_mean),
  1539. "AUC_per_repeat_mean": float(np.mean(auc_per_repeat)),
  1540. "AUC_per_repeat_sd": float(np.std(auc_per_repeat, ddof=1)) if len(auc_per_repeat) > 1 else 0.0,
  1541. "N": int(len(df_ml)),
  1542. "classes": ", ".join(lab_text(c) for c in classes_global) # ✅ L₁,L₂,...
  1543. }])
  1544. auc_rep_df = pd.DataFrame({
  1545. "repeat": np.arange(1, N_REPEATS + 1),
  1546. "auc_oof_macro_ovr": auc_per_repeat
  1547. })
  1548. with pd.ExcelWriter(oof_xlsx, engine="openpyxl") as writer:
  1549. summary_df.to_excel(writer, index=False, sheet_name="summary")
  1550. auc_rep_df.to_excel(writer, index=False, sheet_name="auc_by_repeat")
  1551. oof_df.to_excel(writer, index=False, sheet_name="oof_mean")
  1552. print("[DONE][ML] OOF Excel:", oof_xlsx)
  1553. # ======================================================================================
  1554. # ===================== HEATMAP por INDIVIDUO: subjects (Y) x lesion types (X) ==========
  1555. # ===================== PDF 100% VECTOR (pcolormesh) ====================================
  1556. # ======================================================================================
  1557. import os
  1558. import numpy as np
  1559. import pandas as pd
  1560. import matplotlib as mpl
  1561. import matplotlib.pyplot as plt
  1562. def plot_subject_lesiontype_heatmap_counts_vector(df_in, subj_col, cluster_col,
  1563. out_png,
  1564. out_pdf=None,
  1565. max_subjects=200,
  1566. sort_subjects=True,
  1567. fig_w=3.5, fig_h=6.0,
  1568. cmap_name="OrRd",
  1569. cluster_order=(0, 1, 2)):
  1570. """
  1571. Heatmap de CONTEOS:
  1572. filas = sujetos
  1573. columnas = clusters (en el orden FINAL ya swappeado arriba)
  1574. PDF totalmente vector: usa pcolormesh (no imshow) y guarda directo con savefig.
  1575. """
  1576. # Texto editable en Illustrator/Inkscape
  1577. mpl.rcParams["pdf.fonttype"] = 42
  1578. mpl.rcParams["ps.fonttype"] = 42
  1579. if subj_col not in df_in.columns:
  1580. raise ValueError(f"Missing subject column: {subj_col}")
  1581. if cluster_col not in df_in.columns:
  1582. raise ValueError(f"Missing cluster column: {cluster_col}")
  1583. cl_use = df_in[cluster_col].astype(int)
  1584. tab = pd.crosstab(df_in[subj_col], cl_use).sort_index()
  1585. tab = tab.reindex(columns=list(cluster_order), fill_value=0)
  1586. if sort_subjects:
  1587. tab = tab.loc[tab.sum(axis=1).sort_values(ascending=False).index]
  1588. if (max_subjects is not None) and (len(tab) > max_subjects):
  1589. tab_plot = tab.iloc[:max_subjects].copy()
  1590. else:
  1591. tab_plot = tab.copy()
  1592. x_labels = [lab_text(int(i)) for i in tab_plot.columns]
  1593. Z = tab_plot.to_numpy(dtype=float)
  1594. fig, ax = plt.subplots(figsize=(fig_w, fig_h))
  1595. # --- VECTOR HEATMAP (each cell is a vector rectangle) ---
  1596. x = np.arange(Z.shape[1] + 1)
  1597. y = np.arange(Z.shape[0] + 1)
  1598. mesh = ax.pcolormesh(
  1599. x, y, Z,
  1600. cmap=cmap_name,
  1601. shading="flat",
  1602. edgecolors="none" # pon "k" + linewidth si quieres bordes
  1603. )
  1604. # Ticks centrados
  1605. ax.set_xticks(np.arange(len(x_labels)) + 0.5)
  1606. ax.set_xticklabels(x_labels, rotation=0, fontsize=11)
  1607. # Muchos sujetos => no mostrar ticks
  1608. ax.set_yticks([])
  1609. ax.set_xlabel("Cluster", fontsize=13)
  1610. ax.set_ylabel("Subjects", fontsize=13)
  1611. ax.set_title("")
  1612. # Primera fila arriba (estilo imshow)
  1613. ax.invert_yaxis()
  1614. cbar = fig.colorbar(mesh, ax=ax, fraction=0.08, pad=0.04)
  1615. cbar.set_label("Number of lesions", fontsize=11)
  1616. cbar.ax.tick_params(labelsize=10)
  1617. fig.tight_layout()
  1618. # PNG (raster) para vista rápida
  1619. fig.savefig(out_png, dpi=300, bbox_inches="tight")
  1620. print("[DONE] Heatmap PNG:", out_png)
  1621. # PDF vector (sin PIL)
  1622. if out_pdf is not None:
  1623. fig.savefig(out_pdf, bbox_inches="tight")
  1624. print("[DONE] Heatmap PDF (VECTOR):", out_pdf)
  1625. plt.close(fig)
  1626. return tab, tab_plot
  1627. # --- llamada (PNG + PDF) ---
  1628. heat_png = os.path.join(out_dir, f"{base_name}_HEATMAP_subjects_x_clusters_COUNTS.png")
  1629. heat_pdf = os.path.join(out_dir, f"{base_name}_HEATMAP_subjects_x_clusters_COUNTS.pdf")
  1630. tab_counts_subj, tab_counts_top = plot_subject_lesiontype_heatmap_counts_vector(
  1631. df_out,
  1632. subj_col=SUBJ_COL,
  1633. cluster_col=CLUSTER_COL,
  1634. out_png=heat_png,
  1635. out_pdf=heat_pdf,
  1636. max_subjects=400,
  1637. sort_subjects=True,
  1638. fig_w=3.5, fig_h=6.0,
  1639. cmap_name="OrRd",
  1640. cluster_order=(0, 1, 2)
  1641. )
  1642. # (opcional) exportar la tabla
  1643. heat_xlsx = os.path.join(out_dir, f"{base_name}_HEATMAP_subjects_x_clusters_COUNTS.xlsx")
  1644. with pd.ExcelWriter(heat_xlsx, engine="openpyxl") as writer:
  1645. tab_counts_subj.to_excel(writer, sheet_name="counts_all_subjects")
  1646. tab_counts_top.to_excel(writer, sheet_name=f"counts_top{len(tab_counts_top)}")
  1647. print("[DONE] Heatmap tables Excel:", heat_xlsx)
  1648. # ======================================================================================
  1649. # =================== REPRODUCIBILITY (SUBSAMPLING 80/20 BY SUBJECT) ===================
  1650. # ======================================================================================
  1651. # Objetivo:
  1652. # - Tomar un SUBSAMPLE de sujetos (80% IN) sin reemplazo en cada iteración
  1653. # - Ajustar el GMM (k fijo) usando SOLO sujetos IN
  1654. # - Predecir labels para TODO el dataset
  1655. # - Alinear labels al modelo de referencia (full data) con Hungarian
  1656. # - Medir estabilidad: ARI_all, ARI_in, ARI_oob
  1657. # - Estabilidad por lesión: proporción de veces que cae en su cluster modal (solo predicciones OOB)
  1658. #
  1659. # Mantiene: covariance_type="full" y subsampling 80%
  1660. # Mejora estabilidad: n_init alto + reg_covar + max_iter mayor
  1661. import os
  1662. import numpy as np
  1663. import pandas as pd
  1664. from sklearn.mixture import GaussianMixture
  1665. from sklearn.metrics import adjusted_rand_score
  1666. from scipy.optimize import linear_sum_assignment
  1667. # ---------------- CONFIG ----------------
  1668. N_SUBSAMPLE = 200 # iteraciones (200–1000)
  1669. SUBSEED = 42
  1670. IN_FRAC = 0.80 # 80% IN, 20% OOB
  1671. K_FIXED = best_k_gmm # k elegido (por BIC u otro)
  1672. COV_TYPE = "full" # pedido por ti
  1673. N_INIT = 30 # 🔥 mejora: muchos reinicios
  1674. MAX_ITER = 500 # 🔥 mejora: más iteraciones
  1675. REG_COVAR = 1e-4 # 🔥 mejora: regularización (prueba 1e-4 o 1e-3)
  1676. # Columnas (deben existir en tu script)
  1677. # SUBJ_COL = "subject_id"
  1678. # CLUSTER_COL = "cluster_gmm_auto"
  1679. # df_out = ...
  1680. # X_z = ...
  1681. # gmm_best = ...
  1682. # out_dir = ...
  1683. # base_name = ...
  1684. # ---------------- Reference labels ----------------
  1685. # Modelo entrenado en TODO el dataset (referencia)
  1686. labels_ref = gmm_best.predict(X_z).astype(int)
  1687. K_REF = int(labels_ref.max()) + 1
  1688. # ---------------- Helper: Hungarian alignment ----------------
  1689. def align_labels_hungarian(y_pred, y_ref):
  1690. """
  1691. Alinea etiquetas y_pred a y_ref maximizando acuerdos (Hungarian sobre matriz de conteos).
  1692. Devuelve y_pred_aligned y el mapping.
  1693. """
  1694. y_pred = np.asarray(y_pred).astype(int)
  1695. y_ref = np.asarray(y_ref).astype(int)
  1696. Kp = int(y_pred.max()) + 1
  1697. Kr = int(y_ref.max()) + 1
  1698. K = max(Kp, Kr)
  1699. M = np.zeros((K, K), dtype=int)
  1700. for a, b in zip(y_pred, y_ref):
  1701. if a >= 0 and b >= 0:
  1702. M[a, b] += 1
  1703. row_ind, col_ind = linear_sum_assignment(-M)
  1704. mapping = {int(r): int(c) for r, c in zip(row_ind, col_ind)}
  1705. y_aligned = np.array([mapping.get(int(t), int(t)) for t in y_pred], dtype=int)
  1706. return y_aligned, mapping
  1707. # ---------------- SUBJECT INDEXING ----------------
  1708. # ✅ Importante: subject_id como string
  1709. subj_ids = df_out[SUBJ_COL].astype(str).to_numpy()
  1710. uniq_subj = np.unique(subj_ids)
  1711. # Mapa sujeto -> indices de lesiones
  1712. subj_to_idx = {}
  1713. for i, s in enumerate(subj_ids):
  1714. subj_to_idx.setdefault(s, []).append(i)
  1715. # ---------------- STORAGE ----------------
  1716. rng = np.random.default_rng(SUBSEED)
  1717. ari_all_list, ari_in_list, ari_oob_list = [], [], []
  1718. subs_in_counts, subs_oob_counts = [], []
  1719. les_in_counts, les_oob_counts = [], []
  1720. # Para estabilidad por lesión usando SOLO predicciones OOB:
  1721. # - Contamos votos solo cuando esa lesión estuvo OOB en esa iteración
  1722. label_votes_oob = np.zeros((len(df_out), K_REF), dtype=int)
  1723. oob_seen_count = np.zeros((len(df_out),), dtype=int)
  1724. # ---------------- MAIN LOOP ----------------
  1725. n_in_subj = int(np.round(IN_FRAC * len(uniq_subj)))
  1726. for b in range(N_SUBSAMPLE):
  1727. # ---- Elegir sujetos IN (80%) sin reemplazo ----
  1728. in_subj = rng.choice(uniq_subj, size=n_in_subj, replace=False)
  1729. in_set = set(in_subj)
  1730. oob_subj = np.array([s for s in uniq_subj if s not in in_set], dtype=object)
  1731. # ---- Indices de lesiones IN y OOB ----
  1732. idx_in = []
  1733. for s in in_subj:
  1734. idx_in.extend(subj_to_idx[s])
  1735. idx_in = np.array(idx_in, dtype=int)
  1736. idx_oob = []
  1737. for s in oob_subj:
  1738. idx_oob.extend(subj_to_idx[s])
  1739. idx_oob = np.array(idx_oob, dtype=int)
  1740. # ---- Guardar counts ----
  1741. subs_in_counts.append(len(in_subj))
  1742. subs_oob_counts.append(len(oob_subj))
  1743. les_in_counts.append(len(idx_in))
  1744. les_oob_counts.append(len(idx_oob))
  1745. # ---- Fit GMM en IN ----
  1746. X_in = X_z[idx_in, :]
  1747. gmm_b = GaussianMixture(
  1748. n_components=K_FIXED,
  1749. covariance_type=COV_TYPE, # full (como pediste)
  1750. reg_covar=REG_COVAR, # 🔥 estabilidad numérica
  1751. random_state=SUBSEED + b,
  1752. n_init=N_INIT, # 🔥 muchos reinicios
  1753. max_iter=MAX_ITER
  1754. )
  1755. gmm_b.fit(X_in)
  1756. # ---- Predict en TODO ----
  1757. yb_full = gmm_b.predict(X_z).astype(int)
  1758. # ---- Align a labels_ref ----
  1759. yb_aligned, _ = align_labels_hungarian(yb_full, labels_ref)
  1760. # ---- ARI global ----
  1761. ari_all = adjusted_rand_score(labels_ref, yb_aligned)
  1762. ari_all_list.append(float(ari_all))
  1763. # ---- ARI IN y OOB (comparando solo esas lesiones) ----
  1764. if len(idx_in) > 1:
  1765. ari_in = adjusted_rand_score(labels_ref[idx_in], yb_aligned[idx_in])
  1766. else:
  1767. ari_in = np.nan
  1768. if len(idx_oob) > 1:
  1769. ari_oob = adjusted_rand_score(labels_ref[idx_oob], yb_aligned[idx_oob])
  1770. else:
  1771. ari_oob = np.nan
  1772. ari_in_list.append(float(ari_in) if np.isfinite(ari_in) else np.nan)
  1773. ari_oob_list.append(float(ari_oob) if np.isfinite(ari_oob) else np.nan)
  1774. # ---- Votos OOB por lesión (solo cuando estuvo OOB) ----
  1775. for i in idx_oob:
  1776. lab = int(yb_aligned[i])
  1777. if 0 <= lab < K_REF:
  1778. label_votes_oob[i, lab] += 1
  1779. oob_seen_count[i] += 1
  1780. if (b + 1) % 25 == 0:
  1781. print(f"[SUBSAMPLE] {b+1}/{N_SUBSAMPLE} | ARI_all={ari_all:.3f} | ARI_in={ari_in:.3f} | ARI_oob={ari_oob:.3f}")
  1782. # ---------------- SUMMARY ----------------
  1783. ari_all_arr = np.array(ari_all_list, dtype=float)
  1784. ari_in_arr = np.array(ari_in_list, dtype=float)
  1785. ari_oob_arr = np.array(ari_oob_list, dtype=float)
  1786. def summarize(x):
  1787. x = x[np.isfinite(x)]
  1788. if len(x) == 0:
  1789. return np.nan, np.nan, np.nan, np.nan
  1790. mean = float(np.mean(x))
  1791. sd = float(np.std(x, ddof=1)) if len(x) > 1 else 0.0
  1792. med = float(np.median(x))
  1793. iqr1 = float(np.percentile(x, 25))
  1794. iqr3 = float(np.percentile(x, 75))
  1795. return mean, sd, med, (iqr1, iqr3)
  1796. m_all, sd_all, med_all, (p25_all, p75_all) = summarize(ari_all_arr)
  1797. m_in, sd_in, med_in, (p25_in, p75_in) = summarize(ari_in_arr)
  1798. m_oob, sd_oob, med_oob, (p25_oob, p75_oob) = summarize(ari_oob_arr)
  1799. print("\n[DONE] Subsampling reproducibility (80/20 by subject)")
  1800. print(f"Subjects IN (mean): {np.mean(subs_in_counts):.1f} / {len(uniq_subj)} | OOB (mean): {np.mean(subs_oob_counts):.1f}")
  1801. print(f"Lesions IN (mean): {np.mean(les_in_counts):.1f} / {len(df_out)} | OOB (mean): {np.mean(les_oob_counts):.1f}")
  1802. print(f"ARI_all: {m_all:.3f} ± {sd_all:.3f} | med={med_all:.3f} [IQR {p25_all:.3f}-{p75_all:.3f}]")
  1803. print(f"ARI_in: {m_in:.3f} ± {sd_in:.3f} | med={med_in:.3f} [IQR {p25_in:.3f}-{p75_in:.3f}]")
  1804. print(f"ARI_oob: {m_oob:.3f} ± {sd_oob:.3f} | med={med_oob:.3f} [IQR {p25_oob:.3f}-{p75_oob:.3f}]")
  1805. # ---------------- LESION-LEVEL STABILITY (OOB-ONLY) ----------------
  1806. # Proporción modal SOLO en las veces que la lesión estuvo OOB
  1807. vote_prop_oob = np.zeros_like(label_votes_oob, dtype=float)
  1808. valid = oob_seen_count > 0
  1809. vote_prop_oob[valid] = label_votes_oob[valid] / oob_seen_count[valid, None]
  1810. lesion_stability_oob = np.full((len(df_out),), np.nan, dtype=float)
  1811. lesion_label_mode_oob = np.full((len(df_out),), -1, dtype=int)
  1812. lesion_stability_oob[valid] = vote_prop_oob[valid].max(axis=1)
  1813. lesion_label_mode_oob[valid] = vote_prop_oob[valid].argmax(axis=1).astype(int)
  1814. df_out["oob_seen_count"] = oob_seen_count
  1815. df_out["bootstrap_stability_oob"] = lesion_stability_oob
  1816. df_out["bootstrap_label_mode_oob"] = lesion_label_mode_oob
  1817. stab_summary = pd.DataFrame({
  1818. "IN_FRAC": [IN_FRAC],
  1819. "N_SUBSAMPLE": [N_SUBSAMPLE],
  1820. "K_FIXED": [K_FIXED],
  1821. "covariance_type": [COV_TYPE],
  1822. "n_init": [N_INIT],
  1823. "max_iter": [MAX_ITER],
  1824. "reg_covar": [REG_COVAR],
  1825. "Subjects_IN_mean": [float(np.mean(subs_in_counts))],
  1826. "Subjects_OOB_mean": [float(np.mean(subs_oob_counts))],
  1827. "Lesions_IN_mean": [float(np.mean(les_in_counts))],
  1828. "Lesions_OOB_mean": [float(np.mean(les_oob_counts))],
  1829. "ARI_all_mean": [m_all],
  1830. "ARI_all_sd": [sd_all],
  1831. "ARI_all_median": [med_all],
  1832. "ARI_all_p25": [p25_all],
  1833. "ARI_all_p75": [p75_all],
  1834. "ARI_in_mean": [m_in],
  1835. "ARI_in_sd": [sd_in],
  1836. "ARI_oob_mean": [m_oob],
  1837. "ARI_oob_sd": [sd_oob],
  1838. "lesion_stability_oob_mean": [float(np.nanmean(lesion_stability_oob))],
  1839. "lesion_stability_oob_median": [float(np.nanmedian(lesion_stability_oob))],
  1840. })
  1841. # ---------------- SAVE EXCEL ----------------
  1842. sub_xlsx = os.path.join(out_dir, f"{base_name}_subsampling80_fullcov_reproducibility.xlsx")
  1843. with pd.ExcelWriter(sub_xlsx, engine="openpyxl") as writer:
  1844. pd.DataFrame({
  1845. "ARI_all": ari_all_arr,
  1846. "ARI_in": ari_in_arr,
  1847. "ARI_oob": ari_oob_arr
  1848. }).to_excel(writer, index=False, sheet_name="ARI_per_iter")
  1849. stab_summary.to_excel(writer, index=False, sheet_name="summary")
  1850. df_out[[SUBJ_COL, CLUSTER_COL, "oob_seen_count", "bootstrap_label_mode_oob", "bootstrap_stability_oob"]].to_excel(
  1851. writer, index=False, sheet_name="lesion_stability_OOBonly"
  1852. )
  1853. print("[DONE] Subsampling Excel:", sub_xlsx)
  1854. # ======================================================================================
  1855. # ================= REPRODUCIBILITY (80/20 SUBJECT SUBSAMPLING) ========================
  1856. # ===================== (CONSENSUS STABILITY) ==========================
  1857. # ======================================================================================
  1858. # (label-switching proof, practical):
  1859. # - 80% subjects IN (sin reemplazo) / 20% OOB por repetición
  1860. # - Entrena GMM (covariance_type="full") SOLO con lesiones IN
  1861. # - Predice labels para TODAS las lesiones en cada run
  1862. # - Alinea cada run a un "anchor" usando Hungarian (para hacer comparables los IDs)
  1863. # - Calcula:
  1864. # * label_mode por lesión (consenso)
  1865. # * estabilidad por lesión = proporción de runs donde cae en su label_mode
  1866. # * (opcional) entropía por lesión
  1867. # - Reporta estabilidad global + por cluster
  1868. # - Genera plots y guarda Excel con summaries
  1869. # ------------------------------ CONFIG ------------------------------------
  1870. N_RUNS = 1000 # sube a 1000 como pediste
  1871. SEED = 42
  1872. IN_FRAC = 0.80
  1873. K_FIXED = best_k_gmm # tu k elegido por BIC
  1874. COV_TYPE = "full" # mantener full covariance
  1875. # Columns in df_out
  1876. SUBJ_COL = SUBJ_COL # e.g., "subject_id"
  1877. CLUSTER_COL = CLUSTER_COL # e.g., "cluster_gmm_auto" (del fit original, si existe)
  1878. # Output
  1879. out_dir = out_dir
  1880. base_name = base_name
  1881. os.makedirs(out_dir, exist_ok=True)
  1882. # ------------------------------ INPUTS ------------------------------------
  1883. # X_z: (n_lesions, n_features)
  1884. # df_out: dataframe con una fila por lesión
  1885. n_lesions = X_z.shape[0]
  1886. subj_ids = df_out[SUBJ_COL].astype(str).to_numpy()
  1887. uniq_subj = np.unique(subj_ids)
  1888. n_subj = len(uniq_subj)
  1889. print(f"[INFO] Lesions: {n_lesions} | Subjects: {n_subj} | IN_FRAC={IN_FRAC} | K={K_FIXED} | RUNS={N_RUNS}")
  1890. rng = np.random.default_rng(SEED)
  1891. # --------------------------------------------------------------------------
  1892. # STORAGE
  1893. # --------------------------------------------------------------------------
  1894. labels_runs = np.empty((N_RUNS, n_lesions), dtype=np.int32)
  1895. in_mask_runs = np.zeros((N_RUNS, n_lesions), dtype=bool)
  1896. oob_mask_runs = np.zeros((N_RUNS, n_lesions), dtype=bool)
  1897. n_in_subj_list, n_oob_subj_list = [], []
  1898. n_in_les_list, n_oob_les_list = [], []
  1899. # --------------------------------------------------------------------------
  1900. # MAIN LOOP: 80/20 subject subsampling
  1901. # --------------------------------------------------------------------------
  1902. for b in range(N_RUNS):
  1903. n_in = int(np.round(IN_FRAC * n_subj))
  1904. in_subj = rng.choice(uniq_subj, size=n_in, replace=False)
  1905. in_mask = np.isin(subj_ids, in_subj)
  1906. oob_mask = ~in_mask
  1907. in_mask_runs[b] = in_mask
  1908. oob_mask_runs[b] = oob_mask
  1909. idx_in = np.where(in_mask)[0]
  1910. X_in = X_z[idx_in, :]
  1911. gmm_b = GaussianMixture(
  1912. n_components=K_FIXED,
  1913. covariance_type=COV_TYPE,
  1914. random_state=SEED + b,
  1915. n_init=5,
  1916. max_iter=500,
  1917. reg_covar=1e-6
  1918. )
  1919. gmm_b.fit(X_in)
  1920. # labels para TODO el dataset
  1921. labels_runs[b] = gmm_b.predict(X_z).astype(np.int32)
  1922. # bookkeeping
  1923. n_in_subj_list.append(len(in_subj))
  1924. n_oob_subj_list.append(n_subj - len(in_subj))
  1925. n_in_les_list.append(int(in_mask.sum()))
  1926. n_oob_les_list.append(int(oob_mask.sum()))
  1927. if (b + 1) % 50 == 0:
  1928. print(f"[RUN] {b+1}/{N_RUNS} | IN_subj={len(in_subj)} | IN_les={in_mask.sum()}")
  1929. # --------------------------------------------------------------------------
  1930. # LABEL ALIGNMENT (Hungarian) TO ANCHOR
  1931. # --------------------------------------------------------------------------
  1932. def align_labels_hungarian(y_pred, y_ref):
  1933. """
  1934. Alinea etiquetas de y_pred a y_ref maximizando diagonal (conteos).
  1935. Devuelve y_pred_aligned.
  1936. """
  1937. y_pred = np.asarray(y_pred).astype(int)
  1938. y_ref = np.asarray(y_ref).astype(int)
  1939. K = max(int(y_pred.max()) + 1, int(y_ref.max()) + 1)
  1940. M = np.zeros((K, K), dtype=int)
  1941. for a, b in zip(y_pred, y_ref):
  1942. if a >= 0 and b >= 0:
  1943. M[a, b] += 1
  1944. r, c = linear_sum_assignment(-M)
  1945. mapping = {int(rr): int(cc) for rr, cc in zip(r, c)}
  1946. y_aligned = np.vectorize(lambda t: mapping.get(int(t), int(t)))(y_pred)
  1947. return y_aligned.astype(int)
  1948. anchor_idx = 0
  1949. labels_anchor = labels_runs[anchor_idx].copy()
  1950. labels_aligned_runs = np.empty_like(labels_runs)
  1951. labels_aligned_runs[anchor_idx] = labels_anchor
  1952. for b in range(N_RUNS):
  1953. if b == anchor_idx:
  1954. continue
  1955. labels_aligned_runs[b] = align_labels_hungarian(labels_runs[b], labels_anchor)
  1956. # --------------------------------------------------------------------------
  1957. # CONSENSUS: MODE + STABILITY + (OPTIONAL) ENTROPY
  1958. # --------------------------------------------------------------------------
  1959. label_votes = np.zeros((n_lesions, K_FIXED), dtype=np.int32)
  1960. # más rápido que loop i por i: acumular con np.add.at
  1961. for b in range(N_RUNS):
  1962. yb = labels_aligned_runs[b]
  1963. # proteger por si aparece alguna etiqueta rara fuera de rango
  1964. yb = np.clip(yb, 0, K_FIXED - 1)
  1965. np.add.at(label_votes, (np.arange(n_lesions), yb), 1)
  1966. vote_prop = label_votes / np.maximum(label_votes.sum(axis=1, keepdims=True), 1)
  1967. lesion_mode = vote_prop.argmax(axis=1).astype(int)
  1968. lesion_stability = vote_prop.max(axis=1).astype(float)
  1969. # Entropy (opcional)
  1970. eps = 1e-12
  1971. lesion_entropy = (-np.sum(vote_prop * np.log(vote_prop + eps), axis=1)).astype(float)
  1972. print("\n[APPROACH 2 DONE] Anchor-aligned consensus stability")
  1973. print(f"Stability mean ± SD: {lesion_stability.mean():.3f} ± {lesion_stability.std(ddof=1):.3f}")
  1974. print(f"Stability median [IQR]: {np.median(lesion_stability):.3f} "
  1975. f"[{np.percentile(lesion_stability,25):.3f}-{np.percentile(lesion_stability,75):.3f}]")
  1976. # --------------------------------------------------------------------------
  1977. # IN/OOB SUMMARY
  1978. # --------------------------------------------------------------------------
  1979. print("\n[IN/OOB SUMMARY] (across runs)")
  1980. print(f"Subjects IN (mean): {np.mean(n_in_subj_list):.1f} / {n_subj} | OOB (mean): {np.mean(n_oob_subj_list):.1f}")
  1981. print(f"Lesions IN (mean): {np.mean(n_in_les_list):.1f} / {n_lesions} | OOB (mean): {np.mean(n_oob_les_list):.1f}")
  1982. # --------------------------------------------------------------------------
  1983. # SAVE BACK TO df_out
  1984. # --------------------------------------------------------------------------
  1985. df_out["subsample80_label_mode"] = lesion_mode
  1986. df_out["subsample80_stability"] = lesion_stability
  1987. df_out["subsample80_entropy"] = lesion_entropy
  1988. # --------------------------------------------------------------------------
  1989. # CLUSTER-SPECIFIC STABILITY (por cluster consenso o por cluster original si existe)
  1990. # --------------------------------------------------------------------------
  1991. # Si quieres por cluster original (del fit final), usa CLUSTER_COL si existe.
  1992. # Si no, usamos el consenso (lesion_mode).
  1993. use_original_cluster = (CLUSTER_COL in df_out.columns)
  1994. cluster_group_col = CLUSTER_COL if use_original_cluster else "subsample80_label_mode"
  1995. df_out[cluster_group_col] = df_out[cluster_group_col].astype(int)
  1996. stab_by_cluster = (
  1997. df_out.groupby(cluster_group_col)["subsample80_stability"]
  1998. .agg(["count", "mean", "std", "median"])
  1999. .reset_index()
  2000. .rename(columns={cluster_group_col: "cluster"})
  2001. )
  2002. stab_by_cluster["p25"] = df_out.groupby(cluster_group_col)["subsample80_stability"].quantile(0.25).values
  2003. stab_by_cluster["p75"] = df_out.groupby(cluster_group_col)["subsample80_stability"].quantile(0.75).values
  2004. ent_by_cluster = (
  2005. df_out.groupby(cluster_group_col)["subsample80_entropy"]
  2006. .agg(["count", "mean", "std", "median"])
  2007. .reset_index()
  2008. .rename(columns={cluster_group_col: "cluster"})
  2009. )
  2010. ent_by_cluster["p25"] = df_out.groupby(cluster_group_col)["subsample80_entropy"].quantile(0.25).values
  2011. ent_by_cluster["p75"] = df_out.groupby(cluster_group_col)["subsample80_entropy"].quantile(0.75).values
  2012. print("\n[STABILITY BY CLUSTER]")
  2013. print(stab_by_cluster)
  2014. # --------------------------------------------------------------------------
  2015. # REPLICABILITY
  2016. # --------------------------------------------------------------------------
  2017. import matplotlib.pyplot as plt
  2018. import numpy as np
  2019. import os
  2020. plot_prefix = f"{base_name}_subsample80_runs{N_RUNS}"
  2021. os.makedirs(out_dir, exist_ok=True)
  2022. # ------------------ choose cluster grouping for plots ------------------
  2023. # Recommended: consensus clusters (lesion_mode) because it matches approach 2
  2024. cluster_group_col = "subsample80_label_mode"
  2025. # If you REALLY want the original clusters from full-data fit, uncomment:
  2026. # if (CLUSTER_COL in df_out.columns):
  2027. # cluster_group_col = CLUSTER_COL
  2028. df_out[cluster_group_col] = df_out[cluster_group_col].astype(int)
  2029. # Ensure we have exactly the clusters present (usually 0..K-1)
  2030. clusters_sorted = sorted(df_out[cluster_group_col].unique())
  2031. # --------------------------------------------------------------------------
  2032. # SAVE EXCEL
  2033. # --------------------------------------------------------------------------
  2034. summary_df = pd.DataFrame({
  2035. "N_RUNS": [N_RUNS],
  2036. "IN_FRAC": [IN_FRAC],
  2037. "K_FIXED": [K_FIXED],
  2038. "COV_TYPE": [COV_TYPE],
  2039. "subjects_total": [n_subj],
  2040. "lesions_total": [n_lesions],
  2041. "subjects_in_mean": [float(np.mean(n_in_subj_list))],
  2042. "subjects_oob_mean": [float(np.mean(n_oob_subj_list))],
  2043. "lesions_in_mean": [float(np.mean(n_in_les_list))],
  2044. "lesions_oob_mean": [float(np.mean(n_oob_les_list))],
  2045. "stability_mean": [float(lesion_stability.mean())],
  2046. "stability_sd": [float(lesion_stability.std(ddof=1))],
  2047. "stability_median": [float(np.median(lesion_stability))],
  2048. "stability_p25": [float(np.percentile(lesion_stability, 25))],
  2049. "stability_p75": [float(np.percentile(lesion_stability, 75))],
  2050. "entropy_mean": [float(lesion_entropy.mean())],
  2051. "entropy_sd": [float(lesion_entropy.std(ddof=1))],
  2052. "entropy_median": [float(np.median(lesion_entropy))],
  2053. "entropy_p25": [float(np.percentile(lesion_entropy, 25))],
  2054. "entropy_p75": [float(np.percentile(lesion_entropy, 75))],
  2055. "cluster_grouping": ["original" if use_original_cluster else "consensus"],
  2056. })
  2057. xlsx_path = os.path.join(out_dir, f"{plot_prefix}_approach2_only.xlsx")
  2058. with pd.ExcelWriter(xlsx_path, engine="openpyxl") as writer:
  2059. summary_df.to_excel(writer, index=False, sheet_name="summary")
  2060. stab_by_cluster.to_excel(writer, index=False, sheet_name="stability_by_cluster")
  2061. ent_by_cluster.to_excel(writer, index=False, sheet_name="entropy_by_cluster")
  2062. cols_to_save = [SUBJ_COL, CLUSTER_COL, "subsample80_label_mode", "subsample80_stability", "subsample80_entropy"]
  2063. cols_to_save = [c for c in cols_to_save if c in df_out.columns]
  2064. df_out[cols_to_save].to_excel(writer, index=False, sheet_name="lesion_level")
  2065. print("\n[DONE] Saved Excel:", xlsx_path)
  2066. # ------------------------------------------------------------------
  2067. # RUN-LEVEL REPLICABILITY (ALL vs OOB)
  2068. # ------------------------------------------------------------------
  2069. run_stability_all = []
  2070. run_stability_oob = []
  2071. for b in range(N_RUNS):
  2072. # ALL lesions
  2073. match_all = labels_aligned_runs[b] == lesion_mode
  2074. run_stability_all.append(np.mean(match_all))
  2075. # OOB lesions only
  2076. mask_oob = oob_mask_runs[b]
  2077. if mask_oob.sum() > 0:
  2078. match_oob = labels_aligned_runs[b][mask_oob] == lesion_mode[mask_oob]
  2079. run_stability_oob.append(np.mean(match_oob))
  2080. else:
  2081. run_stability_oob.append(np.nan)
  2082. run_stability_all = np.array(run_stability_all)
  2083. run_stability_oob = np.array(run_stability_oob)
  2084. print("\n[RUN-LEVEL STABILITY]")
  2085. print(f"ALL mean ± SD: {run_stability_all.mean():.3f} ± {run_stability_all.std(ddof=1):.3f}")
  2086. print(f"OOB mean ± SD: {np.nanmean(run_stability_oob):.3f} ± {np.nanstd(run_stability_oob, ddof=1):.3f}")
  2087. # --------------------------------------------------------------------------
  2088. # RUN-LEVEL + LESION-LEVEL STABILITY (ALL vs IN vs OOB) + PLOTS (NO CLUSTERS)
  2089. # --------------------------------------------------------------------------
  2090. import os
  2091. import numpy as np
  2092. import pandas as pd
  2093. import matplotlib.pyplot as plt
  2094. plot_prefix = f"{base_name}_subsample80_runs{N_RUNS}"
  2095. os.makedirs(out_dir, exist_ok=True)
  2096. # -------------------------
  2097. # CONSENSUS LABELS
  2098. # -------------------------
  2099. consensus = lesion_mode.astype(int) # shape: (n_lesions,)
  2100. n_lesions = consensus.shape[0]
  2101. # -------------------------
  2102. # RUN-LEVEL STABILITY (agreement with consensus)
  2103. # - ALL: promedio sobre todas las lesiones
  2104. # - IN: promedio solo sobre lesiones IN de ese run
  2105. # - OOB: promedio solo sobre lesiones OOB de ese run
  2106. # -------------------------
  2107. run_stability_all = np.full(N_RUNS, np.nan, dtype=float)
  2108. run_stability_in = np.full(N_RUNS, np.nan, dtype=float)
  2109. run_stability_oob = np.full(N_RUNS, np.nan, dtype=float)
  2110. n_in_lesions_per_run = np.zeros(N_RUNS, dtype=int)
  2111. n_oob_lesions_per_run = np.zeros(N_RUNS, dtype=int)
  2112. for b in range(N_RUNS):
  2113. yb = labels_aligned_runs[b]
  2114. # ALL
  2115. run_stability_all[b] = np.mean(yb == consensus)
  2116. # IN
  2117. mask_in = in_mask_runs[b]
  2118. n_in = int(mask_in.sum())
  2119. n_in_lesions_per_run[b] = n_in
  2120. if n_in > 0:
  2121. run_stability_in[b] = np.mean(yb[mask_in] == consensus[mask_in])
  2122. # OOB
  2123. mask_oob = oob_mask_runs[b]
  2124. n_oob = int(mask_oob.sum())
  2125. n_oob_lesions_per_run[b] = n_oob
  2126. if n_oob > 0:
  2127. run_stability_oob[b] = np.mean(yb[mask_oob] == consensus[mask_oob])
  2128. print("\n[RUN-LEVEL STABILITY vs CONSENSUS]")
  2129. print(f"ALL mean ± SD: {np.nanmean(run_stability_all):.3f} ± {np.nanstd(run_stability_all, ddof=1):.3f}")
  2130. print(f"IN mean ± SD: {np.nanmean(run_stability_in):.3f} ± {np.nanstd(run_stability_in, ddof=1):.3f}")
  2131. print(f"OOB mean ± SD: {np.nanmean(run_stability_oob):.3f} ± {np.nanstd(run_stability_oob, ddof=1):.3f}")
  2132. # -------------------------
  2133. # LESION-LEVEL STABILITY
  2134. # A) ALL-runs (como antes): max proportion en vote_prop (ya lo puedes tener)
  2135. # si NO lo tienes, lo recalculamos aquí desde labels_aligned_runs
  2136. # B) OOB-only: votos solo cuando esa lesión estuvo OOB en ese run
  2137. # -------------------------
  2138. # A) ALL-runs lesion stability
  2139. label_votes_all = np.zeros((n_lesions, K_FIXED), dtype=np.int32)
  2140. for b in range(N_RUNS):
  2141. yb = np.clip(labels_aligned_runs[b], 0, K_FIXED - 1)
  2142. np.add.at(label_votes_all, (np.arange(n_lesions), yb), 1)
  2143. vote_prop_all = label_votes_all / np.maximum(label_votes_all.sum(axis=1, keepdims=True), 1)
  2144. lesion_label_mode_all = vote_prop_all.argmax(axis=1).astype(int)
  2145. lesion_stability_all = vote_prop_all.max(axis=1).astype(float)
  2146. # B) OOB-only lesion stability
  2147. label_votes_oob = np.zeros((n_lesions, K_FIXED), dtype=np.int32)
  2148. oob_seen_count = np.zeros((n_lesions,), dtype=np.int32)
  2149. for b in range(N_RUNS):
  2150. mask_oob = oob_mask_runs[b]
  2151. if mask_oob.sum() == 0:
  2152. continue
  2153. yb = np.clip(labels_aligned_runs[b], 0, K_FIXED - 1)
  2154. idx = np.where(mask_oob)[0]
  2155. np.add.at(label_votes_oob, (idx, yb[idx]), 1)
  2156. oob_seen_count[idx] += 1
  2157. valid_oob = oob_seen_count > 0
  2158. vote_prop_oob = np.zeros_like(label_votes_oob, dtype=float)
  2159. vote_prop_oob[valid_oob] = label_votes_oob[valid_oob] / oob_seen_count[valid_oob, None]
  2160. lesion_label_mode_oob = np.full(n_lesions, -1, dtype=int)
  2161. lesion_stability_oob = np.full(n_lesions, np.nan, dtype=float)
  2162. lesion_label_mode_oob[valid_oob] = vote_prop_oob[valid_oob].argmax(axis=1).astype(int)
  2163. lesion_stability_oob[valid_oob] = vote_prop_oob[valid_oob].max(axis=1).astype(float)
  2164. print("\n[LESION-LEVEL STABILITY]")
  2165. print(f"ALL-runs mean ± SD: {lesion_stability_all.mean():.3f} ± {lesion_stability_all.std(ddof=1):.3f}")
  2166. print(f"OOB-only mean ± SD: {np.nanmean(lesion_stability_oob):.3f} ± {np.nanstd(lesion_stability_oob, ddof=1):.3f}")
  2167. print(f"OOB-only coverage (mean oob_seen_count): {float(np.mean(oob_seen_count)):.1f} runs per lesion")
  2168. # Guardar en df_out
  2169. df_out["consensus_label_mode_all"] = lesion_label_mode_all
  2170. df_out["stability_allruns"] = lesion_stability_all
  2171. df_out["oob_seen_count"] = oob_seen_count
  2172. df_out["consensus_label_mode_oob"] = lesion_label_mode_oob
  2173. df_out["stability_oobonly"] = lesion_stability_oob
  2174. # -------------------------
  2175. # PLOTS (OOB)
  2176. # -------------------------
  2177. x = np.arange(1, N_RUNS + 1)
  2178. # Plot: solo OOB (línea)
  2179. plt.figure(figsize=(8, 4))
  2180. plt.plot(x, run_stability_oob, label="OOB", linewidth=2)
  2181. plt.ylim(0, 1) # <- pedido
  2182. plt.xlim(0, 1000)
  2183. plt.xlabel("Run")
  2184. plt.ylabel("Proportion matching consensus")
  2185. plt.legend(frameon=False)
  2186. plt.tight_layout()
  2187. run_oob_path = os.path.join(out_dir, f"{plot_prefix}_run_stability_OOB_only.png")
  2188. plt.savefig(run_oob_path, dpi=300)
  2189. plt.show()
  2190. print("\n[PLOTS SAVED]")
  2191. print(run_oob_path)
  2192. # -------------------------
  2193. # 95% CI helper (sobre promedio)
  2194. # -------------------------
  2195. def mean_sd_ci95(x):
  2196. x = np.asarray(x, dtype=float)
  2197. x = x[np.isfinite(x)]
  2198. n = len(x)
  2199. if n == 0:
  2200. return np.nan, np.nan, np.nan, np.nan, 0
  2201. m = float(np.mean(x))
  2202. sd = float(np.std(x, ddof=1)) if n > 1 else 0.0
  2203. se = sd / np.sqrt(n) if n > 0 else np.nan
  2204. ci_lo = m - 1.96 * se
  2205. ci_hi = m + 1.96 * se
  2206. return m, sd, ci_lo, ci_hi, n
  2207. m_all, sd_all, ci_all_lo, ci_all_hi, n_all = mean_sd_ci95(run_stability_all)
  2208. m_in, sd_in, ci_in_lo, ci_in_hi, n_in = mean_sd_ci95(run_stability_in)
  2209. m_oob, sd_oob, ci_oob_lo, ci_oob_hi, n_oob = mean_sd_ci95(run_stability_oob)
  2210. m_lall, sd_lall, ci_lall_lo, ci_lall_hi, n_lall = mean_sd_ci95(lesion_stability_all)
  2211. m_loob, sd_loob, ci_loob_lo, ci_loob_hi, n_loob = mean_sd_ci95(lesion_stability_oob)
  2212. # -------------------------
  2213. # SAVE EXCEL (separado ALL / IN / OOB)
  2214. # -------------------------
  2215. run_level_df = pd.DataFrame({
  2216. "run": np.arange(N_RUNS),
  2217. "n_in_lesions": n_in_lesions_per_run,
  2218. "n_oob_lesions": n_oob_lesions_per_run,
  2219. "stability_ALL": run_stability_all,
  2220. "stability_IN": run_stability_in,
  2221. "stability_OOB": run_stability_oob
  2222. })
  2223. summary_df = pd.DataFrame({
  2224. "N_RUNS": [N_RUNS],
  2225. "IN_FRAC": [IN_FRAC],
  2226. "K_FIXED": [K_FIXED],
  2227. "COV_TYPE": [COV_TYPE],
  2228. # Run-level stability (vs consensus)
  2229. "run_stability_ALL_mean": [m_all],
  2230. "run_stability_ALL_sd": [sd_all],
  2231. "run_stability_ALL_ci95_lo": [ci_all_lo],
  2232. "run_stability_ALL_ci95_hi": [ci_all_hi],
  2233. "run_stability_ALL_n": [n_all],
  2234. "run_stability_IN_mean": [m_in],
  2235. "run_stability_IN_sd": [sd_in],
  2236. "run_stability_IN_ci95_lo": [ci_in_lo],
  2237. "run_stability_IN_ci95_hi": [ci_in_hi],
  2238. "run_stability_IN_n": [n_in],
  2239. "run_stability_OOB_mean": [m_oob],
  2240. "run_stability_OOB_sd": [sd_oob],
  2241. "run_stability_OOB_ci95_lo": [ci_oob_lo],
  2242. "run_stability_OOB_ci95_hi": [ci_oob_hi],
  2243. "run_stability_OOB_n": [n_oob],
  2244. # OOB coverage
  2245. "oob_seen_count_mean": [float(np.mean(oob_seen_count))],
  2246. "oob_seen_count_min": [int(np.min(oob_seen_count))],
  2247. "oob_seen_count_max": [int(np.max(oob_seen_count))],
  2248. })
  2249. lesion_all_df = df_out[[SUBJ_COL]].copy()
  2250. if CLUSTER_COL in df_out.columns:
  2251. lesion_all_df[CLUSTER_COL] = df_out[CLUSTER_COL]
  2252. lesion_all_df["consensus_label_mode_all"] = df_out["consensus_label_mode_all"]
  2253. lesion_all_df["stability_allruns"] = df_out["stability_allruns"]
  2254. lesion_oob_df = df_out[[SUBJ_COL]].copy()
  2255. if CLUSTER_COL in df_out.columns:
  2256. lesion_oob_df[CLUSTER_COL] = df_out[CLUSTER_COL]
  2257. lesion_oob_df["oob_seen_count"] = df_out["oob_seen_count"]
  2258. lesion_oob_df["consensus_label_mode_oob"] = df_out["consensus_label_mode_oob"]
  2259. lesion_oob_df["stability_oobonly"] = df_out["stability_oobonly"]
  2260. xlsx_path = os.path.join(out_dir, f"{plot_prefix}_stability_ALL_IN_OOB.xlsx")
  2261. with pd.ExcelWriter(xlsx_path, engine="openpyxl") as writer:
  2262. summary_df.to_excel(writer, index=False, sheet_name="summary")
  2263. run_level_df.to_excel(writer, index=False, sheet_name="run_level_ALL_IN_OOB")
  2264. lesion_all_df.to_excel(writer, index=False, sheet_name="lesion_level_ALL")
  2265. print("\n[DONE] Saved Excel:", xlsx_path)

clustering.py at commit 1a20057, no license · at the source

Overview

  1. Neuroinformatics for Personalized Medicine Laboratory, Department of Neurology and Neurosurgery, Montreal Neurological Institute-Hospital, McGill University, Quebec, Canada
  2. McConnell Brain Imaging Centre, Montreal Neurological Institute-Hospital, Quebec, Canada
  3. Departamento de Física, Universidad de Buenos Aires, Argentina
  4. Latin American Brain Health Institute (BrainLat), Universidad Adolfo Ibañez, Santiago, Chile
  5. Cognitive Neuroscience Center, Universidad de San Andrés, Victoria, Buenos Aires, Argentina
  6. Consejo Nacional de Investigaciones Científicas y Técnicas (CONICET), Ciudad Autónoma de Buenos Aires, Argentina
  7. Pontificia Universidad Católica de Chile, Santiago, Chile
  8. Rush Alzheimer's Disease Center, Rush University Medical Center, Chicago, IL
  9. Department of Neurological Sciences, Rush University Medical Center, Chicago, IL
  10. Instituto de Assistência Médica ao Servidor Público Estadual, São Paulo, Brazil; and
  11. Ludmer Centre for Neuroinformatics & Mental Health, Montreal, Quebec, Canada
Journal: Neurology, volume 107, issue 6, article e218472
Dates: received 23 March 2026; accepted 23 June 2026; published online 27 August 2026; in print 22 September 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1212/wnl.0000000000218472 · PMID 42659615 · PMCID PMC13528905 · OpenAlex W7204472916
Open access: hybrid, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism), Alzheimer's / dementia (population), stroke (population), Parkinson's (population), clinical / translational (subfield)
Methods: Spectral & time-frequency, Connectivity, Statistics, Complexity, fMRI & imaging, Smoothing, state filtering, decompositions
MeSH: Aging*, Alzheimer Disease*, Brain*, Cognitive Dysfunction*, Parkinson Disease*, White Matter*, Aged, Aged, 80 and over, Female, Humans, Longitudinal Studies, Magnetic Resonance Imaging, Male (* major topic)
Topic: Advanced Neuroimaging Techniques and Applications (Radiology, Nuclear Medicine and Imaging, Medicine), according to OpenAlex
Citations: not cited yet (Europe PMC); 47 references in the paper

Abstract

The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.

Repository

Its files are read in the Code ↔ Paper reader above, with 9 matches between paragraphs and lines of code.

rglezgz/WMH-subtypes

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 1a20057061e82542c37df6249c992b4cb9189d48, 24 May 2026
Languages: Python (7)
Size: 77 files, 7 scripts
Software Heritage: not archived
Found in: “Data Availability”
Holds: README
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: NumPy (7 files), pandas (7 files), SciPy (7 files), NiBabel (5 files), Nilearn (5 files), Matplotlib (4 files), scikit-learn (2 files), statsmodels (2 files), LightGBM (1 file), Plotly (1 file), SHAP (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
8 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 7 scripts, each with its path and the digest of its content;
  • 9 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

No dataset and no data link were found in the paper.

Code and data availability statement

The paper has a code and data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1212/wnl.0000000000218472.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 2, 28 September 2026

  • Publisher: n/a → Lippincott Williams & Wilkins

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 6 authors, 13 MeSH terms, 42 references.

Cite

This paper

Gonzalez-Gomez, R., Tagliazuchi, E., Campo, C. G., Medel, V., Bennett, D. A., & Iturria-Medina, Y. (2026). Lesion-Level Subtypes of White Matter Hyperintensity Evolution Beyond Spatial Location. Neurology, 107(6), e218472. https://doi.org/10.1212/wnl.0000000000218472

BibTeX

@article{gonzalezgomez2026lesion,
author = {Gonzalez-Gomez, Raul and Tagliazuchi, Enzo and Campo, Cecilia Gonzalez and Medel, Vicente and Bennett, David A and Iturria-Medina, Yasser},
title = {{Lesion-Level Subtypes of White Matter Hyperintensity Evolution Beyond Spatial Location}},
journal = {Neurology},
year = {2026},
month = aug,
volume = {107},
number = {6},
pages = {e218472},
publisher = {Lippincott Williams \& Wilkins},
issn = {0028-3878},
doi = {10.1212/wnl.0000000000218472},
url = {https://doi.org/10.1212/wnl.0000000000218472},
pmid = {42659615},
pmcid = {PMC13528905}
}

RIS

TY - JOUR
AU - Gonzalez-Gomez, Raul
AU - Tagliazuchi, Enzo
AU - Campo, Cecilia Gonzalez
AU - Medel, Vicente
AU - Bennett, David A
AU - Iturria-Medina, Yasser
TI - Lesion-Level Subtypes of White Matter Hyperintensity Evolution Beyond Spatial Location
T2 - Neurology
J2 - Neurology
PY - 2026
DA - 2026/08/27
VL - 107
IS - 6
SP - e218472
SN - 0028-3878
PB - Lippincott Williams & Wilkins
DO - 10.1212/wnl.0000000000218472
UR - https://doi.org/10.1212/wnl.0000000000218472
LA - en
ER -

CSL-JSON

{
"id": "10.1212/wnl.0000000000218472",
"type": "article-journal",
"title": "Lesion-Level Subtypes of White Matter Hyperintensity Evolution Beyond Spatial Location",
"container-title": "Neurology",
"author": [
{
"family": "Gonzalez-Gomez",
"given": "Raul"
},
{
"family": "Tagliazuchi",
"given": "Enzo"
},
{
"family": "Campo",
"given": "Cecilia Gonzalez"
},
{
"family": "Medel",
"given": "Vicente"
},
{
"family": "Bennett",
"given": "David A"
},
{
"family": "Iturria-Medina",
"given": "Yasser"
}
],
"container-title-short": "Neurology",
"volume": "107",
"issue": "6",
"page": "e218472",
"DOI": "10.1212/wnl.0000000000218472",
"PMID": "42659615",
"PMCID": "PMC13528905",
"ISSN": "0028-3878",
"publisher": "Lippincott Williams & Wilkins",
"URL": "https://doi.org/10.1212/wnl.0000000000218472",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
27
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1038/s41467-026-73996-z [code]
Genetic architecture of white matter microstructure captured by unsupervised deep representation learning of fractional anisotropy maps.
Journal: Nature communications
In common: LightGBM, Nilearn, Plotly, 7 other tools, structural MRI / diffusion, 1 reference
[2] doi:10.1016/j.nicl.2026.104001 [code]
Effect of vascular lesion preprocessing on Brain Intensity AbNormality Classification Algorithm (BIANCA) white matter hyperintensity segmentation.
Journal: NeuroImage. Clinical
In common: SHAP, NiBabel, scikit-learn, 4 other tools, stroke, structural MRI / diffusion, 3 references
[3] doi:10.1186/s12938-026-01555-0 [code]
Incorporating normal periventricular changes for enhanced pathological white matter hyperintensity segmentation: on multiclass deep learning approaches.
Journal: Biomedical engineering online
In common: scikit-learn, pandas, SciPy, 2 other tools, stroke, structural MRI / diffusion, 4 references
[4] doi:10.7554/elife.103097 [code]
Canonical neurodevelopmental trajectories of structural and functional manifolds.
Journal: eLife
In common: Nilearn, Plotly, NiBabel, 5 other tools, structural MRI / diffusion, 2 references
[5] doi:10.1371/journal.pcbi.1013463 [code]
A multi-frequency whole-brain neural mass model with homeostatic feedback inhibition.
Journal: PLoS computational biology
In common: Nilearn, NiBabel, scikit-learn, 3 other tools, author Raul Gonzalez-Gomez
[6] doi:10.1038/s41598-026-56688-y [code]
On the value of radiomics in addition to clinical measures in emotional conflict fMRI for predicting sertraline response in major depressive disorder.
Journal: Scientific reports
In common: SHAP, Nilearn, NiBabel, 6 other tools, 1 reference
[7] doi:10.1162/imag.a.1269 [code]
From early to contemporary normative modeling: Mapping individual differences in neurophysiological signals.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: Nilearn, Plotly, NiBabel, 6 other tools, 1 reference
[8] doi:10.1162/imag.a.1164 [code]
Bias and generalizability of brain age prediction models: A multi-cohort evaluation with anatomical and interpretability insights.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: Nilearn, Plotly, NiBabel, 6 other tools, Alzheimer's / dementia, structural MRI / diffusion
[9] doi:10.1007/s00415-026-13884-0 [code]
Spatial analysis of paraneoplastic cerebellar degeneration in ovarian cancer with anti-Yo syndrome and SCA1.
Journal: Journal of neurology
In common: Nilearn, Plotly, NiBabel, 6 other tools, structural MRI / diffusion, clinical / translational
[10] doi:10.1371/journal.pbio.3003856 [code]
Aging and metabolism contribute separately to brain-body health.
Journal: PLoS biology
In common: Nilearn, NiBabel, statsmodels, 5 other tools, structural MRI / diffusion, clinical / translational, 1 reference

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.