from pathlib import PurePath import wandb import torch import torch.utils.data import torch.nn as nn from torch.amp import GradScaler from torchmetrics import AUROC, Accuracy, Precision, Specificity, ROC from auto_detect_breast_mri.config import get_config, resolve_path from auto_detect_breast_mri.data.metadata import get_uka_metatensor from auto_detect_breast_mri.data import loaders from auto_detect_breast_mri.models.resnets import model_names from auto_detect_breast_mri.data.transforms import pre_image_shape, minimalAugmentation from auto_detect_breast_mri.models.checkpoints import load_pretrained_model from auto_detect_breast_mri.training.cli import generate_parser from auto_detect_breast_mri.training.loops import eval_epoch_subjects CPU = "cpu" GPU = "cuda" torch.manual_seed(31) DEVICE = GPU if not torch.cuda.is_available(): DEVICE = CPU transform = minimalAugmentation t = 'basic' batch_size = 32 fraction = 1.0 fold = 0 parser = generate_parser() args = parser.parse_args() model_name = args.model_name 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") model_path = args.model_path # only the leaf folder name, so no local path ends up in the run config split_files_folder_name = PurePath(split_files_folder.rstrip('/')).name split_file_path = split_files_folder + f"fold{fold}/stratified_test_set" feature_dataframe = get_uka_metatensor(0, feature_path) test_loader_abrv = loaders.get_subjects_dataloader(path_base, feature_dataframe, pre_image_shape, transform, 'abbreviated', split_file_path, batch_size, fraction, fold, subfold=0) test_loader_full = loaders.get_subjects_dataloader(path_base, feature_dataframe, pre_image_shape, transform, 'full', split_file_path, batch_size, fraction, fold, subfold=0) criterion = nn.BCEWithLogitsLoss() auc_roc = nn.ModuleDict({state: AUROC(**{"task": "binary"}).to(DEVICE) for state in [model_name + "_abrv", model_name + "_full"]}) roc = nn.ModuleDict({state: ROC(**{"task": "binary"}).to(DEVICE) for state in [model_name + "_abrv", model_name + "_full"]}) acc = nn.ModuleDict({state: Accuracy(**{"task": "binary"}).to(DEVICE) for state in [model_name + "_abrv", model_name + "_full"]}) sens = nn.ModuleDict({state: Precision(**{"task": "binary"}).to(DEVICE) for state in [model_name + "_abrv", model_name + "_full"]}) spec = nn.ModuleDict({state: Specificity(**{"task": "binary"}).to(DEVICE) for state in [model_name + "_abrv", model_name + "_full"]}) values = (auc_roc, roc, acc, sens, spec) model_abrv = model_names.get(model_name + '_abrv') model_full = model_names.get(model_name + '_full') scaler = GradScaler() used_mixed_precision = True wandb.init(**get_config().wandb_init_kwargs(), config={ "learning-rate": lr, "model": model_name, "number of test samples abrv": len(test_loader_abrv.dataset), "number of test samples full": len(test_loader_full.dataset), "batch size": batch_size, "Augmentation": t + str(transform), "Machine": "HPC", "Mixed Precision": used_mixed_precision, "Fold": fold, "splits": split_files_folder_name, }, name="{}_comparison_split={}".format(model_name, fold)) model_abrv.to(DEVICE) model_full.to(DEVICE) # load pretrained model: abrv_model_path = model_path.format(model_name + "_abrv", fold) print("load " + abrv_model_path) checkpoint_abrv = load_pretrained_model(model_abrv, abrv_model_path, DEVICE) model_abrv = model_abrv.load_state_dict(checkpoint_abrv) model_abrv.eval() full_model_path = model_path.format(model_name + "_full", fold) print("load " + full_model_path) checkpoint_full = load_pretrained_model(model_full, full_model_path, DEVICE) model_full = model_abrv.load_state_dict(checkpoint_full) model_full.eval() abrv_loss, abrv_results_dict = eval_epoch_subjects(model_abrv, criterion, test_loader_abrv, DEVICE, values, model_name + '_abrv', plot_tpr_fpr=True, scaler=scaler, title_suffix="ABRV") full_loss, full_results_dict = eval_epoch_subjects(model_full, criterion, test_loader_full, DEVICE, values, model_name + '_full', plot_tpr_fpr=True, scaler=scaler, title_suffix="FULL") filename = f"AUC_ROC-{model_name}_{fold}.png" fpr_abrv, tpr_abrv, thresholds_abrv = roc[model_name + '_abrv'].to('cpu').compute() fpr_full, tpr_full, thresholds_full = roc[model_name + '_full'].to('cpu').compute() y_pred = {model_name + '_abrv': tpr_abrv, model_name + '_full': tpr_full} results['fpr'] = fpr results['thresholds'] = thresholds plot_roc_curve(modelname, filename)