Robust disease prognosis via diagnostic knowledge preservation: A sequential learning approach.
The 5 matches
- [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. 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] § 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] § 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] § 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
- import os
- import wandb
- import torch
- import torch.nn as nn
- import torch.optim as optim
- import torch.nn.functional as F
- from torch.utils.data import DataLoader, Dataset
- from sklearn.metrics import roc_auc_score, average_precision_score
- import numpy as np
- from tqdm import tqdm
- import argparse
- import pickle
- from models_torch import BaselineBreastModel
- from torch_transforms import compose_transform
- import load_mammogram
- os.environ['WANDB_API_KEY'] = #
- class MammogramDataset(Dataset):
- """Dataset class for mammogram images with multi-task support"""
- def __init__(self, img_dir, datalist, transform, crop_size=(2944, 1920), tasks=None):
- self.img_dir = img_dir
- self.datalist = datalist
- self.transform = transform
- self.crop_size = crop_size
- self.tasks = tasks if tasks else ["prognosis"] # Default to single task if not specified
- # Filter out samples with missing labels for the required tasks
- self.valid_indices = []
- for idx, metadata in enumerate(self.datalist):
- is_valid = True
- # Check if BI-RADS label is valid if birads task is requested
- if "birads" in self.tasks:
- birads_value = metadata.get('birads', None)
- if birads_value is None or birads_value == 'NA' or birads_value == 'inconclusive':
- is_valid = False
- if is_valid:
- self.valid_indices.append(idx)
- print(f"Dataset initialized with {len(self.valid_indices)}/{len(datalist)} valid samples for tasks: {tasks}")
- def __len__(self):
- return len(self.valid_indices)
- def __getitem__(self, idx):
- original_idx = self.valid_indices[idx]
- metadata = self.datalist[original_idx]
- # Load mammogram exam
- images, scans_info = load_mammogram.load_mammogram_exam(
- img_dir=self.img_dir,
- metadata=metadata,
- transform=self.transform,
- select_index_logic=("random", 4),
- crop_method="best_center",
- crop_size=self.crop_size,
- skip_img=False
- )
- # Prepare image data
- x = {}
- for i, scan in enumerate(scans_info):
- side = scan['laterality']
- if side == "left":
- x[f"L-{scan['view']}"] = images[i,0:1,:,:]
- else:
- x[f"R-{scan['view']}"] = images[i,0:1,:,:]
- # Prepare labels for all requested tasks
- labels = {}
- if "prognosis" in self.tasks:
- labels["prognosis"] = metadata['5yr_cancer_label']['malignant']
- if "birads" in self.tasks and 'birads' in metadata and metadata['birads'] not in [None, 3, 4, 5]:
- labels["birads"] = metadata['birads']
- else:
- labels["birads"] = -1
- # Return different outputs based on the number of tasks
- if len(self.tasks) == 1:
- # For single task, return a single label for backward compatibility
- return x, labels[self.tasks[0]]
- else:
- # For multi-task, return a dictionary of labels
- return x, labels
- def collate_mammogram_batch(batch):
- """
- Custom collate function for mammogram batches with single or multiple labels per sample.
- Args:
- batch: List of tuples (x, label) where:
- x: Dictionary with keys 'L-CC', 'L-MLO', 'R-CC', 'R-MLO' mapping to image tensors
- label: Either a single value or a dictionary of task labels
- Returns:
- Tuple of:
- batch_x: Dictionary mapping view names to batched image tensors
- batch_labels: Dictionary of label tensors or single tensor depending on input format
- """
- # Initialize dictionary for images
- batch_x = {
- 'L-CC': [], 'L-MLO': [],
- 'R-CC': [], 'R-MLO': []
- }
- # Check if using multi-task labels (dictionary) or single task label
- first_label = batch[0][1]
- is_multi_task = isinstance(first_label, dict)
- if is_multi_task:
- # Initialize dictionary for each task
- tasks = first_label.keys()
- batch_labels = {task: [] for task in tasks}
- # Collect images and task labels
- for x, labels in batch:
- for view in batch_x.keys():
- if view in x:
- batch_x[view].append(x[view])
- else:
- # Handle missing views with a tensor of zeros
- # Assuming all tensors in a batch have the same shape
- shape = next(iter(x.values())).shape
- batch_x[view].append(torch.zeros(shape))
- for task, label in labels.items():
- batch_labels[task].append(label)
- # Stack tensors for each view and create label tensors
- batch_x = {view: torch.stack(tensors) for view, tensors in batch_x.items()}
- batch_labels = {task: torch.tensor(labels) for task, labels in batch_labels.items()}
- else:
- # Single task mode - for backward compatibility
- batch_labels = []
- # Collect images and labels
- for x, label in batch:
- for view in batch_x.keys():
- if view in x:
- batch_x[view].append(x[view])
- else:
- # Handle missing views with a tensor of zeros
- shape = next(iter(x.values())).shape
- batch_x[view].append(torch.zeros(shape))
- batch_labels.append(label)
- # Stack tensors for each view and create label tensor
- batch_x = {view: torch.stack(tensors) for view, tensors in batch_x.items()}
- batch_labels = torch.tensor(batch_labels)
- return batch_x, batch_labels
- def unpickle_from_file(file_name):
- with open(file_name, 'rb') as handle:
- try:
- return pickle.load(handle)
- except ImportError:
- return pd.read_pickle(file_name)
- class MultiTaskBreastModel(nn.Module):
- """Wrapper for baseline breast model to support multiple output heads"""
- def __init__(self, base_model, device, tasks=None):
- super(MultiTaskBreastModel, self).__init__()
- self.tasks = tasks if tasks else ["prognosis"]
- self.base_model = base_model
- original_fc2 = self.base_model.fc2
- # Base model
- # Get feature dimension from base model
- self.feature_dim = self.base_model.fc2.in_features
- # Create task-specific output heads
- self.heads = nn.ModuleDict()
- if "prognosis" in self.tasks:
- self.heads["prognosis"] = nn.Linear(self.feature_dim, 1)
- if "birads" in self.tasks:
- self.heads["birads"] = nn.Linear(original_fc2.in_features, original_fc2.out_features)
- # Optionally copy weights and biases
- self.heads["birads"].weight.data = original_fc2.weight.data.clone()
- self.heads["birads"].bias.data = original_fc2.bias.data.clone()
- self.base_model.fc2 = nn.Identity()
- def forward(self, x):
- # Extract features from base model
- features = self.base_model(x)
- # Apply task-specific heads
- outputs = {}
- for task in self.tasks:
- outputs[task] = self.heads[task](features)
- # Reshape binary outputs to match expected format
- if task == "prognosis":
- outputs[task] = outputs[task].squeeze(-1)
- return outputs
- def get_non_requires_grad_params(model):
- """
- Find all parameters in a model that do NOT require gradients.
- Args:
- model: PyTorch model
- Returns:
- List of parameter names that don't require gradients and their shapes
- """
- non_requires_grad_params = []
- for name, param in model.named_parameters():
- if not param.requires_grad:
- non_requires_grad_params.append((name, param.shape))
- return non_requires_grad_params
- def load_pretrained_weights(model, pretrained_path):
- """Load pretrained BI-RADS weights and modify for binary classification"""
- if pretrained_path:
- pretrained_dict = torch.load(pretrained_path, map_location='cpu')
- model_dict = model.state_dict()
- # Filter out fc2 layer from pretrained weights
- pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict and 'fc2' not in k}
- # Update model weights
- model_dict.update(pretrained_dict)
- model.load_state_dict(model_dict)
- return model
- def train_epoch(model, train_loader, criterions, optimizer, device, tasks, diagnosis_loader=None, replay_prob=0.5):
- print(tasks)
- model.train()
- running_loss = 0.0
- task_running_losses = {task: 0.0 for task in tasks}
- # For metrics tracking - only track prognosis for simplicity
- task_metrics = {}
- if 'prognosis' in tasks:
- task_metrics['prognosis'] = {
- 'predictions': [],
- 'targets': []
- }
- if diagnosis_loader is not None:
- task_running_losses['birads'] = 0.0
- # Set up diagnosis iterator if provided
- diagnosis_iterator = iter(diagnosis_loader) if diagnosis_loader else None
- pbar = tqdm(train_loader, desc='Training')
- for batch_idx, (data, targets) in enumerate(pbar):
- # Determine whether to do replay for this batch
- do_replay = diagnosis_loader is not None and np.random.uniform() <= replay_prob
- if do_replay:
- try:
- # Get a batch from diagnosis dataset
- print("its here")
- replay_data, replay_targets = next(diagnosis_iterator)
- # Replace current batch with diagnosis data
- data = replay_data
- # Handle targets format
- if isinstance(replay_targets, dict):
- targets = replay_targets
- else:
- # Convert to dict format if needed
- targets = {"birads": replay_targets}
- # Set current tasks for this batch
- current_tasks = ["birads"]
- except StopIteration:
- # Reset iterator if we've gone through all diagnosis data
- diagnosis_iterator = iter(diagnosis_loader)
- replay_data, replay_targets = next(diagnosis_iterator)
- # Replace current batch with diagnosis data
- data = replay_data
- # Handle targets format
- if isinstance(replay_targets, dict):
- targets = replay_targets
- else:
- # Convert to dict format if needed
- targets = {"birads": replay_targets}
- # Set current tasks for this batch
- current_tasks = ["birads"]
- else:
- # Regular prognosis batch
- current_tasks = tasks
- # Move data to device
- for view in data:
- data[view] = data[view].to(device)
- # Handle targets based on format
- if not isinstance(targets, dict):
- # Convert single task target to dict format for consistent handling
- # Assuming single task is always the first task in the list
- targets = {current_tasks[0]: targets.to(device)}
- else:
- # Move all targets to device
- for task in targets:
- if task in targets: # Only process tasks that have labels
- targets[task] = targets[task].to(device)
- # Convert birads to long for CrossEntropyLoss
- if task == "birads":
- targets[task] = targets[task].long()
- optimizer.zero_grad()
- outputs = model(data)
- # Calculate task-specific losses and total loss
- total_loss = 0.0
- for task in current_tasks:
- if task in targets and task in outputs:
- mask = (targets[task] != -1)
- if isinstance(mask, bool):
- if not mask:
- continue
- mask = torch.tensor([True]) # Create a single-element boolean tensor
- elif mask.sum() == 0:
- continue
- if isinstance(criterions[task], nn.BCEWithLogitsLoss):
- # For binary cross entropy, both outputs and targets should be float
- task_loss = criterions[task](
- outputs[task][mask].float(),
- targets[task][mask].float()
- )
- elif isinstance(criterions[task], nn.CrossEntropyLoss):
- # For cross entropy, outputs should be float but targets should be long
- task_loss = criterions[task](
- outputs[task][mask].float(),
- targets[task][mask].long()
- )
- # print("birads targets ", targets[task][mask].long())
- else:
- # For other loss functions, maintain original types
- task_loss = criterions[task](outputs[task][mask], targets[task][mask])
- task_running_losses[task] += task_loss.item()
- total_loss += task_loss
- # Store predictions for prognosis metrics only (if it's a prognosis batch)
- if task == 'prognosis' and task in task_metrics and task in current_tasks:
- task_metrics[task]['predictions'].extend(torch.sigmoid(outputs[task]).detach().cpu().numpy())
- task_metrics[task]['targets'].extend(targets[task].cpu().numpy())
- total_loss.backward()
- optimizer.step()
- running_loss += total_loss.item()
- # Update progress bar
- pbar.set_postfix({
- 'loss': total_loss.item(),
- **{f"{task}_loss": task_running_losses[task]/(batch_idx+1) for task in tasks if task in task_running_losses}
- })
- # Calculate average losses
- epoch_loss = running_loss / len(train_loader)
- task_losses = {task: task_running_losses[task] / len(train_loader) for task in tasks if task in task_running_losses}
- # Calculate metrics for prognosis task only
- metrics = {}
- if 'prognosis' in task_metrics and len(task_metrics['prognosis']['predictions']) > 0:
- preds = np.array(task_metrics['prognosis']['predictions'])
- targets = np.array(task_metrics['prognosis']['targets'])
- metrics["prognosis_auroc"] = roc_auc_score(targets, preds)
- metrics["prognosis_auprc"] = average_precision_score(targets, preds)
- return epoch_loss, task_losses, metrics
- def validate(model, val_loader, criterions, device, tasks):
- model.eval()
- running_loss = 0.0
- task_running_losses = {task: 0.0 for task in tasks}
- # For metrics tracking - only track prognosis
- task_metrics = {}
- if 'prognosis' in tasks:
- task_metrics['prognosis'] = {
- 'predictions': [],
- 'targets': []
- }
- with torch.no_grad():
- for data, targets in tqdm(val_loader, desc="Validation"):
- # Move data to device
- for view in data:
- data[view] = data[view].to(device)
- # Handle both single-task and multi-task targets
- if not isinstance(targets, dict):
- # Convert single task target to dict format for consistent handling
- targets = {"prognosis": targets.to(device)}
- else:
- # Move all targets to device
- for task in targets:
- if task in targets: # Only process tasks that have labels
- targets[task] = targets[task].to(device)
- # Convert birads to long for CrossEntropyLoss
- if task == "birads":
- targets[task] = targets[task].long()
- outputs = model(data)
- # Calculate task-specific losses and total loss
- total_loss = 0.0
- for task in tasks:
- if task in targets and task in outputs:
- mask = (targets[task] != -1)
- if isinstance(mask, bool):
- if not mask:
- continue
- mask = torch.tensor([True]) # Create a single-element boolean tensor
- elif mask.sum() == 0:
- continue
- if isinstance(criterions[task], nn.BCEWithLogitsLoss):
- # For binary cross entropy, both outputs and targets should be float
- task_loss = criterions[task](
- outputs[task][mask].float(),
- targets[task][mask].float()
- )
- elif isinstance(criterions[task], nn.CrossEntropyLoss):
- # For cross entropy, outputs should be float but targets should be long
- task_loss = criterions[task](
- outputs[task][mask].float(),
- targets[task][mask].long()
- )
- else:
- # For other loss functions, maintain original types
- task_loss = criterions[task](outputs[task][mask], targets[task][mask])
- task_running_losses[task] += task_loss.item()
- total_loss += task_loss
- # Store predictions for prognosis metrics only
- if task == 'prognosis' and task in task_metrics:
- task_metrics[task]['predictions'].extend(torch.sigmoid(outputs[task]).detach().cpu().numpy())
- task_metrics[task]['targets'].extend(targets[task].cpu().numpy())
- running_loss += total_loss.item()
- # Calculate average losses
- epoch_loss = running_loss / len(val_loader)
- task_losses = {task: task_running_losses[task] / len(val_loader) for task in tasks if task in task_running_losses}
- # Calculate metrics for prognosis task only
- metrics = {}
- if 'prognosis' in task_metrics and len(task_metrics['prognosis']['predictions']) > 0:
- preds = np.array(task_metrics['prognosis']['predictions'])
- targets = np.array(task_metrics['prognosis']['targets'])
- metrics["prognosis_auroc"] = roc_auc_score(targets, preds)
- metrics["prognosis_auprc"] = average_precision_score(targets, preds)
- return epoch_loss, task_losses, metrics
- def train(args):
- wandb.init(project=args.wandb_project, name=args.wandb_run_name)
- wandb.config.update(args)
- device = torch.device("cuda" if args.device_type == "gpu" else "cpu")
- # Define tasks based on arguments
- tasks = ["prognosis"]
- if args.multitask:
- tasks.append("birads")
- print(f"Training with tasks: {tasks}")
- # Load prognosis data
- prognosis_train_data = unpickle_from_file(f"{args.prognosis_datalist_path}/Fold_{args.fold}_train.pkl")
- val_data = unpickle_from_file(f"{args.prognosis_datalist_path}/Fold_{args.fold}_val.pkl")
- # np.random.shuffle(prognosis_train_data)
- # prognosis_train_data = prognosis_train_data[:256]
- # np.random.shuffle(val_data)
- # val_data = val_data[:256]
- transform = compose_transform(augmentation="standard", resize=None, image_format="greyscale")
- # Create primary dataset for prognosis
- train_dataset = MammogramDataset(
- img_dir=args.img_dir,
- datalist=prognosis_train_data,
- transform=transform,
- crop_size=args.crop_size,
- tasks=tasks # Primary dataset includes all primary tasks
- )
- val_dataset = MammogramDataset(
- img_dir=args.img_dir,
- datalist=val_data,
- transform=compose_transform(image_format="greyscale"),
- crop_size=args.crop_size,
- tasks=["prognosis", "birads"] # Validation dataset includes all tasks
- )
- # Create data loaders for prognosis
- train_loader = DataLoader(
- train_dataset,
- batch_size=args.batch_size,
- shuffle=True,
- num_workers=args.num_workers,
- collate_fn=collate_mammogram_batch,
- pin_memory=True
- )
- val_loader = DataLoader(
- val_dataset,
- batch_size=args.batch_size,
- shuffle=False,
- num_workers=args.num_workers,
- collate_fn=collate_mammogram_batch,
- pin_memory=True
- )
- # Create diagnosis dataset (separate cohort)
- diagnosis_loader = None
- if args.replay:
- print("Setting up diagnosis loader for batch-level replay")
- # Load diagnosis data from a separate path
- diagnosis_train_data = unpickle_from_file(f"{args.diagnosis_datalist_path}/train.pkl")
- diagnosis_dataset = MammogramDataset(
- img_dir=args.img_dir,
- datalist=diagnosis_train_data,
- transform=transform,
- crop_size=args.crop_size,
- tasks=["birads"] # This dataset is only for diagnosis
- )
- diagnosis_loader = DataLoader(
- diagnosis_dataset,
- batch_size=args.batch_size,
- shuffle=True,
- num_workers=args.num_workers,
- collate_fn=collate_mammogram_batch,
- pin_memory=True
- )
- nodropout_probability = 1 - args.dropout_p
- base_model = BaselineBreastModel(
- device=device,
- nodropout_probability=nodropout_probability,
- gaussian_noise_std=args.gaussian_noise_std,
- )
- if args.pretrained_path:
- base_model = load_pretrained_weights(base_model, args.pretrained_path)
- model = MultiTaskBreastModel(
- base_model=base_model,
- device=device,
- tasks=["prognosis", "birads"] # Always include both tasks in model
- ).to(device)
- frozen_params = get_non_requires_grad_params(model)
- print(f"Total parameters NOT requiring gradients: {len(frozen_params)}")
- for name, shape in frozen_params:
- print(f"{name}: {shape}")
- # Set up class weights for BI-RADS using the provided counts
- class_counts = torch.tensor([6634, 17388,15939], dtype=torch.float).to(device)
- class_weights = 1.0 / class_counts
- class_weights = class_weights / class_weights.sum() * len(class_counts)
- # Setup loss functions for each task
- criterions = {
- "prognosis": nn.BCEWithLogitsLoss(),
- "birads": nn.CrossEntropyLoss(weight=class_weights)
- }
- optimizer = optim.Adam(model.parameters(), lr=args.learning_rate)
- scheduler = optim.lr_scheduler.ReduceLROnPlateau(
- optimizer, mode='max', factor=0.75, patience=10, verbose=True
- )
- # Create save directory
- save_path = os.path.join(args.save_path, f"Fold_{args.fold}", args.wandb_run_name)
- os.makedirs(save_path, exist_ok=True)
- # Track best metrics for each task
- best_metrics = {
- "total_loss": float('inf'),
- "prognosis_auroc": 0,
- "prognosis_auprc": 0,
- }
- patience_counter = 0
- # Training loop
- for epoch in range(args.epochs):
- print(f'\nEpoch {epoch+1}/{args.epochs}')
- # Training phase - with batch-level replay
- if args.replay and diagnosis_loader is not None:
- # Use new train_epoch function with batch-level replay
- total_train_loss, all_train_task_losses, all_train_metrics = train_epoch(
- model=model,
- train_loader=train_loader,
- criterions=criterions,
- optimizer=optimizer,
- device=device,
- tasks=["prognosis"],
- diagnosis_loader=diagnosis_loader,
- replay_prob=args.replay_prob
- )
- # Joint multitask learning
- elif args.multitask:
- # Train on both prognosis and birads tasks in the same dataset
- total_train_loss, all_train_task_losses, all_train_metrics = train_epoch(
- model=model,
- train_loader=train_loader,
- criterions=criterions,
- optimizer=optimizer,
- device=device,
- tasks=["prognosis", "birads"]
- )
- # Single task mode (prognosis only)
- else:
- # Standard single-task training
- total_train_loss, all_train_task_losses, all_train_metrics = train_epoch(
- model=model,
- train_loader=train_loader,
- criterions=criterions,
- optimizer=optimizer,
- device=device,
- tasks=["prognosis"]
- )
- # Validation phase - always validate on all tasks
- val_loss, val_task_losses, val_metrics = validate(
- model=model,
- val_loader=val_loader,
- criterions=criterions,
- device=device,
- tasks=["prognosis", "birads"]
- )
- # Log metrics
- log_dict = {
- 'epoch': epoch,
- 'train_loss': total_train_loss,
- 'val_loss': val_loss,
- **{f"train_{task}_loss": loss for task, loss in all_train_task_losses.items()},
- **{f"val_{task}_loss": loss for task, loss in val_task_losses.items()},
- **{f"train_{metric}": value for metric, value in all_train_metrics.items()},
- **{f"val_{metric}": value for metric, value in val_metrics.items()},
- 'learning_rate': optimizer.param_groups[0]['lr']
- }
- wandb.log(log_dict)
- # Print metrics
- print(f"Epoch {epoch+1} Results:")
- for metric, value in val_metrics.items():
- print(f" Val {metric}: {value:.4f}")
- # Update learning rate based on primary task metric (prognosis AUROC)
- if 'prognosis_auroc' in val_metrics:
- scheduler.step(val_metrics['prognosis_auroc'])
- # Check if this is the best model so far
- improved = False
- # Check each metric
- for metric in best_metrics:
- if metric == 'total_loss':
- if val_loss < best_metrics[metric]:
- best_metrics[metric] = val_loss
- improved = True
- torch.save(model.state_dict(), os.path.join(save_path, f'best_loss_model.pth'))
- elif metric in val_metrics and best_metrics[metric] is not None:
- if val_metrics[metric] > best_metrics[metric]:
- best_metrics[metric] = val_metrics[metric]
- improved = True
- torch.save(model.state_dict(), os.path.join(save_path, f'best_{metric}_model.pth'))
- # Also save a checkpoint for this epoch
- torch.save({
- 'epoch': epoch,
- 'model_state_dict': model.state_dict(),
- 'optimizer_state_dict': optimizer.state_dict(),
- 'scheduler_state_dict': scheduler.state_dict(),
- 'best_metrics': best_metrics,
- **val_metrics
- }, os.path.join(save_path, f'checkpoint_epoch_{epoch}.pth'))
- # Check for early stopping
- if improved:
- patience_counter = 0
- else:
- patience_counter += 1
- if patience_counter >= args.patience:
- print(f'Early stopping triggered after epoch {epoch+1}')
- break
- wandb.finish()
- # Save final model
- torch.save(model.state_dict(), os.path.join(save_path, 'final_model.pth'))
- # Print best results
- print("\nTraining completed. Best results:")
- for metric, value in best_metrics.items():
- if value is not None:
- print(f" Best {metric}: {value:.4f}")
- if __name__ == "__main__":
- parser = argparse.ArgumentParser(description='Train multi-task mammogram classification model')
- # Model parameters
- parser.add_argument('--pretrained-path', type=str, default=None,
- help='Path to pretrained model weights')
- parser.add_argument('--save-path', type=str, default=None,
- help='Path to save models')
- parser.add_argument('--gaussian-noise-std', type=float, default=0.01,
- help='Standard deviation of Gaussian noise')
- parser.add_argument('--dropout-p', type=float, default=0.1,
- help='Dropout probability')
- # Multi-task parameters
- parser.add_argument('--multitask', action='store_true',
- help='Enable multi-task learning with prognosis and diagnosis')
- parser.add_argument('--replay', action='store_true',
- help='Use experience replay for multi-task learning')
- parser.add_argument('--replay-prob', type=float, default=0.5,
- help='Probability of doing experience replay in each epoch')
- # Training parameters
- parser.add_argument('--batch-size', type=int, default=8)
- parser.add_argument('--learning-rate', type=float, default=1e-5)
- parser.add_argument('--epochs', type=int, default=100)
- parser.add_argument('--patience', type=int, default=10,
- help='Early stopping patience')
- parser.add_argument('--num-workers', type=int, default=8)
- # Device parameters
- parser.add_argument('--device-type', type=str, default='gpu')
- parser.add_argument('--gpu-number', type=int, default=0)
- # Prognosis data parameters
- parser.add_argument('--img-dir', type=str, required=True,
- help='Root directory containing the prognosis image data')
- parser.add_argument('--prognosis-datalist-path', type=str, required=True,
- help='Path to the pickle file containing prognosis train/val/test splits')
- # Diagnosis data parameters (for multitask or replay)
- parser.add_argument('--diagnosis-datalist-path', type=str,
- help='Path to the pickle file containing diagnosis train/val/test splits')
- # Common data parameters
- parser.add_argument('--crop-size', type=int, nargs=2, default=[2944, 1920],
- help='Crop size for images (height width)')
- parser.add_argument('--fold', type=int, default=1,
- help='which Fold to train on')
- # Wandb parameters
- parser.add_argument('--wandb-project', type=str, required=True,
- help='WandB project name')
- parser.add_argument('--wandb-run-name', type=str, required=True,
- help='WandB run name')
- args = parser.parse_args()
- train(args)
train.py at commit a1b230f, under CC-BY-NC-4.0 · at the source
Overview
- Center for Data Science, New York University, New York, New York, United States of America
- Department of Radiology, Hospital for Special Surgery, New York, New York, United States of America
- Department of Radiology, New York University Langone Health, New York, New York, United States of America
- Department of Population Health, New York University Langone Health, New York, New York, United States of America
- Bernard and Irene Schwartz Center for Biomedical Imaging, New York University Langone Health, New York, New York, United States of America
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
a1b230ff75f5386020e6e1bce856dabd56d34b72, 7 January 2026Availability: 1 check, the latest on 28 September 2026: the link answers
- 28 September 2026: the link answers
59 files
- AD/
datasets/ , Python, 1 line__init__.py - AD/
datasets/ , Python, 145 linesadni_3d.py - AD/
datasets/ , Python, 111 linesaugmentations.py - AD/
datasets/ , Shell, 20 linesfiles/ run_adni_preprocess.sh - AD/
datasets/ , Shell, 20 linesfiles/ run_adni_preprocess_test .sh - AD/
datasets/ , Shell, 20 linesfiles/ run_adni_preprocess_val. sh - AD/
datasets/ , Shell, 14 linesfiles/ run_convert.sh - AD/
evaluate.py , Python, 212 lines - AD/
lib/ , Python, 106 linesLoss.py - AD/
lib/ , Python, 1 line__init__.py - AD/
lib/ , Python, 63 linesalias_multinomial.py - AD/
lib/ , Python, 436 linescustom_transforms.py - AD/
lib/ , Python, 14 linesnormalize.py - AD/
lib/ , Python, 284 linesutils.py - AD/
main.py , Python, 418 lines - AD/
models/ , Python, 79 linesLinearModel.py - AD/
models/ , Python, 68 linesbuild_model.py - AD/
models/ , Python, 154 linesclassifier.py - AD/
models/ , Python, 142 linesmodels.py - Breast_Cancer/
Multitask/ , Python, 314 lines, 1 matchevaluate.py - Breast_Cancer/
Multitask/ , Python, 110 lineslayers_torch.py - Breast_Cancer/
Multitask/ , Python, 144 linesload_mammogram.py - Breast_Cancer/
Multitask/ , Python, 487 linesloading_mammogram_utils. py - Breast_Cancer/
Multitask/ , Python, 94 linesmodels_torch.py - Breast_Cancer/
Multitask/ , Python, 155 linestorch_transforms.py - Breast_Cancer/
Multitask/ , Python, 796 lines, 2 matchestrain.py - Breast_Cancer/
Single_Task/ , Python, 268 lines, 1 matcheval_birads.py - Breast_Cancer/
Single_Task/ , Python, 246 linesevaluate.py - Breast_Cancer/
Single_Task/ , Python, 110 lineslayers_torch.py - Breast_Cancer/
Single_Task/ , Python, 144 linesload_mammogram.py - Breast_Cancer/
Single_Task/ , Python, 487 linesloading_mammogram_utils. py - Breast_Cancer/
Single_Task/ , Python, 35 linesloss.py - Breast_Cancer/
Single_Task/ , Python, 94 linesmodels_torch.py - Breast_Cancer/
Single_Task/ , Python, 95 linestest_inference.py - Breast_Cancer/
Single_Task/ , Python, 155 linestorch_transforms.py - Breast_Cancer/
Single_Task/ , Python, 483 linestrain.py - Breast_Cancer/
Single_Task/ , Python, 424 linestrain_birads.py - Breast_Cancer/
Single_Task/ , Python, 27 linesutils.py - Knee_OA/
XrayDataLoader.py , Python, 188 lines - Knee_OA/
attn_resnet/ , Python, 46 linesfocal.py - Knee_OA/
attn_resnet/ , Python, 173 linesloss_functions.py - Knee_OA/
attn_resnet/ , Python, 49 linesmodels/ bam.py - Knee_OA/
attn_resnet/ , Python, 99 linesmodels/ cbam.py - Knee_OA/
attn_resnet/ , Python, 223 linesmodels/ model_resnet.py - Knee_OA/
attn_resnet/ , Shell, 9 linesscripts/ train_imagenet_resnet50_ bam.sh - Knee_OA/
attn_resnet/ , Shell, 9 linesscripts/ train_imagenet_resnet50_ cbam.sh - Knee_OA/
attn_resnet/ , Python, 307 linestrain_imagenet.py - Knee_OA/
attn_resnet/ , Python, 354 lineswide_resnet.py - Knee_OA/
attn_resnet/ , Python, 86 lineswide_resnet/ wide_resnet.py - Knee_OA/
augmentation.py , Python, 238 lines - Knee_OA/
data.py , Python, 183 lines - Knee_OA/
evaluate.py , Python, 93 lines - Knee_OA/
helper.py , Python, 258 lines - Knee_OA/
loss.py , Python, 84 lines - Knee_OA/
model.py , Python, 103 lines - Knee_OA/
train.py , Python, 180 lines, 1 match - Knee_OA/
utils.py , Python, 233 lines - LICENSE, License, 330 lines
- README.md, Text, 73 lines
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://
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://
BibTeX
@article{rajamohan2026ro
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/
url = {https://
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/
VL - 21
IS - 5
SP - e0344600
SN - 1932-6203
PB - PLOS
DO - 10.1371/
UR - https://
LA - en
ER -
CSL-JSON
{
"id": "10.1371/
"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":
"volume": "21",
"issue": "5",
"page": "e0344600",
"DOI": "10.1371/
"PMID": "42090385",
"PMCID": "PMC13148697",
"ISSN": "1932-6203",
"publisher": "PLOS",
"URL": "https://
"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: eLifeIn 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 dataIn 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 AssociationIn 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 reportsIn 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 communicationsIn 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 biologyIn 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 intelligenceIn 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 communicationsIn 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 communicationsIn 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.
Claim this paper
Correct its record
Say what each link of this record is, remove the ones that are not the paper's, add the ones that are missing. The correction becomes a new version of the record, in its Versions section.
Validate its tracing map
You validate the map as this page shows it: 1 repository of the authors' code, each at its verified commit and with its license, 57 scripts, and 5 matches between paragraphs and code (see the Code and Map sections). It then receives a DOI on Zenodo, with you (your ORCID iD) and OSCR as its creators; the code itself is not deposited.
The map's fingerprint: sha256:83d2c5bd91bc159c…
Add the badge to its README
The badge links the code to this page. Copy one of these into the README of the paper's code: only you decide where it goes, and nothing is changed for you.
Markdown
[, paste the snippet at the top, then “Commit changes…” and, to review it first, “Create a new branch and start a pull request”. You open the pull request; OSCR asks for no permission.
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.
