OSCR

Improved Multiscale Structural Mapping with Supervertex Vision Transformer for the Detection of Alzheimer's Disease Neurodegeneration.

Code ↔ Paper

1 match 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 1 match
  1. [1] § Material and Methods › AD versus CN Classification ↔ scripts/train-sv-vit.py, lines 194–264 · score 0.80 · cosine annealing learning, AdamW, scheduler, optimized, loss, SV ViT

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 · 694 lines · 24 KB · no license · 1 match

  1. import os
  2. import sys
  3. import argparse
  4. import yaml
  5. from collections import defaultdict
  6. from copy import deepcopy
  7. from sklearn.model_selection import StratifiedKFold, train_test_split
  8. # Set PYTHONPATH
  9. project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
  10. if project_root not in sys.path:
  11. sys.path.append(project_root)
  12. print(f"✅ sys.path: {project_root}")
  13. parser = argparse.ArgumentParser()
  14. parser.add_argument("--config", type=str, required=True, help="Config filename without extension (e.g. All-0-post)")
  15. parser.add_argument(
  16. "--project",
  17. type=str,
  18. default="renewal",
  19. help="Subdirectory under ../config/ where config.yaml is located (default: re)"
  20. )
  21. args = parser.parse_args()
  22. # config path
  23. config_path = os.path.abspath(f"../config/{args.project}/{args.config}.yaml")
  24. print(f"📄 불러오는 config 파일: {config_path}")
  25. # load config
  26. with open(config_path, "r") as f:
  27. CONFIG = yaml.safe_load(f)
  28. # Patch offset for triangle indices, default to 0 if not set
  29. patch_offset = int(CONFIG.get("PATCH_OFFSET", 0))
  30. os.environ["CUDA_VISIBLE_DEVICES"] = str(CONFIG.get("CUDA_DEVICE", 0))
  31. import numpy as np
  32. import pandas as pd
  33. import matplotlib.pyplot as plt
  34. import nibabel as nib
  35. import gc
  36. import random
  37. import re
  38. import wandb
  39. import torch
  40. import torch.nn as nn
  41. import torch.optim as optim
  42. from torch.utils.data import DataLoader, TensorDataset
  43. from torch.optim.lr_scheduler import CosineAnnealingLR
  44. from sklearn.metrics import (
  45. roc_auc_score,
  46. average_precision_score,
  47. roc_curve,
  48. precision_recall_curve,
  49. )
  50. from sklearn.model_selection import StratifiedKFold, train_test_split
  51. # SurfViT 관련 import
  52. from models.sit import SiT
  53. # 시드 고정
  54. def set_seed(seed=42):
  55. random.seed(seed)
  56. np.random.seed(seed)
  57. torch.manual_seed(seed)
  58. torch.cuda.manual_seed_all(seed)
  59. torch.backends.cudnn.deterministic = True
  60. torch.backends.cudnn.benchmark = False
  61. set_seed(CONFIG.get("SEED", 42))
  62. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  63. print("Using device:", device)
  64. base_path = CONFIG["BASE_PATH"]
  65. result_path = CONFIG["RESULT_PATH"]
  66. info_path = CONFIG["INFO_PATH"]
  67. datasets = CONFIG["DATASETS"]
  68. candidates = CONFIG["CANDIDATES"]
  69. triangle_csv_path = CONFIG["TRIANGLE_CSV_PATH"]
  70. is_smoothed = bool(CONFIG["ISSMOOTHED"])
  71. def load_mgh(file_path):
  72. return nib.load(file_path).get_fdata()
  73. def save_mgh(data_array, file_path, reference_mgh):
  74. img = nib.MGHImage(data_array, affine=nib.load(reference_mgh).affine)
  75. nib.save(img, file_path)
  76. # Remove save_dataset_fold_info (now handled in process_candidate)
  77. def load_and_preprocess_surface_data(dataset, candidate):
  78. df = pd.read_csv(os.path.join(info_path))
  79. df = df[df["Split"] == dataset]
  80. labels = (
  81. df["Group"]
  82. .map({"CN": 0, "AD": 1})
  83. .values
  84. )
  85. filename = "_concat"
  86. if is_smoothed:
  87. filename += "_smoothed"
  88. data_lh = load_mgh(
  89. os.path.join(result_path, dataset, candidate, f"lh{filename}.mgh")
  90. ).reshape(len(labels), -1, 1)
  91. data_rh = load_mgh(
  92. os.path.join(result_path, dataset, candidate, f"rh{filename}.mgh")
  93. ).reshape(len(labels), -1, 1)
  94. triangle_indices = load_triangle_indices(triangle_csv_path)
  95. #print('triangle indices shape:', triangle_indices.shape)
  96. # hemi-wise split (LH/RH)
  97. num_patches_total = triangle_indices.shape[0]
  98. half_patches = num_patches_total // 2 # 2560 → 1280
  99. triangle_indices_lh = triangle_indices[:half_patches, :]
  100. triangle_indices_rh = triangle_indices[half_patches:, :] # (1280, 153)
  101. # LH/RH pooling
  102. data_lh = reshape_vertex_to_surface(data_lh, triangle_indices_lh, offset=0)
  103. data_rh = reshape_vertex_to_surface(data_rh, triangle_indices_rh, offset=0)
  104. # (N, 2, 1280)
  105. data_concat = np.stack([data_lh.squeeze(-1), data_rh.squeeze(-1)], axis=1)
  106. return data_concat, labels
  107. def load_triangle_indices(csv_path):
  108. return pd.read_csv(csv_path).values.astype(
  109. np.int64
  110. ).T
  111. def reshape_vertex_to_surface(x, triangle_indices, offset=0):
  112. if isinstance(x, np.ndarray):
  113. x = torch.from_numpy(x).float()
  114. if x.ndim == 3:
  115. x = x.squeeze(-1)
  116. triangle_indices_tensor = torch.from_numpy(triangle_indices).long()
  117. if (triangle_indices_tensor >= 163842).sum():
  118. offset = -163842
  119. triangle_indices_tensor = triangle_indices_tensor + offset
  120. x_tri = x[:, triangle_indices_tensor].mean(dim=2)
  121. x_tri = x_tri.unsqueeze(-1)
  122. return x_tri.numpy()
  123. class SVViT(nn.Module):
  124. def __init__(self):
  125. super().__init__()
  126. self.backbone = SiT(**CONFIG["MODEL_CONFIG"])
  127. def forward(self, x):
  128. x = x.reshape(x.shape[0], 1, -1, 1)
  129. return self.backbone(x).view(-1)
  130. class EarlyStopping:
  131. def __init__(self, patience=CONFIG.get("PATIENCE", 10), verbose=False):
  132. self.patience = patience
  133. self.verbose = verbose
  134. self.counter = 0
  135. self.best_score = None
  136. self.early_stop = False
  137. self.best_model = None
  138. def __call__(self, val_loss, model):
  139. score = -val_loss전
  140. if self.best_score is None:
  141. self.best_score = score
  142. self.best_model = model.state_dict()
  143. elif score < self.best_score:
  144. self.counter += 1
  145. if self.counter >= self.patience:
  146. self.early_stop = True
  147. else:
  148. self.best_score = score
  149. self.best_model = model.state_dict()
  150. self.counter = 0
  151. def train_sit_model(X_train, y_train, X_valid, y_valid, model_cls):
  152. batch_size = CONFIG["DATALOADER_CONFIG"]["batch_size"]
  153. epochs = CONFIG["EPOCHS"]
  154. lr = float(CONFIG["OPTIMIZER_CONFIG"]["lr"])
  155. weight_decay = float(CONFIG["OPTIMIZER_CONFIG"]["weight_decay"])
  156. y_train_tensor = torch.tensor(y_train, dtype=torch.float32)
  157. y_valid_tensor = torch.tensor(y_valid, dtype=torch.float32)
  158. train_dataset = TensorDataset(
  159. torch.tensor(X_train, dtype=torch.float32), y_train_tensor
  160. )
  161. valid_dataset = TensorDataset(
  162. torch.tensor(X_valid, dtype=torch.float32), y_valid_tensor
  163. )
  164. train_loader = DataLoader(
  165. train_dataset, batch_size=batch_size, shuffle=True, num_workers=4
  166. )
  167. valid_loader = DataLoader(
  168. valid_dataset, batch_size=batch_size, shuffle=False, num_workers=4
  169. )
  170. model = model_cls().to(device)
  171. criterion = nn.BCEWithLogitsLoss()
  172. optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=weight_decay)
  173. scheduler = CosineAnnealingLR(optimizer, T_max=epochs)
  174. early_stopping = EarlyStopping(patience=CONFIG.get("PATIENCE", 10), verbose=True)
  175. for epoch in range(1, epochs + 1):
  176. model.train()
  177. train_loss = 0.0
  178. for batch_X, batch_y in train_loader:
  179. batch_X, batch_y = batch_X.to(device), batch_y.to(device)
  180. out = model(batch_X)
  181. loss = criterion(out, batch_y)
  182. optimizer.zero_grad()
  183. loss.backward()
  184. optimizer.step()
  185. train_loss += loss.item()
  186. train_loss /= len(train_loader)
  187. model.eval()
  188. valid_loss = 0.0
  189. with torch.no_grad():
  190. for batch_X, batch_y in valid_loader:
  191. batch_X, batch_y = batch_X.to(device), batch_y.to(device)
  192. out = model(batch_X)
  193. loss = criterion(out, batch_y)
  194. valid_loss += loss.item()
  195. valid_loss /= len(valid_loader)
  196. scheduler.step()
  197. wandb.log({"train_loss": train_loss, "valid_loss": valid_loss, "epoch": epoch})
  198. if epoch % 5 == 0 or epoch == 1:
  199. print(
  200. f"[Epoch {epoch}/{epochs}] Train Loss: {train_loss:.4f} | Valid Loss: {valid_loss:.4f}"
  201. )
  202. early_stopping(val_loss=valid_loss, model=model)
  203. if early_stopping.early_stop:
  204. print(f"Early stopping triggered at epoch {epoch}")
  205. model.load_state_dict(early_stopping.best_model)
  206. break
  207. return model
  208. def test_classifier_sit(X, y, Model, batch_size=48):
  209. # Early check for empty y or y_scores
  210. if len(y) == 0 or len(X) == 0:
  211. print(
  212. "⚠️ Warning: No samples provided for test_classifier_sit. Skipping evaluation."
  213. )
  214. return 0.0, 0.0, np.array([])
  215. dataset = TensorDataset(torch.tensor(X, dtype=torch.float32))
  216. loader = DataLoader(dataset, batch_size=batch_size, shuffle=False, num_workers=4)
  217. Model.eval()
  218. y_scores = []
  219. with torch.no_grad():
  220. for batch in loader:
  221. batch_data = batch[0].to(device)
  222. out = Model(batch_data)
  223. score = torch.sigmoid(out).view(-1).cpu().numpy()
  224. y_scores.extend(score)
  225. y_scores = np.array(y_scores)
  226. if len(y) == 0 or len(y_scores) == 0:
  227. print(
  228. "⚠️ Warning: No samples provided for test_classifier_sit. Skipping evaluation."
  229. )
  230. return 0.0, 0.0, np.array([])
  231. # Try-except for AUROC/AUPRC in case of only one class in y
  232. try:
  233. auroc = roc_auc_score(y, y_scores)
  234. except ValueError:
  235. print("⚠️ AUROC calculation failed due to only one class present in y.")
  236. auroc = 0.0
  237. try:
  238. auprc = average_precision_score(y, y_scores)
  239. except ValueError:
  240. print("⚠️ AUPRC calculation failed due to only one class present in y.")
  241. auprc = 0.0
  242. return auroc, auprc, y_scores
  243. # New process_candidate with new split structure and metrics/fold saving
  244. def generate_fold_assignment(labels, cv_folds, seed=42):
  245. from sklearn.model_selection import StratifiedKFold
  246. import numpy as np
  247. skf = StratifiedKFold(n_splits=cv_folds, shuffle=True, random_state=seed)
  248. fold_assignment = np.zeros(len(labels), dtype=int)
  249. # Use stratification for class balance
  250. for fold_idx, (_, val_idx) in enumerate(skf.split(np.zeros(len(labels)), labels)):
  251. fold_assignment[val_idx] = fold_idx
  252. return fold_assignment
  253. def process_candidate(dataset, candidate):
  254. from collections import defaultdict
  255. from copy import deepcopy
  256. from sklearn.model_selection import StratifiedKFold, train_test_split
  257. # Load data and labels
  258. data_concat, labels = load_and_preprocess_surface_data(dataset, candidate)
  259. labels = np.asarray(labels)
  260. candidate_dir = os.path.join(base_path, dataset, candidate)
  261. os.makedirs(candidate_dir, exist_ok=True)
  262. dataset_dir = os.path.join(base_path, dataset)
  263. os.makedirs(dataset_dir, exist_ok=True)
  264. cv_folds = CONFIG["TRAIN_SPLIT"]["cv_folds"]
  265. # Validate cv_folds value
  266. if not isinstance(cv_folds, int) or cv_folds <= 1:
  267. raise ValueError(
  268. f"Invalid cv_folds={cv_folds} in config. Must be an integer > 1."
  269. )
  270. valid_size = CONFIG["TRAIN_SPLIT"]["valid_size"]
  271. seed = CONFIG.get("SEED", 42)
  272. # Always generate fold_assignment.npy using the new function
  273. fold_file = os.path.join(dataset_dir, "fold_assignment.npy")
  274. fold_assignment = generate_fold_assignment(labels, cv_folds, seed=CONFIG["SEED"])
  275. np.save(fold_file, fold_assignment)
  276. # For summary/metrics
  277. fold_metrics = []
  278. pred_all = np.full(len(labels), np.nan)
  279. # For test predictions, will collect per fold and average at the end
  280. test_preds_accum = []
  281. # For accumulating ROC/PR data for all folds, per split
  282. plot_data = {"train": [], "valid": [], "test": []}
  283. # To accumulate test predictions for each fold for ensemble (mean)
  284. test_pred_folds = []
  285. # For each fold (cross-validation)
  286. for fold in range(cv_folds):
  287. # For this fold: test is fold_assignment == fold
  288. test_idx = np.where(fold_assignment == fold)[0]
  289. trainvalid_idx = np.where(fold_assignment != fold)[0]
  290. # Split trainvalid_idx into train and valid, stratified
  291. X_trainval = data_concat[trainvalid_idx]
  292. y_trainval = labels[trainvalid_idx]
  293. train_idx_sub, valid_idx_sub = train_test_split(
  294. np.arange(len(trainvalid_idx)),
  295. test_size=valid_size,
  296. stratify=y_trainval,
  297. random_state=seed + fold, # ensure reproducibility but different per fold
  298. )
  299. train_idx = trainvalid_idx[train_idx_sub]
  300. valid_idx = trainvalid_idx[valid_idx_sub]
  301. X_train, y_train = data_concat[train_idx], labels[train_idx]
  302. X_valid, y_valid = data_concat[valid_idx], labels[valid_idx]
  303. X_test, y_test = data_concat[test_idx], labels[test_idx]
  304. # For logging
  305. wandb.init(
  306. project=f"hypo2_ad_classification-{args.config}",
  307. name=f"{dataset}-{candidate}-fold{fold}",
  308. config={
  309. "epochs": CONFIG["EPOCHS"],
  310. "batch_size": CONFIG["DATALOADER_CONFIG"]["batch_size"],
  311. "fold": fold,
  312. "dataset": dataset,
  313. "candidate": candidate,
  314. },
  315. )
  316. model = train_sit_model(X_train, y_train, X_valid, y_valid, SVViT)
  317. # Train/Valid predictions
  318. auroc_train, auprc_train, y_scores_train = test_classifier_sit(
  319. X_train,
  320. y_train,
  321. model,
  322. batch_size=CONFIG["DATALOADER_CONFIG"]["batch_size"],
  323. )
  324. auroc_valid, auprc_valid, y_scores_valid = test_classifier_sit(
  325. X_valid,
  326. y_valid,
  327. model,
  328. batch_size=CONFIG["DATALOADER_CONFIG"]["batch_size"],
  329. )
  330. # Test predictions (current fold's test set)
  331. auroc_test, auprc_test, y_scores_test = test_classifier_sit(
  332. X_test,
  333. y_test,
  334. model,
  335. batch_size=CONFIG["DATALOADER_CONFIG"]["batch_size"],
  336. )
  337. # Save predictions for valid (for SiT_predictions.csv)
  338. pred_all[valid_idx] = y_scores_valid
  339. # For test set, accumulate predictions for ensemble
  340. test_pred_folds.append((test_idx, y_scores_test))
  341. # Save metrics for this fold
  342. fold_metrics.append(
  343. {
  344. "fold": fold,
  345. "train_AUROC": round(float(auroc_train), 5),
  346. "train_AUPRC": round(float(auprc_train), 5),
  347. "valid_AUROC": round(float(auroc_valid), 5),
  348. "valid_AUPRC": round(float(auprc_valid), 5),
  349. "test_AUROC": round(float(auroc_test), 5),
  350. "test_AUPRC": round(float(auprc_test), 5),
  351. }
  352. )
  353. print(
  354. f"Fold {fold}: [Train] AUROC {auroc_train:.3f}, AUPRC {auprc_train:.3f} | [Valid] AUROC {auroc_valid:.3f}, AUPRC {auprc_valid:.3f} | [Test] AUROC {auroc_test:.3f}, AUPRC {auprc_test:.3f}"
  355. )
  356. wandb.log(
  357. {
  358. "train_AUROC": auroc_train,
  359. "train_AUPRC": auprc_train,
  360. "valid_AUROC": auroc_valid,
  361. "valid_AUPRC": auprc_valid,
  362. "test_AUROC": auroc_test,
  363. "test_AUPRC": auprc_test,
  364. }
  365. )
  366. # Save model
  367. model_save_path = os.path.join(candidate_dir, f"model_fold{fold}.pt")
  368. torch.save(model.state_dict(), model_save_path)
  369. print(f"Saved model for fold {fold}: {model_save_path}")
  370. # --- ROC/PR curve visualization
  371. fpr_train, tpr_train, _ = roc_curve(y_train, y_scores_train)
  372. fpr_valid, tpr_valid, _ = roc_curve(y_valid, y_scores_valid)
  373. fpr_test, tpr_test, _ = roc_curve(y_test, y_scores_test)
  374. precision_train, recall_train, _ = precision_recall_curve(
  375. y_train, y_scores_train
  376. )
  377. precision_valid, recall_valid, _ = precision_recall_curve(
  378. y_valid, y_scores_valid
  379. )
  380. precision_test, recall_test, _ = precision_recall_curve(y_test, y_scores_test)
  381. # Fold-wise plot (ROC/PR)
  382. fig, axes = plt.subplots(1, 2, figsize=(10, 5))
  383. # ROC
  384. axes[0].plot(
  385. fpr_train,
  386. tpr_train,
  387. color="blue",
  388. lw=2,
  389. label=f"Train AUROC = {auroc_train:.3f}",
  390. )
  391. axes[0].plot(
  392. fpr_valid,
  393. tpr_valid,
  394. color="orange",
  395. lw=2,
  396. label=f"Valid AUROC = {auroc_valid:.3f}",
  397. )
  398. axes[0].plot(
  399. fpr_test,
  400. tpr_test,
  401. color="green",
  402. lw=2,
  403. label=f"Test AUROC = {auroc_test:.3f}",
  404. )
  405. axes[0].plot([0, 1], [0, 1], color="gray", linestyle="--", label="Random")
  406. axes[0].set_xlabel("FPR")
  407. axes[0].set_ylabel("TPR")
  408. axes[0].set_title(f"Fold {fold} ROC curve")
  409. axes[0].legend(loc="lower right")
  410. # PR (오른쪽)
  411. axes[1].plot(
  412. recall_train,
  413. precision_train,
  414. color="blue",
  415. lw=2,
  416. label=f"Train AUPRC = {auprc_train:.3f}",
  417. )
  418. axes[1].plot(
  419. recall_valid,
  420. precision_valid,
  421. color="orange",
  422. lw=2,
  423. label=f"Valid AUPRC = {auprc_valid:.3f}",
  424. )
  425. axes[1].plot(
  426. recall_test,
  427. precision_test,
  428. color="green",
  429. lw=2,
  430. label=f"Test AUPRC = {auprc_test:.3f}",
  431. )
  432. # PR 대각선 기준선 (1-x)
  433. axes[1].plot([0, 1], [1, 0], color="gray", linestyle="--", label="Random")
  434. axes[1].set_xlabel("Recall")
  435. axes[1].set_ylabel("Precision")
  436. axes[1].set_title(f"Fold {fold} PR curve")
  437. axes[1].legend(loc="lower left")
  438. plt.tight_layout()
  439. fold_graph_file = os.path.join(candidate_dir, f"SiT_fold{fold}_AUROC_AUPRC.png")
  440. plt.savefig(fold_graph_file)
  441. plt.close()
  442. print(f"Saved fold {fold} ROC/PR 통합 그림: {fold_graph_file}")
  443. # Instead of per-fold individual plots, accumulate data per split for all folds
  444. plot_data["train"].append(
  445. (
  446. fpr_train,
  447. tpr_train,
  448. recall_train,
  449. precision_train,
  450. auroc_train,
  451. auprc_train,
  452. fold,
  453. )
  454. )
  455. plot_data["valid"].append(
  456. (
  457. fpr_valid,
  458. tpr_valid,
  459. recall_valid,
  460. precision_valid,
  461. auroc_valid,
  462. auprc_valid,
  463. fold,
  464. )
  465. )
  466. plot_data["test"].append(
  467. (
  468. fpr_test,
  469. tpr_test,
  470. recall_test,
  471. precision_test,
  472. auroc_test,
  473. auprc_test,
  474. fold,
  475. )
  476. )
  477. wandb.finish()
  478. del model
  479. torch.cuda.empty_cache()
  480. gc.collect()
  481. # After all folds: plot 1 figure per split (train/valid/test) with all folds as lines
  482. for split in ["train", "valid", "test"]:
  483. fig, axes = plt.subplots(1, 2, figsize=(10, 5))
  484. # ROC curve
  485. for i, (fpr, tpr, recall, precision, auroc, auprc, fold) in enumerate(
  486. plot_data[split]
  487. ):
  488. axes[0].plot(
  489. fpr,
  490. tpr,
  491. lw=2,
  492. label=f"Fold {fold} AUROC = {auroc:.3f}",
  493. )
  494. axes[0].plot([0, 1], [0, 1], color="gray", linestyle="--", label="Random")
  495. axes[0].set_xlabel("FPR")
  496. axes[0].set_ylabel("TPR")
  497. axes[0].set_title(f"All Folds {split.capitalize()} ROC curve")
  498. axes[0].legend(loc="lower right", fontsize=8)
  499. # PR curve
  500. for i, (fpr, tpr, recall, precision, auroc, auprc, fold) in enumerate(
  501. plot_data[split]
  502. ):
  503. axes[1].plot(
  504. recall,
  505. precision,
  506. lw=2,
  507. label=f"Fold {fold} AUPRC = {auprc:.3f}",
  508. )
  509. axes[1].plot([0, 1], [1, 0], color="gray", linestyle="--", label="Random")
  510. axes[1].set_xlabel("Recall")
  511. axes[1].set_ylabel("Precision")
  512. axes[1].set_title(f"All Folds {split.capitalize()} PR curve")
  513. axes[1].legend(loc="lower left", fontsize=8)
  514. plt.tight_layout()
  515. allfold_file = os.path.join(candidate_dir, f"SiT_allfold_{split}.png")
  516. plt.savefig(allfold_file)
  517. plt.close()
  518. print(f"Saved all folds {split} ROC/PR figure: {allfold_file}")
  519. # Save fold_metrics.csv
  520. fold_metrics_df = pd.DataFrame(fold_metrics)
  521. fold_metrics_csv = os.path.join(candidate_dir, "fold_metrics.csv")
  522. fold_metrics_df.to_csv(fold_metrics_csv, index=False)
  523. print(f"Saved fold metrics to {fold_metrics_csv}")
  524. # Save metric_summary.csv (mean/std of each metric)
  525. metric_summary = {}
  526. for metric in [
  527. "train_AUROC",
  528. "train_AUPRC",
  529. "valid_AUROC",
  530. "valid_AUPRC",
  531. "test_AUROC",
  532. "test_AUPRC",
  533. ]:
  534. vals = fold_metrics_df[metric].values.astype(float)
  535. metric_summary[f"{metric}_mean"] = round(np.mean(vals), 5)
  536. metric_summary[f"{metric}_std"] = round(np.std(vals), 5)
  537. metric_summary_df = pd.DataFrame([metric_summary])
  538. metric_summary_csv = os.path.join(candidate_dir, "metric_summary.csv")
  539. metric_summary_df.to_csv(metric_summary_csv, index=False)
  540. print(f"Saved metric summary to {metric_summary_csv}")
  541. # Save SiT_predictions.csv (for all samples: valid predictions for train, test predictions for test)
  542. sit_pred_df = pd.DataFrame(
  543. {
  544. "sample_index": np.arange(len(labels)),
  545. "true_label": labels,
  546. "predicted_prob": pred_all,
  547. "predicted_label": (pred_all >= 0.5).astype(float),
  548. "fold": fold_assignment,
  549. }
  550. )
  551. # For each fold, fill in test predictions for that fold's test indices
  552. for test_idx, y_scores_test in test_pred_folds:
  553. sit_pred_df.loc[test_idx, "predicted_prob"] = y_scores_test
  554. sit_pred_df.loc[test_idx, "predicted_label"] = (
  555. np.array(y_scores_test) >= 0.5
  556. ).astype(float)
  557. sit_pred_df.to_csv(
  558. os.path.join(candidate_dir, "SiT_predictions.csv"),
  559. index=False,
  560. )
  561. print(f"Saved predictions to {os.path.join(candidate_dir, 'SiT_predictions.csv')}")
  562. # Save candidate_results.csv (mean test AUROC/AUPRC)
  563. result_df = pd.DataFrame(
  564. [
  565. {
  566. "dataset": dataset,
  567. "candidate": candidate,
  568. "test_AUROC_mean": metric_summary["test_AUROC_mean"],
  569. "test_AUROC_std": metric_summary["test_AUROC_std"],
  570. "test_AUPRC_mean": metric_summary["test_AUPRC_mean"],
  571. "test_AUPRC_std": metric_summary["test_AUPRC_std"],
  572. }
  573. ]
  574. )
  575. result_df.to_csv(os.path.join(candidate_dir, "candidate_results.csv"), index=False)
  576. print(f"Saved result summary to {candidate_dir}/candidate_results.csv")
  577. def main():
  578. for dataset in datasets:
  579. for candidate in candidates:
  580. if os.path.exists(os.path.join(base_path, dataset, candidate)):
  581. print(f"{dataset} {candidate} exist, skip.")
  582. continue
  583. print(f"Processing candidate {candidate} for dataset {dataset}...")
  584. process_candidate(dataset, candidate)
  585. overall_results = []
  586. for dataset in datasets:
  587. for candidate in candidates:
  588. candidate_dir = os.path.join(base_path, dataset, candidate)
  589. result_file = os.path.join(candidate_dir, "candidate_results.csv")
  590. if os.path.exists(result_file):
  591. df = pd.read_csv(result_file)
  592. # Ensure candidate column is zero-padded string
  593. df["candidate"] = df["candidate"].astype(str).str.zfill(6)
  594. overall_results.append(df)
  595. if overall_results:
  596. overall_df = pd.concat(overall_results, axis=0, ignore_index=True)
  597. overall_csv = os.path.join(base_path, "overall_results.csv")
  598. overall_df.to_csv(overall_csv, index=False)
  599. print(f"\nOverall results saved to {overall_csv}")
  600. if __name__ == "__main__":
  601. main()

