OSCR

Trust-gated synthetic EEG augmentation reduces performance drops when generalizing to new patients.

Code ↔ Paper

15 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 15 matches
  1. [1] § Methods ↔ tga/runner.py, lines 224–236 · score 0.73 · geometry aware teacher, logistic regression, transform, classifier, covariance, trained
  2. [2] § Methods ↔ scripts/legacy/run_TGA_pack_motor_bciiv2a_binary.py, lines 10–115 · score 0.66 · BCI IV, motor imagery, 38 Hz, padding, scoring
  3. [3] § Methods ↔ scripts/tga_painmunich_pack_resting_eeg.py, lines 74–149 · score 0.64 · notch filtered, rejected, Artifact, stride, 45 Hz, channel
  4. [4] § Methods ↔ scripts/legacy/run_TGA_plot_calibration.py, lines 8–85 · score 0.61 · threshold probability, Decision curve, reliability, calibration, bins
  5. [5] § Methods ↔ tga/runner.py, lines 76–119 · score 0.61 · subject stratified validation, class balanced, split
  6. [6] § Methods ↔ tga/runner.py, lines 184–193 · score 0.59 · Ledoit Wolf, shrinkage covariance, windows
  7. [7] § Results ↔ tga/runner.py, lines 224–236 · score 0.58 · geometry aware teacher, covariance features, transform, class, training
  8. [8] § Results ↔ scripts/legacy/run_TGA_summarize_paper_packet.py, lines 76–221 · score 0.54 · paired deltas, harm rate, worst, framing, baseline, seed
  9. [9] § Results ↔ scripts/legacy/run_TGA_negative_control_random_gate.py, lines 202–272 · score 0.54 · random gating, ShallowConvNet, EEGNet, matched
  10. [10] § Methods ↔ scripts/legacy/run_TGA_pain_calibration_replay.py, lines 232–295 · score 0.52 · Decision curve, Brier, ECE, threshold, calibration, bins
  11. [11] § Results ↔ scripts/legacy/run_TGA_summarize_subjectlevel.py, lines 41–185 · score 0.52 · paired deltas, harm rate, metrics, framing, baseline, seed
  12. [12] § Results ↔ scripts/legacy/run_TGA_plot_calibration.py, lines 8–85 · score 0.52 · threshold probabilities, Decision curve, calibration, EEGNet, scarcity, gating
  13. [13] § Methods ↔ scripts/legacy/run_TGA_pack_motor_bciiv2a_binary.py, lines 10–115 · score 0.52 · motor imagery, binary, BCI, IV, classes
  14. [14] § Methods ↔ scripts/legacy/run_TGA_generate_pain_pool_vae_ldm.py, lines 526–576 · score 0.52 · X_pool.npy, y_pool.npy, reused, seed, fold
  15. [15] § Methods ↔ scripts/legacy/run_TGA_build_pain_pool_vae_ldm.py, lines 1296–1324 · score 0.51 · X_pool.npy, y_pool.npy, reused, class

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 · 540 lines · 18 KB · MIT · 4 matches

  1. """TGA runner utilities (reviewer-proof).
  2. This module is intentionally self-contained so a reviewer can:
  3. - install dependencies from `requirements.txt`
  4. - run the analysis/replay scripts in `scripts/`
  5. It implements the minimal APIs expected by the provided scripts.
  6. Important note on reproducibility:
  7. The manuscript uses fixed-fold, subject-disjoint evaluation with a teacher gate
  8. and a fail-closed selection rule. The heavy generator training routines are in
  9. separate scripts. The functions here focus on evaluation discipline and the
  10. teacher gate, and provide reference EEGNet / ShallowConvNet implementations.
  11. If you have a GPU and the required data packaged as `.npz`, you can run end-to-end
  12. experiments. Otherwise, you can still reproduce analysis artifacts from existing
  13. run folders.
  14. """
  15. from __future__ import annotations
  16. import glob
  17. import os
  18. from dataclasses import dataclass
  19. from pathlib import Path
  20. from typing import Iterable, Optional, Tuple
  21. import numpy as np
  22. # ML
  23. from sklearn.covariance import LedoitWolf
  24. from sklearn.linear_model import LogisticRegression
  25. from sklearn.metrics import average_precision_score, roc_auc_score
  26. from sklearn.preprocessing import StandardScaler
  27. # Torch
  28. import torch
  29. import torch.nn as nn
  30. import torch.nn.functional as F
  31. # -----------------------------------------------------------------------------
  32. # Basic helpers
  33. # -----------------------------------------------------------------------------
  34. def seed_all(seed: int) -> None:
  35. """Seed numpy and torch."""
  36. np.random.seed(seed)
  37. torch.manual_seed(seed)
  38. torch.cuda.manual_seed_all(seed)
  39. def load_meta_array(meta_npz: np.lib.npyio.NpzFile, key: str, N: int) -> np.ndarray:
  40. """Load an array from meta.npz and validate length."""
  41. if key not in meta_npz:
  42. raise KeyError(f"meta_npz missing key='{key}'. Available={list(meta_npz.keys())}")
  43. arr = np.asarray(meta_npz[key])
  44. if arr.shape[0] != N:
  45. raise ValueError(f"meta[{key}] length {arr.shape[0]} != N {N}")
  46. return arr
  47. def _subject_majority_labels(y: np.ndarray, groups: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:
  48. """Return (unique_subjects, subject_label) using majority vote over windows."""
  49. groups = np.asarray(groups)
  50. y = np.asarray(y).astype(int)
  51. subs = np.unique(groups)
  52. y_sub = np.zeros(len(subs), dtype=int)
  53. for i, s in enumerate(subs):
  54. m = groups == s
  55. vals, cnt = np.unique(y[m], return_counts=True)
  56. y_sub[i] = int(vals[int(np.argmax(cnt))])
  57. return subs, y_sub
  58. def make_val_split_subject_stratified(
  59. trainval_idx: np.ndarray,
  60. y: np.ndarray,
  61. groups: np.ndarray,
  62. val_frac: float,
  63. rng: np.random.Generator,
  64. ) -> Tuple[np.ndarray, np.ndarray]:
  65. """Subject-stratified validation split.
  66. Splits by subject (group) to avoid leakage. Attempts to preserve class balance
  67. at the subject level.
  68. """
  69. trainval_idx = np.asarray(trainval_idx)
  70. y_tv = np.asarray(y)[trainval_idx]
  71. g_tv = np.asarray(groups)[trainval_idx]
  72. subs, y_sub = _subject_majority_labels(y_tv, g_tv)
  73. # Stratify subjects by label
  74. subs0 = subs[y_sub == 0]
  75. subs1 = subs[y_sub == 1]
  76. n_val0 = max(1, int(round(len(subs0) * val_frac))) if len(subs0) > 0 else 0
  77. n_val1 = max(1, int(round(len(subs1) * val_frac))) if len(subs1) > 0 else 0
  78. val_subs = []
  79. if len(subs0) > 0:
  80. val_subs.append(rng.choice(subs0, size=min(n_val0, len(subs0)), replace=False))
  81. if len(subs1) > 0:
  82. val_subs.append(rng.choice(subs1, size=min(n_val1, len(subs1)), replace=False))
  83. if len(val_subs) == 0:
  84. # Degenerate case; fall back to random split by index
  85. n_val = max(1, int(round(len(trainval_idx) * val_frac)))
  86. perm = rng.permutation(trainval_idx)
  87. return perm[n_val:], perm[:n_val]
  88. val_subs = np.concatenate(val_subs)
  89. is_val = np.isin(g_tv, val_subs)
  90. va_idx = trainval_idx[is_val]
  91. tr_idx = trainval_idx[~is_val]
  92. return tr_idx, va_idx
  93. def apply_scarcity_by_subject_stratified(
  94. tr_idx: np.ndarray,
  95. y: np.ndarray,
  96. groups: np.ndarray,
  97. scarcity: float,
  98. rng: np.random.Generator,
  99. ) -> np.ndarray:
  100. """Subsample training subjects to emulate scarcity."""
  101. tr_idx = np.asarray(tr_idx)
  102. if scarcity >= 0.999:
  103. return tr_idx
  104. y_tr = np.asarray(y)[tr_idx]
  105. g_tr = np.asarray(groups)[tr_idx]
  106. subs, y_sub = _subject_majority_labels(y_tr, g_tr)
  107. subs0 = subs[y_sub == 0]
  108. subs1 = subs[y_sub == 1]
  109. # Choose at least 1 subject per class when possible
  110. n_keep0 = max(1, int(round(len(subs0) * scarcity))) if len(subs0) > 0 else 0
  111. n_keep1 = max(1, int(round(len(subs1) * scarcity))) if len(subs1) > 0 else 0
  112. keep_subs = []
  113. if len(subs0) > 0:
  114. keep_subs.append(rng.choice(subs0, size=min(n_keep0, len(subs0)), replace=False))
  115. if len(subs1) > 0:
  116. keep_subs.append(rng.choice(subs1, size=min(n_keep1, len(subs1)), replace=False))
  117. if len(keep_subs) == 0:
  118. # fallback: uniform subsample by index
  119. n_keep = max(1, int(round(len(tr_idx) * scarcity)))
  120. return rng.choice(tr_idx, size=n_keep, replace=False)
  121. keep_subs = np.concatenate(keep_subs)
  122. mask = np.isin(g_tr, keep_subs)
  123. return tr_idx[mask]
  124. def zscore_per_subject(X: np.ndarray, groups: np.ndarray, eps: float = 1e-8) -> np.ndarray:
  125. """Per-subject, per-channel z-score normalization.
  126. X: (N, C, T)
  127. groups: (N,) subject IDs
  128. """
  129. X = np.asarray(X, dtype=np.float32)
  130. g = np.asarray(groups)
  131. out = np.empty_like(X)
  132. for s in np.unique(g):
  133. m = g == s
  134. Xi = X[m]
  135. # mean/std per channel over all windows and time
  136. mu = Xi.mean(axis=(0, 2), keepdims=True)
  137. sd = Xi.std(axis=(0, 2), keepdims=True)
  138. out[m] = (Xi - mu) / (sd + eps)
  139. return out
  140. # -----------------------------------------------------------------------------
  141. # Teacher: SPD covariance features + logistic regression
  142. # -----------------------------------------------------------------------------
  143. def _covariance_ledoitwolf(x_ct: np.ndarray, eps: float = 1e-6) -> np.ndarray:
  144. """Compute shrinkage covariance for one window.
  145. x_ct: (C, T)
  146. """
  147. x_tc = np.asarray(x_ct, dtype=np.float64).T # (T, C)
  148. lw = LedoitWolf().fit(x_tc)
  149. cov = lw.covariance_
  150. cov = cov + eps * np.eye(cov.shape[0])
  151. return cov
  152. def _logm_spd(cov: np.ndarray) -> np.ndarray:
  153. """Log-Euclidean map for SPD matrix via eigen-decomposition."""
  154. w, V = np.linalg.eigh(cov)
  155. w = np.clip(w, 1e-12, None)
  156. logw = np.log(w)
  157. return (V * logw[None, :]) @ V.T
  158. def _vec_upper(mat: np.ndarray) -> np.ndarray:
  159. iu = np.triu_indices(mat.shape[0])
  160. return mat[iu]
  161. def covariance_features_logeuclid(X: np.ndarray, eps: float = 1e-6) -> np.ndarray:
  162. """Compute log-Euclidean covariance features for a batch.
  163. Returns a 2D array of shape (N, D).
  164. """
  165. X = np.asarray(X)
  166. N, C, T = X.shape
  167. feats = np.zeros((N, C * (C + 1) // 2), dtype=np.float64)
  168. for i in range(N):
  169. cov = _covariance_ledoitwolf(X[i], eps=eps)
  170. logc = _logm_spd(cov)
  171. feats[i] = _vec_upper(logc)
  172. return feats
  173. def train_teacher(Xtr: np.ndarray, ytr: np.ndarray, seed: int = 0, eps: float = 1e-6):
  174. """Train the geometry-aware teacher (scaler + logistic regression)."""
  175. seed_all(seed)
  176. feats = covariance_features_logeuclid(Xtr, eps=eps)
  177. scaler = StandardScaler().fit(feats)
  178. Xs = scaler.transform(feats)
  179. clf = LogisticRegression(
  180. max_iter=1000,
  181. class_weight="balanced",
  182. solver="liblinear",
  183. random_state=seed,
  184. ).fit(Xs, ytr)
  185. return scaler, clf
  186. def teacher_filter_agree_quantile(
  187. Xs: np.ndarray,
  188. ys: np.ndarray,
  189. scaler: StandardScaler,
  190. clf: LogisticRegression,
  191. keep_quantile: float,
  192. min_keep: int = 200,
  193. eps: float = 1e-6,
  194. ):
  195. """Filter synthetic candidates by teacher agreement and confidence quantile.
  196. Returns (X_keep, y_keep).
  197. """
  198. if Xs is None or ys is None or len(ys) == 0:
  199. return np.empty((0,), dtype=np.float32), np.empty((0,), dtype=np.int64)
  200. feats = covariance_features_logeuclid(Xs, eps=eps)
  201. Xf = scaler.transform(feats)
  202. proba = clf.predict_proba(Xf)
  203. pred = np.argmax(proba, axis=1)
  204. conf = np.max(proba, axis=1)
  205. ys = np.asarray(ys).astype(int)
  206. agree = pred == ys
  207. if agree.sum() < min_keep:
  208. # Fail closed: too few label-consistent samples.
  209. return np.empty((0, Xs.shape[1], Xs.shape[2]), dtype=np.float32), np.empty((0,), dtype=np.int64)
  210. conf_agree = conf[agree]
  211. tau = np.quantile(conf_agree, keep_quantile)
  212. keep = agree & (conf >= tau)
  213. # If pruning is too strict but we have enough agreeing samples, keep top-min_keep
  214. if keep.sum() < min_keep:
  215. idx_agree = np.where(agree)[0]
  216. order = np.argsort(-conf[idx_agree])
  217. idx_keep = idx_agree[order[:min_keep]]
  218. else:
  219. idx_keep = np.where(keep)[0]
  220. return np.asarray(Xs[idx_keep], dtype=np.float32), np.asarray(ys[idx_keep], dtype=np.int64)
  221. def find_synth_pool(pool_root: str, seed: int, fold: int) -> Tuple[np.ndarray, np.ndarray]:
  222. """Locate and load (X_pool.npy, y_pool.npy) for a given seed and fold.
  223. This helper tries common directory patterns and falls back to a glob search.
  224. """
  225. root = Path(pool_root)
  226. candidates = [
  227. root / f"seed_{seed}" / f"fold_{fold}" / "X_pool.npy",
  228. root / f"seed{seed}" / f"fold{fold}" / "X_pool.npy",
  229. root / f"seed_{seed}" / f"fold{fold}" / "X_pool.npy",
  230. root / f"seed{seed}" / f"fold_{fold}" / "X_pool.npy",
  231. ]
  232. x_path: Optional[Path] = None
  233. for c in candidates:
  234. if c.exists():
  235. x_path = c
  236. break
  237. if x_path is None:
  238. # Glob fallback
  239. patt = str(root / f"**/*seed*{seed}*/*fold*{fold}*/X_pool.npy")
  240. hits = glob.glob(patt, recursive=True)
  241. if hits:
  242. x_path = Path(hits[0])
  243. if x_path is None:
  244. raise FileNotFoundError(
  245. f"Could not find X_pool.npy for seed={seed}, fold={fold} under {pool_root}. "
  246. f"Tried common patterns and glob." )
  247. y_path = x_path.parent / "y_pool.npy"
  248. if not y_path.exists():
  249. # Some pools use Y_pool.npy
  250. y_path2 = x_path.parent / "Y_pool.npy"
  251. if y_path2.exists():
  252. y_path = y_path2
  253. else:
  254. raise FileNotFoundError(f"Found {x_path} but missing y_pool.npy in {x_path.parent}")
  255. Xs = np.load(x_path)
  256. ys = np.load(y_path)
  257. return np.asarray(Xs, dtype=np.float32), np.asarray(ys, dtype=np.int64)
  258. # -----------------------------------------------------------------------------
  259. # Torch models
  260. # -----------------------------------------------------------------------------
  261. class EEGNet(nn.Module):
  262. """A compact EEGNet-like architecture for (C, T) windows."""
  263. def __init__(self, C: int, T: int, n_classes: int = 2, F1: int = 8, D: int = 2, F2: int = 16, dropout: float = 0.25):
  264. super().__init__()
  265. self.C = C
  266. self.T = T
  267. # Block 1: temporal conv
  268. self.conv_temporal = nn.Conv2d(1, F1, kernel_size=(1, 64), padding=(0, 32), bias=False)
  269. self.bn1 = nn.BatchNorm2d(F1)
  270. # Depthwise spatial conv
  271. self.conv_spatial = nn.Conv2d(F1, F1 * D, kernel_size=(C, 1), groups=F1, bias=False)
  272. self.bn2 = nn.BatchNorm2d(F1 * D)
  273. self.act = nn.ELU()
  274. self.pool1 = nn.AvgPool2d(kernel_size=(1, 4))
  275. self.drop1 = nn.Dropout(dropout)
  276. # Separable conv
  277. self.sep_depth = nn.Conv2d(F1 * D, F1 * D, kernel_size=(1, 16), padding=(0, 8), groups=F1 * D, bias=False)
  278. self.sep_point = nn.Conv2d(F1 * D, F2, kernel_size=(1, 1), bias=False)
  279. self.bn3 = nn.BatchNorm2d(F2)
  280. self.pool2 = nn.AvgPool2d(kernel_size=(1, 8))
  281. self.drop2 = nn.Dropout(dropout)
  282. # Compute flattened dim
  283. with torch.no_grad():
  284. x = torch.zeros(1, 1, C, T)
  285. x = self.drop1(self.pool1(self.act(self.bn2(self.conv_spatial(self.bn1(self.conv_temporal(x)))))))
  286. x = self.drop2(self.pool2(self.act(self.bn3(self.sep_point(self.sep_depth(x))))))
  287. flat = int(np.prod(x.shape[1:]))
  288. self.classifier = nn.Linear(flat, n_classes)
  289. def forward(self, x: torch.Tensor) -> torch.Tensor:
  290. # x: (N, C, T)
  291. x = x.unsqueeze(1) # (N, 1, C, T)
  292. x = self.conv_temporal(x)
  293. x = self.bn1(x)
  294. x = self.conv_spatial(x)
  295. x = self.bn2(x)
  296. x = self.act(x)
  297. x = self.pool1(x)
  298. x = self.drop1(x)
  299. x = self.sep_depth(x)
  300. x = self.sep_point(x)
  301. x = self.bn3(x)
  302. x = self.act(x)
  303. x = self.pool2(x)
  304. x = self.drop2(x)
  305. x = torch.flatten(x, start_dim=1)
  306. return self.classifier(x)
  307. class ShallowConvNet(nn.Module):
  308. """A shallow ConvNet variant for (C, T) windows."""
  309. def __init__(self, C: int, T: int, n_classes: int = 2, F: int = 40, dropout: float = 0.5):
  310. super().__init__()
  311. self.conv_time = nn.Conv2d(1, F, kernel_size=(1, 25), padding=(0, 12), bias=False)
  312. self.conv_spat = nn.Conv2d(F, F, kernel_size=(C, 1), bias=False)
  313. self.bn = nn.BatchNorm2d(F)
  314. self.pool = nn.AvgPool2d(kernel_size=(1, 75), stride=(1, 15))
  315. self.drop = nn.Dropout(dropout)
  316. with torch.no_grad():
  317. x = torch.zeros(1, 1, C, T)
  318. x = self.conv_time(x)
  319. x = self.conv_spat(x)
  320. x = self.bn(x)
  321. x = torch.square(x)
  322. x = self.pool(x)
  323. x = torch.log(torch.clamp(x, min=1e-6))
  324. x = self.drop(x)
  325. flat = int(np.prod(x.shape[1:]))
  326. self.classifier = nn.Linear(flat, n_classes)
  327. def forward(self, x: torch.Tensor) -> torch.Tensor:
  328. x = x.unsqueeze(1)
  329. x = self.conv_time(x)
  330. x = self.conv_spat(x)
  331. x = self.bn(x)
  332. x = torch.square(x)
  333. x = self.pool(x)
  334. x = torch.log(torch.clamp(x, min=1e-6))
  335. x = self.drop(x)
  336. x = torch.flatten(x, start_dim=1)
  337. return self.classifier(x)
  338. # -----------------------------------------------------------------------------
  339. # Training / inference
  340. # -----------------------------------------------------------------------------
  341. @torch.no_grad()
  342. def torch_predict_proba(model: nn.Module, X: np.ndarray, device: str = "cpu", infer_bs: int = 256, num_workers: int = 0) -> np.ndarray:
  343. model.eval()
  344. model.to(device)
  345. X = np.asarray(X, dtype=np.float32)
  346. N = X.shape[0]
  347. out = np.zeros((N, 2), dtype=np.float32)
  348. for i in range(0, N, infer_bs):
  349. xb = torch.from_numpy(X[i:i+infer_bs]).to(device)
  350. logits = model(xb)
  351. probs = torch.softmax(logits, dim=1).cpu().numpy()
  352. out[i:i+infer_bs] = probs
  353. return out
  354. def _safe_auc(y: np.ndarray, p: np.ndarray) -> float:
  355. y = np.asarray(y).reshape(-1)
  356. p = np.asarray(p).reshape(-1)
  357. if len(np.unique(y)) < 2:
  358. return float("nan")
  359. return float(roc_auc_score(y, p))
  360. def _safe_ap(y: np.ndarray, p: np.ndarray) -> float:
  361. y = np.asarray(y).reshape(-1)
  362. p = np.asarray(p).reshape(-1)
  363. if len(np.unique(y)) < 2:
  364. return float("nan")
  365. return float(average_precision_score(y, p))
  366. def train_torch_model(
  367. model: nn.Module,
  368. n_classes: int,
  369. Xtr: np.ndarray,
  370. ytr: np.ndarray,
  371. Xva: np.ndarray,
  372. yva: np.ndarray,
  373. device: str = "cpu",
  374. epochs: int = 100,
  375. patience: int = 15,
  376. lr: float = 1e-3,
  377. weight_decay: float = 0.0,
  378. batch_size: int = 128,
  379. infer_bs: int = 256,
  380. num_workers: int = 0,
  381. amp: bool = False,
  382. seed: int = 0,
  383. log_path: Optional[str] = None,
  384. ):
  385. """Train a torch model with early stopping on validation AUROC.
  386. Returns: (best_model, best_val_auc, best_val_ap, fallback_flag)
  387. """
  388. seed_all(seed)
  389. Xtr = np.asarray(Xtr, dtype=np.float32)
  390. ytr = np.asarray(ytr, dtype=np.int64)
  391. Xva = np.asarray(Xva, dtype=np.float32)
  392. yva = np.asarray(yva, dtype=np.int64)
  393. model = model.to(device)
  394. # Class-balanced weights
  395. cls_counts = np.bincount(ytr, minlength=n_classes).astype(np.float32)
  396. cls_counts = np.maximum(cls_counts, 1.0)
  397. weights = (cls_counts.sum() / (n_classes * cls_counts))
  398. w_t = torch.tensor(weights, dtype=torch.float32, device=device)
  399. opt = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay)
  400. scaler = torch.cuda.amp.GradScaler(enabled=(amp and device.startswith("cuda")))
  401. best_auc = -1e9
  402. best_ap = -1e9
  403. best_state = None
  404. bad = 0
  405. # simple numpy batching
  406. idx = np.arange(len(ytr))
  407. for ep in range(epochs):
  408. model.train()
  409. np.random.shuffle(idx)
  410. for i in range(0, len(idx), batch_size):
  411. b = idx[i:i+batch_size]
  412. xb = torch.from_numpy(Xtr[b]).to(device)
  413. yb = torch.from_numpy(ytr[b]).to(device)
  414. opt.zero_grad(set_to_none=True)
  415. with torch.cuda.amp.autocast(enabled=(amp and device.startswith("cuda"))):
  416. logits = model(xb)
  417. loss = F.cross_entropy(logits, yb, weight=w_t)
  418. scaler.scale(loss).backward()
  419. scaler.step(opt)
  420. scaler.update()
  421. # validation
  422. probs = torch_predict_proba(model, Xva, device=device, infer_bs=infer_bs, num_workers=0)
  423. p1 = probs[:, 1]
  424. vauc = _safe_auc(yva, p1)
  425. vap = _safe_ap(yva, p1)
  426. if np.isfinite(vauc) and vauc > best_auc + 1e-6:
  427. best_auc = float(vauc)
  428. best_ap = float(vap)
  429. best_state = {k: v.detach().cpu().clone() for k, v in model.state_dict().items()}
  430. bad = 0
  431. else:
  432. bad += 1
  433. if bad >= patience:
  434. break
  435. if best_state is not None:
  436. model.load_state_dict(best_state)
  437. return model, float(best_auc), float(best_ap), False

runner.py at commit 1ac3d64, under MIT · at the source

Overview

Authors: Daniel Choi1, Cordelia Yip1, Andrew Choi1, Junho Park1
  1. University of Calgary,Calgary, AB Canada
Institutions: University of Calgary (Canada)
Journal: NPJ digital medicine, volume 9, issue 1, article 634
Dates: received 29 January 2026; accepted 12 May 2026; published online 25 May 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1038/s41746-026-02778-0 · PMID 42185473 · PMCID PMC13482338 · OpenAlex W7162336389
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: EEG (modality), human (organism), pain (population), clinical / translational (subfield)
Methods: Spectral & time-frequency, Preprocessing, Connectivity, Machine learning, Statistics
Keywords: Engineering, Health care, Medical research, Neuroscience
Topic: EEG and Brain-Computer Interfaces (Cognitive Neuroscience, Neuroscience), according to OpenAlex
Funding: University of Calgary (2025 NSERC Discovery Grant Bridge Fund)
Citations: not cited yet (Europe PMC); 22 references in the paper

Abstract

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

Repository

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

danielchoi0315/TGA-repo

License: MIT
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: 1ac3d641a1b63a1514a50dfa77fbe07bde1e250e, 30 August 2026
Languages: Python (49), Shell (3)
Size: 60 files, 52 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, CITATION.cff, environment (requirements.txt), tests, documentation
Not found: continuous integration
Tools: NumPy (26 files), scikit-learn (17 files), pandas (16 files), PyTorch (8 files), Matplotlib (4 files), MOABB (2 files), SciPy (2 files), MNE-Python (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
54 files

Code availability statement

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

Read it in the paper: doi.org/10.1038/s41746-026-02778-0.

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

Datasets cited

Data availability statement

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

Read it in the paper: doi.org/10.1038/s41746-026-02778-0.

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

Recorded: type, language, journal, volume, issue, pages, dates, 4 authors, 4 keywords, 1 funder, 16 references.

Cite

This paper

Choi, D., Yip, C., Choi, A., & Park, J. (2026). Trust-gated synthetic EEG augmentation reduces performance drops when generalizing to new patients. NPJ digital medicine, 9(1), 634. https://doi.org/10.1038/s41746-026-02778-0

BibTeX

@article{choi2026trust,
author = {Choi, Daniel and Yip, Cordelia and Choi, Andrew and Park, Junho},
title = {{Trust-gated synthetic EEG augmentation reduces performance drops when generalizing to new patients}},
journal = {NPJ digital medicine},
year = {2026},
month = may,
volume = {9},
number = {1},
pages = {634},
publisher = {Nature Publishing Group},
issn = {2398-6352},
doi = {10.1038/s41746-026-02778-0},
url = {https://doi.org/10.1038/s41746-026-02778-0},
pmid = {42185473},
pmcid = {PMC13482338}
}

RIS

TY - JOUR
AU - Choi, Daniel
AU - Yip, Cordelia
AU - Choi, Andrew
AU - Park, Junho
TI - Trust-gated synthetic EEG augmentation reduces performance drops when generalizing to new patients
T2 - NPJ digital medicine
J2 - NPJ Digit Med
PY - 2026
DA - 2026/05/25
VL - 9
IS - 1
SP - 634
SN - 2398-6352
PB - Nature Publishing Group
DO - 10.1038/s41746-026-02778-0
UR - https://doi.org/10.1038/s41746-026-02778-0
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41746-026-02778-0",
"type": "article-journal",
"title": "Trust-gated synthetic EEG augmentation reduces performance drops when generalizing to new patients",
"container-title": "NPJ digital medicine",
"author": [
{
"family": "Choi",
"given": "Daniel"
},
{
"family": "Yip",
"given": "Cordelia"
},
{
"family": "Choi",
"given": "Andrew"
},
{
"family": "Park",
"given": "Junho"
}
],
"container-title-short": "NPJ Digit Med",
"volume": "9",
"issue": "1",
"page": "634",
"DOI": "10.1038/s41746-026-02778-0",
"PMID": "42185473",
"PMCID": "PMC13482338",
"ISSN": "2398-6352",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41746-026-02778-0",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
25
]
]
}
}

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

