Download scripts/train/train.py from deboraJ23/AI_MRI: direct link, hf CLI and curl.
- Browser
- Download file 15.8 kB
-
https://huggingface.co/deboraJ23/AI_MRI/resolve/main/scripts/train/train.py
- Command line
-
hf download hf://deboraJ23/AI_MRI/scripts/train/train.py
-
curl -L -o train.py https://huggingface.co/deboraJ23/AI_MRI/resolve/main/scripts/train/train.py
15.8 kB
| #!/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) | |