OSCR

Robust disease prognosis via diagnostic knowledge preservation: A sequential learning approach.

Code ↔ Paper

5 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 5 matches
  1. [1] § 2. Materials and methods › 2.5. Experimental setup › 2.5.2. Training protocol. ↔ Breast_Cancer/Multitask/train.py, lines 592–660 · score 0.77 · cross entropy loss, learning rate scheduler, BI RADS, Adam, breast cancer, optimizer
  2. [2] § 2. Materials and methods › 2.5. Experimental setup › 2.5.3. Evaluation and analysis. ↔ Breast_Cancer/Single_Task/eval_birads.py, lines 108–197 · score 0.76 · Micro AUROC, Macro AUROC, multi class, Balanced accuracy, predictions
  3. [3] § 2. Materials and methods › 2.5. Experimental setup › 2.5.3. Evaluation and analysis. ↔ Breast_Cancer/Multitask/evaluate.py, lines 28–113 · score 0.70 · Macro AUROC, multi class, Balanced accuracy, sensitivity, micro, predictions
  4. [4] § 2. Materials and methods › 2.5. Experimental setup › 2.5.2. Training protocol. ↔ Knee_OA/train.py, lines 90–177 · score 0.67 · cosine annealing learning, Adam, scheduler, augmentation, optimizer, loss
  5. [5] § 2. Materials and methods › 2.5. Experimental setup › 2.5.2. Training protocol. ↔ Breast_Cancer/Multitask/train.py, lines 592–660 · score 0.52 · prognosis AUROC, multitask learning, epoch, probability, loss, replay

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 · 796 lines · 31 KB · CC-BY-NC-4.0 · 2 matches

  1. import os
  2. import wandb
  3. import torch
  4. import torch.nn as nn
  5. import torch.optim as optim
  6. import torch.nn.functional as F
  7. from torch.utils.data import DataLoader, Dataset
  8. from sklearn.metrics import roc_auc_score, average_precision_score
  9. import numpy as np
  10. from tqdm import tqdm
  11. import argparse
  12. import pickle
  13. from models_torch import BaselineBreastModel
  14. from torch_transforms import compose_transform
  15. import load_mammogram
  16. os.environ['WANDB_API_KEY'] = #
  17. class MammogramDataset(Dataset):
  18. """Dataset class for mammogram images with multi-task support"""
  19. def __init__(self, img_dir, datalist, transform, crop_size=(2944, 1920), tasks=None):
  20. self.img_dir = img_dir
  21. self.datalist = datalist
  22. self.transform = transform
  23. self.crop_size = crop_size
  24. self.tasks = tasks if tasks else ["prognosis"] # Default to single task if not specified
  25. # Filter out samples with missing labels for the required tasks
  26. self.valid_indices = []
  27. for idx, metadata in enumerate(self.datalist):
  28. is_valid = True
  29. # Check if BI-RADS label is valid if birads task is requested
  30. if "birads" in self.tasks:
  31. birads_value = metadata.get('birads', None)
  32. if birads_value is None or birads_value == 'NA' or birads_value == 'inconclusive':
  33. is_valid = False
  34. if is_valid:
  35. self.valid_indices.append(idx)
  36. print(f"Dataset initialized with {len(self.valid_indices)}/{len(datalist)} valid samples for tasks: {tasks}")
  37. def __len__(self):
  38. return len(self.valid_indices)
  39. def __getitem__(self, idx):
  40. original_idx = self.valid_indices[idx]
  41. metadata = self.datalist[original_idx]
  42. # Load mammogram exam
  43. images, scans_info = load_mammogram.load_mammogram_exam(
  44. img_dir=self.img_dir,
  45. metadata=metadata,
  46. transform=self.transform,
  47. select_index_logic=("random", 4),
  48. crop_method="best_center",
  49. crop_size=self.crop_size,
  50. skip_img=False
  51. )
  52. # Prepare image data
  53. x = {}
  54. for i, scan in enumerate(scans_info):
  55. side = scan['laterality']
  56. if side == "left":
  57. x[f"L-{scan['view']}"] = images[i,0:1,:,:]
  58. else:
  59. x[f"R-{scan['view']}"] = images[i,0:1,:,:]
  60. # Prepare labels for all requested tasks
  61. labels = {}
  62. if "prognosis" in self.tasks:
  63. labels["prognosis"] = metadata['5yr_cancer_label']['malignant']
  64. if "birads" in self.tasks and 'birads' in metadata and metadata['birads'] not in [None, 3, 4, 5]:
  65. labels["birads"] = metadata['birads']
  66. else:
  67. labels["birads"] = -1
  68. # Return different outputs based on the number of tasks
  69. if len(self.tasks) == 1:
  70. # For single task, return a single label for backward compatibility
  71. return x, labels[self.tasks[0]]
  72. else:
  73. # For multi-task, return a dictionary of labels
  74. return x, labels
  75. def collate_mammogram_batch(batch):
  76. """
  77. Custom collate function for mammogram batches with single or multiple labels per sample.
  78. Args:
  79. batch: List of tuples (x, label) where:
  80. x: Dictionary with keys 'L-CC', 'L-MLO', 'R-CC', 'R-MLO' mapping to image tensors
  81. label: Either a single value or a dictionary of task labels
  82. Returns:
  83. Tuple of:
  84. batch_x: Dictionary mapping view names to batched image tensors
  85. batch_labels: Dictionary of label tensors or single tensor depending on input format
  86. """
  87. # Initialize dictionary for images
  88. batch_x = {
  89. 'L-CC': [], 'L-MLO': [],
  90. 'R-CC': [], 'R-MLO': []
  91. }
  92. # Check if using multi-task labels (dictionary) or single task label
  93. first_label = batch[0][1]
  94. is_multi_task = isinstance(first_label, dict)
  95. if is_multi_task:
  96. # Initialize dictionary for each task
  97. tasks = first_label.keys()
  98. batch_labels = {task: [] for task in tasks}
  99. # Collect images and task labels
  100. for x, labels in batch:
  101. for view in batch_x.keys():
  102. if view in x:
  103. batch_x[view].append(x[view])
  104. else:
  105. # Handle missing views with a tensor of zeros
  106. # Assuming all tensors in a batch have the same shape
  107. shape = next(iter(x.values())).shape
  108. batch_x[view].append(torch.zeros(shape))
  109. for task, label in labels.items():
  110. batch_labels[task].append(label)
  111. # Stack tensors for each view and create label tensors
  112. batch_x = {view: torch.stack(tensors) for view, tensors in batch_x.items()}
  113. batch_labels = {task: torch.tensor(labels) for task, labels in batch_labels.items()}
  114. else:
  115. # Single task mode - for backward compatibility
  116. batch_labels = []
  117. # Collect images and labels
  118. for x, label in batch:
  119. for view in batch_x.keys():
  120. if view in x:
  121. batch_x[view].append(x[view])
  122. else:
  123. # Handle missing views with a tensor of zeros
  124. shape = next(iter(x.values())).shape
  125. batch_x[view].append(torch.zeros(shape))
  126. batch_labels.append(label)
  127. # Stack tensors for each view and create label tensor
  128. batch_x = {view: torch.stack(tensors) for view, tensors in batch_x.items()}
  129. batch_labels = torch.tensor(batch_labels)
  130. return batch_x, batch_labels
  131. def unpickle_from_file(file_name):
  132. with open(file_name, 'rb') as handle:
  133. try:
  134. return pickle.load(handle)
  135. except ImportError:
  136. return pd.read_pickle(file_name)
  137. class MultiTaskBreastModel(nn.Module):
  138. """Wrapper for baseline breast model to support multiple output heads"""
  139. def __init__(self, base_model, device, tasks=None):
  140. super(MultiTaskBreastModel, self).__init__()
  141. self.tasks = tasks if tasks else ["prognosis"]
  142. self.base_model = base_model
  143. original_fc2 = self.base_model.fc2
  144. # Base model
  145. # Get feature dimension from base model
  146. self.feature_dim = self.base_model.fc2.in_features
  147. # Create task-specific output heads
  148. self.heads = nn.ModuleDict()
  149. if "prognosis" in self.tasks:
  150. self.heads["prognosis"] = nn.Linear(self.feature_dim, 1)
  151. if "birads" in self.tasks:
  152. self.heads["birads"] = nn.Linear(original_fc2.in_features, original_fc2.out_features)
  153. # Optionally copy weights and biases
  154. self.heads["birads"].weight.data = original_fc2.weight.data.clone()
  155. self.heads["birads"].bias.data = original_fc2.bias.data.clone()
  156. self.base_model.fc2 = nn.Identity()
  157. def forward(self, x):
  158. # Extract features from base model
  159. features = self.base_model(x)
  160. # Apply task-specific heads
  161. outputs = {}
  162. for task in self.tasks:
  163. outputs[task] = self.heads[task](features)
  164. # Reshape binary outputs to match expected format
  165. if task == "prognosis":
  166. outputs[task] = outputs[task].squeeze(-1)
  167. return outputs
  168. def get_non_requires_grad_params(model):
  169. """
  170. Find all parameters in a model that do NOT require gradients.
  171. Args:
  172. model: PyTorch model
  173. Returns:
  174. List of parameter names that don't require gradients and their shapes
  175. """
  176. non_requires_grad_params = []
  177. for name, param in model.named_parameters():
  178. if not param.requires_grad:
  179. non_requires_grad_params.append((name, param.shape))
  180. return non_requires_grad_params
  181. def load_pretrained_weights(model, pretrained_path):
  182. """Load pretrained BI-RADS weights and modify for binary classification"""
  183. if pretrained_path:
  184. pretrained_dict = torch.load(pretrained_path, map_location='cpu')
  185. model_dict = model.state_dict()
  186. # Filter out fc2 layer from pretrained weights
  187. pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict and 'fc2' not in k}
  188. # Update model weights
  189. model_dict.update(pretrained_dict)
  190. model.load_state_dict(model_dict)
  191. return model
  192. def train_epoch(model, train_loader, criterions, optimizer, device, tasks, diagnosis_loader=None, replay_prob=0.5):
  193. print(tasks)
  194. model.train()
  195. running_loss = 0.0
  196. task_running_losses = {task: 0.0 for task in tasks}
  197. # For metrics tracking - only track prognosis for simplicity
  198. task_metrics = {}
  199. if 'prognosis' in tasks:
  200. task_metrics['prognosis'] = {
  201. 'predictions': [],
  202. 'targets': []
  203. }
  204. if diagnosis_loader is not None:
  205. task_running_losses['birads'] = 0.0
  206. # Set up diagnosis iterator if provided
  207. diagnosis_iterator = iter(diagnosis_loader) if diagnosis_loader else None
  208. pbar = tqdm(train_loader, desc='Training')
  209. for batch_idx, (data, targets) in enumerate(pbar):
  210. # Determine whether to do replay for this batch
  211. do_replay = diagnosis_loader is not None and np.random.uniform() <= replay_prob
  212. if do_replay:
  213. try:
  214. # Get a batch from diagnosis dataset
  215. print("its here")
  216. replay_data, replay_targets = next(diagnosis_iterator)
  217. # Replace current batch with diagnosis data
  218. data = replay_data
  219. # Handle targets format
  220. if isinstance(replay_targets, dict):
  221. targets = replay_targets
  222. else:
  223. # Convert to dict format if needed
  224. targets = {"birads": replay_targets}
  225. # Set current tasks for this batch
  226. current_tasks = ["birads"]
  227. except StopIteration:
  228. # Reset iterator if we've gone through all diagnosis data
  229. diagnosis_iterator = iter(diagnosis_loader)
  230. replay_data, replay_targets = next(diagnosis_iterator)
  231. # Replace current batch with diagnosis data
  232. data = replay_data
  233. # Handle targets format
  234. if isinstance(replay_targets, dict):
  235. targets = replay_targets
  236. else:
  237. # Convert to dict format if needed
  238. targets = {"birads": replay_targets}
  239. # Set current tasks for this batch
  240. current_tasks = ["birads"]
  241. else:
  242. # Regular prognosis batch
  243. current_tasks = tasks
  244. # Move data to device
  245. for view in data:
  246. data[view] = data[view].to(device)
  247. # Handle targets based on format
  248. if not isinstance(targets, dict):
  249. # Convert single task target to dict format for consistent handling
  250. # Assuming single task is always the first task in the list
  251. targets = {current_tasks[0]: targets.to(device)}
  252. else:
  253. # Move all targets to device
  254. for task in targets:
  255. if task in targets: # Only process tasks that have labels
  256. targets[task] = targets[task].to(device)
  257. # Convert birads to long for CrossEntropyLoss
  258. if task == "birads":
  259. targets[task] = targets[task].long()
  260. optimizer.zero_grad()
  261. outputs = model(data)
  262. # Calculate task-specific losses and total loss
  263. total_loss = 0.0
  264. for task in current_tasks:
  265. if task in targets and task in outputs:
  266. mask = (targets[task] != -1)
  267. if isinstance(mask, bool):
  268. if not mask:
  269. continue
  270. mask = torch.tensor([True]) # Create a single-element boolean tensor
  271. elif mask.sum() == 0:
  272. continue
  273. if isinstance(criterions[task], nn.BCEWithLogitsLoss):
  274. # For binary cross entropy, both outputs and targets should be float
  275. task_loss = criterions[task](
  276. outputs[task][mask].float(),
  277. targets[task][mask].float()
  278. )
  279. elif isinstance(criterions[task], nn.CrossEntropyLoss):
  280. # For cross entropy, outputs should be float but targets should be long
  281. task_loss = criterions[task](
  282. outputs[task][mask].float(),
  283. targets[task][mask].long()
  284. )
  285. # print("birads targets ", targets[task][mask].long())
  286. else:
  287. # For other loss functions, maintain original types
  288. task_loss = criterions[task](outputs[task][mask], targets[task][mask])
  289. task_running_losses[task] += task_loss.item()
  290. total_loss += task_loss
  291. # Store predictions for prognosis metrics only (if it's a prognosis batch)
  292. if task == 'prognosis' and task in task_metrics and task in current_tasks:
  293. task_metrics[task]['predictions'].extend(torch.sigmoid(outputs[task]).detach().cpu().numpy())
  294. task_metrics[task]['targets'].extend(targets[task].cpu().numpy())
  295. total_loss.backward()
  296. optimizer.step()
  297. running_loss += total_loss.item()
  298. # Update progress bar
  299. pbar.set_postfix({
  300. 'loss': total_loss.item(),
  301. **{f"{task}_loss": task_running_losses[task]/(batch_idx+1) for task in tasks if task in task_running_losses}
  302. })
  303. # Calculate average losses
  304. epoch_loss = running_loss / len(train_loader)
  305. task_losses = {task: task_running_losses[task] / len(train_loader) for task in tasks if task in task_running_losses}
  306. # Calculate metrics for prognosis task only
  307. metrics = {}
  308. if 'prognosis' in task_metrics and len(task_metrics['prognosis']['predictions']) > 0:
  309. preds = np.array(task_metrics['prognosis']['predictions'])
  310. targets = np.array(task_metrics['prognosis']['targets'])
  311. metrics["prognosis_auroc"] = roc_auc_score(targets, preds)
  312. metrics["prognosis_auprc"] = average_precision_score(targets, preds)
  313. return epoch_loss, task_losses, metrics
  314. def validate(model, val_loader, criterions, device, tasks):
  315. model.eval()
  316. running_loss = 0.0
  317. task_running_losses = {task: 0.0 for task in tasks}
  318. # For metrics tracking - only track prognosis
  319. task_metrics = {}
  320. if 'prognosis' in tasks:
  321. task_metrics['prognosis'] = {
  322. 'predictions': [],
  323. 'targets': []
  324. }
  325. with torch.no_grad():
  326. for data, targets in tqdm(val_loader, desc="Validation"):
  327. # Move data to device
  328. for view in data:
  329. data[view] = data[view].to(device)
  330. # Handle both single-task and multi-task targets
  331. if not isinstance(targets, dict):
  332. # Convert single task target to dict format for consistent handling
  333. targets = {"prognosis": targets.to(device)}
  334. else:
  335. # Move all targets to device
  336. for task in targets:
  337. if task in targets: # Only process tasks that have labels
  338. targets[task] = targets[task].to(device)
  339. # Convert birads to long for CrossEntropyLoss
  340. if task == "birads":
  341. targets[task] = targets[task].long()
  342. outputs = model(data)
  343. # Calculate task-specific losses and total loss
  344. total_loss = 0.0
  345. for task in tasks:
  346. if task in targets and task in outputs:
  347. mask = (targets[task] != -1)
  348. if isinstance(mask, bool):
  349. if not mask:
  350. continue
  351. mask = torch.tensor([True]) # Create a single-element boolean tensor
  352. elif mask.sum() == 0:
  353. continue
  354. if isinstance(criterions[task], nn.BCEWithLogitsLoss):
  355. # For binary cross entropy, both outputs and targets should be float
  356. task_loss = criterions[task](
  357. outputs[task][mask].float(),
  358. targets[task][mask].float()
  359. )
  360. elif isinstance(criterions[task], nn.CrossEntropyLoss):
  361. # For cross entropy, outputs should be float but targets should be long
  362. task_loss = criterions[task](
  363. outputs[task][mask].float(),
  364. targets[task][mask].long()
  365. )
  366. else:
  367. # For other loss functions, maintain original types
  368. task_loss = criterions[task](outputs[task][mask], targets[task][mask])
  369. task_running_losses[task] += task_loss.item()
  370. total_loss += task_loss
  371. # Store predictions for prognosis metrics only
  372. if task == 'prognosis' and task in task_metrics:
  373. task_metrics[task]['predictions'].extend(torch.sigmoid(outputs[task]).detach().cpu().numpy())
  374. task_metrics[task]['targets'].extend(targets[task].cpu().numpy())
  375. running_loss += total_loss.item()
  376. # Calculate average losses
  377. epoch_loss = running_loss / len(val_loader)
  378. task_losses = {task: task_running_losses[task] / len(val_loader) for task in tasks if task in task_running_losses}
  379. # Calculate metrics for prognosis task only
  380. metrics = {}
  381. if 'prognosis' in task_metrics and len(task_metrics['prognosis']['predictions']) > 0:
  382. preds = np.array(task_metrics['prognosis']['predictions'])
  383. targets = np.array(task_metrics['prognosis']['targets'])
  384. metrics["prognosis_auroc"] = roc_auc_score(targets, preds)
  385. metrics["prognosis_auprc"] = average_precision_score(targets, preds)
  386. return epoch_loss, task_losses, metrics
  387. def train(args):
  388. wandb.init(project=args.wandb_project, name=args.wandb_run_name)
  389. wandb.config.update(args)
  390. device = torch.device("cuda" if args.device_type == "gpu" else "cpu")
  391. # Define tasks based on arguments
  392. tasks = ["prognosis"]
  393. if args.multitask:
  394. tasks.append("birads")
  395. print(f"Training with tasks: {tasks}")
  396. # Load prognosis data
  397. prognosis_train_data = unpickle_from_file(f"{args.prognosis_datalist_path}/Fold_{args.fold}_train.pkl")
  398. val_data = unpickle_from_file(f"{args.prognosis_datalist_path}/Fold_{args.fold}_val.pkl")
  399. # np.random.shuffle(prognosis_train_data)
  400. # prognosis_train_data = prognosis_train_data[:256]
  401. # np.random.shuffle(val_data)
  402. # val_data = val_data[:256]
  403. transform = compose_transform(augmentation="standard", resize=None, image_format="greyscale")
  404. # Create primary dataset for prognosis
  405. train_dataset = MammogramDataset(
  406. img_dir=args.img_dir,
  407. datalist=prognosis_train_data,
  408. transform=transform,
  409. crop_size=args.crop_size,
  410. tasks=tasks # Primary dataset includes all primary tasks
  411. )
  412. val_dataset = MammogramDataset(
  413. img_dir=args.img_dir,
  414. datalist=val_data,
  415. transform=compose_transform(image_format="greyscale"),
  416. crop_size=args.crop_size,
  417. tasks=["prognosis", "birads"] # Validation dataset includes all tasks
  418. )
  419. # Create data loaders for prognosis
  420. train_loader = DataLoader(
  421. train_dataset,
  422. batch_size=args.batch_size,
  423. shuffle=True,
  424. num_workers=args.num_workers,
  425. collate_fn=collate_mammogram_batch,
  426. pin_memory=True
  427. )
  428. val_loader = DataLoader(
  429. val_dataset,
  430. batch_size=args.batch_size,
  431. shuffle=False,
  432. num_workers=args.num_workers,
  433. collate_fn=collate_mammogram_batch,
  434. pin_memory=True
  435. )
  436. # Create diagnosis dataset (separate cohort)
  437. diagnosis_loader = None
  438. if args.replay:
  439. print("Setting up diagnosis loader for batch-level replay")
  440. # Load diagnosis data from a separate path
  441. diagnosis_train_data = unpickle_from_file(f"{args.diagnosis_datalist_path}/train.pkl")
  442. diagnosis_dataset = MammogramDataset(
  443. img_dir=args.img_dir,
  444. datalist=diagnosis_train_data,
  445. transform=transform,
  446. crop_size=args.crop_size,
  447. tasks=["birads"] # This dataset is only for diagnosis
  448. )
  449. diagnosis_loader = DataLoader(
  450. diagnosis_dataset,
  451. batch_size=args.batch_size,
  452. shuffle=True,
  453. num_workers=args.num_workers,
  454. collate_fn=collate_mammogram_batch,
  455. pin_memory=True
  456. )
  457. nodropout_probability = 1 - args.dropout_p
  458. base_model = BaselineBreastModel(
  459. device=device,
  460. nodropout_probability=nodropout_probability,
  461. gaussian_noise_std=args.gaussian_noise_std,
  462. )
  463. if args.pretrained_path:
  464. base_model = load_pretrained_weights(base_model, args.pretrained_path)
  465. model = MultiTaskBreastModel(
  466. base_model=base_model,
  467. device=device,
  468. tasks=["prognosis", "birads"] # Always include both tasks in model
  469. ).to(device)
  470. frozen_params = get_non_requires_grad_params(model)
  471. print(f"Total parameters NOT requiring gradients: {len(frozen_params)}")
  472. for name, shape in frozen_params:
  473. print(f"{name}: {shape}")
  474. # Set up class weights for BI-RADS using the provided counts
  475. class_counts = torch.tensor([6634, 17388,15939], dtype=torch.float).to(device)
  476. class_weights = 1.0 / class_counts
  477. class_weights = class_weights / class_weights.sum() * len(class_counts)
  478. # Setup loss functions for each task
  479. criterions = {
  480. "prognosis": nn.BCEWithLogitsLoss(),
  481. "birads": nn.CrossEntropyLoss(weight=class_weights)
  482. }
  483. optimizer = optim.Adam(model.parameters(), lr=args.learning_rate)
  484. scheduler = optim.lr_scheduler.ReduceLROnPlateau(
  485. optimizer, mode='max', factor=0.75, patience=10, verbose=True
  486. )
  487. # Create save directory
  488. save_path = os.path.join(args.save_path, f"Fold_{args.fold}", args.wandb_run_name)
  489. os.makedirs(save_path, exist_ok=True)
  490. # Track best metrics for each task
  491. best_metrics = {
  492. "total_loss": float('inf'),
  493. "prognosis_auroc": 0,
  494. "prognosis_auprc": 0,
  495. }
  496. patience_counter = 0
  497. # Training loop
  498. for epoch in range(args.epochs):
  499. print(f'\nEpoch {epoch+1}/{args.epochs}')
  500. # Training phase - with batch-level replay
  501. if args.replay and diagnosis_loader is not None:
  502. # Use new train_epoch function with batch-level replay
  503. total_train_loss, all_train_task_losses, all_train_metrics = train_epoch(
  504. model=model,
  505. train_loader=train_loader,
  506. criterions=criterions,
  507. optimizer=optimizer,
  508. device=device,
  509. tasks=["prognosis"],
  510. diagnosis_loader=diagnosis_loader,
  511. replay_prob=args.replay_prob
  512. )
  513. # Joint multitask learning
  514. elif args.multitask:
  515. # Train on both prognosis and birads tasks in the same dataset
  516. total_train_loss, all_train_task_losses, all_train_metrics = train_epoch(
  517. model=model,
  518. train_loader=train_loader,
  519. criterions=criterions,
  520. optimizer=optimizer,
  521. device=device,
  522. tasks=["prognosis", "birads"]
  523. )
  524. # Single task mode (prognosis only)
  525. else:
  526. # Standard single-task training
  527. total_train_loss, all_train_task_losses, all_train_metrics = train_epoch(
  528. model=model,
  529. train_loader=train_loader,
  530. criterions=criterions,
  531. optimizer=optimizer,
  532. device=device,
  533. tasks=["prognosis"]
  534. )
  535. # Validation phase - always validate on all tasks
  536. val_loss, val_task_losses, val_metrics = validate(
  537. model=model,
  538. val_loader=val_loader,
  539. criterions=criterions,
  540. device=device,
  541. tasks=["prognosis", "birads"]
  542. )
  543. # Log metrics
  544. log_dict = {
  545. 'epoch': epoch,
  546. 'train_loss': total_train_loss,
  547. 'val_loss': val_loss,
  548. **{f"train_{task}_loss": loss for task, loss in all_train_task_losses.items()},
  549. **{f"val_{task}_loss": loss for task, loss in val_task_losses.items()},
  550. **{f"train_{metric}": value for metric, value in all_train_metrics.items()},
  551. **{f"val_{metric}": value for metric, value in val_metrics.items()},
  552. 'learning_rate': optimizer.param_groups[0]['lr']
  553. }
  554. wandb.log(log_dict)
  555. # Print metrics
  556. print(f"Epoch {epoch+1} Results:")
  557. for metric, value in val_metrics.items():
  558. print(f" Val {metric}: {value:.4f}")
  559. # Update learning rate based on primary task metric (prognosis AUROC)
  560. if 'prognosis_auroc' in val_metrics:
  561. scheduler.step(val_metrics['prognosis_auroc'])
  562. # Check if this is the best model so far
  563. improved = False
  564. # Check each metric
  565. for metric in best_metrics:
  566. if metric == 'total_loss':
  567. if val_loss < best_metrics[metric]:
  568. best_metrics[metric] = val_loss
  569. improved = True
  570. torch.save(model.state_dict(), os.path.join(save_path, f'best_loss_model.pth'))
  571. elif metric in val_metrics and best_metrics[metric] is not None:
  572. if val_metrics[metric] > best_metrics[metric]:
  573. best_metrics[metric] = val_metrics[metric]
  574. improved = True
  575. torch.save(model.state_dict(), os.path.join(save_path, f'best_{metric}_model.pth'))
  576. # Also save a checkpoint for this epoch
  577. torch.save({
  578. 'epoch': epoch,
  579. 'model_state_dict': model.state_dict(),
  580. 'optimizer_state_dict': optimizer.state_dict(),
  581. 'scheduler_state_dict': scheduler.state_dict(),
  582. 'best_metrics': best_metrics,
  583. **val_metrics
  584. }, os.path.join(save_path, f'checkpoint_epoch_{epoch}.pth'))
  585. # Check for early stopping
  586. if improved:
  587. patience_counter = 0
  588. else:
  589. patience_counter += 1
  590. if patience_counter >= args.patience:
  591. print(f'Early stopping triggered after epoch {epoch+1}')
  592. break
  593. wandb.finish()
  594. # Save final model
  595. torch.save(model.state_dict(), os.path.join(save_path, 'final_model.pth'))
  596. # Print best results
  597. print("\nTraining completed. Best results:")
  598. for metric, value in best_metrics.items():
  599. if value is not None:
  600. print(f" Best {metric}: {value:.4f}")
  601. if __name__ == "__main__":
  602. parser = argparse.ArgumentParser(description='Train multi-task mammogram classification model')
  603. # Model parameters
  604. parser.add_argument('--pretrained-path', type=str, default=None,
  605. help='Path to pretrained model weights')
  606. parser.add_argument('--save-path', type=str, default=None,
  607. help='Path to save models')
  608. parser.add_argument('--gaussian-noise-std', type=float, default=0.01,
  609. help='Standard deviation of Gaussian noise')
  610. parser.add_argument('--dropout-p', type=float, default=0.1,
  611. help='Dropout probability')
  612. # Multi-task parameters
  613. parser.add_argument('--multitask', action='store_true',
  614. help='Enable multi-task learning with prognosis and diagnosis')
  615. parser.add_argument('--replay', action='store_true',
  616. help='Use experience replay for multi-task learning')
  617. parser.add_argument('--replay-prob', type=float, default=0.5,
  618. help='Probability of doing experience replay in each epoch')
  619. # Training parameters
  620. parser.add_argument('--batch-size', type=int, default=8)
  621. parser.add_argument('--learning-rate', type=float, default=1e-5)
  622. parser.add_argument('--epochs', type=int, default=100)
  623. parser.add_argument('--patience', type=int, default=10,
  624. help='Early stopping patience')
  625. parser.add_argument('--num-workers', type=int, default=8)
  626. # Device parameters
  627. parser.add_argument('--device-type', type=str, default='gpu')
  628. parser.add_argument('--gpu-number', type=int, default=0)
  629. # Prognosis data parameters
  630. parser.add_argument('--img-dir', type=str, required=True,
  631. help='Root directory containing the prognosis image data')
  632. parser.add_argument('--prognosis-datalist-path', type=str, required=True,
  633. help='Path to the pickle file containing prognosis train/val/test splits')
  634. # Diagnosis data parameters (for multitask or replay)
  635. parser.add_argument('--diagnosis-datalist-path', type=str,
  636. help='Path to the pickle file containing diagnosis train/val/test splits')
  637. # Common data parameters
  638. parser.add_argument('--crop-size', type=int, nargs=2, default=[2944, 1920],
  639. help='Crop size for images (height width)')
  640. parser.add_argument('--fold', type=int, default=1,
  641. help='which Fold to train on')
  642. # Wandb parameters
  643. parser.add_argument('--wandb-project', type=str, required=True,
  644. help='WandB project name')
  645. parser.add_argument('--wandb-run-name', type=str, required=True,
  646. help='WandB run name')
  647. args = parser.parse_args()
  648. train(args)

