OSCR

Multispectral 7 Tesla MRI as a potential predictor of dopamine transporter deficiency in Parkinson's disease.

Code ↔ Paper

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

The 9 matches
  1. [1] § Methods › Contrastive embedder ↔ embedder/trainer.py, lines 138–165 · score 0.94 · additive Gaussian noise, monotone gamma, channel gain, low rank, feature dropout, shifts
  2. [2] § Methods › DaT predictor ↔ nt_mlp/trainer.py, lines 16–64 · score 0.83 · cosine annealing scheduler, mixed precision, AdamW, weight decay, batch, train
  3. [3] § Methods › Contrastive embedder ↔ embedder/model.py, lines 84–100 · score 0.72 · ReLU, layer normalization, predicts subject, subject head, hidden, dropout
  4. [4] § Methods › Contrastive embedder ↔ nt_mlp/trainer.py, lines 16–64 · score 0.69 · mixed precision, AdamW, weight decay, batch, losses, train
  5. [5] § Methods › Contrastive embedder ↔ embedder/trainer.py, lines 22–103 · score 0.69 · weight decay, adversarial losses, mixed precision, adversarial heads, encoder, batch
  6. [6] § Methods › Contrastive embedder ↔ embedder/trainer.py, lines 22–103 · score 0.63 · hidden layers, adversarial heads, coordinate head, architectural, dimensional, encoder
  7. [7] § Methods › Contrastive embedder ↔ embedder/model.py, lines 190–229 · score 0.60 · ReLU, layer normalization, vectors, hidden, MLP, encoder
  8. [8] § Methods › Contrastive embedder ↔ nt_mlp/core_model.py, lines 25–67 · score 0.58 · ReLU, layer normalization, activation, MLP, hidden, dropout
  9. [9] § Results › Quantitative ROI analysis ↔ nt_mlp/cv_evaluator.py, lines 887–1006 · score 0.52 · predicted binding potentials, boxplots, FDR, fold, striatum, HC

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 · 637 lines · 25 KB · MIT · 3 matches

  1. import torch
  2. import torch.optim as optim
  3. from torch.utils.data import DataLoader
  4. import torch.nn.functional as F
  5. import numpy as np
  6. from pathlib import Path
  7. import json
  8. import argparse
  9. import math
  10. import matplotlib.pyplot as plt
  11. from augs import (
  12. Compose, RandomGlobalGainShift, RandomChannelLowRankGain,
  13. BatchPartitionGain, RandomMonotoneGammaWarp, AdditiveGaussianNoise,
  14. FeatureDropout
  15. )
  16. from model import *
  17. from data_handler import ContrastiveDataHandler
  18. def parse_args():
  19. """Parse command-line arguments for contrastive learning training.
  20. Defaults are set to match the values used in run_experiment.py.
  21. """
  22. parser = argparse.ArgumentParser(description='Contrastive Learning for MRI Voxels')
  23. # Required arguments
  24. parser.add_argument('--subjects', nargs='+', required=True,
  25. help='List of training subject IDs')
  26. parser.add_argument('--exp_name', type=str, required=True,
  27. help='Experiment name for saving outputs')
  28. # Data settings
  29. parser.add_argument('--test_subjects', nargs='+', default=None,
  30. help='Test subject IDs (for visualization only)')
  31. parser.add_argument('--fname', type=str, default="train_data.npy",
  32. help='Data filename to load from each subject directory')
  33. # Training hyperparameters (defaults from run_experiment.py)
  34. parser.add_argument('--steps', type=int, default=4000,
  35. help='Number of training steps')
  36. parser.add_argument('--batch_vox', type=int, default=16384,
  37. help='Number of voxels per batch')
  38. parser.add_argument('--lr', type=float, default=1e-4,
  39. help='Learning rate')
  40. parser.add_argument('--weight_decay', type=float, default=1e-2,
  41. help='Weight decay (L2 regularization)')
  42. parser.add_argument('--temperature', type=float, default=0.07,
  43. help='Temperature for contrastive loss')
  44. parser.add_argument('--random_seed', type=int, default=66,
  45. help='Random seed for reproducibility')
  46. # Model architecture (defaults from run_experiment.py)
  47. parser.add_argument('--embedding_dim', type=int, default=16,
  48. help='Embedding dimension')
  49. parser.add_argument('--hidden_dims', nargs='+', type=int, default=[2048, 2048, 2048, 2048],
  50. help='Hidden layer dimensions')
  51. parser.add_argument('--dropout', type=float, default=0.3,
  52. help='Dropout rate for encoder')
  53. parser.add_argument('--hidden_dim_adv', type=int, default=512,
  54. help='Hidden dimension for adversarial heads')
  55. parser.add_argument('--adv_dropout', type=float, default=0.3,
  56. help='Dropout rate for adversarial heads')
  57. # Learning rate warmup
  58. parser.add_argument('--warmup_steps', type=int, default=0,
  59. help='Linear LR warmup steps (0 disables warmup)')
  60. parser.add_argument('--warmup_start_factor', type=float, default=0.1,
  61. help='Warmup start factor relative to lr')
  62. # Loss weights (defaults from run_experiment.py)
  63. parser.add_argument('--w_inst', type=float, default=1.0,
  64. help='Weight for instance contrastive loss')
  65. parser.add_argument('--w_coord_adv', type=float, default=20.0,
  66. help='Weight for coordinate adversary loss')
  67. parser.add_argument('--w_subj_adv', type=float, default=20.0,
  68. help='Weight for subject adversary loss')
  69. # Adversarial training (defaults from run_experiment.py)
  70. parser.add_argument('--use_coord_adv', action='store_true',
  71. help='Use coordinate adversary (enabled in run_experiment)')
  72. parser.add_argument('--use_subject_adv', action='store_true',
  73. help='Use subject adversary (enabled in run_experiment)')
  74. parser.add_argument('--coordHeadBins', type=int, default=16,
  75. help='Number of bins for coordinate head')
  76. # Augmentation (defaults from run_experiment.py)
  77. parser.add_argument('--use_weak_strong_aug', action='store_true',
  78. help='Use weak vs strong augmentations (enabled in run_experiment)')
  79. parser.add_argument('--weak_strength', type=float, default=0.30,
  80. help='Weak view magnitude in [0,1] (scaled into each op range)')
  81. parser.add_argument('--strong_strength', type=float, default=0.75,
  82. help='Strong view magnitude in [0,1] (scaled into each op range)')
  83. parser.add_argument('--aug_ramp', action='store_true',
  84. help='Cosine-ramp augmentation strength (enabled in run_experiment)')
  85. # Training options
  86. parser.add_argument('--amp_scaler', action='store_true',
  87. help='Use mixed precision training (enabled in run_experiment)')
  88. return parser.parse_args()
  89. def setup_experiment_dir(exp_name):
  90. """Create experiment directory and return path."""
  91. exp_dir = Path(f"experiments/{exp_name}")
  92. exp_dir.mkdir(parents=True, exist_ok=True)
  93. return exp_dir
  94. def save_config(exp_dir, args):
  95. """Save experiment configuration."""
  96. config = vars(args)
  97. with open(exp_dir / "config.json", 'w') as f:
  98. json.dump(config, f, indent=2)
  99. def digitize(v, edges, bins):
  100. """Discretize continuous values into bins.
  101. Args:
  102. v: Continuous values to discretize
  103. edges: Bin edges
  104. bins: Number of bins
  105. Returns:
  106. Integer bin indices in [0, bins-1]
  107. """
  108. idx = torch.bucketize(v, edges[1:-1]) # exclude first/last edge
  109. return idx.clamp_(0, bins-1)
  110. def _lerp(a, b, t):
  111. return a + (b - a) * t
  112. def cosine_ramp(t: float) -> float:
  113. t = max(0.0, min(1.0, t))
  114. return 0.5 * (1.0 - math.cos(math.pi * t))
  115. def build_augment_from_strength(m: float) -> Compose:
  116. """Magnitude m in [0,1] controls intensity per operator via interpolation."""
  117. # Global/channel gain-shift
  118. gain_logstd = _lerp(0.05, 0.30, m)
  119. bias_std = _lerp(0.02, 0.1, m)
  120. rank = int(round(_lerp(6, 18, m)))
  121. gain_std = _lerp(0.12, 0.30, m)
  122. # Batch/domain randomization
  123. n_groups = int(round(_lerp(4, 12, m)))
  124. part_logstd = _lerp(0.12, 0.30, m)
  125. # Monotone gamma warp
  126. delta = _lerp(0.15, 0.40, m)
  127. gamma_p = _lerp(0.6, 1.0, m)
  128. # Noise + feature dropout
  129. sigma = _lerp(0.03, 0.10, m)
  130. p_drop = _lerp(0.1, 0.50, m)
  131. return Compose(
  132. RandomGlobalGainShift(gain_logstd=gain_logstd, bias_std=bias_std, p=0.75),
  133. RandomChannelLowRankGain(feat_dim=349, rank=rank, gain_std=gain_std, p=0.75),
  134. BatchPartitionGain(n_groups=n_groups, logstd=part_logstd, p=0.75),
  135. RandomMonotoneGammaWarp(delta=delta, per_feature=True, p=0.75, eps=1e-4),
  136. AdditiveGaussianNoise(sigma=sigma, p=1.0),
  137. FeatureDropout(p_drop=p_drop),
  138. )
  139. def softmax_top1_acc(logits, y):
  140. return (logits.argmax(dim=1) == y).float().mean()
  141. def q_edges(v, bins):
  142. qs = torch.linspace(0, 1, bins+1).to(v.device)
  143. e = torch.quantile(v, qs, interpolation="linear")
  144. e[0], e[-1] = 0.0, 1.0 # clamp ends
  145. # enforce strict monotonicity to avoid bucketize ties
  146. for i in range(1, e.numel()):
  147. if e[i] <= e[i-1]:
  148. e[i] = min(1.0, e[i-1] + 1e-6)
  149. return e
  150. def main():
  151. args = parse_args()
  152. set_random_seeds(args.random_seed)
  153. print(f"Random seed set to: {args.random_seed}")
  154. # Setup experiment
  155. exp_dir = setup_experiment_dir(args.exp_name)
  156. save_config(exp_dir, args)
  157. print(f"Experiment directory: {exp_dir}")
  158. device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  159. print(f"Using device: {device}")
  160. # Enable TF32 for additional speed on Ampere+ (safe for training)
  161. if torch.cuda.is_available():
  162. torch.backends.cuda.matmul.allow_tf32 = True
  163. torch.backends.cudnn.allow_tf32 = True
  164. # Decide AMP dtype based on hardware support (prefer BF16)
  165. bf16_ok = False
  166. if torch.cuda.is_available():
  167. if hasattr(torch.cuda, "is_bf16_supported") and torch.cuda.is_bf16_supported():
  168. bf16_ok = True
  169. else:
  170. try:
  171. major, _ = torch.cuda.get_device_capability()
  172. if major >= 8:
  173. bf16_ok = True
  174. except Exception:
  175. pass
  176. autocast_dtype = torch.bfloat16 if bf16_ok else torch.float16
  177. print(f"AMP: enabled={bool(args.amp_scaler)} | dtype={autocast_dtype} | bf16_supported={bf16_ok}")
  178. # !!! augs assume z-scored features
  179. # Pre-build static augment if not using weak/strong policy
  180. use_ws = args.use_weak_strong_aug
  181. if not use_ws:
  182. augment = Compose(
  183. AdditiveGaussianNoise(sigma=0.05, p=1.0),
  184. FeatureDropout(p_drop=0.3),
  185. )
  186. else:
  187. augment = None # will build per-step weak/strong aug below
  188. # Data setup
  189. dataset = ContrastiveDataHandler(subject_ids=args.subjects,
  190. batch_vox=args.batch_vox,
  191. fname=args.fname, augment=augment)
  192. # Model setup
  193. base_encoder = VoxelEncoder(
  194. in_dim=349,
  195. hidden_dims=args.hidden_dims,
  196. embedding_dim=args.embedding_dim,
  197. dropout=args.dropout
  198. ).to(device)
  199. # Wrap with adversarial heads for GRL training
  200. model = VoxelEncoderAdv(
  201. base_encoder=base_encoder,
  202. embedding_dim=args.embedding_dim,
  203. use_coord_adv=args.use_coord_adv,
  204. use_subject_adv=args.use_subject_adv,
  205. n_subjects=dataset.n_subj,
  206. adv_hidden=args.hidden_dim_adv,
  207. adv_dropout=args.adv_dropout,
  208. coordHeadBins=args.coordHeadBins,
  209. ).to(device)
  210. # Optional compile for speed (PyTorch 2.x)
  211. if hasattr(torch, "compile"):
  212. try:
  213. model = torch.compile(model)
  214. print("Compiled model with torch.compile (max-autotune)")
  215. except Exception:
  216. print("torch.compile not available or failed; continuing without compile")
  217. crit = WeightedSupConLoss(temperature=args.temperature)
  218. try:
  219. optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay, fused=True)
  220. print("Using fused AdamW optimizer")
  221. except TypeError:
  222. optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
  223. #optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
  224. # Learning-rate schedule with optional linear warmup into cosine decay
  225. if args.warmup_steps > 0:
  226. total_after_warmup = max(1, args.steps - args.warmup_steps)
  227. warmup = optim.lr_scheduler.LinearLR(
  228. optimizer,
  229. start_factor=max(1e-6, min(1.0, args.warmup_start_factor)),
  230. total_iters=args.warmup_steps,
  231. )
  232. cosine = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=total_after_warmup)
  233. sched = optim.lr_scheduler.SequentialLR(
  234. optimizer,
  235. schedulers=[warmup, cosine],
  236. milestones=[args.warmup_steps],
  237. )
  238. print(f"Using LR warmup: {args.warmup_steps} steps, start_factor={args.warmup_start_factor}")
  239. else:
  240. sched = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.steps)
  241. print(f"Training with {len(args.subjects)} subjects: {args.subjects}")
  242. # Test BEFORE training
  243. print("=== BEFORE TRAINING ===")
  244. # Training loop
  245. loss_history = []
  246. # Track individual loss components for plotting
  247. contr_history = []
  248. coord_history = []
  249. subj_history = []
  250. log_steps = []
  251. print(f"\n=== TRAINING FOR {args.steps} STEPS ===")
  252. loader = torch.utils.data.DataLoader(
  253. dataset,
  254. batch_size=None, # iterable yields whole batches
  255. num_workers=8, # start with 4–8, tune
  256. pin_memory=True,
  257. persistent_workers=True,
  258. prefetch_factor=2
  259. )
  260. data_iter = iter(loader)
  261. edges_x = edges_y = edges_z = None
  262. if args.use_coord_adv:
  263. #average of some batches for stable bin edges
  264. with torch.no_grad():
  265. all_coords = []
  266. n_avg = 20
  267. for _ in range(n_avg):
  268. try:
  269. _, coords, _ = next(data_iter)
  270. except StopIteration:
  271. data_iter = iter(loader)
  272. _, coords, _ = next(data_iter)
  273. all_coords.append(coords)
  274. all_coords = torch.cat(all_coords,0)
  275. edges_x = q_edges(all_coords[:,0], args.coordHeadBins).to(device)
  276. edges_y = q_edges(all_coords[:,1], args.coordHeadBins).to(device)
  277. edges_z = q_edges(all_coords[:,2], args.coordHeadBins).to(device)
  278. #reset data_iter to start of epoch
  279. data_iter = iter(loader)
  280. warmup_steps = max(1, int(args.steps * 0)) # avoid div by zero
  281. model.set_grl_lambdas(coord=1.0 if args.use_coord_adv else 0.0,
  282. subj=1.0 if args.use_subject_adv else 0.0)
  283. for step in range(args.steps):
  284. # ramp *before* forward so this step uses current weights
  285. ramp = min(1.0, step / warmup_steps)
  286. # Backward-compat: map deprecated lambda_* to weights if new weights are unset
  287. w_coord = args.w_coord_adv if (args.w_coord_adv is not None) else args.lambda_coord
  288. w_subj = args.w_subj_adv if (args.w_subj_adv is not None) else args.lambda_subject
  289. # Effective ramped weights
  290. eff_w_coord = (w_coord if args.use_coord_adv else 0.0) * ramp
  291. eff_w_subj = (w_subj if args.use_subject_adv else 0.0) * ramp
  292. # Set GRL λ to 1.0 (when enabled) so strength is controlled via loss weights
  293. optimizer.zero_grad(set_to_none=True)
  294. try:
  295. feats, coords, sids = next(data_iter)
  296. except StopIteration:
  297. data_iter = iter(loader)
  298. feats, coords, sids = next(data_iter)
  299. # Transfer to GPU with non-blocking for async copy
  300. feats = feats.to(device, non_blocking=True)
  301. coords = coords.to(device, non_blocking=True)
  302. sids = sids.to(device, non_blocking=True)
  303. # Apply augmentations; either static (HA presets) or dynamic weak/strong
  304. if augment is not None:
  305. x1 = augment(feats)
  306. x2 = augment(feats)
  307. else:
  308. # dynamic policy
  309. if args.use_weak_strong_aug:
  310. t = step / float(max(1, args.steps - 1))
  311. ramp_factor = max(cosine_ramp(t) , 0.2) if args.aug_ramp else 1.0
  312. m_w = float(np.clip(args.weak_strength * ramp_factor, 0.0, 1.0))
  313. m_s = float(np.clip(args.strong_strength * ramp_factor, 0.0, 1.0))
  314. aug_w = build_augment_from_strength(m_w)
  315. aug_s = build_augment_from_strength(m_s)
  316. x1 = aug_w(feats)
  317. x2 = aug_s(feats)
  318. else:
  319. x1 = feats
  320. x2 = feats
  321. # Forward pass under autocast (compute in BF16/FP16)
  322. with torch.amp.autocast('cuda', enabled=bool(args.amp_scaler), dtype=autocast_dtype):
  323. compute_heads = (args.use_coord_adv or args.use_subject_adv)
  324. z1, coord1, subj1 = model(x1, compute_heads=compute_heads)
  325. z2, coord2, subj2 = model(x2, compute_heads=compute_heads)
  326. # Compute all losses in float32 for numerical stability
  327. z1f, z2f = z1.float(), z2.float()
  328. contr_loss = crit(z1f, z2f)
  329. # Start with contrastive loss
  330. loss = contr_loss
  331. if args.use_coord_adv:
  332. bins = args.coordHeadBins
  333. cx = digitize(coords[:, 0].contiguous(), edges_x, bins)
  334. cy = digitize(coords[:, 1].contiguous(), edges_y, bins)
  335. cz = digitize(coords[:, 2].contiguous(), edges_z, bins)
  336. logit_x1, logit_y1, logit_z1 = coord1
  337. logit_x2, logit_y2, logit_z2 = coord2
  338. # CE on float32 logits for stability
  339. coord_loss1 = (
  340. F.cross_entropy(logit_x1.float(), cx)
  341. + F.cross_entropy(logit_y1.float(), cy)
  342. + F.cross_entropy(logit_z1.float(), cz)
  343. ) / 3.0
  344. coord_loss2 = (
  345. F.cross_entropy(logit_x2.float(), cx)
  346. + F.cross_entropy(logit_y2.float(), cy)
  347. + F.cross_entropy(logit_z2.float(), cz)
  348. ) / 3.0
  349. coord_loss = (coord_loss1 + coord_loss2) / 2.0
  350. else:
  351. coord_loss = torch.tensor(0.0, device=device)
  352. if args.use_subject_adv:
  353. subj_loss1 = (
  354. F.cross_entropy(subj1.float(), sids)
  355. if subj1 is not None
  356. else torch.tensor(0.0, device=device)
  357. )
  358. subj_loss2 = (
  359. F.cross_entropy(subj2.float(), sids)
  360. if subj2 is not None
  361. else torch.tensor(0.0, device=device)
  362. )
  363. subj_loss = (subj_loss1 + subj_loss2) / 2.0
  364. else:
  365. subj_loss = torch.tensor(0.0, device=device)
  366. # Include head losses so heads train; scale by effective weights; encoder protection via --stop_grad_from_heads
  367. if args.use_coord_adv:
  368. loss = loss + eff_w_coord * coord_loss
  369. if args.use_subject_adv:
  370. loss = loss + eff_w_subj * subj_loss
  371. # Backward + step (no GradScaler for BF16 path)
  372. loss.backward()
  373. optimizer.step()
  374. sched.step()
  375. if step % 250 == 0 or step == args.steps - 1:
  376. loss_history.append(loss.item())
  377. # Store component losses and the step for plotting
  378. contr_history.append(float(contr_loss.item()))
  379. coord_history.append(float(coord_loss.item()))
  380. subj_history.append(float(subj_loss.item()))
  381. log_steps.append(int(step))
  382. # Lightweight monitoring metrics for logging cadence
  383. chance_logs = []
  384. with torch.no_grad():
  385. if args.use_subject_adv and (subj1 is not None and subj2 is not None):
  386. K = dataset.n_subj
  387. subj_acc1 = softmax_top1_acc(subj1, sids).item()
  388. subj_acc2 = softmax_top1_acc(subj2, sids).item()
  389. subj_acc = 0.5*(subj_acc1 + subj_acc2)
  390. subj_ch_acc = 1.0 / K
  391. subj_ch_ce = np.log(K)
  392. chance_logs.append(f"SubjAcc {subj_acc:.3f} (~{subj_ch_acc:.3f}), CE~{subj_ch_ce:.3f}")
  393. if args.use_coord_adv and (coord1 is not None and coord2 is not None):
  394. bins = args.coordHeadBins
  395. cx = digitize(coords[:, 0].contiguous(), edges_x, bins)
  396. cy = digitize(coords[:, 1].contiguous(), edges_y, bins)
  397. cz = digitize(coords[:, 2].contiguous(), edges_z, bins)
  398. (logit_x1, logit_y1, logit_z1) = coord1
  399. (logit_x2, logit_y2, logit_z2) = coord2
  400. ax_accs = []
  401. for (lx, ly, lz), (cx_, cy_, cz_) in [((logit_x1,logit_y1,logit_z1),(cx,cy,cz)),
  402. ((logit_x2,logit_y2,logit_z2),(cx,cy,cz))]:
  403. ax_accs.append((
  404. softmax_top1_acc(lx, cx_).item(),
  405. softmax_top1_acc(ly, cy_).item(),
  406. softmax_top1_acc(lz, cz_).item()
  407. ))
  408. ax_acc = tuple(np.mean([a[i] for a in ax_accs]) for i in range(3))
  409. coord_ch_acc = 1.0 / bins
  410. coord_ch_ce = np.log(bins)
  411. chance_logs.append(f"CoordAcc x/y/z {ax_acc[0]:.3f}/{ax_acc[1]:.3f}/{ax_acc[2]:.3f} (~{coord_ch_acc:.3f}), CE~{coord_ch_ce:.3f}")
  412. chance_str = " | ".join(chance_logs)
  413. print(
  414. f"Step {step:5d} | ramp {ramp:0.2f} | "
  415. f"Wc {eff_w_coord:0.2f} Ws {eff_w_subj:0.2f} | "
  416. f"ContLoss: {contr_loss.item():.4f} | "
  417. f"CoordLoss: {coord_loss.item():.4f} | "
  418. f"SubjLoss: {subj_loss.item():.4f} | "
  419. f"loss {loss.item():.4f} |"
  420. f"Alignment: {alignment(z1, z2):.4f} |"
  421. f"Uniformity z1: {uniformity(z1):.4f} |"
  422. f"Uniformity z2: {uniformity(z2):.4f} |"
  423. f"{chance_str}"
  424. )
  425. # Plot component losses + dotted sum line with legend
  426. save_loss_plot(exp_dir, log_steps, contr_history, coord_history, subj_history)
  427. # Save checkpoint and quick similarity stats
  428. torch.save({
  429. 'adv_model_state_dict': model.state_dict(),
  430. 'base_encoder_state_dict': base_encoder.state_dict(),
  431. 'optimizer_state_dict': optimizer.state_dict(),
  432. 'args': args,
  433. 'loss_history': loss_history
  434. }, exp_dir / "model_checkpoint.pth")
  435. with torch.no_grad():
  436. pos = (z1f * z2f).sum(dim=1) # cosine if z’s are L2-normalized of positive pairs
  437. idx = torch.randperm(z1f.size(0))[:2048] # subset for negatives
  438. z1s, z2s = z1f[idx], z2f[idx]
  439. neg = (z1s @ z2s.t())
  440. neg = neg[~torch.eye(neg.size(0), dtype=bool, device=neg.device)]
  441. print(f"Pos cos μ={pos.mean():.3f}, σ={pos.std():.3f} | Neg cos μ={neg.mean():.3f}, σ={neg.std():.3f}")
  442. S = z1f @ z2f.t()
  443. top1 = S.argmax(dim=1)
  444. inst_r1 = (top1 == torch.arange(S.size(0), device=S.device)).float().mean().item()
  445. print(f"Instance R@1: {inst_r1:.3f}")
  446. # Save loss history
  447. np.save(exp_dir / "loss_history.npy", np.array(loss_history))
  448. print(f"\nExperiment completed! Results saved to: {exp_dir}")
  449. def save_loss_plot(exp_dir: Path, steps: list, contr_hist: list, coord_hist: list, subj_hist: list):
  450. """Save a plot with separate loss components and a dotted sum line.
  451. This version avoids stretching artefacts in SVG viewers by:
  452. - keeping annotations inside the axes range
  453. - using a slightly more square aspect ratio
  454. - relaxing the tight bounding box so the axes, not labels, define width.
  455. Args:
  456. exp_dir: Output directory where the plot will be saved.
  457. steps: Logged training steps.
  458. contr_hist: Contrastive loss values.
  459. coord_hist: Coordinate loss values.
  460. subj_hist: Subject loss values.
  461. """
  462. steps_np = np.asarray(steps, dtype=float)
  463. if steps_np.size == 0:
  464. return
  465. contr_np = np.asarray(contr_hist, dtype=float)
  466. coord_np = np.asarray(coord_hist, dtype=float)
  467. subj_np = np.asarray(subj_hist, dtype=float)
  468. total_np = contr_np + coord_np + subj_np
  469. # Use a restrained, publication-friendly palette (ColorBrewer-inspired)
  470. colors = {
  471. "Contrastive": "#1b9e77",
  472. "Coordinate": "#d95f02",
  473. "Subject": "#7570b3",
  474. "Total": "#000000",
  475. }
  476. # Slightly more square figure for better appearance in wide editors
  477. fig, ax = plt.subplots(figsize=(5.0, 3.8), constrained_layout=True)
  478. # Plot individual components
  479. ax.plot(steps_np, contr_np, label="Contrastive", color=colors["Contrastive"], lw=1.8)
  480. ax.plot(steps_np, coord_np, label="Coordinate", color=colors["Coordinate"], lw=1.8)
  481. ax.plot(steps_np, subj_np, label="Subject", color=colors["Subject"], lw=1.8)
  482. # Dotted total line (subtle)
  483. ax.plot(steps_np, total_np, label="Total", color=colors["Total"], lw=1.4,
  484. linestyle=(0, (3, 2)))
  485. # Axis styling: remove clutter, subtle horizontal grid only
  486. ax.spines['top'].set_visible(False)
  487. ax.spines['right'].set_visible(False)
  488. ax.grid(axis='y', color="#dddddd", linewidth=0.6, alpha=0.6)
  489. ax.grid(axis='x', alpha=0.0)
  490. ax.set_xlabel("Step", labelpad=6)
  491. ax.set_ylabel("Loss", labelpad=6)
  492. ax.set_title("Training Loss Components", pad=10)
  493. # Ticks: modest sizing
  494. ax.tick_params(axis='both', labelsize=10, width=0.8, length=4)
  495. # Annotate last value of each component *inside* the axes to avoid
  496. # expanding the SVG width excessively when viewing the figure.
  497. if steps_np.size > 1:
  498. x_last = steps_np[-1]
  499. x_prev = steps_np[-2]
  500. # small offset but stay within [x_prev, x_last]
  501. x_off = 0.25 * (x_last - x_prev)
  502. x_anno = x_last - x_off
  503. else:
  504. x_anno = steps_np[0]
  505. for arr, name in [(contr_np, "Contrastive"), (coord_np, "Coordinate"), (subj_np, "Subject")]:
  506. if arr.size:
  507. ax.text(x_anno, arr[-1], f"{arr[-1]:.3g}",
  508. color=colors[name], fontsize=9, va='center', ha='right')
  509. # Legend without frame
  510. ax.legend(frameon=False, fontsize=10, handlelength=3, loc='upper right')
  511. # Save without an overly tight bbox so labels don't dominate width
  512. out_path = Path(exp_dir) / "loss_plot.svg"
  513. fig.savefig(out_path, dpi=300)
  514. plt.close(fig)
  515. if __name__ == "__main__":
  516. main()

