OSCR

Advancing fair and explainable machine learning for neuroimaging dementia pattern classification in multi-racial and multi-ethnic populations.

Code ↔ Paper

8 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 8 matches
  1. [1] § Methods › ML classifier and discrepancy mitigation methods ↔ dementia_ai/models/models.py, lines 644–740 · score 0.63 · target domain, RegAlign, binary, shot, entropy, alignment
  2. [2] § Methods › ML classifier and discrepancy mitigation methods ↔ dementia_ai/models/objectives.py, lines 458–592 · score 0.60 · entropy minimization, unlabeled, shot, supervision, objective, alignment
  3. [3] § Methods › ML classifier and discrepancy mitigation methods ↔ dementia_ai/core/ops.py, lines 297–428 · score 0.59 · upper bound, sample weights, sum, RBF, kernel
  4. [4] § Methods › ML classifier and discrepancy mitigation methods ↔ dementia_ai/core/ops.py, lines 438–485 · score 0.58 · Correlation Remover, sensitive feature, vector, transformation, CR
  5. [5] § Results › Performance discrepancies across multi-racial and multi-ethnic groups ↔ dementia_ai/cli/main.py, lines 210–271 · score 0.56 · Pairwise gap, FPR gap, FNR gap
  6. [6] § Methods › ML classifier and discrepancy mitigation methods ↔ dementia_ai/models/objectives.py, lines 458–592 · score 0.53 · entropy minimization, unlabeled, supervised, probabilities, alignment, SSDA
  7. [7] § Methods ↔ dementia_ai/explainer/shap_utils.py, lines 5–30 · score 0.53 · XGBoost, scikit-learn, tree, boosted, NumPy, model
  8. [8] § Methods › ML classifier and discrepancy mitigation methods ↔ dementia_ai/models/models.py, lines 644–740 · score 0.51 · RegAlign, unlabeled, shot, entropy, alignment, SSDA

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,011 lines · 41 KB · no license · 2 matches

  1. import sys
  2. import json
  3. import traceback
  4. import numpy as np
  5. from tqdm import tqdm
  6. import xgboost as xgb
  7. from sklearn.model_selection import StratifiedKFold, ParameterGrid, GridSearchCV
  8. from sklearn.neural_network import MLPClassifier
  9. from sklearn.ensemble import RandomForestClassifier
  10. from sklearn.svm import SVC
  11. from joblib import parallel_backend
  12. from sklearn.base import clone, BaseEstimator, ClassifierMixin
  13. from sklearn.metrics import check_scoring
  14. from sklearn.base import BaseEstimator
  15. from sklearn.metrics import make_scorer, f1_score, balanced_accuracy_score
  16. from sklearn.utils import check_array
  17. from sklearn.metrics import pairwise
  18. from sklearn.metrics.pairwise import KERNEL_PARAMS
  19. from sklearn.preprocessing import StandardScaler
  20. from sklearn.utils.validation import check_is_fitted
  21. from typing import Optional, Dict, List, Tuple, Any
  22. from cvxopt import matrix, solvers
  23. # Reuse shared preprocessing / KMM routines
  24. from dementia_ai.core import ops
  25. #import custom_objectives as CO
  26. from . import objectives as CO
  27. class _IdentityScaler:
  28. def transform(self, X):
  29. import numpy as np
  30. return np.asarray(X, dtype=np.float32)
  31. def _ba_optimal_threshold(y_true, p1, grid=None) -> float:
  32. """Return the threshold that maximizes Balanced Accuracy on y_true."""
  33. import numpy as np
  34. from sklearn.metrics import balanced_accuracy_score
  35. if grid is None:
  36. # dense grid around the middle; robust for imbalanced sets
  37. grid = np.unique(np.clip(np.r_[np.linspace(0.05,0.95,181), p1], 0, 1))
  38. best, thr = -1.0, 0.5
  39. for t in grid:
  40. yhat = (p1 >= t).astype(int)
  41. ba = balanced_accuracy_score(y_true, yhat)
  42. if ba > best:
  43. best, thr = ba, t
  44. return float(thr)
  45. def _split_xgb_custom_objs(paras: dict):
  46. """Split custom-objective keys from plain XGB params."""
  47. p = dict(paras or {})
  48. focal = {
  49. "use": bool(p.pop("use_focal_obj", False)),
  50. "alpha": float(p.pop("focal_alpha", 0.90)),
  51. "gamma": float(p.pop("focal_gamma", 2.0)),
  52. }
  53. cs = {
  54. "use": bool(p.pop("use_cost_sensitive_obj", False)),
  55. "cost_fn": float(p.pop("cs_cost_fn", 5.0)),
  56. "cost_fp": float(p.pop("cs_cost_fp", 1.0)),
  57. }
  58. return p, focal, cs
  59. # def ba_cost_scorer(est, X, y):
  60. # """Scorer used in grid search. Uses t_cost from est params if available."""
  61. # from sklearn.metrics import balanced_accuracy_score
  62. # proba = est.predict_proba(X)[:, 1]
  63. # # Parameters can be in est.paras (CustomML) or in get_params()
  64. # cs_fn = None
  65. # cs_fp = None
  66. # if hasattr(est, "paras"):
  67. # cs_fn = est.paras.get("cs_cost_fn", None)
  68. # cs_fp = est.paras.get("cs_cost_fp", None)
  69. # if (cs_fn is None or cs_fp is None) and hasattr(est, "get_params"):
  70. # p = est.get_params()
  71. # cs_fn = p.get("cs_cost_fn", cs_fn)
  72. # cs_fp = p.get("cs_cost_fp", cs_fp)
  73. # if cs_fn is not None and cs_fp is not None and (cs_fn + cs_fp) > 0:
  74. # t_cost = cs_fp / (cs_fp + cs_fn)
  75. # else:
  76. # t_cost = 0.5
  77. # yhat = (proba >= t_cost).astype(int)
  78. # return balanced_accuracy_score(y, yhat)
  79. class BACostScorer:
  80. def __init__(self, default_threshold: float = 0.5):
  81. self._sign = 1 # sklearn-compatible
  82. self.default_threshold = float(default_threshold)
  83. def _threshold_from_params(self, est):
  84. # Look in both get_params() and .paras, if present
  85. cf = cp = None
  86. if hasattr(est, "get_params"):
  87. try:
  88. p = est.get_params()
  89. cf = p.get("cs_cost_fn", None)
  90. cp = p.get("cs_cost_fp", None)
  91. except Exception:
  92. pass
  93. if (cf is None or cp is None) and hasattr(est, "paras"):
  94. cf = est.paras.get("cs_cost_fn", cf)
  95. cp = est.paras.get("cs_cost_fp", cp)
  96. if cf is not None and cp is not None and (cf + cp) > 0:
  97. return float(cp) / float(cp + cf)
  98. return self.default_threshold
  99. def __call__(self, estimator, X, y):
  100. thr = self._threshold_from_params(estimator)
  101. # 1) Fast path for our custom XGB class that exposes .booster_ & .scaler_
  102. booster = getattr(estimator, "booster_", None)
  103. scaler = getattr(estimator, "scaler_", None)
  104. if booster is not None and scaler is not None:
  105. try:
  106. Xs = scaler.transform(np.asarray(X, dtype=np.float32))
  107. proba = booster.predict(xgb.DMatrix(Xs))
  108. yhat = (proba >= thr).astype(int)
  109. return balanced_accuracy_score(y, yhat)
  110. except Exception:
  111. pass # fall through
  112. # 2) Fallback to predict_proba if available
  113. if hasattr(estimator, "predict_proba"):
  114. try:
  115. proba = estimator.predict_proba(X)[:, 1]
  116. yhat = (proba >= thr).astype(int)
  117. return balanced_accuracy_score(y, yhat)
  118. except Exception:
  119. pass
  120. # 3) Fallback to predict if available (already hard labels)
  121. if hasattr(estimator, "predict"):
  122. try:
  123. yhat = estimator.predict(X)
  124. # If predict returns floats, threshold them
  125. if yhat.dtype != np.int_ and yhat.dtype != np.bool_:
  126. yhat = (np.asarray(yhat, dtype=float) >= thr).astype(int)
  127. return balanced_accuracy_score(y, yhat)
  128. except Exception:
  129. pass
  130. # 4) Last resort: return the worst possible score for this split
  131. # so this candidate is ignored, but we don't crash the whole search
  132. return -np.inf
  133. ba_cost_scorer = BACostScorer()
  134. def custom_f1_score(y_true, y_pred):
  135. # Example: You can modify this function to compute any custom metric
  136. return f1_score(y_true, y_pred, average='weighted')
  137. def custom_ba_score(y_true, y_pred):
  138. # Example: You can modify this function to compute any custom metric
  139. return balanced_accuracy_score(y_true, y_pred)
  140. # def compute_sample_weight_kmm(Xs, Xt):
  141. # def _fit_weights(Xs, Xt, epsilon):
  142. # n_s = len(Xs)
  143. # n_t = len(Xt)
  144. # if epsilon is None:
  145. # epsilon = (np.sqrt(n_s) - 1)/np.sqrt(n_s)
  146. # kernel_params = {'gamma': 1.0}
  147. # # Compute Kernel Matrix
  148. # K = pairwise.pairwise_kernels(Xs, Xs, metric=kernel,
  149. # **kernel_params)
  150. # K = (1/2) * (K + K.transpose())
  151. # # Compute q
  152. # kappa = pairwise.pairwise_kernels(Xs, Xt,
  153. # metric=kernel,
  154. # **kernel_params)
  155. # kappa = (n_s/n_t) * np.dot(kappa, np.ones((n_t, 1)))
  156. # K = np.ascontiguousarray(K, dtype=np.float64)
  157. # kappa = np.ascontiguousarray(kappa, dtype=np.float64)
  158. # P = matrix(K)
  159. # q = -matrix(kappa)
  160. # # Define constraints
  161. # G = np.ones((2*n_s+2, n_s))
  162. # G[1] = -G[1]
  163. # G[2:n_s+2] = np.eye(n_s)
  164. # G[n_s+2:n_s*2+2] = -np.eye(n_s)
  165. # h = np.ones(2*n_s+2)
  166. # h[0] = n_s*(1+epsilon)
  167. # h[1] = n_s*(epsilon-1)
  168. # h[2:n_s+2] = B
  169. # h[n_s+2:] = 0
  170. # G = matrix(G)
  171. # h = matrix(h)
  172. # solvers.options["show_progress"] = bool(verbose)
  173. # solvers.options["maxiters"] = max_iter
  174. # if tol is not None:
  175. # solvers.options['abstol'] = tol
  176. # solvers.options['reltol'] = tol
  177. # solvers.options['feastol'] = tol
  178. # else:
  179. # solvers.options['abstol'] = 1e-7
  180. # solvers.options['reltol'] = 1e-6
  181. # solvers.options['feastol'] = 1e-7
  182. # weights = solvers.qp(P, q, G, h)['x']
  183. # return np.array(weights).ravel()
  184. # kernel = "rbf"
  185. # B = 1000
  186. # epsilon = None
  187. # max_size = 1000
  188. # tol = None
  189. # verbose = 0
  190. # max_iter = 100
  191. # random_state = 777
  192. # Xs = np.ascontiguousarray(check_array(Xs), dtype=np.float64)
  193. # Xt = np.ascontiguousarray(check_array(Xt), dtype=np.float64)
  194. # np.random.seed(random_state)
  195. # if len(Xs) > max_size:
  196. # size = len(Xs)
  197. # power = 0
  198. # while size > max_size:
  199. # size = size / 2
  200. # power += 1
  201. # split = int(len(Xs) / 2**power)
  202. # shuffled_index = np.random.choice(len(Xs), len(Xs), replace=False)
  203. # weights = np.zeros(len(Xs))
  204. # for i in range(2**power):
  205. # index = shuffled_index[split*i:split*(i+1)]
  206. # weights[index] = _fit_weights(Xs[index], Xt, epsilon)
  207. # else:
  208. # weights = _fit_weights(Xs, Xt, epsilon)
  209. # return weights
  210. class CustomML(BaseEstimator):
  211. def __init__(self, model_name, paras=None, sample_weight=None, **kwargs):
  212. self.model_name = model_name
  213. self.paras = paras if paras is not None else {}
  214. #print("Init:", self.paras)
  215. self.sample_weight = sample_weight
  216. self.kwargs = kwargs
  217. self._initialize_predictor()
  218. def _initialize_predictor(self):
  219. import inspect
  220. def _filter_params(params, cls):
  221. if params is None:
  222. return {}
  223. valid = set(inspect.signature(cls.__init__).parameters.keys())
  224. # always ignore our internal dummy knob if present
  225. params = dict(params)
  226. params.pop("__dummy__", None)
  227. return {k: v for k, v in params.items() if k in valid}
  228. if self.model_name == 'svm':
  229. p = _filter_params(self.paras, SVC)
  230. k = _filter_params(self.kwargs, SVC)
  231. # guarantee probabilities unless explicitly overridden
  232. if 'probability' not in p and 'probability' not in k:
  233. p['probability'] = True
  234. self.predictor = SVC(**p, **k)
  235. elif self.model_name == 'xgboost_vanilla':
  236. # Pure, dependable XGB — no custom objective, no fancy knobs
  237. p = dict(self.paras or {})
  238. # normalize param aliases
  239. if 'lambda' in p and 'reg_lambda' not in p: p['reg_lambda'] = p.pop('lambda')
  240. if 'alpha' in p and 'reg_alpha' not in p: p['reg_alpha'] = p.pop('alpha')
  241. if 'eta' not in p and 'learning_rate' in p: p['eta'] = p.pop('learning_rate')
  242. # safe defaults if not provided
  243. p.setdefault('objective', 'binary:logistic')
  244. p.setdefault('eval_metric', 'aucpr')
  245. p.setdefault('tree_method', 'hist')
  246. p.setdefault('device', self.kwargs.get('device', 'cuda'))
  247. p.setdefault('max_depth', 4)
  248. p.setdefault('eta', 0.1)
  249. p.setdefault('subsample', 0.9)
  250. p.setdefault('colsample_bytree', 0.9)
  251. p.setdefault('reg_lambda', 1.0)
  252. p.setdefault('reg_alpha', 0.0)
  253. p.setdefault('n_estimators', self.kwargs.get('total_steps', 300))
  254. p.setdefault('verbosity', 0)
  255. p.setdefault('nthread', 8)
  256. self.predictor = xgb.XGBClassifier(**p)
  257. # vanilla keeps a fold-tuned BA threshold (optional; default 0.5)
  258. self.best_threshold_ = 0.5
  259. # wrap .fit to optionally accept (X_val, y_val) for early-stopping & BA threshold
  260. self._xgb_orig_fit = self.predictor.fit
  261. self._fit_with_es = self._xgb_vanilla_fit_with_es
  262. # helper: class predictions at tuned threshold
  263. self._xgb_predict_proba_fn = self._xgb_vanilla_predict_proba
  264. elif self.model_name == 'random_forest':
  265. self.predictor = RandomForestClassifier(
  266. **_filter_params(self.paras, RandomForestClassifier),
  267. **_filter_params(self.kwargs, RandomForestClassifier)
  268. )
  269. elif self.model_name == 'mlp':
  270. self.predictor = MLPClassifier(
  271. **_filter_params(self.paras, MLPClassifier),
  272. **_filter_params(self.kwargs, MLPClassifier)
  273. )
  274. def set_params(self, **params):
  275. p = dict(params or {})
  276. p.pop("__dummy__", None)
  277. # Preserve/override sample_weight explicitly
  278. if "sample_weight" in p:
  279. self.sample_weight = p.pop("sample_weight")
  280. # keep existing sample_weight if not provided
  281. # Normalize some xgboost aliases so grid params actually change the model
  282. if self.model_name == "xgboost_vanilla":
  283. if "lambda" in p and "reg_lambda" not in p:
  284. p["reg_lambda"] = p.pop("lambda")
  285. if "alpha" in p and "reg_alpha" not in p:
  286. p["reg_alpha"] = p.pop("alpha")
  287. if "learning_rate" in p and "eta" not in p:
  288. p["eta"] = p.pop("learning_rate")
  289. self.paras = p
  290. self._initialize_predictor()
  291. return self
  292. def fit(self, X, y, sample_weight=None):
  293. # ---------------- models (svm / rf / mlp / xgboost_vanilla) ----------------
  294. if sample_weight is not None:
  295. self.predictor.fit(X, y, sample_weight=sample_weight)
  296. else:
  297. self.predictor.fit(X, y)
  298. return self
  299. def _xgb_vanilla_fit_with_es(self, X, y, X_val=None, y_val=None, early_stopping_rounds=50, sample_weight=None):
  300. _orig_fit = getattr(self, "_xgb_orig_fit", None)
  301. if _orig_fit is None:
  302. raise AttributeError("_xgb_orig_fit not set; ensure xgboost_vanilla predictor initialized")
  303. fit_kwargs = {}
  304. if (X_val is not None) and (y_val is not None) and (early_stopping_rounds is not None):
  305. try:
  306. self.predictor.set_params(early_stopping_rounds=int(early_stopping_rounds))
  307. except Exception:
  308. pass
  309. fit_kwargs['eval_set'] = [(X, y), (X_val, y_val)]
  310. if sample_weight is not None:
  311. fit_kwargs["sample_weight"] = sample_weight
  312. _orig_fit(X, y, **fit_kwargs)
  313. if X_val is not None and y_val is not None:
  314. p1_val = self.predictor.predict_proba(X_val)[:, 1]
  315. self.best_threshold_ = _ba_optimal_threshold(y_val, p1_val)
  316. return self
  317. def _xgb_vanilla_predict_proba(self, X):
  318. if hasattr(self.predictor, "predict_proba"):
  319. return self.predictor.predict_proba(X)
  320. import numpy as np
  321. if hasattr(self.predictor, "decision_function"):
  322. s = self.predictor.decision_function(X)
  323. s = np.asarray(s)
  324. if s.ndim == 1:
  325. p1 = 1.0 / (1.0 + np.exp(-s))
  326. return np.vstack([1.0 - p1, p1]).T
  327. else:
  328. e = np.exp(s - s.max(axis=1, keepdims=True))
  329. p = e / e.sum(axis=1, keepdims=True)
  330. return p
  331. raise AttributeError("No predict_proba/decision_function available.")
  332. def predict(self, X):
  333. if self.model_name == 'xgboost_vanilla' and hasattr(self, "best_threshold_"):
  334. p1 = self._xgb_vanilla_predict_proba(X)[:, 1]
  335. thr = getattr(self, "best_threshold_", 0.5)
  336. return (p1 >= thr).astype(int)
  337. return self.predictor.predict(X)
  338. def predict_proba(self, X):
  339. if self.model_name == 'xgboost_vanilla':
  340. return self._xgb_vanilla_predict_proba(X)
  341. # Native path if available
  342. if hasattr(self.predictor, "predict_proba"):
  343. return self.predictor.predict_proba(X)
  344. # Fallback: convert decision scores to [p0, p1] via logistic
  345. import numpy as np
  346. if hasattr(self.predictor, "decision_function"):
  347. s = self.predictor.decision_function(X)
  348. s = np.asarray(s)
  349. if s.ndim == 1:
  350. # binary: map to 2-column probs
  351. p1 = 1.0 / (1.0 + np.exp(-s))
  352. return np.vstack([1.0 - p1, p1]).T
  353. else:
  354. # multiclass one-vs-rest scores -> softmax
  355. e = np.exp(s - s.max(axis=1, keepdims=True))
  356. p = e / e.sum(axis=1, keepdims=True)
  357. return p
  358. # Last resort: hard predictions -> degenerate probabilities
  359. y = self.predictor.predict(X)
  360. y = np.asarray(y)
  361. classes = np.unique(y)
  362. if classes.size == 2:
  363. p1 = (y == classes.max()).astype(float)
  364. return np.vstack([1.0 - p1, p1]).T
  365. else:
  366. # multiclass: one-hot
  367. K = classes.size
  368. idx = {c: i for i, c in enumerate(classes)}
  369. probs = np.zeros((y.shape[0], K), dtype=float)
  370. for i, yy in enumerate(y):
  371. probs[i, idx[yy]] = 1.0
  372. return probs
  373. def compute_shap_values(self, X_train, X_test, K=None):
  374. try:
  375. import shap
  376. except Exception:
  377. print("[SHAP] shap not installed; skipping SHAP explanation.", flush=True)
  378. return None
  379. def model_predict(X):
  380. return self.predict_proba(X)[:, 1]
  381. if K:
  382. X_train_summary = shap.sample(X_train, K)
  383. else:
  384. X_train_summary = X_train
  385. explainer = shap.KernelExplainer(model_predict, X_train_summary)
  386. shap_values = explainer.shap_values(X_test)
  387. return explainer, shap_values
  388. @staticmethod
  389. def hyperparameter_optimization(X_train, y_train, model_name, X_test=None, y_test=None, cv=None,
  390. init_paras=None, sample_weight=None, param_grid=None, n_jobs=-1):
  391. custom_scorer = make_scorer(custom_ba_score, greater_is_better=True)
  392. opt = MyGridSearchCV(estimator=CustomML(model_name, paras=init_paras, sample_weight=sample_weight),
  393. param_grid=param_grid,
  394. scoring=custom_scorer,
  395. n_jobs=n_jobs,
  396. cv=cv,
  397. verbose=0,
  398. refit=True)
  399. #opt.set_separated_test(X_test, y_test)
  400. with parallel_backend('loky', n_jobs=n_jobs):
  401. opt.fit(X_train, y_train)
  402. #print("Best Hyperparameters: ", opt.best_params_)
  403. return opt
  404. def fit_and_score(estimator, X, y, train, test, parameters, scorer, X_test=None, y_test=None):
  405. #est = clone(estimator).set_params(**parameters)
  406. est = clone(estimator)
  407. # preserve sample_weight attribute across clone if present
  408. if hasattr(estimator, "sample_weight"):
  409. try:
  410. est.sample_weight = getattr(estimator, "sample_weight")
  411. except Exception:
  412. pass
  413. if parameters:
  414. # Pass all grid parameters through; CustomML.set_params will normalize/ignore as needed.
  415. try:
  416. est.set_params(**parameters)
  417. except Exception:
  418. # Fallback to legal subset if something goes wrong
  419. legal = est.get_params(deep=True).keys()
  420. clean = {k: v for k, v in parameters.items() if k in legal}
  421. if clean:
  422. est.set_params(**clean)
  423. # slice
  424. X_tr = X.iloc[train] if hasattr(X, "iloc") else X[train]
  425. y_tr = y.iloc[train] if hasattr(y, "iloc") else y[train]
  426. # pass group ids / sample_weight when present (kept from your version)
  427. if hasattr(est, "group_ids_source") and est.group_ids_source is not None:
  428. est.group_ids_source = np.asarray(est.group_ids_source)[train]
  429. sw = getattr(est, "sample_weight", None)
  430. if sw is not None:
  431. sw_tr = sw[train]
  432. est.fit(X_tr, y_tr, sample_weight=sw_tr)
  433. else:
  434. est.fit(X_tr, y_tr)
  435. # --- robust scoring ---
  436. try:
  437. if X_test is not None and y_test is not None:
  438. score_list = [scorer(est, xt, yt) for xt, yt in zip(X_test, y_test)]
  439. score = sum(score_list) / len(score_list)
  440. else:
  441. X_te = X.iloc[test] if hasattr(X, "iloc") else X[test]
  442. y_te = y.iloc[test] if hasattr(y, "iloc") else y[test]
  443. score = scorer(est, X_te, y_te)
  444. except Exception:
  445. # If anything blows up, give this parameter set a terrible score
  446. score = -np.inf
  447. return parameters, score, est
  448. class MyGridSearchCV(GridSearchCV):
  449. def __init__(self, *args, **kwargs):
  450. super(MyGridSearchCV, self).__init__(*args, **kwargs)
  451. self.all_best_parameters_ = {}
  452. self.split_indices = {
  453. 'train': [],
  454. 'test': []
  455. }
  456. self.test_score_fold = []
  457. self.best_test_scores = []
  458. self.X_test = None
  459. self.y_test = None
  460. @staticmethod
  461. def _param_key(params: dict):
  462. return tuple(sorted(params.items()))
  463. def set_separated_test(self, X_test, y_test):
  464. self.X_test = X_test
  465. self.y_test = y_test
  466. def scorer_single(self, estimator, X, y_true):
  467. y_pred = estimator.predict(X)
  468. return balanced_accuracy_score(y_true, y_pred)
  469. def fit(self, X, y=None):
  470. self.scorer_ = self.scoring
  471. # Wrap integer cv into StratifiedKFold
  472. cv = self.cv if not isinstance(self.cv, int) else StratifiedKFold(
  473. n_splits=self.cv, shuffle=True, random_state=42
  474. )
  475. if self.verbose > 0:
  476. print(f"Fitting {len(cv)} folds for each of {len(self.param_grid)} candidates, "
  477. f"totalling {len(cv) * len(self.param_grid)} fits")
  478. # materialize the full list of parameter settings
  479. if not self.param_grid:
  480. self.param_iter = [dict()]
  481. else:
  482. self.param_iter = list(ParameterGrid(self.param_grid))
  483. # for sensitivity: collect scores per param set
  484. score_map = {self._param_key(p): [] for p in self.param_iter}
  485. param_map = {self._param_key(p): p for p in self.param_iter}
  486. for split_idx, (train, test) in enumerate(cv.split(X, y)):
  487. n_tr = len(train)
  488. n_va = len(test)
  489. print(f" =========== Training Split {split_idx} - train {n_tr} samples - val {n_va} samples =========== ")
  490. # reset per-fold best trackers
  491. self.best_score_fold = -np.inf if self.scorer_._sign > 0 else np.inf
  492. self.best_params_fold = None
  493. # grid search with progress bar
  494. with tqdm(total=len(self.param_iter),
  495. desc="Grid Search Progress",
  496. bar_format='{desc:<10}{percentage:3.0f}%|{bar:30}{r_bar}') as pbar:
  497. # sequential evaluation of each parameter setting
  498. for parameters in self.param_iter:
  499. parameters, score, est = fit_and_score(
  500. self.estimator, X, y, train, test,
  501. parameters, self.scorer_, self.X_test, self.y_test
  502. )
  503. # update global best if refit=True
  504. if self.refit and score > getattr(self, 'best_score_', -np.inf):
  505. self.best_score_ = score
  506. self.best_params_ = parameters
  507. self.best_estimator_ = est
  508. # update this fold’s best
  509. if score > self.best_score_fold:
  510. self.best_score_fold = score
  511. self.best_params_fold = parameters
  512. # record fold‐test score & advance bar
  513. score_map[self._param_key(parameters)].append(score)
  514. self.test_score_fold.append(score)
  515. try:
  516. pbar.set_postfix_str(f"val_score:{score:.4f}")
  517. except Exception:
  518. pass
  519. pbar.update(1)
  520. # after finishing this fold
  521. self.best_test_scores.append(self.best_score_fold)
  522. self.all_best_parameters_[f'split{split_idx}_params'] = self.best_params_fold
  523. self.split_indices['train'].append(train)
  524. self.split_indices['test'].append(test)
  525. # Build cv_results_ summary (mean/std per param set)
  526. param_keys = list(score_map.keys())
  527. param_names = sorted({k for p in param_map.values() for k in p.keys()})
  528. self.cv_results_ = {
  529. "mean_test_score": [np.mean(score_map[k]) if score_map[k] else -np.inf for k in param_keys],
  530. "std_test_score": [np.std(score_map[k], ddof=1) if len(score_map[k]) > 1 else 0.0 for k in param_keys],
  531. }
  532. for name in param_names:
  533. self.cv_results_[f"param_{name}"] = [param_map[k].get(name) for k in param_keys]
  534. return self
  535. def predict(self, X):
  536. if hasattr(self.best_estimator_, "predict"):
  537. self._check_is_fitted('predict')
  538. return self.best_estimator_.predict(X)
  539. else:
  540. raise AttributeError("The best estimator does not have a 'predict' method.")
  541. def predict_proba(self, X):
  542. if hasattr(self.best_estimator_, "predict_proba"):
  543. self._check_is_fitted('predict_proba')
  544. return self.best_estimator_.predict_proba(X)
  545. else:
  546. raise AttributeError("The best estimator does not have a 'predict_proba' method.")
  547. def score(self, X, y=None):
  548. if hasattr(self.best_estimator_, "score"):
  549. self._check_is_fitted('score')
  550. return self.best_estimator_.score(X, y)
  551. else:
  552. raise AttributeError("The best estimator does not have a 'score' method.")
  553. class XGBCustomObjectiveEstimator(BaseEstimator, ClassifierMixin):
  554. def __init__(self,
  555. # Custom objective selection
  556. objective_version: str = 'v3',
  557. objective_subtype: str = '1C', # '1C' (RegAlign) or '2B' (SSDA)
  558. total_steps: int = 500,
  559. eval_metric: str = 'auc',
  560. loss_mode: str = None, # 'ce','align','ce+align' (1C) | 'entropy','align','entropy+align' (2B)
  561. # target domain (few-shot/regalign or unlabeled/ssda)
  562. target_X: Optional[np.ndarray] = None,
  563. target_y: Optional[np.ndarray] = None,
  564. # validation for logging
  565. X_valid: Optional[np.ndarray] = None,
  566. y_valid: Optional[np.ndarray] = None,
  567. callbacks: Optional[list] = None,
  568. # --- method hyper-params (tuned by grid) ---
  569. tau: float = 1.0, # target CE / entropy weight
  570. eps: float = 1.0, # alignment/GBA weight
  571. source_focal_gamma: float = 0.0,
  572. target_pos_weight: float = 1.0,
  573. update_every: int = 25,
  574. fair_grad_alpha: float = 0.0, fair_grad_ema: float = 0.9, eq_lambda: float = 0.25, cal_lambda: float = 0.2,
  575. gap_kind: str = 'ba',
  576. class1_boost: float = 1.0,
  577. lambda_align: float = 0.1,
  578. # ----- NEW knobs for V3 -----
  579. target_focal_gamma: float = 0.0,
  580. use_target_reg: bool = False,
  581. target_weight: float = 1.0,
  582. # (ALL-groups homogenization)
  583. bary_align: bool = False,
  584. var_align: bool = False,
  585. center_on: str = 'barycenter', # or 'source'|'target'
  586. group_ids_source: Optional[np.ndarray] = None,
  587. group_ids_target: Optional[np.ndarray] = None,
  588. # --- plain xgb params (also tuned by grid) ---
  589. **xgb_params):
  590. self.objective_version = objective_version
  591. self.objective_subtype = objective_subtype
  592. self.total_steps = total_steps
  593. self.eval_metric = eval_metric
  594. self.target_X = target_X
  595. self.target_y = target_y
  596. self.X_valid = X_valid
  597. self.y_valid = y_valid
  598. self.callbacks = callbacks
  599. self.loss_mode = loss_mode
  600. # method HPs
  601. self.tau = tau
  602. self.eps = eps
  603. self.source_focal_gamma = source_focal_gamma
  604. self.target_pos_weight = target_pos_weight
  605. self.update_every = update_every
  606. self.class1_boost = class1_boost
  607. self.lambda_align = float(lambda_align)
  608. # ---- NEW: V3 knobs ----
  609. self.target_focal_gamma = target_focal_gamma
  610. self.use_target_reg = use_target_reg
  611. self.target_weight = target_weight
  612. self.bary_align = bary_align
  613. self.var_align = var_align
  614. self.center_on = center_on
  615. self.group_ids_source = group_ids_source
  616. self.group_ids_target = group_ids_target
  617. self.fair_grad_alpha = fair_grad_alpha
  618. self.fair_grad_ema = fair_grad_ema
  619. self.eq_lambda = eq_lambda
  620. self.cal_lambda = cal_lambda
  621. self.gap_kind = gap_kind
  622. # defaults + alias normalization (unchanged)
  623. defaults = dict(
  624. max_depth=3, eta=0.3, reg_lambda=1.0, reg_alpha=0.0,
  625. subsample=1.0, colsample_bytree=1.0,
  626. tree_method="hist", device=xgb_params.get("device", "cuda"),
  627. objective="binary:logistic",
  628. scale_pos_weight=1.0, nthread=8, verbosity=0,
  629. eval_metric=self.eval_metric
  630. )
  631. defaults.update(xgb_params)
  632. if 'lambda' in defaults and 'reg_lambda' not in defaults:
  633. defaults['reg_lambda'] = defaults.pop('lambda')
  634. if 'alpha' in defaults and 'reg_alpha' not in defaults:
  635. defaults['reg_alpha'] = defaults.pop('alpha')
  636. if 'eta' not in defaults and 'learning_rate' in defaults:
  637. defaults['eta'] = defaults.pop('learning_rate')
  638. defaults['objective'] = 'binary:logistic'
  639. defaults.setdefault('eval_metric', self.eval_metric)
  640. defaults.setdefault('max_delta_step', 1)
  641. self.xgb_params: Dict = defaults
  642. # fitted state
  643. self.booster_ = None
  644. self.classes_ = np.array([0, 1])
  645. def get_params(self, deep=True):
  646. p = dict(self.xgb_params)
  647. p.update(dict(
  648. objective_version=self.objective_version,
  649. objective_subtype=self.objective_subtype,
  650. total_steps=self.total_steps,
  651. eval_metric=self.eval_metric,
  652. target_X=self.target_X,
  653. target_y=self.target_y,
  654. X_valid=self.X_valid,
  655. y_valid=self.y_valid,
  656. callbacks=self.callbacks,
  657. # method HPs
  658. tau=self.tau, eps=self.eps,
  659. source_focal_gamma=self.source_focal_gamma,
  660. target_pos_weight=self.target_pos_weight,
  661. update_every=self.update_every,
  662. class1_boost=self.class1_boost,
  663. loss_mode=self.loss_mode,
  664. # NEW V3 knobs
  665. target_focal_gamma=self.target_focal_gamma,
  666. use_target_reg=self.use_target_reg,
  667. target_weight=self.target_weight,
  668. lambda_align=self.lambda_align,
  669. bary_align=self.bary_align,
  670. var_align=self.var_align,
  671. center_on=self.center_on,
  672. group_ids_source=self.group_ids_source,
  673. group_ids_target=self.group_ids_target,
  674. gap_lambda=float(getattr(self, "gap_lambda", 1.0)),
  675. ba_tau=float(getattr(self, "ba_tau", 2.0)),
  676. fair_grad_alpha=self.fair_grad_alpha,
  677. fair_grad_ema=self.fair_grad_ema,
  678. eq_lambda=self.eq_lambda,
  679. cal_lambda=self.cal_lambda,
  680. gap_kind=self.gap_kind
  681. ))
  682. return p
  683. def set_params(self, **params):
  684. known = {
  685. 'objective_version','objective_subtype','total_steps','eval_metric',
  686. 'target_X','target_y','X_valid','y_valid','callbacks',
  687. 'tau','eps','source_focal_gamma','target_pos_weight','update_every',
  688. 'class1_boost','loss_mode','target_focal_gamma','use_target_reg','lambda_align',
  689. 'target_weight','bary_align','var_align','center_on',
  690. 'group_ids_source','group_ids_target','gap_lambda','ba_tau','fair_grad_alpha',
  691. 'fair_grad_ema','eq_lambda','cal_lambda','gap_kind'
  692. }
  693. for k, v in params.items():
  694. if k in known:
  695. setattr(self, k, v)
  696. else:
  697. self.xgb_params[k] = v
  698. return self
  699. def _build_custom_obj(self, X_src, y_src, X_tgt, y_tgt):
  700. n_source = X_src.shape[0]
  701. m_target = 0 if X_tgt is None else X_tgt.shape[0]
  702. X_comb = X_src if m_target == 0 else np.vstack([X_src, X_tgt])
  703. kwargs = dict(
  704. total_steps=self.total_steps,
  705. n_source=n_source,
  706. m_target=m_target,
  707. tau=float(self.tau),
  708. eps=float(self.eps),
  709. approach=self.objective_subtype, # '1C' or '2B'
  710. loss_mode=self.loss_mode, # same modes as v2
  711. source_focal_gamma=float(self.source_focal_gamma),
  712. target_pos_weight=float(self.target_pos_weight),
  713. update_every=int(self.update_every),
  714. gap_kind=self.gap_kind,
  715. fair_grad_alpha=float(self.fair_grad_alpha),
  716. fair_grad_ema=float(self.fair_grad_ema),
  717. eq_lambda=float(self.eq_lambda),
  718. cal_lambda=float(self.cal_lambda),
  719. class1_boost=float(self.class1_boost),
  720. lambda_align=float(self.lambda_align),
  721. # v1-style options (restored)
  722. target_focal_gamma=float(self.target_focal_gamma),
  723. use_target_reg=bool(self.use_target_reg),
  724. target_weight=float(self.target_weight),
  725. # ALL-groups homogenization
  726. bary_align=bool(self.bary_align),
  727. var_align=bool(self.var_align),
  728. center_on=str(self.center_on),
  729. group_ids_source=self.group_ids_source,
  730. group_ids_target=self.group_ids_target,
  731. gap_lambda=float(getattr(self, "gap_lambda", 1.0)),
  732. ba_tau=float(getattr(self, "ba_tau", 2.0)),
  733. X_combined=X_comb,
  734. y_source=y_src
  735. )
  736. try:
  737. method_keys = [
  738. #"tau","eps","source_focal_gamma","target_focal_gamma","target_pos_weight",
  739. #"class1_boost","lambda_align",
  740. #"target_weight","eq_lambda","cal_lambda",
  741. "fair_grad_alpha","fair_grad_ema","gap_lambda","ba_tau","update_every",
  742. #"bary_align","var_align","center_on"
  743. ]
  744. meth = {k: kwargs.get(k) for k in method_keys if k in kwargs}
  745. #print("[DEBUG][CustomObjV3] method_params:", json.dumps(meth, default=str))
  746. #print("[DEBUG][CustomObjV3] xgb_params:", json.dumps(self.xgb_params, default=str))
  747. except Exception:
  748. pass
  749. # ---- V3 objective with GBA + smooth BA-gap + v1/v2-compat knobs ----
  750. return CO.CustomObjectiveV3(**kwargs)
  751. def fit(self, X, y):
  752. X_src = np.asarray(X, dtype=np.float32)
  753. y = np.asarray(y, dtype=np.float32).ravel()
  754. self.pos_rate_ = float(np.mean(y)) if y.size else 0.5
  755. X_tgt = None if self.target_X is None else np.asarray(self.target_X, dtype=np.float32)
  756. y_tgt = None if self.target_y is None else np.asarray(self.target_y, dtype=np.float32).ravel()
  757. # Align group ids lengths with source/target
  758. gid_src = None
  759. if getattr(self, "group_ids_source", None) is not None:
  760. gid_src = np.asarray(self.group_ids_source).astype(int)
  761. if gid_src.shape[0] != X_src.shape[0]:
  762. print(f"[XGB-CUSTOM] WARN: group_ids_source len {len(gid_src)} != X_src len {len(X_src)}; resizing.")
  763. if gid_src.shape[0] == 0:
  764. gid_src = np.zeros(X_src.shape[0], dtype=int)
  765. elif gid_src.shape[0] > X_src.shape[0]:
  766. gid_src = gid_src[:X_src.shape[0]]
  767. else:
  768. gid_src = np.resize(gid_src, X_src.shape[0])
  769. gid_tgt = None
  770. if X_tgt is not None and getattr(self, "group_ids_target", None) is not None:
  771. gid_tgt = np.asarray(self.group_ids_target).astype(int)
  772. if gid_tgt.shape[0] != X_tgt.shape[0]:
  773. print(f"[XGB-CUSTOM] WARN: group_ids_target len {len(gid_tgt)} != X_tgt len {len(X_tgt)}; resizing.")
  774. if gid_tgt.shape[0] == 0:
  775. gid_tgt = np.zeros(X_tgt.shape[0], dtype=int)
  776. elif gid_tgt.shape[0] > X_tgt.shape[0]:
  777. gid_tgt = gid_tgt[:X_tgt.shape[0]]
  778. else:
  779. gid_tgt = np.resize(gid_tgt, X_tgt.shape[0])
  780. if X_tgt is None:
  781. dtrain = xgb.DMatrix(X_src, label=y)
  782. else:
  783. if y_tgt is None:
  784. y_comb = np.concatenate([y, np.zeros(X_tgt.shape[0], dtype=np.float32)], axis=0)
  785. else:
  786. y_comb = np.concatenate([y, y_tgt], axis=0)
  787. dtrain = xgb.DMatrix(np.vstack([X_src, X_tgt]), label=y_comb.ravel())
  788. # optional validation for logging
  789. evals = [(dtrain, 'train')]
  790. if self.X_valid is not None and self.y_valid is not None:
  791. Xv = np.asarray(self.X_valid, dtype=np.float32)
  792. yv = np.asarray(self.y_valid, dtype=np.float32).ravel()
  793. dvalid = xgb.DMatrix(Xv, label=yv)
  794. evals.append((dvalid, 'valid'))
  795. # Override gids for this fit
  796. if gid_src is not None:
  797. self.group_ids_source = gid_src
  798. if gid_tgt is not None:
  799. self.group_ids_target = gid_tgt
  800. custom_obj = self._build_custom_obj(X_src, y, X_tgt, y_tgt)
  801. try:
  802. # # Collect eval names safely
  803. # eval_names = [name for (_, name) in (evals or [])]
  804. # # Method HPs we care to see (only include those present on self)
  805. # hp_keys = [
  806. # "tau", "eps", "target_weight",
  807. # "source_focal_gamma", "target_focal_gamma",
  808. # "target_pos_weight", "class1_boost",
  809. # "gap_lambda", "ba_tau", "update_every"
  810. # ]
  811. # method_hparams = {k: getattr(self, k) for k in hp_keys if hasattr(self, k)}
  812. # debug_info = {
  813. # "branch": "custom-xgboost-train",
  814. # "loss_mode": getattr(self, "loss_mode", None),
  815. # "approach": getattr(self, "approach", None),
  816. # "total_steps": int(self.total_steps),
  817. # "xgb_params": self.xgb_params,
  818. # "method_hparams": method_hparams,
  819. # "train_rows": getattr(dtrain, "num_row", lambda: None)(),
  820. # "train_cols": getattr(dtrain, "num_col", lambda: None)(),
  821. # "evals": eval_names,
  822. # "callbacks": [type(cb).__name__ for cb in (self.callbacks or [])],
  823. # }
  824. # print("[XGB-CUSTOM] Starting training with:", flush=True)
  825. # print(json.dumps(debug_info, indent=2, default=str), flush=True)
  826. for bad in ("n_jobs", "n_estimators"):
  827. if bad in self.xgb_params:
  828. self.xgb_params.pop(bad, None)
  829. self.booster_ = xgb.train(
  830. params=self.xgb_params,
  831. dtrain=dtrain,
  832. num_boost_round=int(self.total_steps),
  833. obj=custom_obj,
  834. evals=evals,
  835. callbacks=(self.callbacks if self.callbacks is not None else []),
  836. verbose_eval=False
  837. )
  838. # # Success log
  839. # btype = type(self.booster_).__name__
  840. # try:
  841. # attrs = self.booster_.attributes()
  842. # except Exception:
  843. # attrs = {}
  844. # print(f"[XGB-CUSTOM] Training finished. booster={btype}, attrs_keys={list(attrs.keys())}", flush=True)
  845. except Exception as e:
  846. # Failure path → DummyBooster
  847. print("[XGB-CUSTOM] TRAINING FAILED — falling back to _DummyBooster.", flush=True)
  848. print("[XGB-CUSTOM] Exception:", repr(e), flush=True)
  849. traceback.print_exc()
  850. class _DummyBooster:
  851. def predict(self, dm):
  852. n = dm.num_row()
  853. p = np.full(n, fill_value=max(1e-6, min(1-1e-6, getattr(self, "p_", 0.5))), dtype=np.float32)
  854. return p
  855. db = _DummyBooster()
  856. db.p_ = self.pos_rate_
  857. self.booster_ = db
  858. self.used_dummy_booster_ = True
  859. print(f"[XGB-CUSTOM] Using _DummyBooster with constant p_={self.pos_rate_:.6f}", flush=True)
  860. return self
  861. def predict_proba(self, X):
  862. if getattr(self, "used_dummy_booster_", False) and not getattr(self, "_warned_dummy_", False):
  863. print("[XGB-CUSTOM] WARNING: predictions come from _DummyBooster (constant probs).", flush=True)
  864. self._warned_dummy_ = True
  865. if self.booster_ is None:
  866. # constant fallback
  867. n = len(X)
  868. p = np.full(n, fill_value=max(1e-6, min(1-1e-6, getattr(self, "pos_rate_", 0.5))), dtype=np.float32)
  869. return np.vstack([1 - p, p]).T
  870. Xs = np.asarray(X, dtype=np.float32)
  871. p = self.booster_.predict(xgb.DMatrix(Xs))
  872. p = np.clip(p, 1e-6, 1 - 1e-6)
  873. return np.vstack([1 - p, p]).T
  874. def predict(self, X):
  875. if getattr(self, "used_dummy_booster_", False) and not getattr(self, "_warned_dummy_", False):
  876. print("[XGB-CUSTOM] WARNING: predictions come from _DummyBooster (constant probs).", flush=True)
  877. self._warned_dummy_ = True
  878. return (self.predict_proba(X)[:, 1] >= 0.5).astype(int)
  879. def _xgb_vanilla_predict_proba(self, X):
  880. if hasattr(self.predictor, "predict_proba"):
  881. return self.predictor.predict_proba(X)
  882. import numpy as np
  883. if hasattr(self.predictor, "decision_function"):
  884. s = np.asarray(self.predictor.decision_function(X))
  885. if s.ndim == 1:
  886. p1 = 1.0 / (1.0 + np.exp(-s))
  887. return np.vstack([1.0 - p1, p1]).T
  888. raise AttributeError("No predict_proba/decision_function available.")
  889. def predict(self, X):
  890. p1 = self.predict_proba(X)[:, 1]
  891. thr = getattr(self, "best_threshold_", 0.5)
  892. return (p1 >= thr).astype(int)

