OSCR

Quantifying generalization error in machine learning prediction of cognitive decline.

Code ↔ Paper

4 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 4 matches
  1. [1] § Methods › Structural MRI data ↔ 01_Analysis.py, lines 512–600 · score 0.99 · right cerebellar white, mid anterior, mid posterior, cerebral white matter, cortical thickness, right lateral
  2. [2] § Results › Most predictive features ↔ 01_Analysis.py, lines 512–600 · score 0.86 · word fluency, right thalamus, right amygdala, right hippocampus, corpus callosum, accumbens
  3. [3] § Methods › Predictive analysis ↔ 01_Analysis.py, lines 120–163 · score 0.61 · Scikit learn, random forest, regression, MMSE, SOB, predictive
  4. [4] § Methods › Model evaluation ↔ 01_Analysis.py, lines 760–792 · score 0.54 · squared error, absolute error, metrics, MSE, MAE, splits

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 · 1,371 lines · 65 KB · MIT · 4 matches

  1. # %%
  2. # # 01_Analysis
  3. # %%
  4. # In summary, 01_Analysis.py contains code for predictive modeling and statistical analyses comparing cognitive decline prediction across OASIS-3 and ADNI datasets. The code is structured as follows:
  5. #
  6. # 0. Setup and Helper Functions
  7. # 0.1 Package imports
  8. # 0.2 Helper functions for data processing, model fitting, and evaluation
  9. #
  10. # 1. Data Loading and Within-Dataset Predictions
  11. # 1.1 Data loading and preprocessing
  12. # 1.2 Within-dataset predictions (OASIS-3 and ADNI)
  13. # 1.3 Results aggregation and summary
  14. #
  15. # 2. Permutation Importance Analysis
  16. # 2.1 Permutation importance for individual features
  17. # 2.2 Feature importance visualization and analysis
  18. # 2.3 Coalition-based feature importance
  19. #
  20. # 3. Across-Dataset Predictions
  21. # 3.1 Cross-dataset prediction (clinical, structural, combined, top-15)
  22. # 3.2 Cross-dataset results aggregation
  23. #
  24. # 4. Statistical Model Comparisons
  25. # 4.1 Statistical Testing Framework and Model Comparisons
  26. # 4.2 Absolute Error Comparisons Between Models
  27. #
  28. # 5. Subgroup Analysis
  29. # 5.1 Absolute error analysis preparation
  30. # 5.2 Subgroup comparisons by diagnosis, sex, APOE E4, and age
  31. # 5.3 Results formatting and export
  32. #
  33. # 6. Post-Hoc Matched Samples Analysis
  34. # 6.1 Quantile-based sample matching
  35. # 6.2 Matched sample predictions and evaluation
  36. # 6.3 Results for both combined and top-15 models
  37. #
  38. # Note: Running the predictive models is computationally expensive with n_splits=1000. Consider reducing splits for testing.
  39. # Full results tables are available in the 02_Supplementary_Material folder at https://osf.io/up65f/files/osfstorage.
  40. # %%
  41. # ## 0. Setup and Helper Functions
  42. # ### 0.1 Package Imports
  43. import joblib
  44. import pandas as pd
  45. import numpy as np
  46. import seaborn as sns
  47. import seaborn.objects as so
  48. import sklearn
  49. from sklearn import metrics
  50. from sklearn.pipeline import make_pipeline
  51. from sklearn.ensemble import RandomForestRegressor
  52. from sklearn.multioutput import MultiOutputRegressor
  53. from sklearn.experimental import enable_iterative_imputer
  54. from sklearn.impute import IterativeImputer
  55. from sklearn.metrics import r2_score
  56. from sklearn.model_selection import ShuffleSplit
  57. from sklearn.metrics import r2_score, mean_absolute_error, mean_squared_error
  58. from sklearn.inspection import permutation_importance
  59. import matplotlib
  60. import matplotlib.pyplot as plt
  61. matplotlib.rcParams['font.family'] = ['Arial']
  62. from scipy import stats
  63. from scipy.stats import wilcoxon
  64. from scipy.stats import mannwhitneyu
  65. from tqdm.notebook import tqdm
  66. from pathlib import Path
  67. from joblib import Parallel, delayed
  68. np.random.seed(21) # by setting a seed, you can ensure that every time you run the code, the sequence of random numbers generated will be the same
  69. # %%
  70. # ### 0.2 Helper Functions for Data Processing, Model Fitting, and Evaluation
  71. def get_data(clinical_features_path, data_slopes_path, structural_data_path, missing_cols=[]):
  72. clinical_features = pd.read_pickle(clinical_features_path)
  73. data_slopes = pd.read_pickle(data_slopes_path)
  74. structural_data = pd.read_pickle(structural_data_path)
  75. # clinical_features.columns.str.strip() str --> columns as strings not as indexes
  76. clinical_features.columns = clinical_features.columns.str.removeprefix("clin__npsy__").str.removeprefix("clin__risk__").str.removeprefix("clin__assess__")
  77. clinical_features["demo_sex"] = clinical_features["demo_sex"].map({'F': 0, 'M': 1, 0: 0, 1: 1})
  78. clinical_features["diag"] = clinical_features["diag"].map({"hc": 0, "mci": 1, "dem": 2.0, "NL": 0.0, "MCI": 1.0, "Dementia": 2.0})
  79. # defining the dependent and independent variables and removing the variable "subject" from the dataframes
  80. y = data_slopes.copy().set_index("subject")
  81. X_clin = clinical_features.copy().set_index("subject")
  82. X_fs = structural_data.copy().set_index("subject")
  83. X_clin_fs = pd.concat([X_clin, X_fs], axis=1)
  84. # remove missing columns if specified
  85. if missing_cols:
  86. X_clin = X_clin.drop(columns=missing_cols, errors='ignore')
  87. X_fs = X_fs.drop(columns=missing_cols, errors='ignore')
  88. X_clin_fs = X_clin_fs.drop(columns=missing_cols, errors='ignore')
  89. return y, X_clin, X_fs, X_clin_fs
  90. def get_iterative_imputer(X_train, X_test, random_state):
  91. imputer = IterativeImputer(add_indicator=True, random_state=random_state, n_nearest_features=50)
  92. missing_cols_train = [c + "_missing" for c in X_train.columns[X_train.isna().any()]]
  93. missing_cols_test = [c + "_missing" for c in X_test.columns[X_test.isna().any()]]
  94. feature_names_train = X_train.columns.to_list() + missing_cols_train
  95. feature_names_test = X_test.columns.to_list() + missing_cols_test
  96. return imputer, feature_names_train, feature_names_test
  97. # %%
  98. def fit_pipeline(X, y, train_idx, test_idx, root_dir):
  99. X_train, X_test = X.iloc[train_idx], X.iloc[test_idx]
  100. y_train, y_test = y.iloc[train_idx], y.iloc[test_idx]
  101. y_train = pd.DataFrame(y_train)
  102. y_test = pd.DataFrame(y_test)
  103. # iterative imputation
  104. # define imputer
  105. imputer, feature_names_train, feature_names_test = get_iterative_imputer(X_train, X_test, random_state=0)
  106. # define pipeline
  107. pipeline = make_pipeline(imputer, RandomForestRegressor(random_state=0))
  108. # fit pipeline
  109. pipeline.fit(X_train, y_train)
  110. # predict
  111. y_pred = pd.DataFrame(pipeline.predict(X_test)).rename(columns={0: 'mmse_pred', 1: 'sob_pred'})
  112. y_pred.index = y_test.index.values
  113. # concatenate
  114. df = pd.concat([y_pred, y_test], axis = 1)
  115. # get MAE, MSE, R2 for the predictions with scikit-learn
  116. metrics = {
  117. 'r2': r2_score,
  118. 'mae': mean_absolute_error,
  119. 'mse': mean_squared_error,
  120. }
  121. results_df = pd.DataFrame({
  122. var: {name: func(df[f'{var}_slope'], df[f'{var}_pred']) for name, func in metrics.items()}
  123. for var in ['sob', 'mmse']
  124. }).T
  125. # save pipeline and dataframe
  126. if root_dir:
  127. save_pipeline_and_df(pipeline, df, results_df, root_dir)
  128. return df, pipeline, results_df
  129. def save_pipeline_and_df(pipeline, df, results_df, root_dir):
  130. root_dir = Path(root_dir)
  131. root_dir.mkdir(parents=True, exist_ok=True)
  132. joblib.dump(pipeline, f"{root_dir}/pipeline.joblib", compress=('xz', 3))
  133. df.to_csv(f"{root_dir}/predictions.csv", index=True, header=True)
  134. results_df.to_csv(f"{root_dir}/results.csv", index=True, header=True)
  135. def compress_pipeline_file(root_dir):
  136. pipeline_path = Path(root_dir) / "pipeline.joblib"
  137. if pipeline_path.exists():
  138. with open(pipeline_path, 'rb') as f:
  139. pipeline = joblib.load(f)
  140. with open(pipeline_path, 'wb') as f:
  141. joblib.dump(pipeline, f, compress=('xz', 3))
  142. def split_and_fit(X, y, root_dir, n_splits=1000, test_size=0.2):
  143. ss = ShuffleSplit(n_splits=n_splits, test_size=test_size, random_state=0)
  144. # split data
  145. for iter, (train_idx, test_idx) in tqdm(enumerate(ss.split(X, y)), total=n_splits):
  146. path = Path(f"{root_dir}/iteration_{iter+1}")
  147. path.mkdir(parents=True, exist_ok=True)
  148. # check if both pipeline and df already exist
  149. if (path / "pipeline.joblib").exists() and (path / "predictions.csv").exists():
  150. continue
  151. fit_pipeline(X, y, train_idx, test_idx, path)
  152. def get_results(root, iters=range(1000)):
  153. dfs = [pd.read_csv(f"{root}/iteration_{i+1}/results.csv", index_col=0) for i in iters]
  154. return pd.concat(dfs, keys=iters, names=["iteration", "target"])
  155. def get_results_summary(root, iters=range(1000)):
  156. results = get_results(root, iters)
  157. summary = results.groupby(level=0).agg(['mean', 'median'])
  158. summary.columns = ['_'.join(col).strip() for col in summary.columns.values]
  159. return summary
  160. # %%
  161. # ## 1. Data Loading and Within-Dataset Predictions
  162. # ### 1.1 Data Loading and Preprocessing
  163. # Load data
  164. oasis_y, oasis_X_clin, oasis_X_fs, oasis_X_clin_fs = get_data("OASIS_clinical_features_X.pkl", "OASIS_data_slopes_y.pkl", "OASIS_structGlobScort.pkl")
  165. adni_y, adni_X_clin, adni_X_fs, adni_X_clin_fs = get_data("ADNI_1237_data_clinical_features_X_clinsessionn_2025.pkl", "ADNI_1237_data_slopes_y_2025.pkl", "ADNI_1237_data_structGlobScort_2025.pkl")
  166. missing_adni = ["smoke_PACKSPER", "TMTA_TRAILALI", "smoke_TOBAC100", "smoke_SMOKYRS", "TMTB_TRAILBLI", "smoke_TOBAC30", "WMSds_DIGIF", "WMSds_DIGIFLEN", "WF_VEG", "WMSds_DIGIB", "WMSds_DIGIBLEN", "WAIS_WAIS", "familyHist_DADDEM", "familyHist_MOMDEM"]
  167. missing_oasis = ["familyHist_sumSibDem","familyHist_ratioSibDem","TMTB_TRAILBRR","TMTB_TRAILBLI","TMTA_TRAILALI","TMTA_TRAILARR"]
  168. missing_cols = missing_adni + missing_oasis
  169. oasis_y, oasis_X_clin_filtered, oasis_X_fs_filtered, oasis_X_clin_fs_filtered = get_data("OASIS_clinical_features_X.pkl", "OASIS_data_slopes_y.pkl", "OASIS_structGlobScort.pkl", missing_cols=missing_cols)
  170. adni_y, adni_X_clin_filtered, adni_X_fs_filtered, adni_X_clin_fs_filtered = get_data("ADNI_1237_data_clinical_features_X_clinsessionn_2025.pkl", "ADNI_1237_data_slopes_y_2025.pkl", "ADNI_1237_data_structGlobScort_2025.pkl", missing_cols=missing_cols)
  171. data_combinations = [
  172. (oasis_y, oasis_X_clin, "OASIS_clin_predictions"),
  173. (oasis_y, oasis_X_fs, "OASIS_struct_predictions"),
  174. (oasis_y, oasis_X_clin_fs, "OASIS_clin_struct_predictions"),
  175. (adni_y, adni_X_clin, "ADNI_clin_predictions"),
  176. (adni_y, adni_X_fs, "ADNI_struct_predictions"),
  177. (adni_y, adni_X_clin_fs, "ADNI_clin_struct_predictions"),
  178. (oasis_y, oasis_X_clin_filtered, "OASIS_clin_predictions_filtered"),
  179. (oasis_y, oasis_X_fs_filtered, "OASIS_struct_predictions_filtered"),
  180. (oasis_y, oasis_X_clin_fs_filtered, "OASIS_clin_struct_predictions_filtered"),
  181. (adni_y, adni_X_clin_filtered, "ADNI_clin_predictions_filtered"),
  182. (adni_y, adni_X_fs_filtered, "ADNI_struct_predictions_filtered"),
  183. (adni_y, adni_X_clin_fs_filtered, "ADNI_clin_struct_predictions_filtered"),
  184. ]
  185. oasis_X_clin_filtered.to_csv("OASIS_clin_filtered.csv", index=True, header=True)
  186. adni_X_clin_filtered.to_csv("ADNI_clin_filtered.csv", index=True, header=True)
  187. # ### 1.2 Within-Dataset Predictions (OASIS-3 and ADNI)
  188. for y, X, root_dir in data_combinations:
  189. print(f"Processing {root_dir}...")
  190. split_and_fit(X, y, root_dir, n_splits=1000)
  191. # %%
  192. # ### 1.3 Results Aggregation and Summary
  193. def get_dataset_results(clin_file, struct_file, clin_struct_file, iters=range(1000)):
  194. clin_results = get_results(clin_file, iters)
  195. struct_results = get_results(struct_file, iters)
  196. clin_struct_results = get_results(clin_struct_file, iters)
  197. results = pd.concat([clin_results, struct_results, clin_struct_results], keys=["clin", "struct", "clin_struct"], names=["modality", "iteration", "target"])
  198. return results
  199. oasis_results = get_dataset_results("OASIS_clin_predictions", "OASIS_struct_predictions", "OASIS_clin_struct_predictions")
  200. adni_results = get_dataset_results("ADNI_clin_predictions", "ADNI_struct_predictions", "ADNI_clin_struct_predictions")
  201. oasis_filtered_results = get_dataset_results("OASIS_clin_predictions_filtered", "OASIS_struct_predictions_filtered", "OASIS_clin_struct_predictions_filtered")
  202. adni_filtered_results = get_dataset_results("ADNI_clin_predictions_filtered", "ADNI_struct_predictions_filtered", "ADNI_clin_struct_predictions_filtered")
  203. results = pd.concat([oasis_results, adni_results, oasis_filtered_results, adni_filtered_results], keys=["OASIS", "ADNI", "OASIS_filtered", "ADNI_filtered"], names=["dataset", "modality", "iteration", "target"])
  204. results.to_csv("results_within.csv")
  205. results.groupby(level=["dataset", "modality", "target"]).agg(['mean', 'median'])
  206. # %%
  207. # ## 2. Permutation Importance Analysis
  208. # ### 2.1 Permutation Importance for Individual Features
  209. def load_pipeline(iter, root_dir):
  210. path = Path(f"{root_dir}/iteration_{iter+1}/pipeline.joblib")
  211. if path.exists():
  212. return joblib.load(path)
  213. else:
  214. raise FileNotFoundError(f"Pipeline for iteration {iter+1} not found in {root_dir}.")
  215. def get_permutation_importance_iter(pipeline, X_test, y_test, n_repeats=100, random_state=0):
  216. pi = permutation_importance(pipeline, X_test, y_test, n_repeats=n_repeats, scoring='r2', random_state=random_state)
  217. return pd.DataFrame({"feature": X_test.columns, "permutation_importance": pi.importances_mean})
  218. def get_permutation_importance(X, y, root_dir, n_splits=1000, n_repeats=100):
  219. # quickly check that every path exists
  220. root_dir = Path(root_dir)
  221. for iter in range(n_splits):
  222. path = root_dir / f"iteration_{iter+1}"
  223. if not path.exists():
  224. raise FileNotFoundError(f"Path {path} does not exist. Make sure to run split_and_fit first.")
  225. def process_iteration(iter):
  226. path = root_dir / f"iteration_{iter+1}"
  227. pi_path = path / f"permutation_importance_repeats-{n_repeats}.csv"
  228. if pi_path.exists():
  229. print(f"Skipping iteration {iter+1}, already processed.")
  230. return pd.read_csv(pi_path, index_col=0)
  231. else:
  232. ss = ShuffleSplit(n_splits=n_splits, test_size=0.2, random_state=0)
  233. train_indices = []
  234. test_indices = []
  235. for i, (train_idx, test_idx) in enumerate(ss.split(X, y)):
  236. if i == iter:
  237. train_indices = train_idx
  238. test_indices = test_idx
  239. break
  240. pipeline = load_pipeline(iter, root_dir)
  241. X_test = X.iloc[test_indices]
  242. y_test = y.iloc[test_indices]
  243. pi_df = get_permutation_importance_iter(pipeline, X_test, y_test, n_repeats=n_repeats)
  244. pi_df.to_csv(pi_path)
  245. return pi_df
  246. # Run the processing in parallel with 4 jobs
  247. results = Parallel(n_jobs=4)(
  248. delayed(process_iteration)(iter) for iter in tqdm(range(n_splits), desc="Processing iterations")
  249. )
  250. results = pd.concat(results, keys=range(n_splits), names=["iteration", "feature"])
  251. results = results.droplevel("feature")
  252. return results
  253. oasis_pi = get_permutation_importance(oasis_X_clin_fs, oasis_y,"OASIS_clin_struct_predictions", n_splits=1000, n_repeats=5)
  254. adni_pi = get_permutation_importance(adni_X_clin_fs, adni_y, "ADNI_clin_struct_predictions", n_splits=1000, n_repeats=5)
  255. oasis_filtered_pi = get_permutation_importance(oasis_X_clin_fs_filtered, oasis_y, "OASIS_clin_struct_predictions_filtered", n_splits=1000, n_repeats=5)
  256. adni_filtered_pi = get_permutation_importance(adni_X_clin_fs_filtered, adni_y, "ADNI_clin_struct_predictions_filtered", n_splits=1000, n_repeats=5)
  257. oasis_filtered_pi.set_index("feature", append=True).unstack("feature").droplevel(0, axis=1).to_csv("OASIS_clin_struct_permutation_importance.csv")
  258. adni_filtered_pi.set_index("feature", append=True).unstack("feature").droplevel(0, axis=1).to_csv("ADNI_clin_struct_permutation_importance.csv")
  259. # ### 2.2 Feature Importance Visualization and Analysis
  260. # %%
  261. def get_across_permutation_importance(X, y, root_dir, n_repeats=100, label="permutation_importance"):
  262. # quickly check that every path exists
  263. path = Path(root_dir)
  264. pi_path = path / f"{label}_repeats-{n_repeats}.csv"
  265. if pi_path.exists():
  266. print(f"Skipping {label}, already processed.")
  267. pi_df = pd.read_csv(pi_path, index_col=0)
  268. else:
  269. pipeline = joblib.load(f"{root_dir}/pipeline.joblib")
  270. pi = permutation_importance(pipeline, X, y, n_repeats=n_repeats, scoring='r2', random_state=0)
  271. pi_df = pd.DataFrame({"feature": X.columns, "permutation_importance": pi.importances_mean, "permutation_importance_std": pi.importances_std})
  272. pi_df.to_csv(pi_path)
  273. return pi_df
  274. oasis2adni_pi = get_across_permutation_importance(adni_X_clin_fs_filtered, adni_y, "OASIS_to_ADNI_clin_struct", n_repeats=100)
  275. adni2oasis_pi = get_across_permutation_importance(oasis_X_clin_fs_filtered, oasis_y, "ADNI_to_OASIS_clin_struct", n_repeats=100)
  276. # ### 2.3 Coalition-Based Feature Importance
  277. # %%
  278. feature_coalitions = pd.read_csv("features.csv", header=0)
  279. # group by coalition and feature type
  280. coalition_dict = feature_coalitions.groupby("coalition")["feature"].apply(list).to_dict()
  281. feature_type_dict = feature_coalitions.groupby("feature-type")["feature"].apply(list).to_dict()
  282. # get importance by feature coalition
  283. def importance_by_coalition(pipeline, X_test, y_test, coalition_dict, random_state=0):
  284. # manual reimplementation of permutation_importance
  285. result = {}
  286. for coalition, features_in in coalition_dict.items():
  287. # permute the rows within the coalition
  288. X_test_coalition = X_test.copy()
  289. X_test_coalition[features_in] = X_test_coalition[features_in].sample(frac=1, random_state=random_state).values
  290. pred = pipeline.predict(X_test_coalition)
  291. result[coalition] = metrics.r2_score(y_test, pred)
  292. pred = pipeline.predict(X_test)
  293. all_score = metrics.r2_score(y_test, pred)
  294. result = pd.DataFrame(result, index=['permutation_importance'])
  295. result = all_score - result.T
  296. return result
  297. def get_importance_by_coalition(X, y, root_dir, coalition_dict, n_splits=1000, n_repeats=5, label="coalition_importance"):
  298. # quickly check that every path exists
  299. root_dir = Path(root_dir)
  300. for i in range(n_splits):
  301. path = root_dir / f"iteration_{i+1}"
  302. if not path.exists():
  303. raise FileNotFoundError(f"Path {path} does not exist. Make sure to run split_and_fit first.")
  304. def process_iteration(iter):
  305. path = root_dir / f"iteration_{iter+1}"
  306. pi_path = path / f"{label}_repeats-{n_repeats}.csv"
  307. if pi_path.exists():
  308. return pd.read_csv(pi_path, index_col=0)
  309. else:
  310. ss = ShuffleSplit(n_splits=n_splits, test_size=0.2, random_state=0)
  311. test_indices = []
  312. for i, (train_idx, test_idx) in enumerate(ss.split(X, y)):
  313. if i == iter:
  314. test_indices = test_idx
  315. break
  316. pipeline = load_pipeline(iter, root_dir)
  317. X_test = X.iloc[test_indices]
  318. y_test = y.iloc[test_indices]
  319. pi_df = [importance_by_coalition(pipeline, X_test, y_test, coalition_dict, random_state=i) for i in range(n_repeats)]
  320. pi_df = pd.concat(pi_df).groupby(level=0).mean()
  321. pi_df.to_csv(pi_path)
  322. return pi_df
  323. # Run the processing in parallel with 4 jobs
  324. results = Parallel(n_jobs=4)(
  325. delayed(process_iteration)(iter) for iter in tqdm(range(n_splits), desc="Processing iterations")
  326. )
  327. results = pd.concat(results, keys=range(n_splits), names=["iteration", "coalition"])
  328. return results
  329. oasis_filtered_ci = get_importance_by_coalition(oasis_X_clin_fs_filtered, oasis_y, "OASIS_clin_struct_predictions_filtered", coalition_dict, n_splits=1000, n_repeats=5)
  330. adni_filtered_ci = get_importance_by_coalition(adni_X_clin_fs_filtered, adni_y, "ADNI_clin_struct_predictions_filtered", coalition_dict, n_splits=1000, n_repeats=5)
  331. oasis_filtered_type_ci = get_importance_by_coalition(oasis_X_clin_fs_filtered, oasis_y, "OASIS_clin_struct_predictions_filtered", feature_type_dict, n_splits=1000, n_repeats=5, label="feature_type_importance")
  332. adni_filtered_type_ci = get_importance_by_coalition(adni_X_clin_fs_filtered, adni_y, "ADNI_clin_struct_predictions_filtered", feature_type_dict, n_splits=1000, n_repeats=5, label="feature_type_importance")
  333. # %%
  334. def get_mean_ci(df, alpha=0.05, df_dof=4):
  335. mean = df.mean()
  336. ci = stats.t.interval(alpha, df_dof, loc=mean, scale=stats.sem(df))
  337. ci_lower, ci_upper = ci
  338. return pd.Series({"mean": mean, "ci_lower": ci_lower, "ci_upper": ci_upper})
  339. def get_mean_ci_nb(df, alpha=0.05, VIF=(1 + 0.2/0.8)):
  340. R = len(df)
  341. mean = df.mean()
  342. var = df.var(ddof=1) # unbiased variance
  343. inflation = 1 + VIF
  344. se = np.sqrt(var * inflation / R)
  345. dof = R - 1
  346. t_val = stats.t.ppf(1 - alpha/2, dof)
  347. ci_lower = mean - t_val * se
  348. ci_upper = mean + t_val * se
  349. return pd.Series({"mean": mean, "ci_lower": ci_lower, "ci_upper": ci_upper})
  350. def get_ci_nb(df, alpha=0.05, VIF=(1 + 0.2/0.8)):
  351. res = get_mean_ci_nb(df, alpha, VIF)
  352. ci_lower = res['ci_lower']
  353. ci_upper = res['ci_upper']
  354. return (ci_lower, ci_upper)
  355. # %%
  356. # plot these as point plots with error bars
  357. def plot_permutation_importance(pi_df, title):
  358. plt.figure(figsize=(10, 6))
  359. sns.pointplot(
  360. data=pi_df.reset_index(),
  361. x='feature',
  362. y='permutation_importance',
  363. capsize=0.1,
  364. errorbar=get_ci_nb,
  365. markers='o',
  366. color='blue',
  367. markersize=2,
  368. linestyle='none',
  369. linewidth=1
  370. )
  371. plt.xticks(rotation=90)
  372. plt.title(title)
  373. plt.xlabel('Feature')
  374. plt.ylabel('Permutation Importance (Mean ± CI)')
  375. plt.tight_layout()
  376. plt.show()
  377. top_15_features_oasis = oasis_pi.groupby("feature").mean()["permutation_importance"].sort_values(ascending=False).head(15).index.tolist()
  378. top_15_features_adni = adni_pi.groupby("feature").mean()["permutation_importance"].sort_values(ascending=False).head(15).index.tolist()
  379. top_15_features = list(set(top_15_features_oasis + top_15_features_adni))
  380. plot_permutation_importance(oasis_pi.loc[oasis_pi["feature"].isin(top_15_features)], "OASIS-3: Top 15 Features Permutation Importance")
  381. plot_permutation_importance(adni_pi.loc[adni_pi["feature"].isin(top_15_features)], "ADNI: Top 15 Features Permutation Importance")
  382. # for across datasets:
  383. # %%
  384. def plot_across_permutation_importance(pi_df, title, ax=None):
  385. pi_df = pi_df.sort_values("permutation_importance", ascending=False).head(15).copy()
  386. pi_df['feature'] = pd.Categorical(pi_df['feature'], categories=pi_df['feature'], ordered=True)
  387. pi_df["ymin"] = pi_df["permutation_importance"] - pi_df["permutation_importance_std"]
  388. pi_df["ymax"] = pi_df["permutation_importance"] + pi_df["permutation_importance_std"]
  389. p = so.Plot(data=pi_df, y='feature', x='permutation_importance', xmin='ymin', xmax='ymax')
  390. p = p.add(so.Dot())
  391. p = p.add(so.Range())
  392. p = p.label(title=title, y="Feature", x="Permutation Importance ($R^2$ decrease)")
  393. # Set font to Arial for all text elements
  394. p = p.theme({
  395. "axes.labelweight": "normal",
  396. "axes.labelsize": 10,
  397. "axes.titlesize": 12,
  398. })
  399. p.on(ax).plot()
  400. # Create a figure with two subplots arranged vertically
  401. fig, axes = plt.subplots(2, 1, figsize=(8, 12), constrained_layout=True, sharex=True)
  402. feature_dict = {
  403. 'session_n': "Number of previous visits",
  404. 'diag': "Diagnosis",
  405. 'demo_sex': "Sex",
  406. 'demo_age': "Age",
  407. 'demo_education': "Education",
  408. 'diabetes_DIABETES': "Diabetes",
  409. 'hypercho_HYPERCHO': "Hypercholesterolemia",
  410. 'cvasc_HYPERTEN': "Hypertension",
  411. 'cvasc_CBSTROKE': "Stroke",
  412. 'cvasc_CBTIA': "Transient Ischemic Attack",
  413. 'cvasc_CVHATT': "Heart Attack",
  414. 'cvasc_CVAFIB': "Atrial Fibrillation",
  415. 'cvasc_CVANGIO': "Angiography",
  416. 'cvasc_CVOTHR': "Other Cardiovascular",
  417. 'cdr_commun': "CDR (Communication)",
  418. 'cdr_homehobb': "CDR (Home and Hobbies)",
  419. 'cdr_judgment': "CDR (Judgment)",
  420. 'cdr_memory': "CDR (Memory)",
  421. 'cdr_orient': "CDR (Orientation)",
  422. 'cdr_perscare': "CDR (Personal Care)",
  423. 'cdr_sob': "CDR-SOB",
  424. 'cdr_cdrGlobal': "CDR (Global)",
  425. 'mmse_mmse': "MMSE (Total)",
  426. 'gds_gdsSum': "GDS (Sum)",
  427. 'faq_BILLS': "FAQ (Bills)",
  428. 'faq_TAXES': "FAQ (Taxes)",
  429. 'faq_SHOPPING': "FAQ (Shopping)",
  430. 'faq_GAMES': "FAQ (Games)",
  431. 'faq_STOVE': "FAQ (Stove)",
  432. 'faq_MEALPREP': "FAQ (Meal Preparation)",
  433. 'faq_EVENTS': "FAQ (Events)",
  434. 'faq_PAYATTN': "FAQ (Pay Attention)",
  435. 'faq_REMDATES': "FAQ (Remember Dates)",
  436. 'faq_TRAVEL': "FAQ (Travel)",
  437. 'faq_faqSum': "FAQ (Total)",
  438. 'npiq_npiqPresSum': "NPI-Q (Presence Sum)",
  439. 'npiq_npiqSevSum': "NPI-Q (Severity Sum)",
  440. 'apoe_e2count': "APOE ε2 Count",
  441. 'apoe_e3count': "APOE ε3 Count",
  442. 'apoe_e4count': "APOE ε4 Count",
  443. 'WMSlm_LOGIMEM': "WMS (Logical Memory)",
  444. 'WMSlm_MEMUNITS': "WMS (Memory Units)",
  445. 'WMSlm_MEMTIME': "WMS (Memory Time)",
  446. 'WF_ANIMALS': "Word Fluency (Animals)",
  447. 'TMTA_TRAILA': "Trail Making Test A",
  448. 'TMTB_TRAILB': "Trail Making Test B",
  449. 'TMTB_TRAILBnorm': "Trail Making Test B (normalized)",
  450. 'BOSTON_BOSTON': "Boston Naming Test (Total)",
  451. 'fs__globalVolume__3rd-Ventricle': "Volume of 3rd Ventricle",
  452. 'fs__globalVolume__4th-Ventricle': "Volume of 4th Ventricle",
  453. 'fs__globalVolume__CC_Anterior': "Volume of Corpus Callosum (Anterior)",
  454. 'fs__globalVolume__CC_Central': "Volume of Corpus Callosum (Central)",
  455. 'fs__globalVolume__CC_Mid_Anterior': "Volume of Corpus Callosum (Mid Anterior)",
  456. 'fs__globalVolume__CC_Mid_Posterior': "Volume of Corpus Callosum (Mid Posterior)",
  457. 'fs__globalVolume__CC_Posterior': "Volume of Corpus Callosum (Posterior)",
  458. 'fs__globalVolume__Left-Cerebellum-Cortex': "Volume of Cerebellum Cortex (Left)",
  459. 'fs__globalVolume__Left-Cerebellum-White-Matter': "Volume of Cerebellum White Matter (Left)",
  460. 'fs__globalVolume__Left-Lateral-Ventricle': "Volume of Lateral Ventricle (Left)",
  461. 'fs__globalVolume__Right-Cerebellum-Cortex': "Volume of Cerebellum Cortex (Right)",
  462. 'fs__globalVolume__Right-Cerebellum-White-Matter': "Volume of Cerebellum White Matter (Right)",
  463. 'fs__globalVolume__Right-Lateral-Ventricle': "Volume of Lateral Ventricle (Right)",
  464. 'fs__globalVolume__SubCortGrayVol': "Volume of Subcortical Gray Matter",
  465. 'fs__globalVolume__TotalGrayVol': "Volume of Total Gray Matter",
  466. 'fs__globalVolume__lhCerebralWhiteMatterVol': "Volume of Cerebral White Matter (Left)",
  467. 'fs__globalVolume__lhCortexVol': "Volume of Cortex (Left)",
  468. 'fs__globalVolume__lh_MeanThickness_thickness': "Mean Cortical Thickness (Left)",
  469. 'fs__globalVolume__rhCerebralWhiteMatterVol': "Volume of Cerebral White Matter (Right)",
  470. 'fs__globalVolume__rhCortexVol': "Volume of Cortex (Right)",
  471. 'fs__globalVolume__rh_MeanThickness_thickness': "Mean Cortical Thickness (Right)",
  472. 'fs__subcortVolume__Left-Accumbens-area': "Volume of Accumbens (Left)",
  473. 'fs__subcortVolume__Left-Amygdala': "Volume of Amygdala (Left)",
  474. 'fs__subcortVolume__Left-Caudate': "Volume of Caudate (Left)",
  475. 'fs__subcortVolume__Left-Hippocampus': "Volume of Hippocampus (Left)",
  476. 'fs__subcortVolume__Left-Pallidum': "Volume of Pallidum (Left)",
  477. 'fs__subcortVolume__Left-Putamen': "Volume of Putamen (Left)",
  478. 'fs__subcortVolume__Left-Thalamus-Proper': "Volume of Thalamus (Left)",
  479. 'fs__subcortVolume__Right-Accumbens-area': "Volume of Accumbens (Right)",
  480. 'fs__subcortVolume__Right-Amygdala': "Volume of Amygdala (Right)",
  481. 'fs__subcortVolume__Right-Caudate': "Volume of Caudate (Right)",
  482. 'fs__subcortVolume__Right-Hippocampus': "Volume of Hippocampus (Right)",
  483. 'fs__subcortVolume__Right-Pallidum': "Volume of Pallidum (Right)",
  484. 'fs__subcortVolume__Right-Putamen': "Volume of Putamen (Right)",
  485. 'fs__subcortVolume__Right-Thalamus-Proper': "Volume of Thalamus (Right)",
  486. }
  487. oasis2adni_pi['feature_names'] = oasis2adni_pi['feature'].map(feature_dict)
  488. adni2oasis_pi['feature_names'] = adni2oasis_pi['feature'].map(feature_dict)
  489. # Plot for OASIS to ADNI
  490. plot_across_permutation_importance(oasis2adni_pi.rename(columns={"feature_names": "feature", "feature": "feature_i"}), "A) OASIS-3 → ADNI: Feature Importance", ax=axes[0])
  491. # Plot for ADNI to OASIS
  492. plot_across_permutation_importance(adni2oasis_pi.rename(columns={"feature_names": "feature", "feature": "feature_i"}), "B) ADNI → OASIS-3: Feature Importance", ax=axes[1])
  493. plt.savefig("cross_dataset_feature_importance.png", dpi=300, bbox_inches="tight")
  494. top15_features_oasis2adni = oasis2adni_pi.sort_values("permutation_importance", ascending=False).head(15)["feature"].to_list()
  495. top15_features_adni2oasis = adni2oasis_pi.sort_values("permutation_importance", ascending=False).head(15)["feature"].to_list()
  496. # %%
  497. top_15_features_oasis_filtered = oasis_filtered_pi.groupby("feature").mean()["permutation_importance"].sort_values(ascending=False).head(15).index.tolist()
  498. top_15_features_adni_filtered = adni_filtered_pi.groupby("feature").mean()["permutation_importance"].sort_values(ascending=False).head(15).index.tolist()
  499. top_15_features_filtered = list(set(top_15_features_oasis_filtered + top_15_features_adni_filtered))
  500. top_15_features_intersection = list(set(top_15_features_oasis_filtered) & set(top_15_features_adni_filtered))
  501. adni_overlap = set(top15_features_adni2oasis) & set(top_15_features_adni_filtered)
  502. adni_difference1 = set(top_15_features_adni_filtered) - set(top15_features_adni2oasis)
  503. adni_difference2 = set(top15_features_adni2oasis) - set(top_15_features_adni_filtered)
  504. oasis_overlap = set(top15_features_oasis2adni) & set(top_15_features_oasis_filtered)
  505. oasis_difference1 = set(top_15_features_oasis_filtered) - set(top15_features_oasis2adni)
  506. oasis_difference2 = set(top15_features_oasis2adni) - set(top_15_features_oasis_filtered)
  507. plot_permutation_importance(oasis_filtered_pi.loc[oasis_filtered_pi["feature"].isin(top_15_features)], "OASIS-3 (Filtered): Top 15 Features Permutation Importance")
  508. plot_permutation_importance(adni_filtered_pi.loc[adni_filtered_pi["feature"].isin(top_15_features)], "ADNI (Filtered): Top 15 Features Permutation Importance")
  509. plot_permutation_importance(oasis_filtered_ci.reset_index("coalition").rename(columns={"coalition": "feature"}), "OASIS-3 (Filtered): Permutation Importance by Coalition")
  510. plot_permutation_importance(adni_filtered_ci.reset_index("coalition").rename(columns={"coalition": "feature"}), "ADNI (Filtered): Permutation Importance by Coalition")
  511. # %%
  512. def round(x, n_figs):
  513. x = float(x)
  514. power = 10 ** np.floor(np.log10(abs(x)))
  515. rounded = np.round(x / power, n_figs - 1) * power
  516. rounded = float(rounded)
  517. # Use 'f' format to preserve trailing zeros if needed
  518. digits_after_decimal = max(n_figs - int(np.floor(np.log10(abs(rounded)))) - 1, 0)
  519. format_str = f"{{:.{digits_after_decimal}f}}"
  520. result = format_str.format(rounded)
  521. return result
  522. # oasis_filtered_pi_top15 = oasis_filtered_pi.loc[oasis_filtered_pi["feature"].isin(top_15_features_oasis_filtered)].copy()
  523. oasis_filtered_pi_top15 = oasis_filtered_pi.copy()
  524. oasis_filtered_pi_top15 = oasis_filtered_pi_top15.groupby("feature")["permutation_importance"].apply(get_mean_ci_nb).unstack()
  525. oasis_filtered_pi_top15 = oasis_filtered_pi_top15.sort_values("mean", ascending=False)
  526. oasis_filtered_pi_top15.loc[top_15_features_oasis_filtered].applymap(round, n_figs=2).apply(lambda x: f"{x['mean']} ({x['ci_lower']}, {x['ci_upper']})", axis=1)
  527. oasis_filtered_pi_top15.applymap(round, n_figs=2).apply(lambda x: f"{x['mean']} ({x['ci_lower']}, {x['ci_upper']})", axis=1).to_csv("OASIS_importance.csv")
  528. adni_filtered_pi_top15 = adni_filtered_pi.copy()
  529. adni_filtered_pi_top15 = adni_filtered_pi_top15.groupby("feature")["permutation_importance"].apply(get_mean_ci_nb).unstack()
  530. adni_filtered_pi_top15 = adni_filtered_pi_top15.sort_values("mean", ascending=False)
  531. adni_filtered_pi_top15.loc[top_15_features_adni_filtered].applymap(round, n_figs=2).apply(lambda x: f"{x['mean']} ({x['ci_lower']}, {x['ci_upper']})", axis=1)
  532. adni_filtered_pi_top15.applymap(round, n_figs=2).apply(lambda x: f"{x['mean']} ({x['ci_lower']}, {x['ci_upper']})", axis=1).to_csv("ADNI_importance.csv")
  533. # what is unique to top_15_features_oasis_filtered?
  534. set(top_15_features_oasis_filtered) - set(top_15_features_adni_filtered)
  535. # what is unique to top_15_features_adni_filtered?
  536. set(top_15_features_adni_filtered) - set(top_15_features_oasis_filtered)
  537. oasis_filtered_ci_table = oasis_filtered_type_ci.groupby("coalition")["permutation_importance"].apply(get_mean_ci_nb).unstack()
  538. oasis_filtered_ci_table = oasis_filtered_ci_table.sort_values("mean", ascending=False)
  539. oasis_filtered_ci_table.applymap(round, n_figs=2).apply(lambda x: f"{x['mean']} ({x['ci_lower']}, {x['ci_upper']})", axis=1)
  540. adni_filtered_ci_table = adni_filtered_type_ci.groupby("coalition")["permutation_importance"].apply(get_mean_ci_nb).unstack()
  541. adni_filtered_ci_table = adni_filtered_ci_table.sort_values("mean", ascending=False)
  542. adni_filtered_ci_table.applymap(round, n_figs=2).apply(lambda x: f"{x['mean']} ({x['ci_lower']}, {x['ci_upper']})", axis=1)
  543. # %%
  544. # ## 3. Across-Dataset Predictions
  545. # ### 3.1 Cross-Dataset Prediction Functions and Execution
  546. def predict_across_datasets(X_train, y_train, X_test, y_test, root_dir):
  547. root_dir = Path(root_dir)
  548. root_dir.mkdir(parents=True, exist_ok=True)
  549. if root_dir / "pipeline.joblib" in root_dir.iterdir() and root_dir / "predictions.csv" in root_dir.iterdir():
  550. print(f"Skipping {root_dir}, already processed.")
  551. return
  552. X = pd.concat([X_train, X_test])
  553. y = pd.concat([y_train, y_test])
  554. train_idx = np.arange(len(X_train))
  555. test_idx = np.arange(len(X_train), len(X))
  556. # Train a model on the combined data
  557. fit_pipeline(X, y, train_idx, test_idx, root_dir)
  558. # train on oasis, test on adni
  559. across_data_combinations = [
  560. (oasis_y, oasis_X_clin_filtered, adni_y, adni_X_clin_filtered, "OASIS_to_ADNI_clin"),
  561. (oasis_y, oasis_X_fs_filtered, adni_y, adni_X_fs_filtered, "OASIS_to_ADNI_struct"),
  562. (oasis_y, oasis_X_clin_fs_filtered, adni_y, adni_X_clin_fs_filtered, "OASIS_to_ADNI_clin_struct"),
  563. (adni_y, adni_X_clin_filtered, oasis_y, oasis_X_clin_filtered, "ADNI_to_OASIS_clin"),
  564. (adni_y, adni_X_fs_filtered, oasis_y, oasis_X_fs_filtered, "ADNI_to_OASIS_struct"),
  565. (adni_y, adni_X_clin_fs_filtered, oasis_y, oasis_X_clin_fs_filtered, "ADNI_to_OASIS_clin_struct"),
  566. (oasis_y, oasis_X_clin_fs_filtered[top_15_features_oasis_filtered], adni_y, adni_X_clin_fs_filtered[top_15_features_oasis_filtered], "OASIS_to_ADNI_clin_struct_top15"),
  567. (adni_y, adni_X_clin_fs_filtered[top_15_features_adni_filtered], oasis_y, oasis_X_clin_fs_filtered[top_15_features_adni_filtered], "ADNI_to_OASIS_clin_struct_top15"),
  568. ]
  569. for y_train, X_train, y_test, X_test, root_dir in across_data_combinations:
  570. print(f"Processing {root_dir}...")
  571. predict_across_datasets(X_train, y_train, X_test, y_test, root_dir)
  572. # ### 3.2 Cross-Dataset Results Aggregation
  573. def get_across_data_results(root):
  574. results = pd.read_csv(f"{root}/results.csv", index_col=0)
  575. summary = results.groupby(level=0).agg(['mean', 'median'])
  576. return summary
  577. oasis2adni_results = [get_across_data_results(f"OASIS_to_ADNI_{c}") for c in ["clin", "struct", "clin_struct", "clin_struct_top15"]]
  578. oasis2adni_results = pd.concat(oasis2adni_results, keys=["clin", "struct", "clin_struct", "clin_struct_top15"], names=["modality", "target"])
  579. adni2oasis_results = [get_across_data_results(f"ADNI_to_OASIS_{c}") for c in ["clin", "struct", "clin_struct", "clin_struct_top15"]]
  580. adni2oasis_results = pd.concat(adni2oasis_results, keys=["clin", "struct", "clin_struct", "clin_struct_top15"], names=["modality", "target"])
  581. across_results = pd.concat([oasis2adni_results, adni2oasis_results], axis=0,
  582. keys=["OASIS_to_ADNI", "ADNI_to_OASIS"], names=["dataset", "modality", "target"])
  583. # %%
  584. # ### 3.3 Cross-Dataset Training Size Sensitivity Analysis
  585. def predict_across_datasets_with_train_size(X_train, y_train, X_test, y_test, root_dir, train_size, n_splits=1000):
  586. if train_size > len(X_train):
  587. raise ValueError(f"train_size={train_size} exceeds the available training sample ({len(X_train)}).")
  588. ss = ShuffleSplit(n_splits=n_splits, train_size=train_size, random_state=0)
  589. for iter, (train_idx, _) in tqdm(enumerate(ss.split(X_train, y_train)), total=n_splits):
  590. path = Path(f"{root_dir}/iteration_{iter+1}")
  591. path.mkdir(parents=True, exist_ok=True)
  592. if (path / "pipeline.joblib").exists() and (path / "predictions.csv").exists():
  593. continue
  594. X_train_iter = X_train.iloc[train_idx]
  595. y_train_iter = y_train.iloc[train_idx]
  596. X = pd.concat([X_train_iter, X_test])
  597. y = pd.concat([y_train_iter, y_test])
  598. train_idx_full = np.arange(len(X_train_iter))
  599. test_idx_full = np.arange(len(X_train_iter), len(X))
  600. fit_pipeline(X, y, train_idx_full, test_idx_full, path)
  601. def fit_split_multioutput_pipeline(X_train, y_train, X_test, y_test, root_dir):
  602. root_dir = Path(root_dir)
  603. root_dir.mkdir(parents=True, exist_ok=True)
  604. if (root_dir / "pipeline.joblib").exists() and (root_dir / "predictions.csv").exists():
  605. return
  606. y_train = pd.DataFrame(y_train)
  607. y_test = pd.DataFrame(y_test)
  608. imputer, _, _ = get_iterative_imputer(X_train, X_test, random_state=0)
  609. X_train_imputed = imputer.fit_transform(X_train)
  610. X_test_imputed = imputer.transform(X_test)
  611. estimator = MultiOutputRegressor(RandomForestRegressor(random_state=0, n_estimators=50))
  612. estimator.fit(X_train_imputed, y_train)
  613. y_pred = pd.DataFrame(estimator.predict(X_test_imputed)).rename(columns={0: 'mmse_pred', 1: 'sob_pred'})
  614. y_pred.index = y_test.index.values
  615. df = pd.concat([y_pred, y_test], axis=1)
  616. metrics = {
  617. 'r2': r2_score,
  618. 'mae': mean_absolute_error,
  619. 'mse': mean_squared_error,
  620. }
  621. results_df = pd.DataFrame({
  622. var: {name: func(df[f'{var}_slope'], df[f'{var}_pred']) for name, func in metrics.items()}
  623. for var in ['sob', 'mmse']
  624. }).T
  625. save_pipeline_and_df({"imputer": imputer, "estimator": estimator}, df, results_df, root_dir)
  626. def predict_across_datasets_split_multioutput(X_train, y_train, X_test, y_test, root_dir):
  627. fit_split_multioutput_pipeline(X_train, y_train, X_test, y_test, root_dir)
  628. def predict_across_datasets_with_train_size_split_multioutput(X_train, y_train, X_test, y_test, root_dir, train_size, n_splits=1000):
  629. if train_size > len(X_train):
  630. raise ValueError(f"train_size={train_size} exceeds the available training sample ({len(X_train)}).")
  631. ss = ShuffleSplit(n_splits=n_splits, train_size=train_size, random_state=0)
  632. for iter, (train_idx, _) in tqdm(enumerate(ss.split(X_train, y_train)), total=n_splits):
  633. path = Path(f"{root_dir}/iteration_{iter+1}")
  634. path.mkdir(parents=True, exist_ok=True)
  635. if (path / "pipeline.joblib").exists() and (path / "predictions.csv").exists():
  636. continue
  637. fit_split_multioutput_pipeline(X_train.iloc[train_idx], y_train.iloc[train_idx], X_test, y_test, path)
  638. between_train_sizes = [60, 100, 200, 300, 600]
  639. across_train_size_combinations = [
  640. (oasis_y, oasis_X_clin_fs_filtered, adni_y, adni_X_clin_fs_filtered, "OASIS_to_ADNI_clin_struct_train_size_new"),
  641. (adni_y, adni_X_clin_fs_filtered, oasis_y, oasis_X_clin_fs_filtered, "ADNI_to_OASIS_clin_struct_train_size_new"),
  642. ]
  643. across_train_size_split_multioutput_combinations = [
  644. (oasis_y, oasis_X_clin_fs_filtered, adni_y, adni_X_clin_fs_filtered, "OASIS_to_ADNI_clin_struct_train_size_split_multioutput_new"),
  645. (adni_y, adni_X_clin_fs_filtered, oasis_y, oasis_X_clin_fs_filtered, "ADNI_to_OASIS_clin_struct_train_size_split_multioutput_new"),
  646. ]
  647. for y_train, X_train, y_test, X_test, root_dir in across_train_size_combinations:
  648. for train_size in between_train_sizes:
  649. sized_root_dir = f"{root_dir}_{train_size}"
  650. print(f"Processing {sized_root_dir}...")
  651. predict_across_datasets_with_train_size(X_train, y_train, X_test, y_test, sized_root_dir, train_size=train_size, n_splits=10)
  652. for y_train, X_train, y_test, X_test, root_dir in across_train_size_split_multioutput_combinations:
  653. for train_size in between_train_sizes:
  654. sized_root_dir = f"{root_dir}_{train_size}"
  655. print(f"Processing {sized_root_dir}...")
  656. predict_across_datasets_with_train_size_split_multioutput(X_train, y_train, X_test, y_test, sized_root_dir, train_size=train_size, n_splits=10)
  657. def get_across_train_size_results(root_stub, train_sizes=between_train_sizes, iters=range(10)):
  658. results = [get_results(f"{root_stub}_{train_size}", iters) for train_size in train_sizes]
  659. return pd.concat(results, keys=train_sizes, names=["train_size", "iteration", "target"])
  660. def get_across_train_size_model_results(native_root_stub, split_root_stub, train_sizes=between_train_sizes, iters=range(10)):
  661. native_results = get_across_train_size_results(native_root_stub, train_sizes, iters)
  662. split_results = get_across_train_size_results(split_root_stub, train_sizes, iters)
  663. return pd.concat(
  664. [native_results, split_results],
  665. keys=["native multioutput", "split multioutput"],
  666. names=["model", "train_size", "iteration", "target"]
  667. )
  668. oasis2adni_train_size_results = get_across_train_size_model_results(
  669. "OASIS_to_ADNI_clin_struct_train_size_new",
  670. "OASIS_to_ADNI_clin_struct_train_size_split_multioutput_new"
  671. )
  672. adni2oasis_train_size_results = get_across_train_size_model_results(
  673. "ADNI_to_OASIS_clin_struct_train_size_new",
  674. "ADNI_to_OASIS_clin_struct_train_size_split_multioutput_new"
  675. )
  676. across_train_size_results = pd.concat(
  677. [oasis2adni_train_size_results, adni2oasis_train_size_results],
  678. keys=["OASIS_to_ADNI", "ADNI_to_OASIS"],
  679. names=["dataset", "model", "train_size", "iteration", "target"]
  680. )
  681. across_train_size_results.to_csv("results_between_train_sizes.csv")
  682. def get_train_size_rsq_plot(data, y, x, hue, colors, title="", xlabel="R^2", ax=None, ref_val=None, full_sample_points=None, digits=2, full_sample_label="full sample"):
  683. if ax is None:
  684. ax = plt.gca()
  685. categories = data[y].drop_duplicates().tolist()
  686. hue_order = data[hue].drop_duplicates().tolist()
  687. palette = dict(zip(hue_order, colors))
  688. offsets = np.linspace(-0.2, 0.2, len(hue_order)) if len(hue_order) > 1 else np.array([0.0])
  689. ax.axvline(0, color='black', linestyle='--', linewidth=1.2)
  690. if ref_val is not None:
  691. ax.axvline(ref_val, color='black', linestyle='-', linewidth=1.2)
  692. [ax.axvline(p, color='gray', linestyle='--', linewidth=0.3) for p in np.arange(-0.6, 0.8, 0.2)]
  693. sns.boxplot(data=data, y=y, x=x, hue=hue, palette=palette, dodge=True, ax=ax, showfliers=False, linewidth=0.5, order=categories, hue_order=hue_order)
  694. medians = data.groupby([y, hue], sort=False)[x].median()
  695. for i, category in enumerate(categories):
  696. for offset, model in zip(offsets, hue_order):
  697. key = (category, model)
  698. if key in medians.index:
  699. median = medians.loc[key]
  700. ax.text(median, i + offset, f"{median:.{digits}f}", ha='center', va='center', fontdict={'fontsize': 11, 'fontweight':'bold', 'fontname': 'Arial'})
  701. if full_sample_points is not None:
  702. full_sample_pos = len(categories)
  703. for offset, model in zip(offsets, hue_order):
  704. if model in full_sample_points:
  705. point_val = full_sample_points[model]
  706. ax.scatter(point_val, full_sample_pos + offset, color=palette[model], s=35, zorder=5)
  707. ax.text(point_val, full_sample_pos + offset, f"{point_val:.{digits}f}", ha='left', va='center', fontdict={'fontsize': 11, 'fontweight':'bold', 'fontname': 'Arial'})
  708. ax.set_yticks(range(len(categories) + 1), labels=categories + [full_sample_label])
  709. ax.set_ylim(len(categories) + 0.5, -0.5)
  710. ax.tick_params(axis='both', labelfontfamily='Arial', labelsize=12)
  711. ax.set_xlabel(xlabel, fontdict={'fontsize': 12, 'fontstyle':'italic', 'fontweight':'bold', 'fontname': 'Arial'})
  712. ax.set_ylabel('')
  713. ax.set_title(title, fontdict={'fontsize': 16, 'fontweight':'bold', 'fontname': 'Arial'})
  714. ax.legend(title='')
  715. across_train_size_results_plot = across_train_size_results.reset_index()
  716. across_train_size_results_plot["train_size"] = across_train_size_results_plot["train_size"].map(lambda train_size: f"n = {train_size}")
  717. fig, axs = plt.subplots(2, 2, figsize=(10, 6), sharex=True, sharey=False, gridspec_kw={'hspace': 0.45, 'wspace': 0.4})
  718. plot_specs = [
  719. ("OASIS_to_ADNI", "sob", "OASIS-3 → ADNI: CDR-SOB", axs[0, 0]),
  720. ("OASIS_to_ADNI", "mmse", "OASIS-3 → ADNI: MMSE", axs[0, 1]),
  721. ("ADNI_to_OASIS", "sob", "ADNI → OASIS-3: CDR-SOB", axs[1, 0]),
  722. ("ADNI_to_OASIS", "mmse", "ADNI → OASIS-3: MMSE", axs[1, 1]),
  723. ]
  724. for dataset, target, title, ax in plot_specs:
  725. plot_df = across_train_size_results_plot.query("dataset == @dataset and target == @target")
  726. full_sample_points = {
  727. "native multioutput": pd.read_csv(f"{dataset}_clin_struct/results.csv", index_col=0).loc[target, "r2"],
  728. "split multioutput": pd.read_csv(f"{dataset}_clin_struct_split_multioutput/results.csv", index_col=0).loc[target, "r2"],
  729. }
  730. ref_val = full_sample_points["native multioutput"]
  731. get_train_size_rsq_plot(plot_df, "train_size", "r2", "model", ["#E89D9D", "#7DB7E8"], title, "R\u00B2", ax=ax, ref_val=ref_val, full_sample_points=full_sample_points)
  732. # remove legend from all and add a single legend for the entire figure
  733. handles, labels = axs[0, 0].get_legend_handles_labels()
  734. fig.legend(handles, labels, title='Model', loc='upper center', ncol=2, fontsize=12, title_fontsize=12, frameon=False)
  735. # add some margin to the top of the figure to accommodate the legend
  736. fig.subplots_adjust(top=0.85)
  737. # edit all axes to have same xlim
  738. for ax in axs.flatten():
  739. ax.set_xlim(-0.05, 0.55)
  740. ax.get_legend().remove()
  741. plt.tight_layout()
  742. # plt.savefig("r2_perf_between_train_sizes.pdf", dpi=300)
  743. # %%
  744. # ## 4. Statistical Model Comparisons
  745. # ### 4.1 Statistical Testing Framework and Model Comparisons
  746. def r2_test(df, subset1, subset2=None, metric='r2', null_value=0):
  747. dataset1, modality1, target1 = subset1
  748. test_func = wilcoxon
  749. df1 = df.loc[(dataset1, modality1, slice(None), target1)]
  750. if subset2 is not None:
  751. dataset2, modality2, target2 = subset2
  752. df2 = df.loc[(dataset2, modality2, slice(None), target2)]
  753. res = test_func(df1[metric], df2[metric], alternative='two-sided')
  754. proportion = (df1[metric] > df2[metric]).mean()
  755. else:
  756. res = test_func(df1[metric] - null_value, alternative='two-sided')
  757. proportion = (df1[metric] > null_value).mean()
  758. statistic, p_value = res.statistic, res.pvalue
  759. return [statistic, p_value, proportion, metric]
  760. model_comparisons = []
  761. # Within dataset comparisons
  762. for metric in ["r2", "mse", "mae"]:
  763. for target in ["sob", "mmse"]:
  764. for modality1, modality2 in [("clin", "struct"), ("struct", "clin_struct"), ("clin_struct", "clin")]:
  765. for dataset in ["OASIS_filtered", "ADNI_filtered"]:
  766. comparison_label = f"{dataset}_{modality1}_vs_{modality2}_{target}"
  767. comparison = r2_test(results, (dataset, modality1, target), (dataset, modality2, target), metric=metric)
  768. comparison = pd.Series(comparison, index=["statistic", "p_value", "proportion", "metric"], name=comparison_label)
  769. model_comparisons.append(comparison)
  770. # Between dataset comparisons
  771. for target in ["sob", "mmse"]:
  772. for modality in ["clin", "struct", "clin_struct"]:
  773. comparison_label = f"OASIS_filtered_{modality}_vs_ADNI_filtered_{modality}_{target}"
  774. comparison = r2_test(results, ("OASIS_filtered", modality, target), ("ADNI_filtered", modality, target), metric=metric)
  775. comparison = pd.Series(comparison, index=["statistic", "p_value", "proportion", "metric"], name=comparison_label)
  776. model_comparisons.append(comparison)
  777. # Within vs across comparisons
  778. for target in ["sob", "mmse"]:
  779. for modality in ["clin", "struct", "clin_struct"]:
  780. for dataset_across in ["OASIS_to_ADNI", "ADNI_to_OASIS"]:
  781. for dataset_within in ["OASIS_filtered", "ADNI_filtered"]:
  782. comparison_label = f"{dataset_across}_{modality}_vs_{dataset_within}_{modality}_{target}"
  783. null = across_results.loc[dataset_across, modality, target].loc["r2", "mean"]
  784. comparison = r2_test(results, (dataset_within, modality, target), None, null_value=null, metric=metric)
  785. comparison = pd.Series(comparison, index=["statistic", "p_value", "proportion", "metric"], name=comparison_label)
  786. model_comparisons.append(comparison)
  787. model_comparisons = pd.DataFrame(model_comparisons).rename_axis("comparison")
  788. model_comparisons.round(3).loc["OASIS_filtered_clin_vs_struct_sob"]
  789. model_comparisons.round(3).loc["OASIS_filtered_clin_vs_struct_mmse"]
  790. model_comparisons.loc["OASIS_filtered_clin_struct_vs_clin_sob"]
  791. model_comparisons.loc["OASIS_filtered_clin_struct_vs_clin_mmse"]
  792. model_comparisons.loc["OASIS_filtered_struct_vs_clin_struct_sob"]
  793. model_comparisons.loc["OASIS_filtered_struct_vs_clin_struct_mmse"]
  794. model_comparisons.loc["OASIS_filtered_clin_vs_struct_sob"]
  795. model_comparisons.loc["OASIS_filtered_clin_vs_struct_mmse"]
  796. model_comparisons.loc["ADNI_filtered_clin_struct_vs_clin_sob"]
  797. model_comparisons.loc["ADNI_filtered_clin_struct_vs_clin_mmse"]
  798. model_comparisons.loc["ADNI_filtered_struct_vs_clin_struct_sob"]
  799. model_comparisons.loc["ADNI_filtered_struct_vs_clin_struct_mmse"]
  800. model_comparisons.loc["ADNI_filtered_clin_vs_struct_sob"]
  801. model_comparisons.loc["ADNI_filtered_clin_vs_struct_mmse"]
  802. model_comparisons.loc["OASIS_filtered_clin_struct_vs_ADNI_filtered_clin_struct_sob"]
  803. model_comparisons.loc["OASIS_filtered_clin_struct_vs_ADNI_filtered_clin_struct_mmse"]
  804. model_comparisons.loc["OASIS_filtered_clin_vs_ADNI_filtered_clin_sob"]
  805. model_comparisons.loc["OASIS_filtered_clin_vs_ADNI_filtered_clin_mmse"]
  806. model_comparisons.loc["OASIS_filtered_struct_vs_ADNI_filtered_struct_sob"]
  807. model_comparisons.loc["OASIS_filtered_struct_vs_ADNI_filtered_struct_mmse"]
  808. model_comparisons.loc[
  809. ["OASIS_to_ADNI_clin_struct_vs_OASIS_filtered_clin_struct_sob",
  810. "OASIS_to_ADNI_clin_struct_vs_OASIS_filtered_clin_struct_mmse",
  811. "OASIS_to_ADNI_clin_vs_OASIS_filtered_clin_sob",
  812. "OASIS_to_ADNI_clin_vs_OASIS_filtered_clin_mmse",
  813. "OASIS_to_ADNI_struct_vs_OASIS_filtered_struct_sob",
  814. "OASIS_to_ADNI_struct_vs_OASIS_filtered_struct_mmse",
  815. "ADNI_to_OASIS_clin_struct_vs_OASIS_filtered_clin_struct_sob",
  816. "ADNI_to_OASIS_clin_struct_vs_OASIS_filtered_clin_struct_mmse",
  817. "ADNI_to_OASIS_clin_vs_OASIS_filtered_clin_sob",
  818. "ADNI_to_OASIS_clin_vs_OASIS_filtered_clin_mmse",
  819. "ADNI_to_OASIS_struct_vs_OASIS_filtered_struct_sob",
  820. "ADNI_to_OASIS_struct_vs_OASIS_filtered_struct_mmse"]
  821. ].query("metric == 'r2'")
  822. model_comparisons.loc[
  823. ["OASIS_to_ADNI_clin_struct_vs_ADNI_filtered_clin_struct_sob",
  824. "OASIS_to_ADNI_clin_struct_vs_ADNI_filtered_clin_struct_mmse",
  825. "OASIS_to_ADNI_clin_vs_ADNI_filtered_clin_sob",
  826. "OASIS_to_ADNI_clin_vs_ADNI_filtered_clin_mmse",
  827. "OASIS_to_ADNI_struct_vs_ADNI_filtered_struct_sob",
  828. "OASIS_to_ADNI_struct_vs_ADNI_filtered_struct_mmse",
  829. "ADNI_to_OASIS_clin_struct_vs_ADNI_filtered_clin_struct_sob",
  830. "ADNI_to_OASIS_clin_struct_vs_ADNI_filtered_clin_struct_mmse",
  831. "ADNI_to_OASIS_clin_vs_ADNI_filtered_clin_sob",
  832. "ADNI_to_OASIS_clin_vs_ADNI_filtered_clin_mmse",
  833. "ADNI_to_OASIS_struct_vs_ADNI_filtered_struct_sob",
  834. "ADNI_to_OASIS_struct_vs_ADNI_filtered_struct_mmse"]
  835. ].query("metric == 'r2'")
  836. target="mmse"
  837. (results
  838. .groupby(level=["dataset", "modality", "target"])
  839. .agg(['mean', 'median'])
  840. .xs("median", level=1, axis=1)
  841. .xs(target, level="target", axis=0)
  842. .loc[["OASIS_filtered", "ADNI_filtered"], ["clin", "struct", "clin_struct"], :]
  843. .round(2)
  844. [["r2", "mse", "mae"]])
  845. across_results.round(2).xs(target, level="target", axis=0).xs("median", level=1, axis=1)[["r2", "mse", "mae"]]
  846. r2_comparisons = model_comparisons.set_index("metric", append=True)
  847. r2_comparisons = r2_comparisons.xs("r2", level="metric", axis=0)
  848. ##
  849. r2_comparisons.loc[f"OASIS_filtered_clin_vs_struct_{target}"]
  850. r2_comparisons.loc[f"OASIS_filtered_clin_struct_vs_clin_{target}"]
  851. r2_comparisons.loc[f"OASIS_filtered_struct_vs_clin_struct_{target}"]
  852. ##
  853. r2_comparisons.loc[f"ADNI_filtered_clin_vs_struct_{target}"]
  854. r2_comparisons.loc[f"ADNI_filtered_clin_struct_vs_clin_{target}"]
  855. r2_comparisons.loc[f"ADNI_filtered_struct_vs_clin_struct_{target}"]
  856. ##
  857. r2_comparisons.loc[f"OASIS_filtered_clin_vs_ADNI_filtered_clin_{target}"]
  858. r2_comparisons.loc[f"OASIS_filtered_struct_vs_ADNI_filtered_struct_{target}"]
  859. r2_comparisons.loc[f"OASIS_filtered_clin_struct_vs_ADNI_filtered_clin_struct_{target}"]
  860. ##
  861. r2_comparisons.loc[f"OASIS_to_ADNI_clin_vs_OASIS_filtered_clin_{target}"]
  862. r2_comparisons.loc[f"OASIS_to_ADNI_struct_vs_OASIS_filtered_struct_{target}"]
  863. r2_comparisons.loc[f"OASIS_to_ADNI_clin_struct_vs_OASIS_filtered_clin_struct_{target}"]
  864. ##
  865. r2_comparisons.loc[f"OASIS_to_ADNI_clin_vs_ADNI_filtered_clin_{target}"]
  866. r2_comparisons.loc[f"OASIS_to_ADNI_struct_vs_ADNI_filtered_struct_{target}"]
  867. r2_comparisons.loc[f"OASIS_to_ADNI_clin_struct_vs_ADNI_filtered_clin_struct_{target}"]
  868. ##
  869. r2_comparisons.loc[f"ADNI_to_OASIS_clin_vs_OASIS_filtered_clin_{target}"]
  870. r2_comparisons.loc[f"ADNI_to_OASIS_struct_vs_OASIS_filtered_struct_{target}"]
  871. r2_comparisons.loc[f"ADNI_to_OASIS_clin_struct_vs_OASIS_filtered_clin_struct_{target}"]
  872. ##
  873. r2_comparisons.loc[f"ADNI_to_OASIS_clin_vs_ADNI_filtered_clin_{target}"]
  874. r2_comparisons.loc[f"ADNI_to_OASIS_struct_vs_ADNI_filtered_struct_{target}"]
  875. r2_comparisons.loc[f"ADNI_to_OASIS_clin_struct_vs_ADNI_filtered_clin_struct_{target}"]
  876. ##
  877. # %%
  878. # ### 4.2 Absolute Error Comparisons Between Models
  879. def get_abserr(root_dir):
  880. abserr = pd.read_csv(f"{root_dir}/predictions.csv", index_col=0)
  881. abserr['sob_abs_error'] = (abserr['sob_slope'] - abserr['sob_pred']).abs()
  882. abserr['mmse_abs_error'] = (abserr['mmse_slope'] - abserr['mmse_pred']).abs()
  883. return abserr
  884. def get_abserr_comparison(root_dir1, root_dir2, paired=False):
  885. abserr1 = get_abserr(root_dir1)
  886. abserr2 = get_abserr(root_dir2)
  887. if paired:
  888. test_fun = wilcoxon
  889. # check that abserr1 and abserr2 have the same index
  890. if not abserr1.index.equals(abserr2.index):
  891. raise ValueError("The indices of abserr1 and abserr2 must match for paired tests.")
  892. mmse_proportion = (abserr1['mmse_abs_error'] > abserr2['mmse_abs_error']).mean()
  893. sob_proportion = (abserr1['sob_abs_error'] > abserr2['sob_abs_error']).mean()
  894. else:
  895. test_fun = mannwhitneyu
  896. mmse_proportion = np.nan
  897. sob_proportion = np.nan
  898. mmse_test = test_fun(abserr1['mmse_abs_error'], abserr2['mmse_abs_error'])
  899. sob_test = test_fun(abserr1['sob_abs_error'], abserr2['sob_abs_error'])
  900. res = {
  901. 'mmse_statistic': mmse_test.statistic,
  902. 'mmse_pvalue': mmse_test.pvalue,
  903. 'mmse_proportion': mmse_proportion,
  904. 'sob_statistic': sob_test.statistic,
  905. 'sob_pvalue': sob_test.pvalue,
  906. 'sob_proportion': sob_proportion,
  907. 'paired': paired,
  908. f"{root_dir1}_mmse_mean": abserr1['mmse_abs_error'].mean(),
  909. f"{root_dir2}_mmse_mean": abserr2['mmse_abs_error'].mean(),
  910. f"{root_dir1}_sob_mean": abserr1['sob_abs_error'].mean(),
  911. f"{root_dir2}_sob_mean": abserr2['sob_abs_error'].mean()
  912. }
  913. res = pd.Series(res)
  914. return res
  915. # ### 6.6 OASIS-3 → ADNI: Combined vs. Top 15 (Absolute errors)
  916. get_abserr_comparison("OASIS_to_ADNI_clin_struct", "OASIS_to_ADNI_clin_struct_top15", paired=True)
  917. # ### 6.7 ADNI → OASIS-3: Combined vs. Top 15
  918. get_abserr_comparison("ADNI_to_OASIS_clin_struct", "ADNI_to_OASIS_clin_struct_top15", paired=True)
  919. # ### 6.8 OASIS-3 → ADNI vs. ADNI → OASIS-3: Combined
  920. get_abserr_comparison("OASIS_to_ADNI_clin_struct", "ADNI_to_OASIS_clin_struct", paired=False)
  921. # ### 6.9 OASIS-3 → ADNI vs. ADNI → OASIS-3: Top 15
  922. get_abserr_comparison("OASIS_to_ADNI_clin_struct_top15", "ADNI_to_OASIS_clin_struct_top15", paired=False)
  923. across_results.loc[:, "clin_struct", :]
  924. # %%
  925. # ## 5. Subgroup Analysis
  926. # ### 5.1 Preparation of Data for Subgroup Comparisons
  927. subset_adni = pd.concat([adni_X_clin_filtered, get_abserr("OASIS_to_ADNI_clin_struct")], axis=1)
  928. subset_oasis = pd.concat([oasis_X_clin_fs_filtered, get_abserr("ADNI_to_OASIS_clin_struct")], axis=1)
  929. subset_adni_top15 = pd.concat([adni_X_clin_filtered, get_abserr("OASIS_to_ADNI_clin_struct_top15")], axis=1)
  930. subset_oasis_top15 = pd.concat([oasis_X_clin_fs_filtered, get_abserr("ADNI_to_OASIS_clin_struct_top15")], axis=1)
  931. # ### 5.2 Subgroup Comparison Functions and Analysis
  932. # accessor function to make subset comparisons
  933. def get_subset(variable, value, df, comparison_type='equal'):
  934. if comparison_type == 'equal':
  935. return df[df[variable] == value]
  936. elif comparison_type == 'greater':
  937. return df[df[variable] > value]
  938. elif comparison_type == 'less':
  939. return df[df[variable] < value]
  940. elif comparison_type == 'different':
  941. return df[df[variable] != value]
  942. elif comparison_type == 'greater_equal':
  943. return df[df[variable] >= value]
  944. else:
  945. raise ValueError("Invalid comparison type.")
  946. def get_subset_comparison(df, variable, comparison_1, comparison_2, label=None):
  947. (value1, comparison_type1) = comparison_1
  948. (value2, comparison_type2) = comparison_2
  949. subset1 = get_subset(variable, value1, df, comparison_type1)
  950. subset2 = get_subset(variable, value2, df, comparison_type2)
  951. sob_test = mannwhitneyu(subset1['sob_abs_error'], subset2['sob_abs_error'])
  952. mmse_test = mannwhitneyu(subset1['mmse_abs_error'], subset2['mmse_abs_error'])
  953. sob_mae1 = subset1['sob_abs_error'].mean()
  954. mmse_mae1 = subset1['mmse_abs_error'].mean()
  955. sob_mae2 = subset2['sob_abs_error'].mean()
  956. mmse_mae2 = subset2['mmse_abs_error'].mean()
  957. return pd.Series({
  958. 'sob_statistic': sob_test.statistic,
  959. 'sob_pvalue': sob_test.pvalue,
  960. 'mmse_statistic': mmse_test.statistic,
  961. 'mmse_pvalue': mmse_test.pvalue,
  962. 'sob_mae1': sob_mae1,
  963. 'mmse_mae1': mmse_mae1,
  964. 'sob_mae2': sob_mae2,
  965. 'mmse_mae2': mmse_mae2
  966. }, name=label)
  967. def subgroup_analysis(subset):
  968. res = []
  969. # ### 8.1 Subgroups Segregated by Diagnosis
  970. # ##### HC (0) vs. MCI (1) and AD (2)
  971. res.append(get_subset_comparison(subset, 'diag', (0, 'equal'), (0, 'different'), label='HC vs. MCI and AD'))
  972. # #### MCI vs. HC and AD
  973. res.append(get_subset_comparison(subset, 'diag', (1, 'equal'), (1, 'different'), label='MCI vs. HC and AD'))
  974. # #### AD vs. HC and MCI
  975. res.append(get_subset_comparison(subset, 'diag', (2, 'equal'), (2, 'different'), label='AD vs. HC and MCI'))
  976. # #### HC vs. MCI
  977. res.append(get_subset_comparison(subset, 'diag', (0, 'equal'), (1, 'equal'), label='HC vs. MCI'))
  978. # #### HC vs. AD
  979. res.append(get_subset_comparison(subset, 'diag', (0, 'equal'), (2, 'equal'), label='HC vs. AD'))
  980. # #### MCI vs. AD
  981. res.append(get_subset_comparison(subset, 'diag', (1, 'equal'), (2, 'equal'), label='MCI vs. AD'))
  982. # ### 8.2 Subgroups Segregated by Sex
  983. # #### Female (0) vs. Male (1)
  984. res.append(get_subset_comparison(subset, 'demo_sex', (0, 'equal'), (1, 'equal'), label='Female vs. Male'))
  985. # ### 8.3 Subgroups Segregated by APOE E4
  986. # #### APOE E4: 0 vs. 1 and 2
  987. res.append(get_subset_comparison(subset, 'apoe_e4count', (0, 'equal'), (0, 'greater'), label='APOE E4: 0 vs. 1 and 2'))
  988. # #### APOE E4: 1 vs. 0 and 2
  989. res.append(get_subset_comparison(subset, 'apoe_e4count', (1, 'equal'), (1, 'different'), label='APOE E4: 1 vs. 0 and 2'))
  990. # #### APOE E4: 2 vs. 0 and 1
  991. res.append(get_subset_comparison(subset, 'apoe_e4count', (2, 'equal'), (2, 'less'), label='APOE E4: 2 vs. 0 and 1'))
  992. # #### APOE E4: 0 vs. 1
  993. res.append(get_subset_comparison(subset, 'apoe_e4count', (0, 'equal'), (1, 'equal'), label='APOE E4: 0 vs. 1'))
  994. # #### APOE E4: 0 vs. 2
  995. res.append(get_subset_comparison(subset, 'apoe_e4count', (0, 'equal'), (2, 'equal'), label='APOE E4: 0 vs. 2'))
  996. # #### APOE E4: 1 vs. 2
  997. res.append(get_subset_comparison(subset, 'apoe_e4count', (1, 'equal'), (2, 'equal'), label='APOE E4: 1 vs. 2'))
  998. # ### 8.4 Subgroups Segregated by Age
  999. subset['demo_age_group'] = pd.cut(subset['demo_age'], bins=[-np.inf, 65, 70, 75, 80, np.inf], labels=['<65', '65-70', '70-75', '75-80', '>80'])
  1000. # #### AGE: <65 vs. >=65 (65-70 and 70-75 and 75-80 and and >80)
  1001. res.append(get_subset_comparison(subset, 'demo_age_group', ('65-70', 'less'), ('65-70', 'greater_equal'), label='Age: <65 vs. >=65'))
  1002. # #### AGE: 70-75 vs. <70 and >= 75
  1003. res.append(get_subset_comparison(subset, 'demo_age_group', ('70-75', 'less'), ('70-75', 'greater_equal'), label='Age: <70 vs. >= 70'))
  1004. # #### AGE: 75-80 vs. <75 and >= 80
  1005. res.append(get_subset_comparison(subset, 'demo_age_group', ('75-80', 'less'), ('75-80', 'greater_equal'), label='Age: <75 vs. >= 80'))
  1006. # #### AGE: >= 80 vs. <80
  1007. res.append(get_subset_comparison(subset, 'demo_age_group', ('>80', 'less'), ('>80', 'greater_equal'), label='Age: <80 vs. >= 80'))
  1008. res = pd.DataFrame(res)
  1009. return res
  1010. res_subset_adni = subgroup_analysis(subset_adni)[["sob_pvalue", "sob_mae1", "sob_mae2", "mmse_pvalue", "mmse_mae1", "mmse_mae2"]]
  1011. res_subset_oasis = subgroup_analysis(subset_oasis)[["sob_pvalue", "sob_mae1", "sob_mae2", "mmse_pvalue", "mmse_mae1", "mmse_mae2"]]
  1012. # ### 5.3 Results Formatting and Export
  1013. def apa_format(column):
  1014. """
  1015. Format a dataframe column in APA style.
  1016. """
  1017. # first check if the column name contains 'pvalue'
  1018. if 'pvalue' in column.name:
  1019. # format p-values
  1020. formatted = column.apply(lambda x: f"{x:.3f}" if x > 0.001 else "< .001")
  1021. else:
  1022. formatted = column.apply(lambda x: round(x, 2) if isinstance(x, float) else str(x))
  1023. return formatted
  1024. # correct for multiple comparisons using Holm-Bonferroni method
  1025. from statsmodels.stats.multitest import multipletests
  1026. res_subset_adni["sob_pvalue_corrected"] = multipletests(res_subset_adni["sob_pvalue"], method='holm')[1]
  1027. res_subset_adni["mmse_pvalue_corrected"] = multipletests(res_subset_adni["mmse_pvalue"], method='holm')[1]
  1028. res_subset_oasis["sob_pvalue_corrected"] = multipletests(res_subset_oasis["sob_pvalue"], method='holm')[1]
  1029. res_subset_oasis["mmse_pvalue_corrected"] = multipletests(res_subset_oasis["mmse_pvalue"], method='holm')[1]
  1030. # Format the results in APA style
  1031. res_subset_adni.apply(apa_format).to_csv("ADNI_subgroup_analysis.csv")
  1032. res_subset_oasis.apply(apa_format).to_csv("OASIS_subgroup_analysis.csv")
  1033. res_subset_adni.apply(apa_format)[["sob_pvalue_corrected", "mmse_pvalue_corrected"]]
  1034. res_subset_oasis.apply(apa_format)[["sob_pvalue_corrected", "mmse_pvalue_corrected"]]
  1035. # %%
  1036. # ## 6. Post-Hoc Matched Samples Analysis
  1037. # ### 6.1 Quantile-Based Sample Matching
  1038. def get_matched_samples(source_y, target_y, quantiles=np.arange(0, 1.1, 0.1)):
  1039. """
  1040. Get matched samples based on quantiles of the source dataset.
  1041. """
  1042. quant_source_sob = source_y["sob_slope"].quantile(quantiles)
  1043. quant_source_mmse = source_y["mmse_slope"].quantile(quantiles)
  1044. target_sob_cut = pd.cut(target_y["sob_slope"], quant_source_sob, include_lowest=True, duplicates="drop")
  1045. target_mmse_cut = pd.cut(target_y["mmse_slope"], quant_source_mmse, include_lowest=True, duplicates="drop")
  1046. # number of people in target with each of the categories so that it matches the proportion in source
  1047. source_sob_n = pd.cut(source_y["sob_slope"], quant_source_sob, include_lowest=True, duplicates="drop").value_counts()
  1048. source_mmse_n = pd.cut(source_y["mmse_slope"], quant_source_mmse, include_lowest=True, duplicates="drop").value_counts()
  1049. return target_sob_cut, target_mmse_cut, source_sob_n, source_mmse_n
  1050. # trained on OASIS-3, tested on ADNI
  1051. ADNI_sob_cut, ADNI_mmse_cut, OASIS_sob_n, OASIS_mmse_n = get_matched_samples(oasis_y, adni_y)
  1052. # trained on ADNI, tested on OASIS-3
  1053. OASIS_sob_cut, OASIS_mmse_cut, ADNI_sob_n, ADNI_mmse_n = get_matched_samples(adni_y, oasis_y)
  1054. subset_adni = pd.concat([subset_adni, ADNI_sob_cut.rename("sob_qcat"), ADNI_mmse_cut.rename("mmse_qcat")], axis=1)
  1055. subset_oasis = pd.concat([subset_oasis, OASIS_sob_cut.rename("sob_qcat"), OASIS_mmse_cut.rename("mmse_qcat")], axis=1)
  1056. subset_adni_top15 = pd.concat([subset_adni_top15, ADNI_sob_cut.rename("sob_qcat"), ADNI_mmse_cut.rename("mmse_qcat")], axis=1)
  1057. subset_oasis_top15 = pd.concat([subset_oasis_top15, OASIS_sob_cut.rename("sob_qcat"), OASIS_mmse_cut.rename("mmse_qcat")], axis=1)
  1058. # ### 6.2 Matched Sample Predictions and Evaluation
  1059. # ### 10.3 OASIS-3 → ADNI (Combined)
  1060. # %%
  1061. feature_names_train = oasis_X_clin_fs_filtered.columns.tolist()
  1062. # %%
  1063. # %%
  1064. def subset_results(subset, SOB_n, MMSE_n):
  1065. qcat_results = []
  1066. for _ in tqdm(range(0, 1000)):
  1067. df_sob = subset.groupby("sob_qcat", as_index=False, group_keys=False, observed=True).apply(lambda x: x.sample(n = SOB_n[x.name], replace=True), include_groups=False)
  1068. df_mmse = subset.groupby("mmse_qcat", as_index=False, group_keys=False, observed=True).apply(lambda x: x.sample(n = MMSE_n[x.name], replace=True), include_groups=False)
  1069. sob_test = df_sob['sob_slope']
  1070. mmse_test = df_mmse['mmse_slope']
  1071. sob_pred = df_sob['sob_pred']
  1072. mmse_pred = df_mmse['mmse_pred']
  1073. qcat_results.append({
  1074. 'sob_r2': sklearn.metrics.r2_score(sob_test, sob_pred),
  1075. 'sob_mse': sklearn.metrics.mean_squared_error(sob_test, sob_pred),
  1076. 'sob_mae': sklearn.metrics.mean_absolute_error(sob_test, sob_pred),
  1077. 'mmse_r2': sklearn.metrics.r2_score(mmse_test, mmse_pred),
  1078. 'mmse_mse': sklearn.metrics.mean_squared_error(mmse_test, mmse_pred),
  1079. 'mmse_mae': sklearn.metrics.mean_absolute_error(mmse_test, mmse_pred),
  1080. })
  1081. qcat_results = pd.DataFrame(qcat_results)
  1082. return qcat_results
  1083. # ### 6.3 Results for Both Combined and Top-15 Models
  1084. subset_adni_results = subset_results(subset_adni, OASIS_sob_n, OASIS_mmse_n)
  1085. subset_adni_top15_results = subset_results(subset_adni_top15, OASIS_sob_n, OASIS_mmse_n)
  1086. subset_oasis_results = subset_results(subset_oasis, ADNI_sob_n, ADNI_mmse_n)
  1087. subset_oasis_top15_results = subset_results(subset_oasis_top15, ADNI_sob_n, ADNI_mmse_n)
  1088. subset_adni_results.agg(["median"]).map(round, n_figs=2)
  1089. subset_adni_top15_results.agg(["median"]).map(round, n_figs=2)
  1090. subset_oasis_results.agg(["median"]).map(round, n_figs=2)
  1091. subset_oasis_top15_results.agg(["median"]).map(round, n_figs=2)
  1092. # %%