trainer.py at commit a229605, under MIT · at the source

Overview

Authors: Mert Özer1,2, Bernhard Egger1, Angelika Mennecke3, Armin Nagel4, Moritz Zaiss3, Frederik Bernd Laun4, Arnd Dörfler3, Jürgen Winkler2, Alexander German2,3
  1. Department of Computer Science, Friedrich-Alexander-Universität Erlangen-Nürnberg, Erlangen, Germany
  2. Department of Molecular Neurology, University Hospital Erlangen, Erlangen, Germany
  3. Institute of Neuroradiology, University Hospital Erlangen, Erlangen, Germany
  4. Institute of Radiology, University Hospital Erlangen, Erlangen, Germany
Journal: Imaging neuroscience (Cambridge, Mass.), volume 4, article IMAG.a.1241
Dates: received 28 January 2026; accepted 13 April 2026; published online 26 May 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1162/imag.a.1241 · PMID 42212222 · PMCID PMC13214570 · OpenAlex W7156079084
Open access: diamond, a free copy (OpenAlex)
Status: code verified
Categories: structural MRI / diffusion (modality), human (organism), Parkinson's (population)
Methods: Spectral & time-frequency, Statistics, Machine learning, Preprocessing, fMRI & imaging
Keywords: Parkinson’s disease, dopamine transporter, 7 Tesla MRI, multispectral MRI, contrastive learning
Topic: Parkinson's Disease Mechanisms and Treatments (Neurology, Medicine), according to OpenAlex
Funding: Deutsche Forschungsgemeinschaft (500888779, 505539112, RU5534, KFO5024)
Citations: not cited yet (Europe PMC); 61 references in the paper