models.py at commit 6263968, no license · at the source

Overview

Authors: Ngoc-Huynh Ho1, Sokratis Charisis1, Nicolas Honnorat1, Sachintha Ransara Brandigampala1, Di Wang1, Susan R Heckbert2, Peter T Fox3, David Martinez1, David H Wang1, Timothy M Hughes4, Derek B Archer5, Timothy J Hohman5, Sudha Seshadri1, Christos Davatzikos6, Mohamad Habes1
  1. Glenn Biggs Institute for Neurodegenerative Disorders, Neuroimage Analytics Laboratory and Biggs Institute Neuroimaging Core, University of Texas Health Science Center at San Antonio, San Antonio, TX USA
  2. Department of Epidemiology, University of Washington, Seattle, WA USA
  3. Research Imaging Institute, University of Texas Health Science Center at San Antonio, San Antonio, TX USA
  4. Gerontology and Geriatric Medicine, Wake Forest University School of Medicine, Winston-Salem, NC USA
  5. Vanderbilt Memory and Alzheimer’s Center, Vanderbilt University Medical Center, Nashville, TN USA
  6. Department of Radiology, University of Pennsylvania, Philadelphia, PA USA
Journal: Nature communications, volume 17, issue 1, article 8026
Dates: received 30 May 2025; accepted 4 June 2026; published online 26 June 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1038/s41467-026-74515-w · PMID 42362543 · PMCID PMC13454470 · OpenAlex W7166086908
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, Spectral & time-frequency, fMRI & imaging
Keywords: Alzheimer's disease, Learning algorithms, Computational models
MeSH: Dementia*, Machine Learning*, Neuroimaging*, Black or African American, Classification Algorithms, Ethnicity, Hispanic or Latino, Humans, Magnetic Resonance Imaging, Reproducibility of Results, White (* major topic)
Topic: Dementia and Cognitive Impairment Research (Psychiatry and Mental health, Medicine), according to OpenAlex
Funding: NIA NIH HHS (P20 AG068082, P30 AG062422, P30 AG062429, P30 AG066462, P30 AG066546, P30 AG072975, P20 AG068024, P30 AG072959, P30 AG072972, P30 AG072973, P30 AG072976, P30 AG072979, R01 AG080821, P20 AG068077, P30 AG062715, P30 AG066444, R01 AG079280, U24 AG074855, P30 AG062677, P30 AG066468, P30 AG066507, P30 AG066508, P30 AG066511, P30 AG066518, P30 AG066530, R01 AG083865, U24 AG072122, P30 AG062421, P30 AG066506, P30 AG066512, P30 AG072931, P30 AG072958, P20 AG068053, P30 AG066509, P30 AG066514, P30 AG066515, P30 AG066519, P30 AG072977, P30 AG072978, P30 AG072946, P30 AG072947, R01 AG085571); NIMH NIH HHS (R56 MH074457, R01 MH074457)
Citations: not cited yet (Europe PMC); 56 references in the paper