01_Analysis.py at commit 44acd1b, under MIT · at the source

Overview

Authors: Roya Melanie Hüppi1,2,3, Nicolas Langer1,2, Bruno Hebling Vieira1,2, for the Alzheimer’s Disease Neuroimaging Initiative
  1. Methods of Plasticity Research, Department of Psychology, University of Zurich, 8050, Zurich, Zurich, Switzerland
  2. Neuroscience Center Zurich (ZNZ), University of Zurich & ETH Zurich, 8057, Zurich, Zurich, Switzerland
  3. Department of Adult Psychiatry and Psychotherapy, Psychiatric University Clinic Zurich and University of Zurich, 8032, Zurich, Zurich, Switzerland
Institutions: University of Zurich (Switzerland); ETH Zurich (Switzerland); Psychiatrische Universitätsklinik Zürich (Switzerland)
Journal: The journal of prevention of Alzheimer's disease, volume 13, issue 9, article 100646
Dates: received 14 April 2026; accepted 22 June 2026; published online 8 August 2026; in print August 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1016/j.tjpad.2026.100646 · PMID 42570468 · PMCID PMC13476611 · OpenAlex W7201989932
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism), Alzheimer's / dementia (population)
Methods: Statistics, Machine learning, Preprocessing
Keywords: Cognitive decline, Structural MRI, Generalizability, Machine learning, Predictive modeling
Topic: Neural Networks and Applications (Artificial Intelligence, Computer Science), according to OpenAlex
Funding: Swiss National Science Foundation (10001C_197480); UZH Postdoc (FK-23-086)
Citations: not cited yet (Europe PMC); 66 references in the paper