train-sv-vit.py at commit d22576d, no license · at the source

Overview

Authors: Geonwoo Baek1, David H. Salat2,3,4, Ikbeom Jang1, for the Alzheimer's Disease Neuroimaging Initiative
ORCID iDs: Geonwoo Baek
  1. Department of Computer Science and Engineering Hankuk University of Foreign Studies Seoul Republic of Korea
  2. Athinoula A. Martinos Center for Biomedical Imaging, Department of Radiology Massachusetts General Hospital Charlestown Massachusetts USA
  3. Department of Radiology Harvard Medical School Boston Massachusetts USA
  4. Neuroimaging Research for Veterans (NeRVe) Center, VA Boston Healthcare System Boston Massachusetts USA
Journal: Human brain mapping, volume 47, issue 8, article e70548
Dates: received 12 December 2025; accepted 5 May 2026; published online 2 June 2026; in print June 2026
Type: Research article · Language: English
License: CC BY-NC
Identifiers: DOI 10.1002/hbm.70548 · PMID 42227644 · PMCID PMC13240337 · OpenAlex W7163157794
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism), Alzheimer's / dementia (population), clinical / translational (subfield)
Methods: Spectral & time-frequency, Statistics, Machine learning, fMRI & imaging, Smoothing, state filtering, decompositions, Preprocessing
Keywords: Alzheimer's disease, gray‐to‐white matter contrast, imaging biomarker, MRI, multiscale structural mapping, supervertices
MeSH: Alzheimer Disease*, Gray Matter*, Image Interpretation, Computer-Assisted*, Magnetic Resonance Imaging*, White Matter*, Aged, Aged, 80 and over, Female, Humans, Imaging, Three-Dimensional, Male (* major topic)
Topic: Dementia and Cognitive Impairment Research (Psychiatry and Mental health, Medicine), according to OpenAlex
Funding: NIH (R21AG072431); National Research Foundation of Korea (NRF) (RS‐2024‐00455720); Korea National Institute of Health (2024‐ER0407‐00, 2025‐ER0403‐00, 2026‐ER0904‐00); Korea Institute of Science and Technology Information (KSC‐2024‐CRE‐0021, KSC‐2025‐CRE‐0065); Hankuk University of Foreign Studies (2025)
Citations: not cited yet (Europe PMC); 38 references in the paper