train.py at commit a1b230f, under CC-BY-NC-4.0 · at the source

Overview

Authors: Haresh Rengaraj Rajamohan1, Yanqi Xu1, Weicheng Zhu1, Richard Kijowski2, Kyunghyun Cho1, Krzysztof J. Geras1,3, Narges Razavian3,4, Cem M. Deniz3,5
  1. Center for Data Science, New York University, New York, New York, United States of America
  2. Department of Radiology, Hospital for Special Surgery, New York, New York, United States of America
  3. Department of Radiology, New York University Langone Health, New York, New York, United States of America
  4. Department of Population Health, New York University Langone Health, New York, New York, United States of America
  5. Bernard and Irene Schwartz Center for Biomedical Imaging, New York University Langone Health, New York, New York, United States of America
Institutions: New York University (United States); Hospital for Special Surgery (United States); NYU Langone Health (United States)
Journal: PloS one, volume 21, issue 5, article e0344600
Dates: received 26 September 2025; accepted 23 February 2026; published online 6 May 2026
Type: Research article · Language: English
License: CC BY
Identifiers: DOI 10.1371/journal.pone.0344600 · PMID 42090385 · PMCID PMC13148697 · OpenAlex W4414497706
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: human (organism), other condition (population), Alzheimer's / dementia (population), clinical / translational (subfield)
Methods: Machine learning, Statistics
MeSH: Alzheimer Disease*, Breast Neoplasms*, Deep Learning*, Osteoarthritis, Knee*, Disease Progression, Female, Humans, Magnetic Resonance Imaging, Predictive Learning Models, Prognosis, ROC Curve (* major topic)
Journal subjects: Medicine and Health Sciences, Diagnostic Medicine, Prognosis, Cancer Detection and Diagnosis, Oncology, Mental Health and Psychiatry, Dementia, Alzheimer's Disease, Neurology, Medical Conditions, Neurodegenerative Diseases, Rheumatology, Arthritis, Osteoarthritis, Biology and Life Sciences, Anatomy, Musculoskeletal System, Skeleton, Skeletal Joints, Knees, Body Limbs, Legs, Cancers and Neoplasms, Breast Tumors, Breast Cancer, Neuroscience, Cognitive Science, Cognitive Psychology, Learning, Psychology, Social Sciences, Learning and Memory
Topic: AI in cancer detection (Artificial Intelligence, Computer Science), according to OpenAlex
Citations: cited by 1 paper (Europe PMC); 40 references in the paper