Abstract

The neuropathological hallmark of Parkinson’s disease (PD) is a progressive degeneration of dopaminergic neurons in the substantia nigra resulting in a reduced striatal dopamine transporter (DaT) concentration. Radioligands can detect degeneration of nigrostriatal projections early in the disease course, but their clinical use is constrained by ionizing radiation exposure and the need for radiotracer production and associated logistics. In this study, we investigate whether a healthy DaT atlas can be predicted directly from native 7 Tesla (7T) multicontrast magnetic resonance imaging (MRI) using a normative atlas as the supervised target for voxel-wise DaT density. Based on these results, we estimate the deviation of sporadic PD patients from the normative DaT atlas. Using this approach, we found that estimated DaT distributions in the putamen differ between PD patients and healthy controls, providing preliminary evidence that high-field advanced multispectral MRI could inform on neurochemical alterations in PD.

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

mert-o/dat_prediction

License: MIT
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: a229605f6bc0b7220a6c5e52bcaa13473de3ff74, 28 January 2026
Languages: Python (13)
Size: 18 files, 13 scripts
Software Heritage: not archived
Found in: “Data and Code Availability”
Holds: README, license file, environment (requirements.txt)
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: NumPy (11 files), PyTorch (9 files), NiBabel (5 files), Matplotlib (4 files), Nilearn (2 files), pandas (2 files), scikit-learn (2 files), SciPy (2 files), statsmodels (2 files)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
15 files

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