Abstract

Alzheimer's disease (AD) confirmation often relies on positron emission tomography (PET) or cerebrospinal fluid (CSF) analysis, which are costly and invasive. Consequently, structural MRI biomarkers such as cortical thickness (CT) are widely used for noninvasive AD screening. Multiscale structural mapping (MSSM) was recently proposed to integrate gray–white matter contrasts (GWCs) with CT from a single T1‐weighted MRI (T1w) scan. Building on this framework, we propose MSSM+, together with surface supervertex mapping (SSVM) and a Supervertex Vision Transformer (SV‐ViT). 3D T1w images from individuals with AD and cognitively normal (CN) controls were analyzed. MSSM+ extends MSSM by incorporating sulcal depth and cortical curvature at the vertex level. SSVM partitions the cortical surface into supervertices (surface patches) that effectively represent inter‐ and intra‐regional spatial relationships. SV‐ViT is a Vision Transformer architecture operating on these supervertices, enabling anatomically informed learning from surface mesh representations. Compared with MSSM, MSSM+ identified more spatially extensive and statistically significant group differences between AD and CN. In AD versus CN classification, MSSM+ achieved a 3%p higher area under the precision–recall curve than MSSM. Vendor‐specific analyses further demonstrated reduced signal variability and consistently improved classification performance across MR manufacturers relative to CT, GWCs, and MSSM. These findings suggest that MSSM+ combined with SV‐ViT is a promising MRI‐based imaging marker for AD detection prior to CSF/PET confirmation.

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