Abstract

Accurate disease prognosis is essential for patient care but is often hindered by the scarcity of longitudinal data. This study explores deep learning training strategies that utilize large, accessible diagnostic datasets to pretrain models aimed at predicting future disease progression in knee osteoarthritis (OA), Alzheimer’s disease (AD), and breast cancer (BC). While diagnostic pretraining improves prognostic task performance, naive fine-tuning for prognosis can cause ‘catastrophic forgetting,’ where the model’s original diagnostic accuracy degrades, a significant patient safety concern in real-world settings. To address this, we propose a sequential learning strategy with experience replay. We used cohorts with knee radiographs, brain MRIs, and digital mammograms to predict 4-year structural worsening in OA, 2-year cognitive decline in AD, and 5-year cancer diagnosis in BC. Our results showed that diagnostic pretraining on larger datasets improved prognosis model performance compared to standard baselines, boosting both the Area Under the Receiver Operating Characteristic curve (AUROC) (e.g., Knee OA external: 0.770 vs 0.747; Breast Cancer: 0.874 vs 0.848) and the Area Under the Precision-Recall Curve (AUPRC) (e.g., Alzheimer’s Disease: 0.752 vs 0.683). Additionally, a sequential learning approach with experience replay achieved prognostic performance comparable to dedicated single-task models (e.g., Breast Cancer AUROC 0.876 vs 0.874) while also preserving diagnostic ability. This method maintained high diagnostic accuracy (e.g., Breast Cancer Balanced Accuracy 50.4% vs 50.9% for a dedicated diagnostic model), unlike simpler multitask methods prone to catastrophic forgetting (e.g., 37.7%). Our findings show that leveraging large diagnostic datasets is a reliable and data-efficient way to enhance prognostic models while maintaining essential diagnostic skills.

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

