#!/usr/bin/env python3 """ Train one ResNet on one cross validation fold. Paths may be given as arguments or come from the site config (see config.example.yaml): python scripts/train/train.py resnet18_abrv -b 32 -n 60 -c 0 -d 0 -t -e """ import math import os import pickle import tracemalloc import torch.utils.data import wandb from torch import nn, optim from torch.amp import GradScaler import auto_detect_breast_mri.evaluation.metrics as eva from auto_detect_breast_mri.config import get_config, resolve_path from auto_detect_breast_mri.data import loaders from auto_detect_breast_mri.data.transforms import (intensityAugmentation, minimalAugmentation, pre_image_shape) from auto_detect_breast_mri.models.checkpoints import load_pretrained_model, save_model from auto_detect_breast_mri.models.resnets import model_names from auto_detect_breast_mri.training.cli import generate_parser from auto_detect_breast_mri.training.lr_scheduler import LRScheduler_Manager from auto_detect_breast_mri.training.loops import (CPU, GPU, display_top, eval_epoch_subjects, train_epoch_subjects) from auto_detect_breast_mri.data.metadata import get_uka_metatensor if __name__ == '__main__': torch.manual_seed(31) tracemalloc.start() DEVICE = GPU if not torch.cuda.is_available(): DEVICE = CPU values = None # get_evaluation_metrics() parser = generate_parser() args = parser.parse_args() print(args) model_name = args.model_name # Paths come from the command line; anything omitted is taken from the site config # (see config.example.yaml and src/config.py). path_base = resolve_path(args.data_path, "data_root", "root folder of the NIfTI data") feature_path = resolve_path(args.feature_path, "metadata_file", "metadata export") split_files_folder = resolve_path(args.split_files_folder, "split_root", "folder holding the split files") resulting_folder = resolve_path(args.output_path, "output_root", "output folder") os.makedirs(resulting_folder, exist_ok=True) # Read optional arguments if given ### a - augmentation if args.augmentation and args.augmentation == 'intensity': transform = intensityAugmentation t = args.augmentation else: transform = minimalAugmentation t = 'basic' ### b - batch_size batch_size = 16 if args.batch_size and args.batch_size > 0: batch_size = args.batch_size ### c, d - outer, inner cross validation fold fold = None subfold = None if args.test_fold is not None: fold = args.test_fold if args.validation_fold is not None: subfold = args.validation_fold if args.dataloader is not None and args.dataloader != "subjects": parser.error(f"Dataloadertype {args.dataloader} is not supported. Only 'subjects' is available.") ### e - eval do_eval = args.evaluation ### f - fraction if args.fraction_of_data and 0 < args.fraction_of_data < 1: fraction = args.fraction_of_data else: fraction = -1.0 ### l - learning rate learning_rate = 0.0001 if args.learning_rate and 0 < args.learning_rate: learning_rate = args.learning_rate ### m - mixed precision use_mixed_precision = args.mixed_precision ### n - number of epochs num_epochs = 60 if args.number_of_epochs: num_epochs = args.number_of_epochs ### o - output path model_path = None if args.model_path: if len(args.model_path) > 4 and args.model_path.endswith('.pth'): model_path = args.model_path print("use model: " + model_path) else: raise ValueError("Invalid model path: " + args.model_path) else: model_path = None ### p - patience patience = None # number of epochs to wait without improvement before stopping if args.patience and 0 < args.patience <= num_epochs: patience = args.patience ### s - scheduler for learning rate scheduler = None scheduler_name = "No scheduler" use_scheduler = False if args.scheduler: print("Scheduler parameter is activated.") use_scheduler = True ### t - do training train = args.train ## TRAIN each model model = model_names.get(model_name) print('Prepare model {} for training file under {}'.format(model_name, split_files_folder)) # possible options: model_abrv, model_full if "abrv" in model_name: protocol = 'abbreviated' elif "sub" in model_name: protocol = ['Sub_1'] elif "full" in model_name: protocol = 'full' else: protocol = 'abbreviated' print("Protocol not supported: " + model_name) # Dataset and DataLoader train_loader = [] eval_loader = [] feature_dataframe = get_uka_metatensor(0, feature_path) if fold is not None and subfold is not None and train and do_eval: print("using subjects loader") train_loader, eval_loader, test_loader = loaders.get_multiple_subjects_dataloader(path_base, feature_dataframe, pre_image_shape, protocol=protocol, set_folder=split_files_folder, transform=transform, batch_size=batch_size, stratified=True, fold=fold, subfold=subfold, fraction=fraction) else: print("get individual dataloader") split_file_path_sceleton = split_files_folder + "fold{}/stratified_{}" if train: subfold = 0 if subfold is None else subfold split_file_path = split_file_path_sceleton.format(fold, f"training_set") train_loader = loaders.get_subjects_dataloader(path_base, feature_dataframe, pre_image_shape, transform, protocol, split_file_path, batch_size, fraction, fold, subfold) if do_eval: split_file_path = split_file_path_sceleton.format(fold, f"evaluation_set") # augmentation belongs to the training loader only; None falls back to the # deterministic default transform, so validation metrics are not computed on # randomly flipped and noised volumes eval_loader = loaders.get_subjects_dataloader(path_base, feature_dataframe, pre_image_shape, None, protocol, split_file_path, batch_size, fraction, fold, subfold) split_file_path = split_file_path_sceleton.format(fold, f"test_set") test_loader = loaders.get_subjects_dataloader(path_base, feature_dataframe, pre_image_shape, None, # no augmentation at test time protocol, split_file_path, batch_size, 1, fold) n_train_sample = len(train_loader.dataset) if train else 0 n_eval_sample = len(eval_loader.dataset) if do_eval else 0 snapshot = tracemalloc.take_snapshot() display_top(snapshot=snapshot, title="after dataset creation") weight_decay = 1e-2 # ------------ Initialize Model, Loss Function, Optimizer ------------ print("Use device: " + DEVICE) model.to(DEVICE) criterion = nn.BCEWithLogitsLoss() optimizer = optim.AdamW(model.parameters(), lr=learning_rate, weight_decay=weight_decay) print("use scheduler? " + str(use_scheduler)) if use_scheduler: schedule_manager = LRScheduler_Manager(optimizer, num_epochs) scheduler, scheduler_name = schedule_manager.get_scheduler(args.scheduler) print(f"scheduler: {scheduler_name}") if model_path is not None: # load pretrained model: checkpoint = load_pretrained_model(model, model_path, DEVICE) model.load_state_dict(checkpoint) lr = round(learning_rate, 7) # Project/entity/mode come from the site config; no local path is logged, since run configs # are visible to everybody with access to the wandb project. wandb.init(**get_config().wandb_init_kwargs(), config={ "architecture": model_name, "number of training samples": n_train_sample, "number of evaluation samples": n_eval_sample, "number of test samples": len(test_loader.dataset), "protocol": protocol, "optimizer": str(optimizer), "weight decay": weight_decay, "Augmentation": t + str(transform), "batch size": batch_size, "Fold": fold, "train fraction": fraction, "learning-rate": lr, "Mixed Precision": use_mixed_precision, "epochs": num_epochs, "Pretrained": (model_path is not None), "patience": patience, "scheduler": scheduler_name, "Machine": "HPC" }, name="{}_b={}_l={}_n={}_t={}_split={}".format(model_name, batch_size, lr, num_epochs, t, fold)) print('Start TRAINING model {} with splits under {}'.format(model_name, split_files_folder)) ### ----------------- Main loop ----------------- ### scaler = None if use_mixed_precision: scaler = GradScaler(device=DEVICE) best_auc = -math.inf best_epoch = 0 frac_string = str(fraction).replace('-', '') # -1.0 should be 1.0 if train: for epoch in range(num_epochs): title_suffix = 'EPOCH {}'.format(epoch) wandb.log({'Learning rate': optimizer.param_groups[0]['lr']}) # --- Training loop --- train_loss = train_epoch_subjects(model, optimizer, criterion, train_loader, DEVICE, values, show_figs=False, scaler=scaler, title_suffix=title_suffix) if use_scheduler and not scheduler_name.startswith("ReduceLROnPlateau"): scheduler.step() # Visualization of evaluation measures wandb.log({'Epoch': epoch, "average_loss/train": train_loss}) # --- Validation loop --- if do_eval: eval_loss, results_dict = eval_epoch_subjects(model, criterion, eval_loader, DEVICE, values, 'eval_', scaler=scaler, title_suffix=title_suffix) y_true, y_pred, y_probs = results_dict["y_true"], results_dict["y_pred"], results_dict["y_prob"] sens, spec = eva.compute_sensitivity_specificity(y_true, y_pred, y_scores=y_probs, prefix='Eval') roc_results = eva.compute_roc_curve(y_true, y_probs, plot=False, title_suffix='Evaluation ') roc_auc = roc_results['roc_auc_val'] if "roc_auc_val" in roc_results.keys() else math.nan wandb.log({"sensitivity/Eval": sens, "specificity/Eval": spec, "AUC_ROC/Eval": float(roc_auc)}) print(f"Evaluation AUC(Youden) {roc_auc}.") if use_scheduler and scheduler_name.startswith("ReduceLROnPlateau"): scheduler.step(eval_loss) if patience is not None: if (roc_auc - best_auc) > 0: # reached a better auc than before best_auc = roc_auc best_epoch = epoch elif (epoch - best_epoch) > patience: # patience is reached print("Early stopping: {} was best AUC at epoch {}".format(best_auc, best_epoch)) break if epoch % 10 == 0: save_model(model, resulting_folder, "{}_fold={}_subfold={}_frac={}_n={}".format(model_name, fold, subfold, frac_string, epoch)) if do_eval: with open( f'{resulting_folder}/results_{model_name}_fold{fold}_subfold{subfold}_frac{frac_string}-{epoch}', 'wb') as result_file: pickle.dump(results_dict, result_file) else: title_suffix = 'Final (Evaluation) EPOCH' # check performance on test set test_loss, final_results_dict = eval_epoch_subjects(model, criterion, test_loader, DEVICE, values, 'Test_', plot_tpr_fpr=resulting_folder + f"roc_{model_name}_{fold}.png", scaler=scaler, title_suffix=title_suffix) y_true, y_pred, y_probs = final_results_dict["y_true"], final_results_dict["y_pred"], final_results_dict["y_prob"] sens, spec = eva.compute_sensitivity_specificity(y_true, y_pred, y_scores=y_probs, prefix='Test') roc_results = eva.compute_roc_curve(y_true, y_probs, plot=True, title_suffix='Evaluation ') #roc = roc_results['roc'], fpr = roc_results['fpr'], tpr = roc_results['tpr'], thresholds = roc_results['thresholds'], best_index = roc_results['best_index_val'] roc_auc = roc_results['roc_auc_val'] if "roc_auc_val" in roc_results.keys() else math.nan wandb.log({"sensitivity/Test": sens, "specificity/Test": spec, "AUC_ROC/Test": float(roc_auc)}) print(f"Finished with test AUC (Youden) of {roc_auc}.") # save model parameters and predictions save_model(model, resulting_folder, "{}_fold={}_subfold={}_frac={}_FINAL".format(model_name, fold, subfold, frac_string)) with open(f'{resulting_folder}/results_{model_name}_fold={fold}_frac={frac_string}', 'wb') as result_file: pickle.dump(final_results_dict, result_file)