Abstract

Background: Predicting cognitive decline as a continuum, from healthy age-related decline to mild cognitive impairment and dementia, enables more precise individual-level predictions. However, the practical value of such models for early intervention and prevention depends on their ability to generalize to independent cohorts, a property that is often not evaluated.

Objectives: This study investigated whether adding structural magnetic resonance imaging (MRI) to non-brain data improved machine learning predictions of continuous cognitive decline and analyzed the models’ generalizability.

Design: Multi-target random forest regression models predicted annual decline in the Clinical Dementia Rating Scale Sum of Boxes (CDR-SOB) and Mini-Mental State Examination (MMSE) using non-brain data, structural MRI data, or their combination from the Alzheimer's Disease Neuroimaging Initiative (ADNI; N = 1237) and Open Access Series of Imaging Studies (OASIS-3; N = 662) datasets. Cross-site generalizability was evaluated.

Setting: Data from ADNI and OASIS-3 were used for this study.

Participants: A total of 1899 participants who had demographic, clinical, and brain imaging data from a baseline session and clinical data from at least 2 follow-up sessions were included.

Measurements: Baseline non-brain (demographics, clinical and neuropsychological scores, information on APOE genotype, cognitive diagnosis, health, and number of sessions before baseline) and/or structural MRI data were used to predict the yearly rate of change in CDR-SOB and MMSE scores.