denizlab/diag-to-prog-replay

License: CC-BY-NC-4.0
State: the link answers, verified on 28 September 2026
Evidence: files inventoried
Commit: a1b230ff75f5386020e6e1bce856dabd56d34b72, 7 January 2026
Languages: Python (51), Shell (6)
Size: 179 files, 57 scripts
Software Heritage: not archived
Found in: “Data Availability”
Holds: README, license file, environment (requirements.txt)
Not found: CITATION.cff, tests, continuous integration, documentation
Tools: PyTorch (44 files), NumPy (28 files), scikit-learn (11 files), pandas (9 files), Pillow (7 files), SciPy (6 files), h5py (5 files), FreeSurfer (4 files), Matplotlib (3 files), OpenCV (3 files), scikit-image (2 files), NiBabel (1 file)
Availability: 1 check, the latest on 28 September 2026: the link answers
  • 28 September 2026: the link answers
59 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;
  • 57 scripts, each with its path and the digest of its content;
  • 5 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

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

Data

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

Data Availability

- Osteoarthritis Initiative (OAI): OAI data are publicly available through the NIMH Data Archive (NDA). Access requires creating an NDA account, agreeing to the OAI data access terms, and requesting the desired collections through the NDA portal. Full instructions are provided at https://nda.nih.gov/oai. - Multicenter Osteoarthritis Study (MOST): MOST data are publicly available through the NIA Aging Research Biobank. Access requires creating a Biobank account and submitting a data request in accordance with the Biobank’s terms and conditions. More information is available at https://agingresearchbiobank.nia.nih.gov/. - Alzheimer’s Disease Neuroimaging Initiative (ADNI): ADNI data are publicly available from the LONI Image & Data Archive upon completion of a web application and acceptance of the ADNI Data Use Agreement. Data can be accessed at https://adni.loni.usc.edu/. - Breast Cancer Cohort (NYU Langone Health): This institutional dataset contains protected health information and cannot be made publicly available. De-identified data may be shared upon reasonable request, pending approval by NYU Langone’s data governance committee. Inquiries can be directed to Krzysztof Geras (). - Code and Model Availability: All code for preprocessing and model training sufficient to reproduce the analyses reported in this study are available at https://github.com/denizlab/diag-to-prog-replay.

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

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 28 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 8 authors, 11 MeSH terms, 1 funder, 35 references.