Tracing map

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

What the map holds:

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

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

Data

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

Data and Code Availability

The data used in this study are available from the corresponding authors upon reasonable request. The preprocessing, training, and evaluation code are available at: https://github.com/mert-o/dat_prediction.

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

  • Funding: added Deutsche Forschungsgemeinschaft: 500888779, 505539112, RU5534, KFO5024

Version 1, 28 September 2026: the first record

Recorded: type, language, journal, volume, pages, dates, 9 authors, 5 keywords, 54 references.

Cite

This paper

Özer, M., Egger, B., Mennecke, A., Nagel, A., Zaiss, M., Laun, F. B., Dörfler, A., Winkler, J., & German, A. (2026). Multispectral 7 Tesla MRI as a potential predictor of dopamine transporter deficiency in Parkinson's disease. Imaging neuroscience (Cambridge, Mass.), 4, IMAG.a.1241. https://doi.org/10.1162/imag.a.1241

BibTeX

@article{ozer2026multispectral,
author = {Özer, Mert and Egger, Bernhard and Mennecke, Angelika and Nagel, Armin and Zaiss, Moritz and Laun, Frederik Bernd and Dörfler, Arnd and Winkler, Jürgen and German, Alexander},
title = {{Multispectral 7 Tesla MRI as a potential predictor of dopamine transporter deficiency in Parkinson's disease}},
journal = {Imaging neuroscience (Cambridge, Mass.)},
year = {2026},
month = may,
volume = {4},
pages = {IMAG.a.1241},
publisher = {MIT Press},
issn = {2837-6056},
doi = {10.1162/imag.a.1241},
url = {https://doi.org/10.1162/imag.a.1241},
pmid = {42212222},
pmcid = {PMC13214570}
}

RIS

TY - JOUR
AU - Özer, Mert
AU - Egger, Bernhard
AU - Mennecke, Angelika
AU - Nagel, Armin
AU - Zaiss, Moritz
AU - Laun, Frederik Bernd
AU - Dörfler, Arnd
AU - Winkler, Jürgen
AU - German, Alexander
TI - Multispectral 7 Tesla MRI as a potential predictor of dopamine transporter deficiency in Parkinson's disease
T2 - Imaging neuroscience (Cambridge, Mass.)
J2 - Imaging Neurosci (Camb)
PY - 2026
DA - 2026/05/26
VL - 4
SP - IMAG.a.1241
SN - 2837-6056
PB - MIT Press
DO - 10.1162/imag.a.1241
UR - https://doi.org/10.1162/imag.a.1241
LA - en
ER -

CSL-JSON

{
"id": "10.1162/imag.a.1241",
"type": "article-journal",
"title": "Multispectral 7 Tesla MRI as a potential predictor of dopamine transporter deficiency in Parkinson's disease",
"container-title": "Imaging neuroscience (Cambridge, Mass.)",
"author": [
{
"family": "Özer",
"given": "Mert"
},
{
"family": "Egger",
"given": "Bernhard"
},
{
"family": "Mennecke",
"given": "Angelika"
},
{
"family": "Nagel",
"given": "Armin"
},
{
"family": "Zaiss",
"given": "Moritz"
},
{
"family": "Laun",
"given": "Frederik Bernd"
},
{
"family": "Dörfler",
"given": "Arnd"
},
{
"family": "Winkler",
"given": "Jürgen"
},
{
"family": "German",
"given": "Alexander"
}
],
"container-title-short": "Imaging Neurosci (Camb)",
"volume": "4",
"page": "IMAG.a.1241",
"DOI": "10.1162/imag.a.1241",
"PMID": "42212222",
"PMCID": "PMC13214570",
"ISSN": "2837-6056",
"publisher": "MIT Press",
"URL": "https://doi.org/10.1162/imag.a.1241",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
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.1371/journal.pbio.3003856 [code]
Aging and metabolism contribute separately to brain-body health.
Journal: PLoS biology
In common: Nilearn, NiBabel, statsmodels, 6 other tools, structural MRI / diffusion, 1 reference
[2] doi:10.1007/s00234-026-04103-8 [code]
Enhanced detection of subtle cortical abnormalities in focal epilepsy using 7 T MRI surface-based models and graph neural networks.
Journal: Neuroradiology
In common: Nilearn, NiBabel, statsmodels, 6 other tools, structural MRI / diffusion, 1 reference
[3] doi:10.1002/agm2.70073 [code]
Foreign Language Learning in Older Adults Modifies Resting-State Functional Connectivity Between the Subcortical Structures and the Cortex.
Journal: Aging medicine (Milton (N.S.W))
In common: Nilearn, NiBabel, statsmodels, 6 other tools, 1 reference
[4] doi:10.1038/s41467-026-76452-0 [code]
Music evokes shared neural representations of imagined narratives across sensory modalities.
Journal: Nature communications
In common: Nilearn, NiBabel, statsmodels, 6 other tools, 1 reference
[5] doi:10.1016/j.isci.2026.117180 [code]
Developmental changes in similarity between neural representations of mental arithmetic and artificial neural networks.
Journal: iScience
In common: Nilearn, NiBabel, statsmodels, 6 other tools, 1 reference
[6] doi:10.1162/imag.a.1256 [code]
Gamer in the scanner: Event-related analysis of fMRI activity during retro videogame play guided by automated annotations of game content.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: Nilearn, NiBabel, statsmodels, 6 other tools, 1 reference
[7] doi:10.7554/elife.107933 [code]
Modality-agnostic decoding of vision and language from fMRI.
Journal: eLife
In common: Nilearn, NiBabel, statsmodels, 6 other tools, 1 reference
[8] doi:10.1126/sciadv.adu9309 [code]
Variations of global brain asymmetry are associated with aging and related diseases.
Journal: Science advances
In common: Nilearn, NiBabel, PyTorch, 4 other tools, Parkinson's, 2 references
[9] doi:10.1038/s41531-026-01354-3 [code]
Neuromodulation-induced normalization of cortical metastable dynamics signatures in Parkinson's disease.
Journal: NPJ Parkinson's disease
In common: Nilearn, NiBabel, statsmodels, 5 other tools, Parkinson's, 1 reference
[10] doi:10.64898/2026.08.18.26360725 [code]
Temporal pole blurring in hippocampal sclerosis reflects seizure-disrupted myelination
Journal: medRxiv (preprint)
In common: Nilearn, NiBabel, statsmodels, 6 other tools, structural MRI / diffusion

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.