Results: Including structural MRI data improved prediction of CDR-SOB and MMSE change, reaching respective R2 values of .41 and .33 in ADNI and .42 and .33 in OASIS-3. Model performance for across-dataset predictions was reduced (R2 between .18 and .35), unexplained by distributional shifts of target variables. Models using only top predictive features performed similarly to full models when tested externally (R2 between .18 and .34), suggesting predictor redundancy.

Conclusions: Incorporating structural MRI data enhances within-dataset prediction of continuous cognitive decline, allowing for more precise individual-level prediction and advancing towards precision medicine. Even though external validation remains limited, quantifying the generalizability gap is a crucial step towards the responsible use of ML models in clinical intervention and prevention.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repository

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

rhuepp/gemacode

License: MIT
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 44acd1bbbf67d866b6f432c8b701b6de391fc6a8, 21 September 2026
Languages: Python (2)
Size: 4 files, 2 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file
Not found: CITATION.cff, environment file, tests, continuous integration, documentation
Tools: Matplotlib (2 files), NumPy (2 files), pandas (2 files), scikit-learn (2 files), seaborn (2 files), SciPy (1 file), statsmodels (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
4 files

Code availability

The code is available on github.com/rhuepp/gemacode (https://github.com/rhuepp/gemacode.git).

Reproduced under the paper's license (CC BY), from the paper cited above.

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;
  • 2 scripts, each with its path and the digest of its content;
  • 4 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.

Data availability statement

ADNI data can be obtained from https://adni.loni.usc.edu/, OASIS-3 data from https://www.oasis-brains.org.

Reproduced under the paper's license (CC BY), from the paper cited above.

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 3, 28 September 2026

  • Authors: added Nicolas Langer (0000-0002-6038-9471); removed Nicolas Langer

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 4 authors, 5 keywords, 2 funders, 53 references.

Cite

This paper

Hüppi, R. M., Langer, N., Hebling Vieira, B., & for the Alzheimer’s Disease Neuroimaging Initiative. (2026). Quantifying generalization error in machine learning prediction of cognitive decline. The journal of prevention of Alzheimer's disease, 13(9), 100646. https://doi.org/10.1016/j.tjpad.2026.100646

BibTeX

@article{huppi2026quantifying,
author = {Hüppi, Roya Melanie and Langer, Nicolas and Hebling Vieira, Bruno and {for the Alzheimer’s Disease Neuroimaging Initiative}},
title = {{Quantifying generalization error in machine learning prediction of cognitive decline}},
journal = {The journal of prevention of Alzheimer's disease},
year = {2026},
month = aug,
volume = {13},
number = {9},
pages = {100646},
publisher = {Elsevier},
issn = {2274-5807},
doi = {10.1016/j.tjpad.2026.100646},
url = {https://doi.org/10.1016/j.tjpad.2026.100646},
pmid = {42570468},
pmcid = {PMC13476611}
}

RIS

TY - JOUR
AU - Hüppi, Roya Melanie
AU - Langer, Nicolas
AU - Hebling Vieira, Bruno
AU - for the Alzheimer’s Disease Neuroimaging Initiative
TI - Quantifying generalization error in machine learning prediction of cognitive decline
T2 - The journal of prevention of Alzheimer's disease
J2 - J Prev Alzheimers Dis
PY - 2026
DA - 2026/08/08
VL - 13
IS - 9
SP - 100646
SN - 2274-5807
PB - Elsevier
DO - 10.1016/j.tjpad.2026.100646
UR - https://doi.org/10.1016/j.tjpad.2026.100646
LA - en
ER -

CSL-JSON

{
"id": "10.1016/j.tjpad.2026.100646",
"type": "article-journal",
"title": "Quantifying generalization error in machine learning prediction of cognitive decline",
"container-title": "The journal of prevention of Alzheimer's disease",
"author": [
{
"family": "Hüppi",
"given": "Roya Melanie"
},
{
"family": "Langer",
"given": "Nicolas"
},
{
"family": "Hebling Vieira",
"given": "Bruno"
},
{
"literal": "for the Alzheimer’s Disease Neuroimaging Initiative"
}
],
"container-title-short": "J Prev Alzheimers Dis",
"volume": "13",
"issue": "9",
"page": "100646",
"DOI": "10.1016/j.tjpad.2026.100646",
"PMID": "42570468",
"PMCID": "PMC13476611",
"ISSN": "2274-5807",
"publisher": "Elsevier",
"URL": "https://doi.org/10.1016/j.tjpad.2026.100646",
"language": "en",
"issued": {
"date-parts": [
[
2026,
8,
8
]
]
}
}

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/s41514-026-00390-w [code]
Mild cognitive impairment cases affect the predictive power of Alzheimer's disease diagnostic models using routine clinical variables.
Journal: npj aging
In common: seaborn, scikit-learn, pandas, 3 other tools, Alzheimer's / dementia, 2 references
[2] doi:10.1038/s42003-026-10205-z [code]
Source-space EEG alpha activity reveals brain age gaps due to neurodegeneration and disparity.
Journal: Communications biology
In common: statsmodels, seaborn, scikit-learn, 4 other tools, Alzheimer's / dementia, structural MRI / diffusion, 1 reference
[3] doi:10.21203/rs.3.rs-9914920/v1 [code]
Prediction of cognitive performance by demographics, sleep, and brain morphometry: machine learning findings from ENIGMA-Sleep Working Group
Journal: Research Square (preprint)
In common: statsmodels, seaborn, scikit-learn, 4 other tools, structural MRI / diffusion, 1 reference
[4] doi:10.1111/ejn.70480 [code]
Astrocyte Proximity Protects Synapses From Human Amyloid-Beta Induced Degeneration in a Mouse Ex Vivo Model of Early Alzheimer's Disease.
Journal: The European journal of neuroscience
In common: statsmodels, seaborn, scikit-learn, 4 other tools, Alzheimer's / dementia, 1 reference
[5] doi:10.1002/advs.202600020 [code]
Early Retinal UCHL1 Dysregulation Coupled With Synaptic Loss Reflects Alzheimer's Disease Severity.
Journal: Advanced science (Weinheim, Baden-Wurttemberg, Germany)
In common: seaborn, scikit-learn, pandas, 3 other tools, Alzheimer's / dementia, 2 references
[6] doi:10.1162/imag.a.1208 [code]
Brain network analysis in Alzheimer's disease and mild cognitive impairment using high-density diffuse optical tomography.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: statsmodels, seaborn, scikit-learn, 4 other tools, Alzheimer's / dementia, 1 reference
[7] doi:10.1371/journal.pone.0343722 [code]
Comprehensive methodology for sample enrichment in EEG biomarker studies for Alzheimer's risk classification.
Journal: PloS one
In common: statsmodels, seaborn, scikit-learn, 4 other tools, Alzheimer's / dementia, 1 reference
[8] doi:10.3389/fnagi.2026.1847611 [code]
APOE ε4-associated hippocampal atrophy trajectories across the Alzheimer's disease continuum: a systematic review, meta-analysis, and longitudinal validation.
Journal: Frontiers in aging neuroscience
In common: statsmodels, seaborn, pandas, 3 other tools, Alzheimer's / dementia, structural MRI / diffusion, 1 reference
[9] doi:10.1038/s41467-026-75661-x [code]
A neural signature of sleep deprivation in the human brain.
Journal: Nature communications
In common: statsmodels, seaborn, scikit-learn, 4 other tools, 1 reference
[10] doi:10.1093/nc/niag029 [code]
A data-driven approach to identifying and evaluating connectivity-based neural correlates of conscious visual perception.
Journal: Neuroscience of consciousness
In common: statsmodels, seaborn, scikit-learn, 4 other tools, 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.