Similar papers

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

[1] doi:10.1038/s41598-026-47627-y [code]
QuantumNeuroXAI: a quantum-inspired deep learning framework with explainability for brain signal analysis and neurological disorder detection.
Journal: Scientific reports
In common: MNE-Python, PyTorch, scikit-learn, 4 other tools, bbci.de/competition/iv, EEG, 2 references
[2] doi:10.1002/hbm.70528 [code]
Explainable AI Insights Into EEG Classification and Its Alignment to Neural Correlates.
Journal: Human brain mapping
In common: MOABB, MNE-Python, PyTorch, 5 other tools, EEG, 1 reference
[3] doi: [code]
Self-Rectifying Integrate-and-Fire Neuron and Collaborative Trim Training Framework for SNN-Based EEG Motor Imagery Classification
Journal: Brain sciences
In common: PyTorch, SciPy, Matplotlib, 1 other tool, bbci.de/competition/iv, EEG, 3 references
[4] doi:10.1371/journal.pone.0354976 [code]
Improved motor imagery BCI performance via task-unaware compression in the BELT Bayesian Edge-Cloud architecture.
Journal: PloS one
In common: MNE-Python, pandas, SciPy, 2 other tools, bbci.de/competition/iv, EEG, 1 reference
[5] doi:10.1371/journal.pone.0347671 [code]
RMETNet: A cross-subject motor imagery EEG signal classification model based on TSLANet and riemannian geometry features.
Journal: PloS one
In common: MNE-Python, PyTorch, scikit-learn, 4 other tools, EEG, 2 references
[6] doi:10.1186/s13634-026-01330-2 [code]
Leednet: a lightweight network for event detection in EEG signals.
Journal: Journal on advances in signal processing
In common: MNE-Python, PyTorch, scikit-learn, 4 other tools, EEG, 2 references
[7] doi:10.3390/s26051730 [code]
SFE-GAT: Structure-Feature Evolution Graph Attention Network for Motor Imagery Decoding.
Journal: Sensors (Basel, Switzerland)
In common: PyTorch, SciPy, Matplotlib, 1 other tool, bbci.de/competition/iv, EEG, 1 reference
[8] doi:10.3389/fnhum.2026.1895016 [code]
A confidence-gated source selection strategy for cross-session transfer in brain-computer interfaces.
Journal: Frontiers in human neuroscience
In common: MOABB, MNE-Python, scikit-learn, 2 other tools, EEG, 1 reference
[9] doi:10.3389/fnins.2026.1874302 [code]
Architecture-data matching for EEG-EMG decoding: compact deep models match classical spectral decoders on the WAY-EEG-GAL grasp-and-lift dataset.
Journal: Frontiers in neuroscience
In common: PyTorch, scikit-learn, pandas, 3 other tools, EEG, 3 references
[10] doi:10.3390/bios16080437 [code]
Deep Learning for Ear-EEG-Based Brain-Computer Interface: A Systematic Comparison and Design Insights.
Journal: Biosensors
In common: MNE-Python, PyTorch, scikit-learn, 3 other tools, EEG, 2 references

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.