Abstract

Dementia, a degenerative disease affecting millions globally, is projected to triple by 2050. Early and precise diagnosis is essential for effective treatment and improved quality of life. However, current diagnostic approaches often show inconsistent performance across multi-racial and multi-ethnic groups, raising concerns about fairness and clinical reliability. This study investigates performance discrepancies in dementia classification among 6584 Non-Hispanic White, 1263 Non-Hispanic African American, and 713 Hispanic White populations. We observed significant cross-group bias, particularly when models trained on one group are tested on another. To address this, we evaluated RegAlign, a few-shot domain adaptation objective that combines source-side focal learning, target-side class-weighted supervision, and class-conditional alignment to improve adaptation to underrepresented populations. Our results show that this approach substantially reduces inter-group performance gaps, especially between Non-Hispanic White and Hispanic populations. Here, we show the importance of fairness-aware learning strategies and diverse training data for improving the accuracy and equity of MRI-based dementia classification.

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 8 matches between paragraphs and lines of code.

UTHSCSA-NAL/explainable_and_fair_ML_for_ADRD

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 6263968edb231adf6d1ec76f65605617727fa923, 21 January 2026
Languages: Python (14), Shell (1)
Size: 24 files, 15 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, environment (pyproject.toml)
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (10 files), pandas (7 files), scikit-learn (7 files), XGBoost (3 files), SHAP (2 files), PyTorch (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
16 files

Code availability

The code supporting this study is available at: https://github.com/UTHSCSA-NAL/explainable_and_fair_ML_for_ADRD.

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

The minimum dataset used in this study was obtained from four third-party clinical research resources: the National Alzheimer’s Coordinating Center (NACC), the Standardized Centralized Alzheimer’s & Related Dementias Neuroimaging initiative (SCAN), the Alzheimer’s Disease Neuroimaging Initiative (ADNI), and the South Texas Alzheimer’s Disease Research Center (STAC). Because these are human-participant datasets, access is restricted to protect participant privacy and is governed by each resource’s data-use procedures and agreements. NACC clinical and related research data are available to qualified researchers worldwide through the NACC Data Request Process (https://www.naccdata.org/data-request-process). Investigators must submit a data request describing the proposed project and accept the NACC Data Use Agreement. It will acknowledge requests within three business days and may request clarification if needed. Approved datasets are provided through secure download mechanisms maintained by NACC. SCAN harmonized neuroimaging-derived data, including summary, quality-control, and analysis variables, and defaced SCAN images are available through NACC’s data request system (https://scan.naccdata.org/; see also NACC Data Request Process at https://www.naccdata.org/data-request-process). Access is restricted because the resource contains human neuroimaging data and is subject to de-identification, quality-control, and data-use requirements. All submitted image data may be immediately available because additional processing steps, including defacing and QC, may be required. ADNI data are available to qualified investigators through the Laboratory of Neuro Imaging (LONI) Image and Data Archive (IDA) after acceptance of the ADNI Data Use Agreement and submission of an online application describing the investigator’s institutional affiliation and proposed data use (https://adni.loni.usc.edu/data-samples/adni-data/). Applications are generally reviewed within approximately two weeks. Approved users receive login credentials to access and download the requested data through IDA. To access STAC data, investigators and research groups may request patients, control subjects, biospecimens, and data, including through its Population Neuroscience request process (https://stxadrc.org/for-scientists/). Access is therefore restricted and managed directly by STAC. Researchers seeking access should submit a request through the STAC scientist portal. The authors do not have permission to redistribute these third-party datasets. Researchers must obtain access directly from the corresponding data custodians under the conditions described above. Source data are provided with this paper.

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 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 15 authors, 3 keywords, 11 MeSH terms, 2 funders, 41 references.

Cite

This paper

Ho, N.-H., Charisis, S., Honnorat, N., Brandigampala, S. R., Wang, D., Heckbert, S. R., Fox, P. T., Martinez, D., Wang, D. H., Hughes, T. M., Archer, D. B., Hohman, T. J., Seshadri, S., Davatzikos, C., & Habes, M. (2026). Advancing fair and explainable machine learning for neuroimaging dementia pattern classification in multi-racial and multi-ethnic populations. Nature communications, 17(1), 8026. https://doi.org/10.1038/s41467-026-74515-w

BibTeX

@article{ho2026advancing,
author = {Ho, Ngoc-Huynh and Charisis, Sokratis and Honnorat, Nicolas and Brandigampala, Sachintha Ransara and Wang, Di and Heckbert, Susan R and Fox, Peter T and Martinez, David and Wang, David H and Hughes, Timothy M and Archer, Derek B and Hohman, Timothy J and Seshadri, Sudha and Davatzikos, Christos and Habes, Mohamad},
title = {{Advancing fair and explainable machine learning for neuroimaging dementia pattern classification in multi-racial and multi-ethnic populations}},
journal = {Nature communications},
year = {2026},
month = jun,
volume = {17},
number = {1},
pages = {8026},
publisher = {Nature Publishing Group},
issn = {2041-1723},
doi = {10.1038/s41467-026-74515-w},
url = {https://doi.org/10.1038/s41467-026-74515-w},
pmid = {42362543},
pmcid = {PMC13454470}
}

RIS

TY - JOUR
AU - Ho, Ngoc-Huynh
AU - Charisis, Sokratis
AU - Honnorat, Nicolas
AU - Brandigampala, Sachintha Ransara
AU - Wang, Di
AU - Heckbert, Susan R
AU - Fox, Peter T
AU - Martinez, David
AU - Wang, David H
AU - Hughes, Timothy M
AU - Archer, Derek B
AU - Hohman, Timothy J
AU - Seshadri, Sudha
AU - Davatzikos, Christos
AU - Habes, Mohamad
TI - Advancing fair and explainable machine learning for neuroimaging dementia pattern classification in multi-racial and multi-ethnic populations
T2 - Nature communications
J2 - Nat Commun
PY - 2026
DA - 2026/06/26
VL - 17
IS - 1
SP - 8026
SN - 2041-1723
PB - Nature Publishing Group
DO - 10.1038/s41467-026-74515-w
UR - https://doi.org/10.1038/s41467-026-74515-w
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41467-026-74515-w",
"type": "article-journal",
"title": "Advancing fair and explainable machine learning for neuroimaging dementia pattern classification in multi-racial and multi-ethnic populations",
"container-title": "Nature communications",
"author": [
{
"family": "Ho",
"given": "Ngoc-Huynh"
},
{
"family": "Charisis",
"given": "Sokratis"
},
{
"family": "Honnorat",
"given": "Nicolas"
},
{
"family": "Brandigampala",
"given": "Sachintha Ransara"
},
{
"family": "Wang",
"given": "Di"
},
{
"family": "Heckbert",
"given": "Susan R"
},
{
"family": "Fox",
"given": "Peter T"
},
{
"family": "Martinez",
"given": "David"
},
{
"family": "Wang",
"given": "David H"
},
{
"family": "Hughes",
"given": "Timothy M"
},
{
"family": "Archer",
"given": "Derek B"
},
{
"family": "Hohman",
"given": "Timothy J"
},
{
"family": "Seshadri",
"given": "Sudha"
},
{
"family": "Davatzikos",
"given": "Christos"
},
{
"family": "Habes",
"given": "Mohamad"
}
],
"container-title-short": "Nat Commun",
"volume": "17",
"issue": "1",
"page": "8026",
"DOI": "10.1038/s41467-026-74515-w",
"PMID": "42362543",
"PMCID": "PMC13454470",
"ISSN": "2041-1723",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41467-026-74515-w",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
26
]
]
}
}

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.1016/j.eclinm.2026.104143 [code]
Intensive versus standard blood pressure control and overall brain small vessel disease burden: a post-hoc analysis of the SPRINT randomized clinical trial.
Journal: EClinicalMedicine
In common: structural MRI / diffusion, 2 references, 2 authors
[2] doi:10.1038/s41467-026-72091-7 [code]
Coupled cross-sectional and longitudinal non-negative matrix factorization reveals dominant brain aging trajectories in 48,949 individuals.
Journal: Nature communications
In common: scikit-learn, pandas, NumPy, Alzheimer's / dementia, structural MRI / diffusion, 3 references, author Christos Davatzikos
[3] doi:10.1038/s41467-026-71555-0 [code]
A deep representation learning model to predict response to vagus nerve stimulation.
Journal: Nature communications
In common: SHAP, XGBoost, PyTorch, 3 other tools, structural MRI / diffusion, 1 reference
[4] doi:10.1038/s41467-026-71682-8 [code]
GWAS meta-analysis of cerebrospinal fluid Alzheimer's biomarkers reveals loci regulating lipids, brain volume and autophagy.
Journal: Nature communications
In common: scikit-learn, pandas, NumPy, Alzheimer's / dementia, structural MRI / diffusion, author Timothy Hohman
[5] doi:10.1038/s41467-026-76837-1 [code]
Drug screen and machine learning predict neuroprotective agents in a preclinical human model of childhood dementia.
Journal: Nature communications
In common: SHAP, XGBoost, PyTorch, 3 other tools, Alzheimer's / dementia
[6] 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: SHAP, XGBoost, scikit-learn, 2 other tools, structural MRI / diffusion, 1 reference
[7] doi:10.1038/s41586-026-10454-2 [code]
White matter micro- and macrostructure brain charts for the human lifespan.
Journal: Nature
In common: pandas, NumPy, 1 reference, author Timothy Hohman
[8] doi:10.1038/s41591-026-04485-5 [code]
Blood-based circular RNAs for early diagnosis of Alzheimer's disease.
Journal: Nature medicine
In common: NumPy, Alzheimer's / dementia, 1 reference, author Timothy Hohman
[9] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: SHAP, XGBoost, PyTorch, 3 other tools
[10] doi:10.1371/journal.pcbi.1014615 [code]
Toward reliable machine learning models for neural circuit inference: A diagnostic study of CNNs on spike trains.
Journal: PLoS computational biology
In common: SHAP, XGBoost, PyTorch, 3 other tools

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.