Repositories

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

labhai/MSSMplus

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: d22576d6e1155f8a333e1c3cb76b833d62b46ab9, 11 June 2026
Languages: Python (3), Shell (2)
Size: 10 files, 5 scripts
Software Heritage: not archived
Found in: “Data Availability Statement”
Holds: README, environment (requirements.txt)
Not found: license file, CITATION.cff, tests, continuous integration, documentation
Tools: FreeSurfer (3 files), NiBabel (3 files), NumPy (3 files), pandas (3 files), scikit-learn (3 files), Matplotlib (1 file), PyTorch (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
6 files

adni.loni.usc.edu/wp-content/uploads

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: the link answers
Software Heritage: not checked
Found in: “Data Availability Statement”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 27 September 2026: the link answers (HTTP 200)
  • 27 September 2026: the link answers (HTTP 200)

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

Tracing map

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

What the map holds:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 5 scripts, each with its path and the digest of its content;
  • 1 match 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

Data used in the preparation of this article were obtained from the ADNI (Jack et al. 2008; Jack, Bernstein, Borowski, et al. 2010) and OASIS databases (LaMontagne et al. 2019; Koenig et al. 2020). The ADNI and OASIS data are available at https://adni.loni.usc.edu and https://www.oasis‐brains.org (https://www.oasis-brains.org), respectively. As such, the investigators within the ADNI contributed to the design and implementation of ADNI and/or provided data but did not participate in analysis or writing of this report. A complete listing of ADNI investigators can be found at: http://adni.loni.usc.edu/wp‐content/uploads/how_to_apply/ADNI_Acknowledgement_List.pdf (http://adni.loni.usc.edu/wp-content/uploads/how_to_apply/ADNI_Acknowledgement_List.pdf). The code used in this study is available at https://github.com/labhai/MSSMplus.

Reproduced under the paper's license (CC BY-NC), 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, 4 authors, 6 keywords, 11 MeSH terms, 5 funders, 36 references.

Cite

This paper

Baek, G., Salat, D. H., Jang, I., & for the Alzheimer's Disease Neuroimaging Initiative. (2026). Improved Multiscale Structural Mapping with Supervertex Vision Transformer for the Detection of Alzheimer's Disease Neurodegeneration. Human brain mapping, 47(8), e70548. https://doi.org/10.1002/hbm.70548

BibTeX

@article{baek2026improved,
author = {Baek, Geonwoo and Salat, David H. and Jang, Ikbeom and {for the Alzheimer's Disease Neuroimaging Initiative}},
title = {{Improved Multiscale Structural Mapping with Supervertex Vision Transformer for the Detection of Alzheimer's Disease Neurodegeneration}},
journal = {Human brain mapping},
year = {2026},
month = jun,
volume = {47},
number = {8},
pages = {e70548},
publisher = {Wiley},
issn = {1065-9471},
doi = {10.1002/hbm.70548},
url = {https://doi.org/10.1002/hbm.70548},
pmid = {42227644},
pmcid = {PMC13240337}
}

RIS

TY - JOUR
AU - Baek, Geonwoo
AU - Salat, David H.
AU - Jang, Ikbeom
AU - for the Alzheimer's Disease Neuroimaging Initiative
TI - Improved Multiscale Structural Mapping with Supervertex Vision Transformer for the Detection of Alzheimer's Disease Neurodegeneration
T2 - Human brain mapping
J2 - Hum Brain Mapp
PY - 2026
DA - 2026/06/01
VL - 47
IS - 8
SP - e70548
SN - 1065-9471
PB - Wiley
DO - 10.1002/hbm.70548
UR - https://doi.org/10.1002/hbm.70548
LA - en
ER -

CSL-JSON

{
"id": "10.1002/hbm.70548",
"type": "article-journal",
"title": "Improved Multiscale Structural Mapping with Supervertex Vision Transformer for the Detection of Alzheimer's Disease Neurodegeneration",
"container-title": "Human brain mapping",
"author": [
{
"family": "Baek",
"given": "Geonwoo"
},
{
"family": "Salat",
"given": "David H."
},
{
"family": "Jang",
"given": "Ikbeom"
},
{
"literal": "for the Alzheimer's Disease Neuroimaging Initiative"
}
],
"container-title-short": "Hum Brain Mapp",
"volume": "47",
"issue": "8",
"page": "e70548",
"DOI": "10.1002/hbm.70548",
"PMID": "42227644",
"PMCID": "PMC13240337",
"ISSN": "1065-9471",
"publisher": "Wiley",
"URL": "https://doi.org/10.1002/hbm.70548",
"language": "en",
"issued": {
"date-parts": [
[
2026,
6,
1
]
]
}
}

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.1371/journal.pbio.3003856 [code]
Aging and metabolism contribute separately to brain-body health.
Journal: PLoS biology
In common: FreeSurfer, NiBabel, PyTorch, 4 other tools, structural MRI / diffusion, clinical / translational, 2 references
[2] doi:10.1038/s41598-026-55397-w [code]
Fast surface reconstruction of human brain MRI: benchmarking deep-learning based morphometry tools.
Journal: Scientific reports
In common: FreeSurfer, NiBabel, PyTorch, 3 other tools, structural MRI / diffusion, 3 references
[3] doi:10.1038/s41398-026-04081-8 [code]
Functional system-specific brain aging across the Alzheimer's disease continuum.
Journal: Translational psychiatry
In common: FreeSurfer, NiBabel, scikit-learn, 3 other tools, Alzheimer's / dementia, structural MRI / diffusion, clinical / translational, 2 references
[4] doi:10.64898/2026.08.18.26360725 [code]
Temporal pole blurring in hippocampal sclerosis reflects seizure-disrupted myelination
Journal: medRxiv (preprint)
In common: FreeSurfer, NiBabel, PyTorch, 4 other tools, structural MRI / diffusion, 1 reference
[5] doi:10.1002/alz.71530 [code]
Differential associations of plasma biomarkers with Alzheimer's disease and small vessel disease: A multimodal imaging study.
Journal: Alzheimer's & dementia : the journal of the Alzheimer's Association
In common: NiBabel, PyTorch, pandas, 2 other tools, Alzheimer's / dementia, structural MRI / diffusion, clinical / translational, 2 references
[6] doi:10.3389/frai.2026.1771088 [code]
Few-shot deployment of pretrained MRI transformers in brain imaging tasks.
Journal: Frontiers in artificial intelligence
In common: NiBabel, PyTorch, scikit-learn, 3 other tools, structural MRI / diffusion, 2 references
[7] doi:10.64898/2026.05.06.26352540 [code]
Generating synthetic tau-PET scans in Alzheimer’s disease from MRI, blood biomarkers and demographics with deep learning
Journal: medRxiv (preprint)
In common: NiBabel, scikit-learn, pandas, 2 other tools, Alzheimer's / dementia, structural MRI / diffusion, clinical / translational, 2 references
[8] doi:10.1002/alz.71649 [code]
Postmortem brain MRI reveals differential associations of subcortical and limbic volumes with cortical thinning and neurodegenerative pathologies.
Journal: Alzheimer's & dementia : the journal of the Alzheimer's Association
In common: FreeSurfer, NiBabel, PyTorch, 4 other tools, Alzheimer's / dementia, structural MRI / diffusion
[9] doi:10.1016/j.xcrm.2026.102943 [code]
Parent-of-origin effects in Alzheimer's liability dissociate neurocognitive and cardiovascular traits in at-risk individuals.
Journal: Cell reports. Medicine
In common: FreeSurfer, NiBabel, PyTorch, 4 other tools, Alzheimer's / dementia, clinical / translational
[10] doi:10.1371/journal.pone.0344600 [code]
Robust disease prognosis via diagnostic knowledge preservation: A sequential learning approach.
Journal: PloS one
In common: FreeSurfer, NiBabel, PyTorch, 4 other tools, Alzheimer's / dementia, clinical / translational

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.