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