Cite

This paper

Rajamohan, H. R., Xu, Y., Zhu, W., Kijowski, R., Cho, K., Geras, K. J., Razavian, N., & Deniz, C. M. (2026). Robust disease prognosis via diagnostic knowledge preservation: A sequential learning approach. PloS one, 21(5), e0344600. https://doi.org/10.1371/journal.pone.0344600

BibTeX

@article{rajamohan2026robust,
author = {Rajamohan, Haresh Rengaraj and Xu, Yanqi and Zhu, Weicheng and Kijowski, Richard and Cho, Kyunghyun and Geras, Krzysztof J. and Razavian, Narges and Deniz, Cem M.},
title = {{Robust disease prognosis via diagnostic knowledge preservation: A sequential learning approach}},
journal = {PloS one},
year = {2026},
month = may,
volume = {21},
number = {5},
pages = {e0344600},
publisher = {PLOS},
issn = {1932-6203},
doi = {10.1371/journal.pone.0344600},
url = {https://doi.org/10.1371/journal.pone.0344600},
pmid = {42090385},
pmcid = {PMC13148697}
}

RIS

TY - JOUR
AU - Rajamohan, Haresh Rengaraj
AU - Xu, Yanqi
AU - Zhu, Weicheng
AU - Kijowski, Richard
AU - Cho, Kyunghyun
AU - Geras, Krzysztof J.
AU - Razavian, Narges
AU - Deniz, Cem M.
TI - Robust disease prognosis via diagnostic knowledge preservation: A sequential learning approach
T2 - PloS one
J2 - PLoS One
PY - 2026
DA - 2026/05/06
VL - 21
IS - 5
SP - e0344600
SN - 1932-6203
PB - PLOS
DO - 10.1371/journal.pone.0344600
UR - https://doi.org/10.1371/journal.pone.0344600
LA - en
ER -

CSL-JSON

{
"id": "10.1371/journal.pone.0344600",
"type": "article-journal",
"title": "Robust disease prognosis via diagnostic knowledge preservation: A sequential learning approach",
"container-title": "PloS one",
"author": [
{
"family": "Rajamohan",
"given": "Haresh Rengaraj"
},
{
"family": "Xu",
"given": "Yanqi"
},
{
"family": "Zhu",
"given": "Weicheng"
},
{
"family": "Kijowski",
"given": "Richard"
},
{
"family": "Cho",
"given": "Kyunghyun"
},
{
"family": "Geras",
"given": "Krzysztof J."
},
{
"family": "Razavian",
"given": "Narges"
},
{
"family": "Deniz",
"given": "Cem M."
}
],
"container-title-short": "PLoS One",
"volume": "21",
"issue": "5",
"page": "e0344600",
"DOI": "10.1371/journal.pone.0344600",
"PMID": "42090385",
"PMCID": "PMC13148697",
"ISSN": "1932-6203",
"publisher": "PLOS",
"URL": "https://doi.org/10.1371/journal.pone.0344600",
"language": "en",
"issued": {
"date-parts": [
[
2026,
5,
6
]
]
}
}

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.7554/elife.107933 [code]
Modality-agnostic decoding of vision and language from fMRI.
Journal: eLife
In common: FreeSurfer, OpenCV, scikit-image, 9 other tools, 1 reference
[2] doi:10.1016/j.patter.2026.101538 [code]
A multi-modal foundation model for brain disease diagnosis and medical imaging.
Journal: Patterns (New York, N.Y.)
In common: OpenCV, scikit-image, h5py, 8 other tools, clinical / translational, 1 reference
[3] doi:10.1038/s41597-026-07248-6 [code]
A large-scale fMRI dataset for vision-language semantic association.
Journal: Scientific data
In common: FreeSurfer, OpenCV, scikit-image, 9 other tools
[4] 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, OpenCV, scikit-image, 8 other tools, Alzheimer's / dementia, other condition
[5] doi:10.1038/s41598-026-57519-w [code]
Automated segmentation of neurons and spinal cord structures in immunofluorescence images using SpineDL.
Journal: Scientific reports
In common: OpenCV, scikit-image, h5py, 8 other tools, other condition
[6] doi:10.1038/s41467-026-76098-y [code]
A single computational objective can produce specialization of streams in visual cortex.
Journal: Nature communications
In common: FreeSurfer, scikit-image, h5py, 8 other tools
[7] doi:10.1371/journal.pcbi.1014263 [code]
MIRAGE: Robust multi-modal architectures translate fMRI-to-image models from vision to mental imagery.
Journal: PLoS computational biology
In common: OpenCV, scikit-image, h5py, 8 other tools
[8] 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: OpenCV, scikit-image, h5py, 8 other tools
[9] doi:10.1038/s41467-026-73373-w [code]
Mapping neuro-vascular unit communications reveals distinct angiogenic programs across developing mouse brain regions.
Journal: Nature communications
In common: OpenCV, scikit-image, h5py, 8 other tools
[10] doi:10.1038/s41467-026-71555-0 [code]
A deep representation learning model to predict response to vagus nerve stimulation.
Journal: Nature communications
In common: FreeSurfer, NiBabel, PyTorch, 5 other tools, clinical / translational, 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.